from __future__ import annotations import argparse import json import os import numpy as np import soundfile as sf import torch from PIL import Image from diffusers.pipelines.deprecated.audio_diffusion.mel import Mel from safetensors.torch import load_file from transformers import AutoTokenizer, ClapModel from audio_dit import AudioDiT MAX_TOKENS = 32 @torch.no_grad() def generate(model, seq, pool, null_seq, null_pool, steps, cfg, dev): B = seq.shape[0] x = torch.randn(B, 1, model.y_res, model.x_res, device=dev) ns = null_seq.expand(B, -1, -1) npool = null_pool.expand(B, -1) dt = 1.0 / steps for i in range(steps): t = torch.full((B,), i * dt, device=dev) with torch.autocast("cuda", dtype=torch.bfloat16): vc = model(x, t, seq, pool) vu = model(x, t, ns, npool) x = x + (vu + cfg * (vc - vu)).float() * dt return x def image_from_tensor(row): arr = ((row.clamp(-1, 1).float().cpu().numpy() + 1) * 127.5 + 0.5).astype(np.uint8) return Image.fromarray(arr) def load_model(weights, config, dev): cfgd = json.load(open(config)) dit_cfg = cfgd["dit"] model = AudioDiT(x_res=dit_cfg["x_res"], y_res=dit_cfg["y_res"], text_seq_dim=dit_cfg["text_seq_dim"], text_pool_dim=dit_cfg["text_pool_dim"]).to(dev).eval() model.load_state_dict(load_file(weights)) mel = Mel(x_res=dit_cfg["x_res"], y_res=dit_cfg["y_res"], sample_rate=cfgd["mel"]["sample_rate"], n_fft=cfgd["mel"]["n_fft"], hop_length=cfgd["mel"]["hop_length"], top_db=cfgd["mel"]["top_db"]) return model, mel, cfgd def load_clap(name, dev): tok = AutoTokenizer.from_pretrained(name) clap = ClapModel.from_pretrained(name).to(dev).eval() @torch.no_grad() def enc(strings): t = tok(strings, padding="max_length", max_length=MAX_TOKENS, truncation=True, return_tensors="pt").to(dev) out = clap.text_model(**t) return out.last_hidden_state.float(), clap.text_projection(out.pooler_output).float() return enc def main(): ap = argparse.ArgumentParser() ap.add_argument("prompts", nargs="+") ap.add_argument("--out-dir", default="samples") ap.add_argument("--cfg", type=float, default=4.0) ap.add_argument("--steps", type=int, default=50) ap.add_argument("--device", default="cuda") ap.add_argument("--weights", default="/root/runs/audio_v1/model.safetensors") ap.add_argument("--config", default="/root/runs/audio_v1/config.json") ap.add_argument("--clap", default="laion/clap-htsat-unfused") args = ap.parse_args() dev = args.device os.makedirs(args.out_dir, exist_ok=True) model, mel, _ = load_model(args.weights, args.config, dev) enc = load_clap(args.clap, dev) seq, pool = enc(args.prompts) null_seq, null_pool = enc([""]) x = generate(model, seq, pool, null_seq, null_pool, args.steps, args.cfg, dev) for prompt, row in zip(args.prompts, x[:, 0]): img = image_from_tensor(row) audio = mel.image_to_audio(img) name = "_".join(prompt.lower().split())[:40] sf.write(os.path.join(args.out_dir, f"{name}.wav"), audio, mel.get_sample_rate()) print(f" wrote {name}.wav ({len(audio)/mel.get_sample_rate():.1f}s) {prompt}", flush=True) if __name__ == "__main__": main()