mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-09-18 14:23:20 +08:00
feat: final_train_optuna_cv 支持 MoE + Qwen2.5-7B QLoRA + soft-RAG,产出可部署 checkpoint
This commit is contained in:
parent
bab8203c98
commit
866941039f
18
app/api.py
18
app/api.py
@ -87,6 +87,7 @@ class OptimizeRequest(BaseModel):
|
||||
top_k: int = Field(default=20, ge=1, le=100, description="Number of top formulations to return")
|
||||
num_seeds: Optional[int] = Field(default=None, ge=1, le=500, description="Number of seed points from first iteration (default: top_k * 5)")
|
||||
top_per_seed: int = Field(default=1, ge=1, le=10, description="Number of local best to keep per seed in refinement")
|
||||
rerank_top_n: int = Field(default=0, ge=0, le=1000, description="两阶段推理:粗筛后用完整模型重排的候选数,0 表示关闭")
|
||||
step_sizes: Optional[List[float]] = Field(default=None, description="Mol ratio step sizes for each iteration (default: [10, 2, 1])")
|
||||
wr_step_sizes: Optional[List[float]] = Field(default=None, description="Weight ratio step sizes for each iteration (default: [5, 2, 1])")
|
||||
comp_ranges: Optional[CompRangesRequest] = Field(default=None, description="组分范围配置(默认使用标准范围)")
|
||||
@ -188,6 +189,20 @@ async def lifespan(app: FastAPI):
|
||||
try:
|
||||
state.model = load_model(model_path, state.device)
|
||||
logger.success("Model loaded successfully!")
|
||||
# RAG 检索池:服务期用全部内部数据(无泄漏顾虑,查询分子会被 _retrieve_topk 自动排除)
|
||||
_llm = getattr(state.model, "llm_prompt", None)
|
||||
if _llm is not None and getattr(_llm, "use_rag", False):
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from lnp_ml.dataset import LNPDataset, process_dataframe
|
||||
from lnp_ml.modeling.nested_cv_optuna import _build_rag_pool
|
||||
|
||||
rag_csv = Path(os.environ.get("RAG_POOL_CSV", "data/interim/internal.csv"))
|
||||
logger.info(f"Building RAG retrieval pool from {rag_csv}...")
|
||||
_ds = LNPDataset(process_dataframe(pd.read_csv(rag_csv)))
|
||||
_s, _d, _ex = _build_rag_pool(_ds, np.arange(len(_ds)))
|
||||
_llm.set_retrieval_pool(_s, _d, pool_id="serve", extra_labels=_ex)
|
||||
logger.success(f"RAG pool ready: {len(_s)} molecules")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load model: {e}")
|
||||
raise
|
||||
@ -305,6 +320,7 @@ async def optimize_formulation(request: OptimizeRequest):
|
||||
comp_ranges=comp_ranges,
|
||||
routes=request.routes,
|
||||
scoring_weights=scoring_weights,
|
||||
rerank_top_n=request.rerank_top_n,
|
||||
batch_size=256,
|
||||
)
|
||||
|
||||
@ -350,6 +366,8 @@ async def optimize_formulation(request: OptimizeRequest):
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Optimization failed: {e}")
|
||||
if getattr(state.model, "llm_prompt", None) is not None:
|
||||
state.model.set_llm_enabled(True)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
|
||||
@ -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 分类标签
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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'}")
|
||||
|
||||
# 保存训练历史
|
||||
|
||||
@ -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 分类
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -12,6 +12,15 @@ DEFAULT_MOLT5_PATH = os.environ.get("MOLT5_PATH", "models/molt5-base")
|
||||
# 检索池支持的额外多任务(除 delivery 外)
|
||||
_EXTRA_TASKS = ["size", "pdi", "ee", "toxic", "biodist"]
|
||||
|
||||
# 连续量的分位桶数。边界只从训练检索池统计,不引入测试集分布。
|
||||
_N_QBINS = 5
|
||||
|
||||
# 离散标签的语义,避免 LLM 只看到裸整数
|
||||
_PDI_LABELS = {0: "<0.2", 1: ">=0.2"}
|
||||
_EE_LABELS = {0: "<50%", 1: "50-80%", 2: ">80%"}
|
||||
_TOX_LABELS = {0: "non-toxic", 1: "toxic"}
|
||||
_ORGANS = ["lymph_nodes", "heart", "liver", "spleen", "lung", "kidney", "muscle"]
|
||||
|
||||
|
||||
class LLMPromptEncoder(nn.Module):
|
||||
"""用 LLM 编码分子,输出 F_llm [B, d_model]。
|
||||
@ -35,7 +44,7 @@ class LLMPromptEncoder(nn.Module):
|
||||
lora_r: int = 8,
|
||||
lora_alpha: int = 16,
|
||||
lora_dropout: float = 0.05,
|
||||
max_length: int = 256,
|
||||
max_length: int = 1536,
|
||||
use_rag: bool = False,
|
||||
rag_top_k: int = 4,
|
||||
use_soft_prompt: bool = False,
|
||||
@ -64,6 +73,9 @@ class LLMPromptEncoder(nn.Module):
|
||||
)
|
||||
if _is_qwen and self.tokenizer.pad_token is None:
|
||||
self.tokenizer.pad_token = self.tokenizer.eos_token
|
||||
if use_soft_prompt:
|
||||
# soft token 后置要求文本右对齐,否则 soft 会落在 PAD 之后
|
||||
self.tokenizer.padding_side = "left"
|
||||
|
||||
if _is_t5:
|
||||
self.encoder = T5EncoderModel.from_pretrained(model_name_or_path)
|
||||
@ -128,6 +140,7 @@ class LLMPromptEncoder(nn.Module):
|
||||
self._rag_pool_extra = None # dict: task -> (values, valid)
|
||||
self._rag_pool_fps = None
|
||||
self._rag_pool_id: str = "none"
|
||||
self._rag_pool_qedges: Dict[str, List[float]] = {}
|
||||
|
||||
def _apply_lora(self, r, alpha, dropout, prepare_kbit=False):
|
||||
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
|
||||
@ -167,6 +180,19 @@ class LLMPromptEncoder(nn.Module):
|
||||
self._rag_pool_labels = np.asarray(labels, dtype=np.float32).reshape(-1)
|
||||
self._rag_pool_extra = extra_labels
|
||||
self._rag_pool_fps = [_smiles_to_fp(s) for s in self._rag_pool_smiles]
|
||||
|
||||
self._rag_pool_qedges = {}
|
||||
_probs = np.linspace(0.0, 1.0, _N_QBINS + 1)[1:-1]
|
||||
_cont = {"delivery": self._rag_pool_labels}
|
||||
if extra_labels is not None and "size" in extra_labels:
|
||||
_v, _ok = extra_labels["size"]
|
||||
_cont["size"] = np.asarray(_v, dtype=np.float32)[np.asarray(_ok, dtype=bool)]
|
||||
for _name, _vals in _cont.items():
|
||||
_vals = np.asarray(_vals, dtype=np.float32)
|
||||
_vals = _vals[np.isfinite(_vals)]
|
||||
if _vals.size >= _N_QBINS:
|
||||
self._rag_pool_qedges[_name] = np.quantile(_vals, _probs).tolist()
|
||||
|
||||
if pool_id != self._rag_pool_id:
|
||||
self._cache = {k: v for k, v in self._cache.items() if not k.startswith("RAG::")}
|
||||
self._prompt_cache.clear()
|
||||
@ -213,6 +239,21 @@ class LLMPromptEncoder(nn.Module):
|
||||
return "[" + ", ".join(f"{x:.3f}" for x in v) + "]"
|
||||
return f"{float(v):.3f}"
|
||||
|
||||
def _qbin(self, name: str, v) -> str:
|
||||
"""v 落在训练池经验分布的第几桶,如 '(Q3/5)';无边界或缺失时返回空串。"""
|
||||
import bisect
|
||||
edges = self._rag_pool_qedges.get(name)
|
||||
if v is None or not edges:
|
||||
return ""
|
||||
return f"(Q{bisect.bisect_right(edges, float(v)) + 1}/{len(edges) + 1})"
|
||||
|
||||
@staticmethod
|
||||
def _fmt_class(v, labels: Dict[int, str]) -> str:
|
||||
"""离散标签附带语义,例如 '1(>=0.2)'。"""
|
||||
if v is None:
|
||||
return "unknown"
|
||||
return f"{int(v)}({labels.get(int(v), '?')})"
|
||||
|
||||
def _fmt_mol(self, smiles: str) -> str:
|
||||
"""按 backbone 期望格式化分子。
|
||||
BioT5:SMILES -> SELFIES,用 <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())
|
||||
|
||||
# ---------- 旧路径(保留,向后兼容)----------
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -35,65 +35,86 @@ def load_model(
|
||||
) -> Union[LNPModel, LNPModelWithoutMPNN]:
|
||||
"""
|
||||
加载训练好的模型。
|
||||
|
||||
自动根据 checkpoint 的 config.use_mpnn 选择模型类型。
|
||||
|
||||
根据 checkpoint 的 config 决定模型类型与全部结构开关
|
||||
(MoE / LLM / RAG / soft-prompt / reg_bypass)。
|
||||
"""
|
||||
checkpoint = torch.load(model_path, map_location=device, weights_only=False)
|
||||
config = checkpoint["config"]
|
||||
use_mpnn = config.get("use_mpnn", False)
|
||||
|
||||
|
||||
common_kwargs = dict(
|
||||
d_model=config["d_model"],
|
||||
num_heads=config["num_heads"],
|
||||
n_attn_layers=config["n_attn_layers"],
|
||||
set_transformer_block=config.get("set_transformer_block", "sab"),
|
||||
fusion_strategy=config["fusion_strategy"],
|
||||
head_hidden_dim=config["head_hidden_dim"],
|
||||
dropout=config["dropout"],
|
||||
reg_bypass=config.get("reg_bypass", "on"),
|
||||
use_moe=config.get("use_moe", False),
|
||||
moe_n_experts=config.get("moe_n_experts", 4),
|
||||
moe_top_k=config.get("moe_top_k", 2),
|
||||
moe_expert_hidden_mult=config.get("moe_expert_hidden_mult", 2),
|
||||
moe_jitter_noise=config.get("moe_jitter_noise", 0.0),
|
||||
use_llm=config.get("use_llm", False),
|
||||
llm_model_path=config.get("llm_model_path", "models/molt5-base"),
|
||||
llm_freeze=config.get("llm_freeze", True),
|
||||
llm_use_lora=config.get("llm_use_lora", False),
|
||||
llm_use_qlora=config.get("llm_use_qlora", False),
|
||||
use_soft_prompt=config.get("use_soft_prompt", False),
|
||||
llm_lora_r=config.get("llm_lora_r", 8),
|
||||
llm_lora_alpha=config.get("llm_lora_alpha", 16),
|
||||
llm_lora_dropout=config.get("llm_lora_dropout", 0.05),
|
||||
use_rag=config.get("use_rag", False),
|
||||
rag_top_k=config.get("rag_top_k", 4),
|
||||
llm_max_length=config.get("llm_max_length", 256),
|
||||
use_retrieval=config.get("use_retrieval", False),
|
||||
retr_feature_dim=config.get("retr_feature_dim", 0),
|
||||
chemeleon_cache_path=config.get("chemeleon_cache"),
|
||||
unimol_cache_path=config.get("unimol_cache"),
|
||||
)
|
||||
|
||||
if use_mpnn:
|
||||
# 总是自动查找 MPNN ensemble,避免使用 checkpoint 中的旧绝对路径(可能来自其他机器)
|
||||
logger.info("Model was trained with MPNN, auto-detecting ensemble...")
|
||||
ensemble_paths = find_mpnn_ensemble_paths()
|
||||
logger.info(f"Found {len(ensemble_paths)} MPNN models")
|
||||
|
||||
model = LNPModel(
|
||||
d_model=config["d_model"],
|
||||
num_heads=config["num_heads"],
|
||||
n_attn_layers=config["n_attn_layers"],
|
||||
fusion_strategy=config["fusion_strategy"],
|
||||
head_hidden_dim=config["head_hidden_dim"],
|
||||
dropout=config["dropout"],
|
||||
use_llm=config.get("use_llm", False),
|
||||
llm_model_path=config.get("llm_model_path", "models/molt5-base"),
|
||||
llm_freeze=config.get("llm_freeze", True),
|
||||
llm_use_lora=config.get("llm_use_lora", False),
|
||||
llm_lora_r=config.get("llm_lora_r", 8),
|
||||
llm_lora_alpha=config.get("llm_lora_alpha", 16),
|
||||
llm_lora_dropout=config.get("llm_lora_dropout", 0.05),
|
||||
mpnn_ensemble_paths=ensemble_paths,
|
||||
mpnn_device=mpnn_device,
|
||||
chemeleon_cache_path=config.get("chemeleon_cache"),
|
||||
unimol_cache_path=config.get("unimol_cache"),
|
||||
**common_kwargs,
|
||||
)
|
||||
else:
|
||||
model = LNPModelWithoutMPNN(
|
||||
d_model=config["d_model"],
|
||||
num_heads=config["num_heads"],
|
||||
n_attn_layers=config["n_attn_layers"],
|
||||
fusion_strategy=config["fusion_strategy"],
|
||||
head_hidden_dim=config["head_hidden_dim"],
|
||||
dropout=config["dropout"],
|
||||
use_llm=config.get("use_llm", False),
|
||||
llm_model_path=config.get("llm_model_path", "models/molt5-base"),
|
||||
llm_freeze=config.get("llm_freeze", True),
|
||||
llm_use_lora=config.get("llm_use_lora", False),
|
||||
llm_lora_r=config.get("llm_lora_r", 8),
|
||||
llm_lora_alpha=config.get("llm_lora_alpha", 16),
|
||||
llm_lora_dropout=config.get("llm_lora_dropout", 0.05),
|
||||
chemeleon_cache_path=config.get("chemeleon_cache"),
|
||||
unimol_cache_path=config.get("unimol_cache"),
|
||||
model = LNPModelWithoutMPNN(**common_kwargs)
|
||||
|
||||
# 兼容旧 checkpoint:fusion 门从标量升级为逐维向量后,需把 0 维张量展开
|
||||
_sd = checkpoint["model_state_dict"]
|
||||
for _k in ("fusion.g_moe", "fusion.g_llm", "fusion.g_retr"):
|
||||
if _k in _sd and _sd[_k].dim() == 0:
|
||||
_sd[_k] = _sd[_k].reshape(1).expand(model.fusion.d_model).clone()
|
||||
|
||||
missing, unexpected = model.load_state_dict(
|
||||
checkpoint["model_state_dict"], strict=False
|
||||
)
|
||||
# strict=False 不会因结构不匹配报错,这里手动兜底
|
||||
if unexpected:
|
||||
raise RuntimeError(
|
||||
f"checkpoint 中有 {len(unexpected)} 个权重找不到对应模块,"
|
||||
f"结构开关可能未对齐: {unexpected[:5]}"
|
||||
)
|
||||
|
||||
model.load_state_dict(checkpoint["model_state_dict"], strict=False)
|
||||
if missing:
|
||||
logger.warning(
|
||||
f"{len(missing)} 个参数未从 checkpoint 恢复,将使用随机初始化: {missing[:5]}"
|
||||
)
|
||||
|
||||
model.to(device)
|
||||
model.eval()
|
||||
|
||||
|
||||
logger.info(f"Loaded model from {model_path}")
|
||||
logger.info(f"Model config: {config}")
|
||||
logger.info(f"Best val_loss: {checkpoint.get('best_val_loss', 'N/A')}")
|
||||
|
||||
|
||||
return model
|
||||
|
||||
|
||||
@ -152,7 +173,7 @@ def predictions_to_dataframe(predictions: Dict) -> pd.DataFrame:
|
||||
})
|
||||
|
||||
# PDI 类别映射
|
||||
pdi_labels = ["0_0to0_2", "0_2to0_3", "0_3to0_4", "0_4to0_5"]
|
||||
pdi_labels = ["PDI<0.2", "PDI>=0.2"]
|
||||
df["pred_pdi_label"] = df["pred_pdi_class"].map(lambda x: pdi_labels[x])
|
||||
|
||||
# EE 类别映射
|
||||
@ -311,10 +332,10 @@ def test(
|
||||
"r2": float(r2_score(y_true, y_pred)),
|
||||
}
|
||||
|
||||
# 分类指标:PDI
|
||||
# 分类指标:PDI(真值需与 dataset.py 同样折叠为二分类)
|
||||
pdi_cols = ["PDI_0_0to0_2", "PDI_0_2to0_3", "PDI_0_3to0_4", "PDI_0_4to0_5"]
|
||||
if all(c in test_df.columns for c in pdi_cols):
|
||||
pdi_true = test_df[pdi_cols].values.argmax(axis=1)
|
||||
pdi_true = (test_df[pdi_cols].values.argmax(axis=1) >= 1).astype(np.int64)
|
||||
mask = test_df[pdi_cols].sum(axis=1) > 0
|
||||
if mask.any():
|
||||
y_true = pdi_true[mask]
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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"
|
||||
}
|
||||
@ -1,9 +1,7 @@
|
||||
{
|
||||
"pdi": [
|
||||
0.011065198108553886,
|
||||
0.03264812380075455,
|
||||
0.8369068503379822,
|
||||
3.119379997253418
|
||||
0.524036169052124,
|
||||
1.475963830947876
|
||||
],
|
||||
"ee": [
|
||||
1.600011944770813,
|
||||
|
||||
@ -1 +1 @@
|
||||
{"epoch_mean": 23}
|
||||
{"epoch_mean": 11}
|
||||
@ -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.
@ -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
11
models/final_backup_backbone/best_params.json
Normal file
11
models/final_backup_backbone/best_params.json
Normal 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
|
||||
}
|
||||
17
models/final_backup_backbone/class_weights.json
Normal file
17
models/final_backup_backbone/class_weights.json
Normal 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
|
||||
]
|
||||
}
|
||||
1
models/final_backup_backbone/epoch_mean.json
Normal file
1
models/final_backup_backbone/epoch_mean.json
Normal file
@ -0,0 +1 @@
|
||||
{"epoch_mean": 23}
|
||||
212
models/final_backup_backbone/history.json
Normal file
212
models/final_backup_backbone/history.json
Normal 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": []
|
||||
}
|
||||
BIN
models/final_backup_backbone/loss_curves.png
Normal file
BIN
models/final_backup_backbone/loss_curves.png
Normal file
Binary file not shown.
3
models/final_backup_backbone/model.pt
Normal file
3
models/final_backup_backbone/model.pt
Normal file
@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:241b22c5170dd7d5a82f409228589d61d16ca14c3ae1991a03bda8ca547ce743
|
||||
size 28583340
|
||||
962
models/final_backup_backbone/optuna_trials.json
Normal file
962
models/final_backup_backbone/optuna_trials.json
Normal 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"
|
||||
}
|
||||
]
|
||||
60
models/final_backup_backbone/strata_info.json
Normal file
60
models/final_backup_backbone/strata_info.json
Normal 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"
|
||||
}
|
||||
57
models/pretrain/mpnn/pretrain.log
Normal file
57
models/pretrain/mpnn/pretrain.log
Normal file
File diff suppressed because one or more lines are too long
3
models/pretrain/mpnn/pretrain_delivery.pt
Normal file
3
models/pretrain/mpnn/pretrain_delivery.pt
Normal file
@ -0,0 +1,3 @@
|
||||
version https://git-lfs.github.com/spec/v1
|
||||
oid sha256:0ea823b87a22a3a87cb487a9e876a6605c27a138abe53eaf228bb348dad43667
|
||||
size 15970234
|
||||
254
models/pretrain/mpnn/pretrain_history.json
Normal file
254
models/pretrain/mpnn/pretrain_history.json
Normal 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
|
||||
}
|
||||
]
|
||||
}
|
||||
BIN
models/pretrain/mpnn/pretrain_loss_curves.png
Normal file
BIN
models/pretrain/mpnn/pretrain_loss_curves.png
Normal file
Binary file not shown.
38
scripts/check_prompt_len.py
Normal file
38
scripts/check_prompt_len.py
Normal 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}")
|
||||
8
scripts/fixed_hparams.json
Normal file
8
scripts/fixed_hparams.json
Normal 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
89
scripts/measure_cost.py
Normal 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()
|
||||
21
scripts_run/launch_final_cv.sh
Normal file
21
scripts_run/launch_final_cv.sh
Normal 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"
|
||||
95
scripts_run/run_final_cv.sh
Normal file
95
scripts_run/run_final_cv.sh
Normal file
@ -0,0 +1,95 @@
|
||||
#!/usr/bin/env bash
|
||||
# 完整 nested CV:MPNN + MoE + Qwen2.5-7B QLoRA + soft-RAG
|
||||
# 单卡一个分片。被 OOM / 抢占 / 断连杀掉后自动等显存并续跑。
|
||||
set -uo pipefail
|
||||
|
||||
GPU=${GPU:?必须指定,例如 GPU=0}
|
||||
SEED=${SEED:-42}
|
||||
FOLDS=${FOLDS:-} # 空=跑全部 5 折;"0,1,2"=只跑这三折
|
||||
REG_BYPASS=${REG_BYPASS:-on} # on=delivery/size 绕开 MoE+LLM;off=接入
|
||||
BATCH=${BATCH:-8} # 显存紧就设 4
|
||||
N_OUTER=${N_OUTER:-5}
|
||||
N_INNER=${N_INNER:-3}
|
||||
N_TRIALS=${N_TRIALS:-20}
|
||||
EPOCHS=${EPOCHS:-20}
|
||||
PATIENCE=${PATIENCE:-5}
|
||||
MIN_FREE_MB=${MIN_FREE_MB:-14000}
|
||||
MAX_RETRY=${MAX_RETRY:-50}
|
||||
RETRY_WAIT=${RETRY_WAIT:-300}
|
||||
MIN_OK_SEC=${MIN_OK_SEC:-180} # 存活不足这么久就挂 → 判为配置/代码错误,立即停止重试
|
||||
PRETRAIN=${PRETRAIN:-} # 非空且文件存在才注入 --init-from-pretrain
|
||||
|
||||
RUN_DIR="models/final_cv/${REG_BYPASS}_seed${SEED}"
|
||||
TAG="gpu${GPU}$([ -n "${FOLDS}" ] && echo "_folds$(echo "${FOLDS}" | tr -d ',')")"
|
||||
LOG="${RUN_DIR}/${TAG}.log"
|
||||
|
||||
export TRANSFORMERS_OFFLINE=1
|
||||
export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export CUDA_VISIBLE_DEVICES=${GPU}
|
||||
|
||||
mkdir -p "${RUN_DIR}"
|
||||
|
||||
# nvidia-smi 不受 CUDA_VISIBLE_DEVICES 影响,-i 用物理卡号
|
||||
wait_free() {
|
||||
while :; do
|
||||
local free
|
||||
free=$(nvidia-smi --query-gpu=memory.free --format=csv,noheader,nounits -i "${GPU}")
|
||||
[ "${free}" -ge "${MIN_FREE_MB}" ] && return 0
|
||||
echo "[$(date '+%F %T')] GPU${GPU} 仅空 ${free}MiB < ${MIN_FREE_MB}MiB,60s 后重试" >>"${LOG}"
|
||||
sleep 60
|
||||
done
|
||||
}
|
||||
|
||||
EXTRA=()
|
||||
if [ -n "${FOLDS}" ]; then
|
||||
# 注意:需先给 nested_cv_optuna.main 增加 --only-folds 选项,否则 Typer 直接报错
|
||||
EXTRA+=(--only-folds "${FOLDS}")
|
||||
fi
|
||||
if [ -n "${PRETRAIN}" ] && [ -f "${PRETRAIN}" ]; then
|
||||
EXTRA+=(--init-from-pretrain "${PRETRAIN}")
|
||||
elif [ -n "${PRETRAIN}" ]; then
|
||||
echo "[$(date '+%F %T')] 警告:${PRETRAIN} 不存在,跳过预训练初始化" >>"${LOG}"
|
||||
fi
|
||||
|
||||
for attempt in $(seq 1 "${MAX_RETRY}"); do
|
||||
wait_free
|
||||
echo "[$(date '+%F %T')] ATTEMPT ${attempt}/${MAX_RETRY} gpu=${GPU} seed=${SEED} folds=${FOLDS:-all} reg_bypass=${REG_BYPASS} batch=${BATCH} trials=${N_TRIALS}" >>"${LOG}"
|
||||
t0=${SECONDS}
|
||||
|
||||
python -u -m lnp_ml.modeling.nested_cv_optuna \
|
||||
--input-path data/interim/internal.csv \
|
||||
--output-dir models/final_cv \
|
||||
--resume-dir "${RUN_DIR}" \
|
||||
--seed ${SEED} \
|
||||
--use-mpnn \
|
||||
--use-moe \
|
||||
--reg-bypass ${REG_BYPASS} \
|
||||
--use-llm --use-rag --rag-top-k 4 \
|
||||
--use-soft-prompt --llm-use-qlora \
|
||||
--llm-model-path models/qwen2.5-7b-instruct \
|
||||
--no-llm-freeze \
|
||||
--llm-max-length 1536 \
|
||||
--n-outer-folds ${N_OUTER} --n-inner-folds ${N_INNER} \
|
||||
--n-trials ${N_TRIALS} --epochs-per-trial ${EPOCHS} \
|
||||
--inner-patience ${PATIENCE} \
|
||||
--batch-size ${BATCH} --device cuda \
|
||||
${EXTRA[@]+"${EXTRA[@]}"} \
|
||||
>>"${LOG}" 2>&1
|
||||
rc=$?
|
||||
dt=$((SECONDS - t0))
|
||||
|
||||
if [ ${rc} -eq 0 ]; then
|
||||
echo "[$(date '+%F %T')] DONE gpu=${GPU} folds=${FOLDS:-all}(用时 ${dt}s)" >>"${LOG}"
|
||||
exit 0
|
||||
fi
|
||||
if [ ${dt} -lt ${MIN_OK_SEC} ]; then
|
||||
echo "[$(date '+%F %T')] 仅存活 ${dt}s 即以 ${rc} 退出,判定为配置/代码错误而非抢占,停止重试" >>"${LOG}"
|
||||
exit ${rc}
|
||||
fi
|
||||
echo "[$(date '+%F %T')] 运行 ${dt}s 后以 ${rc} 退出,${RETRY_WAIT}s 后带 --resume-dir 续跑" >>"${LOG}"
|
||||
sleep "${RETRY_WAIT}"
|
||||
done
|
||||
|
||||
echo "[$(date '+%F %T')] 超过 MAX_RETRY 次仍失败,放弃" >>"${LOG}"
|
||||
exit 1
|
||||
109
scripts_run/run_final_full.sh
Normal file
109
scripts_run/run_final_full.sh
Normal 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}MiB,60s 后重试" >>"${LOG}"
|
||||
sleep 60
|
||||
done
|
||||
}
|
||||
|
||||
# Optuna 的 study.optimize(n_trials=N) 在续跑时语义是"再跑 N 个",不是"补到 N 个"。
|
||||
# 所以每次重启前先数已完成的 trial 只补差额,否则被抢占几次就会多跑几十个 trial。
|
||||
remaining_trials() {
|
||||
if [ ! -f "${STUDY}" ]; then echo "${N_TRIALS}"; return; fi
|
||||
python - "${STUDY}" "${N_TRIALS}" <<'PY' 2>/dev/null || echo "${N_TRIALS}"
|
||||
import sys, optuna
|
||||
optuna.logging.set_verbosity(optuna.logging.WARNING)
|
||||
try:
|
||||
st = optuna.load_study(study_name="final_optuna_cv", storage=f"sqlite:///{sys.argv[1]}")
|
||||
done = sum(1 for t in st.trials if t.state == optuna.trial.TrialState.COMPLETE)
|
||||
except Exception:
|
||||
done = 0
|
||||
print(max(0, int(sys.argv[2]) - done))
|
||||
PY
|
||||
}
|
||||
|
||||
EXTRA=()
|
||||
if [ -n "${PRETRAIN}" ] && [ -f "${PRETRAIN}" ]; then
|
||||
EXTRA+=(--init-from-pretrain "${PRETRAIN}")
|
||||
elif [ -n "${PRETRAIN}" ]; then
|
||||
echo "[$(date '+%F %T')] 警告:${PRETRAIN} 不存在,跳过预训练初始化" >>"${LOG}"
|
||||
fi
|
||||
|
||||
for attempt in $(seq 1 "${MAX_RETRY}"); do
|
||||
wait_free
|
||||
NT=$(remaining_trials)
|
||||
echo "[$(date '+%F %T')] ATTEMPT ${attempt}/${MAX_RETRY} gpu=${GPU} 本次补 ${NT} 个 trial(目标 ${N_TRIALS})reg_bypass=${REG_BYPASS} batch=${BATCH}" >>"${LOG}"
|
||||
t0=${SECONDS}
|
||||
|
||||
python -u -m lnp_ml.modeling.final_train_optuna_cv \
|
||||
--input-path data/interim/internal.csv \
|
||||
--output-dir "${OUT}" \
|
||||
--seed ${SEED} \
|
||||
--n-folds ${N_FOLDS} \
|
||||
--n-trials ${NT} \
|
||||
--epochs-per-trial ${EPOCHS} \
|
||||
--patience ${PATIENCE} \
|
||||
--batch-size ${BATCH} \
|
||||
--use-mpnn --use-moe \
|
||||
--reg-bypass ${REG_BYPASS} \
|
||||
--use-llm --use-rag --rag-top-k 4 \
|
||||
--use-soft-prompt --llm-use-qlora \
|
||||
--llm-model-path models/qwen2.5-7b-instruct \
|
||||
--no-llm-freeze \
|
||||
--llm-max-length 1536 \
|
||||
--freeze-backbone-epochs ${FREEZE} \
|
||||
--device cuda \
|
||||
${EXTRA[@]+"${EXTRA[@]}"} \
|
||||
>>"${LOG}" 2>&1
|
||||
rc=$?
|
||||
dt=$((SECONDS - t0))
|
||||
|
||||
if [ ${rc} -eq 0 ]; then
|
||||
echo "[$(date '+%F %T')] DONE(用时 ${dt}s)→ ${OUT}/model.pt" >>"${LOG}"
|
||||
exit 0
|
||||
fi
|
||||
if [ ${dt} -lt ${MIN_OK_SEC} ]; then
|
||||
echo "[$(date '+%F %T')] 仅存活 ${dt}s 即以 ${rc} 退出,判为配置/代码错误而非抢占,停止重试" >>"${LOG}"
|
||||
exit ${rc}
|
||||
fi
|
||||
echo "[$(date '+%F %T')] 运行 ${dt}s 后以 ${rc} 退出,${RETRY_WAIT}s 后续跑" >>"${LOG}"
|
||||
sleep "${RETRY_WAIT}"
|
||||
done
|
||||
|
||||
echo "[$(date '+%F %T')] 超过 MAX_RETRY 次仍失败,放弃" >>"${LOG}"
|
||||
exit 1
|
||||
@ -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:
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user