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.hf_config.architectures
|
||||||
)
|
)
|
||||||
self.use_ngram_embedding = getattr(self.hf_config, "use_ngram_embedding", False)
|
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.
|
# A multimodal arch is piecewise-incompatible until its LM prefill is validated.
|
||||||
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)
|
||||||
|
|||||||
@@ -1474,6 +1474,15 @@ class Envs:
|
|||||||
# picks the buffer dtype and this switch is inert.
|
# picks the buffer dtype and this switch is inert.
|
||||||
SGLANG_DSV4_UNIFIED_KV_FP8 = EnvBool(False)
|
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
|
# Kernels and indexer
|
||||||
SGLANG_OPT_DEEPGEMM_HC_PRENORM = EnvBool(True)
|
SGLANG_OPT_DEEPGEMM_HC_PRENORM = EnvBool(True)
|
||||||
SGLANG_OPT_USE_TILELANG_MHC_PRE = 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
|
# Mask marking chunked (not-yet-finished) prefill requests whose sampled
|
||||||
# pseudo next-token must NOT be written into the ngram token table.
|
# pseudo next-token must NOT be written into the ngram token table.
|
||||||
ne_skip_token_table_update: torch.Tensor = None
|
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
|
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
|
||||||
|
|||||||
@@ -699,6 +699,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
# For ngram embedding
|
# For ngram embedding
|
||||||
ngram_embedding_info: Optional[NgramEmbeddingInfo] = None
|
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)
|
# For dumper: int-hashed request / bootstrap-room IDs (derived from rids)
|
||||||
rids_int: Optional[torch.Tensor] = None
|
rids_int: Optional[torch.Tensor] = None
|
||||||
bootstrap_room_ids_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,
|
seq_lens_cpu=seq_lens_cpu,
|
||||||
orig_seq_lens=batch.orig_seq_lens,
|
orig_seq_lens=batch.orig_seq_lens,
|
||||||
out_cache_loc_dsv4=batch.out_cache_loc_dsv4,
|
out_cache_loc_dsv4=batch.out_cache_loc_dsv4,
|
||||||
|
engram_history=batch.engram_history,
|
||||||
mamba_track_indices=batch.mamba_track_indices,
|
mamba_track_indices=batch.mamba_track_indices,
|
||||||
mamba_track_mask=batch.mamba_track_mask,
|
mamba_track_mask=batch.mamba_track_mask,
|
||||||
mamba_track_seqlens=batch.mamba_track_seqlens,
|
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
|
from sglang.srt.runtime_context import get_schedule
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.layers.engram import EngramHasher
|
||||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
@@ -23,7 +24,8 @@ class NgramEmbeddingManager:
|
|||||||
enabled: bool
|
enabled: bool
|
||||||
table: Optional[torch.Tensor]
|
table: Optional[torch.Tensor]
|
||||||
n: int
|
n: int
|
||||||
k: int
|
# Draft runners have no hasher even when their config has engram layers.
|
||||||
|
engram_hasher: Optional[EngramHasher] = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_model(
|
def from_model(
|
||||||
@@ -36,8 +38,6 @@ class NgramEmbeddingManager:
|
|||||||
device: str,
|
device: str,
|
||||||
):
|
):
|
||||||
token_table = None
|
token_table = None
|
||||||
ngram_embedding_n = 0
|
|
||||||
ngram_embedding_k = 0
|
|
||||||
use_ngram_embedding = model_config.use_ngram_embedding
|
use_ngram_embedding = model_config.use_ngram_embedding
|
||||||
if use_ngram_embedding:
|
if use_ngram_embedding:
|
||||||
from sglang.srt.layers.n_gram_embedding import NgramEmbedding
|
from sglang.srt.layers.n_gram_embedding import NgramEmbedding
|
||||||
@@ -58,14 +58,20 @@ class NgramEmbeddingManager:
|
|||||||
module.init_buffers(
|
module.init_buffers(
|
||||||
max_running_requests, chunked_prefill_size, device
|
max_running_requests, chunked_prefill_size, device
|
||||||
)
|
)
|
||||||
hf_config = model_config.hf_config
|
engram_hasher = None
|
||||||
ngram_embedding_n = hf_config.ngram_embedding_n
|
if model_config.use_engram:
|
||||||
ngram_embedding_k = hf_config.ngram_embedding_k
|
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(
|
return cls(
|
||||||
enabled=use_ngram_embedding,
|
enabled=use_ngram_embedding,
|
||||||
table=token_table,
|
table=token_table,
|
||||||
n=ngram_embedding_n,
|
n=model_config.ngram_embedding_n,
|
||||||
k=ngram_embedding_k,
|
engram_hasher=engram_hasher,
|
||||||
)
|
)
|
||||||
|
|
||||||
def update_after_decode(
|
def update_after_decode(
|
||||||
@@ -85,14 +91,31 @@ class NgramEmbeddingManager:
|
|||||||
batch_size=forward_batch.batch_size,
|
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(
|
def prepare_for_forward(
|
||||||
self,
|
self,
|
||||||
batch: Optional[ScheduleBatch],
|
batch: Optional[ScheduleBatch],
|
||||||
*,
|
*,
|
||||||
chunked_req: Optional[Req],
|
chunked_req: Optional[Req],
|
||||||
) -> Optional[ScheduleBatch]:
|
) -> Optional[ScheduleBatch]:
|
||||||
"""Fill the token table for ngram embedding before a forward pass."""
|
"""Fill the ngram token table and engram history before a forward pass."""
|
||||||
if batch is None or not self.enabled:
|
if batch is None:
|
||||||
|
return batch
|
||||||
|
if self.engram_hasher is not None:
|
||||||
|
self._prepare_engram_history(batch)
|
||||||
|
if not self.enabled:
|
||||||
return batch
|
return batch
|
||||||
batch.ne_token_table = self.table
|
batch.ne_token_table = self.table
|
||||||
if batch.forward_mode == ForwardMode.EXTEND:
|
if batch.forward_mode == ForwardMode.EXTEND:
|
||||||
@@ -144,6 +167,43 @@ class NgramEmbeddingManager:
|
|||||||
)
|
)
|
||||||
return batch
|
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(
|
def update_ngram_token_table_after_sampling(
|
||||||
*,
|
*,
|
||||||
|
|||||||
@@ -73,6 +73,8 @@ def _stub_for_initialize(
|
|||||||
num_attention_layers=1,
|
num_attention_layers=1,
|
||||||
context_len=64,
|
context_len=64,
|
||||||
use_ngram_embedding=False,
|
use_ngram_embedding=False,
|
||||||
|
ngram_embedding_n=0,
|
||||||
|
use_engram=False,
|
||||||
)
|
)
|
||||||
return stub
|
return stub
|
||||||
|
|
||||||
|
|||||||
@@ -111,6 +111,8 @@ def _hybrid_stub_for_initialize(
|
|||||||
num_attention_layers=1,
|
num_attention_layers=1,
|
||||||
context_len=64,
|
context_len=64,
|
||||||
use_ngram_embedding=False, # short-circuits NgramEmbeddingManager
|
use_ngram_embedding=False, # short-circuits NgramEmbeddingManager
|
||||||
|
ngram_embedding_n=0,
|
||||||
|
use_engram=False,
|
||||||
)
|
)
|
||||||
return stub
|
return stub
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user