File size: 6,486 Bytes
e14d114
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
"""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
    )