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