Skip to content

Source: examples/inference/eval/mmaudio

MMAudio dataset inference and evaluation¶

This example evaluates a trained MMAudio checkpoint without importing the upstream MMAudio package. Generation uses FastVideo's native MMAudioPipeline; the default metrics use fastvideo.eval.

Relation to official MMAudio¶

Official batch_eval.py launches one complete model replica per GPU with torchrun, uses a DistributedSampler, runs batched bf16 inference, and writes audio without composing videos. Metric calculation is a separate step through av_bench.extract and av_bench.evaluate.

This FastVideo path keeps that high-level separation:

  1. one complete MMAudio pipeline replica per GPU;
  2. data-parallel dataset inference to WAV;
  3. a canonical JSONL manifest that binds id, source video, generated audio, caption, seed, and reference-audio source;
  4. independent multi-GPU metric workers over the complete corpus.

There are two intentional differences. Rank-strided indices avoid the padded duplicates that PyTorch DistributedSampler(drop_last=False) can introduce, and cached CLIP/Synchformer/text features skip video decoding and conditioning encoder work already completed during FastVideo preprocessing. The current cached-feature stage accepts one sample per pipeline call; four GPUs still process four samples concurrently.

Environment¶

Use the same FastVideo environment as training:

cd /mnt/lustre/vlm-kai/FastVideo
source .venv/bin/activate
uv pip install --python .venv/bin/python hear21passt pyloudnorm

The four default native metrics are:

FastVideo metric Protocol
audio.frechet_distance corpus PaSST FAD / FD_PaSST
audio.kl_divergence paired PaSST KL, KL(gt || pred)
audio.clap_score text-audio CLAP similarity
audio.desync source-video/generated-audio Synchformer DeSync

FAD and KL port the av-benchmark mathematics. DeSync vendors the same Synchformer family. FastVideo CLAP uses the Hugging Face LAION CLAP checkpoint and is not byte-identical to the external package's CLAP implementation; use the optional official backend below when exact paper-table reproduction is required.

1. Export the final EMA¶

The training result is a transformer state dict. Inference also needs the frozen audio VAE, vocoder, scheduler, and (for raw-video inference) the frozen conditioning encoders. Export a complete FastVideo model once:

bash examples/inference/eval/mmaudio/export_final_ema.sh

Default input:

outputs/mmaudio_small_44k_ddp_from_scratch/posthoc_ema/official_ddp/
  mmaudio_ema_final_sigma_0p05_step_000300000.pth

Default output:

converted_weights/mmaudio/small_44k_ema_300000/

Override EMA_CHECKPOINT, OUTPUT_MODEL, or ASSET_ROOT before the command for another run or variant.

2. Multi-GPU inference¶

The launcher defaults to the VGGSound test feature cache and four GPUs:

bash examples/inference/eval/mmaudio/run_inference_vggsound.sh

For a four-sample smoke test:

MAX_SAMPLES=4 \
OUTPUT_DIR=outputs/mmaudio_small_44k_ema_300000_vggsound_smoke \
bash examples/inference/eval/mmaudio/run_inference_vggsound.sh

For a background full run with a persistent log:

mkdir -p logs
nohup bash examples/inference/eval/mmaudio/run_inference_vggsound.sh \
  > logs/mmaudio_vggsound_inference.log 2>&1 &
tail -f logs/mmaudio_vggsound_inference.log

The output directory contains:

audio/<sample-id>.wav
eval_manifest.jsonl
failures.jsonl
summary.json
manifest_rank_*.jsonl
failures_rank_*.jsonl

Existing WAV files are resumed rather than regenerated. Add --overwrite to the Python runner when intentional regeneration is required. Seeds are base_seed + global_index, so results do not change with GPU count.

3. Native FastVideo evaluation¶

Run FAD, KL, CLAP, and DeSync with four metric replicas:

bash examples/inference/eval/mmaudio/run_evaluate_vggsound.sh

The first run extracts ground-truth audio from each source MP4 into <OUTPUT_DIR>/reference_audio and caches metric weights. Later runs reuse both. Results are written to fastvideo_eval_results.json; per-sample metrics are under samples, while FAD is under corpus.

The underlying generic command is:

