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

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

View File

@ -87,6 +87,7 @@ class OptimizeRequest(BaseModel):
top_k: int = Field(default=20, ge=1, le=100, description="Number of top formulations to return")
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))

View File

@ -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 分类标签

View File

@ -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:

View File

@ -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'}")
# 保存训练历史

View File

@ -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 分类

View File

@ -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:

View File

@ -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 期望格式化分子。
BioT5SMILES -> SELFIES <bom>...<eom> 紧贴包裹官方格式token 间无空格
@ -227,36 +268,54 @@ class LLMPromptEncoder(nn.Module):
return f"<bom>{sfs}<eom>"
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())
# ---------- 旧路径(保留,向后兼容)----------

View File

@ -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,

View File

@ -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:

View File

@ -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)
# 兼容旧 checkpointfusion 门从标量升级为逐维向量后,需把 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]

View File

@ -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)

View File

@ -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"
}

View File

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

View File

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

View File

@ -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": []

Binary file not shown.

View File

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

Binary file not shown.

File diff suppressed because it is too large Load Diff

View File

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

View File

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

View File

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

View File

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

Binary file not shown.

View File

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

View File

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

View File

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

File diff suppressed because one or more lines are too long

View File

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

View File

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

Binary file not shown.

View File

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

View File

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

89
scripts/measure_cost.py Normal file
View File

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

View File

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

View File

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

View File

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

View File

@ -4,7 +4,8 @@ import os
import numpy as np
import 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: