"""Finetune a transformer for stance detection. Picks the checkpoint off dev Favg2, not loss or a fixed epoch count. Loss fn is configurable -- plain CE, inverse-frequency weighted, or focal -- to deal with the class imbalance. python -m src.train --config configs/track1.yaml """ import argparse import json import os import random import numpy as np import torch import torch.nn.functional as F import yaml from torch.utils.data import DataLoader from transformers import ( AutoModelForSequenceClassification, AutoTokenizer, get_linear_schedule_with_warmup, ) from src.data import ID2LABEL, LABEL2ID, StanceDataset, load_split from src.scorer import score def set_seed(seed): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) def focal_loss(logits, targets, gamma, weight=None): ce = F.cross_entropy(logits, targets, weight=weight, reduction="none") pt = torch.exp(-ce) return ((1 - pt) ** gamma * ce).mean() def class_weights(df, device): counts = np.array( [(df["label"] == i).sum() for i in range(3)], dtype=np.float64 ) counts = np.clip(counts, 1, None) w = counts.sum() / (3.0 * counts) return torch.tensor(w, dtype=torch.float, device=device) @torch.no_grad() def predict_logits(model, loader, device): model.eval() out = [] for batch in loader: batch = { k: v.to(device) for k, v in batch.items() if k != "labels" } out.append(model(**batch).logits.float().cpu().numpy()) return np.concatenate(out, axis=0) def logits_to_labels(logits, none_bias=0.0): """A negative none_bias lowers the None logit before argmax.""" adj = logits.copy() adj[:, LABEL2ID["None"]] += none_bias return [ID2LABEL[i] for i in adj.argmax(axis=1)] def build_config(): ap = argparse.ArgumentParser() ap.add_argument("--config", required=True) ap.add_argument("--overrides", default="", help="k=v,k=v pairs") args = ap.parse_args() cfg = yaml.safe_load(open(args.config)) for kv in [x for x in args.overrides.split(",") if x]: k, v = kv.split("=", 1) cfg[k] = yaml.safe_load(v) return cfg def main(): cfg = build_config() print("[config]", json.dumps(cfg, ensure_ascii=False)) set_seed(cfg.get("seed", 42)) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") prep = cfg.get("prep_mode", "preserve") train_df = load_split(cfg["train_csv"], prep) dev_df = load_split(cfg["dev_csv"], prep) trust = cfg.get("trust_remote_code", False) tok = AutoTokenizer.from_pretrained( cfg["model_hf"], trust_remote_code=trust ) max_len = cfg.get("max_len", 128) use_desc = cfg.get("use_description", False) train_ds = StanceDataset(train_df, tok, max_len, use_desc) dev_ds = StanceDataset(dev_df, tok, max_len, use_desc) train_loader = DataLoader( train_ds, batch_size=cfg.get("batch_size", 16), shuffle=True ) dev_loader = DataLoader( dev_ds, batch_size=cfg.get("eval_batch_size", 64), shuffle=False ) model = AutoModelForSequenceClassification.from_pretrained( cfg["model_hf"], num_labels=3, id2label=ID2LABEL, label2id=LABEL2ID, trust_remote_code=trust, ).to(device) optim = torch.optim.AdamW( model.parameters(), lr=cfg.get("lr", 2e-5), weight_decay=cfg.get("weight_decay", 0.01), ) epochs = cfg.get("epochs", 10) total_steps = len(train_loader) * epochs sched = get_linear_schedule_with_warmup( optim, int(0.06 * total_steps), total_steps ) loss_type = cfg.get("loss", "ce") weight = None if loss_type in ("weighted", "focal_weighted"): weight = class_weights(train_df, device) gamma = cfg.get("focal_gamma", 2.0) out_dir = cfg["out_dir"] os.makedirs(out_dir, exist_ok=True) best_favg2, best_epoch = -1.0, -1 patience = cfg.get("patience", 3) for epoch in range(epochs): model.train() running = 0.0 for batch in train_loader: batch = {k: v.to(device) for k, v in batch.items()} labels = batch.pop("labels") logits = model(**batch).logits if loss_type.startswith("focal"): loss = focal_loss(logits, labels, gamma, weight) else: loss = F.cross_entropy(logits, labels, weight=weight) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optim.step() sched.step() optim.zero_grad() running += loss.item() logits = predict_logits(model, dev_loader, device) preds = logits_to_labels(logits, cfg.get("none_bias", 0.0)) print( f"\n=== epoch {epoch + 1}/{epochs} " f"train_loss={running / len(train_loader):.4f} ===" ) res = score(dev_df[["target", "stance"]], preds) favg2 = res["overall"]["Favg2"] if favg2 > best_favg2: best_favg2, best_epoch = favg2, epoch + 1 model.save_pretrained(out_dir) tok.save_pretrained(out_dir) np.save(os.path.join(out_dir, "best_dev_logits.npy"), logits) json.dump( { "best_epoch": best_epoch, "best_favg2": best_favg2, "config": cfg, }, open(os.path.join(out_dir, "best.json"), "w"), ensure_ascii=False, indent=2, ) print(f" new best Favg2={best_favg2:.4f} (saved)") elif epoch + 1 - best_epoch >= patience: print(f" early stop after {patience} epochs without gain") break print(f"\nBEST dev Favg2={best_favg2:.4f} @ epoch {best_epoch} " f"-> {out_dir}") if __name__ == "__main__": main()