from __future__ import annotations import argparse import os import sys from pathlib import Path import h5py import numpy as np import torch import torch.distributed as dist from tqdm import tqdm PROJECT_ROOT = Path(__file__).resolve().parent.parent TRAIN_ROOT = Path(__file__).resolve().parent / "training" if str(TRAIN_ROOT) not in sys.path: sys.path.insert(0, str(TRAIN_ROOT)) from config import resolve_path from diffusion import create_diffusion from models import EmptyConditionSWaG, SWaG def distributed_context() -> tuple[int, int, int]: world_size = int(os.environ.get("WORLD_SIZE", "1")) rank = int(os.environ.get("RANK", "0")) local_rank = int(os.environ.get("LOCAL_RANK", "0")) if world_size > 1: dist.init_process_group("nccl") return rank, local_rank, world_size def shard_path(output_path: Path, rank: int) -> Path: return output_path.with_name(f".{output_path.name}.rank{rank}.partial.h5") def main() -> None: parser = argparse.ArgumentParser(description="Generate waveforms with SWaG") parser.add_argument("--checkpoint", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--num-samples", type=int, required=True) parser.add_argument("--batch-size", type=int, default=256) parser.add_argument("--sampling-steps", type=int, default=250) parser.add_argument("--sampler", choices=("ddpm", "ddim"), default="ddpm") parser.add_argument("--seed", type=int, default=0) parser.add_argument( "--clip-denoised", action="store_true", help="Clip predicted waveforms to [-1, 1]. Leave disabled for standardized STEAD waveforms.", ) args = parser.parse_args() if args.num_samples < 1 or args.batch_size < 1 or args.sampling_steps < 1: raise ValueError("num-samples, batch-size, and sampling-steps must be positive") if not torch.cuda.is_available(): raise RuntimeError("CUDA is required for sampling") rank, local_rank, world_size = distributed_context() device = torch.device("cuda", local_rank) torch.cuda.set_device(device) if args.batch_size % world_size: raise ValueError(f"batch-size must be divisible by world size ({world_size})") local_batch = args.batch_size // world_size checkpoint_path = resolve_path(PROJECT_ROOT, str(args.checkpoint)) output_path = resolve_path(PROJECT_ROOT, str(args.output)) checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False) config = checkpoint["config"] model_class = EmptyConditionSWaG if int(config["model"].get("condition_slot_count", 0)) == 8 else SWaG model = model_class(**config["model"]).to(device).eval() model.load_state_dict(checkpoint["ema"], strict=True) del checkpoint diffusion_config = dict(config["diffusion"]) respacing = f"ddim{args.sampling_steps}" if args.sampler == "ddim" else str(args.sampling_steps) diffusion = create_diffusion(timestep_respacing=respacing, **diffusion_config) generator = torch.Generator(device=device).manual_seed(args.seed + rank) if rank == 0: output_path.parent.mkdir(parents=True, exist_ok=True) for shard_rank in range(world_size): shard_path(output_path, shard_rank).unlink(missing_ok=True) if world_size > 1: dist.barrier() shape = (args.num_samples, int(config["model"]["length"]), int(config["model"]["in_channels"])) assigned = np.array_split(np.arange(args.num_samples, dtype=np.int64), world_size)[rank] local_shape = (len(assigned), shape[1], shape[2]) with h5py.File(shard_path(output_path, rank), "w") as handle: output = handle.create_dataset("data", shape=local_shape, dtype="float32", compression="gzip", shuffle=True) handle.create_dataset("indices", data=assigned, dtype="int64") for start in tqdm( range(0, len(assigned), local_batch), desc=f"sampling rank {rank}", disable=rank != 0, ): count = min(local_batch, len(assigned) - start) noise = torch.randn( count, int(config["model"]["in_channels"]), int(config["model"]["length"]), device=device, generator=generator, ) sample_loop = diffusion.ddim_sample_loop if args.sampler == "ddim" else diffusion.p_sample_loop model_kwargs = None if int(config["model"].get("condition_slot_count", 0)) == 8: model_kwargs = { "conditions": torch.zeros(count, 8, device=device, dtype=torch.float32) } samples = sample_loop( model, noise.shape, noise=noise, device=device, progress=False, clip_denoised=args.clip_denoised, model_kwargs=model_kwargs, ) output[start : start + count] = samples.transpose(1, 2).cpu().numpy().astype(np.float32) if world_size > 1: dist.barrier() if rank == 0: partial = output_path.with_name(f".{output_path.name}.partial") partial.unlink(missing_ok=True) with h5py.File(partial, "w") as handle: output = handle.create_dataset("data", shape=shape, dtype="float32", compression="gzip", shuffle=True) handle.attrs["checkpoint"] = str(checkpoint_path) handle.attrs["sampling_steps"] = args.sampling_steps handle.attrs["sampler"] = args.sampler handle.attrs["seed"] = args.seed handle.attrs["clip_denoised"] = args.clip_denoised handle.attrs["world_size"] = world_size for shard_rank in range(world_size): current = shard_path(output_path, shard_rank) with h5py.File(current, "r") as shard: indices = np.asarray(shard["indices"], dtype=np.int64) output[indices] = shard["data"][:] current.unlink() partial.replace(output_path) if world_size > 1: dist.barrier() dist.destroy_process_group() if __name__ == "__main__": main()