96 lines
3.5 KiB
Python
96 lines
3.5 KiB
Python
# -*- 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
|