Feature/support longcat flash lite (#17838)
Co-authored-by: sunjiaqi11 <sunjiaqi11@meituan.com> Co-authored-by: ispobock <ispobaoke@gmail.com>
This commit is contained in:
co-authored by
sunjiaqi11
ispobock
parent
11b76d24dc
commit
eb4ba1bde2
@@ -0,0 +1,295 @@
|
|||||||
|
#include <sgl_kernel/tensor.h>
|
||||||
|
#include <sgl_kernel/utils.h>
|
||||||
|
|
||||||
|
#include <sgl_kernel/utils.cuh>
|
||||||
|
|
||||||
|
#include <dlpack/dlpack.h>
|
||||||
|
|
||||||
|
#include <algorithm>
|
||||||
|
#include <concepts>
|
||||||
|
#include <cstddef>
|
||||||
|
#include <cstdint>
|
||||||
|
#include <type_traits>
|
||||||
|
|
||||||
|
namespace device::ngram_embedding {
|
||||||
|
|
||||||
|
__global__ void ComputeNGramIdsKernel(
|
||||||
|
int batch_size,
|
||||||
|
int ne_n,
|
||||||
|
int ne_k,
|
||||||
|
int* ne_weights, // [ne_n-1,ne_k,ne_n]
|
||||||
|
int* ne_mods, // [ne_n-1,ne_k]
|
||||||
|
int* exclusive_ne_embeder_size_sums, // [(ne_n-1)*ne_k]
|
||||||
|
int* tokens, // [token_num]
|
||||||
|
int* exclusive_req_len_sums, // [batch_size+1]
|
||||||
|
int* ne_token_table, // [max_running_reqs, max_context_len]
|
||||||
|
int max_context_len, // max_context_len
|
||||||
|
long* row_indices, // [batch_size]
|
||||||
|
int* column_starts, // [batch_size]
|
||||||
|
int* n_gram_ids // [ne_n-1,ne_k,token_num]
|
||||||
|
) {
|
||||||
|
// Determine which n, k, and request this block handles.
|
||||||
|
/**
|
||||||
|
Example: [req0, req1, req2] with n=3, k=2
|
||||||
|
n k req_id blockIdx.x config_id (combination of n and k)
|
||||||
|
2 1 0 0 0
|
||||||
|
2 1 1 1 0
|
||||||
|
2 1 2 2 0
|
||||||
|
2 2 0 3 1
|
||||||
|
2 2 1 4 1
|
||||||
|
2 2 2 5 1
|
||||||
|
3 1 0 0 2
|
||||||
|
3 1 1 1 2
|
||||||
|
3 1 2 2 2
|
||||||
|
3 2 0 3 3
|
||||||
|
3 2 1 4 3
|
||||||
|
3 2 2 5 3
|
||||||
|
*/
|
||||||
|
const int req_id = blockIdx.x % batch_size;
|
||||||
|
const int config_id = (blockIdx.x - req_id) / batch_size;
|
||||||
|
// n and k here are offset from their physical meanings: n = real_n - 2, k = real_k - 1.
|
||||||
|
// This offset exists because n and k are used as indices into ne_weights and ne_mods.
|
||||||
|
const int k = config_id % ne_k;
|
||||||
|
const int n = (config_id - config_id % ne_k) / ne_k;
|
||||||
|
// ne_weights has shape [ne_n-1, ne_k, ne_n]; last dim is token distance, so compute base index first
|
||||||
|
const int ne_weight_base_idx = n * ne_k * ne_n + k * ne_n;
|
||||||
|
// ne_mods has shape [ne_n-1, ne_k]
|
||||||
|
const int ne_mod = ne_mods[n * ne_k + k];
|
||||||
|
// stride loop
|
||||||
|
for (int i = exclusive_req_len_sums[req_id] + threadIdx.x; i < exclusive_req_len_sums[req_id + 1]; i += blockDim.x) {
|
||||||
|
uint64_t n_gram_id = 0;
|
||||||
|
// Token offset within the current request
|
||||||
|
int current_token_offset = i - exclusive_req_len_sums[req_id];
|
||||||
|
// Start index of this request in the token table; tokens before this belong to other requests
|
||||||
|
int req_token_table_index = row_indices[req_id] * max_context_len;
|
||||||
|
// Position of the current token in the token table
|
||||||
|
int current_token_table_index = req_token_table_index + column_starts[req_id] + current_token_offset;
|
||||||
|
for (int j = 0; j < n + 2; j++) {
|
||||||
|
if (current_token_table_index - j < req_token_table_index) {
|
||||||
|
// Out of this request's range, stop computing n_gram_id
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
if (ne_token_table[current_token_table_index - j] < 0) {
|
||||||
|
// Token was marked as ignored during write
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
const uint64_t term =
|
||||||
|
(uint64_t)ne_token_table[current_token_table_index - j] * (uint64_t)ne_weights[ne_weight_base_idx + j];
|
||||||
|
n_gram_id += term % ne_mod;
|
||||||
|
}
|
||||||
|
n_gram_id %= ne_mod;
|
||||||
|
n_gram_id += exclusive_ne_embeder_size_sums[n * ne_k + k];
|
||||||
|
// [token_num, ne_n-1, ne_k]
|
||||||
|
n_gram_ids[i * (ne_n - 1) * ne_k + n * ne_k + k] = (int)(n_gram_id);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
__global__ void UpdateTokenTableKernel(
|
||||||
|
int batch_size,
|
||||||
|
int* tokens, // [token_num]
|
||||||
|
int* ne_token_table, // [max_running_reqs, max_context_len]
|
||||||
|
int max_context_len, // max_context_len
|
||||||
|
long* row_indices, // [batch_size]
|
||||||
|
int* column_starts, // [batch_size]
|
||||||
|
int* req_lens, // [batch_size]
|
||||||
|
int ignore_token_num, // number of tokens to ignore
|
||||||
|
int* ignore_tokens // [ignore_token_num]
|
||||||
|
) {
|
||||||
|
// Each block processes one request.
|
||||||
|
const int req_id = blockIdx.x % batch_size;
|
||||||
|
int start = 0;
|
||||||
|
int end = 0;
|
||||||
|
for (int i = 0; i < req_id; i++) {
|
||||||
|
start += req_lens[i];
|
||||||
|
}
|
||||||
|
end = start + req_lens[req_id];
|
||||||
|
// stride loop
|
||||||
|
for (int i = start + threadIdx.x; i < end; i += blockDim.x) {
|
||||||
|
// Token offset within the current request
|
||||||
|
int current_token_offset = i - start;
|
||||||
|
// Start index of this request in the token table
|
||||||
|
int req_token_table_index = row_indices[req_id] * max_context_len;
|
||||||
|
// Position of the current token in the token table
|
||||||
|
int current_token_table_index = req_token_table_index + column_starts[req_id] + current_token_offset;
|
||||||
|
ne_token_table[current_token_table_index] = tokens[i];
|
||||||
|
for (int j = 0; j < ignore_token_num; j++) {
|
||||||
|
if (ignore_tokens[j] == tokens[i]) {
|
||||||
|
ne_token_table[current_token_table_index] = -tokens[i];
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
} // namespace device::ngram_embedding
|
||||||
|
|
||||||
|
namespace {
|
||||||
|
|
||||||
|
struct NgramEmbeddingKernel {
|
||||||
|
static void compute_n_gram_ids(
|
||||||
|
const int64_t ne_n,
|
||||||
|
const int64_t ne_k,
|
||||||
|
const tvm::ffi::TensorView ne_weights,
|
||||||
|
const tvm::ffi::TensorView ne_mods,
|
||||||
|
const tvm::ffi::TensorView exclusive_ne_embeder_size_sums,
|
||||||
|
const tvm::ffi::TensorView tokens,
|
||||||
|
const tvm::ffi::TensorView exclusive_req_len_sums,
|
||||||
|
const tvm::ffi::TensorView ne_token_table,
|
||||||
|
const tvm::ffi::TensorView row_indices,
|
||||||
|
const tvm::ffi::TensorView column_starts,
|
||||||
|
const tvm::ffi::TensorView n_gram_ids) {
|
||||||
|
using namespace host;
|
||||||
|
|
||||||
|
auto device_ = SymbolicDevice{};
|
||||||
|
|
||||||
|
// Verify tensor shapes and types using -1 (kAnySize) for dynamic dimensions
|
||||||
|
TensorMatcher({-1, -1, -1}) // [ne_n-1, ne_k, ne_n]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>(device_)
|
||||||
|
.verify(ne_weights);
|
||||||
|
|
||||||
|
TensorMatcher({-1, -1}) // [ne_n-1, ne_k]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(ne_mods);
|
||||||
|
|
||||||
|
TensorMatcher({-1}) // [(ne_n-1)*ne_k + 1]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(exclusive_ne_embeder_size_sums);
|
||||||
|
|
||||||
|
TensorMatcher({-1}) // [token_num]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(tokens);
|
||||||
|
|
||||||
|
TensorMatcher({-1}) // [batch_size+1]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(exclusive_req_len_sums);
|
||||||
|
|
||||||
|
TensorMatcher({-1, -1}) // [max_running_reqs, max_context_len]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(ne_token_table);
|
||||||
|
|
||||||
|
TensorMatcher({-1}) // [batch_size]
|
||||||
|
.with_dtype<int64_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(row_indices);
|
||||||
|
|
||||||
|
TensorMatcher({-1}) // [batch_size]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(column_starts);
|
||||||
|
|
||||||
|
TensorMatcher({-1, -1}) // [token_num, (ne_n-1)*ne_k]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(n_gram_ids);
|
||||||
|
|
||||||
|
const int batch_size = static_cast<int>(exclusive_req_len_sums.size(0) - 1);
|
||||||
|
const int max_context_len = static_cast<int>(ne_token_table.size(1));
|
||||||
|
const auto stream = LaunchKernel::resolve_device(device_.unwrap());
|
||||||
|
|
||||||
|
constexpr int BLOCK_THREADS = 256;
|
||||||
|
const int num_configs = (static_cast<int>(ne_n) - 1) * static_cast<int>(ne_k);
|
||||||
|
const int grid_size = num_configs * batch_size;
|
||||||
|
|
||||||
|
LaunchKernel(grid_size, BLOCK_THREADS, stream)(
|
||||||
|
device::ngram_embedding::ComputeNGramIdsKernel,
|
||||||
|
batch_size,
|
||||||
|
static_cast<int>(ne_n),
|
||||||
|
static_cast<int>(ne_k),
|
||||||
|
static_cast<int*>(ne_weights.data_ptr()),
|
||||||
|
static_cast<int*>(ne_mods.data_ptr()),
|
||||||
|
static_cast<int*>(exclusive_ne_embeder_size_sums.data_ptr()),
|
||||||
|
static_cast<int*>(tokens.data_ptr()),
|
||||||
|
static_cast<int*>(exclusive_req_len_sums.data_ptr()),
|
||||||
|
static_cast<int*>(ne_token_table.data_ptr()),
|
||||||
|
max_context_len,
|
||||||
|
static_cast<long*>(row_indices.data_ptr()),
|
||||||
|
static_cast<int*>(column_starts.data_ptr()),
|
||||||
|
static_cast<int*>(n_gram_ids.data_ptr()));
|
||||||
|
}
|
||||||
|
|
||||||
|
static void update_token_table(
|
||||||
|
const tvm::ffi::TensorView tokens,
|
||||||
|
const tvm::ffi::TensorView ne_token_table,
|
||||||
|
const tvm::ffi::TensorView row_indices,
|
||||||
|
const tvm::ffi::TensorView column_starts,
|
||||||
|
const tvm::ffi::TensorView req_lens,
|
||||||
|
const tvm::ffi::TensorView ignore_tokens) {
|
||||||
|
using namespace host;
|
||||||
|
|
||||||
|
auto device_ = SymbolicDevice{};
|
||||||
|
|
||||||
|
// Verify tensor shapes and types using -1 (kAnySize) for dynamic dimensions
|
||||||
|
TensorMatcher({-1}) // [token_num]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>(device_)
|
||||||
|
.verify(tokens);
|
||||||
|
|
||||||
|
TensorMatcher({-1, -1}) // [max_running_reqs, max_context_len]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(ne_token_table);
|
||||||
|
|
||||||
|
TensorMatcher({-1}) // [batch_size]
|
||||||
|
.with_dtype<int64_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(row_indices);
|
||||||
|
|
||||||
|
TensorMatcher({-1}) // [batch_size]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(column_starts);
|
||||||
|
|
||||||
|
TensorMatcher({-1}) // [batch_size]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(req_lens);
|
||||||
|
|
||||||
|
// ignore_tokens can be empty or have values
|
||||||
|
void* ignore_tokens_ptr = ignore_tokens.data_ptr();
|
||||||
|
const bool has_ignore_tokens = ignore_tokens_ptr != nullptr && ignore_tokens.numel() > 0;
|
||||||
|
if (has_ignore_tokens) {
|
||||||
|
TensorMatcher({-1}) // [ignore_token_num]
|
||||||
|
.with_dtype<int32_t>()
|
||||||
|
.with_device<kDLCUDA>()
|
||||||
|
.verify(ignore_tokens);
|
||||||
|
}
|
||||||
|
|
||||||
|
const int batch_size = static_cast<int>(req_lens.size(0));
|
||||||
|
if (batch_size <= 0) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
const int max_context_len = static_cast<int>(ne_token_table.size(1));
|
||||||
|
const auto stream = LaunchKernel::resolve_device(device_.unwrap());
|
||||||
|
|
||||||
|
constexpr int BLOCK_THREADS = 256;
|
||||||
|
const int grid_size = batch_size;
|
||||||
|
|
||||||
|
int ignore_token_num = 0;
|
||||||
|
int* ignore_tokens_typed_ptr = nullptr;
|
||||||
|
if (has_ignore_tokens) {
|
||||||
|
ignore_token_num = static_cast<int>(ignore_tokens.numel());
|
||||||
|
ignore_tokens_typed_ptr = static_cast<int*>(ignore_tokens_ptr);
|
||||||
|
}
|
||||||
|
|
||||||
|
LaunchKernel(grid_size, BLOCK_THREADS, stream)(
|
||||||
|
device::ngram_embedding::UpdateTokenTableKernel,
|
||||||
|
batch_size,
|
||||||
|
static_cast<int*>(tokens.data_ptr()),
|
||||||
|
static_cast<int*>(ne_token_table.data_ptr()),
|
||||||
|
max_context_len,
|
||||||
|
static_cast<long*>(row_indices.data_ptr()),
|
||||||
|
static_cast<int*>(column_starts.data_ptr()),
|
||||||
|
static_cast<int*>(req_lens.data_ptr()),
|
||||||
|
ignore_token_num,
|
||||||
|
ignore_tokens_typed_ptr);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
} // namespace
|
||||||
@@ -0,0 +1,100 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from functools import lru_cache
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from sglang.jit_kernel.utils import load_jit
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import torch
|
||||||
|
from tvm_ffi.module import Module
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache(maxsize=None)
|
||||||
|
def _jit_ngram_embedding_module() -> Module:
|
||||||
|
return load_jit(
|
||||||
|
"ngram_embedding",
|
||||||
|
cuda_files=["ngram_embedding.cuh"],
|
||||||
|
cuda_wrappers=[
|
||||||
|
("compute_n_gram_ids", "&NgramEmbeddingKernel::compute_n_gram_ids"),
|
||||||
|
("update_token_table", "&NgramEmbeddingKernel::update_token_table"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def compute_n_gram_ids(
|
||||||
|
ne_n: int,
|
||||||
|
ne_k: int,
|
||||||
|
ne_weights: torch.Tensor,
|
||||||
|
ne_mods: torch.Tensor,
|
||||||
|
exclusive_ne_embeder_size_sums: torch.Tensor,
|
||||||
|
tokens: torch.Tensor,
|
||||||
|
exclusive_req_len_sums: torch.Tensor,
|
||||||
|
ne_token_table: torch.Tensor,
|
||||||
|
row_indices: torch.Tensor,
|
||||||
|
column_starts: torch.Tensor,
|
||||||
|
n_gram_ids: torch.Tensor,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Compute n-gram IDs for embedding.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
ne_n: n value for n-gram
|
||||||
|
ne_k: k value for n-gram configurations
|
||||||
|
ne_weights: weights tensor with shape [ne_n-1, ne_k, ne_n]
|
||||||
|
ne_mods: mods tensor with shape [ne_n-1, ne_k]
|
||||||
|
exclusive_ne_embeder_size_sums: exclusive sum of embedder sizes
|
||||||
|
tokens: input token ids
|
||||||
|
exclusive_req_len_sums: exclusive sum of request lengths
|
||||||
|
ne_token_table: token table for all requests
|
||||||
|
row_indices: row indices for each request
|
||||||
|
column_starts: column start positions for each request
|
||||||
|
n_gram_ids: output tensor for n-gram ids
|
||||||
|
"""
|
||||||
|
module = _jit_ngram_embedding_module()
|
||||||
|
module.compute_n_gram_ids(
|
||||||
|
ne_n,
|
||||||
|
ne_k,
|
||||||
|
ne_weights,
|
||||||
|
ne_mods,
|
||||||
|
exclusive_ne_embeder_size_sums,
|
||||||
|
tokens,
|
||||||
|
exclusive_req_len_sums,
|
||||||
|
ne_token_table,
|
||||||
|
row_indices,
|
||||||
|
column_starts,
|
||||||
|
n_gram_ids,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def update_token_table(
|
||||||
|
tokens: torch.Tensor,
|
||||||
|
ne_token_table: torch.Tensor,
|
||||||
|
row_indices: torch.Tensor,
|
||||||
|
column_starts: torch.Tensor,
|
||||||
|
req_lens: torch.Tensor,
|
||||||
|
ignore_tokens: torch.Tensor | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Update the token table with new tokens.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
tokens: input token ids
|
||||||
|
ne_token_table: token table for all requests
|
||||||
|
row_indices: row indices for each request
|
||||||
|
column_starts: column start positions for each request
|
||||||
|
req_lens: request lengths
|
||||||
|
ignore_tokens: tokens to be ignored (marked as negative in table)
|
||||||
|
"""
|
||||||
|
module = _jit_ngram_embedding_module()
|
||||||
|
if ignore_tokens is None:
|
||||||
|
# Create an empty tensor for ignore_tokens
|
||||||
|
ignore_tokens = tokens.new_empty(0, dtype=tokens.dtype)
|
||||||
|
module.update_token_table(
|
||||||
|
tokens,
|
||||||
|
ne_token_table,
|
||||||
|
row_indices,
|
||||||
|
column_starts,
|
||||||
|
req_lens,
|
||||||
|
ignore_tokens,
|
||||||
|
)
|
||||||
@@ -53,6 +53,9 @@ class LongcatFlashConfig(PretrainedConfig):
|
|||||||
zero_expert_type="identity",
|
zero_expert_type="identity",
|
||||||
nextn_use_scmoe=False,
|
nextn_use_scmoe=False,
|
||||||
num_nextn_predict_layers=1,
|
num_nextn_predict_layers=1,
|
||||||
|
ngram_vocab_size_ratio=None,
|
||||||
|
emb_neighbor_num=None,
|
||||||
|
emb_split_num=None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(
|
super().__init__(
|
||||||
@@ -102,3 +105,8 @@ class LongcatFlashConfig(PretrainedConfig):
|
|||||||
self.zero_expert_type = zero_expert_type
|
self.zero_expert_type = zero_expert_type
|
||||||
self.routed_scaling_factor = routed_scaling_factor
|
self.routed_scaling_factor = routed_scaling_factor
|
||||||
self.hidden_act = "silu"
|
self.hidden_act = "silu"
|
||||||
|
self.use_ngram_embedding = ngram_vocab_size_ratio is not None
|
||||||
|
if self.use_ngram_embedding:
|
||||||
|
self.ngram_embedding_m = int(ngram_vocab_size_ratio * vocab_size)
|
||||||
|
self.ngram_embedding_n = emb_neighbor_num
|
||||||
|
self.ngram_embedding_k = emb_split_num
|
||||||
|
|||||||
@@ -193,6 +193,7 @@ class ModelConfig:
|
|||||||
self.is_local_attention_model = is_local_attention_model(
|
self.is_local_attention_model = is_local_attention_model(
|
||||||
self.hf_config.architectures
|
self.hf_config.architectures
|
||||||
)
|
)
|
||||||
|
self.use_ngram_embedding = getattr(self.hf_config, "use_ngram_embedding", False)
|
||||||
self.is_piecewise_cuda_graph_disabled_model = (
|
self.is_piecewise_cuda_graph_disabled_model = (
|
||||||
is_piecewise_cuda_graph_disabled_model(self.hf_config.architectures)
|
is_piecewise_cuda_graph_disabled_model(self.hf_config.architectures)
|
||||||
or is_deepseek_nsa(self.hf_text_config)
|
or is_deepseek_nsa(self.hf_text_config)
|
||||||
|
|||||||
@@ -0,0 +1,172 @@
|
|||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
from torch.nn import Parameter
|
||||||
|
|
||||||
|
from sglang.jit_kernel.ngram_embedding import compute_n_gram_ids
|
||||||
|
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
|
||||||
|
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
|
|
||||||
|
class NgramEmbedding(torch.nn.Module):
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
num_embeddings: int,
|
||||||
|
embedding_dim: int,
|
||||||
|
over_embedding_m: int,
|
||||||
|
over_embedding_k: int,
|
||||||
|
over_embedding_n: int,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
assert (
|
||||||
|
over_embedding_n > 1
|
||||||
|
), f"over_embedding_n must be > 1, got {over_embedding_n}"
|
||||||
|
self.num_embeddings = num_embeddings
|
||||||
|
self.embedding_dim = embedding_dim
|
||||||
|
self.over_embedding_m = over_embedding_m
|
||||||
|
self.over_embedding_k = over_embedding_k
|
||||||
|
self.over_embedding_n = over_embedding_n
|
||||||
|
|
||||||
|
self.word_embeder = VocabParallelEmbedding(
|
||||||
|
num_embeddings,
|
||||||
|
embedding_dim,
|
||||||
|
enable_tp=is_dp_attention_enabled(),
|
||||||
|
)
|
||||||
|
self.n_grams = (over_embedding_n - 1) * over_embedding_k
|
||||||
|
oe_hidden_dim = embedding_dim // (over_embedding_k * (over_embedding_n - 1))
|
||||||
|
self.exclusive_oe_embeder_size_sums = torch.zeros(
|
||||||
|
[over_embedding_k * (over_embedding_n - 1) + 1],
|
||||||
|
dtype=torch.int32,
|
||||||
|
device="cuda",
|
||||||
|
)
|
||||||
|
for i in range(over_embedding_k * (over_embedding_n - 1)):
|
||||||
|
self.exclusive_oe_embeder_size_sums[i + 1] = (
|
||||||
|
self.exclusive_oe_embeder_size_sums[i]
|
||||||
|
+ int(over_embedding_m + i * 2 + 1)
|
||||||
|
)
|
||||||
|
self.oe_embeder = VocabParallelEmbedding(
|
||||||
|
num_embeddings=self.exclusive_oe_embeder_size_sums[-1],
|
||||||
|
embedding_dim=oe_hidden_dim,
|
||||||
|
enable_tp=is_dp_attention_enabled(),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.oe_projection = nn.Parameter(
|
||||||
|
torch.empty(
|
||||||
|
(over_embedding_n - 1) * over_embedding_k, oe_hidden_dim, embedding_dim
|
||||||
|
),
|
||||||
|
requires_grad=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.oe_mods = torch.zeros(
|
||||||
|
[self.over_embedding_n - 1, self.over_embedding_k], dtype=torch.int32
|
||||||
|
)
|
||||||
|
self.oe_weights = torch.zeros(
|
||||||
|
[self.over_embedding_n - 1, self.over_embedding_k, self.over_embedding_n],
|
||||||
|
dtype=torch.int32,
|
||||||
|
)
|
||||||
|
for n in range(2, self.over_embedding_n + 1):
|
||||||
|
for k in range(self.over_embedding_k):
|
||||||
|
mod = (
|
||||||
|
self.over_embedding_m
|
||||||
|
+ 2 * ((n - 2) * self.over_embedding_k + k)
|
||||||
|
+ 1
|
||||||
|
)
|
||||||
|
self.oe_mods[n - 2][k] = mod
|
||||||
|
for delta in range(self.over_embedding_n):
|
||||||
|
self.oe_weights[n - 2][k][delta] = pow(num_embeddings, delta, mod)
|
||||||
|
|
||||||
|
def init_buffers(
|
||||||
|
self, max_running_requests: int, chunked_prefill_size: int, device: str
|
||||||
|
):
|
||||||
|
max_tokens = max(chunked_prefill_size, max_running_requests)
|
||||||
|
self.oe_n_gram_ids = torch.zeros(
|
||||||
|
[max_tokens, self.n_grams],
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
self.exclusive_req_len_sums = torch.zeros(
|
||||||
|
max_running_requests + 1, dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
def load_weight(
|
||||||
|
self, param: Parameter, weight_name: str, loaded_weight: torch.Tensor
|
||||||
|
):
|
||||||
|
if ".embed_tokens." in weight_name:
|
||||||
|
param.weight_loader(param, loaded_weight)
|
||||||
|
elif "model.ngram_embeddings.embedders." in weight_name:
|
||||||
|
index = int(
|
||||||
|
weight_name.replace("model.ngram_embeddings.embedders.", "").replace(
|
||||||
|
".weight", ""
|
||||||
|
)
|
||||||
|
)
|
||||||
|
oe_weight_start = self.exclusive_oe_embeder_size_sums[index]
|
||||||
|
oe_weight_end = self.exclusive_oe_embeder_size_sums[index + 1]
|
||||||
|
assert (
|
||||||
|
oe_weight_end - oe_weight_start == loaded_weight.shape[0]
|
||||||
|
), f"{oe_weight_end - oe_weight_start=} {loaded_weight.shape[0]=}"
|
||||||
|
tp_start = self.oe_embeder.shard_indices.org_vocab_start_index
|
||||||
|
tp_end = self.oe_embeder.shard_indices.org_vocab_end_index
|
||||||
|
to_load_start = max(oe_weight_start, tp_start)
|
||||||
|
to_load_end = min(oe_weight_end, tp_end)
|
||||||
|
if to_load_start < to_load_end:
|
||||||
|
src_start = to_load_start - oe_weight_start
|
||||||
|
src_end = to_load_end - oe_weight_start
|
||||||
|
dest_start = to_load_start - tp_start
|
||||||
|
dest_end = to_load_end - tp_start
|
||||||
|
self.oe_embeder.weight.data[dest_start:dest_end] = loaded_weight[
|
||||||
|
src_start:src_end
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
return
|
||||||
|
elif "model.ngram_embeddings.post_projs." in weight_name:
|
||||||
|
index = int(
|
||||||
|
weight_name.replace("model.ngram_embeddings.post_projs.", "").replace(
|
||||||
|
".weight", ""
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.oe_projection[index].copy_(loaded_weight.data.t())
|
||||||
|
else:
|
||||||
|
assert False, f"Unknown ngram embedding weight name: {weight_name}"
|
||||||
|
|
||||||
|
def forward(self, input_ids: torch.Tensor, forward_batch: ForwardBatch):
|
||||||
|
if (
|
||||||
|
forward_batch.forward_mode.is_extend()
|
||||||
|
or forward_batch.forward_mode.is_decode()
|
||||||
|
):
|
||||||
|
ngram_embedding_info = forward_batch.ngram_embedding_info
|
||||||
|
torch.cumsum(
|
||||||
|
ngram_embedding_info.req_lens,
|
||||||
|
dim=0,
|
||||||
|
dtype=torch.int32,
|
||||||
|
out=self.exclusive_req_len_sums[1 : 1 + forward_batch.batch_size],
|
||||||
|
)
|
||||||
|
compute_n_gram_ids(
|
||||||
|
ne_n=self.over_embedding_n,
|
||||||
|
ne_k=self.over_embedding_k,
|
||||||
|
ne_weights=self.oe_weights,
|
||||||
|
ne_mods=self.oe_mods,
|
||||||
|
tokens=input_ids.to(torch.int32),
|
||||||
|
exclusive_ne_embeder_size_sums=self.exclusive_oe_embeder_size_sums,
|
||||||
|
exclusive_req_len_sums=self.exclusive_req_len_sums[
|
||||||
|
: forward_batch.batch_size + 1
|
||||||
|
],
|
||||||
|
ne_token_table=ngram_embedding_info.token_table,
|
||||||
|
row_indices=forward_batch.req_pool_indices,
|
||||||
|
column_starts=ngram_embedding_info.column_starts,
|
||||||
|
n_gram_ids=self.oe_n_gram_ids[: len(input_ids)],
|
||||||
|
)
|
||||||
|
|
||||||
|
# [13, seq_len, hidden_dim]
|
||||||
|
all_hidden_states = torch.empty(
|
||||||
|
[self.n_grams + 1, len(input_ids), self.embedding_dim],
|
||||||
|
dtype=self.oe_projection.dtype,
|
||||||
|
device=input_ids.device,
|
||||||
|
)
|
||||||
|
all_hidden_states[0] = self.word_embeder(input_ids)
|
||||||
|
# oe_hidden_states: [12, seq_len, hidden_dim / 12]
|
||||||
|
oe_hidden_states = self.oe_embeder(
|
||||||
|
self.oe_n_gram_ids[: len(input_ids)].permute(1, 0).contiguous()
|
||||||
|
)
|
||||||
|
torch.bmm(oe_hidden_states, self.oe_projection, out=all_hidden_states[1:])
|
||||||
|
return all_hidden_states.mean(dim=0)
|
||||||
@@ -1216,6 +1216,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# Batched arguments to model runner
|
# Batched arguments to model runner
|
||||||
input_ids: torch.Tensor = None # shape: [b], int64
|
input_ids: torch.Tensor = None # shape: [b], int64
|
||||||
input_embeds: torch.Tensor = None # shape: [b, hidden_size], float32
|
input_embeds: torch.Tensor = None # shape: [b, hidden_size], float32
|
||||||
|
ne_token_table: torch.Tensor = None
|
||||||
token_type_ids: torch.Tensor = None # shape: [b], int64
|
token_type_ids: torch.Tensor = None # shape: [b], int64
|
||||||
req_pool_indices: torch.Tensor = None # shape: [b], int64
|
req_pool_indices: torch.Tensor = None # shape: [b], int64
|
||||||
seq_lens: torch.Tensor = None # shape: [b], int64
|
seq_lens: torch.Tensor = None # shape: [b], int64
|
||||||
@@ -1932,7 +1933,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
self.seq_lens_cpu = torch.empty(0, dtype=torch.int64)
|
self.seq_lens_cpu = torch.empty(0, dtype=torch.int64)
|
||||||
self.orig_seq_lens = torch.empty(0, dtype=torch.int32, device=self.device)
|
self.orig_seq_lens = torch.empty(0, dtype=torch.int32, device=self.device)
|
||||||
self.out_cache_loc = torch.empty(0, dtype=torch.int64, device=self.device)
|
self.out_cache_loc = torch.empty(0, dtype=torch.int64, device=self.device)
|
||||||
self.req_pool_indices = torch.empty(0, dtype=torch.int32, device=self.device)
|
self.req_pool_indices = torch.empty(0, dtype=torch.int64, device=self.device)
|
||||||
self.seq_lens_sum = 0
|
self.seq_lens_sum = 0
|
||||||
self.extend_num_tokens = 0
|
self.extend_num_tokens = 0
|
||||||
self.sampling_info = SamplingBatchInfo.from_schedule_batch(
|
self.sampling_info = SamplingBatchInfo.from_schedule_batch(
|
||||||
@@ -2234,6 +2235,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
lora_ids=[req.lora_id for req in self.reqs],
|
lora_ids=[req.lora_id for req in self.reqs],
|
||||||
sampling_info=self.sampling_info,
|
sampling_info=self.sampling_info,
|
||||||
input_embeds=self.input_embeds,
|
input_embeds=self.input_embeds,
|
||||||
|
ne_token_table=self.ne_token_table,
|
||||||
token_type_ids=self.token_type_ids,
|
token_type_ids=self.token_type_ids,
|
||||||
spec_algorithm=self.spec_algorithm,
|
spec_algorithm=self.spec_algorithm,
|
||||||
spec_info=self.spec_info,
|
spec_info=self.spec_info,
|
||||||
@@ -2416,6 +2418,9 @@ class ModelWorkerBatch:
|
|||||||
# The input Embeds
|
# The input Embeds
|
||||||
input_embeds: Optional[torch.Tensor] = None
|
input_embeds: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
# token table for ngram embedding
|
||||||
|
ne_token_table: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# For corss-encoder model
|
# For corss-encoder model
|
||||||
token_type_ids: Optional[torch.Tensor] = None
|
token_type_ids: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
|||||||
@@ -33,6 +33,7 @@ from torch.cuda import Stream as CudaStream
|
|||||||
from torch.cuda import StreamContext as CudaStreamContext
|
from torch.cuda import StreamContext as CudaStreamContext
|
||||||
from torch.distributed import barrier
|
from torch.distributed import barrier
|
||||||
|
|
||||||
|
from sglang.jit_kernel.ngram_embedding import update_token_table
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
from sglang.srt.constrained.grammar_manager import GrammarManager
|
from sglang.srt.constrained.grammar_manager import GrammarManager
|
||||||
from sglang.srt.disaggregation.decode import (
|
from sglang.srt.disaggregation.decode import (
|
||||||
@@ -394,6 +395,9 @@ class Scheduler(
|
|||||||
# Init overlap schedule
|
# Init overlap schedule
|
||||||
self.init_overlap()
|
self.init_overlap()
|
||||||
|
|
||||||
|
# Init Ngram Embedding
|
||||||
|
self.maybe_init_ngram_embedding()
|
||||||
|
|
||||||
# Init prefill kv split size when deterministic inference is enabled with various attention backends
|
# Init prefill kv split size when deterministic inference is enabled with various attention backends
|
||||||
self.init_deterministic_inference_config()
|
self.init_deterministic_inference_config()
|
||||||
|
|
||||||
@@ -1050,6 +1054,57 @@ class Scheduler(
|
|||||||
self.batch_record_buf = [None] * 2
|
self.batch_record_buf = [None] * 2
|
||||||
self.batch_record_ct = 0
|
self.batch_record_ct = 0
|
||||||
|
|
||||||
|
def maybe_init_ngram_embedding(self):
|
||||||
|
self.use_ngram_embedding = self.tp_worker.model_config.use_ngram_embedding
|
||||||
|
if self.use_ngram_embedding:
|
||||||
|
self.token_table = self.tp_worker.model_runner.token_table
|
||||||
|
hf_config = self.tp_worker.model_config.hf_config
|
||||||
|
self.ngram_embedding_n = hf_config.ngram_embedding_n
|
||||||
|
self.ngram_embedding_k = hf_config.ngram_embedding_k
|
||||||
|
|
||||||
|
def _maybe_prepare_ngram_embedding(
|
||||||
|
self, batch: Optional[ScheduleBatch]
|
||||||
|
) -> Optional[ScheduleBatch]:
|
||||||
|
"""Fill the token table for ngram embedding before a forward pass."""
|
||||||
|
if batch is None or not self.use_ngram_embedding:
|
||||||
|
return batch
|
||||||
|
batch.ne_token_table = self.token_table
|
||||||
|
if batch.forward_mode == ForwardMode.EXTEND:
|
||||||
|
all_tokens = []
|
||||||
|
column_starts = []
|
||||||
|
request_lengths = []
|
||||||
|
for req in batch.reqs:
|
||||||
|
start = len(req.prefix_indices)
|
||||||
|
end = start + req.extend_input_len
|
||||||
|
fill_ids = req.origin_input_ids + req.output_ids
|
||||||
|
if start == 0:
|
||||||
|
tokens = fill_ids[start:end]
|
||||||
|
column_starts.append(0)
|
||||||
|
elif start < self.ngram_embedding_n:
|
||||||
|
tokens = fill_ids[0:end]
|
||||||
|
column_starts.append(0)
|
||||||
|
else:
|
||||||
|
# Prepend n-1 tokens before prefix_len for n-gram context
|
||||||
|
tokens = fill_ids[start - self.ngram_embedding_n + 1 : end]
|
||||||
|
column_starts.append(start - self.ngram_embedding_n + 1)
|
||||||
|
all_tokens.extend(tokens)
|
||||||
|
request_lengths.append(len(tokens))
|
||||||
|
dtype = self.token_table.dtype
|
||||||
|
device = self.token_table.device
|
||||||
|
update_token_table(
|
||||||
|
ne_token_table=self.token_table,
|
||||||
|
tokens=torch.tensor(all_tokens, dtype=dtype, device=device),
|
||||||
|
row_indices=batch.req_pool_indices,
|
||||||
|
column_starts=torch.tensor(
|
||||||
|
column_starts, dtype=torch.int32, device=device
|
||||||
|
),
|
||||||
|
req_lens=torch.tensor(
|
||||||
|
request_lengths, dtype=torch.int32, device=device
|
||||||
|
),
|
||||||
|
ignore_tokens=None,
|
||||||
|
)
|
||||||
|
return batch
|
||||||
|
|
||||||
def init_deterministic_inference_config(self):
|
def init_deterministic_inference_config(self):
|
||||||
"""Initialize deterministic inference configuration for different attention backends."""
|
"""Initialize deterministic inference configuration for different attention backends."""
|
||||||
if not self.server_args.enable_deterministic_inference:
|
if not self.server_args.enable_deterministic_inference:
|
||||||
@@ -2035,6 +2090,9 @@ class Scheduler(
|
|||||||
# Handle DP attention and log stats
|
# Handle DP attention and log stats
|
||||||
ret = self.maybe_prepare_mlp_sync_batch(ret, need_sync=need_mlp_sync)
|
ret = self.maybe_prepare_mlp_sync_batch(ret, need_sync=need_mlp_sync)
|
||||||
|
|
||||||
|
# Handle ngram embedding
|
||||||
|
ret = self._maybe_prepare_ngram_embedding(ret)
|
||||||
|
|
||||||
if ret:
|
if ret:
|
||||||
set_schedule_time_batch(ret)
|
set_schedule_time_batch(ret)
|
||||||
|
|
||||||
|
|||||||
@@ -58,6 +58,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
CaptureHiddenMode,
|
CaptureHiddenMode,
|
||||||
ForwardBatch,
|
ForwardBatch,
|
||||||
ForwardMode,
|
ForwardMode,
|
||||||
|
NgramEmbeddingInfo,
|
||||||
PPProxyTensors,
|
PPProxyTensors,
|
||||||
compute_local_num_token_non_padded,
|
compute_local_num_token_non_padded,
|
||||||
enable_num_token_non_padded,
|
enable_num_token_non_padded,
|
||||||
@@ -126,6 +127,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
global_num_tokens_for_logprob_gpu: torch.Tensor
|
global_num_tokens_for_logprob_gpu: torch.Tensor
|
||||||
encoder_lens: Optional[torch.Tensor]
|
encoder_lens: Optional[torch.Tensor]
|
||||||
pp_proxy_tensors: Optional[Dict[str, torch.Tensor]]
|
pp_proxy_tensors: Optional[Dict[str, torch.Tensor]]
|
||||||
|
ngram_embedding_info: Optional["NgramEmbeddingInfo"]
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def create(
|
def create(
|
||||||
@@ -146,11 +148,12 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
num_tokens_per_bs: int,
|
num_tokens_per_bs: int,
|
||||||
cache_loc_dtype: torch.dtype,
|
cache_loc_dtype: torch.dtype,
|
||||||
enable_mamba_track: bool,
|
enable_mamba_track: bool,
|
||||||
|
ne_token_table: Optional[torch.Tensor] = None,
|
||||||
) -> "DecodeInputBuffers":
|
) -> "DecodeInputBuffers":
|
||||||
with torch.device(device):
|
with torch.device(device):
|
||||||
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
|
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
|
||||||
input_embeds = torch.zeros((max_num_token, hidden_size), dtype=dtype)
|
input_embeds = torch.zeros((max_num_token, hidden_size), dtype=dtype)
|
||||||
req_pool_indices = torch.zeros((max_bs,), dtype=torch.int32)
|
req_pool_indices = torch.zeros((max_bs,), dtype=torch.int64)
|
||||||
seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int32)
|
seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int32)
|
||||||
out_cache_loc = torch.zeros((max_num_token,), dtype=cache_loc_dtype)
|
out_cache_loc = torch.zeros((max_num_token,), dtype=cache_loc_dtype)
|
||||||
positions = torch.zeros((max_num_token,), dtype=torch.int64)
|
positions = torch.zeros((max_num_token,), dtype=torch.int64)
|
||||||
@@ -197,6 +200,18 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
|
global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||||
global_num_tokens_for_logprob_gpu = torch.zeros((1,), dtype=torch.int32)
|
global_num_tokens_for_logprob_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||||
|
|
||||||
|
ngram_embedding_info = (
|
||||||
|
NgramEmbeddingInfo(
|
||||||
|
token_table=ne_token_table,
|
||||||
|
column_starts=torch.zeros([max_bs], dtype=torch.int32),
|
||||||
|
req_lens=torch.ones([max_bs], dtype=torch.int32),
|
||||||
|
out_column_starts=torch.zeros([max_bs], dtype=torch.int32),
|
||||||
|
out_req_lens=torch.ones([max_bs], dtype=torch.int32),
|
||||||
|
)
|
||||||
|
if ne_token_table is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
# Keep seq_lens_cpu as a true CPU tensor, like the old implementation.
|
# Keep seq_lens_cpu as a true CPU tensor, like the old implementation.
|
||||||
seq_lens_cpu = torch.full(
|
seq_lens_cpu = torch.full(
|
||||||
(max_bs,),
|
(max_bs,),
|
||||||
@@ -223,6 +238,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
global_num_tokens_gpu=global_num_tokens_gpu,
|
global_num_tokens_gpu=global_num_tokens_gpu,
|
||||||
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
||||||
pp_proxy_tensors=pp_proxy_tensors,
|
pp_proxy_tensors=pp_proxy_tensors,
|
||||||
|
ngram_embedding_info=ngram_embedding_info,
|
||||||
)
|
)
|
||||||
|
|
||||||
def populate_from_forward_batch(
|
def populate_from_forward_batch(
|
||||||
@@ -263,6 +279,15 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
forward_batch.positions,
|
forward_batch.positions,
|
||||||
]
|
]
|
||||||
|
|
||||||
|
if self.ngram_embedding_info is not None:
|
||||||
|
ngram_embedding_info = forward_batch.ngram_embedding_info
|
||||||
|
self.ngram_embedding_info.column_starts[:raw_bs].copy_(
|
||||||
|
ngram_embedding_info.column_starts
|
||||||
|
)
|
||||||
|
self.ngram_embedding_info.req_lens[:raw_bs].copy_(
|
||||||
|
ngram_embedding_info.req_lens
|
||||||
|
)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.mamba_track_indices is not None
|
self.mamba_track_indices is not None
|
||||||
and forward_batch.mamba_track_indices is not None
|
and forward_batch.mamba_track_indices is not None
|
||||||
@@ -484,6 +509,11 @@ class CudaGraphRunner:
|
|||||||
self.enable_two_batch_overlap = (
|
self.enable_two_batch_overlap = (
|
||||||
model_runner.server_args.enable_two_batch_overlap
|
model_runner.server_args.enable_two_batch_overlap
|
||||||
)
|
)
|
||||||
|
self.use_ngram_embedding = model_runner.use_ngram_embedding
|
||||||
|
if self.use_ngram_embedding:
|
||||||
|
hf_config = model_runner.model_config.hf_config
|
||||||
|
self.ngram_embedding_n = hf_config.ngram_embedding_n
|
||||||
|
self.ngram_embedding_k = hf_config.ngram_embedding_k
|
||||||
self.speculative_algorithm = model_runner.server_args.speculative_algorithm
|
self.speculative_algorithm = model_runner.server_args.speculative_algorithm
|
||||||
self.enable_profile_cuda_graph = (
|
self.enable_profile_cuda_graph = (
|
||||||
model_runner.server_args.enable_profile_cuda_graph
|
model_runner.server_args.enable_profile_cuda_graph
|
||||||
@@ -582,6 +612,9 @@ class CudaGraphRunner:
|
|||||||
num_tokens_per_bs=self.num_tokens_per_bs,
|
num_tokens_per_bs=self.num_tokens_per_bs,
|
||||||
cache_loc_dtype=self._cache_loc_dtype(),
|
cache_loc_dtype=self._cache_loc_dtype(),
|
||||||
enable_mamba_track=enable_mamba_track,
|
enable_mamba_track=enable_mamba_track,
|
||||||
|
ne_token_table=(
|
||||||
|
model_runner.token_table if self.use_ngram_embedding else None
|
||||||
|
),
|
||||||
)
|
)
|
||||||
self.buffers.share_buffers()
|
self.buffers.share_buffers()
|
||||||
|
|
||||||
@@ -910,6 +943,9 @@ class CudaGraphRunner:
|
|||||||
global_forward_mode=self.capture_forward_mode,
|
global_forward_mode=self.capture_forward_mode,
|
||||||
lora_ids=lora_ids,
|
lora_ids=lora_ids,
|
||||||
)
|
)
|
||||||
|
if buffers.ngram_embedding_info is not None:
|
||||||
|
forward_batch.ngram_embedding_info = buffers.ngram_embedding_info.slice(bs)
|
||||||
|
|
||||||
self.tbo_plugin.capture_one_batch_size(forward_batch, num_tokens=num_tokens)
|
self.tbo_plugin.capture_one_batch_size(forward_batch, num_tokens=num_tokens)
|
||||||
|
|
||||||
if lora_ids is not None:
|
if lora_ids is not None:
|
||||||
|
|||||||
@@ -227,6 +227,48 @@ def compute_local_num_token_non_padded(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class NgramEmbeddingInfo:
|
||||||
|
"""Ngram embedding state for LongCat models."""
|
||||||
|
|
||||||
|
token_table: torch.Tensor
|
||||||
|
column_starts: torch.Tensor
|
||||||
|
req_lens: torch.Tensor
|
||||||
|
out_column_starts: torch.Tensor
|
||||||
|
out_req_lens: torch.Tensor
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(
|
||||||
|
cls,
|
||||||
|
token_table: torch.Tensor,
|
||||||
|
batch_size: int,
|
||||||
|
device: torch.device,
|
||||||
|
column_starts=None,
|
||||||
|
req_lens=None,
|
||||||
|
) -> NgramEmbeddingInfo:
|
||||||
|
info = cls(
|
||||||
|
token_table=token_table,
|
||||||
|
column_starts=torch.empty(batch_size, dtype=torch.int32, device=device),
|
||||||
|
req_lens=torch.empty(batch_size, dtype=torch.int32, device=device),
|
||||||
|
out_column_starts=torch.empty(batch_size, dtype=torch.int32, device=device),
|
||||||
|
out_req_lens=torch.empty(batch_size, dtype=torch.int32, device=device),
|
||||||
|
)
|
||||||
|
if column_starts is not None:
|
||||||
|
info.column_starts[:] = column_starts
|
||||||
|
if req_lens is not None:
|
||||||
|
info.req_lens[:] = req_lens
|
||||||
|
return info
|
||||||
|
|
||||||
|
def slice(self, bs: int) -> NgramEmbeddingInfo:
|
||||||
|
return NgramEmbeddingInfo(
|
||||||
|
token_table=self.token_table,
|
||||||
|
column_starts=self.column_starts[:bs],
|
||||||
|
req_lens=self.req_lens[:bs],
|
||||||
|
out_column_starts=self.out_column_starts[:bs],
|
||||||
|
out_req_lens=self.out_req_lens[:bs],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||||
"""Store all inputs of a forward pass."""
|
"""Store all inputs of a forward pass."""
|
||||||
@@ -374,6 +416,9 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
# For hidden states before normal
|
# For hidden states before normal
|
||||||
return_hidden_states_before_norm: bool = False
|
return_hidden_states_before_norm: bool = False
|
||||||
|
|
||||||
|
# For ngram embedding
|
||||||
|
ngram_embedding_info: Optional[NgramEmbeddingInfo] = None
|
||||||
|
|
||||||
# For dumper: request IDs for cross-step sequence tracking
|
# For dumper: request IDs for cross-step sequence tracking
|
||||||
rids: Optional[List[str]] = None
|
rids: Optional[List[str]] = None
|
||||||
|
|
||||||
@@ -511,6 +556,9 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
ret.extend_seq_lens_cpu = batch.extend_seq_lens
|
ret.extend_seq_lens_cpu = batch.extend_seq_lens
|
||||||
ret.extend_logprob_start_lens_cpu = batch.extend_logprob_start_lens
|
ret.extend_logprob_start_lens_cpu = batch.extend_logprob_start_lens
|
||||||
|
|
||||||
|
if model_runner.use_ngram_embedding:
|
||||||
|
ret._init_ngram_embedding_info(batch, model_runner, device)
|
||||||
|
|
||||||
if model_runner.model_is_mrope:
|
if model_runner.model_is_mrope:
|
||||||
if (
|
if (
|
||||||
ret.spec_info is not None
|
ret.spec_info is not None
|
||||||
@@ -602,6 +650,21 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
or self.contains_image_inputs()
|
or self.contains_image_inputs()
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _init_ngram_embedding_info(
|
||||||
|
self, batch: ModelWorkerBatch, model_runner: ModelRunner, device: torch.device
|
||||||
|
):
|
||||||
|
if self.forward_mode.is_decode():
|
||||||
|
column_starts, req_lens = self.seq_lens - 1, 1
|
||||||
|
else:
|
||||||
|
column_starts, req_lens = self.extend_prefix_lens, self.extend_seq_lens
|
||||||
|
self.ngram_embedding_info = NgramEmbeddingInfo.create(
|
||||||
|
batch.ne_token_table,
|
||||||
|
self.batch_size,
|
||||||
|
device,
|
||||||
|
column_starts=column_starts,
|
||||||
|
req_lens=req_lens,
|
||||||
|
)
|
||||||
|
|
||||||
def _compute_spec_mrope_positions(
|
def _compute_spec_mrope_positions(
|
||||||
self, model_runner: ModelRunner, batch: ModelWorkerBatch
|
self, model_runner: ModelRunner, batch: ModelWorkerBatch
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import dataclasses
|
||||||
from dataclasses import dataclass, fields
|
from dataclasses import dataclass, fields
|
||||||
from typing import Dict
|
from typing import Dict
|
||||||
|
|
||||||
@@ -38,7 +39,11 @@ class ForwardInputBuffers:
|
|||||||
|
|
||||||
if buffer is None:
|
if buffer is None:
|
||||||
continue
|
continue
|
||||||
elif isinstance(buffer, dict):
|
|
||||||
|
if dataclasses.is_dataclass(buffer):
|
||||||
|
buffer = vars(buffer)
|
||||||
|
|
||||||
|
if isinstance(buffer, dict):
|
||||||
for sub_name, sub_buffer in buffer.items():
|
for sub_name, sub_buffer in buffer.items():
|
||||||
assert isinstance(
|
assert isinstance(
|
||||||
sub_buffer, torch.Tensor
|
sub_buffer, torch.Tensor
|
||||||
@@ -50,6 +55,6 @@ class ForwardInputBuffers:
|
|||||||
else:
|
else:
|
||||||
assert isinstance(
|
assert isinstance(
|
||||||
buffer, torch.Tensor
|
buffer, torch.Tensor
|
||||||
), f"Field {name} is expected to be a torch.Tensor or a dict of torch.Tensor, but got {type(buffer)}."
|
), f"Field {name} is expected to be a torch.Tensor, a dict of torch.Tensor, or a dataclass of torch.Tensor, but got {type(buffer)}."
|
||||||
new_buffer = self._share_one_buffer(name, buffer)
|
new_buffer = self._share_one_buffer(name, buffer)
|
||||||
setattr(self, name, new_buffer)
|
setattr(self, name, new_buffer)
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ import torch
|
|||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
|
from sglang.jit_kernel.ngram_embedding import update_token_table
|
||||||
from sglang.srt.configs import (
|
from sglang.srt.configs import (
|
||||||
BailingHybridConfig,
|
BailingHybridConfig,
|
||||||
FalconH1Config,
|
FalconH1Config,
|
||||||
@@ -602,6 +603,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
# Init memory pool and attention backends
|
# Init memory pool and attention backends
|
||||||
self.init_memory_pool(pre_model_load_memory)
|
self.init_memory_pool(pre_model_load_memory)
|
||||||
|
|
||||||
|
# Init ngram embedding token table
|
||||||
|
self.maybe_init_ngram_embedding()
|
||||||
|
|
||||||
# Init routed experts capturer
|
# Init routed experts capturer
|
||||||
self.init_routed_experts_capturer()
|
self.init_routed_experts_capturer()
|
||||||
|
|
||||||
@@ -2108,6 +2112,49 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
with torch.inference_mode(), run_ctx or empty_context():
|
with torch.inference_mode(), run_ctx or empty_context():
|
||||||
run_once()
|
run_once()
|
||||||
|
|
||||||
|
def maybe_init_ngram_embedding(self):
|
||||||
|
self.use_ngram_embedding = self.model_config.use_ngram_embedding
|
||||||
|
if self.use_ngram_embedding:
|
||||||
|
from sglang.srt.layers.n_gram_embedding import NgramEmbedding
|
||||||
|
|
||||||
|
self.token_table = torch.empty(
|
||||||
|
self.req_to_token_pool.size,
|
||||||
|
self.model_config.context_len,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
chunked_prefill_size = self.server_args.chunked_prefill_size
|
||||||
|
assert (
|
||||||
|
chunked_prefill_size is not None and chunked_prefill_size > 0
|
||||||
|
), "Ngram embedding requires chunked prefill to be enabled (chunked_prefill_size > 0)"
|
||||||
|
for module in self.model.modules():
|
||||||
|
if isinstance(module, NgramEmbedding):
|
||||||
|
module.init_buffers(
|
||||||
|
self.max_running_requests, chunked_prefill_size, self.device
|
||||||
|
)
|
||||||
|
|
||||||
|
def maybe_update_ngram_token_table(
|
||||||
|
self,
|
||||||
|
next_token_ids: torch.Tensor,
|
||||||
|
forward_batch: "ForwardBatch",
|
||||||
|
):
|
||||||
|
"""Update the ngram embedding token table after sampling."""
|
||||||
|
ngram_embedding_info = forward_batch.ngram_embedding_info
|
||||||
|
if ngram_embedding_info is None:
|
||||||
|
return
|
||||||
|
ngram_embedding_info.out_column_starts[: forward_batch.batch_size] = (
|
||||||
|
forward_batch.seq_lens
|
||||||
|
)
|
||||||
|
ngram_embedding_info.out_req_lens[: forward_batch.batch_size] = 1
|
||||||
|
update_token_table(
|
||||||
|
ne_token_table=ngram_embedding_info.token_table,
|
||||||
|
tokens=next_token_ids.to(torch.int32),
|
||||||
|
row_indices=forward_batch.req_pool_indices,
|
||||||
|
column_starts=ngram_embedding_info.out_column_starts,
|
||||||
|
req_lens=torch.ones_like(ngram_embedding_info.out_column_starts),
|
||||||
|
ignore_tokens=None,
|
||||||
|
)
|
||||||
|
|
||||||
def init_device_graphs(self):
|
def init_device_graphs(self):
|
||||||
"""Capture device graphs."""
|
"""Capture device graphs."""
|
||||||
self.graph_runner = None
|
self.graph_runner = None
|
||||||
@@ -2570,6 +2617,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
else forward_batch.seq_lens - 1
|
else forward_batch.seq_lens - 1
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
self.maybe_update_ngram_token_table(next_token_ids, forward_batch)
|
||||||
return next_token_ids
|
return next_token_ids
|
||||||
|
|
||||||
def compute_logprobs_only(
|
def compute_logprobs_only(
|
||||||
|
|||||||
@@ -64,6 +64,7 @@ from sglang.srt.layers.moe.ep_moe.layer import DeepEPMoE, get_moe_impl_class
|
|||||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||||
from sglang.srt.layers.moe.topk import StandardTopKOutput, TopK
|
from sglang.srt.layers.moe.topk import StandardTopKOutput, TopK
|
||||||
from sglang.srt.layers.moe.utils import filter_moe_weight_param_global_expert
|
from sglang.srt.layers.moe.utils import filter_moe_weight_param_global_expert
|
||||||
|
from sglang.srt.layers.n_gram_embedding import NgramEmbedding
|
||||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz
|
from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz
|
||||||
from sglang.srt.layers.quantization.fp8_utils import (
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
@@ -329,7 +330,7 @@ class LongcatFlashDecoderLayer(nn.Module):
|
|||||||
q_lora_rank=config.q_lora_rank,
|
q_lora_rank=config.q_lora_rank,
|
||||||
kv_lora_rank=config.kv_lora_rank,
|
kv_lora_rank=config.kv_lora_rank,
|
||||||
rope_theta=config.rope_theta,
|
rope_theta=config.rope_theta,
|
||||||
rope_scaling=None,
|
rope_scaling=getattr(config, "rope_scaling", None),
|
||||||
max_position_embeddings=config.max_position_embeddings,
|
max_position_embeddings=config.max_position_embeddings,
|
||||||
quant_config=(
|
quant_config=(
|
||||||
None
|
None
|
||||||
@@ -500,6 +501,17 @@ class LongcatFlashModel(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.vocab_size = config.vocab_size
|
self.vocab_size = config.vocab_size
|
||||||
|
|
||||||
|
if config.use_ngram_embedding:
|
||||||
|
self.use_ngram_embedding = True
|
||||||
|
self.embed_tokens = NgramEmbedding(
|
||||||
|
num_embeddings=config.vocab_size,
|
||||||
|
embedding_dim=config.hidden_size,
|
||||||
|
over_embedding_m=config.ngram_embedding_m,
|
||||||
|
over_embedding_k=config.ngram_embedding_k,
|
||||||
|
over_embedding_n=config.ngram_embedding_n,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.use_ngram_embedding = False
|
||||||
self.embed_tokens = VocabParallelEmbedding(
|
self.embed_tokens = VocabParallelEmbedding(
|
||||||
config.vocab_size,
|
config.vocab_size,
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
@@ -540,6 +552,9 @@ class LongcatFlashModel(nn.Module):
|
|||||||
device=device,
|
device=device,
|
||||||
)
|
)
|
||||||
if input_embeds is None:
|
if input_embeds is None:
|
||||||
|
if self.use_ngram_embedding:
|
||||||
|
hidden_states = self.embed_tokens(input_ids, forward_batch)
|
||||||
|
else:
|
||||||
hidden_states = self.embed_tokens(input_ids)
|
hidden_states = self.embed_tokens(input_ids)
|
||||||
else:
|
else:
|
||||||
hidden_states = input_embeds
|
hidden_states = input_embeds
|
||||||
@@ -597,6 +612,7 @@ class LongcatFlashForCausalLM(nn.Module):
|
|||||||
self.model = LongcatFlashModel(
|
self.model = LongcatFlashModel(
|
||||||
config, quant_config, prefix=add_prefix("model", prefix)
|
config, quant_config, prefix=add_prefix("model", prefix)
|
||||||
)
|
)
|
||||||
|
self.use_ngram_embedding = config.use_ngram_embedding
|
||||||
self.lm_head = ParallelLMHead(
|
self.lm_head = ParallelLMHead(
|
||||||
config.vocab_size,
|
config.vocab_size,
|
||||||
config.hidden_size,
|
config.hidden_size,
|
||||||
@@ -874,7 +890,6 @@ class LongcatFlashForCausalLM(nn.Module):
|
|||||||
self.config.q_lora_rank is not None
|
self.config.q_lora_rank is not None
|
||||||
)
|
)
|
||||||
cached_a_proj = {} if fuse_qkv_a_proj else None
|
cached_a_proj = {} if fuse_qkv_a_proj else None
|
||||||
|
|
||||||
with concurrent.futures.ThreadPoolExecutor() as executor:
|
with concurrent.futures.ThreadPoolExecutor() as executor:
|
||||||
futures = []
|
futures = []
|
||||||
params_dict = dict(self.named_parameters())
|
params_dict = dict(self.named_parameters())
|
||||||
@@ -883,6 +898,12 @@ class LongcatFlashForCausalLM(nn.Module):
|
|||||||
use_async_loading = should_async_load(loaded_weight)
|
use_async_loading = should_async_load(loaded_weight)
|
||||||
if "mtp" in name:
|
if "mtp" in name:
|
||||||
continue
|
continue
|
||||||
|
if self.use_ngram_embedding:
|
||||||
|
if ".embed_tokens." in name:
|
||||||
|
name = "model.embed_tokens.word_embeder.weight"
|
||||||
|
if ".ngram_embeddings" in name:
|
||||||
|
self.model.embed_tokens.load_weight(None, name, loaded_weight)
|
||||||
|
continue
|
||||||
weight_names.append(name)
|
weight_names.append(name)
|
||||||
if "rotary_emb.inv_freq" in name:
|
if "rotary_emb.inv_freq" in name:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -105,7 +105,7 @@ class EAGLEDraftCudaGraphRunner:
|
|||||||
# Graph inputs
|
# Graph inputs
|
||||||
with torch.device(model_runner.device):
|
with torch.device(model_runner.device):
|
||||||
input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
||||||
req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32)
|
req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int64)
|
||||||
out_cache_loc = torch.zeros(
|
out_cache_loc = torch.zeros(
|
||||||
(self.max_num_token * self.speculative_num_steps,),
|
(self.max_num_token * self.speculative_num_steps,),
|
||||||
dtype=self._cache_loc_dtype(),
|
dtype=self._cache_loc_dtype(),
|
||||||
|
|||||||
@@ -110,7 +110,7 @@ class EAGLEDraftExtendCudaGraphRunner:
|
|||||||
# Graph inputs
|
# Graph inputs
|
||||||
with torch.device(model_runner.device):
|
with torch.device(model_runner.device):
|
||||||
input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
input_ids = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
||||||
req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int32)
|
req_pool_indices = torch.zeros((self.max_bs,), dtype=torch.int64)
|
||||||
out_cache_loc = torch.ones(
|
out_cache_loc = torch.ones(
|
||||||
(self.max_num_token,), dtype=self._cache_loc_dtype()
|
(self.max_num_token,), dtype=self._cache_loc_dtype()
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -669,7 +669,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
|
|||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
)
|
)
|
||||||
self.cuda_graph_buffers["req_pool_indices"] = torch.zeros(
|
self.cuda_graph_buffers["req_pool_indices"] = torch.zeros(
|
||||||
(self.max_bs,), dtype=torch.int32
|
(self.max_bs,), dtype=torch.int64
|
||||||
)
|
)
|
||||||
self.cuda_graph_buffers["accept_length"] = torch.full(
|
self.cuda_graph_buffers["accept_length"] = torch.full(
|
||||||
(self.max_bs,), 1, dtype=torch.int32
|
(self.max_bs,), 1, dtype=torch.int32
|
||||||
|
|||||||
@@ -342,6 +342,14 @@ def get_config(
|
|||||||
"patch_size": 14,
|
"patch_size": 14,
|
||||||
}
|
}
|
||||||
config.vision_config = SiglipVisionConfig(**vision_config)
|
config.vision_config = SiglipVisionConfig(**vision_config)
|
||||||
|
|
||||||
|
if config.architectures in [
|
||||||
|
["LongcatCausalLM"],
|
||||||
|
["LongcatFlashForCausalLM"],
|
||||||
|
["LongcatFlashNgramForCausalLM"],
|
||||||
|
]:
|
||||||
|
config.model_type = "longcat_flash"
|
||||||
|
|
||||||
text_config = get_hf_text_config(config=config)
|
text_config = get_hf_text_config(config=config)
|
||||||
|
|
||||||
if isinstance(model, str) and text_config is not None:
|
if isinstance(model, str) and text_config is not None:
|
||||||
@@ -389,6 +397,9 @@ def get_config(
|
|||||||
if config.model_type == "multi_modality":
|
if config.model_type == "multi_modality":
|
||||||
config.update({"architectures": ["MultiModalityCausalLM"]})
|
config.update({"architectures": ["MultiModalityCausalLM"]})
|
||||||
|
|
||||||
|
if config.model_type == "longcat_flash":
|
||||||
|
config.update({"architectures": ["LongcatFlashForCausalLM"]})
|
||||||
|
|
||||||
if model_override_args:
|
if model_override_args:
|
||||||
config.update(model_override_args)
|
config.update(model_override_args)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user