| 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() |
|
|