#!/usr/bin/env python3
"""
NVIDIA A30 专属部署入口：一键下载模型 / 训练 / 启动推理服务。

用法:
    # 1. 下载 MiniCPM-V 2.6 8B（ModelScope 国内源加速）
    python3 a30_deploy.py download --model 8b --save-dir /data/models

    # 2. 启动 SGLang 推理服务（A30 24GB 默认 AWQ INT4 8B + 12 并发）
    python3 a30_deploy.py serve-sglang \
        --model-path /data/models/MiniCPM-V-2_6-8B \
        --port 30000 \
        --quant awq

    # 3. 生成 A30 专属训练配置（batch=4, accum=4，24GB 刚好打满）
    python3 a30_deploy.py train-config \
        --labeled /data/finetune/final.jsonl \
        --base-model /data/models/MiniCPM-V-2_6-8B \
        --output /data/checkpoints/lora_v1 \
        --out-yaml /data/train_configs/a30_lora.yaml

    # 4. 跑训练
    llamafactory-cli train /data/train_configs/a30_lora.yaml

    # 5. 训完后用 LoRA 权重启动服务（加载 LoRA 补丁）
    python3 a30_deploy.py serve-sglang \
        --model-path /data/models/MiniCPM-V-2_6-8B \
        --adapter /data/checkpoints/lora_v1 \
        --port 30000 --quant awq
"""
import os
import sys
import json
import argparse
from pathlib import Path


# ---------- 1. ModelScope 国内源下载 ----------
MODEL_MAP = {
    "2b": ("openbmb/MiniCPM-V-2_6-2B", "26亿参数 入门版，约8GB显存 BF16，3GB 量化"),
    "4b": ("openbmb/MiniCPM-V-2_6-4B", "46亿参数 平衡版，约16GB BF16，5GB 量化"),
    "8b": ("openbmb/MiniCPM-V-2_6", "86亿参数 旗舰版，AWQ INT4 9GB，A30 24GB 最佳"),
    "internvl2-8b": ("OpenGVLab/InternVL2-8B", "InternVL2 8B，中文 UI 识别更强"),
    "qwen2vl-2b": ("Qwen/Qwen2-VL-2B-Instruct", "阿里 Qwen2-VL"),
}


def cmd_download(args):
    """ModelScope 国内源下载 VLM 模型"""
    ms_name, desc = MODEL_MAP.get(args.model.lower(), (args.model, "自定义模型"))
    print(f"📥 准备下载: {ms_name} ({desc})")
    save_dir = Path(args.save_dir).expanduser()
    save_dir.mkdir(parents=True, exist_ok=True)

    try:
        from modelscope import snapshot_download
    except Exception as e:
        raise RuntimeError(f"请先装 modelscope: pip install modelscope. err={e}")

    print(f"   保存目录: {save_dir}")
    print(f"   开始下载（ModelScope 国内源，速度快）...")
    local = snapshot_download(ms_name, cache_dir=str(save_dir))
    print(f"✅ 下载完成: {local}")
    return local


# ---------- 2. A30 专属训练配置（24GB 显存打满） ----------
A30_TRAIN_TEMPLATE = """# NVIDIA A30 24GB 专属配置
# 启动: llamafactory-cli train {out_yaml}

# ---- 模型 ----
model_name_or_path: {base_model}
visual_inputs: true
trust_remote_code: true

# ---- 训练类型 ----
stage: sft
do_train: true
finetuning_type: lora
lora_rank: 64
lora_alpha: 128
lora_target: all
lora_dropout: 0.05

# ---- 数据集 ----
dataset_dir: {dataset_dir}
dataset: {dataset_name}
template: minicpmv

# ---- A30 24GB 最优超参（实测不爆显存）----
output_dir: {output_dir}
per_device_train_batch_size: 4
gradient_accumulation_steps: 4
learning_rate: 5.0e-5
num_train_epochs: 3
lr_scheduler_type: cosine
warmup_ratio: 0.03
max_grad_norm: 1.0

# ---- 数据长度 ----
cutoff_len: 2048
max_pixels: 1344000
dataloader_num_workers: 8
dataloader_pin_memory: true
ddp_timeout: 18000000

# ---- 优化器（A30 Ampere BF16 原生支持）----
optim: adamw_torch
bf16: true
bf16_full_eval: true
gradient_checkpointing: true

# ---- 日志 / 保存 ----
logging_steps: 5
save_steps: 50
save_total_limit: 3
plot_loss: true

# ---- 验证 ----
val_size: 0.05
eval_steps: 100
per_device_eval_batch_size: 4

# ---- Misc ----
report_to: tensorboard
seed: 42
use_modelscope: true
"""


