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