已合并
feat: add fused GDN backend for linear attention CP #1114
feat: add fused GDN backend for linear attention CP #1114
已合并
xu-xianliang创建于 8月5日
共 20 个文件变更+5004-31
@@ -3,3 +3,14 @@ hyper-parallel/hyper_parallel/core/hsdp/api.py:hsdp
3hyper-parallel/hyper_parallel/core/multicore/ops/mega_moe/op_host/mega_moe_def.cpp:ops::MegaMoe::MegaMoe3hyper-parallel/hyper_parallel/core/multicore/ops/mega_moe/op_host/mega_moe_def.cpp:ops::MegaMoe::MegaMoe
4hyper-parallel/hyper_parallel/core/multicore/ops/mega_moe_grad/op_host/mega_moe_grad_def.cpp:ops::MegaMoeGrad::MegaMoeGrad4hyper-parallel/hyper_parallel/core/multicore/ops/mega_moe_grad/op_host/mega_moe_grad_def.cpp:ops::MegaMoeGrad::MegaMoeGrad
5hyper-parallel/hyper_parallel/integration/llamafactory/context_parallel/models/qwen3_vl/qwen3vl_forward.py:forward5hyper-parallel/hyper_parallel/integration/llamafactory/context_parallel/models/qwen3_vl/qwen3vl_forward.py:forward
6+hyper-parallel/hyper_parallel/core/context_parallel/linear_attention_context_parallel.py:forward
7+hyper-parallel/hyper_parallel/core/context_parallel/linear_attention_context_parallel.py:backward
8+hyper-parallel/hyper_parallel/platform/torch/custom_ops/gdn/chunk_gated_delta_rule.py:chunk_gated_delta_rule
9+hyper-parallel/hyper_parallel/platform/torch/custom_ops/gdn/triton/chunk_delta_h.py:chunk_gated_delta_rule_fwd_kernel_h_blockdim64
10+hyper-parallel/hyper_parallel/platform/torch/custom_ops/gdn/triton/chunk_delta_h.py:chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64
11+hyper-parallel/hyper_parallel/platform/torch/custom_ops/gdn/triton/chunk_o.py:chunk_bwd_kernel_dqkwg
12+hyper-parallel/hyper_parallel/platform/torch/custom_ops/gdn/triton/chunk_o.py:chunk_fwd_kernel_o
13+hyper-parallel/hyper_parallel/platform/torch/custom_ops/gdn/triton/solve_tril.py:solve_tril_16x16_loop_kernel_paral_v3
14+hyper-parallel/hyper_parallel/platform/torch/custom_ops/gdn/triton/solve_tril.py:solve_tril
15+hyper-parallel/hyper_parallel/platform/torch/custom_ops/gdn/triton/state_summary.py:gdn_state_grad_ext_kernel
16+hyper-parallel/hyper_parallel/platform/torch/custom_ops/gdn/triton/wy_fast.py:prepare_wy_repr_bwd_kernel
@@ -13,6 +13,12 @@
13# limitations under the License.13# limitations under the License.
14# ============================================================================14# ============================================================================
15"""Context parallel execution for Qwen3.5-style Gated DeltaNet layers."""15"""Context parallel execution for Qwen3.5-style Gated DeltaNet layers."""
16+ 
17+# This module is the Torch implementation of the public CP style.
18+# pylint: disable=forbidden-backend-import,missing-public-type-hints
19+# pylint: disable=missing-public-docstring,not-callable
20+# PyTorch autograd.Function intentionally defines framework-specific signatures.
21+# pylint: disable=abstract-method,arguments-differ
16from __future__ import annotations22from __future__ import annotations
17 23 
18from typing import NamedTuple, Optional24from typing import NamedTuple, Optional
@@ -29,7 +35,11 @@ from hyper_parallel.core.context_parallel.context_parallel import (
29from hyper_parallel.core.dtensor.device_mesh import DeviceMesh35from hyper_parallel.core.dtensor.device_mesh import DeviceMesh
30from hyper_parallel.core.dtensor.dtensor import DTensor36from hyper_parallel.core.dtensor.dtensor import DTensor
31from hyper_parallel.core.tensor_parallel.style import ParallelStyle37from hyper_parallel.core.tensor_parallel.style import ParallelStyle
32-from hyper_parallel.models.modules.linear_attention import torch_chunk_gated_delta_rule38+from hyper_parallel.models.modules.linear_attention import (
39+ chunk_gated_delta_rule,
40+ is_triton_gdn_available,
41+ torch_chunk_gated_delta_rule,
42+)
33from hyper_parallel.platform import get_platform43from hyper_parallel.platform import get_platform
34 44 
35 45 
@@ -646,6 +656,290 @@ def _gdn_state_p2p_summary(
646 return core_attn_out656 return core_attn_out
647 657 
648 658 
659+class _GDNStateP2PTritonFunction(torch.autograd.Function):
660+ """Pipeline fused affine GDN states over sequence-sharded CP ranks."""
661+ 
662+ @staticmethod
663+ def forward( # pylint: disable=arguments-differ,too-many-locals
664+ ctx,
665+ query: torch.Tensor,
666+ key: torch.Tensor,
667+ value: torch.Tensor,
668+ g: torch.Tensor,
669+ beta: torch.Tensor,
670+ cp_rank: int,
671+ cp_size: int,
672+ cp_group,
673+ prev_rank: int,
674+ next_rank: int,
675+ ) -> torch.Tensor:
676+ """Run fused local GDN and forward its affine state across CP ranks."""
677+ from hyper_parallel.platform.torch.custom_ops.gdn.chunk_gated_delta_rule import ( # pylint: disable=import-outside-toplevel
678+ chunk_gated_delta_rule_fwd_apply_state_saved,
679+ chunk_gated_delta_rule_fwd_output_saved,
680+ chunk_gated_delta_rule_fwd_prepare_saved,
681+ )
682+ from hyper_parallel.platform.torch.custom_ops.gdn.state_summary import ( # pylint: disable=import-outside-toplevel
683+ apply_gdn_state_summary,
684+ chunk_gated_delta_rule_state_summary_fwd,
685+ )
686+ 
687+ (
688+ query_norm,
689+ key_norm,
690+ _,
691+ _,
692+ g_cumsum,
693+ matrix_a,
694+ w,
695+ u,
696+ scale,
697+ ) = chunk_gated_delta_rule_fwd_prepare_saved(
698+ query,
699+ key,
700+ value,
701+ g,
702+ beta,
703+ use_qk_l2norm_in_kernel=False,
704+ )
705+ initial_state = None
706+ recv_buffer = None
707+ recv_work = None
708+ if cp_rank > 0:
709+ recv_buffer = torch.empty(
710+ (query.shape[0], query.shape[2], query.shape[3], value.shape[3]),
711+ device=query.device,
712+ dtype=torch.float32,
713+ )
714+ recv_work = dist.irecv(recv_buffer, src=prev_rank, group=cp_group)
715+ 
716+ state_ext = None
717+ transition = None
718+ if cp_rank < cp_size - 1:
719+ state_ext, transition = chunk_gated_delta_rule_state_summary_fwd(
720+ key_norm,
721+ w,
722+ u,
723+ g_cumsum,
724+ )
725+ 
726+ if recv_work is not None:
727+ recv_work.wait()
728+ initial_state = recv_buffer
729+ 
730+ send_buffer = None
731+ send_work = None
732+ if cp_rank < cp_size - 1:
733+ send_buffer = apply_gdn_state_summary(
734+ state_ext,
735+ transition,
736+ initial_state,
737+ ).contiguous()
738+ send_work = dist.isend(send_buffer, dst=next_rank, group=cp_group)
739+ 
740+ h, v_new, _ = chunk_gated_delta_rule_fwd_apply_state_saved(
741+ key_norm,
742+ g_cumsum,
743+ w,
744+ u,
745+ initial_state=initial_state,
746+ output_final_state=False,
747+ )
748+ output = chunk_gated_delta_rule_fwd_output_saved(
749+ query_norm,
750+ key_norm,
751+ g_cumsum,
752+ h,
753+ v_new,
754+ scale,
755+ ).to(query.dtype)
756+ 
757+ if send_work is not None:
758+ send_work.wait()
759+ 
760+ empty = query.new_empty(0)
761+ ctx.save_for_backward(
762+ query_norm,
763+ key_norm,
764+ value,
765+ g_cumsum,
766+ beta,
767+ matrix_a,
768+ initial_state if initial_state is not None else empty,
769+ transition if transition is not None else empty,
770+ )
771+ ctx.has_initial_state = initial_state is not None
772+ ctx.cp_rank = cp_rank
773+ ctx.cp_size = cp_size
774+ ctx.cp_group = cp_group
775+ ctx.prev_rank = prev_rank
776+ ctx.next_rank = next_rank
777+ ctx.scale = scale
778+ return output
779+ 
780+ @staticmethod
781+ def backward(ctx, grad_output: torch.Tensor): # pylint: disable=too-many-locals
782+ """Backpropagate local GDN tensors and the state gradient wavefront."""
783+ from hyper_parallel.platform.torch.custom_ops.gdn.chunk_gated_delta_rule import ( # pylint: disable=import-outside-toplevel
784+ chunk_gated_delta_rule_bwd_finish_saved,
785+ chunk_gated_delta_rule_bwd_prepare_saved,
786+ chunk_gated_delta_rule_bwd_state_saved,
787+ )
788+ from hyper_parallel.platform.torch.custom_ops.gdn.state_summary import ( # pylint: disable=import-outside-toplevel
789+ apply_gdn_state_gradient_summary,
790+ chunk_gated_delta_rule_state_gradient_summary_bwd,
791+ )
792+ 
793+ (
794+ query,
795+ key,
796+ value,
797+ g_cumsum,
798+ beta,
799+ matrix_a,
800+ initial_state,
801+ transition,
802+ ) = ctx.saved_tensors
803+ if not ctx.has_initial_state:
804+ initial_state = None
805+ 
806+ w, h, v_new, dv = chunk_gated_delta_rule_bwd_prepare_saved(
807+ query,
808+ key,
809+ value,
810+ g_cumsum,
811+ beta,
812+ matrix_a,
813+ initial_state,
814+ grad_output,
815+ ctx.scale,
816+ )
817+ 
818+ grad_state_ext = None
819+ if ctx.cp_rank > 0:
820+ grad_state_ext = chunk_gated_delta_rule_state_gradient_summary_bwd(
821+ query,
822+ key,
823+ w,
824+ g_cumsum,
825+ grad_output,
826+ dv,
827+ ctx.scale,
828+ )
829+ 
830+ grad_final_state = None
831+ recv_work = None
832+ if ctx.cp_rank < ctx.cp_size - 1:
833+ recv_buffer = torch.empty(
834+ (query.shape[0], query.shape[2], query.shape[3], value.shape[3]),
835+ device=grad_output.device,
836+ dtype=torch.float32,
837+ )
838+ recv_work = dist.irecv(
839+ recv_buffer,
840+ src=ctx.next_rank,
841+ group=ctx.cp_group,
842+ )
843+ if recv_work is not None:
844+ recv_work.wait()
845+ grad_final_state = recv_buffer
846+ 
847+ send_buffer = None
848+ send_work = None
849+ if ctx.cp_rank > 0:
850+ send_buffer = apply_gdn_state_gradient_summary(
851+ grad_state_ext,
852+ transition,
853+ grad_final_state,
854+ ).contiguous()
855+ send_work = dist.isend(
856+ send_buffer,
857+ dst=ctx.prev_rank,
858+ group=ctx.cp_group,
859+ )
860+ 
861+ dh, _, dv = chunk_gated_delta_rule_bwd_state_saved(
862+ query,
863+ key,
864+ g_cumsum,
865+ w,
866+ initial_state,
867+ grad_final_state,
868+ grad_output,
869+ dv,
870+ ctx.scale,
871+ )
872+ empty = query.new_empty(0)
873+ dq, dk, dv, dg, dbeta = chunk_gated_delta_rule_bwd_finish_saved(
874+ query,
875+ key,
876+ query,
877+ key,
878+ value,
879+ g_cumsum,
880+ beta,
881+ matrix_a,
882+ w,
883+ h,
884+ v_new,
885+ dv,
886+ grad_output,
887+ dh,
888+ empty,
889+ empty,
890+ ctx.scale,
891+ use_qk_l2norm_in_kernel=False,
892+ )
893+ 
894+ if send_work is not None:
895+ send_work.wait()
896+ return dq, dk, dv, dg, dbeta, None, None, None, None, None
897+ 
898+ 
899+def _gdn_state_p2p_triton(
900+ query: torch.Tensor,
901+ key: torch.Tensor,
902+ value: torch.Tensor,
903+ g: torch.Tensor,
904+ beta: torch.Tensor,
905+ cp_mesh: DeviceMesh,
906+ cp_rank: int,
907+ cp_size: int,
908+) -> torch.Tensor:
909+ """Run fused local GDN with an affine state wavefront."""
910+ if cp_size == 1:
911+ output, _ = chunk_gated_delta_rule(
912+ query,
913+ key,
914+ value,
915+ g=g,
916+ beta=beta,
917+ output_final_state=False,
918+ use_qk_l2norm_in_kernel=True,
919+ backend="triton",
920+ )
921+ return output
922+ 
923+ query = _l2norm_torch(query)
924+ key = _l2norm_torch(key)
925+ prev_rank = _global_peer_rank(cp_mesh, cp_rank - 1) if cp_rank > 0 else -1
926+ next_rank = (
927+ _global_peer_rank(cp_mesh, cp_rank + 1) if cp_rank < cp_size - 1 else -1
928+ )
929+ return _GDNStateP2PTritonFunction.apply(
930+ query,
931+ key,
932+ value,
933+ g,
934+ beta,
935+ cp_rank,
936+ cp_size,
937+ cp_mesh.get_group(),
938+ prev_rank,
939+ next_rank,
940+ )
941+ 
942+ 
649def _differentiable_all_to_all_shard(943def _differentiable_all_to_all_shard(
650 tensor: torch.Tensor,944 tensor: torch.Tensor,
651 device_mesh: DeviceMesh,945 device_mesh: DeviceMesh,
@@ -736,9 +1030,16 @@ class LinearAttentionUlyssesCPWrapper(nn.Module):
736 [B, S_local, full_heads]``.1030 [B, S_local, full_heads]``.
737 """1031 """
738 1032 
739- def __init__(self, module: nn.Module, device_mesh: DeviceMesh):1033+ def __init__(
1034+ self,
1035+ module: nn.Module,
1036+ device_mesh: DeviceMesh,
1037+ *,
1038+ backend: str = "eager",
1039+ ):
740 super().__init__()1040 super().__init__()
741 self.module = module1041 self.module = module
1042+ self.gdn_backend = backend
742 self.cp_mesh = _ensure_1d(device_mesh)1043 self.cp_mesh = _ensure_1d(device_mesh)
743 self.cp_size = self.cp_mesh.size()1044 self.cp_size = self.cp_mesh.size()
744 self.cp_rank = self.cp_mesh.get_local_rank()1045 self.cp_rank = self.cp_mesh.get_local_rank()
@@ -871,7 +1172,9 @@ class LinearAttentionUlyssesCPWrapper(nn.Module):
871 [base.key_dim, base.key_dim, base.value_dim],1172 [base.key_dim, base.key_dim, base.value_dim],
872 dim=-1,1173 dim=-1,
873 )1174 )
874- q_proj, k_proj, v_proj, b, a = self._seq_to_head_qkvba(q_proj, k_proj, v_proj, b, a)1175+ q_proj, k_proj, v_proj, b, a = self._seq_to_head_qkvba(
1176+ q_proj, k_proj, v_proj, b, a
1177+ )
875 1178 
876 full_seq_len = q_proj.shape[1]1179 full_seq_len = q_proj.shape[1]
877 local_key_dim = base.key_dim // self.cp_size1180 local_key_dim = base.key_dim // self.cp_size
@@ -910,7 +1213,7 @@ class LinearAttentionUlyssesCPWrapper(nn.Module):
910 query = query.repeat_interleave(base.kv_groups, dim=2)1213 query = query.repeat_interleave(base.kv_groups, dim=2)
911 key = key.repeat_interleave(base.kv_groups, dim=2)1214 key = key.repeat_interleave(base.kv_groups, dim=2)
912 1215 
913- core_attn_out, _ = torch_chunk_gated_delta_rule(1216+ core_attn_out, _ = chunk_gated_delta_rule(
914 query,1217 query,
915 key,1218 key,
916 value,1219 value,
@@ -919,6 +1222,7 @@ class LinearAttentionUlyssesCPWrapper(nn.Module):
919 initial_state=None,1222 initial_state=None,
920 output_final_state=False,1223 output_final_state=False,
921 use_qk_l2norm_in_kernel=True,1224 use_qk_l2norm_in_kernel=True,
1225+ backend=self.gdn_backend,
922 )1226 )
923 1227 
924 core_attn_out = self._head_to_seq(core_attn_out)1228 core_attn_out = self._head_to_seq(core_attn_out)
@@ -934,9 +1238,16 @@ class LinearAttentionUlyssesCPWrapper(nn.Module):
934class LinearAttentionP2PCPWrapper(nn.Module):1238class LinearAttentionP2PCPWrapper(nn.Module):
935 """Sequence-sharded GDN CP with an affine-summary state wavefront."""1239 """Sequence-sharded GDN CP with an affine-summary state wavefront."""
936 1240 
937- def __init__(self, module: nn.Module, device_mesh: DeviceMesh):1241+ def __init__(
1242+ self,
1243+ module: nn.Module,
1244+ device_mesh: DeviceMesh,
1245+ *,
1246+ backend: str = "eager",
1247+ ):
938 super().__init__()1248 super().__init__()
939 self.module = module1249 self.module = module
1250+ self.gdn_backend = backend
940 self.cp_mesh = _ensure_1d(device_mesh)1251 self.cp_mesh = _ensure_1d(device_mesh)
941 self.cp_size = self.cp_mesh.size()1252 self.cp_size = self.cp_mesh.size()
942 self.cp_rank = self.cp_mesh.get_local_rank()1253 self.cp_rank = self.cp_mesh.get_local_rank()
@@ -944,6 +1255,13 @@ class LinearAttentionP2PCPWrapper(nn.Module):
944 1255 
945 def _validate_module(self) -> None:1256 def _validate_module(self) -> None:
946 """Validate the Conv1d requirements of the P2P CP path."""1257 """Validate the Conv1d requirements of the P2P CP path."""
1258+ if self.gdn_backend == "triton" and (
1259+ self.module.head_k_dim != 128 or self.module.head_v_dim != 128
1260+ ):
1261+ raise NotImplementedError(
1262+ "linear attention P2P Triton backend requires "
1263+ "head_k_dim=head_v_dim=128."
1264+ )
947 conv = self.module.conv1d1265 conv = self.module.conv1d
948 if conv.stride != (1,):1266 if conv.stride != (1,):
949 raise ValueError(1267 raise ValueError(
@@ -1015,17 +1333,41 @@ class LinearAttentionP2PCPWrapper(nn.Module):
1015 query = query.repeat_interleave(base.kv_groups, dim=2)1333 query = query.repeat_interleave(base.kv_groups, dim=2)
1016 key = key.repeat_interleave(base.kv_groups, dim=2)1334 key = key.repeat_interleave(base.kv_groups, dim=2)
1017 1335 
1018- core_attn_out = _gdn_state_p2p_summary(1336+ if self.gdn_backend == "triton":
1019- query,1337+ if local_seq_len % 64 != 0:
1020- key,1338+ raise NotImplementedError(
1021- value,1339+ "linear attention P2P Triton backend requires each CP "
1022- g,1340+ f"rank's local sequence length ({local_seq_len}) to be "
1023- beta,1341+ "divisible by 64."
1024- self.cp_mesh,1342+ )
1025- self.cp_rank,1343+ if not is_triton_gdn_available(query, key, value, g, beta):
1026- self.cp_size,1344+ raise RuntimeError(
1027- use_qk_l2norm_in_kernel=True,1345+ "linear attention P2P Triton backend requires an NPU "
1028- )1346+ "input satisfying the fixed GDN contract and a validated "
1347+ "triton-ascend 3.2.x installation."
1348+ )
1349+ core_attn_out = _gdn_state_p2p_triton(
1350+ query,
1351+ key,
1352+ value,
1353+ g,
1354+ beta,
1355+ self.cp_mesh,
1356+ self.cp_rank,
1357+ self.cp_size,
1358+ )
1359+ else:
1360+ core_attn_out = _gdn_state_p2p_summary(
1361+ query,
1362+ key,
1363+ value,
1364+ g,
1365+ beta,
1366+ self.cp_mesh,
1367+ self.cp_rank,
1368+ self.cp_size,
1369+ use_qk_l2norm_in_kernel=True,
1370+ )
1029 1371 
1030 core_attn_out = core_attn_out.reshape(-1, base.head_v_dim)1372 core_attn_out = core_attn_out.reshape(-1, base.head_v_dim)
1031 z_flat = z.reshape(-1, base.head_v_dim)1373 z_flat = z.reshape(-1, base.head_v_dim)
@@ -1147,22 +1489,41 @@ class LinearAttentionAllGatherCPWrapper(nn.Module):
1147class LinearAttentionContextParallel(ParallelStyle):1489class LinearAttentionContextParallel(ParallelStyle):
1148 """Apply context parallel execution to a Gated DeltaNet module."""1490 """Apply context parallel execution to a Gated DeltaNet module."""
1149 1491 
1150- def __init__(self, *, mode: str = "ulysses") -> None:1492+ def __init__(self, *, mode: str = "ulysses", backend: str = "eager") -> None:
1151 if mode not in {"ulysses", "p2p", "all_gather"}:1493 if mode not in {"ulysses", "p2p", "all_gather"}:
1152 raise NotImplementedError(1494 raise NotImplementedError(
1153 "LinearAttentionContextParallel currently supports mode='ulysses', "1495 "LinearAttentionContextParallel currently supports mode='ulysses', "
1154 "mode='p2p', and mode='all_gather'."1496 "mode='p2p', and mode='all_gather'."
1155 )1497 )
1498+ if backend not in {"eager", "triton"}:
1499+ raise ValueError(
1500+ "LinearAttentionContextParallel backend must be 'eager' or "
1501+ f"'triton', got {backend!r}."
1502+ )
1503+ if mode == "all_gather" and backend == "triton":
1504+ raise NotImplementedError(
1505+ "linear attention all-gather CP does not yet support the "
1506+ "Triton backend."
1507+ )
1156 self.mode = mode1508 self.mode = mode
1509+ self.backend = backend
1157 1510 
1158 def apply(self, module: nn.Module, device_mesh: DeviceMesh) -> nn.Module:1511 def apply(self, module: nn.Module, device_mesh: DeviceMesh) -> nn.Module:
1159 """Patch ``module.forward`` with a linear-attention CP executor."""1512 """Patch ``module.forward`` with a linear-attention CP executor."""
1160 if self.mode == "ulysses":1513 if self.mode == "ulysses":
1161- executor = LinearAttentionUlyssesCPWrapper(module, device_mesh)1514+ executor = LinearAttentionUlyssesCPWrapper(
1515+ module,
1516+ device_mesh,
1517+ backend=self.backend,
1518+ )
1162 elif self.mode == "all_gather":1519 elif self.mode == "all_gather":
1163 executor = LinearAttentionAllGatherCPWrapper(module, device_mesh)1520 executor = LinearAttentionAllGatherCPWrapper(module, device_mesh)
1164 else:1521 else:
1165- executor = LinearAttentionP2PCPWrapper(module, device_mesh)1522+ executor = LinearAttentionP2PCPWrapper(
1523+ module,
1524+ device_mesh,
1525+ backend=self.backend,
1526+ )
1166 object.__setattr__(module, "_hp_linear_attention_cp_executor", executor)1527 object.__setattr__(module, "_hp_linear_attention_cp_executor", executor)
1167 object.__setattr__(module, "_hp_linear_attention_original_forward", module.forward)1528 object.__setattr__(module, "_hp_linear_attention_original_forward", module.forward)
1168 1529 
@@ -33,8 +33,15 @@ This module is for **training** (no KV-cache, no chunk recurrence reuse).
33For inference with cache, use a kernel-optimised path or the recurrent33For inference with cache, use a kernel-optimised path or the recurrent
34variant from ``transformers.models.qwen3_next``.34variant from ``transformers.models.qwen3_next``.
35"""35"""
36+# This model module currently has a Torch-only implementation.
37+# pylint: disable=forbidden-backend-import,missing-public-type-hints
38+# pylint: disable=missing-public-docstring,not-callable
36# pylint: disable=C0103 # SSM/state-space convention: A_log, A39# pylint: disable=C0103 # SSM/state-space convention: A_log, A
37 40 
41+import importlib
42+import importlib.metadata
43+import importlib.util
44+import re
38from typing import Optional45from typing import Optional
39 46 
40import torch47import torch
@@ -44,6 +51,14 @@ from torch.nn import functional as F
44from hyper_parallel.models.modules.rmsnorm import RMSNormGated51from hyper_parallel.models.modules.rmsnorm import RMSNormGated
45 52 
46 53 
54+_GDN_BACKENDS = frozenset({"eager", "triton"})
55+_MIN_TRITON_ASCEND_VERSION = (3, 2, 1)
56+_MAX_TRITON_ASCEND_VERSION = (3, 3, 0)
57+_SUPPORTED_TRITON_MODULE_SERIES = (3, 2)
58+_TRITON_GDN_HEAD_DIM = 128
59+_TRITON_GDN_CHUNK_SIZE = 64
60+ 
61+ 
47def _l2norm(x: torch.Tensor, dim: int = -1, eps: float = 1e-6) -> torch.Tensor:62def _l2norm(x: torch.Tensor, dim: int = -1, eps: float = 1e-6) -> torch.Tensor:
48 """L2-normalize along ``dim``."""63 """L2-normalize along ``dim``."""
49 # ``(x * x).sum`` (MulBackward) instead of ``x.pow(2).sum`` (PowBackward):64 # ``(x * x).sum`` (MulBackward) instead of ``x.pow(2).sum`` (PowBackward):
@@ -149,6 +164,133 @@ def torch_chunk_gated_delta_rule(
149 return core_attn_out, last_recurrent_state164 return core_attn_out, last_recurrent_state
150 165 
151 166 
167+def _parse_version(version_text: str) -> tuple[int, int, int]:
168+ """Return a three-component numeric version tuple."""
169+ parts = [
170+ int(match.group())
171+ for part in version_text.split("+")[0].split(".")[:3]
172+ if (match := re.match(r"\d+", part)) is not None
173+ ]
174+ return tuple((parts + [0, 0, 0])[:3])
175+ 
176+ 
177+def _is_triton_gdn_input_supported(
178+ query: torch.Tensor,
179+ key: Optional[torch.Tensor],
180+ value: Optional[torch.Tensor],
181+ g: Optional[torch.Tensor],
182+ beta: Optional[torch.Tensor],
183+) -> bool:
184+ """Check the fixed Qwen3.5 GDN contract validated by this backend."""
185+ if key is None or value is None or g is None or beta is None:
186+ return False
187+ if not (
188+ query.device.type == "npu"
189+ and query.dtype == key.dtype == value.dtype == beta.dtype == torch.bfloat16
190+ and g.dtype == torch.float32
191+ and query.ndim == key.ndim == value.ndim == 4
192+ and g.ndim == beta.ndim == 3
193+ and query.shape == key.shape
194+ and query.shape[:3] == value.shape[:3] == g.shape == beta.shape
195+ and query.shape[-1] == value.shape[-1] == _TRITON_GDN_HEAD_DIM
196+ ):
197+ return False
198+ return all(
199+ tensor.device == query.device
200+ for tensor in (key, value, g, beta)
201+ )
202+ 
203+ 
204+def is_triton_gdn_available(
205+ query: Optional[torch.Tensor] = None,
206+ key: Optional[torch.Tensor] = None,
207+ value: Optional[torch.Tensor] = None,
208+ g: Optional[torch.Tensor] = None,
209+ beta: Optional[torch.Tensor] = None,
210+ chunk_size: int = _TRITON_GDN_CHUNK_SIZE,
211+) -> bool:
212+ """Return whether the validated Triton-Ascend GDN backend is available."""
213+ if chunk_size != _TRITON_GDN_CHUNK_SIZE:
214+ return False
215+ if query is not None and not _is_triton_gdn_input_supported(
216+ query, key, value, g, beta
217+ ):
218+ return False
219+ try:
220+ version_text = importlib.metadata.version("triton-ascend")
221+ except importlib.metadata.PackageNotFoundError:
222+ return False
223+ version = _parse_version(version_text)
224+ if not _MIN_TRITON_ASCEND_VERSION <= version < _MAX_TRITON_ASCEND_VERSION:
225+ return False
226+ try:
227+ triton_module = importlib.import_module("triton")
228+ module_version = _parse_version(triton_module.__version__)
229+ if module_version[:2] != _SUPPORTED_TRITON_MODULE_SERIES:
230+ return False
231+ return importlib.util.find_spec("triton.backends.ascend") is not None
232+ except (AttributeError, ImportError, ModuleNotFoundError):
233+ return False
234+ 
235+ 
236+def chunk_gated_delta_rule(
237+ query: torch.Tensor,
238+ key: torch.Tensor,
239+ value: torch.Tensor,
240+ g: torch.Tensor,
241+ beta: torch.Tensor,
242+ chunk_size: int = 64,
243+ initial_state: Optional[torch.Tensor] = None,
244+ output_final_state: bool = False,
245+ use_qk_l2norm_in_kernel: bool = False,
246+ backend: str = "eager",
247+):
248+ """Dispatch GDN to an explicitly selected eager or Triton backend."""
249+ backend = backend.lower()
250+ if backend not in _GDN_BACKENDS:
251+ raise ValueError(
252+ f"unsupported GDN backend {backend!r}; "
253+ f"expected one of {sorted(_GDN_BACKENDS)}."
254+ )
255+ if backend == "eager":
256+ return torch_chunk_gated_delta_rule(
257+ query,
258+ key,
259+ value,
260+ g=g,
261+ beta=beta,
262+ chunk_size=chunk_size,
263+ initial_state=initial_state,
264+ output_final_state=output_final_state,
265+ use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
266+ )
267+ if not is_triton_gdn_available(
268+ query, key, value, g, beta, chunk_size=chunk_size
269+ ):
270+ raise RuntimeError(
271+ "GDN backend='triton' requires a validated triton-ascend 3.2.x "
272+ "installation with the Ascend backend and NPU inputs "
273+ "q/k/v/beta=bf16, g=fp32, head_k_dim=head_v_dim=128, "
274+ "and chunk_size=64."
275+ )
276+ 
277+ from hyper_parallel.platform.torch.custom_ops.gdn.chunk_gated_delta_rule import ( # pylint: disable=import-outside-toplevel
278+ chunk_gated_delta_rule as triton_chunk_gated_delta_rule,
279+ )
280+ 
281+ return triton_chunk_gated_delta_rule(
282+ query,
283+ key,
284+ value,
285+ g,
286+ beta,
287+ initial_state=initial_state,
288+ output_final_state=output_final_state,
289+ use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
290+ chunk_size=chunk_size,
291+ )
292+ 
293+ 
152class GatedDeltaNet(nn.Module):294class GatedDeltaNet(nn.Module):
153 """Gated DeltaNet linear-attention block (Qwen3.5 / Qwen3-Next style).295 """Gated DeltaNet linear-attention block (Qwen3.5 / Qwen3-Next style).
154 296 
@@ -19,6 +19,8 @@ parameter shards on ``head`` or ``feature`` dimensions and norms on the
19sequence axis so per-step gradients stay slice-faithful to the single-card19sequence axis so per-step gradients stay slice-faithful to the single-card
20run.20run.
21"""21"""
22+# Qwen3.5 is currently registered as a Torch model implementation.
23+# pylint: disable=forbidden-backend-import
22from dataclasses import replace24from dataclasses import replace
23from types import SimpleNamespace25from types import SimpleNamespace
24from typing import Optional, TYPE_CHECKING26from typing import Optional, TYPE_CHECKING
@@ -476,9 +478,14 @@ def qwen3_5_tp_load_transforms(
476 return transforms478 return transforms
477 479 
478 480 
479-def _apply_linear_attention_cp(module: nn.Module, cp_mesh: DeviceMesh, mode: str) -> None:481+def _apply_linear_attention_cp(
482+ module: nn.Module,
483+ cp_mesh: DeviceMesh,
484+ mode: str,
485+ backend: str,
486+) -> None:
480 """Apply CP to a Qwen3.5 linear-attention module."""487 """Apply CP to a Qwen3.5 linear-attention module."""
481- LinearAttentionContextParallel(mode=mode).apply(module, cp_mesh)488+ LinearAttentionContextParallel(mode=mode, backend=backend).apply(module, cp_mesh)
482 489 
483 490 
484def _validate_qwen3_5_tp_config(model: Qwen3_5ForCausalLM, tp_world: int) -> None:491def _validate_qwen3_5_tp_config(model: Qwen3_5ForCausalLM, tp_world: int) -> None:
@@ -680,6 +687,7 @@ def parallelize_qwen3_5_cp(
680 *,687 *,
681 ulysses_degree: Optional[int] = None,688 ulysses_degree: Optional[int] = None,
682 linear_attention_cp_mode: str = "ulysses",689 linear_attention_cp_mode: str = "ulysses",
690+ linear_attention_gdn_backend: str = "eager",
683) -> Qwen3_5ForCausalLM:691) -> Qwen3_5ForCausalLM:
684 """Apply context parallelism across the Qwen3.5 hybrid decoder.692 """Apply context parallelism across the Qwen3.5 hybrid decoder.
685 693 
@@ -697,11 +705,11 @@ def parallelize_qwen3_5_cp(
697 rank ends up with the full sequence on a head-shard and a square causal705 rank ends up with the full sequence on a head-shard and a square causal
698 mask is correct again.706 mask is correct again.
699 707 
700- Linear-attention (:class:`Qwen3_5GatedDeltaNet`) layers use a matching708+ Linear-attention (:class:`Qwen3_5GatedDeltaNet`) layers select Ulysses,
701- pure-Ulysses execution wrapper: project local sequence shards, all-to-all709+ State-P2P, or all-gather execution with ``linear_attention_cp_mode``.
702- the projected Q/K/V/B/A tensors to full-sequence local-head shards, run710+ ``linear_attention_gdn_backend`` explicitly selects the eager or Triton
703- the per-head conv and gated delta rule on local heads, then all-to-all the711+ local GDN implementation; unsupported combinations fail instead of
704- result back to sequence shards before the output projection.712+ silently falling back.
705 """713 """
706 # Only pure Ulysses is wired here; a smaller ``ulysses_degree`` makes each714 # Only pure Ulysses is wired here; a smaller ``ulysses_degree`` makes each
707 # rank attend over gathered K/V with ``is_causal=True`` but without a715 # rank attend over gathered K/V with ``is_causal=True`` but without a
@@ -724,12 +732,18 @@ def parallelize_qwen3_5_cp(
724 cp_plan.apply(block.self_attn.sdpa_core, cp_mesh)732 cp_plan.apply(block.self_attn.sdpa_core, cp_mesh)
725 full_attached += 1733 full_attached += 1
726 else:734 else:
727- _apply_linear_attention_cp(block.linear_attn, cp_mesh, linear_attention_cp_mode)735+ _apply_linear_attention_cp(
736+ block.linear_attn,
737+ cp_mesh,
738+ linear_attention_cp_mode,
739+ linear_attention_gdn_backend,
740+ )
728 linear_attached += 1741 linear_attached += 1
729 logger.info_rank0(742 logger.info_rank0(
730 "CP applied to Qwen3.5: cp_size=%d, ulysses_degree=%s, full-attn hooks=%d, "743 "CP applied to Qwen3.5: cp_size=%d, ulysses_degree=%s, full-attn hooks=%d, "
731- "linear-attn %s hooks=%d",744+ "linear-attn %s/%s hooks=%d",
732- cp_mesh.size(), ulysses_degree, full_attached, linear_attention_cp_mode, linear_attached,745+ cp_mesh.size(), ulysses_degree, full_attached, linear_attention_cp_mode,
746+ linear_attention_gdn_backend, linear_attached,
733 )747 )
734 return model748 return model
735 749 
@@ -1292,11 +1306,17 @@ def parallelize_qwen3_5(
1292 "linear_attention_cp_mode",1306 "linear_attention_cp_mode",
1293 "ulysses",1307 "ulysses",
1294 )1308 )
1309+ linear_attention_gdn_backend = getattr(
1310+ cfg.train.accelerator,
1311+ "linear_attention_gdn_backend",
1312+ "eager",
1313+ )
1295 parallelize_qwen3_5_cp(1314 parallelize_qwen3_5_cp(
1296 model,1315 model,
1297 cp_mesh,1316 cp_mesh,
1298 ulysses_degree=ulysses_degree,1317 ulysses_degree=ulysses_degree,
1299 linear_attention_cp_mode=linear_attention_cp_mode,1318 linear_attention_cp_mode=linear_attention_cp_mode,
1319+ linear_attention_gdn_backend=linear_attention_gdn_backend,
1300 )1320 )
1301 1321 
1302 if cfg.train.accelerator.ep > 1:1322 if cfg.train.accelerator.ep > 1:
@@ -0,0 +1,21 @@
1+MIT License
2+ 
3+Copyright (c) 2023-2026 Songlin Yang, Yu Zhang, Zhiyuan Li
4+ 
5+Permission is hereby granted, free of charge, to any person obtaining a copy
6+of this software and associated documentation files (the "Software"), to deal
7+in the Software without restriction, including without limitation the rights
8+to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9+copies of the Software, and to permit persons to whom the Software is
10+furnished to do so, subject to the following conditions:
11+ 
12+The above copyright notice and this permission notice shall be included in all
13+copies or substantial portions of the Software.
14+ 
15+THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16+IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17+FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18+AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19+LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20+OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21+SOFTWARE.
@@ -0,0 +1,19 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""Triton-Ascend Gated DeltaNet operators for the Torch platform.
16+ 
17+Callers import concrete submodules lazily so importing Hyper-Parallel on CPU
18+or MindSpore installations does not require Triton.
19+"""
@@ -0,0 +1,196 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""Affine state-summary operations used by GDN State-P2P."""
16+ 
17+from typing import Optional
18+ 
19+import torch
20+import triton
21+ 
22+from .triton.state_summary import (
23+ gdn_packed_state_summary_kernel,
24+ gdn_state_grad_ext_kernel,
25+)
26+ 
27+ 
28+def _validate_fixed_summary_shape(
29+ key_dim: int,
30+ value_dim: int,
31+ chunk_size: int,
32+) -> None:
33+ if key_dim != 128 or value_dim != 128 or chunk_size != 64:
34+ raise NotImplementedError(
35+ "Triton GDN state summary requires key_dim=value_dim=128 and "
36+ "chunk_size=64."
37+ )
38+ 
39+ 
40+@torch.compiler.disable
41+def chunk_gated_delta_rule_state_summary_fwd(
42+ key: torch.Tensor,
43+ w: torch.Tensor,
44+ u: torch.Tensor,
45+ g: torch.Tensor,
46+ *,
47+ chunk_size: int = 64,
48+ block_size: int = 128,
49+) -> tuple[torch.Tensor, torch.Tensor]:
50+ """Return the local affine map ``state_out = M @ state_in + S``."""
51+ if key.ndim != 4 or w.ndim != 4 or u.ndim != 4 or g.ndim != 3:
52+ raise ValueError("GDN state summary expects key/w/u [B,T,H,D] and g [B,T,H].")
53+ batch, seq_len, heads, key_dim = key.shape
54+ value_dim = u.shape[-1]
55+ _validate_fixed_summary_shape(key_dim, value_dim, chunk_size)
56+ if block_size not in (64, 128):
57+ raise ValueError(f"GDN state-summary block_size must be 64 or 128, got {block_size}.")
58+ if seq_len % chunk_size != 0:
59+ raise ValueError(
60+ f"GDN state-summary sequence length {seq_len} must be divisible by {chunk_size}."
61+ )
62+ if w.shape != key.shape or u.shape[:3] != key.shape[:3] or g.shape != key.shape[:3]:
63+ raise ValueError(
64+ "Incompatible GDN state-summary shapes: "
65+ f"key={tuple(key.shape)}, w={tuple(w.shape)}, "
66+ f"u={tuple(u.shape)}, g={tuple(g.shape)}."
67+ )
68+ 
69+ key, w, u, g = (tensor.contiguous() for tensor in (key, w, u, g))
70+ packed_summary = torch.empty(
71+ batch,
72+ heads,
73+ key_dim,
74+ value_dim + key_dim,
75+ device=key.device,
76+ dtype=torch.float32,
77+ )
78+ gdn_packed_state_summary_kernel[
79+ (triton.cdiv(value_dim + key_dim, block_size), batch * heads)
80+ ](
81+ key,
82+ w,
83+ u,
84+ g,
85+ packed_summary,
86+ seq_len,
87+ H=heads,
88+ K=key_dim,
89+ V=value_dim,
90+ BT=chunk_size,
91+ BV=block_size,
92+ NT=seq_len // chunk_size,
93+ )
94+ state_ext = packed_summary[..., :value_dim].contiguous()
95+ transition = packed_summary[..., value_dim:].contiguous()
96+ return state_ext, transition
97+ 
98+ 
99+@torch.compiler.disable
100+def chunk_gated_delta_rule_state_gradient_summary_bwd(
101+ query: torch.Tensor,
102+ key: torch.Tensor,
103+ w: torch.Tensor,
104+ g: torch.Tensor,
105+ grad_output: torch.Tensor,
106+ dv: torch.Tensor,
107+ scale: float,
108+ *,
109+ chunk_size: int = 64,
110+) -> torch.Tensor:
111+ """Return the local-loss contribution to the incoming state gradient."""
112+ batch, seq_len, heads, key_dim = query.shape
113+ value_dim = grad_output.shape[-1]
114+ _validate_fixed_summary_shape(key_dim, value_dim, chunk_size)
115+ if seq_len % chunk_size != 0:
116+ raise ValueError(
117+ f"GDN state-gradient sequence length {seq_len} must be divisible by {chunk_size}."
118+ )
119+ qk_shape = (batch, seq_len, heads, key_dim)
120+ value_shape = (batch, seq_len, heads, value_dim)
121+ if (
122+ key.shape != qk_shape
123+ or w.shape != qk_shape
124+ or g.shape != qk_shape[:3]
125+ or grad_output.shape != value_shape
126+ or dv.shape != value_shape
127+ ):
128+ raise ValueError(
129+ "Incompatible GDN state-gradient summary shapes: "
130+ f"query={tuple(query.shape)}, key={tuple(key.shape)}, "
131+ f"w={tuple(w.shape)}, g={tuple(g.shape)}, "
132+ f"grad_output={tuple(grad_output.shape)}, dv={tuple(dv.shape)}."
133+ )
134+ 
135+ query, key, w, g, grad_output, dv = (
136+ tensor.contiguous() for tensor in (query, key, w, g, grad_output, dv)
137+ )
138+ grad_state_ext = torch.empty(
139+ batch,
140+ heads,
141+ key_dim,
142+ value_dim,
143+ device=query.device,
144+ dtype=torch.float32,
145+ )
146+ gdn_state_grad_ext_kernel[(1, batch * heads)](
147+ query,
148+ key,
149+ w,
150+ g,
151+ grad_output,
152+ dv,
153+ grad_state_ext,
154+ scale,
155+ seq_len,
156+ H=heads,
157+ K=key_dim,
158+ V=value_dim,
159+ BT=chunk_size,
160+ BV=128,
161+ NT=seq_len // chunk_size,
162+ )
163+ return grad_state_ext
164+ 
165+ 
166+def apply_gdn_state_summary(
167+ state_ext: torch.Tensor,
168+ transition: torch.Tensor,
169+ initial_state: Optional[torch.Tensor],
170+) -> torch.Tensor:
171+ """Apply a local affine state summary in FP32."""
172+ if initial_state is None:
173+ return state_ext
174+ return torch.matmul(transition, initial_state.float()) + state_ext
175+ 
176+ 
177+def apply_gdn_state_gradient_summary(
178+ grad_state_ext: torch.Tensor,
179+ transition: torch.Tensor,
180+ grad_final_state: Optional[torch.Tensor],
181+) -> torch.Tensor:
182+ """Apply the adjoint affine summary to a gradient from the next rank."""
183+ if grad_final_state is None:
184+ return grad_state_ext
185+ return (
186+ torch.matmul(transition.transpose(-2, -1), grad_final_state.float())
187+ + grad_state_ext
188+ )
189+ 
190+ 
191+__all__ = [
192+ "apply_gdn_state_gradient_summary",
193+ "apply_gdn_state_summary",
194+ "chunk_gated_delta_rule_state_gradient_summary_bwd",
195+ "chunk_gated_delta_rule_state_summary_fwd",
196+]
@@ -0,0 +1,15 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+"""Internal Triton-Ascend kernels for Gated DeltaNet."""
@@ -0,0 +1,592 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+# -*- coding: utf-8 -*-
16+# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
17+ 
18+# pylint: disable=line-too-long,missing-public-type-hints,missing-public-docstring
19+# pylint: disable=used-before-assignment,unsupported-binary-operation,unused-argument
20+# pylint: disable=invalid-name,missing-module-docstring,missing-function-docstring
21+ 
22+from typing import Optional, Tuple
23+ 
24+import torch
25+import triton
26+import triton.language as tl
27+ 
28+from .utils import prepare_chunk_indices, prepare_chunk_offsets, get_autotune_config, get_npu_properties
29+ 
30+CUBE_CORE_NUM = get_npu_properties()['num_aicore']
31+ 
32+ 
33+@triton.heuristics({
34+ 'USE_G': lambda args: args['g'] is not None,
35+ 'USE_GK': lambda args: args['gk'] is not None,
36+ 'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
37+ 'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
38+ 'SAVE_NEW_VALUE': lambda args: args['v_new'] is not None,
39+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
40+})
41+@triton.autotune(
42+ configs=get_autotune_config(multibuffer_list=(False,)),
43+ key=['H', 'K', 'V', 'BT'],
44+)
45+@triton.jit(do_not_specialize=['T'])
46+def chunk_gated_delta_rule_fwd_kernel_h_blockdim64(
47+ k,
48+ v,
49+ w,
50+ v_new,
51+ g,
52+ gk,
53+ h,
54+ h0,
55+ ht,
56+ cu_seqlens,
57+ chunk_offsets,
58+ T,
59+ H: tl.constexpr,
60+ K: tl.constexpr,
61+ V: tl.constexpr,
62+ BT: tl.constexpr,
63+ BV: tl.constexpr,
64+ NT: tl.constexpr,
65+ USE_G: tl.constexpr,
66+ USE_GK: tl.constexpr,
67+ USE_INITIAL_STATE: tl.constexpr,
68+ STORE_FINAL_STATE: tl.constexpr,
69+ SAVE_NEW_VALUE: tl.constexpr,
70+ IS_VARLEN: tl.constexpr,
71+):
72+ T_all = T
73+ NT_all = NT
74+ i_v, i_nh = tl.program_id(0), tl.program_id(1)
75+ i_n, i_h = i_nh // H, i_nh % H
76+ if IS_VARLEN:
77+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
78+ T = eos - bos
79+ NT = tl.cdiv(T, BT)
80+ boh = tl.load(chunk_offsets + i_n).to(tl.int32)
81+ else:
82+ bos, eos = i_n * T, i_n * T + T
83+ NT = tl.cdiv(T, BT)
84+ boh = i_n * NT
85+ 
86+ # Initialize hidden states
87+ b_h1 = tl.zeros([64, BV], dtype=tl.float32)
88+ if K > 64:
89+ b_h2 = tl.zeros([64, BV], dtype=tl.float32)
90+ if K > 128:
91+ b_h3 = tl.zeros([64, BV], dtype=tl.float32)
92+ if K > 192:
93+ b_h4 = tl.zeros([64, BV], dtype=tl.float32)
94+ 
95+ if IS_VARLEN:
96+ v = v + (i_h * T_all + bos) * V
97+ k = k + (i_h * T_all + bos) * K
98+ w = w + (i_h * T_all + bos) * K
99+ g = g + i_h * T_all + bos
100+ h = h + (i_h * NT_all + boh) * K * V
101+ if SAVE_NEW_VALUE:
102+ v_new_base = v_new + (i_h * T_all + bos) * V
103+ else:
104+ v = v + (i_n * H + i_h) * T * V
105+ k = k + (i_n * H + i_h) * T * K
106+ w = w + (i_n * H + i_h) * T * K
107+ g = g + (i_n * H + i_h) * T
108+ h = h + (i_n * H + i_h) * NT * K * V
109+ if SAVE_NEW_VALUE:
110+ v_new_base = v_new + (i_n * H + i_h) * T * V
111+ 
112+ if USE_INITIAL_STATE:
113+ h0_ptr = h0 + i_nh * K * V
114+ if STORE_FINAL_STATE:
115+ ht_ptr = ht + i_nh * K * V
116+ 
117+ # Load initial state
118+ if USE_INITIAL_STATE:
119+ p_h0_1 = tl.make_block_ptr(h0_ptr, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
120+ b_h1 += tl.load(p_h0_1, boundary_check=(0, 1)).to(tl.float32)
121+ if K > 64:
122+ p_h0_2 = tl.make_block_ptr(h0_ptr, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
123+ b_h2 += tl.load(p_h0_2, boundary_check=(0, 1)).to(tl.float32)
124+ if K > 128:
125+ p_h0_3 = tl.make_block_ptr(h0_ptr, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
126+ b_h3 += tl.load(p_h0_3, boundary_check=(0, 1)).to(tl.float32)
127+ if K > 192:
128+ p_h0_4 = tl.make_block_ptr(h0_ptr, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
129+ b_h4 += tl.load(p_h0_4, boundary_check=(0, 1)).to(tl.float32)
130+ 
131+ # Main recurrence over chunks
132+ for i_t in range(NT):
133+ # Store current hidden state h_t
134+ p_h1 = tl.make_block_ptr(h + i_t * K * V, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
135+ tl.store(p_h1, b_h1.to(p_h1.dtype.element_ty), boundary_check=(0, 1))
136+ if K > 64:
137+ p_h2 = tl.make_block_ptr(h + i_t * K * V, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
138+ tl.store(p_h2, b_h2.to(p_h2.dtype.element_ty), boundary_check=(0, 1))
139+ if K > 128:
140+ p_h3 = tl.make_block_ptr(h + i_t * K * V, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
141+ tl.store(p_h3, b_h3.to(p_h3.dtype.element_ty), boundary_check=(0, 1))
142+ if K > 192:
143+ p_h4 = tl.make_block_ptr(h + i_t * K * V, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
144+ tl.store(p_h4, b_h4.to(p_h4.dtype.element_ty), boundary_check=(0, 1))
145+ 
146+ # Compute v_residual = v - w @ h
147+ p_w = tl.make_block_ptr(w, (T, K), (K, 1), (i_t * BT, 0), (BT, 64), (1, 0))
148+ b_w = tl.load(p_w, boundary_check=(0, 1))
149+ b_v = tl.dot(b_w, b_h1.to(b_w.dtype))
150+ if K > 64:
151+ p_w = tl.make_block_ptr(w, (T, K), (K, 1), (i_t * BT, 64), (BT, 64), (1, 0))
152+ b_w = tl.load(p_w, boundary_check=(0, 1))
153+ b_v += tl.dot(b_w, b_h2.to(b_w.dtype))
154+ if K > 128:
155+ p_w = tl.make_block_ptr(w, (T, K), (K, 1), (i_t * BT, 128), (BT, 64), (1, 0))
156+ b_w = tl.load(p_w, boundary_check=(0, 1))
157+ b_v += tl.dot(b_w, b_h3.to(b_w.dtype))
158+ if K > 192:
159+ p_w = tl.make_block_ptr(w, (T, K), (K, 1), (i_t * BT, 192), (BT, 64), (1, 0))
160+ b_w = tl.load(p_w, boundary_check=(0, 1))
161+ b_v += tl.dot(b_w, b_h4.to(b_w.dtype))
162+ 
163+ p_v = tl.make_block_ptr(v, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
164+ b_v = tl.load(p_v, boundary_check=(0, 1)) - b_v
165+ 
166+ if SAVE_NEW_VALUE:
167+ p_v_new = tl.make_block_ptr(v_new_base, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
168+ tl.store(p_v_new, b_v.to(p_v_new.dtype.element_ty), boundary_check=(0, 1))
169+ 
170+ last_idx = min((i_t + 1) * BT, T) - 1
171+ 
172+ # Apply output gate g
173+ if USE_G:
174+ m_t = (i_t * BT + tl.arange(0, BT)).to(tl.float32) < T
175+ b_g_last = tl.load(g + last_idx)
176+ p_g = tl.make_block_ptr(g, (T,), (1,), (i_t * BT,), (BT,), (0,))
177+ b_g = tl.load(p_g, boundary_check=(0,))
178+ b_v *= (m_t * tl.exp(b_g_last - b_g))[:, None]
179+ b_g_last_exp = tl.exp(b_g_last)
180+ b_h1 *= b_g_last_exp
181+ if K > 64:
182+ b_h2 *= b_g_last_exp
183+ if K > 128:
184+ b_h3 *= b_g_last_exp
185+ if K > 192:
186+ b_h4 *= b_g_last_exp
187+ 
188+ # Apply key gate gk
189+ if USE_GK:
190+ o_k1 = tl.arange(0, 64).to(tl.float32)
191+ gk_base_ptr = gk + (i_n * H + i_h) * T * K
192+ b_gk_last1 = tl.load(gk_base_ptr + last_idx * K + o_k1, mask=(o_k1 < K), other=0.)
193+ b_h1 *= tl.exp(b_gk_last1)[:, None]
194+ if K > 64:
195+ o_k2 = 64 + o_k1
196+ b_gk_last2 = tl.load(gk_base_ptr + last_idx * K + o_k2, mask=(o_k2 < K), other=0.)
197+ b_h2 *= tl.exp(b_gk_last2)[:, None]
198+ if K > 128:
199+ o_k3 = 128 + o_k1
200+ b_gk_last3 = tl.load(gk_base_ptr + last_idx * K + o_k3, mask=(o_k3 < K), other=0.)
201+ b_h3 *= tl.exp(b_gk_last3)[:, None]
202+ if K > 192:
203+ o_k4 = 192 + o_k1
204+ b_gk_last4 = tl.load(gk_base_ptr + last_idx * K + o_k4, mask=(o_k4 < K), other=0.)
205+ b_h4 *= tl.exp(b_gk_last4)[:, None]
206+ 
207+ b_v = b_v.to(k.dtype.element_ty)
208+ 
209+ # Update hidden state: h += k @ v
210+ p_k = tl.make_block_ptr(k, (K, T), (1, K), (0, i_t * BT), (64, BT), (0, 1))
211+ b_k = tl.load(p_k, boundary_check=(0, 1))
212+ if USE_GK:
213+ p_gk = tl.make_block_ptr(gk_base_ptr, (K, T), (1, K), (0, i_t * BT), (64, BT), (0, 1))
214+ b_k = (b_k * tl.exp(b_gk_last1[:, None] - tl.load(p_gk, boundary_check=(0, 1)))).to(b_k.dtype)
215+ b_h1 += tl.dot(b_k, b_v)
216+ 
217+ if K > 64:
218+ p_k = tl.make_block_ptr(k, (K, T), (1, K), (64, i_t * BT), (64, BT), (0, 1))
219+ b_k = tl.load(p_k, boundary_check=(0, 1))
220+ if USE_GK:
221+ p_gk = tl.make_block_ptr(gk_base_ptr, (K, T), (1, K), (64, i_t * BT), (64, BT), (0, 1))
222+ b_k = (b_k * tl.exp(b_gk_last2[:, None] - tl.load(p_gk, boundary_check=(0, 1)))).to(b_k.dtype)
223+ b_h2 += tl.dot(b_k, b_v)
224+ 
225+ if K > 128:
226+ p_k = tl.make_block_ptr(k, (K, T), (1, K), (128, i_t * BT), (64, BT), (0, 1))
227+ b_k = tl.load(p_k, boundary_check=(0, 1))
228+ if USE_GK:
229+ p_gk = tl.make_block_ptr(gk_base_ptr, (K, T), (1, K), (128, i_t * BT), (64, BT), (0, 1))
230+ b_k = (b_k * tl.exp(b_gk_last3[:, None] - tl.load(p_gk, boundary_check=(0, 1)))).to(b_k.dtype)
231+ b_h3 += tl.dot(b_k, b_v)
232+ 
233+ if K > 192:
234+ p_k = tl.make_block_ptr(k, (K, T), (1, K), (192, i_t * BT), (64, BT), (0, 1))
235+ b_k = tl.load(p_k, boundary_check=(0, 1))
236+ if USE_GK:
237+ p_gk = tl.make_block_ptr(gk_base_ptr, (K, T), (1, K), (192, i_t * BT), (64, BT), (0, 1))
238+ b_k = (b_k * tl.exp(b_gk_last4[:, None] - tl.load(p_gk, boundary_check=(0, 1)))).to(b_k.dtype)
239+ b_h4 += tl.dot(b_k, b_v)
240+ 
241+ # Store final state
242+ if STORE_FINAL_STATE:
243+ p_ht = tl.make_block_ptr(ht_ptr, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
244+ tl.store(p_ht, b_h1.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
245+ if K > 64:
246+ p_ht = tl.make_block_ptr(ht_ptr, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
247+ tl.store(p_ht, b_h2.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
248+ if K > 128:
249+ p_ht = tl.make_block_ptr(ht_ptr, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
250+ tl.store(p_ht, b_h3.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
251+ if K > 192:
252+ p_ht = tl.make_block_ptr(ht_ptr, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
253+ tl.store(p_ht, b_h4.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
254+ 
255+ 
256+def chunk_gated_delta_rule_fwd_h(
257+ k: torch.Tensor,
258+ w: torch.Tensor,
259+ u: torch.Tensor,
260+ g: Optional[torch.Tensor] = None,
261+ gk: Optional[torch.Tensor] = None,
262+ initial_state: Optional[torch.Tensor] = None,
263+ output_final_state: bool = False,
264+ chunk_size: int = 64, # default:64
265+ save_new_value: bool = True,
266+ cu_seqlens: Optional[torch.LongTensor] = None,
atomgit-bot
atomgit-botatomgit-bot8月5日

🔵 Low Priority

函数声明返回类型为 Tuple[torch.Tensor, torch.Tensor](两个张量),但第 310 行实际返回的是 h, v_new, final_state 三个张量。

changed line(第 266 行):-> Tuple[torch.Tensor, torch.Tensor] → 受影响的行为:静态类型检查器(mypy/pyright)和 IDE 会认为该函数只返回两个值,导致调用方类型推断错误。虽然当前所有调用方(chunk_gated_delta_rule_fwd_apply_state 等)都正确使用三元组解包,但标注与实现不一致会在重构或新调用方引入运行时报错(ValueError: too many values to unpack)。

建议:将返回类型从 Tuple[torch.Tensor, torch.Tensor] 改为 Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]],或直接改为 tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]](Python 3.9+ 风格)。

改动建议
266
- cu_seqlens: Optional[torch.LongTensor] = None,
266
+ ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
应用建议
likedislike
不准确?
267+) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
268+ B, T, H, K, V = *k.shape, u.shape[-1]
269+ BT = chunk_size
270+ 
271+ chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size) if cu_seqlens is not None else None
272+ # N: the actual number of sequences in the batch with either equal or variable lengths
273+ if cu_seqlens is None:
274+ N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
275+ else:
276+ N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
277+ assert K <= 256, "current kernel does not support head dimension larger than 256."
278+ 
279+ h = k.new_empty(B, NT, H, K, V).permute(0, 2, 1, 3, 4).contiguous()
280+ final_state = k.new_empty(N, H, K, V, dtype=torch.float32) if output_final_state else None
281+ 
282+ BV = 128
283+ 
284+ v_new = torch.empty_like(u).permute(0, 2, 1, 3).contiguous() if save_new_value else None
285+ k = k.permute(0, 2, 1, 3).contiguous()
286+ w = w.permute(0, 2, 1, 3).contiguous()
287+ u = u.permute(0, 2, 1, 3).contiguous()
288+ g = g.permute(0, 2, 1).contiguous()
289+ chunk_gated_delta_rule_fwd_kernel_h_blockdim64[(triton.cdiv(V, BV), N * H)](
290+ k=k,
291+ v=u,
292+ w=w,
293+ v_new=v_new,
294+ g=g,
295+ gk=gk,
296+ h=h,
297+ h0=initial_state,
298+ ht=final_state,
299+ cu_seqlens=cu_seqlens,
300+ chunk_offsets=chunk_offsets,
301+ T=T,
302+ H=H,
303+ K=K,
304+ V=V,
305+ BT=BT,
306+ BV=BV,
307+ NT=NT,
308+ )
309+ h = h.permute(0, 2, 1, 3, 4).contiguous()
310+ v_new = v_new.permute(0, 2, 1, 3).contiguous()
311+ return h, v_new, final_state
312+ 
313+ 
314+@triton.heuristics({
315+ 'USE_G': lambda args: args['g'] is not None,
316+ 'USE_GK': lambda args: args['gk'] is not None,
317+ 'USE_INITIAL_STATE': lambda args: args['dh0'] is not None,
318+ 'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
319+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
320+})
321+@triton.autotune(
322+ configs=get_autotune_config(multibuffer_list=(True, False)),
323+ key=['H', 'K', 'V', 'BT', 'BV', 'USE_G', 'IS_VARLEN'],
324+)
325+@triton.jit(do_not_specialize=['T'])
326+def chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64(
327+ q,
328+ k,
329+ w,
330+ g,
331+ gk,
332+ dht,
333+ dh0,
334+ do,
335+ dh,
336+ dv,
337+ dv2,
338+ cu_seqlens,
339+ chunk_offsets,
340+ scale,
341+ T,
342+ H: tl.constexpr,
343+ K: tl.constexpr,
344+ V: tl.constexpr,
345+ BT: tl.constexpr,
346+ BV: tl.constexpr,
347+ USE_G: tl.constexpr,
348+ USE_GK: tl.constexpr,
349+ USE_INITIAL_STATE: tl.constexpr,
350+ USE_FINAL_STATE_GRADIENT: tl.constexpr,
351+ IS_VARLEN: tl.constexpr,
352+):
353+ T_all = T
354+ i_v, i_nh = tl.program_id(0), tl.program_id(1)
355+ i_n, i_h = i_nh // H, i_nh % H
356+ if IS_VARLEN:
357+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
358+ T = eos - bos
359+ NT = tl.cdiv(T, BT)
360+ boh = tl.load(chunk_offsets + i_n).to(tl.int32)
361+ else:
362+ bos, eos = i_n * T, i_n * T + T
363+ NT = tl.cdiv(T, BT)
364+ boh = i_n * NT
365+ 
366+ b_dh1 = tl.zeros([64, BV], dtype=tl.float32)
367+ if K > 64:
368+ b_dh2 = tl.zeros([64, BV], dtype=tl.float32)
369+ if K > 128:
370+ b_dh3 = tl.zeros([64, BV], dtype=tl.float32)
371+ if K > 192:
372+ b_dh4 = tl.zeros([64, BV], dtype=tl.float32)
373+ 
374+ q += (bos * H + i_h) * K
375+ k += (bos * H + i_h) * K
376+ w += (bos * H + i_h) * K
377+ do += (bos * H + i_h) * V
378+ dv += (bos * H + i_h) * V
379+ dv2 += (bos * H + i_h) * V
380+ dh += (boh * H + i_h) * K * V
381+ if USE_GK:
382+ gk += (bos * H + i_h) * K
383+ 
384+ if USE_INITIAL_STATE:
385+ dh0 += i_nh * K * V
386+ if USE_FINAL_STATE_GRADIENT:
387+ dht += i_nh * K * V
388+ 
389+ stride_v = H * V
390+ stride_h = H * K * V
391+ stride_k = H * K
392+ 
393+ if USE_FINAL_STATE_GRADIENT:
394+ p_dht1 = tl.make_block_ptr(dht, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
395+ b_dh1 += tl.load(p_dht1, boundary_check=(0, 1))
396+ if K > 64:
397+ p_dht2 = tl.make_block_ptr(dht, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
398+ b_dh2 += tl.load(p_dht2, boundary_check=(0, 1))
399+ if K > 128:
400+ p_dht3 = tl.make_block_ptr(dht, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
401+ b_dh3 += tl.load(p_dht3, boundary_check=(0, 1))
402+ if K > 192:
403+ p_dht4 = tl.make_block_ptr(dht, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
404+ b_dh4 += tl.load(p_dht4, boundary_check=(0, 1))
405+ 
406+ for i_t in range(NT - 1, -1, -1):
407+ p_dh1 = tl.make_block_ptr(dh + i_t * stride_h, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
408+ tl.store(p_dh1, b_dh1.to(p_dh1.dtype.element_ty), boundary_check=(0, 1))
409+ if K > 64:
410+ p_dh2 = tl.make_block_ptr(dh + i_t * stride_h, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
411+ tl.store(p_dh2, b_dh2.to(p_dh2.dtype.element_ty), boundary_check=(0, 1))
412+ if K > 128:
413+ p_dh3 = tl.make_block_ptr(dh + i_t * stride_h, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
414+ tl.store(p_dh3, b_dh3.to(p_dh3.dtype.element_ty), boundary_check=(0, 1))
415+ if K > 192:
416+ p_dh4 = tl.make_block_ptr(dh + i_t * stride_h, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
417+ tl.store(p_dh4, b_dh4.to(p_dh4.dtype.element_ty), boundary_check=(0, 1))
418+ 
419+ last_idx = min((i_t + 1) * BT, T) - 1
420+ if USE_G:
421+ if IS_VARLEN:
422+ bos_g = i_h * T_all + bos
423+ else:
424+ bos_g = (i_n * H + i_h) * T_all
425+ bg_last = tl.load(g + bos_g + last_idx)
426+ bg_last_exp = tl.exp(bg_last)
427+ p_g = tl.make_block_ptr(base=g + bos_g, shape=(T,), strides=(1,), offsets=(i_t * BT,), block_shape=(BT,), order=(0,))
428+ b_g = tl.load(p_g, boundary_check=(0,))
429+ b_g_exp = tl.exp(b_g)
430+ 
431+ p_dv = tl.make_block_ptr(dv, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
432+ p_dv2 = tl.make_block_ptr(dv2, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
433+ p_do = tl.make_block_ptr(do, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
434+ 
435+ b_do = tl.load(p_do, boundary_check=(0, 1))
436+ 
437+ # Update dv
438+ p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 0), (BT, 64), (1, 0))
439+ b_k = tl.load(p_k, boundary_check=(0, 1))
440+ if USE_GK:
441+ o_k1 = tl.arange(0, 64)
442+ b_gk_last1 = tl.load(gk + last_idx * H * K + o_k1, mask=(o_k1 < K), other=0.)
443+ b_dv = tl.dot(b_k, b_dh1.to(b_k.dtype))
444+ 
445+ if K > 64:
446+ p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 64), (BT, 64), (1, 0))
447+ b_k = tl.load(p_k, boundary_check=(0, 1))
448+ if USE_GK:
449+ o_k2 = 64 + o_k1
450+ b_gk_last2 = tl.load(gk + last_idx * H * K + o_k2, mask=(o_k2 < K), other=0.)
451+ b_dv += tl.dot(b_k, b_dh2.to(b_k.dtype))
452+ 
453+ if K > 128:
454+ p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 128), (BT, 64), (1, 0))
455+ b_k = tl.load(p_k, boundary_check=(0, 1))
456+ if USE_GK:
457+ o_k3 = 128 + o_k1
458+ b_gk_last3 = tl.load(gk + last_idx * H * K + o_k3, mask=(o_k3 < K), other=0.)
459+ b_dv += tl.dot(b_k, b_dh3.to(b_k.dtype))
460+ 
461+ if K > 192:
462+ p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 192), (BT, 64), (1, 0))
463+ b_k = tl.load(p_k, boundary_check=(0, 1))
464+ if USE_GK:
465+ o_k4 = 192 + o_k1
466+ b_gk_last4 = tl.load(gk + last_idx * H * K + o_k4, mask=(o_k4 < K), other=0.)
467+ b_dv += tl.dot(b_k, b_dh4.to(b_k.dtype))
468+ 
469+ if USE_G:
470+ m_t = (i_t * BT + tl.arange(0, BT)).to(tl.float32) < T
471+ b_dv *= (m_t * tl.exp(bg_last - b_g))[:, None]
472+ b_dv += tl.load(p_dv, boundary_check=(0, 1))
473+ 
474+ tl.store(p_dv2, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
475+ # Update dh
476+ p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (0, i_t * BT), (64, BT), (0, 1))
477+ p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (0, i_t * BT), (64, BT), (0, 1))
478+ b_w = tl.load(p_w, boundary_check=(0, 1))
479+ b_q = tl.load(p_q, boundary_check=(0, 1))
480+ if USE_G:
481+ b_dh1 *= bg_last_exp
482+ b_q = b_q * b_g_exp[None, :]
483+ if USE_GK:
484+ b_dh1 *= tl.exp(b_gk_last1[:, None])
485+ b_dh1 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
486+ if K > 64:
487+ p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (64, i_t * BT), (64, BT), (0, 1))
488+ p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (64, i_t * BT), (64, BT), (0, 1))
489+ b_q = tl.load(p_q, boundary_check=(0, 1))
490+ b_w = tl.load(p_w, boundary_check=(0, 1))
491+ if USE_G:
492+ b_dh2 *= bg_last_exp
493+ b_q = b_q * b_g_exp[None, :]
494+ if USE_GK:
495+ b_dh2 *= tl.exp(b_gk_last2[:, None])
496+ b_dh2 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
497+ if K > 128:
498+ p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (128, i_t * BT), (64, BT), (0, 1))
499+ p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (128, i_t * BT), (64, BT), (0, 1))
500+ b_q = tl.load(p_q, boundary_check=(0, 1))
501+ b_w = tl.load(p_w, boundary_check=(0, 1))
502+ if USE_G:
503+ b_dh3 *= bg_last_exp
504+ b_q = b_q * b_g_exp[None, :]
505+ if USE_GK:
506+ b_dh3 *= tl.exp(b_gk_last3[:, None])
507+ b_dh3 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
508+ if K > 192:
509+ p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (192, i_t * BT), (64, BT), (0, 1))
510+ p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (192, i_t * BT), (64, BT), (0, 1))
511+ b_q = tl.load(p_q, boundary_check=(0, 1))
512+ b_w = tl.load(p_w, boundary_check=(0, 1))
513+ if USE_G:
514+ b_dh4 *= bg_last_exp
515+ b_q = b_q * b_g_exp[None, :]
516+ if USE_GK:
517+ b_dh4 *= tl.exp(b_gk_last4[:, None])
518+ b_dh4 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
519+ 
520+ if USE_INITIAL_STATE:
521+ p_dh0 = tl.make_block_ptr(dh0, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
522+ tl.store(p_dh0, b_dh1.to(p_dh0.dtype.element_ty), boundary_check=(0, 1))
523+ if K > 64:
524+ p_dh1 = tl.make_block_ptr(dh0, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
525+ tl.store(p_dh1, b_dh2.to(p_dh1.dtype.element_ty), boundary_check=(0, 1))
526+ if K > 128:
527+ p_dh2 = tl.make_block_ptr(dh0, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
528+ tl.store(p_dh2, b_dh3.to(p_dh2.dtype.element_ty), boundary_check=(0, 1))
529+ if K > 192:
530+ p_dh3 = tl.make_block_ptr(dh0, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
531+ tl.store(p_dh3, b_dh4.to(p_dh3.dtype.element_ty), boundary_check=(0, 1))
532+ 
533+ 
534+def chunk_gated_delta_rule_bwd_dhu(
535+ q: torch.Tensor,
536+ k: torch.Tensor,
537+ w: torch.Tensor,
538+ do: torch.Tensor,
539+ dv: torch.Tensor,
540+ g: torch.Tensor | None = None,
541+ gk: torch.Tensor | None = None,
542+ h0: torch.Tensor | None = None,
543+ dht: torch.Tensor | None = None,
544+ scale: float | None = None,
545+ cu_seqlens: torch.LongTensor | None = None,
546+ chunk_size: int = 64, # SY: remove this argument and force chunk size 64?
547+ chunk_indices: torch.LongTensor | None = None,
548+ use_exp2: bool = False,
549+) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
550+ B, T, H, K, V = *q.shape, do.shape[-1]
551+ # N: the actual number of sequences in the batch with either equal or variable lengths
552+ BT = 64
553+ assert K <= 256, "current kernel does not support head dimension being larger than 256."
554+ 
555+ if chunk_indices is None and cu_seqlens is not None:
556+ chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size)
557+ if cu_seqlens is None:
558+ N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
559+ else:
560+ N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
561+ 
562+ dh = q.new_empty(B, NT, H, K, V)
563+ dh0 = torch.empty_like(h0, dtype=torch.float32) if h0 is not None else None
564+ dv2 = torch.empty_like(dv)
565+ 
566+ BV = 128
567+ 
568+ g = g.permute(0, 2, 1).contiguous()
569+ 
570+ chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64[(triton.cdiv(V, BV), N * H)](
571+ q=q,
572+ k=k,
573+ w=w,
574+ g=g,
575+ gk=gk,
576+ dht=dht,
577+ dh0=dh0,
578+ do=do,
579+ dh=dh,
580+ dv=dv,
581+ dv2=dv2,
582+ cu_seqlens=cu_seqlens,
583+ chunk_offsets=chunk_offsets,
584+ scale=scale,
585+ T=T,
586+ H=H,
587+ K=K,
588+ V=V,
589+ BT=BT,
590+ BV=BV,
591+ )
592+ return dh, dh0, dv2
@@ -0,0 +1,607 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+# -*- coding: utf-8 -*-
16+# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
17+ 
18+# pylint: disable=missing-public-type-hints,missing-public-docstring,disallowed-name
19+# pylint: disable=invalid-name,missing-module-docstring,missing-function-docstring
20+# pylint: disable=unused-variable,too-many-nested-blocks
21+ 
22+from typing import Optional, Tuple
23+ 
24+import torch
25+import triton
26+import triton.language as tl
27+ 
28+from .utils import prepare_chunk_indices, exp, prepare_chunk_offsets
29+ 
30+ 
31+@triton.heuristics({
32+ 'USE_G': lambda args: args['g'] is not None,
33+ 'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
34+ 'USE_DW': lambda args: args['dw'] is not None,
35+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
36+})
37+@triton.jit(do_not_specialize=['T'])
38+def chunk_bwd_kernel_dqkwg(
39+ q,
40+ k,
41+ v,
42+ h,
43+ g,
44+ g_gamma,
45+ do,
46+ dh,
47+ dq,
48+ dk,
49+ dg,
50+ w,
51+ dv,
52+ dw,
53+ cu_seqlens,
54+ chunk_indices,
55+ scale,
56+ B: tl.constexpr,
57+ T,
58+ H: tl.constexpr,
59+ K: tl.constexpr,
60+ V: tl.constexpr,
61+ BT: tl.constexpr,
62+ BK: tl.constexpr,
63+ BV: tl.constexpr,
64+ USE_G: tl.constexpr,
65+ USE_G_GAMMA: tl.constexpr,
66+ USE_DW: tl.constexpr,
67+ IS_VARLEN: tl.constexpr,
68+ gdiff,
69+):
70+ i_t, i_b = tl.program_id(0), tl.program_id(1)
71+ T_max = T
72+ if IS_VARLEN:
73+ i_tg = i_t
74+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
75+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
76+ total = B * T_max
77+ T = eos - bos
78+ else:
79+ NT = tl.cdiv(T, BT)
80+ i_tg = i_b * NT + i_t
81+ bos, eos = i_b * T, i_b * T + T
82+ total = B * T_max
83+ 
84+ NK = tl.cdiv(K, BK)
85+ for i_k in range(NK):
86+ if USE_G:
87+ dg_k = dg + i_k * total * H
88+ 
89+ for i_h in range(H):
90+ v_h = v + (bos * H + i_h) * V
91+ do_h = do + (bos * H + i_h) * V
92+ h_h = h + (i_tg * H + i_h).to(tl.int64) * K * V
93+ dh_h = dh + (i_tg * H + i_h).to(tl.int64) * K * V
94+ q_h = q + (bos * H + i_h) * K
95+ k_h = k + (bos * H + i_h) * K
96+ dq_h = dq + (bos * H + i_h) * K
97+ dk_h = dk + (bos * H + i_h) * K
98+ 
99+ if USE_DW:
100+ w_h = w + (bos * H + i_h) * K
101+ dw_h = dw + (bos * H + i_h) * K
102+ dv_h = dv + (bos * H + i_h) * V
103+ 
104+ if USE_G:
105+ if IS_VARLEN:
106+ dg_h = dg_k + i_h * T_max + bos
107+ g_h = g + i_h * T_max + bos
108+ else:
109+ dg_h = dg_k + (i_b * H + i_h) * T_max
110+ g_h = g + (i_b * H + i_h) * T_max
111+ b_dg_last = tl.zeros([1, ], dtype=tl.float32)
112+ 
113+ if USE_G_GAMMA:
114+ b_gamma = tl.load(g_gamma + i_h)
115+ b_g = b_gamma * (tl.arange(0, BT) + 1)
116+ b_g_last = b_gamma * min(BT, T - i_t * BT)
117+ 
118+ b_dq = tl.zeros([BT, BK], dtype=tl.float32)
119+ b_dk = tl.zeros([BT, BK], dtype=tl.float32)
120+ b_ds = tl.zeros([BT, BT], dtype=tl.float32)
121+ b_dw = tl.zeros([BT, BK], dtype=tl.float32) if USE_DW else None
122+ 
123+ for i_v in range(tl.cdiv(V, BV)):
124+ p_v = tl.make_block_ptr(v_h, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
125+ p_do = tl.make_block_ptr(do_h, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
126+ p_h = tl.make_block_ptr(h_h, (V, K), (1, V), (i_v * BV, i_k * BK), (BV, BK), (0, 1))
127+ p_dh = tl.make_block_ptr(dh_h, (V, K), (1, V), (i_v * BV, i_k * BK), (BV, BK), (0, 1))
128+ 
129+ b_v = tl.load(p_v, boundary_check=(0, 1))
130+ b_do = tl.load(p_do, boundary_check=(0, 1))
131+ b_h = tl.load(p_h, boundary_check=(0, 1))
132+ b_dh = tl.load(p_dh, boundary_check=(0, 1))
133+ 
134+ if USE_G:
135+ b_dg_last += (tl.sum(b_h * b_dh))
136+ 
137+ b_ds += tl.dot(b_do, tl.trans(b_v))
138+ b_dq += tl.dot(b_do, b_h.to(b_do.dtype))
139+ b_dk += tl.dot(b_v, b_dh.to(b_v.dtype))
140+ 
141+ if USE_DW:
142+ p_dv = tl.make_block_ptr(dv_h, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
143+ b_dv = tl.load(p_dv, boundary_check=(0, 1))
144+ b_dw += tl.dot(b_dv.to(b_v.dtype), b_h.to(b_v.dtype))
145+ 
146+ if USE_DW:
147+ p_dw = tl.make_block_ptr(dw_h, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
148+ tl.store(p_dw, -b_dw.to(p_dw.dtype.element_ty), boundary_check=(0, 1))
149+ 
150+ tl.debug_barrier()
151+ 
152+ p_q = tl.make_block_ptr(q_h, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
153+ p_k = tl.make_block_ptr(k_h, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
154+ b_q = tl.load(p_q, boundary_check=(0, 1))
155+ b_k = tl.load(p_k, boundary_check=(0, 1))
156+ 
157+ p_dq = tl.make_block_ptr(dq_h, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
158+ p_dk = tl.make_block_ptr(dk_h, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
159+ 
160+ o_t = i_t * BT + tl.arange(0, BT)
161+ m_t = o_t < T
162+ m_A = (o_t[:, None] >= o_t[None, :]) & (m_t[:, None] & m_t)
163+ 
164+ if USE_G:
165+ b_dg = tl.zeros([BT, ], dtype=tl.float32)
166+ p_g = tl.make_block_ptr(g_h, (T,), (1,), (i_t * BT,), (BT,), (0,))
167+ b_g = tl.load(p_g, boundary_check=(0,))
168+ b_g_last = tl.load(g_h + (min(i_t * BT + BT, T) - 1) * 1)
169+ b_dg_last *= tl.exp(b_g_last)
170+ 
171+ b_dq = b_dq * tl.exp(b_g)[:, None] * scale
172+ b_dg += tl.sum(b_dq * b_q, axis=1)
173+ 
174+ b_dk = b_dk * tl.where(m_t, tl.exp(-b_g + b_g_last), 0)[:, None]
175+ b_dg -= tl.sum(b_k * b_dk, axis=1)
176+ b_dg_last += tl.sum(b_dk * b_k)
177+ 
178+ if IS_VARLEN:
179+ b_ds = tl.where(m_A, b_ds * exp(b_g[:, None] - b_g[None, :]), 0) * scale
180+ else:
181+ p_gdiff = tl.make_block_ptr(gdiff + i_b * H * NT * BT * BT + i_h * NT * BT * BT + i_t * BT * BT,
182+ (BT, BT), (BT, 1), (0, 0), (BT, BT), (1, 0))
183+ gdiff_ = tl.load(p_gdiff)
184+ b_ds = b_ds * gdiff_ * scale
185+ 
186+ b_ds2 = b_ds * tl.dot(b_q, tl.trans(b_k))
187+ b_dg += tl.sum(b_ds2, axis=1)
188+ b_dg -= tl.sum(b_ds2, axis=0)
189+ 
190+ b_ds = b_ds.to(b_k.dtype)
191+ b_dq += tl.dot(b_ds, b_k)
192+ b_dk += tl.dot(tl.trans(b_ds), b_q)
193+ p_dg = tl.make_block_ptr(dg_h, (T,), (1,), (i_t * BT,), (BT,), (0,))
194+ 
195+ last_index_local = min(BT, T - i_t * BT) - 1
196+ if last_index_local >= 0:
197+ is_last_mask = tl.arange(0, BT) == last_index_local
198+ b_dg = tl.where(is_last_mask, b_dg + b_dg_last, b_dg)
199+ else:
200+ pass
201+ 
202+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
203+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
204+ tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), boundary_check=(0,))
205+ 
206+ elif USE_G_GAMMA:
207+ b_dq = b_dq * exp(b_g)[:, None] * scale
208+ b_dk = b_dk * tl.where(m_t, exp(-b_g + b_g_last), 0)[:, None]
209+ b_ds = tl.where(m_A, b_ds * exp(b_g[:, None] - b_g[None, :]), 0) * scale
210+ b_ds = b_ds.to(b_k.dtype)
211+ b_dq += tl.dot(b_ds, b_k)
212+ b_dk += tl.dot(tl.trans(b_ds), b_q)
213+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
214+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
215+ 
216+ else:
217+ b_ds = tl.where(m_A, b_ds, 0)
218+ b_ds = b_ds.to(b_k.dtype)
219+ b_dq += tl.dot(b_ds, b_k)
220+ b_dk += tl.dot(tl.trans(b_ds), b_q) * scale
221+ b_dq *= scale
222+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
223+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
224+ 
225+ 
226+@triton.heuristics({
227+ 'USE_G': lambda args: args['g'] is not None,
228+ 'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
229+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
230+})
231+@triton.jit(do_not_specialize=['T'])
232+def chunk_bwd_kernel_dv_local(
233+ q,
234+ k,
235+ g,
236+ g_gamma,
237+ do,
238+ dv,
239+ cu_seqlens,
240+ chunk_indices,
241+ scale,
242+ T,
243+ H: tl.constexpr,
244+ K: tl.constexpr,
245+ V: tl.constexpr,
246+ BT: tl.constexpr,
247+ BK: tl.constexpr,
248+ BV: tl.constexpr,
249+ USE_G: tl.constexpr,
250+ USE_G_GAMMA: tl.constexpr,
251+ IS_VARLEN: tl.constexpr,
252+):
253+ i_t, i_b = tl.program_id(0), tl.program_id(1)
254+ T_max = T
255+ 
256+ if IS_VARLEN:
257+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
258+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
259+ T = eos - bos
260+ else:
261+ bos, eos = i_b * T, i_b * T + T
262+ 
263+ for i_h in range(H):
264+ offset_kh = (bos * H + i_h) * K
265+ offset_vh = (bos * H + i_h) * V
266+ 
267+ b_A = tl.zeros([BT, BT], dtype=tl.float32)
268+ for i_k in range(tl.cdiv(K, BK)):
269+ p_k = tl.make_block_ptr(k + offset_kh, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
270+ p_q = tl.make_block_ptr(q + offset_kh, (K, T), (1, H * K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
271+ b_q = tl.load(p_q, boundary_check=(0, 1))
272+ b_k = tl.load(p_k, boundary_check=(0, 1))
273+ b_A += tl.dot(b_k, b_q)
274+ 
275+ if USE_G:
276+ if IS_VARLEN:
277+ offset_g = i_h * T_max + bos
278+ else:
279+ offset_g = i_b * H * T_max + i_h * T_max
280+ 
281+ p_g = tl.make_block_ptr(g + offset_g, (T,), (1,), (i_t * BT,), (BT,), (0,))
282+ b_g = tl.load(p_g, boundary_check=(0,))
283+ 
284+ if USE_G_GAMMA:
285+ b_gamma = tl.load(g_gamma + i_h)
286+ b_g = b_gamma * (tl.arange(0, BT) + 1)
287+ 
288+ o_t = i_t * BT + tl.arange(0, BT)
289+ m_t = o_t < T
290+ m_A = (o_t[:, None] <= o_t[None, :]) & (m_t[:, None] & m_t)
291+ 
292+ if USE_G:
293+ b_A = tl.where(m_A, b_A * tl.exp(b_g[None, :] - b_g[:, None]) * scale, 0).to(do.dtype.element_ty)
294+ else:
295+ b_A = tl.where(m_A, b_A * scale, 0).to(do.dtype.element_ty)
296+ 
297+ for i_v in range(tl.cdiv(V, BV)):
298+ p_do = tl.make_block_ptr(do + offset_vh, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
299+ p_dv = tl.make_block_ptr(dv + offset_vh, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
300+ b_do = tl.load(p_do, boundary_check=(0, 1))
301+ b_dv = tl.dot(b_A.to(b_do.dtype), b_do)
302+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
303+ 
304+ 
305+@triton.heuristics({
306+ 'USE_G': lambda args: args['g'] is not None,
307+ 'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
308+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None
309+})
310+@triton.jit(do_not_specialize=['T'])
311+def chunk_fwd_kernel_o(
312+ q,
313+ k,
314+ v,
315+ h,
316+ g,
317+ g_gamma,
318+ o,
319+ cu_seqlens,
320+ chunk_offsets,
321+ scale,
322+ T,
323+ H: tl.constexpr,
324+ N: tl.constexpr,
325+ Hg: tl.constexpr,
326+ K: tl.constexpr,
327+ V: tl.constexpr,
328+ BT: tl.constexpr,
329+ BK: tl.constexpr,
330+ BV: tl.constexpr,
331+ USE_G: tl.constexpr,
332+ USE_G_GAMMA: tl.constexpr,
333+ IS_VARLEN: tl.constexpr,
334+):
335+ T_max = T
336+ for i_v in range(tl.cdiv(V, BV)):
337+ for i_n in range(N):
338+ if IS_VARLEN:
339+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(
340+ cu_seqlens + i_n + 1
341+ ).to(tl.int32)
342+ T = eos - bos
343+ NT = tl.cdiv(T, BT)
344+ boh = tl.load(chunk_offsets + i_n).to(tl.int64)
345+ else:
346+ bos, eos = i_n * T, i_n * T + T
347+ NT = tl.cdiv(T, BT)
348+ boh = i_n * NT
349+ 
350+ core_id = tl.program_id(0)
351+ total_cores = tl.num_programs(0)
352+ base_chunks_per_pid = NT // total_cores
353+ remainder = NT % total_cores
354+ 
355+ if core_id < remainder:
356+ chunks_this_pid = base_chunks_per_pid + 1
357+ start_idx = core_id * chunks_this_pid
358+ else:
359+ chunks_this_pid = base_chunks_per_pid
360+ start_idx = core_id * base_chunks_per_pid + remainder
361+ 
362+ # offset calculation
363+ for i_h in range(0, H):
364+ q_offset = (bos * Hg + i_h // (H // Hg)) * K
365+ k_offset = (bos * Hg + i_h // (H // Hg)) * K
366+ v_offset = (bos * H + i_h) * V
367+ o_offset = (bos * H + i_h) * V
368+ 
369+ for i_t in range(start_idx, start_idx + chunks_this_pid):
370+ i_tg = boh + i_t
371+ h_base = h + (i_tg * H + i_h).to(tl.int64) * K * V
372+ b_o = tl.zeros([BT, BV], dtype=tl.float32)
373+ b_A = tl.zeros([BT, BT], dtype=tl.float32)
374+ for i_k in range(tl.cdiv(K, BK)):
375+ p_q = tl.make_block_ptr(
376+ q + q_offset, (T, K), (Hg * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0)
377+ )
378+ p_k = tl.make_block_ptr(
379+ k + k_offset, (K, T), (1, Hg * K), (i_k * BK, i_t * BT), (BK, BT), (0, 1)
380+ )
381+ p_h = tl.make_block_ptr(
382+ h_base, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0)
383+ )
384+ b_q = tl.load(p_q, boundary_check=(0, 1))
385+ b_k = tl.load(p_k, boundary_check=(0, 1))
386+ b_h = tl.load(p_h, boundary_check=(0, 1))
387+ 
388+ # [BT, BK] @ [BK, BV] -> [BT, BV]
389+ b_o += tl.dot(b_q, b_h)
390+ # [BT, BK] @ [BK, BT] -> [BT, BT]
391+ b_A += tl.dot(b_q, b_k)
392+ 
393+ if USE_G:
394+ if IS_VARLEN:
395+ p_g = tl.make_block_ptr(g + bos + i_h * T_max, (T,), (1,), (i_t * BT,), (BT,), (0,))
396+ else:
397+ p_g = tl.make_block_ptr(g + bos * H + i_h * T_max, (T,), (1,), (i_t * BT,), (BT,), (0,))
398+ b_g = tl.load(p_g, boundary_check=(0,))
399+ b_o = b_o * exp(b_g)[:, None]
400+ b_A = b_A * exp(b_g[:, None] - b_g[None, :])
401+ if USE_G_GAMMA:
402+ b_gamma = tl.load(g_gamma + i_h)
403+ b_g = b_gamma * (tl.arange(0, BT) + 1)
404+ 
405+ o_i = tl.arange(0, BT)
406+ m_A = o_i[:, None] >= o_i[None, :]
407+ b_A = tl.where(m_A, b_A, 0)
408+ 
409+ p_v = tl.make_block_ptr(
410+ v + v_offset, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0)
411+ )
412+ p_o = tl.make_block_ptr(
413+ o + o_offset, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0)
414+ )
415+ b_v = tl.load(p_v, boundary_check=(0, 1))
416+ 
417+ # to fix mma -> mma layout conversion
418+ # already solved by triton v3.2 or higher
419+ b_o = b_o * scale + tl.dot(b_A.to(b_v.dtype), b_v) * scale
420+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))
421+ 
422+ 
423+def chunk_bwd_dqkwg(
424+ q: torch.Tensor,
425+ k: torch.Tensor,
426+ v: torch.Tensor,
427+ do: torch.Tensor,
428+ h: torch.Tensor,
429+ dh: torch.Tensor,
430+ g: Optional[torch.Tensor] = None,
431+ g_gamma: Optional[torch.Tensor] = None,
432+ dv: Optional[torch.Tensor] = None,
433+ w: Optional[torch.Tensor] = None,
434+ cu_seqlens: Optional[torch.LongTensor] = None,
435+ chunk_size: int = 64,
436+ scale: float = 1.0,
437+) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
438+ B, T, H, K, V = *k.shape, v.shape[-1]
439+ BT = min(chunk_size, max(16, triton.next_power_of_2(T)))
440+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
441+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
442+ 
443+ BK = 128 if cu_seqlens is None else 64
444+ BV = 64
445+ NK = triton.cdiv(K, BK)
446+ dq = torch.empty_like(q)
447+ dk = torch.empty_like(k)
448+ g = g.transpose(1, 2).contiguous()
449+ dg = torch.empty(NK, *g.shape, dtype=torch.float32, device=g.device) if g is not None else None
450+ dw = torch.empty_like(w) if w is not None else None
451+ grid = (NT, B)
452+ 
453+ if cu_seqlens is None:
454+ if NT * BT == T:
455+ g_ = g.reshape(B, H, NT, BT)
456+ g_diff = g_[:, :, :, :, None] - g_[:, :, :, None, :]
457+ g_diff = g_diff.clamp(-60, 60).exp()
458+ g_diff[:, :, :] *= torch.tril(torch.ones(BT, BT), diagonal=0).to(g.device)
459+ else:
460+ diff = NT * BT - T
461+ g_ = torch.cat((g, torch.zeros(B, H, diff).to(g.device)), dim=-1).reshape(B, H, NT, BT)
462+ g_diff = g_[:, :, :, :, None] - g_[:, :, :, None, :]
463+ g_diff = g_diff.clamp(-60, 60).exp()
464+ g_diff[:, :, :] *= torch.tril(torch.ones(BT, BT), diagonal=0).to(g.device)
465+ bias = torch.arange(0, BT).to(g.device)
466+ o_t = (NT - 1) * BT + bias
467+ m_t = o_t < T
468+ m_A = (m_t[:, None] & m_t)
469+ g_diff[:, :, -1] *= m_A
470+ else:
471+ g_diff = None
472+ 
473+ chunk_bwd_kernel_dqkwg[grid](
474+ q=q,
475+ k=k,
476+ v=v,
477+ h=h,
478+ g=g,
479+ g_gamma=g_gamma,
480+ do=do,
481+ dh=dh,
482+ dv=dv,
483+ w=w,
484+ dw=dw,
485+ dq=dq,
486+ dk=dk,
487+ dg=dg,
488+ cu_seqlens=cu_seqlens,
489+ chunk_indices=chunk_indices,
490+ scale=scale,
491+ B=B,
492+ T=T,
493+ H=H,
494+ K=K,
495+ V=V,
496+ BT=BT,
497+ BK=BK,
498+ BV=BV,
499+ gdiff=g_diff,
500+ )
501+ 
502+ if dg is not None:
503+ dg = dg.sum(0)
504+ dg = dg.transpose(1, 2).contiguous()
505+ return dq, dk, dw, dg
506+ 
507+ 
508+def chunk_bwd_dv_local(
509+ q: torch.Tensor,
510+ k: torch.Tensor,
511+ do: torch.Tensor,
512+ g: Optional[torch.Tensor] = None,
513+ g_gamma: Optional[torch.Tensor] = None,
514+ scale: float = None,
515+ cu_seqlens: Optional[torch.LongTensor] = None,
516+ chunk_size: int = 64
517+) -> torch.Tensor:
518+ B, T, H, K, V = *k.shape, do.shape[-1]
519+ BT = min(chunk_size, max(16, triton.next_power_of_2(T)))
520+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
521+ 
522+ BK = 128
523+ BV = 128
524+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
525+ 
526+ g = g.transpose(1, 2).contiguous()
527+ dv = torch.empty_like(do)
528+ grid = (NT, B)
529+ chunk_bwd_kernel_dv_local[grid](
530+ q=q,
531+ k=k,
532+ g=g,
533+ g_gamma=g_gamma,
534+ do=do,
535+ dv=dv,
536+ cu_seqlens=cu_seqlens,
537+ chunk_indices=chunk_indices,
538+ scale=scale,
539+ T=T,
540+ H=H,
541+ K=K,
542+ V=V,
543+ BT=BT,
544+ BK=BK,
545+ BV=BV,
546+ )
547+ return dv
548+ 
549+ 
550+def chunk_fwd_o(
551+ q: torch.Tensor,
552+ k: torch.Tensor,
553+ v: torch.Tensor,
554+ h: torch.Tensor,
555+ g: Optional[torch.Tensor] = None,
556+ g_gamma: Optional[torch.Tensor] = None,
557+ scale: Optional[float] = None,
558+ cu_seqlens: Optional[torch.LongTensor] = None,
559+ chunk_size: int = 64
560+) -> torch.Tensor:
561+ B, T, Hg, K, V = *q.shape, v.shape[-1]
562+ H = v.shape[-2]
563+ BT = min(chunk_size, max(16, triton.next_power_of_2(T)))
564+ chunk_indices = (
565+ prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
566+ )
567+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
568+ if scale is None:
569+ scale = k.shape[-1] ** -0.5
570+ 
571+ o = torch.empty_like(v)
572+ if cu_seqlens is None:
573+ N, chunk_offsets = B, None
574+ else:
575+ N, chunk_offsets = (
576+ len(cu_seqlens) - 1,
577+ prepare_chunk_offsets(cu_seqlens, BT),
578+ )
579+ 
atomgit-bot
atomgit-botatomgit-bot8月5日

🔵 Low Priority

第 578-579 行定义了一个 grid 函数用于计算 kernel launch grid,但实际 kernel 启动时(第 584 行)使用的是硬编码的 (CV_kernel_num,) = (24,),而非调用该 grid 函数。

changed line(第 578-579 行):def grid(meta): return (triton.cdiv(V, meta["BV"]), N * H) → 该函数从未被调用,属于死代码。虽然不影响运行时行为(kernel 内部通过 tl.num_programs(0) 动态分配工作),但死代码会造成维护困惑,且 N 和 H 变量被函数闭包捕获但实际上并未使用。

建议:删除未使用的 grid 函数定义(第 578-579 行)。如果未来需要动态 grid,应该在 kernel 启动处调用它。

likedislike
不准确?
580+ g = g.transpose(1, 2).contiguous()
581+ h = h.contiguous()
582+ CV_kernel_num = 24
583+ chunk_fwd_kernel_o[(CV_kernel_num,)](
584+ q,
585+ k,
586+ v,
587+ h,
588+ g,
589+ g_gamma,
590+ o,
591+ cu_seqlens,
592+ chunk_offsets,
593+ scale,
594+ T=T,
595+ H=H,
596+ N=N,
597+ Hg=Hg,
598+ K=K,
599+ V=V,
600+ BT=BT,
601+ BK=128,
602+ BV=128,
603+ )
604+ return o
605+ 
606+bwd_chunk_dqkwg = chunk_bwd_dqkwg
607+bwd_chunk_dv_local = chunk_bwd_dv_local
@@ -0,0 +1,355 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+# -*- coding: utf-8 -*-
16+# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
17+ 
18+# pylint: disable=line-too-long,missing-public-type-hints,missing-public-docstring
19+# pylint: disable=unused-argument,invalid-name,missing-module-docstring
20+# pylint: disable=missing-function-docstring
21+ 
22+from typing import Optional
23+ 
24+import torch
25+import triton
26+import triton.language as tl
27+ 
28+from .utils import prepare_chunk_indices
29+ 
30+ 
31+@triton.heuristics({
32+ 'USE_G': lambda args: args['g'] is not None,
33+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
34+})
35+@triton.jit(do_not_specialize=['T', 'NT', 'TOTAL_TASKS'])
36+def chunk_scaled_dot_kkt_fwd_kernel(
37+ k,
38+ g,
39+ beta,
40+ A,
41+ cu_seqlens,
42+ chunk_indices,
43+ T,
44+ H: tl.constexpr,
45+ K: tl.constexpr,
46+ BT: tl.constexpr,
47+ BK: tl.constexpr,
48+ IS_VARLEN: tl.constexpr,
49+ USE_G: tl.constexpr,
50+ NT,
51+ B,
52+ TOTAL_TASKS,
53+):
54+ core_id = tl.program_id(0)
55+ num_blocks = tl.num_programs(0)
56+ T_max = T
57+ 
58+ base_tasks_per_block = TOTAL_TASKS // num_blocks
59+ remainder_tasks = TOTAL_TASKS % num_blocks
60+ 
61+ if core_id < remainder_tasks:
62+ tasks_this_core = base_tasks_per_block + 1
63+ start_idx = core_id * tasks_this_core
64+ else:
65+ tasks_this_core = base_tasks_per_block
66+ start_idx = core_id * base_tasks_per_block + remainder_tasks
67+ 
68+ for idx in range(start_idx, start_idx + tasks_this_core):
69+ i_b = idx // NT
70+ local_idx = idx % NT
71+ 
72+ if IS_VARLEN:
73+ i_n = tl.load(chunk_indices + local_idx * 2).to(tl.int32)
74+ i_t = tl.load(chunk_indices + local_idx * 2 + 1).to(tl.int32)
75+ bos = tl.load(cu_seqlens + i_n).to(tl.int32)
76+ eos = tl.load(cu_seqlens + i_n + 1).to(tl.int32)
77+ T_local = eos - bos
78+ else:
79+ bos, eos = 0, T
80+ i_t = local_idx
81+ T_local = T
82+ 
83+ for i_h in range(H):
84+ k_batch_off = i_b * T_max * H * K
85+ beta_batch_off = i_b * H * T_max
86+ g_batch_off = i_b * H * T_max
87+ A_batch_off = i_b * T_max * H * BT
88+ 
89+ p_beta = tl.make_block_ptr(beta + beta_batch_off + bos + i_h * T_max, (T_local,), (1,), (i_t * BT,), (BT,), (0,))
90+ b_beta = tl.load(p_beta, boundary_check=(0,))
91+ 
92+ b_A = tl.zeros([BT, BT], dtype=tl.float32)
93+ for i_k in range(tl.cdiv(K, BK)):
94+ p_k = tl.make_block_ptr(k + k_batch_off + (bos * H + i_h) * K, (T_local, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
95+ b_k = tl.load(p_k, boundary_check=(0, 1))
96+ dot_product = tl.dot(b_k, tl.trans(b_k))
97+ 
98+ o_t = i_t * BT + tl.arange(0, BT)
99+ o_t = o_t.to(tl.float32)
100+ T_mask = (o_t < T_local).to(tl.float32)
101+ 
102+ row_indices = tl.arange(0, BT)[:, None]
103+ col_indices = tl.arange(0, BT)[None, :]
104+ tril_mask = (row_indices > col_indices).to(tl.float32)
105+ tril_mask = tril_mask * T_mask[:, None]
106+ masked_dot = dot_product * tril_mask
107+ b_A += masked_dot
108+ 
109+ if USE_G:
110+ p_g = tl.make_block_ptr(g + g_batch_off + bos + i_h * T_max, (T_local,), (1,), (i_t * BT,), (BT,), (0,))
111+ b_g = tl.load(p_g, boundary_check=(0,))
112+ b_g_diff = b_g[:, None] - b_g[None, :]
113+ b_g_diff = tl.minimum(tl.maximum(b_g_diff, -50.0), 50.0)
114+ b_A *= tl.exp(b_g_diff)
115+ b_A *= b_beta[:, None]
116+ 
117+ p_A = tl.make_block_ptr(A + A_batch_off + (bos * H + i_h) * BT, (T_local, BT), (BT * H, 1), (i_t * BT, 0), (BT, BT), (1, 0))
118+ tl.store(p_A, b_A.to(p_A.dtype.element_ty), boundary_check=(0, 1))
119+ 
120+ 
121+@triton.heuristics({
122+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None
123+})
124+@triton.autotune(
125+ configs=[
126+ triton.Config({'BK': BK})
127+ for BK in [32, 64]
128+ ],
129+ key=["BC"]
130+)
131+@triton.jit(do_not_specialize=['T'])
132+def chunk_scaled_dot_kkt_fwd_kernel_intra_sub_inter(
133+ k,
134+ g,
135+ beta,
136+ A,
137+ cu_seqlens,
138+ chunk_indices,
139+ T,
140+ H: tl.constexpr,
141+ K: tl.constexpr,
142+ BT: tl.constexpr,
143+ BC: tl.constexpr,
144+ BK: tl.constexpr,
145+ NC: tl.constexpr,
146+ IS_VARLEN: tl.constexpr,
147+):
148+ i_t, i_c, i_b = tl.program_id(0), tl.program_id(1), tl.program_id(2)
149+ i_i, i_j = i_c // NC, i_c % NC
150+ 
151+ for i_h in range(H):
152+ if IS_VARLEN:
153+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
154+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
155+ T_val = eos - bos
156+ else:
157+ bos, eos = i_b * T, i_b * T + T
158+ T_val = T
159+ 
160+ should_compute = (i_t * BT + i_i * BC < T_val) and (i_i > i_j)
161+ 
162+ if should_compute:
163+ k_ptr = k + (bos * H + i_h) * K
164+ g_ptr = g + (bos * H + i_h) * K
165+ A_ptr = A + (bos * H + i_h) * BT
166+ 
167+ p_beta = tl.make_block_ptr(beta + bos * H + i_h, (T_val,), (H,), (i_t * BT + i_i * BC,), (BC,), (0,))
168+ b_beta = tl.load(p_beta, boundary_check=(0,))
169+ 
170+ b_A = tl.zeros([BC, BC], dtype=tl.float32)
171+ for i_k in range(tl.cdiv(K, BK)):
172+ p_k = tl.make_block_ptr(k_ptr, (T_val, K), (H * K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK),
173+ (1, 0))
174+ p_g = tl.make_block_ptr(g_ptr, (T_val, K), (H * K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK),
175+ (1, 0))
176+ b_kt = tl.make_block_ptr(k_ptr, (K, T_val), (1, H * K), (i_k * BK, i_t * BT + i_j * BC), (BK, BC),
177+ (0, 1))
178+ p_gk = tl.make_block_ptr(g_ptr, (K, T_val), (1, H * K), (i_k * BK, i_t * BT + i_j * BC), (BK, BC),
179+ (0, 1))
180+ 
181+ o_k = i_k * BK + tl.arange(0, BK)
182+ m_k = o_k < K
183+ b_gn = tl.load(g_ptr + (i_t * BT + i_i * BC) * H * K + o_k, mask=m_k, other=0)
184+ b_g = tl.load(p_g, boundary_check=(0, 1))
185+ b_k = tl.load(p_k, boundary_check=(0, 1)) * tl.exp(b_g - b_gn[None, :])
186+ b_gk = tl.load(p_gk, boundary_check=(0, 1))
187+ b_kt = tl.load(b_kt, boundary_check=(0, 1)) * tl.exp(b_gn[:, None] - b_gk)
188+ b_A += tl.dot(b_k, b_kt)
189+ b_A *= b_beta[:, None]
190+ 
191+ p_A = tl.make_block_ptr(A_ptr, (T_val, BT), (H * BT, 1), (i_t * BT + i_i * BC, i_j * BC), (BC, BC), (1, 0))
192+ tl.store(p_A, b_A.to(A.dtype.element_ty), boundary_check=(0, 1))
193+ 
194+ 
195+@triton.heuristics({
196+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None
197+})
198+@triton.jit(do_not_specialize=['T'])
199+def chunk_scaled_dot_kkt_fwd_kernel_intra_sub_intra(
200+ k,
201+ g,
202+ beta,
203+ A,
204+ cu_seqlens,
205+ chunk_indices,
206+ T,
207+ H: tl.constexpr,
208+ K: tl.constexpr,
209+ BT: tl.constexpr,
210+ BC: tl.constexpr,
211+ BK: tl.constexpr,
212+ IS_VARLEN: tl.constexpr,
213+):
214+ i_t, i_i, i_b = tl.program_id(0), tl.program_id(1), tl.program_id(2)
215+ 
216+ for i_h in range(H):
217+ if IS_VARLEN:
218+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
219+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
220+ T_val = eos - bos
221+ else:
222+ bos, eos = i_b * T, i_b * T + T
223+ T_val = T
224+ 
225+ should_compute = i_t * BT + i_i * BC < T_val
226+ 
227+ if should_compute:
228+ o_i = tl.arange(0, BC)
229+ o_k = tl.arange(0, BK)
230+ m_k = o_k < K
231+ m_A = (i_t * BT + i_i * BC + o_i) < T_val
232+ o_A = (bos + i_t * BT + i_i * BC + o_i) * H * BT + i_h * BT + i_i * BC
233+ 
234+ p_k = tl.make_block_ptr(k + (bos * H + i_h) * K, (T_val, K), (H * K, 1), (i_t * BT + i_i * BC, 0), (BC, BK),
235+ (1, 0))
236+ p_g = tl.make_block_ptr(g + (bos * H + i_h) * K, (T_val, K), (H * K, 1), (i_t * BT + i_i * BC, 0), (BC, BK),
237+ (1, 0))
238+ p_beta = beta + (bos + i_t * BT + i_i * BC + o_i) * H + i_h
239+ 
240+ b_k = tl.load(p_k, boundary_check=(0, 1)) * tl.load(p_beta, mask=m_A, other=0)[:, None]
241+ b_g = tl.load(p_g, boundary_check=(0, 1))
242+ 
243+ p_kt = k + (bos + i_t * BT + i_i * BC) * H * K + i_h * K + o_k
244+ p_gk = g + (bos + i_t * BT + i_i * BC) * H * K + i_h * K + o_k
245+ 
246+ for j in range(0, min(BC, T_val - i_t * BT - i_i * BC)):
247+ b_kt = tl.load(p_kt, mask=m_k, other=0).to(tl.float32)
248+ b_gk = tl.load(p_gk, mask=m_k, other=0).to(tl.float32)
249+ b_A = tl.sum(b_k * b_kt[None, :] * tl.exp(b_g - b_gk[None, :]), 1)
250+ # 转化成f32
251+ o_i_tmp = o_i.to(tl.float32)
252+ b_A = tl.where(o_i_tmp > j, b_A, 0.)
253+ 
254+ tl.store(A + o_A + j, b_A, mask=m_A)
255+ p_kt += H * K
256+ p_gk += H * K
257+ 
258+ 
259+def chunk_scaled_dot_kkt_fwd(
260+ k: torch.Tensor,
261+ g: Optional[torch.Tensor] = None,
262+ gk: Optional[torch.Tensor] = None,
263+ beta: Optional[torch.Tensor] = None,
264+ cu_seqlens: Optional[torch.LongTensor] = None,
265+ chunk_size: int = 64,
266+ output_dtype: torch.dtype = torch.float32
267+) -> torch.Tensor:
268+ r"""
269+ Compute beta * K * K^T.
270+ 
271+ Args:
272+ k (torch.Tensor):
273+ The key tensor of shape `[B, T, H, K]`.
274+ beta (torch.Tensor):
275+ The beta tensor of shape `[B, T, H]`.
276+ g (torch.Tensor):
277+ The cumulative sum of the gate tensor of shape `[B, T, H]`. Default: `None`.
278+ gk (torch.Tensor):
279+ The cumulative sum of the gate tensor of shape `[B, T, H, K]` applied to the key tensor. Default: `None`.
280+ cu_seqlens (torch.LongTensor):
281+ The cumulative sequence lengths of the input tensor.
282+ Default: None
283+ chunk_size (int):
284+ The chunk size. Default: 64.
285+ output_dtype (torch.dtype):
286+ The dtype of the output tensor. Default: `torch.float32`
287+ 
288+ Returns:
289+ beta * K * K^T of shape `[B, T, H, BT]` where `BT` is the chunk size.
290+ """
291+ B, T, H, K = k.shape
292+ BT = chunk_size
293+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
294+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
295+ beta = beta.transpose(1, 2).contiguous()
296+ g = g.transpose(1, 2).contiguous()
297+ BK = 128
298+ kernel_num = 24
299+ 
300+ if gk is None:
301+ A = torch.empty(B, T, H, BT, device=k.device, dtype=output_dtype)
302+ chunk_scaled_dot_kkt_fwd_kernel[(kernel_num,)](
303+ k=k,
304+ g=g,
305+ beta=beta,
306+ A=A,
307+ cu_seqlens=cu_seqlens,
308+ chunk_indices=chunk_indices,
309+ T=T,
310+ H=H,
311+ K=K,
312+ BT=BT,
313+ BK=BK,
314+ NT=NT,
315+ B=B,
316+ TOTAL_TASKS=B * NT,
317+ )
318+ return A
319+ 
320+ BC = min(16, BT)
321+ NC = triton.cdiv(BT, BC)
322+ BK = max(triton.next_power_of_2(K), 16)
323+ A = torch.zeros(B, T, H, BT, device=k.device, dtype=output_dtype)
324+ grid = (NT, NC * NC, B)
325+ chunk_scaled_dot_kkt_fwd_kernel_intra_sub_inter[grid](
326+ k=k,
327+ g=gk,
328+ beta=beta,
329+ A=A,
330+ cu_seqlens=cu_seqlens,
331+ chunk_indices=chunk_indices,
332+ T=T,
333+ H=H,
334+ K=K,
335+ BT=BT,
336+ BC=BC,
337+ NC=NC,
338+ )
339+ 
340+ grid = (NT, NC, B)
341+ chunk_scaled_dot_kkt_fwd_kernel_intra_sub_intra[grid](
342+ k=k,
343+ g=gk,
344+ beta=beta,
345+ A=A,
346+ cu_seqlens=cu_seqlens,
347+ chunk_indices=chunk_indices,
348+ T=T,
349+ H=H,
350+ K=K,
351+ BT=BT,
352+ BC=BC,
353+ BK=BK,
354+ )
355+ return A
@@ -0,0 +1,163 @@
1+# Copyright 2026 Huawei Technologies Co., Ltd
2+#
3+# Licensed under the Apache License, Version 2.0 (the "License");
4+# you may not use this file except in compliance with the License.
5+# You may obtain a copy of the License at
6+#
7+# http://www.apache.org/licenses/LICENSE-2.0
8+#
9+# Unless required by applicable law or agreed to in writing, software
10+# distributed under the License is distributed on an "AS IS" BASIS,
11+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+# See the License for the specific language governing permissions and
13+# limitations under the License.
14+# ============================================================================
15+# -*- coding: utf-8 -*-
16+# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
17+ 
18+# pylint: disable=missing-public-type-hints,missing-public-docstring,disallowed-name
19+# pylint: disable=useless-return,unused-argument,no-else-return,invalid-name
20+# pylint: disable=missing-module-docstring,missing-function-docstring
21+ 
22+from typing import Optional
23+ 
24+import torch
25+import triton
26+import triton.language as tl
27+ 
28+from .utils import prepare_chunk_indices
29+ 
30+ 
31+@triton.heuristics({
32+ 'HAS_SCALE': lambda args: args['scale'] is not None,
33+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None
34+})
35+@triton.jit(do_not_specialize=['T'])
36+def chunk_local_cumsum_scalar_kernel(
37+ s,
38+ o,
39+ scale,
40+ cu_seqlens,
41+ chunk_indices,
42+ T,
43+ B: tl.constexpr,
44+ H: tl.constexpr,
45+ BLOCK_T: tl.constexpr,
46+ REVERSE: tl.constexpr,
47+ HAS_SCALE: tl.constexpr,
48+ IS_VARLEN: tl.constexpr,
49+ HEAD_FIRST: tl.constexpr,
50+ CHUNK_SIZE: tl.constexpr = 64,
51+):
52+ i_block, i_b = tl.program_id(0), tl.program_id(1)
53+ N_CHUNKS: tl.constexpr = BLOCK_T // CHUNK_SIZE
54+ 
55+ if IS_VARLEN:
56+ i_s, i_block = tl.load(chunk_indices + i_block * 2).to(tl.int32), tl.load(
57+ chunk_indices + i_block * 2 + 1
58+ ).to(tl.int32)
59+ 
60+ bos, eos = tl.load(cu_seqlens + i_s).to(tl.int32), tl.load(
61+ cu_seqlens + i_s + 1
62+ ).to(tl.int32)
63+ T = eos - bos
64+ else:
65+ bos, eos = i_b * T, i_b * T + T
66+ 
67+ ptr_s = tl.make_block_ptr(
68+ s + bos * H, (T, H), (H, 1), (i_block * BLOCK_T, 0), (BLOCK_T, H), (1, 0)
69+ )
70+ ptr_o = tl.make_block_ptr(
71+ o + bos * H, (T, H), (H, 1), (i_block * BLOCK_T, 0), (BLOCK_T, H), (1, 0)
72+ )
73+ b_s = tl.load(ptr_s, boundary_check=(0,)).to(tl.float32)
74+ b_s = tl.reshape(b_s, (N_CHUNKS, CHUNK_SIZE, H))
75+ b_s = tl.trans(b_s, (1, 0, 2))
76+ b_o = tl.cumsum(b_s, axis=0)
77+ if REVERSE:
78+ b_z = tl.sum(b_s, axis=0)
79+ b_o = -b_o + b_z[None] + b_s
80+ if HAS_SCALE:
81+ b_o *= scale
82+ b_o = tl.trans(b_o, (1, 0, 2))
83+ b_o = tl.reshape(b_o, (BLOCK_T, H))
84+ 
85+ tl.store(ptr_o, b_o.to(ptr_o.dtype.element_ty), boundary_check=(0,))
86+ return
87+ 
88+ 
89+def chunk_local_cumsum_scalar(
90+ g: torch.Tensor,
91+ chunk_size: int,
92+ reverse: bool = False,
93+ scale: float = None,
94+ cu_seqlens: Optional[torch.Tensor] = None,
95+ head_first: bool = False,
96+ output_dtype: Optional[torch.dtype] = torch.float
97+) -> torch.Tensor:
98+ 
99+ B, T, H = g.shape
100+ if chunk_size != 2 ** (chunk_size.bit_length() - 1):
101+ raise ValueError(
102+ f"chunk_size must be a power of 2, chunk_size is {chunk_size}"
103+ )
104+ # We adjust the tiling strategy to prevent overflow in in backward passes and context parallel scenarios
105+ # while maximizing UB utilization where possible.
106+ # The tiling strategy is as follows:
107+ # 1. BT must be greater than or equal to chunk_size.
108+ # 2. UB estimation varies directly with H.
109+ # 3. BT in reverse mode is smaller than in forward mode.
110+ BT = max(chunk_size, triton.next_power_of_2((1 << 11 if reverse else 1 << 12) // H))
111+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
112+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
113+ g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
114+ grid = (NT, B)
115+ chunk_local_cumsum_scalar_kernel[grid](
116+ s=g_org,
117+ o=g,
118+ scale=scale,
119+ cu_seqlens=cu_seqlens,
120+ chunk_indices=chunk_indices,
121+ T=T,
122+ B=B,
123+ H=H,
124+ BLOCK_T=BT,
125+ HEAD_FIRST=head_first,
126+ REVERSE=reverse,
127+ CHUNK_SIZE=chunk_size,
128+ )
129+ return g
130+ 
131+ 
132+def chunk_local_cumsum(
133+ g: torch.Tensor,
134+ chunk_size: int,
135+ reverse: bool = False,
136+ scale: float = None,
137+ cu_seqlens: Optional[torch.Tensor] = None,
138+ head_first: bool = False,
139+ output_dtype: Optional[torch.dtype] = torch.float,
140+ **kwargs
141+) -> torch.Tensor:
142+ if cu_seqlens is not None:
143+ if g.shape[0] != 1:
144+ raise ValueError(
145+ "Only batch size 1 is supported when cu_seqlens are provided, "
146+ f"current size is {g.shape[0]}"
147+ )
148+ if len(g.shape) == 3:
149+ return chunk_local_cumsum_scalar(
150+ g=g,
151+ chunk_size=chunk_size,
152+ reverse=reverse,
153+ scale=scale,
154+ cu_seqlens=cu_seqlens,
155+ head_first=head_first,
156+ output_dtype=output_dtype
157+ )
158+ else:
159+ raise ValueError(
160+ f"Unsupported input shape {g.shape}, "
161+ f"which should be (B, T, H, D) if `head_first=False` "
162+ f"or (B, H, T, D) otherwise"
163+ )