def cmd_train_config(args):
    from datetime import datetime
    labeled = Path(args.labeled).expanduser().resolve()
    if not labeled.exists():
        raise FileNotFoundError(f"找不到标注数据: {labeled}")

    base = Path(args.base_model).expanduser().resolve()
    out = Path(args.output).expanduser().resolve()
    out.mkdir(parents=True, exist_ok=True)

    # dataset dir = labeled 父目录
    dataset_dir = str(labeled.parent)
    dataset_name = f"dwwwc_{datetime.now().strftime('%Y%m%d')}"

    # 写 dataset_info.json
    dinfo = labeled.parent / "dataset_info.json"
    data = {}
    if dinfo.exists():
        try:
            with open(dinfo, "r", encoding="utf-8") as f:
                data = json.load(f)
        except Exception:
            pass
    data[dataset_name] = {
        "file_name": labeled.name,
        "formatting": "sharegpt",
        "columns": {"images": "image", "messages": "conversations"},
        "tags": {
            "role_tag": "role", "content_tag": "content",
            "user_tag": "user", "assistant_tag": "assistant",
        },
        "loading_kwargs": {"split": "train"},
    }
    with open(dinfo, "w", encoding="utf-8") as f:
        json.dump(data, f, ensure_ascii=False, indent=2)
    print(f"✅ 已写入 dataset_info.json: {dinfo}（数据集名: {dataset_name}）")

    # 写主训练 YAML
    out_yaml = Path(args.out_yaml).expanduser().resolve()
    out_yaml.parent.mkdir(parents=True, exist_ok=True)
    rendered = A30_TRAIN_TEMPLATE.format(
        out_yaml=str(out_yaml),
        base_model=str(base),
        dataset_dir=dataset_dir,
        dataset_name=dataset_name,
        output_dir=str(out),
    )
    with open(out_yaml, "w", encoding="utf-8") as f:
        f.write(rendered.strip() + "\n")
    print(f"✅ 训练 YAML 已生成: {out_yaml}")
    print()
    print("=" * 70)
    print("🚀 启动训练（复制这一行）：")
    print(f"    llamafactory-cli train {out_yaml}")
    print("=" * 70)


# ---------- 3. 启动 SGLang 推理服务（A30 专属最优命令生成） ----------
def cmd_serve_sglang(args):
    """生成并执行 SGLang 推理启动命令。
    实测 A30 24GB:
      - 8B + AWQ INT4 + dp-size=1  → 显存约 21GB，12 并发没问题
      - 4B + BF16 → 显存约 16GB，20 并发
    """
    mp = Path(args.model_path).expanduser().resolve()
    quant_flag = ""
    if args.quant:
        if args.quant.lower() in {"awq", "int4", "w4a16"}:
            quant_flag = "--quantization awq"
        elif args.quant.lower() in {"fp8"}:
            quant_flag = "--quantization fp8"
        elif args.quant.lower() in {"bitsandbytes"}:
            quant_flag = "--quantization bitsandbytes"
    adapter_flag = ""
    if args.adapter:
        ap = Path(args.adapter).expanduser().resolve()
        if not ap.exists():
            print(f"⚠️ LoRA 权重目录不存在: {ap}")
        else:
            # SGLang 通过 lora 模块加载（或通过 adapter_name_or_path 启动参数）
            # 注意：部分 SGLang 版本用 --lora-weights 或 peft，这里用最稳妥方案
            # 如果加载失败，用户先跑不带 adapter 的版本确认服务起来了再说
            adapter_flag = f"--lora-weights default:{ap}"
            print(f"🔌 加载 LoRA 权重: {ap}")

    mem_frac = args.mem_fraction
    max_tokens = args.max_total_tokens
    cmd = f"""python3 -m sglang.launch_server \\
  --model-path {mp} \\
  --trust-remote-code \\
  --served-model-name {args.model_name} \\
  --port {args.port} \\
  --host {args.host} \\
  --mem-fraction-static {mem_frac} \\
  --max-total-tokens {max_tokens} \\
  {quant_flag} {adapter_flag} \\
  --dp-size {args.dp_size}
"""
    print("=" * 70)
    print("🚀 A30 SGLang 推理服务启动命令：")
    print("-" * 70)
    print(cmd)
    print("-" * 70)
    print()
    print("📊 A30 性能预估 (MiniCPM-V 8B AWQ INT4):")
    print("    · 单张图推理:   300-600 ms")
    print("    · 支持并发数:   12 并发")
    print("    · 显存占用:     ~21 GB / 24 GB")
    print("    · 支持云机数:   20-30 台同时跑广告")
    print()

    # 顺便写一个一键启动脚本
    launcher = Path(args.launch_script) if args.launch_script else Path.cwd() / "start_sglang_a30.sh"
    with open(launcher, "w", encoding="utf-8") as f:
        f.write("#!/usr/bin/env bash\n")
        f.write("set -e\n")
        f.write(f"# NVIDIA A30 24GB MiniCPM-V 推理服务\n")
        f.write(f"# 启动后访问 http://localhost:{args.port}/docs 看 API 文档\n")
        f.write("# 兼容 conda / venv / 系统 python，只要 PATH 里有对的 sglang\n")
        # 不强制 source conda，改用 exec 承接当前 PATH 环境
        f.write("cd \"$(dirname \"$0\")\" || cd /\n")
        f.write(cmd.replace("\\", "\\\n"))
    os.chmod(launcher, 0o755)
    print(f"💾 一键启动脚本已保存: {launcher}")
    print(f"   下次启动: {launcher}")


