File size: 5,409 Bytes
7a76816
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
from __future__ import annotations

import argparse
import json
import os
import shutil

import numpy as np
import soundfile as sf
import torch
from diffusers.pipelines.deprecated.audio_diffusion.mel import Mel
from frechet_audio_distance import FrechetAudioDistance

from sample import generate, image_from_tensor, load_clap, load_model

PANN_EMBED_DIM = 2048

def fad_score(frechet, background_dir, eval_dir):
    audio_bg = frechet._FrechetAudioDistance__load_audio_files(background_dir, dtype="float32")
    embds_bg = frechet.get_embeddings(audio_bg, sr=frechet.sample_rate).reshape(-1, PANN_EMBED_DIM)
    audio_ev = frechet._FrechetAudioDistance__load_audio_files(eval_dir, dtype="float32")
    embds_ev = frechet.get_embeddings(audio_ev, sr=frechet.sample_rate).reshape(-1, PANN_EMBED_DIM)
    mu1, sigma1 = frechet.calculate_embd_statistics(embds_bg)
    mu2, sigma2 = frechet.calculate_embd_statistics(embds_ev)
    return frechet.calculate_frechet_distance(mu1, sigma1, mu2, sigma2)

def first_caption_per_clip(meta):
    seen, order = {}, []
    for pair_idx, ci in enumerate(meta["pair_clip_idx"]):
        if ci not in seen:
            seen[ci] = meta["captions"][pair_idx]
            order.append(ci)
    return order, seen

def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--data", default="/root/data")
    ap.add_argument("--audio-dir", default="/root/clotho_raw/evaluation")
    ap.add_argument("--weights", default="/root/runs/audio_v1/model_best.safetensors")
    ap.add_argument("--config", default="/root/runs/audio_v1/config.json")
    ap.add_argument("--clap", default="laion/clap-htsat-unfused")
    ap.add_argument("--n-eval", type=int, default=300)
    ap.add_argument("--cfg", type=float, default=4.0)
    ap.add_argument("--steps", type=int, default=50)
    ap.add_argument("--seed", type=int, default=0)
    ap.add_argument("--work", default="/root/fad_work")
    ap.add_argument("--out", default="/root/fad_results.json")
    ap.add_argument("--pann-sr", type=int, default=32000)
    args = ap.parse_args()
    dev = "cuda"

    meta = json.load(open(f"{args.data}/evaluation_meta.json"))
    mel_arr = np.load(f"{args.data}/evaluation_mel.npy")
    clip_order, caption_by_clip = first_caption_per_clip(meta)
    rng = np.random.RandomState(args.seed)
    n = min(args.n_eval, len(clip_order))
    chosen = [clip_order[i] for i in rng.permutation(len(clip_order))[:n]]
    prompts = [caption_by_clip[ci] for ci in chosen]
    files = [meta["clip_files"][ci] for ci in chosen]
    print(f"[fad] evaluating on {n} held-out clips from {args.audio_dir}", flush=True)

    real_dir = os.path.join(args.work, "real")
    gen_dir = os.path.join(args.work, "generated")
    floor_dir = os.path.join(args.work, "griffinlim_floor")
    uncond_dir = os.path.join(args.work, "no_prompt")
    for d in (real_dir, gen_dir, floor_dir, uncond_dir):
        shutil.rmtree(d, ignore_errors=True)
        os.makedirs(d, exist_ok=True)

    for fn in files:
        shutil.copy(os.path.join(args.audio_dir, fn), os.path.join(real_dir, fn))

    model, mel, _ = load_model(args.weights, args.config, dev)
    enc = load_clap(args.clap, dev)
    null_seq, null_pool = enc([""])

    print("[fad] writing griffin-lim floor (real mel -> griffin-lim, isolates vocoder loss)", flush=True)
    for ci, fn in zip(chosen, files):
        img = image_from_tensor(torch.from_numpy(mel_arr[ci].astype(np.float32) / 127.5 - 1.0))
        audio = mel.image_to_audio(img)
        sf.write(os.path.join(floor_dir, fn), audio, mel.get_sample_rate())

    print("[fad] generating from trained model", flush=True)
    batch = 32
    for i in range(0, n, batch):
        p = prompts[i:i + batch]
        seq, pool = enc(p)
        x = generate(model, seq, pool, null_seq, null_pool, args.steps, args.cfg, dev)
        for j, row in enumerate(x[:, 0]):
            img = image_from_tensor(row)
            audio = mel.image_to_audio(img)
            sf.write(os.path.join(gen_dir, files[i + j]), audio, mel.get_sample_rate())
        print(f"[fad] generated {min(i+batch,n)}/{n}", flush=True)

    print("[fad] generating unconditional (no prompt) baseline", flush=True)
    for i in range(0, n, batch):
        m = min(batch, n - i)
        ns = null_seq.expand(m, -1, -1)
        npool = null_pool.expand(m, -1)
        x = generate(model, ns, npool, null_seq, null_pool, args.steps, 1.0, dev)
        for j, row in enumerate(x[:, 0]):
            img = image_from_tensor(row)
            audio = mel.image_to_audio(img)
            sf.write(os.path.join(uncond_dir, files[i + j]), audio, mel.get_sample_rate())

    frechet = FrechetAudioDistance(model_name="pann", sample_rate=args.pann_sr,
                                    use_pca=False, use_activation=False, verbose=False)

    results = {}
    for name, d in [("generated_vs_real", gen_dir), ("griffinlim_floor_vs_real", floor_dir),
                    ("no_prompt_vs_real", uncond_dir)]:
        score = fad_score(frechet, real_dir, d)
        results[name] = score
        print(f"[fad] FAD {name:<26} = {score:.4f}", flush=True)

    results["config"] = {"n_eval": n, "cfg": args.cfg, "steps": args.steps,
                          "pann_sample_rate": args.pann_sr, "weights": args.weights}
    json.dump(results, open(args.out, "w"), indent=2)
    print(f"[fad] wrote {args.out}", flush=True)

if __name__ == "__main__":
    main()