import json import copy import shutil from pathlib import Path import optuna from optuna.samplers import TPESampler from optuna.pruners import MedianPruner from transformers import TrainerCallback from src.training.train_model import ( load_yaml_config, resolve_runtime_config, train_once, ) class OptunaPruningCallback(TrainerCallback): """Report dev-metric mỗi lần eval -> MedianPruner cắt sớm trial tệ (Phase 02). Pruner cũ vô dụng vì objective chỉ trả best cuối, KHÔNG report intermediate.""" def __init__(self, trial, metric_name: str = "eval_macro_f1"): self.trial = trial self.metric_name = metric_name self.step = 0 def on_evaluate(self, args, state, control, metrics=None, **kwargs): if not metrics or self.metric_name not in metrics: return self.trial.report(metrics[self.metric_name], step=self.step) self.step += 1 if self.trial.should_prune(): raise optuna.TrialPruned() def objective(trial: optuna.Trial, base_config: dict) -> float: config = copy.deepcopy(base_config) model_type = config["model"].get("type", "encoder") if model_type == "decoder_lora": # LEAN: mỗi trial = train LLM 7B (~20-30 phút) -> không sweep batch/max_length. config["training"]["learning_rate"] = trial.suggest_float("learning_rate", 5e-5, 3e-4, log=True) config["training"]["epochs"] = trial.suggest_int("epochs", 2, 4) lora_r = trial.suggest_categorical("lora_r", [8, 16, 32]) config["model"].setdefault("lora", {}) config["model"]["lora"]["r"] = lora_r config["model"]["lora"]["alpha"] = 2 * lora_r else: # encoder (PhoBERT) — search LR/batch/wd/warmup (Phase 02: BỎ epochs, THÊM warmup). # max_length KHÔNG search: set per-task từ data (topic=128, subst=256) trong config/train.yml. config["training"]["learning_rate"] = trial.suggest_float("learning_rate", 1e-5, 5e-5, log=True) # batch search per-task từ config (subst data ít -> [8,16]; topic -> [8,16,32]). Xem config/train.yml. batch_choices = config.get("tune", {}).get("train_batch_size", [8, 16, 32]) config["training"]["train_batch_size"] = trial.suggest_categorical("train_batch_size", batch_choices) config["training"]["weight_decay"] = trial.suggest_float("weight_decay", 0.0, 0.10, step=0.0001) # warmup_ratio: lever ổn định fine-tune (Mosbach/Dodge). KHÔNG tune dropout (chốt 2026-06-13). config["training"]["warmup_ratio"] = trial.suggest_float("warmup_ratio", 0.0, 0.2) # epochs KHÔNG tune: cố định = trần config (10) + EarlyStopping quyết best-epoch. trial_dir = Path(config["paths"]["output_dir"]) / f"trial_{trial.number}" config["paths"]["output_dir"] = str(trial_dir) # Không lưu checkpoint trong quá trình tuning để tiết kiệm disk. # EarlyStoppingCallback vẫn track best_metric qua state nên metric vẫn đúng. config["training"]["save_strategy"] = "no" config["training"]["load_best_model_at_end"] = False print(f"\n{'='*50}") print(f"TRIAL {trial.number} [{model_type}] — {trial.params}") print(f"{'='*50}") metric_name = "eval_" + config["training"].get("metric_for_best_model", "macro_f1") pruning_cb = OptunaPruningCallback(trial, metric_name=metric_name) # try/finally: trial bị PRUNE sẽ raise TrialPruned giữa chừng -> phải dọn trial_dir ở finally # (nếu không, mỗi trial pruned để lại 1 thư mục trial_N rỗng trong outputs/models/