Merge remote-tracking branch 'origin/feat/moe-layer' into feat/moe-layer

This commit is contained in:
Michelle0574 2026-07-07 17:55:13 +00:00
commit 0c6828076b
4 changed files with 452 additions and 384 deletions

View File

@ -1,155 +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(())) self.g_retr = nn.Parameter(torch.zeros(())) # 检索旁路零初始化门控
def forward( def forward(
self, self,
chem: torch.Tensor, chem: torch.Tensor,
tab: torch.Tensor, tab: torch.Tensor,
f_moe: Optional[torch.Tensor] = None, f_moe: Optional[torch.Tensor] = None,
f_llm: Optional[torch.Tensor] = None, f_llm: Optional[torch.Tensor] = None,
f_retr: Optional[torch.Tensor] = None, f_retr: Optional[torch.Tensor] = None,
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]]:
seq = torch.cat([chem, tab], dim=1) # [B, n_chem + n_cond, d_model] # 只对真实 token(chem + tab)做注意力池化,旁路不参与 softmax 竞争
pooled = self.pool(seq, return_attn_weights=return_attn_weights) seq = torch.cat([chem, tab], dim=1) # [B, n_chem + n_cond, d_model]
if return_attn_weights: pooled = self.pool(seq, return_attn_weights=return_attn_weights)
pooled, attn = pooled if return_attn_weights:
pooled, attn = pooled
out = pooled
if f_moe is not None: out = pooled
out = out + self.g_moe * f_moe # 残差 + 零初始化门 if f_moe is not None:
if f_llm is not None: out = out + self.g_moe * f_moe # 残差 + 零初始化门
out = out + self.g_llm * f_llm if f_llm is not None:
if f_retr is not None: out = out + self.g_llm * f_llm
out = out + self.g_retr * f_retr 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

