From a2c9a23c95628692224e997ee61cd1e3670c1c25 Mon Sep 17 00:00:00 2001 From: Youhezhen Date: Thu, 24 Sep 2026 01:36:13 -0700 Subject: [PATCH 1/2] Perf: score DSpark CSA indexer candidates on Cube indexer_score_topk_leaf dominates decode_csa at 128k. It kept a 64-head INT32 score on Vector and reduced it there with a cast, a ReLU, a broadcast multiply and a column sum. Move that reduction onto Cube, following the AscendC quant_lightning_indexer_v2 arch22 kernel. - FIXPIPE drains the INT32 QK accumulator from L0C straight into an FP16 Mat tile, applying ReLU on the accumulator and a fixed 1/1024 dequantization in the same writeback. This replaces Vector's cast and maximum, and the 64-head intermediate never reaches UB or GM. - A second Cube matmul contracts the heads against an FP16 coefficient, replacing Vector's row_expand_mul and col_sum. The dequantized score moves from Mat to L0B and the coefficient from Mat to L0A, so the reduction happens where the score already is. Vector keeps only the Key-scale multiply and Top-K, on a tile that is 64x narrower. - Two query tokens share one Key tile. Their coefficients sit block diagonally in one FP16 matrix, so a batch of GEMVs becomes a single GEMM and each Key tile is loaded into L0B once for two tokens instead of twice. This geometry follows #1380. - The coefficient is built by a loop inside weights_proj_reduce alongside the weight that task already computes, so its preparation costs no task of its own. It carries a 1024x factor that cancels the writeback scale exactly and keeps small coefficients clear of FP16 subnormals. A2/A3, TP1, B16, S6, swimlane 0, 100 rounds after 5 warmups, A-B-B-A against 07b5d2f, one session on one device against one frozen fixture. At start_pos = 16 x 131072 the median goes 1499.6 -> 1241.3 us, -258.3 us (-17.2%). At 16 x 8192 it goes 726.5 -> 707.2 us, -19.3 us (-2.7%). On both workloads each arm's median is below both baseline arms'. The level-4 trace at 128k puts the score task at 833.6 -> 547.5 us and the whole operation at 1623.6 -> 1358.7 us over an unchanged 983 blocks. topk_idxs keeps main's comparator, topk_pair_compare, which admits an index swap only between tied scores and rejects any strictly worse pick. Contracting on Cube rounds both operands to FP16, which puts about a quarter of one percent of the diagnostic score values past the existing relative bound, so topk_scores' outlier allowance goes from 0.1% to 0.5%. That is the only gate this change touches; the decode_csa x_out comparator and every cache gate are untouched and pass as they are. a2a3sim segfaults on this kernel: the CPU model of pto.tinsert sizes the scalar pre-quant vector by rows while indexing it by columns, so a tail tile whose valid rows fall below its 64 columns reads out of bounds. Hardware never runs that code. hw-native-sys/pto-isa#334 carries a repro and a one-line fix. Co-Authored-By: Claude Opus 5 (1M context) --- .../decode_indexer.py | 245 ++++++++++++++---- 1 file changed, 197 insertions(+), 48 deletions(-) diff --git a/models/deepseek_v4_flash_dspark/decode_indexer.py b/models/deepseek_v4_flash_dspark/decode_indexer.py index 92dc656e5..992e00ff4 100644 --- a/models/deepseek_v4_flash_dspark/decode_indexer.py +++ b/models/deepseek_v4_flash_dspark/decode_indexer.py @@ -95,14 +95,27 @@ TOPK_QUERY_WORKERS = 48 # Top-K query-merge workers TOPK_ARENA_ROWS = T_PAD * TOPK_ROWS_PER_QUERY TOPK_SCORE_WORKERS = 24 # Top-K score workers -SCORE_TILE = 384 +SCORE_TILE = 256 BUFFERED_SCORE_TILE = 768 BUFFERED_SCORE_LANE_TILE = BUFFERED_SCORE_TILE // 2 BUFFERED_SCORE_SCALE = 1024.0 SCORE_READY_EVENT = 0 SCORE_CONSUMED_EVENT = 1 SCORE_LANE_ROWS = SCORE_TILE // 2 -SCORE_ARENA_ROWS = max(T_PAD, TOPK_SCORE_WORKERS * 2) +# Query tokens scored against one Key tile. Their head blocks are adjacent rows +# of qr_hadamard_i8, so one QK matmul produces both tokens' scores side by side +# and one block-diagonal weight matrix reduces both in a single matmul. T is +# B * S, so an even S pairs for every batch size; an odd S scores one token at +# a time and the block-diagonal weight degenerates to a single row. +SCORE_QUERY_TILE = 2 if S % 2 == 0 else 1 +SCORE_ARENA_ROWS = max(T_PAD, TOPK_SCORE_WORKERS * SCORE_QUERY_TILE * 2) +# Cube score reduction. FIXPIPE drains the INT32 QK accumulator into an FP16 Mat +# tile, applying ReLU on the accumulator and this dequantization scale on the way +# out; 1/1024 keeps the widest INT8x128 dot product inside the FP16 range. The +# head coefficient carries the reciprocal so no post-scale is needed, and the +# 1024x lift keeps small coefficients clear of FP16 subnormals. +SCORE_FIXPIPE_SCALE = 1.0 / 1024 +SCORE_COEF_SCALE = 1024.0 @pl.jit.inline @@ -340,6 +353,9 @@ def indexer_score_topk_forest( qr_hadamard_i8: pl.Tensor[[T_PAD * IDX_N_HEADS, IDX_HEAD_DIM], pl.INT8], qr_hadamard_scale_dq: pl.Tensor[[T_PAD * IDX_N_HEADS, 1], pl.FP32], weights: pl.Tensor[[T_PAD, IDX_N_HEADS], pl.FP32], + score_coefficient: pl.Tensor[ + [T_PAD // SCORE_QUERY_TILE * MM_ROW_TILE, IDX_N_HEADS * SCORE_QUERY_TILE], pl.FP16 + ], idx_kv_cache: pl.Tensor[[IDX_CACHE_BLOCK_NUM_DYN, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8], idx_kv_scale: pl.Tensor[[IDX_CACHE_BLOCK_NUM_DYN, BLOCK_SIZE, 1, 1], pl.FP32], idx_block_table: pl.Tensor[[B_DYN, IDX_MAX_BLOCKS], pl.INT32], @@ -516,11 +532,13 @@ def indexer_score_topk_forest( max_cache_len = pl.max(max_cache_len, batch_cache_len) max_leaves = pl.max((pl.min(max_cache_len, TOPK_MAX_CANDIDATES) + TOPK_CANDIDATES_PER_LEAF - 1) // TOPK_CANDIDATES_PER_LEAF, 1) single_leaf = pl.cast(max_leaves == 1, pl.INDEX) - for item in pl.range(worker, query_count * max_leaves, TOPK_SCORE_WORKERS): - query = item // max_leaves + for item in pl.range(worker, query_count // SCORE_QUERY_TILE * max_leaves, TOPK_SCORE_WORKERS): + query = (item // max_leaves) * SCORE_QUERY_TILE leaf = item % max_leaves batch_idx = query // S - position = pl.read(position_ids, [query]) + # The pair's last token sees the most candidates; each token + # masks its own tail when it stores. + position = pl.read(position_ids, [query + SCORE_QUERY_TILE - 1]) cache_len = pl.read(kv_seq_lens, [batch_idx]) // COMPRESS_RATIO cache_bound = pl.min(cache_len, (position + 1) // COMPRESS_RATIO) visible_count = pl.max(pl.min(cache_bound, TOPK_MAX_CANDIDATES), 0) @@ -534,18 +552,30 @@ def indexer_score_topk_forest( ) lane_stride = single_leaf * SCORE_LANE_ROWS + (1 - single_leaf) * lane_span query_head_begin = query * IDX_N_HEADS - query_vector = qr_hadamard_i8[query_head_begin : query_head_begin + IDX_N_HEADS, 0:IDX_HEAD_DIM] - # Both Vector lanes share the head coefficients for this query. - for _aiv_coeff in pl.split_aiv(2, mode=pl.SplitMode.NONE): - query_scale = pl.reshape( - qr_hadamard_scale_dq[query_head_begin : query_head_begin + IDX_N_HEADS, 0:1], - [1, IDX_N_HEADS], - ) - query_weight = weights[query : query + 1, 0:IDX_N_HEADS] - head_coefficient = pl.reshape(pl.mul(query_scale, query_weight), [IDX_N_HEADS, 1]) + # Both left operands are invariant across this leaf's + # candidate tiles, so they stay resident in L0A instead of + # being reloaded for every tile. + query_vector = pl.tile.load( + qr_hadamard_i8, + [query_head_begin, 0], + [SCORE_QUERY_TILE * IDX_N_HEADS, IDX_HEAD_DIM], + target_memory=pl.Mem.Mat, + ) + query_left = pl.tile.move(query_vector, target_memory=pl.Mem.Left) + coefficient_mat = pl.tile.load( + score_coefficient, + [query // SCORE_QUERY_TILE * MM_ROW_TILE, 0], + [MM_ROW_TILE, SCORE_QUERY_TILE * IDX_N_HEADS], + target_memory=pl.Mem.Mat, + ) + coefficient_left = pl.tile.move( + coefficient_mat, target_memory=pl.Mem.Left + ) for score_begin in pl.pipeline(0, lane_span, SCORE_LANE_ROWS, stage=2): read_begin = score_begin * (1 + single_leaf) - kv_i8 = pl.create_l1([SCORE_TILE, IDX_HEAD_DIM], pl.INT8) + kv_i8 = pl.tile.create( + [SCORE_TILE, IDX_HEAD_DIM], pl.INT8, target_memory=pl.Mem.Mat + ) for page in pl.unroll(SCORE_TILE // BLOCK_SIZE): page_begin = page * BLOCK_SIZE lane_page = (page_begin // SCORE_LANE_ROWS) * lane_stride + page_begin % SCORE_LANE_ROWS @@ -553,17 +583,49 @@ def indexer_score_topk_forest( logical_page = (logical_begin + safe_page_begin) // BLOCK_SIZE physical_block = pl.cast(pl.read(idx_block_table_flat, [batch_idx * IDX_MAX_BLOCKS + logical_page]), pl.INDEX) physical_row = physical_block * BLOCK_SIZE - kv_i8 = pl.gather_row( + kv_i8 = pl.tile.gather_row( kv_i8, kv_cache_i8_flat, [page_begin, 0], [physical_row, 0], [BLOCK_SIZE, IDX_HEAD_DIM], ) - # Reduce over heads with col_sum to avoid the large row_sum UB scratch. - score_i32 = pl.matmul(query_vector, kv_i8, out_dtype=pl.INT32, b_trans=True) - # Each lane keeps all heads and owns a contiguous candidate-column range. + # QK stays on Cube. FIXPIPE then drains the INT32 + # accumulator straight into an FP16 Mat tile, applying + # ReLU and the fixed dequantization in the same + # writeback, so the 64-head intermediate never reaches + # Vector or GM. Pairing two query tokens halves how + # often this Key tile has to be gathered and pushed + # through L0B. + score_i32 = pl.tile.matmul( + query_left, + pl.tile.move( + pl.tile.transpose_view(kv_i8), target_memory=pl.Mem.Right + ), + ) + score_relu = pl.tile.create( + [SCORE_QUERY_TILE * IDX_N_HEADS, SCORE_TILE], + pl.FP16, target_memory=pl.Mem.Mat, + ) + score_relu = pl.tile.assemble( + score_relu, + score_i32, + [0, 0], + pre_quant=SCORE_FIXPIPE_SCALE, + pre_relu=True, + ) + # One Cube matmul replaces the Vector broadcast, + # multiply and head reduction. The coefficient is block + # diagonal, so the pair's two token rows reduce over + # disjoint head ranges in a single GEMM. + score_acc = pl.tile.matmul( + coefficient_left, + pl.tile.move(score_relu, target_memory=pl.Mem.Right), + ) + # Each lane owns a contiguous candidate-column range of + # the reduced score row. for aiv_id in pl.split_aiv(2, mode=pl.SplitMode.LEFT_RIGHT): lane_begin = aiv_id * lane_stride - lane_valid_rows = pl.max(pl.min(valid_count - read_begin - lane_begin, SCORE_LANE_ROWS), 0) - kv_scale = pl.create_tensor([1, SCORE_LANE_ROWS], dtype=pl.FP32) + kv_scale = pl.tile.create( + [1, SCORE_LANE_ROWS], pl.FP32, target_memory=pl.Mem.Vec + ) for scale_page in pl.unroll(SCORE_TILE // (2 * BLOCK_SIZE)): scale_page_begin = scale_page * BLOCK_SIZE safe_scale_begin = pl.min( @@ -575,34 +637,59 @@ def indexer_score_topk_forest( idx_block_table_flat, [batch_idx * IDX_MAX_BLOCKS + scale_logical_page] ), pl.INDEX) scale_physical_row = scale_block * BLOCK_SIZE - kv_scale = pl.gather_row( + kv_scale = pl.tile.gather_row( kv_scale, kv_scale_row, [0, scale_page_begin], [0, scale_physical_row], [1, BLOCK_SIZE], ) - score_shard = pl.aiv_shard(score_i32) - score_fp32 = pl.cast(score_shard, target_type=pl.FP32, mode="none") - score_fp32 = pl.maximum(score_fp32, 0.0) - score_fp32 = pl.row_expand_mul(score_fp32, head_coefficient) - score_sum = pl.col_sum(score_fp32) - score_row = pl.reshape(score_sum, [1, SCORE_LANE_ROWS]) - score_row = pl.mul(score_row, kv_scale) - score_row_id = single_leaf * query + (1 - single_leaf) * (worker * 2 + aiv_id) - score_col = single_leaf * (read_begin + lane_begin) + (1 - single_leaf) * score_begin - # Store only valid scores; the Top-K load pads its own tail. - if lane_valid_rows > 0: - score_valid = pl.set_validshape(score_row, 1, lane_valid_rows) - score_arena[score_row_id : score_row_id + 1, score_col : score_col + SCORE_LANE_ROWS] = score_valid + score_shard = pl.aiv_shard(score_acc) + # Row t of the reduction carries token t of the pair. + for pair_token in pl.unroll(SCORE_QUERY_TILE): + pair_position = pl.read(position_ids, [query + pair_token]) + pair_bound = pl.min(cache_len, (pair_position + 1) // COMPRESS_RATIO) + pair_visible = pl.max(pl.min(pair_bound, TOPK_MAX_CANDIDATES), 0) + pair_valid_rows = pl.max( + pl.min( + pair_visible - logical_begin - read_begin - lane_begin, + SCORE_LANE_ROWS, + ), + 0, + ) + score_row = pl.tile.mul( + pl.tile.slice( + score_shard, [1, SCORE_LANE_ROWS], [pair_token, 0] + ), + kv_scale, + ) + score_row_id = ( + single_leaf * (query + pair_token) + + (1 - single_leaf) + * (worker * SCORE_QUERY_TILE * 2 + pair_token * 2 + aiv_id) + ) + score_col = single_leaf * (read_begin + lane_begin) + (1 - single_leaf) * score_begin + # Store only valid scores; the Top-K load pads its own tail. + if pair_valid_rows > 0: + score_valid = pl.tile.set_validshape( + score_row, 1, pair_valid_rows + ) + pl.tile.store( + score_valid, [score_row_id, score_col], score_arena + ) if single_leaf == 0: for sort_lane in pl.split_aiv(2, mode=pl.SplitMode.NONE): - half_begin = logical_begin + sort_lane * lane_span - half_valid = pl.max(pl.min(valid_count - sort_lane * lane_span, lane_span), 0) - half_slot = query * TOPK_ROWS_PER_QUERY + leaf * 2 + sort_lane - if half_valid > 0: - indexer_topk_half_leaf(score_arena, pair_arena, worker * 2 + sort_lane, half_begin, half_valid, half_slot) - else: - empty_pairs = pl.tile.full([1, TOPK_PAIR_WIDTH], dtype=pl.FP32, value=FP32_NEG_INF) - pl.store(empty_pairs, [half_slot, 0], pair_arena) + for sort_token in pl.unroll(SCORE_QUERY_TILE): + sort_position = pl.read(position_ids, [query + sort_token]) + sort_bound = pl.min(cache_len, (sort_position + 1) // COMPRESS_RATIO) + sort_visible = pl.max(pl.min(sort_bound, TOPK_MAX_CANDIDATES), 0) + half_begin = logical_begin + sort_lane * lane_span + half_valid = pl.max(pl.min(sort_visible - half_begin, lane_span), 0) + half_slot = (query + sort_token) * TOPK_ROWS_PER_QUERY + leaf * 2 + sort_lane + sort_row = worker * SCORE_QUERY_TILE * 2 + sort_token * 2 + sort_lane + if half_valid > 0: + indexer_topk_half_leaf(score_arena, pair_arena, sort_row, half_begin, half_valid, half_slot) + else: + empty_pairs = pl.tile.full([1, TOPK_PAIR_WIDTH], dtype=pl.FP32, value=FP32_NEG_INF) + pl.store(empty_pairs, [half_slot, 0], pair_arena) score_tid = direct_leaf_tid with pl.scope(): @@ -823,21 +910,79 @@ def indexer_weights_score( weights_acc = pl.matmul_acc(weights_acc, x_tile, weights_proj_tile, init_cond=(db == 0)) weights_partial[kb * T_PAD + w_r0 : kb * T_PAD + w_r0 + MM_ROW_TILE, :] = weights_acc + # Left operand of the Cube head reduction, one MM_ROW_TILE fractal per query + # pair. Row t holds token t's head weight times its query dequantization + # scale, in token t's head block; every other entry is zero. This task + # already holds the weight, so the coefficient costs no task of its own. + score_coefficient = pl.create_tensor( + [T_PAD // SCORE_QUERY_TILE * MM_ROW_TILE, IDX_N_HEADS * SCORE_QUERY_TILE], + dtype=pl.FP16, + ) with pl.spmd( row_blocks, name_hint="weights_proj_reduce", + deps=[qh_quant_tid], allow_early_resolve=True, ) as weights_tid: w_rb = pl.tile.get_block_idx() w_r0 = w_rb * MM_ROW_TILE - w_sum = weights_partial[w_r0 : w_r0 + MM_ROW_TILE, :] + w_sum = pl.tile.load( + weights_partial, [w_r0, 0], [MM_ROW_TILE, IDX_N_HEADS], + target_memory=pl.Mem.Vec, + ) for kb in pl.unroll(1, WEIGHTS_OK): partial_r0 = kb * T_PAD + w_r0 - w_sum = pl.add(w_sum, weights_partial[partial_r0 : partial_r0 + MM_ROW_TILE, :]) - weights[w_r0 : w_r0 + MM_ROW_TILE, :] = pl.mul(w_sum, WEIGHTS_SCALE) + w_sum = pl.tile.add( + w_sum, + pl.tile.load( + weights_partial, [partial_r0, 0], [MM_ROW_TILE, IDX_N_HEADS], + target_memory=pl.Mem.Vec, + ), + ) + w_scaled = pl.tile.muls(w_sum, WEIGHTS_SCALE) + pl.tile.store(w_scaled, [w_r0, 0], weights) + # One block-diagonal weight matrix per query pair: token t occupies row t + # and the head block [t*IDX_N_HEADS, (t+1)*IDX_N_HEADS), everything else + # zero, so a single matmul reduces both tokens of the pair. + for w_pair in pl.unroll(MM_ROW_TILE // SCORE_QUERY_TILE): + w_pair_row = (w_r0 // SCORE_QUERY_TILE + w_pair) * MM_ROW_TILE + pl.tile.store( + pl.tile.full( + [MM_ROW_TILE, IDX_N_HEADS * SCORE_QUERY_TILE], + dtype=pl.FP16, value=0.0, + ), + [w_pair_row, 0], + score_coefficient, + ) + for w_token in pl.unroll(SCORE_QUERY_TILE): + w_member = w_pair * SCORE_QUERY_TILE + w_token + w_query = w_r0 + w_member + coefficient_scale = pl.tile.reshape( + pl.tile.load( + qr_hadamard_scale_dq, [w_query * IDX_N_HEADS, 0], [IDX_N_HEADS, 1], + target_memory=pl.Mem.Vec, + ), + [1, IDX_N_HEADS], + ) + coefficient_row = pl.tile.cast( + pl.tile.muls( + pl.tile.mul( + coefficient_scale, + pl.tile.slice(w_scaled, [1, IDX_N_HEADS], [w_member, 0]), + ), + SCORE_COEF_SCALE, + ), + target_type=pl.FP16, + mode="rint", + ) + pl.tile.store( + coefficient_row, + [w_pair_row + w_token, w_token * IDX_N_HEADS], + score_coefficient, + ) topk_scores, topk_idxs, leaf_tid = indexer_score_topk_forest( - qr_hadamard_i8, qr_hadamard_scale_dq, weights, + qr_hadamard_i8, qr_hadamard_scale_dq, weights, score_coefficient, idx_kv_cache, idx_kv_scale, idx_block_table, position_ids, kv_seq_lens, topk_scores, topk_idxs, @@ -1381,7 +1526,11 @@ def init_idx_slot_mapping(): # Scores are diagnostic; sparse attention consumes the selected # indices checked below. A3 reduction order may perturb one # root score per query without changing the selected set. - atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.001 + # The Cube head reduction rounds its two operands to FP16, + # which puts about a quarter of one percent of the scores past + # the relative bound. topk_idxs below keeps main's comparator, + # which admits a swap only between tied scores. + atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.005 ), "topk_idxs": topk_pair_compare("topk_scores"), # C8 cache: history is exact; only the <=B boundary rows the compressor rewrote may From 59571f43fe711bb30a8071cf448e1cdbdaaaa940 Mon Sep 17 00:00:00 2001 From: Youhezhen Date: Mon, 28 Sep 2026 00:41:05 -0700 Subject: [PATCH 2/2] CI: re-run against the pypto orchestration codegen fix The previous run failed native orchestration compilation on every DSpark decode entry, with the emitted C++ assigning from an SSA return-value version it never declared. It was not this change: the same tree passes decode_csa on device against an older pypto, and an unrelated branch failed identically. pypto #2933 preserves Tensor carries across AUTO scopes and is now on pypto main. Co-Authored-By: Claude Opus 5 (1M context)