Files
enginex-mrv100-compat/shims_ixformer_layers.py

96 lines
3.5 KiB
Python
Raw Normal View History

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