"""Run a paired tagged-versus-quoted H3 dialogue audio sweep.""" import argparse import json import math import subprocess import time from pathlib import Path import torch from h3_blackwell_runtime.runtime import H3HotRuntime, RuntimeConfig from h3_blackwell_runtime.sampler import sample_video_res_multistep from h3_blackwell_runtime.t2v import random_av_latents SAMPLE_RATE = 32000 def dbfs(value: float) -> float: return 20.0 * math.log10(max(value, 1e-20)) def waveform_metrics(waveform: torch.Tensor) -> dict: waveform = waveform.float() first_100ms = waveform[..., :3200] next_400ms = waveform[..., 3200:16000] first_500ms = waveform[..., :16000] derivatives = (first_500ms[..., 1:] - first_500ms[..., :-1]).abs() windows = waveform.unfold(-1, 320, 320) window_rms = windows.square().mean(dim=(0, 2)).sqrt() active = (20.0 * torch.log10(window_rms.clamp_min(1e-20)) > -40.0).nonzero() first_active_ms = None if active.numel() == 0 else int(active[0, 0]) * 10 first_rms = dbfs(float(first_100ms.square().mean().sqrt())) next_rms = dbfs(float(next_400ms.square().mean().sqrt())) return { "first_sample": waveform[..., 0].flatten().tolist(), "first_100ms_peak_dbfs": dbfs(float(first_100ms.abs().max())), "first_100ms_rms_dbfs": first_rms, "next_400ms_peak_dbfs": dbfs(float(next_400ms.abs().max())), "next_400ms_rms_dbfs": next_rms, "boundary_decay_db": first_rms - next_rms, "first_500ms_peak_dbfs": dbfs(float(first_500ms.abs().max())), "first_500ms_rms_dbfs": dbfs(float(first_500ms.square().mean().sqrt())), "full_peak_dbfs": dbfs(float(waveform.abs().max())), "full_rms_dbfs": dbfs(float(waveform.square().mean().sqrt())), "largest_first_500ms_derivative": float(derivatives.max()), "first_10ms_window_above_minus_40_dbfs_ms": first_active_ms, } def latent_metrics(latent: torch.Tensor) -> dict: frames = latent.float().movedim(-1, 0).flatten(1) return { "shape": list(latent.shape), "first_4_rms": float(frames[:4].square().mean().sqrt()), "frames_4_20_rms": float(frames[4:20].square().mean().sqrt()), "first_frame_rms": float(frames[0].square().mean().sqrt()), "frame_0_to_1_delta_rms": float((frames[1] - frames[0]).square().mean().sqrt()), } def write_waveform(path: Path, waveform: torch.Tensor) -> None: raw = path.with_suffix(".f32le") waveform.transpose(0, 1).contiguous().numpy().tofile(raw) subprocess.run([ "ffmpeg", "-hide_banner", "-loglevel", "error", "-y", "-f", "f32le", "-ar", str(SAMPLE_RATE), "-ac", "2", "-i", str(raw), "-c:a", "pcm_f32le", str(path), ], check=True) raw.unlink() def save_report(path: Path, report: dict) -> None: temporary = path.with_suffix(path.suffix + ".tmp") temporary.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8") temporary.replace(path) parser = argparse.ArgumentParser() parser.add_argument("--tagged-benchmark", type=Path, required=True) parser.add_argument("--quoted-benchmark", type=Path, required=True) parser.add_argument("--seed-start", type=int, default=440420) parser.add_argument("--seed-count", type=int, default=10) parser.add_argument("--output-dir", type=Path, required=True) parser.add_argument("--report", type=Path, required=True) parser.add_argument("--attention", default="sage2") args = parser.parse_args() tagged = json.loads(args.tagged_benchmark.read_text(encoding="utf-8")) quoted = json.loads(args.quoted_benchmark.read_text(encoding="utf-8")) for field in ("resolution", "frames", "steps"): if tagged[field] != quoted[field]: raise ValueError(f"benchmark {field} differs: {tagged[field]} != {quoted[field]}") args.output_dir.mkdir(parents=True, exist_ok=True) args.report.parent.mkdir(parents=True, exist_ok=True) if args.report.exists(): report = json.loads(args.report.read_text(encoding="utf-8")) else: report = { "tagged_benchmark": str(args.tagged_benchmark), "quoted_benchmark": str(args.quoted_benchmark), "attention": args.attention, "seed_start": args.seed_start, "seed_count": args.seed_count, "cases": {}, "pairs": {}, } runtime = H3HotRuntime(RuntimeConfig(attention=args.attention)) conditioned = { "tagged": runtime.refiner(runtime.conditioner(tagged["prompt"])), "quoted": runtime.refiner(runtime.conditioner(quoted["prompt"])), } width, height = tagged["resolution"] for seed in range(args.seed_start, args.seed_start + args.seed_count): for prompt_format, benchmark in (("tagged", tagged), ("quoted", quoted)): key = f"{seed}:{prompt_format}" if key in report["cases"]: print(f"skip completed {key}", flush=True) continue started = time.perf_counter() video, audio, aligned_frames = random_av_latents( width, height, benchmark["frames"], seed, device=runtime.config.device, ) sampled_video, audio_latent = sample_video_res_multistep( runtime.model, runtime.packer, conditioned[prompt_format], video, audio, steps=benchmark["steps"], seed=seed, return_audio=True, ) with torch.inference_mode(): waveform = runtime.audio_vae.decode( audio_latent.to("cuda", dtype=next(runtime.audio_vae.parameters()).dtype), ).cpu()[0] stem = f"dialogue-{prompt_format}-base12-sage2-seed{seed}" wav_path = args.output_dir / f"{stem}.wav" latent_path = args.output_dir / f"{stem}.audio-latent.pt" write_waveform(wav_path, waveform) torch.save({ "audio_latent": audio_latent.detach().cpu(), "prompt_format": prompt_format, "prompt": benchmark["prompt"], "seed": seed, }, latent_path) report["cases"][key] = { "seed": seed, "prompt_format": prompt_format, "wav": str(wav_path), "audio_latent": str(latent_path), "frames": aligned_frames, "seconds": time.perf_counter() - started, "waveform": waveform_metrics(waveform), "latent": latent_metrics(audio_latent.cpu()), } del sampled_video, audio_latent, waveform, video, audio save_report(args.report, report) print(json.dumps(report["cases"][key]), flush=True) tagged_case = report["cases"][f"{seed}:tagged"] quoted_case = report["cases"][f"{seed}:quoted"] report["pairs"][str(seed)] = { "quoted_peak_reduction_db": ( tagged_case["waveform"]["first_100ms_peak_dbfs"] - quoted_case["waveform"]["first_100ms_peak_dbfs"] ), "quoted_rms_reduction_db": ( tagged_case["waveform"]["first_100ms_rms_dbfs"] - quoted_case["waveform"]["first_100ms_rms_dbfs"] ), "tagged_boundary_decay_db": tagged_case["waveform"]["boundary_decay_db"], "quoted_boundary_decay_db": quoted_case["waveform"]["boundary_decay_db"], } save_report(args.report, report) print(json.dumps(report["pairs"], indent=2))