fastvideo eval run \
  --manifest outputs/mmaudio_small_44k_ema_300000_vggsound_test/eval_manifest.jsonl \
  --metrics audio.frechet_distance,audio.kl_divergence,audio.clap_score,audio.desync \
  --num-gpus 4 \
  --extract-audio outputs/mmaudio_small_44k_ema_300000_vggsound_test/reference_audio \
  --extract-workers 16 \
  --output-format full \
  --output outputs/mmaudio_small_44k_ema_300000_vggsound_test/fastvideo_eval_results.json

Do not evaluate FAD by invoking the evaluator once per file. FAD is a set-vs-set statistic; this CLI submits the complete manifest in one call and serializes EvalResults.corpus.

Exact external av-benchmark backend¶

For strict comparison to official MMAudio numbers, FastVideo provides the dedicated fastvideo eval v2a command. It launches the exact av_bench.extract and av_bench.evaluate functions with a Python interpreter from an isolated environment. This keeps av-benchmark's CLAP, ImageBind, and PyTorch dependencies out of the main FastVideo environment.

Create the isolated environment once:

cd /mnt/lustre/vlm-kai
git clone https://github.com/hkchengrex/av-benchmark.git
uv venv av-benchmark/.venv --python 3.12
uv pip install --python av-benchmark/.venv/bin/python -e av-benchmark
# PyTorch 2.9+ routes torchaudio decoding through TorchCodec. Transformers 5
# removed an API used by the vendored Synchformer AST, while SciPy 1.17 removed
# an argument used by the official FAD implementation. Keep these constraints
# in this isolated environment; select the wheel matching its PyTorch CUDA build.
UV_TORCH_BACKEND=cu130 uv pip install \
  --python av-benchmark/.venv/bin/python \
  torchcodec==0.15.0 'transformers<5' 'scipy<1.17'

mkdir -p av-benchmark/weights
curl -L https://huggingface.co/lukewys/laion_clap/resolve/main/music_speech_audioset_epoch_15_esc_89.98.pt \
  -o av-benchmark/weights/music_speech_audioset_epoch_15_esc_89.98.pt
curl -L https://github.com/hkchengrex/MMAudio/releases/download/v0.1/synchformer_state_dict.pth \
  -o av-benchmark/weights/synchformer_state_dict.pth

Download only the official VGGSound ground-truth cache (about 16 GB):

cd /mnt/lustre/vlm-kai/FastVideo
.venv/bin/hf download hkchengrex/MMAudio-precomputed-results \
  --repo-type dataset \
  --include "vggsound-test-eval-cache/*" \
  --local-dir /mnt/lustre/vlm-kai/datasets/VGGSound/av_benchmark

Then run the dedicated launcher:

bash examples/inference/eval/mmaudio/run_evaluate_vggsound_av_benchmark.sh

For a background run:

nohup bash examples/inference/eval/mmaudio/run_evaluate_vggsound_av_benchmark.sh \
  > logs/mmaudio_vggsound_av_benchmark.log 2>&1 &
tail -f logs/mmaudio_vggsound_av_benchmark.log

The official VGGSound cache produces the complete paper-style metric set:

FD-VGG, FD-PANN, FD-PASST
KL-PANNS-softmax, KL-PASST-softmax
ISC-PANNS-mean/std, ISC-PASST-mean/std
IB-Score, DeSync

The official VGGSound cache does not contain CLAP text features, so the launcher defaults to SKIP_CLAP=1; this saves prediction feature passes without changing the metric set. Set SKIP_CLAP=0 for a custom GT cache that includes CLAP text features. Prediction features are cached and reused; set RECOMPUTE=1 to force a clean extraction.

The launcher also defaults to ALIGN_PREDICTION_KEYS=1. The official VGGSound cache prepends underscores to some sanitized ids; FastVideo keeps the original ids. The backend changes only unique matches in prediction feature caches, leaves WAV filenames untouched, and records exact/remapped/ambiguous/ unmatched counts under prediction_key_alignment in the result JSON.

The launcher defaults to NUM_WORKERS=0: this avoids both the machine's 64 MB /dev/shm limit and forking workers after the official extractor initializes CUDA. Raise it only after validating the container and backend combination.

Evaluation protocol warning¶

The current FastVideo test cache contains the filtered natural-language descriptions used in this training run. Official MMAudio's VGGSound batch_eval.py reads class-label captions from the original VGGSound CSV. FAD, KL, and DeSync formulas do not read captions directly, but changing the generation caption changes the generated waveform and can therefore change all scores. CLAP additionally consumes the caption during evaluation. Record the generation manifest, test split, and ground-truth cache with reported results.

