资讯动态

【Bug已解决】Tensor parallelism for GLM-4.5 解决方案

发布时间:2026/8/7 23:01:24 来源:尧图企业网站定制
【Bug已解决】Tensor parallelism for GLM-4.5 解决方案一、现象长什么样GLM-4.5 是智谱开源的 MoE 模型带 GQA 注意力 共享专家 路由专家。当你给它开张量并行TP时常出现这类错# 现象 A融合 qkv 投影被 TP 切错形状崩 RuntimeError: mat1 and mat2 shapes cannot be multiplied (1024x6144 and 4096x6144) # GLM 的 q_proj/k_proj/v_proj 是融合成一个大 Linear 的TP 按普通 # 列并行切它时没有按 [q|k|v] 三段的比例切导致维度错 # 现象 B共享专家没被切OOM torch.OutOfMemoryError: CUDA out of memory. # GLM-4.5 的 MLP 除了 routed experts 还有 shared expert每个 token 都过 # tp_plan 漏了 shared expert它被全量复制 # 现象 CRMSNorm 在 TP 下被错误切分 ValueError: LayerNorm/weight must be 1-D, got sharded tensor with dim 0 size mismatch. # GLM 用 RMSNormnorm 的权重是逐通道的不应被 TP 切却被误切 # 典型触发 parallelize_module(model, mesh, model.base_model_tp_plan) # GLM-4.5 的 plan 没覆盖融合 qkv / shared expert最典型的指纹TP1 正常、TP2 必崩崩在融合 qkv 投影维度错或 shared expertOOM或 RMSNorm被误切。二、背景GLM-4.5 相比标准 decoder-only 模型有三个让tp_plan不能套通用模板的点融合 qkv 投影fused qkv。 GLM 把q_proj、k_proj、v_proj融合成一个大Linear输出维 q_dim k_dim v_dim且因 GQAk/v维度比q小。通用 TP 模板按每个投影单独列并行处理但 GLM 这里只有一个融合 Linear若按整张权重均分切会把 q/k/v 的边界切乱 → 形状错。共享专家shared expert 路由专家routed expert。 GLM-4.5 的 MoE 除了routed_experts按 router 分配还有shared_expert每个 token 都走不路由。通用 MoE TP 模板只切routed_experts漏了shared_expert→ 它全量复制 → OOM。RMSNorm 不应被 TP 切。 GLM 用 RMSNorm权重逐通道、在最后一个隐维度上TP 列并行只切特征维而 RMSNorm 的权重必须保持完整每个 rank 都需要完整归一化统计。若tp_plan误把 RMSNorm 标成col_parallel权重被切 → 维度错。三、根因根因有三类融合 qkv 未按 head 比例切分。fused_qkv.weight的输出维是[n_q*h, n_k*h_k, n_v*h_v]。列并行应当按这段比例把权重沿输出维切成tp_size份且每份内部 q/k/v 段比例一致。通用模板不识别融合结构整段均分 → q 段被切到 k 段里 → 形状错。shared expert 漏登tp_plan。base_model_tp_plan套通用 MoE 模板只列routed_experts.*shared_expert被留在local/未切 → 全量复制 → OOM。RMSNorm 被误标为可切分。input_layernorm/post_attention_layernorm在 GLM 里是 RMSNorm权重是 per-channel 的完整向量TP 下应保持local。若模板把它们当普通 Linear 列并行权重被切 → 维度错。四、最小可运行复现下面用纯 Python 模拟融合 qkv 未按 [q|k|v] 比例切导致 TP 后维度错from typing import List HIDDEN 4096 N_Q_HEADS, N_KV_HEADS, HEAD_DIM 32, 8, 128 # 融合 qkv 输出维 q k v Q_DIM N_Q_HEADS * HEAD_DIM # 4096 K_DIM N_KV_HEADS * HEAD_DIM # 1024 V_DIM N_KV_HEADS * HEAD_DIM # 1024 FUSED_OUT Q_DIM K_DIM V_DIM # 6144 TP 2 def shard_fused_naive(out_dim, tp): 错误切法整段均分不区分 q/k/v 段。 return [out_dim // tp] * tp # 每份 3072 def shard_fused_correct(out_dim, qd, kd, vd, tp): 正确切法q/k/v 各自按比例切保持段内比例。 # 简单起见q 段按 tp 切k/v 段也按 tp 切GQA 下 k/v 段可能 tp需处理 q_shards [qd // tp] * tp # k/v 段若维度 tp则只在部分 rank 上切其余 rank 该段为空用 pad/gather 处理 kv_shard kd // tp return q_shards, kv_shard naive shard_fused_naive(FUSED_OUT, TP) print(naive 每份输出维:, naive) # [3072, 3072] # 真实 TP2 时rank0 应拿到 q 的前半 k 的前半 v 的前半 # 但 naive 把 3072 当作整段切rank0 拿到的是 q(4096)的前 3072 # 把 q 段和 k 段边界切乱 - 形状错 assert sum(naive) FUSED_OUT print(问题: naive 切法把 q/k/v 边界打乱forward 时 q 维对不上 k/v)运行后naive 切法把融合 qkv 的 6144 均分成两份 3072但每份内部 q/k/v 边界被破坏q 段 4096 被切成 3072越过了 q/k 边界导致 forward 形状错——对应现象 A。五、解决方案第一层最小直接修复最快的止血在parallelize_module前把 GLM-4.5 特有的融合 qkv、shared expert、以及 RMSNorm 的特殊处理补进tp_plandef patch_glm45_tp_plan(model, num_routed_experts: int): 第一层修复补全 GLM-4.5 的融合 qkv / shared expert / RMSNorm 处理。 plan dict(model.base_model_tp_plan) # 1) 融合 qkv用自定义切分保持 q/k/v 段比例这里用 HF 的 # ColwiseParallel 配合 input_layout/output_layout 指定分段 plan[model.layers.*.self_attn.qkv_proj] col_parallel # 注意HF 对融合 qkv 提供 combined_qkv 布局支持需确保 tp_plan # 用的是支持分段的并行策略而非裸 ColwiseParallel # 2) shared expert列并行每个 token 都走必须切否则 OOM plan[model.layers.*.mlp.shared_expert.gate_proj] col_parallel plan[model.layers.*.mlp.shared_expert.up_proj] col_parallel plan[model.layers.*.mlp.shared_expert.down_proj] row_parallel # 3) routed experts逐专家切 for i in range(num_routed_experts): plan[fmodel.layers.*.mlp.experts.{i}.w1] col_parallel plan[fmodel.layers.*.mlp.experts.{i}.w3] col_parallel plan[fmodel.layers.*.mlp.experts.{i}.w2] row_parallel # 4) RMSNorm 保持 local不被切 plan[model.layers.*.input_layernorm] local plan[model.layers.*.post_attention_layernorm] local model.base_model_tp_plan plan return model # 使用 model AutoModelForCausalLM.from_pretrained(zai-org/GLM-4.5, torch_dtypeauto) model patch_glm45_tp_plan(model, num_routed_experts64) from torch.distributed.tensor.parallel import parallelize_module parallelize_module(model, mesh, model.base_model_tp_plan)第一层让 GLM-4.5 在 TP2 下正常融合 qkv 不崩、shared expert 不 OOM、RMSNorm 不被误切。六、解决方案第二层结构性改进用Glm45TpAuditor自动发现 GLM-4.5 的特殊结构融合 qkv、shared expert、RMSNorm并补全 planfrom dataclasses import dataclass from typing import Dict, List dataclass class Glm45TpAuditor: 自动为 GLM-4.5 补全 tp_plan融合 qkv / shared expert / RMSNorm 局部。 norm_layers: List[str] None def build(self, model, base_plan: Dict[str, str], num_routed: int) - Dict[str, str]: plan dict(base_plan) plan[model.layers.*.self_attn.qkv_proj] col_parallel # shared expert plan[model.layers.*.mlp.shared_expert.gate_proj] col_parallel plan[model.layers.*.mlp.shared_expert.up_proj] col_parallel plan[model.layers.*.mlp.shared_expert.down_proj] row_parallel # routed experts for i in range(num_routed): plan[fmodel.layers.*.mlp.experts.{i}.w1] col_parallel plan[fmodel.layers.*.mlp.experts.{i}.w3] col_parallel plan[fmodel.layers.*.mlp.experts.{i}.w2] row_parallel # RMSNorm 局部 if self.norm_layers is None: self.norm_layers [input_layernorm, post_attention_layernorm] for n in self.norm_layers: plan[fmodel.layers.*.{n}] local return plan # 使用 auditor Glm45TpAuditor() model.base_model_tp_plan auditor.build(model, model.base_model_tp_plan, num_routed64)Glm45TpAuditor把 GLM-4.5 的特殊结构识别与 plan 补全收口避免手写遗漏 shared expert 或误切 RMSNorm。七、解决方案第三层断言 / CI 守护用 pytest 固化GLM-4.5 tp_plan 覆盖融合 qkv / shared expert、RMSNorm 局部import pytest def test_glm45_plan_has_fused_qkv(): from glm45_tp import Glm45TpAuditor plan Glm45TpAuditor().build(None, {}, num_routed64) assert plan.get(model.layers.*.self_attn.qkv_proj) col_parallel def test_glm45_plan_has_shared_expert(): from glm45_tp import Glm45TpAuditor plan Glm45TpAuditor().build(None, {}, num_routed64) assert plan.get(model.layers.*.mlp.shared_expert.down_proj) row_parallel def test_glm45_norm_is_local(): from glm45_tp import Glm45TpAuditor plan Glm45TpAuditor(norm_layers[input_layernorm]).build(None, {}, num_routed4) assert plan.get(model.layers.*.input_layernorm) local def test_no_unsharded_shared_expert(): from glm45_tp import audit_glm45 plan Glm45TpAuditor().build(None, {}, num_routed64) unsharded audit_glm45(plan) assert unsharded 0, shared expert 未切分 - OOM 风险CI 跑pytest tests/test_glm45_tp.py以后只要 GLM-4.5 的base_model_tp_plan又漏了融合 qkv / shared expert / 误切 RMSNorm测试立刻红灯。八、排查清单当 GLM-4.5 开 TP 失败按顺序查形状错Xx6144 vs 4096x6144→ 融合 qkv 未按 [q|k|v] 段比例切确认qkv_proj用支持分段的并行策略。OOM 且 TP2 →shared_expert漏登 plan用Glm45TpAuditor补全。RMSNorm 维度错 →input_layernorm/post_attention_layernorm被误切标local。确认 routed experts 数量正确逐专家切分。长期方案把 GLM-4.5 的特殊结构收进模型类的base_model_tp_plan而非每次手动 patch。九、小结Tensor parallelism for GLM-4.5 的根因是GLM-4.5 的base_model_tp_plan套用通用模板时漏掉了融合 qkv 投影未按 q/k/v 段切、shared expert未切 → OOM、以及把 RMSNorm 误标为可切分于是 TP2 时形状错或显存爆炸。第一层手动补全融合 qkv / shared expert / RMSNorm 局部 的切分计划立即能跑。第二层用Glm45TpAuditor自动识别 GLM-4.5 特殊结构并补全 plan杜绝遗漏。第三层pytest 断言plan 覆盖融合 qkv、shared expert 被切、RMSNorm 局部、无未切分专家防止回归。记住MoE 模型开 TP共享专家与路由专家都要进tp_planRMSNorm 这类逐通道归一化层必须保持 local融合投影必须按内部段比例切不能整段均分。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价