Spaces:
Running on Zero
Running on Zero
Codex commited on
Commit ·
e1518d1
1
Parent(s): cb0290b
Vendor the pinned Mamba Triton runtime
Browse files- README.md +3 -0
- app.py +3 -2
- mamba_ssm/LICENSE.upstream +201 -0
- mamba_ssm/PROVENANCE.md +19 -0
- mamba_ssm/__init__.py +9 -0
- mamba_ssm/ops/__init__.py +1 -0
- mamba_ssm/ops/triton/__init__.py +1 -0
- mamba_ssm/ops/triton/k_activations.py +171 -0
- mamba_ssm/ops/triton/layernorm_gated.py +437 -0
- mamba_ssm/ops/triton/softplus.py +15 -0
- mamba_ssm/ops/triton/ssd_bmm.py +264 -0
- mamba_ssm/ops/triton/ssd_chunk_scan.py +0 -0
- mamba_ssm/ops/triton/ssd_chunk_state.py +1122 -0
- mamba_ssm/ops/triton/ssd_combined.py +1047 -0
- mamba_ssm/ops/triton/ssd_state_passing.py +350 -0
- mamba_ssm/utils/__init__.py +1 -0
- mamba_ssm/utils/determinism.py +96 -0
- mamba_ssm/utils/torch.py +21 -0
- requirements.txt +1 -0
- tests/test_release_pins.py +57 -1
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":
|
| 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"] ==
|
|
|
|
|
|
|
| 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
|