class MambaBlock(nn.Module):
"""
Mamba Block - Selective State Space Model.
Based on: "Mamba: Linear-Time Sequence Modeling with Selective State Spaces"
Gu & Dao, arXiv 2312.00752
This is a simplified PyTorch implementation for educational purposes.
For production, use: https://github.com/state-spaces/mamba
"""
def __init__(
self,
d_model: int,
d_state: int = 16,
d_conv: int = 4,
expand: int = 2,
dropout: float = 0.0
):
"""
Parameters:
-----------
d_model : int
Model dimension
d_state : int
SSM state dimension (N)
d_conv : int
Local convolution width
expand : int
Inner dimension expansion factor
dropout : float
Dropout rate
"""
super().__init__()
self.d_model = d_model
self.d_state = d_state
self.d_conv = d_conv
self.d_inner = d_model * expand
# Input projection (to 2x for gating)
self.in_proj = nn.Linear(d_model, self.d_inner * 2, bias=False)
# 1D Convolution for local context
self.conv1d = nn.Conv1d(
in_channels=self.d_inner,
out_channels=self.d_inner,
kernel_size=d_conv,
padding=d_conv - 1,
groups=self.d_inner # Depthwise
)
# SSM parameters (input-dependent)
# x_proj projects to (delta, B, C)
self.x_proj = nn.Linear(self.d_inner, d_state * 2 + 1, bias=False)
# Delta (discretization step) projection
self.dt_proj = nn.Linear(1, self.d_inner, bias=True)
# A parameter (log scale for stability)
A = torch.arange(1, d_state + 1, dtype=torch.float32)
self.A_log = nn.Parameter(torch.log(A.repeat(self.d_inner, 1)))
# D skip connection
self.D = nn.Parameter(torch.ones(self.d_inner))
# Output projection
self.out_proj = nn.Linear(self.d_inner, d_model, bias=False)
# Dropout
self.dropout = nn.Dropout(dropout)
def ssm_step(
self,
x: torch.Tensor,
h: torch.Tensor,
delta: torch.Tensor,
A: torch.Tensor,
B: torch.Tensor,
C: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Single SSM step with selective parameters.
h_t = exp(delta * A) * h_{t-1} + delta * B * x_t
y_t = C * h_t
"""
# Discretize A
A_bar = torch.exp(delta.unsqueeze(-1) * A) # (batch, d_inner, d_state)
# Discretize B
B_bar = delta.unsqueeze(-1) * B # (batch, d_inner, d_state)
# State update
h = A_bar * h + B_bar * x.unsqueeze(-1)
# Output
y = torch.sum(C * h, dim=-1) # (batch, d_inner)
return y, h
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass.
Parameters:
-----------
x : tensor
Input of shape (batch, seq_len, d_model)
Returns:
--------
tensor : Output of shape (batch, seq_len, d_model)
"""
batch, seq_len, _ = x.shape
# Input projection with gating
xz = self.in_proj(x) # (batch, seq_len, d_inner * 2)
x_main, z = xz.chunk(2, dim=-1) # Each: (batch, seq_len, d_inner)
# 1D convolution for local context
x_conv = x_main.transpose(1, 2) # (batch, d_inner, seq_len)
x_conv = self.conv1d(x_conv)[:, :, :seq_len] # Causal
x_conv = x_conv.transpose(1, 2) # (batch, seq_len, d_inner)
x_main = F.silu(x_conv)
# Project to SSM parameters
x_ssm = self.x_proj(x_main) # (batch, seq_len, d_state*2 + 1)
# Split into delta, B, C
delta_raw = x_ssm[:, :, :1] # (batch, seq_len, 1)
B = x_ssm[:, :, 1:1+self.d_state] # (batch, seq_len, d_state)
C = x_ssm[:, :, 1+self.d_state:] # (batch, seq_len, d_state)
# Delta projection + softplus
delta = F.softplus(self.dt_proj(delta_raw)) # (batch, seq_len, d_inner)
# A (negative for stability)
A = -torch.exp(self.A_log) # (d_inner, d_state)
# Initialize state
h = torch.zeros(batch, self.d_inner, self.d_state, device=x.device)
# Recurrent SSM (simplified, not optimized)
outputs = []
for t in range(seq_len):
x_t = x_main[:, t, :] # (batch, d_inner)
delta_t = delta[:, t, :] # (batch, d_inner)
B_t = B[:, t, :].unsqueeze(1).expand(-1, self.d_inner, -1) # (batch, d_inner, d_state)
C_t = C[:, t, :].unsqueeze(1).expand(-1, self.d_inner, -1) # (batch, d_inner, d_state)
y_t, h = self.ssm_step(x_t, h, delta_t, A, B_t, C_t)
# Skip connection
y_t = y_t + self.D * x_t
outputs.append(y_t)
y = torch.stack(outputs, dim=1) # (batch, seq_len, d_inner)
# Gating with z
y = y * F.silu(z)
# Output projection
y = self.out_proj(y)
y = self.dropout(y)
return y
# Test Mamba Block
print("Test de MambaBlock:")
mamba_block = MambaBlock(d_model=32, d_state=16, d_conv=4, expand=2)
test_input = torch.randn(2, 60, 32)
output = mamba_block(test_input)
print(f" Input shape: {test_input.shape}")
print(f" Output shape: {output.shape}")
print(f" Parameters: {sum(p.numel() for p in mamba_block.parameters()):,}")