#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""把 YuE2 示例工作流改造成本机(RTX 5060 Ti 16GB)可用版，并装进 ComfyUI 的 Workflows 侧边栏。

关键改动（否则 16GB 卡会崩）：
  1) YuE2MemoryPreset.vram            : "48 GB" -> "16 GB"
  2) YuE2Loader.memory_budget_gib     : 24 -> 16
  3) YuE2Loader.offload_ar            : False -> True
  4) YuE2Loader.attention_backend     : "auto" -> "sdpa"   (auto 会转 cuDNN，本机 cuDNN 建不出执行计划)
  5) YuE2Sampler.vae_decode           : -> "tiled"
  6) YuE2Sampler.vae_tile_frames      : -> 384
并在每个工作流里放一个 Note 节点写清原因。
"""
from __future__ import annotations
import json, os, pathlib, shutil, sys

SRC = pathlib.Path("/home/zyw/ComfyUI/custom_nodes/ComfyUI-YuE2/example_workflows")
DST = pathlib.Path("/home/zyw/ComfyUI/user/default/workflows/YuE2")

NOTE = """【本机适配版 · RTX 5060 Ti 16GB】

已按 16GB 显存预调好，别改这三处：
  • YuE2MemoryPreset = 16 GB
  • YuE2Loader.offload_ar = True
  • YuE2Loader.attention_backend = sdpa
    （auto 会切 cuDNN SDPA，本机 cuDNN 报
      "No valid execution plans built" 必崩）
  • YuE2Sampler.vae_decode = tiled, tile=384

实测（本机）：
  cot=off  63秒歌 → 66秒，峰值显存 9.26 GiB
  cot=full 172秒歌 → 222秒，峰值显存 10.05 GiB

显存档位对应（fork 的 Memory Preset）：
  12GB→offload+tiled/256  16GB→offload+tiled/384
  24/32GB→不offload+tiled 48GB+→不offload+full
用之前的 48GB 档在 16GB 卡上就会"爆显存"。

用参考音频(SheetSage2)前，先把音频放进
ComfyUI/input/ 目录，再在 LoadAudio 里选。
"""


def find(nodes, t):
    return [n for n in nodes if n.get("type") == t]


def patch(doc, *, name, chinese_lyrics=None, chinese_style=None, note=True):
    nodes = doc.get("nodes", [])
    log = [f"--- {name} ---"]

    for n in find(nodes, "YuE2MemoryPreset"):
        wv = n.setdefault("widgets_values", [])
        if wv:
            log.append(f"  MemoryPreset[{n['id']}]: {wv[0]} -> 16 GB")
            wv[0] = "16 GB"

    for n in find(nodes, "YuE2Loader"):
        wv = n.setdefault("widgets_values", [])
        if len(wv) >= 8:
            log.append(f"  Loader[{n['id']}]: budget {wv[3]}->16, offload_ar {wv[4]}->True, "
                       f"attn {wv[6]!r}->'sdpa'")
            wv[3] = 16
            wv[4] = True
            wv[6] = "sdpa"

    for n in find(nodes, "YuE2Sampler"):
        wv = n.setdefault("widgets_values", [])
        if len(wv) >= 14:
            log.append(f"  Sampler[{n['id']}]: vae_decode {wv[12]!r}->'tiled', tile {wv[13]}->384")
            wv[12] = "tiled"
            wv[13] = 384
        if chinese_style is not None and len(wv) >= 1:
            wv[0] = chinese_style
        if chinese_lyrics is not None and len(wv) >= 2:
            wv[1] = chinese_lyrics

    if note:
        nid = max([n.get("id", 0) for n in nodes] + [0]) + 1
        xs = [n.get("pos", [0, 0])[0] for n in nodes] or [0]
        ys = [n.get("pos", [0, 0])[1] for n in nodes] or [0]
        nodes.append({
            "id": nid, "type": "Note", "pos": [min(xs), min(ys) - 320], "size": [520, 320],
            "flags": {}, "order": 0, "mode": 0, "inputs": [], "outputs": [],
            "widgets_values": [NOTE], "color": "#432", "bgcolor": "#653",
        })
        log.append(f"  加说明 Note[{nid}]")
    return log


def main():
    DST.mkdir(parents=True, exist_ok=True)
    zh_lyrics = open("/home/zyw/yue2-lyrics-zh.txt", encoding="utf-8").read().strip()
    zh_style = ("Mandarin pop ballad, warm female vocal, clean electric guitar, soft piano, "
                "restrained drums, intimate delivery, hopeful nocturnal city mood, 78 BPM")
    lines = []
    plan = [
        ("YuE2_Text_to_Song.json",   "YuE2-1-文生音乐-中文-16GB.json",  dict(chinese_lyrics=zh_lyrics, chinese_style=zh_style)),
        ("YuE2_Text_to_Song.json",   "YuE2-2-文生音乐-英文示例-16GB.json", {}),
        ("YuE2_Reference_Remix.json", "YuE2-3-参考音频改编-16GB.json", {}),
        ("YuE2_All_Nodes_Showcase.json", "YuE2-4-全节点总览-16GB.json", {}),
    ]
    for src, dst, kw in plan:
        doc = json.loads((SRC / src).read_text(encoding="utf-8"))
        lines += patch(doc, name=dst, **kw)
        out = DST / dst
        out.write_text(json.dumps(doc, ensure_ascii=False, indent=2), encoding="utf-8")
        lines.append(f"  写入 {out} ({out.stat().st_size} 字节)")
    print("\n".join(lines))
    print("\n目录内容:")
    for p in sorted(DST.iterdir()):
        print("  ", p.name, p.stat().st_size)


if __name__ == "__main__":
    main()
