[AMD] Enable FP4 indexer for Deepseek V4 (#37353)
Co-authored-by: 1am9trash <1am9trash@gmail.com> Co-authored-by: AMD-yanfeiwang <256076023+AMD-yanfeiwang@users.noreply.github.com> Co-authored-by: Thomas Wang <thomawan@amd.com>
This commit is contained in:
co-authored by
1am9trash
AMD-yanfeiwang
Thomas Wang
parent
f6aed6ec53
commit
f8cbf000f4
@@ -1,6 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Literal, NamedTuple, Optional, Union
|
||||
from typing import TYPE_CHECKING, Literal, NamedTuple, Optional, Union, cast
|
||||
|
||||
import torch
|
||||
|
||||
@@ -430,9 +430,33 @@ def compress_norm_rope_store(
|
||||
page_size: int,
|
||||
use_fp4: bool = False,
|
||||
bf16_store: bool = False,
|
||||
# HIP FP4 uses split scale storage and precomputed BF16 RoPE tables.
|
||||
kvcache_scale: Optional[torch.Tensor] = None,
|
||||
rope_cache: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
fp4_k_write_metadata=None,
|
||||
) -> None:
|
||||
if use_fp4:
|
||||
assert kv.shape[-1] == 128
|
||||
if is_hip() and use_fp4:
|
||||
from sglang.kernels.ops.attention.dsv4.fp4_indexer_hip import (
|
||||
aiter_k_indexer_fp4_cache_write,
|
||||
)
|
||||
|
||||
cos, sin = cast(tuple[torch.Tensor, torch.Tensor], rope_cache)
|
||||
aiter_k_indexer_fp4_cache_write(
|
||||
k=kv,
|
||||
norm_weight=norm_weight,
|
||||
norm_epsilon=norm_eps,
|
||||
cos=cos,
|
||||
sin=sin,
|
||||
plan=plan,
|
||||
out_loc=out_loc,
|
||||
k_payload=kvcache,
|
||||
k_scale=cast(torch.Tensor, kvcache_scale),
|
||||
write_metadata=fp4_k_write_metadata,
|
||||
)
|
||||
return
|
||||
|
||||
freq_cis = torch.view_as_real(freq_cis).flatten(-2)
|
||||
if _is_xpu:
|
||||
compress_norm_rope_store_xpu(
|
||||
|
||||
@@ -0,0 +1,400 @@
|
||||
"""AITER adapters for the DeepSeek-V4 FP4 indexer on HIP."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, NamedTuple, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.kernels.ops.attention.dsv4.compress import (
|
||||
CompressorDecodePlan,
|
||||
CompressorPrefillPlan,
|
||||
)
|
||||
|
||||
|
||||
_HEADS = 64
|
||||
_HEAD_DIM = 128
|
||||
_ROPE_DIM = 64
|
||||
_GROUP_SIZE = 32
|
||||
_KV_BLOCK_SIZE = 64
|
||||
_Q_SCALE_SHAPE = (1, 4, 16, 4)
|
||||
# gfx950 has 256 CUs; target four persistent CTAs per CU.
|
||||
_DECODE_BASE_CTA_TARGET = 1024
|
||||
# Preserve per-query parallelism when the batch itself exceeds one CTA per CU.
|
||||
_DECODE_CTAS_PER_QUERY = 4
|
||||
_PREFILL_BASE_CTA_TARGET = 1024
|
||||
# AITER varctx cta_info row: [batch_packed, chunk_start, chunk_count, ctx_len].
|
||||
_DECODE_CTA_INFO_WIDTH = 4
|
||||
|
||||
|
||||
class FP4DecodeWorkspace(NamedTuple):
|
||||
guarded_page_table: torch.Tensor
|
||||
c4_seq_lens: torch.Tensor
|
||||
cta_info: torch.Tensor
|
||||
cta_count: int
|
||||
max_seq_len: int
|
||||
# Held only so AITER's schedule scratch never returns to the graph memory
|
||||
# pool: the captured builder writes it again on every replay.
|
||||
schedule_scratch: torch.Tensor
|
||||
|
||||
|
||||
class FP4PrefillWorkspace(NamedTuple):
|
||||
guarded_page_table: torch.Tensor
|
||||
row_to_batch: torch.Tensor
|
||||
local_starts: torch.Tensor
|
||||
cta_info: torch.Tensor
|
||||
cta_count: int
|
||||
max_seq_len: int
|
||||
|
||||
|
||||
class FP4KWriteMetadata(NamedTuple):
|
||||
positions: torch.Tensor
|
||||
slots: torch.Tensor
|
||||
|
||||
|
||||
def aiter_q_indexer_fp4(
|
||||
q: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Apply RoPE and Hadamard rotation, then quantize indexer Q to FP4."""
|
||||
import aiter
|
||||
|
||||
num_tokens = q.shape[0]
|
||||
# AITER asserts int64 positions; the caller normally widens once per forward.
|
||||
if (
|
||||
positions.dtype is not torch.int64
|
||||
or positions.device != q.device
|
||||
or not positions.is_contiguous()
|
||||
):
|
||||
positions = positions.to(device=q.device, dtype=torch.int64).contiguous()
|
||||
q_fp4 = torch.empty(
|
||||
(num_tokens, _HEADS, _HEAD_DIM // 2),
|
||||
dtype=aiter.dtypes.fp4x2,
|
||||
device=q.device,
|
||||
)
|
||||
q_scale = torch.empty(
|
||||
(num_tokens, *_Q_SCALE_SHAPE), dtype=torch.uint8, device=q.device
|
||||
)
|
||||
aiter.rope_rotate_activation(
|
||||
q_fp4,
|
||||
q,
|
||||
cos,
|
||||
sin,
|
||||
positions,
|
||||
rope_dim=_ROPE_DIM,
|
||||
out_scale=q_scale,
|
||||
group_size=_GROUP_SIZE,
|
||||
shuffle_scale=True,
|
||||
do_rotate_act=True,
|
||||
)
|
||||
return q_fp4, q_scale
|
||||
|
||||
|
||||
def _as_int32_1d(t: torch.Tensor) -> torch.Tensor:
|
||||
"""Normalize a length vector without dispatching when it already matches.
|
||||
|
||||
Called once per C4 layer, so the no-op fast path matters: the metadata
|
||||
builder already hands us a 1-D contiguous int32 tensor.
|
||||
"""
|
||||
if t.dim() == 1 and t.dtype is torch.int32 and t.is_contiguous():
|
||||
return t
|
||||
return t.reshape(-1).to(torch.int32).contiguous()
|
||||
|
||||
|
||||
def _decode_cta_count(num_queries: int, max_seq_len: int) -> int:
|
||||
"""Choose a bounded persistent grid without exceeding available KV chunks."""
|
||||
chunks_per_seq = max(1, (max_seq_len + 255) // 256)
|
||||
available_ctas = num_queries * chunks_per_seq
|
||||
target_ctas = max(_DECODE_BASE_CTA_TARGET, num_queries * _DECODE_CTAS_PER_QUERY)
|
||||
return min(available_ctas, target_ctas)
|
||||
|
||||
|
||||
def _guard_page_table(page_table: torch.Tensor, out: Optional[torch.Tensor] = None):
|
||||
"""Pad page tables for 256-token scheduling and one-chunk lookahead."""
|
||||
page_table = page_table.to(dtype=torch.int32).contiguous()
|
||||
rows, logical_width = page_table.shape
|
||||
padded_width = max(4, (logical_width + 3) // 4 * 4)
|
||||
if out is None:
|
||||
out = page_table.new_zeros((rows, padded_width + 4))
|
||||
else:
|
||||
assert out.shape == (rows, padded_width + 4), f"{out.shape=} {rows=}"
|
||||
out[:, :logical_width].copy_(page_table)
|
||||
return out, padded_width * _KV_BLOCK_SIZE
|
||||
|
||||
|
||||
def prepare_fp4_decode_workspace(
|
||||
page_table: torch.Tensor,
|
||||
c4_seq_lens: torch.Tensor,
|
||||
) -> FP4DecodeWorkspace:
|
||||
"""Build the decode page-table, schedule, and logits buffers.
|
||||
|
||||
Safe to run under CUDA-graph capture: every tensor the captured schedule
|
||||
kernel touches is reachable from the returned workspace, so none of it can
|
||||
be handed out again by a later capture sharing the graph memory pool.
|
||||
"""
|
||||
from aiter.ops.flydsl.kernels.mqa_logits.pa_mqa_logits_fp4 import (
|
||||
compute_varctx_schedule,
|
||||
)
|
||||
|
||||
guarded, max_seq_len = _guard_page_table(page_table)
|
||||
c4_seq_lens = _as_int32_1d(c4_seq_lens)
|
||||
num_queries = guarded.shape[0]
|
||||
cta_count = _decode_cta_count(num_queries, max_seq_len)
|
||||
cta_info = torch.empty(
|
||||
(cta_count, _DECODE_CTA_INFO_WIDTH),
|
||||
dtype=torch.int32,
|
||||
device=guarded.device,
|
||||
)
|
||||
schedule_scratch, _, _ = compute_varctx_schedule(
|
||||
c4_seq_lens,
|
||||
block_k=256,
|
||||
parallel_unit_num=cta_count,
|
||||
max_seq_len=max_seq_len,
|
||||
next_n=1,
|
||||
cta_info_out=cta_info,
|
||||
)
|
||||
return FP4DecodeWorkspace(
|
||||
guarded, c4_seq_lens, cta_info, cta_count, max_seq_len, schedule_scratch
|
||||
)
|
||||
|
||||
|
||||
def prepare_fp4_prefill_workspace(
|
||||
page_table: torch.Tensor,
|
||||
c4_seq_lens: torch.Tensor,
|
||||
workspace: Optional[FP4PrefillWorkspace] = None,
|
||||
) -> FP4PrefillWorkspace:
|
||||
"""Build or refresh the prefill page-table, schedule, and logits buffers.
|
||||
|
||||
Must run OUTSIDE CUDA-graph capture. AITER's prefill scheduler frees its own
|
||||
scratch when it returns, and its schedule kernel reads that scratch, so a
|
||||
captured build would replay against recycled graph-pool memory. Callers
|
||||
instead refresh this workspace per step and let the graph read only the
|
||||
pinned ``cta_info`` / ``logits`` / page-table buffers.
|
||||
"""
|
||||
from aiter.ops.flydsl.kernels.mqa_logits.pa_mqa_logits_fp4_prefill import (
|
||||
CTA_INFO_WIDTH,
|
||||
compute_prefill_schedule,
|
||||
)
|
||||
|
||||
c4_seq_lens = _as_int32_1d(c4_seq_lens)
|
||||
if workspace is None:
|
||||
guarded, max_seq_len = _guard_page_table(page_table)
|
||||
num_queries = guarded.shape[0]
|
||||
cta_count = max(_PREFILL_BASE_CTA_TARGET, num_queries)
|
||||
workspace = FP4PrefillWorkspace(
|
||||
guarded_page_table=guarded,
|
||||
row_to_batch=torch.arange(
|
||||
num_queries, device=guarded.device, dtype=torch.int32
|
||||
),
|
||||
local_starts=torch.zeros(
|
||||
num_queries, device=guarded.device, dtype=torch.int32
|
||||
),
|
||||
cta_info=torch.empty(
|
||||
(cta_count, CTA_INFO_WIDTH), dtype=torch.int32, device=guarded.device
|
||||
),
|
||||
cta_count=cta_count,
|
||||
max_seq_len=max_seq_len,
|
||||
)
|
||||
else:
|
||||
_guard_page_table(page_table, out=workspace.guarded_page_table)
|
||||
|
||||
assert c4_seq_lens.shape[0] == workspace.row_to_batch.shape[0], (
|
||||
f"c4_seq_lens rows {c4_seq_lens.shape[0]} do not match the workspace's "
|
||||
f"{workspace.row_to_batch.shape[0]}; the schedule kernel indexes both by row"
|
||||
)
|
||||
compute_prefill_schedule(
|
||||
workspace.row_to_batch,
|
||||
workspace.local_starts,
|
||||
c4_seq_lens,
|
||||
block_k=256,
|
||||
parallel_unit_num=workspace.cta_count,
|
||||
max_seq_len=workspace.max_seq_len,
|
||||
cta_info_out=workspace.cta_info,
|
||||
)
|
||||
return workspace
|
||||
|
||||
|
||||
def aiter_fp4_paged_mqa_logits(
|
||||
*,
|
||||
q_fp4: torch.Tensor,
|
||||
q_scale: torch.Tensor,
|
||||
k_payload: torch.Tensor,
|
||||
k_scale: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
page_table: torch.Tensor,
|
||||
c4_seq_lens: torch.Tensor,
|
||||
weight_scale: float,
|
||||
is_decode: bool,
|
||||
decode_workspace: Optional[FP4DecodeWorkspace] = None,
|
||||
prefill_workspace: Optional[FP4PrefillWorkspace] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Compute FP4 Q/K indexer logits with the decode or prefill FlyDSL kernel."""
|
||||
from aiter.ops.flydsl import (
|
||||
flydsl_pa_mqa_logits_fp4,
|
||||
flydsl_pa_mqa_logits_fp4_prefill,
|
||||
)
|
||||
|
||||
num_tokens = q_fp4.shape[0]
|
||||
c4_seq_lens = _as_int32_1d(c4_seq_lens)
|
||||
workspace = decode_workspace if is_decode else prefill_workspace
|
||||
# A workspace is bound to one row count. DP padding or truncated activations
|
||||
# can leave it stale, in which case fall back to building the schedule here.
|
||||
if workspace is not None and workspace.guarded_page_table.shape[0] != num_tokens:
|
||||
workspace = None
|
||||
if workspace is not None:
|
||||
page_table = workspace.guarded_page_table
|
||||
max_seq_len = workspace.max_seq_len
|
||||
else:
|
||||
page_table, max_seq_len = _guard_page_table(page_table)
|
||||
q_payload = q_fp4.view(torch.uint8)
|
||||
k_payload = k_payload.view(torch.uint8)
|
||||
# Scored write-once and freed with this call. Recycling it through the
|
||||
# allocator costs nothing because a pinned cta_info makes the kernel skip
|
||||
# its -inf pre-fill and the length-aware top-k reads only [0, c4_seq_len).
|
||||
logits = torch.empty(
|
||||
(num_tokens, max_seq_len), dtype=torch.float32, device=q_fp4.device
|
||||
)
|
||||
common = {
|
||||
"weight_scale": weight_scale,
|
||||
"block_k": 256,
|
||||
"kv_block_size": _KV_BLOCK_SIZE,
|
||||
"num_warps": 4,
|
||||
"out": logits,
|
||||
}
|
||||
|
||||
if is_decode:
|
||||
pinned = (
|
||||
{}
|
||||
if workspace is None
|
||||
else {
|
||||
"cta_info": workspace.cta_info,
|
||||
"total_ctas": workspace.cta_count,
|
||||
}
|
||||
)
|
||||
logits = flydsl_pa_mqa_logits_fp4(
|
||||
q_payload.reshape(num_tokens, 1, _HEADS, _HEAD_DIM // 2),
|
||||
q_scale.reshape(num_tokens, 1, *_Q_SCALE_SHAPE),
|
||||
k_payload,
|
||||
k_scale,
|
||||
page_table,
|
||||
weights,
|
||||
c4_seq_lens,
|
||||
max_seq_len,
|
||||
next_n=1,
|
||||
parallel_unit_num=None,
|
||||
**pinned,
|
||||
**common,
|
||||
)
|
||||
else:
|
||||
if workspace is None:
|
||||
pinned = {}
|
||||
row_to_batch = torch.arange(
|
||||
num_tokens, device=q_fp4.device, dtype=torch.int32
|
||||
)
|
||||
local_starts = torch.zeros(
|
||||
num_tokens, device=q_fp4.device, dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
pinned = {
|
||||
"cta_info": workspace.cta_info,
|
||||
"n_ctas": workspace.cta_count,
|
||||
}
|
||||
row_to_batch = workspace.row_to_batch
|
||||
local_starts = workspace.local_starts
|
||||
logits = flydsl_pa_mqa_logits_fp4_prefill(
|
||||
q_payload,
|
||||
q_scale,
|
||||
k_payload,
|
||||
k_scale,
|
||||
page_table,
|
||||
weights,
|
||||
row_to_batch,
|
||||
local_starts,
|
||||
c4_seq_lens,
|
||||
max_seq_len,
|
||||
parallel_unit_num=max(_PREFILL_BASE_CTA_TARGET, num_tokens),
|
||||
**pinned,
|
||||
**common,
|
||||
)
|
||||
|
||||
return logits
|
||||
|
||||
|
||||
def prepare_fp4_k_write_metadata(
|
||||
plan: Union[CompressorDecodePlan, CompressorPrefillPlan],
|
||||
out_loc: torch.Tensor,
|
||||
rope_table_len: int,
|
||||
) -> FP4KWriteMetadata:
|
||||
"""
|
||||
Build RoPE positions and cache slots from a compressor plan.
|
||||
"""
|
||||
plan_words = plan[1].view(torch.int32)
|
||||
seq_lens = plan_words[:, 0].to(torch.int64)
|
||||
positions = seq_lens - plan.compress_ratio
|
||||
valid = (positions >= 0) & (positions < rope_table_len)
|
||||
positions = torch.where(valid, positions, torch.zeros_like(positions))
|
||||
valid &= seq_lens % plan.compress_ratio == 0
|
||||
|
||||
out_loc = out_loc.to(dtype=torch.int64)
|
||||
if plan.is_decode:
|
||||
slots = out_loc
|
||||
elif out_loc.shape[0] == 0:
|
||||
slots = torch.full_like(seq_lens, -1)
|
||||
valid.zero_()
|
||||
else:
|
||||
ragged_ids = plan_words[:, 1].bitwise_and(0xFFFF).to(torch.int64)
|
||||
valid &= ragged_ids < out_loc.shape[0]
|
||||
slots = out_loc[ragged_ids.clamp(max=out_loc.shape[0] - 1)]
|
||||
slots = torch.where(valid, slots, torch.full_like(slots, -1))
|
||||
return FP4KWriteMetadata(positions.contiguous(), slots.contiguous())
|
||||
|
||||
|
||||
def aiter_k_indexer_fp4_cache_write(
|
||||
*,
|
||||
k: torch.Tensor,
|
||||
norm_weight: torch.Tensor,
|
||||
norm_epsilon: float,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
plan: Union[CompressorDecodePlan, CompressorPrefillPlan],
|
||||
out_loc: torch.Tensor,
|
||||
k_payload: torch.Tensor,
|
||||
k_scale: torch.Tensor,
|
||||
write_metadata: Optional[FP4KWriteMetadata] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Map compressed K rows to cache slots and run the fused AITER FP4 writer.
|
||||
"""
|
||||
num_rows = k.shape[0]
|
||||
if num_rows == 0:
|
||||
return
|
||||
|
||||
assert write_metadata is not None, "FP4 K-write metadata is missing."
|
||||
|
||||
positions, slots = write_metadata
|
||||
# The compressor normally hands over its BF16 mirror; convert only when some
|
||||
# caller still passes the FP32 parameter.
|
||||
if norm_weight.dtype is not torch.bfloat16 or norm_weight.device != k.device:
|
||||
norm_weight = norm_weight.to(device=k.device, dtype=torch.bfloat16).contiguous()
|
||||
|
||||
import aiter
|
||||
|
||||
aiter.rmsnorm_rope_rotate_activation_fp4quant_kvcache(
|
||||
k_payload,
|
||||
k_scale,
|
||||
k.view(num_rows, 1, _HEAD_DIM),
|
||||
norm_weight,
|
||||
cos,
|
||||
sin,
|
||||
positions,
|
||||
slots,
|
||||
norm_epsilon,
|
||||
rope_dim=_ROPE_DIM,
|
||||
kv_block_size=_KV_BLOCK_SIZE,
|
||||
group_size=_GROUP_SIZE,
|
||||
shuffle_scale=True,
|
||||
do_rotate_act=True,
|
||||
)
|
||||
@@ -22,6 +22,7 @@ from sglang.srt.runtime_context import get_platform
|
||||
from sglang.srt.utils.common import (
|
||||
configure_media_url_security,
|
||||
get_device,
|
||||
is_gfx95_supported,
|
||||
is_mnnvl_fabric_device,
|
||||
)
|
||||
from sglang.utils import is_in_ci
|
||||
@@ -419,11 +420,11 @@ def handle_environment_variables(server_args: Any):
|
||||
"All operations will run eagerly through the graph capture/replay path."
|
||||
)
|
||||
if cfg.enable_deepseek_v4_fp4_indexer and not (
|
||||
get_platform().is_sm100 or get_platform().is_sm120
|
||||
get_platform().is_sm100 or get_platform().is_sm120 or is_gfx95_supported()
|
||||
):
|
||||
raise ValueError(
|
||||
"--enable-deepseek-v4-fp4-indexer requires SM100 or SM120 GPUs with "
|
||||
"DeepGEMM FP4 indexer support."
|
||||
"--enable-deepseek-v4-fp4-indexer requires SM100, SM120, or gfx95 GPUs "
|
||||
"with FP4 indexer support."
|
||||
)
|
||||
# FP8 W_o GEMM needs DeepGEMM JIT. Enable exactly where the runtime can run
|
||||
# it, mirroring the forward scale split: the ue8m0 path
|
||||
|
||||
@@ -47,6 +47,11 @@ from sglang.srt.utils import ceil_align
|
||||
if TYPE_CHECKING:
|
||||
from sgl_kernel.flash_mla import FlashMLASchedMeta
|
||||
|
||||
from sglang.kernels.ops.attention.dsv4.fp4_indexer_hip import (
|
||||
FP4DecodeWorkspace,
|
||||
FP4KWriteMetadata,
|
||||
FP4PrefillWorkspace,
|
||||
)
|
||||
from sglang.srt.layers.radix_attention import RadixAttention
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
|
||||
@@ -353,6 +358,18 @@ class DSV4Metadata:
|
||||
|
||||
c4_compress_metadata: Optional[FusedCompressMetadata] = None
|
||||
c128_compress_metadata: Optional[FusedCompressMetadata] = None
|
||||
# FP4 indexer buffers that captured kernels bind by address. Deliberately
|
||||
# absent from copy_: the addresses must stay pinned across replays, and the
|
||||
# workspace builders refresh their contents instead.
|
||||
fp4_decode_workspace: Optional[FP4DecodeWorkspace] = field(default=None, repr=False)
|
||||
fp4_prefill_workspace: Optional[FP4PrefillWorkspace] = field(
|
||||
default=None, repr=False
|
||||
)
|
||||
# Derived by the first C4 layer of a forward and reused by the rest.
|
||||
fp4_k_write_metadata: Optional[FP4KWriteMetadata] = field(default=None, repr=False)
|
||||
# AITER's rope kernels require int64 positions while the core metadata keeps
|
||||
# them int32, so widen once per forward instead of once per C4 layer.
|
||||
fp4_q_positions: Optional[torch.Tensor] = field(default=None, repr=False)
|
||||
|
||||
@property
|
||||
def core_metadata(self) -> DSV4AttnMetadata:
|
||||
@@ -655,6 +672,8 @@ class DeepseekV4HipRadixBackend(
|
||||
ragged_layout=None,
|
||||
) -> DSV4Metadata:
|
||||
batch_size = len(seq_lens)
|
||||
# Verify tokens may cross into the next page beyond the accepted prefix.
|
||||
max_seq_len += self.target_verify_num_draft_tokens
|
||||
extend_start_loc = None
|
||||
if ragged_layout is not None:
|
||||
verify_lens_dev = ragged_layout.verify_lens.to(
|
||||
@@ -760,6 +779,10 @@ class DeepseekV4HipRadixBackend(
|
||||
req_pool_indices = raw_metadata.req_pool_indices
|
||||
seq_lens = raw_metadata.seq_lens
|
||||
out_cache_loc = raw_metadata.out_cache_loc
|
||||
if self.topk > 0 and self.speculative_num_steps > 1:
|
||||
# Each EAGLE draft step appends one token while ForwardBatch keeps
|
||||
# the accepted-prefix lengths unchanged across the captured loop.
|
||||
seq_lens = seq_lens + self.speculative_step_id + 1
|
||||
|
||||
core_attn_metadata = self.make_core_attn_metadata(
|
||||
req_to_token=self.req_to_token,
|
||||
@@ -860,6 +883,83 @@ class DeepseekV4HipRadixBackend(
|
||||
)
|
||||
)
|
||||
|
||||
if self._fp4_workspaces_enabled(metadata):
|
||||
from sglang.kernels.ops.attention.dsv4.fp4_indexer_hip import (
|
||||
prepare_fp4_k_write_metadata,
|
||||
)
|
||||
|
||||
metadata.fp4_k_write_metadata = prepare_fp4_k_write_metadata(
|
||||
metadata.c4_compress_metadata,
|
||||
metadata.core_attn_metadata.c4_out_loc,
|
||||
self.MAX_SEQ_LEN_FOR_CAPTURE,
|
||||
)
|
||||
metadata.fp4_q_positions = metadata.core_attn_metadata.positions.to(
|
||||
torch.int64
|
||||
)
|
||||
|
||||
# Decode's schedule builder is capture-safe because the workspace pins
|
||||
# the scratch it reads, so it can stay next to the metadata it consumes.
|
||||
# Prefill/target-verify cannot; see _refresh_fp4_prefill_workspace.
|
||||
if self._fp4_workspaces_enabled(metadata) and (
|
||||
forward_batch.forward_mode.is_decode()
|
||||
):
|
||||
from sglang.kernels.ops.attention.dsv4.fp4_indexer_hip import (
|
||||
prepare_fp4_decode_workspace,
|
||||
)
|
||||
|
||||
indexer_metadata = metadata.indexer_metadata
|
||||
metadata.fp4_decode_workspace = prepare_fp4_decode_workspace(
|
||||
indexer_metadata.page_table,
|
||||
indexer_metadata.c4_seq_lens,
|
||||
)
|
||||
|
||||
def _fp4_workspaces_enabled(self, metadata) -> bool:
|
||||
return (
|
||||
self.enable_deepseek_v4_fp4_indexer
|
||||
# Draft-step backends drive the NextN layer, which is built with
|
||||
# compress_ratio_override=0 and so owns no C4 indexer. Their
|
||||
# workspaces would be built, scheduled, and never read.
|
||||
and self.speculative_num_steps == 0
|
||||
and isinstance(metadata, DSV4Metadata)
|
||||
and metadata.indexer_metadata is not None
|
||||
and metadata.c4_compress_metadata is not None
|
||||
)
|
||||
|
||||
def _refresh_fp4_prefill_workspace(self, forward_batch: ForwardBatch) -> None:
|
||||
"""Rebuild the FP4 prefill schedule outside CUDA-graph capture.
|
||||
|
||||
AITER's prefill scheduler frees the scratch that its schedule kernel
|
||||
reads, so recording the build into a graph would leave every replay
|
||||
reading recycled graph-pool memory. Only the pinned buffers it fills
|
||||
(cta_info / logits / guarded page table) may be read from the graph.
|
||||
"""
|
||||
metadata = self.forward_metadata
|
||||
if not self._fp4_workspaces_enabled(metadata):
|
||||
return
|
||||
if forward_batch.forward_mode not in (
|
||||
ForwardMode.EXTEND,
|
||||
ForwardMode.MIXED,
|
||||
ForwardMode.TARGET_VERIFY,
|
||||
):
|
||||
return
|
||||
if (
|
||||
get_parallel().attn_cp_size != 1
|
||||
or getattr(forward_batch, "tbo_children", None)
|
||||
or getattr(forward_batch, "tbo_parent_token_range", None) is not None
|
||||
):
|
||||
return
|
||||
|
||||
from sglang.kernels.ops.attention.dsv4.fp4_indexer_hip import (
|
||||
prepare_fp4_prefill_workspace,
|
||||
)
|
||||
|
||||
indexer_metadata = metadata.indexer_metadata
|
||||
metadata.fp4_prefill_workspace = prepare_fp4_prefill_workspace(
|
||||
indexer_metadata.page_table,
|
||||
indexer_metadata.c4_seq_lens,
|
||||
workspace=metadata.fp4_prefill_workspace,
|
||||
)
|
||||
|
||||
def init_forward_metadata_out_graph(
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
@@ -984,6 +1084,7 @@ class DeepseekV4HipRadixBackend(
|
||||
self.replay_cuda_graph_metadata_from(
|
||||
bs=graph_key, temp_metadata=temp_metadata, bucket=bucket
|
||||
)
|
||||
self._refresh_fp4_prefill_workspace(forward_batch)
|
||||
|
||||
if in_capture:
|
||||
metadata = self.forward_metadata
|
||||
@@ -1065,6 +1166,7 @@ class DeepseekV4HipRadixBackend(
|
||||
|
||||
self.forward_metadata = metadata
|
||||
self.init_forward_metadata_in_graph(forward_batch)
|
||||
self._refresh_fp4_prefill_workspace(forward_batch)
|
||||
|
||||
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int) -> None:
|
||||
self.cuda_graph_metadata_of_bucket_and_bs: Dict[
|
||||
@@ -1767,14 +1869,32 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend):
|
||||
):
|
||||
from types import SimpleNamespace
|
||||
|
||||
actual_forward_mode = getattr(
|
||||
forward_batch, "actual_forward_mode", forward_batch.forward_mode
|
||||
)
|
||||
out_cache_loc = getattr(forward_batch, "out_cache_loc", None)
|
||||
step_out_cache_locs = None
|
||||
# C4/C128 write locations are baked into each child backend's metadata,
|
||||
# so every speculative step must consume its own cache-location row.
|
||||
if (
|
||||
actual_forward_mode != ForwardMode.IDLE
|
||||
and out_cache_loc is not None
|
||||
and self.topk > 0
|
||||
and self.speculative_num_steps > 1
|
||||
):
|
||||
step_out_cache_locs = per_step_draft_out_cache_loc(
|
||||
out_cache_loc,
|
||||
forward_batch.batch_size,
|
||||
self.topk,
|
||||
self.speculative_num_steps,
|
||||
)
|
||||
|
||||
inner_fb = SimpleNamespace(
|
||||
batch_size=forward_batch.batch_size,
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
# Propagate the real runtime mode so inner backends can detect IDLE
|
||||
# and apply their idle substitution.
|
||||
actual_forward_mode=getattr(
|
||||
forward_batch, "actual_forward_mode", forward_batch.forward_mode
|
||||
),
|
||||
actual_forward_mode=actual_forward_mode,
|
||||
input_ids=getattr(forward_batch, "input_ids", None),
|
||||
positions=getattr(forward_batch, "positions", None),
|
||||
req_pool_indices=forward_batch.req_pool_indices,
|
||||
@@ -1782,23 +1902,37 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend):
|
||||
seq_lens_sum=forward_batch.seq_lens_sum,
|
||||
seq_lens_cpu=forward_batch.seq_lens_cpu,
|
||||
encoder_lens=None,
|
||||
out_cache_loc=getattr(forward_batch, "out_cache_loc", None),
|
||||
out_cache_loc=out_cache_loc,
|
||||
spec_info=forward_batch.spec_info,
|
||||
)
|
||||
if in_capture:
|
||||
for i in range(self.speculative_num_steps):
|
||||
if step_out_cache_locs is not None:
|
||||
inner_fb.out_cache_loc = step_out_cache_locs[i]
|
||||
self.attn_backends[i].init_forward_metadata_out_graph(
|
||||
inner_fb, in_capture=True
|
||||
)
|
||||
else:
|
||||
if self.speculative_num_steps == 1:
|
||||
return
|
||||
if step_out_cache_locs is not None:
|
||||
inner_fb.out_cache_loc = step_out_cache_locs[0]
|
||||
self.attn_backends[0].init_forward_metadata_out_graph(inner_fb)
|
||||
temp_metadata = self.attn_backends[0].forward_metadata
|
||||
if step_out_cache_locs is not None:
|
||||
assert isinstance(temp_metadata, DSV4RawDecodeMetadata)
|
||||
for i in range(1, self.speculative_num_steps - 1):
|
||||
if step_out_cache_locs is None:
|
||||
step_metadata = temp_metadata
|
||||
else:
|
||||
step_metadata = DSV4RawDecodeMetadata(
|
||||
req_pool_indices=temp_metadata.req_pool_indices,
|
||||
seq_lens=temp_metadata.seq_lens,
|
||||
out_cache_loc=step_out_cache_locs[i],
|
||||
)
|
||||
self.attn_backends[i].replay_cuda_graph_metadata_from(
|
||||
bs=forward_batch.batch_size,
|
||||
temp_metadata=temp_metadata,
|
||||
temp_metadata=step_metadata,
|
||||
bucket=_GraphBucket.DECODE_OR_IDLE,
|
||||
)
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ from sglang.srt.mem_cache.deepseek_v4_compress_state import (
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||
from sglang.srt.models.deepseek_v2 import _is_hip
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils import add_prefix, is_npu, set_weight_attrs
|
||||
|
||||
_is_npu = is_npu()
|
||||
@@ -374,11 +374,44 @@ class Compressor(BaseFusedOp):
|
||||
self.norm = RMSNorm(
|
||||
self.head_dim, eps=config.rms_norm_eps, weight_dtype=torch.float32
|
||||
)
|
||||
if (
|
||||
is_in_indexer
|
||||
and _is_hip
|
||||
and get_exec().kernel.enable_deepseek_v4_fp4_indexer
|
||||
):
|
||||
self._init_fp4_norm_weight()
|
||||
self.rotary_emb = rotary_emb
|
||||
self.freqs_cis = freqs_cis
|
||||
|
||||
self.ape_converted = False
|
||||
|
||||
def _init_fp4_norm_weight(self) -> None:
|
||||
"""Mirror the FP32 norm weight in BF16 for the AITER FP4 K writer.
|
||||
|
||||
The FP8 path feeds the FP32 weight straight to its kernel; AITER wants
|
||||
BF16, and converting at the call site costs one copy per C4 layer per
|
||||
forward. A buffer keeps the conversion out of the forward and survives
|
||||
module ``_apply``, and the loader below re-derives it so online weight
|
||||
updates propagate -- they land as ``param.data.copy_``, which leaves the
|
||||
parameter's identity and ``_version`` untouched and would silently
|
||||
defeat any cache keyed on those. Same reach as ``load_ape_weight``:
|
||||
``update_weights_from_tensor(load_format="direct")`` calls
|
||||
``default_weight_loader`` itself and so skips both hooks.
|
||||
"""
|
||||
self.norm.register_buffer(
|
||||
"fp4_weight_bf16",
|
||||
self.norm.weight.detach().to(torch.bfloat16).contiguous(),
|
||||
persistent=False,
|
||||
)
|
||||
set_weight_attrs(self.norm.weight, {"weight_loader": self.load_norm_weight})
|
||||
|
||||
def load_norm_weight(
|
||||
self, param: torch.Tensor, loaded_weight: torch.Tensor
|
||||
) -> None:
|
||||
assert param is self.norm.weight
|
||||
param.data.copy_(loaded_weight)
|
||||
self.norm.fp4_weight_bf16.copy_(param.data)
|
||||
|
||||
def _apply_ape_hotfix(self):
|
||||
self.ape_converted = True
|
||||
|
||||
|
||||
@@ -156,6 +156,8 @@ class CompressorBackendMixin:
|
||||
out_loc: torch.Tensor,
|
||||
use_fp4_indexer: bool = False,
|
||||
bf16_store: bool = False,
|
||||
kv_scale_cache: Optional[torch.Tensor] = None,
|
||||
rope_cache: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
) -> None:
|
||||
assert compress_ratio == 4 or compress_ratio == 128
|
||||
assert rotate == is_indexer == (head_dim == 128)
|
||||
@@ -175,6 +177,11 @@ class CompressorBackendMixin:
|
||||
kv_score_buffer = kv_score_buffer.view(-1, compress_ratio, last_dim)
|
||||
|
||||
# Step 1: compress_forward
|
||||
compress_out = None
|
||||
if _is_hip and use_fp4_indexer:
|
||||
compress_out = kv_score_input.new_empty(
|
||||
(plan[1].shape[0], head_dim), dtype=torch.bfloat16
|
||||
)
|
||||
kv_compressed = compress_forward(
|
||||
kv_score_buffer=kv_score_buffer,
|
||||
kv_score_input=kv_score_input,
|
||||
@@ -182,14 +189,21 @@ class CompressorBackendMixin:
|
||||
plan=plan,
|
||||
compress_ratio=compress_ratio,
|
||||
head_dim=head_dim,
|
||||
out=compress_out,
|
||||
is_online=is_online,
|
||||
)
|
||||
|
||||
# The AITER FP4 writer takes BF16; the compressor mirrors its FP32 norm
|
||||
# weight so the conversion stays out of the per-layer forward.
|
||||
norm_weight = norm.weight
|
||||
if _is_hip and use_fp4_indexer:
|
||||
norm_weight = getattr(norm, "fp4_weight_bf16", norm_weight)
|
||||
|
||||
# Step 2: norm + rope + store
|
||||
compress_norm_rope_store(
|
||||
kv_compressed,
|
||||
plan,
|
||||
norm_weight=norm.weight,
|
||||
norm_weight=norm_weight,
|
||||
norm_eps=norm.variance_epsilon,
|
||||
freq_cis=freqs_cis_cache,
|
||||
out_loc=out_loc,
|
||||
@@ -197,6 +211,15 @@ class CompressorBackendMixin:
|
||||
page_size=page_size,
|
||||
use_fp4=use_fp4_indexer,
|
||||
bf16_store=bf16_store,
|
||||
kvcache_scale=kv_scale_cache,
|
||||
rope_cache=rope_cache,
|
||||
# Derived once per forward by the backend; every C4 layer writes the
|
||||
# same rows to the same slots.
|
||||
fp4_k_write_metadata=(
|
||||
getattr(self.forward_metadata, "fp4_k_write_metadata", None)
|
||||
if _is_hip and use_fp4_indexer
|
||||
else None
|
||||
),
|
||||
)
|
||||
|
||||
def forward_unified(
|
||||
@@ -222,10 +245,16 @@ class CompressorBackendMixin:
|
||||
use_fp4_indexer = (
|
||||
compressor.is_in_indexer and self.enable_deepseek_v4_fp4_indexer
|
||||
)
|
||||
use_hip_fp4 = _is_hip and use_fp4_indexer
|
||||
bf16_store = False
|
||||
kv_scale_cache = None
|
||||
if compressor.is_in_indexer:
|
||||
kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer(layer_id)
|
||||
page_size = token_to_kv_pool.get_index_k_page_size()
|
||||
if use_hip_fp4:
|
||||
kv_cache = token_to_kv_pool.get_index_k_fp4_payload_buffer(layer_id)
|
||||
kv_scale_cache = token_to_kv_pool.get_index_k_fp4_scale_buffer(layer_id)
|
||||
else:
|
||||
kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer(layer_id)
|
||||
elif is_unified_kv_triton():
|
||||
kv_cache = token_to_kv_pool.get_unified_kv(layer_id)
|
||||
page_size = 1
|
||||
@@ -248,7 +277,7 @@ class CompressorBackendMixin:
|
||||
head_dim=compressor.head_dim,
|
||||
norm=compressor.norm,
|
||||
freqs_cis_cache=compressor.freqs_cis,
|
||||
kv_cache=kv_cache.view(dtype=torch.uint8),
|
||||
kv_cache=kv_cache if use_hip_fp4 else kv_cache.view(dtype=torch.uint8),
|
||||
is_indexer=compressor.is_in_indexer,
|
||||
rotate=compressor.rotate,
|
||||
compress_ratio=compressor.ratio,
|
||||
@@ -256,6 +285,10 @@ class CompressorBackendMixin:
|
||||
out_loc=out_loc,
|
||||
use_fp4_indexer=use_fp4_indexer,
|
||||
bf16_store=bf16_store,
|
||||
kv_scale_cache=kv_scale_cache,
|
||||
rope_cache=(
|
||||
(compressor.fp4_cos, compressor.fp4_sin) if use_hip_fp4 else None
|
||||
),
|
||||
)
|
||||
online_c128_mtp = getattr(self, "online_c128_mtp", None)
|
||||
if online_c128_mtp is not None:
|
||||
|
||||
@@ -22,6 +22,10 @@ from sglang.kernels.ops.attention.dsv4 import (
|
||||
topk_transform_paged,
|
||||
topk_transform_paged_v2,
|
||||
)
|
||||
from sglang.kernels.ops.attention.dsv4.fp4_indexer_hip import (
|
||||
aiter_fp4_paged_mqa_logits,
|
||||
aiter_q_indexer_fp4,
|
||||
)
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz
|
||||
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
|
||||
from sglang.srt.environ import envs
|
||||
@@ -694,8 +698,13 @@ class C4IndexerBackendMixin:
|
||||
core_metadata = metadata.core_metadata
|
||||
|
||||
assert isinstance(indexer_metadata, PagedIndexerMetadata)
|
||||
use_aiter_fp4 = c4_indexer.use_fp4_indexer and is_hip()
|
||||
|
||||
positions = core_metadata.positions
|
||||
if use_aiter_fp4:
|
||||
widened = getattr(metadata, "fp4_q_positions", None)
|
||||
if widened is not None and widened.shape == positions.shape:
|
||||
positions = widened
|
||||
num_queries = min(x.shape[0], q_lora.shape[0], positions.shape[0])
|
||||
if x.shape[0] != num_queries:
|
||||
x = x[:num_queries]
|
||||
@@ -727,7 +736,9 @@ class C4IndexerBackendMixin:
|
||||
|
||||
use_fp4_indexer = c4_indexer.use_fp4_indexer
|
||||
|
||||
if use_fp4_indexer:
|
||||
if use_aiter_fp4:
|
||||
q = q_indexer
|
||||
elif use_fp4_indexer:
|
||||
q_fp4, q_sf = q_indexer
|
||||
assert len(q_fp4.shape) == 3
|
||||
assert len(q_sf.shape) == 2
|
||||
@@ -736,13 +747,18 @@ class C4IndexerBackendMixin:
|
||||
assert len(q_indexer.shape) == 3
|
||||
q = q_indexer.unsqueeze(1)
|
||||
|
||||
assert len(weights.shape) == 3
|
||||
weights = weights.squeeze(2)
|
||||
if use_fp4_indexer:
|
||||
if use_aiter_fp4:
|
||||
weights = weights.contiguous()
|
||||
else:
|
||||
assert len(weights.shape) == 3
|
||||
weights = weights.squeeze(2)
|
||||
if use_fp4_indexer and not use_aiter_fp4:
|
||||
weights = weights.float()
|
||||
if envs.SGLANG_OPT_USE_TILELANG_INDEXER.get():
|
||||
raise RuntimeError("DeepSeek V4 FP4 indexer requires DeepGEMM indexer.")
|
||||
from deep_gemm import fp8_fp4_paged_mqa_logits as fn
|
||||
elif use_aiter_fp4:
|
||||
fn = None
|
||||
elif envs.SGLANG_OPT_USE_TILELANG_INDEXER.get():
|
||||
from sglang.kernels.ops.attention.dsa.tilelang_kernel import (
|
||||
tilelang_fp8_paged_mqa_logits as fn,
|
||||
@@ -773,7 +789,9 @@ class C4IndexerBackendMixin:
|
||||
pad = (0, 0) * (tensor.dim() - 1) + (0, query_rows - tensor.shape[0])
|
||||
return F.pad(tensor, pad, value=value)
|
||||
|
||||
c4_seq_lens = match_num_queries(indexer_metadata.c4_seq_lens, value=1)
|
||||
c4_seq_lens = match_num_queries(
|
||||
indexer_metadata.c4_seq_lens, value=0 if use_aiter_fp4 else 1
|
||||
)
|
||||
_c4sl = c4_seq_lens
|
||||
page_table = match_num_queries(indexer_metadata.page_table, value=0)
|
||||
c4_sparse_page_indices = match_num_queries(
|
||||
@@ -791,17 +809,39 @@ class C4IndexerBackendMixin:
|
||||
and not _use_tilelang
|
||||
and not _use_aiter
|
||||
and not _use_torch_fn
|
||||
and not use_aiter_fp4
|
||||
):
|
||||
_c4sl = _c4sl.unsqueeze(-1)
|
||||
nonpaged_plan = self._get_nonpaged_indexer_plan(
|
||||
c4_indexer=c4_indexer,
|
||||
forward_batch=forward_batch,
|
||||
indexer_metadata=indexer_metadata,
|
||||
page_table=page_table,
|
||||
c4_seq_lens=c4_seq_lens,
|
||||
query_rows=query_rows,
|
||||
)
|
||||
if nonpaged_plan is not None:
|
||||
nonpaged_plan = None
|
||||
if not use_aiter_fp4:
|
||||
nonpaged_plan = self._get_nonpaged_indexer_plan(
|
||||
c4_indexer=c4_indexer,
|
||||
forward_batch=forward_batch,
|
||||
indexer_metadata=indexer_metadata,
|
||||
page_table=page_table,
|
||||
c4_seq_lens=c4_seq_lens,
|
||||
query_rows=query_rows,
|
||||
)
|
||||
if use_aiter_fp4:
|
||||
q_fp4, q_scale = q
|
||||
logits = aiter_fp4_paged_mqa_logits(
|
||||
q_fp4=q_fp4,
|
||||
q_scale=q_scale,
|
||||
k_payload=token_to_kv_pool.get_index_k_fp4_payload_buffer(
|
||||
c4_indexer.layer_id
|
||||
),
|
||||
k_scale=token_to_kv_pool.get_index_k_fp4_scale_buffer(
|
||||
c4_indexer.layer_id
|
||||
),
|
||||
weights=weights,
|
||||
page_table=page_table,
|
||||
c4_seq_lens=c4_seq_lens,
|
||||
weight_scale=c4_indexer.weight_scale,
|
||||
is_decode=forward_batch.forward_mode.is_decode(),
|
||||
decode_workspace=metadata.fp4_decode_workspace,
|
||||
prefill_workspace=metadata.fp4_prefill_workspace,
|
||||
)
|
||||
elif nonpaged_plan is not None:
|
||||
assert isinstance(q_indexer, torch.Tensor)
|
||||
logits = self._forward_nonpaged_indexer(
|
||||
q_indexer=q_indexer,
|
||||
@@ -926,6 +966,8 @@ class C4Indexer(nn.Module):
|
||||
prefix: str = "",
|
||||
alt_streams: Optional[List[torch.cuda.Stream]] = None,
|
||||
rotary_emb=None,
|
||||
fp4_cos: Optional[torch.Tensor] = None,
|
||||
fp4_sin: Optional[torch.Tensor] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.layer_id = layer_id
|
||||
@@ -937,6 +979,7 @@ class C4Indexer(nn.Module):
|
||||
self.q_lora_rank = config.q_lora_rank
|
||||
self.softmax_scale = self.head_dim**-0.5
|
||||
self.n_local_heads = self.n_heads
|
||||
self.use_fp4_indexer = get_exec().kernel.enable_deepseek_v4_fp4_indexer
|
||||
self.wq_b = ReplicatedLinear(
|
||||
self.q_lora_rank,
|
||||
self.n_heads * self.head_dim,
|
||||
@@ -970,11 +1013,13 @@ class C4Indexer(nn.Module):
|
||||
quant_config=expert_pack_quant_config,
|
||||
rotary_emb=rotary_emb,
|
||||
)
|
||||
if self.use_fp4_indexer and is_hip():
|
||||
self.compressor.fp4_cos = fp4_cos
|
||||
self.compressor.fp4_sin = fp4_sin
|
||||
self.rotary_emb = rotary_emb
|
||||
self.freqs_cis = freqs_cis
|
||||
self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5
|
||||
|
||||
self.use_fp4_indexer = get_exec().kernel.enable_deepseek_v4_fp4_indexer
|
||||
self.alt_streams = alt_streams
|
||||
|
||||
def compute_q(
|
||||
@@ -985,6 +1030,14 @@ class C4Indexer(nn.Module):
|
||||
) -> Tuple[IndexerQuery, torch.Tensor]:
|
||||
q, _ = self.wq_b(q_lora)
|
||||
q = q.view(-1, self.n_local_heads, self.head_dim)
|
||||
if self.use_fp4_indexer and is_hip():
|
||||
q_fp4, q_scale = aiter_q_indexer_fp4(
|
||||
q.contiguous(),
|
||||
self.compressor.fp4_cos,
|
||||
self.compressor.fp4_sin,
|
||||
positions,
|
||||
)
|
||||
return (q_fp4, q_scale), weight
|
||||
if self.use_fp4_indexer:
|
||||
return fused_q_indexer_rope_hadamard_fp4_quant(
|
||||
q.contiguous(), weight, self.weight_scale, self.freqs_cis, positions
|
||||
|
||||
@@ -123,9 +123,7 @@ class PagedIndexerMetadata:
|
||||
|
||||
def __post_init__(self):
|
||||
if (
|
||||
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.get()
|
||||
or is_xpu()
|
||||
or envs.SGLANG_OPT_USE_AITER_INDEXER.get()
|
||||
is_hip() or is_xpu() or envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.get()
|
||||
) and not self.force_deep_gemm_metadata:
|
||||
self.deep_gemm_metadata = None
|
||||
else:
|
||||
|
||||
@@ -31,6 +31,13 @@ _is_hip = is_hip()
|
||||
ONLINE_C128 = not _is_hip and envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get()
|
||||
|
||||
|
||||
def get_dsv4_indexer_bytes_per_token(index_head_dim: int, use_fp4_indexer: bool) -> int:
|
||||
"""Return payload and quant-scale bytes for one compressed indexer token."""
|
||||
if use_fp4_indexer:
|
||||
return index_head_dim // 2 + index_head_dim // 32
|
||||
return index_head_dim + index_head_dim // 128 * 4
|
||||
|
||||
|
||||
def get_compress_state_ring_size(
|
||||
compress_ratio: int, is_speculative: bool = False
|
||||
) -> int:
|
||||
@@ -285,25 +292,47 @@ class DeepSeekV4IndexerPool(KVCache):
|
||||
)
|
||||
self.index_head_dim = index_head_dim
|
||||
self.use_fp4_indexer = get_exec().kernel.enable_deepseek_v4_fp4_indexer
|
||||
self.uses_aiter_fp4_layout = _is_hip and self.use_fp4_indexer
|
||||
|
||||
self._create_buffer()
|
||||
|
||||
def get_bytes_per_token(self) -> int:
|
||||
if self.use_fp4_indexer:
|
||||
return self.index_head_dim // 2 + 4
|
||||
return self.index_head_dim + 4
|
||||
return get_dsv4_indexer_bytes_per_token(
|
||||
self.index_head_dim, self.use_fp4_indexer
|
||||
)
|
||||
|
||||
def _create_buffer(self):
|
||||
page_bytes = self.page_size * self.get_bytes_per_token()
|
||||
num_pages = (self.size + self.page_size + 1) // self.page_size
|
||||
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
||||
with (
|
||||
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
||||
if self.custom_mem_pool
|
||||
else nullcontext()
|
||||
):
|
||||
if self.uses_aiter_fp4_layout:
|
||||
self.index_k_payload_buffer = [
|
||||
torch.zeros(
|
||||
(num_pages, 1, 4, self.page_size, 16),
|
||||
dtype=torch.uint8,
|
||||
device=self.device,
|
||||
).view(torch.float4_e2m1fn_x2)
|
||||
for _ in range(self.layer_num)
|
||||
]
|
||||
self.index_k_scale_buffer = [
|
||||
torch.zeros(
|
||||
(num_pages, 1, 4, self.page_size),
|
||||
dtype=torch.uint8,
|
||||
device=self.device,
|
||||
)
|
||||
for _ in range(self.layer_num)
|
||||
]
|
||||
self.index_k_with_scale_buffer = None
|
||||
return
|
||||
|
||||
self.index_k_with_scale_buffer = [
|
||||
torch.zeros(
|
||||
(self.size + self.page_size + 1) // self.page_size,
|
||||
num_pages,
|
||||
page_bytes,
|
||||
dtype=self.index_k_with_scale_buffer_dtype,
|
||||
device=self.device,
|
||||
@@ -326,6 +355,25 @@ class DeepSeekV4IndexerPool(KVCache):
|
||||
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
return self.index_k_with_scale_buffer[layer_id]
|
||||
|
||||
def contiguous_page_row_buffers(self) -> List[torch.Tensor]:
|
||||
"""Every indexer buffer as 2D page rows, for PD and HiCache transfer.
|
||||
|
||||
FP8 keeps key and scale fused in one buffer per layer; the FP4 layout
|
||||
stores payload and scale separately, so it yields two buffers per layer.
|
||||
"""
|
||||
if self.index_k_with_scale_buffer is not None:
|
||||
return self.index_k_with_scale_buffer
|
||||
return [
|
||||
buf.view(torch.uint8).flatten(1)
|
||||
for buf in (*self.index_k_payload_buffer, *self.index_k_scale_buffer)
|
||||
]
|
||||
|
||||
def get_index_k_fp4_payload_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
return self.index_k_payload_buffer[layer_id]
|
||||
|
||||
def get_index_k_fp4_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
return self.index_k_scale_buffer[layer_id]
|
||||
|
||||
def get_index_k_scale_buffer(
|
||||
self,
|
||||
layer_id: int,
|
||||
@@ -702,7 +750,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
||||
|
||||
for i in c4_locals:
|
||||
_append_compressed_entry(i, 4)
|
||||
for buf in self.c4_indexer_kv_pool.index_k_with_scale_buffer:
|
||||
for buf in self.c4_indexer_kv_pool.contiguous_page_row_buffers():
|
||||
assert buf.ndim == 2, f"expected 2D buffer, got {buf.ndim}D"
|
||||
data_ptrs.append(buf.data_ptr())
|
||||
data_lens.append(buf.nbytes)
|
||||
@@ -714,7 +762,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
||||
|
||||
buf_groups = [
|
||||
self.c4_kv_pool.kv_buffer,
|
||||
self.c4_indexer_kv_pool.index_k_with_scale_buffer,
|
||||
self.c4_indexer_kv_pool.contiguous_page_row_buffers(),
|
||||
self.c128_kv_pool.kv_buffer,
|
||||
]
|
||||
|
||||
@@ -1109,6 +1157,18 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
||||
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }"
|
||||
return self.c4_indexer_kv_pool.get_index_k_with_scale_buffer(compress_layer_id)
|
||||
|
||||
def get_index_k_fp4_payload_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
self.wait_layer_transfer(layer_id)
|
||||
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }"
|
||||
return self.c4_indexer_kv_pool.get_index_k_fp4_payload_buffer(compress_layer_id)
|
||||
|
||||
def get_index_k_fp4_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
self.wait_layer_transfer(layer_id)
|
||||
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }"
|
||||
return self.c4_indexer_kv_pool.get_index_k_fp4_scale_buffer(compress_layer_id)
|
||||
|
||||
def get_index_k_scale_buffer(
|
||||
self,
|
||||
layer_id: int,
|
||||
|
||||
@@ -66,6 +66,9 @@ class PoolName(str, Enum):
|
||||
# 'COMPRESSED_KV / COMPRESSED_INDEXER / COMPRESSED_STATE' in the next PR.
|
||||
DEEPSEEK_V4_C4 = "deepseek_v4_c4"
|
||||
DEEPSEEK_V4_C4_INDEXER = "deepseek_v4_c4_indexer"
|
||||
# FP4 indexer splits the indexer cache into separate payload/scale buffers,
|
||||
# so it needs a second pool alongside DEEPSEEK_V4_C4_INDEXER.
|
||||
DEEPSEEK_V4_C4_INDEXER_SCALE = "deepseek_v4_c4_indexer_scale"
|
||||
DEEPSEEK_V4_C128 = "deepseek_v4_c128"
|
||||
DEEPSEEK_V4_C4_STATE = "deepseek_v4_c4_state"
|
||||
DEEPSEEK_V4_C4_INDEXER_STATE = "deepseek_v4_c4_indexer_state"
|
||||
|
||||
@@ -470,6 +470,72 @@ def _dsv4_compressed_region_buffers(kvcache: Any, ratio: int) -> tuple[list, int
|
||||
return pool.kv_buffer, pool.bytes_per_page_padded
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _IndexerRegion:
|
||||
"""One page-contiguous indexer buffer group to mirror on the host."""
|
||||
|
||||
name: PoolName
|
||||
device_buffers: list
|
||||
item_bytes: int
|
||||
# FP4 page rows group their slots instead of laying tokens out flat, so the
|
||||
# fused-row token-granular copy does not apply and transfers must be whole
|
||||
# pages. The fused FP8 row has no such restriction.
|
||||
page_aligned_only: bool
|
||||
|
||||
|
||||
def _dsv4_indexer_regions(kvcache: Any, page_size: int) -> list[_IndexerRegion]:
|
||||
"""
|
||||
Resolve the indexer HiCache regions, hiding the FP8/FP4 split from the
|
||||
stack builder. FP8 keeps key and scale fused in one buffer, while FP4
|
||||
stores payload and scale separately, so it maps to two host pools.
|
||||
"""
|
||||
import torch
|
||||
|
||||
pool = kvcache.c4_indexer_kv_pool
|
||||
fused = pool.index_k_with_scale_buffer
|
||||
if fused is not None:
|
||||
return [
|
||||
_IndexerRegion(
|
||||
name=PoolName.DEEPSEEK_V4_C4_INDEXER,
|
||||
device_buffers=fused,
|
||||
item_bytes=fused[0].shape[1] * fused[0].element_size(),
|
||||
page_aligned_only=False,
|
||||
)
|
||||
]
|
||||
|
||||
payload_ref = pool.index_k_payload_buffer[0]
|
||||
scale_ref = pool.index_k_scale_buffer[0]
|
||||
# A page row covers ``page_slots`` C4 slots, i.e. one tree page of tokens
|
||||
# after 4:1 compression.
|
||||
page_slots = payload_ref.shape[3]
|
||||
if scale_ref.shape[3] != page_slots:
|
||||
raise ValueError(
|
||||
"FP4 indexer payload and scale must agree on slots per page: "
|
||||
f"payload={page_slots}, scale={scale_ref.shape[3]}"
|
||||
)
|
||||
if page_size % page_slots != 0:
|
||||
raise ValueError(
|
||||
f"Tree page size {page_size} must be a multiple of the FP4 indexer "
|
||||
f"slots per page {page_slots}"
|
||||
)
|
||||
payload = [b.view(torch.uint8).flatten(1) for b in pool.index_k_payload_buffer]
|
||||
scale = [b.view(torch.uint8).flatten(1) for b in pool.index_k_scale_buffer]
|
||||
return [
|
||||
_IndexerRegion(
|
||||
name=PoolName.DEEPSEEK_V4_C4_INDEXER,
|
||||
device_buffers=payload,
|
||||
item_bytes=payload[0].shape[1],
|
||||
page_aligned_only=True,
|
||||
),
|
||||
_IndexerRegion(
|
||||
name=PoolName.DEEPSEEK_V4_C4_INDEXER_SCALE,
|
||||
device_buffers=scale,
|
||||
item_bytes=scale[0].shape[1],
|
||||
page_aligned_only=True,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def build_deepseek_v4_hicache_stack(
|
||||
*,
|
||||
params: CacheInitParams,
|
||||
@@ -582,36 +648,34 @@ def build_deepseek_v4_hicache_stack(
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(),
|
||||
)
|
||||
c4_indexer_host_pool = DeepSeekV4PagedHostPool(
|
||||
pool_name=str(PoolName.DEEPSEEK_V4_C4_INDEXER),
|
||||
device_buffers=kvcache.c4_indexer_kv_pool.index_k_with_scale_buffer,
|
||||
item_bytes=(
|
||||
kvcache.c4_indexer_kv_pool.index_k_with_scale_buffer[0].shape[1]
|
||||
* kvcache.c4_indexer_kv_pool.index_k_with_scale_buffer[0].element_size()
|
||||
),
|
||||
num_host_pages=num_host_pages,
|
||||
slot_page_size=page_size,
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(),
|
||||
entries.append(
|
||||
build_pool_entry(
|
||||
name=PoolName.DEEPSEEK_V4_C4,
|
||||
host_pool=c4_host_pool,
|
||||
device_pool=kvcache.c4_kv_pool,
|
||||
layer_mapping=c4_layer_mapping,
|
||||
transfer_layer_num=transfer_layer_num,
|
||||
)
|
||||
)
|
||||
entries.extend(
|
||||
[
|
||||
for region in _dsv4_indexer_regions(kvcache, page_size):
|
||||
entries.append(
|
||||
build_pool_entry(
|
||||
name=PoolName.DEEPSEEK_V4_C4,
|
||||
host_pool=c4_host_pool,
|
||||
device_pool=kvcache.c4_kv_pool,
|
||||
layer_mapping=c4_layer_mapping,
|
||||
transfer_layer_num=transfer_layer_num,
|
||||
),
|
||||
build_pool_entry(
|
||||
name=PoolName.DEEPSEEK_V4_C4_INDEXER,
|
||||
host_pool=c4_indexer_host_pool,
|
||||
name=region.name,
|
||||
host_pool=DeepSeekV4PagedHostPool(
|
||||
pool_name=str(region.name),
|
||||
device_buffers=region.device_buffers,
|
||||
item_bytes=region.item_bytes,
|
||||
num_host_pages=num_host_pages,
|
||||
slot_page_size=page_size,
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(),
|
||||
page_aligned_only=region.page_aligned_only,
|
||||
),
|
||||
device_pool=kvcache.c4_indexer_kv_pool,
|
||||
layer_mapping=c4_layer_mapping,
|
||||
transfer_layer_num=transfer_layer_num,
|
||||
),
|
||||
]
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
if not is_unified_kv:
|
||||
c4_state_host_pool = DeepSeekV4StateHostPool(
|
||||
@@ -1276,6 +1340,7 @@ class _DeepSeekV4Strategy(StackStrategy):
|
||||
for name, src in (
|
||||
(PoolName.DEEPSEEK_V4_C4, PoolName.KV),
|
||||
(PoolName.DEEPSEEK_V4_C4_INDEXER, PoolName.KV),
|
||||
(PoolName.DEEPSEEK_V4_C4_INDEXER_SCALE, PoolName.KV),
|
||||
(PoolName.DEEPSEEK_V4_C128, PoolName.KV),
|
||||
(PoolName.DEEPSEEK_V4_C4_STATE, PoolName.SWA),
|
||||
(PoolName.DEEPSEEK_V4_C4_INDEXER_STATE, PoolName.SWA),
|
||||
|
||||
@@ -239,6 +239,7 @@ def _build_deepseek_v4_device_pool_group(
|
||||
) -> DevicePoolGroup:
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import HiSparseC4DevicePool
|
||||
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
|
||||
_dsv4_indexer_regions,
|
||||
_resolve_deepseek_v4_layer_mappings,
|
||||
)
|
||||
|
||||
@@ -288,13 +289,14 @@ def _build_deepseek_v4_device_pool_group(
|
||||
kvcache.c4_kv_pool.kv_buffer,
|
||||
mappings.c4,
|
||||
)
|
||||
add(
|
||||
PoolName.DEEPSEEK_V4_C4_INDEXER,
|
||||
PoolName.KV,
|
||||
kvcache.c4_indexer_kv_pool,
|
||||
kvcache.c4_indexer_kv_pool.index_k_with_scale_buffer,
|
||||
mappings.c4,
|
||||
)
|
||||
for region in _dsv4_indexer_regions(kvcache, page_size):
|
||||
add(
|
||||
region.name,
|
||||
PoolName.KV,
|
||||
kvcache.c4_indexer_kv_pool,
|
||||
region.device_buffers,
|
||||
mappings.c4,
|
||||
)
|
||||
add(
|
||||
PoolName.DEEPSEEK_V4_C128,
|
||||
PoolName.KV,
|
||||
|
||||
@@ -185,10 +185,15 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache):
|
||||
device: str = "cpu",
|
||||
pin_memory: bool = True,
|
||||
allocator_type: str = "default",
|
||||
page_aligned_only: bool = False,
|
||||
):
|
||||
self.pool_name = pool_name
|
||||
self.layer_num = len(device_buffers)
|
||||
self.item_bytes = item_bytes
|
||||
# A page row of the FP4 indexer buffers is a grouped slot layout rather
|
||||
# than a flat token array, so the token-granular copy used for fused
|
||||
# DSv4 C4 rows does not apply and only whole pages may move.
|
||||
self.page_aligned_only = page_aligned_only
|
||||
self.num_host_pages = num_host_pages
|
||||
self.slot_page_size = slot_page_size
|
||||
self.dtype = torch.uint8
|
||||
@@ -305,6 +310,15 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache):
|
||||
def _to_page_indices(self, indices: torch.Tensor) -> torch.Tensor:
|
||||
return indices.reshape(-1, self.slot_page_size)[:, 0] // self.slot_page_size
|
||||
|
||||
def _unaligned_transfer_error(
|
||||
self, host_indices: torch.Tensor, device_indices: torch.Tensor
|
||||
) -> ValueError:
|
||||
return ValueError(
|
||||
f"{self.pool_name} expects page-aligned indices: got "
|
||||
f"{host_indices.numel()} host and {device_indices.numel()} device "
|
||||
f"indices for page size {self.slot_page_size}."
|
||||
)
|
||||
|
||||
def _has_transfer_indices(
|
||||
self, host_indices: torch.Tensor | None, device_indices: torch.Tensor | None
|
||||
) -> bool:
|
||||
@@ -375,6 +389,8 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache):
|
||||
# Token-granular DSV4 C4 copy needs this helper because a token is
|
||||
# not one contiguous byte range in the paged row:
|
||||
# [value0..value63][scale0..scale63].
|
||||
if self.page_aligned_only:
|
||||
raise self._unaligned_transfer_error(host_indices, device_indices)
|
||||
transfer_cache_dsv4_mla(
|
||||
src_ptrs=self.device_ptrs,
|
||||
dst_ptrs=self.data_ptrs,
|
||||
@@ -453,6 +469,8 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache):
|
||||
):
|
||||
# Same DSV4 C4 layout issue as backup: this is token-granular
|
||||
# preload, so it cannot use the normal HiCache page-row copy.
|
||||
if self.page_aligned_only:
|
||||
raise self._unaligned_transfer_error(host_indices, device_indices)
|
||||
transfer_cache_dsv4_mla(
|
||||
src_ptrs=self.data_ptrs[layer_id : layer_id + 1],
|
||||
dst_ptrs=self.device_ptrs[layer_id : layer_id + 1],
|
||||
|
||||
@@ -799,6 +799,7 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
|
||||
PoolName.DRAFT_INDEXER,
|
||||
PoolName.DEEPSEEK_V4_C4,
|
||||
PoolName.DEEPSEEK_V4_C4_INDEXER,
|
||||
PoolName.DEEPSEEK_V4_C4_INDEXER_SCALE,
|
||||
PoolName.DEEPSEEK_V4_C128,
|
||||
PoolName.DEEPSEEK_V4_C4_STATE,
|
||||
PoolName.DEEPSEEK_V4_C4_INDEXER_STATE,
|
||||
|
||||
@@ -36,6 +36,7 @@ from sglang.srt.mem_cache.allocation_sizing import get_alloc_len_per_decode
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
|
||||
get_compress_state_ring_size,
|
||||
get_compress_state_write_pad,
|
||||
get_dsv4_indexer_bytes_per_token,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
|
||||
from sglang.srt.runtime_context import (
|
||||
@@ -50,9 +51,12 @@ from sglang.srt.utils.common import (
|
||||
ceil_align,
|
||||
ceil_div,
|
||||
is_float4_e2m1fn_x2,
|
||||
is_hip,
|
||||
spec_decode_alloc_len_per_request,
|
||||
)
|
||||
|
||||
_is_hip = is_hip()
|
||||
|
||||
|
||||
@dataclass
|
||||
class MemoryPoolConfig:
|
||||
@@ -774,6 +778,12 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
|
||||
self.qk_nope_head_dim = cfg.qk_nope_head_dim
|
||||
self.qk_rope_head_dim = cfg.qk_rope_head_dim
|
||||
self.indexer_head_dim = cfg.index_head_dim
|
||||
# HIP takes the FP4-accurate byte count here. The NVIDIA FP4 path
|
||||
# keeps the FP8 estimate.
|
||||
self.indexer_bytes_per_token = get_dsv4_indexer_bytes_per_token(
|
||||
self.indexer_head_dim,
|
||||
_is_hip and kvc.server_args.enable_deepseek_v4_fp4_indexer,
|
||||
)
|
||||
self.context_len = kvc.model_config.context_len
|
||||
# PP-local slice; matches DeepSeekV4TokenToKVPool's stage_ratios.
|
||||
self.compression_ratios = cfg.compress_ratios[
|
||||
@@ -884,11 +894,6 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
|
||||
def _get_bytes_per_full_token(self) -> float:
|
||||
kv_bytes = self.qk_nope_head_dim + self.qk_rope_head_dim * 2 + 8
|
||||
|
||||
quant_block_size = 128
|
||||
indexer_bytes = (
|
||||
self.indexer_head_dim + self.indexer_head_dim // quant_block_size * 4
|
||||
)
|
||||
|
||||
attn_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim
|
||||
c4_state_dtype_size, c128_state_dtype_size = (
|
||||
_get_dsv4_compress_state_dtype_sizes()
|
||||
@@ -914,7 +919,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
|
||||
self.swa_ratio * kv_bytes * self.num_layers_total
|
||||
+ c4_frac * kv_bytes * self.num_layers_ca4
|
||||
+ 1 / 128 * kv_bytes * self.num_layers_ca128
|
||||
+ 1 / 4 * indexer_bytes * self.num_layers_ca4
|
||||
+ 1 / 4 * self.indexer_bytes_per_token * self.num_layers_ca4
|
||||
+ self.swa_ratio * c4_state_ratio * c4_state_bytes * self.num_layers_ca4
|
||||
+ c128_state_ratio * c128_state_bytes * self.num_layers_ca128
|
||||
+ self.swa_ratio
|
||||
|
||||
@@ -891,6 +891,8 @@ class MQALayer(MqaAttentionBase):
|
||||
prefix=add_prefix("indexer", prefix),
|
||||
alt_streams=self.alt_streams_indexer,
|
||||
rotary_emb=self.rotary_emb,
|
||||
fp4_cos=(self.cos_cache[:, 0, 0, :] if _is_hip else None),
|
||||
fp4_sin=(self.sin_cache[:, 0, 0, :] if _is_hip else None),
|
||||
)
|
||||
|
||||
self.attn_mqa = RadixAttention(
|
||||
@@ -911,6 +913,13 @@ class MQALayer(MqaAttentionBase):
|
||||
# (`_compute_kv_to_cache`), so the legacy "overlap store cache" flag
|
||||
# has no effect here -- the fused path is on by default.
|
||||
|
||||
def _apply(self, fn, recurse=True):
|
||||
result = super()._apply(fn, recurse=recurse)
|
||||
if self.indexer is not None and hasattr(self.indexer.compressor, "fp4_cos"):
|
||||
self.indexer.compressor.fp4_cos = self.cos_cache[:, 0, 0, :]
|
||||
self.indexer.compressor.fp4_sin = self.sin_cache[:, 0, 0, :]
|
||||
return result
|
||||
|
||||
def _get_npu_rope_position_cache(
|
||||
self, positions: torch.Tensor, dtype: torch.dtype, inverse: bool = False
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
|
||||
Reference in New Issue
Block a user