h3-blackwell-runtime/tools/sweep_dialogue_audio_boundary.py
2026-08-22 14:09:45 +07:00

186 lines
7.1 KiB
Python

"""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))