diff --git a/atom/model_ops/attention_mla.py b/atom/model_ops/attention_mla.py index 18f2814528..b8605d332f 100644 --- a/atom/model_ops/attention_mla.py +++ b/atom/model_ops/attention_mla.py @@ -95,8 +95,27 @@ def indexer_qk_rope_quant_and_cache( weights_scale: float, preshuffle: bool = False, is_neox: bool = True, + q_scale_out: torch.Tensor | None = None, + kv_cache_scale: torch.Tensor | None = None, ) -> None: - """Run the fused indexer cache op with ATOM's DCP query semantics.""" + """Run the fused indexer cache op with ATOM's DCP query semantics. + + The two scale buffers together switch the op to packed E2M1 + e8m0 outputs, + and are forwarded only when given so that an aiter predating them still + accepts the FP8 call -- the same version tolerance the seg-variant import + above keeps. One without the other would fall back to FP8 and write that + layout into FP4-shaped buffers, so it is refused here rather than passed on. + """ + if (q_scale_out is None) != (kv_cache_scale is None): + raise ValueError( + "FP4 output needs both q_scale_out and kv_cache_scale, got only " + + ("q_scale_out" if q_scale_out is not None else "kv_cache_scale") + ) + fp4_out = ( + {"q_scale_out": q_scale_out, "kv_cache_scale": kv_cache_scale} + if q_scale_out is not None + else {} + ) _indexer_qk_rope_quant_and_cache( q, q_out, @@ -117,6 +136,7 @@ def indexer_qk_rope_quant_and_cache( preshuffle=preshuffle, is_neox=is_neox, compute_all_q_rope=get_dcp_world_size() > 1, + **fp4_out, ) @@ -2993,6 +3013,9 @@ def _convert_req_index_to_global_index_dsa_prefill_kernel( MAX_NUM_BLOCKS_PER_REQ: tl.constexpr, PAGE_SIZE: tl.constexpr, BLOCK_N: tl.constexpr, # tile width along columns + # The FP8 indexer scores one concatenated KV plane, so a request's own + # position is `indice - cu_seqlens_q[req]`; the paged FP4 scorer emits it. + SEQ_LOCAL: tl.constexpr, # strides (in elements) ti_stride0: tl.int64, # topk_indices stride 0 ti_stride1: tl.constexpr, # topk_indices stride 1 @@ -3019,7 +3042,7 @@ def _convert_req_index_to_global_index_dsa_prefill_kernel( req_kv_end = tl.load(cu_seqlens_q + req_id + 1, mask=valid_req, other=0) req_kv_len = req_kv_end - pre_seqlens_q - seq_token_idx = indice - pre_seqlens_q + seq_token_idx = indice if SEQ_LOCAL else indice - pre_seqlens_q block_id = seq_token_idx // PAGE_SIZE inblock_offset = seq_token_idx % PAGE_SIZE @@ -3069,6 +3092,7 @@ def triton_convert_req_index_to_global_index_dsa_prefill( NUM_TOPK_TOKENS: int = 2048, BLOCK_N: int = 1024, # tile width along columns out: torch.Tensor | None = None, + seq_local: bool = False, ): assert topk_indices.shape[1] == NUM_TOPK_TOKENS @@ -3120,6 +3144,7 @@ def triton_convert_req_index_to_global_index_dsa_prefill( max_num_blocks_per_req, PAGE_SIZE, BLOCK_N, + seq_local, # strides ti_stride0, ti_stride1, diff --git a/atom/model_ops/attentions/aiter_mla.py b/atom/model_ops/attentions/aiter_mla.py index 9613dfb83e..35927ba7ed 100644 --- a/atom/model_ops/attentions/aiter_mla.py +++ b/atom/model_ops/attentions/aiter_mla.py @@ -48,6 +48,14 @@ effective_kpool_size, topk_output_width, ) +from atom.model_ops.sparse_indexer_fp4 import ( + FP4_MQA_BLOCK_K, + FP4_MQA_PARALLEL_UNIT_NUM, + fp4_decode_parallel_units, + fp4_decode_schedule, + fp4_prefill_schedule, + sparse_indexer_fp4_enabled, +) from atom.utils import CpuGpuBuffer, envs, upload_numpy from atom.utils.block_convert import ( kv_indices_generate_triton, @@ -257,6 +265,10 @@ class AiterMLAMetadataBuilder(CommonAttentionBuilder): # backend). The fused kernel handles both sparse and dense MLA. fuse_mtp_decode_position_update = True + # `__init__` decides this; the default serves builders assembled field by + # field, as tests do. + _indexer_fp4 = False + def _global_num_draft_layers(self) -> int: """Return draft layers in the target MLA pool across all PP stages.""" runner = self.model_runner @@ -381,6 +393,9 @@ def __init__(self, model_runner): ) self.index_kpool = effective_kpool_size(configured_kpool) self.index_topk_out = topk_output_width(self.index_topk, configured_kpool) + self._indexer_fp4 = self.is_sparse and sparse_indexer_fp4_enabled( + config.index_cache_dtype, hf_config, warn=True + ) self.dtype_kv = dtypes.d_dtypes[config.kv_cache_dtype] self.dtype_q = self.dtype_kv @@ -522,6 +537,23 @@ def __init__(self, model_runner): dtype=torch.int32, device=self.device, ) + # `[P, 4]` int32 read by the mqa-logits kernel during a CUDAGraph + # replay, so it has to be a fixed address refreshed in place, and + # one per ubatch since TBO has two in flight. + self._indexer_fp4_cta_info = ( + [ + torch.zeros( + fp4_decode_parallel_units(self.max_bs, max_seqlen_qo), + 4, + **i32_kwargs, + ) + for _ in range( + self._NUM_TBO_UBATCHES + 1 if config.enable_tbo else 1 + ) + ] + if self._indexer_fp4 + else [] + ) # One block-table row per query token; only MTP verify needs a # copy. Built once per step, not in the indexer, where every # full-index layer would allocate one into the CUDAGraph pool. @@ -1181,6 +1213,8 @@ def _declare_kv_pool(self) -> MlaKvPool: "index_rows_per_block": self._index_rows_per_block(), "index_dim": aligned_index_cache_dim(hf_config), "index_dtype": dtypes.fp8, + "index_head_dim": hf_config.index_head_dim, + "index_fp4": self._indexer_fp4, } if runner.has_mla_indexer else {} @@ -1264,10 +1298,27 @@ def build_kv_cache_tensor(self, module): ) index_cache_layer_id = runner.index_cache_layer_map[global_layer_id] index_cache = self.kv_pool.layer("index", index_cache_layer_id) - # Use aligned dimension to avoid memory copy in torch inductor - module.indexer.k_cache.kv_cache[0] = index_cache.view( - -1, 1, runner.aligned_index_dim - ) + # `is_sparse` is this branch's own condition, so both sides + # reduce to the same `sparse_indexer_fp4_enabled` call on the same + # two inputs and cannot legitimately disagree. Compared rather than + # assigned: the Indexer built `k_cache` from its own answer before + # this runs, so writing the builder's over it would leave the flag + # describing an object it no longer matches. + assert module.indexer._indexer_fp4 == self._indexer_fp4, ( + "FP4 sparse indexer verdict diverged: metadata builder says " + f"{self._indexer_fp4}, layer {global_layer_id}'s Indexer says " + f"{module.indexer._indexer_fp4}" + ) + if self._indexer_fp4: + module.indexer.k_cache.kv_cache[0] = index_cache + module.indexer.k_cache.kv_cache_scale = self.kv_pool.layer( + "index_scale", index_cache_layer_id + ) + else: + # Use aligned dimension to avoid memory copy in torch inductor + module.indexer.k_cache.kv_cache[0] = index_cache.view( + -1, 1, runner.aligned_index_dim + ) module.kv_cache = kv_cache return KVCacheTensor( layer_num=module.layer_num, @@ -1289,6 +1340,17 @@ def get_kv_transfer_tensors(self): runner = self.model_runner if self.kv_pool is None: return None + if self._indexer_fp4: + # A connector is handed one `INDEX_CACHE_ROLE` region per layer -- + # the whole vocabulary it and its DCP shard plan have for the + # indexer -- and the FP4 cache is two planes neither parses. + if runner.config.kv_transfer_config: + raise NotImplementedError( + "KV transfer with the FP4 sparse indexer is unsupported: " + "the region map cannot describe its separate e8m0 scale " + "plane. Pass --index_cache_dtype fp8 to use a connector." + ) + return None # What the pool was built with, not what the hook would recompute: a # hybrid caches for fewer layers than the model has, and the consumer # indices below are positions in the allocated rows. @@ -1556,6 +1618,43 @@ def _build_dcp_indexer_prefill_meta(self, attn_metadata, bs: int, counts, var): attn_metadata.dcp_indexer_gather_index = torch.from_numpy( src.astype(np.int32) ).to(dev, non_blocking=True) + if self._indexer_fp4: + self._build_dcp_indexer_fp4_prefill_meta( + attn_metadata, bs, lpad, cu_pad, total_kv, var + ) + + def _build_dcp_indexer_fp4_prefill_meta( + self, attn_metadata, bs: int, lpad, cu_pad, total_kv: int, var + ): + """Publish the local slot list and identity page table FP4 staging reads. + + The slot formula must track `cp_gather_indexer_k_quant_cache`, the FP8 + plane's gather, which the FP4 planes have no equivalent of. Its padded + indices stay inside the sequence's own blocks, so the staging read + cannot leave the allocation. + """ + block = self.model_runner.block_size + dev = self.device + seq_of = np.repeat(np.arange(bs, dtype=np.int64), lpad) + j = np.arange(int(cu_pad[bs]), dtype=np.int64) - np.repeat(cu_pad[:bs], lpad) + table = var["block_tables"].np[:bs].astype(np.int64) + slots = table[seq_of, j // block] * block + (j % block) + attn_metadata.dcp_indexer_fp4_local_slots = torch.from_numpy( + slots.astype(np.int32) + ).to(dev, non_blocking=True) + + # Fixed width, not `pages`: the scorer specializes on this table's + # stride, so a per-batch width recompiles it. Sized at a whole batch's + # summed context rather than `block_tables`' one-sequence allowance -- + # co-scheduled prefills run past that allowance whenever prefix caching + # keeps their cached tokens off `max_num_batched_tokens`, which is a + # legal schedule, and this table is the only thing that would have + # bounded it. Buys one scorer variant against the non-DCP width. + pages = -(-total_kv // block) + cols = self.max_bs * self.block_table_cols + staged_tables = torch.zeros(bs, cols, dtype=torch.int32, device=dev) + staged_tables[:, :pages] = torch.arange(pages, dtype=torch.int32, device=dev) + attn_metadata.dcp_indexer_fp4_block_tables = staged_tables def _sparse_selected_counts(self, seq_lens): """How many KV entries the indexer actually selects for each row. @@ -1662,6 +1761,9 @@ def prepare_prefill(self, batch: ScheduledBatch, running_bs: int): attn_metadata.sparse_kv_indptr = var["sparse_kv_indptr"].copy_to_gpu( scheduled_tokens + 1 ) + self._publish_indexer_fp4_prefill_schedule( + attn_metadata, sparse_counts, int(full_seq_lens.sum()) + ) if self.dcp_world_size > 1: self._build_dcp_indexer_prefill_meta(attn_metadata, bs, counts, var) get_mla_metadata_v1( @@ -2131,6 +2233,105 @@ def _publish_dcp_token_block_tables( rows.view(running_bs, max_seqlen_q, -1).copy_(block_tables.unsqueeze(1)) attn_metadata.dcp_token_block_tables = rows + def _publish_indexer_fp4_decode_schedule( + self, attn_metadata: AttentionMetaData, bs: int, next_n: int, ubatch: int = 0 + ) -> None: + """Refresh the FP4 decode CTA schedule; the captured kernel replays off + it, and every indexer layer of a step shares the one answer. + + A replay reads the buffer's CONTENTS, so every decode forward has to + call this -- including the MTP draft's, whose rows are its own. The + lengths come from the published buffer rather than the metadata's view + of it, because a draft following a prefill carries prefill metadata. + """ + if not self._indexer_fp4: + return + parallel_units = fp4_decode_parallel_units(bs, next_n) + cta_info = self._indexer_fp4_cta_info[ubatch][:parallel_units] + # `[:n]` past the end truncates rather than raising, while + # `indexer_fp4_n_ctas` still hands the kernel `n`. The rows past the end + # would then go unscheduled, and the logits buffer they should have + # written is `torch.empty` -- aiter only sentinels it when it builds the + # schedule itself, which it never does here. + assert cta_info.shape[0] == parallel_units + if self.dcp_world_size > 1: + # DCP scores this rank's shard. Its rows are query tokens carrying + # their own local window, so next_n is already flattened out of the + # row axis and the width is the sharded one -- `parallel_units` + # stays the non-DCP count, which keeps the captured grid identical. + from atom.model_ops.dcp_ops import dcp_local_logits_width + from atom.utils.forward_context import ( + get_published_dcp_local_context_lens, + ) + + # A TBO ubatch owns rows from its own offset, not the batch's first + # `bs`, and it already publishes that slice on its metadata -- the + # same tensor the scorer reads. The global buffer stays the fallback + # so callers that publish nothing keep the prefix they had. + context_lens = get_published_dcp_local_context_lens( + attn_metadata, bs * next_n + ) + if context_lens is None: + context_lens = self.model_runner.forward_vars[ + "dcp_local_context_lens" + ].gpu[: bs * next_n] + schedule_next_n = 1 + width = dcp_local_logits_width( + self.model_runner.config.max_model_len, self.dcp_world_size + ) + else: + context_lens = attn_metadata.context_lens[:bs] + schedule_next_n = next_n + width = self.model_runner.config.max_model_len + fp4_decode_schedule( + context_lens, + FP4_MQA_BLOCK_K, + parallel_units, + width, + schedule_next_n, + cta_info, + ) + attn_metadata.indexer_fp4_cta_info = cta_info + attn_metadata.indexer_fp4_n_ctas = parallel_units + + def _publish_indexer_fp4_prefill_schedule( + self, attn_metadata: AttentionMetaData, local_ends_np: np.ndarray, total_kv: int + ) -> None: + """Build the FP4 ragged-prefill schedule and the width it scores in. + + The width comes off `sparse_counts`, which the host already holds, so + right-sizing the logits buffer to this batch rather than max_model_len + costs no device sync. + + Under DCP the scorer reads a staged copy of every sequence's keys rather + than the in-place shard, so its columns are the flat concatenated ones + the FP8 path ranks in and the width is the whole key set. + """ + if not self._indexer_fp4: + return + if self.dcp_world_size > 1: + local_starts = attn_metadata.cu_seqlen_ks + local_ends = attn_metadata.cu_seqlen_ke + max_seq_len = max(int(total_kv), 1) + else: + local_starts = None + local_ends = attn_metadata.cu_seqlen_ke - attn_metadata.cu_seqlen_ks + max_seq_len = max(int(local_ends_np.max(initial=0)), 1) + ( + attn_metadata.indexer_fp4_cta_info, + attn_metadata.indexer_fp4_n_ctas, + attn_metadata.indexer_fp4_local_starts, + ) = fp4_prefill_schedule( + attn_metadata.batch_id_per_q_token, + local_ends, + FP4_MQA_BLOCK_K, + FP4_MQA_PARALLEL_UNIT_NUM, + max_seq_len, + local_starts, + ) + attn_metadata.indexer_fp4_local_ends = local_ends + attn_metadata.indexer_fp4_max_seq_len = max_seq_len + def prepare_decode( self, batch: ScheduledBatch, @@ -2437,6 +2638,9 @@ def prepare_decode( "sparse_kv_last_page_lens" ].gpu[:running_bs] self._publish_dcp_token_block_tables(attn_metadata, running_bs, max_seqlen_q) + self._publish_indexer_fp4_decode_schedule( + attn_metadata, running_bs, max_seqlen_q + ) # running_bs, not scheduled_bs: the padded rows have to be split into the # ubatches too, or accuracy drifts. @@ -2742,6 +2946,7 @@ def build_for_cudagraph_capture(self, bs: int) -> AttentionMetaData: "sparse_kv_last_page_lens" ].gpu[:bs] self._publish_dcp_token_block_tables(attn_matadata, bs, max_q_len) + self._publish_indexer_fp4_decode_schedule(attn_matadata, bs, max_q_len) positions = var["positions"].copy_to_gpu(scheduled_tokens) context = Context( positions=positions, @@ -2819,6 +3024,11 @@ def build_ubatch_metadata( # DCP, so a ubatch always runs one query per sequence. assert max_q_len == 1 attn.dcp_token_block_tables = attn.block_tables + # Slot 0 stays the full-batch schedule, so a ubatch never overwrites + # one the other's kernels still read. + self._publish_indexer_fp4_decode_schedule( + attn, running_bs, max_q_len, ubatch=ubatch_idx + 1 + ) return attn def build_ubatch_prefill_metadata( diff --git a/atom/model_ops/attentions/backends.py b/atom/model_ops/attentions/backends.py index 50142427cf..0b03240179 100644 --- a/atom/model_ops/attentions/backends.py +++ b/atom/model_ops/attentions/backends.py @@ -479,6 +479,17 @@ def __init__(self, model_runner): self.model_runner.forward_vars.update(attn_metadata) self.has_sliding_window = hasattr(hf_config, "sliding_window") + def _publish_indexer_fp4_decode_schedule( + self, attn_metadata, bs: int, next_n: int, ubatch: int = 0 + ) -> None: + """Nothing to refresh: this backend has no FP4 sparse indexer. + + `EagleProposer` publishes on whatever builder the target uses, so every + backend a draft can run against has to answer. Only the MLA one + overrides. Inert by contract, not just by accident -- the draft reuses + the target's metadata, so a write here would reach the verify step. + """ + def prepare_block_tables(self, batch: ScheduledBatch): """Marshal the batch's block tables into `forward_vars["block_tables"]`. diff --git a/atom/model_ops/attentions/mla_kv_pool.py b/atom/model_ops/attentions/mla_kv_pool.py index 9067438a1d..a12ea1e6e8 100644 --- a/atom/model_ops/attentions/mla_kv_pool.py +++ b/atom/model_ops/attentions/mla_kv_pool.py @@ -31,6 +31,7 @@ carve_layer_major, entry_bytes_for, ) +from atom.model_ops.sparse_indexer_fp4 import fp4_index_block_shapes class MlaKvPool: @@ -55,6 +56,8 @@ def __init__( index_rows_per_block: int = 0, index_dim: int = 0, index_dtype: torch.dtype | None = None, + index_head_dim: int = 0, + index_fp4: bool = False, ): self.layers = layers self.block_size = block_size @@ -62,8 +65,17 @@ def __init__( self.cache_fields = [ EntryField("kv", layers, (block_size, entry_dim), kv_dtype) ] - self.index_fields = ( - [ + self.index_fields: list[EntryField] = [] + if index_layers and index_fp4: + data_shape, scale_shape = fp4_index_block_shapes( + index_rows_per_block, index_head_dim + ) + self.index_fields = [ + EntryField("index", index_layers, data_shape, torch.uint8), + EntryField("index_scale", index_layers, scale_shape, torch.uint8), + ] + elif index_layers: + self.index_fields = [ EntryField( "index", index_layers, @@ -71,9 +83,6 @@ def __init__( index_dtype, ) ] - if index_layers - else [] - ) self.index_dim = index_dim # The regions a block is charged for, in layout order; one list for the # price and the allocation both. @@ -95,9 +104,10 @@ def allocate(self, entries: int, device, buf: torch.Tensor | None = None) -> Non self.field_groups, entries, device, buf ) self._views = { - name: arena.view(name) - for name, arena in (("kv", self.cache), ("index", self.index)) + field.name: arena.view(field.name) + for arena in (self.cache, self.index) if arena is not None + for field in arena.fields } def release(self) -> None: diff --git a/atom/model_ops/dcp_ops.py b/atom/model_ops/dcp_ops.py index 7b9ce7377a..fd8c7aaf66 100644 --- a/atom/model_ops/dcp_ops.py +++ b/atom/model_ops/dcp_ops.py @@ -823,6 +823,15 @@ def dcp_prefill_slot_mapping( return slot_mapping +def dcp_local_logits_width(max_model_len: int, dcp_world_size: int) -> int: + """Column count of one DCP rank's local indexer logits plane. + + The FP4 metadata builder bakes this into a CUDAGraph-captured schedule while + the scorer allocates the plane, so the two have to read it off one formula. + """ + return -(-max_model_len // dcp_world_size) + + def dcp_local_context_lens( attn_metadata, dcp_rank: int, @@ -900,6 +909,9 @@ def dcp_decode_candidate_exchange_fused( out_kv_indices: torch.Tensor, out_kv_indptr: torch.Tensor, owned_counts: torch.Tensor, + q_scale: torch.Tensor | None = None, + kv_scale: torch.Tensor | None = None, + weights_scale: float = 1.0, ) -> None: """Score the local shard, exchange scores, emit this rank's owned KV slots. @@ -914,12 +926,18 @@ def dcp_decode_candidate_exchange_fused( a gid-ordered tie-break, so when the local boundary (2048th) sits on a score tie a tied token may be dropped before the exchange. Exact fp32 score ties are rare, so this is a negligible boundary effect, not a systematic loss. + + Only the scoring below is dtype-bound: pass `q_scale`/`kv_scale` to score the + shard in FP4 instead of FP8. Everything after it -- the local top-k, the + exchange and the merge -- reads fp32 logits and is the same either way. """ # Imported here, not at module scope, so this module stays importable on a # machine without aiter (collection of CPU-only tests, doc builds). from aiter.ops.topk import flydsl_dcp_topk_merge, top_k_per_row_decode from aiter.ops.triton.pa_mqa_logits import deepgemm_fp8_paged_mqa_logits + from atom.model_ops.sparse_indexer_fp4 import FP4_MQA_BLOCK_K + dcp_world_size = get_dcp_world_size() # Size everything off num_decode_tokens, the rows this rank actually # scheduled -- NOT padded_q_fp8_decode_tokens.shape, which is the padded @@ -944,21 +962,45 @@ def dcp_decode_candidate_exchange_fused( cp_kv_cache_interleave_size, num_decode_tokens, ) - l_max = (max_model_len + dcp_world_size - 1) // dcp_world_size + l_max = dcp_local_logits_width(max_model_len, dcp_world_size) local_logits = torch.empty( [num_decode_tokens, l_max], dtype=torch.float32, device="cuda" ) - deepgemm_fp8_paged_mqa_logits( - q_rows, - kv_cache, - weights[:num_decode_tokens], - local_logits, - local_ctx, - block_tables, - l_max, - KVBlockSize=runner_block_size, - Preshuffle=True, - ) + if q_scale is None: + deepgemm_fp8_paged_mqa_logits( + q_rows, + kv_cache, + weights[:num_decode_tokens], + local_logits, + local_ctx, + block_tables, + l_max, + KVBlockSize=runner_block_size, + Preshuffle=True, + ) + else: + from aiter.ops.flydsl import flydsl_pa_mqa_logits_fp4 + + # The schedule the metadata builder published is the one this kernel was + # captured with, and it is built off `local_ctx` at `l_max` -- the same + # two the call below scores in. Neither may be re-derived here. + flydsl_pa_mqa_logits_fp4( + q_rows, + q_scale[:num_decode_tokens].unsqueeze(1), + kv_cache, + kv_scale, + block_tables, + weights[:num_decode_tokens], + local_ctx, + l_max, + weight_scale=weights_scale, + next_n=1, + block_k=FP4_MQA_BLOCK_K, + kv_block_size=runner_block_size, + out=local_logits, + cta_info=attn_metadata.indexer_fp4_cta_info, + total_ctas=attn_metadata.indexer_fp4_n_ctas, + ) # k_loc is the constant `topk_tokens`, never the live local length: the # exchanged size must be static for CUDAGraph. Short contexts therefore ship diff --git a/atom/model_ops/sparse_indexer_fp4.py b/atom/model_ops/sparse_indexer_fp4.py new file mode 100644 index 0000000000..9620348431 --- /dev/null +++ b/atom/model_ops/sparse_indexer_fp4.py @@ -0,0 +1,264 @@ +# SPDX-License-Identifier: MIT +# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved. + +"""FP4 scoring for a DSA sparse indexer: the predicate, the ABI, the schedule. + +FP4 is a change of dtype, not a second cache: `indexer_qk_rope_quant_and_cache` +writes packed E2M1 keys plus their e8m0 scales in place of the FP8 row, and +`flydsl_pa_mqa_logits_fp4[_prefill]` score them straight out of the paged cache. +Every number here is one those kernels fix, so it drifts with them. +""" + +from __future__ import annotations + +import logging +from typing import Any + +import torch + +logger = logging.getLogger("atom") + +# Persistent-grid schedule params for the `pa_mqa_logits_fp4*` kernels, which +# decode and prefill both score through. A metadata builder precomputes each +# path's cta_info with these and the scorer passes the matching block_k, so +# layout and grid agree. DeepSeek-V4 drives the same two kernels and carries its +# own copy of this pair in `v4_kernels`; the two must be changed together until +# something merges them, which is not this module's to do. +# +# Both floors are CTA-count targets, not kernel defaults: every consumer takes +# `max(floor, rows)`, so a floor only adds split-K to grids too small to fill +# the GPU and is an identity for the wide ones. Splits are numerically inert -- +# each CTA gets a disjoint KV-column range, no cross-CTA partial sums. +# +# Neither can be fitted to the work. A CUDAGraph bakes the grid at capture, +# while the work is `rows * ceil(ctx / block_k)` and the context term is only +# known per replay; of what capture does know, the batch size measures flat, so +# there is nothing to derive one from. Both are therefore the smallest +# worst-case regret over the served context range rather than any shape's +# optimum -- re-tune against a real workload mix. At prefill rows=1024 and +# W~32768, 4096 measured 206.2us a call against 512's 224.6us. +FP4_MQA_PARALLEL_UNIT_NUM = 4096 +# The varctx scorer takes a lower floor than the prefill one. Its rows are +# sequences rather than query tokens, so `max(floor, rows)` never lifts off the +# floor and every CTA past the work is pure setup: 8.75us a layer at 4096 +# against 2.99 at 512, flat in context. Going lower still gives back more at +# 131k+ than it wins at 4k. This is not "the decode floor" -- DeepSeek-V4's +# decode scores through the prefill kernel and keeps 4096. +FP4_MQA_VARCTX_PARALLEL_UNIT_NUM = 512 +FP4_MQA_BLOCK_K = 256 + +# The fused writer's FP4 group width, against 128 on the FP8 path. +FP4_QUANT_BLOCK_SIZE = 32 +# `pa_mqa_logits_fp4*` pack four N-tiles' e8m0 bytes into one dword and read +# them with N_PHYS == 1, which holds only where a paged block covers exactly +# NTPW(4) x MFMA_N(16) indexer rows. +FP4_KV_BLOCK_SIZE = 64 +_MFMA_M = 16 +_K_TILE = 128 + + +def _gfx() -> str: + """The chip this is running on, imported on use rather than at module scope: + a non-GPU install carries `aiter` without `aiter.jit`, and the rest of this + module is ABI constants such an install still has to be able to read.""" + from aiter.jit.utils.chip_info import get_gfx + + return get_gfx() + + +def sparse_indexer_fp4_enabled( + index_cache_dtype: str | None, config: Any, *, warn: bool = False +) -> bool: + """Does this layer's DSA indexer score in FP4? The one place that decides. + + Structural, never per-model: the kernels need an indexer whose head dim is a + single 128-wide K tile and whose head count tiles by MFMA_M, on a chip that + has them. That also reaches the right verdict for an MTP draft, whose + `model_type` `SpeculativeConfig._MTP_TYPE_MAP` has rewritten but whose + indexer geometry -- and cache, shared with the target -- is unchanged. + + The metadata builder and `Indexer.__init__` must agree: the builder picks the + cache-pool width, but warmup traces the indexer's graph piece before + `allocate_kv_cache` runs, so a layer that guessed wrong bakes in the other + branch. Pass `warn=True` from the builder only; the Indexer runs per layer. + """ + if index_cache_dtype != "fp4": + return False + head_dim = getattr(config, "index_head_dim", None) + n_heads = getattr(config, "index_n_heads", None) + if not getattr(config, "index_topk", None): + why = "the layer has no sparse indexer" + elif head_dim != _K_TILE: + why = f"index_head_dim is {head_dim}, not {_K_TILE}" + elif not n_heads or n_heads % _MFMA_M: + why = f"index_n_heads is {n_heads}, not a multiple of {_MFMA_M}" + elif (getattr(config, "index_kpool", 1) or 1) != 1: + # A pooled indexer scores through `sparse_attn_indexer_kpool`, which has + # no FP4 arm, and keeps `block_size // index_kpool` rows per block -- + # so it would also hand the scorer a block size the cache never used. + why = f"index_kpool is {config.index_kpool}, and pooling has no FP4 scorer" + elif (gfx := _gfx()) != "gfx950": + why = f"{gfx} does not have the FP4 mqa-logits kernels" + else: + return True + if warn: + logger.warning("FP4 sparse indexer unavailable (%s); scoring in FP8.", why) + return False + + +def assert_fp4_indexer_supported( + *, fused_writer: bool, prefill_context_parallel: bool, prefill_ubatching: bool +) -> None: + """Reject the FP4 requests this build cannot serve. + + Both are knobs the user set, not geometry we can read off the config, so + they raise instead of falling back: quietly ignoring `--index-cache-dtype + fp4` would hide which of the two cost the memory the flag was asked for. + """ + if not fused_writer: + raise ValueError( + "The FP4 sparse indexer requires the fused QK/RoPE/cache kernel, " + "the only writer of the packed E2M1 planes. Either " + "ATOM_DISABLE_DS_INDEXER_QK_ROPE_CACHE_FUSION is set, which unset " + "restores it, or this indexer has no fused kernel to begin with: " + "that needs head_dim == quant_block_size and rope_dim == " + "head_dim // 2, so a NoPE indexer has no FP4 route at all. Pass " + "--index-cache-dtype fp8." + ) + if prefill_context_parallel: + raise ValueError( + "The FP4 sparse indexer does not support PCP, whose candidate " + "exchange is the one reader of the fp32 `weights` the FP4 writer " + "does not produce. Pass --index-cache-dtype fp8." + ) + if prefill_ubatching: + raise ValueError( + "The FP4 sparse indexer does not support prefill micro-batching. " + "`split_attn_metadata` rebuilds the metadata from its declared " + "fields, and the FP4 schedule rides on undeclared ones whose row " + "ids a ubatch rebases. Decode micro-batching (--enable-tbo-decode) " + "is supported. Pass --index-cache-dtype fp8." + ) + + +def fp4_decode_parallel_units(max_bs: int, next_n: int) -> int: + """CTAs the varctx schedule is built for -- the grid a CUDAGraph captures. + + `compute_varctx_schedule` needs a multiple of `next_n` leaving at least one + slot per sequence, so the floor rounds up to both. Not monotonic in `next_n` + -- at 512 units `f(3)` is 513 against `f(4)`'s 512. What lets one buffer + serve every width is the weaker `f(next_n) >= f(1)`, which holds because + both arms of the max scale with `next_n`: the buffer is sized at + `max_seqlen_qo` and the draft, the only caller asking for a different width, + asks for 1. + """ + return next_n * max(-(-FP4_MQA_VARCTX_PARALLEL_UNIT_NUM // next_n), max_bs) + + +def fp4_q_scale_shape(tokens: int, heads: int, head_dim: int) -> tuple: + """`q_scale_out`'s shape: `[T, k_tiles, 4, 16, round_up(H // 16, 4)]`. + + The trailing pad is the dword a lane loads its four M-tile scale bytes with, + so it is four even where the head count needs two. + """ + m_tiles = heads // _MFMA_M + return (tokens, head_dim // _K_TILE, 4, _MFMA_M, -(-m_tiles // 4) * 4) + + +def fp4_index_scale_rows(rows: torch.Tensor, block_size: int) -> torch.Tensor: + """Where logical block rows `rows` sit on the e8m0 plane's row axis. + + `indexer_qk_rope_quant_and_cache` stores that axis as an `_MFMA_M`-wide + transpose of the packed plane's, which is flat. Everything else addresses + both planes by the same logical row, so any reader that moves the two + together has to bend exactly here or it mixes exponents across a block -- + silently, since every index stays in bounds. + + `block_size` is the caller's own page size, taken rather than assumed: the + lane count is a property of the plane the kernel wrote, and a caller paging + the cache differently would otherwise get a wrong mapping that is still + in-bounds. `fp4_index_block_shapes` is what holds the two equal. + """ + if block_size != FP4_KV_BLOCK_SIZE: + raise ValueError( + f"the FP4 e8m0 row swizzle describes {FP4_KV_BLOCK_SIZE}-row blocks, " + f"got {block_size}" + ) + lanes = FP4_KV_BLOCK_SIZE // _MFMA_M + return (rows % _MFMA_M) * lanes + rows // _MFMA_M + + +def fp4_index_block_shapes(rows: int, head_dim: int) -> tuple[tuple, tuple]: + """One block's packed-E2M1 and e8m0 shapes, for one indexer layer.""" + if rows != FP4_KV_BLOCK_SIZE: + raise ValueError( + f"The FP4 sparse indexer requires --block-size {FP4_KV_BLOCK_SIZE}, " + f"got {rows}" + ) + k_tiles = head_dim // _K_TILE + return (k_tiles, 4, rows, 16), (k_tiles, 4, rows) + + +def fp4_prefill_schedule( + row_to_batch: torch.Tensor, + local_ends: torch.Tensor, + block_k: int, + parallel_floor: int, + max_seq_len: int, + local_starts: torch.Tensor | None = None, +) -> tuple[torch.Tensor, int, torch.Tensor]: + """Schedule one ragged-prefill forward, and the `local_starts` it scores from. + + `local_ends` is each row's causal upper bound in the column space the paged + FP4 scorer emits -- DeepSeek-V4's `visible_end`, MLA's `cu_seqlen_ke - + cu_seqlen_ks`. Rows start at column 0 there, so `local_starts` defaults to + zeros; pass it when the columns are one concatenated plane instead, as they + are under DCP, where the scorer reads a gathered copy of every sequence. + + `parallel_floor` is raised to the row count: prefill has one row per query + token and every (row, chunk-split) needs a slot. `max_seq_len` has to be the + width the scorer allocates its logits at, which the schedule bakes in. + """ + from aiter.ops.flydsl.kernels.mqa_logits.pa_mqa_logits_fp4_prefill import ( + compute_prefill_schedule, + ) + + if local_starts is None: + local_starts = torch.zeros_like(local_ends) + _, cta_info, n_ctas = compute_prefill_schedule( + row_to_batch.to(torch.int32), + local_starts, + local_ends, + block_k, + max(parallel_floor, local_ends.shape[0]), + max_seq_len, + ) + return cta_info, n_ctas, local_starts + + +def fp4_decode_schedule( + context_lens: torch.Tensor, + block_k: int, + parallel_units: int, + max_seq_len: int, + next_n: int, + cta_info_out: torch.Tensor, +) -> None: + """Refresh a decode step's schedule in place, at a CUDAGraph-stable address. + + `cta_info_out` has to be the buffer the captured kernel was handed and + `parallel_units` its row count: the grid is baked at capture, so only the + contents may change between replays. + """ + from aiter.ops.flydsl.kernels.mqa_logits.pa_mqa_logits_fp4 import ( + compute_varctx_schedule, + ) + + compute_varctx_schedule( + context_lens, + block_k, + parallel_units, + max_seq_len, + next_n=next_n, + cta_info_out=cta_info_out, + ) diff --git a/atom/models/deepseek_v2.py b/atom/models/deepseek_v2.py index 78e3b8b331..34b15e8408 100644 --- a/atom/models/deepseek_v2.py +++ b/atom/models/deepseek_v2.py @@ -108,6 +108,15 @@ use_triton_gemm, ) from atom.model_ops.moe import FusedMoE +from atom.model_ops.sparse_indexer_fp4 import ( + FP4_MQA_BLOCK_K, + FP4_MQA_PARALLEL_UNIT_NUM, + FP4_QUANT_BLOCK_SIZE, + assert_fp4_indexer_supported, + fp4_index_scale_rows, + fp4_q_scale_shape, + sparse_indexer_fp4_enabled, +) from atom.model_ops.topK import is_rocm_aiter_fusion_shared_expert_enabled from atom.model_ops.utils import MXFP4_QUANT_BLOCK_SIZE, atom_parameter from atom.models.utils import ( @@ -1456,6 +1465,108 @@ def _dcp_gather_indexer_k_prefill( return k_fp8, k_scale +def _dcp_stage_indexer_fp4_prefill( + kv_cache: torch.Tensor, + kv_cache_scale: torch.Tensor, + prefill_metadata, + total_kv: int, + block_size: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Stage the whole key set's FP4 index planes into pages of this rank's own. + + Same three steps as `_dcp_gather_indexer_k_prefill` -- read the local shard, + all-gather it, de-interleave back to global order -- on the two E2M1/e8m0 + planes instead of the one FP8 one. It ends in a paged buffer rather than a + flat one because every FP4 mqa-logits kernel is paged; over an identity + block table, column j of the scores is then flat KV index j, the space + `cu_seqlen_ks/ke` and the DCP prefill filter already speak. + + The two planes disagree on their row axis, so both the read and the write + bend through `fp4_index_scale_rows`; see it for what goes wrong otherwise. + """ + slots = prefill_metadata.dcp_indexer_fp4_local_slots + page, row = slots // block_size, slots % block_size + data = kv_cache[page, :, :, row, :] + scale = kv_cache_scale[page, :, :, fp4_index_scale_rows(row, block_size)] + + dcp_group = get_dcp_group() + gather_index = prefill_metadata.dcp_indexer_gather_index + data = dcp_group.all_gather(data, dim=0).index_select(0, gather_index) + scale = dcp_group.all_gather(scale, dim=0).index_select(0, gather_index) + + token = torch.arange(total_kv, device=data.device) + page, row = token // block_size, token % block_size + pages = -(-total_kv // block_size) + staged = kv_cache.new_zeros(pages, *kv_cache.shape[1:]) + staged[page, :, :, row, :] = data + staged_scale = kv_cache_scale.new_zeros(pages, *kv_cache_scale.shape[1:]) + staged_scale[page, :, :, fp4_index_scale_rows(row, block_size)] = scale + return staged, staged_scale + + +def _prefill_mqa_logits_fp4( + prefill_metadata, + chunk: slice, + whole_batch: bool, + q_fp4: torch.Tensor, + q_scale: torch.Tensor, + weights: torch.Tensor, + kv_cache: torch.Tensor, + kv_scale: torch.Tensor, + weights_scale: float, + kv_block_size: int, + block_tables: torch.Tensor, +) -> torch.Tensor: + """One chunk of ragged-prefill FP4 logits, scored out of the paged cache. + + A split chunk rebuilds the schedule rather than slicing the forward's: + `cta_info` encodes absolute row ids, so a slice would drop its rows. + """ + from aiter.ops.flydsl.kernels.mqa_logits.pa_mqa_logits_fp4_prefill import ( + compute_prefill_schedule, + flydsl_pa_mqa_logits_fp4_prefill, + ) + + rows = prefill_metadata.batch_id_per_q_token[chunk] + starts = prefill_metadata.indexer_fp4_local_starts[chunk] + ends = prefill_metadata.indexer_fp4_local_ends[chunk] + width = prefill_metadata.indexer_fp4_max_seq_len + if whole_batch: + cta_info = prefill_metadata.indexer_fp4_cta_info + n_ctas = prefill_metadata.indexer_fp4_n_ctas + else: + _, cta_info, n_ctas = compute_prefill_schedule( + rows, + starts, + ends, + FP4_MQA_BLOCK_K, + max(FP4_MQA_PARALLEL_UNIT_NUM, q_fp4.shape[0]), + width, + ) + logits = torch.empty( + q_fp4.shape[0], width, dtype=torch.float32, device=q_fp4.device + ) + flydsl_pa_mqa_logits_fp4_prefill( + q_fp4, + q_scale, + kv_cache, + kv_scale, + block_tables, + weights, + rows, + starts, + ends, + width, + weight_scale=weights_scale, + block_k=FP4_MQA_BLOCK_K, + kv_block_size=kv_block_size, + out=logits, + cta_info=cta_info, + n_ctas=n_ctas, + ) + return logits + + def sparse_attn_indexer( hidden_states: torch.Tensor, k_cache_prefix: str, @@ -1506,7 +1617,17 @@ def sparse_attn_indexer( ) runner_block_size = get_current_atom_config().kv_cache_block_size cp_kv_cache_interleave_size = get_current_atom_config().dcp_config.interleave_size - kv_cache = kv_cache.view(-1, runner_block_size, kv_cache.shape[-1]) + # Which width this indexer's cache is in -- the one thing the arguments + # cannot say, and a thing the traced graph piece already committed to. + fwd_ctx = get_current_atom_config().compilation_config.static_forward_context + # A custom op cannot take a module, so the Indexer is reached by undoing the + # one name `Indexer.__init__` builds: `prefix` + ".k_cache". Both spellings + # have to move together; a miss is a KeyError here, at the first forward. + indexer_module = fwd_ctx[k_cache_prefix.rsplit(".k_cache", 1)[0]] + indexer_fp4 = indexer_module._indexer_fp4 + q_fp4_scale = None + if not indexer_fp4: + kv_cache = kv_cache.view(-1, runner_block_size, kv_cache.shape[-1]) # PCP prefill: `k` (and `positions`) arrive as the full PADDED key set # [S_pad] produced by an all-gather of the round-robin shards. The KV-cache # write (driven by slot_mapping) and the gathered-KV sizing (total_kv = @@ -1523,15 +1644,63 @@ def sparse_attn_indexer( k = k[:n_real] if positions is not None: positions = positions[:n_real] - if use_qk_rope_cache_fusion: + if indexer_fp4: + # `weights_out` comes back unscaled in q's dtype: the FP4 mqa-logits + # kernels apply the per-group q scale inside the MFMA and take + # `weights_scale` as their own fp32 scalar, so either folded in here + # would be applied twice. `preshuffle` has no FP4 spelling. + q_bf16 = q_input + q_quant = torch.empty( + (*q_bf16.shape[:-1], head_dim // 2), device=q_bf16.device, dtype=torch.uint8 + ) + q_fp4_scale = torch.empty( + fp4_q_scale_shape(q_bf16.shape[0], q_bf16.shape[1], head_dim), + device=q_bf16.device, + dtype=torch.uint8, + ) + # Zeroed, not empty: this call leaves `compute_all_q_rope` default, so + # the op skips `slot < 0` rows outright, while the decode scorer reads + # the full `batch_size * next_n`. A sequence short of the speculation + # width would otherwise weight its pad rows with whatever the allocator + # held. Zero is also the right weight for a row whose logits go unread. + weights_mqa = torch.zeros_like(weights) + indexer_qk_rope_quant_and_cache( + q_bf16, + q_quant, + weights, + weights_mqa, + k, + kv_cache, + slot_mapping, + k_norm_weight, + k_norm_bias, + positions, + cos_cache, + sin_cache, + k_norm_eps, + FP4_QUANT_BLOCK_SIZE, + scale_fmt, + weights_scale, + is_neox=is_neox_style, + q_scale_out=q_fp4_scale, + kv_cache_scale=indexer_module.k_cache.kv_cache_scale, + ) + # Only this op's fp32 *return* is synthesised. The kernel's `weights_out` + # must stay `q.dtype` under FP4 (`aiter/ops/cache.py`), so `weights_mqa` + # cannot simply be allocated fp32; converting it instead would put a copy + # kernel in all 21 captured layers. Zeroed rather than empty: what makes + # it unread is a refusal three files away in `Indexer.__init__`, while + # `sparse_attn_indexer_fake` promises torch.compile a real tensor. + weights = torch.zeros(weights.shape, device=weights.device, dtype=torch.float32) + elif use_qk_rope_cache_fusion: q_bf16 = q_input - q_fp8 = torch.empty_like(q_bf16, dtype=dtypes.fp8) + q_quant = torch.empty_like(q_bf16, dtype=dtypes.fp8) weights_out = torch.empty( weights.shape, device=weights.device, dtype=torch.float32 ) indexer_qk_rope_quant_and_cache( q_bf16, - q_fp8, + q_quant, weights, weights_out, k, @@ -1551,7 +1720,7 @@ def sparse_attn_indexer( ) weights = weights_out else: - q_fp8 = q_input + q_quant = q_input indexer_k_quant_and_cache( k, kv_cache, @@ -1560,6 +1729,8 @@ def sparse_attn_indexer( scale_fmt, preshuffle=True, ) + if not indexer_fp4: + weights_mqa = weights if context.is_prefill: # Below index_topk the indexer is a no-op: top-k would select every token, # so prefill runs dense (attention_mla.use_prefill_mla gates on the same @@ -1584,7 +1755,23 @@ def sparse_attn_indexer( dtype=torch.long, device=prefill_metadata.block_tables.device, ) - if get_dcp_world_size() > 1: + if indexer_fp4: + # The paged FP4 scorer reads the cache in place -- except under DCP, + # where in place is only this rank's 1/W of the sequence. + k_fp8 = k_scale = None + fp4_kv_cache = kv_cache + fp4_kv_scale = indexer_module.k_cache.kv_cache_scale + fp4_block_tables = prefill_metadata.block_tables + if get_dcp_world_size() > 1: + fp4_kv_cache, fp4_kv_scale = _dcp_stage_indexer_fp4_prefill( + kv_cache, + fp4_kv_scale, + prefill_metadata, + total_kv, + runner_block_size, + ) + fp4_block_tables = prefill_metadata.dcp_indexer_fp4_block_tables + elif get_dcp_world_size() > 1: k_fp8, k_scale = _dcp_gather_indexer_k_prefill( kv_cache, prefill_metadata, head_dim, k.device ) @@ -1603,22 +1790,30 @@ def sparse_attn_indexer( ), preshuffle=True, ) - cu_seqlen_ks = prefill_metadata.cu_seqlen_ks - cu_seqlen_ke = prefill_metadata.cu_seqlen_ke + # Per-row window bounds, in the column space this scorer emits. + if indexer_fp4: + cu_seqlen_ks = prefill_metadata.indexer_fp4_local_starts + cu_seqlen_ke = prefill_metadata.indexer_fp4_local_ends + else: + cu_seqlen_ks = prefill_metadata.cu_seqlen_ks + cu_seqlen_ke = prefill_metadata.cu_seqlen_ke num_tokens = hidden_states.shape[0] - q_prefill = q_fp8[num_decode_tokens:num_tokens] - weights_prefill = weights[num_decode_tokens:num_tokens] + q_prefill = q_quant[num_decode_tokens:num_tokens] + weights_prefill = weights_mqa[num_decode_tokens:num_tokens] num_rows = q_prefill.shape[0] assert topk_tokens == 2048, "top_k_per_row assumes size 2048" topk_indices_prefill = topk_indices[num_decode_tokens:num_tokens, :topk_tokens] - # The dense logits buffer is [num_rows, total_kv] fp32. total_kv is the - # sum of all co-scheduled prefill contexts and is unbounded by - # max_num_batched_tokens, so a burst of long-context requests can push a - # single allocation to tens of GiB (#1376). Under chunked prefill - # num_rows is already capped by max_num_batched_tokens, so the OOM is - # driven by total_kv (the column dim). Chunk along the Q (query-row) - # dimension with q_chunk sized so the buffer [q_chunk, total_kv] fp32 - # stays within the memory budget — q_chunk shrinks as total_kv grows. + row_width = ( + prefill_metadata.indexer_fp4_max_seq_len if indexer_fp4 else total_kv + ) + # The dense logits buffer is [num_rows, row_width] fp32. For FP8 that + # width is total_kv, the sum of all co-scheduled prefill contexts, and + # is unbounded by max_num_batched_tokens, so a burst of long-context + # requests can push a single allocation to tens of GiB (#1376). Under + # chunked prefill num_rows is already capped by max_num_batched_tokens, + # so the OOM is driven by the column dim. Chunk along the Q (query-row) + # dimension with q_chunk sized so the buffer [q_chunk, row_width] fp32 + # stays within the memory budget — q_chunk shrinks as row_width grows. # Each chunk still scores the FULL KV, so every row's top-k is computed # completely in one shot: the result is exact with no cross-chunk merge, # the kernel's column indices are already global (no remapping), and each @@ -1628,17 +1823,17 @@ def sparse_attn_indexer( budget_bytes = SPARSE_INDEXER_LOGITS_BUDGET_MB * 1024 * 1024 if ( budget_bytes > 0 - and total_kv > 0 - and budget_bytes // (total_kv * 4) < num_rows + and row_width > 0 + and budget_bytes // (row_width * 4) < num_rows ): - # 4 bytes per fp32 logit; total_kv * 4 is one query row's footprint. + # 4 bytes per fp32 logit; row_width * 4 is one query row's footprint. # Round the budget-derived row count DOWN to keep the buffer within # budget: a multiple of 128 (aligned to the kernel's row tiling) in # the normal regime, avoiding the coarse power-of-2 doubling. When - # the budget affords < 128 rows (extreme total_kv), fall back to a + # the budget affords < 128 rows (extreme row_width), fall back to a # power-of-2 floor so it degrades to 64/32/.../1 instead of # collapsing straight to 1. - budget_rows = budget_bytes // (total_kv * 4) + budget_rows = budget_bytes // (row_width * 4) if budget_rows >= 128: chunk_tokens = (budget_rows // 128) * 128 else: @@ -1651,14 +1846,30 @@ def sparse_attn_indexer( # Per-row window bounds slice 1:1 with this chunk's rows. row_starts = cu_seqlen_ks[chunk_start:chunk_end] row_ends = cu_seqlen_ke[chunk_start:chunk_end] - logits = fp8_mqa_logits( - Q=q_prefill[chunk_start:chunk_end], - KV=k_fp8, - kv_scales=k_scale, - weights=weights_prefill[chunk_start:chunk_end], - cu_starts=row_starts, - cu_ends=row_ends, - ) + if indexer_fp4: + chunk = slice(chunk_start, chunk_end) + logits = _prefill_mqa_logits_fp4( + prefill_metadata, + chunk, + chunk_tokens == num_rows, + q_prefill[chunk], + q_fp4_scale[num_decode_tokens:num_tokens][chunk], + weights_prefill[chunk], + fp4_kv_cache, + fp4_kv_scale, + weights_scale, + runner_block_size, + fp4_block_tables, + ) + else: + logits = fp8_mqa_logits( + Q=q_prefill[chunk_start:chunk_end], + KV=k_fp8, + kv_scales=k_scale, + weights=weights_prefill[chunk_start:chunk_end], + cu_starts=row_starts, + cu_ends=row_ends, + ) top_k_per_row_prefill( logits=logits, rowStarts=row_starts, @@ -1703,21 +1914,23 @@ def sparse_attn_indexer( NUM_TOPK_TOKENS=topk_tokens, PAGE_SIZE=runner_block_size, out=sparse_kv_indices_buffer, + seq_local=indexer_fp4, ) else: decode_metadata = attn_metadata - # kv_cache size requirement [num_block, block_size, n_head, head_dim], - # we only have [num_block, block_size, head_dim], - kv_cache = kv_cache.unsqueeze(-2) - padded_q_fp8_decode_tokens = q_fp8[:num_decode_tokens].reshape( - context.scheduled_bs, -1, *q_fp8.shape[1:] + if not indexer_fp4: + # kv_cache size requirement [num_block, block_size, n_head, head_dim], + # we only have [num_block, block_size, head_dim], + kv_cache = kv_cache.unsqueeze(-2) + padded_q_decode_tokens = q_quant[:num_decode_tokens].reshape( + context.scheduled_bs, -1, *q_quant.shape[1:] ) # TODO: move and optimize below logic with triton kernels - batch_size = padded_q_fp8_decode_tokens.shape[0] - next_n = padded_q_fp8_decode_tokens.shape[1] + batch_size = padded_q_decode_tokens.shape[0] + next_n = padded_q_decode_tokens.shape[1] assert batch_size == context.scheduled_bs num_padded_tokens = batch_size * next_n - batch_size, next_n, _heads, _ = padded_q_fp8_decode_tokens.shape + batch_size, next_n, _heads, _ = padded_q_decode_tokens.shape num_rows = batch_size * next_n dcp_world_size = get_dcp_world_size() assert topk_tokens == 2048, "top_k_per_row assumes size 2048" @@ -1730,9 +1943,9 @@ def sparse_attn_indexer( # non-DCP path does below has already happened, and we return here. dcp_decode_candidate_exchange_fused( attn_metadata, - padded_q_fp8_decode_tokens, + padded_q_decode_tokens, kv_cache, - weights, + weights_mqa, get_dcp_rank(), num_decode_tokens, topk_tokens, @@ -1743,6 +1956,11 @@ def sparse_attn_indexer( out_kv_indices=sparse_kv_indices_buffer, out_kv_indptr=dcp_sparse_kv_indptr_buffer, owned_counts=dcp_owned_counts_buffer, + q_scale=q_fp4_scale, + kv_scale=( + indexer_module.k_cache.kv_cache_scale if indexer_fp4 else None + ), + weights_scale=weights_scale, ) return weights # Non-DCP: this rank holds the whole plane, so its top-k is already the @@ -1750,17 +1968,40 @@ def sparse_attn_indexer( logits = torch.empty( [num_rows, max_model_len], dtype=torch.float32, device="cuda" ) - deepgemm_fp8_paged_mqa_logits( - padded_q_fp8_decode_tokens, - kv_cache, - weights[:num_padded_tokens], - logits, - decode_metadata.context_lens, - attn_metadata.block_tables, - max_model_len, - KVBlockSize=runner_block_size, - Preshuffle=True, - ) + if indexer_fp4: + from aiter.ops.flydsl import flydsl_pa_mqa_logits_fp4 + + flydsl_pa_mqa_logits_fp4( + padded_q_decode_tokens, + q_fp4_scale[:num_decode_tokens].reshape( + batch_size, next_n, *q_fp4_scale.shape[1:] + ), + kv_cache, + indexer_module.k_cache.kv_cache_scale, + attn_metadata.block_tables, + weights_mqa[:num_padded_tokens], + decode_metadata.context_lens, + max_model_len, + weight_scale=weights_scale, + next_n=next_n, + block_k=FP4_MQA_BLOCK_K, + kv_block_size=runner_block_size, + out=logits, + cta_info=decode_metadata.indexer_fp4_cta_info, + total_ctas=decode_metadata.indexer_fp4_n_ctas, + ) + else: + deepgemm_fp8_paged_mqa_logits( + padded_q_decode_tokens, + kv_cache, + weights[:num_padded_tokens], + logits, + decode_metadata.context_lens, + attn_metadata.block_tables, + max_model_len, + KVBlockSize=runner_block_size, + Preshuffle=True, + ) topk_indices_decode = topk_indices[:num_decode_tokens, :topk_tokens] top_k_per_row_decode( logits, @@ -2075,6 +2316,9 @@ def __init__( super().__init__() self.atom_config = atom_config self.config = config + self._indexer_fp4 = sparse_indexer_fp4_enabled( + atom_config.index_cache_dtype, config + ) # self.indexer_cfg = config.attn_module_list_cfg[0]["attn_index"] self.topk_tokens = config.index_topk self.n_head = config.index_n_heads # 64 @@ -2125,6 +2369,12 @@ def __init__( self.k_norm = LayerNorm(self.head_dim, eps=1e-6, dtype=torch.float32) self.softmax_scale = self.head_dim**-0.5 self._weights_scale = self.softmax_scale * self.n_head**-0.5 + if self._indexer_fp4: + assert_fp4_indexer_supported( + fused_writer=self.use_qk_rope_cache_fusion, + prefill_context_parallel=pcp_is_enabled(), + prefill_ubatching=get_current_atom_config().enable_tbo, + ) # TODO (zyongye) change dim to fp8 later to (self.head_dim + 4) self.k_cache = DeepseekV32IndexerCache( diff --git a/atom/spec_decode/eagle_proposer.py b/atom/spec_decode/eagle_proposer.py index 913607d904..5abf91ed34 100644 --- a/atom/spec_decode/eagle_proposer.py +++ b/atom/spec_decode/eagle_proposer.py @@ -476,6 +476,11 @@ def _enter_decode_metadata( # itself; the verify step's has the right row count, wrong rows. attn_metadata.dcp_token_block_tables = attn_metadata.block_tables attn_metadata.context_lens = var["context_lens"].gpu[:running_bs] + # Unguarded because the base builder answers it: this runs for every + # target a draft can have, and only the MLA one has an FP4 indexer. The + # draft's rows are its own, and a replay addresses whatever row count + # the schedule buffer holds. + builder._publish_indexer_fp4_decode_schedule(attn_metadata, running_bs, 1) if "sparse_kv_indptr" in var: attn_metadata.sparse_kv_indptr = var["sparse_kv_indptr"].gpu[ : running_bs + 1 @@ -780,6 +785,15 @@ def propose( ) for k, v in workinfos.items(): attn_metadata.__dict__[k] = v + # Every step, and only after `prepare_mtp_decode`: that call + # is what refreshes the DCP local lengths the schedule is + # built from, and a replay reads the buffer's contents. The + # publish in `_enter_decode_metadata` runs before the + # refresh, and steps 1+ reached their forward with step 0's. + builder = self.runner.attn_metadata_builder + builder._publish_indexer_fp4_decode_schedule( + attn_metadata, running_bs, 1 + ) if has_flat_kv and "slot_mapping" not in workinfos: # MLA/MHA path: slot derived from flat kv_indices. Both, # and the slot_mapping written below, are the ones @@ -789,7 +803,6 @@ def propose( raw_slots = attn_metadata.kv_indices[ attn_metadata.kv_indptr[1 : running_bs + 1] - 1 ] - builder = self.runner.attn_metadata_builder if getattr(builder, "dcp_world_size", 1) > 1: # DCP interleave-S: only rank ((ctx-1)//S) % W owns this # draft token; other ranks' kv_indptr didn't grow, so diff --git a/docs/configuration_guide.md b/docs/configuration_guide.md index 9a452d46c3..10502eff73 100644 --- a/docs/configuration_guide.md +++ b/docs/configuration_guide.md @@ -37,10 +37,10 @@ Defined in `atom/config.py`. The root dataclass that the engine consumes. | `tensor_parallel_size` | `int` | `1` | Number of tensor-parallel GPUs (1 — 8) | | `enforce_eager` | `bool` | `False` | Disable compilation and CUDA graphs; run in eager mode | | `parallel_config` | `ParallelConfig` | `ParallelConfig()` | Data-parallel configuration (see Section 4) | -| `kv_cache_block_size` | `int` | `16` | Block size for paged KV cache; must be a multiple of 16 or exactly 1 | +| `kv_cache_block_size` | `int` | `16` | Block size for paged KV cache; must be a multiple of 16 or exactly 1. The FP4 sparse indexer requires exactly 64 | | `num_kvcache_blocks` | `int` | `-1` | Number of KV cache blocks (`-1` = auto) | | `kv_cache_dtype` | `str` | `"bf16"` | KV cache data type (`"bf16"` or `"fp8"`) | -| `index_cache_dtype` | `str \| None` | `None` | Indexer-cache dtype, resolved after model detection. Native single-node DeepSeek-V4 defaults to `"fp4"` except on gfx942; plugin and KV-transfer integrations retain `"fp8"`. Other models inherit `kv_cache_dtype`. An explicit `"bf16"`, `"fp8"`, or `"fp4"` value is preserved. | +| `index_cache_dtype` | `str \| None` | `None` | Indexer-cache dtype, resolved after model detection. Native single-node DeepSeek-V4 defaults to `"fp4"` except on gfx942; plugin and KV-transfer integrations retain `"fp8"`. Other models inherit `kv_cache_dtype`. An explicit `"bf16"`, `"fp8"`, or `"fp4"` value is preserved. `"fp4"` is honoured only where the FP4 mqa-logits kernels apply -- gfx950, an indexer head dim of 128, an indexer head count that is a multiple of 16, and `--block-size 64`; anything else logs one warning and scores in FP8. It runs under DCP, and is refused outright under PCP and KV transfer rather than falling back. | | `enable_prefix_caching` | `bool` | `False` | Enable prefix caching to reuse KV blocks across requests sharing the same prefix | | `enable_log_stats` | `bool` | `True` | Emit the periodic engine-status line (running/waiting reqs, KV usage, prefix-cache hit rate, prompt/generation throughput). Applies to offline `LLM(...)` as well as to the server. Scoped to that line: `[MTP Stats]` and `[Cache Stats]` have their own gates | | `throughput_log_interval` | `float` | `10.0` | Seconds between engine-status lines. Must be > 0 | @@ -364,7 +364,7 @@ all flags via `add_cli_args()` and converts them into a `Config` via | `--port` | | `int` | `8006` | Engine internal port | | `--kv_cache_dtype` | | `str` | `"bf16"` | KV cache dtype; choices: `bf16`, `fp8` | | `--index-cache-dtype`, `--index_cache_dtype` | | `str` | `None` | Indexer-cache dtype; choices: `bf16`, `fp8`, `fp4`. When omitted, uses the architecture- and integration-aware `Config.index_cache_dtype` defaults described above. | -| `--block-size` | | `int` | `16` | KV cache block size (maps to `kv_cache_block_size`) | +| `--block-size` | | `int` | `16` | KV cache block size (maps to `kv_cache_block_size`); the FP4 sparse indexer requires exactly 64 | | `--max-model-len` | | `int` | `None` | Maximum model context length; defaults to `hf_config.max_position_embeddings` | | `--cudagraph-capture-sizes` | | `str` | `"[1,2,4,8,16,32,48,64,128,256]"` | CUDA graph capture sizes as a Python list string | | `--level` | | `int` | `3` | Compilation level (0 — 3) | diff --git a/docs/model_ops_guide.md b/docs/model_ops_guide.md index 3d9fbb6122..07b79528a0 100644 --- a/docs/model_ops_guide.md +++ b/docs/model_ops_guide.md @@ -14,6 +14,7 @@ ATOM (AiTer Optimized Model) wraps AITER kernels with model-level abstractions f | `Attention` | `base_attention.py` | `unified_attention_with_output_base` (custom op) | Unified attention entry | | MHA `Attention` | `attention_mha.py` | `flash_attn_varlen_func`, `pa_fwd_asm`, `pa_persistent_fwd`, `pa_decode_gluon` | Multi-head attention | | `MLAAttention` | `attention_mla.py` | `mla_decode_fwd`, `mla_prefill_fwd`, `concat_and_cache_mla`, `fused_qk_rope_concat_and_cache_mla` | Multi-head latent attention | +| DSA sparse indexer | `sparse_indexer_fp4.py` | `flydsl_pa_mqa_logits_fp4`, `indexer_qk_rope_quant_and_cache`, `top_k_per_row_decode` | FP4 (E2M1 + e8m0) index-cache geometry, scoring and CTA schedules | | `FusedMoE` | `moe.py` | `aiter.fused_moe.fused_moe`, `asm_moe` | Mixture of experts | | `RMSNorm` | `layernorm.py` | `rmsnorm2d_fwd`, `rmsnorm2d_fwd_with_add`, `fused_add_rmsnorm_pad` | RMS normalization | | `LayerNorm` | `layernorm.py` | `layernorm2d_fwd`, `layernorm2d_fwd_with_add` | Layer normalization | diff --git a/docs/scheduling_kv_cache_guide.md b/docs/scheduling_kv_cache_guide.md index c9b9cd19dd..cf9ac10c91 100644 --- a/docs/scheduling_kv_cache_guide.md +++ b/docs/scheduling_kv_cache_guide.md @@ -24,7 +24,7 @@ ATOM (AiTer Optimized Model) uses a prefill-first scheduler with paged KV cache |---|---|---| | `max_num_seqs` | 512 | Maximum sequences in a single batch | | `max_num_batched_tokens` | 16384 | Maximum tokens scheduled in a single step | -| `kv_cache_block_size` | 16 | Tokens per KV cache block (must be multiple of 16, or 1) | +| `kv_cache_block_size` | 16 | Tokens per KV cache block (must be multiple of 16, or 1; exactly 64 for the FP4 sparse indexer) | | `enable_prefix_caching` | `False` | Enable hash-based prefix block sharing | | `scheduler_delay_factor` | 0.0 | Delay factor for batching prompt requests (0 = no delay) | | `gpu_memory_utilization` | 0.9 | Fraction of GPU memory for KV cache | diff --git a/tests/test_dsv4_aiter_fp4_imports.py b/tests/test_dsv4_aiter_fp4_imports.py index 204d5cb3e8..9a8a6d8549 100644 --- a/tests/test_dsv4_aiter_fp4_imports.py +++ b/tests/test_dsv4_aiter_fp4_imports.py @@ -5,10 +5,11 @@ from pathlib import Path ATOM_ROOT = Path(__file__).resolve().parents[1] / "atom" -# One entry, not two: the decode and prefill FP4 scorers were unified onto the -# varqlen (`_prefill`) kernel family, so nothing imports the rectangular -# `pa_mqa_logits_fp4` any more. Widen this only alongside a caller. +# DeepSeek-V4 scores decode and prefill both through the varqlen (`_prefill`) +# family; the paged MLA indexer takes the rectangular kernel for decode, where +# every row of a step shares one width. Widen this only alongside a caller. CURRENT_MODULES = { + "aiter.ops.flydsl.kernels.mqa_logits.pa_mqa_logits_fp4", "aiter.ops.flydsl.kernels.mqa_logits.pa_mqa_logits_fp4_prefill", } diff --git a/tests/test_sparse_indexer_fp4.py b/tests/test_sparse_indexer_fp4.py new file mode 100644 index 0000000000..17510b7151 --- /dev/null +++ b/tests/test_sparse_indexer_fp4.py @@ -0,0 +1,645 @@ +"""The FP4 sparse indexer: the predicate, the ABI shapes, the KV pool, and one +component cross-check against the bytes the production writer emits. + +`indexer_qk_rope_quant_and_cache` in FP4 mode is the only writer of the packed +E2M1 Q/K and their e8m0 planes, and `flydsl_pa_mqa_logits_fp4[_prefill]` the +only readers, so what is worth checking is that the two agree on a real DSA +indexer's shapes (H=32, D=128, kv_block=64, block_k=256) -- decode at a +speculation width and at the DCP path's one row per query token, prefill down +both schedule paths. The reference dequantizes +exactly what the writer produced, so a disagreement is a layout bug, not +rounding. The FP8 default has to come out of all of it untouched. +""" + +import importlib +from types import SimpleNamespace + +import pytest +import torch + +from atom.model_ops import sparse_indexer_fp4 +from atom.model_ops.attentions.mla_kv_pool import MlaKvPool +from atom.model_ops.sparse_indexer_fp4 import ( + FP4_KV_BLOCK_SIZE, + FP4_MQA_BLOCK_K, + FP4_QUANT_BLOCK_SIZE, + assert_fp4_indexer_supported, + fp4_decode_parallel_units, + fp4_decode_schedule, + fp4_index_scale_rows, + fp4_prefill_schedule, + fp4_q_scale_shape, + sparse_indexer_fp4_enabled, +) + +# The DSA indexer geometry GLM-5.2 and DeepSeek-V3.2 share. +DSA = SimpleNamespace(index_topk=2048, index_n_heads=32, index_head_dim=128) + +HEADS, HEAD_DIM, _BLOCK = 32, 128, FP4_KV_BLOCK_SIZE +WEIGHTS_SCALE = HEAD_DIM**-0.5 * HEADS**-0.5 +_E2M1_MAG = torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0]) +_E2M1 = torch.cat([_E2M1_MAG, -_E2M1_MAG]) + + +def _import_or_skip(name: str, reason: str | None = None): + """Import `name`, skipping the test where this box cannot. + + Not `pytest.importorskip`: it warns on a plain `ImportError` rather than a + missing module -- which CI escalates, and pytest 9.1 will too -- and on a + CPU runner that is exactly how these fail, `aiter` importing but carrying no + `QuantType`. + """ + try: + return importlib.import_module(name) + except ImportError as exc: + pytest.skip(reason or f"{name} is unavailable here: {exc}") + + +def _pool(**overrides): + args = { + "layers": 2, + "block_size": 64, + "entry_dim": 576, + "kv_dtype": torch.bfloat16, + "index_layers": 3, + "index_rows_per_block": 64, + "index_dim": 144, + "index_dtype": torch.uint8, + "index_head_dim": 128, + } + args.update(overrides) + return MlaKvPool(**args) + + +def test_predicate_is_structural_and_never_probes_the_chip_for_fp8(monkeypatch): + monkeypatch.setattr(sparse_indexer_fp4, "_gfx", lambda: "gfx950") + assert sparse_indexer_fp4_enabled("fp4", DSA) + # An MTP draft: `_MTP_TYPE_MAP` rewrote its model_type but not its indexer, + # and it shares the target's cache, so it must reach the target's verdict. + assert sparse_indexer_fp4_enabled( + "fp4", SimpleNamespace(model_type="deepseek_mtp", **vars(DSA)) + ) + + monkeypatch.delattr(sparse_indexer_fp4, "_gfx") + assert not sparse_indexer_fp4_enabled("fp8", DSA) + assert not sparse_indexer_fp4_enabled(None, DSA) + + +@pytest.mark.parametrize( + ("override", "gfx", "why"), + [ + ({"index_topk": 0}, "gfx950", "no sparse indexer"), + ({"index_head_dim": 64}, "gfx950", "index_head_dim is 64"), + ({"index_n_heads": 24}, "gfx950", "index_n_heads is 24"), + ({"index_kpool": 2}, "gfx950", "index_kpool is 2"), + ({}, "gfx942", "gfx942"), + ], +) +def test_predicate_falls_back_and_names_what_blocked_it( + monkeypatch, caplog, override, gfx, why +): + monkeypatch.setattr(sparse_indexer_fp4, "_gfx", lambda: gfx) + config = SimpleNamespace(**{**vars(DSA), **override}) + assert not sparse_indexer_fp4_enabled("fp4", config) + with caplog.at_level("WARNING", logger="atom"): + assert not sparse_indexer_fp4_enabled("fp4", config, warn=True) + assert why in caplog.text + + +def test_unsupported_fp4_requests_name_the_knob_that_blocked_them(): + def check(**overrides): + assert_fp4_indexer_supported( + **{ + "fused_writer": True, + "prefill_context_parallel": False, + "prefill_ubatching": False, + **overrides, + } + ) + + check() + with pytest.raises(ValueError, match="fused QK/RoPE/cache"): + check(fused_writer=False) + # PCP's candidate exchange is the only reader of the indexer op's return, so + # the FP4 path may leave that tensor unwritten only while this refusal holds. + with pytest.raises(ValueError, match="does not support PCP"): + check(prefill_context_parallel=True) + # Decode micro-batching is supported; only the prefill split is not, because + # it rebuilds metadata from declared fields and the schedule rides on + # undeclared ones. A blanket TBO refusal would take the decode path with it. + with pytest.raises(ValueError, match="prefill micro-batching"): + check(prefill_ubatching=True) + + +def test_every_backend_answers_the_draft_s_fp4_schedule_publish(): + """`EagleProposer` refreshes this on whatever builder the target uses, and + only the MLA one has an FP4 indexer to refresh. Every other backend has to + answer it anyway: EAGLE3 on Llama-3, MTP on Qwen3-Next and on DeepSeek-V4 + (whose builder is a `CommonAttentionBuilder` sibling, not an MLA subclass) + all reach that line with FP4 nowhere in the picture.""" + backends = _import_or_skip("atom.model_ops.attentions.backends") + base = backends.CommonAttentionBuilder._publish_indexer_fp4_decode_schedule + + mla = _import_or_skip("atom.model_ops.attentions.aiter_mla") + assert mla.AiterMLAMetadataBuilder._publish_indexer_fp4_decode_schedule is not base + + for module, name in ( + ("atom.model_ops.attentions.aiter_attention", "AiterAttentionMetadataBuilder"), + ( + "atom.model_ops.attentions.deepseek_v4_attn", + "DeepseekV4AttentionMetadataBuilder", + ), + ("atom.model_ops.attentions.gdn_attn", "GDNAttentionMetadataBuilder"), + ("atom.model_ops.attentions.triton_mha", "TritonMHAMetadataBuilder"), + ): + builder = getattr(_import_or_skip(module), name) + assert builder._publish_indexer_fp4_decode_schedule is base, name + + # Inert, not merely present: the draft reuses the target's metadata object, + # so anything written here would reach the verify step. + metadata = SimpleNamespace() + base(object(), metadata, 4, 1) + assert not vars(metadata) + + +def test_the_builder_compares_the_indexer_s_fp4_verdict_instead_of_setting_it(): + """The builder and `Indexer.__init__` answer the same predicate from the + same two inputs, and the Indexer has already built `k_cache` from its answer + by the time the builder reaches it. Assigning over it cannot fix a + divergence -- the object is already built -- it only hides one, to surface + later as a graph/eager dtype mismatch. Read off the source because reaching + that line needs an allocated pool and a loaded model.""" + import ast + import pathlib + + root = pathlib.Path(sparse_indexer_fp4.__file__).parent + tree = ast.parse((root / "attentions" / "aiter_mla.py").read_text(encoding="utf-8")) + overwrites = [ + node.lineno + for node in ast.walk(tree) + if isinstance(node, ast.Assign) + for target in node.targets + if isinstance(target, ast.Attribute) + and target.attr == "_indexer_fp4" + and isinstance(target.value, ast.Attribute) + and target.value.attr == "indexer" + ] + assert not overwrites, f"builder overwrites the Indexer's verdict at {overwrites}" + + +def test_decode_parallel_units_cover_the_batch_at_every_speculation_width(): + """Everything the one captured buffer rests on: a slot per sequence per + step, the varctx floor, and `f(n) >= f(1)` -- the buffer is sized at + `max_seqlen_qo` and the draft then asks at 1. That last one is the weakest + true statement, not the obvious one: `f` is NOT monotonic in `next_n`.""" + for max_bs in (1, 7, 16, 64, 128, 300, 512, 8192): + floor = fp4_decode_parallel_units(max_bs, 1) + for next_n in range(1, 17): + units = fp4_decode_parallel_units(max_bs, next_n) + assert units % next_n == 0 + assert units // next_n >= max_bs + assert units >= sparse_indexer_fp4.FP4_MQA_VARCTX_PARALLEL_UNIT_NUM + assert units >= floor + + # The counterexample the docstring names, so it cannot rot back. + assert fp4_decode_parallel_units(1, 3) > fp4_decode_parallel_units(1, 4) + + +def test_q_scale_shape_pads_the_m_tile_axis_to_one_dword(): + # H=32 is two M-tiles, still loaded as one dword of four scale bytes. + assert fp4_q_scale_shape(7, 32, 128) == (7, 1, 4, 16, 4) + assert fp4_q_scale_shape(7, 64, 128) == (7, 1, 4, 16, 4) + assert fp4_q_scale_shape(7, 128, 128) == (7, 1, 4, 16, 8) + + +def test_index_field_narrows_under_fp4_without_adding_an_arena(): + fp8, fp4 = _pool(), _pool(index_fp4=True) + assert len(fp8.field_groups) == len(fp4.field_groups) == 2 + assert [f.name for f in fp8.index_fields] == ["index"] + assert [f.name for f in fp4.index_fields] == ["index", "index_scale"] + assert fp8.entry_bytes == 2 * 64 * 576 * 2 + 3 * 64 * 144 + assert fp4.entry_bytes == 2 * 64 * 576 * 2 + 3 * (4 * 64 * 16) + 3 * (4 * 64) + assert fp4.entry_bytes < fp8.entry_bytes + + fp8.allocate(3, "cpu") + fp4.allocate(3, "cpu") + assert fp8.layer("index", 0).shape == (3, 64, 144) + data, scale = fp4.layer("index", 0), fp4.layer("index_scale", 0) + assert data.shape == (3, 1, 4, 64, 16) and data.dtype is torch.uint8 + assert scale.shape == (3, 1, 4, 64) and scale.dtype is torch.uint8 + spans = [ + (t.data_ptr(), t.data_ptr() + t.numel() * t.element_size()) + for t in (fp4.layer("kv", 0), data, scale) + ] + for i, lhs in enumerate(spans): + for rhs in spans[i + 1 :]: + assert lhs[1] <= rhs[0] or rhs[1] <= lhs[0] + + # The indexer rows per block are constrained, the KV block size is not. + _pool(index_fp4=True, block_size=32, index_rows_per_block=64) + with pytest.raises(ValueError, match="--block-size 64"): + _pool(index_fp4=True, index_rows_per_block=32) + + +@pytest.fixture +def on_gfx950(monkeypatch): + """Gate the cross-check on a chip that has the kernels, single-rank.""" + if not torch.cuda.is_available(): + pytest.skip("requires a ROCm GPU") + from aiter.jit.utils.chip_info import get_gfx + + if get_gfx() != "gfx950": + pytest.skip("the FP4 paged-MQA-logits kernels are gfx950-only") + # ATOM's shim reads the DCP world size off the global config, which a unit + # test has no reason to build. + from atom.model_ops import attention_mla + + monkeypatch.setattr(attention_mla, "get_dcp_world_size", lambda: 1) + + +def _dequant(packed: torch.Tensor, e8m0: torch.Tensor) -> torch.Tensor: + """`[..., D // 2]` E2M1 pairs plus `[..., D // 32]` e8m0 -> fp32 `[..., D]`.""" + nibbles = torch.empty( + *packed.shape[:-1], packed.shape[-1] * 2, dtype=torch.long, device=packed.device + ) + nibbles[..., 0::2] = packed & 0xF + nibbles[..., 1::2] = packed >> 4 + scale = torch.exp2(e8m0.float() - 127.0).repeat_interleave( + FP4_QUANT_BLOCK_SIZE, dim=-1 + ) + return _E2M1.to(packed.device)[nibbles] * scale + + +def _paged_layout(batch: int, ctx_len: int): + """A shuffled block table plus the slot of every KV token.""" + blocks_per_seq = ctx_len // _BLOCK + num_blocks = batch * blocks_per_seq + table = torch.randperm(num_blocks, device="cuda").to(torch.int32) + table = table.reshape(batch, blocks_per_seq) + token = torch.arange(ctx_len, device="cuda").repeat(batch) + seq = torch.arange(batch, device="cuda").repeat_interleave(ctx_len) + slots = table[seq, token // _BLOCK].long() * _BLOCK + token % _BLOCK + return table, num_blocks, token, seq, slots + + +def _fused_fp4(slots, positions, num_blocks, weight_gain=1.0): + """The production writer, in FP4 mode. Returns everything it emits. + + Through ATOM's own shim rather than `aiter.` directly: the shim is what + decides `compute_all_q_rope` and forwards the two scale buffers, so it is + the seam worth covering. + """ + from atom.model_ops import attention_mla + + rows = slots.shape[0] + u8 = {"dtype": torch.uint8, "device": "cuda"} + bf16 = {"dtype": torch.bfloat16, "device": "cuda"} + angles = torch.randn(4096, 32, device="cuda") + norm = torch.randn(HEAD_DIM, dtype=torch.float32, device="cuda") + weights = (torch.randn(rows, HEADS, device="cuda") * weight_gain).bfloat16() + q_fp4 = torch.zeros(rows, HEADS, HEAD_DIM // 2, **u8) + q_scale = torch.zeros(fp4_q_scale_shape(rows, HEADS, HEAD_DIM), **u8) + weights_out = torch.zeros_like(weights) + kv_cache = torch.zeros(num_blocks, 1, 4, _BLOCK, 16, **u8) + kv_scale = torch.zeros(num_blocks, 1, 4, _BLOCK, **u8) + attention_mla.indexer_qk_rope_quant_and_cache( + torch.randn(rows, HEADS, HEAD_DIM, **bf16), + q_fp4, + weights, + weights_out, + torch.randn(rows, HEAD_DIM, **bf16), + kv_cache, + slots, + norm, + norm, + positions, + angles.cos().bfloat16(), + angles.sin().bfloat16(), + 1e-6, + FP4_QUANT_BLOCK_SIZE, + "ue8m0", + WEIGHTS_SCALE, + is_neox=True, + q_scale_out=q_scale, + kv_cache_scale=kv_scale, + ) + return q_fp4, q_scale, weights_out, kv_cache, kv_scale + + +def _written_cache_and_queries(batch, next_n, ctx_len, seed): + """A paged cache plus `batch * next_n` query rows, both from the writer. + + The query rows get real slots of their own, which is what a decode step + passes and what makes the writer compute Q at all. + """ + torch.manual_seed(seed) + rows = batch * next_n + table, num_blocks, token, _, slots = _paged_layout(batch, ctx_len) + *_, kv_cache, kv_scale = _fused_fp4(slots, token, num_blocks) + q_fp4, q_scale, weights_out, *_ = _fused_fp4( + torch.arange(rows, dtype=torch.int64, device="cuda"), + torch.full((rows,), ctx_len - 1, dtype=torch.int64, device="cuda"), + num_blocks, + weight_gain=0.1, + ) + return table, kv_cache, kv_scale, q_fp4, q_scale, weights_out, rows + + +def _oracle(q_fp4, q_scale, kv_cache, kv_scale, table, ctx_len, weights, rows_of): + """The scorer's math in fp32 over the cache as written: per-head ReLU(q.k), + weighted and summed. `rows_of` maps the per-sequence keys onto query rows.""" + batch = table.shape[0] + token = torch.arange(ctx_len, device=kv_cache.device) + phys = table[:, token // _BLOCK].long().unsqueeze(-1) + pos = (token % _BLOCK).expand(batch, ctx_len).unsqueeze(-1) + group = torch.arange(4, device=kv_cache.device) + packed = kv_cache[phys, 0, group, pos].reshape(batch, ctx_len, HEAD_DIM // 2) + keys = _dequant(packed, kv_scale[phys, 0, group, fp4_index_scale_rows(pos, _BLOCK)]) + # `[T, k_tiles, 4, 16, qs_pad]` -> the dense `[T, H, D // 32]` a reader sees. + dense = ( + q_scale[..., : HEADS // 16] + .permute(0, 4, 3, 1, 2) + .reshape(q_scale.shape[0], HEADS, HEAD_DIM // FP4_QUANT_BLOCK_SIZE) + ) + scores = torch.einsum( + "rhd,rtd->rht", _dequant(q_fp4, dense.contiguous()), rows_of(keys) + ) + return (torch.relu(scores) * weights.float().unsqueeze(-1)).sum(1) * WEIGHTS_SCALE + + +def _assert_agrees(got, want, visible, topk): + """Cosine over the visible window and the worst row's top-k overlap: the two + numbers that say whether the selection this feeds would differ.""" + mask = torch.arange(want.shape[1], device=want.device)[None, :] < visible[:, None] + a, b = got[mask].double(), want[mask].double() + cosine = (a @ b / (a.norm() * b.norm())).item() + lens = [min(topk, int(n)) for n in visible] + overlap = min( + len( + set(got[r, : int(visible[r])].topk(k).indices.tolist()) + & set(want[r, : int(visible[r])].topk(k).indices.tolist()) + ) + / k + for r, k in enumerate(lens) + if k + ) + assert cosine > 0.9999, cosine + assert overlap > 0.99, overlap + + +def test_decode_scores_the_cache_the_fused_writer_wrote(on_gfx950): + """The rectangular kernel at a speculation width: one row per (seq, step), + each seeing one token less than the step after it. `next_n=1` is the DCP + test's shape, so what this one holds down is the `next_n > 1` reshape.""" + from aiter.ops.flydsl import flydsl_pa_mqa_logits_fp4 + from aiter.ops.flydsl.kernels.mqa_logits.pa_mqa_logits_fp4 import ( + compute_varctx_schedule, + ) + + batch, next_n, ctx_len = 2, 4, 768 + table, kv_cache, kv_scale, q_fp4, q_scale, weights_out, rows = ( + _written_cache_and_queries(batch, next_n, ctx_len, seed=0) + ) + ctx_lens = torch.full((batch,), ctx_len, dtype=torch.int32, device="cuda") + _, cta_info, n_ctas = compute_varctx_schedule( + ctx_lens, FP4_MQA_BLOCK_K, None, ctx_len, next_n=next_n + ) + logits = torch.empty(rows, ctx_len, dtype=torch.float32, device="cuda") + flydsl_pa_mqa_logits_fp4( + q_fp4.reshape(batch, next_n, HEADS, HEAD_DIM // 2), + q_scale.reshape(batch, next_n, *q_scale.shape[1:]), + kv_cache, + kv_scale, + table, + weights_out, + ctx_lens, + ctx_len, + weight_scale=WEIGHTS_SCALE, + next_n=next_n, + block_k=FP4_MQA_BLOCK_K, + kv_block_size=_BLOCK, + out=logits, + cta_info=cta_info, + total_ctas=n_ctas, + ) + + want = _oracle( + q_fp4, + q_scale, + kv_cache, + kv_scale, + table, + ctx_len, + weights_out, + lambda keys: keys.repeat_interleave(next_n, dim=0), + ) + # Each of a request's next_n rows sees one token less than the one after it. + row = torch.arange(rows, device="cuda") + visible = ctx_lens.repeat_interleave(next_n) - (next_n - 1 - row % next_n) + _assert_agrees(logits, want, visible, topk=512) + + +def test_dcp_decode_scores_each_query_token_over_its_own_local_window(on_gfx950): + """The geometry `dcp_decode_candidate_exchange_fused` hands the scorer: one + row per query token over this rank's shard, so the windows are ragged and + the schedule is built at next_n=1 whatever the speculation width.""" + from aiter.ops.flydsl import flydsl_pa_mqa_logits_fp4 + + batch, next_n, width = 3, 4, 1024 + table, kv_cache, kv_scale, q_fp4, q_scale, weights_out, rows = ( + _written_cache_and_queries(batch, next_n, width, seed=2) + ) + + # Ragged on purpose: a draft position's extra token lands on ONE rank, so + # the local lengths of a request's next_n rows do not all advance together. + local_ctx = torch.tensor( + [width - (r % 7) * 37 for r in range(rows)], dtype=torch.int32, device="cuda" + ) + units = fp4_decode_parallel_units(batch, next_n) + cta_info = torch.zeros(units, 4, dtype=torch.int32, device="cuda") + fp4_decode_schedule(local_ctx, FP4_MQA_BLOCK_K, units, width, 1, cta_info) + + logits = torch.empty(rows, width, dtype=torch.float32, device="cuda") + flydsl_pa_mqa_logits_fp4( + q_fp4.reshape(rows, 1, HEADS, HEAD_DIM // 2), + q_scale.reshape(rows, 1, *q_scale.shape[1:]), + kv_cache, + kv_scale, + table.repeat_interleave(next_n, dim=0), + weights_out, + local_ctx, + width, + weight_scale=WEIGHTS_SCALE, + next_n=1, + block_k=FP4_MQA_BLOCK_K, + kv_block_size=_BLOCK, + out=logits, + cta_info=cta_info, + total_ctas=units, + ) + + want = _oracle( + q_fp4, + q_scale, + kv_cache, + kv_scale, + table, + width, + weights_out, + lambda keys: keys.repeat_interleave(next_n, dim=0), + ) + _assert_agrees(logits, want, local_ctx, topk=512) + + +def test_dcp_prefill_staging_keeps_every_key_with_its_own_exponent(monkeypatch): + """Staging moves two planes whose row axes disagree -- the packed one flat, + the e8m0 one transposed -- so it cannot address them with a single index. + + Every index stays in bounds either way, so the failure is silent: keys come + back wearing another row's exponent. Real exponents are nearly uniform + inside a block because `k_norm` precedes the quantizer, which is why an + end-to-end accuracy run can pass with this broken; the random planes here + remove that cover. + """ + dsv2 = _import_or_skip("atom.models.deepseek_v2") + + torch.manual_seed(0) + block, world, src_pages = _BLOCK, 2, 4 + local = 64 + total_kv = world * local + shape = {"dtype": torch.uint8, "device": "cpu"} + data_src = torch.randint(0, 256, (src_pages, 1, 4, block, 16), **shape) + scale_src = torch.randint(0, 256, (src_pages, 1, 4, block), **shape) + + # This rank's slots, spread over pages and rows so no source row equals the + # destination row it lands on; the gather then reorders them again. + slots = torch.randperm(src_pages * block)[:local].to(torch.int32) + gather_index = torch.randperm(total_kv).to(torch.int32) + monkeypatch.setattr( + dsv2, + "get_dcp_group", + lambda: SimpleNamespace( + all_gather=lambda t, dim: t.repeat(world, *([1] * (t.dim() - 1))) + ), + ) + + staged, staged_scale = dsv2._dcp_stage_indexer_fp4_prefill( + data_src, + scale_src, + SimpleNamespace( + dcp_indexer_fp4_local_slots=slots, + dcp_indexer_gather_index=gather_index, + ), + total_kv, + block, + ) + + src = slots[gather_index.long() % local].long() + got_rows = torch.arange(total_kv) + q_dst = fp4_index_scale_rows(got_rows % block, block) + q_src = fp4_index_scale_rows(src % block, block) + assert torch.equal( + staged[got_rows // block, 0, :, got_rows % block, :], + data_src[src // block, 0, :, src % block, :], + ) + assert torch.equal( + staged_scale[got_rows // block, 0, :, q_dst], + scale_src[src // block, 0, :, q_src], + ) + + +def test_staged_page_table_spans_a_whole_batch_not_one_sequence(): + """`pages` counts the summed co-scheduled prefill context, and prefix + caching lets that run past any one sequence's block allowance -- + `max_num_batched_tokens` bounds only the uncached tokens. Sized at that + allowance the table would turn a legal schedule into a mid-serving raise, + so it spans a full batch; the tail the scorer never reads stays zero rather + than aliasing a real page.""" + import numpy as np + + aiter_mla = _import_or_skip( + "atom.model_ops.attentions.aiter_mla", + reason="the MLA builder imports triton at module scope", + ) + + build = aiter_mla.AiterMLAMetadataBuilder._build_dcp_indexer_fp4_prefill_meta + block, bs, per_seq = 64, 2, 6 + builder = SimpleNamespace( + model_runner=SimpleNamespace(block_size=block), + device=torch.device("cpu"), + max_bs=4, + block_table_cols=per_seq, + ) + cols = builder.max_bs * per_seq + lpad = np.full(bs, block, dtype=np.int64) + cu_pad = np.concatenate([[0], np.cumsum(lpad)]).astype(np.int64) + var = {"block_tables": SimpleNamespace(np=np.zeros((bs, 8), dtype=np.int32))} + meta = SimpleNamespace() + + def staged_for(total_kv): + build(builder, meta, bs, lpad, cu_pad, total_kv, var) + return meta.dcp_indexer_fp4_block_tables + + for pages in (4, per_seq + 1, cols): + staged = staged_for(pages * block) + # Past `per_seq` is the case a one-sequence width used to raise on. + assert staged.shape == (bs, cols), pages + want = torch.arange(pages, dtype=torch.int32).expand(bs, pages) + assert torch.equal(staged[:, :pages], want), pages + assert not staged[:, pages:].any(), pages + + +@pytest.mark.parametrize("whole_batch", [True, False]) +def test_prefill_scores_the_same_cache_seq_locally(on_gfx950, whole_batch): + from atom.models.deepseek_v2 import _prefill_mqa_logits_fp4 + + torch.manual_seed(1) + batch, ctx_len = 2, 512 + table, num_blocks, token, seq, slots = _paged_layout(batch, ctx_len) + q_fp4, q_scale, weights_out, kv_cache, kv_scale = _fused_fp4( + slots, token, num_blocks, weight_gain=0.1 + ) + + # One row per query token, each seeing `[0, its own position]` of its own + # sequence -- the seq-local windows the metadata builder publishes. + rows = batch * ctx_len + local_ends = (token + 1).to(torch.int32) + row_to_batch = seq.to(torch.int32) + cta_info, n_ctas, local_starts = fp4_prefill_schedule( + row_to_batch, local_ends, FP4_MQA_BLOCK_K, rows, ctx_len + ) + # With `whole_batch=False` the model rebuilds the schedule per chunk instead + # of reusing the one the builder left on the metadata. + logits = _prefill_mqa_logits_fp4( + SimpleNamespace( + batch_id_per_q_token=row_to_batch, + block_tables=table, + indexer_fp4_local_starts=local_starts, + indexer_fp4_local_ends=local_ends, + indexer_fp4_max_seq_len=ctx_len, + indexer_fp4_cta_info=cta_info, + indexer_fp4_n_ctas=n_ctas, + ), + slice(0, rows), + whole_batch, + q_fp4, + q_scale, + weights_out, + kv_cache, + kv_scale, + WEIGHTS_SCALE, + _BLOCK, + table, + ) + + want = _oracle( + q_fp4, + q_scale, + kv_cache, + kv_scale, + table, + ctx_len, + weights_out, + lambda keys: keys[row_to_batch.long()], + ) + _assert_agrees(logits, want, local_ends, topk=256)