mslovenberg's picture
Upload 2 files
8c59a81 verified
Raw
History Blame Contribute Delete
5.22 kB
import argparse
import json
from pathlib import Path
import numpy as np
import pandas as pd
from datasets import Dataset
from sklearn.metrics import accuracy_score, f1_score
from sklearn.model_selection import train_test_split
from transformers import (
AutoModelForSequenceClassification,
AutoTokenizer,
DataCollatorWithPadding,
Trainer,
TrainingArguments,
set_seed,
)
def read_table(path: Path) -> pd.DataFrame:
if path.suffix.lower() == ".jsonl":
return pd.read_json(path, lines=True)
if path.suffix.lower() == ".json":
return pd.read_json(path)
return pd.read_csv(path)
def first_existing(columns, candidates):
for candidate in candidates:
if candidate in columns:
return candidate
raise ValueError(f"Не нашёл ни одну из колонок: {candidates}")
def prepare_dataframe(path: Path, max_samples: int | None) -> tuple[pd.DataFrame, list[str]]:
df = read_table(path)
title_col = first_existing(df.columns, ["title", "Title"])
abstract_col = first_existing(df.columns, ["abstract", "summary", "Abstract"])
label_col = first_existing(df.columns, ["tag", "categories", "category", "label"])
df = df[[title_col, abstract_col, label_col]].dropna()
df = df.rename(columns={title_col: "title", abstract_col: "abstract", label_col: "label_name"})
df["label_name"] = df["label_name"].astype(str).str.split().str[0]
counts = df["label_name"].value_counts()
valid_labels = counts[counts >= 20].index
df = df[df["label_name"].isin(valid_labels)].copy()
if max_samples is not None:
df = df.sample(min(max_samples, len(df)), random_state=42)
labels = sorted(df["label_name"].unique())
label2id = {label: i for i, label in enumerate(labels)}
df["label"] = df["label_name"].map(label2id)
df["text"] = "Title: " + df["title"].astype(str) + "\nAbstract: " + df["abstract"].astype(str)
return df[["text", "label"]], labels
def compute_metrics(eval_pred):
logits, labels = eval_pred
predictions = np.argmax(logits, axis=-1)
return {
"accuracy": accuracy_score(labels, predictions),
"macro_f1": f1_score(labels, predictions, average="macro"),
}
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--data", type=Path, required=True, help="CSV/JSON/JSONL с title, abstract и tag/categories.")
parser.add_argument("--model-name", default="prajjwal1/bert-tiny")
parser.add_argument("--output-dir", type=Path, default=Path("model"))
parser.add_argument("--max-samples", type=int, default=5000)
parser.add_argument("--epochs", type=float, default=3.0)
parser.add_argument("--batch-size", type=int, default=16)
parser.add_argument("--lr", type=float, default=2e-5)
args = parser.parse_args()
set_seed(42)
df, labels = prepare_dataframe(args.data, args.max_samples)
train_df, valid_df = train_test_split(
df,
test_size=0.15,
random_state=42,
stratify=df["label"],
)
tokenizer = AutoTokenizer.from_pretrained(args.model_name, use_fast=False)
def tokenize(batch):
return tokenizer(batch["text"], truncation=True, max_length=384)
train_ds = Dataset.from_pandas(train_df, preserve_index=False).map(tokenize, batched=True)
valid_ds = Dataset.from_pandas(valid_df, preserve_index=False).map(tokenize, batched=True)
id2label = {i: label for i, label in enumerate(labels)}
label2id = {label: i for i, label in id2label.items()}
model = AutoModelForSequenceClassification.from_pretrained(
args.model_name,
num_labels=len(labels),
id2label=id2label,
label2id=label2id,
)
training_kwargs = {
"output_dir": str(args.output_dir),
"learning_rate": args.lr,
"per_device_train_batch_size": args.batch_size,
"per_device_eval_batch_size": args.batch_size,
"num_train_epochs": args.epochs,
"weight_decay": 0.01,
"load_best_model_at_end": True,
"metric_for_best_model": "macro_f1",
"report_to": "none",
"save_safetensors": False,
}
try:
training_args = TrainingArguments(
evaluation_strategy="epoch",
save_strategy="epoch",
**training_kwargs,
)
except TypeError:
training_args = TrainingArguments(
eval_strategy="epoch",
save_strategy="epoch",
**training_kwargs,
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_ds,
eval_dataset=valid_ds,
tokenizer=tokenizer,
data_collator=DataCollatorWithPadding(tokenizer),
compute_metrics=compute_metrics,
)
trainer.train()
metrics = trainer.evaluate()
args.output_dir.mkdir(parents=True, exist_ok=True)
trainer.save_model(args.output_dir)
tokenizer.save_pretrained(args.output_dir)
with (args.output_dir / "metrics.json").open("w", encoding="utf-8") as f:
json.dump(metrics, f, ensure_ascii=False, indent=2)
print(json.dumps(metrics, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()