#!/usr/bin/env python3 """ Fine-tune an EchoModelForSentenceEmbedding → EchoForSequenceClassification on Amazon MASSIVE intent classification. Usage: PYTHONPATH=. python training/train_intent_clf.py \ --embed_model ethicalabs/Echo-DSRN-v0.1.3-Embed-Intent \ --num_labels 60 \ --batch_size 32 --lr 2e-5 --epochs 5 \ --output_dir models/Echo-DSRN-v0.1.3-Embed-Intent-CLF """ import argparse import os import sys import numpy as np import torch from datasets import load_dataset from transformers import ( AutoTokenizer, EarlyStoppingCallback, Trainer, TrainerCallback, TrainingArguments, ) sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import echo_dsrn # noqa: F401 — registers AutoModel classes from echo_dsrn.modeling_echo import EchoForSequenceClassification from echo_embedding.modeling_embedding import EchoModelForSentenceEmbedding MASSIVE_REVISION = "refs/convert/parquet" class NaNCheckCallback(TrainerCallback): """Stop training if a NaN gradient is detected.""" def on_pre_optimizer_step(self, args, state, control, model=None, **kwargs): if model is None: return for name, param in model.named_parameters(): if param.grad is not None and not torch.isfinite(param.grad).all(): print(f"❌ NaN gradient in {name} — stopping training.", flush=True) control.should_training_stop = True return INTENT_NAMES = [] # loaded from MASSIVE dataset in main() def main(): parser = argparse.ArgumentParser(description="Fine-tune embed→classifier on MASSIVE") parser.add_argument( "--embed_model", default="ethicalabs/Echo-DSRN-v0.1.3-Embed-Intent", help="HF path to the embedding model", ) parser.add_argument("--num_labels", type=int, default=60) parser.add_argument("--freeze_backbone", action="store_true", help="Freeze backbone, train head only") parser.add_argument( "--sklearn_init", action="store_true", help="Precompute embeddings, fit sklearn SGDClassifier, and copy weights into the head before training", ) parser.add_argument("--max_train_samples", type=int, default=0) parser.add_argument("--max_eval_samples", type=int, default=0) parser.add_argument("--output_dir", default="models/Echo-DSRN-v0.1.3-Embed-Intent-CLF") parser.add_argument("--batch_size", type=int, default=32) parser.add_argument("--lr", type=float, default=2e-5) parser.add_argument("--epochs", type=int, default=5) parser.add_argument("--eval_steps", type=int, default=1000) parser.add_argument("--early_stopping_patience", type=int, default=3) parser.add_argument("--resume_from_checkpoint", type=str, default=None) parser.add_argument("--bf16", action="store_true", default=True) parser.add_argument("--seed", type=int, default=42) args = parser.parse_args() # ── 1. Load embedding model and convert to classifier ───────── print(f"⚡ Loading embedding model: {args.embed_model}") embed_model = EchoModelForSentenceEmbedding.from_pretrained( args.embed_model, trust_remote_code=True, ) # SentenceTransformer wrapper (needed for .encode() in sklearn_init) if args.sklearn_init: from sentence_transformers import SentenceTransformer st_model = SentenceTransformer(args.embed_model, trust_remote_code=True) tokenizer = AutoTokenizer.from_pretrained(args.embed_model, trust_remote_code=True) print(f" dim={embed_model.config.hidden_size}x{embed_model.config.num_heads}") # Load canonical intent names from MASSIVE dataset (never hardcode ordering) global INTENT_NAMES from datasets import load_dataset as _load_ds intent_ds = _load_ds("AmazonScience/massive", split="train", revision=MASSIVE_REVISION) INTENT_NAMES = intent_ds.features["intent"].names print(f" Loaded {len(INTENT_NAMES)} intent names from MASSIVE") id2label = {i: name for i, name in enumerate(INTENT_NAMES[:args.num_labels])} label2id = {v: k for k, v in id2label.items()} print(f"⚡ Converting to sequence classifier ({args.num_labels} labels)...") model = EchoForSequenceClassification.from_embedding( embed_model, num_labels=args.num_labels, id2label=id2label, label2id=label2id, ) # ── 2. Load MASSIVE dataset ────────────────────────────────── print("⚡ Loading MASSIVE dataset...") raw_train = load_dataset("AmazonScience/massive", split="train", revision=MASSIVE_REVISION) raw_val = load_dataset("AmazonScience/massive", split="validation", revision=MASSIVE_REVISION) raw_train = raw_train.select_columns(["utt", "intent"]) raw_val = raw_val.select_columns(["utt", "intent"]) # ── 2b. Sklearn init ───────────────────────────────────── if args.sklearn_init: print("🧠 Fitting sklearn SGDClassifier on precomputed embeddings...") from sklearn.linear_model import SGDClassifier # Use SentenceTransformer for fast encoding emb_train = st_model.encode( raw_train["utt"], batch_size=64, show_progress_bar=True, convert_to_numpy=True, ) emb_train = emb_train.astype(np.float32) labels_train = np.array(raw_train["intent"], dtype=np.int64) clf = SGDClassifier(loss="log_loss", max_iter=1000, tol=1e-3, random_state=args.seed) clf.fit(emb_train, labels_train) # Copy coefficients with torch.no_grad(): model.classifier.weight.copy_(torch.from_numpy(clf.coef_)) model.classifier.bias.copy_(torch.from_numpy(clf.intercept_)) acc = clf.score(emb_train, labels_train) print(f" Sklearn training accuracy: {acc:.4f}") # Save immediately so re-runs skip the expensive encode+fit step os.makedirs(args.output_dir, exist_ok=True) model.save_pretrained(args.output_dir) print(f" Sklearn-initialized model saved to {args.output_dir}") del st_model # free SentenceTransformer wrapper del emb_train, labels_train, clf # free numpy arrays del embed_model # free memory if args.freeze_backbone: print(" ❄ Freezing backbone, training classifier head only") for param in model.model.parameters(): param.requires_grad = False model.model.eval() else: print(" 🔥 Full fine-tuning (backbone + head)") def tokenize_fn(examples): return tokenizer( examples["utt"], truncation=True, padding="max_length", max_length=128, ) if args.max_train_samples > 0: raw_train = raw_train.shuffle(seed=args.seed).select( range(min(args.max_train_samples, len(raw_train))) ) if args.max_eval_samples > 0: raw_val = raw_val.shuffle(seed=args.seed).select( range(min(args.max_eval_samples, len(raw_val))) ) print(f" Train: {len(raw_train)} | Val: {len(raw_val)}") # ── 3. Create HF datasets with tokenized inputs ────────────── train_ds = raw_train.map(tokenize_fn, batched=True, remove_columns=["utt"]) train_ds = train_ds.rename_column("intent", "labels") eval_ds = raw_val.map(tokenize_fn, batched=True, remove_columns=["utt"]) eval_ds = eval_ds.rename_column("intent", "labels") # ── 4. Training arguments ──────────────────────────────────── eval_strategy = "steps" if args.eval_steps > 0 else "epoch" training_args = TrainingArguments( output_dir=args.output_dir, num_train_epochs=args.epochs, per_device_train_batch_size=args.batch_size, per_device_eval_batch_size=args.batch_size, learning_rate=args.lr, warmup_ratio=0.1, weight_decay=0.01, logging_steps=100, eval_strategy=eval_strategy, eval_steps=args.eval_steps if args.eval_steps > 0 else None, save_strategy=eval_strategy, save_steps=args.eval_steps if args.eval_steps > 0 else None, save_total_limit=2, load_best_model_at_end=True, metric_for_best_model="eval_loss", greater_is_better=False, bf16=args.bf16, report_to="none", seed=args.seed, dataloader_drop_last=False, remove_unused_columns=False, ) callbacks = [NaNCheckCallback()] if args.early_stopping_patience > 0: callbacks.append(EarlyStoppingCallback(early_stopping_patience=args.early_stopping_patience)) # ── 5. Train ──────────────────────────────────────────────── print("🚀 Starting training...") trainer = Trainer( model=model, args=training_args, train_dataset=train_ds, eval_dataset=eval_ds, callbacks=callbacks, ) resume = args.resume_from_checkpoint if resume and resume.lower() == "true": resume = True trainer.train(resume_from_checkpoint=resume or None) # ── 6. Save ───────────────────────────────────────────────── print(f"💾 Saving to {args.output_dir}") model.save_pretrained(args.output_dir) print(f"✅ Done — model saved to {args.output_dir}") if __name__ == "__main__": main()