Ashiedu's picture
Initial kernel bundle upload: selective-update
e14d114 verified
Raw
History Blame Contribute Delete
6.49 kB
"""Python autograd wrapper for selective-update Triton kernel.
Matches upstream ``mamba_ssm.ops.triton.selective_state_update`` ABI exactly:
state: (B, D, N) or (B, H, D, N)
x: (B, D) or (B, H, D)
dt: (B, D) or (B, H, D)
A: (D, N) or (H, D, N)
B: (B, N) or (B, G, N)
C: (B, N) or (B, G, N)
D: (D,) or (H, D) or None
z: (B, D) or (B, H, D) or None
dt_bias: (D,) or (H, D) or None
state_batch_indices: (B,) or None
"""
from __future__ import annotations
import importlib.util
from pathlib import Path
import torch
import triton
_TRITON_KERNELS_PATH = (
Path(__file__).resolve().parent / "triton_kernels.py"
)
def _load_triton_kernel():
spec = importlib.util.spec_from_file_location(
"selective_update.triton_kernels",
_TRITON_KERNELS_PATH,
)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.selective_state_update_kernel
selective_state_update_kernel = None # populated on first call
def _get_kernel():
global selective_state_update_kernel
if selective_state_update_kernel is None:
selective_state_update_kernel = _load_triton_kernel()
return selective_state_update_kernel
def _resolve_tie_hdim(A, dt, dt_bias):
if A.dim() == 2:
a0 = A[0]
return a0.stride(0) == 0 and a0.stride(1) == 0 and dt.stride(-1) == 0 and (dt_bias.stride(-1) == 0 if dt_bias is not None else True)
return A.stride(-1) == 0 and A.stride(-2) == 0 and dt.stride(-1) == 0 and (dt_bias.stride(-1) == 0 if dt_bias is not None else True)
class SelectiveUpdateFunction(torch.autograd.Function):
"""Differentiable single-token selective state update.
Mirrors upstream ``selective_state_update`` shape contract exactly.
Backward raises NotImplementedError — training uses the patched
selective_scan reverse path instead.
"""
@staticmethod
def forward(ctx, state, x, dt, A, B, C, D=None, z=None, dt_bias=None, dt_softplus=False, state_batch_indices=None):
orig_state = state
has_heads = state.dim() > 3
if state.dim() == 3:
state = state.unsqueeze(1)
if x.dim() == 2:
x = x.unsqueeze(1)
if dt.dim() == 2:
dt = dt.unsqueeze(1)
if A.dim() == 2:
A = A.unsqueeze(0)
if B.dim() == 2:
B = B.unsqueeze(1)
if C.dim() == 2:
C = C.unsqueeze(1)
if D is not None and D.dim() == 1:
D = D.unsqueeze(0)
if z is not None and z.dim() == 2:
z = z.unsqueeze(1)
if dt_bias is not None and dt_bias.dim() == 1:
dt_bias = dt_bias.unsqueeze(0)
B_sz, H, D_sz, N_sz = state.shape
x_sz = x.shape[0]
if x.shape != (x_sz, H, D_sz):
raise ValueError(f"x shape {x.shape} does not match state batch/heads/dim")
if dt.shape != (x_sz, H, D_sz):
raise ValueError(f"dt shape {dt.shape} does not match x shape")
if A.shape != (H, D_sz, N_sz):
raise ValueError(f"A shape {A.shape} does not match (H, D, N)")
ngroups = B.shape[1]
if H % ngroups != 0:
raise ValueError(f"nheads {H} must be divisible by ngroups {ngroups}")
if B.shape[0] != x_sz or B.shape[2] != N_sz or C.shape[0] != x_sz or C.shape[2] != N_sz:
raise ValueError(f"B/C batch or dstate mismatch")
if D is not None and D.shape not in {(H, D_sz), (H,)}:
raise ValueError(f"D shape {D.shape} does not match (H, D) or (H,)")
if z is not None and z.shape != (x_sz, H, D_sz):
raise ValueError(f"z shape {z.shape} does not match x shape")
if dt_bias is not None and dt_bias.shape not in {(H, D_sz), (H,)}:
raise ValueError(f"dt_bias shape {dt_bias.shape} does not match (H, D) or (H,)")
if state_batch_indices is not None and state_batch_indices.shape != (x_sz,):
raise ValueError(f"state_batch_indices shape {state_batch_indices.shape} does not match (B,)")
tie_hdim = _resolve_tie_hdim(A, dt, dt_bias)
nheads_ratio = H // ngroups
out = torch.empty_like(x)
grid = lambda META: (triton.cdiv(D_sz, META["BLOCK_SIZE_M"]), x_sz, H)
BLOCK_M = 32 if N_sz <= 16 else (16 if N_sz <= 32 else (8 if N_sz <= 64 else (4 if N_sz <= 128 else 4)))
num_warps = 4 if N_sz <= 64 else 8
z_strides = ((z.stride(0), z.stride(1), z.stride(2)) if z is not None else (0, 0, 0))
_get_kernel()[grid](
state, x, dt, dt_bias, A, B, C, D, z, out, state_batch_indices,
x_sz, H, D_sz, N_sz, nheads_ratio,
state.stride(0), state.stride(1), state.stride(2), state.stride(3),
x.stride(0), x.stride(1), x.stride(2),
dt.stride(0), dt.stride(1), dt.stride(2),
dt_bias.stride(0) if dt_bias is not None else 0, dt_bias.stride(1) if dt_bias is not None else 0,
A.stride(0), A.stride(1), A.stride(2),
B.stride(0), B.stride(1), B.stride(2),
C.stride(0), C.stride(1), C.stride(2),
D.stride(0) if D is not None else 0, D.stride(1) if D is not None else 0,
z_strides[0], z_strides[1], z_strides[2],
out.stride(0), out.stride(1), out.stride(2),
dt_softplus,
tie_hdim,
BLOCK_M,
num_warps=num_warps,
num_stages=2,
)
if not has_heads:
out = out.squeeze(1)
state = state.squeeze(1)
ctx.save_for_backward(orig_state, x, dt, A, B, C, D, z, dt_bias, state, state_batch_indices)
ctx.tie_hdim = tie_hdim
ctx.dt_softplus = dt_softplus
ctx.has_heads = has_heads
ctx.nheads_ratio = nheads_ratio
return out
@staticmethod
def backward(ctx, grad_out):
raise NotImplementedError(
"SelectiveUpdateFunction backward is not implemented. "
"Use autograd through the patched selective_scan path for training."
)
def selective_state_update(state, x, dt, A, B, C, D=None, z=None, dt_bias=None, dt_softplus=False, state_batch_indices=None, **kwargs):
"""Dispatches single-token state updates to the custom Triton backend.
Shapes mirror upstream ``mamba_ssm.ops.triton.selective_state_update`` exactly.
"""
return SelectiveUpdateFunction.apply(
state, x, dt, A, B, C, D, z, dt_bias, dt_softplus, state_batch_indices
)