Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 21 additions & 2 deletions docs/source/user_guide/examples/experimental-server.md
Original file line number Diff line number Diff line change
Expand Up @@ -348,13 +348,32 @@ API.

## Runtime Concurrency

The current high-level runtime has one mutable generation state. The server
therefore admits one request at a time and uses a bounded async queue configured
The current high-level runtime has one mutable generation state. By default the
server admits one request at a time and uses a bounded async queue configured
by `--max-queued-requests` and `--queue-timeout`. Queue overflow and timeout
return HTTP 429 (Anthropic 529). Streaming disconnects cancel the native channel
immediately, wait for the native worker to exit, and then release the runtime
lease. Engines stay resident across HTTP connections; graceful server shutdown
drains active work and releases the runtime and its device resources.

`--enable-batching` merges concurrent non-streaming requests into one runtime
call. Requests join a batch only when their generation-level settings match
(sampling parameters, `max_tokens`, chat-template and thinking flags,
speculative-decoding and context-cache options); each request keeps its own
messages, media, stop strings and logit bias. A request waits at most
`--batch-timeout-ms` (default 10) for others to arrive, and one call carries at
most `--max-queue-batch-size` requests, capped by the engine's
`--max-batch-size`. Decode is memory-bandwidth bound, so a batch of four costs
little more than a single request and aggregate throughput scales almost
linearly. Streaming requests still run one at a time; a streaming request in
flight delays batched requests until it completes. `/health` reports the
active batching settings.

`--max-verify-tree-size` and `--max-draft-tree-size` bound the speculative
tree the compiled bundle supports. The builder defaults both to 60; for linear
MTP drafting a bound of one more than the largest `num_speculative_tokens` you
intend to run is sufficient, and the runtime's spec-verify state buffers
scale with it (about 0.9 GB per position for a 27B hybrid model at batch 4).

