feat: support EPD disaggregation (#12263)

Co-authored-by: liusy58 <liusy58@linux.alibaba.com>
Co-authored-by: ZhengWG <zwg0606@gmail.com>
Co-authored-by: Nicholas <45984215+liusy58@users.noreply.github.com>
Co-authored-by: Shangming Cai <csmthu@gmail.com>
Co-authored-by: Yuhao Yang <47235274+yhyang201@users.noreply.github.com>
This commit is contained in:
Tianyu Guo
2025-12-14 22:30:08 +08:00
committed by GitHub
co-authored by liusy58 ZhengWG Nicholas Shangming Cai Yuhao Yang
parent a9ce1623cd
commit 9acb21ae27
19 changed files with 1910 additions and 68 deletions
+31 -14
View File
@@ -601,6 +601,7 @@ class Qwen3VLForConditionalGeneration(nn.Module):
super().__init__()
self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder
self.visual = Qwen3VLMoeVisionModel(
config.vision_config,
# NOTE: Qwen3-VL vision encoder currently supports BitsAndBytes 4-bit quantization.
@@ -616,22 +617,28 @@ class Qwen3VLForConditionalGeneration(nn.Module):
self.config: Qwen3VLConfig = config # for qwen3-vl
else:
self.config = config.text_config # for qwen3-omni
self.config.encoder_only = getattr(config, "encoder_only", False)
self.config.language_only = getattr(config, "language_only", False)
self.model = language_model_cls(
config=self.config,
quant_config=quant_config,
prefix=add_prefix("model", prefix),
)
if self.config.tie_word_embeddings:
self.lm_head = self.model.embed_tokens
else:
self.lm_head = ParallelLMHead(
self.config.vocab_size,
self.config.hidden_size,
if not hasattr(config, "encoder_only") or not config.encoder_only:
self.model = language_model_cls(
config=self.config,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
prefix=add_prefix("model", prefix),
)
if self.config.tie_word_embeddings:
self.lm_head = self.model.embed_tokens
else:
self.lm_head = ParallelLMHead(
self.config.vocab_size,
self.config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
)
else:
# encoder_only mode: no language model, so no lm_head needed
self.lm_head = None
self.is_mrope_enabled = "mrope_section" in self.config.rope_scaling
self.logits_processor = LogitsProcessor(self.config)
@@ -640,7 +647,7 @@ class Qwen3VLForConditionalGeneration(nn.Module):
# 8, 16, 24 layer will be merged to 0, 1, 2 layer of decoder output hidden_states
# deepstack
self.deepstack_visual_indexes = self.visual.deepstack_visual_indexes
self.deepstack_visual_indexes = config.vision_config.deepstack_visual_indexes
self.num_deepstack_embeddings = len(self.deepstack_visual_indexes)
self.use_deepstack = {Modality.IMAGE: True, Modality.VIDEO: True}
@@ -774,6 +781,11 @@ class Qwen3VLForConditionalGeneration(nn.Module):
# Skip loading extra bias for GPTQ models.
if name.endswith(".bias") and name not in params_dict:
continue
# Skip loading visual/language model weights
if (
self.config.encoder_only or self.config.language_only
) and name not in params_dict:
continue
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
@@ -788,6 +800,11 @@ class Qwen3VLForConditionalGeneration(nn.Module):
# Skip loading extra bias for GPTQ models.
if name.endswith(".bias") and name not in params_dict:
continue
# Skip loading visual/language model weights
if (
self.config.encoder_only or self.config.language_only
) and name not in params_dict:
continue
param = params_dict[name]
except KeyError:
print(params_dict.keys())