| """Replay pinned short/multichunk Phase-0 latents through the frozen vocoder.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import os |
| from pathlib import Path |
| import shutil |
| import tempfile |
|
|
| import torch |
|
|
| from music3lab.manifests import ( |
| atomic_write_bytes, |
| canonical_json_bytes, |
| semantic_digest, |
| ) |
| from music3lab.vocoder import ( |
| build_replay_manifest, |
| load_frozen_vocoder, |
| load_phase0_vocoder_oracle, |
| module_state_sha256, |
| probe_latent_gradient, |
| publish_replay_bundle, |
| publish_replay_session_root, |
| replay_vocoder_oracle, |
| secure_replay_output_parent, |
| verify_replay_bundle, |
| verify_replay_session_tree, |
| ) |
|
|
|
|
| def build_parser() -> argparse.ArgumentParser: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--snapshot", type=Path, required=True) |
| parser.add_argument("--base-manifest", type=Path, required=True) |
| parser.add_argument("--diffusers-root", type=Path, required=True) |
| parser.add_argument("--phase0-artifacts", type=Path, required=True) |
| parser.add_argument("--output-root", type=Path, required=True) |
| parser.add_argument("--device", default="cuda") |
| return parser |
|
|
|
|
| def main() -> int: |
| args = build_parser().parse_args() |
| device = torch.device(args.device) |
| if device.type != "cuda" or not torch.cuda.is_available(): |
| raise RuntimeError( |
| "oracle replay requires the explicitly authorized CUDA GPU" |
| ) |
| device_name = torch.cuda.get_device_name(device) |
| if "H100" not in device_name: |
| raise RuntimeError(f"oracle replay requires H100, observed {device_name}") |
| device_capability = tuple(torch.cuda.get_device_capability(device)) |
| cuda_runtime = torch.version.cuda |
| if not cuda_runtime: |
| raise RuntimeError("PyTorch does not expose a CUDA runtime identity") |
|
|
| output_root = Path( |
| os.path.abspath(os.fspath(args.output_root.expanduser())) |
| ) |
| if os.path.lexists(output_root): |
| raise FileExistsError(f"replay output already exists: {output_root}") |
| output_root.parent.mkdir(parents=True, exist_ok=True) |
| secure_replay_output_parent(output_root.parent) |
|
|
| oracles = { |
| kind: load_phase0_vocoder_oracle(args.phase0_artifacts, kind) |
| for kind in ("short", "multi") |
| } |
| adapter = load_frozen_vocoder( |
| snapshot=args.snapshot, |
| base_manifest=args.base_manifest, |
| diffusers_root=args.diffusers_root, |
| ) |
| if adapter.report.project_git_dirty: |
| raise RuntimeError("oracle replay refuses dirty project source") |
|
|
| pending = [] |
| summaries = [] |
| for mode, dtype in ( |
| ("fp32_reference", torch.float32), |
| ("bf16_pipeline", torch.bfloat16), |
| ): |
| adapter.to(device=device, dtype=dtype) |
| torch.cuda.reset_peak_memory_stats(device) |
| mode_weight_sha256 = module_state_sha256(adapter.model) |
| gradient_probe = probe_latent_gradient( |
| adapter, |
| oracles["short"].latents[0], |
| device=device, |
| dtype=dtype, |
| ) |
| torch.cuda.synchronize(device) |
| if gradient_probe.weight_state_sha256_before != mode_weight_sha256: |
| raise RuntimeError("gradient probe started from different weights") |
|
|
| for kind, oracle in oracles.items(): |
| with torch.inference_mode(): |
| result = replay_vocoder_oracle( |
| adapter, |
| oracle, |
| device=device, |
| dtype=dtype, |
| ) |
| torch.cuda.synchronize(device) |
| peak_allocated = torch.cuda.max_memory_allocated(device) |
| peak_reserved = torch.cuda.max_memory_reserved(device) |
| manifest, audio_bytes, wav_bytes = build_replay_manifest( |
| mode=mode, |
| adapter=adapter, |
| oracle=oracle, |
| result=result, |
| weight_state_sha256_before=mode_weight_sha256, |
| gradient_probe=gradient_probe, |
| device_name=device_name, |
| device_capability=device_capability, |
| cuda_runtime=cuda_runtime, |
| peak_cuda_allocated_bytes=peak_allocated, |
| peak_cuda_reserved_bytes=peak_reserved, |
| ) |
| if mode == "bf16_pipeline" and not manifest.exact_oracle_audio: |
| raise RuntimeError( |
| f"BF16 official replay differs for {kind}; stopping" |
| ) |
| relative = Path(kind) / mode |
| pending.append((relative, manifest, audio_bytes, wav_bytes)) |
| summary = { |
| "oracle": kind, |
| "mode": mode, |
| "output": str(output_root / relative), |
| "manifest_semantic_digest": manifest.semantic_digest, |
| "output_content_sha256": manifest.output_content_sha256, |
| "expected_content_sha256": manifest.expected_content_sha256, |
| "output_wav_sha256": manifest.output_wav_artifact.sha256, |
| "exact_oracle_audio": manifest.exact_oracle_audio, |
| "max_abs_error": manifest.max_abs_error, |
| "mean_abs_error": manifest.mean_abs_error, |
| "weight_state_sha256": manifest.weight_state_sha256_after, |
| "latent_gradient_sha256": ( |
| manifest.gradient_probe.gradient_content_sha256 |
| ), |
| "peak_cuda_allocated_bytes": peak_allocated, |
| "peak_cuda_reserved_bytes": peak_reserved, |
| } |
| print(json.dumps(summary, sort_keys=True), flush=True) |
| summaries.append(summary) |
| del result |
| if module_state_sha256(adapter.model) != mode_weight_sha256: |
| raise RuntimeError("vocoder weights changed across replay mode") |
|
|
| temporary_root = Path( |
| tempfile.mkdtemp( |
| prefix=f".{output_root.name}.", |
| suffix=".tmp", |
| dir=output_root.parent, |
| ) |
| ) |
| published = False |
| try: |
| verified_summaries = [] |
| for relative, manifest, audio_bytes, wav_bytes in pending: |
| leaf = temporary_root / relative |
| publish_replay_bundle( |
| leaf, |
| manifest=manifest, |
| audio_bytes=audio_bytes, |
| wav_bytes=wav_bytes, |
| ) |
| verified = verify_replay_bundle( |
| leaf, |
| expected_adapter_semantic_digest=adapter.report.semantic_digest, |
| expected_oracle_semantic_digest=manifest.oracle_semantic_digest, |
| ) |
| verified_summaries.append( |
| { |
| "path": relative.as_posix(), |
| "manifest_file_sha256": verified.manifest_file_sha256, |
| "manifest_semantic_digest": manifest.semantic_digest, |
| "audio_file_sha256": manifest.output_artifact.sha256, |
| "wav_file_sha256": manifest.output_wav_artifact.sha256, |
| "audio_content_sha256": manifest.output_content_sha256, |
| } |
| ) |
| session_payload = { |
| "schema_version": "music3lab.vocoder-replay-session.v2", |
| "status": "PASS", |
| "adapter_report": adapter.report.model_dump(mode="json"), |
| "device": str(device), |
| "device_name": device_name, |
| "device_capability": list(device_capability), |
| "cuda_runtime": cuda_runtime, |
| "replays": verified_summaries, |
| } |
| session_payload["semantic_digest"] = semantic_digest(session_payload) |
| atomic_write_bytes( |
| temporary_root / "session.json", |
| canonical_json_bytes(session_payload), |
| mode=0o644, |
| ) |
| publish_replay_session_root(temporary_root, output_root) |
| published = True |
| finally: |
| if not published and temporary_root.exists(): |
| shutil.rmtree(temporary_root) |
|
|
| for relative, manifest, _audio_bytes, _wav_bytes in pending: |
| verify_replay_bundle( |
| output_root / relative, |
| expected_adapter_semantic_digest=adapter.report.semantic_digest, |
| expected_oracle_semantic_digest=manifest.oracle_semantic_digest, |
| ) |
| verify_replay_session_tree(output_root) |
| print( |
| json.dumps( |
| { |
| "status": "PASS", |
| "adapter_semantic_digest": adapter.report.semantic_digest, |
| "project_git_commit": adapter.report.project_git_commit, |
| "project_source_sha256": adapter.report.project_source_sha256, |
| "session_semantic_digest": session_payload["semantic_digest"], |
| "output_root": str(output_root), |
| "replays": summaries, |
| }, |
| sort_keys=True, |
| ) |
| ) |
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|