diff --git a/shims_ixformer_layers.py b/shims_ixformer_layers.py new file mode 100644 index 0000000..cfd338f --- /dev/null +++ b/shims_ixformer_layers.py @@ -0,0 +1,95 @@ +# -*- coding: utf-8 -*- +""" +ixformer.contrib.vllm.layers —— 纯 PyTorch 回退(R7 兜底) + +存在原因(实证):Iluvatar bi-100 引擎镜像里没有 `ixformer.contrib.vllm.layers`, +MoE 模型实现里 + from ixformer.contrib.vllm.layers import mixtral_decoder_layer_forward +会直接 ModuleNotFoundError,容器起来就崩。缺的是**引擎镜像的依赖**,不是模型的问题。 + +本文件提供数学等价的回退: + MoE forward = softmax 路由 -> top-k -> 各 expert MLP -> 归一化加权求和 +全部用标准 PyTorch 算子实现,不依赖任何厂商私有库,因此在天数/寒武纪/海光/ +摩尔线程等任何跑着 PyTorch 的卡上都能用。 + +定位说明:这是"让它能跑"的兜底,不是"跑得快"的优化。追求吞吐仍应用厂商原生算子。 +只有真缺模块时才会被 import 到(经 entrypoint 注入 PYTHONPATH=/opt/shims)。 +""" +import torch +import torch.nn.functional as F + +__all__ = ["moe_forward", "mixtral_decoder_layer_forward"] + + +def _flatten(x): + """[B,S,H] -> [B*S,H],返回还原所需的形状;已是 2D 则原样返回。""" + if x.dim() == 3: + return x.reshape(-1, x.shape[-1]), x.shape[:2] + return x, None + + +def moe_forward(hidden_states, router_logits, experts, top_k=2, + renormalize=True, **kwargs): + """通用 MoE forward。 + + 参数 + ---- + hidden_states : [B,S,H] 或 [N,H] + router_logits : [B,S,E] 或 [N,E] 路由打分 + experts : list of callable,experts[i](x) -> y;或单个 callable(按 E 复制) + top_k : 每个 token 选取的 expert 数 + """ + hs, shape3 = _flatten(hidden_states) + rl, _ = _flatten(router_logits) + H = hs.shape[-1] + + if not isinstance(experts, (list, tuple)): + experts = [experts] * rl.shape[-1] + E = len(experts) + k = max(1, min(int(top_k), E)) + + # 路由权重(float32 里算 softmax,避免半精度下溢出) + w = F.softmax(rl.float(), dim=-1) + top_w, top_i = torch.topk(w, k, dim=-1) + if renormalize: + top_w = top_w / top_w.sum(dim=-1, keepdim=True).clamp_min(1e-20) + top_w = top_w.to(hs.dtype) + + out = torch.zeros_like(hs) + for j in range(k): + idx = top_i[:, j] + wj = top_w[:, j].unsqueeze(-1) + for e in range(E): + mask = idx == e + if not bool(mask.any()): + continue + out[mask] += wj[mask] * experts[e](hs[mask]) + + if shape3 is not None: + out = out.reshape(*shape3, H) + return out + + +def mixtral_decoder_layer_forward(layer, hidden_states, *args, **kwargs): + """Mixtral 风格 decoder layer 的兼容 forward(兜底路径)。 + + 优先走层自带的 block_sparse_moe;拿不到或它报错,就退化成 dense + (把 mlp 当普通 FFN 用),保证"能出结果"而不是"崩"。 + """ + residual = hidden_states + h = layer.input_layernorm(hidden_states) + h = layer.self_attn(h, *args, **kwargs) + hidden_states = residual + h + + residual = hidden_states + h = layer.post_attention_layernorm(hidden_states) + moe = getattr(layer, "block_sparse_moe", None) or getattr(layer, "mlp", None) + if moe is None: + raise AttributeError("layer 上既没有 block_sparse_moe 也没有 mlp") + try: + return residual + moe(h) + except Exception: # noqa: BLE001 原生 MoE 路径失败 -> 退化 dense + out = moe(h) + if isinstance(out, tuple): + out = out[0] + return residual + out