From d9ce99c034b116e5d9ba567d81f2aa2ee14024e6 Mon Sep 17 00:00:00 2001 From: sinle4cat Date: Thu, 16 Jul 2026 10:21:15 +0800 Subject: [PATCH] perf: optimize causal conv1d and mega chunk gdn for qwen3.5. --- test/python_test/RegisterOps.cpp | 118 ++++++++++++ test/python_test/custom_ops.py | 20 +- test/python_test/test_causal_conv1d_qkv.py | 150 +++++++++++++++ test/python_test/test_mega_chunk_gdn.py | 40 ++-- xllm_ops/build_aclnn.sh | 2 + .../op_host/causal_conv1d_tiling.cpp | 59 +++++- .../op_host/causal_conv1d_tiling.h | 13 ++ .../op_host/causal_conv1d_tiling_planner.h | 12 +- .../op_host/causal_conv1d_tiling_utils.h | 8 + .../op_host/causal_conv1d_tiling_validation.h | 44 ++++- .../causal_conv1d/op_kernel/causal_conv1d.h | 151 ++++++++++++++- .../op_kernel/causal_conv1d_common.h | 6 + .../op_kernel/causal_conv1d_fn.h | 1 + .../op_kernel/causal_conv1d_fn_tasks.h | 44 ++++- .../op_kernel/causal_conv1d_tiling_data.h | 15 ++ .../op_kernel/causal_conv1d_tiling_key.h | 6 + xllm_ops/causal_conv1d_qkv/CMakeLists.txt | 6 + .../causal_conv1d_qkv/op_host/CMakeLists.txt | 19 ++ .../op_host/causal_conv1d_qkv_def.cpp | 82 +++++++++ .../op_host/causal_conv1d_qkv_proto.cpp | 18 ++ .../op_host/causal_conv1d_qkv_tiling.cpp | 10 + .../op_kernel/causal_conv1d_qkv.cpp | 41 +++++ .../op_host/mega_chunk_gdn_def.cpp | 53 +++--- xllm_ops/mega_chunk_gdn/op_kernel/chunk_h.cpp | 134 +++++++------- xllm_ops/mega_chunk_gdn/op_kernel/chunk_o.cpp | 172 +++++++++--------- .../op_kernel/mega_chunk_gdn.cpp | 69 ++++--- .../op_kernel/scaled_dot_kkt.cpp | 68 +++---- xllm_ops/mega_chunk_gdn/op_kernel/wy_fast.cpp | 92 +++++----- 28 files changed, 1126 insertions(+), 327 deletions(-) create mode 100644 test/python_test/test_causal_conv1d_qkv.py create mode 100644 xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling.h create mode 100644 xllm_ops/causal_conv1d_qkv/CMakeLists.txt create mode 100644 xllm_ops/causal_conv1d_qkv/op_host/CMakeLists.txt create mode 100644 xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_def.cpp create mode 100644 xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_proto.cpp create mode 100644 xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_tiling.cpp create mode 100644 xllm_ops/causal_conv1d_qkv/op_kernel/causal_conv1d_qkv.cpp diff --git a/test/python_test/RegisterOps.cpp b/test/python_test/RegisterOps.cpp index 9a760b9..c315082 100644 --- a/test/python_test/RegisterOps.cpp +++ b/test/python_test/RegisterOps.cpp @@ -173,6 +173,115 @@ at::Tensor causal_conv1d( return output; } +at::Tensor causal_conv1d_packed_qkv_general_impl( + const at::Tensor& x, + const at::Tensor& weight, + const at::Tensor& conv_state, + at::IntArrayRef query_start_loc_opt, + at::IntArrayRef cache_indices_opt, + at::IntArrayRef initial_state_mode_opt, + int64_t q_dim, + int64_t k_dim, + int64_t v_dim, + int64_t head_dim, + at::ScalarType output_dtype) +{ + constexpr int64_t kPackedQkvActivationMode = 2; + constexpr int64_t kPadSlotId = -1; + constexpr int64_t kForwardRunMode = 0; + at::Tensor output = at::empty(x.sizes(), x.options().dtype(output_dtype)); + c10::optional bias_opt = c10::nullopt; + at::IntArrayRef num_accepted_tokens_opt; + EXEC_NPU_CMD(aclnnCausalConv1dQkv, + x, + weight, + bias_opt, + conv_state, + query_start_loc_opt, + cache_indices_opt, + initial_state_mode_opt, + num_accepted_tokens_opt, + kPackedQkvActivationMode, + kPadSlotId, + kForwardRunMode, + q_dim, + k_dim, + v_dim, + head_dim, + output); + return output; +} + +at::Tensor causal_conv1d_packed_qkv_general( + const at::Tensor& x, + const at::Tensor& weight, + const at::Tensor& conv_state, + at::IntArrayRef query_start_loc_opt, + at::IntArrayRef cache_indices_opt, + at::IntArrayRef initial_state_mode_opt, + int64_t q_dim, + int64_t k_dim, + int64_t v_dim, + int64_t head_dim) +{ + return causal_conv1d_packed_qkv_general_impl(x, + weight, + conv_state, + query_start_loc_opt, + cache_indices_opt, + initial_state_mode_opt, + q_dim, + k_dim, + v_dim, + head_dim, + at::kHalf); +} + +at::Tensor causal_conv1d_packed_qkv_general_bf16( + const at::Tensor& x, + const at::Tensor& weight, + const at::Tensor& conv_state, + at::IntArrayRef query_start_loc_opt, + at::IntArrayRef cache_indices_opt, + at::IntArrayRef initial_state_mode_opt, + int64_t q_dim, + int64_t k_dim, + int64_t v_dim, + int64_t head_dim) +{ + return causal_conv1d_packed_qkv_general_impl(x, + weight, + conv_state, + query_start_loc_opt, + cache_indices_opt, + initial_state_mode_opt, + q_dim, + k_dim, + v_dim, + head_dim, + at::kBFloat16); +} + +at::Tensor causal_conv1d_packed_qkv( + const at::Tensor& x, + const at::Tensor& weight, + const at::Tensor& conv_state, + at::IntArrayRef query_start_loc_opt, + at::IntArrayRef cache_indices_opt, + at::IntArrayRef initial_state_mode_opt) +{ + return causal_conv1d_packed_qkv_general(x, + weight, + conv_state, + query_start_loc_opt, + cache_indices_opt, + initial_state_mode_opt, + 1024, + 1024, + 3072, + 128); +} + at::Tensor recurrent_gated_delta_rule( const at::Tensor &query, const at::Tensor &key, @@ -1026,6 +1135,15 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { &beam_search_rec_final_select_impl_npu, "beam_search_rec_final_select"); m.def("causal_conv1d", &causal_conv1d, "causal_conv1d"); + m.def("causal_conv1d_packed_qkv", + &causal_conv1d_packed_qkv, + "causal_conv1d_packed_qkv"); + m.def("causal_conv1d_packed_qkv_general", + &causal_conv1d_packed_qkv_general, + "causal_conv1d_packed_qkv_general"); + m.def("causal_conv1d_packed_qkv_general_bf16", + &causal_conv1d_packed_qkv_general_bf16, + "causal_conv1d_packed_qkv_general_bf16"); m.def("recurrent_gated_delta_rule", &recurrent_gated_delta_rule, "recurrent_gated_delta_rule"); m.def("rec_constrained_topk", &rec_constrained_topk_impl_npu, "rec_constrained_topk"); m.def("mega_chunk_gdn", &mega_chunk_gdn, "mega_chunk_gdn"); diff --git a/test/python_test/custom_ops.py b/test/python_test/custom_ops.py index 18e51cc..52eceb1 100644 --- a/test/python_test/custom_ops.py +++ b/test/python_test/custom_ops.py @@ -114,10 +114,10 @@ def _mega_get_masks(device): return _MEGA_MASK_CACHE[key] -def _mega_get_minus_identity(device): - key = _mega_device_key(device) +def _mega_get_minus_identity(device, dtype): + key = (*_mega_device_key(device), dtype) if key not in _MEGA_MINUS_IDENTITY_CACHE: - minus_identity = torch.zeros(CHUNK_SIZE, CHUNK_SIZE, device=device, dtype=torch.float16) + minus_identity = torch.zeros(CHUNK_SIZE, CHUNK_SIZE, device=device, dtype=dtype) minus_identity.fill_diagonal_(-1) _MEGA_MINUS_IDENTITY_CACHE[key] = minus_identity return _MEGA_MINUS_IDENTITY_CACHE[key] @@ -265,7 +265,13 @@ def mega_chunk_gdn_npu(q, k, v, g, beta, scale=None, initial_state=None, scale = k.shape[-1] ** -0.5 q_dtype, k_dtype, v_dtype = q.dtype, k.dtype, v.dtype - q, k, v, beta = (t.half() for t in (q, k, v, beta)) + supported_compute_dtypes = (torch.float16, torch.bfloat16) + compute_dtype = ( + q.dtype + if q.dtype in supported_compute_dtypes and k.dtype == q.dtype and v.dtype == q.dtype + else torch.float16 + ) + q, k, v, beta = (t.to(compute_dtype) for t in (q, k, v, beta)) _, total_tokens, _, _ = q.shape num_value_heads = v.shape[-2] @@ -279,7 +285,7 @@ def mega_chunk_gdn_npu(q, k, v, g, beta, scale=None, initial_state=None, num_matrices = num_chunks * num_value_heads has_initial_state = initial_state is not None if has_initial_state: - initial_state_arg = initial_state.half() + initial_state_arg = initial_state.to(compute_dtype) else: initial_state_arg = torch.zeros( num_sequences, @@ -287,11 +293,11 @@ def mega_chunk_gdn_npu(q, k, v, g, beta, scale=None, initial_state=None, q.shape[-1], q.shape[-1], device=q.device, - dtype=torch.float16, + dtype=compute_dtype, ) mask_lower, mask_full = _mega_get_masks(q.device) - minus_identity = _mega_get_minus_identity(q.device) + minus_identity = _mega_get_minus_identity(q.device, compute_dtype) (out, g_sum, _, _, _, _, a_inv, w, _, h, v_new, final_state) = custom_ops_lib.mega_chunk_gdn( q, diff --git a/test/python_test/test_causal_conv1d_qkv.py b/test/python_test/test_causal_conv1d_qkv.py new file mode 100644 index 0000000..efb590e --- /dev/null +++ b/test/python_test/test_causal_conv1d_qkv.py @@ -0,0 +1,150 @@ +import pytest +import torch +import torch.nn.functional as F + + +torch_npu = pytest.importorskip("torch_npu") +custom_ops_lib = pytest.importorskip("custom_ops_lib") + +HEAD_DIM = 128 +WIDTH = 4 + + +def _cumulative_lengths(lengths): + result = [0] + for length in lengths: + result.append(result[-1] + length) + return tuple(result) + + +def _conv_reference( + x, + weight, + conv_state, + lengths, + cache_indices, + initial_state_mode, +): + outputs = [] + token_offset = 0 + weight_fp32 = weight.float() + for batch_idx, length in enumerate(lengths): + slot = cache_indices[batch_idx] + sequence = x[token_offset : token_offset + length] + history = ( + conv_state[slot].clone() + if initial_state_mode[batch_idx] + else torch.zeros((WIDTH - 1, x.size(1)), dtype=x.dtype) + ) + padded = torch.cat((history, sequence), dim=0) + accumulator = torch.zeros_like(sequence, dtype=torch.float32) + for tap in range(WIDTH): + accumulator.add_( + padded[tap : tap + length].float() * weight_fp32[tap] + ) + outputs.append(F.silu(accumulator).to(x.dtype)) + conv_state[slot].copy_(padded[-(WIDTH - 1) :]) + token_offset += length + return torch.cat(outputs, dim=0) + + +def _packed_reference(conv_output, q_dim, k_dim, v_dim, output_dtype): + q, k, v = torch.split(conv_output, (q_dim, k_dim, v_dim), dim=-1) + + def normalize_qk(value): + heads = value.view(value.size(0), -1, HEAD_DIM).float() + norm = torch.sqrt((heads * heads).sum(dim=-1, keepdim=True) + 1.0e-6) + return (heads / norm).to(torch.bfloat16).to(output_dtype) + + return torch.cat( + ( + normalize_qk(q).reshape(-1), + normalize_qk(k).reshape(-1), + v.to(output_dtype).reshape(-1), + ) + ).view_as(conv_output) + + +@pytest.mark.parametrize( + ("q_heads", "v_heads", "lengths", "initial_state_mode"), + [ + pytest.param(16, 48, (64,), (1,), id="tp1"), + pytest.param(8, 24, (127, 129), (0, 1), id="tp2-ragged"), + pytest.param(4, 12, (256,), (0,), id="tp4"), + pytest.param(2, 6, (63, 65), (1, 0), id="tp8-ragged"), + pytest.param(1, 3, (129,), (1,), id="tp16"), + ], +) +@pytest.mark.parametrize( + "output_dtype", [torch.float16, torch.bfloat16], ids=["fp16", "bf16"] +) +def test_causal_conv1d_qkv_general( + q_heads, + v_heads, + lengths, + initial_state_mode, + output_dtype, +): + generator = torch.Generator(device="cpu") + generator.manual_seed(20260716 + q_heads + v_heads) + q_dim = q_heads * HEAD_DIM + k_dim = q_dim + v_dim = v_heads * HEAD_DIM + conv_dim = q_dim + k_dim + v_dim + total_tokens = sum(lengths) + num_slots = len(lengths) + 2 + cache_indices = tuple(range(1, len(lengths) + 1)) + + x = (torch.randn((total_tokens, conv_dim), generator=generator) * 0.25).to( + torch.bfloat16 + ) + weight = (torch.randn((WIDTH, conv_dim), generator=generator) * 0.25).to( + torch.bfloat16 + ) + conv_state = ( + torch.randn((num_slots, WIDTH - 1, conv_dim), generator=generator) * 0.1 + ).to(torch.bfloat16) + conv_state_ref = conv_state.clone() + + conv_output = _conv_reference( + x, + weight, + conv_state_ref, + lengths, + cache_indices, + initial_state_mode, + ) + expected = _packed_reference( + conv_output, q_dim, k_dim, v_dim, output_dtype + ) + + device = torch.device("npu") + x_npu = x.to(device) + weight_npu = weight.to(device) + conv_state_npu = conv_state.to(device) + op = ( + custom_ops_lib.causal_conv1d_packed_qkv_general_bf16 + if output_dtype == torch.bfloat16 + else custom_ops_lib.causal_conv1d_packed_qkv_general + ) + actual = op( + x_npu, + weight_npu, + conv_state_npu, + _cumulative_lengths(lengths), + cache_indices, + initial_state_mode, + q_dim, + k_dim, + v_dim, + HEAD_DIM, + ) + torch.npu.synchronize() + + assert actual.dtype == output_dtype + atol = 1.0e-2 if output_dtype == torch.bfloat16 else 1.0e-3 + rtol = atol + torch.testing.assert_close(actual.cpu(), expected, atol=atol, rtol=rtol) + torch.testing.assert_close( + conv_state_npu.cpu(), conv_state_ref, atol=0.0, rtol=0.0 + ) diff --git a/test/python_test/test_mega_chunk_gdn.py b/test/python_test/test_mega_chunk_gdn.py index a1893d9..20309c5 100644 --- a/test/python_test/test_mega_chunk_gdn.py +++ b/test/python_test/test_mega_chunk_gdn.py @@ -143,14 +143,14 @@ def _native_reference(q, k, v, g, beta, cu_seqlens=None, initial_state=None, out return torch.cat(outs, dim=1), torch.cat(final_states, dim=0) if output_final_state else None -def _make_inputs(total_tokens, num_value_heads, num_key_heads, seed): +def _make_inputs(total_tokens, num_value_heads, num_key_heads, seed, dtype=torch.float16): torch.manual_seed(seed) head_dim = 128 - q = F.normalize(torch.randn(1, total_tokens, num_key_heads, head_dim), p=2, dim=-1).half() - k = F.normalize(torch.randn(1, total_tokens, num_key_heads, head_dim), p=2, dim=-1).half() - v = torch.randn(1, total_tokens, num_value_heads, head_dim, dtype=torch.float16) + q = F.normalize(torch.randn(1, total_tokens, num_key_heads, head_dim), p=2, dim=-1).to(dtype) + k = F.normalize(torch.randn(1, total_tokens, num_key_heads, head_dim), p=2, dim=-1).to(dtype) + v = torch.randn(1, total_tokens, num_value_heads, head_dim, dtype=dtype) g = F.logsigmoid(torch.randn(1, total_tokens, num_value_heads, dtype=torch.float32)) - beta = torch.rand(1, total_tokens, num_value_heads, dtype=torch.float16) + beta = torch.rand(1, total_tokens, num_value_heads, dtype=dtype) return q, k, v, g, beta @@ -184,10 +184,14 @@ def _run_mega(q_cpu, k_cpu, v_cpu, g_cpu, beta_cpu, cu_list=None, initial_state_ pytest.param(512, [0, 96, 128, 512], 48, 16, id="long-varlen-H48-Hg16"), ], ) -def test_mega_chunk_gdn_e2e(total_tokens, cu_list, num_value_heads, num_key_heads): - q, k, v, g, beta = _make_inputs(total_tokens, num_value_heads, num_key_heads, seed=0) +@pytest.mark.parametrize("compute_dtype", [torch.float16, torch.bfloat16], ids=["fp16", "bf16"]) +def test_mega_chunk_gdn_e2e(total_tokens, cu_list, num_value_heads, num_key_heads, compute_dtype): + q, k, v, g, beta = _make_inputs( + total_tokens, num_value_heads, num_key_heads, seed=0, dtype=compute_dtype + ) actual, _ = _run_mega(q, k, v, g, beta, cu_list) expected, _ = _native_reference(q, k, v, g, beta, cu_list) + assert actual.dtype == compute_dtype _assert_close("mega_vs_native", actual, expected) @@ -199,13 +203,18 @@ def test_mega_chunk_gdn_e2e(total_tokens, cu_list, num_value_heads, num_key_head ], ) @pytest.mark.parametrize("state_kind", ["zero", "random"]) -def test_mega_chunk_gdn_initial_state(total_tokens, cu_list, num_value_heads, num_key_heads, state_kind): - q, k, v, g, beta = _make_inputs(total_tokens, num_value_heads, num_key_heads, seed=1) +@pytest.mark.parametrize("compute_dtype", [torch.float16, torch.bfloat16], ids=["fp16", "bf16"]) +def test_mega_chunk_gdn_initial_state( + total_tokens, cu_list, num_value_heads, num_key_heads, state_kind, compute_dtype +): + q, k, v, g, beta = _make_inputs( + total_tokens, num_value_heads, num_key_heads, seed=1, dtype=compute_dtype + ) num_sequences = 1 if cu_list is None else len(cu_list) - 1 if state_kind == "zero": - h0 = torch.zeros(num_sequences, num_value_heads, 128, 128, dtype=torch.float16) + h0 = torch.zeros(num_sequences, num_value_heads, 128, 128, dtype=compute_dtype) else: - h0 = (0.1 * torch.randn(num_sequences, num_value_heads, 128, 128)).half() + h0 = (0.1 * torch.randn(num_sequences, num_value_heads, 128, 128)).to(compute_dtype) actual, actual_final_state = _run_mega(q, k, v, g, beta, cu_list, h0, output_final_state=True) expected, expected_final_state = _native_reference(q, k, v, g, beta, cu_list, h0, output_final_state=True) @@ -214,11 +223,14 @@ def test_mega_chunk_gdn_initial_state(total_tokens, cu_list, num_value_heads, nu @pytest.mark.parametrize(("num_value_heads", "num_key_heads"), SUPPORTED_HEAD_CONFIGS) -def test_mega_chunk_gdn_supported_head_configs(num_value_heads, num_key_heads): +@pytest.mark.parametrize("compute_dtype", [torch.float16, torch.bfloat16], ids=["fp16", "bf16"]) +def test_mega_chunk_gdn_supported_head_configs(num_value_heads, num_key_heads, compute_dtype): total_tokens = 129 cu_list = [0, 64, total_tokens] - q, k, v, g, beta = _make_inputs(total_tokens, num_value_heads, num_key_heads, seed=2) - h0 = (0.05 * torch.randn(len(cu_list) - 1, num_value_heads, 128, 128)).half() + q, k, v, g, beta = _make_inputs( + total_tokens, num_value_heads, num_key_heads, seed=2, dtype=compute_dtype + ) + h0 = (0.05 * torch.randn(len(cu_list) - 1, num_value_heads, 128, 128)).to(compute_dtype) actual, actual_final_state = _run_mega(q, k, v, g, beta, cu_list, h0, output_final_state=True) expected, expected_final_state = _native_reference(q, k, v, g, beta, cu_list, h0, output_final_state=True) diff --git a/xllm_ops/build_aclnn.sh b/xllm_ops/build_aclnn.sh index 6bddae7..faba676 100644 --- a/xllm_ops/build_aclnn.sh +++ b/xllm_ops/build_aclnn.sh @@ -131,6 +131,7 @@ elif [[ "$SOC_VERSION" =~ ^(ascend)?910b ]]; then "x_attention" "cache_unshared_kv" "causal_conv1d" + "causal_conv1d_qkv" "convert_kv_cache_format" "beam_search" "index_group_matmul" @@ -239,6 +240,7 @@ elif [[ "$SOC_VERSION" =~ ^ascend910_93 ]]; then "x_attention" "cache_unshared_kv" "causal_conv1d" + "causal_conv1d_qkv" "convert_kv_cache_format" "beam_search" "index_group_matmul" diff --git a/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling.cpp b/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling.cpp index c164278..1b4329a 100644 --- a/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling.cpp +++ b/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling.cpp @@ -15,6 +15,7 @@ */ #include "causal_conv1d_tiling_planner.h" +#include "causal_conv1d_tiling.h" #include "causal_conv1d_tiling_validation.h" #include #include @@ -32,9 +33,59 @@ bool HostTilingDebugEnabled() return value != nullptr && value[0] != '\0' && value[0] != '0'; } +FnHostPlan ChoosePackedQkvFnHostPlan(gert::TilingContext *context, const CausalConv1dTilingData &tiling, + uint64_t ubSize, uint32_t coreNum) +{ + FnHostPlan plan; + if (tiling.inputMode != 0 || tiling.batch <= 0 || tiling.cuSeqlen <= 0 || tiling.dim <= 0 || + tiling.packedHeadDim <= 0 || coreNum == 0) { + return plan; + } + + if (tiling.dim == CAUSAL_CONV1D_PACKED_QKV_LEGACY_DIM && + tiling.packedQDim + tiling.packedKDim == CAUSAL_CONV1D_PACKED_QK_LEGACY_DIM && + tiling.packedVDim == CAUSAL_CONV1D_PACKED_V_LEGACY_DIM) { + plan.baseDimChoice.baseDim = CAUSAL_CONV1D_PACKED_V_LEGACY_DIM; + plan.baseDimChoice.baseDimCnt = 2; + } else if (tiling.dim <= MAX_DIM_TILE_SIZE) { + plan.baseDimChoice.baseDim = tiling.dim; + plan.baseDimChoice.baseDimCnt = 1; + } else { + const int64_t ubLimitedBaseDim = + AlignDownInt64(ComputeFnUbLimitedBaseDim(ubSize), tiling.packedHeadDim); + if (ubLimitedBaseDim <= 0) { + return plan; + } + const int64_t baseDimCnt = CeilDivInt64(tiling.dim, ubLimitedBaseDim); + const int64_t balancedBaseDim = + AlignUpInt64(CeilDivInt64(tiling.dim, baseDimCnt), tiling.packedHeadDim); + plan.baseDimChoice.baseDim = + (balancedBaseDim > 0 && balancedBaseDim <= ubLimitedBaseDim) ? balancedBaseDim : ubLimitedBaseDim; + plan.baseDimChoice.baseDimCnt = CeilDivInt64(tiling.dim, plan.baseDimChoice.baseDim); + } + + plan.executionPlan = + static_cast(ResolveFnExecutionPlan(plan.baseDimChoice.baseDimCnt)); + plan.caseKind = (plan.baseDimChoice.baseDimCnt == 1) ? FN_TILING_CASE_TOKEN_FIRST + : FN_TILING_CASE_TOKEN_DIM_CO_SPLIT; + plan.baseDimChoice.gridSize = tiling.batch * plan.baseDimChoice.baseDimCnt; + plan.tokenBlockChoice = + ChooseUnifiedFnTokenBlockPlan(context, tiling, plan.baseDimChoice, plan.executionPlan, coreNum); + if (!plan.tokenBlockChoice.enabled) { + return {}; + } + plan.tokenCoreMapping = BuildFnTokenCoreMappingChoice( + plan.tokenBlockChoice.tokenBlockCnt, plan.baseDimChoice.baseDimCnt, plan.executionPlan, coreNum); + if (plan.tokenCoreMapping.tokenCoreBudget <= 0 || plan.tokenCoreMapping.blockDim <= 0) { + return {}; + } + plan.tokenCoreMapping.blockDim = std::min(plan.tokenCoreMapping.blockDim, coreNum); + return plan; +} + } // namespace -static ge::graphStatus CausalConv1dTilingFunc(gert::TilingContext *context) +ge::graphStatus CausalConv1dTilingFunc(gert::TilingContext *context) { uint64_t ubSize = 0; uint32_t coreNum = 0; @@ -71,7 +122,9 @@ static ge::graphStatus CausalConv1dTilingFunc(gert::TilingContext *context) } if (isFn) { - fnHostPlan = ChooseFnHostPlan(context, *tiling, ubSize, coreNum); + fnHostPlan = attrInfo.activationMode == CAUSAL_CONV1D_ACTIVATION_SILU_PACKED_QKV + ? ChoosePackedQkvFnHostPlan(context, *tiling, ubSize, coreNum) + : ChooseFnHostPlan(context, *tiling, ubSize, coreNum); plannerModeTag = GetFnTilingCaseName(fnHostPlan.caseKind); baseDimChoice = fnHostPlan.baseDimChoice; fnExecutionPlan = fnHostPlan.executionPlan; @@ -168,7 +221,7 @@ static ge::graphStatus CausalConv1dTilingFunc(gert::TilingContext *context) return ge::GRAPH_SUCCESS; } -static ge::graphStatus TilingParseForCausalConv1d(gert::TilingParseContext *context) +ge::graphStatus TilingParseForCausalConv1d(gert::TilingParseContext *context) { OP_LOGD(context, "Enter TilingParseForCausalConv1d."); return ge::GRAPH_SUCCESS; diff --git a/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling.h b/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling.h new file mode 100644 index 0000000..4d66e98 --- /dev/null +++ b/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling.h @@ -0,0 +1,13 @@ +#ifndef CAUSAL_CONV1D_HOST_TILING_H +#define CAUSAL_CONV1D_HOST_TILING_H + +#include "register/op_impl_registry.h" + +namespace optiling { + +ge::graphStatus CausalConv1dTilingFunc(gert::TilingContext *context); +ge::graphStatus TilingParseForCausalConv1d(gert::TilingParseContext *context); + +} // namespace optiling + +#endif // CAUSAL_CONV1D_HOST_TILING_H diff --git a/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_planner.h b/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_planner.h index 8d21bb1..34aec9e 100644 --- a/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_planner.h +++ b/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_planner.h @@ -126,11 +126,19 @@ inline DimTileChoice ChooseFnTokenDimCoSplitBaseDimChoice(gert::TilingContext *c result.baseDimCnt = CeilDivInt64(dim, result.baseDim); result.gridSize = result.baseDimCnt; + const int64_t balancedBaseDim = AlignUpInt64(CeilDivInt64(dim, result.baseDimCnt), DIM_ALIGN_ELEMS); + if (balancedBaseDim > 0 && balancedBaseDim <= ubLimitedBaseDim) { + result.baseDim = balancedBaseDim; + result.baseDimCnt = CeilDivInt64(dim, result.baseDim); + result.gridSize = result.baseDimCnt; + } + if (coreNum == 0 || result.baseDimCnt <= 1 || result.baseDimCnt >= static_cast(coreNum) || (coreNum % result.baseDimCnt == 0)) { OP_LOGD(context, - "FnDimCoSplit: dim[%ld], ubLimitedBaseDim[%ld], baseDimCnt[%ld], coreNum[%u], adjusted[%d].", dim, - result.baseDim, result.baseDimCnt, coreNum, 0); + "FnDimCoSplit: dim[%ld], ubLimitedBaseDim[%ld], balancedBaseDim[%ld], baseDimCnt[%ld], coreNum[%u], " + "adjusted[%d].", + dim, ubLimitedBaseDim, result.baseDim, result.baseDimCnt, coreNum, 0); return result; } diff --git a/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_utils.h b/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_utils.h index 90af198..88731e6 100644 --- a/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_utils.h +++ b/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_utils.h @@ -31,6 +31,10 @@ constexpr uint32_t NUM_ACCEPTED_TOKENS_INDEX = 7; constexpr int32_t ATTR_ACTIVATION_MODE_INDEX = 0; constexpr int32_t ATTR_PAD_SLOT_ID_INDEX = 1; constexpr int32_t ATTR_RUN_MODE_INDEX = 2; +constexpr int32_t ATTR_PACKED_Q_DIM_INDEX = 3; +constexpr int32_t ATTR_PACKED_K_DIM_INDEX = 4; +constexpr int32_t ATTR_PACKED_V_DIM_INDEX = 5; +constexpr int32_t ATTR_PACKED_HEAD_DIM_INDEX = 6; constexpr int64_t ASCENDC_RESERVED_WORKSPACE_SIZE = 16 * 1024 * 1024; @@ -43,6 +47,10 @@ struct CausalConv1dAttrInfo { int64_t activationMode = 0; int64_t padSlotId = -1; int64_t runMode = 0; + int64_t packedQDim = 0; + int64_t packedKDim = 0; + int64_t packedVDim = 0; + int64_t packedHeadDim = 0; }; struct DimTileChoice { diff --git a/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_validation.h b/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_validation.h index 6a630e8..9ce4b59 100644 --- a/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_validation.h +++ b/xllm_ops/causal_conv1d/op_host/causal_conv1d_tiling_validation.h @@ -60,8 +60,10 @@ inline ge::graphStatus GetAttrsInfo(gert::TilingContext *context, CausalConv1dAt const int64_t *activationModePtr = attrs->GetAttrPointer(ATTR_ACTIVATION_MODE_INDEX); OP_CHECK_NULL_WITH_CONTEXT(context, activationModePtr); attrInfo.activationMode = *activationModePtr; - OP_CHECK_IF(attrInfo.activationMode != 0 && attrInfo.activationMode != 1, - OP_LOGE(context, "activationMode only supports 0/1"), + OP_CHECK_IF(attrInfo.activationMode != CAUSAL_CONV1D_ACTIVATION_NONE && + attrInfo.activationMode != CAUSAL_CONV1D_ACTIVATION_SILU && + attrInfo.activationMode != CAUSAL_CONV1D_ACTIVATION_SILU_PACKED_QKV, + OP_LOGE(context, "activationMode only supports 0/1/2"), return ge::GRAPH_FAILED); const int64_t *padSlotIdPtr = attrs->GetAttrPointer(ATTR_PAD_SLOT_ID_INDEX); @@ -72,6 +74,21 @@ inline ge::graphStatus GetAttrsInfo(gert::TilingContext *context, CausalConv1dAt attrInfo.runMode = (runModePtr == nullptr) ? 0 : *runModePtr; OP_CHECK_IF(attrInfo.runMode != 0 && attrInfo.runMode != 1, OP_LOGE(context, "runMode only supports 0/1"), return ge::GRAPH_FAILED); + + if (attrInfo.activationMode == CAUSAL_CONV1D_ACTIVATION_SILU_PACKED_QKV) { + const int64_t *packedQDimPtr = attrs->GetAttrPointer(ATTR_PACKED_Q_DIM_INDEX); + const int64_t *packedKDimPtr = attrs->GetAttrPointer(ATTR_PACKED_K_DIM_INDEX); + const int64_t *packedVDimPtr = attrs->GetAttrPointer(ATTR_PACKED_V_DIM_INDEX); + const int64_t *packedHeadDimPtr = attrs->GetAttrPointer(ATTR_PACKED_HEAD_DIM_INDEX); + OP_CHECK_NULL_WITH_CONTEXT(context, packedQDimPtr); + OP_CHECK_NULL_WITH_CONTEXT(context, packedKDimPtr); + OP_CHECK_NULL_WITH_CONTEXT(context, packedVDimPtr); + OP_CHECK_NULL_WITH_CONTEXT(context, packedHeadDimPtr); + attrInfo.packedQDim = *packedQDimPtr; + attrInfo.packedKDim = *packedKDimPtr; + attrInfo.packedVDim = *packedVDimPtr; + attrInfo.packedHeadDim = *packedHeadDimPtr; + } return ge::GRAPH_SUCCESS; } @@ -93,6 +110,10 @@ inline ge::graphStatus GetShapeDtypeInfo(gert::TilingContext *context, const Cau const bool isDecodeMode = (attrInfo.runMode == 1); tiling.activationMode = attrInfo.activationMode; tiling.padSlotId = attrInfo.padSlotId; + tiling.packedQDim = attrInfo.packedQDim; + tiling.packedKDim = attrInfo.packedKDim; + tiling.packedVDim = attrInfo.packedVDim; + tiling.packedHeadDim = attrInfo.packedHeadDim; auto xShapePtr = context->GetInputShape(X_INDEX); OP_CHECK_NULL_WITH_CONTEXT(context, xShapePtr); @@ -419,6 +440,25 @@ inline ge::graphStatus GetShapeDtypeInfo(gert::TilingContext *context, const Cau OP_CHECK_IF(sDesc->GetDataType() != xDtype, OP_LOGE(context, "convStates dtype must equal x dtype"), return ge::GRAPH_FAILED); + if (attrInfo.activationMode == CAUSAL_CONV1D_ACTIVATION_SILU_PACKED_QKV) { + auto yDesc = context->GetOutputDesc(0); + OP_CHECK_NULL_WITH_CONTEXT(context, yDesc); + const int64_t qDim = attrInfo.packedQDim; + const int64_t kDim = attrInfo.packedKDim; + const int64_t vDim = attrInfo.packedVDim; + const int64_t headDim = attrInfo.packedHeadDim; + const ge::DataType yDtype = yDesc->GetDataType(); + const bool validPackedLayout = qDim > 0 && qDim == kDim && vDim > 0 && headDim == 128 && + qDim % headDim == 0 && kDim % headDim == 0 && vDim % headDim == 0 && + qDim + kDim + vDim == dim; + OP_CHECK_IF(attrInfo.runMode != 0 || inputMode != 0 || width != 4 || !validPackedLayout || hasBias || + xDtype != ge::DT_BF16 || (yDtype != ge::DT_FLOAT16 && yDtype != ge::DT_BF16), + OP_LOGE(context, + "activationMode=2 requires BF16 2D varlen prefill, width=4, qDim=kDim, " + "128-aligned Q/K/V dimensions whose sum equals dim, no bias, and FP16/BF16 output"), + return ge::GRAPH_FAILED); + } + if (!qslAbsent) { auto qslDesc2 = context->GetOptionalInputDesc(QUERY_START_LOC_INDEX); OP_CHECK_NULL_WITH_CONTEXT(context, qslDesc2); diff --git a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d.h b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d.h index 49f501c..a6f4e7c 100644 --- a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d.h +++ b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d.h @@ -19,6 +19,7 @@ #include "kernel_operator.h" #include "kernel_tiling/kernel_tiling.h" +#include "adv_api/pad/broadcast.h" #include "causal_conv1d_tiling_data.h" #include "causal_conv1d_tiling_key.h" #include "causal_conv1d_common.h" @@ -28,6 +29,10 @@ namespace NsCausalConv1d { using namespace AscendC; using namespace NsCausalConv1dCommon; +using PackedQkvT = DTYPE_Y; +static_assert(IsSameType::value || IsSameType::value, + "Packed QKV output supports FP16 or BF16."); + #define CAUSAL_CONV1D_TEMPLATE_ARGS typename T, uint32_t runModeKey, uint32_t widthKey, uint32_t fnPlanKey #define CAUSAL_CONV1D_CLASS CausalConv1d @@ -127,6 +132,10 @@ class CausalConv1d { __aicore__ inline void RestoreFnLocalPartials(int32_t baseDim); __aicore__ inline void ComputeFnRollingOutput(int32_t slotCurr, int32_t baseDim); __aicore__ inline void AdvanceFnLocalPartials(int32_t slotCurr, int32_t baseDim); + __aicore__ inline void PreparePackedQkvOutput(const LocalTensor &outSlotT, int32_t channelStart, + int32_t baseDim); + __aicore__ inline void WritePackedQkvOutput(const LocalTensor &outSlot, int32_t tokenIdx, + int32_t channelStart, int32_t baseDim, int32_t numTokens); __aicore__ inline void RunSeqFnRolling(int32_t start, int32_t len, int32_t channelStart, int32_t baseDim, int32_t dim); __aicore__ inline void RunSeq(int32_t start, int32_t len, int32_t channelStart, int32_t baseDim, int32_t dim); @@ -157,6 +166,7 @@ class CausalConv1d { int32_t baseDim, int32_t dim); __aicore__ inline const CausalConv1dTilingData *GetTilingData() const; __aicore__ inline bool HasActivation() const; + __aicore__ inline bool IsPackedQkvOutput() const; __aicore__ inline bool HasBias() const; __aicore__ inline bool IsUpdateMode() const; __aicore__ inline bool IsFnRollingFastPathEnabled() const; @@ -168,6 +178,7 @@ class CausalConv1d { TBuf inBuf; TBuf outBuf; TBuf calcBuf; + TBuf packedQkvNormBuf; TEventID weightBiasMte2ToVEvent_; TEventID stateMte2ToVEvent_; @@ -196,6 +207,7 @@ class CausalConv1d { GlobalTensor initialStateModeGm; GlobalTensor numAcceptedTokensGm; GlobalTensor yGm; + GlobalTensor packedQkvYGm; GlobalTensor initStateSyncGm_; GlobalTensor initStateWorkspaceGm_; @@ -212,8 +224,14 @@ template __aicore__ inline void CAUSAL_CONV1D_CLASS::InitSharedBuffersAndEvents() { pipe.InitBuffer(inBuf, RING_SLOTS * MAX_BLOCK_DIM * sizeof(T)); - pipe.InitBuffer(outBuf, 2 * MAX_BLOCK_DIM * sizeof(T)); + const int32_t outSlotCount = IsPackedQkvOutput() ? 1 : 2; + pipe.InitBuffer(outBuf, outSlotCount * MAX_BLOCK_DIM * sizeof(T)); pipe.InitBuffer(calcBuf, (MAX_WIDTH + 4) * MAX_BLOCK_DIM * sizeof(float)); + if (IsPackedQkvOutput()) { + pipe.InitBuffer(packedQkvNormBuf, + (PACKED_QKV_REDUCE_TMP_ELEMS + PACKED_QKV_REDUCE_SUM_ELEMS + + PACKED_QKV_NORM_ELEMS) * sizeof(float)); + } AllocEvents(); } @@ -609,6 +627,110 @@ __aicore__ inline void CAUSAL_CONV1D_CLASS::AdvanceFnLocalPartials(int32_t slotC } } +template +__aicore__ inline void CAUSAL_CONV1D_CLASS::PreparePackedQkvOutput(const LocalTensor &outSlotT, + int32_t channelStart, int32_t baseDim) +{ + if constexpr (!IsSameType::value) { + return; + } + + auto cl = CalcBufLayout::FromCalcBuf(calcBuf); + LocalTensor &headF = cl.currF; + LocalTensor normScratch = packedQkvNormBuf.Get(); + LocalTensor reduceTmpF = normScratch; + LocalTensor sumF = normScratch[PACKED_QKV_REDUCE_TMP_ELEMS]; + LocalTensor normF = normScratch[PACKED_QKV_REDUCE_TMP_ELEMS + PACKED_QKV_REDUCE_SUM_ELEMS]; + + const int32_t headDim = static_cast(tilingData_->packedHeadDim); + const int32_t qkDim = static_cast(tilingData_->packedQDim + tilingData_->packedKDim); + const int32_t channelEnd = channelStart + baseDim; + const int32_t qkStart = (channelStart > 0) ? channelStart : 0; + const int32_t qkEnd = (channelEnd < qkDim) ? channelEnd : qkDim; + const int32_t totalHeadCount = (headDim > 0) ? ((qkEnd - qkStart) / headDim) : 0; + for (int32_t headOffset = 0; headOffset < totalHeadCount; headOffset += PACKED_QKV_NORM_HEADS_PER_PASS) { + const int32_t remainingHeads = totalHeadCount - headOffset; + const int32_t headCount = (remainingHeads < PACKED_QKV_NORM_HEADS_PER_PASS) + ? remainingHeads + : PACKED_QKV_NORM_HEADS_PER_PASS; + const int32_t localOffset = qkStart - channelStart + headOffset * headDim; + const int32_t qkElems = headCount * headDim; + Cast(headF, outSlotT[localOffset], RoundMode::CAST_NONE, qkElems); + PipeBarrier(); + + BinaryRepeatParams reduceParams; + reduceParams.src0BlkStride = 1; + reduceParams.src1BlkStride = 1; + reduceParams.dstBlkStride = 1; + reduceParams.src0RepStride = headDim / 8; + reduceParams.src1RepStride = headDim / 8; + reduceParams.dstRepStride = 64 / 8; + Mul(reduceTmpF, headF, headF, 64, headCount, reduceParams); + PipeBarrier(); + MulAddDst(reduceTmpF, headF[64], headF[64], 64, headCount, reduceParams); + PipeBarrier(); + AscendCUtils::SetMask(64); + WholeReduceSum(sumF, reduceTmpF, 64, headCount, 1, 1, 64 / 8); + PipeBarrier(); + Adds(sumF, sumF, 1.0e-6f, headCount); + PipeBarrier(); + Sqrt(sumF, sumF, headCount); + PipeBarrier(); + + const uint32_t dstShape[2] = {static_cast(headCount), static_cast(headDim)}; + const uint32_t srcShape[2] = {static_cast(headCount), 1}; + Broadcast(normF, sumF, dstShape, srcShape); + PipeBarrier(); + Div(headF, headF, normF, qkElems); + PipeBarrier(); + Cast(outSlotT[localOffset], headF, RoundMode::CAST_RINT, qkElems); + PipeBarrier(); + } + + if constexpr (IsSameType::value) { + Cast(headF, outSlotT, RoundMode::CAST_NONE, baseDim); + PipeBarrier(); + Cast(outSlotT.template ReinterpretCast(), headF, RoundMode::CAST_RINT, baseDim); + PipeBarrier(); + } +} + +template +__aicore__ inline void CAUSAL_CONV1D_CLASS::WritePackedQkvOutput(const LocalTensor &outSlot, + int32_t tokenIdx, int32_t channelStart, + int32_t baseDim, int32_t numTokens) +{ + const int32_t qDim = static_cast(tilingData_->packedQDim); + const int32_t kDim = static_cast(tilingData_->packedKDim); + const int32_t vDim = static_cast(tilingData_->packedVDim); + const int32_t qkDim = qDim + kDim; + const int32_t totalDim = qkDim + vDim; + const int32_t channelEnd = channelStart + baseDim; + + const int32_t qStart = (channelStart > 0) ? channelStart : 0; + const int32_t qEnd = (channelEnd < qDim) ? channelEnd : qDim; + if (qStart < qEnd) { + const int64_t dstOffset = static_cast(tokenIdx) * qDim + qStart; + DataCopy(packedQkvYGm[dstOffset], outSlot[qStart - channelStart], qEnd - qStart); + } + + const int32_t kStart = (channelStart > qDim) ? channelStart : qDim; + const int32_t kEnd = (channelEnd < qkDim) ? channelEnd : qkDim; + if (kStart < kEnd) { + const int64_t dstOffset = static_cast(numTokens) * qDim + + static_cast(tokenIdx) * kDim + (kStart - qDim); + DataCopy(packedQkvYGm[dstOffset], outSlot[kStart - channelStart], kEnd - kStart); + } + + const int32_t vStart = (channelStart > qkDim) ? channelStart : qkDim; + const int32_t vEnd = (channelEnd < totalDim) ? channelEnd : totalDim; + if (vStart < vEnd) { + const int64_t dstOffset = static_cast(numTokens) * qkDim + + static_cast(tokenIdx) * vDim + (vStart - qkDim); + DataCopy(packedQkvYGm[dstOffset], outSlot[vStart - channelStart], vEnd - vStart); + } +} + template __aicore__ inline void CAUSAL_CONV1D_CLASS::RunSeqFnRolling(int32_t start, int32_t len, int32_t channelStart, int32_t baseDim, int32_t dim) @@ -640,9 +762,10 @@ __aicore__ inline void CAUSAL_CONV1D_CLASS::RunSeqFnRolling(int32_t start, int32 ComputeFnRollingOutput(slotCurr, baseDim); - const int32_t outSlot = t & 1; + const bool packedQkvOutput = IsPackedQkvOutput(); + const int32_t outSlot = packedQkvOutput ? 0 : (t & 1); LocalTensor outSlotT = outT[outSlot * MAX_BLOCK_DIM]; - if (t >= 2) { + if ((packedQkvOutput && t >= 1) || (!packedQkvOutput && t >= 2)) { WaitFlag(outMte3ToVEvent_[outSlot]); } @@ -662,12 +785,21 @@ __aicore__ inline void CAUSAL_CONV1D_CLASS::RunSeqFnRolling(int32_t start, int32 AdvanceFnLocalPartials(slotCurr, baseDim); + if (packedQkvOutput) { + PreparePackedQkvOutput(outSlotT, channelStart, baseDim); + } + SetFlag(outVToMte3Event_[outSlot]); - const int64_t outOffset = static_cast(start + t) * dim + channelStart; WaitFlag(outVToMte3Event_[outSlot]); - DataCopy(yGm[outOffset], outSlotT, baseDim); - if (t + 2 < len) { + if (packedQkvOutput) { + WritePackedQkvOutput(outSlotT.template ReinterpretCast(), start + t, channelStart, baseDim, + static_cast(tilingData_->cuSeqlen)); + } else { + const int64_t outOffset = static_cast(start + t) * dim + channelStart; + DataCopy(yGm[outOffset], outSlotT, baseDim); + } + if ((packedQkvOutput && t + 1 < len) || (!packedQkvOutput && t + 2 < len)) { SetFlag(outMte3ToVEvent_[outSlot]); } @@ -962,6 +1094,13 @@ __aicore__ inline bool CAUSAL_CONV1D_CLASS::HasActivation() const return (tilingData_ != nullptr) && (tilingData_->activationMode != 0); } +template +__aicore__ inline bool CAUSAL_CONV1D_CLASS::IsPackedQkvOutput() const +{ + return (tilingData_ != nullptr) && + (tilingData_->activationMode == CAUSAL_CONV1D_ACTIVATION_SILU_PACKED_QKV); +} + template __aicore__ inline bool CAUSAL_CONV1D_CLASS::HasBias() const { diff --git a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_common.h b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_common.h index 7e5a845..17f5e04 100644 --- a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_common.h +++ b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_common.h @@ -24,6 +24,12 @@ constexpr int32_t MAX_WIDTH = 4; constexpr int32_t MAX_BLOCK_DIM = 4096; constexpr int32_t RING_SLOTS = 5; +constexpr int32_t PACKED_QKV_HEAD_DIM = 128; +constexpr int32_t PACKED_QKV_NORM_HEADS_PER_PASS = 16; +constexpr int32_t PACKED_QKV_REDUCE_TMP_ELEMS = PACKED_QKV_NORM_HEADS_PER_PASS * 64; +constexpr int32_t PACKED_QKV_REDUCE_SUM_ELEMS = 64; +constexpr int32_t PACKED_QKV_NORM_ELEMS = PACKED_QKV_NORM_HEADS_PER_PASS * PACKED_QKV_HEAD_DIM; + __aicore__ inline int32_t SlotCurr(int32_t t) { return (t + 3) % RING_SLOTS; diff --git a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_fn.h b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_fn.h index 5b25cb3..0552469 100644 --- a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_fn.h +++ b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_fn.h @@ -33,6 +33,7 @@ class CausalConv1dFn : public CausalConv1dcacheIndicesGm.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(cacheIndices)); this->initialStateModeGm.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(initialStateMode)); this->yGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(y)); + this->packedQkvYGm.SetGlobalBuffer(reinterpret_cast<__gm__ PackedQkvT *>(y)); if (tilingData->hasInitStateWorkspace != 0) { const uint64_t syncElems = static_cast(GetBlockNum()) * INIT_STATE_SYNCALL_NEED_SIZE; diff --git a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_fn_tasks.h b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_fn_tasks.h index ce34900..2b9a8b0 100644 --- a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_fn_tasks.h +++ b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_fn_tasks.h @@ -60,6 +60,40 @@ __aicore__ inline FnDirectBlockTask ResolveFnDirectBlockTask(int32_t blockIdx, i return task; } +__aicore__ inline FnDirectBlockTask ResolveFnPackedQkvBlockTask(int32_t blockIdx, int32_t tokenBlockCnt, + int32_t tokenBlockSize, int32_t cuSeqlen, + int32_t baseDimCnt, int32_t baseDim, int32_t dim, + int32_t qDim, int32_t kDim, int32_t vDim) +{ + auto task = ResolveFnDirectBlockTask(blockIdx, tokenBlockCnt, tokenBlockSize, cuSeqlen, baseDimCnt, baseDim, dim); + if (!task.valid) { + return task; + } + const bool useLegacyTwoBlockMapping = baseDimCnt == 2 && dim == CAUSAL_CONV1D_PACKED_QKV_LEGACY_DIM && + qDim + kDim == CAUSAL_CONV1D_PACKED_QK_LEGACY_DIM && + vDim == CAUSAL_CONV1D_PACKED_V_LEGACY_DIM; + if (useLegacyTwoBlockMapping) { + if (task.baseDimIdx == 0) { + task.channelStart = 0; + task.baseDimSize = CAUSAL_CONV1D_PACKED_QK_LEGACY_DIM; + } else { + task.channelStart = CAUSAL_CONV1D_PACKED_QK_LEGACY_DIM; + task.baseDimSize = CAUSAL_CONV1D_PACKED_V_LEGACY_DIM; + } + } + return task; +} + +__aicore__ inline FnDirectBlockTask ResolveFnPackedQkvBlockTask(int32_t blockIdx, int32_t tokenBlockCnt, + int32_t tokenBlockSize, int32_t cuSeqlen, + int32_t baseDimCnt, int32_t baseDim, int32_t dim) +{ + return ResolveFnPackedQkvBlockTask(blockIdx, tokenBlockCnt, tokenBlockSize, cuSeqlen, baseDimCnt, baseDim, dim, + CAUSAL_CONV1D_PACKED_QK_LEGACY_DIM / 2, + CAUSAL_CONV1D_PACKED_QK_LEGACY_DIM / 2, + CAUSAL_CONV1D_PACKED_V_LEGACY_DIM); +} + __aicore__ inline bool IsFnInitStateSnapshotOwnerBlock(const FnDirectBlockTask &task) { return task.valid && task.tokenTileId == 0; @@ -223,8 +257,14 @@ __aicore__ inline void CAUSAL_CONV1D_CLASS::ProcessVarlenTokenTiled() const bool isVarlenMode = (tilingData_->inputMode == 0); const int32_t blockIdx = static_cast(GetBlockIdx()); - const auto blockTask = ResolveFnDirectBlockTask(blockIdx, tokenBlockCnt, tokenBlockSize, cuSeqlen, baseDimCnt, - baseDim, dim); + const auto blockTask = IsPackedQkvOutput() + ? ResolveFnPackedQkvBlockTask(blockIdx, tokenBlockCnt, tokenBlockSize, cuSeqlen, + baseDimCnt, baseDim, dim, + static_cast(tilingData_->packedQDim), + static_cast(tilingData_->packedKDim), + static_cast(tilingData_->packedVDim)) + : ResolveFnDirectBlockTask(blockIdx, tokenBlockCnt, tokenBlockSize, cuSeqlen, + baseDimCnt, baseDim, dim); if (tilingData_->hasInitStateWorkspace != 0) { if (IsFnInitStateSnapshotOwnerBlock(blockTask)) { PrefetchInitStatesToWorkspace(blockTask.channelStart, blockTask.baseDimSize); diff --git a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_tiling_data.h b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_tiling_data.h index 5f88d26..18301a5 100644 --- a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_tiling_data.h +++ b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_tiling_data.h @@ -25,6 +25,16 @@ enum FnExecutionPlan : int64_t { FN_EXECUTION_PLAN_CUTBSD = 2, }; +enum CausalConv1dActivationMode : int64_t { + CAUSAL_CONV1D_ACTIVATION_NONE = 0, + CAUSAL_CONV1D_ACTIVATION_SILU = 1, + CAUSAL_CONV1D_ACTIVATION_SILU_PACKED_QKV = 2, +}; + +inline constexpr int64_t CAUSAL_CONV1D_PACKED_QKV_LEGACY_DIM = 5120; +inline constexpr int64_t CAUSAL_CONV1D_PACKED_QK_LEGACY_DIM = 2048; +inline constexpr int64_t CAUSAL_CONV1D_PACKED_V_LEGACY_DIM = 3072; + inline constexpr int64_t ResolveFnExecutionPlan(int64_t baseDimCnt) { if (baseDimCnt <= 0) { @@ -54,6 +64,11 @@ struct CausalConv1dTilingData { int64_t padSlotId; int64_t hasBias; + int64_t packedQDim; + int64_t packedKDim; + int64_t packedVDim; + int64_t packedHeadDim; + int64_t baseDim; int64_t baseDimCnt; diff --git a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_tiling_key.h b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_tiling_key.h index c4ed711..8fe7a1a 100644 --- a/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_tiling_key.h +++ b/xllm_ops/causal_conv1d/op_kernel/causal_conv1d_tiling_key.h @@ -18,7 +18,10 @@ #define __CAUSAL_CONV1D_TILING_KEY_H__ #include "causal_conv1d_tiling_data.h" + +#ifndef CAUSAL_CONV1D_SKIP_TPL_REGISTRATION #include "ascendc/host_api/tiling/template_argument.h" +#endif #define CAUSAL_CONV1D_TPL_RUN_MODE_FN 0 #define CAUSAL_CONV1D_TPL_RUN_MODE_UPDATE 1 @@ -29,6 +32,8 @@ #define CAUSAL_CONV1D_TPL_FN_PLAN_INVALID 0 #define CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS 1 #define CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD 2 + +#ifndef CAUSAL_CONV1D_SKIP_TPL_REGISTRATION ASCENDC_TPL_ARGS_DECL(CausalConv1d, ASCENDC_TPL_UINT_DECL(runModeKey, 1, ASCENDC_TPL_UI_LIST, CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_RUN_MODE_UPDATE), @@ -61,5 +66,6 @@ ASCENDC_TPL_SEL( CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD)); #undef CAUSAL_CONV1D_TPL_SEL_ENTRY +#endif #endif // __CAUSAL_CONV1D_TILING_KEY_H__ diff --git a/xllm_ops/causal_conv1d_qkv/CMakeLists.txt b/xllm_ops/causal_conv1d_qkv/CMakeLists.txt new file mode 100644 index 0000000..32c5c3d --- /dev/null +++ b/xllm_ops/causal_conv1d_qkv/CMakeLists.txt @@ -0,0 +1,6 @@ +file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*) +foreach(SUB_DIR ${CURRENT_DIRS}) + if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt") + add_subdirectory(${SUB_DIR}) + endif() +endforeach() diff --git a/xllm_ops/causal_conv1d_qkv/op_host/CMakeLists.txt b/xllm_ops/causal_conv1d_qkv/op_host/CMakeLists.txt new file mode 100644 index 0000000..9beb073 --- /dev/null +++ b/xllm_ops/causal_conv1d_qkv/op_host/CMakeLists.txt @@ -0,0 +1,19 @@ +add_op_to_compiled_list() +set(causal_conv1d_qkv_depends causal_conv1d PARENT_SCOPE) + +if (BUILD_OPEN_PROJECT) + target_sources(op_host_aclnn PRIVATE + causal_conv1d_qkv_def.cpp + ) +endif() + +add_ops_compile_options( + OP_NAME CausalConv1dQkv + OPTIONS --cce-auto-sync=on + -Wno-deprecated-declarations + -Werror +) + +if (NOT BUILD_OPS_RTY_KERNEL) + add_modules_sources(OPTYPE causal_conv1d_qkv ACLNNTYPE aclnn) +endif() diff --git a/xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_def.cpp b/xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_def.cpp new file mode 100644 index 0000000..c0de108 --- /dev/null +++ b/xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_def.cpp @@ -0,0 +1,82 @@ +#include "register/op_def_registry.h" +#include "../../causal_conv1d/op_kernel/causal_conv1d_tiling_data.h" + +namespace ops { + +class CausalConv1dQkv : public OpDef { +public: + explicit CausalConv1dQkv(const char* name) : OpDef(name) + { + this->Input("x") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("weight") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("bias") + .ParamType(OPTIONAL) + .DataType({ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("convStates") + .ParamType(REQUIRED) + .DataType({ge::DT_BF16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + this->Input("queryStartLoc") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT64}) + .FormatList({ge::FORMAT_ND}) + .ValueDepend(OPTIONAL) + .AutoContiguous(); + this->Input("cacheIndices") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT64}) + .FormatList({ge::FORMAT_ND}) + .ValueDepend(OPTIONAL) + .AutoContiguous(); + this->Input("initialStateMode") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT64}) + .FormatList({ge::FORMAT_ND}) + .ValueDepend(OPTIONAL) + .AutoContiguous(); + this->Input("numAcceptedTokens") + .ParamType(OPTIONAL) + .DataTypeList({ge::DT_INT64}) + .FormatList({ge::FORMAT_ND}) + .ValueDepend(OPTIONAL) + .AutoContiguous(); + this->Output("y") + .ParamType(REQUIRED) + .DataType({ge::DT_FLOAT16, ge::DT_BF16}) + .FormatList({ge::FORMAT_ND}) + .AutoContiguous(); + + this->Attr("activationMode").AttrType(OPTIONAL).Int(CAUSAL_CONV1D_ACTIVATION_SILU_PACKED_QKV); + this->Attr("padSlotId").AttrType(OPTIONAL).Int(-1); + this->Attr("runMode").AttrType(OPTIONAL).Int(0); + this->Attr("qDim").AttrType(OPTIONAL).Int(1024); + this->Attr("kDim").AttrType(OPTIONAL).Int(1024); + this->Attr("vDim").AttrType(OPTIONAL).Int(3072); + this->Attr("headDim").AttrType(OPTIONAL).Int(128); + + OpAICoreConfig aicoreConfig; + aicoreConfig.DynamicCompileStaticFlag(true) + .DynamicFormatFlag(false) + .DynamicRankSupportFlag(true) + .DynamicShapeSupportFlag(true) + .NeedCheckSupportFlag(false) + .PrecisionReduceFlag(true) + .ExtendCfgInfo("coreType.value", "AiCore"); + this->AICore().AddConfig("ascend910b", aicoreConfig); + this->AICore().AddConfig("ascend910_93", aicoreConfig); + } +}; +OP_ADD(CausalConv1dQkv); + +} // namespace ops diff --git a/xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_proto.cpp b/xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_proto.cpp new file mode 100644 index 0000000..005e9d4 --- /dev/null +++ b/xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_proto.cpp @@ -0,0 +1,18 @@ +#include "register/op_impl_registry.h" +#include "tiling_base/error_log.h" + +namespace ops { + +static ge::graphStatus InferShapeCausalConv1dQkv(gert::InferShapeContext* context) +{ + const gert::Shape* xShape = context->GetInputShape(0); + OP_CHECK_NULL_WITH_CONTEXT(context, xShape); + gert::Shape* yShape = context->GetOutputShape(0); + OP_CHECK_NULL_WITH_CONTEXT(context, yShape); + *yShape = *xShape; + return ge::GRAPH_SUCCESS; +} + +IMPL_OP_INFERSHAPE(CausalConv1dQkv).InferShape(InferShapeCausalConv1dQkv); + +} // namespace ops diff --git a/xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_tiling.cpp b/xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_tiling.cpp new file mode 100644 index 0000000..fdade4c --- /dev/null +++ b/xllm_ops/causal_conv1d_qkv/op_host/causal_conv1d_qkv_tiling.cpp @@ -0,0 +1,10 @@ +#include "../../causal_conv1d/op_host/causal_conv1d_tiling.h" +#include "../../causal_conv1d/op_host/causal_conv1d_tiling_utils.h" + +namespace optiling { + +IMPL_OP_OPTILING(CausalConv1dQkv) + .Tiling(CausalConv1dTilingFunc) + .TilingParse(TilingParseForCausalConv1d); + +} // namespace optiling diff --git a/xllm_ops/causal_conv1d_qkv/op_kernel/causal_conv1d_qkv.cpp b/xllm_ops/causal_conv1d_qkv/op_kernel/causal_conv1d_qkv.cpp new file mode 100644 index 0000000..a68321e --- /dev/null +++ b/xllm_ops/causal_conv1d_qkv/op_kernel/causal_conv1d_qkv.cpp @@ -0,0 +1,41 @@ +#include "../../causal_conv1d/op_kernel/causal_conv1d_fn.h" +#include "../../causal_conv1d/op_kernel/causal_conv1d_update.h" + +namespace { + +template +__aicore__ inline void RunCausalConv1dQkv(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates, + GM_ADDR queryStartLoc, GM_ADDR cacheIndices, GM_ADDR initialStateMode, + GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace, + const CausalConv1dTilingData *tilingData) +{ + if constexpr (runModeKey == CAUSAL_CONV1D_TPL_RUN_MODE_FN) { + NsCausalConv1d::RunCausalConv1dFn( + x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y, + workspace, tilingData); + } else { + NsCausalConv1d::RunCausalConv1dUpdate( + x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y, + workspace, tilingData); + } +} + +} // namespace + +template +__global__ __aicore__ void causal_conv1d_qkv(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates, + GM_ADDR queryStartLoc, GM_ADDR cacheIndices, GM_ADDR initialStateMode, + GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling) +{ + REGISTER_TILING_DEFAULT(CausalConv1dTilingData); + GET_TILING_DATA(tilingData, tiling); + KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIV_1_0); + GM_ADDR userWorkspace = workspace; + if (workspace != nullptr) { + userWorkspace = AscendC::GetUserWorkspace(workspace); + } + + RunCausalConv1dQkv( + x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y, + userWorkspace, &tilingData); +} diff --git a/xllm_ops/mega_chunk_gdn/op_host/mega_chunk_gdn_def.cpp b/xllm_ops/mega_chunk_gdn/op_host/mega_chunk_gdn_def.cpp index 45e5c76..aa7226e 100644 --- a/xllm_ops/mega_chunk_gdn/op_host/mega_chunk_gdn_def.cpp +++ b/xllm_ops/mega_chunk_gdn/op_host/mega_chunk_gdn_def.cpp @@ -15,34 +15,41 @@ limitations under the License. #include "register/op_def_registry.h" +#include + namespace ops { class MegaChunkGdn : public OpDef { public: explicit MegaChunkGdn(const char *name) : OpDef(name) { - this->Input("q").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}).AutoContiguous(); - this->Input("k").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}).AutoContiguous(); - this->Input("v").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}).AutoContiguous(); - this->Input("g").ParamType(REQUIRED).DataType({ge::DT_FLOAT}).Format({ge::FORMAT_ND}).AutoContiguous(); - this->Input("beta").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}).AutoContiguous(); - this->Input("mask_lower").ParamType(REQUIRED).DataType({ge::DT_FLOAT}).Format({ge::FORMAT_ND}); - this->Input("mask_full").ParamType(REQUIRED).DataType({ge::DT_FLOAT}).Format({ge::FORMAT_ND}); - this->Input("minus_identity").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}); - this->Input("cu_seqlens").ParamType(REQUIRED).DataType({ge::DT_INT32}).Format({ge::FORMAT_ND}).AutoContiguous(); - this->Input("initial_state").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}).AutoContiguous(); - - this->Output("out").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}); - this->Output("g_sum").ParamType(REQUIRED).DataType({ge::DT_FLOAT}).Format({ge::FORMAT_ND}); - this->Output("g_t").ParamType(REQUIRED).DataType({ge::DT_FLOAT}).Format({ge::FORMAT_ND}); - this->Output("beta_t").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}); - this->Output("a").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}); - this->Output("a_inv_f32").ParamType(REQUIRED).DataType({ge::DT_FLOAT}).Format({ge::FORMAT_ND}); - this->Output("a_inv").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}); - this->Output("w").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}); - this->Output("u").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}); - this->Output("h").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}); - this->Output("v_new").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}); - this->Output("final_state").ParamType(REQUIRED).DataType({ge::DT_FLOAT16}).Format({ge::FORMAT_ND}); + const std::vector computeTypes = {ge::DT_FLOAT16, ge::DT_BF16}; + const std::vector floatTypes = {ge::DT_FLOAT, ge::DT_FLOAT}; + const std::vector int32Types = {ge::DT_INT32, ge::DT_INT32}; + const std::vector formats = {ge::FORMAT_ND, ge::FORMAT_ND}; + + this->Input("q").ParamType(REQUIRED).DataType(computeTypes).Format(formats).AutoContiguous(); + this->Input("k").ParamType(REQUIRED).DataType(computeTypes).Format(formats).AutoContiguous(); + this->Input("v").ParamType(REQUIRED).DataType(computeTypes).Format(formats).AutoContiguous(); + this->Input("g").ParamType(REQUIRED).DataType(floatTypes).Format(formats).AutoContiguous(); + this->Input("beta").ParamType(REQUIRED).DataType(computeTypes).Format(formats).AutoContiguous(); + this->Input("mask_lower").ParamType(REQUIRED).DataType(floatTypes).Format(formats); + this->Input("mask_full").ParamType(REQUIRED).DataType(floatTypes).Format(formats); + this->Input("minus_identity").ParamType(REQUIRED).DataType(computeTypes).Format(formats); + this->Input("cu_seqlens").ParamType(REQUIRED).DataType(int32Types).Format(formats).AutoContiguous(); + this->Input("initial_state").ParamType(REQUIRED).DataType(computeTypes).Format(formats).AutoContiguous(); + + this->Output("out").ParamType(REQUIRED).DataType(computeTypes).Format(formats); + this->Output("g_sum").ParamType(REQUIRED).DataType(floatTypes).Format(formats); + this->Output("g_t").ParamType(REQUIRED).DataType(floatTypes).Format(formats); + this->Output("beta_t").ParamType(REQUIRED).DataType(computeTypes).Format(formats); + this->Output("a").ParamType(REQUIRED).DataType(computeTypes).Format(formats); + this->Output("a_inv_f32").ParamType(REQUIRED).DataType(floatTypes).Format(formats); + this->Output("a_inv").ParamType(REQUIRED).DataType(computeTypes).Format(formats); + this->Output("w").ParamType(REQUIRED).DataType(computeTypes).Format(formats); + this->Output("u").ParamType(REQUIRED).DataType(computeTypes).Format(formats); + this->Output("h").ParamType(REQUIRED).DataType(computeTypes).Format(formats); + this->Output("v_new").ParamType(REQUIRED).DataType(computeTypes).Format(formats); + this->Output("final_state").ParamType(REQUIRED).DataType(computeTypes).Format(formats); this->Attr("num_matrices").AttrType(OPTIONAL).Int(0); this->Attr("has_initial_state").AttrType(OPTIONAL).Bool(false); diff --git a/xllm_ops/mega_chunk_gdn/op_kernel/chunk_h.cpp b/xllm_ops/mega_chunk_gdn/op_kernel/chunk_h.cpp index 430c6b8..70806b7 100644 --- a/xllm_ops/mega_chunk_gdn/op_kernel/chunk_h.cpp +++ b/xllm_ops/mega_chunk_gdn/op_kernel/chunk_h.cpp @@ -29,14 +29,14 @@ // K/S ready. // // Inputs: -// K [total_tokens, Hg, D] half — keys (BSND layout; GQA/MQA group heads) -// W [total_tokens, H, D] half — wy_fast output (BSND layout) -// U [total_tokens, H, D] half — values pre-residual (BSND layout) +// K [total_tokens, Hg, D] DTYPE_Q — keys (BSND layout; GQA/MQA group heads) +// W [total_tokens, H, D] DTYPE_Q — wy_fast output (BSND layout) +// U [total_tokens, H, D] DTYPE_Q — values pre-residual (BSND layout) // G [H, total_tokens] float — pre-transposed cumulative gates -// S [total_chunks, H, D, D] half — per-chunk state snapshots (output) -// V [total_tokens, H, D] half — residual-corrected values (output) -// FS [batch, H, D, D] half — final state per sequence (output) -// H0 [batch, H, D, D] half — optional initial state per sequence +// S [total_chunks, H, D, D] DTYPE_Q — per-chunk state snapshots (output) +// V [total_tokens, H, D] DTYPE_Q — residual-corrected values (output) +// FS [batch, H, D, D] DTYPE_Q — final state per sequence (output) +// H0 [batch, H, D, D] DTYPE_Q — optional initial state per sequence // workspace [per-core scratch] — Cube↔Vec communication buffer // // NPU memory hierarchy: @@ -53,7 +53,7 @@ // TLOAD(dst, gm) — dst = gm_data (DMA: GM→L1 or GM→UB) // TSTORE(gm, src) — gm_data = src (DMA: UB/L0C→GM) // TASSIGN(tile, addr) — tile = memory[addr] (bind tile to buffer address) -// TCVT(dst, src, mode) — dst = src.float()/.half() +// TCVT(dst, src, mode) — converts between float and DTYPE_Q // TMOV(dst, src) — dst = src.clone() // TADD(d, a, b) — d = a + b // TSUB(d, a, b) — d = a - b @@ -66,7 +66,7 @@ // TFILLPAD(dst, src) — zero-fill L1 tile padding (for tail chunks) // TEXTRACT(l0, l1, r, c) — L1 sub-tile → L0A/L0B // TRESHAPE(zn, nz) — reinterpret layout NZ↔ZN (logical transpose, free) -// TMATMUL(C, A, B) — C = A @ B (Cube GEMM, fp16 inputs → fp32 accum) +// TMATMUL(C, A, B) — C = A @ B (Cube GEMM, DTYPE_Q inputs → FP32 accum) // set_flag/wait_flag — pipe sync within same core // ffts_cross_core_sync — cross-core signal Cube↔Vec // wait_flag_dev(flag) — wait for cross-core signal @@ -296,13 +296,13 @@ gemm_v0(std::conditional_t, template AICORE void chunk_h_kernel( - __gm__ half *K_handle, __gm__ half *W_handle, __gm__ half *U_handle, + __gm__ DTYPE_Q *K_handle, __gm__ DTYPE_Q *W_handle, __gm__ DTYPE_Q *U_handle, __gm__ float *G_handle, - __gm__ half *S_handle, __gm__ half *V_handle, __gm__ half *FS_handle, - __gm__ half *H0_handle, + __gm__ DTYPE_Q *S_handle, __gm__ DTYPE_Q *V_handle, __gm__ DTYPE_Q *FS_handle, + __gm__ DTYPE_Q *H0_handle, int64_t has_initial_state, int64_t output_final_state, - __gm__ half *workspace_handle, + __gm__ DTYPE_Q *workspace_handle, __gm__ int32_t *cu_seqlens, int64_t batch_size, int64_t seq_len, int64_t total_tokens, uint32_t num_heads, @@ -351,16 +351,16 @@ AICORE void chunk_h_kernel( constexpr int32_t WS_KV = DD * 3; constexpr int32_t WS_PER_CORE = DD * 4; - TileMatL1 s_l1; + TileMatL1 s_l1; TASSIGN(s_l1, 0); - TileMatL1 w_l1; - TASSIGN(w_l1, D * D * sizeof(half)); + TileMatL1 w_l1; + TASSIGN(w_l1, D * D * sizeof(DTYPE_Q)); TileAcc ws_l0; TASSIGN(ws_l0, 0); - TileMatL1 k_l1; - TASSIGN(k_l1, (DD + C * D) * sizeof(half)); - TileMatL1 v_l1; - TASSIGN(v_l1, (DD + C * D + D * C) * sizeof(half)); + TileMatL1 k_l1; + TASSIGN(k_l1, (DD + C * D) * sizeof(DTYPE_Q)); + TileMatL1 v_l1; + TASSIGN(v_l1, (DD + C * D + D * C) * sizeof(DTYPE_Q)); TileAcc kv_l0; TASSIGN(kv_l0, C * D * sizeof(float)); @@ -371,9 +371,9 @@ AICORE void chunk_h_kernel( ChunkSize * 16 * static_cast(sizeof(float)); constexpr int32_t S_UB = ZERO_UB + 64 * sizeof(float); constexpr int32_t K_UB_HALF = S_UB + HalfC * D * sizeof(float); - constexpr int32_t G_UB = K_UB_HALF + HalfC * D * sizeof(half); + constexpr int32_t G_UB = K_UB_HALF + HalfC * D * sizeof(DTYPE_Q); constexpr int32_t U_UB_HALF = G_UB + C * sizeof(float); - constexpr int32_t K_UB = U_UB_HALF + HalfC * D * sizeof(half); + constexpr int32_t K_UB = U_UB_HALF + HalfC * D * sizeof(DTYPE_Q); constexpr int32_t G_V_UB = K_UB + HalfC * D * sizeof(float); constexpr int32_t COEFF_UB = G_V_UB + 64 * sizeof(float); constexpr int32_t U_UB = COEFF_UB + 64 * sizeof(float); @@ -385,13 +385,13 @@ AICORE void chunk_h_kernel( TASSIGN(zero_ub, ZERO_UB); TileUbDataND s_ub; TASSIGN(s_ub, S_UB); - TileUbDataND k_ub_half; + TileUbDataND k_ub_half; TASSIGN(k_ub_half, K_UB_HALF); TileUbDataND g_ub; TASSIGN(g_ub, G_UB); - TileUbDataND s_ub_half; + TileUbDataND s_ub_half; TASSIGN(s_ub_half, S_UB_HALF); - TileUbDataND u_ub_half; + TileUbDataND u_ub_half; TASSIGN(u_ub_half, U_UB_HALF); TileUbDataND k_ub; TASSIGN(k_ub, K_UB); @@ -453,9 +453,9 @@ AICORE void chunk_h_kernel( { GmShape2D s_shape(D, D); GmStride2D s_stride(D); - GmTensor2D s_global(workspace_handle + ws_base + WS_S, s_shape, + GmTensor2D s_global(workspace_handle + ws_base + WS_S, s_shape, s_stride); - DynMatL1 s_l1_load(D, D); + DynMatL1 s_l1_load(D, D); TASSIGN(s_l1_load, 0); // Load the previous recurrent state S_i from per-core workspace. TLOAD(s_l1_load, s_global); @@ -465,9 +465,9 @@ AICORE void chunk_h_kernel( { GmShape2D w_shape(static_cast(valid), D); GmStride2D w_stride(BSND_QKV_STRIDE); - GmTensor2D w_global(W_handle + w_offset, w_shape, w_stride); - DynMatL1 w_l1_load(static_cast(valid), D); - TASSIGN(w_l1_load, D * D * static_cast(sizeof(half))); + GmTensor2D w_global(W_handle + w_offset, w_shape, w_stride); + DynMatL1 w_l1_load(static_cast(valid), D); + TASSIGN(w_l1_load, D * D * static_cast(sizeof(DTYPE_Q))); TLOAD(w_l1_load, w_global); if (valid != C) { TFILLPAD(w_l1_load, w_l1_load); @@ -477,13 +477,13 @@ AICORE void chunk_h_kernel( set_flag(PIPE_FIX, PIPE_M, EVENT_ID0); wait_flag(PIPE_FIX, PIPE_M, EVENT_ID0); // Apply the carried recurrent state to every token in this chunk. - gemm_v0( + gemm_v0( w_l1, s_l1, ws_l0, (bool)1); { GmShape2D ws_shape(C, D); GmStride2D ws_stride(D); - GmTensor2D ws_global(workspace_handle + ws_base + WS_WS, + GmTensor2D ws_global(workspace_handle + ws_base + WS_WS, ws_shape, ws_stride); DynAccTile ws_store(C, D); TASSIGN(ws_store, 0); @@ -497,10 +497,10 @@ AICORE void chunk_h_kernel( { GmShape2D k_shape(D, C); GmStride2D k_stride(C); - GmTensor2D k_global(workspace_handle + ws_base + WS_K, k_shape, + GmTensor2D k_global(workspace_handle + ws_base + WS_K, k_shape, k_stride); - DynMatL1 k_l1_load(D, C); - TASSIGN(k_l1_load, (DD + C * D) * static_cast(sizeof(half))); + DynMatL1 k_l1_load(D, C); + TASSIGN(k_l1_load, (DD + C * D) * static_cast(sizeof(DTYPE_Q))); TLOAD(k_l1_load, k_global); } @@ -508,10 +508,10 @@ AICORE void chunk_h_kernel( { GmShape2D v_shape(static_cast(valid), D); GmStride2D v_stride(BSND_QKV_STRIDE); - GmTensor2D v_global(V_handle + v_offset, v_shape, v_stride); - DynMatL1 v_l1_load(static_cast(valid), D); + GmTensor2D v_global(V_handle + v_offset, v_shape, v_stride); + DynMatL1 v_l1_load(static_cast(valid), D); TASSIGN(v_l1_load, - (DD + C * D + D * C) * static_cast(sizeof(half))); + (DD + C * D + D * C) * static_cast(sizeof(DTYPE_Q))); TLOAD(v_l1_load, v_global); if (valid != C) { TFILLPAD(v_l1_load, v_l1_load); @@ -521,13 +521,13 @@ AICORE void chunk_h_kernel( set_flag(PIPE_FIX, PIPE_M, EVENT_ID0); wait_flag(PIPE_FIX, PIPE_M, EVENT_ID0); // This chunk contributes the additive update K_i^T V_i to the state recurrence. - gemm_v0( + gemm_v0( k_l1, v_l1, kv_l0, (bool)1); { GmShape2D kv_shape(D, D); GmStride2D kv_stride(D); - GmTensor2D kv_global(workspace_handle + ws_base + WS_KV, + GmTensor2D kv_global(workspace_handle + ws_base + WS_KV, kv_shape, kv_stride); DynAccTile kv_store(D, D); TASSIGN(kv_store, C * D * static_cast(sizeof(float))); @@ -579,8 +579,8 @@ AICORE void chunk_h_kernel( int64_t h0_offset = (seq_idx * H + head) * DD + vid * HalfC * D; GmShape2D h0_shape(HalfC, D); GmStride2D h0_stride(D); - GmTensor2D h0_global(H0_handle + h0_offset, h0_shape, h0_stride); - DynVecTile h0_load(HalfC, D); + GmTensor2D h0_global(H0_handle + h0_offset, h0_shape, h0_stride); + DynVecTile h0_load(HalfC, D); TASSIGN(h0_load, S_UB_HALF); TLOAD(h0_load, h0_global); set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); @@ -594,13 +594,13 @@ AICORE void chunk_h_kernel( set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); { - // `workspace_handle` is a `half*`, so all offsets here are in half elements. + // `workspace_handle` is a `DTYPE_Q*`, so all offsets here are in DTYPE_Q elements. GmShape2D s_shape(HalfC, D); GmStride2D s_stride(D); - GmTensor2D s_global( + GmTensor2D s_global( workspace_handle + ws_base + WS_S + vid * HalfC * D, s_shape, s_stride); - DynVecTile s_store(HalfC, D); + DynVecTile s_store(HalfC, D); TASSIGN(s_store, S_UB_HALF); TSTORE(s_global, s_store); } @@ -608,10 +608,10 @@ AICORE void chunk_h_kernel( int64_t s_out_offset = (chunk_offset * H + head) * DD; GmShape2D s_out_shape(HalfC, D); GmStride2D s_out_stride(D); - GmTensor2D s_out_global( + GmTensor2D s_out_global( S_handle + s_out_offset + vid * HalfC * D, s_out_shape, s_out_stride); - DynVecTile s_out_store(HalfC, D); + DynVecTile s_out_store(HalfC, D); TASSIGN(s_out_store, S_UB_HALF); TSTORE(s_out_global, s_out_store); } @@ -633,8 +633,8 @@ AICORE void chunk_h_kernel( if (valid_rows_0 > 0) { GmShape2D k_shape(valid_rows_0, D); GmStride2D k_stride(BSND_K_STRIDE); - GmTensor2D k_global(K_handle + k_offset_0, k_shape, k_stride); - DynVecTile k_load(valid_rows_0, D); + GmTensor2D k_global(K_handle + k_offset_0, k_shape, k_stride); + DynVecTile k_load(valid_rows_0, D); TASSIGN(k_load, K_UB_HALF); TLOAD(k_load, k_global); if (valid_rows_0 != HalfC) { @@ -682,8 +682,8 @@ AICORE void chunk_h_kernel( if (valid_rows > 0) { GmShape2D u_shape(valid_rows, D); GmStride2D u_stride(BSND_QKV_STRIDE); - GmTensor2D u_global(U_handle + u_offset, u_shape, u_stride); - DynVecTile u_load(valid_rows, D); + GmTensor2D u_global(U_handle + u_offset, u_shape, u_stride); + DynVecTile u_load(valid_rows, D); TASSIGN(u_load, U_UB_HALF); TLOAD(u_load, u_global); if (valid_rows != HalfC) { @@ -737,10 +737,10 @@ AICORE void chunk_h_kernel( { GmShape2D ws_shape(HalfC, D); GmStride2D ws_stride(D); - GmTensor2D ws_global( + GmTensor2D ws_global( workspace_handle + ws_base + WS_WS + vid * HalfC * D, ws_shape, ws_stride); - DynVecTile ws_load(HalfC, D); + DynVecTile ws_load(HalfC, D); TASSIGN(ws_load, U_UB_HALF); TLOAD(ws_load, ws_global); } @@ -762,8 +762,8 @@ AICORE void chunk_h_kernel( if (valid_rows > 0) { GmShape2D v_shape(valid_rows, D); GmStride2D v_stride(BSND_QKV_STRIDE); - GmTensor2D v_global(V_handle + v_offset, v_shape, v_stride); - DynVecTile v_store(valid_rows, D); + GmTensor2D v_global(V_handle + v_offset, v_shape, v_stride); + DynVecTile v_store(valid_rows, D); TASSIGN(v_store, U_UB_HALF); TSTORE(v_global, v_store); } @@ -773,10 +773,10 @@ AICORE void chunk_h_kernel( { GmShape2D k_shape(HalfC, D); GmStride2D k_stride(D); - GmTensor2D k_global( + GmTensor2D k_global( workspace_handle + ws_base + WS_K + vid * HalfC * D, k_shape, k_stride); - DynVecTile k_store(HalfC, D); + DynVecTile k_store(HalfC, D); TASSIGN(k_store, K_UB_HALF); TSTORE(k_global, k_store); } @@ -805,8 +805,8 @@ AICORE void chunk_h_kernel( if (next_valid_rows > 0) { GmShape2D k_shape(next_valid_rows, D); GmStride2D k_stride(BSND_K_STRIDE); - GmTensor2D k_global(K_handle + nk_off, k_shape, k_stride); - DynVecTile k_load( + GmTensor2D k_global(K_handle + nk_off, k_shape, k_stride); + DynVecTile k_load( next_valid_rows, D); TASSIGN(k_load, K_UB_HALF); TLOAD(k_load, k_global); @@ -840,10 +840,10 @@ AICORE void chunk_h_kernel( { GmShape2D kv_shape(HalfC, D); GmStride2D kv_stride(D); - GmTensor2D kv_global( + GmTensor2D kv_global( workspace_handle + ws_base + WS_KV + vid * HalfC * D, kv_shape, kv_stride); - DynVecTile kv_load(HalfC, D); + DynVecTile kv_load(HalfC, D); TASSIGN(kv_load, S_UB_HALF); TLOAD(kv_load, kv_global); } @@ -864,10 +864,10 @@ AICORE void chunk_h_kernel( { GmShape2D s_shape(HalfC, D); GmStride2D s_stride(D); - GmTensor2D s_global( + GmTensor2D s_global( workspace_handle + ws_base + WS_S + vid * HalfC * D, s_shape, s_stride); - DynVecTile s_store(HalfC, D); + DynVecTile s_store(HalfC, D); TASSIGN(s_store, S_UB_HALF); TSTORE(s_global, s_store); } @@ -879,10 +879,10 @@ AICORE void chunk_h_kernel( { GmShape2D s_out_shape(HalfC, D); GmStride2D s_out_stride(D); - GmTensor2D s_out_global( + GmTensor2D s_out_global( S_handle + s_out_offset + vid * HalfC * D, s_out_shape, s_out_stride); - DynVecTile s_out_store(HalfC, D); + DynVecTile s_out_store(HalfC, D); TASSIGN(s_out_store, S_UB_HALF); TSTORE(s_out_global, s_out_store); } @@ -902,9 +902,9 @@ AICORE void chunk_h_kernel( { GmShape2D fs_shape(HalfC, D); GmStride2D fs_stride(D); - GmTensor2D fs_global(FS_handle + fs_offset + vid * HalfC * D, + GmTensor2D fs_global(FS_handle + fs_offset + vid * HalfC * D, fs_shape, fs_stride); - DynVecTile fs_store(HalfC, D); + DynVecTile fs_store(HalfC, D); TASSIGN(fs_store, S_UB_HALF); TSTORE(fs_global, fs_store); } diff --git a/xllm_ops/mega_chunk_gdn/op_kernel/chunk_o.cpp b/xllm_ops/mega_chunk_gdn/op_kernel/chunk_o.cpp index e951b3b..d1bc875 100644 --- a/xllm_ops/mega_chunk_gdn/op_kernel/chunk_o.cpp +++ b/xllm_ops/mega_chunk_gdn/op_kernel/chunk_o.cpp @@ -61,7 +61,7 @@ // TLOAD(dst, gm) — dst = gm_data (DMA: GM→UB/L1, async) // TSTORE(gm, src) — gm = src (DMA: UB/L0C→GM, async) // TASSIGN(tile, addr) — bind tile descriptor to buffer address -// TCVT(dst, src, mode) — type cast: dst = src.float() or .half() +// TCVT(dst, src, mode) — converts between float and DTYPE_Q // TMOV(dst, src) — copy: dst = src.clone() // TADD(d, a, b) — d = a + b // TSUB(d, a, b) — d = a - b @@ -72,7 +72,7 @@ // TCOLEXPAND(2d, row) — 2d[i,j] = row[j] (broadcast row→columns) // TEXTRACT(l0, l1, r, c) — copy L1 sub-tile → L0A/L0B (Cube input regs) // TRESHAPE(zn, nz) — reinterpret L1 fractal layout (transpose, free) -// TMATMUL(C, A, B) — C = A @ B (Cube engine, fp16→fp32 accum) +// TMATMUL(C, A, B) — C = A @ B (Cube engine, DTYPE_Q→FP32 accum) // set_flag / wait_flag — synchronize pipes within same AI core // ffts_cross_core_sync — signal across Cube↔Vec cores // wait_flag_dev(flag) — wait for cross-core signal @@ -147,13 +147,13 @@ using GmTensor2D = pto::GlobalTensor; template AICORE void GDN_CHUNK_O_KERNEL( - __gm__ half *Q_handle, __gm__ half *K_handle, __gm__ half *V_handle, - __gm__ half *S_handle, __gm__ float *G_handle, + __gm__ DTYPE_Q *Q_handle, __gm__ DTYPE_Q *K_handle, __gm__ DTYPE_Q *V_handle, + __gm__ DTYPE_Q *S_handle, __gm__ float *G_handle, __gm__ float *Msk_handle, - __gm__ half *workspace_qk_handle, - __gm__ half *workspace_qs_qkv_handle, - __gm__ half *workspace_qk_gated_handle, - __gm__ half *O_handle, + __gm__ DTYPE_Q *workspace_qk_handle, + __gm__ DTYPE_Q *workspace_qs_qkv_handle, + __gm__ DTYPE_Q *workspace_qk_gated_handle, + __gm__ DTYPE_Q *O_handle, __gm__ int32_t *cu_seqlens, int64_t batch_size, int64_t seq_len, int64_t total_tokens, @@ -217,21 +217,21 @@ AICORE void GDN_CHUNK_O_KERNEL( // s_l1 at 65536: S [D×D] — accumulated state, used in GEMM 2 // qk_gated at 98304: QK_gated [C×C] — from Vec, used in GEMM 3 // v_l1 at 131072: V [C×D] — values, used in GEMM 3 - L1Mat q_l1; + L1Mat q_l1; TASSIGN(q_l1, 0); - L1Mat k_l1; + L1Mat k_l1; TASSIGN(k_l1, 32768); TileAcc qk_l0; TASSIGN(qk_l0, 0); - L1Mat s_l1; + L1Mat s_l1; TASSIGN(s_l1, 65536); TileAcc qs_l0; TASSIGN(qs_l0, 65536); - L1Mat qk_gated_l1; + L1Mat qk_gated_l1; TASSIGN(qk_gated_l1, 98304); - L1Mat v_l1; + L1Mat v_l1; TASSIGN(v_l1, 131072); TileAcc qkv_l0; @@ -247,13 +247,13 @@ AICORE void GDN_CHUNK_O_KERNEL( // fit simultaneously in the ~256KB UB without overlapping: // g_ub: gate values [1, C] float @ 0 // msk_ub: causal mask [C/2, C] float @ 512 (loaded once, reused) - // qk_ub: QK scores in float [C/2, C] @ 33280 (after cast from half) + // qk_ub: QK scores in float [C/2, C] @ 33280 (after cast from DTYPE_Q) // g_v_ub: this sub-block's gate slice [1, C/2] @ 66048 // coeff_ub: gating coefficients [C/2, C] float @ 66304 - // qk_ub_half: QK in half [C/2, C] @ 99072 - // qs_ub_half: QS in half [C/2, D] @ 115456 + // qk_ub_half: QK in DTYPE_Q [C/2, C] @ 99072 + // qs_ub_half: QS in DTYPE_Q [C/2, D] @ 115456 // qs_ub: QS in float [C/2, D] @ 131840 - // o_ub_half: output O in half [C/2, D] @ 164608 + // o_ub_half: output O in DTYPE_Q [C/2, D] @ 164608 // o_ub: output O in float [C/2, D] @ QKUbAddr (reuses qk_ub space) UbND g_ub; TASSIGN(g_ub, GUbAddr); @@ -265,13 +265,13 @@ AICORE void GDN_CHUNK_O_KERNEL( TASSIGN(g_v_ub, GvUbAddr); UbND coeff_ub; TASSIGN(coeff_ub, CoeffUbAddr); - UbND qk_ub_half; + UbND qk_ub_half; TASSIGN(qk_ub_half, QKHalfUbAddr); - UbND qs_ub_half; + UbND qs_ub_half; TASSIGN(qs_ub_half, QSHalfUbAddr); UbND qs_ub; TASSIGN(qs_ub, QSUbAddr); - UbND o_ub_half; + UbND o_ub_half; TASSIGN(o_ub_half, OHalfUbAddr); UbND o_ub; TASSIGN(o_ub, OUbAddr); @@ -343,21 +343,21 @@ AICORE void GDN_CHUNK_O_KERNEL( // TLOAD performs DMA (MTE2 pipe). TFILLPAD zero-pads tail rows so // downstream GEMMs see a clean C×D matrix. { - L1Mat _l1(valid_rows, HiddenSize); + L1Mat _l1(valid_rows, HiddenSize); TASSIGN(_l1, 0); GmShape2D _gs(valid_rows, HiddenSize); GmStride2D _stride(BSND_QK_STRIDE); - GmTensor2D _gm(Q_handle + qk_off, _gs, _stride); + GmTensor2D _gm(Q_handle + qk_off, _gs, _stride); TLOAD(_l1, _gm); if (valid_rows != ChunkSize) TFILLPAD(_l1, _l1); } // ── Load K [valid_rows × D] from GM → L1 ──────────────────────── { - L1Mat _l1(valid_rows, HiddenSize); + L1Mat _l1(valid_rows, HiddenSize); TASSIGN(_l1, 32768); GmShape2D _gs(valid_rows, HiddenSize); GmStride2D _stride(BSND_QK_STRIDE); - GmTensor2D _gm(K_handle + qk_off, _gs, _stride); + GmTensor2D _gm(K_handle + qk_off, _gs, _stride); TLOAD(_l1, _gm); if (valid_rows != ChunkSize) TFILLPAD(_l1, _l1); } @@ -380,14 +380,14 @@ AICORE void GDN_CHUNK_O_KERNEL( // transpose_B: TRESHAPE converts k_l1 from NZ → ZN fractal layout, // effectively transposing K before TEXTRACT loads it into L0B. { - TileLeft _l0a; - TileRight _l0b; + TileLeft _l0a; + TileRight _l0b; TASSIGN(_l0a, 0x0); TASSIGN(_l0b, 0x0); auto _we = EVENT_ID1; set_flag(PIPE_MTE2, PIPE_MTE1, _we); wait_flag(PIPE_MTE2, PIPE_MTE1, _we); set_flag(PIPE_M, PIPE_MTE1, _we); wait_flag(PIPE_M, PIPE_MTE1, _we); TEXTRACT(_l0a, q_l1, 0, 0); - L1MatZN _bzn; TRESHAPE(_bzn, k_l1); TEXTRACT(_l0b, _bzn, 0, 0); + L1MatZN _bzn; TRESHAPE(_bzn, k_l1); TEXTRACT(_l0b, _bzn, 0, 0); set_flag(PIPE_MTE1, PIPE_M, _we); wait_flag(PIPE_MTE1, PIPE_M, _we); TMATMUL(qk_l0, _l0a, _l0b); set_flag(PIPE_MTE1, PIPE_MTE2, _we); wait_flag(PIPE_MTE1, PIPE_MTE2, _we); @@ -396,18 +396,18 @@ AICORE void GDN_CHUNK_O_KERNEL( // ── Load S [D × D] from GM → L1 (accumulated hidden state) ───── { - L1Mat _l1(HiddenSize, HiddenSize); + L1Mat _l1(HiddenSize, HiddenSize); TASSIGN(_l1, 65536); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = HiddenSize; _gs.shape[4] = HiddenSize; - GlobalTensor> _gm(S_handle + s_offset, _gs); + GlobalTensor> _gm(S_handle + s_offset, _gs); TLOAD(_l1, _gm); } // ── GEMM 2: QS = Q @ S (query applied to accumulated state) ──── { - TileLeft _l0a; - TileRight _l0b; + TileLeft _l0a; + TileRight _l0b; TASSIGN(_l0a, 0x0); TASSIGN(_l0b, 0x0); auto _we = EVENT_ID1; set_flag(PIPE_MTE2, PIPE_MTE1, _we); wait_flag(PIPE_MTE2, PIPE_MTE1, _we); @@ -420,14 +420,14 @@ AICORE void GDN_CHUNK_O_KERNEL( set_flag(PIPE_M, PIPE_FIX, _we); wait_flag(PIPE_M, PIPE_FIX, _we); } - // ── Store QK [C × C] from L0C → GM workspace (fp32→fp16 cast) ─── + // ── Store QK [C × C] from L0C → GM workspace (FP32→DTYPE_Q cast) ─── // TSTORE on TileAcc triggers MTE3 DMA with implicit type conversion. { TileAcc _l0(ChunkSize, ChunkSize); TASSIGN(_l0, 0); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = ChunkSize; _gs.shape[4] = ChunkSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qk_handle + static_cast(cid) * WsQKSize, _gs); TSTORE(_gm, _l0); @@ -439,7 +439,7 @@ AICORE void GDN_CHUNK_O_KERNEL( TASSIGN(_l0, 65536); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = ChunkSize; _gs.shape[4] = HiddenSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qs_qkv_handle + static_cast(cid) * WsQSSize, _gs); TSTORE(_gm, _l0); @@ -471,31 +471,31 @@ AICORE void GDN_CHUNK_O_KERNEL( // ── Load QK_gated [C × C] from GM workspace → L1 ──────────────── { - L1Mat _l1(ChunkSize, ChunkSize); + L1Mat _l1(ChunkSize, ChunkSize); TASSIGN(_l1, 98304); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = ChunkSize; _gs.shape[4] = ChunkSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qk_gated_handle + static_cast(cid) * WsGatedSize, _gs); TLOAD(_l1, _gm); } // ── Load V [valid_rows × D] from GM → L1 ──────────────────────── { - L1Mat _l1(valid_rows, HiddenSize); + L1Mat _l1(valid_rows, HiddenSize); TASSIGN(_l1, 131072); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = valid_rows; _gs.shape[4] = HiddenSize; GmStride2D _stride(BSND_V_STRIDE); - GmTensor2D _gm(V_handle + v_off, _gs, _stride); + GmTensor2D _gm(V_handle + v_off, _gs, _stride); TLOAD(_l1, _gm); if (valid_rows != ChunkSize) TFILLPAD(_l1, _l1); } // ── GEMM 3: QKV = QK_gated @ V (gated attention → values) ────── { - TileLeft _l0a; - TileRight _l0b; + TileLeft _l0a; + TileRight _l0b; TASSIGN(_l0a, 0x0); TASSIGN(_l0b, 0x0); auto _we = EVENT_ID1; set_flag(PIPE_MTE2, PIPE_MTE1, _we); wait_flag(PIPE_MTE2, PIPE_MTE1, _we); @@ -522,7 +522,7 @@ AICORE void GDN_CHUNK_O_KERNEL( TASSIGN(_l0, 0); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = ChunkSize; _gs.shape[4] = HiddenSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qs_qkv_handle + static_cast(cid) * WsQSSize, _gs); TSTORE(_gm, _l0); @@ -574,35 +574,35 @@ AICORE void GDN_CHUNK_O_KERNEL( // Load Q { - L1Mat _l1(valid_rows, HiddenSize); + L1Mat _l1(valid_rows, HiddenSize); TASSIGN(_l1, 0); GmShape2D _gs(valid_rows, HiddenSize); GmStride2D _stride(BSND_QK_STRIDE); - GmTensor2D _gm(Q_handle + qk_off, _gs, _stride); + GmTensor2D _gm(Q_handle + qk_off, _gs, _stride); TLOAD(_l1, _gm); if (valid_rows != ChunkSize) TFILLPAD(_l1, _l1); } // Load K { - L1Mat _l1(valid_rows, HiddenSize); + L1Mat _l1(valid_rows, HiddenSize); TASSIGN(_l1, 32768); GmShape2D _gs(valid_rows, HiddenSize); GmStride2D _stride(BSND_QK_STRIDE); - GmTensor2D _gm(K_handle + qk_off, _gs, _stride); + GmTensor2D _gm(K_handle + qk_off, _gs, _stride); TLOAD(_l1, _gm); if (valid_rows != ChunkSize) TFILLPAD(_l1, _l1); } // GEMM 1: QK = Q @ K^T (transpose_B via TRESHAPE NZ→ZN) { - TileLeft _l0a; - TileRight _l0b; + TileLeft _l0a; + TileRight _l0b; TASSIGN(_l0a, 0x0); TASSIGN(_l0b, 0x0); auto _we = EVENT_ID1; set_flag(PIPE_MTE2, PIPE_MTE1, _we); wait_flag(PIPE_MTE2, PIPE_MTE1, _we); set_flag(PIPE_M, PIPE_MTE1, _we); wait_flag(PIPE_M, PIPE_MTE1, _we); TEXTRACT(_l0a, q_l1, 0, 0); - L1MatZN _bzn; TRESHAPE(_bzn, k_l1); TEXTRACT(_l0b, _bzn, 0, 0); + L1MatZN _bzn; TRESHAPE(_bzn, k_l1); TEXTRACT(_l0b, _bzn, 0, 0); set_flag(PIPE_MTE1, PIPE_M, _we); wait_flag(PIPE_MTE1, PIPE_M, _we); TMATMUL(qk_l0, _l0a, _l0b); set_flag(PIPE_MTE1, PIPE_MTE2, _we); wait_flag(PIPE_MTE1, PIPE_MTE2, _we); @@ -611,18 +611,18 @@ AICORE void GDN_CHUNK_O_KERNEL( // Load S { - L1Mat _l1(HiddenSize, HiddenSize); + L1Mat _l1(HiddenSize, HiddenSize); TASSIGN(_l1, 65536); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = HiddenSize; _gs.shape[4] = HiddenSize; - GlobalTensor> _gm(S_handle + s_offset, _gs); + GlobalTensor> _gm(S_handle + s_offset, _gs); TLOAD(_l1, _gm); } // GEMM 2: QS = Q @ S { - TileLeft _l0a; - TileRight _l0b; + TileLeft _l0a; + TileRight _l0b; TASSIGN(_l0a, 0x0); TASSIGN(_l0b, 0x0); auto _we = EVENT_ID1; set_flag(PIPE_MTE2, PIPE_MTE1, _we); wait_flag(PIPE_MTE2, PIPE_MTE1, _we); @@ -641,7 +641,7 @@ AICORE void GDN_CHUNK_O_KERNEL( TASSIGN(_l0, 0); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = ChunkSize; _gs.shape[4] = ChunkSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qk_handle + static_cast(cid) * WsQKSize, _gs); TSTORE(_gm, _l0); @@ -653,7 +653,7 @@ AICORE void GDN_CHUNK_O_KERNEL( TASSIGN(_l0, 65536); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = ChunkSize; _gs.shape[4] = HiddenSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qs_qkv_handle + static_cast(cid) * WsQSSize, _gs); TSTORE(_gm, _l0); @@ -670,31 +670,31 @@ AICORE void GDN_CHUNK_O_KERNEL( // Load QK_gated { - L1Mat _l1(ChunkSize, ChunkSize); + L1Mat _l1(ChunkSize, ChunkSize); TASSIGN(_l1, 98304); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = ChunkSize; _gs.shape[4] = ChunkSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qk_gated_handle + static_cast(cid) * WsGatedSize, _gs); TLOAD(_l1, _gm); } // Load V { - L1Mat _l1(valid_rows, HiddenSize); + L1Mat _l1(valid_rows, HiddenSize); TASSIGN(_l1, 131072); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = valid_rows; _gs.shape[4] = HiddenSize; GmStride2D _stride(BSND_V_STRIDE); - GmTensor2D _gm(V_handle + v_off, _gs, _stride); + GmTensor2D _gm(V_handle + v_off, _gs, _stride); TLOAD(_l1, _gm); if (valid_rows != ChunkSize) TFILLPAD(_l1, _l1); } // GEMM 3: QKV = QK_gated @ V { - TileLeft _l0a; - TileRight _l0b; + TileLeft _l0a; + TileRight _l0b; TASSIGN(_l0a, 0x0); TASSIGN(_l0b, 0x0); auto _we = EVENT_ID1; set_flag(PIPE_MTE2, PIPE_MTE1, _we); wait_flag(PIPE_MTE2, PIPE_MTE1, _we); @@ -712,7 +712,7 @@ AICORE void GDN_CHUNK_O_KERNEL( TASSIGN(_l0, 0); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = ChunkSize; _gs.shape[4] = HiddenSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qs_qkv_handle + static_cast(cid) * WsQSSize, _gs); TSTORE(_gm, _l0); @@ -861,11 +861,11 @@ AICORE void GDN_CHUNK_O_KERNEL( { Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = local_rows; _gs.shape[4] = ChunkSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qk_handle + static_cast(cid) * WsQKSize + static_cast(vid) * HalfChunk * ChunkSize, _gs); - UbND _ld(local_rows, ChunkSize); + UbND _ld(local_rows, ChunkSize); TASSIGN(_ld, QKHalfUbAddr); TLOAD(_ld, _gm); if (local_rows != HalfChunk) { @@ -884,11 +884,11 @@ AICORE void GDN_CHUNK_O_KERNEL( { Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = local_rows; _gs.shape[4] = HiddenSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qs_qkv_handle + static_cast(cid) * WsQSSize + static_cast(vid) * HalfChunk * HiddenSize, _gs); - UbND _ld(local_rows, HiddenSize); + UbND _ld(local_rows, HiddenSize); TASSIGN(_ld, QSHalfUbAddr); TLOAD(_ld, _gm); if (local_rows != HalfChunk) { @@ -906,11 +906,11 @@ AICORE void GDN_CHUNK_O_KERNEL( { Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = local_rows; _gs.shape[4] = ChunkSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qk_gated_handle + static_cast(cid) * WsGatedSize + static_cast(vid) * HalfChunk * ChunkSize, _gs); - UbND _st(local_rows, ChunkSize); + UbND _st(local_rows, ChunkSize); TASSIGN(_st, QKHalfUbAddr); TSTORE(_gm, _st); } @@ -941,11 +941,11 @@ AICORE void GDN_CHUNK_O_KERNEL( { Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = local_rows; _gs.shape[4] = HiddenSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qs_qkv_handle + static_cast(cid) * WsQSSize + static_cast(vid) * HalfChunk * HiddenSize, _gs); - UbND _ld(local_rows, HiddenSize); + UbND _ld(local_rows, HiddenSize); TASSIGN(_ld, OHalfUbAddr); TLOAD(_ld, _gm); if (local_rows != HalfChunk) { @@ -960,7 +960,7 @@ AICORE void GDN_CHUNK_O_KERNEL( // ── Final output: O = QKV + QS_scaled ───────────────────────────── // numpy: O = (QK_gated @ V) + (Q @ S) * exp(g)[:, None] // = intra_chunk_attention + inter_chunk_state_contribution - // TCVT half→float for QKV, then TADD, then TCVT float→half for output. + // TCVT DTYPE_Q→float for QKV, then TADD, then TCVT float→DTYPE_Q for output. TCVT(o_ub, o_ub_half, pto::RoundMode::CAST_NONE); TADD(o_ub, qs_ub, o_ub); TCVT(o_ub_half, o_ub, pto::RoundMode::CAST_NONE); @@ -980,8 +980,8 @@ AICORE void GDN_CHUNK_O_KERNEL( Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = local_rows; _gs.shape[4] = HiddenSize; GmStride2D _stride(BSND_V_STRIDE); - GmTensor2D _gm(O_handle + o_offset, _gs, _stride); - UbND _st(local_rows, HiddenSize); + GmTensor2D _gm(O_handle + o_offset, _gs, _stride); + UbND _st(local_rows, HiddenSize); TASSIGN(_st, OHalfUbAddr); TSTORE(_gm, _st); } @@ -1068,11 +1068,11 @@ AICORE void GDN_CHUNK_O_KERNEL( { Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = local_rows; _gs.shape[4] = ChunkSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qk_handle + static_cast(cid) * WsQKSize + static_cast(vid) * HalfChunk * ChunkSize, _gs); - UbND _ld(local_rows, ChunkSize); + UbND _ld(local_rows, ChunkSize); TASSIGN(_ld, QKHalfUbAddr); TLOAD(_ld, _gm); if (local_rows != HalfChunk) { @@ -1091,11 +1091,11 @@ AICORE void GDN_CHUNK_O_KERNEL( { Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = local_rows; _gs.shape[4] = HiddenSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qs_qkv_handle + static_cast(cid) * WsQSSize + static_cast(vid) * HalfChunk * HiddenSize, _gs); - UbND _ld(local_rows, HiddenSize); + UbND _ld(local_rows, HiddenSize); TASSIGN(_ld, QSHalfUbAddr); TLOAD(_ld, _gm); if (local_rows != HalfChunk) { @@ -1104,7 +1104,7 @@ AICORE void GDN_CHUNK_O_KERNEL( } TMUL(qk_ub, qk_ub, coeff_ub); - TCVT(qk_ub_half, qk_ub, pto::RoundMode::CAST_NONE); // float→half for GM store + TCVT(qk_ub_half, qk_ub, pto::RoundMode::CAST_NONE); // float→DTYPE_Q for GM store // Store QK_gated → workspace set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); @@ -1112,11 +1112,11 @@ AICORE void GDN_CHUNK_O_KERNEL( { Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = local_rows; _gs.shape[4] = ChunkSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qk_gated_handle + static_cast(cid) * WsGatedSize + static_cast(vid) * HalfChunk * ChunkSize, _gs); - UbND _st(local_rows, ChunkSize); + UbND _st(local_rows, ChunkSize); TASSIGN(_st, QKHalfUbAddr); TSTORE(_gm, _st); } @@ -1127,7 +1127,7 @@ AICORE void GDN_CHUNK_O_KERNEL( // (same inter-chunk state scaling as fixed-length path) set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); - TCVT(qs_ub, qs_ub_half, pto::RoundMode::CAST_NONE); // half→float for Vec math + TCVT(qs_ub, qs_ub_half, pto::RoundMode::CAST_NONE); // DTYPE_Q→float for Vec math UbND g_exp_2d_v; TASSIGN(g_exp_2d_v, CoeffUbAddr); @@ -1143,11 +1143,11 @@ AICORE void GDN_CHUNK_O_KERNEL( { Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = local_rows; _gs.shape[4] = HiddenSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_qs_qkv_handle + static_cast(cid) * WsQSSize + static_cast(vid) * HalfChunk * HiddenSize, _gs); - UbND _ld(local_rows, HiddenSize); + UbND _ld(local_rows, HiddenSize); TASSIGN(_ld, OHalfUbAddr); TLOAD(_ld, _gm); if (local_rows != HalfChunk) { @@ -1159,9 +1159,9 @@ AICORE void GDN_CHUNK_O_KERNEL( wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); // O = QS_gated + QKV (final output: intra-chunk attention + inter-chunk state) - TCVT(o_ub, o_ub_half, pto::RoundMode::CAST_NONE); // half→float + TCVT(o_ub, o_ub_half, pto::RoundMode::CAST_NONE); // DTYPE_Q→float TADD(o_ub, qs_ub, o_ub); // O = QS_scaled + QKV - TCVT(o_ub_half, o_ub, pto::RoundMode::CAST_NONE); // float→half for GM store + TCVT(o_ub_half, o_ub, pto::RoundMode::CAST_NONE); // float→DTYPE_Q for GM store // Store O → GM set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); @@ -1178,8 +1178,8 @@ AICORE void GDN_CHUNK_O_KERNEL( Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = local_rows; _gs.shape[4] = HiddenSize; GmStride2D _stride(BSND_V_STRIDE); - GmTensor2D _gm(O_handle + o_offset, _gs, _stride); - UbND _st(local_rows, HiddenSize); + GmTensor2D _gm(O_handle + o_offset, _gs, _stride); + UbND _st(local_rows, HiddenSize); TASSIGN(_st, OHalfUbAddr); TSTORE(_gm, _st); } diff --git a/xllm_ops/mega_chunk_gdn/op_kernel/mega_chunk_gdn.cpp b/xllm_ops/mega_chunk_gdn/op_kernel/mega_chunk_gdn.cpp index 057c84a..04f7326 100644 --- a/xllm_ops/mega_chunk_gdn/op_kernel/mega_chunk_gdn.cpp +++ b/xllm_ops/mega_chunk_gdn/op_kernel/mega_chunk_gdn.cpp @@ -39,6 +39,9 @@ #include using namespace pto; +static_assert(std::is_same_v || std::is_same_v, + "MegaChunkGdn supports FP16 or BF16 compute tensors."); + struct MegaChunkGdnKernelTilingData { uint32_t block_dim; uint32_t num_matrices; @@ -56,12 +59,8 @@ struct MegaChunkGdnKernelTilingData { // =================================================================== #ifdef __CCE_AICORE__ -constexpr uint16_t SYNC_AIV_FLAG = 12; -constexpr uint16_t SYNC_AIC_FLAG = 11; -constexpr uint16_t SYNC_AIC_AIV_FLAG = 13; -// NOTE: SYNC_AIV_ONLY_ALL (==14) is provided by pto (pto::SYNC_AIV_ONLY_ALL) via -// `using namespace pto;`. Defining it here again causes an ambiguous reference, -// so rely on the pto-provided constant instead. +// Sync flag ids are provided by PTO. Keep this kernel on the shared sync +// contract so updates to the runtime synchronization framework stay aligned. constexpr uint16_t SYNC_MODE_SHIFT_VALUE = 4; constexpr uint16_t SYNC_FLAG_SHIFT_VALUE = 8; @@ -173,7 +172,7 @@ AICORE inline void mega_transpose_TH_to_HT(__gm__ T *src, __gm__ T *dst, int64_t } template -AICORE inline void mega_cast_fp32_to_fp16_bsnd(__gm__ float *src, __gm__ half *dst, uint32_t num_matrices, +AICORE inline void mega_cast_fp32_to_dtype_bsnd(__gm__ float *src, __gm__ DTYPE_Q *dst, uint32_t num_matrices, int64_t total_tokens) { #if defined(__DAV_C220_VEC__) @@ -190,8 +189,8 @@ AICORE inline void mega_cast_fp32_to_fp16_bsnd(__gm__ float *src, __gm__ half *d using SrcUB = Tile; using DynSrcUB = Tile; - using DstUB = Tile; - using DynDstUB = Tile; + using DstUB = Tile; + using DynDstUB = Tile; using Gm1D = Shape<1, 1, 1, 1, DYNAMIC>; using GmS1 = Stride<1, 1, 1, 1, 1>; @@ -225,7 +224,7 @@ AICORE inline void mega_cast_fp32_to_fp16_bsnd(__gm__ float *src, __gm__ half *d { Gm1D gs; gs.shape[4] = C; - GlobalTensor gm(dst + off, gs); + GlobalTensor gm(dst + off, gs); DstUB st; TASSIGN(st, F16_UB); TSTORE(gm, st); @@ -278,18 +277,18 @@ namespace mk_o { #define GDN_CHUNK_O_CALL chunk_o_kernel #endif -AICORE inline void mega_solve_tril(__gm__ half *out, __gm__ half *in, __gm__ half *minus_id, uint32_t matrix_size, +AICORE inline void mega_solve_tril(__gm__ DTYPE_Q *out, __gm__ DTYPE_Q *in, __gm__ DTYPE_Q *minus_id, uint32_t matrix_size, uint32_t num_matrices, uint32_t num_bsnd_heads, __gm__ int32_t *cu_seqlens, uint32_t is_lower) { if (num_matrices <= get_block_num()) - mk_solve::runKernelTriInvRecUnroll(out, in, minus_id, num_matrices, + mk_solve::runKernelTriInvRecUnroll(out, in, minus_id, num_matrices, num_bsnd_heads, cu_seqlens, is_lower); else if (num_matrices <= 2u * get_block_num()) - mk_solve::runKernelTriInvRecUnroll(out, in, minus_id, num_matrices, + mk_solve::runKernelTriInvRecUnroll(out, in, minus_id, num_matrices, num_bsnd_heads, cu_seqlens, is_lower); else - mk_solve::runKernelTriInvRecUnroll(out, in, minus_id, num_matrices, + mk_solve::runKernelTriInvRecUnroll(out, in, minus_id, num_matrices, num_bsnd_heads, cu_seqlens, is_lower); } @@ -329,8 +328,8 @@ AICORE inline void mega_kernel_impl(GM_ADDR q_ptr, GM_ADDR k_ptr, GM_ADDR v_ptr, mega_transpose_TH_to_HT(reinterpret_cast<__gm__ float *>(g_sum_ptr), reinterpret_cast<__gm__ float *>(g_t_ptr), total_tokens, H); - mega_transpose_TH_to_HT(reinterpret_cast<__gm__ half *>(beta_ptr), - reinterpret_cast<__gm__ half *>(beta_t_ptr), total_tokens, H); + mega_transpose_TH_to_HT(reinterpret_cast<__gm__ DTYPE_Q *>(beta_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(beta_t_ptr), total_tokens, H); #ifdef MEGA_STOP_AFTER_TRANSPOSE pipe_barrier(PIPE_ALL); @@ -340,9 +339,9 @@ AICORE inline void mega_kernel_impl(GM_ADDR q_ptr, GM_ADDR k_ptr, GM_ADDR v_ptr, SyncAllImpl(); mk_kkt::kkt_kernel( - reinterpret_cast<__gm__ half *>(k_ptr), reinterpret_cast<__gm__ half *>(beta_t_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(k_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(beta_t_ptr), reinterpret_cast<__gm__ float *>(g_t_ptr), reinterpret_cast<__gm__ float *>(msk_lower_ptr), - reinterpret_cast<__gm__ half *>(kkt_ws_ptr), reinterpret_cast<__gm__ half *>(A_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(kkt_ws_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(A_ptr), reinterpret_cast<__gm__ int32_t *>(cu_seqlens_ptr), batch_size, seq_len, total_tokens, static_cast(H), num_key_heads, ffts_addr); @@ -359,8 +358,8 @@ AICORE inline void mega_kernel_impl(GM_ADDR q_ptr, GM_ADDR k_ptr, GM_ADDR v_ptr, SyncAllImpl(); - mega_solve_tril(reinterpret_cast<__gm__ half *>(A_inv_ptr), reinterpret_cast<__gm__ half *>(A_ptr), - reinterpret_cast<__gm__ half *>(minus_id_ptr), C, num_matrices, H, + mega_solve_tril(reinterpret_cast<__gm__ DTYPE_Q *>(A_inv_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(A_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(minus_id_ptr), C, num_matrices, H, reinterpret_cast<__gm__ int32_t *>(cu_seqlens_ptr), 1); #ifdef MEGA_STOP_AFTER_SOLVE @@ -382,11 +381,11 @@ AICORE inline void mega_kernel_impl(GM_ADDR q_ptr, GM_ADDR k_ptr, GM_ADDR v_ptr, #endif mk_wy::GDN_WY_FAST_CALL( - reinterpret_cast<__gm__ half *>(k_ptr), reinterpret_cast<__gm__ half *>(v_ptr), - reinterpret_cast<__gm__ half *>(beta_t_ptr), reinterpret_cast<__gm__ float *>(g_t_ptr), - reinterpret_cast<__gm__ half *>(A_inv_ptr), reinterpret_cast<__gm__ half *>(wy_ws_a1_ptr), - reinterpret_cast<__gm__ half *>(wy_ws_a2_ptr), reinterpret_cast<__gm__ half *>(w_ptr), - reinterpret_cast<__gm__ half *>(u_ptr), reinterpret_cast<__gm__ int32_t *>(cu_seqlens_ptr), batch_size, seq_len, + reinterpret_cast<__gm__ DTYPE_Q *>(k_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(v_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(beta_t_ptr), reinterpret_cast<__gm__ float *>(g_t_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(A_inv_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(wy_ws_a1_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(wy_ws_a2_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(w_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(u_ptr), reinterpret_cast<__gm__ int32_t *>(cu_seqlens_ptr), batch_size, seq_len, total_tokens, static_cast(H), num_key_heads, ffts_addr); #if defined(__DAV_C220_VEC__) @@ -405,12 +404,12 @@ AICORE inline void mega_kernel_impl(GM_ADDR q_ptr, GM_ADDR k_ptr, GM_ADDR v_ptr, SyncAllImpl(); mk_h::chunk_h_kernel( - reinterpret_cast<__gm__ half *>(k_ptr), reinterpret_cast<__gm__ half *>(w_ptr), - reinterpret_cast<__gm__ half *>(u_ptr), reinterpret_cast<__gm__ float *>(g_t_ptr), - reinterpret_cast<__gm__ half *>(s_ptr), reinterpret_cast<__gm__ half *>(v_new_ptr), - reinterpret_cast<__gm__ half *>(fs_ptr), reinterpret_cast<__gm__ half *>(h0_ptr), has_initial_state, + reinterpret_cast<__gm__ DTYPE_Q *>(k_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(w_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(u_ptr), reinterpret_cast<__gm__ float *>(g_t_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(s_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(v_new_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(fs_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(h0_ptr), has_initial_state, 1, - reinterpret_cast<__gm__ half *>(h_ws_ptr), reinterpret_cast<__gm__ int32_t *>(cu_seqlens_ptr), batch_size, + reinterpret_cast<__gm__ DTYPE_Q *>(h_ws_ptr), reinterpret_cast<__gm__ int32_t *>(cu_seqlens_ptr), batch_size, seq_len, total_tokens, static_cast(H), num_key_heads, ffts_addr); #ifdef MEGA_STOP_AFTER_H @@ -421,11 +420,11 @@ AICORE inline void mega_kernel_impl(GM_ADDR q_ptr, GM_ADDR k_ptr, GM_ADDR v_ptr, SyncAllImpl(); mk_o::GDN_CHUNK_O_CALL( - reinterpret_cast<__gm__ half *>(q_ptr), reinterpret_cast<__gm__ half *>(k_ptr), - reinterpret_cast<__gm__ half *>(v_new_ptr), reinterpret_cast<__gm__ half *>(s_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(q_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(k_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(v_new_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(s_ptr), reinterpret_cast<__gm__ float *>(g_t_ptr), reinterpret_cast<__gm__ float *>(msk_full_ptr), - reinterpret_cast<__gm__ half *>(o_ws_qk_ptr), reinterpret_cast<__gm__ half *>(o_ws_qs_ptr), - reinterpret_cast<__gm__ half *>(o_ws_gated_ptr), reinterpret_cast<__gm__ half *>(o_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(o_ws_qk_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(o_ws_qs_ptr), + reinterpret_cast<__gm__ DTYPE_Q *>(o_ws_gated_ptr), reinterpret_cast<__gm__ DTYPE_Q *>(o_ptr), reinterpret_cast<__gm__ int32_t *>(cu_seqlens_ptr), batch_size, seq_len, total_tokens, static_cast(H), num_key_heads, ffts_addr); @@ -450,7 +449,7 @@ GDN_KERNEL_NAME(GM_ADDR q_ptr, GM_ADDR k_ptr, GM_ADDR v_ptr, GM_ADDR g_in_ptr, G REGISTER_TILING_DEFAULT(MegaChunkGdnKernelTilingData); GET_TILING_DATA_WITH_STRUCT(MegaChunkGdnKernelTilingData, tiling_data, tiling); GM_ADDR user_ws = AscendC::GetUserWorkspace(workspace); - const uint64_t tile_bytes = static_cast(GDN_C) * GDN_C * sizeof(half); + const uint64_t tile_bytes = static_cast(GDN_C) * GDN_C * sizeof(DTYPE_Q); GM_ADDR kkt_ws_ptr = user_ws; GM_ADDR wy_ws_a1_ptr = kkt_ws_ptr + static_cast(tiling_data.block_dim) * 2 * tile_bytes; diff --git a/xllm_ops/mega_chunk_gdn/op_kernel/scaled_dot_kkt.cpp b/xllm_ops/mega_chunk_gdn/op_kernel/scaled_dot_kkt.cpp index 47d882a..669057e 100644 --- a/xllm_ops/mega_chunk_gdn/op_kernel/scaled_dot_kkt.cpp +++ b/xllm_ops/mega_chunk_gdn/op_kernel/scaled_dot_kkt.cpp @@ -7,13 +7,13 @@ // A[i,j] = KK^T[i,j] · coeff[i,j] · causal_mask[i,j] // // Inputs: -// K [total_tokens, Hg, D] half — key vectors (BSND along seq; stride Hg * D) -// Beta [H, total_tokens] half — gate bias per **value** head (pre-transposed) +// K [total_tokens, Hg, D] DTYPE_Q — key vectors (BSND along seq; stride Hg * D) +// Beta [H, total_tokens] DTYPE_Q — gate bias per **value** head (pre-transposed) // G [H, total_tokens] float — cumulative gate sum per **value** head // Msk [C, C] float — lower-triangular causal mask // // Output: -// A [total_tokens, H, C] half — gated attention matrix in BSND +// A [total_tokens, H, C] DTYPE_Q — gated attention matrix in BSND // // Architecture: Cube + Vec cross-core kernel. // Cube phase: K→L1, GEMM K@K^T→L0C, store to workspace (GM) @@ -50,7 +50,7 @@ // TEXTRACT(l0, l1, r, c) — Copy L1 sub-block → L0A or L0B (MTE1 pipe) // TRESHAPE(dst, src) — Reinterpret L1 tile layout (NZ↔ZN for transpose) // TMATMUL(C, A, B) — Matrix multiply: C = A @ B in Cube engine -// TCVT(dst, src, mode) — Type conversion: like dst = src.float() or src.half() +// TCVT(dst, src, mode) — converts between float and DTYPE_Q // TMOV(dst, src) — Copy: dst = src.clone() // TADD(d, a, b) — Element-wise add: d = a + b // TSUB(d, a, b) — Element-wise subtract: d = a - b @@ -134,9 +134,9 @@ using GmTensor2D = pto::GlobalTensor; // CUDA's device memory. All input/output tensors live in GM. template AICORE void kkt_kernel( - __gm__ half *K_handle, __gm__ half *Beta_handle, + __gm__ DTYPE_Q *K_handle, __gm__ DTYPE_Q *Beta_handle, __gm__ float *G_handle, __gm__ float *Msk_handle, - __gm__ half *workspace_handle, __gm__ half *A_handle, + __gm__ DTYPE_Q *workspace_handle, __gm__ DTYPE_Q *A_handle, __gm__ int32_t *cu_seqlens, int64_t batch_size, int64_t seq_len, int64_t total_tokens, @@ -161,7 +161,7 @@ AICORE void kkt_kernel( // The UB is a flat SRAM; we manually assign byte offsets for each tile. // This is like malloc'ing fixed regions — no dynamic allocator on NPU. constexpr int32_t GUbAddr = 0; // g_ub: cumulative gates [1×C] - constexpr int32_t BetaHalfUbAddr = 512; // beta_ub_half: gate bias fp16 [1×C/2] + constexpr int32_t BetaHalfUbAddr = 512; // beta_ub_half: gate bias DTYPE_Q [1×C/2] constexpr int32_t BetaUbAddr = 640; // beta_ub: gate bias fp32 [1×C/2] constexpr int32_t GvUbAddr = 896; // g_v_ub: combined gate+bias [1×C/2] constexpr int32_t AUbAddr = 1152; // a_ub: attention sub-block fp32 [C/2×C] @@ -194,13 +194,13 @@ AICORE void kkt_kernel( // ── Cube-side tile declarations ───────────────────────────────────── // Cube-side tiles: K in L1 (NZ format), accumulator in L0C - L1Mat k_l1; TASSIGN(k_l1, 0); // TileAcc: L0C accumulator tile for GEMM results. // The Cube engine always accumulates in float32 for precision, even when - // inputs are fp16. Think of it as: result = torch.matmul(a.half(), b.half()).float() - // When stored to GM via TSTORE with a half GlobalTensor, automatic fp32→fp16 cast occurs. + // Inputs use DTYPE_Q. Conceptually: result = torch.matmul(a, b).float(). + // Storing through a DTYPE_Q GlobalTensor converts the FP32 accumulator to DTYPE_Q. TileAcc a_l0; TASSIGN(a_l0, 0); @@ -210,7 +210,7 @@ AICORE void kkt_kernel( // Vec-side UB tiles for gating computation UbND g_ub; TASSIGN(g_ub, GUbAddr); - UbND beta_ub_half; + UbND beta_ub_half; TASSIGN(beta_ub_half, BetaHalfUbAddr); UbND beta_ub; TASSIGN(beta_ub, BetaUbAddr); @@ -235,7 +235,7 @@ AICORE void kkt_kernel( UbND coeff_ub; TASSIGN(coeff_ub, CoeffUbAddr); - UbND a_ub_half; TASSIGN(a_ub_half, AUbHalfAddr); @@ -253,7 +253,7 @@ AICORE void kkt_kernel( // Left operand: K [C×D] loaded into L1 in NZ format // Right operand: K^T — same data, but we TRESHAPE to ZN format // (TRESHAPE is FREE — it just reinterprets the fractal layout as transposed) - // Result: KK^T [C×C] in L0C (float32 accumulator, even though inputs are fp16) + // Result: KK^T [C×C] in L0C (FP32 accumulator with DTYPE_Q inputs) // ======================================================================== // __DAV_C220_CUBE__: This code only compiles for the Cube core. // On NPU, Cube and Vec are separate compilation targets (like two different GPUs). @@ -314,11 +314,11 @@ AICORE void kkt_kernel( // If the chunk is partial, TFILLPAD zero-fills the padding region // so the GEMM doesn't produce garbage from uninitialized memory. { - L1Mat _l1(valid_rows, HiddenSize); + L1Mat _l1(valid_rows, HiddenSize); TASSIGN(_l1, 0); GmShape2D _gs(valid_rows, HiddenSize); GmStride2D _stride(bsnd_qk_stride); - GmTensor2D _gm(K_handle + k_offset, _gs, _stride); + GmTensor2D _gm(K_handle + k_offset, _gs, _stride); TLOAD(_l1, _gm); if (valid_rows != ChunkSize) TFILLPAD(_l1, _l1); } @@ -337,8 +337,8 @@ AICORE void kkt_kernel( // This is like ensuring a producer-consumer chain is properly ordered. // WAR sync: MTE2→MTE1, M→MTE1 before extract; MTE1→M before matmul. { - TileLeft _l0a; - TileRight _l0b; + TileLeft _l0a; + TileRight _l0b; TASSIGN(_l0a, 0x0); TASSIGN(_l0b, 0x0); auto _we = EVENT_ID1; @@ -349,7 +349,7 @@ AICORE void kkt_kernel( // Left operand: K in NZ format, extract directly to L0A TEXTRACT(_l0a, k_l1, 0, 0); // Right operand: K^T via ZN reshape of same L1 tile, extract to L0B - L1MatZN _bzn; + L1MatZN _bzn; TRESHAPE(_bzn, k_l1); TEXTRACT(_l0b, _bzn, 0, 0); set_flag(PIPE_MTE1, PIPE_M, _we); @@ -361,13 +361,13 @@ AICORE void kkt_kernel( wait_flag(PIPE_M, PIPE_FIX, _we); } - // ── Store KK^T from L0C → workspace GM (with fp32→fp16 cast) ─── + // ── Store KK^T from L0C → workspace GM (with FP32→DTYPE_Q cast) ─── { TileAcc _l0(ChunkSize, ChunkSize); TASSIGN(_l0, 0); Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = ChunkSize; _gs.shape[4] = ChunkSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_handle + (static_cast(cid) * 2 + slot) * ChunkSquare, _gs); @@ -402,7 +402,7 @@ AICORE void kkt_kernel( // Each sub-block (vid=0,1) handles HalfChunk rows of the C×C matrix. // // ── Gating computation (numpy pseudocode) ───────────────────────────── - // # For each sub-block's C/2 rows (vid selects upper or lower half): + // # For each sub-block's C/2 rows (vid selects the upper or lower half): // g_row = g_sum[row_offset:row_offset+C/2] # this sub-block's gates // g_v = g_row + np.log(beta[row_offset:row_offset+C/2]) # combined gate+bias // g_col = g_sum[0:C] # full chunk gates @@ -426,7 +426,7 @@ AICORE void kkt_kernel( set_vector_mask(-1, -1); // ── Load causal mask (lower triangular) once, reused across all chunks ── - // vid=0 loads the top half (rows 0..C/2-1), vid=1 loads the bottom half. + // vid=0 loads the upper half (rows 0..C/2-1), vid=1 loads the lower half. // The mask is [C×C] in GM; each sub-block loads its [C/2×C] portion. { Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; @@ -507,19 +507,19 @@ AICORE void kkt_kernel( } } - // Beta is [H, total_tokens] half — contiguous per head + // Beta is [H, total_tokens] DTYPE_Q — contiguous per head { Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = 1; _gs.shape[4] = local_valid; - GlobalTensor> _gm( + GlobalTensor> _gm( Beta_handle + static_cast(head_idx) * total_tokens + (bos + chunk_start + row_offset), _gs); - UbND _ld(1, local_valid); + UbND _ld(1, local_valid); TASSIGN(_ld, BetaHalfUbAddr); TLOAD(_ld, _gm); if (local_valid != HalfChunk) { - UbND _pd; + UbND _pd; TASSIGN(_pd, BetaHalfUbAddr); TFILLPAD_INPLACE(_pd, _ld); } @@ -532,7 +532,7 @@ AICORE void kkt_kernel( if (local_valid > 0) { // ── Compute gating coefficient ──────────────────────────────── - // Step 1: Convert beta from fp16→fp32 for precision + // Step 1: Convert beta from DTYPE_Q→FP32 for precision // Step 2: g_v[i] = g[row_offset+i] + log(β[i]) — combined row gate // Step 3: Broadcast g_v (rows) and g (cols) to 2D matrices // Step 4: coeff = exp(min(g_v_2d - g_2d, 0)) — clamped exponential gating @@ -579,18 +579,18 @@ AICORE void kkt_kernel( set_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); - // ── Load KK^T sub-block from workspace (fp16) ──────────────── + // ── Load KK^T sub-block from workspace (DTYPE_Q) ───────────── // workspace layout: [core_id * 2 + slot][C×C], we load our sub-block's // [C/2×C] portion (offset by vid * HalfChunk * ChunkSize elements). { Shape<1, 1, 1, DYNAMIC, DYNAMIC> _gs; _gs.shape[3] = HalfChunk; _gs.shape[4] = ChunkSize; - GlobalTensor> _gm( + GlobalTensor> _gm( workspace_handle + (static_cast(cid) * 2 + slot) * ChunkSquare + static_cast(vid) * HalfChunk * ChunkSize, _gs); - UbND _ld(HalfChunk, ChunkSize); + UbND _ld(HalfChunk, ChunkSize); TASSIGN(_ld, AUbHalfAddr); TLOAD(_ld, _gm); } @@ -600,13 +600,13 @@ AICORE void kkt_kernel( wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); // ── Apply gating and mask: A = KK^T · coeff · mask ─────────── - // 1. Convert KK^T from fp16 → fp32 (Cube stored it as fp16 to save GM bandwidth) + // 1. Convert KK^T from DTYPE_Q → FP32 (stored compactly to save GM bandwidth) TCVT(a_ub, a_ub_half, pto::RoundMode::CAST_NONE); // 2. Element-wise multiply by gating coefficient TMUL(a_ub, a_ub, coeff_ub); // 3. Element-wise multiply by causal mask (lower triangular, zeros above diagonal) TMUL(a_ub, a_ub, msk_ub); - // 4. Convert result back to fp16 for output + // 4. Convert result back to DTYPE_Q for output TCVT(a_ub_half, a_ub, pto::RoundMode::CAST_NONE); // V→MTE3 sync: Vec computation done, safe for DMA store to begin @@ -625,8 +625,8 @@ AICORE void kkt_kernel( { GmShape2D _gs(local_valid, ChunkSize); GmStride2D _stride(H * ChunkSize); - GmTensor2D _gm(A_handle + a_gm_offset, _gs, _stride); - UbND _st(local_valid, ChunkSize); + GmTensor2D _gm(A_handle + a_gm_offset, _gs, _stride); + UbND _st(local_valid, ChunkSize); TASSIGN(_st, AUbHalfAddr); TSTORE(_gm, _st); } diff --git a/xllm_ops/mega_chunk_gdn/op_kernel/wy_fast.cpp b/xllm_ops/mega_chunk_gdn/op_kernel/wy_fast.cpp index 021a875..b46ae9a 100644 --- a/xllm_ops/mega_chunk_gdn/op_kernel/wy_fast.cpp +++ b/xllm_ops/mega_chunk_gdn/op_kernel/wy_fast.cpp @@ -39,13 +39,13 @@ // Key PTO APIs (with numpy/torch equivalents): // TLOAD(ub_tile, gm) — ub_tile = gm[...] (DMA: GM→UB, async MTE2) // TSTORE(gm, ub_tile) — gm[...] = ub_tile (DMA: UB→GM, async MTE3) -// TCVT(dst, src, mode) — dst = src.float() or .half() (type conversion) +// TCVT(dst, src, mode) — converts between float and DTYPE_Q // TMOV(dst, src) — dst = src.clone() // TMUL(d, a, b) — d = a * b (element-wise) // TEXP(d, s) — d = torch.exp(s) // TCOLEXPAND(2d, row) — 2d[i,j] = row[j] (broadcast row across all rows) // TEXTRACT(l0, l1, r, c) — L1 sub-block → L0A/L0B (MTE1 for Cube GEMM) -// TMATMUL(C, A, B) — C = A @ B in Cube engine (fp16→fp32 accumulate) +// TMATMUL(C, A, B) — C = A @ B in Cube engine (DTYPE_Q→FP32 accumulate) // set_flag / wait_flag — sync between pipes on SAME core // ffts_cross_core_sync — signal ACROSS Cube↔Vec cores // wait_flag_dev(flag) — wait for cross-core signal @@ -265,11 +265,11 @@ gemm_v0(std::conditional_t, template AICORE void GDN_WY_FAST_KERNEL( - __gm__ half *K_handle, __gm__ half *V_handle, - __gm__ half *Beta_handle, __gm__ float *G_handle, - __gm__ half *A_handle, - __gm__ half *workspace_a1_handle, __gm__ half *workspace_a2_handle, - __gm__ half *W_handle, __gm__ half *U_handle, + __gm__ DTYPE_Q *K_handle, __gm__ DTYPE_Q *V_handle, + __gm__ DTYPE_Q *Beta_handle, __gm__ float *G_handle, + __gm__ DTYPE_Q *A_handle, + __gm__ DTYPE_Q *workspace_a1_handle, __gm__ DTYPE_Q *workspace_a2_handle, + __gm__ DTYPE_Q *W_handle, __gm__ DTYPE_Q *U_handle, __gm__ int32_t *cu_seqlens, int64_t batch_size, int64_t seq_len, int64_t total_tokens, uint32_t num_heads, @@ -334,10 +334,10 @@ AICORE void GDN_WY_FAST_KERNEL( int64_t num_seqs = batch_size; - TileUbDataND beta_ub_half; TASSIGN(beta_ub_half, BetaHalfUbAddr); - TileUbDataND a1_ub_half; TASSIGN(a1_ub_half, A1HalfUbAddr); TileUbDataND beta_ub; @@ -355,7 +355,7 @@ AICORE void GDN_WY_FAST_KERNEL( TileUbDataND a2_ub; TASSIGN(a2_ub, A2UbAddr); - TileUbDataND a2_ub_half; TASSIGN(a2_ub_half, A2HalfUbAddr); TileUbDataND g_2d_ub; TASSIGN(g_2d_ub, G2dUbAddr); - TileMatL1 k_l1; TASSIGN(k_l1, 0); - TileMatL1 v_l1; TASSIGN(v_l1, 32768); - TileMatL1 a2_l1; TASSIGN(a2_l1, 65536); TileAcc u_l0; TASSIGN(u_l0, 0); - TileMatL1 a1_l1; TASSIGN(a1_l1, 98304); TileAcc beta_global( + GmTensor2D beta_global( Beta_handle + static_cast(head_idx) * total_tokens + chunk_token_start, beta_shape, beta_stride); - DynVecTile beta_load( + DynVecTile beta_load( 1, valid_rows); TASSIGN(beta_load, BetaHalfUbAddr); TLOAD(beta_load, beta_global); @@ -464,9 +464,9 @@ AICORE void GDN_WY_FAST_KERNEL( static_cast(ChunkSize); GmShape2D a_shape(local_rows, ChunkSize); GmStride2D a_stride(H * ChunkSize); - GmTensor2D a_global(A_handle + a_gm_offset, a_shape, + GmTensor2D a_global(A_handle + a_gm_offset, a_shape, a_stride); - DynVecTile a_load( + DynVecTile a_load( local_rows, ChunkSize); TASSIGN(a_load, A1HalfUbAddr); TLOAD(a_load, a_global); @@ -505,7 +505,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D a2_shape(HalfChunk, ChunkSize); GmStride2D a2_stride(ChunkSize); - GmTensor2D workspace_a2_global( + GmTensor2D workspace_a2_global( workspace_a2_handle + static_cast(cid) * WsA2Size + static_cast(vid) * HalfChunk * ChunkSize, @@ -556,7 +556,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D a1_shape(HalfChunk, ChunkSize); GmStride2D a1_stride(ChunkSize); - GmTensor2D workspace_a1_global( + GmTensor2D workspace_a1_global( workspace_a1_handle + static_cast(cid) * WsA1Size + static_cast(vid) * HalfChunk * ChunkSize, @@ -613,11 +613,11 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D beta_shape(1, valid_rows); GmStride2D beta_stride(1); - GmTensor2D beta_global( + GmTensor2D beta_global( Beta_handle + static_cast(head_idx) * total_tokens + chunk_token_start, beta_shape, beta_stride); - DynVecTile beta_load( + DynVecTile beta_load( 1, valid_rows); TASSIGN(beta_load, BetaHalfUbAddr); TLOAD(beta_load, beta_global); @@ -637,9 +637,9 @@ AICORE void GDN_WY_FAST_KERNEL( static_cast(ChunkSize); GmShape2D a_shape(local_rows, ChunkSize); GmStride2D a_stride(H * ChunkSize); - GmTensor2D a_global(A_handle + a_gm_offset, a_shape, + GmTensor2D a_global(A_handle + a_gm_offset, a_shape, a_stride); - DynVecTile a_load( + DynVecTile a_load( local_rows, ChunkSize); TASSIGN(a_load, A1HalfUbAddr); TLOAD(a_load, a_global); @@ -674,7 +674,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D a2_shape(HalfChunk, ChunkSize); GmStride2D a2_stride(ChunkSize); - GmTensor2D workspace_a2_global( + GmTensor2D workspace_a2_global( workspace_a2_handle + static_cast(cid) * WsA2Size + static_cast(vid) * HalfChunk * ChunkSize, @@ -721,7 +721,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D a1_shape(HalfChunk, ChunkSize); GmStride2D a1_stride(ChunkSize); - GmTensor2D workspace_a1_global( + GmTensor2D workspace_a1_global( workspace_a1_handle + static_cast(cid) * WsA1Size + static_cast(vid) * HalfChunk * ChunkSize, @@ -772,8 +772,8 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D k_shape(valid_rows, HiddenSize); GmStride2D k_stride(BSND_QK_STRIDE); - GmTensor2D k_global(K_handle + k_off, k_shape, k_stride); - DynMatL1 k_l1_load(valid_rows, + GmTensor2D k_global(K_handle + k_off, k_shape, k_stride); + DynMatL1 k_l1_load(valid_rows, HiddenSize); TASSIGN(k_l1_load, 0); TLOAD(k_l1_load, k_global); @@ -784,8 +784,8 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D v_shape(valid_rows, HiddenSize); GmStride2D v_stride(BSND_V_STRIDE); - GmTensor2D v_global(V_handle + v_off, v_shape, v_stride); - DynMatL1 v_l1_load(valid_rows, + GmTensor2D v_global(V_handle + v_off, v_shape, v_stride); + DynMatL1 v_l1_load(valid_rows, HiddenSize); TASSIGN(v_l1_load, 32768); TLOAD(v_l1_load, v_global); @@ -798,7 +798,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D a2_shape(ChunkSize, ChunkSize); GmStride2D a2_stride(ChunkSize); - GmTensor2D workspace_a2_global( + GmTensor2D workspace_a2_global( workspace_a2_handle + static_cast(cid) * WsA2Size, a2_shape, a2_stride); // Load the Vec-prepared A2 tile: @@ -809,7 +809,7 @@ AICORE void GDN_WY_FAST_KERNEL( set_flag(PIPE_FIX, PIPE_M, EVENT_ID0); wait_flag(PIPE_FIX, PIPE_M, EVENT_ID0); // U = A2 * V keeps the beta-scaled path separate from the K-side update. - gemm_v0(a2_l1, v_l1, u_l0, true); @@ -817,7 +817,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D u_shape(valid_rows, HiddenSize); GmStride2D u_stride(BSND_V_STRIDE); - GmTensor2D u_global(U_handle + v_off, u_shape, u_stride); + GmTensor2D u_global(U_handle + v_off, u_shape, u_stride); DynAccTile u_store(valid_rows, HiddenSize); TASSIGN(u_store, 0); @@ -831,7 +831,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D a1_shape(ChunkSize, ChunkSize); GmStride2D a1_stride(ChunkSize); - GmTensor2D workspace_a1_global( + GmTensor2D workspace_a1_global( workspace_a1_handle + static_cast(cid) * WsA1Size, a1_shape, a1_stride); // Load the Vec-prepared A1 tile: @@ -842,7 +842,7 @@ AICORE void GDN_WY_FAST_KERNEL( set_flag(PIPE_FIX, PIPE_M, EVENT_ID0); wait_flag(PIPE_FIX, PIPE_M, EVENT_ID0); // W = A1 * K uses the g-reweighted path for the complementary WY factor. - gemm_v0(a1_l1, k_l1, w_l0, true); @@ -850,7 +850,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D w_shape(valid_rows, HiddenSize); GmStride2D w_stride(BSND_V_STRIDE); - GmTensor2D w_global(W_handle + v_off, w_shape, w_stride); + GmTensor2D w_global(W_handle + v_off, w_shape, w_stride); DynAccTile w_store(valid_rows, HiddenSize); TASSIGN(w_store, 65536); @@ -894,9 +894,9 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D k_shape(valid_rows, HiddenSize); GmStride2D k_stride(BSND_QK_STRIDE); - GmTensor2D k_global(K_handle + k_off, k_shape, + GmTensor2D k_global(K_handle + k_off, k_shape, k_stride); - DynMatL1 k_l1_load(valid_rows, + DynMatL1 k_l1_load(valid_rows, HiddenSize); TASSIGN(k_l1_load, 0); TLOAD(k_l1_load, k_global); @@ -907,9 +907,9 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D v_shape(valid_rows, HiddenSize); GmStride2D v_stride(BSND_V_STRIDE); - GmTensor2D v_global(V_handle + v_off, v_shape, + GmTensor2D v_global(V_handle + v_off, v_shape, v_stride); - DynMatL1 v_l1_load(valid_rows, + DynMatL1 v_l1_load(valid_rows, HiddenSize); TASSIGN(v_l1_load, 32768); TLOAD(v_l1_load, v_global); @@ -922,7 +922,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D a2_shape(ChunkSize, ChunkSize); GmStride2D a2_stride(ChunkSize); - GmTensor2D workspace_a2_global( + GmTensor2D workspace_a2_global( workspace_a2_handle + static_cast(cid) * WsA2Size, a2_shape, a2_stride); TLOAD(a2_l1, workspace_a2_global); @@ -931,7 +931,7 @@ AICORE void GDN_WY_FAST_KERNEL( set_flag(PIPE_FIX, PIPE_M, EVENT_ID0); wait_flag(PIPE_FIX, PIPE_M, EVENT_ID0); // U = A2 * V keeps the beta-scaled path separate from the K-side update. - gemm_v0(a2_l1, v_l1, u_l0, true); @@ -939,7 +939,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D u_shape(valid_rows, HiddenSize); GmStride2D u_stride(BSND_V_STRIDE); - GmTensor2D u_global(U_handle + v_off, u_shape, + GmTensor2D u_global(U_handle + v_off, u_shape, u_stride); DynAccTile u_store(valid_rows, HiddenSize); @@ -952,7 +952,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D a1_shape(ChunkSize, ChunkSize); GmStride2D a1_stride(ChunkSize); - GmTensor2D workspace_a1_global( + GmTensor2D workspace_a1_global( workspace_a1_handle + static_cast(cid) * WsA1Size, a1_shape, a1_stride); TLOAD(a1_l1, workspace_a1_global); @@ -961,7 +961,7 @@ AICORE void GDN_WY_FAST_KERNEL( set_flag(PIPE_FIX, PIPE_M, EVENT_ID0); wait_flag(PIPE_FIX, PIPE_M, EVENT_ID0); // W = A1 * K uses the g-reweighted path for the complementary WY factor. - gemm_v0(a1_l1, k_l1, w_l0, true); @@ -969,7 +969,7 @@ AICORE void GDN_WY_FAST_KERNEL( { GmShape2D w_shape(valid_rows, HiddenSize); GmStride2D w_stride(BSND_V_STRIDE); - GmTensor2D w_global(W_handle + v_off, w_shape, + GmTensor2D w_global(W_handle + v_off, w_shape, w_stride); DynAccTile w_store(valid_rows, HiddenSize);