feat: 同步 RAG 融合与训练核心代码(fusion 支持 f_retr、models/trainer/moe/dataset)

This commit is contained in:
DicongLi 2026-07-05 01:07:26 +08:00
parent 3898fb8973
commit 987ecf6386
2 changed files with 383 additions and 379 deletions

View File

@ -1,152 +1,156 @@
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from typing import Dict, List, Literal, Optional, Tuple, Union from typing import Dict, List, Literal, Optional, Tuple, Union
PoolingStrategy = Literal["concat", "avg", "max", "attention"] PoolingStrategy = Literal["concat", "avg", "max", "attention"]
class FusionLayer(nn.Module): class FusionLayer(nn.Module):
""" """
将多个 token 融合成单个向量。 将多个 token 融合成单个向量。
输入: Dict[str, Tensor] 或 [B, n_tokens, d_model] 输入: Dict[str, Tensor] 或 [B, n_tokens, d_model]
输出: [B, fusion_dim] 输出: [B, fusion_dim]
策略: 策略:
- concat: [B, n_tokens, d_model] -> [B, n_tokens * d_model] - concat: [B, n_tokens, d_model] -> [B, n_tokens * d_model]
- avg: [B, n_tokens, d_model] -> [B, d_model] - avg: [B, n_tokens, d_model] -> [B, d_model]
- max: [B, n_tokens, d_model] -> [B, d_model] - max: [B, n_tokens, d_model] -> [B, d_model]
- attention: [B, n_tokens, d_model] -> [B, d_model] (learnable attention pooling) - attention: [B, n_tokens, d_model] -> [B, d_model] (learnable attention pooling)
""" """
def __init__( def __init__(
self, self,
d_model: int, d_model: int,
n_tokens: int, n_tokens: int,
strategy: PoolingStrategy = "attention", strategy: PoolingStrategy = "attention",
) -> None: ) -> None:
""" """
Args: Args:
d_model: 每个 token 的维度 d_model: 每个 token 的维度
n_tokens: token 数量(如 8) n_tokens: token 数量(如 8)
strategy: 融合策略 strategy: 融合策略
""" """
super().__init__() super().__init__()
self.d_model = d_model self.d_model = d_model
self.n_tokens = n_tokens self.n_tokens = n_tokens
self.strategy = strategy self.strategy = strategy
if strategy == "concat": if strategy == "concat":
self.fusion_dim = n_tokens * d_model self.fusion_dim = n_tokens * d_model
else: else:
self.fusion_dim = d_model self.fusion_dim = d_model
# Attention pooling: learnable query # Attention pooling: learnable query
if strategy == "attention": if strategy == "attention":
self.attn_query = nn.Parameter(torch.randn(1, 1, d_model)) self.attn_query = nn.Parameter(torch.randn(1, 1, d_model))
self.attn_proj = nn.Linear(d_model, d_model) self.attn_proj = nn.Linear(d_model, d_model)
def forward( def forward(
self, self,
x: Union[Dict[str, torch.Tensor], torch.Tensor], x: Union[Dict[str, torch.Tensor], torch.Tensor],
return_attn_weights: bool = False, return_attn_weights: bool = False,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
""" """
Args: Args:
x: Dict[str, Tensor] 每个 [B, d_model],或已 stack 的 [B, n_tokens, d_model] x: Dict[str, Tensor] 每个 [B, d_model],或已 stack 的 [B, n_tokens, d_model]
return_attn_weights: 若为 True 且策略为 attention,额外返回 attn_weights [B, n_tokens] return_attn_weights: 若为 True 且策略为 attention,额外返回 attn_weights [B, n_tokens]
Returns: Returns:
return_attn_weights=False: [B, fusion_dim] return_attn_weights=False: [B, fusion_dim]
return_attn_weights=True: ([B, fusion_dim], [B, n_tokens]) return_attn_weights=True: ([B, fusion_dim], [B, n_tokens])
""" """
if isinstance(x, dict): if isinstance(x, dict):
x = torch.stack(list(x.values()), dim=1) x = torch.stack(list(x.values()), dim=1)
if self.strategy == "concat": if self.strategy == "concat":
out = x.flatten(start_dim=1) out = x.flatten(start_dim=1)
return (out, None) if return_attn_weights else out return (out, None) if return_attn_weights else out
elif self.strategy == "avg": elif self.strategy == "avg":
out = x.mean(dim=1) out = x.mean(dim=1)
return (out, None) if return_attn_weights else out return (out, None) if return_attn_weights else out
elif self.strategy == "max": elif self.strategy == "max":
out = x.max(dim=1).values out = x.max(dim=1).values
return (out, None) if return_attn_weights else out return (out, None) if return_attn_weights else out
elif self.strategy == "attention": elif self.strategy == "attention":
return self._attention_pooling(x, return_attn_weights) return self._attention_pooling(x, return_attn_weights)
else: else:
raise ValueError(f"Unknown strategy: {self.strategy}") raise ValueError(f"Unknown strategy: {self.strategy}")
def _attention_pooling( def _attention_pooling(
self, x: torch.Tensor, return_attn_weights: bool = False, self, x: torch.Tensor, return_attn_weights: bool = False,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
""" """
Attention pooling: 用可学习 query 对 tokens 做加权求和 Attention pooling: 用可学习 query 对 tokens 做加权求和
Args: Args:
x: [B, n_tokens, d_model] x: [B, n_tokens, d_model]
return_attn_weights: 是否返回权重 return_attn_weights: 是否返回权重
Returns: Returns:
return_attn_weights=False: [B, d_model] return_attn_weights=False: [B, d_model]
return_attn_weights=True: ([B, d_model], [B, n_tokens]) return_attn_weights=True: ([B, d_model], [B, n_tokens])
""" """
B = x.size(0) B = x.size(0)
query = self.attn_query.expand(B, -1, -1) query = self.attn_query.expand(B, -1, -1)
keys = self.attn_proj(x) keys = self.attn_proj(x)
scores = torch.bmm(query, keys.transpose(1, 2)) / (self.d_model ** 0.5) scores = torch.bmm(query, keys.transpose(1, 2)) / (self.d_model ** 0.5)
attn_weights = F.softmax(scores, dim=-1) # [B, 1, n_tokens] attn_weights = F.softmax(scores, dim=-1) # [B, 1, n_tokens]
out = torch.bmm(attn_weights, x).squeeze(1) # [B, d_model] out = torch.bmm(attn_weights, x).squeeze(1) # [B, d_model]
if return_attn_weights: if return_attn_weights:
return out, attn_weights.squeeze(1) # [B, n_tokens] return out, attn_weights.squeeze(1) # [B, n_tokens]
return out return out
class ResidualConcatFusion(nn.Module): class ResidualConcatFusion(nn.Module):
"""对真实 token 做 attention pooling,再用零初始化门把 MoE/LLM 旁路以残差方式加入。 """对真实 token 做 attention pooling,再用零初始化门把 MoE/LLM 旁路以残差方式加入。
g_moe / g_llm 初始为 0 → +moe/+llm 起点严格等于 baseline; g_moe / g_llm 初始为 0 → +moe/+llm 起点严格等于 baseline;
旁路只有确实有用时才会被训练打开,从机制上保证“加了不会更差”。 旁路只有确实有用时才会被训练打开,从机制上保证“加了不会更差”。
""" """
def __init__(self, d_model: int, strategy: PoolingStrategy = "attention") -> None: def __init__(self, d_model: int, strategy: PoolingStrategy = "attention") -> None:
super().__init__() super().__init__()
if strategy == "concat": if strategy == "concat":
raise ValueError("ResidualConcatFusion 不支持 concat(token 数随开关变化)") raise ValueError("ResidualConcatFusion 不支持 concat(token 数随开关变化)")
self.d_model = d_model self.d_model = d_model
self.pool = FusionLayer(d_model=d_model, n_tokens=1, strategy=strategy) self.pool = FusionLayer(d_model=d_model, n_tokens=1, strategy=strategy)
self.fusion_dim = self.pool.fusion_dim self.fusion_dim = self.pool.fusion_dim
# 零初始化门控(可学习标量),旁路初始不参与 # 零初始化门控(可学习标量),旁路初始不参与
self.g_moe = nn.Parameter(torch.zeros(())) self.g_moe = nn.Parameter(torch.zeros(()))
self.g_llm = nn.Parameter(torch.zeros(())) self.g_llm = nn.Parameter(torch.zeros(()))
self.g_retr = nn.Parameter(torch.zeros(())) # 检索旁路零初始化门控
def forward(
self, def forward(
chem: torch.Tensor, self,
tab: torch.Tensor, chem: torch.Tensor,
f_moe: Optional[torch.Tensor] = None, tab: torch.Tensor,
f_llm: Optional[torch.Tensor] = None, f_moe: Optional[torch.Tensor] = None,
return_attn_weights: bool = False, f_llm: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: f_retr: Optional[torch.Tensor] = None,
# 只对真实 token(chem + tab)做注意力池化,旁路不参与 softmax 竞争 return_attn_weights: bool = False,
seq = torch.cat([chem, tab], dim=1) # [B, n_chem + n_cond, d_model] ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
pooled = self.pool(seq, return_attn_weights=return_attn_weights) # 只对真实 token(chem + tab)做注意力池化,旁路不参与 softmax 竞争
if return_attn_weights: seq = torch.cat([chem, tab], dim=1) # [B, n_chem + n_cond, d_model]
pooled, attn = pooled pooled = self.pool(seq, return_attn_weights=return_attn_weights)
if return_attn_weights:
out = pooled pooled, attn = pooled
if f_moe is not None:
out = out + self.g_moe * f_moe # 残差 + 零初始化门 out = pooled
if f_llm is not None: if f_moe is not None:
out = out + self.g_llm * f_llm out = out + self.g_moe * f_moe # 残差 + 零初始化门
if f_llm is not None:
out = out + self.g_llm * f_llm
if f_retr is not None:
out = out + self.g_retr * f_retr # 检索旁路,零初始化门保证起点=不开
return (out, attn) if return_attn_weights else out return (out, attn) if return_attn_weights else out

View File

@ -1,229 +1,229 @@
""" """
MoE (Mixture-of-Experts) 模块:sample-level + 跨模态路由。 MoE (Mixture-of-Experts) 模块:sample-level + 跨模态路由。
设计要点(与现有 8-token 架构对齐): 设计要点(与现有 8-token 架构对齐):
- Router 输入:tab token 池化后的 [B, d](即配方/实验条件向量)。 - Router 输入:tab token 池化后的 [B, d](即配方/实验条件向量)。
- Expert 输入:chem token flatten 后的 [B, T_chem * d](化学侧 token)。 - Expert 输入:chem token flatten 后的 [B, T_chem * d](化学侧 token)。
- 路由粒度:每个样本只算一次 router → gates [B, K]。 - 路由粒度:每个样本只算一次 router → gates [B, K]。
- Top-k 稀疏激活 + 可选训练态 jitter 噪声。 - Top-k 稀疏激活 + 可选训练态 jitter 噪声。
- 返回 load-balancing aux loss,由 trainer 端按权重加进总 loss。 - 返回 load-balancing aux loss,由 trainer 端按权重加进总 loss。
输出: 输出:
F_moe: [B, d_model] 与单个 token 同维度,便于追加到 fusion 序列中。 F_moe: [B, d_model] 与单个 token 同维度,便于追加到 fusion 序列中。
extras: dict 诊断与监控信息(aux loss、gates 等)。 extras: dict 诊断与监控信息(aux loss、gates 等)。
""" """
from typing import Dict, Tuple from typing import Dict, Tuple
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
class MoEAttentionPool(nn.Module): class MoEAttentionPool(nn.Module):
"""与 FusionLayer 同款的 attention pooling,参数独立。 """与 FusionLayer 同款的 attention pooling,参数独立。
将一组 token [B, T, d] 池化成单个查询向量 [B, d],用作 router 的输入。 将一组 token [B, T, d] 池化成单个查询向量 [B, d],用作 router 的输入。
""" """
def __init__(self, d_model: int) -> None: def __init__(self, d_model: int) -> None:
super().__init__() super().__init__()
self.d_model = d_model self.d_model = d_model
self.query = nn.Parameter(torch.randn(1, 1, d_model)) self.query = nn.Parameter(torch.randn(1, 1, d_model))
self.proj = nn.Linear(d_model, d_model) self.proj = nn.Linear(d_model, d_model)
def forward(self, x: torch.Tensor) -> torch.Tensor: def forward(self, x: torch.Tensor) -> torch.Tensor:
""" """
Args: Args:
x: [B, T, d_model] x: [B, T, d_model]
Returns: Returns:
[B, d_model] [B, d_model]
""" """
B = x.size(0) B = x.size(0)
q = self.query.expand(B, -1, -1) # [B, 1, d] q = self.query.expand(B, -1, -1) # [B, 1, d]
k = self.proj(x) # [B, T, d] k = self.proj(x) # [B, T, d]
scores = torch.bmm(q, k.transpose(1, 2)) / (self.d_model ** 0.5) scores = torch.bmm(q, k.transpose(1, 2)) / (self.d_model ** 0.5)
weights = F.softmax(scores, dim=-1) # [B, 1, T] weights = F.softmax(scores, dim=-1) # [B, 1, T]
return torch.bmm(weights, x).squeeze(1) # [B, d] return torch.bmm(weights, x).squeeze(1) # [B, d]
class MoERouter(nn.Module): class MoERouter(nn.Module):
"""单层 softmax router + Top-k 稀疏激活。 """单层 softmax router + Top-k 稀疏激活。
Args: Args:
d_model: router 输入维度。 d_model: router 输入维度。
n_experts: 专家数量 K。 n_experts: 专家数量 K。
top_k: 每个样本激活的专家数。 top_k: 每个样本激活的专家数。
jitter_noise: 训练态加在 logits 上的均匀噪声幅度,用于探索;评估态自动关闭。 jitter_noise: 训练态加在 logits 上的均匀噪声幅度,用于探索;评估态自动关闭。
""" """
def __init__( def __init__(
self, self,
d_model: int, d_model: int,
n_experts: int, n_experts: int,
top_k: int = 2, top_k: int = 2,
jitter_noise: float = 0.0, jitter_noise: float = 0.0,
) -> None: ) -> None:
super().__init__() super().__init__()
if not 1 <= top_k <= n_experts: if not 1 <= top_k <= n_experts:
raise ValueError(f"top_k 必须在 [1, n_experts={n_experts}],收到 {top_k}") raise ValueError(f"top_k 必须在 [1, n_experts={n_experts}],收到 {top_k}")
self.n_experts = n_experts self.n_experts = n_experts
self.top_k = top_k self.top_k = top_k
self.jitter_noise = jitter_noise self.jitter_noise = jitter_noise
self.linear = nn.Linear(d_model, n_experts) self.linear = nn.Linear(d_model, n_experts)
def forward( def forward(
self, q: torch.Tensor, self, q: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
""" """
Args: Args:
q: [B, d_model] 路由查询向量。 q: [B, d_model] 路由查询向量。
Returns: Returns:
gates: [B, n_experts] top-k 后再归一化的稀疏概率。 gates: [B, n_experts] top-k 后再归一化的稀疏概率。
probs_full: [B, n_experts] 未掩码的 softmax 概率(用于 aux loss)。 probs_full: [B, n_experts] 未掩码的 softmax 概率(用于 aux loss)。
expert_mask: [B, n_experts] 0/1 掩码,标记被激活的专家。 expert_mask: [B, n_experts] 0/1 掩码,标记被激活的专家。
""" """
logits = self.linear(q) # [B, K] logits = self.linear(q) # [B, K]
if self.training and self.jitter_noise > 0.0: if self.training and self.jitter_noise > 0.0:
noise = (torch.rand_like(logits) - 0.5) * self.jitter_noise noise = (torch.rand_like(logits) - 0.5) * self.jitter_noise
logits = logits + noise logits = logits + noise
probs_full = F.softmax(logits, dim=-1) # [B, K] probs_full = F.softmax(logits, dim=-1) # [B, K]
topk_vals, topk_idx = probs_full.topk(self.top_k, dim=-1) # [B, k] topk_vals, topk_idx = probs_full.topk(self.top_k, dim=-1) # [B, k]
expert_mask = torch.zeros_like(probs_full) expert_mask = torch.zeros_like(probs_full)
expert_mask.scatter_(1, topk_idx, 1.0) expert_mask.scatter_(1, topk_idx, 1.0)
gates = probs_full * expert_mask gates = probs_full * expert_mask
gates = gates / gates.sum(dim=-1, keepdim=True).clamp(min=1e-9) gates = gates / gates.sum(dim=-1, keepdim=True).clamp(min=1e-9)
return gates, probs_full, expert_mask return gates, probs_full, expert_mask
class MoEExpert(nn.Module): class MoEExpert(nn.Module):
"""单个专家 MLP:in_dim → hidden_dim → out_dim。""" """单个专家 MLP:in_dim → hidden_dim → out_dim。"""
def __init__( def __init__(
self, in_dim: int, hidden_dim: int, out_dim: int, dropout: float = 0.1, self, in_dim: int, hidden_dim: int, out_dim: int, dropout: float = 0.1,
) -> None: ) -> None:
super().__init__() super().__init__()
self.net = nn.Sequential( self.net = nn.Sequential(
nn.Linear(in_dim, hidden_dim), nn.Linear(in_dim, hidden_dim),
nn.GELU(), nn.GELU(),
nn.Dropout(dropout), nn.Dropout(dropout),
nn.Linear(hidden_dim, out_dim), nn.Linear(hidden_dim, out_dim),
) )
def forward(self, x: torch.Tensor) -> torch.Tensor: def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.net(x) return self.net(x)
class MoEBlock(nn.Module): class MoEBlock(nn.Module):
""" """
Sample-level 跨模态 MoE。 Sample-level 跨模态 MoE。
流程: 流程:
1. tab tokens [B, T_tab, d] -> attention pool -> q_tab [B, d] 1. tab tokens [B, T_tab, d] -> attention pool -> q_tab [B, d]
2. q_tab -> Router -> gates [B, K] (+ probs_full / expert_mask) 2. q_tab -> Router -> gates [B, K] (+ probs_full / expert_mask)
3. chem tokens [B, T_chem, d] -> flatten -> [B, T_chem * d] 3. chem tokens [B, T_chem, d] -> flatten -> [B, T_chem * d]
4. K 个 Expert MLP 并行处理 -> stack [B, K, d] 4. K 个 Expert MLP 并行处理 -> stack [B, K, d]
5. 加权求和 -> F_moe [B, d] 5. 加权求和 -> F_moe [B, d]
6. 计算 load-balancing aux loss 6. 计算 load-balancing aux loss
Args: Args:
d_model: token 维度。 d_model: token 维度。
n_chem_tokens: chem 侧 token 数(决定 expert 输入维度)。 n_chem_tokens: chem 侧 token 数(决定 expert 输入维度)。
n_experts: 专家数量 K。 n_experts: 专家数量 K。
top_k: 每个样本激活的专家数。 top_k: 每个样本激活的专家数。
expert_hidden_mult: expert 中间层维度 = expert_hidden_mult * d_model。 expert_hidden_mult: expert 中间层维度 = expert_hidden_mult * d_model。
dropout: expert 内部 dropout。 dropout: expert 内部 dropout。
jitter_noise: router 训练态噪声幅度,0.0 表示关闭。 jitter_noise: router 训练态噪声幅度,0.0 表示关闭。
""" """
def __init__( def __init__(
self, self,
d_model: int, d_model: int,
n_chem_tokens: int = 4, n_chem_tokens: int = 4,
n_experts: int = 4, n_experts: int = 4,
top_k: int = 2, top_k: int = 2,
expert_hidden_mult: int = 2, expert_hidden_mult: int = 2,
dropout: float = 0.1, dropout: float = 0.1,
jitter_noise: float = 0.0, jitter_noise: float = 0.0,
) -> None: ) -> None:
super().__init__() super().__init__()
self.d_model = d_model self.d_model = d_model
self.n_chem_tokens = n_chem_tokens self.n_chem_tokens = n_chem_tokens
self.n_experts = n_experts self.n_experts = n_experts
self.top_k = top_k self.top_k = top_k
self.tab_pool = MoEAttentionPool(d_model) self.tab_pool = MoEAttentionPool(d_model)
self.router = MoERouter( self.router = MoERouter(
d_model, n_experts, top_k=top_k, jitter_noise=jitter_noise, d_model, n_experts, top_k=top_k, jitter_noise=jitter_noise,
) )
expert_in = d_model * n_chem_tokens expert_in = d_model * n_chem_tokens
expert_hidden = d_model * expert_hidden_mult expert_hidden = d_model * expert_hidden_mult
self.experts = nn.ModuleList([ self.experts = nn.ModuleList([
MoEExpert(expert_in, expert_hidden, d_model, dropout=dropout) MoEExpert(expert_in, expert_hidden, d_model, dropout=dropout)
for _ in range(n_experts) for _ in range(n_experts)
]) ])
def forward( def forward(
self, self,
chem: torch.Tensor, chem: torch.Tensor,
tab: torch.Tensor, tab: torch.Tensor,
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]: ) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
""" """
Args: Args:
chem: [B, T_chem, d_model] 化学侧 token(router 不看)。 chem: [B, T_chem, d_model] 化学侧 token(router 不看)。
tab: [B, T_tab, d_model] 配方/实验侧 token(router 看)。 tab: [B, T_tab, d_model] 配方/实验侧 token(router 看)。
Returns: Returns:
F_moe: [B, d_model] F_moe: [B, d_model]
extras: { extras: {
"lb_loss": 标量 load-balancing aux loss(带梯度), "lb_loss": 标量 load-balancing aux loss(带梯度),
"gates": [B, K] 稀疏归一后的门控(detached), "gates": [B, K] 稀疏归一后的门控(detached),
"probs": [B, K] 原始 softmax 概率(detached), "probs": [B, K] 原始 softmax 概率(detached),
} }
""" """
if chem.size(1) != self.n_chem_tokens: if chem.size(1) != self.n_chem_tokens:
raise ValueError( raise ValueError(
f"chem token 数不匹配:期望 {self.n_chem_tokens}, " f"chem token 数不匹配:期望 {self.n_chem_tokens}, "
f"实际 {chem.size(1)}" f"实际 {chem.size(1)}"
) )
q_tab = self.tab_pool(tab) # [B, d] q_tab = self.tab_pool(tab) # [B, d]
gates, probs_full, expert_mask = self.router(q_tab) # 各 [B, K] gates, probs_full, expert_mask = self.router(q_tab) # 各 [B, K]
flat = chem.flatten(start_dim=1) # [B, T_chem * d] flat = chem.flatten(start_dim=1) # [B, T_chem * d]
expert_outs = torch.stack( expert_outs = torch.stack(
[expert(flat) for expert in self.experts], dim=1, [expert(flat) for expert in self.experts], dim=1,
) # [B, K, d] ) # [B, K, d]
F_moe = (gates.unsqueeze(-1) * expert_outs).sum(dim=1) # [B, d] F_moe = (gates.unsqueeze(-1) * expert_outs).sum(dim=1) # [B, d]
lb_loss = self._load_balancing_loss(probs_full, expert_mask) lb_loss = self._load_balancing_loss(probs_full, expert_mask)
return F_moe, { return F_moe, {
"lb_loss": lb_loss, "lb_loss": lb_loss,
"gates": gates.detach(), "gates": gates.detach(),
"probs": probs_full.detach(), "probs": probs_full.detach(),
} }
def _load_balancing_loss( def _load_balancing_loss(
self, self,
probs_full: torch.Tensor, probs_full: torch.Tensor,
expert_mask: torch.Tensor, expert_mask: torch.Tensor,
) -> torch.Tensor: ) -> torch.Tensor:
"""Switch Transformer 风格的 load-balancing loss。 """Switch Transformer 风格的 load-balancing loss。
f_i = 该 batch 中 expert_i 被激活的样本占比(top-k mask 的均值) f_i = 该 batch 中 expert_i 被激活的样本占比(top-k mask 的均值)
p_i = 该 batch 中 expert_i 的 softmax 概率均值 p_i = 该 batch 中 expert_i 的 softmax 概率均值
loss = K * Σ_i f_i * p_i loss = K * Σ_i f_i * p_i
理想情况下 f_i 与 p_i 都接近 1/K,loss ≈ 1。 理想情况下 f_i 与 p_i 都接近 1/K,loss ≈ 1。
""" """
f = expert_mask.mean(dim=0) # [K] f = expert_mask.mean(dim=0) # [K]
p = probs_full.mean(dim=0) # [K] p = probs_full.mean(dim=0) # [K]
return self.n_experts * (f * p).sum() return self.n_experts * (f * p).sum()