mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-07-21 21:22:05 +08:00
Merge remote-tracking branch 'origin/feat/moe-layer' into feat/moe-layer
This commit is contained in:
commit
0c6828076b
@ -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]]:
|
||||
# 只对真实 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:
|
||||
@ -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
|
||||
@ -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/o;Roberta/ChemBERTa: query/key/value;Qwen/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:
|
||||
|
||||
@ -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}")
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user