Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
18 commits
Select commit Hold shift + click to select a range
93815ce
feat(deepseek_v4): Add fp4 indexer
junhaha666 Jun 30, 2026
a57cfbb
feat(deepseek_v4): gate FP4 indexer behind --enable-deepseek-v4-fp4-i…
junhaha666 Jun 30, 2026
40502bc
perf(deepseek_v4): write FP4 decode cta_info via fused single-kernel …
junhaha666 Jul 1, 2026
367267e
update rope_rotate_activation args
junhaha666 Jul 13, 2026
66d50c9
Merge remote-tracking branch 'origin/main' into jun/fp4_indexer
junhaha666 Jul 13, 2026
e23f31c
Fix FP4 indexer forward_pre to return (q, weights) after merge
junhaha666 Jul 13, 2026
881ad85
perf(dsv4): right-size FP4 indexer prefill logits + share chunk helper
junhaha666 Jul 13, 2026
533e8af
fix(dsv4): HCA paged-gather must account for k2_hca entries per block
junhaha666 Jul 15, 2026
1caedba
Merge remote-tracking branch 'origin/main' into jun/fp4_indexer
junhaha666 Jul 21, 2026
4d9c0ef
fix merge
junhaha666 Jul 21, 2026
29ebfaa
add dsv4 fp4 indexer support DSpark
junhaha666 Jul 25, 2026
f3b1fb4
Fix FP4 indexer OOB under --cudagraph-mode FULL + DSpark ragged
junhaha666 Jul 27, 2026
032c0eb
Merge remote-tracking branch 'origin/main' into jun/fp4_indexer_nextn
junhaha666 Jul 27, 2026
588c96c
Drop the FP4 rectangular-decode ragged-window refresh, unreachable si…
junhaha666 Jul 27, 2026
2ce176c
Apply black and the ruff fixes CI reported
junhaha666 Jul 27, 2026
7899087
Fix Ruff warnings from CI
junhaha666 Jul 28, 2026
5afee50
Remove unused Optional import
junhaha666 Jul 28, 2026
99e9947
refactor(dsv4): share the FP4-indexer predicate between builder and I…
junhaha666 Jul 28, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 8 additions & 5 deletions atom/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -1463,18 +1463,21 @@ def __post_init__(self):
)

# DeepSeek V4: paper §3.6.1 mandates classical KV cache block_size =
# lcm(m, m'). For V4-Pro / V4-Flash this is lcm(4, 128) = 128 original
# tokens. ATOM's BlockManager + slot_mapping math assume one global
# block_size, so we override `kv_cache_block_size` here when V4 is
# detected; the V4 attention builder enforces the same value.
# a multiple of lcm(m, m'). For V4-Pro / V4-Flash lcm(4, 128) = 128;
# we use 2*lcm = 256 so each block holds k1=256/4=64 CSA entries — the
# FP4 paged-MQA-logits indexer kernels require kv_block_size=64 (so
# NTPW=4 N-tiles share one physical block, N_PHYS=1). ATOM's
# BlockManager + slot_mapping math assume one global block_size, so we
# override `kv_cache_block_size` here when V4 is detected; the V4
# attention builder enforces the same value.
#
# NOTE: cannot use `hf_config.model_type` for detection — `_CONFIG_REGISTRY`
# maps "deepseek_v4" → "deepseek_v3" so model_type reads as "deepseek_v3".
# Use the preserved `architectures` field (re-injected by get_hf_config,
# line 567) which keeps the original "DeepseekV4ForCausalLM[NextN]" name.
arches = getattr(self.hf_config, "architectures", None) or []
if any("DeepseekV4" in str(a) for a in arches):
v4_block_size = 128
v4_block_size = 256
if self.kv_cache_block_size != v4_block_size:
self.kv_cache_block_size = v4_block_size

