"""Merge the LoRA adapter into Gemma 4 and quantize to int4 for serving.

    python scripts/merge_and_quantize.py --config config/config.yaml     # (GPU required)

Steps:
  1. Load the 4-bit base + LoRA adapter via Unsloth and merge to a standalone 16-bit model.
  2. Quantize per model.quantization:
       - w4a16 -> compressed-tensors int4 via llm-compressor  (recommended: day-0 vLLM support,
                  same format as Gemma 4's official QAT weights)
       - gguf  -> Q4_0 GGUF via Unsloth  (for Ollama / llama.cpp; Q4_0 matches Gemma 4 QAT)
       - awq   -> AutoAWQ int4
  3. Write to model.quantized_dir.

int4 serving is the biggest inference-time energy win: an E4B model in int4 needs only a few GB of
weights. Tip: if you don't need custom fine-tuned behavior, you can skip training entirely and
serve Google's official QAT checkpoint (google/gemma-4-E4B-it-qat-*), which is already w4a16.
"""
from __future__ import annotations

import argparse
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))

from agent.settings import load_config  # noqa: E402


def load_merged(base_id: str, adapter_dir: str, max_seq_len: int):
    """Load base + adapter via Unsloth (returns model, tokenizer ready to export)."""
    from unsloth import FastModel

    print("Loading base + LoRA adapter (Unsloth)...")
    model, tokenizer = FastModel.from_pretrained(
        model_name=adapter_dir,          # Unsloth resolves the base from the adapter config
        max_seq_length=max_seq_len,
        load_in_4bit=False,              # load in 16-bit so we can merge cleanly
    )
    return model, tokenizer


def export_gguf(model, tokenizer, out_dir: str):
    print("Exporting merged model to GGUF (Q4_0) for Ollama/llama.cpp...")
    model.save_pretrained_gguf(out_dir, tokenizer, quantization_method="q4_0")
    print(f"GGUF -> {out_dir}")
    print("Serve with Ollama: create a Modelfile `FROM ./<model>.gguf` then `ollama create`. ")


def export_merged_16bit(model, tokenizer, merged_dir: str):
    print("Merging LoRA into base (16-bit)...")
    model.save_pretrained_merged(merged_dir, tokenizer, save_method="merged_16bit")
    print(f"Merged 16-bit -> {merged_dir}")
    return merged_dir


def quantize_w4a16(merged_dir: str, out_dir: str, max_seq_len: int):
    """compressed-tensors W4A16 via llm-compressor (GPTQ-style), served by vLLM."""
    from llmcompressor.modifiers.quantization import GPTQModifier
    from llmcompressor.transformers import oneshot

    print("Quantizing to W4A16 (compressed-tensors) via llm-compressor...")
    recipe = GPTQModifier(
        targets="Linear",
        scheme="W4A16",
        # Don't quantize the LM head or the vision/audio towers of the E-series model.
        ignore=["lm_head", "re:.*vision_tower.*", "re:.*audio_tower.*", "re:.*multi_modal.*"],
    )
    oneshot(
        model=merged_dir,
        dataset="open_platypus",       # small calibration set; swap for your own domain text
        recipe=recipe,
        output_dir=out_dir,
        max_seq_length=max_seq_len,
        num_calibration_samples=256,
    )
    print(f"W4A16 model -> {out_dir}")


def quantize_awq(merged_dir: str, out_dir: str):
    from awq import AutoAWQForCausalLM
    from transformers import AutoTokenizer

    print("Quantizing with AWQ (int4)...")
    model = AutoAWQForCausalLM.from_pretrained(merged_dir)
    tokenizer = AutoTokenizer.from_pretrained(merged_dir)
    model.quantize(tokenizer, quant_config={"zero_point": True, "q_group_size": 128, "w_bit": 4, "version": "GEMM"})
    model.save_quantized(out_dir)
    tokenizer.save_pretrained(out_dir)
    print(f"AWQ int4 model -> {out_dir}")


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--config", default="config/config.yaml")
    args = ap.parse_args()
    cfg = load_config(args.config)

    base_id = cfg.get("model.base_id")
    adapter_dir = str(cfg.resolve("train.output_adapter"))
    out_dir = str(cfg.resolve("model.quantized_dir"))
    method = str(cfg.get("model.quantization", "w4a16")).lower()
    max_seq_len = int(cfg.get("model.max_model_len", 4096))

    model, tokenizer = load_merged(base_id, adapter_dir, max_seq_len)

    if method == "gguf":
        export_gguf(model, tokenizer, out_dir)
        return

    merged_dir = str(cfg.resolve("model.quantized_dir").parent / "model-merged-16bit")
    export_merged_16bit(model, tokenizer, merged_dir)

    if method == "w4a16":
        quantize_w4a16(merged_dir, out_dir, max_seq_len)
        print("\nServe with vLLM:")
        print(f"  python -m vllm.entrypoints.openai.api_server --model {out_dir} "
              f"--max-model-len {max_seq_len} --port 8001")
    elif method == "awq":
        quantize_awq(merged_dir, out_dir)
        print(f"\nServe with vLLM: ... --model {out_dir} --quantization awq --port 8001")
    else:
        raise SystemExit(f"Unknown quantization method: {method}")


if __name__ == "__main__":
    main()