@ -76,8 +76,19 @@ class LLMPromptEncoder(nn.Module):
quantization_config=bnb, torch_dtype=torch.bfloat16) quantization_config=bnb, torch_dtype=torch.bfloat16)
self.hidden_size = self.encoder.config.hidden_size self.hidden_size = self.encoder.config.hidden_size
elif _is_qwen: elif _is_qwen:
self.encoder = AutoModel.from_pretrained( # 微调(use_lora)时用 4-bit 量化(QLoRA)省显存;纯冻结时用 fp16
model_name_or_path, trust_remote_code=True, torch_dtype=torch.float16) if use_lora:
from transformers import BitsAndBytesConfig
_bnb = BitsAndBytesConfig(
load_in_4bit=True, bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.float16,
bnb_4bit_use_double_quant=True)
self.encoder = AutoModel.from_pretrained(
model_name_or_path, trust_remote_code=True,
quantization_config=_bnb, device_map={"": 0})
else:
self.encoder = AutoModel.from_pretrained(
model_name_or_path, trust_remote_code=True, torch_dtype=torch.float16)
self.hidden_size = self.encoder.config.hidden_size self.hidden_size = self.encoder.config.hidden_size
else: else:
self.encoder = AutoModel.from_pretrained(model_name_or_path) self.encoder = AutoModel.from_pretrained(model_name_or_path)
@ -122,6 +133,14 @@ class LLMPromptEncoder(nn.Module):
for p in self.encoder.parameters(): for p in self.encoder.parameters():
p.requires_grad = False p.requires_grad = False
# 4-bit 量化模型(QLoRA)需先 prepare,才能正确接收梯度
if getattr(self.encoder, "is_loaded_in_4bit", False) or getattr(self.encoder, "is_loaded_in_8bit", False):
from peft import prepare_model_for_kbit_training
self.encoder = prepare_model_for_kbit_training(self.encoder)
for p in self.encoder.parameters():
p.requires_grad = False
# 不同架构注意力层命名不同:
# T5: q/k/v/o;Roberta/ChemBERTa: query/key/value;Qwen/Llama: q_proj/k_proj/v_proj/o_proj
_enc_name = type(self.encoder).__name__.lower() _enc_name = type(self.encoder).__name__.lower()
if getattr(self, "_is_qwen", False) or "qwen" in _enc_name or "llama" in _enc_name: if getattr(self, "_is_qwen", False) or "qwen" in _enc_name or "llama" in _enc_name:
_targets = ["q_proj", "k_proj", "v_proj", "o_proj"] _targets = ["q_proj", "k_proj", "v_proj", "o_proj"]
@ -227,6 +246,39 @@ class LLMPromptEncoder(nn.Module):
if key not in self._prompt_cache: if key not in self._prompt_cache:
self._prompt_cache[key] = self._build_rag_prompt(s, self._retrieve_topk(s)) self._prompt_cache[key] = self._build_rag_prompt(s, self._retrieve_topk(s))
return self._prompt_cache[key] return self._prompt_cache[key]
def _rag_encode_batch(self, prompts, device):
"""编码一批 RAG prompt,取最后有效 token。grad 由调用方上下文决定。"""
outs = []
for i in range(0, len(prompts), 4):
bt = prompts[i:i+4]
enc = self.tokenizer(
bt, padding=True, truncation=True,
max_length=self.max_length, return_tensors="pt",
).to(device)
out = self.encoder(**enc).last_hidden_state # [B,L,H]
lengths = enc["attention_mask"].sum(1) - 1
b = torch.arange(out.size(0), device=device)
outs.append(out[b, lengths.long(), :].float()) # [B,H]
return torch.cat(outs, 0)
def _encode_rag(self, smiles, device):
"""RAG 编码。冻结时:no_grad + 缓存(快)。微调时:算梯度 + 不缓存。"""
if self._frozen:
# 冻结路径:缓存复用,no_grad
keys = [f"RAG::{self._rag_pool_id}::{s}" for s in smiles]
to_compute = [s for s, k in zip(smiles, keys) if k not in self._cache]
if to_compute:
uniq = list(dict.fromkeys(to_compute))
prompts = [self._build_rag_prompt(s, self._retrieve_topk(s)) for s in uniq]
with torch.no_grad():
feat = self._rag_encode_batch(prompts, device)
for s, v in zip(uniq, feat):
self._cache[f"RAG::{self._rag_pool_id}::{s}"] = v.float().cpu()
return torch.stack([self._cache[k] for k in keys]).to(device)
else:
# 微调路径:每次重新编码,保留计算图(算梯度),不缓存
prompts = [self._build_rag_prompt(s, self._retrieve_topk(s)) for s in smiles]
return self._rag_encode_batch(prompts, device)
# ---------- soft-prompt 编码(带梯度,不缓存特征)---------- # ---------- soft-prompt 编码(带梯度,不缓存特征)----------
def _encode_softrag(self, smiles, chem, tab, device) -> torch.Tensor: def _encode_softrag(self, smiles, chem, tab, device) -> torch.Tensor:

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()

View File

@ -682,6 +682,21 @@ def _run_single_outer_fold(
logger.info(f"[RESUME] fold {outer_fold} 复用已存 best_params,跳过内层 Optuna。") logger.info(f"[RESUME] fold {outer_fold} 复用已存 best_params,跳过内层 Optuna。")
full_dataset = LNPDataset(df) full_dataset = LNPDataset(df)
# === 断点续跑:已完成的 fold 直接跳过(读回磁盘结果)===
_tm = fold_dir / "test_metrics.json"
_bp = fold_dir / "best_params.json"
_em = fold_dir / "epoch_mean.json"
if precomputed_best_params is None and _tm.exists() and _bp.exists() and _em.exists():
logger.success(f"[SKIP] Outer fold {outer_fold} already done, loading cached results.")
with open(_tm) as _f: _tmd = json.load(_f)
with open(_bp) as _f: _bpd = json.load(_f)
with open(_em) as _f: _emd = json.load(_f)
return {
"fold": outer_fold,
"best_params": _bpd,
"epoch_mean": int(_emd.get("epoch_mean", _emd) if isinstance(_emd, dict) else _emd),
"test_metrics": _tmd,
}
logger.info(f"\n{'='*60}") logger.info(f"\n{'='*60}")
logger.info(f"OUTER FOLD {outer_fold}") logger.info(f"OUTER FOLD {outer_fold}")