Spaces:
Running
Running
Download src/train.py from asriel14/article_classifier: direct link, hf CLI and curl.
- Browser
- Download file 5.31 kB
-
https://huggingface.co/spaces/asriel14/article_classifier/resolve/main/src/train.py
- Command line
-
hf download hf://spaces/asriel14/article_classifier/src/train.py
-
curl -L -o train.py https://huggingface.co/spaces/asriel14/article_classifier/resolve/main/src/train.py
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() | |