# -*- 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