#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""h3_official_shot.py — 用**官方 MinimaxH3Director 工作流**（无任何加速节点/LoRA）逐镜生成 MV 片段。

严格遵循的三条运行纪律（用户 2026-09-25 明确要求）：
  1. 只用官方工作流（MiniMaxH3Director + ref2va/fl2va 官方权重 + 25 步 res_multistep/simple），
     **绝不使用** Director-加速版 / 二采加速 / DSW V18 等含 MiniMaxH3MemoryEfficientSageAttentionPatch
     或 turbo LoRA 的加速链路（那是 GPU 掉总线的元凶）。
  2. 单镜时长 ≤ 11 秒（帧数取 17k+5 网格）。
  3. 每个生成任务结束后强制冷却 90 秒，让 GPU 释放显存与降温。

用法:
  python3 h3_official_shot.py --shot S02 [--size 960x544] [--dur 5.2] [--seed 123] \
      [--prompt-file <文件> | --keyframe <图>] [--extra-ref <图>] [--dry-run] [--cooldown 90]
"""
import argparse, json, os, shutil, subprocess, sys, time
from pathlib import Path

sys.path.insert(0, "/home/zyw/Downloads/dasiwa-h3-research")
import v18_pipeline as p  # 仅复用上传/提交/轮询/下载这几个工具函数（不碰它的加速模板）

ROOT = Path("/home/zyw/Downloads/dl-hub/54-九九八十一MV")
OUTBASE = ROOT / "06-H3片段"
OUTBASE.mkdir(parents=True, exist_ok=True)

TASK_R2V = "r2v — 参考主体生视频(Reference to Video)"
TASK_T2V = "t2v — 文生视频(Text to Video)"
UNET_R2V = "minimax_h3_ref2va_pruned_int8_convrot.safetensors"
UNET_TEXT = "minimax_h3_fl2va_pruned_int8_convrot.safetensors"
CLIP = "qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors"
VAE_VIDEO = "minimax_h3_video_vae_fp16.safetensors"
VAE_AUDIO = "minimax_h3_audio_vae_fp32.safetensors"

STEPS = 12            # 用户指定：12 步足够（官方流 25 步为基准，实测步数与耗时近线性）
SAMPLER = "res_multistep"
SCHEDULER = "simple"
CFG = 1.0
SHIFT_VIDEO = 12.0
SHIFT_AUDIO = 3.0
FPS = 24.0
MAX_SEC = 11.0
MIN_FRAMES = 124      # 17*7+5 = 5.17s，H3 官方流实测可用的最短稳定档


def frames_for(dur):
    """取 ≤ dur 秒的最大 17k+5 帧数（≤11s）。"""
    want = min(dur, MAX_SEC) * FPS
    k = int((want - 5) // 17)
    n = 17 * k + 5
    while n > MAX_SEC * FPS:
        k -= 1
        n = 17 * k + 5
    return max(n, MIN_FRAMES)


def build_timeline(task_label, prompt_text, remote_refs, w, h, frames):
    refs = [{"index": i, "imageFile": fn, "subfolder": "", "type": "input"}
            for i, fn in enumerate(remote_refs or [])]
    return {
        "version": 4, "editMode": "global", "timelineMode": "gen_blank",
        "totalFrames": frames, "frameRate": FPS, "width": w, "height": h,
        "refMaxSize": max(w, h),
        "output": {"mode": "fixed", "longEdge": max(w, h), "width": w, "height": h,
                   "maxExportFrames": 0, "exportMode": "all",
                   "continuityEnabled": False, "continuityOverlapFrames": 9},
        "videoClips": [],
        "video": {"fileName": "", "videoFile": "", "subfolder": "", "type": "input",
                  "frames": [], "frameMap": []},
        "global": {"taskType": task_label, "prompt": prompt_text, "commonEnabled": True,
                   "refs": refs, "referenceVideo": {}, "continuousReference": False,
                   "genImage": {"imageFile": ""}},
        "segments": [{"id": "s0", "start": 0, "length": frames, "frameCount": frames,
                      "durationSec": round(frames / FPS, 3), "prompt": "", "taskType": "",
                      "refs": [], "referenceVideo": {}, "genImage": {"imageFile": ""},
                      "negativePrompt": ""}],
        "gen": {"defaultFrameCount": frames},
        "runSelectEnabled": False, "runSelection": [],
    }


def build_api(prompt_text, remote_refs, seed, task, w, h, frames):
    unet = UNET_R2V if task == "r2v" else UNET_TEXT
    label = TASK_R2V if task == "r2v" else TASK_T2V
    timeline = build_timeline(label, prompt_text, remote_refs, w, h, frames)
    return {
        "1": {"class_type": "UNETLoader", "inputs": {"unet_name": unet, "weight_dtype": "default"}},
        "2": {"class_type": "CLIPLoader", "inputs": {"clip_name": CLIP, "type": "minimax", "device": "default"}},
        "3": {"class_type": "VAELoader", "inputs": {"vae_name": VAE_VIDEO}},
        "4": {"class_type": "VAELoader", "inputs": {"vae_name": VAE_AUDIO}},
        "5": {"class_type": "MiniMaxH3Director", "inputs": {
            "model": ["1", 0], "video_vae": ["3", 0], "audio_vae": ["4", 0], "clip": ["2", 0],
            "task_type": label, "global_prompt": prompt_text,
            "bd_grp_sample": "采样设置", "cfg": CFG, "seed": seed, "frame_rate": FPS,
            "width": w, "height": h, "ref_max_size": max(w, h), "total_frames": frames,
            "timeline_data": json.dumps(timeline, ensure_ascii=False),
            "bd_grp_advanced": "高级采样", "steps": STEPS, "sampler": SAMPLER,
            "scheduler": SCHEDULER, "shift_video": SHIFT_VIDEO, "shift_audio": SHIFT_AUDIO,
            "bd_grp_perf": "性能 Performance", "clear_vram_between_segments": True,
            "export_source_images": False}},
        "6": {"class_type": "CreateVideo", "inputs": {"images": ["5", 0], "fps": ["5", 2], "audio": ["5", 1]}},
        "7": {"class_type": "SaveVideo", "inputs": {"video": ["6", 0],
              "filename_prefix": "video/OfficialDirector", "format": "auto", "codec": "h264"}},
    }


def gpu_temp():
    """读 GPU 机温度（通过 ssh_exec.sh，失败返回 None）。"""
    try:
        r = subprocess.run(["bash", os.path.expanduser("~/Downloads/ssh_exec.sh"),
                            "nvidia-smi --query-gpu=temperature.gpu --format=csv,noheader"],
                           capture_output=True, text=True, timeout=60)
        for line in r.stdout.splitlines():
            line = line.strip()
            if line.isdigit():
                return int(line)
    except Exception:
        pass
    return None


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--shot", required=True)
    ap.add_argument("--prompt-file", default=None)
    ap.add_argument("--prompt", default=None)
    ap.add_argument("--keyframe", default=None, help="主参考图（通常是该镜关键帧）")
    ap.add_argument("--extra-ref", action="append", default=[], help="附加参考图（如角色设定图）")
    ap.add_argument("--size", default="1344x768")
    ap.add_argument("--dur", type=float, default=5.2)
    ap.add_argument("--seed", type=int, default=20260925)
    ap.add_argument("--task", default="auto", choices=["auto", "r2v", "t2v"])
    ap.add_argument("--cooldown", type=int, default=300, help="段间冷却秒数（记录①：13 小时零挂机配方 = 段间 5 分钟）")
    ap.add_argument("--max-temp", type=int, default=70, help="冷却后仍需等待 GPU 降到该温度再开下一镜")
    ap.add_argument("--worker", default="kf", help="输出子目录（默认关键帧目录名）")
    ap.add_argument("--dry-run", action="store_true")
    a = ap.parse_args()

    w, h = (int(x) for x in a.size.lower().split("x"))
    frames = frames_for(a.dur)
    prompt = (Path(a.prompt_file).read_text(encoding="utf-8").strip() if a.prompt_file else (a.prompt or ""))
    if not prompt:
        sys.exit("需要 --prompt-file 或 --prompt")
    task = a.task
    if task == "auto":
        task = "r2v" if a.keyframe else "t2v"

    shot_dir = OUTBASE / a.shot
    shot_dir.mkdir(parents=True, exist_ok=True)
    (shot_dir / "prompt.txt").write_text(prompt, encoding="utf-8")

    refs_local = ([Path(a.keyframe)] if a.keyframe else []) + [Path(x) for x in a.extra_ref]
    remote = []
    for i, lp in enumerate(refs_local, 1):
        rn = f"{a.shot}_ref{i}_{lp.name}"
        if not a.dry_run:
            p.upload_image_to_comfy(str(lp), rn)
        remote.append(rn)
    print(f"[{a.shot}] 参考图 {remote} | {w}x{h} | {frames}帧={frames/FPS:.2f}s | steps={STEPS} | seed={a.seed} | task={task}", flush=True)

    api = build_api(prompt, remote, a.seed, task, w, h, frames)
    (shot_dir / "api_prompt.json").write_text(json.dumps(api, ensure_ascii=False, indent=1), encoding="utf-8")
    if a.dry_run:
        print("dry-run：已写出 api_prompt.json"); return 0

    t0 = time.time()
    entry = p.submit_and_wait(api, timeout_min=180)
    info = p.extract_video_output(entry)
    video = p.download_video(info, shot_dir)
    final = shot_dir / f"{a.shot}.mp4"
    if final.exists():
        final.unlink()
    shutil.move(str(video), str(final))
    dt = time.time() - t0
    print(f"[{a.shot}] 完成 {final} | 生成耗时 {dt/60:.1f} 分钟", flush=True)

    with open(ROOT / "logs" / "h3_official.log", "a", encoding="utf-8") as f:
        f.write(json.dumps({"shot": a.shot, "size": a.size, "frames": frames, "steps": STEPS,
                            "seed": a.seed, "task": task, "seconds": round(dt, 1),
                            "output": str(final)}, ensure_ascii=False) + "\n")
    if a.cooldown > 0:
        print(f"[{a.shot}] 冷却 {a.cooldown}s（释放显存/降温）…", flush=True)
        time.sleep(a.cooldown)
        # 温度闸门：等 GPU 降到 max_temp 以下再允许下一镜（记录显示 80℃ 长期满载会导致硬挂）
        for k in range(10):
            t = gpu_temp()
            if t is None or t <= a.max_temp:
                print(f"[{a.shot}] 温度闸门通过：{t}℃ ≤ {a.max_temp}℃", flush=True)
                break
            print(f"[{a.shot}] 当前 {t}℃ > {a.max_temp}℃，继续等 30s…", flush=True)
            time.sleep(30)
    return 0


if __name__ == "__main__":
    sys.exit(main())
