extract adjust_hybrid_swa_layers_for_pp (#24334)
This commit is contained in:
@@ -665,6 +665,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.end_layer = getattr(self.model, "end_layer", model_num_layers)
|
self.end_layer = getattr(self.model, "end_layer", model_num_layers)
|
||||||
self.num_effective_layers = self.end_layer - self.start_layer
|
self.num_effective_layers = self.end_layer - self.start_layer
|
||||||
|
|
||||||
|
self.adjust_hybrid_swa_layers_for_pp()
|
||||||
|
|
||||||
# For LoopCoder models, each loop has its own layer_id, so we need to multiply by loop_num
|
# For LoopCoder models, each loop has its own layer_id, so we need to multiply by loop_num
|
||||||
loop_num = getattr(self.model_config.hf_config, "loop_num", 1)
|
loop_num = getattr(self.model_config.hf_config, "loop_num", 1)
|
||||||
if loop_num > 1:
|
if loop_num > 1:
|
||||||
@@ -679,23 +681,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
), "PP is not compatible with MTP models."
|
), "PP is not compatible with MTP models."
|
||||||
|
|
||||||
# Consider PP, so use start_layer and end_layer.
|
|
||||||
full_attention_layer_ids = [
|
|
||||||
layer_idx
|
|
||||||
for layer_idx in range(self.start_layer, self.end_layer + 1)
|
|
||||||
if hasattr(self.model_config, "full_attention_layer_ids")
|
|
||||||
and layer_idx in self.model_config.full_attention_layer_ids
|
|
||||||
]
|
|
||||||
swa_attention_layer_ids = [
|
|
||||||
layer_idx
|
|
||||||
for layer_idx in range(self.start_layer, self.end_layer + 1)
|
|
||||||
if hasattr(self.model_config, "swa_attention_layer_ids")
|
|
||||||
and layer_idx in self.model_config.swa_attention_layer_ids
|
|
||||||
]
|
|
||||||
# Update back to model_config.
|
|
||||||
self.model_config.swa_attention_layer_ids = swa_attention_layer_ids
|
|
||||||
self.model_config.full_attention_layer_ids = full_attention_layer_ids
|
|
||||||
|
|
||||||
# Apply torchao quantization
|
# Apply torchao quantization
|
||||||
torchao_applied = getattr(self.model, "torchao_applied", False)
|
torchao_applied = getattr(self.model, "torchao_applied", False)
|
||||||
# In layered loading, torchao may have been applied
|
# In layered loading, torchao may have been applied
|
||||||
@@ -792,6 +777,25 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
|
|
||||||
self.prealloc_symmetric_memory_pool()
|
self.prealloc_symmetric_memory_pool()
|
||||||
|
|
||||||
|
def adjust_hybrid_swa_layers_for_pp(self):
|
||||||
|
if not self.is_hybrid_swa:
|
||||||
|
return
|
||||||
|
|
||||||
|
full_attention_layer_ids = [
|
||||||
|
layer_idx
|
||||||
|
for layer_idx in range(self.start_layer, self.end_layer + 1)
|
||||||
|
if hasattr(self.model_config, "full_attention_layer_ids")
|
||||||
|
and layer_idx in self.model_config.full_attention_layer_ids
|
||||||
|
]
|
||||||
|
swa_attention_layer_ids = [
|
||||||
|
layer_idx
|
||||||
|
for layer_idx in range(self.start_layer, self.end_layer + 1)
|
||||||
|
if hasattr(self.model_config, "swa_attention_layer_ids")
|
||||||
|
and layer_idx in self.model_config.swa_attention_layer_ids
|
||||||
|
]
|
||||||
|
self.model_config.swa_attention_layer_ids = swa_attention_layer_ids
|
||||||
|
self.model_config.full_attention_layer_ids = full_attention_layer_ids
|
||||||
|
|
||||||
def init_routed_experts_capturer(self):
|
def init_routed_experts_capturer(self):
|
||||||
if not self.server_args.disable_shared_experts_fusion and hasattr(
|
if not self.server_args.disable_shared_experts_fusion and hasattr(
|
||||||
self.model, "num_fused_shared_experts"
|
self.model, "num_fused_shared_experts"
|
||||||
|
|||||||
Reference in New Issue
Block a user