stanceeval2026 / code /src /train.py
zaher-m's picture
Add files using upload-large-folder tool
7e9cfd1 verified
Raw
History Blame Contribute Delete
5.95 kB
"""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()