[HiSparse] Add MHA hisparse support for MiniMax M3 (#31446)
Co-authored-by: Guangda Liu <bingps@users.noreply.github.com>
This commit is contained in:
co-authored by
Guangda Liu
parent
095e45100b
commit
04c0913434
@@ -129,6 +129,7 @@ The NVIDIA Blackwell recipes are validated single-node: **B200 at `--tp 8`** and
|
|||||||
|
|
||||||
- **Memory**: `--mem-fraction-static` reserves GPU memory for weights + KV pool; the rest is prefill **activation headroom**. The value scales with *free* memory per GPU (card capacity minus per-GPU weight), so it tracks the card more than the TP degree: **`0.65` on B200** (180 GB — less headroom once weights are resident) and **`0.75` on the larger-memory B300 / GB300** (`0.80` on AMD). Lower TP packs more weight per GPU, so a tighter config needs a *lower* value — B200 needs `0.65` even at `--tp 4`. Raising it past the validated value is fine only for low-concurrency single-stream serving; it OOMs under high concurrency or long context.
|
- **Memory**: `--mem-fraction-static` reserves GPU memory for weights + KV pool; the rest is prefill **activation headroom**. The value scales with *free* memory per GPU (card capacity minus per-GPU weight), so it tracks the card more than the TP degree: **`0.65` on B200** (180 GB — less headroom once weights are resident) and **`0.75` on the larger-memory B300 / GB300** (`0.80` on AMD). Lower TP packs more weight per GPU, so a tighter config needs a *lower* value — B200 needs `0.65` even at `--tp 4`. Raising it past the validated value is fine only for low-concurrency single-stream serving; it OOMs under high concurrency or long context.
|
||||||
- **Long context (32K+)**: keep `--mem-fraction-static` at the platform default and raise `--chunked-prefill-size` to `16384`. Decode TPOT stays roughly flat in context length thanks to sparse attention; 1K–128K prompts are validated.
|
- **Long context (32K+)**: keep `--mem-fraction-static` at the platform default and raise `--chunked-prefill-size` to `16384`. Decode TPOT stays roughly flat in context length thanks to sparse attention; 1K–128K prompts are validated.
|
||||||
|
- **HiSparse for decode capacity**: on NVIDIA CUDA, HiSparse keeps the three dense layers on GPU, moves the 57 sparse-layer K/V caches to pinned host memory, and feeds selected block IDs directly to the swap-in kernel. For the released four-KV-head model, use `--tp 4` or greater, `--disable-radix-cache`, and `device_buffer_size >= 2048`. Enable it with `--enable-hisparse --hisparse-config='{"device_buffer_size":4096,"host_to_device_ratio":2}'` on the Triton launch command.
|
||||||
- **Scaling TP**: B200 is documented at `--tp 8`; B300 / GB200 / GB300 at `--tp 4` (the single-node cross-family common denominator). On an 8-GPU B300 host you can also raise to `--tp 8` for more throughput / KV headroom.
|
- **Scaling TP**: B200 is documented at `--tp 8`; B300 / GB200 / GB300 at `--tp 4` (the single-node cross-family common denominator). On an 8-GPU B300 host you can also raise to `--tp 8` for more throughput / KV headroom.
|
||||||
- **Expert parallelism**: to trade latency for throughput add `--ep` (see [Expert Parallelism Deployment](../../../docs/advanced_features/expert_parallelism)). On AMD, set `--ep` equal to `--tp`. Shared-experts fusion is automatically disabled when EP > 1; on AMD standard EP the server also disables `--enable-aiter-allreduce-fusion` automatically to preserve accuracy.
|
- **Expert parallelism**: to trade latency for throughput add `--ep` (see [Expert Parallelism Deployment](../../../docs/advanced_features/expert_parallelism)). On AMD, set `--ep` equal to `--tp`. Shared-experts fusion is automatically disabled when EP > 1; on AMD standard EP the server also disables `--enable-aiter-allreduce-fusion` automatically to preserve accuracy.
|
||||||
- `--trust-remote-code` is required to load the MiniMax config / processor classes.
|
- `--trust-remote-code` is required to load the MiniMax config / processor classes.
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ metatags:
|
|||||||
|
|
||||||
HiSparse reduces per-request GPU memory consumption during the decode phase by maintaining only a small "hot" KV buffer on GPU while keeping complete KV data in CPU pinned memory. Combined with PD disaggregation, it enables significantly higher decode concurrency.
|
HiSparse reduces per-request GPU memory consumption during the decode phase by maintaining only a small "hot" KV buffer on GPU while keeping complete KV data in CPU pinned memory. Combined with PD disaggregation, it enables significantly higher decode concurrency.
|
||||||
|
|
||||||
> **Prerequisites**: HiSparse works with models that use **DeepSeek Sparse Attention (DSA)** architectures (e.g., DeepSeek-V3.2, GLM-5.1) and **DeepSeek V4**. These models natively select a subset of tokens for attention, making it possible to keep only the top-k KV on GPU while storing the full KV in host memory — without accuracy loss. Additionally, HiSparse currently requires **PD disaggregation mode** and is enabled on the **decode instance** only.
|
> **Prerequisites**: HiSparse works with models that use **DeepSeek Sparse Attention (DSA)** architectures (e.g., DeepSeek-V3.2, GLM-5.1), **DeepSeek V4**, and **MiniMax M3**. These models natively select a subset of tokens for attention, making it possible to keep only the top-k KV on GPU while storing the full KV in host memory — without accuracy loss. Additionally, HiSparse currently requires **PD disaggregation mode** and is enabled on the **decode instance** only.
|
||||||
|
|
||||||
## Why HiSparse?
|
## Why HiSparse?
|
||||||
|
|
||||||
@@ -165,6 +165,15 @@ python3 -m sglang.launch_server \
|
|||||||
|
|
||||||
> **Note**: For DSA models, `--kv-cache-dtype` defaults to `auto`, which resolves to `fp8_e4m3` on SM100+ (Blackwell) and `bfloat16` on older architectures. The DSA decode backend is automatically selected based on KV dtype (`bfloat16` → `flashmla_sparse`, `fp8_e4m3` → `flashmla_kv`), except for GLM DSA models on SM120/SM121 with `fp8_e4m3`, which use `flashinfer_sparse_mla`. DSA backend flags apply only to DSA models; DeepSeek V4 uses its own `dsv4` attention backend.
|
> **Note**: For DSA models, `--kv-cache-dtype` defaults to `auto`, which resolves to `fp8_e4m3` on SM100+ (Blackwell) and `bfloat16` on older architectures. The DSA decode backend is automatically selected based on KV dtype (`bfloat16` → `flashmla_sparse`, `fp8_e4m3` → `flashmla_kv`), except for GLM DSA models on SM120/SM121 with `fp8_e4m3`, which use `flashinfer_sparse_mla`. DSA backend flags apply only to DSA models; DeepSeek V4 uses its own `dsv4` attention backend.
|
||||||
|
|
||||||
|
### MiniMax M3
|
||||||
|
|
||||||
|
Dense-layer K/V and index K stay on GPU; sparse-layer K/V use host memory plus a GPU working set.
|
||||||
|
|
||||||
|
- Use TP 4 or greater, with the same TP size and PP 1 on both PD instances.
|
||||||
|
- Use `--attention-backend triton`, `--mm-attention-backend triton_attn`, `--disable-prefill-cuda-graph`, and `--disable-radix-cache`.
|
||||||
|
- Set `device_buffer_size` to at least 2048 in `--hisparse-config`; `top_k` does not override the model's selection width.
|
||||||
|
- PD retraction backup is not supported. Use `--num-reserved-decode-tokens` to reserve capacity for the expected output length.
|
||||||
|
|
||||||
### Benchmark
|
### Benchmark
|
||||||
|
|
||||||
```bash Command
|
```bash Command
|
||||||
|
|||||||
@@ -303,13 +303,26 @@ template <int NUM_TOP_K, int HOT_BUFFER_SIZE>
|
|||||||
struct SmemLayout {
|
struct SmemLayout {
|
||||||
static constexpr int HASH_SIZE = NUM_TOP_K * 2;
|
static constexpr int HASH_SIZE = NUM_TOP_K * 2;
|
||||||
static constexpr int NUM_BUFFER_CHUNKS = (HOT_BUFFER_SIZE + WARP_SIZE - 1) / WARP_SIZE;
|
static constexpr int NUM_BUFFER_CHUNKS = (HOT_BUFFER_SIZE + WARP_SIZE - 1) / WARP_SIZE;
|
||||||
// int32_t region: top_k_tokens + chunk_offset + evict_chunk_offset + hash_keys + total_hits + newest_hit
|
// int32_t region: top_k_tokens + chunk offsets + hash keys + hit counters
|
||||||
static constexpr int TOTAL_INT32 = NUM_TOP_K + (NUM_BUFFER_CHUNKS + 1) + (NUM_BUFFER_CHUNKS + 1) + HASH_SIZE + 2;
|
static constexpr int TOTAL_INT32 = NUM_TOP_K + (NUM_BUFFER_CHUNKS + 1) + (NUM_BUFFER_CHUNKS + 1) + HASH_SIZE + 2;
|
||||||
// int16_t region: lru_slots_out + hash_vals
|
// int16_t region: lru_slots_out + hash_vals
|
||||||
static constexpr int TOTAL_INT16 = HOT_BUFFER_SIZE + HASH_SIZE;
|
static constexpr int TOTAL_INT16 = HOT_BUFFER_SIZE + HASH_SIZE;
|
||||||
static constexpr size_t BYTES = TOTAL_INT32 * sizeof(int32_t) + TOTAL_INT16 * sizeof(int16_t);
|
static constexpr size_t BYTES = TOTAL_INT32 * sizeof(int32_t) + TOTAL_INT16 * sizeof(int16_t);
|
||||||
};
|
};
|
||||||
|
|
||||||
|
template <int SPARSE_BLOCK_SIZE, bool TopKIsBlocks>
|
||||||
|
__device__ __forceinline__ int32_t resolve_selected_token(const int32_t* top_k, int32_t token_index) {
|
||||||
|
if constexpr (TopKIsBlocks) {
|
||||||
|
const int32_t block_index = top_k[token_index / SPARSE_BLOCK_SIZE];
|
||||||
|
if (block_index < 0) {
|
||||||
|
return -1;
|
||||||
|
}
|
||||||
|
return block_index * SPARSE_BLOCK_SIZE + token_index % SPARSE_BLOCK_SIZE;
|
||||||
|
} else {
|
||||||
|
return top_k[token_index];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Each block processes one request
|
// Each block processes one request
|
||||||
// req_pool_indices and seq_lens can each be int32_t or int64_t
|
// req_pool_indices and seq_lens can each be int32_t or int64_t
|
||||||
// Layout: [HOT_BUFFER_SIZE slots for LRU] + [page_size slots for newest token]
|
// Layout: [HOT_BUFFER_SIZE slots for LRU] + [page_size slots for newest token]
|
||||||
@@ -319,23 +332,28 @@ struct SmemLayout {
|
|||||||
// false -> generic byte-stride: device + host both linear, stride = item_size_bytes
|
// false -> generic byte-stride: device + host both linear, stride = item_size_bytes
|
||||||
// true -> DSv4 page-padded device + page-padded host (kvcacheio.cuh constants)
|
// true -> DSv4 page-padded device + page-padded host (kvcacheio.cuh constants)
|
||||||
//
|
//
|
||||||
|
// TopKIsBlocks makes the kernel consume block ids directly. It resolves token
|
||||||
|
// positions in registers and writes the flattened token-slot table expected by
|
||||||
|
// sparse attention without materializing an intermediate token-index tensor.
|
||||||
// RecordMissPlan records this step's miss plan (miss_src/dst = host/device loc
|
// RecordMissPlan records this step's miss plan (miss_src/dst = host/device loc
|
||||||
// per miss, miss_count per request) for shared-index skip layers to replay via
|
// per miss, miss_count per request) for shared-index skip layers to replay via
|
||||||
// copy_cache_planned_kernel. SkipIO elides only the KV byte movement (timing
|
// copy_cache_planned_kernel. SkipIO elides only the KV byte movement (timing
|
||||||
// probe; output is garbage). Both are compile-time flags so the production
|
// probe; output is garbage). These are compile-time flags, so inactive paths
|
||||||
// (false, false) instantiation stays byte-identical.
|
// are removed from each specialization.
|
||||||
template <
|
template <
|
||||||
int BLOCK_SIZE,
|
int BLOCK_SIZE,
|
||||||
int NUM_TOP_K,
|
int NUM_TOP_K,
|
||||||
int HOT_BUFFER_SIZE,
|
int HOT_BUFFER_SIZE,
|
||||||
bool IsMLA,
|
bool IsMLA,
|
||||||
bool IsDsv4Layout,
|
bool IsDsv4Layout,
|
||||||
|
int SPARSE_BLOCK_SIZE,
|
||||||
|
bool TopKIsBlocks,
|
||||||
bool RecordMissPlan,
|
bool RecordMissPlan,
|
||||||
bool SkipIO,
|
bool SkipIO,
|
||||||
typename SeqLensT,
|
typename SeqLensT,
|
||||||
typename ReqPoolIndicesT>
|
typename ReqPoolIndicesT>
|
||||||
__global__ void load_cache_to_device_buffer_kernel(
|
__global__ void load_cache_to_device_buffer_kernel(
|
||||||
const int32_t* __restrict__ top_k_tokens,
|
const int32_t* __restrict__ top_k,
|
||||||
int32_t* __restrict__ device_buffer_tokens,
|
int32_t* __restrict__ device_buffer_tokens,
|
||||||
const int64_t* __restrict__ host_cache_locs,
|
const int64_t* __restrict__ host_cache_locs,
|
||||||
const int32_t* __restrict__ device_buffer_locs,
|
const int32_t* __restrict__ device_buffer_locs,
|
||||||
@@ -351,7 +369,7 @@ __global__ void load_cache_to_device_buffer_kernel(
|
|||||||
int64_t buffer_stride_0,
|
int64_t buffer_stride_0,
|
||||||
int64_t host_stride,
|
int64_t host_stride,
|
||||||
int64_t lru_slot_stride_0,
|
int64_t lru_slot_stride_0,
|
||||||
int64_t top_k_tokens_stride,
|
int64_t top_k_stride,
|
||||||
int64_t top_k_device_locs_stride,
|
int64_t top_k_device_locs_stride,
|
||||||
int64_t page_size,
|
int64_t page_size,
|
||||||
int64_t item_size_bytes,
|
int64_t item_size_bytes,
|
||||||
@@ -360,9 +378,12 @@ __global__ void load_cache_to_device_buffer_kernel(
|
|||||||
int32_t* __restrict__ miss_count_out,
|
int32_t* __restrict__ miss_count_out,
|
||||||
int64_t plan_stride) {
|
int64_t plan_stride) {
|
||||||
static_assert(!IsDsv4Layout || IsMLA, "DSv4 page-padded layout is K-only (MLA).");
|
static_assert(!IsDsv4Layout || IsMLA, "DSv4 page-padded layout is K-only (MLA).");
|
||||||
// todo hisparse: support page wise sparsity
|
static_assert(SPARSE_BLOCK_SIZE > 0, "SPARSE_BLOCK_SIZE must be positive.");
|
||||||
|
// Cache residency and LRU replacement remain token-granular even when the
|
||||||
|
// sparse-attention selection arrives as block ids.
|
||||||
|
constexpr int NUM_TOP_K_TOKENS = NUM_TOP_K * (TopKIsBlocks ? SPARSE_BLOCK_SIZE : 1);
|
||||||
constexpr int NUM_WARPS = BLOCK_SIZE / WARP_SIZE;
|
constexpr int NUM_WARPS = BLOCK_SIZE / WARP_SIZE;
|
||||||
constexpr int NUM_TOKEN_CHUNKS = (NUM_TOP_K + WARP_SIZE - 1) / WARP_SIZE;
|
constexpr int NUM_TOKEN_CHUNKS = (NUM_TOP_K_TOKENS + WARP_SIZE - 1) / WARP_SIZE;
|
||||||
constexpr int NUM_BUFFER_CHUNKS = (HOT_BUFFER_SIZE + WARP_SIZE - 1) / WARP_SIZE;
|
constexpr int NUM_BUFFER_CHUNKS = (HOT_BUFFER_SIZE + WARP_SIZE - 1) / WARP_SIZE;
|
||||||
|
|
||||||
const int bid = blockIdx.x;
|
const int bid = blockIdx.x;
|
||||||
@@ -372,7 +393,7 @@ __global__ void load_cache_to_device_buffer_kernel(
|
|||||||
// CUDA graph pads the batch to a captured size. Keep padded output rows
|
// CUDA graph pads the batch to a captured size. Keep padded output rows
|
||||||
// invalid without a separate fill kernel.
|
// invalid without a separate fill kernel.
|
||||||
if (bid >= num_real_reqs[0]) {
|
if (bid >= num_real_reqs[0]) {
|
||||||
for (int i = tid; i < NUM_TOP_K; i += BLOCK_SIZE) {
|
for (int i = tid; i < NUM_TOP_K_TOKENS; i += BLOCK_SIZE) {
|
||||||
req_top_k_device_locs[i] = -1;
|
req_top_k_device_locs[i] = -1;
|
||||||
}
|
}
|
||||||
return;
|
return;
|
||||||
@@ -386,7 +407,7 @@ __global__ void load_cache_to_device_buffer_kernel(
|
|||||||
const int64_t seq_len = seq_lens[bid];
|
const int64_t seq_len = seq_lens[bid];
|
||||||
|
|
||||||
// Calculate offsets for this request
|
// Calculate offsets for this request
|
||||||
const int32_t* req_top_k_tokens = top_k_tokens + bid * top_k_tokens_stride;
|
const int32_t* req_top_k = top_k + bid * top_k_stride;
|
||||||
|
|
||||||
const int64_t buffer_offset = rid * buffer_stride_0;
|
const int64_t buffer_offset = rid * buffer_stride_0;
|
||||||
int32_t* req_device_buffer_tokens = device_buffer_tokens + buffer_offset;
|
int32_t* req_device_buffer_tokens = device_buffer_tokens + buffer_offset;
|
||||||
@@ -396,14 +417,16 @@ __global__ void load_cache_to_device_buffer_kernel(
|
|||||||
|
|
||||||
// Fast path: short sequences have all tokens in the device buffer in order.
|
// Fast path: short sequences have all tokens in the device buffer in order.
|
||||||
if (seq_len <= HOT_BUFFER_SIZE) {
|
if (seq_len <= HOT_BUFFER_SIZE) {
|
||||||
const int count = (seq_len < NUM_TOP_K) ? static_cast<int>(seq_len) : NUM_TOP_K;
|
const int count = (seq_len < NUM_TOP_K_TOKENS) ? static_cast<int>(seq_len) : NUM_TOP_K_TOKENS;
|
||||||
for (int i = tid; i < NUM_TOP_K; i += BLOCK_SIZE) {
|
for (int i = tid; i < NUM_TOP_K_TOKENS; i += BLOCK_SIZE) {
|
||||||
int32_t device_loc = -1;
|
int32_t device_loc = -1;
|
||||||
if (i < count) {
|
const int32_t token_pos = resolve_selected_token<SPARSE_BLOCK_SIZE, TopKIsBlocks>(req_top_k, i);
|
||||||
int32_t token_pos = req_top_k_tokens[i];
|
if constexpr (TopKIsBlocks) {
|
||||||
if (token_pos >= 0) {
|
if (token_pos >= 0 && token_pos < seq_len) {
|
||||||
device_loc = req_device_buffer_locs[token_pos];
|
device_loc = req_device_buffer_locs[token_pos];
|
||||||
}
|
}
|
||||||
|
} else if (i < count && token_pos >= 0) {
|
||||||
|
device_loc = req_device_buffer_locs[token_pos];
|
||||||
}
|
}
|
||||||
req_top_k_device_locs[i] = device_loc;
|
req_top_k_device_locs[i] = device_loc;
|
||||||
}
|
}
|
||||||
@@ -418,21 +441,21 @@ __global__ void load_cache_to_device_buffer_kernel(
|
|||||||
|
|
||||||
// Dynamic shared memory layout: int32_t arrays first, then int16_t arrays.
|
// Dynamic shared memory layout: int32_t arrays first, then int16_t arrays.
|
||||||
extern __shared__ char smem_raw[];
|
extern __shared__ char smem_raw[];
|
||||||
using Layout = SmemLayout<NUM_TOP_K, HOT_BUFFER_SIZE>;
|
using Layout = SmemLayout<NUM_TOP_K_TOKENS, HOT_BUFFER_SIZE>;
|
||||||
constexpr int HASH_SIZE = Layout::HASH_SIZE;
|
constexpr int HASH_SIZE = Layout::HASH_SIZE;
|
||||||
|
|
||||||
int32_t* smem_i32 = reinterpret_cast<int32_t*>(smem_raw);
|
int32_t* smem_i32 = reinterpret_cast<int32_t*>(smem_raw);
|
||||||
// Top-k token positions; reused as miss-token scratch in the copy phase
|
// Top-k token positions; reused as miss-token scratch in the copy phase
|
||||||
int32_t* s_top_k_tokens = smem_i32;
|
int32_t* s_top_k_tokens = smem_i32;
|
||||||
// Prefix-sum offsets for hit counting and miss counting
|
// Prefix-sum offsets for hit counting and miss counting
|
||||||
int32_t* s_chunk_offset = s_top_k_tokens + NUM_TOP_K;
|
int32_t* s_chunk_offset = s_top_k_tokens + NUM_TOP_K_TOKENS;
|
||||||
// Prefix-sum offsets for evictable counting
|
// Prefix-sum offsets for evictable counting
|
||||||
int32_t* s_evict_chunk_offset = s_chunk_offset + (NUM_BUFFER_CHUNKS + 1);
|
int32_t* s_evict_chunk_offset = s_chunk_offset + (NUM_BUFFER_CHUNKS + 1);
|
||||||
// Open-addressing hash table: top-k token_id -> top-k index (keys)
|
// Open-addressing hash table: top-k token_id -> top-k index (keys)
|
||||||
int32_t* s_hash_keys = s_evict_chunk_offset + (NUM_BUFFER_CHUNKS + 1);
|
int32_t* s_hash_keys = s_evict_chunk_offset + (NUM_BUFFER_CHUNKS + 1);
|
||||||
// Scalar counters
|
// Scalar counters
|
||||||
int32_t& s_total_hits = s_hash_keys[HASH_SIZE];
|
int32_t& s_total_hits = s_hash_keys[HASH_SIZE];
|
||||||
int32_t& s_newest_hit = s_hash_keys[HASH_SIZE + 1];
|
int32_t& s_total_misses = s_hash_keys[HASH_SIZE + 1];
|
||||||
|
|
||||||
int16_t* smem_i16 = reinterpret_cast<int16_t*>(smem_i32 + Layout::TOTAL_INT32);
|
int16_t* smem_i16 = reinterpret_cast<int16_t*>(smem_i32 + Layout::TOTAL_INT32);
|
||||||
// Compacted slot ordering: [hits fwd-> ... <-evictables bwd]
|
// Compacted slot ordering: [hits fwd-> ... <-evictables bwd]
|
||||||
@@ -443,7 +466,7 @@ __global__ void load_cache_to_device_buffer_kernel(
|
|||||||
// Initialize shared memory: counters, hash table, prefix-sum offsets.
|
// Initialize shared memory: counters, hash table, prefix-sum offsets.
|
||||||
if (tid == 0) {
|
if (tid == 0) {
|
||||||
s_total_hits = 0;
|
s_total_hits = 0;
|
||||||
s_newest_hit = 0;
|
s_total_misses = 0;
|
||||||
}
|
}
|
||||||
for (int i = tid; i < HASH_SIZE; i += BLOCK_SIZE) {
|
for (int i = tid; i < HASH_SIZE; i += BLOCK_SIZE) {
|
||||||
s_hash_keys[i] = HASH_EMPTY;
|
s_hash_keys[i] = HASH_EMPTY;
|
||||||
@@ -458,14 +481,20 @@ __global__ void load_cache_to_device_buffer_kernel(
|
|||||||
const int32_t newest_token = seq_len - 1;
|
const int32_t newest_token = seq_len - 1;
|
||||||
|
|
||||||
// Insert top-k tokens into shared-memory hash table.
|
// Insert top-k tokens into shared-memory hash table.
|
||||||
for (int i = tid; i < NUM_TOP_K; i += BLOCK_SIZE) {
|
for (int i = tid; i < NUM_TOP_K_TOKENS; i += BLOCK_SIZE) {
|
||||||
int32_t token_idx = req_top_k_tokens[i];
|
const int32_t token_idx = resolve_selected_token<SPARSE_BLOCK_SIZE, TopKIsBlocks>(req_top_k, i);
|
||||||
|
if constexpr (TopKIsBlocks) {
|
||||||
|
if (token_idx < 0 || token_idx >= seq_len) {
|
||||||
|
s_top_k_tokens[i] = TOKEN_HIT;
|
||||||
|
req_top_k_device_locs[i] = -1;
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
}
|
||||||
if (token_idx == newest_token) {
|
if (token_idx == newest_token) {
|
||||||
// If topk includes the latest token, bind its canonical occurrence to newest_slot (at HOT_BUFFER_SIZE) and mark
|
// If topk includes the latest token, bind its canonical occurrence to newest_slot (at HOT_BUFFER_SIZE) and mark
|
||||||
// it as a hit. newest_slot is at the first position of the extra page, excluded from LRU tracking.
|
// it as a hit. newest_slot is at the first position of the extra page, excluded from LRU tracking.
|
||||||
s_top_k_tokens[i] = TOKEN_HIT;
|
s_top_k_tokens[i] = TOKEN_HIT;
|
||||||
req_top_k_device_locs[i] = req_device_buffer_locs[newest_slot];
|
req_top_k_device_locs[i] = req_device_buffer_locs[newest_slot];
|
||||||
s_newest_hit = 1;
|
|
||||||
} else {
|
} else {
|
||||||
int slot = hash_slot(token_idx, HASH_SIZE);
|
int slot = hash_slot(token_idx, HASH_SIZE);
|
||||||
while (true) {
|
while (true) {
|
||||||
@@ -580,7 +609,7 @@ __global__ void load_cache_to_device_buffer_kernel(
|
|||||||
|
|
||||||
const int chunk_token_start = chunk_idx * WARP_SIZE;
|
const int chunk_token_start = chunk_idx * WARP_SIZE;
|
||||||
const int my_token_idx = chunk_token_start + lane_id;
|
const int my_token_idx = chunk_token_start + lane_id;
|
||||||
const bool has_valid_token = has_valid_chunk && (my_token_idx < NUM_TOP_K);
|
const bool has_valid_token = has_valid_chunk && (my_token_idx < NUM_TOP_K_TOKENS);
|
||||||
|
|
||||||
int32_t my_token = 0;
|
int32_t my_token = 0;
|
||||||
bool is_miss = false;
|
bool is_miss = false;
|
||||||
@@ -611,6 +640,9 @@ __global__ void load_cache_to_device_buffer_kernel(
|
|||||||
#else
|
#else
|
||||||
total_misses = warp_inclusive_scan(s_chunk_offset, lane_id, chunk_idx + 1, NUM_TOKEN_CHUNKS + 1, total_misses);
|
total_misses = warp_inclusive_scan(s_chunk_offset, lane_id, chunk_idx + 1, NUM_TOKEN_CHUNKS + 1, total_misses);
|
||||||
#endif
|
#endif
|
||||||
|
if (tid == 0) {
|
||||||
|
s_total_misses = total_misses;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
__syncthreads();
|
__syncthreads();
|
||||||
|
|
||||||
@@ -632,7 +664,7 @@ __global__ void load_cache_to_device_buffer_kernel(
|
|||||||
}
|
}
|
||||||
__syncthreads();
|
__syncthreads();
|
||||||
|
|
||||||
total_misses = NUM_TOP_K - s_total_hits - s_newest_hit;
|
total_misses = s_total_misses;
|
||||||
if constexpr (RecordMissPlan) {
|
if constexpr (RecordMissPlan) {
|
||||||
if (tid == 0) {
|
if (tid == 0) {
|
||||||
miss_count_out[bid] = total_misses;
|
miss_count_out[bid] = total_misses;
|
||||||
@@ -695,10 +727,12 @@ template <
|
|||||||
int HOT_BUFFER_SIZE,
|
int HOT_BUFFER_SIZE,
|
||||||
bool IsMLA,
|
bool IsMLA,
|
||||||
bool IsDsv4Layout,
|
bool IsDsv4Layout,
|
||||||
|
int SPARSE_BLOCK_SIZE,
|
||||||
|
bool TopKIsBlocks,
|
||||||
bool RecordMissPlan,
|
bool RecordMissPlan,
|
||||||
bool SkipIO>
|
bool SkipIO>
|
||||||
void load_cache_to_device_buffer(
|
void load_cache_to_device_buffer(
|
||||||
tvm::ffi::TensorView top_k_tokens,
|
tvm::ffi::TensorView top_k,
|
||||||
tvm::ffi::TensorView device_buffer_tokens,
|
tvm::ffi::TensorView device_buffer_tokens,
|
||||||
tvm::ffi::TensorView host_cache_locs,
|
tvm::ffi::TensorView host_cache_locs,
|
||||||
tvm::ffi::TensorView device_buffer_locs,
|
tvm::ffi::TensorView device_buffer_locs,
|
||||||
@@ -718,7 +752,8 @@ void load_cache_to_device_buffer(
|
|||||||
tvm::ffi::TensorView miss_count_out) {
|
tvm::ffi::TensorView miss_count_out) {
|
||||||
using namespace host;
|
using namespace host;
|
||||||
|
|
||||||
const int64_t bs = top_k_tokens.shape()[0];
|
constexpr int NUM_TOP_K_TOKENS = NUM_TOP_K * (TopKIsBlocks ? SPARSE_BLOCK_SIZE : 1);
|
||||||
|
const int64_t bs = top_k.shape()[0];
|
||||||
const int64_t host_stride = host_cache_locs.shape()[1];
|
const int64_t host_stride = host_cache_locs.shape()[1];
|
||||||
// Miss-plan side outputs; 0-dim sentinels when RecordMissPlan is false.
|
// Miss-plan side outputs; 0-dim sentinels when RecordMissPlan is false.
|
||||||
int64_t* const miss_src_ptr = RecordMissPlan ? static_cast<int64_t*>(miss_src_out.data_ptr()) : nullptr;
|
int64_t* const miss_src_ptr = RecordMissPlan ? static_cast<int64_t*>(miss_src_out.data_ptr()) : nullptr;
|
||||||
@@ -730,9 +765,9 @@ void load_cache_to_device_buffer(
|
|||||||
}
|
}
|
||||||
const int64_t buffer_stride_0 = device_buffer_tokens.strides()[0];
|
const int64_t buffer_stride_0 = device_buffer_tokens.strides()[0];
|
||||||
const int64_t lru_slot_stride_0 = lru_slots.strides()[0];
|
const int64_t lru_slot_stride_0 = lru_slots.strides()[0];
|
||||||
const int64_t top_k_tokens_stride = top_k_tokens.strides()[0];
|
const int64_t top_k_stride = top_k.strides()[0];
|
||||||
const int64_t top_k_device_locs_stride = top_k_device_locs.strides()[0];
|
const int64_t top_k_device_locs_stride = top_k_device_locs.strides()[0];
|
||||||
const auto kernel_device = top_k_tokens.device();
|
const auto kernel_device = top_k.device();
|
||||||
const auto device = LaunchKernel::resolve_device(kernel_device);
|
const auto device = LaunchKernel::resolve_device(kernel_device);
|
||||||
const void* const host_cache_k_ptr = runtime::get_device_accessible_ptr(host_cache_k);
|
const void* const host_cache_k_ptr = runtime::get_device_accessible_ptr(host_cache_k);
|
||||||
const void* const host_cache_v_ptr =
|
const void* const host_cache_v_ptr =
|
||||||
@@ -741,7 +776,7 @@ void load_cache_to_device_buffer(
|
|||||||
// Generic lambda: int32/int64 kernel variants are compiled for both
|
// Generic lambda: int32/int64 kernel variants are compiled for both
|
||||||
// seq_lens and req_pool_indices; the correct combo is selected at runtime.
|
// seq_lens and req_pool_indices; the correct combo is selected at runtime.
|
||||||
auto launch = [&](auto kernel_fn, const auto* seq_lens_ptr, const auto* req_pool_indices_ptr) {
|
auto launch = [&](auto kernel_fn, const auto* seq_lens_ptr, const auto* req_pool_indices_ptr) {
|
||||||
constexpr size_t smem_bytes = SmemLayout<NUM_TOP_K, HOT_BUFFER_SIZE>::BYTES;
|
constexpr size_t smem_bytes = SmemLayout<NUM_TOP_K_TOKENS, HOT_BUFFER_SIZE>::BYTES;
|
||||||
#ifndef USE_ROCM
|
#ifndef USE_ROCM
|
||||||
if constexpr (smem_bytes > 48u * 1024u) {
|
if constexpr (smem_bytes > 48u * 1024u) {
|
||||||
cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes);
|
cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes);
|
||||||
@@ -749,7 +784,7 @@ void load_cache_to_device_buffer(
|
|||||||
#endif
|
#endif
|
||||||
LaunchKernel(bs, BLOCK_SIZE, device, smem_bytes)(
|
LaunchKernel(bs, BLOCK_SIZE, device, smem_bytes)(
|
||||||
kernel_fn,
|
kernel_fn,
|
||||||
static_cast<const int32_t*>(top_k_tokens.data_ptr()),
|
static_cast<const int32_t*>(top_k.data_ptr()),
|
||||||
static_cast<int32_t*>(device_buffer_tokens.data_ptr()),
|
static_cast<int32_t*>(device_buffer_tokens.data_ptr()),
|
||||||
static_cast<const int64_t*>(host_cache_locs.data_ptr()),
|
static_cast<const int64_t*>(host_cache_locs.data_ptr()),
|
||||||
static_cast<const int32_t*>(device_buffer_locs.data_ptr()),
|
static_cast<const int32_t*>(device_buffer_locs.data_ptr()),
|
||||||
@@ -765,7 +800,7 @@ void load_cache_to_device_buffer(
|
|||||||
buffer_stride_0,
|
buffer_stride_0,
|
||||||
host_stride,
|
host_stride,
|
||||||
lru_slot_stride_0,
|
lru_slot_stride_0,
|
||||||
top_k_tokens_stride,
|
top_k_stride,
|
||||||
top_k_device_locs_stride,
|
top_k_device_locs_stride,
|
||||||
page_size,
|
page_size,
|
||||||
item_size_bytes,
|
item_size_bytes,
|
||||||
@@ -788,6 +823,8 @@ void load_cache_to_device_buffer(
|
|||||||
HOT_BUFFER_SIZE,
|
HOT_BUFFER_SIZE,
|
||||||
IsMLA,
|
IsMLA,
|
||||||
IsDsv4Layout,
|
IsDsv4Layout,
|
||||||
|
SPARSE_BLOCK_SIZE,
|
||||||
|
TopKIsBlocks,
|
||||||
RecordMissPlan,
|
RecordMissPlan,
|
||||||
SkipIO,
|
SkipIO,
|
||||||
int64_t,
|
int64_t,
|
||||||
@@ -802,6 +839,8 @@ void load_cache_to_device_buffer(
|
|||||||
HOT_BUFFER_SIZE,
|
HOT_BUFFER_SIZE,
|
||||||
IsMLA,
|
IsMLA,
|
||||||
IsDsv4Layout,
|
IsDsv4Layout,
|
||||||
|
SPARSE_BLOCK_SIZE,
|
||||||
|
TopKIsBlocks,
|
||||||
RecordMissPlan,
|
RecordMissPlan,
|
||||||
SkipIO,
|
SkipIO,
|
||||||
int64_t,
|
int64_t,
|
||||||
@@ -816,6 +855,8 @@ void load_cache_to_device_buffer(
|
|||||||
HOT_BUFFER_SIZE,
|
HOT_BUFFER_SIZE,
|
||||||
IsMLA,
|
IsMLA,
|
||||||
IsDsv4Layout,
|
IsDsv4Layout,
|
||||||
|
SPARSE_BLOCK_SIZE,
|
||||||
|
TopKIsBlocks,
|
||||||
RecordMissPlan,
|
RecordMissPlan,
|
||||||
SkipIO,
|
SkipIO,
|
||||||
int32_t,
|
int32_t,
|
||||||
@@ -830,6 +871,8 @@ void load_cache_to_device_buffer(
|
|||||||
HOT_BUFFER_SIZE,
|
HOT_BUFFER_SIZE,
|
||||||
IsMLA,
|
IsMLA,
|
||||||
IsDsv4Layout,
|
IsDsv4Layout,
|
||||||
|
SPARSE_BLOCK_SIZE,
|
||||||
|
TopKIsBlocks,
|
||||||
RecordMissPlan,
|
RecordMissPlan,
|
||||||
SkipIO,
|
SkipIO,
|
||||||
int32_t,
|
int32_t,
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from ..common.utils import (
|
|||||||
"BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["max_topk"]),
|
"BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["max_topk"]),
|
||||||
"HAS_SINK": lambda args: args["sink_ptr"] is not None,
|
"HAS_SINK": lambda args: args["sink_ptr"] is not None,
|
||||||
"BATCH_SIZE_BUCKET": lambda args: triton.next_power_of_2(args["batch_size"]),
|
"BATCH_SIZE_BUCKET": lambda args: triton.next_power_of_2(args["batch_size"]),
|
||||||
|
"HAS_HISPARSE_SLOTS": lambda args: args["hisparse_slots_ptr"] is not None,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
@triton.autotune(
|
@triton.autotune(
|
||||||
@@ -43,6 +44,7 @@ def _gqa_share_sparse_decode_kernel(
|
|||||||
idx_ptr, # topk index: qh x b x topk
|
idx_ptr, # topk index: qh x b x topk
|
||||||
o_ptr, # O partial: c x b x qh x d
|
o_ptr, # O partial: c x b x qh x d
|
||||||
lse_ptr, # lse partial: c x b x qh
|
lse_ptr, # lse partial: c x b x qh
|
||||||
|
hisparse_slots_ptr, # pre-resolved device slots: kh x b x (topk * block)
|
||||||
seq_lens,
|
seq_lens,
|
||||||
slot_ids,
|
slot_ids,
|
||||||
# shape
|
# shape
|
||||||
@@ -52,6 +54,8 @@ def _gqa_share_sparse_decode_kernel(
|
|||||||
head_dim,
|
head_dim,
|
||||||
max_topk,
|
max_topk,
|
||||||
max_kv_len,
|
max_kv_len,
|
||||||
|
hisparse_slots_stride_h,
|
||||||
|
hisparse_slots_stride_b,
|
||||||
# sm_scale
|
# sm_scale
|
||||||
sm_scale,
|
sm_scale,
|
||||||
# per-tensor KV dequant scales (1.0 when the cache is unit-scaled)
|
# per-tensor KV dequant scales (1.0 when the cache is unit-scaled)
|
||||||
@@ -89,6 +93,7 @@ def _gqa_share_sparse_decode_kernel(
|
|||||||
NUM_TOPK_CHUNKS: tl.constexpr,
|
NUM_TOPK_CHUNKS: tl.constexpr,
|
||||||
HAS_SINK: tl.constexpr,
|
HAS_SINK: tl.constexpr,
|
||||||
IS_FP8: tl.constexpr,
|
IS_FP8: tl.constexpr,
|
||||||
|
HAS_HISPARSE_SLOTS: tl.constexpr,
|
||||||
):
|
):
|
||||||
# decode program ids: split-K over the topk dimension to give every SM
|
# decode program ids: split-K over the topk dimension to give every SM
|
||||||
# something to do at small batch. pid(0) folds (batch, chunk) together so
|
# something to do at small batch. pid(0) folds (batch, chunk) together so
|
||||||
@@ -161,18 +166,30 @@ def _gqa_share_sparse_decode_kernel(
|
|||||||
# only iterate over this chunk's topk slice. the load must respect the
|
# only iterate over this chunk's topk slice. the load must respect the
|
||||||
# per-chunk start offset.
|
# per-chunk start offset.
|
||||||
cur_idx_ptr = idx_base + chunk_start_topk * stride_ti_t
|
cur_idx_ptr = idx_base + chunk_start_topk * stride_ti_t
|
||||||
|
hisparse_topk_counter = chunk_start_topk
|
||||||
for _ in tl.range(chunk_start_topk, chunk_end_topk):
|
for _ in tl.range(chunk_start_topk, chunk_end_topk):
|
||||||
# load index
|
# load index
|
||||||
c = tl.load(cur_idx_ptr).to(tl.int32) * BLOCK_SIZE_N
|
c = tl.load(cur_idx_ptr).to(tl.int32) * BLOCK_SIZE_N
|
||||||
cur_idx_ptr = cur_idx_ptr + stride_ti_t
|
cur_idx_ptr = cur_idx_ptr + stride_ti_t
|
||||||
# resolve slots for this block via req_to_token
|
|
||||||
pos = c + off_n
|
pos = c + off_n
|
||||||
pos_mask = pos < seq_len
|
pos_mask = pos < seq_len
|
||||||
slots = tl.load(
|
if HAS_HISPARSE_SLOTS:
|
||||||
req_to_token_ptr + sid * stride_r2t_b + pos,
|
slots = tl.load(
|
||||||
mask=pos_mask,
|
hisparse_slots_ptr
|
||||||
other=0,
|
+ pid_kh * hisparse_slots_stride_h
|
||||||
).to(tl.int64)
|
+ pid_b * hisparse_slots_stride_b
|
||||||
|
+ hisparse_topk_counter * BLOCK_SIZE_N
|
||||||
|
+ off_n,
|
||||||
|
mask=off_n < BLOCK_SIZE_N,
|
||||||
|
other=0,
|
||||||
|
).to(tl.int64)
|
||||||
|
hisparse_topk_counter = hisparse_topk_counter + 1
|
||||||
|
else:
|
||||||
|
slots = tl.load(
|
||||||
|
req_to_token_ptr + sid * stride_r2t_b + pos,
|
||||||
|
mask=pos_mask,
|
||||||
|
other=0,
|
||||||
|
).to(tl.int64)
|
||||||
slots = (slots + max_slots) % max_slots # safety against negative
|
slots = (slots + max_slots) % max_slots # safety against negative
|
||||||
# load K as (head_dim, BLOCK_SIZE_N) via indirect addressing
|
# load K as (head_dim, BLOCK_SIZE_N) via indirect addressing
|
||||||
k_off = (
|
k_off = (
|
||||||
@@ -321,6 +338,7 @@ def flash_decode_with_gqa_share_sparse(
|
|||||||
q_scale: Optional[float] = None,
|
q_scale: Optional[float] = None,
|
||||||
k_scale: Optional[float] = None,
|
k_scale: Optional[float] = None,
|
||||||
v_scale: Optional[float] = None,
|
v_scale: Optional[float] = None,
|
||||||
|
hisparse_slots: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
triton.set_allocator(robust_allocator)
|
triton.set_allocator(robust_allocator)
|
||||||
is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="decode")
|
is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="decode")
|
||||||
@@ -384,6 +402,7 @@ def flash_decode_with_gqa_share_sparse(
|
|||||||
topk_idx,
|
topk_idx,
|
||||||
o_partial,
|
o_partial,
|
||||||
lse_partial,
|
lse_partial,
|
||||||
|
hisparse_slots,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
slot_ids,
|
slot_ids,
|
||||||
max_slots,
|
max_slots,
|
||||||
@@ -392,6 +411,8 @@ def flash_decode_with_gqa_share_sparse(
|
|||||||
head_dim,
|
head_dim,
|
||||||
max_topk,
|
max_topk,
|
||||||
max_kv_len,
|
max_kv_len,
|
||||||
|
hisparse_slots.stride(0) if hisparse_slots is not None else 0,
|
||||||
|
hisparse_slots.stride(1) if hisparse_slots is not None else 0,
|
||||||
sm_scale,
|
sm_scale,
|
||||||
k_scale,
|
k_scale,
|
||||||
v_scale,
|
v_scale,
|
||||||
|
|||||||
@@ -28,6 +28,7 @@ from ..common.utils import (
|
|||||||
"BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["max_topk"]),
|
"BLOCK_SIZE_T": lambda args: triton.next_power_of_2(args["max_topk"]),
|
||||||
"BLOCK_SIZE_QH": lambda args: args["BLOCK_SIZE_Q"] * args["BLOCK_SIZE_H"],
|
"BLOCK_SIZE_QH": lambda args: args["BLOCK_SIZE_Q"] * args["BLOCK_SIZE_H"],
|
||||||
"HAS_SINK": lambda args: args["sink_ptr"] is not None,
|
"HAS_SINK": lambda args: args["sink_ptr"] is not None,
|
||||||
|
"HAS_LOC_MAPPING": lambda args: args["loc_mapping_ptr"] is not None,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
@triton.autotune(
|
@triton.autotune(
|
||||||
@@ -55,6 +56,7 @@ def _gqa_share_sparse_fwd_kernel(
|
|||||||
t_ptr, # topk_idx: kh x n x k
|
t_ptr, # topk_idx: kh x n x k
|
||||||
o_ptr, # O: n x h x d
|
o_ptr, # O: n x h x d
|
||||||
req_to_token_ptr, # req_to_token: max_reqs x max_kv_len
|
req_to_token_ptr, # req_to_token: max_reqs x max_kv_len
|
||||||
|
loc_mapping_ptr, # logical slot to HiSparse device slot
|
||||||
# seqlens
|
# seqlens
|
||||||
cu_seqlens_q,
|
cu_seqlens_q,
|
||||||
cu_seqblocks_q,
|
cu_seqblocks_q,
|
||||||
@@ -106,6 +108,7 @@ def _gqa_share_sparse_fwd_kernel(
|
|||||||
HAS_SINK: tl.constexpr,
|
HAS_SINK: tl.constexpr,
|
||||||
USE_TMA: tl.constexpr,
|
USE_TMA: tl.constexpr,
|
||||||
IS_FP8: tl.constexpr,
|
IS_FP8: tl.constexpr,
|
||||||
|
HAS_LOC_MAPPING: tl.constexpr,
|
||||||
):
|
):
|
||||||
sm_scale_log2e = sm_scale * 1.4426950409
|
sm_scale_log2e = sm_scale * 1.4426950409
|
||||||
# get batch id and head id
|
# get batch id and head id
|
||||||
@@ -199,6 +202,12 @@ def _gqa_share_sparse_fwd_kernel(
|
|||||||
mask=pos_mask,
|
mask=pos_mask,
|
||||||
other=0,
|
other=0,
|
||||||
).to(tl.int64)
|
).to(tl.int64)
|
||||||
|
if HAS_LOC_MAPPING:
|
||||||
|
slots = tl.load(
|
||||||
|
loc_mapping_ptr + slots,
|
||||||
|
mask=pos_mask,
|
||||||
|
other=0,
|
||||||
|
).to(tl.int64)
|
||||||
slots = (slots + max_slots) % max_slots # safety against negative
|
slots = (slots + max_slots) % max_slots # safety against negative
|
||||||
# k shape: [BLOCK_SIZE_KD, BLOCK_SIZE_K] (transposed for tl.dot)
|
# k shape: [BLOCK_SIZE_KD, BLOCK_SIZE_K] (transposed for tl.dot)
|
||||||
k = tl.load(
|
k = tl.load(
|
||||||
@@ -289,6 +298,7 @@ def flash_prefill_with_gqa_share_sparse(
|
|||||||
q_scale: Optional[float] = None,
|
q_scale: Optional[float] = None,
|
||||||
k_scale: Optional[float] = None,
|
k_scale: Optional[float] = None,
|
||||||
v_scale: Optional[float] = None,
|
v_scale: Optional[float] = None,
|
||||||
|
loc_mapping: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
triton.set_allocator(robust_allocator)
|
triton.set_allocator(robust_allocator)
|
||||||
is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="prefill")
|
is_fp8 = check_sparse_kv_fp8(q, k_cache, v_cache, label="prefill")
|
||||||
@@ -340,6 +350,7 @@ def flash_prefill_with_gqa_share_sparse(
|
|||||||
topk_idx,
|
topk_idx,
|
||||||
o,
|
o,
|
||||||
req_to_token,
|
req_to_token,
|
||||||
|
loc_mapping,
|
||||||
cu_seqlens,
|
cu_seqlens,
|
||||||
cu_seqblocks_q,
|
cu_seqblocks_q,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
|
|||||||
@@ -174,6 +174,8 @@ def _jit_sparse_module(
|
|||||||
hot_buffer_size: int,
|
hot_buffer_size: int,
|
||||||
is_mla: bool = False,
|
is_mla: bool = False,
|
||||||
is_dsv4_layout: bool = False,
|
is_dsv4_layout: bool = False,
|
||||||
|
top_k_block_size: int = 1,
|
||||||
|
top_k_is_blocks: bool = False,
|
||||||
record_miss_plan: bool = False,
|
record_miss_plan: bool = False,
|
||||||
skip_io: bool = False,
|
skip_io: bool = False,
|
||||||
) -> Module:
|
) -> Module:
|
||||||
@@ -185,6 +187,8 @@ def _jit_sparse_module(
|
|||||||
hot_buffer_size,
|
hot_buffer_size,
|
||||||
is_mla,
|
is_mla,
|
||||||
is_dsv4_layout,
|
is_dsv4_layout,
|
||||||
|
top_k_block_size,
|
||||||
|
top_k_is_blocks,
|
||||||
record_miss_plan,
|
record_miss_plan,
|
||||||
skip_io,
|
skip_io,
|
||||||
)
|
)
|
||||||
@@ -195,6 +199,8 @@ def _jit_sparse_module(
|
|||||||
hot_buffer_size,
|
hot_buffer_size,
|
||||||
is_mla,
|
is_mla,
|
||||||
is_dsv4_layout,
|
is_dsv4_layout,
|
||||||
|
top_k_block_size,
|
||||||
|
top_k_is_blocks,
|
||||||
record_miss_plan,
|
record_miss_plan,
|
||||||
skip_io,
|
skip_io,
|
||||||
)
|
)
|
||||||
@@ -308,7 +314,7 @@ def _load_cache_to_device_buffer_mla(
|
|||||||
skip_io=skip_io,
|
skip_io=skip_io,
|
||||||
)
|
)
|
||||||
|
|
||||||
empty = torch.empty(0)
|
empty = torch.empty(0, device=top_k_tokens.device)
|
||||||
|
|
||||||
if num_real_reqs is None:
|
if num_real_reqs is None:
|
||||||
num_real_reqs = torch.tensor(
|
num_real_reqs = torch.tensor(
|
||||||
@@ -399,6 +405,83 @@ def load_cache_to_device_buffer_mla(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def load_blocks_to_device_buffer_mha(
|
||||||
|
top_k_blocks: torch.Tensor,
|
||||||
|
device_buffer_tokens: torch.Tensor,
|
||||||
|
host_cache_locs: torch.Tensor,
|
||||||
|
device_buffer_locs: torch.Tensor,
|
||||||
|
host_cache_k: torch.Tensor,
|
||||||
|
host_cache_v: torch.Tensor,
|
||||||
|
device_buffer_k: torch.Tensor,
|
||||||
|
device_buffer_v: torch.Tensor,
|
||||||
|
top_k_device_locs: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
lru_slots: torch.Tensor,
|
||||||
|
item_size_bytes: int,
|
||||||
|
hot_buffer_size: int,
|
||||||
|
sparse_block_size: int,
|
||||||
|
page_size: int = 1,
|
||||||
|
block_size: int = 256,
|
||||||
|
num_real_reqs: torch.Tensor | None = None,
|
||||||
|
skip_io: bool = False,
|
||||||
|
) -> None:
|
||||||
|
"""Swap block-selected MHA K/V into the HiSparse device pool."""
|
||||||
|
num_top_k_blocks = top_k_blocks.size(1)
|
||||||
|
num_top_k_tokens = num_top_k_blocks * sparse_block_size
|
||||||
|
assert hot_buffer_size >= num_top_k_tokens, (
|
||||||
|
f"hot_buffer_size ({hot_buffer_size}) must be >= selected tokens "
|
||||||
|
f"({num_top_k_tokens})"
|
||||||
|
)
|
||||||
|
assert top_k_device_locs.size(1) >= num_top_k_tokens
|
||||||
|
k_stride = host_cache_k.stride(0) * host_cache_k.element_size()
|
||||||
|
v_stride = host_cache_v.stride(0) * host_cache_v.element_size()
|
||||||
|
assert k_stride == v_stride == item_size_bytes, (
|
||||||
|
"K/V token strides must equal item_size_bytes: "
|
||||||
|
f"k_stride={k_stride}, v_stride={v_stride}, "
|
||||||
|
f"item_size_bytes={item_size_bytes}"
|
||||||
|
)
|
||||||
|
|
||||||
|
module = _jit_sparse_module(
|
||||||
|
item_size_bytes,
|
||||||
|
block_size,
|
||||||
|
num_top_k_blocks,
|
||||||
|
hot_buffer_size,
|
||||||
|
is_mla=False,
|
||||||
|
is_dsv4_layout=False,
|
||||||
|
top_k_block_size=sparse_block_size,
|
||||||
|
top_k_is_blocks=True,
|
||||||
|
record_miss_plan=False,
|
||||||
|
skip_io=skip_io,
|
||||||
|
)
|
||||||
|
empty = torch.empty(0, device=top_k_blocks.device)
|
||||||
|
if num_real_reqs is None:
|
||||||
|
num_real_reqs = torch.tensor(
|
||||||
|
[top_k_blocks.size(0)], dtype=torch.int32, device=top_k_blocks.device
|
||||||
|
)
|
||||||
|
|
||||||
|
module.load_cache_to_device_buffer(
|
||||||
|
top_k_blocks,
|
||||||
|
device_buffer_tokens,
|
||||||
|
host_cache_locs,
|
||||||
|
device_buffer_locs,
|
||||||
|
host_cache_k,
|
||||||
|
host_cache_v,
|
||||||
|
device_buffer_k,
|
||||||
|
device_buffer_v,
|
||||||
|
top_k_device_locs,
|
||||||
|
req_pool_indices,
|
||||||
|
seq_lens,
|
||||||
|
lru_slots,
|
||||||
|
num_real_reqs,
|
||||||
|
page_size,
|
||||||
|
item_size_bytes,
|
||||||
|
empty,
|
||||||
|
empty,
|
||||||
|
empty,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def copy_cache_planned_mla(
|
def copy_cache_planned_mla(
|
||||||
*,
|
*,
|
||||||
miss_src: torch.Tensor,
|
miss_src: torch.Tensor,
|
||||||
|
|||||||
@@ -86,14 +86,16 @@ def validate_hisparse(server_args: ServerArgs) -> None:
|
|||||||
from sglang.srt.configs.model_config import (
|
from sglang.srt.configs.model_config import (
|
||||||
is_deepseek_dsa,
|
is_deepseek_dsa,
|
||||||
is_deepseek_v4,
|
is_deepseek_v4,
|
||||||
|
is_minimax_sparse,
|
||||||
)
|
)
|
||||||
|
|
||||||
hf_config = model_config_of(server_args).hf_config
|
hf_config = model_config_of(server_args).hf_config
|
||||||
is_v4_hisparse = is_deepseek_v4(hf_config)
|
is_v4_hisparse = is_deepseek_v4(hf_config)
|
||||||
|
is_m3_hisparse = is_minimax_sparse(hf_config)
|
||||||
is_hip = get_platform().is_hip
|
is_hip = get_platform().is_hip
|
||||||
assert is_deepseek_dsa(hf_config) or is_v4_hisparse, (
|
assert is_deepseek_dsa(hf_config) or is_v4_hisparse or is_m3_hisparse, (
|
||||||
"--enable-hisparse is only supported for DSA (DeepSeek Sparse Attention) "
|
"--enable-hisparse is only supported for DSA (DeepSeek Sparse Attention) "
|
||||||
"models (e.g., DeepSeek V3.2, GLM-5) and DeepSeek V4 now. "
|
"models (e.g., DeepSeek V3.2, GLM-5), DeepSeek V4, and MiniMax M3 now. "
|
||||||
)
|
)
|
||||||
|
|
||||||
assert cfg.disable_radix_cache, (
|
assert cfg.disable_radix_cache, (
|
||||||
@@ -121,6 +123,10 @@ def validate_hisparse(server_args: ServerArgs) -> None:
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# MiniMax M3 uses its own Triton sparse kernels.
|
||||||
|
if is_m3_hisparse:
|
||||||
|
return
|
||||||
|
|
||||||
if resolved_view(server_args).kv_cache_dtype not in (
|
if resolved_view(server_args).kv_cache_dtype not in (
|
||||||
"bfloat16",
|
"bfloat16",
|
||||||
"auto",
|
"auto",
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ class StateType(str, enum.Enum):
|
|||||||
# only the live subrange of that row for the current open pool.
|
# only the live subrange of that row for the current open pool.
|
||||||
DSA_TAIL = "dsa_tail"
|
DSA_TAIL = "dsa_tail"
|
||||||
MINIMAX_INDEX_K = "minimax_index_k"
|
MINIMAX_INDEX_K = "minimax_index_k"
|
||||||
|
MINIMAX_DENSE_KV = "minimax_dense_kv"
|
||||||
# DeepSeek-V4 unified_kv SWA ring: addressed per-row by ring slot
|
# DeepSeek-V4 unified_kv SWA ring: addressed per-row by ring slot
|
||||||
# (req_pool_idx * ring_stride + pos % ring_stride), needs its own component.
|
# (req_pool_idx * ring_stride + pos % ring_stride), needs its own component.
|
||||||
SWA_RING = "swa_ring"
|
SWA_RING = "swa_ring"
|
||||||
|
|||||||
@@ -62,6 +62,7 @@ from sglang.srt.disaggregation.utils import (
|
|||||||
build_staging_slot_metadata,
|
build_staging_slot_metadata,
|
||||||
get_dsa_tail_state_indices,
|
get_dsa_tail_state_indices,
|
||||||
get_kv_class,
|
get_kv_class,
|
||||||
|
get_kv_transfer_buf_infos,
|
||||||
get_qsa_pending_state_indices,
|
get_qsa_pending_state_indices,
|
||||||
is_mla_backend,
|
is_mla_backend,
|
||||||
is_unadmitted_reject,
|
is_unadmitted_reject,
|
||||||
@@ -575,8 +576,8 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
if self.scheduler.enable_hisparse
|
if self.scheduler.enable_hisparse
|
||||||
else self.token_to_kv_pool
|
else self.token_to_kv_pool
|
||||||
)
|
)
|
||||||
kv_data_ptrs, kv_data_lens, kv_item_lens = (
|
kv_data_ptrs, kv_data_lens, kv_item_lens = get_kv_transfer_buf_infos(
|
||||||
transfer_kv_pool.get_contiguous_buf_infos()
|
transfer_kv_pool
|
||||||
)
|
)
|
||||||
kv_data_mem_kinds = (
|
kv_data_mem_kinds = (
|
||||||
["DRAM"] * len(kv_data_ptrs)
|
["DRAM"] * len(kv_data_ptrs)
|
||||||
@@ -1579,6 +1580,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
StateType.DSA: _full_kv_pages_payload,
|
StateType.DSA: _full_kv_pages_payload,
|
||||||
StateType.DSA_TAIL: _dsa_tail_payload,
|
StateType.DSA_TAIL: _dsa_tail_payload,
|
||||||
StateType.MINIMAX_INDEX_K: _full_kv_pages_payload,
|
StateType.MINIMAX_INDEX_K: _full_kv_pages_payload,
|
||||||
|
StateType.MINIMAX_DENSE_KV: _full_kv_pages_payload,
|
||||||
StateType.SWA_RING: _swa_ring_payload,
|
StateType.SWA_RING: _swa_ring_payload,
|
||||||
StateType.DSV4_REQUEST_STATE: _request_state_payload,
|
StateType.DSV4_REQUEST_STATE: _request_state_payload,
|
||||||
StateType.BLOCK_SCALE: _full_kv_pages_payload,
|
StateType.BLOCK_SCALE: _full_kv_pages_payload,
|
||||||
|
|||||||
@@ -1398,7 +1398,9 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
if src_layer_ids or dst_layer_ids:
|
if src_layer_ids or dst_layer_ids:
|
||||||
# Draft buffers break the flat [K block, V block] layout, so pair by
|
# Draft buffers break the flat [K block, V block] layout, so pair by
|
||||||
# layer ID instead of the half-split used by get_mha_kv_ptrs_with_pp.
|
# layer ID instead of the half-split used by get_mha_kv_ptrs_with_pp.
|
||||||
if any(l != src_kv_item_len for l in self.kv_args.kv_item_lens):
|
if any(
|
||||||
|
item_len != src_kv_item_len for item_len in self.kv_args.kv_item_lens
|
||||||
|
):
|
||||||
logger.error(
|
logger.error(
|
||||||
f"[{mooncake_session_id}] head-sliced transfer assumes one item "
|
f"[{mooncake_session_id}] head-sliced transfer assumes one item "
|
||||||
f"length for every KV entry, got {set(self.kv_args.kv_item_lens)}"
|
f"length for every KV entry, got {set(self.kv_args.kv_item_lens)}"
|
||||||
@@ -1852,12 +1854,11 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
)
|
)
|
||||||
or rc
|
or rc
|
||||||
)
|
)
|
||||||
elif st == StateType.MINIMAX_INDEX_K:
|
elif st in (StateType.MINIMAX_INDEX_K, StateType.MINIMAX_DENSE_KV):
|
||||||
# Equal-TP / PP=1 only. Sub-pools are compacted sparse-layer
|
# Compacted layer lists require equal TP and PP=1 on both peers.
|
||||||
# lists, so PP>1 mis-slices and heterogeneous TP is unsupported.
|
|
||||||
if self.pp_size is not None and self.pp_size > 1:
|
if self.pp_size is not None and self.pp_size > 1:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"PD disagg: PP>1 not supported for MiniMax sparse index yet."
|
"PD disagg: PP>1 not supported for MiniMax state yet."
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
target_rank_registration_info is not None
|
target_rank_registration_info is not None
|
||||||
@@ -1866,11 +1867,17 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
):
|
):
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"PD disagg: heterogeneous TP not supported for MiniMax "
|
"PD disagg: heterogeneous TP not supported for MiniMax "
|
||||||
"sparse index yet."
|
"state yet."
|
||||||
)
|
)
|
||||||
src_indices = list(indices)
|
src_indices = list(indices)
|
||||||
dst_indices_local = list(dst_indices)
|
dst_indices_local = list(dst_indices)
|
||||||
if len(src_indices) > len(dst_indices_local):
|
if st == StateType.MINIMAX_DENSE_KV:
|
||||||
|
if len(src_indices) != len(dst_indices_local):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"{st.value} state index length mismatch: "
|
||||||
|
f"prefill={len(src_indices)}, dst={len(dst_indices_local)}"
|
||||||
|
)
|
||||||
|
elif len(src_indices) > len(dst_indices_local):
|
||||||
src_indices = src_indices[: len(dst_indices_local)]
|
src_indices = src_indices[: len(dst_indices_local)]
|
||||||
elif len(src_indices) < len(dst_indices_local):
|
elif len(src_indices) < len(dst_indices_local):
|
||||||
dst_indices_local = dst_indices_local[: len(src_indices)]
|
dst_indices_local = dst_indices_local[: len(src_indices)]
|
||||||
|
|||||||
@@ -1282,6 +1282,7 @@ class MoriKVManager(CommonKVManager):
|
|||||||
"swa_ring",
|
"swa_ring",
|
||||||
"c128_state",
|
"c128_state",
|
||||||
"minimax_index_k",
|
"minimax_index_k",
|
||||||
|
"minimax_dense_kv",
|
||||||
):
|
):
|
||||||
statuses.extend(
|
statuses.extend(
|
||||||
self._send_swa_dsa_state(
|
self._send_swa_dsa_state(
|
||||||
@@ -1409,7 +1410,12 @@ class MoriKVManager(CommonKVManager):
|
|||||||
f"PD state transfer does not support TP-mismatched non-MLA SWA models "
|
f"PD state transfer does not support TP-mismatched non-MLA SWA models "
|
||||||
f"(prefill_tp_size={self.attn_tp_size}, decode_tp_size={peer_info.decode_tp_size})"
|
f"(prefill_tp_size={self.attn_tp_size}, decode_tp_size={peer_info.decode_tp_size})"
|
||||||
)
|
)
|
||||||
if state_type in ("qsa_pending", "qsa_compressed", "minimax_index_k"):
|
if state_type in (
|
||||||
|
"qsa_pending",
|
||||||
|
"qsa_compressed",
|
||||||
|
"minimax_index_k",
|
||||||
|
"minimax_dense_kv",
|
||||||
|
):
|
||||||
if self.pp_size is not None and self.pp_size > 1:
|
if self.pp_size is not None and self.pp_size > 1:
|
||||||
# MORI registration does not exchange state_layer_ids. Compact
|
# MORI registration does not exchange state_layer_ids. Compact
|
||||||
# sparse-state lists therefore cannot be paired safely across
|
# sparse-state lists therefore cannot be paired safely across
|
||||||
@@ -1445,6 +1451,7 @@ class MoriKVManager(CommonKVManager):
|
|||||||
"qsa_compressed",
|
"qsa_compressed",
|
||||||
"swa_ring",
|
"swa_ring",
|
||||||
"c128_state",
|
"c128_state",
|
||||||
|
"minimax_dense_kv",
|
||||||
):
|
):
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"{state_type.upper()} state index length mismatch: "
|
f"{state_type.upper()} state index length mismatch: "
|
||||||
|
|||||||
@@ -2696,17 +2696,16 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
|
|||||||
dst_layer_ids=dst_lids,
|
dst_layer_ids=dst_lids,
|
||||||
dst_item_lens=dst_lens,
|
dst_item_lens=dst_lens,
|
||||||
)
|
)
|
||||||
elif st == StateType.MINIMAX_INDEX_K:
|
elif st in (StateType.MINIMAX_INDEX_K, StateType.MINIMAX_DENSE_KV):
|
||||||
# Equal-TP / PP=1 only. Sub-pools are compacted sparse-layer
|
# Compacted layer lists require equal TP and PP=1 on both peers.
|
||||||
# lists, so PP>1 mis-slices and heterogeneous TP is unsupported.
|
|
||||||
if self.pp_size is not None and self.pp_size > 1:
|
if self.pp_size is not None and self.pp_size > 1:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"PD disagg: PP>1 not supported for MiniMax sparse index yet."
|
"PD disagg: PP>1 not supported for MiniMax state yet."
|
||||||
)
|
)
|
||||||
if self.attn_tp_size != decode_tp_size:
|
if self.attn_tp_size != decode_tp_size:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"PD disagg: heterogeneous TP not supported for MiniMax "
|
"PD disagg: heterogeneous TP not supported for MiniMax "
|
||||||
"sparse index yet."
|
"state yet."
|
||||||
)
|
)
|
||||||
if len(src_indices) != len(dst_indices):
|
if len(src_indices) != len(dst_indices):
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
|
|||||||
@@ -53,6 +53,7 @@ from sglang.srt.disaggregation.utils import (
|
|||||||
build_staging_slot_metadata,
|
build_staging_slot_metadata,
|
||||||
get_dsa_tail_state_indices,
|
get_dsa_tail_state_indices,
|
||||||
get_kv_class,
|
get_kv_class,
|
||||||
|
get_kv_transfer_buf_infos,
|
||||||
get_qsa_pending_state_indices,
|
get_qsa_pending_state_indices,
|
||||||
is_aborted,
|
is_aborted,
|
||||||
is_mla_backend,
|
is_mla_backend,
|
||||||
@@ -256,8 +257,8 @@ class PrefillBootstrapQueue:
|
|||||||
hf_text_config=self.scheduler.model_config.hf_text_config,
|
hf_text_config=self.scheduler.model_config.hf_text_config,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
kv_data_ptrs, kv_data_lens, kv_item_lens = (
|
kv_data_ptrs, kv_data_lens, kv_item_lens = get_kv_transfer_buf_infos(
|
||||||
self.token_to_kv_pool.get_contiguous_buf_infos()
|
self.token_to_kv_pool
|
||||||
)
|
)
|
||||||
kv_args.prefill_end_layer = (
|
kv_args.prefill_end_layer = (
|
||||||
kv_args.prefill_start_layer + len(kv_data_ptrs)
|
kv_args.prefill_start_layer + len(kv_data_ptrs)
|
||||||
@@ -1424,6 +1425,7 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
StateType.DSA: _full_kv_pages_payload,
|
StateType.DSA: _full_kv_pages_payload,
|
||||||
StateType.DSA_TAIL: _dsa_tail_payload,
|
StateType.DSA_TAIL: _dsa_tail_payload,
|
||||||
StateType.MINIMAX_INDEX_K: _full_kv_pages_payload,
|
StateType.MINIMAX_INDEX_K: _full_kv_pages_payload,
|
||||||
|
StateType.MINIMAX_DENSE_KV: _full_kv_pages_payload,
|
||||||
StateType.SWA_RING: _swa_ring_payload,
|
StateType.SWA_RING: _swa_ring_payload,
|
||||||
StateType.DSV4_REQUEST_STATE: _request_state_payload,
|
StateType.DSV4_REQUEST_STATE: _request_state_payload,
|
||||||
StateType.BLOCK_SCALE: _full_kv_pages_payload,
|
StateType.BLOCK_SCALE: _full_kv_pages_payload,
|
||||||
|
|||||||
@@ -1310,6 +1310,14 @@ def build_dsa_tail_transfer_blocks(
|
|||||||
return transfer_blocks
|
return transfer_blocks
|
||||||
|
|
||||||
|
|
||||||
|
def get_kv_transfer_buf_infos(pool):
|
||||||
|
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
|
||||||
|
|
||||||
|
if isinstance(pool, MiniMaxSparseKVPool):
|
||||||
|
return pool.get_sparse_kv_buf_infos()
|
||||||
|
return pool.get_contiguous_buf_infos()
|
||||||
|
|
||||||
|
|
||||||
def setup_state_kv_args(
|
def setup_state_kv_args(
|
||||||
kv_args: KVArgs,
|
kv_args: KVArgs,
|
||||||
token_to_kv_pool,
|
token_to_kv_pool,
|
||||||
@@ -1375,6 +1383,11 @@ def setup_state_kv_args(
|
|||||||
if token_to_kv_pool.index_k_pool is not None:
|
if token_to_kv_pool.index_k_pool is not None:
|
||||||
dp, dl, il = token_to_kv_pool.get_index_k_state_buf_infos()
|
dp, dl, il = token_to_kv_pool.get_index_k_state_buf_infos()
|
||||||
append_state_component(kv_args, StateType.MINIMAX_INDEX_K, dp, dl, il)
|
append_state_component(kv_args, StateType.MINIMAX_INDEX_K, dp, dl, il)
|
||||||
|
append_state_component(
|
||||||
|
kv_args,
|
||||||
|
StateType.MINIMAX_DENSE_KV,
|
||||||
|
*token_to_kv_pool.get_dense_kv_state_buf_infos(),
|
||||||
|
)
|
||||||
elif hasattr(token_to_kv_pool, "get_state_buf_infos"):
|
elif hasattr(token_to_kv_pool, "get_state_buf_infos"):
|
||||||
data_ptrs, data_lens, item_lens = token_to_kv_pool.get_state_buf_infos()
|
data_ptrs, data_lens, item_lens = token_to_kv_pool.get_state_buf_infos()
|
||||||
|
|
||||||
|
|||||||
@@ -116,6 +116,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
assert isinstance(runner.token_to_kv_pool, MiniMaxSparseKVPool)
|
assert isinstance(runner.token_to_kv_pool, MiniMaxSparseKVPool)
|
||||||
self.is_npu = is_npu()
|
self.is_npu = is_npu()
|
||||||
self.kv_pool = runner.token_to_kv_pool
|
self.kv_pool = runner.token_to_kv_pool
|
||||||
|
self.hisparse_coordinator = runner.hisparse_coordinator
|
||||||
self.token_to_kv_pool = runner.token_to_kv_pool # alias for TboAttnBackend
|
self.token_to_kv_pool = runner.token_to_kv_pool # alias for TboAttnBackend
|
||||||
self.req_to_token_pool = runner.req_to_token_pool # pool obj for TboAttnBackend
|
self.req_to_token_pool = runner.req_to_token_pool # pool obj for TboAttnBackend
|
||||||
self.req_to_token = runner.req_to_token_pool.req_to_token
|
self.req_to_token = runner.req_to_token_pool.req_to_token
|
||||||
@@ -176,6 +177,18 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
local_tokens + self.block_size_k - 1
|
local_tokens + self.block_size_k - 1
|
||||||
) // self.block_size_k + 1
|
) // self.block_size_k + 1
|
||||||
self.topk_blocks = sparse_cfg["sparse_topk_blocks"]
|
self.topk_blocks = sparse_cfg["sparse_topk_blocks"]
|
||||||
|
if self.hisparse_coordinator is not None:
|
||||||
|
selected_tokens = self.topk_blocks * self.block_size_k
|
||||||
|
assert selected_tokens <= self.hisparse_coordinator.device_buffer_size, (
|
||||||
|
f"MiniMax M3 selects {selected_tokens} sparse-attention tokens, "
|
||||||
|
"but the HiSparse device buffer holds only "
|
||||||
|
f"{self.hisparse_coordinator.device_buffer_size}."
|
||||||
|
)
|
||||||
|
self._loc_mapping = (
|
||||||
|
self.kv_pool.main_pool.full_to_hisparse_device_index_mapping
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self._loc_mapping = None
|
||||||
|
|
||||||
# MSA (fmha_sm100) is SM100-only; fall back to the Triton sparse path when
|
# MSA (fmha_sm100) is SM100-only; fall back to the Triton sparse path when
|
||||||
# the kernel is unavailable or its constraints don't hold.
|
# the kernel is unavailable or its constraints don't hold.
|
||||||
@@ -209,6 +222,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
self.use_msa = (
|
self.use_msa = (
|
||||||
not envs.SGLANG_DISABLE_MSA.get()
|
not envs.SGLANG_DISABLE_MSA.get()
|
||||||
|
and self.hisparse_coordinator is None
|
||||||
and msa_available()
|
and msa_available()
|
||||||
and self.block_size_k == 128
|
and self.block_size_k == 128
|
||||||
and self.kv_pool.page_size == self.block_size_k
|
and self.kv_pool.page_size == self.block_size_k
|
||||||
@@ -245,6 +259,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
self.page_size = self.kv_pool.page_size
|
self.page_size = self.kv_pool.page_size
|
||||||
self.use_dense_sparse_decode = (
|
self.use_dense_sparse_decode = (
|
||||||
(not self.is_npu)
|
(not self.is_npu)
|
||||||
|
and self.hisparse_coordinator is None
|
||||||
and envs.SGLANG_OPT_USE_MINIMAX_DENSE_SPARSE_DECODE.get()
|
and envs.SGLANG_OPT_USE_MINIMAX_DENSE_SPARSE_DECODE.get()
|
||||||
and self.block_size_k % self.page_size == 0
|
and self.block_size_k % self.page_size == 0
|
||||||
# _dense_sparse_main_decode calls trtllm decode with a bf16 q and
|
# _dense_sparse_main_decode calls trtllm decode with a bf16 q and
|
||||||
@@ -326,6 +341,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
f"msa_owns_decode={self._msa_owns_decode}, "
|
f"msa_owns_decode={self._msa_owns_decode}, "
|
||||||
f"decode_cuda_graph={_decode_cuda_graph}, "
|
f"decode_cuda_graph={_decode_cuda_graph}, "
|
||||||
f"fp8_attn_gemm={self.fp8_attn_gemm}, "
|
f"fp8_attn_gemm={self.fp8_attn_gemm}, "
|
||||||
|
f"hisparse={'enabled' if self._loc_mapping is not None else 'disabled'}, "
|
||||||
f"npu_native_attn={'on' if (self._native_sparse_ok and _native_attn_enabled()) else 'off'}, "
|
f"npu_native_attn={'on' if (self._native_sparse_ok and _native_attn_enabled()) else 'off'}, "
|
||||||
f"disable_value_layers={sorted(self.disable_value_layer_ids)})"
|
f"disable_value_layers={sorted(self.disable_value_layer_ids)})"
|
||||||
)
|
)
|
||||||
@@ -336,6 +352,22 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
"take minutes; compiles serialize across TP ranks)."
|
"take minutes; compiles serialize across TP ranks)."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _hisparse_swap_in_blocks(
|
||||||
|
self,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
topk_idx: torch.Tensor,
|
||||||
|
layer_id: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
assert topk_idx.size(0) == 1
|
||||||
|
top_k_device_locs = self.hisparse_coordinator.swap_in_selected_blocks(
|
||||||
|
req_pool_indices=forward_batch.req_pool_indices,
|
||||||
|
seq_lens=forward_batch.seq_lens,
|
||||||
|
top_k_blocks=topk_idx[0],
|
||||||
|
layer_id=layer_id,
|
||||||
|
sparse_block_size=self.block_size_k,
|
||||||
|
)
|
||||||
|
return top_k_device_locs.unsqueeze(0)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _choose_decode_score_max_chunks(batch_size: int) -> int:
|
def _choose_decode_score_max_chunks(batch_size: int) -> int:
|
||||||
"""Score chunk count per graph bucket.
|
"""Score chunk count per graph bucket.
|
||||||
@@ -1549,6 +1581,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
idx_v_scale=layer.idx_v_scale_float,
|
idx_v_scale=layer.idx_v_scale_float,
|
||||||
cached_topk_idx=cached_topk_idx,
|
cached_topk_idx=cached_topk_idx,
|
||||||
return_topk_idx=want_topk,
|
return_topk_idx=want_topk,
|
||||||
|
loc_mapping=self._loc_mapping,
|
||||||
)
|
)
|
||||||
if want_topk:
|
if want_topk:
|
||||||
idx_o, o, reduced_topk_idx = result
|
idx_o, o, reduced_topk_idx = result
|
||||||
@@ -1702,6 +1735,16 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
else:
|
else:
|
||||||
_cached_topk = _topk_buf
|
_cached_topk = _topk_buf
|
||||||
|
|
||||||
|
hisparse_swap_in_fn = None
|
||||||
|
if self.hisparse_coordinator is not None:
|
||||||
|
|
||||||
|
def hisparse_swap_in_fn(topk_idx):
|
||||||
|
return self._hisparse_swap_in_blocks(
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
topk_idx=topk_idx,
|
||||||
|
layer_id=layer.layer_id,
|
||||||
|
)
|
||||||
|
|
||||||
idx_o, o = minimax_sparse_decode(
|
idx_o, o = minimax_sparse_decode(
|
||||||
q,
|
q,
|
||||||
None,
|
None,
|
||||||
@@ -1735,6 +1778,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
|
|||||||
idx_v_scale=layer.idx_v_scale_float,
|
idx_v_scale=layer.idx_v_scale_float,
|
||||||
cached_topk_idx=_cached_topk,
|
cached_topk_idx=_cached_topk,
|
||||||
topk_out=_topk_buf if _want_topk else None,
|
topk_out=_topk_buf if _want_topk else None,
|
||||||
|
hisparse_swap_in_fn=hisparse_swap_in_fn,
|
||||||
)
|
)
|
||||||
return (
|
return (
|
||||||
None if idx_o is None else idx_o.reshape(q.shape[0], -1).contiguous(),
|
None if idx_o is None else idx_o.reshape(q.shape[0], -1).contiguous(),
|
||||||
|
|||||||
@@ -75,6 +75,7 @@ def minimax_sparse_prefill(
|
|||||||
idx_v_scale: Optional[float] = None,
|
idx_v_scale: Optional[float] = None,
|
||||||
cached_topk_idx: Optional[torch.Tensor] = None,
|
cached_topk_idx: Optional[torch.Tensor] = None,
|
||||||
return_topk_idx: bool = False,
|
return_topk_idx: bool = False,
|
||||||
|
loc_mapping: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
"""Run MiniMax-M3 sparse prefill.
|
"""Run MiniMax-M3 sparse prefill.
|
||||||
|
|
||||||
@@ -146,7 +147,7 @@ def minimax_sparse_prefill(
|
|||||||
# Step 3: Sparse attention using topk index (main head). The MSA path only
|
# Step 3: Sparse attention using topk index (main head). The MSA path only
|
||||||
# replaces this step; the indexer above is unchanged. MSA has no attn-sink
|
# replaces this step; the indexer above is unchanged. MSA has no attn-sink
|
||||||
# input, so keep the Triton path when sink is present.
|
# input, so keep the Triton path when sink is present.
|
||||||
if use_msa and sink is None:
|
if use_msa and sink is None and loc_mapping is None:
|
||||||
from .msa import MSAUnavailableError, msa_sparse_prefill_main
|
from .msa import MSAUnavailableError, msa_sparse_prefill_main
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -188,6 +189,7 @@ def minimax_sparse_prefill(
|
|||||||
q_scale=q_scale,
|
q_scale=q_scale,
|
||||||
k_scale=k_scale,
|
k_scale=k_scale,
|
||||||
v_scale=v_scale,
|
v_scale=v_scale,
|
||||||
|
loc_mapping=loc_mapping,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
o = flash_prefill_with_gqa_share_sparse(
|
o = flash_prefill_with_gqa_share_sparse(
|
||||||
@@ -210,6 +212,7 @@ def minimax_sparse_prefill(
|
|||||||
q_scale=q_scale,
|
q_scale=q_scale,
|
||||||
k_scale=k_scale,
|
k_scale=k_scale,
|
||||||
v_scale=v_scale,
|
v_scale=v_scale,
|
||||||
|
loc_mapping=loc_mapping,
|
||||||
)
|
)
|
||||||
if return_topk_idx:
|
if return_topk_idx:
|
||||||
return idx_o, o, reduced_topk_idx
|
return idx_o, o, reduced_topk_idx
|
||||||
@@ -255,6 +258,7 @@ def minimax_sparse_decode(
|
|||||||
idx_v_scale: Optional[float] = None,
|
idx_v_scale: Optional[float] = None,
|
||||||
cached_topk_idx: Optional[torch.Tensor] = None,
|
cached_topk_idx: Optional[torch.Tensor] = None,
|
||||||
topk_out: Optional[torch.Tensor] = None,
|
topk_out: Optional[torch.Tensor] = None,
|
||||||
|
hisparse_swap_in_fn: Optional[Callable] = None,
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
# Index top-k sharing for DECODE. A group's source layer passes ``topk_out``
|
# Index top-k sharing for DECODE. A group's source layer passes ``topk_out``
|
||||||
# (a persistent buffer) and publishes its reduced top-k there; the group's
|
# (a persistent buffer) and publishes its reduced top-k there; the group's
|
||||||
@@ -319,9 +323,12 @@ def minimax_sparse_decode(
|
|||||||
f"reduced top-k shape {tuple(topk_idx.shape)}"
|
f"reduced top-k shape {tuple(topk_idx.shape)}"
|
||||||
)
|
)
|
||||||
topk_out.copy_(topk_idx)
|
topk_out.copy_(topk_idx)
|
||||||
|
hisparse_slots = (
|
||||||
|
hisparse_swap_in_fn(topk_idx) if hisparse_swap_in_fn is not None else None
|
||||||
|
)
|
||||||
# Step 3: Sparse attention using topk index (main head). The MSA path
|
# Step 3: Sparse attention using topk index (main head). The MSA path
|
||||||
# only replaces this step; keep the Triton path when sink is present.
|
# only replaces this step; keep the Triton path when sink is present.
|
||||||
if use_msa and sink is None:
|
if use_msa and sink is None and hisparse_slots is None:
|
||||||
from .msa import MSAUnavailableError, msa_sparse_decode_main
|
from .msa import MSAUnavailableError, msa_sparse_decode_main
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -357,6 +364,7 @@ def minimax_sparse_decode(
|
|||||||
q_scale=q_scale,
|
q_scale=q_scale,
|
||||||
k_scale=k_scale,
|
k_scale=k_scale,
|
||||||
v_scale=v_scale,
|
v_scale=v_scale,
|
||||||
|
hisparse_slots=hisparse_slots,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
o = flash_decode_with_gqa_share_sparse(
|
o = flash_decode_with_gqa_share_sparse(
|
||||||
@@ -373,5 +381,6 @@ def minimax_sparse_decode(
|
|||||||
q_scale=q_scale,
|
q_scale=q_scale,
|
||||||
k_scale=k_scale,
|
k_scale=k_scale,
|
||||||
v_scale=v_scale,
|
v_scale=v_scale,
|
||||||
|
hisparse_slots=hisparse_slots,
|
||||||
)
|
)
|
||||||
return idx_o, o
|
return idx_o, o
|
||||||
|
|||||||
@@ -19,9 +19,16 @@ if is_xpu():
|
|||||||
"copy_cache_planned_mla has no AOT sgl_kernel implementation."
|
"copy_cache_planned_mla has no AOT sgl_kernel implementation."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def load_blocks_to_device_buffer_mha(*args, **kwargs):
|
||||||
|
raise RuntimeError(
|
||||||
|
"MiniMax M3 HiSparse block swap-in is unsupported on XPU: "
|
||||||
|
"load_blocks_to_device_buffer_mha has no AOT sgl_kernel implementation."
|
||||||
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
from sglang.kernels.ops.kvcache.hisparse import (
|
from sglang.kernels.ops.kvcache.hisparse import (
|
||||||
copy_cache_planned_mla,
|
copy_cache_planned_mla,
|
||||||
|
load_blocks_to_device_buffer_mha,
|
||||||
load_cache_to_device_buffer_dsv4_mla,
|
load_cache_to_device_buffer_dsv4_mla,
|
||||||
load_cache_to_device_buffer_mla,
|
load_cache_to_device_buffer_mla,
|
||||||
)
|
)
|
||||||
@@ -36,8 +43,9 @@ from sglang.srt.mem_cache.allocator.hisparse import (
|
|||||||
from sglang.srt.mem_cache.hisparse_memory_pool import (
|
from sglang.srt.mem_cache.hisparse_memory_pool import (
|
||||||
HiSparseDSATokenToKVPool,
|
HiSparseDSATokenToKVPool,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool, ReqToTokenPool
|
||||||
from sglang.srt.mem_cache.memory_pool_host import DeepSeekV4PagedHostPool
|
from sglang.srt.mem_cache.memory_pool_host import DeepSeekV4PagedHostPool
|
||||||
|
from sglang.srt.mem_cache.pool_host.mha import HiSparseMHATokenToKVPoolHost
|
||||||
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
|
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
|
||||||
|
|
||||||
device_module = get_device_module()
|
device_module = get_device_module()
|
||||||
@@ -157,9 +165,11 @@ class HiSparseCoordinator:
|
|||||||
)
|
)
|
||||||
self.compress_ratio = self.token_to_kv_pool_allocator.compress_ratio
|
self.compress_ratio = self.token_to_kv_pool_allocator.compress_ratio
|
||||||
|
|
||||||
|
kvcache = self.token_to_kv_pool_allocator.get_kvcache()
|
||||||
self.is_dsv4_hisparse = isinstance(
|
self.is_dsv4_hisparse = isinstance(
|
||||||
self.token_to_kv_pool_allocator, DeepSeekV4HiSparseTokenToKVPoolAllocator
|
self.token_to_kv_pool_allocator, DeepSeekV4HiSparseTokenToKVPoolAllocator
|
||||||
)
|
)
|
||||||
|
self.is_m3_hisparse = isinstance(kvcache, MiniMaxSparseKVPool)
|
||||||
if self.is_dsv4_hisparse:
|
if self.is_dsv4_hisparse:
|
||||||
self.mem_pool_device = self.token_to_kv_pool_allocator.hisparse_kvcache
|
self.mem_pool_device = self.token_to_kv_pool_allocator.hisparse_kvcache
|
||||||
page_size = self.mem_pool_device.page_size
|
page_size = self.mem_pool_device.page_size
|
||||||
@@ -184,18 +194,30 @@ class HiSparseCoordinator:
|
|||||||
assert isinstance(
|
assert isinstance(
|
||||||
self.token_to_kv_pool_allocator, HiSparseTokenToKVPoolAllocator
|
self.token_to_kv_pool_allocator, HiSparseTokenToKVPoolAllocator
|
||||||
)
|
)
|
||||||
self.mem_pool_device: HiSparseDSATokenToKVPool = (
|
if self.is_m3_hisparse:
|
||||||
self.token_to_kv_pool_allocator.get_kvcache()
|
self.mem_pool_device = kvcache.main_pool
|
||||||
)
|
assert self.mem_pool_device.head_num == 1, (
|
||||||
self.mem_pool_host = MLATokenToKVPoolHost(
|
"MiniMax M3 HiSparse requires one KV head per TP rank, "
|
||||||
device_pool=self.mem_pool_device,
|
f"got {self.mem_pool_device.head_num}. Increase the "
|
||||||
host_to_device_ratio=host_to_device_ratio,
|
"tensor-parallel size."
|
||||||
host_size=0,
|
)
|
||||||
page_size=self.mem_pool_device.page_size,
|
self.mem_pool_host = HiSparseMHATokenToKVPoolHost(
|
||||||
layout="layer_first",
|
device_pool=self.mem_pool_device,
|
||||||
override_kv_cache_dim=self.mem_pool_device.kv_cache_dim,
|
host_to_device_ratio=host_to_device_ratio,
|
||||||
)
|
page_size=self.mem_pool_device.page_size,
|
||||||
self.item_size_bytes = self.mem_pool_host.token_stride_size
|
)
|
||||||
|
self.item_size_bytes = self.mem_pool_device.bytes_per_token_k
|
||||||
|
else:
|
||||||
|
self.mem_pool_device: HiSparseDSATokenToKVPool = kvcache
|
||||||
|
self.mem_pool_host = MLATokenToKVPoolHost(
|
||||||
|
device_pool=self.mem_pool_device,
|
||||||
|
host_to_device_ratio=host_to_device_ratio,
|
||||||
|
host_size=0,
|
||||||
|
page_size=self.mem_pool_device.page_size,
|
||||||
|
layout="layer_first",
|
||||||
|
override_kv_cache_dim=self.mem_pool_device.kv_cache_dim,
|
||||||
|
)
|
||||||
|
self.item_size_bytes = self.mem_pool_host.token_stride_size
|
||||||
self.page_size = self.mem_pool_device.page_size
|
self.page_size = self.mem_pool_device.page_size
|
||||||
|
|
||||||
max_num_req_slots = req_to_token_pool.req_to_token.shape[0]
|
max_num_req_slots = req_to_token_pool.req_to_token.shape[0]
|
||||||
@@ -263,9 +285,17 @@ class HiSparseCoordinator:
|
|||||||
self.device_buffer_size, dtype=torch.int32, device=device
|
self.device_buffer_size, dtype=torch.int32, device=device
|
||||||
)
|
)
|
||||||
|
|
||||||
# Pre-allocated output buffer for swap_in_selected_pages (CUDA-graph safe)
|
# Pre-allocated output buffer for swap-in (CUDA-graph safe). MiniMax
|
||||||
|
# selects blocks, so its flattened token-slot output can occupy any
|
||||||
|
# prefix up to the full device working-set size.
|
||||||
|
swap_output_width = (
|
||||||
|
self.device_buffer_size if self.is_m3_hisparse else self.top_k
|
||||||
|
)
|
||||||
self.top_k_device_locs_buffer = torch.full(
|
self.top_k_device_locs_buffer = torch.full(
|
||||||
(max_num_req_slots, self.top_k), -1, dtype=torch.int32, device=device
|
(max_num_req_slots, swap_output_width),
|
||||||
|
-1,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=device,
|
||||||
)
|
)
|
||||||
self.raw_indices_buffer = torch.full(
|
self.raw_indices_buffer = torch.full(
|
||||||
(max_num_req_slots, self.top_k), -1, dtype=torch.int32, device=device
|
(max_num_req_slots, self.top_k), -1, dtype=torch.int32, device=device
|
||||||
@@ -457,7 +487,12 @@ class HiSparseCoordinator:
|
|||||||
host_indices = self.req_to_host_pool[req.kv.req_pool_idx, :n]
|
host_indices = self.req_to_host_pool[req.kv.req_pool_idx, :n]
|
||||||
device_locs = self.req_to_device_buffer[req.kv.req_pool_idx, :n]
|
device_locs = self.req_to_device_buffer[req.kv.req_pool_idx, :n]
|
||||||
|
|
||||||
for layer_id in range(self.mem_pool_device.layer_num):
|
layer_ids = (
|
||||||
|
range(self.mem_pool_device.start_layer, self.mem_pool_device.end_layer)
|
||||||
|
if self.is_m3_hisparse
|
||||||
|
else range(self.mem_pool_device.layer_num)
|
||||||
|
)
|
||||||
|
for layer_id in layer_ids:
|
||||||
self.mem_pool_host.load_to_device_per_layer(
|
self.mem_pool_host.load_to_device_per_layer(
|
||||||
self.mem_pool_device,
|
self.mem_pool_device,
|
||||||
host_indices,
|
host_indices,
|
||||||
@@ -642,13 +677,9 @@ class HiSparseCoordinator:
|
|||||||
compressed_locs = self.token_to_kv_pool_allocator.get_last_loc_compressed(
|
compressed_locs = self.token_to_kv_pool_allocator.get_last_loc_compressed(
|
||||||
out_cache_loc
|
out_cache_loc
|
||||||
)
|
)
|
||||||
# ROCm: the decode remap creates a temporary hisparse device slot per
|
# Page-size-one allocation creates a temporary slot before remapping
|
||||||
# new token (via the page_size==1 allocator path). Free the stale
|
# the new token into the request's reserved device-buffer slot.
|
||||||
# slot before pointing the mapping at the reserved device-buffer slot,
|
if _is_hip or self.mem_pool_device.page_size == 1:
|
||||||
# otherwise the temporary slots leak and corrupt later swap-in lookups.
|
|
||||||
# CUDA keeps the original behavior: the swap-in kernel consumes only
|
|
||||||
# top_k_device_locs, so stale mapping entries are harmless there.
|
|
||||||
if _is_hip:
|
|
||||||
previous_locs = self.mem_pool_device._translate_loc_to_hisparse_device(
|
previous_locs = self.mem_pool_device._translate_loc_to_hisparse_device(
|
||||||
compressed_locs
|
compressed_locs
|
||||||
)
|
)
|
||||||
@@ -976,8 +1007,7 @@ class HiSparseCoordinator:
|
|||||||
miss plan into self._miss_{src,dst,count} for the skip layers to replay.
|
miss plan into self._miss_{src,dst,count} for the skip layers to replay.
|
||||||
"""
|
"""
|
||||||
num_reqs = req_pool_indices.size(0)
|
num_reqs = req_pool_indices.size(0)
|
||||||
top_k_indices = self.top_k_device_locs_buffer[:num_reqs]
|
top_k_indices = self.top_k_device_locs_buffer[:num_reqs, : self.top_k]
|
||||||
|
|
||||||
swap_in_fn = (
|
swap_in_fn = (
|
||||||
load_cache_to_device_buffer_dsv4_mla
|
load_cache_to_device_buffer_dsv4_mla
|
||||||
if self.is_dsv4_hisparse
|
if self.is_dsv4_hisparse
|
||||||
@@ -1015,6 +1045,46 @@ class HiSparseCoordinator:
|
|||||||
)
|
)
|
||||||
return top_k_indices
|
return top_k_indices
|
||||||
|
|
||||||
|
def swap_in_selected_blocks(
|
||||||
|
self,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
top_k_blocks: torch.Tensor,
|
||||||
|
layer_id: int,
|
||||||
|
sparse_block_size: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
assert self.is_m3_hisparse
|
||||||
|
num_reqs = req_pool_indices.size(0)
|
||||||
|
num_selected_tokens = top_k_blocks.size(1) * sparse_block_size
|
||||||
|
assert num_selected_tokens <= self.device_buffer_size, (
|
||||||
|
f"MiniMax M3 selected {num_selected_tokens} tokens, but the "
|
||||||
|
f"HiSparse device buffer holds only {self.device_buffer_size}."
|
||||||
|
)
|
||||||
|
top_k_indices = self.top_k_device_locs_buffer[:num_reqs, :num_selected_tokens]
|
||||||
|
host_layer = layer_id - self.mem_pool_device.start_layer
|
||||||
|
load_blocks_to_device_buffer_mha(
|
||||||
|
top_k_blocks=top_k_blocks,
|
||||||
|
device_buffer_tokens=self.req_device_buffer_tokens[host_layer],
|
||||||
|
host_cache_locs=self.req_to_host_pool,
|
||||||
|
device_buffer_locs=self.req_device_buffer_token_locs[host_layer],
|
||||||
|
host_cache_k=self.mem_pool_host.k_buffer[host_layer],
|
||||||
|
host_cache_v=self.mem_pool_host.v_buffer[host_layer],
|
||||||
|
device_buffer_k=self.mem_pool_device.get_key_buffer(layer_id),
|
||||||
|
device_buffer_v=self.mem_pool_device.get_value_buffer(layer_id),
|
||||||
|
top_k_device_locs=top_k_indices,
|
||||||
|
req_pool_indices=req_pool_indices,
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
lru_slots=self.lru_slots[host_layer],
|
||||||
|
item_size_bytes=self.item_size_bytes,
|
||||||
|
hot_buffer_size=self.device_buffer_size,
|
||||||
|
sparse_block_size=sparse_block_size,
|
||||||
|
page_size=1,
|
||||||
|
block_size=self.swap_in_block_size,
|
||||||
|
num_real_reqs=self.num_real_reqs,
|
||||||
|
skip_io=self.skip_io,
|
||||||
|
)
|
||||||
|
return top_k_indices
|
||||||
|
|
||||||
def _run_copy_only_kernel(self, num_reqs: int, skip_layer: int) -> None:
|
def _run_copy_only_kernel(self, num_reqs: int, skip_layer: int) -> None:
|
||||||
"""Replay the anchor's recorded miss plan into a skip layer's buffers
|
"""Replay the anchor's recorded miss plan into a skip layer's buffers
|
||||||
(IO-only; the anchor's slot table stays valid -- lockstep layout)."""
|
(IO-only; the anchor's slot table stays valid -- lockstep layout)."""
|
||||||
@@ -1045,7 +1115,10 @@ class HiSparseCoordinator:
|
|||||||
"""
|
"""
|
||||||
if not self.enable_prefetch:
|
if not self.enable_prefetch:
|
||||||
return self._run_swap_in_kernel(
|
return self._run_swap_in_kernel(
|
||||||
req_pool_indices, compressed_seq_lens, top_k_result, layer_id
|
req_pool_indices,
|
||||||
|
compressed_seq_lens,
|
||||||
|
top_k_result,
|
||||||
|
layer_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
num_reqs = req_pool_indices.size(0)
|
num_reqs = req_pool_indices.size(0)
|
||||||
@@ -1054,7 +1127,7 @@ class HiSparseCoordinator:
|
|||||||
# applies (shared index + lockstep buffers).
|
# applies (shared index + lockstep buffers).
|
||||||
slot = self._prefetch_slot[layer_id]
|
slot = self._prefetch_slot[layer_id]
|
||||||
self._prefetch_events[slot].wait(device_module.current_stream())
|
self._prefetch_events[slot].wait(device_module.current_stream())
|
||||||
return self.top_k_device_locs_buffer[:num_reqs]
|
return self.top_k_device_locs_buffer[:num_reqs, : self.top_k]
|
||||||
|
|
||||||
# Anchor: swap in synchronously (recording the plan), then prefetch the
|
# Anchor: swap in synchronously (recording the plan), then prefetch the
|
||||||
# skip layers' copies on the side stream.
|
# skip layers' copies on the side stream.
|
||||||
|
|||||||
@@ -1,4 +1,7 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import weakref
|
import weakref
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -11,6 +14,9 @@ from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
|
|||||||
from sglang.srt.mem_cache.hisparse_memory_pool import HiSparseDSATokenToKVPool
|
from sglang.srt.mem_cache.hisparse_memory_pool import HiSparseDSATokenToKVPool
|
||||||
from sglang.srt.utils.common import get_num_new_pages
|
from sglang.srt.utils.common import get_num_new_pages
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
|
||||||
|
|
||||||
|
|
||||||
class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -19,7 +25,7 @@ class HiSparseTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
page_size: int,
|
page_size: int,
|
||||||
dtype: torch.dtype,
|
dtype: torch.dtype,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
kvcache: HiSparseDSATokenToKVPool,
|
kvcache: HiSparseDSATokenToKVPool | MiniMaxSparseKVPool,
|
||||||
need_sort: bool,
|
need_sort: bool,
|
||||||
host_to_device_ratio: int = 2,
|
host_to_device_ratio: int = 2,
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from sglang.kernels.ops.kvcache.hisparse_slot_mapping import (
|
|||||||
translate_padded_hisparse_locations,
|
translate_padded_hisparse_locations,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
|
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool, MHATokenToKVPool
|
||||||
from sglang.srt.utils import is_cuda, is_hip, is_xpu
|
from sglang.srt.utils import is_cuda, is_hip, is_xpu
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -147,3 +147,111 @@ class HiSparseDSATokenToKVPool(DSATokenToKVPool):
|
|||||||
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||||||
):
|
):
|
||||||
raise NotImplementedError("HiSparseDevicePool does not support load_cpu_copy")
|
raise NotImplementedError("HiSparseDevicePool does not support load_cpu_copy")
|
||||||
|
|
||||||
|
|
||||||
|
class HiSparseMHAMainPool(MHATokenToKVPool):
|
||||||
|
"""MHA KV pool with HiSparse logical-to-device mapping.
|
||||||
|
|
||||||
|
Used by MiniMax M3 HiSparse. The index pools (index_kv_pool, index_k_pool)
|
||||||
|
stay fully resident on the device and do not use this mapping.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
size: int,
|
||||||
|
page_size: int,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
head_num: int,
|
||||||
|
head_dim: int,
|
||||||
|
layer_num: int,
|
||||||
|
device: str,
|
||||||
|
enable_memory_saver: bool,
|
||||||
|
start_layer: Optional[int] = None,
|
||||||
|
end_layer: Optional[int] = None,
|
||||||
|
):
|
||||||
|
super().__init__(
|
||||||
|
size=size,
|
||||||
|
page_size=page_size,
|
||||||
|
dtype=dtype,
|
||||||
|
head_num=head_num,
|
||||||
|
head_dim=head_dim,
|
||||||
|
layer_num=layer_num,
|
||||||
|
device=device,
|
||||||
|
enable_memory_saver=enable_memory_saver,
|
||||||
|
start_layer=start_layer,
|
||||||
|
end_layer=end_layer,
|
||||||
|
)
|
||||||
|
self.full_to_hisparse_device_index_mapping: Optional[torch.Tensor] = None
|
||||||
|
self.bytes_per_token_k = head_num * head_dim * self.store_dtype.itemsize
|
||||||
|
self.bytes_per_token_v = head_num * self.v_head_dim * self.store_dtype.itemsize
|
||||||
|
|
||||||
|
def register_mapping(
|
||||||
|
self, full_to_hisparse_device_index_mapping: torch.Tensor
|
||||||
|
) -> None:
|
||||||
|
self.full_to_hisparse_device_index_mapping = (
|
||||||
|
full_to_hisparse_device_index_mapping
|
||||||
|
)
|
||||||
|
|
||||||
|
def translate_loc_to_hisparse_device(self, indices: torch.Tensor) -> torch.Tensor:
|
||||||
|
assert self.full_to_hisparse_device_index_mapping is not None
|
||||||
|
return self.full_to_hisparse_device_index_mapping[indices]
|
||||||
|
|
||||||
|
def _translate_loc_to_hisparse_device(self, indices: torch.Tensor) -> torch.Tensor:
|
||||||
|
assert self.full_to_hisparse_device_index_mapping is not None
|
||||||
|
return self.full_to_hisparse_device_index_mapping[indices]
|
||||||
|
|
||||||
|
def translate_loc_from_full_to_hisparse_device(
|
||||||
|
self, full_indices: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
assert self.full_to_hisparse_device_index_mapping is not None
|
||||||
|
return self.full_to_hisparse_device_index_mapping[full_indices]
|
||||||
|
|
||||||
|
def translate_loc_from_full_to_compressed(
|
||||||
|
self, full_indices: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
return full_indices
|
||||||
|
|
||||||
|
def set_kv_buffer(
|
||||||
|
self,
|
||||||
|
layer: RadixAttention,
|
||||||
|
loc,
|
||||||
|
cache_k: torch.Tensor,
|
||||||
|
cache_v: torch.Tensor,
|
||||||
|
*args,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
from sglang.srt.mem_cache.memory_pool import unwrap_write_loc
|
||||||
|
|
||||||
|
raw_loc, _, _ = unwrap_write_loc(loc)
|
||||||
|
translated = self.translate_loc_to_hisparse_device(raw_loc)
|
||||||
|
super().set_kv_buffer(layer, translated, cache_k, cache_v, *args, **kwargs)
|
||||||
|
|
||||||
|
def transfer_values_on_device(
|
||||||
|
self,
|
||||||
|
dst_indices: torch.Tensor,
|
||||||
|
src_indices: torch.Tensor,
|
||||||
|
) -> None:
|
||||||
|
transfer_kv_all_layer_mla(
|
||||||
|
src_layers=self.k_data_ptrs,
|
||||||
|
dst_layers=self.k_data_ptrs,
|
||||||
|
src_indices=src_indices,
|
||||||
|
dst_indices=dst_indices,
|
||||||
|
item_size=self.bytes_per_token_k,
|
||||||
|
num_layers=self.layer_num,
|
||||||
|
)
|
||||||
|
transfer_kv_all_layer_mla(
|
||||||
|
src_layers=self.v_data_ptrs,
|
||||||
|
dst_layers=self.v_data_ptrs,
|
||||||
|
src_indices=src_indices,
|
||||||
|
dst_indices=dst_indices,
|
||||||
|
item_size=self.bytes_per_token_v,
|
||||||
|
num_layers=self.layer_num,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
|
||||||
|
raise NotImplementedError("HiSparseMHAMainPool does not support get_cpu_copy")
|
||||||
|
|
||||||
|
def load_cpu_copy(
|
||||||
|
self, kv_cache_cpu, indices, mamba_indices=None, req_pool_index=None
|
||||||
|
):
|
||||||
|
raise NotImplementedError("HiSparseMHAMainPool does not support load_cpu_copy")
|
||||||
|
|||||||
@@ -1836,6 +1836,14 @@ class KVCacheConfigurator:
|
|||||||
disable_value_sparse_layer_ids = get_minimax_sparse_disable_value_layer_ids(
|
disable_value_sparse_layer_ids = get_minimax_sparse_disable_value_layer_ids(
|
||||||
sparse_cfg
|
sparse_cfg
|
||||||
)
|
)
|
||||||
|
enable_hisparse = get_memory().enable_hisparse
|
||||||
|
hisparse_kwargs = {}
|
||||||
|
if enable_hisparse:
|
||||||
|
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
|
||||||
|
|
||||||
|
hisparse_kwargs["host_to_device_ratio"] = (
|
||||||
|
parse_hisparse_config().host_to_device_ratio
|
||||||
|
)
|
||||||
token_to_kv_pool = MiniMaxSparseKVPool(
|
token_to_kv_pool = MiniMaxSparseKVPool(
|
||||||
size=max_total_num_tokens,
|
size=max_total_num_tokens,
|
||||||
page_size=self.pool_page_size,
|
page_size=self.pool_page_size,
|
||||||
@@ -1861,6 +1869,8 @@ class KVCacheConfigurator:
|
|||||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||||
start_layer=self.layer_info.start_layer,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
|
enable_hisparse=enable_hisparse,
|
||||||
|
**hisparse_kwargs,
|
||||||
)
|
)
|
||||||
return token_to_kv_pool
|
return token_to_kv_pool
|
||||||
|
|
||||||
|
|||||||
@@ -5448,6 +5448,8 @@ class MiniMaxSparseKVPool(KVCache):
|
|||||||
main_pool_cls=MHATokenToKVPool,
|
main_pool_cls=MHATokenToKVPool,
|
||||||
index_kv_pool_cls=MHATokenToKVPool,
|
index_kv_pool_cls=MHATokenToKVPool,
|
||||||
index_k_pool_cls=MHATokenToKOnlyPool,
|
index_k_pool_cls=MHATokenToKOnlyPool,
|
||||||
|
enable_hisparse: bool = False,
|
||||||
|
host_to_device_ratio: int = 2,
|
||||||
):
|
):
|
||||||
# Do not call super().__init__() — delegate to sub-pools instead.
|
# Do not call super().__init__() — delegate to sub-pools instead.
|
||||||
self.size = size
|
self.size = size
|
||||||
@@ -5466,6 +5468,7 @@ class MiniMaxSparseKVPool(KVCache):
|
|||||||
]
|
]
|
||||||
|
|
||||||
index_dtype = index_dtype if index_dtype is not None else dtype
|
index_dtype = index_dtype if index_dtype is not None else dtype
|
||||||
|
index_pool_size = size * host_to_device_ratio if enable_hisparse else size
|
||||||
|
|
||||||
# Split sparse layers by V policy: kv_sparse (index_kv_pool holds K+V) vs
|
# Split sparse layers by V policy: kv_sparse (index_kv_pool holds K+V) vs
|
||||||
# k_only_sparse (index_k_pool holds only K; V is never read).
|
# k_only_sparse (index_k_pool holds only K; V is never read).
|
||||||
@@ -5489,22 +5492,59 @@ class MiniMaxSparseKVPool(KVCache):
|
|||||||
gid: i for i, gid in enumerate(local_k_only_sparse_layer_ids)
|
gid: i for i, gid in enumerate(local_k_only_sparse_layer_ids)
|
||||||
}
|
}
|
||||||
|
|
||||||
self.main_pool = main_pool_cls(
|
self._dense_layer_ids = set(local_dense_layer_ids)
|
||||||
size=size,
|
main_layer_num = len(local_dense_layer_ids) + len(local_sparse_layer_ids)
|
||||||
page_size=page_size,
|
if enable_hisparse:
|
||||||
dtype=dtype,
|
from sglang.srt.mem_cache.hisparse_memory_pool import (
|
||||||
head_num=head_num,
|
HiSparseMHAMainPool,
|
||||||
head_dim=head_dim,
|
)
|
||||||
layer_num=len(local_dense_layer_ids) + len(local_sparse_layer_ids),
|
|
||||||
device=device,
|
self.dense_pool = (
|
||||||
enable_memory_saver=enable_memory_saver,
|
main_pool_cls(
|
||||||
start_layer=start_layer,
|
size=index_pool_size,
|
||||||
end_layer=end_layer,
|
page_size=page_size,
|
||||||
)
|
dtype=dtype,
|
||||||
|
head_num=head_num,
|
||||||
|
head_dim=head_dim,
|
||||||
|
layer_num=len(local_dense_layer_ids),
|
||||||
|
device=device,
|
||||||
|
enable_memory_saver=enable_memory_saver,
|
||||||
|
start_layer=start_layer,
|
||||||
|
end_layer=start_layer + len(local_dense_layer_ids),
|
||||||
|
)
|
||||||
|
if local_dense_layer_ids
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
self.main_pool = HiSparseMHAMainPool(
|
||||||
|
size=size,
|
||||||
|
page_size=page_size,
|
||||||
|
dtype=dtype,
|
||||||
|
head_num=head_num,
|
||||||
|
head_dim=head_dim,
|
||||||
|
layer_num=len(local_sparse_layer_ids),
|
||||||
|
device=device,
|
||||||
|
enable_memory_saver=enable_memory_saver,
|
||||||
|
start_layer=local_sparse_layer_ids[0],
|
||||||
|
end_layer=end_layer,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.dense_pool = None
|
||||||
|
self.main_pool = main_pool_cls(
|
||||||
|
size=size,
|
||||||
|
page_size=page_size,
|
||||||
|
dtype=dtype,
|
||||||
|
head_num=head_num,
|
||||||
|
head_dim=head_dim,
|
||||||
|
layer_num=main_layer_num,
|
||||||
|
device=device,
|
||||||
|
enable_memory_saver=enable_memory_saver,
|
||||||
|
start_layer=start_layer,
|
||||||
|
end_layer=end_layer,
|
||||||
|
)
|
||||||
|
|
||||||
self.index_kv_pool: Optional[MHATokenToKVPool] = (
|
self.index_kv_pool: Optional[MHATokenToKVPool] = (
|
||||||
index_kv_pool_cls(
|
index_kv_pool_cls(
|
||||||
size=size,
|
size=index_pool_size,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
dtype=index_dtype,
|
dtype=index_dtype,
|
||||||
head_num=1,
|
head_num=1,
|
||||||
@@ -5519,7 +5559,7 @@ class MiniMaxSparseKVPool(KVCache):
|
|||||||
|
|
||||||
self.index_k_pool: Optional[MHATokenToKOnlyPool] = (
|
self.index_k_pool: Optional[MHATokenToKOnlyPool] = (
|
||||||
index_k_pool_cls(
|
index_k_pool_cls(
|
||||||
size=size,
|
size=index_pool_size,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
dtype=index_dtype,
|
dtype=index_dtype,
|
||||||
head_num=1,
|
head_num=1,
|
||||||
@@ -5533,19 +5573,58 @@ class MiniMaxSparseKVPool(KVCache):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.mem_usage = self.main_pool.mem_usage
|
self.mem_usage = self.main_pool.mem_usage
|
||||||
|
if self.dense_pool is not None:
|
||||||
|
self.mem_usage += self.dense_pool.mem_usage
|
||||||
if self.index_kv_pool is not None:
|
if self.index_kv_pool is not None:
|
||||||
self.mem_usage += self.index_kv_pool.mem_usage
|
self.mem_usage += self.index_kv_pool.mem_usage
|
||||||
if self.index_k_pool is not None:
|
if self.index_k_pool is not None:
|
||||||
self.mem_usage += self.index_k_pool.mem_usage
|
self.mem_usage += self.index_k_pool.mem_usage
|
||||||
|
|
||||||
# HiCacheController reads these from the top-level KV pool wrapper.
|
# HiCacheController reads these from the top-level KV pool wrapper.
|
||||||
self.layer_num = self.main_pool.layer_num
|
self.layer_num = main_layer_num
|
||||||
self.start_layer = self.main_pool.start_layer
|
self.start_layer = start_layer
|
||||||
self.end_layer = self.main_pool.end_layer
|
self.end_layer = end_layer
|
||||||
# PD disaggregation reads these directly (no fallback) off the wrapper.
|
# PD disaggregation reads these directly (no fallback) off the wrapper.
|
||||||
self.head_num = self.main_pool.head_num
|
self.head_num = self.main_pool.head_num
|
||||||
self.head_dim = self.main_pool.head_dim
|
self.head_dim = self.main_pool.head_dim
|
||||||
|
self.v_head_dim = self.main_pool.v_head_dim
|
||||||
|
self.store_dtype = self.main_pool.store_dtype
|
||||||
self.layer_transfer_counter = None
|
self.layer_transfer_counter = None
|
||||||
|
self._enable_hisparse = enable_hisparse
|
||||||
|
|
||||||
|
def register_mapping(self, mapping: torch.Tensor) -> None:
|
||||||
|
assert self._enable_hisparse
|
||||||
|
self.main_pool.register_mapping(mapping)
|
||||||
|
|
||||||
|
def _translate_loc_to_hisparse_device(self, indices: torch.Tensor) -> torch.Tensor:
|
||||||
|
assert self._enable_hisparse
|
||||||
|
return self.main_pool._translate_loc_to_hisparse_device(indices)
|
||||||
|
|
||||||
|
def translate_loc_to_hisparse_device(self, indices: torch.Tensor) -> torch.Tensor:
|
||||||
|
assert self._enable_hisparse
|
||||||
|
return self.main_pool.translate_loc_to_hisparse_device(indices)
|
||||||
|
|
||||||
|
def translate_loc_from_full_to_hisparse_device(
|
||||||
|
self, indices: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
assert self._enable_hisparse
|
||||||
|
return self.main_pool.translate_loc_from_full_to_hisparse_device(indices)
|
||||||
|
|
||||||
|
def translate_loc_from_full_to_compressed(
|
||||||
|
self, indices: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
assert self._enable_hisparse
|
||||||
|
return self.main_pool.translate_loc_from_full_to_compressed(indices)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def bytes_per_token_k(self) -> int:
|
||||||
|
assert self._enable_hisparse
|
||||||
|
return self.main_pool.bytes_per_token_k
|
||||||
|
|
||||||
|
@property
|
||||||
|
def full_to_hisparse_device_index_mapping(self):
|
||||||
|
assert self._enable_hisparse
|
||||||
|
return self.main_pool.full_to_hisparse_device_index_mapping
|
||||||
|
|
||||||
def register_layer_transfer_counter(
|
def register_layer_transfer_counter(
|
||||||
self, layer_transfer_counter: LayerDoneCounter
|
self, layer_transfer_counter: LayerDoneCounter
|
||||||
@@ -5561,17 +5640,22 @@ class MiniMaxSparseKVPool(KVCache):
|
|||||||
if self.layer_transfer_counter is not None:
|
if self.layer_transfer_counter is not None:
|
||||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||||
|
|
||||||
|
def _pool_for(self, layer_id: int) -> MHATokenToKVPool:
|
||||||
|
if self.dense_pool is not None and layer_id in self._dense_layer_ids:
|
||||||
|
return self.dense_pool
|
||||||
|
return self.main_pool
|
||||||
|
|
||||||
def get_key_buffer(self, layer_id: int) -> torch.Tensor:
|
def get_key_buffer(self, layer_id: int) -> torch.Tensor:
|
||||||
self._wait_for_layer(layer_id)
|
self._wait_for_layer(layer_id)
|
||||||
return self.main_pool.get_key_buffer(layer_id)
|
return self._pool_for(layer_id).get_key_buffer(layer_id)
|
||||||
|
|
||||||
def get_value_buffer(self, layer_id: int) -> torch.Tensor:
|
def get_value_buffer(self, layer_id: int) -> torch.Tensor:
|
||||||
self._wait_for_layer(layer_id)
|
self._wait_for_layer(layer_id)
|
||||||
return self.main_pool.get_value_buffer(layer_id)
|
return self._pool_for(layer_id).get_value_buffer(layer_id)
|
||||||
|
|
||||||
def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
self._wait_for_layer(layer_id)
|
self._wait_for_layer(layer_id)
|
||||||
return self.main_pool.get_kv_buffer(layer_id)
|
return self._pool_for(layer_id).get_kv_buffer(layer_id)
|
||||||
|
|
||||||
def get_index_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
def get_index_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
self._wait_for_layer(layer_id)
|
self._wait_for_layer(layer_id)
|
||||||
@@ -5613,7 +5697,7 @@ class MiniMaxSparseKVPool(KVCache):
|
|||||||
Scale semantics follow MHATokenToKVPool: None means unit scale;
|
Scale semantics follow MHATokenToKVPool: None means unit scale;
|
||||||
a non-None scale is applied with an in-place div_ before the fp8 cast.
|
a non-None scale is applied with an in-place div_ before the fp8 cast.
|
||||||
"""
|
"""
|
||||||
self.main_pool.set_kv_buffer(
|
self._pool_for(layer.layer_id).set_kv_buffer(
|
||||||
layer,
|
layer,
|
||||||
loc,
|
loc,
|
||||||
cache_k,
|
cache_k,
|
||||||
@@ -5711,8 +5795,10 @@ class MiniMaxSparseKVPool(KVCache):
|
|||||||
disable_value = cache_idx_v is None
|
disable_value = cache_idx_v is None
|
||||||
index_pool = self.index_k_pool if disable_value else self.index_kv_pool
|
index_pool = self.index_k_pool if disable_value else self.index_kv_pool
|
||||||
|
|
||||||
if index_pool is not None and self._can_fuse_kv_index_store(
|
if (
|
||||||
index_pool, cache_k, cache_idx_k
|
index_pool is not None
|
||||||
|
and not self._enable_hisparse
|
||||||
|
and self._can_fuse_kv_index_store(index_pool, cache_k, cache_idx_k)
|
||||||
):
|
):
|
||||||
from sglang.kernels.ops.kvcache.minimax_store_kv_index import store_kv_index
|
from sglang.kernels.ops.kvcache.minimax_store_kv_index import store_kv_index
|
||||||
|
|
||||||
@@ -5759,7 +5845,12 @@ class MiniMaxSparseKVPool(KVCache):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_kv_size_bytes(self):
|
def get_kv_size_bytes(self):
|
||||||
sub_pools = [self.main_pool, self.index_kv_pool, self.index_k_pool]
|
sub_pools = [
|
||||||
|
self.main_pool,
|
||||||
|
self.dense_pool,
|
||||||
|
self.index_kv_pool,
|
||||||
|
self.index_k_pool,
|
||||||
|
]
|
||||||
sizes = [p.get_kv_size_bytes() for p in sub_pools if p is not None]
|
sizes = [p.get_kv_size_bytes() for p in sub_pools if p is not None]
|
||||||
return sum(k for k, _ in sizes), sum(v for _, v in sizes)
|
return sum(k for k, _ in sizes), sum(v for _, v in sizes)
|
||||||
|
|
||||||
@@ -5767,6 +5858,25 @@ class MiniMaxSparseKVPool(KVCache):
|
|||||||
# Main K/V only; index buffers ride the state-buffer channel.
|
# Main K/V only; index buffers ride the state-buffer channel.
|
||||||
return self.main_pool.get_contiguous_buf_infos()
|
return self.main_pool.get_contiguous_buf_infos()
|
||||||
|
|
||||||
|
def get_sparse_kv_buf_infos(self):
|
||||||
|
return self._get_layer_kv_buf_infos(
|
||||||
|
layer_ids=sorted(self.sparse_layer_id_mapping)
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_dense_kv_state_buf_infos(self):
|
||||||
|
# Dense KV uses logical device slots, independently of sparse host slots.
|
||||||
|
return self._get_layer_kv_buf_infos(layer_ids=sorted(self._dense_layer_ids))
|
||||||
|
|
||||||
|
def _get_layer_kv_buf_infos(self, *, layer_ids):
|
||||||
|
buffers = [self.get_key_buffer(layer_id) for layer_id in layer_ids] + [
|
||||||
|
self.get_value_buffer(layer_id) for layer_id in layer_ids
|
||||||
|
]
|
||||||
|
return (
|
||||||
|
[buffer.data_ptr() for buffer in buffers],
|
||||||
|
[buffer.nbytes for buffer in buffers],
|
||||||
|
[buffer[0].nbytes * self.page_size for buffer in buffers],
|
||||||
|
)
|
||||||
|
|
||||||
def get_index_k_state_buf_infos(self):
|
def get_index_k_state_buf_infos(self):
|
||||||
# Per-page item_len (MHATokenToKVPool convention); index rows share the
|
# Per-page item_len (MHATokenToKVPool convention); index rows share the
|
||||||
# main-KV `loc`, so the transfer reuses the same page-ids.
|
# main-KV `loc`, so the transfer reuses the same page-ids.
|
||||||
|
|||||||
@@ -43,6 +43,7 @@ from sglang.srt.mem_cache.pool_host.common import (
|
|||||||
get_allocator_from_storage,
|
get_allocator_from_storage,
|
||||||
make_kernel_ptr_table,
|
make_kernel_ptr_table,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.pool_host.hisparse import HiSparseHostPoolMixin
|
||||||
from sglang.srt.utils import is_cuda, is_hip, is_mps, is_npu, is_xpu
|
from sglang.srt.utils import is_cuda, is_hip, is_mps, is_npu, is_xpu
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
@@ -1098,6 +1099,76 @@ class MHATokenToKOnlyPoolHost(HostKVCache):
|
|||||||
return ptr_list, element_size_list
|
return ptr_list, element_size_list
|
||||||
|
|
||||||
|
|
||||||
|
class HiSparseMHATokenToKVPoolHost(HiSparseHostPoolMixin, MHATokenToKVPoolHost):
|
||||||
|
"""Layer-first MHA host pool with page-granular HiSparse allocation."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
device_pool: MHATokenToKVPool,
|
||||||
|
host_to_device_ratio: float,
|
||||||
|
page_size: int,
|
||||||
|
):
|
||||||
|
super().__init__(
|
||||||
|
device_pool=device_pool,
|
||||||
|
host_to_device_ratio=host_to_device_ratio,
|
||||||
|
host_size=0,
|
||||||
|
page_size=page_size,
|
||||||
|
layout="layer_first",
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_contiguous_buf_infos(self):
|
||||||
|
buffers = self.k_data_refs + self.v_data_refs
|
||||||
|
return (
|
||||||
|
[buffer.data_ptr() for buffer in buffers],
|
||||||
|
[buffer.nbytes for buffer in buffers],
|
||||||
|
[buffer[0].nbytes * self.page_size for buffer in buffers],
|
||||||
|
)
|
||||||
|
|
||||||
|
def load_to_device_per_layer(
|
||||||
|
self,
|
||||||
|
device_pool,
|
||||||
|
host_indices,
|
||||||
|
device_indices,
|
||||||
|
layer_id,
|
||||||
|
io_backend,
|
||||||
|
*,
|
||||||
|
is_draft: bool = False,
|
||||||
|
):
|
||||||
|
if io_backend != "kernel" or is_draft:
|
||||||
|
raise ValueError(
|
||||||
|
"MiniMax M3 HiSparse host transfers require the kernel backend."
|
||||||
|
)
|
||||||
|
host_layer = layer_id - device_pool.start_layer
|
||||||
|
transfer_kv_per_layer(
|
||||||
|
src_k=self.k_buffer[host_layer],
|
||||||
|
dst_k=device_pool.get_key_buffer(layer_id),
|
||||||
|
src_v=self.v_buffer[host_layer],
|
||||||
|
dst_v=device_pool.get_value_buffer(layer_id),
|
||||||
|
src_indices=host_indices,
|
||||||
|
dst_indices=device_indices,
|
||||||
|
item_size=self.token_stride_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
def backup_from_device_all_layer(
|
||||||
|
self, device_pool, host_indices, device_indices, io_backend
|
||||||
|
):
|
||||||
|
if io_backend != "kernel":
|
||||||
|
raise ValueError(
|
||||||
|
"MiniMax M3 HiSparse host transfers require the kernel backend."
|
||||||
|
)
|
||||||
|
for layer_id in range(device_pool.start_layer, device_pool.end_layer):
|
||||||
|
host_layer = layer_id - device_pool.start_layer
|
||||||
|
transfer_kv_per_layer(
|
||||||
|
src_k=device_pool.get_key_buffer(layer_id),
|
||||||
|
dst_k=self.k_buffer[host_layer],
|
||||||
|
src_v=device_pool.get_value_buffer(layer_id),
|
||||||
|
dst_v=self.v_buffer[host_layer],
|
||||||
|
src_indices=device_indices,
|
||||||
|
dst_indices=host_indices,
|
||||||
|
item_size=self.token_stride_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
||||||
"""Host KV pool for MHA models whose K and V have different head dims
|
"""Host KV pool for MHA models whose K and V have different head dims
|
||||||
(``head_dim != v_head_dim``), e.g. MiMo-V2.
|
(``head_dim != v_head_dim``), e.g. MiMo-V2.
|
||||||
|
|||||||
@@ -28,6 +28,7 @@ if is_xpu():
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
from sglang.kernels.ops.kvcache.hisparse import (
|
from sglang.kernels.ops.kvcache.hisparse import (
|
||||||
|
load_blocks_to_device_buffer_mha,
|
||||||
load_cache_to_device_buffer_dsv4_mla,
|
load_cache_to_device_buffer_dsv4_mla,
|
||||||
load_cache_to_device_buffer_mla,
|
load_cache_to_device_buffer_mla,
|
||||||
transfer_cache_dsv4_mla,
|
transfer_cache_dsv4_mla,
|
||||||
@@ -368,6 +369,84 @@ def test_load_cache_to_device_buffer_hits_newest_and_updates_lru() -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(is_xpu(), reason="MiniMax MHA block swap-in has no XPU kernel.")
|
||||||
|
def test_load_blocks_to_device_buffer_mha_handles_partial_newest_block() -> None:
|
||||||
|
"""A partial newest block must not consume slots for its invalid tail."""
|
||||||
|
sparse_block_size = 4
|
||||||
|
hot_buffer_size = 8
|
||||||
|
host_k = _host_cache()
|
||||||
|
host_v = _host_cache()
|
||||||
|
host_v.add_(1000)
|
||||||
|
device_k = torch.full(
|
||||||
|
(DEVICE_CACHE_SIZE, 1, KV_DIM), -1, dtype=DTYPE, device=DEVICE
|
||||||
|
)
|
||||||
|
device_v = torch.full_like(device_k, -1)
|
||||||
|
device_buffer_locs = torch.arange(
|
||||||
|
hot_buffer_size + 1, dtype=torch.int32, device=DEVICE
|
||||||
|
).view(1, -1)
|
||||||
|
device_buffer_tokens = torch.tensor(
|
||||||
|
[[0, 1, 2, 3, -1, -1, -1, -1, -1]],
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=DEVICE,
|
||||||
|
)
|
||||||
|
for slot, token in enumerate([0, 1, 2, 3]):
|
||||||
|
device_k[device_buffer_locs[0, slot]].copy_(host_k[token], non_blocking=True)
|
||||||
|
device_v[device_buffer_locs[0, slot]].copy_(host_v[token], non_blocking=True)
|
||||||
|
device_k[device_buffer_locs[0, hot_buffer_size]].copy_(
|
||||||
|
host_k[10], non_blocking=True
|
||||||
|
)
|
||||||
|
device_v[device_buffer_locs[0, hot_buffer_size]].copy_(
|
||||||
|
host_v[10], non_blocking=True
|
||||||
|
)
|
||||||
|
|
||||||
|
top_k_blocks = torch.tensor([[0, 2]], dtype=torch.int32, device=DEVICE)
|
||||||
|
out = torch.full(
|
||||||
|
(1, top_k_blocks.size(1) * sparse_block_size),
|
||||||
|
-1,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=DEVICE,
|
||||||
|
)
|
||||||
|
lru_slots = torch.arange(hot_buffer_size, dtype=torch.int16, device=DEVICE).view(
|
||||||
|
1, -1
|
||||||
|
)
|
||||||
|
load_blocks_to_device_buffer_mha(
|
||||||
|
top_k_blocks=top_k_blocks,
|
||||||
|
device_buffer_tokens=device_buffer_tokens,
|
||||||
|
host_cache_locs=torch.arange(
|
||||||
|
HOST_CACHE_SIZE, dtype=torch.int64, device=DEVICE
|
||||||
|
).view(1, -1),
|
||||||
|
device_buffer_locs=device_buffer_locs,
|
||||||
|
host_cache_k=host_k,
|
||||||
|
host_cache_v=host_v,
|
||||||
|
device_buffer_k=device_k,
|
||||||
|
device_buffer_v=device_v,
|
||||||
|
top_k_device_locs=out,
|
||||||
|
req_pool_indices=torch.tensor([0], dtype=torch.int64, device=DEVICE),
|
||||||
|
seq_lens=torch.tensor([11], dtype=torch.int32, device=DEVICE),
|
||||||
|
lru_slots=lru_slots,
|
||||||
|
item_size_bytes=ITEM_SIZE_BYTES,
|
||||||
|
hot_buffer_size=hot_buffer_size,
|
||||||
|
sparse_block_size=sparse_block_size,
|
||||||
|
num_real_reqs=torch.tensor([1], dtype=torch.int32, device=DEVICE),
|
||||||
|
)
|
||||||
|
get_device_module().synchronize()
|
||||||
|
|
||||||
|
assert torch.equal(
|
||||||
|
out.cpu(), torch.tensor([[0, 1, 2, 3, 4, 5, 8, -1]], dtype=torch.int32)
|
||||||
|
)
|
||||||
|
assert torch.equal(device_k[4].cpu(), host_k[8])
|
||||||
|
assert torch.equal(device_v[4].cpu(), host_v[8])
|
||||||
|
assert torch.equal(device_k[5].cpu(), host_k[9])
|
||||||
|
assert torch.equal(device_v[5].cpu(), host_v[9])
|
||||||
|
assert torch.equal(
|
||||||
|
device_buffer_tokens.cpu(),
|
||||||
|
torch.tensor([[0, 1, 2, 3, 8, 9, -1, -1, -1]], dtype=torch.int32),
|
||||||
|
)
|
||||||
|
assert torch.equal(
|
||||||
|
lru_slots.cpu(), torch.tensor([[6, 7, 4, 5, 0, 1, 2, 3]], dtype=torch.int16)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_load_cache_to_device_buffer_miss_uses_updated_lru_slot() -> None:
|
def test_load_cache_to_device_buffer_miss_uses_updated_lru_slot() -> None:
|
||||||
state = _long_case()
|
state = _long_case()
|
||||||
|
|
||||||
|
|||||||
@@ -53,14 +53,19 @@ def _make_kv_pool(start_layer: int = 0) -> MiniMaxSparseKVPool:
|
|||||||
|
|
||||||
|
|
||||||
class TestMiniMaxSparseDisaggStateKvArgs(unittest.TestCase):
|
class TestMiniMaxSparseDisaggStateKvArgs(unittest.TestCase):
|
||||||
def test_setup_state_kv_args_single_minimax_component(self):
|
def test_setup_state_kv_args_minimax_components(self):
|
||||||
pool = _make_k_only_pool()
|
pool = _make_k_only_pool()
|
||||||
kv_args = KVArgs()
|
kv_args = KVArgs()
|
||||||
setup_state_kv_args(kv_args, pool)
|
setup_state_kv_args(kv_args, pool)
|
||||||
self.assertEqual(kv_args.state_types, [StateType.MINIMAX_INDEX_K])
|
self.assertEqual(
|
||||||
self.assertEqual(len(kv_args.state_data_ptrs), 1)
|
kv_args.state_types,
|
||||||
|
[StateType.MINIMAX_INDEX_K, StateType.MINIMAX_DENSE_KV],
|
||||||
|
)
|
||||||
|
self.assertEqual(len(kv_args.state_data_ptrs), 2)
|
||||||
self.assertEqual(len(kv_args.state_data_ptrs[0]), pool.index_k_pool.layer_num)
|
self.assertEqual(len(kv_args.state_data_ptrs[0]), pool.index_k_pool.layer_num)
|
||||||
self.assertEqual(len(kv_args.state_item_lens[0]), pool.index_k_pool.layer_num)
|
self.assertEqual(len(kv_args.state_item_lens[0]), pool.index_k_pool.layer_num)
|
||||||
|
self.assertEqual(len(kv_args.state_data_ptrs[1]), 6)
|
||||||
|
self.assertEqual(len(kv_args.state_item_lens[1]), 6)
|
||||||
|
|
||||||
def test_index_kv_pool_raises(self):
|
def test_index_kv_pool_raises(self):
|
||||||
pool = _make_kv_pool()
|
pool = _make_kv_pool()
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
|
import ctypes
|
||||||
import unittest
|
import unittest
|
||||||
from threading import Event
|
from threading import Event
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
@@ -6,6 +7,7 @@ from unittest.mock import MagicMock, call, patch
|
|||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
|
from sglang.srt.disaggregation.base.conn import StateType
|
||||||
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager
|
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
@@ -132,6 +134,61 @@ class TestMooncakeTransferBatching(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestMiniMaxStateTransfer(CustomTestCase):
|
||||||
|
def test_index_truncates_but_dense_rejects_mismatched_page_lists(self):
|
||||||
|
"""Legacy index transfers copy the common prefix; incomplete dense KV must fail."""
|
||||||
|
|
||||||
|
def copy_bytes(session, sources, destinations, lengths):
|
||||||
|
for src, dst, length in zip(sources, destinations, lengths, strict=True):
|
||||||
|
ctypes.memmove(dst, src, length)
|
||||||
|
return 0
|
||||||
|
|
||||||
|
for state in (StateType.MINIMAX_INDEX_K, StateType.MINIMAX_DENSE_KV):
|
||||||
|
for src_pages, dst_pages in (([1], [0]), ([1, 2], [0]), ([1], [0, 2])):
|
||||||
|
with self.subTest(state=state, src=src_pages, dst=dst_pages):
|
||||||
|
src = np.arange(3, dtype=np.int32)
|
||||||
|
dst = np.full(3, -1, dtype=np.int32)
|
||||||
|
manager = MooncakeKVManager.__new__(MooncakeKVManager)
|
||||||
|
manager.kv_args = SimpleNamespace(
|
||||||
|
state_types=[state],
|
||||||
|
state_data_ptrs=[[src.ctypes.data]],
|
||||||
|
state_item_lens=[[src.itemsize]],
|
||||||
|
state_dim_per_tensor=[[]],
|
||||||
|
state_layer_ids=[[]],
|
||||||
|
)
|
||||||
|
manager.engine = SimpleNamespace(batch_transfer_sync=copy_bytes)
|
||||||
|
manager.pp_size = manager.attn_tp_size = 1
|
||||||
|
manager.is_mla_backend = manager.is_hybrid_mla_backend = False
|
||||||
|
manager.enable_custom_mem_pool = False
|
||||||
|
manager.max_transfer_batch_indices = 0
|
||||||
|
peer = SimpleNamespace(
|
||||||
|
dst_state_data_ptrs=[[dst.ctypes.data]],
|
||||||
|
dst_state_item_lens=[[dst.itemsize]],
|
||||||
|
dst_state_dim_per_tensor=[[]],
|
||||||
|
dst_state_layer_ids=[[]],
|
||||||
|
dst_attn_tp_size=1,
|
||||||
|
)
|
||||||
|
kwargs = dict(
|
||||||
|
req=SimpleNamespace(
|
||||||
|
mooncake_session_id="cpu", dst_state_indices=[dst_pages]
|
||||||
|
),
|
||||||
|
prefill_state_indices=[src_pages],
|
||||||
|
executor=None,
|
||||||
|
target_rank_registration_info=peer,
|
||||||
|
)
|
||||||
|
if state == StateType.MINIMAX_DENSE_KV and len(src_pages) != len(
|
||||||
|
dst_pages
|
||||||
|
):
|
||||||
|
with self.assertRaisesRegex(
|
||||||
|
RuntimeError, "state index length mismatch"
|
||||||
|
):
|
||||||
|
manager.maybe_send_extra(**kwargs)
|
||||||
|
np.testing.assert_array_equal(dst, [-1, -1, -1])
|
||||||
|
else:
|
||||||
|
self.assertEqual(manager.maybe_send_extra(**kwargs), 0)
|
||||||
|
np.testing.assert_array_equal(dst, [1, -1, -1])
|
||||||
|
|
||||||
|
|
||||||
class TestDcpDraftHeadTransfer(unittest.TestCase):
|
class TestDcpDraftHeadTransfer(unittest.TestCase):
|
||||||
def test_transfers_draft_heads_to_logical_destination_rows(self):
|
def test_transfers_draft_heads_to_logical_destination_rows(self):
|
||||||
for src_tp, dst_tp in ((4, 8), (8, 4), (8, 8), (4, 32), (32, 4)):
|
for src_tp, dst_tp in ((4, 8), (8, 4), (8, 8), (4, 32), (32, 4)):
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import unittest
|
import unittest
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
@@ -18,6 +18,71 @@ from sglang.test.test_utils import CustomTestCase
|
|||||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class TestHiSparseDecodeRemap(CustomTestCase):
|
||||||
|
def test_page_size_one_reclaims_temporary_device_slot(self):
|
||||||
|
"""Decode remapping must reclaim its temporary slot without freeing the live slot."""
|
||||||
|
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator
|
||||||
|
from sglang.srt.mem_cache.allocator.hisparse import (
|
||||||
|
HiSparseTokenToKVPoolAllocator,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
|
||||||
|
|
||||||
|
pool = MiniMaxSparseKVPool(
|
||||||
|
size=8,
|
||||||
|
page_size=1,
|
||||||
|
dtype=torch.float32,
|
||||||
|
head_num=1,
|
||||||
|
head_dim=8,
|
||||||
|
idx_head_dim=16,
|
||||||
|
dense_layer_ids=[0],
|
||||||
|
sparse_layer_ids=[1],
|
||||||
|
disable_value_sparse_layer_ids=[1],
|
||||||
|
device="cpu",
|
||||||
|
start_layer=0,
|
||||||
|
end_layer=2,
|
||||||
|
enable_hisparse=True,
|
||||||
|
)
|
||||||
|
allocator = HiSparseTokenToKVPoolAllocator(
|
||||||
|
size=pool.size,
|
||||||
|
page_size=1,
|
||||||
|
dtype=pool.dtype,
|
||||||
|
device="cpu",
|
||||||
|
kvcache=pool,
|
||||||
|
need_sort=False,
|
||||||
|
)
|
||||||
|
coordinator = HiSparseCoordinator.__new__(HiSparseCoordinator)
|
||||||
|
coordinator.is_dsv4_hisparse = False
|
||||||
|
coordinator.mem_pool_device = pool.main_pool
|
||||||
|
coordinator.token_to_kv_pool_allocator = allocator
|
||||||
|
coordinator.device_buffer_size = 2
|
||||||
|
coordinator.req_to_device_buffer = allocator.hisparse_attn_allocator.alloc(
|
||||||
|
3
|
||||||
|
).reshape(1, 3)
|
||||||
|
coordinator.req_device_buffer_size = torch.tensor([3])
|
||||||
|
coordinator.req_device_buffer_token_locs = torch.zeros(
|
||||||
|
(1, 1, 3), dtype=torch.int32
|
||||||
|
)
|
||||||
|
coordinator._skip_first_backup = [True]
|
||||||
|
out_loc = allocator.alloc(1)
|
||||||
|
with patch("sglang.srt.managers.hisparse_coordinator._is_hip", False):
|
||||||
|
for _ in range(2):
|
||||||
|
coordinator._skip_first_backup[0] = True
|
||||||
|
coordinator.map_last_loc_to_buffer(
|
||||||
|
seq_lens=torch.tensor([3]),
|
||||||
|
out_cache_loc=out_loc,
|
||||||
|
req_pool_indices=torch.tensor([0]),
|
||||||
|
seq_lens_cpu=torch.tensor([3]),
|
||||||
|
req_pool_indices_cpu=torch.tensor([0]),
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
allocator.hisparse_attn_allocator.available_size(), pool.size - 3
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(
|
||||||
|
allocator.full_to_hisparse_device_index_mapping[out_loc],
|
||||||
|
coordinator.req_to_device_buffer[:, 2],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestDeepSeekV4HiSparseAllocator(CustomTestCase):
|
class TestDeepSeekV4HiSparseAllocator(CustomTestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
# The code under test reads its config from the bags.
|
# The code under test reads its config from the bags.
|
||||||
|
|||||||
@@ -1,14 +1,23 @@
|
|||||||
import unittest
|
import unittest
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.disaggregation.utils import get_kv_transfer_buf_infos
|
||||||
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
|
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
|
||||||
|
from sglang.srt.mem_cache.pool_host.mha import (
|
||||||
|
HiSparseMHATokenToKVPoolHost,
|
||||||
|
MHATokenToKVPoolHost,
|
||||||
|
)
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
def _make_k_only_pool(start_layer: int = 0) -> MiniMaxSparseKVPool:
|
def _make_k_only_pool(
|
||||||
|
start_layer: int = 0, *, enable_hisparse: bool = False
|
||||||
|
) -> MiniMaxSparseKVPool:
|
||||||
"""Mirror the released MiniMax-M3 config shape: all sparse layers K-only."""
|
"""Mirror the released MiniMax-M3 config shape: all sparse layers K-only."""
|
||||||
dense_layer_ids = [start_layer, start_layer + 1, start_layer + 2]
|
dense_layer_ids = [start_layer, start_layer + 1, start_layer + 2]
|
||||||
sparse_layer_ids = [start_layer + 3 + i for i in range(4)]
|
sparse_layer_ids = [start_layer + 3 + i for i in range(4)]
|
||||||
@@ -26,10 +35,11 @@ def _make_k_only_pool(start_layer: int = 0) -> MiniMaxSparseKVPool:
|
|||||||
device="cpu",
|
device="cpu",
|
||||||
start_layer=start_layer,
|
start_layer=start_layer,
|
||||||
end_layer=end_layer,
|
end_layer=end_layer,
|
||||||
|
enable_hisparse=enable_hisparse,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestMiniMaxSparsePoolPD(unittest.TestCase):
|
class TestMiniMaxSparsePoolPD(CustomTestCase):
|
||||||
def test_contiguous_buf_infos_main_only(self):
|
def test_contiguous_buf_infos_main_only(self):
|
||||||
pool = _make_k_only_pool()
|
pool = _make_k_only_pool()
|
||||||
ptrs, lens, item_lens = pool.get_contiguous_buf_infos()
|
ptrs, lens, item_lens = pool.get_contiguous_buf_infos()
|
||||||
@@ -53,6 +63,52 @@ class TestMiniMaxSparsePoolPD(unittest.TestCase):
|
|||||||
self.assertEqual(lens[i], buf.nbytes)
|
self.assertEqual(lens[i], buf.nbytes)
|
||||||
self.assertEqual(item_lens[i], buf[0].nbytes * pool.page_size)
|
self.assertEqual(item_lens[i], buf[0].nbytes * pool.page_size)
|
||||||
|
|
||||||
|
def test_hisparse_host_registration(self):
|
||||||
|
"""PD startup must register every sparse host K/V buffer with page strides."""
|
||||||
|
pool = _make_k_only_pool(enable_hisparse=True)
|
||||||
|
host = HiSparseMHATokenToKVPoolHost.__new__(HiSparseMHATokenToKVPoolHost)
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.mem_cache.pool_host.base.host_memory_budget_bytes",
|
||||||
|
return_value=1 << 30,
|
||||||
|
):
|
||||||
|
MHATokenToKVPoolHost.__init__(
|
||||||
|
host,
|
||||||
|
device_pool=pool.main_pool,
|
||||||
|
host_to_device_ratio=2,
|
||||||
|
host_size=0,
|
||||||
|
page_size=pool.page_size,
|
||||||
|
layout="layer_first",
|
||||||
|
pin_memory=False,
|
||||||
|
)
|
||||||
|
ptrs, lens, item_lens = get_kv_transfer_buf_infos(host)
|
||||||
|
buffers = list(host.k_buffer.unbind()) + list(host.v_buffer.unbind())
|
||||||
|
self.assertEqual(len(buffers), 8)
|
||||||
|
self.assertEqual(ptrs, [buffer.data_ptr() for buffer in buffers])
|
||||||
|
self.assertEqual(lens, [buffer.nbytes for buffer in buffers])
|
||||||
|
self.assertEqual(
|
||||||
|
item_lens, [buffer[0].nbytes * pool.page_size for buffer in buffers]
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_pd_registration_separates_dense_and_sparse_layers(self):
|
||||||
|
"""Both PD peers must keep dense device KV out of the sparse transfer list."""
|
||||||
|
for hisparse in (False, True):
|
||||||
|
pool = _make_k_only_pool(enable_hisparse=hisparse)
|
||||||
|
for layers, infos in (
|
||||||
|
(range(3, 7), get_kv_transfer_buf_infos(pool)),
|
||||||
|
(range(3), pool.get_dense_kv_state_buf_infos()),
|
||||||
|
):
|
||||||
|
with self.subTest(hisparse=hisparse, layers=layers):
|
||||||
|
buffers = [pool.get_key_buffer(i) for i in layers] + [
|
||||||
|
pool.get_value_buffer(i) for i in layers
|
||||||
|
]
|
||||||
|
ptrs, lens, item_lens = infos
|
||||||
|
self.assertEqual(ptrs, [buffer.data_ptr() for buffer in buffers])
|
||||||
|
self.assertEqual(lens, [buffer.nbytes for buffer in buffers])
|
||||||
|
self.assertEqual(
|
||||||
|
item_lens,
|
||||||
|
[buffer[0].nbytes * pool.page_size for buffer in buffers],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user