1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89
| import torch import torch.nn as nn import torch.nn.functional as F
class MambaBlock(nn.Module): def __init__(self, dim, d_state=16, d_conv=4, expand=2): super().__init__() self.dim = dim self.d_state = d_state self.d_conv = d_conv self.expand = expand self.in_proj = nn.Linear(dim, dim * expand * 2, bias=False) self.conv1d = nn.Conv1d( in_channels=dim * expand, out_channels=dim * expand, kernel_size=d_conv, groups=dim * expand, padding=d_conv - 1 ) self.x_proj = nn.Linear(dim * expand, d_state * 2, bias=False) self.dt_proj = nn.Linear(d_state, dim * expand, bias=True) self.A_log = nn.Parameter(torch.randn(dim * expand, d_state)) self.D = nn.Parameter(torch.ones(dim * expand)) self.out_proj = nn.Linear(dim * expand, dim, bias=False) def forward(self, x): batch_size, seq_len, _ = x.shape xz = self.in_proj(x) x, z = xz.chunk(2, dim=-1) x = x.transpose(1, 2) x = self.conv1d(x)[:, :, :seq_len] x = x.transpose(1, 2) x = F.silu(x) x_dbl = self.x_proj(x) x_db, x_dt = x_dbl.chunk(2, dim=-1) A = -torch.exp(self.A_log) dt = self.dt_proj(x_dt) dt = F.softplus(dt) y = self.ssm(x, x_db, A, dt, self.D) y = y * F.silu(z) return self.out_proj(y) def ssm(self, x, x_db, A, dt, D): """状态空间模型计算""" batch_size, seq_len, dim = x.shape dA = torch.exp(dt.unsqueeze(-1) * A) dB = dt.unsqueeze(-1) * x_db.unsqueeze(2) y = torch.zeros_like(x) h = torch.zeros(batch_size, dim, self.d_state, device=x.device) for i in range(seq_len): h = dA[:, i] * h + dB[:, i] * x[:, i].unsqueeze(-1) y[:, i] = (h * D).sum(-1) return y
|