import collections import collections.abc import math from functools import partial from typing import Any, Callable, Literal, Optional, Tuple, Type, Union import torch import torch.nn as nn import torchaudio from einops import pack, rearrange def to_2tuple(x: Any) -> Tuple[Any, Any]: if isinstance(x, collections.abc.Iterable): return x return (x, x) Conv_Kernel = Union[int, Tuple[int, int]] class AudioPatchEmbed(nn.Module): def __init__( self, input_size: Conv_Kernel = 64, patch_size: Conv_Kernel = 16, patch_stride: Conv_Kernel = 16, in_chans: int = 1, embed_dim: int = 768, norm_layer: Optional[Callable] = None, flatten: bool = False, ): super().__init__() self.input_size = to_2tuple(input_size) self.patch_size = to_2tuple(patch_size) self.patch_stride = to_2tuple(patch_stride) self.grid_size = ( self.input_size[0] // self.patch_stride[0], self.input_size[1] // self.patch_stride[1], ) self.num_patches = self.grid_size[0] * self.grid_size[1] self.flatten = flatten self.proj = nn.Conv2d( in_chans, embed_dim, kernel_size=patch_size, stride=patch_stride ) self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity() def forward(self, x): x = self.proj(x) if self.flatten: x = rearrange(x, "b c f t -> b (f t) c") x = self.norm(x) return x class Mlp(nn.Module): def __init__( self, in_features: int, hidden_features: Optional[int] = None, out_features: Optional[int] = None, act_layer: Type[torch.nn.Module] = nn.GELU, drop: float = 0.0, ): super().__init__() out_features = out_features or in_features hidden_features = hidden_features or in_features self.fc1 = nn.Linear(in_features, hidden_features) self.act = act_layer() self.fc2 = nn.Linear(hidden_features, out_features) self.drop = nn.Dropout(drop) def forward(self, x): x = self.fc1(x) x = self.act(x) x = self.drop(x) x = self.fc2(x) x = self.drop(x) return x def drop_path(x, drop_prob: float = 0.0, training: bool = False, scale_by_keep: bool = True): if drop_prob == 0.0 or not training: return x keep_prob = 1 - drop_prob shape = (x.shape[0],) + (1,) * (x.ndim - 1) random_tensor = x.new_empty(shape).bernoulli_(keep_prob) if keep_prob > 0.0 and scale_by_keep: random_tensor.div_(keep_prob) return x * random_tensor class DropPath(nn.Module): def __init__(self, drop_prob: float = 0.0, scale_by_keep: bool = True): super().__init__() self.drop_prob = drop_prob self.scale_by_keep = scale_by_keep def forward(self, x): return drop_path(x, self.drop_prob, self.training, self.scale_by_keep) def extra_repr(self): return f"drop_prob={round(self.drop_prob, 3):0.3f}" def _no_grad_trunc_normal_(tensor, mean, std, a, b): def norm_cdf(x): return (1.0 + math.erf(x / math.sqrt(2.0))) / 2.0 with torch.no_grad(): l = norm_cdf((a - mean) / std) u = norm_cdf((b - mean) / std) tensor.uniform_(2 * l - 1, 2 * u - 1) tensor.erfinv_() tensor.mul_(std * math.sqrt(2.0)) tensor.add_(mean) tensor.clamp_(min=a, max=b) return tensor def trunc_normal_(tensor, mean=0.0, std=1.0, a=-2.0, b=2.0): return _no_grad_trunc_normal_(tensor, mean, std, a, b) class LayerScale(nn.Module): def __init__(self, dim, init_values=1e-5, inplace=False): super().__init__() self.inplace = inplace self.gamma = nn.Parameter(init_values * torch.ones(dim)) def forward(self, x): return x.mul_(self.gamma) if self.inplace else x * self.gamma class KwargsSequential(nn.Sequential): def forward(self, x, **kwargs): for module in self._modules.values(): x = module(x, **kwargs) return x class Attention(nn.Module): def __init__( self, dim: int, num_heads: int = 8, qkv_bias: bool = False, attn_drop: float = 0.0, proj_drop: float = 0.0, causal: bool = False, ): super().__init__() assert dim % num_heads == 0, "dim should be divisible by num_heads" self.num_heads = num_heads head_dim = dim // num_heads self.scale = head_dim ** -0.5 self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) self.attn_drop = nn.Dropout(attn_drop) self.proj = nn.Linear(dim, dim) self.proj_drop = nn.Dropout(proj_drop) self.causal = causal def forward(self, x, mask: Optional[torch.Tensor] = None): B, N, C = x.shape qkv = ( self.qkv(x) .reshape(B, N, 3, self.num_heads, C // self.num_heads) .permute(2, 0, 3, 1, 4) ) q, k, v = qkv[0], qkv[1], qkv[2] attn = (q @ k.transpose(-2, -1)) * self.scale if self.causal: mask_value = -torch.finfo(attn.dtype).max i, j = attn.shape[-2:] causal_mask = torch.ones(i, j, device=q.device, dtype=torch.bool).triu(j - i + 1) attn = attn.masked_fill(causal_mask, mask_value) if mask is not None: mask_value = torch.finfo(attn.dtype).min attn_mask = mask[:, None, None, :].expand(B, 1, N, N) attn = attn.masked_fill(attn_mask, mask_value) attn = attn.softmax(dim=-1) attn = torch.nan_to_num(attn) attn = self.attn_drop(attn) x = (attn @ v).transpose(1, 2).reshape(B, N, C) x = self.proj(x) x = self.proj_drop(x) return x class Block(nn.Module): def __init__( self, dim: int, num_heads: int, mlp_ratio: float = 4.0, qkv_bias: bool = False, drop: float = 0.0, attn_drop: float = 0.0, init_values=None, drop_path: float = 0.0, act_layer: Type[torch.nn.Module] = nn.GELU, norm_layer: Type[torch.nn.Module] = nn.LayerNorm, attention_type: Type[torch.nn.Module] = Attention, attention_kwargs={}, **kwargs, ): super().__init__() self.norm1 = norm_layer(dim) self.attn = attention_type( dim, num_heads=num_heads, qkv_bias=qkv_bias, attn_drop=attn_drop, proj_drop=drop, **attention_kwargs, ) self.ls1 = LayerScale(dim, init_values=init_values) if init_values else nn.Identity() self.drop_path1 = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() self.norm2 = norm_layer(dim) self.mlp = Mlp( in_features=dim, hidden_features=int(dim * mlp_ratio), act_layer=act_layer, drop=drop, ) self.ls2 = LayerScale(dim, init_values=init_values) if init_values else nn.Identity() self.drop_path2 = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() def forward(self, x, **kwargs): x = x + self.drop_path1(self.ls1(self.attn(self.norm1(x), **kwargs))) x = x + self.drop_path2(self.ls2(self.mlp(self.norm2(x)))) return x def drop_patches(x: torch.Tensor, dim: int, frac: float) -> torch.Tensor: N = x.shape[dim] to_keep = N - int(N * frac) random_mask = torch.randperm(N, device=x.device)[:to_keep].sort().values return x.index_select(dim=dim, index=random_mask) def calculate_padding( input_length: Union[int, torch.Tensor], target_length: Union[int, torch.Tensor] ) -> Union[int, torch.Tensor]: return (target_length - (input_length % target_length)) % target_length class AudioTransformer(nn.Module): def __init__( self, outputdim: int = 527, patch_size: Union[int, Tuple[int, int]] = 16, patch_stride: Union[int, Tuple[int, int]] = 16, embed_dim: int = 768, depth: int = 12, num_heads: int = 12, mlp_ratio: float = 4.0, qkv_bias: bool = True, drop_rate: float = 0.0, attn_drop_rate: float = 0.0, drop_path_rate: float = 0.0, init_bn: bool = True, norm_layer: Optional[torch.nn.Module] = None, act_layer: Type[torch.nn.Module] = nn.GELU, init_values=None, target_length: int = 1012, input_channels: int = 1, pooling: Optional[Literal["mean", "token", "dm", "logit", "cat"]] = "token", wavtransforms: Optional[Callable] = None, spectransforms: Optional[Callable] = None, time_patch_out: Optional[float] = None, freq_patch_out: Optional[float] = None, block_type: Type[torch.nn.Module] = Block, attention_type: Type[torch.nn.Module] = Attention, eval_avg: Literal["mean", "max", "cat"] = "mean", **kwargs, ): super().__init__() assert pooling in ("mean", "token", "dm", "logit", "cat", None) self.outputdim = outputdim self.pooling = pooling self.embed_dim = embed_dim self.depth = depth self.patch_stride = patch_stride self.patch_size = patch_size self.n_mels = kwargs.get("n_mels", 64) self.n_fft = kwargs.get("n_fft", 512) self.hop_size = kwargs.get("hop_size", 160) self.win_size = kwargs.get("win_size", 512) self.f_min = kwargs.get("f_min", 0) self.f_max = kwargs.get("f_max", 8000) self.sample_rate = kwargs.get("sample_rate", 16000) self.center = kwargs.get("center", True) self.pad_last = kwargs.get("pad_last", True) self.input_channels = input_channels self.eval_avg = eval_avg self.time_patch_out = time_patch_out self.freq_patch_out = freq_patch_out self.target_length = target_length patch_stride = to_2tuple(self.patch_stride)[-1] self.maximal_allowed_length = self.target_length self.patch_embed = AudioPatchEmbed( input_size=(self.n_mels, target_length), embed_dim=self.embed_dim, in_chans=self.input_channels, patch_size=self.patch_size, flatten=False, patch_stride=self.patch_stride, ) self.spectransforms = nn.Sequential() if spectransforms is None else spectransforms self.wavtransforms = nn.Sequential() if wavtransforms is None else wavtransforms if self.pooling == "token": self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.token_pos_embed = nn.Parameter(torch.randn(1, embed_dim) * 0.02) self.time_pos_embed = nn.Parameter( torch.randn(1, embed_dim, 1, self.patch_embed.grid_size[1]) * 0.02 ) self.freq_pos_embed = nn.Parameter( torch.randn(1, embed_dim, self.patch_embed.grid_size[0], 1) * 0.02 ) norm_layer = norm_layer or partial(nn.LayerNorm, eps=1e-6) act_layer = act_layer or nn.GELU dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)] self.pos_drop = nn.Dropout(p=drop_rate) self.blocks = KwargsSequential( *[ block_type( dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, init_values=init_values, drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[i], norm_layer=norm_layer, act_layer=act_layer, attention_type=attention_type, ) for i in range(depth) ] ) self.norm = norm_layer(embed_dim) self.outputlayer = nn.Identity() self.apply(self.init_weights) if hasattr(self, "cls_token"): nn.init.normal_(self.cls_token, std=1e-6) def init_weights(self, module): if isinstance(module, nn.Linear): trunc_normal_(module.weight, std=0.02) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.LayerNorm): nn.init.constant_(module.bias, 0) nn.init.constant_(module.weight, 1.0) def forward_features(self, x: torch.Tensor, **kwargs) -> torch.Tensor: x = self.patch_embed(x) b, c, f, t = x.shape x = x + self.time_pos_embed[:, :, :, :t] x = x + self.freq_pos_embed[:, :, :, :] if self.training and self.time_patch_out is not None: x = drop_patches(x, dim=-1, frac=self.time_patch_out) if self.training and self.freq_patch_out is not None: x = drop_patches(x, dim=-2, frac=self.freq_patch_out) x = rearrange(x, "b c f t -> b (f t) c") if self.pooling == "token": cls_token = self.cls_token.expand(x.shape[0], -1, -1) cls_token = cls_token + self.token_pos_embed x = torch.cat((cls_token, x), dim=1) x = self.pos_drop(x) x = self.blocks(x, **kwargs) x = self.norm(x) return x def forward_head(self, x: torch.Tensor, **kwargs) -> torch.Tensor: mask = kwargs.get("mask", None) if self.pooling == "token": x = x[:, 0] return self.outputlayer(x).sigmoid() elif self.pooling == "mean": if mask is not None: m = (1.0 - mask.float()).unsqueeze(-1) x = torch.nan_to_num((x * m).sum(1) / m.sum(1)) else: x = x.mean(1) return self.outputlayer(x).sigmoid() elif self.pooling == "logit": if mask is not None: m = (1.0 - mask.float()).unsqueeze(-1) x = torch.nan_to_num((x * m).sum(1) / m.sum(1)) else: x = x.mean(1) return self.outputlayer(x) elif self.pooling == "dm": x = rearrange(x, "b (f t) d -> b f t d", f=self.patch_embed.grid_size[0]) x = self.outputlayer(x.mean(1)).sigmoid() return x.mean(1) elif self.pooling is None: return x else: return x.mean(1) def _audiosample_to_mellength(self, lengths: torch.Tensor) -> torch.Tensor: if self.center: lengths = lengths + self.win_size lengths = 1 + ((lengths - self.win_size) / self.hop_size).long() return lengths def _audiosample_to_patchlength(self, lengths: torch.Tensor) -> torch.Tensor: lengths = self._audiosample_to_mellength(lengths) return self._frames_to_patchlength(lengths) def _frames_to_patchlength(self, lengths: torch.Tensor) -> torch.Tensor: patch_stride = to_2tuple(self.patch_stride) patch_size = to_2tuple(self.patch_size) frequency_patch_size = self.n_mels // patch_stride[0] time_patch_size = patch_stride[1] time_window_size = patch_size[1] number_of_tokens = ( torch.floor((lengths - time_window_size) / time_patch_size) + 1 ) * frequency_patch_size if self.pooling == "token": number_of_tokens += 1 return number_of_tokens def _reshape_mask_to_ft_format(self, mask: torch.Tensor) -> torch.Tensor: n_freq_patches = self.n_mels // to_2tuple(self.patch_stride)[0] mask = mask.reshape(-1, n_freq_patches).transpose(-2, -1).flatten(-2).reshape_as(mask) return mask def _to_binary_mask(self, lengths: torch.Tensor, max_length: int) -> torch.Tensor: batch_size = len(lengths) lengths = self._audiosample_to_patchlength(lengths) idx = torch.arange(max_length, device=lengths.device) idx = idx.repeat(batch_size).view(batch_size, max_length) mask = (idx >= lengths.unsqueeze(-1)).bool() return mask def _create_mask(self, x_length, audio_length_in_spec_frames: int): max_length_in_patches = self._frames_to_patchlength( torch.tensor(audio_length_in_spec_frames) ) mask_1d = self._to_binary_mask(x_length, max_length=int(max_length_in_patches)) return mask_1d def _forward_spec0(self, x: torch.Tensor, x_length: Optional[torch.Tensor] = None): input_length_in_frames = x.shape[-1] if input_length_in_frames > self.maximal_allowed_length: if self.pad_last: to_pad = int(calculate_padding(input_length_in_frames, self.target_length)) x = torch.nn.functional.pad(x, (0, to_pad), value=0) else: valid_length_frames = input_length_in_frames - ( input_length_in_frames % self.target_length ) x = x[..., :valid_length_frames] forward_kwargs, mask = {}, None if x_length is not None: assert len(x_length) == len(x), "batchsizes of input x and x_length need to be same" assert x_length.ndim == 1, "Lengths are of size (B,)" mask = self._create_mask( x_length=x_length, audio_length_in_spec_frames=x.shape[-1] ) if input_length_in_frames > self.maximal_allowed_length: target_length_in_patches = self._frames_to_patchlength( torch.tensor([self.target_length]) ) valid_length = int(mask.shape[-1] - (mask.shape[-1] % target_length_in_patches)) mask = mask[..., :valid_length] mask = map( self._reshape_mask_to_ft_format, mask.split( self._frames_to_patchlength(torch.tensor(self.target_length)), dim=-1 ), ) mask, ps_mask = pack(list(mask), "* d") forward_kwargs["mask"] = mask has_splits = False splits = None if x.shape[-1] > self.target_length: splits = x.split(self.target_length, dim=-1) has_splits = len(splits) > 1 x, x_pack_size = pack(splits, "* c f t") x = self.patch_embed(x) b, c, f, t = x.shape x = x + self.time_pos_embed[:, :, :, :t] x = x + self.freq_pos_embed[:, :, :, :] if self.training and self.time_patch_out is not None: x = drop_patches(x, dim=-1, frac=self.time_patch_out) if self.training and self.freq_patch_out is not None: x = drop_patches(x, dim=-2, frac=self.freq_patch_out) x = rearrange(x, "b c f t -> b (f t) c") if self.pooling == "token": cls_token = self.cls_token.expand(x.shape[0], -1, -1) cls_token = cls_token + self.token_pos_embed x = torch.cat((cls_token, x), dim=1) x = self.pos_drop(x) return x, has_splits, forward_kwargs, mask, splits def _forward_spec1( self, x: torch.Tensor, x_length: Optional[torch.Tensor] = None, has_splits: bool = False, masks: Optional[torch.Tensor] = None, splits=None, **forward_kwargs, ): x = self.norm(x) x = self.forward_head(x, **forward_kwargs) if has_splits: if self.eval_avg == "mean": if x_length is not None and masks is not None: mask_ = masks.all(-1).reshape(x.shape[0], *((1,) * (x.ndim - 1))) x.masked_fill_(mask_, float("nan")) x = rearrange(x, "(spl b) ... -> spl b ...", spl=len(splits)) x = x.nanmean(0) elif self.eval_avg == "max": if x_length is not None and masks is not None: mask_ = masks.all(-1).reshape(x.shape[0], *((1,) * (x.ndim - 1))) x.masked_fill_(mask_, -float("inf")) x = rearrange(x, "(spl b) ... -> spl b ...", spl=len(splits)) x = x.max(0)[0] elif self.eval_avg == "cat": if x_length is not None and masks is not None: mask_ = masks.all(-1).reshape(x.shape[0], *((1,) * (x.ndim - 1))) x.masked_fill_(mask_, 0.0) x = rearrange(x, "(spl b) ... d -> b (spl ...) d", spl=len(splits)) else: raise ValueError(f"Unknown eval_avg function {self.eval_avg}") return x def _forward_spectrogram(self, x: torch.Tensor, x_length: Optional[torch.Tensor] = None): x, has_splits, forward_kwargs, mask, splits = self._forward_spec0(x, x_length) x = self.blocks(x, **forward_kwargs) x = self._forward_spec1(x, x_length, has_splits, mask, splits, **forward_kwargs) return x def forward_spectrogram( self, x: torch.Tensor, x_length: Optional[torch.Tensor] = None ) -> torch.Tensor: return self._forward_spectrogram(x, x_length) def forward(self, x: torch.Tensor, x_length: Optional[torch.Tensor] = None) -> torch.Tensor: x = self.forward_spectrogram(x, x_length=x_length) return x OTHER_KWARGS = { "target_length": 1008, "pooling": None, "eval_avg": "cat", "sample_rate": 16000, "input_size_spectrogram": (1, 64, 1008), "patch_size": [64, 4], "patch_stride": [64, 4], "init_bn": False, } def audiotransformer_base(**kwargs) -> AudioTransformer: model_kwargs = dict( embed_dim=768, depth=12, num_heads=12, pooling="mean", init_bn=True, drop_path_rate=0.0 ) model_kwargs.update(OTHER_KWARGS) model_kwargs.update(kwargs) return AudioTransformer(**model_kwargs) def audiotransformer_huge(**kwargs) -> AudioTransformer: model_kwargs = dict( embed_dim=1280, depth=32, num_heads=16, pooling="mean", init_bn=True, drop_path_rate=0.0 ) model_kwargs.update(OTHER_KWARGS) model_kwargs.update(kwargs) return AudioTransformer(**model_kwargs) class PatchedDasheng(AudioTransformer): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) def forward_features(self, x: torch.Tensor, **kwargs) -> torch.Tensor: t = x.shape[-1] x = x + self.time_pos_embed[:, :, :, :t] x = x + self.freq_pos_embed[:, :, :, :] x = rearrange(x, "b c f t -> b (f t) c") if self.pooling == "token": cls_token = self.cls_token.expand(x.shape[0], -1, -1) cls_token = cls_token + self.token_pos_embed x = torch.cat((cls_token, x), dim=1) x = self.pos_drop(x) x = self.blocks(x, **kwargs) x = self.norm(x) return x def _to_mask(self, lengths: torch.Tensor, max_length: int) -> torch.Tensor: batch_size = len(lengths) idx = torch.arange(max_length, device=lengths.device) idx = idx.repeat(batch_size).view(batch_size, max_length) mask = (idx >= lengths.unsqueeze(-1)).bool() return mask def _forward_spectrogram(self, x: torch.Tensor, x_length: Optional[torch.Tensor] = None): target_length_in_patches = self.target_length // 4 x = self.patch_embed(x) b, c, f, t = x.shape input_splits = x.split(target_length_in_patches, dim=-1) mask = None masks = [None for _ in range(len(input_splits))] if x_length is not None: assert len(x_length) == len(x), "batchsizes of input x and x_length need to be same" assert x_length.ndim == 1, "Lengths are of size (B,)" scaled_lengths = (x_length / (self.hop_size * 4)).long() mask = self._to_mask(max_length=t, lengths=scaled_lengths) masks = mask.split(target_length_in_patches, dim=-1) outputs = [] for split_x, mask in zip(input_splits, masks): forward_kwargs = {} forward_kwargs["mask"] = mask split_x = self.forward_features(split_x, **forward_kwargs) split_x = self.forward_head(split_x, **forward_kwargs) outputs.append(split_x) x = torch.cat(outputs, dim=1) return x class DashengAudioEncoder(nn.Module): def __init__(self, append_cls_token: bool = False, transformer: Optional[AudioTransformer] = None): super().__init__() self.append_cls_token = append_cls_token # `transformer` lets the HF package build the inner AudioTransformer from # config (variable dim/depth); the default reproduces the training repo's # `DashengAudioEncoder()` (Dasheng-Huge) exactly. self.model = transformer if transformer is not None else audiotransformer_huge() self.embed_dim = self.model.embed_dim self.model.outputlayer = torch.nn.Identity() self.model.__class__ = PatchedDasheng def _to_mask(self, lengths: torch.Tensor, max_length: int) -> torch.Tensor: batch_size = len(lengths) idx = torch.arange(max_length, device=lengths.device) idx = idx.repeat(batch_size).view(batch_size, max_length) mask = (idx < lengths.unsqueeze(-1)).long() return mask def _create_encoder_attention_mask(self, model_output: torch.Tensor, input_lengths: torch.Tensor): scaled_lengths = (input_lengths / (self.model.hop_size * 4)).long() return self._to_mask(scaled_lengths, max_length=model_output.shape[1]) def forward( self, input: torch.Tensor, input_length: Optional[torch.Tensor] = None, return_attention_mask: bool = False, ): emb = self.model(input, input_length) if input_length is not None: input_length = input_length + self.model.n_fft scaled_lengths = ( (1 + (input_length - self.model.n_fft) / self.model.hop_size) // 4 ).long() max_length = torch.max(scaled_lengths) emb = emb[:, :max_length, :] if self.append_cls_token: emb = torch.cat([emb.mean(1, keepdims=True), emb], dim=1) if return_attention_mask and input_length is not None: return emb, self._create_encoder_attention_mask(emb, input_length) return emb class FrontEndFeatureExtractor(nn.Module): """Feature extractor for audio front-end processing.""" def __init__(self, feature="LogMel", sample_rate=16000, n_mels=64, n_fft=512, hop_length=160): super().__init__() self.feature = feature self.stft_extractor = torchaudio.transforms.Spectrogram( n_fft=n_fft, hop_length=hop_length, win_length=n_fft, power=None ) self.mel_scale = torchaudio.transforms.MelScale( n_mels=n_mels, sample_rate=sample_rate, norm="slaney", n_stft=n_fft // 2 + 1 ) self.amp2db = torchaudio.transforms.AmplitudeToDB(stype="power", top_db=120) if feature in ["LogMel"]: self.num_channels = 2 elif feature in ["MonoLogMel"]: self.num_channels = 1 else: raise ValueError(f"Unsupported feature type: {feature}") @torch.compiler.disable(recursive=True) def forward(self, waveform): x = self.stft_extractor(waveform) # (B, C, F, T), complex mel_spec = self.mel_scale(torch.abs(x) ** 2) # (B, C, n_mels, T) mel_spec = self.amp2db(mel_spec) if self.feature in ["LogMel"]: x = mel_spec elif self.feature in ["MonoLogMel"]: x = mel_spec[:, :1, :, :] return x class AudioProjectorSubsample(nn.Module): def __init__(self, in_dim: int, out_dim: int, downsample_rate=5): super().__init__() self.k = downsample_rate self.net = nn.Sequential( nn.Linear(in_dim * self.k, out_dim), nn.GELU(), nn.Linear(out_dim, out_dim) ) def forward(self, x, mask=None): batch_size, seq_len, dim = x.shape num_frames_to_discard = seq_len % self.k if num_frames_to_discard > 0: x = x[:, :-num_frames_to_discard, :] if mask is not None: mask = mask[:, :-num_frames_to_discard] if mask is None: mask = torch.ones(x.shape[:-1], dtype=torch.long, device=x.device) x = rearrange(x, "b (s k) d -> b s (k d)", k=self.k) x = self.net(x) mask = rearrange(mask, "b (s k) -> b s k", k=self.k) mask = mask.any(dim=-1).long() return x, mask