#!/usr/bin/env python3
"""Standalone LoRA training script for the CUDA backend — invoked as a subprocess by
app/services/finetune/cuda_backend.py, never imported by the FastAPI/Celery app itself
(that's deliberate: keeps torch/unsloth, both huge and GPU-specific, out of the main
process's import graph entirely).

Trains against the same train.jsonl/valid.jsonl format
app/services/finetune/shared.py::export_training_data already produces
({"messages": [{"role": "user", ...}, {"role": "assistant", ...}]} per line) — no
export changes needed to support this backend.

Prints one JSON object per line for progress (`{"type": "train"|"val", "iter": N,
"loss": L}`), which cuda_backend.py::parse_progress_line reads directly — no text
parsing needed since this script owns its own output format.

NOTE: this has not been run against real NVIDIA hardware yet (built in an environment
with no GPU available) — the API calls below follow Unsloth/TRL's documented usage,
but exact keyword arguments can drift between library versions. If this errors out,
the traceback printed to stdout is the fastest way to pin down what changed.
"""

import argparse
import json
import sys


def emit(event: dict) -> None:
    print(json.dumps(event), flush=True)


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--model", required=True)
    parser.add_argument("--data", required=True, help="Directory with train.jsonl/valid.jsonl")
    parser.add_argument("--adapter-path", required=True)
    parser.add_argument("--iters", type=int, required=True)
    parser.add_argument("--batch-size", type=int, required=True)
    parser.add_argument("--learning-rate", type=float, required=True)
    parser.add_argument("--lora-r", type=int, default=16)
    parser.add_argument("--lora-alpha", type=int, default=16)
    parser.add_argument("--lora-dropout", type=float, default=0.0)
    parser.add_argument("--steps-per-report", type=int, default=5)
    parser.add_argument("--steps-per-eval", type=int, default=20)
    parser.add_argument("--max-seq-length", type=int, default=2048)
    args = parser.parse_args()

    from datasets import load_dataset
    from transformers import TrainerCallback
    from trl import SFTConfig, SFTTrainer
    from unsloth import FastLanguageModel

    print(f"Loading base model {args.model}...", flush=True)
    model, tokenizer = FastLanguageModel.from_pretrained(
        model_name=args.model,
        max_seq_length=args.max_seq_length,
        dtype=None,  # auto-detect the right compute dtype for this GPU
        load_in_4bit=True,
    )

    model = FastLanguageModel.get_peft_model(
        model,
        r=args.lora_r,
        target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
        lora_alpha=args.lora_alpha,
        lora_dropout=args.lora_dropout,
        bias="none",
        use_gradient_checkpointing="unsloth",
        random_state=3407,
    )

    def to_text(example: dict) -> dict:
        return {"text": tokenizer.apply_chat_template(example["messages"], tokenize=False, add_generation_prompt=False)}

    print("Loading datasets...", flush=True)
    train_ds = load_dataset("json", data_files=f"{args.data}/train.jsonl", split="train").map(to_text)
    valid_ds = load_dataset("json", data_files=f"{args.data}/valid.jsonl", split="train").map(to_text)

    class ProgressCallback(TrainerCallback):
        def on_log(self, _args, state, _control, logs=None, **_kwargs):
            if not logs:
                return
            if "eval_loss" in logs:
                emit({"type": "val", "iter": state.global_step, "loss": round(logs["eval_loss"], 4)})
            elif "loss" in logs:
                emit({"type": "train", "iter": state.global_step, "loss": round(logs["loss"], 4)})

    print("Starting training...", flush=True)
    trainer = SFTTrainer(
        model=model,
        tokenizer=tokenizer,
        train_dataset=train_ds,
        eval_dataset=valid_ds,
        dataset_text_field="text",
        max_seq_length=args.max_seq_length,
        callbacks=[ProgressCallback()],
        args=SFTConfig(
            per_device_train_batch_size=args.batch_size,
            gradient_accumulation_steps=1,
            max_steps=args.iters,
            learning_rate=args.learning_rate,
            logging_steps=args.steps_per_report,
            eval_strategy="steps",
            eval_steps=args.steps_per_eval,
            output_dir=f"{args.adapter_path}/_trainer_output",
            optim="adamw_8bit",
            seed=3407,
            report_to="none",
        ),
    )
    trainer.train()

    print(f"Saving adapter to {args.adapter_path}...", flush=True)
    model.save_pretrained(args.adapter_path)
    tokenizer.save_pretrained(args.adapter_path)
    print("Done.", flush=True)


if __name__ == "__main__":
    try:
        main()
    except Exception as exc:  # noqa: BLE001
        import traceback

        traceback.print_exc()
        print(f"FATAL: {exc}", flush=True)
        sys.exit(1)
