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