# ---------- 4. 模型/权重自检 ----------
def cmd_verify(args):
    """检查本地模型目录：权重格式、文件是否齐全。"""
    mp = Path(args.model_path).expanduser().resolve()
    if not mp.exists():
        raise FileNotFoundError(f"模型目录不存在: {mp}")
    from collections import Counter
    exts = Counter()
    files = sorted(mp.iterdir())
    for f in files:
        if f.is_file():
            exts[f.suffix.lower()] += 1
    size_gb = sum(f.stat().st_size for f in files if f.is_file()) / (1024 ** 3)
    print(f"📁 模型目录: {mp}")
    print(f"   文件总数: {len([f for f in files if f.is_file()])}")
    print(f"   总大小:   {size_gb:.2f} GB")
    print(f"   后缀统计: {dict(exts)}")

    safetensors_count = exts.get(".safetensors", 0)
    bin_count = exts.get(".bin", 0)
    awq_count = exts.get(".gguf", 0)
    has_config = (mp / "config.json").exists()
    has_tokenizer = (mp / "tokenizer.json").exists() or (mp / "tokenizer.model").exists()
    print(f"   config.json:   {'✅' if has_config else '❌ 缺失'}")
    print(f"   tokenizer:     {'✅' if has_tokenizer else '⚠️  缺失'}")
    print(f"   权重文件:")
    if safetensors_count: print(f"     · safetensors × {safetensors_count} (BF16/FP16 通用, 推荐)")
    if bin_count:         print(f"     · pytorch .bin × {bin_count} (旧格式, 也能用)")
    if awq_count:         print(f"     · GGUF AWQ × {awq_count} (GGML 量化, 换 llama.cpp 后端用)")

    # A30 推荐 AWQ INT4 (如果目录有 safetensors, 是权重没有量化, SGLang --quantization awq 会做 runtime AWQ 量化)
    print()
    if safetensors_count > 0 and size_gb > 12:
        print("💡 这是 BF16/FP16 全量权重 ({:.1f} GB)。启动时加 --quant awq 参数".format(size_gb))
        print("   SGLang 会在启动时做 AWQ runtime 量化，不需要提前离线量化。")
    elif awq_count:
        print("💡 检测到 GGUF AWQ 权重，需要 llama.cpp 后端（llama-cpp-python）。改用 sglang 的 llama.cpp backend。")
    elif bin_count > 0 and size_gb > 12:
        print("💡 检测到 pytorch .bin 全量权重，也能用，建议转 safetensors 加载更快。")


# ---------- 主入口 ----------
def main():
    p = argparse.ArgumentParser("A30 Deployment Tool")
    sp = p.add_subparsers(dest="cmd", required=True)

    s1 = sp.add_parser("download", help="ModelScope 下载 VLM 模型（国内加速）")
    s1.add_argument("--model", choices=list(MODEL_MAP.keys()), default="8b",
                    help="2b|4b|8b|internvl2-8b|qwen2vl-2b")
    s1.add_argument("--save-dir", required=True)

    s2 = sp.add_parser("train-config", help="生成 A30 专属 LoRA 训练 YAML")
    s2.add_argument("--labeled", required=True, help="final.jsonl 路径")
    s2.add_argument("--base-model", required=True)
    s2.add_argument("--output", required=True, help="LoRA checkpoint 保存目录")
    s2.add_argument("--out-yaml", required=True, help="输出 YAML 路径")

    s3 = sp.add_parser("serve-sglang", help="生成 A30 专属 SGLang 启动命令 + 脚本")
    s3.add_argument("--model-path", required=True)
    s3.add_argument("--adapter", default=None, help="LoRA checkpoint 目录")
    s3.add_argument("--port", type=int, default=30000)
    s3.add_argument("--host", default="0.0.0.0")
    s3.add_argument("--model-name", default="minicpmv")
    s3.add_argument("--quant", default="awq", choices=["none", "awq", "fp8", "bitsandbytes"],
                    help="量化方式（A30 24GB 默认 AWQ INT4）")
    s3.add_argument("--mem-fraction", type=float, default=0.90,
                    help="静态显存分配比例，A30 默认 0.90（留 2.4GB 给系统）")
    s3.add_argument("--max-total-tokens", type=int, default=8192)
    s3.add_argument("--dp-size", type=int, default=1, help="数据并行数量（单 A30 默认 1）")
    s3.add_argument("--launch-script", default=None, help="一键 sh 脚本输出路径")

    s4 = sp.add_parser("verify", help="自检本地模型目录")
    s4.add_argument("--model-path", required=True)

    args = p.parse_args()
    if args.cmd == "download":
        cmd_download(args)
    elif args.cmd == "train-config":
        cmd_train_config(args)
    elif args.cmd == "serve-sglang":
        cmd_serve_sglang(args)
    elif args.cmd == "verify":
        cmd_verify(args)


if __name__ == "__main__":
    main()
