#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""提交 YuE2 文生音乐工作流到 GPU 机 ComfyUI，轮询进度并回传产物。

用法示例：
  python3 yue2-run.py --cot off --sem-mx 3000 --tag smoke
"""
from __future__ import annotations
import argparse, json, os, sys, time, urllib.request, urllib.error, shutil

HOST = "http://192.168.31.31:8189"
OUTDIR = "/home/zyw/Downloads/dl-hub/18-YuE2音乐生成"
OPENER = urllib.request.build_opener(urllib.request.ProxyHandler({}))  # 绕过本机 socks 代理


def api(path: str, data=None, timeout=60):
    url = HOST + path
    body = json.dumps(data).encode() if data is not None else None
    req = urllib.request.Request(url, data=body,
                                headers={"Content-Type": "application/json"} if body else {})
    with OPENER.open(req, timeout=timeout) as r:
        raw = r.read()
    try:
        return json.loads(raw)
    except Exception:
        return raw


def build(style, lyrics, cot, seed, sem_mx, abc_mx, budget, tag, ode_steps=32,
          vae_decode="tiled", tile=384, offload_ar=True, attn="auto"):
    return {
        "1": {"class_type": "YuE2Loader", "inputs": {
            "model": "YuE2-3B", "vae": "YuE2-Vae", "device": "cuda",
            "memory_budget_gib": budget, "offload_ar": offload_ar, "offline": True,
            "attention_backend": attn, "quantization": "none"}},
        "2": {"class_type": "YuE2Sampler", "inputs": {
            "pipeline": ["1", 0], "style": style, "lyrics": lyrics, "cot": cot,
            "seed": seed, "cfg_scale": -1.0, "ode_steps": ode_steps,
            "abc_max_tokens": abc_mx, "semantic_max_tokens": sem_mx,
            "abc_temperature": 0.7, "semantic_temperature": 1.0,
            "save_flac": True, "save_abc": True,
            "vae_decode": vae_decode, "vae_tile_frames": tile}},
        "3": {"class_type": "SaveAudio", "inputs": {
            "audio": ["2", 0], "filename_prefix": f"YuE2/{tag}"}},
    }


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--style", default="Mandarin pop ballad, warm female vocal, clean electric "
                                     "guitar, soft piano, restrained drums, intimate, hopeful night mood")
    p.add_argument("--lyrics-file", default=None)
    p.add_argument("--lyrics", default=None)
    p.add_argument("--cot", default="off", choices=["full", "melody", "off"])
    p.add_argument("--seed", type=int, default=831001)
    p.add_argument("--sem-mx", type=int, default=3000, help="semantic_max_tokens（越大歌越长）")
    p.add_argument("--abc-mx", type=int, default=4096)
    p.add_argument("--budget", type=int, default=16)
    p.add_argument("--ode-steps", type=int, default=32)
    p.add_argument("--vae-decode", default="tiled")
    p.add_argument("--tile", type=int, default=384)
    p.add_argument("--tag", default="run")
    p.add_argument("--attn", default="auto",
                   choices=["auto", "sdpa", "cudnn", "external-flash", "flash"],
                   help="AR attention 后端；本机 cuDNN 建不出执行计划，用 sdpa")
    p.add_argument("--no-download", action="store_true")
    a = p.parse_args()

    if a.lyrics_file:
        lyrics = open(a.lyrics_file, encoding="utf-8").read().strip()
    elif a.lyrics:
        lyrics = a.lyrics
    else:
        lyrics = ("[verse]\n路灯把影子拉得很长\n我数着脚步走回家\n"
                  "[chorus]\n把心事交给晚风吧\n明天会有新的回答")

    prompt = build(a.style, lyrics, a.cot, a.seed, a.sem_mx, a.abc_mx,
                   a.budget, a.tag, a.ode_steps, a.vae_decode, a.tile, attn=a.attn)

    print(f"[提交] cot={a.cot} seed={a.seed} sem_max_tokens={a.sem_mx} budget={a.budget}GiB "
          f"tile={a.tile} attn={a.attn} tag={a.tag}", flush=True)
    print(f"[歌词] {lyrics[:100].replace(chr(10), ' / ')}", flush=True)
    res = api("/prompt", {"prompt": prompt, "client_id": "dsh-yue2"})
    if "prompt_id" not in res:
        print("!! 提交失败:", json.dumps(res, ensure_ascii=False)[:2000])
        sys.exit(1)
    pid = res["prompt_id"]
    print(f"[排队] prompt_id={pid}", flush=True)

    t0 = time.time()
    last = 0
    while True:
        time.sleep(10)
        el = time.time() - t0
        try:
            q = api("/queue")
            running = len(q.get("queue_running", []))
            pending = len(q.get("queue_pending", []))
        except Exception:
            running = pending = -1
        h = api(f"/history/{pid}")
        if pid in h:
            rec = h[pid]
            status = rec.get("status", {})
            print(f"\n[完成] 用时 {el/60:.1f} 分钟 status={status.get('status_str')} "
                  f"completed={status.get('completed')}", flush=True)
            if status.get("messages"):
                for m in status["messages"]:
                    if m[0] in ("execution_error", "execution_interrupted"):
                        print("!! 错误:", json.dumps(m[1], ensure_ascii=False)[:1500])
            outs = rec.get("outputs", {})
            print("[产物]", json.dumps(outs, ensure_ascii=False)[:1200], flush=True)
            if not a.no_download:
                os.makedirs(OUTDIR, exist_ok=True)
                n = 0
                for nid, o in outs.items():
                    for key in ("audio", "images", "gifs"):
                        for item in o.get(key, []) or []:
                            fn = item.get("filename")
                            if not fn:
                                continue
                            sub = item.get("subfolder", "")
                            url = (f"/view?filename={urllib.parse.quote(fn)}"
                                   f"&subfolder={urllib.parse.quote(sub)}&type={item.get('type','output')}")
                            dst = os.path.join(OUTDIR, f"{a.tag}_{fn}")
                            with OPENER.open(HOST + url, timeout=300) as r, open(dst, "wb") as w:
                                shutil.copyfileobj(r, w)
                            print(f"[下载] {dst} {os.path.getsize(dst)/1e6:.1f} MB", flush=True)
                            n += 1
                print(f"[下载完成] {n} 个文件 -> {OUTDIR}", flush=True)
            return
        print(f"  等待中… {el/60:.1f} 分钟 (running={running} pending={pending})", flush=True)


if __name__ == "__main__":
    main()
