article_classifier / src /train.py
asriel14's picture
Upload 4 files
5499d76 verified
Raw History Blame Contribute Delete
5.31 kB
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()