Audio-Text-to-Text
Transformers
Safetensors
midashenglm_spatial
text-generation
spatial-audio-understanding
multimodal
audio-language-model
audio
dasheng
custom_code
Instructions to use mispeech/midashenglm-spatial with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use mispeech/midashenglm-spatial with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("mispeech/midashenglm-spatial", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download spatial_audio_encoder.py from mispeech/midashenglm-spatial: direct link, hf CLI and curl.
- Browser
- Download file 28.5 kB
-
https://huggingface.co/mispeech/midashenglm-spatial/resolve/main/spatial_audio_encoder.py
- Command line
-
hf download hf://mispeech/midashenglm-spatial/spatial_audio_encoder.py
-
curl -L -o spatial_audio_encoder.py https://huggingface.co/mispeech/midashenglm-spatial/resolve/main/spatial_audio_encoder.py
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}") | |
| 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 | |