from __future__ import annotations import argparse import json from pathlib import Path import evaluate import numpy as np from datasets import Dataset, DatasetDict from sklearn.model_selection import train_test_split from transformers import ( AutoModelForSequenceClassification, AutoTokenizer, DataCollatorWithPadding, Trainer, TrainingArguments, ) from src.data_utils import filter_rare_classes, load_dataset_frame, save_json def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Train article topic classifier.") parser.add_argument("--data-path", type=str, required=True, help="Path to CSV/JSONL/Parquet dataset.") parser.add_argument("--text-cols", nargs="*", default=None, help="Optional explicit text columns: title abstract") parser.add_argument("--label-col", type=str, default=None, help="Optional explicit label column.") parser.add_argument("--output-dir", type=str, default="artifacts/article_topic_model") parser.add_argument("--model-name", type=str, default="allenai/scibert_scivocab_uncased") parser.add_argument("--max-length", type=int, default=256) parser.add_argument("--epochs", type=int, default=3) parser.add_argument("--lr", type=float, default=2e-5) parser.add_argument("--weight-decay", type=float, default=0.01) parser.add_argument("--train-batch-size", type=int, default=8) parser.add_argument("--eval-batch-size", type=int, default=16) parser.add_argument("--test-size", type=float, default=0.15) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--min-examples-per-class", type=int, default=20) return parser.parse_args() def main() -> None: args = parse_args() output_dir = Path(args.output_dir) output_dir.mkdir(parents=True, exist_ok=True) df = load_dataset_frame(args.data_path, text_cols=args.text_cols, label_col=args.label_col) df = filter_rare_classes(df, min_examples_per_class=args.min_examples_per_class) labels = sorted(df["label"].unique().tolist()) label2id = {label: idx for idx, label in enumerate(labels)} id2label = {idx: label for label, idx in label2id.items()} df["label_id"] = df["label"].map(label2id) train_df, valid_df = train_test_split( df, test_size=args.test_size, random_state=args.seed, stratify=df["label_id"], ) ds = DatasetDict( { "train": Dataset.from_pandas(train_df[["text", "label_id"]], preserve_index=False), "validation": Dataset.from_pandas(valid_df[["text", "label_id"]], preserve_index=False), } ) tokenizer = AutoTokenizer.from_pretrained(args.model_name) def tokenize_batch(batch: dict) -> dict: return tokenizer(batch["text"], truncation=True, max_length=args.max_length) tokenized = ds.map(tokenize_batch, batched=True) tokenized = tokenized.rename_column("label_id", "labels") tokenized.set_format(type="torch", columns=["input_ids", "attention_mask", "labels"]) model = AutoModelForSequenceClassification.from_pretrained( args.model_name, num_labels=len(labels), id2label={int(k): v for k, v in id2label.items()}, label2id=label2id, ) accuracy_metric = evaluate.load("accuracy") f1_metric = evaluate.load("f1") def compute_metrics(eval_pred): logits, labels_np = eval_pred preds = np.argmax(logits, axis=-1) accuracy = accuracy_metric.compute(predictions=preds, references=labels_np)["accuracy"] f1_macro = f1_metric.compute(predictions=preds, references=labels_np, average="macro")["f1"] f1_weighted = f1_metric.compute(predictions=preds, references=labels_np, average="weighted")["f1"] return { "accuracy": accuracy, "f1_macro": f1_macro, "f1_weighted": f1_weighted, } training_args = TrainingArguments( output_dir=str(output_dir / "checkpoints"), learning_rate=args.lr, per_device_train_batch_size=args.train_batch_size, per_device_eval_batch_size=args.eval_batch_size, num_train_epochs=args.epochs, weight_decay=args.weight_decay, eval_strategy="epoch", save_strategy="epoch", logging_strategy="steps", logging_steps=50, load_best_model_at_end=True, metric_for_best_model="f1_macro", greater_is_better=True, save_total_limit=2, report_to="none", seed=args.seed, ) trainer = Trainer( model=model, args=training_args, train_dataset=tokenized["train"], eval_dataset=tokenized["validation"], processing_class=tokenizer, data_collator=DataCollatorWithPadding(tokenizer=tokenizer), compute_metrics=compute_metrics, ) trainer.train() metrics = trainer.evaluate() model.save_pretrained(output_dir) tokenizer.save_pretrained(output_dir) save_json( { "label2id": label2id, "id2label": {str(k): v for k, v in id2label.items()}, }, output_dir / "label_mapping.json", ) save_json(metrics, output_dir / "metrics.json") print(json.dumps(metrics, indent=2)) if __name__ == "__main__": main()