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