mrs83's picture
Upload folder using huggingface_hub
bc67aa1 verified
Raw
History Blame Contribute Delete
9.72 kB
#!/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()