""" MATPAC model definition. This is a `transformers`-native port of the inference code from https://github.com/aurianworld/matpac/tree/main/inference_matpac (module `matpac`), following the same `trust_remote_code` pattern used by m-a-p/MERT-v0 (https://huggingface.co/m-a-p/MERT-v0). The module tree below (`patch_embed`, `cls_token`, `pos_embed`, `student_encoder`, `head_norm`/`head`, `probe`) intentionally mirrors the attribute names of the original `matpac_wrapper` so that the raw `.pt` checkpoints released on https://github.com/aurianworld/matpac/releases load with their keys unchanged. """ from dataclasses import dataclass from functools import partial from typing import Optional, Tuple, Union import numpy as np import torch import torch.nn as nn import torchaudio from einops import rearrange from timm.models.layers import trunc_normal_ from timm.models.vision_transformer import Block from transformers.modeling_outputs import ModelOutput from transformers.modeling_utils import PreTrainedModel from .configuration_matpac import MatpacConfig class MatpacLogMelSpectrogram(nn.Module): """Log-mel spectrogram front-end. Wraps `torchaudio.transforms.MelSpectrogram` with librosa-style defaults and takes `win_length`/`hop_length` in seconds. """ def __init__( self, sample_rate=16000, n_fft=400, win_length=0.025, hop_length=0.01, f_min=0.0, f_max=None, log_offset=0.001, n_mels=128, center=False, ) -> None: super().__init__() if f_max is None: f_max = sample_rate // 2 win_length = int(np.round(sample_rate * win_length)) hop_length = int(np.round(sample_rate * hop_length)) # Built explicitly on CPU: `torchaudio`'s filterbank construction mixes # tensors created with and without an explicit device, which breaks under # transformers' meta-device fast init (some end up on "meta", some on the # ambient device, and torch refuses to compare across them). with torch.device("cpu"): self.MelSpectrogram = torchaudio.transforms.MelSpectrogram( sample_rate=sample_rate, n_fft=n_fft, win_length=win_length, hop_length=hop_length, f_min=f_min, f_max=f_max, n_mels=n_mels, center=center, pad_mode="reflect", norm="slaney", mel_scale="slaney", ) self.log_offset = log_offset def forward(self, waveform): mel_spectrogram = self.MelSpectrogram(waveform) return torch.log(mel_spectrogram + self.log_offset) def _expand_size(sz): if isinstance(sz, int): return [sz, sz] return sz class MatpacPatchEmbed(nn.Module): """2D image (log-mel spectrogram) to patch embedding, borrowed from https://pypi.org/project/timm/0.4.12/. """ def __init__(self, img_size, patch_size, in_chans=1, embed_dim=768): super().__init__() img_size = _expand_size(img_size) patch_size = _expand_size(patch_size) self.img_size = img_size self.patch_size = patch_size self.grid_size = (img_size[0] // patch_size[0], img_size[1] // patch_size[1]) self.num_patches = self.grid_size[0] * self.grid_size[1] self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) def forward(self, x): x = self.proj(x) return x.flatten(2).transpose(1, 2) # BCHW -> BNC class MatpacEncoderLayers(nn.Module): """Vision transformer encoder layers (the student encoder).""" def __init__(self, config: MatpacConfig): super().__init__() self.blocks = nn.ModuleList( [ Block( config.embed_dim, config.num_heads, config.mlp_ratio, qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), ) for _ in range(config.depth) ] ) self.norm = nn.LayerNorm(config.embed_dim, eps=1e-6) def forward(self, x, return_layers=False): layers = [] for blk in self.blocks: x = blk(x) if return_layers: layers.append(x.unsqueeze(dim=1)) x = self.norm(x) if return_layers: layers[-1] = x.unsqueeze(dim=1) return torch.cat(layers, dim=1) return x @dataclass class MatpacModelOutput(ModelOutput): """ Args: last_hidden_state (`torch.FloatTensor`): Output of the last encoder layer. Shape `(batch, embed_dim)` when `pull_time_dimension=True`, else `(batch, time, embed_dim)`. hidden_states (`tuple(torch.FloatTensor)`): Output of every encoder layer, one tensor per layer (same convention as `torch.stack(outputs.hidden_states)` in the MERT model card). logits (`torch.FloatTensor`, *optional*): Classification logits, only set when `config.as_class_head=True` or `config.probe_out_features` is not `None`. """ last_hidden_state: torch.FloatTensor = None hidden_states: Optional[Tuple[torch.FloatTensor]] = None logits: Optional[torch.FloatTensor] = None class MatpacPreTrainedModel(PreTrainedModel): config_class = MatpacConfig base_model_prefix = "matpac" main_input_name = "input_values" supports_gradient_checkpointing = False def _init_weights(self, module): if isinstance(module, (nn.Linear, nn.Conv2d)): trunc_normal_(module.weight, std=0.02) if getattr(module, "bias", None) is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.LayerNorm): nn.init.ones_(module.weight) nn.init.zeros_(module.bias) class MatpacModel(MatpacPreTrainedModel): """MATPAC / MATPAC++ audio encoder. Accepts raw, mono, 16kHz waveforms as `input_values` of shape `(batch, n_samples)` and returns per-layer embeddings. See the model card for usage examples. """ def __init__(self, config: MatpacConfig): super().__init__(config) self.config = config self.patch_embed = MatpacPatchEmbed( img_size=[config.n_freq, config.n_t], patch_size=config.patch_size, embed_dim=config.embed_dim, ) num_patches = self.patch_embed.num_patches self.cls_token = nn.Parameter(torch.zeros(1, 1, config.embed_dim)) self.pos_embed = nn.Parameter( torch.zeros(1, num_patches + 1, config.embed_dim), requires_grad=False ) self.student_encoder = MatpacEncoderLayers(config) self.log_mel = MatpacLogMelSpectrogram( sample_rate=config.sample_rate, n_fft=400, win_length=0.025, hop_length=0.01, f_min=50, f_max=config.sample_rate // 2, log_offset=torch.finfo().eps, n_mels=config.n_freq, center=False, ) if config.as_class_head: self.head_norm = nn.BatchNorm1d(config.embed_dim, affine=False) self.head = nn.Linear(config.embed_dim, config.num_labels) if config.probe_out_features is not None: patch_fbins = self.grid_size()[0] probe_in_features = config.embed_dim * patch_fbins if config.concat_freq else config.embed_dim self.probe = nn.Linear(in_features=probe_in_features, out_features=config.probe_out_features) self.post_init() def grid_size(self): img_size = np.array(self.patch_embed.img_size) patch_size = np.array(self.patch_embed.patch_size) return img_size // patch_size def preprocess(self, input_values): x = input_values if x.ndim < 2: x = x.unsqueeze(dim=0) x = self.log_mel(x) x = (x - self.config.lms_mean) / self.config.lms_std return x def extract_features(self, x): if x.ndim <= 3: x = x.unsqueeze(dim=1) x = self.patch_embed(x) pos_embed = self.pos_embed[:, 1:, :] if x.shape[1] < pos_embed.shape[1]: # shorten pos_embed for a shorter-than-usual input dims = pos_embed.shape[-1] fbins = self.grid_size()[0] frames = x.shape[1] // fbins pos_embed = pos_embed.reshape(1, fbins, -1, dims)[:, :, :frames, :].reshape(1, fbins * frames, dims) x = x + pos_embed cls_token = self.cls_token + self.pos_embed[:, :1, :] cls_tokens = cls_token.expand(x.shape[0], -1, -1) x = torch.cat((cls_tokens, x), dim=1) x = self.student_encoder(x, return_layers=True) return x[:, -1, :, :], x # last layer embedding, all layers def forward_fast(self, x): """Faster, padding-heavy inference: better suited to fine-tuning or large batches, less precise than `forward_precise`. """ bs, _, _ = x.shape patch_fbins = self.grid_size()[0] unit_frames = self.config.n_t embed_d = self.patch_embed.proj.out_channels n_chunk = (x.shape[-1] + unit_frames - 1) // unit_frames pad_frames = n_chunk * unit_frames - x.shape[-1] if pad_frames > 0: x = torch.nn.functional.pad(x, (0, pad_frames)) x_full = rearrange(x, "b f (n u) -> (b n) f u", n=n_chunk, f=x.shape[-2], b=bs) _, layer_results_full = self.extract_features(x_full.unsqueeze(dim=1)) layer_results_full = layer_results_full[..., 1:, :] if self.config.concat_freq: layer_results_full = rearrange( layer_results_full, "b l (f t) d -> b l t (f d)", f=patch_fbins, d=embed_d ) layer_results_full = rearrange( layer_results_full, "(b n) l t d -> b l (t n) d", b=bs, n=n_chunk, d=embed_d * patch_fbins ) else: layer_results_full = rearrange( layer_results_full, "(b n) l t d -> b l (t n) d", b=bs, n=n_chunk, d=embed_d ) emb = layer_results_full[:, -1] return emb, layer_results_full def forward_precise(self, x): """Precise but potentially slower (loop over chunks) inference. This is the forward pass used to obtain the results reported in the MATPAC papers. """ patch_fbins = self.grid_size()[0] unit_frames = self.config.n_t patch_frames = self.patch_embed.patch_size[1] embed_d = self.patch_embed.proj.out_channels n_chunk = (x.shape[-1] + unit_frames - 1) // unit_frames pad_frames = (patch_frames - x.shape[-1] % patch_frames) % patch_frames if pad_frames > 0: x = torch.nn.functional.pad(x, (0, pad_frames)) x = x.unsqueeze(dim=1) embeddings = [] for i in range(n_chunk): _, layer_results = self.extract_features(x[..., i * unit_frames : (i + 1) * unit_frames]) layer_results = layer_results[..., 1:, :] if self.config.concat_freq: layer_results = rearrange(layer_results, "b n (f t) d -> b n t (f d)", f=patch_fbins, d=embed_d) elif self.config.as_class_head: layer_results = rearrange( layer_results, "b n (f t) d -> b n t d f", f=patch_fbins, d=embed_d ).mean(-1) embeddings.append(layer_results) layer_results = torch.cat(embeddings, axis=-2) emb = layer_results[:, -1] return emb, layer_results def forward( self, input_values: torch.Tensor, inference_type: Optional[str] = None, pull_time_dimension: Optional[bool] = None, return_dict: Optional[bool] = None, **kwargs, ) -> Union[Tuple, MatpacModelOutput]: """ Args: input_values (`torch.FloatTensor` of shape `(batch, n_samples)`): Raw, mono, 16kHz waveform. inference_type (`str`, *optional*): `"precise"` (loop over 6s/10s chunks, no padding artifacts, used for the paper's results) or `"fast"` (vectorized, some padding, faster on large batches). Defaults to `config.inference_type`. pull_time_dimension (`bool`, *optional*): Whether to mean-pool the time dimension of `last_hidden_state` and `hidden_states`. Defaults to `config.pull_time_dimension`. """ inference_type = inference_type if inference_type is not None else self.config.inference_type pull_time_dimension = ( pull_time_dimension if pull_time_dimension is not None else self.config.pull_time_dimension ) return_dict = return_dict if return_dict is not None else self.config.use_return_dict x = self.preprocess(input_values) if x.ndim > 3: x = x.squeeze(dim=1) if inference_type == "precise": last_layer, layer_results = self.forward_precise(x) elif inference_type == "fast": last_layer, layer_results = self.forward_fast(x) else: raise ValueError(f"inference_type must be 'precise' or 'fast', got {inference_type!r}") # The classification/probe heads always consume the time-pooled embedding of # the last layer, regardless of `pull_time_dimension` (mirrors the original # `probe_forward`). logits = None pooled_last_layer = last_layer.mean(dim=1) if last_layer.ndim == 3 else last_layer if self.config.as_class_head: pooled = self.head_norm(pooled_last_layer.unsqueeze(-1)).squeeze(-1) logits = self.head(pooled) elif self.config.probe_out_features is not None: logits = self.probe(pooled_last_layer) if pull_time_dimension: last_layer = last_layer.mean(dim=1) layer_results = layer_results.mean(dim=2) hidden_states = tuple(layer_results.unbind(dim=1)) last_hidden_state = hidden_states[-1] if not return_dict: return tuple(v for v in (last_hidden_state, hidden_states, logits) if v is not None) return MatpacModelOutput( last_hidden_state=last_hidden_state, hidden_states=hidden_states, logits=logits, )