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

@ -128,7 +128,7 @@ class ResidualConcatFusion(nn.Module):
# 零初始化门控(可学习标量),旁路初始不参与
self.g_moe = 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(
self,
@ -139,6 +139,7 @@ class ResidualConcatFusion(nn.Module):
f_retr: Optional[torch.Tensor] = None,
return_attn_weights: bool = False,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
# 只对真实 tokenchem + 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:
@ -150,6 +151,6 @@ class ResidualConcatFusion(nn.Module):
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
out = out + self.g_retr * f_retr # 检索旁路,零初始化门保证起点=不开
return (out, attn) if return_attn_weights else out

View File

@ -76,6 +76,17 @@ class LLMPromptEncoder(nn.Module):
quantization_config=bnb, torch_dtype=torch.bfloat16)
self.hidden_size = self.encoder.config.hidden_size
elif _is_qwen:
# 微调(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
@ -122,6 +133,14 @@ class LLMPromptEncoder(nn.Module):
for p in self.encoder.parameters():
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/oRoberta/ChemBERTa: query/key/valueQwen/Llama: q_proj/k_proj/v_proj/o_proj
_enc_name = type(self.encoder).__name__.lower()
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"]
@ -227,6 +246,39 @@ class LLMPromptEncoder(nn.Module):
if key not in self._prompt_cache:
self._prompt_cache[key] = self._build_rag_prompt(s, self._retrieve_topk(s))
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 编码(带梯度,不缓存特征)----------
def _encode_softrag(self, smiles, chem, tab, device) -> torch.Tensor:

View File

@ -682,6 +682,21 @@ def _run_single_outer_fold(
logger.info(f"[RESUME] fold {outer_fold} 复用已存 best_params跳过内层 Optuna。")
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}")