#!/usr/bin/env python3
"""Fine-tune a transformer model for text classification.

This script provides a practical, vendor-neutral workflow for:
- loading labeled text data from CSV or JSONL
- validating schema and label consistency
- creating train/validation/test splits with leakage checks
- tokenizing text for transformer training
- fine-tuning a Hugging Face transformer classifier
- evaluating the model and saving artifacts

Expected input data format:
- text: the input text to classify
- label: the target class label
- optional id: a stable record identifier

Examples:
    python fine_tune_text_classifier.py \
        --data-path data/train.csv \
        --text-column text \
        --label-column label \
        --model-name distilbert-base-uncased \
        --output-dir outputs/text-classifier

    python fine_tune_text_classifier.py \
        --data-path data/train.jsonl \
        --format jsonl \
        --max-length 256 \
        --epochs 3
"""

from __future__ import annotations

import argparse
import json
import logging
import os
import random
from dataclasses import dataclass
from pathlib import Path
from typing import Dict, List, Optional, Sequence, Tuple

import numpy as np
import pandas as pd
from sklearn.metrics import accuracy_score, classification_report, f1_score
from sklearn.model_selection import train_test_split

try:
    from datasets import Dataset, DatasetDict
    from transformers import (
        AutoModelForSequenceClassification,
        AutoTokenizer,
        DataCollatorWithPadding,
        Trainer,
        TrainingArguments,
        set_seed,
    )
except ImportError as exc:
    raise SystemExit(
        "Missing dependencies. Install: pip install transformers datasets torch scikit-learn pandas numpy"
    ) from exc


LOGGER = logging.getLogger("fine_tune_text_classifier")


@dataclass
class Config:
    data_path: Path
    output_dir: Path
    model_name: str
    text_column: str
    label_column: str
    id_column: Optional[str]
    format: str
    max_length: int
    test_size: float
    validation_size: float
    seed: int
    epochs: float
    batch_size: int
    learning_rate: float
    weight_decay: float
    max_samples: Optional[int]
    remove_duplicates: bool
    do_lower_case: bool


def parse_args() -> Config:
    parser = argparse.ArgumentParser(
        description="Fine-tune a transformer model for text classification."
    )
    parser.add_argument("--data-path", required=True, help="Path to CSV or JSONL data file.")
    parser.add_argument("--output-dir", default="outputs/text-classifier", help="Directory for model and metrics.")
    parser.add_argument(
        "--model-name",
        default="distilbert-base-uncased",
        help="Pretrained transformer model name or local path.",
    )
    parser.add_argument("--text-column", default="text", help="Name of the text column.")
    parser.add_argument("--label-column", default="label", help="Name of the label column.")
    parser.add_argument("--id-column", default=None, help="Optional stable record ID column.")
    parser.add_argument("--format", choices=["csv", "jsonl"], default="csv", help="Input file format.")
    parser.add_argument("--max-length", type=int, default=256, help="Maximum token length.")
    parser.add_argument("--test-size", type=float, default=0.15, help="Test split fraction.")
    parser.add_argument("--validation-size", type=float, default=0.15, help="Validation split fraction of full dataset.")
    parser.add_argument("--seed", type=int, default=42, help="Random seed.")
    parser.add_argument("--epochs", type=float, default=3.0, help="Training epochs.")
    parser.add_argument("--batch-size", type=int, default=8, help="Per-device batch size.")
    parser.add_argument("--learning-rate", type=float, default=2e-5, help="Learning rate.")
    parser.add_argument("--weight-decay", type=float, default=0.01, help="Weight decay.")
    parser.add_argument("--max-samples", type=int, default=None, help="Optional cap for quick experiments.")
    parser.add_argument(
        "--remove-duplicates",
        action="store_true",
        help="Remove duplicate text rows before splitting.",
    )
    parser.add_argument(
        "--do-lower-case",
        action="store_true",
        help="Lowercase text before training if appropriate for your model and use case.",
    )

    args = parser.parse_args()
    return Config(
        data_path=Path(args.data_path),
        output_dir=Path(args.output_dir),
        model_name=args.model_name,
        text_column=args.text_column,
        label_column=args.label_column,
        id_column=args.id_column,
        format=args.format,
        max_length=args.max_length,
        test_size=args.test_size,
        validation_size=args.validation_size,
        seed=args.seed,
        epochs=args.epochs,
        batch_size=args.batch_size,
        learning_rate=args.learning_rate,
        weight_decay=args.weight_decay,
        max_samples=args.max_samples,
        remove_duplicates=args.remove_duplicates,
        do_lower_case=args.do_lower_case,
    )


