Spaces:
Running on Zero
Running on Zero
Download models/selective_scan.py from HanzhouLiu/XYScanNet_Demo: direct link, hf CLI and curl.
- Browser
- Download file 2.59 kB
-
https://huggingface.co/spaces/HanzhouLiu/XYScanNet_Demo/resolve/main/models/selective_scan.py
- Command line
-
hf download hf://spaces/HanzhouLiu/XYScanNet_Demo/models/selective_scan.py
-
curl -L -o selective_scan.py https://huggingface.co/spaces/HanzhouLiu/XYScanNet_Demo/resolve/main/models/selective_scan.py
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." | |
| ) | |