# 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", } @dataclass(frozen=True) 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) @dataclass(frozen=True) 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 @property def sample_rate(self) -> int: return self.spec.sample_rate @property 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