Expand Down
6 changes: 4 additions & 2 deletions atom/model_engine/arg_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,10 +152,12 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser:
parser.add_argument(
"--index-cache-dtype",
"--index_cache_dtype",
choices=["bf16", "fp8"],
choices=["bf16", "fp8", "fp4"],
type=str,
default=None,
help="Index cache type. Defaults to --kv_cache_dtype.",
help="Index cache type. Defaults to --kv_cache_dtype. 'fp4' selects "
"the DeepSeek-V4 FP4 CSA indexer (gfx950 only; falls back to fp8 "
"elsewhere).",
)
parser.add_argument(
"--block-size", type=int, default=16, help="KV cache block size."
Expand Down
427 changes: 374 additions & 53 deletions atom/model_ops/attentions/deepseek_v4_attn.py

Large diffs are not rendered by default.

16 changes: 10 additions & 6 deletions atom/model_ops/module_dispatch_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,9 @@ def _maybe_dual_stream_forward_fake(
# ---------------------------------------------------------------------------
#
# Caller contract (the Indexer module looked up by `layer_name`):
# - `indexer_score_topk(q_fp8, weights, topk) -> Tensor` — real impl,
# - `indexer_score_topk(q_quant, weights, q_scale, topk) -> Tensor` — real impl,
# (q_quant is FP8, or packed FP4 uint8 when the FP4 indexer is on; q_scale is
# the paired e8m0 Q scale on the FP4 path, None for FP8)
# must return `[total_tokens, topk] int32` indices
#
# `topk` is on the op signature (not derived from the module) so the fake
Expand All @@ -115,27 +117,29 @@ def _maybe_dual_stream_forward_fake(


def indexer_score_topk(
q_fp8: torch.Tensor,
q_quant: torch.Tensor,
weights: torch.Tensor,
q_scale: torch.Tensor | None,
layer_name: str,
topk: int,
) -> torch.Tensor:
indexer = get_current_atom_config().compilation_config.static_forward_context[
layer_name
]
return indexer.indexer_score_topk(q_fp8, weights, topk)
return indexer.indexer_score_topk(q_quant, weights, q_scale, topk)


def _indexer_score_topk_fake(
q_fp8: torch.Tensor,
q_quant: torch.Tensor,
weights: torch.Tensor,
q_scale: torch.Tensor | None,
layer_name: str,
topk: int,
) -> torch.Tensor:
return torch.empty(
(q_fp8.shape[0], topk),
(q_quant.shape[0], topk),
dtype=torch.int32,
device=q_fp8.device,
device=q_quant.device,
)


Expand Down
105 changes: 84 additions & 21 deletions atom/model_ops/v4_kernels/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,11 @@
derived from device data via `.item()`.
"""

import logging
from typing import Any

from aiter.jit.utils.chip_info import get_gfx

from atom.model_ops.v4_kernels.compress_plan import (
CompressPlan,
make_compress_plans,
Expand All @@ -21,58 +26,116 @@
fused_compress_attn,
fused_compress_attn_reference,
)
from atom.model_ops.v4_kernels.indexer_weights import scale_indexer_weights
from atom.model_ops.v4_kernels.indexer_weights import (
scale_indexer_weights,
)
from atom.model_ops.v4_kernels.inverse_rope import inverse_rope_inplace
from atom.model_ops.v4_kernels.paged_decode import (
sparse_attn_v4_paged_decode,
sparse_attn_v4_paged_decode_reference,
)
from atom.model_ops.v4_kernels.paged_prefill import (
sparse_attn_v4_paged_prefill,
sparse_attn_v4_paged_prefill_reference,
)
from atom.model_ops.v4_kernels.inverse_rope import inverse_rope_inplace
from atom.model_ops.v4_kernels.paged_decode_indices import (
hca_compress_paged_offsets,
write_v4_paged_decode_indices,
write_v4_paged_decode_indices_reference,
)
from atom.model_ops.v4_kernels.paged_prefill import (
sparse_attn_v4_paged_prefill,
sparse_attn_v4_paged_prefill_reference,
)
from atom.model_ops.v4_kernels.paged_prefill_indices import (
write_v4_paged_prefill_indices,
write_v4_paged_prefill_indices_reference,
)
from atom.model_ops.v4_kernels.qk_norm_rope_maybe_quant import (
QKNormRopeOut,
qk_norm_rope_maybe_quant,
qk_norm_rope_maybe_quant_reference,
qk_norm_rope_maybe_quant_fp8_2buff,
qk_norm_rope_maybe_quant_reference,
)
from atom.model_ops.v4_kernels.state_writes import (
update_compressor_states,
swa_write,
swa_write_2buff_prepacked,
update_compressor_states,
)

__all__ = [
"update_compressor_states",
"swa_write",
"swa_write_2buff_prepacked",
"FP4_MQA_BLOCK_K",
"FP4_MQA_PARALLEL_UNIT_NUM",
"CompressPlan",
"QKNormRopeOut",
"csa_translate_pack",
"csa_translate_pack_reference",
"fp4_indexer_enabled",
"fused_compress_attn",
"fused_compress_attn_reference",
"hca_compress_paged_offsets",
"inverse_rope_inplace",
"make_compress_plans",
"qk_norm_rope_maybe_quant",
"qk_norm_rope_maybe_quant_fp8_2buff",
"qk_norm_rope_maybe_quant_reference",
"scale_indexer_weights",
"sparse_attn_v4_paged_decode",
"sparse_attn_v4_paged_decode_reference",
"sparse_attn_v4_paged_prefill",
"sparse_attn_v4_paged_prefill_reference",
"csa_translate_pack",
"csa_translate_pack_reference",
"CompressPlan",
"make_compress_plans",
"inverse_rope_inplace",
"scale_indexer_weights",
"swa_write",
"swa_write_2buff_prepacked",
"update_compressor_states",
"write_v4_paged_decode_indices",
"write_v4_paged_decode_indices_reference",
"write_v4_paged_prefill_indices",
"write_v4_paged_prefill_indices_reference",
"QKNormRopeOut",
"qk_norm_rope_maybe_quant",
"qk_norm_rope_maybe_quant_reference",
"qk_norm_rope_maybe_quant_fp8_2buff",
]

logger = logging.getLogger("atom")

# FP4 indexer persistent-grid schedule params, shared by the decode
# (`pa_mqa_logits_fp4`) and prefill (`pa_mqa_logits_fp4_prefill`) kernels.
# The attention metadata builder precomputes each path's cta_info with these
# and the scorer passes the matching block_k, so layout and grid agree. They
# live here (rather than in either caller) because both the builder and the
# model-side scorer must use the SAME values. Mirrors the kernel defaults.
FP4_MQA_PARALLEL_UNIT_NUM = 512
FP4_MQA_BLOCK_K = 256


def fp4_indexer_enabled(index_cache_dtype: Any, *, warn: bool = False) -> bool:
"""Is the FP4 CSA indexer active? Single source of truth for the predicate.

Two call sites must reach the SAME verdict, and neither can be dropped:

* `DeepseekV4AttentionMetadataBuilder.__init__` — authoritative. Picks the
cache-pool layout and re-asserts the flag onto each Indexer in
`build_kv_cache_tensor`.
* `Indexer.__init__` — must already be correct BEFORE that re-assert.
`model_runner._maybe_warmup()` traces the graphed `_attn_pre`/`forward_pre`
piece before `allocate_kv_cache()` -> `build_kv_cache_tensor()` runs, so an
Indexer that defaulted to False would bake the FP8 branch (`q_scale=None`)
into the graph while the eager `indexer_score_topk` later took the FP4
branch. It is also the ONLY setter under the vLLM / SGLang plugins, which
never call `build_kv_cache_tensor`.

Keeping the predicate in one place is what stops the two from drifting —
a divergence is silent at startup and only surfaces as a graph/eager dtype
mismatch. Lives here rather than in either caller for the same reason as
`FP4_MQA_*` above.

The FP4 mqa-logits / scatter kernels are gfx950 (MI355X / CDNA4) only; on any
other arch fall back to the FP8 indexer instead of failing. Pass `warn=True`
from the builder only — it runs once, while `Indexer.__init__` runs per CSA
layer and would repeat the message.
"""
if index_cache_dtype != "fp4":
return False
gfx = get_gfx()
if gfx != "gfx950":
if warn:
logger.warning(
"--index_cache_dtype fp4 requires a gfx950 (MI355X / CDNA4) GPU; "
"current arch is %r. Falling back to the FP8 indexer.",
gfx,
)
return False
return True
2 changes: 1 addition & 1 deletion atom/model_ops/v4_kernels/csa_translate_pack.py
Original file line number Diff line number Diff line change
Expand Up @@ -200,7 +200,7 @@ def csa_translate_pack(
`[indptr[t], indptr[t]+valid_k[t])`.
swa_pages: SWA region size — `num_slots * window_size`,
fixed at CG capture time. Keyword-only.
csa_block_capacity: `block_size // ratio = 128 // 4 = 32`
csa_block_capacity: `block_size // ratio = 256 // 4 = 64`
(constexpr; triton can strength-reduce
// and %). Keyword-only.
window_size: SWA window. When > 0 the kernel computes
Expand Down
51 changes: 36 additions & 15 deletions atom/model_ops/v4_kernels/fused_compress.py
Original file line number Diff line number Diff line change
Expand Up @@ -389,10 +389,10 @@ def fused_compress_attn(
cos_cache: torch.Tensor, # [max_seq, ..., rope_head_dim/2] bf16/fp16
sin_cache: torch.Tensor, # same shape
# KV cache scatter
kv_cache: Optional[
torch.Tensor
], # bf16: [NB, k_per_block, head_dim] / fp8: same shape, fp8
block_tables: Optional[torch.Tensor], # [bs, max_blocks_per_seq] int32
kv_cache: (
torch.Tensor | None
), # bf16: [NB, k_per_block, head_dim] / fp8: same shape, fp8
block_tables: torch.Tensor | None, # [bs, max_blocks_per_seq] int32
k_per_block: int,
# Geometry
overlap: bool,
Expand All @@ -401,20 +401,19 @@ def fused_compress_attn(
rope_head_dim: int,
# FP8 quant fusion (Indexer-inner Compressor path)
quant: bool = False,
cache_scale: Optional[
torch.Tensor
] = None, # fp32 [NB, k_per_block]; required when quant=True
cache_scale: (
torch.Tensor | None
) = None, # fp32 [NB, k_per_block] (FP8) / uint8 (FP4); required when quant
use_ue8m0: bool = True, # round scale to power-of-2 (UE8M0); only when quant=True
preshuffle: bool = True, # MFMA 16x16 preshuffled FP8 layout; only when quant=True
fp8_max: Optional[
float
] = None, # E4M3 max; required only on the quant (indexer_fp8) path
fp8_max: float | None = None, # E4M3 max; required for FP8 quant
quant_mode: str | None = None, # "none"|"fp8"|"fp4"; default from `quant`
# V4-Main native fp8 2buff path (CSA/HCA Main under --kv_cache_dtype fp8).
# Distinct from `quant` (Indexer-inner per-row preshuffle): writes per-64-tile
# e8m0 nope-fp8 + inline dup-scale into `kv_cache` (fp8 [NB,k,512]) and bf16
# rope into `kv_cache_rope` (bf16 [NB,k,64]) via the flydsl group_fp8 scatter.
main_2buff_fp8: bool = False,
kv_cache_rope: Optional[torch.Tensor] = None, # bf16 [NB,k_per_block,64]
kv_cache_rope: torch.Tensor | None = None, # bf16 [NB,k_per_block,64]
prefix: str = "",
) -> None:
"""Batched fused per-source-position pool + RMSNorm + RoPE + cache scatter,
Expand Down Expand Up @@ -445,6 +444,18 @@ def fused_compress_attn(
if plan_capacity == 0:
return # nothing to do — no plan rows ever populated.

# Resolve quant mode (single source of truth). `quant_mode` is the unified
# selector; the legacy `quant` bool is only a fallback for callers that don't
# pass quant_mode ("fp8" if quant else "none").
_mode = quant_mode if quant_mode is not None else ("fp8" if quant else "none")
_fp4 = _mode == "fp4"
# `quant` = the Indexer-inner per-row fp8/fp4 scatter (needs cache_scale /
# fp8_max, takes the quant validation + Triton fallback path). Derive it from
# the mode so callers only pass quant_mode. CSA/HCA Main group_fp8 is NOT
# `quant` — its 2buff scatter is driven by `main_2buff_fp8` — and `none`/bf16
# is plain. (Back-compat: legacy quant=True → quant_mode None → _mode "fp8".)
quant = _mode in ("fp8", "per_row_fp8", "fp4")

# ------------------------------------------------------------------
# flydsl dispatch. Pure-GPU time on V4-Pro beats Triton 0.9x→2.9x
# across the relevant N_compress range; the small-N gap is bridged
Expand Down Expand Up @@ -526,11 +537,21 @@ def fused_compress_attn(
cache_scale=cache_scale,
use_ue8m0=use_ue8m0,
preshuffle=preshuffle and not main_2buff_fp8,
quant_mode="group_fp8" if main_2buff_fp8 else "per_row_fp8",
quant_mode="group_fp8" if main_2buff_fp8 else _mode,
k_rope_cache=kv_cache_rope if main_2buff_fp8 else None,
)
return

if _fp4:
# FP4 indexer scatter only exists in the flydsl kernel — the Triton
# fallback below has no FP4 path. Reaching here means flydsl is
# unavailable or the shape is unsupported.
raise RuntimeError(
"quant_mode='fp4' requires the flydsl fused_compress_attn kernel "
f"(available={flydsl_fused_compress_attn is not None}, "
f"shape_ok={_flydsl_shape_ok}, mode={_flydsl_mode})."
)

# Validate shapes
dim_full = (2 if overlap else 1) * head_dim
K_pool = (2 if overlap else 1) * ratio # pool window size (algorithm-defined)
Expand Down Expand Up @@ -672,15 +693,15 @@ def fused_compress_attn_reference(
rms_eps: float,
cos_cache: torch.Tensor,
sin_cache: torch.Tensor,
kv_cache: Optional[torch.Tensor],
block_tables: Optional[torch.Tensor],
kv_cache: torch.Tensor | None,
block_tables: torch.Tensor | None,
k_per_block: int,
overlap: bool,
ratio: int,
head_dim: int,
rope_head_dim: int,
out_dtype: torch.dtype = torch.bfloat16,
) -> Optional[torch.Tensor]:
) -> torch.Tensor | None:
"""Pure-PyTorch reference equivalent of `fused_compress_attn` (plan path).

Returns `[num_compress, head_dim]` BF16 in plan order. None if num_compress=0.
Expand Down
24 changes: 24 additions & 0 deletions atom/model_ops/v4_kernels/paged_decode_indices.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,13 +31,37 @@
per layer).
"""

import numpy as np
import torch
import triton
import triton.language as tl

from atom.utils.decorators import mark_trace


def hca_compress_paged_offsets(
entry_idx, bid_per_entry, block_tables_np, swa_pages, k2_hca
):
"""HCA compress entry -> unified paged row (numpy, decode index build).

The compressor packs ``k2_hca = block_size // hca_ratio`` HCA compress entries
per physical block (cache view ``[num_blocks, k2_hca, head_dim]``): entry ``e``
lives in physical block ``block_tables[bid, e // k2_hca]`` at slot
``e % k2_hca``, so its unified row is ``swa_pages + phys * k2_hca + slot``.
(Reduces to ``swa_pages + block_tables[bid, e]`` at ``k2_hca == 1``.)

``entry_idx`` / ``bid_per_entry`` are int arrays of equal length; returns an
int32 array of the same length. Shared by ``_attach_v4_paged_decode_meta`` and
covered by ``tests/test_decode_indices_paged.py`` so the packing stays correct
for V4 ``block_size=256`` (``k2_hca=2``).
"""
blk = entry_idx // k2_hca
slot = entry_idx % k2_hca
return (swa_pages + block_tables_np[bid_per_entry, blk] * k2_hca + slot).astype(
np.int32
)


@triton.jit
def _v4_paged_decode_indices_kernel(
block_tables_ptr, # [bs, max_blocks] int32 — logical→physical block
Expand Down
Loading
Loading