186 lines
7.1 KiB
Python
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))
|