admesh/agentic-intent-classifier
251
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 