# coding=utf-8
# Qwen3-ASR SFT script — single-GPU Colab edition
# Adapted from: https://github.com/QwenLM/Qwen3-ASR
import argparse, os, re, shutil
from dataclasses import dataclass
from typing import Any, Dict, List, Optional

import librosa
import torch
from datasets import load_dataset
from qwen_asr import Qwen3ASRModel
from transformers import (GenerationConfig, Trainer, TrainerCallback,
                          TrainingArguments)


# ── Forward patch ─────────────────────────────────────────────────────────────
def patch_outer_forward(model):
    cls = model.__class__
    if getattr(cls, "_forward_patched", False):
        return
    if not hasattr(model, "thinker") or not hasattr(model.thinker, "forward"):
        raise RuntimeError(
            "Cannot patch forward: model has no .thinker.forward. "
            "Check your qwen-asr version."
        )
    def forward(self, input_ids=None, attention_mask=None, input_features=None,
                feature_attention_mask=None, labels=None, **kwargs):
        return self.thinker.forward(
            input_ids=input_ids, attention_mask=attention_mask,
            input_features=input_features,
            feature_attention_mask=feature_attention_mask,
            labels=labels, **kwargs,
        )
    cls.forward = forward
    cls._forward_patched = True


# ── Checkpoint utils ──────────────────────────────────────────────────────────
_CKPT_RE = re.compile(r"^checkpoint-(\d+)$")

def find_latest_checkpoint(output_dir: str) -> Optional[str]:
    if not output_dir or not os.path.isdir(output_dir):
        return None
    best_step, best_path = None, None
    for name in os.listdir(output_dir):
        m = _CKPT_RE.match(name)
        if not m:
            continue
        step = int(m.group(1))
        path = os.path.join(output_dir, name)
        if os.path.isdir(path) and (best_step is None or step > best_step):
            best_step, best_path = step, path
    return best_path


# ── Audio ─────────────────────────────────────────────────────────────────────
def load_audio(path: str, sr: int = 16000):
    wav, _ = librosa.load(path, sr=sr, mono=True)
    return wav


# ── Preprocessing ─────────────────────────────────────────────────────────────
def build_prefix_messages(prompt: str, audio_array):
    return [
        {"role": "system", "content": prompt or ""},
        {"role": "user",   "content": [{"type": "audio", "audio": audio_array}]},
    ]

def make_preprocess_fn(processor):
    def _preprocess(ex: Dict[str, Any]) -> Dict[str, Any]:
        prompt = ex.get("prompt", "")
        prefix_msgs = build_prefix_messages(prompt, None)
        prefix_text = processor.apply_chat_template(
            [prefix_msgs], add_generation_prompt=True, tokenize=False
        )[0]
        return {
            "prompt":      prompt,
            "audio":       ex["audio"],
            "target":      ex["text"],
            "prefix_text": prefix_text,
        }
    return _preprocess


# ── Collator ──────────────────────────────────────────────────────────────────
@dataclass
class DataCollatorQwen3ASR:
    processor:     Any
    sampling_rate: int = 16000

    def __call__(self, features: List[Dict[str, Any]]) -> Dict[str, torch.Tensor]:
        audio_paths  = [f["audio"]       for f in features]
        prefix_texts = [f["prefix_text"] for f in features]
        targets      = [f["target"]      for f in features]

        eos        = self.processor.tokenizer.eos_token or ""
        full_texts = [p + t + eos for p, t in zip(prefix_texts, targets)]
        audios     = [load_audio(p, self.sampling_rate) for p in audio_paths]

        full_inp   = self.processor(text=full_texts,   audio=audios,
                                    return_tensors="pt", padding=True, truncation=False)
        prefix_inp = self.processor(text=prefix_texts, audio=audios,
                                    return_tensors="pt", padding=True, truncation=False)

        prefix_lens = prefix_inp["attention_mask"].sum(dim=1).tolist()
        labels = full_inp["input_ids"].clone()
        for i, pl in enumerate(prefix_lens):
            labels[i, :pl] = -100           # mask prompt tokens from loss

        pad_id = self.processor.tokenizer.pad_token_id
        if pad_id is not None:
            labels[labels == pad_id] = -100

        full_inp["labels"] = labels
        return full_inp


# ── Trainer ───────────────────────────────────────────────────────────────────
class CastFloatTrainer(Trainer):
    """Cast all float inputs to the model's dtype (handles fp16 / bf16 mismatches)."""
    def _prepare_inputs(self, inputs):
        inputs = super()._prepare_inputs(inputs)
        model_dtype = getattr(self.model, "dtype", None)
        if model_dtype is not None:
            for k, v in list(inputs.items()):
                if torch.is_tensor(v) and v.is_floating_point():
                    inputs[k] = v.to(dtype=model_dtype)
        return inputs


