Initial commit: FunASR Speech Recognition Toolkit
Update API Documentation / build-api-docs (push) Has been cancelled

Add complete FunASR codebase including models, runtime, and documentation.
This commit is contained in:
freedakgmail
2026-07-09 22:38:58 +08:00
commit 6116b1f3c6
3683 changed files with 990984 additions and 0 deletions
+113
View File
@@ -0,0 +1,113 @@
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
import copy
from funasr.models.base_model import FunASRModel
from funasr.models.encoder.mossformer_encoder import MossFormerEncoder, MossFormer_MaskNet
from funasr.models.decoder.mossformer_decoder import MossFormerDecoder
class MossFormer(FunASRModel):
"""The MossFormer model for separating input mixed speech into different speaker's speech.
Arguments
---------
in_channels : int
Number of channels at the output of the encoder.
out_channels : int
Number of channels that would be inputted to the intra and inter blocks.
num_blocks : int
Number of layers of Dual Computation Block.
norm : str
Normalization type.
num_spks : int
Number of sources (speakers).
skip_around_intra : bool
Skip connection around intra.
use_global_pos_enc : bool
Global positional encodings.
max_length : int
Maximum sequence length.
kernel_size: int
Encoder and decoder kernel size
"""
def __init__(
self,
in_channels=512,
out_channels=512,
num_blocks=24,
kernel_size=16,
norm="ln",
num_spks=2,
skip_around_intra=True,
use_global_pos_enc=True,
max_length=20000,
):
"""Initialize MossFormer.
Args:
in_channels: TODO.
out_channels: TODO.
num_blocks: TODO.
kernel_size: Size/dimension parameter.
norm: TODO.
num_spks: TODO.
skip_around_intra: TODO.
use_global_pos_enc: TODO.
max_length: TODO.
"""
super(MossFormer, self).__init__()
self.num_spks = num_spks
# Encoding
self.enc = MossFormerEncoder(
kernel_size=kernel_size, out_channels=in_channels, in_channels=1
)
##Compute Mask
self.mask_net = MossFormer_MaskNet(
in_channels=in_channels,
out_channels=out_channels,
num_blocks=num_blocks,
norm=norm,
num_spks=num_spks,
skip_around_intra=skip_around_intra,
use_global_pos_enc=use_global_pos_enc,
max_length=max_length,
)
self.dec = MossFormerDecoder(
in_channels=out_channels,
out_channels=1,
kernel_size=kernel_size,
stride=kernel_size // 2,
bias=False,
)
def forward(self, input):
"""Forward pass for training.
Args:
input: Input audio/text data.
"""
x = self.enc(input)
mask = self.mask_net(x)
x = torch.stack([x] * self.num_spks)
sep_x = x * mask
# Decoding
est_source = torch.cat(
[self.dec(sep_x[i]).unsqueeze(-1) for i in range(self.num_spks)],
dim=-1,
)
T_origin = input.size(1)
T_est = est_source.size(1)
if T_origin > T_est:
est_source = F.pad(est_source, (0, 0, 0, T_origin - T_est))
else:
est_source = est_source[:, :T_origin, :]
out = []
for spk in range(self.num_spks):
out.append(est_source[:, :, spk])
return out
+422
View File
@@ -0,0 +1,422 @@
import torch
import torch.nn.functional as F
from torch import nn, einsum
from einops import rearrange
def identity(t, *args, **kwargs):
"""Identity.
Args:
t: TODO.
*args: Variable positional arguments.
**kwargs: Additional keyword arguments.
"""
return t
def append_dims(x, num_dims):
"""Append dims.
Args:
x: TODO.
num_dims: TODO.
"""
if num_dims <= 0:
return x
return x.view(*x.shape, *((1,) * num_dims))
def exists(val):
"""Exists.
Args:
val: TODO.
"""
return val is not None
def default(val, d):
"""Default.
Args:
val: TODO.
d: TODO.
"""
return val if exists(val) else d
def padding_to_multiple_of(n, mult):
"""Padding to multiple of.
Args:
n: TODO.
mult: TODO.
"""
remainder = n % mult
if remainder == 0:
return 0
return mult - remainder
class Transpose(nn.Module):
"""Wrapper class of torch.transpose() for Sequential module."""
def __init__(self, shape: tuple):
"""Initialize Transpose.
Args:
shape: TODO.
"""
super(Transpose, self).__init__()
self.shape = shape
def forward(self, x):
"""Forward pass for training.
Args:
x: TODO.
"""
return x.transpose(*self.shape)
class DepthwiseConv1d(nn.Module):
"""
When groups == in_channels and out_channels == K * in_channels, where K is a positive integer,
this operation is termed in literature as depthwise convolution.
Args:
in_channels (int): Number of channels in the input
out_channels (int): Number of channels produced by the convolution
kernel_size (int or tuple): Size of the convolving kernel
stride (int, optional): Stride of the convolution. Default: 1
padding (int or tuple, optional): Zero-padding added to both sides of the input. Default: 0
bias (bool, optional): If True, adds a learnable bias to the output. Default: True
Inputs: inputs
- **inputs** (batch, in_channels, time): Tensor containing input vector
Returns: outputs
- **outputs** (batch, out_channels, time): Tensor produces by depthwise 1-D convolution.
"""
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: int,
stride: int = 1,
padding: int = 0,
bias: bool = False,
) -> None:
"""Initialize DepthwiseConv1d.
Args:
in_channels: TODO.
out_channels: TODO.
kernel_size: Size/dimension parameter.
stride: TODO.
padding: TODO.
bias: TODO.
"""
super(DepthwiseConv1d, self).__init__()
assert (
out_channels % in_channels == 0
), "out_channels should be constant multiple of in_channels"
self.conv = nn.Conv1d(
in_channels=in_channels,
out_channels=out_channels,
kernel_size=kernel_size,
groups=in_channels,
stride=stride,
padding=padding,
bias=bias,
)
def forward(self, inputs):
"""Forward pass for training.
Args:
inputs: TODO.
"""
return self.conv(inputs)
class ConvModule(nn.Module):
"""
Conformer convolution module starts with a pointwise convolution and a gated linear unit (GLU).
This is followed by a single 1-D depthwise convolution layer. Batchnorm is deployed just after the convolution
to aid training deep models.
Args:
in_channels (int): Number of channels in the input
kernel_size (int or tuple, optional): Size of the convolving kernel Default: 31
dropout_p (float, optional): probability of dropout
Inputs: inputs
inputs (batch, time, dim): Tensor contains input sequences
Outputs: outputs
outputs (batch, time, dim): Tensor produces by conformer convolution module.
"""
def __init__(
self,
in_channels: int,
kernel_size: int = 17,
expansion_factor: int = 2,
dropout_p: float = 0.1,
) -> None:
"""Initialize ConvModule.
Args:
in_channels: TODO.
kernel_size: Size/dimension parameter.
expansion_factor: TODO.
dropout_p: TODO.
"""
super(ConvModule, self).__init__()
assert (kernel_size - 1) % 2 == 0, "kernel_size should be a odd number for 'SAME' padding"
assert expansion_factor == 2, "Currently, Only Supports expansion_factor 2"
self.sequential = nn.Sequential(
Transpose(shape=(1, 2)),
DepthwiseConv1d(
in_channels, in_channels, kernel_size, stride=1, padding=(kernel_size - 1) // 2
),
)
def forward(self, inputs):
"""Forward pass for training.
Args:
inputs: TODO.
"""
return inputs + self.sequential(inputs).transpose(1, 2)
class OffsetScale(nn.Module):
def __init__(self, dim, heads=1):
"""Initialize OffsetScale.
Args:
dim: TODO.
heads: TODO.
"""
super().__init__()
self.gamma = nn.Parameter(torch.ones(heads, dim))
self.beta = nn.Parameter(torch.zeros(heads, dim))
nn.init.normal_(self.gamma, std=0.02)
def forward(self, x):
"""Forward pass for training.
Args:
x: TODO.
"""
out = einsum("... d, h d -> ... h d", x, self.gamma) + self.beta
return out.unbind(dim=-2)
class FFConvM(nn.Module):
def __init__(self, dim_in, dim_out, norm_klass=nn.LayerNorm, dropout=0.1):
"""Initialize FFConvM.
Args:
dim_in: TODO.
dim_out: TODO.
norm_klass: TODO.
dropout: TODO.
"""
super().__init__()
self.mdl = nn.Sequential(
norm_klass(dim_in),
nn.Linear(dim_in, dim_out),
nn.SiLU(),
ConvModule(dim_out),
nn.Dropout(dropout),
)
def forward(
self,
x,
):
"""Forward pass for training.
Args:
x: TODO.
"""
output = self.mdl(x)
return output
class FLASH_ShareA_FFConvM(nn.Module):
def __init__(
self,
*,
dim,
group_size=256,
query_key_dim=128,
expansion_factor=1.0,
causal=False,
dropout=0.1,
rotary_pos_emb=None,
norm_klass=nn.LayerNorm,
shift_tokens=True
):
"""Initialize FLASH_ShareA_FFConvM."""
super().__init__()
hidden_dim = int(dim * expansion_factor)
self.group_size = group_size
self.causal = causal
self.shift_tokens = shift_tokens
# positional embeddings
self.rotary_pos_emb = rotary_pos_emb
# norm
self.dropout = nn.Dropout(dropout)
# projections
self.to_hidden = FFConvM(
dim_in=dim,
dim_out=hidden_dim,
norm_klass=norm_klass,
dropout=dropout,
)
self.to_qk = FFConvM(
dim_in=dim,
dim_out=query_key_dim,
norm_klass=norm_klass,
dropout=dropout,
)
self.qk_offset_scale = OffsetScale(query_key_dim, heads=4)
self.to_out = FFConvM(
dim_in=dim * 2,
dim_out=dim,
norm_klass=norm_klass,
dropout=dropout,
)
self.gateActivate = nn.Sigmoid()
def forward(self, x, *, mask=None):
"""
b - batch
n - sequence length (within groups)
g - group dimension
d - feature dimension (keys)
e - feature dimension (values)
i - sequence dimension (source)
j - sequence dimension (target)
"""
normed_x = x
# do token shift - a great, costless trick from an independent AI researcher in Shenzhen
residual = x
if self.shift_tokens:
x_shift, x_pass = normed_x.chunk(2, dim=-1)
x_shift = F.pad(x_shift, (0, 0, 1, -1), value=0.0)
normed_x = torch.cat((x_shift, x_pass), dim=-1)
# initial projections
v, u = self.to_hidden(normed_x).chunk(2, dim=-1)
qk = self.to_qk(normed_x)
# offset and scale
quad_q, lin_q, quad_k, lin_k = self.qk_offset_scale(qk)
att_v, att_u = self.cal_attention(x, quad_q, lin_q, quad_k, lin_k, v, u)
out = (att_u * v) * self.gateActivate(att_v * u)
x = x + self.to_out(out)
return x
def cal_attention(self, x, quad_q, lin_q, quad_k, lin_k, v, u, mask=None):
"""Cal attention.
Args:
x: TODO.
quad_q: TODO.
lin_q: TODO.
quad_k: TODO.
lin_k: TODO.
v: TODO.
u: TODO.
mask: TODO.
"""
b, n, device, g = x.shape[0], x.shape[-2], x.device, self.group_size
if exists(mask):
lin_mask = rearrange(mask, "... -> ... 1")
lin_k = lin_k.masked_fill(~lin_mask, 0.0)
# rotate queries and keys
if exists(self.rotary_pos_emb):
quad_q, lin_q, quad_k, lin_k = map(
self.rotary_pos_emb.rotate_queries_or_keys, (quad_q, lin_q, quad_k, lin_k)
)
# padding for groups
padding = padding_to_multiple_of(n, g)
if padding > 0:
quad_q, quad_k, lin_q, lin_k, v, u = map(
lambda t: F.pad(t, (0, 0, 0, padding), value=0.0),
(quad_q, quad_k, lin_q, lin_k, v, u),
)
mask = default(mask, torch.ones((b, n), device=device, dtype=torch.bool))
mask = F.pad(mask, (0, padding), value=False)
# group along sequence
quad_q, quad_k, lin_q, lin_k, v, u = map(
lambda t: rearrange(t, "b (g n) d -> b g n d", n=self.group_size),
(quad_q, quad_k, lin_q, lin_k, v, u),
)
if exists(mask):
mask = rearrange(mask, "b (g j) -> b g 1 j", j=g)
# calculate quadratic attention output
sim = einsum("... i d, ... j d -> ... i j", quad_q, quad_k) / g
attn = F.relu(sim) ** 2
attn = self.dropout(attn)
if exists(mask):
attn = attn.masked_fill(~mask, 0.0)
if self.causal:
causal_mask = torch.ones((g, g), dtype=torch.bool, device=device).triu(1)
attn = attn.masked_fill(causal_mask, 0.0)
quad_out_v = einsum("... i j, ... j d -> ... i d", attn, v)
quad_out_u = einsum("... i j, ... j d -> ... i d", attn, u)
# calculate linear attention output
if self.causal:
lin_kv = einsum("b g n d, b g n e -> b g d e", lin_k, v) / g
# exclusive cumulative sum along group dimension
lin_kv = lin_kv.cumsum(dim=1)
lin_kv = F.pad(lin_kv, (0, 0, 0, 0, 1, -1), value=0.0)
lin_out_v = einsum("b g d e, b g n d -> b g n e", lin_kv, lin_q)
lin_ku = einsum("b g n d, b g n e -> b g d e", lin_k, u) / g
# exclusive cumulative sum along group dimension
lin_ku = lin_ku.cumsum(dim=1)
lin_ku = F.pad(lin_ku, (0, 0, 0, 0, 1, -1), value=0.0)
lin_out_u = einsum("b g d e, b g n d -> b g n e", lin_ku, lin_q)
else:
lin_kv = einsum("b g n d, b g n e -> b d e", lin_k, v) / n
lin_out_v = einsum("b g n d, b d e -> b g n e", lin_q, lin_kv)
lin_ku = einsum("b g n d, b g n e -> b d e", lin_k, u) / n
lin_out_u = einsum("b g n d, b d e -> b g n e", lin_q, lin_ku)
# fold back groups into full sequence, and excise out padding
return map(
lambda t: rearrange(t, "b g n d -> b (g n) d")[:, :n],
(quad_out_v + lin_out_v, quad_out_u + lin_out_u),
)
@@ -0,0 +1,56 @@
import torch
import torch.nn as nn
class MossFormerDecoder(nn.ConvTranspose1d):
"""A decoder layer that consists of ConvTranspose1d.
Arguments
---------
kernel_size : int
Length of filters.
in_channels : int
Number of input channels.
out_channels : int
Number of output channels.
Example
---------
>>> x = torch.randn(2, 100, 1000)
>>> decoder = Decoder(kernel_size=4, in_channels=100, out_channels=1)
>>> h = decoder(x)
>>> h.shape
torch.Size([2, 1003])
"""
def __init__(self, *args, **kwargs):
"""Initialize MossFormerDecoder.
Args:
*args: Variable positional arguments.
**kwargs: Additional keyword arguments.
"""
super(MossFormerDecoder, self).__init__(*args, **kwargs)
def forward(self, x):
"""Return the decoded output.
Arguments
---------
x : torch.Tensor
Input tensor with dimensionality [B, N, L].
where, B = Batchsize,
N = number of filters
L = time points
"""
if x.dim() not in [2, 3]:
raise RuntimeError("{} accept 3/4D tensor as input".format(self.__name__))
x = super().forward(x if x.dim() == 3 else torch.unsqueeze(x, 1))
if torch.squeeze(x).dim() == 1:
x = torch.squeeze(x, dim=1)
else:
x = torch.squeeze(x)
return x
@@ -0,0 +1,473 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
try:
from rotary_embedding_torch import RotaryEmbedding
except:
# print(
# "If you want use mossformer, lease install rotary_embedding_torch by: \n pip install -U rotary_embedding_torch"
# )
pass
from funasr.models.transformer.layer_norm import GlobalLayerNorm, CumulativeLayerNorm, ScaleNorm
from funasr.models.transformer.embedding import ScaledSinuEmbedding
from funasr.models.mossformer.mossformer import FLASH_ShareA_FFConvM
def select_norm(norm, dim, shape):
"""Just a wrapper to select the normalization type."""
if norm == "gln":
return GlobalLayerNorm(dim, shape, elementwise_affine=True)
if norm == "cln":
return CumulativeLayerNorm(dim, elementwise_affine=True)
if norm == "ln":
return nn.GroupNorm(1, dim, eps=1e-8)
else:
return nn.BatchNorm1d(dim)
class MossformerBlock(nn.Module):
def __init__(
self,
*,
dim,
depth,
group_size=256,
query_key_dim=128,
expansion_factor=4.0,
causal=False,
attn_dropout=0.1,
norm_type="scalenorm",
shift_tokens=True
):
"""Initialize MossformerBlock."""
super().__init__()
assert norm_type in (
"scalenorm",
"layernorm",
), "norm_type must be one of scalenorm or layernorm"
if norm_type == "scalenorm":
norm_klass = ScaleNorm
elif norm_type == "layernorm":
norm_klass = nn.LayerNorm
self.group_size = group_size
rotary_pos_emb = RotaryEmbedding(dim=min(32, query_key_dim))
# max rotary embedding dimensions of 32, partial Rotary embeddings, from Wang et al - GPT-J
self.layers = nn.ModuleList(
[
FLASH_ShareA_FFConvM(
dim=dim,
group_size=group_size,
query_key_dim=query_key_dim,
expansion_factor=expansion_factor,
causal=causal,
dropout=attn_dropout,
rotary_pos_emb=rotary_pos_emb,
norm_klass=norm_klass,
shift_tokens=shift_tokens,
)
for _ in range(depth)
]
)
def forward(self, x, *, mask=None):
"""Forward pass for training.
Args:
x: TODO.
"""
ii = 0
for flash in self.layers:
x = flash(x, mask=mask)
ii = ii + 1
return x
class MossFormer_MaskNet(nn.Module):
"""The MossFormer module for computing output masks.
Arguments
---------
in_channels : int
Number of channels at the output of the encoder.
out_channels : int
Number of channels that would be inputted to the intra and inter blocks.
num_blocks : int
Number of layers of Dual Computation Block.
norm : str
Normalization type.
num_spks : int
Number of sources (speakers).
skip_around_intra : bool
Skip connection around intra.
use_global_pos_enc : bool
Global positional encodings.
max_length : int
Maximum sequence length.
Example
---------
>>> mossformer_block = MossFormerM(1, 64, 8)
>>> mossformer_masknet = MossFormer_MaskNet(64, 64, intra_block, num_spks=2)
>>> x = torch.randn(10, 64, 2000)
>>> x = mossformer_masknet(x)
>>> x.shape
torch.Size([2, 10, 64, 2000])
"""
def __init__(
self,
in_channels,
out_channels,
num_blocks=24,
norm="ln",
num_spks=2,
skip_around_intra=True,
use_global_pos_enc=True,
max_length=20000,
):
"""Initialize MossFormer_MaskNet.
Args:
in_channels: TODO.
out_channels: TODO.
num_blocks: TODO.
norm: TODO.
num_spks: TODO.
skip_around_intra: TODO.
use_global_pos_enc: TODO.
max_length: TODO.
"""
super(MossFormer_MaskNet, self).__init__()
self.num_spks = num_spks
self.num_blocks = num_blocks
self.norm = select_norm(norm, in_channels, 3)
self.conv1d_encoder = nn.Conv1d(in_channels, out_channels, 1, bias=False)
self.use_global_pos_enc = use_global_pos_enc
if self.use_global_pos_enc:
self.pos_enc = ScaledSinuEmbedding(out_channels)
self.mdl = Computation_Block(
num_blocks,
out_channels,
norm,
skip_around_intra=skip_around_intra,
)
self.conv1d_out = nn.Conv1d(out_channels, out_channels * num_spks, kernel_size=1)
self.conv1_decoder = nn.Conv1d(out_channels, in_channels, 1, bias=False)
self.prelu = nn.PReLU()
self.activation = nn.ReLU()
# gated output layer
self.output = nn.Sequential(nn.Conv1d(out_channels, out_channels, 1), nn.Tanh())
self.output_gate = nn.Sequential(nn.Conv1d(out_channels, out_channels, 1), nn.Sigmoid())
def forward(self, x):
"""Returns the output tensor.
Arguments
---------
x : torch.Tensor
Input tensor of dimension [B, N, S].
Returns
-------
out : torch.Tensor
Output tensor of dimension [spks, B, N, S]
where, spks = Number of speakers
B = Batchsize,
N = number of filters
S = the number of time frames
"""
# before each line we indicate the shape after executing the line
# [B, N, L]
x = self.norm(x)
# [B, N, L]
x = self.conv1d_encoder(x)
if self.use_global_pos_enc:
# x = self.pos_enc(x.transpose(1, -1)).transpose(1, -1) + x * (
# x.size(1) ** 0.5)
base = x
x = x.transpose(1, -1)
emb = self.pos_enc(x)
emb = emb.transpose(0, -1)
# print('base: {}, emb: {}'.format(base.shape, emb.shape))
x = base + emb
# [B, N, S]
# for i in range(self.num_modules):
# x = self.dual_mdl[i](x)
x = self.mdl(x)
x = self.prelu(x)
# [B, N*spks, S]
x = self.conv1d_out(x)
B, _, S = x.shape
# [B*spks, N, S]
x = x.view(B * self.num_spks, -1, S)
# [B*spks, N, S]
x = self.output(x) * self.output_gate(x)
# [B*spks, N, S]
x = self.conv1_decoder(x)
# [B, spks, N, S]
_, N, L = x.shape
x = x.view(B, self.num_spks, N, L)
x = self.activation(x)
# [spks, B, N, S]
x = x.transpose(0, 1)
return x
class MossFormerEncoder(nn.Module):
"""Convolutional Encoder Layer.
Arguments
---------
kernel_size : int
Length of filters.
in_channels : int
Number of input channels.
out_channels : int
Number of output channels.
Example
-------
>>> x = torch.randn(2, 1000)
>>> encoder = Encoder(kernel_size=4, out_channels=64)
>>> h = encoder(x)
>>> h.shape
torch.Size([2, 64, 499])
"""
def __init__(self, kernel_size=2, out_channels=64, in_channels=1):
"""Initialize MossFormerEncoder.
Args:
kernel_size: Size/dimension parameter.
out_channels: TODO.
in_channels: TODO.
"""
super(MossFormerEncoder, self).__init__()
self.conv1d = nn.Conv1d(
in_channels=in_channels,
out_channels=out_channels,
kernel_size=kernel_size,
stride=kernel_size // 2,
groups=1,
bias=False,
)
self.in_channels = in_channels
def forward(self, x):
"""Return the encoded output.
Arguments
---------
x : torch.Tensor
Input tensor with dimensionality [B, L].
Return
------
x : torch.Tensor
Encoded tensor with dimensionality [B, N, T_out].
where B = Batchsize
L = Number of timepoints
N = Number of filters
T_out = Number of timepoints at the output of the encoder
"""
# B x L -> B x 1 x L
if self.in_channels == 1:
x = torch.unsqueeze(x, dim=1)
# B x 1 x L -> B x N x T_out
x = self.conv1d(x)
x = F.relu(x)
return x
class MossFormerM(nn.Module):
"""This class implements the transformer encoder.
Arguments
---------
num_blocks : int
Number of mossformer blocks to include.
d_model : int
The dimension of the input embedding.
attn_dropout : float
Dropout for the self-attention (Optional).
group_size: int
the chunk size
query_key_dim: int
the attention vector dimension
expansion_factor: int
the expansion factor for the linear projection in conv module
causal: bool
true for causal / false for non causal
Example
-------
>>> import torch
>>> x = torch.rand((8, 60, 512))
>>> net = TransformerEncoder_MossFormerM(num_blocks=8, d_model=512)
>>> output, _ = net(x)
>>> output.shape
torch.Size([8, 60, 512])
"""
def __init__(
self,
num_blocks,
d_model=None,
causal=False,
group_size=256,
query_key_dim=128,
expansion_factor=4.0,
attn_dropout=0.1,
):
"""Initialize MossFormerM.
Args:
num_blocks: TODO.
d_model: D Model instance.
causal: TODO.
group_size: Size/dimension parameter.
query_key_dim: Size/dimension parameter.
expansion_factor: TODO.
attn_dropout: TODO.
"""
super().__init__()
self.mossformerM = MossformerBlock(
dim=d_model,
depth=num_blocks,
group_size=group_size,
query_key_dim=query_key_dim,
expansion_factor=expansion_factor,
causal=causal,
attn_dropout=attn_dropout,
)
self.norm = nn.LayerNorm(d_model, eps=1e-6)
def forward(
self,
src,
):
"""
Arguments
----------
src : torch.Tensor
Tensor shape [B, L, N],
where, B = Batchsize,
L = time points
N = number of filters
The sequence to the encoder layer (required).
src_mask : tensor
The mask for the src sequence (optional).
src_key_padding_mask : tensor
The mask for the src keys per batch (optional).
"""
output = self.mossformerM(src)
output = self.norm(output)
return output
class Computation_Block(nn.Module):
"""Computation block for dual-path processing.
Arguments
---------
out_channels : int
Dimensionality of inter/intra model.
norm : str
Normalization type.
skip_around_intra : bool
Skip connection around the intra layer.
Example
---------
>>> comp_block = Computation_Block(64)
>>> x = torch.randn(10, 64, 100)
>>> x = comp_block(x)
>>> x.shape
torch.Size([10, 64, 100])
"""
def __init__(
self,
num_blocks,
out_channels,
norm="ln",
skip_around_intra=True,
):
"""Initialize Computation_Block.
Args:
num_blocks: TODO.
out_channels: TODO.
norm: TODO.
skip_around_intra: TODO.
"""
super(Computation_Block, self).__init__()
##MossFormer2M: MossFormer with recurrence
# self.intra_mdl = MossFormer2M(num_blocks=num_blocks, d_model=out_channels)
##MossFormerM: the orignal MossFormer
self.intra_mdl = MossFormerM(num_blocks=num_blocks, d_model=out_channels)
self.skip_around_intra = skip_around_intra
# Norm
self.norm = norm
if norm is not None:
self.intra_norm = select_norm(norm, out_channels, 3)
def forward(self, x):
"""Returns the output tensor.
Arguments
---------
x : torch.Tensor
Input tensor of dimension [B, N, S].
Return
---------
out: torch.Tensor
Output tensor of dimension [B, N, S].
where, B = Batchsize,
N = number of filters
S = sequence time index
"""
B, N, S = x.shape
# intra RNN
# [B, S, N]
intra = x.permute(0, 2, 1).contiguous() # .view(B, S, N)
intra = self.intra_mdl(intra)
# [B, N, S]
intra = intra.permute(0, 2, 1).contiguous()
if self.norm is not None:
intra = self.intra_norm(intra)
# [B, N, S]
if self.skip_around_intra:
intra = intra + x
out = intra
return out