BEST-RQ-2 / adapters.py
ltuncay's picture
Add Transformers loading for the existing AECC 2026 encoder
86dc2b6 verified
Raw
History Blame Contribute Delete
18.5 kB
# 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