feat: final_train_optuna_cv 支持 MoE + Qwen2.5-7B QLoRA + soft-RAG,产出可部署 checkpoint

This commit is contained in:
Michelle0574 2026-08-13 10:05:00 +00:00
parent bab8203c98
commit 866941039f
38 changed files with 2956 additions and 1088 deletions

View File

@ -87,6 +87,7 @@ class OptimizeRequest(BaseModel):
top_k: int = Field(default=20, ge=1, le=100, description="Number of top formulations to return") top_k: int = Field(default=20, ge=1, le=100, description="Number of top formulations to return")
num_seeds: Optional[int] = Field(default=None, ge=1, le=500, description="Number of seed points from first iteration (default: top_k * 5)") num_seeds: Optional[int] = Field(default=None, ge=1, le=500, description="Number of seed points from first iteration (default: top_k * 5)")
top_per_seed: int = Field(default=1, ge=1, le=10, description="Number of local best to keep per seed in refinement") top_per_seed: int = Field(default=1, ge=1, le=10, description="Number of local best to keep per seed in refinement")
rerank_top_n: int = Field(default=0, ge=0, le=1000, description="两阶段推理粗筛后用完整模型重排的候选数0 表示关闭")
step_sizes: Optional[List[float]] = Field(default=None, description="Mol ratio step sizes for each iteration (default: [10, 2, 1])") step_sizes: Optional[List[float]] = Field(default=None, description="Mol ratio step sizes for each iteration (default: [10, 2, 1])")
wr_step_sizes: Optional[List[float]] = Field(default=None, description="Weight ratio step sizes for each iteration (default: [5, 2, 1])") wr_step_sizes: Optional[List[float]] = Field(default=None, description="Weight ratio step sizes for each iteration (default: [5, 2, 1])")
comp_ranges: Optional[CompRangesRequest] = Field(default=None, description="组分范围配置(默认使用标准范围)") comp_ranges: Optional[CompRangesRequest] = Field(default=None, description="组分范围配置(默认使用标准范围)")
@ -188,6 +189,20 @@ async def lifespan(app: FastAPI):
try: try:
state.model = load_model(model_path, state.device) state.model = load_model(model_path, state.device)
logger.success("Model loaded successfully!") logger.success("Model loaded successfully!")
# RAG 检索池:服务期用全部内部数据(无泄漏顾虑,查询分子会被 _retrieve_topk 自动排除)
_llm = getattr(state.model, "llm_prompt", None)
if _llm is not None and getattr(_llm, "use_rag", False):
import numpy as np
import pandas as pd
from lnp_ml.dataset import LNPDataset, process_dataframe
from lnp_ml.modeling.nested_cv_optuna import _build_rag_pool
rag_csv = Path(os.environ.get("RAG_POOL_CSV", "data/interim/internal.csv"))
logger.info(f"Building RAG retrieval pool from {rag_csv}...")
_ds = LNPDataset(process_dataframe(pd.read_csv(rag_csv)))
_s, _d, _ex = _build_rag_pool(_ds, np.arange(len(_ds)))
_llm.set_retrieval_pool(_s, _d, pool_id="serve", extra_labels=_ex)
logger.success(f"RAG pool ready: {len(_s)} molecules")
except Exception as e: except Exception as e:
logger.error(f"Failed to load model: {e}") logger.error(f"Failed to load model: {e}")
raise raise
@ -305,6 +320,7 @@ async def optimize_formulation(request: OptimizeRequest):
comp_ranges=comp_ranges, comp_ranges=comp_ranges,
routes=request.routes, routes=request.routes,
scoring_weights=scoring_weights, scoring_weights=scoring_weights,
rerank_top_n=request.rerank_top_n,
batch_size=256, batch_size=256,
) )
@ -350,6 +366,8 @@ async def optimize_formulation(request: OptimizeRequest):
except Exception as e: except Exception as e:
logger.error(f"Optimization failed: {e}") logger.error(f"Optimization failed: {e}")
if getattr(state.model, "llm_prompt", None) is not None:
state.model.set_llm_enabled(True)
raise HTTPException(status_code=500, detail=str(e)) raise HTTPException(status_code=500, detail=str(e))

View File

