midashenglm-spatial / spatial_audio_encoder.py
Jinbo-HU's picture
Upload folder using huggingface_hub
241c29f verified
Raw History Blame Contribute Delete
28.5 kB
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