已合并
feat: add fused GDN backend for linear attention CP #1114
xu-xianliang创建于 8月5日
feat: add fused GDN backend for linear attention CP #1114
已合并
共 20 个文件变更+5004-31
| @@ -3,3 +3,14 @@ hyper-parallel/hyper_parallel/core/hsdp/api.py:hsdp | |||
| 3 | hyper-parallel/hyper_parallel/core/multicore/ops/mega_moe/op_host/mega_moe_def.cpp:ops::MegaMoe::MegaMoe | 3 | hyper-parallel/hyper_parallel/core/multicore/ops/mega_moe/op_host/mega_moe_def.cpp:ops::MegaMoe::MegaMoe |
| 4 | hyper-parallel/hyper_parallel/core/multicore/ops/mega_moe_grad/op_host/mega_moe_grad_def.cpp:ops::MegaMoeGrad::MegaMoeGrad | 4 | hyper-parallel/hyper_parallel/core/multicore/ops/mega_moe_grad/op_host/mega_moe_grad_def.cpp:ops::MegaMoeGrad::MegaMoeGrad |
| 5 | hyper-parallel/hyper_parallel/integration/llamafactory/context_parallel/models/qwen3_vl/qwen3vl_forward.py:forward | 5 | hyper-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 | ||
| 16 | from __future__ import annotations | 22 | from __future__ import annotations |
| 17 | 23 | ||
| 18 | from typing import NamedTuple, Optional | 24 | from typing import NamedTuple, Optional |
| @@ -29,7 +35,11 @@ from hyper_parallel.core.context_parallel.context_parallel import ( | |||
| 29 | from hyper_parallel.core.dtensor.device_mesh import DeviceMesh | 35 | from hyper_parallel.core.dtensor.device_mesh import DeviceMesh |
| 30 | from hyper_parallel.core.dtensor.dtensor import DTensor | 36 | from hyper_parallel.core.dtensor.dtensor import DTensor |
| 31 | from hyper_parallel.core.tensor_parallel.style import ParallelStyle | 37 | from hyper_parallel.core.tensor_parallel.style import ParallelStyle |
| 32 | -from hyper_parallel.models.modules.linear_attention import torch_chunk_gated_delta_rule | 38 | +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 | +) | ||
| 33 | from hyper_parallel.platform import get_platform | 43 | from hyper_parallel.platform import get_platform |
| 34 | 44 | ||
| 35 | 45 | ||
| @@ -646,6 +656,290 @@ def _gdn_state_p2p_summary( | |||
| 646 | return core_attn_out | 656 | 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 649 | def _differentiable_all_to_all_shard( | 943 | def _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 = module | 1041 | 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_size | 1180 | 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): | |||
| 934 | class LinearAttentionP2PCPWrapper(nn.Module): | 1238 | class 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 = module | 1249 | 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.conv1d | 1265 | 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): | |||
| 1147 | class LinearAttentionContextParallel(ParallelStyle): | 1489 | class 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 = mode | 1508 | 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). | |||
| 33 | For inference with cache, use a kernel-optimised path or the recurrent | 33 | For inference with cache, use a kernel-optimised path or the recurrent |
| 34 | variant from ``transformers.models.qwen3_next``. | 34 | variant 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, A | 39 | # 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 | ||
| 38 | from typing import Optional | 45 | from typing import Optional |
| 39 | 46 | ||
| 40 | import torch | 47 | import torch |
| @@ -44,6 +51,14 @@ from torch.nn import functional as F | |||
| 44 | from hyper_parallel.models.modules.rmsnorm import RMSNormGated | 51 | from 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 | + | ||
| 47 | def _l2norm(x: torch.Tensor, dim: int = -1, eps: float = 1e-6) -> torch.Tensor: | 62 | def _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_state | 164 | 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 | + | ||
| 152 | class GatedDeltaNet(nn.Module): | 294 | class 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 | |||
| 19 | sequence axis so per-step gradients stay slice-faithful to the single-card | 19 | sequence axis so per-step gradients stay slice-faithful to the single-card |
| 20 | run. | 20 | run. |
| 21 | """ | 21 | """ |
| 22 | +# Qwen3.5 is currently registered as a Torch model implementation. | ||
| 23 | +# pylint: disable=forbidden-backend-import | ||
| 22 | from dataclasses import replace | 24 | from dataclasses import replace |
| 23 | from types import SimpleNamespace | 25 | from types import SimpleNamespace |
| 24 | from typing import Optional, TYPE_CHECKING | 26 | from typing import Optional, TYPE_CHECKING |
| @@ -476,9 +478,14 @@ def qwen3_5_tp_load_transforms( | |||
| 476 | return transforms | 478 | 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 | ||
| 484 | def _validate_qwen3_5_tp_config(model: Qwen3_5ForCausalLM, tp_world: int) -> None: | 491 | def _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 causal | 705 | 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 matching | 708 | + Linear-attention (:class:`Qwen3_5GatedDeltaNet`) layers select Ulysses, |
| 701 | - pure-Ulysses execution wrapper: project local sequence shards, all-to-all | 709 | + 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, run | 710 | + ``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 the | 711 | + 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 each | 714 | # 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 a | 715 | # 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 += 1 | 733 | 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 += 1 | 741 | 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 model | 748 | 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 42 | + configs=get_autotune_config(multibuffer_list=(False,)), | ||
| 43 | + key=['H', 'K', 'V', 'BT'], | ||
| 44 | +) | ||
| 45 | + | ||
| 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, | ||
| 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 | + | ||
| 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 | + | ||
| 322 | + configs=get_autotune_config(multibuffer_list=(True, False)), | ||
| 323 | + key=['H', 'K', 'V', 'BT', 'BV', 'USE_G', 'IS_VARLEN'], | ||
| 324 | +) | ||
| 325 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
🔵 Low Priority 第 578-579 行定义了一个 changed line(第 578-579 行): 建议:删除未使用的 ![]() ![]() 不准确? | |||
| 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 | + | ||
| 32 | + 'USE_G': lambda args: args['g'] is not None, | ||
| 33 | + 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None, | ||
| 34 | +}) | ||
| 35 | + | ||
| 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 | + | ||
| 122 | + 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None | ||
| 123 | +}) | ||
| 124 | + | ||
| 125 | + configs=[ | ||
| 126 | + triton.Config({'BK': BK}) | ||
| 127 | + for BK in [32, 64] | ||
| 128 | + ], | ||
| 129 | + key=["BC"] | ||
| 130 | +) | ||
| 131 | + | ||
| 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 | + | ||
| 196 | + 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None | ||
| 197 | +}) | ||
| 198 | + | ||
| 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 | + | ||
| 32 | + 'HAS_SCALE': lambda args: args['scale'] is not None, | ||
| 33 | + 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None | ||
| 34 | +}) | ||
| 35 | + | ||
| 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 | + ) | ||


🔵 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+ 风格)。cu_seqlens:Optional[torch.LongTensor]= None,