From 724814851594dbef7389807a175a2f3d935150e6 Mon Sep 17 00:00:00 2001 From: Michelle0474 <2170308303@qq.com> Date: Wed, 1 Jul 2026 15:52:20 +0800 Subject: [PATCH 1/4] =?UTF-8?q?feat(llm):=20Qwen=20QLoRA=204-bit=20?= =?UTF-8?q?=E7=9C=9F=E5=BE=AE=E8=B0=83=E8=B7=AF=E5=BE=84=20+=20RAG=20?= =?UTF-8?q?=E7=BC=96=E7=A0=81=E5=86=BB=E7=BB=93/=E5=BE=AE=E8=B0=83?= =?UTF-8?q?=E5=88=86=E6=B5=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - use_lora 时 Qwen 改 4-bit(NF4) 加载,配 prepare_model_for_kbit_training 使 LoRA 可回传梯度 - _encode_rag 拆分:冻结走 no_grad+缓存,微调走保留计算图+不缓存 - fold_0/1 delivery R2 0.44/0.24,较冻结版 0.185 显著提升 --- lnp_ml/modeling/layers/llm_prompt.py | 76 ++++++++++++++++++---------- 1 file changed, 48 insertions(+), 28 deletions(-) diff --git a/lnp_ml/modeling/layers/llm_prompt.py b/lnp_ml/modeling/layers/llm_prompt.py index 824dfef..7500873 100644 --- a/lnp_ml/modeling/layers/llm_prompt.py +++ b/lnp_ml/modeling/layers/llm_prompt.py @@ -54,8 +54,19 @@ class LLMPromptEncoder(nn.Module): self.encoder = T5EncoderModel.from_pretrained(model_name_or_path) self.hidden_size = self.encoder.config.d_model elif _is_qwen: - self.encoder = AutoModel.from_pretrained( - model_name_or_path, trust_remote_code=True, torch_dtype=torch.float16) + # 微调(use_lora)时用 4-bit 量化(QLoRA)省显存;纯冻结时用 fp16 + 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 else: self.encoder = AutoModel.from_pretrained(model_name_or_path) @@ -82,6 +93,10 @@ class LLMPromptEncoder(nn.Module): def _apply_lora(self, r: int, alpha: int, dropout: float) -> None: from peft import LoraConfig, get_peft_model + # 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 # 不同架构注意力层命名不同: @@ -153,34 +168,39 @@ class LLMPromptEncoder(nn.Module): "into your internal representation." ) - @torch.no_grad() + 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 编码:每个分子→检索→构造prompt→Qwen编码→取最后有效token。带缓存。""" - feats = [] - to_compute = [] - keys = [] - for s in smiles: - key = f"RAG::{self._rag_pool_id}::{s}" - keys.append(key) - if key not in self._cache: - to_compute.append(s) - # 逐个构造 prompt(检索是 CPU 操作) - if to_compute: - uniq = list(dict.fromkeys(to_compute)) - prompts = [self._build_rag_prompt(s, self._retrieve_topk(s)) for s in uniq] - for i in range(0, len(uniq), 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 # 最后有效token位置 - b = torch.arange(out.size(0), device=device) - last_tok = out[b, lengths.long(), :] # [B,H] - for s, v in zip(uniq[i:i+4], last_tok): + """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) + 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) def _mean_pool(self, last_hidden: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: m = mask.unsqueeze(-1).float() # [B, L, 1] From 07b6b44f50e2a110531c4fc1ae74c544ed1ab247 Mon Sep 17 00:00:00 2001 From: Michelle0474 <2170308303@qq.com> Date: Fri, 3 Jul 2026 17:37:34 +0800 Subject: [PATCH 2/4] =?UTF-8?q?feat(cv):=20=E6=96=AD=E7=82=B9=E7=BB=AD?= =?UTF-8?q?=E8=B7=91=20-=20=E4=B8=BB=E5=BE=AA=E7=8E=AF=E8=B7=B3=E8=BF=87?= =?UTF-8?q?=E5=B7=B2=E5=AE=8C=E6=88=90=E7=9A=84=20outer=20fold?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 启动时若 fold_dir 已有 test_metrics/best_params/epoch_mean,直接读回结果跳过训练。 仅对主循环生效(precomputed_best_params is None),不影响 repeat 阶段。 应对服务器偶发崩溃,崩溃重启后重跑同一 output-dir 可接续未完成的 fold。 --- lnp_ml/modeling/nested_cv_optuna.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/lnp_ml/modeling/nested_cv_optuna.py b/lnp_ml/modeling/nested_cv_optuna.py index b235d10..5fd94c3 100644 --- a/lnp_ml/modeling/nested_cv_optuna.py +++ b/lnp_ml/modeling/nested_cv_optuna.py @@ -632,6 +632,21 @@ def _run_single_outer_fold( fold_dir.mkdir(parents=True, exist_ok=True) 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"OUTER FOLD {outer_fold}") From 3898fb897325ddef98484758e9333bf9b97d278a Mon Sep 17 00:00:00 2001 From: DicongLi <2024379585@qq.com> Date: Sun, 5 Jul 2026 01:01:22 +0800 Subject: [PATCH 3/4] feat(rag): add retrieval module (_smiles_to_fp, Top-K neighbor search) --- lnp_ml/modeling/retrieval.py | 193 +++++++++++++++++++++++++++++++++++ 1 file changed, 193 insertions(+) create mode 100644 lnp_ml/modeling/retrieval.py diff --git a/lnp_ml/modeling/retrieval.py b/lnp_ml/modeling/retrieval.py new file mode 100644 index 0000000..45b74e4 --- /dev/null +++ b/lnp_ml/modeling/retrieval.py @@ -0,0 +1,193 @@ +""" +检索增强特征模块 (第一层优化)。 + +核心思想 (借鉴 MolRAG 的化学硬先验「结构相似 → 性质相似」): + 对每个分子,用 Morgan 指纹在「训练集检索池」中找最相似的 k 个邻居, + 把它们的已知标签 (已 z-score 的 quantified_delivery 等) 聚合成一个特征向量, + 作为额外旁路注入融合层 (配零初始化门控,保证「加了不会更差」)。 + +防数据泄漏铁律 (本模块通过接口与逻辑强制保证): + 1. 检索池 (pool_smiles / pool_labels) 只能由调用方传入「训练集」分子。 + 本模块自身不接触任何全局数据,只用传进来的池。 + 2. exclude_self=True 时,查询分子若在池中 (训练分子查询自己), + 排除指纹完全相同的「自己」(相似度=1.0 的那一个),否则等于直接看答案。 + 3. 测试分子查询时,池里全是训练集,测试分子不在池中, + 天然不会检索到自己或其他测试分子。 + +设计原则: + - 纯特征工程,不引入可训练参数 (投影层在 models.py 侧)。 + - 池构建时一次性预计算所有指纹,查询时只算查询分子指纹 + 相似度,高效。 + - 标签已在 dataset.py 标准化 (z-score),聚合时直接 (加权) 平均即可。 +""" + +from __future__ import annotations + +from typing import List, Optional, Sequence + +import numpy as np + +from rdkit import Chem +from rdkit.Chem import AllChem, DataStructs +from rdkit import RDLogger + +# 关闭 RDKit 的噪音日志 (无效 SMILES 警告等) +RDLogger.DisableLog("rdApp.*") + + +def _smiles_to_fp(smiles: str, radius: int = 2, n_bits: int = 2048): + """单个 SMILES -> Morgan 指纹 (ExplicitBitVect)。无效 SMILES 返回 None。""" + mol = Chem.MolFromSmiles(smiles) + if mol is None: + return None + return AllChem.GetMorganFingerprintAsBitVect(mol, radius=radius, nBits=n_bits) + + +class MorganRetriever: + """基于 Morgan 指纹 + Tanimoto 相似度的分子检索器。 + + 用法: + # 只用训练集分子建池 + retr = MorganRetriever( + pool_smiles=train_smiles, # List[str] + pool_labels=train_labels, # np.ndarray [N_pool, label_dim] + radius=2, n_bits=2048, k=5, + ) + # 查询 (训练分子查询时 exclude_self=True 防泄漏) + feat = retr.query(query_smiles, exclude_self=True) # np.ndarray [label_dim*?] + """ + + def __init__( + self, + pool_smiles: Sequence[str], + pool_labels: np.ndarray, + radius: int = 2, + n_bits: int = 2048, + k: int = 5, + sim_weighted: bool = True, + include_sim_stats: bool = True, + ) -> None: + """ + Args: + pool_smiles: 检索池分子的 SMILES (必须只含训练集分子)。 + pool_labels: 检索池分子的标签, shape [N_pool, label_dim], 已标准化。 + radius: Morgan 指纹半径 (MolRAG 用 2)。 + n_bits: 指纹位数。 + k: 检索近邻数。 + sim_weighted: 聚合邻居标签时是否按相似度加权 (否则等权平均)。 + include_sim_stats: 是否在输出特征里附带相似度统计 (top1/mean 相似度), + 让模型知道「邻居有多可靠」。 + """ + assert len(pool_smiles) == len(pool_labels), "池 SMILES 与标签数量不一致" + self.radius = radius + self.n_bits = n_bits + self.k = k + self.sim_weighted = sim_weighted + self.include_sim_stats = include_sim_stats + + _smiles_arr = list(pool_smiles) + _labels_arr = np.asarray(pool_labels, dtype=np.float32) + # 过滤掉标签含 NaN 的分子:它们没有有效标签,不能作为性质参考邻居 + _valid_mask = ~np.isnan(_labels_arr).any(axis=1) + self.pool_smiles: List[str] = [_smiles_arr[i] for i in range(len(_smiles_arr)) if _valid_mask[i]] + self.pool_labels = _labels_arr[_valid_mask] + self.label_dim = self.pool_labels.shape[1] + _n_dropped = int((~_valid_mask).sum()) + if _n_dropped > 0: + import sys + print(f'[MorganRetriever] 过滤 {_n_dropped} 个 NaN 标签分子,有效检索池={len(self.pool_smiles)}', file=sys.stderr) + + # 池内 SMILES 集合,用于 exclude_self="auto" 时判断查询分子是否在池中 + self.pool_smiles_set = set(self.pool_smiles) + + # 预计算池内所有指纹 (无效分子记 None,检索时跳过) + self.pool_fps = [_smiles_to_fp(s, radius, n_bits) for s in self.pool_smiles] + + # 标签均值,用于无有效邻居时的回退 (用训练集均值,不泄漏) + self.label_mean = self.pool_labels.mean(axis=0) + + @property + def feature_dim(self) -> int: + """输出特征维度: 聚合标签 (label_dim) [+ 相似度统计 2 维]。""" + return self.label_dim + (2 if self.include_sim_stats else 0) + + def query(self, query_smiles: str, exclude_self="auto") -> np.ndarray: + """检索查询分子的 top-k 邻居并聚合其标签为特征向量。 + + Args: + query_smiles: 查询分子 SMILES。 + exclude_self: 排除自身策略,防数据泄漏的关键。 + - "auto" (默认, 推荐): 查询分子若在检索池中(训练分子),自动排除自己; + 不在池中(测试分子)则不排除。无需调用方区分训练/测试。 + - True: 强制排除相似度=1.0 的分子。 + - False: 不排除(仅当确定查询分子不在池中时使用)。 + + Returns: + 特征向量 np.ndarray [feature_dim]。无有效邻居时回退到训练集均值。 + """ + # auto 模式:查询分子在池中 → 训练分子 → 排除自己;否则不排除 + if exclude_self == "auto": + do_exclude = query_smiles in self.pool_smiles_set + else: + do_exclude = bool(exclude_self) + + q_fp = _smiles_to_fp(query_smiles, self.radius, self.n_bits) + if q_fp is None: + # 查询分子无效:回退到训练集均值 + 零相似度 + agg = self.label_mean.copy() + if self.include_sim_stats: + agg = np.concatenate([agg, np.array([0.0, 0.0], dtype=np.float32)]) + return agg.astype(np.float32) + + # 与池内每个分子算 Tanimoto 相似度 + sims = np.full(len(self.pool_fps), -1.0, dtype=np.float32) + for i, fp in enumerate(self.pool_fps): + if fp is None: + continue + sims[i] = DataStructs.TanimotoSimilarity(q_fp, fp) + + # 排除自己: 指纹完全相同 (相似度 >= 1.0 - eps) 的那一个 + if do_exclude: + eps = 1e-6 + self_mask = sims >= (1.0 - eps) + # 保守起见:把所有相似度=1.0 的都视作潜在自身并排除, + # 因为结构完全相同的分子标签也应相同,留着等于泄漏答案。 + sims[self_mask] = -1.0 + + # 取 top-k (相似度降序),过滤掉无效 (-1) 的 + valid = np.where(sims >= 0.0)[0] + if len(valid) == 0: + # 没有有效邻居:回退到训练集均值 + agg = self.label_mean.copy() + if self.include_sim_stats: + agg = np.concatenate([agg, np.array([0.0, 0.0], dtype=np.float32)]) + return agg.astype(np.float32) + + order = valid[np.argsort(-sims[valid])] + topk_idx = order[: self.k] + topk_sims = sims[topk_idx] + topk_labels = self.pool_labels[topk_idx] # [k', label_dim] + + # 聚合邻居标签 + if self.sim_weighted and topk_sims.sum() > 1e-8: + w = topk_sims / topk_sims.sum() + agg_label = (w[:, None] * topk_labels).sum(axis=0) + else: + agg_label = topk_labels.mean(axis=0) + + if self.include_sim_stats: + sim_stats = np.array( + [float(topk_sims[0]), float(topk_sims.mean())], dtype=np.float32 + ) + agg = np.concatenate([agg_label.astype(np.float32), sim_stats]) + else: + agg = agg_label.astype(np.float32) + + return agg.astype(np.float32) + + def query_batch( + self, query_smiles_list, exclude_self="auto" + ) -> np.ndarray: + """批量检索。返回 [N_query, feature_dim]。""" + return np.stack( + [self.query(s, exclude_self=exclude_self) for s in query_smiles_list] + ) From 987ecf6386e505ee747642ce972f2c01fa25b4dc Mon Sep 17 00:00:00 2001 From: DicongLi <2024379585@qq.com> Date: Sun, 5 Jul 2026 01:07:26 +0800 Subject: [PATCH 4/4] =?UTF-8?q?feat:=20=E5=90=8C=E6=AD=A5=20RAG=20?= =?UTF-8?q?=E8=9E=8D=E5=90=88=E4=B8=8E=E8=AE=AD=E7=BB=83=E6=A0=B8=E5=BF=83?= =?UTF-8?q?=E4=BB=A3=E7=A0=81(fusion=20=E6=94=AF=E6=8C=81=20f=5Fretr?= =?UTF-8?q?=E3=80=81models/trainer/moe/dataset)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- lnp_ml/modeling/layers/fusion.py | 306 +++++++++++---------- lnp_ml/modeling/layers/moe.py | 456 +++++++++++++++---------------- 2 files changed, 383 insertions(+), 379 deletions(-) diff --git a/lnp_ml/modeling/layers/fusion.py b/lnp_ml/modeling/layers/fusion.py index 70ef066..3824eb8 100644 --- a/lnp_ml/modeling/layers/fusion.py +++ b/lnp_ml/modeling/layers/fusion.py @@ -1,152 +1,156 @@ -import torch -import torch.nn as nn -import torch.nn.functional as F -from typing import Dict, List, Literal, Optional, Tuple, Union - - -PoolingStrategy = Literal["concat", "avg", "max", "attention"] - - -class FusionLayer(nn.Module): - """ - 将多个 token 融合成单个向量。 - - 输入: Dict[str, Tensor] 或 [B, n_tokens, d_model] - 输出: [B, fusion_dim] - - 策略: - - concat: [B, n_tokens, d_model] -> [B, n_tokens * d_model] - - avg: [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) - """ - - def __init__( - self, - d_model: int, - n_tokens: int, - strategy: PoolingStrategy = "attention", - ) -> None: - """ - Args: - d_model: 每个 token 的维度 - n_tokens: token 数量(如 8) - strategy: 融合策略 - """ - super().__init__() - self.d_model = d_model - self.n_tokens = n_tokens - self.strategy = strategy - - if strategy == "concat": - self.fusion_dim = n_tokens * d_model - else: - self.fusion_dim = d_model - - # Attention pooling: learnable query - if strategy == "attention": - self.attn_query = nn.Parameter(torch.randn(1, 1, d_model)) - self.attn_proj = nn.Linear(d_model, d_model) - - def forward( - self, - x: Union[Dict[str, torch.Tensor], torch.Tensor], - return_attn_weights: bool = False, - ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: - """ - Args: - x: Dict[str, Tensor] 每个 [B, d_model],或已 stack 的 [B, n_tokens, d_model] - return_attn_weights: 若为 True 且策略为 attention,额外返回 attn_weights [B, n_tokens] - - Returns: - return_attn_weights=False: [B, fusion_dim] - return_attn_weights=True: ([B, fusion_dim], [B, n_tokens]) - """ - if isinstance(x, dict): - x = torch.stack(list(x.values()), dim=1) - - if self.strategy == "concat": - out = x.flatten(start_dim=1) - return (out, None) if return_attn_weights else out - - elif self.strategy == "avg": - out = x.mean(dim=1) - return (out, None) if return_attn_weights else out - - elif self.strategy == "max": - out = x.max(dim=1).values - return (out, None) if return_attn_weights else out - - elif self.strategy == "attention": - return self._attention_pooling(x, return_attn_weights) - - else: - raise ValueError(f"Unknown strategy: {self.strategy}") - - def _attention_pooling( - self, x: torch.Tensor, return_attn_weights: bool = False, - ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: - """ - Attention pooling: 用可学习 query 对 tokens 做加权求和 - - Args: - x: [B, n_tokens, d_model] - return_attn_weights: 是否返回权重 - - Returns: - return_attn_weights=False: [B, d_model] - return_attn_weights=True: ([B, d_model], [B, n_tokens]) - """ - B = x.size(0) - query = self.attn_query.expand(B, -1, -1) - - keys = self.attn_proj(x) - scores = torch.bmm(query, keys.transpose(1, 2)) / (self.d_model ** 0.5) - attn_weights = F.softmax(scores, dim=-1) # [B, 1, n_tokens] - - out = torch.bmm(attn_weights, x).squeeze(1) # [B, d_model] - - if return_attn_weights: - return out, attn_weights.squeeze(1) # [B, n_tokens] - return out - - -class ResidualConcatFusion(nn.Module): - """对真实 token 做 attention pooling,再用零初始化门把 MoE/LLM 旁路以残差方式加入。 - - g_moe / g_llm 初始为 0 → +moe/+llm 起点严格等于 baseline; - 旁路只有确实有用时才会被训练打开,从机制上保证“加了不会更差”。 - """ - - def __init__(self, d_model: int, strategy: PoolingStrategy = "attention") -> None: - super().__init__() - if strategy == "concat": - raise ValueError("ResidualConcatFusion 不支持 concat(token 数随开关变化)") - self.d_model = d_model - self.pool = FusionLayer(d_model=d_model, n_tokens=1, strategy=strategy) - self.fusion_dim = self.pool.fusion_dim - # 零初始化门控(可学习标量),旁路初始不参与 - self.g_moe = nn.Parameter(torch.zeros(())) - self.g_llm = nn.Parameter(torch.zeros(())) - - def forward( - self, - chem: torch.Tensor, - tab: torch.Tensor, - f_moe: Optional[torch.Tensor] = None, - f_llm: Optional[torch.Tensor] = None, - return_attn_weights: bool = False, - ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: - # 只对真实 token(chem + tab)做注意力池化,旁路不参与 softmax 竞争 - seq = torch.cat([chem, tab], dim=1) # [B, n_chem + n_cond, d_model] - pooled = self.pool(seq, return_attn_weights=return_attn_weights) - if return_attn_weights: - pooled, attn = pooled - - out = pooled - if f_moe is not None: - out = out + self.g_moe * f_moe # 残差 + 零初始化门 - if f_llm is not None: - out = out + self.g_llm * f_llm - +import torch +import torch.nn as nn +import torch.nn.functional as F +from typing import Dict, List, Literal, Optional, Tuple, Union + + +PoolingStrategy = Literal["concat", "avg", "max", "attention"] + + +class FusionLayer(nn.Module): + """ + 将多个 token 融合成单个向量。 + + 输入: Dict[str, Tensor] 或 [B, n_tokens, d_model] + 输出: [B, fusion_dim] + + 策略: + - concat: [B, n_tokens, d_model] -> [B, n_tokens * d_model] + - avg: [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) + """ + + def __init__( + self, + d_model: int, + n_tokens: int, + strategy: PoolingStrategy = "attention", + ) -> None: + """ + Args: + d_model: 每个 token 的维度 + n_tokens: token 数量(如 8) + strategy: 融合策略 + """ + super().__init__() + self.d_model = d_model + self.n_tokens = n_tokens + self.strategy = strategy + + if strategy == "concat": + self.fusion_dim = n_tokens * d_model + else: + self.fusion_dim = d_model + + # Attention pooling: learnable query + if strategy == "attention": + self.attn_query = nn.Parameter(torch.randn(1, 1, d_model)) + self.attn_proj = nn.Linear(d_model, d_model) + + def forward( + self, + x: Union[Dict[str, torch.Tensor], torch.Tensor], + return_attn_weights: bool = False, + ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: + """ + Args: + x: Dict[str, Tensor] 每个 [B, d_model],或已 stack 的 [B, n_tokens, d_model] + return_attn_weights: 若为 True 且策略为 attention,额外返回 attn_weights [B, n_tokens] + + Returns: + return_attn_weights=False: [B, fusion_dim] + return_attn_weights=True: ([B, fusion_dim], [B, n_tokens]) + """ + if isinstance(x, dict): + x = torch.stack(list(x.values()), dim=1) + + if self.strategy == "concat": + out = x.flatten(start_dim=1) + return (out, None) if return_attn_weights else out + + elif self.strategy == "avg": + out = x.mean(dim=1) + return (out, None) if return_attn_weights else out + + elif self.strategy == "max": + out = x.max(dim=1).values + return (out, None) if return_attn_weights else out + + elif self.strategy == "attention": + return self._attention_pooling(x, return_attn_weights) + + else: + raise ValueError(f"Unknown strategy: {self.strategy}") + + def _attention_pooling( + self, x: torch.Tensor, return_attn_weights: bool = False, + ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: + """ + Attention pooling: 用可学习 query 对 tokens 做加权求和 + + Args: + x: [B, n_tokens, d_model] + return_attn_weights: 是否返回权重 + + Returns: + return_attn_weights=False: [B, d_model] + return_attn_weights=True: ([B, d_model], [B, n_tokens]) + """ + B = x.size(0) + query = self.attn_query.expand(B, -1, -1) + + keys = self.attn_proj(x) + scores = torch.bmm(query, keys.transpose(1, 2)) / (self.d_model ** 0.5) + attn_weights = F.softmax(scores, dim=-1) # [B, 1, n_tokens] + + out = torch.bmm(attn_weights, x).squeeze(1) # [B, d_model] + + if return_attn_weights: + return out, attn_weights.squeeze(1) # [B, n_tokens] + return out + + +class ResidualConcatFusion(nn.Module): + """对真实 token 做 attention pooling,再用零初始化门把 MoE/LLM 旁路以残差方式加入。 + + g_moe / g_llm 初始为 0 → +moe/+llm 起点严格等于 baseline; + 旁路只有确实有用时才会被训练打开,从机制上保证“加了不会更差”。 + """ + + def __init__(self, d_model: int, strategy: PoolingStrategy = "attention") -> None: + super().__init__() + if strategy == "concat": + raise ValueError("ResidualConcatFusion 不支持 concat(token 数随开关变化)") + self.d_model = d_model + self.pool = FusionLayer(d_model=d_model, n_tokens=1, strategy=strategy) + self.fusion_dim = self.pool.fusion_dim + # 零初始化门控(可学习标量),旁路初始不参与 + self.g_moe = nn.Parameter(torch.zeros(())) + self.g_llm = nn.Parameter(torch.zeros(())) + self.g_retr = nn.Parameter(torch.zeros(())) # 检索旁路零初始化门控 + + def forward( + self, + chem: torch.Tensor, + tab: torch.Tensor, + f_moe: Optional[torch.Tensor] = None, + f_llm: Optional[torch.Tensor] = None, + f_retr: Optional[torch.Tensor] = None, + return_attn_weights: bool = False, + ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: + # 只对真实 token(chem + tab)做注意力池化,旁路不参与 softmax 竞争 + seq = torch.cat([chem, tab], dim=1) # [B, n_chem + n_cond, d_model] + pooled = self.pool(seq, return_attn_weights=return_attn_weights) + if return_attn_weights: + pooled, attn = pooled + + out = pooled + if f_moe is not None: + 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 \ No newline at end of file diff --git a/lnp_ml/modeling/layers/moe.py b/lnp_ml/modeling/layers/moe.py index 8a07c05..aab389d 100644 --- a/lnp_ml/modeling/layers/moe.py +++ b/lnp_ml/modeling/layers/moe.py @@ -1,229 +1,229 @@ -""" -MoE (Mixture-of-Experts) 模块:sample-level + 跨模态路由。 - -设计要点(与现有 8-token 架构对齐): - - Router 输入:tab token 池化后的 [B, d](即配方/实验条件向量)。 - - Expert 输入:chem token flatten 后的 [B, T_chem * d](化学侧 token)。 - - 路由粒度:每个样本只算一次 router → gates [B, K]。 - - Top-k 稀疏激活 + 可选训练态 jitter 噪声。 - - 返回 load-balancing aux loss,由 trainer 端按权重加进总 loss。 - -输出: - F_moe: [B, d_model] 与单个 token 同维度,便于追加到 fusion 序列中。 - extras: dict 诊断与监控信息(aux loss、gates 等)。 -""" - -from typing import Dict, Tuple - -import torch -import torch.nn as nn -import torch.nn.functional as F - - -class MoEAttentionPool(nn.Module): - """与 FusionLayer 同款的 attention pooling,参数独立。 - - 将一组 token [B, T, d] 池化成单个查询向量 [B, d],用作 router 的输入。 - """ - - def __init__(self, d_model: int) -> None: - super().__init__() - self.d_model = d_model - self.query = nn.Parameter(torch.randn(1, 1, d_model)) - self.proj = nn.Linear(d_model, d_model) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - """ - Args: - x: [B, T, d_model] - - Returns: - [B, d_model] - """ - B = x.size(0) - q = self.query.expand(B, -1, -1) # [B, 1, d] - k = self.proj(x) # [B, T, d] - scores = torch.bmm(q, k.transpose(1, 2)) / (self.d_model ** 0.5) - weights = F.softmax(scores, dim=-1) # [B, 1, T] - return torch.bmm(weights, x).squeeze(1) # [B, d] - - -class MoERouter(nn.Module): - """单层 softmax router + Top-k 稀疏激活。 - - Args: - d_model: router 输入维度。 - n_experts: 专家数量 K。 - top_k: 每个样本激活的专家数。 - jitter_noise: 训练态加在 logits 上的均匀噪声幅度,用于探索;评估态自动关闭。 - """ - - def __init__( - self, - d_model: int, - n_experts: int, - top_k: int = 2, - jitter_noise: float = 0.0, - ) -> None: - super().__init__() - if not 1 <= top_k <= n_experts: - raise ValueError(f"top_k 必须在 [1, n_experts={n_experts}],收到 {top_k}") - self.n_experts = n_experts - self.top_k = top_k - self.jitter_noise = jitter_noise - self.linear = nn.Linear(d_model, n_experts) - - def forward( - self, q: torch.Tensor, - ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """ - Args: - q: [B, d_model] 路由查询向量。 - - Returns: - gates: [B, n_experts] top-k 后再归一化的稀疏概率。 - probs_full: [B, n_experts] 未掩码的 softmax 概率(用于 aux loss)。 - expert_mask: [B, n_experts] 0/1 掩码,标记被激活的专家。 - """ - logits = self.linear(q) # [B, K] - if self.training and self.jitter_noise > 0.0: - noise = (torch.rand_like(logits) - 0.5) * self.jitter_noise - logits = logits + noise - - probs_full = F.softmax(logits, 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.scatter_(1, topk_idx, 1.0) - - gates = probs_full * expert_mask - gates = gates / gates.sum(dim=-1, keepdim=True).clamp(min=1e-9) - return gates, probs_full, expert_mask - - -class MoEExpert(nn.Module): - """单个专家 MLP:in_dim → hidden_dim → out_dim。""" - - def __init__( - self, in_dim: int, hidden_dim: int, out_dim: int, dropout: float = 0.1, - ) -> None: - super().__init__() - self.net = nn.Sequential( - nn.Linear(in_dim, hidden_dim), - nn.GELU(), - nn.Dropout(dropout), - nn.Linear(hidden_dim, out_dim), - ) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.net(x) - - -class MoEBlock(nn.Module): - """ - Sample-level 跨模态 MoE。 - - 流程: - 1. tab tokens [B, T_tab, d] -> attention pool -> q_tab [B, d] - 2. q_tab -> Router -> gates [B, K] (+ probs_full / expert_mask) - 3. chem tokens [B, T_chem, d] -> flatten -> [B, T_chem * d] - 4. K 个 Expert MLP 并行处理 -> stack [B, K, d] - 5. 加权求和 -> F_moe [B, d] - 6. 计算 load-balancing aux loss - - Args: - d_model: token 维度。 - n_chem_tokens: chem 侧 token 数(决定 expert 输入维度)。 - n_experts: 专家数量 K。 - top_k: 每个样本激活的专家数。 - expert_hidden_mult: expert 中间层维度 = expert_hidden_mult * d_model。 - dropout: expert 内部 dropout。 - jitter_noise: router 训练态噪声幅度,0.0 表示关闭。 - """ - - def __init__( - self, - d_model: int, - n_chem_tokens: int = 4, - n_experts: int = 4, - top_k: int = 2, - expert_hidden_mult: int = 2, - dropout: float = 0.1, - jitter_noise: float = 0.0, - ) -> None: - super().__init__() - self.d_model = d_model - self.n_chem_tokens = n_chem_tokens - self.n_experts = n_experts - self.top_k = top_k - - self.tab_pool = MoEAttentionPool(d_model) - self.router = MoERouter( - d_model, n_experts, top_k=top_k, jitter_noise=jitter_noise, - ) - - expert_in = d_model * n_chem_tokens - expert_hidden = d_model * expert_hidden_mult - self.experts = nn.ModuleList([ - MoEExpert(expert_in, expert_hidden, d_model, dropout=dropout) - for _ in range(n_experts) - ]) - - def forward( - self, - chem: torch.Tensor, - tab: torch.Tensor, - ) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]: - """ - Args: - chem: [B, T_chem, d_model] 化学侧 token(router 不看)。 - tab: [B, T_tab, d_model] 配方/实验侧 token(router 看)。 - - Returns: - F_moe: [B, d_model] - extras: { - "lb_loss": 标量 load-balancing aux loss(带梯度), - "gates": [B, K] 稀疏归一后的门控(detached), - "probs": [B, K] 原始 softmax 概率(detached), - } - """ - if chem.size(1) != self.n_chem_tokens: - raise ValueError( - f"chem token 数不匹配:期望 {self.n_chem_tokens}, " - f"实际 {chem.size(1)}" - ) - - q_tab = self.tab_pool(tab) # [B, d] - gates, probs_full, expert_mask = self.router(q_tab) # 各 [B, K] - - flat = chem.flatten(start_dim=1) # [B, T_chem * d] - expert_outs = torch.stack( - [expert(flat) for expert in self.experts], dim=1, - ) # [B, K, d] - - F_moe = (gates.unsqueeze(-1) * expert_outs).sum(dim=1) # [B, d] - - lb_loss = self._load_balancing_loss(probs_full, expert_mask) - - return F_moe, { - "lb_loss": lb_loss, - "gates": gates.detach(), - "probs": probs_full.detach(), - } - - def _load_balancing_loss( - self, - probs_full: torch.Tensor, - expert_mask: torch.Tensor, - ) -> torch.Tensor: - """Switch Transformer 风格的 load-balancing loss。 - - f_i = 该 batch 中 expert_i 被激活的样本占比(top-k mask 的均值) - p_i = 该 batch 中 expert_i 的 softmax 概率均值 - loss = K * Σ_i f_i * p_i - - 理想情况下 f_i 与 p_i 都接近 1/K,loss ≈ 1。 - """ - f = expert_mask.mean(dim=0) # [K] - p = probs_full.mean(dim=0) # [K] +""" +MoE (Mixture-of-Experts) 模块:sample-level + 跨模态路由。 + +设计要点(与现有 8-token 架构对齐): + - Router 输入:tab token 池化后的 [B, d](即配方/实验条件向量)。 + - Expert 输入:chem token flatten 后的 [B, T_chem * d](化学侧 token)。 + - 路由粒度:每个样本只算一次 router → gates [B, K]。 + - Top-k 稀疏激活 + 可选训练态 jitter 噪声。 + - 返回 load-balancing aux loss,由 trainer 端按权重加进总 loss。 + +输出: + F_moe: [B, d_model] 与单个 token 同维度,便于追加到 fusion 序列中。 + extras: dict 诊断与监控信息(aux loss、gates 等)。 +""" + +from typing import Dict, Tuple + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +class MoEAttentionPool(nn.Module): + """与 FusionLayer 同款的 attention pooling,参数独立。 + + 将一组 token [B, T, d] 池化成单个查询向量 [B, d],用作 router 的输入。 + """ + + def __init__(self, d_model: int) -> None: + super().__init__() + self.d_model = d_model + self.query = nn.Parameter(torch.randn(1, 1, d_model)) + self.proj = nn.Linear(d_model, d_model) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """ + Args: + x: [B, T, d_model] + + Returns: + [B, d_model] + """ + B = x.size(0) + q = self.query.expand(B, -1, -1) # [B, 1, d] + k = self.proj(x) # [B, T, d] + scores = torch.bmm(q, k.transpose(1, 2)) / (self.d_model ** 0.5) + weights = F.softmax(scores, dim=-1) # [B, 1, T] + return torch.bmm(weights, x).squeeze(1) # [B, d] + + +class MoERouter(nn.Module): + """单层 softmax router + Top-k 稀疏激活。 + + Args: + d_model: router 输入维度。 + n_experts: 专家数量 K。 + top_k: 每个样本激活的专家数。 + jitter_noise: 训练态加在 logits 上的均匀噪声幅度,用于探索;评估态自动关闭。 + """ + + def __init__( + self, + d_model: int, + n_experts: int, + top_k: int = 2, + jitter_noise: float = 0.0, + ) -> None: + super().__init__() + if not 1 <= top_k <= n_experts: + raise ValueError(f"top_k 必须在 [1, n_experts={n_experts}],收到 {top_k}") + self.n_experts = n_experts + self.top_k = top_k + self.jitter_noise = jitter_noise + self.linear = nn.Linear(d_model, n_experts) + + def forward( + self, q: torch.Tensor, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """ + Args: + q: [B, d_model] 路由查询向量。 + + Returns: + gates: [B, n_experts] top-k 后再归一化的稀疏概率。 + probs_full: [B, n_experts] 未掩码的 softmax 概率(用于 aux loss)。 + expert_mask: [B, n_experts] 0/1 掩码,标记被激活的专家。 + """ + logits = self.linear(q) # [B, K] + if self.training and self.jitter_noise > 0.0: + noise = (torch.rand_like(logits) - 0.5) * self.jitter_noise + logits = logits + noise + + probs_full = F.softmax(logits, 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.scatter_(1, topk_idx, 1.0) + + gates = probs_full * expert_mask + gates = gates / gates.sum(dim=-1, keepdim=True).clamp(min=1e-9) + return gates, probs_full, expert_mask + + +class MoEExpert(nn.Module): + """单个专家 MLP:in_dim → hidden_dim → out_dim。""" + + def __init__( + self, in_dim: int, hidden_dim: int, out_dim: int, dropout: float = 0.1, + ) -> None: + super().__init__() + self.net = nn.Sequential( + nn.Linear(in_dim, hidden_dim), + nn.GELU(), + nn.Dropout(dropout), + nn.Linear(hidden_dim, out_dim), + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.net(x) + + +class MoEBlock(nn.Module): + """ + Sample-level 跨模态 MoE。 + + 流程: + 1. tab tokens [B, T_tab, d] -> attention pool -> q_tab [B, d] + 2. q_tab -> Router -> gates [B, K] (+ probs_full / expert_mask) + 3. chem tokens [B, T_chem, d] -> flatten -> [B, T_chem * d] + 4. K 个 Expert MLP 并行处理 -> stack [B, K, d] + 5. 加权求和 -> F_moe [B, d] + 6. 计算 load-balancing aux loss + + Args: + d_model: token 维度。 + n_chem_tokens: chem 侧 token 数(决定 expert 输入维度)。 + n_experts: 专家数量 K。 + top_k: 每个样本激活的专家数。 + expert_hidden_mult: expert 中间层维度 = expert_hidden_mult * d_model。 + dropout: expert 内部 dropout。 + jitter_noise: router 训练态噪声幅度,0.0 表示关闭。 + """ + + def __init__( + self, + d_model: int, + n_chem_tokens: int = 4, + n_experts: int = 4, + top_k: int = 2, + expert_hidden_mult: int = 2, + dropout: float = 0.1, + jitter_noise: float = 0.0, + ) -> None: + super().__init__() + self.d_model = d_model + self.n_chem_tokens = n_chem_tokens + self.n_experts = n_experts + self.top_k = top_k + + self.tab_pool = MoEAttentionPool(d_model) + self.router = MoERouter( + d_model, n_experts, top_k=top_k, jitter_noise=jitter_noise, + ) + + expert_in = d_model * n_chem_tokens + expert_hidden = d_model * expert_hidden_mult + self.experts = nn.ModuleList([ + MoEExpert(expert_in, expert_hidden, d_model, dropout=dropout) + for _ in range(n_experts) + ]) + + def forward( + self, + chem: torch.Tensor, + tab: torch.Tensor, + ) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]: + """ + Args: + chem: [B, T_chem, d_model] 化学侧 token(router 不看)。 + tab: [B, T_tab, d_model] 配方/实验侧 token(router 看)。 + + Returns: + F_moe: [B, d_model] + extras: { + "lb_loss": 标量 load-balancing aux loss(带梯度), + "gates": [B, K] 稀疏归一后的门控(detached), + "probs": [B, K] 原始 softmax 概率(detached), + } + """ + if chem.size(1) != self.n_chem_tokens: + raise ValueError( + f"chem token 数不匹配:期望 {self.n_chem_tokens}, " + f"实际 {chem.size(1)}" + ) + + q_tab = self.tab_pool(tab) # [B, d] + gates, probs_full, expert_mask = self.router(q_tab) # 各 [B, K] + + flat = chem.flatten(start_dim=1) # [B, T_chem * d] + expert_outs = torch.stack( + [expert(flat) for expert in self.experts], dim=1, + ) # [B, K, d] + + F_moe = (gates.unsqueeze(-1) * expert_outs).sum(dim=1) # [B, d] + + lb_loss = self._load_balancing_loss(probs_full, expert_mask) + + return F_moe, { + "lb_loss": lb_loss, + "gates": gates.detach(), + "probs": probs_full.detach(), + } + + def _load_balancing_loss( + self, + probs_full: torch.Tensor, + expert_mask: torch.Tensor, + ) -> torch.Tensor: + """Switch Transformer 风格的 load-balancing loss。 + + f_i = 该 batch 中 expert_i 被激活的样本占比(top-k mask 的均值) + p_i = 该 batch 中 expert_i 的 softmax 概率均值 + loss = K * Σ_i f_i * p_i + + 理想情况下 f_i 与 p_i 都接近 1/K,loss ≈ 1。 + """ + f = expert_mask.mean(dim=0) # [K] + p = probs_full.mean(dim=0) # [K] return self.n_experts * (f * p).sum() \ No newline at end of file