diff --git a/app/api.py b/app/api.py index 6fed55a..82ec2a7 100644 --- a/app/api.py +++ b/app/api.py @@ -87,6 +87,7 @@ class OptimizeRequest(BaseModel): 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)") 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])") 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="组分范围配置(默认使用标准范围)") @@ -188,6 +189,20 @@ async def lifespan(app: FastAPI): try: state.model = load_model(model_path, state.device) 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: logger.error(f"Failed to load model: {e}") raise @@ -305,6 +320,7 @@ async def optimize_formulation(request: OptimizeRequest): comp_ranges=comp_ranges, routes=request.routes, scoring_weights=scoring_weights, + rerank_top_n=request.rerank_top_n, batch_size=256, ) @@ -350,6 +366,8 @@ async def optimize_formulation(request: OptimizeRequest): except Exception as 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)) diff --git a/app/app.py b/app/app.py index ec819ea..affec88 100644 --- a/app/app.py +++ b/app/app.py @@ -175,9 +175,7 @@ def call_optimize_api( # PDI 分类标签 PDI_CLASS_LABELS = { 0: "<0.2 (优)", - 1: "0.2-0.3 (良)", - 2: "0.3-0.4 (中)", - 3: ">0.4 (差)", + 1: "≥0.2 (欠佳)", } # EE 分类标签 diff --git a/app/optimize.py b/app/optimize.py index 5a362ab..5809f8a 100644 --- a/app/optimize.py +++ b/app/optimize.py @@ -701,6 +701,43 @@ def select_top_k( 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( seed: Formulation, mol_step: float, @@ -780,6 +817,8 @@ def optimize( routes: Optional[List[str]] = None, scoring_weights: Optional[ScoringWeights] = None, batch_size: int = 256, + rerank_top_n: int = 0, + rerank_batch_size: int = 16, ) -> List[Formulation]: """ 执行配方优化(层级搜索策略)。 @@ -843,6 +882,14 @@ def optimize( logger.info(f"Comp ranges: {comp_ranges.to_dict()}") 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)): logger.info(f"\n{'='*60}") @@ -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"Best formulation: {best.to_dict()}") - # 最终去重、按综合评分排序并返回 top_k + # 粗筛结果:去重、按综合评分排序 seeds_sorted = sorted(seeds, key=_score, reverse=True) - - # 去重:保留每个唯一配方中得分最高的(已排序,所以第一个出现的就是最高的) + seen_keys = set() unique_results = [] for f in seeds_sorted: @@ -935,10 +981,31 @@ def optimize( if key not in seen_keys: seen_keys.add(key) unique_results.append(f) - - logger.info(f"Final results: {len(unique_results)} unique formulations (from {len(seeds)} candidates)") - - return unique_results[:top_k] + + logger.info(f"Stage-1: {len(unique_results)} unique formulations (from {len(seeds)} candidates)") + + 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: diff --git a/lnp_ml/modeling/final_train_optuna_cv.py b/lnp_ml/modeling/final_train_optuna_cv.py index 4b0a8b8..ede7852 100644 --- a/lnp_ml/modeling/final_train_optuna_cv.py +++ b/lnp_ml/modeling/final_train_optuna_cv.py @@ -53,6 +53,7 @@ from lnp_ml.modeling.trainer_balanced import ( train_fixed_epochs, ) from lnp_ml.modeling.visualization import plot_multitask_loss_curves +from lnp_ml.modeling.nested_cv_optuna import _build_rag_pool # MPNN ensemble 默认路径 DEFAULT_MPNN_ENSEMBLE_DIR = MODELS_DIR / "mpnn" / "all_amine_split_for_LiON" @@ -178,20 +179,36 @@ def create_model( mpnn_device: str = "cpu", chemeleon_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, moe_n_experts: int = 4, moe_top_k: int = 2, moe_expert_hidden_mult: int = 2, 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]: - """创建模型""" - moe_kwargs = dict( + """创建模型。llm_kwargs 含 use_llm / llm_model_path / llm_freeze / llm_use_lora / llm_lora_* / reg_bypass。""" + extra_kwargs = dict( + set_transformer_block=set_transformer_block, 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, + 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: @@ -205,9 +222,7 @@ def create_model( dropout=dropout, mpnn_ensemble_paths=ensemble_paths, mpnn_device=mpnn_device, - chemeleon_cache_path=chemeleon_cache, - unimol_cache_path=unimol_cache, - **moe_kwargs, + **extra_kwargs, ) else: return LNPModelWithoutMPNN( @@ -217,9 +232,7 @@ def create_model( fusion_strategy=fusion_strategy, head_hidden_dim=head_hidden_dim, dropout=dropout, - chemeleon_cache_path=chemeleon_cache, - unimol_cache_path=unimol_cache, - **moe_kwargs, + **extra_kwargs, ) @@ -277,6 +290,10 @@ def run_optuna_cv( moe_top_k: int = 2, moe_expert_hidden_mult: int = 2, 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]: """ 使用全量数据做 3-fold CV Optuna 超参搜索。 @@ -331,29 +348,43 @@ def run_optuna_cv( lr = trial.suggest_float("lr", 1e-5, 1e-3, 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) - + + # 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 cv = StratifiedKFold(n_splits=n_folds, shuffle=True, random_state=seed) - + fold_val_losses = [] fold_best_epochs = [] - + for fold, (train_idx, val_idx) in enumerate(cv.split(indices, strata)): - # 创建 DataLoader train_subset = Subset(full_dataset, train_idx.tolist()) val_subset = Subset(full_dataset, val_idx.tolist()) - + train_loader = DataLoader( train_subset, batch_size=batch_size, shuffle=True, collate_fn=collate_fn ) val_loader = DataLoader( val_subset, batch_size=batch_size, shuffle=False, collate_fn=collate_fn ) - - # 计算类权重 + class_weights = compute_class_weights_from_loader(train_loader) - - # 创建模型 + model = create_model( d_model=d_model, num_heads=num_heads, @@ -365,22 +396,29 @@ def run_optuna_cv( mpnn_device=device.type, chemeleon_cache=chemeleon_cache, unimol_cache=unimol_cache, + set_transformer_block=set_transformer_block, 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, + 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: 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: load_pretrain_weights_to_model( model, pretrain_state_dict, d_model, pretrain_config, load_delivery_head ) - - # 训练(带早停) + result = train_with_early_stopping( model=model, train_loader=train_loader, @@ -392,8 +430,9 @@ def run_optuna_cv( patience=patience, class_weights=class_weights, backbone_lr_ratio=backbone_lr_ratio, + freeze_backbone_epochs=freeze_backbone_epochs, ) - + fold_val_losses.append(result["best_val_loss"]) fold_best_epochs.append(result["best_epoch"]) @@ -427,6 +466,7 @@ def run_optuna_cv( "n_attn_layers": fixed_n_attn_layers, "fusion_strategy": fixed_fusion_strategy, "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) @@ -472,6 +512,27 @@ def main( moe_top_k: int = 2, moe_expert_hidden_mult: int = 2, 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", ): @@ -491,7 +552,27 @@ def main( logger.info(f"Using 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_config = None @@ -559,6 +640,10 @@ def main( moe_top_k=moe_top_k, moe_expert_hidden_mult=moe_expert_hidden_mult, 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: 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( d_model=best_params["d_model"], num_heads=best_params["num_heads"], @@ -619,13 +718,21 @@ def main( chemeleon_cache=(chemeleon_cache if use_chemeleon else None), unimol_cache=(unimol_cache if use_unimol else None), 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, + llm_kwargs=llm_kwargs_resolved, + use_retrieval=use_retrieval, + retr_feature_dim=(3 if use_retrieval else 0), + **arch, ) 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: loaded = load_pretrain_weights_to_model( @@ -634,19 +741,17 @@ def main( ) if loaded: logger.info("Loaded pretrain weights for final training") - - # 打印模型信息 + 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) 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 - + train_result = train_fixed_epochs( model=model, train_loader=full_loader, - val_loader=None, # 全量训练,无验证集 + val_loader=None, device=device, lr=best_params["lr"], weight_decay=best_params["weight_decay"], @@ -656,12 +761,12 @@ def main( use_swa=use_swa, swa_start_epoch=swa_start, backbone_lr_ratio=best_params.get("backbone_lr_ratio", 1.0), + freeze_backbone_epochs=freeze_backbone_epochs, ) - - # 加载最终权重 - model.load_state_dict(train_result["final_state"]) - - # 保存模型 + + # QLoRA 4-bit 基座的量化元数据键不在 final_state 里,必须 strict=False + model.load_state_dict(train_result["final_state"], strict=False) + config = { "d_model": best_params["d_model"], "num_heads": best_params["num_heads"], @@ -675,20 +780,29 @@ def main( "use_unimol": use_unimol, "unimol_cache": unimol_cache if use_unimol else None, "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, + "use_retrieval": use_retrieval, + "retr_feature_dim": (3 if use_retrieval else 0), + **arch, + **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({ - "model_state_dict": train_result["final_state"], + "model_state_dict": _slim_state, "config": config, "best_params": best_params, "epoch_mean": epoch_mean, "use_swa": use_swa, }, output_dir / "model.pt") - + logger.success(f"Saved model to {output_dir / 'model.pt'}") # 保存训练历史 diff --git a/lnp_ml/modeling/heads.py b/lnp_ml/modeling/heads.py index 2c4c89a..329b0ca 100644 --- a/lnp_ml/modeling/heads.py +++ b/lnp_ml/modeling/heads.py @@ -80,7 +80,7 @@ class MultiTaskHead(nn.Module): size_dropout = min(0.5, dropout + 0.2) 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) # Encapsulation Efficiency: 3 分类 diff --git a/lnp_ml/modeling/layers/fusion.py b/lnp_ml/modeling/layers/fusion.py index 7edeed5..54beeb3 100644 --- a/lnp_ml/modeling/layers/fusion.py +++ b/lnp_ml/modeling/layers/fusion.py @@ -114,8 +114,10 @@ class FusionLayer(nn.Module): class ResidualConcatFusion(nn.Module): """对真实 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: @@ -125,10 +127,18 @@ class ResidualConcatFusion(nn.Module): self.d_model = d_model self.pool = FusionLayer(d_model=d_model, n_tokens=1, strategy=strategy) self.fusion_dim = self.pool.fusion_dim - # 零初始化门控(可学习标量),旁路初始不参与 - self.g_moe = nn.Parameter(torch.zeros(())) - self.g_llm = nn.Parameter(torch.zeros(())) - self.g_retr = nn.Parameter(torch.zeros(())) # 检索旁路零初始化门控 + # 零初始化逐维门控,旁路初始不参与 + self.g_moe = nn.Parameter(torch.zeros(d_model)) + self.g_llm = nn.Parameter(torch.zeros(d_model)) + 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( self, @@ -145,13 +155,14 @@ class ResidualConcatFusion(nn.Module): if return_attn_weights: pooled, attn = pooled + # [d_model] 与 [B, d_model] 自动广播 out = pooled 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: out = out + self.g_llm * f_llm if f_retr is not None: - out = out + self.g_retr * f_retr # 检索旁路,零初始化门保证起点=不开 + out = out + self.g_retr * f_retr # 回归旁路:额外返回 pooled(纯 chem+tab,不含 f_llm/f_moe/f_retr) if return_attn_weights: diff --git a/lnp_ml/modeling/layers/llm_prompt.py b/lnp_ml/modeling/layers/llm_prompt.py index b5c6987..c530f00 100644 --- a/lnp_ml/modeling/layers/llm_prompt.py +++ b/lnp_ml/modeling/layers/llm_prompt.py @@ -12,6 +12,15 @@ DEFAULT_MOLT5_PATH = os.environ.get("MOLT5_PATH", "models/molt5-base") # 检索池支持的额外多任务(除 delivery 外) _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): """用 LLM 编码分子,输出 F_llm [B, d_model]。 @@ -35,7 +44,7 @@ class LLMPromptEncoder(nn.Module): lora_r: int = 8, lora_alpha: int = 16, lora_dropout: float = 0.05, - max_length: int = 256, + max_length: int = 1536, use_rag: bool = False, rag_top_k: int = 4, use_soft_prompt: bool = False, @@ -64,6 +73,9 @@ class LLMPromptEncoder(nn.Module): ) if _is_qwen and self.tokenizer.pad_token is None: self.tokenizer.pad_token = self.tokenizer.eos_token + if use_soft_prompt: + # soft token 后置要求文本右对齐,否则 soft 会落在 PAD 之后 + self.tokenizer.padding_side = "left" if _is_t5: 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_fps = 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): 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_extra = extra_labels 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: self._cache = {k: v for k, v in self._cache.items() if not k.startswith("RAG::")} self._prompt_cache.clear() @@ -213,6 +239,21 @@ class LLMPromptEncoder(nn.Module): return "[" + ", ".join(f"{x:.3f}" for x in v) + "]" 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: """按 backbone 期望格式化分子。 BioT5:SMILES -> SELFIES,用 ... 紧贴包裹(官方格式,token 间无空格); @@ -227,36 +268,54 @@ class LLMPromptEncoder(nn.Module): return f"{sfs}" def _build_rag_prompt(self, target_smiles: str, neighbors) -> str: - """构造 RAG prompt:原始 SMILES + 邻居多任务结果(numeric)。""" + """构造 RAG prompt:原始 SMILES + 邻居多任务结果(数值 + 训练池分位桶)。""" blocks = [] for rank, nb in enumerate(neighbors, 1): 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( - f"Retrieved sample {rank}:\n" - f"Molecule: {self._fmt_mol(nb['smiles'])}\n" - f"Similarity score: {nb['sim']:.3f}\n" - f"delivery_log: {self._fmt(nb['delivery'])}\n" - f"size_z: {self._fmt(ex.get('size'))}\n" - f"pdi_class: {self._fmt(ex.get('pdi'), 'int')}\n" - 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')}" + f"#{rank} sim={nb['sim']:.2f} {self._fmt_mol(nb['smiles'])} " + f"deliv={self._fmt(nb['delivery'])}{self._qbin('delivery', nb['delivery'])} " + f"size={self._fmt(ex.get('size'))}{self._qbin('size', ex.get('size'))} " + f"pdi={self._fmt_class(ex.get('pdi'), _PDI_LABELS)} " + f"ee={self._fmt_class(ex.get('ee'), _EE_LABELS)} " + f"tox={self._fmt_class(ex.get('toxic'), _TOX_LABELS)} bio={bio_s}" ) - retrieved_block = "\n\n".join(blocks) if blocks else "(no retrieved samples)" + retrieved_block = "\n".join(blocks) if blocks else "(none)" return ( - "Task: Encode the target LNP molecule into a retrieval-aware representation " - "for downstream multi-task property prediction. Do not output predictions.\n\n" - f"[Target Molecule]\nMolecule: {self._fmt_mol(target_smiles)}\n\n" - "[Retrieved Similar LNP Samples]\n" - "Retrieved from the training set by fingerprint similarity, with their known " - "multi-task outcomes (delivery_log, size_z, pdi_class, ee_class, toxic, biodist; " - "'unknown' means the measurement is missing):\n\n" - f"{retrieved_block}\n\n" - "[Encoding Instructions]\n" - "Capture the target structure and the retrieval evidence (structural similarity " - "and consistency of retrieved outcomes) into your internal representation." + "Encode the target LNP molecule into a retrieval-aware representation " + "for multi-task property prediction. Do not output predictions.\n" + f"[Target] {self._fmt_mol(target_smiles)}\n" + "[Retrieved] Nearest training molecules by fingerprint similarity, with " + "known outcomes. deliv=delivery_log and size=size_z are raw values, each " + f"followed by (Qi/{_N_QBINS}) = which {_N_QBINS}-quantile bin it falls into " + "among training molecules, Q1=lowest. pdi/ee/tox give the class index with " + "its meaning. bio=fraction in [lymph_nodes,heart,liver,spleen,lung,kidney," + "muscle] with the dominant organ named. 'unknown' = missing:\n" + f"{retrieved_block}\n" + "[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: if not self.use_rag: return self._fmt_mol(s) @@ -273,6 +332,7 @@ class LLMPromptEncoder(nn.Module): bt, padding=True, truncation=True, max_length=self.max_length, return_tensors="pt", ).to(device) + self._warn_if_truncated(enc) out = self.encoder(**enc).last_hidden_state # [B,L,H] lengths = enc["attention_mask"].sum(1) - 1 b = torch.arange(out.size(0), device=device) @@ -318,18 +378,33 @@ class LLMPromptEncoder(nn.Module): if soft_list: 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), 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: inputs_embeds = text_embeds attn_mask = text_mask - out = self.encoder(inputs_embeds=inputs_embeds, attention_mask=attn_mask).last_hidden_state - lengths = attn_mask.sum(1) - 1 - b = torch.arange(out.size(0), device=device) - feat = out[b, lengths.long(), :] # [B, H] 最后有效 token + # 左填充下位置编码必须由 mask 推出,否则 PAD 会把真实 token 的位置顶偏 + position_ids = attn_mask.long().cumsum(-1) - 1 + position_ids = position_ids.masked_fill(attn_mask == 0, 1) + + 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()) # ---------- 旧路径(保留,向后兼容)---------- diff --git a/lnp_ml/modeling/models.py b/lnp_ml/modeling/models.py index 395f083..2b5a136 100644 --- a/lnp_ml/modeling/models.py +++ b/lnp_ml/modeling/models.py @@ -20,7 +20,7 @@ from lnp_ml.modeling.layers import ( LLMPromptEncoder, ) 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"] @@ -128,6 +128,7 @@ class LNPModel(nn.Module): llm_lora_dropout: float = 0.05, use_rag: bool = False, rag_top_k: int = 4, + llm_max_length: int = 1536, use_retrieval: bool = False, retr_feature_dim: int = 0, ) -> None: @@ -246,6 +247,7 @@ class LNPModel(nn.Module): lora_dropout=llm_lora_dropout, use_rag=use_rag, rag_top_k=rag_top_k, + max_length=llm_max_length, use_soft_prompt=use_soft_prompt, ) else: @@ -269,6 +271,14 @@ class LNPModel(nn.Module): 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( self, smiles: List[str], @@ -324,7 +334,8 @@ class LNPModel(nn.Module): self._last_moe_extras = 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): f_llm = self.llm_prompt(smiles, chem=chem, tab=tab) else: @@ -336,10 +347,27 @@ class LNPModel(nn.Module): _feats_t = torch.as_tensor(_feats, dtype=chem.dtype, device=chem.device) 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) self._last_pooled = pooled # 纯数值向量,供回归 head 使用 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( self, stacked: torch.Tensor, @@ -492,13 +520,22 @@ class LNPModel(nn.Module): } unexpected = [] + loaded = [] model_state = self.state_dict() for k, v in filtered_state_dict.items(): if k in model_state and model_state[k].shape == v.shape: model_state[k] = v + loaded.append(k) else: 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) @@ -543,6 +580,7 @@ class LNPModelWithoutMPNN(LNPModel): llm_lora_dropout: float = 0.05, use_rag: bool = False, rag_top_k: int = 4, + llm_max_length: int = 1536, use_retrieval: bool = False, retr_feature_dim: int = 0, ) -> None: @@ -573,6 +611,7 @@ class LNPModelWithoutMPNN(LNPModel): use_llm=use_llm, use_rag=use_rag, rag_top_k=rag_top_k, + llm_max_length=llm_max_length, use_retrieval=use_retrieval, retr_feature_dim=retr_feature_dim, llm_model_path=llm_model_path, diff --git a/lnp_ml/modeling/nested_cv_optuna.py b/lnp_ml/modeling/nested_cv_optuna.py index e2ab035..fb1ca94 100644 --- a/lnp_ml/modeling/nested_cv_optuna.py +++ b/lnp_ml/modeling/nested_cv_optuna.py @@ -804,6 +804,19 @@ def _run_single_outer_fold( 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( d_model=best_params["d_model"], num_heads=best_params["num_heads"], @@ -818,15 +831,10 @@ def _run_single_outer_fold( moleculestm_cache=moleculestm_cache, mole_cache=mole_cache, use_moe=use_moe, - 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={**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 {})}, + llm_kwargs=llm_kwargs_resolved, use_retrieval=use_retrieval, retr_feature_dim=(3 if use_retrieval else 0), + **arch, ) model.rdkit_encoder._cache = rdkit_cache if use_retrieval: @@ -878,7 +886,6 @@ def _run_single_outer_fold( "n_attn_layers": best_params["n_attn_layers"], "fusion_strategy": best_params["fusion_strategy"], "head_hidden_dim": best_params["head_hidden_dim"], - "set_transformer_block": best_params.get("set_transformer_block", "sab"), "dropout": best_params["dropout"], "use_mpnn": use_mpnn, "use_chemeleon": chemeleon_cache is not None, @@ -890,16 +897,19 @@ def _run_single_outer_fold( "use_mole": mole_cache is not None, "mole_cache": mole_cache, "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, - **(llm_kwargs or {}), + "use_retrieval": use_retrieval, + "retr_feature_dim": (3 if use_retrieval else 0), + **arch, + **llm_kwargs_resolved, } _full_state = train_result["final_state"] - _slim_state = {k: v for k, v in _full_state.items() - if not k.startswith("llm_prompt.encoder")} + _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({ "model_state_dict": _slim_state, "config": config, @@ -981,6 +991,7 @@ def main( use_llm: bool = False, use_rag: bool = False, rag_top_k: int = 4, + llm_max_length: int = 1536, use_retrieval: bool = False, retrieval_source: str = "internal", llm_model_path: str = DEFAULT_MOLT5_PATH, @@ -991,6 +1002,8 @@ def main( llm_lora_r: int = 8, llm_lora_alpha: int = 16, llm_lora_dropout: float = 0.05, + fix_hparams_json: Optional[Path] = None, + fix_epoch_mean: int = 12, # 并行 parallel: bool = False, # 设备 @@ -1017,6 +1030,7 @@ def main( 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, @@ -1027,8 +1041,18 @@ def main( llm_lora_alpha=llm_lora_alpha, 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_config = None if init_from_pretrain is not None: @@ -1113,6 +1137,8 @@ def main( moe_jitter_noise=moe_jitter_noise, set_transformer_block=set_transformer_block, llm_kwargs=llm_kwargs, + precomputed_best_params=_fixed_bp, + precomputed_epoch_mean=_fixed_em, )) if parallel: diff --git a/lnp_ml/modeling/predict.py b/lnp_ml/modeling/predict.py index 4b02f44..1b8cc3c 100644 --- a/lnp_ml/modeling/predict.py +++ b/lnp_ml/modeling/predict.py @@ -35,65 +35,86 @@ def load_model( ) -> Union[LNPModel, LNPModelWithoutMPNN]: """ 加载训练好的模型。 - - 自动根据 checkpoint 的 config.use_mpnn 选择模型类型。 + + 根据 checkpoint 的 config 决定模型类型与全部结构开关 + (MoE / LLM / RAG / soft-prompt / reg_bypass)。 """ checkpoint = torch.load(model_path, map_location=device, weights_only=False) config = checkpoint["config"] 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: # 总是自动查找 MPNN ensemble,避免使用 checkpoint 中的旧绝对路径(可能来自其他机器) logger.info("Model was trained with MPNN, auto-detecting ensemble...") ensemble_paths = find_mpnn_ensemble_paths() logger.info(f"Found {len(ensemble_paths)} MPNN models") - 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_device=mpnn_device, - chemeleon_cache_path=config.get("chemeleon_cache"), - unimol_cache_path=config.get("unimol_cache"), + **common_kwargs, ) else: - model = LNPModelWithoutMPNN( - 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), - chemeleon_cache_path=config.get("chemeleon_cache"), - unimol_cache_path=config.get("unimol_cache"), + model = LNPModelWithoutMPNN(**common_kwargs) + + # 兼容旧 checkpoint:fusion 门从标量升级为逐维向量后,需把 0 维张量展开 + _sd = checkpoint["model_state_dict"] + for _k in ("fusion.g_moe", "fusion.g_llm", "fusion.g_retr"): + if _k in _sd and _sd[_k].dim() == 0: + _sd[_k] = _sd[_k].reshape(1).expand(model.fusion.d_model).clone() + + missing, unexpected = model.load_state_dict( + checkpoint["model_state_dict"], strict=False + ) + # strict=False 不会因结构不匹配报错,这里手动兜底 + if unexpected: + raise RuntimeError( + f"checkpoint 中有 {len(unexpected)} 个权重找不到对应模块," + f"结构开关可能未对齐: {unexpected[:5]}" ) - - model.load_state_dict(checkpoint["model_state_dict"], strict=False) + if missing: + logger.warning( + f"{len(missing)} 个参数未从 checkpoint 恢复,将使用随机初始化: {missing[:5]}" + ) + model.to(device) model.eval() - + logger.info(f"Loaded model from {model_path}") logger.info(f"Model config: {config}") logger.info(f"Best val_loss: {checkpoint.get('best_val_loss', 'N/A')}") - + return model @@ -152,7 +173,7 @@ def predictions_to_dataframe(predictions: Dict) -> pd.DataFrame: }) # 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]) # EE 类别映射 @@ -311,10 +332,10 @@ def test( "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"] 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 if mask.any(): y_true = pdi_true[mask] diff --git a/lnp_ml/modeling/trainer_balanced.py b/lnp_ml/modeling/trainer_balanced.py index a54ef58..2981530 100644 --- a/lnp_ml/modeling/trainer_balanced.py +++ b/lnp_ml/modeling/trainer_balanced.py @@ -1,7 +1,7 @@ """带类权重的训练器:处理分类任务的数据不均衡问题""" from typing import Dict, List, Optional, Tuple -from dataclasses import dataclass, field +from dataclasses import dataclass, field, replace import numpy as np import torch @@ -14,7 +14,7 @@ from tqdm import tqdm @dataclass 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 toxic: Optional[torch.Tensor] = None # [2] for binary toxic @@ -29,6 +29,10 @@ class LossWeightsBalanced: biodist: float = 1.0 toxic: float = 1.0 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( @@ -188,6 +192,16 @@ def compute_multitask_loss_balanced( losses["moe_lb"] = extras["lb_loss"] 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 @@ -202,7 +216,8 @@ def train_epoch_balanced( """带类权重的训练一个 epoch""" model.train() 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 for batch in tqdm(loader, desc="Training", leave=False): @@ -394,10 +409,14 @@ def train_with_early_stopping( task_weights: Optional[LossWeightsBalanced] = None, class_weights: Optional[ClassWeights] = None, backbone_lr_ratio: float = 1.0, + freeze_backbone_epochs: int = 0, ) -> Dict: """ 带早停的完整训练流程。 - + + freeze_backbone_epochs 必须与 train_fixed_epochs 取同一个值:best_epoch 是在 + 这个调度下测出来的,最终训练换了调度这个数就不可迁移。 + Returns: Dict with keys: history, best_val_loss, best_epoch, best_state """ @@ -407,12 +426,28 @@ def train_with_early_stopping( optimizer, mode="min", factor=0.5, patience=5 ) 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": []} best_val_loss = float("inf") best_state = None - + 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_metrics = train_epoch_balanced( model, train_loader, optimizer, device, task_weights, class_weights @@ -515,10 +550,14 @@ def train_fixed_epochs( swa_model = AveragedModel(model) swa_start = swa_start_epoch or int(epochs * 0.75) 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": []} - + 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) diff --git a/models/final/best_params.json b/models/final/best_params.json index 7c28cbd..df17e17 100644 --- a/models/final/best_params.json +++ b/models/final/best_params.json @@ -1,11 +1,16 @@ { - "dropout": 0.14731648771233646, - "lr": 0.00010399668603456206, - "weight_decay": 0.0007696017715340258, - "backbone_lr_ratio": 0.7472270478607838, + "dropout": 0.19323529834582587, + "lr": 0.000980484886752915, + "weight_decay": 0.00014367509113739864, + "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, "num_heads": 8, "n_attn_layers": 4, "fusion_strategy": "attention", - "head_hidden_dim": 128 + "head_hidden_dim": 128, + "set_transformer_block": "sab" } \ No newline at end of file diff --git a/models/final/class_weights.json b/models/final/class_weights.json index 73cc564..82420f6 100644 --- a/models/final/class_weights.json +++ b/models/final/class_weights.json @@ -1,9 +1,7 @@ { "pdi": [ - 0.011065198108553886, - 0.03264812380075455, - 0.8369068503379822, - 3.119379997253418 + 0.524036169052124, + 1.475963830947876 ], "ee": [ 1.600011944770813, diff --git a/models/final/epoch_mean.json b/models/final/epoch_mean.json index 65f6bc1..735a6b9 100644 --- a/models/final/epoch_mean.json +++ b/models/final/epoch_mean.json @@ -1 +1 @@ -{"epoch_mean": 23} \ No newline at end of file +{"epoch_mean": 11} \ No newline at end of file diff --git a/models/final/history.json b/models/final/history.json index d9eadef..2d1eabb 100644 --- a/models/final/history.json +++ b/models/final/history.json @@ -1,211 +1,126 @@ { "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": 5.662022154286222, + "loss_size": 0.9464401806581695, + "loss_pdi": 0.6576592483610477, + "loss_ee": 1.0441054870497506, + "loss_delivery": 1.0691580114499577, + "loss_biodist": 0.7718977461446006, + "loss_toxic": 0.47160032813279135, + "loss_moe_lb": 1.9999999887538407, + "loss_aux_moe": 1.096376850709038, + "loss_aux_llm": 1.2355621509113401 }, { - "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.549718969273117, + "loss_size": 0.9265210998227011, + "loss_pdi": 0.6320651516599475, + "loss_ee": 0.9564013177493833, + "loss_delivery": 0.8463795804682205, + "loss_biodist": 0.48631338606465535, + "loss_toxic": 0.25031149114991696, + "loss_moe_lb": 1.9999999910030726, + "loss_aux_moe": 0.8730144033928946, + "loss_aux_llm": 1.0696029401612732 }, { - "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.183341939494295, + "loss_size": 0.8995583306927726, + "loss_pdi": 0.6350391439671786, + "loss_ee": 0.9038635166186206, + "loss_delivery": 0.8323076782080362, + "loss_biodist": 0.44391101247297143, + "loss_toxic": 0.1671132598740031, + "loss_moe_lb": 1.9999999865046088, + "loss_aux_moe": 0.8456886384003567, + "loss_aux_llm": 1.0173823231796049 }, { - "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.9456988685535936, + "loss_size": 0.925582969104344, + "loss_pdi": 0.586557297211773, + "loss_ee": 0.867810919599713, + "loss_delivery": 0.7656616399521535, + "loss_biodist": 0.4162127935099152, + "loss_toxic": 0.2137195658123226, + "loss_moe_lb": 1.9999999887538407, + "loss_aux_moe": 0.7506198146784643, + "loss_aux_llm": 0.9585366686981804 }, { - "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.4442428957741216, + "loss_size": 0.8259714532573268, + "loss_pdi": 0.5488103347004585, + "loss_ee": 0.8086997242468708, + "loss_delivery": 0.7543281304808158, + "loss_biodist": 0.3298270061330975, + "loss_toxic": 0.09599270570566351, + "loss_moe_lb": 1.9999999955015362, + "loss_aux_moe": 0.8137425728282839, + "loss_aux_llm": 1.044384686873769 }, { - "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": 3.108477974837681, + "loss_size": 0.7671496489981435, + "loss_pdi": 0.5366503028374798, + "loss_ee": 0.7869719379353073, + "loss_delivery": 0.685029683248052, + "loss_biodist": 0.2828179093183212, + "loss_toxic": 0.08221106674908749, + "loss_moe_lb": 1.9999999910030726, + "loss_aux_moe": 0.7168521201413758, + "loss_aux_llm": 0.9446638939234445 }, { - "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.983976681277437, + "loss_size": 0.7495474740159962, + "loss_pdi": 0.5227343659355955, + "loss_ee": 0.7467741780685928, + "loss_delivery": 0.6811002301368511, + "loss_biodist": 0.24798926155803339, + "loss_toxic": 0.11302385037876929, + "loss_moe_lb": 1.9999999955015362 }, { - "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": 3.014890724757932, + "loss_size": 0.7314784642098084, + "loss_pdi": 0.48637263083233023, + "loss_ee": 0.769600696158859, + "loss_delivery": 0.8273954027912246, + "loss_biodist": 0.21140338072799286, + "loss_toxic": 0.08676587497249338, + "loss_moe_lb": 1.9999999955015362 }, { - "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.774192355713754, + "loss_size": 0.6206410539881239, + "loss_pdi": 0.4826473706173447, + "loss_ee": 0.7343562110415045, + "loss_delivery": 0.6946425725758638, + "loss_biodist": 0.19993741290186937, + "loss_toxic": 0.1272664367724298, + "loss_moe_lb": 1.9999999932523043 }, { - "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.683339330385316, + "loss_size": 0.5581827477885867, + "loss_pdi": 0.4963476042140205, + "loss_ee": 0.7145765797709519, + "loss_delivery": 0.6466561276817097, + "loss_biodist": 0.1851363978436533, + "loss_toxic": 0.1616253986507698, + "loss_moe_lb": 1.9999999932523043 }, { - "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 + "loss": 2.5718357652988075, + "loss_size": 0.6112437141391466, + "loss_pdi": 0.4928069238392812, + "loss_ee": 0.7109674007262824, + "loss_delivery": 0.6155119884829476, + "loss_biodist": 0.18839857628885306, + "loss_toxic": 0.05047845465344166, + "loss_moe_lb": 1.9999999842553768 } ], "val": [] diff --git a/models/final/loss_curves.png b/models/final/loss_curves.png index bd8f556..59f3b98 100644 Binary files a/models/final/loss_curves.png and b/models/final/loss_curves.png differ diff --git a/models/final/model.pt b/models/final/model.pt index e8c5230..9ea1e33 100644 --- a/models/final/model.pt +++ b/models/final/model.pt @@ -1,3 +1,3 @@ version https://git-lfs.github.com/spec/v1 -oid sha256:241b22c5170dd7d5a82f409228589d61d16ca14c3ae1991a03bda8ca547ce743 -size 28583340 +oid sha256:b96d3867be53a3e23abdf291baeec4a9d610da58dc790220a345e7012187aaa0 +size 53008646 diff --git a/models/final/optuna_study.sqlite3 b/models/final/optuna_study.sqlite3 index 262a85a..9748533 100644 Binary files a/models/final/optuna_study.sqlite3 and b/models/final/optuna_study.sqlite3 differ diff --git a/models/final/optuna_trials.json b/models/final/optuna_trials.json index f99c776..f96b9f0 100644 --- a/models/final/optuna_trials.json +++ b/models/final/optuna_trials.json @@ -1,960 +1,560 @@ [ { "number": 0, - "value": 2.8156970342000327, + "value": 2.9030250906944275, "params": { "dropout": 0.249816047538945, "lr": 0.0007969454818643932, "weight_decay": 0.008471801418819975, - "backbone_lr_ratio": 0.15751320499779725 + "backbone_lr_ratio": 0.15751320499779725, + "moe_n_experts": 2, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 16 }, "user_attrs": { - "epoch_mean": 14, + "epoch_mean": 10, "fold_best_epochs": [ - 10, - 12, - 21 + 7, + 11, + 13 ], "fold_val_losses": [ - 3.733495855331421, - 2.5024956703186034, - 2.2110995769500734 + 3.346657862265905, + 2.727404405673345, + 2.635013004144033 ] }, "state": "TrialState.COMPLETE" }, { "number": 1, - "value": 3.3739805857340492, + "value": 4.1861379808849755, "params": { - "dropout": 0.1624074561769746, - "lr": 2.0511104188433963e-05, - "weight_decay": 1.7073967431528103e-05, - "backbone_lr_ratio": 0.5399484409787431 + "dropout": 0.18493564427131048, + "lr": 2.3102018878452926e-05, + "weight_decay": 5.415244119402538e-05, + "backbone_lr_ratio": 0.04059611610484305, + "moe_n_experts": 2, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 32 }, "user_attrs": { - "epoch_mean": 30, + "epoch_mean": 20, "fold_best_epochs": [ - 30, - 30, - 30 + 20, + 20, + 20 ], "fold_val_losses": [ - 3.7005380153656007, - 3.236070442199707, - 3.1853332996368406 + 4.24272620677948, + 4.012106723255581, + 4.303581012619866 ] }, "state": "TrialState.COMPLETE" }, { "number": 2, - "value": 2.6023060957590736, + "value": 3.5133936824621976, "params": { - "dropout": 0.34044600469728353, - "lr": 0.0002607024758370766, - "weight_decay": 1.2087541473056957e-05, - "backbone_lr_ratio": 0.8706020878304853 + "dropout": 0.1798695128633439, + "lr": 0.00010677482709481354, + "weight_decay": 0.00234238498471129, + "backbone_lr_ratio": 0.012385137298860933, + "moe_n_experts": 2, + "moe_top_k": 2, + "moe_expert_hidden_mult": 1, + "llm_lora_r": 32 }, "user_attrs": { - "epoch_mean": 16, + "epoch_mean": 19, "fold_best_epochs": [ - 14, - 19, - 16 + 18, + 20, + 20 ], "fold_val_losses": [ - 3.283898639678955, - 2.366805338859558, - 2.1562143087387087 + 3.6705237362119885, + 3.35566207435396, + 3.513995236820645 ] }, "state": "TrialState.COMPLETE" }, { "number": 3, - "value": 10.838356399536133, + "value": 4.693677743275961, "params": { - "dropout": 0.4329770563201687, - "lr": 2.6587543983272695e-05, - "weight_decay": 5.3370327626039544e-05, - "backbone_lr_ratio": 0.023270677083837805 + "dropout": 0.27606099749584057, + "lr": 1.7541893487450798e-05, + "weight_decay": 0.0009565499215943821, + "backbone_lr_ratio": 0.011715937392307063, + "moe_n_experts": 2, + "moe_top_k": 1, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 16 }, "user_attrs": { - "epoch_mean": 30, + "epoch_mean": 20, "fold_best_epochs": [ - 30, - 30, - 30 + 20, + 20, + 20 ], "fold_val_losses": [ - 9.760257911682128, - 10.611498069763183, - 12.143313217163087 + 4.833255463176304, + 4.549637953440349, + 4.698139813211229 ] }, "state": "TrialState.COMPLETE" }, { "number": 4, - "value": 3.0307216803232833, + "value": 3.2253691399538957, "params": { - "dropout": 0.2216968971838151, - "lr": 0.00011207606211860574, - "weight_decay": 0.0005342937261279777, - "backbone_lr_ratio": 0.038234752246751866 + "dropout": 0.4757995766256756, + "lr": 0.0006161049539380963, + "weight_decay": 0.002463768595899745, + "backbone_lr_ratio": 0.697828126512603, + "moe_n_experts": 4, + "moe_top_k": 1, + "moe_expert_hidden_mult": 1, + "llm_lora_r": 8 }, "user_attrs": { - "epoch_mean": 30, + "epoch_mean": 13, "fold_best_epochs": [ - 29, - 30, - 30 + 7, + 13, + 19 ], "fold_val_losses": [ - 3.6398903369903564, - 2.7796840190887453, - 2.6725906848907472 + 3.5457335942321353, + 3.1522699031564922, + 2.9781039224730597 ] }, "state": "TrialState.COMPLETE" }, { "number": 5, - "value": 12.962510363260904, + "value": 4.624290854842574, "params": { - "dropout": 0.34474115788895177, - "lr": 1.9010245319870364e-05, - "weight_decay": 0.00014742753159914678, - "backbone_lr_ratio": 0.05404103854647329 + "dropout": 0.31707843326329943, + "lr": 1.913588048769229e-05, + "weight_decay": 0.016172900811143146, + "backbone_lr_ratio": 0.01409617514981587, + "moe_n_experts": 2, + "moe_top_k": 1, + "moe_expert_hidden_mult": 1, + "llm_lora_r": 16 }, "user_attrs": { - "epoch_mean": 30, + "epoch_mean": 20, "fold_best_epochs": [ - 30, - 30, - 30 + 20, + 20, + 20 ], "fold_val_losses": [ - 11.658656311035156, - 13.798492240905762, - 13.430382537841798 + 4.722254474957784, + 4.507743954658508, + 4.642874134911431 ] }, "state": "TrialState.COMPLETE" }, { "number": 6, - "value": 2.8127863725026447, + "value": 4.131107242019088, "params": { - "dropout": 0.28242799368681437, - "lr": 0.00037183641805732076, - "weight_decay": 6.290644294586152e-05, - "backbone_lr_ratio": 0.10677482709481352 + "dropout": 0.24338629141770907, + "lr": 1.70505392602693e-05, + "weight_decay": 0.028340904295147733, + "backbone_lr_ratio": 0.17643967683381545, + "moe_n_experts": 2, + "moe_top_k": 1, + "moe_expert_hidden_mult": 1, + "llm_lora_r": 8 }, "user_attrs": { - "epoch_mean": 21, + "epoch_mean": 20, "fold_best_epochs": [ - 15, 20, - 28 + 20, + 20 ], "fold_val_losses": [ - 3.527592992782593, - 2.6037728786468506, - 2.306993246078491 + 4.218653745121426, + 3.9478391938739352, + 4.226828787061903 ] }, "state": "TrialState.COMPLETE" }, { "number": 7, - "value": 18.49700838724772, + "value": 3.0308966758074583, "params": { - "dropout": 0.33696582754481696, - "lr": 1.2385137298860926e-05, - "weight_decay": 0.0026926469100861782, - "backbone_lr_ratio": 0.021930485556643693 + "dropout": 0.38529791488919807, + "lr": 0.0003323304206226791, + "weight_decay": 0.0017583640270008513, + "backbone_lr_ratio": 0.3482846706526883, + "moe_n_experts": 4, + "moe_top_k": 1, + "moe_expert_hidden_mult": 1, + "llm_lora_r": 8 }, "user_attrs": { - "epoch_mean": 30, + "epoch_mean": 14, "fold_best_epochs": [ - 30, - 30, - 30 + 8, + 15, + 19 ], "fold_val_losses": [ - 18.425870513916017, - 18.542006301879884, - 18.523148345947266 + 3.5245150327682495, + 2.85073944595125, + 2.7174355487028756 ] }, "state": "TrialState.COMPLETE" }, { "number": 8, - "value": 2.6576756556828816, + "value": 3.889537049664391, "params": { - "dropout": 0.12602063719411183, - "lr": 0.000790261954970823, - "weight_decay": 0.07286653737491042, - "backbone_lr_ratio": 0.4138040112561014 + "dropout": 0.4630265895704372, + "lr": 3.151987295193886e-05, + "weight_decay": 0.00043805807679056546, + "backbone_lr_ratio": 0.3244160088734159, + "moe_n_experts": 8, + "moe_top_k": 1, + "moe_expert_hidden_mult": 1, + "llm_lora_r": 16 }, "user_attrs": { - "epoch_mean": 10, + "epoch_mean": 20, "fold_best_epochs": [ - 9, - 11, - 10 + 20, + 20, + 20 ], "fold_val_losses": [ - 3.56644024848938, - 2.2665807485580443, - 2.1400059700012206 + 4.002199537224239, + 3.755847387843662, + 3.9105642239252725 ] }, "state": "TrialState.COMPLETE" }, { "number": 9, - "value": 11.92358964284261, + "value": 2.9054288014217664, "params": { - "dropout": 0.2218455076693483, - "lr": 1.5679933916722995e-05, - "weight_decay": 0.005456725485601475, - "backbone_lr_ratio": 0.07591104805282695 + "dropout": 0.17462802355441434, + "lr": 0.0006097025297491432, + "weight_decay": 0.0014367095138664223, + "backbone_lr_ratio": 0.4119839624605187, + "moe_n_experts": 2, + "moe_top_k": 1, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 8 }, "user_attrs": { - "epoch_mean": 30, + "epoch_mean": 13, "fold_best_epochs": [ - 30, - 30, - 30 + 10, + 11, + 17 ], "fold_val_losses": [ - 12.178006362915038, - 12.020879745483398, - 11.571882820129394 + 3.114242351717419, + 2.8655375705824957, + 2.736506481965383 ] }, "state": "TrialState.COMPLETE" }, { "number": 10, - "value": 2.744228076934814, + "value": 3.2267529169718423, "params": { - "dropout": 0.48159581757804304, + "dropout": 0.11691600370687807, "lr": 0.00013112992007873722, - "weight_decay": 1.0719864040714885e-05, - "backbone_lr_ratio": 0.8131735182005935 + "weight_decay": 0.07553503645583189, + "backbone_lr_ratio": 0.05739049944367051, + "moe_n_experts": 8, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 16 }, "user_attrs": { - "epoch_mean": 25, + "epoch_mean": 20, "fold_best_epochs": [ - 21, - 29, - 24 + 20, + 20, + 19 ], "fold_val_losses": [ - 3.2794203758239746, - 2.4811975240707396, - 2.472066330909729 + 3.380473878648546, + 3.0783044166035123, + 3.2214804556634693 ] }, "state": "TrialState.COMPLETE" }, { "number": 11, - "value": 2.600485173861186, + "value": 2.8969200551509857, "params": { - "dropout": 0.1398149921035961, - "lr": 0.0003676766469931975, - "weight_decay": 0.0814481577934799, - "backbone_lr_ratio": 0.416600060135216 + "dropout": 0.19323529834582587, + "lr": 0.000980484886752915, + "weight_decay": 0.00014367509113739864, + "backbone_lr_ratio": 0.13093310334961422, + "moe_n_experts": 2, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 8 }, "user_attrs": { - "epoch_mean": 14, + "epoch_mean": 11, "fold_best_epochs": [ - 8, - 14, - 21 + 11, + 6, + 17 ], "fold_val_losses": [ - 3.2930724143981935, - 2.312326765060425, - 2.196056342124939 + 3.2465109096633062, + 2.925493541691038, + 2.5187557140986123 ] }, "state": "TrialState.COMPLETE" }, { "number": 12, - "value": 2.856875479221344, + "value": 2.9026516896707037, "params": { - "dropout": 0.39786435978029433, - "lr": 0.00027104165213579705, - "weight_decay": 0.09324179801968951, - "backbone_lr_ratio": 0.24792013326598586 + "dropout": 0.34287370961329033, + "lr": 0.0009954383775681419, + "weight_decay": 1.2791436634859164e-05, + "backbone_lr_ratio": 0.11798172454857601, + "moe_n_experts": 2, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 8 }, "user_attrs": { - "epoch_mean": 17, + "epoch_mean": 14, "fold_best_epochs": [ - 14, - 12, - 25 + 8, + 20, + 15 ], "fold_val_losses": [ - 3.5214386940002442, - 2.671292519569397, - 2.3778952240943907 + 3.2743234402603574, + 2.9065700074036918, + 2.527061621348063 ] }, "state": "TrialState.COMPLETE" }, { "number": 13, - "value": 2.5929951747258504, + "value": 3.0620501229056605, "params": { - "dropout": 0.1708198460450869, - "lr": 0.0002771027430797957, - "weight_decay": 0.022371906875975713, - "backbone_lr_ratio": 0.8529032090320552 + "dropout": 0.35582764453212956, + "lr": 0.00025048271616507544, + "weight_decay": 1.0326835200348013e-05, + "backbone_lr_ratio": 0.08476798466510083, + "moe_n_experts": 8, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 8 }, "user_attrs": { - "epoch_mean": 11, + "epoch_mean": 19, "fold_best_epochs": [ - 9, - 12, - 13 + 16, + 20, + 20 ], "fold_val_losses": [ - 3.211256170272827, - 2.29467031955719, - 2.273059034347534 + 3.3088861074712543, + 2.8279807567596436, + 3.049283504486084 ] }, "state": "TrialState.COMPLETE" }, { "number": 14, - "value": 2.7000243584314982, + "value": 3.0037260243186243, "params": { - "dropout": 0.10319674017011438, - "lr": 5.318587054778435e-05, - "weight_decay": 0.026386070022612545, - "backbone_lr_ratio": 0.3335324680299515 + "dropout": 0.37640619632798594, + "lr": 0.000983992977213618, + "weight_decay": 8.560558113400214e-05, + "backbone_lr_ratio": 0.029140753854155613, + "moe_n_experts": 2, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 8 }, "user_attrs": { - "epoch_mean": 28, + "epoch_mean": 13, "fold_best_epochs": [ - 26, - 29, - 30 + 12, + 16, + 11 ], "fold_val_losses": [ - 3.095314884185791, - 2.559102272987366, - 2.445655918121338 + 3.3224475847350226, + 2.900958306259579, + 2.7877721819612713 ] }, "state": "TrialState.COMPLETE" }, { "number": 15, - "value": 2.811073144276937, + "value": 3.057953337828318, "params": { - "dropout": 0.17491537723907288, - "lr": 0.0004715550190685932, - "weight_decay": 0.020503728356529433, - "backbone_lr_ratio": 0.202005130829578 + "dropout": 0.4243012196347461, + "lr": 0.00031858611917698923, + "weight_decay": 1.457235563612646e-05, + "backbone_lr_ratio": 0.13443508445661562, + "moe_n_experts": 4, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 8 }, "user_attrs": { - "epoch_mean": 15, + "epoch_mean": 18, "fold_best_epochs": [ - 11, - 13, - 22 + 18, + 18, + 17 ], "fold_val_losses": [ - 3.665318727493286, - 2.5160216808319094, - 2.251879024505615 + 3.3081277741326227, + 2.9643009768591986, + 2.9014312624931335 ] }, "state": "TrialState.COMPLETE" }, { "number": 16, - "value": 2.51602889696757, + "value": 3.5674776236216226, "params": { - "dropout": 0.157877413230548, - "lr": 0.00016221844139316182, - "weight_decay": 0.0011695139073230128, - "backbone_lr_ratio": 0.9818226878189111 + "dropout": 0.10439567569750857, + "lr": 5.047116346834323e-05, + "weight_decay": 0.00012030457421620538, + "backbone_lr_ratio": 0.08403842125743101, + "moe_n_experts": 2, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 32 }, "user_attrs": { - "epoch_mean": 19, + "epoch_mean": 20, "fold_best_epochs": [ - 12, - 21, - 24 + 20, + 20, + 19 ], "fold_val_losses": [ - 2.967289590835571, - 2.41019024848938, - 2.170606851577759 + 3.61348729663425, + 3.4652820825576782, + 3.6236634916729398 ] }, "state": "TrialState.COMPLETE" }, { "number": 17, - "value": 2.617725872993469, + "value": 3.2637910324114343, "params": { - "dropout": 0.1975499843454273, - "lr": 6.0680145665816345e-05, - "weight_decay": 0.0006054421043842105, - "backbone_lr_ratio": 0.9370662382902515 + "dropout": 0.3121830815678499, + "lr": 0.00019221762216693134, + "weight_decay": 3.3001861372005635e-05, + "backbone_lr_ratio": 0.026623228883155377, + "moe_n_experts": 2, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 8 }, "user_attrs": { - "epoch_mean": 28, + "epoch_mean": 19, "fold_best_epochs": [ - 23, - 30, - 30 + 20, + 19, + 18 ], "fold_val_losses": [ - 3.112912178039551, - 2.4265938997268677, - 2.3136715412139894 + 3.35848887430297, + 3.204534967740377, + 3.228349255190955 ] }, "state": "TrialState.COMPLETE" }, { "number": 18, - "value": 2.6030575354894006, + "value": 2.932276170562815, "params": { - "dropout": 0.2840632229983807, - "lr": 0.00016641263642332344, - "weight_decay": 0.0015861751549934148, - "backbone_lr_ratio": 0.5813736146009224 + "dropout": 0.20651560054419343, + "lr": 0.0004218332710981938, + "weight_decay": 0.00022119130962064922, + "backbone_lr_ratio": 0.2167274117185541, + "moe_n_experts": 8, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 8 }, "user_attrs": { - "epoch_mean": 23, + "epoch_mean": 15, "fold_best_epochs": [ - 21, - 21, - 27 + 12, + 16, + 17 ], "fold_val_losses": [ - 3.163059639930725, - 2.3479116916656495, - 2.2982012748718263 + 3.367874321010378, + 2.7223513391282825, + 2.706602851549784 ] }, "state": "TrialState.COMPLETE" }, { "number": 19, - "value": 2.865672874450684, + "value": 3.2775669528378386, "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 + "dropout": 0.34109501751685645, + "lr": 5.26451532204125e-05, + "weight_decay": 1.8996419945400298e-05, + "backbone_lr_ratio": 0.6387255624731281, + "moe_n_experts": 4, + "moe_top_k": 2, + "moe_expert_hidden_mult": 2, + "llm_lora_r": 8 }, "user_attrs": { "epoch_mean": 20, "fold_best_epochs": [ - 14, - 19, - 28 + 20, + 20, + 20 ], "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 + 3.412749081850052, + 3.104751957787408, + 3.3151998188760547 ] }, "state": "TrialState.COMPLETE" diff --git a/models/final_backup_backbone/best_params.json b/models/final_backup_backbone/best_params.json new file mode 100644 index 0000000..7c28cbd --- /dev/null +++ b/models/final_backup_backbone/best_params.json @@ -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 +} \ No newline at end of file diff --git a/models/final_backup_backbone/class_weights.json b/models/final_backup_backbone/class_weights.json new file mode 100644 index 0000000..73cc564 --- /dev/null +++ b/models/final_backup_backbone/class_weights.json @@ -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 + ] +} \ No newline at end of file diff --git a/models/final_backup_backbone/epoch_mean.json b/models/final_backup_backbone/epoch_mean.json new file mode 100644 index 0000000..65f6bc1 --- /dev/null +++ b/models/final_backup_backbone/epoch_mean.json @@ -0,0 +1 @@ +{"epoch_mean": 23} \ No newline at end of file diff --git a/models/final_backup_backbone/history.json b/models/final_backup_backbone/history.json new file mode 100644 index 0000000..d9eadef --- /dev/null +++ b/models/final_backup_backbone/history.json @@ -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": [] +} \ No newline at end of file diff --git a/models/final_backup_backbone/loss_curves.png b/models/final_backup_backbone/loss_curves.png new file mode 100644 index 0000000..bd8f556 Binary files /dev/null and b/models/final_backup_backbone/loss_curves.png differ diff --git a/models/final_backup_backbone/model.pt b/models/final_backup_backbone/model.pt new file mode 100644 index 0000000..e8c5230 --- /dev/null +++ b/models/final_backup_backbone/model.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:241b22c5170dd7d5a82f409228589d61d16ca14c3ae1991a03bda8ca547ce743 +size 28583340 diff --git a/models/final_backup_backbone/optuna_trials.json b/models/final_backup_backbone/optuna_trials.json new file mode 100644 index 0000000..f99c776 --- /dev/null +++ b/models/final_backup_backbone/optuna_trials.json @@ -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" + } +] \ No newline at end of file diff --git a/models/final_backup_backbone/strata_info.json b/models/final_backup_backbone/strata_info.json new file mode 100644 index 0000000..ef3be58 --- /dev/null +++ b/models/final_backup_backbone/strata_info.json @@ -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" +} \ No newline at end of file diff --git a/models/pretrain/mpnn/pretrain.log b/models/pretrain/mpnn/pretrain.log new file mode 100644 index 0000000..7bffb80 --- /dev/null +++ b/models/pretrain/mpnn/pretrain.log @@ -0,0 +1,57 @@ +nohup: ignoring input +2026-08-07 06:43:53.431 | INFO | lnp_ml.config::11 - PROJ_ROOT path is: /home/gongruyi/lnp_ml +2026-08-07 06:43:53.749 | INFO  | __main__:main:273 - Using device: cuda | seed: 42 +2026-08-07 06:43:53.749 | INFO  | __main__:main:277 - Loading train data from /home/gongruyi/lnp_ml/data/processed/train_pretrain.parquet +2026-08-07 06:43:54.598 | INFO  | __main__:main:281 - Loading val data from /home/gongruyi/lnp_ml/data/processed/val_pretrain.parquet +2026-08-07 06:43:54.622 | INFO  | __main__:main:285 - Train samples: 8236, Val samples: 1454 +2026-08-07 06:43:54.623 | INFO  | __main__:main:301 - Auto-detecting MPNN ensemble from /home/gongruyi/lnp_ml/models/mpnn/all_amine_split_for_LiON +2026-08-07 06:43:54.624 | INFO  | __main__:main:303 - Found 5 MPNN models +2026-08-07 06:43:54.624 | INFO  | __main__:main:308 - Creating model (use_mpnn=True, use_moe=False, use_llm=False)... +2026-08-07 06:43:54.661 | INFO  | __main__:main:346 - Model parameters: 3,961,386 total, 3,961,386 trainable +2026-08-07 06:43:54.663 | INFO  | __main__:warmup_cache:64 - Warming up RDKit cache for 7493 unique SMILES... + Cache warmup: 0%| | 0/30 [00:00 New best model (val_loss=0.7837) + Epoch 2 [Train]: 0%| | 0/129 [00:00 New best model (val_loss=0.7153) + Epoch 3 [Train]: 0%| | 0/129 [00:00 New best model (val_loss=0.6935) + Epoch 5 [Train]: 0%| | 0/129 [00:00 New best model (val_loss=0.6741) + Epoch 6 [Train]: 0%| | 0/129 [00:00 New best model (val_loss=0.6682) + Epoch 11 [Train]: 0%| | 0/129 [00:00 New best model (val_loss=0.6577) + Epoch 18 [Train]: 0%| | 0/129 [00:00 New best model (val_loss=0.6451) + Epoch 22 [Train]: 0%| | 0/129 [00:00 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() \ No newline at end of file diff --git a/scripts_run/launch_final_cv.sh b/scripts_run/launch_final_cv.sh new file mode 100644 index 0000000..7936bad --- /dev/null +++ b/scripts_run/launch_final_cv.sh @@ -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" \ No newline at end of file diff --git a/scripts_run/run_final_cv.sh b/scripts_run/run_final_cv.sh new file mode 100644 index 0000000..f370e57 --- /dev/null +++ b/scripts_run/run_final_cv.sh @@ -0,0 +1,95 @@ +#!/usr/bin/env bash +# 完整 nested CV:MPNN + 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+LLM;off=接入 +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}MiB,60s 后重试" >>"${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 \ No newline at end of file diff --git a/scripts_run/run_final_full.sh b/scripts_run/run_final_full.sh new file mode 100644 index 0000000..a0f30de --- /dev/null +++ b/scripts_run/run_final_full.sh @@ -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}MiB,60s 后重试" >>"${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 \ No newline at end of file diff --git a/tests/check_gates.py b/tests/check_gates.py index c2cf8c8..04ea661 100644 --- a/tests/check_gates.py +++ b/tests/check_gates.py @@ -4,7 +4,8 @@ import os import numpy as np 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"] @@ -13,7 +14,7 @@ def latest(v: str): return dirs[-1] if dirs else None -for v in ["moe", "llm", "both"]: +for v in VARIANTS: run = latest(v) if run is None: print(f"\n=== {v}: 无 run 目录 ===") @@ -28,14 +29,15 @@ for v in ["moe", "llm", "both"]: gm, gl = [], [] for f in files: sd = torch.load(f, map_location="cpu", weights_only=False)["model_state_dict"] - if "fusion.g_moe" in sd: - gm.append(sd["fusion.g_moe"].item()) - if "fusion.g_llm" in sd: - gl.append(sd["fusion.g_llm"].item()) - if gm: - print(f" g_moe: mean={np.mean(gm):+.4f} |abs|=[{min(map(abs, gm)):.4f},{max(map(abs, gm)):.4f}]") - if gl: - print(f" g_llm: mean={np.mean(gl):+.4f} |abs|=[{min(map(abs, gl)):.4f},{max(map(abs, gl)):.4f}]") + for key, acc in (("fusion.g_moe", gm), ("fusion.g_llm", gl)): + if key not in sd: + continue + t = sd[key] + acc.append(abs(t.item()) if t.dim() == 0 + else float(t.norm()) / (t.numel() ** 0.5)) + for name, acc in (("g_moe", gm), ("g_llm", gl)): + 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"] for t in TASKS: