mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-07-22 05:49:55 +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_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,
|
||||||
@ -139,6 +139,7 @@ class ResidualConcatFusion(nn.Module):
|
|||||||
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]]:
|
||||||
|
# 只对真实 token(chem + tab)做注意力池化,旁路不参与 softmax 竞争
|
||||||
seq = torch.cat([chem, tab], dim=1) # [B, n_chem + n_cond, d_model]
|
seq = torch.cat([chem, tab], dim=1) # [B, n_chem + n_cond, d_model]
|
||||||
pooled = self.pool(seq, return_attn_weights=return_attn_weights)
|
pooled = self.pool(seq, return_attn_weights=return_attn_weights)
|
||||||
if return_attn_weights:
|
if return_attn_weights:
|
||||||
@ -146,10 +147,10 @@ class ResidualConcatFusion(nn.Module):
|
|||||||
|
|
||||||
out = pooled
|
out = pooled
|
||||||
if f_moe is not None:
|
if f_moe is not None:
|
||||||
out = out + self.g_moe * f_moe # 残差 + 零初始化门
|
out = out + self.g_moe * f_moe # 残差 + 零初始化门
|
||||||
if f_llm is not None:
|
if f_llm is not None:
|
||||||
out = out + self.g_llm * f_llm
|
out = out + self.g_llm * f_llm
|
||||||
if f_retr is not None:
|
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
|
return (out, attn) if return_attn_weights else out
|
||||||
@ -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:
|
||||||
|
|||||||
@ -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}")
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user