[Bug Fix] GLM-V / GLM-OCR: field detection for transformers 5.x and MTP omission fix (#21134)

This commit is contained in:
Yuxuan Zhang
2026-03-23 13:19:48 -07:00
committed by GitHub
parent 4779755eb9
commit fcaad42b00
3 changed files with 16 additions and 9 deletions
@@ -604,7 +604,11 @@ def maybe_add_mtp_safetensors(
"""
# Only apply for GLM4Moe architecture with nextn layers
arch = getattr(hf_config, "architectures", [None])[0]
num_nextn_layers = getattr(hf_config, "num_nextn_predict_layers", 0)
num_nextn_layers = getattr(
getattr(hf_config, "text_config", hf_config),
"num_nextn_predict_layers",
getattr(hf_config, "num_nextn_predict_layers", 0),
)
if not (
arch in ["Glm4MoeForCausalLM", "Glm4MoeForCausalLMNextN"]
and num_nextn_layers > 0
+7 -7
View File
@@ -158,6 +158,13 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration):
params_dict = dict(self.named_parameters())
weight_names = []
for name, loaded_weight in weights:
if "language_model." in name:
name = name.replace("language_model.", "")
if "model.visual." in name:
name = name.replace("model.visual.", "visual.")
if "rotary_emb.inv_freq" in name:
continue
weight_names.append(name)
if self.num_fused_shared_experts > 0 and "mlp.shared_experts" in name:
@@ -196,13 +203,6 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration):
if is_decoder:
name = name.replace(nextn_layer_prefix, "model.decoder")
if "language_model." in name:
name = name.replace("language_model.", "")
if "model.visual." in name:
name = name.replace("model.visual.", "visual.")
if "rotary_emb.inv_freq" in name:
continue
for param_name, weight_name, shard_id in stacked_params_mapping:
# Skip non-stacked layers and experts (experts handled below).
if weight_name not in name:
+4 -1
View File
@@ -26,6 +26,7 @@ import torch.nn as nn
from einops import rearrange
from transformers.models.glm_ocr.configuration_glm_ocr import (
GlmOcrConfig,
GlmOcrTextConfig,
GlmOcrVisionConfig,
)
@@ -151,6 +152,7 @@ class GlmOcrVisionModel(Glm4vVisionModel):
def __init__(
self,
vision_config: GlmOcrVisionConfig,
text_config: GlmOcrTextConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
use_data_parallel: bool = False,
@@ -203,7 +205,7 @@ class GlmOcrVisionModel(Glm4vVisionModel):
)
self.merger = GlmOcrVisionPatchMerger(
d_model=vision_config.out_hidden_size,
context_dim=vision_config.out_hidden_size * vision_config.in_channels,
context_dim=text_config.intermediate_size,
quant_config=quant_config,
bias=False,
prefix=add_prefix("merger", prefix),
@@ -273,6 +275,7 @@ class GlmOcrForConditionalGeneration(Glm4vForConditionalGeneration):
self.use_data_parallel = get_global_server_args().mm_enable_dp_encoder
self.visual = GlmOcrVisionModel(
vision_config=config.vision_config,
text_config=config.text_config,
quant_config=quant_config,
prefix=add_prefix("visual", prefix),
use_data_parallel=self.use_data_parallel,