def setup_logging() -> None:
    logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")


def load_data(path: Path, fmt: str) -> pd.DataFrame:
    if not path.exists():
        raise FileNotFoundError(f"Data file not found: {path}")
    if fmt == "csv":
        df = pd.read_csv(path)
    else:
        df = pd.read_json(path, lines=True)
    if df.empty:
        raise ValueError("Input data is empty.")
    return df


def validate_columns(df: pd.DataFrame, text_column: str, label_column: str, id_column: Optional[str]) -> None:
    missing = [col for col in [text_column, label_column] if col not in df.columns]
    if missing:
        raise ValueError(f"Missing required columns: {missing}")
    if id_column and id_column not in df.columns:
        raise ValueError(f"ID column not found: {id_column}")

    if df[text_column].isna().any():
        raise ValueError(f"Text column '{text_column}' contains null values.")
    if df[label_column].isna().any():
        raise ValueError(f"Label column '{label_column}' contains null values.")

    text_lengths = df[text_column].astype(str).str.len()
    if text_lengths.min() == 0:
        raise ValueError("At least one text row is empty.")


def normalize_data(df: pd.DataFrame, config: Config) -> pd.DataFrame:
    df = df.copy()
    df[config.text_column] = df[config.text_column].astype(str)
    df[config.label_column] = df[config.label_column].astype(str)

    if config.do_lower_case:
        df[config.text_column] = df[config.text_column].str.lower()

    if config.remove_duplicates:
        subset_cols = [config.text_column]
        if config.id_column:
            subset_cols.append(config.id_column)
        before = len(df)
        df = df.drop_duplicates(subset=subset_cols).reset_index(drop=True)
        LOGGER.info("Removed %d duplicate rows.", before - len(df))

    if config.max_samples is not None:
        if config.max_samples <= 0:
            raise ValueError("--max-samples must be greater than zero when provided.")
        df = df.sample(n=min(config.max_samples, len(df)), random_state=config.seed).reset_index(drop=True)

    return df


def make_label_map(labels: Sequence[str]) -> Tuple[Dict[str, int], Dict[int, str]]:
    unique_labels = sorted(set(labels))
    if len(unique_labels) < 2:
        raise ValueError("Need at least two distinct labels for classification.")
    label2id = {label: idx for idx, label in enumerate(unique_labels)}
    id2label = {idx: label for label, idx in label2id.items()}
    return label2id, id2label


def split_data(df: pd.DataFrame, config: Config) -> Tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:
    if not 0 < config.test_size < 1:
        raise ValueError("--test-size must be between 0 and 1.")
    if not 0 < config.validation_size < 1:
        raise ValueError("--validation-size must be between 0 and 1.")
    if config.test_size + config.validation_size >= 1:
        raise ValueError("Test and validation fractions must leave room for training data.")

    stratify = df[config.label_column] if df[config.label_column].nunique() > 1 else None
    train_val_df, test_df = train_test_split(
        df,
        test_size=config.test_size,
        random_state=config.seed,
        stratify=stratify,
    )

    remaining_fraction = 1 - config.test_size
    adjusted_val_size = config.validation_size / remaining_fraction
    stratify_tv = train_val_df[config.label_column] if train_val_df[config.label_column].nunique() > 1 else None
    train_df, val_df = train_test_split(
        train_val_df,
        test_size=adjusted_val_size,
        random_state=config.seed,
        stratify=stratify_tv,
    )
    return train_df.reset_index(drop=True), val_df.reset_index(drop=True), test_df.reset_index(drop=True)


def check_for_leakage(train_df: pd.DataFrame, test_df: pd.DataFrame, text_column: str) -> None:
    overlap = set(train_df[text_column]).intersection(set(test_df[text_column]))
    if overlap:
        raise ValueError(f"Potential leakage detected: {len(overlap)} identical texts appear in both train and test.")


def to_hf_dataset(df: pd.DataFrame, label2id: Dict[str, int], text_column: str, label_column: str) -> Dataset:
    mapped = df.copy()
    mapped["labels"] = mapped[label_column].map(label2id)
    if mapped["labels"].isna().any():
        raise ValueError("Found labels not present in label map.")
    return Dataset.from_pandas(mapped[[text_column, "labels"]], preserve_index=False)


def tokenize_datasets(dataset_dict: DatasetDict, tokenizer, text_column: str, max_length: int) -> DatasetDict:
    def tokenize_batch(batch):
        return tokenizer(
            batch[text_column],
            padding="max_length",
            truncation=True,
            max_length=max_length,
        )

    return dataset_dict.map(tokenize_batch, batched=True)