Additional Files¶

eval_mmaudio_dataset.py
# SPDX-License-Identifier: Apache-2.0
"""Distributed MMAudio inference over FastVideo preprocessing caches.

Launch with ``torchrun``. Each process owns one complete MMAudio pipeline and
processes a non-padded, rank-strided subset of the dataset. The output manifest
is directly consumable by ``fastvideo eval run --manifest ...``.
"""

from __future__ import annotations

import argparse
import json
import os
import re
from pathlib import Path
from typing import Any

import numpy as np
import torch
from scipy.io import wavfile
from torch.utils.data import DataLoader, Dataset, Sampler
from tqdm import tqdm

from fastvideo.dataset.mmaudio_feature_dataset import build_mmaudio_feature_dataset
from fastvideo.distributed import cleanup_dist_env_and_memory
from fastvideo.pipelines.basic.mmaudio import MMAudioPipeline
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch


class _RankStrideSampler(Sampler[int]):
    """Assign every global index exactly once without tail padding."""

    def __init__(self, size: int, rank: int, world_size: int) -> None:
        self.indices = range(rank, size, world_size)

    def __iter__(self):
        return iter(self.indices)

    def __len__(self) -> int:
        return len(self.indices)


class _IndexedDataset(Dataset):

    def __init__(self, dataset: Dataset) -> None:
        self.dataset = dataset

    def __len__(self) -> int:
        return len(self.dataset)

    def __getitem__(self, index: int) -> tuple[int, dict[str, Any]]:
        return index, self.dataset[index]


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--model-path", type=Path, required=True)
    parser.add_argument("--feature-root", type=Path, required=True)
    parser.add_argument("--output-dir", type=Path, required=True)
    parser.add_argument("--duration-seconds", type=float, default=8.0)
    parser.add_argument("--num-inference-steps", type=int, default=25)
    parser.add_argument("--guidance-scale", type=float, default=4.5)
    parser.add_argument("--seed", type=int, default=14159265)
    parser.add_argument("--negative-prompt", default="")
    parser.add_argument("--num-workers", type=int, default=2)
    parser.add_argument("--max-samples", type=int, default=0)
    parser.add_argument("--overwrite", action="store_true")
    parser.add_argument("--compile", action="store_true")
    return parser.parse_args()


def _safe_name(value: str) -> str:
    value = re.sub(r"[^A-Za-z0-9._-]+", "_", value).strip("._")
    return value or "sample"


def _feature_shapes(transformer: torch.nn.Module) -> dict[str, int]:
    arch = transformer.config.arch_config
    return {
        "latent_seq_len": int(transformer.latent_seq_len),
        "latent_dim": int(transformer.latent_dim),
        "clip_seq_len": int(transformer.clip_seq_len),
        "clip_dim": int(arch.clip_dim),
        "sync_seq_len": int(transformer.sync_seq_len),
        "sync_dim": int(arch.sync_dim),
        "text_seq_len": int(arch.text_seq_len),
        "text_dim": int(arch.text_dim),
    }


class _IgnoredComponent:
    """Placeholder for a required module the direct-feature path never touches.

    ``MMAudioPipeline`` requires the CLIP text/vision and Synchformer encoders,
    but its direct-feature conditioning path returns before dereferencing them.
    Callers that pass ``video_path`` (or omit the cached features) must supply
    the real modules instead.
    """

    def __repr__(self) -> str:
        return "<ignored component>"


def _build_pipeline(args: argparse.Namespace, world_size: int) -> MMAudioPipeline:
    # Cached features make the three conditioning encoders unnecessary. The
    # stages still exist, but consume direct tensors and never dereference
    # these sentinels.
    ignored_component = _IgnoredComponent()
    return MMAudioPipeline.from_pretrained(
        str(args.model_path),
        inference_mode=True,
        loaded_modules={
            "text_encoder": ignored_component,
            "tokenizer": ignored_component,
            "image_encoder": ignored_component,
            "image_encoder_2": ignored_component,
        },
        workload_type="v2a",
        num_gpus=world_size,
        tp_size=1,
        sp_size=1,
        hsdp_replicate_dim=1,
        hsdp_shard_dim=1,
        dit_cpu_offload=False,
        dit_layerwise_offload=False,
        text_encoder_cpu_offload=False,
        image_encoder_cpu_offload=False,
        vae_cpu_offload=False,
        enable_torch_compile=args.compile,
    )


