dsv4.1: Engram module and request history support (#39666)
Co-authored-by: BBuf <1182563586@qq.com>
This commit is contained in:
@@ -688,6 +688,10 @@ class ModelConfig:
|
||||
self.hf_config.architectures
|
||||
)
|
||||
self.use_ngram_embedding = getattr(self.hf_config, "use_ngram_embedding", False)
|
||||
self.ngram_embedding_n = (
|
||||
self.hf_config.ngram_embedding_n if self.use_ngram_embedding else 0
|
||||
)
|
||||
self.use_engram = bool(getattr(self.hf_config, "engram_layer_ids", ()))
|
||||
# A multimodal arch is piecewise-incompatible until its LM prefill is validated.
|
||||
self.is_piecewise_cuda_graph_disabled_model = (
|
||||
is_piecewise_cuda_graph_disabled_model(self.hf_config.architectures)
|
||||
|
||||
@@ -1474,6 +1474,15 @@ class Envs:
|
||||
# picks the buffer dtype and this switch is inert.
|
||||
SGLANG_DSV4_UNIFIED_KV_FP8 = EnvBool(False)
|
||||
|
||||
# DeepSeek-V4.1 engram host table: keep the tables in host memory (layout
|
||||
# below) and gather rows from the GPU instead of sharding them over HBM.
|
||||
SGLANG_ENABLE_DSV41_ENGRAM_HOST_TABLE = EnvBool(False)
|
||||
# "shared" is one buffer for the whole TP group, mapped by every rank, with no
|
||||
# lookup all-reduce (the ranks must share a PID namespace); "per_rank" is one
|
||||
# anonymous mapping per rank holding only its rows, gathered with the
|
||||
# all-reduce, and the only layout that gets huge pages without shmem THP.
|
||||
SGLANG_DSV41_ENGRAM_HOST_TABLE_LAYOUT = EnvStr("shared")
|
||||
|
||||
# Kernels and indexer
|
||||
SGLANG_OPT_DEEPGEMM_HC_PRENORM = EnvBool(True)
|
||||
SGLANG_OPT_USE_TILELANG_MHC_PRE = EnvBool(True)
|
||||
|
||||
@@ -0,0 +1,927 @@
|
||||
"""Engram: gated n-gram hash memory added to the hc residual stream.
|
||||
|
||||
The token map, hash multipliers and prime layout must match
|
||||
sglang.kernels.ops.embeddings.engram_hash to produce identical hash ids.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ctypes
|
||||
import errno
|
||||
import glob
|
||||
import logging
|
||||
import mmap
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
import msgspec
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from sglang.kernels.ops.attention.dsv4.torch_quant import FP8_BLOCK_SIZE
|
||||
from sglang.kernels.ops.embeddings.engram_gate import fused_engram_gate
|
||||
from sglang.kernels.ops.embeddings.engram_gather import engram_gather
|
||||
from sglang.kernels.ops.embeddings.engram_hash import (
|
||||
MODE_DECODE,
|
||||
MODE_EXTEND,
|
||||
MODE_VERIFY,
|
||||
engram_commit_history,
|
||||
engram_hash_ids,
|
||||
engram_hash_ids_and_commit,
|
||||
)
|
||||
from sglang.srt.distributed import tensor_model_parallel_all_reduce
|
||||
from sglang.srt.distributed.parallel_state import get_tp_group
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
attn_cp_all_gather_into_tensor,
|
||||
dp_gather_replicate,
|
||||
dp_reduce_scatter_tensor,
|
||||
dp_scatter,
|
||||
get_attention_dp_size,
|
||||
get_global_dp_buffer_len,
|
||||
is_dp_gatherv_active,
|
||||
)
|
||||
from sglang.srt.layers.linear import ReplicatedLinear
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.managers.schedule_batch import MM_PAD_SHIFT_VALUE
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.runtime_context import get_model, get_parallel, get_serving
|
||||
from sglang.srt.utils import add_prefix, is_cuda
|
||||
from sglang.srt.utils.hf_transformers.tokenizer import get_tokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_MILLER_RABIN_WITNESSES = (2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37)
|
||||
|
||||
|
||||
def _cuda_kernels(t: torch.Tensor) -> bool:
|
||||
"""True where the Triton kernels apply; ROCm and CPU take the torch paths."""
|
||||
return t.is_cuda and is_cuda()
|
||||
|
||||
|
||||
def _is_prime(n: int) -> bool:
|
||||
"""Deterministic Miller-Rabin; exact for n < 3.3e24 with these witnesses."""
|
||||
if n < 2:
|
||||
return False
|
||||
for p in _MILLER_RABIN_WITNESSES:
|
||||
if n % p == 0:
|
||||
return n == p
|
||||
d, r = n - 1, 0
|
||||
while d % 2 == 0:
|
||||
d, r = d // 2, r + 1
|
||||
for a in _MILLER_RABIN_WITNESSES:
|
||||
x = pow(a, d, n)
|
||||
if x == 1 or x == n - 1:
|
||||
continue
|
||||
for _ in range(r - 1):
|
||||
x = x * x % n
|
||||
if x == n - 1:
|
||||
break
|
||||
else:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _find_next_prime(start: int, seen_primes: set[int]) -> int:
|
||||
candidate = start + 1
|
||||
while not _is_prime(candidate) or candidate in seen_primes:
|
||||
candidate += 1
|
||||
return candidate
|
||||
|
||||
|
||||
def build_compressed_token_map(tokenizer) -> tuple[list[int], int]:
|
||||
"""Token id -> compressed id, plus the compressed vocab size. The size feeds every
|
||||
hash multiplier, so a mismatch rehashes the whole table."""
|
||||
from tokenizers import Regex, normalizers
|
||||
|
||||
# A private-use sentinel keeps a lone-space token from collapsing to "" under Strip().
|
||||
sentinel = "\ue000"
|
||||
normalizer = normalizers.Sequence(
|
||||
[
|
||||
normalizers.NFKC(),
|
||||
normalizers.NFD(),
|
||||
normalizers.StripAccents(),
|
||||
normalizers.Lowercase(),
|
||||
normalizers.Replace(Regex(r"[ \t\r\n]+"), " "),
|
||||
normalizers.Replace(Regex(r"^ $"), sentinel),
|
||||
normalizers.Strip(),
|
||||
normalizers.Replace(sentinel, " "),
|
||||
]
|
||||
)
|
||||
backend = tokenizer.backend_tokenizer
|
||||
key_to_new: dict[str, int] = {}
|
||||
lookup = [0] * len(tokenizer)
|
||||
for token_id in range(len(tokenizer)):
|
||||
text = backend.decode([token_id], skip_special_tokens=False)
|
||||
if "\ufffd" in text:
|
||||
# A partial UTF-8 byte token has nothing to normalize; key it by its raw form.
|
||||
key = backend.id_to_token(token_id)
|
||||
else:
|
||||
normalized = normalizer.normalize_str(text)
|
||||
key = normalized if normalized else text
|
||||
new_id = key_to_new.get(key)
|
||||
if new_id is None:
|
||||
new_id = len(key_to_new)
|
||||
key_to_new[key] = new_id
|
||||
lookup[token_id] = new_id
|
||||
return lookup, len(key_to_new)
|
||||
|
||||
|
||||
def compute_hash_multipliers(
|
||||
layer_ids: tuple[int, ...], max_ngram_size: int, vocab_size: int
|
||||
) -> torch.Tensor:
|
||||
"""One odd multiplier per (layer, lookback) from a per-layer RNG, bounded so that
|
||||
token_id * multiplier fits in int64."""
|
||||
max_long = np.iinfo(np.int64).max
|
||||
multiplier_bound = max(1, (max_long // vocab_size) // 2)
|
||||
rows = []
|
||||
for layer_id in layer_ids:
|
||||
generator = np.random.default_rng(10007 * layer_id)
|
||||
values = generator.integers(
|
||||
low=0, high=multiplier_bound, size=(max_ngram_size,), dtype=np.int64
|
||||
)
|
||||
rows.append(torch.tensor(values * 2 + 1))
|
||||
return torch.stack(rows)
|
||||
|
||||
|
||||
class EngramLayout(msgspec.Struct, frozen=True):
|
||||
max_ngram_size: int
|
||||
layer_ids: tuple[int, ...]
|
||||
num_embeddings: tuple[int, ...]
|
||||
primes: tuple[tuple[tuple[int, ...], ...], ...] # [layer][n-gram size][head]
|
||||
n_heads: int
|
||||
head_dim: int
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config) -> Optional[EngramLayout]:
|
||||
"""Primes are drawn in (layer, n-gram size, head) order from one shared
|
||||
ascending sequence starting above engram_vocab_size - 1."""
|
||||
layer_ids = tuple(config.engram_layer_ids)
|
||||
if not layer_ids:
|
||||
return None
|
||||
max_ngram_size, n_heads = config.engram_max_ngram_size, config.engram_n_heads
|
||||
vocab_size = config.engram_vocab_size
|
||||
primes, seen = [], set()
|
||||
for _ in layer_ids:
|
||||
per_ngram = []
|
||||
for _ in range(max_ngram_size - 1):
|
||||
sizes, current = [], vocab_size - 1
|
||||
for _ in range(n_heads):
|
||||
current = _find_next_prime(current, seen)
|
||||
seen.add(current)
|
||||
sizes.append(current)
|
||||
per_ngram.append(tuple(sizes))
|
||||
primes.append(tuple(per_ngram))
|
||||
return cls(
|
||||
max_ngram_size=max_ngram_size,
|
||||
layer_ids=layer_ids,
|
||||
num_embeddings=tuple(config.engram_num_embeddings),
|
||||
primes=tuple(primes),
|
||||
n_heads=n_heads,
|
||||
head_dim=config.engram_head_dim,
|
||||
)
|
||||
|
||||
|
||||
def compute_engram_hash_ids(
|
||||
tokens: torch.Tensor,
|
||||
blocked: torch.Tensor,
|
||||
pad_id: int,
|
||||
token_map: torch.Tensor,
|
||||
multipliers: torch.Tensor,
|
||||
primes: torch.Tensor,
|
||||
offsets: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""tokens [T, n]: column 0 is the token itself, column s its s-th predecessor;
|
||||
blocked [T, n] marks look-back that ran off the sequence start. Returns
|
||||
[T, n_engram_layers, (n - 1) * n_heads] row ids into each layer's table."""
|
||||
compressed = torch.where(blocked, pad_id, token_map[tokens])
|
||||
products = compressed.unsqueeze(1) * multipliers
|
||||
# After step i the running xor is the (i + 1)-gram hash, bucketed by that
|
||||
# n-gram size's primes.
|
||||
rolling, hashes = products[..., 0], []
|
||||
for i in range(1, tokens.shape[-1]):
|
||||
rolling = torch.bitwise_xor(rolling, products[..., i])
|
||||
hashes.append(rolling.unsqueeze(-1) % primes[:, i - 1])
|
||||
return torch.cat(hashes, dim=-1) + offsets
|
||||
|
||||
|
||||
class EngramHasher(nn.Module):
|
||||
"""Hash ids for every token of a forward batch, [T, n_engram_layers, n_hash_cols]."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
layout: EngramLayout,
|
||||
tokenizer,
|
||||
pad_id: int,
|
||||
compressed_vocab_size: int,
|
||||
):
|
||||
super().__init__()
|
||||
self.max_ngram_size = layout.max_ngram_size
|
||||
token_map, vocab_size = build_compressed_token_map(tokenizer)
|
||||
assert vocab_size == compressed_vocab_size, (
|
||||
f"the tokenizer normalizes to {vocab_size} distinct tokens but the config "
|
||||
f"expects {compressed_vocab_size}; every hash multiplier depends on it"
|
||||
)
|
||||
self.pad_id = token_map[pad_id]
|
||||
flat = [
|
||||
[p for per_ngram in layer for p in per_ngram] for layer in layout.primes
|
||||
]
|
||||
offsets = np.array([np.cumsum([0, *sizes[:-1]]) for sizes in flat])
|
||||
multipliers = compute_hash_multipliers(
|
||||
layout.layer_ids, layout.max_ngram_size, vocab_size
|
||||
)
|
||||
self.register_buffer("token_map", torch.tensor(token_map), persistent=False)
|
||||
self.register_buffer("multipliers", multipliers, persistent=False)
|
||||
self.register_buffer("primes", torch.tensor(layout.primes), persistent=False)
|
||||
self.register_buffer("offsets", torch.tensor(offsets), persistent=False)
|
||||
self.image_token_id: Optional[int] = None
|
||||
self.history: Optional[torch.Tensor] = None
|
||||
self.pad_row = 0
|
||||
|
||||
def init_history(self, num_req_slots: int, device) -> None:
|
||||
"""Allocate oldest-first history with a spare row for graph padding."""
|
||||
self.history = torch.zeros(
|
||||
num_req_slots + 1,
|
||||
self.max_ngram_size - 1,
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
self.pad_row = num_req_slots
|
||||
|
||||
@classmethod
|
||||
def from_config(
|
||||
cls, config, layout: EngramLayout, *, image_token_id: Optional[int] = None
|
||||
) -> EngramHasher:
|
||||
# The compressed token map is built with the HF normalizers, so the HF
|
||||
# tokenizer backend is used here whatever the serving backend is.
|
||||
tokenizer = get_tokenizer(
|
||||
get_serving().tokenizer_path,
|
||||
tokenizer_mode=get_serving().tokenizer_mode,
|
||||
trust_remote_code=get_model().trust_remote_code,
|
||||
revision=get_model().revision,
|
||||
tokenizer_backend="huggingface",
|
||||
)
|
||||
result = cls(
|
||||
layout,
|
||||
tokenizer,
|
||||
config.engram_pad_token_id,
|
||||
config.engram_compressed_vocab_size,
|
||||
)
|
||||
result.image_token_id = image_token_id
|
||||
return result
|
||||
|
||||
def forward(
|
||||
self, input_ids: torch.Tensor, forward_batch: ForwardBatch
|
||||
) -> torch.Tensor:
|
||||
assert self.history is not None, "EngramHasher.init_history was not called"
|
||||
n = self.max_ngram_size
|
||||
num_tokens = input_ids.shape[0]
|
||||
if num_tokens == 0:
|
||||
return torch.empty(
|
||||
(0, self.primes.shape[0], self.offsets.shape[1]),
|
||||
dtype=torch.int64,
|
||||
device=input_ids.device,
|
||||
)
|
||||
mode = forward_batch.forward_mode
|
||||
req_slots = forward_batch.req_pool_indices
|
||||
bs = req_slots.shape[0]
|
||||
device = input_ids.device
|
||||
# Tokens at or past num_real are graph padding. History comes from
|
||||
# self.history via req_slots unless the scheduler supplied this extend's rows.
|
||||
num_real, block, row, starts = num_tokens, 1, None, None
|
||||
history, hist_via_slots = self.history, True
|
||||
if mode.is_decode():
|
||||
kmode = MODE_DECODE
|
||||
commit_rows, commit_last = req_slots, None
|
||||
elif mode.is_target_verify():
|
||||
block = int(forward_batch.spec_info.draft_token_num)
|
||||
assert num_tokens == bs * block, (
|
||||
"engram target-verify expects one equal block per request, got "
|
||||
f"{num_tokens} tokens for {bs} requests of {block}"
|
||||
)
|
||||
kmode = MODE_VERIFY
|
||||
commit_rows = commit_last = None
|
||||
else:
|
||||
assert mode.is_extend(), (
|
||||
f"engram serves extend, target-verify and decode, not {mode}"
|
||||
)
|
||||
lens = forward_batch.extend_seq_lens.to(torch.int64)
|
||||
starts = forward_batch.extend_start_loc.to(torch.int64)
|
||||
row = torch.repeat_interleave(torch.arange(bs, device=device), lens)
|
||||
num_real = row.shape[0]
|
||||
kmode = MODE_EXTEND
|
||||
if forward_batch.engram_history is not None:
|
||||
history, hist_via_slots = forward_batch.engram_history, False
|
||||
commit_rows = torch.where(lens > 0, req_slots, self.pad_row)
|
||||
commit_last = (starts + lens - 1).clamp(0, num_tokens - 1)
|
||||
|
||||
if _cuda_kernels(input_ids):
|
||||
if kmode == MODE_DECODE:
|
||||
# out_cache_loc 0 marks the CUDA-graph padded rows that must not commit.
|
||||
assert forward_batch.out_cache_loc is not None
|
||||
return engram_hash_ids_and_commit(
|
||||
input_ids,
|
||||
forward_batch.positions,
|
||||
history=self.history,
|
||||
req_slots=req_slots,
|
||||
out_cache_loc=forward_batch.out_cache_loc,
|
||||
token_map=self.token_map,
|
||||
multipliers=self.multipliers,
|
||||
primes=self.primes,
|
||||
offsets=self.offsets,
|
||||
pad_id=self.pad_id,
|
||||
image_token_id=self.image_token_id,
|
||||
mm_pad_shift=MM_PAD_SHIFT_VALUE,
|
||||
)
|
||||
hash_ids, tokens = engram_hash_ids(
|
||||
input_ids,
|
||||
forward_batch.positions,
|
||||
mode=kmode,
|
||||
history=history,
|
||||
token_map=self.token_map,
|
||||
multipliers=self.multipliers,
|
||||
primes=self.primes,
|
||||
offsets=self.offsets,
|
||||
pad_id=self.pad_id,
|
||||
num_real=num_real,
|
||||
req_slots=req_slots if hist_via_slots else None,
|
||||
block=block,
|
||||
row=row,
|
||||
starts=starts,
|
||||
image_token_id=self.image_token_id,
|
||||
mm_pad_shift=MM_PAD_SHIFT_VALUE,
|
||||
)
|
||||
else:
|
||||
hash_ids, tokens = self._torch_hash_ids(
|
||||
input_ids,
|
||||
forward_batch.positions,
|
||||
kmode,
|
||||
history[req_slots] if hist_via_slots else history,
|
||||
num_real,
|
||||
block,
|
||||
row,
|
||||
starts,
|
||||
)
|
||||
if commit_rows is not None:
|
||||
# Padded rows must not overwrite a live request's history.
|
||||
last_tokens = tokens if commit_last is None else tokens[commit_last]
|
||||
out_loc = forward_batch.out_cache_loc
|
||||
if out_loc is not None:
|
||||
if commit_last is not None:
|
||||
out_loc = out_loc[commit_last]
|
||||
commit_rows = torch.where(out_loc == 0, self.pad_row, commit_rows)
|
||||
self.history[commit_rows] = (
|
||||
last_tokens[:, : n - 1].flip(-1).to(self.history.dtype)
|
||||
)
|
||||
return hash_ids
|
||||
|
||||
def _torch_hash_ids(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
kmode: int,
|
||||
history: torch.Tensor,
|
||||
num_real: int,
|
||||
block: int,
|
||||
row: Optional[torch.Tensor],
|
||||
starts: Optional[torch.Tensor],
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Non-CUDA fallback of the hash kernel: predecessor table [T, n] and hash ids.
|
||||
``history`` is already the per-request [bs, n - 1] rows of this batch."""
|
||||
n = self.max_ngram_size
|
||||
num_tokens = input_ids.shape[0]
|
||||
device = input_ids.device
|
||||
input_ids = input_ids.to(torch.int64)
|
||||
shifts = torch.arange(n, device=device)
|
||||
if kmode == MODE_DECODE:
|
||||
tokens = torch.cat(
|
||||
[input_ids.unsqueeze(-1), history.to(torch.int64).flip(-1)], dim=-1
|
||||
)
|
||||
else:
|
||||
t = torch.arange(num_real, device=device)
|
||||
if kmode == MODE_VERIFY:
|
||||
row = t // block
|
||||
offset = t - row * block
|
||||
else:
|
||||
offset = t - starts[row]
|
||||
# Predecessor at shift s of token t: an earlier token of the same run
|
||||
# when s <= offset, else history[row, -(s - offset)].
|
||||
from_batch = shifts.unsqueeze(0) <= offset.unsqueeze(-1)
|
||||
in_batch = input_ids[(t.unsqueeze(-1) - shifts).clamp_min(0)]
|
||||
hist_col = (n - 2 - (shifts.unsqueeze(0) - offset.unsqueeze(-1) - 1)).clamp(
|
||||
0, n - 2
|
||||
)
|
||||
from_hist = history.to(torch.int64)[row].gather(1, hist_col)
|
||||
tokens = torch.where(from_batch, in_batch, from_hist)
|
||||
if num_real < num_tokens:
|
||||
tokens = torch.cat([tokens, tokens.new_zeros(num_tokens - num_real, n)])
|
||||
positions = positions.to(torch.int64)
|
||||
blocked = positions.unsqueeze(-1) < shifts
|
||||
if num_real < num_tokens:
|
||||
blocked[num_real:] = True
|
||||
if self.image_token_id is not None:
|
||||
# Scheduler-provided history still carries the multimodal pad ids.
|
||||
tokens = tokens.masked_fill(
|
||||
tokens >= MM_PAD_SHIFT_VALUE, self.image_token_id
|
||||
)
|
||||
# Once a lookback hits an image, every older predecessor is PAD.
|
||||
blocked = (
|
||||
(blocked | (tokens == self.image_token_id))
|
||||
.to(torch.int32)
|
||||
.cummax(-1)
|
||||
.values.bool()
|
||||
)
|
||||
hash_ids = compute_engram_hash_ids(
|
||||
tokens,
|
||||
blocked,
|
||||
self.pad_id,
|
||||
self.token_map,
|
||||
self.multipliers,
|
||||
self.primes,
|
||||
self.offsets,
|
||||
)
|
||||
return hash_ids, tokens
|
||||
|
||||
def commit_after_verify(
|
||||
self,
|
||||
verify_ids_2d: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
commit_lens: torch.Tensor,
|
||||
) -> None:
|
||||
"""Commit anchor + accepted drafts; the bonus is the next block's anchor."""
|
||||
assert self.history is not None, "EngramHasher.init_history was not called"
|
||||
if _cuda_kernels(self.history):
|
||||
engram_commit_history(
|
||||
self.history, verify_ids_2d, req_pool_indices, commit_lens
|
||||
)
|
||||
return
|
||||
n1 = self.max_ngram_size - 1
|
||||
req = req_pool_indices.to(torch.int64)
|
||||
window = torch.cat(
|
||||
[self.history[req], verify_ids_2d.to(self.history.dtype)], dim=1
|
||||
)
|
||||
cols = commit_lens.to(torch.int64).unsqueeze(-1) + torch.arange(
|
||||
n1, device=window.device
|
||||
)
|
||||
self.history[req] = window.gather(1, cols)
|
||||
|
||||
|
||||
_THP_DIR = "/sys/kernel/mm/transparent_hugepage"
|
||||
|
||||
|
||||
def _thp_mode(knob: str) -> str:
|
||||
"""Active mode of a transparent_hugepage sysfs knob ("" if unreadable)."""
|
||||
try:
|
||||
with open(f"{_THP_DIR}/{knob}") as f:
|
||||
m = re.search(r"\[(\w+)\]", f.read())
|
||||
return m.group(1) if m else ""
|
||||
except OSError:
|
||||
return ""
|
||||
|
||||
|
||||
def _huge_pages_backing(addr: int) -> tuple[int, int]:
|
||||
"""(mapped_kB, huge_kB) of the VMA holding addr, from /proc/self/smaps.
|
||||
The only evidence that the kernel really handed out huge pages."""
|
||||
mapped = huge = 0
|
||||
inside = False
|
||||
try:
|
||||
with open("/proc/self/smaps") as f:
|
||||
for line in f:
|
||||
m = re.match(r"^([0-9a-f]+)-([0-9a-f]+) ", line)
|
||||
if m:
|
||||
if inside:
|
||||
break
|
||||
inside = int(m.group(1), 16) <= addr < int(m.group(2), 16)
|
||||
elif inside:
|
||||
key, _, rest = line.partition(":")
|
||||
if key == "Rss":
|
||||
mapped = int(rest.split()[0])
|
||||
elif key in ("AnonHugePages", "ShmemPmdMapped", "FilePmdMapped"):
|
||||
huge += int(rest.split()[0])
|
||||
except OSError:
|
||||
pass
|
||||
return mapped, huge
|
||||
|
||||
|
||||
_page_cache_dropped = False
|
||||
|
||||
|
||||
def drop_checkpoint_page_cache() -> tuple[int, int]:
|
||||
"""posix_fadvise(DONTNEED) on the checkpoint files; returns (files, bytes)."""
|
||||
try:
|
||||
model_path = get_model().model_path
|
||||
except (ValueError, AttributeError):
|
||||
# No published runtime context (unit tests, offline tools): nothing to drop.
|
||||
return 0, 0
|
||||
files, nbytes = 0, 0
|
||||
for f in sorted(glob.glob(os.path.join(model_path, "*.safetensors"))):
|
||||
try:
|
||||
fd = os.open(f, os.O_RDONLY)
|
||||
except OSError:
|
||||
continue
|
||||
try:
|
||||
nbytes += os.fstat(fd).st_size
|
||||
os.posix_fadvise(fd, 0, 0, os.POSIX_FADV_DONTNEED)
|
||||
files += 1
|
||||
finally:
|
||||
os.close(fd)
|
||||
return files, nbytes
|
||||
|
||||
|
||||
def _drop_page_cache_once(reason: str) -> None:
|
||||
global _page_cache_dropped
|
||||
if _page_cache_dropped:
|
||||
return
|
||||
_page_cache_dropped = True
|
||||
files, nbytes = drop_checkpoint_page_cache()
|
||||
logger.info(
|
||||
"engram host table: dropped the page cache of %d checkpoint files (%.0f GiB) %s",
|
||||
files,
|
||||
nbytes / 2**30,
|
||||
reason,
|
||||
)
|
||||
|
||||
|
||||
class _HostTable:
|
||||
"""Host-memory backing for one engram table ('shared' or 'per_rank' layout).
|
||||
|
||||
Lives for the whole process: the mapping, the memfd and the cudaHostRegister
|
||||
pin are never released because the table is read by every forward.
|
||||
"""
|
||||
|
||||
def __init__(self, layout: str, nbytes: int, name: str, group):
|
||||
if layout not in ("shared", "per_rank"):
|
||||
raise ValueError(
|
||||
f"Invalid SGLANG_DSV41_ENGRAM_HOST_TABLE_LAYOUT={layout!r}; expected "
|
||||
"'shared' or 'per_rank'"
|
||||
)
|
||||
self.layout = layout
|
||||
self.nbytes = nbytes
|
||||
self.group = group
|
||||
self.dirty = False
|
||||
if layout == "shared":
|
||||
self.fd = self._open_shared_fd(nbytes, name)
|
||||
self.mm = mmap.mmap(
|
||||
self.fd,
|
||||
nbytes,
|
||||
flags=mmap.MAP_SHARED,
|
||||
prot=mmap.PROT_READ | mmap.PROT_WRITE,
|
||||
)
|
||||
else:
|
||||
self.fd = None
|
||||
self.mm = mmap.mmap(
|
||||
-1,
|
||||
nbytes,
|
||||
flags=mmap.MAP_PRIVATE | mmap.MAP_ANONYMOUS,
|
||||
prot=mmap.PROT_READ | mmap.PROT_WRITE,
|
||||
)
|
||||
# Advisory before the first touch: pages are allocated huge at fault time.
|
||||
self.mm.madvise(mmap.MADV_HUGEPAGE)
|
||||
self.bytes = torch.frombuffer(self.mm, dtype=torch.uint8)
|
||||
if layout == "per_rank":
|
||||
# Cached checkpoint pages, left by a previous server or by the loader,
|
||||
# make the 512 MiB huge-page faults fall back, so empty them first.
|
||||
_drop_page_cache_once("before pre-faulting the per-rank shard")
|
||||
np.frombuffer(self.mm, dtype=np.uint8)[:: mmap.PAGESIZE] = 0
|
||||
if layout == "shared":
|
||||
# Every rank holds the fd before rank 0 continues; the /proc path only
|
||||
# resolves while rank 0 keeps its descriptor.
|
||||
group.barrier()
|
||||
err = torch.cuda.cudart().cudaHostRegister(self.bytes.data_ptr(), nbytes, 0)
|
||||
if int(err) != 0:
|
||||
raise RuntimeError(f"cudaHostRegister({nbytes} bytes) failed: {err}")
|
||||
|
||||
def _open_shared_fd(self, nbytes: int, name: str) -> int:
|
||||
owner = None
|
||||
if self.group.rank_in_group == 0:
|
||||
fd = os.memfd_create(name, 0)
|
||||
os.ftruncate(fd, nbytes)
|
||||
owner = (os.getpid(), fd)
|
||||
pid, owner_fd = self.group.broadcast_object(owner, src=0)
|
||||
if self.group.rank_in_group == 0:
|
||||
return fd
|
||||
try:
|
||||
return os.open(f"/proc/{pid}/fd/{owner_fd}", os.O_RDWR)
|
||||
except OSError as e:
|
||||
raise RuntimeError(
|
||||
"engram host table: cannot open rank 0's memfd through /proc; the "
|
||||
"TP ranks must share a PID namespace"
|
||||
) from e
|
||||
|
||||
def _collapse(self, tries: int = 3) -> None:
|
||||
"""Synchronously fold whatever is still on base pages into huge pages.
|
||||
Anonymous memory only; shmem obeys shmem_enabled and refuses."""
|
||||
MADV_COLLAPSE = 25 # Linux >= 6.1; not in Python's mmap module
|
||||
libc = ctypes.CDLL(None, use_errno=True)
|
||||
libc.madvise.argtypes = (ctypes.c_void_p, ctypes.c_size_t, ctypes.c_int)
|
||||
for attempt in range(tries):
|
||||
rc = libc.madvise(
|
||||
ctypes.c_void_p(self.bytes.data_ptr()),
|
||||
ctypes.c_size_t(self.nbytes),
|
||||
MADV_COLLAPSE,
|
||||
)
|
||||
if rc == 0:
|
||||
return
|
||||
err = ctypes.get_errno()
|
||||
if (
|
||||
err != errno.EAGAIN or attempt == tries - 1
|
||||
): # EAGAIN is the only one worth retrying
|
||||
logger.info("engram host table: MADV_COLLAPSE errno %d", err)
|
||||
return
|
||||
time.sleep(1.0)
|
||||
|
||||
def finish_load(self, label: str):
|
||||
if not self.dirty:
|
||||
return
|
||||
self.dirty = False
|
||||
if self.layout == "shared":
|
||||
self.group.barrier()
|
||||
mapped_kb, huge_kb = _huge_pages_backing(self.bytes.data_ptr())
|
||||
if self.layout == "per_rank" and huge_kb < mapped_kb * 0.98:
|
||||
# The loader's own reads refilled the page cache; empty it again so the
|
||||
# collapse can find contiguous memory.
|
||||
drop_checkpoint_page_cache()
|
||||
self._collapse()
|
||||
mapped_kb, huge_kb = _huge_pages_backing(self.bytes.data_ptr())
|
||||
pct = 100.0 * huge_kb / mapped_kb if mapped_kb else 0.0
|
||||
msg = (
|
||||
f"engram host table {label}: layout={self.layout}, "
|
||||
f"{mapped_kb / 2**10:.0f} MiB resident, {huge_kb / 2**10:.0f} MiB in huge pages "
|
||||
f"({pct:.0f}%)"
|
||||
)
|
||||
if huge_kb == 0:
|
||||
knob = "shmem_enabled" if self.layout == "shared" else "enabled"
|
||||
logger.warning(
|
||||
"%s. No huge pages: expect ~10x slower lookups (one TLB miss per row); "
|
||||
"transparent_hugepage/%s is '%s'",
|
||||
msg,
|
||||
knob,
|
||||
_thp_mode(knob) or "unreadable",
|
||||
)
|
||||
else:
|
||||
logger.info(msg)
|
||||
|
||||
|
||||
class EngramEmbedding(nn.Module):
|
||||
"""One layer's fp8 hash table with e8m0 block scales, dequantized on lookup.
|
||||
|
||||
Rows are sharded over the TP group in device memory; with
|
||||
SGLANG_ENABLE_DSV41_ENGRAM_HOST_TABLE they live in host memory instead, as
|
||||
one shared copy or one shard per rank (see _HostTable). Loading is sharded
|
||||
in every layout: a rank writes only its own row range.
|
||||
"""
|
||||
|
||||
def __init__(self, num_embeddings: int, dim: int, layer_id: int):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.tp_size = get_parallel().tp_size
|
||||
tp_rank = get_parallel().tp_rank
|
||||
self.row_start = num_embeddings * tp_rank // self.tp_size
|
||||
row_end = num_embeddings * (tp_rank + 1) // self.tp_size
|
||||
self.rows = row_end - self.row_start
|
||||
self.host_table: Optional[_HostTable] = None
|
||||
if envs.SGLANG_ENABLE_DSV41_ENGRAM_HOST_TABLE.get():
|
||||
self._init_host_table(num_embeddings, dim, layer_id)
|
||||
else:
|
||||
self.weight = nn.Parameter(
|
||||
torch.empty(self.rows, dim, dtype=torch.float8_e4m3fn),
|
||||
requires_grad=False,
|
||||
)
|
||||
self.scale = nn.Parameter(
|
||||
torch.empty(
|
||||
self.rows, dim // FP8_BLOCK_SIZE, dtype=torch.float8_e8m0fnu
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
self.weight.weight_loader = self._load_rows
|
||||
self.scale.weight_loader = self._load_rows
|
||||
|
||||
def _init_host_table(self, num_embeddings: int, dim: int, layer_id: int):
|
||||
layout = envs.SGLANG_DSV41_ENGRAM_HOST_TABLE_LAYOUT.get()
|
||||
n = num_embeddings if layout == "shared" else self.rows
|
||||
w_bytes = n * dim
|
||||
s_bytes = n * (dim // FP8_BLOCK_SIZE)
|
||||
self.host_table = _HostTable(
|
||||
layout,
|
||||
max(1, w_bytes + s_bytes), # mmap requires storage even for an empty shard.
|
||||
f"sglang_engram_{layer_id}",
|
||||
get_tp_group(),
|
||||
)
|
||||
raw = self.host_table.bytes[: w_bytes + s_bytes]
|
||||
weight = raw[:w_bytes].view(torch.float8_e4m3fn).view(n, dim)
|
||||
scale = raw[w_bytes:].view(torch.float8_e8m0fnu).view(n, dim // FP8_BLOCK_SIZE)
|
||||
self.weight = nn.Parameter(weight, requires_grad=False)
|
||||
self.scale = nn.Parameter(scale, requires_grad=False)
|
||||
|
||||
@property
|
||||
def _shared(self) -> bool:
|
||||
return self.host_table is not None and self.host_table.layout == "shared"
|
||||
|
||||
def _load_rows(self, param: nn.Parameter, loaded_weight: torch.Tensor):
|
||||
rows = slice(self.row_start, self.row_start + self.rows)
|
||||
if self._shared:
|
||||
param.data[rows].copy_(loaded_weight[rows])
|
||||
else:
|
||||
param.data.copy_(loaded_weight[rows])
|
||||
if self.host_table is not None:
|
||||
self.host_table.dirty = True
|
||||
|
||||
def finish_load(self, label: str):
|
||||
"""Collective in the shared layout: every rank calls it after loading."""
|
||||
if self.host_table is not None:
|
||||
self.host_table.finish_load(label)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
indices: torch.Tensor,
|
||||
forward_batch: Optional[ForwardBatch] = None,
|
||||
*,
|
||||
cp_all_tokens: bool = False,
|
||||
) -> torch.Tensor:
|
||||
if self._shared:
|
||||
if indices.shape[0] == 0:
|
||||
return self._empty(indices)
|
||||
out = self._empty(indices)
|
||||
engram_gather(
|
||||
self.weight.data_ptr(),
|
||||
self.scale.data_ptr(),
|
||||
indices.reshape(-1),
|
||||
out.view(-1, self.dim),
|
||||
self.dim,
|
||||
FP8_BLOCK_SIZE,
|
||||
)
|
||||
return out
|
||||
if cp_all_tokens and self.tp_size > 1:
|
||||
# Prefill CP: gather the hash ids over the CP group first so every
|
||||
# TP rank looks up the same indices, then keep this rank's slice.
|
||||
parallel = get_parallel()
|
||||
local_rows = indices.shape[0]
|
||||
all_indices = indices.new_empty(
|
||||
(parallel.attn_cp_size * local_rows, *indices.shape[1:])
|
||||
)
|
||||
attn_cp_all_gather_into_tensor(all_indices, indices.contiguous())
|
||||
start = parallel.attn_cp_rank * local_rows
|
||||
return self._lookup(all_indices)[start : start + local_rows]
|
||||
if self.tp_size > 1 and get_attention_dp_size() > 1:
|
||||
return self._dp_sharded_lookup(indices, forward_batch)
|
||||
return self._lookup(indices)
|
||||
|
||||
def _lookup(self, indices: torch.Tensor) -> torch.Tensor:
|
||||
"""Lookup when every TP rank holds the same indices: the device and
|
||||
per-rank host shards zero unowned rows and the all-reduce reassembles."""
|
||||
if indices.shape[0] == 0:
|
||||
return self._empty(indices)
|
||||
values = self._owned_rows(indices)
|
||||
if self.tp_size > 1:
|
||||
values = tensor_model_parallel_all_reduce(values)
|
||||
return values
|
||||
|
||||
def _empty(self, indices: torch.Tensor) -> torch.Tensor:
|
||||
return torch.empty(
|
||||
*indices.shape, self.dim, dtype=torch.bfloat16, device=indices.device
|
||||
)
|
||||
|
||||
def _owned_rows(self, indices: torch.Tensor) -> torch.Tensor:
|
||||
"""Rows of `indices` this rank's shard holds, zero for the rest."""
|
||||
if self.rows == 0:
|
||||
return self._empty(indices).zero_()
|
||||
if self.host_table is None and not _cuda_kernels(indices):
|
||||
local = indices - self.row_start
|
||||
owned = (local >= 0) & (local < self.rows)
|
||||
local = local.masked_fill(~owned, 0)
|
||||
rows = self.weight[local].float().unflatten(-1, (-1, FP8_BLOCK_SIZE))
|
||||
values = (rows * self.scale[local].float().unsqueeze(-1)).flatten(-2)
|
||||
return values.to(torch.bfloat16).masked_fill(~owned.unsqueeze(-1), 0)
|
||||
out = self._empty(indices)
|
||||
engram_gather(
|
||||
self.weight.data_ptr(),
|
||||
self.scale.data_ptr(),
|
||||
indices.reshape(-1),
|
||||
out.view(-1, self.dim),
|
||||
self.dim,
|
||||
FP8_BLOCK_SIZE,
|
||||
row_lo=self.row_start,
|
||||
row_hi=self.row_start + self.rows,
|
||||
)
|
||||
return out
|
||||
|
||||
def _dp_sharded_lookup(
|
||||
self, indices: torch.Tensor, forward_batch: Optional[ForwardBatch]
|
||||
) -> torch.Tensor:
|
||||
"""Gather DP ranks' indices before looking up TP-sharded rows."""
|
||||
assert forward_batch is not None, "the DP engram lookup needs the batch"
|
||||
rows = get_global_dp_buffer_len()
|
||||
ids_global = torch.empty(
|
||||
(rows, *indices.shape[1:]), dtype=indices.dtype, device=indices.device
|
||||
)
|
||||
# The MAX_LEN gather may zero its local input in place, hence the clone.
|
||||
dp_gather_replicate(ids_global, indices.clone(), forward_batch)
|
||||
if rows == 0:
|
||||
return self._empty(indices)
|
||||
values = self._owned_rows(ids_global).view(rows, -1)
|
||||
local = torch.empty(
|
||||
(indices.shape[0], values.shape[1]),
|
||||
dtype=values.dtype,
|
||||
device=values.device,
|
||||
)
|
||||
padding = forward_batch.dp_padding_mode
|
||||
if (
|
||||
padding is not None
|
||||
and padding.is_max_len()
|
||||
and self.tp_size == get_attention_dp_size()
|
||||
and rows == self.tp_size * local.shape[0]
|
||||
) or is_dp_gatherv_active():
|
||||
dp_reduce_scatter_tensor(local, values)
|
||||
else:
|
||||
dp_scatter(local, tensor_model_parallel_all_reduce(values), forward_batch)
|
||||
return local.view(*indices.shape, self.dim)
|
||||
|
||||
|
||||
def engram_gate(
|
||||
x: torch.Tensor,
|
||||
kv: torch.Tensor,
|
||||
q_weight: torch.Tensor,
|
||||
k_weight: torch.Tensor,
|
||||
eps: float,
|
||||
clamp_value: float,
|
||||
) -> torch.Tensor:
|
||||
"""x [T, hc_mult, dim]; kv [T, (hc_mult + 1) * dim] holds one key per hc copy
|
||||
followed by the shared value. Adds the gated value to every copy."""
|
||||
if (
|
||||
_cuda_kernels(x)
|
||||
and x.ndim == 3
|
||||
and kv.shape == (x.shape[0], (x.shape[1] + 1) * x.shape[2])
|
||||
and x.dtype == kv.dtype
|
||||
and x.dtype in (torch.bfloat16, torch.float32)
|
||||
and q_weight.dtype in (torch.bfloat16, torch.float32)
|
||||
and k_weight.dtype in (torch.bfloat16, torch.float32)
|
||||
and all(t.is_contiguous() for t in (x, kv, q_weight, k_weight))
|
||||
):
|
||||
return fused_engram_gate(x, kv, q_weight, k_weight, eps, clamp_value)
|
||||
hc_mult, dim = x.shape[-2:]
|
||||
key, value = kv.split([hc_mult * dim, dim], dim=-1)
|
||||
key = key.float().unflatten(-1, (hc_mult, dim))
|
||||
weight = q_weight.float() * k_weight.float()
|
||||
h = x.float()
|
||||
# Normalized per (token, hc copy) over dim, not jointly over the copies.
|
||||
rstd = torch.rsqrt(h.square().mean(-1) + eps) * torch.rsqrt(
|
||||
key.square().mean(-1) + eps
|
||||
)
|
||||
dot = (h * weight * key).sum(-1) * rstd * dim**-0.5
|
||||
# Signed square root before the sigmoid, matching the training kernel.
|
||||
gate = torch.sigmoid(torch.copysign(dot.abs().clamp_min(clamp_value).sqrt(), dot))
|
||||
return (h + gate.unsqueeze(-1) * value.float().unsqueeze(-2)).to(x.dtype)
|
||||
|
||||
|
||||
class Engram(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
layer_id: int,
|
||||
layout: EngramLayout,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.layer_hash_index = layout.layer_ids.index(layer_id)
|
||||
self.eps = config.rms_norm_eps
|
||||
self.clamp_value = 1e-6
|
||||
dim, hc_mult = config.hidden_size, config.hc_mult
|
||||
self.embed = EngramEmbedding(
|
||||
layout.num_embeddings[self.layer_hash_index], layout.head_dim, layer_id
|
||||
)
|
||||
n_hash_cols = (layout.max_ngram_size - 1) * layout.n_heads
|
||||
self.wkv = ReplicatedLinear(
|
||||
n_hash_cols * layout.head_dim,
|
||||
dim * (hc_mult + 1),
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("wkv", prefix),
|
||||
)
|
||||
self.q_weight = nn.Parameter(torch.ones(hc_mult, dim), requires_grad=False)
|
||||
self.k_weight = nn.Parameter(torch.ones(hc_mult, dim), requires_grad=False)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
hash_ids: torch.Tensor,
|
||||
forward_batch: Optional[ForwardBatch] = None,
|
||||
*,
|
||||
cp_all_tokens: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""x [T, hc_mult, dim]; hash_ids [T, n_hash_cols] for this layer."""
|
||||
# The lookup runs first even for an idle DP-attention batch: under DP
|
||||
# attention it is a collective every rank has to join.
|
||||
emb = self.embed(hash_ids, forward_batch, cp_all_tokens=cp_all_tokens)
|
||||
if x.shape[0] == 0:
|
||||
# Nothing to gate, and the MXFP8 quantize behind wkv rejects an
|
||||
# empty M.
|
||||
return x
|
||||
kv, _ = self.wkv(emb.flatten(-2))
|
||||
return engram_gate(
|
||||
x, kv, self.q_weight, self.k_weight, self.eps, self.clamp_value
|
||||
)
|
||||
@@ -2382,6 +2382,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
# Mask marking chunked (not-yet-finished) prefill requests whose sampled
|
||||
# pseudo next-token must NOT be written into the ngram token table.
|
||||
ne_skip_token_table_update: torch.Tensor = None
|
||||
# DeepSeek-V4.1 engram, extend batches only: [bs, n - 1] int32 predecessors
|
||||
# of each request's first extend token (NgramEmbeddingManager).
|
||||
engram_history: Optional[torch.Tensor] = None
|
||||
|
||||
req_pool_indices: torch.Tensor = None # shape: [b], int64
|
||||
seq_lens: torch.Tensor = None # shape: [b], int64
|
||||
|
||||
@@ -699,6 +699,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
# For ngram embedding
|
||||
ngram_embedding_info: Optional[NgramEmbeddingInfo] = None
|
||||
|
||||
# DeepSeek-V4.1 engram, extend only: the n - 1 tokens before each request's
|
||||
# first extend token, oldest first, [bs, n - 1] int32 (see EngramHasher).
|
||||
engram_history: Optional[torch.Tensor] = None
|
||||
|
||||
# For dumper: int-hashed request / bootstrap-room IDs (derived from rids)
|
||||
rids_int: Optional[torch.Tensor] = None
|
||||
bootstrap_room_ids_int: Optional[torch.Tensor] = None
|
||||
@@ -900,6 +904,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
orig_seq_lens=batch.orig_seq_lens,
|
||||
out_cache_loc_dsv4=batch.out_cache_loc_dsv4,
|
||||
engram_history=batch.engram_history,
|
||||
mamba_track_indices=batch.mamba_track_indices,
|
||||
mamba_track_mask=batch.mamba_track_mask,
|
||||
mamba_track_seqlens=batch.mamba_track_seqlens,
|
||||
|
||||
+70
-10
@@ -14,6 +14,7 @@ from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||
from sglang.srt.runtime_context import get_schedule
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.engram import EngramHasher
|
||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
|
||||
@@ -23,7 +24,8 @@ class NgramEmbeddingManager:
|
||||
enabled: bool
|
||||
table: Optional[torch.Tensor]
|
||||
n: int
|
||||
k: int
|
||||
# Draft runners have no hasher even when their config has engram layers.
|
||||
engram_hasher: Optional[EngramHasher] = None
|
||||
|
||||
@classmethod
|
||||
def from_model(
|
||||
@@ -36,8 +38,6 @@ class NgramEmbeddingManager:
|
||||
device: str,
|
||||
):
|
||||
token_table = None
|
||||
ngram_embedding_n = 0
|
||||
ngram_embedding_k = 0
|
||||
use_ngram_embedding = model_config.use_ngram_embedding
|
||||
if use_ngram_embedding:
|
||||
from sglang.srt.layers.n_gram_embedding import NgramEmbedding
|
||||
@@ -58,14 +58,20 @@ class NgramEmbeddingManager:
|
||||
module.init_buffers(
|
||||
max_running_requests, chunked_prefill_size, device
|
||||
)
|
||||
hf_config = model_config.hf_config
|
||||
ngram_embedding_n = hf_config.ngram_embedding_n
|
||||
ngram_embedding_k = hf_config.ngram_embedding_k
|
||||
engram_hasher = None
|
||||
if model_config.use_engram:
|
||||
from sglang.srt.layers.engram import EngramHasher
|
||||
|
||||
for module in model.modules():
|
||||
if isinstance(module, EngramHasher):
|
||||
assert engram_hasher is None, "one engram hasher per model"
|
||||
module.init_history(req_to_token_pool.req_to_token.shape[0], device)
|
||||
engram_hasher = module
|
||||
return cls(
|
||||
enabled=use_ngram_embedding,
|
||||
table=token_table,
|
||||
n=ngram_embedding_n,
|
||||
k=ngram_embedding_k,
|
||||
n=model_config.ngram_embedding_n,
|
||||
engram_hasher=engram_hasher,
|
||||
)
|
||||
|
||||
def update_after_decode(
|
||||
@@ -85,14 +91,31 @@ class NgramEmbeddingManager:
|
||||
batch_size=forward_batch.batch_size,
|
||||
)
|
||||
|
||||
def update_after_verify(
|
||||
self,
|
||||
*,
|
||||
verify_ids_2d: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
commit_lens: torch.Tensor,
|
||||
) -> None:
|
||||
if self.engram_hasher is None:
|
||||
return
|
||||
self.engram_hasher.commit_after_verify(
|
||||
verify_ids_2d, req_pool_indices, commit_lens
|
||||
)
|
||||
|
||||
def prepare_for_forward(
|
||||
self,
|
||||
batch: Optional[ScheduleBatch],
|
||||
*,
|
||||
chunked_req: Optional[Req],
|
||||
) -> Optional[ScheduleBatch]:
|
||||
"""Fill the token table for ngram embedding before a forward pass."""
|
||||
if batch is None or not self.enabled:
|
||||
"""Fill the ngram token table and engram history before a forward pass."""
|
||||
if batch is None:
|
||||
return batch
|
||||
if self.engram_hasher is not None:
|
||||
self._prepare_engram_history(batch)
|
||||
if not self.enabled:
|
||||
return batch
|
||||
batch.ne_token_table = self.table
|
||||
if batch.forward_mode == ForwardMode.EXTEND:
|
||||
@@ -144,6 +167,43 @@ class NgramEmbeddingManager:
|
||||
)
|
||||
return batch
|
||||
|
||||
def _prepare_engram_history(self, batch: ScheduleBatch) -> None:
|
||||
"""Refresh extend predecessors after prefix hits, retraction, or slot reuse."""
|
||||
n1 = self.engram_hasher.max_ngram_size - 1
|
||||
if batch.forward_mode.is_prebuilt():
|
||||
# PD decode runs no EXTEND for this request, so the row its first
|
||||
# DECODE reads is written here: the n - 1 tokens before the one
|
||||
# prefill sampled, which is the token decode feeds next.
|
||||
history = self.engram_hasher.history
|
||||
rows = []
|
||||
for req in batch.reqs:
|
||||
# full_untruncated_fill_ids is only refreshed on the extend and
|
||||
# decode paths, which a PD decode request has not run yet.
|
||||
fill_ids = req.origin_input_ids + req.output_ids
|
||||
end = len(fill_ids) - 1
|
||||
ids = fill_ids[max(0, end - n1) : end]
|
||||
rows.append([0] * (n1 - len(ids)) + list(ids))
|
||||
slots = torch.tensor(
|
||||
[req.kv.req_pool_idx for req in batch.reqs],
|
||||
dtype=torch.int64,
|
||||
device=history.device,
|
||||
)
|
||||
history[slots] = torch.tensor(
|
||||
rows, dtype=history.dtype, device=history.device
|
||||
).view(len(rows), n1)
|
||||
return
|
||||
if not batch.forward_mode.is_extend_without_speculative():
|
||||
return
|
||||
rows = []
|
||||
for req in batch.reqs:
|
||||
start = req.extend_range.start
|
||||
lo = max(0, start - n1)
|
||||
ids = req.full_untruncated_fill_ids[lo:start]
|
||||
rows.append([0] * (n1 - len(ids)) + list(ids))
|
||||
batch.engram_history = torch.tensor(
|
||||
rows, dtype=torch.int32, device=self.engram_hasher.history.device
|
||||
).view(len(rows), n1)
|
||||
|
||||
|
||||
def update_ngram_token_table_after_sampling(
|
||||
*,
|
||||
|
||||
@@ -73,6 +73,8 @@ def _stub_for_initialize(
|
||||
num_attention_layers=1,
|
||||
context_len=64,
|
||||
use_ngram_embedding=False,
|
||||
ngram_embedding_n=0,
|
||||
use_engram=False,
|
||||
)
|
||||
return stub
|
||||
|
||||
|
||||
@@ -111,6 +111,8 @@ def _hybrid_stub_for_initialize(
|
||||
num_attention_layers=1,
|
||||
context_len=64,
|
||||
use_ngram_embedding=False, # short-circuits NgramEmbeddingManager
|
||||
ngram_embedding_n=0,
|
||||
use_engram=False,
|
||||
)
|
||||
return stub
|
||||
|
||||
|
||||
Reference in New Issue
Block a user