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 不支持 concattoken 数随开关变化)") raise ValueError("ResidualConcatFusion 不支持 concattoken 数随开关变化)")
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,
# 只对真实 tokenchem + 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) # 只对真实 tokenchem + 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 lossgates extras: dict 诊断与监控信息aux lossgates
""" """
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):
"""单个专家 MLPin_dim → hidden_dim → out_dim。""" """单个专家 MLPin_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] 化学侧 tokenrouter 不看 chem: [B, T_chem, d_model] 化学侧 tokenrouter 不看
tab: [B, T_tab, d_model] 配方/实验侧 tokenrouter tab: [B, T_tab, d_model] 配方/实验侧 tokenrouter
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/Kloss 1 理想情况下 f_i p_i 都接近 1/Kloss 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()