def _forward_batch(
    sample: dict[str, Any],
    *,
    duration_seconds: float,
    num_inference_steps: int,
    guidance_scale: float,
    seed: int,
    negative_prompt: str,
) -> ForwardBatch:
    return ForwardBatch(
        data_type="video",
        prompt="",
        negative_prompt=negative_prompt,
        audio_start_in_s=0.0,
        audio_end_in_s=duration_seconds,
        num_inference_steps=num_inference_steps,
        guidance_scale=guidance_scale,
        seed=seed,
        num_videos_per_prompt=1,
        height=8,
        width=8,
        num_frames=1,
        save_video=False,
        return_frames=False,
        extra={
            "mmaudio_clip_features": sample["clip_features"].unsqueeze(0),
            "mmaudio_sync_features": sample["sync_features"].unsqueeze(0),
            "mmaudio_text_features": sample["text_features"].unsqueeze(0),
        },
    )


def _write_jsonl(path: Path, rows: list[dict[str, Any]]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    temporary = path.with_suffix(path.suffix + ".tmp")
    with temporary.open("w", encoding="utf-8") as handle:
        for row in rows:
            handle.write(json.dumps(row, sort_keys=True) + "\n")
    temporary.replace(path)


def _merge_rank_outputs(output_dir: Path, world_size: int) -> None:
    rows: list[dict[str, Any]] = []
    failures: list[dict[str, Any]] = []
    for rank in range(world_size):
        manifest = output_dir / f"manifest_rank_{rank:05d}.jsonl"
        failure_path = output_dir / f"failures_rank_{rank:05d}.jsonl"
        if manifest.is_file():
            rows.extend(json.loads(line) for line in manifest.read_text(encoding="utf-8").splitlines() if line.strip())
        if failure_path.is_file():
            failures.extend(
                json.loads(line) for line in failure_path.read_text(encoding="utf-8").splitlines() if line.strip())
    rows.sort(key=lambda row: int(row["index"]))
    failures.sort(key=lambda row: int(row["index"]))
    ids = [str(row["id"]) for row in rows]
    if len(ids) != len(set(ids)):
        raise RuntimeError("Duplicate sample ids found while merging inference manifests")
    _write_jsonl(output_dir / "eval_manifest.jsonl", rows)
    _write_jsonl(output_dir / "failures.jsonl", failures)
    summary = {
        "num_succeeded": len(rows),
        "num_failed": len(failures),
        "world_size": world_size,
    }
    (output_dir / "summary.json").write_text(
        json.dumps(summary, indent=2, sort_keys=True) + "\n",
        encoding="utf-8",
    )


@torch.inference_mode()
def main() -> None:
    args = parse_args()
    rank = int(os.environ.get("RANK", "0"))
    local_rank = int(os.environ.get("LOCAL_RANK", "0"))
    world_size = int(os.environ.get("WORLD_SIZE", "1"))
    if not torch.cuda.is_available():
        raise RuntimeError("MMAudio dataset inference requires CUDA")
    torch.cuda.set_device(local_rank)

    output_dir = args.output_dir.expanduser().resolve()
    audio_dir = output_dir / "audio"
    audio_dir.mkdir(parents=True, exist_ok=True)
    pipeline = _build_pipeline(args, world_size)
    transformer = pipeline.get_module("transformer")
    dataset = build_mmaudio_feature_dataset(
        args.feature_root,
        feature_shapes=_feature_shapes(transformer),
        include_metadata=True,
    )
    dataset_size = len(dataset)
    if args.max_samples > 0:
        dataset_size = min(dataset_size, args.max_samples)
    sampler = _RankStrideSampler(dataset_size, rank, world_size)
    loader = DataLoader(
        _IndexedDataset(dataset),
        batch_size=None,
        sampler=sampler,
        num_workers=max(0, args.num_workers),
        pin_memory=True,
        persistent_workers=args.num_workers > 0,
    )

    manifest_rows: list[dict[str, Any]] = []
    failure_rows: list[dict[str, Any]] = []
    progress = tqdm(
        loader,
        total=len(sampler),
        desc=f"MMAudio inference rank {rank}",
        position=local_rank,
    )
    for index, sample in progress:
        index = int(index)
        sample_id = _safe_name(str(sample.get("sample_id", index)))
        source_path = str(sample.get("source_path", ""))
        caption = str(sample.get("caption", ""))
        output_path = audio_dir / f"{sample_id}.wav"
        sample_seed = args.seed + index
        try:
            if not source_path or not Path(source_path).is_file():
                raise FileNotFoundError(f"source video is missing for sample {sample_id}: {source_path}")
            if args.overwrite or not output_path.is_file():
                batch = _forward_batch(
                    sample,
                    duration_seconds=args.duration_seconds,
                    num_inference_steps=args.num_inference_steps,
                    guidance_scale=args.guidance_scale,
                    seed=sample_seed,
                    negative_prompt=args.negative_prompt,
                )
                result = pipeline.forward(batch, pipeline.fastvideo_args)
                audio = result.extra.get("audio")
                sample_rate = result.extra.get("audio_sample_rate")
                if not isinstance(audio, np.ndarray) or not isinstance(sample_rate, int):
                    raise RuntimeError("MMAudio pipeline did not return decoded audio")
                temporary = output_path.with_suffix(".tmp.wav")
                wavfile.write(temporary, sample_rate, audio.astype(np.float32))
                temporary.replace(output_path)
            manifest_rows.append({
                "id": sample_id,
                "index": index,
                "video": str(Path(source_path).resolve()),
                "audio": str(output_path.resolve()),
                "reference_audio_source": str(Path(source_path).resolve()),
                "text_prompt": caption,
                "seed": sample_seed,
            })
        except Exception as error:  # noqa: BLE001 - isolate corrupt benchmark samples
            failure_rows.append({
                "id": sample_id,
                "index": index,
                "source": source_path,
                "error_type": type(error).__name__,
                "error": str(error),
            })

    _write_jsonl(output_dir / f"manifest_rank_{rank:05d}.jsonl", manifest_rows)
    _write_jsonl(output_dir / f"failures_rank_{rank:05d}.jsonl", failure_rows)
    if torch.distributed.is_initialized():
        torch.distributed.barrier()
    if rank == 0:
        _merge_rank_outputs(output_dir, world_size)
        print(f"Merged eval manifest: {output_dir / 'eval_manifest.jsonl'}")
        print(f"Inference summary: {output_dir / 'summary.json'}")
    if torch.distributed.is_initialized():
        torch.distributed.barrier()
    pipeline.close()
    cleanup_dist_env_and_memory()


if __name__ == "__main__":
    main()
evaluate_official_av_benchmark.py
# SPDX-License-Identifier: Apache-2.0
"""Run the exact external av-benchmark API used by official MMAudio."""

from __future__ import annotations

import argparse
import json
from pathlib import Path


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--audio-dir", type=Path, required=True)
    parser.add_argument("--gt-cache", type=Path, required=True)
    parser.add_argument("--prediction-cache", type=Path, required=True)
    parser.add_argument("--output", type=Path, required=True)
    parser.add_argument("--device", default="cuda")
    parser.add_argument("--batch-size", type=int, default=32)
    parser.add_argument("--audio-length", type=float, default=8.0)
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    try:
        from av_bench.evaluate import evaluate
        from av_bench.extract import extract
    except ImportError as error:
        raise ImportError("The exact official backend is optional. Install hkchengrex/"
                          "av-benchmark in a separate environment before running this file.") from error

    args.prediction_cache.mkdir(parents=True, exist_ok=True)
    extract(
        audio_path=args.audio_dir,
        output_path=args.prediction_cache,
        device=args.device,
        batch_size=args.batch_size,
        audio_length=args.audio_length,
    )
    metrics = evaluate(
        gt_audio_cache=args.gt_cache,
        pred_audio_cache=args.prediction_cache,
    )
    payload = {
        "backend": "hkchengrex/av-benchmark",
        "audio_dir": str(args.audio_dir.resolve()),
        "gt_cache": str(args.gt_cache.resolve()),
        "prediction_cache": str(args.prediction_cache.resolve()),
        "metrics": metrics,
    }
    args.output.parent.mkdir(parents=True, exist_ok=True)
    args.output.write_text(
        json.dumps(payload, indent=2, sort_keys=True) + "\n",
        encoding="utf-8",
    )
    print(json.dumps(payload, indent=2, sort_keys=True))


if __name__ == "__main__":
    main()
export_final_ema.sh
#!/usr/bin/env bash
set -euo pipefail

# One-time export of the synthesized PostHocEMA transformer plus the frozen
# official conditioning/decoder components into a complete FastVideo model.
ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../../.." && pwd)"
cd "${ROOT_DIR}"

EMA_CHECKPOINT="${EMA_CHECKPOINT:-outputs/mmaudio_small_44k_ddp_from_scratch/posthoc_ema/official_ddp/mmaudio_ema_final_sigma_0p05_step_000300000.pth}"
OUTPUT_MODEL="${OUTPUT_MODEL:-converted_weights/mmaudio/small_44k_ema_300000}"
ASSET_ROOT="${ASSET_ROOT:-official_weights/mmaudio}"

if [[ -e "${OUTPUT_MODEL}/model_index.json" ]]; then
  echo "Model already exported: ${OUTPUT_MODEL}"
  exit 0
fi

source .venv/bin/activate
python scripts/checkpoint_conversion/convert_mmaudio_to_diffusers.py \
  --variant small_44k \
  --transformer-checkpoint "${EMA_CHECKPOINT}" \
  --audio-vae-checkpoint "${ASSET_ROOT}/raw/ext_weights/v1-44.pth" \
  --synchformer-checkpoint "${ASSET_ROOT}/raw/ext_weights/synchformer_state_dict.pth" \
  --dfn5b-dir "${ASSET_ROOT}/DFN5B-CLIP-ViT-H-14-384" \
  --bigvgan-dir "${ASSET_ROOT}/bigvgan_v2_44khz_128band_512x" \
  --output "${OUTPUT_MODEL}"

echo "Exported FastVideo MMAudio model: ${OUTPUT_MODEL}"
run_evaluate_vggsound.sh
#!/usr/bin/env bash
set -euo pipefail

ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../../.." && pwd)"
cd "${ROOT_DIR}"
source .venv/bin/activate

OUTPUT_DIR="${OUTPUT_DIR:-outputs/mmaudio_small_44k_ema_300000_vggsound_test}"
MANIFEST="${MANIFEST:-${OUTPUT_DIR}/eval_manifest.jsonl}"
REFERENCE_AUDIO_CACHE="${REFERENCE_AUDIO_CACHE:-${OUTPUT_DIR}/reference_audio}"
RESULTS="${RESULTS:-${OUTPUT_DIR}/fastvideo_eval_results.json}"
NUM_GPUS="${NUM_GPUS:-4}"
EXTRACT_WORKERS="${EXTRACT_WORKERS:-16}"
METRICS="${METRICS:-audio.frechet_distance,audio.kl_divergence,audio.clap_score,audio.desync}"

fastvideo eval run \
  --manifest "${MANIFEST}" \
  --metrics "${METRICS}" \
  --num-gpus "${NUM_GPUS}" \
  --extract-audio "${REFERENCE_AUDIO_CACHE}" \
  --extract-workers "${EXTRACT_WORKERS}" \
  --output-format full \
  --output "${RESULTS}"

echo "FastVideo evaluation results: ${RESULTS}"
run_evaluate_vggsound_av_benchmark.sh
#!/usr/bin/env bash
set -euo pipefail

ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../../.." && pwd)"
cd "${ROOT_DIR}"

OUTPUT_DIR="${OUTPUT_DIR:-outputs/mmaudio_small_44k_ema_300000_vggsound_test}"
AUDIO_DIR="${AUDIO_DIR:-${OUTPUT_DIR}/audio}"
GT_CACHE="${GT_CACHE:-/mnt/lustre/vlm-kai/datasets/VGGSound/av_benchmark/vggsound-test-eval-cache}"
PREDICTION_CACHE="${PREDICTION_CACHE:-${OUTPUT_DIR}/av_benchmark_cache}"
RESULTS="${RESULTS:-${OUTPUT_DIR}/av_benchmark_results.json}"
AV_BENCH_PYTHON="${AV_BENCH_PYTHON:-/mnt/lustre/vlm-kai/av-benchmark/.venv/bin/python}"
BATCH_SIZE="${BATCH_SIZE:-32}"
# The official extractor initializes CUDA models before constructing its
# DataLoader. Zero avoids CUDA-after-fork hangs and this container's 64 MB shm.
NUM_WORKERS="${NUM_WORKERS:-0}"
DEVICE="${DEVICE:-cuda}"
RECOMPUTE="${RECOMPUTE:-0}"
SKIP_VIDEO_RELATED="${SKIP_VIDEO_RELATED:-0}"
# The official VGGSound GT cache has no CLAP text features, so CLAP is not an
# official VGGSound output metric. Skip its prediction passes by default.
SKIP_CLAP="${SKIP_CLAP:-1}"
# Official VGGSound precomputed cache sanitizes a subset of sample ids by
# prefixing underscores. Align only unique matches inside prediction caches.
ALIGN_PREDICTION_KEYS="${ALIGN_PREDICTION_KEYS:-1}"

if [[ ! -x "${AV_BENCH_PYTHON}" ]]; then
  echo "Missing isolated av-benchmark Python: ${AV_BENCH_PYTHON}" >&2
  echo "Set AV_BENCH_PYTHON=/path/to/av-benchmark/.venv/bin/python" >&2
  exit 2
fi
if [[ ! -d "${GT_CACHE}" ]]; then
  echo "Missing official VGGSound GT cache: ${GT_CACHE}" >&2
  echo "Set GT_CACHE=/path/to/vggsound-test-eval-cache" >&2
  exit 2
fi

EXTRA_ARGS=()
if [[ "${RECOMPUTE}" == "1" ]]; then
  EXTRA_ARGS+=(--recompute)
fi
if [[ "${SKIP_VIDEO_RELATED}" == "1" ]]; then
  EXTRA_ARGS+=(--skip-video-related)
fi
if [[ "${SKIP_CLAP}" == "1" ]]; then
  EXTRA_ARGS+=(--skip-clap)
fi
if [[ "${ALIGN_PREDICTION_KEYS}" == "1" ]]; then
  EXTRA_ARGS+=(--align-prediction-keys)
fi

.venv/bin/fastvideo eval v2a \
  --backend av-benchmark \
  --audio-dir "${AUDIO_DIR}" \
  --gt-cache "${GT_CACHE}" \
  --prediction-cache "${PREDICTION_CACHE}" \
  --output "${RESULTS}" \
  --python-executable "${AV_BENCH_PYTHON}" \
  --device "${DEVICE}" \
  --batch-size "${BATCH_SIZE}" \
  --num-workers "${NUM_WORKERS}" \
  --audio-length 8 \
  "${EXTRA_ARGS[@]}"

echo "Official av-benchmark results: ${RESULTS}"
run_inference_vggsound.sh
#!/usr/bin/env bash
set -euo pipefail

ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../../.." && pwd)"
cd "${ROOT_DIR}"
source .venv/bin/activate

MODEL_PATH="${MODEL_PATH:-converted_weights/mmaudio/small_44k_ema_300000}"
FEATURE_ROOT="${FEATURE_ROOT:-/mnt/lustre/vlm-kai/datasets/VGGSound/mmaudio_features_torio/test}"
OUTPUT_DIR="${OUTPUT_DIR:-outputs/mmaudio_small_44k_ema_300000_vggsound_test}"
NUM_GPUS="${NUM_GPUS:-4}"
NUM_WORKERS="${NUM_WORKERS:-2}"
MAX_SAMPLES="${MAX_SAMPLES:-0}"
COMPILE="${COMPILE:-1}"
MASTER_PORT="${MASTER_PORT:-29513}"
export OMP_NUM_THREADS="${OMP_NUM_THREADS:-4}"

EXTRA_ARGS=()
if [[ "${MAX_SAMPLES}" -gt 0 ]]; then
  EXTRA_ARGS+=(--max-samples "${MAX_SAMPLES}")
fi
if [[ "${COMPILE}" == "1" ]]; then
  EXTRA_ARGS+=(--compile)
fi

torchrun \
  --standalone \
  --nproc_per_node="${NUM_GPUS}" \
  --master_port="${MASTER_PORT}" \
  examples/inference/eval/mmaudio/eval_mmaudio_dataset.py \
  --model-path "${MODEL_PATH}" \
  --feature-root "${FEATURE_ROOT}" \
  --output-dir "${OUTPUT_DIR}" \
  --duration-seconds 8 \
  --num-inference-steps 25 \
  --guidance-scale 4.5 \
  --seed 14159265 \
  --num-workers "${NUM_WORKERS}" \
  "${EXTRA_ARGS[@]}"

echo "Generated audio: ${OUTPUT_DIR}/audio"
echo "Evaluation manifest: ${OUTPUT_DIR}/eval_manifest.jsonl"