Download inference_example.py from ltuncay/Audio-JEPA: direct link, hf CLI and curl.
- Browser
- Download file 7.77 kB
-
https://huggingface.co/ltuncay/Audio-JEPA/resolve/main/inference_example.py
- Command line
-
hf download hf://ltuncay/Audio-JEPA/inference_example.py
-
curl -L -o inference_example.py https://huggingface.co/ltuncay/Audio-JEPA/resolve/main/inference_example.py
7.77 kB
| """Minimal standalone inference for Audio-JEPA (ltuncay/Audio-JEPA). | |
| Loads the JEPA.ckpt weights, instantiates the ViT encoder with the correct | |
| input shape, and produces (T, D) embeddings from a mono waveform. | |
| Runs on CPU without installing flash-attn CUDA kernels: a small shim | |
| substitutes flash_attn.modules.mha.MHA with a torch-native equivalent that | |
| matches the checkpoint parameter names (qkv, proj). | |
| Requires the upstream code cloned locally: | |
| git clone --depth 1 https://github.com/LudovicTuncay/Audio-JEPA.git /path/to/audio-jepa | |
| Usage: | |
| python inference_example.py --audio-jepa-src /path/to/audio-jepa | |
| or, from a directory containing this file next to ``audio-jepa/``: | |
| python inference_example.py | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import sys | |
| import time | |
| import types | |
| from importlib.machinery import ModuleSpec | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torchaudio | |
| from huggingface_hub import hf_hub_download | |
| # Config values matching the shipped checkpoint (confirmed by the author). | |
| SAMPLE_RATE = 32_000 | |
| CLIP_LENGTH_S = 10 | |
| N_MELS = 128 | |
| TARGET_TIME_BINS = 256 | |
| PATCH_SIZE = (16, 16) | |
| EMBED_DIM = 768 | |
| DEPTH = 12 | |
| NUM_HEADS = 12 | |
| MLP_RATIO = 4.0 | |
| def install_flash_attn_shim() -> None: | |
| """Replace ``flash_attn.modules.mha.MHA`` with a CPU-friendly torch class. | |
| Matches the checkpoint's parameter naming (``qkv``, ``proj``) so | |
| ``load_state_dict(..., strict=True)`` succeeds without any CUDA build. | |
| """ | |
| class CpuMultiHeadAttention(nn.Module): | |
| def __init__( | |
| self, | |
| embed_dim: int, | |
| num_heads: int, | |
| dropout: float = 0.0, | |
| qkv_proj_bias: bool = True, | |
| use_flash_attn: bool = False, | |
| **_ignore, | |
| ) -> None: | |
| super().__init__() | |
| assert embed_dim % num_heads == 0 | |
| self.embed_dim = embed_dim | |
| self.num_heads = num_heads | |
| self.head_dim = embed_dim // num_heads | |
| self.dropout = dropout | |
| self.qkv = nn.Linear(embed_dim, 3 * embed_dim, bias=qkv_proj_bias) | |
| self.proj = nn.Linear(embed_dim, embed_dim, bias=True) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| b, n, d = x.shape | |
| qkv = self.qkv(x).reshape(b, n, 3, self.num_heads, self.head_dim) | |
| q, k, v = qkv.permute(2, 0, 3, 1, 4).unbind(0) | |
| out = torch.nn.functional.scaled_dot_product_attention( | |
| q, k, v, dropout_p=self.dropout if self.training else 0.0 | |
| ) | |
| return self.proj(out.transpose(1, 2).reshape(b, n, d)) | |
| def make(name: str) -> types.ModuleType: | |
| m = types.ModuleType(name) | |
| m.__spec__ = ModuleSpec(name, loader=None) | |
| return m | |
| flash_attn = make("flash_attn") | |
| flash_attn.__version__ = "0.0.0-cpu-shim" | |
| modules = make("flash_attn.modules") | |
| mha = make("flash_attn.modules.mha") | |
| mha.MHA = CpuMultiHeadAttention | |
| sys.modules.update( | |
| { | |
| "flash_attn": flash_attn, | |
| "flash_attn.modules": modules, | |
| "flash_attn.modules.mha": mha, | |
| } | |
| ) | |
| def stub_upstream_inits(root: Path) -> None: | |
| """Pre-empt heavy ``__init__.py`` files in the upstream that pull hydra/wandb. | |
| Only the leaf model file and the mel-spec transform are needed for inference. | |
| """ | |
| for cached in list(sys.modules): | |
| if cached == "src" or cached.startswith("src."): | |
| del sys.modules[cached] | |
| for name in ( | |
| "src", | |
| "src.utils", | |
| "src.models", | |
| "src.models.components", | |
| "src.masks", | |
| "src.masks.components", | |
| "src.data", | |
| "src.data.components", | |
| ): | |
| m = types.ModuleType(name) | |
| m.__path__ = [str(root / name.replace(".", "/"))] | |
| sys.modules[name] = m | |
| def compute_mel_spec(waveform: torch.Tensor) -> torch.Tensor: | |
| """Kaldi-fbank mel spectrogram of shape (1, TARGET_TIME_BINS, N_MELS).""" | |
| hop_length_ms = (CLIP_LENGTH_S * 1000) / TARGET_TIME_BINS | |
| frame_length_ms = 2.5 * hop_length_ms | |
| spec = torchaudio.compliance.kaldi.fbank( | |
| waveform - waveform.mean(), | |
| sample_frequency=SAMPLE_RATE, | |
| frame_length=frame_length_ms, | |
| frame_shift=hop_length_ms, | |
| num_mel_bins=N_MELS, | |
| low_freq=20, | |
| high_freq=SAMPLE_RATE // 2, | |
| use_log_fbank=True, | |
| window_type="hanning", | |
| ) | |
| if spec.shape[0] < TARGET_TIME_BINS: | |
| pad = TARGET_TIME_BINS - spec.shape[0] | |
| spec = torch.cat([spec, torch.zeros(pad, N_MELS)], dim=0) | |
| elif spec.shape[0] > TARGET_TIME_BINS: | |
| spec = spec[:TARGET_TIME_BINS] | |
| return spec.unsqueeze(0) | |
| def main() -> int: | |
| parser = argparse.ArgumentParser(description="Audio-JEPA CPU inference example.") | |
| parser.add_argument( | |
| "--audio-jepa-src", | |
| default="./audio-jepa", | |
| help="Path to the cloned LudovicTuncay/Audio-JEPA repo (default: ./audio-jepa).", | |
| ) | |
| parser.add_argument( | |
| "--wav", | |
| default=None, | |
| help="Optional WAV file to encode (mono, will be resampled to 32 kHz).", | |
| ) | |
| args = parser.parse_args() | |
| root = Path(args.audio_jepa_src).resolve() | |
| if not (root / "src" / "models" / "components" / "vision_transformer.py").exists(): | |
| print(f"ERROR: Audio-JEPA source not found at {root}", file=sys.stderr) | |
| print("Run: git clone --depth 1 https://github.com/LudovicTuncay/Audio-JEPA.git", file=sys.stderr) | |
| return 2 | |
| install_flash_attn_shim() | |
| if str(root) not in sys.path: | |
| sys.path.insert(0, str(root)) | |
| stub_upstream_inits(root) | |
| from src.models.components.vision_transformer import VisionTransformer | |
| print("Building encoder (input_size=(256, 128), patch=(16, 16))...") | |
| encoder = VisionTransformer( | |
| input_size=(TARGET_TIME_BINS, N_MELS), | |
| patch_size=PATCH_SIZE, | |
| in_chans=1, | |
| embed_dim=EMBED_DIM, | |
| depth=DEPTH, | |
| num_heads=NUM_HEADS, | |
| mlp_ratio=MLP_RATIO, | |
| use_flash_attn=False, | |
| ) | |
| print("Downloading checkpoint (JEPA.ckpt, ~350 MB, first run only)...") | |
| ckpt_path = hf_hub_download("ltuncay/Audio-JEPA", "JEPA.ckpt") | |
| state = torch.load(ckpt_path, map_location="cpu", weights_only=False) | |
| raw_sd = state.get("state_dict", state) | |
| encoder_sd = { | |
| k[len("encoder.") :]: v | |
| for k, v in raw_sd.items() | |
| if k.startswith("encoder.") and not k.startswith("encoder_") | |
| } | |
| missing, unexpected = encoder.load_state_dict(encoder_sd, strict=True) | |
| print(f"Loaded: {len(encoder_sd)} keys, {len(missing)} missing, {len(unexpected)} unexpected.") | |
| encoder.eval() | |
| if args.wav: | |
| waveform, sr = torchaudio.load(args.wav) | |
| if waveform.shape[0] > 1: | |
| waveform = waveform.mean(dim=0, keepdim=True) | |
| if sr != SAMPLE_RATE: | |
| waveform = torchaudio.functional.resample(waveform, sr, SAMPLE_RATE) | |
| else: | |
| n = SAMPLE_RATE * 5 | |
| t = torch.arange(n).float() / SAMPLE_RATE | |
| waveform = (0.3 * torch.sin(2 * torch.pi * 440 * t)).unsqueeze(0) | |
| print(f"No --wav given, using a 5 s synthetic tone at {SAMPLE_RATE} Hz.") | |
| duration_s = waveform.shape[1] / SAMPLE_RATE | |
| print(f"Waveform: {waveform.shape[1]} samples ({duration_s:.2f} s)") | |
| spec = compute_mel_spec(waveform).unsqueeze(0) | |
| t0 = time.perf_counter() | |
| with torch.inference_mode(): | |
| emb = encoder(spec) | |
| wall = time.perf_counter() - t0 | |
| print(f"Embeddings: shape={tuple(emb.shape)} wall={wall*1000:.1f} ms RTF={wall/duration_s:.3f}") | |
| print("Note: the encoder always produces 128 patches (8 temporal x 16 frequency).") | |
| return 0 | |
| if __name__ == "__main__": | |
| sys.exit(main()) | |