# ── Callback: copy HF config files into every checkpoint ─────────────────────
def copy_hf_config_files(src_dir: str, dst_dir: str):
    os.makedirs(dst_dir, exist_ok=True)
    for fn in [
        "config.json", "generation_config.json", "preprocessor_config.json",
        "processor_config.json", "tokenizer_config.json", "tokenizer.json",
        "special_tokens_map.json", "chat_template.json", "merges.txt", "vocab.json",
    ]:
        src = os.path.join(src_dir, fn)
        if os.path.exists(src):
            shutil.copy2(src, os.path.join(dst_dir, fn))

class MakeCheckpointInferableCallback(TrainerCallback):
    def __init__(self, base_model_path: str):
        self.base_model_path = base_model_path

    def on_save(self, args: TrainingArguments, state, control, **kwargs):
        if args.process_index != 0:
            return control
        ckpt_dir = os.path.join(args.output_dir, f"checkpoint-{state.global_step}")
        copy_hf_config_files(self.base_model_path, ckpt_dir)
        return control


# ── Argument parsing ──────────────────────────────────────────────────────────
def parse_args():
    p = argparse.ArgumentParser("Qwen3-ASR SFT")
    p.add_argument("--model_path",      default="Qwen/Qwen3-ASR-0.6B")
    p.add_argument("--train_file",      required=True)
    p.add_argument("--eval_file",       default="")
    p.add_argument("--output_dir",      default="./qwen3-asr-out")
    p.add_argument("--sr",              type=int,   default=16000)
    p.add_argument("--batch_size",      type=int,   default=4)
    p.add_argument("--grad_acc",        type=int,   default=16)
    p.add_argument("--lr",              type=float, default=2e-5)
    p.add_argument("--epochs",          type=float, default=10)
    p.add_argument("--log_steps",       type=int,   default=20)
    p.add_argument("--warmup_ratio",    type=float, default=0.05)
    p.add_argument("--save_strategy",   default="steps")
    p.add_argument("--save_steps",      type=int,   default=500)
    p.add_argument("--save_total_limit",type=int,   default=100)
    p.add_argument("--num_workers",     type=int,   default=4)
    p.add_argument("--resume_from",     default="")
    p.add_argument("--resume",          type=int,   default=0)
    return p.parse_args()


# ── Main ──────────────────────────────────────────────────────────────────────
def main():
    args = parse_args()

    use_bf16 = (
        torch.cuda.is_available()
        and torch.cuda.get_device_capability(0)[0] >= 8
    )
    print(f"Using {'bf16' if use_bf16 else 'fp16'}")

    asr_wrapper = Qwen3ASRModel.from_pretrained(
        args.model_path,
        dtype=torch.bfloat16 if use_bf16 else torch.float16,
        device_map=None,
    )
    model     = asr_wrapper.model
    processor = asr_wrapper.processor

    patch_outer_forward(model)
    model.generation_config = GenerationConfig.from_model_config(model.config)

    # ── Dataset ───────────────────────────────────────────────────────────────
    data_files = {"train": args.train_file}
    if args.eval_file:
        data_files["validation"] = args.eval_file

    raw_ds = load_dataset("json", data_files=data_files)
    ds     = raw_ds.map(make_preprocess_fn(processor), num_proc=1)

    keep = {"prompt", "audio", "target", "prefix_text"}
    for split in ds.keys():
        drop = [c for c in ds[split].column_names if c not in keep]
        if drop:
            ds[split] = ds[split].remove_columns(drop)

    # ── Training args ─────────────────────────────────────────────────────────
    training_args = TrainingArguments(
        output_dir                  = args.output_dir,
        per_device_train_batch_size = args.batch_size,
        gradient_accumulation_steps = args.grad_acc,
        learning_rate               = args.lr,
        num_train_epochs            = args.epochs,
        logging_steps               = args.log_steps,
        lr_scheduler_type           = "cosine",
        warmup_ratio                = args.warmup_ratio,
        dataloader_num_workers      = args.num_workers,
        save_strategy               = args.save_strategy,
        save_steps                  = args.save_steps,
        save_total_limit            = args.save_total_limit,
        save_safetensors            = True,
        eval_strategy               = "steps" if args.eval_file else "no",
        eval_steps                  = args.save_steps if args.eval_file else None,
        do_eval                     = bool(args.eval_file),
        bf16                        = use_bf16,
        fp16                        = not use_bf16,
        ddp_find_unused_parameters  = False,
        remove_unused_columns       = False,
        report_to                   = "none",
    )

    collator = DataCollatorQwen3ASR(processor=processor, sampling_rate=args.sr)

    trainer = CastFloatTrainer(
        model         = model,
        args          = training_args,
        train_dataset = ds["train"],
        eval_dataset  = ds.get("validation", None),
        data_collator = collator,
        tokenizer     = processor.tokenizer,
        callbacks     = [MakeCheckpointInferableCallback(args.model_path)],
    )

    resume = (args.resume_from or "").strip()
    if not resume and args.resume == 1:
        resume = find_latest_checkpoint(training_args.output_dir) or ""
    if resume:
        print(f"Resuming from: {resume}")
        trainer.train(resume_from_checkpoint=resume)
    else:
        trainer.train()


if __name__ == "__main__":
    main()
