File size: 3,378 Bytes
d5d6c88 | 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 | 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()
|