Continuous batching, chunked prefill scheduling, and tensor parallelism require
additional native scheduler support and are rejected at launch.
28 changes: 25 additions & 3 deletions experimental/builder/core/artifacts/chat_template.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,8 +92,23 @@ def try_write_packaged_chat_template(model_dir: str,


def _format_chat(tokenizer, messages, **kwargs) -> str:
"""Apply a tokenizer chat template and return text."""
return tokenizer.apply_chat_template(messages, tokenize=False, **kwargs)
"""Apply a tokenizer chat template and return text.

Probes render with thinking disabled unless told otherwise: templates
that rewrite the conversation in thinking mode (Qwen3.8 injects reasoning
instructions into the system block and treats an undefined flag as
enabled) would otherwise yield prefixes that never match a request.
"""
kwargs.setdefault("enable_thinking", False)
try:
return tokenizer.apply_chat_template(messages,
tokenize=False,
**kwargs)
except TypeError:
kwargs.pop("enable_thinking", None)
return tokenizer.apply_chat_template(messages,
tokenize=False,
**kwargs)


def _format_chat_generation(tokenizer, messages, enable_thinking: bool) -> str:
Expand Down Expand Up @@ -213,10 +228,17 @@ def process_chat_template(model_dir: str,

generation_prompt_thinking = None
try:
# Slice against a thinking-mode baseline: a template may render a
# longer conversation in thinking mode, so the non-thinking
# baseline offset would cut into the wrong place.
thinking_base = _format_chat(tokenizer,
[system_message, user_message],
add_generation_prompt=False,
enable_thinking=True)
thinking_formatted = _format_chat_generation(
tokenizer, [system_message, user_message],
enable_thinking=True)
candidate = thinking_formatted[len(user_formatted):]
candidate = thinking_formatted[len(thinking_base):]
if candidate != generation_prompt:
generation_prompt_thinking = candidate
except (TypeError, ValueError, KeyError):
Expand Down
58 changes: 56 additions & 2 deletions experimental/builder/core/weights.py
Original file line number Diff line number Diff line change
Expand Up @@ -188,13 +188,46 @@ def _resolve(self, name: str, required: bool = True) -> Optional[str]:
if not candidate.endswith("embed_tokens.weight")
]
candidates = list(dict.fromkeys(candidates))
owned = self._sidecar_in_owner_namespace(name)
if owned is not None:
candidates = [owned]
for candidate in candidates:
if self.store.has(candidate):
return candidate
if required:
raise KeyError(f"checkpoint tensor not found: {name!r}")
return None

# Per-module tensors that must live beside the module's weight. A draft
# resolves its modules in a separate checkpoint namespace (Qwen3.5 MTP
# ``mtp.``); without this pin, a sidecar missing there would be answered
# by the base model's identically named module.
_MODULE_SIDECARS = (".weight_scale", ".weight_scale_2",
".weight_global_scale", ".input_scale",
".input_global_scale", ".bias", ".pre_quant_scale",
".qzeros", ".scales", ".g_idx", ".q_scale", ".k_scale",
".v_scale")
_PRIMARY_WEIGHTS = (".weight", ".qweight", ".weight_packed")

def _sidecar_in_owner_namespace(self, name: str) -> Optional[str]:
"""Concrete key for a module sidecar, pinned to its weight's namespace.

Returns None when *name* is not a module sidecar or the module has no
resolvable primary weight, leaving ordinary candidate resolution to
run.
"""
for suffix in self._MODULE_SIDECARS:
if name.endswith(suffix) and len(name) > len(suffix):
prefix = name[:-len(suffix)]
break
else:
return None
for weight_suffix in self._PRIMARY_WEIGHTS:
key = self._resolve(prefix + weight_suffix, required=False)
if key is not None:
return key[:-len(weight_suffix)] + suffix
return None

def checkpoint_key(self, name: str) -> str:
"""Return the concrete checkpoint key backing a model tensor name."""
return self._resolve(name, required=False) or name
Expand Down Expand Up @@ -241,15 +274,36 @@ def checkpoint_locations(self, names: Sequence[str]) -> dict:
locations[name] = location
return locations

_QUANT_SIDECAR_SUFFIXES = (".weight_scale", ".weight_scale_2",
".weight_global_scale", ".qweight",
".weight_packed", ".scales")

def module_quant_type(self,
name: str,
*,
tie_word_embeddings: bool = False) -> str:
"""Return the checkpoint precision owned by one model projection."""
lookup = name
normalize = getattr(self.conversion, "normalize_checkpoint_name", None)
if normalize is not None:
name = normalize(name)
return self.quant.module_type(name, tie_word_embeddings)
lookup = normalize(lookup)
quant_type = self.quant.module_type(lookup, tie_word_embeddings)
# A draft resolves its projections in a separate checkpoint namespace
# (Qwen3.5 MTP stores them under ``mtp.``), where the same short name
# may be unquantized even though the base layer carries a quantized
# override. The tensor actually read decides: a bare weight with no
# quantization sidecar is FP16. ``lm_head`` is skipped because its
# resolution consults this method for embedding tying.
if (quant_type != quantization.QUANT_FP16 and name != "lm_head"
and self._has_plain_weight(name)):
return quantization.QUANT_FP16
return quant_type

def _has_plain_weight(self, name: str) -> bool:
if not self.has(name + ".weight"):
return False
return not any(
self.has(name + suffix) for suffix in self._QUANT_SIDECAR_SUFFIXES)

def parameter_spec(self,
name: str,
Expand Down
4 changes: 3 additions & 1 deletion experimental/builder/models/dflash/modeling_dflash_draft.py
Original file line number Diff line number Diff line change
Expand Up @@ -163,13 +163,15 @@ def __init__(self, ctx, lm_head=None) -> None:

def input_tensors(self) -> Dict[str, object]:
cfg = self.cfg
kv_dtype = (trt.DataType.FP8
if cfg.kv_cache_quant == "fp8" else trt.float16)
target_layers = cfg.dflash_target_layer_ids or [1, 8, 15, 22, 29]
return {
"inputs_embeds":
self.add_input("inputs_embeds", trt.float16,
(-1, -1, cfg.hidden_size)),
"past_key_values": [
self.add_input(f"past_key_values_{index}", trt.float16,
self.add_input(f"past_key_values_{index}", kv_dtype,
(2, -1, F.KV_PAGE_SIZE, cfg.num_key_value_heads,
cfg.head_dim))
for index in range(cfg.num_hidden_layers)
Expand Down
4 changes: 3 additions & 1 deletion experimental/builder/models/dspark/modeling_dspark_draft.py
Original file line number Diff line number Diff line change
Expand Up @@ -128,13 +128,15 @@ def __init__(self, ctx) -> None:

def input_tensors(self) -> Dict[str, object]:
cfg = self.cfg
kv_dtype = (trt.DataType.FP8
if cfg.kv_cache_quant == "fp8" else trt.float16)
target_layers = cfg.dspark_target_layer_ids
return {
"inputs_embeds":
self.add_input("inputs_embeds", trt.float16,
(-1, -1, cfg.hidden_size)),
"past_key_values": [
self.add_input(f"past_key_values_{index}", trt.float16,
self.add_input(f"past_key_values_{index}", kv_dtype,
(2, -1, F.KV_PAGE_SIZE, cfg.num_key_value_heads,
cfg.head_dim))
for index in range(cfg.num_hidden_layers)
Expand Down
4 changes: 3 additions & 1 deletion experimental/builder/models/eagle3/modeling_eagle3_draft.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,8 @@ def __init__(self, ctx, lm_head=None) -> None:

def input_tensors(self) -> Dict[str, object]:
cfg = self.cfg
kv_dtype = (trt.DataType.FP8
if cfg.kv_cache_quant == "fp8" else trt.float16)
target_hidden = int(cfg.target_hidden_size or cfg.raw_component.get(
"eagle3_target_hidden_size", cfg.hidden_size))
target_layers = len(cfg.eagle3_target_layer_ids)
Expand All @@ -112,7 +114,7 @@ def input_tensors(self) -> Dict[str, object]:
self.add_input("inputs_embeds", trt.float16,
(-1, -1, cfg.hidden_size)),
"past_key_values": [
self.add_input(f"past_key_values_{index}", trt.float16,
self.add_input(f"past_key_values_{index}", kv_dtype,
(2, -1, F.KV_PAGE_SIZE, cfg.num_key_value_heads,
cfg.head_dim))
for index in range(cfg.num_hidden_layers)
Expand Down
4 changes: 3 additions & 1 deletion experimental/builder/models/qwen3_5/modeling_qwen3_5_mtp.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,12 +77,14 @@ def __init__(self, ctx) -> None:

def input_tensors(self) -> Dict[str, object]:
cfg = self.cfg
kv_dtype = (trt.DataType.FP8
if cfg.kv_cache_quant == "fp8" else trt.float16)
return {
"inputs_embeds":
self.add_input("inputs_embeds", trt.float16,
(-1, -1, cfg.hidden_size)),
"past_key_values": [
self.add_input(f"past_key_values_{index}", trt.float16,
self.add_input(f"past_key_values_{index}", kv_dtype,
(2, -1, F.KV_PAGE_SIZE, cfg.num_key_value_heads,
cfg.head_dim))
for index in range(cfg.num_hidden_layers)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -127,12 +127,14 @@ def __init__(self, ctx) -> None:

def input_tensors(self) -> Dict[str, object]:
cfg = self.cfg
kv_dtype = (trt.DataType.FP8
if cfg.kv_cache_quant == "fp8" else trt.float16)
return {
"inputs_embeds":
self.add_input("inputs_embeds", trt.float16,
(-1, -1, cfg.hidden_size)),
"past_key_values": [
self.add_input(f"past_key_values_{index}", trt.float16,
self.add_input(f"past_key_values_{index}", kv_dtype,
(2, -1, F.KV_PAGE_SIZE, cfg.num_key_value_heads,
cfg.head_dim))
for index in range(cfg.num_hidden_layers)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -85,12 +85,14 @@ def __init__(self, ctx) -> None:

def input_tensors(self) -> Dict[str, object]:
cfg = self.cfg
kv_dtype = (trt.DataType.FP8
if cfg.kv_cache_quant == "fp8" else trt.float16)
return {
"inputs_embeds":
self.add_input("inputs_embeds", trt.float16,
(-1, -1, cfg.hidden_size)),
"past_key_values": [
self.add_input(f"past_key_values_{index}", trt.float16,
self.add_input(f"past_key_values_{index}", kv_dtype,
(2, -1, F.KV_PAGE_SIZE, cfg.num_key_value_heads,
cfg.head_dim))
for index in range(cfg.num_hidden_layers)
Expand Down
1 change: 1 addition & 0 deletions experimental/server/api/routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ async def health(request: Request):
"model": client.model_name,
"active_requests": client.active_requests,
"queued_requests": client.queued_requests,
"batching": client.batching,
"capabilities": {
"chat": caps.chat,
"transcription": caps.transcription,
Expand Down
58 changes: 58 additions & 0 deletions experimental/server/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -174,6 +174,11 @@ class ModelConfig:
draft_top_k: Optional[int] = None
draft_step: Optional[int] = None
verify_tree_size: Optional[int] = None
max_verify_tree_size: Optional[int] = None
max_draft_tree_size: Optional[int] = None
enable_batching: bool = False
batch_timeout_ms: float = 10.0
max_queue_batch_size: Optional[int] = None
speculative_config: Optional[SpeculativeConfig] = None
context_cache_config: ContextCacheConfig = field(
default_factory=ContextCacheConfig)
Expand All @@ -184,6 +189,10 @@ def __post_init__(self) -> None:
or self.engine_cache_max_size_gb <= 0):
raise ServerConfigError(
"engine_cache_max_size_gb must be positive")
if (isinstance(self.batch_timeout_ms, bool)
or not math.isfinite(self.batch_timeout_ms)
or self.batch_timeout_ms < 0):
raise ServerConfigError("batch_timeout_ms must be non-negative")

def llm_kwargs(self) -> Dict[str, Any]:
return {
Expand All @@ -197,6 +206,11 @@ def llm_kwargs(self) -> Dict[str, Any]:
"draft_top_k": self.draft_top_k,
"draft_step": self.draft_step,
"verify_tree_size": self.verify_tree_size,
"max_verify_tree_size": self.max_verify_tree_size,
"max_draft_tree_size": self.max_draft_tree_size,
"enable_batching": self.enable_batching,
"batch_timeout_ms": self.batch_timeout_ms,
"max_queue_batch_size": self.max_queue_batch_size,
"speculative_config": self.speculative_config,
"context_cache_config": self.context_cache_config,
}
Expand Down Expand Up @@ -254,6 +268,13 @@ def _positive_float(value: str) -> float:
return parsed


def _non_negative_float(value: str) -> float:
parsed = float(value)
if not math.isfinite(parsed) or parsed < 0:
raise argparse.ArgumentTypeError("must be non-negative")
return parsed


def _bounded_port(value: str) -> int:
parsed = int(value)
if not 1 <= parsed <= 65535:
Expand Down Expand Up @@ -325,7 +346,39 @@ def create_argument_parser() -> argparse.ArgumentParser:
model.add_argument("--draft-top-k", type=_positive_int)
model.add_argument("--draft-step", type=_positive_int)
model.add_argument("--verify-tree-size", type=_positive_int)
model.add_argument(
"--max-verify-tree-size",
type=_positive_int,
help="Largest base verification input the compiled bundle supports. "
"Spec-verify state buffers scale with it; the builder default is 60.",
)
model.add_argument(
"--max-draft-tree-size",
type=_positive_int,
help="Largest draft proposal the compiled bundle supports; the "
"builder default is 60.",
)
model.add_argument("--speculative-config", default="")
model.add_argument(
"--enable-batching",
action="store_true",
help="Merge concurrent non-streaming requests with matching sampling "
"settings into one runtime call, up to the engine's max batch size. "
"Streaming requests still run one at a time.",
)
model.add_argument(
"--batch-timeout-ms",
type=_non_negative_float,
default=10.0,
help="How long a request waits for compatible requests to join its "
"batch before running.",
)
model.add_argument(
"--max-queue-batch-size",
type=_positive_int,
help="Cap on requests merged per runtime call; defaults to the "
"engine's max batch size.",
)
model.add_argument(
"--enable-context-reuse",
action="store_true",
Expand Down Expand Up @@ -377,6 +430,11 @@ def parse_server_config(argv: Optional[Sequence[str]] = None) -> ServerConfig:
draft_top_k=args.draft_top_k,
draft_step=draft_step,
verify_tree_size=args.verify_tree_size,
max_verify_tree_size=args.max_verify_tree_size,
max_draft_tree_size=args.max_draft_tree_size,
enable_batching=args.enable_batching,
batch_timeout_ms=args.batch_timeout_ms,
max_queue_batch_size=args.max_queue_batch_size,
speculative_config=speculative,
context_cache_config=context_cache,
)
Expand Down
Loading