File size: 6,230 Bytes
275b5a1 923fdff 275b5a1 923fdff 275b5a1 923fdff 275b5a1 923fdff 275b5a1 923fdff 275b5a1 923fdff 275b5a1 923fdff 275b5a1 923fdff 275b5a1 923fdff 275b5a1 923fdff 275b5a1 923fdff 275b5a1 923fdff 275b5a1 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 | 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()
|