| """
|
| bert/train.py
|
| Fine-tuning loop for the DeBERTa-v3 Cross-Encoder entailment verifier.
|
|
|
| Pipeline:
|
| 1. Load DeBERTa-v3-base with 3-label classification head
|
| 2. ANLI (R1+R2+R3) + TrueTeacher + MNLI combined dataset
|
| 3. Weighted CrossEntropyLoss (Contradiction=2.0, Neutral=1.5, Entailment=1.0)
|
| 4. AdamW + linear warmup + cosine decay
|
| 5. Checkpoint every epoch; early stop on validation loss
|
| """
|
|
|
| from __future__ import annotations
|
|
|
| import argparse
|
| import math
|
| import os
|
| from pathlib import Path
|
| from typing import Optional
|
|
|
| import torch
|
| import torch.nn as nn
|
| from torch.optim import AdamW
|
| from torch.optim.lr_scheduler import LambdaLR
|
| from tqdm import tqdm
|
|
|
| from bert.dataset import build_combined_dataset, make_dataloader
|
| from bert.model import BertCrossEncoderVerifier, CrossEncoderConfig, build_model, load_tokenizer
|
|
|
|
|
| def get_linear_warmup_cosine_schedule(
|
| optimizer: AdamW,
|
| num_warmup_steps: int,
|
| num_training_steps: int,
|
| ) -> LambdaLR:
|
| def lr_lambda(current_step: int) -> float:
|
| if current_step < num_warmup_steps:
|
| return float(current_step) / float(max(1, num_warmup_steps))
|
| progress = float(current_step - num_warmup_steps) / float(
|
| max(1, num_training_steps - num_warmup_steps)
|
| )
|
| return max(0.0, 0.5 * (1.0 + math.cos(math.pi * progress)))
|
|
|
| return LambdaLR(optimizer, lr_lambda)
|
|
|
|
|
| def train_epoch(
|
| model: BertCrossEncoderVerifier,
|
| loader: torch.utils.data.DataLoader,
|
| optimizer: AdamW,
|
| scheduler: LambdaLR,
|
| device: torch.device,
|
| grad_accum_steps: int = 4,
|
| max_grad_norm: float = 1.0,
|
| ) -> float:
|
| model.train()
|
| total_loss = 0.0
|
| optimizer.zero_grad()
|
|
|
| for step, batch in enumerate(tqdm(loader, desc="train")):
|
| input_ids = batch["input_ids"].to(device)
|
| attention_mask = batch["attention_mask"].to(device)
|
| token_type_ids = batch.get("token_type_ids")
|
| if token_type_ids is not None:
|
| token_type_ids = token_type_ids.to(device)
|
| labels = batch["label"].to(device)
|
|
|
| out = model(input_ids, attention_mask, token_type_ids, labels)
|
| loss = out["loss"] / grad_accum_steps
|
| loss.backward()
|
| total_loss += loss.item() * grad_accum_steps
|
|
|
| if (step + 1) % grad_accum_steps == 0:
|
| nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)
|
| optimizer.step()
|
| scheduler.step()
|
| optimizer.zero_grad()
|
|
|
| return total_loss / len(loader)
|
|
|
|
|
| @torch.no_grad()
|
| def evaluate(
|
| model: BertCrossEncoderVerifier,
|
| loader: torch.utils.data.DataLoader,
|
| device: torch.device,
|
| ) -> dict:
|
| model.eval()
|
| total_loss = 0.0
|
| correct = 0
|
| total = 0
|
|
|
| class_correct = [0, 0, 0]
|
| class_total = [0, 0, 0]
|
|
|
| for batch in tqdm(loader, desc="eval"):
|
| input_ids = batch["input_ids"].to(device)
|
| attention_mask = batch["attention_mask"].to(device)
|
| token_type_ids = batch.get("token_type_ids")
|
| if token_type_ids is not None:
|
| token_type_ids = token_type_ids.to(device)
|
| labels = batch["label"].to(device)
|
|
|
| out = model(input_ids, attention_mask, token_type_ids, labels)
|
| total_loss += out["loss"].item()
|
|
|
| preds = out["logits"].argmax(dim=-1)
|
| correct += (preds == labels).sum().item()
|
| total += labels.size(0)
|
| for c in range(3):
|
| mask = labels == c
|
| class_correct[c] += (preds[mask] == labels[mask]).sum().item()
|
| class_total[c] += mask.sum().item()
|
|
|
| acc = correct / total if total > 0 else 0.0
|
| per_class = {
|
| c: class_correct[c] / class_total[c] if class_total[c] > 0 else 0.0
|
| for c in range(3)
|
| }
|
| label_names = {0: "contradiction", 1: "neutral", 2: "entailment"}
|
| return {
|
| "loss": total_loss / len(loader),
|
| "accuracy": acc,
|
| "per_class_accuracy": {label_names[k]: v for k, v in per_class.items()},
|
| }
|
|
|
|
|
| def train(
|
| data_dir: Path,
|
| output_dir: Path,
|
| backbone: str = "microsoft/deberta-v3-base",
|
| epochs: int = 5,
|
| batch_size: int = 32,
|
| lr: float = 2e-5,
|
| warmup_ratio: float = 0.06,
|
| max_length: int = 512,
|
| grad_accum: int = 4,
|
| seed: int = 42,
|
| ) -> None:
|
| torch.manual_seed(seed)
|
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| print(f"[train] device={device}, backbone={backbone}")
|
|
|
| tokenizer = load_tokenizer(backbone)
|
| config = CrossEncoderConfig(backbone=backbone, max_length=max_length)
|
| model = build_model(config).to(device)
|
|
|
| train_ds = build_combined_dataset(data_dir, tokenizer, "train", max_length, seed)
|
| val_ds = build_combined_dataset(data_dir, tokenizer, "dev", max_length, seed)
|
| train_loader = make_dataloader(train_ds, batch_size, shuffle=True)
|
| val_loader = make_dataloader(val_ds, batch_size, shuffle=False)
|
|
|
| num_training_steps = epochs * len(train_loader) // grad_accum
|
| num_warmup_steps = int(warmup_ratio * num_training_steps)
|
|
|
| optimizer = AdamW(model.parameters(), lr=lr, weight_decay=0.01, eps=1e-8)
|
| scheduler = get_linear_warmup_cosine_schedule(optimizer, num_warmup_steps, num_training_steps)
|
|
|
| output_dir.mkdir(parents=True, exist_ok=True)
|
| best_val_loss = float("inf")
|
|
|
| for epoch in range(1, epochs + 1):
|
| print(f"\nββ Epoch {epoch}/{epochs} ββββββββββββββββββββββ")
|
| train_loss = train_epoch(model, train_loader, optimizer, scheduler, device, grad_accum)
|
| val_metrics = evaluate(model, val_loader, device)
|
|
|
| print(f" train_loss={train_loss:.4f}")
|
| print(f" val_loss={val_metrics['loss']:.4f} val_acc={val_metrics['accuracy']:.4f}")
|
| print(f" per_class={val_metrics['per_class_accuracy']}")
|
|
|
| ckpt_path = output_dir / f"checkpoint_epoch{epoch}.pt"
|
| torch.save({
|
| "epoch": epoch,
|
| "model_state_dict": model.state_dict(),
|
| "optimizer_state_dict": optimizer.state_dict(),
|
| "val_loss": val_metrics["loss"],
|
| "config": config,
|
| }, ckpt_path)
|
| print(f" saved β {ckpt_path}")
|
|
|
| if val_metrics["loss"] < best_val_loss:
|
| best_val_loss = val_metrics["loss"]
|
| best_path = output_dir / "best_model.pt"
|
| torch.save(model.state_dict(), best_path)
|
| print(f" β
new best β {best_path}")
|
|
|
|
|
| tokenizer.save_pretrained(output_dir / "tokenizer")
|
| print(f"\n[train] complete. Best val_loss={best_val_loss:.4f}")
|
|
|
|
|
| if __name__ == "__main__":
|
| parser = argparse.ArgumentParser()
|
| parser.add_argument("--data_dir", type=Path, required=True)
|
| parser.add_argument("--output_dir", type=Path, default=Path("checkpoints"))
|
| parser.add_argument("--backbone", type=str, default="microsoft/deberta-v3-base")
|
| parser.add_argument("--epochs", type=int, default=5)
|
| parser.add_argument("--batch_size", type=int, default=32)
|
| parser.add_argument("--lr", type=float, default=2e-5)
|
| parser.add_argument("--max_length", type=int, default=512)
|
| parser.add_argument("--grad_accum", type=int, default=4)
|
| args = parser.parse_args()
|
| train(**vars(args))
|
|
|