enginex-mrv100-compat: shims_ixformer_layers.py
This commit is contained in:
95
shims_ixformer_layers.py
Normal file
95
shims_ixformer_layers.py
Normal file
@@ -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
|
||||
Reference in New Issue
Block a user