Feature Extraction
Transformers
Safetensors
audio_embeddings
audio
custom_code
self-supervised-learning
audio-embeddings
best-rq-2
audioset
Instructions to use ltuncay/BEST-RQ-2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use ltuncay/BEST-RQ-2 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="ltuncay/BEST-RQ-2", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("ltuncay/BEST-RQ-2", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| # MIT License | |
| # | |
| # Copyright (c) 2026 audio-embeddings contributors | |
| # | |
| # Permission is hereby granted, free of charge, to any person obtaining a copy | |
| # of this software and associated documentation files (the "Software"), to deal | |
| # in the Software without restriction, including without limitation the rights | |
| # to use, copy, modify, merge, publish, distribute, sublicense, and/or sell | |
| # copies of the Software, and to permit persons to whom the Software is | |
| # furnished to do so, subject to the following conditions: | |
| # | |
| # The above copyright notice and this permission notice shall be included in all | |
| # copies or substantial portions of the Software. | |
| # | |
| # THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR | |
| # IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, | |
| # FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE | |
| # AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER | |
| # LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, | |
| # OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE | |
| # SOFTWARE. | |
| from __future__ import annotations | |
| from dataclasses import asdict, dataclass | |
| from typing import Any, Mapping | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from .extraction import fuse_context_windows | |
| from .extraction import get_preset | |
| from .extraction import merge_phases | |
| from .patch_embed import PatchEmbed | |
| from .spectrogram import Spectrogram | |
| from .vit import ViT | |
| from .vit import vit_config_with_patch_geometry | |
| from .waveform_feature_encoder import WaveformFeatureEncoder | |
| SPECTROGRAM_TARGETS = { | |
| "src.models.audio_jepa_module.AudioJEPAModule": "student", | |
| "src.models.rqa_jepa_module.RQAJEPAModule": "student", | |
| "src.models.best_rq_module.BestRQModule": "encoder", | |
| "src.models.best_rq2_module.BestRQ2Module": "encoder", | |
| "src.models.best_rq22_module.BestRQ22Module": "encoder", | |
| "src.models.best_rq23_module.BestRQ23Module": "encoder", | |
| } | |
| WAVEFORM_TARGETS = { | |
| "src.models.best_rq3_module.BestRQ3Module": "encoder", | |
| } | |
| class AdapterSpec: | |
| adapter_key: str | |
| model_target: str | |
| encoder_prefix: str | |
| sample_rate: int | |
| embedding_dim: int | |
| max_context_tokens: int | |
| temporal_grid_tokens: int | |
| token_hop_samples: int | |
| receptive_field_samples: int | |
| supported_phase_offsets_samples: tuple[int, ...] | |
| channel_policy: str = "mono_mean" | |
| def to_dict(self) -> dict[str, Any]: | |
| return asdict(self) | |
| class EmbeddingOutput: | |
| timestamp_embeddings: torch.Tensor | |
| timestamps_ms: torch.Tensor | |
| scene_embedding: torch.Tensor | |
| def _mapping(value: Any, path: str) -> Mapping[str, Any]: | |
| if not isinstance(value, Mapping): | |
| raise ValueError(f"Expected mapping at {path}, got {type(value).__name__}") | |
| return value | |
| def _positive_int(value: Any, path: str) -> int: | |
| try: | |
| normalized = int(value) | |
| except (TypeError, ValueError) as error: | |
| raise ValueError(f"Expected integer at {path}, got {value!r}") from error | |
| if normalized <= 0: | |
| raise ValueError(f"Expected positive integer at {path}, got {normalized}") | |
| return normalized | |
| def resolve_adapter_spec(config: Mapping[str, Any]) -> AdapterSpec: | |
| model = _mapping(config.get("model"), "model") | |
| target = str(model.get("_target_", "")) | |
| net = _mapping(model.get("net"), "model.net") | |
| encoder = _mapping(net.get("encoder"), "model.net.encoder") | |
| if target in SPECTROGRAM_TARGETS: | |
| adapter_key = "spectrogram_patch" | |
| encoder_prefix = SPECTROGRAM_TARGETS[target] | |
| frontend = _mapping( | |
| net.get("spectrogram"), | |
| "model.net.spectrogram", | |
| ) | |
| sample_rate_path = "model.net.spectrogram.sample_rate" | |
| patch = _mapping(net.get("patch_embed"), "model.net.patch_embed") | |
| patch_size = tuple(patch.get("patch_size", ())) | |
| image_size = tuple(patch.get("img_size", ())) | |
| if len(patch_size) != 2 or len(image_size) != 2: | |
| raise ValueError("Spectrogram HEAR adapters require 2-D patch/image sizes") | |
| patch_height, patch_width = map(int, patch_size) | |
| frequency_tokens = int(image_size[0]) // patch_height | |
| n_fft = _positive_int( | |
| frontend.get("n_fft", 4096), "model.net.spectrogram.n_fft" | |
| ) | |
| if frontend.get("win_length") is not None: | |
| win_length = _positive_int( | |
| frontend.get("win_length"), "model.net.spectrogram.win_length" | |
| ) | |
| elif frontend.get("win_length_ms") is not None: | |
| win_length = int( | |
| int(frontend["sample_rate"]) * float(frontend["win_length_ms"]) / 1000 | |
| ) | |
| else: | |
| win_length = n_fft | |
| if frontend.get("hop_length") is not None: | |
| frontend_hop = _positive_int( | |
| frontend.get("hop_length"), "model.net.spectrogram.hop_length" | |
| ) | |
| elif frontend.get("hop_length_ms") is not None: | |
| frontend_hop = int( | |
| int(frontend["sample_rate"]) * float(frontend["hop_length_ms"]) / 1000 | |
| ) | |
| else: | |
| frontend_hop = win_length // 2 | |
| token_hop = frontend_hop * patch_width | |
| receptive_field = n_fft + (patch_width - 1) * frontend_hop | |
| elif target in WAVEFORM_TARGETS: | |
| adapter_key = "waveform_conv" | |
| encoder_prefix = WAVEFORM_TARGETS[target] | |
| frontend = _mapping(net.get("sampling"), "model.net.sampling") | |
| sample_rate_path = "model.net.sampling.sample_rate" | |
| feature_config = dict( | |
| _mapping(net.get("feature_encoder"), "model.net.feature_encoder") | |
| ) | |
| feature_encoder = WaveformFeatureEncoder(**feature_config) | |
| receptive_field = 1 | |
| token_hop = 1 | |
| for _, kernel, stride in feature_encoder.conv_layers_spec: | |
| receptive_field += (kernel - 1) * token_hop | |
| token_hop *= stride | |
| frequency_tokens = 1 | |
| else: | |
| supported = sorted((*SPECTROGRAM_TARGETS, *WAVEFORM_TARGETS)) | |
| raise ValueError( | |
| f"No HEAR adapter is registered for model target {target!r}. " | |
| f"Register one for the new model. Current targets: {supported}" | |
| ) | |
| sample_rate = _positive_int(frontend.get("sample_rate"), sample_rate_path) | |
| data = config.get("data") | |
| if isinstance(data, Mapping) and data.get("target_sample_rate") is not None: | |
| data_sample_rate = _positive_int( | |
| data.get("target_sample_rate"), | |
| "data.target_sample_rate", | |
| ) | |
| if data_sample_rate != sample_rate: | |
| raise ValueError( | |
| "Model/data sampling-rate mismatch: " | |
| f"{sample_rate_path}={sample_rate}, " | |
| f"data.target_sample_rate={data_sample_rate}" | |
| ) | |
| max_context_tokens = _positive_int( | |
| encoder.get("num_patches"), | |
| "model.net.encoder.num_patches", | |
| ) | |
| if frequency_tokens <= 0 or max_context_tokens < frequency_tokens: | |
| raise ValueError( | |
| "Encoder context cannot hold one complete frequency-token column" | |
| ) | |
| return AdapterSpec( | |
| adapter_key=adapter_key, | |
| model_target=target, | |
| encoder_prefix=encoder_prefix, | |
| sample_rate=sample_rate, | |
| embedding_dim=_positive_int( | |
| encoder.get("embed_dim"), | |
| "model.net.encoder.embed_dim", | |
| ), | |
| max_context_tokens=max_context_tokens, | |
| temporal_grid_tokens=max_context_tokens // frequency_tokens, | |
| token_hop_samples=token_hop, | |
| receptive_field_samples=receptive_field, | |
| supported_phase_offsets_samples=( | |
| (0, token_hop // 2) if token_hop % 2 == 0 else (0,) | |
| ), | |
| ) | |
| class HearEncoderAdapter(nn.Module): | |
| spec: AdapterSpec | |
| def sample_rate(self) -> int: | |
| return self.spec.sample_rate | |
| def embedding_dim(self) -> int: | |
| return self.spec.embedding_dim | |
| def extract(self, waveform: torch.Tensor, *, preset_name: str) -> EmbeddingOutput: | |
| raise NotImplementedError | |
| def _single_waveform(waveform: torch.Tensor) -> torch.Tensor: | |
| if waveform.ndim == 1: | |
| waveform = waveform.unsqueeze(0) | |
| if waveform.ndim != 2 or waveform.shape[0] != 1: | |
| raise ValueError( | |
| "Adapter extraction expects one mono waveform [samples] or [1, samples], " | |
| f"got {tuple(waveform.shape)}" | |
| ) | |
| if waveform.shape[-1] == 0: | |
| raise ValueError("Cannot embed an empty waveform") | |
| return waveform.unsqueeze(0) | |
| class SpectrogramPatchAdapter(HearEncoderAdapter): | |
| def __init__(self, config: Mapping[str, Any], spec: AdapterSpec) -> None: | |
| super().__init__() | |
| self.spec = spec | |
| model = _mapping(config.get("model"), "model") | |
| net = _mapping(model.get("net"), "model.net") | |
| spectrogram_config = dict( | |
| _mapping(net.get("spectrogram"), "model.net.spectrogram") | |
| ) | |
| patch_config = dict(_mapping(net.get("patch_embed"), "model.net.patch_embed")) | |
| encoder_config = dict(_mapping(net.get("encoder"), "model.net.encoder")) | |
| self.spectrogram = Spectrogram(**spectrogram_config) | |
| self.patch_embed = PatchEmbed(**patch_config) | |
| self.encoder = ViT( | |
| **vit_config_with_patch_geometry( | |
| encoder_config, | |
| img_size=self.patch_embed.img_size, | |
| patch_size=self.patch_embed.patch_size, | |
| ) | |
| ) | |
| self.adjustment_mode = str(model.get("spectrogram_adjustment_mode", "pad")) | |
| if self.adjustment_mode not in {"pad", "truncate"}: | |
| raise ValueError( | |
| f"Unknown spectrogram_adjustment_mode {self.adjustment_mode!r}" | |
| ) | |
| def _phase_tokens( | |
| self, | |
| spectrogram: torch.Tensor, | |
| *, | |
| frame_offset: int, | |
| ) -> tuple[torch.Tensor, int]: | |
| patch_height, patch_width = self.patch_embed.patch_size | |
| phase = spectrogram[..., frame_offset:] | |
| original_frames = phase.shape[-1] | |
| if original_frames < patch_width: | |
| phase = F.pad(phase, (0, patch_width - original_frames)) | |
| else: | |
| remainder = original_frames % patch_width | |
| if remainder: | |
| if self.adjustment_mode == "pad": | |
| phase = F.pad(phase, (0, patch_width - remainder)) | |
| else: | |
| phase = phase[..., : original_frames - remainder] | |
| tokens = self.patch_embed(phase) | |
| frequency = phase.shape[-2] // patch_height | |
| time = phase.shape[-1] // patch_width | |
| return tokens.reshape(frequency, time, -1), original_frames | |
| def extract(self, waveform: torch.Tensor, *, preset_name: str) -> EmbeddingOutput: | |
| preset = get_preset(preset_name) | |
| waveform = _single_waveform(waveform) | |
| duration_samples = waveform.shape[-1] | |
| spectrogram = self.spectrogram(waveform) | |
| patch_width = self.patch_embed.patch_size[1] | |
| if preset.num_phases == 2 and patch_width % 2: | |
| raise ValueError( | |
| "Two-phase extraction requires an exact half-hop temporal patch offset" | |
| ) | |
| frame_offsets = [0] | |
| if preset.num_phases == 2: | |
| frame_offsets.append(patch_width // 2) | |
| hop_samples = int(self.spectrogram.mel_spec.hop_length) | |
| phases: list[tuple[torch.Tensor, torch.Tensor]] = [] | |
| for frame_offset in frame_offsets: | |
| token_grid, _ = self._phase_tokens( | |
| spectrogram, | |
| frame_offset=frame_offset, | |
| ) | |
| frequency = token_grid.shape[0] | |
| def encode_window( | |
| window: torch.Tensor, | |
| position_ids: torch.Tensor, | |
| ) -> torch.Tensor: | |
| width = window.shape[1] // frequency | |
| return self.encoder( | |
| window, | |
| pos_ids=position_ids, | |
| grid_size=(frequency, width), | |
| ) | |
| embeddings = fuse_context_windows( | |
| token_grid, | |
| max_context_tokens=self.spec.max_context_tokens, | |
| overlap=preset.overlap, | |
| encode_window=encode_window, | |
| ) | |
| centers_in_frames = ( | |
| torch.arange(embeddings.shape[0], device=embeddings.device) | |
| * patch_width | |
| + frame_offset | |
| + (patch_width - 1) / 2.0 | |
| ) | |
| centers_in_samples = centers_in_frames * hop_samples | |
| centers_in_samples = torch.clamp( | |
| centers_in_samples, | |
| max=max(0, duration_samples - 1), | |
| ) | |
| timestamps_ms = centers_in_samples * (1000.0 / self.sample_rate) | |
| phases.append((embeddings, timestamps_ms)) | |
| merged, timestamps, scene = merge_phases(phases) | |
| return EmbeddingOutput(merged, timestamps, scene) | |
| class WaveformConvAdapter(HearEncoderAdapter): | |
| def __init__(self, config: Mapping[str, Any], spec: AdapterSpec) -> None: | |
| super().__init__() | |
| self.spec = spec | |
| model = _mapping(config.get("model"), "model") | |
| net = _mapping(model.get("net"), "model.net") | |
| feature_config = dict( | |
| _mapping(net.get("feature_encoder"), "model.net.feature_encoder") | |
| ) | |
| encoder_config = dict(_mapping(net.get("encoder"), "model.net.encoder")) | |
| self.feature_encoder = WaveformFeatureEncoder(**feature_config) | |
| feature_dim = self.feature_encoder.embedding_dim | |
| self.encoder_input_proj: nn.Module | |
| if feature_dim == spec.embedding_dim: | |
| self.encoder_input_proj = nn.Identity() | |
| else: | |
| self.encoder_input_proj = nn.Linear(feature_dim, spec.embedding_dim) | |
| self.encoder = ViT(**encoder_config) | |
| receptive_field = 1 | |
| stride = 1 | |
| for _, kernel, layer_stride in self.feature_encoder.conv_layers_spec: | |
| receptive_field += (kernel - 1) * stride | |
| stride *= layer_stride | |
| self.receptive_field_samples = receptive_field | |
| self.token_hop_samples = stride | |
| if self.receptive_field_samples != spec.receptive_field_samples: | |
| raise ValueError("Waveform adapter receptive-field metadata mismatch") | |
| if self.token_hop_samples != spec.token_hop_samples: | |
| raise ValueError("Waveform adapter hop metadata mismatch") | |
| def extract(self, waveform: torch.Tensor, *, preset_name: str) -> EmbeddingOutput: | |
| preset = get_preset(preset_name) | |
| waveform = _single_waveform(waveform) | |
| duration_samples = waveform.shape[-1] | |
| offsets = [0] | |
| if preset.num_phases == 2: | |
| if self.token_hop_samples % 2: | |
| raise ValueError( | |
| "Two-phase extraction requires an exact half-hop sample offset" | |
| ) | |
| offsets.append(self.token_hop_samples // 2) | |
| phases: list[tuple[torch.Tensor, torch.Tensor]] = [] | |
| for offset in offsets: | |
| phase = waveform[..., offset:] | |
| if phase.shape[-1] < self.receptive_field_samples: | |
| phase = F.pad( | |
| phase, | |
| (0, self.receptive_field_samples - phase.shape[-1]), | |
| ) | |
| local_features = self.feature_encoder(phase) | |
| tokens = self.encoder_input_proj(local_features).squeeze(0).unsqueeze(0) | |
| def encode_window( | |
| window: torch.Tensor, | |
| position_ids: torch.Tensor, | |
| ) -> torch.Tensor: | |
| return self.encoder( | |
| window, | |
| pos_ids=position_ids, | |
| grid_size=(1, window.shape[1]), | |
| ) | |
| embeddings = fuse_context_windows( | |
| tokens, | |
| max_context_tokens=self.spec.max_context_tokens, | |
| overlap=preset.overlap, | |
| encode_window=encode_window, | |
| ) | |
| centers = ( | |
| torch.arange(embeddings.shape[0], device=embeddings.device) | |
| * self.token_hop_samples | |
| + offset | |
| + (self.receptive_field_samples - 1) / 2.0 | |
| ) | |
| centers = torch.clamp(centers, max=max(0, duration_samples - 1)) | |
| phases.append((embeddings, centers * (1000.0 / self.sample_rate))) | |
| merged, timestamps, scene = merge_phases(phases) | |
| return EmbeddingOutput(merged, timestamps, scene) | |
| def _normalized_source_state( | |
| state_dict: Mapping[str, torch.Tensor], | |
| ) -> dict[str, torch.Tensor]: | |
| normalized: dict[str, torch.Tensor] = {} | |
| for key, value in state_dict.items(): | |
| name = str(key) | |
| for prefix in ("module.", "model."): | |
| if name.startswith(prefix): | |
| name = name.removeprefix(prefix) | |
| normalized[name] = value | |
| return normalized | |
| def _load_inference_weights( | |
| adapter: HearEncoderAdapter, | |
| source_state: Mapping[str, torch.Tensor], | |
| ) -> None: | |
| source = _normalized_source_state(source_state) | |
| canonical: dict[str, torch.Tensor] = {} | |
| missing: list[str] = [] | |
| for expected_key in adapter.state_dict(): | |
| source_key = expected_key | |
| if expected_key not in source and expected_key.startswith("encoder."): | |
| source_key = ( | |
| f"{adapter.spec.encoder_prefix}.{expected_key.removeprefix('encoder.')}" | |
| ) | |
| if source_key not in source: | |
| missing.append(source_key) | |
| else: | |
| canonical[expected_key] = source[source_key] | |
| if missing: | |
| raise ValueError( | |
| "Checkpoint is missing inference weights: " + ", ".join(missing[:12]) | |
| ) | |
| adapter.load_state_dict(canonical, strict=True) | |
| def build_encoder_adapter( | |
| config: Mapping[str, Any], | |
| state_dict: Mapping[str, torch.Tensor], | |
| ) -> HearEncoderAdapter: | |
| spec = resolve_adapter_spec(config) | |
| if spec.adapter_key == "spectrogram_patch": | |
| adapter: HearEncoderAdapter = SpectrogramPatchAdapter(config, spec) | |
| elif spec.adapter_key == "waveform_conv": | |
| adapter = WaveformConvAdapter(config, spec) | |
| else: | |
| raise AssertionError(f"Unsupported registered adapter {spec.adapter_key}") | |
| _load_inference_weights(adapter, state_dict) | |
| adapter.eval() | |
| for parameter in adapter.parameters(): | |
| parameter.requires_grad = False | |
| return adapter | |