XYScanNet_Demo / models /selective_scan.py
Hanzhou Liu
Add pure PyTorch selective scan fallback and remove mamba-ssm wheel dependency
c463a26
Raw History Blame Contribute Delete
2.59 kB
import torch
import torch.nn.functional as F
def selective_scan_fn(
u,
delta,
A,
B,
C,
D=None,
z=None,
delta_bias=None,
delta_softplus=False,
return_last_state=False,
):
"""
Pure PyTorch fallback for selective_scan_fn.
u: (B, D, L)
delta: (B, D, L)
A: (D, N)
B: (B, N, L) or (B, G, N, L)
C: (B, N, L) or (B, G, N, L)
D: (D,) optional
z: (B, D, L) optional
delta_bias: (D,) optional
delta_softplus: bool
return_last_state: bool
"""
dtype_in = u.dtype
u = u.float()
delta = delta.float()
if delta_bias is not None:
delta = delta + delta_bias[..., None].float()
if delta_softplus:
delta = F.softplus(delta)
batch, dim, dstate = u.shape[0], A.shape[0], A.shape[1]
seqlen = u.shape[2]
# Discretize A: shape (B, D, L, N)
deltaA = torch.exp(torch.einsum("bdl,dn->bdln", delta, A))
# Discretize B * u: shape (B, D, L, N)
if B.dim() == 3:
deltaB_u = torch.einsum("bdl,bnl,bdl->bdln", delta, B.float(), u)
elif B.dim() == 4:
if B.shape[1] == 1:
deltaB_u = torch.einsum("bdl,b1nl,bdl->bdln", delta, B.float(), u)
else:
G = B.shape[1]
d = dim // G
deltaB_u = torch.einsum("b(g d)l,bgnl,b(g d)l->bdln", delta, B.float(), u, d=d)
else:
raise ValueError(f"Unsupported B shape: {B.shape}")
# Recurrent scan over sequence length L
x = torch.zeros((batch, dim, dstate), device=u.device, dtype=deltaA.dtype)
ys = []
for i in range(seqlen):
x = deltaA[:, :, i] * x + deltaB_u[:, :, i]
if C.dim() == 3:
y = torch.einsum("bdn,bn->bd", x, C[:, :, i].float())
elif C.dim() == 4:
if C.shape[1] == 1:
y = torch.einsum("bdn,bn->bd", x, C[:, 0, :, i].float())
else:
G = C.shape[1]
d = dim // G
x_g = x.view(batch, G, d, dstate)
y = torch.einsum("bgdn,bgn->bgd", x_g, C[:, :, :, i].float())
y = y.view(batch, dim)
ys.append(y)
y = torch.stack(ys, dim=-1) # (B, D, L)
if D is not None:
y = y + u * D[..., None].float()
if z is not None:
y = y * F.silu(z.float())
out = y.to(dtype=dtype_in)
if return_last_state:
return out, x.to(dtype=dtype_in)
return out
def mamba_inner_fn(*args, **kwargs):
raise NotImplementedError(
"mamba_inner_fn is not available without compiled CUDA extensions; use fallback selective_scan path."
)