Files
enginex-mrv100-compat/shims_ixformer_layers.py

96 lines
3.5 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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