def compute_metrics(pred) -> Dict[str, float]:
    logits, labels = pred
    predictions = np.argmax(logits, axis=-1)
    return {
        "accuracy": accuracy_score(labels, predictions),
        "macro_f1": f1_score(labels, predictions, average="macro", zero_division=0),
    }


def save_label_map(output_dir: Path, label2id: Dict[str, int], id2label: Dict[int, str]) -> None:
    payload = {
        "label2id": label2id,
        "id2label": {str(k): v for k, v in id2label.items()},
    }
    with (output_dir / "label_map.json").open("w", encoding="utf-8") as f:
        json.dump(payload, f, indent=2, sort_keys=True)


def main() -> None:
    setup_logging()
    config = parse_args()
    set_seed(config.seed)
    random.seed(config.seed)
    np.random.seed(config.seed)

    config.output_dir.mkdir(parents=True, exist_ok=True)

    LOGGER.info("Loading data from %s", config.data_path)
    df = load_data(config.data_path, config.format)
    validate_columns(df, config.text_column, config.label_column, config.id_column)
    df = normalize_data(df, config)

    LOGGER.info("Preparing label map")
    label2id, id2label = make_label_map(df[config.label_column].tolist())
    LOGGER.info("Labels: %s", ", ".join(label2id.keys()))

    train_df, val_df, test_df = split_data(df, config)
    check_for_leakage(train_df, test_df, config.text_column)

    LOGGER.info("Loading tokenizer and model: %s", config.model_name)
    tokenizer = AutoTokenizer.from_pretrained(config.model_name)
    model = AutoModelForSequenceClassification.from_pretrained(
        config.model_name,
        num_labels=len(label2id),
        label2id=label2id,
        id2label={str(k): v for k, v in id2label.items()},
    )

    ds = DatasetDict(
        {
            "train": to_hf_dataset(train_df, label2id, config.text_column, config.label_column),
            "validation": to_hf_dataset(val_df, label2id, config.text_column, config.label_column),
            "test": to_hf_dataset(test_df, label2id, config.text_column, config.label_column),
        }
    )
    ds = tokenize_datasets(ds, tokenizer, config.text_column, config.max_length)
    ds = ds.remove_columns([config.text_column])

    data_collator = DataCollatorWithPadding(tokenizer=tokenizer)
    training_args = TrainingArguments(
        output_dir=str(config.output_dir),
        learning_rate=config.learning_rate,
        per_device_train_batch_size=config.batch_size,
        per_device_eval_batch_size=config.batch_size,
        num_train_epochs=config.epochs,
        weight_decay=config.weight_decay,
        evaluation_strategy="epoch",
        save_strategy="epoch",
        load_best_model_at_end=True,
        metric_for_best_model="macro_f1",
        greater_is_better=True,
        seed=config.seed,
        logging_dir=str(config.output_dir / "logs"),
        report_to=[],
    )

    trainer = Trainer(
        model=model,
        args=training_args,
        train_dataset=ds["train"],
        eval_dataset=ds["validation"],
        tokenizer=tokenizer,
        data_collator=data_collator,
        compute_metrics=compute_metrics,
    )

    LOGGER.info("Starting training")
    trainer.train()

    LOGGER.info("Evaluating on test set")
    test_metrics = trainer.evaluate(ds["test"])
    predictions = trainer.predict(ds["test"])
    y_true = predictions.label_ids
    y_pred = np.argmax(predictions.predictions, axis=-1)
    report = classification_report(y_true, y_pred, target_names=[id2label[i] for i in range(len(id2label))], zero_division=0)

    trainer.save_model(str(config.output_dir))
    tokenizer.save_pretrained(str(config.output_dir))
    save_label_map(config.output_dir, label2id, id2label)

    metrics_path = config.output_dir / "metrics.json"
    with metrics_path.open("w", encoding="utf-8") as f:
        json.dump(
            {
                "test_metrics": test_metrics,
                "classification_report": report,
                "config": {
                    "model_name": config.model_name,
                    "max_length": config.max_length,
                    "seed": config.seed,
                    "epochs": config.epochs,
                    "batch_size": config.batch_size,
                    "learning_rate": config.learning_rate,
                    "weight_decay": config.weight_decay,
                },
            },
            f,
            indent=2,
        )

    LOGGER.info("Training complete")
    LOGGER.info("Artifacts saved to %s", config.output_dir)
    LOGGER.info("Test metrics: %s", test_metrics)
    print(report)


if __name__ == "__main__":
    main()