diff --git a/models/deepseek/v4/compressor_common.py b/models/deepseek/v4/compressor_common.py new file mode 100644 index 00000000..c4d936c8 --- /dev/null +++ b/models/deepseek/v4/compressor_common.py @@ -0,0 +1,270 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""Shared DeepSeek-V4 compressor scheduling and compute primitives.""" + +import pypto.language as pl + +from config import FLASH as M + + +EPS = M.rms_norm_eps +MAX_SEQ_LEN = M.max_position_embeddings +HEAD_DIM = M.head_dim +HEAD_DIM_INV = 1.0 / HEAD_DIM +ROPE_HEAD_DIM = M.qk_rope_head_dim +ROPE_HALF = ROPE_HEAD_DIM // 2 +NOPE_HEAD_DIM = M.nope_head_dim + +COMPRESSOR_RMS_ROW_TILE = 8 +COMPRESSOR_HEAD_TILE = 64 +COMPRESSOR_RMS_ROWS = pl.dynamic("COMPRESSOR_COMMON_RMS_ROWS") +COMPRESSOR_FINALIZE_ROWS = pl.dynamic("COMPRESSOR_FINALIZE_ROWS") +COMPRESSOR_FINALIZE_CACHE_ROWS = pl.dynamic("COMPRESSOR_FINALIZE_CACHE_ROWS") +SCHEDULE_TOKENS = pl.dynamic("COMPRESSOR_SCHEDULE_TOKENS") +SCHEDULE_WRITES = pl.dynamic("COMPRESSOR_SCHEDULE_WRITES") +SCHEDULE_ROPE_ROWS = pl.dynamic("COMPRESSOR_SCHEDULE_ROPE_ROWS") + + +@pl.jit.inline +def build_prefill_write_schedule( + position_ids: pl.Tensor[[SCHEDULE_TOKENS], pl.INT32], + cmp_slot_mapping: pl.Tensor[[SCHEDULE_TOKENS], pl.INT64], + num_tokens: pl.Scalar[pl.INT32], + write_pos_map: pl.Tensor[[1, SCHEDULE_WRITES], pl.INT32], + write_dst_map: pl.Tensor[[1, SCHEDULE_WRITES], pl.INT32], + state_table_row_map: pl.Tensor[[1, SCHEDULE_WRITES], pl.INT32], +): + token_rows = pl.tensor.dim(position_ids, 0) + write_rows = pl.tensor.dim(write_pos_map, 1) + with pl.at(level=pl.Level.CORE_GROUP, name_hint="prefill_cmp_write_schedule"): + for init_i in pl.range(write_rows): + pl.write(write_pos_map, [0, init_i], pl.cast(0, pl.INT32)) + pl.write(write_dst_map, [0, init_i], pl.cast(-1, pl.INT32)) + pl.write(state_table_row_map, [0, init_i], pl.cast(-1, pl.INT32)) + map_seen = pl.cast(0, pl.INDEX) + for map_t in pl.range(token_rows): + if map_t < num_tokens: + map_slot_raw = pl.read(cmp_slot_mapping, [map_t]) + if map_slot_raw >= 0: + if map_seen < write_rows: + pl.write(write_pos_map, [0, map_seen], pl.read(position_ids, [map_t])) + pl.write(write_dst_map, [0, map_seen], pl.cast(map_slot_raw, pl.INT32)) + pl.write(state_table_row_map, [0, map_seen], pl.cast(0, pl.INT32)) + map_seen = map_seen + 1 + + +@pl.jit.inline +def build_decode_padded_write_schedule( + position_ids: pl.Tensor[[SCHEDULE_TOKENS], pl.INT32], + cmp_slot_mapping: pl.Tensor[[SCHEDULE_TOKENS], pl.INT64], + seq_len: pl.Scalar[pl.INT32], + compress_ratio: pl.Scalar[pl.INT32], + rms_tile: pl.Scalar[pl.INT32], + rms_pad_tile: pl.Scalar[pl.INT32], + write_pos_map: pl.Tensor[[1, SCHEDULE_WRITES], pl.INT32], + write_dst_map: pl.Tensor[[1, SCHEDULE_WRITES], pl.INT32], + kv_out_row_map: pl.Tensor[[1, SCHEDULE_WRITES], pl.INT32], + state_table_row_map: pl.Tensor[[1, SCHEDULE_WRITES], pl.INT32], +): + token_rows = pl.tensor.dim(position_ids, 0) + write_rows = pl.tensor.dim(write_pos_map, 1) + seq_len_i = pl.cast(seq_len, pl.INDEX) + compress_ratio_i = pl.cast(compress_ratio, pl.INDEX) + rms_tile_i = pl.cast(rms_tile, pl.INDEX) + rms_pad_tile_i = pl.cast(rms_pad_tile, pl.INDEX) + batch_rows = token_rows // seq_len_i + with pl.at(level=pl.Level.CORE_GROUP, name_hint="decode_cmp_write_schedule"): + for init_i in pl.range(write_rows): + pl.write(write_pos_map, [0, init_i], pl.cast(0, pl.INT32)) + pl.write(write_dst_map, [0, init_i], pl.cast(-1, pl.INT32)) + pl.write(kv_out_row_map, [0, init_i], pl.cast(-1, pl.INT32)) + pl.write(state_table_row_map, [0, init_i], pl.cast(-1, pl.INT32)) + for b in pl.range(batch_rows): + base_t = b * seq_len_i + first_pos = pl.read(position_ids, [base_t]) + pos_in_window = pl.cast(first_pos % compress_ratio, pl.INDEX) + if pos_in_window + seq_len_i >= compress_ratio_i: + boundary_s = compress_ratio_i - 1 - pos_in_window + token_t = base_t + boundary_s + dst_raw = pl.read(cmp_slot_mapping, [token_t]) + if dst_raw >= 0: + pad_row = (b // rms_tile_i) * rms_pad_tile_i + (b % rms_tile_i) + if pad_row < write_rows: + pl.write( + write_pos_map, + [0, pad_row], + pl.cast(first_pos + pl.cast(boundary_s, pl.INT32), pl.INT32), + ) + pl.write(write_dst_map, [0, pad_row], pl.cast(dst_raw, pl.INT32)) + pl.write(kv_out_row_map, [0, pad_row], pl.cast(base_t, pl.INT32)) + pl.write(state_table_row_map, [0, pad_row], pl.cast(b, pl.INT32)) + + +@pl.jit.inline +def gather_compressor_rope_rows( + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + write_pos_map: pl.Tensor[[1, SCHEDULE_WRITES], pl.INT32], + write_dst_map: pl.Tensor[[1, SCHEDULE_WRITES], pl.INT32], + compress_ratio: pl.Scalar[pl.INT32], + cos_b: pl.Tensor[[SCHEDULE_ROPE_ROWS, ROPE_HALF], pl.FP32], + sin_b: pl.Tensor[[SCHEDULE_ROPE_ROWS, ROPE_HALF], pl.FP32], +): + write_rows = pl.tensor.dim(write_pos_map, 1) + rope_rows = pl.tensor.dim(cos_b, 0) + with pl.at(level=pl.Level.CORE_GROUP, name_hint="cmp_rope_schedule"): + for rope_i in pl.range(rope_rows): + cos_b[rope_i : rope_i + 1, 0:ROPE_HALF] = pl.full([1, ROPE_HALF], dtype=pl.FP32, value=0.0) + sin_b[rope_i : rope_i + 1, 0:ROPE_HALF] = pl.full([1, ROPE_HALF], dtype=pl.FP32, value=0.0) + if rope_i < write_rows: + write_slot_raw = pl.read(write_dst_map, [0, rope_i]) + if write_slot_raw >= 0: + cmp_pos = pl.cast(pl.read(write_pos_map, [0, rope_i]) + 1 - compress_ratio, pl.INDEX) + cos_b[rope_i : rope_i + 1, 0:ROPE_HALF] = pl.cast( + freqs_cos[cmp_pos : cmp_pos + 1, 0:ROPE_HALF], + target_type=pl.FP32, + ) + sin_b[rope_i : rope_i + 1, 0:ROPE_HALF] = pl.cast( + freqs_sin[cmp_pos : cmp_pos + 1, 0:ROPE_HALF], + target_type=pl.FP32, + ) + + +@pl.jit.inline +def compressor_rmsnorm_rope( + pooled_kv: pl.Tensor[[COMPRESSOR_RMS_ROWS, HEAD_DIM], pl.FP32], + norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], + cos_b: pl.Tensor[[COMPRESSOR_RMS_ROWS, ROPE_HALF], pl.FP32], + sin_b: pl.Tensor[[COMPRESSOR_RMS_ROWS, ROPE_HALF], pl.FP32], + normed_kv: pl.Tensor[[COMPRESSOR_RMS_ROWS, HEAD_DIM], pl.FP32], +): + norm_w_2d = pl.reshape(norm_w, [1, HEAD_DIM]) + rows = pl.tensor.dim(pooled_kv, 0) + for rt in pl.spmd(rows // COMPRESSOR_RMS_ROW_TILE, name_hint="rmsnorm_rope"): + r0 = rt * COMPRESSOR_RMS_ROW_TILE + partial_sq = pl.full([1, COMPRESSOR_RMS_ROW_TILE], dtype=pl.FP32, value=0.0) + for rms_kb in pl.pipeline(HEAD_DIM // COMPRESSOR_HEAD_TILE, stage=2): + rms_h0 = rms_kb * COMPRESSOR_HEAD_TILE + kv_rms_chunk = pooled_kv[ + r0 : r0 + COMPRESSOR_RMS_ROW_TILE, + rms_h0 : rms_h0 + COMPRESSOR_HEAD_TILE, + ] + kv_rms_sq = pl.mul(kv_rms_chunk, kv_rms_chunk) + partial_sq = pl.add( + partial_sq, + pl.reshape(pl.row_sum(kv_rms_sq), [1, COMPRESSOR_RMS_ROW_TILE]), + ) + variance = pl.reshape( + pl.add(pl.mul(partial_sq, HEAD_DIM_INV), EPS), + [COMPRESSOR_RMS_ROW_TILE, 1], + ) + inv_rms = pl.recip(pl.sqrt(variance)) + for rms_kb in pl.pipeline(NOPE_HEAD_DIM // COMPRESSOR_HEAD_TILE, stage=2): + norm_h0 = rms_kb * COMPRESSOR_HEAD_TILE + kv_norm_chunk = pooled_kv[ + r0 : r0 + COMPRESSOR_RMS_ROW_TILE, + norm_h0 : norm_h0 + COMPRESSOR_HEAD_TILE, + ] + gamma = pl.cast( + norm_w_2d[:, norm_h0 : norm_h0 + COMPRESSOR_HEAD_TILE], + pl.FP32, + ) + normed_chunk = pl.col_expand_mul( + pl.row_expand_mul(kv_norm_chunk, inv_rms), + gamma, + ) + normed_kv[ + r0 : r0 + COMPRESSOR_RMS_ROW_TILE, + norm_h0 : norm_h0 + COMPRESSOR_HEAD_TILE, + ] = normed_chunk + + kv_rope_norm = pooled_kv[ + r0 : r0 + COMPRESSOR_RMS_ROW_TILE, + NOPE_HEAD_DIM:HEAD_DIM, + ] + gamma_rope = pl.cast(norm_w_2d[:, NOPE_HEAD_DIM:HEAD_DIM], pl.FP32) + rope_normed = pl.col_expand_mul(pl.row_expand_mul(kv_rope_norm, inv_rms), gamma_rope) + rope_ones = pl.full( + [COMPRESSOR_RMS_ROW_TILE, ROPE_HEAD_DIM], + dtype=pl.FP32, + value=1.0, + ) + rope_col = pl.col_expand_mul( + rope_ones, + pl.cast(pl.arange(0, [1, ROPE_HEAD_DIM], dtype=pl.INT32), target_type=pl.FP32), + ) + rope_dup_f = pl.cast( + pl.cast(pl.mul(rope_col, 0.5), target_type=pl.INT32, mode="trunc"), + target_type=pl.FP32, + ) + rope_dup_idx = pl.cast(rope_dup_f, target_type=pl.INT32) + rope_lane = pl.sub(rope_col, pl.mul(rope_dup_f, 2.0)) + rope_swap_idx = pl.cast( + pl.sub(pl.add(rope_col, 1.0), pl.mul(rope_lane, 2.0)), + target_type=pl.INT32, + ) + rope_sign = pl.sub(pl.mul(rope_lane, 2.0), 1.0) + cos_il = pl.gather( + cos_b[r0 : r0 + COMPRESSOR_RMS_ROW_TILE, 0:ROPE_HALF], + dim=-1, + index=rope_dup_idx, + ) + sin_il = pl.gather( + sin_b[r0 : r0 + COMPRESSOR_RMS_ROW_TILE, 0:ROPE_HALF], + dim=-1, + index=rope_dup_idx, + ) + swapped = pl.gather(rope_normed, dim=-1, index=rope_swap_idx) + rope_rot = pl.add( + pl.mul(rope_normed, cos_il), + pl.mul(pl.mul(swapped, rope_sign), sin_il), + ) + normed_kv[ + r0 : r0 + COMPRESSOR_RMS_ROW_TILE, + NOPE_HEAD_DIM:HEAD_DIM, + ] = rope_rot + return normed_kv + + +@pl.jit.inline +def finalize_compressor_writes( + normed_kv: pl.Tensor[[COMPRESSOR_FINALIZE_ROWS, HEAD_DIM], pl.FP32], + write_dst_map: pl.Tensor[[1, COMPRESSOR_FINALIZE_ROWS], pl.INT32], + cmp_kv_cache_flat: pl.Tensor[[COMPRESSOR_FINALIZE_CACHE_ROWS, HEAD_DIM], pl.BF16], + rows_per_task: pl.Scalar[pl.INT32], + keepalive_invalid: pl.Scalar[pl.INT32], +): + """Write scheduled compressed rows while preserving each caller's task tiling.""" + write_rows = pl.tensor.dim(write_dst_map, 1) + cache_rows = pl.tensor.dim(cmp_kv_cache_flat, 0) + row_tile = pl.cast(rows_per_task, pl.INDEX) + for final_block in pl.spmd((write_rows + row_tile - 1) // row_tile, name_hint="compressor_cache_write"): + final_base = final_block * row_tile + for final_dt in pl.range(row_tile): + final_row = final_base + final_dt + if final_row < write_rows: + dst_row_raw = pl.read(write_dst_map, [0, final_row]) + if dst_row_raw >= 0: + dst_row = pl.cast(dst_row_raw, pl.INDEX) + cmp_kv_cache_flat[dst_row : dst_row + 1, 0:HEAD_DIM] = pl.cast( + normed_kv[final_row : final_row + 1, 0:HEAD_DIM], + target_type=pl.BF16, + mode="rint", + ) + else: + if keepalive_invalid != 0: + keepalive_row = cache_rows - write_rows + final_row + cmp_kv_cache_flat[ + keepalive_row : keepalive_row + 1, + 0:HEAD_DIM, + ] = cmp_kv_cache_flat[ + keepalive_row : keepalive_row + 1, + 0:HEAD_DIM, + ] + return cmp_kv_cache_flat diff --git a/models/deepseek/v4/compressor_ratio128.py b/models/deepseek/v4/compressor_ratio128.py new file mode 100644 index 00000000..c6694dbb --- /dev/null +++ b/models/deepseek/v4/compressor_ratio128.py @@ -0,0 +1,997 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""DeepSeek-V4 ratio-128 compressor decode and prefill paths.""" + +import pypto.language as pl + +from config import ( + FLASH as M, + BLOCK_SIZE, + C128_COMPRESSOR_BLOCK_SIZE, + DECODE_BATCH, + DECODE_SEQ, + DECODE_CMP_BLOCK_NUM, + FP32_NEG_INF, + KV_CMP_MAX_BLOCKS, + PREFILL_BATCH, + PREFILL_SEQ, + PREFILL_CMP_BLOCK_NUM, + PREFILL_CMP_MAX_BLOCKS, +) +from compressor_common import ( + COMPRESSOR_RMS_ROW_TILE as RMS_ROW_TILE, + build_decode_padded_write_schedule, + build_prefill_write_schedule, + compressor_rmsnorm_rope, + finalize_compressor_writes, + gather_compressor_rope_rows, +) + + +EPS = M.rms_norm_eps +D = M.hidden_size +HEAD_DIM = M.head_dim +ROPE_HEAD_DIM = M.qk_rope_head_dim +ROPE_HALF = ROPE_HEAD_DIM // 2 +NOPE_HEAD_DIM = M.nope_head_dim +MAX_SEQ_LEN = M.max_position_embeddings + +COMPRESS_RATIO = 128 +OUT_DIM = HEAD_DIM +STATE_LEN = COMPRESS_RATIO +COMPRESS_STATE_DIM = 2 * OUT_DIM +POOL_HEAD_TILE = 128 +RATIO128_STATE_BLOCK_SIZE = C128_COMPRESSOR_BLOCK_SIZE + +# Decode shape and paging contract. +DECODE_B = DECODE_BATCH +DECODE_S = DECODE_SEQ +DECODE_T = DECODE_B * DECODE_S +DECODE_IDX_KV_LEN = MAX_SEQ_LEN // COMPRESS_RATIO +DECODE_COMPRESS_STATE_BLOCK_SIZE = RATIO128_STATE_BLOCK_SIZE +DECODE_COMPRESS_STATE_PHYSICAL_BLOCKS = 64 +DECODE_COMPRESS_STATE_MAX_BLOCKS = (MAX_SEQ_LEN + DECODE_COMPRESS_STATE_BLOCK_SIZE - 1) // DECODE_COMPRESS_STATE_BLOCK_SIZE +DECODE_COMPRESS_STATE_BLOCK_NUM = DECODE_B * DECODE_COMPRESS_STATE_PHYSICAL_BLOCKS +DECODE_COMPRESSOR_CMP_MAX_BLOCKS = KV_CMP_MAX_BLOCKS +DECODE_COMPRESSOR_CMP_BLOCK_NUM = DECODE_CMP_BLOCK_NUM +if DECODE_IDX_KV_LEN > DECODE_COMPRESSOR_CMP_MAX_BLOCKS * BLOCK_SIZE: + raise ValueError("ratio128 compressed KV cache capacity is smaller than max compressed sequence length") + +# Decode tiling. +DECODE_ROPE_TILE = 32 +DECODE_K_TILE = 512 +DECODE_OUT_TILE = 64 +DECODE_HEAD_TILE = 64 +DECODE_B_TILE = 8 +DECODE_MM_B_TILE = 16 +DECODE_BS_PAD = ((DECODE_B * DECODE_S + DECODE_MM_B_TILE - 1) // DECODE_MM_B_TILE) * DECODE_MM_B_TILE +DECODE_RMS_TILE = 4 +DECODE_RMS_PAD_TILE = 16 +DECODE_RMS_PAD_TAIL = DECODE_RMS_PAD_TILE - DECODE_RMS_TILE +DECODE_RMS_PAD_ROWS = (DECODE_B // DECODE_RMS_TILE) * DECODE_RMS_PAD_TILE +DECODE_POOL_HEAD_TILE = 128 + +# Prefill shape and paging contract. +PREFILL_B = PREFILL_BATCH +PREFILL_S = PREFILL_SEQ +PREFILL_T = PREFILL_B * PREFILL_S +PREFILL_START_POS = 0 + +K_TILE = 512 +OUT_TILE = 32 # prefill (large M): finer tiles fill the array +HEAD_TILE = 64 +K_BLOCKS = D // K_TILE +OUT_BLOCKS = OUT_DIM // OUT_TILE +HEAD_BLOCKS = HEAD_DIM // HEAD_TILE + +assert PREFILL_S == COMPRESS_RATIO, "ratio128 prefill compressor bring-up expects one full compression chunk" + +HCA_STATE_BLOCK_SIZE = RATIO128_STATE_BLOCK_SIZE +HCA_STATE_MAX_BLOCKS = (MAX_SEQ_LEN + HCA_STATE_BLOCK_SIZE - 1) // HCA_STATE_BLOCK_SIZE +HCA_STATE_BLOCK_NUM = HCA_STATE_MAX_BLOCKS +MAX_CMP_WRITES = max(1, PREFILL_T // COMPRESS_RATIO) +HCA_CMP_MAX_BLOCKS = PREFILL_CMP_MAX_BLOCKS +HCA_CMP_BLOCK_NUM = PREFILL_CMP_BLOCK_NUM +HCA_KV_STORE_TILE = 16 +HCA_C128_RMS_TILE = 8 +HCA_C128_RMS_PAD_ROWS = HCA_C128_RMS_TILE + +PACKED_C128_PROJ_BLOCKS = OUT_BLOCKS +POOL_HEAD_BLOCKS = HEAD_DIM // POOL_HEAD_TILE +PACKED_C128_POOL_BLOCKS = MAX_CMP_WRITES * POOL_HEAD_BLOCKS + + +# Shared ratio128 sub-kernel tiling. +PROJ_ROWS = pl.dynamic("COMPRESSOR_PROJ_ROWS") +PROJ_ROWS_PAD = pl.dynamic("COMPRESSOR_PROJ_ROWS_PAD") +PROJ_MM_B_TILE = 16 +PROJ_OUT_TILE = 64 # decode (small M): coarse tiles avoid dispatch overhead +PROJ_K_TILE = 512 +POOL_STATE_ROWS = pl.dynamic("COMPRESSOR128_POOL_STATE_ROWS") +POOL_TABLE_ROWS = pl.dynamic("COMPRESSOR128_POOL_TABLE_ROWS") +POOL_TABLE_BLOCKS = pl.dynamic("COMPRESSOR128_POOL_TABLE_BLOCKS") +POOL_STATE_BLOCKS = STATE_LEN // RATIO128_STATE_BLOCK_SIZE + +# Shared-core dynamic shapes (bind per caller: decode vs prefill). State tensors +# are passed as pre-reshaped flat views so their shapes stay statically inferable +# per call site (dim()-derived reshapes inside the core are not). +CORE_PROJ_ROWS = pl.dynamic("COMPRESSOR_CORE_PROJ_ROWS") +CORE_TOKENS = pl.dynamic("COMPRESSOR_CORE_TOKENS") +CORE_WRITE_ROWS = pl.dynamic("COMPRESSOR_CORE_WRITE_ROWS") +CORE_STATE_ROWS = pl.dynamic("COMPRESSOR_CORE_STATE_ROWS") +CORE_TABLE_ROWS = pl.dynamic("COMPRESSOR_CORE_TABLE_ROWS") +CORE_TABLE_BLOCKS = pl.dynamic("COMPRESSOR_CORE_TABLE_BLOCKS") +# Compact per-regime pool enumeration length: decode binds it to the batch-row +# count (real windows), prefill to its write-row count. Pooling only the real +# windows — instead of every RMS-padded write row — keeps decode off the padded +# grid that otherwise inflates its pool/init dispatch. +CORE_POOL_ROWS = pl.dynamic("COMPRESSOR_CORE_POOL_ROWS") +POOL_HEAD_BLOCKS_CORE = HEAD_DIM // POOL_HEAD_TILE + + +@pl.jit.inline +def compressor_ratio128_pool_math( + score_state: pl.Tensor[[STATE_LEN, POOL_HEAD_TILE], pl.FP32], + kv_state: pl.Tensor[[STATE_LEN, POOL_HEAD_TILE], pl.FP32], +): + score_max = pl.col_max(score_state) + score_exp = pl.col_expand_expdif(score_state, score_max) + score_sum = pl.col_sum(score_exp) + score_prob = pl.col_expand_mul(score_exp, pl.recip(score_sum)) + pooled_chunk = pl.col_sum(pl.mul(kv_state, score_prob)) + return pooled_chunk + + +@pl.jit.inline +def compressor_ratio128_pool_window( + compress_state_rows: pl.Tensor[[POOL_STATE_ROWS, COMPRESS_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[POOL_TABLE_ROWS, POOL_TABLE_BLOCKS], pl.INT32], + table_row: pl.Scalar[pl.INDEX], + write_pos: pl.Scalar[pl.INT32], + h0: pl.Scalar[pl.INDEX], +): + softmax_score_state = pl.create_tensor([STATE_LEN, POOL_HEAD_TILE], dtype=pl.FP32) + softmax_kv_state = pl.create_tensor([STATE_LEN, POOL_HEAD_TILE], dtype=pl.FP32) + state_pos0 = write_pos + 1 - COMPRESS_RATIO + base_logical_blk = pl.cast(state_pos0 // RATIO128_STATE_BLOCK_SIZE, pl.INDEX) + for blk_i in pl.pipeline(POOL_STATE_BLOCKS, stage=2): + s0 = blk_i * RATIO128_STATE_BLOCK_SIZE + slot_score = pl.full([RATIO128_STATE_BLOCK_SIZE, POOL_HEAD_TILE], dtype=pl.FP32, value=FP32_NEG_INF) + slot_kv = pl.full([RATIO128_STATE_BLOCK_SIZE, POOL_HEAD_TILE], dtype=pl.FP32, value=0.0) + state_blk_raw = pl.read(compress_state_block_table, [table_row, base_logical_blk + blk_i]) + if state_blk_raw >= 0: + state_blk_id = pl.cast(state_blk_raw, target_type=pl.INDEX) + row0 = state_blk_id * RATIO128_STATE_BLOCK_SIZE + slot_score = compress_state_rows[ + row0 : row0 + RATIO128_STATE_BLOCK_SIZE, + OUT_DIM + h0 : OUT_DIM + h0 + POOL_HEAD_TILE, + ] + slot_kv = compress_state_rows[ + row0 : row0 + RATIO128_STATE_BLOCK_SIZE, + h0 : h0 + POOL_HEAD_TILE, + ] + softmax_score_state[s0 : s0 + RATIO128_STATE_BLOCK_SIZE, :] = slot_score + softmax_kv_state[s0 : s0 + RATIO128_STATE_BLOCK_SIZE, :] = slot_kv + return compressor_ratio128_pool_math(softmax_score_state, softmax_kv_state) + + +@pl.jit.inline +def compressor_ratio128_proj( + x: pl.Tensor[[PROJ_ROWS, D], pl.BF16], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + kv_proj_out: pl.Tensor[[PROJ_ROWS_PAD, OUT_DIM], pl.FP32], + score_proj_out: pl.Tensor[[PROJ_ROWS_PAD, OUT_DIM], pl.FP32], +): + t_dim = pl.tensor.dim(x, 0) + t_matmul = pl.tensor.dim(kv_proj_out, 0) + for idx in pl.spmd(t_matmul * OUT_DIM // (PROJ_MM_B_TILE * PROJ_OUT_TILE), name_hint="kv_score_proj"): + global_row0 = (idx // (OUT_DIM // PROJ_OUT_TILE)) * PROJ_MM_B_TILE + o0 = (idx % (OUT_DIM // PROJ_OUT_TILE)) * PROJ_OUT_TILE + kv_acc = pl.create_tensor([PROJ_MM_B_TILE, PROJ_OUT_TILE], dtype=pl.FP32) + score_acc = pl.create_tensor([PROJ_MM_B_TILE, PROJ_OUT_TILE], dtype=pl.FP32) + for kb in pl.pipeline(0, D // PROJ_K_TILE, stage=2): + k0 = kb * PROJ_K_TILE + x_rows = pl.min(PROJ_MM_B_TILE, t_dim - global_row0) + x_tile = pl.slice(x, [PROJ_MM_B_TILE, PROJ_K_TILE], [global_row0, k0], valid_shape=[x_rows, PROJ_K_TILE]) + wkv_tile = wkv[o0 : o0 + PROJ_OUT_TILE, k0 : k0 + PROJ_K_TILE] + wgate_tile = wgate[o0 : o0 + PROJ_OUT_TILE, k0 : k0 + PROJ_K_TILE] + if k0 == 0: + kv_acc = pl.matmul(x_tile, wkv_tile, out_dtype=pl.FP32, b_trans=True) + score_acc = pl.matmul(x_tile, wgate_tile, out_dtype=pl.FP32, b_trans=True) + else: + kv_acc = pl.matmul_acc(kv_acc, x_tile, wkv_tile, b_trans=True) + score_acc = pl.matmul_acc(score_acc, x_tile, wgate_tile, b_trans=True) + kv_proj_out[global_row0 : global_row0 + PROJ_MM_B_TILE, o0 : o0 + PROJ_OUT_TILE] = kv_acc + score_proj_out[global_row0 : global_row0 + PROJ_MM_B_TILE, o0 : o0 + PROJ_OUT_TILE] = score_acc + + +@pl.jit.inline +def prefill_compressor_ratio128_proj( + x: pl.Tensor[[PREFILL_T, D], pl.BF16], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + kv_proj_out: pl.Tensor[[PREFILL_T, OUT_DIM], pl.FP32], + score_proj_out: pl.Tensor[[PREFILL_T, OUT_DIM], pl.FP32], +): + for proj_idx in pl.spmd(PACKED_C128_PROJ_BLOCKS, name_hint="prefill_hca_c128_kv_score_proj"): + o0 = proj_idx * OUT_TILE + kv_acc = pl.create_tensor([PREFILL_T, OUT_TILE], dtype=pl.FP32) + score_acc = pl.create_tensor([PREFILL_T, OUT_TILE], dtype=pl.FP32) + for kb in pl.pipeline(0, K_BLOCKS, stage=2): + k0 = kb * K_TILE + x_tile = x[0:PREFILL_T, k0 : k0 + K_TILE] + wkv_tile = wkv[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] + wgate_tile = wgate[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] + if k0 == 0: + kv_acc = pl.matmul(x_tile, wkv_tile, out_dtype=pl.FP32, b_trans=True) + score_acc = pl.matmul(x_tile, wgate_tile, out_dtype=pl.FP32, b_trans=True) + else: + kv_acc = pl.matmul_acc(kv_acc, x_tile, wkv_tile, b_trans=True) + score_acc = pl.matmul_acc(score_acc, x_tile, wgate_tile, b_trans=True) + kv_proj_out[0:PREFILL_T, o0 : o0 + OUT_TILE] = kv_acc + score_proj_out[0:PREFILL_T, o0 : o0 + OUT_TILE] = score_acc + + +@pl.jit.inline +def compressor_core_ratio128( + kv_proj: pl.Tensor[[CORE_PROJ_ROWS, OUT_DIM], pl.FP32], + score_proj: pl.Tensor[[CORE_PROJ_ROWS, OUT_DIM], pl.FP32], + position_ids: pl.Tensor[[CORE_TOKENS], pl.INT32], + state_slot_mapping: pl.Tensor[[CORE_TOKENS], pl.INT64], + ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], + norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], + compress_state_rows: pl.Tensor[[CORE_STATE_ROWS, COMPRESS_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[CORE_TABLE_ROWS, CORE_TABLE_BLOCKS], pl.INT32], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + write_pos_map: pl.Tensor[[1, CORE_WRITE_ROWS], pl.INT32], + write_dst_map: pl.Tensor[[1, CORE_WRITE_ROWS], pl.INT32], + state_table_row_map: pl.Tensor[[1, CORE_WRITE_ROWS], pl.INT32], + pool_row_map: pl.Tensor[[1, CORE_POOL_ROWS], pl.INT32], + pooled_kv: pl.Tensor[[CORE_WRITE_ROWS, HEAD_DIM], pl.FP32], + normed_kv: pl.Tensor[[CORE_WRITE_ROWS, HEAD_DIM], pl.FP32], + cos_b: pl.Tensor[[CORE_WRITE_ROWS, ROPE_HALF], pl.FP32], + sin_b: pl.Tensor[[CORE_WRITE_ROWS, ROPE_HALF], pl.FP32], +): + """Shared decode/prefill ratio128 compression math pipeline. + + Regime-specific projection tiling, write scheduling, cache finalization, and + optional per-token output remain in the caller. The core scatters projected + state, pools every complete window, applies RMSNorm and RoPE, then returns the + per-write FP32 rows in `normed_kv`.""" + token_rows = pl.tensor.dim(state_slot_mapping, 0) + write_rows = pl.tensor.dim(write_dst_map, 1) + + # 1. Scatter projected (kv, score+APE) into the paged state buffer. Padding / + # non-writing tokens carry state_slot_mapping < 0 and are skipped (uniform + # decode/prefill validity contract; num_tokens is folded into the schedule). + with pl.spmd(token_rows, name_hint="state_scatter_pre") as scatter_tid: + scatter_t = pl.tile.get_block_idx() + state_row_i64 = pl.read(state_slot_mapping, [scatter_t]) + if state_row_i64 >= 0: + state_row = pl.cast(state_row_i64, target_type=pl.INDEX) + token_pos = pl.read(position_ids, [scatter_t]) + token_ape_row = pl.cast(token_pos % COMPRESS_RATIO, target_type=pl.INDEX) + ape_row = ape[token_ape_row : token_ape_row + 1, 0:OUT_DIM] + kv_row = kv_proj[scatter_t : scatter_t + 1, 0:OUT_DIM] + score_row = pl.add(score_proj[scatter_t : scatter_t + 1, 0:OUT_DIM], ape_row) + compress_state_rows[state_row : state_row + 1, 0:OUT_DIM] = kv_row + compress_state_rows[state_row : state_row + 1, OUT_DIM:COMPRESS_STATE_DIM] = score_row + + # 2. Zero the pooled scratch (invalid write rows must not feed rmsnorm garbage). + # Coarse RMS_ROW_TILE-row full-width tiles: write_rows is a multiple of + # RMS_ROW_TILE (rmsnorm tiles the same rows), so this covers every row exactly + # while keeping the dispatch count tiny. Real rows are overwritten by the pool + # below (ordered via init_tid dep), so zeroing them first is harmless. + with pl.spmd(write_rows // RMS_ROW_TILE, name_hint="pooled_pad_init") as init_tid: + init_r0 = pl.tile.get_block_idx() * RMS_ROW_TILE + pooled_kv[init_r0 : init_r0 + RMS_ROW_TILE, 0:HEAD_DIM] = pl.full( + [RMS_ROW_TILE, HEAD_DIM], dtype=pl.FP32, value=0.0 + ) + + # 3. Softmax-pool each complete window from the paged state. Iterate the compact + # pool_row_map (one entry per real window) rather than every padded write row: + # pool_row_map[p] -> padded write row, from which write_dst/write_pos/table_row + # are read exactly as the padded grid would. Decode's scattered RMS padding thus + # costs no extra pool tasks; prefill passes an identity map (no change). + pool_rows = pl.tensor.dim(pool_row_map, 1) + with pl.spmd(pool_rows * POOL_HEAD_BLOCKS_CORE, name_hint="softmax_pool", deps=[scatter_tid, init_tid]) as pool_tid: + idx = pl.tile.get_block_idx() + pool_p = idx // POOL_HEAD_BLOCKS_CORE + h0 = (idx % POOL_HEAD_BLOCKS_CORE) * POOL_HEAD_TILE + write_row = pl.cast(pl.read(pool_row_map, [0, pool_p]), target_type=pl.INDEX) + write_slot_raw = pl.read(write_dst_map, [0, write_row]) + if write_slot_raw >= 0: + write_pos = pl.read(write_pos_map, [0, write_row]) + table_row = pl.cast(pl.read(state_table_row_map, [0, write_row]), target_type=pl.INDEX) + pooled_chunk = compressor_ratio128_pool_window( + compress_state_rows, + compress_state_block_table, + table_row, + write_pos, + h0, + ) + pooled_kv[write_row : write_row + 1, h0 : h0 + POOL_HEAD_TILE] = pooled_chunk + + # 4. RoPE tables for each window position, then rmsnorm + rope. + gather_compressor_rope_rows( + freqs_cos, + freqs_sin, + write_pos_map, + write_dst_map, + pl.const(COMPRESS_RATIO, pl.INT32), + cos_b, + sin_b, + ) + normed_kv = compressor_rmsnorm_rope(pooled_kv, norm_w, cos_b, sin_b, normed_kv) + return normed_kv + + +@pl.jit.inline +def decode_compressor_ratio128( + x: pl.Tensor[[DECODE_T, D], pl.BF16], + kv: pl.Tensor[[DECODE_T, HEAD_DIM], pl.FP32], + compress_state: pl.Tensor[[DECODE_COMPRESS_STATE_BLOCK_NUM, DECODE_COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[DECODE_B, DECODE_COMPRESS_STATE_MAX_BLOCKS], pl.INT32], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], + norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + cmp_kv_cache: pl.Tensor[[DECODE_COMPRESSOR_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16], + position_ids: pl.Tensor[[DECODE_T], pl.INT32], + cmp_slot_mapping: pl.Tensor[[DECODE_T], pl.INT64], + state_slot_mapping: pl.Tensor[[DECODE_T], pl.INT64], +): + # Thin decode wrapper: build the padded per-batch write schedule, project + # (dynamic-M tiling), then run the shared core with the paged decode state + # and per-token kv output. See compressor_core_ratio128. + write_pos_map = pl.create_tensor([1, DECODE_RMS_PAD_ROWS], dtype=pl.INT32) + write_dst_map = pl.create_tensor([1, DECODE_RMS_PAD_ROWS], dtype=pl.INT32) + kv_out_row_map = pl.create_tensor([1, DECODE_RMS_PAD_ROWS], dtype=pl.INT32) + state_table_row_map = pl.create_tensor([1, DECODE_RMS_PAD_ROWS], dtype=pl.INT32) + build_decode_padded_write_schedule( + position_ids, + cmp_slot_mapping, + pl.const(DECODE_S, pl.INT32), + pl.const(COMPRESS_RATIO, pl.INT32), + pl.const(DECODE_RMS_TILE, pl.INT32), + pl.const(DECODE_RMS_PAD_TILE, pl.INT32), + write_pos_map, + write_dst_map, + kv_out_row_map, + state_table_row_map, + ) + + # Compact pool enumeration: one entry per batch row, mapping to its scattered + # RMS-padded write row (pad_row = (b // RMS_TILE) * RMS_PAD_TILE + b % RMS_TILE, + # the same placement build_decode_padded_write_schedule uses). The core pools + # only these DECODE_B rows instead of all DECODE_RMS_PAD_ROWS padded rows. + pool_row_map = pl.create_tensor([1, DECODE_B], dtype=pl.INT32) + with pl.at(level=pl.Level.CORE_GROUP, name_hint="decode_pool_row_map"): + for b in pl.range(DECODE_B): + pad_row = (b // DECODE_RMS_TILE) * DECODE_RMS_PAD_TILE + (b % DECODE_RMS_TILE) + pl.write(pool_row_map, [0, b], pl.cast(pad_row, pl.INT32)) + + kv_proj_pad = pl.create_tensor([DECODE_BS_PAD, OUT_DIM], dtype=pl.FP32) + score_proj_pad = pl.create_tensor([DECODE_BS_PAD, OUT_DIM], dtype=pl.FP32) + compressor_ratio128_proj(x, wkv, wgate, kv_proj_pad, score_proj_pad) + + compress_state_rows = pl.reshape( + compress_state, + [DECODE_COMPRESS_STATE_BLOCK_NUM * DECODE_COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], + ) + cmp_kv_cache_flat = pl.reshape(cmp_kv_cache, [DECODE_COMPRESSOR_CMP_BLOCK_NUM * BLOCK_SIZE, HEAD_DIM]) + + pooled_kv = pl.create_tensor([DECODE_RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) + normed_kv = pl.create_tensor([DECODE_RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) + cos_b = pl.create_tensor([DECODE_RMS_PAD_ROWS, ROPE_HALF], dtype=pl.FP32) + sin_b = pl.create_tensor([DECODE_RMS_PAD_ROWS, ROPE_HALF], dtype=pl.FP32) + normed_kv = compressor_core_ratio128( + kv_proj_pad, + score_proj_pad, + position_ids, + state_slot_mapping, + ape, + norm_w, + compress_state_rows, + compress_state_block_table, + freqs_cos, + freqs_sin, + write_pos_map, + write_dst_map, + state_table_row_map, + pool_row_map, + pooled_kv, + normed_kv, + cos_b, + sin_b, + ) + finalize_compressor_writes( + normed_kv, + write_dst_map, + cmp_kv_cache_flat, + pl.const(DECODE_RMS_PAD_ROWS, pl.INT32), + pl.const(0, pl.INT32), + ) + + # Decode-only: scatter the compressed rows to the per-token kv output. + with pl.at(level=pl.Level.CORE_GROUP, name_hint="decode_kv_out"): + for kv_write_row in pl.range(DECODE_RMS_PAD_ROWS): + cmp_row_raw = pl.read(write_dst_map, [0, kv_write_row]) + if cmp_row_raw >= 0: + kv_out_raw = pl.read(kv_out_row_map, [0, kv_write_row]) + if kv_out_raw >= 0: + kv_out_row = pl.cast(kv_out_raw, target_type=pl.INDEX) + kv[kv_out_row : kv_out_row + 1, :] = normed_kv[kv_write_row : kv_write_row + 1, 0:HEAD_DIM] + return kv + + +@pl.jit +def decode_compressor_ratio128_test( + x: pl.Tensor[[DECODE_T, D], pl.BF16], + kv: pl.Out[pl.Tensor[[DECODE_T, HEAD_DIM], pl.FP32]], + compress_state: pl.InOut[pl.Tensor[[DECODE_COMPRESS_STATE_BLOCK_NUM, DECODE_COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32]], + compress_state_block_table: pl.Tensor[[DECODE_B, DECODE_COMPRESS_STATE_MAX_BLOCKS], pl.INT32], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], + norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + cmp_kv_cache: pl.InOut[pl.Tensor[[DECODE_COMPRESSOR_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], + position_ids: pl.Tensor[[DECODE_T], pl.INT32], + cmp_slot_mapping: pl.Tensor[[DECODE_T], pl.INT64], + state_slot_mapping: pl.Tensor[[DECODE_T], pl.INT64], +): + decode_compressor_ratio128( + x, kv, compress_state, compress_state_block_table, wkv, wgate, ape, norm_w, freqs_cos, freqs_sin, + cmp_kv_cache, position_ids, cmp_slot_mapping, state_slot_mapping, + ) + return kv, compress_state, cmp_kv_cache + + +@pl.jit.inline +def prefill_compressor_ratio128( + x: pl.Tensor[[PREFILL_T, D], pl.BF16], + compress_state: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[PREFILL_B, HCA_STATE_MAX_BLOCKS], pl.INT32], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], + norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + cmp_kv: pl.Out[pl.Tensor[[HCA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], + position_ids: pl.Tensor[[PREFILL_T], pl.INT32], + num_tokens: pl.Scalar[pl.INT32], + cmp_slot_mapping: pl.Tensor[[PREFILL_T], pl.INT64], + state_slot_mapping: pl.Tensor[[PREFILL_T], pl.INT64], +): + # Thin prefill wrapper: build the num_tokens-bounded write schedule, project + # (static whole-M tiling), then run the shared core with the fresh prefill + # state and no per-token kv output. num_tokens is consumed only by the + # schedule builder; the core gates scatter on state_slot_mapping >= 0 (padding + # is marked -1 by the host). See compressor_core_ratio128. + write_pos_map = pl.create_tensor([1, HCA_C128_RMS_TILE], dtype=pl.INT32) + write_dst_map = pl.create_tensor([1, HCA_C128_RMS_TILE], dtype=pl.INT32) + state_table_row_map = pl.create_tensor([1, HCA_C128_RMS_TILE], dtype=pl.INT32) + build_prefill_write_schedule( + position_ids, + cmp_slot_mapping, + num_tokens, + write_pos_map, + write_dst_map, + state_table_row_map, + ) + + # Prefill has no scattered RMS padding: its write rows are already compact, so + # the pool enumeration is the identity over the write rows. The core then pools + # exactly the same rows it did before this fix (behaviour unchanged for prefill). + pool_row_map = pl.create_tensor([1, HCA_C128_RMS_TILE], dtype=pl.INT32) + with pl.at(level=pl.Level.CORE_GROUP, name_hint="prefill_pool_row_map"): + for p in pl.range(HCA_C128_RMS_TILE): + pl.write(pool_row_map, [0, p], pl.cast(p, pl.INT32)) + + kv_proj_scratch = pl.create_tensor([PREFILL_T, OUT_DIM], dtype=pl.FP32) + score_proj_scratch = pl.create_tensor([PREFILL_T, OUT_DIM], dtype=pl.FP32) + prefill_compressor_ratio128_proj(x, wkv, wgate, kv_proj_scratch, score_proj_scratch) + + compress_state_rows = pl.reshape( + compress_state, [HCA_STATE_BLOCK_NUM * HCA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM] + ) + cmp_kv_flat = pl.reshape(cmp_kv, [HCA_CMP_BLOCK_NUM * BLOCK_SIZE, HEAD_DIM]) + + pooled_kv_pad = pl.create_tensor([HCA_C128_RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) + normed_kv_pad = pl.create_tensor([HCA_C128_RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) + cos_b = pl.create_tensor([HCA_C128_RMS_PAD_ROWS, ROPE_HALF], dtype=pl.FP32) + sin_b = pl.create_tensor([HCA_C128_RMS_PAD_ROWS, ROPE_HALF], dtype=pl.FP32) + normed_kv_pad = compressor_core_ratio128( + kv_proj_scratch, + score_proj_scratch, + position_ids, + state_slot_mapping, + ape, + norm_w, + compress_state_rows, + compress_state_block_table, + freqs_cos, + freqs_sin, + write_pos_map, + write_dst_map, + state_table_row_map, + pool_row_map, + pooled_kv_pad, + normed_kv_pad, + cos_b, + sin_b, + ) + finalize_compressor_writes( + normed_kv_pad, + write_dst_map, + cmp_kv_flat, + pl.const(HCA_C128_RMS_TILE, pl.INT32), + pl.const(0, pl.INT32), + ) + return cmp_kv, compress_state + + +@pl.jit +def prefill_compressor_ratio128_test( + x: pl.Tensor[[PREFILL_T, D], pl.BF16], + compress_state: pl.InOut[pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32]], + compress_state_block_table: pl.Tensor[[PREFILL_B, HCA_STATE_MAX_BLOCKS], pl.INT32], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], + norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + cmp_kv: pl.InOut[pl.Tensor[[HCA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], + position_ids: pl.Tensor[[PREFILL_T], pl.INT32], + num_tokens: pl.Scalar[pl.INT32], + cmp_slot_mapping: pl.Tensor[[PREFILL_T], pl.INT64], + state_slot_mapping: pl.Tensor[[PREFILL_T], pl.INT64], +): + return prefill_compressor_ratio128( + x, compress_state, compress_state_block_table, wkv, wgate, ape, norm_w, freqs_cos, freqs_sin, + cmp_kv, position_ids, num_tokens, cmp_slot_mapping, state_slot_mapping, + ) + + +def _golden_compressor_ratio128_pipeline( + tensors, + *, + write_pos_map, + write_dst_map, + kv_out_row_map, + state_table_row_map, + cmp_cache_name, + kv_out_name=None, +): + import torch + + x = tensors["x"].view(-1, D).float() + position_ids = tensors["position_ids"].view(-1).to(torch.int64) + state_slot_mapping = tensors["state_slot_mapping"].view(-1).to(torch.int64) + compress_state = tensors["compress_state"] + compress_state_rows = compress_state.view(-1, COMPRESS_STATE_DIM) + compress_state_block_table = tensors["compress_state_block_table"] + cmp_kv_cache = tensors[cmp_cache_name] + cmp_kv_cache_flat = cmp_kv_cache.view(cmp_kv_cache.shape[0] * BLOCK_SIZE, HEAD_DIM) + + kv_proj = x @ tensors["wkv"].float().t() + score_proj = x @ tensors["wgate"].float().t() + ape = tensors["ape"] + for token_id in range(x.shape[0]): + dst_row = int(state_slot_mapping[token_id].item()) + if dst_row < 0: + continue + pos = int(position_ids[token_id].item()) + ape_slot = pos % COMPRESS_RATIO + compress_state_rows[dst_row, 0:OUT_DIM] = kv_proj[token_id] + compress_state_rows[dst_row, OUT_DIM:COMPRESS_STATE_DIM] = score_proj[token_id] + ape[ape_slot] + + def rmsnorm(x, w): + var = x.square().mean(-1, keepdim=True) + return x * torch.rsqrt(var + EPS) * w.float().view(1, HEAD_DIM) + + state_block_size = compress_state.shape[1] + num_state_blocks = STATE_LEN // state_block_size + kv_out = tensors[kv_out_name] if kv_out_name is not None else None + for write_i, dst_row in enumerate(write_dst_map): + if dst_row < 0: + continue + kv_out_row = kv_out_row_map[write_i] + if kv_out_row < 0: + continue + state_table_row = state_table_row_map[write_i] + if state_table_row < 0: + continue + write_pos = write_pos_map[write_i] + state_pos0 = write_pos + 1 - COMPRESS_RATIO + base_logical_blk = state_pos0 // state_block_size + pool_kv_state = torch.zeros(STATE_LEN, OUT_DIM, dtype=torch.float32, device=x.device) + pool_score_state = torch.full((STATE_LEN, OUT_DIM), float("-inf"), dtype=torch.float32, device=x.device) + for blk_i in range(num_state_blocks): + logical_blk = base_logical_blk + blk_i + if logical_blk < 0 or logical_blk >= compress_state_block_table.shape[1]: + continue + state_blk = int(compress_state_block_table[state_table_row, logical_blk].item()) + if state_blk < 0: + continue + row0 = state_blk * state_block_size + s0 = blk_i * state_block_size + pool_kv_state[s0 : s0 + state_block_size] = compress_state_rows[row0 : row0 + state_block_size, 0:OUT_DIM] + pool_score_state[s0 : s0 + state_block_size] = compress_state_rows[ + row0 : row0 + state_block_size, OUT_DIM:COMPRESS_STATE_DIM + ] + + pooled = (pool_kv_state * pool_score_state.softmax(dim=0)).sum(dim=0, keepdim=True) + normed = rmsnorm(pooled, tensors["norm_w"]) + rope_pair = normed[..., NOPE_HEAD_DIM:].unflatten(-1, (-1, 2)) + even = rope_pair[..., 0].float() + odd = rope_pair[..., 1].float() + cmp_pos = write_pos + 1 - COMPRESS_RATIO + cos = tensors["freqs_cos"][cmp_pos : cmp_pos + 1, 0:ROPE_HALF].float() + sin = tensors["freqs_sin"][cmp_pos : cmp_pos + 1, 0:ROPE_HALF].float() + rot_even = even * cos - odd * sin + rot_odd = even * sin + odd * cos + normed[:, NOPE_HEAD_DIM:] = torch.stack([rot_even, rot_odd], dim=-1).flatten(-2) + + if kv_out is not None: + kv_out[kv_out_row : kv_out_row + 1, :] = normed.reshape(1, HEAD_DIM) + cmp_kv_cache_flat[dst_row] = normed[0] + + tensors["compress_state"][:] = compress_state_rows.view_as(compress_state) + tensors[cmp_cache_name][:] = cmp_kv_cache_flat.view_as(cmp_kv_cache) + + +def golden_decode_compressor_ratio128(tensors): + """Torch reference for Compressor.forward (decode branch, ratio=128 non-overlap).""" + position_ids = tensors["position_ids"].view(-1).to("cpu") + cmp_slot_mapping = tensors["cmp_slot_mapping"].view(-1).to("cpu") + write_pos_map = [0] * DECODE_RMS_PAD_ROWS + write_dst_map = [-1] * DECODE_RMS_PAD_ROWS + kv_out_row_map = [-1] * DECODE_RMS_PAD_ROWS + state_table_row_map = [-1] * DECODE_RMS_PAD_ROWS + for b in range(DECODE_B): + base_t = b * DECODE_S + first_pos = int(position_ids[base_t].item()) + pos_in_window = first_pos % COMPRESS_RATIO + if pos_in_window + DECODE_S >= COMPRESS_RATIO: + boundary_s = COMPRESS_RATIO - 1 - pos_in_window + token_t = base_t + boundary_s + dst_row = int(cmp_slot_mapping[token_t].item()) + if dst_row >= 0: + pad_row = (b // DECODE_RMS_TILE) * DECODE_RMS_PAD_TILE + (b % DECODE_RMS_TILE) + if pad_row < DECODE_RMS_PAD_ROWS: + write_pos_map[pad_row] = first_pos + boundary_s + write_dst_map[pad_row] = dst_row + kv_out_row_map[pad_row] = base_t + state_table_row_map[pad_row] = b + + _golden_compressor_ratio128_pipeline( + tensors, + write_pos_map=write_pos_map, + write_dst_map=write_dst_map, + kv_out_row_map=kv_out_row_map, + state_table_row_map=state_table_row_map, + cmp_cache_name="cmp_kv_cache", + kv_out_name="kv", + ) + + +def golden_prefill_compressor_ratio128(tensors): + num_tokens = int(tensors["num_tokens"]) + position_ids = tensors["position_ids"].view(-1).to("cpu") + cmp_slot_mapping = tensors["cmp_slot_mapping"].view(-1).to("cpu") + write_pos_map = [0] * HCA_C128_RMS_TILE + write_dst_map = [-1] * HCA_C128_RMS_TILE + kv_out_row_map = [-1] * HCA_C128_RMS_TILE + state_table_row_map = [-1] * HCA_C128_RMS_TILE + write_i = 0 + for token_id in range(position_ids.numel()): + if token_id >= num_tokens: + break + dst_row = int(cmp_slot_mapping[token_id].item()) + if dst_row < 0: + continue + if write_i >= HCA_C128_RMS_TILE: + break + write_pos_map[write_i] = int(position_ids[token_id].item()) + write_dst_map[write_i] = dst_row + kv_out_row_map[write_i] = token_id + state_table_row_map[write_i] = 0 + write_i += 1 + + _golden_compressor_ratio128_pipeline( + tensors, + write_pos_map=write_pos_map, + write_dst_map=write_dst_map, + kv_out_row_map=kv_out_row_map, + state_table_row_map=state_table_row_map, + cmp_cache_name="cmp_kv", + ) + + +def build_decode_tensor_specs(start_pos=None): + import torch # type: ignore[import] + from decode_metadata import ( + block_table, + compressed_slot_mapping, + hca_decode_start_set, + position_ids_from_starts, + resolve_start_positions, + state_slot_mapping, + ) + from golden import TensorSpec + from rope_tables import build_deepseek_v4_rope_tables + + shared_freqs_cos, shared_freqs_sin = build_deepseek_v4_rope_tables(M, COMPRESS_RATIO, dtype=torch.bfloat16) + + def init_x(): + return torch.rand(DECODE_B, DECODE_S, D).reshape(DECODE_T, D) + def init_compress_state(): + return torch.zeros(DECODE_COMPRESS_STATE_BLOCK_NUM, DECODE_COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM) + # Calibrated to the real DeepSeek-V4-Flash 150 + # (ratio-128) main compressor (mean l7/l9 of + # extract_weights_flash): zero-mean Gaussian BF16 weights at the measured std; the RMSNorm + # gamma centers near the measured mean (not ones / not uniform). + def init_wkv(): + return torch.randn(OUT_DIM, D) * 0.0246 + def init_wgate(): + return torch.randn(OUT_DIM, D) * 0.0316 + def init_ape(): + return torch.randn(COMPRESS_RATIO, OUT_DIM) * 0.0340 + def init_norm_w(): + return 0.1001 + 0.0549 * torch.randn(HEAD_DIM) + def init_freqs_cos(): + return shared_freqs_cos.clone() + def init_freqs_sin(): + return shared_freqs_sin.clone() + def init_cmp_kv_cache(): + return torch.zeros(DECODE_COMPRESSOR_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM) + def init_compress_state_block_table(): + return block_table( + batch=DECODE_B, + table_blocks=DECODE_COMPRESS_STATE_MAX_BLOCKS, + physical_blocks=DECODE_COMPRESS_STATE_PHYSICAL_BLOCKS, + permuted=True, + ) + def init_cmp_block_table(): + return block_table( + batch=DECODE_B, + table_blocks=DECODE_COMPRESSOR_CMP_MAX_BLOCKS, + physical_blocks=DECODE_COMPRESSOR_CMP_MAX_BLOCKS, + permuted=True, + ) + def init_default_start_pos(): + # Canonical HCA start-position set (ratio-128 compressor branches + 8k long-context). + return hca_decode_start_set( + batch=DECODE_B, compress_ratio=COMPRESS_RATIO, state_block_size=DECODE_COMPRESS_STATE_BLOCK_SIZE) + def init_start_pos(): + return resolve_start_positions( + start_pos, + batch=DECODE_B, + seq=DECODE_S, + max_seq_len=MAX_SEQ_LEN, + default_fn=init_default_start_pos, + ) + def _position_ids_bs(): + # [DECODE_B, DECODE_S] positions; token-major init_position_ids reshapes this to [DECODE_T]. + return position_ids_from_starts(init_start_pos(), seq=DECODE_S) + def init_position_ids(): + # token-major [DECODE_T]; row t == (b * DECODE_S + s) carries position of sequence b, intra-token s + return _position_ids_bs().reshape(DECODE_T) + def init_state_slot_mapping(): + return state_slot_mapping( + _position_ids_bs(), + init_compress_state_block_table(), + state_block_size=DECODE_COMPRESS_STATE_BLOCK_SIZE, + ).reshape(DECODE_T) + def init_cmp_slot_mapping(): + return compressed_slot_mapping( + _position_ids_bs(), + init_cmp_block_table(), + compress_ratio=COMPRESS_RATIO, + block_size=BLOCK_SIZE, + ).reshape(DECODE_T) + return [ + TensorSpec("x", [DECODE_T, D], torch.bfloat16, init_value=init_x), + TensorSpec("kv", [DECODE_T, HEAD_DIM], torch.float32, is_output=True), + TensorSpec("compress_state", [DECODE_COMPRESS_STATE_BLOCK_NUM, DECODE_COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], torch.float32, init_value=init_compress_state, is_output=True), + TensorSpec("compress_state_block_table", [DECODE_B, DECODE_COMPRESS_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), + TensorSpec("wkv", [OUT_DIM, D], torch.bfloat16, init_value=init_wkv), + TensorSpec("wgate", [OUT_DIM, D], torch.bfloat16, init_value=init_wgate), + TensorSpec("ape", [COMPRESS_RATIO, OUT_DIM], torch.float32, init_value=init_ape), + TensorSpec("norm_w", [HEAD_DIM], torch.bfloat16, init_value=init_norm_w), + TensorSpec("freqs_cos", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_cos), + TensorSpec("freqs_sin", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_sin), + TensorSpec("cmp_kv_cache", [DECODE_COMPRESSOR_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], torch.bfloat16, init_value=init_cmp_kv_cache, is_output=True), + TensorSpec("position_ids", [DECODE_T], torch.int32, init_value=init_position_ids), + TensorSpec("cmp_slot_mapping", [DECODE_T], torch.int64, init_value=init_cmp_slot_mapping), + TensorSpec("state_slot_mapping", [DECODE_T], torch.int64, init_value=init_state_slot_mapping), + ] + + +def build_prefill_tensor_specs(start_pos: int = PREFILL_START_POS): + import torch + from golden import ScalarSpec, TensorSpec + from rope_tables import build_deepseek_v4_rope_tables + + shared_freqs_cos, shared_freqs_sin = build_deepseek_v4_rope_tables(M, COMPRESS_RATIO, dtype=torch.bfloat16) + + num_tokens = PREFILL_T + if start_pos < 0: + raise ValueError("start_pos must be non-negative") + if start_pos + num_tokens > MAX_SEQ_LEN: + raise ValueError("start_pos + num_tokens exceeds max_position_embeddings") + + def init_compress_state_block_table(): + table = torch.full((PREFILL_B, HCA_STATE_MAX_BLOCKS), -1, dtype=torch.int32) + for block in range(HCA_STATE_MAX_BLOCKS): + table[0, block] = (block * 17 + 3) % HCA_STATE_MAX_BLOCKS + return table + def state_row(abs_pos): + if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: + return -1 + table = init_compress_state_block_table() + block = abs_pos // HCA_STATE_BLOCK_SIZE + intra = abs_pos % HCA_STATE_BLOCK_SIZE + return int(table[0, block].item()) * HCA_STATE_BLOCK_SIZE + intra + def init_x(): + return ((torch.rand(PREFILL_T, D) - 0.5) * 0.1).to(torch.bfloat16) + def init_compress_state(): + state = torch.zeros(HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM) + flat = state.view(-1, COMPRESS_STATE_DIM) + for abs_pos in range(max(0, start_pos - COMPRESS_RATIO), start_pos): + row = state_row(abs_pos) + if row >= 0: + flat[row, 0:OUT_DIM] = (torch.rand(OUT_DIM) - 0.5) * 0.05 + flat[row, OUT_DIM:COMPRESS_STATE_DIM] = (torch.rand(OUT_DIM) - 0.5) * 0.05 + return state + # Calibrated to the real DeepSeek-V4-Flash HCA (ratio-128) main compressor (mean l7/l9 of + # extract_weights_flash): zero-mean Gaussian BF16 weights at the measured std; the RMSNorm + # gamma centers near the measured mean (not ones / not uniform). Mirrors the decode path. + def init_wkv(): + return torch.randn(OUT_DIM, D) * 0.0246 + def init_wgate(): + return torch.randn(OUT_DIM, D) * 0.0316 + def init_ape(): + return torch.randn(COMPRESS_RATIO, OUT_DIM) * 0.0340 + def init_norm_w(): + return 0.1001 + 0.0549 * torch.randn(HEAD_DIM) + def init_freqs_cos(): + return shared_freqs_cos.clone() + def init_freqs_sin(): + return shared_freqs_sin.clone() + def init_cmp_kv(): + return torch.zeros(HCA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM, dtype=torch.bfloat16) + def init_position_ids(): + return torch.arange(start_pos, start_pos + PREFILL_T, dtype=torch.int32) + def init_cmp_slot_mapping(): + mapping = torch.full((PREFILL_T,), -1, dtype=torch.int64) + for token_id in range(num_tokens): + pos = start_pos + token_id + if pos + 1 >= COMPRESS_RATIO and (pos + 1) % COMPRESS_RATIO == 0: + mapping[token_id] = (pos + 1) // COMPRESS_RATIO - 1 + return mapping + def init_state_slot_mapping(): + mapping = torch.full((PREFILL_T,), -1, dtype=torch.int64) + for token_id in range(num_tokens): + mapping[token_id] = state_row(start_pos + token_id) + return mapping + + return [ + TensorSpec("x", [PREFILL_T, D], torch.bfloat16, init_value=init_x), + TensorSpec("compress_state", [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], torch.float32, init_value=init_compress_state, is_output=True), + TensorSpec("compress_state_block_table", [PREFILL_B, HCA_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), + TensorSpec("wkv", [OUT_DIM, D], torch.bfloat16, init_value=init_wkv), + TensorSpec("wgate", [OUT_DIM, D], torch.bfloat16, init_value=init_wgate), + TensorSpec("ape", [COMPRESS_RATIO, OUT_DIM], torch.float32, init_value=init_ape), + TensorSpec("norm_w", [HEAD_DIM], torch.bfloat16, init_value=init_norm_w), + TensorSpec("freqs_cos", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_cos), + TensorSpec("freqs_sin", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_sin), + TensorSpec("cmp_kv", [HCA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], torch.bfloat16, init_value=init_cmp_kv, is_output=True), + TensorSpec("position_ids", [PREFILL_T], torch.int32, init_value=init_position_ids), + ScalarSpec("num_tokens", torch.int32, num_tokens), + TensorSpec("cmp_slot_mapping", [PREFILL_T], torch.int64, init_value=init_cmp_slot_mapping), + TensorSpec("state_slot_mapping", [PREFILL_T], torch.int64, init_value=init_state_slot_mapping), + ] + + +def _run_decode_validation(args): + from golden import ratio_allclose, run_jit + + return run_jit( + fn=decode_compressor_ratio128_test, + specs=build_decode_tensor_specs(args.start_pos), + golden_fn=golden_decode_compressor_ratio128, + runtime_dir=args.runtime_dir, + golden_data=args.golden_data, + compile_cfg=dict(dump_passes=args.dump_passes), + runtime_cfg=dict( + platform=args.platform, + device_id=args.device, + enable_l2_swimlane=args.enable_l2_swimlane, + enable_dep_gen=args.enable_dep_gen, + ), + compile_only=args.compile_only, + rtol=1e-3, + atol=1e-3, + compare_fn={ + "kv": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), + "compress_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), + "cmp_kv_cache": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), + }, + ) + + +def _run_prefill_validation(args): + from golden import ratio_allclose, run_jit + + start_pos = PREFILL_START_POS if args.start_pos is None else args.start_pos + return run_jit( + fn=prefill_compressor_ratio128_test, + specs=build_prefill_tensor_specs(start_pos), + golden_fn=golden_prefill_compressor_ratio128, + runtime_dir=args.runtime_dir, + golden_data=args.golden_data, + compile_cfg=dict(dump_passes=args.dump_passes), + runtime_cfg=dict( + platform=args.platform, + device_id=args.device, + enable_l2_swimlane=args.enable_l2_swimlane, + enable_dep_gen=args.enable_dep_gen, + ), + rtol=1e-3, + atol=1e-3, + compile_only=args.compile_only, + compare_fn={ + "cmp_kv": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), + "compress_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), + }, + ) + + +def main(): + import argparse + + parser = argparse.ArgumentParser(description="Standalone DeepSeek V4 compressor ratio128 validation.") + parser.add_argument("--mode", choices=["decode", "prefill", "both"], default="both") + parser.add_argument("-p", "--platform", type=str, default="a2a3", choices=["a2a3", "a2a3sim", "a5", "a5sim"]) + parser.add_argument("-d", "--device", type=int, default=0) + parser.add_argument("--compile-only", action="store_true", default=False) + parser.add_argument( + "--start-pos", + type=int, + default=None, + help="Fixture-only start position override. Decode defaults to its canonical batch set; prefill defaults to 0.", + ) + parser.add_argument("--enable-l2-swimlane", action="store_true", default=False) + parser.add_argument("--enable-dep-gen", action="store_true", default=False) + parser.add_argument("--runtime-dir", type=str, default=None) + parser.add_argument("--golden-data", type=str, default=None) + parser.add_argument("--dump-passes", action="store_true", default=False) + args = parser.parse_args() + + modes = ("decode", "prefill") if args.mode == "both" else (args.mode,) + for mode in modes: + result = _run_decode_validation(args) if mode == "decode" else _run_prefill_validation(args) + if not result.passed: + if result.error: + print(result.error) + raise SystemExit(1) + + +if __name__ == "__main__": + main() diff --git a/models/deepseek/v4/compressor_ratio4.py b/models/deepseek/v4/compressor_ratio4.py new file mode 100644 index 00000000..538e07e4 --- /dev/null +++ b/models/deepseek/v4/compressor_ratio4.py @@ -0,0 +1,1227 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""DeepSeek-V4 ratio-4 compressor decode and prefill paths.""" + +import pypto.language as pl + +from config import ( + FLASH as M, + DECODE_BATCH, + DECODE_SEQ, + BLOCK_SIZE, + C4A_COMPRESSOR_BLOCK_SIZE, + DECODE_CMP_BLOCK_NUM, + KV_CMP_MAX_BLOCKS, + FP32_NEG_INF, + PREFILL_CMP_BLOCK_NUM, +) +from compressor_common import ( + build_prefill_write_schedule, + compressor_rmsnorm_rope, + finalize_compressor_writes, + gather_compressor_rope_rows, +) + + +EPS = M.rms_norm_eps +D = M.hidden_size +HEAD_DIM = M.head_dim +HEAD_DIM_INV = 1.0 / HEAD_DIM +ROPE_HEAD_DIM = M.qk_rope_head_dim +NOPE_HEAD_DIM = M.nope_head_dim +MAX_SEQ_LEN = M.max_position_embeddings + +COMPRESS_RATIO = 4 +OVERLAP = COMPRESS_RATIO == 4 +COFF = 1 + int(OVERLAP) +OUT_DIM = COFF * HEAD_DIM +STATE_LEN = COFF * COMPRESS_RATIO +COMPRESS_STATE_DIM = 2 * OUT_DIM +POOL_HEAD_TILE = HEAD_DIM +RATIO4_STATE_BLOCK_SIZE = C4A_COMPRESSOR_BLOCK_SIZE + +# Shared ratio4 sub-kernel tiling. +PROJ_ROWS = pl.dynamic("COMPRESSOR4_PROJ_ROWS") +PROJ_ROWS_PAD = pl.dynamic("COMPRESSOR4_PROJ_ROWS_PAD") +PROJ_MM_B_TILE = 16 +PROJ_OUT_TILE = 64 +PROJ_K_TILE = 512 +POOL_STATE_ROWS = pl.dynamic("COMPRESSOR4_POOL_STATE_ROWS") +POOL_TABLE_ROWS = pl.dynamic("COMPRESSOR4_POOL_TABLE_ROWS") +POOL_TABLE_BLOCKS = pl.dynamic("COMPRESSOR4_POOL_TABLE_BLOCKS") + + +@pl.jit.inline +def compressor_ratio4_proj( + x: pl.Tensor[[PROJ_ROWS, D], pl.BF16], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + kv_proj_out: pl.Tensor[[PROJ_ROWS_PAD, OUT_DIM], pl.FP32], + score_proj_out: pl.Tensor[[PROJ_ROWS_PAD, OUT_DIM], pl.FP32], +): + t_dim = pl.tensor.dim(x, 0) + t_matmul = pl.tensor.dim(kv_proj_out, 0) + for idx in pl.spmd(t_matmul * OUT_DIM // (PROJ_MM_B_TILE * PROJ_OUT_TILE), name_hint="kv_score_proj"): + global_row0 = (idx // (OUT_DIM // PROJ_OUT_TILE)) * PROJ_MM_B_TILE + o0 = (idx % (OUT_DIM // PROJ_OUT_TILE)) * PROJ_OUT_TILE + kv_acc = pl.create_tensor([PROJ_MM_B_TILE, PROJ_OUT_TILE], dtype=pl.FP32) + score_acc = pl.create_tensor([PROJ_MM_B_TILE, PROJ_OUT_TILE], dtype=pl.FP32) + for kb in pl.pipeline(0, D // PROJ_K_TILE, stage=2): + k0 = kb * PROJ_K_TILE + x_rows = pl.min(PROJ_MM_B_TILE, t_dim - global_row0) + x_tile = pl.slice(x, [PROJ_MM_B_TILE, PROJ_K_TILE], [global_row0, k0], valid_shape=[x_rows, PROJ_K_TILE]) + wkv_tile = wkv[o0 : o0 + PROJ_OUT_TILE, k0 : k0 + PROJ_K_TILE] + wgate_tile = wgate[o0 : o0 + PROJ_OUT_TILE, k0 : k0 + PROJ_K_TILE] + if k0 == 0: + kv_acc = pl.matmul(x_tile, wkv_tile, out_dtype=pl.FP32, b_trans=True) + score_acc = pl.matmul(x_tile, wgate_tile, out_dtype=pl.FP32, b_trans=True) + else: + kv_acc = pl.matmul_acc(kv_acc, x_tile, wkv_tile, b_trans=True) + score_acc = pl.matmul_acc(score_acc, x_tile, wgate_tile, b_trans=True) + kv_proj_out[global_row0 : global_row0 + PROJ_MM_B_TILE, o0 : o0 + PROJ_OUT_TILE] = kv_acc + score_proj_out[global_row0 : global_row0 + PROJ_MM_B_TILE, o0 : o0 + PROJ_OUT_TILE] = score_acc + + +@pl.jit.inline +def compressor_ratio4_pool_math( + score_state: pl.Tensor[[STATE_LEN, POOL_HEAD_TILE], pl.FP32], + kv_state: pl.Tensor[[STATE_LEN, POOL_HEAD_TILE], pl.FP32], +): + init_slot = STATE_LEN - 1 + mi_buf = pl.create_tensor([1, POOL_HEAD_TILE], dtype=pl.FP32) + li_buf = pl.create_tensor([1, POOL_HEAD_TILE], dtype=pl.FP32) + oi_buf = pl.create_tensor([1, POOL_HEAD_TILE], dtype=pl.FP32) + mi_buf[0:1, 0:POOL_HEAD_TILE] = score_state[init_slot : init_slot + 1, 0 : POOL_HEAD_TILE] + li_buf[0:1, 0:POOL_HEAD_TILE] = pl.exp(pl.sub(mi_buf[0:1, 0:POOL_HEAD_TILE], mi_buf[0:1, 0:POOL_HEAD_TILE])) + oi_buf[0:1, 0:POOL_HEAD_TILE] = kv_state[init_slot : init_slot + 1, 0 : POOL_HEAD_TILE] + for slot_i in pl.range(STATE_LEN - 1): + mi = mi_buf[0:1, 0:POOL_HEAD_TILE] + li = li_buf[0:1, 0:POOL_HEAD_TILE] + oi = oi_buf[0:1, 0:POOL_HEAD_TILE] + slot_score = score_state[slot_i : slot_i + 1, 0 : POOL_HEAD_TILE] + slot_kv = kv_state[slot_i : slot_i + 1, 0 : POOL_HEAD_TILE] + mi_next = pl.maximum(mi, slot_score) + alpha = pl.exp(pl.sub(mi, mi_next)) + beta = pl.exp(pl.sub(slot_score, mi_next)) + li_buf[0:1, 0:POOL_HEAD_TILE] = pl.add(pl.mul(alpha, li), beta) + oi_buf[0:1, 0:POOL_HEAD_TILE] = pl.add(pl.mul(oi, alpha), pl.mul(slot_kv, beta)) + mi_buf[0:1, 0:POOL_HEAD_TILE] = mi_next + return pl.div(oi_buf[0:1, 0:POOL_HEAD_TILE], li_buf[0:1, 0:POOL_HEAD_TILE]) + + +@pl.jit.inline +def compressor_ratio4_pool_window( + compress_state_rows: pl.Tensor[[POOL_STATE_ROWS, COMPRESS_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[POOL_TABLE_ROWS, POOL_TABLE_BLOCKS], pl.INT32], + table_row: pl.Scalar[pl.INDEX], + write_pos: pl.Scalar[pl.INT32], + h0: pl.Scalar[pl.INDEX], +): + pool_score_tile = pl.create_tensor([STATE_LEN, POOL_HEAD_TILE], dtype=pl.FP32) + pool_kv_tile = pl.create_tensor([STATE_LEN, POOL_HEAD_TILE], dtype=pl.FP32) + cur_start = write_pos + 1 - COMPRESS_RATIO + prev_start = cur_start - COMPRESS_RATIO + + if write_pos >= 2 * COMPRESS_RATIO - 1: + for pool_s in pl.range(COMPRESS_RATIO): + prev_abs = prev_start + pool_s + prev_state_block = pl.cast(prev_abs // RATIO4_STATE_BLOCK_SIZE, pl.INDEX) + prev_state_intra = pl.cast(prev_abs - prev_state_block * RATIO4_STATE_BLOCK_SIZE, pl.INDEX) + prev_phys_block_raw = pl.read(compress_state_block_table, [table_row, prev_state_block]) + if prev_phys_block_raw >= 0: + prev_phys_block = pl.cast(prev_phys_block_raw, pl.INDEX) + prev_state_row = prev_phys_block * RATIO4_STATE_BLOCK_SIZE + prev_state_intra + pool_kv_tile[pool_s : pool_s + 1, 0:POOL_HEAD_TILE] = compress_state_rows[ + prev_state_row : prev_state_row + 1, + h0 : h0 + POOL_HEAD_TILE, + ] + pool_score_tile[pool_s : pool_s + 1, 0:POOL_HEAD_TILE] = compress_state_rows[ + prev_state_row : prev_state_row + 1, + OUT_DIM + h0 : OUT_DIM + h0 + POOL_HEAD_TILE, + ] + else: + pool_kv_tile[pool_s : pool_s + 1, 0:POOL_HEAD_TILE] = pl.full( + [1, POOL_HEAD_TILE], + dtype=pl.FP32, + value=0.0, + ) + pool_score_tile[pool_s : pool_s + 1, 0:POOL_HEAD_TILE] = pl.full( + [1, POOL_HEAD_TILE], + dtype=pl.FP32, + value=FP32_NEG_INF, + ) + else: + pool_kv_tile[0:COMPRESS_RATIO, 0:POOL_HEAD_TILE] = pl.full( + [COMPRESS_RATIO, POOL_HEAD_TILE], + dtype=pl.FP32, + value=0.0, + ) + pool_score_tile[0:COMPRESS_RATIO, 0:POOL_HEAD_TILE] = pl.full( + [COMPRESS_RATIO, POOL_HEAD_TILE], + dtype=pl.FP32, + value=FP32_NEG_INF, + ) + + for pool_s in pl.range(COMPRESS_RATIO): + cur_abs = cur_start + pool_s + back_slot = COMPRESS_RATIO + pool_s + cur_state_block = pl.cast(cur_abs // RATIO4_STATE_BLOCK_SIZE, pl.INDEX) + cur_state_intra = pl.cast(cur_abs - cur_state_block * RATIO4_STATE_BLOCK_SIZE, pl.INDEX) + cur_phys_block_raw = pl.read(compress_state_block_table, [table_row, cur_state_block]) + if cur_phys_block_raw >= 0: + cur_phys_block = pl.cast(cur_phys_block_raw, pl.INDEX) + cur_state_row = cur_phys_block * RATIO4_STATE_BLOCK_SIZE + cur_state_intra + pool_kv_tile[back_slot : back_slot + 1, 0:POOL_HEAD_TILE] = compress_state_rows[ + cur_state_row : cur_state_row + 1, + HEAD_DIM + h0 : HEAD_DIM + h0 + POOL_HEAD_TILE, + ] + pool_score_tile[back_slot : back_slot + 1, 0:POOL_HEAD_TILE] = compress_state_rows[ + cur_state_row : cur_state_row + 1, + OUT_DIM + HEAD_DIM + h0 : OUT_DIM + HEAD_DIM + h0 + POOL_HEAD_TILE, + ] + else: + pool_kv_tile[back_slot : back_slot + 1, 0:POOL_HEAD_TILE] = pl.full( + [1, POOL_HEAD_TILE], + dtype=pl.FP32, + value=0.0, + ) + pool_score_tile[back_slot : back_slot + 1, 0:POOL_HEAD_TILE] = pl.full( + [1, POOL_HEAD_TILE], + dtype=pl.FP32, + value=FP32_NEG_INF, + ) + + return compressor_ratio4_pool_math(pool_score_tile, pool_kv_tile) + + +# Shared-core dynamic shapes (bind per caller: decode vs prefill). State and +# cmp-cache tensors are passed as pre-reshaped flat views so their shapes stay +# statically inferable per call site. +CORE_PROJ_ROWS = pl.dynamic("COMPRESSOR4_CORE_PROJ_ROWS") +CORE_TOKENS = pl.dynamic("COMPRESSOR4_CORE_TOKENS") +CORE_WRITE_ROWS = pl.dynamic("COMPRESSOR4_CORE_WRITE_ROWS") +CORE_STATE_ROWS = pl.dynamic("COMPRESSOR4_CORE_STATE_ROWS") +CORE_TABLE_ROWS = pl.dynamic("COMPRESSOR4_CORE_TABLE_ROWS") +CORE_TABLE_BLOCKS = pl.dynamic("COMPRESSOR4_CORE_TABLE_BLOCKS") +# Compact per-regime pool enumeration: decode binds it to the batch-row count +# (real windows scattered across the RMS-padded schedule), prefill to its compact +# write-row count. Pooling only real windows keeps decode off the padded grid. +CORE_POOL_ROWS = pl.dynamic("COMPRESSOR4_CORE_POOL_ROWS") +POOL_HEAD_BLOCKS_CORE = HEAD_DIM // POOL_HEAD_TILE # 1: ratio4 pools the whole head +CORE_INIT_ROW_TILE = 8 # divides both regimes' write_rows (== shared rmsnorm tile) +ROPE_HALF = ROPE_HEAD_DIM // 2 + + +@pl.jit.inline +def compressor_core_ratio4( + kv_proj: pl.Tensor[[CORE_PROJ_ROWS, OUT_DIM], pl.FP32], + score_proj: pl.Tensor[[CORE_PROJ_ROWS, OUT_DIM], pl.FP32], + position_ids: pl.Tensor[[CORE_TOKENS], pl.INT32], + state_slot_mapping: pl.Tensor[[CORE_TOKENS], pl.INT64], + ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], + norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], + compress_state_rows: pl.Tensor[[CORE_STATE_ROWS, COMPRESS_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[CORE_TABLE_ROWS, CORE_TABLE_BLOCKS], pl.INT32], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + write_pos_map: pl.Tensor[[1, CORE_WRITE_ROWS], pl.INT32], + write_dst_map: pl.Tensor[[1, CORE_WRITE_ROWS], pl.INT32], + state_table_row_map: pl.Tensor[[1, CORE_WRITE_ROWS], pl.INT32], + pool_row_map: pl.Tensor[[1, CORE_POOL_ROWS], pl.INT32], + pooled_kv: pl.Tensor[[CORE_WRITE_ROWS, HEAD_DIM], pl.FP32], + normed_kv: pl.Tensor[[CORE_WRITE_ROWS, HEAD_DIM], pl.FP32], + cos_b: pl.Tensor[[CORE_WRITE_ROWS, ROPE_HALF], pl.FP32], + sin_b: pl.Tensor[[CORE_WRITE_ROWS, ROPE_HALF], pl.FP32], +): + """Prefill ratio4 compression math core (mirrors compressor_core_ratio128, + adapted for the overlap window: COFF=2 state, STATE_LEN=8 online-softmax pool + over the whole head, POOL_HEAD_BLOCKS=1). + + scatter projected (kv, score+APE) into paged state (skip slot < 0) + -> softmax-pool each real window (compact pool_row_map, block-table gather) + -> rmsnorm + rope at the window position. + + Stops at normed_kv; the prefill wrapper finalizes cmp_kv + keepalive. Decode no + longer shares this core: it fuses scatter+pool and rmsnorm+rope+cache-write into + two tasks over a single paged 16-row block (see decode_compressor_ratio4), which + a packed multi-write prefill schedule cannot reuse.""" + token_rows = pl.tensor.dim(state_slot_mapping, 0) + write_rows = pl.tensor.dim(write_dst_map, 1) + pool_rows = pl.tensor.dim(pool_row_map, 1) + + # 1. Scatter projected (kv, score+APE) into the paged state buffer. Padding / + # non-writing tokens carry state_slot_mapping < 0 and are skipped (uniform + # decode/prefill validity contract; num_tokens is folded into the schedule). + with pl.spmd(token_rows, name_hint="state_scatter_pre") as scatter_tid: + scatter_t = pl.tile.get_block_idx() + state_row_i64 = pl.read(state_slot_mapping, [scatter_t]) + if state_row_i64 >= 0: + state_row = pl.cast(state_row_i64, target_type=pl.INDEX) + token_pos = pl.read(position_ids, [scatter_t]) + token_ape_row = pl.cast(token_pos % COMPRESS_RATIO, target_type=pl.INDEX) + ape_row = ape[token_ape_row : token_ape_row + 1, 0:OUT_DIM] + kv_row = kv_proj[scatter_t : scatter_t + 1, 0:OUT_DIM] + score_row = pl.add(score_proj[scatter_t : scatter_t + 1, 0:OUT_DIM], ape_row) + compress_state_rows[state_row : state_row + 1, 0:OUT_DIM] = kv_row + compress_state_rows[state_row : state_row + 1, OUT_DIM:COMPRESS_STATE_DIM] = score_row + + # 2. Zero the pooled scratch coarsely (real rows overwritten by the pool via + # the init_tid dep). write_rows is a multiple of CORE_INIT_ROW_TILE. + with pl.spmd(write_rows // CORE_INIT_ROW_TILE, name_hint="pooled_pad_init") as init_tid: + init_r0 = pl.tile.get_block_idx() * CORE_INIT_ROW_TILE + pooled_kv[init_r0 : init_r0 + CORE_INIT_ROW_TILE, 0:HEAD_DIM] = pl.full( + [CORE_INIT_ROW_TILE, HEAD_DIM], dtype=pl.FP32, value=0.0 + ) + + # 3. Softmax-pool each real window from the paged state. Iterate the compact + # pool_row_map (one entry per real window) -> padded write row; decode's + # scattered RMS padding thus costs no extra pool tasks, prefill passes identity. + with pl.spmd(pool_rows * POOL_HEAD_BLOCKS_CORE, name_hint="softmax_pool", deps=[scatter_tid, init_tid]) as pool_tid: + idx = pl.tile.get_block_idx() + pool_p = idx // POOL_HEAD_BLOCKS_CORE + h0 = (idx % POOL_HEAD_BLOCKS_CORE) * POOL_HEAD_TILE + write_row = pl.cast(pl.read(pool_row_map, [0, pool_p]), target_type=pl.INDEX) + write_slot_raw = pl.read(write_dst_map, [0, write_row]) + if write_slot_raw >= 0: + write_pos = pl.read(write_pos_map, [0, write_row]) + table_row = pl.cast(pl.read(state_table_row_map, [0, write_row]), target_type=pl.INDEX) + pooled_chunk = compressor_ratio4_pool_window( + compress_state_rows, + compress_state_block_table, + table_row, + write_pos, + h0, + ) + pooled_kv[write_row : write_row + 1, h0 : h0 + POOL_HEAD_TILE] = pooled_chunk + + # 4. RoPE tables for each window position, then rmsnorm + rope -> normed_kv. + gather_compressor_rope_rows( + freqs_cos, + freqs_sin, + write_pos_map, + write_dst_map, + pl.const(COMPRESS_RATIO, pl.INT32), + cos_b, + sin_b, + ) + normed_kv = compressor_rmsnorm_rope(pooled_kv, norm_w, cos_b, sin_b, normed_kv) + return normed_kv + + +# Decode shape and paging contract. +DECODE_B = DECODE_BATCH +DECODE_S = DECODE_SEQ +DECODE_T = DECODE_B * DECODE_S +DECODE_IDX_KV_LEN = MAX_SEQ_LEN // COMPRESS_RATIO +DECODE_COMPRESS_STATE_BLOCK_SIZE = RATIO4_STATE_BLOCK_SIZE +DECODE_COMPRESS_STATE_PHYSICAL_BLOCKS = 65 +DECODE_COMPRESS_STATE_MAX_BLOCKS = (MAX_SEQ_LEN + DECODE_COMPRESS_STATE_BLOCK_SIZE - 1) // DECODE_COMPRESS_STATE_BLOCK_SIZE +DECODE_COMPRESS_STATE_BLOCK_NUM = DECODE_B * DECODE_COMPRESS_STATE_PHYSICAL_BLOCKS +DECODE_COMPRESSOR_CMP_MAX_BLOCKS = KV_CMP_MAX_BLOCKS +DECODE_COMPRESSOR_CMP_BLOCK_NUM = DECODE_CMP_BLOCK_NUM + +# Decode tiling. +DECODE_ROPE_TILE = 32 +DECODE_K_TILE = 512 +DECODE_OUT_TILE = 64 +DECODE_B_TILE = 8 +DECODE_MM_B_TILE = 16 +DECODE_BS_PAD = ((DECODE_B * DECODE_S + DECODE_MM_B_TILE - 1) // DECODE_MM_B_TILE) * DECODE_MM_B_TILE +DECODE_HEAD_TILE = 64 +DECODE_HEAD_DIM_TILE = 128 +DECODE_RMS_PAD_TILE = 16 # pad DECODE_B rows into one 16-row block (min M for FP32 vec ops) +DECODE_RMS_PAD_ROWS = DECODE_RMS_PAD_TILE # single block; requires DECODE_B <= RMS_PAD_TILE +assert DECODE_B <= DECODE_RMS_PAD_TILE + +@pl.jit.inline +def decode_compressor_ratio4( + x: pl.Tensor[[DECODE_T, D], pl.BF16], + kv: pl.Tensor[[DECODE_T, HEAD_DIM], pl.FP32], + compress_state: pl.Tensor[[DECODE_COMPRESS_STATE_BLOCK_NUM, DECODE_COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[DECODE_B, DECODE_COMPRESS_STATE_MAX_BLOCKS], pl.INT32], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], + norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + cmp_kv_cache: pl.Tensor[[DECODE_COMPRESSOR_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16], + position_ids: pl.Tensor[[DECODE_T], pl.INT32], + cmp_slot_mapping: pl.Tensor[[DECODE_T], pl.INT64], + state_slot_mapping: pl.Tensor[[DECODE_T], pl.INT64], +): + # Decode fuses the paged CSA compressor into two tasks (upstream #734 single-block + # form). scatter_softmax_pool computes each batch's overlap window from its own + # just-scattered paged state (per-batch block table, so no cross-task barrier); + # rmsnorm_rope_cache_write normalizes the single 16-row block and writes the + # per-token kv output + paged cmp_kv_cache. Prefill keeps the multi-write shared + # core (compressor_core_ratio4): a single paged block and packed multi-writes + # diverge too much to share the fused path. Reads stay token-major ([T]). + kv_flat = kv + compress_state_flat = pl.reshape(compress_state, [DECODE_COMPRESS_STATE_BLOCK_NUM * DECODE_COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM]) + cmp_kv_cache_flat = pl.reshape(cmp_kv_cache, [DECODE_COMPRESSOR_CMP_BLOCK_NUM * BLOCK_SIZE, HEAD_DIM]) + + cmp4_kv_proj_pad = pl.create_tensor([DECODE_BS_PAD, OUT_DIM], dtype=pl.FP32) + cmp4_score_proj_pad = pl.create_tensor([DECODE_BS_PAD, OUT_DIM], dtype=pl.FP32) + compressor_ratio4_proj(x, wkv, wgate, cmp4_kv_proj_pad, cmp4_score_proj_pad) + + # scatter_softmax_pool: per batch, scatter the padded proj rows into the paged + # compress_state, then online-softmax pool that batch's window into pooled_kv. + # One region -- each batch's pool reads only its own just-scattered state. + pooled_kv = pl.create_tensor([DECODE_RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) + with pl.at(level=pl.Level.CORE_GROUP, name_hint="scatter_softmax_pool"): + for c_idx in pl.range(DECODE_B): + for s_sc in pl.pipeline(DECODE_S, stage=2): + proj_row = c_idx * DECODE_S + s_sc + token_pos = pl.read(position_ids, [proj_row]) + state_row_i64 = pl.read(state_slot_mapping, [proj_row]) + token_ape_row = pl.cast(token_pos % COMPRESS_RATIO, target_type=pl.INDEX) + if state_row_i64 >= 0: + state_row = pl.cast(state_row_i64, pl.INDEX) + kv_tile = cmp4_kv_proj_pad[proj_row : proj_row + 1, 0 : OUT_DIM] + score_tile = cmp4_score_proj_pad[proj_row : proj_row + 1, 0 : OUT_DIM] + ape_tile = ape[token_ape_row : token_ape_row + 1, 0 : OUT_DIM] + score_tile = pl.add(score_tile, ape_tile) + compress_state_flat[state_row : state_row + 1, 0 : OUT_DIM] = kv_tile + compress_state_flat[state_row : state_row + 1, OUT_DIM : COMPRESS_STATE_DIM] = score_tile + + pad_idx = c_idx + first_pos_b = pl.read(position_ids, [c_idx * DECODE_S]) + pos_b = first_pos_b % COMPRESS_RATIO + pre_tokens_b = COMPRESS_RATIO - pos_b + boundary_end_b = first_pos_b + pre_tokens_b - 1 + cur_window_start_b = boundary_end_b - COMPRESS_RATIO + 1 + prev_window_start_b = cur_window_start_b - COMPRESS_RATIO + + if pos_b + DECODE_S >= COMPRESS_RATIO: + # Head-chunk loop collapsed to one [1, HEAD_DIM] tile: the online + # softmax is per-column elementwise, so widening is bit-identical. + last_abs = cur_window_start_b + COMPRESS_RATIO - 1 + last_blk_off = last_abs // DECODE_COMPRESS_STATE_BLOCK_SIZE + last_intra = last_abs % DECODE_COMPRESS_STATE_BLOCK_SIZE + last_blk_id = pl.cast(pl.read(compress_state_block_table, [c_idx, last_blk_off]), pl.INDEX) + last_row = last_blk_id * DECODE_COMPRESS_STATE_BLOCK_SIZE + last_intra + mi = compress_state_flat[last_row : last_row + 1, OUT_DIM + HEAD_DIM : COMPRESS_STATE_DIM] + li = pl.exp(pl.sub(mi, mi)) + oi = compress_state_flat[last_row : last_row + 1, HEAD_DIM : OUT_DIM] + + for s in pl.range(0, COMPRESS_RATIO): + prev_abs = prev_window_start_b + s + front_score = pl.full([1, HEAD_DIM], dtype=pl.FP32, value=FP32_NEG_INF) + front_kv = pl.full([1, HEAD_DIM], dtype=pl.FP32, value=0.0) + if first_pos_b >= COMPRESS_RATIO: + prev_blk_off = prev_abs // DECODE_COMPRESS_STATE_BLOCK_SIZE + prev_intra = prev_abs % DECODE_COMPRESS_STATE_BLOCK_SIZE + prev_blk_id = pl.cast(pl.read(compress_state_block_table, [c_idx, prev_blk_off]), pl.INDEX) + prev_row = prev_blk_id * DECODE_COMPRESS_STATE_BLOCK_SIZE + prev_intra + front_score = compress_state_flat[prev_row : prev_row + 1, OUT_DIM : OUT_DIM + HEAD_DIM] + front_kv = compress_state_flat[prev_row : prev_row + 1, 0 : HEAD_DIM] + mi_next_front = pl.maximum(mi, front_score) + alpha_front = pl.exp(pl.sub(mi, mi_next_front)) + beta_front = pl.exp(pl.sub(front_score, mi_next_front)) + li = pl.add(pl.mul(alpha_front, li), beta_front) + oi = pl.add(pl.mul(oi, alpha_front), pl.mul(front_kv, beta_front)) + mi = mi_next_front + + for s in pl.range(0, COMPRESS_RATIO - 1): + cur_abs = cur_window_start_b + s + cur_blk_off = cur_abs // DECODE_COMPRESS_STATE_BLOCK_SIZE + cur_intra = cur_abs % DECODE_COMPRESS_STATE_BLOCK_SIZE + cur_blk_id = pl.cast(pl.read(compress_state_block_table, [c_idx, cur_blk_off]), pl.INDEX) + cur_row = cur_blk_id * DECODE_COMPRESS_STATE_BLOCK_SIZE + cur_intra + back_score = compress_state_flat[cur_row : cur_row + 1, OUT_DIM + HEAD_DIM : COMPRESS_STATE_DIM] + back_kv = compress_state_flat[cur_row : cur_row + 1, HEAD_DIM : OUT_DIM] + mi_next_back = pl.maximum(mi, back_score) + alpha_back = pl.exp(pl.sub(mi, mi_next_back)) + beta_back = pl.exp(pl.sub(back_score, mi_next_back)) + li = pl.add(pl.mul(alpha_back, li), beta_back) + oi = pl.add(pl.mul(oi, alpha_back), pl.mul(back_kv, beta_back)) + mi = mi_next_back + + pooled_chunk = pl.div(oi, li) + pooled_kv[pad_idx : pad_idx + 1, 0 : HEAD_DIM] = pooled_chunk + + # rmsnorm_rope_cache_write: normalize + rope the single 16-row block, then write + # the per-token kv output + paged cmp_kv_cache. The cache write reads back only + # this block's own normed_kv rows (intra-block RAW) -- no separate scope needed. + normed_kv = pl.create_tensor([DECODE_RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) + norm_w_2d = pl.reshape(norm_w, [1, HEAD_DIM]) + with pl.at(level=pl.Level.CORE_GROUP, name_hint="rmsnorm_rope_cache_write"): + # single 16-row block: DECODE_B real rows at 0..DECODE_B-1, rest are pad. + # In-kernel token-major gather of each batch's compressor rope row. + cos_b = pl.full([DECODE_RMS_PAD_TILE, ROPE_HALF], dtype=pl.FP32, value=0.0) + sin_b = pl.full([DECODE_RMS_PAD_TILE, ROPE_HALF], dtype=pl.FP32, value=0.0) + for inner in pl.range(DECODE_B): + first_pos_b = pl.read(position_ids, [inner * DECODE_S]) + cmp_pos_b = pl.cast(first_pos_b - (first_pos_b % COMPRESS_RATIO), pl.INDEX) + cos_b[inner : inner + 1, 0 : ROPE_HALF] = pl.cast( + freqs_cos[cmp_pos_b : cmp_pos_b + 1, 0 : ROPE_HALF], target_type=pl.FP32) + sin_b[inner : inner + 1, 0 : ROPE_HALF] = pl.cast( + freqs_sin[cmp_pos_b : cmp_pos_b + 1, 0 : ROPE_HALF], target_type=pl.FP32) + + partial_sq = pl.full([1, DECODE_RMS_PAD_TILE], dtype=pl.FP32, value=0.0) + for k0 in pl.range(0, HEAD_DIM, DECODE_HEAD_TILE): + kv_rms_chunk = pooled_kv[0 : DECODE_RMS_PAD_TILE, k0 : k0 + DECODE_HEAD_TILE] + kv_rms_sq = pl.mul(kv_rms_chunk, kv_rms_chunk) + kv_rms_rowsum = pl.reshape(pl.row_sum(kv_rms_sq), [1, DECODE_RMS_PAD_TILE]) + partial_sq = pl.add(partial_sq, kv_rms_rowsum) + + variance = pl.reshape(pl.add(pl.mul(partial_sq, HEAD_DIM_INV), EPS), [DECODE_RMS_PAD_TILE, 1]) + inv_rms = pl.recip(pl.sqrt(variance)) + for k0 in pl.range(0, NOPE_HEAD_DIM, DECODE_HEAD_TILE): + kv_norm_chunk = pooled_kv[0 : DECODE_RMS_PAD_TILE, k0 : k0 + DECODE_HEAD_TILE] + gamma = pl.cast(norm_w_2d[:, k0 : k0 + DECODE_HEAD_TILE], pl.FP32) + normed_chunk = pl.col_expand_mul(pl.row_expand_mul(kv_norm_chunk, inv_rms), gamma) + normed_kv[0 : DECODE_RMS_PAD_TILE, k0 : k0 + DECODE_HEAD_TILE] = normed_chunk + + kv_rope_norm = pooled_kv[0 : DECODE_RMS_PAD_TILE, NOPE_HEAD_DIM : HEAD_DIM] + gamma_rope = pl.cast(norm_w_2d[:, NOPE_HEAD_DIM : HEAD_DIM], pl.FP32) + # A3 interleaved swap-gather rope (swap/sign/dup indices built in-kernel from + # pl.arange): out[j] = n[j]*cos_il[j] + n[j^1]*sign[j]*sin_il[j]. normed_kv is + # FP32 so rope_rot is written directly. + rope_normed = pl.col_expand_mul(pl.row_expand_mul(kv_rope_norm, inv_rms), gamma_rope) + rope_ones = pl.full([DECODE_RMS_PAD_TILE, ROPE_HEAD_DIM], dtype=pl.FP32, value=1.0) + rope_col = pl.col_expand_mul(rope_ones, pl.cast(pl.arange(0, [1, ROPE_HEAD_DIM], dtype=pl.INT32), target_type=pl.FP32)) + rope_dup_f = pl.cast(pl.cast(pl.mul(rope_col, 0.5), target_type=pl.INT32, mode="trunc"), target_type=pl.FP32) + rope_dup_idx = pl.cast(rope_dup_f, target_type=pl.INT32) + rope_lane = pl.sub(rope_col, pl.mul(rope_dup_f, 2.0)) + rope_swap_idx = pl.cast(pl.sub(pl.add(rope_col, 1.0), pl.mul(rope_lane, 2.0)), target_type=pl.INT32) + rope_sign = pl.sub(pl.mul(rope_lane, 2.0), 1.0) + cos_il = pl.gather(cos_b, dim=-1, index=rope_dup_idx) + sin_il = pl.gather(sin_b, dim=-1, index=rope_dup_idx) + swapped = pl.gather(rope_normed, dim=-1, index=rope_swap_idx) + rope_rot = pl.add(pl.mul(rope_normed, cos_il), pl.mul(pl.mul(swapped, rope_sign), sin_il)) + normed_kv[0 : DECODE_RMS_PAD_TILE, NOPE_HEAD_DIM : HEAD_DIM] = rope_rot + + for inner in pl.range(DECODE_B): + c_idx = inner + first_pos_b = pl.read(position_ids, [c_idx * DECODE_S]) + pos_b = first_pos_b % COMPRESS_RATIO + if pos_b + DECODE_S >= COMPRESS_RATIO: + boundary_s = COMPRESS_RATIO - 1 - pos_b + kv_row_fp32 = normed_kv[inner : inner + 1, 0 : HEAD_DIM] + cache_row_i64 = pl.read(cmp_slot_mapping, [c_idx * DECODE_S + boundary_s]) + if cache_row_i64 >= 0: + cache_row = pl.cast(cache_row_i64, pl.INDEX) + kv_flat[c_idx * DECODE_S : c_idx * DECODE_S + 1, :] = kv_row_fp32 + cmp_kv_cache_flat[cache_row : cache_row + 1, :] = pl.cast(kv_row_fp32, target_type=pl.BF16, mode="rint") + + return kv_flat + + +@pl.jit +def decode_compressor_ratio4_test( + x: pl.Tensor[[DECODE_T, D], pl.BF16], + kv: pl.Out[pl.Tensor[[DECODE_T, HEAD_DIM], pl.FP32]], + compress_state: pl.InOut[pl.Tensor[[DECODE_COMPRESS_STATE_BLOCK_NUM, DECODE_COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32]], + compress_state_block_table: pl.Tensor[[DECODE_B, DECODE_COMPRESS_STATE_MAX_BLOCKS], pl.INT32], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], + norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + cmp_kv_cache: pl.InOut[pl.Tensor[[DECODE_COMPRESSOR_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], + position_ids: pl.Tensor[[DECODE_T], pl.INT32], + cmp_slot_mapping: pl.Tensor[[DECODE_T], pl.INT64], + state_slot_mapping: pl.Tensor[[DECODE_T], pl.INT64], +): + decode_compressor_ratio4( + x, + kv, + compress_state, + compress_state_block_table, + wkv, + wgate, + ape, + norm_w, + freqs_cos, + freqs_sin, + cmp_kv_cache, + position_ids, + cmp_slot_mapping, + state_slot_mapping, + ) + return kv, compress_state, cmp_kv_cache + + +def golden_decode_compressor_ratio4(tensors): + """Torch reference for Compressor.forward (decode branch, ratio=4 overlap).""" + import torch + + x = tensors["x"].float().reshape(DECODE_B, DECODE_S, D) + compress_state = tensors["compress_state"] + compress_state_block_table = tensors["compress_state_block_table"] + wkv = tensors["wkv"].float() + wgate = tensors["wgate"].float() + ape = tensors["ape"] + norm_w = tensors["norm_w"] + freqs_cos = tensors["freqs_cos"] + freqs_sin = tensors["freqs_sin"] + cmp_kv_cache = tensors["cmp_kv_cache"] + position_ids = tensors["position_ids"].reshape(DECODE_B, DECODE_S).to(torch.int64) + cmp_slot_mapping = tensors["cmp_slot_mapping"].reshape(DECODE_B, DECODE_S).to(torch.int64) + state_slot_mapping = tensors["state_slot_mapping"].reshape(DECODE_B, DECODE_S).to(torch.int64) + bsz, _, _ = x.shape + ratio, rd = COMPRESS_RATIO, ROPE_HEAD_DIM + + kv = x @ wkv.t() # [DECODE_B, DECODE_S, OUT_DIM] (wkv stored [OUT_DIM, D] for b_trans) + score = x @ wgate.t() # [DECODE_B, DECODE_S, OUT_DIM] + + pooled = torch.zeros(bsz, 1, HEAD_DIM, dtype=torch.float32, device=x.device) + should_compress_rows = torch.zeros(bsz, dtype=torch.bool, device=x.device) + + def read_front_state(b, abs_pos): + blk_id = int(compress_state_block_table[b, abs_pos // DECODE_COMPRESS_STATE_BLOCK_SIZE].item()) + if blk_id < 0: + return ( + torch.zeros(HEAD_DIM, dtype=torch.float32, device=x.device), + torch.full((HEAD_DIM,), float("-inf"), dtype=torch.float32, device=x.device), + ) + intra = abs_pos % DECODE_COMPRESS_STATE_BLOCK_SIZE + return ( + compress_state[blk_id, intra, :HEAD_DIM], + compress_state[blk_id, intra, OUT_DIM:OUT_DIM + HEAD_DIM], + ) + + def read_back_state(b, abs_pos): + blk_id = int(compress_state_block_table[b, abs_pos // DECODE_COMPRESS_STATE_BLOCK_SIZE].item()) + if blk_id < 0: + return ( + torch.zeros(HEAD_DIM, dtype=torch.float32, device=x.device), + torch.full((HEAD_DIM,), float("-inf"), dtype=torch.float32, device=x.device), + ) + intra = abs_pos % DECODE_COMPRESS_STATE_BLOCK_SIZE + return ( + compress_state[blk_id, intra, HEAD_DIM:OUT_DIM], + compress_state[blk_id, intra, OUT_DIM + HEAD_DIM:], + ) + + for b in range(bsz): + first_pos = int(position_ids[b, 0].item()) + pre_tokens = min(DECODE_S, ratio - (first_pos % ratio)) + boundary_s = ratio - 1 - (first_pos % ratio) + should_compress = 0 <= boundary_s < DECODE_S + boundary_end = first_pos + pre_tokens - 1 + cur_window_start = boundary_end - ratio + 1 + prev_window_start = cur_window_start - ratio + + # Per-token ape add + state scatter through explicit token-major slots. + for s in range(DECODE_S): + pos = int(position_ids[b, s].item()) + token_ape_row = pos % ratio + score[b, s, :] = score[b, s, :] + ape[token_ape_row] + state_row = int(state_slot_mapping[b, s].item()) + if state_row >= 0: + blk_id = state_row // DECODE_COMPRESS_STATE_BLOCK_SIZE + intra = state_row % DECODE_COMPRESS_STATE_BLOCK_SIZE + compress_state[blk_id, intra, :OUT_DIM] = kv[b, s, :] + compress_state[blk_id, intra, OUT_DIM:] = score[b, s, :] + + if should_compress: + should_compress_rows[b] = True + kv_rows = [] + score_rows = [] + for s in range(ratio): + abs_pos = prev_window_start + s + if abs_pos < 0: + kv_rows.append(torch.zeros(HEAD_DIM, dtype=torch.float32, device=x.device)) + score_rows.append(torch.full((HEAD_DIM,), float("-inf"), dtype=torch.float32, device=x.device)) + continue + kv_row, score_row = read_front_state(b, abs_pos) + kv_rows.append(kv_row) + score_rows.append(score_row) + for s in range(ratio): + abs_pos = cur_window_start + s + kv_row, score_row = read_back_state(b, abs_pos) + kv_rows.append(kv_row) + score_rows.append(score_row) + kvs = torch.stack(kv_rows, dim=0).unsqueeze(0) + scs = torch.stack(score_rows, dim=0).unsqueeze(0) + pooled[b : b + 1] = (kvs * scs.softmax(dim=1)).sum(dim=1, keepdim=True) + + tensors["compress_state"][:] = compress_state + + if not bool(should_compress_rows.any()): + return + + def rmsnorm(x, w): + x = x.float() + var = x.square().mean(-1, keepdim=True) + x = x * torch.rsqrt(var + EPS) + return w * x + + for b in range(bsz): + if not bool(should_compress_rows[b]): + continue + first_pos = int(position_ids[b, 0].item()) + boundary_s = ratio - 1 - (first_pos % ratio) + kv_b = rmsnorm(pooled[b : b + 1], norm_w) + + x_pair = kv_b[..., -rd:].unflatten(-1, (-1, 2)) + x0, x1 = x_pair[..., 0], x_pair[..., 1] + # cos/sin from the shared freqs table at the compression-window origin, matching + # the in-kernel gather (window_start = first_pos - first_pos % ratio). freqs_* is + # BF16; .float() replicates the kernel's BF16->FP32 cast. + window_start_b = first_pos - (first_pos % ratio) + cos_v = freqs_cos[window_start_b, : rd // 2].float().view(-1) + sin_v = freqs_sin[window_start_b, : rd // 2].float().view(-1) + y0 = x0 * cos_v - x1 * sin_v + y1 = x0 * sin_v + x1 * cos_v + + kv_b = torch.cat([kv_b[..., :-rd], torch.stack([y0, y1], dim=-1).flatten(-2)], dim=-1) + + cmp_row = int(cmp_slot_mapping[b, boundary_s].item()) + if cmp_row >= 0: + # Kernel writes the committed pooled result to the sequence's first + # token row (t = b * DECODE_S); non-boundary batches and other token rows + # stay at the output tensor's zero-init. + tensors["kv"][b * DECODE_S : b * DECODE_S + 1, :] = kv_b.reshape(1, HEAD_DIM) + blk_id = cmp_row // BLOCK_SIZE + cmp_kv_cache[blk_id, cmp_row % BLOCK_SIZE, 0] = kv_b[0, 0] + + tensors["cmp_kv_cache"][:] = cmp_kv_cache + + +def build_decode_tensor_specs(start_pos=None): + import torch # type: ignore[import] + from decode_metadata import ( + block_table, + compressed_slot_mapping, + csa_decode_start_set, + position_ids_from_starts, + resolve_start_positions, + state_slot_mapping, + ) + from golden import TensorSpec + from rope_tables import build_deepseek_v4_rope_tables + + shared_freqs_cos, shared_freqs_sin = build_deepseek_v4_rope_tables(M, COMPRESS_RATIO, dtype=torch.bfloat16) + + def init_x(): + return torch.rand(DECODE_B, DECODE_S, D).reshape(DECODE_T, D) + def init_compress_state(): + state = torch.zeros(DECODE_COMPRESS_STATE_BLOCK_NUM, DECODE_COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM) + state[:, :, OUT_DIM:] = FP32_NEG_INF + return state + def init_compress_state_block_table(): + return block_table( + batch=DECODE_B, + table_blocks=DECODE_COMPRESS_STATE_MAX_BLOCKS, + physical_blocks=DECODE_COMPRESS_STATE_PHYSICAL_BLOCKS, + ) + # Calibrated to the real DeepSeek-V4-Flash CSA (ratio-4) main compressor (mean l8/l32 of + # extract_weights_flash): zero-mean Gaussian BF16 weights at the measured std; the RMSNorm + # gamma centers near the measured mean (not ones / not uniform). + def init_wkv(): + return torch.randn(OUT_DIM, D) * 0.0245 + def init_wgate(): + return torch.randn(OUT_DIM, D) * 0.0388 + def init_ape(): + return torch.randn(COMPRESS_RATIO, OUT_DIM) * 0.1243 + def init_norm_w(): + return 0.9666 + 0.1929 * torch.randn(HEAD_DIM) + def init_freqs_cos(): + return shared_freqs_cos.clone() + def init_freqs_sin(): + return shared_freqs_sin.clone() + def init_cmp_kv_cache(): + return torch.zeros(DECODE_COMPRESSOR_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM) + def init_cmp_block_table(): + tbl = torch.full((DECODE_B, DECODE_COMPRESSOR_CMP_MAX_BLOCKS), -1, dtype=torch.int32) + for b in range(DECODE_B): + for j in range(DECODE_COMPRESSOR_CMP_MAX_BLOCKS): + tbl[b, j] = b * DECODE_COMPRESSOR_CMP_MAX_BLOCKS + j + return tbl + def init_default_start_pos(): + # Canonical CSA start-position set (ratio-4 compressor + indexer + sliding-window + 8k). + return csa_decode_start_set( + batch=DECODE_B, seq=DECODE_S, compress_ratio=COMPRESS_RATIO, + state_block_size=DECODE_COMPRESS_STATE_BLOCK_SIZE) + def init_start_pos(): + return resolve_start_positions( + start_pos, + batch=DECODE_B, + seq=DECODE_S, + max_seq_len=MAX_SEQ_LEN, + default_fn=init_default_start_pos, + ) + def _position_ids_bs(): + # [DECODE_B, DECODE_S] positions; token-major init_position_ids reshapes this to [DECODE_T]. + return position_ids_from_starts(init_start_pos(), seq=DECODE_S) + def init_position_ids(): + # token-major [DECODE_T]; row t == (b * DECODE_S + s) carries position of sequence b, intra-token s + return _position_ids_bs().reshape(DECODE_T) + def init_state_slot_mapping(): + return state_slot_mapping( + _position_ids_bs(), + init_compress_state_block_table(), + state_block_size=DECODE_COMPRESS_STATE_BLOCK_SIZE, + ).reshape(DECODE_T) + def init_cmp_slot_mapping(): + return compressed_slot_mapping( + _position_ids_bs(), + init_cmp_block_table(), + compress_ratio=COMPRESS_RATIO, + block_size=BLOCK_SIZE, + ).reshape(DECODE_T) + + return [ + TensorSpec("x", [DECODE_T, D], torch.bfloat16, init_value=init_x), + TensorSpec("kv", [DECODE_T, HEAD_DIM], torch.float32, is_output=True), + TensorSpec("compress_state", [DECODE_COMPRESS_STATE_BLOCK_NUM, DECODE_COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], torch.float32, init_value=init_compress_state, is_output=True), + TensorSpec("compress_state_block_table", [DECODE_B, DECODE_COMPRESS_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), + TensorSpec("wkv", [OUT_DIM, D], torch.bfloat16, init_value=init_wkv), + TensorSpec("wgate", [OUT_DIM, D], torch.bfloat16, init_value=init_wgate), + TensorSpec("ape", [COMPRESS_RATIO, OUT_DIM], torch.float32, init_value=init_ape), + TensorSpec("norm_w", [HEAD_DIM], torch.bfloat16, init_value=init_norm_w), + TensorSpec("freqs_cos", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_cos), + TensorSpec("freqs_sin", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_sin), + TensorSpec("cmp_kv_cache", [DECODE_COMPRESSOR_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], torch.bfloat16, init_value=init_cmp_kv_cache, is_output=True), + TensorSpec("position_ids", [DECODE_T], torch.int32, init_value=init_position_ids), + TensorSpec("cmp_slot_mapping", [DECODE_T], torch.int64, init_value=init_cmp_slot_mapping), + TensorSpec("state_slot_mapping", [DECODE_T], torch.int64, init_value=init_state_slot_mapping), + ] + + +# Prefill shape and paging contract. +PREFILL_B = 1 +PREFILL_S = 128 +PREFILL_START_POS = 0 +PREFILL_COMPRESSED_LEN = PREFILL_S // COMPRESS_RATIO +PREFILL_ROWS = PREFILL_B * PREFILL_COMPRESSED_LEN +assert HEAD_DIM % POOL_HEAD_TILE == 0 +POOL_HEAD_BLOCKS = HEAD_DIM // POOL_HEAD_TILE +K_TILE = 512 +OUT_TILE = 32 +HEAD_TILE = 64 +RMS_TILE = 16 + +PREFILL_T = PREFILL_B * PREFILL_S +CSA_STATE_BLOCK_SIZE = RATIO4_STATE_BLOCK_SIZE +CSA_STATE_MAX_BLOCKS = (MAX_SEQ_LEN + CSA_STATE_BLOCK_SIZE - 1) // CSA_STATE_BLOCK_SIZE +CSA_STATE_BLOCK_NUM = CSA_STATE_MAX_BLOCKS +MAX_CMP_WRITES = max(1, PREFILL_T // COMPRESS_RATIO) +PACKED_PROJ_BLOCKS = OUT_DIM // OUT_TILE +PACKED_POOL_BLOCKS = MAX_CMP_WRITES * POOL_HEAD_BLOCKS +PACKED_STATE_UPDATE_TILE = 16 +PACKED_RMS_TILE = 16 + + +@pl.jit.inline +def prefill_compressor_ratio4_proj( + x: pl.Tensor[[PREFILL_T, D], pl.BF16], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + kv_proj_out: pl.Tensor[[PREFILL_T, OUT_DIM], pl.FP32], + score_proj_out: pl.Tensor[[PREFILL_T, OUT_DIM], pl.FP32], +): + for proj_idx in pl.spmd(PACKED_PROJ_BLOCKS, name_hint="prefill_c4_kv_score_proj"): + o0 = proj_idx * OUT_TILE + kv_acc = pl.create_tensor([PREFILL_T, OUT_TILE], dtype=pl.FP32) + score_acc = pl.create_tensor([PREFILL_T, OUT_TILE], dtype=pl.FP32) + for kb in pl.pipeline(0, D // K_TILE, stage=2): + k0 = kb * K_TILE + x_tile = x[0:PREFILL_T, k0 : k0 + K_TILE] + wkv_tile = wkv[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] + wgate_tile = wgate[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] + if k0 == 0: + kv_acc = pl.matmul(x_tile, wkv_tile, out_dtype=pl.FP32, b_trans=True) + score_acc = pl.matmul(x_tile, wgate_tile, out_dtype=pl.FP32, b_trans=True) + else: + kv_acc = pl.matmul_acc(kv_acc, x_tile, wkv_tile, b_trans=True) + score_acc = pl.matmul_acc(score_acc, x_tile, wgate_tile, b_trans=True) + kv_proj_out[0:PREFILL_T, o0 : o0 + OUT_TILE] = kv_acc + score_proj_out[0:PREFILL_T, o0 : o0 + OUT_TILE] = score_acc + + +@pl.jit.inline +def prefill_compressor_ratio4( + x: pl.Tensor[[PREFILL_T, D], pl.BF16], + compress_state: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[PREFILL_B, CSA_STATE_MAX_BLOCKS], pl.INT32], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], + norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + cmp_kv: pl.Tensor[[PREFILL_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16], + position_ids: pl.Tensor[[PREFILL_T], pl.INT32], + num_tokens: pl.Scalar[pl.INT32], + cmp_slot_mapping: pl.Tensor[[PREFILL_T], pl.INT64], + state_slot_mapping: pl.Tensor[[PREFILL_T], pl.INT64], +): + # Thin prefill wrapper: build the num_tokens-bounded write schedule, project + # (static whole-M tiling), run the shared core with the fresh prefill state, + # then write cmp_kv (+ keepalive). num_tokens is consumed only by the schedule + # builder; the core gates scatter on state_slot_mapping >= 0. See + # compressor_core_ratio4. + compress_state_flat = pl.reshape(compress_state, [CSA_STATE_BLOCK_NUM * CSA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM]) + cmp_kv_flat = pl.reshape(cmp_kv, [PREFILL_CMP_BLOCK_NUM * BLOCK_SIZE, HEAD_DIM]) + + write_pos_map = pl.create_tensor([1, MAX_CMP_WRITES], dtype=pl.INT32) + write_dst_map = pl.create_tensor([1, MAX_CMP_WRITES], dtype=pl.INT32) + state_table_row_map = pl.create_tensor([1, MAX_CMP_WRITES], dtype=pl.INT32) + build_prefill_write_schedule( + position_ids, + cmp_slot_mapping, + num_tokens, + write_pos_map, + write_dst_map, + state_table_row_map, + ) + + # Prefill has no scattered RMS padding: its write rows are already compact, so + # the pool enumeration is the identity over the write rows (core behaviour + # unchanged for prefill). + pool_row_map = pl.create_tensor([1, MAX_CMP_WRITES], dtype=pl.INT32) + with pl.at(level=pl.Level.CORE_GROUP, name_hint="prefill_pool_row_map"): + for p in pl.range(MAX_CMP_WRITES): + pl.write(pool_row_map, [0, p], pl.cast(p, pl.INT32)) + + cmp4_kv_proj_scratch = pl.create_tensor([PREFILL_T, OUT_DIM], dtype=pl.FP32) + cmp4_score_proj_scratch = pl.create_tensor([PREFILL_T, OUT_DIM], dtype=pl.FP32) + prefill_compressor_ratio4_proj(x, wkv, wgate, cmp4_kv_proj_scratch, cmp4_score_proj_scratch) + + pooled_kv = pl.create_tensor([MAX_CMP_WRITES, HEAD_DIM], dtype=pl.FP32) + normed_kv = pl.create_tensor([MAX_CMP_WRITES, HEAD_DIM], dtype=pl.FP32) + cos_b = pl.create_tensor([MAX_CMP_WRITES, ROPE_HEAD_DIM // 2], dtype=pl.FP32) + sin_b = pl.create_tensor([MAX_CMP_WRITES, ROPE_HEAD_DIM // 2], dtype=pl.FP32) + normed_kv = compressor_core_ratio4( + cmp4_kv_proj_scratch, + cmp4_score_proj_scratch, + position_ids, + state_slot_mapping, + ape, + norm_w, + compress_state_flat, + compress_state_block_table, + freqs_cos, + freqs_sin, + write_pos_map, + write_dst_map, + state_table_row_map, + pool_row_map, + pooled_kv, + normed_kv, + cos_b, + sin_b, + ) + + cmp_kv_flat = finalize_compressor_writes( + normed_kv, + write_dst_map, + cmp_kv_flat, + pl.const(PACKED_RMS_TILE, pl.INT32), + pl.const(1, pl.INT32), + ) + + cmp_kv = pl.reshape(cmp_kv_flat, [PREFILL_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM]) + compress_state = pl.reshape(compress_state_flat, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM]) + return cmp_kv, compress_state + + +def golden_prefill_compressor_ratio4(tensors): + """Packed token-major torch reference for ratio-4 prefill compressor.""" + import torch + + x = tensors["x"].view(PREFILL_T, D).float() + compress_state_flat = tensors["compress_state"].view(CSA_STATE_BLOCK_NUM * CSA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM) + state_block_table = tensors["compress_state_block_table"] + wkv = tensors["wkv"].float() + wgate = tensors["wgate"].float() + ape = tensors["ape"] + norm_w = tensors["norm_w"] + cmp_kv = tensors["cmp_kv"] + cache_rows = cmp_kv.view(cmp_kv.shape[0] * BLOCK_SIZE, 1, HEAD_DIM)[:, 0, :] + position_ids = tensors["position_ids"] + + kv_proj = x @ wkv.t() # wkv stored [OUT_DIM, D] for b_trans + score_proj = x @ wgate.t() + + def state_row(abs_pos): + if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: + return -1 + block = abs_pos // CSA_STATE_BLOCK_SIZE + intra = abs_pos % CSA_STATE_BLOCK_SIZE + phys_block = int(state_block_table[0, block].item()) + if phys_block < 0: + return -1 + return phys_block * CSA_STATE_BLOCK_SIZE + intra + + for token_id in range(int(tensors["num_tokens"])): + dst_row = int(tensors["cmp_slot_mapping"][token_id].item()) + if dst_row < 0: + continue + write_pos = int(position_ids[token_id].item()) + cur_start = write_pos + 1 - COMPRESS_RATIO + prev_start = cur_start - COMPRESS_RATIO + pool_kv = torch.zeros(STATE_LEN, HEAD_DIM, dtype=torch.float32) + pool_score = torch.full((STATE_LEN, HEAD_DIM), float("-inf"), dtype=torch.float32) + + for s in range(COMPRESS_RATIO): + prev_abs = prev_start + s + if write_pos >= 2 * COMPRESS_RATIO - 1: + prev_row = state_row(prev_abs) + if prev_row >= 0: + pool_kv[s] = compress_state_flat[prev_row, :HEAD_DIM] + pool_score[s] = compress_state_flat[prev_row, OUT_DIM : OUT_DIM + HEAD_DIM] + + cur_abs = cur_start + s + cur_row = state_row(cur_abs) + if cur_row >= 0: + pool_kv[COMPRESS_RATIO + s] = compress_state_flat[cur_row, HEAD_DIM:OUT_DIM] + pool_score[COMPRESS_RATIO + s] = compress_state_flat[cur_row, OUT_DIM + HEAD_DIM : COMPRESS_STATE_DIM] + + for t in range(int(tensors["num_tokens"])): + pos = int(position_ids[t].item()) + if pos < prev_start or pos > write_pos: + continue + if pos < cur_start: + pool_slot = pos - prev_start + col0 = 0 + else: + pool_slot = COMPRESS_RATIO + pos - cur_start + col0 = HEAD_DIM + ape_slot = pos % COMPRESS_RATIO + pool_kv[pool_slot] = kv_proj[t, col0 : col0 + HEAD_DIM] + pool_score[pool_slot] = score_proj[t, col0 : col0 + HEAD_DIM] + ape[ape_slot, col0 : col0 + HEAD_DIM] + + init_slot = STATE_LEN - 1 + mi = pool_score[init_slot : init_slot + 1].clone() + li = torch.exp(mi - mi) + oi = pool_kv[init_slot : init_slot + 1].clone() + for slot_i in range(STATE_LEN - 1): + if slot_i < COMPRESS_RATIO and write_pos < 2 * COMPRESS_RATIO - 1: + continue + slot_score = pool_score[slot_i : slot_i + 1] + slot_kv = pool_kv[slot_i : slot_i + 1] + mi_next = torch.maximum(mi, slot_score) + alpha = torch.exp(mi - mi_next) + beta = torch.exp(slot_score - mi_next) + li = alpha * li + beta + oi = oi * alpha + slot_kv * beta + mi = mi_next + pooled = oi / li + inv_rms = torch.rsqrt(pooled.square().mean(dim=-1, keepdim=True) + EPS) + normed = pooled * inv_rms * norm_w.float().view(1, HEAD_DIM) + rope_pair = normed[..., NOPE_HEAD_DIM:HEAD_DIM].unflatten(-1, (-1, 2)) + rope_even = rope_pair[..., 0] + rope_odd = rope_pair[..., 1] + cmp_pos = write_pos + 1 - COMPRESS_RATIO + cos = tensors["freqs_cos"][cmp_pos : cmp_pos + 1, 0 : ROPE_HEAD_DIM // 2].float() + sin = tensors["freqs_sin"][cmp_pos : cmp_pos + 1, 0 : ROPE_HEAD_DIM // 2].float() + rot_even = rope_even * cos - rope_odd * sin + rot_odd = rope_even * sin + rope_odd * cos + normed[:, NOPE_HEAD_DIM:HEAD_DIM] = torch.stack([rot_even, rot_odd], dim=-1).flatten(-2) + cache_rows[dst_row] = normed.to(torch.bfloat16)[0] + + for t in range(int(tensors["num_tokens"])): + pos = int(tensors["position_ids"][t].item()) + dst_row = int(tensors["state_slot_mapping"][t].item()) + if dst_row < 0: + continue + ape_slot = pos % COMPRESS_RATIO + compress_state_flat[dst_row, 0:OUT_DIM] = kv_proj[t] + compress_state_flat[dst_row, OUT_DIM:COMPRESS_STATE_DIM] = score_proj[t] + tensors["ape"][ape_slot] + tensors["cmp_kv"][:] = cmp_kv + tensors["compress_state"][:] = compress_state_flat.view_as(tensors["compress_state"]) + + +@pl.jit +def prefill_compressor_ratio4_test( + x: pl.Tensor[[PREFILL_T, D], pl.BF16], + compress_state: pl.InOut[pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32]], + compress_state_block_table: pl.Tensor[[PREFILL_B, CSA_STATE_MAX_BLOCKS], pl.INT32], + wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], + wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], + ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], + norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + cmp_kv: pl.InOut[pl.Tensor[[PREFILL_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], + position_ids: pl.Tensor[[PREFILL_T], pl.INT32], + num_tokens: pl.Scalar[pl.INT32], + cmp_slot_mapping: pl.Tensor[[PREFILL_T], pl.INT64], + state_slot_mapping: pl.Tensor[[PREFILL_T], pl.INT64], +): + return prefill_compressor_ratio4( + x, compress_state, compress_state_block_table, wkv, wgate, ape, norm_w, freqs_cos, freqs_sin, + cmp_kv, position_ids, num_tokens, cmp_slot_mapping, state_slot_mapping, + ) + + +def build_prefill_tensor_specs(start_pos: int = PREFILL_START_POS): + import torch + from golden import ScalarSpec, TensorSpec + from rope_tables import build_deepseek_v4_rope_tables + + shared_freqs_cos, shared_freqs_sin = build_deepseek_v4_rope_tables(M, COMPRESS_RATIO, dtype=torch.bfloat16) + + if start_pos < 0 or start_pos + PREFILL_T > MAX_SEQ_LEN: + raise ValueError(f"start_pos must satisfy 0 <= start_pos <= {MAX_SEQ_LEN - PREFILL_T}, got {start_pos}") + + def init_compress_state_block_table(): + table = torch.full((PREFILL_B, CSA_STATE_MAX_BLOCKS), -1, dtype=torch.int32) + for block in range(CSA_STATE_MAX_BLOCKS): + table[0, block] = (block * 17 + 3) % CSA_STATE_MAX_BLOCKS + return table + def state_row(abs_pos): + if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: + return -1 + table = init_compress_state_block_table() + block = abs_pos // CSA_STATE_BLOCK_SIZE + intra = abs_pos % CSA_STATE_BLOCK_SIZE + return int(table[0, block].item()) * CSA_STATE_BLOCK_SIZE + intra + def init_x(): + return ((torch.rand(PREFILL_T, D) - 0.5) * 0.1).to(torch.bfloat16) + def init_compress_state(): + state = torch.zeros(CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM) + flat = state.view(-1, COMPRESS_STATE_DIM) + for abs_pos in range(max(0, start_pos - STATE_LEN), start_pos): + row = state_row(abs_pos) + if row >= 0: + flat[row, 0:OUT_DIM] = (torch.rand(OUT_DIM) - 0.5) * 0.05 + flat[row, OUT_DIM:COMPRESS_STATE_DIM] = (torch.rand(OUT_DIM) - 0.5) * 0.05 + return state + # Calibrated to the real DeepSeek-V4-Flash CSA (ratio-4) compressor (mean l8/l32 of + # extract_weights_flash): zero-mean Gaussian BF16 weights at the measured std; the RMSNorm + # gamma centers near the measured mean (not ones / not uniform). Mirrors the decode path. + def init_wkv(): + return torch.randn(OUT_DIM, D) * 0.0245 + def init_wgate(): + return torch.randn(OUT_DIM, D) * 0.0388 + def init_ape(): + return torch.randn(COMPRESS_RATIO, OUT_DIM) * 0.1243 + def init_norm_w(): + return 0.9666 + 0.1929 * torch.randn(HEAD_DIM) + def init_freqs_cos(): + return shared_freqs_cos.clone() + def init_freqs_sin(): + return shared_freqs_sin.clone() + def init_cmp_kv(): + return torch.zeros(PREFILL_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM, dtype=torch.bfloat16) + def init_position_ids(): + return torch.arange(start_pos, start_pos + PREFILL_T, dtype=torch.int32) + def init_cmp_slot_mapping(): + mapping = torch.full((PREFILL_T,), -1, dtype=torch.int64) + for t in range(PREFILL_T): + pos = start_pos + t + if (pos + 1) % COMPRESS_RATIO == 0: + dst_row = (pos + 1) // COMPRESS_RATIO - 1 + if dst_row >= PREFILL_CMP_BLOCK_NUM * BLOCK_SIZE: + raise ValueError("fixture compressed slot exceeds standalone cmp_kv capacity") + mapping[t] = dst_row + return mapping + def init_state_slot_mapping(): + mapping = torch.full((PREFILL_T,), -1, dtype=torch.int64) + for t in range(PREFILL_T): + mapping[t] = state_row(start_pos + t) + return mapping + + return [ + TensorSpec("x", [PREFILL_T, D], torch.bfloat16, init_value=init_x), + TensorSpec("compress_state", [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], torch.float32, init_value=init_compress_state, is_output=True), + TensorSpec("compress_state_block_table", [PREFILL_B, CSA_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), + TensorSpec("wkv", [OUT_DIM, D], torch.bfloat16, init_value=init_wkv), + TensorSpec("wgate", [OUT_DIM, D], torch.bfloat16, init_value=init_wgate), + TensorSpec("ape", [COMPRESS_RATIO, OUT_DIM], torch.float32, init_value=init_ape), + TensorSpec("norm_w", [HEAD_DIM], torch.bfloat16, init_value=init_norm_w), + TensorSpec("freqs_cos", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_cos), + TensorSpec("freqs_sin", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_sin), + TensorSpec("cmp_kv", [PREFILL_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], torch.bfloat16, init_value=init_cmp_kv, is_output=True), + TensorSpec("position_ids", [PREFILL_T], torch.int32, init_value=init_position_ids), + ScalarSpec("num_tokens", torch.int32, PREFILL_T), + TensorSpec("cmp_slot_mapping", [PREFILL_T], torch.int64, init_value=init_cmp_slot_mapping), + TensorSpec("state_slot_mapping", [PREFILL_T], torch.int64, init_value=init_state_slot_mapping), + ] + + +def _run_decode_validation(args): + from golden import ratio_allclose, run_jit + + return run_jit( + fn=decode_compressor_ratio4_test, + specs=build_decode_tensor_specs(args.start_pos), + golden_fn=golden_decode_compressor_ratio4, + runtime_dir=args.runtime_dir, + golden_data=args.golden_data, + compile_cfg=dict(dump_passes=args.dump_passes), + runtime_cfg=dict( + platform=args.platform, + device_id=args.device, + enable_l2_swimlane=args.enable_l2_swimlane, + ), + compile_only=args.compile_only, + rtol=1e-3, + atol=1e-3, + compare_fn={ + "kv": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), + "compress_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), + "cmp_kv_cache": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), + }, + ) + + +def _run_prefill_validation(args): + from golden import ratio_allclose, run_jit + + start_pos = PREFILL_START_POS if args.start_pos is None else args.start_pos + return run_jit( + fn=prefill_compressor_ratio4_test, + specs=build_prefill_tensor_specs(start_pos), + golden_fn=golden_prefill_compressor_ratio4, + runtime_dir=args.runtime_dir, + golden_data=args.golden_data, + compile_cfg=dict(dump_passes=args.dump_passes), + runtime_cfg=dict( + platform=args.platform, + device_id=args.device, + enable_l2_swimlane=args.enable_l2_swimlane, + ), + compile_only=args.compile_only, + compare_fn={ + "compress_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), + "cmp_kv": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), + }, + ) + + +def main(): + import argparse + + parser = argparse.ArgumentParser(description="Standalone DeepSeek V4 compressor ratio4 validation.") + parser.add_argument("--mode", choices=["decode", "prefill", "both"], default="both") + parser.add_argument("-p", "--platform", type=str, default="a2a3", choices=["a2a3", "a2a3sim", "a5", "a5sim"]) + parser.add_argument("-d", "--device", type=int, default=0) + parser.add_argument("--compile-only", action="store_true", default=False) + parser.add_argument( + "--start-pos", + type=int, + default=None, + help="Fixture-only start position override. Decode defaults to its canonical batch set; prefill defaults to 0.", + ) + parser.add_argument("--enable-l2-swimlane", action="store_true", default=False) + parser.add_argument("--runtime-dir", type=str, default=None) + parser.add_argument("--golden-data", type=str, default=None) + parser.add_argument("--dump-passes", action="store_true", default=False) + args = parser.parse_args() + + modes = ("decode", "prefill") if args.mode == "both" else (args.mode,) + for mode in modes: + result = _run_decode_validation(args) if mode == "decode" else _run_prefill_validation(args) + if not result.passed: + if result.error: + print(result.error) + raise SystemExit(1) + + +if __name__ == "__main__": + main() diff --git a/models/deepseek/v4/decode_attention_csa.py b/models/deepseek/v4/decode_attention_csa.py index e456abf5..9b3439bb 100644 --- a/models/deepseek/v4/decode_attention_csa.py +++ b/models/deepseek/v4/decode_attention_csa.py @@ -45,7 +45,7 @@ INT8_SCALE_MAX, INT8_AMAX_EPS, ) -from decode_compressor_ratio4 import compressor_ratio4 +from compressor_ratio4 import decode_compressor_ratio4 from hc_post import hc_post from hc_pre import hc_pre from decode_indexer import indexer @@ -62,7 +62,6 @@ H = M.num_attention_heads HEAD_DIM = M.head_dim ROPE_HEAD_DIM = M.qk_rope_head_dim -HALF_ROPE = ROPE_HEAD_DIM // 2 Q_LORA = M.q_lora_rank WIN = M.sliding_window MAX_SEQ_LEN = M.max_position_embeddings @@ -162,13 +161,8 @@ def attention_csa( rope_cos_t = pl.create_tensor([T, ROPE_HEAD_DIM], dtype=pl.BF16) rope_sin_t = pl.create_tensor([T, ROPE_HEAD_DIM], dtype=pl.BF16) - step_cos = pl.create_tensor([B, HALF_ROPE], dtype=pl.FP32) - step_sin = pl.create_tensor([B, HALF_ROPE], dtype=pl.FP32) with pl.at(level=pl.Level.CORE_GROUP, name_hint="csa_rope_step"): for b in pl.range(B): - first_t = b * S - first_pos_b = pl.read(position_ids, [first_t]) - step_pos_b = pl.cast(first_pos_b, pl.INDEX) for s in pl.range(S): t = b * S + s pos_b = pl.cast(pl.read(position_ids, [t]), pl.INDEX) @@ -176,20 +170,9 @@ def attention_csa( sin_row = pl.cast(freqs_sin[pos_b : pos_b + 1, 0 : ROPE_HEAD_DIM], target_type=pl.FP32) rope_cos_t[t : t + 1, 0 : ROPE_HEAD_DIM] = pl.cast(cos_row, target_type=pl.BF16) rope_sin_t[t : t + 1, 0 : ROPE_HEAD_DIM] = pl.cast(sin_row, target_type=pl.BF16) - step_cos[b : b + 1, 0 : HALF_ROPE] = pl.cast(freqs_cos[step_pos_b : step_pos_b + 1, 0 : HALF_ROPE], target_type=pl.FP32) - step_sin[b : b + 1, 0 : HALF_ROPE] = pl.cast(freqs_sin[step_pos_b : step_pos_b + 1, 0 : HALF_ROPE], target_type=pl.FP32) - - cmp_cos = pl.create_tensor([B, HALF_ROPE], dtype=pl.FP32) - cmp_sin = pl.create_tensor([B, HALF_ROPE], dtype=pl.FP32) - with pl.at(level=pl.Level.CORE_GROUP, name_hint="csa_cmp_rope"): - for b in pl.range(B): - first_t = b * S - first_pos_b = pl.read(position_ids, [first_t]) - cmp_offset_b = COMPRESS_RATIO - (first_pos_b % COMPRESS_RATIO) - cmp_pos_b = pl.cast(first_pos_b + cmp_offset_b - COMPRESS_RATIO, pl.INDEX) - cmp_cos[b : b + 1, 0 : HALF_ROPE] = pl.cast(freqs_cos[cmp_pos_b : cmp_pos_b + 1, 0 : HALF_ROPE], target_type=pl.FP32) - cmp_sin[b : b + 1, 0 : HALF_ROPE] = pl.cast(freqs_sin[cmp_pos_b : cmp_pos_b + 1, 0 : HALF_ROPE], target_type=pl.FP32) + # The ratio-4 compressor now indexes the shared freqs table in-kernel (same rope path + # as decode/prefill), so the host no longer precomputes a per-batch cmp_cos/cmp_sin. x_normed_t = pl.create_tensor([T, D], dtype=pl.BF16) rms_norm(x_mixed, attn_norm_w, x_normed_t) q = pl.create_tensor([T, H, HEAD_DIM], dtype=pl.BF16) @@ -212,32 +195,27 @@ def attention_csa( write_row = pl.cast(write_row_i64, pl.INDEX) kv_cache_flat[write_row : write_row + 1, 0 : HEAD_DIM] = kv[write_t : write_t + 1, 0 : HEAD_DIM] - x_normed = pl.reshape(x_normed_t, [B, S, D]) - cmp_out = pl.create_tensor([B, S, HEAD_DIM], dtype=pl.FP32) - position_ids_bsd = pl.reshape(position_ids, [B, S]) - cmp_slot_mapping_bsd = pl.reshape(cmp_slot_mapping, [B, S]) - idx_slot_mapping_bsd = pl.reshape(idx_slot_mapping, [B, S]) - state_slot_mapping_bsd = pl.reshape(state_slot_mapping, [B, S]) - inner_state_slot_mapping_bsd = pl.reshape(inner_state_slot_mapping, [B, S]) - compressor_ratio4( - x_normed, cmp_out, + cmp_out = pl.create_tensor([T, HEAD_DIM], dtype=pl.FP32) + # decode_compressor_ratio4 is token-major: pass the native [T] tensors directly (no [B, S] adapter). + decode_compressor_ratio4( + x_normed_t, cmp_out, compress_state, compress_state_block_table, cmp_wkv, cmp_wgate, cmp_ape, cmp_norm_w, - cmp_cos, cmp_sin, cmp_kv, - position_ids_bsd, cmp_slot_mapping_bsd, state_slot_mapping_bsd, + freqs_cos, freqs_sin, cmp_kv, + position_ids, cmp_slot_mapping, state_slot_mapping, ) - idx_kv_unused = pl.create_tensor([B, S, IDX_HEAD_DIM], dtype=pl.FP32) - idx_score_unused = pl.create_tensor([B, S, INDEXER_SCORE_LEN], dtype=pl.FP32) - idx_topk_full = pl.create_tensor([B, S, INDEXER_SCORE_LEN], dtype=pl.INT32) + idx_kv_unused = pl.create_tensor([T, IDX_HEAD_DIM], dtype=pl.FP32) + idx_score_unused = pl.create_tensor([T, INDEXER_SCORE_LEN], dtype=pl.FP32) + idx_topk_full = pl.create_tensor([T, INDEXER_SCORE_LEN], dtype=pl.INT32) indexer( - x_normed, qr, qr_scale, idx_wq_b, idx_wq_b_scale, - weights_proj, step_cos, step_sin, hadamard_idx, + x_normed_t, qr, qr_scale, idx_wq_b, idx_wq_b_scale, + weights_proj, freqs_cos, freqs_sin, hadamard_idx, idx_kv_unused, inner_compress_state, inner_compress_state_block_table, inner_wkv, inner_wgate, inner_ape, inner_norm_w, idx_kv_cache, idx_kv_scale, idx_block_table, idx_score_unused, idx_topk_full, - position_ids_bsd, idx_slot_mapping_bsd, inner_state_slot_mapping_bsd, + position_ids, idx_slot_mapping, inner_state_slot_mapping, kv_seq_lens, 0, ) @@ -336,7 +314,7 @@ def golden_attention_csa(tensors): """Torch reference for the ratio-4 compression-step CSA orchestration.""" import torch - from decode_compressor_ratio4 import golden_compressor + from compressor_ratio4 import golden_decode_compressor_ratio4 from hc_pre import golden_hc_pre from decode_indexer import golden_indexer from qkv_proj_rope import golden_qkv_proj_rope @@ -358,23 +336,11 @@ def golden_attention_csa(tensors): }) position_ids = tensors["position_ids"].to(torch.int64) - position_ids_bsd = position_ids.reshape(B, S).to(torch.int32).contiguous() - cmp_slot_mapping_bsd = tensors["cmp_slot_mapping"].reshape(B, S).to(torch.int64).contiguous() - idx_slot_mapping_bsd = tensors["idx_slot_mapping"].reshape(B, S).to(torch.int64).contiguous() - state_slot_mapping_bsd = tensors["state_slot_mapping"].reshape(B, S).to(torch.int64).contiguous() - inner_state_slot_mapping_bsd = tensors["inner_state_slot_mapping"].reshape(B, S).to(torch.int64).contiguous() freqs_cos = tensors["freqs_cos"] freqs_sin = tensors["freqs_sin"] rope_cos_t = freqs_cos[position_ids].contiguous() rope_sin_t = freqs_sin[position_ids].contiguous() - first_pos = position_ids.reshape(B, S)[:, 0] - step_cos = freqs_cos[first_pos, :HALF_ROPE].float().contiguous() - step_sin = freqs_sin[first_pos, :HALF_ROPE].float().contiguous() - cmp_pos = first_pos + (COMPRESS_RATIO - (first_pos % COMPRESS_RATIO)) - COMPRESS_RATIO - cmp_cos = freqs_cos[cmp_pos, :HALF_ROPE].float().contiguous() - cmp_sin = freqs_sin[cmp_pos, :HALF_ROPE].float().contiguous() - q = torch.zeros(T, H, HEAD_DIM, dtype=torch.bfloat16) kv = torch.zeros(T, HEAD_DIM, dtype=torch.bfloat16) qr_i8 = torch.zeros(T, Q_LORA, dtype=torch.int8) @@ -402,9 +368,9 @@ def golden_attention_csa(tensors): cmp_kv = tensors["cmp_kv"] cmp_block_table = tensors["cmp_block_table"] - cmp_out = torch.zeros(B, S, HEAD_DIM, dtype=torch.float32) - golden_compressor({ - "x": x_normed.reshape(B, S, D), + cmp_out = torch.zeros(T, HEAD_DIM, dtype=torch.float32) + golden_decode_compressor_ratio4({ + "x": x_normed, "kv": cmp_out, "compress_state": tensors["compress_state"], "compress_state_block_table": tensors["compress_state_block_table"], @@ -412,26 +378,26 @@ def golden_attention_csa(tensors): "wgate": tensors["cmp_wgate"], "ape": tensors["cmp_ape"], "norm_w": tensors["cmp_norm_w"], - "cos": cmp_cos, - "sin": cmp_sin, + "freqs_cos": freqs_cos, + "freqs_sin": freqs_sin, "cmp_kv_cache": cmp_kv, - "position_ids": position_ids_bsd, - "cmp_slot_mapping": cmp_slot_mapping_bsd, - "state_slot_mapping": state_slot_mapping_bsd, + "position_ids": tensors["position_ids"], + "cmp_slot_mapping": tensors["cmp_slot_mapping"], + "state_slot_mapping": tensors["state_slot_mapping"], }) - idx_kv = torch.zeros(B, S, IDX_HEAD_DIM, dtype=torch.float32) - idx_score = torch.zeros(B, S, INDEXER_SCORE_LEN, dtype=torch.float32) - idx_topk_full = torch.full((B, S, INDEXER_SCORE_LEN), -1, dtype=torch.int32) + idx_kv = torch.zeros(T, IDX_HEAD_DIM, dtype=torch.float32) + idx_score = torch.zeros(T, INDEXER_SCORE_LEN, dtype=torch.float32) + idx_topk_full = torch.full((T, INDEXER_SCORE_LEN), -1, dtype=torch.int32) golden_indexer({ - "x": x_normed.reshape(B, S, D), + "x": x_normed, "qr": qr_i8, "qr_scale": qr_scale, "wq_b": tensors["idx_wq_b"], "wq_b_scale": tensors["idx_wq_b_scale"], "weights_proj": tensors["weights_proj"], - "cos": step_cos, - "sin": step_sin, + "freqs_cos": freqs_cos, + "freqs_sin": freqs_sin, "hadamard": tensors["hadamard_idx"], "inner_kv": idx_kv, "inner_compress_state": tensors["inner_compress_state"], @@ -445,9 +411,9 @@ def golden_attention_csa(tensors): "idx_block_table": tensors["idx_block_table"], "score": idx_score, "topk_idxs": idx_topk_full, - "position_ids": position_ids_bsd, - "idx_slot_mapping": idx_slot_mapping_bsd, - "inner_state_slot_mapping": inner_state_slot_mapping_bsd, + "position_ids": position_ids.to(torch.int32).contiguous(), + "idx_slot_mapping": tensors["idx_slot_mapping"].to(torch.int64).contiguous(), + "inner_state_slot_mapping": tensors["inner_state_slot_mapping"].to(torch.int64).contiguous(), "kv_seq_lens": tensors["kv_seq_lens"], "offset": torch.tensor(0, dtype=torch.int32), }) diff --git a/models/deepseek/v4/decode_attention_hca.py b/models/deepseek/v4/decode_attention_hca.py index 333e8d0b..90243d0d 100644 --- a/models/deepseek/v4/decode_attention_hca.py +++ b/models/deepseek/v4/decode_attention_hca.py @@ -34,7 +34,7 @@ from hc_post import hc_post from qkv_proj_rope import qkv_proj_rope from rmsnorm import rms_norm -from decode_compressor_ratio128 import compressor_ratio128 +from compressor_ratio128 import decode_compressor_ratio128 from decode_sparse_attn_hca import sparse_attn_hca, CMP_TOPK as HCA_SPARSE_CMP_TOPK @@ -140,20 +140,13 @@ def attention_hca( comb_t = pl.create_tensor([T, HC_MULT * HC_MULT], dtype=pl.FP32) hc_pre(x_hc, hc_attn_fn, hc_attn_scale, hc_attn_base, x_mixed, post_t, comb_t) + # The compressor now indexes the shared freqs table in-kernel (decode and prefill use + # the same rope path), so the host no longer precomputes a per-batch cmp_cos/cmp_sin. + # This loop only builds the per-token step rope table for qkv_proj_rope. rope_cos_t = pl.create_tensor([T, ROPE_HEAD_DIM], dtype=pl.BF16) rope_sin_t = pl.create_tensor([T, ROPE_HEAD_DIM], dtype=pl.BF16) - cmp_cos = pl.create_tensor([B, ROPE_HEAD_DIM // 2], dtype=pl.FP32) - cmp_sin = pl.create_tensor([B, ROPE_HEAD_DIM // 2], dtype=pl.FP32) with pl.at(level=pl.Level.CORE_GROUP, name_hint="hca_rope"): for b in pl.range(B): - first_t = b * S - first_pos_b = pl.read(position_ids, [first_t]) - cmp_offset_b = COMPRESS_RATIO - (first_pos_b % COMPRESS_RATIO) - cmp_pos_b = pl.cast(first_pos_b + cmp_offset_b - COMPRESS_RATIO, pl.INDEX) - cmp_cos_row = freqs_cos[cmp_pos_b : cmp_pos_b + 1, 0 : ROPE_HEAD_DIM // 2] - cmp_sin_row = freqs_sin[cmp_pos_b : cmp_pos_b + 1, 0 : ROPE_HEAD_DIM // 2] - cmp_cos[b : b + 1, 0 : ROPE_HEAD_DIM // 2] = pl.cast(cmp_cos_row, target_type=pl.FP32) - cmp_sin[b : b + 1, 0 : ROPE_HEAD_DIM // 2] = pl.cast(cmp_sin_row, target_type=pl.FP32) for s in pl.range(S): t = b * S + s pos_b = pl.cast(pl.read(position_ids, [t]), pl.INDEX) @@ -184,17 +177,14 @@ def attention_hca( write_row = pl.cast(write_row_i64, pl.INDEX) kv_cache_flat[write_row : write_row + 1, 0 : HEAD_DIM] = kv[write_t : write_t + 1, 0 : HEAD_DIM] - x_normed_bsd = pl.reshape(x_normed, [B, S, D]) - cmp_kv_proj = pl.create_tensor([B, S, HEAD_DIM], dtype=pl.FP32) - position_ids_bsd = pl.reshape(position_ids, [B, S]) - cmp_slot_mapping_bsd = pl.reshape(cmp_slot_mapping, [B, S]) - state_slot_mapping_bsd = pl.reshape(state_slot_mapping, [B, S]) - compressor_ratio128( - x_normed_bsd, cmp_kv_proj, + cmp_kv_proj = pl.create_tensor([T, HEAD_DIM], dtype=pl.FP32) + # decode_compressor_ratio128 is token-major: pass the native [T] tensors directly (no [B, S] adapter). + decode_compressor_ratio128( + x_normed, cmp_kv_proj, compress_state, compress_state_block_table, cmp_wkv, cmp_wgate, cmp_ape, cmp_norm_w, - cmp_cos, cmp_sin, cmp_kv, - position_ids_bsd, cmp_slot_mapping_bsd, state_slot_mapping_bsd, + freqs_cos, freqs_sin, cmp_kv, + position_ids, cmp_slot_mapping, state_slot_mapping, ) # Sparse-index build fanned out over an SPMD (8 tokens/block) instead of one @@ -297,7 +287,7 @@ def golden_attention_hca(tensors): from hc_pre import golden_hc_pre from qkv_proj_rope import golden_qkv_proj_rope from rmsnorm import golden_rms_norm - from decode_compressor_ratio128 import golden_compressor + from compressor_ratio128 import golden_decode_compressor_ratio128 from decode_sparse_attn_hca import golden_sparse_attn from hc_post import golden_hc_post @@ -360,22 +350,9 @@ def golden_attention_hca(tensors): cmp_block_table = tensors["cmp_block_table"] attn_out = torch.zeros(T, D, dtype=torch.bfloat16) - half_rd = rd // 2 - cmp_cos = torch.empty(B, half_rd, dtype=torch.float32) - cmp_sin = torch.empty(B, half_rd, dtype=torch.float32) - for b in range(B): - first_pos_b = int(position_ids[b * S].item()) - cmp_offset_b = ratio - (first_pos_b % ratio) - cmp_pos_b = first_pos_b + cmp_offset_b - ratio - cmp_cos[b] = freqs_cos[cmp_pos_b, :half_rd].float() - cmp_sin[b] = freqs_sin[cmp_pos_b, :half_rd].float() - - cmp_kv_proj = torch.zeros(B, S, HEAD_DIM, dtype=torch.float32) - position_ids_bsd = position_ids.reshape(B, S).to(torch.int32).contiguous() - cmp_slot_mapping_bsd = tensors["cmp_slot_mapping"].reshape(B, S).to(torch.int64).contiguous() - state_slot_mapping_bsd = tensors["state_slot_mapping"].reshape(B, S).to(torch.int64).contiguous() - golden_compressor({ - "x": x_normed.reshape(B, S, D), + cmp_kv_proj = torch.zeros(T, HEAD_DIM, dtype=torch.float32) + golden_decode_compressor_ratio128({ + "x": x_normed, "kv": cmp_kv_proj, "compress_state": tensors["compress_state"], "compress_state_block_table": tensors["compress_state_block_table"], @@ -383,12 +360,12 @@ def golden_attention_hca(tensors): "wgate": tensors["cmp_wgate"], "ape": tensors["cmp_ape"], "norm_w": tensors["cmp_norm_w"], - "cos": cmp_cos, - "sin": cmp_sin, + "freqs_cos": freqs_cos, + "freqs_sin": freqs_sin, "cmp_kv_cache": cmp_kv, - "position_ids": position_ids_bsd, - "cmp_slot_mapping": cmp_slot_mapping_bsd, - "state_slot_mapping": state_slot_mapping_bsd, + "position_ids": tensors["position_ids"], + "cmp_slot_mapping": tensors["cmp_slot_mapping"], + "state_slot_mapping": tensors["state_slot_mapping"], }) ori_slot_mapping = tensors["ori_slot_mapping"].to(torch.int64) diff --git a/models/deepseek/v4/decode_compressor_ratio128.py b/models/deepseek/v4/decode_compressor_ratio128.py deleted file mode 100644 index adb35535..00000000 --- a/models/deepseek/v4/decode_compressor_ratio128.py +++ /dev/null @@ -1,588 +0,0 @@ -# Copyright (c) PyPTO Contributors. -# This program is free software, you can redistribute it and/or modify it under the terms and conditions of -# CANN Open Software License Agreement Version 2.0 (the "License"). -# Please refer to the License for details. You may not use this file except in compliance with the License. -# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, -# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. -# See LICENSE in the root of the software repository for the full text of the License. -# ----------------------------------------------------------------------------------------------------------- -"""DeepSeek-V4 KV Compressor (decode incremental, ratio=128 non-overlap). - -Uses non-overlapping state layout with 128 slots. -Softmax+pool over all slots. No state shift needed.""" - -import pypto.language as pl - -from config import ( - FLASH as M, - BLOCK_SIZE, - C128_COMPRESSOR_BLOCK_SIZE, - DECODE_BATCH, - DECODE_SEQ, - DECODE_CMP_BLOCK_NUM, - FP32_NEG_INF, - KV_CMP_MAX_BLOCKS, -) - -# Dynamic shape variables. -B_DYN = pl.dynamic("B_DYN") -S_DYN = pl.dynamic("S_DYN") -COMPRESS_STATE_MAX_BLOCKS_DYN = pl.dynamic("COMPRESS_STATE_MAX_BLOCKS_DYN") -COMPRESS_STATE_BLOCK_NUM_DYN = pl.dynamic("COMPRESS_STATE_BLOCK_NUM_DYN") -CMP_BLOCK_NUM_DYN = pl.dynamic("CMP_BLOCK_NUM_DYN") - -# model config -B = DECODE_BATCH -S = DECODE_SEQ -EPS = M.rms_norm_eps -D = M.hidden_size -HEAD_DIM = M.head_dim -HEAD_DIM_INV = 1.0 / HEAD_DIM -ROPE_HEAD_DIM = M.qk_rope_head_dim -NOPE_HEAD_DIM = M.nope_head_dim -MAX_SEQ_LEN = M.max_position_embeddings - -# kernel-local (ratio-128 non-overlap compressor) -COMPRESS_RATIO = 128 -IDX_KV_LEN = MAX_SEQ_LEN // COMPRESS_RATIO -COFF = 1 -OUT_DIM = COFF * HEAD_DIM -STATE_LEN = COFF * COMPRESS_RATIO -# Paged read contract: -# - compress_state_block_table is still used to read the historical ratio window -# by absolute token position. -# - Persistent writes are explicit token-major contracts: -# state_slot_mapping[b, s] -> flattened compressor-state row, -1 means no-write -# cmp_slot_mapping[b, s] -> flattened compressed-KV row, -1 means no-write -# - APE remains ratio-local: -# ape_row = position_ids[b, s] % COMPRESS_RATIO -COMPRESS_STATE_BLOCK_SIZE = C128_COMPRESSOR_BLOCK_SIZE -# Logical state block tables cover MAX_SEQ_LEN while the physical state pool -# remains bounded to the per-request rolling state capacity. -COMPRESS_STATE_PHYSICAL_BLOCKS = 64 -COMPRESS_STATE_MAX_BLOCKS = (MAX_SEQ_LEN + COMPRESS_STATE_BLOCK_SIZE - 1) // COMPRESS_STATE_BLOCK_SIZE -COMPRESS_STATE_BLOCK_NUM = B * COMPRESS_STATE_PHYSICAL_BLOCKS -COMPRESS_STATE_DIM = 2 * OUT_DIM -CMP_MAX_BLOCKS = KV_CMP_MAX_BLOCKS -CMP_BLOCK_NUM = DECODE_CMP_BLOCK_NUM -if IDX_KV_LEN > CMP_MAX_BLOCKS * BLOCK_SIZE: - raise ValueError("ratio128 compressed KV cache capacity is smaller than max compressed sequence length") - -# tiling -ROPE_TILE = 32 -K_TILE = 512 -OUT_TILE = 64 -HEAD_TILE = 64 -B_TILE = 8 -MM_B_TILE = 16 -BS_PAD = ((B * S + MM_B_TILE - 1) // MM_B_TILE) * MM_B_TILE -RMS_TILE = 4 -RMS_PAD_TILE = 16 -RMS_PAD_TAIL = RMS_PAD_TILE - RMS_TILE -RMS_PAD_ROWS = (B // RMS_TILE) * RMS_PAD_TILE -# softmax_pool reduces over the state axis with column reductions (no transpose), so it can -# afford a wider head tile than HEAD_TILE: each wider tile loads each state block fewer times -# (HEAD_DIM/POOL_HEAD_TILE tiles/batch instead of HEAD_DIM/HEAD_TILE), cutting load redundancy. -POOL_HEAD_TILE = 128 - - -@pl.jit.inline -def compressor_ratio128( - x: pl.Tensor[[B_DYN, S_DYN, D], pl.BF16], - kv: pl.Tensor[[B_DYN, S_DYN, HEAD_DIM], pl.FP32], - compress_state: pl.Tensor[[COMPRESS_STATE_BLOCK_NUM_DYN, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32], - compress_state_block_table: pl.Tensor[[B_DYN, COMPRESS_STATE_MAX_BLOCKS_DYN], pl.INT32], - wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], - wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], - ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], - norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - cos: pl.Tensor[[B_DYN, ROPE_HEAD_DIM // 2], pl.FP32], - sin: pl.Tensor[[B_DYN, ROPE_HEAD_DIM // 2], pl.FP32], - cmp_kv_cache: pl.Tensor[[CMP_BLOCK_NUM_DYN, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16], - position_ids: pl.Tensor[[B_DYN, S_DYN], pl.INT32], - cmp_slot_mapping: pl.Tensor[[B_DYN, S_DYN], pl.INT64], - state_slot_mapping: pl.Tensor[[B_DYN, S_DYN], pl.INT64], -): - b_dim = pl.tensor.dim(x, 0) - s_dim = pl.tensor.dim(x, 1) - bs = b_dim * s_dim - compress_state_block_num = pl.tensor.dim(compress_state, 0) - cmp_block_num = pl.tensor.dim(cmp_kv_cache, 0) - - x_flat = pl.reshape(x, [bs, D]) - t_matmul = pl.max(bs, MM_B_TILE) - kv_proj_pad = pl.create_tensor([BS_PAD, OUT_DIM], dtype=pl.FP32) - score_proj_pad = pl.create_tensor([BS_PAD, OUT_DIM], dtype=pl.FP32) - - for idx in pl.spmd(t_matmul * OUT_DIM // (MM_B_TILE * OUT_TILE), name_hint="kv_score_proj"): - global_row0 = (idx // (OUT_DIM // OUT_TILE)) * MM_B_TILE - o0 = (idx % (OUT_DIM // OUT_TILE)) * OUT_TILE - kv_acc = pl.create_tensor([MM_B_TILE, OUT_TILE], dtype=pl.FP32) - score_acc = pl.create_tensor([MM_B_TILE, OUT_TILE], dtype=pl.FP32) - for kb in pl.pipeline(0, D // K_TILE, stage=2): - k0 = kb * K_TILE - x_rows = pl.min(MM_B_TILE, bs - global_row0) - x_tile = pl.slice(x_flat, [MM_B_TILE, K_TILE], [global_row0, k0], valid_shape=[x_rows, K_TILE]) - # Weights stored transposed [OUT_DIM, D] and consumed via b_trans=True so the - # GM->L1 load is a DN2ZN (each [OUT_TILE, K_TILE] row is K-contiguous = long - # bursts) instead of ND2NZ on [K_TILE, OUT_TILE] (K strided = many short - # bursts). Cuts the transaction-bound MTE2 cost. Matches ratio4/CSA layout. - wkv_tile = wkv[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] - wgate_tile = wgate[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] - if k0 == 0: - kv_acc = pl.matmul(x_tile, wkv_tile, out_dtype=pl.FP32, b_trans=True) - score_acc = pl.matmul(x_tile, wgate_tile, out_dtype=pl.FP32, b_trans=True) - else: - kv_acc = pl.matmul_acc(kv_acc, x_tile, wkv_tile, b_trans=True) - score_acc = pl.matmul_acc(score_acc, x_tile, wgate_tile, b_trans=True) - - kv_proj_pad[global_row0 : global_row0 + MM_B_TILE, o0 : o0 + OUT_TILE] = kv_acc - score_proj_pad[global_row0 : global_row0 + MM_B_TILE, o0 : o0 + OUT_TILE] = score_acc - - compress_state_flat = pl.reshape(compress_state, [compress_state_block_num, COMPRESS_STATE_BLOCK_SIZE * COMPRESS_STATE_DIM]) - - # state scatter reads the padded proj tensors directly by flat token row (no unpad pass). - with pl.at(level=pl.Level.CORE_GROUP, name_hint="state_scatter_pre") as scatter_tid: - for global_c_idx in pl.range(b_dim): - for s in pl.pipeline(s_dim, stage=2): - proj_row = global_c_idx * s_dim + s - token_pos = pl.read(position_ids, [global_c_idx, s]) - token_ape_row = pl.cast(token_pos % COMPRESS_RATIO, target_type=pl.INDEX) - state_row_i64 = pl.read(state_slot_mapping, [global_c_idx, s]) - if state_row_i64 >= 0: - state_row = pl.cast(state_row_i64, target_type=pl.INDEX) - state_blk_id = state_row // COMPRESS_STATE_BLOCK_SIZE - state_intra = state_row % COMPRESS_STATE_BLOCK_SIZE - slot_col0_s = state_intra * COMPRESS_STATE_DIM - ape_row = ape[token_ape_row : token_ape_row + 1, 0 : OUT_DIM] - kv_row = kv_proj_pad[proj_row : proj_row + 1, 0 : OUT_DIM] - score_row = score_proj_pad[proj_row : proj_row + 1, 0 : OUT_DIM] - score_row = pl.add(score_row, ape_row) - compress_state_flat[state_blk_id : state_blk_id + 1, slot_col0_s : slot_col0_s + OUT_DIM] = kv_row - compress_state_flat[state_blk_id : state_blk_id + 1, slot_col0_s + OUT_DIM : slot_col0_s + 2 * OUT_DIM] = score_row - - pooled_kv = pl.create_tensor([RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) - # One GM row per compressed-state slot (block * BLOCK_SIZE + intra). This lets - # softmax_pool fetch a whole physical block's BLOCK_SIZE state rows in a single - # strided MTE2 instead of BLOCK_SIZE single-row gathers. - compress_state_rows = pl.reshape( - compress_state, [compress_state_block_num * COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM] - ) - NUM_STATE_BLOCKS = STATE_LEN // COMPRESS_STATE_BLOCK_SIZE - with pl.spmd(b_dim * HEAD_DIM // POOL_HEAD_TILE, name_hint="softmax_pool", deps=[scatter_tid]) as pool_tid: - idx = pl.tile.get_block_idx() - global_c_idx = idx // (HEAD_DIM // POOL_HEAD_TILE) - pad_idx = (global_c_idx // RMS_TILE) * RMS_PAD_TILE + (global_c_idx % RMS_TILE) - h0 = (idx % (HEAD_DIM // POOL_HEAD_TILE)) * POOL_HEAD_TILE - first_pos_gate = pl.read(position_ids, [global_c_idx, 0]) - pos_gate = first_pos_gate % COMPRESS_RATIO - if pos_gate + S >= COMPRESS_RATIO: - softmax_score_state = pl.create_tensor([STATE_LEN, POOL_HEAD_TILE], dtype=pl.FP32) - softmax_kv_state = pl.create_tensor([STATE_LEN, POOL_HEAD_TILE], dtype=pl.FP32) - # The STATE_LEN contiguous state positions begin at a multiple of COMPRESS_RATIO - # (hence a multiple of COMPRESS_STATE_BLOCK_SIZE), so the window is exactly - # NUM_STATE_BLOCKS full physical blocks with no partial head/tail. Load each - # block's BLOCK_SIZE rows in ONE [BLOCK_SIZE, HEAD_TILE] strided MTE2 instead of - # BLOCK_SIZE separate [1, HEAD_TILE] row loads: 8x fewer transactions and no - # per-row UB staging. Bit-identical to the per-row gather. - compress_pos = first_pos_gate + (COMPRESS_RATIO - 1 - pos_gate) - state_pos0 = compress_pos - (COMPRESS_RATIO - 1) - base_logical_blk = state_pos0 // COMPRESS_STATE_BLOCK_SIZE - for blk_i in pl.pipeline(NUM_STATE_BLOCKS, stage=2): - s0 = blk_i * COMPRESS_STATE_BLOCK_SIZE - slot_score = pl.full([COMPRESS_STATE_BLOCK_SIZE, POOL_HEAD_TILE], dtype=pl.FP32, value=FP32_NEG_INF) - slot_kv = pl.full([COMPRESS_STATE_BLOCK_SIZE, POOL_HEAD_TILE], dtype=pl.FP32, value=0.0) - state_blk_raw = pl.read(compress_state_block_table, [global_c_idx, base_logical_blk + blk_i]) - if state_blk_raw >= 0: - state_blk_id = pl.cast(state_blk_raw, target_type=pl.INDEX) - row0 = state_blk_id * COMPRESS_STATE_BLOCK_SIZE - slot_score = compress_state_rows[row0 : row0 + COMPRESS_STATE_BLOCK_SIZE, OUT_DIM + h0 : OUT_DIM + h0 + POOL_HEAD_TILE] - slot_kv = compress_state_rows[row0 : row0 + COMPRESS_STATE_BLOCK_SIZE, h0 : h0 + POOL_HEAD_TILE] - softmax_score_state[s0 : s0 + COMPRESS_STATE_BLOCK_SIZE, :] = slot_score - softmax_kv_state[s0 : s0 + COMPRESS_STATE_BLOCK_SIZE, :] = slot_kv - - # Softmax over the state axis (rows) directly via column reductions, avoiding the two - # [STATE_LEN, *] transposes (VNCHWCONV) the row-reduce form needed. col_max/col_sum - # reduce over rows -> [1, POOL_HEAD_TILE]; col_expand_expdif fuses exp(x - col_max) - # and col_expand_mul broadcasts recip(sum) back over rows (col_expand_sub/div have no - # codegen; mul/expdif do). Same reduction over the same STATE_LEN values per head col. - score_max = pl.col_max(softmax_score_state) - score_exp = pl.col_expand_expdif(softmax_score_state, score_max) - score_sum = pl.col_sum(score_exp) - score_prob = pl.col_expand_mul(score_exp, pl.recip(score_sum)) - pooled_chunk = pl.col_sum(pl.mul(softmax_kv_state, score_prob)) - pooled_kv[pad_idx : pad_idx + 1, h0 : h0 + POOL_HEAD_TILE] = pooled_chunk - - norm_w_2d = pl.reshape(norm_w, [1, HEAD_DIM]) - normed_kv = pl.create_tensor([RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) - - with pl.spmd(b_dim // RMS_TILE, name_hint="rmsnorm_rope", deps=[pool_tid]) as rms_tid: - batch_base_idx = pl.tile.get_block_idx() - batch_base = batch_base_idx * RMS_TILE - pad_base = batch_base_idx * RMS_PAD_TILE - cos_b = pl.full([RMS_PAD_TILE, ROPE_HEAD_DIM // 2], dtype=pl.FP32, value=0.0) - sin_b = pl.full([RMS_PAD_TILE, ROPE_HEAD_DIM // 2], dtype=pl.FP32, value=0.0) - cos_b[0:RMS_TILE, 0 : ROPE_HEAD_DIM // 2] = cos[batch_base : batch_base + RMS_TILE, 0 : ROPE_HEAD_DIM // 2] - sin_b[0:RMS_TILE, 0 : ROPE_HEAD_DIM // 2] = sin[batch_base : batch_base + RMS_TILE, 0 : ROPE_HEAD_DIM // 2] - partial_sq = pl.full([1, RMS_PAD_TILE], dtype=pl.FP32, value=0.0) - for rms_kb in pl.pipeline(HEAD_DIM // HEAD_TILE, stage=2): - rms_h0 = rms_kb * HEAD_TILE - kv_rms_chunk = pooled_kv[pad_base : pad_base + RMS_PAD_TILE, rms_h0 : rms_h0 + HEAD_TILE] - kv_rms_sq = pl.mul(kv_rms_chunk, kv_rms_chunk) - kv_rms_rowsum = pl.reshape(pl.row_sum(kv_rms_sq), [1, RMS_PAD_TILE]) - partial_sq = pl.add(partial_sq, kv_rms_rowsum) - - variance = pl.reshape(pl.add(pl.mul(partial_sq, HEAD_DIM_INV), EPS), [RMS_PAD_TILE, 1]) - inv_rms = pl.recip(pl.sqrt(variance)) - for rms_kb in pl.pipeline(NOPE_HEAD_DIM // HEAD_TILE, stage=2): - norm_h0 = rms_kb * HEAD_TILE - kv_norm_chunk = pooled_kv[pad_base : pad_base + RMS_PAD_TILE, norm_h0 : norm_h0 + HEAD_TILE] - gamma = pl.cast(norm_w_2d[:, norm_h0 : norm_h0 + HEAD_TILE], pl.FP32) - normed_chunk = pl.col_expand_mul(pl.row_expand_mul(kv_norm_chunk, inv_rms), gamma) - normed_kv[pad_base : pad_base + RMS_PAD_TILE, norm_h0 : norm_h0 + HEAD_TILE] = normed_chunk - - kv_rope_norm = pooled_kv[pad_base : pad_base + RMS_PAD_TILE, NOPE_HEAD_DIM : HEAD_DIM] - gamma_rope = pl.cast(norm_w_2d[:, NOPE_HEAD_DIM : HEAD_DIM], pl.FP32) - # A3 interleaved swap-gather (same form as kv_rope_fused in qkv_proj_rope), - # replacing the de-interleave gather + rotate + re-interleave scatter. gamma+inv_rms - # are folded into rope_normed BEFORE the swap, so the swapped lane n[j^1] correctly - # carries gamma[j^1]; inv_rms is per-row so it commutes. swap_idx (j^1), sign - # ([-1,+1,...]) and dup_idx (j>>1) are built IN-KERNEL from pl.arange; cos_il/sin_il - # are dup-gathered from the per-batch cos/sin rows. normed_kv is FP32 -> write directly. - # out[j] = n[j]*cos_il[j] + n[j^1]*sign[j]*sin_il[j] - rope_normed = pl.col_expand_mul(pl.row_expand_mul(kv_rope_norm, inv_rms), gamma_rope) - rope_ones = pl.full([RMS_PAD_TILE, ROPE_HEAD_DIM], dtype=pl.FP32, value=1.0) - rope_col = pl.col_expand_mul(rope_ones, pl.cast(pl.arange(0, [1, ROPE_HEAD_DIM], dtype=pl.INT32), target_type=pl.FP32)) - rope_dup_f = pl.cast(pl.cast(pl.mul(rope_col, 0.5), target_type=pl.INT32, mode="trunc"), target_type=pl.FP32) - rope_dup_idx = pl.cast(rope_dup_f, target_type=pl.INT32) # j>>1 - rope_lane = pl.sub(rope_col, pl.mul(rope_dup_f, 2.0)) # j%2 - rope_swap_idx = pl.cast(pl.sub(pl.add(rope_col, 1.0), pl.mul(rope_lane, 2.0)), target_type=pl.INT32) # j^1 - rope_sign = pl.sub(pl.mul(rope_lane, 2.0), 1.0) # [-1,+1,...] - cos_il = pl.gather(cos_b, dim=-1, index=rope_dup_idx) - sin_il = pl.gather(sin_b, dim=-1, index=rope_dup_idx) - swapped = pl.gather(rope_normed, dim=-1, index=rope_swap_idx) - rope_rot = pl.add(pl.mul(rope_normed, cos_il), pl.mul(pl.mul(swapped, rope_sign), sin_il)) - normed_kv[pad_base : pad_base + RMS_PAD_TILE, NOPE_HEAD_DIM : HEAD_DIM] = rope_rot - - kv_flat = pl.reshape(kv, [bs, HEAD_DIM]) - cmp_flat_rows = cmp_block_num * BLOCK_SIZE - cmp_kv_cache_flat = pl.reshape(cmp_kv_cache, [cmp_flat_rows, HEAD_DIM]) - - with pl.spmd(b_dim // RMS_TILE, name_hint="kv_finalize", deps=[rms_tid]) as _write_tid: - batch_base_idx = pl.tile.get_block_idx() - batch_base = batch_base_idx * RMS_TILE - pad_base = batch_base_idx * RMS_PAD_TILE - for inner in pl.range(RMS_TILE): - global_c_idx = batch_base + inner - first_pos_b = pl.read(position_ids, [global_c_idx, 0]) - pos_b = first_pos_b % COMPRESS_RATIO - if pos_b + s_dim >= COMPRESS_RATIO: - boundary_s = COMPRESS_RATIO - 1 - pos_b - kv_row = normed_kv[pad_base + inner : pad_base + inner + 1, 0 : HEAD_DIM] - cmp_row_i64 = pl.read(cmp_slot_mapping, [global_c_idx, boundary_s]) - if cmp_row_i64 >= 0: - cmp_row = pl.cast(cmp_row_i64, target_type=pl.INDEX) - kv_flat[global_c_idx * s_dim : global_c_idx * s_dim + 1, :] = kv_row - cmp_kv_cache_flat[cmp_row : cmp_row + 1, :] = pl.cast(kv_row, target_type=pl.BF16, mode="rint") - - kv = pl.reshape(kv_flat, [b_dim, s_dim, HEAD_DIM]) - return kv - - -@pl.jit -def compressor_test( - x: pl.Tensor[[B_DYN, S_DYN, D], pl.BF16], - kv: pl.Out[pl.Tensor[[B_DYN, S_DYN, HEAD_DIM], pl.FP32]], - compress_state: pl.InOut[pl.Tensor[[COMPRESS_STATE_BLOCK_NUM_DYN, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32]], - compress_state_block_table: pl.Tensor[[B_DYN, COMPRESS_STATE_MAX_BLOCKS_DYN], pl.INT32], - wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], - wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], - ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], - norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - cos: pl.Tensor[[B_DYN, ROPE_HEAD_DIM // 2], pl.FP32], - sin: pl.Tensor[[B_DYN, ROPE_HEAD_DIM // 2], pl.FP32], - cmp_kv_cache: pl.InOut[pl.Tensor[[CMP_BLOCK_NUM_DYN, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], - position_ids: pl.Tensor[[B_DYN, S_DYN], pl.INT32], - cmp_slot_mapping: pl.Tensor[[B_DYN, S_DYN], pl.INT64], - state_slot_mapping: pl.Tensor[[B_DYN, S_DYN], pl.INT64], -): - x.bind_dynamic(0, B_DYN) - x.bind_dynamic(1, S_DYN) - kv.bind_dynamic(0, B_DYN) - kv.bind_dynamic(1, S_DYN) - compress_state.bind_dynamic(0, COMPRESS_STATE_BLOCK_NUM_DYN) - compress_state_block_table.bind_dynamic(0, B_DYN) - compress_state_block_table.bind_dynamic(1, COMPRESS_STATE_MAX_BLOCKS_DYN) - cos.bind_dynamic(0, B_DYN) - sin.bind_dynamic(0, B_DYN) - cmp_kv_cache.bind_dynamic(0, CMP_BLOCK_NUM_DYN) - position_ids.bind_dynamic(0, B_DYN) - position_ids.bind_dynamic(1, S_DYN) - cmp_slot_mapping.bind_dynamic(0, B_DYN) - cmp_slot_mapping.bind_dynamic(1, S_DYN) - state_slot_mapping.bind_dynamic(0, B_DYN) - state_slot_mapping.bind_dynamic(1, S_DYN) - - compressor_ratio128( - x, kv, compress_state, compress_state_block_table, wkv, wgate, ape, norm_w, cos, sin, - cmp_kv_cache, position_ids, cmp_slot_mapping, state_slot_mapping, - ) - return kv, compress_state, cmp_kv_cache - - -def golden_compressor(tensors): - """Torch reference for Compressor.forward (decode branch, ratio=128 non-overlap). - - Operates on paged caches: compress_state (kv + score channels merged) and cmp_kv_cache, - each addressed via the corresponding block_table. - """ - import torch - - x = tensors["x"].float() - compress_state_block_table = tensors["compress_state_block_table"] - position_ids = tensors["position_ids"].to(torch.int64) - cmp_slot_mapping = tensors["cmp_slot_mapping"].to(torch.int64) - state_slot_mapping = tensors["state_slot_mapping"].to(torch.int64) - # Historical state reads still use absolute-position block-table addressing. - # Persistent writes use token-major slot mappings. APE remains modulo ratio. - compress_state = tensors["compress_state"] - - def read_state_row(b, pos): - logical_blk = pos // COMPRESS_STATE_BLOCK_SIZE - intra = pos % COMPRESS_STATE_BLOCK_SIZE - sblk = int(compress_state_block_table[b, logical_blk].item()) - if sblk < 0: - return ( - torch.zeros(OUT_DIM, dtype=torch.float32, device=compress_state.device), - torch.full((OUT_DIM,), float("-inf"), dtype=torch.float32, device=compress_state.device), - ) - return ( - compress_state[sblk, intra, :OUT_DIM], - compress_state[sblk, intra, OUT_DIM:2 * OUT_DIM], - ) - - def write_state_row(slot, kv_row, score_row): - if slot < 0: - return - sblk = slot // COMPRESS_STATE_BLOCK_SIZE - intra = slot % COMPRESS_STATE_BLOCK_SIZE - compress_state[sblk, intra, :OUT_DIM] = kv_row - compress_state[sblk, intra, OUT_DIM:2 * OUT_DIM] = score_row - - wkv = tensors["wkv"].float() - wgate = tensors["wgate"].float() - ape = tensors["ape"] - norm_w = tensors["norm_w"] - cos = tensors["cos"] - sin = tensors["sin"] - cmp_kv_cache = tensors["cmp_kv_cache"] - bsz, _, _ = x.shape - ratio, rd = COMPRESS_RATIO, ROPE_HEAD_DIM - - kv = x @ wkv.t() # [B, S, OUT_DIM] (wkv stored [OUT_DIM, D] for b_trans) - score = x @ wgate.t() # [B, S, OUT_DIM] - pooled = torch.zeros(bsz, 1, HEAD_DIM, dtype=torch.float32, device=x.device) - should_compress_rows = torch.zeros(bsz, dtype=torch.bool, device=x.device) - - for b in range(bsz): - boundary_s = None - for s in range(S): - pos = int(position_ids[b, s].item()) - token_ape_row = pos % ratio - score[b, s, :] = score[b, s, :] + ape[token_ape_row] - write_state_row(int(state_slot_mapping[b, s].item()), kv[b, s, :], score[b, s, :]) - if (pos + 1) % ratio == 0: - boundary_s = s - - if boundary_s is not None: - should_compress_rows[b] = True - compress_pos = int(position_ids[b, boundary_s].item()) - kv_rows = [] - score_rows = [] - for pos in range(compress_pos - ratio + 1, compress_pos + 1): - kv_row, score_row = read_state_row(b, pos) - kv_rows.append(kv_row) - score_rows.append(score_row) - kv_state = torch.stack(kv_rows, dim=0).unsqueeze(0) - score_state = torch.stack(score_rows, dim=0).unsqueeze(0) - pooled[b : b + 1] = (kv_state * score_state.softmax(dim=1)).sum(dim=1, keepdim=True) - tensors["compress_state"][:] = compress_state - - if not bool(should_compress_rows.any()): - return - - def rmsnorm(x, w): - x = x.float() - var = x.square().mean(-1, keepdim=True) - x = x * torch.rsqrt(var + EPS) - return w * x - - for b in range(bsz): - if not bool(should_compress_rows[b]): - continue - kv_b = rmsnorm(pooled[b : b + 1], norm_w) - - x_pair = kv_b[..., -rd:].unflatten(-1, (-1, 2)) - x0, x1 = x_pair[..., 0], x_pair[..., 1] - cos_v, sin_v = cos[b].view(-1), sin[b].view(-1) - y0 = x0 * cos_v - x1 * sin_v - y1 = x0 * sin_v + x1 * cos_v - - kv_b = torch.cat([kv_b[..., :-rd], torch.stack([y0, y1], dim=-1).flatten(-2)], dim=-1) - - boundary_positions = torch.nonzero((position_ids[b, :S] + 1) % ratio == 0, as_tuple=False).flatten() - if int(boundary_positions.numel()) == 0: - continue - boundary_s = int(boundary_positions[0].item()) - cmp_row = int(cmp_slot_mapping[b, boundary_s].item()) - if cmp_row >= 0: - # Kernel writes committed pooled result only to kv[:, 0, :]; leave - # speculative-boundary rows and kv[:, 1:, :] zero-initialized. - tensors["kv"][b : b + 1, 0:1, :] = kv_b - cblk = cmp_row // BLOCK_SIZE - intra_offset = cmp_row % BLOCK_SIZE - cmp_kv_cache[cblk, intra_offset, 0] = kv_b[0, 0] - - tensors["cmp_kv_cache"][:] = cmp_kv_cache - - -def build_tensor_specs(start_pos=None): - import torch # type: ignore[import] - from decode_metadata import ( - block_table, - compressed_slot_mapping, - hca_decode_start_set, - position_ids_from_starts, - resolve_start_positions, - state_slot_mapping, - ) - from golden import TensorSpec - from rope_tables import build_deepseek_v4_rope_tables, materialize_half_rope_tables - - shared_freqs_cos, shared_freqs_sin = build_deepseek_v4_rope_tables(M, COMPRESS_RATIO, dtype=torch.bfloat16) - - def init_x(): - return torch.rand(B, S, D) - def init_compress_state(): - return torch.zeros(COMPRESS_STATE_BLOCK_NUM, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM) - # Calibrated to the real DeepSeek-V4-Flash 150 - # (ratio-128) main compressor (mean l7/l9 of - # extract_weights_flash): zero-mean Gaussian BF16 weights at the measured std; the RMSNorm - # gamma centers near the measured mean (not ones / not uniform). - def init_wkv(): - return torch.randn(OUT_DIM, D) * 0.0246 - def init_wgate(): - return torch.randn(OUT_DIM, D) * 0.0316 - def init_ape(): - return torch.randn(COMPRESS_RATIO, OUT_DIM) * 0.0340 - def init_norm_w(): - return 0.1001 + 0.0549 * torch.randn(HEAD_DIM) - def init_rope_positions(): - first_pos = init_position_ids().to(torch.int64)[:, 0] - cmp_offset = COMPRESS_RATIO - (first_pos % COMPRESS_RATIO) - return (first_pos + cmp_offset - COMPRESS_RATIO).to(torch.int64) - def init_cos(): - return materialize_half_rope_tables(shared_freqs_cos, shared_freqs_sin, init_rope_positions())[0] - def init_sin(): - return materialize_half_rope_tables(shared_freqs_cos, shared_freqs_sin, init_rope_positions())[1] - def init_cmp_kv_cache(): - return torch.zeros(CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM) - def init_compress_state_block_table(): - return block_table( - batch=B, - table_blocks=COMPRESS_STATE_MAX_BLOCKS, - physical_blocks=COMPRESS_STATE_PHYSICAL_BLOCKS, - permuted=True, - ) - def init_cmp_block_table(): - return block_table( - batch=B, - table_blocks=CMP_MAX_BLOCKS, - physical_blocks=CMP_MAX_BLOCKS, - permuted=True, - ) - def init_default_start_pos(): - # Canonical HCA start-position set (ratio-128 compressor branches + 8k long-context). - return hca_decode_start_set( - batch=B, compress_ratio=COMPRESS_RATIO, state_block_size=COMPRESS_STATE_BLOCK_SIZE) - def init_start_pos(): - return resolve_start_positions( - start_pos, - batch=B, - seq=S, - max_seq_len=MAX_SEQ_LEN, - default_fn=init_default_start_pos, - ) - def init_position_ids(): - return position_ids_from_starts(init_start_pos(), seq=S) - def init_state_slot_mapping(): - return state_slot_mapping( - init_position_ids(), - init_compress_state_block_table(), - state_block_size=COMPRESS_STATE_BLOCK_SIZE, - ) - def init_cmp_slot_mapping(): - positions = init_position_ids() - return compressed_slot_mapping( - positions, - init_cmp_block_table(), - compress_ratio=COMPRESS_RATIO, - block_size=BLOCK_SIZE, - ) - return [ - TensorSpec("x", [B, S, D], torch.bfloat16, init_value=init_x), - TensorSpec("kv", [B, S, HEAD_DIM], torch.float32, is_output=True), - TensorSpec("compress_state", [COMPRESS_STATE_BLOCK_NUM, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], torch.float32, init_value=init_compress_state, is_output=True), - TensorSpec("compress_state_block_table", [B, COMPRESS_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), - TensorSpec("wkv", [OUT_DIM, D], torch.bfloat16, init_value=init_wkv), - TensorSpec("wgate", [OUT_DIM, D], torch.bfloat16, init_value=init_wgate), - TensorSpec("ape", [COMPRESS_RATIO, OUT_DIM], torch.float32, init_value=init_ape), - TensorSpec("norm_w", [HEAD_DIM], torch.bfloat16, init_value=init_norm_w), - TensorSpec("cos", [B, ROPE_HEAD_DIM // 2], torch.float32, init_value=init_cos), - TensorSpec("sin", [B, ROPE_HEAD_DIM // 2], torch.float32, init_value=init_sin), - TensorSpec("cmp_kv_cache", [CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], torch.bfloat16, init_value=init_cmp_kv_cache, is_output=True), - TensorSpec("position_ids", [B, S], torch.int32, init_value=init_position_ids), - TensorSpec("cmp_slot_mapping", [B, S], torch.int64, init_value=init_cmp_slot_mapping), - TensorSpec("state_slot_mapping", [B, S], torch.int64, init_value=init_state_slot_mapping), - ] - - -if __name__ == "__main__": - import argparse - from golden import ratio_allclose, run_jit - - parser = argparse.ArgumentParser() - parser.add_argument("-p", "--platform", type=str, default="a2a3", - choices=["a2a3", "a2a3sim", "a5", "a5sim"]) - parser.add_argument("-d", "--device", type=int, default=0) - parser.add_argument("--start-pos", type=int, default=None, - help="Uniform fixture-only start_pos override for all batches; " - "default (unset) uses the canonical per-batch HCA set that includes the 8k point.") - parser.add_argument("--enable-l2-swimlane", action="store_true", default=False) - parser.add_argument("--dump-passes", action="store_true", default=False) - args = parser.parse_args() - - result = run_jit( - fn=compressor_test, - specs=build_tensor_specs(args.start_pos), - golden_fn=golden_compressor, - compile_cfg=dict(dump_passes=args.dump_passes), - runtime_cfg=dict( - platform=args.platform, - device_id=args.device, - enable_l2_swimlane=args.enable_l2_swimlane, - ), - rtol=1e-3, - atol=1e-3, - # Precision reference: AscendC torch.ops.custom.compressor — - # ops-transformer/experimental/attention/compressor/tests/pytest/compressor_golden.py - compare_fn={ - "kv": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), - "compress_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), - "cmp_kv_cache": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), - }, - ) - if not result.passed: - if result.error: - print(result.error) - raise SystemExit(1) diff --git a/models/deepseek/v4/decode_compressor_ratio4.py b/models/deepseek/v4/decode_compressor_ratio4.py deleted file mode 100644 index 7bc58928..00000000 --- a/models/deepseek/v4/decode_compressor_ratio4.py +++ /dev/null @@ -1,568 +0,0 @@ -# Copyright (c) PyPTO Contributors. -# This program is free software, you can redistribute it and/or modify it under the terms and conditions of -# CANN Open Software License Agreement Version 2.0 (the "License"). -# Please refer to the License for details. You may not use this file except in compliance with the License. -# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, -# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. -# See LICENSE in the root of the software repository for the full text of the License. -# ----------------------------------------------------------------------------------------------------------- -"""DeepSeek-V4 KV Compressor (decode incremental, ratio=4 overlap). - -Uses overlapping state layout with 8 slots. -Front slots 0-3 at columns [0:HEAD_DIM], back slots 4-7 at columns [HEAD_DIM:OUT_DIM]. -Tree reduction for softmax+pool. State shift after compression.""" - - -import pypto.language as pl - -from config import ( - FLASH as M, - DECODE_BATCH, - DECODE_SEQ, - BLOCK_SIZE, - C4A_COMPRESSOR_BLOCK_SIZE, - DECODE_CMP_BLOCK_NUM, - KV_CMP_MAX_BLOCKS, - FP32_NEG_INF, -) - - -# model config -B = DECODE_BATCH -S = DECODE_SEQ -EPS = M.rms_norm_eps -D = M.hidden_size -HEAD_DIM = M.head_dim -HEAD_DIM_INV = 1.0 / HEAD_DIM -ROPE_HEAD_DIM = M.qk_rope_head_dim -NOPE_HEAD_DIM = M.nope_head_dim -MAX_SEQ_LEN = M.max_position_embeddings - -# kernel-local (ratio-4 overlapping compressor) -COMPRESS_RATIO = 4 -OVERLAP = COMPRESS_RATIO == 4 -COFF = 1 + int(OVERLAP) -OUT_DIM = COFF * HEAD_DIM -STATE_LEN = COFF * COMPRESS_RATIO -IDX_KV_LEN = MAX_SEQ_LEN // COMPRESS_RATIO -COMPRESS_STATE_BLOCK_SIZE = C4A_COMPRESSOR_BLOCK_SIZE -COMPRESS_STATE_PHYSICAL_BLOCKS = 65 -COMPRESS_STATE_MAX_BLOCKS = (MAX_SEQ_LEN + COMPRESS_STATE_BLOCK_SIZE - 1) // COMPRESS_STATE_BLOCK_SIZE -COMPRESS_STATE_BLOCK_NUM = B * COMPRESS_STATE_PHYSICAL_BLOCKS -COMPRESS_STATE_DIM = 2 * OUT_DIM -CMP_MAX_BLOCKS = KV_CMP_MAX_BLOCKS -CMP_BLOCK_NUM = DECODE_CMP_BLOCK_NUM - -# tiling -ROPE_TILE = 32 -K_TILE = 512 -OUT_TILE = 64 -B_TILE = 8 -MM_B_TILE = 16 -BS_PAD = ((B * S + MM_B_TILE - 1) // MM_B_TILE) * MM_B_TILE -HEAD_TILE = 64 -HEAD_DIM_TILE = 128 -RMS_PAD_TILE = 16 # pad B rows up to one 16-row block (min M for FP32 vec ops) -RMS_PAD_ROWS = RMS_PAD_TILE # single block; requires B <= RMS_PAD_TILE -assert B <= RMS_PAD_TILE - - -@pl.jit.inline -def compressor_ratio4( - x: pl.Tensor[[B, S, D], pl.BF16], - kv: pl.Tensor[[B, S, HEAD_DIM], pl.FP32], - compress_state: pl.Tensor[[COMPRESS_STATE_BLOCK_NUM, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32], - compress_state_block_table: pl.Tensor[[B, COMPRESS_STATE_MAX_BLOCKS], pl.INT32], - wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], - wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], - ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], - norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - cos: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], - sin: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], - cmp_kv_cache: pl.Tensor[[CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16], - position_ids: pl.Tensor[[B, S], pl.INT32], - cmp_slot_mapping: pl.Tensor[[B, S], pl.INT64], - state_slot_mapping: pl.Tensor[[B, S], pl.INT64], -): - x_flat = pl.reshape(x, [B * S, D]) - cmp4_kv_proj_pad = pl.create_tensor([BS_PAD, OUT_DIM], dtype=pl.FP32) - cmp4_score_proj_pad = pl.create_tensor([BS_PAD, OUT_DIM], dtype=pl.FP32) - compress_state_flat = pl.reshape(compress_state, [COMPRESS_STATE_BLOCK_NUM * COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM]) - kv_flat = pl.reshape(kv, [B * S, HEAD_DIM]) - cmp_kv_cache_flat = pl.reshape(cmp_kv_cache, [CMP_BLOCK_NUM * BLOCK_SIZE, HEAD_DIM]) - - for idx in pl.spmd(BS_PAD * OUT_DIM // (MM_B_TILE * OUT_TILE), name_hint="kv_score_proj"): - global_row0 = (idx // (OUT_DIM // OUT_TILE)) * MM_B_TILE - o0 = (idx % (OUT_DIM // OUT_TILE)) * OUT_TILE - kv_acc = pl.create_tensor([MM_B_TILE, OUT_TILE], dtype=pl.FP32) - score_acc = pl.create_tensor([MM_B_TILE, OUT_TILE], dtype=pl.FP32) - for kb in pl.pipeline(0, D // K_TILE, stage=2): - k0 = kb * K_TILE - x_rows = pl.min(MM_B_TILE, B * S - global_row0) - x_tile = pl.slice(x_flat, [MM_B_TILE, K_TILE], [global_row0, k0], valid_shape=[x_rows, K_TILE]) - # Weights stored transposed [OUT_DIM, D] and consumed via b_trans=True so the - # GM->L1 load is a DN2ZN (each [OUT_TILE, K_TILE] row is K-contiguous = long - # bursts) instead of ND2NZ on [K_TILE, OUT_TILE] (K strided = many short - # bursts). Cuts the transaction-bound MTE2 cost ~14% busy / ~7% compressor wall. - wkv_tile = wkv[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] - wgate_tile = wgate[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] - if k0 == 0: - kv_acc = pl.matmul(x_tile, wkv_tile, out_dtype=pl.FP32, b_trans=True) - score_acc = pl.matmul(x_tile, wgate_tile, out_dtype=pl.FP32, b_trans=True) - else: - kv_acc = pl.matmul_acc(kv_acc, x_tile, wkv_tile, b_trans=True) - score_acc = pl.matmul_acc(score_acc, x_tile, wgate_tile, b_trans=True) - - cmp4_kv_proj_pad[global_row0 : global_row0 + MM_B_TILE, o0 : o0 + OUT_TILE] = kv_acc - cmp4_score_proj_pad[global_row0 : global_row0 + MM_B_TILE, o0 : o0 + OUT_TILE] = score_acc - - # scatter_softmax_pool: per batch, scatter the padded proj rows into compress_state, then - # online-softmax pool that batch's window into pooled_kv. One region -- each batch's pool - # reads only its own just-scattered state (per-batch block table), so no cross-task barrier. - pooled_kv = pl.create_tensor([RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) - with pl.at(level=pl.Level.CORE_GROUP, name_hint="scatter_softmax_pool"): - for c_idx in pl.range(B): - for s_sc in pl.pipeline(S, stage=2): - token_pos = pl.read(position_ids, [c_idx, s_sc]) - state_row_i64 = pl.read(state_slot_mapping, [c_idx, s_sc]) - proj_row = c_idx * S + s_sc - token_ape_row = pl.cast(token_pos % COMPRESS_RATIO, target_type=pl.INDEX) - if state_row_i64 >= 0: - state_row = pl.cast(state_row_i64, pl.INDEX) - kv_tile = cmp4_kv_proj_pad[proj_row : proj_row + 1, 0 : OUT_DIM] - score_tile = cmp4_score_proj_pad[proj_row : proj_row + 1, 0 : OUT_DIM] - ape_tile = ape[token_ape_row : token_ape_row + 1, 0 : OUT_DIM] - score_tile = pl.add(score_tile, ape_tile) - compress_state_flat[state_row : state_row + 1, 0 : OUT_DIM] = kv_tile - compress_state_flat[state_row : state_row + 1, OUT_DIM : 2 * OUT_DIM] = score_tile - - pad_idx = c_idx - first_pos_b = pl.read(position_ids, [c_idx, 0]) - pos_b = first_pos_b % COMPRESS_RATIO - pre_tokens_b = COMPRESS_RATIO - pos_b - boundary_end_b = first_pos_b + pre_tokens_b - 1 - cur_window_start_b = boundary_end_b - COMPRESS_RATIO + 1 - prev_window_start_b = cur_window_start_b - COMPRESS_RATIO - - if pos_b + S >= COMPRESS_RATIO: - last_abs = cur_window_start_b + COMPRESS_RATIO - 1 - last_blk_off = last_abs // COMPRESS_STATE_BLOCK_SIZE - last_intra = last_abs % COMPRESS_STATE_BLOCK_SIZE - last_blk_id = pl.cast(pl.read(compress_state_block_table, [c_idx, last_blk_off]), pl.INDEX) - last_row = last_blk_id * COMPRESS_STATE_BLOCK_SIZE + last_intra - mi = compress_state_flat[last_row : last_row + 1, OUT_DIM + HEAD_DIM : COMPRESS_STATE_DIM] - li = pl.exp(pl.sub(mi, mi)) - oi = compress_state_flat[last_row : last_row + 1, HEAD_DIM : OUT_DIM] - - for s in pl.range(0, COMPRESS_RATIO): - prev_abs = prev_window_start_b + s - front_score = pl.full([1, HEAD_DIM], dtype=pl.FP32, value=FP32_NEG_INF) - front_kv = pl.full([1, HEAD_DIM], dtype=pl.FP32, value=0.0) - if first_pos_b >= COMPRESS_RATIO: - prev_blk_off = prev_abs // COMPRESS_STATE_BLOCK_SIZE - prev_intra = prev_abs % COMPRESS_STATE_BLOCK_SIZE - prev_blk_id = pl.cast(pl.read(compress_state_block_table, [c_idx, prev_blk_off]), pl.INDEX) - prev_row = prev_blk_id * COMPRESS_STATE_BLOCK_SIZE + prev_intra - front_score = compress_state_flat[prev_row : prev_row + 1, OUT_DIM : OUT_DIM + HEAD_DIM] - front_kv = compress_state_flat[prev_row : prev_row + 1, 0 : HEAD_DIM] - mi_next_front = pl.maximum(mi, front_score) - alpha_front = pl.exp(pl.sub(mi, mi_next_front)) - beta_front = pl.exp(pl.sub(front_score, mi_next_front)) - li = pl.add(pl.mul(alpha_front, li), beta_front) - oi = pl.add(pl.mul(oi, alpha_front), pl.mul(front_kv, beta_front)) - mi = mi_next_front - - for s in pl.range(0, COMPRESS_RATIO - 1): - cur_abs = cur_window_start_b + s - cur_blk_off = cur_abs // COMPRESS_STATE_BLOCK_SIZE - cur_intra = cur_abs % COMPRESS_STATE_BLOCK_SIZE - cur_blk_id = pl.cast(pl.read(compress_state_block_table, [c_idx, cur_blk_off]), pl.INDEX) - cur_row = cur_blk_id * COMPRESS_STATE_BLOCK_SIZE + cur_intra - back_score = compress_state_flat[cur_row : cur_row + 1, OUT_DIM + HEAD_DIM : COMPRESS_STATE_DIM] - back_kv = compress_state_flat[cur_row : cur_row + 1, HEAD_DIM : OUT_DIM] - mi_next_back = pl.maximum(mi, back_score) - alpha_back = pl.exp(pl.sub(mi, mi_next_back)) - beta_back = pl.exp(pl.sub(back_score, mi_next_back)) - li = pl.add(pl.mul(alpha_back, li), beta_back) - oi = pl.add(pl.mul(oi, alpha_back), pl.mul(back_kv, beta_back)) - mi = mi_next_back - - pooled_chunk = pl.div(oi, li) - pooled_kv[pad_idx : pad_idx + 1, 0 : HEAD_DIM] = pooled_chunk - - normed_kv = pl.create_tensor([RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) - norm_w_2d = pl.reshape(norm_w, [1, HEAD_DIM]) - with pl.at(level=pl.Level.CORE_GROUP, name_hint="rmsnorm_rope_cache_write"): - # single 16-row block: B real rows at rows 0..B-1, rows B..15 are pad - cos_b = pl.full([RMS_PAD_TILE, ROPE_HEAD_DIM // 2], dtype=pl.FP32, value=0.0) - sin_b = pl.full([RMS_PAD_TILE, ROPE_HEAD_DIM // 2], dtype=pl.FP32, value=0.0) - cos_b[0:B, 0 : ROPE_HEAD_DIM // 2] = cos[0:B, 0 : ROPE_HEAD_DIM // 2] - sin_b[0:B, 0 : ROPE_HEAD_DIM // 2] = sin[0:B, 0 : ROPE_HEAD_DIM // 2] - partial_sq = pl.full([1, RMS_PAD_TILE], dtype=pl.FP32, value=0.0) - for k0 in pl.range(0, HEAD_DIM, HEAD_TILE): - kv_rms_chunk = pooled_kv[0 : RMS_PAD_TILE, k0 : k0 + HEAD_TILE] - kv_rms_sq = pl.mul(kv_rms_chunk, kv_rms_chunk) - kv_rms_rowsum = pl.reshape(pl.row_sum(kv_rms_sq), [1, RMS_PAD_TILE]) - partial_sq = pl.add(partial_sq, kv_rms_rowsum) - - variance = pl.reshape(pl.add(pl.mul(partial_sq, HEAD_DIM_INV), EPS), [RMS_PAD_TILE, 1]) - inv_rms = pl.recip(pl.sqrt(variance)) - for k0 in pl.range(0, NOPE_HEAD_DIM, HEAD_TILE): - kv_norm_chunk = pooled_kv[0 : RMS_PAD_TILE, k0 : k0 + HEAD_TILE] - gamma = pl.cast(norm_w_2d[:, k0 : k0 + HEAD_TILE], pl.FP32) - normed_chunk = pl.col_expand_mul(pl.row_expand_mul(kv_norm_chunk, inv_rms), gamma) - normed_kv[0 : RMS_PAD_TILE, k0 : k0 + HEAD_TILE] = normed_chunk - - kv_rope_norm = pooled_kv[0 : RMS_PAD_TILE, NOPE_HEAD_DIM : HEAD_DIM] - gamma_rope = pl.cast(norm_w_2d[:, NOPE_HEAD_DIM : HEAD_DIM], pl.FP32) - # A3 interleaved swap-gather (same form as kv_rope_fused in qkv_proj_rope), - # replacing the de-interleave gather + rotate + re-interleave scatter. gamma+inv_rms - # are folded into rope_normed BEFORE the swap, so the swapped lane n[j^1] correctly - # carries gamma[j^1]; inv_rms is per-row so it commutes. swap_idx (j^1), sign - # ([-1,+1,...]) and dup_idx (j>>1) are built IN-KERNEL from pl.arange; cos_il/sin_il - # are dup-gathered from the per-batch cos/sin rows. normed_kv is FP32 -> write directly. - # out[j] = n[j]*cos_il[j] + n[j^1]*sign[j]*sin_il[j] - rope_normed = pl.col_expand_mul(pl.row_expand_mul(kv_rope_norm, inv_rms), gamma_rope) - rope_ones = pl.full([RMS_PAD_TILE, ROPE_HEAD_DIM], dtype=pl.FP32, value=1.0) - rope_col = pl.col_expand_mul(rope_ones, pl.cast(pl.arange(0, [1, ROPE_HEAD_DIM], dtype=pl.INT32), target_type=pl.FP32)) - rope_dup_f = pl.cast(pl.cast(pl.mul(rope_col, 0.5), target_type=pl.INT32, mode="trunc"), target_type=pl.FP32) - rope_dup_idx = pl.cast(rope_dup_f, target_type=pl.INT32) # j>>1 - rope_lane = pl.sub(rope_col, pl.mul(rope_dup_f, 2.0)) # j%2 - rope_swap_idx = pl.cast(pl.sub(pl.add(rope_col, 1.0), pl.mul(rope_lane, 2.0)), target_type=pl.INT32) # j^1 - rope_sign = pl.sub(pl.mul(rope_lane, 2.0), 1.0) # [-1,+1,...] - cos_il = pl.gather(cos_b, dim=-1, index=rope_dup_idx) - sin_il = pl.gather(sin_b, dim=-1, index=rope_dup_idx) - swapped = pl.gather(rope_normed, dim=-1, index=rope_swap_idx) - rope_rot = pl.add(pl.mul(rope_normed, cos_il), pl.mul(pl.mul(swapped, rope_sign), sin_il)) - normed_kv[0 : RMS_PAD_TILE, NOPE_HEAD_DIM : HEAD_DIM] = rope_rot - - # cache write: reads back only this block's own normed_kv rows, so the normed_kv - # RAW is intra-block -- no separate scope / cross-task barrier needed. - for inner in pl.range(B): - c_idx = inner - first_pos_b = pl.read(position_ids, [c_idx, 0]) - pos_b = first_pos_b % COMPRESS_RATIO - if pos_b + S >= COMPRESS_RATIO: - boundary_s = COMPRESS_RATIO - 1 - pos_b - kv_row_fp32 = normed_kv[inner : inner + 1, 0 : HEAD_DIM] - cache_row_i64 = pl.read(cmp_slot_mapping, [c_idx, boundary_s]) - if cache_row_i64 >= 0: - cache_row = pl.cast(cache_row_i64, pl.INDEX) - kv_flat[c_idx * S : c_idx * S + 1, :] = kv_row_fp32 - cmp_kv_cache_flat[cache_row : cache_row + 1, :] = pl.cast(kv_row_fp32, target_type=pl.BF16, mode="rint") - - kv = pl.reshape(kv_flat, [B, S, HEAD_DIM]) - return kv - - -@pl.jit -def compressor_test( - x: pl.Tensor[[B, S, D], pl.BF16], - kv: pl.Out[pl.Tensor[[B, S, HEAD_DIM], pl.FP32]], - compress_state: pl.InOut[pl.Tensor[[COMPRESS_STATE_BLOCK_NUM, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32]], - compress_state_block_table: pl.Tensor[[B, COMPRESS_STATE_MAX_BLOCKS], pl.INT32], - wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], - wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], - ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], - norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - cos: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], - sin: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], - cmp_kv_cache: pl.InOut[pl.Tensor[[CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], - position_ids: pl.Tensor[[B, S], pl.INT32], - cmp_slot_mapping: pl.Tensor[[B, S], pl.INT64], - state_slot_mapping: pl.Tensor[[B, S], pl.INT64], -): - compressor_ratio4( - x, - kv, - compress_state, - compress_state_block_table, - wkv, - wgate, - ape, - norm_w, - cos, - sin, - cmp_kv_cache, - position_ids, - cmp_slot_mapping, - state_slot_mapping, - ) - return kv, compress_state, cmp_kv_cache - - -def golden_compressor(tensors): - """Torch reference for Compressor.forward (decode branch, ratio=4 overlap).""" - import torch - - x = tensors["x"].float() - compress_state = tensors["compress_state"] - compress_state_block_table = tensors["compress_state_block_table"] - wkv = tensors["wkv"].float() - wgate = tensors["wgate"].float() - ape = tensors["ape"] - norm_w = tensors["norm_w"] - cos = tensors["cos"] - sin = tensors["sin"] - cmp_kv_cache = tensors["cmp_kv_cache"] - position_ids = tensors["position_ids"].to(torch.int64) - cmp_slot_mapping = tensors["cmp_slot_mapping"].to(torch.int64) - state_slot_mapping = tensors["state_slot_mapping"].to(torch.int64) - bsz, _, _ = x.shape - ratio, rd = COMPRESS_RATIO, ROPE_HEAD_DIM - - kv = x @ wkv.t() # [B, S, OUT_DIM] (wkv stored [OUT_DIM, D] for b_trans) - score = x @ wgate.t() # [B, S, OUT_DIM] - - pooled = torch.zeros(bsz, 1, HEAD_DIM, dtype=torch.float32, device=x.device) - should_compress_rows = torch.zeros(bsz, dtype=torch.bool, device=x.device) - - def read_front_state(b, abs_pos): - blk_id = int(compress_state_block_table[b, abs_pos // COMPRESS_STATE_BLOCK_SIZE].item()) - if blk_id < 0: - return ( - torch.zeros(HEAD_DIM, dtype=torch.float32, device=x.device), - torch.full((HEAD_DIM,), float("-inf"), dtype=torch.float32, device=x.device), - ) - intra = abs_pos % COMPRESS_STATE_BLOCK_SIZE - return ( - compress_state[blk_id, intra, :HEAD_DIM], - compress_state[blk_id, intra, OUT_DIM:OUT_DIM + HEAD_DIM], - ) - - def read_back_state(b, abs_pos): - blk_id = int(compress_state_block_table[b, abs_pos // COMPRESS_STATE_BLOCK_SIZE].item()) - if blk_id < 0: - return ( - torch.zeros(HEAD_DIM, dtype=torch.float32, device=x.device), - torch.full((HEAD_DIM,), float("-inf"), dtype=torch.float32, device=x.device), - ) - intra = abs_pos % COMPRESS_STATE_BLOCK_SIZE - return ( - compress_state[blk_id, intra, HEAD_DIM:OUT_DIM], - compress_state[blk_id, intra, OUT_DIM + HEAD_DIM:], - ) - - for b in range(bsz): - first_pos = int(position_ids[b, 0].item()) - pre_tokens = min(S, ratio - (first_pos % ratio)) - boundary_s = ratio - 1 - (first_pos % ratio) - should_compress = 0 <= boundary_s < S - boundary_end = first_pos + pre_tokens - 1 - cur_window_start = boundary_end - ratio + 1 - prev_window_start = cur_window_start - ratio - - # Per-token ape add + state scatter through explicit token-major slots. - for s in range(S): - pos = int(position_ids[b, s].item()) - token_ape_row = pos % ratio - score[b, s, :] = score[b, s, :] + ape[token_ape_row] - state_row = int(state_slot_mapping[b, s].item()) - if state_row >= 0: - blk_id = state_row // COMPRESS_STATE_BLOCK_SIZE - intra = state_row % COMPRESS_STATE_BLOCK_SIZE - compress_state[blk_id, intra, :OUT_DIM] = kv[b, s, :] - compress_state[blk_id, intra, OUT_DIM:] = score[b, s, :] - - if should_compress: - should_compress_rows[b] = True - kv_rows = [] - score_rows = [] - for s in range(ratio): - abs_pos = prev_window_start + s - if abs_pos < 0: - kv_rows.append(torch.zeros(HEAD_DIM, dtype=torch.float32, device=x.device)) - score_rows.append(torch.full((HEAD_DIM,), float("-inf"), dtype=torch.float32, device=x.device)) - continue - kv_row, score_row = read_front_state(b, abs_pos) - kv_rows.append(kv_row) - score_rows.append(score_row) - for s in range(ratio): - abs_pos = cur_window_start + s - kv_row, score_row = read_back_state(b, abs_pos) - kv_rows.append(kv_row) - score_rows.append(score_row) - kvs = torch.stack(kv_rows, dim=0).unsqueeze(0) - scs = torch.stack(score_rows, dim=0).unsqueeze(0) - pooled[b : b + 1] = (kvs * scs.softmax(dim=1)).sum(dim=1, keepdim=True) - - tensors["compress_state"][:] = compress_state - - if not bool(should_compress_rows.any()): - return - - def rmsnorm(x, w): - x = x.float() - var = x.square().mean(-1, keepdim=True) - x = x * torch.rsqrt(var + EPS) - return w * x - - for b in range(bsz): - if not bool(should_compress_rows[b]): - continue - first_pos = int(position_ids[b, 0].item()) - boundary_s = ratio - 1 - (first_pos % ratio) - kv_b = rmsnorm(pooled[b : b + 1], norm_w) - - x_pair = kv_b[..., -rd:].unflatten(-1, (-1, 2)) - x0, x1 = x_pair[..., 0], x_pair[..., 1] - cos_v, sin_v = cos[b].view(-1), sin[b].view(-1) - y0 = x0 * cos_v - x1 * sin_v - y1 = x0 * sin_v + x1 * cos_v - - kv_b = torch.cat([kv_b[..., :-rd], torch.stack([y0, y1], dim=-1).flatten(-2)], dim=-1) - - cmp_row = int(cmp_slot_mapping[b, boundary_s].item()) - if cmp_row >= 0: - # Kernel writes committed pooled result only to kv[:, 0, :]; leave - # speculative-boundary rows and kv[:, 1:, :] zero-initialized. - tensors["kv"][b : b + 1, 0:1, :] = kv_b - blk_id = cmp_row // BLOCK_SIZE - cmp_kv_cache[blk_id, cmp_row % BLOCK_SIZE, 0] = kv_b[0, 0] - - tensors["cmp_kv_cache"][:] = cmp_kv_cache - - -def build_tensor_specs(start_pos=None): - import torch # type: ignore[import] - from decode_metadata import ( - block_table, - compressed_slot_mapping, - csa_decode_start_set, - position_ids_from_starts, - resolve_start_positions, - state_slot_mapping, - ) - from golden import TensorSpec - from rope_tables import build_deepseek_v4_rope_tables, materialize_half_rope_tables - - shared_freqs_cos, shared_freqs_sin = build_deepseek_v4_rope_tables(M, COMPRESS_RATIO, dtype=torch.bfloat16) - - def init_x(): - return torch.rand(B, S, D) - def init_compress_state(): - state = torch.zeros(COMPRESS_STATE_BLOCK_NUM, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM) - state[:, :, OUT_DIM:] = FP32_NEG_INF - return state - def init_compress_state_block_table(): - return block_table( - batch=B, - table_blocks=COMPRESS_STATE_MAX_BLOCKS, - physical_blocks=COMPRESS_STATE_PHYSICAL_BLOCKS, - ) - # Calibrated to the real DeepSeek-V4-Flash CSA (ratio-4) main compressor (mean l8/l32 of - # extract_weights_flash): zero-mean Gaussian BF16 weights at the measured std; the RMSNorm - # gamma centers near the measured mean (not ones / not uniform). - def init_wkv(): - return torch.randn(OUT_DIM, D) * 0.0245 - def init_wgate(): - return torch.randn(OUT_DIM, D) * 0.0388 - def init_ape(): - return torch.randn(COMPRESS_RATIO, OUT_DIM) * 0.1243 - def init_norm_w(): - return 0.9666 + 0.1929 * torch.randn(HEAD_DIM) - def init_rope_positions(): - first_pos = init_position_ids().to(torch.int64)[:, 0] - cmp_offset = COMPRESS_RATIO - (first_pos % COMPRESS_RATIO) - return (first_pos + cmp_offset - COMPRESS_RATIO).to(torch.int64) - def init_cos(): - return materialize_half_rope_tables(shared_freqs_cos, shared_freqs_sin, init_rope_positions())[0] - def init_sin(): - return materialize_half_rope_tables(shared_freqs_cos, shared_freqs_sin, init_rope_positions())[1] - def init_cmp_kv_cache(): - return torch.zeros(CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM) - def init_cmp_block_table(): - tbl = torch.full((B, CMP_MAX_BLOCKS), -1, dtype=torch.int32) - for b in range(B): - for j in range(CMP_MAX_BLOCKS): - tbl[b, j] = b * CMP_MAX_BLOCKS + j - return tbl - def init_default_start_pos(): - # Canonical CSA start-position set (ratio-4 compressor + indexer + sliding-window + 8k). - return csa_decode_start_set( - batch=B, seq=S, compress_ratio=COMPRESS_RATIO, - state_block_size=COMPRESS_STATE_BLOCK_SIZE) - def init_start_pos(): - return resolve_start_positions( - start_pos, - batch=B, - seq=S, - max_seq_len=MAX_SEQ_LEN, - default_fn=init_default_start_pos, - ) - def init_position_ids(): - return position_ids_from_starts(init_start_pos(), seq=S) - def init_state_slot_mapping(): - return state_slot_mapping( - init_position_ids(), - init_compress_state_block_table(), - state_block_size=COMPRESS_STATE_BLOCK_SIZE, - ) - def init_cmp_slot_mapping(): - positions = init_position_ids() - return compressed_slot_mapping( - positions, - init_cmp_block_table(), - compress_ratio=COMPRESS_RATIO, - block_size=BLOCK_SIZE, - ) - - return [ - TensorSpec("x", [B, S, D], torch.bfloat16, init_value=init_x), - TensorSpec("kv", [B, S, HEAD_DIM], torch.float32, is_output=True), - TensorSpec("compress_state", [COMPRESS_STATE_BLOCK_NUM, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], torch.float32, init_value=init_compress_state, is_output=True), - TensorSpec("compress_state_block_table", [B, COMPRESS_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), - TensorSpec("wkv", [OUT_DIM, D], torch.bfloat16, init_value=init_wkv), - TensorSpec("wgate", [OUT_DIM, D], torch.bfloat16, init_value=init_wgate), - TensorSpec("ape", [COMPRESS_RATIO, OUT_DIM], torch.float32, init_value=init_ape), - TensorSpec("norm_w", [HEAD_DIM], torch.bfloat16, init_value=init_norm_w), - TensorSpec("cos", [B, ROPE_HEAD_DIM // 2], torch.float32, init_value=init_cos), - TensorSpec("sin", [B, ROPE_HEAD_DIM // 2], torch.float32, init_value=init_sin), - TensorSpec("cmp_kv_cache", [CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], torch.bfloat16, init_value=init_cmp_kv_cache, is_output=True), - TensorSpec("position_ids", [B, S], torch.int32, init_value=init_position_ids), - TensorSpec("cmp_slot_mapping", [B, S], torch.int64, init_value=init_cmp_slot_mapping), - TensorSpec("state_slot_mapping", [B, S], torch.int64, init_value=init_state_slot_mapping), - ] - - -if __name__ == "__main__": - import argparse - from golden import ratio_allclose, run_jit - - parser = argparse.ArgumentParser() - parser.add_argument("-p", "--platform", type=str, default="a2a3", - choices=["a2a3", "a2a3sim", "a5", "a5sim"]) - parser.add_argument("-d", "--device", type=int, default=0) - parser.add_argument("--start-pos", type=int, default=None, - help="Uniform fixture-only start_pos override for all batches; " - "default (unset) uses the canonical per-batch CSA set that includes the 8k point.") - parser.add_argument("--enable-l2-swimlane", action="store_true", default=False) - parser.add_argument("--runtime-dir", type=str, default=None) - parser.add_argument("--golden-data", type=str, default=None) - parser.add_argument("--dump-passes", action="store_true", default=False) - args = parser.parse_args() - - result = run_jit( - fn=compressor_test, - specs=build_tensor_specs(args.start_pos), - golden_fn=golden_compressor, - runtime_dir=args.runtime_dir, - golden_data=args.golden_data, - compile_cfg=dict(dump_passes=args.dump_passes), - runtime_cfg=dict( - platform=args.platform, - device_id=args.device, - enable_l2_swimlane=args.enable_l2_swimlane, - ), - rtol=1e-3, - atol=1e-3, - compare_fn={ - "kv": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), - "compress_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), - "cmp_kv_cache": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), - }, - ) - if not result.passed: - if result.error: - print(result.error) - raise SystemExit(1) diff --git a/models/deepseek/v4/decode_indexer.py b/models/deepseek/v4/decode_indexer.py index c9e4fd48..061f90fa 100644 --- a/models/deepseek/v4/decode_indexer.py +++ b/models/deepseek/v4/decode_indexer.py @@ -96,16 +96,16 @@ @pl.jit.inline def indexer( - x: pl.Tensor[[B, S, D], pl.BF16], + x: pl.Tensor[[T, D], pl.BF16], qr: pl.Tensor[[T, Q_LORA], pl.INT8], qr_scale: pl.Tensor[[T, 1], pl.FP32], wq_b: pl.Tensor[[Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8], wq_b_scale: pl.Tensor[[IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32], weights_proj: pl.Tensor[[D, IDX_N_HEADS], pl.BF16], - cos: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], - sin: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], hadamard: pl.Tensor[[IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16], # shared by q rotation and inner Compressor - inner_kv: pl.Tensor[[B, S, INNER_HEAD_DIM], pl.FP32], + inner_kv: pl.Tensor[[T, INNER_HEAD_DIM], pl.FP32], inner_compress_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32], inner_compress_state_block_table: pl.Tensor[[B, INNER_STATE_MAX_BLOCKS], pl.INT32], inner_wkv: pl.Tensor[[INNER_OUT_DIM, D], pl.BF16], @@ -116,11 +116,11 @@ def indexer( idx_kv_cache: pl.InOut[pl.Tensor[[IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8]], idx_kv_scale: pl.InOut[pl.Tensor[[IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32]], idx_block_table: pl.Tensor[[B, IDX_CACHE_MAX_BLOCKS], pl.INT32], - score: pl.Tensor[[B, S, SCORE_LEN], pl.FP32], - topk_idxs: pl.Tensor[[B, S, SCORE_LEN], pl.INT32], - position_ids: pl.Tensor[[B, S], pl.INT32], - idx_slot_mapping: pl.Tensor[[B, S], pl.INT64], - inner_state_slot_mapping: pl.Tensor[[B, S], pl.INT64], + score: pl.Tensor[[T, SCORE_LEN], pl.FP32], + topk_idxs: pl.Tensor[[T, SCORE_LEN], pl.INT32], + position_ids: pl.Tensor[[T], pl.INT32], + idx_slot_mapping: pl.Tensor[[T], pl.INT64], + inner_state_slot_mapping: pl.Tensor[[T], pl.INT64], kv_seq_lens: pl.Tensor[[B], pl.INT32], offset: pl.Scalar[pl.INT32], ): @@ -148,15 +148,16 @@ def indexer( qr_proj_flat = pl.reshape(qr_proj, [T * IDX_N_HEADS, IDX_HEAD_DIM]) qr_rope_out = pl.create_tensor([T * IDX_N_HEADS, ROPE_HEAD_DIM], dtype=pl.BF16) - # spmd over ROPE_ROW_TILE-row blocks; batch_idx = block base // ROPE_ROW_BLOCK - # picks the per-batch cos/sin row. Rotation indices/sign and cos_il/sin_il are - # built once per block. + # spmd over ROPE_ROW_TILE-row blocks; token_idx = row base // IDX_N_HEADS + # picks the token-major freqs row. Rotation indices/sign and cos_il/sin_il + # are built once per block. # out[j] = x[j]*cos_il[j] + x[j^1]*sign[j]*sin_il[j] (sign folded into sin_il_signed) for idx in pl.spmd(T * IDX_N_HEADS // ROPE_ROW_TILE, name_hint="qr_rope"): o0 = idx * ROPE_ROW_TILE - batch_idx = o0 // ROPE_ROW_BLOCK - cos_b = cos[batch_idx : batch_idx + 1, 0 : ROPE_HEAD_DIM // 2] - sin_b = sin[batch_idx : batch_idx + 1, 0 : ROPE_HEAD_DIM // 2] + token_idx = o0 // IDX_N_HEADS + pos_b = pl.cast(pl.read(position_ids, [token_idx]), pl.INDEX) + cos_b = pl.cast(freqs_cos[pos_b : pos_b + 1, 0 : ROPE_HEAD_DIM // 2], target_type=pl.FP32) + sin_b = pl.cast(freqs_sin[pos_b : pos_b + 1, 0 : ROPE_HEAD_DIM // 2], target_type=pl.FP32) rope_ones = pl.full([ROPE_ROW_TILE, ROPE_HEAD_DIM], dtype=pl.FP32, value=1.0) rope_col = pl.col_expand_mul(rope_ones, pl.cast(pl.arange(0, [1, ROPE_HEAD_DIM], dtype=pl.INT32), target_type=pl.FP32)) rope_dup_f = pl.cast(pl.cast(pl.mul(rope_col, 0.5), target_type=pl.INT32, mode="trunc"), target_type=pl.FP32) @@ -215,18 +216,21 @@ def indexer( weights_acc = pl.matmul_acc(weights_acc, x_tile, weights_proj_tile) weights[0:MM_ROW_TILE, :] = pl.mul(weights_acc, WEIGHTS_SCALE) + inner_kv_flat = pl.reshape(inner_kv, [T, INNER_HEAD_DIM]) + idx_slot_mapping_flat = pl.reshape(idx_slot_mapping, [T]) + inner_state_slot_mapping_flat = pl.reshape(inner_state_slot_mapping, [T]) indexer_compressor( - x, inner_kv, + x_flat, inner_kv_flat, inner_compress_state, inner_compress_state_block_table, inner_wkv, inner_wgate, inner_ape, inner_norm_w, - cos, sin, hadamard, idx_kv_cache, idx_kv_scale, - position_ids, idx_slot_mapping, inner_state_slot_mapping, + freqs_cos, freqs_sin, hadamard, idx_kv_cache, idx_kv_scale, + position_ids, idx_slot_mapping_flat, inner_state_slot_mapping_flat, ) kv_cache_i8_flat = pl.reshape(idx_kv_cache, [IDX_CACHE_BLOCK_NUM * BLOCK_SIZE, IDX_HEAD_DIM]) kv_scale_flat = pl.reshape(idx_kv_scale, [IDX_CACHE_BLOCK_NUM * BLOCK_SIZE, 1]) idx_block_table_flat = pl.reshape(idx_block_table, [B * IDX_CACHE_MAX_BLOCKS]) - score_flat = pl.reshape(score, [T, SCORE_LEN]) + score_flat = score # No score_init: reduce writes the valid region; the tail is never read (topk re-masks). # Two GM-handoff stages: matmul (cube, reads paged C8 directly) -> reduce (vec). @@ -258,7 +262,7 @@ def indexer( b = tg // S s = tg - b * S clen_b = pl.read(kv_seq_lens, [b]) // COMPRESS_RATIO - pos_t = pl.read(position_ids, [b, s]) + pos_t = pl.read(position_ids, [tg]) visible_len_t = pl.min(pl.min(clen_b, (pos_t + 1) // COMPRESS_RATIO), SCORE_LEN) cblk_t = (visible_len_t + REDUCE_TILE - 1) // REDUCE_TILE tb = b * S @@ -292,14 +296,13 @@ def indexer( ) score_flat[tb + s : tb + s + 1, cache0 : cache0 + REDUCE_TILE] = weighted_score_valid_s - topk_idxs_flat = pl.reshape(topk_idxs, [T, SCORE_LEN]) + topk_idxs_flat = topk_idxs for t in pl.spmd(T, name_hint="topk"): invalid_idxs = pl.full([1, SCORE_LEN], dtype=pl.INT32, value=-1) topk_idxs_flat[t : t + 1, :] = invalid_idxs batch_idx = t // S - token_s = t - batch_idx * S cache_len_b = pl.read(kv_seq_lens, [batch_idx]) // COMPRESS_RATIO - pos_t = pl.read(position_ids, [batch_idx, token_s]) + pos_t = pl.read(position_ids, [t]) visible_len_t = pl.min(pl.min(cache_len_b, (pos_t + 1) // COMPRESS_RATIO), SCORE_LEN) if visible_len_t > 0: offset_i32 = pl.cast(offset, target_type=pl.INT32) @@ -329,16 +332,16 @@ def indexer( @pl.jit def indexer_test( - x: pl.Tensor[[B, S, D], pl.BF16], + x: pl.Tensor[[T, D], pl.BF16], qr: pl.Tensor[[T, Q_LORA], pl.INT8], qr_scale: pl.Tensor[[T, 1], pl.FP32], wq_b: pl.Tensor[[Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8], wq_b_scale: pl.Tensor[[IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32], weights_proj: pl.Tensor[[D, IDX_N_HEADS], pl.BF16], - cos: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], - sin: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], hadamard: pl.Tensor[[IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16], - inner_kv: pl.Tensor[[B, S, INNER_HEAD_DIM], pl.FP32], + inner_kv: pl.Tensor[[T, INNER_HEAD_DIM], pl.FP32], inner_compress_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32], inner_compress_state_block_table: pl.Tensor[[B, INNER_STATE_MAX_BLOCKS], pl.INT32], inner_wkv: pl.Tensor[[INNER_OUT_DIM, D], pl.BF16], @@ -348,11 +351,11 @@ def indexer_test( idx_kv_cache: pl.InOut[pl.Tensor[[IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8]], idx_kv_scale: pl.InOut[pl.Tensor[[IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32]], idx_block_table: pl.Tensor[[B, IDX_CACHE_MAX_BLOCKS], pl.INT32], - score: pl.Out[pl.Tensor[[B, S, SCORE_LEN], pl.FP32]], - topk_idxs: pl.Out[pl.Tensor[[B, S, SCORE_LEN], pl.INT32]], - position_ids: pl.Tensor[[B, S], pl.INT32], - idx_slot_mapping: pl.Tensor[[B, S], pl.INT64], - inner_state_slot_mapping: pl.Tensor[[B, S], pl.INT64], + score: pl.Out[pl.Tensor[[T, SCORE_LEN], pl.FP32]], + topk_idxs: pl.Out[pl.Tensor[[T, SCORE_LEN], pl.INT32]], + position_ids: pl.Tensor[[T], pl.INT32], + idx_slot_mapping: pl.Tensor[[T], pl.INT64], + inner_state_slot_mapping: pl.Tensor[[T], pl.INT64], kv_seq_lens: pl.Tensor[[B], pl.INT32], offset: pl.Scalar[pl.INT32], ): @@ -363,8 +366,8 @@ def indexer_test( wq_b, wq_b_scale, weights_proj, - cos, - sin, + freqs_cos, + freqs_sin, hadamard, inner_kv, inner_compress_state, @@ -444,14 +447,13 @@ def golden_indexer(tensors): wq_b = tensors["wq_b"] wq_b_scale = tensors["wq_b_scale"].float() weights_proj = tensors["weights_proj"].float() - cos = tensors["cos"] - sin = tensors["sin"] + freqs_cos = tensors["freqs_cos"] + freqs_sin = tensors["freqs_sin"] hadamard = tensors["hadamard"].float() kv_seq_lens = tensors["kv_seq_lens"].to(torch.int64) offset = int(tensors["offset"]) - bsz, seqlen, _ = x.shape ratio, rd = COMPRESS_RATIO, ROPE_HEAD_DIM q_i32 = qr.to(torch.int32) @ wq_b.to(torch.int32) @@ -459,8 +461,9 @@ def golden_indexer(tensors): x_pair = q[..., -rd:].unflatten(-1, (-1, 2)) x0, x1 = x_pair[..., 0], x_pair[..., 1] - cos_v = cos.view(B, 1, 1, -1) - sin_v = sin.view(B, 1, 1, -1) + rope_pos = tensors["position_ids"].to(torch.int64).reshape(T) + cos_v = freqs_cos.index_select(0, rope_pos)[:, : rd // 2].float().view(B, S, 1, -1) + sin_v = freqs_sin.index_select(0, rope_pos)[:, : rd // 2].float().view(B, S, 1, -1) y0 = (x0 * cos_v - x1 * sin_v).to(torch.bfloat16) y1 = (x0 * sin_v + x1 * cos_v).to(torch.bfloat16) @@ -471,39 +474,41 @@ def golden_indexer(tensors): # then dequantized with q_scale * kv_scale. # flash: fp4_act_quant on q (FP4 simulation). + inner_kv_flat = tensors["inner_kv"].contiguous() inner_tensors = { - "x": tensors["x"], - "kv": tensors["inner_kv"], + "x": tensors["x"].contiguous(), + "kv": inner_kv_flat, "wkv": tensors["inner_wkv"], "wgate": tensors["inner_wgate"], "ape": tensors["inner_ape"], "norm_w": tensors["inner_norm_w"], - "cos": tensors["cos"], - "sin": tensors["sin"], + "freqs_cos": tensors["freqs_cos"], + "freqs_sin": tensors["freqs_sin"], "hadamard": tensors["hadamard"], "compress_state": tensors["inner_compress_state"], "compress_state_block_table": tensors["inner_compress_state_block_table"], "idx_kv_cache": tensors["idx_kv_cache"], "idx_kv_scale": tensors["idx_kv_scale"], - "position_ids": tensors["position_ids"], - "idx_slot_mapping": tensors["idx_slot_mapping"], - "inner_state_slot_mapping": tensors["inner_state_slot_mapping"], + "position_ids": tensors["position_ids"].reshape(T).contiguous(), + "idx_slot_mapping": tensors["idx_slot_mapping"].reshape(T).contiguous(), + "inner_state_slot_mapping": tensors["inner_state_slot_mapping"].reshape(T).contiguous(), } golden_compressor(inner_tensors) + tensors["inner_kv"][:] = inner_kv_flat - weights = (x @ weights_proj) * WEIGHTS_SCALE + weights = ((x @ weights_proj) * WEIGHTS_SCALE).view(B, S, IDX_N_HEADS) # C8 cache: pre-quantized INT8 KV + per-position dequant scale (no score-time re-quant) idx_kv_cache_i8 = tensors["idx_kv_cache"] idx_kv_scale = tensors["idx_kv_scale"].float() idx_block_table = tensors["idx_block_table"] - score_full = torch.full((bsz, seqlen, SCORE_LEN), FP32_NEG_INF, dtype=torch.float32) - topk_idxs = torch.full((bsz, seqlen, SCORE_LEN), -1, dtype=torch.int32) + score_full = torch.full((T, SCORE_LEN), FP32_NEG_INF, dtype=torch.float32) + topk_idxs = torch.full((T, SCORE_LEN), -1, dtype=torch.int32) q_i8, q_scale = _int8_quant_per_row(q.reshape(B * S * IDX_N_HEADS, IDX_HEAD_DIM)) q_i8 = q_i8.view(B, S, IDX_N_HEADS, IDX_HEAD_DIM) q_scale = q_scale.view(B, S, IDX_N_HEADS, 1) - for b in range(bsz): + for b in range(B): cache_len = int(kv_seq_lens[b].item()) // ratio if cache_len <= 0: continue @@ -520,19 +525,21 @@ def golden_indexer(tensors): score = score_i32.float() * q_scale[b] score = (torch.relu(score) * weights[b].unsqueeze(-1)).sum(dim=1) score = score * kv_scale.view(1, cache_len) - for s in range(seqlen): - visible_len = min(cache_len, int(tensors["position_ids"][b, s].item() + 1) // ratio, SCORE_LEN) + row0 = b * S + pos_rows = tensors["position_ids"].reshape(B, S) + for s in range(S): + visible_len = min(cache_len, int(pos_rows[b, s].item() + 1) // ratio, SCORE_LEN) if visible_len <= 0: continue - score_full[b, s, :visible_len] = score[s, :visible_len].to(torch.float32) + score_full[row0 + s, :visible_len] = score[s, :visible_len].to(torch.float32) k = min(IDX_TOPK, visible_len) _, idx = score[s, :visible_len].topk(k, dim=-1) - topk_idxs[b, s, :k] = idx.to(torch.int32) - topk_idxs[b, s, :k] += offset + topk_idxs[row0 + s, :k] = idx.to(torch.int32) + topk_idxs[row0 + s, :k] += offset tensors["score"][:] = score_full - tensors["topk_idxs"][:] = topk_idxs.view(B, S, SCORE_LEN) + tensors["topk_idxs"][:] = topk_idxs def build_tensor_specs(start_pos=None): @@ -547,12 +554,12 @@ def build_tensor_specs(start_pos=None): state_slot_mapping, ) from golden import ScalarSpec, TensorSpec - from rope_tables import build_deepseek_v4_rope_tables, materialize_half_rope_tables + from rope_tables import build_deepseek_v4_rope_tables shared_freqs_cos, shared_freqs_sin = build_deepseek_v4_rope_tables(M, COMPRESS_RATIO, dtype=torch.bfloat16) def init_x(): - return torch.rand(B, S, D) + return torch.rand(T, D) def init_qr(): return torch.rand(T, Q_LORA) # weights_proj / inner compressor calibrated to the real DeepSeek-V4-Flash CSA indexer @@ -560,12 +567,10 @@ def init_qr(): # near the measured mean. idx wq_b uses the MXFP8 grid below (not a benign randn INT8). def init_weights_proj(): return torch.randn(D, IDX_N_HEADS) * 0.2313 - def init_rope_positions(): - return init_position_ids().to(torch.int64)[:, 0] - def init_cos(): - return materialize_half_rope_tables(shared_freqs_cos, shared_freqs_sin, init_rope_positions())[0] - def init_sin(): - return materialize_half_rope_tables(shared_freqs_cos, shared_freqs_sin, init_rope_positions())[1] + def init_freqs_cos(): + return shared_freqs_cos.clone() + def init_freqs_sin(): + return shared_freqs_sin.clone() def init_hadamard(): return torch.rand(IDX_HEAD_DIM, IDX_HEAD_DIM) * (IDX_HEAD_DIM ** -0.5) def init_inner_compress_state(): @@ -606,23 +611,24 @@ def init_start_pos(): default_fn=init_default_start_pos, ) def init_position_ids(): - return position_ids_from_starts(init_start_pos(), seq=S) + return position_ids_from_starts(init_start_pos(), seq=S).reshape(T).contiguous() def init_kv_seq_lens(): return kv_seq_lens_from_starts(init_start_pos(), seq=S) def init_inner_state_slot_mapping(): + positions = position_ids_from_starts(init_start_pos(), seq=S) return state_slot_mapping( - init_position_ids(), + positions, init_inner_compress_state_block_table(), state_block_size=INNER_STATE_BLOCK_SIZE, - ) + ).reshape(T).contiguous() def init_idx_slot_mapping(): - positions = init_position_ids() + positions = position_ids_from_starts(init_start_pos(), seq=S) return compressed_slot_mapping( positions, init_idx_block_table(), compress_ratio=COMPRESS_RATIO, block_size=BLOCK_SIZE, - ) + ).reshape(T).contiguous() # idx wq_b: simulate the real MXFP8 (e4m3 + 128x128-block E8M0) grid (~200 levels, scaleCV # ~0.61, ~1.1% zero spike) instead of a benign randn INT8. gen_shared_weight reduces over @@ -640,16 +646,16 @@ def init_idx_slot_mapping(): idx_kv_sc = idx_kv_sc.view(IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, 1) return [ - TensorSpec("x", [B, S, D], torch.bfloat16, init_value=init_x), + TensorSpec("x", [T, D], torch.bfloat16, init_value=init_x), TensorSpec("qr", [T, Q_LORA], torch.int8, init_value=lambda: qr_i8), TensorSpec("qr_scale", [T, 1], torch.float32, init_value=lambda: qr_scale), TensorSpec("wq_b", [Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], torch.int8, init_value=lambda: wq_b_i8), TensorSpec("wq_b_scale", [IDX_N_HEADS * IDX_HEAD_DIM], torch.float32, init_value=lambda: wq_b_scale), TensorSpec("weights_proj", [D, IDX_N_HEADS], torch.bfloat16, init_value=init_weights_proj), - TensorSpec("cos", [B, ROPE_HEAD_DIM // 2], torch.float32, init_value=init_cos), - TensorSpec("sin", [B, ROPE_HEAD_DIM // 2], torch.float32, init_value=init_sin), + TensorSpec("freqs_cos", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_cos), + TensorSpec("freqs_sin", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_sin), TensorSpec("hadamard", [IDX_HEAD_DIM, IDX_HEAD_DIM], torch.bfloat16, init_value=init_hadamard), - TensorSpec("inner_kv", [B, S, INNER_HEAD_DIM], torch.float32), + TensorSpec("inner_kv", [T, INNER_HEAD_DIM], torch.float32), TensorSpec("inner_compress_state", [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], torch.float32, init_value=init_inner_compress_state), TensorSpec("inner_compress_state_block_table", [B, INNER_STATE_MAX_BLOCKS], torch.int32, init_value=init_inner_compress_state_block_table), TensorSpec("inner_wkv", [INNER_OUT_DIM, D], torch.bfloat16, init_value=init_inner_wkv), @@ -660,11 +666,11 @@ def init_idx_slot_mapping(): TensorSpec("idx_kv_scale", [IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, 1], torch.float32, init_value=lambda: idx_kv_sc, is_output=True), TensorSpec("idx_block_table", [B, IDX_CACHE_MAX_BLOCKS], torch.int32, init_value=init_idx_block_table), # Outputs are fixed to SCORE_LEN; positions past cache_len are -inf for score and -1 for topk_idxs. - TensorSpec("score", [B, S, SCORE_LEN], torch.float32, is_output=True), - TensorSpec("topk_idxs", [B, S, SCORE_LEN], torch.int32, is_output=True), - TensorSpec("position_ids", [B, S], torch.int32, init_value=init_position_ids), - TensorSpec("idx_slot_mapping", [B, S], torch.int64, init_value=init_idx_slot_mapping), - TensorSpec("inner_state_slot_mapping", [B, S], torch.int64, init_value=init_inner_state_slot_mapping), + TensorSpec("score", [T, SCORE_LEN], torch.float32, is_output=True), + TensorSpec("topk_idxs", [T, SCORE_LEN], torch.int32, is_output=True), + TensorSpec("position_ids", [T], torch.int32, init_value=init_position_ids), + TensorSpec("idx_slot_mapping", [T], torch.int64, init_value=init_idx_slot_mapping), + TensorSpec("inner_state_slot_mapping", [T], torch.int64, init_value=init_inner_state_slot_mapping), TensorSpec("kv_seq_lens", [B], torch.int32, init_value=init_kv_seq_lens), ScalarSpec("offset", torch.int32, OFFSET), ] diff --git a/models/deepseek/v4/decode_indexer_compressor.py b/models/deepseek/v4/decode_indexer_compressor.py index fa31fe56..a74f6317 100644 --- a/models/deepseek/v4/decode_indexer_compressor.py +++ b/models/deepseek/v4/decode_indexer_compressor.py @@ -28,6 +28,7 @@ # model config B = DECODE_BATCH S = DECODE_SEQ +T = B * S EPS = M.rms_norm_eps D = M.hidden_size HEAD_DIM = M.index_head_dim @@ -66,28 +67,26 @@ @pl.jit.inline def indexer_compressor( - x: pl.Tensor[[B, S, D], pl.BF16], - kv: pl.Tensor[[B, S, HEAD_DIM], pl.FP32], + x: pl.Tensor[[T, D], pl.BF16], + kv: pl.Tensor[[T, HEAD_DIM], pl.FP32], compress_state: pl.Tensor[[COMPRESS_STATE_BLOCK_NUM, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32], compress_state_block_table: pl.Tensor[[B, COMPRESS_STATE_MAX_BLOCKS], pl.INT32], wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - cos: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], - sin: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], hadamard: pl.Tensor[[HEAD_DIM, HEAD_DIM], pl.BF16], idx_kv_cache: pl.Tensor[[IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.INT8], idx_kv_scale: pl.Tensor[[IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32], - position_ids: pl.Tensor[[B, S], pl.INT32], - idx_slot_mapping: pl.Tensor[[B, S], pl.INT64], - inner_state_slot_mapping: pl.Tensor[[B, S], pl.INT64], + position_ids: pl.Tensor[[T], pl.INT32], + idx_slot_mapping: pl.Tensor[[T], pl.INT64], + inner_state_slot_mapping: pl.Tensor[[T], pl.INT64], ): - x_flat = pl.reshape(x, [B * S, D]) kv_proj_pad = pl.create_tensor([BS_PAD, OUT_DIM], dtype=pl.FP32) score_proj_pad = pl.create_tensor([BS_PAD, OUT_DIM], dtype=pl.FP32) compress_state_flat = pl.reshape(compress_state, [COMPRESS_STATE_BLOCK_NUM * COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM]) - kv_flat = pl.reshape(kv, [B * S, HEAD_DIM]) idx_kv_cache_flat = pl.reshape(idx_kv_cache, [IDX_CACHE_BLOCK_NUM * BLOCK_SIZE, HEAD_DIM]) idx_kv_scale_flat = pl.reshape(idx_kv_scale, [IDX_CACHE_BLOCK_NUM * BLOCK_SIZE, 1]) @@ -99,11 +98,11 @@ def indexer_compressor( for kb in pl.pipeline(0, D // K_TILE, stage=2): k0 = kb * K_TILE x_rows = pl.min(MM_B_TILE, B * S - global_row0) - x_tile = pl.slice(x_flat, [MM_B_TILE, K_TILE], [global_row0, k0], valid_shape=[x_rows, K_TILE]) + x_tile = pl.slice(x, [MM_B_TILE, K_TILE], [global_row0, k0], valid_shape=[x_rows, K_TILE]) # Weights stored transposed [OUT_DIM, D] and consumed via b_trans=True so the # GM->L1 load is a DN2ZN (each [OUT_TILE, K_TILE] row is K-contiguous = long # bursts) instead of ND2NZ on [K_TILE, OUT_TILE] (K strided = many short - # bursts). Mirrors the main compressor (decode_compressor_ratio4); the strided + # bursts). Mirrors compressor_ratio4 decode mode; the strided # ND2NZ form here was ~2x slower on this matmul (43us -> ~20us per task). wkv_tile = wkv[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] wgate_tile = wgate[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] @@ -123,10 +122,10 @@ def indexer_compressor( pooled_kv = pl.create_tensor([RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) with pl.at(level=pl.Level.CORE_GROUP, name_hint="scatter_softmax_pool"): for c_idx in pl.range(B): - for s_sc in pl.pipeline(S, stage=2): - token_pos = pl.read(position_ids, [c_idx, s_sc]) - state_row_i64 = pl.read(inner_state_slot_mapping, [c_idx, s_sc]) - proj_row = c_idx * S + s_sc + for s in pl.pipeline(S, stage=2): + proj_row = c_idx * S + s + token_pos = pl.read(position_ids, [proj_row]) + state_row_i64 = pl.read(inner_state_slot_mapping, [proj_row]) token_ape_row = pl.cast(token_pos % COMPRESS_RATIO, target_type=pl.INDEX) if state_row_i64 >= 0: state_row = pl.cast(state_row_i64, pl.INDEX) @@ -138,7 +137,7 @@ def indexer_compressor( compress_state_flat[state_row : state_row + 1, OUT_DIM : 2 * OUT_DIM] = score_tile pad_idx = c_idx - first_pos_b = pl.read(position_ids, [c_idx, 0]) + first_pos_b = pl.read(position_ids, [c_idx * S]) pos_b = first_pos_b % COMPRESS_RATIO pre_tokens_b = COMPRESS_RATIO - pos_b boundary_end_b = first_pos_b + pre_tokens_b - 1 @@ -201,8 +200,19 @@ def indexer_compressor( # single 16-row block: B real rows at rows 0..B-1, rows B..15 are pad cos_b = pl.full([RMS_PAD_TILE, ROPE_HEAD_DIM // 2], dtype=pl.FP32, value=0.0) sin_b = pl.full([RMS_PAD_TILE, ROPE_HEAD_DIM // 2], dtype=pl.FP32, value=0.0) - cos_b[0:B, 0 : ROPE_HEAD_DIM // 2] = cos[0:B, 0 : ROPE_HEAD_DIM // 2] - sin_b[0:B, 0 : ROPE_HEAD_DIM // 2] = sin[0:B, 0 : ROPE_HEAD_DIM // 2] + # in-kernel gather of the per-batch compressor rope rows (single 16-row block: + # real rows 0..B-1, pad rows B..15 keep the zero cos/sin from the pl.full above) + for inner in pl.range(B): + first_pos_b = pl.read(position_ids, [inner * S]) + cmp_pos_b = pl.cast(first_pos_b - (first_pos_b % COMPRESS_RATIO), pl.INDEX) + cos_b[inner : inner + 1, 0 : ROPE_HEAD_DIM // 2] = pl.cast( + freqs_cos[cmp_pos_b : cmp_pos_b + 1, 0 : ROPE_HEAD_DIM // 2], + target_type=pl.FP32, + ) + sin_b[inner : inner + 1, 0 : ROPE_HEAD_DIM // 2] = pl.cast( + freqs_sin[cmp_pos_b : cmp_pos_b + 1, 0 : ROPE_HEAD_DIM // 2], + target_type=pl.FP32, + ) partial_sq = pl.full([1, RMS_PAD_TILE], dtype=pl.FP32, value=0.0) for k0 in pl.range(0, HEAD_DIM, HEAD_TILE): kv_rms_chunk = pooled_kv[0 : RMS_PAD_TILE, k0 : k0 + HEAD_TILE] @@ -275,41 +285,40 @@ def indexer_compressor( kv_i8_blk = pl.cast(kv_half, target_type=pl.INT8, mode="trunc") for inner in pl.range(B): c_idx = inner - first_pos_b = pl.read(position_ids, [c_idx, 0]) + first_pos_b = pl.read(position_ids, [c_idx * S]) pos_b = first_pos_b % COMPRESS_RATIO if pos_b + S >= COMPRESS_RATIO: boundary_s = COMPRESS_RATIO - 1 - pos_b kv_row_fp32 = kv_final[inner : inner + 1, 0 : HEAD_DIM] - cache_row_i64 = pl.read(idx_slot_mapping, [c_idx, boundary_s]) + cache_row_i64 = pl.read(idx_slot_mapping, [c_idx * S + boundary_s]) if cache_row_i64 >= 0: cache_row = pl.cast(cache_row_i64, pl.INDEX) - kv_flat[c_idx * S : c_idx * S + 1, :] = kv_row_fp32 + kv[c_idx * S : c_idx * S + 1, :] = kv_row_fp32 idx_kv_cache_flat[cache_row : cache_row + 1, :] = kv_i8_blk[inner : inner + 1, :] # scale is one value per position; a [1,1] tile store is sub-32B, so scalar-write it pl.write(idx_kv_scale_flat, [cache_row, 0], pl.read(kv_scale_dq_col, [inner, 0])) - kv = pl.reshape(kv_flat, [B, S, HEAD_DIM]) return kv @pl.jit def compressor_test( - x: pl.Tensor[[B, S, D], pl.BF16], - kv: pl.Out[pl.Tensor[[B, S, HEAD_DIM], pl.FP32]], + x: pl.Tensor[[T, D], pl.BF16], + kv: pl.Out[pl.Tensor[[T, HEAD_DIM], pl.FP32]], compress_state: pl.InOut[pl.Tensor[[COMPRESS_STATE_BLOCK_NUM, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32]], compress_state_block_table: pl.Tensor[[B, COMPRESS_STATE_MAX_BLOCKS], pl.INT32], wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - cos: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], - sin: pl.Tensor[[B, ROPE_HEAD_DIM // 2], pl.FP32], + freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], + freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], hadamard: pl.Tensor[[HEAD_DIM, HEAD_DIM], pl.BF16], idx_kv_cache: pl.InOut[pl.Tensor[[IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.INT8]], idx_kv_scale: pl.InOut[pl.Tensor[[IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32]], - position_ids: pl.Tensor[[B, S], pl.INT32], - idx_slot_mapping: pl.Tensor[[B, S], pl.INT64], - inner_state_slot_mapping: pl.Tensor[[B, S], pl.INT64], + position_ids: pl.Tensor[[T], pl.INT32], + idx_slot_mapping: pl.Tensor[[T], pl.INT64], + inner_state_slot_mapping: pl.Tensor[[T], pl.INT64], ): indexer_compressor( x, @@ -320,8 +329,8 @@ def compressor_test( wgate, ape, norm_w, - cos, - sin, + freqs_cos, + freqs_sin, hadamard, idx_kv_cache, idx_kv_scale, @@ -336,21 +345,21 @@ def golden_compressor(tensors): """Torch reference for Compressor.forward (decode branch, ratio=4 overlap).""" import torch - x = tensors["x"].float() + x = tensors["x"].float().reshape(B, S, D) compress_state = tensors["compress_state"] compress_state_block_table = tensors["compress_state_block_table"] wkv = tensors["wkv"].float() wgate = tensors["wgate"].float() ape = tensors["ape"] norm_w = tensors["norm_w"] - cos = tensors["cos"] - sin = tensors["sin"] + freqs_cos = tensors["freqs_cos"] + freqs_sin = tensors["freqs_sin"] hadamard = tensors["hadamard"].float() idx_kv_cache = tensors["idx_kv_cache"] idx_kv_scale = tensors["idx_kv_scale"] - position_ids = tensors["position_ids"].to(torch.int64) - idx_slot_mapping = tensors["idx_slot_mapping"].to(torch.int64) - inner_state_slot_mapping = tensors["inner_state_slot_mapping"].to(torch.int64) + position_ids = tensors["position_ids"].to(torch.int64).reshape(B, S) + idx_slot_mapping = tensors["idx_slot_mapping"].to(torch.int64).reshape(B, S) + inner_state_slot_mapping = tensors["inner_state_slot_mapping"].to(torch.int64).reshape(B, S) bsz, _, _ = x.shape ratio, rd = COMPRESS_RATIO, ROPE_HEAD_DIM @@ -449,7 +458,9 @@ def rmsnorm(x, w): x_pair = kv_b[..., -rd:].unflatten(-1, (-1, 2)) x0, x1 = x_pair[..., 0], x_pair[..., 1] - cos_v, sin_v = cos[b].view(-1), sin[b].view(-1) + cmp_pos = first_pos - (first_pos % ratio) + cos_v = freqs_cos[cmp_pos, : rd // 2].float().view(-1) + sin_v = freqs_sin[cmp_pos, : rd // 2].float().view(-1) y0 = x0 * cos_v - x1 * sin_v y1 = x0 * sin_v + x1 * cos_v @@ -461,7 +472,7 @@ def rmsnorm(x, w): if cache_row >= 0: # Kernel writes committed pooled result only to kv[:, 0, :]; leave # speculative-boundary rows and kv[:, 1:, :] zero-initialized. - tensors["kv"][b : b + 1, 0:1, :] = kv_b + tensors["kv"][b * S : b * S + 1, :] = kv_b.reshape(1, HEAD_DIM) blk_id = cache_row // BLOCK_SIZE intra = cache_row % BLOCK_SIZE # C8 quant-on-write: quantize the bf16-rounded compressed row to int8 + per-position scale @@ -486,12 +497,12 @@ def build_tensor_specs(start_pos=None): state_slot_mapping, ) from golden import TensorSpec - from rope_tables import build_deepseek_v4_rope_tables, materialize_half_rope_tables + from rope_tables import build_deepseek_v4_rope_tables shared_freqs_cos, shared_freqs_sin = build_deepseek_v4_rope_tables(M, COMPRESS_RATIO, dtype=torch.bfloat16) def init_x(): - return torch.rand(B, S, D) + return torch.rand(T, D) def init_compress_state(): state = torch.zeros(COMPRESS_STATE_BLOCK_NUM, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM) state[:, :, OUT_DIM:] = FP32_NEG_INF @@ -513,14 +524,10 @@ def init_ape(): return torch.randn(COMPRESS_RATIO, OUT_DIM) * 0.1528 def init_norm_w(): return 0.6850 + 0.2610 * torch.randn(HEAD_DIM) - def init_rope_positions(): - first_pos = init_position_ids().to(torch.int64)[:, 0] - cmp_offset = COMPRESS_RATIO - (first_pos % COMPRESS_RATIO) - return (first_pos + cmp_offset - COMPRESS_RATIO).to(torch.int64) - def init_cos(): - return materialize_half_rope_tables(shared_freqs_cos, shared_freqs_sin, init_rope_positions())[0] - def init_sin(): - return materialize_half_rope_tables(shared_freqs_cos, shared_freqs_sin, init_rope_positions())[1] + def init_freqs_cos(): + return shared_freqs_cos.clone() + def init_freqs_sin(): + return shared_freqs_sin.clone() def init_hadamard(): return torch.rand(HEAD_DIM, HEAD_DIM) * (HEAD_DIM ** -0.5) def init_idx_kv_cache(): @@ -546,40 +553,42 @@ def init_start_pos(): max_seq_len=MAX_SEQ_LEN, default_fn=init_default_start_pos, ) - def init_position_ids(): + def init_position_ids_2d(): return position_ids_from_starts(init_start_pos(), seq=S) + def init_position_ids(): + return init_position_ids_2d().reshape(T).contiguous() def init_inner_state_slot_mapping(): return state_slot_mapping( - init_position_ids(), + init_position_ids_2d(), init_compress_state_block_table(), state_block_size=COMPRESS_STATE_BLOCK_SIZE, - ) + ).reshape(T).contiguous() def init_idx_slot_mapping(): - positions = init_position_ids() + positions = init_position_ids_2d() return compressed_slot_mapping( positions, init_idx_block_table(), compress_ratio=COMPRESS_RATIO, block_size=BLOCK_SIZE, - ) + ).reshape(T).contiguous() return [ - TensorSpec("x", [B, S, D], torch.bfloat16, init_value=init_x), - TensorSpec("kv", [B, S, HEAD_DIM], torch.float32, is_output=True), + TensorSpec("x", [T, D], torch.bfloat16, init_value=init_x), + TensorSpec("kv", [T, HEAD_DIM], torch.float32, is_output=True), TensorSpec("compress_state", [COMPRESS_STATE_BLOCK_NUM, COMPRESS_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], torch.float32, init_value=init_compress_state, is_output=True), TensorSpec("compress_state_block_table", [B, COMPRESS_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), TensorSpec("wkv", [OUT_DIM, D], torch.bfloat16, init_value=init_wkv), TensorSpec("wgate", [OUT_DIM, D], torch.bfloat16, init_value=init_wgate), TensorSpec("ape", [COMPRESS_RATIO, OUT_DIM], torch.float32, init_value=init_ape), TensorSpec("norm_w", [HEAD_DIM], torch.bfloat16, init_value=init_norm_w), - TensorSpec("cos", [B, ROPE_HEAD_DIM // 2], torch.float32, init_value=init_cos), - TensorSpec("sin", [B, ROPE_HEAD_DIM // 2], torch.float32, init_value=init_sin), + TensorSpec("freqs_cos", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_cos), + TensorSpec("freqs_sin", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_sin), TensorSpec("hadamard", [HEAD_DIM, HEAD_DIM], torch.bfloat16, init_value=init_hadamard), TensorSpec("idx_kv_cache", [IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], torch.int8, init_value=init_idx_kv_cache, is_output=True), TensorSpec("idx_kv_scale", [IDX_CACHE_BLOCK_NUM, BLOCK_SIZE, 1, 1], torch.float32, init_value=init_idx_kv_scale, is_output=True), - TensorSpec("position_ids", [B, S], torch.int32, init_value=init_position_ids), - TensorSpec("idx_slot_mapping", [B, S], torch.int64, init_value=init_idx_slot_mapping), - TensorSpec("inner_state_slot_mapping", [B, S], torch.int64, init_value=init_inner_state_slot_mapping), + TensorSpec("position_ids", [T], torch.int32, init_value=init_position_ids), + TensorSpec("idx_slot_mapping", [T], torch.int64, init_value=init_idx_slot_mapping), + TensorSpec("inner_state_slot_mapping", [T], torch.int64, init_value=init_inner_state_slot_mapping), ] diff --git a/models/deepseek/v4/prefill_attention_csa.py b/models/deepseek/v4/prefill_attention_csa.py index 7015c870..37703627 100644 --- a/models/deepseek/v4/prefill_attention_csa.py +++ b/models/deepseek/v4/prefill_attention_csa.py @@ -31,7 +31,7 @@ PREFILL_SEQ, ) -from prefill_compressor_ratio4 import ( +from compressor_ratio4 import ( CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_STATE_MAX_BLOCKS, @@ -49,6 +49,7 @@ prefill_indexer, ) from prefill_indexer_compressor import ( + COMPRESS_STATE_DIM as INNER_STATE_DIM, INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_MAX_BLOCKS, @@ -68,7 +69,6 @@ H = M.num_attention_heads HEAD_DIM = M.head_dim ROPE_HEAD_DIM = M.qk_rope_head_dim -HALF_ROPE = ROPE_HEAD_DIM // 2 Q_LORA = M.q_lora_rank MAX_SEQ_LEN = M.max_position_embeddings WIN = M.sliding_window @@ -89,6 +89,7 @@ COFF = 2 MAIN_OUT_DIM = COFF * HEAD_DIM MAIN_STATE_LEN = COFF * COMPRESS_RATIO +MAIN_STATE_DIM = 2 * MAIN_OUT_DIM INNER_OUT_DIM = COFF * IDX_HEAD_DIM INNER_STATE_LEN = COFF * COMPRESS_RATIO ORI_MAX_BLOCKS = PREFILL_ORI_MAX_BLOCKS @@ -137,9 +138,8 @@ def prefill_attention_csa( cmp_wgate: pl.Tensor[[MAIN_OUT_DIM, D], pl.BF16], cmp_ape: pl.Tensor[[COMPRESS_RATIO, MAIN_OUT_DIM], pl.FP32], cmp_norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - cmp_kv_state: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], pl.FP32], - cmp_score_state: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], pl.FP32], - compress_state_block_table: pl.Tensor[[CSA_STATE_MAX_BLOCKS], pl.INT32], + compress_state: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[B, CSA_STATE_MAX_BLOCKS], pl.INT32], hadamard_idx: pl.Tensor[[IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16], idx_wq_b: pl.Tensor[[Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8], idx_wq_b_scale: pl.Tensor[[IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32], @@ -148,9 +148,8 @@ def prefill_attention_csa( inner_wgate: pl.Tensor[[INNER_OUT_DIM, D], pl.BF16], inner_ape: pl.Tensor[[COMPRESS_RATIO, INNER_OUT_DIM], pl.FP32], inner_norm_w: pl.Tensor[[IDX_HEAD_DIM], pl.BF16], - inner_kv_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], - inner_score_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], - inner_compress_state_block_table: pl.Tensor[[INNER_STATE_MAX_BLOCKS], pl.INT32], + inner_compress_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32], + inner_compress_state_block_table: pl.Tensor[[B, INNER_STATE_MAX_BLOCKS], pl.INT32], kv_cache: pl.InOut[pl.Tensor[[CSA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], ori_block_table: pl.Tensor[[SPARSE_ORI_MAX_BLOCKS], pl.INT32], ori_slot_mapping: pl.Tensor[[T], pl.INT64], @@ -158,7 +157,7 @@ def prefill_attention_csa( cmp_block_table: pl.Tensor[[SPARSE_CMP_MAX_BLOCKS], pl.INT32], idx_kv_cache: pl.Out[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8]], idx_kv_scale: pl.Out[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32]], - idx_block_table: pl.Tensor[[IDX_CACHE_MAX_BLOCKS], pl.INT32], + idx_block_table: pl.Tensor[[B, IDX_CACHE_MAX_BLOCKS], pl.INT32], position_ids: pl.Tensor[[T], pl.INT32], cmp_slot_mapping: pl.Tensor[[T], pl.INT64], idx_slot_mapping: pl.Tensor[[T], pl.INT64], @@ -209,28 +208,18 @@ def prefill_attention_csa( kv_cache_flat[write_row : write_row + 1, :] = kv[write_t : write_t + 1, :] prefill_compressor_ratio4( - x_normed, cmp_kv_state, cmp_score_state, compress_state_block_table, + x_normed, compress_state, compress_state_block_table, cmp_wkv, cmp_wgate, cmp_ape, cmp_norm_w, freqs_cos, freqs_sin, cmp_kv, position_ids, num_tokens, cmp_slot_mapping, state_slot_mapping, ) - # Half-width FP32 cos/sin rows for the indexer Q-RoPE: gather freqs at each token's position - # and take the first HALF_ROPE columns (matches the golden's materialize_half_rope_tables). - idx_cos = pl.create_tensor([T, HALF_ROPE], dtype=pl.FP32) - idx_sin = pl.create_tensor([T, HALF_ROPE], dtype=pl.FP32) - with pl.at(level=pl.Level.CORE_GROUP, name_hint="prefill_csa_idx_halfrope"): - for idx_t in pl.range(T): - idx_pos = pl.cast(pl.read(position_ids, [idx_t]), pl.INDEX) - idx_cos = pl.assemble(idx_cos, pl.cast(pl.slice(freqs_cos, [1, HALF_ROPE], [idx_pos, 0]), target_type=pl.FP32), [idx_t, 0]) - idx_sin = pl.assemble(idx_sin, pl.cast(pl.slice(freqs_sin, [1, HALF_ROPE], [idx_pos, 0]), target_type=pl.FP32), [idx_t, 0]) - cmp_topk_indices = pl.create_tensor([T, IDX_TOPK], dtype=pl.INT32) idx_score_unused = pl.create_tensor([T, INDEXER_SCORE_CAP], dtype=pl.FP32) prefill_indexer( x_normed, qr, qr_scale, idx_wq_b, idx_wq_b_scale, idx_weights_proj, - idx_cos, idx_sin, freqs_cos, freqs_sin, hadamard_idx, - inner_kv_state, inner_score_state, inner_compress_state_block_table, + freqs_cos, freqs_sin, hadamard_idx, + inner_compress_state, inner_compress_state_block_table, inner_wkv, inner_wgate, inner_ape, inner_norm_w, idx_kv_cache, idx_kv_scale, idx_block_table, idx_score_unused, cmp_topk_indices, @@ -282,7 +271,7 @@ def prefill_attention_csa( ) hc_post(attn_out, x_hc, post, comb, x_out) - return kv_cache, cmp_kv, cmp_kv_state, cmp_score_state, idx_kv_cache, idx_kv_scale, inner_kv_state, inner_score_state, x_out + return kv_cache, cmp_kv, compress_state, idx_kv_cache, idx_kv_scale, inner_compress_state, x_out @pl.jit @@ -304,9 +293,8 @@ def prefill_attention_csa_test( cmp_wgate: pl.Tensor[[MAIN_OUT_DIM, D], pl.BF16], cmp_ape: pl.Tensor[[COMPRESS_RATIO, MAIN_OUT_DIM], pl.FP32], cmp_norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - cmp_kv_state: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], pl.FP32], - cmp_score_state: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], pl.FP32], - compress_state_block_table: pl.Tensor[[CSA_STATE_MAX_BLOCKS], pl.INT32], + compress_state: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[B, CSA_STATE_MAX_BLOCKS], pl.INT32], hadamard_idx: pl.Tensor[[IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16], idx_wq_b: pl.Tensor[[Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8], idx_wq_b_scale: pl.Tensor[[IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32], @@ -315,9 +303,8 @@ def prefill_attention_csa_test( inner_wgate: pl.Tensor[[INNER_OUT_DIM, D], pl.BF16], inner_ape: pl.Tensor[[COMPRESS_RATIO, INNER_OUT_DIM], pl.FP32], inner_norm_w: pl.Tensor[[IDX_HEAD_DIM], pl.BF16], - inner_kv_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], - inner_score_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], - inner_compress_state_block_table: pl.Tensor[[INNER_STATE_MAX_BLOCKS], pl.INT32], + inner_compress_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32], + inner_compress_state_block_table: pl.Tensor[[B, INNER_STATE_MAX_BLOCKS], pl.INT32], kv_cache: pl.InOut[pl.Tensor[[CSA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], ori_block_table: pl.Tensor[[SPARSE_ORI_MAX_BLOCKS], pl.INT32], ori_slot_mapping: pl.Tensor[[T], pl.INT64], @@ -325,7 +312,7 @@ def prefill_attention_csa_test( cmp_block_table: pl.Tensor[[SPARSE_CMP_MAX_BLOCKS], pl.INT32], idx_kv_cache: pl.InOut[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8]], idx_kv_scale: pl.InOut[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32]], - idx_block_table: pl.Tensor[[IDX_CACHE_MAX_BLOCKS], pl.INT32], + idx_block_table: pl.Tensor[[B, IDX_CACHE_MAX_BLOCKS], pl.INT32], position_ids: pl.Tensor[[T], pl.INT32], cmp_slot_mapping: pl.Tensor[[T], pl.INT64], idx_slot_mapping: pl.Tensor[[T], pl.INT64], @@ -344,10 +331,10 @@ def prefill_attention_csa_test( attn_norm_w, wq_a, wq_b, wq_b_scale, wkv, gamma_cq, gamma_ckv, freqs_cos, freqs_sin, cmp_wkv, cmp_wgate, cmp_ape, cmp_norm_w, - cmp_kv_state, cmp_score_state, compress_state_block_table, + compress_state, compress_state_block_table, hadamard_idx, idx_wq_b, idx_wq_b_scale, idx_weights_proj, inner_wkv, inner_wgate, inner_ape, inner_norm_w, - inner_kv_state, inner_score_state, inner_compress_state_block_table, + inner_compress_state, inner_compress_state_block_table, kv_cache, ori_block_table, ori_slot_mapping, cmp_kv, cmp_block_table, idx_kv_cache, idx_kv_scale, idx_block_table, position_ids, cmp_slot_mapping, idx_slot_mapping, @@ -355,7 +342,7 @@ def prefill_attention_csa_test( attn_sink, wo_a, wo_b, wo_b_scale, x_out, num_tokens, ) - return kv_cache, cmp_kv, cmp_kv_state, cmp_score_state, idx_kv_cache, idx_kv_scale, inner_kv_state, inner_score_state, x_out + return kv_cache, cmp_kv, compress_state, idx_kv_cache, idx_kv_scale, inner_compress_state, x_out def golden_prefill_attention_csa(tensors): @@ -406,8 +393,7 @@ def golden_prefill_attention_csa(tensors): golden_prefill_compressor_ratio4({ "x": x_normed.view(T, D), - "kv_state": tensors["cmp_kv_state"], - "score_state": tensors["cmp_score_state"], + "compress_state": tensors["compress_state"], "compress_state_block_table": tensors["compress_state_block_table"], "wkv": tensors["cmp_wkv"], "wgate": tensors["cmp_wgate"], @@ -421,8 +407,6 @@ def golden_prefill_attention_csa(tensors): "cmp_slot_mapping": tensors["cmp_slot_mapping"], "state_slot_mapping": tensors["state_slot_mapping"], }) - idx_cos = rope_cos_t[:, :HALF_ROPE].float().contiguous() - idx_sin = rope_sin_t[:, :HALF_ROPE].float().contiguous() cmp_topk_indices, _idx_score = golden_prefill_indexer_core({ "x": x_normed.view(T, D), "qr": qr, @@ -430,13 +414,10 @@ def golden_prefill_attention_csa(tensors): "wq_b": tensors["idx_wq_b"], "wq_b_scale": tensors["idx_wq_b_scale"], "weights_proj": tensors["idx_weights_proj"], - "cos": idx_cos, - "sin": idx_sin, "freqs_cos": tensors["freqs_cos"], "freqs_sin": tensors["freqs_sin"], "hadamard": tensors["hadamard_idx"], - "inner_kv_state": tensors["inner_kv_state"], - "inner_score_state": tensors["inner_score_state"], + "inner_compress_state": tensors["inner_compress_state"], "inner_compress_state_block_table": tensors["inner_compress_state_block_table"], "inner_wkv": tensors["inner_wkv"], "inner_wgate": tensors["inner_wgate"], @@ -618,7 +599,7 @@ def init_freqs_sin(): return shared_freqs_sin.clone() # Quant-faithful CSA (ratio-4) main compressor fixtures (mean l8/l32 of extract_weights_flash): # zero-mean Gaussian BF16 weights at the measured std; RMSNorm gamma near the measured mean. - # Mirrors decode_attention_csa / decode_compressor_ratio4. + # Mirrors decode_attention_csa / compressor_ratio4 decode mode. def init_cmp_wkv(): return torch.randn(MAIN_OUT_DIM, D) * 0.0245 def init_cmp_wgate(): @@ -629,28 +610,21 @@ def init_cmp_norm_w(): return 0.9666 + torch.randn(HEAD_DIM,) * 0.1929 state_table = _state_block_table(CSA_STATE_MAX_BLOCKS) def init_compress_state_block_table(): - return state_table.clone() + return state_table.unsqueeze(0).clone() def state_row(abs_pos): if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: return -1 block = abs_pos // CSA_STATE_BLOCK_SIZE intra = abs_pos % CSA_STATE_BLOCK_SIZE return int(state_table[block].item()) * CSA_STATE_BLOCK_SIZE + intra - def init_cmp_state(): - state = torch.zeros(CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_OUT_DIM) - flat = state.view(-1, MAIN_OUT_DIM) - for abs_pos in range(max(0, context_len - MAIN_STATE_LEN), context_len): - row = state_row(abs_pos) - if row >= 0: - flat[row] = (torch.rand(MAIN_OUT_DIM,) - 0.5) * 0.05 - return state - def init_cmp_score_state(): - state = torch.zeros(CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_OUT_DIM) - flat = state.view(-1, MAIN_OUT_DIM) + def init_compress_state(): + state = torch.zeros(CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_STATE_DIM) + flat = state.view(-1, MAIN_STATE_DIM) for abs_pos in range(max(0, context_len - MAIN_STATE_LEN), context_len): row = state_row(abs_pos) if row >= 0: - flat[row] = (torch.rand(MAIN_OUT_DIM,) - 0.5) * 0.05 + flat[row, 0:MAIN_OUT_DIM] = (torch.rand(MAIN_OUT_DIM,) - 0.5) * 0.05 + flat[row, MAIN_OUT_DIM:MAIN_STATE_DIM] = (torch.rand(MAIN_OUT_DIM,) - 0.5) * 0.05 return state def init_hadamard_idx(): h = torch.ones((1, 1)) @@ -670,28 +644,21 @@ def init_inner_norm_w(): return 0.6850 + torch.randn(IDX_HEAD_DIM,) * 0.2610 inner_state_table = _state_block_table(INNER_STATE_MAX_BLOCKS) def init_inner_compress_state_block_table(): - return inner_state_table.clone() + return inner_state_table.unsqueeze(0).clone() def inner_state_row(abs_pos): if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: return -1 block = abs_pos // INNER_STATE_BLOCK_SIZE intra = abs_pos % INNER_STATE_BLOCK_SIZE return int(inner_state_table[block].item()) * INNER_STATE_BLOCK_SIZE + intra - def init_inner_kv_state(): - state = torch.zeros(INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM) - flat = state.view(-1, INNER_OUT_DIM) - for abs_pos in range(max(0, context_len - INNER_STATE_LEN), context_len): - row = inner_state_row(abs_pos) - if row >= 0: - flat[row] = (torch.rand(INNER_OUT_DIM,) - 0.5) * 0.05 - return state - def init_inner_score_state(): - state = torch.zeros(INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM) - flat = state.view(-1, INNER_OUT_DIM) + def init_inner_compress_state(): + state = torch.zeros(INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM) + flat = state.view(-1, INNER_STATE_DIM) for abs_pos in range(max(0, context_len - INNER_STATE_LEN), context_len): row = inner_state_row(abs_pos) if row >= 0: - flat[row] = (torch.rand(INNER_OUT_DIM,) - 0.5) * 0.05 + flat[row, 0:INNER_OUT_DIM] = (torch.rand(INNER_OUT_DIM,) - 0.5) * 0.05 + flat[row, INNER_OUT_DIM:INNER_STATE_DIM] = (torch.rand(INNER_OUT_DIM,) - 0.5) * 0.05 return state # C8 historical index cache: completed compressed slots hold INT8 + a per-position dequant scale. # Build both from one bf16-rounded random draw so cache and scale stay consistent. @@ -764,14 +731,14 @@ def init_cmp_block_table(): table[block] = block return table def init_idx_block_table(): - table = torch.full((IDX_CACHE_MAX_BLOCKS,), -1, dtype=torch.int32) + table = torch.full((B, IDX_CACHE_MAX_BLOCKS), -1, dtype=torch.int32) for block in range(IDX_CACHE_MAX_BLOCKS): - table[block] = block + table[0, block] = block return table def cache_row_from_table(table, slot): block = slot // BLOCK_SIZE intra = slot % BLOCK_SIZE - phys_block = int(table[block].item()) + phys_block = int(table.reshape(-1)[block].item()) if phys_block < 0: return -1 return phys_block * BLOCK_SIZE + intra @@ -837,18 +804,12 @@ def init_wo_b(): TensorSpec("cmp_ape", [COMPRESS_RATIO, MAIN_OUT_DIM], torch.float32, init_value=init_cmp_ape), TensorSpec("cmp_norm_w", [HEAD_DIM], torch.bfloat16, init_value=init_cmp_norm_w), TensorSpec( - "cmp_kv_state", - [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], - torch.float32, - init_value=init_cmp_state, - ), - TensorSpec( - "cmp_score_state", - [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], + "compress_state", + [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, MAIN_STATE_DIM], torch.float32, - init_value=init_cmp_score_state, + init_value=init_compress_state, ), - TensorSpec("compress_state_block_table", [CSA_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), + TensorSpec("compress_state_block_table", [B, CSA_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), TensorSpec("hadamard_idx", [IDX_HEAD_DIM, IDX_HEAD_DIM], torch.bfloat16, init_value=init_hadamard_idx), TensorSpec("idx_wq_b", [Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], torch.int8, init_value=lambda: idx_wq_b_i8), TensorSpec("idx_wq_b_scale", [IDX_N_HEADS * IDX_HEAD_DIM], torch.float32, init_value=lambda: idx_wq_b_scale), @@ -858,24 +819,18 @@ def init_wo_b(): TensorSpec("inner_ape", [COMPRESS_RATIO, INNER_OUT_DIM], torch.float32, init_value=init_inner_ape), TensorSpec("inner_norm_w", [IDX_HEAD_DIM], torch.bfloat16, init_value=init_inner_norm_w), TensorSpec( - "inner_kv_state", - [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], - torch.float32, - init_value=init_inner_kv_state, - ), - TensorSpec( - "inner_score_state", - [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], + "inner_compress_state", + [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], torch.float32, - init_value=init_inner_score_state, + init_value=init_inner_compress_state, ), - TensorSpec("inner_compress_state_block_table", [INNER_STATE_MAX_BLOCKS], torch.int32, init_value=init_inner_compress_state_block_table), + TensorSpec("inner_compress_state_block_table", [B, INNER_STATE_MAX_BLOCKS], torch.int32, init_value=init_inner_compress_state_block_table), TensorSpec("kv_cache", [CSA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], torch.bfloat16, init_value=init_kv_cache, is_output=True), TensorSpec("ori_block_table", [SPARSE_ORI_MAX_BLOCKS], torch.int32, init_value=init_ori_block_table), TensorSpec("ori_slot_mapping", [T], torch.int64, init_value=init_ori_slot_mapping), # Compressor / indexer caches are written in-place but not validated here - # (decode parity); the dedicated prefill_compressor_ratio4 and + # (decode parity); the dedicated compressor_ratio4 prefill mode and # prefill_indexer tests cover them. TensorSpec( "cmp_kv", @@ -896,7 +851,7 @@ def init_wo_b(): torch.float32, init_value=init_idx_kv_scale, ), - TensorSpec("idx_block_table", [IDX_CACHE_MAX_BLOCKS], torch.int32, init_value=init_idx_block_table), + TensorSpec("idx_block_table", [B, IDX_CACHE_MAX_BLOCKS], torch.int32, init_value=init_idx_block_table), TensorSpec("position_ids", [T], torch.int32, init_value=init_position_ids), TensorSpec("cmp_slot_mapping", [T], torch.int64, init_value=init_cmp_slot_mapping), TensorSpec("idx_slot_mapping", [T], torch.int64, init_value=init_idx_slot_mapping), diff --git a/models/deepseek/v4/prefill_attention_hca.py b/models/deepseek/v4/prefill_attention_hca.py index 6ed72ba7..8b65125c 100644 --- a/models/deepseek/v4/prefill_attention_hca.py +++ b/models/deepseek/v4/prefill_attention_hca.py @@ -32,7 +32,7 @@ ) from hc_post import golden_hc_post, hc_post from hc_pre import golden_hc_pre, hc_pre -from prefill_compressor_ratio128 import ( +from compressor_ratio128 import ( HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_STATE_MAX_BLOCKS, @@ -83,6 +83,7 @@ COMPRESS_RATIO = 128 MAIN_OUT_DIM = HEAD_DIM MAIN_STATE_LEN = COMPRESS_RATIO +MAIN_STATE_DIM = 2 * MAIN_OUT_DIM PREFILL_COMPRESSED_LEN = S // COMPRESS_RATIO START_POS = 0 HCA_ORI_BLOCK_NUM = SPARSE_ORI_MAX_BLOCKS @@ -116,9 +117,8 @@ def prefill_attention_hca( cmp_wgate: pl.Tensor[[MAIN_OUT_DIM, D], pl.BF16], cmp_ape: pl.Tensor[[COMPRESS_RATIO, MAIN_OUT_DIM], pl.FP32], cmp_norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - cmp_kv_state: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], pl.FP32], - cmp_score_state: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], pl.FP32], - compress_state_block_table: pl.Tensor[[HCA_STATE_MAX_BLOCKS], pl.INT32], + compress_state: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[B, HCA_STATE_MAX_BLOCKS], pl.INT32], kv_cache: pl.InOut[pl.Tensor[[HCA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], ori_slot_mapping: pl.Tensor[[T], pl.INT64], ori_block_table: pl.Tensor[[SPARSE_ORI_MAX_BLOCKS], pl.INT32], @@ -173,7 +173,7 @@ def prefill_attention_hca( kv_cache_flat[write_row : write_row + 1, :] = kv[write_t : write_t + 1, :] prefill_compressor_ratio128( - x_normed, cmp_kv_state, cmp_score_state, compress_state_block_table, + x_normed, compress_state, compress_state_block_table, cmp_wkv, cmp_wgate, cmp_ape, cmp_norm_w, freqs_cos, freqs_sin, cmp_kv, position_ids, num_tokens, cmp_slot_mapping, state_slot_mapping, @@ -240,9 +240,8 @@ def prefill_attention_hca_test( cmp_wgate: pl.Tensor[[MAIN_OUT_DIM, D], pl.BF16], cmp_ape: pl.Tensor[[COMPRESS_RATIO, MAIN_OUT_DIM], pl.FP32], cmp_norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - cmp_kv_state: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], pl.FP32], - cmp_score_state: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], pl.FP32], - compress_state_block_table: pl.Tensor[[HCA_STATE_MAX_BLOCKS], pl.INT32], + compress_state: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_STATE_DIM], pl.FP32], + compress_state_block_table: pl.Tensor[[B, HCA_STATE_MAX_BLOCKS], pl.INT32], kv_cache: pl.InOut[pl.Tensor[[HCA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], ori_slot_mapping: pl.Tensor[[T], pl.INT64], ori_block_table: pl.Tensor[[SPARSE_ORI_MAX_BLOCKS], pl.INT32], @@ -264,7 +263,7 @@ def prefill_attention_hca_test( attn_norm_w, wq_a, wq_b, wq_b_scale, wkv, gamma_cq, gamma_ckv, freqs_cos, freqs_sin, cmp_wkv, cmp_wgate, cmp_ape, cmp_norm_w, - cmp_kv_state, cmp_score_state, compress_state_block_table, + compress_state, compress_state_block_table, kv_cache, ori_slot_mapping, ori_block_table, cmp_kv, cmp_block_table, position_ids, cmp_slot_mapping, state_slot_mapping, @@ -341,8 +340,7 @@ def golden_prefill_attention_hca(tensors): cmp_kv = tensors["cmp_kv"] golden_prefill_compressor_ratio128({ "x": x_normed.view(T, D), - "kv_state": tensors["cmp_kv_state"], - "score_state": tensors["cmp_score_state"], + "compress_state": tensors["compress_state"], "compress_state_block_table": tensors["compress_state_block_table"], "wkv": tensors["cmp_wkv"], "wgate": tensors["cmp_wgate"], @@ -504,7 +502,7 @@ def init_freqs_sin(): return shared_freqs_sin.clone() # Quant-faithful HCA (ratio-128) main compressor fixtures (mean l7/l9 of extract_weights_flash): # zero-mean Gaussian BF16 weights at the measured std; RMSNorm gamma near the measured mean. - # Mirrors decode_attention_hca / decode_compressor_ratio128. + # Mirrors decode_attention_hca / compressor_ratio128 decode mode. def init_cmp_wkv(): return torch.randn(MAIN_OUT_DIM, D) * 0.0246 def init_cmp_wgate(): @@ -515,28 +513,21 @@ def init_cmp_norm_w(): return 0.1001 + torch.randn(HEAD_DIM,) * 0.0549 state_table = _state_block_table(HCA_STATE_MAX_BLOCKS) def init_compress_state_block_table(): - return state_table.clone() + return state_table.unsqueeze(0).clone() def state_row(abs_pos): if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: return -1 block = abs_pos // HCA_STATE_BLOCK_SIZE intra = abs_pos % HCA_STATE_BLOCK_SIZE return int(state_table[block].item()) * HCA_STATE_BLOCK_SIZE + intra - def init_cmp_state(): - state = torch.zeros(HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_OUT_DIM) - flat = state.view(-1, MAIN_OUT_DIM) + def init_compress_state(): + state = torch.zeros(HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_STATE_DIM) + flat = state.view(-1, MAIN_STATE_DIM) for abs_pos in range(max(0, context_len - COMPRESS_RATIO), context_len): row = state_row(abs_pos) if row >= 0: - flat[row] = (torch.rand(MAIN_OUT_DIM,) - 0.5) * 0.05 - return state - def init_cmp_score_state(): - state = torch.zeros(HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_OUT_DIM) - flat = state.view(-1, MAIN_OUT_DIM) - for abs_pos in range(max(0, context_len - COMPRESS_RATIO), context_len): - row = state_row(abs_pos) - if row >= 0: - flat[row] = (torch.rand(MAIN_OUT_DIM,) - 0.5) * 0.05 + flat[row, 0:MAIN_OUT_DIM] = (torch.rand(MAIN_OUT_DIM,) - 0.5) * 0.05 + flat[row, MAIN_OUT_DIM:MAIN_STATE_DIM] = (torch.rand(MAIN_OUT_DIM,) - 0.5) * 0.05 return state def cache_row_from_table(table, slot): block = slot // BLOCK_SIZE @@ -633,20 +624,14 @@ def init_wo_b(): TensorSpec("cmp_ape", [COMPRESS_RATIO, MAIN_OUT_DIM], torch.float32, init_value=init_cmp_ape), TensorSpec("cmp_norm_w", [HEAD_DIM], torch.bfloat16, init_value=init_cmp_norm_w), # Compressor caches are written in-place but not validated here (decode - # parity); the dedicated prefill_compressor_ratio128 test covers them. - TensorSpec( - "cmp_kv_state", - [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], - torch.float32, - init_value=init_cmp_state, - ), + # parity); the dedicated compressor_ratio128 prefill mode covers them. TensorSpec( - "cmp_score_state", - [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_OUT_DIM], + "compress_state", + [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, MAIN_STATE_DIM], torch.float32, - init_value=init_cmp_score_state, + init_value=init_compress_state, ), - TensorSpec("compress_state_block_table", [HCA_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), + TensorSpec("compress_state_block_table", [B, HCA_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), TensorSpec( "kv_cache", [HCA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], diff --git a/models/deepseek/v4/prefill_compressor_ratio128.py b/models/deepseek/v4/prefill_compressor_ratio128.py deleted file mode 100644 index 93947f8d..00000000 --- a/models/deepseek/v4/prefill_compressor_ratio128.py +++ /dev/null @@ -1,486 +0,0 @@ -# Copyright (c) PyPTO Contributors. -# This program is free software, you can redistribute it and/or modify it under the terms and conditions of -# CANN Open Software License Agreement Version 2.0 (the "License"). -# Please refer to the License for details. You may not use this file except in compliance with the License. -# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, -# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. -# See LICENSE in the root of the software repository for the full text of the License. -# ----------------------------------------------------------------------------------------------------------- -"""DeepSeek-V4 single-request token-major prefill compressor, ratio=128. - -The public contract is single-request token-major prefill: the -layer owns the per-request loop and feeds this op one contiguous run of <=T -tokens. -""" - -import pypto.language as pl - -from config import BLOCK_SIZE, FLASH as M, PREFILL_BATCH, PREFILL_SEQ, PREFILL_CMP_BLOCK_NUM, PREFILL_CMP_MAX_BLOCKS - - -B = PREFILL_BATCH -S = PREFILL_SEQ -T = B * S -EPS = M.rms_norm_eps -D = M.hidden_size -HEAD_DIM = M.head_dim -HEAD_DIM_INV = 1.0 / HEAD_DIM -ROPE_HEAD_DIM = M.qk_rope_head_dim -ROPE_HALF = ROPE_HEAD_DIM // 2 -NOPE_HEAD_DIM = HEAD_DIM - ROPE_HEAD_DIM -MAX_SEQ_LEN = M.max_position_embeddings - -COMPRESS_RATIO = 128 -OUT_DIM = HEAD_DIM -STATE_LEN = COMPRESS_RATIO -START_POS = 0 - -K_TILE = 512 -OUT_TILE = 32 -HEAD_TILE = 64 -K_BLOCKS = D // K_TILE -OUT_BLOCKS = OUT_DIM // OUT_TILE -HEAD_BLOCKS = HEAD_DIM // HEAD_TILE - -assert S == COMPRESS_RATIO, "ratio128 prefill compressor bring-up expects one full compression chunk" - - -HCA_STATE_BLOCK_SIZE = 8 -HCA_STATE_MAX_BLOCKS = (MAX_SEQ_LEN + HCA_STATE_BLOCK_SIZE - 1) // HCA_STATE_BLOCK_SIZE -HCA_STATE_BLOCK_NUM = HCA_STATE_MAX_BLOCKS -MAX_CMP_WRITES = max(1, T // COMPRESS_RATIO) -HCA_CMP_MAX_BLOCKS = PREFILL_CMP_MAX_BLOCKS -HCA_CMP_BLOCK_NUM = PREFILL_CMP_BLOCK_NUM -HCA_KV_STORE_TILE = 16 -HCA_C128_RMS_TILE = 8 -HCA_C128_RMS_PAD_ROWS = HCA_C128_RMS_TILE - -PACKED_C128_PROJ_BLOCKS = OUT_BLOCKS -PACKED_C128_POOL_BLOCKS = MAX_CMP_WRITES * HEAD_BLOCKS - - -@pl.jit.inline -def prefill_compressor_ratio128( - x: pl.Tensor[[T, D], pl.BF16], - kv_state: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32], - score_state: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32], - compress_state_block_table: pl.Tensor[[HCA_STATE_MAX_BLOCKS], pl.INT32], - wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], - wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], - ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], - norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], - freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], - cmp_kv: pl.Out[pl.Tensor[[HCA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], - position_ids: pl.Tensor[[T], pl.INT32], - num_tokens: pl.Scalar[pl.INT32], - cmp_slot_mapping: pl.Tensor[[T], pl.INT64], - state_slot_mapping: pl.Tensor[[T], pl.INT64], -): - x_flat = x - kv_proj_scratch = pl.create_tensor([T, OUT_DIM], dtype=pl.FP32) - score_proj_scratch = pl.create_tensor([T, OUT_DIM], dtype=pl.FP32) - kv_state_flat = pl.reshape(kv_state, [HCA_STATE_BLOCK_NUM * HCA_STATE_BLOCK_SIZE, OUT_DIM]) - score_state_flat = pl.reshape(score_state, [HCA_STATE_BLOCK_NUM * HCA_STATE_BLOCK_SIZE, OUT_DIM]) - cmp_kv_flat = pl.reshape(cmp_kv, [HCA_CMP_BLOCK_NUM * BLOCK_SIZE, HEAD_DIM]) - pooled_kv_pad = pl.create_tensor([HCA_C128_RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) - normed_kv_pad = pl.create_tensor([HCA_C128_RMS_PAD_ROWS, HEAD_DIM], dtype=pl.FP32) - - with pl.at(level=pl.Level.CORE_GROUP, name_hint="prefill_hca_c128_norm_pad_init"): - for init_hb in pl.pipeline(HEAD_BLOCKS, stage=2): - init_h0 = init_hb * HEAD_TILE - zero_chunk = pl.full([HCA_C128_RMS_TILE, HEAD_TILE], dtype=pl.FP32, value=0.0) - pooled_kv_pad[0:HCA_C128_RMS_TILE, init_h0 : init_h0 + HEAD_TILE] = zero_chunk - normed_kv_pad[0:HCA_C128_RMS_TILE, init_h0 : init_h0 + HEAD_TILE] = zero_chunk - - for proj_idx in pl.spmd(PACKED_C128_PROJ_BLOCKS, name_hint="prefill_hca_c128_kv_score_proj"): - o0 = proj_idx * OUT_TILE - kv_acc = pl.create_tensor([T, OUT_TILE], dtype=pl.FP32) - score_acc = pl.create_tensor([T, OUT_TILE], dtype=pl.FP32) - for kb in pl.pipeline(0, K_BLOCKS, stage=2): - k0 = kb * K_TILE - x_tile = x_flat[0:T, k0 : k0 + K_TILE] - # Weights stored transposed [OUT_DIM, D] + b_trans=True -> DN2ZN load (K-contiguous - # long bursts) instead of ND2NZ (strided short bursts). Matches ratio4/CSA/decode-HCA. - wkv_tile = wkv[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] - wgate_tile = wgate[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] - if k0 == 0: - kv_acc = pl.matmul(x_tile, wkv_tile, out_dtype=pl.FP32, b_trans=True) - score_acc = pl.matmul(x_tile, wgate_tile, out_dtype=pl.FP32, b_trans=True) - else: - kv_acc = pl.matmul_acc(kv_acc, x_tile, wkv_tile, b_trans=True) - score_acc = pl.matmul_acc(score_acc, x_tile, wgate_tile, b_trans=True) - kv_proj_scratch[0:T, o0 : o0 + OUT_TILE] = kv_acc - score_proj_scratch[0:T, o0 : o0 + OUT_TILE] = score_acc - - # Precompute write_i -> (position, dst cache row) once (input-only deps -> overlaps the matmul), - # replacing the O(T) write-discovery scan in pool / rmsnorm_rope / kv_finalize. Sized to - # HCA_C128_RMS_TILE because rmsnorm_rope indexes padded rows beyond MAX_CMP_WRITES (rest stay -1). - write_pos_map = pl.create_tensor([1, HCA_C128_RMS_TILE], dtype=pl.INT32) - write_dst_map = pl.create_tensor([1, HCA_C128_RMS_TILE], dtype=pl.INT32) - with pl.at(level=pl.Level.CORE_GROUP, name_hint="prefill_hca_c128_write_map"): - write_pos_map[0:1, 0:HCA_C128_RMS_TILE] = pl.full([1, HCA_C128_RMS_TILE], dtype=pl.INT32, value=0) - write_dst_map[0:1, 0:HCA_C128_RMS_TILE] = pl.full([1, HCA_C128_RMS_TILE], dtype=pl.INT32, value=-1) - map_seen = pl.cast(0, pl.INDEX) - for map_w in pl.range(T): - if map_w < num_tokens: - map_slot_raw = pl.read(cmp_slot_mapping, [map_w]) - if map_slot_raw >= 0: - pl.write(write_pos_map, [0, map_seen], pl.read(position_ids, [map_w])) - pl.write(write_dst_map, [0, map_seen], pl.cast(map_slot_raw, pl.INT32)) - map_seen = map_seen + 1 - - # State scatter (decode order): write every token's raw projection (+APE on score) into - # paged kv_state/score_state BEFORE pooling, so softmax_pool reads its window straight from - # state (no seed+overlay, no pool_dep ordering hack). pool depends on this via kv_state RAW. - for scatter_t in pl.spmd(T, name_hint="prefill_hca_c128_state_scatter_pre"): - if scatter_t < num_tokens: - scatter_row_raw = pl.read(state_slot_mapping, [scatter_t]) - if scatter_row_raw >= 0: - scatter_row = pl.cast(scatter_row_raw, pl.INDEX) - scatter_pos = pl.read(position_ids, [scatter_t]) - scatter_ape_slot = pl.cast(scatter_pos % COMPRESS_RATIO, pl.INDEX) - kv_state_flat[scatter_row : scatter_row + 1, 0:OUT_DIM] = kv_proj_scratch[scatter_t : scatter_t + 1, 0:OUT_DIM] - score_state_flat[scatter_row : scatter_row + 1, 0:OUT_DIM] = pl.add( - score_proj_scratch[scatter_t : scatter_t + 1, 0:OUT_DIM], - ape[scatter_ape_slot : scatter_ape_slot + 1, 0:OUT_DIM], - ) - - for pool_idx in pl.spmd(PACKED_C128_POOL_BLOCKS, name_hint="prefill_hca_c128_softmax_pool"): - write_i = pool_idx // HEAD_BLOCKS - hb = pool_idx - write_i * HEAD_BLOCKS - h0 = hb * HEAD_TILE - pool_kv_tile = pl.create_tensor([STATE_LEN, HEAD_TILE], dtype=pl.FP32) - pool_score_tile = pl.create_tensor([STATE_LEN, HEAD_TILE], dtype=pl.FP32) - write_slot_raw = pl.read(write_dst_map, [0, write_i]) - if write_slot_raw >= 0: - write_pos = pl.read(write_pos_map, [0, write_i]) - for pool_state_i in pl.range(STATE_LEN): - pool_kv_tile[pool_state_i : pool_state_i + 1, 0:HEAD_TILE] = pl.full( - [1, HEAD_TILE], - dtype=pl.FP32, - value=0.0, - ) - pool_score_tile[pool_state_i : pool_state_i + 1, 0:HEAD_TILE] = pl.full( - [1, HEAD_TILE], - dtype=pl.FP32, - value=0.0, - ) - pool_abs = write_pos + 1 - COMPRESS_RATIO + pool_state_i - pool_state_block = pl.cast(pool_abs // HCA_STATE_BLOCK_SIZE, pl.INDEX) - pool_state_intra = pl.cast(pool_abs - pool_state_block * HCA_STATE_BLOCK_SIZE, pl.INDEX) - pool_phys_block_raw = pl.read(compress_state_block_table, [pool_state_block]) - if pool_phys_block_raw >= 0: - pool_phys_block = pl.cast(pool_phys_block_raw, pl.INDEX) - pool_state_row = pool_phys_block * HCA_STATE_BLOCK_SIZE + pool_state_intra - pool_kv_tile[pool_state_i : pool_state_i + 1, 0:HEAD_TILE] = kv_state_flat[ - pool_state_row : pool_state_row + 1, - h0 : h0 + HEAD_TILE, - ] - pool_score_tile[pool_state_i : pool_state_i + 1, 0:HEAD_TILE] = score_state_flat[ - pool_state_row : pool_state_row + 1, - h0 : h0 + HEAD_TILE, - ] - # Vectorized softmax over all STATE_LEN slots (matches decode128): transpose the - # assembled [STATE_LEN, HEAD_TILE] tile and do row_max/exp/sum/div + weighted sum, - # replacing the STATE_LEN-1 serial online-flash fold. Same result, no long chain. - pool_score_t = pl.transpose(pool_score_tile, axis1=0, axis2=1) - pool_kv_t = pl.transpose(pool_kv_tile, axis1=0, axis2=1) - score_max = pl.row_max(pool_score_t) - score_exp = pl.exp(pl.row_expand_sub(pool_score_t, score_max)) - score_sum = pl.row_sum(score_exp) - score_prob = pl.row_expand_div(score_exp, score_sum) - pooled_chunk_t = pl.row_sum(pl.mul(pool_kv_t, score_prob)) - pooled_kv_pad[write_i : write_i + 1, h0 : h0 + HEAD_TILE] = pl.reshape(pooled_chunk_t, [1, HEAD_TILE]) - else: - pooled_kv_pad[write_i : write_i + 1, h0 : h0 + HEAD_TILE] = pooled_kv_pad[ - write_i : write_i + 1, - h0 : h0 + HEAD_TILE, - ] - - norm_w_2d = pl.reshape(norm_w, [1, HEAD_DIM]) - with pl.at(level=pl.Level.CORE_GROUP, name_hint="prefill_hca_c128_rmsnorm_rope"): - cos_b = pl.full([HCA_C128_RMS_TILE, ROPE_HALF], dtype=pl.FP32, value=0.0) - sin_b = pl.full([HCA_C128_RMS_TILE, ROPE_HALF], dtype=pl.FP32, value=0.0) - for norm_i in pl.range(HCA_C128_RMS_TILE): - norm_slot_raw = pl.read(write_dst_map, [0, norm_i]) - if norm_slot_raw >= 0: - norm_cmp_pos = pl.cast(pl.read(write_pos_map, [0, norm_i]) + 1 - COMPRESS_RATIO, pl.INDEX) - cos_row = pl.cast(freqs_cos[norm_cmp_pos : norm_cmp_pos + 1, 0:ROPE_HALF], target_type=pl.FP32) - sin_row = pl.cast(freqs_sin[norm_cmp_pos : norm_cmp_pos + 1, 0:ROPE_HALF], target_type=pl.FP32) - cos_b[norm_i : norm_i + 1, 0:ROPE_HALF] = cos_row - sin_b[norm_i : norm_i + 1, 0:ROPE_HALF] = sin_row - partial_sq = pl.full([1, HCA_C128_RMS_TILE], dtype=pl.FP32, value=0.0) - for rms_kb in pl.pipeline(HEAD_BLOCKS, stage=2): - rms_h0 = rms_kb * HEAD_TILE - kv_rms_chunk = pooled_kv_pad[0:HCA_C128_RMS_TILE, rms_h0 : rms_h0 + HEAD_TILE] - kv_rms_sq = pl.mul(kv_rms_chunk, kv_rms_chunk) - partial_sq = pl.add(partial_sq, pl.reshape(pl.row_sum(kv_rms_sq), [1, HCA_C128_RMS_TILE])) - - variance = pl.reshape(pl.add(pl.mul(partial_sq, 1.0 / HEAD_DIM), EPS), [HCA_C128_RMS_TILE, 1]) - inv_rms = pl.recip(pl.sqrt(variance)) - for norm_kb in pl.pipeline(NOPE_HEAD_DIM // HEAD_TILE, stage=2): - norm_h0 = norm_kb * HEAD_TILE - kv_norm_chunk = pooled_kv_pad[0:HCA_C128_RMS_TILE, norm_h0 : norm_h0 + HEAD_TILE] - gamma = pl.cast(norm_w_2d[:, norm_h0 : norm_h0 + HEAD_TILE], pl.FP32) - normed_chunk = pl.col_expand_mul(pl.row_expand_mul(kv_norm_chunk, inv_rms), gamma) - normed_kv_pad[0:HCA_C128_RMS_TILE, norm_h0 : norm_h0 + HEAD_TILE] = normed_chunk - - kv_rope = pooled_kv_pad[0:HCA_C128_RMS_TILE, NOPE_HEAD_DIM:HEAD_DIM] - gamma_rope = pl.cast(norm_w_2d[:, NOPE_HEAD_DIM:HEAD_DIM], pl.FP32) - rope_normed = pl.col_expand_mul(pl.row_expand_mul(kv_rope, inv_rms), gamma_rope) - # A3 interleaved swap-gather (matches decode): single data gather + sign trick instead of - # the P0101/P1010 de-interleave gather + rotate + re-interleave scatter. - # out[j] = n[j]*cos_il[j] + n[j^1]*sign[j]*sin_il[j]; idx built in-kernel from pl.arange. - rope_ones = pl.full([HCA_C128_RMS_TILE, ROPE_HEAD_DIM], dtype=pl.FP32, value=1.0) - rope_col = pl.col_expand_mul(rope_ones, pl.cast(pl.arange(0, [1, ROPE_HEAD_DIM], dtype=pl.INT32), target_type=pl.FP32)) - rope_dup_f = pl.cast(pl.cast(pl.mul(rope_col, 0.5), target_type=pl.INT32, mode="trunc"), target_type=pl.FP32) - rope_dup_idx = pl.cast(rope_dup_f, target_type=pl.INT32) # j>>1 - rope_lane = pl.sub(rope_col, pl.mul(rope_dup_f, 2.0)) # j%2 - rope_swap_idx = pl.cast(pl.sub(pl.add(rope_col, 1.0), pl.mul(rope_lane, 2.0)), target_type=pl.INT32) # j^1 - rope_sign = pl.sub(pl.mul(rope_lane, 2.0), 1.0) # [-1,+1,...] - cos_il = pl.gather(cos_b, dim=-1, index=rope_dup_idx) - sin_il = pl.gather(sin_b, dim=-1, index=rope_dup_idx) - swapped = pl.gather(rope_normed, dim=-1, index=rope_swap_idx) - rope_rot = pl.add(pl.mul(rope_normed, cos_il), pl.mul(pl.mul(swapped, rope_sign), sin_il)) - normed_kv_pad[0:HCA_C128_RMS_TILE, NOPE_HEAD_DIM:HEAD_DIM] = rope_rot - - with pl.at(level=pl.Level.CORE_GROUP, name_hint="prefill_hca_c128_kv_finalize"): - for final_i in pl.range(MAX_CMP_WRITES): - final_cmp_row_raw = pl.read(write_dst_map, [0, final_i]) - if final_cmp_row_raw >= 0: - final_cmp_row = pl.cast(final_cmp_row_raw, pl.INDEX) - for final_hb in pl.range(HEAD_BLOCKS): - final_h0 = final_hb * HEAD_TILE - final_chunk = normed_kv_pad[final_i : final_i + 1, final_h0 : final_h0 + HEAD_TILE] - cmp_kv_flat[final_cmp_row : final_cmp_row + 1, final_h0 : final_h0 + HEAD_TILE] = pl.cast( - final_chunk, - target_type=pl.BF16, - mode="rint", - ) - - cmp_kv = pl.reshape(cmp_kv_flat, [HCA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM]) - kv_state = pl.reshape(kv_state_flat, [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, OUT_DIM]) - score_state = pl.reshape(score_state_flat, [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, OUT_DIM]) - return cmp_kv, kv_state, score_state - - -@pl.jit -def prefill_compressor_ratio128_test( - x: pl.Tensor[[T, D], pl.BF16], - kv_state: pl.InOut[pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32]], - score_state: pl.InOut[pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32]], - compress_state_block_table: pl.Tensor[[HCA_STATE_MAX_BLOCKS], pl.INT32], - wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], - wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], - ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], - norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], - freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], - cmp_kv: pl.InOut[pl.Tensor[[HCA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], - position_ids: pl.Tensor[[T], pl.INT32], - num_tokens: pl.Scalar[pl.INT32], - cmp_slot_mapping: pl.Tensor[[T], pl.INT64], - state_slot_mapping: pl.Tensor[[T], pl.INT64], -): - return prefill_compressor_ratio128( - x, kv_state, score_state, compress_state_block_table, wkv, wgate, ape, norm_w, freqs_cos, freqs_sin, - cmp_kv, position_ids, num_tokens, cmp_slot_mapping, state_slot_mapping, - ) - - -def golden_prefill_compressor_ratio128(tensors): - import torch - - num_tokens = int(tensors["num_tokens"]) - kv_proj = tensors["x"].float() @ tensors["wkv"].float().t() # wkv stored [OUT_DIM, D] for b_trans - score_proj = tensors["x"].float() @ tensors["wgate"].float().t() - kv_state_flat = tensors["kv_state"].view(HCA_STATE_BLOCK_NUM * HCA_STATE_BLOCK_SIZE, OUT_DIM) - score_state_flat = tensors["score_state"].view(HCA_STATE_BLOCK_NUM * HCA_STATE_BLOCK_SIZE, OUT_DIM) - state_block_table = tensors["compress_state_block_table"] - cmp_kv_flat = tensors["cmp_kv"].view(HCA_CMP_BLOCK_NUM * BLOCK_SIZE, HEAD_DIM) - - def state_row(abs_pos): - if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: - return -1 - block = abs_pos // HCA_STATE_BLOCK_SIZE - intra = abs_pos % HCA_STATE_BLOCK_SIZE - phys_block = int(state_block_table[block].item()) - if phys_block < 0: - return -1 - return phys_block * HCA_STATE_BLOCK_SIZE + intra - - for token_id in range(num_tokens): - dst_row = int(tensors["cmp_slot_mapping"][token_id].item()) - if dst_row < 0: - continue - write_pos = int(tensors["position_ids"][token_id].item()) - pool_kv_state = torch.zeros(STATE_LEN, OUT_DIM, dtype=torch.float32) - pool_score_state = torch.zeros(STATE_LEN, OUT_DIM, dtype=torch.float32) - for slot in range(STATE_LEN): - row = state_row(write_pos + 1 - COMPRESS_RATIO + slot) - if row >= 0: - pool_kv_state[slot] = kv_state_flat[row] - pool_score_state[slot] = score_state_flat[row] - for t in range(num_tokens): - pos = int(tensors["position_ids"][t].item()) - if pos > write_pos: - continue - slot = pos % COMPRESS_RATIO - pool_kv_state[slot] = kv_proj[t] - pool_score_state[slot] = score_proj[t] + tensors["ape"][slot] - pooled = (pool_kv_state * pool_score_state.softmax(dim=0)).sum(dim=0, keepdim=True) - inv = torch.rsqrt(pooled.square().mean(dim=-1, keepdim=True) + EPS) - normed = pooled * inv * tensors["norm_w"].float().view(1, HEAD_DIM) - rope_pair = normed[..., NOPE_HEAD_DIM:].unflatten(-1, (-1, 2)) - even = rope_pair[..., 0].float() - odd = rope_pair[..., 1].float() - cmp_pos = write_pos + 1 - COMPRESS_RATIO - cos = tensors["freqs_cos"][cmp_pos : cmp_pos + 1, 0:ROPE_HALF].float() - sin = tensors["freqs_sin"][cmp_pos : cmp_pos + 1, 0:ROPE_HALF].float() - rot_even = even * cos - odd * sin - rot_odd = even * sin + odd * cos - normed[:, NOPE_HEAD_DIM:] = torch.stack([rot_even, rot_odd], dim=-1).flatten(-2) - cmp_kv_flat[dst_row] = normed[0] - - for t in range(num_tokens): - pos = int(tensors["position_ids"][t].item()) - dst_row = int(tensors["state_slot_mapping"][t].item()) - if dst_row < 0: - continue - slot = pos % COMPRESS_RATIO - kv_state_flat[dst_row] = kv_proj[t] - score_state_flat[dst_row] = score_proj[t] + tensors["ape"][slot] - - -def build_tensor_specs(start_pos: int = START_POS): - import torch - from golden import ScalarSpec, TensorSpec - from rope_tables import build_deepseek_v4_rope_tables - - shared_freqs_cos, shared_freqs_sin = build_deepseek_v4_rope_tables(M, COMPRESS_RATIO, dtype=torch.bfloat16) - - num_tokens = T - if start_pos < 0: - raise ValueError("start_pos must be non-negative") - if start_pos + num_tokens > MAX_SEQ_LEN: - raise ValueError("start_pos + num_tokens exceeds max_position_embeddings") - - def init_compress_state_block_table(): - table = torch.full((HCA_STATE_MAX_BLOCKS,), -1, dtype=torch.int32) - for block in range(HCA_STATE_MAX_BLOCKS): - table[block] = (block * 17 + 3) % HCA_STATE_MAX_BLOCKS - return table - def state_row(abs_pos): - if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: - return -1 - table = init_compress_state_block_table() - block = abs_pos // HCA_STATE_BLOCK_SIZE - intra = abs_pos % HCA_STATE_BLOCK_SIZE - return int(table[block].item()) * HCA_STATE_BLOCK_SIZE + intra - def init_x(): - return ((torch.rand(T, D) - 0.5) * 0.1).to(torch.bfloat16) - def init_state(): - state = torch.zeros(HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, OUT_DIM) - for abs_pos in range(max(0, start_pos - COMPRESS_RATIO), start_pos): - row = state_row(abs_pos) - if row >= 0: - state.view(-1, OUT_DIM)[row] = (torch.rand(OUT_DIM) - 0.5) * 0.05 - return state - # Calibrated to the real DeepSeek-V4-Flash HCA (ratio-128) main compressor (mean l7/l9 of - # extract_weights_flash): zero-mean Gaussian BF16 weights at the measured std; the RMSNorm - # gamma centers near the measured mean (not ones / not uniform). Mirrors decode_compressor_ratio128. - def init_wkv(): - return torch.randn(OUT_DIM, D) * 0.0246 - def init_wgate(): - return torch.randn(OUT_DIM, D) * 0.0316 - def init_ape(): - return torch.randn(COMPRESS_RATIO, OUT_DIM) * 0.0340 - def init_norm_w(): - return 0.1001 + 0.0549 * torch.randn(HEAD_DIM) - def init_freqs_cos(): - return shared_freqs_cos.clone() - def init_freqs_sin(): - return shared_freqs_sin.clone() - def init_cmp_kv(): - return torch.zeros(HCA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM, dtype=torch.bfloat16) - def init_position_ids(): - return torch.arange(start_pos, start_pos + T, dtype=torch.int32) - def init_cmp_slot_mapping(): - mapping = torch.full((T,), -1, dtype=torch.int64) - for token_id in range(num_tokens): - pos = start_pos + token_id - if pos + 1 >= COMPRESS_RATIO and (pos + 1) % COMPRESS_RATIO == 0: - mapping[token_id] = (pos + 1) // COMPRESS_RATIO - 1 - return mapping - def init_state_slot_mapping(): - mapping = torch.full((T,), -1, dtype=torch.int64) - for token_id in range(num_tokens): - mapping[token_id] = state_row(start_pos + token_id) - return mapping - - return [ - TensorSpec("x", [T, D], torch.bfloat16, init_value=init_x), - TensorSpec("kv_state", [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, OUT_DIM], torch.float32, init_value=init_state, is_output=True), - TensorSpec("score_state", [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, OUT_DIM], torch.float32, init_value=init_state, is_output=True), - TensorSpec("compress_state_block_table", [HCA_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), - TensorSpec("wkv", [OUT_DIM, D], torch.bfloat16, init_value=init_wkv), - TensorSpec("wgate", [OUT_DIM, D], torch.bfloat16, init_value=init_wgate), - TensorSpec("ape", [COMPRESS_RATIO, OUT_DIM], torch.float32, init_value=init_ape), - TensorSpec("norm_w", [HEAD_DIM], torch.bfloat16, init_value=init_norm_w), - TensorSpec("freqs_cos", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_cos), - TensorSpec("freqs_sin", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_sin), - TensorSpec("cmp_kv", [HCA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], torch.bfloat16, init_value=init_cmp_kv, is_output=True), - TensorSpec("position_ids", [T], torch.int32, init_value=init_position_ids), - ScalarSpec("num_tokens", torch.int32, num_tokens), - TensorSpec("cmp_slot_mapping", [T], torch.int64, init_value=init_cmp_slot_mapping), - TensorSpec("state_slot_mapping", [T], torch.int64, init_value=init_state_slot_mapping), - ] - - -if __name__ == "__main__": - import argparse - from golden import ratio_allclose, run_jit - - parser = argparse.ArgumentParser(description="Standalone token-major DeepSeek V4 prefill compressor ratio128 validation.") - parser.add_argument("-p", "--platform", type=str, default="a2a3", - choices=["a2a3", "a2a3sim", "a5", "a5sim"]) - parser.add_argument("-d", "--device", type=int, default=0) - parser.add_argument( - "--compile-only", - action="store_true", - default=False, - help="Compile/codegen only. This is also the implicit behavior on *sim platforms used by CI.", - ) - parser.add_argument( - "--start-pos", - type=int, - default=START_POS, - help=( - "Fixture-only absolute position for token 0. It is lowered into position_ids and compressed write " - "slot mapping; it is not a JIT kernel parameter." - ), - ) - parser.add_argument("--enable-l2-swimlane", action="store_true", default=False) - parser.add_argument("--dump-passes", action="store_true", default=False) - args = parser.parse_args() - - result = run_jit( - fn=prefill_compressor_ratio128_test, - specs=build_tensor_specs(args.start_pos), - golden_fn=golden_prefill_compressor_ratio128, - compile_cfg=dict(dump_passes=args.dump_passes), - runtime_cfg=dict(platform=args.platform, device_id=args.device, enable_l2_swimlane=args.enable_l2_swimlane), - rtol=1e-3, - atol=1e-3, - compile_only=args.compile_only, - compare_fn={ - "cmp_kv": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), - "kv_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), - "score_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), - }, - ) - if not result.passed: - if result.error: - print(result.error) - raise SystemExit(1) diff --git a/models/deepseek/v4/prefill_compressor_ratio4.py b/models/deepseek/v4/prefill_compressor_ratio4.py deleted file mode 100644 index 6c5017fb..00000000 --- a/models/deepseek/v4/prefill_compressor_ratio4.py +++ /dev/null @@ -1,601 +0,0 @@ -# Copyright (c) PyPTO Contributors. -# This program is free software, you can redistribute it and/or modify it under the terms and conditions of -# CANN Open Software License Agreement Version 2.0 (the "License"). -# Please refer to the License for details. You may not use this file except in compliance with the License. -# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, -# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. -# See LICENSE in the root of the software repository for the full text of the License. -# ----------------------------------------------------------------------------------------------------------- -"""DeepSeek-V4 prefill attention compressor for ratio-4 overlapping KV cache (rotate=False).""" - -import pypto.language as pl - -from config import FLASH as M, BLOCK_SIZE, FP32_NEG_INF, PREFILL_CMP_BLOCK_NUM - -# model config (mirrors decode_compressor_ratio4) -EPS = M.rms_norm_eps -D = M.hidden_size -HEAD_DIM = M.head_dim -HEAD_DIM_INV = 1.0 / HEAD_DIM -ROPE_HEAD_DIM = M.qk_rope_head_dim -NOPE_HEAD_DIM = M.nope_head_dim -MAX_SEQ_LEN = M.max_position_embeddings - -# kernel-local (ratio-4 overlapping compressor) -COMPRESS_RATIO = 4 -OVERLAP = COMPRESS_RATIO == 4 -COFF = 1 + int(OVERLAP) -OUT_DIM = COFF * HEAD_DIM -STATE_LEN = COFF * COMPRESS_RATIO - -B = 1 -S = 128 -START_POS = 0 -PREFILL_COMPRESSED_LEN = S // COMPRESS_RATIO -PREFILL_ROWS = B * PREFILL_COMPRESSED_LEN -HEAD_CHUNK = 256 -assert HEAD_DIM % HEAD_CHUNK == 0 -HEAD_BLOCKS = HEAD_DIM // HEAD_CHUNK -K_TILE = 512 -OUT_TILE = 32 -HEAD_TILE = 64 -RMS_TILE = 16 - -T = B * S -CSA_STATE_BLOCK_SIZE = 4 -CSA_STATE_MAX_BLOCKS = (MAX_SEQ_LEN + CSA_STATE_BLOCK_SIZE - 1) // CSA_STATE_BLOCK_SIZE -CSA_STATE_BLOCK_NUM = CSA_STATE_MAX_BLOCKS -MAX_CMP_WRITES = max(1, T // COMPRESS_RATIO) -PACKED_PROJ_BLOCKS = OUT_DIM // OUT_TILE -PACKED_POOL_BLOCKS = MAX_CMP_WRITES * HEAD_BLOCKS -PACKED_STATE_UPDATE_TILE = 16 -PACKED_RMS_TILE = 16 - - -@pl.jit.inline -def prefill_compressor_ratio4( - x: pl.Tensor[[T, D], pl.BF16], - kv_state: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32], - score_state: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32], - compress_state_block_table: pl.Tensor[[CSA_STATE_MAX_BLOCKS], pl.INT32], - wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], - wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], - ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], - norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], - freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], - cmp_kv: pl.Tensor[[PREFILL_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16], - position_ids: pl.Tensor[[T], pl.INT32], - num_tokens: pl.Scalar[pl.INT32], - cmp_slot_mapping: pl.Tensor[[T], pl.INT64], - state_slot_mapping: pl.Tensor[[T], pl.INT64], -): - cmp4_kv_proj_scratch = pl.create_tensor([T, OUT_DIM], dtype=pl.FP32) - cmp4_score_proj_scratch = pl.create_tensor([T, OUT_DIM], dtype=pl.FP32) - kv_state_flat = pl.reshape(kv_state, [CSA_STATE_BLOCK_NUM * CSA_STATE_BLOCK_SIZE, OUT_DIM]) - score_state_flat = pl.reshape(score_state, [CSA_STATE_BLOCK_NUM * CSA_STATE_BLOCK_SIZE, OUT_DIM]) - cmp_kv_flat = pl.reshape(cmp_kv, [PREFILL_CMP_BLOCK_NUM * BLOCK_SIZE, HEAD_DIM]) - pooled_kv = pl.create_tensor([MAX_CMP_WRITES, HEAD_DIM], dtype=pl.FP32) - normed_kv = pl.create_tensor([MAX_CMP_WRITES, HEAD_DIM], dtype=pl.FP32) - - for proj_idx in pl.spmd(PACKED_PROJ_BLOCKS, name_hint="prefill_c4_kv_score_proj"): - o0 = proj_idx * OUT_TILE - kv_acc = pl.create_tensor([T, OUT_TILE], dtype=pl.FP32) - score_acc = pl.create_tensor([T, OUT_TILE], dtype=pl.FP32) - for kb in pl.pipeline(0, D // K_TILE, stage=2): - k0 = kb * K_TILE - x_tile = x[0:T, k0 : k0 + K_TILE] - # Weights stored transposed [OUT_DIM, D] + b_trans=True -> DN2ZN load (K-contiguous - # long bursts) instead of ND2NZ (strided short bursts). Matches ratio4/CSA/HCA layout. - wkv_tile = wkv[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] - wgate_tile = wgate[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] - if k0 == 0: - kv_acc = pl.matmul(x_tile, wkv_tile, out_dtype=pl.FP32, b_trans=True) - score_acc = pl.matmul(x_tile, wgate_tile, out_dtype=pl.FP32, b_trans=True) - else: - kv_acc = pl.matmul_acc(kv_acc, x_tile, wkv_tile, b_trans=True) - score_acc = pl.matmul_acc(score_acc, x_tile, wgate_tile, b_trans=True) - cmp4_kv_proj_scratch[0:T, o0 : o0 + OUT_TILE] = kv_acc - cmp4_score_proj_scratch[0:T, o0 : o0 + OUT_TILE] = score_acc - - # Precompute write_i -> (position, dst cache row) once. Depends only on the slot-mapping and - # position inputs, so it overlaps the projection matmul, replacing the O(T) write-discovery - # scan that every later stage (pool / rmsnorm_rope / cache_write) otherwise repeats. - write_pos_map = pl.create_tensor([1, MAX_CMP_WRITES], dtype=pl.INT32) - write_dst_map = pl.create_tensor([1, MAX_CMP_WRITES], dtype=pl.INT32) - with pl.at(level=pl.Level.CORE_GROUP, name_hint="prefill_c4_write_map"): - write_pos_map[0:1, 0:MAX_CMP_WRITES] = pl.full([1, MAX_CMP_WRITES], dtype=pl.INT32, value=0) - write_dst_map[0:1, 0:MAX_CMP_WRITES] = pl.full([1, MAX_CMP_WRITES], dtype=pl.INT32, value=-1) - map_seen = pl.cast(0, pl.INDEX) - for map_w in pl.range(T): - if map_w < num_tokens: - map_slot_raw = pl.read(cmp_slot_mapping, [map_w]) - if map_slot_raw >= 0: - pl.write(write_pos_map, [0, map_seen], pl.read(position_ids, [map_w])) - pl.write(write_dst_map, [0, map_seen], pl.cast(map_slot_raw, pl.INT32)) - map_seen = map_seen + 1 - - for pool_idx in pl.spmd(PACKED_POOL_BLOCKS, name_hint="prefill_c4_softmax_pool"): - write_i = pool_idx // HEAD_BLOCKS - hb = pool_idx - write_i * HEAD_BLOCKS - h0 = hb * HEAD_CHUNK - pool_kv_tile = pl.create_tensor([STATE_LEN, HEAD_CHUNK], dtype=pl.FP32) - pool_score_tile = pl.create_tensor([STATE_LEN, HEAD_CHUNK], dtype=pl.FP32) - write_slot_raw = pl.read(write_dst_map, [0, write_i]) - if write_slot_raw >= 0: - write_pos = pl.read(write_pos_map, [0, write_i]) - cur_start = write_pos + 1 - COMPRESS_RATIO - prev_start = cur_start - COMPRESS_RATIO - for pool_s in pl.range(COMPRESS_RATIO): - prev_abs = prev_start + pool_s - front_slot = pool_s - pool_kv_tile[front_slot : front_slot + 1, 0:HEAD_CHUNK] = pl.full( - [1, HEAD_CHUNK], - dtype=pl.FP32, - value=0.0, - ) - pool_score_tile[front_slot : front_slot + 1, 0:HEAD_CHUNK] = pl.full( - [1, HEAD_CHUNK], - dtype=pl.FP32, - value=FP32_NEG_INF, - ) - if write_pos >= 2 * COMPRESS_RATIO - 1: - prev_state_block = pl.cast(prev_abs // CSA_STATE_BLOCK_SIZE, pl.INDEX) - prev_state_intra = pl.cast(prev_abs - prev_state_block * CSA_STATE_BLOCK_SIZE, pl.INDEX) - prev_phys_block_raw = pl.read(compress_state_block_table, [prev_state_block]) - if prev_phys_block_raw >= 0: - prev_phys_block = pl.cast(prev_phys_block_raw, pl.INDEX) - prev_state_row = prev_phys_block * CSA_STATE_BLOCK_SIZE + prev_state_intra - pool_kv_tile[front_slot : front_slot + 1, 0:HEAD_CHUNK] = kv_state_flat[ - prev_state_row : prev_state_row + 1, - h0 : h0 + HEAD_CHUNK, - ] - pool_score_tile[front_slot : front_slot + 1, 0:HEAD_CHUNK] = score_state_flat[ - prev_state_row : prev_state_row + 1, - h0 : h0 + HEAD_CHUNK, - ] - - cur_abs = cur_start + pool_s - back_slot = COMPRESS_RATIO + pool_s - pool_kv_tile[back_slot : back_slot + 1, 0:HEAD_CHUNK] = pl.full( - [1, HEAD_CHUNK], - dtype=pl.FP32, - value=0.0, - ) - pool_score_tile[back_slot : back_slot + 1, 0:HEAD_CHUNK] = pl.full( - [1, HEAD_CHUNK], - dtype=pl.FP32, - value=FP32_NEG_INF, - ) - cur_state_block = pl.cast(cur_abs // CSA_STATE_BLOCK_SIZE, pl.INDEX) - cur_state_intra = pl.cast(cur_abs - cur_state_block * CSA_STATE_BLOCK_SIZE, pl.INDEX) - cur_phys_block_raw = pl.read(compress_state_block_table, [cur_state_block]) - if cur_phys_block_raw >= 0: - cur_phys_block = pl.cast(cur_phys_block_raw, pl.INDEX) - cur_state_row = cur_phys_block * CSA_STATE_BLOCK_SIZE + cur_state_intra - pool_kv_tile[back_slot : back_slot + 1, 0:HEAD_CHUNK] = kv_state_flat[ - cur_state_row : cur_state_row + 1, - HEAD_DIM + h0 : HEAD_DIM + h0 + HEAD_CHUNK, - ] - pool_score_tile[back_slot : back_slot + 1, 0:HEAD_CHUNK] = score_state_flat[ - cur_state_row : cur_state_row + 1, - HEAD_DIM + h0 : HEAD_DIM + h0 + HEAD_CHUNK, - ] - - for pool_t in pl.range(T): - if pool_t < num_tokens: - pool_pos = pl.read(position_ids, [pool_t]) - if pool_pos <= write_pos: - if pool_pos >= prev_start: - if pool_pos < cur_start: - pool_slot = pl.cast(pool_pos - prev_start, pl.INDEX) - pool_col0 = h0 - else: - pool_slot = pl.cast(COMPRESS_RATIO + pool_pos - cur_start, pl.INDEX) - pool_col0 = HEAD_DIM + h0 - pool_ape_slot = pl.cast(pool_pos % COMPRESS_RATIO, pl.INDEX) - pool_ape = ape[pool_ape_slot : pool_ape_slot + 1, pool_col0 : pool_col0 + HEAD_CHUNK] - pool_score = pl.add( - cmp4_score_proj_scratch[pool_t : pool_t + 1, pool_col0 : pool_col0 + HEAD_CHUNK], - pool_ape, - ) - pool_kv_tile[pool_slot : pool_slot + 1, 0:HEAD_CHUNK] = cmp4_kv_proj_scratch[ - pool_t : pool_t + 1, - pool_col0 : pool_col0 + HEAD_CHUNK, - ] - pool_score_tile[pool_slot : pool_slot + 1, 0:HEAD_CHUNK] = pool_score - - init_slot = STATE_LEN - 1 - mi_buf = pl.create_tensor([1, HEAD_CHUNK], dtype=pl.FP32) - li_buf = pl.create_tensor([1, HEAD_CHUNK], dtype=pl.FP32) - oi_buf = pl.create_tensor([1, HEAD_CHUNK], dtype=pl.FP32) - mi_buf[0:1, 0:HEAD_CHUNK] = pool_score_tile[init_slot : init_slot + 1, 0:HEAD_CHUNK] - li_buf[0:1, 0:HEAD_CHUNK] = pl.exp(pl.sub(mi_buf[0:1, 0:HEAD_CHUNK], mi_buf[0:1, 0:HEAD_CHUNK])) - oi_buf[0:1, 0:HEAD_CHUNK] = pool_kv_tile[init_slot : init_slot + 1, 0:HEAD_CHUNK] - for pool_slot_i in pl.range(STATE_LEN - 1): - if pool_slot_i >= COMPRESS_RATIO or write_pos >= 2 * COMPRESS_RATIO - 1: - mi = mi_buf[0:1, 0:HEAD_CHUNK] - li = li_buf[0:1, 0:HEAD_CHUNK] - oi = oi_buf[0:1, 0:HEAD_CHUNK] - slot_score = pool_score_tile[pool_slot_i : pool_slot_i + 1, 0:HEAD_CHUNK] - slot_kv = pool_kv_tile[pool_slot_i : pool_slot_i + 1, 0:HEAD_CHUNK] - mi_next = pl.maximum(mi, slot_score) - alpha = pl.exp(pl.sub(mi, mi_next)) - beta = pl.exp(pl.sub(slot_score, mi_next)) - li_next = pl.add(pl.mul(alpha, li), beta) - oi_next = pl.add(pl.mul(oi, alpha), pl.mul(slot_kv, beta)) - mi_buf[0:1, 0:HEAD_CHUNK] = mi_next - li_buf[0:1, 0:HEAD_CHUNK] = li_next - oi_buf[0:1, 0:HEAD_CHUNK] = oi_next - pooled_kv[write_i : write_i + 1, h0 : h0 + HEAD_CHUNK] = pl.div( - oi_buf[0:1, 0:HEAD_CHUNK], - li_buf[0:1, 0:HEAD_CHUNK], - ) - else: - pooled_kv[write_i : write_i + 1, h0 : h0 + HEAD_CHUNK] = pl.full([1, HEAD_CHUNK], dtype=pl.FP32, value=0.0) - - norm_w_2d = pl.reshape(norm_w, [1, HEAD_DIM]) - for final_block in pl.spmd(MAX_CMP_WRITES // PACKED_RMS_TILE, name_hint="prefill_c4_rmsnorm_rope"): - final_base = final_block * PACKED_RMS_TILE - cos_b = pl.full([PACKED_RMS_TILE, ROPE_HEAD_DIM // 2], dtype=pl.FP32, value=0.0) - sin_b = pl.full([PACKED_RMS_TILE, ROPE_HEAD_DIM // 2], dtype=pl.FP32, value=0.0) - for final_dt in pl.range(PACKED_RMS_TILE): - final_i = final_base + final_dt - write_slot_raw = pl.read(write_dst_map, [0, final_i]) - if write_slot_raw >= 0: - write_pos = pl.read(write_pos_map, [0, final_i]) - cmp_pos = pl.cast(write_pos + 1 - COMPRESS_RATIO, pl.INDEX) - cos_b[final_dt : final_dt + 1, 0 : ROPE_HEAD_DIM // 2] = pl.cast( - freqs_cos[cmp_pos : cmp_pos + 1, 0 : ROPE_HEAD_DIM // 2], - target_type=pl.FP32, - ) - sin_b[final_dt : final_dt + 1, 0 : ROPE_HEAD_DIM // 2] = pl.cast( - freqs_sin[cmp_pos : cmp_pos + 1, 0 : ROPE_HEAD_DIM // 2], - target_type=pl.FP32, - ) - - partial_sq = pl.full([1, PACKED_RMS_TILE], dtype=pl.FP32, value=0.0) - for k0 in pl.range(0, HEAD_DIM, HEAD_TILE): - kv_rms_chunk = pooled_kv[final_base : final_base + PACKED_RMS_TILE, k0 : k0 + HEAD_TILE] - kv_rms_sq = pl.mul(kv_rms_chunk, kv_rms_chunk) - partial_sq = pl.add(partial_sq, pl.reshape(pl.row_sum(kv_rms_sq), [1, PACKED_RMS_TILE])) - variance = pl.reshape(pl.add(pl.mul(partial_sq, HEAD_DIM_INV), EPS), [PACKED_RMS_TILE, 1]) - inv_rms = pl.recip(pl.sqrt(variance)) - for k0 in pl.range(0, NOPE_HEAD_DIM, HEAD_TILE): - kv_norm_chunk = pooled_kv[final_base : final_base + PACKED_RMS_TILE, k0 : k0 + HEAD_TILE] - gamma = pl.cast(norm_w_2d[:, k0 : k0 + HEAD_TILE], pl.FP32) - normed_chunk = pl.col_expand_mul(pl.row_expand_mul(kv_norm_chunk, inv_rms), gamma) - normed_kv[final_base : final_base + PACKED_RMS_TILE, k0 : k0 + HEAD_TILE] = normed_chunk - kv_rope_norm = pooled_kv[final_base : final_base + PACKED_RMS_TILE, NOPE_HEAD_DIM : HEAD_DIM] - gamma_rope = pl.cast(norm_w_2d[:, NOPE_HEAD_DIM : HEAD_DIM], pl.FP32) - rope_normed = pl.col_expand_mul(pl.row_expand_mul(kv_rope_norm, inv_rms), gamma_rope) - # A3 interleaved swap-gather (matches decode): single data gather + sign trick instead of - # the P0101/P1010 de-interleave gather + rotate + re-interleave scatter. swap_idx (j^1), - # sign ([-1,+1,...]) and dup_idx (j>>1) are built in-kernel from pl.arange; cos_il/sin_il - # dup-gather the half-width cos_b/sin_b. out[j] = n[j]*cos_il[j] + n[j^1]*sign[j]*sin_il[j]. - rope_ones = pl.full([PACKED_RMS_TILE, ROPE_HEAD_DIM], dtype=pl.FP32, value=1.0) - rope_col = pl.col_expand_mul(rope_ones, pl.cast(pl.arange(0, [1, ROPE_HEAD_DIM], dtype=pl.INT32), target_type=pl.FP32)) - rope_dup_f = pl.cast(pl.cast(pl.mul(rope_col, 0.5), target_type=pl.INT32, mode="trunc"), target_type=pl.FP32) - rope_dup_idx = pl.cast(rope_dup_f, target_type=pl.INT32) # j>>1 - rope_lane = pl.sub(rope_col, pl.mul(rope_dup_f, 2.0)) # j%2 - rope_swap_idx = pl.cast(pl.sub(pl.add(rope_col, 1.0), pl.mul(rope_lane, 2.0)), target_type=pl.INT32) # j^1 - rope_sign = pl.sub(pl.mul(rope_lane, 2.0), 1.0) # [-1,+1,...] - cos_il = pl.gather(cos_b, dim=-1, index=rope_dup_idx) - sin_il = pl.gather(sin_b, dim=-1, index=rope_dup_idx) - swapped = pl.gather(rope_normed, dim=-1, index=rope_swap_idx) - rope_rot = pl.add(pl.mul(rope_normed, cos_il), pl.mul(pl.mul(swapped, rope_sign), sin_il)) - normed_kv[final_base : final_base + PACKED_RMS_TILE, NOPE_HEAD_DIM : HEAD_DIM] = rope_rot - - for final_block in pl.spmd(MAX_CMP_WRITES // PACKED_RMS_TILE, name_hint="prefill_c4_cache_write"): - final_base = final_block * PACKED_RMS_TILE - for final_dt in pl.range(PACKED_RMS_TILE): - final_i = final_base + final_dt - dst_row_raw = pl.read(write_dst_map, [0, final_i]) - if dst_row_raw >= 0: - dst_row = pl.cast(dst_row_raw, pl.INDEX) - cmp_kv_flat[dst_row : dst_row + 1, 0:HEAD_DIM] = pl.cast( - normed_kv[final_i : final_i + 1, 0:HEAD_DIM], - target_type=pl.BF16, - mode="rint", - ) - else: - keepalive_row = PREFILL_CMP_BLOCK_NUM * BLOCK_SIZE - MAX_CMP_WRITES + final_i - cmp_kv_flat[keepalive_row : keepalive_row + 1, 0:HEAD_DIM] = cmp_kv_flat[ - keepalive_row : keepalive_row + 1, - 0:HEAD_DIM, - ] - - # State writeback: one SPMD task per token (was per token x out-block = - # T*PACKED_PROJ_BLOCKS tiny tasks). The per-token guard is checked - # once so a skipped token costs one empty task instead of PACKED_PROJ_BLOCKS - # of them; out-blocks are looped inside the task at the OUT_TILE width. - # pool_dep keeps the (zero-weighted) ordering after the pool and is hoisted - # to once per token. - for update_t in pl.spmd(T, name_hint="prefill_c4_state_update"): - if update_t < num_tokens: - state_row_raw = pl.read(state_slot_mapping, [update_t]) - if state_row_raw >= 0: - state_row = pl.cast(state_row_raw, pl.INDEX) - update_pos = pl.read(position_ids, [update_t]) - ape_slot = pl.cast(update_pos % COMPRESS_RATIO, pl.INDEX) - pool_dep = pl.mul(pooled_kv[0:1, 0:OUT_TILE], 0.0) - for update_ob in pl.range(PACKED_PROJ_BLOCKS): - update_o0 = update_ob * OUT_TILE - ape_row = ape[ape_slot : ape_slot + 1, update_o0 : update_o0 + OUT_TILE] - kv_state_flat[state_row : state_row + 1, update_o0 : update_o0 + OUT_TILE] = pl.add( - cmp4_kv_proj_scratch[ - update_t : update_t + 1, - update_o0 : update_o0 + OUT_TILE, - ], - pool_dep, - ) - score_state_flat[state_row : state_row + 1, update_o0 : update_o0 + OUT_TILE] = pl.add( - pl.add( - cmp4_score_proj_scratch[update_t : update_t + 1, update_o0 : update_o0 + OUT_TILE], - ape_row, - ), - pool_dep, - ) - - cmp_kv = pl.reshape(cmp_kv_flat, [PREFILL_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM]) - kv_state = pl.reshape(kv_state_flat, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, OUT_DIM]) - score_state = pl.reshape(score_state_flat, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, OUT_DIM]) - return cmp_kv, kv_state, score_state - - -def golden_prefill_compressor_ratio4(tensors): - """Packed token-major torch reference for ratio-4 prefill compressor.""" - import torch - - x = tensors["x"].view(T, D).float() - kv_state_flat = tensors["kv_state"].view(CSA_STATE_BLOCK_NUM * CSA_STATE_BLOCK_SIZE, OUT_DIM) - score_state_flat = tensors["score_state"].view(CSA_STATE_BLOCK_NUM * CSA_STATE_BLOCK_SIZE, OUT_DIM) - state_block_table = tensors["compress_state_block_table"] - wkv = tensors["wkv"].float() - wgate = tensors["wgate"].float() - ape = tensors["ape"] - norm_w = tensors["norm_w"] - cmp_kv = tensors["cmp_kv"] - cache_rows = cmp_kv.view(cmp_kv.shape[0] * BLOCK_SIZE, 1, HEAD_DIM)[:, 0, :] - position_ids = tensors["position_ids"] - - kv_proj = x @ wkv.t() # wkv stored [OUT_DIM, D] for b_trans - score_proj = x @ wgate.t() - - def state_row(abs_pos): - if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: - return -1 - block = abs_pos // CSA_STATE_BLOCK_SIZE - intra = abs_pos % CSA_STATE_BLOCK_SIZE - phys_block = int(state_block_table[block].item()) - if phys_block < 0: - return -1 - return phys_block * CSA_STATE_BLOCK_SIZE + intra - - for token_id in range(int(tensors["num_tokens"])): - dst_row = int(tensors["cmp_slot_mapping"][token_id].item()) - if dst_row < 0: - continue - write_pos = int(position_ids[token_id].item()) - cur_start = write_pos + 1 - COMPRESS_RATIO - prev_start = cur_start - COMPRESS_RATIO - pool_kv = torch.zeros(STATE_LEN, HEAD_DIM, dtype=torch.float32) - pool_score = torch.full((STATE_LEN, HEAD_DIM), float("-inf"), dtype=torch.float32) - - for s in range(COMPRESS_RATIO): - prev_abs = prev_start + s - if write_pos >= 2 * COMPRESS_RATIO - 1: - prev_row = state_row(prev_abs) - if prev_row >= 0: - pool_kv[s] = kv_state_flat[prev_row, :HEAD_DIM] - pool_score[s] = score_state_flat[prev_row, :HEAD_DIM] - - cur_abs = cur_start + s - cur_row = state_row(cur_abs) - if cur_row >= 0: - pool_kv[COMPRESS_RATIO + s] = kv_state_flat[cur_row, HEAD_DIM:OUT_DIM] - pool_score[COMPRESS_RATIO + s] = score_state_flat[cur_row, HEAD_DIM:OUT_DIM] - - for t in range(int(tensors["num_tokens"])): - pos = int(position_ids[t].item()) - if pos < prev_start or pos > write_pos: - continue - if pos < cur_start: - pool_slot = pos - prev_start - col0 = 0 - else: - pool_slot = COMPRESS_RATIO + pos - cur_start - col0 = HEAD_DIM - ape_slot = pos % COMPRESS_RATIO - pool_kv[pool_slot] = kv_proj[t, col0 : col0 + HEAD_DIM] - pool_score[pool_slot] = score_proj[t, col0 : col0 + HEAD_DIM] + ape[ape_slot, col0 : col0 + HEAD_DIM] - - init_slot = STATE_LEN - 1 - mi = pool_score[init_slot : init_slot + 1].clone() - li = torch.exp(mi - mi) - oi = pool_kv[init_slot : init_slot + 1].clone() - for slot_i in range(STATE_LEN - 1): - if slot_i < COMPRESS_RATIO and write_pos < 2 * COMPRESS_RATIO - 1: - continue - slot_score = pool_score[slot_i : slot_i + 1] - slot_kv = pool_kv[slot_i : slot_i + 1] - mi_next = torch.maximum(mi, slot_score) - alpha = torch.exp(mi - mi_next) - beta = torch.exp(slot_score - mi_next) - li = alpha * li + beta - oi = oi * alpha + slot_kv * beta - mi = mi_next - pooled = oi / li - inv_rms = torch.rsqrt(pooled.square().mean(dim=-1, keepdim=True) + EPS) - normed = pooled * inv_rms * norm_w.float().view(1, HEAD_DIM) - rope_pair = normed[..., NOPE_HEAD_DIM:HEAD_DIM].unflatten(-1, (-1, 2)) - rope_even = rope_pair[..., 0] - rope_odd = rope_pair[..., 1] - cmp_pos = write_pos + 1 - COMPRESS_RATIO - cos = tensors["freqs_cos"][cmp_pos : cmp_pos + 1, 0 : ROPE_HEAD_DIM // 2].float() - sin = tensors["freqs_sin"][cmp_pos : cmp_pos + 1, 0 : ROPE_HEAD_DIM // 2].float() - rot_even = rope_even * cos - rope_odd * sin - rot_odd = rope_even * sin + rope_odd * cos - normed[:, NOPE_HEAD_DIM:HEAD_DIM] = torch.stack([rot_even, rot_odd], dim=-1).flatten(-2) - cache_rows[dst_row] = normed.to(torch.bfloat16)[0] - - for t in range(int(tensors["num_tokens"])): - pos = int(tensors["position_ids"][t].item()) - dst_row = int(tensors["state_slot_mapping"][t].item()) - if dst_row < 0: - continue - ape_slot = pos % COMPRESS_RATIO - kv_state_flat[dst_row] = kv_proj[t] - score_state_flat[dst_row] = score_proj[t] + tensors["ape"][ape_slot] - tensors["cmp_kv"][:] = cmp_kv - tensors["kv_state"][:] = kv_state_flat.view_as(tensors["kv_state"]) - tensors["score_state"][:] = score_state_flat.view_as(tensors["score_state"]) - - -@pl.jit -def prefill_compressor_ratio4_test( - x: pl.Tensor[[T, D], pl.BF16], - kv_state: pl.InOut[pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32]], - score_state: pl.InOut[pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32]], - compress_state_block_table: pl.Tensor[[CSA_STATE_MAX_BLOCKS], pl.INT32], - wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], - wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], - ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], - norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], - freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], - cmp_kv: pl.InOut[pl.Tensor[[PREFILL_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], - position_ids: pl.Tensor[[T], pl.INT32], - num_tokens: pl.Scalar[pl.INT32], - cmp_slot_mapping: pl.Tensor[[T], pl.INT64], - state_slot_mapping: pl.Tensor[[T], pl.INT64], -): - return prefill_compressor_ratio4( - x, kv_state, score_state, compress_state_block_table, wkv, wgate, ape, norm_w, freqs_cos, freqs_sin, - cmp_kv, position_ids, num_tokens, cmp_slot_mapping, state_slot_mapping, - ) - - -def build_tensor_specs(start_pos: int = START_POS): - import torch - from golden import ScalarSpec, TensorSpec - from rope_tables import build_deepseek_v4_rope_tables - - shared_freqs_cos, shared_freqs_sin = build_deepseek_v4_rope_tables(M, COMPRESS_RATIO, dtype=torch.bfloat16) - - if start_pos < 0 or start_pos + T > MAX_SEQ_LEN: - raise ValueError(f"start_pos must satisfy 0 <= start_pos <= {MAX_SEQ_LEN - T}, got {start_pos}") - - def init_compress_state_block_table(): - table = torch.full((CSA_STATE_MAX_BLOCKS,), -1, dtype=torch.int32) - for block in range(CSA_STATE_MAX_BLOCKS): - table[block] = (block * 17 + 3) % CSA_STATE_MAX_BLOCKS - return table - def state_row(abs_pos): - if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: - return -1 - table = init_compress_state_block_table() - block = abs_pos // CSA_STATE_BLOCK_SIZE - intra = abs_pos % CSA_STATE_BLOCK_SIZE - return int(table[block].item()) * CSA_STATE_BLOCK_SIZE + intra - def init_x(): - return ((torch.rand(T, D) - 0.5) * 0.1).to(torch.bfloat16) - def init_state(): - state = torch.zeros(CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, OUT_DIM) - flat = state.view(-1, OUT_DIM) - for abs_pos in range(max(0, start_pos - STATE_LEN), start_pos): - row = state_row(abs_pos) - if row >= 0: - flat[row] = (torch.rand(OUT_DIM) - 0.5) * 0.05 - return state - # Calibrated to the real DeepSeek-V4-Flash CSA (ratio-4) compressor (mean l8/l32 of - # extract_weights_flash): zero-mean Gaussian BF16 weights at the measured std; the RMSNorm - # gamma centers near the measured mean (not ones / not uniform). Mirrors decode_compressor_ratio4. - def init_wkv(): - return torch.randn(OUT_DIM, D) * 0.0245 - def init_wgate(): - return torch.randn(OUT_DIM, D) * 0.0388 - def init_ape(): - return torch.randn(COMPRESS_RATIO, OUT_DIM) * 0.1243 - def init_norm_w(): - return 0.9666 + 0.1929 * torch.randn(HEAD_DIM) - def init_freqs_cos(): - return shared_freqs_cos.clone() - def init_freqs_sin(): - return shared_freqs_sin.clone() - def init_cmp_kv(): - return torch.zeros(PREFILL_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM, dtype=torch.bfloat16) - def init_position_ids(): - return torch.arange(start_pos, start_pos + T, dtype=torch.int32) - def init_cmp_slot_mapping(): - mapping = torch.full((T,), -1, dtype=torch.int64) - for t in range(T): - pos = start_pos + t - if (pos + 1) % COMPRESS_RATIO == 0: - dst_row = (pos + 1) // COMPRESS_RATIO - 1 - if dst_row >= PREFILL_CMP_BLOCK_NUM * BLOCK_SIZE: - raise ValueError("fixture compressed slot exceeds standalone cmp_kv capacity") - mapping[t] = dst_row - return mapping - def init_state_slot_mapping(): - mapping = torch.full((T,), -1, dtype=torch.int64) - for t in range(T): - mapping[t] = state_row(start_pos + t) - return mapping - - return [ - TensorSpec("x", [T, D], torch.bfloat16, init_value=init_x), - TensorSpec("kv_state", [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, OUT_DIM], torch.float32, init_value=init_state, is_output=True), - TensorSpec("score_state", [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, OUT_DIM], torch.float32, init_value=init_state, is_output=True), - TensorSpec("compress_state_block_table", [CSA_STATE_MAX_BLOCKS], torch.int32, init_value=init_compress_state_block_table), - TensorSpec("wkv", [OUT_DIM, D], torch.bfloat16, init_value=init_wkv), - TensorSpec("wgate", [OUT_DIM, D], torch.bfloat16, init_value=init_wgate), - TensorSpec("ape", [COMPRESS_RATIO, OUT_DIM], torch.float32, init_value=init_ape), - TensorSpec("norm_w", [HEAD_DIM], torch.bfloat16, init_value=init_norm_w), - TensorSpec("freqs_cos", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_cos), - TensorSpec("freqs_sin", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_sin), - TensorSpec("cmp_kv", [PREFILL_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], torch.bfloat16, init_value=init_cmp_kv, is_output=True), - TensorSpec("position_ids", [T], torch.int32, init_value=init_position_ids), - ScalarSpec("num_tokens", torch.int32, T), - TensorSpec("cmp_slot_mapping", [T], torch.int64, init_value=init_cmp_slot_mapping), - TensorSpec("state_slot_mapping", [T], torch.int64, init_value=init_state_slot_mapping), - ] - - -if __name__ == "__main__": - import argparse - from golden import ratio_allclose, run_jit - - parser = argparse.ArgumentParser(description="Standalone token-major DeepSeek V4 prefill compressor ratio4 validation.") - parser.add_argument("-p", "--platform", type=str, default="a2a3", - choices=["a2a3", "a2a3sim", "a5", "a5sim"]) - parser.add_argument("-d", "--device", type=int, default=0) - parser.add_argument( - "--compile-only", - action="store_true", - default=False, - help="Compile/codegen only. This is also the implicit behavior on *sim platforms used by CI.", - ) - parser.add_argument("--start-pos", type=int, default=START_POS, - help="Fixture-only absolute position for token 0; lowered into position_ids and dense cmp_slot_mapping.") - parser.add_argument("--enable-l2-swimlane", action="store_true", default=False) - parser.add_argument("--dump-passes", action="store_true", default=False) - args = parser.parse_args() - - result = run_jit( - fn=prefill_compressor_ratio4_test, - specs=build_tensor_specs(args.start_pos), - golden_fn=golden_prefill_compressor_ratio4, - compile_cfg=dict(dump_passes=args.dump_passes), - runtime_cfg=dict(platform=args.platform, device_id=args.device, enable_l2_swimlane=args.enable_l2_swimlane), - compile_only=args.compile_only, - compare_fn={ - "kv_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), - "score_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), - "cmp_kv": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.0), - }, - ) - if not result.passed: - if result.error: - print(result.error) - raise SystemExit(1) diff --git a/models/deepseek/v4/prefill_fwd.py b/models/deepseek/v4/prefill_fwd.py index de83e63a..6a8508c2 100644 --- a/models/deepseek/v4/prefill_fwd.py +++ b/models/deepseek/v4/prefill_fwd.py @@ -57,7 +57,7 @@ build_tensor_specs as build_moe_tensor_specs, moe, ) -from config import FLASH as MODEL_CONFIG +from config import FLASH as MODEL_CONFIG, PREFILL_BATCH as B from prefill_attention_swa import ( build_tensor_specs as build_swa_attention_tensor_specs, prefill_attention_swa, @@ -68,6 +68,7 @@ HCA_STATE_BLOCK_SIZE, HCA_STATE_MAX_BLOCKS, MAIN_OUT_DIM as HCA_MAIN_OUT_DIM, + MAIN_STATE_DIM as HCA_MAIN_STATE_DIM, build_tensor_specs as build_hca_attention_tensor_specs, prefill_attention_hca, ) @@ -85,10 +86,12 @@ IDX_HEAD_DIM, IDX_N_HEADS, INNER_OUT_DIM, + INNER_STATE_DIM, INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_MAX_BLOCKS, MAIN_OUT_DIM as CSA_MAIN_OUT_DIM, + MAIN_STATE_DIM as CSA_MAIN_STATE_DIM, MAX_SEQ_LEN, O_GROUPS, O_GROUP_IN, @@ -142,15 +145,15 @@ # CSA-compact stacked weights (sliced by the CSA order index 0..20). CSA_LAYER_STACKED_NAMES = [ "csa_cmp_wkv", "csa_cmp_wgate", "csa_cmp_ape", "csa_cmp_norm_w", - "csa_cmp_kv_state", "csa_cmp_score_state", + "csa_compress_state", "csa_hadamard_idx", "csa_idx_wq_b", "csa_idx_wq_b_scale", "csa_weights_proj", "csa_inner_wkv", "csa_inner_wgate", "csa_inner_ape", "csa_inner_norm_w", - "csa_inner_kv_state", "csa_inner_score_state", "idx_kv_cache", "idx_kv_scale", + "csa_inner_compress_state", "idx_kv_cache", "idx_kv_scale", ] # HCA-compact stacked weights (sliced by the HCA order index 0..19). HCA_LAYER_STACKED_NAMES = [ "hca_cmp_wkv", "hca_cmp_wgate", "hca_cmp_ape", "hca_cmp_norm_w", - "hca_cmp_kv_state", "hca_cmp_score_state", + "hca_compress_state", ] # Replicated once and passed whole to every layer (block tables are smoke zeros; # slot mappings depend only on token position + a fixed per-kind compress ratio, @@ -170,9 +173,9 @@ # tensors (re-bound each dispatch) rather than device-resident. CACHE_NAMES = { "kv_cache", "cmp_kv", - "hca_cmp_kv_state", "hca_cmp_score_state", - "csa_cmp_kv_state", "csa_cmp_score_state", - "csa_inner_kv_state", "csa_inner_score_state", "idx_kv_cache", "idx_kv_scale", + "hca_compress_state", + "csa_compress_state", + "csa_inner_compress_state", "idx_kv_cache", "idx_kv_scale", } # Static weight parameters to keep device-resident, sharded per rank. Every host @@ -204,7 +207,7 @@ # signature) is additionally read back once at the end for validation # (``is_output=True`` -> ``copy_stacked_from``). The other caches are plain # ``pl.Tensor`` inputs — read-only here, so resident but not read back. Only -# ``kv_cache`` is ``pl.InOut``; the rest (cmp_kv, *_kv_state, *_score_state, +# ``kv_cache`` is ``pl.InOut``; the rest (cmp_kv, *_compress_state, # idx_kv_cache) are read-only inputs. RESIDENT_CACHE_OUTPUT_NAMES = frozenset(["kv_cache"]) @@ -232,14 +235,12 @@ def prefill_fwd( hca_cmp_wgate: pl.Tensor[[HCA_NUM_LAYERS * HCA_MAIN_OUT_DIM, D], pl.BF16], hca_cmp_ape: pl.Tensor[[HCA_NUM_LAYERS * HCA_COMPRESS_RATIO, HCA_MAIN_OUT_DIM], pl.FP32], hca_cmp_norm_w: pl.Tensor[[HCA_NUM_LAYERS * HEAD_DIM], pl.BF16], - hca_cmp_kv_state: pl.Tensor[[HCA_NUM_LAYERS * HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], pl.FP32], - hca_cmp_score_state: pl.Tensor[[HCA_NUM_LAYERS * HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], pl.FP32], + hca_compress_state: pl.Tensor[[HCA_NUM_LAYERS * HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_STATE_DIM], pl.FP32], csa_cmp_wkv: pl.Tensor[[CSA_NUM_LAYERS * CSA_MAIN_OUT_DIM, D], pl.BF16], csa_cmp_wgate: pl.Tensor[[CSA_NUM_LAYERS * CSA_MAIN_OUT_DIM, D], pl.BF16], csa_cmp_ape: pl.Tensor[[CSA_NUM_LAYERS * CSA_COMPRESS_RATIO, CSA_MAIN_OUT_DIM], pl.FP32], csa_cmp_norm_w: pl.Tensor[[CSA_NUM_LAYERS * HEAD_DIM], pl.BF16], - csa_cmp_kv_state: pl.Tensor[[CSA_NUM_LAYERS * CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32], - csa_cmp_score_state: pl.Tensor[[CSA_NUM_LAYERS * CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32], + csa_compress_state: pl.Tensor[[CSA_NUM_LAYERS * CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_STATE_DIM], pl.FP32], csa_hadamard_idx: pl.Tensor[[CSA_NUM_LAYERS * IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16], csa_idx_wq_b: pl.Tensor[[CSA_NUM_LAYERS * Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8], csa_idx_wq_b_scale: pl.Tensor[[CSA_NUM_LAYERS * IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32], @@ -248,18 +249,17 @@ def prefill_fwd( csa_inner_wgate: pl.Tensor[[CSA_NUM_LAYERS * INNER_OUT_DIM, D], pl.BF16], csa_inner_ape: pl.Tensor[[CSA_NUM_LAYERS * CSA_COMPRESS_RATIO, INNER_OUT_DIM], pl.FP32], csa_inner_norm_w: pl.Tensor[[CSA_NUM_LAYERS * IDX_HEAD_DIM], pl.BF16], - csa_inner_kv_state: pl.Tensor[[CSA_NUM_LAYERS * INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], - csa_inner_score_state: pl.Tensor[[CSA_NUM_LAYERS * INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], + csa_inner_compress_state: pl.Tensor[[CSA_NUM_LAYERS * INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32], idx_kv_cache: pl.Tensor[[CSA_NUM_LAYERS * PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8], idx_kv_scale: pl.Tensor[[CSA_NUM_LAYERS * PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32], - hca_compress_state_block_table: pl.Tensor[[HCA_STATE_MAX_BLOCKS], pl.INT32], - csa_compress_state_block_table: pl.Tensor[[CSA_STATE_MAX_BLOCKS], pl.INT32], - csa_inner_compress_state_block_table: pl.Tensor[[INNER_STATE_MAX_BLOCKS], pl.INT32], + hca_compress_state_block_table: pl.Tensor[[B, HCA_STATE_MAX_BLOCKS], pl.INT32], + csa_compress_state_block_table: pl.Tensor[[B, CSA_STATE_MAX_BLOCKS], pl.INT32], + csa_inner_compress_state_block_table: pl.Tensor[[B, INNER_STATE_MAX_BLOCKS], pl.INT32], freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], ori_block_table: pl.Tensor[[SPARSE_ORI_MAX_BLOCKS], pl.INT32], cmp_block_table: pl.Tensor[[SPARSE_CMP_MAX_BLOCKS], pl.INT32], - idx_block_table: pl.Tensor[[IDX_CACHE_MAX_BLOCKS], pl.INT32], + idx_block_table: pl.Tensor[[B, IDX_CACHE_MAX_BLOCKS], pl.INT32], ori_slot_mapping: pl.Tensor[[T], pl.INT64], position_ids: pl.Tensor[[T], pl.INT32], input_ids: pl.Tensor[[T], pl.INT64], @@ -453,8 +453,7 @@ def prefill_fwd( csa_cmp_wgate_csa: pl.Tensor[[CSA_MAIN_OUT_DIM, D], pl.BF16] = pl.slice(csa_cmp_wgate, [CSA_MAIN_OUT_DIM, D], [loop_i * CSA_MAIN_OUT_DIM, 0]) csa_cmp_ape_csa: pl.Tensor[[CSA_COMPRESS_RATIO, CSA_MAIN_OUT_DIM], pl.FP32] = pl.slice(csa_cmp_ape, [CSA_COMPRESS_RATIO, CSA_MAIN_OUT_DIM], [loop_i * CSA_COMPRESS_RATIO, 0]) csa_cmp_norm_w_csa: pl.Tensor[[HEAD_DIM], pl.BF16] = pl.slice(csa_cmp_norm_w, [HEAD_DIM], [loop_i * HEAD_DIM]) - csa_cmp_kv_state_csa: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32] = pl.slice(csa_cmp_kv_state, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], [loop_i * CSA_STATE_BLOCK_NUM, 0, 0]) - csa_cmp_score_state_csa: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32] = pl.slice(csa_cmp_score_state, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], [loop_i * CSA_STATE_BLOCK_NUM, 0, 0]) + csa_compress_state_csa: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_STATE_DIM], pl.FP32] = pl.slice(csa_compress_state, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_STATE_DIM], [loop_i * CSA_STATE_BLOCK_NUM, 0, 0]) csa_hadamard_idx_csa: pl.Tensor[[IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16] = pl.slice(csa_hadamard_idx, [IDX_HEAD_DIM, IDX_HEAD_DIM], [loop_i * IDX_HEAD_DIM, 0]) csa_idx_wq_b_csa: pl.Tensor[[Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8] = pl.slice(csa_idx_wq_b, [Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], [loop_i * Q_LORA, 0]) csa_idx_wq_b_scale_csa: pl.Tensor[[IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32] = pl.slice(csa_idx_wq_b_scale, [IDX_N_HEADS * IDX_HEAD_DIM], [loop_i * IDX_N_HEADS * IDX_HEAD_DIM]) @@ -463,8 +462,7 @@ def prefill_fwd( csa_inner_wgate_csa: pl.Tensor[[INNER_OUT_DIM, D], pl.BF16] = pl.slice(csa_inner_wgate, [INNER_OUT_DIM, D], [loop_i * INNER_OUT_DIM, 0]) csa_inner_ape_csa: pl.Tensor[[CSA_COMPRESS_RATIO, INNER_OUT_DIM], pl.FP32] = pl.slice(csa_inner_ape, [CSA_COMPRESS_RATIO, INNER_OUT_DIM], [loop_i * CSA_COMPRESS_RATIO, 0]) csa_inner_norm_w_csa: pl.Tensor[[IDX_HEAD_DIM], pl.BF16] = pl.slice(csa_inner_norm_w, [IDX_HEAD_DIM], [loop_i * IDX_HEAD_DIM]) - csa_inner_kv_state_csa: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32] = pl.slice(csa_inner_kv_state, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], [loop_i * INNER_STATE_BLOCK_NUM, 0, 0]) - csa_inner_score_state_csa: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32] = pl.slice(csa_inner_score_state, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], [loop_i * INNER_STATE_BLOCK_NUM, 0, 0]) + csa_inner_compress_state_csa: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32] = pl.slice(csa_inner_compress_state, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], [loop_i * INNER_STATE_BLOCK_NUM, 0, 0]) kv_cache_csa: pl.Tensor[[CSA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16] = pl.slice(kv_cache, [CSA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], [csa_layer * CSA_ORI_BLOCK_NUM, 0, 0, 0]) cmp_kv_csa: pl.Tensor[[CSA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16] = pl.slice(cmp_kv, [CSA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], [csa_layer * CSA_CMP_BLOCK_NUM, 0, 0, 0]) idx_kv_cache_csa: pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8] = pl.slice(idx_kv_cache, [PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], [loop_i * PREFILL_IDX_BLOCK_NUM, 0, 0, 0]) @@ -501,11 +499,11 @@ def prefill_fwd( wq_a_csa, wq_b_csa, wq_b_scale_csa, wkv_csa, gamma_cq_csa, gamma_ckv_csa, freqs_cos, freqs_sin, csa_cmp_wkv_csa, csa_cmp_wgate_csa, csa_cmp_ape_csa, csa_cmp_norm_w_csa, - csa_cmp_kv_state_csa, csa_cmp_score_state_csa, csa_compress_state_block_table, + csa_compress_state_csa, csa_compress_state_block_table, csa_hadamard_idx_csa, csa_idx_wq_b_csa, csa_idx_wq_b_scale_csa, csa_weights_proj_csa, csa_inner_wkv_csa, csa_inner_wgate_csa, csa_inner_ape_csa, csa_inner_norm_w_csa, - csa_inner_kv_state_csa, csa_inner_score_state_csa, csa_inner_compress_state_block_table, + csa_inner_compress_state_csa, csa_inner_compress_state_block_table, kv_cache_csa, ori_block_table, ori_slot_mapping, cmp_kv_csa, cmp_block_table, idx_kv_cache_csa, idx_kv_scale_csa, idx_block_table, position_ids, csa_cmp_slot_mapping, csa_idx_slot_mapping, @@ -543,8 +541,7 @@ def prefill_fwd( hca_cmp_wgate_hca: pl.Tensor[[HCA_MAIN_OUT_DIM, D], pl.BF16] = pl.slice(hca_cmp_wgate, [HCA_MAIN_OUT_DIM, D], [loop_i * HCA_MAIN_OUT_DIM, 0]) hca_cmp_ape_hca: pl.Tensor[[HCA_COMPRESS_RATIO, HCA_MAIN_OUT_DIM], pl.FP32] = pl.slice(hca_cmp_ape, [HCA_COMPRESS_RATIO, HCA_MAIN_OUT_DIM], [loop_i * HCA_COMPRESS_RATIO, 0]) hca_cmp_norm_w_hca: pl.Tensor[[HEAD_DIM], pl.BF16] = pl.slice(hca_cmp_norm_w, [HEAD_DIM], [loop_i * HEAD_DIM]) - hca_cmp_kv_state_hca: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], pl.FP32] = pl.slice(hca_cmp_kv_state, [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], [loop_i * HCA_STATE_BLOCK_NUM, 0, 0]) - hca_cmp_score_state_hca: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], pl.FP32] = pl.slice(hca_cmp_score_state, [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], [loop_i * HCA_STATE_BLOCK_NUM, 0, 0]) + hca_compress_state_hca: pl.Tensor[[HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_STATE_DIM], pl.FP32] = pl.slice(hca_compress_state, [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_STATE_DIM], [loop_i * HCA_STATE_BLOCK_NUM, 0, 0]) kv_cache_hca: pl.Tensor[[CSA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16] = pl.slice(kv_cache, [CSA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], [hca_layer * CSA_ORI_BLOCK_NUM, 0, 0, 0]) cmp_kv_hca: pl.Tensor[[CSA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16] = pl.slice(cmp_kv, [CSA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], [hca_layer * CSA_CMP_BLOCK_NUM, 0, 0, 0]) attn_sink_hca: pl.Tensor[[H], pl.FP32] = pl.slice(attn_sink, [H], [hca_layer * H]) @@ -578,7 +575,7 @@ def prefill_fwd( wq_a_hca, wq_b_hca, wq_b_scale_hca, wkv_hca, gamma_cq_hca, gamma_ckv_hca, freqs_cos, freqs_sin, hca_cmp_wkv_hca, hca_cmp_wgate_hca, hca_cmp_ape_hca, hca_cmp_norm_w_hca, - hca_cmp_kv_state_hca, hca_cmp_score_state_hca, hca_compress_state_block_table, + hca_compress_state_hca, hca_compress_state_block_table, kv_cache_hca, ori_slot_mapping, ori_block_table, cmp_kv_hca, cmp_block_table, position_ids, hca_cmp_slot_mapping, hca_state_slot_mapping, @@ -618,8 +615,7 @@ def prefill_fwd( csa_cmp_wgate_last: pl.Tensor[[CSA_MAIN_OUT_DIM, D], pl.BF16] = pl.slice(csa_cmp_wgate, [CSA_MAIN_OUT_DIM, D], [csa_order_last * CSA_MAIN_OUT_DIM, 0]) csa_cmp_ape_last: pl.Tensor[[CSA_COMPRESS_RATIO, CSA_MAIN_OUT_DIM], pl.FP32] = pl.slice(csa_cmp_ape, [CSA_COMPRESS_RATIO, CSA_MAIN_OUT_DIM], [csa_order_last * CSA_COMPRESS_RATIO, 0]) csa_cmp_norm_w_last: pl.Tensor[[HEAD_DIM], pl.BF16] = pl.slice(csa_cmp_norm_w, [HEAD_DIM], [csa_order_last * HEAD_DIM]) - csa_cmp_kv_state_last: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32] = pl.slice(csa_cmp_kv_state, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], [csa_order_last * CSA_STATE_BLOCK_NUM, 0, 0]) - csa_cmp_score_state_last: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32] = pl.slice(csa_cmp_score_state, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], [csa_order_last * CSA_STATE_BLOCK_NUM, 0, 0]) + csa_compress_state_last: pl.Tensor[[CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_STATE_DIM], pl.FP32] = pl.slice(csa_compress_state, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_STATE_DIM], [csa_order_last * CSA_STATE_BLOCK_NUM, 0, 0]) csa_hadamard_idx_last: pl.Tensor[[IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16] = pl.slice(csa_hadamard_idx, [IDX_HEAD_DIM, IDX_HEAD_DIM], [csa_order_last * IDX_HEAD_DIM, 0]) csa_idx_wq_b_last: pl.Tensor[[Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8] = pl.slice(csa_idx_wq_b, [Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], [csa_order_last * Q_LORA, 0]) csa_idx_wq_b_scale_last: pl.Tensor[[IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32] = pl.slice(csa_idx_wq_b_scale, [IDX_N_HEADS * IDX_HEAD_DIM], [csa_order_last * IDX_N_HEADS * IDX_HEAD_DIM]) @@ -628,8 +624,7 @@ def prefill_fwd( csa_inner_wgate_last: pl.Tensor[[INNER_OUT_DIM, D], pl.BF16] = pl.slice(csa_inner_wgate, [INNER_OUT_DIM, D], [csa_order_last * INNER_OUT_DIM, 0]) csa_inner_ape_last: pl.Tensor[[CSA_COMPRESS_RATIO, INNER_OUT_DIM], pl.FP32] = pl.slice(csa_inner_ape, [CSA_COMPRESS_RATIO, INNER_OUT_DIM], [csa_order_last * CSA_COMPRESS_RATIO, 0]) csa_inner_norm_w_last: pl.Tensor[[IDX_HEAD_DIM], pl.BF16] = pl.slice(csa_inner_norm_w, [IDX_HEAD_DIM], [csa_order_last * IDX_HEAD_DIM]) - csa_inner_kv_state_last: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32] = pl.slice(csa_inner_kv_state, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], [csa_order_last * INNER_STATE_BLOCK_NUM, 0, 0]) - csa_inner_score_state_last: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32] = pl.slice(csa_inner_score_state, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], [csa_order_last * INNER_STATE_BLOCK_NUM, 0, 0]) + csa_inner_compress_state_last: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32] = pl.slice(csa_inner_compress_state, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], [csa_order_last * INNER_STATE_BLOCK_NUM, 0, 0]) kv_cache_last: pl.Tensor[[CSA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16] = pl.slice(kv_cache, [CSA_ORI_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], [csa_layer_last * CSA_ORI_BLOCK_NUM, 0, 0, 0]) cmp_kv_last: pl.Tensor[[CSA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16] = pl.slice(cmp_kv, [CSA_CMP_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], [csa_layer_last * CSA_CMP_BLOCK_NUM, 0, 0, 0]) idx_kv_cache_last: pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8] = pl.slice(idx_kv_cache, [PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], [csa_order_last * PREFILL_IDX_BLOCK_NUM, 0, 0, 0]) @@ -666,11 +661,11 @@ def prefill_fwd( wq_a_last, wq_b_last, wq_b_scale_last, wkv_last, gamma_cq_last, gamma_ckv_last, freqs_cos, freqs_sin, csa_cmp_wkv_last, csa_cmp_wgate_last, csa_cmp_ape_last, csa_cmp_norm_w_last, - csa_cmp_kv_state_last, csa_cmp_score_state_last, csa_compress_state_block_table, + csa_compress_state_last, csa_compress_state_block_table, csa_hadamard_idx_last, csa_idx_wq_b_last, csa_idx_wq_b_scale_last, csa_weights_proj_last, csa_inner_wkv_last, csa_inner_wgate_last, csa_inner_ape_last, csa_inner_norm_w_last, - csa_inner_kv_state_last, csa_inner_score_state_last, csa_inner_compress_state_block_table, + csa_inner_compress_state_last, csa_inner_compress_state_block_table, kv_cache_last, ori_block_table, ori_slot_mapping, cmp_kv_last, cmp_block_table, idx_kv_cache_last, idx_kv_scale_last, idx_block_table, position_ids, csa_cmp_slot_mapping, csa_idx_slot_mapping, @@ -722,14 +717,12 @@ def l3_prefill_fwd( hca_cmp_wgate: pl.Tensor[[N_RANKS, HCA_NUM_LAYERS * HCA_MAIN_OUT_DIM, D], pl.BF16], hca_cmp_ape: pl.Tensor[[N_RANKS, HCA_NUM_LAYERS * HCA_COMPRESS_RATIO, HCA_MAIN_OUT_DIM], pl.FP32], hca_cmp_norm_w: pl.Tensor[[N_RANKS, HCA_NUM_LAYERS * HEAD_DIM], pl.BF16], - hca_cmp_kv_state: pl.Tensor[[N_RANKS, HCA_NUM_LAYERS * HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], pl.FP32], - hca_cmp_score_state: pl.Tensor[[N_RANKS, HCA_NUM_LAYERS * HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], pl.FP32], + hca_compress_state: pl.Tensor[[N_RANKS, HCA_NUM_LAYERS * HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_STATE_DIM], pl.FP32], csa_cmp_wkv: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * CSA_MAIN_OUT_DIM, D], pl.BF16], csa_cmp_wgate: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * CSA_MAIN_OUT_DIM, D], pl.BF16], csa_cmp_ape: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * CSA_COMPRESS_RATIO, CSA_MAIN_OUT_DIM], pl.FP32], csa_cmp_norm_w: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * HEAD_DIM], pl.BF16], - csa_cmp_kv_state: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32], - csa_cmp_score_state: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32], + csa_compress_state: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_STATE_DIM], pl.FP32], csa_hadamard_idx: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16], csa_idx_wq_b: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8], csa_idx_wq_b_scale: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32], @@ -738,18 +731,17 @@ def l3_prefill_fwd( csa_inner_wgate: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * INNER_OUT_DIM, D], pl.BF16], csa_inner_ape: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * CSA_COMPRESS_RATIO, INNER_OUT_DIM], pl.FP32], csa_inner_norm_w: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * IDX_HEAD_DIM], pl.BF16], - csa_inner_kv_state: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], - csa_inner_score_state: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], + csa_inner_compress_state: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32], idx_kv_cache: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8], idx_kv_scale: pl.Tensor[[N_RANKS, CSA_NUM_LAYERS * PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32], - hca_compress_state_block_table: pl.Tensor[[N_RANKS, HCA_STATE_MAX_BLOCKS], pl.INT32], - csa_compress_state_block_table: pl.Tensor[[N_RANKS, CSA_STATE_MAX_BLOCKS], pl.INT32], - csa_inner_compress_state_block_table: pl.Tensor[[N_RANKS, INNER_STATE_MAX_BLOCKS], pl.INT32], + hca_compress_state_block_table: pl.Tensor[[N_RANKS, B, HCA_STATE_MAX_BLOCKS], pl.INT32], + csa_compress_state_block_table: pl.Tensor[[N_RANKS, B, CSA_STATE_MAX_BLOCKS], pl.INT32], + csa_inner_compress_state_block_table: pl.Tensor[[N_RANKS, B, INNER_STATE_MAX_BLOCKS], pl.INT32], freqs_cos: pl.Tensor[[N_RANKS, MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], freqs_sin: pl.Tensor[[N_RANKS, MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], ori_block_table: pl.Tensor[[N_RANKS, SPARSE_ORI_MAX_BLOCKS], pl.INT32], cmp_block_table: pl.Tensor[[N_RANKS, SPARSE_CMP_MAX_BLOCKS], pl.INT32], - idx_block_table: pl.Tensor[[N_RANKS, IDX_CACHE_MAX_BLOCKS], pl.INT32], + idx_block_table: pl.Tensor[[N_RANKS, B, IDX_CACHE_MAX_BLOCKS], pl.INT32], ori_slot_mapping: pl.Tensor[[N_RANKS, T], pl.INT64], position_ids: pl.Tensor[[N_RANKS, T], pl.INT32], input_ids: pl.Tensor[[N_RANKS, T], pl.INT64], @@ -809,12 +801,12 @@ def l3_prefill_fwd( wq_a[r], wq_b[r], wq_b_scale[r], wkv[r], gamma_cq[r], gamma_ckv[r], kv_cache[r], attn_sink[r], wo_a[r], wo_b[r], wo_b_scale[r], cmp_kv[r], hca_cmp_wkv[r], hca_cmp_wgate[r], hca_cmp_ape[r], hca_cmp_norm_w[r], - hca_cmp_kv_state[r], hca_cmp_score_state[r], + hca_compress_state[r], csa_cmp_wkv[r], csa_cmp_wgate[r], csa_cmp_ape[r], csa_cmp_norm_w[r], - csa_cmp_kv_state[r], csa_cmp_score_state[r], + csa_compress_state[r], csa_hadamard_idx[r], csa_idx_wq_b[r], csa_idx_wq_b_scale[r], csa_weights_proj[r], csa_inner_wkv[r], csa_inner_wgate[r], csa_inner_ape[r], csa_inner_norm_w[r], - csa_inner_kv_state[r], csa_inner_score_state[r], idx_kv_cache[r], idx_kv_scale[r], + csa_inner_compress_state[r], idx_kv_cache[r], idx_kv_scale[r], hca_compress_state_block_table[r], csa_compress_state_block_table[r], csa_inner_compress_state_block_table[r], freqs_cos[r], freqs_sin[r], @@ -913,6 +905,7 @@ def init_value(): return ranked(out) if name in ("ori_block_table", "cmp_block_table", "idx_block_table"): out = torch.arange(spec.shape[-1], dtype=spec.dtype) + out = out.expand(*spec.shape[1:]).contiguous() return ranked(out) # Any remaining shared metadata: smoke zeros. return torch.zeros(list(spec.shape), dtype=spec.dtype) @@ -982,15 +975,13 @@ def _make_final_norm_spec(name): "hca_cmp_wgate", "hca_cmp_ape", "hca_cmp_norm_w", - "hca_cmp_kv_state", - "hca_cmp_score_state", + "hca_compress_state", "hca_compress_state_block_table", "csa_cmp_wkv", "csa_cmp_wgate", "csa_cmp_ape", "csa_cmp_norm_w", - "csa_cmp_kv_state", - "csa_cmp_score_state", + "csa_compress_state", "csa_compress_state_block_table", "csa_hadamard_idx", "csa_idx_wq_b", @@ -1000,8 +991,7 @@ def _make_final_norm_spec(name): "csa_inner_wgate", "csa_inner_ape", "csa_inner_norm_w", - "csa_inner_kv_state", - "csa_inner_score_state", + "csa_inner_compress_state", "csa_inner_compress_state_block_table", "kv_cache", "ori_block_table", @@ -1125,15 +1115,13 @@ def kind_specs(build_fn): ("hca_cmp_wgate", hca["cmp_wgate"]), ("hca_cmp_ape", hca["cmp_ape"]), ("hca_cmp_norm_w", hca["cmp_norm_w"]), - ("hca_cmp_kv_state", hca["cmp_kv_state"]), - ("hca_cmp_score_state", hca["cmp_score_state"]), + ("hca_compress_state", hca["compress_state"]), ("hca_compress_state_block_table", hca["compress_state_block_table"]), ("csa_cmp_wkv", csa["cmp_wkv"]), ("csa_cmp_wgate", csa["cmp_wgate"]), ("csa_cmp_ape", csa["cmp_ape"]), ("csa_cmp_norm_w", csa["cmp_norm_w"]), - ("csa_cmp_kv_state", csa["cmp_kv_state"]), - ("csa_cmp_score_state", csa["cmp_score_state"]), + ("csa_compress_state", csa["compress_state"]), ("csa_compress_state_block_table", csa["compress_state_block_table"]), ("csa_hadamard_idx", csa["hadamard_idx"]), ("csa_idx_wq_b", csa["idx_wq_b"]), @@ -1143,8 +1131,7 @@ def kind_specs(build_fn): ("csa_inner_wgate", csa["inner_wgate"]), ("csa_inner_ape", csa["inner_ape"]), ("csa_inner_norm_w", csa["inner_norm_w"]), - ("csa_inner_kv_state", csa["inner_kv_state"]), - ("csa_inner_score_state", csa["inner_score_state"]), + ("csa_inner_compress_state", csa["inner_compress_state"]), ("csa_inner_compress_state_block_table", csa["inner_compress_state_block_table"]), ("kv_cache", active["kv_cache"]), ("ori_block_table", active.get("ori_block_table", swa.get("block_table"))), @@ -1234,12 +1221,12 @@ def build_tensor_specs(start_pos=0, num_tokens=T): "wq_a", "wq_b", "wq_b_scale", "wkv", "gamma_cq", "gamma_ckv", "kv_cache", "attn_sink", "wo_a", "wo_b", "wo_b_scale", "cmp_kv", "hca_cmp_wkv", "hca_cmp_wgate", "hca_cmp_ape", "hca_cmp_norm_w", - "hca_cmp_kv_state", "hca_cmp_score_state", + "hca_compress_state", "csa_cmp_wkv", "csa_cmp_wgate", "csa_cmp_ape", "csa_cmp_norm_w", - "csa_cmp_kv_state", "csa_cmp_score_state", + "csa_compress_state", "csa_hadamard_idx", "csa_idx_wq_b", "csa_idx_wq_b_scale", "csa_weights_proj", "csa_inner_wkv", "csa_inner_wgate", "csa_inner_ape", "csa_inner_norm_w", - "csa_inner_kv_state", "csa_inner_score_state", "idx_kv_cache", "idx_kv_scale", + "csa_inner_compress_state", "idx_kv_cache", "idx_kv_scale", "hca_compress_state_block_table", "csa_compress_state_block_table", "csa_inner_compress_state_block_table", "freqs_cos", "freqs_sin", diff --git a/models/deepseek/v4/prefill_indexer.py b/models/deepseek/v4/prefill_indexer.py index ee9d49f6..f9ec949f 100644 --- a/models/deepseek/v4/prefill_indexer.py +++ b/models/deepseek/v4/prefill_indexer.py @@ -24,6 +24,7 @@ PREFILL_IDX_BLOCK_NUM, ) from prefill_indexer_compressor import ( + COMPRESS_STATE_DIM as INNER_STATE_DIM, INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_MAX_BLOCKS, @@ -107,14 +108,11 @@ def prefill_indexer( wq_b: pl.Tensor[[Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8], wq_b_scale: pl.Tensor[[IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32], weights_proj: pl.Tensor[[D, IDX_N_HEADS], pl.BF16], - cos: pl.Tensor[[T, ROPE_HEAD_DIM // 2], pl.FP32], - sin: pl.Tensor[[T, ROPE_HEAD_DIM // 2], pl.FP32], freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], hadamard: pl.Tensor[[IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16], - inner_kv_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], - inner_score_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], - inner_compress_state_block_table: pl.Tensor[[INNER_STATE_MAX_BLOCKS], pl.INT32], + inner_compress_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32], + inner_compress_state_block_table: pl.Tensor[[B, INNER_STATE_MAX_BLOCKS], pl.INT32], inner_wkv: pl.Tensor[[INNER_OUT_DIM, D], pl.BF16], inner_wgate: pl.Tensor[[INNER_OUT_DIM, D], pl.BF16], inner_ape: pl.Tensor[[COMPRESS_RATIO, INNER_OUT_DIM], pl.FP32], @@ -122,7 +120,7 @@ def prefill_indexer( # C8 indexer cache: INT8 KV (quant-on-write) + per-position FP32 dequant scale; no bf16 cache. idx_kv_cache: pl.Out[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8]], idx_kv_scale: pl.Out[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32]], - idx_block_table: pl.Tensor[[IDX_CACHE_MAX_BLOCKS], pl.INT32], + idx_block_table: pl.Tensor[[B, IDX_CACHE_MAX_BLOCKS], pl.INT32], score: pl.Out[pl.Tensor[[T, INDEXER_SCORE_CAP], pl.FP32]], cmp_topk_indices: pl.Out[pl.Tensor[[T, IDX_TOPK], pl.INT32]], position_ids: pl.Tensor[[T], pl.INT32], @@ -156,8 +154,9 @@ def prefill_indexer( for idx in pl.spmd(T * IDX_N_HEADS // ROPE_ROW_BLOCK, name_hint="prefill_idx_qr_rope"): o0 = idx * ROPE_ROW_BLOCK token_idx = idx # ROPE_ROW_BLOCK == IDX_N_HEADS, so one task == one token - cos_b = cos[token_idx : token_idx + 1, 0 : ROPE_HEAD_DIM // 2] - sin_b = sin[token_idx : token_idx + 1, 0 : ROPE_HEAD_DIM // 2] + pos_t = pl.cast(pl.read(position_ids, [token_idx]), pl.INDEX) + cos_b = pl.cast(freqs_cos[pos_t : pos_t + 1, 0 : ROPE_HEAD_DIM // 2], target_type=pl.FP32) + sin_b = pl.cast(freqs_sin[pos_t : pos_t + 1, 0 : ROPE_HEAD_DIM // 2], target_type=pl.FP32) rope_ones = pl.full([ROPE_ROW_TILE, ROPE_HEAD_DIM], dtype=pl.FP32, value=1.0) rope_col = pl.col_expand_mul(rope_ones, pl.cast(pl.arange(0, [1, ROPE_HEAD_DIM], dtype=pl.INT32), target_type=pl.FP32)) rope_dup_f = pl.cast(pl.cast(pl.mul(rope_col, 0.5), target_type=pl.INT32, mode="trunc"), target_type=pl.FP32) @@ -225,7 +224,7 @@ def prefill_indexer( # === inner compressor: build the paged compressed index KV cache === prefill_indexer_compressor( x, - inner_kv_state, inner_score_state, inner_compress_state_block_table, + inner_compress_state, inner_compress_state_block_table, inner_wkv, inner_wgate, inner_ape, inner_norm_w, freqs_cos, freqs_sin, hadamard, idx_kv_cache, idx_kv_scale, idx_block_table, @@ -251,7 +250,7 @@ def prefill_indexer( for cb in pl.range(INDEXER_SCORE_BLOCKS): cache0 = cb * CACHE_TILE if max_visible > cache0: - idx_blk_id = pl.cast(pl.read(idx_block_table, [cache0 // BLOCK_SIZE]), pl.INDEX) + idx_blk_id = pl.cast(pl.read(idx_block_table, [0, cache0 // BLOCK_SIZE]), pl.INDEX) kv_row0 = idx_blk_id * BLOCK_SIZE + (cache0 % BLOCK_SIZE) # C8: the compressor stored this block as INT8 + a per-position dequant scale; read # both from the paged cache directly (no score-time re-quant). @@ -331,8 +330,7 @@ def golden_prefill_indexer_core(tensors): compressor_tensors = { "x": tensors["x"], "kv": torch.zeros(MAX_CMP_WRITES, IDX_HEAD_DIM, dtype=torch.bfloat16), - "kv_state": tensors["inner_kv_state"], - "score_state": tensors["inner_score_state"], + "compress_state": tensors["inner_compress_state"], "inner_compress_state_block_table": tensors["inner_compress_state_block_table"], "wkv": tensors["inner_wkv"], "wgate": tensors["inner_wgate"], @@ -375,8 +373,8 @@ def golden_prefill_indexer_core(tensors): wq_b = tensors["wq_b"] wq_b_scale = tensors["wq_b_scale"].float() hadamard = tensors["hadamard"].float() - cos = tensors["cos"].float().view(T, 1, -1) - sin = tensors["sin"].float().view(T, 1, -1) + cos = tensors["freqs_cos"][position_ids, : rd // 2].float().view(T, 1, -1) + sin = tensors["freqs_sin"][position_ids, : rd // 2].float().view(T, 1, -1) q_i32 = qr.to(torch.int32) @ wq_b.to(torch.int32) q = (q_i32.float() * qr_scale * wq_b_scale.view(1, -1)).view(T, IDX_N_HEADS, IDX_HEAD_DIM) q_pair = q[..., -rd:].unflatten(-1, (-1, 2)) @@ -394,7 +392,7 @@ def golden_prefill_indexer_core(tensors): scale_flat = tensors["idx_kv_scale"].float().reshape(PREFILL_IDX_BLOCK_NUM * BLOCK_SIZE, 1) idx_block_table = tensors["idx_block_table"] rows = [ - int(idx_block_table[c // BLOCK_SIZE].item()) * BLOCK_SIZE + (c % BLOCK_SIZE) + int(idx_block_table[0, c // BLOCK_SIZE].item()) * BLOCK_SIZE + (c % BLOCK_SIZE) for c in range(max_visible) ] kv_i8 = torch.stack([cache_flat_i8[r] for r in rows], dim=0).to(torch.int32) # [max_visible, dim] @@ -440,21 +438,18 @@ def prefill_indexer_test( wq_b: pl.Tensor[[Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8], wq_b_scale: pl.Tensor[[IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32], weights_proj: pl.Tensor[[D, IDX_N_HEADS], pl.BF16], - cos: pl.Tensor[[T, ROPE_HEAD_DIM // 2], pl.FP32], - sin: pl.Tensor[[T, ROPE_HEAD_DIM // 2], pl.FP32], freqs_cos: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], freqs_sin: pl.Tensor[[MAX_SEQ_LEN, ROPE_HEAD_DIM], pl.BF16], hadamard: pl.Tensor[[IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16], - inner_kv_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], - inner_score_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32], - inner_compress_state_block_table: pl.Tensor[[INNER_STATE_MAX_BLOCKS], pl.INT32], + inner_compress_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32], + inner_compress_state_block_table: pl.Tensor[[B, INNER_STATE_MAX_BLOCKS], pl.INT32], inner_wkv: pl.Tensor[[INNER_OUT_DIM, D], pl.BF16], inner_wgate: pl.Tensor[[INNER_OUT_DIM, D], pl.BF16], inner_ape: pl.Tensor[[COMPRESS_RATIO, INNER_OUT_DIM], pl.FP32], inner_norm_w: pl.Tensor[[INNER_HEAD_DIM], pl.BF16], idx_kv_cache: pl.InOut[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8]], idx_kv_scale: pl.InOut[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32]], - idx_block_table: pl.Tensor[[IDX_CACHE_MAX_BLOCKS], pl.INT32], + idx_block_table: pl.Tensor[[B, IDX_CACHE_MAX_BLOCKS], pl.INT32], score: pl.Out[pl.Tensor[[T, INDEXER_SCORE_CAP], pl.FP32]], topk_idxs: pl.Out[pl.Tensor[[T, INDEXER_SCORE_CAP], pl.INT32]], position_ids: pl.Tensor[[T], pl.INT32], @@ -465,8 +460,8 @@ def prefill_indexer_test( cmp_topk_indices = pl.create_tensor([T, IDX_TOPK], dtype=pl.INT32) prefill_indexer( x, qr, qr_scale, wq_b, wq_b_scale, weights_proj, - cos, sin, freqs_cos, freqs_sin, hadamard, - inner_kv_state, inner_score_state, inner_compress_state_block_table, + freqs_cos, freqs_sin, hadamard, + inner_compress_state, inner_compress_state_block_table, inner_wkv, inner_wgate, inner_ape, inner_norm_w, idx_kv_cache, idx_kv_scale, idx_block_table, score, cmp_topk_indices, @@ -511,7 +506,7 @@ def sim_fp8(W, block=128): def build_tensor_specs(start_pos: int = START_POS): import torch from golden import ScalarSpec, TensorSpec - from rope_tables import build_deepseek_v4_rope_tables, materialize_half_rope_tables + from rope_tables import build_deepseek_v4_rope_tables shared_freqs_cos, shared_freqs_sin = build_deepseek_v4_rope_tables(M, COMPRESS_RATIO, dtype=torch.bfloat16) @@ -529,9 +524,9 @@ def build_tensor_specs(start_pos: int = START_POS): raise ValueError(f"fixture generated {write_count} compressed writes, cap is {MAX_CMP_WRITES}") def init_inner_compress_state_block_table(): - table = torch.full((INNER_STATE_MAX_BLOCKS,), -1, dtype=torch.int32) + table = torch.full((B, INNER_STATE_MAX_BLOCKS), -1, dtype=torch.int32) for block in range(INNER_STATE_MAX_BLOCKS): - table[block] = (block * 17 + 3) % INNER_STATE_MAX_BLOCKS + table[0, block] = (block * 17 + 3) % INNER_STATE_MAX_BLOCKS return table def state_row(abs_pos): if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: @@ -539,7 +534,7 @@ def state_row(abs_pos): table = init_inner_compress_state_block_table() block = abs_pos // INNER_STATE_BLOCK_SIZE intra = abs_pos % INNER_STATE_BLOCK_SIZE - return int(table[block].item()) * INNER_STATE_BLOCK_SIZE + intra + return int(table[0, block].item()) * INNER_STATE_BLOCK_SIZE + intra def init_x(): return ((torch.rand(T, D) - 0.5) * 0.1).to(torch.bfloat16) def init_freqs_cos(): @@ -551,13 +546,14 @@ def init_hadamard(): while h.shape[0] < IDX_HEAD_DIM: h = torch.cat([torch.cat([h, h], dim=1), torch.cat([h, -h], dim=1)], dim=0) return (h * (IDX_HEAD_DIM ** -0.5)).to(torch.bfloat16) - def init_inner_state(): - state = torch.zeros(INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM) - flat = state.view(-1, INNER_OUT_DIM) + def init_inner_compress_state(): + state = torch.zeros(INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM) + flat = state.view(-1, INNER_STATE_DIM) for abs_pos in range(max(0, start_pos - INNER_STATE_LEN), start_pos): row = state_row(abs_pos) if row >= 0: - flat[row] = (torch.rand(INNER_OUT_DIM) - 0.5) * 0.05 + flat[row, 0:INNER_OUT_DIM] = (torch.rand(INNER_OUT_DIM) - 0.5) * 0.05 + flat[row, INNER_OUT_DIM:INNER_STATE_DIM] = (torch.rand(INNER_OUT_DIM) - 0.5) * 0.05 return state # Calibrated to the real DeepSeek-V4-Flash indexer inner compressor (mean l8/l32 of # extract_weights_flash): zero-mean Gaussian BF16 weights at the measured std; the RMSNorm @@ -599,15 +595,15 @@ def init_idx_kv_scale(): _build_idx_hist() return _idx_hist["scale"].clone() def init_idx_block_table(): - table = torch.full((IDX_CACHE_MAX_BLOCKS,), -1, dtype=torch.int32) + table = torch.full((B, IDX_CACHE_MAX_BLOCKS), -1, dtype=torch.int32) for block in range(IDX_CACHE_MAX_BLOCKS): - table[block] = block + table[0, block] = block return table def idx_row(cmp_slot): table = init_idx_block_table() block = cmp_slot // BLOCK_SIZE intra = cmp_slot % BLOCK_SIZE - phys_block = int(table[block].item()) + phys_block = int(table[0, block].item()) if phys_block < 0: return -1 return phys_block * BLOCK_SIZE + intra @@ -631,11 +627,6 @@ def init_inner_state_slot_mapping(): def init_weights_proj(): # weights_proj calibrated to the real DeepSeek-V4-Flash indexer weights projection. return torch.randn(D, IDX_N_HEADS) * 0.2313 - def init_cos(): - return materialize_half_rope_tables(shared_freqs_cos, shared_freqs_sin, init_position_ids().to(torch.int64))[0] - def init_sin(): - return materialize_half_rope_tables(shared_freqs_cos, shared_freqs_sin, init_position_ids().to(torch.int64))[1] - # idx wq_b uses the real MXFP8 grid (not a benign randn int8); qr is per-row int8 like the # runtime W8A8C16 activation path. wq_b_i8_T, wq_b_scale = gen_shared_weight((IDX_N_HEADS * IDX_HEAD_DIM, Q_LORA), dequant_std=0.108, chan_cv=0.56) @@ -649,21 +640,18 @@ def init_sin(): TensorSpec("wq_b", [Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], torch.int8, init_value=lambda: wq_b_i8), TensorSpec("wq_b_scale", [IDX_N_HEADS * IDX_HEAD_DIM], torch.float32, init_value=lambda: wq_b_scale), TensorSpec("weights_proj", [D, IDX_N_HEADS], torch.bfloat16, init_value=init_weights_proj), - TensorSpec("cos", [T, ROPE_HEAD_DIM // 2], torch.float32, init_value=init_cos), - TensorSpec("sin", [T, ROPE_HEAD_DIM // 2], torch.float32, init_value=init_sin), TensorSpec("freqs_cos", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_cos), TensorSpec("freqs_sin", [MAX_SEQ_LEN, ROPE_HEAD_DIM], torch.bfloat16, init_value=init_freqs_sin), TensorSpec("hadamard", [IDX_HEAD_DIM, IDX_HEAD_DIM], torch.bfloat16, init_value=init_hadamard), - TensorSpec("inner_kv_state", [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], torch.float32, init_value=init_inner_state), - TensorSpec("inner_score_state", [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], torch.float32, init_value=init_inner_state), - TensorSpec("inner_compress_state_block_table", [INNER_STATE_MAX_BLOCKS], torch.int32, init_value=init_inner_compress_state_block_table), + TensorSpec("inner_compress_state", [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], torch.float32, init_value=init_inner_compress_state), + TensorSpec("inner_compress_state_block_table", [B, INNER_STATE_MAX_BLOCKS], torch.int32, init_value=init_inner_compress_state_block_table), TensorSpec("inner_wkv", [INNER_OUT_DIM, D], torch.bfloat16, init_value=init_inner_wkv), TensorSpec("inner_wgate", [INNER_OUT_DIM, D], torch.bfloat16, init_value=init_inner_wgate), TensorSpec("inner_ape", [COMPRESS_RATIO, INNER_OUT_DIM], torch.float32, init_value=init_inner_ape), TensorSpec("inner_norm_w", [INNER_HEAD_DIM], torch.bfloat16, init_value=init_inner_norm_w), TensorSpec("idx_kv_cache", [PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, IDX_HEAD_DIM], torch.int8, init_value=init_idx_kv_cache, is_output=True), TensorSpec("idx_kv_scale", [PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1], torch.float32, init_value=init_idx_kv_scale, is_output=True), - TensorSpec("idx_block_table", [IDX_CACHE_MAX_BLOCKS], torch.int32, init_value=init_idx_block_table), + TensorSpec("idx_block_table", [B, IDX_CACHE_MAX_BLOCKS], torch.int32, init_value=init_idx_block_table), TensorSpec("score", [T, INDEXER_SCORE_CAP], torch.float32, is_output=True), TensorSpec("topk_idxs", [T, INDEXER_SCORE_CAP], torch.int32, is_output=True), TensorSpec("position_ids", [T], torch.int32, init_value=init_position_ids), diff --git a/models/deepseek/v4/prefill_indexer_compressor.py b/models/deepseek/v4/prefill_indexer_compressor.py index 8b928a46..38bb97a4 100644 --- a/models/deepseek/v4/prefill_indexer_compressor.py +++ b/models/deepseek/v4/prefill_indexer_compressor.py @@ -35,6 +35,7 @@ COFF = 1 + int(OVERLAP) OUT_DIM = COFF * HEAD_DIM STATE_LEN = COFF * COMPRESS_RATIO +COMPRESS_STATE_DIM = 2 * OUT_DIM IDX_CACHE_MAX_BLOCKS = PREFILL_IDX_MAX_BLOCKS @@ -63,9 +64,8 @@ @pl.jit.inline def prefill_indexer_compressor( x: pl.Tensor[[T, D], pl.BF16], - kv_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32], - score_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32], - inner_compress_state_block_table: pl.Tensor[[INNER_STATE_MAX_BLOCKS], pl.INT32], + compress_state: pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32], + inner_compress_state_block_table: pl.Tensor[[B, INNER_STATE_MAX_BLOCKS], pl.INT32], wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], @@ -76,7 +76,7 @@ def prefill_indexer_compressor( # C8 indexer cache: INT8 KV (quant-on-write) + per-position FP32 dequant scale; no bf16 cache. idx_kv_cache: pl.Out[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.INT8]], idx_kv_scale: pl.Out[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32]], - idx_block_table: pl.Tensor[[IDX_CACHE_MAX_BLOCKS], pl.INT32], + idx_block_table: pl.Tensor[[B, IDX_CACHE_MAX_BLOCKS], pl.INT32], position_ids: pl.Tensor[[T], pl.INT32], num_tokens: pl.Scalar[pl.INT32], idx_slot_mapping: pl.Tensor[[T], pl.INT64], @@ -84,8 +84,7 @@ def prefill_indexer_compressor( ): kv_proj_scratch = pl.create_tensor([T, OUT_DIM], dtype=pl.FP32) score_proj_scratch = pl.create_tensor([T, OUT_DIM], dtype=pl.FP32) - kv_state_flat = pl.reshape(kv_state, [INNER_STATE_BLOCK_NUM * INNER_STATE_BLOCK_SIZE, OUT_DIM]) - score_state_flat = pl.reshape(score_state, [INNER_STATE_BLOCK_NUM * INNER_STATE_BLOCK_SIZE, OUT_DIM]) + compress_state_flat = pl.reshape(compress_state, [INNER_STATE_BLOCK_NUM * INNER_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM]) idx_kv_cache_flat = pl.reshape(idx_kv_cache, [PREFILL_IDX_BLOCK_NUM * BLOCK_SIZE, HEAD_DIM]) idx_kv_scale_flat = pl.reshape(idx_kv_scale, [PREFILL_IDX_BLOCK_NUM * BLOCK_SIZE, 1]) pooled_kv = pl.create_tensor([MAX_CMP_WRITES, HEAD_DIM], dtype=pl.FP32) @@ -101,7 +100,7 @@ def prefill_indexer_compressor( x_tile = x[0:T, k0 : k0 + K_TILE] # Weights stored transposed [OUT_DIM, D] + b_trans=True -> DN2ZN load # (K-contiguous long bursts) instead of ND2NZ strided; mirrors the main - # compressor (prefill_compressor_ratio4) and the decode indexer compressor. + # compressor_ratio4 prefill mode and the decode indexer compressor. wkv_tile = wkv[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] wgate_tile = wgate[o0 : o0 + OUT_TILE, k0 : k0 + K_TILE] if k0 == 0: @@ -157,17 +156,17 @@ def prefill_indexer_compressor( if write_pos >= 2 * COMPRESS_RATIO - 1: prev_state_block = pl.cast(prev_abs // INNER_STATE_BLOCK_SIZE, pl.INDEX) prev_state_intra = pl.cast(prev_abs - prev_state_block * INNER_STATE_BLOCK_SIZE, pl.INDEX) - prev_phys_block_raw = pl.read(inner_compress_state_block_table, [prev_state_block]) + prev_phys_block_raw = pl.read(inner_compress_state_block_table, [0, prev_state_block]) if prev_phys_block_raw >= 0: prev_phys_block = pl.cast(prev_phys_block_raw, pl.INDEX) prev_state_row = prev_phys_block * INNER_STATE_BLOCK_SIZE + prev_state_intra - pool_kv_tile[front_slot : front_slot + 1, 0:HEAD_CHUNK] = kv_state_flat[ + pool_kv_tile[front_slot : front_slot + 1, 0:HEAD_CHUNK] = compress_state_flat[ prev_state_row : prev_state_row + 1, h0 : h0 + HEAD_CHUNK, ] - pool_score_tile[front_slot : front_slot + 1, 0:HEAD_CHUNK] = score_state_flat[ + pool_score_tile[front_slot : front_slot + 1, 0:HEAD_CHUNK] = compress_state_flat[ prev_state_row : prev_state_row + 1, - h0 : h0 + HEAD_CHUNK, + OUT_DIM + h0 : OUT_DIM + h0 + HEAD_CHUNK, ] cur_abs = cur_start + pool_s @@ -184,17 +183,17 @@ def prefill_indexer_compressor( ) cur_state_block = pl.cast(cur_abs // INNER_STATE_BLOCK_SIZE, pl.INDEX) cur_state_intra = pl.cast(cur_abs - cur_state_block * INNER_STATE_BLOCK_SIZE, pl.INDEX) - cur_phys_block_raw = pl.read(inner_compress_state_block_table, [cur_state_block]) + cur_phys_block_raw = pl.read(inner_compress_state_block_table, [0, cur_state_block]) if cur_phys_block_raw >= 0: cur_phys_block = pl.cast(cur_phys_block_raw, pl.INDEX) cur_state_row = cur_phys_block * INNER_STATE_BLOCK_SIZE + cur_state_intra - pool_kv_tile[back_slot : back_slot + 1, 0:HEAD_CHUNK] = kv_state_flat[ + pool_kv_tile[back_slot : back_slot + 1, 0:HEAD_CHUNK] = compress_state_flat[ cur_state_row : cur_state_row + 1, HEAD_DIM + h0 : HEAD_DIM + h0 + HEAD_CHUNK, ] - pool_score_tile[back_slot : back_slot + 1, 0:HEAD_CHUNK] = score_state_flat[ + pool_score_tile[back_slot : back_slot + 1, 0:HEAD_CHUNK] = compress_state_flat[ cur_state_row : cur_state_row + 1, - HEAD_DIM + h0 : HEAD_DIM + h0 + HEAD_CHUNK, + OUT_DIM + HEAD_DIM + h0 : OUT_DIM + HEAD_DIM + h0 + HEAD_CHUNK, ] for pool_t in pl.range(T): @@ -375,14 +374,14 @@ def prefill_indexer_compressor( ape_slot = pl.cast(update_pos % COMPRESS_RATIO, pl.INDEX) ape_row = ape[ape_slot : ape_slot + 1, update_o0 : update_o0 + OUT_TILE] pool_dep = pl.mul(pooled_kv[0:1, 0:OUT_TILE], 0.0) - kv_state_flat[state_row : state_row + 1, update_o0 : update_o0 + OUT_TILE] = pl.add( + compress_state_flat[state_row : state_row + 1, update_o0 : update_o0 + OUT_TILE] = pl.add( kv_proj_scratch[ update_t : update_t + 1, update_o0 : update_o0 + OUT_TILE, ], pool_dep, ) - score_state_flat[state_row : state_row + 1, update_o0 : update_o0 + OUT_TILE] = pl.add( + compress_state_flat[state_row : state_row + 1, OUT_DIM + update_o0 : OUT_DIM + update_o0 + OUT_TILE] = pl.add( pl.add( score_proj_scratch[update_t : update_t + 1, update_o0 : update_o0 + OUT_TILE], ape_row, @@ -392,18 +391,16 @@ def prefill_indexer_compressor( idx_kv_cache = pl.reshape(idx_kv_cache_flat, [PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM]) idx_kv_scale = pl.reshape(idx_kv_scale_flat, [PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1]) - kv_state = pl.reshape(kv_state_flat, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, OUT_DIM]) - score_state = pl.reshape(score_state_flat, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, OUT_DIM]) - return idx_kv_cache, idx_kv_scale, kv_state, score_state + compress_state = pl.reshape(compress_state_flat, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM]) + return idx_kv_cache, idx_kv_scale, compress_state @pl.jit def prefill_indexer_compressor_test( x: pl.Tensor[[T, D], pl.BF16], kv: pl.Out[pl.Tensor[[MAX_CMP_WRITES, HEAD_DIM], pl.INT8]], - kv_state: pl.InOut[pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32]], - score_state: pl.InOut[pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, OUT_DIM], pl.FP32]], - inner_compress_state_block_table: pl.Tensor[[INNER_STATE_MAX_BLOCKS], pl.INT32], + compress_state: pl.InOut[pl.Tensor[[INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], pl.FP32]], + inner_compress_state_block_table: pl.Tensor[[B, INNER_STATE_MAX_BLOCKS], pl.INT32], wkv: pl.Tensor[[OUT_DIM, D], pl.BF16], wgate: pl.Tensor[[OUT_DIM, D], pl.BF16], ape: pl.Tensor[[COMPRESS_RATIO, OUT_DIM], pl.FP32], @@ -413,14 +410,14 @@ def prefill_indexer_compressor_test( hadamard: pl.Tensor[[HEAD_DIM, HEAD_DIM], pl.BF16], idx_kv_cache: pl.InOut[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], pl.INT8]], idx_kv_scale: pl.InOut[pl.Tensor[[PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1], pl.FP32]], - idx_block_table: pl.Tensor[[IDX_CACHE_MAX_BLOCKS], pl.INT32], + idx_block_table: pl.Tensor[[B, IDX_CACHE_MAX_BLOCKS], pl.INT32], position_ids: pl.Tensor[[T], pl.INT32], num_tokens: pl.Scalar[pl.INT32], idx_slot_mapping: pl.Tensor[[T], pl.INT64], inner_state_slot_mapping: pl.Tensor[[T], pl.INT64], ): prefill_indexer_compressor( - x, kv_state, score_state, inner_compress_state_block_table, wkv, wgate, ape, norm_w, freqs_cos, freqs_sin, + x, compress_state, inner_compress_state_block_table, wkv, wgate, ape, norm_w, freqs_cos, freqs_sin, hadamard, idx_kv_cache, idx_kv_scale, idx_block_table, position_ids, num_tokens, idx_slot_mapping, inner_state_slot_mapping, ) @@ -447,7 +444,7 @@ def prefill_indexer_compressor_test( # INT8 zero via the fp16->int8 cast (a direct pl.full INT8 hits an i8 texpands wall) kv[kv_i : kv_i + 1, 0:HEAD_DIM] = pl.cast( pl.full([1, HEAD_DIM], dtype=pl.FP16, value=0.0), target_type=pl.INT8, mode="trunc") - return kv, kv_state, score_state, idx_kv_cache, idx_kv_scale + return kv, compress_state, idx_kv_cache, idx_kv_scale def golden_prefill_indexer_compressor(tensors): @@ -455,8 +452,7 @@ def golden_prefill_indexer_compressor(tensors): kv_proj = tensors["x"].float() @ tensors["wkv"].float().t() # wkv stored [OUT_DIM, D] for b_trans score_proj = tensors["x"].float() @ tensors["wgate"].float().t() - kv_state_flat = tensors["kv_state"].view(INNER_STATE_BLOCK_NUM * INNER_STATE_BLOCK_SIZE, OUT_DIM) - score_state_flat = tensors["score_state"].view(INNER_STATE_BLOCK_NUM * INNER_STATE_BLOCK_SIZE, OUT_DIM) + compress_state_flat = tensors["compress_state"].view(INNER_STATE_BLOCK_NUM * INNER_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM) state_block_table = tensors["inner_compress_state_block_table"] idx_kv_cache = tensors["idx_kv_cache"] # C8: INT8 KV idx_kv_scale = tensors["idx_kv_scale"] # C8: per-position FP32 dequant scale @@ -473,7 +469,7 @@ def state_row(abs_pos): return -1 block = abs_pos // INNER_STATE_BLOCK_SIZE intra = abs_pos % INNER_STATE_BLOCK_SIZE - phys_block = int(state_block_table[block].item()) + phys_block = int(state_block_table[0, block].item()) if phys_block < 0: return -1 return phys_block * INNER_STATE_BLOCK_SIZE + intra @@ -494,14 +490,14 @@ def state_row(abs_pos): if write_pos >= 2 * COMPRESS_RATIO - 1: prev_row = state_row(prev_abs) if prev_row >= 0: - pool_kv[s] = kv_state_flat[prev_row, :HEAD_DIM] - pool_score[s] = score_state_flat[prev_row, :HEAD_DIM] + pool_kv[s] = compress_state_flat[prev_row, :HEAD_DIM] + pool_score[s] = compress_state_flat[prev_row, OUT_DIM : OUT_DIM + HEAD_DIM] cur_abs = cur_start + s cur_row = state_row(cur_abs) if cur_row >= 0: - pool_kv[COMPRESS_RATIO + s] = kv_state_flat[cur_row, HEAD_DIM:OUT_DIM] - pool_score[COMPRESS_RATIO + s] = score_state_flat[cur_row, HEAD_DIM:OUT_DIM] + pool_kv[COMPRESS_RATIO + s] = compress_state_flat[cur_row, HEAD_DIM:OUT_DIM] + pool_score[COMPRESS_RATIO + s] = compress_state_flat[cur_row, OUT_DIM + HEAD_DIM : COMPRESS_STATE_DIM] for t in range(int(tensors["num_tokens"])): pos = int(position_ids[t].item()) @@ -565,11 +561,10 @@ def state_row(abs_pos): if dst_row < 0: continue ape_slot = pos % COMPRESS_RATIO - kv_state_flat[dst_row] = kv_proj[t] - score_state_flat[dst_row] = score_proj[t] + tensors["ape"][ape_slot] + compress_state_flat[dst_row, 0:OUT_DIM] = kv_proj[t] + compress_state_flat[dst_row, OUT_DIM:COMPRESS_STATE_DIM] = score_proj[t] + tensors["ape"][ape_slot] tensors["kv"][:] = kv - tensors["kv_state"][:] = kv_state_flat.view_as(tensors["kv_state"]) - tensors["score_state"][:] = score_state_flat.view_as(tensors["score_state"]) + tensors["compress_state"][:] = compress_state_flat.view_as(tensors["compress_state"]) tensors["idx_kv_cache"][:] = idx_kv_cache tensors["idx_kv_scale"][:] = idx_kv_scale @@ -589,9 +584,9 @@ def build_tensor_specs(start_pos: int = START_POS): raise ValueError(f"fixture generated {write_count} compressed writes, cap is {MAX_CMP_WRITES}") def init_inner_compress_state_block_table(): - table = torch.full((INNER_STATE_MAX_BLOCKS,), -1, dtype=torch.int32) + table = torch.full((B, INNER_STATE_MAX_BLOCKS), -1, dtype=torch.int32) for block in range(INNER_STATE_MAX_BLOCKS): - table[block] = (block * 17 + 3) % INNER_STATE_MAX_BLOCKS + table[0, block] = (block * 17 + 3) % INNER_STATE_MAX_BLOCKS return table def state_row(abs_pos): if abs_pos < 0 or abs_pos >= MAX_SEQ_LEN: @@ -599,16 +594,17 @@ def state_row(abs_pos): table = init_inner_compress_state_block_table() block = abs_pos // INNER_STATE_BLOCK_SIZE intra = abs_pos % INNER_STATE_BLOCK_SIZE - return int(table[block].item()) * INNER_STATE_BLOCK_SIZE + intra + return int(table[0, block].item()) * INNER_STATE_BLOCK_SIZE + intra def init_x(): return ((torch.rand(T, D) - 0.5) * 0.1).to(torch.bfloat16) - def init_state(): - state = torch.zeros(INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, OUT_DIM) - flat = state.view(-1, OUT_DIM) + def init_compress_state(): + state = torch.zeros(INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM) + flat = state.view(-1, COMPRESS_STATE_DIM) for abs_pos in range(max(0, start_pos - STATE_LEN), start_pos): row = state_row(abs_pos) if row >= 0: - flat[row] = (torch.rand(OUT_DIM) - 0.5) * 0.05 + flat[row, 0:OUT_DIM] = (torch.rand(OUT_DIM) - 0.5) * 0.05 + flat[row, OUT_DIM:COMPRESS_STATE_DIM] = (torch.rand(OUT_DIM) - 0.5) * 0.05 return state # Calibrated to the real DeepSeek-V4-Flash indexer inner compressor (mean l8/l32 of # extract_weights_flash): zero-mean Gaussian BF16 weights at the measured std; the RMSNorm @@ -635,18 +631,18 @@ def init_idx_kv_cache(): def init_idx_kv_scale(): return torch.zeros(PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1) def init_idx_block_table(): - table = torch.full((IDX_CACHE_MAX_BLOCKS,), -1, dtype=torch.int32) + table = torch.full((B, IDX_CACHE_MAX_BLOCKS), -1, dtype=torch.int32) for block in range(IDX_CACHE_MAX_BLOCKS): phys = block if IDX_CACHE_MAX_BLOCKS > 1: phys = (block * 5 + 1) % IDX_CACHE_MAX_BLOCKS - table[block] = phys + table[0, block] = phys return table def idx_row(cmp_slot): table = init_idx_block_table() block = cmp_slot // BLOCK_SIZE intra = cmp_slot % BLOCK_SIZE - phys_block = int(table[block].item()) + phys_block = int(table[0, block].item()) if phys_block < 0: return -1 return phys_block * BLOCK_SIZE + intra @@ -671,9 +667,8 @@ def init_inner_state_slot_mapping(): return [ TensorSpec("x", [T, D], torch.bfloat16, init_value=init_x), TensorSpec("kv", [MAX_CMP_WRITES, HEAD_DIM], torch.int8, is_output=True), - TensorSpec("kv_state", [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, OUT_DIM], torch.float32, init_value=init_state, is_output=True), - TensorSpec("score_state", [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, OUT_DIM], torch.float32, init_value=init_state, is_output=True), - TensorSpec("inner_compress_state_block_table", [INNER_STATE_MAX_BLOCKS], torch.int32, init_value=init_inner_compress_state_block_table), + TensorSpec("compress_state", [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, COMPRESS_STATE_DIM], torch.float32, init_value=init_compress_state, is_output=True), + TensorSpec("inner_compress_state_block_table", [B, INNER_STATE_MAX_BLOCKS], torch.int32, init_value=init_inner_compress_state_block_table), TensorSpec("wkv", [OUT_DIM, D], torch.bfloat16, init_value=init_wkv), TensorSpec("wgate", [OUT_DIM, D], torch.bfloat16, init_value=init_wgate), TensorSpec("ape", [COMPRESS_RATIO, OUT_DIM], torch.float32, init_value=init_ape), @@ -683,7 +678,7 @@ def init_inner_state_slot_mapping(): TensorSpec("hadamard", [HEAD_DIM, HEAD_DIM], torch.bfloat16, init_value=init_hadamard), TensorSpec("idx_kv_cache", [PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, HEAD_DIM], torch.int8, init_value=init_idx_kv_cache, is_output=True), TensorSpec("idx_kv_scale", [PREFILL_IDX_BLOCK_NUM, BLOCK_SIZE, 1, 1], torch.float32, init_value=init_idx_kv_scale, is_output=True), - TensorSpec("idx_block_table", [IDX_CACHE_MAX_BLOCKS], torch.int32, init_value=init_idx_block_table), + TensorSpec("idx_block_table", [B, IDX_CACHE_MAX_BLOCKS], torch.int32, init_value=init_idx_block_table), TensorSpec("position_ids", [T], torch.int32, init_value=init_position_ids), ScalarSpec("num_tokens", torch.int32, T), TensorSpec("idx_slot_mapping", [T], torch.int64, init_value=init_idx_slot_mapping), @@ -721,8 +716,7 @@ def init_inner_state_slot_mapping(): compare_fn={ # C8: raw INT8 compressed rows (+/-1 LSB on the boundary rows the compressor rewrote). "kv": ratio_allclose(atol=1, rtol=0, max_error_ratio=0.01), - "kv_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), - "score_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), + "compress_state": ratio_allclose(atol=1e-3, rtol=1e-3, max_error_ratio=0.0), # C8 cache: INT8 rows exact bar the <=B boundary rows the compressor rewrote (+/-1 LSB). "idx_kv_cache": ratio_allclose(atol=1, rtol=0, max_error_ratio=0.01), "idx_kv_scale": ratio_allclose(atol=1e-4, rtol=1.0 / 128, max_error_ratio=0.01), diff --git a/models/deepseek/v4/prefill_layer.py b/models/deepseek/v4/prefill_layer.py index 2921ba25..d55acba9 100644 --- a/models/deepseek/v4/prefill_layer.py +++ b/models/deepseek/v4/prefill_layer.py @@ -75,6 +75,7 @@ HCA_STATE_BLOCK_SIZE, HCA_STATE_MAX_BLOCKS, MAIN_OUT_DIM as HCA_MAIN_OUT_DIM, + MAIN_STATE_DIM as HCA_MAIN_STATE_DIM, build_tensor_specs as build_hca_attention_tensor_specs, golden_prefill_attention_hca, prefill_attention_hca, @@ -93,10 +94,12 @@ IDX_HEAD_DIM, IDX_N_HEADS, INNER_OUT_DIM, + INNER_STATE_DIM, INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_MAX_BLOCKS, MAIN_OUT_DIM as CSA_MAIN_OUT_DIM, + MAIN_STATE_DIM as CSA_MAIN_STATE_DIM, MAX_SEQ_LEN, O_GROUPS, O_GROUP_IN, @@ -119,6 +122,7 @@ # fixed ``[T, ...]`` tile at a time. TOK_TILE = T PREFILL_CHUNK_TOKENS = T +CHILD_BATCH = config.PREFILL_BATCH DEFAULT_CHUNK_LENS = (T, T + T // 2) DEFAULT_USER_BATCH = len(DEFAULT_CHUNK_LENS) @@ -148,6 +152,9 @@ PREFILL_HCA_STATE_BLOCKS_DYN = pl.dynamic("DEEPSEEK_PREFILL_HCA_STATE_BLOCKS_DYN") PREFILL_CSA_STATE_BLOCKS_DYN = pl.dynamic("DEEPSEEK_PREFILL_CSA_STATE_BLOCKS_DYN") PREFILL_INNER_STATE_BLOCKS_DYN = pl.dynamic("DEEPSEEK_PREFILL_INNER_STATE_BLOCKS_DYN") +PREFILL_HCA_STATE_TABLE_ROWS_DYN = pl.dynamic("DEEPSEEK_PREFILL_HCA_STATE_TABLE_ROWS_DYN") +PREFILL_CSA_STATE_TABLE_ROWS_DYN = pl.dynamic("DEEPSEEK_PREFILL_CSA_STATE_TABLE_ROWS_DYN") +PREFILL_INNER_STATE_TABLE_ROWS_DYN = pl.dynamic("DEEPSEEK_PREFILL_INNER_STATE_TABLE_ROWS_DYN") @pl.jit @@ -173,26 +180,22 @@ def prefill_layer_core( hca_cmp_wgate: pl.Tensor[[HCA_MAIN_OUT_DIM, D], pl.BF16], hca_cmp_ape: pl.Tensor[[HCA_COMPRESS_RATIO, HCA_MAIN_OUT_DIM], pl.FP32], hca_cmp_norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - hca_cmp_kv_state: pl.InOut[pl.Tensor[ - [PREFILL_HCA_STATE_BLOCKS_DYN, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], + hca_compress_state: pl.InOut[pl.Tensor[ + [PREFILL_HCA_STATE_BLOCKS_DYN, HCA_STATE_BLOCK_SIZE, HCA_MAIN_STATE_DIM], pl.FP32, ]], - hca_cmp_score_state: pl.InOut[pl.Tensor[ - [PREFILL_HCA_STATE_BLOCKS_DYN, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], - pl.FP32, - ]], - hca_compress_state_block_table: pl.Tensor[[PREFILL_HCA_STATE_BLOCKS_DYN], pl.INT32], + hca_compress_state_block_table: pl.Tensor[ + [PREFILL_HCA_STATE_TABLE_ROWS_DYN, HCA_STATE_MAX_BLOCKS], + pl.INT32, + ], csa_cmp_wkv: pl.Tensor[[CSA_MAIN_OUT_DIM, D], pl.BF16], csa_cmp_wgate: pl.Tensor[[CSA_MAIN_OUT_DIM, D], pl.BF16], csa_cmp_ape: pl.Tensor[[CSA_COMPRESS_RATIO, CSA_MAIN_OUT_DIM], pl.FP32], csa_cmp_norm_w: pl.Tensor[[HEAD_DIM], pl.BF16], - csa_cmp_kv_state: pl.InOut[ - pl.Tensor[[PREFILL_CSA_STATE_BLOCKS_DYN, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32] + csa_compress_state: pl.InOut[ + pl.Tensor[[PREFILL_CSA_STATE_BLOCKS_DYN, CSA_STATE_BLOCK_SIZE, CSA_MAIN_STATE_DIM], pl.FP32] ], - csa_cmp_score_state: pl.InOut[ - pl.Tensor[[PREFILL_CSA_STATE_BLOCKS_DYN, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32] - ], - csa_compress_state_block_table: pl.Tensor[[PREFILL_CSA_STATE_BLOCKS_DYN], pl.INT32], + csa_compress_state_block_table: pl.Tensor[[PREFILL_CSA_STATE_TABLE_ROWS_DYN, CSA_STATE_MAX_BLOCKS], pl.INT32], csa_hadamard_idx: pl.Tensor[[IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16], csa_idx_wq_b: pl.Tensor[[Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8], csa_idx_wq_b_scale: pl.Tensor[[IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32], @@ -201,13 +204,13 @@ def prefill_layer_core( csa_inner_wgate: pl.Tensor[[INNER_OUT_DIM, D], pl.BF16], csa_inner_ape: pl.Tensor[[CSA_COMPRESS_RATIO, INNER_OUT_DIM], pl.FP32], csa_inner_norm_w: pl.Tensor[[IDX_HEAD_DIM], pl.BF16], - csa_inner_kv_state: pl.InOut[ - pl.Tensor[[PREFILL_INNER_STATE_BLOCKS_DYN, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32] + csa_inner_compress_state: pl.InOut[ + pl.Tensor[[PREFILL_INNER_STATE_BLOCKS_DYN, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32] ], - csa_inner_score_state: pl.InOut[ - pl.Tensor[[PREFILL_INNER_STATE_BLOCKS_DYN, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32] + csa_inner_compress_state_block_table: pl.Tensor[ + [PREFILL_INNER_STATE_TABLE_ROWS_DYN, INNER_STATE_MAX_BLOCKS], + pl.INT32, ], - csa_inner_compress_state_block_table: pl.Tensor[[PREFILL_INNER_STATE_BLOCKS_DYN], pl.INT32], kv_cache: pl.InOut[pl.Tensor[[PREFILL_ORI_CACHE_BLOCKS_DYN, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], ori_block_table: pl.Tensor[[PREFILL_ORI_BLOCK_TABLE_DYN], pl.INT32], ori_slot_mapping: pl.Tensor[[PREFILL_TOKENS_DYN], pl.INT64], @@ -215,7 +218,7 @@ def prefill_layer_core( cmp_block_table: pl.Tensor[[PREFILL_CMP_BLOCK_TABLE_DYN], pl.INT32], idx_kv_cache: pl.InOut[pl.Tensor[[PREFILL_IDX_CACHE_BLOCKS_DYN, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8]], idx_kv_scale: pl.InOut[pl.Tensor[[PREFILL_IDX_CACHE_BLOCKS_DYN, BLOCK_SIZE, 1, 1], pl.FP32]], - idx_block_table: pl.Tensor[[PREFILL_IDX_BLOCK_TABLE_DYN], pl.INT32], + idx_block_table: pl.Tensor[[PREFILL_IDX_BLOCK_TABLE_DYN, IDX_TABLE_BLOCKS], pl.INT32], position_ids: pl.Tensor[[PREFILL_TOKENS_DYN], pl.INT32], hca_cmp_slot_mapping: pl.Tensor[[PREFILL_TOKENS_DYN], pl.INT64], hca_state_slot_mapping: pl.Tensor[[PREFILL_TOKENS_DYN], pl.INT64], @@ -281,15 +284,12 @@ def prefill_layer_core( idx_kv_cache.bind_dynamic(0, PREFILL_IDX_CACHE_BLOCKS_DYN) idx_kv_scale.bind_dynamic(0, PREFILL_IDX_CACHE_BLOCKS_DYN) idx_block_table.bind_dynamic(0, PREFILL_IDX_BLOCK_TABLE_DYN) - hca_cmp_kv_state.bind_dynamic(0, PREFILL_HCA_STATE_BLOCKS_DYN) - hca_cmp_score_state.bind_dynamic(0, PREFILL_HCA_STATE_BLOCKS_DYN) - hca_compress_state_block_table.bind_dynamic(0, PREFILL_HCA_STATE_BLOCKS_DYN) - csa_cmp_kv_state.bind_dynamic(0, PREFILL_CSA_STATE_BLOCKS_DYN) - csa_cmp_score_state.bind_dynamic(0, PREFILL_CSA_STATE_BLOCKS_DYN) - csa_compress_state_block_table.bind_dynamic(0, PREFILL_CSA_STATE_BLOCKS_DYN) - csa_inner_kv_state.bind_dynamic(0, PREFILL_INNER_STATE_BLOCKS_DYN) - csa_inner_score_state.bind_dynamic(0, PREFILL_INNER_STATE_BLOCKS_DYN) - csa_inner_compress_state_block_table.bind_dynamic(0, PREFILL_INNER_STATE_BLOCKS_DYN) + hca_compress_state.bind_dynamic(0, PREFILL_HCA_STATE_BLOCKS_DYN) + hca_compress_state_block_table.bind_dynamic(0, PREFILL_HCA_STATE_TABLE_ROWS_DYN) + csa_compress_state.bind_dynamic(0, PREFILL_CSA_STATE_BLOCKS_DYN) + csa_compress_state_block_table.bind_dynamic(0, PREFILL_CSA_STATE_TABLE_ROWS_DYN) + csa_inner_compress_state.bind_dynamic(0, PREFILL_INNER_STATE_BLOCKS_DYN) + csa_inner_compress_state_block_table.bind_dynamic(0, PREFILL_INNER_STATE_TABLE_ROWS_DYN) user_batch = pl.tensor.dim(seq_lens, 0) for request_id in pl.range(user_batch): chunk_len_b = pl.tensor.read(chunk_lens, [request_id]) @@ -310,25 +310,28 @@ def prefill_layer_core( [ridx * IDX_CACHE_BLOCKS, 0, 0, 0]) idx_kv_scale_req = pl.slice(idx_kv_scale, [IDX_CACHE_BLOCKS, BLOCK_SIZE, 1, 1], [ridx * IDX_CACHE_BLOCKS, 0, 0, 0]) - idx_block_table_req = pl.slice(idx_block_table, [IDX_TABLE_BLOCKS], [ridx * IDX_TABLE_BLOCKS]) - hca_kv_state_req = pl.slice(hca_cmp_kv_state, [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], - [ridx * HCA_STATE_BLOCK_NUM, 0, 0]) - hca_score_state_req = pl.slice(hca_cmp_score_state, [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], - [ridx * HCA_STATE_BLOCK_NUM, 0, 0]) - hca_state_table_req = pl.slice(hca_compress_state_block_table, [HCA_STATE_MAX_BLOCKS], - [ridx * HCA_STATE_MAX_BLOCKS]) - csa_kv_state_req = pl.slice(csa_cmp_kv_state, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], - [ridx * CSA_STATE_BLOCK_NUM, 0, 0]) - csa_score_state_req = pl.slice(csa_cmp_score_state, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], - [ridx * CSA_STATE_BLOCK_NUM, 0, 0]) - csa_state_table_req = pl.slice(csa_compress_state_block_table, [CSA_STATE_MAX_BLOCKS], - [ridx * CSA_STATE_MAX_BLOCKS]) - csa_inner_kv_state_req = pl.slice(csa_inner_kv_state, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], - [ridx * INNER_STATE_BLOCK_NUM, 0, 0]) - csa_inner_score_state_req = pl.slice(csa_inner_score_state, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], - [ridx * INNER_STATE_BLOCK_NUM, 0, 0]) - csa_inner_state_table_req = pl.slice(csa_inner_compress_state_block_table, [INNER_STATE_MAX_BLOCKS], - [ridx * INNER_STATE_MAX_BLOCKS]) + idx_block_table_req = pl.slice(idx_block_table, [CHILD_BATCH, IDX_TABLE_BLOCKS], [ridx * CHILD_BATCH, 0]) + hca_compress_state_req = pl.slice(hca_compress_state, [HCA_STATE_BLOCK_NUM, HCA_STATE_BLOCK_SIZE, HCA_MAIN_STATE_DIM], + [ridx * HCA_STATE_BLOCK_NUM, 0, 0]) + hca_state_table_req = pl.slice( + hca_compress_state_block_table, + [CHILD_BATCH, HCA_STATE_MAX_BLOCKS], + [ridx * CHILD_BATCH, 0], + ) + csa_compress_state_req = pl.slice(csa_compress_state, [CSA_STATE_BLOCK_NUM, CSA_STATE_BLOCK_SIZE, CSA_MAIN_STATE_DIM], + [ridx * CSA_STATE_BLOCK_NUM, 0, 0]) + csa_state_table_req = pl.slice( + csa_compress_state_block_table, + [CHILD_BATCH, CSA_STATE_MAX_BLOCKS], + [ridx * CHILD_BATCH, 0], + ) + csa_inner_compress_state_req = pl.slice(csa_inner_compress_state, [INNER_STATE_BLOCK_NUM, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], + [ridx * INNER_STATE_BLOCK_NUM, 0, 0]) + csa_inner_state_table_req = pl.slice( + csa_inner_compress_state_block_table, + [CHILD_BATCH, INNER_STATE_MAX_BLOCKS], + [ridx * CHILD_BATCH, 0], + ) for tile_id in pl.range(tok_blocks): p0 = tile_id * TOK_TILE @@ -377,7 +380,7 @@ def prefill_layer_core( attn_norm_w, wq_a, wq_b, wq_b_scale, wkv, gamma_cq, gamma_ckv, freqs_cos, freqs_sin, hca_cmp_wkv, hca_cmp_wgate, hca_cmp_ape, hca_cmp_norm_w, - hca_kv_state_req, hca_score_state_req, hca_state_table_req, + hca_compress_state_req, hca_state_table_req, kv_cache_req, ori_slot_tile, ori_block_table_req, cmp_kv_req, cmp_block_table_req, position_ids_tile, hca_cmp_slot_tile, hca_state_slot_tile, @@ -390,11 +393,11 @@ def prefill_layer_core( attn_norm_w, wq_a, wq_b, wq_b_scale, wkv, gamma_cq, gamma_ckv, freqs_cos, freqs_sin, csa_cmp_wkv, csa_cmp_wgate, csa_cmp_ape, csa_cmp_norm_w, - csa_kv_state_req, csa_score_state_req, csa_state_table_req, + csa_compress_state_req, csa_state_table_req, csa_hadamard_idx, csa_idx_wq_b, csa_idx_wq_b_scale, csa_weights_proj, csa_inner_wkv, csa_inner_wgate, csa_inner_ape, csa_inner_norm_w, - csa_inner_kv_state_req, csa_inner_score_state_req, csa_inner_state_table_req, + csa_inner_compress_state_req, csa_inner_state_table_req, kv_cache_req, ori_block_table_req, ori_slot_tile, cmp_kv_req, cmp_block_table_req, idx_kv_cache_req, idx_kv_scale_req, idx_block_table_req, position_ids_tile, csa_cmp_slot_tile, csa_idx_slot_tile, @@ -448,26 +451,25 @@ def l3_prefill_layer( hca_cmp_wgate: pl.Tensor[[N_RANKS, HCA_MAIN_OUT_DIM, D], pl.BF16], hca_cmp_ape: pl.Tensor[[N_RANKS, HCA_COMPRESS_RATIO, HCA_MAIN_OUT_DIM], pl.FP32], hca_cmp_norm_w: pl.Tensor[[N_RANKS, HEAD_DIM], pl.BF16], - hca_cmp_kv_state: pl.InOut[pl.Tensor[ - [N_RANKS, PREFILL_HCA_STATE_BLOCKS_DYN, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], + hca_compress_state: pl.InOut[pl.Tensor[ + [N_RANKS, PREFILL_HCA_STATE_BLOCKS_DYN, HCA_STATE_BLOCK_SIZE, HCA_MAIN_STATE_DIM], pl.FP32, ]], - hca_cmp_score_state: pl.InOut[pl.Tensor[ - [N_RANKS, PREFILL_HCA_STATE_BLOCKS_DYN, HCA_STATE_BLOCK_SIZE, HCA_MAIN_OUT_DIM], - pl.FP32, - ]], - hca_compress_state_block_table: pl.Tensor[[N_RANKS, PREFILL_HCA_STATE_BLOCKS_DYN], pl.INT32], + hca_compress_state_block_table: pl.Tensor[ + [N_RANKS, PREFILL_HCA_STATE_TABLE_ROWS_DYN, HCA_STATE_MAX_BLOCKS], + pl.INT32, + ], csa_cmp_wkv: pl.Tensor[[N_RANKS, CSA_MAIN_OUT_DIM, D], pl.BF16], csa_cmp_wgate: pl.Tensor[[N_RANKS, CSA_MAIN_OUT_DIM, D], pl.BF16], csa_cmp_ape: pl.Tensor[[N_RANKS, CSA_COMPRESS_RATIO, CSA_MAIN_OUT_DIM], pl.FP32], csa_cmp_norm_w: pl.Tensor[[N_RANKS, HEAD_DIM], pl.BF16], - csa_cmp_kv_state: pl.InOut[ - pl.Tensor[[N_RANKS, PREFILL_CSA_STATE_BLOCKS_DYN, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32] + csa_compress_state: pl.InOut[ + pl.Tensor[[N_RANKS, PREFILL_CSA_STATE_BLOCKS_DYN, CSA_STATE_BLOCK_SIZE, CSA_MAIN_STATE_DIM], pl.FP32] ], - csa_cmp_score_state: pl.InOut[ - pl.Tensor[[N_RANKS, PREFILL_CSA_STATE_BLOCKS_DYN, CSA_STATE_BLOCK_SIZE, CSA_MAIN_OUT_DIM], pl.FP32] + csa_compress_state_block_table: pl.Tensor[ + [N_RANKS, PREFILL_CSA_STATE_TABLE_ROWS_DYN, CSA_STATE_MAX_BLOCKS], + pl.INT32, ], - csa_compress_state_block_table: pl.Tensor[[N_RANKS, PREFILL_CSA_STATE_BLOCKS_DYN], pl.INT32], csa_hadamard_idx: pl.Tensor[[N_RANKS, IDX_HEAD_DIM, IDX_HEAD_DIM], pl.BF16], csa_idx_wq_b: pl.Tensor[[N_RANKS, Q_LORA, IDX_N_HEADS * IDX_HEAD_DIM], pl.INT8], csa_idx_wq_b_scale: pl.Tensor[[N_RANKS, IDX_N_HEADS * IDX_HEAD_DIM], pl.FP32], @@ -476,13 +478,13 @@ def l3_prefill_layer( csa_inner_wgate: pl.Tensor[[N_RANKS, INNER_OUT_DIM, D], pl.BF16], csa_inner_ape: pl.Tensor[[N_RANKS, CSA_COMPRESS_RATIO, INNER_OUT_DIM], pl.FP32], csa_inner_norm_w: pl.Tensor[[N_RANKS, IDX_HEAD_DIM], pl.BF16], - csa_inner_kv_state: pl.InOut[ - pl.Tensor[[N_RANKS, PREFILL_INNER_STATE_BLOCKS_DYN, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32] + csa_inner_compress_state: pl.InOut[ + pl.Tensor[[N_RANKS, PREFILL_INNER_STATE_BLOCKS_DYN, INNER_STATE_BLOCK_SIZE, INNER_STATE_DIM], pl.FP32] ], - csa_inner_score_state: pl.InOut[ - pl.Tensor[[N_RANKS, PREFILL_INNER_STATE_BLOCKS_DYN, INNER_STATE_BLOCK_SIZE, INNER_OUT_DIM], pl.FP32] + csa_inner_compress_state_block_table: pl.Tensor[ + [N_RANKS, PREFILL_INNER_STATE_TABLE_ROWS_DYN, INNER_STATE_MAX_BLOCKS], + pl.INT32, ], - csa_inner_compress_state_block_table: pl.Tensor[[N_RANKS, PREFILL_INNER_STATE_BLOCKS_DYN], pl.INT32], kv_cache: pl.InOut[pl.Tensor[[N_RANKS, PREFILL_ORI_CACHE_BLOCKS_DYN, BLOCK_SIZE, 1, HEAD_DIM], pl.BF16]], ori_block_table: pl.Tensor[[N_RANKS, PREFILL_ORI_BLOCK_TABLE_DYN], pl.INT32], ori_slot_mapping: pl.Tensor[[N_RANKS, PREFILL_TOKENS_DYN], pl.INT64], @@ -490,7 +492,7 @@ def l3_prefill_layer( cmp_block_table: pl.Tensor[[N_RANKS, PREFILL_CMP_BLOCK_TABLE_DYN], pl.INT32], idx_kv_cache: pl.InOut[pl.Tensor[[N_RANKS, PREFILL_IDX_CACHE_BLOCKS_DYN, BLOCK_SIZE, 1, IDX_HEAD_DIM], pl.INT8]], idx_kv_scale: pl.InOut[pl.Tensor[[N_RANKS, PREFILL_IDX_CACHE_BLOCKS_DYN, BLOCK_SIZE, 1, 1], pl.FP32]], - idx_block_table: pl.Tensor[[N_RANKS, PREFILL_IDX_BLOCK_TABLE_DYN], pl.INT32], + idx_block_table: pl.Tensor[[N_RANKS, PREFILL_IDX_BLOCK_TABLE_DYN, IDX_TABLE_BLOCKS], pl.INT32], position_ids: pl.Tensor[[N_RANKS, PREFILL_TOKENS_DYN], pl.INT32], hca_cmp_slot_mapping: pl.Tensor[[N_RANKS, PREFILL_TOKENS_DYN], pl.INT64], hca_state_slot_mapping: pl.Tensor[[N_RANKS, PREFILL_TOKENS_DYN], pl.INT64], @@ -550,14 +552,13 @@ def l3_prefill_layer( attn_norm_w[rank], wq_a[rank], wq_b[rank], wq_b_scale[rank], wkv[rank], gamma_cq[rank], gamma_ckv[rank], freqs_cos[rank], freqs_sin[rank], hca_cmp_wkv[rank], hca_cmp_wgate[rank], hca_cmp_ape[rank], hca_cmp_norm_w[rank], - hca_cmp_kv_state[rank], hca_cmp_score_state[rank], hca_compress_state_block_table[rank], + hca_compress_state[rank], hca_compress_state_block_table[rank], csa_cmp_wkv[rank], csa_cmp_wgate[rank], csa_cmp_ape[rank], csa_cmp_norm_w[rank], - csa_cmp_kv_state[rank], csa_cmp_score_state[rank], csa_compress_state_block_table[rank], + csa_compress_state[rank], csa_compress_state_block_table[rank], csa_hadamard_idx[rank], csa_idx_wq_b[rank], csa_idx_wq_b_scale[rank], csa_weights_proj[rank], csa_inner_wkv[rank], csa_inner_wgate[rank], csa_inner_ape[rank], csa_inner_norm_w[rank], - csa_inner_kv_state[rank], csa_inner_score_state[rank], - csa_inner_compress_state_block_table[rank], + csa_inner_compress_state[rank], csa_inner_compress_state_block_table[rank], kv_cache[rank], ori_block_table[rank], ori_slot_mapping[rank], cmp_kv[rank], cmp_block_table[rank], idx_kv_cache[rank], idx_kv_scale[rank], idx_block_table[rank], @@ -602,15 +603,13 @@ def l3_prefill_layer( "hca_cmp_wgate", "hca_cmp_ape", "hca_cmp_norm_w", - "hca_cmp_kv_state", - "hca_cmp_score_state", + "hca_compress_state", "hca_compress_state_block_table", "csa_cmp_wkv", "csa_cmp_wgate", "csa_cmp_ape", "csa_cmp_norm_w", - "csa_cmp_kv_state", - "csa_cmp_score_state", + "csa_compress_state", "csa_compress_state_block_table", "csa_hadamard_idx", "csa_idx_wq_b", @@ -620,8 +619,7 @@ def l3_prefill_layer( "csa_inner_wgate", "csa_inner_ape", "csa_inner_norm_w", - "csa_inner_kv_state", - "csa_inner_score_state", + "csa_inner_compress_state", "csa_inner_compress_state_block_table", "kv_cache", "ori_block_table", @@ -685,8 +683,8 @@ def l3_prefill_layer( _CACHE_STATE_NAMES = { "kv_cache", "block_table", "ori_block_table", "cmp_kv", "cmp_block_table", "idx_kv_cache", "idx_kv_scale", "idx_block_table", - "cmp_kv_state", "cmp_score_state", "compress_state_block_table", - "inner_kv_state", "inner_score_state", "inner_compress_state_block_table", + "compress_state", "compress_state_block_table", + "inner_compress_state", "inner_compress_state_block_table", } # Packed per-request cache/state/table tensors (packed-name -> child-local name, @@ -699,22 +697,19 @@ def l3_prefill_layer( "idx_kv_cache": "idx_kv_cache", "idx_kv_scale": "idx_kv_scale", "idx_block_table": "idx_block_table", - "hca_cmp_kv_state": ("hca", "cmp_kv_state"), - "hca_cmp_score_state": ("hca", "cmp_score_state"), + "hca_compress_state": ("hca", "compress_state"), "hca_compress_state_block_table": ("hca", "compress_state_block_table"), - "csa_cmp_kv_state": ("csa", "cmp_kv_state"), - "csa_cmp_score_state": ("csa", "cmp_score_state"), + "csa_compress_state": ("csa", "compress_state"), "csa_compress_state_block_table": ("csa", "compress_state_block_table"), - "csa_inner_kv_state": ("csa", "inner_kv_state"), - "csa_inner_score_state": ("csa", "inner_score_state"), + "csa_inner_compress_state": ("csa", "inner_compress_state"), "csa_inner_compress_state_block_table": ("csa", "inner_compress_state_block_table"), } _HISTORY_CACHE_NAMES = { "kv_cache", "cmp_kv", "idx_kv_cache", - "hca_cmp_kv_state", "hca_cmp_score_state", - "csa_cmp_kv_state", "csa_cmp_score_state", - "csa_inner_kv_state", "csa_inner_score_state", + "hca_compress_state", + "csa_compress_state", + "csa_inner_compress_state", } @@ -731,15 +726,15 @@ def _req_block_count(kind, child_name): if child_name in ("idx_kv_cache", "idx_kv_scale"): return IDX_CACHE_BLOCKS if child_name == "idx_block_table": - return IDX_TABLE_BLOCKS - if child_name in ("cmp_kv_state", "cmp_score_state"): + return CHILD_BATCH + if child_name == "compress_state": return HCA_STATE_BLOCK_NUM if kind == "hca" else CSA_STATE_BLOCK_NUM if child_name == "compress_state_block_table": - return HCA_STATE_MAX_BLOCKS if kind == "hca" else CSA_STATE_MAX_BLOCKS - if child_name in ("inner_kv_state", "inner_score_state"): + return CHILD_BATCH + if child_name == "inner_compress_state": return INNER_STATE_BLOCK_NUM if child_name == "inner_compress_state_block_table": - return INNER_STATE_MAX_BLOCKS + return CHILD_BATCH raise KeyError(child_name) @@ -1108,7 +1103,7 @@ def golden_prefill_layer(tensors): mapped.update({ "cmp_wkv": tensors["hca_cmp_wkv"], "cmp_wgate": tensors["hca_cmp_wgate"], "cmp_ape": tensors["hca_cmp_ape"], "cmp_norm_w": tensors["hca_cmp_norm_w"], - "cmp_kv_state": tensors["hca_cmp_kv_state"], "cmp_score_state": tensors["hca_cmp_score_state"], + "compress_state": tensors["hca_compress_state"], "compress_state_block_table": tensors["hca_compress_state_block_table"], "cmp_slot_mapping": tensors["hca_cmp_slot_mapping"], "state_slot_mapping": tensors["hca_state_slot_mapping"], }) @@ -1117,13 +1112,13 @@ def golden_prefill_layer(tensors): mapped.update({ "cmp_wkv": tensors["csa_cmp_wkv"], "cmp_wgate": tensors["csa_cmp_wgate"], "cmp_ape": tensors["csa_cmp_ape"], "cmp_norm_w": tensors["csa_cmp_norm_w"], - "cmp_kv_state": tensors["csa_cmp_kv_state"], "cmp_score_state": tensors["csa_cmp_score_state"], + "compress_state": tensors["csa_compress_state"], "compress_state_block_table": tensors["csa_compress_state_block_table"], "hadamard_idx": tensors["csa_hadamard_idx"], "idx_wq_b": tensors["csa_idx_wq_b"], "idx_wq_b_scale": tensors["csa_idx_wq_b_scale"], "idx_weights_proj": tensors["csa_weights_proj"], "inner_wkv": tensors["csa_inner_wkv"], "inner_wgate": tensors["csa_inner_wgate"], "inner_ape": tensors["csa_inner_ape"], "inner_norm_w": tensors["csa_inner_norm_w"], - "inner_kv_state": tensors["csa_inner_kv_state"], "inner_score_state": tensors["csa_inner_score_state"], + "inner_compress_state": tensors["csa_inner_compress_state"], "inner_compress_state_block_table": tensors["csa_inner_compress_state_block_table"], "cmp_slot_mapping": tensors["csa_cmp_slot_mapping"], "idx_slot_mapping": tensors["csa_idx_slot_mapping"], "state_slot_mapping": tensors["csa_state_slot_mapping"],