@ -175,9 +175,7 @@ def call_optimize_api(
# PDI 分类标签 # PDI 分类标签
PDI_CLASS_LABELS = { PDI_CLASS_LABELS = {
0: "<0.2 (优)", 0: "<0.2 (优)",
1: "0.2-0.3 (良)", 1: "≥0.2 (欠佳)",
2: "0.3-0.4 (中)",
3: ">0.4 (差)",
} }
# EE 分类标签 # EE 分类标签

View File

@ -701,6 +701,43 @@ def select_top_k(
return formulations return formulations
def create_dataframe_from_candidates(
smiles: str,
candidates: List[Formulation],
) -> pd.DataFrame:
"""为一批已确定 helper_lipid / route 的候选配方构建 DataFrame一行一个候选"""
frames = []
for f in candidates:
comp = (
f.cationic_lipid_to_mrna_ratio,
f.cationic_lipid_mol_ratio,
f.phospholipid_mol_ratio,
f.cholesterol_mol_ratio,
f.peg_lipid_mol_ratio,
)
frames.append(
create_dataframe_from_formulations(smiles, [comp], [f.helper_lipid], [f.route])
)
return pd.concat(frames, ignore_index=True)
def rerank_with_llm(
smiles: str,
organ: str,
candidates: List[Formulation],
model: torch.nn.Module,
device: torch.device,
scoring_weights: ScoringWeights,
batch_size: int = 16,
) -> List[Formulation]:
"""第二阶段:用完整模型(含 LLM 旁路)对候选集重新预测并排序。"""
if not candidates:
return list(candidates)
df = create_dataframe_from_candidates(smiles, candidates)
df = predict_all(model, df, device, batch_size)
return select_top_k(df, organ, k=len(df), scoring_weights=scoring_weights)
def generate_single_seed_grid( def generate_single_seed_grid(
seed: Formulation, seed: Formulation,
mol_step: float, mol_step: float,
@ -780,6 +817,8 @@ def optimize(
routes: Optional[List[str]] = None, routes: Optional[List[str]] = None,
scoring_weights: Optional[ScoringWeights] = None, scoring_weights: Optional[ScoringWeights] = None,
batch_size: int = 256, batch_size: int = 256,
rerank_top_n: int = 0,
rerank_batch_size: int = 16,
) -> List[Formulation]: ) -> List[Formulation]:
""" """
执行配方优化层级搜索策略 执行配方优化层级搜索策略
@ -844,6 +883,14 @@ def optimize(
seeds = None seeds = None
_has_llm = getattr(model, "llm_prompt", None) is not None
_do_rerank = bool(_has_llm and rerank_top_n and rerank_top_n > 0)
if _has_llm:
model.set_llm_enabled(not _do_rerank)
if _do_rerank:
logger.info(f"Two-stage inference: coarse search w/o LLM, rerank top-{rerank_top_n} w/ LLM")
for iteration, (mol_step, wr_step) in enumerate(zip(step_sizes, wr_step_sizes)): for iteration, (mol_step, wr_step) in enumerate(zip(step_sizes, wr_step_sizes)):
logger.info(f"\n{'='*60}") logger.info(f"\n{'='*60}")
logger.info(f"Iteration {iteration + 1}/{len(step_sizes)}, mol_step={mol_step}, wr_step={wr_step}") logger.info(f"Iteration {iteration + 1}/{len(step_sizes)}, mol_step={mol_step}, wr_step={wr_step}")
@ -924,10 +971,9 @@ def optimize(
logger.info(f"Current best score: {_score(best):.4f} (biodist_{organ}={best.get_biodist(organ):.4f})") logger.info(f"Current best score: {_score(best):.4f} (biodist_{organ}={best.get_biodist(organ):.4f})")
logger.info(f"Best formulation: {best.to_dict()}") logger.info(f"Best formulation: {best.to_dict()}")
# 最终去重、按综合评分排序并返回 top_k # 粗筛结果:去重、按综合评分排序
seeds_sorted = sorted(seeds, key=_score, reverse=True) seeds_sorted = sorted(seeds, key=_score, reverse=True)
# 去重:保留每个唯一配方中得分最高的(已排序,所以第一个出现的就是最高的)
seen_keys = set() seen_keys = set()
unique_results = [] unique_results = []
for f in seeds_sorted: for f in seeds_sorted:
@ -936,9 +982,30 @@ def optimize(
seen_keys.add(key) seen_keys.add(key)
unique_results.append(f) unique_results.append(f)
logger.info(f"Final results: {len(unique_results)} unique formulations (from {len(seeds)} candidates)") logger.info(f"Stage-1: {len(unique_results)} unique formulations (from {len(seeds)} candidates)")
return unique_results[:top_k] if not _do_rerank:
return unique_results[:top_k]
# ==================== 第二阶段:完整模型重排 ====================
pool = unique_results[: max(rerank_top_n, top_k)]
logger.info(f"Stage-2: reranking {len(pool)} candidates with LLM enabled...")
model.set_llm_enabled(True)
reranked = rerank_with_llm(
smiles, organ, pool, model, device, scoring_weights, rerank_batch_size,
)
# 监控:粗筛与重排的秩相关。若长期接近 1说明 LLM 旁路对排序无贡献,可直接关闭
try:
from scipy.stats import spearmanr
before = {f.unique_key(): i for i, f in enumerate(pool)}
after = [before[f.unique_key()] for f in reranked if f.unique_key() in before]
rho = spearmanr(after, range(len(after))).correlation
logger.info(f"Stage-2 rank correlation with stage-1: rho={rho:.4f}")
except Exception:
pass
return reranked[:top_k]
def format_results(formulations: List[Formulation], organ: str) -> pd.DataFrame: def format_results(formulations: List[Formulation], organ: str) -> pd.DataFrame:

View File

@ -53,6 +53,7 @@ from lnp_ml.modeling.trainer_balanced import (
train_fixed_epochs, train_fixed_epochs,
) )
from lnp_ml.modeling.visualization import plot_multitask_loss_curves from lnp_ml.modeling.visualization import plot_multitask_loss_curves
from lnp_ml.modeling.nested_cv_optuna import _build_rag_pool
# MPNN ensemble 默认路径 # MPNN ensemble 默认路径
DEFAULT_MPNN_ENSEMBLE_DIR = MODELS_DIR / "mpnn" / "all_amine_split_for_LiON" DEFAULT_MPNN_ENSEMBLE_DIR = MODELS_DIR / "mpnn" / "all_amine_split_for_LiON"
@ -178,20 +179,36 @@ def create_model(
mpnn_device: str = "cpu", mpnn_device: str = "cpu",
chemeleon_cache: Optional[str] = None, chemeleon_cache: Optional[str] = None,
unimol_cache: Optional[str] = None, unimol_cache: Optional[str] = None,
# ============ MoE 相关(新增) ============ moleculestm_cache: Optional[str] = None,
mole_cache: Optional[str] = None,
set_transformer_block: str = "sab",
# MoE
use_moe: bool = False, use_moe: bool = False,
moe_n_experts: int = 4, moe_n_experts: int = 4,
moe_top_k: int = 2, moe_top_k: int = 2,
moe_expert_hidden_mult: int = 2, moe_expert_hidden_mult: int = 2,
moe_jitter_noise: float = 0.0, moe_jitter_noise: float = 0.0,
# LLM含 reg_bypass / use_rag / soft prompt / QLoRA
llm_kwargs: Optional[Dict] = None,
# 检索增强
use_retrieval: bool = False,
retr_feature_dim: int = 0,
) -> Union[LNPModel, LNPModelWithoutMPNN]: ) -> Union[LNPModel, LNPModelWithoutMPNN]:
"""创建模型""" """创建模型。llm_kwargs 含 use_llm / llm_model_path / llm_freeze / llm_use_lora / llm_lora_* / reg_bypass。"""
moe_kwargs = dict( extra_kwargs = dict(
set_transformer_block=set_transformer_block,
use_moe=use_moe, use_moe=use_moe,
moe_n_experts=moe_n_experts, moe_n_experts=moe_n_experts,
moe_top_k=moe_top_k, moe_top_k=moe_top_k,
moe_expert_hidden_mult=moe_expert_hidden_mult, moe_expert_hidden_mult=moe_expert_hidden_mult,
moe_jitter_noise=moe_jitter_noise, moe_jitter_noise=moe_jitter_noise,
use_retrieval=use_retrieval,
retr_feature_dim=retr_feature_dim,
chemeleon_cache_path=chemeleon_cache,
unimol_cache_path=unimol_cache,
moleculestm_cache_path=moleculestm_cache,
mole_cache_path=mole_cache,
**(llm_kwargs or {}),
) )
if use_mpnn: if use_mpnn:
@ -205,9 +222,7 @@ def create_model(
dropout=dropout, dropout=dropout,
mpnn_ensemble_paths=ensemble_paths, mpnn_ensemble_paths=ensemble_paths,
mpnn_device=mpnn_device, mpnn_device=mpnn_device,
chemeleon_cache_path=chemeleon_cache, **extra_kwargs,
unimol_cache_path=unimol_cache,
**moe_kwargs,
) )
else: else:
return LNPModelWithoutMPNN( return LNPModelWithoutMPNN(
@ -217,9 +232,7 @@ def create_model(
fusion_strategy=fusion_strategy, fusion_strategy=fusion_strategy,
head_hidden_dim=head_hidden_dim, head_hidden_dim=head_hidden_dim,
dropout=dropout, dropout=dropout,
chemeleon_cache_path=chemeleon_cache, **extra_kwargs,
unimol_cache_path=unimol_cache,
**moe_kwargs,
) )
@ -277,6 +290,10 @@ def run_optuna_cv(
moe_top_k: int = 2, moe_top_k: int = 2,
moe_expert_hidden_mult: int = 2, moe_expert_hidden_mult: int = 2,
moe_jitter_noise: float = 0.0, moe_jitter_noise: float = 0.0,
set_transformer_block: str = "sab",
llm_kwargs: Optional[Dict] = None,
use_retrieval: bool = False,
freeze_backbone_epochs: int = 3,
) -> Tuple[Dict, int, optuna.Study]: ) -> Tuple[Dict, int, optuna.Study]:
""" """
使用全量数据做 3-fold CV Optuna 超参搜索 使用全量数据做 3-fold CV Optuna 超参搜索
@ -332,6 +349,23 @@ def run_optuna_cv(
weight_decay = trial.suggest_float("weight_decay", 1e-5, 1e-1, log=True) weight_decay = trial.suggest_float("weight_decay", 1e-5, 1e-1, log=True)
backbone_lr_ratio = trial.suggest_float("backbone_lr_ratio", 0.01, 1.0, log=True) backbone_lr_ratio = trial.suggest_float("backbone_lr_ratio", 0.01, 1.0, log=True)
# MoE / LoRA 结构超参:搜索空间与 nested_cv_optuna 保持一致,便于横向比较
moe_t = dict(
moe_n_experts=moe_n_experts,
moe_top_k=moe_top_k,
moe_expert_hidden_mult=moe_expert_hidden_mult,
)
if use_moe:
moe_t["moe_n_experts"] = trial.suggest_categorical("moe_n_experts", [2, 4, 8])
moe_t["moe_top_k"] = trial.suggest_int("moe_top_k", 1, 2)
moe_t["moe_expert_hidden_mult"] = trial.suggest_categorical(
"moe_expert_hidden_mult", [1, 2])
llm_kwargs_t = dict(llm_kwargs or {})
if llm_kwargs_t.get("use_llm", False):
llm_kwargs_t["llm_use_lora"] = True
llm_kwargs_t["llm_lora_r"] = trial.suggest_categorical("llm_lora_r", [8, 16, 32])
# 3-fold CV # 3-fold CV
cv = StratifiedKFold(n_splits=n_folds, shuffle=True, random_state=seed) cv = StratifiedKFold(n_splits=n_folds, shuffle=True, random_state=seed)
@ -339,7 +373,6 @@ def run_optuna_cv(
fold_best_epochs = [] fold_best_epochs = []
for fold, (train_idx, val_idx) in enumerate(cv.split(indices, strata)): for fold, (train_idx, val_idx) in enumerate(cv.split(indices, strata)):
# 创建 DataLoader
train_subset = Subset(full_dataset, train_idx.tolist()) train_subset = Subset(full_dataset, train_idx.tolist())
val_subset = Subset(full_dataset, val_idx.tolist()) val_subset = Subset(full_dataset, val_idx.tolist())
@ -350,10 +383,8 @@ def run_optuna_cv(
val_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn val_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn
) )
# 计算类权重
class_weights = compute_class_weights_from_loader(train_loader) class_weights = compute_class_weights_from_loader(train_loader)
# 创建模型
model = create_model( model = create_model(
d_model=d_model, d_model=d_model,
num_heads=num_heads, num_heads=num_heads,
@ -365,22 +396,29 @@ def run_optuna_cv(
mpnn_device=device.type, mpnn_device=device.type,
chemeleon_cache=chemeleon_cache, chemeleon_cache=chemeleon_cache,
unimol_cache=unimol_cache, unimol_cache=unimol_cache,
set_transformer_block=set_transformer_block,
use_moe=use_moe, use_moe=use_moe,
moe_n_experts=moe_n_experts,
moe_top_k=moe_top_k,
moe_expert_hidden_mult=moe_expert_hidden_mult,
moe_jitter_noise=moe_jitter_noise, moe_jitter_noise=moe_jitter_noise,
llm_kwargs=llm_kwargs_t,
use_retrieval=use_retrieval,
retr_feature_dim=(3 if use_retrieval else 0),
**moe_t,
) )
if rdkit_cache is not None: if rdkit_cache is not None:
model.rdkit_encoder._cache = rdkit_cache model.rdkit_encoder._cache = rdkit_cache
# 加载预训练权重 # 检索池只用当前 train fold用全量会让验证损失被自己的标签污染
# 选出的超参和 epoch_mean 都会有偏
if llm_kwargs_t.get("use_rag", False) and getattr(model, "llm_prompt", None) is not None:
_s, _d, _ex = _build_rag_pool(full_dataset, train_idx)
model.llm_prompt.set_retrieval_pool(
_s, _d, pool_id=f"opt{fold}", extra_labels=_ex)
if pretrain_state_dict is not None and pretrain_config is not None: if pretrain_state_dict is not None and pretrain_config is not None:
load_pretrain_weights_to_model( load_pretrain_weights_to_model(
model, pretrain_state_dict, d_model, pretrain_config, load_delivery_head model, pretrain_state_dict, d_model, pretrain_config, load_delivery_head
) )
# 训练(带早停)
result = train_with_early_stopping( result = train_with_early_stopping(
model=model, model=model,
train_loader=train_loader, train_loader=train_loader,
@ -392,6 +430,7 @@ def run_optuna_cv(
patience=patience, patience=patience,
class_weights=class_weights, class_weights=class_weights,
backbone_lr_ratio=backbone_lr_ratio, backbone_lr_ratio=backbone_lr_ratio,
freeze_backbone_epochs=freeze_backbone_epochs,
) )
fold_val_losses.append(result["best_val_loss"]) fold_val_losses.append(result["best_val_loss"])
@ -427,6 +466,7 @@ def run_optuna_cv(
"n_attn_layers": fixed_n_attn_layers, "n_attn_layers": fixed_n_attn_layers,
"fusion_strategy": fixed_fusion_strategy, "fusion_strategy": fixed_fusion_strategy,
"head_hidden_dim": fixed_head_hidden_dim, "head_hidden_dim": fixed_head_hidden_dim,
"set_transformer_block": set_transformer_block,
}) })
epoch_mean = study.best_trial.user_attrs.get("epoch_mean", epochs_per_trial) epoch_mean = study.best_trial.user_attrs.get("epoch_mean", epochs_per_trial)
@ -472,6 +512,27 @@ def main(
moe_top_k: int = 2, moe_top_k: int = 2,
moe_expert_hidden_mult: int = 2, moe_expert_hidden_mult: int = 2,
moe_jitter_noise: float = 0.0, moe_jitter_noise: float = 0.0,
# Set Transformer
set_transformer_block: str = "sab",
# 回归旁路
reg_bypass: str = "on",
# LLM
use_llm: bool = False,
use_rag: bool = False,
rag_top_k: int = 4,
use_soft_prompt: bool = False,
llm_model_path: str = "models/qwen2.5-7b-instruct",
llm_freeze: bool = True,
llm_use_lora: bool = False,
llm_use_qlora: bool = False,
llm_lora_r: int = 8,
llm_lora_alpha: int = 16,
llm_lora_dropout: float = 0.05,
llm_max_length: int = 1536,
# 检索增强
use_retrieval: bool = False,
# backbone 冻结轮数(必须与 Optuna 阶段一致)
freeze_backbone_epochs: int = 3,
# 设备 # 设备
device: str = "cuda" if torch.cuda.is_available() else "cpu", device: str = "cuda" if torch.cuda.is_available() else "cpu",
): ):
@ -492,6 +553,26 @@ def main(
logger.info(f"Using device: {device}") logger.info(f"Using device: {device}")
device = torch.device(device) device = torch.device(device)
if use_swa and llm_use_qlora:
logger.error("use_swa 与 QLoRA 不兼容SWA 要对全部参数做滑动平均4bit 量化基座无法参与")
raise typer.Exit(1)
llm_kwargs = dict(
reg_bypass=reg_bypass,
use_rag=use_rag,
rag_top_k=rag_top_k,
llm_max_length=llm_max_length,
use_llm=use_llm,
llm_model_path=llm_model_path,
llm_freeze=llm_freeze,
llm_use_lora=llm_use_lora,
llm_use_qlora=llm_use_qlora,
use_soft_prompt=use_soft_prompt,
llm_lora_r=llm_lora_r,
llm_lora_alpha=llm_lora_alpha,
llm_lora_dropout=llm_lora_dropout,
)
# 加载预训练权重(如果指定) # 加载预训练权重(如果指定)
pretrain_state_dict = None pretrain_state_dict = None
pretrain_config = None pretrain_config = None
@ -559,6 +640,10 @@ def main(
moe_top_k=moe_top_k, moe_top_k=moe_top_k,
moe_expert_hidden_mult=moe_expert_hidden_mult, moe_expert_hidden_mult=moe_expert_hidden_mult,
moe_jitter_noise=moe_jitter_noise, moe_jitter_noise=moe_jitter_noise,
set_transformer_block=set_transformer_block,
llm_kwargs=llm_kwargs,
use_retrieval=use_retrieval,
freeze_backbone_epochs=freeze_backbone_epochs,
) )
# 保存最佳参数 # 保存最佳参数
@ -606,7 +691,21 @@ def main(
with open(output_dir / "class_weights.json", "w") as f: with open(output_dir / "class_weights.json", "w") as f:
json.dump(class_weights_info, f, indent=2) json.dump(class_weights_info, f, indent=2)
# 创建模型 # 架构超参一次性解析Optuna 采样值优先于 CLI 默认值。
# 解析结果必须同时喂给 create_model 和 config否则 load_model 回读时结构对不上。
arch = {
"moe_n_experts": best_params.get("moe_n_experts", moe_n_experts),
"moe_top_k": best_params.get("moe_top_k", moe_top_k),
"moe_expert_hidden_mult": best_params.get("moe_expert_hidden_mult", moe_expert_hidden_mult),
"moe_jitter_noise": moe_jitter_noise,
"set_transformer_block": best_params.get("set_transformer_block", set_transformer_block),
}
llm_kwargs_resolved = {
**llm_kwargs,
**({"llm_use_lora": True, "llm_lora_r": best_params["llm_lora_r"]}
if (use_llm and "llm_lora_r" in best_params) else {}),
}
model = create_model( model = create_model(
d_model=best_params["d_model"], d_model=best_params["d_model"],
num_heads=best_params["num_heads"], num_heads=best_params["num_heads"],
@ -619,13 +718,21 @@ def main(
chemeleon_cache=(chemeleon_cache if use_chemeleon else None), chemeleon_cache=(chemeleon_cache if use_chemeleon else None),
unimol_cache=(unimol_cache if use_unimol else None), unimol_cache=(unimol_cache if use_unimol else None),
use_moe=use_moe, use_moe=use_moe,
moe_n_experts=moe_n_experts, llm_kwargs=llm_kwargs_resolved,
moe_top_k=moe_top_k, use_retrieval=use_retrieval,
moe_expert_hidden_mult=moe_expert_hidden_mult, retr_feature_dim=(3 if use_retrieval else 0),
moe_jitter_noise=moe_jitter_noise, **arch,
) )
model.rdkit_encoder._cache = rdkit_cache model.rdkit_encoder._cache = rdkit_cache
# 全量训练阶段检索池用全量_retrieve_topk 按 SMILES 精确匹配排除查询分子自己
# (含它的所有重复行),所以不存在抄自己标签的退化解。这也让训练期的池子规模
# 与服务期完全一致。
if use_rag and getattr(model, "llm_prompt", None) is not None:
_s, _d, _ex = _build_rag_pool(full_dataset, np.arange(len(full_dataset)))
model.llm_prompt.set_retrieval_pool(_s, _d, pool_id="full", extra_labels=_ex)
logger.info(f"[RAG] 检索池={len(_s)} 全量分子")
# 加载预训练权重 # 加载预训练权重
if pretrain_state_dict is not None and pretrain_config is not None: if pretrain_state_dict is not None and pretrain_config is not None:
loaded = load_pretrain_weights_to_model( loaded = load_pretrain_weights_to_model(
@ -635,18 +742,16 @@ def main(
if loaded: if loaded:
logger.info("Loaded pretrain weights for final training") logger.info("Loaded pretrain weights for final training")
# 打印模型信息
n_params_total = sum(p.numel() for p in model.parameters()) n_params_total = sum(p.numel() for p in model.parameters())
n_params_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) n_params_trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
logger.info(f"Model parameters: {n_params_total:,} total, {n_params_trainable:,} trainable") logger.info(f"Model parameters: {n_params_total:,} total, {n_params_trainable:,} trainable")
# 训练(固定 epoch不 early-stop
swa_start = int(epoch_mean * swa_start_ratio) if use_swa else None swa_start = int(epoch_mean * swa_start_ratio) if use_swa else None
train_result = train_fixed_epochs( train_result = train_fixed_epochs(
model=model, model=model,
train_loader=full_loader, train_loader=full_loader,
val_loader=None, # 全量训练,无验证集 val_loader=None,
device=device, device=device,
lr=best_params["lr"], lr=best_params["lr"],
weight_decay=best_params["weight_decay"], weight_decay=best_params["weight_decay"],
@ -656,12 +761,12 @@ def main(
use_swa=use_swa, use_swa=use_swa,
swa_start_epoch=swa_start, swa_start_epoch=swa_start,
backbone_lr_ratio=best_params.get("backbone_lr_ratio", 1.0), backbone_lr_ratio=best_params.get("backbone_lr_ratio", 1.0),
freeze_backbone_epochs=freeze_backbone_epochs,
) )
# 加载最终权重 # QLoRA 4-bit 基座的量化元数据键不在 final_state 里,必须 strict=False
model.load_state_dict(train_result["final_state"]) model.load_state_dict(train_result["final_state"], strict=False)
# 保存模型
config = { config = {
"d_model": best_params["d_model"], "d_model": best_params["d_model"],
"num_heads": best_params["num_heads"], "num_heads": best_params["num_heads"],
@ -675,14 +780,23 @@ def main(
"use_unimol": use_unimol, "use_unimol": use_unimol,
"unimol_cache": unimol_cache if use_unimol else None, "unimol_cache": unimol_cache if use_unimol else None,
"use_moe": use_moe, "use_moe": use_moe,
"moe_n_experts": moe_n_experts, "use_retrieval": use_retrieval,
"moe_top_k": moe_top_k, "retr_feature_dim": (3 if use_retrieval else 0),
"moe_expert_hidden_mult": moe_expert_hidden_mult, **arch,
"moe_jitter_noise": moe_jitter_noise, **llm_kwargs_resolved,
} }
# 只丢 LLM 基座权重(几 GB加载时从 llm_model_path 现读LoRA 适配器必须留下
_full_state = train_result["final_state"]
_slim_state = {
k: v for k, v in _full_state.items()
if (not k.startswith("llm_prompt.encoder")) or ("lora_" in k)
}
_n_lora = sum(1 for k in _slim_state if "lora_" in k)
logger.info(f"保存 checkpoint: {len(_slim_state)} 个张量,其中 LoRA 适配器 {_n_lora}")
torch.save({ torch.save({
"model_state_dict": train_result["final_state"], "model_state_dict": _slim_state,
"config": config, "config": config,
"best_params": best_params, "best_params": best_params,
"epoch_mean": epoch_mean, "epoch_mean": epoch_mean,

View File

@ -80,7 +80,7 @@ class MultiTaskHead(nn.Module):
size_dropout = min(0.5, dropout + 0.2) size_dropout = min(0.5, dropout + 0.2)
self.size_head = RegressionHead(in_dim, hidden_dim, size_dropout) self.size_head = RegressionHead(in_dim, hidden_dim, size_dropout)
# PDI: 2 分类 # PDI: 2 分类dataset.py 已将 4 段 one-hot 折叠为 pdi_4 >= 1
self.pdi_head = ClassificationHead(in_dim, num_classes=2, hidden_dim=hidden_dim, dropout=dropout) self.pdi_head = ClassificationHead(in_dim, num_classes=2, hidden_dim=hidden_dim, dropout=dropout)
# Encapsulation Efficiency: 3 分类 # Encapsulation Efficiency: 3 分类

View File

@ -114,8 +114,10 @@ class FusionLayer(nn.Module):
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 门为逐维向量初始全零
旁路只有确实有用时才会被训练打开从机制上保证加了不会更差 - 起点仍严格等于 baseline"加了不会更差"的保证不变
- 标量门只能整体调音量最优值是"有用维度的收益""噪声维度的损害"之间的妥协
- 标量门的梯度 L/g = Σ_i u_i f_i 会自我抵消逐维门 L/g_i = u_i f_i 不会
""" """
def __init__(self, d_model: int, strategy: PoolingStrategy = "attention") -> None: def __init__(self, d_model: int, strategy: PoolingStrategy = "attention") -> None:
@ -125,10 +127,18 @@ class ResidualConcatFusion(nn.Module):
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(d_model))
self.g_llm = nn.Parameter(torch.zeros(())) self.g_llm = nn.Parameter(torch.zeros(d_model))
self.g_retr = nn.Parameter(torch.zeros(())) # 检索旁路零初始化门控 self.g_retr = nn.Parameter(torch.zeros(d_model))
def gate_norms(self) -> Dict[str, float]:
"""诊断用。标量门时代的 |g| 对应这里的 norm / sqrt(d_model)。"""
return {
"g_moe": float(self.g_moe.norm()),
"g_llm": float(self.g_llm.norm()),
"g_retr": float(self.g_retr.norm()),
}
def forward( def forward(
self, self,
@ -145,13 +155,14 @@ class ResidualConcatFusion(nn.Module):
if return_attn_weights: if return_attn_weights:
pooled, attn = pooled pooled, attn = pooled
# [d_model] 与 [B, d_model] 自动广播
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
# 回归旁路:额外返回 pooled纯 chem+tab不含 f_llm/f_moe/f_retr # 回归旁路:额外返回 pooled纯 chem+tab不含 f_llm/f_moe/f_retr
if return_attn_weights: if return_attn_weights:

View File

@ -12,6 +12,15 @@ DEFAULT_MOLT5_PATH = os.environ.get("MOLT5_PATH", "models/molt5-base")
# 检索池支持的额外多任务(除 delivery 外) # 检索池支持的额外多任务(除 delivery 外)
_EXTRA_TASKS = ["size", "pdi", "ee", "toxic", "biodist"] _EXTRA_TASKS = ["size", "pdi", "ee", "toxic", "biodist"]
# 连续量的分位桶数。边界只从训练检索池统计,不引入测试集分布。
_N_QBINS = 5
# 离散标签的语义,避免 LLM 只看到裸整数
_PDI_LABELS = {0: "<0.2", 1: ">=0.2"}
_EE_LABELS = {0: "<50%", 1: "50-80%", 2: ">80%"}
_TOX_LABELS = {0: "non-toxic", 1: "toxic"}
_ORGANS = ["lymph_nodes", "heart", "liver", "spleen", "lung", "kidney", "muscle"]
class LLMPromptEncoder(nn.Module): class LLMPromptEncoder(nn.Module):
"""用 LLM 编码分子,输出 F_llm [B, d_model]。 """用 LLM 编码分子,输出 F_llm [B, d_model]。
@ -35,7 +44,7 @@ class LLMPromptEncoder(nn.Module):
lora_r: int = 8, lora_r: int = 8,
lora_alpha: int = 16, lora_alpha: int = 16,
lora_dropout: float = 0.05, lora_dropout: float = 0.05,
max_length: int = 256, max_length: int = 1536,
use_rag: bool = False, use_rag: bool = False,
rag_top_k: int = 4, rag_top_k: int = 4,
use_soft_prompt: bool = False, use_soft_prompt: bool = False,
@ -64,6 +73,9 @@ class LLMPromptEncoder(nn.Module):
) )
if _is_qwen and self.tokenizer.pad_token is None: if _is_qwen and self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token self.tokenizer.pad_token = self.tokenizer.eos_token
if use_soft_prompt:
# soft token 后置要求文本右对齐,否则 soft 会落在 PAD 之后
self.tokenizer.padding_side = "left"
if _is_t5: if _is_t5:
self.encoder = T5EncoderModel.from_pretrained(model_name_or_path) self.encoder = T5EncoderModel.from_pretrained(model_name_or_path)
@ -128,6 +140,7 @@ class LLMPromptEncoder(nn.Module):
self._rag_pool_extra = None # dict: task -> (values, valid) self._rag_pool_extra = None # dict: task -> (values, valid)
self._rag_pool_fps = None self._rag_pool_fps = None
self._rag_pool_id: str = "none" self._rag_pool_id: str = "none"
self._rag_pool_qedges: Dict[str, List[float]] = {}
def _apply_lora(self, r, alpha, dropout, prepare_kbit=False): def _apply_lora(self, r, alpha, dropout, prepare_kbit=False):
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
@ -167,6 +180,19 @@ class LLMPromptEncoder(nn.Module):
self._rag_pool_labels = np.asarray(labels, dtype=np.float32).reshape(-1) self._rag_pool_labels = np.asarray(labels, dtype=np.float32).reshape(-1)
self._rag_pool_extra = extra_labels self._rag_pool_extra = extra_labels
self._rag_pool_fps = [_smiles_to_fp(s) for s in self._rag_pool_smiles] self._rag_pool_fps = [_smiles_to_fp(s) for s in self._rag_pool_smiles]
self._rag_pool_qedges = {}
_probs = np.linspace(0.0, 1.0, _N_QBINS + 1)[1:-1]
_cont = {"delivery": self._rag_pool_labels}
if extra_labels is not None and "size" in extra_labels:
_v, _ok = extra_labels["size"]
_cont["size"] = np.asarray(_v, dtype=np.float32)[np.asarray(_ok, dtype=bool)]
for _name, _vals in _cont.items():
_vals = np.asarray(_vals, dtype=np.float32)
_vals = _vals[np.isfinite(_vals)]
if _vals.size >= _N_QBINS:
self._rag_pool_qedges[_name] = np.quantile(_vals, _probs).tolist()
if pool_id != self._rag_pool_id: if pool_id != self._rag_pool_id:
self._cache = {k: v for k, v in self._cache.items() if not k.startswith("RAG::")} self._cache = {k: v for k, v in self._cache.items() if not k.startswith("RAG::")}
self._prompt_cache.clear() self._prompt_cache.clear()
@ -213,6 +239,21 @@ class LLMPromptEncoder(nn.Module):
return "[" + ", ".join(f"{x:.3f}" for x in v) + "]" return "[" + ", ".join(f"{x:.3f}" for x in v) + "]"
return f"{float(v):.3f}" return f"{float(v):.3f}"
def _qbin(self, name: str, v) -> str:
"""v 落在训练池经验分布的第几桶,如 '(Q3/5)';无边界或缺失时返回空串。"""
import bisect
edges = self._rag_pool_qedges.get(name)
if v is None or not edges:
return ""
return f"(Q{bisect.bisect_right(edges, float(v)) + 1}/{len(edges) + 1})"
@staticmethod
def _fmt_class(v, labels: Dict[int, str]) -> str:
"""离散标签附带语义,例如 '1(>=0.2)'"""
if v is None:
return "unknown"
return f"{int(v)}({labels.get(int(v), '?')})"
def _fmt_mol(self, smiles: str) -> str: def _fmt_mol(self, smiles: str) -> str:
"""按 backbone 期望格式化分子。 """按 backbone 期望格式化分子。
BioT5SMILES -> SELFIES <bom>...<eom> 紧贴包裹官方格式token 间无空格 BioT5SMILES -> SELFIES <bom>...<eom> 紧贴包裹官方格式token 间无空格
@ -227,36 +268,54 @@ class LLMPromptEncoder(nn.Module):
return f"<bom>{sfs}<eom>" return f"<bom>{sfs}<eom>"
def _build_rag_prompt(self, target_smiles: str, neighbors) -> str: def _build_rag_prompt(self, target_smiles: str, neighbors) -> str:
"""构造 RAG prompt原始 SMILES + 邻居多任务结果(numeric)。""" """构造 RAG prompt原始 SMILES + 邻居多任务结果(数值 + 训练池分位桶)。"""
blocks = [] blocks = []
for rank, nb in enumerate(neighbors, 1): for rank, nb in enumerate(neighbors, 1):
ex = nb["extra"] ex = nb["extra"]
bio = ex.get("biodist")
if bio is not None:
_top = max(range(len(bio)), key=lambda i: bio[i])
bio_s = (f"{_ORGANS[_top]}-dominant ["
+ ",".join(f"{x:.2f}" for x in bio) + "]")
else:
bio_s = "unknown"
blocks.append( blocks.append(
f"Retrieved sample {rank}:\n" f"#{rank} sim={nb['sim']:.2f} {self._fmt_mol(nb['smiles'])} "
f"Molecule: {self._fmt_mol(nb['smiles'])}\n" f"deliv={self._fmt(nb['delivery'])}{self._qbin('delivery', nb['delivery'])} "
f"Similarity score: {nb['sim']:.3f}\n" f"size={self._fmt(ex.get('size'))}{self._qbin('size', ex.get('size'))} "
f"delivery_log: {self._fmt(nb['delivery'])}\n" f"pdi={self._fmt_class(ex.get('pdi'), _PDI_LABELS)} "
f"size_z: {self._fmt(ex.get('size'))}\n" f"ee={self._fmt_class(ex.get('ee'), _EE_LABELS)} "
f"pdi_class: {self._fmt(ex.get('pdi'), 'int')}\n" f"tox={self._fmt_class(ex.get('toxic'), _TOX_LABELS)} bio={bio_s}"
f"ee_class: {self._fmt(ex.get('ee'), 'int')}\n"
f"toxic: {self._fmt(ex.get('toxic'), 'int')}\n"
f"biodist: {self._fmt(ex.get('biodist'), 'vec')}"
) )
retrieved_block = "\n\n".join(blocks) if blocks else "(no retrieved samples)" retrieved_block = "\n".join(blocks) if blocks else "(none)"
return ( return (
"Task: Encode the target LNP molecule into a retrieval-aware representation " "Encode the target LNP molecule into a retrieval-aware representation "
"for downstream multi-task property prediction. Do not output predictions.\n\n" "for multi-task property prediction. Do not output predictions.\n"
f"[Target Molecule]\nMolecule: {self._fmt_mol(target_smiles)}\n\n" f"[Target] {self._fmt_mol(target_smiles)}\n"
"[Retrieved Similar LNP Samples]\n" "[Retrieved] Nearest training molecules by fingerprint similarity, with "
"Retrieved from the training set by fingerprint similarity, with their known " "known outcomes. deliv=delivery_log and size=size_z are raw values, each "
"multi-task outcomes (delivery_log, size_z, pdi_class, ee_class, toxic, biodist; " f"followed by (Qi/{_N_QBINS}) = which {_N_QBINS}-quantile bin it falls into "
"'unknown' means the measurement is missing):\n\n" "among training molecules, Q1=lowest. pdi/ee/tox give the class index with "
f"{retrieved_block}\n\n" "its meaning. bio=fraction in [lymph_nodes,heart,liver,spleen,lung,kidney,"
"[Encoding Instructions]\n" "muscle] with the dominant organ named. 'unknown' = missing:\n"
"Capture the target structure and the retrieval evidence (structural similarity " f"{retrieved_block}\n"
"and consistency of retrieved outcomes) into your internal representation." "[Instruction] Capture the target structure and the retrieval evidence "
"(similarity, and both the level and the consistency of retrieved outcomes) "
"into your representation."
) )
def _warn_if_truncated(self, enc) -> None:
if getattr(self, "_trunc_warned", False):
return
if int(enc["attention_mask"].sum(1).max()) >= self.max_length:
from loguru import logger
logger.warning(
f"RAG prompt 触达 max_length={self.max_length} 被截断:"
f"尾部邻居与 [Instruction] 丢失,读出位置偏移。"
f"请调大 max_length 或减小 rag_top_k当前 {self.rag_top_k})。"
)
self._trunc_warned = True
def _get_prompt(self, s: str) -> str: def _get_prompt(self, s: str) -> str:
if not self.use_rag: if not self.use_rag:
return self._fmt_mol(s) return self._fmt_mol(s)
@ -273,6 +332,7 @@ class LLMPromptEncoder(nn.Module):
bt, padding=True, truncation=True, bt, padding=True, truncation=True,
max_length=self.max_length, return_tensors="pt", max_length=self.max_length, return_tensors="pt",
).to(device) ).to(device)
self._warn_if_truncated(enc)
out = self.encoder(**enc).last_hidden_state # [B,L,H] out = self.encoder(**enc).last_hidden_state # [B,L,H]
lengths = enc["attention_mask"].sum(1) - 1 lengths = enc["attention_mask"].sum(1) - 1
b = torch.arange(out.size(0), device=device) b = torch.arange(out.size(0), device=device)
@ -318,18 +378,33 @@ class LLMPromptEncoder(nn.Module):
if soft_list: if soft_list:
soft = torch.cat(soft_list, dim=1).to(text_embeds.dtype) soft = torch.cat(soft_list, dim=1).to(text_embeds.dtype)
inputs_embeds = torch.cat([soft, text_embeds], dim=1) # soft 后置:因果注意力下文本表示不受右侧影响,文本段因此只依赖 SMILES
# 对同一分子的所有配方保持不变(可缓存);读出点落在最后一个 soft token
# 它同时看到全部文本与全部 soft整合能力不损失。
inputs_embeds = torch.cat([text_embeds, soft], dim=1)
soft_mask = torch.ones(soft.size(0), soft.size(1), soft_mask = torch.ones(soft.size(0), soft.size(1),
device=device, dtype=text_mask.dtype) device=device, dtype=text_mask.dtype)
attn_mask = torch.cat([soft_mask, text_mask], dim=1) attn_mask = torch.cat([text_mask, soft_mask], dim=1)
else: else:
inputs_embeds = text_embeds inputs_embeds = text_embeds
attn_mask = text_mask attn_mask = text_mask
out = self.encoder(inputs_embeds=inputs_embeds, attention_mask=attn_mask).last_hidden_state # 左填充下位置编码必须由 mask 推出,否则 PAD 会把真实 token 的位置顶偏
lengths = attn_mask.sum(1) - 1 position_ids = attn_mask.long().cumsum(-1) - 1
b = torch.arange(out.size(0), device=device) position_ids = position_ids.masked_fill(attn_mask == 0, 1)
feat = out[b, lengths.long(), :] # [B, H] 最后有效 token
out = self.encoder(
inputs_embeds=inputs_embeds,
attention_mask=attn_mask,
position_ids=position_ids,
).last_hidden_state
if soft_list:
feat = out[:, -1, :] # 后置时末位必为最后一个 soft token
else:
lengths = attn_mask.sum(1) - 1
b = torch.arange(out.size(0), device=device)
feat = out[b, lengths.long(), :]
return self.proj_down(feat.float()) return self.proj_down(feat.float())
# ---------- 旧路径(保留,向后兼容)---------- # ---------- 旧路径(保留,向后兼容)----------

View File

@ -20,7 +20,7 @@ from lnp_ml.modeling.layers import (
LLMPromptEncoder, LLMPromptEncoder,
) )
from lnp_ml.modeling.layers.llm_prompt import DEFAULT_MOLT5_PATH from lnp_ml.modeling.layers.llm_prompt import DEFAULT_MOLT5_PATH
from lnp_ml.modeling.heads import MultiTaskHead from lnp_ml.modeling.heads import MultiTaskHead, RegressionHead
PoolingStrategy = Literal["attention", "avg", "max"] PoolingStrategy = Literal["attention", "avg", "max"]
@ -128,6 +128,7 @@ class LNPModel(nn.Module):
llm_lora_dropout: float = 0.05, llm_lora_dropout: float = 0.05,
use_rag: bool = False, use_rag: bool = False,
rag_top_k: int = 4, rag_top_k: int = 4,
llm_max_length: int = 1536,
use_retrieval: bool = False, use_retrieval: bool = False,
retr_feature_dim: int = 0, retr_feature_dim: int = 0,
) -> None: ) -> None:
@ -246,6 +247,7 @@ class LNPModel(nn.Module):
lora_dropout=llm_lora_dropout, lora_dropout=llm_lora_dropout,
use_rag=use_rag, use_rag=use_rag,
rag_top_k=rag_top_k, rag_top_k=rag_top_k,
max_length=llm_max_length,
use_soft_prompt=use_soft_prompt, use_soft_prompt=use_soft_prompt,
) )
else: else:
@ -269,6 +271,14 @@ class LNPModel(nn.Module):
dropout=dropout, dropout=dropout,
) )
# ============ 分支辅助头(仅训练期参与 loss不进 forward============
# f_moe / f_llm 要经过零初始化门 g 才进入主路径g=0 时这两条分支梯度为零。
# 辅助头提供一条绕开 g 的通路:先把分支练成对 delivery 有判别力的方向,
# ⟨u, f⟩ 从随机噪声变成明确的正数后g 才拿得到强梯度。
# 它不参与 forward(),所以 g=0 时主路径输出仍逐比特等于 baseline。
self.aux_head_moe = RegressionHead(d_model, head_hidden_dim, dropout) if use_moe else None
self.aux_head_llm = RegressionHead(d_model, head_hidden_dim, dropout) if use_llm else None
def _encode_and_project( def _encode_and_project(
self, self,
smiles: List[str], smiles: List[str],
@ -324,7 +334,8 @@ class LNPModel(nn.Module):
self._last_moe_extras = None self._last_moe_extras = None
f_llm = None f_llm = None
if self.llm_prompt is not None and smiles is not None: if (self.llm_prompt is not None and smiles is not None
and getattr(self, "_llm_enabled", True)):
if getattr(self.llm_prompt, "use_soft_prompt", False): if getattr(self.llm_prompt, "use_soft_prompt", False):
f_llm = self.llm_prompt(smiles, chem=chem, tab=tab) f_llm = self.llm_prompt(smiles, chem=chem, tab=tab)
else: else:
@ -336,10 +347,27 @@ class LNPModel(nn.Module):
_feats_t = torch.as_tensor(_feats, dtype=chem.dtype, device=chem.device) _feats_t = torch.as_tensor(_feats, dtype=chem.dtype, device=chem.device)
f_retr = self.retr_proj(_feats_t) f_retr = self.retr_proj(_feats_t)
# 供辅助监督使用(训练期),推理期不消费
self._last_f_moe = f_moe
self._last_f_llm = f_llm
fused, pooled = self.fusion(chem, tab, f_moe=f_moe, f_llm=f_llm, f_retr=f_retr) fused, pooled = self.fusion(chem, tab, f_moe=f_moe, f_llm=f_llm, f_retr=f_retr)
self._last_pooled = pooled # 纯数值向量,供回归 head 使用 self._last_pooled = pooled # 纯数值向量,供回归 head 使用
return fused return fused
def get_aux_outputs(self) -> Dict[str, torch.Tensor]:
"""返回各旁路分支对 delivery 的辅助预测。仅训练期由 loss 消费。"""
out: Dict[str, torch.Tensor] = {}
if not self.training:
return out
f_moe = getattr(self, "_last_f_moe", None)
f_llm = getattr(self, "_last_f_llm", None)
if getattr(self, "aux_head_moe", None) is not None and f_moe is not None:
out["aux_moe"] = self.aux_head_moe(f_moe)
if getattr(self, "aux_head_llm", None) is not None and f_llm is not None:
out["aux_llm"] = self.aux_head_llm(f_llm)
return out
def forward_from_projected( def forward_from_projected(
self, self,
stacked: torch.Tensor, stacked: torch.Tensor,
@ -492,13 +520,22 @@ class LNPModel(nn.Module):
} }
unexpected = [] unexpected = []
loaded = []
model_state = self.state_dict() model_state = self.state_dict()
for k, v in filtered_state_dict.items(): for k, v in filtered_state_dict.items():
if k in model_state and model_state[k].shape == v.shape: if k in model_state and model_state[k].shape == v.shape:
model_state[k] = v model_state[k] = v
loaded.append(k)
else: else:
unexpected.append(k) unexpected.append(k)
from loguru import logger
logger.info(
f"预训练权重:迁移 {len(loaded)} 个张量"
f"{sum(model_state[k].numel() for k in loaded) / 1e6:.2f}M 参数),"
f"跳过 {len(unexpected)}"
+ (f",例如 {unexpected[:3]}" if unexpected else "")
)
self.load_state_dict(model_state, strict=False) self.load_state_dict(model_state, strict=False)
@ -543,6 +580,7 @@ class LNPModelWithoutMPNN(LNPModel):
llm_lora_dropout: float = 0.05, llm_lora_dropout: float = 0.05,
use_rag: bool = False, use_rag: bool = False,
rag_top_k: int = 4, rag_top_k: int = 4,
llm_max_length: int = 1536,
use_retrieval: bool = False, use_retrieval: bool = False,
retr_feature_dim: int = 0, retr_feature_dim: int = 0,
) -> None: ) -> None:
@ -573,6 +611,7 @@ class LNPModelWithoutMPNN(LNPModel):
use_llm=use_llm, use_llm=use_llm,
use_rag=use_rag, use_rag=use_rag,
rag_top_k=rag_top_k, rag_top_k=rag_top_k,
llm_max_length=llm_max_length,
use_retrieval=use_retrieval, use_retrieval=use_retrieval,
retr_feature_dim=retr_feature_dim, retr_feature_dim=retr_feature_dim,
llm_model_path=llm_model_path, llm_model_path=llm_model_path,

View File

@ -804,6 +804,19 @@ def _run_single_outer_fold(
class_weights = compute_class_weights_from_loader(train_loader) class_weights = compute_class_weights_from_loader(train_loader)
arch = {
"moe_n_experts": best_params.get("moe_n_experts", moe_n_experts),
"moe_top_k": best_params.get("moe_top_k", moe_top_k),
"moe_expert_hidden_mult": best_params.get("moe_expert_hidden_mult", moe_expert_hidden_mult),
"moe_jitter_noise": moe_jitter_noise,
"set_transformer_block": best_params.get("set_transformer_block", "sab"),
}
llm_kwargs_resolved = {
**(llm_kwargs or {}),
**({"llm_use_lora": True, "llm_lora_r": best_params["llm_lora_r"]}
if ((llm_kwargs or {}).get("use_llm", False) and "llm_lora_r" in best_params) else {}),
}
model = create_model( model = create_model(
d_model=best_params["d_model"], d_model=best_params["d_model"],
num_heads=best_params["num_heads"], num_heads=best_params["num_heads"],
@ -818,15 +831,10 @@ def _run_single_outer_fold(
moleculestm_cache=moleculestm_cache, moleculestm_cache=moleculestm_cache,
mole_cache=mole_cache, mole_cache=mole_cache,
use_moe=use_moe, use_moe=use_moe,
moe_n_experts=best_params.get("moe_n_experts", moe_n_experts), llm_kwargs=llm_kwargs_resolved,
moe_top_k=best_params.get("moe_top_k", moe_top_k),
moe_expert_hidden_mult=best_params.get("moe_expert_hidden_mult", moe_expert_hidden_mult),
moe_jitter_noise=moe_jitter_noise,
set_transformer_block=best_params.get("set_transformer_block", "sab"),
llm_kwargs={**llm_kwargs, **({"llm_use_lora": True, "llm_lora_r": best_params["llm_lora_r"]}
if (llm_kwargs.get("use_llm", False) and "llm_lora_r" in best_params) else {})},
use_retrieval=use_retrieval, use_retrieval=use_retrieval,
retr_feature_dim=(3 if use_retrieval else 0), retr_feature_dim=(3 if use_retrieval else 0),
**arch,
) )
model.rdkit_encoder._cache = rdkit_cache model.rdkit_encoder._cache = rdkit_cache
if use_retrieval: if use_retrieval:
@ -878,7 +886,6 @@ def _run_single_outer_fold(
"n_attn_layers": best_params["n_attn_layers"], "n_attn_layers": best_params["n_attn_layers"],
"fusion_strategy": best_params["fusion_strategy"], "fusion_strategy": best_params["fusion_strategy"],
"head_hidden_dim": best_params["head_hidden_dim"], "head_hidden_dim": best_params["head_hidden_dim"],
"set_transformer_block": best_params.get("set_transformer_block", "sab"),
"dropout": best_params["dropout"], "dropout": best_params["dropout"],
"use_mpnn": use_mpnn, "use_mpnn": use_mpnn,
"use_chemeleon": chemeleon_cache is not None, "use_chemeleon": chemeleon_cache is not None,
@ -890,16 +897,19 @@ def _run_single_outer_fold(
"use_mole": mole_cache is not None, "use_mole": mole_cache is not None,
"mole_cache": mole_cache, "mole_cache": mole_cache,
"use_moe": use_moe, "use_moe": use_moe,
"moe_n_experts": moe_n_experts, "use_retrieval": use_retrieval,
"moe_top_k": moe_top_k, "retr_feature_dim": (3 if use_retrieval else 0),
"moe_expert_hidden_mult": moe_expert_hidden_mult, **arch,
"moe_jitter_noise": moe_jitter_noise, **llm_kwargs_resolved,
**(llm_kwargs or {}),
} }
_full_state = train_result["final_state"] _full_state = train_result["final_state"]
_slim_state = {k: v for k, v in _full_state.items() _slim_state = {
if not k.startswith("llm_prompt.encoder")} k: v for k, v in _full_state.items()
if (not k.startswith("llm_prompt.encoder")) or ("lora_" in k)
}
_n_lora = sum(1 for k in _slim_state if "lora_" in k)
logger.info(f"保存 checkpoint: {len(_slim_state)} 个张量,其中 LoRA 适配器 {_n_lora}")
torch.save({ torch.save({
"model_state_dict": _slim_state, "model_state_dict": _slim_state,
"config": config, "config": config,
@ -981,6 +991,7 @@ def main(
use_llm: bool = False, use_llm: bool = False,
use_rag: bool = False, use_rag: bool = False,
rag_top_k: int = 4, rag_top_k: int = 4,
llm_max_length: int = 1536,
use_retrieval: bool = False, use_retrieval: bool = False,
retrieval_source: str = "internal", retrieval_source: str = "internal",
llm_model_path: str = DEFAULT_MOLT5_PATH, llm_model_path: str = DEFAULT_MOLT5_PATH,
@ -991,6 +1002,8 @@ def main(
llm_lora_r: int = 8, llm_lora_r: int = 8,
llm_lora_alpha: int = 16, llm_lora_alpha: int = 16,
llm_lora_dropout: float = 0.05, llm_lora_dropout: float = 0.05,
fix_hparams_json: Optional[Path] = None,
fix_epoch_mean: int = 12,
# 并行 # 并行
parallel: bool = False, parallel: bool = False,
# 设备 # 设备
@ -1017,6 +1030,7 @@ def main(
reg_bypass=reg_bypass, reg_bypass=reg_bypass,
use_rag=use_rag, use_rag=use_rag,
rag_top_k=rag_top_k, rag_top_k=rag_top_k,
llm_max_length=llm_max_length,
use_llm=use_llm, use_llm=use_llm,
llm_model_path=llm_model_path, llm_model_path=llm_model_path,
llm_freeze=llm_freeze, llm_freeze=llm_freeze,
@ -1028,7 +1042,17 @@ def main(
llm_lora_dropout=llm_lora_dropout, llm_lora_dropout=llm_lora_dropout,
) )
# 加载预训练权重(如果指定) _fixed_bp = None
_fixed_em = None
if fix_hparams_json is not None:
_fixed_bp = json.loads(Path(fix_hparams_json).read_text())
_fixed_em = int(fix_epoch_mean)
if use_moe:
_fixed_bp["moe_n_experts"] = moe_n_experts
_fixed_bp["moe_top_k"] = moe_top_k
logger.info(f"[FIX] 固定超参={_fixed_bp}, epoch_mean={_fixed_em}, 跳过内层 Optuna")
# 加载预训练权重
pretrain_state_dict = None pretrain_state_dict = None
pretrain_config = None pretrain_config = None
if init_from_pretrain is not None: if init_from_pretrain is not None:
@ -1113,6 +1137,8 @@ def main(
moe_jitter_noise=moe_jitter_noise, moe_jitter_noise=moe_jitter_noise,
set_transformer_block=set_transformer_block, set_transformer_block=set_transformer_block,
llm_kwargs=llm_kwargs, llm_kwargs=llm_kwargs,
precomputed_best_params=_fixed_bp,
precomputed_epoch_mean=_fixed_em,
)) ))
if parallel: if parallel:

View File

@ -36,57 +36,78 @@ def load_model(
""" """
加载训练好的模型 加载训练好的模型
自动根据 checkpoint config.use_mpnn 选择模型类型 根据 checkpoint config 决定模型类型与全部结构开关
(MoE / LLM / RAG / soft-prompt / reg_bypass)
""" """
checkpoint = torch.load(model_path, map_location=device, weights_only=False) checkpoint = torch.load(model_path, map_location=device, weights_only=False)
config = checkpoint["config"] config = checkpoint["config"]
use_mpnn = config.get("use_mpnn", False) use_mpnn = config.get("use_mpnn", False)
common_kwargs = dict(
d_model=config["d_model"],
num_heads=config["num_heads"],
n_attn_layers=config["n_attn_layers"],
set_transformer_block=config.get("set_transformer_block", "sab"),
fusion_strategy=config["fusion_strategy"],
head_hidden_dim=config["head_hidden_dim"],
dropout=config["dropout"],
reg_bypass=config.get("reg_bypass", "on"),
use_moe=config.get("use_moe", False),
moe_n_experts=config.get("moe_n_experts", 4),
moe_top_k=config.get("moe_top_k", 2),
moe_expert_hidden_mult=config.get("moe_expert_hidden_mult", 2),
moe_jitter_noise=config.get("moe_jitter_noise", 0.0),
use_llm=config.get("use_llm", False),
llm_model_path=config.get("llm_model_path", "models/molt5-base"),
llm_freeze=config.get("llm_freeze", True),
llm_use_lora=config.get("llm_use_lora", False),
llm_use_qlora=config.get("llm_use_qlora", False),
use_soft_prompt=config.get("use_soft_prompt", False),
llm_lora_r=config.get("llm_lora_r", 8),
llm_lora_alpha=config.get("llm_lora_alpha", 16),
llm_lora_dropout=config.get("llm_lora_dropout", 0.05),
use_rag=config.get("use_rag", False),
rag_top_k=config.get("rag_top_k", 4),
llm_max_length=config.get("llm_max_length", 256),
use_retrieval=config.get("use_retrieval", False),
retr_feature_dim=config.get("retr_feature_dim", 0),
chemeleon_cache_path=config.get("chemeleon_cache"),
unimol_cache_path=config.get("unimol_cache"),
)
if use_mpnn: if use_mpnn:
# 总是自动查找 MPNN ensemble避免使用 checkpoint 中的旧绝对路径(可能来自其他机器) # 总是自动查找 MPNN ensemble避免使用 checkpoint 中的旧绝对路径(可能来自其他机器)
logger.info("Model was trained with MPNN, auto-detecting ensemble...") logger.info("Model was trained with MPNN, auto-detecting ensemble...")
ensemble_paths = find_mpnn_ensemble_paths() ensemble_paths = find_mpnn_ensemble_paths()
logger.info(f"Found {len(ensemble_paths)} MPNN models") logger.info(f"Found {len(ensemble_paths)} MPNN models")
model = LNPModel( model = LNPModel(
d_model=config["d_model"],
num_heads=config["num_heads"],
n_attn_layers=config["n_attn_layers"],
fusion_strategy=config["fusion_strategy"],
head_hidden_dim=config["head_hidden_dim"],
dropout=config["dropout"],
use_llm=config.get("use_llm", False),
llm_model_path=config.get("llm_model_path", "models/molt5-base"),
llm_freeze=config.get("llm_freeze", True),
llm_use_lora=config.get("llm_use_lora", False),
llm_lora_r=config.get("llm_lora_r", 8),
llm_lora_alpha=config.get("llm_lora_alpha", 16),
llm_lora_dropout=config.get("llm_lora_dropout", 0.05),
mpnn_ensemble_paths=ensemble_paths, mpnn_ensemble_paths=ensemble_paths,
mpnn_device=mpnn_device, mpnn_device=mpnn_device,
chemeleon_cache_path=config.get("chemeleon_cache"), **common_kwargs,
unimol_cache_path=config.get("unimol_cache"),
) )
else: else:
model = LNPModelWithoutMPNN( model = LNPModelWithoutMPNN(**common_kwargs)
d_model=config["d_model"],
num_heads=config["num_heads"], # 兼容旧 checkpointfusion 门从标量升级为逐维向量后,需把 0 维张量展开
n_attn_layers=config["n_attn_layers"], _sd = checkpoint["model_state_dict"]
fusion_strategy=config["fusion_strategy"], for _k in ("fusion.g_moe", "fusion.g_llm", "fusion.g_retr"):
head_hidden_dim=config["head_hidden_dim"], if _k in _sd and _sd[_k].dim() == 0:
dropout=config["dropout"], _sd[_k] = _sd[_k].reshape(1).expand(model.fusion.d_model).clone()
use_llm=config.get("use_llm", False),
llm_model_path=config.get("llm_model_path", "models/molt5-base"), missing, unexpected = model.load_state_dict(
llm_freeze=config.get("llm_freeze", True), checkpoint["model_state_dict"], strict=False
llm_use_lora=config.get("llm_use_lora", False), )
llm_lora_r=config.get("llm_lora_r", 8), # strict=False 不会因结构不匹配报错,这里手动兜底
llm_lora_alpha=config.get("llm_lora_alpha", 16), if unexpected:
llm_lora_dropout=config.get("llm_lora_dropout", 0.05), raise RuntimeError(
chemeleon_cache_path=config.get("chemeleon_cache"), f"checkpoint 中有 {len(unexpected)} 个权重找不到对应模块,"
unimol_cache_path=config.get("unimol_cache"), f"结构开关可能未对齐: {unexpected[:5]}"
)
if missing:
logger.warning(
f"{len(missing)} 个参数未从 checkpoint 恢复,将使用随机初始化: {missing[:5]}"
) )
model.load_state_dict(checkpoint["model_state_dict"], strict=False)
model.to(device) model.to(device)
model.eval() model.eval()
@ -152,7 +173,7 @@ def predictions_to_dataframe(predictions: Dict) -> pd.DataFrame:
}) })
# PDI 类别映射 # PDI 类别映射
pdi_labels = ["0_0to0_2", "0_2to0_3", "0_3to0_4", "0_4to0_5"] pdi_labels = ["PDI<0.2", "PDI>=0.2"]
df["pred_pdi_label"] = df["pred_pdi_class"].map(lambda x: pdi_labels[x]) df["pred_pdi_label"] = df["pred_pdi_class"].map(lambda x: pdi_labels[x])
# EE 类别映射 # EE 类别映射
@ -311,10 +332,10 @@ def test(
"r2": float(r2_score(y_true, y_pred)), "r2": float(r2_score(y_true, y_pred)),
} }
# 分类指标PDI # 分类指标PDI(真值需与 dataset.py 同样折叠为二分类)
pdi_cols = ["PDI_0_0to0_2", "PDI_0_2to0_3", "PDI_0_3to0_4", "PDI_0_4to0_5"] pdi_cols = ["PDI_0_0to0_2", "PDI_0_2to0_3", "PDI_0_3to0_4", "PDI_0_4to0_5"]
if all(c in test_df.columns for c in pdi_cols): if all(c in test_df.columns for c in pdi_cols):
pdi_true = test_df[pdi_cols].values.argmax(axis=1) pdi_true = (test_df[pdi_cols].values.argmax(axis=1) >= 1).astype(np.int64)
mask = test_df[pdi_cols].sum(axis=1) > 0 mask = test_df[pdi_cols].sum(axis=1) > 0
if mask.any(): if mask.any():
y_true = pdi_true[mask] y_true = pdi_true[mask]

View File

@ -1,7 +1,7 @@
"""带类权重的训练器:处理分类任务的数据不均衡问题""" """带类权重的训练器:处理分类任务的数据不均衡问题"""
from typing import Dict, List, Optional, Tuple from typing import Dict, List, Optional, Tuple
from dataclasses import dataclass, field from dataclasses import dataclass, field, replace
import numpy as np import numpy as np
import torch import torch
@ -14,7 +14,7 @@ from tqdm import tqdm
@dataclass @dataclass
class ClassWeights: class ClassWeights:
"""分类任务的类权重""" """分类任务的类权重"""
pdi: Optional[torch.Tensor] = None # [4] for 4 PDI classes pdi: Optional[torch.Tensor] = None # [2] 二分类 (0: PDI<0.2, 1: PDI>=0.2)
ee: Optional[torch.Tensor] = None # [3] for 3 EE classes ee: Optional[torch.Tensor] = None # [3] for 3 EE classes
toxic: Optional[torch.Tensor] = None # [2] for binary toxic toxic: Optional[torch.Tensor] = None # [2] for binary toxic
@ -29,6 +29,10 @@ class LossWeightsBalanced:
biodist: float = 1.0 biodist: float = 1.0
toxic: float = 1.0 toxic: float = 1.0
moe_lb: float = 0.01 # MoE load-balancing 系数(仅在 use_moe=True 时生效) moe_lb: float = 0.01 # MoE load-balancing 系数(仅在 use_moe=True 时生效)
# 旁路分支辅助监督系数。GoogLeNet 辅助分类器用的也是 0.3。
# 训练后期应退火到 0辅助 loss 优化的是"f 单独能预测 delivery"
# 而主 loss 要的是"f 对 pooled 有增量价值",两者并不等价。
aux_branch: float = 0.3
def compute_class_weights_from_loader( def compute_class_weights_from_loader(
@ -188,6 +192,16 @@ def compute_multitask_loss_balanced(
losses["moe_lb"] = extras["lb_loss"] losses["moe_lb"] = extras["lb_loss"]
total_loss = total_loss + task_weights.moe_lb * losses["moe_lb"] total_loss = total_loss + task_weights.moe_lb * losses["moe_lb"]
# 旁路分支辅助监督:绕开零初始化门,直接用 f_moe / f_llm 预测 delivery。
# 不纳入不确定性加权——它是优化脚手架,不是一个真实任务。
if (model is not None and hasattr(model, "get_aux_outputs")
and task_weights.aux_branch > 0
and "delivery" in targets and mask["delivery"].any()):
m = mask["delivery"]
for name, pred in model.get_aux_outputs().items():
losses[name] = F.mse_loss(pred[m].squeeze(-1), targets["delivery"][m])
total_loss = total_loss + task_weights.aux_branch * losses[name]
return total_loss, losses return total_loss, losses
@ -202,7 +216,8 @@ def train_epoch_balanced(
"""带类权重的训练一个 epoch""" """带类权重的训练一个 epoch"""
model.train() model.train()
total_loss = 0.0 total_loss = 0.0
task_losses = {k: 0.0 for k in ["size", "pdi", "ee", "delivery", "biodist", "toxic", "moe_lb"]} task_losses = {k: 0.0 for k in ["size", "pdi", "ee", "delivery", "biodist", "toxic",
"moe_lb", "aux_moe", "aux_llm"]}
n_batches = 0 n_batches = 0
for batch in tqdm(loader, desc="Training", leave=False): for batch in tqdm(loader, desc="Training", leave=False):
@ -394,10 +409,14 @@ def train_with_early_stopping(
task_weights: Optional[LossWeightsBalanced] = None, task_weights: Optional[LossWeightsBalanced] = None,
class_weights: Optional[ClassWeights] = None, class_weights: Optional[ClassWeights] = None,
backbone_lr_ratio: float = 1.0, backbone_lr_ratio: float = 1.0,
freeze_backbone_epochs: int = 0,
) -> Dict: ) -> Dict:
""" """
带早停的完整训练流程 带早停的完整训练流程
freeze_backbone_epochs 必须与 train_fixed_epochs 取同一个值best_epoch 是在
这个调度下测出来的最终训练换了调度这个数就不可迁移
Returns: Returns:
Dict with keys: history, best_val_loss, best_epoch, best_state Dict with keys: history, best_val_loss, best_epoch, best_state
""" """
@ -408,11 +427,27 @@ def train_with_early_stopping(
) )
early_stopping = EarlyStoppingBalanced(patience=patience, min_delta=1e-3) early_stopping = EarlyStoppingBalanced(patience=patience, min_delta=1e-3)
backbone_named = [
(n, p) for n, p in model.named_parameters()
if n.startswith(BACKBONE_PREFIXES) and not n.startswith(FROM_SCRATCH_PREFIXES)
and p.requires_grad
]
if freeze_backbone_epochs > 0:
for _, p in backbone_named:
p.requires_grad_(False)
task_weights = replace(task_weights) if task_weights is not None else LossWeightsBalanced()
_aux_w0 = task_weights.aux_branch
history = {"train": [], "val": []} history = {"train": [], "val": []}
best_val_loss = float("inf") best_val_loss = float("inf")
best_state = None best_state = None
for epoch in range(epochs): for epoch in range(epochs):
task_weights.aux_branch = _aux_w0 * max(0.0, 1.0 - epoch / (0.5 * epochs))
if freeze_backbone_epochs > 0 and epoch == freeze_backbone_epochs:
for _, p in backbone_named:
p.requires_grad_(True)
# Train # Train
train_metrics = train_epoch_balanced( train_metrics = train_epoch_balanced(
model, train_loader, optimizer, device, task_weights, class_weights model, train_loader, optimizer, device, task_weights, class_weights
@ -516,9 +551,13 @@ def train_fixed_epochs(
swa_start = swa_start_epoch or int(epochs * 0.75) swa_start = swa_start_epoch or int(epochs * 0.75)
swa_scheduler = SWALR(optimizer, swa_lr=lr * 0.1) swa_scheduler = SWALR(optimizer, swa_lr=lr * 0.1)
task_weights = replace(task_weights) if task_weights is not None else LossWeightsBalanced()
_aux_w0 = task_weights.aux_branch
history = {"train": [], "val": []} history = {"train": [], "val": []}
for epoch in range(epochs): for epoch in range(epochs):
task_weights.aux_branch = _aux_w0 * max(0.0, 1.0 - epoch / (0.5 * epochs))
if freeze_backbone_epochs > 0 and epoch == freeze_backbone_epochs: if freeze_backbone_epochs > 0 and epoch == freeze_backbone_epochs:
for _, p in backbone_named: for _, p in backbone_named:
p.requires_grad_(True) p.requires_grad_(True)

View File

@ -1,11 +1,16 @@
{ {
"dropout": 0.14731648771233646, "dropout": 0.19323529834582587,
"lr": 0.00010399668603456206, "lr": 0.000980484886752915,
"weight_decay": 0.0007696017715340258, "weight_decay": 0.00014367509113739864,
"backbone_lr_ratio": 0.7472270478607838, "backbone_lr_ratio": 0.13093310334961422,
"moe_n_experts": 2,
"moe_top_k": 2,
"moe_expert_hidden_mult": 2,
"llm_lora_r": 8,
"d_model": 256, "d_model": 256,
"num_heads": 8, "num_heads": 8,
"n_attn_layers": 4, "n_attn_layers": 4,
"fusion_strategy": "attention", "fusion_strategy": "attention",
"head_hidden_dim": 128 "head_hidden_dim": 128,
"set_transformer_block": "sab"
} }

View File

@ -1,9 +1,7 @@
{ {
"pdi": [ "pdi": [
0.011065198108553886, 0.524036169052124,
0.03264812380075455, 1.475963830947876
0.8369068503379822,
3.119379997253418
], ],
"ee": [ "ee": [
1.600011944770813, 1.600011944770813,

View File

@ -1 +1 @@
{"epoch_mean": 23} {"epoch_mean": 11}

View File

@ -1,211 +1,126 @@
{ {
"train": [ "train": [
{ {
"loss": 17.60348994391305, "loss": 5.662022154286222,
"loss_size": 12.354392801012311, "loss_size": 0.9464401806581695,
"loss_pdi": 1.4643298983573914, "loss_pdi": 0.6576592483610477,
"loss_ee": 1.090087456362588, "loss_ee": 1.0441054870497506,
"loss_delivery": 0.8421656531946999, "loss_delivery": 1.0691580114499577,
"loss_biodist": 1.23681207214083, "loss_biodist": 0.7718977461446006,
"loss_toxic": 0.6157020671027047 "loss_toxic": 0.47160032813279135,
"loss_moe_lb": 1.9999999887538407,
"loss_aux_moe": 1.096376850709038,
"loss_aux_llm": 1.2355621509113401
}, },
{ {
"loss": 6.5880933318819315, "loss": 4.549718969273117,
"loss_size": 1.750367718083518, "loss_size": 0.9265210998227011,
"loss_pdi": 1.3177664024489266, "loss_pdi": 0.6320651516599475,
"loss_ee": 1.068630073751722, "loss_ee": 0.9564013177493833,
"loss_delivery": 0.8820391488926751, "loss_delivery": 0.8463795804682205,
"loss_biodist": 1.0014605011258806, "loss_biodist": 0.48631338606465535,
"loss_toxic": 0.567829430103302 "loss_toxic": 0.25031149114991696,
"loss_moe_lb": 1.9999999910030726,
"loss_aux_moe": 0.8730144033928946,
"loss_aux_llm": 1.0696029401612732
}, },
{ {
"loss": 4.678432379450117, "loss": 4.183341939494295,
"loss_size": 0.3203257258449282, "loss_size": 0.8995583306927726,
"loss_pdi": 1.180948099919728, "loss_pdi": 0.6350391439671786,
"loss_ee": 1.0613888587270464, "loss_ee": 0.9038635166186206,
"loss_delivery": 0.8234607481530735, "loss_delivery": 0.8323076782080362,
"loss_biodist": 0.7979316924299512, "loss_biodist": 0.44391101247297143,
"loss_toxic": 0.49437726182597025 "loss_toxic": 0.1671132598740031,
"loss_moe_lb": 1.9999999865046088,
"loss_aux_moe": 0.8456886384003567,
"loss_aux_llm": 1.0173823231796049
}, },
{ {
"loss": 4.174277833529881, "loss": 3.9456988685535936,
"loss_size": 0.4864327609539032, "loss_size": 0.925582969104344,
"loss_pdi": 1.014147652047021, "loss_pdi": 0.586557297211773,
"loss_ee": 1.0116309651306696, "loss_ee": 0.867810919599713,
"loss_delivery": 0.8047538192144462, "loss_delivery": 0.7656616399521535,
"loss_biodist": 0.5864214556557792, "loss_biodist": 0.4162127935099152,
"loss_toxic": 0.2708912321499416 "loss_toxic": 0.2137195658123226,
"loss_moe_lb": 1.9999999887538407,
"loss_aux_moe": 0.7506198146784643,
"loss_aux_llm": 0.9585366686981804
}, },
{ {
"loss": 3.4547910520008633, "loss": 3.4442428957741216,
"loss_size": 0.31596518414361136, "loss_size": 0.8259714532573268,
"loss_pdi": 0.9084383249282837, "loss_pdi": 0.5488103347004585,
"loss_ee": 0.9215172486645835, "loss_ee": 0.8086997242468708,
"loss_delivery": 0.7124462143651077, "loss_delivery": 0.7543281304808158,
"loss_biodist": 0.42033374096666065, "loss_biodist": 0.3298270061330975,
"loss_toxic": 0.17609038363609994 "loss_toxic": 0.09599270570566351,
"loss_moe_lb": 1.9999999955015362,
"loss_aux_moe": 0.8137425728282839,
"loss_aux_llm": 1.044384686873769
}, },
{ {
"loss": 3.2108109337942943, "loss": 3.108477974837681,
"loss_size": 0.2688417855118002, "loss_size": 0.7671496489981435,
"loss_pdi": 0.8587345012596675, "loss_pdi": 0.5366503028374798,
"loss_ee": 0.875971223626818, "loss_ee": 0.7869719379353073,
"loss_delivery": 0.7385377415588924, "loss_delivery": 0.685029683248052,
"loss_biodist": 0.3365256999220167, "loss_biodist": 0.2828179093183212,
"loss_toxic": 0.1321999922926937 "loss_toxic": 0.08221106674908749,
"loss_moe_lb": 1.9999999910030726,
"loss_aux_moe": 0.7168521201413758,
"loss_aux_llm": 0.9446638939234445
}, },
{ {
"loss": 2.9460756949016025, "loss": 2.983976681277437,
"loss_size": 0.2779192647763661, "loss_size": 0.7495474740159962,
"loss_pdi": 0.7961623072624207, "loss_pdi": 0.5227343659355955,
"loss_ee": 0.8504888585635594, "loss_ee": 0.7467741780685928,
"loss_delivery": 0.6400965334448431, "loss_delivery": 0.6811002301368511,
"loss_biodist": 0.2687057022537504, "loss_biodist": 0.24798926155803339,
"loss_toxic": 0.112703027070633 "loss_toxic": 0.11302385037876929,
"loss_moe_lb": 1.9999999955015362
}, },
{ {
"loss": 2.7530925273895264, "loss": 3.014890724757932,
"loss_size": 0.25332520263535635, "loss_size": 0.7314784642098084,
"loss_pdi": 0.7469035514763424, "loss_pdi": 0.48637263083233023,
"loss_ee": 0.8471435691629138, "loss_ee": 0.769600696158859,
"loss_delivery": 0.5530903546937874, "loss_delivery": 0.8273954027912246,
"loss_biodist": 0.2454697404588972, "loss_biodist": 0.21140338072799286,
"loss_toxic": 0.10716012333120618 "loss_toxic": 0.08676587497249338,
"loss_moe_lb": 1.9999999955015362
}, },
{ {
"loss": 2.742221474647522, "loss": 2.774192355713754,
"loss_size": 0.2413133829832077, "loss_size": 0.6206410539881239,
"loss_pdi": 0.6884610418762479, "loss_pdi": 0.4826473706173447,
"loss_ee": 0.8217804133892059, "loss_ee": 0.7343562110415045,
"loss_delivery": 0.6628774159720966, "loss_delivery": 0.6946425725758638,
"loss_biodist": 0.23072052534137452, "loss_biodist": 0.19993741290186937,
"loss_toxic": 0.09706869375492845 "loss_toxic": 0.1272664367724298,
"loss_moe_lb": 1.9999999932523043
}, },
{ {
"loss": 2.550677682672228, "loss": 2.683339330385316,
"loss_size": 0.2765064158343843, "loss_size": 0.5581827477885867,
"loss_pdi": 0.6737975222723824, "loss_pdi": 0.4963476042140205,
"loss_ee": 0.7778623870440892, "loss_ee": 0.7145765797709519,
"loss_delivery": 0.5299507762704577, "loss_delivery": 0.6466561276817097,
"loss_biodist": 0.1921106672712735, "loss_biodist": 0.1851363978436533,
"loss_toxic": 0.10044991970062256 "loss_toxic": 0.1616253986507698,
"loss_moe_lb": 1.9999999932523043
}, },
{ {
"loss": 2.481816521712712, "loss": 2.5718357652988075,
"loss_size": 0.23023176033582007, "loss_size": 0.6112437141391466,
"loss_pdi": 0.6636314690113068, "loss_pdi": 0.4928069238392812,
"loss_ee": 0.7728129327297211, "loss_ee": 0.7109674007262824,
"loss_delivery": 0.5704377527747836, "loss_delivery": 0.6155119884829476,
"loss_biodist": 0.1672141200729779, "loss_biodist": 0.18839857628885306,
"loss_toxic": 0.07748848732028689 "loss_toxic": 0.05047845465344166,
}, "loss_moe_lb": 1.9999999842553768
{
"loss": 2.360851458140782,
"loss_size": 0.21546032758695738,
"loss_pdi": 0.6131412557193211,
"loss_ee": 0.830133033650262,
"loss_delivery": 0.4567379740016934,
"loss_biodist": 0.17073182123047964,
"loss_toxic": 0.07464704542819943
},
{
"loss": 2.2426107866423473,
"loss_size": 0.19084072799887508,
"loss_pdi": 0.6170341287340436,
"loss_ee": 0.7295470791203635,
"loss_delivery": 0.4835586505276816,
"loss_biodist": 0.1602931246161461,
"loss_toxic": 0.061337063022490056
},
{
"loss": 2.2947227358818054,
"loss_size": 0.21879314631223679,
"loss_pdi": 0.6161722391843796,
"loss_ee": 0.762788040297372,
"loss_delivery": 0.45067436354500906,
"loss_biodist": 0.18453349598816463,
"loss_toxic": 0.061761434456067424
},
{
"loss": 2.3243510808263506,
"loss_size": 0.23437534804855073,
"loss_pdi": 0.594099121434348,
"loss_ee": 0.77819305232593,
"loss_delivery": 0.4997542394059045,
"loss_biodist": 0.1680830376488822,
"loss_toxic": 0.049846263602375984
},
{
"loss": 2.2657968997955322,
"loss_size": 0.2230603982295309,
"loss_pdi": 0.6490842210394996,
"loss_ee": 0.7079165577888489,
"loss_delivery": 0.4766663594969681,
"loss_biodist": 0.16171239422900335,
"loss_toxic": 0.0473570313770324
},
{
"loss": 2.2227368525096347,
"loss_size": 0.22425250709056854,
"loss_pdi": 0.6083124216113772,
"loss_ee": 0.7293048288140979,
"loss_delivery": 0.4598711388451712,
"loss_biodist": 0.1544684906091009,
"loss_toxic": 0.04652748556275453
},
{
"loss": 2.1243858337402344,
"loss_size": 0.21263282213892257,
"loss_pdi": 0.573635390826634,
"loss_ee": 0.6900361818926675,
"loss_delivery": 0.4512601649122579,
"loss_biodist": 0.15764176366584642,
"loss_toxic": 0.03917955420911312
},
{
"loss": 2.289846863065447,
"loss_size": 0.21625602032457078,
"loss_pdi": 0.596784600189754,
"loss_ee": 0.7933975628444127,
"loss_delivery": 0.4503451883792877,
"loss_biodist": 0.18038264129843032,
"loss_toxic": 0.05268087850085327
},
{
"loss": 2.1171253493853976,
"loss_size": 0.21727498407874787,
"loss_pdi": 0.5871099403926304,
"loss_ee": 0.6942037897450584,
"loss_delivery": 0.4114475510349231,
"loss_biodist": 0.16847609941448485,
"loss_toxic": 0.038613016584089825
},
{
"loss": 2.2243686744144986,
"loss_size": 0.23041224852204323,
"loss_pdi": 0.5950902572699955,
"loss_ee": 0.7303137523787362,
"loss_delivery": 0.4817630084497588,
"loss_biodist": 0.13934307438986643,
"loss_toxic": 0.047446319966443946
},
{
"loss": 2.142884841987065,
"loss_size": 0.18598268926143646,
"loss_pdi": 0.5736084473984582,
"loss_ee": 0.7191481994731086,
"loss_delivery": 0.4928105877978461,
"loss_biodist": 0.13822122237512044,
"loss_toxic": 0.03311372648126313
},
{
"loss": 2.0640590446335927,
"loss_size": 0.1981512185718332,
"loss_pdi": 0.5239957081420081,
"loss_ee": 0.6807411738804409,
"loss_delivery": 0.45748772472143173,
"loss_biodist": 0.14945154903190477,
"loss_toxic": 0.0542316823931677
} }
], ],
"val": [] "val": []

Binary file not shown.

View File

@ -1,3 +1,3 @@
version https://git-lfs.github.com/spec/v1 version https://git-lfs.github.com/spec/v1
oid sha256:241b22c5170dd7d5a82f409228589d61d16ca14c3ae1991a03bda8ca547ce743 oid sha256:b96d3867be53a3e23abdf291baeec4a9d610da58dc790220a345e7012187aaa0
size 28583340 size 53008646

Binary file not shown.

File diff suppressed because it is too large Load Diff

View File

@ -0,0 +1,11 @@
{
"dropout": 0.14731648771233646,
"lr": 0.00010399668603456206,
"weight_decay": 0.0007696017715340258,
"backbone_lr_ratio": 0.7472270478607838,
"d_model": 256,
"num_heads": 8,
"n_attn_layers": 4,
"fusion_strategy": "attention",
"head_hidden_dim": 128
}

View File

@ -0,0 +1,17 @@
{
"pdi": [
0.011065198108553886,
0.03264812380075455,
0.8369068503379822,
3.119379997253418
],
"ee": [
1.600011944770813,
0.9947696924209595,
0.40521833300590515
],
"toxic": [
0.09425132721662521,
1.9057486057281494
]
}

View File

@ -0,0 +1 @@
{"epoch_mean": 23}

View File

@ -0,0 +1,212 @@
{
"train": [
{
"loss": 17.60348994391305,
"loss_size": 12.354392801012311,
"loss_pdi": 1.4643298983573914,
"loss_ee": 1.090087456362588,
"loss_delivery": 0.8421656531946999,
"loss_biodist": 1.23681207214083,
"loss_toxic": 0.6157020671027047
},
{
"loss": 6.5880933318819315,
"loss_size": 1.750367718083518,
"loss_pdi": 1.3177664024489266,
"loss_ee": 1.068630073751722,
"loss_delivery": 0.8820391488926751,
"loss_biodist": 1.0014605011258806,
"loss_toxic": 0.567829430103302
},
{
"loss": 4.678432379450117,
"loss_size": 0.3203257258449282,
"loss_pdi": 1.180948099919728,
"loss_ee": 1.0613888587270464,
"loss_delivery": 0.8234607481530735,
"loss_biodist": 0.7979316924299512,
"loss_toxic": 0.49437726182597025
},
{
"loss": 4.174277833529881,
"loss_size": 0.4864327609539032,
"loss_pdi": 1.014147652047021,
"loss_ee": 1.0116309651306696,
"loss_delivery": 0.8047538192144462,
"loss_biodist": 0.5864214556557792,
"loss_toxic": 0.2708912321499416
},
{
"loss": 3.4547910520008633,
"loss_size": 0.31596518414361136,
"loss_pdi": 0.9084383249282837,
"loss_ee": 0.9215172486645835,
"loss_delivery": 0.7124462143651077,
"loss_biodist": 0.42033374096666065,
"loss_toxic": 0.17609038363609994
},
{
"loss": 3.2108109337942943,
"loss_size": 0.2688417855118002,
"loss_pdi": 0.8587345012596675,
"loss_ee": 0.875971223626818,
"loss_delivery": 0.7385377415588924,
"loss_biodist": 0.3365256999220167,
"loss_toxic": 0.1321999922926937
},
{
"loss": 2.9460756949016025,
"loss_size": 0.2779192647763661,
"loss_pdi": 0.7961623072624207,
"loss_ee": 0.8504888585635594,
"loss_delivery": 0.6400965334448431,
"loss_biodist": 0.2687057022537504,
"loss_toxic": 0.112703027070633
},
{
"loss": 2.7530925273895264,
"loss_size": 0.25332520263535635,
"loss_pdi": 0.7469035514763424,
"loss_ee": 0.8471435691629138,
"loss_delivery": 0.5530903546937874,
"loss_biodist": 0.2454697404588972,
"loss_toxic": 0.10716012333120618
},
{
"loss": 2.742221474647522,
"loss_size": 0.2413133829832077,
"loss_pdi": 0.6884610418762479,
"loss_ee": 0.8217804133892059,
"loss_delivery": 0.6628774159720966,
"loss_biodist": 0.23072052534137452,
"loss_toxic": 0.09706869375492845
},
{
"loss": 2.550677682672228,
"loss_size": 0.2765064158343843,
"loss_pdi": 0.6737975222723824,
"loss_ee": 0.7778623870440892,
"loss_delivery": 0.5299507762704577,
"loss_biodist": 0.1921106672712735,
"loss_toxic": 0.10044991970062256
},
{
"loss": 2.481816521712712,
"loss_size": 0.23023176033582007,
"loss_pdi": 0.6636314690113068,
"loss_ee": 0.7728129327297211,
"loss_delivery": 0.5704377527747836,
"loss_biodist": 0.1672141200729779,
"loss_toxic": 0.07748848732028689
},
{
"loss": 2.360851458140782,
"loss_size": 0.21546032758695738,
"loss_pdi": 0.6131412557193211,
"loss_ee": 0.830133033650262,
"loss_delivery": 0.4567379740016934,
"loss_biodist": 0.17073182123047964,
"loss_toxic": 0.07464704542819943
},
{
"loss": 2.2426107866423473,
"loss_size": 0.19084072799887508,
"loss_pdi": 0.6170341287340436,
"loss_ee": 0.7295470791203635,
"loss_delivery": 0.4835586505276816,
"loss_biodist": 0.1602931246161461,
"loss_toxic": 0.061337063022490056
},
{
"loss": 2.2947227358818054,
"loss_size": 0.21879314631223679,
"loss_pdi": 0.6161722391843796,
"loss_ee": 0.762788040297372,
"loss_delivery": 0.45067436354500906,
"loss_biodist": 0.18453349598816463,
"loss_toxic": 0.061761434456067424
},
{
"loss": 2.3243510808263506,
"loss_size": 0.23437534804855073,
"loss_pdi": 0.594099121434348,
"loss_ee": 0.77819305232593,
"loss_delivery": 0.4997542394059045,
"loss_biodist": 0.1680830376488822,
"loss_toxic": 0.049846263602375984
},
{
"loss": 2.2657968997955322,
"loss_size": 0.2230603982295309,
"loss_pdi": 0.6490842210394996,
"loss_ee": 0.7079165577888489,
"loss_delivery": 0.4766663594969681,
"loss_biodist": 0.16171239422900335,
"loss_toxic": 0.0473570313770324
},
{
"loss": 2.2227368525096347,
"loss_size": 0.22425250709056854,
"loss_pdi": 0.6083124216113772,
"loss_ee": 0.7293048288140979,
"loss_delivery": 0.4598711388451712,
"loss_biodist": 0.1544684906091009,
"loss_toxic": 0.04652748556275453
},
{
"loss": 2.1243858337402344,
"loss_size": 0.21263282213892257,
"loss_pdi": 0.573635390826634,
"loss_ee": 0.6900361818926675,
"loss_delivery": 0.4512601649122579,
"loss_biodist": 0.15764176366584642,
"loss_toxic": 0.03917955420911312
},
{
"loss": 2.289846863065447,
"loss_size": 0.21625602032457078,
"loss_pdi": 0.596784600189754,
"loss_ee": 0.7933975628444127,
"loss_delivery": 0.4503451883792877,
"loss_biodist": 0.18038264129843032,
"loss_toxic": 0.05268087850085327
},
{
"loss": 2.1171253493853976,
"loss_size": 0.21727498407874787,
"loss_pdi": 0.5871099403926304,
"loss_ee": 0.6942037897450584,
"loss_delivery": 0.4114475510349231,
"loss_biodist": 0.16847609941448485,
"loss_toxic": 0.038613016584089825
},
{
"loss": 2.2243686744144986,
"loss_size": 0.23041224852204323,
"loss_pdi": 0.5950902572699955,
"loss_ee": 0.7303137523787362,
"loss_delivery": 0.4817630084497588,
"loss_biodist": 0.13934307438986643,
"loss_toxic": 0.047446319966443946
},
{
"loss": 2.142884841987065,
"loss_size": 0.18598268926143646,
"loss_pdi": 0.5736084473984582,
"loss_ee": 0.7191481994731086,
"loss_delivery": 0.4928105877978461,
"loss_biodist": 0.13822122237512044,
"loss_toxic": 0.03311372648126313
},
{
"loss": 2.0640590446335927,
"loss_size": 0.1981512185718332,
"loss_pdi": 0.5239957081420081,
"loss_ee": 0.6807411738804409,
"loss_delivery": 0.45748772472143173,
"loss_biodist": 0.14945154903190477,
"loss_toxic": 0.0542316823931677
}
],
"val": []
}

Binary file not shown.

View File

@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:241b22c5170dd7d5a82f409228589d61d16ca14c3ae1991a03bda8ca547ce743
size 28583340

View File

@ -0,0 +1,962 @@
[
{
"number": 0,
"value": 2.8156970342000327,
"params": {
"dropout": 0.249816047538945,
"lr": 0.0007969454818643932,
"weight_decay": 0.008471801418819975,
"backbone_lr_ratio": 0.15751320499779725
},
"user_attrs": {
"epoch_mean": 14,
"fold_best_epochs": [
10,
12,
21
],
"fold_val_losses": [
3.733495855331421,
2.5024956703186034,
2.2110995769500734
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 1,
"value": 3.3739805857340492,
"params": {
"dropout": 0.1624074561769746,
"lr": 2.0511104188433963e-05,
"weight_decay": 1.7073967431528103e-05,
"backbone_lr_ratio": 0.5399484409787431
},
"user_attrs": {
"epoch_mean": 30,
"fold_best_epochs": [
30,
30,
30
],
"fold_val_losses": [
3.7005380153656007,
3.236070442199707,
3.1853332996368406
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 2,
"value": 2.6023060957590736,
"params": {
"dropout": 0.34044600469728353,
"lr": 0.0002607024758370766,
"weight_decay": 1.2087541473056957e-05,
"backbone_lr_ratio": 0.8706020878304853
},
"user_attrs": {
"epoch_mean": 16,
"fold_best_epochs": [
14,
19,
16
],
"fold_val_losses": [
3.283898639678955,
2.366805338859558,
2.1562143087387087
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 3,
"value": 10.838356399536133,
"params": {
"dropout": 0.4329770563201687,
"lr": 2.6587543983272695e-05,
"weight_decay": 5.3370327626039544e-05,
"backbone_lr_ratio": 0.023270677083837805
},
"user_attrs": {
"epoch_mean": 30,
"fold_best_epochs": [
30,
30,
30
],
"fold_val_losses": [
9.760257911682128,
10.611498069763183,
12.143313217163087
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 4,
"value": 3.0307216803232833,
"params": {
"dropout": 0.2216968971838151,
"lr": 0.00011207606211860574,
"weight_decay": 0.0005342937261279777,
"backbone_lr_ratio": 0.038234752246751866
},
"user_attrs": {
"epoch_mean": 30,
"fold_best_epochs": [
29,
30,
30
],
"fold_val_losses": [
3.6398903369903564,
2.7796840190887453,
2.6725906848907472
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 5,
"value": 12.962510363260904,
"params": {
"dropout": 0.34474115788895177,
"lr": 1.9010245319870364e-05,
"weight_decay": 0.00014742753159914678,
"backbone_lr_ratio": 0.05404103854647329
},
"user_attrs": {
"epoch_mean": 30,
"fold_best_epochs": [
30,
30,
30
],
"fold_val_losses": [
11.658656311035156,
13.798492240905762,
13.430382537841798
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 6,
"value": 2.8127863725026447,
"params": {
"dropout": 0.28242799368681437,
"lr": 0.00037183641805732076,
"weight_decay": 6.290644294586152e-05,
"backbone_lr_ratio": 0.10677482709481352
},
"user_attrs": {
"epoch_mean": 21,
"fold_best_epochs": [
15,
20,
28
],
"fold_val_losses": [
3.527592992782593,
2.6037728786468506,
2.306993246078491
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 7,
"value": 18.49700838724772,
"params": {
"dropout": 0.33696582754481696,
"lr": 1.2385137298860926e-05,
"weight_decay": 0.0026926469100861782,
"backbone_lr_ratio": 0.021930485556643693
},
"user_attrs": {
"epoch_mean": 30,
"fold_best_epochs": [
30,
30,
30
],
"fold_val_losses": [
18.425870513916017,
18.542006301879884,
18.523148345947266
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 8,
"value": 2.6576756556828816,
"params": {
"dropout": 0.12602063719411183,
"lr": 0.000790261954970823,
"weight_decay": 0.07286653737491042,
"backbone_lr_ratio": 0.4138040112561014
},
"user_attrs": {
"epoch_mean": 10,
"fold_best_epochs": [
9,
11,
10
],
"fold_val_losses": [
3.56644024848938,
2.2665807485580443,
2.1400059700012206
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 9,
"value": 11.92358964284261,
"params": {
"dropout": 0.2218455076693483,
"lr": 1.5679933916722995e-05,
"weight_decay": 0.005456725485601475,
"backbone_lr_ratio": 0.07591104805282695
},
"user_attrs": {
"epoch_mean": 30,
"fold_best_epochs": [
30,
30,
30
],
"fold_val_losses": [
12.178006362915038,
12.020879745483398,
11.571882820129394
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 10,
"value": 2.744228076934814,
"params": {
"dropout": 0.48159581757804304,
"lr": 0.00013112992007873722,
"weight_decay": 1.0719864040714885e-05,
"backbone_lr_ratio": 0.8131735182005935
},
"user_attrs": {
"epoch_mean": 25,
"fold_best_epochs": [
21,
29,
24
],
"fold_val_losses": [
3.2794203758239746,
2.4811975240707396,
2.472066330909729
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 11,
"value": 2.600485173861186,
"params": {
"dropout": 0.1398149921035961,
"lr": 0.0003676766469931975,
"weight_decay": 0.0814481577934799,
"backbone_lr_ratio": 0.416600060135216
},
"user_attrs": {
"epoch_mean": 14,
"fold_best_epochs": [
8,
14,
21
],
"fold_val_losses": [
3.2930724143981935,
2.312326765060425,
2.196056342124939
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 12,
"value": 2.856875479221344,
"params": {
"dropout": 0.39786435978029433,
"lr": 0.00027104165213579705,
"weight_decay": 0.09324179801968951,
"backbone_lr_ratio": 0.24792013326598586
},
"user_attrs": {
"epoch_mean": 17,
"fold_best_epochs": [
14,
12,
25
],
"fold_val_losses": [
3.5214386940002442,
2.671292519569397,
2.3778952240943907
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 13,
"value": 2.5929951747258504,
"params": {
"dropout": 0.1708198460450869,
"lr": 0.0002771027430797957,
"weight_decay": 0.022371906875975713,
"backbone_lr_ratio": 0.8529032090320552
},
"user_attrs": {
"epoch_mean": 11,
"fold_best_epochs": [
9,
12,
13
],
"fold_val_losses": [
3.211256170272827,
2.29467031955719,
2.273059034347534
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 14,
"value": 2.7000243584314982,
"params": {
"dropout": 0.10319674017011438,
"lr": 5.318587054778435e-05,
"weight_decay": 0.026386070022612545,
"backbone_lr_ratio": 0.3335324680299515
},
"user_attrs": {
"epoch_mean": 28,
"fold_best_epochs": [
26,
29,
30
],
"fold_val_losses": [
3.095314884185791,
2.559102272987366,
2.445655918121338
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 15,
"value": 2.811073144276937,
"params": {
"dropout": 0.17491537723907288,
"lr": 0.0004715550190685932,
"weight_decay": 0.020503728356529433,
"backbone_lr_ratio": 0.202005130829578
},
"user_attrs": {
"epoch_mean": 15,
"fold_best_epochs": [
11,
13,
22
],
"fold_val_losses": [
3.665318727493286,
2.5160216808319094,
2.251879024505615
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 16,
"value": 2.51602889696757,
"params": {
"dropout": 0.157877413230548,
"lr": 0.00016221844139316182,
"weight_decay": 0.0011695139073230128,
"backbone_lr_ratio": 0.9818226878189111
},
"user_attrs": {
"epoch_mean": 19,
"fold_best_epochs": [
12,
21,
24
],
"fold_val_losses": [
2.967289590835571,
2.41019024848938,
2.170606851577759
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 17,
"value": 2.617725872993469,
"params": {
"dropout": 0.1975499843454273,
"lr": 6.0680145665816345e-05,
"weight_decay": 0.0006054421043842105,
"backbone_lr_ratio": 0.9370662382902515
},
"user_attrs": {
"epoch_mean": 28,
"fold_best_epochs": [
23,
30,
30
],
"fold_val_losses": [
3.112912178039551,
2.4265938997268677,
2.3136715412139894
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 18,
"value": 2.6030575354894006,
"params": {
"dropout": 0.2840632229983807,
"lr": 0.00016641263642332344,
"weight_decay": 0.0015861751549934148,
"backbone_lr_ratio": 0.5813736146009224
},
"user_attrs": {
"epoch_mean": 23,
"fold_best_epochs": [
21,
21,
27
],
"fold_val_losses": [
3.163059639930725,
2.3479116916656495,
2.2982012748718263
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 19,
"value": 2.865672874450684,
"params": {
"dropout": 0.10525612628712336,
"lr": 6.627976747186786e-05,
"weight_decay": 0.00020126720225868128,
"backbone_lr_ratio": 0.12579626206936234
},
"user_attrs": {
"epoch_mean": 30,
"fold_best_epochs": [
30,
30,
30
],
"fold_val_losses": [
3.413493585586548,
2.58580904006958,
2.5977159976959228
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 20,
"value": 2.541643071174622,
"params": {
"dropout": 0.24563678706364706,
"lr": 0.0002197507105689095,
"weight_decay": 0.01585496331337175,
"backbone_lr_ratio": 0.2693624761288439
},
"user_attrs": {
"epoch_mean": 19,
"fold_best_epochs": [
15,
18,
24
],
"fold_val_losses": [
3.2464155673980715,
2.2392319440841675,
2.139281702041626
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 21,
"value": 2.5363871733347576,
"params": {
"dropout": 0.24822650356184767,
"lr": 0.00018488380020633278,
"weight_decay": 0.016124874730349143,
"backbone_lr_ratio": 0.30071816195793755
},
"user_attrs": {
"epoch_mean": 21,
"fold_best_epochs": [
13,
26,
24
],
"fold_val_losses": [
3.1824454784393312,
2.337935042381287,
2.088780999183655
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 22,
"value": 2.608898858229319,
"params": {
"dropout": 0.23999254948655954,
"lr": 0.00019074146892949267,
"weight_decay": 0.007473456802129268,
"backbone_lr_ratio": 0.24792013326598586
},
"user_attrs": {
"epoch_mean": 18,
"fold_best_epochs": [
12,
18,
25
],
"fold_val_losses": [
3.299156427383423,
2.3785044670104982,
2.1490356802940367
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 23,
"value": 2.6372965017954506,
"params": {
"dropout": 0.2732819324439319,
"lr": 8.559450244862353e-05,
"weight_decay": 0.003126023459511234,
"backbone_lr_ratio": 0.3168235172426432
},
"user_attrs": {
"epoch_mean": 27,
"fold_best_epochs": [
28,
22,
30
],
"fold_val_losses": [
3.0834681034088134,
2.466027092933655,
2.3623943090438844
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 24,
"value": 2.482991655667623,
"params": {
"dropout": 0.30169410062526664,
"lr": 0.000178917590650789,
"weight_decay": 0.013098468879191921,
"backbone_lr_ratio": 0.5322698244647086
},
"user_attrs": {
"epoch_mean": 19,
"fold_best_epochs": [
18,
16,
23
],
"fold_val_losses": [
2.952860403060913,
2.3297120332717896,
2.166402530670166
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 25,
"value": 2.767317620913188,
"params": {
"dropout": 0.3194613602400577,
"lr": 4.342839523653342e-05,
"weight_decay": 0.001138022773453874,
"backbone_lr_ratio": 0.575108040368093
},
"user_attrs": {
"epoch_mean": 30,
"fold_best_epochs": [
30,
30,
29
],
"fold_val_losses": [
3.0871256828308105,
2.6297523975372314,
2.585074782371521
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 26,
"value": 2.6014270941416426,
"params": {
"dropout": 0.380465038157504,
"lr": 0.000131818792969444,
"weight_decay": 0.031283173085439854,
"backbone_lr_ratio": 0.5275684824553009
},
"user_attrs": {
"epoch_mean": 23,
"fold_best_epochs": [
21,
25,
24
],
"fold_val_losses": [
3.164607048034668,
2.430464816093445,
2.209209418296814
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 27,
"value": 2.627247397104899,
"params": {
"dropout": 0.20197960767997714,
"lr": 8.186282230228678e-05,
"weight_decay": 0.003540606605897514,
"backbone_lr_ratio": 0.17090703553802636
},
"user_attrs": {
"epoch_mean": 30,
"fold_best_epochs": [
29,
30,
30
],
"fold_val_losses": [
3.2599289894104,
2.3099464178085327,
2.311866784095764
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 28,
"value": 2.6080079317092895,
"params": {
"dropout": 0.29875866523156275,
"lr": 0.0004703885391544067,
"weight_decay": 0.045327015648549004,
"backbone_lr_ratio": 0.6534099986720939
},
"user_attrs": {
"epoch_mean": 15,
"fold_best_epochs": [
9,
12,
25
],
"fold_val_losses": [
3.350850296020508,
2.2914549112319946,
2.181718587875366
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 29,
"value": 2.6760944962501525,
"params": {
"dropout": 0.26572299607512456,
"lr": 0.0007379244149640388,
"weight_decay": 0.008589710887330617,
"backbone_lr_ratio": 0.40188133372791734
},
"user_attrs": {
"epoch_mean": 13,
"fold_best_epochs": [
9,
13,
17
],
"fold_val_losses": [
3.524716091156006,
2.295975375175476,
2.207592022418976
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 30,
"value": 3.544327608744304,
"params": {
"dropout": 0.3647346234383415,
"lr": 3.7911482246002866e-05,
"weight_decay": 0.011264417139383486,
"backbone_lr_ratio": 0.18158532922746584
},
"user_attrs": {
"epoch_mean": 30,
"fold_best_epochs": [
30,
30,
30
],
"fold_val_losses": [
3.795345687866211,
3.331618309020996,
3.5060188293457033
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 31,
"value": 2.553473687171936,
"params": {
"dropout": 0.24271079071613583,
"lr": 0.0002079970663621355,
"weight_decay": 0.01383692673296454,
"backbone_lr_ratio": 0.2928581085972686
},
"user_attrs": {
"epoch_mean": 20,
"fold_best_epochs": [
14,
19,
28
],
"fold_val_losses": [
3.171020269393921,
2.3204617738723754,
2.1689390182495116
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 32,
"value": 2.7147754907608026,
"params": {
"dropout": 0.3081417266050286,
"lr": 0.00015961988076370855,
"weight_decay": 0.005647416385662339,
"backbone_lr_ratio": 0.13126369611144165
},
"user_attrs": {
"epoch_mean": 25,
"fold_best_epochs": [
19,
28,
28
],
"fold_val_losses": [
3.3805705070495606,
2.4436901807785034,
2.3200657844543455
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 33,
"value": 2.8405439217885338,
"params": {
"dropout": 0.19697712607188347,
"lr": 0.00022483991847689944,
"weight_decay": 0.04396257638256477,
"backbone_lr_ratio": 0.01101100938924494
},
"user_attrs": {
"epoch_mean": 26,
"fold_best_epochs": [
19,
30,
30
],
"fold_val_losses": [
3.8356207847595214,
2.3643120765686034,
2.3216989040374756
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 34,
"value": 2.503487364451091,
"params": {
"dropout": 0.2561341328718051,
"lr": 0.00010300416486373038,
"weight_decay": 0.012808904827686988,
"backbone_lr_ratio": 0.6727535050162337
},
"user_attrs": {
"epoch_mean": 24,
"fold_best_epochs": [
17,
27,
28
],
"fold_val_losses": [
3.041924476623535,
2.3080700874328612,
2.160467529296875
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 35,
"value": 2.4608302156130475,
"params": {
"dropout": 0.15463920314801388,
"lr": 0.00011604934996458467,
"weight_decay": 0.001746884129583056,
"backbone_lr_ratio": 0.7073818888939638
},
"user_attrs": {
"epoch_mean": 18,
"fold_best_epochs": [
16,
16,
23
],
"fold_val_losses": [
2.9413320541381838,
2.2545334100723267,
2.1866251826286316
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 36,
"value": 2.452152466773987,
"params": {
"dropout": 0.14731648771233646,
"lr": 0.00010399668603456206,
"weight_decay": 0.0007696017715340258,
"backbone_lr_ratio": 0.7472270478607838
},
"user_attrs": {
"epoch_mean": 23,
"fold_best_epochs": [
19,
21,
29
],
"fold_val_losses": [
2.921252131462097,
2.245358777046204,
2.1898464918136598
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 37,
"value": 2.5273226141929626,
"params": {
"dropout": 0.14194609604277394,
"lr": 0.00010021053483316553,
"weight_decay": 0.00038614379642781024,
"backbone_lr_ratio": 0.7094861329105961
},
"user_attrs": {
"epoch_mean": 21,
"fold_best_epochs": [
17,
19,
28
],
"fold_val_losses": [
3.076446294784546,
2.3063220262527464,
2.1991995215415954
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 38,
"value": 2.8850494384765626,
"params": {
"dropout": 0.2174541133163097,
"lr": 3.3739662416512915e-05,
"weight_decay": 0.002017506096408397,
"backbone_lr_ratio": 0.4899201254858191
},
"user_attrs": {
"epoch_mean": 30,
"fold_best_epochs": [
30,
30,
29
],
"fold_val_losses": [
3.2505751609802247,
2.640340280532837,
2.764232873916626
]
},
"state": "TrialState.COMPLETE"
},
{
"number": 39,
"value": 2.597004270553589,
"params": {
"dropout": 0.4276269504938004,
"lr": 0.0001196820298534061,
"weight_decay": 0.00035992145065850994,
"backbone_lr_ratio": 0.7607765936389007
},
"user_attrs": {
"epoch_mean": 25,
"fold_best_epochs": [
24,
24,
27
],
"fold_val_losses": [
3.053041934967041,
2.4371867895126345,
2.3007840871810914
]
},
"state": "TrialState.COMPLETE"
}
]

View File

@ -0,0 +1,60 @@
{
"original_strata_counts": {
"T0|P0|E0": "5",
"T0|P0|E1": "58",
"T0|P0|E2": "169",
"T0|P1|E0": "1",
"T0|P1|E1": "17",
"T0|P1|E2": "32",
"T0|P2|E2": "3",
"T1|P0|E2": "9",
"T1|P1|E2": "5",
"TNA|P0|E0": "28",
"TNA|P0|E1": "21",
"TNA|P0|E2": "20",
"TNA|P1|E0": "29",
"TNA|P1|E1": "7",
"TNA|P1|E2": "14",
"TNA|P2|E2": "1",
"TNA|P3|E0": "1"
},
"rare_strata": [
"T0|P1|E0",
"T0|P2|E2",
"TNA|P2|E2",
"TNA|P3|E0"
],
"final_strata": [
"RARE",
"T0|P0|E0",
"T0|P0|E1",
"T0|P0|E2",
"T0|P1|E1",
"T0|P1|E2",
"T1|P0|E2",
"T1|P1|E2",
"TNA|P0|E0",
"TNA|P0|E1",
"TNA|P0|E2",
"TNA|P1|E0",
"TNA|P1|E1",
"TNA|P1|E2"
],
"final_strata_counts": {
"RARE": "6",
"T0|P0|E0": "5",
"T0|P0|E1": "58",
"T0|P0|E2": "169",
"T0|P1|E1": "17",
"T0|P1|E2": "32",
"T1|P0|E2": "9",
"T1|P1|E2": "5",
"TNA|P0|E0": "28",
"TNA|P0|E1": "21",
"TNA|P0|E2": "20",
"TNA|P1|E0": "29",
"TNA|P1|E1": "7",
"TNA|P1|E2": "14"
},
"n_rare_merged": "6"
}

File diff suppressed because one or more lines are too long

View File

@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:0ea823b87a22a3a87cb487a9e876a6605c27a138abe53eaf228bb348dad43667
size 15970234

View File

@ -0,0 +1,254 @@
{
"train": [
{
"loss": 0.8207380050244176,
"n_samples": 8236
},
{
"loss": 0.7037654354365026,
"n_samples": 8236
},
{
"loss": 0.641996113549974,
"n_samples": 8236
},
{
"loss": 0.6070571800584409,
"n_samples": 8236
},
{
"loss": 0.5664057664468255,
"n_samples": 8236
},
{
"loss": 0.5478297611675892,
"n_samples": 8236
},
{
"loss": 0.5225721440771472,
"n_samples": 8236
},
{
"loss": 0.5069099566408475,
"n_samples": 8236
},
{
"loss": 0.4883425224024211,
"n_samples": 8236
},
{
"loss": 0.47361071388254816,
"n_samples": 8236
},
{
"loss": 0.4710996970063914,
"n_samples": 8236
},
{
"loss": 0.4520228576202772,
"n_samples": 8236
},
{
"loss": 0.45232121680418574,
"n_samples": 8236
},
{
"loss": 0.4500566631222651,
"n_samples": 8236
},
{
"loss": 0.4274507362961827,
"n_samples": 8236
},
{
"loss": 0.41204158604984553,
"n_samples": 8236
},
{
"loss": 0.39526375005000886,
"n_samples": 8236
},
{
"loss": 0.3738111386011387,
"n_samples": 8236
},
{
"loss": 0.3710037660054526,
"n_samples": 8236
},
{
"loss": 0.36222530734255337,
"n_samples": 8236
},
{
"loss": 0.3630701918634412,
"n_samples": 8236
},
{
"loss": 0.36140536376143495,
"n_samples": 8236
},
{
"loss": 0.3487493085334346,
"n_samples": 8236
},
{
"loss": 0.3467851376116652,
"n_samples": 8236
},
{
"loss": 0.34048351907452196,
"n_samples": 8236
},
{
"loss": 0.3541595635627183,
"n_samples": 8236
},
{
"loss": 0.33575822587381465,
"n_samples": 8236
},
{
"loss": 0.3219511436570785,
"n_samples": 8236
},
{
"loss": 0.30643386410068923,
"n_samples": 8236
},
{
"loss": 0.3101392985258709,
"n_samples": 8236
},
{
"loss": 0.31629752242895165,
"n_samples": 8236
}
],
"val": [
{
"loss": 0.7836992233950629,
"n_samples": 1454
},
{
"loss": 0.7153084917606317,
"n_samples": 1454
},
{
"loss": 0.723381241791186,
"n_samples": 1454
},
{
"loss": 0.6935155311345726,
"n_samples": 1454
},
{
"loss": 0.6741305548682337,
"n_samples": 1454
},
{
"loss": 0.6921254230824593,
"n_samples": 1454
},
{
"loss": 0.6984183775181619,
"n_samples": 1454
},
{
"loss": 0.6875423658664813,
"n_samples": 1454
},
{
"loss": 0.6951277488855417,
"n_samples": 1454
},
{
"loss": 0.6681769074567246,
"n_samples": 1454
},
{
"loss": 0.7054754480863044,
"n_samples": 1454
},
{
"loss": 0.7176714997016088,
"n_samples": 1454
},
{
"loss": 0.676661756540427,
"n_samples": 1454
},
{
"loss": 1.1059426443120963,
"n_samples": 1454
},
{
"loss": 0.7007260356185525,
"n_samples": 1454
},
{
"loss": 0.7131651905576661,
"n_samples": 1454
},
{
"loss": 0.6576622851613135,
"n_samples": 1454
},
{
"loss": 0.6736111434322604,
"n_samples": 1454
},
{
"loss": 0.7036862504531461,
"n_samples": 1454
},
{
"loss": 0.6835714435970931,
"n_samples": 1454
},
{
"loss": 0.6450633911679831,
"n_samples": 1454
},
{
"loss": 0.6777869927014084,
"n_samples": 1454
},
{
"loss": 0.7330788161600472,
"n_samples": 1454
},
{
"loss": 0.6992498341911925,
"n_samples": 1454
},
{
"loss": 0.6878705129662767,
"n_samples": 1454
},
{
"loss": 0.6829031571397427,
"n_samples": 1454
},
{
"loss": 0.6679887015685091,
"n_samples": 1454
},
{
"loss": 0.730571337382436,
"n_samples": 1454
},
{
"loss": 0.902666473815005,
"n_samples": 1454
},
{
"loss": 0.6748524632053822,
"n_samples": 1454
},
{
"loss": 0.7610602847811608,
"n_samples": 1454
}
]
}

Binary file not shown.

View File

@ -0,0 +1,38 @@
import inspect
import pandas as pd
from transformers import AutoTokenizer
from lnp_ml.dataset import process_dataframe, SMILES_COL
from lnp_ml.modeling.layers.llm_prompt import LLMPromptEncoder
# 从源头读,避免脚本和模型的阈值各说各话
MAXLEN = inspect.signature(LLMPromptEncoder.__init__).parameters["max_length"].default
tok = AutoTokenizer.from_pretrained("models/qwen2.5-7b-instruct", trust_remote_code=True)
df = process_dataframe(pd.read_csv("data/interim/internal.csv"))
smis = sorted(set(df[SMILES_COL].dropna()), key=len)
enc = LLMPromptEncoder.__new__(LLMPromptEncoder) # 不加载 7B 权重,只借用格式化方法
enc._is_biot5 = False
# _build_rag_prompt 现在还要读分位边界。这里必须给非空值:留空会让 _qbin 返回空串,
# 测出来的 token 数比真实 prompt 少约 8/邻居,等于白测。
# 具体数值不影响长度Q1/5 和 Q5/5 token 数相同),只决定落在哪个桶。
enc._rag_pool_qedges = {
"delivery": [-0.80, -0.30, 0.20, 0.70],
"size": [-0.90, -0.20, 0.40, 1.10],
}
def nb_of(s):
return {"smiles": s, "sim": 0.812, "delivery": 1.234,
"extra": {"size": -0.456, "pdi": 1, "ee": 2, "toxic": 0,
"biodist": [0.12, 0.34, 0.21, 0.05, 0.18, 0.07, 0.03]}}
worst = smis[-1]
print(f"分子数={len(smis)} SMILES 长度 {len(smis[0])}~{len(worst)} 字符")
print(f"分子数={len(smis)} SMILES 长度 {len(smis[0])}~{len(worst)} 字符 max_length={MAXLEN}")
for k in (2, 4, 8):
# 最坏情况:目标分子和全部邻居都取最长的那条
p = LLMPromptEncoder._build_rag_prompt(enc, worst, [nb_of(worst)] * k)
n = len(tok(p)["input_ids"])
status = "OK" if n <= MAXLEN else f"超出 {n - MAXLEN} tokens会被截断"
print(f"最坏 rag_top_k={k}: {n:5d} tokens 余量 {MAXLEN - n:5d} {status}")

View File

@ -0,0 +1,8 @@
{
"dropout": 0.33, "lr": 0.0005, "weight_decay": 0.0001, "backbone_lr_ratio": 0.3,
"moe_n_experts": 4, "moe_top_k": 2, "moe_expert_hidden_mult": 1,
"llm_lora_r": 8,
"d_model": 256, "num_heads": 8, "n_attn_layers": 4,
"fusion_strategy": "attention", "head_hidden_dim": 128,
"set_transformer_block": "sab"
}

89
scripts/measure_cost.py Normal file
View File

@ -0,0 +1,89 @@
"""测量 MARLIN / Backbone 的参数量、峰值显存与单-LNP 延迟。"""
import json, time, argparse
from pathlib import Path
import numpy as np
import pandas as pd
import torch
from torch.utils.data import DataLoader, Subset
from lnp_ml.dataset import LNPDataset, collate_fn, process_dataframe
from lnp_ml.modeling.nested_cv_optuna import create_model, _build_rag_pool
def build_marlin(bp, use_llm):
"""use_llm=True -> MARLIN(s1_both); False -> Backbone(s1_baseline)。"""
llm_kwargs = dict(
reg_bypass="on",
use_rag=use_llm, rag_top_k=4, use_llm=use_llm,
llm_model_path="models/qwen2.5-7b-instruct",
llm_freeze=False, llm_use_lora=False, llm_use_qlora=use_llm,
use_soft_prompt=use_llm,
llm_lora_r=bp.get("llm_lora_r", 8), llm_lora_alpha=16, llm_lora_dropout=0.05,
)
return create_model(
d_model=bp["d_model"], num_heads=bp["num_heads"], n_attn_layers=bp["n_attn_layers"],
fusion_strategy=bp["fusion_strategy"], head_hidden_dim=bp["head_hidden_dim"],
dropout=bp["dropout"], use_mpnn=True, mpnn_device="cuda",
set_transformer_block=bp["set_transformer_block"],
use_moe=use_llm, moe_n_experts=bp["moe_n_experts"], moe_top_k=bp["moe_top_k"],
moe_expert_hidden_mult=bp["moe_expert_hidden_mult"],
llm_kwargs=llm_kwargs,
)
@torch.no_grad()
def measure(model, full, train_idx, test_idx, device, use_llm, warmup=3):
model.eval().to(device)
tot = sum(p.numel() for p in model.parameters())
tr = sum(p.numel() for p in model.parameters() if p.requires_grad)
if use_llm: # RAG 前向需要检索池(只用训练集,防泄漏)
s, d, ex = _build_rag_pool(full, train_idx)
model.llm_prompt.set_retrieval_pool(s, d, pool_id="cost", extra_labels=ex)
loader = DataLoader(Subset(full, test_idx.tolist()), batch_size=1,
shuffle=False, collate_fn=collate_fn)
it = iter(loader)
for _ in range(warmup):
b = next(it)
model(b["smiles"], {k: v.to(device) for k, v in b["tabular"].items()})
torch.cuda.synchronize()
torch.cuda.reset_peak_memory_stats()
t0 = time.time(); n = 0
for b in loader:
model(b["smiles"], {k: v.to(device) for k, v in b["tabular"].items()}); n += 1
torch.cuda.synchronize()
dt_ms = (time.time() - t0) / n * 1000
peak_gb = torch.cuda.max_memory_allocated() / 1e9
return tot, tr, peak_gb, dt_ms
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--input", default="data/interim/internal.csv")
ap.add_argument("--run", default="models/abl_full/s1_both/seed42")
ap.add_argument("--fold", type=int, default=0)
args = ap.parse_args()
device = torch.device("cuda")
df = process_dataframe(pd.read_csv(args.input))
full = LNPDataset(df)
run = Path(args.run)
bp = json.load(open(run / "summary.json"))["fold_results"][args.fold]["best_params"]
sp = json.load(open(run / f"outer_fold_{args.fold}" / "splits.json"))
tr_idx = np.array(sp["outer_train_idx"]); te_idx = np.array(sp["outer_test_idx"])
for name, use_llm in [("Backbone", False), ("MARLIN", True)]:
model = build_marlin(bp, use_llm)
tot, trn, mem, lat = measure(model, full, tr_idx, te_idx, device, use_llm)
print(f"{name:9s} Total={tot/1e6:8.2f}M Trainable={trn/1e6:7.3f}M "
f"({100*trn/tot:.3f}%) PeakMem={mem:5.2f}GB Latency/LNP={lat:6.1f}ms")
del model; torch.cuda.empty_cache()
if __name__ == "__main__":
main()

View File

@ -0,0 +1,21 @@
#!/usr/bin/env bash
set -uo pipefail
SEED=${SEED:-42}
MODE=${MODE:-variant} # variant=两卡各跑一个 reg_bypass 变体shard=两卡分 fold
if [ "${MODE}" = "shard" ]; then
RB=${REG_BYPASS:-off}
GPU=0 FOLDS=0,1,2 SEED=${SEED} REG_BYPASS=${RB} BATCH=8 \
setsid nohup bash scripts_run/run_final_cv.sh >/dev/null 2>&1 &
GPU=1 FOLDS=3,4 SEED=${SEED} REG_BYPASS=${RB} BATCH=4 \
setsid nohup bash scripts_run/run_final_cv.sh >/dev/null 2>&1 &
else
GPU=0 SEED=${SEED} REG_BYPASS=off BATCH=8 MIN_FREE_MB=14000 \
setsid nohup bash scripts_run/run_final_cv.sh >/dev/null 2>&1 &
GPU=1 SEED=${SEED} REG_BYPASS=on BATCH=4 MIN_FREE_MB=11000 \
setsid nohup bash scripts_run/run_final_cv.sh >/dev/null 2>&1 &
fi
sleep 2
echo "已拉起 MODE=${MODE} SEED=${SEED} PRETRAIN=${PRETRAIN:-<未设置>}"
echo "日志models/final_cv/*/gpu*.log"

View File

@ -0,0 +1,95 @@
#!/usr/bin/env bash
# 完整 nested CVMPNN + MoE + Qwen2.5-7B QLoRA + soft-RAG
# 单卡一个分片。被 OOM / 抢占 / 断连杀掉后自动等显存并续跑。
set -uo pipefail
GPU=${GPU:?必须指定,例如 GPU=0}
SEED=${SEED:-42}
FOLDS=${FOLDS:-} # 空=跑全部 5 折;"0,1,2"=只跑这三折
REG_BYPASS=${REG_BYPASS:-on} # on=delivery/size 绕开 MoE+LLMoff=接入
BATCH=${BATCH:-8} # 显存紧就设 4
N_OUTER=${N_OUTER:-5}
N_INNER=${N_INNER:-3}
N_TRIALS=${N_TRIALS:-20}
EPOCHS=${EPOCHS:-20}
PATIENCE=${PATIENCE:-5}
MIN_FREE_MB=${MIN_FREE_MB:-14000}
MAX_RETRY=${MAX_RETRY:-50}
RETRY_WAIT=${RETRY_WAIT:-300}
MIN_OK_SEC=${MIN_OK_SEC:-180} # 存活不足这么久就挂 → 判为配置/代码错误,立即停止重试
PRETRAIN=${PRETRAIN:-} # 非空且文件存在才注入 --init-from-pretrain
RUN_DIR="models/final_cv/${REG_BYPASS}_seed${SEED}"
TAG="gpu${GPU}$([ -n "${FOLDS}" ] && echo "_folds$(echo "${FOLDS}" | tr -d ',')")"
LOG="${RUN_DIR}/${TAG}.log"
export TRANSFORMERS_OFFLINE=1
export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
export TOKENIZERS_PARALLELISM=false
export CUDA_VISIBLE_DEVICES=${GPU}
mkdir -p "${RUN_DIR}"
# nvidia-smi 不受 CUDA_VISIBLE_DEVICES 影响,-i 用物理卡号
wait_free() {
while :; do
local free
free=$(nvidia-smi --query-gpu=memory.free --format=csv,noheader,nounits -i "${GPU}")
[ "${free}" -ge "${MIN_FREE_MB}" ] && return 0
echo "[$(date '+%F %T')] GPU${GPU} 仅空 ${free}MiB < ${MIN_FREE_MB}MiB60s 后重试" >>"${LOG}"
sleep 60
done
}
EXTRA=()
if [ -n "${FOLDS}" ]; then
# 注意:需先给 nested_cv_optuna.main 增加 --only-folds 选项,否则 Typer 直接报错
EXTRA+=(--only-folds "${FOLDS}")
fi
if [ -n "${PRETRAIN}" ] && [ -f "${PRETRAIN}" ]; then
EXTRA+=(--init-from-pretrain "${PRETRAIN}")
elif [ -n "${PRETRAIN}" ]; then
echo "[$(date '+%F %T')] 警告:${PRETRAIN} 不存在,跳过预训练初始化" >>"${LOG}"
fi
for attempt in $(seq 1 "${MAX_RETRY}"); do
wait_free
echo "[$(date '+%F %T')] ATTEMPT ${attempt}/${MAX_RETRY} gpu=${GPU} seed=${SEED} folds=${FOLDS:-all} reg_bypass=${REG_BYPASS} batch=${BATCH} trials=${N_TRIALS}" >>"${LOG}"
t0=${SECONDS}
python -u -m lnp_ml.modeling.nested_cv_optuna \
--input-path data/interim/internal.csv \
--output-dir models/final_cv \
--resume-dir "${RUN_DIR}" \
--seed ${SEED} \
--use-mpnn \
--use-moe \
--reg-bypass ${REG_BYPASS} \
--use-llm --use-rag --rag-top-k 4 \
--use-soft-prompt --llm-use-qlora \
--llm-model-path models/qwen2.5-7b-instruct \
--no-llm-freeze \
--llm-max-length 1536 \
--n-outer-folds ${N_OUTER} --n-inner-folds ${N_INNER} \
--n-trials ${N_TRIALS} --epochs-per-trial ${EPOCHS} \
--inner-patience ${PATIENCE} \
--batch-size ${BATCH} --device cuda \
${EXTRA[@]+"${EXTRA[@]}"} \
>>"${LOG}" 2>&1
rc=$?
dt=$((SECONDS - t0))
if [ ${rc} -eq 0 ]; then
echo "[$(date '+%F %T')] DONE gpu=${GPU} folds=${FOLDS:-all}(用时 ${dt}s" >>"${LOG}"
exit 0
fi
if [ ${dt} -lt ${MIN_OK_SEC} ]; then
echo "[$(date '+%F %T')] 仅存活 ${dt}s 即以 ${rc} 退出,判定为配置/代码错误而非抢占,停止重试" >>"${LOG}"
exit ${rc}
fi
echo "[$(date '+%F %T')] 运行 ${dt}s 后以 ${rc} 退出,${RETRY_WAIT}s 后带 --resume-dir 续跑" >>"${LOG}"
sleep "${RETRY_WAIT}"
done
echo "[$(date '+%F %T')] 超过 MAX_RETRY 次仍失败,放弃" >>"${LOG}"
exit 1

View File

@ -0,0 +1,109 @@
#!/usr/bin/env bash
# 全量数据 final 模型3-fold Optuna 调参 + 全量固定 epoch 重训
# MPNN + MoE + Qwen2.5-7B QLoRA + soft-RAG
# 被 OOM / 抢占 / SSH 断连杀掉后自动等显存并续跑Optuna 从 sqlite 恢复)
set -uo pipefail
GPU=${GPU:-0}
SEED=${SEED:-42}
OUT=${OUT:-models/final}
N_TRIALS=${N_TRIALS:-20}
EPOCHS=${EPOCHS:-20} # 与 nested CV 的 EPOCHS 保持一致
PATIENCE=${PATIENCE:-5} # 与 nested CV 的 PATIENCE 保持一致
N_FOLDS=${N_FOLDS:-3}
BATCH=${BATCH:-8}
REG_BYPASS=${REG_BYPASS:-off}
FREEZE=${FREEZE:-3}
PRETRAIN=${PRETRAIN:-models/pretrain/mpnn/pretrain_delivery.pt}
MIN_FREE_MB=${MIN_FREE_MB:-14000}
MAX_RETRY=${MAX_RETRY:-50}
RETRY_WAIT=${RETRY_WAIT:-300}
MIN_OK_SEC=${MIN_OK_SEC:-180} # 存活不足这么久就挂 → 判为配置/代码错误,立即停止重试
LOG="${OUT}/train.log"
STUDY="${OUT}/optuna_study.sqlite3"
export TRANSFORMERS_OFFLINE=1
export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
export TOKENIZERS_PARALLELISM=false
export CUDA_VISIBLE_DEVICES=${GPU}
mkdir -p "${OUT}"
# nvidia-smi 不受 CUDA_VISIBLE_DEVICES 影响,-i 用物理卡号
wait_free() {
while :; do
local free
free=$(nvidia-smi --query-gpu=memory.free --format=csv,noheader,nounits -i "${GPU}")
[ "${free}" -ge "${MIN_FREE_MB}" ] && return 0
echo "[$(date '+%F %T')] GPU${GPU} 仅空 ${free}MiB < ${MIN_FREE_MB}MiB60s 后重试" >>"${LOG}"
sleep 60
done
}
# Optuna 的 study.optimize(n_trials=N) 在续跑时语义是"再跑 N 个",不是"补到 N 个"。
# 所以每次重启前先数已完成的 trial 只补差额,否则被抢占几次就会多跑几十个 trial。
remaining_trials() {
if [ ! -f "${STUDY}" ]; then echo "${N_TRIALS}"; return; fi
python - "${STUDY}" "${N_TRIALS}" <<'PY' 2>/dev/null || echo "${N_TRIALS}"
import sys, optuna
optuna.logging.set_verbosity(optuna.logging.WARNING)
try:
st = optuna.load_study(study_name="final_optuna_cv", storage=f"sqlite:///{sys.argv[1]}")
done = sum(1 for t in st.trials if t.state == optuna.trial.TrialState.COMPLETE)
except Exception:
done = 0
print(max(0, int(sys.argv[2]) - done))
PY
}
EXTRA=()
if [ -n "${PRETRAIN}" ] && [ -f "${PRETRAIN}" ]; then
EXTRA+=(--init-from-pretrain "${PRETRAIN}")
elif [ -n "${PRETRAIN}" ]; then
echo "[$(date '+%F %T')] 警告:${PRETRAIN} 不存在,跳过预训练初始化" >>"${LOG}"
fi
for attempt in $(seq 1 "${MAX_RETRY}"); do
wait_free
NT=$(remaining_trials)
echo "[$(date '+%F %T')] ATTEMPT ${attempt}/${MAX_RETRY} gpu=${GPU} 本次补 ${NT} 个 trial目标 ${N_TRIALS}reg_bypass=${REG_BYPASS} batch=${BATCH}" >>"${LOG}"
t0=${SECONDS}
python -u -m lnp_ml.modeling.final_train_optuna_cv \
--input-path data/interim/internal.csv \
--output-dir "${OUT}" \
--seed ${SEED} \
--n-folds ${N_FOLDS} \
--n-trials ${NT} \
--epochs-per-trial ${EPOCHS} \
--patience ${PATIENCE} \
--batch-size ${BATCH} \
--use-mpnn --use-moe \
--reg-bypass ${REG_BYPASS} \
--use-llm --use-rag --rag-top-k 4 \
--use-soft-prompt --llm-use-qlora \
--llm-model-path models/qwen2.5-7b-instruct \
--no-llm-freeze \
--llm-max-length 1536 \
--freeze-backbone-epochs ${FREEZE} \
--device cuda \
${EXTRA[@]+"${EXTRA[@]}"} \
>>"${LOG}" 2>&1
rc=$?
dt=$((SECONDS - t0))
if [ ${rc} -eq 0 ]; then
echo "[$(date '+%F %T')] DONE用时 ${dt}s${OUT}/model.pt" >>"${LOG}"
exit 0
fi
if [ ${dt} -lt ${MIN_OK_SEC} ]; then
echo "[$(date '+%F %T')] 仅存活 ${dt}s 即以 ${rc} 退出,判为配置/代码错误而非抢占,停止重试" >>"${LOG}"
exit ${rc}
fi
echo "[$(date '+%F %T')] 运行 ${dt}s 后以 ${rc} 退出,${RETRY_WAIT}s 后续跑" >>"${LOG}"
sleep "${RETRY_WAIT}"
done
echo "[$(date '+%F %T')] 超过 MAX_RETRY 次仍失败,放弃" >>"${LOG}"
exit 1

View File

@ -4,7 +4,8 @@ import os
import numpy as np import numpy as np
import torch import torch
BASE = "models/abl" BASE = os.environ.get("GATES_BASE", "models/abl")
VARIANTS = os.environ.get("GATES_VARIANTS", "moe,llm,both").split(",")
TASKS = ["size", "delivery", "pdi", "ee", "toxic", "biodist"] TASKS = ["size", "delivery", "pdi", "ee", "toxic", "biodist"]
@ -13,7 +14,7 @@ def latest(v: str):
return dirs[-1] if dirs else None return dirs[-1] if dirs else None
for v in ["moe", "llm", "both"]: for v in VARIANTS:
run = latest(v) run = latest(v)
if run is None: if run is None:
print(f"\n=== {v}: 无 run 目录 ===") print(f"\n=== {v}: 无 run 目录 ===")
@ -28,14 +29,15 @@ for v in ["moe", "llm", "both"]:
gm, gl = [], [] gm, gl = [], []
for f in files: for f in files:
sd = torch.load(f, map_location="cpu", weights_only=False)["model_state_dict"] sd = torch.load(f, map_location="cpu", weights_only=False)["model_state_dict"]
if "fusion.g_moe" in sd: for key, acc in (("fusion.g_moe", gm), ("fusion.g_llm", gl)):
gm.append(sd["fusion.g_moe"].item()) if key not in sd:
if "fusion.g_llm" in sd: continue
gl.append(sd["fusion.g_llm"].item()) t = sd[key]
if gm: acc.append(abs(t.item()) if t.dim() == 0
print(f" g_moe: mean={np.mean(gm):+.4f} |abs|=[{min(map(abs, gm)):.4f},{max(map(abs, gm)):.4f}]") else float(t.norm()) / (t.numel() ** 0.5))
if gl: for name, acc in (("g_moe", gm), ("g_llm", gl)):
print(f" g_llm: mean={np.mean(gl):+.4f} |abs|=[{min(map(abs, gl)):.4f},{max(map(abs, gl)):.4f}]") if acc:
print(f" {name}: mean={np.mean(acc):.4f} range=[{min(acc):.4f},{max(acc):.4f}]")
sd = torch.load(files[0], map_location="cpu", weights_only=False)["model_state_dict"] sd = torch.load(files[0], map_location="cpu", weights_only=False)["model_state_dict"]
for t in TASKS: for t in TASKS: