Codex commited on
Commit
e1518d1
·
1 Parent(s): cb0290b

Vendor the pinned Mamba Triton runtime

Browse files
README.md CHANGED
@@ -37,6 +37,9 @@ speaker-reference 模式使用的 ECAPA encoder 固定在
37
  `0f99f2d0ebe89ac095bcc5903c4dd8f72b367286`;TTS runtime 以
38
  `ce384c8cc54efea1aaba7b9f1d7ded6c1c99aa9a` 的已驗證原始碼隨 Space 保存,
39
  Barbet 另固定在 `6fcd7ce4aa37f2250a3242995bef0fbc3b026ba8`,
 
 
 
40
  避免 Space 重啟後在沒有程式版本變更的情況下改變模型或 speaker embedding 空間。
41
  候選語意驗證固定使用 Whisper large-v3-turbo revision
42
  `41f01f3fe87f28c78e2fbf8b568835947dd65ed9`;exact-assembled whole-output
 
37
  `0f99f2d0ebe89ac095bcc5903c4dd8f72b367286`;TTS runtime 以
38
  `ce384c8cc54efea1aaba7b9f1d7ded6c1c99aa9a` 的已驗證原始碼隨 Space 保存,
39
  Barbet 另固定在 `6fcd7ce4aa37f2250a3242995bef0fbc3b026ba8`,
40
+ Mamba2 所需的 Triton runtime 以
41
+ `v2.3.2.post1` 的 Apache-2.0 上游原始碼隨 Space 保存(不綁定
42
+ `selective_scan_cuda` 編譯產物),
43
  避免 Space 重啟後在沒有程式版本變更的情況下改變模型或 speaker embedding 空間。
44
  候選語意驗證固定使用 Whisper large-v3-turbo revision
45
  `41f01f3fe87f28c78e2fbf8b568835947dd65ed9`;exact-assembled whole-output
app.py CHANGED
@@ -12,6 +12,7 @@ from pathlib import Path
12
  import secrets
13
  import threading
14
  from types import MappingProxyType
 
15
 
16
  # Cross-process release reproducibility is a hard contract. Configure CUDA
17
  # before importing torch so cuBLAS and Mamba choose their deterministic paths.
@@ -25,7 +26,7 @@ CUDNN_ALLOW_TF32 = False
25
  EXPECTED_TORCH_VERSION = "2.11.0+cu130"
26
  EXPECTED_TORCH_CUDA_VERSION = "13.0"
27
  EXPECTED_CUDNN_VERSION = 91900
28
- EXPECTED_MAMBA_SSM_VERSION = "2.3.2.post1"
29
  EXPECTED_TRITON_VERSION = "3.6.0"
30
  if (
31
  os.environ.setdefault(
@@ -55,7 +56,7 @@ runtime_versions = {
55
  "torch": str(torch.__version__),
56
  "cuda": str(torch.version.cuda),
57
  "cudnn": torch.backends.cudnn.version(),
58
- "mamba-ssm": importlib_metadata.version("mamba-ssm"),
59
  "triton": importlib_metadata.version("triton"),
60
  }
61
  expected_runtime_versions = {
 
12
  import secrets
13
  import threading
14
  from types import MappingProxyType
15
+ from mamba_ssm import __version__ as mamba_ssm_runtime_version
16
 
17
  # Cross-process release reproducibility is a hard contract. Configure CUDA
18
  # before importing torch so cuBLAS and Mamba choose their deterministic paths.
 
26
  EXPECTED_TORCH_VERSION = "2.11.0+cu130"
27
  EXPECTED_TORCH_CUDA_VERSION = "13.0"
28
  EXPECTED_CUDNN_VERSION = 91900
29
+ EXPECTED_MAMBA_SSM_VERSION = "2.3.2.post1+bluemagpie.triton1"
30
  EXPECTED_TRITON_VERSION = "3.6.0"
31
  if (
32
  os.environ.setdefault(
 
56
  "torch": str(torch.__version__),
57
  "cuda": str(torch.version.cuda),
58
  "cudnn": torch.backends.cudnn.version(),
59
+ "mamba-ssm": mamba_ssm_runtime_version,
60
  "triton": importlib_metadata.version("triton"),
61
  }
62
  expected_runtime_versions = {
mamba_ssm/LICENSE.upstream ADDED
@@ -0,0 +1,201 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Apache License
2
+ Version 2.0, January 2004
3
+ http://www.apache.org/licenses/
4
+
5
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
6
+
7
+ 1. Definitions.
8
+
9
+ "License" shall mean the terms and conditions for use, reproduction,
10
+ and distribution as defined by Sections 1 through 9 of this document.
11
+
12
+ "Licensor" shall mean the copyright owner or entity authorized by
13
+ the copyright owner that is granting the License.
14
+
15
+ "Legal Entity" shall mean the union of the acting entity and all
16
+ other entities that control, are controlled by, or are under common
17
+ control with that entity. For the purposes of this definition,
18
+ "control" means (i) the power, direct or indirect, to cause the
19
+ direction or management of such entity, whether by contract or
20
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
21
+ outstanding shares, or (iii) beneficial ownership of such entity.
22
+
23
+ "You" (or "Your") shall mean an individual or Legal Entity
24
+ exercising permissions granted by this License.
25
+
26
+ "Source" form shall mean the preferred form for making modifications,
27
+ including but not limited to software source code, documentation
28
+ source, and configuration files.
29
+
30
+ "Object" form shall mean any form resulting from mechanical
31
+ transformation or translation of a Source form, including but
32
+ not limited to compiled object code, generated documentation,
33
+ and conversions to other media types.
34
+
35
+ "Work" shall mean the work of authorship, whether in Source or
36
+ Object form, made available under the License, as indicated by a
37
+ copyright notice that is included in or attached to the work
38
+ (an example is provided in the Appendix below).
39
+
40
+ "Derivative Works" shall mean any work, whether in Source or Object
41
+ form, that is based on (or derived from) the Work and for which the
42
+ editorial revisions, annotations, elaborations, or other modifications
43
+ represent, as a whole, an original work of authorship. For the purposes
44
+ of this License, Derivative Works shall not include works that remain
45
+ separable from, or merely link (or bind by name) to the interfaces of,
46
+ the Work and Derivative Works thereof.
47
+
48
+ "Contribution" shall mean any work of authorship, including
49
+ the original version of the Work and any modifications or additions
50
+ to that Work or Derivative Works thereof, that is intentionally
51
+ submitted to Licensor for inclusion in the Work by the copyright owner
52
+ or by an individual or Legal Entity authorized to submit on behalf of
53
+ the copyright owner. For the purposes of this definition, "submitted"
54
+ means any form of electronic, verbal, or written communication sent
55
+ to the Licensor or its representatives, including but not limited to
56
+ communication on electronic mailing lists, source code control systems,
57
+ and issue tracking systems that are managed by, or on behalf of, the
58
+ Licensor for the purpose of discussing and improving the Work, but
59
+ excluding communication that is conspicuously marked or otherwise
60
+ designated in writing by the copyright owner as "Not a Contribution."
61
+
62
+ "Contributor" shall mean Licensor and any individual or Legal Entity
63
+ on behalf of whom a Contribution has been received by Licensor and
64
+ subsequently incorporated within the Work.
65
+
66
+ 2. Grant of Copyright License. Subject to the terms and conditions of
67
+ this License, each Contributor hereby grants to You a perpetual,
68
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
69
+ copyright license to reproduce, prepare Derivative Works of,
70
+ publicly display, publicly perform, sublicense, and distribute the
71
+ Work and such Derivative Works in Source or Object form.
72
+
73
+ 3. Grant of Patent License. Subject to the terms and conditions of
74
+ this License, each Contributor hereby grants to You a perpetual,
75
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
76
+ (except as stated in this section) patent license to make, have made,
77
+ use, offer to sell, sell, import, and otherwise transfer the Work,
78
+ where such license applies only to those patent claims licensable
79
+ by such Contributor that are necessarily infringed by their
80
+ Contribution(s) alone or by combination of their Contribution(s)
81
+ with the Work to which such Contribution(s) was submitted. If You
82
+ institute patent litigation against any entity (including a
83
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
84
+ or a Contribution incorporated within the Work constitutes direct
85
+ or contributory patent infringement, then any patent licenses
86
+ granted to You under this License for that Work shall terminate
87
+ as of the date such litigation is filed.
88
+
89
+ 4. Redistribution. You may reproduce and distribute copies of the
90
+ Work or Derivative Works thereof in any medium, with or without
91
+ modifications, and in Source or Object form, provided that You
92
+ meet the following conditions:
93
+
94
+ (a) You must give any other recipients of the Work or
95
+ Derivative Works a copy of this License; and
96
+
97
+ (b) You must cause any modified files to carry prominent notices
98
+ stating that You changed the files; and
99
+
100
+ (c) You must retain, in the Source form of any Derivative Works
101
+ that You distribute, all copyright, patent, trademark, and
102
+ attribution notices from the Source form of the Work,
103
+ excluding those notices that do not pertain to any part of
104
+ the Derivative Works; and
105
+
106
+ (d) If the Work includes a "NOTICE" text file as part of its
107
+ distribution, then any Derivative Works that You distribute must
108
+ include a readable copy of the attribution notices contained
109
+ within such NOTICE file, excluding those notices that do not
110
+ pertain to any part of the Derivative Works, in at least one
111
+ of the following places: within a NOTICE text file distributed
112
+ as part of the Derivative Works; within the Source form or
113
+ documentation, if provided along with the Derivative Works; or,
114
+ within a display generated by the Derivative Works, if and
115
+ wherever such third-party notices normally appear. The contents
116
+ of the NOTICE file are for informational purposes only and
117
+ do not modify the License. You may add Your own attribution
118
+ notices within Derivative Works that You distribute, alongside
119
+ or as an addendum to the NOTICE text from the Work, provided
120
+ that such additional attribution notices cannot be construed
121
+ as modifying the License.
122
+
123
+ You may add Your own copyright statement to Your modifications and
124
+ may provide additional or different license terms and conditions
125
+ for use, reproduction, or distribution of Your modifications, or
126
+ for any such Derivative Works as a whole, provided Your use,
127
+ reproduction, and distribution of the Work otherwise complies with
128
+ the conditions stated in this License.
129
+
130
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
131
+ any Contribution intentionally submitted for inclusion in the Work
132
+ by You to the Licensor shall be under the terms and conditions of
133
+ this License, without any additional terms or conditions.
134
+ Notwithstanding the above, nothing herein shall supersede or modify
135
+ the terms of any separate license agreement you may have executed
136
+ with Licensor regarding such Contributions.
137
+
138
+ 6. Trademarks. This License does not grant permission to use the trade
139
+ names, trademarks, service marks, or product names of the Licensor,
140
+ except as required for reasonable and customary use in describing the
141
+ origin of the Work and reproducing the content of the NOTICE file.
142
+
143
+ 7. Disclaimer of Warranty. Unless required by applicable law or
144
+ agreed to in writing, Licensor provides the Work (and each
145
+ Contributor provides its Contributions) on an "AS IS" BASIS,
146
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
147
+ implied, including, without limitation, any warranties or conditions
148
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
149
+ PARTICULAR PURPOSE. You are solely responsible for determining the
150
+ appropriateness of using or redistributing the Work and assume any
151
+ risks associated with Your exercise of permissions under this License.
152
+
153
+ 8. Limitation of Liability. In no event and under no legal theory,
154
+ whether in tort (including negligence), contract, or otherwise,
155
+ unless required by applicable law (such as deliberate and grossly
156
+ negligent acts) or agreed to in writing, shall any Contributor be
157
+ liable to You for damages, including any direct, indirect, special,
158
+ incidental, or consequential damages of any character arising as a
159
+ result of this License or out of the use or inability to use the
160
+ Work (including but not limited to damages for loss of goodwill,
161
+ work stoppage, computer failure or malfunction, or any and all
162
+ other commercial damages or losses), even if such Contributor
163
+ has been advised of the possibility of such damages.
164
+
165
+ 9. Accepting Warranty or Additional Liability. While redistributing
166
+ the Work or Derivative Works thereof, You may choose to offer,
167
+ and charge a fee for, acceptance of support, warranty, indemnity,
168
+ or other liability obligations and/or rights consistent with this
169
+ License. However, in accepting such obligations, You may act only
170
+ on Your own behalf and on Your sole responsibility, not on behalf
171
+ of any other Contributor, and only if You agree to indemnify,
172
+ defend, and hold each Contributor harmless for any liability
173
+ incurred by, or claims asserted against, such Contributor by reason
174
+ of your accepting any such warranty or additional liability.
175
+
176
+ END OF TERMS AND CONDITIONS
177
+
178
+ APPENDIX: How to apply the Apache License to your work.
179
+
180
+ To apply the Apache License to your work, attach the following
181
+ boilerplate notice, with the fields enclosed by brackets "[]"
182
+ replaced with your own identifying information. (Don't include
183
+ the brackets!) The text should be enclosed in the appropriate
184
+ comment syntax for the file format. We also recommend that a
185
+ file or class name and description of purpose be included on the
186
+ same "printed page" as the copyright notice for easier
187
+ identification within third-party archives.
188
+
189
+ Copyright 2023 Tri Dao, Albert Gu
190
+
191
+ Licensed under the Apache License, Version 2.0 (the "License");
192
+ you may not use this file except in compliance with the License.
193
+ You may obtain a copy of the License at
194
+
195
+ http://www.apache.org/licenses/LICENSE-2.0
196
+
197
+ Unless required by applicable law or agreed to in writing, software
198
+ distributed under the License is distributed on an "AS IS" BASIS,
199
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
200
+ See the License for the specific language governing permissions and
201
+ limitations under the License.
mamba_ssm/PROVENANCE.md ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Vendored Mamba Triton runtime
2
+
3
+ This directory contains the Mamba 2 Triton runtime required by the pinned
4
+ OpenFormosa Barbet revision.
5
+
6
+ - Upstream: `https://github.com/state-spaces/mamba`
7
+ - Tag: `v2.3.2.post1`
8
+ - Commit: `a14b1dff0454a3bc27d9eb31355dc01e4b2490ec`
9
+ - Source archive SHA-256:
10
+ `104cc47e9101e5401a675fa2b784f2952b9b037f3b1dd83b5ac544394e95d028`
11
+ - Upstream license: Apache-2.0, retained in `LICENSE.upstream`
12
+
13
+ The files under `ops/triton/` and `utils/` are copied from that source archive;
14
+ `softplus.py` only normalizes its missing final newline. The local
15
+ `__init__.py` is intentionally reduced so importing a
16
+ Triton submodule does not import the unused `selective_scan_cuda` extension.
17
+ BlueMagpie's Barbet revision calls `mamba_chunk_scan_combined` and
18
+ `RMSNorm` from these Triton modules and implements its causal convolution
19
+ itself. No generated CUDA binary is bundled.
mamba_ssm/__init__.py ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ """Pinned Triton-only Mamba runtime used by the BlueMagpie Space.
2
+
3
+ The upstream package imports its optional compiled selective-scan extension
4
+ from ``__init__``. BlueMagpie's Barbet runtime uses the Mamba2 Triton SSD
5
+ modules directly, so this vendored entry point intentionally exposes only the
6
+ release version and leaves the unused compiled extension out of the Space.
7
+ """
8
+
9
+ __version__ = "2.3.2.post1+bluemagpie.triton1"
mamba_ssm/ops/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Operations for the pinned BlueMagpie Mamba runtime."""
mamba_ssm/ops/triton/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Exact upstream Mamba 2 Triton operations used by Barbet."""
mamba_ssm/ops/triton/k_activations.py ADDED
@@ -0,0 +1,171 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2024, Tri Dao, Albert Gu.
2
+
3
+ import torch
4
+
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from mamba_ssm.utils.determinism import autotune_configs
9
+
10
+
11
+ @triton.autotune(
12
+ configs=autotune_configs([
13
+ triton.Config({'BLOCK_N': 32}),
14
+ triton.Config({'BLOCK_N': 64}),
15
+ triton.Config({'BLOCK_N': 128}),
16
+ triton.Config({'BLOCK_N': 256}),
17
+ triton.Config({'BLOCK_N': 512}),
18
+ triton.Config({'BLOCK_N': 1024}),
19
+ ]),
20
+ key=['ncols'],
21
+ )
22
+ @triton.jit
23
+ def _swiglu_fwd_kernel(
24
+ X,
25
+ Y,
26
+ OUT,
27
+ stride_x_row, # how much to increase the pointer when moving by 1 row
28
+ stride_y_row,
29
+ stride_out_row,
30
+ ncols,
31
+ BLOCK_N: tl.constexpr,
32
+ ):
33
+ # Map the program id to the row of X and Y it should compute.
34
+ row = tl.program_id(0)
35
+ start_col = tl.program_id(1) * BLOCK_N
36
+ X += row * stride_x_row
37
+ Y += row * stride_y_row
38
+ OUT += row * stride_out_row
39
+ cols = start_col + tl.arange(0, BLOCK_N)
40
+ x = tl.load(X + cols, mask=cols < ncols, other=0.).to(tl.float32)
41
+ y = tl.load(Y + cols, mask=cols < ncols, other=0.).to(tl.float32)
42
+ out = x * tl.sigmoid(x) * y
43
+ tl.store(OUT + cols, out, mask=cols < ncols)
44
+
45
+
46
+ def _swiglu_fwd(xy, out=None):
47
+ if xy.stride(-1) != 1:
48
+ xy = xy.contiguous()
49
+ batch_shape = xy.shape[:-1]
50
+ xy = xy.reshape(-1, xy.shape[-1])
51
+ x, y = xy.chunk(2, dim=-1)
52
+ if out is None:
53
+ out = torch.empty_like(x)
54
+ else:
55
+ out = out.reshape(-1, out.shape[-1])
56
+ assert out.shape == x.shape
57
+ assert out.stride(-1) == 1
58
+ M, N = x.shape
59
+ grid = lambda META: (M, triton.cdiv(N, META['BLOCK_N']))
60
+ with torch.cuda.device(x.device.index):
61
+ _swiglu_fwd_kernel[grid](x, y, out, x.stride(0), y.stride(0), out.stride(0), N)
62
+ return out.reshape(*batch_shape, out.shape[-1])
63
+
64
+
65
+ @triton.autotune(
66
+ configs=autotune_configs([
67
+ triton.Config({'BLOCK_N': 32}),
68
+ triton.Config({'BLOCK_N': 64}),
69
+ triton.Config({'BLOCK_N': 128}),
70
+ triton.Config({'BLOCK_N': 256}),
71
+ triton.Config({'BLOCK_N': 512}),
72
+ triton.Config({'BLOCK_N': 1024}),
73
+ ]),
74
+ key=['ncols'],
75
+ )
76
+ @triton.heuristics({"RECOMPUTE_OUTPUT": lambda args: args["OUT"] is not None})
77
+ @triton.jit
78
+ def _swiglu_bwd_kernel(
79
+ X,
80
+ Y,
81
+ DOUT,
82
+ OUT,
83
+ DX,
84
+ DY,
85
+ stride_x_row, # how much to increase the pointer when moving by 1 row
86
+ stride_y_row,
87
+ stride_dout_row,
88
+ stride_out_row,
89
+ stride_dx_row,
90
+ stride_dy_row,
91
+ ncols,
92
+ BLOCK_N: tl.constexpr,
93
+ RECOMPUTE_OUTPUT: tl.constexpr,
94
+ ):
95
+ # Map the program id to the row of X and Y it should compute.
96
+ row = tl.program_id(0)
97
+ start_col = tl.program_id(1) * BLOCK_N
98
+ X += row * stride_x_row
99
+ Y += row * stride_y_row
100
+ DOUT += row * stride_dout_row
101
+ if RECOMPUTE_OUTPUT:
102
+ OUT += row * stride_out_row
103
+ DX += row * stride_dx_row
104
+ DY += row * stride_dy_row
105
+ cols = start_col + tl.arange(0, BLOCK_N)
106
+ x = tl.load(X + cols, mask=cols < ncols, other=0.).to(tl.float32)
107
+ y = tl.load(Y + cols, mask=cols < ncols, other=0.).to(tl.float32)
108
+ dout = tl.load(DOUT + cols, mask=cols < ncols, other=0.).to(tl.float32)
109
+ x_sigmoid = tl.sigmoid(x)
110
+ dx = x_sigmoid * (1 + x * (1 - x_sigmoid)) * y * dout
111
+ dy = x * x_sigmoid * dout
112
+ tl.store(DX + cols, dx, mask=cols < ncols)
113
+ tl.store(DY + cols, dy, mask=cols < ncols)
114
+ if RECOMPUTE_OUTPUT:
115
+ out = x * x_sigmoid * y
116
+ tl.store(OUT + cols, out, mask=cols < ncols)
117
+
118
+
119
+ def _swiglu_bwd(xy, dout, dxy=None, recompute_output=False, out=None):
120
+ if xy.stride(-1) != 1:
121
+ xy = xy.contiguous()
122
+ if dout.stride(-1) != 1:
123
+ dout = dout.contiguous()
124
+ batch_shape = xy.shape[:-1]
125
+ xy = xy.reshape(-1, xy.shape[-1])
126
+ x, y = xy.chunk(2, dim=-1)
127
+ dout = dout.reshape(-1, dout.shape[-1])
128
+ assert dout.shape == x.shape
129
+ if dxy is None:
130
+ dxy = torch.empty_like(xy)
131
+ else:
132
+ dxy = dxy.reshape(-1, dxy.shape[-1])
133
+ assert dxy.shape == xy.shape
134
+ dx, dy = dxy.chunk(2, dim=-1)
135
+ assert dx.stride(-1) == 1
136
+ assert dy.stride(-1) == 1
137
+ if recompute_output:
138
+ if out is None:
139
+ out = torch.empty_like(x)
140
+ else:
141
+ out = out.reshape(-1, out.shape[-1])
142
+ assert out.shape == x.shape
143
+ assert out.stride(-1) == 1
144
+ M, N = x.shape
145
+ grid = lambda META: (M, triton.cdiv(N, META['BLOCK_N']))
146
+ with torch.cuda.device(x.device.index):
147
+ _swiglu_bwd_kernel[grid](x, y, dout, out if recompute_output else None, dx, dy,
148
+ x.stride(0), y.stride(0), dout.stride(0),
149
+ out.stride(0) if recompute_output else 0,
150
+ dx.stride(0), dy.stride(0),
151
+ N)
152
+ if not recompute_output:
153
+ return dxy.reshape(*batch_shape, dxy.shape[-1])
154
+ else:
155
+ return dxy.reshape(*batch_shape, dxy.shape[-1]), out.reshape(*batch_shape, out.shape[-1])
156
+
157
+
158
+ class SwiGLU(torch.autograd.Function):
159
+
160
+ @staticmethod
161
+ def forward(ctx, xy):
162
+ ctx.save_for_backward(xy)
163
+ return _swiglu_fwd(xy)
164
+
165
+ @staticmethod
166
+ def backward(ctx, dout):
167
+ xy, = ctx.saved_tensors
168
+ return _swiglu_bwd(xy, dout)
169
+
170
+
171
+ swiglu = SwiGLU.apply
mamba_ssm/ops/triton/layernorm_gated.py ADDED
@@ -0,0 +1,437 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2024, Tri Dao.
2
+ # Based on the Triton LayerNorm tutorial: https://triton-lang.org/main/getting-started/tutorials/05-layer-norm.html
3
+ # For the backward pass, we keep weight_grad and bias_grad in registers and accumulate.
4
+ # This backward pass is faster for dimensions up to 8k, but after that it's much slower due to register spilling.
5
+ # The models we train have hidden dim up to 8k anyway (e.g. Llama 70B), so this is fine.
6
+
7
+ import math
8
+
9
+ import torch
10
+ import torch.nn.functional as F
11
+
12
+ import triton
13
+ import triton.language as tl
14
+
15
+ from einops import rearrange
16
+
17
+
18
+ def rms_norm_ref(x, weight, bias, z=None, eps=1e-6, group_size=None, norm_before_gate=True, upcast=True):
19
+ dtype = x.dtype
20
+ N = x.shape[-1]
21
+ weight = weight.float()
22
+ bias = bias.float() if bias is not None else None
23
+ if upcast:
24
+ x = x.float()
25
+ z = z.float() if z is not None else z
26
+ if z is not None and not norm_before_gate:
27
+ x = x * F.silu(z)
28
+ if group_size is None:
29
+ rstd = 1 / torch.sqrt((x.square()).mean(dim=-1, keepdim=True) + eps)
30
+ out = (x * rstd * weight) + bias if bias is not None else (x * rstd * weight)
31
+ else:
32
+ x_group = rearrange(x, "... (g d) -> ... g d", d=group_size)
33
+ rstd = 1 / torch.sqrt((x_group.square()).mean(dim=-1, keepdim=True) + eps)
34
+ out = rearrange(x_group * rstd, "... g d -> ... (g d)") * weight
35
+ if bias is not None:
36
+ out = out + bias
37
+ if z is not None and norm_before_gate:
38
+ out *= F.silu(z)
39
+ return out.to(dtype)
40
+
41
+
42
+ @triton.heuristics({"HAS_BIAS": lambda args: args["B"] is not None})
43
+ @triton.heuristics({"HAS_Z": lambda args: args["Z"] is not None})
44
+ @triton.jit
45
+ def _layer_norm_fwd_1pass_kernel(
46
+ X, # pointer to the input
47
+ Y, # pointer to the output
48
+ W, # pointer to the weights
49
+ B, # pointer to the biases
50
+ Z, # pointer to the other branch
51
+ Mean, # pointer to the mean
52
+ Rstd, # pointer to the 1/std
53
+ stride_x_row, # how much to increase the pointer when moving by 1 row
54
+ stride_y_row,
55
+ stride_z_row,
56
+ M, # number of rows in X
57
+ N, # number of columns in X
58
+ eps, # epsilon to avoid division by zero
59
+ BLOCK_N: tl.constexpr,
60
+ HAS_BIAS: tl.constexpr,
61
+ HAS_Z: tl.constexpr,
62
+ NORM_BEFORE_GATE: tl.constexpr,
63
+ IS_RMS_NORM: tl.constexpr,
64
+ ):
65
+ # Map the program id to the row of X and Y it should compute.
66
+ row = tl.program_id(0)
67
+ group = tl.program_id(1)
68
+ X += row * stride_x_row + group * N
69
+ Y += row * stride_y_row + group * N
70
+ if HAS_Z:
71
+ Z += row * stride_z_row + group * N
72
+ if not IS_RMS_NORM:
73
+ Mean += group * M
74
+ Rstd += group * M
75
+ W += group * N
76
+ if HAS_BIAS:
77
+ B += group * N
78
+ # Compute mean and variance
79
+ cols = tl.arange(0, BLOCK_N)
80
+ x = tl.load(X + cols, mask=cols < N, other=0.).to(tl.float32)
81
+ if HAS_Z and not NORM_BEFORE_GATE:
82
+ z = tl.load(Z + cols, mask=cols < N).to(tl.float32)
83
+ x *= z * tl.sigmoid(z)
84
+ if not IS_RMS_NORM:
85
+ mean = tl.sum(x, axis=0) / N
86
+ tl.store(Mean + row, mean)
87
+ xbar = tl.where(cols < N, x - mean, 0.)
88
+ var = tl.sum(xbar * xbar, axis=0) / N
89
+ else:
90
+ xbar = tl.where(cols < N, x, 0.)
91
+ var = tl.sum(xbar * xbar, axis=0) / N
92
+ rstd = 1 / tl.sqrt(var + eps)
93
+ tl.store(Rstd + row, rstd)
94
+ # Normalize and apply linear transformation
95
+ mask = cols < N
96
+ w = tl.load(W + cols, mask=mask).to(tl.float32)
97
+ if HAS_BIAS:
98
+ b = tl.load(B + cols, mask=mask).to(tl.float32)
99
+ x_hat = (x - mean) * rstd if not IS_RMS_NORM else x * rstd
100
+ y = x_hat * w + b if HAS_BIAS else x_hat * w
101
+ if HAS_Z and NORM_BEFORE_GATE:
102
+ z = tl.load(Z + cols, mask=mask).to(tl.float32)
103
+ y *= z * tl.sigmoid(z)
104
+ # Write output
105
+ tl.store(Y + cols, y, mask=mask)
106
+
107
+
108
+ def _layer_norm_fwd(x, weight, bias, eps, z=None, out=None, group_size=None, norm_before_gate=True, is_rms_norm=False):
109
+ M, N = x.shape
110
+ if group_size is None:
111
+ group_size = N
112
+ assert N % group_size == 0
113
+ ngroups = N // group_size
114
+ assert x.stride(-1) == 1
115
+ if z is not None:
116
+ assert z.stride(-1) == 1
117
+ assert z.shape == (M, N)
118
+ assert weight.shape == (N,)
119
+ assert weight.stride(-1) == 1
120
+ if bias is not None:
121
+ assert bias.stride(-1) == 1
122
+ assert bias.shape == (N,)
123
+ # allocate output
124
+ if out is not None:
125
+ assert out.shape == x.shape
126
+ else:
127
+ out = torch.empty_like(x)
128
+ assert out.stride(-1) == 1
129
+ mean = torch.empty((ngroups * M, ), dtype=torch.float32, device=x.device) if not is_rms_norm else None
130
+ rstd = torch.empty((ngroups * M, ), dtype=torch.float32, device=x.device)
131
+ # Less than 64KB per feature: enqueue fused kernel
132
+ MAX_FUSED_SIZE = 65536 // x.element_size()
133
+ BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(group_size))
134
+ if group_size > BLOCK_N:
135
+ raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
136
+ # heuristics for number of warps
137
+ num_warps = min(max(BLOCK_N // 256, 1), 8)
138
+ grid = (M, ngroups)
139
+ with torch.cuda.device(x.device.index):
140
+ _layer_norm_fwd_1pass_kernel[grid](x, out, weight, bias, z, mean, rstd,
141
+ x.stride(0), out.stride(0), z.stride(0) if z is not None else 0,
142
+ M, group_size, eps,
143
+ BLOCK_N=BLOCK_N,
144
+ NORM_BEFORE_GATE=norm_before_gate,
145
+ IS_RMS_NORM=is_rms_norm,
146
+ num_warps=num_warps)
147
+ return out, mean, rstd
148
+
149
+
150
+
151
+ @triton.heuristics({"HAS_BIAS": lambda args: args["B"] is not None})
152
+ @triton.heuristics({"HAS_Z": lambda args: args["Z"] is not None})
153
+ @triton.heuristics({"RECOMPUTE_OUTPUT": lambda args: args["Y"] is not None})
154
+ @triton.jit
155
+ def _layer_norm_bwd_kernel(
156
+ X, # pointer to the input
157
+ W, # pointer to the weights
158
+ B, # pointer to the biases
159
+ Z, # pointer to the other branch
160
+ Y, # pointer to the output to be recomputed
161
+ DY, # pointer to the output gradient
162
+ DX, # pointer to the input gradient
163
+ DW, # pointer to the partial sum of weights gradient
164
+ DB, # pointer to the partial sum of biases gradient
165
+ DZ, # pointer to the other branch
166
+ Mean, # pointer to the mean
167
+ Rstd, # pointer to the 1/std
168
+ stride_x_row, # how much to increase the pointer when moving by 1 row
169
+ stride_z_row,
170
+ stride_y_row,
171
+ stride_dy_row,
172
+ stride_dx_row,
173
+ stride_dz_row,
174
+ stride_dw_row,
175
+ stride_db_row,
176
+ M, # number of rows in X
177
+ N, # number of columns in X
178
+ eps, # epsilon to avoid division by zero
179
+ rows_per_program,
180
+ NORM_BEFORE_GATE: tl.constexpr,
181
+ IS_RMS_NORM: tl.constexpr,
182
+ HAS_BIAS: tl.constexpr,
183
+ HAS_Z: tl.constexpr,
184
+ RECOMPUTE_OUTPUT: tl.constexpr,
185
+ BLOCK_N: tl.constexpr,
186
+ ):
187
+ # Map the program id to the elements of X, DX, and DY it should compute.
188
+ row_block_id = tl.program_id(0)
189
+ group = tl.program_id(1)
190
+ row_start = row_block_id * rows_per_program
191
+ cols = tl.arange(0, BLOCK_N)
192
+ mask = cols < N
193
+ X += row_start * stride_x_row + group * N
194
+ if HAS_Z:
195
+ Z += row_start * stride_z_row + group * N
196
+ DZ += row_start * stride_dz_row + group * N
197
+ DY += row_start * stride_dy_row + group * N
198
+ DX += row_start * stride_dx_row + group * N
199
+ if RECOMPUTE_OUTPUT:
200
+ Y += row_start * stride_y_row + group * N
201
+ if not IS_RMS_NORM:
202
+ Mean += group * M
203
+ Rstd += group * M
204
+ W += group * N
205
+ w = tl.load(W + cols, mask=mask).to(tl.float32)
206
+ if (RECOMPUTE_OUTPUT or HAS_Z) and HAS_BIAS:
207
+ B += group * N
208
+ b = tl.load(B + cols, mask=mask, other=0.).to(tl.float32)
209
+ dw = tl.zeros((BLOCK_N,), dtype=tl.float32)
210
+ if HAS_BIAS:
211
+ db = tl.zeros((BLOCK_N,), dtype=tl.float32)
212
+ row_end = min((row_block_id + 1) * rows_per_program, M)
213
+ for row in range(row_start, row_end):
214
+ # Load data to SRAM
215
+ x = tl.load(X + cols, mask=mask, other=0).to(tl.float32)
216
+ dy = tl.load(DY + cols, mask=mask, other=0).to(tl.float32)
217
+ if not IS_RMS_NORM:
218
+ mean = tl.load(Mean + row)
219
+ if HAS_Z and not NORM_BEFORE_GATE:
220
+ z = tl.load(Z + cols, mask=mask, other=0.).to(tl.float32)
221
+ x_og = x
222
+ x = x_og * z * tl.sigmoid(z)
223
+ rstd = tl.load(Rstd + row)
224
+ # Compute dx
225
+ xhat = (x - mean) * rstd if not IS_RMS_NORM else x * rstd
226
+ xhat = tl.where(mask, xhat, 0.)
227
+ if HAS_Z and NORM_BEFORE_GATE:
228
+ z = tl.load(Z + cols, mask=mask, other=0.).to(tl.float32)
229
+ z_sigmoid = tl.sigmoid(z)
230
+ y = xhat * w + b if HAS_BIAS else xhat * w
231
+ if RECOMPUTE_OUTPUT:
232
+ tl.store(Y + cols, y * z * z_sigmoid, mask=mask)
233
+ dz = dy * y * z_sigmoid * (1 + z * (1 - z_sigmoid))
234
+ tl.store(DZ + cols, dz, mask=mask)
235
+ dy *= z * z_sigmoid
236
+ else:
237
+ if RECOMPUTE_OUTPUT:
238
+ y = xhat * w + b if HAS_BIAS else xhat * w
239
+ tl.store(Y + cols, y, mask=mask)
240
+ wdy = w * dy
241
+ c1 = tl.sum(xhat * wdy, axis=0) / N
242
+ if not IS_RMS_NORM:
243
+ c2 = tl.sum(wdy, axis=0) / N
244
+ dx = (wdy - (xhat * c1 + c2)) * rstd
245
+ else:
246
+ dx = (wdy - xhat * c1) * rstd
247
+ dw += dy * xhat
248
+ if HAS_BIAS:
249
+ db += dy
250
+ if HAS_Z and not NORM_BEFORE_GATE:
251
+ z_sigmoid = tl.sigmoid(z)
252
+ dz = dx * x_og * z_sigmoid * (1 + z * (1 - z_sigmoid))
253
+ tl.store(DZ + cols, dz, mask=mask)
254
+ dx *= z * z_sigmoid
255
+ # Write dx
256
+ tl.store(DX + cols, dx, mask=mask)
257
+
258
+ X += stride_x_row
259
+ if HAS_Z:
260
+ Z += stride_z_row
261
+ DZ += stride_dz_row
262
+ if RECOMPUTE_OUTPUT:
263
+ Y += stride_y_row
264
+ DY += stride_dy_row
265
+ DX += stride_dx_row
266
+ tl.store(DW + row_block_id * stride_dw_row + group * N + cols, dw, mask=mask)
267
+ if HAS_BIAS:
268
+ tl.store(DB + row_block_id * stride_db_row + group * N + cols, db, mask=mask)
269
+
270
+
271
+ def _layer_norm_bwd(dy, x, weight, bias, eps, mean, rstd, z=None, group_size=None,
272
+ norm_before_gate=True, is_rms_norm=False, recompute_output=False, dz=None, out=None):
273
+ M, N = x.shape
274
+ if group_size is None:
275
+ group_size = N
276
+ assert N % group_size == 0
277
+ ngroups = N // group_size
278
+ assert x.stride(-1) == 1
279
+ assert dy.stride(-1) == 1
280
+ assert dy.shape == (M, N)
281
+ if z is not None:
282
+ assert z.stride(-1) == 1
283
+ assert z.shape == (M, N)
284
+ assert weight.shape == (N,)
285
+ assert weight.stride(-1) == 1
286
+ if bias is not None:
287
+ assert bias.stride(-1) == 1
288
+ assert bias.shape == (N,)
289
+ # allocate output
290
+ dx = torch.empty_like(x)
291
+ if dz is not None:
292
+ assert z is not None
293
+ assert dz.shape == z.shape
294
+ assert dz.stride(-1) == 1
295
+ else:
296
+ dz = torch.empty_like(z) if z is not None else None
297
+ if recompute_output:
298
+ if out is None:
299
+ out = torch.empty_like(x)
300
+ assert out.shape == x.shape
301
+
302
+ # Less than 64KB per feature: enqueue fused kernel
303
+ MAX_FUSED_SIZE = 65536 // x.element_size()
304
+ BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(group_size))
305
+ if group_size > BLOCK_N:
306
+ raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
307
+ # heuristics for number of warps
308
+ num_warps = min(max(BLOCK_N // 256, 1), 8)
309
+ sm_count = torch.cuda.get_device_properties(x.device).multi_processor_count
310
+ # If group size is small (e.g., 64), we're only using 1 warp. So having just 108 programs
311
+ # would limit the occupancy.
312
+ nrow_groups = math.ceil(sm_count * math.ceil(4 / num_warps) / ngroups)
313
+ _dw = torch.empty((nrow_groups, N), dtype=torch.float32, device=weight.device)
314
+ _db = torch.empty((nrow_groups, N), dtype=torch.float32, device=bias.device) if bias is not None else None
315
+ rows_per_program = math.ceil(M / nrow_groups)
316
+ grid = (nrow_groups, ngroups)
317
+ with torch.cuda.device(x.device.index):
318
+ _layer_norm_bwd_kernel[grid](x, weight, bias, z, out if recompute_output else None,
319
+ dy, dx, _dw, _db, dz, mean, rstd,
320
+ x.stride(0),
321
+ z.stride(0) if z is not None else 0,
322
+ 0 if not recompute_output else out.stride(0),
323
+ dy.stride(0), dx.stride(0),
324
+ dz.stride(0) if dz is not None else 0,
325
+ _dw.stride(0),
326
+ _db.stride(0) if _db is not None else 0,
327
+ M, group_size, eps,
328
+ rows_per_program,
329
+ BLOCK_N=BLOCK_N,
330
+ NORM_BEFORE_GATE=norm_before_gate,
331
+ IS_RMS_NORM=is_rms_norm,
332
+ num_warps=num_warps)
333
+ dw = _dw.sum(0).to(weight.dtype)
334
+ db = _db.sum(0).to(bias.dtype) if bias is not None else None
335
+ return (dx, dw, db, dz) if not recompute_output else (dx, dw, db, dz, out)
336
+
337
+
338
+ class LayerNormFn(torch.autograd.Function):
339
+
340
+ @staticmethod
341
+ def forward(ctx, x, weight, bias, z=None, eps=1e-6, group_size=None, norm_before_gate=True,
342
+ is_rms_norm=False):
343
+ """If z is not None, we do norm(x) * silu(z) if norm_before_gate, else norm(x * silu(z))
344
+ """
345
+
346
+ x_shape_og = x.shape
347
+ # reshape input data into 2D tensor
348
+ x = x.reshape(-1, x.shape[-1])
349
+ if x.stride(-1) != 1:
350
+ x = x.contiguous()
351
+ if z is not None:
352
+ assert z.shape == x_shape_og
353
+ z = z.reshape(-1, z.shape[-1])
354
+ if z.stride(-1) != 1:
355
+ z = z.contiguous()
356
+ weight = weight.contiguous()
357
+ if bias is not None:
358
+ bias = bias.contiguous()
359
+ y, mean, rstd = _layer_norm_fwd(x, weight, bias, eps, z=z, group_size=group_size, norm_before_gate=norm_before_gate, is_rms_norm=is_rms_norm)
360
+ ctx.save_for_backward(x, weight, bias, mean, rstd, z)
361
+ ctx.x_shape_og = x_shape_og
362
+ ctx.eps = eps
363
+ ctx.group_size = group_size
364
+ ctx.norm_before_gate = norm_before_gate
365
+ ctx.is_rms_norm = is_rms_norm
366
+ return y.reshape(x_shape_og)
367
+
368
+ @staticmethod
369
+ def backward(ctx, dy):
370
+ x, weight, bias, mean, rstd, z = ctx.saved_tensors
371
+ dy = dy.reshape(-1, dy.shape[-1])
372
+ if dy.stride(-1) != 1:
373
+ dy = dy.contiguous()
374
+ assert dy.shape == x.shape
375
+ dx, dw, db, dz = _layer_norm_bwd(dy, x, weight, bias, ctx.eps, mean, rstd, z, ctx.group_size,
376
+ ctx.norm_before_gate, ctx.is_rms_norm)
377
+ return dx.reshape(ctx.x_shape_og), dw, db, dz.reshape(ctx.x_shape_og) if dz is not None else None, None, None, None, None
378
+
379
+
380
+ def layernorm_fn(x, weight, bias, z=None, eps=1e-6, group_size=None, norm_before_gate=True, is_rms_norm=False):
381
+ return LayerNormFn.apply(x, weight, bias, z, eps, group_size, norm_before_gate, is_rms_norm)
382
+
383
+
384
+ def rmsnorm_fn(x, weight, bias, z=None, eps=1e-6, group_size=None, norm_before_gate=True):
385
+ return LayerNormFn.apply(x, weight, bias, z, eps, group_size, norm_before_gate, True)
386
+
387
+
388
+ class LayerNorm(torch.nn.Module):
389
+
390
+ def __init__(self, hidden_size, eps=1e-5, group_size=None, norm_before_gate=True, device=None, dtype=None):
391
+ """If group_size is not None, we do GroupNorm with each group having group_size elements.
392
+ group_size=None is equivalent to group_size=hidden_size (i.e. there's only 1 group).
393
+ """
394
+
395
+ factory_kwargs = {"device": device, "dtype": dtype}
396
+ super().__init__()
397
+ self.eps = eps
398
+ self.weight = torch.nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
399
+ self.bias = torch.nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
400
+ self.group_size = group_size
401
+ self.norm_before_gate = norm_before_gate
402
+ self.reset_parameters()
403
+
404
+ def reset_parameters(self):
405
+ torch.nn.init.ones_(self.weight)
406
+ torch.nn.init.zeros_(self.bias)
407
+
408
+ def forward(self, x, z=None):
409
+ """If z is not None, we do norm(x) * silu(z) if norm_before_gate, else norm(x * silu(z))
410
+ """
411
+ return layernorm_fn(x, self.weight, self.bias, z=z, group_size=self.group_size, eps=self.eps,
412
+ norm_before_gate=self.norm_before_gate)
413
+
414
+
415
+ class RMSNorm(torch.nn.Module):
416
+
417
+ def __init__(self, hidden_size, eps=1e-5, group_size=None, norm_before_gate=True, device=None, dtype=None):
418
+ """If group_size is not None, we do GroupNorm with each group having group_size elements.
419
+ group_size=None is equivalent to group_size=hidden_size (i.e. there's only 1 group).
420
+ """
421
+ factory_kwargs = {"device": device, "dtype": dtype}
422
+ super().__init__()
423
+ self.eps = eps
424
+ self.weight = torch.nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
425
+ self.register_parameter("bias", None)
426
+ self.group_size = group_size
427
+ self.norm_before_gate = norm_before_gate
428
+ self.reset_parameters()
429
+
430
+ def reset_parameters(self):
431
+ torch.nn.init.ones_(self.weight)
432
+
433
+ def forward(self, x, z=None):
434
+ """If z is not None, we do norm(x) * silu(z) if norm_before_gate, else norm(x * silu(z))
435
+ """
436
+ return rmsnorm_fn(x, self.weight, self.bias, z=z, eps=self.eps, group_size=self.group_size,
437
+ norm_before_gate=self.norm_before_gate)
mamba_ssm/ops/triton/softplus.py ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import triton
2
+ import triton.language as tl
3
+ from packaging import version
4
+
5
+ TRITON3 = version.parse(triton.__version__) >= version.parse("3.0.0")
6
+
7
+
8
+ if TRITON3:
9
+ @triton.jit
10
+ def softplus(dt):
11
+ return tl.math.log(tl.math.exp(dt) + 1)
12
+ else:
13
+ @triton.jit
14
+ def softplus(dt):
15
+ return tl.math.log1p(tl.exp(dt))
mamba_ssm/ops/triton/ssd_bmm.py ADDED
@@ -0,0 +1,264 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2024, Tri Dao, Albert Gu.
2
+
3
+ """We want triton==2.1.0 or 2.2.0 for this
4
+ """
5
+
6
+ import math
7
+ import torch
8
+ import torch.nn.functional as F
9
+
10
+ import triton
11
+ import triton.language as tl
12
+
13
+ from einops import rearrange, repeat
14
+
15
+ from mamba_ssm.utils.determinism import autotune_configs
16
+
17
+
18
+ def init_to_zero(names):
19
+ return lambda nargs: [nargs[name].zero_() for name in names if nargs[name] is not None]
20
+
21
+
22
+ @triton.autotune(
23
+ configs=autotune_configs([
24
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64}, num_stages=3, num_warps=8),
25
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
26
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
27
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
28
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
29
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
30
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=2),
31
+ triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=2),
32
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=2),
33
+ ]),
34
+ key=['chunk_size', 'K', 'IS_CAUSAL'],
35
+ )
36
+ @triton.jit
37
+ def _bmm_chunk_fwd_kernel(
38
+ # Pointers to matrices
39
+ a_ptr, b_ptr, out_ptr, seq_idx_ptr,
40
+ # Matrix dimensions
41
+ seqlen, chunk_size, K, ngroups,
42
+ stride_a_batch, stride_a_seqlen, stride_a_head, stride_ak,
43
+ stride_b_batch, stride_b_seqlen, stride_b_head, stride_bk,
44
+ stride_out_batch, stride_out_chunk, stride_out_head, stride_outm, stride_outn,
45
+ stride_seq_idx_batch, stride_seq_idx_seqlen,
46
+ # Meta-parameters
47
+ IS_CAUSAL: tl.constexpr,
48
+ dot_dtype: tl.constexpr,
49
+ HAS_SEQ_IDX: tl.constexpr,
50
+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
51
+ ):
52
+ pid_b = tl.program_id(axis=1)
53
+ pid_ch = tl.program_id(axis=2)
54
+ pid_c = pid_ch // ngroups
55
+ pid_h = pid_ch - pid_c * ngroups
56
+ num_pid_n = tl.cdiv(chunk_size, BLOCK_SIZE_N)
57
+ pid_m = tl.program_id(axis=0) // num_pid_n
58
+ pid_n = tl.program_id(axis=0) % num_pid_n
59
+ if IS_CAUSAL:
60
+ if pid_n * BLOCK_SIZE_N >= (pid_m + 1) * BLOCK_SIZE_M:
61
+ return
62
+ a_ptr += pid_b * stride_a_batch + pid_c * chunk_size * stride_a_seqlen + pid_h * stride_a_head
63
+ b_ptr += pid_b * stride_b_batch + pid_c * chunk_size * stride_b_seqlen + pid_h * stride_b_head
64
+ if HAS_SEQ_IDX:
65
+ seq_idx_ptr += pid_b * stride_seq_idx_batch + pid_c * chunk_size * stride_seq_idx_seqlen
66
+
67
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
68
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
69
+ offs_k = tl.arange(0, BLOCK_SIZE_K)
70
+ a_ptrs = a_ptr + (offs_m[:, None] * stride_a_seqlen + offs_k[None, :] * stride_ak)
71
+ b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_n[None, :] * stride_b_seqlen)
72
+ chunk_size_limit = min(chunk_size, seqlen - pid_c * chunk_size)
73
+
74
+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
75
+ for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
76
+ a = tl.load(a_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_k[None, :] < K - k * BLOCK_SIZE_K), other=0.0).to(dot_dtype)
77
+ b = tl.load(b_ptrs, mask=(offs_k[:, None] < K - k * BLOCK_SIZE_K) & (offs_n[None, :] < chunk_size_limit), other=0.0).to(dot_dtype)
78
+ acc += tl.dot(a, b)
79
+ a_ptrs += BLOCK_SIZE_K * stride_ak
80
+ b_ptrs += BLOCK_SIZE_K * stride_bk
81
+
82
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
83
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
84
+ if HAS_SEQ_IDX:
85
+ chunk_size_limit = min(chunk_size, seqlen - pid_c * chunk_size)
86
+ seq_idx_m = tl.load(seq_idx_ptr + offs_m * stride_seq_idx_seqlen, mask=offs_m < chunk_size_limit, other=-1)
87
+ seq_idx_n = tl.load(seq_idx_ptr + offs_n * stride_seq_idx_seqlen, mask=offs_n < chunk_size_limit, other=-2)
88
+ acc = tl.where(seq_idx_m[:, None] == seq_idx_n[None, :], acc, 0.0)
89
+ out = acc.to(out_ptr.dtype.element_ty)
90
+
91
+ out_ptr += pid_b * stride_out_batch + pid_c * stride_out_chunk + pid_h * stride_out_head
92
+ out_ptrs = out_ptr + (stride_outm * offs_m[:, None] + offs_n[None, :] * stride_outn)
93
+ tl.store(out_ptrs, out, mask=(offs_m[:, None] < chunk_size) & (offs_n[None, :] < chunk_size))
94
+
95
+
96
+ @triton.autotune(
97
+ configs=autotune_configs([
98
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_CS': 64}, num_stages=3, num_warps=8),
99
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_CS': 32}, num_stages=4, num_warps=4),
100
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_CS': 32}, num_stages=4, num_warps=4),
101
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_CS': 32}, num_stages=4, num_warps=4),
102
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_CS': 32}, num_stages=4, num_warps=4),
103
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_CS': 32}, num_stages=4, num_warps=4),
104
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_CS': 32}, num_stages=5, num_warps=2),
105
+ triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_CS': 32}, num_stages=5, num_warps=2),
106
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_CS': 32}, num_stages=4, num_warps=2),
107
+ ]),
108
+ key=['chunk_size', 'K'],
109
+ )
110
+ @triton.jit
111
+ def _bmm_chunk_bwd_kernel(
112
+ # Pointers to matrices
113
+ a_ptr, dout_ptr, db_ptr, res_ptr,
114
+ # Matrix dimensions
115
+ seqlen, chunk_size, K, ngroups,
116
+ stride_a_batch, stride_a_seqlen, stride_a_head, stride_ak,
117
+ stride_dout_batch, stride_dout_chunk, stride_dout_head, stride_dout_csize_m, stride_dout_csize_n,
118
+ stride_db_batch, stride_db_seqlen, stride_db_head, stride_db_k,
119
+ stride_res_batch, stride_res_seqlen, stride_res_head, stride_res_k,
120
+ # Meta-parameters
121
+ dot_dtype: tl.constexpr,
122
+ HAS_RESIDUAL: tl.constexpr,
123
+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_CS: tl.constexpr,
124
+ ):
125
+ pid_b = tl.program_id(axis=1)
126
+ pid_ch = tl.program_id(axis=2)
127
+ pid_c = pid_ch // ngroups
128
+ pid_h = pid_ch - pid_c * ngroups
129
+ num_pid_n = tl.cdiv(K, BLOCK_SIZE_N)
130
+ pid_m = tl.program_id(axis=0) // num_pid_n
131
+ pid_n = tl.program_id(axis=0) % num_pid_n
132
+
133
+ a_ptr += pid_b * stride_a_batch + pid_c * chunk_size * stride_a_seqlen + pid_h * stride_a_head
134
+ dout_ptr += pid_b * stride_dout_batch + pid_c * stride_dout_chunk + pid_h * stride_dout_head
135
+
136
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
137
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
138
+ offs_cs = tl.arange(0, BLOCK_SIZE_CS)
139
+ dout_ptrs = dout_ptr + (offs_m[:, None] * stride_dout_csize_n + offs_cs[None, :] * stride_dout_csize_m)
140
+ a_ptrs = a_ptr + (offs_cs[:, None] * stride_a_seqlen + offs_n[None, :] * stride_ak)
141
+ chunk_size_limit = min(chunk_size, seqlen - pid_c * chunk_size)
142
+
143
+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
144
+ for cs in range(0, tl.cdiv(chunk_size_limit, BLOCK_SIZE_CS)):
145
+ dout = tl.load(dout_ptrs, mask=(offs_m[:, None] < chunk_size) & (offs_cs[None, :] < chunk_size_limit - cs * BLOCK_SIZE_CS), other=0.0).to(dot_dtype)
146
+ a = tl.load(a_ptrs, mask=(offs_cs[:, None] < chunk_size_limit - cs * BLOCK_SIZE_CS) & (offs_n[None, :] < K), other=0.0).to(dot_dtype)
147
+ acc += tl.dot(dout, a)
148
+ dout_ptrs += BLOCK_SIZE_CS * stride_dout_csize_m
149
+ a_ptrs += BLOCK_SIZE_CS * stride_a_seqlen
150
+
151
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
152
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
153
+ if HAS_RESIDUAL:
154
+ res_ptr += pid_b * stride_res_batch + pid_c * chunk_size * stride_res_seqlen + pid_h * stride_res_head
155
+ res_ptrs = res_ptr + (offs_m[:, None] * stride_res_seqlen + offs_n[None, :] * stride_res_k)
156
+ res = tl.load(res_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_n[None, :] < K)).to(tl.float32)
157
+ acc += res
158
+ db = acc.to(db_ptr.dtype.element_ty)
159
+
160
+ db_ptr += pid_b * stride_db_batch + pid_c * chunk_size * stride_db_seqlen + pid_h * stride_db_head
161
+ db_ptrs = db_ptr + (offs_m[:, None] * stride_db_seqlen + offs_n[None, :] * stride_db_k)
162
+ tl.store(db_ptrs, db, mask=(offs_m[:, None] < chunk_size_limit) & (offs_n[None, :] < K))
163
+
164
+
165
+ def _bmm_chunk_fwd(a, b, chunk_size, seq_idx=None, causal=False, output_dtype=None):
166
+ """
167
+ Argument:
168
+ a: (batch, seqlen, k) or (batch, seqlen, ngroups, k)
169
+ b: (batch, seqlen, k) or (batch, seqlen, ngroups, k)
170
+ seq_idx: (batch, seqlen) or None. out[i, j] for seq_idx[i] != seq_idx[j] will be zeroed out.
171
+ causal: if True, then out[i, j] for i > j will be arbitrary, only out[i, j] for i <= j are
172
+ guaranteed to be correct.
173
+ Return:
174
+ out: (batch, nchunks, chunk_size, chunk_size) or (batch, nchunks, ngroups, chunk_size, chunk_size)
175
+ """
176
+ # Check constraints.
177
+ has_groups = a.dim() == 4
178
+ if not has_groups:
179
+ batch, seqlen, k = a.shape
180
+ else:
181
+ batch, seqlen, ngroups, k = a.shape
182
+ assert b.shape == a.shape
183
+ if seq_idx is not None:
184
+ assert seq_idx.shape == (batch, seqlen)
185
+ if a.stride(-1) != 1 and a.stride(1) != 1:
186
+ a = a.contiguous()
187
+ if b.stride(-1) != 1 and b.stride(1) != 1:
188
+ b = b.contiguous()
189
+ nchunks = math.ceil(seqlen / chunk_size)
190
+ # Allocates output.
191
+ out_dtype = a.dtype if output_dtype is None else output_dtype
192
+ out = torch.empty((batch, nchunks, chunk_size, chunk_size) if not has_groups else (batch, nchunks, ngroups, chunk_size, chunk_size),
193
+ device=a.device, dtype=out_dtype)
194
+ dot_dtype = (tl.bfloat16 if a.dtype == torch.bfloat16 or b.dtype == torch.bfloat16 else
195
+ (tl.float16 if a.dtype == torch.float16 or b.dtype == torch.float16 else tl.float32))
196
+ grid = lambda META: (triton.cdiv(chunk_size, META['BLOCK_SIZE_M']) * triton.cdiv(chunk_size, META['BLOCK_SIZE_N']),
197
+ batch, nchunks if not has_groups else nchunks * ngroups)
198
+ with torch.cuda.device(a.device.index):
199
+ _bmm_chunk_fwd_kernel[grid](
200
+ a, b, out, seq_idx,
201
+ seqlen, chunk_size, k, ngroups if has_groups else 1,
202
+ a.stride(0), a.stride(1), 0 if not has_groups else a.stride(2), a.stride(-1),
203
+ b.stride(0), b.stride(1), 0 if not has_groups else b.stride(2), b.stride(-1),
204
+ out.stride(0), out.stride(1), 0 if not has_groups else out.stride(2), out.stride(-2), out.stride(-1),
205
+ *((seq_idx.stride(0), seq_idx.stride(1)) if seq_idx is not None else (0, 0)),
206
+ causal,
207
+ dot_dtype,
208
+ HAS_SEQ_IDX=seq_idx is not None,
209
+ )
210
+ return out
211
+
212
+
213
+ def _bmm_chunk_bwd(a, dout, residual=None, out=None):
214
+ """
215
+ Argument:
216
+ a: (batch, seqlen, k) or (batch, seqlen, ngroups, k)
217
+ dout: (batch, nchunks, chunk_size, chunk_size) or (batch, nchunks, ngroups, chunk_size, chunk_size)
218
+ residual: (batch, seqlen, k) or (batch, seqlen, ngroups, k)
219
+ Return:
220
+ out: (batch, seqlen, k) or (batch, seqlen, ngroups, k)
221
+
222
+ If there was seq_idx in the fwd pass, then dout[i, j] for seq_idx[i] != seq_idx[j] should already be
223
+ zeroed out before calling this function.
224
+ """
225
+ # Check constraints.
226
+ has_groups = a.dim() == 4
227
+ if not has_groups:
228
+ batch, seqlen, k = a.shape
229
+ else:
230
+ batch, seqlen, ngroups, k = a.shape
231
+ nchunks, chunk_size = dout.shape[1], dout.shape[-1]
232
+ if a.stride(-1) != 1 and a.stride(-2) != 1:
233
+ a = a.contiguous()
234
+ if dout.stride(-1) != 1 and dout.stride(-2) != 1:
235
+ dout = dout.contiguous()
236
+ if residual is not None:
237
+ assert residual.shape == (batch, seqlen, k) if not has_groups else (batch, seqlen, ngroups, k)
238
+ if residual.stride(-1) != 1 and residual.stride(1) != 1:
239
+ residual = residual.contiguous()
240
+ # Allocates output.
241
+ if out is not None:
242
+ assert out.shape == a.shape
243
+ assert out.stride(-1) == 1 or out.stride(1) == 1
244
+ else:
245
+ out = torch.empty_like(a)
246
+ dot_dtype = (tl.bfloat16 if a.dtype == torch.bfloat16 or dout.dtype == torch.bfloat16 else
247
+ (tl.float16 if a.dtype == torch.float16 or dout.dtype == torch.float16 else tl.float32))
248
+ grid = lambda META: (triton.cdiv(chunk_size, META['BLOCK_SIZE_M']) * triton.cdiv(k, META['BLOCK_SIZE_N']), batch,
249
+ nchunks if not has_groups else nchunks * ngroups)
250
+ residual_strides = ((residual.stride(0), residual.stride(1), 0 if not has_groups else residual.stride(2),
251
+ residual.stride(-1))
252
+ if residual is not None else (0, 0, 0, 0))
253
+ with torch.cuda.device(a.device.index):
254
+ _bmm_chunk_bwd_kernel[grid](
255
+ a, dout, out, residual,
256
+ seqlen, chunk_size, k, ngroups if has_groups else 1,
257
+ a.stride(0), a.stride(1), 0 if not has_groups else a.stride(2), a.stride(-1),
258
+ dout.stride(0), dout.stride(1), 0 if not has_groups else dout.stride(2), dout.stride(-2), dout.stride(-1),
259
+ out.stride(0), out.stride(1), 0 if not has_groups else out.stride(2), out.stride(-1),
260
+ residual_strides[0], residual_strides[1], residual_strides[2], residual_strides[3],
261
+ dot_dtype,
262
+ HAS_RESIDUAL=residual is not None,
263
+ )
264
+ return out
mamba_ssm/ops/triton/ssd_chunk_scan.py ADDED
The diff for this file is too large to render. See raw diff
 
mamba_ssm/ops/triton/ssd_chunk_state.py ADDED
@@ -0,0 +1,1122 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2024, Tri Dao, Albert Gu.
2
+
3
+ """We want triton==2.1.0 or 2.2.0 for this
4
+ """
5
+
6
+ import math
7
+ import torch
8
+ import torch.nn.functional as F
9
+
10
+ import triton
11
+ import triton.language as tl
12
+
13
+ from einops import rearrange, repeat
14
+
15
+ from mamba_ssm.ops.triton.softplus import softplus
16
+ from mamba_ssm.utils.determinism import (
17
+ alloc_tile_workspace,
18
+ finalize_tile_workspace,
19
+ use_deterministic_mode,
20
+ autotune_configs,
21
+ )
22
+
23
+
24
+ def init_to_zero(names):
25
+ return lambda nargs: [nargs[name].zero_() for name in names if nargs[name] is not None]
26
+
27
+ @triton.autotune(
28
+ configs=autotune_configs([
29
+ triton.Config({'BLOCK_SIZE_H': 1}),
30
+ triton.Config({'BLOCK_SIZE_H': 2}),
31
+ triton.Config({'BLOCK_SIZE_H': 4}),
32
+ triton.Config({'BLOCK_SIZE_H': 8}),
33
+ triton.Config({'BLOCK_SIZE_H': 16}),
34
+ triton.Config({'BLOCK_SIZE_H': 32}),
35
+ triton.Config({'BLOCK_SIZE_H': 64}),
36
+ ]),
37
+ key=['chunk_size', 'nheads'],
38
+ )
39
+ @triton.jit
40
+ def _chunk_cumsum_fwd_kernel(
41
+ # Pointers to matrices
42
+ dt_ptr, A_ptr, dt_bias_ptr, dt_out_ptr, dA_cumsum_ptr,
43
+ # Matrix dimension
44
+ batch, seqlen, nheads, chunk_size,
45
+ dt_min, dt_max,
46
+ # Strides
47
+ stride_dt_batch, stride_dt_seqlen, stride_dt_head,
48
+ stride_A_head,
49
+ stride_dt_bias_head,
50
+ stride_dt_out_batch, stride_dt_out_chunk, stride_dt_out_head, stride_dt_out_csize,
51
+ stride_dA_cs_batch, stride_dA_cs_chunk, stride_dA_cs_head, stride_dA_cs_csize,
52
+ # Meta-parameters
53
+ DT_SOFTPLUS: tl.constexpr,
54
+ HAS_DT_BIAS: tl.constexpr,
55
+ BLOCK_SIZE_H: tl.constexpr, BLOCK_SIZE_CHUNK: tl.constexpr,
56
+ ):
57
+ pid_b = tl.program_id(axis=0)
58
+ pid_c = tl.program_id(axis=1)
59
+ pid_h = tl.program_id(axis=2)
60
+ dt_ptr += pid_b * stride_dt_batch + pid_c * chunk_size * stride_dt_seqlen
61
+ dt_out_ptr += pid_b * stride_dt_out_batch + pid_c * stride_dt_out_chunk
62
+ dA_cumsum_ptr += pid_b * stride_dA_cs_batch + pid_c * stride_dA_cs_chunk
63
+
64
+ offs_h = pid_h * BLOCK_SIZE_H + tl.arange(0, BLOCK_SIZE_H)
65
+ offs_c = tl.arange(0, BLOCK_SIZE_CHUNK)
66
+ dt_ptrs = dt_ptr + (offs_h[:, None] * stride_dt_head + offs_c[None, :] * stride_dt_seqlen)
67
+ A_ptrs = A_ptr + offs_h * stride_A_head
68
+ dt_out_ptrs = dt_out_ptr + (offs_h[:, None] * stride_dt_out_head + offs_c[None, :] * stride_dt_out_csize)
69
+ dA_cs_ptrs = dA_cumsum_ptr + (offs_h[:, None] * stride_dA_cs_head + offs_c[None, :] * stride_dA_cs_csize)
70
+ chunk_size_limit = min(chunk_size, seqlen - pid_c * chunk_size)
71
+
72
+ dt = tl.load(dt_ptrs, mask=(offs_h[:, None] < nheads) & (offs_c[None, :] < chunk_size_limit), other=0.0).to(tl.float32)
73
+ if HAS_DT_BIAS:
74
+ dt_bias = tl.load(dt_bias_ptr + offs_h * stride_dt_bias_head, mask=offs_h < nheads, other=0.0).to(tl.float32)
75
+ dt += dt_bias[:, None]
76
+ if DT_SOFTPLUS:
77
+ dt = tl.where(dt <= 20.0, softplus(dt), dt)
78
+ # As of Triton 2.2.0, tl.clamp is not available yet
79
+ # dt = tl.clamp(dt, dt_min, dt_max)
80
+ dt = tl.minimum(tl.maximum(dt, dt_min), dt_max)
81
+ dt = tl.where((offs_h[:, None] < nheads) & (offs_c[None, :] < chunk_size_limit), dt, 0.0)
82
+ tl.store(dt_out_ptrs, dt, mask=(offs_h[:, None] < nheads) & (offs_c[None, :] < chunk_size))
83
+ A = tl.load(A_ptrs, mask=offs_h < nheads, other=0.0).to(tl.float32)
84
+ dA = dt * A[:, None]
85
+ dA_cs = tl.cumsum(dA, axis=1)
86
+ tl.store(dA_cs_ptrs, dA_cs, mask=(offs_h[:, None] < nheads) & (offs_c[None, :] < chunk_size))
87
+
88
+
89
+ @triton.autotune(
90
+ configs=autotune_configs([
91
+ triton.Config({'BLOCK_SIZE_H': 1}, pre_hook=init_to_zero(["dA_ptr", "ddt_bias_ptr"])),
92
+ triton.Config({'BLOCK_SIZE_H': 2}, pre_hook=init_to_zero(["dA_ptr", "ddt_bias_ptr"])),
93
+ triton.Config({'BLOCK_SIZE_H': 4}, pre_hook=init_to_zero(["dA_ptr", "ddt_bias_ptr"])),
94
+ triton.Config({'BLOCK_SIZE_H': 8}, pre_hook=init_to_zero(["dA_ptr", "ddt_bias_ptr"])),
95
+ triton.Config({'BLOCK_SIZE_H': 16}, pre_hook=init_to_zero(["dA_ptr", "ddt_bias_ptr"])),
96
+ triton.Config({'BLOCK_SIZE_H': 32}, pre_hook=init_to_zero(["dA_ptr", "ddt_bias_ptr"])),
97
+ triton.Config({'BLOCK_SIZE_H': 64}, pre_hook=init_to_zero(["dA_ptr", "ddt_bias_ptr"])),
98
+ ]),
99
+ key=['chunk_size', 'nheads'],
100
+ )
101
+ @triton.jit
102
+ def _chunk_cumsum_bwd_kernel(
103
+ # Pointers to matrices
104
+ ddA_ptr, ddt_out_ptr, dt_ptr, A_ptr, dt_bias_ptr,
105
+ ddt_ptr, dA_ptr, ddt_bias_ptr,
106
+ # Matrix dimensions
107
+ batch, seqlen, nheads, chunk_size,
108
+ dt_min, dt_max,
109
+ # Strides
110
+ stride_ddA_batch, stride_ddA_chunk, stride_ddA_head, stride_ddA_csize,
111
+ stride_ddt_out_batch, stride_ddt_out_chunk, stride_ddt_out_head, stride_ddt_out_csize,
112
+ stride_dt_batch, stride_dt_seqlen, stride_dt_head,
113
+ stride_A_head,
114
+ stride_dt_bias_head,
115
+ stride_ddt_batch, stride_ddt_seqlen, stride_ddt_head,
116
+ stride_dA_batch, stride_dA_chunk, stride_dA_head,
117
+ stride_ddt_bias_batch, stride_ddt_bias_chunk, stride_ddt_bias_head,
118
+ # Meta-parameters
119
+ DT_SOFTPLUS: tl.constexpr,
120
+ HAS_DT_BIAS: tl.constexpr,
121
+ BLOCK_SIZE_H: tl.constexpr, BLOCK_SIZE_CHUNK: tl.constexpr,
122
+ DETERMINISTIC_REDUCTION: tl.constexpr,
123
+ ):
124
+ pid_b = tl.program_id(axis=0)
125
+ pid_c = tl.program_id(axis=1)
126
+ pid_h = tl.program_id(axis=2)
127
+ ddt_out_ptr += pid_b * stride_ddt_out_batch + pid_c * stride_ddt_out_chunk
128
+ ddA_ptr += pid_b * stride_ddA_batch + pid_c * stride_ddA_chunk
129
+ dt_ptr += pid_b * stride_dt_batch + pid_c * chunk_size * stride_dt_seqlen
130
+ ddt_ptr += pid_b * stride_ddt_batch + pid_c * chunk_size * stride_ddt_seqlen
131
+
132
+ offs_h = pid_h * BLOCK_SIZE_H + tl.arange(0, BLOCK_SIZE_H)
133
+ offs_c = tl.arange(0, BLOCK_SIZE_CHUNK)
134
+ ddt_out_ptrs = ddt_out_ptr + (offs_h[:, None] * stride_ddt_out_head + offs_c[None, :] * stride_ddt_out_csize)
135
+ ddA_ptrs = ddA_ptr + (offs_h[:, None] * stride_ddA_head + offs_c[None, :] * stride_ddA_csize)
136
+ dt_ptrs = dt_ptr + (offs_h[:, None] * stride_dt_head + offs_c[None, :] * stride_dt_seqlen)
137
+ ddt_ptrs = ddt_ptr + (offs_h[:, None] * stride_ddt_head + offs_c[None, :] * stride_ddt_seqlen)
138
+ A_ptrs = A_ptr + offs_h * stride_A_head
139
+ chunk_size_limit = min(chunk_size, seqlen - pid_c * chunk_size)
140
+
141
+ ddA = tl.load(ddA_ptrs, mask=(offs_h[:, None] < nheads) & (offs_c[None, :] < chunk_size_limit), other=0.0).to(tl.float32)
142
+ ddt_out = tl.load(ddt_out_ptrs, mask=(offs_h[:, None] < nheads) & (offs_c[None, :] < chunk_size_limit), other=0.0).to(tl.float32)
143
+ A = tl.load(A_ptrs, mask=offs_h < nheads, other=0.0).to(tl.float32)
144
+ ddt = ddA * A[:, None] + ddt_out
145
+ dt = tl.load(dt_ptrs, mask=(offs_h[:, None] < nheads) & (offs_c[None, :] < chunk_size_limit), other=0.0).to(tl.float32)
146
+ if HAS_DT_BIAS:
147
+ dt_bias = tl.load(dt_bias_ptr + offs_h * stride_dt_bias_head, mask=offs_h < nheads, other=0.0).to(tl.float32)
148
+ dt += dt_bias[:, None]
149
+ if DT_SOFTPLUS:
150
+ dt_presoftplus = dt
151
+ dt = tl.where(dt <= 20.0, softplus(dt), dt)
152
+ clamp_mask = (dt < dt_min) | (dt > dt_max)
153
+ # As of Triton 2.2.0, tl.clamp is not available yet
154
+ # dt = tl.clamp(dt, dt_min, dt_max)
155
+ dt = tl.minimum(tl.maximum(dt, dt_min), dt_max)
156
+ dt = tl.where((offs_h[:, None] < nheads) & (offs_c[None, :] < chunk_size_limit), dt, 0.0)
157
+ ddt = tl.where((offs_h[:, None] < nheads) & (offs_c[None, :] < chunk_size_limit), ddt, 0.0)
158
+ ddt = tl.where(clamp_mask, 0.0, ddt)
159
+ if DT_SOFTPLUS:
160
+ ddt = tl.where(dt_presoftplus <= 20.0, ddt * tl.sigmoid(dt_presoftplus), ddt)
161
+ tl.store(ddt_ptrs, ddt, mask=(offs_h[:, None] < nheads) & (offs_c[None, :] < chunk_size_limit))
162
+ dA = tl.sum(ddA * dt, axis=1)
163
+ dA_ptr += pid_b * stride_dA_batch + pid_c * stride_dA_chunk
164
+ if DETERMINISTIC_REDUCTION:
165
+ tl.store(dA_ptr + offs_h * stride_dA_head, dA, mask=offs_h < nheads)
166
+ else:
167
+ tl.atomic_add(dA_ptr + offs_h * stride_dA_head, dA, mask=offs_h < nheads)
168
+ if HAS_DT_BIAS:
169
+ ddt_bias = tl.sum(ddt, axis=1)
170
+ ddt_bias_ptr += pid_b * stride_ddt_bias_batch + pid_c * stride_ddt_bias_chunk
171
+ if DETERMINISTIC_REDUCTION:
172
+ tl.store(ddt_bias_ptr + offs_h * stride_ddt_bias_head, ddt_bias, mask=offs_h < nheads)
173
+ else:
174
+ tl.atomic_add(ddt_bias_ptr + offs_h * stride_ddt_bias_head, ddt_bias, mask=offs_h < nheads)
175
+
176
+
177
+ @triton.autotune(
178
+ configs=autotune_configs([
179
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64}, num_stages=3, num_warps=8),
180
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
181
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
182
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
183
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
184
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
185
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=2),
186
+ triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=2),
187
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=2),
188
+ ]),
189
+ key=['hdim', 'dstate', 'chunk_size'],
190
+ )
191
+ @triton.jit
192
+ def _chunk_state_fwd_kernel(
193
+ # Pointers to matrices
194
+ x_ptr, b_ptr, states_ptr, dt_ptr, dA_cumsum_ptr, seq_idx_ptr,
195
+ # Matrix dimensions
196
+ hdim, dstate, chunk_size,
197
+ batch, seqlen, nheads_ngroups_ratio,
198
+ # Strides
199
+ stride_x_batch, stride_x_seqlen, stride_x_head, stride_x_hdim,
200
+ stride_b_batch, stride_b_seqlen, stride_b_head, stride_b_dstate,
201
+ stride_states_batch, stride_states_chunk, stride_states_head, stride_states_hdim, stride_states_dstate,
202
+ stride_dt_batch, stride_dt_chunk, stride_dt_head, stride_dt_csize,
203
+ stride_dA_cs_batch, stride_dA_cs_chunk, stride_dA_cs_head, stride_dA_cs_csize,
204
+ stride_seq_idx_batch, stride_seq_idx_seqlen,
205
+ # Meta-parameters
206
+ HAS_SEQ_IDX: tl.constexpr,
207
+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
208
+ ):
209
+ pid_bc = tl.program_id(axis=1)
210
+ pid_c = pid_bc // batch
211
+ pid_b = pid_bc - pid_c * batch
212
+ pid_h = tl.program_id(axis=2)
213
+ num_pid_n = tl.cdiv(dstate, BLOCK_SIZE_N)
214
+ pid_m = tl.program_id(axis=0) // num_pid_n
215
+ pid_n = tl.program_id(axis=0) % num_pid_n
216
+ b_ptr += pid_b * stride_b_batch + pid_c * chunk_size * stride_b_seqlen + (pid_h // nheads_ngroups_ratio) * stride_b_head
217
+ x_ptr += pid_b * stride_x_batch + pid_c * chunk_size * stride_x_seqlen + pid_h * stride_x_head
218
+ dt_ptr += pid_b * stride_dt_batch + pid_c * stride_dt_chunk + pid_h * stride_dt_head
219
+ dA_cumsum_ptr += pid_b * stride_dA_cs_batch + pid_c * stride_dA_cs_chunk + pid_h * stride_dA_cs_head
220
+ if HAS_SEQ_IDX:
221
+ seq_idx_ptr += pid_b * stride_seq_idx_batch + pid_c * chunk_size * stride_seq_idx_seqlen
222
+
223
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
224
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
225
+ offs_k = tl.arange(0, BLOCK_SIZE_K)
226
+ x_ptrs = x_ptr + (offs_m[:, None] * stride_x_hdim + offs_k[None, :] * stride_x_seqlen)
227
+ b_ptrs = b_ptr + (offs_n[None, :] * stride_b_dstate + offs_k[:, None] * stride_b_seqlen)
228
+ dt_ptrs = dt_ptr + offs_k * stride_dt_csize
229
+ dA_cs_last = tl.load(dA_cumsum_ptr + (chunk_size - 1) * stride_dA_cs_csize).to(tl.float32)
230
+ dA_cumsum_ptrs = dA_cumsum_ptr + offs_k * stride_dA_cs_csize
231
+ if HAS_SEQ_IDX:
232
+ seq_idx_ptrs = seq_idx_ptr + offs_k * stride_seq_idx_seqlen
233
+
234
+ chunk_size_limit = min(chunk_size, seqlen - pid_c * chunk_size)
235
+ if HAS_SEQ_IDX:
236
+ seq_idx_last = tl.load(seq_idx_ptr + (chunk_size_limit - 1) * stride_seq_idx_seqlen)
237
+
238
+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
239
+ for k in range(0, chunk_size_limit, BLOCK_SIZE_K):
240
+ x = tl.load(x_ptrs, mask=(offs_m[:, None] < hdim) & (offs_k[None, :] < chunk_size_limit - k), other=0.0)
241
+ b = tl.load(b_ptrs, mask=(offs_k[:, None] < chunk_size_limit - k) & (offs_n[None, :] < dstate), other=0.0).to(tl.float32)
242
+ dA_cs_k = tl.load(dA_cumsum_ptrs, mask=offs_k < chunk_size_limit - k, other=0.0).to(tl.float32)
243
+ if HAS_SEQ_IDX:
244
+ seq_idx_k = tl.load(seq_idx_ptrs, mask=offs_k < chunk_size_limit - k, other=-1)
245
+ dt_k = tl.load(dt_ptrs, mask=offs_k < chunk_size_limit - k, other=0.0).to(tl.float32)
246
+ if not HAS_SEQ_IDX:
247
+ # scale = tl.exp((dA_cs_last - dA_cs_k)) * dt_k
248
+ scale = tl.exp(tl.minimum((dA_cs_last - dA_cs_k), 0.0)) * dt_k
249
+ else:
250
+ # scale = tl.where(seq_idx_k == seq_idx_last, tl.exp((dA_cs_last - dA_cs_k)) * dt_k, 0.0)
251
+ scale = tl.where((seq_idx_last >= 0) & (seq_idx_k == seq_idx_last), tl.exp(tl.minimum((dA_cs_last - dA_cs_k), 0.0)) * dt_k, 0.0)
252
+ b *= scale[:, None]
253
+ b = b.to(x_ptr.dtype.element_ty)
254
+ acc += tl.dot(x, b)
255
+ x_ptrs += BLOCK_SIZE_K * stride_x_seqlen
256
+ b_ptrs += BLOCK_SIZE_K * stride_b_seqlen
257
+ dt_ptrs += BLOCK_SIZE_K * stride_dt_csize
258
+ dA_cumsum_ptrs += BLOCK_SIZE_K * stride_dA_cs_csize
259
+ if HAS_SEQ_IDX:
260
+ seq_idx_ptrs += BLOCK_SIZE_K * stride_seq_idx_seqlen
261
+ states = acc.to(states_ptr.dtype.element_ty)
262
+
263
+ states_ptr += pid_b * stride_states_batch + pid_c * stride_states_chunk + pid_h * stride_states_head
264
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
265
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
266
+ states_ptrs = states_ptr + (offs_m[:, None] * stride_states_hdim + offs_n[None, :] * stride_states_dstate)
267
+ c_mask = (offs_m[:, None] < hdim) & (offs_n[None, :] < dstate)
268
+ tl.store(states_ptrs, states, mask=c_mask)
269
+
270
+
271
+ @triton.autotune(
272
+ configs=autotune_configs([
273
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64}, num_stages=3, num_warps=8, pre_hook=init_to_zero(["ddt_ptr", "ddA_cumsum_ptr"])),
274
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "ddA_cumsum_ptr"])),
275
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "ddA_cumsum_ptr"])),
276
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "ddA_cumsum_ptr"])),
277
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "ddA_cumsum_ptr"])),
278
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "ddA_cumsum_ptr"])),
279
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "ddA_cumsum_ptr"])),
280
+ triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "ddA_cumsum_ptr"])),
281
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "ddA_cumsum_ptr"])),
282
+ ]),
283
+ key=['chunk_size', 'hdim', 'dstate'],
284
+ )
285
+ @triton.jit
286
+ def _chunk_state_bwd_dx_kernel(
287
+ # Pointers to matrices
288
+ x_ptr, b_ptr, dstates_ptr, dt_ptr, dA_cumsum_ptr,
289
+ dx_ptr, ddt_ptr, ddA_cumsum_ptr,
290
+ # Matrix dimensions
291
+ chunk_size, hdim, dstate,
292
+ batch, seqlen, nheads_ngroups_ratio,
293
+ # Strides
294
+ stride_x_batch, stride_x_seqlen, stride_x_head, stride_x_hdim,
295
+ stride_b_batch, stride_b_seqlen, stride_b_head, stride_b_dstate,
296
+ stride_dstates_batch, stride_dstates_chunk, stride_states_head, stride_states_hdim, stride_states_dstate,
297
+ stride_dt_batch, stride_dt_chunk, stride_dt_head, stride_dt_csize,
298
+ stride_dA_cs_batch, stride_dA_cs_chunk, stride_dA_cs_head, stride_dA_cs_csize,
299
+ stride_dx_batch, stride_dx_seqlen, stride_dx_head, stride_dx_hdim,
300
+ stride_ddt_batch, stride_ddt_chunk, stride_ddt_head, stride_ddt_csize, stride_ddt_tile,
301
+ stride_ddA_cs_batch, stride_ddA_cs_chunk, stride_ddA_cs_head, stride_ddA_cs_csize, stride_ddA_tile,
302
+ # Meta-parameters
303
+ DETERMINISTIC_REDUCTION: tl.constexpr,
304
+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
305
+ BLOCK_SIZE_DSTATE: tl.constexpr,
306
+ ):
307
+ pid_bc = tl.program_id(axis=1)
308
+ pid_c = pid_bc // batch
309
+ pid_b = pid_bc - pid_c * batch
310
+ pid_h = tl.program_id(axis=2)
311
+ num_pid_n = tl.cdiv(hdim, BLOCK_SIZE_N)
312
+ pid_m = tl.program_id(axis=0) // num_pid_n
313
+ pid_n = tl.program_id(axis=0) % num_pid_n
314
+ x_ptr += pid_b * stride_x_batch + pid_c * chunk_size * stride_x_seqlen + pid_h * stride_x_head
315
+ b_ptr += pid_b * stride_b_batch + pid_c * chunk_size * stride_b_seqlen + (pid_h // nheads_ngroups_ratio) * stride_b_head
316
+ dstates_ptr += pid_b * stride_dstates_batch + pid_c * stride_dstates_chunk + pid_h * stride_states_head
317
+ dt_ptr += pid_b * stride_dt_batch + pid_c * stride_dt_chunk + pid_h * stride_dt_head
318
+ ddt_ptr += pid_b * stride_ddt_batch + pid_c * stride_ddt_chunk + pid_h * stride_ddt_head + pid_n * stride_ddt_tile
319
+ ddA_cumsum_ptr += pid_b * stride_ddA_cs_batch + pid_c * stride_ddA_cs_chunk + pid_h * stride_ddA_cs_head + pid_n * stride_ddA_tile
320
+ dA_cumsum_ptr += pid_b * stride_dA_cs_batch + pid_c * stride_dA_cs_chunk + pid_h * stride_dA_cs_head
321
+
322
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
323
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
324
+
325
+ chunk_size_limit = min(chunk_size, seqlen - pid_c * chunk_size)
326
+ # Faster to just do 1 iteration with larger BLOCK_SIZE_K, up to block size 128
327
+ offs_k = tl.arange(0, BLOCK_SIZE_DSTATE if BLOCK_SIZE_DSTATE <= 128 else BLOCK_SIZE_K)
328
+ b_ptrs = b_ptr + (offs_m[:, None] * stride_b_seqlen + offs_k[None, :] * stride_b_dstate)
329
+ dstates_ptrs = dstates_ptr + (offs_n[None, :] * stride_states_hdim + offs_k[:, None] * stride_states_dstate)
330
+ if BLOCK_SIZE_DSTATE <= 128:
331
+ b = tl.load(b_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_k[None, :] < dstate), other=0.0)
332
+ dstates = tl.load(dstates_ptrs, mask=(offs_k[:, None] < dstate) & (offs_n[None, :] < hdim), other=0.0)
333
+ dstates = dstates.to(b_ptr.dtype.element_ty)
334
+ acc = tl.dot(b, dstates)
335
+ else:
336
+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
337
+ for k in range(0, dstate, BLOCK_SIZE_K):
338
+ b = tl.load(b_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_k[None, :] < dstate - k), other=0.0)
339
+ dstates = tl.load(dstates_ptrs, mask=(offs_k[:, None] < dstate - k) & (offs_n[None, :] < hdim), other=0.0)
340
+ dstates = dstates.to(b_ptr.dtype.element_ty)
341
+ acc += tl.dot(b, dstates)
342
+ b_ptrs += BLOCK_SIZE_K * stride_b_dstate
343
+ dstates_ptrs += BLOCK_SIZE_K * stride_states_dstate
344
+
345
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
346
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
347
+
348
+ dA_cs_last = tl.load(dA_cumsum_ptr + (chunk_size - 1) * stride_dA_cs_csize).to(tl.float32)
349
+ dt_ptrs = dt_ptr + offs_m * stride_dt_csize
350
+ dA_cumsum_ptrs = dA_cumsum_ptr + offs_m * stride_dA_cs_csize
351
+ dA_cs_m = tl.load(dA_cumsum_ptrs, mask=offs_m < chunk_size, other=0.0).to(tl.float32)
352
+ dt_m = tl.load(dt_ptrs, mask=offs_m < chunk_size, other=0.0).to(tl.float32)
353
+ # acc *= tl.exp(dA_cs_last - dA_cs_m)[:, None]
354
+ acc *= tl.exp(tl.minimum((dA_cs_last - dA_cs_m), 0.0))[:, None]
355
+
356
+ x_ptrs = x_ptr + (offs_m[:, None] * stride_x_seqlen + offs_n[None, :] * stride_x_hdim)
357
+ x = tl.load(x_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_n[None, :] < hdim), other=0.0).to(tl.float32)
358
+ ddt = tl.sum(acc * x, axis=1)
359
+ ddt_ptrs = ddt_ptr + offs_m * stride_ddt_csize
360
+ if DETERMINISTIC_REDUCTION:
361
+ tl.store(ddt_ptrs, ddt, mask=offs_m < chunk_size)
362
+ else:
363
+ tl.atomic_add(ddt_ptrs, ddt, mask=offs_m < chunk_size)
364
+ ddA_cs = -(ddt * dt_m)
365
+ ddA_cs_last = -tl.sum(ddA_cs)
366
+ ddA_cumsum_ptrs = ddA_cumsum_ptr + offs_m * stride_ddA_cs_csize
367
+ if DETERMINISTIC_REDUCTION:
368
+ tl.store(ddA_cumsum_ptrs, ddA_cs, mask=offs_m < chunk_size)
369
+ else:
370
+ tl.atomic_add(ddA_cumsum_ptrs, ddA_cs, mask=offs_m < chunk_size)
371
+ tl.atomic_add(ddA_cumsum_ptr + (chunk_size - 1) * stride_ddA_cs_csize, ddA_cs_last)
372
+
373
+ dx = (acc * dt_m[:, None]).to(dx_ptr.dtype.element_ty)
374
+ dx_ptr += pid_b * stride_dx_batch + pid_c * chunk_size * stride_dx_seqlen + pid_h * stride_dx_head
375
+ dx_ptrs = dx_ptr + (offs_m[:, None] * stride_dx_seqlen + offs_n[None, :] * stride_dx_hdim)
376
+ tl.store(dx_ptrs, dx, mask=(offs_m[:, None] < chunk_size_limit) & (offs_n[None, :] < hdim))
377
+
378
+
379
+ _CHUNK_STATE_BWD_DX_MIN_BLOCK_N = min(
380
+ cfg.kwargs['BLOCK_SIZE_N'] for cfg in _chunk_state_bwd_dx_kernel.configs
381
+ )
382
+
383
+
384
+ @triton.autotune(
385
+ configs=autotune_configs([
386
+ triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 128}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
387
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
388
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
389
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
390
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
391
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
392
+ triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
393
+ triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 32}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
394
+ ]),
395
+ key=['chunk_size', 'dstate', 'hdim'],
396
+ )
397
+ @triton.jit
398
+ def _chunk_state_bwd_db_kernel(
399
+ # Pointers to matrices
400
+ x_ptr, dstates_ptr, b_ptr, dt_ptr, dA_cumsum_ptr, seq_idx_ptr,
401
+ db_ptr, ddA_cumsum_ptr,
402
+ # Matrix dimensions
403
+ chunk_size, dstate, hdim,
404
+ batch, seqlen, nheads, nheads_per_program, ngroups,
405
+ # Strides
406
+ stride_x_batch, stride_x_seqlen, stride_x_head, stride_x_hdim,
407
+ stride_dstates_batch, stride_dstates_chunk, stride_states_head, stride_states_hdim, stride_states_dstate,
408
+ stride_b_batch, stride_b_seqlen, stride_b_head, stride_b_dstate,
409
+ stride_dt_batch, stride_dt_chunk, stride_dt_head, stride_dt_csize,
410
+ stride_dA_cs_batch, stride_dA_cs_chunk, stride_dA_cs_head, stride_dA_cs_csize,
411
+ stride_seq_idx_batch, stride_seq_idx_seqlen,
412
+ stride_db_batch, stride_db_seqlen, stride_db_split, stride_db_group, stride_db_dstate,
413
+ stride_ddA_cs_batch, stride_ddA_cs_chunk, stride_ddA_cs_head, stride_ddA_cs_csize, stride_ddA_tile,
414
+ # Meta-parameters
415
+ HAS_DDA_CS: tl.constexpr,
416
+ HAS_SEQ_IDX: tl.constexpr,
417
+ DETERMINISTIC_REDUCTION: tl.constexpr,
418
+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
419
+ ):
420
+ pid_bc = tl.program_id(axis=1)
421
+ pid_c = pid_bc // batch
422
+ pid_b = pid_bc - pid_c * batch
423
+ pid_sg = tl.program_id(axis=2)
424
+ pid_s = pid_sg // ngroups
425
+ pid_g = pid_sg - pid_s * ngroups
426
+ num_pid_n = tl.cdiv(dstate, BLOCK_SIZE_N)
427
+ pid_m = tl.program_id(axis=0) // num_pid_n
428
+ pid_n = tl.program_id(axis=0) % num_pid_n
429
+ x_ptr += pid_b * stride_x_batch + pid_c * chunk_size * stride_x_seqlen + (pid_g * (nheads // ngroups) + pid_s * nheads_per_program) * stride_x_head
430
+ db_ptr += pid_b * stride_db_batch + pid_c * chunk_size * stride_db_seqlen + pid_g * stride_db_group + pid_s * stride_db_split
431
+ dstates_ptr += pid_b * stride_dstates_batch + pid_c * stride_dstates_chunk + (pid_g * (nheads // ngroups) + pid_s * nheads_per_program) * stride_states_head
432
+ dt_ptr += pid_b * stride_dt_batch + pid_c * stride_dt_chunk + (pid_g * (nheads // ngroups) + pid_s * nheads_per_program) * stride_dt_head
433
+ dA_cumsum_ptr += pid_b * stride_dA_cs_batch + pid_c * stride_dA_cs_chunk + (pid_g * (nheads // ngroups) + pid_s * nheads_per_program) * stride_dA_cs_head
434
+ if HAS_DDA_CS:
435
+ b_ptr += pid_b * stride_b_batch + pid_c * chunk_size * stride_b_seqlen + pid_g * stride_b_head
436
+ ddA_cumsum_ptr += pid_b * stride_ddA_cs_batch + pid_c * stride_ddA_cs_chunk + (pid_g * (nheads // ngroups) + pid_s * nheads_per_program) * stride_ddA_cs_head + pid_n * stride_ddA_tile
437
+ if HAS_SEQ_IDX:
438
+ seq_idx_ptr += pid_b * stride_seq_idx_batch + pid_c * chunk_size * stride_seq_idx_seqlen
439
+
440
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
441
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
442
+ offs_k = tl.arange(0, BLOCK_SIZE_K)
443
+ x_ptrs = x_ptr + (offs_m[:, None] * stride_x_seqlen + offs_k[None, :] * stride_x_hdim)
444
+ dstates_ptrs = dstates_ptr + (offs_n[None, :] * stride_states_dstate + offs_k[:, None] * stride_states_hdim)
445
+ dt_ptrs = dt_ptr + offs_m * stride_dt_csize
446
+ dA_cumsum_ptrs = dA_cumsum_ptr + offs_m * stride_dA_cs_csize
447
+ if HAS_DDA_CS:
448
+ b_ptrs = b_ptr + (offs_m[:, None] * stride_b_seqlen + offs_n[None, :] * stride_b_dstate)
449
+ ddA_cumsum_ptrs = ddA_cumsum_ptr + offs_m * stride_ddA_cs_csize
450
+
451
+ chunk_size_limit = min(chunk_size, seqlen - pid_c * chunk_size)
452
+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
453
+ if HAS_DDA_CS:
454
+ b = tl.load(b_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_n[None, :] < dstate), other=0.0).to(tl.float32)
455
+ if HAS_SEQ_IDX:
456
+ seq_idx_m = tl.load(seq_idx_ptr + offs_m * stride_seq_idx_seqlen, mask=offs_m < chunk_size_limit, other=-1)
457
+ seq_idx_last = tl.load(seq_idx_ptr + (chunk_size_limit - 1) * stride_seq_idx_seqlen)
458
+ nheads_iter = min(nheads_per_program, nheads // ngroups - pid_s * nheads_per_program)
459
+ for h in range(nheads_iter):
460
+ x = tl.load(x_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_k[None, :] < hdim), other=0.0)
461
+ dstates = tl.load(dstates_ptrs, mask=(offs_k[:, None] < hdim) & (offs_n[None, :] < dstate), other=0.0)
462
+ dstates = dstates.to(x_ptrs.dtype.element_ty)
463
+ db = tl.dot(x, dstates)
464
+ dA_cs_last = tl.load(dA_cumsum_ptr + (chunk_size - 1) * stride_dA_cs_csize).to(tl.float32)
465
+ dA_cs_m = tl.load(dA_cumsum_ptrs, mask=offs_m < chunk_size, other=0.0).to(tl.float32)
466
+ dt_m = tl.load(dt_ptrs, mask=offs_m < chunk_size, other=0.0).to(tl.float32)
467
+ if not HAS_SEQ_IDX:
468
+ # scale = tl.exp(dA_cs_last - dA_cs_m)
469
+ scale = tl.exp(tl.minimum((dA_cs_last - dA_cs_m), 0.0))
470
+ else:
471
+ # scale = tl.where(seq_idx_m == seq_idx_last, tl.exp(dA_cs_last - dA_cs_m), 0.0)
472
+ scale = tl.where(seq_idx_m == seq_idx_last, tl.exp(tl.minimum((dA_cs_last - dA_cs_m), 0.0)), 0.0)
473
+ db *= (scale * dt_m)[:, None]
474
+ if HAS_DDA_CS:
475
+ # This is the gradient wrt (dA_cs_last - dA_cs_m), i.e. the exclusive reverse cumsum
476
+ ddA_cs = tl.sum(db * b, axis=1)
477
+ if DETERMINISTIC_REDUCTION:
478
+ tl.store(ddA_cumsum_ptrs + stride_ddA_cs_csize, ddA_cs, mask=offs_m < chunk_size - 1)
479
+ else:
480
+ tl.atomic_add(ddA_cumsum_ptrs + stride_ddA_cs_csize, ddA_cs, mask=offs_m < chunk_size - 1)
481
+ acc += db
482
+ x_ptrs += stride_x_head
483
+ dstates_ptrs += stride_states_head
484
+ dt_ptrs += stride_dt_head
485
+ dA_cumsum_ptr += stride_dA_cs_head
486
+ dA_cumsum_ptrs += stride_dA_cs_head
487
+ if HAS_DDA_CS:
488
+ ddA_cumsum_ptrs += stride_ddA_cs_head
489
+
490
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
491
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
492
+ # if HAS_SEQ_IDX:
493
+ # seq_idx_last = tl.load(seq_idx_ptr + (chunk_size_limit - 1) * stride_seq_idx_seqlen)
494
+ # seq_idx_m = tl.load(seq_idx_ptr + offs_m * stride_seq_idx_seqlen, mask=offs_m < chunk_size_limit, other=-1)
495
+ # acc = tl.where(seq_idx_m[:, None] == seq_idx_last, acc, 0.0)
496
+ db_ptrs = db_ptr + (offs_m[:, None] * stride_db_seqlen + offs_n[None, :] * stride_db_dstate)
497
+ tl.store(db_ptrs, acc, mask=(offs_m[:, None] < chunk_size_limit) & (offs_n[None, :] < dstate))
498
+
499
+
500
+ _CHUNK_STATE_BWD_DB_MIN_BLOCK_N = min(
501
+ cfg.kwargs['BLOCK_SIZE_N'] for cfg in _chunk_state_bwd_db_kernel.configs
502
+ )
503
+
504
+
505
+ @triton.autotune(
506
+ configs=autotune_configs([
507
+ # triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64}, num_stages=3, num_warps=8, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
508
+ # triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
509
+ # triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
510
+ # triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
511
+ # triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
512
+ # triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
513
+ # triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
514
+ # triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
515
+ # triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
516
+ triton.Config({'BLOCK_SIZE_N': 16, 'BLOCK_SIZE_K': 32}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
517
+ triton.Config({'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
518
+ triton.Config({'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
519
+ triton.Config({'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=3, num_warps=4, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
520
+ triton.Config({'BLOCK_SIZE_N': 16, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=8, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
521
+ triton.Config({'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=8, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
522
+ triton.Config({'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=8, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
523
+ triton.Config({'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=8, pre_hook=init_to_zero(["ddA_cumsum_ptr"])),
524
+ ]),
525
+ key=['chunk_size', 'hdim', 'dstate'],
526
+ )
527
+ @triton.jit
528
+ def _chunk_state_bwd_ddAcs_stable_kernel(
529
+ # Pointers to matrices
530
+ x_ptr, b_ptr, dstates_ptr, dt_ptr, dA_cumsum_ptr, seq_idx_ptr,
531
+ ddA_cumsum_ptr,
532
+ # Matrix dimensions
533
+ chunk_size, hdim, dstate,
534
+ batch, seqlen, nheads_ngroups_ratio,
535
+ # Strides
536
+ stride_x_batch, stride_x_seqlen, stride_x_head, stride_x_hdim,
537
+ stride_b_batch, stride_b_seqlen, stride_b_head, stride_b_dstate,
538
+ stride_dstates_batch, stride_dstates_chunk, stride_states_head, stride_states_hdim, stride_states_dstate,
539
+ stride_dt_batch, stride_dt_chunk, stride_dt_head, stride_dt_csize,
540
+ stride_dA_cs_batch, stride_dA_cs_chunk, stride_dA_cs_head, stride_dA_cs_csize,
541
+ stride_seq_idx_batch, stride_seq_idx_seqlen,
542
+ stride_ddA_cs_batch, stride_ddA_cs_chunk, stride_ddA_cs_head, stride_ddA_cs_csize, stride_ddA_tile,
543
+ # Meta-parameters
544
+ HAS_SEQ_IDX: tl.constexpr,
545
+ DETERMINISTIC_REDUCTION: tl.constexpr,
546
+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
547
+ BLOCK_SIZE_DSTATE: tl.constexpr,
548
+ ):
549
+ pid_bc = tl.program_id(axis=1)
550
+ pid_c = pid_bc // batch
551
+ pid_b = pid_bc - pid_c * batch
552
+ pid_h = tl.program_id(axis=2)
553
+ num_pid_n = tl.cdiv(hdim, BLOCK_SIZE_N)
554
+ pid_m = tl.program_id(axis=0) // num_pid_n
555
+ pid_n = tl.program_id(axis=0) % num_pid_n
556
+ x_ptr += pid_b * stride_x_batch + pid_c * chunk_size * stride_x_seqlen + pid_h * stride_x_head
557
+ b_ptr += pid_b * stride_b_batch + pid_c * chunk_size * stride_b_seqlen + (pid_h // nheads_ngroups_ratio) * stride_b_head
558
+ dstates_ptr += pid_b * stride_dstates_batch + pid_c * stride_dstates_chunk + pid_h * stride_states_head
559
+ dt_ptr += pid_b * stride_dt_batch + pid_c * stride_dt_chunk + pid_h * stride_dt_head
560
+ ddA_cumsum_ptr += pid_b * stride_ddA_cs_batch + pid_c * stride_ddA_cs_chunk + pid_h * stride_ddA_cs_head + pid_n * stride_ddA_tile
561
+ dA_cumsum_ptr += pid_b * stride_dA_cs_batch + pid_c * stride_dA_cs_chunk + pid_h * stride_dA_cs_head
562
+ if HAS_SEQ_IDX:
563
+ seq_idx_ptr += pid_b * stride_seq_idx_batch + pid_c * chunk_size * stride_seq_idx_seqlen
564
+
565
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
566
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
567
+
568
+ chunk_size_limit = min(chunk_size, seqlen - pid_c * chunk_size)
569
+ # Faster to just do 1 iteration with larger BLOCK_SIZE_K, up to block size 128
570
+ offs_k = tl.arange(0, BLOCK_SIZE_DSTATE if BLOCK_SIZE_DSTATE <= 128 else BLOCK_SIZE_K)
571
+ b_ptrs = b_ptr + (offs_m[:, None] * stride_b_seqlen + offs_k[None, :] * stride_b_dstate)
572
+ dstates_ptrs = dstates_ptr + (offs_n[None, :] * stride_states_hdim + offs_k[:, None] * stride_states_dstate)
573
+ if BLOCK_SIZE_DSTATE <= 128:
574
+ b = tl.load(b_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_k[None, :] < dstate), other=0.0)
575
+ dstates = tl.load(dstates_ptrs, mask=(offs_k[:, None] < dstate) & (offs_n[None, :] < hdim), other=0.0)
576
+ dstates = dstates.to(b_ptr.dtype.element_ty)
577
+ acc = tl.dot(b, dstates)
578
+ else:
579
+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
580
+ for k in range(0, dstate, BLOCK_SIZE_K):
581
+ b = tl.load(b_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_k[None, :] < dstate - k), other=0.0)
582
+ dstates = tl.load(dstates_ptrs, mask=(offs_k[:, None] < dstate - k) & (offs_n[None, :] < hdim), other=0.0)
583
+ dstates = dstates.to(b_ptr.dtype.element_ty)
584
+ acc += tl.dot(b, dstates)
585
+ b_ptrs += BLOCK_SIZE_K * stride_b_dstate
586
+ dstates_ptrs += BLOCK_SIZE_K * stride_states_dstate
587
+
588
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
589
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
590
+
591
+ dA_cs_m = tl.load(dA_cumsum_ptr + offs_m * stride_dA_cs_csize, mask=offs_m < chunk_size, other=0.0).to(tl.float32)
592
+ dA_cs_last = tl.load(dA_cumsum_ptr + (chunk_size - 1) * stride_dA_cs_csize).to(tl.float32)
593
+ if not HAS_SEQ_IDX:
594
+ # scale = tl.exp(dA_cs_last - dA_cs_m)
595
+ scale = tl.exp(tl.minimum((dA_cs_last - dA_cs_m), 0.0))
596
+ else:
597
+ seq_idx_m = tl.load(seq_idx_ptr + offs_m * stride_seq_idx_seqlen, mask=offs_m < chunk_size_limit, other=-1)
598
+ seq_idx_last = tl.load(seq_idx_ptr + (chunk_size_limit - 1) * stride_seq_idx_seqlen)
599
+ # scale = tl.where(seq_idx_m == seq_idx_last, tl.exp(dA_cs_last - dA_cs_m), 0.0)
600
+ scale = tl.where(seq_idx_m == seq_idx_last, tl.exp(tl.minimum((dA_cs_last - dA_cs_m), 0.0)), 0.0)
601
+ acc *= scale[:, None]
602
+
603
+ x_ptrs = x_ptr + (offs_m[:, None] * stride_x_seqlen + offs_n[None, :] * stride_x_hdim)
604
+ x = tl.load(x_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_n[None, :] < hdim), other=0.0).to(tl.float32)
605
+ dt_ptrs = dt_ptr + offs_m * stride_dt_csize
606
+ dt_m = tl.load(dt_ptrs, mask=offs_m < chunk_size, other=0.0).to(tl.float32)
607
+ ddt = tl.sum(acc * x, axis=1)
608
+ # ddA_cs = -(ddt * dt_m)
609
+ # Triton 2.2.0 errors if we have the cumsum here, so we just write it out
610
+ # then call torch.cumsum outside this kernel.
611
+ # ddA_cs = tl.cumsum(ddt * dt_m)
612
+ ddA_cs = ddt * dt_m
613
+ ddA_cumsum_ptrs = ddA_cumsum_ptr + offs_m * stride_ddA_cs_csize
614
+ if DETERMINISTIC_REDUCTION:
615
+ tl.store(ddA_cumsum_ptrs + stride_ddA_cs_csize, ddA_cs, mask=offs_m < chunk_size - 1)
616
+ else:
617
+ tl.atomic_add(ddA_cumsum_ptrs + stride_ddA_cs_csize, ddA_cs, mask=offs_m < chunk_size - 1)
618
+
619
+
620
+ _CHUNK_STATE_BWD_DDACS_MIN_BLOCK_N = min(
621
+ cfg.kwargs['BLOCK_SIZE_N'] for cfg in _chunk_state_bwd_ddAcs_stable_kernel.configs
622
+ )
623
+
624
+
625
+ @triton.autotune(
626
+ configs=autotune_configs([
627
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64}, num_stages=3, num_warps=8),
628
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
629
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
630
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
631
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
632
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4),
633
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=2),
634
+ triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=2),
635
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=2),
636
+ ]),
637
+ key=['hdim', 'dstate', 'chunk_size'],
638
+ )
639
+ @triton.jit
640
+ def _chunk_state_varlen_kernel(
641
+ # Pointers to matrices
642
+ x_ptr, b_ptr, dt_ptr, dA_cumsum_ptr, chunk_states_ptr, cu_seqlens_ptr, states_ptr,
643
+ # Matrix dimensions
644
+ hdim, dstate, chunk_size,
645
+ seqlen, nheads_ngroups_ratio,
646
+ # Strides
647
+ stride_x_seqlen, stride_x_head, stride_x_hdim,
648
+ stride_b_seqlen, stride_b_head, stride_b_dstate,
649
+ stride_dt_chunk, stride_dt_head, stride_dt_csize,
650
+ stride_dA_cs_chunk, stride_dA_cs_head, stride_dA_cs_csize,
651
+ stride_chunk_states_chunk, stride_chunk_states_head, stride_chunk_states_hdim, stride_chunk_states_dstate,
652
+ stride_states_batch, stride_states_head, stride_states_hdim, stride_states_dstate,
653
+ # Meta-parameters
654
+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
655
+ ):
656
+ pid_b = tl.program_id(axis=1)
657
+ pid_h = tl.program_id(axis=2)
658
+ num_pid_n = tl.cdiv(dstate, BLOCK_SIZE_N)
659
+ pid_m = tl.program_id(axis=0) // num_pid_n
660
+ pid_n = tl.program_id(axis=0) % num_pid_n
661
+ end_idx = tl.load(cu_seqlens_ptr + pid_b + 1)
662
+ pid_c = (end_idx - 1) // chunk_size
663
+ b_ptr += pid_c * chunk_size * stride_b_seqlen + (pid_h // nheads_ngroups_ratio) * stride_b_head
664
+ x_ptr += pid_c * chunk_size * stride_x_seqlen + pid_h * stride_x_head
665
+ dt_ptr += pid_c * stride_dt_chunk + pid_h * stride_dt_head
666
+ dA_cumsum_ptr += pid_c * stride_dA_cs_chunk + pid_h * stride_dA_cs_head
667
+ chunk_states_ptr += pid_c * stride_chunk_states_chunk + pid_h * stride_chunk_states_head
668
+
669
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
670
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
671
+ offs_k = tl.arange(0, BLOCK_SIZE_K)
672
+ x_ptrs = x_ptr + (offs_m[:, None] * stride_x_hdim + offs_k[None, :] * stride_x_seqlen)
673
+ b_ptrs = b_ptr + (offs_n[None, :] * stride_b_dstate + offs_k[:, None] * stride_b_seqlen)
674
+ dt_ptrs = dt_ptr + offs_k * stride_dt_csize
675
+ dA_cs_last = tl.load(dA_cumsum_ptr + (end_idx - pid_c * chunk_size - 1) * stride_dA_cs_csize).to(tl.float32)
676
+ dA_cumsum_ptrs = dA_cumsum_ptr + offs_k * stride_dA_cs_csize
677
+
678
+ chunk_size_limit = end_idx - pid_c * chunk_size
679
+ start_idx = tl.load(cu_seqlens_ptr + pid_b)
680
+ start_idx_cur = tl.maximum(start_idx - pid_c * chunk_size, 0)
681
+
682
+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
683
+ for k in range(0, chunk_size_limit, BLOCK_SIZE_K):
684
+ x = tl.load(x_ptrs, mask=(offs_m[:, None] < hdim) & (offs_k[None, :] < chunk_size_limit - k) & (offs_k[None, :] >= start_idx_cur - k), other=0.0)
685
+ b = tl.load(b_ptrs, mask=(offs_k[:, None] < chunk_size_limit - k) & (offs_n[None, :] < dstate) & (offs_k[:, None] >= start_idx_cur - k), other=0.0).to(tl.float32)
686
+ dA_cs_k = tl.load(dA_cumsum_ptrs, mask=offs_k < chunk_size_limit - k, other=0.0).to(tl.float32)
687
+ dt_k = tl.load(dt_ptrs, mask=offs_k < chunk_size_limit - k, other=0.0).to(tl.float32)
688
+ # scale = tl.where((offs_k >= start_idx_cur - k) & (offs_k < chunk_size_limit - k),
689
+ # tl.exp((dA_cs_last - dA_cs_k)) * dt_k, 0.0)
690
+ scale = tl.where((offs_k >= start_idx_cur - k) & (offs_k < chunk_size_limit - k),
691
+ tl.exp(tl.minimum((dA_cs_last - dA_cs_k), 0.0)) * dt_k, 0.0)
692
+ b *= scale[:, None]
693
+ b = b.to(x_ptr.dtype.element_ty)
694
+ acc += tl.dot(x, b)
695
+ x_ptrs += BLOCK_SIZE_K * stride_x_seqlen
696
+ b_ptrs += BLOCK_SIZE_K * stride_b_seqlen
697
+ dt_ptrs += BLOCK_SIZE_K * stride_dt_csize
698
+ dA_cumsum_ptrs += BLOCK_SIZE_K * stride_dA_cs_csize
699
+
700
+ # If the sequence starts after the last chunk idx, we don't need to add the contribution from the last chunk
701
+ if start_idx < pid_c * chunk_size:
702
+ chunk_states_ptrs = chunk_states_ptr + (offs_m[:, None] * stride_chunk_states_hdim + offs_n[None, :] * stride_chunk_states_dstate)
703
+ chunk_states = tl.load(chunk_states_ptrs, mask=(offs_m[:, None] < hdim) & (offs_n[None, :] < dstate), other=0.0).to(tl.float32)
704
+ # scale = tl.where(start_idx < pid_c * chunk_size, tl.exp(dA_cs_last), 0.0)
705
+ scale = tl.exp(dA_cs_last)
706
+ acc += chunk_states * scale
707
+
708
+ states = acc.to(states_ptr.dtype.element_ty)
709
+
710
+ states_ptr += pid_b * stride_states_batch + pid_h * stride_states_head
711
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
712
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
713
+ states_ptrs = states_ptr + (offs_m[:, None] * stride_states_hdim + offs_n[None, :] * stride_states_dstate)
714
+ c_mask = (offs_m[:, None] < hdim) & (offs_n[None, :] < dstate)
715
+ tl.store(states_ptrs, states, mask=c_mask)
716
+
717
+
718
+ def _chunk_cumsum_fwd(dt, A, chunk_size, dt_bias=None, dt_softplus=False, dt_limit=(0.0, float("inf"))):
719
+ batch, seqlen, nheads = dt.shape
720
+ assert A.shape == (nheads,)
721
+ if dt_bias is not None:
722
+ assert dt_bias.shape == (nheads,)
723
+ nchunks = math.ceil(seqlen / chunk_size)
724
+ dt_out = torch.empty(batch, nheads, nchunks, chunk_size, device=dt.device, dtype=torch.float32)
725
+ dA_cumsum = torch.empty(batch, nheads, nchunks, chunk_size, device=dt.device, dtype=torch.float32)
726
+ grid_chunk_cs = lambda META: (batch, nchunks, triton.cdiv(nheads, META['BLOCK_SIZE_H']))
727
+ with torch.cuda.device(dt.device.index):
728
+ _chunk_cumsum_fwd_kernel[grid_chunk_cs](
729
+ dt, A, dt_bias, dt_out, dA_cumsum,
730
+ batch, seqlen, nheads, chunk_size,
731
+ dt_limit[0], dt_limit[1],
732
+ dt.stride(0), dt.stride(1), dt.stride(2),
733
+ A.stride(0),
734
+ dt_bias.stride(0) if dt_bias is not None else 0,
735
+ dt_out.stride(0), dt_out.stride(2), dt_out.stride(1), dt_out.stride(3),
736
+ dA_cumsum.stride(0), dA_cumsum.stride(2), dA_cumsum.stride(1), dA_cumsum.stride(3),
737
+ dt_softplus,
738
+ HAS_DT_BIAS=dt_bias is not None,
739
+ BLOCK_SIZE_CHUNK=triton.next_power_of_2(chunk_size),
740
+ )
741
+ return dA_cumsum, dt_out
742
+
743
+
744
+ def _chunk_cumsum_bwd(ddA, ddt_out, dt, A, dt_bias=None, dt_softplus=False, dt_limit=(0.0, float("inf")), ddt=None):
745
+ batch, seqlen, nheads = dt.shape
746
+ _, _, nchunks, chunk_size = ddA.shape
747
+ assert ddA.shape == (batch, nheads, nchunks, chunk_size)
748
+ assert ddt_out.shape == (batch, nheads, nchunks, chunk_size)
749
+ assert A.shape == (nheads,)
750
+ deterministic = use_deterministic_mode()
751
+ if dt_bias is not None:
752
+ assert dt_bias.shape == (nheads,)
753
+ if deterministic:
754
+ ddt_bias_workspace = torch.zeros(
755
+ batch, nchunks, nheads, device=dt.device, dtype=torch.float32
756
+ )
757
+ ddt_bias = torch.empty_like(dt_bias, dtype=torch.float32)
758
+ stride_ddt_bias_batch = ddt_bias_workspace.stride(0)
759
+ stride_ddt_bias_chunk = ddt_bias_workspace.stride(1)
760
+ else:
761
+ ddt_bias_workspace = ddt_bias = torch.empty_like(
762
+ dt_bias, dtype=torch.float32
763
+ )
764
+ stride_ddt_bias_batch = 0
765
+ stride_ddt_bias_chunk = 0
766
+ else:
767
+ ddt_bias = None
768
+ ddt_bias_workspace = None
769
+ stride_ddt_bias_batch = 0
770
+ stride_ddt_bias_chunk = 0
771
+ if ddt is not None:
772
+ assert ddt.shape == dt.shape
773
+ else:
774
+ ddt = torch.empty_like(dt)
775
+ dA = torch.empty_like(A, dtype=torch.float32)
776
+ if deterministic:
777
+ dA_workspace = torch.zeros(
778
+ batch, nchunks, nheads, device=dt.device, dtype=torch.float32
779
+ )
780
+ stride_dA_batch = dA_workspace.stride(0)
781
+ stride_dA_chunk = dA_workspace.stride(1)
782
+ else:
783
+ dA_workspace = dA
784
+ stride_dA_batch = 0
785
+ stride_dA_chunk = 0
786
+ grid_chunk_cs = lambda META: (batch, nchunks, triton.cdiv(nheads, META['BLOCK_SIZE_H']))
787
+ with torch.cuda.device(dt.device.index):
788
+ _chunk_cumsum_bwd_kernel[grid_chunk_cs](
789
+ ddA, ddt_out, dt, A, dt_bias, ddt, dA_workspace, ddt_bias_workspace if ddt_bias is not None else None,
790
+ batch, seqlen, nheads, chunk_size,
791
+ dt_limit[0], dt_limit[1],
792
+ ddA.stride(0), ddA.stride(2), ddA.stride(1), ddA.stride(3),
793
+ ddt_out.stride(0), ddt_out.stride(2), ddt_out.stride(1), ddt_out.stride(3),
794
+ dt.stride(0), dt.stride(1), dt.stride(2),
795
+ A.stride(0),
796
+ dt_bias.stride(0) if dt_bias is not None else 0,
797
+ ddt.stride(0), ddt.stride(1), ddt.stride(2),
798
+ stride_dA_batch, stride_dA_chunk, dA_workspace.stride(-1),
799
+ stride_ddt_bias_batch, stride_ddt_bias_chunk, (ddt_bias_workspace.stride(-1) if ddt_bias is not None else 0),
800
+ dt_softplus,
801
+ HAS_DT_BIAS=dt_bias is not None,
802
+ BLOCK_SIZE_CHUNK=triton.next_power_of_2(chunk_size),
803
+ DETERMINISTIC_REDUCTION=deterministic,
804
+ )
805
+ if deterministic:
806
+ dA.copy_(dA_workspace.sum(dim=(0, 1)))
807
+ if ddt_bias is not None:
808
+ ddt_bias.copy_(ddt_bias_workspace.sum(dim=(0, 1)))
809
+ return ddt, dA, ddt_bias
810
+
811
+
812
+ def _chunk_state_fwd(B, x, dt, dA_cumsum, seq_idx=None, states=None, states_in_fp32=True):
813
+ batch, seqlen, nheads, headdim = x.shape
814
+ _, _, nchunks, chunk_size = dt.shape
815
+ _, _, ngroups, dstate = B.shape
816
+ assert nheads % ngroups == 0
817
+ assert B.shape == (batch, seqlen, ngroups, dstate)
818
+ assert dt.shape == (batch, nheads, nchunks, chunk_size)
819
+ assert dA_cumsum.shape == dt.shape
820
+ if seq_idx is not None:
821
+ assert seq_idx.shape == (batch, seqlen)
822
+ if states is not None:
823
+ assert states.shape == (batch, nchunks, nheads, headdim, dstate)
824
+ else:
825
+ states_dtype = torch.float32 if states_in_fp32 else B.dtype
826
+ states = torch.empty((batch, nchunks, nheads, headdim, dstate), device=x.device, dtype=states_dtype)
827
+ grid = lambda META: (triton.cdiv(headdim, META['BLOCK_SIZE_M']) * triton.cdiv(dstate, META['BLOCK_SIZE_N']),
828
+ batch * nchunks, nheads)
829
+ with torch.cuda.device(x.device.index):
830
+ _chunk_state_fwd_kernel[grid](
831
+ x, B, states, dt, dA_cumsum, seq_idx,
832
+ headdim, dstate, chunk_size,
833
+ batch, seqlen, nheads // ngroups,
834
+ x.stride(0), x.stride(1), x.stride(2), x.stride(3),
835
+ B.stride(0), B.stride(1), B.stride(2), B.stride(-1),
836
+ states.stride(0), states.stride(1), states.stride(2), states.stride(3), states.stride(4),
837
+ dt.stride(0), dt.stride(2), dt.stride(1), dt.stride(3),
838
+ dA_cumsum.stride(0), dA_cumsum.stride(2), dA_cumsum.stride(1), dA_cumsum.stride(3),
839
+ *((seq_idx.stride(0), seq_idx.stride(1)) if seq_idx is not None else (0, 0)),
840
+ HAS_SEQ_IDX=seq_idx is not None,
841
+ )
842
+ return states
843
+
844
+
845
+ def _chunk_state_bwd_dx(B, x, dt, dA_cumsum, dstates, dx=None):
846
+ batch, seqlen, nheads, headdim = x.shape
847
+ _, _, nchunks, chunk_size = dt.shape
848
+ _, _, ngroups, dstate = B.shape
849
+ assert nheads % ngroups == 0
850
+ assert B.shape == (batch, seqlen, ngroups, dstate)
851
+ assert dt.shape == (batch, nheads, nchunks, chunk_size)
852
+ assert dA_cumsum.shape == dt.shape
853
+ assert dstates.shape == (batch, nchunks, nheads, headdim, dstate)
854
+ if dx is not None:
855
+ assert dx.shape == x.shape
856
+ else:
857
+ dx = torch.empty_like(x)
858
+ deterministic = use_deterministic_mode()
859
+ tile_count = math.ceil(headdim / _CHUNK_STATE_BWD_DX_MIN_BLOCK_N)
860
+ ddt, stride_ddt_tile = alloc_tile_workspace(
861
+ (batch, nheads, nchunks, chunk_size),
862
+ tile_count,
863
+ torch.float32,
864
+ dt.device,
865
+ deterministic,
866
+ zero_init=True,
867
+ )
868
+ ddA_cumsum, stride_ddA_tile = alloc_tile_workspace(
869
+ (batch, nheads, nchunks, chunk_size),
870
+ tile_count,
871
+ torch.float32,
872
+ dA_cumsum.device,
873
+ deterministic,
874
+ zero_init=True,
875
+ )
876
+ grid_dx = lambda META: (triton.cdiv(chunk_size, META['BLOCK_SIZE_M']) * triton.cdiv(headdim, META['BLOCK_SIZE_N']),
877
+ batch * nchunks, nheads)
878
+ with torch.cuda.device(x.device.index):
879
+ _chunk_state_bwd_dx_kernel[grid_dx](
880
+ x, B, dstates, dt, dA_cumsum, dx, ddt, ddA_cumsum,
881
+ chunk_size, headdim, dstate,
882
+ batch, seqlen, nheads // ngroups,
883
+ x.stride(0), x.stride(1), x.stride(2), x.stride(3),
884
+ B.stride(0), B.stride(1), B.stride(2), B.stride(-1),
885
+ dstates.stride(0), dstates.stride(1), dstates.stride(2), dstates.stride(3), dstates.stride(4),
886
+ dt.stride(0), dt.stride(2), dt.stride(1), dt.stride(3),
887
+ dA_cumsum.stride(0), dA_cumsum.stride(2), dA_cumsum.stride(1), dA_cumsum.stride(3),
888
+ dx.stride(0), dx.stride(1), dx.stride(2), dx.stride(3),
889
+ ddt.stride(0), ddt.stride(2), ddt.stride(1), ddt.stride(3), stride_ddt_tile,
890
+ ddA_cumsum.stride(0), ddA_cumsum.stride(2), ddA_cumsum.stride(1), ddA_cumsum.stride(3), stride_ddA_tile,
891
+ DETERMINISTIC_REDUCTION=deterministic,
892
+ BLOCK_SIZE_DSTATE=max(triton.next_power_of_2(dstate), 16),
893
+ )
894
+ ddt = finalize_tile_workspace(ddt, deterministic)
895
+ ddA_cumsum = finalize_tile_workspace(ddA_cumsum, deterministic)
896
+ if deterministic:
897
+ # Match `_chunk_state_bwd_dx_kernel` atomic path (`tl.atomic_add(..., ddA_cs_last)` into last element).
898
+ ddA_cumsum[..., -1] -= ddA_cumsum.sum(dim=-1)
899
+ return dx, ddt, ddA_cumsum
900
+
901
+
902
+ def _chunk_state_bwd_db(x, dt, dA_cumsum, dstates, seq_idx=None, B=None, ngroups=1):
903
+ batch, seqlen, nheads, headdim = x.shape
904
+ _, _, nchunks, chunk_size = dt.shape
905
+ dstate = dstates.shape[-1]
906
+ assert dt.shape == (batch, nheads, nchunks, chunk_size)
907
+ assert dA_cumsum.shape == dt.shape
908
+ assert dstates.shape == (batch, nchunks, nheads, headdim, dstate)
909
+ if seq_idx is not None:
910
+ assert seq_idx.shape == (batch, seqlen)
911
+ deterministic = use_deterministic_mode()
912
+ if B is not None:
913
+ assert B.shape == (batch, seqlen, ngroups, dstate)
914
+ B_strides = (B.stride(0), B.stride(1), B.stride(2), B.stride(3))
915
+ # Use torch.empty since the Triton kernel will call init_to_zero
916
+ tile_count = math.ceil(dstate / _CHUNK_STATE_BWD_DB_MIN_BLOCK_N)
917
+ ddA_cumsum, stride_ddA_tile = alloc_tile_workspace(
918
+ (batch, nheads, nchunks, chunk_size),
919
+ tile_count,
920
+ torch.float32,
921
+ x.device,
922
+ deterministic,
923
+ zero_init=True,
924
+ )
925
+ ddA_cumsum_strides = (
926
+ ddA_cumsum.stride(0),
927
+ ddA_cumsum.stride(2),
928
+ ddA_cumsum.stride(1),
929
+ ddA_cumsum.stride(3),
930
+ )
931
+ else:
932
+ B_strides = (0, 0, 0, 0)
933
+ ddA_cumsum = None
934
+ ddA_cumsum_strides = (0, 0, 0, 0)
935
+ stride_ddA_tile = 0
936
+ nheads_ngroups_ratio = nheads // ngroups
937
+ sm_count = torch.cuda.get_device_properties(x.device).multi_processor_count
938
+ nheads_per_program = max(min(math.ceil(batch * nchunks * nheads / sm_count), nheads_ngroups_ratio), 1)
939
+ nsplits = triton.cdiv(nheads_ngroups_ratio, nheads_per_program)
940
+ dB = torch.empty(batch, seqlen, nsplits, ngroups, dstate, device=x.device, dtype=torch.float32)
941
+ grid_db = lambda META: (triton.cdiv(chunk_size, META['BLOCK_SIZE_M']) * triton.cdiv(dstate, META['BLOCK_SIZE_N']),
942
+ batch * nchunks, nsplits * ngroups)
943
+ with torch.cuda.device(x.device.index):
944
+ _chunk_state_bwd_db_kernel[grid_db](
945
+ x, dstates, B, dt, dA_cumsum, seq_idx, dB, ddA_cumsum,
946
+ chunk_size, dstate, headdim,
947
+ batch, seqlen, nheads, nheads_per_program, ngroups,
948
+ x.stride(0), x.stride(1), x.stride(2), x.stride(3),
949
+ dstates.stride(0), dstates.stride(1), dstates.stride(2), dstates.stride(3), dstates.stride(4),
950
+ *B_strides,
951
+ dt.stride(0), dt.stride(2), dt.stride(1), dt.stride(3),
952
+ dA_cumsum.stride(0), dA_cumsum.stride(2), dA_cumsum.stride(1), dA_cumsum.stride(3),
953
+ *((seq_idx.stride(0), seq_idx.stride(1)) if seq_idx is not None else (0, 0)),
954
+ dB.stride(0), dB.stride(1), dB.stride(2), dB.stride(3), dB.stride(4),
955
+ *ddA_cumsum_strides, stride_ddA_tile,
956
+ HAS_DDA_CS=ddA_cumsum is not None,
957
+ HAS_SEQ_IDX=seq_idx is not None,
958
+ DETERMINISTIC_REDUCTION=deterministic,
959
+ BLOCK_SIZE_K=max(triton.next_power_of_2(headdim), 16),
960
+ )
961
+ dB = dB.sum(2)
962
+ if ddA_cumsum is not None:
963
+ ddA_cumsum = finalize_tile_workspace(ddA_cumsum, deterministic)
964
+ # The first element of ddA_cumsum is always zero, since that dA_cumsum does not contribute
965
+ # to the state of the chunk.
966
+ # torch.cumsum(ddA_cumsum[..., 1:], dim=-1, out=ddA_cumsum[..., 1:])
967
+ # But it's easier to just do the cumsum for all elements, the result will be the same.
968
+ torch.cumsum(ddA_cumsum, dim=-1, out=ddA_cumsum)
969
+ return dB if B is None else (dB, ddA_cumsum)
970
+
971
+
972
+ def _chunk_state_bwd_ddAcs_stable(B, x, dt, dA_cumsum, dstates, seq_idx=None):
973
+ batch, seqlen, nheads, headdim = x.shape
974
+ _, _, nchunks, chunk_size = dt.shape
975
+ _, _, ngroups, dstate = B.shape
976
+ assert nheads % ngroups == 0
977
+ assert B.shape == (batch, seqlen, ngroups, dstate)
978
+ assert dt.shape == (batch, nheads, nchunks, chunk_size)
979
+ assert dA_cumsum.shape == dt.shape
980
+ assert dstates.shape == (batch, nchunks, nheads, headdim, dstate)
981
+ if seq_idx is not None:
982
+ assert seq_idx.shape == (batch, seqlen)
983
+ # Use torch.empty since the Triton kernel will call init_to_zero
984
+ deterministic = use_deterministic_mode()
985
+ tile_count = math.ceil(headdim / _CHUNK_STATE_BWD_DDACS_MIN_BLOCK_N)
986
+ ddA_cumsum, stride_ddA_tile = alloc_tile_workspace(
987
+ (batch, nheads, nchunks, chunk_size),
988
+ tile_count,
989
+ torch.float32,
990
+ x.device,
991
+ deterministic,
992
+ zero_init=True,
993
+ )
994
+ grid_ddtcs = lambda META: (triton.cdiv(chunk_size, META['BLOCK_SIZE_M']) * triton.cdiv(headdim, META['BLOCK_SIZE_N']),
995
+ batch * nchunks, nheads)
996
+ with torch.cuda.device(x.device.index):
997
+ _chunk_state_bwd_ddAcs_stable_kernel[grid_ddtcs](
998
+ x, B, dstates, dt, dA_cumsum, seq_idx, ddA_cumsum,
999
+ chunk_size, headdim, dstate,
1000
+ batch, seqlen, nheads // ngroups,
1001
+ x.stride(0), x.stride(1), x.stride(2), x.stride(3),
1002
+ B.stride(0), B.stride(1), B.stride(2), B.stride(-1),
1003
+ dstates.stride(0), dstates.stride(1), dstates.stride(2), dstates.stride(3), dstates.stride(4),
1004
+ dt.stride(0), dt.stride(2), dt.stride(1), dt.stride(3),
1005
+ dA_cumsum.stride(0), dA_cumsum.stride(2), dA_cumsum.stride(1), dA_cumsum.stride(3),
1006
+ *((seq_idx.stride(0), seq_idx.stride(1)) if seq_idx is not None else (0, 0)),
1007
+ ddA_cumsum.stride(0), ddA_cumsum.stride(2), ddA_cumsum.stride(1), ddA_cumsum.stride(3), stride_ddA_tile,
1008
+ HAS_SEQ_IDX=seq_idx is not None,
1009
+ DETERMINISTIC_REDUCTION=deterministic,
1010
+ BLOCK_SIZE_M=max(triton.next_power_of_2(chunk_size), 16),
1011
+ BLOCK_SIZE_DSTATE=max(triton.next_power_of_2(dstate), 16),
1012
+ )
1013
+ ddA_cumsum = finalize_tile_workspace(ddA_cumsum, deterministic)
1014
+ torch.cumsum(ddA_cumsum[..., 1:], dim=-1, out=ddA_cumsum[..., 1:])
1015
+ return ddA_cumsum
1016
+
1017
+
1018
+ def chunk_state_varlen(B, x, dt, dA_cumsum, cu_seqlens, chunk_states):
1019
+ total_seqlen, nheads, headdim = x.shape
1020
+ _, nchunks, chunk_size = dt.shape
1021
+ _, ngroups, dstate = B.shape
1022
+ batch = cu_seqlens.shape[0] - 1
1023
+ cu_seqlens = cu_seqlens.contiguous()
1024
+ assert nheads % ngroups == 0
1025
+ assert B.shape == (total_seqlen, ngroups, dstate)
1026
+ assert dt.shape == (nheads, nchunks, chunk_size)
1027
+ assert dA_cumsum.shape == dt.shape
1028
+ assert chunk_states.shape == (nchunks, nheads, headdim, dstate)
1029
+ states = torch.empty(batch, nheads, headdim, dstate, dtype=chunk_states.dtype, device=chunk_states.device)
1030
+ grid = lambda META: (triton.cdiv(headdim, META['BLOCK_SIZE_M']) * triton.cdiv(dstate, META['BLOCK_SIZE_N']),
1031
+ batch, nheads)
1032
+ with torch.cuda.device(x.device.index):
1033
+ _chunk_state_varlen_kernel[grid](
1034
+ x, B, dt, dA_cumsum, chunk_states, cu_seqlens, states,
1035
+ headdim, dstate, chunk_size,
1036
+ total_seqlen, nheads // ngroups,
1037
+ x.stride(0), x.stride(1), x.stride(2),
1038
+ B.stride(0), B.stride(1), B.stride(2),
1039
+ dt.stride(1), dt.stride(0), dt.stride(2),
1040
+ dA_cumsum.stride(1), dA_cumsum.stride(0), dA_cumsum.stride(2),
1041
+ chunk_states.stride(0), chunk_states.stride(1), chunk_states.stride(2), chunk_states.stride(3),
1042
+ states.stride(0), states.stride(1), states.stride(2), states.stride(3),
1043
+ )
1044
+ return states
1045
+
1046
+
1047
+ class ChunkStateFn(torch.autograd.Function):
1048
+
1049
+ @staticmethod
1050
+ def forward(ctx, B, x, dt, dA_cumsum, states_in_fp32=True):
1051
+ batch, seqlen, nheads, headdim = x.shape
1052
+ _, _, nchunks, chunk_size = dt.shape
1053
+ assert seqlen <= nchunks * chunk_size
1054
+ _, _, ngroups, dstate = B.shape
1055
+ assert B.shape == (batch, seqlen, ngroups, dstate)
1056
+ assert dt.shape == (batch, nheads, nchunks, chunk_size)
1057
+ assert dA_cumsum.shape == (batch, nheads, nchunks, chunk_size)
1058
+ if B.stride(-1) != 1:
1059
+ B = B.contiguous()
1060
+ if x.stride(-1) != 1 and x.stride(1) != 1: # Either M or K dimension should be contiguous
1061
+ x = x.contiguous()
1062
+ states = _chunk_state_fwd(B, x, dt, dA_cumsum, states_in_fp32=states_in_fp32)
1063
+ ctx.save_for_backward(B, x, dt, dA_cumsum)
1064
+ return states
1065
+
1066
+ @staticmethod
1067
+ def backward(ctx, dstates):
1068
+ B, x, dt, dA_cumsum = ctx.saved_tensors
1069
+ batch, seqlen, nheads, headdim = x.shape
1070
+ _, _, nchunks, chunk_size = dt.shape
1071
+ _, _, ngroups, dstate = B.shape
1072
+ assert dstates.shape == (batch, nchunks, nheads, headdim, dstate)
1073
+ if dstates.stride(-1) != 1:
1074
+ dstates = dstates.contiguous()
1075
+ dx, ddt, ddA_cumsum = _chunk_state_bwd_dx(B, x, dt, dA_cumsum, dstates)
1076
+ dB = _chunk_state_bwd_db(x, dt, dA_cumsum, dstates, ngroups=ngroups)
1077
+ dB = dB.to(B.dtype)
1078
+ return dB, dx, ddt, ddA_cumsum, None
1079
+
1080
+
1081
+ def chunk_state(B, x, dt, dA_cumsum, states_in_fp32=True):
1082
+ """
1083
+ Argument:
1084
+ B: (batch, seqlen, ngroups, dstate)
1085
+ x: (batch, seqlen, nheads, headdim)
1086
+ dt: (batch, nheads, nchunks, chunk_size)
1087
+ dA_cumsum: (batch, nheads, nchunks, chunk_size)
1088
+ Return:
1089
+ states: (batch, nchunks, nheads, headdim, dstate)
1090
+ """
1091
+ return ChunkStateFn.apply(B, x, dt, dA_cumsum, states_in_fp32)
1092
+
1093
+
1094
+ def chunk_state_ref(B, x, dt, dA_cumsum):
1095
+ """
1096
+ Argument:
1097
+ B: (batch, seqlen, ngroups, dstate)
1098
+ x: (batch, seqlen, nheads, headdim)
1099
+ dt: (batch, nheads, nchunks, chunk_size)
1100
+ dA_cumsum: (batch, nheads, nchunks, chunk_size)
1101
+ Return:
1102
+ states: (batch, nchunks, nheads, headdim, dstate)
1103
+ """
1104
+ # Check constraints.
1105
+ batch, seqlen, nheads, headdim = x.shape
1106
+ dstate = B.shape[-1]
1107
+ _, _, nchunks, chunk_size = dt.shape
1108
+ assert seqlen <= nchunks * chunk_size
1109
+ assert x.shape == (batch, seqlen, nheads, headdim)
1110
+ assert dt.shape == (batch, nheads, nchunks, chunk_size)
1111
+ ngroups = B.shape[2]
1112
+ assert nheads % ngroups == 0
1113
+ assert B.shape == (batch, seqlen, ngroups, dstate)
1114
+ B = repeat(B, "b l g d -> b l (g h) d", h=nheads // ngroups)
1115
+ assert dA_cumsum.shape == (batch, nheads, nchunks, chunk_size)
1116
+ if seqlen < nchunks * chunk_size:
1117
+ x = F.pad(x, (0, 0, 0, 0, 0, nchunks * chunk_size - seqlen))
1118
+ B = F.pad(B, (0, 0, 0, 0, 0, nchunks * chunk_size - seqlen))
1119
+ x = rearrange(x, "b (c l) h p -> b c l h p", l=chunk_size)
1120
+ B = rearrange(B, "b (c l) ... -> b c l ...", l=chunk_size)
1121
+ decay_states = torch.exp((dA_cumsum[:, :, :, -1:] - dA_cumsum))
1122
+ return torch.einsum("bclhn,bhcl,bhcl,bclhp->bchpn", B.to(x.dtype), decay_states.to(x.dtype), dt.to(x.dtype), x)
mamba_ssm/ops/triton/ssd_combined.py ADDED
@@ -0,0 +1,1047 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2024, Tri Dao, Albert Gu.
2
+
3
+ """We want triton==2.1.0 or 2.2.0 for this
4
+ """
5
+
6
+ from typing import Optional
7
+
8
+ import math
9
+ from packaging import version
10
+
11
+ import torch
12
+ import torch.nn.functional as F
13
+ from torch import Tensor
14
+ from mamba_ssm.utils.torch import custom_bwd, custom_fwd
15
+
16
+ import triton
17
+ import triton.language as tl
18
+
19
+ from einops import rearrange, repeat
20
+
21
+ try:
22
+ from causal_conv1d import causal_conv1d_fn
23
+ from causal_conv1d.cpp_functions import causal_conv1d_fwd_function, causal_conv1d_bwd_function, causal_conv1d_update_function
24
+ except ImportError:
25
+ causal_conv1d_fn = None
26
+ causal_conv1d_fwd_function = None
27
+ causal_conv1d_bwd_function = None
28
+ causal_conv1d_update_function = None
29
+
30
+ from mamba_ssm.ops.triton.ssd_bmm import _bmm_chunk_fwd, _bmm_chunk_bwd
31
+ from mamba_ssm.ops.triton.ssd_chunk_state import _chunk_cumsum_fwd, _chunk_cumsum_bwd
32
+ from mamba_ssm.ops.triton.ssd_chunk_state import _chunk_state_fwd, _chunk_state_bwd_db
33
+ from mamba_ssm.ops.triton.ssd_chunk_state import _chunk_state_bwd_ddAcs_stable
34
+ from mamba_ssm.ops.triton.ssd_chunk_state import chunk_state, chunk_state_ref
35
+ from mamba_ssm.ops.triton.ssd_chunk_state import chunk_state_varlen
36
+ from mamba_ssm.ops.triton.ssd_state_passing import _state_passing_fwd, _state_passing_bwd
37
+ from mamba_ssm.ops.triton.ssd_state_passing import state_passing, state_passing_ref
38
+ from mamba_ssm.ops.triton.ssd_chunk_scan import _chunk_scan_fwd, _chunk_scan_bwd_dz, _chunk_scan_bwd_dstates
39
+ from mamba_ssm.ops.triton.ssd_chunk_scan import _chunk_scan_bwd_dC, _chunk_scan_bwd_dcb
40
+ from mamba_ssm.ops.triton.ssd_chunk_scan import _chunk_scan_bwd_ddAcs_stable
41
+ from mamba_ssm.ops.triton.ssd_chunk_scan import chunk_scan, chunk_scan_ref
42
+ from mamba_ssm.ops.triton.ssd_chunk_scan import _chunk_scan_bwd_ddAcs_prev
43
+ from mamba_ssm.ops.triton.layernorm_gated import rmsnorm_fn, _layer_norm_fwd, _layer_norm_bwd
44
+ from mamba_ssm.ops.triton.k_activations import _swiglu_fwd, _swiglu_bwd
45
+ from mamba_ssm.utils.determinism import (
46
+ alloc_tile_workspace,
47
+ autotune_configs,
48
+ finalize_tile_workspace,
49
+ use_deterministic_mode,
50
+ )
51
+
52
+ TRITON_22 = version.parse(triton.__version__) >= version.parse('2.2.0')
53
+
54
+
55
+ def init_to_zero(names):
56
+ return lambda nargs: [nargs[name].zero_() for name in names if nargs[name] is not None]
57
+
58
+
59
+ def ensure_stride(inp):
60
+ """
61
+ Return inp, while ensuring that stride(1) of the returned tensor is a multiple of 8.
62
+
63
+ The inp tensor is of shape [batch, length, channels], where channels is assumed, and tested, to be
64
+ a multiple of 8. If it is contiguous, inp will have strides [length*channels, channels, 1]. The
65
+ output of this function will be rearranged to shape [batch, channels, length] before being passed to
66
+ causal_conv1d. That rearranged tensor will have strides [length*channels, 1, channels].
67
+ causal_conv1d handles this stride configuration (which it calls channels_last) directly and
68
+ efficiently, after first recognizing it (when stride[1]==1 and stride[2]>1). causal_conv1d cannot
69
+ operate on a channels_last tensor for which stride[2] is not a multiple of 8, and in that case will
70
+ raise an exception. This function prevents the aforementioned exception by returning a tensor with
71
+ stride(1) equal to channels, by making the returned tensor contiguous, if inp.stride(1) is not
72
+ already a multiple of 8.
73
+ """
74
+ assert inp.shape[2] % 8 == 0, "Number of convolution channels is required to be a multiple of 8."
75
+ return inp if inp.stride(1) % 8 == 0 else inp.contiguous()
76
+
77
+
78
+ @triton.autotune(
79
+ configs=autotune_configs([
80
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 64}, num_stages=3, num_warps=8, pre_hook=init_to_zero(["ddt_ptr", "dD_ptr"])),
81
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "dD_ptr"])),
82
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "dD_ptr"])),
83
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "dD_ptr"])),
84
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "dD_ptr"])),
85
+ triton.Config({'BLOCK_SIZE_M': 128, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "dD_ptr"])),
86
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 32, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "dD_ptr"])),
87
+ triton.Config({'BLOCK_SIZE_M': 32, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=5, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "dD_ptr"])),
88
+ triton.Config({'BLOCK_SIZE_M': 64, 'BLOCK_SIZE_N': 64, 'BLOCK_SIZE_K': 32}, num_stages=4, num_warps=4, pre_hook=init_to_zero(["ddt_ptr", "dD_ptr"])),
89
+ ]),
90
+ key=['chunk_size', 'hdim', 'dstate'],
91
+ )
92
+ @triton.jit
93
+ def _chunk_scan_chunk_state_bwd_dx_kernel(
94
+ # Pointers to matrices
95
+ x_ptr, cb_ptr, dout_ptr, dt_ptr, dA_cumsum_ptr, seq_idx_ptr, D_ptr,
96
+ b_ptr, dstates_ptr,
97
+ dx_ptr, ddt_ptr, dD_ptr,
98
+ # Matrix dimensions
99
+ chunk_size, hdim, dstate,
100
+ batch, seqlen, nheads_ngroups_ratio,
101
+ # Strides
102
+ stride_x_batch, stride_x_seqlen, stride_x_head, stride_x_hdim,
103
+ stride_cb_batch, stride_cb_chunk, stride_cb_head, stride_cb_csize_m, stride_cb_csize_k,
104
+ stride_dout_batch, stride_dout_seqlen, stride_dout_head, stride_dout_hdim,
105
+ stride_dt_batch, stride_dt_chunk, stride_dt_head, stride_dt_csize,
106
+ stride_dA_cs_batch, stride_dA_cs_chunk, stride_dA_cs_head, stride_dA_cs_csize,
107
+ stride_seq_idx_batch, stride_seq_idx_seqlen,
108
+ stride_D_head,
109
+ stride_b_batch, stride_b_seqlen, stride_b_head, stride_b_dstate,
110
+ stride_dstates_batch, stride_dstates_chunk, stride_dstates_head, stride_dstates_hdim, stride_dstates_dstate,
111
+ stride_dx_batch, stride_dx_seqlen, stride_dx_head, stride_dx_hdim,
112
+ stride_ddt_batch, stride_ddt_chunk, stride_ddt_head, stride_ddt_csize, stride_ddt_tile,
113
+ stride_dD_batch, stride_dD_chunk, stride_dD_head, stride_dD_csize, stride_dD_hdim,
114
+ # Meta-parameters
115
+ HAS_D: tl.constexpr,
116
+ D_HAS_HDIM: tl.constexpr,
117
+ HAS_SEQ_IDX: tl.constexpr,
118
+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
119
+ BLOCK_SIZE_DSTATE: tl.constexpr,
120
+ IS_TRITON_22: tl.constexpr,
121
+ DETERMINISTIC_REDUCTION: tl.constexpr,
122
+ ):
123
+ pid_bc = tl.program_id(axis=1)
124
+ pid_c = pid_bc // batch
125
+ pid_b = pid_bc - pid_c * batch
126
+ pid_h = tl.program_id(axis=2)
127
+ num_pid_n = tl.cdiv(hdim, BLOCK_SIZE_N)
128
+ pid_m = tl.program_id(axis=0) // num_pid_n
129
+ pid_n = tl.program_id(axis=0) % num_pid_n
130
+ x_ptr += pid_b * stride_x_batch + pid_c * chunk_size * stride_x_seqlen + pid_h * stride_x_head
131
+ cb_ptr += pid_b * stride_cb_batch + pid_c * stride_cb_chunk + (pid_h // nheads_ngroups_ratio) * stride_cb_head
132
+ dout_ptr += pid_b * stride_dout_batch + pid_c * chunk_size * stride_dout_seqlen + pid_h * stride_dout_head
133
+ dt_ptr += pid_b * stride_dt_batch + pid_c * stride_dt_chunk + pid_h * stride_dt_head
134
+ ddt_ptr += pid_b * stride_ddt_batch + pid_c * stride_ddt_chunk + pid_h * stride_ddt_head + pid_n * stride_ddt_tile
135
+ dA_cumsum_ptr += pid_b * stride_dA_cs_batch + pid_c * stride_dA_cs_chunk + pid_h * stride_dA_cs_head
136
+ b_ptr += pid_b * stride_b_batch + pid_c * chunk_size * stride_b_seqlen + (pid_h // nheads_ngroups_ratio) * stride_b_head
137
+ dstates_ptr += pid_b * stride_dstates_batch + pid_c * stride_dstates_chunk + pid_h * stride_dstates_head
138
+ if HAS_SEQ_IDX:
139
+ seq_idx_ptr += pid_b * stride_seq_idx_batch + pid_c * chunk_size * stride_seq_idx_seqlen
140
+
141
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
142
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
143
+
144
+ chunk_size_limit = min(chunk_size, seqlen - pid_c * chunk_size)
145
+
146
+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
147
+
148
+ dA_cs_m = tl.load(dA_cumsum_ptr + offs_m * stride_dA_cs_csize, mask=offs_m < chunk_size_limit, other=0.0).to(tl.float32)
149
+
150
+ dA_cs_last = tl.load(dA_cumsum_ptr + (chunk_size - 1) * stride_dA_cs_csize).to(tl.float32)
151
+ if not HAS_SEQ_IDX:
152
+ # scale = tl.exp(dA_cs_last - dA_cs_m)
153
+ scale = tl.exp(tl.minimum((dA_cs_last - dA_cs_m), 0.0))
154
+ else:
155
+ seq_idx_m = tl.load(seq_idx_ptr + offs_m * stride_seq_idx_seqlen, mask=offs_m < chunk_size_limit, other=-1)
156
+ seq_idx_last = tl.load(seq_idx_ptr + (chunk_size_limit - 1) * stride_seq_idx_seqlen)
157
+ # scale = tl.where(seq_idx_m == seq_idx_last, tl.exp(dA_cs_last - dA_cs_m), 0.0)
158
+ scale = tl.where(seq_idx_m == seq_idx_last, tl.exp(tl.minimum((dA_cs_last - dA_cs_m), 0.0)), 0.0)
159
+ # Might be faster to just do 1 iteration with larger BLOCK_SIZE_K, up to block size 128
160
+ # However, we're getting error with the Triton compiler 2.1.0 for that code path:
161
+ # Unexpected mma -> mma layout conversion
162
+ # Triton 2.2.0 fixes this
163
+ offs_dstate = tl.arange(0, BLOCK_SIZE_DSTATE if IS_TRITON_22 and BLOCK_SIZE_DSTATE <= 128 else BLOCK_SIZE_K)
164
+ b_ptrs = b_ptr + (offs_m[:, None] * stride_b_seqlen + offs_dstate[None, :] * stride_b_dstate)
165
+ dstates_ptrs = dstates_ptr + (offs_n[None, :] * stride_dstates_hdim + offs_dstate[:, None] * stride_dstates_dstate)
166
+ if IS_TRITON_22 and BLOCK_SIZE_DSTATE <= 128:
167
+ b = tl.load(b_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_dstate[None, :] < dstate), other=0.0)
168
+ dstates = tl.load(dstates_ptrs, mask=(offs_dstate[:, None] < dstate) & (offs_n[None, :] < hdim), other=0.0)
169
+ dstates = dstates.to(b_ptr.dtype.element_ty)
170
+ acc = tl.dot(b, dstates) * scale[:, None]
171
+ else:
172
+ for k in range(0, dstate, BLOCK_SIZE_K):
173
+ b = tl.load(b_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_dstate[None, :] < dstate - k), other=0.0)
174
+ dstates = tl.load(dstates_ptrs, mask=(offs_dstate[:, None] < dstate - k) & (offs_n[None, :] < hdim), other=0.0)
175
+ dstates = dstates.to(b_ptr.dtype.element_ty)
176
+ acc += tl.dot(b, dstates)
177
+ b_ptrs += BLOCK_SIZE_K * stride_b_dstate
178
+ dstates_ptrs += BLOCK_SIZE_K * stride_dstates_dstate
179
+ acc *= scale[:, None]
180
+
181
+ # x_ptrs = x_ptr + (offs_m[:, None] * stride_x_seqlen + offs_n[None, :] * stride_x_hdim)
182
+ # x = tl.load(x_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_n[None, :] < hdim), other=0.0).to(tl.float32)
183
+ # dt_ptrs = dt_ptr + offs_m * stride_dt_csize
184
+ # dt_m = tl.load(dt_ptrs, mask=offs_m < chunk_size_limit, other=0.0).to(tl.float32)
185
+ # ddt = tl.sum(acc * x, axis=1) * dt_m
186
+ # ddt_ptrs = ddt_ptr + offs_m * stride_ddt_csize
187
+ # tl.atomic_add(ddt_ptrs, ddt, mask=offs_m < chunk_size)
188
+
189
+ offs_k = tl.arange(0, BLOCK_SIZE_K)
190
+ cb_ptrs = cb_ptr + (offs_m[:, None] * stride_cb_csize_m + offs_k[None, :] * stride_cb_csize_k)
191
+ dout_ptrs = dout_ptr + (offs_k[:, None] * stride_dout_seqlen + offs_n[None, :] * stride_dout_hdim)
192
+ dA_cumsum_ptrs = dA_cumsum_ptr + offs_k * stride_dA_cs_csize
193
+ K_MAX = chunk_size_limit
194
+ K_MIN = pid_m * BLOCK_SIZE_M
195
+ cb_ptrs += K_MIN * stride_cb_csize_k
196
+ dout_ptrs += K_MIN * stride_dout_seqlen
197
+ dA_cumsum_ptrs += K_MIN * stride_dA_cs_csize
198
+ for k in range(K_MIN, K_MAX, BLOCK_SIZE_K):
199
+ k = tl.multiple_of(k, BLOCK_SIZE_K)
200
+ # For some reason setting mask to (offs_m[:, None] < chunk_size_limit) is much slower
201
+ cb = tl.load(cb_ptrs, mask=(offs_m[:, None] < chunk_size) & (offs_k[None, :] < K_MAX - k), other=0.0)
202
+ dout = tl.load(dout_ptrs, mask=(offs_k[:, None] < K_MAX - k) & (offs_n[None, :] < hdim), other=0.0)
203
+ dA_cs_k = tl.load(dA_cumsum_ptrs, mask=offs_k < K_MAX - k, other=0.0).to(tl.float32)
204
+ # cb *= tl.exp(dA_cs_k[None, :] - dA_cs_m[:, None])
205
+ cb *= tl.exp(tl.minimum((dA_cs_k[None, :] - dA_cs_m[:, None]), 0.0))
206
+ # If we don't have the (k + offs_k[None, :] < K_MAX) mask, for indices outside this range,
207
+ # we might have dA_cs_m = 0.0 and dA_cs_k very negative, and tl.exp will return inf.
208
+ # Multiplying with cb, which is 0.0 outside the range, will make the result NaN.
209
+ # This will cause NaN in acc, and hence NaN in dx and ddt.
210
+ mask = (k + offs_k[None, :] >= offs_m[:, None]) & (k + offs_k[None, :] < K_MAX)
211
+ cb = tl.where(mask, cb, 0.0)
212
+ cb = cb.to(dout_ptr.dtype.element_ty)
213
+ acc += tl.dot(cb, dout)
214
+ cb_ptrs += BLOCK_SIZE_K * stride_cb_csize_k
215
+ dout_ptrs += BLOCK_SIZE_K * stride_dout_seqlen
216
+ dA_cumsum_ptrs += BLOCK_SIZE_K * stride_dA_cs_csize
217
+
218
+ offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
219
+ offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
220
+ dt_ptrs = dt_ptr + offs_m * stride_dt_csize
221
+ dt_m = tl.load(dt_ptrs, mask=offs_m < chunk_size_limit, other=0.0).to(tl.float32)
222
+ dx = acc * dt_m[:, None]
223
+ dx_ptr += pid_b * stride_dx_batch + pid_c * chunk_size * stride_dx_seqlen + pid_h * stride_dx_head
224
+ dx_ptrs = dx_ptr + (offs_m[:, None] * stride_dx_seqlen + offs_n[None, :] * stride_dx_hdim)
225
+ if HAS_D:
226
+ dout_res_ptrs = dout_ptr + (offs_m[:, None] * stride_dout_seqlen + offs_n[None, :] * stride_dout_hdim)
227
+ dout_res = tl.load(dout_res_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_n[None, :] < hdim), other=0.0).to(tl.float32)
228
+ if D_HAS_HDIM:
229
+ D = tl.load(D_ptr + pid_h * stride_D_head + offs_n, mask=offs_n < hdim, other=0.0).to(tl.float32)
230
+ else:
231
+ D = tl.load(D_ptr + pid_h * stride_D_head).to(tl.float32)
232
+ dx += dout_res * D
233
+ tl.store(dx_ptrs, dx, mask=(offs_m[:, None] < chunk_size_limit) & (offs_n[None, :] < hdim))
234
+
235
+ x_ptrs = x_ptr + (offs_m[:, None] * stride_x_seqlen + offs_n[None, :] * stride_x_hdim)
236
+ x = tl.load(x_ptrs, mask=(offs_m[:, None] < chunk_size_limit) & (offs_n[None, :] < hdim), other=0.0).to(tl.float32)
237
+ if HAS_D:
238
+ dD_ptr += pid_b * stride_dD_batch + pid_c * stride_dD_chunk + pid_h * stride_dD_head + pid_m * stride_dD_csize
239
+ if D_HAS_HDIM:
240
+ dD_ptrs = dD_ptr + offs_n * stride_dD_hdim
241
+ dD = tl.sum(dout_res * x, axis=0)
242
+ tl.store(dD_ptrs, dD, mask=offs_n < hdim)
243
+ else:
244
+ dD = tl.sum(dout_res * x)
245
+ if DETERMINISTIC_REDUCTION:
246
+ tl.store(dD_ptr + pid_n * stride_dD_hdim, dD)
247
+ else:
248
+ tl.atomic_add(dD_ptr, dD)
249
+ ddt = tl.sum(acc * x, axis=1)
250
+ ddt_ptrs = ddt_ptr + offs_m * stride_ddt_csize
251
+ if DETERMINISTIC_REDUCTION:
252
+ tl.store(ddt_ptrs, ddt, mask=offs_m < chunk_size)
253
+ else:
254
+ tl.atomic_add(ddt_ptrs, ddt, mask=offs_m < chunk_size)
255
+
256
+
257
+ _CHUNK_SCAN_CHUNK_STATE_BWD_DX_MIN_BLOCK_N = min(
258
+ cfg.kwargs['BLOCK_SIZE_N'] for cfg in _chunk_scan_chunk_state_bwd_dx_kernel.configs
259
+ )
260
+
261
+
262
+ def _chunk_scan_chunk_state_bwd_dx(x, dt, dA_cumsum, B, CB, dout, dstates, D=None, seq_idx=None, dx=None):
263
+ batch, seqlen, nheads, headdim = x.shape
264
+ _, _, nchunks, chunk_size = dt.shape
265
+ _, _, ngroups, dstate = B.shape
266
+ assert nheads % ngroups == 0
267
+ assert B.shape == (batch, seqlen, ngroups, dstate)
268
+ assert CB.shape == (batch, nchunks, ngroups, chunk_size, chunk_size)
269
+ assert dt.shape == (batch, nheads, nchunks, chunk_size)
270
+ assert dA_cumsum.shape == dt.shape
271
+ assert dout.shape == x.shape
272
+ assert dstates.shape == (batch, nchunks, nheads, headdim, dstate)
273
+ if seq_idx is not None:
274
+ assert seq_idx.shape == (batch, seqlen)
275
+ deterministic = use_deterministic_mode()
276
+ if D is not None:
277
+ assert D.shape == (nheads, headdim) or D.shape == (nheads,)
278
+ assert D.stride(-1) == 1
279
+ BLOCK_SIZE_min = 32
280
+ pid_m_tiles = triton.cdiv(chunk_size, BLOCK_SIZE_min)
281
+ pid_n_tiles = math.ceil(headdim / _CHUNK_SCAN_CHUNK_STATE_BWD_DX_MIN_BLOCK_N)
282
+ if D.dim() == 2:
283
+ dD_hdim = headdim
284
+ elif deterministic:
285
+ dD_hdim = pid_n_tiles
286
+ else:
287
+ dD_hdim = 1
288
+ dD = torch.zeros(pid_m_tiles, batch, nchunks, nheads, dD_hdim, device=D.device, dtype=torch.float32)
289
+ dD_strides = (dD.stride(0), dD.stride(1), dD.stride(2), dD.stride(3), dD.stride(4))
290
+ else:
291
+ dD = None
292
+ dD_strides = (0, 0, 0, 0, 0)
293
+ if dx is None:
294
+ dx = torch.empty_like(x)
295
+ else:
296
+ assert dx.shape == x.shape
297
+ tile_count = math.ceil(headdim / _CHUNK_SCAN_CHUNK_STATE_BWD_DX_MIN_BLOCK_N)
298
+ ddt, stride_ddt_tile = alloc_tile_workspace(
299
+ (batch, nheads, nchunks, chunk_size),
300
+ tile_count,
301
+ torch.float32,
302
+ dout.device,
303
+ deterministic,
304
+ zero_init=True,
305
+ )
306
+ grid_dx = lambda META: (triton.cdiv(chunk_size, META['BLOCK_SIZE_M']) * triton.cdiv(headdim, META['BLOCK_SIZE_N']),
307
+ batch * nchunks, nheads)
308
+ with torch.cuda.device(x.device.index):
309
+ _chunk_scan_chunk_state_bwd_dx_kernel[grid_dx](
310
+ x, CB, dout, dt, dA_cumsum, seq_idx, D, B, dstates, dx, ddt, dD,
311
+ chunk_size, headdim, dstate,
312
+ batch, seqlen, nheads // ngroups,
313
+ x.stride(0), x.stride(1), x.stride(2), x.stride(3),
314
+ CB.stride(0), CB.stride(1), CB.stride(2), CB.stride(-1), CB.stride(-2),
315
+ dout.stride(0), dout.stride(1), dout.stride(2), dout.stride(3),
316
+ dt.stride(0), dt.stride(2), dt.stride(1), dt.stride(3),
317
+ dA_cumsum.stride(0), dA_cumsum.stride(2), dA_cumsum.stride(1), dA_cumsum.stride(3),
318
+ *((seq_idx.stride(0), seq_idx.stride(1)) if seq_idx is not None else (0, 0)),
319
+ D.stride(0) if D is not None else 0,
320
+ B.stride(0), B.stride(1), B.stride(2), B.stride(3),
321
+ dstates.stride(0), dstates.stride(1), dstates.stride(2), dstates.stride(3), dstates.stride(4),
322
+ dx.stride(0), dx.stride(1), dx.stride(2), dx.stride(3),
323
+ ddt.stride(0), ddt.stride(2), ddt.stride(1), ddt.stride(3), stride_ddt_tile,
324
+ dD_strides[1], dD_strides[2], dD_strides[3], dD_strides[0], dD_strides[4],
325
+ D is not None,
326
+ D.dim() == 2 if D is not None else True,
327
+ HAS_SEQ_IDX=seq_idx is not None,
328
+ BLOCK_SIZE_DSTATE=max(triton.next_power_of_2(dstate), 16),
329
+ IS_TRITON_22=TRITON_22,
330
+ DETERMINISTIC_REDUCTION=deterministic,
331
+ )
332
+ if D is not None:
333
+ BLOCK_SIZE_actual = _chunk_scan_chunk_state_bwd_dx_kernel.best_config.kwargs["BLOCK_SIZE_M"]
334
+ n_valid_blocks = (chunk_size + BLOCK_SIZE_actual - 1) // BLOCK_SIZE_actual
335
+ dD = dD[:n_valid_blocks].sum(dim=(0, 1, 2))
336
+ if D.dim() == 1:
337
+ dD = dD.sum(dim=-1)
338
+ dD = dD.to(dtype=D.dtype)
339
+ ddt = finalize_tile_workspace(ddt, deterministic)
340
+ return dx, ddt, dD
341
+
342
+
343
+ def _mamba_chunk_scan_combined_fwd(x, dt, A, B, C, chunk_size, D=None, z=None, dt_bias=None, initial_states=None, seq_idx=None, cu_seqlens=None, dt_softplus=False, dt_limit=(0.0, float("inf"))):
344
+ batch, seqlen, nheads, headdim = x.shape
345
+ _, _, ngroups, dstate = B.shape
346
+ assert nheads % ngroups == 0
347
+ assert B.shape == (batch, seqlen, ngroups, dstate)
348
+ assert x.shape == (batch, seqlen, nheads, headdim)
349
+ assert dt.shape == (batch, seqlen, nheads)
350
+ assert A.shape == (nheads,)
351
+ assert C.shape == B.shape
352
+ if z is not None:
353
+ assert z.shape == x.shape
354
+ if D is not None:
355
+ assert D.shape == (nheads, headdim) or D.shape == (nheads,)
356
+ if seq_idx is not None:
357
+ assert seq_idx.shape == (batch, seqlen)
358
+ if B.stride(-1) != 1:
359
+ B = B.contiguous()
360
+ if C.stride(-1) != 1:
361
+ C = C.contiguous()
362
+ if x.stride(-1) != 1 and x.stride(1) != 1: # Either M or K dimension should be contiguous
363
+ x = x.contiguous()
364
+ if z is not None and z.stride(-1) != 1 and z.stride(1) != 1: # Either M or K dimension should be contiguous
365
+ z = z.contiguous()
366
+ if D is not None and D.stride(-1) != 1:
367
+ D = D.contiguous()
368
+ if initial_states is not None:
369
+ assert initial_states.shape == (batch, nheads, headdim, dstate)
370
+ # # (batch, nchunks, chunk_size, chunk_size) or (batch, nchunks, nheads, chunk_size, chunk_size)
371
+ # dA_cumsum_tmp0, dt_tmp0 = _chunk_cumsum_fwd(dt[:, :147], A, chunk_size, dt_bias=dt_bias, dt_softplus=dt_softplus)
372
+ # dA_cumsum_tmp1, dt_tmp1 = _chunk_cumsum_fwd(dt[:, 147:], A, chunk_size, dt_bias=dt_bias, dt_softplus=dt_softplus)
373
+ # dA_cumsum_tmp2, dt_tmp2 = _chunk_cumsum_fwd(dt[:, 147:256], A, chunk_size, dt_bias=dt_bias, dt_softplus=dt_softplus)
374
+ dA_cumsum, dt = _chunk_cumsum_fwd(dt, A, chunk_size, dt_bias=dt_bias, dt_softplus=dt_softplus, dt_limit=dt_limit)
375
+ states = _chunk_state_fwd(B, x, dt, dA_cumsum, seq_idx=seq_idx, states_in_fp32=True)
376
+ # states_tmp0 = _chunk_state_fwd(B[:, :147], x[:, :147], dt_tmp0, dA_cumsum_tmp0, states_in_fp32=True)
377
+ # states_tmp1 = _chunk_state_fwd(B[:, 147:], x[:, 147:], dt_tmp1, dA_cumsum_tmp1, states_in_fp32=True)
378
+ # states_tmp2 = _chunk_state_fwd(B[:, 147:256], x[:, 147:256], dt_tmp2, dA_cumsum_tmp2, states_in_fp32=True)
379
+ states, final_states = _state_passing_fwd(rearrange(states, "... p n -> ... (p n)"), dA_cumsum[:, :, :, -1],
380
+ initial_states=rearrange(initial_states, "... p n -> ... (p n)") if initial_states is not None else None,
381
+ seq_idx=seq_idx, chunk_size=chunk_size, out_dtype=C.dtype)
382
+ states, final_states = [rearrange(t, "... (p n) -> ... p n", n=dstate) for t in [states, final_states]]
383
+ # states_tmp0 = rearrange(_state_passing_fwd(rearrange(states_tmp0, "... p n -> ... (p n)"), dA_cumsum_tmp0[:, :, :, -1], chunk_size=chunk_size), "... (p n) -> ... p n", n=dstate)
384
+ # states_tmp1 = rearrange(_state_passing_fwd(rearrange(states_tmp1, "... p n -> ... (p n)"), dA_cumsum_tmp1[:, :, :, -1], chunk_size=chunk_size), "... (p n) -> ... p n", n=dstate)
385
+ CB = _bmm_chunk_fwd(C, B, chunk_size, seq_idx=seq_idx, output_dtype=torch.float32)
386
+ out, out_x = _chunk_scan_fwd(CB, x, dt, dA_cumsum, C, states, D=D, z=z, seq_idx=seq_idx)
387
+ if cu_seqlens is None:
388
+ return out, out_x, dt, dA_cumsum, states, final_states
389
+ else:
390
+ assert batch == 1, "passing cu_seqlens to get the varlen states is only supported if batch dimension is 1"
391
+ varlen_states = chunk_state_varlen(B.squeeze(0), x.squeeze(0), dt.squeeze(0), dA_cumsum.squeeze(0),
392
+ cu_seqlens, states.squeeze(0))
393
+ return out, out_x, dt, dA_cumsum, states, final_states, varlen_states
394
+
395
+
396
+ def _mamba_chunk_scan_combined_bwd(dout, x, dt, A, B, C, out, chunk_size, D=None, z=None,
397
+ dt_bias=None, initial_states=None, dfinal_states=None, seq_idx=None, dt_softplus=False,
398
+ dt_limit=(0.0, float("inf")),
399
+ dx=None, ddt=None, dB=None, dC=None, dz=None, recompute_output=False):
400
+ if dout.stride(-1) != 1:
401
+ dout = dout.contiguous()
402
+ batch, seqlen, nheads, headdim = x.shape
403
+ nchunks = math.ceil(seqlen / chunk_size)
404
+ _, _, ngroups, dstate = B.shape
405
+ assert dout.shape == (batch, seqlen, nheads, headdim)
406
+ assert dt.shape == (batch, seqlen, nheads)
407
+ assert A.shape == (nheads,)
408
+ assert nheads % ngroups == 0
409
+ assert B.shape == (batch, seqlen, ngroups, dstate)
410
+ assert C.shape == B.shape
411
+ assert out.shape == x.shape
412
+ if initial_states is not None:
413
+ assert initial_states.shape == (batch, nheads, headdim, dstate)
414
+ if seq_idx is not None:
415
+ assert seq_idx.shape == (batch, seqlen)
416
+ if dx is not None:
417
+ assert dx.shape == x.shape
418
+ if dB is not None:
419
+ assert dB.shape == B.shape
420
+ dB_given = dB
421
+ else:
422
+ dB_given = torch.empty_like(B)
423
+ if dC is not None:
424
+ assert dC.shape == C.shape
425
+ dC_given = dC
426
+ else:
427
+ dC_given = torch.empty_like(C)
428
+ if dz is not None:
429
+ assert z is not None
430
+ assert dz.shape == z.shape
431
+ if ddt is not None:
432
+ assert ddt.shape == dt.shape
433
+ ddt_given = ddt
434
+ else:
435
+ ddt_given = torch.empty_like(dt)
436
+ # TD: For some reason Triton (2.1.0 and 2.2.0) errors with
437
+ # "[CUDA]: invalid device context" (e.g. during varlne test), and cloning makes it work. Idk why.
438
+ dt_in = dt.clone()
439
+ dA_cumsum, dt = _chunk_cumsum_fwd(dt_in, A, chunk_size, dt_bias=dt_bias, dt_softplus=dt_softplus,
440
+ dt_limit=dt_limit)
441
+ CB = _bmm_chunk_fwd(C, B, chunk_size, seq_idx=seq_idx, output_dtype=torch.float32)
442
+ states = _chunk_state_fwd(B, x, dt, dA_cumsum, seq_idx=seq_idx, states_in_fp32=True)
443
+ states, _ = _state_passing_fwd(rearrange(states, "... p n -> ... (p n)"), dA_cumsum[:, :, :, -1],
444
+ initial_states=rearrange(initial_states, "... p n -> ... (p n)") if initial_states is not None else None,
445
+ seq_idx=seq_idx, chunk_size=chunk_size)
446
+ states = rearrange(states, "... (p n) -> ... p n", n=dstate)
447
+ if z is not None:
448
+ dz, dout, dD, *rest = _chunk_scan_bwd_dz(x, z, out, dout, chunk_size=chunk_size, has_ddAcs=False, D=D, dz=dz, recompute_output=recompute_output)
449
+ outz = rest[0] if recompute_output else out
450
+ else:
451
+ dz = None
452
+ outz = out
453
+ dstates = _chunk_scan_bwd_dstates(C, dA_cumsum, dout, seq_idx=seq_idx, dtype=states.dtype)
454
+ # dstates has length nchunks, containing the gradient to initial states at index 0 and
455
+ # gradient to the states of chunk (nchunks - 2) at index (nchunks - 1)
456
+ # Do computation in fp32 but convert dstates and states to fp16/bf16 since dstates and states
457
+ # will be used in matmul in the next kernels.
458
+ dstates, ddA_chunk_cumsum, dinitial_states, states = _state_passing_bwd(
459
+ rearrange(states, "... p n -> ... (p n)"),
460
+ dA_cumsum[:, :, :, -1],
461
+ rearrange(dstates, "... p n -> ... (p n)"),
462
+ dfinal_states=rearrange(dfinal_states, "... p n -> ... (p n)") if dfinal_states is not None else None,
463
+ seq_idx=seq_idx,
464
+ has_initial_states=initial_states is not None,
465
+ dstates_dtype=x.dtype,
466
+ states_dtype=x.dtype,
467
+ chunk_size=chunk_size,
468
+ )
469
+ # dstates has length nchunks, containing the gradient to states of chunk 0 at index 0 and
470
+ # gradient to the final states at index (nchunks - 1)
471
+ # states has length nchunks, containing the initial states at index 0 and the state for chunk (nchunks - 2) at index (nchunks - 1)
472
+ # The final states is not stored.
473
+ states = rearrange(states, "... (p n) -> ... p n", n=dstate)
474
+ dstates = rearrange(dstates, "... (p n) -> ... p n", n=dstate)
475
+ dinitial_states = rearrange(dinitial_states, "... (p n) -> ... p n", n=dstate) if dinitial_states is not None else None
476
+ dx, ddt, dD_from_x = _chunk_scan_chunk_state_bwd_dx(x, dt, dA_cumsum, B, CB, dout, dstates, D=D, seq_idx=seq_idx, dx=dx)
477
+ # dB = _chunk_state_bwd_db(x, dt, dA_cumsum, dstates, seq_idx=seq_idx, ngroups=ngroups)
478
+ dB, ddA_next = _chunk_state_bwd_db(x, dt, dA_cumsum, dstates, seq_idx=seq_idx, B=B, ngroups=ngroups)
479
+ # dC = _chunk_scan_bwd_dC(states[:, :-1].to(x.dtype), dA_cumsum, dout, seq_idx=seq_idx, ngroups=ngroups)
480
+ dC, ddA_cumsum_prev = _chunk_scan_bwd_dC(states.to(x.dtype), dA_cumsum, dout, seq_idx=seq_idx, C=C, ngroups=ngroups)
481
+ # Computing ddA with the dcb kernel is much slower, so we're not using it for now
482
+ dCB = _chunk_scan_bwd_dcb(x, dt, dA_cumsum, dout, seq_idx=seq_idx, ngroups=ngroups)
483
+ # dCB, ddA_tmp = _chunk_scan_bwd_dcb(x, dt, dA_cumsum, dout, seq_idx=seq_idx, CB=CB, ngroups=ngroups)
484
+ dCB = dCB.to(CB.dtype)
485
+ _bmm_chunk_bwd(C, dCB, residual=dB, out=dB_given)
486
+ _bmm_chunk_bwd(B, rearrange(dCB, "... l s -> ... s l"), residual=dC, out=dC_given)
487
+ # If we have z, then dout_x is recomputed in fp32 so dD = (dout_x * x).sum() is more accurate
488
+ # than dD_from_x = (dout_x * x).sum() where dout_x is in fp16/bf16
489
+ if z is None:
490
+ dD = dD_from_x
491
+ # Formula for ddA_cumsum, assuming out is the output of the forward pass before adding x * D.
492
+ # ddA_cumsum = torch.einsum("bclhp,bclhp->bhcl", out.float(), dout.float()) - ddt * dt
493
+ # However, this is numerically unstable: when we do the reverse cumsum on ddA_cumsum, there might
494
+ # be a lot of underflow.
495
+
496
+ # This is already done as part of bwd_dC kernel
497
+ # ddA_cumsum_prev = _chunk_scan_bwd_ddAcs_prev(states[:, :-1], C, dout, dA_cumsum, seq_idx=seq_idx)
498
+ ddA_cumsum_prev[..., -1] += ddA_chunk_cumsum
499
+ ddA_prev = ddA_cumsum_prev.flip([-1]).cumsum(dim=-1).flip([-1])
500
+ # This is already done as part of bwd_dB kernel
501
+ # ddA_next = _chunk_state_bwd_ddAcs_stable(B, x, dt, dA_cumsum, dstates, seq_idx=seq_idx)
502
+ # We don't need to pass in seq_idx because CB also zeros out entries where seq_idx[i] != seq_idx[j]
503
+ ddA = _chunk_scan_bwd_ddAcs_stable(x, dt, dA_cumsum, dout, CB)
504
+ ddA += ddA_next + ddA_prev
505
+
506
+ ddt_given, dA, ddt_bias = _chunk_cumsum_bwd(ddA, ddt, dt_in, A, dt_bias=dt_bias, dt_softplus=dt_softplus, dt_limit=dt_limit, ddt=ddt_given)
507
+
508
+ # These 2 lines are just to test ddt and dA being computed by old code
509
+ # _, dA = selective_scan_bwd(dout, x, dt, A, B, C, D=D.float(), z=z)
510
+ # ddt_given.copy_(ddt)
511
+
512
+ return_vals = (dx, ddt_given, dA, dB_given, dC_given, dD, dz, ddt_bias, dinitial_states)
513
+ return return_vals if not recompute_output else (*return_vals, outz)
514
+
515
+
516
+ def selective_scan_bwd(dout, x, dt, A, B, C, D=None, z=None):
517
+ """
518
+ Argument:
519
+ dout: (batch, seqlen, nheads, headdim)
520
+ x: (batch, seqlen, nheads, headdim)
521
+ dt: (batch, nheads, nchunks, chunk_size) or (batch, nheads, headdim, nchunks, chunk_size)
522
+ A: (nheads) or (dim, dstate)
523
+ B: (batch, seqlen, ngroups, dstate)
524
+ C: (batch, seqlen, ngroups, dstate)
525
+ D: (nheads, headdim) or (nheads,)
526
+ z: (batch, seqlen, nheads, headdim)
527
+ Return:
528
+ out: (batch, seqlen, nheads, headdim)
529
+ """
530
+ import selective_scan
531
+
532
+ batch, seqlen, nheads, headdim = x.shape
533
+ chunk_size = dt.shape[-1]
534
+ _, _, ngroups, dstate = B.shape
535
+ assert nheads % ngroups == 0
536
+ x = rearrange(x, "b l h p -> b (h p) l")
537
+ squeeze_dt = dt.dim() == 4
538
+ if dt.dim() == 4:
539
+ dt = repeat(dt, "b h c l -> b h p c l", p=headdim)
540
+ dt = rearrange(dt, "b h p c l -> b (h p) (c l)", p=headdim)
541
+ squeeze_A = A.dim() == 1
542
+ if A.dim() == 1:
543
+ A = repeat(A, "h -> (h p) n", p=headdim, n=dstate).to(dtype=torch.float32)
544
+ else:
545
+ A = A.to(dtype=torch.float32)
546
+ B = rearrange(B, "b l g n -> b g n l")
547
+ C = rearrange(C, "b l g n -> b g n l")
548
+ if D is not None:
549
+ if D.dim() == 2:
550
+ D = rearrange(D, "h p -> (h p)")
551
+ else:
552
+ D = repeat(D, "h -> (h p)", p=headdim)
553
+ if z is not None:
554
+ z = rearrange(z, "b l h p -> b (h p) l")
555
+
556
+ if x.stride(-1) != 1:
557
+ x = x.contiguous()
558
+ if dt.stride(-1) != 1:
559
+ dt = dt.contiguous()
560
+ if D is not None:
561
+ D = D.contiguous()
562
+ if B.stride(-1) != 1:
563
+ B = B.contiguous()
564
+ if C.stride(-1) != 1:
565
+ C = C.contiguous()
566
+ if z is not None and z.stride(-1) != 1:
567
+ z = z.contiguous()
568
+ _, intermediate, *rest = selective_scan.fwd(x, dt.to(dtype=x.dtype), A, B, C, D, z, None, False)
569
+ if z is not None:
570
+ out = rest[0]
571
+ else:
572
+ out = None
573
+
574
+ dout = rearrange(dout, "b l h p -> b (h p) l")
575
+
576
+ if dout.stride(-1) != 1:
577
+ dout = dout.contiguous()
578
+ # The kernel supports passing in a pre-allocated dz (e.g., in case we want to fuse the
579
+ # backward of selective_scan with the backward of chunk).
580
+ # Here we just pass in None and dz will be allocated in the C++ code.
581
+ _, ddt, dA, *rest = selective_scan.bwd(
582
+ x, dt.to(dtype=x.dtype), A, B, C, D, z, None, dout, intermediate, out, None, False,
583
+ False # option to recompute out_z, not used here
584
+ )
585
+ ddt = rearrange(ddt, "b (h p) (c l) -> b h p c l", p=headdim, l=chunk_size)
586
+ if squeeze_dt:
587
+ ddt = ddt.float().sum(dim=2)
588
+ if squeeze_A:
589
+ dA = rearrange(dA, "(h p) n -> h p n", p=headdim).sum(dim=(1, 2))
590
+ return ddt, dA
591
+
592
+
593
+ class MambaChunkScanCombinedFn(torch.autograd.Function):
594
+
595
+ @staticmethod
596
+ def forward(ctx, x, dt, A, B, C, chunk_size, D=None, z=None, dt_bias=None, initial_states=None, seq_idx=None, cu_seqlens=None, dt_softplus=False, dt_limit=(0.0, float("inf")), return_final_states=False, return_varlen_states=False):
597
+ ctx.dt_dtype = dt.dtype
598
+ if not return_varlen_states:
599
+ cu_seqlens = None
600
+ else:
601
+ assert cu_seqlens is not None, "cu_seqlens must be provided if return_varlen_states is True"
602
+ out, out_x, dt_out, dA_cumsum, states, final_states, *rest = _mamba_chunk_scan_combined_fwd(x, dt, A, B, C, chunk_size, D=D, z=z, dt_bias=dt_bias, initial_states=initial_states, seq_idx=seq_idx, cu_seqlens=cu_seqlens, dt_softplus=dt_softplus, dt_limit=dt_limit)
603
+ ctx.save_for_backward(out if z is None else out_x, x, dt, dA_cumsum, A, B, C, D, z, dt_bias, initial_states, seq_idx)
604
+ ctx.dt_softplus = dt_softplus
605
+ ctx.chunk_size = chunk_size
606
+ ctx.dt_limit = dt_limit
607
+ ctx.return_final_states = return_final_states
608
+ ctx.return_varlen_states = return_varlen_states
609
+ if not return_varlen_states:
610
+ return out if not return_final_states else (out, final_states)
611
+ else:
612
+ varlen_states = rest[0]
613
+ return (out, varlen_states) if not return_final_states else (out, final_states, varlen_states)
614
+
615
+ @staticmethod
616
+ def backward(ctx, dout, *args):
617
+ out, x, dt, dA_cumsum, A, B, C, D, z, dt_bias, initial_states, seq_idx = ctx.saved_tensors
618
+ assert not ctx.return_varlen_states, "return_varlen_states is not supported in backward"
619
+ dfinal_states = args[0] if ctx.return_final_states else None
620
+ dx, ddt, dA, dB, dC, dD, dz, ddt_bias, dinitial_states = _mamba_chunk_scan_combined_bwd(dout, x, dt, A, B, C, out, ctx.chunk_size, D=D, z=z, dt_bias=dt_bias, initial_states=initial_states, dfinal_states=dfinal_states, seq_idx=seq_idx, dt_softplus=ctx.dt_softplus, dt_limit=ctx.dt_limit)
621
+ return dx, ddt, dA, dB, dC, None, dD, dz, ddt_bias, dinitial_states, None, None, None, None, None, None
622
+
623
+
624
+ def mamba_chunk_scan_combined(x, dt, A, B, C, chunk_size, D=None, z=None, dt_bias=None, initial_states=None, seq_idx=None, cu_seqlens=None, dt_softplus=False, dt_limit=(0.0, float("inf")), return_final_states=False, return_varlen_states=False):
625
+ """
626
+ Argument:
627
+ x: (batch, seqlen, nheads, headdim)
628
+ dt: (batch, seqlen, nheads)
629
+ A: (nheads)
630
+ B: (batch, seqlen, ngroups, dstate)
631
+ C: (batch, seqlen, ngroups, dstate)
632
+ chunk_size: int
633
+ D: (nheads, headdim) or (nheads,)
634
+ z: (batch, seqlen, nheads, headdim)
635
+ dt_bias: (nheads,)
636
+ initial_states: (batch, nheads, headdim, dstate)
637
+ seq_idx: (batch, seqlen)
638
+ cu_seqlens: (num_sequences + 1) or None, only used if return_varlen_states is True
639
+ dt_softplus: Whether to apply softplus to dt
640
+ Return:
641
+ out: (batch, seqlen, nheads, headdim)
642
+ """
643
+ return MambaChunkScanCombinedFn.apply(x, dt, A, B, C, chunk_size, D, z, dt_bias, initial_states, seq_idx, cu_seqlens, dt_softplus, dt_limit, return_final_states, return_varlen_states)
644
+
645
+
646
+ def mamba_chunk_scan(x, dt, A, B, C, chunk_size, D=None, z=None, dt_bias=None, dt_softplus=False):
647
+ """
648
+ Argument:
649
+ x: (batch, seqlen, nheads, headdim)
650
+ dt: (batch, seqlen, nheads)
651
+ A: (nheads)
652
+ B: (batch, seqlen, ngroups, dstate)
653
+ C: (batch, seqlen, ngroups, dstate)
654
+ D: (nheads, headdim) or (nheads,)
655
+ z: (batch, seqlen, nheads, headdim)
656
+ dt_bias: (nheads,)
657
+ Return:
658
+ out: (batch, seqlen, nheads, headdim)
659
+ """
660
+ batch, seqlen, nheads, headdim = x.shape
661
+ dstate = B.shape[-1]
662
+ if seqlen % chunk_size != 0:
663
+ dt = F.pad(dt, (0, 0, 0, chunk_size - seqlen % chunk_size))
664
+ dt = rearrange(dt, "b (c l) h -> b h c l", l=chunk_size)
665
+ dt = dt.float() # We want high precision for this before cumsum
666
+ if dt_bias is not None:
667
+ dt = dt + rearrange(dt_bias, "h -> h 1 1")
668
+ if dt_softplus:
669
+ dt = F.softplus(dt)
670
+ dA = dt * rearrange(A, "h -> h 1 1")
671
+ dA_cumsum = torch.cumsum(dA, dim=-1)
672
+ # 1. Compute the state for each chunk
673
+ states = chunk_state(B, x, dt, dA_cumsum, states_in_fp32=True)
674
+ # 2. Pass the state to all the chunks by weighted cumsum.
675
+ states = rearrange(state_passing(rearrange(states, "... p n -> ... (p n)"), dA_cumsum[:, :, :, -1])[0],
676
+ "... (p n) -> ... p n", n=dstate)
677
+ # 3. Compute the output for each chunk
678
+ out = chunk_scan(B, C, x, dt, dA_cumsum, states, D=D, z=z)
679
+ return out
680
+
681
+
682
+ def ssd_chunk_scan_combined_ref(x, dt, A, B, C, chunk_size, D=None, z=None, dt_bias=None, dt_softplus=False):
683
+ """
684
+ Argument:
685
+ x: (batch, seqlen, nheads, headdim)
686
+ dt: (batch, seqlen, nheads)
687
+ A: (nheads)
688
+ B: (batch, seqlen, ngroups, dstate)
689
+ C: (batch, seqlen, ngroups, dstate)
690
+ D: (nheads, headdim) or (nheads,)
691
+ z: (batch, seqlen, nheads, headdim)
692
+ dt_bias: (nheads,)
693
+ Return:
694
+ out: (batch, seqlen, nheads, headdim)
695
+ """
696
+ batch, seqlen, nheads, headdim = x.shape
697
+ dstate = B.shape[-1]
698
+ if seqlen % chunk_size != 0:
699
+ dt = F.pad(dt, (0, 0, 0, chunk_size - seqlen % chunk_size))
700
+ dt = rearrange(dt, "b (c l) h -> b h c l", l=chunk_size)
701
+ dt = dt.float() # We want high precision for this before cumsum
702
+ if dt_bias is not None:
703
+ dt = dt + rearrange(dt_bias, "h -> h 1 1")
704
+ if dt_softplus:
705
+ dt = F.softplus(dt)
706
+ dA = dt * rearrange(A, "h -> h 1 1")
707
+ dA_cumsum = torch.cumsum(dA, dim=-1)
708
+ # 1. Compute the state for each chunk
709
+ states = chunk_state_ref(B, x, dt, dA_cumsum)
710
+ states_dtype = states.dtype
711
+ if states.dtype not in [torch.float32, torch.float64]:
712
+ states = states.to(torch.float32)
713
+ # 2. Pass the state to all the chunks by weighted cumsum.
714
+ # state_passing_ref is much less numerically stable
715
+ states = rearrange(state_passing_ref(rearrange(states, "... p n -> ... (p n)"), dA_cumsum[:, :, :, -1])[0],
716
+ "... (p n) -> ... p n", n=dstate)
717
+ states = states.to(states_dtype)
718
+ # 3. Compute the output for each chunk
719
+ out = chunk_scan_ref(B, C, x, dt, dA_cumsum, states, D=D, z=z)
720
+ return out
721
+
722
+
723
+ def ssd_selective_scan(x, dt, A, B, C, D=None, z=None, dt_bias=None, dt_softplus=False, dt_limit=(0.0, float("inf"))):
724
+ """
725
+ Argument:
726
+ x: (batch, seqlen, nheads, headdim)
727
+ dt: (batch, seqlen, nheads) or (batch, seqlen, nheads, headdim)
728
+ A: (nheads) or (dim, dstate)
729
+ B: (batch, seqlen, ngroups, dstate)
730
+ C: (batch, seqlen, ngroups, dstate)
731
+ D: (nheads, headdim) or (nheads,)
732
+ z: (batch, seqlen, nheads, headdim)
733
+ dt_bias: (nheads,) or (nheads, headdim)
734
+ Return:
735
+ out: (batch, seqlen, nheads, headdim)
736
+ """
737
+ from mamba_ssm.ops.selective_scan_interface import selective_scan_fn
738
+
739
+ batch, seqlen, nheads, headdim = x.shape
740
+ _, _, ngroups, dstate = B.shape
741
+ x = rearrange(x, "b l h p -> b (h p) l")
742
+ if dt.dim() == 3:
743
+ dt = repeat(dt, "b l h -> b l h p", p=headdim)
744
+ dt = rearrange(dt, "b l h p -> b (h p) l")
745
+ if A.dim() == 1:
746
+ A = repeat(A, "h -> (h p) n", p=headdim, n=dstate).to(dtype=torch.float32)
747
+ else:
748
+ A = A.to(dtype=torch.float32)
749
+ B = rearrange(B, "b l g n -> b g n l")
750
+ C = rearrange(C, "b l g n -> b g n l")
751
+ if D is not None:
752
+ if D.dim() == 2:
753
+ D = rearrange(D, "h p -> (h p)")
754
+ else:
755
+ D = repeat(D, "h -> (h p)", p=headdim)
756
+ if z is not None:
757
+ z = rearrange(z, "b l h p -> b (h p) l")
758
+ if dt_bias is not None:
759
+ if dt_bias.dim() == 1:
760
+ dt_bias = repeat(dt_bias, "h -> h p", p=headdim)
761
+ dt_bias = rearrange(dt_bias, "h p -> (h p)")
762
+ if dt_limit != (0.0, float("inf")):
763
+ if dt_bias is not None:
764
+ dt = dt + rearrange(dt_bias, "d -> d 1")
765
+ if dt_softplus:
766
+ dt = F.softplus(dt)
767
+ dt = dt.clamp(min=dt_limit[0], max=dt_limit[1]).to(x.dtype)
768
+ dt_bias = None
769
+ dt_softplus = None
770
+ out = selective_scan_fn(x, dt, A, B, C, D=D, z=z, delta_bias=dt_bias, delta_softplus=dt_softplus)
771
+ return rearrange(out, "b (h p) l -> b l h p", p=headdim)
772
+
773
+
774
+ def mamba_conv1d_scan_ref(xBC, conv1d_weight, conv1d_bias, dt, A, chunk_size, D=None, z=None,
775
+ dt_bias=None, dt_softplus=False, dt_limit=(0.0, float("inf")),
776
+ activation="silu", headdim=None, ngroups=1):
777
+ """
778
+ Argument:
779
+ xBC: (batch, seqlen, dim + 2 * ngroups * dstate) where dim == nheads * headdim
780
+ conv1d_weight: (dim + 2 * ngroups * dstate, width)
781
+ conv1d_bias: (dim + 2 * ngroups * dstate,)
782
+ dt: (batch, seqlen, nheads) or (batch, seqlen, nheads, headdim)
783
+ A: (nheads)
784
+ D: (nheads, headdim) or (nheads,)
785
+ z: (batch, seqlen, dim)
786
+ dt_bias: (nheads) or (nheads, headdim)
787
+ headdim: if D is 1D and z is None, headdim must be passed in
788
+ Return:
789
+ out: (batch, seqlen, dim)
790
+ """
791
+ batch, seqlen, nheads = dt.shape[:3]
792
+ assert nheads % ngroups == 0
793
+ if z is not None:
794
+ dim = z.shape[-1]
795
+ assert dim % nheads == 0
796
+ headdim = dim // nheads
797
+ else:
798
+ if D.dim() == 1:
799
+ assert headdim is not None
800
+ else:
801
+ headdim = D.shape[1]
802
+ dim = nheads * headdim
803
+ xBC = rearrange(causal_conv1d_fn(rearrange(xBC, "b s d -> b d s"), conv1d_weight, conv1d_bias, activation=activation),
804
+ "b d s -> b s d")
805
+ dstate = (xBC.shape[-1] - dim) // ngroups // 2
806
+ x, B, C = torch.split(xBC, [dim, ngroups * dstate, ngroups * dstate], dim=-1)
807
+ x = rearrange(x, "b l (h p) -> b l h p", h=nheads)
808
+ B = rearrange(B, "b l (g n) -> b l g n", g=ngroups)
809
+ C = rearrange(C, "b l (g n) -> b l g n", g=ngroups)
810
+ z = rearrange(z, "b l (h p) -> b l h p", h=nheads) if z is not None else None
811
+ out = ssd_selective_scan(x, dt.to(x.dtype), A, B, C, D=D.float(), z=z, dt_bias=dt_bias, dt_softplus=dt_softplus, dt_limit=dt_limit)
812
+ return rearrange(out, "b s h p -> b s (h p)")
813
+
814
+
815
+ class MambaSplitConv1dScanCombinedFn(torch.autograd.Function):
816
+
817
+ @staticmethod
818
+ @custom_fwd
819
+ def forward(ctx, zxbcdt, conv1d_weight, conv1d_bias, dt_bias, A, D, chunk_size, initial_states=None, seq_idx=None, dt_limit=(0.0, float("inf")), return_final_states=False, activation="silu",
820
+ rmsnorm_weight=None, rmsnorm_eps=1e-6, outproj_weight=None, outproj_bias=None, headdim=None,
821
+ ngroups=1, norm_before_gate=True):
822
+ assert activation in [None, "silu", "swish"]
823
+ if D.dim() == 1:
824
+ assert headdim is not None
825
+ nheads, = D.shape
826
+ else:
827
+ nheads, headdim = D.shape
828
+ batch, seqlen, _ = zxbcdt.shape
829
+ dim = nheads * headdim
830
+ assert nheads % ngroups == 0
831
+ dstate = (conv1d_weight.shape[0] - dim) // ngroups // 2
832
+ d_nonssm = (zxbcdt.shape[-1] - 2 * dim - 2 * ngroups * dstate - nheads) // 2
833
+ assert d_nonssm >= 0
834
+ assert zxbcdt.shape == (batch, seqlen, 2 * d_nonssm + 2 * dim + 2 * ngroups * dstate + nheads)
835
+ assert dt_bias.shape == (nheads,)
836
+ assert A.shape == (nheads,)
837
+ zx0, z, xBC, dt = torch.split(zxbcdt, [2 * d_nonssm, dim, dim + ngroups * dstate * 2, nheads], dim=-1)
838
+ seq_idx = seq_idx.contiguous() if seq_idx is not None else None
839
+ xBC_conv = rearrange(
840
+ causal_conv1d_fwd_function(rearrange(ensure_stride(xBC), "b s d -> b d s"),
841
+ conv1d_weight, conv1d_bias, seq_idx, None, None, activation in ["silu", "swish"]),
842
+ "b d s -> b s d"
843
+ )
844
+ x, B, C = torch.split(xBC_conv, [dim, ngroups * dstate, ngroups * dstate], dim=-1)
845
+ x = rearrange(x, "b l (h p) -> b l h p", h=nheads)
846
+ B = rearrange(B, "b l (g n) -> b l g n", g=ngroups)
847
+ C = rearrange(C, "b l (g n) -> b l g n", g=ngroups)
848
+ z = rearrange(z, "b l (h p) -> b l h p", h=nheads) if z is not None else None
849
+ if rmsnorm_weight is None:
850
+ out, out_x, dt_out, dA_cumsum, states, final_states = _mamba_chunk_scan_combined_fwd(x, dt, A, B, C, chunk_size=chunk_size, D=D, z=z, dt_bias=dt_bias, initial_states=initial_states, seq_idx=seq_idx, dt_softplus=True, dt_limit=dt_limit)
851
+ out = rearrange(out, "b s h p -> b s (h p)")
852
+ rstd = None
853
+ if d_nonssm > 0:
854
+ out = torch.cat([_swiglu_fwd(zx0), out], dim=-1)
855
+ else:
856
+ out_x, _, dt_out, dA_cumsum, states, final_states = _mamba_chunk_scan_combined_fwd(x, dt, A, B, C, chunk_size=chunk_size, D=D, z=None, dt_bias=dt_bias, initial_states=initial_states, seq_idx=seq_idx, dt_softplus=True, dt_limit=dt_limit)
857
+ # reshape input data into 2D tensor
858
+ x_rms = rearrange(out_x, "b s h p -> (b s) (h p)")
859
+ z_rms = rearrange(z, "b s h p -> (b s) (h p)")
860
+ rmsnorm_weight = rmsnorm_weight.contiguous()
861
+ if d_nonssm == 0:
862
+ out = None
863
+ else:
864
+ out01 = torch.empty((batch, seqlen, d_nonssm + dim), dtype=x_rms.dtype, device=x_rms.device)
865
+ out = rearrange(out01[..., d_nonssm:], "b s d -> (b s) d")
866
+ _swiglu_fwd(zx0, out=out01[..., :d_nonssm])
867
+ out, _, rstd = _layer_norm_fwd(x_rms, rmsnorm_weight, None, rmsnorm_eps, z_rms, out=out,
868
+ group_size=dim // ngroups,
869
+ norm_before_gate=norm_before_gate, is_rms_norm=True)
870
+ if d_nonssm == 0:
871
+ out = rearrange(out, "(b s) d -> b s d", b=batch)
872
+ else:
873
+ out = out01
874
+ ctx.outproj_weight_dtype = outproj_weight.dtype if outproj_weight is not None else None
875
+ if outproj_weight is not None:
876
+ if torch.is_autocast_enabled():
877
+ dtype = torch.get_autocast_gpu_dtype()
878
+ out, outproj_weight = out.to(dtype), outproj_weight.to(dtype)
879
+ outproj_bias = outproj_bias.to(dtype) if outproj_bias is not None else None
880
+ out = F.linear(out, outproj_weight, outproj_bias)
881
+ else:
882
+ assert outproj_bias is None
883
+ ctx.save_for_backward(zxbcdt, conv1d_weight, conv1d_bias,
884
+ out_x, A, D, dt_bias, initial_states, seq_idx, rmsnorm_weight, rstd, outproj_weight, outproj_bias)
885
+ ctx.dt_limit = dt_limit
886
+ ctx.return_final_states = return_final_states
887
+ ctx.activation = activation
888
+ ctx.rmsnorm_eps = rmsnorm_eps
889
+ ctx.norm_before_gate = norm_before_gate
890
+ ctx.chunk_size = chunk_size
891
+ ctx.headdim = headdim
892
+ ctx.ngroups = ngroups
893
+ return out if not return_final_states else (out, final_states)
894
+
895
+ @staticmethod
896
+ @custom_bwd
897
+ def backward(ctx, dout, *args):
898
+ zxbcdt, conv1d_weight, conv1d_bias, out, A, D, dt_bias, initial_states, seq_idx, rmsnorm_weight, rstd, outproj_weight, outproj_bias = ctx.saved_tensors
899
+ dfinal_states = args[0] if ctx.return_final_states else None
900
+ headdim = ctx.headdim
901
+ nheads = D.shape[0]
902
+ dim = nheads * headdim
903
+ assert nheads % ctx.ngroups == 0
904
+ dstate = (conv1d_weight.shape[0] - dim) // ctx.ngroups // 2
905
+ d_nonssm = (zxbcdt.shape[-1] - 2 * dim - 2 * ctx.ngroups * dstate - nheads) // 2
906
+ assert d_nonssm >= 0
907
+ recompute_output = outproj_weight is not None
908
+ if recompute_output:
909
+ out_recompute = torch.empty(*out.shape[:2], d_nonssm + dim, device=out.device, dtype=out.dtype)
910
+ out0_recompute, out1_recompute = out_recompute.split([d_nonssm, dim], dim=-1)
911
+ zx0, z, xBC, dt = torch.split(zxbcdt, [2 * d_nonssm, dim, dim + 2 * ctx.ngroups * dstate, nheads], dim=-1)
912
+ # Recompute x, B, C
913
+ xBC_conv = rearrange(
914
+ causal_conv1d_fwd_function(rearrange(ensure_stride(xBC), "b s d -> b d s"),
915
+ conv1d_weight, conv1d_bias, seq_idx, None, None, ctx.activation in ["silu", "swish"]),
916
+ "b d s -> b s d"
917
+ )
918
+ x, B, C = torch.split(xBC_conv, [dim, ctx.ngroups * dstate, ctx.ngroups * dstate], dim=-1)
919
+ x = rearrange(x, "b l (h p) -> b l h p", h=nheads)
920
+ B = rearrange(B, "b l (g n) -> b l g n", g=ctx.ngroups)
921
+ C = rearrange(C, "b l (g n) -> b l g n", g=ctx.ngroups)
922
+ dzxbcdt = torch.empty_like(zxbcdt)
923
+ dzx0, dz, dxBC_given, ddt_given = torch.split(dzxbcdt, [2 * d_nonssm, dim, dim + 2 * ctx.ngroups * dstate, nheads], dim=-1)
924
+ dxBC = torch.empty_like(xBC)
925
+ dx, dB, dC = torch.split(dxBC, [dim, ctx.ngroups * dstate, ctx.ngroups * dstate], dim=-1)
926
+ z = rearrange(z, "b l (h p) -> b l h p", h=nheads)
927
+ dx = rearrange(dx, "b l (h p) -> b l h p", h=nheads)
928
+ dB = rearrange(dB, "b l (g n) -> b l g n", g=ctx.ngroups)
929
+ dC = rearrange(dC, "b l (g n) -> b l g n", g=ctx.ngroups)
930
+ if outproj_weight is not None:
931
+ dout_og = dout
932
+ dout = F.linear(dout, outproj_weight.t())
933
+ if d_nonssm > 0:
934
+ dout0, dout = dout.split([d_nonssm, dim], dim=-1)
935
+ _swiglu_bwd(zx0, dout0, dxy=dzx0, recompute_output=True, out=out0_recompute)
936
+ dout = rearrange(dout, "b s (h p) -> b s h p", p=headdim)
937
+ if rmsnorm_weight is None:
938
+ dz = rearrange(dz, "b l (h p) -> b l h p", h=nheads)
939
+ dx, ddt, dA, dB, dC, dD, dz, ddt_bias, dinitial_states, *rest = _mamba_chunk_scan_combined_bwd(
940
+ dout, x, dt, A, B, C, out, ctx.chunk_size, D=D, z=z, dt_bias=dt_bias, initial_states=initial_states, dfinal_states=dfinal_states, seq_idx=seq_idx, dt_softplus=True, dt_limit=ctx.dt_limit, dx=dx, ddt=ddt_given, dB=dB, dC=dC, dz=dz, recompute_output=recompute_output
941
+ )
942
+ out_for_linear = rearrange(rest[0], "b s h p -> b s (h p)") if recompute_output else None
943
+ drmsnorm_weight = None
944
+ else:
945
+ batch = dout.shape[0]
946
+ dy_rms = rearrange(dout, "b s h p -> (b s) (h p)")
947
+ dz = rearrange(dz, "b l d -> (b l) d")
948
+ x_rms = rearrange(out, "b s h p -> (b s) (h p)")
949
+ z_rms = rearrange(z, "b s h p -> (b s) (h p)")
950
+ out1_recompute = rearrange(out1_recompute, "b s d -> (b s) d") if recompute_output else None
951
+ dout, drmsnorm_weight, _, dz, *rest = _layer_norm_bwd(dy_rms, x_rms, rmsnorm_weight, None, ctx.rmsnorm_eps, None, rstd, z_rms, group_size=dim//ctx.ngroups, norm_before_gate=ctx.norm_before_gate, is_rms_norm=True, recompute_output=recompute_output, dz=dz, out=out1_recompute if recompute_output else None)
952
+ out_for_linear = out_recompute if recompute_output else None
953
+ dout = rearrange(dout, "(b s) (h p) -> b s h p", b=batch, p=headdim)
954
+ dx, ddt, dA, dB, dC, dD, _, ddt_bias, dinitial_states = _mamba_chunk_scan_combined_bwd(
955
+ dout, x, dt, A, B, C, out, ctx.chunk_size, D=D, z=None, dt_bias=dt_bias, initial_states=initial_states, dfinal_states=dfinal_states, seq_idx=seq_idx, dt_softplus=True, dt_limit=ctx.dt_limit, dx=dx, ddt=ddt_given, dB=dB, dC=dC
956
+ )
957
+
958
+ if outproj_weight is not None:
959
+ doutproj_weight = torch.einsum("bso,bsd->od", dout_og, out_for_linear)
960
+ doutproj_bias = dout_og.sum(dim=(0, 1)) if outproj_bias is not None else None
961
+ else:
962
+ doutproj_weight, doutproj_bias = None, None
963
+ dxBC_given_update, dweight, dbias, *_ = causal_conv1d_bwd_function(
964
+ rearrange(ensure_stride(xBC), "b s d -> b d s"), conv1d_weight, conv1d_bias,
965
+ # It might be okay to not run ensure_stride on dxBC, but we're not sure. So playing safe here.
966
+ rearrange(ensure_stride(dxBC), "b s d -> b d s"), seq_idx, None, None,
967
+ rearrange(ensure_stride(dxBC_given), "b s d -> b d s"), False, ctx.activation in ["silu", "swish"]
968
+ )
969
+ dxBC_given_update = rearrange(dxBC_given_update, "b d s -> b s d")
970
+ if dxBC_given.stride() != dxBC_given_update.stride():
971
+ dxBC_given.copy_(dxBC_given_update)
972
+ else:
973
+ dxBC_given = dxBC_given_update
974
+ return dzxbcdt, dweight, dbias, ddt_bias, dA, dD, None, dinitial_states, None, None, None, None, drmsnorm_weight, None, doutproj_weight, doutproj_bias, None, None, None
975
+
976
+
977
+ def mamba_split_conv1d_scan_combined(zxbcdt, conv1d_weight, conv1d_bias, dt_bias, A, D, chunk_size, initial_states=None, seq_idx=None, dt_limit=(0.0, float("inf")), return_final_states=False, activation="silu", rmsnorm_weight=None, rmsnorm_eps=1e-6, outproj_weight=None, outproj_bias=None, headdim=None, ngroups=1, norm_before_gate=True):
978
+ """
979
+ Argument:
980
+ zxbcdt: (batch, seqlen, 2 * dim + 2 * ngroups * dstate + nheads) where dim == nheads * headdim
981
+ conv1d_weight: (dim + 2 * ngroups * dstate, width)
982
+ conv1d_bias: (dim + 2 * ngroups * dstate,)
983
+ dt_bias: (nheads,)
984
+ A: (nheads)
985
+ D: (nheads, headdim) or (nheads,)
986
+ initial_states: (batch, nheads, headdim, dstate)
987
+ seq_idx: (batch, seqlen), int32
988
+ rmsnorm_weight: (dim,)
989
+ outproj_weight: (out_dim, dim)
990
+ outproj_bias: (out_dim,)
991
+ headdim: if D is 1D, headdim must be passed in
992
+ norm_before_gate: if True, we do RMSNorm(x) * F.silu(z). If False, we do RMSNorm(x * F.silu(z))
993
+ Return:
994
+ out: (batch, seqlen, dim)
995
+ """
996
+ return MambaSplitConv1dScanCombinedFn.apply(zxbcdt, conv1d_weight, conv1d_bias, dt_bias, A, D, chunk_size, initial_states, seq_idx, dt_limit, return_final_states, activation, rmsnorm_weight, rmsnorm_eps, outproj_weight, outproj_bias, headdim, ngroups, norm_before_gate)
997
+
998
+
999
+ def mamba_split_conv1d_scan_ref(zxbcdt, conv1d_weight, conv1d_bias, dt_bias, A, D, chunk_size, dt_limit=(0.0, float("inf")), activation="silu", rmsnorm_weight=None, rmsnorm_eps=1e-6, outproj_weight=None, outproj_bias=None, headdim=None, ngroups=1, norm_before_gate=True):
1000
+ """
1001
+ Argument:
1002
+ zxbcdt: (batch, seqlen, 2 * dim + 2 * ngroups * dstate + nheads) where dim == nheads * headdim
1003
+ conv1d_weight: (dim + 2 * ngroups * dstate, width)
1004
+ conv1d_bias: (dim + 2 * ngroups * dstate,)
1005
+ dt_bias: (nheads,)
1006
+ A: (nheads)
1007
+ D: (nheads, headdim) or (nheads,)
1008
+ rmsnorm_weight: (dim,)
1009
+ outproj_weight: (out_dim, dim)
1010
+ outproj_bias: (out_dim,)
1011
+ headdim: if D is 1D, headdim must be passed in
1012
+ norm_before_gate: if True, we do RMSNorm(x) * F.silu(z). If False, we do RMSNorm(x * F.silu(z))
1013
+ Return:
1014
+ out: (batch, seqlen, dim)
1015
+ """
1016
+ if D.dim() == 1:
1017
+ assert headdim is not None
1018
+ nheads, = D.shape
1019
+ else:
1020
+ nheads, headdim = D.shape
1021
+ assert nheads % ngroups == 0
1022
+ batch, seqlen, _ = zxbcdt.shape
1023
+ dim = nheads * headdim
1024
+ dstate = (zxbcdt.shape[-1] - 2 * dim - nheads) // ngroups // 2
1025
+ assert zxbcdt.shape == (batch, seqlen, 2 * dim + 2 * ngroups * dstate + nheads)
1026
+ assert dt_bias.shape == (nheads,)
1027
+ assert A.shape == (nheads,)
1028
+ if rmsnorm_weight is not None:
1029
+ assert rmsnorm_weight.shape == (dim,)
1030
+ z, xBC, dt = torch.split(zxbcdt, [dim, dim + 2 * ngroups * dstate, nheads], dim=-1)
1031
+ xBC = rearrange(causal_conv1d_fn(rearrange(xBC, "b s d -> b d s"), conv1d_weight, conv1d_bias, activation=activation),
1032
+ "b d s -> b s d")
1033
+ x, B, C = torch.split(xBC, [dim, ngroups * dstate, ngroups * dstate], dim=-1)
1034
+ x = rearrange(x, "b l (h p) -> b l h p", h=nheads)
1035
+ B = rearrange(B, "b l (g n) -> b l g n", g=ngroups)
1036
+ C = rearrange(C, "b l (g n) -> b l g n", g=ngroups)
1037
+ z = rearrange(z, "b l (h p) -> b l h p", h=nheads)
1038
+ out = ssd_selective_scan(x, dt.to(x.dtype), A, B, C, D=D.float(),
1039
+ z=z if rmsnorm_weight is None else None, dt_bias=dt_bias, dt_softplus=True, dt_limit=dt_limit)
1040
+ out = rearrange(out, "b s h p -> b s (h p)")
1041
+ if rmsnorm_weight is not None:
1042
+ out = rmsnorm_fn(out, rmsnorm_weight, None, z=rearrange(z, "b l h p -> b l (h p)"), eps=rmsnorm_eps,
1043
+ norm_before_gate=norm_before_gate)
1044
+ if outproj_weight is not None:
1045
+ out = F.linear(out, outproj_weight, outproj_bias)
1046
+ return out
1047
+
mamba_ssm/ops/triton/ssd_state_passing.py ADDED
@@ -0,0 +1,350 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2024, Tri Dao, Albert Gu.
2
+
3
+ """We want triton==2.1.0 or 2.2.0 for this
4
+ """
5
+
6
+ import math
7
+ import torch
8
+ import torch.nn.functional as F
9
+
10
+ import triton
11
+ import triton.language as tl
12
+
13
+ from einops import rearrange, repeat
14
+
15
+ from mamba_ssm.utils.determinism import autotune_configs
16
+
17
+
18
+ @triton.autotune(
19
+ configs=autotune_configs([
20
+ triton.Config({'BLOCK_SIZE': 64}),
21
+ triton.Config({'BLOCK_SIZE': 128}),
22
+ triton.Config({'BLOCK_SIZE': 256}),
23
+ triton.Config({'BLOCK_SIZE': 512}),
24
+ triton.Config({'BLOCK_SIZE': 1024}),
25
+ triton.Config({'BLOCK_SIZE': 2048}),
26
+ ]),
27
+ key=['dim'],
28
+ )
29
+ @triton.jit
30
+ def _state_passing_fwd_kernel(
31
+ # Pointers to matrices
32
+ states_ptr, out_ptr, final_states_ptr, dA_cs_ptr, initstates_ptr, seq_idx_ptr,
33
+ # Matrix dimensions
34
+ dim, nchunks, seqlen, chunk_size,
35
+ # Strides
36
+ stride_states_batch, stride_states_chunk, stride_states_head, stride_states_dim,
37
+ stride_out_batch, stride_out_chunk, stride_out_head, stride_out_dim,
38
+ stride_final_states_batch, stride_final_states_head, stride_final_states_dim,
39
+ stride_dA_cs_batch, stride_dA_cs_chunk, stride_dA_cs_head,
40
+ stride_initstates_batch, stride_initstates_head, stride_initstates_dim,
41
+ stride_seq_idx_batch, stride_seq_idx_seqlen,
42
+ # Meta-parameters
43
+ HAS_INITSTATES: tl.constexpr,
44
+ HAS_SEQ_IDX: tl.constexpr,
45
+ BLOCK_SIZE: tl.constexpr,
46
+ ):
47
+ pid_b = tl.program_id(axis=1)
48
+ pid_h = tl.program_id(axis=2)
49
+ pid_m = tl.program_id(axis=0)
50
+ states_ptr += pid_b * stride_states_batch + pid_h * stride_states_head
51
+ dA_cs_ptr += pid_b * stride_dA_cs_batch + pid_h * stride_dA_cs_head
52
+ out_ptr += pid_b * stride_out_batch + pid_h * stride_out_head
53
+ final_states_ptr += pid_b * stride_final_states_batch + pid_h * stride_final_states_head
54
+ if HAS_INITSTATES:
55
+ initstates_ptr += pid_b * stride_initstates_batch + pid_h * stride_initstates_head
56
+ if HAS_SEQ_IDX:
57
+ seq_idx_ptr += pid_b * stride_seq_idx_batch
58
+
59
+ offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
60
+ states_ptrs = states_ptr + offs_m * stride_states_dim
61
+ out_ptrs = out_ptr + offs_m * stride_out_dim
62
+ final_states_ptrs = final_states_ptr + offs_m * stride_final_states_dim
63
+
64
+ if not HAS_INITSTATES:
65
+ states = tl.zeros((BLOCK_SIZE, ), dtype=tl.float32)
66
+ else:
67
+ initstates_ptrs = initstates_ptr + offs_m * stride_initstates_dim
68
+ states = tl.load(initstates_ptrs, mask=offs_m < dim, other=0.0).to(tl.float32)
69
+ tl.store(out_ptrs, states, mask=offs_m < dim)
70
+ out_ptrs += stride_out_chunk
71
+ seq_idx = 0
72
+ for c in range(nchunks):
73
+ new_states = tl.load(states_ptrs, mask=offs_m < dim, other=0.0).to(tl.float32)
74
+ dA_cs = tl.load(dA_cs_ptr).to(tl.float32)
75
+ scale = tl.exp(dA_cs)
76
+ if HAS_SEQ_IDX:
77
+ seq_idx_new = tl.load(seq_idx_ptr + (min((c + 1) * chunk_size, seqlen) - 1) * stride_seq_idx_seqlen)
78
+ scale = tl.where(seq_idx_new == seq_idx, scale, 0.0)
79
+ seq_idx = seq_idx_new
80
+ states = scale * states + new_states
81
+ if c < nchunks - 1:
82
+ tl.store(out_ptrs, states, mask=offs_m < dim)
83
+ else:
84
+ tl.store(final_states_ptrs, states, mask=offs_m < dim)
85
+ states_ptrs += stride_states_chunk
86
+ dA_cs_ptr += stride_dA_cs_chunk
87
+ out_ptrs += stride_out_chunk
88
+
89
+
90
+ @triton.autotune(
91
+ configs=autotune_configs([
92
+ triton.Config({'BLOCK_SIZE': 64}),
93
+ triton.Config({'BLOCK_SIZE': 128}),
94
+ triton.Config({'BLOCK_SIZE': 256}),
95
+ triton.Config({'BLOCK_SIZE': 512}),
96
+ triton.Config({'BLOCK_SIZE': 1024}),
97
+ triton.Config({'BLOCK_SIZE': 2048}),
98
+ ]),
99
+ key=['dim'],
100
+ )
101
+ @triton.jit
102
+ def _state_passing_bwd_kernel(
103
+ # Pointers to matrices
104
+ dout_ptr, out_ptr, dA_cs_ptr, dfinal_states_ptr, seq_idx_ptr,
105
+ dstates_ptr, ddA_cs_ptr, dinitstates_ptr, states_converted_ptr,
106
+ # Matrix dimensions
107
+ dim, nchunks, seqlen, chunk_size,
108
+ # Strides
109
+ stride_dout_batch, stride_dout_chunk, stride_dout_head, stride_dout_dim,
110
+ stride_out_batch, stride_out_chunk, stride_out_head, stride_out_dim,
111
+ stride_dA_cs_batch, stride_dA_cs_chunk, stride_dA_cs_head,
112
+ stride_dfinal_states_batch, stride_dfinal_states_head, stride_dfinal_states_dim,
113
+ stride_seq_idx_batch, stride_seq_idx_seqlen,
114
+ stride_dstates_batch, stride_dstates_chunk, stride_dstates_head, stride_dstates_dim,
115
+ stride_ddA_cs_batch, stride_ddA_cs_chunk, stride_ddA_cs_head,
116
+ stride_dinitstates_batch, stride_dinitstates_head, stride_dinitstates_dim,
117
+ # Meta-parameters
118
+ CONVERT_STATES: tl.constexpr,
119
+ HAS_DFINAL_STATES: tl.constexpr,
120
+ HAS_DINITSTATES: tl.constexpr,
121
+ HAS_SEQ_IDX: tl.constexpr,
122
+ BLOCK_SIZE: tl.constexpr,
123
+ ):
124
+ pid_b = tl.program_id(axis=1)
125
+ pid_h = tl.program_id(axis=2)
126
+ pid_m = tl.program_id(axis=0)
127
+ dstates_ptr += pid_b * stride_dstates_batch + pid_h * stride_dstates_head + (nchunks - 1) * stride_dstates_chunk
128
+ dA_cs_ptr += pid_b * stride_dA_cs_batch + pid_h * stride_dA_cs_head + (nchunks - 1) * stride_dA_cs_chunk
129
+ ddA_cs_ptr += pid_b * stride_ddA_cs_batch + pid_h * stride_ddA_cs_head + (nchunks - 1) * stride_ddA_cs_chunk + pid_m
130
+ out_ptr += pid_b * stride_out_batch + pid_h * stride_out_head + (nchunks - 1) * stride_out_chunk
131
+ dout_ptr += pid_b * stride_dout_batch + pid_h * stride_dout_head + (nchunks - 1) * stride_dout_chunk
132
+ if CONVERT_STATES:
133
+ states_converted_ptr += pid_b * stride_out_batch + pid_h * stride_out_head + (nchunks - 1) * stride_out_chunk
134
+ if HAS_DFINAL_STATES:
135
+ dfinal_states_ptr += pid_b * stride_dfinal_states_batch + pid_h * stride_dfinal_states_head
136
+ if HAS_DINITSTATES:
137
+ dinitstates_ptr += pid_b * stride_dinitstates_batch + pid_h * stride_dinitstates_head
138
+ if HAS_SEQ_IDX:
139
+ seq_idx_ptr += pid_b * stride_seq_idx_batch
140
+
141
+ offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
142
+ dstates_ptrs = dstates_ptr + offs_m * stride_dstates_dim
143
+ out_ptrs = out_ptr + offs_m * stride_out_dim
144
+ dout_ptrs = dout_ptr + offs_m * stride_dout_dim
145
+ if CONVERT_STATES:
146
+ states_converted_ptrs = states_converted_ptr + offs_m * stride_out_dim
147
+
148
+ if HAS_DFINAL_STATES:
149
+ dstates = tl.load(dfinal_states_ptr + offs_m * stride_dfinal_states_dim, mask=offs_m < dim, other=0.0).to(tl.float32)
150
+ else:
151
+ dstates = tl.zeros((BLOCK_SIZE, ), dtype=tl.float32)
152
+ tl.store(dstates_ptrs, dstates, mask=offs_m < dim)
153
+ if HAS_SEQ_IDX:
154
+ seq_idx = tl.load(seq_idx_ptr + (seqlen - 1) * stride_seq_idx_seqlen)
155
+ dstates_ptrs -= stride_dstates_chunk
156
+ for c in range(nchunks - 1):
157
+ dA_cs = tl.load(dA_cs_ptr).to(tl.float32)
158
+ scale = tl.exp(dA_cs)
159
+ if HAS_SEQ_IDX:
160
+ seq_idx_new = tl.load(seq_idx_ptr + (((nchunks - c - 1) * chunk_size - 1) * stride_seq_idx_seqlen))
161
+ scale = tl.where(seq_idx_new == seq_idx, scale, 0.0)
162
+ seq_idx = seq_idx_new
163
+ out = tl.load(out_ptrs, mask=offs_m < dim, other=0.0).to(tl.float32)
164
+ if CONVERT_STATES:
165
+ tl.store(states_converted_ptrs, out, mask=offs_m < dim)
166
+ ddA = tl.sum(out * dstates) * scale
167
+ tl.store(ddA_cs_ptr, ddA)
168
+ dout = tl.load(dout_ptrs, mask=offs_m < dim, other=0.0).to(tl.float32)
169
+ dstates = scale * dstates + dout
170
+ tl.store(dstates_ptrs, dstates, mask=offs_m < dim)
171
+ dout_ptrs -= stride_dout_chunk
172
+ dstates_ptrs -= stride_dstates_chunk
173
+ dA_cs_ptr -= stride_dA_cs_chunk
174
+ ddA_cs_ptr -= stride_ddA_cs_chunk
175
+ out_ptrs -= stride_out_chunk
176
+ if CONVERT_STATES:
177
+ states_converted_ptrs -= stride_out_chunk
178
+ if CONVERT_STATES:
179
+ out = tl.load(out_ptrs, mask=offs_m < dim, other=0.0).to(tl.float32)
180
+ tl.store(states_converted_ptrs, out, mask=offs_m < dim)
181
+ if not HAS_DINITSTATES:
182
+ tl.store(ddA_cs_ptr, 0.0)
183
+ else:
184
+ dA_cs = tl.load(dA_cs_ptr).to(tl.float32)
185
+ scale = tl.exp(dA_cs)
186
+ if HAS_SEQ_IDX:
187
+ scale = tl.where(seq_idx == 0, scale, 0.0)
188
+ out = tl.load(out_ptrs, mask=offs_m < dim, other=0.0).to(tl.float32)
189
+ ddA = tl.sum(out * dstates) * scale
190
+ tl.store(ddA_cs_ptr, ddA)
191
+ dout = tl.load(dout_ptrs, mask=offs_m < dim, other=0.0).to(tl.float32)
192
+ dstates = scale * dstates + dout
193
+ tl.store(dinitstates_ptr + offs_m * stride_dinitstates_dim, dstates, mask=offs_m < dim)
194
+
195
+
196
+ def _state_passing_fwd(states, dA_chunk_cumsum, initial_states=None, seq_idx=None, chunk_size=None,
197
+ out_dtype=None):
198
+ batch, nchunks, nheads, dim = states.shape
199
+ assert dA_chunk_cumsum.shape == (batch, nheads, nchunks)
200
+ if initial_states is not None:
201
+ assert initial_states.shape == (batch, nheads, dim)
202
+ if seq_idx is not None:
203
+ assert chunk_size is not None
204
+ seqlen = seq_idx.shape[-1]
205
+ assert seq_idx.shape == (batch, seqlen)
206
+ out_dtype = states.dtype if out_dtype is None else out_dtype
207
+ out = torch.empty((batch, nchunks, nheads, dim), device=states.device, dtype=out_dtype)
208
+ final_states = torch.empty((batch, nheads, dim), device=states.device, dtype=torch.float32)
209
+ grid = lambda META: (triton.cdiv(dim, META['BLOCK_SIZE']), batch, nheads)
210
+ with torch.cuda.device(states.device.index):
211
+ _state_passing_fwd_kernel[grid](
212
+ states, out, final_states, dA_chunk_cumsum, initial_states, seq_idx,
213
+ dim, nchunks, seqlen if seq_idx is not None else 0, chunk_size if seq_idx is not None else 0,
214
+ states.stride(0), states.stride(1), states.stride(2), states.stride(3),
215
+ out.stride(0), out.stride(1), out.stride(2), out.stride(3),
216
+ final_states.stride(0), final_states.stride(1), final_states.stride(2),
217
+ dA_chunk_cumsum.stride(0), dA_chunk_cumsum.stride(2), dA_chunk_cumsum.stride(1),
218
+ *((initial_states.stride(0), initial_states.stride(1), initial_states.stride(2))
219
+ if initial_states is not None else (0, 0, 0)),
220
+ *((seq_idx.stride(0), seq_idx.stride(1)) if seq_idx is not None else (0, 0)),
221
+ HAS_INITSTATES=initial_states is not None,
222
+ HAS_SEQ_IDX=seq_idx is not None,
223
+ )
224
+ return out, final_states
225
+
226
+
227
+ def _state_passing_bwd(
228
+ states, dA_chunk_cumsum, dout, dfinal_states=None, seq_idx=None, has_initial_states=None,
229
+ dstates_dtype=None, states_dtype=None, chunk_size=None
230
+ ):
231
+ """
232
+ states contains the initial_states at index 0. The final states are not included in states.
233
+ """
234
+ batch, nchunks, nheads, dim = states.shape
235
+ assert dA_chunk_cumsum.shape == (batch, nheads, nchunks)
236
+ assert dout.shape == (batch, nchunks, nheads, dim)
237
+ if seq_idx is not None:
238
+ assert chunk_size is not None
239
+ seqlen = seq_idx.shape[-1]
240
+ assert seq_idx.shape == (batch, seqlen)
241
+ dstates = torch.empty_like(dout, dtype=dstates_dtype if dstates_dtype is not None else dout.dtype)
242
+ if states_dtype is not None and states_dtype != states.dtype:
243
+ states_converted = torch.empty_like(states, dtype=dstates_dtype if dstates_dtype is not None else dout.dtype)
244
+ assert states_converted.stride() == states.stride()
245
+ else:
246
+ states_converted = None
247
+ if has_initial_states:
248
+ dinitstates = torch.empty_like(dstates[:, 0])
249
+ else:
250
+ dinitstates = None
251
+ if dfinal_states is not None:
252
+ assert dfinal_states.shape == (batch, nheads, dim)
253
+ BLOCK_SIZE_min = 64
254
+ n_blocks = (dim + BLOCK_SIZE_min - 1) // BLOCK_SIZE_min
255
+ ddA_chunk_cumsum = torch.empty(batch, nheads, nchunks, n_blocks,
256
+ dtype=torch.float32, device=dA_chunk_cumsum.device)
257
+ grid = lambda META: (triton.cdiv(dim, META['BLOCK_SIZE']), batch, nheads)
258
+ with torch.cuda.device(dout.device.index):
259
+ _state_passing_bwd_kernel[grid](
260
+ dout, states, dA_chunk_cumsum, dfinal_states, seq_idx,
261
+ dstates, ddA_chunk_cumsum, dinitstates, states_converted,
262
+ dim, nchunks, seqlen if seq_idx is not None else 0, chunk_size if seq_idx is not None else 0,
263
+ dout.stride(0), dout.stride(1), dout.stride(2), dout.stride(3),
264
+ states.stride(0), states.stride(1), states.stride(2), states.stride(3),
265
+ dA_chunk_cumsum.stride(0), dA_chunk_cumsum.stride(2), dA_chunk_cumsum.stride(1),
266
+ *((dfinal_states.stride(0), dfinal_states.stride(1), dfinal_states.stride(2))
267
+ if dfinal_states is not None else (0, 0, 0)),
268
+ *((seq_idx.stride(0), seq_idx.stride(1)) if seq_idx is not None else (0, 0)),
269
+ dstates.stride(0), dstates.stride(1), dstates.stride(2), dstates.stride(3),
270
+ ddA_chunk_cumsum.stride(0), ddA_chunk_cumsum.stride(2), ddA_chunk_cumsum.stride(1),
271
+ *((dinitstates.stride(0), dinitstates.stride(1), dinitstates.stride(2))
272
+ if dinitstates is not None else (0, 0, 0)),
273
+ CONVERT_STATES=states_converted is not None,
274
+ HAS_DFINAL_STATES=dfinal_states is not None,
275
+ HAS_DINITSTATES=dinitstates is not None,
276
+ HAS_SEQ_IDX=seq_idx is not None,
277
+ )
278
+ BLOCK_SIZE_actual = _state_passing_bwd_kernel.best_config.kwargs["BLOCK_SIZE"]
279
+ n_valid_blocks = (dim + BLOCK_SIZE_actual - 1) // BLOCK_SIZE_actual
280
+ ddA_chunk_cumsum = ddA_chunk_cumsum[..., :n_valid_blocks].sum(dim=-1).to(dtype=dA_chunk_cumsum.dtype)
281
+ if states_dtype is not None and states_dtype == states.dtype:
282
+ states_converted = states
283
+ return (dstates, ddA_chunk_cumsum, dinitstates) if states_dtype is None else (dstates, ddA_chunk_cumsum, dinitstates, states_converted)
284
+
285
+
286
+ class StatePassingFn(torch.autograd.Function):
287
+
288
+ @staticmethod
289
+ def forward(ctx, states, dA_chunk_cumsum, initial_states=None):
290
+ batch, nchunks, nheads, dim = states.shape
291
+ assert dA_chunk_cumsum.shape == (batch, nheads, nchunks)
292
+ if states.stride(-1) != 1:
293
+ states = states.contiguous()
294
+ out, final_states = _state_passing_fwd(states, dA_chunk_cumsum, initial_states)
295
+ ctx.save_for_backward(out, dA_chunk_cumsum)
296
+ ctx.has_initial_states = initial_states is not None
297
+ return out, final_states
298
+
299
+ @staticmethod
300
+ def backward(ctx, dout, dfinal_states):
301
+ out, dA_chunk_cumsum = ctx.saved_tensors
302
+ batch, nchunks, nheads, dim = out.shape
303
+ assert dout.shape == (batch, nchunks, nheads, dim)
304
+ assert dA_chunk_cumsum.shape == (batch, nheads, nchunks)
305
+ assert dfinal_states.shape == (batch, nheads, dim)
306
+ if dout.stride(-1) != 1:
307
+ dout = dout.contiguous()
308
+ dstates, ddA_chunk_cumsum, dinitstates = _state_passing_bwd(
309
+ out, dA_chunk_cumsum, dout, dfinal_states=dfinal_states , has_initial_states=ctx.has_initial_states
310
+ )
311
+ return dstates, ddA_chunk_cumsum, dinitstates
312
+
313
+
314
+ def state_passing(states, dA_chunk_cumsum, initial_states=None):
315
+ """
316
+ Argument:
317
+ states: (batch, nchunks, nheads, dim)
318
+ dA_chunk_cumsum: (batch, nheads, nchunks)
319
+ initial_states: (batch, nheads, dim)
320
+ Return:
321
+ out: (batch, nchunks, nheads, dim)
322
+ final_states: (batch, nheads, dim)
323
+ """
324
+ return StatePassingFn.apply(states, dA_chunk_cumsum, initial_states)
325
+
326
+
327
+ def state_passing_ref(states, dA_chunk_cumsum, initial_states=None):
328
+ """
329
+ Argument:
330
+ states: (batch, nchunks, nheads, dim)
331
+ dA_chunk_cumsum: (batch, nheads, nchunks)
332
+ initial_states: (batch, nheads, dim)
333
+ Return:
334
+ out: (batch, nchunks, nheads, dim)
335
+ final_states: (batch, nheads, dim)
336
+ """
337
+ if initial_states is None:
338
+ initial_states = torch.zeros_like(states[:, 0])
339
+ states = torch.cat([rearrange(initial_states, "b h d -> b 1 h d"), states], dim=1)
340
+ dA_chunk_cumsum = F.pad(dA_chunk_cumsum, (1, 0))
341
+ dA_chunk_cumsum = torch.cumsum(dA_chunk_cumsum, dim=-1)
342
+ nchunks = dA_chunk_cumsum.shape[-1]
343
+ # (batch, nheads, nchunks, nchunks)
344
+ dt_chunk_segment_sum = dA_chunk_cumsum[:, :, :, None] - dA_chunk_cumsum[:, :, None, :]
345
+ # (batch, nheads, nchunks, nchunks)
346
+ decay_chunk = torch.exp(dt_chunk_segment_sum)
347
+ causal_mask = torch.tril(torch.ones(nchunks, nchunks, device=states.device, dtype=bool), diagonal=0)
348
+ decay_chunk = decay_chunk.masked_fill(~causal_mask, 0)
349
+ out = torch.einsum("bhzc,bchd->bzhd", decay_chunk.to(dtype=states.dtype), states)
350
+ return out[:, :-1], out[:, -1]
mamba_ssm/utils/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Utilities for the pinned BlueMagpie Mamba runtime."""
mamba_ssm/utils/determinism.py ADDED
@@ -0,0 +1,96 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2024, Tri Dao, Albert Gu.
2
+
3
+ import os
4
+ import warnings
5
+ from packaging import version
6
+
7
+ import torch
8
+
9
+ try:
10
+ import triton
11
+ TRITON_VERSION = version.parse(triton.__version__)
12
+ except ImportError:
13
+ TRITON_VERSION = version.parse("0.0.0")
14
+
15
+ TRITON_HAS_CACHE_RESULTS = TRITON_VERSION >= version.parse("3.4.0")
16
+ _autotune_warning_issued = False
17
+
18
+ _deterministic_override = None
19
+
20
+
21
+ def use_deterministic_mode():
22
+ if _deterministic_override is not None:
23
+ return _deterministic_override
24
+ env = os.environ.get('MAMBA_DETERMINISTIC')
25
+ if env:
26
+ return env[0] == '1'
27
+ return torch.are_deterministic_algorithms_enabled()
28
+
29
+
30
+ def set_deterministic_mode(value):
31
+ global _deterministic_override
32
+ _deterministic_override = value
33
+
34
+
35
+ def _estimate_config_cost(cfg):
36
+ """Estimate shared memory cost of a config. Lower is cheaper."""
37
+ block_product = 1
38
+ for key, val in cfg.kwargs.items():
39
+ if key.startswith('BLOCK_SIZE_'):
40
+ block_product *= val
41
+ return block_product * (getattr(cfg, 'num_stages', 1) or 1)
42
+
43
+
44
+ def _filter_configs_by_block_sizes(configs):
45
+ """Filter configs by TRITON_AUTOTUNE_BLOCK_SIZE_* env vars."""
46
+ env_filters = {}
47
+ for suffix in ('M', 'N', 'K', 'DSTATE'):
48
+ env_val = os.environ.get(f"TRITON_AUTOTUNE_BLOCK_SIZE_{suffix}")
49
+ if env_val is not None:
50
+ env_filters[f'BLOCK_SIZE_{suffix}'] = int(env_val)
51
+ if not env_filters:
52
+ return None
53
+ matching = configs
54
+ for key, target in env_filters.items():
55
+ matching = [c for c in matching if c.kwargs.get(key) == target]
56
+ return matching[:1] if matching else None
57
+
58
+
59
+ def autotune_configs(configs):
60
+ """Select autotune configs for deterministic mode.
61
+
62
+ Uses cached autotuning (TRITON_CACHE_AUTOTUNING=1) if Triton >= 3.4.0,
63
+ otherwise auto-selects the cheapest config by block size * stages.
64
+ """
65
+ if not configs or not use_deterministic_mode():
66
+ return configs
67
+ if TRITON_HAS_CACHE_RESULTS and os.environ.get("TRITON_CACHE_AUTOTUNING") == "1":
68
+ return configs
69
+ global _autotune_warning_issued
70
+ if not _autotune_warning_issued:
71
+ _autotune_warning_issued = True
72
+ msg = "Deterministic mode: set TRITON_CACHE_AUTOTUNING=1 for cached autotuning." if TRITON_HAS_CACHE_RESULTS else "Deterministic mode: upgrade to Triton >= 3.4.0 for cached autotuning."
73
+ warnings.warn(msg)
74
+ filtered = _filter_configs_by_block_sizes(configs)
75
+ if filtered:
76
+ return filtered
77
+ return [min(configs, key=_estimate_config_cost)]
78
+
79
+
80
+ def alloc_tile_workspace(base_shape, tile_dim, dtype, device, deterministic, *, zero_init=True):
81
+ """Allocate buffer for deterministic per-program reductions."""
82
+ if base_shape is None:
83
+ return None, 0
84
+ if deterministic:
85
+ factory = torch.zeros if zero_init else torch.empty
86
+ tensor = factory(*base_shape, tile_dim, device=device, dtype=dtype)
87
+ return tensor, tensor.stride(-1)
88
+ return torch.empty(*base_shape, device=device, dtype=dtype), 0
89
+
90
+
91
+ def finalize_tile_workspace(tensor, deterministic):
92
+ if tensor is None:
93
+ return None
94
+ if deterministic:
95
+ tensor = tensor.sum(dim=-1)
96
+ return tensor
mamba_ssm/utils/torch.py ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from functools import partial
3
+ from typing import Callable
4
+
5
+ def custom_amp_decorator(dec: Callable, cuda_amp_deprecated: bool):
6
+ def decorator(*args, **kwargs):
7
+ if cuda_amp_deprecated:
8
+ kwargs["device_type"] = "cuda"
9
+ return dec(*args, **kwargs)
10
+ return decorator
11
+
12
+
13
+ if hasattr(torch.amp, "custom_fwd"): # type: ignore[attr-defined]
14
+ deprecated = True
15
+ from torch.amp import custom_fwd, custom_bwd # type: ignore[attr-defined]
16
+ else:
17
+ deprecated = False
18
+ from torch.cuda.amp import custom_fwd, custom_bwd
19
+
20
+ custom_fwd = custom_amp_decorator(custom_fwd, deprecated)
21
+ custom_bwd = custom_amp_decorator(custom_bwd, deprecated)
requirements.txt CHANGED
@@ -7,6 +7,7 @@ transformers==4.57.6
7
  accelerate==1.12.0
8
  einops==0.8.2
9
  pydantic==2.11.10
 
10
  # Hugging Face ZeroGPU currently builds this Space on CPython 3.10.
11
  numpy==2.2.6
12
  scipy==1.15.3
 
7
  accelerate==1.12.0
8
  einops==0.8.2
9
  pydantic==2.11.10
10
+ packaging==26.0
11
  # Hugging Face ZeroGPU currently builds this Space on CPython 3.10.
12
  numpy==2.2.6
13
  scipy==1.15.3
tests/test_release_pins.py CHANGED
@@ -225,7 +225,9 @@ def test_cuda_runtime_is_pinned_deterministic_before_model_import():
225
  assert constants["EXPECTED_TORCH_VERSION"] == "2.11.0+cu130"
226
  assert constants["EXPECTED_TORCH_CUDA_VERSION"] == "13.0"
227
  assert constants["EXPECTED_CUDNN_VERSION"] == 91900
228
- assert constants["EXPECTED_MAMBA_SSM_VERSION"] == "2.3.2.post1"
 
 
229
  assert constants["EXPECTED_TRITON_VERSION"] == "3.6.0"
230
  assert source.index('os.environ.setdefault(\n "CUBLAS_WORKSPACE_CONFIG"') < (
231
  source.index("import torch")
@@ -271,6 +273,59 @@ def test_tts_runtime_is_vendored_from_the_frozen_commit():
271
  assert hashlib.sha256(payload).hexdigest() == expected_hash
272
 
273
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
274
  def test_barbet_runtime_dependency_is_commit_pinned():
275
  requirements = (ROOT / "requirements.txt").read_text(encoding="utf-8").splitlines()
276
  barbet_lines = [line for line in requirements if "OpenFormosa/Barbet.git" in line]
@@ -315,6 +370,7 @@ def test_quality_runtime_dependencies_are_version_pinned():
315
  assert "accelerate==1.12.0" in requirements
316
  assert "einops==0.8.2" in requirements
317
  assert "pydantic==2.11.10" in requirements
 
318
  assert "numpy==2.2.6" in requirements
319
  assert "scipy==1.15.3" in requirements
320
  assert "numexpr==2.10.0" in requirements
 
225
  assert constants["EXPECTED_TORCH_VERSION"] == "2.11.0+cu130"
226
  assert constants["EXPECTED_TORCH_CUDA_VERSION"] == "13.0"
227
  assert constants["EXPECTED_CUDNN_VERSION"] == 91900
228
+ assert constants["EXPECTED_MAMBA_SSM_VERSION"] == (
229
+ "2.3.2.post1+bluemagpie.triton1"
230
+ )
231
  assert constants["EXPECTED_TRITON_VERSION"] == "3.6.0"
232
  assert source.index('os.environ.setdefault(\n "CUBLAS_WORKSPACE_CONFIG"') < (
233
  source.index("import torch")
 
273
  assert hashlib.sha256(payload).hexdigest() == expected_hash
274
 
275
 
276
+ def test_mamba_triton_runtime_is_vendored_and_importable_without_extension():
277
+ import sys
278
+
279
+ import mamba_ssm
280
+ from mamba_ssm.ops.triton.layernorm_gated import RMSNorm
281
+ from mamba_ssm.ops.triton.ssd_combined import mamba_chunk_scan_combined
282
+
283
+ assert Path(mamba_ssm.__file__).resolve().is_relative_to(ROOT)
284
+ assert mamba_ssm.__version__ == "2.3.2.post1+bluemagpie.triton1"
285
+ assert RMSNorm.__name__ == "RMSNorm"
286
+ assert mamba_chunk_scan_combined.__name__ == "mamba_chunk_scan_combined"
287
+ assert "selective_scan_cuda" not in sys.modules
288
+
289
+ pinned_hashes = {
290
+ "LICENSE.upstream": (
291
+ "760939b000194d04548ede6a857bbe735d1695d8422ec85955c8e2bd7f4b95c5"
292
+ ),
293
+ "ops/triton/k_activations.py": (
294
+ "ede8d75600b1b01fd867df00eae4a727df9a34ef8091baf965522c36b3dfdf7b"
295
+ ),
296
+ "ops/triton/layernorm_gated.py": (
297
+ "eb6252e247b90f1c8a75946efbc1a221e0c4da701b6757ddae49f3495cf7a42f"
298
+ ),
299
+ "ops/triton/softplus.py": (
300
+ "989f7667ad7f8866dfafa2783d553555e235478cbce349905b23d8fb66b7c5ab"
301
+ ),
302
+ "ops/triton/ssd_bmm.py": (
303
+ "5059b16f8fa269cd84e8159ee3c6880e4ccf90e4bedf9c2d733e6bc2baff9587"
304
+ ),
305
+ "ops/triton/ssd_chunk_scan.py": (
306
+ "055b5ce4cb0f30c84c3e031d775f218227fa0b793359254a2465777097b224f9"
307
+ ),
308
+ "ops/triton/ssd_chunk_state.py": (
309
+ "6d82771b2f62a7bf84c381c8ef8e4b579a95e0352021d117507a1072e4d71dcd"
310
+ ),
311
+ "ops/triton/ssd_combined.py": (
312
+ "39c19c1e4c8e36982847079bc12b4f07ed3af03ac68f6a6653200ef80df5515d"
313
+ ),
314
+ "ops/triton/ssd_state_passing.py": (
315
+ "ae1fab6c680cb5312ef746862a2cf8e74157033c80878490cc64fa1f24f88b70"
316
+ ),
317
+ "utils/determinism.py": (
318
+ "cb6e1c30392c11200425c2a23ad9fa3d47f50b556d15e9b0caf79b7d483d6f1d"
319
+ ),
320
+ "utils/torch.py": (
321
+ "1c3132a1a747e914874e84b5f300ea80d0297faba6eb9256cb5ac4b3d6413c06"
322
+ ),
323
+ }
324
+ for relative_path, expected_hash in pinned_hashes.items():
325
+ payload = (ROOT / "mamba_ssm" / relative_path).read_bytes()
326
+ assert hashlib.sha256(payload).hexdigest() == expected_hash
327
+
328
+
329
  def test_barbet_runtime_dependency_is_commit_pinned():
330
  requirements = (ROOT / "requirements.txt").read_text(encoding="utf-8").splitlines()
331
  barbet_lines = [line for line in requirements if "OpenFormosa/Barbet.git" in line]
 
370
  assert "accelerate==1.12.0" in requirements
371
  assert "einops==0.8.2" in requirements
372
  assert "pydantic==2.11.10" in requirements
373
+ assert "packaging==26.0" in requirements
374
  assert "numpy==2.2.6" in requirements
375
  assert "scipy==1.15.3" in requirements
376
  assert "numexpr==2.10.0" in requirements