Download generate.py from DancingNow/swag-train-bundle: direct link, hf CLI and curl.
- Browser
- Download file 6.23 kB
-
https://huggingface.co/DancingNow/swag-train-bundle/resolve/main/generate.py
- Command line
-
hf download hf://DancingNow/swag-train-bundle/generate.py
-
curl -L -o generate.py https://huggingface.co/DancingNow/swag-train-bundle/resolve/main/generate.py
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() | |