swag-train-bundle / generate.py
DancingNow's picture
Add files using upload-large-folder tool
923fdff verified
Raw History Blame Contribute Delete
6.23 kB
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()