[Kernel] Move CUDA and ROCm speculative kernels to JIT (#40033)
This commit is contained in:
@@ -74,9 +74,9 @@ jobs:
|
|||||||
if: steps.gate.outputs.run_job == 'true'
|
if: steps.gate.outputs.run_job == 'true'
|
||||||
timeout-minutes: 30
|
timeout-minutes: 30
|
||||||
run: |
|
run: |
|
||||||
pytest python/sglang/kernels/aot/tests/speculative/test_eagle_utils.py
|
pytest test/registered/kernels/ops/speculative/test_eagle_utils.py
|
||||||
pytest python/sglang/kernels/aot/tests/speculative/test_ngram_utils.py
|
pytest test/registered/kernels/ops/speculative/test_ngram_utils.py
|
||||||
pytest python/sglang/kernels/aot/tests/speculative/test_speculative_sampling.py
|
pytest test/registered/kernels/ops/speculative/test_speculative_sampling.py
|
||||||
pytest python/sglang/kernels/aot/tests/test_torch_defaults_reset.py
|
pytest python/sglang/kernels/aot/tests/test_torch_defaults_reset.py
|
||||||
|
|
||||||
# ==================== General: multimodal layer ====================
|
# ==================== General: multimodal layer ====================
|
||||||
|
|||||||
@@ -321,7 +321,7 @@ jobs:
|
|||||||
}
|
}
|
||||||
run_pytest docker exec -w /sglang-checkout/python/sglang/kernels/aot/tests ci_sglang python3 -m pytest test_moe_align.py
|
run_pytest docker exec -w /sglang-checkout/python/sglang/kernels/aot/tests ci_sglang python3 -m pytest test_moe_align.py
|
||||||
run_pytest docker exec -w /sglang-checkout/python/sglang/kernels/aot/tests ci_sglang python3 -m pytest test_moe_topk_softmax.py
|
run_pytest docker exec -w /sglang-checkout/python/sglang/kernels/aot/tests ci_sglang python3 -m pytest test_moe_topk_softmax.py
|
||||||
run_pytest docker exec -w /sglang-checkout/python/sglang/kernels/aot/tests/speculative ci_sglang python3 -m pytest test_eagle_utils.py
|
run_pytest docker exec -w /sglang-checkout/test/registered/kernels/ops/speculative ci_sglang python3 -m pytest test_eagle_utils.py
|
||||||
run_pytest docker exec -w /sglang-checkout/python/sglang/kernels/aot/tests ci_sglang python3 -m pytest test_apply_token_bitmask_inplace.py
|
run_pytest docker exec -w /sglang-checkout/python/sglang/kernels/aot/tests ci_sglang python3 -m pytest test_apply_token_bitmask_inplace.py
|
||||||
run_pytest docker exec -w /sglang-checkout/python/sglang/kernels/aot/tests ci_sglang python3 -m pytest test_activation.py
|
run_pytest docker exec -w /sglang-checkout/python/sglang/kernels/aot/tests ci_sglang python3 -m pytest test_activation.py
|
||||||
run_pytest docker exec -w /sglang-checkout/python/sglang/kernels/aot/tests ci_sglang python3 -m pytest test_topk.py
|
run_pytest docker exec -w /sglang-checkout/python/sglang/kernels/aot/tests ci_sglang python3 -m pytest test_topk.py
|
||||||
|
|||||||
@@ -222,9 +222,9 @@ jobs:
|
|||||||
- name: Run sgl-kernel test
|
- name: Run sgl-kernel test
|
||||||
timeout-minutes: 20
|
timeout-minutes: 20
|
||||||
run: |
|
run: |
|
||||||
pytest python/sglang/kernels/aot/tests/speculative/test_eagle_utils.py
|
pytest test/registered/kernels/ops/speculative/test_eagle_utils.py
|
||||||
pytest python/sglang/kernels/aot/tests/speculative/test_ngram_utils.py
|
pytest test/registered/kernels/ops/speculative/test_ngram_utils.py
|
||||||
pytest python/sglang/kernels/aot/tests/speculative/test_speculative_sampling.py
|
pytest test/registered/kernels/ops/speculative/test_speculative_sampling.py
|
||||||
pytest python/sglang/kernels/aot/tests/test_torch_defaults_reset.py
|
pytest python/sglang/kernels/aot/tests/test_torch_defaults_reset.py
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -295,10 +295,6 @@ set(SOURCES
|
|||||||
"csrc/moe/prepare_moe_input.cu"
|
"csrc/moe/prepare_moe_input.cu"
|
||||||
|
|
||||||
"csrc/quantization/gguf/gguf_kernel.cu"
|
"csrc/quantization/gguf/gguf_kernel.cu"
|
||||||
"csrc/speculative/eagle_utils.cu"
|
|
||||||
"csrc/speculative/ngram_utils.cu"
|
|
||||||
"csrc/speculative/packbit.cu"
|
|
||||||
"csrc/speculative/speculative_sampling.cu"
|
|
||||||
|
|
||||||
"${repo-flashinfer_SOURCE_DIR}/csrc/norm.cu"
|
"${repo-flashinfer_SOURCE_DIR}/csrc/norm.cu"
|
||||||
"${repo-flashinfer_SOURCE_DIR}/csrc/renorm.cu"
|
"${repo-flashinfer_SOURCE_DIR}/csrc/renorm.cu"
|
||||||
|
|||||||
@@ -220,41 +220,6 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
|||||||
" int chunk_size, int topk) -> ()");
|
" int chunk_size, int topk) -> ()");
|
||||||
m.impl("cutlass_w4a8_moe_mm", torch::kCUDA, &cutlass_w4a8_moe_mm);
|
m.impl("cutlass_w4a8_moe_mm", torch::kCUDA, &cutlass_w4a8_moe_mm);
|
||||||
|
|
||||||
/*
|
|
||||||
* From csrc/speculative
|
|
||||||
*/
|
|
||||||
m.def(
|
|
||||||
"tree_speculative_sampling_target_only(Tensor! predicts, Tensor! accept_index, Tensor! accept_token_num, "
|
|
||||||
"Tensor candidates, Tensor retrive_index, Tensor retrive_next_token, Tensor retrive_next_sibling, "
|
|
||||||
"Tensor uniform_samples, Tensor uniform_samples_for_final_sampling, Tensor target_probs, Tensor draft_probs, "
|
|
||||||
"float threshold_single, float threshold_acc, "
|
|
||||||
"bool deterministic) -> ()");
|
|
||||||
m.impl("tree_speculative_sampling_target_only", torch::kCUDA, &tree_speculative_sampling_target_only);
|
|
||||||
|
|
||||||
m.def(
|
|
||||||
"verify_tree_greedy(Tensor! predicts, Tensor! accept_index, Tensor! accept_token_num, "
|
|
||||||
"Tensor candidates, Tensor retrive_index, Tensor retrive_next_token, Tensor retrive_next_sibling, "
|
|
||||||
"Tensor target_predict) -> ()");
|
|
||||||
m.impl("verify_tree_greedy", torch::kCUDA, &verify_tree_greedy);
|
|
||||||
|
|
||||||
m.def(
|
|
||||||
"reconstruct_indices_from_tree_mask(Tensor tree_mask, Tensor verified_seq_len, Tensor positions, "
|
|
||||||
"Tensor retrive_index, Tensor retrive_next_token, Tensor retrive_next_sibling, "
|
|
||||||
"int batch_size, int draft_token_num) -> ()");
|
|
||||||
m.impl("reconstruct_indices_from_tree_mask", torch::kCUDA, &reconstruct_indices_from_tree_mask);
|
|
||||||
|
|
||||||
m.def(
|
|
||||||
"build_tree_kernel_efficient(Tensor parent_list, Tensor selected_index, Tensor verified_seq_len, "
|
|
||||||
"Tensor! tree_mask, Tensor! positions, Tensor! retrive_index, Tensor! retrive_next_token, "
|
|
||||||
"Tensor! retrive_next_sibling, int topk, int depth, int draft_token_num, int tree_mask_mode) -> "
|
|
||||||
"()");
|
|
||||||
m.impl("build_tree_kernel_efficient", torch::kCUDA, &build_tree_kernel_efficient);
|
|
||||||
|
|
||||||
m.def(
|
|
||||||
"segment_packbits(Tensor x, Tensor input_indptr, Tensor output_indptr, Tensor! y, int batch_size, "
|
|
||||||
"int cuda_stream) -> ()");
|
|
||||||
m.impl("segment_packbits", torch::kCUDA, &segment_packbits);
|
|
||||||
|
|
||||||
/*
|
/*
|
||||||
* From csrc/kvcacheio
|
* From csrc/kvcacheio
|
||||||
*/
|
*/
|
||||||
|
|||||||
@@ -148,22 +148,6 @@ TORCH_LIBRARY_EXPAND(sgl_kernel, m) {
|
|||||||
"correction_bias) -> ()");
|
"correction_bias) -> ()");
|
||||||
m.impl("topk_sigmoid", torch::kCUDA, &topk_sigmoid);
|
m.impl("topk_sigmoid", torch::kCUDA, &topk_sigmoid);
|
||||||
|
|
||||||
/*
|
|
||||||
* From csrc/speculative
|
|
||||||
*/
|
|
||||||
m.def(
|
|
||||||
"verify_tree_greedy(Tensor! predicts, Tensor! accept_index, Tensor! accept_token_num, "
|
|
||||||
"Tensor candidates, Tensor retrive_index, Tensor retrive_next_token, Tensor retrive_next_sibling, "
|
|
||||||
"Tensor target_predict) -> ()");
|
|
||||||
m.impl("verify_tree_greedy", torch::kCUDA, &verify_tree_greedy);
|
|
||||||
|
|
||||||
m.def(
|
|
||||||
"build_tree_kernel_efficient(Tensor parent_list, Tensor selected_index, Tensor verified_seq_len, "
|
|
||||||
"Tensor! tree_mask, Tensor! positions, Tensor! retrive_index, Tensor! retrive_next_token, "
|
|
||||||
"Tensor! retrive_next_sibling, int topk, int depth, int draft_token_num, int tree_mask_mode) -> "
|
|
||||||
"()");
|
|
||||||
m.impl("build_tree_kernel_efficient", torch::kCUDA, &build_tree_kernel_efficient);
|
|
||||||
|
|
||||||
/*
|
/*
|
||||||
* From csrc/kvcacheio
|
* From csrc/kvcacheio
|
||||||
*/
|
*/
|
||||||
|
|||||||
@@ -23,193 +23,12 @@
|
|||||||
#include "pytorch_extension_utils_rocm.h"
|
#include "pytorch_extension_utils_rocm.h"
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
typedef enum { FULL_MASK = 0, QLEN_ONLY = 1, QLEN_ONLY_BITPACKING = 2 } TreeMaskMode;
|
#include <sgl_kernel/speculative/eagle.cuh>
|
||||||
|
|
||||||
// parent_list [bs, topk * (depth - 1) + 1)]
|
using sglang::speculative::build_tree_efficient;
|
||||||
// selected_index [bs, draft_token_num - 1]
|
using sglang::speculative::build_tree_efficient_partial_packed;
|
||||||
// verified_seq_len [bs]
|
using sglang::speculative::QLEN_ONLY_BITPACKING;
|
||||||
// tree_mask [draft_token*(seq_len[0]+draft_token) | draft_token*(seq_len[1]+draft_token) | ..] =
|
using sglang::speculative::VerifyTreeGreedy;
|
||||||
// [sum(verified_seq_len)*draft_token+bs*draft_token*draft_token] positions [bs * draft_token] retrive_index [b,
|
|
||||||
// draft_token] retrive_next_token [b, draft_token] retrive_next_sibling [b, draft_token]
|
|
||||||
__global__ void build_tree_efficient(
|
|
||||||
int64_t* parent_list,
|
|
||||||
int64_t* selected_index,
|
|
||||||
int64_t* verified_seq_len,
|
|
||||||
bool* tree_mask,
|
|
||||||
int64_t* positions,
|
|
||||||
int64_t* retrive_index,
|
|
||||||
int64_t* retrive_next_token,
|
|
||||||
int64_t* retrive_next_sibling,
|
|
||||||
int topk,
|
|
||||||
int depth,
|
|
||||||
int draft_token_num,
|
|
||||||
int tree_mask_mode) {
|
|
||||||
int bid = blockIdx.x;
|
|
||||||
int tid = threadIdx.x;
|
|
||||||
|
|
||||||
if (tid >= draft_token_num) {
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
int seq_tree_idx = draft_token_num * draft_token_num * bid;
|
|
||||||
for (int i = 0; i < bid; i++) {
|
|
||||||
seq_tree_idx += verified_seq_len[i] * draft_token_num;
|
|
||||||
}
|
|
||||||
int seq_len = verified_seq_len[bid];
|
|
||||||
int token_tree_idx;
|
|
||||||
if (tree_mask_mode == FULL_MASK) {
|
|
||||||
token_tree_idx = seq_tree_idx + (seq_len + draft_token_num) * tid + seq_len + 1;
|
|
||||||
} else {
|
|
||||||
token_tree_idx = draft_token_num * draft_token_num * bid + draft_token_num * tid + 1;
|
|
||||||
}
|
|
||||||
tree_mask[token_tree_idx - 1] = true;
|
|
||||||
for (int i = 0; i < draft_token_num - 1; i++) {
|
|
||||||
tree_mask[token_tree_idx + i] = false;
|
|
||||||
}
|
|
||||||
|
|
||||||
int position = 0;
|
|
||||||
if (tid == 0) {
|
|
||||||
positions[bid * draft_token_num] = seq_len;
|
|
||||||
|
|
||||||
int retrive_index_offset = bid * draft_token_num;
|
|
||||||
for (int i = draft_token_num - 1; i > 0; --i) {
|
|
||||||
int current_token_idx = retrive_index_offset + i;
|
|
||||||
retrive_index[bid * draft_token_num + i] = current_token_idx;
|
|
||||||
int parent_tb_idx = selected_index[bid * (draft_token_num - 1) + i - 1] / topk;
|
|
||||||
int parent_position = 0;
|
|
||||||
if (parent_tb_idx > 0) {
|
|
||||||
int parent_token_idx = parent_list[bid * (topk * (depth - 1) + 1) + parent_tb_idx];
|
|
||||||
for (; parent_position < draft_token_num; ++parent_position) {
|
|
||||||
if (selected_index[bid * (draft_token_num - 1) + parent_position] == parent_token_idx) {
|
|
||||||
++parent_position;
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if (parent_position == draft_token_num) {
|
|
||||||
printf(
|
|
||||||
"WARNING: invalid eagle tree!!! Detected a token with no parent token selected. "
|
|
||||||
"Please check if the logprob has nan. The token will be ignored to keep proceeding.\n");
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
|
|
||||||
if (retrive_next_token[bid * draft_token_num + parent_position] == -1) {
|
|
||||||
retrive_next_token[bid * draft_token_num + parent_position] = i;
|
|
||||||
} else {
|
|
||||||
int origin_next_token = retrive_next_token[bid * draft_token_num + parent_position];
|
|
||||||
retrive_next_token[bid * draft_token_num + parent_position] = i;
|
|
||||||
retrive_next_sibling[bid * draft_token_num + i] = origin_next_token;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
retrive_index[bid * draft_token_num] = bid * draft_token_num;
|
|
||||||
} else {
|
|
||||||
int cur_position = tid - 1;
|
|
||||||
while (true) {
|
|
||||||
position += 1;
|
|
||||||
tree_mask[token_tree_idx + cur_position] = true;
|
|
||||||
int parent_tb_idx = selected_index[bid * (draft_token_num - 1) + cur_position] / topk;
|
|
||||||
if (parent_tb_idx == 0) {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
|
|
||||||
int token_idx = parent_list[bid * (topk * (depth - 1) + 1) + parent_tb_idx];
|
|
||||||
for (cur_position = 0; cur_position < draft_token_num; ++cur_position) {
|
|
||||||
if (selected_index[bid * (draft_token_num - 1) + cur_position] == token_idx) {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
positions[bid * draft_token_num + tid] = position + seq_len;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// parent_list [bs, topk * (depth - 1) + 1)]
|
|
||||||
// selected_index [bs, draft_token_num - 1]
|
|
||||||
// verified_seq_len [bs]
|
|
||||||
// tree_mask: [draft_token*num_bytes_per_item | .. ] = [bs*draft_token*num_bytes_per_item]
|
|
||||||
// positions [bs * draft_token]
|
|
||||||
// retrive_index [bs, draft_token]
|
|
||||||
// retrive_next_token [bs, draft_token]
|
|
||||||
// retrive_next_sibling [bs, draft_token]
|
|
||||||
__global__ void build_tree_efficient_partial_packed(
|
|
||||||
int64_t* parent_list,
|
|
||||||
int64_t* selected_index,
|
|
||||||
int64_t* verified_seq_len,
|
|
||||||
uint8_t* tree_mask,
|
|
||||||
int64_t* positions,
|
|
||||||
int64_t* retrive_index,
|
|
||||||
int64_t* retrive_next_token,
|
|
||||||
int64_t* retrive_next_sibling,
|
|
||||||
int topk,
|
|
||||||
int depth,
|
|
||||||
int draft_token_num,
|
|
||||||
size_t num_bytes_per_item) {
|
|
||||||
int bid = blockIdx.x;
|
|
||||||
int tid = threadIdx.x;
|
|
||||||
|
|
||||||
if (tid >= draft_token_num) {
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
int seq_len = verified_seq_len[bid];
|
|
||||||
int token_tree_idx = (bid * draft_token_num + tid) * num_bytes_per_item;
|
|
||||||
tree_mask[token_tree_idx] = 1; // little endian
|
|
||||||
|
|
||||||
int position = 0;
|
|
||||||
if (tid == 0) {
|
|
||||||
positions[bid * draft_token_num] = seq_len;
|
|
||||||
|
|
||||||
int retrive_index_offset = bid * draft_token_num;
|
|
||||||
for (int i = draft_token_num - 1; i > 0; --i) {
|
|
||||||
int current_token_idx = retrive_index_offset + i;
|
|
||||||
retrive_index[bid * draft_token_num + i] = current_token_idx;
|
|
||||||
int parent_tb_idx = selected_index[bid * (draft_token_num - 1) + i - 1] / topk;
|
|
||||||
int parent_position = 0;
|
|
||||||
if (parent_tb_idx > 0) {
|
|
||||||
int parent_token_idx = parent_list[bid * (topk * (depth - 1) + 1) + parent_tb_idx];
|
|
||||||
for (; parent_position < draft_token_num; ++parent_position) {
|
|
||||||
if (selected_index[bid * (draft_token_num - 1) + parent_position] == parent_token_idx) {
|
|
||||||
++parent_position;
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if (parent_position == draft_token_num) {
|
|
||||||
printf(
|
|
||||||
"WARNING: invalid eagle tree!!! Detected a token with no parent token selected. "
|
|
||||||
"Please check if the logprob has nan. The token will be ignored to keep proceeding.\n");
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
|
|
||||||
if (retrive_next_token[bid * draft_token_num + parent_position] == -1) {
|
|
||||||
retrive_next_token[bid * draft_token_num + parent_position] = i;
|
|
||||||
} else {
|
|
||||||
int origin_next_token = retrive_next_token[bid * draft_token_num + parent_position];
|
|
||||||
retrive_next_token[bid * draft_token_num + parent_position] = i;
|
|
||||||
retrive_next_sibling[bid * draft_token_num + i] = origin_next_token;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
retrive_index[bid * draft_token_num] = bid * draft_token_num;
|
|
||||||
} else {
|
|
||||||
int cur_position = tid - 1;
|
|
||||||
while (true) {
|
|
||||||
position += 1;
|
|
||||||
int byte_idx = (cur_position + 1) / 8;
|
|
||||||
int bit_idx = (cur_position + 1) % 8;
|
|
||||||
tree_mask[token_tree_idx + byte_idx] |= (1 << bit_idx);
|
|
||||||
int parent_tb_idx = selected_index[bid * (draft_token_num - 1) + cur_position] / topk;
|
|
||||||
if (parent_tb_idx == 0) {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
|
|
||||||
int token_idx = parent_list[bid * (topk * (depth - 1) + 1) + parent_tb_idx];
|
|
||||||
for (cur_position = 0; cur_position < draft_token_num; ++cur_position) {
|
|
||||||
if (selected_index[bid * (draft_token_num - 1) + cur_position] == token_idx) {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
positions[bid * draft_token_num + tid] = position + seq_len;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
void build_tree_kernel_efficient(
|
void build_tree_kernel_efficient(
|
||||||
at::Tensor parent_list,
|
at::Tensor parent_list,
|
||||||
@@ -268,50 +87,6 @@ void build_tree_kernel_efficient(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
template <typename IdType, typename IdType2>
|
|
||||||
__global__ void VerifyTreeGreedy(
|
|
||||||
IdType* predicts,
|
|
||||||
IdType* accept_index,
|
|
||||||
IdType* accept_token_num, // mutable
|
|
||||||
IdType2* candidates,
|
|
||||||
IdType2* retrive_index,
|
|
||||||
IdType2* retrive_next_token,
|
|
||||||
IdType2* retrive_next_sibling,
|
|
||||||
IdType2* target_predict,
|
|
||||||
uint32_t batch_size,
|
|
||||||
uint32_t num_speculative_tokens,
|
|
||||||
uint32_t num_draft_tokens) {
|
|
||||||
uint32_t bx = blockIdx.x;
|
|
||||||
|
|
||||||
IdType2 last_accepted_retrive_idx = retrive_index[bx * num_draft_tokens];
|
|
||||||
accept_index[bx * num_speculative_tokens] = last_accepted_retrive_idx;
|
|
||||||
uint32_t num_accepted_tokens = 0;
|
|
||||||
IdType2 cur_index = 0;
|
|
||||||
|
|
||||||
for (uint32_t j = 1; j < num_speculative_tokens; ++j) {
|
|
||||||
cur_index = retrive_next_token[bx * num_draft_tokens + cur_index];
|
|
||||||
while (cur_index != -1) {
|
|
||||||
IdType2 draft_index = retrive_index[bx * num_draft_tokens + cur_index];
|
|
||||||
IdType2 draft_token_id = candidates[bx * num_draft_tokens + cur_index];
|
|
||||||
IdType2 target_token_id = target_predict[last_accepted_retrive_idx];
|
|
||||||
|
|
||||||
if (draft_token_id == target_token_id) {
|
|
||||||
// accept token
|
|
||||||
predicts[last_accepted_retrive_idx] = target_token_id;
|
|
||||||
++num_accepted_tokens;
|
|
||||||
accept_index[bx * num_speculative_tokens + num_accepted_tokens] = draft_index;
|
|
||||||
last_accepted_retrive_idx = draft_index;
|
|
||||||
break;
|
|
||||||
} else {
|
|
||||||
cur_index = retrive_next_sibling[bx * num_draft_tokens + cur_index];
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if (cur_index == -1) break;
|
|
||||||
}
|
|
||||||
accept_token_num[bx] = num_accepted_tokens;
|
|
||||||
predicts[last_accepted_retrive_idx] = target_predict[last_accepted_retrive_idx];
|
|
||||||
}
|
|
||||||
|
|
||||||
// predicts: [tot_num_draft_tokens]
|
// predicts: [tot_num_draft_tokens]
|
||||||
// accept_index: [bs, num_spec_step]
|
// accept_index: [bs, num_spec_step]
|
||||||
// accept_token_num: [bs]
|
// accept_token_num: [bs]
|
||||||
|
|||||||
@@ -7,78 +7,9 @@
|
|||||||
#include "pytorch_extension_utils_rocm.h"
|
#include "pytorch_extension_utils_rocm.h"
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
// tree_mask: [bs * draft_token_num * draft_token_num]
|
#include <sgl_kernel/speculative/ngram.cuh>
|
||||||
// verified_seq_len: [bs]
|
|
||||||
// positions: [bs * draft_token_num]
|
|
||||||
// retrive_index: [bs, draft_token_num]
|
|
||||||
// retrive_next_token: [bs, draft_token_num]
|
|
||||||
// retrive_next_sibling: [bs, draft_token_num]
|
|
||||||
__global__ void reconstructIndicesFromTreeMask(
|
|
||||||
bool* tree_mask,
|
|
||||||
int64_t* verified_seq_len,
|
|
||||||
int64_t* positions,
|
|
||||||
int64_t* retrive_index,
|
|
||||||
int64_t* retrive_next_token,
|
|
||||||
int64_t* retrive_next_sibling,
|
|
||||||
int batch_size,
|
|
||||||
int draft_token_num) {
|
|
||||||
int bid = blockIdx.x;
|
|
||||||
int tid = threadIdx.x;
|
|
||||||
|
|
||||||
if (bid >= batch_size || tid >= draft_token_num) {
|
using sglang::speculative::reconstructIndicesFromTreeMask;
|
||||||
return;
|
|
||||||
}
|
|
||||||
int base_offset = draft_token_num * draft_token_num;
|
|
||||||
// token_idx: [bid * draft_token_num, (bid + 1) * draft_token_num)
|
|
||||||
int token_idx = bid * draft_token_num;
|
|
||||||
// tree_mask_idx: [bid * base_offset, (bid + 1) * base_offset)
|
|
||||||
int tree_mask_offset = bid * base_offset;
|
|
||||||
|
|
||||||
int depth = 0;
|
|
||||||
int parent_idx = -1;
|
|
||||||
|
|
||||||
for (int i = tid - 1, start_idx = tree_mask_offset + tid * draft_token_num; i >= 0; i--) {
|
|
||||||
if (tree_mask[start_idx + i]) {
|
|
||||||
depth++;
|
|
||||||
if (parent_idx == -1) {
|
|
||||||
parent_idx = i;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
retrive_index[token_idx + tid] = token_idx + tid;
|
|
||||||
positions[token_idx + tid] = depth + verified_seq_len[bid];
|
|
||||||
|
|
||||||
int next_token_idx = -1;
|
|
||||||
for (int i = tid + 1; i < draft_token_num; i++) {
|
|
||||||
if (tree_mask[tree_mask_offset + i * draft_token_num + tid]) {
|
|
||||||
next_token_idx = i;
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
retrive_next_token[token_idx + tid] = next_token_idx;
|
|
||||||
|
|
||||||
int next_sibling_idx = -1;
|
|
||||||
if (parent_idx != -1) {
|
|
||||||
for (int i = tid + 1; i < draft_token_num; i++) {
|
|
||||||
int start_idx = tree_mask_offset + i * draft_token_num + parent_idx;
|
|
||||||
if (tree_mask[start_idx]) {
|
|
||||||
bool is_sibling = true;
|
|
||||||
int end_idx = tree_mask_offset + i * draft_token_num + i;
|
|
||||||
for (int j = start_idx + 1; j < end_idx; ++j) {
|
|
||||||
if (tree_mask[j]) {
|
|
||||||
is_sibling = false;
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if (is_sibling) {
|
|
||||||
next_sibling_idx = i;
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
retrive_next_sibling[token_idx + tid] = next_sibling_idx;
|
|
||||||
}
|
|
||||||
|
|
||||||
void reconstruct_indices_from_tree_mask(
|
void reconstruct_indices_from_tree_mask(
|
||||||
at::Tensor tree_mask,
|
at::Tensor tree_mask,
|
||||||
|
|||||||
@@ -19,7 +19,7 @@
|
|||||||
|
|
||||||
#include <assert.h>
|
#include <assert.h>
|
||||||
|
|
||||||
#include <flashinfer/sampling.cuh>
|
#include <sgl_kernel/speculative/sampling.cuh>
|
||||||
|
|
||||||
namespace flashinfer {
|
namespace flashinfer {
|
||||||
|
|
||||||
@@ -27,152 +27,6 @@ namespace sampling {
|
|||||||
|
|
||||||
using namespace cub;
|
using namespace cub;
|
||||||
|
|
||||||
template <
|
|
||||||
uint32_t BLOCK_THREADS,
|
|
||||||
BlockScanAlgorithm SCAN_ALGORITHM,
|
|
||||||
BlockReduceAlgorithm REDUCE_ALGORITHM,
|
|
||||||
uint32_t VEC_SIZE,
|
|
||||||
bool DETERMINISTIC,
|
|
||||||
typename DType,
|
|
||||||
typename IdType,
|
|
||||||
typename IdType2>
|
|
||||||
__global__ void TreeSpeculativeSamplingTargetOnly(
|
|
||||||
IdType* predicts, // mutable
|
|
||||||
IdType* accept_index, // mutable
|
|
||||||
IdType* accept_token_num, // mutable
|
|
||||||
IdType2* candidates,
|
|
||||||
IdType2* retrive_index,
|
|
||||||
IdType2* retrive_next_token,
|
|
||||||
IdType2* retrive_next_sibling,
|
|
||||||
DType* uniform_samples,
|
|
||||||
DType* uniform_samples_for_final_sampling,
|
|
||||||
DType* target_probs,
|
|
||||||
DType* draft_probs,
|
|
||||||
uint32_t batch_size,
|
|
||||||
uint32_t num_speculative_tokens,
|
|
||||||
uint32_t num_draft_tokens,
|
|
||||||
uint32_t d,
|
|
||||||
DType threshold_single,
|
|
||||||
DType threshold_acc) {
|
|
||||||
const uint32_t bx = blockIdx.x, tx = threadIdx.x;
|
|
||||||
|
|
||||||
extern __shared__ __align__(alignof(SamplingTempStorage<BLOCK_THREADS, SCAN_ALGORITHM, REDUCE_ALGORITHM>))
|
|
||||||
uint8_t smem_sampling[];
|
|
||||||
auto& temp_storage =
|
|
||||||
reinterpret_cast<SamplingTempStorage<BLOCK_THREADS, SCAN_ALGORITHM, REDUCE_ALGORITHM>&>(smem_sampling);
|
|
||||||
|
|
||||||
DType prob_acc = 0.0;
|
|
||||||
uint32_t cur_prob_offset = bx * num_draft_tokens * d;
|
|
||||||
DType coin = uniform_samples[bx * num_draft_tokens];
|
|
||||||
IdType2 last_accepted_retrive_idx = retrive_index[bx * num_draft_tokens];
|
|
||||||
accept_index[bx * num_speculative_tokens] = last_accepted_retrive_idx;
|
|
||||||
uint32_t num_accepted_tokens = 0;
|
|
||||||
IdType2 cur_index = 0;
|
|
||||||
|
|
||||||
for (uint32_t j = 1; j < num_speculative_tokens; ++j) {
|
|
||||||
cur_index = retrive_next_token[bx * num_draft_tokens + cur_index];
|
|
||||||
while (cur_index != -1) {
|
|
||||||
IdType2 draft_index = retrive_index[bx * num_draft_tokens + cur_index];
|
|
||||||
IdType2 draft_token_id = candidates[bx * num_draft_tokens + cur_index];
|
|
||||||
DType target_prob_single = target_probs[cur_prob_offset + draft_token_id];
|
|
||||||
prob_acc += target_prob_single;
|
|
||||||
|
|
||||||
if (coin <= prob_acc / threshold_acc || target_prob_single >= threshold_single) {
|
|
||||||
// accept token
|
|
||||||
prob_acc = 0.;
|
|
||||||
cur_prob_offset = (bx * num_draft_tokens + cur_index) * d;
|
|
||||||
coin = uniform_samples[bx * num_draft_tokens + cur_index];
|
|
||||||
predicts[last_accepted_retrive_idx] = draft_token_id;
|
|
||||||
++num_accepted_tokens;
|
|
||||||
accept_index[bx * num_speculative_tokens + num_accepted_tokens] = draft_index;
|
|
||||||
last_accepted_retrive_idx = draft_index;
|
|
||||||
break;
|
|
||||||
} else {
|
|
||||||
// FIXME: leverage draft probs
|
|
||||||
draft_probs[cur_prob_offset + draft_token_id] = target_probs[cur_prob_offset + draft_token_id];
|
|
||||||
cur_index = retrive_next_sibling[bx * num_draft_tokens + cur_index];
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if (cur_index == -1) break;
|
|
||||||
}
|
|
||||||
accept_token_num[bx] = num_accepted_tokens;
|
|
||||||
|
|
||||||
// we need a different coin for the final sampling
|
|
||||||
coin = uniform_samples_for_final_sampling[bx];
|
|
||||||
|
|
||||||
// sample from relu(target_probs - draft_probs)
|
|
||||||
DType sum_relu_q_minus_p(0);
|
|
||||||
vec_t<DType, VEC_SIZE> q_vec, p_vec;
|
|
||||||
DType relu_q_minus_p[VEC_SIZE];
|
|
||||||
for (uint32_t i = 0; i < ceil_div(d, BLOCK_THREADS * VEC_SIZE); ++i) {
|
|
||||||
q_vec.fill(DType(0));
|
|
||||||
p_vec.fill(DType(0));
|
|
||||||
if ((i * BLOCK_THREADS + tx) * VEC_SIZE < d) {
|
|
||||||
q_vec.load(target_probs + cur_prob_offset + i * BLOCK_THREADS * VEC_SIZE + tx * VEC_SIZE);
|
|
||||||
if (num_accepted_tokens != num_speculative_tokens - 1) {
|
|
||||||
// there is no draft_probs for the bonus token
|
|
||||||
p_vec.load(draft_probs + cur_prob_offset + i * BLOCK_THREADS * VEC_SIZE + tx * VEC_SIZE);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
#pragma unroll
|
|
||||||
for (uint32_t j = 0; j < VEC_SIZE; ++j) {
|
|
||||||
relu_q_minus_p[j] = max(q_vec[j] - p_vec[j], DType(0));
|
|
||||||
}
|
|
||||||
sum_relu_q_minus_p += BlockReduce<DType, BLOCK_THREADS, REDUCE_ALGORITHM>(temp_storage.block_prim.reduce)
|
|
||||||
.Sum<VEC_SIZE>(relu_q_minus_p);
|
|
||||||
__syncthreads();
|
|
||||||
}
|
|
||||||
if (tx == 0) {
|
|
||||||
temp_storage.block_aggregate.value = sum_relu_q_minus_p;
|
|
||||||
}
|
|
||||||
|
|
||||||
temp_storage.sampled_id = d;
|
|
||||||
temp_storage.last_valid_id = -1;
|
|
||||||
__syncthreads();
|
|
||||||
sum_relu_q_minus_p = temp_storage.block_aggregate.value;
|
|
||||||
DType u = coin * sum_relu_q_minus_p;
|
|
||||||
|
|
||||||
DType aggregate_relu_q_minus_p(0);
|
|
||||||
for (uint32_t i = 0; i < ceil_div(d, BLOCK_THREADS * VEC_SIZE); ++i) {
|
|
||||||
q_vec.fill(DType(0));
|
|
||||||
p_vec.fill(DType(0));
|
|
||||||
if ((i * BLOCK_THREADS + tx) * VEC_SIZE < d) {
|
|
||||||
q_vec.load(target_probs + cur_prob_offset + i * BLOCK_THREADS * VEC_SIZE + tx * VEC_SIZE);
|
|
||||||
if (num_accepted_tokens != num_speculative_tokens - 1) {
|
|
||||||
// there is no draft_probs for the bonus token
|
|
||||||
p_vec.load(draft_probs + cur_prob_offset + i * BLOCK_THREADS * VEC_SIZE + tx * VEC_SIZE);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
vec_t<DType, VEC_SIZE> relu_q_minus_p_vec;
|
|
||||||
#pragma unroll
|
|
||||||
for (uint32_t j = 0; j < VEC_SIZE; ++j) {
|
|
||||||
relu_q_minus_p_vec[j] = max(q_vec[j] - p_vec[j], DType(0));
|
|
||||||
}
|
|
||||||
|
|
||||||
DeviceSamplingFromProb<VEC_SIZE, BLOCK_THREADS, SCAN_ALGORITHM, REDUCE_ALGORITHM, DETERMINISTIC>(
|
|
||||||
i, d, [&](DType x) { return x > 0; }, u, relu_q_minus_p_vec, aggregate_relu_q_minus_p, &temp_storage);
|
|
||||||
if (aggregate_relu_q_minus_p > u) {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
__syncthreads();
|
|
||||||
// This would happen when u is very close to 1
|
|
||||||
// and the sum of probabilities is smaller than u
|
|
||||||
// In this case, we use the last valid index as the sampled id
|
|
||||||
int sampled_id = temp_storage.sampled_id;
|
|
||||||
if (sampled_id == d) {
|
|
||||||
if (temp_storage.last_valid_id == -1) {
|
|
||||||
sampled_id = d - 1;
|
|
||||||
} else {
|
|
||||||
sampled_id = temp_storage.last_valid_id;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// set the first rejected token
|
|
||||||
predicts[last_accepted_retrive_idx] = sampled_id;
|
|
||||||
// value at not used indices are undefined
|
|
||||||
}
|
|
||||||
|
|
||||||
template <typename DType, typename IdType, typename IdType2>
|
template <typename DType, typename IdType, typename IdType2>
|
||||||
cudaError_t TreeSpeculativeSamplingTargetOnly(
|
cudaError_t TreeSpeculativeSamplingTargetOnly(
|
||||||
IdType* predicts, // mutable
|
IdType* predicts, // mutable
|
||||||
@@ -221,7 +75,7 @@ cudaError_t TreeSpeculativeSamplingTargetOnly(
|
|||||||
&capped_threshold_acc};
|
&capped_threshold_acc};
|
||||||
DISPATCH_ALIGNED_VEC_SIZE(
|
DISPATCH_ALIGNED_VEC_SIZE(
|
||||||
vec_size, VEC_SIZE, {DISPATCH_DETERMINISTIC(deterministic, DETERMINISTIC, {
|
vec_size, VEC_SIZE, {DISPATCH_DETERMINISTIC(deterministic, DETERMINISTIC, {
|
||||||
auto kernel = TreeSpeculativeSamplingTargetOnly<
|
auto kernel = sglang::speculative_sampling::TreeSpeculativeSamplingTargetOnly<
|
||||||
BLOCK_THREADS,
|
BLOCK_THREADS,
|
||||||
SCAN_ALGO,
|
SCAN_ALGO,
|
||||||
REDUCE_ALGO,
|
REDUCE_ALGO,
|
||||||
|
|||||||
@@ -106,17 +106,23 @@ else:
|
|||||||
assign_extend_cache_locs_cpu,
|
assign_extend_cache_locs_cpu,
|
||||||
assign_req_to_token_pool_cpu,
|
assign_req_to_token_pool_cpu,
|
||||||
build_draft_decode_metadata_cpu,
|
build_draft_decode_metadata_cpu,
|
||||||
build_tree_kernel_efficient,
|
|
||||||
build_tree_kernel_efficient_cpu,
|
build_tree_kernel_efficient_cpu,
|
||||||
fill_accept_out_cache_loc_cpu,
|
fill_accept_out_cache_loc_cpu,
|
||||||
fill_bonus_tokens_cpu,
|
fill_bonus_tokens_cpu,
|
||||||
reconstruct_indices_from_tree_mask,
|
|
||||||
rotate_input_ids_cpu,
|
rotate_input_ids_cpu,
|
||||||
|
verify_tree_greedy_cpu,
|
||||||
|
)
|
||||||
|
|
||||||
|
# CUDA and ROCm route these through sglang.kernels.ops.speculative (JIT).
|
||||||
|
if torch.version.cuda is None:
|
||||||
|
from sgl_kernel.speculative import (
|
||||||
|
build_tree_kernel_efficient,
|
||||||
|
reconstruct_indices_from_tree_mask,
|
||||||
segment_packbits,
|
segment_packbits,
|
||||||
tree_speculative_sampling_target_only,
|
tree_speculative_sampling_target_only,
|
||||||
verify_tree_greedy,
|
verify_tree_greedy,
|
||||||
verify_tree_greedy_cpu,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
from sgl_kernel.top_k import (
|
from sgl_kernel.top_k import (
|
||||||
fast_topk,
|
fast_topk,
|
||||||
fast_topk_transform_fused,
|
fast_topk_transform_fused,
|
||||||
|
|||||||
@@ -70,6 +70,9 @@ include_dirs = [
|
|||||||
root / "include",
|
root / "include",
|
||||||
root / "include" / "impl",
|
root / "include" / "impl",
|
||||||
root / "csrc",
|
root / "csrc",
|
||||||
|
# Speculative device kernels are shared with the JIT launchers, which own
|
||||||
|
# the header tree they live in.
|
||||||
|
root.parent / "jit" / "include",
|
||||||
root / _FLASHINFER_REPO.source_dir / "include",
|
root / _FLASHINFER_REPO.source_dir / "include",
|
||||||
root / _FLASHINFER_REPO.source_dir / "csrc",
|
root / _FLASHINFER_REPO.source_dir / "csrc",
|
||||||
root / _MUTLASS_REPO.source_dir / "include",
|
root / _MUTLASS_REPO.source_dir / "include",
|
||||||
|
|||||||
@@ -54,7 +54,6 @@ sources = [
|
|||||||
"csrc/moe/moe_align_kernel.cu",
|
"csrc/moe/moe_align_kernel.cu",
|
||||||
"csrc/moe/moe_topk_softmax_kernels.cu",
|
"csrc/moe/moe_topk_softmax_kernels.cu",
|
||||||
"csrc/moe/moe_topk_sigmoid_kernels.cu",
|
"csrc/moe/moe_topk_sigmoid_kernels.cu",
|
||||||
"csrc/speculative/eagle_utils.cu",
|
|
||||||
"csrc/kvcacheio/transfer.cu",
|
"csrc/kvcacheio/transfer.cu",
|
||||||
"csrc/memory/weak_ref_tensor.cpp",
|
"csrc/memory/weak_ref_tensor.cpp",
|
||||||
"csrc/elementwise/pos_enc.cu",
|
"csrc/elementwise/pos_enc.cu",
|
||||||
|
|||||||
@@ -0,0 +1,112 @@
|
|||||||
|
/*
|
||||||
|
* Copyright (c) 2025 by SGLang team.
|
||||||
|
* Copyright (c) 2024-2025 by FlashInfer team.
|
||||||
|
*
|
||||||
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
* you may not use this file except in compliance with the License.
|
||||||
|
* You may obtain a copy of the License at
|
||||||
|
*
|
||||||
|
* http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
*
|
||||||
|
* Unless required by applicable law or agreed to in writing, software
|
||||||
|
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
* See the License for the specific language governing permissions and
|
||||||
|
* limitations under the License.
|
||||||
|
*/
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
#include <sgl_kernel/tensor.h>
|
||||||
|
|
||||||
|
#include <sgl_kernel/utils.cuh>
|
||||||
|
|
||||||
|
#include <sgl_kernel/speculative/sampling.cuh>
|
||||||
|
|
||||||
|
#include <flashinfer/sampling.cuh>
|
||||||
|
|
||||||
|
#include <numeric>
|
||||||
|
|
||||||
|
namespace sglang {
|
||||||
|
|
||||||
|
/// \brief Sample and verify a draft tree using only target probabilities.
|
||||||
|
template <uint32_t VecSize, bool Deterministic>
|
||||||
|
void tree_speculative_sampling_target_only(
|
||||||
|
tvm::ffi::TensorView predicts,
|
||||||
|
tvm::ffi::TensorView accept_index,
|
||||||
|
tvm::ffi::TensorView accept_token_num,
|
||||||
|
tvm::ffi::TensorView candidates,
|
||||||
|
tvm::ffi::TensorView retrive_index,
|
||||||
|
tvm::ffi::TensorView retrive_next_token,
|
||||||
|
tvm::ffi::TensorView retrive_next_sibling,
|
||||||
|
tvm::ffi::TensorView uniform_samples,
|
||||||
|
tvm::ffi::TensorView uniform_samples_for_final_sampling,
|
||||||
|
tvm::ffi::TensorView target_probs,
|
||||||
|
tvm::ffi::TensorView draft_probs,
|
||||||
|
double threshold_single,
|
||||||
|
double threshold_acc) {
|
||||||
|
using namespace host;
|
||||||
|
using namespace flashinfer::sampling;
|
||||||
|
static_assert(VecSize == 1 || VecSize == 2 || VecSize == 4);
|
||||||
|
SymbolicSize batch_size{"batch_size"}, draft_tokens{"draft_tokens"}, spec_tokens{"spec_tokens"},
|
||||||
|
vocab_size{"vocab_size"};
|
||||||
|
SymbolicDevice device;
|
||||||
|
TensorMatcher({batch_size, draft_tokens})
|
||||||
|
.with_dtype<int64_t>()
|
||||||
|
.with_device<kDLCUDA>(device)
|
||||||
|
.verify(candidates)
|
||||||
|
.verify(retrive_index)
|
||||||
|
.verify(retrive_next_token)
|
||||||
|
.verify(retrive_next_sibling);
|
||||||
|
TensorMatcher({batch_size, spec_tokens}).with_dtype<int32_t>().with_device(device).verify(accept_index);
|
||||||
|
TensorMatcher({batch_size}).with_dtype<int32_t>().with_device(device).verify(accept_token_num);
|
||||||
|
TensorMatcher({batch_size.unwrap() * draft_tokens.unwrap()})
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device(device)
|
||||||
|
.verify(predicts);
|
||||||
|
TensorMatcher({batch_size, draft_tokens}).with_dtype<float>().with_device(device).verify(uniform_samples);
|
||||||
|
TensorMatcher({batch_size}).with_dtype<float>().with_device(device).verify(uniform_samples_for_final_sampling);
|
||||||
|
TensorMatcher({batch_size, draft_tokens, vocab_size})
|
||||||
|
.with_dtype<float>()
|
||||||
|
.with_device(device)
|
||||||
|
.verify(target_probs)
|
||||||
|
.verify(draft_probs);
|
||||||
|
CHECK_HOST(draft_tokens.unwrap() > 0 && spec_tokens.unwrap() > 0 && vocab_size.unwrap() > 0);
|
||||||
|
CHECK_HOST(std::gcd(int64_t{4}, vocab_size.unwrap()) == VecSize);
|
||||||
|
CHECK_HOST(threshold_single >= 0 && threshold_single <= 1);
|
||||||
|
CHECK_HOST(threshold_acc >= 0 && threshold_acc <= 1);
|
||||||
|
if (batch_size.unwrap() == 0) return;
|
||||||
|
|
||||||
|
constexpr uint32_t block_threads = 1024;
|
||||||
|
constexpr size_t smem_size = sizeof(SamplingTempStorage<block_threads, SCAN_ALGO, REDUCE_ALGO>);
|
||||||
|
auto kernel = speculative_sampling::TreeSpeculativeSamplingTargetOnly<
|
||||||
|
block_threads,
|
||||||
|
SCAN_ALGO,
|
||||||
|
REDUCE_ALGO,
|
||||||
|
VecSize,
|
||||||
|
Deterministic,
|
||||||
|
float,
|
||||||
|
int32_t,
|
||||||
|
int64_t>;
|
||||||
|
CHECK_CUDA(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
|
||||||
|
LaunchKernel(static_cast<uint32_t>(batch_size.unwrap()), block_threads, device.unwrap(), smem_size)(
|
||||||
|
kernel,
|
||||||
|
static_cast<int32_t*>(predicts.data_ptr()),
|
||||||
|
static_cast<int32_t*>(accept_index.data_ptr()),
|
||||||
|
static_cast<int32_t*>(accept_token_num.data_ptr()),
|
||||||
|
static_cast<int64_t*>(candidates.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_index.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_next_token.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_next_sibling.data_ptr()),
|
||||||
|
static_cast<float*>(uniform_samples.data_ptr()),
|
||||||
|
static_cast<float*>(uniform_samples_for_final_sampling.data_ptr()),
|
||||||
|
static_cast<float*>(target_probs.data_ptr()),
|
||||||
|
static_cast<float*>(draft_probs.data_ptr()),
|
||||||
|
static_cast<uint32_t>(batch_size.unwrap()),
|
||||||
|
static_cast<uint32_t>(spec_tokens.unwrap()),
|
||||||
|
static_cast<uint32_t>(draft_tokens.unwrap()),
|
||||||
|
static_cast<uint32_t>(vocab_size.unwrap()),
|
||||||
|
static_cast<float>(threshold_single),
|
||||||
|
std::max(static_cast<float>(threshold_acc), 1e-9f));
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace sglang
|
||||||
@@ -0,0 +1,210 @@
|
|||||||
|
/*
|
||||||
|
* Copyright (c) 2025 by SGLang team.
|
||||||
|
*
|
||||||
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
* you may not use this file except in compliance with the License.
|
||||||
|
* You may obtain a copy of the License at
|
||||||
|
*
|
||||||
|
* http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
*
|
||||||
|
* Unless required by applicable law or agreed to in writing, software
|
||||||
|
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
* See the License for the specific language governing permissions and
|
||||||
|
* limitations under the License.
|
||||||
|
*/
|
||||||
|
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
#include <sgl_kernel/tensor.h>
|
||||||
|
|
||||||
|
#include <sgl_kernel/utils.cuh>
|
||||||
|
|
||||||
|
#include <sgl_kernel/speculative/eagle.cuh>
|
||||||
|
#include <sgl_kernel/speculative/ngram.cuh>
|
||||||
|
|
||||||
|
namespace sglang {
|
||||||
|
|
||||||
|
/// \brief Build the EAGLE tree and its retrieval links in preallocated buffers.
|
||||||
|
inline void build_tree_kernel_efficient(
|
||||||
|
tvm::ffi::TensorView parent_list,
|
||||||
|
tvm::ffi::TensorView selected_index,
|
||||||
|
tvm::ffi::TensorView verified_seq_len,
|
||||||
|
tvm::ffi::TensorView tree_mask,
|
||||||
|
tvm::ffi::TensorView positions,
|
||||||
|
tvm::ffi::TensorView retrive_index,
|
||||||
|
tvm::ffi::TensorView retrive_next_token,
|
||||||
|
tvm::ffi::TensorView retrive_next_sibling,
|
||||||
|
int64_t topk,
|
||||||
|
int64_t depth,
|
||||||
|
int64_t draft_token_num,
|
||||||
|
int64_t tree_mask_mode) {
|
||||||
|
using namespace host;
|
||||||
|
SymbolicSize batch_size{"batch_size"}, parent_width{"parent_width"};
|
||||||
|
SymbolicDevice device;
|
||||||
|
CHECK_HOST(topk > 0 && depth > 0);
|
||||||
|
CHECK_HOST(draft_token_num > 0 && draft_token_num <= 1024);
|
||||||
|
CHECK_HOST(tree_mask_mode >= 0 && tree_mask_mode <= 2);
|
||||||
|
TensorMatcher({batch_size, parent_width})
|
||||||
|
.with_strides({parent_list.size(1) == 0 ? -1 : parent_list.size(1), 1})
|
||||||
|
.with_dtype<int64_t>()
|
||||||
|
.with_device<kDLCUDA>(device)
|
||||||
|
.verify(parent_list);
|
||||||
|
CHECK_HOST(depth == 1 || parent_width.unwrap() == topk * (depth - 1) + 1);
|
||||||
|
TensorMatcher({batch_size, draft_token_num - 1})
|
||||||
|
.with_strides({draft_token_num == 1 ? -1 : draft_token_num - 1, 1})
|
||||||
|
.with_dtype<int64_t>()
|
||||||
|
.with_device(device)
|
||||||
|
.verify(selected_index);
|
||||||
|
TensorMatcher({batch_size}).with_dtype<int64_t>().with_device(device).verify(verified_seq_len);
|
||||||
|
TensorMatcher({batch_size, draft_token_num})
|
||||||
|
.with_dtype<int64_t>()
|
||||||
|
.with_device(device)
|
||||||
|
.verify(retrive_index)
|
||||||
|
.verify(retrive_next_token)
|
||||||
|
.verify(retrive_next_sibling);
|
||||||
|
const int64_t bs = batch_size.unwrap();
|
||||||
|
TensorMatcher({bs * draft_token_num}).with_dtype<int64_t>().with_device(device).verify(positions);
|
||||||
|
// The mask is raw bytes to these kernels, and callers disagree on how they
|
||||||
|
// spell that: bool for an inline allocation or the default preallocated
|
||||||
|
// buffer, uint8 for the triton backend's, uint{8,16,32} for a bit-packed one.
|
||||||
|
// So constrain its BYTES, never its element type. The bound is a lower one --
|
||||||
|
// a buffer preallocated for the captured max batch is larger than what this
|
||||||
|
// batch writes, and FULL_MASK spans each request's context length so it has
|
||||||
|
// no host-side bound at all.
|
||||||
|
size_t num_bytes_per_item = 1;
|
||||||
|
int64_t min_mask_bytes = 0;
|
||||||
|
if (tree_mask_mode == speculative::QLEN_ONLY_BITPACKING) {
|
||||||
|
CHECK_HOST(draft_token_num <= 32);
|
||||||
|
// Width comes from draft_token_num alone, exactly as the AOT launcher does.
|
||||||
|
num_bytes_per_item = draft_token_num > 16 ? 4 : (draft_token_num > 8 ? 2 : 1);
|
||||||
|
min_mask_bytes = bs * draft_token_num * static_cast<int64_t>(num_bytes_per_item);
|
||||||
|
} else if (tree_mask_mode == speculative::QLEN_ONLY) {
|
||||||
|
min_mask_bytes = bs * draft_token_num * draft_token_num;
|
||||||
|
}
|
||||||
|
// -1 leaves the extent unconstrained; the byte bound below is the real guard.
|
||||||
|
TensorMatcher({-1}).with_device(device).verify(tree_mask);
|
||||||
|
const int64_t mask_elem_bytes = tree_mask.dtype().bits / 8;
|
||||||
|
CHECK_HOST(tree_mask.dtype().bits % 8 == 0 && mask_elem_bytes > 0);
|
||||||
|
CHECK_HOST(tree_mask.numel() * mask_elem_bytes >= min_mask_bytes);
|
||||||
|
if (bs == 0) return;
|
||||||
|
|
||||||
|
auto launch = LaunchKernel(static_cast<uint32_t>(bs), static_cast<uint32_t>(draft_token_num), device.unwrap());
|
||||||
|
if (tree_mask_mode == speculative::QLEN_ONLY_BITPACKING) {
|
||||||
|
launch(
|
||||||
|
speculative::build_tree_efficient_partial_packed,
|
||||||
|
static_cast<int64_t*>(parent_list.data_ptr()),
|
||||||
|
static_cast<int64_t*>(selected_index.data_ptr()),
|
||||||
|
static_cast<int64_t*>(verified_seq_len.data_ptr()),
|
||||||
|
static_cast<uint8_t*>(tree_mask.data_ptr()),
|
||||||
|
static_cast<int64_t*>(positions.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_index.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_next_token.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_next_sibling.data_ptr()),
|
||||||
|
static_cast<int32_t>(topk),
|
||||||
|
static_cast<int32_t>(depth),
|
||||||
|
static_cast<int32_t>(draft_token_num),
|
||||||
|
num_bytes_per_item);
|
||||||
|
} else {
|
||||||
|
launch(
|
||||||
|
speculative::build_tree_efficient,
|
||||||
|
static_cast<int64_t*>(parent_list.data_ptr()),
|
||||||
|
static_cast<int64_t*>(selected_index.data_ptr()),
|
||||||
|
static_cast<int64_t*>(verified_seq_len.data_ptr()),
|
||||||
|
static_cast<bool*>(tree_mask.data_ptr()),
|
||||||
|
static_cast<int64_t*>(positions.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_index.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_next_token.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_next_sibling.data_ptr()),
|
||||||
|
static_cast<int32_t>(topk),
|
||||||
|
static_cast<int32_t>(depth),
|
||||||
|
static_cast<int32_t>(draft_token_num),
|
||||||
|
static_cast<int32_t>(tree_mask_mode));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// \brief Verify a draft tree against greedy target predictions.
|
||||||
|
inline void verify_tree_greedy(
|
||||||
|
tvm::ffi::TensorView predicts,
|
||||||
|
tvm::ffi::TensorView accept_index,
|
||||||
|
tvm::ffi::TensorView accept_token_num,
|
||||||
|
tvm::ffi::TensorView candidates,
|
||||||
|
tvm::ffi::TensorView retrive_index,
|
||||||
|
tvm::ffi::TensorView retrive_next_token,
|
||||||
|
tvm::ffi::TensorView retrive_next_sibling,
|
||||||
|
tvm::ffi::TensorView target_predict) {
|
||||||
|
using namespace host;
|
||||||
|
SymbolicSize batch_size{"batch_size"}, draft_tokens{"draft_tokens"}, spec_tokens{"spec_tokens"};
|
||||||
|
SymbolicDevice device;
|
||||||
|
TensorMatcher({batch_size, draft_tokens})
|
||||||
|
.with_dtype<int64_t>()
|
||||||
|
.with_device<kDLCUDA>(device)
|
||||||
|
.verify(candidates)
|
||||||
|
.verify(retrive_index)
|
||||||
|
.verify(retrive_next_token)
|
||||||
|
.verify(retrive_next_sibling)
|
||||||
|
.verify(target_predict);
|
||||||
|
TensorMatcher({batch_size, spec_tokens}).with_dtype<int32_t>().with_device(device).verify(accept_index);
|
||||||
|
TensorMatcher({batch_size}).with_dtype<int32_t>().with_device(device).verify(accept_token_num);
|
||||||
|
TensorMatcher({batch_size.unwrap() * draft_tokens.unwrap()})
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device(device)
|
||||||
|
.verify(predicts);
|
||||||
|
CHECK_HOST(draft_tokens.unwrap() > 0 && spec_tokens.unwrap() > 0);
|
||||||
|
if (batch_size.unwrap() == 0) return;
|
||||||
|
LaunchKernel(static_cast<uint32_t>(batch_size.unwrap()), 1, device.unwrap())(
|
||||||
|
speculative::VerifyTreeGreedy<int32_t, int64_t>,
|
||||||
|
static_cast<int32_t*>(predicts.data_ptr()),
|
||||||
|
static_cast<int32_t*>(accept_index.data_ptr()),
|
||||||
|
static_cast<int32_t*>(accept_token_num.data_ptr()),
|
||||||
|
static_cast<int64_t*>(candidates.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_index.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_next_token.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_next_sibling.data_ptr()),
|
||||||
|
static_cast<int64_t*>(target_predict.data_ptr()),
|
||||||
|
static_cast<uint32_t>(batch_size.unwrap()),
|
||||||
|
static_cast<uint32_t>(spec_tokens.unwrap()),
|
||||||
|
static_cast<uint32_t>(draft_tokens.unwrap()));
|
||||||
|
}
|
||||||
|
|
||||||
|
/// \brief Reconstruct retrieval links and positions from an ngram tree mask.
|
||||||
|
inline void reconstruct_indices_from_tree_mask(
|
||||||
|
tvm::ffi::TensorView tree_mask,
|
||||||
|
tvm::ffi::TensorView verified_seq_len,
|
||||||
|
tvm::ffi::TensorView positions,
|
||||||
|
tvm::ffi::TensorView retrive_index,
|
||||||
|
tvm::ffi::TensorView retrive_next_token,
|
||||||
|
tvm::ffi::TensorView retrive_next_sibling,
|
||||||
|
int64_t batch_size,
|
||||||
|
int64_t draft_token_num) {
|
||||||
|
using namespace host;
|
||||||
|
SymbolicDevice device;
|
||||||
|
CHECK_HOST(batch_size >= 0 && draft_token_num > 0 && draft_token_num <= 1024);
|
||||||
|
// Bytes, not element type -- same reasoning as build_tree_kernel_efficient:
|
||||||
|
// the kernel casts straight to bool* and callers are free to spell a 1-byte
|
||||||
|
// mask as bool or uint8.
|
||||||
|
TensorMatcher({-1}).with_device<kDLCUDA>(device).verify(tree_mask);
|
||||||
|
CHECK_HOST(tree_mask.dtype().bits == 8);
|
||||||
|
CHECK_HOST(tree_mask.numel() >= batch_size * draft_token_num * draft_token_num);
|
||||||
|
TensorMatcher({batch_size}).with_dtype<int64_t>().with_device(device).verify(verified_seq_len);
|
||||||
|
TensorMatcher({batch_size * draft_token_num}).with_dtype<int64_t>().with_device(device).verify(positions);
|
||||||
|
TensorMatcher({batch_size, draft_token_num})
|
||||||
|
.with_dtype<int64_t>()
|
||||||
|
.with_device(device)
|
||||||
|
.verify(retrive_index)
|
||||||
|
.verify(retrive_next_token)
|
||||||
|
.verify(retrive_next_sibling);
|
||||||
|
if (batch_size == 0) return;
|
||||||
|
LaunchKernel(static_cast<uint32_t>(batch_size), static_cast<uint32_t>(draft_token_num), device.unwrap())(
|
||||||
|
speculative::reconstructIndicesFromTreeMask,
|
||||||
|
static_cast<bool*>(tree_mask.data_ptr()),
|
||||||
|
static_cast<int64_t*>(verified_seq_len.data_ptr()),
|
||||||
|
static_cast<int64_t*>(positions.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_index.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_next_token.data_ptr()),
|
||||||
|
static_cast<int64_t*>(retrive_next_sibling.data_ptr()),
|
||||||
|
static_cast<int32_t>(batch_size),
|
||||||
|
static_cast<int32_t>(draft_token_num));
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace sglang
|
||||||
@@ -0,0 +1,259 @@
|
|||||||
|
/*
|
||||||
|
* Copyright (c) 2025 by SGLang team.
|
||||||
|
*
|
||||||
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
* you may not use this file except in compliance with the License.
|
||||||
|
* You may obtain a copy of the License at
|
||||||
|
*
|
||||||
|
* http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
*
|
||||||
|
* Unless required by applicable law or agreed to in writing, software
|
||||||
|
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
* See the License for the specific language governing permissions and
|
||||||
|
* limitations under the License.
|
||||||
|
*/
|
||||||
|
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
#include <cstdint>
|
||||||
|
#include <cstdio>
|
||||||
|
|
||||||
|
namespace sglang {
|
||||||
|
namespace speculative {
|
||||||
|
|
||||||
|
typedef enum { FULL_MASK = 0, QLEN_ONLY = 1, QLEN_ONLY_BITPACKING = 2 } TreeMaskMode;
|
||||||
|
|
||||||
|
// parent_list [bs, topk * (depth - 1) + 1)]
|
||||||
|
// selected_index [bs, draft_token_num - 1]
|
||||||
|
// verified_seq_len [bs]
|
||||||
|
// tree_mask [draft_token*(seq_len[0]+draft_token) | draft_token*(seq_len[1]+draft_token) | ..] =
|
||||||
|
// [sum(verified_seq_len)*draft_token+bs*draft_token*draft_token] positions [bs * draft_token] retrive_index [b,
|
||||||
|
// draft_token] retrive_next_token [b, draft_token] retrive_next_sibling [b, draft_token]
|
||||||
|
__global__ void build_tree_efficient(
|
||||||
|
int64_t* parent_list,
|
||||||
|
int64_t* selected_index,
|
||||||
|
int64_t* verified_seq_len,
|
||||||
|
bool* tree_mask,
|
||||||
|
int64_t* positions,
|
||||||
|
int64_t* retrive_index,
|
||||||
|
int64_t* retrive_next_token,
|
||||||
|
int64_t* retrive_next_sibling,
|
||||||
|
int topk,
|
||||||
|
int depth,
|
||||||
|
int draft_token_num,
|
||||||
|
int tree_mask_mode) {
|
||||||
|
int bid = blockIdx.x;
|
||||||
|
int tid = threadIdx.x;
|
||||||
|
|
||||||
|
if (tid >= draft_token_num) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
int seq_tree_idx = draft_token_num * draft_token_num * bid;
|
||||||
|
for (int i = 0; i < bid; i++) {
|
||||||
|
seq_tree_idx += verified_seq_len[i] * draft_token_num;
|
||||||
|
}
|
||||||
|
int seq_len = verified_seq_len[bid];
|
||||||
|
int token_tree_idx;
|
||||||
|
if (tree_mask_mode == FULL_MASK) {
|
||||||
|
token_tree_idx = seq_tree_idx + (seq_len + draft_token_num) * tid + seq_len + 1;
|
||||||
|
} else {
|
||||||
|
token_tree_idx = draft_token_num * draft_token_num * bid + draft_token_num * tid + 1;
|
||||||
|
}
|
||||||
|
tree_mask[token_tree_idx - 1] = true;
|
||||||
|
for (int i = 0; i < draft_token_num - 1; i++) {
|
||||||
|
tree_mask[token_tree_idx + i] = false;
|
||||||
|
}
|
||||||
|
|
||||||
|
int position = 0;
|
||||||
|
if (tid == 0) {
|
||||||
|
positions[bid * draft_token_num] = seq_len;
|
||||||
|
|
||||||
|
int retrive_index_offset = bid * draft_token_num;
|
||||||
|
for (int i = draft_token_num - 1; i > 0; --i) {
|
||||||
|
int current_token_idx = retrive_index_offset + i;
|
||||||
|
retrive_index[bid * draft_token_num + i] = current_token_idx;
|
||||||
|
int parent_tb_idx = selected_index[bid * (draft_token_num - 1) + i - 1] / topk;
|
||||||
|
int parent_position = 0;
|
||||||
|
if (parent_tb_idx > 0) {
|
||||||
|
int parent_token_idx = parent_list[bid * (topk * (depth - 1) + 1) + parent_tb_idx];
|
||||||
|
for (; parent_position < draft_token_num; ++parent_position) {
|
||||||
|
if (selected_index[bid * (draft_token_num - 1) + parent_position] == parent_token_idx) {
|
||||||
|
++parent_position;
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (parent_position == draft_token_num) {
|
||||||
|
printf(
|
||||||
|
"WARNING: invalid eagle tree!!! Detected a token with no parent token selected. "
|
||||||
|
"Please check if the logprob has nan. The token will be ignored to keep proceeding.\n");
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (retrive_next_token[bid * draft_token_num + parent_position] == -1) {
|
||||||
|
retrive_next_token[bid * draft_token_num + parent_position] = i;
|
||||||
|
} else {
|
||||||
|
int origin_next_token = retrive_next_token[bid * draft_token_num + parent_position];
|
||||||
|
retrive_next_token[bid * draft_token_num + parent_position] = i;
|
||||||
|
retrive_next_sibling[bid * draft_token_num + i] = origin_next_token;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
retrive_index[bid * draft_token_num] = bid * draft_token_num;
|
||||||
|
} else {
|
||||||
|
int cur_position = tid - 1;
|
||||||
|
while (true) {
|
||||||
|
position += 1;
|
||||||
|
tree_mask[token_tree_idx + cur_position] = true;
|
||||||
|
int parent_tb_idx = selected_index[bid * (draft_token_num - 1) + cur_position] / topk;
|
||||||
|
if (parent_tb_idx == 0) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
int token_idx = parent_list[bid * (topk * (depth - 1) + 1) + parent_tb_idx];
|
||||||
|
for (cur_position = 0; cur_position < draft_token_num; ++cur_position) {
|
||||||
|
if (selected_index[bid * (draft_token_num - 1) + cur_position] == token_idx) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
positions[bid * draft_token_num + tid] = position + seq_len;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// parent_list [bs, topk * (depth - 1) + 1)]
|
||||||
|
// selected_index [bs, draft_token_num - 1]
|
||||||
|
// verified_seq_len [bs]
|
||||||
|
// tree_mask: [draft_token*num_bytes_per_item | .. ] = [bs*draft_token*num_bytes_per_item]
|
||||||
|
// positions [bs * draft_token]
|
||||||
|
// retrive_index [bs, draft_token]
|
||||||
|
// retrive_next_token [bs, draft_token]
|
||||||
|
// retrive_next_sibling [bs, draft_token]
|
||||||
|
__global__ void build_tree_efficient_partial_packed(
|
||||||
|
int64_t* parent_list,
|
||||||
|
int64_t* selected_index,
|
||||||
|
int64_t* verified_seq_len,
|
||||||
|
uint8_t* tree_mask,
|
||||||
|
int64_t* positions,
|
||||||
|
int64_t* retrive_index,
|
||||||
|
int64_t* retrive_next_token,
|
||||||
|
int64_t* retrive_next_sibling,
|
||||||
|
int topk,
|
||||||
|
int depth,
|
||||||
|
int draft_token_num,
|
||||||
|
size_t num_bytes_per_item) {
|
||||||
|
int bid = blockIdx.x;
|
||||||
|
int tid = threadIdx.x;
|
||||||
|
|
||||||
|
if (tid >= draft_token_num) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
int seq_len = verified_seq_len[bid];
|
||||||
|
int token_tree_idx = (bid * draft_token_num + tid) * num_bytes_per_item;
|
||||||
|
tree_mask[token_tree_idx] = 1; // little endian
|
||||||
|
|
||||||
|
int position = 0;
|
||||||
|
if (tid == 0) {
|
||||||
|
positions[bid * draft_token_num] = seq_len;
|
||||||
|
|
||||||
|
int retrive_index_offset = bid * draft_token_num;
|
||||||
|
for (int i = draft_token_num - 1; i > 0; --i) {
|
||||||
|
int current_token_idx = retrive_index_offset + i;
|
||||||
|
retrive_index[bid * draft_token_num + i] = current_token_idx;
|
||||||
|
int parent_tb_idx = selected_index[bid * (draft_token_num - 1) + i - 1] / topk;
|
||||||
|
int parent_position = 0;
|
||||||
|
if (parent_tb_idx > 0) {
|
||||||
|
int parent_token_idx = parent_list[bid * (topk * (depth - 1) + 1) + parent_tb_idx];
|
||||||
|
for (; parent_position < draft_token_num; ++parent_position) {
|
||||||
|
if (selected_index[bid * (draft_token_num - 1) + parent_position] == parent_token_idx) {
|
||||||
|
++parent_position;
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (parent_position == draft_token_num) {
|
||||||
|
printf(
|
||||||
|
"WARNING: invalid eagle tree!!! Detected a token with no parent token selected. "
|
||||||
|
"Please check if the logprob has nan. The token will be ignored to keep proceeding.\n");
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (retrive_next_token[bid * draft_token_num + parent_position] == -1) {
|
||||||
|
retrive_next_token[bid * draft_token_num + parent_position] = i;
|
||||||
|
} else {
|
||||||
|
int origin_next_token = retrive_next_token[bid * draft_token_num + parent_position];
|
||||||
|
retrive_next_token[bid * draft_token_num + parent_position] = i;
|
||||||
|
retrive_next_sibling[bid * draft_token_num + i] = origin_next_token;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
retrive_index[bid * draft_token_num] = bid * draft_token_num;
|
||||||
|
} else {
|
||||||
|
int cur_position = tid - 1;
|
||||||
|
while (true) {
|
||||||
|
position += 1;
|
||||||
|
int byte_idx = (cur_position + 1) / 8;
|
||||||
|
int bit_idx = (cur_position + 1) % 8;
|
||||||
|
tree_mask[token_tree_idx + byte_idx] |= (1 << bit_idx);
|
||||||
|
int parent_tb_idx = selected_index[bid * (draft_token_num - 1) + cur_position] / topk;
|
||||||
|
if (parent_tb_idx == 0) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
int token_idx = parent_list[bid * (topk * (depth - 1) + 1) + parent_tb_idx];
|
||||||
|
for (cur_position = 0; cur_position < draft_token_num; ++cur_position) {
|
||||||
|
if (selected_index[bid * (draft_token_num - 1) + cur_position] == token_idx) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
positions[bid * draft_token_num + tid] = position + seq_len;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
template <typename IdType, typename IdType2>
|
||||||
|
__global__ void VerifyTreeGreedy(
|
||||||
|
IdType* predicts,
|
||||||
|
IdType* accept_index,
|
||||||
|
IdType* accept_token_num, // mutable
|
||||||
|
IdType2* candidates,
|
||||||
|
IdType2* retrive_index,
|
||||||
|
IdType2* retrive_next_token,
|
||||||
|
IdType2* retrive_next_sibling,
|
||||||
|
IdType2* target_predict,
|
||||||
|
uint32_t batch_size,
|
||||||
|
uint32_t num_speculative_tokens,
|
||||||
|
uint32_t num_draft_tokens) {
|
||||||
|
uint32_t bx = blockIdx.x;
|
||||||
|
|
||||||
|
IdType2 last_accepted_retrive_idx = retrive_index[bx * num_draft_tokens];
|
||||||
|
accept_index[bx * num_speculative_tokens] = last_accepted_retrive_idx;
|
||||||
|
uint32_t num_accepted_tokens = 0;
|
||||||
|
IdType2 cur_index = 0;
|
||||||
|
|
||||||
|
for (uint32_t j = 1; j < num_speculative_tokens; ++j) {
|
||||||
|
cur_index = retrive_next_token[bx * num_draft_tokens + cur_index];
|
||||||
|
while (cur_index != -1) {
|
||||||
|
IdType2 draft_index = retrive_index[bx * num_draft_tokens + cur_index];
|
||||||
|
IdType2 draft_token_id = candidates[bx * num_draft_tokens + cur_index];
|
||||||
|
IdType2 target_token_id = target_predict[last_accepted_retrive_idx];
|
||||||
|
|
||||||
|
if (draft_token_id == target_token_id) {
|
||||||
|
// accept token
|
||||||
|
predicts[last_accepted_retrive_idx] = target_token_id;
|
||||||
|
++num_accepted_tokens;
|
||||||
|
accept_index[bx * num_speculative_tokens + num_accepted_tokens] = draft_index;
|
||||||
|
last_accepted_retrive_idx = draft_index;
|
||||||
|
break;
|
||||||
|
} else {
|
||||||
|
cur_index = retrive_next_sibling[bx * num_draft_tokens + cur_index];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (cur_index == -1) break;
|
||||||
|
}
|
||||||
|
accept_token_num[bx] = num_accepted_tokens;
|
||||||
|
predicts[last_accepted_retrive_idx] = target_predict[last_accepted_retrive_idx];
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace speculative
|
||||||
|
} // namespace sglang
|
||||||
@@ -0,0 +1,94 @@
|
|||||||
|
/*
|
||||||
|
* Copyright (c) 2025 by SGLang team.
|
||||||
|
*
|
||||||
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
* you may not use this file except in compliance with the License.
|
||||||
|
* You may obtain a copy of the License at
|
||||||
|
*
|
||||||
|
* http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
*
|
||||||
|
* Unless required by applicable law or agreed to in writing, software
|
||||||
|
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
* See the License for the specific language governing permissions and
|
||||||
|
* limitations under the License.
|
||||||
|
*/
|
||||||
|
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
#include <cstdint>
|
||||||
|
|
||||||
|
namespace sglang {
|
||||||
|
namespace speculative {
|
||||||
|
|
||||||
|
__global__ void reconstructIndicesFromTreeMask(
|
||||||
|
bool* tree_mask,
|
||||||
|
int64_t* verified_seq_len,
|
||||||
|
int64_t* positions,
|
||||||
|
int64_t* retrive_index,
|
||||||
|
int64_t* retrive_next_token,
|
||||||
|
int64_t* retrive_next_sibling,
|
||||||
|
int batch_size,
|
||||||
|
int draft_token_num) {
|
||||||
|
int bid = blockIdx.x;
|
||||||
|
int tid = threadIdx.x;
|
||||||
|
|
||||||
|
if (bid >= batch_size || tid >= draft_token_num) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
int base_offset = draft_token_num * draft_token_num;
|
||||||
|
// token_idx: [bid * draft_token_num, (bid + 1) * draft_token_num)
|
||||||
|
int token_idx = bid * draft_token_num;
|
||||||
|
// tree_mask_idx: [bid * base_offset, (bid + 1) * base_offset)
|
||||||
|
int tree_mask_offset = bid * base_offset;
|
||||||
|
|
||||||
|
int depth = 0;
|
||||||
|
int parent_idx = -1;
|
||||||
|
|
||||||
|
for (int i = tid - 1, start_idx = tree_mask_offset + tid * draft_token_num; i >= 0; i--) {
|
||||||
|
if (tree_mask[start_idx + i]) {
|
||||||
|
depth++;
|
||||||
|
if (parent_idx == -1) {
|
||||||
|
parent_idx = i;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
retrive_index[token_idx + tid] = token_idx + tid;
|
||||||
|
positions[token_idx + tid] = depth + verified_seq_len[bid];
|
||||||
|
|
||||||
|
int next_token_idx = -1;
|
||||||
|
for (int i = tid + 1; i < draft_token_num; i++) {
|
||||||
|
if (tree_mask[tree_mask_offset + i * draft_token_num + tid]) {
|
||||||
|
next_token_idx = i;
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
retrive_next_token[token_idx + tid] = next_token_idx;
|
||||||
|
|
||||||
|
int next_sibling_idx = -1;
|
||||||
|
if (parent_idx != -1) {
|
||||||
|
for (int i = tid + 1; i < draft_token_num; i++) {
|
||||||
|
int start_idx = tree_mask_offset + i * draft_token_num + parent_idx;
|
||||||
|
if (tree_mask[start_idx]) {
|
||||||
|
bool is_sibling = true;
|
||||||
|
int end_idx = tree_mask_offset + i * draft_token_num + i;
|
||||||
|
for (int j = start_idx + 1; j < end_idx; ++j) {
|
||||||
|
if (tree_mask[j]) {
|
||||||
|
is_sibling = false;
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (is_sibling) {
|
||||||
|
next_sibling_idx = i;
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
retrive_next_sibling[token_idx + tid] = next_sibling_idx;
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace speculative
|
||||||
|
} // namespace sglang
|
||||||
@@ -0,0 +1,177 @@
|
|||||||
|
/*
|
||||||
|
* Copyright (c) 2025 by SGLang team.
|
||||||
|
* Copyright (c) 2024-2025 by FlashInfer team.
|
||||||
|
*
|
||||||
|
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
* you may not use this file except in compliance with the License.
|
||||||
|
* You may obtain a copy of the License at
|
||||||
|
*
|
||||||
|
* http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
*
|
||||||
|
* Unless required by applicable law or agreed to in writing, software
|
||||||
|
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
* See the License for the specific language governing permissions and
|
||||||
|
* limitations under the License.
|
||||||
|
*/
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
#pragma once
|
||||||
|
|
||||||
|
#include <flashinfer/sampling.cuh>
|
||||||
|
|
||||||
|
namespace sglang {
|
||||||
|
namespace speculative_sampling {
|
||||||
|
|
||||||
|
using namespace flashinfer;
|
||||||
|
using namespace flashinfer::sampling;
|
||||||
|
using namespace cub;
|
||||||
|
|
||||||
|
template <
|
||||||
|
uint32_t BLOCK_THREADS,
|
||||||
|
BlockScanAlgorithm SCAN_ALGORITHM,
|
||||||
|
BlockReduceAlgorithm REDUCE_ALGORITHM,
|
||||||
|
uint32_t VEC_SIZE,
|
||||||
|
bool DETERMINISTIC,
|
||||||
|
typename DType,
|
||||||
|
typename IdType,
|
||||||
|
typename IdType2>
|
||||||
|
__global__ void TreeSpeculativeSamplingTargetOnly(
|
||||||
|
IdType* predicts, // mutable
|
||||||
|
IdType* accept_index, // mutable
|
||||||
|
IdType* accept_token_num, // mutable
|
||||||
|
IdType2* candidates,
|
||||||
|
IdType2* retrive_index,
|
||||||
|
IdType2* retrive_next_token,
|
||||||
|
IdType2* retrive_next_sibling,
|
||||||
|
DType* uniform_samples,
|
||||||
|
DType* uniform_samples_for_final_sampling,
|
||||||
|
DType* target_probs,
|
||||||
|
DType* draft_probs,
|
||||||
|
uint32_t batch_size,
|
||||||
|
uint32_t num_speculative_tokens,
|
||||||
|
uint32_t num_draft_tokens,
|
||||||
|
uint32_t d,
|
||||||
|
DType threshold_single,
|
||||||
|
DType threshold_acc) {
|
||||||
|
const uint32_t bx = blockIdx.x, tx = threadIdx.x;
|
||||||
|
|
||||||
|
extern __shared__ __align__(alignof(SamplingTempStorage<BLOCK_THREADS, SCAN_ALGORITHM, REDUCE_ALGORITHM>))
|
||||||
|
uint8_t smem_sampling[];
|
||||||
|
auto& temp_storage =
|
||||||
|
reinterpret_cast<SamplingTempStorage<BLOCK_THREADS, SCAN_ALGORITHM, REDUCE_ALGORITHM>&>(smem_sampling);
|
||||||
|
|
||||||
|
DType prob_acc = 0.0;
|
||||||
|
uint32_t cur_prob_offset = bx * num_draft_tokens * d;
|
||||||
|
DType coin = uniform_samples[bx * num_draft_tokens];
|
||||||
|
IdType2 last_accepted_retrive_idx = retrive_index[bx * num_draft_tokens];
|
||||||
|
accept_index[bx * num_speculative_tokens] = last_accepted_retrive_idx;
|
||||||
|
uint32_t num_accepted_tokens = 0;
|
||||||
|
IdType2 cur_index = 0;
|
||||||
|
|
||||||
|
for (uint32_t j = 1; j < num_speculative_tokens; ++j) {
|
||||||
|
cur_index = retrive_next_token[bx * num_draft_tokens + cur_index];
|
||||||
|
while (cur_index != -1) {
|
||||||
|
IdType2 draft_index = retrive_index[bx * num_draft_tokens + cur_index];
|
||||||
|
IdType2 draft_token_id = candidates[bx * num_draft_tokens + cur_index];
|
||||||
|
DType target_prob_single = target_probs[cur_prob_offset + draft_token_id];
|
||||||
|
prob_acc += target_prob_single;
|
||||||
|
|
||||||
|
if (coin <= prob_acc / threshold_acc || target_prob_single >= threshold_single) {
|
||||||
|
// accept token
|
||||||
|
prob_acc = 0.;
|
||||||
|
cur_prob_offset = (bx * num_draft_tokens + cur_index) * d;
|
||||||
|
coin = uniform_samples[bx * num_draft_tokens + cur_index];
|
||||||
|
predicts[last_accepted_retrive_idx] = draft_token_id;
|
||||||
|
++num_accepted_tokens;
|
||||||
|
accept_index[bx * num_speculative_tokens + num_accepted_tokens] = draft_index;
|
||||||
|
last_accepted_retrive_idx = draft_index;
|
||||||
|
break;
|
||||||
|
} else {
|
||||||
|
// FIXME: leverage draft probs
|
||||||
|
draft_probs[cur_prob_offset + draft_token_id] = target_probs[cur_prob_offset + draft_token_id];
|
||||||
|
cur_index = retrive_next_sibling[bx * num_draft_tokens + cur_index];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (cur_index == -1) break;
|
||||||
|
}
|
||||||
|
accept_token_num[bx] = num_accepted_tokens;
|
||||||
|
|
||||||
|
// we need a different coin for the final sampling
|
||||||
|
coin = uniform_samples_for_final_sampling[bx];
|
||||||
|
|
||||||
|
// sample from relu(target_probs - draft_probs)
|
||||||
|
DType sum_relu_q_minus_p(0);
|
||||||
|
vec_t<DType, VEC_SIZE> q_vec, p_vec;
|
||||||
|
DType relu_q_minus_p[VEC_SIZE];
|
||||||
|
for (uint32_t i = 0; i < ceil_div(d, BLOCK_THREADS * VEC_SIZE); ++i) {
|
||||||
|
q_vec.fill(DType(0));
|
||||||
|
p_vec.fill(DType(0));
|
||||||
|
if ((i * BLOCK_THREADS + tx) * VEC_SIZE < d) {
|
||||||
|
q_vec.load(target_probs + cur_prob_offset + i * BLOCK_THREADS * VEC_SIZE + tx * VEC_SIZE);
|
||||||
|
if (num_accepted_tokens != num_speculative_tokens - 1) {
|
||||||
|
// there is no draft_probs for the bonus token
|
||||||
|
p_vec.load(draft_probs + cur_prob_offset + i * BLOCK_THREADS * VEC_SIZE + tx * VEC_SIZE);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
#pragma unroll
|
||||||
|
for (uint32_t j = 0; j < VEC_SIZE; ++j) {
|
||||||
|
relu_q_minus_p[j] = max(q_vec[j] - p_vec[j], DType(0));
|
||||||
|
}
|
||||||
|
sum_relu_q_minus_p += BlockReduce<DType, BLOCK_THREADS, REDUCE_ALGORITHM>(temp_storage.block_prim.reduce)
|
||||||
|
.Sum<VEC_SIZE>(relu_q_minus_p);
|
||||||
|
__syncthreads();
|
||||||
|
}
|
||||||
|
if (tx == 0) {
|
||||||
|
temp_storage.block_aggregate.value = sum_relu_q_minus_p;
|
||||||
|
}
|
||||||
|
|
||||||
|
temp_storage.sampled_id = d;
|
||||||
|
temp_storage.last_valid_id = -1;
|
||||||
|
__syncthreads();
|
||||||
|
sum_relu_q_minus_p = temp_storage.block_aggregate.value;
|
||||||
|
DType u = coin * sum_relu_q_minus_p;
|
||||||
|
|
||||||
|
DType aggregate_relu_q_minus_p(0);
|
||||||
|
for (uint32_t i = 0; i < ceil_div(d, BLOCK_THREADS * VEC_SIZE); ++i) {
|
||||||
|
q_vec.fill(DType(0));
|
||||||
|
p_vec.fill(DType(0));
|
||||||
|
if ((i * BLOCK_THREADS + tx) * VEC_SIZE < d) {
|
||||||
|
q_vec.load(target_probs + cur_prob_offset + i * BLOCK_THREADS * VEC_SIZE + tx * VEC_SIZE);
|
||||||
|
if (num_accepted_tokens != num_speculative_tokens - 1) {
|
||||||
|
// there is no draft_probs for the bonus token
|
||||||
|
p_vec.load(draft_probs + cur_prob_offset + i * BLOCK_THREADS * VEC_SIZE + tx * VEC_SIZE);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
vec_t<DType, VEC_SIZE> relu_q_minus_p_vec;
|
||||||
|
#pragma unroll
|
||||||
|
for (uint32_t j = 0; j < VEC_SIZE; ++j) {
|
||||||
|
relu_q_minus_p_vec[j] = max(q_vec[j] - p_vec[j], DType(0));
|
||||||
|
}
|
||||||
|
|
||||||
|
DeviceSamplingFromProb<VEC_SIZE, BLOCK_THREADS, SCAN_ALGORITHM, REDUCE_ALGORITHM, DETERMINISTIC>(
|
||||||
|
i, d, [&](DType x) { return x > 0; }, u, relu_q_minus_p_vec, aggregate_relu_q_minus_p, &temp_storage);
|
||||||
|
if (aggregate_relu_q_minus_p > u) {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
// This would happen when u is very close to 1
|
||||||
|
// and the sum of probabilities is smaller than u
|
||||||
|
// In this case, we use the last valid index as the sampled id
|
||||||
|
int sampled_id = temp_storage.sampled_id;
|
||||||
|
if (sampled_id == d) {
|
||||||
|
if (temp_storage.last_valid_id == -1) {
|
||||||
|
sampled_id = d - 1;
|
||||||
|
} else {
|
||||||
|
sampled_id = temp_storage.last_valid_id;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// set the first rejected token
|
||||||
|
predicts[last_accepted_retrive_idx] = sampled_id;
|
||||||
|
// value at not used indices are undefined
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace speculative_sampling
|
||||||
|
} // namespace sglang
|
||||||
@@ -0,0 +1,65 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import math
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from sglang.kernels.jit.utils import cache_once, load_jit
|
||||||
|
from sglang.kernels.kernel_api_logging import debug_kernel_api
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import torch
|
||||||
|
from tvm_ffi.module import Module
|
||||||
|
|
||||||
|
|
||||||
|
@cache_once
|
||||||
|
def _jit_sampling_module(vec_size: int, deterministic: bool) -> Module:
|
||||||
|
deterministic_cpp = "true" if deterministic else "false"
|
||||||
|
return load_jit(
|
||||||
|
"speculative_sampling",
|
||||||
|
vec_size,
|
||||||
|
deterministic,
|
||||||
|
cuda_files=["speculative/sampling.cuh"],
|
||||||
|
cuda_wrappers=[
|
||||||
|
(
|
||||||
|
"sample",
|
||||||
|
f"tree_speculative_sampling_target_only<{vec_size}, {deterministic_cpp}>",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
extra_dependencies=["flashinfer"],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@debug_kernel_api
|
||||||
|
def tree_speculative_sampling_target_only(
|
||||||
|
predicts: torch.Tensor, # mutable
|
||||||
|
accept_index: torch.Tensor, # mutable
|
||||||
|
accept_token_num: torch.Tensor, # mutable
|
||||||
|
candidates: torch.Tensor,
|
||||||
|
retrive_index: torch.Tensor,
|
||||||
|
retrive_next_token: torch.Tensor,
|
||||||
|
retrive_next_sibling: torch.Tensor,
|
||||||
|
uniform_samples: torch.Tensor,
|
||||||
|
uniform_samples_for_final_sampling: torch.Tensor,
|
||||||
|
target_probs: torch.Tensor,
|
||||||
|
draft_probs: torch.Tensor,
|
||||||
|
threshold_single: float = 1.0,
|
||||||
|
threshold_acc: float = 1.0,
|
||||||
|
deterministic: bool = True,
|
||||||
|
) -> None:
|
||||||
|
vec_size = math.gcd(4, target_probs.shape[-1])
|
||||||
|
module = _jit_sampling_module(vec_size, deterministic)
|
||||||
|
module.sample(
|
||||||
|
predicts,
|
||||||
|
accept_index,
|
||||||
|
accept_token_num,
|
||||||
|
candidates,
|
||||||
|
retrive_index,
|
||||||
|
retrive_next_token,
|
||||||
|
retrive_next_sibling,
|
||||||
|
uniform_samples,
|
||||||
|
uniform_samples_for_final_sampling,
|
||||||
|
target_probs,
|
||||||
|
draft_probs,
|
||||||
|
threshold_single,
|
||||||
|
threshold_acc,
|
||||||
|
)
|
||||||
@@ -0,0 +1,103 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from sglang.kernels.jit.utils import cache_once, load_jit
|
||||||
|
from sglang.kernels.kernel_api_logging import debug_kernel_api
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import torch
|
||||||
|
from tvm_ffi.module import Module
|
||||||
|
|
||||||
|
|
||||||
|
@cache_once
|
||||||
|
def _jit_tree_module() -> Module:
|
||||||
|
return load_jit(
|
||||||
|
"speculative_tree",
|
||||||
|
cuda_files=["speculative/tree.cuh"],
|
||||||
|
cuda_wrappers=[
|
||||||
|
("build_tree_kernel_efficient", "build_tree_kernel_efficient"),
|
||||||
|
("verify_tree_greedy", "verify_tree_greedy"),
|
||||||
|
(
|
||||||
|
"reconstruct_indices_from_tree_mask",
|
||||||
|
"reconstruct_indices_from_tree_mask",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@debug_kernel_api
|
||||||
|
def build_tree_kernel_efficient(
|
||||||
|
parent_list: torch.Tensor,
|
||||||
|
selected_index: torch.Tensor,
|
||||||
|
verified_seq_len: torch.Tensor,
|
||||||
|
tree_mask: torch.Tensor,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
retrive_index: torch.Tensor,
|
||||||
|
retrive_next_token: torch.Tensor,
|
||||||
|
retrive_next_sibling: torch.Tensor,
|
||||||
|
topk: int,
|
||||||
|
depth: int,
|
||||||
|
draft_token_num: int,
|
||||||
|
tree_mask_mode: int,
|
||||||
|
) -> None:
|
||||||
|
_jit_tree_module().build_tree_kernel_efficient(
|
||||||
|
parent_list,
|
||||||
|
selected_index,
|
||||||
|
verified_seq_len,
|
||||||
|
tree_mask,
|
||||||
|
positions,
|
||||||
|
retrive_index,
|
||||||
|
retrive_next_token,
|
||||||
|
retrive_next_sibling,
|
||||||
|
topk,
|
||||||
|
depth,
|
||||||
|
draft_token_num,
|
||||||
|
tree_mask_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@debug_kernel_api
|
||||||
|
def verify_tree_greedy(
|
||||||
|
predicts: torch.Tensor, # mutable
|
||||||
|
accept_index: torch.Tensor, # mutable
|
||||||
|
accept_token_num: torch.Tensor, # mutable
|
||||||
|
candidates: torch.Tensor,
|
||||||
|
retrive_index: torch.Tensor,
|
||||||
|
retrive_next_token: torch.Tensor,
|
||||||
|
retrive_next_sibling: torch.Tensor,
|
||||||
|
target_predict: torch.Tensor,
|
||||||
|
) -> None:
|
||||||
|
_jit_tree_module().verify_tree_greedy(
|
||||||
|
predicts,
|
||||||
|
accept_index,
|
||||||
|
accept_token_num,
|
||||||
|
candidates,
|
||||||
|
retrive_index,
|
||||||
|
retrive_next_token,
|
||||||
|
retrive_next_sibling,
|
||||||
|
target_predict,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@debug_kernel_api
|
||||||
|
def reconstruct_indices_from_tree_mask(
|
||||||
|
tree_mask: torch.Tensor,
|
||||||
|
verified_seq_len: torch.Tensor,
|
||||||
|
positions: torch.Tensor,
|
||||||
|
retrive_index: torch.Tensor,
|
||||||
|
retrive_next_token: torch.Tensor,
|
||||||
|
retrive_next_sibling: torch.Tensor,
|
||||||
|
batch_size: int,
|
||||||
|
draft_token_num: int,
|
||||||
|
) -> None:
|
||||||
|
_jit_tree_module().reconstruct_indices_from_tree_mask(
|
||||||
|
tree_mask,
|
||||||
|
verified_seq_len,
|
||||||
|
positions,
|
||||||
|
retrive_index,
|
||||||
|
retrive_next_token,
|
||||||
|
retrive_next_sibling,
|
||||||
|
batch_size,
|
||||||
|
draft_token_num,
|
||||||
|
)
|
||||||
@@ -41,7 +41,16 @@ _DFLASH_VERIFY_SKIP_CUSTOM_MASK_BACKENDS = frozenset(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
if is_cuda() or is_musa():
|
if is_cuda():
|
||||||
|
from flashinfer.sampling import top_k_renorm_probs as top_k_renorm_prob
|
||||||
|
from flashinfer.sampling import top_p_renorm_probs as top_p_renorm_prob
|
||||||
|
|
||||||
|
from sglang.kernels.ops.speculative.sampling import (
|
||||||
|
tree_speculative_sampling_target_only,
|
||||||
|
)
|
||||||
|
|
||||||
|
_DFLASH_SAMPLING_VERIFY_AVAILABLE = True
|
||||||
|
elif is_musa():
|
||||||
try:
|
try:
|
||||||
from sgl_kernel import (
|
from sgl_kernel import (
|
||||||
top_k_renorm_prob,
|
top_k_renorm_prob,
|
||||||
|
|||||||
@@ -49,7 +49,11 @@ _is_cpu = is_cpu()
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
if _is_cuda or _is_hip or _is_musa:
|
if _is_cuda or _is_hip:
|
||||||
|
from sglang.kernels.ops.speculative.tree import (
|
||||||
|
build_tree_kernel_efficient as sgl_build_tree_kernel_efficient,
|
||||||
|
)
|
||||||
|
elif _is_musa:
|
||||||
from sgl_kernel import (
|
from sgl_kernel import (
|
||||||
build_tree_kernel_efficient as sgl_build_tree_kernel_efficient,
|
build_tree_kernel_efficient as sgl_build_tree_kernel_efficient,
|
||||||
)
|
)
|
||||||
@@ -386,6 +390,9 @@ def verify_tree_greedy_func(
|
|||||||
topk: int = -1,
|
topk: int = -1,
|
||||||
):
|
):
|
||||||
if _is_cuda or _is_hip or _is_musa:
|
if _is_cuda or _is_hip or _is_musa:
|
||||||
|
if _is_cuda or _is_hip:
|
||||||
|
from sglang.kernels.ops.speculative.tree import verify_tree_greedy
|
||||||
|
else:
|
||||||
from sgl_kernel import verify_tree_greedy
|
from sgl_kernel import verify_tree_greedy
|
||||||
|
|
||||||
verify_tree_greedy(
|
verify_tree_greedy(
|
||||||
@@ -889,7 +896,11 @@ def eagle_sample(
|
|||||||
if use_rejection_sampling:
|
if use_rejection_sampling:
|
||||||
sampling_fn = chain_speculative_sampling_triton
|
sampling_fn = chain_speculative_sampling_triton
|
||||||
else:
|
else:
|
||||||
if not _is_npu:
|
if _is_cuda:
|
||||||
|
from sglang.kernels.ops.speculative.sampling import (
|
||||||
|
tree_speculative_sampling_target_only,
|
||||||
|
)
|
||||||
|
elif not _is_npu:
|
||||||
from sgl_kernel import tree_speculative_sampling_target_only
|
from sgl_kernel import tree_speculative_sampling_target_only
|
||||||
|
|
||||||
sampling_fn = tree_speculative_sampling_target_only
|
sampling_fn = tree_speculative_sampling_target_only
|
||||||
@@ -902,6 +913,9 @@ def eagle_sample(
|
|||||||
from sglang.kernels.ops.sampling.renorm_triton import (
|
from sglang.kernels.ops.sampling.renorm_triton import (
|
||||||
top_p_renorm_probs_triton as top_p_renorm_prob,
|
top_p_renorm_probs_triton as top_p_renorm_prob,
|
||||||
)
|
)
|
||||||
|
elif _is_cuda:
|
||||||
|
from flashinfer.sampling import top_k_renorm_probs as top_k_renorm_prob
|
||||||
|
from flashinfer.sampling import top_p_renorm_probs as top_p_renorm_prob
|
||||||
elif not _is_npu:
|
elif not _is_npu:
|
||||||
from sgl_kernel import top_k_renorm_prob, top_p_renorm_prob
|
from sgl_kernel import top_k_renorm_prob, top_p_renorm_prob
|
||||||
|
|
||||||
|
|||||||
@@ -3,7 +3,6 @@ from typing import List, Optional
|
|||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from sgl_kernel.speculative import reconstruct_indices_from_tree_mask
|
|
||||||
|
|
||||||
from sglang.kernels.ops.speculative.cache_locs import (
|
from sglang.kernels.ops.speculative.cache_locs import (
|
||||||
assign_extend_cache_locs_func as assign_extend_cache_locs_func,
|
assign_extend_cache_locs_func as assign_extend_cache_locs_func,
|
||||||
@@ -33,11 +32,16 @@ from sglang.srt.speculative.spec_utils import (
|
|||||||
prepare_mamba_track_for_verify,
|
prepare_mamba_track_for_verify,
|
||||||
record_stream_for_v2_verify,
|
record_stream_for_v2_verify,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import is_cpu
|
from sglang.srt.utils import is_cpu, is_cuda
|
||||||
from sglang.srt.utils.async_probe import maybe_detect_inf, maybe_detect_nan
|
from sglang.srt.utils.async_probe import maybe_detect_inf, maybe_detect_nan
|
||||||
|
|
||||||
_is_cpu = is_cpu()
|
_is_cpu = is_cpu()
|
||||||
|
|
||||||
|
if is_cuda():
|
||||||
|
from sglang.kernels.ops.speculative.tree import reconstruct_indices_from_tree_mask
|
||||||
|
else:
|
||||||
|
from sgl_kernel.speculative import reconstruct_indices_from_tree_mask
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -64,6 +64,7 @@ from sglang.srt.utils import (
|
|||||||
next_power_of_2,
|
next_power_of_2,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.async_probe import maybe_detect_oob
|
from sglang.srt.utils.async_probe import maybe_detect_oob
|
||||||
|
from sglang.srt.utils.common import fast_topk
|
||||||
from sglang.srt.utils.nvtx_utils import profile_range
|
from sglang.srt.utils.nvtx_utils import profile_range
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
@@ -81,13 +82,6 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
|
|
||||||
|
|
||||||
if _is_cuda:
|
|
||||||
from sgl_kernel import fast_topk
|
|
||||||
elif _is_hip:
|
|
||||||
from sgl_kernel import fast_topk
|
|
||||||
else:
|
|
||||||
from sglang.srt.utils.common import fast_topk
|
|
||||||
|
|
||||||
if _is_cpu:
|
if _is_cpu:
|
||||||
from sgl_kernel import assign_extend_cache_locs_cpu
|
from sgl_kernel import assign_extend_cache_locs_cpu
|
||||||
|
|
||||||
|
|||||||
@@ -28,10 +28,7 @@ from sglang.srt.speculative.eagle_utils import default_tree_mask_mode
|
|||||||
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
|
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.speculative.spec_utils import draft_tp_context, get_plan_stream
|
from sglang.srt.speculative.spec_utils import draft_tp_context, get_plan_stream
|
||||||
from sglang.srt.utils import empty_context, get_bool_env_var, is_cuda
|
from sglang.srt.utils import empty_context, get_bool_env_var
|
||||||
|
|
||||||
if is_cuda():
|
|
||||||
from sgl_kernel import segment_packbits # noqa: F401
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
SGLANG_RETURN_ORIGINAL_LOGPROB = get_bool_env_var("SGLANG_RETURN_ORIGINAL_LOGPROB")
|
SGLANG_RETURN_ORIGINAL_LOGPROB = get_bool_env_var("SGLANG_RETURN_ORIGINAL_LOGPROB")
|
||||||
|
|||||||
@@ -294,7 +294,7 @@ class TestVerifyTreeGreedy(CustomTestCase):
|
|||||||
|
|
||||||
def test_verify_tree_greedy_upstream_golden(self):
|
def test_verify_tree_greedy_upstream_golden(self):
|
||||||
# Golden fixture ported from the CUDA kernel UT
|
# Golden fixture ported from the CUDA kernel UT
|
||||||
# python/sglang/kernels/aot/tests/speculative/test_eagle_utils.py::test_verify_tree_greedy
|
# test/registered/kernels/ops/speculative/test_eagle_utils.py::test_verify_tree_greedy
|
||||||
# (device swapped to CPU); expected outputs are the CUDA kernel's.
|
# (device swapped to CPU); expected outputs are the CUDA kernel's.
|
||||||
candidates = torch.tensor(
|
candidates = torch.tensor(
|
||||||
[
|
[
|
||||||
@@ -863,7 +863,7 @@ class TestReconstructIndicesFromTreeMask(CustomTestCase):
|
|||||||
bs, draft_token_num = 2, 4
|
bs, draft_token_num = 2, 4
|
||||||
seq_lens = torch.tensor([12, 5], dtype=torch.int64)
|
seq_lens = torch.tensor([12, 5], dtype=torch.int64)
|
||||||
# Request 0: root(0) -> {1, 2}, 2 -> 3 (golden case from
|
# Request 0: root(0) -> {1, 2}, 2 -> 3 (golden case from
|
||||||
# python/sglang/kernels/aot/tests/speculative/test_ngram_utils.py).
|
# test/registered/kernels/ops/speculative/test_ngram_utils.py).
|
||||||
# Request 1: plain chain 0 -> 1 -> 2 -> 3.
|
# Request 1: plain chain 0 -> 1 -> 2 -> 3.
|
||||||
tree_mask = torch.tensor(
|
tree_mask = torch.tensor(
|
||||||
# fmt: off
|
# fmt: off
|
||||||
|
|||||||
+12
-3
@@ -2,8 +2,17 @@ import sys
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
|
||||||
from sgl_kernel import verify_tree_greedy
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
# CUDA and ROCm both build this kernel through the JIT launcher; MUSA still
|
||||||
|
# gets it from the AOT extension.
|
||||||
|
if torch.version.cuda is not None or torch.version.hip is not None:
|
||||||
|
from sglang.kernels.ops.speculative.tree import verify_tree_greedy
|
||||||
|
else:
|
||||||
|
from sgl_kernel import verify_tree_greedy
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=15, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||||
|
|
||||||
|
|
||||||
def test_verify_tree_greedy():
|
def test_verify_tree_greedy():
|
||||||
@@ -86,4 +95,4 @@ def test_verify_tree_greedy():
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__, "-v"]))
|
||||||
+10
-4
@@ -2,8 +2,15 @@ import sys
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
|
||||||
from sgl_kernel import reconstruct_indices_from_tree_mask
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
if torch.version.cuda is not None:
|
||||||
|
from sglang.kernels.ops.speculative.tree import reconstruct_indices_from_tree_mask
|
||||||
|
else:
|
||||||
|
from sgl_kernel import reconstruct_indices_from_tree_mask
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=15, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||||
|
|
||||||
|
|
||||||
def test_reconstruct_indices_from_tree_mask():
|
def test_reconstruct_indices_from_tree_mask():
|
||||||
@@ -74,5 +81,4 @@ def test_reconstruct_indices_from_tree_mask():
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
test_reconstruct_indices_from_tree_mask()
|
sys.exit(pytest.main([__file__, "-v"]))
|
||||||
sys.exit(pytest.main([__file__]))
|
|
||||||
+12
-2
@@ -3,7 +3,17 @@ import sys
|
|||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from sgl_kernel import tree_speculative_sampling_target_only
|
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
if torch.version.cuda is not None:
|
||||||
|
from sglang.kernels.ops.speculative.sampling import (
|
||||||
|
tree_speculative_sampling_target_only,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
from sgl_kernel import tree_speculative_sampling_target_only
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=15, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||||
|
|
||||||
test_cases = [
|
test_cases = [
|
||||||
(
|
(
|
||||||
@@ -128,4 +138,4 @@ def test_tree_speculative_sampling_target_only(
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__, "-v"]))
|
||||||
@@ -219,7 +219,7 @@ class TestGraphPoolBorrow(CustomTestCase):
|
|||||||
),
|
),
|
||||||
patch("sglang.srt.distributed.get_tp_group", return_value=tp_group),
|
patch("sglang.srt.distributed.get_tp_group", return_value=tp_group),
|
||||||
patch(
|
patch(
|
||||||
"sgl_kernel.tree_speculative_sampling_target_only",
|
"sglang.kernels.ops.speculative.sampling.tree_speculative_sampling_target_only",
|
||||||
side_effect=fake_sampling,
|
side_effect=fake_sampling,
|
||||||
),
|
),
|
||||||
):
|
):
|
||||||
|
|||||||
Reference in New Issue
Block a user