CoolFace
Modelpublic

admesh/agentic-intent-classifier

sourceHugging Faceapache-2.0updated 6mo agoView on Hugging Face
2likes51downloads
train_subtype.py126 linesDownload Raw Back to training
1import sys2from pathlib import Path3 4import torch5from transformers import AutoModelForSequenceClassification, AutoTokenizer, Trainer, TrainingArguments6 7BASE_DIR = Path(__file__).resolve().parent.parent8if str(BASE_DIR) not in sys.path:9    sys.path.insert(0, str(BASE_DIR))10 11from config import (12    FULL_INTENT_TAXONOMY_DATA_DIR,13    SUBTYPE_DIFFICULTY_DATA_DIR,14    SUBTYPE_HEAD_CONFIG,15    SUBTYPE_TRAINING_WEIGHTS,16)17from training.common import (18    build_label_weight_tensor,19    compute_classification_metrics,20    load_labeled_rows,21    load_labeled_rows_from_paths,22    prepare_dataset,23    write_json,24)25 26 27class WeightedTrainer(Trainer):28    def __init__(self, *args, class_weights: torch.Tensor | None = None, **kwargs):29        super().__init__(*args, **kwargs)30        self.class_weights = class_weights31 32    def compute_loss(self, model, inputs, return_outputs=False, **kwargs):33        labels = inputs.pop("labels")34        outputs = model(**inputs)35        logits = outputs.get("logits")36        weight = self.class_weights.to(logits.device) if self.class_weights is not None else None37        loss_fct = torch.nn.CrossEntropyLoss(weight=weight)38        loss = loss_fct(logits.view(-1, model.config.num_labels), labels.view(-1))39        return (loss, outputs) if return_outputs else loss40 41 42train_rows = load_labeled_rows_from_paths(43    [44        SUBTYPE_HEAD_CONFIG.split_paths["train"],45        FULL_INTENT_TAXONOMY_DATA_DIR / "train.jsonl",46        SUBTYPE_DIFFICULTY_DATA_DIR / "train.jsonl",47    ],48    SUBTYPE_HEAD_CONFIG.label_field,49    SUBTYPE_HEAD_CONFIG.label2id,50)51val_rows = load_labeled_rows_from_paths(52    [53        SUBTYPE_HEAD_CONFIG.split_paths["val"],54        FULL_INTENT_TAXONOMY_DATA_DIR / "val.jsonl",55        SUBTYPE_DIFFICULTY_DATA_DIR / "val.jsonl",56    ],57    SUBTYPE_HEAD_CONFIG.label_field,58    SUBTYPE_HEAD_CONFIG.label2id,59)60test_rows = load_labeled_rows(61    SUBTYPE_HEAD_CONFIG.split_paths["test"],62    SUBTYPE_HEAD_CONFIG.label_field,63    SUBTYPE_HEAD_CONFIG.label2id,64)65 66tokenizer = AutoTokenizer.from_pretrained(SUBTYPE_HEAD_CONFIG.model_name)67 68train_dataset = prepare_dataset(train_rows, tokenizer, SUBTYPE_HEAD_CONFIG.max_length)69val_dataset = prepare_dataset(val_rows, tokenizer, SUBTYPE_HEAD_CONFIG.max_length)70test_dataset = prepare_dataset(test_rows, tokenizer, SUBTYPE_HEAD_CONFIG.max_length)71class_weights = build_label_weight_tensor(SUBTYPE_HEAD_CONFIG.labels, SUBTYPE_TRAINING_WEIGHTS)72 73model = AutoModelForSequenceClassification.from_pretrained(74    SUBTYPE_HEAD_CONFIG.model_name,75    num_labels=len(SUBTYPE_HEAD_CONFIG.labels),76    id2label=SUBTYPE_HEAD_CONFIG.id2label,77    label2id=SUBTYPE_HEAD_CONFIG.label2id,78)79 80training_args = TrainingArguments(81    output_dir=str(SUBTYPE_HEAD_CONFIG.model_dir),82    eval_strategy="epoch",83    save_strategy="no",84    logging_strategy="epoch",85    num_train_epochs=4,86    per_device_train_batch_size=4,87    per_device_eval_batch_size=4,88    learning_rate=2e-5,89    weight_decay=0.01,90    report_to="none",91)92 93trainer = WeightedTrainer(94    model=model,95    args=training_args,96    train_dataset=train_dataset,97    eval_dataset=val_dataset,98    compute_metrics=compute_classification_metrics,99    class_weights=class_weights,100)101 102print(103    f"Loaded subtype splits: train={len(train_rows)} val={len(val_rows)} test={len(test_rows)}"104)105print(f"Subtype weights: {[round(float(x), 3) for x in class_weights.tolist()]}")106trainer.train()107val_metrics = trainer.evaluate(eval_dataset=val_dataset, metric_key_prefix="val")108test_metrics = trainer.evaluate(eval_dataset=test_dataset, metric_key_prefix="test")109print(val_metrics)110print(test_metrics)111 112SUBTYPE_HEAD_CONFIG.model_dir.mkdir(parents=True, exist_ok=True)113model.save_pretrained(SUBTYPE_HEAD_CONFIG.model_dir)114tokenizer.save_pretrained(SUBTYPE_HEAD_CONFIG.model_dir)115write_json(116    SUBTYPE_HEAD_CONFIG.model_dir / "train_metrics.json",117    {118        "head": SUBTYPE_HEAD_CONFIG.slug,119        "train_count": len(train_rows),120        "val_count": len(val_rows),121        "test_count": len(test_rows),122        "val_metrics": val_metrics,123        "test_metrics": test_metrics,124    },125)126