[Speculative Decoding] Add native UNO serving support (#37667)

Co-authored-by: drproduck <drproduck@MacBook-Air-2.local>
Co-authored-by: BBuf <1182563586@qq.com>
This commit is contained in:
Yash Akhauri
2026-09-03 20:08:41 +08:00
committed by GitHub
co-authored by drproduck BBuf
parent 354ed6d66b
commit 2bb25dc18b
51 changed files with 6300 additions and 74 deletions
+148
View File
@@ -0,0 +1,148 @@
# UNO full-dataset math evaluation
`run_math_eval.py` evaluates AR, UNO, DFLASH, EAGLE, or EAGLE3 with identical
datasets, prompts, sampling parameters, and grading. It creates an in-process
`sgl.Engine`; there is no separate server process. Engine startup is excluded
from the timed interval, and no additional request warmup is run.
The runner downloads pinned revisions of GSM8K, MATH-500, AIME 2024, AIME
2025, and AIME 2026. It applies the same boxed-answer instruction and Qwen
reasoning chat template to every engine, then grades with `math_verify`.
Install SGLang with its evaluation dependencies:
```bash
pip install -e "python[test]"
```
## Reproduce the H200 table
Run from the SGLang repository root. Each invocation below produces one row of
the PR table. GSM8K and MATH-500 use one sample per problem; AIME 2025 uses ten
samples per problem, or 300 completions.
```bash
export MODEL_PATH=Qwen/Qwen3-8B
export TOKENIZER_PATH=Qwen/Qwen3-8B
export UNO_LORA_PATH=s-sahoo/uno-qwen3-8B
export DATA_ROOT=/path/to/math-eval-data
export RESULT_ROOT=/path/to/math-eval-results
COMMON_ARGS=(
--model-path "$MODEL_PATH"
--tokenizer-path "$TOKENIZER_PATH"
--data-root "$DATA_ROOT"
--context-length 40960
--max-tokens 32768
--temperature 1
--top-k 50
--top-p 0.95
--random-seed 42
)
run_ar() {
local benchmark=$1 samples=$2 requests=$3 output_name=$4
PYTHONPATH=python python -m benchmark.uno.run_math_eval \
"${COMMON_ARGS[@]}" \
--benchmark "$benchmark" \
--num-samples "$samples" \
--max-running-requests "$requests" \
--output-dir "$RESULT_ROOT/$output_name"
}
run_linear_uno() {
local benchmark=$1 samples=$2 requests=$3 output_name=$4
PYTHONPATH=python python -m benchmark.uno.run_math_eval \
"${COMMON_ARGS[@]}" \
--benchmark "$benchmark" \
--num-samples "$samples" \
--max-running-requests "$requests" \
--output-dir "$RESULT_ROOT/$output_name" \
--speculative-algorithm UNO \
--uno-lora-path "$UNO_LORA_PATH" \
--speculative-num-steps 1 \
--speculative-eagle-topk 1 \
--speculative-num-draft-tokens 8
}
run_tree_uno() {
local benchmark=$1 samples=$2 requests=$3 output_name=$4
PYTHONPATH=python python -m benchmark.uno.run_math_eval \
"${COMMON_ARGS[@]}" \
--benchmark "$benchmark" \
--num-samples "$samples" \
--max-running-requests "$requests" \
--output-dir "$RESULT_ROOT/$output_name" \
--speculative-algorithm UNO \
--uno-lora-path "$UNO_LORA_PATH" \
--speculative-num-steps 15 \
--speculative-eagle-topk 32 \
--speculative-num-draft-tokens 32
}
```
Run the six batch-64 AR and linear `B/K/V = 8/1/8` rows:
```bash
run_ar gsm8k 1 64 ar-gsm8k-c64
run_linear_uno gsm8k 1 64 uno-linear-b8-k1-v8-gsm8k-c64
run_ar math500 1 64 ar-math500-c64
run_linear_uno math500 1 64 uno-linear-b8-k1-v8-math500-c64
run_ar aime25 10 64 ar-aime25-c64
run_linear_uno aime25 10 64 uno-linear-b8-k1-v8-aime25-c64
```
Run the six batch-1 AR and tree `B/K/V = 16/32/32` rows:
```bash
run_ar gsm8k 1 1 ar-gsm8k-c1
run_tree_uno gsm8k 1 1 uno-tree-b16-k32-v32-gsm8k-c1
run_ar math500 1 1 ar-math500-c1
run_tree_uno math500 1 1 uno-tree-b16-k32-v32-math500-c1
run_ar aime25 10 1 ar-aime25-c1
run_tree_uno aime25 10 1 uno-tree-b16-k32-v32-aime25-c1
```
Each output directory contains raw generations, per-answer grades, and
`summary.json` and `summary.md`. AR TPF is one. UNO TPF counts both full
target-model forwards in each cycle: the diffusion-pathway draft and
AR-pathway verification forwards.
## Other speculative decoders
The runner uses the same public option names as `sglang serve`. For example,
DFLASH can be evaluated with:
```bash
PYTHONPATH=python python -m benchmark.uno.run_math_eval \
"${COMMON_ARGS[@]}" \
--benchmark math500 \
--num-samples 1 \
--output-dir "$RESULT_ROOT/dflash-b8-math500-c64" \
--max-running-requests 64 \
--speculative-algorithm DFLASH \
--speculative-draft-model-path z-lab/Qwen3-8B-DFlash-b16 \
--speculative-dflash-block-size 8 \
--speculative-draft-attention-backend fa3
```
EAGLE or EAGLE3 can be evaluated with the corresponding draft model:
```bash
export EAGLE_DRAFT_MODEL=/path/to/compatible-eagle-draft-model
PYTHONPATH=python python -m benchmark.uno.run_math_eval \
"${COMMON_ARGS[@]}" \
--benchmark math500 \
--num-samples 1 \
--output-dir "$RESULT_ROOT/eagle3-b8-math500-c64" \
--max-running-requests 64 \
--speculative-algorithm EAGLE3 \
--speculative-draft-model-path "$EAGLE_DRAFT_MODEL" \
--speculative-num-steps 7 \
--speculative-eagle-topk 1 \
--speculative-num-draft-tokens 8
```
For EAGLE and DFLASH, TPF follows SGLang's acceptance-length convention and
counts generated tokens per target verification forward.
+1
View File
@@ -0,0 +1 @@
"""Speculative-decoding benchmark utilities."""
+159
View File
@@ -0,0 +1,159 @@
"""Pinned dataset preparation for speculative-decoding math evaluation."""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any, NamedTuple
MATH_INSTRUCTION = "Please reason step by step and put your final answer in \\boxed{}."
class BenchmarkConfig(NamedTuple):
name: str
expected_rows: int
instruction: str
chat_template_kwargs: dict[str, object]
BENCHMARKS = {
name: BenchmarkConfig(
name=name,
expected_rows=expected_rows,
instruction=MATH_INSTRUCTION,
chat_template_kwargs={"reasoning_effort": "high"},
)
for name, expected_rows in (
("gsm8k", 1319),
("math500", 500),
("aime24", 30),
("aime25", 30),
("aime26", 30),
)
}
DATASET_REVISIONS = {
"openai/gsm8k": "740312add88f781978c0658806c59bc2815b9866",
"HuggingFaceH4/MATH-500": "6e4ed1a2a79af7d8630a6b768ec859cb5af4d3be",
"hypaai/Hypa_AIME2024": "11ab79f0eed5f4fdf3d469b466663ab86bbd77c8",
"math-ai/aime25": "563bb8404243c5f09de6ec262f2db674fe5bce9b",
"math-ai/aime26": "79037aebdb6580008fb960d17cb21fd3099083e3",
}
def get_benchmark(name: str) -> BenchmarkConfig:
try:
return BENCHMARKS[name]
except KeyError as exc:
available = ", ".join(BENCHMARKS)
raise KeyError(f"Unknown benchmark {name!r}. Available: {available}") from exc
def _load_dataset(
repo_id: str,
config_name: str | None = None,
*,
split: str,
) -> list[dict[str, Any]]:
from datasets import load_dataset
dataset = load_dataset(
repo_id,
config_name,
split=split,
revision=DATASET_REVISIONS[repo_id],
)
return [dict(row) for row in dataset]
def _prepare_gsm8k() -> list[dict[str, Any]]:
rows = _load_dataset("openai/gsm8k", "main", split="test")
records = []
for index, row in enumerate(rows):
prompt = f"Q: {row['question']}\nA: Let's think step by step."
records.append(
{
"row": index,
"ground_truth": row["answer"],
"chat_input": [{"role": "user", "content": prompt}],
}
)
return records
def _prepare_math500() -> list[dict[str, Any]]:
rows = _load_dataset("HuggingFaceH4/MATH-500", split="test")
return [
{
"row": index,
"ground_truth": row["answer"],
"chat_input": [{"role": "user", "content": row["problem"]}],
}
for index, row in enumerate(rows)
]
def _prepare_aime(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
records = []
for index, row in enumerate(rows):
prompt = f"Question: {row['problem']}\nAnswer:"
records.append(
{
"row": index,
"ground_truth": str(row.get("answer", row.get("solution"))).strip(),
"chat_input": [{"role": "user", "content": prompt}],
}
)
return records
def _prepare_aime24() -> list[dict[str, Any]]:
return _prepare_aime(_load_dataset("hypaai/Hypa_AIME2024", split="english"))
def _prepare_aime25() -> list[dict[str, Any]]:
return _prepare_aime(_load_dataset("math-ai/aime25", split="test"))
def _prepare_aime26() -> list[dict[str, Any]]:
return _prepare_aime(_load_dataset("math-ai/aime26", split="test"))
BUILDERS = {
"gsm8k": _prepare_gsm8k,
"math500": _prepare_math500,
"aime24": _prepare_aime24,
"aime25": _prepare_aime25,
"aime26": _prepare_aime26,
}
def _row_count(path: Path) -> int:
with path.open("rb") as handle:
return sum(bool(line.strip()) for line in handle)
def prepare_benchmark_data(name: str, *, output_dir: Path) -> Path:
benchmark = get_benchmark(name)
output = output_dir / f"{name}.jsonl"
if output.is_file():
count = _row_count(output)
if count == benchmark.expected_rows:
return output
raise ValueError(
f"{output} has {count} rows; expected {benchmark.expected_rows}"
)
records = BUILDERS[name]()
if len(records) != benchmark.expected_rows:
raise ValueError(
f"Expected {benchmark.expected_rows} rows for {name}, got {len(records)}"
)
output.parent.mkdir(parents=True, exist_ok=True)
temporary = output.with_suffix(".jsonl.tmp")
with temporary.open("w", encoding="utf-8") as handle:
for record in records:
handle.write(json.dumps(record, ensure_ascii=False) + "\n")
temporary.replace(output)
return output
+310
View File
@@ -0,0 +1,310 @@
"""Minimal math scorer adapted from Nano-vLLM-UNO's Eval360 grader."""
from __future__ import annotations
import re
from typing import Any
try:
from math_verify import parse, verify
except ModuleNotFoundError: # Allow CLI help without optional eval dependencies.
parse = None
verify = None
def _require_math_verify() -> None:
if parse is None or verify is None:
raise RuntimeError("math scoring requires the 'math_verify' package")
def extract_last_boxed_content(text: str | None) -> str | None:
if not text:
return None
matches = list(re.finditer(r"\\(?:boxed|fbox)\s*\{", text, re.IGNORECASE))
for match in reversed(matches):
start = match.end()
depth = 1
index = start
while index < len(text) and depth:
depth += (text[index] == "{") - (text[index] == "}")
index += 1
if depth == 0:
return text[start : index - 1].strip().strip("$").strip()
return None
def _generation_list(row: dict[str, Any]) -> list[str]:
if isinstance(row.get("generations"), list):
return [str(value or "") for value in row["generations"]]
return [str(row.get("generation") or "")]
def _get_accuracy(correct: list[bool]) -> float:
return sum(correct) / len(correct) if correct else float("nan")
def _source_key(row: dict[str, Any], fallback: int) -> str:
return str(row.get("source_row", row.get("row", fallback)))
def normalize_answer_text(text: str | None) -> str:
if text is None:
return ""
text = str(text).strip().strip("$").strip()
gsm8k_answer = re.search(r"####\s*([^\n]+)", text)
if gsm8k_answer:
text = gsm8k_answer.group(1).strip()
text = re.sub(r"\\(?:boxed|fbox)\s*\{(.+)\}\s*$", r"\1", text)
text = text.replace("\\dfrac", "\\frac").replace("\\tfrac", "\\frac")
text = re.sub(r"\\(?:left|right|bigl|bigr|Bigl|Bigr|big|Big)", "", text)
text = text.replace("\\displaystyle", "")
text = re.sub(r"\\mathbf\s*\{([^{}]+)\}", r"\1", text)
text = re.sub(r"\\mathbf\s+([A-Za-z])", r"\1", text)
text = re.sub(r"\\\\\s*\[[^\]]+\]", r"\\\\", text)
text = re.sub(r"\\frac\s*\{([^{}]+)\}\s*\{([^{}]+)\}", r"\1/\2", text)
text = re.sub(r"\\frac\s*([+-]?\d+)\s*([+-]?\d+)", r"\1/\2", text)
text = text.replace("{,}", "")
text = re.sub(r"\\(?:,|;|!|:)", "", text)
text = text.replace("\\%", "%")
text = text.replace("\\$", "").replace("$", "")
text = re.sub(r"\\(?:text|mathrm)\s*\{([^{}]*)\}", r"\1", text)
text = re.sub(r"\^\{([^{}])\}", r"^\1", text)
text = re.sub(r"_\{([^{}])\}", r"_\1", text)
text = re.sub(r"\s+", "", text)
text = re.sub(r"(?<=\d),(?=\d{3}(?:\D|$))", "", text)
return re.sub(r"^[A-Za-z]+=(?=\\begin\{pmatrix\})", "", text)
def _has_plain_variable(text: str) -> bool:
text = re.sub(r"\\(?:text|mathrm)\s*\{[^{}]*\}", "", text)
text = re.sub(r"\\[A-Za-z]+", "", text)
return bool(re.search(r"[A-Za-z]", text))
def _parse_boxed_content(boxed: str) -> Any:
_require_math_verify()
boxed = normalize_answer_text(boxed)
leading_number = re.match(r"\s*([-+]?[0-9][0-9,]*(?:\.[0-9]+)?)", boxed)
if leading_number and re.fullmatch(r"[A-Za-z]+", boxed[leading_number.end() :]):
answer = leading_number.group(1).replace(",", "")
try:
return parse(answer)
except Exception:
return answer
if _has_plain_variable(boxed):
return boxed
try:
parsed = parse(f"${boxed}$") or parse(boxed)
if parsed:
return parsed
except Exception:
pass
if leading_number:
answer = leading_number.group(1).replace(",", "")
try:
return parse(answer)
except Exception:
return answer
return boxed
def _parse_unboxed_answer(text: str) -> Any:
_require_math_verify()
try:
parsed = parse(f"${text}$") or parse(text)
if parsed:
return parsed
except Exception:
pass
patterns = [
r"The answer is:?\s*\$?([\-0-9\.,]+)",
r"#### ?\$?([\-0-9\.,]+)",
r"Therefore,? the answer is:?\s*\$?([\-0-9\.,]+)",
r"So,? the answer is:?\s*\$?([\-0-9\.,]+)",
r"Thus,? the answer is:?\s*\$?([\-0-9\.,]+)",
r"Hence,? the answer is:?\s*\$?([\-0-9\.,]+)",
r"Final answer:?\s*\$?([\-0-9\.,]+)",
r"The final answer is:?\s*\$?([\-0-9\.,]+)",
r"The answer is:?\s*\$?([\-0-9\.,]+)\s*(?:miles?|minutes?|hours?|dollars?|GB)?",
(
r"=\s*\$?([\-0-9\.,]+)"
r"\s*(?:miles?|minutes?|hours?|dollars?|GB)?"
r"\.?\s*(?:The answer|$)"
),
]
for pattern in patterns:
matches = re.findall(pattern, text, re.IGNORECASE)
if matches:
answer = matches[-1].replace(",", "").strip().rstrip(".")
try:
return parse(answer)
except Exception:
return answer
sentence_end = (
r"(?:is|are|equals?|makes?|has|have|gets?|arrives?|covers?|travels?)"
r"\s+\$?([\-0-9\.,]+)(?:\s*(?:miles?|minutes?|hours?|dollars?|GB))?"
r"\.?\s*$"
)
match = re.search(sentence_end, text, re.MULTILINE | re.IGNORECASE)
if match:
answer = match.group(1).replace(",", "").strip().rstrip(".")
try:
return parse(answer)
except Exception:
return answer
for sentence in reversed(text.split(".")):
if "Human:" in sentence or "Assistant:" in sentence:
continue
numbers = re.findall(r"[-+]?[0-9]*\.?[0-9]+", sentence)
if numbers:
answer = numbers[-1].lstrip("0") or "0"
try:
return parse(answer)
except Exception:
return answer
return None
def _parse_answer(text: str) -> Any:
boxed = extract_last_boxed_content(text)
return _parse_boxed_content(boxed) if boxed else _parse_unboxed_answer(text)
def _answer_candidates(text: str) -> list[Any]:
boxed = extract_last_boxed_content(text)
candidates = (
[_parse_boxed_content(boxed), normalize_answer_text(boxed), boxed]
if boxed
else [_parse_unboxed_answer(text)]
)
unique = []
seen = set()
for candidate in candidates:
if candidate is not None and repr(candidate) not in seen:
seen.add(repr(candidate))
unique.append(candidate)
return unique
def _vector_components(text: str) -> tuple[str, ...] | None:
matrix = re.fullmatch(r"\\begin\{pmatrix\}(.+)\\end\{pmatrix\}", text)
if matrix:
return tuple(part for part in matrix.group(1).split(r"\\") if part)
tuple_match = re.fullmatch(r"\(([^()]+)\)", text)
if tuple_match and "," in tuple_match.group(1):
return tuple(part for part in tuple_match.group(1).split(",") if part)
return None
def _text_answers_match(answer: str, gold: str) -> bool:
answer_norm = normalize_answer_text(answer)
gold_norm = normalize_answer_text(gold)
if answer_norm == gold_norm:
return True
try:
if float(answer_norm.rstrip("%")) == float(gold_norm.rstrip("%")):
return True
except ValueError:
pass
if re.fullmatch(r"\([A-Za-z]\)", gold_norm) and answer_norm == gold_norm[1:-1]:
return True
answer_parts = _vector_components(answer_norm)
gold_parts = _vector_components(gold_norm)
return bool(answer_parts and answer_parts == gold_parts)
def _compare_answers(answer: Any, gold: str | None) -> bool:
_require_math_verify()
if not answer or gold is None:
return False
if isinstance(answer, str) and _text_answers_match(answer, gold):
return True
try:
if verify(answer, gold):
return True
gold_answer = _parse_answer(gold)
if gold_answer:
return bool(verify(gold_answer, answer))
except Exception:
return isinstance(answer, str) and _text_answers_match(answer, gold)
def grade_math_row(row: dict[str, Any]) -> dict[str, Any]:
expected = row.get("ground_truth")
expected_values = expected if isinstance(expected, list) else [expected]
correct = []
parsed_generations = []
for generation in _generation_list(row):
candidates = _answer_candidates(generation)
parsed_generations.append([str(candidate) for candidate in candidates])
correct.append(
any(
_compare_answers(candidate, str(gold))
for candidate in candidates
for gold in expected_values
if gold is not None
)
)
return {
**row,
"parsed_generations": parsed_generations,
"correct": correct,
"accuracy": _get_accuracy(correct),
"grader": "math",
}
def score_math(
rows: list[dict[str, Any]],
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
graded = [grade_math_row(row) for row in rows]
correct = [bool(value) for row in graded for value in row["correct"]]
by_source: dict[str, list[tuple[int, bool]]] = {}
for index, row in enumerate(graded):
try:
sample_index = int(row.get("sample_index", 0))
except (TypeError, ValueError):
sample_index = 0
samples = by_source.setdefault(_source_key(row, index), [])
samples.extend(
(sample_index + offset, bool(value))
for offset, value in enumerate(row["correct"])
)
per_problem = [
[value for _, value in sorted(samples)]
for samples in by_source.values()
if samples
]
samples_per_problem = max((len(samples) for samples in per_problem), default=0)
sample0 = [samples[0] for samples in per_problem]
summary = {
"grader": "math",
"num_rows": len(graded),
"num_problems": len(per_problem),
"samples_per_problem": samples_per_problem,
"num_correct": sum(correct),
"accuracy": _get_accuracy(correct),
}
if per_problem:
summary.update(
avg_at_1=_get_accuracy(sample0),
pass_at_1=_get_accuracy(sample0),
num_correct_at_1=sum(sample0),
)
if samples_per_problem > 1:
summary.update(
{
f"avg_at_{samples_per_problem}": _get_accuracy(correct),
f"pass_at_{samples_per_problem}": _get_accuracy(
[any(samples) for samples in per_problem]
),
f"num_pass_at_{samples_per_problem}": sum(
any(samples) for samples in per_problem
),
}
)
return graded, summary
+517
View File
@@ -0,0 +1,517 @@
"""Run full-dataset AR or speculative-decoding math evaluation offline."""
from __future__ import annotations
import argparse
import json
from collections import defaultdict
from pathlib import Path
from time import perf_counter
from typing import Any
from benchmark.uno.math_data import (
BENCHMARKS,
BenchmarkConfig,
get_benchmark,
prepare_benchmark_data,
)
from benchmark.uno.math_grader import score_math
CONTEXT_LENGTH = 40960
MAX_TOKENS = 2**15
SPECULATIVE_ALGORITHMS = ("EAGLE", "EAGLE3", "DFLASH", "UNO")
SPECULATIVE_OPTION_NAMES = (
"speculative_algorithm",
"speculative_draft_model_path",
"speculative_draft_model_revision",
"speculative_num_steps",
"speculative_eagle_topk",
"speculative_num_draft_tokens",
"speculative_dflash_block_size",
"speculative_draft_attention_backend",
"uno_lora_path",
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-path", default="Qwen/Qwen3-8B")
parser.add_argument("--tokenizer-path")
parser.add_argument("--revision")
parser.add_argument("--dtype", default="bfloat16")
parser.add_argument("--attention-backend", default="fa3")
parser.add_argument("--data-root", type=Path, required=True)
parser.add_argument("--output-dir", type=Path, required=True)
parser.add_argument(
"--speculative-algorithm",
type=str.upper,
choices=SPECULATIVE_ALGORITHMS,
)
parser.add_argument("--speculative-draft-model-path")
parser.add_argument("--speculative-draft-model-revision")
parser.add_argument("--speculative-num-steps", type=int)
parser.add_argument("--speculative-eagle-topk", type=int)
parser.add_argument("--speculative-num-draft-tokens", type=int)
parser.add_argument("--speculative-dflash-block-size", type=int)
parser.add_argument("--speculative-draft-attention-backend")
parser.add_argument("--uno-lora-path")
parser.add_argument(
"--benchmark",
action="append",
choices=tuple(BENCHMARKS),
help="Benchmark to run; repeat to select multiple (default: all).",
)
parser.add_argument("--limit", type=int, help="Maximum problems per benchmark.")
parser.add_argument("--num-samples", type=int, default=1)
parser.add_argument("--max-running-requests", type=int, default=4)
parser.add_argument("--context-length", type=int, default=CONTEXT_LENGTH)
parser.add_argument("--max-tokens", type=int, default=MAX_TOKENS)
parser.add_argument("--temperature", type=float, default=1.0)
parser.add_argument("--top-k", type=int, default=50)
parser.add_argument("--top-p", type=float, default=0.95)
parser.add_argument("--random-seed", type=int, default=42)
args = parser.parse_args()
_validate_args(parser=parser, args=args)
return args
def _validate_args(
*, parser: argparse.ArgumentParser, args: argparse.Namespace
) -> None:
for name in (
"num_samples",
"max_running_requests",
"context_length",
"max_tokens",
"top_k",
):
if getattr(args, name) <= 0:
parser.error(f"--{name.replace('_', '-')} must be positive")
if args.limit is not None and args.limit <= 0:
parser.error("--limit must be positive")
if not 0 < args.top_p <= 1:
parser.error("--top-p must be in (0, 1]")
if args.temperature < 0:
parser.error("--temperature must be non-negative")
def _context_reserve(args: argparse.Namespace) -> int:
width = args.speculative_num_draft_tokens
if width is None:
width = args.speculative_dflash_block_size or 1
if args.speculative_algorithm == "DFLASH":
return 2 * width
if args.speculative_algorithm == "UNO" and (args.speculative_eagle_topk or 1) > 1:
draft_width = (args.speculative_num_steps or 0) + 1
return max(width, draft_width) + 1
return width
def _model_ref(value: str) -> str:
path = Path(value).expanduser()
return str(path.resolve()) if path.exists() else value
def _read_jsonl(path: Path) -> list[dict[str, Any]]:
with path.open(encoding="utf-8") as handle:
return [json.loads(line) for line in handle if line.strip()]
def _first_value(row: dict[str, Any], *names: str) -> Any | None:
return next((row[name] for name in names if row.get(name) is not None), None)
def _format_prompt(
*,
tokenizer: Any,
messages: list[dict[str, Any]],
benchmark: BenchmarkConfig,
) -> tuple[list[int], str]:
messages = [dict(message) for message in messages]
system_message = next(
(message for message in messages if message.get("role") == "system"),
None,
)
if system_message is None:
messages.insert(0, {"role": "system", "content": benchmark.instruction})
elif benchmark.instruction not in system_message["content"]:
existing = system_message["content"].strip()
system_message["content"] = f"{benchmark.instruction}\n\n{existing}"
rendered = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
**benchmark.chat_template_kwargs,
)
token_ids = tokenizer.apply_chat_template(
messages,
tokenize=True,
add_generation_prompt=True,
return_dict=False,
**benchmark.chat_template_kwargs,
)
return list(token_ids), str(rendered)
def _prepare_prompts(
*,
args: argparse.Namespace,
tokenizer: Any,
context_reserve: int,
) -> list[dict[str, Any]]:
prompts = []
for benchmark_name in args.benchmark or BENCHMARKS:
benchmark = get_benchmark(benchmark_name)
data_path = prepare_benchmark_data(
benchmark.name,
output_dir=args.data_root,
)
rows = _read_jsonl(data_path)
rows = rows[: args.limit] if args.limit is not None else rows
for row_index, row in enumerate(rows):
input_ids, rendered = _format_prompt(
tokenizer=tokenizer,
messages=row["chat_input"],
benchmark=benchmark,
)
available = args.context_length - len(input_ids) - context_reserve
if available < args.max_tokens:
source = _first_value(row, "id", "problem_id", "index", "row")
source = row_index if source is None else source
raise ValueError(
f"{benchmark.name}:{source} has only {max(0, available)} "
"completion tokens available; increase --context-length"
)
source = _first_value(row, "id", "problem_id", "index", "row")
source = row_index if source is None else source
for sample_index in range(args.num_samples):
prompt_id = f"{source}:sample{sample_index}"
prompts.append(
{
"id": prompt_id,
"benchmark": benchmark.name,
"input_ids": input_ids,
"row": {
"id": prompt_id,
"source_row": source,
"sample_index": sample_index,
"problem": rendered,
"ground_truth": row["ground_truth"],
"prompt_token_count": len(input_ids),
"resolved_max_tokens": args.max_tokens,
},
}
)
if not prompts:
raise ValueError("No prompts were prepared")
return prompts
def _engine_options(
*,
args: argparse.Namespace,
) -> dict[str, Any]:
tokenizer_path = _model_ref(args.tokenizer_path or args.model_path)
options: dict[str, Any] = {
"model_path": _model_ref(args.model_path),
"tokenizer_path": tokenizer_path,
"skip_tokenizer_init": True,
"context_length": args.context_length,
"dtype": args.dtype,
"random_seed": args.random_seed,
"max_running_requests": args.max_running_requests,
"attention_backend": args.attention_backend,
"log_level": "info",
}
if args.revision is not None:
options["revision"] = args.revision
for name in SPECULATIVE_OPTION_NAMES:
value = getattr(args, name)
if value is not None:
if name in (
"speculative_draft_model_path",
"uno_lora_path",
):
value = _model_ref(value)
options[name] = value
return options
def _generate(
*,
args: argparse.Namespace,
prompts: list[dict[str, Any]],
engine_options: dict[str, Any],
) -> tuple[list[dict[str, Any]], float]:
import sglang as sgl
engine = sgl.Engine(**engine_options)
sampling = [
{
"max_new_tokens": args.max_tokens,
"temperature": args.temperature,
"top_k": args.top_k,
"top_p": args.top_p,
}
for _ in prompts
]
try:
start = perf_counter()
outputs = engine.generate(
input_ids=[prompt["input_ids"] for prompt in prompts],
sampling_params=sampling,
rid=[f"{prompt['benchmark']}:{prompt['id']}" for prompt in prompts],
)
elapsed = perf_counter() - start
finally:
engine.shutdown()
return outputs, elapsed
def _build_rows(
*,
prompts: list[dict[str, Any]],
outputs: list[dict[str, Any]],
tokenizer: Any,
speculative_algorithm: str | None,
) -> list[dict[str, Any]]:
rows = []
for prompt, output in zip(prompts, outputs, strict=True):
token_ids = list(output["output_ids"])
metadata = output["meta_info"]
verify_forwards = int(metadata.get("spec_verify_ct", 0))
if speculative_algorithm == "UNO":
# Both UNO pathways are full target-model forward passes.
num_forwards = 2 * verify_forwards
elif speculative_algorithm is not None:
# EAGLE and DFLASH TPF conventionally count target verification.
num_forwards = verify_forwards
else:
num_forwards = len(token_ids)
rows.append(
{
"benchmark": prompt["benchmark"],
**prompt["row"],
"output_ids": token_ids,
"generation": tokenizer.decode(token_ids, skip_special_tokens=True),
"num_tokens": len(token_ids),
"num_forwards": num_forwards,
"tokens_per_forward": _divide(len(token_ids), num_forwards),
"sglang_meta_info": metadata,
}
)
return rows
def _divide(numerator: int | float, denominator: int | float) -> float | None:
return numerator / denominator if denominator else None
def _write_json(path: Path, value: dict[str, Any]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
json.dumps(value, indent=2, ensure_ascii=False) + "\n",
encoding="utf-8",
)
def _write_jsonl(path: Path, rows: list[dict[str, Any]]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as handle:
for row in rows:
handle.write(json.dumps(row, ensure_ascii=False) + "\n")
def _benchmark_summary(rows: list[dict[str, Any]]) -> dict[str, Any]:
tokens = sum(row["num_tokens"] for row in rows)
forwards = sum(row["num_forwards"] for row in rows)
tpfs = [
row["tokens_per_forward"]
for row in rows
if row["tokens_per_forward"] is not None
]
return {
"num_tokens": tokens,
"num_forwards": forwards,
"tokens_per_forward": _divide(tokens, forwards),
"unweighted_mean_tokens_per_forward": (sum(tpfs) / len(tpfs) if tpfs else None),
}
def _write_markdown(path: Path, summary: dict[str, Any]) -> None:
lines = [
"| Dataset | Accuracy | TPF | tok/s | tok/s/request |",
"| --- | ---: | ---: | ---: | ---: |",
]
for name, metrics in summary["by_benchmark"].items():
tokens_per_second = metrics.get("tokens_per_second")
per_request = metrics.get("tokens_per_second_per_request")
tokens_per_forward = metrics.get("tokens_per_forward")
tps = f"{tokens_per_second:.2f}" if tokens_per_second is not None else ""
request_tps = f"{per_request:.2f}" if per_request is not None else ""
tpf = f"{tokens_per_forward:.3f}" if tokens_per_forward is not None else ""
lines.append(
f"| {name} | {metrics['accuracy']:.2%} | {tpf} | {tps} | {request_tps} |"
)
lines.extend(
[
"| **Average** | "
f"**{summary['macro_average_accuracy']:.2%}** | "
f"**{summary['macro_average_tokens_per_forward']:.3f}** | "
f"**{summary['tokens_per_second']:.2f}** | "
f"**{summary['tokens_per_second_per_request']:.2f}** |",
"",
"Accuracy and TPF in the Average row are unweighted means across "
"datasets. Throughput is aggregate output tokens divided by timed "
"generation seconds; tok/s/request divides it by max running requests.",
]
)
path.write_text("\n".join(lines) + "\n", encoding="utf-8")
def _write_results(
*,
args: argparse.Namespace,
rows: list[dict[str, Any]],
generation_seconds: float,
engine_options: dict[str, Any],
) -> None:
rows_by_benchmark: dict[str, list[dict[str, Any]]] = defaultdict(list)
for row in rows:
output_row = {key: value for key, value in row.items() if key != "benchmark"}
rows_by_benchmark[row["benchmark"]].append(output_row)
benchmark_metrics = {}
totals = {"prompts": 0, "completions": 0, "correct": 0, "tokens": 0, "forwards": 0}
request_tpfs = []
for name, benchmark_rows in rows_by_benchmark.items():
output_dir = args.output_dir / name
_write_jsonl(output_dir / "generations.jsonl", benchmark_rows)
graded, scores = score_math(benchmark_rows)
_write_jsonl(output_dir / "grades.jsonl", graded)
_write_json(output_dir / "scores.json", scores)
metrics = {
"num_prompts": scores["num_problems"],
"num_completions": scores["num_rows"],
"num_correct": scores["num_correct"],
"accuracy": scores["accuracy"],
**_benchmark_summary(benchmark_rows),
}
benchmark_metrics[name] = metrics
totals["prompts"] += scores["num_problems"]
totals["completions"] += scores["num_rows"]
totals["correct"] += scores["num_correct"]
totals["tokens"] += metrics["num_tokens"]
totals["forwards"] += metrics["num_forwards"]
request_tpfs.extend(
row["tokens_per_forward"]
for row in benchmark_rows
if row["tokens_per_forward"] is not None
)
tokens_per_second = totals["tokens"] / generation_seconds
tokens_per_second_per_request = tokens_per_second / args.max_running_requests
if len(benchmark_metrics) == 1:
only_metrics = next(iter(benchmark_metrics.values()))
only_metrics["tokens_per_second"] = tokens_per_second
only_metrics["tokens_per_second_per_request"] = tokens_per_second_per_request
dataset_accuracies = [value["accuracy"] for value in benchmark_metrics.values()]
dataset_tpfs = [
value["tokens_per_forward"]
for value in benchmark_metrics.values()
if value["tokens_per_forward"] is not None
]
summary = {
"engine": "sglang-offline",
"mode": _mode_name(args),
"model_path": engine_options["model_path"],
"tokenizer_path": engine_options["tokenizer_path"],
"revision": args.revision,
"dtype": args.dtype,
"attention_backend": args.attention_backend,
"speculative_parameters": {
name: engine_options[name]
for name in SPECULATIVE_OPTION_NAMES
if name in engine_options
},
"max_running_requests": args.max_running_requests,
"benchmarks": list(rows_by_benchmark),
"context_length": args.context_length,
"max_tokens": args.max_tokens,
"num_samples": args.num_samples,
"random_seed": args.random_seed,
"sampling": {
"temperature": args.temperature,
"top_k": args.top_k,
"top_p": args.top_p,
},
"num_prompts": totals["prompts"],
"num_completions": totals["completions"],
"num_correct": totals["correct"],
"accuracy": _divide(totals["correct"], totals["completions"]),
"macro_average_accuracy": sum(dataset_accuracies) / len(dataset_accuracies),
"generation_seconds": generation_seconds,
"num_tokens": totals["tokens"],
"tokens_per_second": tokens_per_second,
"tokens_per_second_per_request": tokens_per_second_per_request,
"num_forwards": totals["forwards"],
"tokens_per_forward": _divide(totals["tokens"], totals["forwards"]),
"macro_average_tokens_per_forward": sum(dataset_tpfs) / len(dataset_tpfs),
"unweighted_mean_tokens_per_forward": (
sum(request_tpfs) / len(request_tpfs) if request_tpfs else None
),
"by_benchmark": benchmark_metrics,
}
_write_json(args.output_dir / "summary.json", summary)
_write_markdown(args.output_dir / "summary.md", summary)
print(json.dumps(summary, indent=2))
def _mode_name(args: argparse.Namespace) -> str:
algorithm = args.speculative_algorithm
if algorithm is None:
return "ar"
if algorithm == "UNO":
return "uno-tree" if (args.speculative_eagle_topk or 1) > 1 else "uno-linear"
return algorithm.lower()
def main() -> None:
args = parse_args()
from transformers import AutoTokenizer
engine_options = _engine_options(args=args)
tokenizer = AutoTokenizer.from_pretrained(
engine_options["tokenizer_path"],
use_fast=True,
trust_remote_code=True,
)
prompts = _prepare_prompts(
args=args,
tokenizer=tokenizer,
context_reserve=_context_reserve(args),
)
outputs, elapsed = _generate(
args=args,
prompts=prompts,
engine_options=engine_options,
)
rows = _build_rows(
prompts=prompts,
outputs=outputs,
tokenizer=tokenizer,
speculative_algorithm=args.speculative_algorithm,
)
_write_results(
args=args,
rows=rows,
generation_seconds=elapsed,
engine_options=engine_options,
)
if __name__ == "__main__":
main()
@@ -1,9 +1,9 @@
---
title: "Speculative Decoding"
metatags:
description: "SGLang speculative decoding: EAGLE-2/EAGLE-3, MTP, DFLASH, draft model configuration, and overlap-scheduler guidance."
description: "SGLang speculative decoding: EAGLE-2/EAGLE-3, MTP, UNO, DFLASH, draft model configuration, and overlap-scheduler guidance."
---
SGLang provides several speculative decoding options, including EAGLE-2/EAGLE-3, MTP, DFLASH, classic draft-model decoding, and an NGRAM-based variant. Our implementation aims to maximize speed and efficiency and is considered to be among the fastest in open-source LLM engines.
SGLang provides several speculative decoding options, including EAGLE-2/EAGLE-3, MTP, UNO, DFLASH, classic draft-model decoding, and an NGRAM-based variant. Our implementation aims to maximize speed and efficiency and is considered to be among the fastest in open-source LLM engines.
## Summary
@@ -15,6 +15,7 @@ SGLang provides several speculative decoding options, including EAGLE-2/EAGLE-3,
- [EAGLE-2 Decoding via Frequency-Ranked Speculative Sampling](#eagle-2-decoding-via-frequency-ranked-speculative-sampling)
- [EAGLE-3 Decoding](#eagle-3-decoding)
- [Multi Token Prediction](#multi-token-prediction)
- [UNO decoding](#uno-decoding)
- [DFlash Decoding](#dflash-decoding)
- [Standalone Speculative Decoding (Small Draft Model)](#standalone-speculative-decoding-small-draft-model)
- [Speculative Decoding V2 (Overlap Scheduler)](#speculative-decoding-v2-overlap-scheduler)
@@ -30,6 +31,7 @@ SGLang provides several speculative decoding options, including EAGLE-2/EAGLE-3,
- **Workload acceptance changes over time**: Use [**Adaptive speculative decoding**](./adaptive_speculative_decoding) on top of **EAGLE** with `--speculative-eagle-topk 1`.
- **Lower `lm_head` overhead for EAGLE-2**: Enable **FR-Spec** with `--speculative-token-map`.
- **Model is MTP-enabled**: Use **MTP via speculative decoding** (often with small `speculative_num_steps/topk/num_draft_tokens`, see the example section).
- **You have a UNO adapter for the target model**: Use **UNO** with `--speculative-algorithm UNO` and `--uno-lora-path ...`.
- **You have a DFlash draft checkpoint**: Use **DFLASH** with `--speculative-algorithm DFLASH` and `--speculative-draft-model-path ...`.
- **You have a smaller draft LLM**: Use **STANDALONE** (`--speculative-algorithm STANDALONE`).
- **No extra model available**: Use **NGRAM** (`--speculative-algorithm NGRAM`, CUDA-only).
@@ -86,6 +88,13 @@ SGLang provides several speculative decoding options, including EAGLE-2/EAGLE-3,
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>See <strong>Multi Token Prediction</strong> section</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Uses speculative workflow; draft path may be auto-handled for some models</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>UNO</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Target model with a trained UNO LoRA (linear chain or tree)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>No</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>--speculative-algorithm UNO</code> + <code>--uno-lora-path ...</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>CUDA and FA3; currently requires TP=PP=1</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>DFLASH</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>DFlash draft model (linear block verification)</td>
@@ -434,6 +443,88 @@ print(response.json())
---
## UNO decoding
UNO reuses the target transformer for both passes of each speculative decode cycle instead of loading a separate draft model. During the draft forward, the first row of each request uses the base weights and the remaining `B - 1` rows use a trained UNO LoRA. Target verification is entirely base-only.
UNO provides two sampling modes:
- **Linear** constructs and verifies one chain of draft tokens, similar to DFlash's proposal layout.
- **Tree** expands multiple candidates at each draft depth and uses SGLang's EAGLE tree-verification path.
For UNO modes, the three numbers are `B/K/V`: draft-forward width, candidates kept per expansion, and verification width. Thus, `Linear 8/1/8` uses an 8-token draft width, one candidate per expansion, and an 8-token verification width. UNO's `K` controls proposal-tree breadth and is unrelated to the request sampling parameter `top_k`. The command-line mapping depends on the mode:
- In linear mode, `B` and `V` are both `--speculative-num-draft-tokens`; set `--speculative-num-steps 1` and `--speculative-eagle-topk 1`.
- In tree mode, `B` is `--speculative-num-steps + 1`, `K` is `--speculative-eagle-topk`, and `V` is `--speculative-num-draft-tokens`.
Use an adapter trained for the exact target-model checkpoint. The current integration has been validated with `Qwen/Qwen3-8B`; unsupported LoRA target layers are rejected at startup. Set the adapter's local path or Hugging Face repository before starting the server:
### Linear 8/1/8
```bash Command
export UNO_LORA_PATH="s-sahoo/uno-qwen3-8B"
sglang serve \
--model-path Qwen/Qwen3-8B \
--speculative-algorithm UNO \
--uno-lora-path "$UNO_LORA_PATH" \
--speculative-num-steps 1 \
--speculative-eagle-topk 1 \
--speculative-num-draft-tokens 8 \
--attention-backend fa3 \
--tp 1 \
--host 0.0.0.0 \
--port 30000
```
### Tree 8/32/8
The following command selects tree mode because `K` is greater than one:
```bash Command
export UNO_LORA_PATH="s-sahoo/uno-qwen3-8B"
sglang serve \
--model-path Qwen/Qwen3-8B \
--speculative-algorithm UNO \
--uno-lora-path "$UNO_LORA_PATH" \
--speculative-num-steps 7 \
--speculative-eagle-topk 32 \
--speculative-num-draft-tokens 8 \
--attention-backend fa3 \
--tp 1 \
--host 0.0.0.0 \
--port 30000
```
Send requests through the standard OpenAI-compatible API:
```python Example
import openai
client = openai.Client(base_url="http://127.0.0.1:30000/v1", api_key="None")
response = client.chat.completions.create(
model="Qwen/Qwen3-8B",
messages=[{"role": "user", "content": "Explain speculative decoding."}],
max_tokens=128,
)
print(response.choices[0].message.content)
```
### Key requirements and limitations
- UNO requires CUDA with FA3 for prefill and decode, tensor and pipeline parallel sizes of 1, and no DP attention or context parallelism.
- `--uno-lora-path` loads UNO's fixed internal adapter and cannot be combined with request-selectable Multi-LoRA serving.
- UNO manages its own stochastic verification; do not set `--speculative-use-rejection-sampling`.
- Grammar decoding, returned logprobs or hidden states, sampling penalties, `min_p`, logit bias, custom logit processors, strict thinking, and deterministic inference are not yet supported.
- In tree mode, both speculative acceptance thresholds must remain at 1.0, `V` must be at least `B` and at most 128, and `V * K` must not exceed 2048. SGLang validates the remaining tree-capacity and EAGLE parent-representation constraints at startup.
- Mixed chunked prefill is disabled for UNO.
- Ordinary overlap scheduling is supported. Tree mode does not yet support PDMux or the separate `--enable-two-batch-overlap` feature.
---
## DFlash Decoding
SGLang also supports **DFLASH** speculative decoding using a dedicated draft model checkpoint. Compared with EAGLE-style tree verification, DFLASH verifies a linear draft block and is configured around a block size / draft window. This path is useful when the target model has a matching DFlash draft checkpoint, such as `meta-llama/Llama-3.1-8B-Instruct` with `z-lab/LLaMA3.1-8B-Instruct-DFlash-UltraChat`.
@@ -753,7 +844,7 @@ Below is a comprehensive list of all speculative decoding parameters available i
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--speculative-algorithm</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>str</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>None</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Algorithm to use: <code>DFLASH</code>, <code>EAGLE</code>, <code>EAGLE3</code>, <code>STANDALONE</code>, <code>NGRAM</code>, <code>NEXTN</code> (alias of <code>EAGLE</code>)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Algorithm to use: <code>UNO</code>, <code>DFLASH</code>, <code>EAGLE</code>, <code>EAGLE3</code>, <code>STANDALONE</code>, <code>NGRAM</code>, <code>NEXTN</code> (alias of <code>EAGLE</code>)</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--speculative-draft-model-path</code></td>
@@ -761,6 +852,12 @@ Below is a comprehensive list of all speculative decoding parameters available i
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>None</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Path to the draft model weights</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--uno-lora-path</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>str</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>None</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>UNO-only path or Hugging Face repository for the draft LoRA checkpoint</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--speculative-draft-model-revision</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>str</code></td>
@@ -777,19 +874,19 @@ Below is a comprehensive list of all speculative decoding parameters available i
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--speculative-num-steps</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>int</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>None</code> (auto-chosen when omitted)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Autoregressive drafting depth</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Autoregressive drafting depth. Some algorithms auto-tune it; tree UNO requires it and uses <code>B = speculative_num_steps + 1</code>.</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--speculative-eagle-topk</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>int</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>None</code> (auto-chosen when omitted)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Branching factor per drafting step</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Branching factor per drafting step. Some algorithms auto-tune it; UNO resolves an omitted value to <code>1</code>, and values greater than one select UNO tree mode.</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--speculative-num-draft-tokens</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>int</code></td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>None</code> (auto-chosen when omitted)</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Maximum number of draft tokens for verification</td>
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Maximum number of draft tokens for verification. Some algorithms auto-tune it; UNO requires it and uses it as linear width or tree verification width <code>V</code>.</td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--speculative-dflash-block-size</code></td>
@@ -16,6 +16,7 @@ _TRITON_KERNELS = [
("extend_attention", "build_unified_kv_indices"),
("prefill_attention", "context_attention_fwd"),
("merge_state", "merge_state_triton"),
("suffix_attention_merge", "merge_suffix_attention_in_place"),
("metadata", "get_num_kv_splits_triton"),
("metadata", "prepare_swa_spec_page_table_triton"),
("metadata", "normal_decode_set_metadata"),
@@ -0,0 +1,203 @@
"""Fused attention over a short sparse suffix and merge with a prefix state."""
from __future__ import annotations
import torch
import triton
import triton.language as tl
def can_use_fused_suffix_attention_merge(
*,
layer,
q: torch.Tensor,
key_cache: torch.Tensor,
value_cache: torch.Tensor,
extra_kwargs: dict,
) -> bool:
"""Whether attention can use the specialized suffix merge."""
return bool(
q.dtype in (torch.float16, torch.bfloat16)
and key_cache.dtype == q.dtype
and value_cache.dtype == q.dtype
and layer.head_dim == layer.v_head_dim
and not layer.is_cross_attention
and not layer.logit_cap
and not extra_kwargs
)
@triton.jit
def _fused_suffix_attention_merge_kernel(
q_ptr,
k_cache_ptr,
v_cache_ptr,
page_table_ptr,
suffix_seqlens_ptr,
prefix_ptr,
prefix_lse_ptr,
scale,
q_stride_t,
q_stride_h,
q_stride_d,
k_stride_t,
k_stride_h,
k_stride_d,
v_stride_t,
v_stride_h,
v_stride_d,
page_stride_t,
page_stride_s,
prefix_stride_t,
prefix_stride_h,
prefix_stride_d,
prefix_lse_stride_h,
prefix_lse_stride_t,
NUM_Q_HEADS: tl.constexpr,
NUM_KV_HEADS: tl.constexpr,
HEAD_DIM: tl.constexpr,
BLOCK_SUFFIX: tl.constexpr,
BLOCK_D: tl.constexpr,
):
token = tl.program_id(0)
q_head = tl.program_id(1)
kv_head = q_head // (NUM_Q_HEADS // NUM_KV_HEADS)
suffix_offsets = tl.arange(0, BLOCK_SUFFIX)
suffix_length = tl.load(suffix_seqlens_ptr + token)
suffix_valid = suffix_offsets < suffix_length
slots = tl.load(
page_table_ptr + token * page_stride_t + suffix_offsets * page_stride_s,
mask=suffix_valid,
other=0,
).to(tl.int64)
dims = tl.arange(0, BLOCK_D)
dim_valid = dims < HEAD_DIM
q = tl.load(
q_ptr + token * q_stride_t + q_head * q_stride_h + dims * q_stride_d,
mask=dim_valid,
other=0.0,
).to(tl.float32)
k = tl.load(
k_cache_ptr
+ slots[:, None] * k_stride_t
+ kv_head * k_stride_h
+ dims[None, :] * k_stride_d,
mask=suffix_valid[:, None] & dim_valid[None, :],
other=0.0,
).to(tl.float32)
scores = tl.sum(k * q[None, :], axis=1) * scale
scores = tl.where(suffix_valid, scores, -float("inf"))
suffix_max = tl.max(scores, axis=0)
prefix_lse = tl.load(
prefix_lse_ptr + q_head * prefix_lse_stride_h + token * prefix_lse_stride_t
).to(tl.float32)
global_max = tl.maximum(prefix_lse, suffix_max)
prefix_weight = tl.exp(prefix_lse - global_max)
suffix_weights = tl.exp(scores - global_max)
denominator = prefix_weight + tl.sum(suffix_weights, axis=0)
v = tl.load(
v_cache_ptr
+ slots[:, None] * v_stride_t
+ kv_head * v_stride_h
+ dims[None, :] * v_stride_d,
mask=suffix_valid[:, None] & dim_valid[None, :],
other=0.0,
).to(tl.float32)
suffix_numerator = tl.sum(suffix_weights[:, None] * v, axis=0)
prefix = tl.load(
prefix_ptr
+ token * prefix_stride_t
+ q_head * prefix_stride_h
+ dims * prefix_stride_d,
mask=dim_valid,
other=0.0,
).to(tl.float32)
output = (prefix * prefix_weight + suffix_numerator) / denominator
tl.store(
prefix_ptr
+ token * prefix_stride_t
+ q_head * prefix_stride_h
+ dims * prefix_stride_d,
output,
mask=dim_valid,
)
def merge_suffix_attention_in_place(
q: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
suffix_page_table: torch.Tensor,
suffix_cache_seqlens: torch.Tensor,
prefix: torch.Tensor,
prefix_lse: torch.Tensor,
softmax_scale: float,
) -> torch.Tensor:
"""Compute a short sparse suffix and merge it into ``prefix`` in place.
``prefix_lse`` uses FlashAttention's varlen layout ``[num_q_heads,
num_queries]``. The suffix page table contains physical token slots, one
row per query, and only its first ``suffix_cache_seqlens[row]`` entries are
visible.
"""
if q.ndim != 3:
raise ValueError("q must have shape [num_queries, num_q_heads, head_dim]")
num_queries, num_q_heads, head_dim = q.shape
if prefix.shape != q.shape:
raise ValueError("prefix output must have the same shape as q")
if prefix_lse.shape != (num_q_heads, num_queries):
raise ValueError("prefix_lse must have shape [num_q_heads, num_queries]")
if k_cache.ndim != 3 or v_cache.ndim != 3:
raise ValueError("flattened KV caches must have shape [slots, heads, dim]")
if k_cache.shape != v_cache.shape:
raise ValueError("K and V caches must have matching shapes")
num_kv_heads = k_cache.shape[1]
if k_cache.shape[2] != head_dim:
raise ValueError("K/V and query head dimensions must match")
if num_q_heads % num_kv_heads:
raise ValueError("query heads must be divisible by KV heads")
if suffix_page_table.ndim != 2 or suffix_page_table.shape[0] != num_queries:
raise ValueError("suffix page table must have one row per query")
if suffix_cache_seqlens.numel() != num_queries:
raise ValueError("suffix cache lengths must have one value per query")
if suffix_page_table.shape[1] == 0 or num_queries == 0:
return prefix
_fused_suffix_attention_merge_kernel[(num_queries, num_q_heads)](
q,
k_cache,
v_cache,
suffix_page_table,
suffix_cache_seqlens,
prefix,
prefix_lse,
softmax_scale,
q.stride(0),
q.stride(1),
q.stride(2),
k_cache.stride(0),
k_cache.stride(1),
k_cache.stride(2),
v_cache.stride(0),
v_cache.stride(1),
v_cache.stride(2),
suffix_page_table.stride(0),
suffix_page_table.stride(1),
prefix.stride(0),
prefix.stride(1),
prefix.stride(2),
prefix_lse.stride(0),
prefix_lse.stride(1),
NUM_Q_HEADS=num_q_heads,
NUM_KV_HEADS=num_kv_heads,
HEAD_DIM=head_dim,
BLOCK_SUFFIX=triton.next_power_of_2(suffix_page_table.shape[1]),
BLOCK_D=triton.next_power_of_2(head_dim),
num_warps=4,
num_stages=1,
)
return prefix
@@ -327,6 +327,171 @@ def _handle_dflash(server_args: ServerArgs) -> None:
)
def _handle_uno(server_args: ServerArgs) -> None:
cfg = resolving_view(server_args)
if not cfg.device.startswith("cuda"):
raise ValueError("UNO only supports CUDA.")
if cfg.speculative_draft_model_path is not None:
raise ValueError(
"UNO reuses the target model and does not accept "
"--speculative-draft-model-path."
)
if cfg.uno_lora_path is None:
raise ValueError("UNO requires --uno-lora-path.")
if cfg.enable_deterministic_inference:
raise ValueError(
"UNO does not support --enable-deterministic-inference because its "
"sampling path does not use per-request seeds."
)
if cfg.enable_strict_thinking:
raise ValueError(
"UNO does not support --enable-strict-thinking because it requires "
"grammar decoding."
)
verify_width = cfg.speculative_num_draft_tokens
if verify_width is None or int(verify_width) < 1:
raise ValueError(
"UNO requires --speculative-num-draft-tokens to be a positive "
"integer denoting the linear width or tree verify width Q."
)
verify_width = int(verify_width)
declare_resolution(
server_args,
"_handle_uno",
speculative_num_draft_tokens=verify_width,
)
candidate_top_k = (
1 if cfg.speculative_eagle_topk is None else int(cfg.speculative_eagle_topk)
)
if candidate_top_k < 1:
raise ValueError(
"UNO requires --speculative-eagle-topk to be at least 1, "
f"got {candidate_top_k}."
)
if candidate_top_k > 1:
if cfg.speculative_num_steps is None:
raise ValueError(
"UNO tree mode requires --speculative-num-steps so its draft "
"width F can be derived as speculative_num_steps + 1."
)
speculative_num_steps = int(cfg.speculative_num_steps)
if speculative_num_steps < 1:
raise ValueError(
"UNO tree mode requires --speculative-num-steps to be positive, "
f"got {speculative_num_steps}."
)
draft_width = speculative_num_steps + 1
if verify_width < draft_width:
raise ValueError(
f"UNO tree mode requires Q >= F; got Q={verify_width}, F={draft_width}."
)
if verify_width > 128:
raise ValueError(
"UNO tree mode currently supports at most Q=128 verify nodes, "
f"got Q={verify_width}."
)
frontier_slots = verify_width * candidate_top_k
if frontier_slots > 2048:
raise ValueError(
"UNO tree mode currently supports Q*K <= 2048; got "
f"Q*K={verify_width}*{candidate_top_k}={frontier_slots}."
)
tree_capacity = 1
nodes_at_depth = 1
for _ in range(speculative_num_steps):
nodes_at_depth *= candidate_top_k
tree_capacity += nodes_at_depth
if tree_capacity >= verify_width:
break
if verify_width > tree_capacity:
raise ValueError(
"UNO tree mode cannot build the requested Q from F and K: "
f"Q={verify_width} exceeds capacity={tree_capacity} for "
f"F={draft_width}, K={candidate_top_k}."
)
parent_width = candidate_top_k * max(speculative_num_steps - 1, 0) + 1
if verify_width - 1 > parent_width:
raise ValueError(
"UNO tree mode cannot represent the requested Q in EAGLE's "
"parent-list ABI: "
f"Q-1={verify_width - 1} exceeds "
f"K*(F-2)+1={parent_width} for "
f"F={draft_width}, K={candidate_top_k}."
)
if cfg.enable_pdmux:
raise ValueError("UNO tree mode does not yet support PDMux.")
if cfg.enable_two_batch_overlap:
raise ValueError("UNO tree mode does not yet support two-batch overlap.")
if (
cfg.speculative_accept_threshold_single != 1.0
or cfg.speculative_accept_threshold_acc != 1.0
):
raise ValueError(
"UNO tree mode reuses EAGLE target-only sampling and requires "
"both speculative accept thresholds to be 1.0."
)
declare_resolution(
server_args,
"_handle_uno",
speculative_num_steps=speculative_num_steps,
speculative_eagle_topk=candidate_top_k,
)
else:
for field in ("speculative_num_steps", "speculative_eagle_topk"):
old_value = getattr(cfg, field)
if old_value not in (None, 1):
logger.warning("UNO uses %s=1; overriding %s.", field, old_value)
declare_resolution(
server_args,
"_handle_uno",
speculative_num_steps=1,
speculative_eagle_topk=1,
)
if (cfg.tp_size, cfg.pp_size) != (1, 1):
raise ValueError("UNO requires TP=PP=1.")
if cfg.enable_dp_attention or cfg.attn_cp_size != 1:
raise ValueError("UNO does not support DP attention or context parallelism.")
if cfg.enable_lora or cfg.lora_paths:
raise ValueError("UNO does not support public Multi-LoRA serving.")
declare_resolution(
server_args,
"_handle_uno",
enable_lora_overlap_loading=False,
lora_strict_loading=True,
)
if cfg.speculative_use_rejection_sampling:
raise ValueError(
"UNO manages its own stochastic verification and does not use "
"--speculative-use-rejection-sampling."
)
if cfg.enable_mixed_chunk:
declare_resolution(
server_args,
"_handle_uno",
enable_mixed_chunk=False,
)
logger.warning(
"Mixed chunked prefill is disabled for UNO speculative decoding."
)
prefill_backend, decode_backend = attention_backends_of(resolved_view(server_args))
if (prefill_backend, decode_backend) != ("fa3", "fa3"):
raise ValueError(
"UNO requires FA3 for both prefill and decode attention; "
f"got prefill={prefill_backend!r}, decode={decode_backend!r}."
)
def _target_checkpoint_bundles_dspark_draft(server_args: ServerArgs) -> bool:
from sglang.srt.speculative.dspark_components.dspark_config import (
checkpoint_bundles_dspark_draft,
@@ -12,6 +12,10 @@ from sglang.kernels.ops.attention.metadata import (
prepare_swa_spec_page_table_triton,
)
from sglang.kernels.ops.attention.pa_page_table import _build_pa_page_table
from sglang.kernels.ops.attention.suffix_attention_merge import (
can_use_fused_suffix_attention_merge,
merge_suffix_attention_in_place,
)
from sglang.kernels.ops.attention.utils import assert_buffer_fits
from sglang.kernels.ops.kvcache.trtllm_mha_page_table import (
build_trtllm_mha_page_table,
@@ -1579,14 +1583,36 @@ class FlashAttentionBackend(AttentionBackend):
if use_cascade_attn:
o, softmax_lse, *rest = result
if (
use_cascade_attn
and forward_batch.spec_algorithm.is_uno()
and can_use_fused_suffix_attention_merge(
layer=layer,
q=q,
key_cache=key_cache,
value_cache=value_cache,
extra_kwargs=kwargs,
)
):
suffix_metadata = self.forward_metadata_spec_decode_expand
o = merge_suffix_attention_in_place(
q=q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
k_cache=key_cache.view(-1, layer.tp_k_head_num, layer.head_dim),
v_cache=value_cache.view(-1, layer.tp_v_head_num, layer.v_head_dim),
suffix_page_table=suffix_metadata.page_table,
suffix_cache_seqlens=suffix_metadata.cache_seqlens_int32,
prefix=o,
prefix_lse=softmax_lse,
softmax_scale=layer.scaling,
)
elif use_cascade_attn:
o_expand, softmax_lse_expand, *rest_expand = flash_attn_with_kvcache(
q=q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
# Here metadata_expand.page_table is not divided with page_size.
# This is because we loose the fine control of what token to attend,
# but has to attend to some block completely.
# The suffix table stores physical token slots, so expose
# the paged cache as page-size-one blocks.
k_cache=key_cache.view(-1, 1, layer.tp_k_head_num, layer.head_dim),
v_cache=value_cache.view(
-1, 1, layer.tp_v_head_num, layer.head_dim
-1, 1, layer.tp_v_head_num, layer.v_head_dim
),
page_table=self.forward_metadata_spec_decode_expand.page_table,
cache_seqlens=self.forward_metadata_spec_decode_expand.cache_seqlens_int32,
@@ -23,6 +23,9 @@ class BaseLoRABackend(LoRABackendLmHeadMixing):
device: the device where the backend runs.
"""
supports_lora_a_overlap = False
skip_inactive_lora_batches = False
# Supporting backends implement init_prefill_cuda_graph_batch_info() and
# honor use_prefill_cuda_graph in prepare_lora_batch().
supports_prefill_cuda_graph: bool = False
@@ -52,6 +55,14 @@ class BaseLoRABackend(LoRABackendLmHeadMixing):
self.lm_head_pass_batch_infos = None
self._lm_head_pass_idx = None
def validate_lora_targets(
self,
base_model: torch.nn.Module,
target_modules: set[str],
) -> None:
"""Raise before wrapping when this backend cannot execute its targets."""
pass
def run_lora_a_embedding(
self,
input_ids: torch.Tensor,
@@ -351,6 +362,20 @@ class BaseLoRABackend(LoRABackendLmHeadMixing):
"""
pass
def prepare_lora_token_segments(
self,
*,
segment_lens: list[int],
weight_indices: list[int],
lora_ranks: list[int],
scalings: list[float],
) -> None:
"""Prepare explicit eager token-row LoRA segments."""
raise NotImplementedError(
f"LoRA backend {type(self).__name__} does not support explicit "
"token segments."
)
@triton.jit
def _compute_moe_lora_info_kernel(
@@ -44,6 +44,13 @@ def create_torch_native_backend():
return TorchNativeLoRABackend
@register_lora_backend("uno_cublas")
def create_uno_cublas_backend():
from sglang.srt.lora.backend.uno_cublas_backend import UnoCublasLoRABackend
return UnoCublasLoRABackend
@register_lora_backend("flashinfer")
def create_flashinfer_backend():
raise ValueError(
@@ -362,6 +362,44 @@ class TritonLoRABackend(BaseLoRABackend):
self._prepare_lm_head_batch_info(forward_batch, weight_indices, batch_info)
)
def prepare_lora_token_segments(
self,
*,
segment_lens: list[int],
weight_indices: list[int],
lora_ranks: list[int],
scalings: list[float],
) -> None:
"""Install explicit eager token-row routing metadata."""
self.reset_batch_state()
segment_lens_tensor = torch.tensor(
segment_lens, dtype=torch.int32, device=self.device
)
segment_indptr = torch.zeros(
len(segment_lens) + 1, dtype=torch.int32, device=self.device
)
segment_indptr[1:] = torch.cumsum(segment_lens_tensor, dim=0)
self.batch_info = LoRABatchInfo(
use_cuda_graph=False,
bs=len(segment_lens),
num_segments=len(segment_lens),
seg_lens=segment_lens_tensor,
seg_indptr=segment_indptr,
max_len=max(segment_lens),
weight_indices=torch.tensor(
weight_indices, dtype=torch.int32, device=self.device
),
lora_ranks=torch.tensor(lora_ranks, dtype=torch.int64, device=self.device),
scalings=torch.tensor(scalings, dtype=torch.float, device=self.device),
permutation=None,
expected_tokens=sum(segment_lens),
)
# These segments already describe physical token-row order.
self.sgemm_batch_info = None
def _prepare_lm_head_batch_info(
self,
forward_batch: ForwardBatch,
@@ -0,0 +1,455 @@
"""Single-adapter LoRA backend for UNO draft forwards.
Parallel-linear layers overlap LoRA-A with the base GEMM on an auxiliary CUDA
stream. Other dense LoRA layers use the inherited Triton implementation.
Single-request cuBLAS batches operate only on draft rows; larger batches use
one ``mm``/``addmm_`` over all rows and zero the seed-row LoRA hidden states
between the two GEMMs.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Optional
import torch
from sglang.srt.lora.backend.triton_backend import TritonLoRABackend
from sglang.srt.lora.utils import LoRABatchInfo
@dataclass(frozen=True)
class _UnoSingleAdapterRoute:
weight_index: int
rank: int
scaling: float
batch_size: int
forward_width: int
total_rows: int
active_rows: int
@property
def lora_rows(self) -> int:
# For C=1, skipping the seed row makes both GEMMs smaller. For C>1,
# one GEMM across C*F rows is faster than C tiny (F-1)-row GEMMs.
return self.active_rows if self.batch_size == 1 else self.total_rows
@dataclass(frozen=True)
class _PendingLoRAA:
output: torch.Tensor
producer_stream: torch.cuda.Stream
class UnoCublasLoRABackend(TritonLoRABackend):
"""Fast LoRA backend for UNO's draft forwards.
Use cuBLAS and CUDA streams.
Reuse multi-LoRA batch metadata from the Triton parent.
"""
name = "uno_cublas"
supports_lora_a_overlap = True
# K2's runners prepare base-only LoRA metadata whenever any internal
# manager exists. UNO never exposes request-selectable adapters, so those
# prefill/warmup batches must stay on the plain base-model path.
skip_inactive_lora_batches = True
def __init__(
self,
max_loras_per_batch: int,
device: torch.device,
**kwargs,
):
super().__init__(max_loras_per_batch, device, **kwargs)
# Different CUDA-graph runners may capture and replay on different main
# streams. Give each main stream its own LoRA-A side stream so concurrently
# replayed graphs do not serialize or interfere through one shared stream.
self._lora_a_streams: dict[torch.cuda.Stream, torch.cuda.Stream] = {}
self._pending_lora_a: Optional[_PendingLoRAA] = None
self._use_cublas_lora_b = False
def reset_batch_state(self):
self._pending_lora_a = None
self._use_cublas_lora_b = False
super().reset_batch_state()
def validate_lora_targets(
self,
base_model: torch.nn.Module,
target_modules: set[str],
) -> None:
"""Reject target layers that cannot honor UNO's token-row routing."""
from sglang.srt.layers.linear import (
ColumnParallelLinear,
ReplicatedLinear,
RowParallelLinear,
)
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.layers.utils import get_layer_id
from sglang.srt.models.inkling_common.dense_mlp import InklingBatchDenseMLP
unsupported: list[str] = []
# Embedding and LM-head wrappers use the inherited Triton kernels and
# are handled separately by LoRAManager. Decoder-layer projections use
# either the overlapped cuBLAS path or the Triton ReplicatedLinear path.
supported = (ColumnParallelLinear, RowParallelLinear, ReplicatedLinear)
target_moe = {"gate_up_proj", "down_proj"}.issubset(target_modules)
for module_name, module in base_model.named_modules():
parts = module_name.split(".")
named_target = bool(parts) and (
parts[-1] in target_modules or ".".join(parts[-2:]) in target_modules
)
special_moe_target = target_moe and isinstance(
module, (FusedMoE, InklingBatchDenseMLP)
)
if not (named_target or special_moe_target):
continue
if get_layer_id(module_name) is None:
continue
if not isinstance(module, supported):
unsupported.append(f"{module_name} ({type(module).__name__})")
if unsupported:
raise ValueError(
"UNO's LoRA backend cannot execute these target modules: "
+ ", ".join(sorted(set(unsupported)))
)
def prepare_lora_token_segments(
self,
*,
segment_lens: list[int],
weight_indices: list[int],
lora_ranks: list[int],
scalings: list[float],
) -> None:
super().prepare_lora_token_segments(
segment_lens=segment_lens,
weight_indices=weight_indices,
lora_ranks=lora_ranks,
scalings=scalings,
)
route: Optional[_UnoSingleAdapterRoute] = None
if (
len(segment_lens) >= 2
and len(segment_lens) % 2 == 0
and len(weight_indices) == len(segment_lens)
and segment_lens[0] == 1
and segment_lens[1] > 0
):
batch_size = len(segment_lens) // 2
forward_width = segment_lens[1] + 1
base_index, adapter_index = weight_indices[:2]
adapter_rank = lora_ranks[adapter_index]
if (
base_index != adapter_index
and lora_ranks[base_index] == 0
and adapter_rank > 0
and all(
segment_lens[2 * index] == 1
and segment_lens[2 * index + 1] == forward_width - 1
and weight_indices[2 * index] == base_index
and weight_indices[2 * index + 1] == adapter_index
for index in range(batch_size)
)
):
route = _UnoSingleAdapterRoute(
weight_index=adapter_index,
rank=adapter_rank,
scaling=float(scalings[adapter_index]),
batch_size=batch_size,
forward_width=forward_width,
total_rows=sum(segment_lens),
active_rows=batch_size * (forward_width - 1),
)
# Each CUDA-graph bucket retains its own batch_info. Store the immutable
# UNO route there so switching back to a captured bucket restores the
# route corresponding to that bucket, rather than using metadata last
# written by another bucket.
self.batch_info.uno_single_adapter_route = route
def _route(self) -> _UnoSingleAdapterRoute:
route = getattr(self.batch_info, "uno_single_adapter_route", None)
if route is None:
raise RuntimeError("UNO cuBLAS execution requires an active UNO route.")
return route
@staticmethod
def _output_offsets(output_offset_cpu, output_offset) -> list[int]:
offsets = output_offset_cpu if output_offset_cpu is not None else output_offset
return [int(offset) for offset in offsets.tolist()]
@staticmethod
def _lora_a_input(
x: torch.Tensor,
route: _UnoSingleAdapterRoute,
) -> torch.Tensor:
if route.batch_size == 1:
return x[1:]
return x
def _compute_lora_a(
self,
lora_input: torch.Tensor,
active_a: torch.Tensor,
route: _UnoSingleAdapterRoute,
*,
output: Optional[torch.Tensor] = None,
) -> torch.Tensor:
hidden = torch.mm(lora_input, active_a.t(), out=output)
if route.batch_size > 1:
# LoRA must not affect each request's seed token. Zeroing C rows
# is cheaper than masking every hidden element or running C tiny
# strided-batched GEMMs.
hidden.view(route.batch_size, route.forward_width, -1)[:, 0].zero_()
return hidden
def _accumulate_lora_b(
self,
*,
hidden: torch.Tensor,
active_b: torch.Tensor,
base_output: torch.Tensor,
output_start: int,
output_end: int,
route: _UnoSingleAdapterRoute,
) -> None:
if route.batch_size == 1:
base_output[1:, output_start:output_end].addmm_(
hidden,
active_b.t(),
beta=1.0,
alpha=route.scaling,
)
return
base_output[:, output_start:output_end].addmm_(
hidden,
active_b.t(),
beta=1.0,
alpha=route.scaling,
)
def _run_lora_b(
self,
x: torch.Tensor,
weights: torch.Tensor,
base_output: torch.Tensor,
) -> torch.Tensor:
route = self._route()
active_b = weights[route.weight_index, :, : route.rank]
self._accumulate_lora_b(
hidden=x[:, : route.rank],
active_b=active_b,
base_output=base_output,
output_start=0,
output_end=active_b.shape[0],
route=route,
)
return base_output
def run_lora_a_sgemm(
self,
x: torch.Tensor,
weights: torch.Tensor,
pruned_batch_info: LoRABatchInfo = None,
stack_num: int = 1,
*args,
**kwargs,
) -> torch.Tensor:
pending = self._pending_lora_a
if pending is None:
self._use_cublas_lora_b = False
return super().run_lora_a_sgemm(
x,
weights,
pruned_batch_info,
stack_num,
*args,
**kwargs,
)
output = self._consume_lora_a_overlap(pending)
self._use_cublas_lora_b = True
return output
def run_lora_b_sgemm(
self,
x: torch.Tensor,
weights: torch.Tensor,
base_output: torch.Tensor = None,
pruned_batch_info: LoRABatchInfo = None,
*args,
**kwargs,
) -> torch.Tensor:
if not self._use_cublas_lora_b:
return super().run_lora_b_sgemm(
x,
weights,
base_output,
pruned_batch_info,
*args,
**kwargs,
)
self._use_cublas_lora_b = False
return self._run_lora_b(x, weights, base_output)
def _run_stacked_lora(
self,
*,
lora_b: torch.Tensor,
base_output: torch.Tensor,
output_offset,
output_offset_cpu,
num_slices: int,
pending: _PendingLoRAA,
) -> torch.Tensor:
route = self._route()
hidden = self._consume_lora_a_overlap(pending)
offsets = self._output_offsets(output_offset_cpu, output_offset)
for slice_index in range(num_slices):
input_start = slice_index * route.rank
input_end = input_start + route.rank
output_start = offsets[slice_index]
output_end = offsets[slice_index + 1]
active_b = lora_b[
route.weight_index,
output_start:output_end,
: route.rank,
]
self._accumulate_lora_b(
hidden=hidden[:, input_start:input_end],
active_b=active_b,
base_output=base_output,
output_start=output_start,
output_end=output_end,
route=route,
)
return base_output
def run_qkv_lora(
self,
x: torch.Tensor,
qkv_lora_a: torch.Tensor,
qkv_lora_b: torch.Tensor,
output_offset: torch.Tensor,
max_qkv_out_dim: int,
base_output: torch.Tensor = None,
n_slices: int = 3,
*args,
output_offset_cpu=None,
**kwargs,
) -> torch.Tensor:
pending = self._pending_lora_a
if pending is None:
return super().run_qkv_lora(
x,
qkv_lora_a,
qkv_lora_b,
output_offset,
max_qkv_out_dim,
base_output,
n_slices,
*args,
**kwargs,
)
return self._run_stacked_lora(
lora_b=qkv_lora_b,
base_output=base_output,
output_offset=output_offset,
output_offset_cpu=output_offset_cpu,
num_slices=n_slices,
pending=pending,
)
def run_gate_up_lora(
self,
x: torch.Tensor,
gate_up_lora_a: torch.Tensor,
gate_up_lora_b: torch.Tensor,
base_output: torch.Tensor = None,
*args,
output_offset=None,
output_offset_cpu=None,
**kwargs,
) -> torch.Tensor:
pending = self._pending_lora_a
if pending is None:
return super().run_gate_up_lora(
x,
gate_up_lora_a,
gate_up_lora_b,
base_output,
*args,
**kwargs,
)
return self._run_stacked_lora(
lora_b=gate_up_lora_b,
base_output=base_output,
output_offset=output_offset,
output_offset_cpu=output_offset_cpu,
num_slices=2,
pending=pending,
)
def start_lora_a_overlap(
self,
x: torch.Tensor,
weights: torch.Tensor,
*,
num_slices: int = 1,
) -> None:
"""Launch LoRA-A on the auxiliary stream before the base GEMM."""
if self._pending_lora_a is not None:
raise RuntimeError("Previous UNO LoRA-A overlap was not consumed.")
route = self._route()
out_dim = num_slices * route.rank
lora_input = self._lora_a_input(x, route)
active_a = weights[route.weight_index, :out_dim]
main_stream = torch.cuda.current_stream(x.device)
stream = self._lora_a_streams.get(main_stream)
if stream is None:
if torch.cuda.is_current_stream_capturing():
raise RuntimeError(
"UNO LoRA overlap side stream was not created during "
"CUDA-graph warmup."
)
stream = torch.cuda.Stream(device=x.device)
self._lora_a_streams[main_stream] = stream
# Allocate on the main/consumer stream before handing the buffer to
# the auxiliary stream. This keeps CUDA-graph allocator ownership and
# the eventual LoRA-B consumer on the same stream.
output = torch.empty(
(route.lora_rows, out_dim),
dtype=x.dtype,
device=x.device,
)
stream.wait_stream(main_stream)
with torch.cuda.stream(stream):
self._compute_lora_a(lora_input, active_a, route, output=output)
self._pending_lora_a = _PendingLoRAA(
output=output,
producer_stream=stream,
)
def _consume_lora_a_overlap(
self,
pending: _PendingLoRAA,
) -> torch.Tensor:
self._pending_lora_a = None
torch.cuda.current_stream(pending.output.device).wait_stream(
pending.producer_stream
)
return pending.output
+36 -3
View File
@@ -66,7 +66,15 @@ class BaseLayerWithLoRA(nn.Module):
has LoRA batch metadata. batch_info is None on DP-attention idle
forwards (see LoRAManager.prepare_lora_batch), so idle forwards take
the base path."""
return self.set_lora and self.lora_backend.batch_info is not None
batch_info = self.lora_backend.batch_info
return (
self.set_lora
and batch_info is not None
and (
not self.lora_backend.skip_inactive_lora_batches
or batch_info.has_active_lora
)
)
def set_lora_info(self, *args):
pass
@@ -482,14 +490,22 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
)
return lora_output
def start_lora_a_overlap(self, x: torch.Tensor) -> None:
if self.lora_backend.supports_lora_a_overlap:
self.lora_backend.start_lora_a_overlap(x, self.A_buffer)
def forward(self, input_: torch.Tensor):
# duplicate the logic in ColumnParallelLinear
lora_active = self.lora_active
if lora_active:
self.start_lora_a_overlap(input_)
bias = self.base_layer.bias if not self.base_layer.skip_bias_add else None
output_parallel = self.base_layer.quant_method.apply(
self.base_layer, input_, bias
)
if self.lora_active:
if lora_active:
output_parallel = self.apply_lora(output_parallel, input_)
if self.base_layer.gather_output:
@@ -596,6 +612,12 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
)
return lora_output
def start_lora_a_overlap(self, x: torch.Tensor) -> None:
if self.lora_backend.supports_lora_a_overlap:
self.lora_backend.start_lora_a_overlap(
x, self.A_buffer, num_slices=self._get_lora_n_slices()
)
def slice_lora_a_weights(self, A: torch.Tensor):
return A
@@ -703,6 +725,10 @@ class QKVParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
return lora_output
def start_lora_a_overlap(self, x: torch.Tensor) -> None:
if self.lora_backend.supports_lora_a_overlap:
self.lora_backend.start_lora_a_overlap(x, self.A_buffer_qkv, num_slices=3)
def slice_lora_a_weights(self, A: torch.Tensor):
return A
@@ -773,6 +799,10 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
)
return lora_output
def start_lora_a_overlap(self, x: torch.Tensor) -> None:
if self.lora_backend.supports_lora_a_overlap:
self.lora_backend.start_lora_a_overlap(x, self.A_buffer)
def forward(self, input_: torch.Tensor, skip_all_reduce=False, forward_batch=None):
if self.base_layer.input_is_parallel:
input_parallel = input_
@@ -783,6 +813,10 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
)
input_parallel = splitted_input[tp_rank].contiguous()
lora_active = self.lora_active
if lora_active:
self.start_lora_a_overlap(input_parallel)
bias_ = (
None
if (self.base_layer.tp_rank > 0 or self.base_layer.skip_bias_add)
@@ -806,7 +840,6 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
all_reduce = get_parallel().attn_tp_group.all_reduce
else:
all_reduce = tensor_model_parallel_all_reduce
lora_active = self.lora_active
if lora_active and should_reduce:
lora_a_output = self.lora_backend.run_lora_a_sgemm(
input_parallel, self.A_buffer
+48 -1
View File
@@ -17,7 +17,7 @@
import logging
import re
from typing import Dict, Iterable, List, Optional
from typing import Dict, Iterable, List, Optional, Sequence
import torch
@@ -431,6 +431,16 @@ class LoRAManager:
self.lora_backend.reset_batch_state()
def prepare_lora_batch(self, forward_batch: ForwardBatch):
# Some internal-only backends (currently UNO) use explicit token-row
# routing for their adapted forwards and want all-base batches to run
# through the plain model path. Clear any routing retained by the
# preceding adapted forward before inspecting CUDA-graph metadata.
if self.lora_backend.skip_inactive_lora_batches and not any(
uid is not None for uid in forward_batch.lora_ids
):
self.reset_lora_batch()
return
# set up batch info shared by all lora modules
bs = forward_batch.batch_size
@@ -470,6 +480,39 @@ class LoRAManager:
lora_ranks[wi] > 0 for wi in weight_indices
)
def prepare_lora_token_segments(
self,
*,
lora_ids: Sequence[Optional[str]],
segment_lens: Sequence[int],
) -> None:
"""Prepare eager LoRA routing independently of request batching."""
lora_ids = list(lora_ids)
segment_lens = list(segment_lens)
if len(lora_ids) != len(segment_lens):
raise ValueError("LoRA ids and segment lengths must have equal length.")
weight_indices = []
lora_ranks = [0] * self.max_loras_per_batch
scalings = [0.0] * self.max_loras_per_batch
for lora_id in lora_ids:
weight_index = self.memory_pool.get_buffer_id(lora_id)
weight_indices.append(weight_index)
if lora_id is not None:
lora = self.loras[lora_id]
lora_ranks[weight_index] = lora.config.r
scalings[weight_index] = lora.scaling
self.lora_backend.prepare_lora_token_segments(
segment_lens=segment_lens,
weight_indices=weight_indices,
lora_ranks=lora_ranks,
scalings=scalings,
)
self.lora_backend.batch_info.has_active_lora = any(
lora_ranks[index] > 0 for index in weight_indices
)
def update_lora_info(self):
"""
Update all LoRA modules to associate them with the latest memory buffer.
@@ -579,6 +622,10 @@ class LoRAManager:
max_lora_rank=max_lora_rank,
target_modules=target_modules,
)
self.lora_backend.validate_lora_targets(
base_model=self.base_model,
target_modules=self.target_modules,
)
if self._experts_shared_outer_override is not None:
self.experts_shared_outer_loras = self._experts_shared_outer_override
+10 -1
View File
@@ -308,6 +308,7 @@ from sglang.srt.speculative.eagle_utils import (
get_draft_recurrent_hidden_state_spec_from_config,
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.uno_validation import validate_uno_request
from sglang.srt.utils import (
DynamicGradMode,
configure_gc_logger,
@@ -2813,6 +2814,14 @@ class Scheduler(
self._add_request_to_queue(req)
return
if self.spec_algorithm.is_uno():
error_msg = validate_uno_request(req)
if error_msg is not None:
req.set_finish_with_abort(error_msg)
self.init_req_max_new_tokens(req)
self._add_request_to_queue(req)
return
if (
req.return_sampling_mask
and self.disaggregation_mode != DisaggregationMode.NULL
@@ -4484,7 +4493,7 @@ class Scheduler(
self.decode_moment_totals,
batch_size,
step_us,
batch_size + result.num_correct_drafts,
result.get_num_generated_tokens(batch_size),
)
def maybe_send_health_check_signal(self):
@@ -78,6 +78,16 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
def _get_speculative_output_stride(result: GenerationBatchResult) -> int:
"""Return the padded per-request width in flattened speculative output."""
stride = result.speculative_output_stride
if stride is None:
stride = result.speculative_num_draft_tokens
if stride is None or stride < 1:
raise RuntimeError("speculative result is missing a positive output row stride")
return stride
@dataclass(kw_only=True, slots=True, frozen=True)
class SchedulerBatchResultProcessor:
is_generation: bool
@@ -711,8 +721,12 @@ class SchedulerBatchResultProcessor:
next_token_ids = result.next_token_ids.tolist()
accept_lens = result.accept_lens.tolist()
result.num_correct_drafts = sum(accept_lens) - len(batch.reqs)
result.num_correct_drafts_per_req_cpu = [x - 1 for x in accept_lens]
stride = _get_speculative_output_stride(result)
num_non_draft = result.num_non_draft_tokens_per_req
result.num_correct_drafts_per_req_cpu = [
length - num_non_draft for length in accept_lens
]
result.num_correct_drafts = sum(result.num_correct_drafts_per_req_cpu)
block_accept_lens = (
result.block_accept_lens.tolist()
@@ -740,11 +754,6 @@ class SchedulerBatchResultProcessor:
self.advance_grammar_fsm(result, batch)
predict_tokens = []
# In adaptive spec-v2, the worker state may already have switched when this
# delayed result is processed. Use the draft token count recorded on result.
stride = result.speculative_num_draft_tokens
assert stride is not None, "spec-v2 result missing speculative_num_draft_tokens"
for i, req in enumerate(batch.reqs):
accept_tokens = next_token_ids[i * stride : i * stride + accept_lens[i]]
@@ -853,8 +862,7 @@ class SchedulerBatchResultProcessor:
if result.accept_lens is None:
return
accept_lens = result.accept_lens.tolist()
stride = result.speculative_num_draft_tokens
assert stride is not None, "spec-v2 result missing speculative_num_draft_tokens"
stride = _get_speculative_output_stride(result)
retained = [None] * len(batch.reqs)
for i, req in enumerate(batch.reqs):
if req.grammar is None or req.is_retracted or req.finished():
@@ -907,11 +915,14 @@ class SchedulerBatchResultProcessor:
next_token_ids=next_token_ids,
)
self.metrics_reporter.num_generated_tokens += len(batch.reqs)
batch_size = batch.batch_size()
num_generated_tokens = result.get_num_generated_tokens(batch_size)
self.metrics_reporter.num_generated_tokens += num_generated_tokens
if not batch.spec_algorithm.is_none():
self.metrics_reporter.update_spec_metrics(
batch.batch_size(),
batch_size,
result.num_correct_drafts,
num_accept_tokens=num_generated_tokens,
num_block_accept_tokens=result.num_block_accept_tokens,
num_cap_tokens=result.num_cap_tokens,
)
@@ -982,8 +993,8 @@ class SchedulerBatchResultProcessor:
if req.return_hidden_states and logits_output.hidden_states is not None:
# hidden_states is [bs * stride, hidden_dim], one row per emitted
# token; stride = speculative_num_draft_tokens for spec, 1 for non-spec.
stride = result.speculative_num_draft_tokens or 1
# token; speculative workers record their padded row width.
stride = _get_speculative_output_stride(result) if is_spec else 1
accept_len = len(next_token_id)
start = i * stride
self._append_decode_hidden_states(
@@ -1016,7 +1027,7 @@ class SchedulerBatchResultProcessor:
self.metrics_reporter.report_decode_stats(
can_run_cuda_graph,
running_batch=batch,
num_correct_drafts=result.num_correct_drafts,
num_generated_tokens=num_generated_tokens,
)
def _normalize_decode_outputs(
@@ -175,9 +175,10 @@ class SchedulerMetricsReporter:
}.get(getattr(self.scheduler, "device", ""), "cuda graph")
# Cumulative spec-decoding counters (reset every decode_log_interval).
# Each update adds (num_correct_drafts + bs, bs).
# `*_accept_tokens` = drafts + bonus; `*_correct_drafts` = drafts-only.
# `*_accept_tokens` includes accepted drafts and non-draft output tokens;
# `*_correct_drafts` counts accepted draft proposals only.
self.spec_num_accept_tokens = 0 # per-log-interval
self.spec_num_correct_drafts = 0
self.spec_num_forward_ct = 0
self.spec_total_num_accept_tokens = 0 # lifetime
self.spec_total_num_forward_ct = 0
@@ -398,17 +399,16 @@ class SchedulerMetricsReporter:
self,
bs: int,
num_correct_drafts: int,
num_accept_tokens: int,
num_block_accept_tokens: int = 0,
num_cap_tokens: int = 0,
):
self.spec_num_accept_tokens += num_correct_drafts + bs
self.spec_num_accept_tokens += num_accept_tokens
self.spec_num_correct_drafts += num_correct_drafts
self.spec_num_forward_ct += bs
self.spec_num_block_accept_tokens += num_block_accept_tokens
self.spec_num_cap_tokens += num_cap_tokens
# Bonus tokens updated elsewhere
self.num_generated_tokens += num_correct_drafts
def _init_estimated_perf_constants(self) -> None:
model_config = self.scheduler.model_config
hf_text_config = model_config.hf_text_config
@@ -572,6 +572,7 @@ class SchedulerMetricsReporter:
self.forward_ct_decode = 0
self.num_generated_tokens = 0
self.spec_num_accept_tokens = 0
self.spec_num_correct_drafts = 0
self.spec_num_forward_ct = 0
self.spec_total_num_accept_tokens = 0
self.spec_total_num_forward_ct = 0
@@ -757,13 +758,13 @@ class SchedulerMetricsReporter:
self,
can_run_cuda_graph: bool,
running_batch: ScheduleBatch = None,
num_correct_drafts: int = 0,
num_generated_tokens: int = 0,
):
batch = running_batch or self.scheduler.running_batch
# Every-iteration work: realtime token counting + status logger
if self.current_scheduler_metrics_enabled:
decode_tokens = batch.batch_size() + num_correct_drafts
decode_tokens = num_generated_tokens
self.metrics_collector.increment_realtime_tokens(
# TODO unify this w/ the bumping logic in `Scheduler.num_generated_tokens` accumulator
decode_tokens=decode_tokens,
@@ -826,7 +827,7 @@ class SchedulerMetricsReporter:
spec_block_accept_length = 0
else:
spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct
num_correct_drafts = self.spec_num_accept_tokens - self.spec_num_forward_ct
num_correct_drafts = self.spec_num_correct_drafts
if get_spec().speculative_num_draft_tokens:
draft_per_round = get_spec().speculative_num_draft_tokens - 1
else:
@@ -853,7 +854,8 @@ class SchedulerMetricsReporter:
)
self.spec_total_num_accept_tokens += self.spec_num_accept_tokens
self.spec_total_num_forward_ct += self.spec_num_forward_ct
self.spec_num_accept_tokens = self.spec_num_forward_ct = 0
self.spec_num_accept_tokens = self.spec_num_correct_drafts = 0
self.spec_num_forward_ct = 0
self.spec_num_block_accept_tokens = 0
self.spec_num_cap_tokens = 0
msg += f"accept len: {spec_accept_length:.2f}, accept rate: {spec_accept_rate:.2f}, "
+11 -2
View File
@@ -22,7 +22,7 @@ from sglang.srt.state_capturer.base import TopkCaptureOutput
if TYPE_CHECKING:
from sglang.srt.managers.scheduler import GenerationBatchResult
from sglang.srt.sampling.sampling_observer import HostAuxiliaryOutput
from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.speculative.spec_info import SpecInput
logger = logging.getLogger(__name__)
@@ -71,6 +71,12 @@ class GenerationBatchResult:
delay_sample_func: Optional[callable] = None
future_indices: Optional[torch.Tensor] = None
speculative_num_draft_tokens: Optional[int] = None
# Padded row width in flattened speculative output. Existing algorithms
# default to speculative_num_draft_tokens; linear UNO emits F + 1 columns.
speculative_output_stride: Optional[int] = None
# Valid output tokens that are not accepted draft proposals. Existing
# algorithms have one bonus token; UNO also emits its clean root.
num_non_draft_tokens_per_req: int = 1
# Grammar FSM advance memoization (spec-v2 overlap). advance_grammar_fsm sets
# these once — eagerly via the scheduler's grammar barrier inside verify(), or
@@ -91,7 +97,7 @@ class GenerationBatchResult:
new_seq_lens: Optional[torch.Tensor] = None
# relay path: forward stream -> next step forward
next_draft_input: Optional[EagleDraftInput] = None
next_draft_input: Optional[SpecInput] = None
# Refs the worker wants scheduler to keep alive for the same 2-iter window
# as batch_record_buf. Used for cross-stream tensor lifetime (e.g. a spec
@@ -117,6 +123,9 @@ class GenerationBatchResult:
this rank/split (a non-last PP rank or a non-final prefill split)."""
return isinstance(self.next_token_ids, torch.Tensor)
def get_num_generated_tokens(self, batch_size: int) -> int:
return self.num_correct_drafts + batch_size * self.num_non_draft_tokens_per_req
@torch.profiler.record_function("copy_result_to_cpu")
def copy_to_cpu(self, return_logprob: bool, return_hidden_states: bool = True):
"""Copy tensors to CPU in overlap scheduling.
@@ -33,6 +33,11 @@ def get_alloc_len_per_decode() -> int:
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
spec_algo = SpeculativeAlgorithm.from_string(spec.speculative_algorithm)
if spec_algo.is_uno():
if spec_tokens is None:
raise RuntimeError("UNO requires speculative_num_draft_tokens")
# UNO retains an additional clean-root position beside Q/F draft slots.
return spec_tokens + 1
if page_size == 1 or spec_topk == 1 or not spec_algo.has_draft_kv():
return max(spec_steps * spec_topk, spec_tokens)
else:
@@ -88,9 +93,13 @@ def get_req_to_token_extra_context_len() -> int:
# FIXME(lsyin): temporary fix for the context length issue under spec decoding
extra = 4 + (max_speculative_num_draft_tokens() or 0)
page_size = get_alloc_page_size()
if get_spec().speculative_algorithm is not None and page_size > 1:
# kv_allocated_len is page-aligned (eagle_prepare_for_decode), so near
# the context limit the aligned reserve can overshoot by page_size - 1;
# without the headroom the row write silently lands in the neighbor row.
extra = max(extra, get_alloc_reserve_per_decode() + page_size - 1)
spec_algorithm = get_spec().speculative_algorithm
if spec_algorithm is not None:
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
spec_algo = SpeculativeAlgorithm.from_string(spec_algorithm)
if page_size > 1 or spec_algo.is_uno():
# UNO's double-buffer reserve applies at every page size. Larger
# pages may additionally round the allocation up by page_size - 1.
extra = max(extra, get_alloc_reserve_per_decode() + page_size - 1)
return extra
@@ -777,6 +777,12 @@ class ModelRunner:
self.apply_torch_tp()
def maybe_init_lora_manager(self):
if self.spec_algorithm.is_uno():
from sglang.srt.speculative.uno_lora import init_uno_lora_manager
self.lora_manager, self.uno_lora_id = init_uno_lora_manager(self)
return
# Adapters apply to the target model only; the draft runs unadapted.
if get_lora().enable_lora and not self.is_draft_worker:
self.init_lora_manager()
@@ -1430,6 +1436,13 @@ class ModelRunner:
Subclasses can override this to install specialized decode graph runners.
"""
if self.spec_algorithm.is_uno():
from sglang.srt.speculative.uno_cuda_graph_runner import (
UnoDecodeCudaGraphRunner,
)
return UnoDecodeCudaGraphRunner
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
DecodeCudaGraphRunner,
)
+6 -1
View File
@@ -2067,7 +2067,12 @@ class ServerArgs:
# -------------------------------------------------------------------------
speculative_algorithm: A[
Optional[str],
"Speculative algorithm. Builtins: EAGLE, EAGLE3, NEXTN, STANDALONE, NGRAM, DFLASH, DSPARK. Or any name registered via `SpeculativeAlgorithm.register`.",
"Speculative algorithm. Builtins: EAGLE, EAGLE3, NEXTN, STANDALONE, NGRAM, DFLASH, DSPARK, UNO. Or any name registered via `SpeculativeAlgorithm.register`.",
NS("spec"),
] = None
uno_lora_path: A[
Optional[str],
"Path to the UNO draft LoRA checkpoint.",
NS("spec"),
] = None
speculative_draft_model_path: A[
@@ -663,11 +663,30 @@ def _verify_coins(
return coins, coins_for_final_sampling
def _can_use_sparse_uno_tree_target_sampling(
max_top_k: Optional[int],
sampling_info: SamplingBatchInfo,
) -> bool:
if max_top_k is None:
return False
from sglang.srt.speculative.uno_utils import _SPARSE_TOP_K_LIMIT
return bool(
_is_cuda
and max_top_k <= _SPARSE_TOP_K_LIMIT
and sampling_info.sampling_seed is None
and not sampling_info.need_min_p_sampling
and not get_spec().speculative_use_rejection_sampling
)
def eagle_sample(
verify_input: EagleVerifyInput,
batch: ScheduleBatch,
logits_output: LogitsProcessorOutput,
grammar_mask: Optional[GrammarMask] = None,
uno_target_max_top_k: Optional[int] = None,
):
"""
Verify and find accepted tokens based on logits output and batch
@@ -769,6 +788,42 @@ def eagle_sample(
tp_group.broadcast(predict, src=0)
tp_group.broadcast(accept_index, src=0)
tp_group.broadcast(num_correct_drafts, src=0)
elif _can_use_sparse_uno_tree_target_sampling(
uno_target_max_top_k,
sampling_info,
):
from sglang.srt.speculative.uno_utils import (
sample_uno_tree_target_tokens,
)
target_predict = sample_uno_tree_target_tokens(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
batch_size=bs,
verify_width=verify_input.draft_token_num,
max_top_k=uno_target_max_top_k,
)
predict, accept_index, num_correct_drafts = verify_tree_greedy_func(
predicts=predict,
accept_index=accept_index,
accept_token_num=num_correct_drafts,
candidates=candidates,
retrieve_index=verify_input.retrieve_index,
retrieve_next_token=verify_input.retrieve_next_token,
retrieve_next_sibling=verify_input.retrieve_next_sibling,
target_predict=target_predict,
topk=verify_input.tree_topk,
)
tp_group = (
get_parallel().attn_tp_group
if is_dp_attention_enabled()
else get_tp_group()
)
if tp_group.world_size > 1:
tp_group.broadcast(predict, src=0)
tp_group.broadcast(accept_index, src=0)
tp_group.broadcast(num_correct_drafts, src=0)
else:
from sgl_kernel import (
top_k_renorm_prob,
@@ -472,6 +472,7 @@ def run_eagle_verify(
metadata_ready_pre_pad: bool,
finalize_tree_path: bool,
grammar_barrier=None,
uno_target_max_top_k: Optional[int] = None,
) -> GenerationBatchResult:
"""Shared verify step: target-verify forward, sampling, acceptance bookkeeping.
@@ -584,7 +585,13 @@ def run_eagle_verify(
predict,
accept_lens,
accept_index,
) = eagle_sample(verify_input, batch, logits_output, grammar_mask)
) = eagle_sample(
verify_input,
batch,
logits_output,
grammar_mask,
uno_target_max_top_k=uno_target_max_top_k,
)
new_seq_lens = batch.seq_lens + accept_lens
clear_unaccepted_c128 = getattr(
token_to_kv_pool_allocator.get_kvcache(),
@@ -37,6 +37,7 @@ class SpeculativeAlgorithm(Enum):
"""
DFLASH = auto()
UNO = auto()
DSPARK = auto()
EAGLE = auto()
EAGLE3 = auto()
@@ -114,6 +115,9 @@ class SpeculativeAlgorithm(Enum):
def is_dflash(self) -> bool:
return self == SpeculativeAlgorithm.DFLASH
def is_uno(self) -> bool:
return self == SpeculativeAlgorithm.UNO
def is_dspark(self) -> bool:
return self == SpeculativeAlgorithm.DSPARK
@@ -220,6 +224,7 @@ class SpeculativeAlgorithm(Enum):
_handle_eagle_family,
_handle_frozen_kv_mtp,
_handle_ngram,
_handle_uno,
)
# Validate for every algorithm at startup: the metrics paths read the
@@ -230,6 +235,8 @@ class SpeculativeAlgorithm(Enum):
if self.is_dflash():
_handle_dflash(server_args)
elif self.is_uno():
_handle_uno(server_args)
elif self.is_dspark():
_handle_dspark(server_args)
elif self.is_frozen_kv_mtp():
@@ -304,6 +311,11 @@ class SpeculativeAlgorithm(Enum):
return DFlashWorkerV2
if self.is_uno():
from sglang.srt.speculative.uno_worker_v2 import UnoWorkerV2
return UnoWorkerV2
if self.is_dspark():
from sglang.srt.speculative.dspark_components.dspark_worker_v2 import (
DSparkWorkerV2,
@@ -356,6 +368,9 @@ class SpecInputType(IntEnum):
DFLASH_DRAFT = auto()
DFLASH_VERIFY = auto()
NGRAM_VERIFY = auto()
UNO_STATE = auto()
UNO_DRAFT = auto()
UNO_VERIFY = auto()
class SpecInput(ABC):
@@ -393,6 +408,7 @@ class SpecInput(ABC):
SpecInputType.EAGLE_DRAFT_EXTEND,
SpecInputType.FROZEN_KV_MTP_DRAFT,
SpecInputType.DFLASH_DRAFT,
SpecInputType.UNO_DRAFT,
}
def is_verify_input(self) -> bool:
@@ -401,6 +417,7 @@ class SpecInput(ABC):
SpecInputType.FROZEN_KV_MTP_VERIFY,
SpecInputType.DFLASH_VERIFY,
SpecInputType.NGRAM_VERIFY,
SpecInputType.UNO_VERIFY,
}
@@ -82,6 +82,9 @@ class CustomSpecAlgo:
def is_dflash(self) -> bool:
return False
def is_uno(self) -> bool:
return False
def is_dspark(self) -> bool:
return False
@@ -1051,6 +1051,12 @@ def spec_prepare_for_decode(batch: ScheduleBatch) -> None:
)
if batch.spec_algorithm.is_dflash_family():
batch.spec_info.prepare_for_decode(batch)
elif batch.spec_algorithm.is_uno():
from sglang.srt.speculative.uno_info import UnoDraftInput
if not isinstance(batch.spec_info, UnoDraftInput):
raise RuntimeError("UNO decode preparation requires UnoDraftInput")
batch.spec_info.prepare_for_decode(batch)
else:
from sglang.srt.speculative.eagle_utils import eagle_prepare_for_decode
@@ -0,0 +1,214 @@
import torch
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
DecodeCudaGraphRunner,
)
from sglang.srt.runtime_context import get_spec
from sglang.srt.speculative.eagle_info import EagleVerifyInput
from sglang.srt.speculative.spec_info import SpecInputType
from sglang.srt.speculative.uno_info import UnoForwardInput
from sglang.srt.speculative.uno_lora import UnoCudaGraphLoRAState
class UnoDecodeCudaGraphRunner(DecodeCudaGraphRunner):
"""Decode graph runner for linear UNO and both tree forward roles.
Linear UNO uses two-variant F-wide capture. Tree UNO uses
separate runner instances for its F-wide LoRA draft
and native Q/K EAGLE target verification.
"""
def __init__(
self,
model_runner,
*,
tree_draft_attn_backend=None,
tree_draft_width=None,
**kwargs,
):
candidate_top_k = get_spec().speculative_eagle_topk
self._tree_mode = candidate_top_k > 1
self._tree_draft_mode = tree_draft_width is not None
if self._tree_draft_mode:
self.record_nolora_graph = False
self._capture_spec_input_type = SpecInputType.UNO_DRAFT
self._lora_state = UnoCudaGraphLoRAState(
model_runner.lora_manager,
model_runner.uno_lora_id,
tree_draft_width,
)
model_runner.lora_manager.reset_lora_batch()
kwargs.update(
attn_backend=tree_draft_attn_backend,
speculative_num_steps=1,
speculative_num_draft_tokens=tree_draft_width,
)
super().__init__(model_runner, **kwargs)
model_runner.lora_manager.reset_lora_batch()
return
if self._tree_mode:
# Capture exactly one base-model target graph. The internal UNO
# adapter is active only in the rejected F-wide draft phase.
self.record_nolora_graph = False
model_runner.lora_manager.reset_lora_batch()
kwargs.update(
attn_backend=model_runner.attn_backend,
speculative_num_steps=get_spec().speculative_num_steps,
speculative_num_draft_tokens=get_spec().speculative_num_draft_tokens,
)
super().__init__(model_runner, **kwargs)
model_runner.lora_manager.reset_lora_batch()
return
forward_width = model_runner.decode_num_tokens_per_req()
self.record_nolora_graph = forward_width > 1
self._capture_spec_input_type = SpecInputType.UNO_VERIFY
self._lora_state = UnoCudaGraphLoRAState(
model_runner.lora_manager,
model_runner.uno_lora_id,
forward_width,
)
model_runner.lora_manager.reset_lora_batch()
super().__init__(model_runner, **kwargs)
def capture_prepare(self, size, stream_idx=None, num_tokens=None):
forward_batch, attn_backend, pp_proxy_tensors = super().capture_prepare(
size, stream_idx=stream_idx, num_tokens=num_tokens
)
# UNO owns token-row routing directly. K2's generic graph runner now
# keys LoRA setup off lora_manager presence, so suppress its synthetic
# request-level base-adapter routing during UNO capture.
forward_batch.lora_ids = None
return forward_batch, attn_backend, pp_proxy_tensors
def can_run_graph(self, forward_batch):
spec_info = forward_batch.spec_info
if self._tree_draft_mode:
if not isinstance(spec_info, UnoForwardInput):
return False
if spec_info.spec_input_type != SpecInputType.UNO_DRAFT:
return False
return super().can_run_graph(forward_batch)
if self._tree_mode:
if not isinstance(spec_info, EagleVerifyInput):
return False
return super().can_run_graph(forward_batch)
if not isinstance(spec_info, UnoForwardInput):
return False
# At F=1 both phases are base-only and share variant_label=None.
if spec_info.spec_input_type not in {
SpecInputType.UNO_DRAFT,
SpecInputType.UNO_VERIFY,
}:
return False
return super().can_run_graph(forward_batch)
def _resolve_lora_variant(self, forward_batch):
"""borrowed technique from multi-LoRA serving
to capture separate graph for each step."""
if self._tree_mode:
# Tree verification is always the clean target model. Keeping the
# graph keys unlabeled is sufficient because draft and verify own
# separate runners.
return None
if not self.record_nolora_graph:
return None
if forward_batch.spec_info.spec_input_type == SpecInputType.UNO_DRAFT:
return "lora"
return "nolora"
def capture_one_shape(
self,
size,
forward,
stream_idx=None,
variant_label=None,
dsa_variant=None,
):
"""capture one CUDA graph with/out UNO LoRA."""
if self._tree_draft_mode:
self._lora_state.capture_draft(size)
try:
return super().capture_one_shape(
size,
forward,
stream_idx,
None,
dsa_variant,
)
finally:
self._lora_state.reset()
if self._tree_mode:
self.model_runner.lora_manager.reset_lora_batch()
try:
return super().capture_one_shape(
size,
forward,
stream_idx,
None,
dsa_variant,
)
finally:
self.model_runner.lora_manager.reset_lora_batch()
if variant_label == "lora":
self._capture_spec_input_type = SpecInputType.UNO_DRAFT
self._lora_state.capture_draft(size)
else:
self._capture_spec_input_type = SpecInputType.UNO_VERIFY
self._lora_state.reset()
super().capture_one_shape(
size,
forward,
stream_idx,
variant_label,
dsa_variant,
)
self._lora_state.reset()
def get_spec_info(self, num_tokens: int):
if self._tree_draft_mode:
return UnoForwardInput(
spec_input_type=SpecInputType.UNO_DRAFT,
positions=self.buffers.positions[:num_tokens],
draft_token_num=self.captured_req_width,
)
if self._tree_mode:
# This deliberately mirrors DecodeCudaGraphRunner's EAGLE capture
# input. Current eagle_prepare_for_verify requests FULL hidden
# capture for every non-STANDALONE algorithm, including UNO.
spec_info = EagleVerifyInput(
draft_token=None,
custom_mask=self.buffers.custom_mask,
positions=None,
retrieve_index=None,
retrieve_next_token=None,
retrieve_next_sibling=None,
retrieve_cum_len=None,
spec_steps=self.speculative_num_steps,
topk=get_spec().speculative_eagle_topk,
draft_token_num=self.speculative_num_draft_tokens,
capture_hidden_mode=CaptureHiddenMode.FULL,
seq_lens_sum=None,
seq_lens_cpu=None,
)
spec_info.hidden_states = torch.zeros(
(num_tokens, self.model_runner.model_config.hidden_size),
dtype=self.model_runner.dtype,
device=self.model_runner.device,
)
return spec_info
return UnoForwardInput(
spec_input_type=self._capture_spec_input_type,
positions=self.buffers.positions[:num_tokens],
draft_token_num=self.captured_req_width,
)
+276
View File
@@ -0,0 +1,276 @@
from dataclasses import dataclass, field
from typing import ClassVar, List, Optional
import torch
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.mem_cache.allocation import alloc_for_spec_decode
from sglang.srt.mem_cache.allocation_sizing import (
get_alloc_reserve_per_decode,
page_aligned_decode_alloc_lens,
)
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
@dataclass
class UnoDraftInput(SpecInput):
"""UNO state carried from one engine iteration to the next."""
# Previously emitted token at logical position C. It has no KV yet.
bonus_tokens: torch.Tensor
# Target-correct KV frontier C.
new_seq_lens: torch.Tensor
# Number of queries in each UNO forward.
forward_width: int
# FutureMap compatibility. UNO relays only bonus_tokens, but the generic
# relay payload reads these optional Eagle-shaped fields.
topk_p: ClassVar[Optional[torch.Tensor]] = None
topk_index: ClassVar[Optional[torch.Tensor]] = None
hidden_states: ClassVar[Optional[torch.Tensor]] = None
# Filled by the scheduler after an overlapped dispatch.
future_indices: Optional[torch.Tensor] = None
# Host-side upper bound for the allocator mapping prepared for this step.
reserved_seq_lens_cpu: Optional[torch.Tensor] = None
reserved_seq_lens_sum: Optional[int] = None
# Sampling metadata prepared before either internal forward.
max_top_k: int = 1
uniform_top_k_value: Optional[int] = None
def __post_init__(self):
super().__init__(SpecInputType.UNO_STATE)
if self.forward_width < 1:
raise ValueError("UNO forward_width must be positive.")
# The carried state represents one request row, not an F-row forward.
self.num_tokens_per_req = 1
self.num_tokens_for_logprob_per_req = 1
@property
def tail_width(self) -> int:
return self.forward_width + 1
@classmethod
def create_idle_input(
cls,
*,
device,
forward_width: int,
) -> "UnoDraftInput":
return cls(
bonus_tokens=torch.empty(
(0,),
dtype=torch.int64,
device=device,
),
new_seq_lens=torch.empty(
(0,),
dtype=torch.int64,
device=device,
),
forward_width=forward_width,
)
def prepare_for_decode(self, batch: ScheduleBatch) -> None:
batch.maybe_evict_swa()
batch_size = batch.batch_size()
if batch_size == 0:
return
if self.future_indices is None:
if self.bonus_tokens.numel() != batch_size:
raise RuntimeError("UNO seed count does not match decode batch size.")
if self.new_seq_lens.numel() != batch_size:
raise RuntimeError(
"UNO frontier count does not match decode batch size."
)
elif self.future_indices.numel() != batch_size:
raise RuntimeError(
"UNO future-index count does not match decode batch size."
)
committed_lengths: list[int] = []
reserve_width = int(get_alloc_reserve_per_decode())
max_top_k = 1
uniform_top_k_value = None
uniform_top_k = True
for index, req in enumerate(batch.reqs):
if req.kv is None:
raise RuntimeError("UNO decode request has no KV allocation.")
committed = int(req.kv.kv_committed_len)
allocated = int(req.kv.kv_allocated_len)
if allocated < committed:
raise RuntimeError(
"UNO encountered an invalid KV watermark: "
f"committed={committed}, allocated={allocated}."
)
committed_lengths.append(committed)
top_k = int(req.sampling_params.top_k)
max_top_k = max(max_top_k, top_k)
if index == 0:
uniform_top_k_value = top_k
elif uniform_top_k and top_k != uniform_top_k_value:
uniform_top_k = False
self.max_top_k = max_top_k
self.uniform_top_k_value = uniform_top_k_value if uniform_top_k else None
page_size = batch.token_to_kv_pool_allocator.page_size
current_lengths, next_lengths, num_needed_tokens = (
page_aligned_decode_alloc_lens(
batch.reqs,
reserve=reserve_width,
page_size=page_size,
)
)
row_width = int(batch.req_to_token_pool.req_to_token.shape[1])
if max(next_lengths) > row_width:
raise RuntimeError(
"UNO allocation exceeds the req_to_token row: "
f"needed={max(next_lengths)}, available={row_width}."
)
current_cpu = torch.tensor(
current_lengths,
dtype=torch.int32,
device="cpu",
)
next_cpu = torch.tensor(
next_lengths,
dtype=torch.int32,
device="cpu",
)
current_device = current_cpu.to(
batch.device,
non_blocking=True,
)
next_device = next_cpu.to(
batch.device,
non_blocking=True,
)
alloc_for_spec_decode(
batch.tree_cache,
batch.req_to_token_pool,
reqs=batch.reqs,
req_pool_indices=batch.req_pool_indices,
cur_kv_lens=current_device,
cur_kv_lens_cpu=current_cpu,
nxt_kv_lens=next_device,
nxt_kv_lens_cpu=next_cpu,
num_needed_tokens=num_needed_tokens,
batch=batch,
)
for req in batch.reqs:
req.decode_batch_idx += 1
batch.seq_lens_cpu = torch.tensor(
committed_lengths,
dtype=torch.int64,
device="cpu",
)
batch.seq_lens_sum = sum(committed_lengths)
self.reserved_seq_lens_cpu = next_cpu
self.reserved_seq_lens_sum = sum(next_lengths)
def filter_batch(
self,
new_indices: torch.Tensor,
new_indices_cpu: Optional[List[int]] = None,
) -> None:
if self.reserved_seq_lens_cpu is not None:
host_indices = (
new_indices_cpu if new_indices_cpu is not None else new_indices.cpu()
)
self.reserved_seq_lens_cpu = self.reserved_seq_lens_cpu[host_indices]
self.reserved_seq_lens_sum = int(self.reserved_seq_lens_cpu.sum().item())
if self.future_indices is not None:
self.future_indices = self.future_indices[new_indices]
return
self.bonus_tokens = self.bonus_tokens[new_indices]
self.new_seq_lens = self.new_seq_lens[new_indices]
def merge_batch(self, other: "UnoDraftInput") -> None:
if not isinstance(other, UnoDraftInput):
raise TypeError(f"Cannot merge UnoDraftInput with {type(other).__name__}.")
if self.forward_width != other.forward_width:
raise RuntimeError("Cannot merge UNO states with different forward widths.")
self_has_reservation = self.reserved_seq_lens_cpu is not None
other_has_reservation = other.reserved_seq_lens_cpu is not None
if self_has_reservation != other_has_reservation:
raise RuntimeError("Cannot merge prepared and unprepared UNO states.")
if self_has_reservation:
self.reserved_seq_lens_cpu = torch.cat(
(
self.reserved_seq_lens_cpu,
other.reserved_seq_lens_cpu,
)
)
self.reserved_seq_lens_sum = int(self.reserved_seq_lens_cpu.sum().item())
if self.future_indices is not None:
assert other.future_indices is not None
self.future_indices = torch.cat((self.future_indices, other.future_indices))
return
self.bonus_tokens = torch.cat((self.bonus_tokens, other.bonus_tokens))
self.new_seq_lens = torch.cat((self.new_seq_lens, other.new_seq_lens))
@dataclass
class UnoForwardInput(SpecInput):
"""Metadata for one fixed-width UNO forward."""
# constructor
spec_input_type: SpecInputType
positions: torch.Tensor
# For UNO, this is forward width F, not proposal count F - 1.
draft_token_num: int
# expected by interface
custom_mask: Optional[torch.Tensor] = None
capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.NULL
hidden_states: Optional[torch.Tensor] = None
# derived
num_tokens_per_req: int = field(init=False)
num_tokens_for_logprob_per_req: int = field(init=False)
def __post_init__(self):
if self.spec_input_type not in {
SpecInputType.UNO_DRAFT,
SpecInputType.UNO_VERIFY,
}:
raise ValueError(f"Invalid UNO input type: {self.spec_input_type}")
if self.draft_token_num < 1:
raise ValueError("UNO forward width must be positive.")
# Dataclass-generated __init__ does not call the non-dataclass base
# initializer. This currently reassigns the same field, while also
# preserving the SpecInput initialization contract.
super().__init__(self.spec_input_type)
self.num_tokens_per_req = self.draft_token_num
self.num_tokens_for_logprob_per_req = self.draft_token_num
+97
View File
@@ -0,0 +1,97 @@
"""Loading for UNO's draft LoRA."""
from __future__ import annotations
from typing import TYPE_CHECKING
from sglang.srt.lora.lora_manager import LoRAManager
from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.runtime_context import get_spec
if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner
# This name is internal to the model-execution process. It is never exposed as
# a request-selectable serving adapter.
_UNO_INTERNAL_LORA_NAME = "__uno_draft__"
# LoRA pool capacity includes the base-model slot.
_UNO_LORA_POOL_CAPACITY = 2
def init_uno_lora_manager(
model_runner: ModelRunner,
) -> tuple[LoRAManager, str]:
"""Load and pin the single UNO draft adapter."""
lora_path = get_spec().uno_lora_path
uno_ref = LoRARef(
lora_id=LoRARef.deterministic_id(
_UNO_INTERNAL_LORA_NAME,
lora_path,
),
lora_name=_UNO_INTERNAL_LORA_NAME,
lora_path=lora_path,
pinned=True,
)
manager = LoRAManager(
base_model=model_runner.model,
base_hf_config=model_runner.model_config.hf_config,
max_loras_per_batch=_UNO_LORA_POOL_CAPACITY,
load_config=model_runner.load_config,
dtype=model_runner.dtype,
server_args=model_runner.server_args,
lora_backend="uno_cublas", # fast path
tp_size=model_runner.ps.tp_size,
tp_rank=model_runner.ps.tp_rank,
# Infer these from the one trained adapter.
max_lora_rank=None,
target_modules=None,
lora_paths=[uno_ref],
)
# LoRAManager construction initially makes only the base slot resident.
# UNO always needs both fixed choices resident:
#
# None -> base model
# uno_ref.lora_id -> base model + UNO draft LoRA
manager.fetch_new_loras({None, uno_ref.lora_id})
return manager, uno_ref.lora_id
class UnoCudaGraphLoRAState:
"""Retained token-row LoRA routing for UNO draft graph buckets."""
def __init__(
self,
manager: LoRAManager,
uno_lora_id: str,
forward_width: int,
):
self.manager = manager
self.uno_lora_id = uno_lora_id
self.forward_width = forward_width
self._draft_batch_infos = {}
def capture_draft(self, batch_size: int) -> None:
if batch_size not in self._draft_batch_infos:
self.manager.prepare_lora_token_segments(
lora_ids=[None, self.uno_lora_id] * batch_size,
segment_lens=[1, self.forward_width - 1] * batch_size,
)
batch_info = self.manager.lora_backend.batch_info
batch_info.use_cuda_graph = True
self._draft_batch_infos[batch_size] = batch_info
self.activate_draft(batch_size)
def activate_draft(self, batch_size: int) -> None:
self.manager.reset_lora_batch()
self.manager.lora_backend.batch_info = self._draft_batch_infos[batch_size]
def reset(self) -> None:
self.manager.reset_lora_batch()
+584
View File
@@ -0,0 +1,584 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from nano-vllm-uno's fixed-budget draft-tree builder for SGLang.
"""GPU-native UNO proposal-tree construction for SGLang spec-v2.
UNO owns only proposal ranking and fixed-budget best-first selection. The
result is expressed directly in EAGLE's candidate-lineage ABI so the existing
EAGLE implementation can build masks, positions, traversal links, verify the
tree, sample the accepted path, and compact KV state.
"""
from __future__ import annotations
from dataclasses import dataclass
import torch
import triton
import triton.language as tl
from flashinfer import top_k as _flashinfer_top_k
from torch import Tensor
@dataclass(frozen=True)
class UnoTreeProposal:
"""EAGLE-native representation of one fixed-width UNO tree batch.
For non-root node ``i``, ``top_scores_index[:, i - 1]`` is the implicit
candidate edge ``parent_node * candidate_top_k + candidate_rank``.
``parent_list`` maps an EAGLE candidate row back to the edge that created
that row's node. Its rows use EAGLE's native
``candidate_top_k * (max_depth - 1) + 1`` stride, including unused
padding. The tensors can therefore be passed directly to
``build_tree_kernel_efficient`` without constructing direct parent arrays.
Tensors may be backed by the caller's reusable workspace and remain valid
only until that workspace is reused.
"""
root_tokens: Tensor
draft_tokens: Tensor
parent_list: Tensor
top_scores_index: Tensor
candidate_top_k: int
max_depth: int
@property
def num_verify_tokens(self) -> int:
return int(self.draft_tokens.shape[1]) + 1
@triton.jit
def _candidate_lse_partials_kernel(
logits,
partial_max,
partial_sum,
temperature_values,
inverse_temperature,
BATCH_STRIDE: tl.constexpr,
DEPTH_STRIDE: tl.constexpr,
VOCAB_STRIDE: tl.constexpr,
NUM_DEPTHS: tl.constexpr,
VOCAB_SIZE: tl.constexpr,
NUM_BLOCKS: tl.constexpr,
BLOCK_VOCAB: tl.constexpr,
TEMPERATURE_BATCH_STRIDE: tl.constexpr,
TEMPERATURE_IS_TENSOR: tl.constexpr,
):
batch = tl.program_id(0)
depth = tl.program_id(1)
block = tl.program_id(2)
row = batch * NUM_DEPTHS + depth
offsets = block * BLOCK_VOCAB + tl.arange(0, BLOCK_VOCAB)
values = tl.load(
logits + batch * BATCH_STRIDE + depth * DEPTH_STRIDE + offsets * VOCAB_STRIDE,
mask=offsets < VOCAB_SIZE,
other=-float("inf"),
).to(tl.float32)
if TEMPERATURE_IS_TENSOR:
row_temperature = tl.load(
temperature_values + batch * TEMPERATURE_BATCH_STRIDE
).to(tl.float32)
row_inverse_temperature = tl.where(
row_temperature > 0.0,
1.0 / row_temperature,
1.0,
)
else:
row_inverse_temperature = inverse_temperature
values *= row_inverse_temperature
maximum = tl.max(values, axis=0)
total = tl.sum(tl.exp(values - maximum), axis=0)
output = row * NUM_BLOCKS + block
tl.store(partial_max + output, maximum)
tl.store(partial_sum + output, total)
@triton.jit
def _candidate_lse_finalize_kernel(
top_values,
partial_max,
partial_sum,
top_log_probs,
temperature_values,
inverse_temperature,
K: tl.constexpr,
NUM_DEPTHS: tl.constexpr,
NUM_BLOCKS: tl.constexpr,
BLOCK_K: tl.constexpr,
BLOCK_PARTIALS: tl.constexpr,
TEMPERATURE_BATCH_STRIDE: tl.constexpr,
TEMPERATURE_IS_TENSOR: tl.constexpr,
):
row = tl.program_id(0)
batch = row // NUM_DEPTHS
blocks = tl.arange(0, BLOCK_PARTIALS)
block_max = tl.load(
partial_max + row * NUM_BLOCKS + blocks,
mask=blocks < NUM_BLOCKS,
other=-float("inf"),
)
maximum = tl.max(block_max, axis=0)
block_sum = tl.load(
partial_sum + row * NUM_BLOCKS + blocks,
mask=blocks < NUM_BLOCKS,
other=0.0,
)
total = tl.sum(block_sum * tl.exp(block_max - maximum), axis=0)
normalizer = maximum + tl.log(total)
ranks = tl.arange(0, BLOCK_K)
candidates = tl.load(
top_values + row * K + ranks,
mask=ranks < K,
other=-float("inf"),
).to(tl.float32)
if TEMPERATURE_IS_TENSOR:
row_temperature = tl.load(
temperature_values + batch * TEMPERATURE_BATCH_STRIDE
).to(tl.float32)
row_inverse_temperature = tl.where(
row_temperature > 0.0,
1.0 / row_temperature,
1.0,
)
else:
row_inverse_temperature = inverse_temperature
log_probs = candidates * row_inverse_temperature - normalizer
# Proposal mass affects efficiency, not correctness. A malformed row
# must not leave the best-first frontier without a deterministic winner.
log_probs = tl.where(log_probs == log_probs, log_probs, -float("inf"))
log_probs = tl.where(log_probs > 0.0, 0.0, log_probs)
tl.store(
top_log_probs + row * K + ranks,
log_probs,
mask=ranks < K,
)
def _temperature_rows(
temperature: float | Tensor,
*,
batch_size: int,
device: torch.device,
) -> tuple[Tensor, int] | None:
if not isinstance(temperature, Tensor):
return None
if temperature.device != device:
raise ValueError("temperature and logits must share a device")
if not temperature.is_floating_point():
raise TypeError("temperature tensor must use a floating dtype")
if temperature.ndim == 0:
return temperature.reshape(1), 0
if temperature.shape not in ((batch_size,), (batch_size, 1)):
raise ValueError(
"temperature tensor must be scalar, [B], or [B, 1]; got "
f"{tuple(temperature.shape)}"
)
values = temperature.reshape(batch_size)
return values, values.stride(0)
def _build_candidate_log_probs(
logits: Tensor,
top_values: Tensor,
top_log_probs: Tensor,
partial_max: Tensor,
partial_sum: Tensor,
temperature: float | Tensor,
) -> None:
"""Normalize selected logits over the full vocabulary in FP32."""
batch_size, num_depths, vocab_size = logits.shape
num_rows = batch_size * num_depths
candidate_top_k = int(top_values.size(-1))
block_vocab = 8192
num_blocks = triton.cdiv(vocab_size, block_vocab)
temperature_rows = _temperature_rows(
temperature,
batch_size=batch_size,
device=logits.device,
)
if temperature_rows is None:
temperature_values = logits
temperature_batch_stride = 0
temperature_is_tensor = False
scalar_temperature = float(temperature)
inverse_temperature = (
1.0 / scalar_temperature if scalar_temperature > 0.0 else 1.0
)
else:
temperature_values, temperature_batch_stride = temperature_rows
temperature_is_tensor = True
inverse_temperature = 1.0
_candidate_lse_partials_kernel[(batch_size, num_depths, num_blocks)](
logits,
partial_max,
partial_sum,
temperature_values,
inverse_temperature,
BATCH_STRIDE=logits.stride(0),
DEPTH_STRIDE=logits.stride(1),
VOCAB_STRIDE=logits.stride(-1),
NUM_DEPTHS=num_depths,
VOCAB_SIZE=vocab_size,
NUM_BLOCKS=num_blocks,
BLOCK_VOCAB=block_vocab,
TEMPERATURE_BATCH_STRIDE=temperature_batch_stride,
TEMPERATURE_IS_TENSOR=temperature_is_tensor,
num_warps=4,
num_stages=1,
)
_candidate_lse_finalize_kernel[(num_rows,)](
top_values,
partial_max,
partial_sum,
top_log_probs,
temperature_values,
inverse_temperature,
K=candidate_top_k,
NUM_DEPTHS=num_depths,
NUM_BLOCKS=num_blocks,
BLOCK_K=triton.next_power_of_2(candidate_top_k),
BLOCK_PARTIALS=triton.next_power_of_2(num_blocks),
TEMPERATURE_BATCH_STRIDE=temperature_batch_stride,
TEMPERATURE_IS_TENSOR=temperature_is_tensor,
num_warps=1,
num_stages=1,
)
@triton.jit
def _build_tree_kernel(
top_token_ids,
top_log_probs,
draft_tokens,
selected_edges,
parent_list,
search_depths,
search_log_masses,
TOKEN_BATCH_STRIDE: tl.constexpr,
TOKEN_DEPTH_STRIDE: tl.constexpr,
TOKEN_RANK_STRIDE: tl.constexpr,
PROB_BATCH_STRIDE: tl.constexpr,
PROB_DEPTH_STRIDE: tl.constexpr,
PROB_RANK_STRIDE: tl.constexpr,
PARENT_BATCH_STRIDE: tl.constexpr,
NUM_DEPTHS: tl.constexpr,
K: tl.constexpr,
Q: tl.constexpr,
BLOCK_CANDIDATES: tl.constexpr,
):
batch = tl.program_id(0)
search_offset = batch * Q
edge_output_offset = batch * (Q - 1)
parent_output_offset = batch * PARENT_BATCH_STRIDE
candidate_slots = tl.arange(0, BLOCK_CANDIDATES)
candidate_parents = candidate_slots // K
candidate_ranks = candidate_slots % K
candidate_in_bounds = candidate_slots < Q * K
used = tl.zeros((BLOCK_CANDIDATES,), dtype=tl.int1)
# The root token itself is returned by reference. Only its private search
# state is stored here; EAGLE prepends it as the verify-tree root later.
tl.store(
parent_list + parent_output_offset + candidate_slots,
-1,
mask=candidate_slots < PARENT_BATCH_STRIDE,
)
tl.store(search_depths + search_offset, 0)
tl.store(search_log_masses + search_offset, 0.0)
tl.debug_barrier()
for node_index in range(1, Q):
parent_valid = candidate_in_bounds & (candidate_parents < node_index)
safe_parent = tl.where(parent_valid, candidate_parents, 0)
parent_depth = tl.load(
search_depths + search_offset + safe_parent,
mask=parent_valid,
other=NUM_DEPTHS,
)
parent_mass = tl.load(
search_log_masses + search_offset + safe_parent,
mask=parent_valid,
other=-float("inf"),
).to(tl.float32)
valid = parent_valid & (parent_depth < NUM_DEPTHS) & ~used
safe_depth = tl.where(valid, parent_depth, 0)
token = tl.load(
top_token_ids
+ batch * TOKEN_BATCH_STRIDE
+ safe_depth * TOKEN_DEPTH_STRIDE
+ candidate_ranks * TOKEN_RANK_STRIDE,
mask=valid,
other=0,
)
log_prob = tl.load(
top_log_probs
+ batch * PROB_BATCH_STRIDE
+ safe_depth * PROB_DEPTH_STRIDE
+ candidate_ranks * PROB_RANK_STRIDE,
mask=valid,
other=-float("inf"),
).to(tl.float32)
mass = parent_mass + log_prob
child_depth = parent_depth + 1
# Deterministic best-first order: mass, shallower depth, lower rank,
# lower token ID, then lower parent node.
best_mass = tl.max(tl.where(valid, mass, -float("inf")), axis=0)
winner = valid & (mass == best_mass)
best_depth = tl.min(tl.where(winner, child_depth, 1 << 30), axis=0)
winner &= child_depth == best_depth
best_rank = tl.min(tl.where(winner, candidate_ranks, 1 << 30), axis=0)
winner &= candidate_ranks == best_rank
best_token = tl.min(tl.where(winner, token, 1 << 30), axis=0)
winner &= token == best_token
best_parent = tl.min(tl.where(winner, candidate_parents, 1 << 30), axis=0)
winner &= candidate_parents == best_parent
best_slot = tl.min(tl.where(winner, candidate_slots, 1 << 30), axis=0)
selected_parent = best_slot // K
selected_rank = best_slot % K
selected_depth = tl.load(search_depths + search_offset + selected_parent)
selected_mass = tl.load(search_log_masses + search_offset + selected_parent).to(
tl.float32
) + tl.load(
top_log_probs
+ batch * PROB_BATCH_STRIDE
+ selected_depth * PROB_DEPTH_STRIDE
+ selected_rank * PROB_RANK_STRIDE
).to(tl.float32)
selected_token = tl.load(
top_token_ids
+ batch * TOKEN_BATCH_STRIDE
+ selected_depth * TOKEN_DEPTH_STRIDE
+ selected_rank * TOKEN_RANK_STRIDE
)
output_index = edge_output_offset + node_index - 1
tl.store(draft_tokens + output_index, selected_token)
# This is already EAGLE's implicit selected-edge encoding.
tl.store(selected_edges + output_index, best_slot)
if node_index < Q - 1:
# EAGLE shifts each selected edge by one candidate row. Fusing
# this write avoids separate fill/copy launches on every step.
tl.store(
parent_list + parent_output_offset + node_index,
best_slot,
)
tl.store(
search_depths + search_offset + node_index,
selected_depth + 1,
)
tl.store(
search_log_masses + search_offset + node_index,
selected_mass,
)
used |= candidate_slots == best_slot
tl.debug_barrier()
def _candidate_tree_capacity(
num_depths: int,
candidate_top_k: int,
stop_at: int,
) -> int:
capacity = 1
width = 1
for _ in range(num_depths):
width *= candidate_top_k
capacity += width
if capacity >= stop_at:
break
return capacity
@torch.inference_mode()
def build_uno_tree_proposal(
root_tokens: Tensor,
draft_logits: Tensor,
*,
max_nodes: int,
candidate_top_k: int,
temperature: float | Tensor,
workspace: dict[str, Tensor] | None = None,
) -> UnoTreeProposal:
"""Build fixed-``Q`` UNO trees directly in EAGLE's proposal ABI.
``root_tokens`` is ``[B]`` and ``draft_logits`` is ``[B, F-1, V]``.
Candidate log probabilities are normalized over the full vocabulary. No
CPU reference/fallback, direct parent array, attention mask, traversal
structure, acceptance walk, or KV operation is implemented here.
"""
if root_tokens.ndim != 1:
raise ValueError(
f"root_tokens must have shape [B], got {tuple(root_tokens.shape)}"
)
if draft_logits.ndim != 3 or draft_logits.size(0) != root_tokens.size(0):
raise ValueError(
"draft_logits must have shape [B, depth, vocab] with the same B "
f"as root_tokens; got {tuple(draft_logits.shape)}"
)
if not root_tokens.is_cuda or not draft_logits.is_cuda:
raise ValueError("UNO tree construction requires CUDA tensors")
if root_tokens.device != draft_logits.device:
raise ValueError("tree roots and draft logits must share a device")
if root_tokens.dtype not in (torch.int32, torch.int64):
raise TypeError("tree roots must use an integer dtype")
if not draft_logits.is_floating_point():
raise TypeError("draft logits must use a floating dtype")
if root_tokens.numel() == 0:
raise ValueError("UNO tree construction requires a non-empty batch")
if max_nodes < 1:
raise ValueError("max_nodes must include at least the root")
if candidate_top_k < 1:
raise ValueError("candidate_top_k must be >= 1")
batch_size = int(root_tokens.size(0))
num_depths = int(draft_logits.size(1))
vocab_size = int(draft_logits.size(2))
candidate_top_k = int(candidate_top_k)
draft_width = num_depths + 1
parent_width = candidate_top_k * max(num_depths - 1, 0) + 1
if max_nodes < draft_width:
raise ValueError(
f"max_nodes Q must be >= draft width F; got Q={max_nodes}, F={draft_width}"
)
if candidate_top_k > vocab_size:
raise ValueError(
f"candidate_top_k ({candidate_top_k}) exceeds vocabulary size "
f"({vocab_size})"
)
if max_nodes > 128 or max_nodes * candidate_top_k > 2048:
raise ValueError(
"the initial single-program UNO builder requires Q <= 128 and "
f"Q*K <= 2048; got Q={max_nodes}, K={candidate_top_k}"
)
capacity = _candidate_tree_capacity(
num_depths,
candidate_top_k,
max_nodes,
)
if capacity < max_nodes:
raise ValueError(
f"candidate set can produce only {capacity} tree nodes, but "
f"fixed tree verification requires {max_nodes}"
)
if max_nodes - 1 > parent_width:
raise ValueError(
"EAGLE's parent-list ABI cannot represent this UNO tree: "
f"Q-1={max_nodes - 1} exceeds K*(depth-1)+1={parent_width}"
)
_temperature_rows(
temperature,
batch_size=batch_size,
device=draft_logits.device,
)
def buffer(
name: str,
shape: tuple[int, ...],
dtype: torch.dtype,
) -> Tensor:
if workspace is None:
return torch.empty(
shape,
dtype=dtype,
device=draft_logits.device,
)
value = workspace.get(name)
if (
value is None
or value.shape != shape
or value.dtype != dtype
or value.device != draft_logits.device
):
value = torch.empty(
shape,
dtype=dtype,
device=draft_logits.device,
)
workspace[name] = value
return value
edge_shape = (batch_size, max_nodes - 1)
draft_tokens = buffer("draft_tokens", edge_shape, torch.long)
selected_edges = buffer("top_scores_index", edge_shape, torch.long)
parent_list = buffer(
"parent_list",
(batch_size, parent_width),
torch.long,
)
if max_nodes == 1:
parent_list.fill_(-1)
return UnoTreeProposal(
root_tokens=root_tokens,
draft_tokens=draft_tokens,
parent_list=parent_list,
top_scores_index=selected_edges,
candidate_top_k=candidate_top_k,
max_depth=num_depths,
)
candidate_shape = (batch_size, num_depths, candidate_top_k)
flat_logits = draft_logits.contiguous().view(batch_size * num_depths, vocab_size)
flat_top_values, flat_top_token_ids = _flashinfer_top_k(
flat_logits,
candidate_top_k,
sorted=True,
deterministic=False,
)
top_values = flat_top_values.view(candidate_shape)
top_token_ids = buffer("top_token_ids", candidate_shape, torch.long)
top_token_ids.copy_(flat_top_token_ids.view(candidate_shape))
top_log_probs = buffer("top_log_probs", candidate_shape, torch.float32)
num_partial_blocks = (vocab_size + 8191) // 8192
partial_shape = (batch_size * num_depths, num_partial_blocks)
_build_candidate_log_probs(
draft_logits,
top_values,
top_log_probs,
buffer("partial_lse_max", partial_shape, torch.float32),
buffer("partial_lse_sum", partial_shape, torch.float32),
temperature,
)
search_shape = (batch_size, max_nodes)
search_depths = buffer("search_depths", search_shape, torch.int32)
search_log_masses = buffer("search_log_masses", search_shape, torch.float32)
_build_tree_kernel[(batch_size,)](
top_token_ids,
top_log_probs,
draft_tokens,
selected_edges,
parent_list,
search_depths,
search_log_masses,
TOKEN_BATCH_STRIDE=top_token_ids.stride(0),
TOKEN_DEPTH_STRIDE=top_token_ids.stride(1),
TOKEN_RANK_STRIDE=top_token_ids.stride(2),
PROB_BATCH_STRIDE=top_log_probs.stride(0),
PROB_DEPTH_STRIDE=top_log_probs.stride(1),
PROB_RANK_STRIDE=top_log_probs.stride(2),
PARENT_BATCH_STRIDE=parent_list.stride(0),
NUM_DEPTHS=num_depths,
K=candidate_top_k,
Q=max_nodes,
BLOCK_CANDIDATES=triton.next_power_of_2(max_nodes * candidate_top_k),
num_warps=8,
num_stages=1,
)
return UnoTreeProposal(
root_tokens=root_tokens,
draft_tokens=draft_tokens,
parent_list=parent_list,
top_scores_index=selected_edges,
candidate_top_k=candidate_top_k,
max_depth=num_depths,
)
+587
View File
@@ -0,0 +1,587 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
import torch
from flashinfer import top_k as _flashinfer_top_k
from sglang.kernels.ops.speculative.reject_sampling import (
chain_speculative_sampling_triton,
)
from sglang.srt.speculative.dflash_utils import (
_get_or_create_chain_verify_buffers,
build_dflash_verify_target_probs,
)
from sglang.srt.speculative.spec_utils import fast_sample
_SPARSE_TOP_K_LIMIT = 128
def _normalize_sparse_topk_probs(
topk_logits: torch.Tensor,
temperatures: torch.Tensor,
valid: torch.Tensor,
top_ps: torch.Tensor,
) -> torch.Tensor:
"""Normalize a compact top-k support with top-k-first top-p semantics."""
scaled = topk_logits.float() / temperatures
scaled = scaled.masked_fill(~valid, float("-inf"))
probs = torch.softmax(scaled, dim=-1)
cdf = torch.cumsum(probs, dim=-1)
probs = probs.masked_fill((cdf - probs) > top_ps, 0.0)
return probs / probs.sum(dim=-1, keepdim=True).clamp_min(1e-12)
@torch.compile(dynamic=True)
def _sparse_rejection_from_support(
candidates: torch.Tensor,
target_ids: torch.Tensor,
target_probs: torch.Tensor,
draft_ids: torch.Tensor,
draft_probs: torch.Tensor,
accept_uniforms: torch.Tensor,
final_uniforms: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Run exact p/q rejection sampling on compact linear-chain supports."""
batch_size, forward_width = candidates.shape
num_proposals = forward_width - 1
if num_proposals == 0:
final_ids = target_ids[:, 0]
final_probs = target_probs[:, 0]
accepted_counts = torch.zeros(
batch_size,
dtype=torch.int32,
device=candidates.device,
)
else:
proposal_ids = candidates[:, 1:]
p_proposal = torch.where(
target_ids[:, :num_proposals].eq(proposal_ids.unsqueeze(-1)),
target_probs[:, :num_proposals],
torch.zeros(
(),
dtype=target_probs.dtype,
device=target_probs.device,
),
).sum(dim=-1)
q_proposal = torch.where(
draft_ids.eq(proposal_ids.unsqueeze(-1)),
draft_probs,
torch.zeros(
(),
dtype=draft_probs.dtype,
device=draft_probs.device,
),
).sum(dim=-1)
ratios = torch.where(
q_proposal > 0,
p_proposal / q_proposal,
torch.zeros_like(p_proposal),
).clamp_(max=1.0)
accepted_flags = accept_uniforms[:, :num_proposals] < ratios
accepted_counts = accepted_flags.to(torch.int32).cumprod(dim=1).sum(dim=1)
batch_indices = torch.arange(batch_size, device=candidates.device)
final_rows = accepted_counts.to(torch.long)
final_ids = target_ids[batch_indices, final_rows]
final_probs = target_probs[batch_indices, final_rows]
rejected = final_rows < num_proposals
draft_rows = final_rows.clamp(max=num_proposals - 1)
final_draft_ids = draft_ids[batch_indices, draft_rows]
final_draft_probs = draft_probs[batch_indices, draft_rows]
q_on_target = torch.where(
final_ids.unsqueeze(2).eq(final_draft_ids.unsqueeze(1)),
final_draft_probs.unsqueeze(1),
torch.zeros(
(),
dtype=final_draft_probs.dtype,
device=final_draft_probs.device,
),
).sum(dim=2)
correction_probs = (final_probs - q_on_target).clamp_min_(0.0)
correction_sum = correction_probs.sum(dim=1, keepdim=True)
correction_probs = torch.where(
correction_sum > 0,
correction_probs / correction_sum.clamp_min(1e-12),
final_probs,
)
final_probs = torch.where(
rejected[:, None],
correction_probs,
final_probs,
)
cdf = torch.cumsum(final_probs, dim=-1)
thresholds = final_uniforms * final_probs.sum(dim=-1)
sampled_offsets = (cdf <= thresholds[:, None]).sum(dim=-1)
sampled_offsets.clamp_(max=final_ids.shape[-1] - 1)
bonus = (
final_ids.gather(1, sampled_offsets[:, None]).squeeze(1).to(candidates.dtype)
)
return accepted_counts, bonus
@torch.compile(dynamic=True)
def _build_sparse_target_support_tensors(
next_token_logits: torch.Tensor,
temperatures: torch.Tensor,
top_ks: torch.Tensor,
top_ps: torch.Tensor,
batch_size: int,
forward_width: int,
max_top_k: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Build compact support using the fastest available top-k primitive."""
rows = batch_size * forward_width
topk_logits, topk_ids = _flashinfer_top_k(
next_token_logits.contiguous(),
max_top_k,
sorted=True,
deterministic=False,
)
expanded_temperatures = torch.repeat_interleave(
temperatures,
forward_width,
dim=0,
).reshape(rows, -1)
expanded_top_ks = torch.repeat_interleave(
top_ks,
forward_width,
dim=0,
).reshape(rows, 1)
expanded_top_ps = torch.repeat_interleave(
top_ps,
forward_width,
dim=0,
).reshape(rows, 1)
ranks = torch.arange(
max_top_k,
dtype=expanded_top_ks.dtype,
device=next_token_logits.device,
)[None, :]
probs = _normalize_sparse_topk_probs(
topk_logits,
expanded_temperatures,
ranks < expanded_top_ks,
expanded_top_ps,
)
return (
topk_ids.view(batch_size, forward_width, max_top_k),
probs.view(batch_size, forward_width, max_top_k),
)
def _build_sparse_target_support(
*,
next_token_logits: torch.Tensor,
sampling_info: Any,
batch_size: int,
forward_width: int,
max_top_k: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Return compact target token IDs/probabilities without a dense scatter."""
if bool(getattr(sampling_info, "need_top_p_sampling", False)):
top_ps = sampling_info.top_ps
else:
top_ps = torch.ones(
(batch_size,),
dtype=torch.float32,
device=next_token_logits.device,
)
return _build_sparse_target_support_tensors(
next_token_logits,
sampling_info.temperatures,
sampling_info.top_ks,
top_ps,
batch_size,
forward_width,
max_top_k,
)
def _sample_from_support(
support_ids: torch.Tensor,
support_probs: torch.Tensor,
) -> torch.Tensor:
"""Sample one token from every compact support row."""
flat_ids = support_ids.flatten(0, 1)
flat_probs = support_probs.flatten(0, 1)
_, offsets = fast_sample(flat_probs)
return flat_ids.gather(1, offsets).view(support_ids.shape[:2])
@dataclass(frozen=True)
class UnoDraftDistribution:
"""The exact q used to sample future UNO proposal rows."""
probs: torch.Tensor
token_ids: torch.Tensor | None = None
def _run_sparse_rejection(
*,
candidates: torch.Tensor,
next_token_logits: torch.Tensor,
sampling_info: Any,
max_top_k: int,
draft_distribution: UnoDraftDistribution,
) -> tuple[torch.Tensor, torch.Tensor]:
if draft_distribution.token_ids is None:
raise RuntimeError("Sparse UNO verification requires sparse draft q.")
batch_size, forward_width = candidates.shape
support_ids, support_probs = _build_sparse_target_support(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
batch_size=batch_size,
forward_width=forward_width,
max_top_k=max_top_k,
)
# Preserve the legacy path's two RNG draws and tensor shapes. The final
# uniforms select the correction/bonus; the first F - 1 acceptance coins
# are consumed by a linear chain.
accept_uniforms = torch.rand(
(batch_size, forward_width),
dtype=torch.float32,
device=next_token_logits.device,
)
final_uniforms = torch.rand(
(batch_size,),
dtype=torch.float32,
device=next_token_logits.device,
)
return _sparse_rejection_from_support(
candidates,
support_ids,
support_probs,
draft_distribution.token_ids,
draft_distribution.probs,
accept_uniforms,
final_uniforms,
)
def _build_dense_probs(
*,
next_token_logits: torch.Tensor,
sampling_info: Any,
batch_size: int,
forward_width: int,
max_top_k: int,
uniform_top_k_value: int | None,
) -> torch.Tensor:
"""Build the dense sampling distribution used by SGLang verification."""
return build_dflash_verify_target_probs(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
draft_token_num=forward_width,
bs=batch_size,
max_top_k=max_top_k,
uniform_top_k_value=uniform_top_k_value,
use_sparse_topk=True,
)
def _sample_from_dense_probs(probs: torch.Tensor) -> torch.Tensor:
"""Sample one token from every dense distribution row."""
_, token_ids = fast_sample(probs.flatten(0, 1))
return token_ids.view(probs.shape[:2])
def _run_dense_rejection(
*,
candidates: torch.Tensor,
next_token_logits: torch.Tensor,
sampling_info: Any,
max_top_k: int,
uniform_top_k_value: int | None,
draft_distribution: UnoDraftDistribution,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Run SGLang's fused linear-chain p/q rejection kernel."""
if draft_distribution.token_ids is not None:
raise RuntimeError("Dense UNO verification requires dense draft q.")
batch_size, forward_width = candidates.shape
target_probs = _build_dense_probs(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
batch_size=batch_size,
forward_width=forward_width,
max_top_k=max_top_k,
uniform_top_k_value=uniform_top_k_value,
)
accept_uniforms = torch.rand(
(batch_size, forward_width),
dtype=torch.float32,
device=next_token_logits.device,
)
final_uniforms = torch.rand(
(batch_size,),
dtype=torch.float32,
device=next_token_logits.device,
)
(
retrieve_index,
retrieve_next_token,
retrieve_next_sibling,
predicts,
accept_index,
accepted_counts,
) = _get_or_create_chain_verify_buffers(
bs=batch_size,
draft_token_num=forward_width,
device=next_token_logits.device,
)
chain_speculative_sampling_triton(
predicts=predicts,
accept_index=accept_index,
accept_token_num=accepted_counts,
candidates=candidates,
retrive_index=retrieve_index,
retrive_next_token=retrieve_next_token,
retrive_next_sibling=retrieve_next_sibling,
uniform_samples=accept_uniforms,
uniform_samples_for_final_sampling=final_uniforms,
target_probs=target_probs,
draft_probs=draft_distribution.probs,
threshold_single=1.0,
threshold_acc=1.0,
deterministic=True,
)
rows = torch.arange(batch_size, device=candidates.device)
bonus_positions = accept_index[
rows,
accepted_counts.to(torch.long),
].to(torch.long)
bonus = predicts[bonus_positions].to(candidates.dtype)
return accepted_counts, bonus
@dataclass
class UnoSamplingResult:
output_ids: torch.Tensor
accept_lens: torch.Tensor
new_seq_lens: torch.Tensor
next_seed_tokens: torch.Tensor
@dataclass
class UnoTreeSamplingResult:
output_ids: torch.Tensor
accept_lens: torch.Tensor
def build_uno_draft_input(
*,
seed_tokens: torch.Tensor,
forward_width: int,
vocab_size: int,
noise_tokens: torch.Tensor | None = None, # for testing
) -> torch.Tensor:
"""Build one ``[seed, uniform noise...]`` row per request.
Random noise is sampled independently from ``[0, vocab_size)``.
Supplying ``noise_tokens`` bypasses random generation. This is used
by deterministic tests and must have shape
``(batch_size, forward_width - 1)``.
``forward_width == 1`` never generates noise or consumes RNG state.
"""
seed_tokens = seed_tokens.reshape(-1).to(dtype=torch.int64)
batch_size = seed_tokens.numel()
noise_shape = (batch_size, forward_width - 1)
if forward_width == 1:
return seed_tokens[:, None]
if noise_tokens is None:
noise_tokens = torch.randint(
low=0,
high=vocab_size,
size=noise_shape,
dtype=torch.int64,
device=seed_tokens.device,
)
else:
noise_tokens = noise_tokens.to(
device=seed_tokens.device,
dtype=torch.int64,
)
draft_input_ids = seed_tokens.new_empty((batch_size, forward_width))
draft_input_ids[:, 0].copy_(seed_tokens)
draft_input_ids[:, 1:].copy_(noise_tokens)
return draft_input_ids
def sample_uno_candidates(
*,
draft_logits: torch.Tensor, # [B, F, V]
sampling_info: Any,
max_top_k: int,
uniform_top_k_value: int | None = None,
) -> tuple[torch.Tensor, UnoDraftDistribution]:
"""Sample 1 clean token from seed and F-1 draft tokens.
Depending on max_top_k value, may use sparse representations for efficiency.
candidates: [B, F] sampled tokens including clean and draft.
draft_distribution: probabilities of the draft tokens for rejection sampling.
"""
batch_size, forward_width, vocab_size = draft_logits.shape
flat_logits = draft_logits.reshape(-1, vocab_size)
if max_top_k <= _SPARSE_TOP_K_LIMIT:
support_ids, support_probs = _build_sparse_target_support(
next_token_logits=flat_logits,
sampling_info=sampling_info,
batch_size=batch_size,
forward_width=forward_width,
max_top_k=max_top_k,
)
candidates = _sample_from_support(support_ids, support_probs)
draft_distribution = UnoDraftDistribution(
token_ids=support_ids[:, 1:],
probs=support_probs[:, 1:],
)
else:
probs = _build_dense_probs(
next_token_logits=flat_logits,
sampling_info=sampling_info,
batch_size=batch_size,
forward_width=forward_width,
max_top_k=max_top_k,
uniform_top_k_value=uniform_top_k_value,
)
candidates = _sample_from_dense_probs(probs)
draft_distribution = UnoDraftDistribution(probs=probs[:, 1:])
return candidates, draft_distribution
def sample_uno_clean_root(
*,
seed_tokens: torch.Tensor,
draft_logits: torch.Tensor,
sampling_info: Any,
max_top_k: int,
uniform_top_k_value: int | None = None,
) -> torch.Tensor:
"""Sample the clean root using the current UNO sampling path."""
del seed_tokens
candidates, _ = sample_uno_candidates(
draft_logits=draft_logits[:, :1, :].contiguous(),
sampling_info=sampling_info,
max_top_k=max_top_k,
uniform_top_k_value=uniform_top_k_value,
)
return candidates[:, 0]
def sample_uno_tree_target_tokens(
*,
next_token_logits: torch.Tensor,
sampling_info: Any,
batch_size: int,
verify_width: int,
max_top_k: int,
) -> torch.Tensor:
"""Sample one target token per verify node from compact top-k support."""
support_ids, support_probs = _build_sparse_target_support(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
batch_size=batch_size,
forward_width=verify_width,
max_top_k=max_top_k,
)
return _sample_from_support(support_ids, support_probs)
def pack_uno_tree_result(
*,
clean_root_tokens: torch.Tensor,
eagle_predict: torch.Tensor,
eagle_accept_lens: torch.Tensor,
draft_width: int,
) -> UnoTreeSamplingResult:
"""Convert an internal EAGLE tree result into UNO's public row."""
clean_root_tokens = clean_root_tokens.reshape(-1).to(dtype=torch.int64)
batch_size = clean_root_tokens.numel()
verify_width = eagle_predict.numel() // batch_size
predict_rows = eagle_predict.reshape(batch_size, verify_width)
output_ids = clean_root_tokens.new_zeros((batch_size, draft_width + 1))
output_ids[:, 0].copy_(clean_root_tokens)
output_ids[:, 1:].copy_(predict_rows[:, :draft_width].to(dtype=output_ids.dtype))
return UnoTreeSamplingResult(
output_ids=output_ids,
accept_lens=eagle_accept_lens + 1,
)
def pack_uno_result(
*,
candidates: torch.Tensor, # [B, F]
accepted_proposal_counts: torch.Tensor, # [B]
bonus_tokens: torch.Tensor, # [B]
committed_frontiers: torch.Tensor, # [B]
) -> UnoSamplingResult:
"""Pack acceptance into fixed-width UNO output rows."""
batch_size, forward_width = candidates.shape
output_ids = candidates.new_zeros((batch_size, forward_width + 1))
output_ids[:, :forward_width].copy_(candidates)
output_ids.scatter_(
1,
(accepted_proposal_counts.to(torch.long) + 1)[:, None],
bonus_tokens[:, None],
)
accept_lens = accepted_proposal_counts + 2
new_seq_lens = committed_frontiers + accept_lens.to(committed_frontiers.dtype)
return UnoSamplingResult(
output_ids=output_ids,
accept_lens=accept_lens,
new_seq_lens=new_seq_lens,
next_seed_tokens=bonus_tokens,
)
def run_uno_sampling(
*,
candidates: torch.Tensor, # [B, F]
next_token_logits: torch.Tensor, # [B x F, V]
sampling_info: Any,
committed_frontiers: torch.Tensor, # [B]
draft_distribution: UnoDraftDistribution,
max_top_k: int,
uniform_top_k_value: int | None = None,
) -> UnoSamplingResult:
"""Verify sampled UNO proposals against target p and pack the result."""
if draft_distribution.token_ids is not None:
accepted, bonus = _run_sparse_rejection(
candidates=candidates,
next_token_logits=next_token_logits,
sampling_info=sampling_info,
max_top_k=max_top_k,
draft_distribution=draft_distribution,
)
else:
accepted, bonus = _run_dense_rejection(
candidates=candidates,
next_token_logits=next_token_logits,
sampling_info=sampling_info,
draft_distribution=draft_distribution,
max_top_k=max_top_k,
uniform_top_k_value=uniform_top_k_value,
)
return pack_uno_result(
candidates=candidates,
accepted_proposal_counts=accepted,
bonus_tokens=bonus,
committed_frontiers=committed_frontiers,
)
@@ -0,0 +1,50 @@
"""Request-admission validation for UNO speculative decoding."""
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
def validate_uno_request(req: Req) -> Optional[str]:
"""Return an error for request features that UNO cannot execute."""
sampling_params = req.sampling_params
if sampling_params.min_p > 0.0:
return "UNO speculative decoding does not support min_p sampling."
has_grammar = req.grammar is not None or any(
getattr(sampling_params, field) is not None
for field in ("json_schema", "regex", "ebnf", "structural_tag")
)
if has_grammar:
return "UNO speculative decoding does not support grammar decoding."
if req.return_logprob:
return "UNO speculative decoding does not support returned logprobs."
if req.return_hidden_states_mode.need_capture():
return "UNO speculative decoding does not support return_hidden_states."
has_penalties = (
sampling_params.frequency_penalty != 0.0
or sampling_params.presence_penalty != 0.0
or sampling_params.repetition_penalty != 1.0
or sampling_params.min_new_tokens > 0
)
if has_penalties:
return "UNO speculative decoding does not support sampling penalties."
if sampling_params.logit_bias is not None:
return "UNO speculative decoding does not support logit_bias."
if req.custom_logit_processor:
return "UNO speculative decoding does not support custom logit processors."
if req.lora_id is not None:
return "UNO speculative decoding does not support request-selectable LoRA."
return None
@@ -0,0 +1,867 @@
from __future__ import annotations
import contextlib
import copy
import logging
import time
from typing import TYPE_CHECKING, Optional
import torch
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
ForwardMode,
)
from sglang.srt.model_executor.forward_context import (
ForwardContext,
forward_context,
)
from sglang.srt.runtime_context import get_schedule, get_spec
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.speculative.eagle_utils import default_tree_mask_mode
from sglang.srt.speculative.eagle_worker_common import (
build_eagle_verify_input,
run_eagle_verify,
)
from sglang.srt.speculative.spec_info import SpecInputType, SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import get_plan_stream
from sglang.srt.speculative.uno_cuda_graph_runner import (
UnoDecodeCudaGraphRunner,
)
from sglang.srt.speculative.uno_info import UnoDraftInput, UnoForwardInput
from sglang.srt.speculative.uno_tree import build_uno_tree_proposal
from sglang.srt.speculative.uno_utils import (
build_uno_draft_input,
pack_uno_tree_result,
run_uno_sampling,
sample_uno_candidates,
sample_uno_clean_root,
)
from sglang.srt.utils.common import (
get_available_gpu_memory,
log_info_on_rank0,
)
if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
class UnoWorkerV2(BaseSpecWorker):
"""Single-model UNO worker with linear and native-EAGLE tree decode."""
def __init__(
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
super().__init__()
self.server_args = server_args
self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port
self._target_worker = target_worker
self._draft_worker = None
self.model_runner = target_worker.model_runner
self.lora_manager = self.model_runner.lora_manager
self.uno_lora_id = self.model_runner.uno_lora_id
self.device = target_worker.device
self.enable_overlap = not get_schedule().disable_overlap_schedule
configured_topk = int(get_spec().speculative_eagle_topk or 1)
self.tree_mode = configured_topk > 1
# Linear UNO stores F in speculative_num_draft_tokens. Tree UNO reuses
# EAGLE's native dimensions: F=steps+1, K=eagle_topk, Q=draft_tokens.
default_forward_width = (
int(get_spec().speculative_num_steps) + 1
if self.tree_mode
else int(get_spec().speculative_num_draft_tokens)
)
self.forward_width = default_forward_width
self.verify_width = int(get_spec().speculative_num_draft_tokens)
self.candidate_top_k = (
int(get_spec().speculative_eagle_topk) if self.tree_mode else 1
)
self.tree_depth = int(get_spec().speculative_num_steps) if self.tree_mode else 1
self.num_speculative_proposals = self.forward_width - 1
self.tail_width = self.forward_width + 1
# Compatibility fields read by speculative infrastructure.
self.speculative_num_draft_tokens = self.verify_width
self.speculative_num_steps = self.tree_depth
self.topk = self.candidate_top_k
# Ordinary scheduler overlap still serializes model work on the
# forward stream, so one persistent proposal workspace is sufficient:
# the next reuse is ordered after this step's tree build and verify.
self._uno_tree_workspace = {} if self.tree_mode else None
# The scheduler constructs speculative workers before allocating the
# target KV pools. Build the private tree-draft backend later, from
# init_attention_backends(), after those pools exist.
self._uno_draft_attn_backend = None
self._uno_draft_cuda_graph_runner = None
self.plan_stream, self.plan_stream_ctx = (
get_plan_stream(self.device)
if self.tree_mode
else (None, contextlib.nullcontext())
)
self._tail_offsets = torch.arange(
self.tail_width,
dtype=torch.int64,
device=self.device,
)
def _build_uno_draft_attn_backend(self):
"""Build only the F/1 backend absent from the native Q/K target role."""
model_runner = self.model_runner
original_workspace_flag = model_runner.init_new_workspace
try:
with get_spec().override(
speculative_num_steps=1,
speculative_eagle_topk=1,
speculative_num_draft_tokens=self.forward_width,
):
return model_runner._get_attention_backend(init_new_workspace=True)
finally:
model_runner.init_new_workspace = original_workspace_flag
def init_attention_backends(self):
"""Initialize only UNO's private backend after target pool allocation."""
if self.tree_mode:
self._uno_draft_attn_backend = self._build_uno_draft_attn_backend()
def init_cuda_graphs(self):
"""Capture only the private F-wide tree-draft graph."""
self._uno_draft_cuda_graph_runner = None
if not self.tree_mode or self.model_runner.decode_cuda_graph_runner is None:
return None
if self._uno_draft_attn_backend is None:
raise RuntimeError(
"UNO tree draft graph capture requires its attention backend."
)
tic = time.perf_counter()
before_mem = get_available_gpu_memory(
self.device,
self.gpu_id,
empty_cache=False,
)
log_info_on_rank0(
logger,
"Capture UNO tree draft CUDA graph begin. "
f"num_tokens_per_req={self.forward_width}, "
f"avail mem={before_mem:.2f} GB",
)
with self._bind_uno_draft_runtime():
self._uno_draft_cuda_graph_runner = UnoDecodeCudaGraphRunner(
self.model_runner,
tree_draft_attn_backend=self._uno_draft_attn_backend,
tree_draft_width=self.forward_width,
)
after_mem = get_available_gpu_memory(
self.device,
self.gpu_id,
empty_cache=False,
)
capture_time = time.perf_counter() - tic
self._additional_graph_memory_usage["draft_decode"] = before_mem - after_mem
self._additional_graph_time_usage["draft_decode"] = capture_time
log_info_on_rank0(
logger,
"Capture UNO tree draft CUDA graph end. "
f"elapsed={capture_time:.2f} s, "
f"mem usage={(before_mem - after_mem):.2f} GB, "
f"avail mem={after_mem:.2f} GB.",
)
return None
@property
def draft_worker(self):
# Both passes use the target runner and its KV pool.
return None
@property
def last_shared_read_runner(self):
# The target verify is the final phase that reads shared scheduler
# buffers, so its runner owns the WAR-barrier completion event.
return self._target_worker.model_runner
@property
def spec_v2_attn_backends(self) -> tuple:
"""Return every attention backend touched by one UNO step.
Linear UNO uses only the target runner's native backend. Tree UNO adds
one private F-wide draft backend before finishing on the native Q-wide
target backend. The scheduler ORs these capabilities when deciding
whether FutureMap must carry a CPU sequence-length mirror.
"""
target_backend = self._target_worker.model_runner.attn_backend
if not self.tree_mode:
return (target_backend,)
return (target_backend, self._uno_draft_attn_backend)
def __getattr__(self, name):
# Scheduler-facing methods not implemented by this wrapper belong to
# the target worker. Guard initialization to avoid recursive lookup.
if name == "_target_worker":
raise AttributeError(name)
return getattr(self.target_worker, name)
def _validate_batch(self, batch: ScheduleBatch) -> None:
if batch.forward_mode.is_idle():
raise NotImplementedError("UNO does not support idle batches.")
if batch.forward_mode.is_mixed() or (
batch.forward_mode.is_decode() and batch.is_extend_in_batch
):
raise NotImplementedError(
"UNO does not support mixed extend/decode batches."
)
if not batch.spec_algorithm.is_uno():
raise RuntimeError(
"UnoWorkerV2 received a batch whose speculative algorithm is not UNO."
)
sampling_info = batch.sampling_info
if sampling_info is None:
raise RuntimeError("UNO requires sampling metadata.")
if sampling_info.need_min_p_sampling:
raise NotImplementedError("UNO does not support min-p sampling.")
if batch.has_grammar:
raise NotImplementedError("UNO does not support grammar decoding.")
if batch.return_logprob:
raise NotImplementedError("UNO does not support returned logprobs.")
if batch.return_hidden_states:
raise NotImplementedError("UNO does not support returned hidden states.")
penalizer = sampling_info.penalizer_orchestrator
penalties_active = (
(penalizer is not None and penalizer.is_required)
or sampling_info.acc_additive_penalties is not None
or sampling_info.acc_scaling_penalties is not None
)
if penalties_active:
raise NotImplementedError("UNO does not support sampling penalties.")
if sampling_info.logit_bias is not None:
raise NotImplementedError("UNO does not support logit bias.")
if sampling_info.has_custom_logit_processor:
raise NotImplementedError("UNO does not support custom logit processors.")
if any(req.lora_id is not None for req in batch.reqs):
raise NotImplementedError("UNO does not support multi-LoRA.")
def _make_forward_batch(
self,
*,
spec_input_type: SpecInputType,
input_ids: torch.Tensor,
positions: torch.Tensor,
out_cache_loc: torch.Tensor,
prefix_lens: torch.Tensor,
seq_lens_cpu: Optional[torch.Tensor],
seq_lens_sum: Optional[int],
req_pool_indices: torch.Tensor,
) -> ForwardBatch:
input_ids = input_ids.reshape(-1)
positions = positions.reshape(-1)
out_cache_loc = out_cache_loc.reshape(-1)
spec_info = UnoForwardInput(
spec_input_type=spec_input_type,
positions=positions,
draft_token_num=self.forward_width,
)
return ForwardBatch(
forward_mode=ForwardMode.TARGET_VERIFY,
batch_size=len(prefix_lens),
input_ids=input_ids,
req_pool_indices=req_pool_indices,
seq_lens=prefix_lens,
out_cache_loc=out_cache_loc,
seq_lens_sum=seq_lens_sum,
seq_lens_cpu=seq_lens_cpu,
positions=positions,
spec_algorithm=SpeculativeAlgorithm.UNO,
spec_info=spec_info,
capture_hidden_mode=CaptureHiddenMode.NULL,
return_hidden_states_before_norm=False,
)
def _run_target_block(
self,
forward_batch: ForwardBatch,
*,
need_top1: bool = True,
) -> tuple:
result = self.target_worker.forward_batch_generation(
batch=None,
forward_batch=forward_batch,
is_verify=True,
)
if result.logits_output is None:
raise RuntimeError("UNO target block returned no logits output.")
logits = result.logits_output.next_token_logits
if logits is None:
raise RuntimeError("UNO target block returned no next-token logits.")
expected_rows = forward_batch.batch_size * self.forward_width
if logits.ndim != 2 or logits.shape[0] != expected_rows:
raise RuntimeError(
"UNO target block returned an invalid logits shape: "
f"expected ({expected_rows}, vocab_size), got "
f"{tuple(logits.shape)}."
)
# DFlash consumes logits directly; only greedy acceptance needs top-1.
if not need_top1:
return result, None
predictions = torch.argmax(logits, dim=-1).view(
forward_batch.batch_size,
self.forward_width,
)
return result, predictions
@staticmethod
def _accept_and_pack(
*,
candidates: torch.Tensor,
target_top1: torch.Tensor,
committed_seq_lens: torch.Tensor,
) -> tuple:
if candidates.ndim != 2:
raise RuntimeError(
f"UNO candidates must be rank 2, got shape={tuple(candidates.shape)}."
)
if target_top1.shape != candidates.shape:
raise RuntimeError(
"UNO candidate and target shapes differ: "
f"{tuple(candidates.shape)} versus {tuple(target_top1.shape)}."
)
batch_size, forward_width = candidates.shape
device = candidates.device
# forward_width is static configuration, so this branch does not inspect
# or synchronize a device tensor.
if forward_width == 1:
accepted_specs = torch.zeros(
batch_size,
dtype=torch.int32,
device=device,
)
else:
matches = candidates[:, 1:] == target_top1[:, :-1]
accepted_specs = (
matches.to(torch.int32).cumprod(dim=1).sum(dim=1).to(torch.int32)
)
accepted_specs_long = accepted_specs.to(torch.int64)
correction = target_top1.gather(
1,
accepted_specs_long[:, None],
).squeeze(1)
output_ids = torch.zeros(
(batch_size, forward_width + 1),
dtype=torch.int64,
device=device,
)
output_ids[:, :forward_width].copy_(candidates)
output_ids.scatter_(
1,
(accepted_specs_long + 1)[:, None],
correction[:, None],
)
accept_lens = accepted_specs + 2
new_seq_lens = committed_seq_lens + accept_lens.to(committed_seq_lens.dtype)
return output_ids, accept_lens, new_seq_lens, correction
def _forward_prefill(
self,
batch: ScheduleBatch,
on_publish,
) -> GenerationBatchResult:
result = self.target_worker.forward_batch_generation(batch)
if not isinstance(result.next_token_ids, torch.Tensor):
raise RuntimeError("UNO target prefill returned no sampled seed tensor.")
seed_tokens = result.next_token_ids.reshape(-1)
if seed_tokens.shape[0] != len(batch.reqs):
raise RuntimeError(
"UNO prefill seed count does not match batch size: "
f"{seed_tokens.shape[0]} versus {len(batch.reqs)}."
)
result.new_seq_lens = batch.seq_lens
result.next_draft_input = UnoDraftInput(
bonus_tokens=seed_tokens,
new_seq_lens=batch.seq_lens,
forward_width=self.forward_width,
)
if on_publish is not None:
on_publish(result.new_seq_lens)
return result
def _forward_decode_tree(
self,
batch: ScheduleBatch,
on_publish,
) -> GenerationBatchResult:
"""Run UNO's F-wide proposal pass, then native EAGLE Q-node verify."""
draft_state = batch.spec_info
if batch.seq_lens.is_cuda:
batch.seq_lens.record_stream(
torch.get_device_module(self.device).current_stream()
)
batch_size = len(batch.seq_lens)
committed_seq_lens = batch.seq_lens.clone()
seed_tokens = draft_state.bonus_tokens.reshape(-1).to(
device=self.device,
dtype=torch.int64,
)
committed_seq_lens_cpu = None
if batch.seq_lens_cpu is not None:
committed_seq_lens_cpu = batch.seq_lens_cpu.to(
device="cpu",
dtype=torch.int64,
)
draft_seq_lens_cpu = committed_seq_lens_cpu + self.forward_width
draft_seq_lens_sum = int(draft_seq_lens_cpu.sum())
elif draft_state.reserved_seq_lens_cpu is not None:
# This host tensor is a planning upper bound only. The device
# frontier below remains the exact committed length.
draft_seq_lens_cpu = draft_state.reserved_seq_lens_cpu
draft_seq_lens_sum = draft_state.reserved_seq_lens_sum
else:
draft_seq_lens_cpu = None
draft_seq_lens_sum = None
draft_positions = (
committed_seq_lens.to(torch.int64)[:, None]
+ (self._tail_offsets[None, : self.forward_width])
)
req_pool_indices_long = batch.req_pool_indices.to(torch.int64)
req_to_token = self.model_runner.req_to_token_pool.req_to_token
draft_locs = req_to_token[
req_pool_indices_long[:, None],
draft_positions,
].to(torch.int64)
draft_input_ids = build_uno_draft_input(
seed_tokens=seed_tokens,
forward_width=self.forward_width,
vocab_size=self.model_runner.model_config.vocab_size,
)
draft_forward_batch = self._make_forward_batch(
spec_input_type=SpecInputType.UNO_DRAFT,
input_ids=draft_input_ids,
positions=draft_positions,
out_cache_loc=draft_locs,
prefix_lens=committed_seq_lens,
seq_lens_cpu=draft_seq_lens_cpu,
seq_lens_sum=draft_seq_lens_sum,
req_pool_indices=batch.req_pool_indices,
)
draft_result, _ = self._run_draft_block(
draft_forward_batch,
need_top1=False,
)
draft_logits = draft_result.logits_output.next_token_logits.reshape(
batch_size,
self.forward_width,
-1,
)
sampling_info = batch.sampling_info
if sampling_info.is_all_greedy:
clean_root_tokens = torch.argmax(
draft_logits[:, 0, :],
dim=-1,
)
else:
clean_root_tokens = sample_uno_clean_root(
seed_tokens=seed_tokens,
draft_logits=draft_logits,
sampling_info=sampling_info,
max_top_k=draft_state.max_top_k,
uniform_top_k_value=draft_state.uniform_top_k_value,
)
proposal = build_uno_tree_proposal(
clean_root_tokens,
draft_logits[:, 1:, :],
max_nodes=self.verify_width,
candidate_top_k=self.candidate_top_k,
temperature=sampling_info.temperatures,
workspace=self._uno_tree_workspace,
)
# The first pass wrote the carried seed at C. Give EAGLE a shallow
# batch whose KV-ready prefix is therefore C+1; its existing allocator
# assigns all Q tree slots at C+1 and its compactor can stay unchanged.
verify_batch = copy.copy(batch)
verify_batch.seq_lens = committed_seq_lens + 1
if committed_seq_lens_cpu is None:
verify_batch.seq_lens_cpu = None
verify_batch.seq_lens_sum = None
else:
verify_batch.seq_lens_cpu = committed_seq_lens_cpu + 1
verify_batch.seq_lens_sum = int(verify_batch.seq_lens_cpu.sum())
verify_input = build_eagle_verify_input(
verify_batch,
EagleDraftInput(bonus_tokens=proposal.root_tokens),
proposal.parent_list,
proposal.top_scores_index,
proposal.draft_tokens,
None,
target_worker=self.target_worker,
topk=self.candidate_top_k,
num_steps=self.tree_depth,
num_draft_tokens=self.verify_width,
tree_mask_mode=default_tree_mask_mode(),
device=self.device,
)
verify_batch.spec_info = verify_input
if self.plan_stream is not None:
# C+1 was produced on the forward stream immediately above. The
# generic EAGLE path receives an older, already-visible frontier;
# UNO must explicitly order its freshly derived tensor before the
# plan stream assigns Q verify cache locations from it.
self.plan_stream.wait_stream(
torch.get_device_module(self.device).current_stream()
)
eagle_result = run_eagle_verify(
verify_batch,
target_worker=self.target_worker,
req_to_token_pool=self.model_runner.req_to_token_pool,
token_to_kv_pool_allocator=(self.model_runner.token_to_kv_pool_allocator),
plan_stream=self.plan_stream,
plan_stream_ctx=self.plan_stream_ctx,
topk=self.candidate_top_k,
num_draft_tokens=self.verify_width,
device=self.device,
metadata_ready_pre_pad=False,
finalize_tree_path=True,
uno_target_max_top_k=draft_state.max_top_k,
)
packed = pack_uno_tree_result(
clean_root_tokens=clean_root_tokens,
eagle_predict=eagle_result.next_token_ids,
eagle_accept_lens=eagle_result.accept_lens,
draft_width=self.forward_width,
)
new_seq_lens = eagle_result.new_seq_lens
next_draft_input = UnoDraftInput(
bonus_tokens=eagle_result.next_draft_input.bonus_tokens,
new_seq_lens=new_seq_lens,
forward_width=self.forward_width,
)
if on_publish is not None:
on_publish(new_seq_lens)
# Preserve EAGLE's verify ForwardBatch keep-alive refs verbatim. FutureMap
# relays the wrapped bonus; on_publish above relays the new frontier.
return GenerationBatchResult(
logits_output=eagle_result.logits_output,
next_token_ids=packed.output_ids.reshape(-1),
accept_lens=packed.accept_lens,
next_draft_input=next_draft_input,
speculative_num_draft_tokens=self.forward_width,
speculative_output_stride=self.forward_width + 1,
num_non_draft_tokens_per_req=2,
new_seq_lens=new_seq_lens,
can_run_cuda_graph=eagle_result.can_run_cuda_graph,
routed_experts_output=eagle_result.routed_experts_output,
indexer_topk_output=eagle_result.indexer_topk_output,
extra_keep_alive_refs=eagle_result.extra_keep_alive_refs,
)
def _forward_decode(
self,
batch: ScheduleBatch,
on_publish,
) -> GenerationBatchResult:
if self.tree_mode:
return self._forward_decode_tree(batch, on_publish)
draft_state = batch.spec_info
if batch.seq_lens.is_cuda:
batch.seq_lens.record_stream(
torch.get_device_module(self.device).current_stream()
)
batch_size = len(batch.seq_lens)
committed_seq_lens = batch.seq_lens.clone()
seed_tokens = draft_state.bonus_tokens.reshape(-1).to(
device=self.device,
dtype=torch.int64,
)
if batch.seq_lens_cpu is not None:
committed_seq_lens_cpu = batch.seq_lens_cpu.to(
device="cpu",
dtype=torch.int32,
)
draft_seq_lens_cpu = committed_seq_lens_cpu + self.forward_width
verify_seq_lens_cpu = committed_seq_lens_cpu + self.tail_width
draft_seq_lens_sum = int(draft_seq_lens_cpu.sum())
verify_seq_lens_sum = int(verify_seq_lens_cpu.sum())
elif draft_state.reserved_seq_lens_cpu is not None:
# Triton only needs a safe host planning bound. The allocator's
# retained reservation avoids a D2H copy when FutureMap keeps the
# exact committed frontier on GPU.
draft_seq_lens_cpu = draft_state.reserved_seq_lens_cpu
verify_seq_lens_cpu = draft_state.reserved_seq_lens_cpu
draft_seq_lens_sum = draft_state.reserved_seq_lens_sum
verify_seq_lens_sum = draft_state.reserved_seq_lens_sum
else:
draft_seq_lens_cpu = None
verify_seq_lens_cpu = None
draft_seq_lens_sum = None
verify_seq_lens_sum = None
logical_positions = (
committed_seq_lens.to(torch.int64)[:, None] + (self._tail_offsets[None, :])
)
req_pool_indices_long = batch.req_pool_indices.to(torch.int64)
req_to_token = self.model_runner.req_to_token_pool.req_to_token
tail_locs = req_to_token[
req_pool_indices_long[:, None],
logical_positions,
].to(torch.int64)
draft_input_ids = build_uno_draft_input(
seed_tokens=seed_tokens,
forward_width=self.forward_width,
vocab_size=self.model_runner.model_config.vocab_size,
)
draft_forward_batch = self._make_forward_batch(
spec_input_type=SpecInputType.UNO_DRAFT,
input_ids=draft_input_ids,
positions=logical_positions[:, : self.forward_width],
out_cache_loc=tail_locs[:, : self.forward_width],
prefix_lens=committed_seq_lens,
seq_lens_cpu=draft_seq_lens_cpu,
seq_lens_sum=draft_seq_lens_sum,
req_pool_indices=batch.req_pool_indices,
)
sampling_info = batch.sampling_info
all_greedy = sampling_info.is_all_greedy
max_top_k = draft_state.max_top_k
uniform_top_k_value = draft_state.uniform_top_k_value
draft_result, candidates = self._run_draft_block(
draft_forward_batch,
need_top1=all_greedy,
)
if not all_greedy:
draft_logits = draft_result.logits_output.next_token_logits.reshape(
batch_size,
self.forward_width,
-1,
)
candidates, draft_distribution = sample_uno_candidates(
draft_logits=draft_logits,
sampling_info=sampling_info,
max_top_k=max_top_k,
uniform_top_k_value=uniform_top_k_value,
)
verify_prefix_lens = committed_seq_lens + 1
verify_forward_batch = self._make_forward_batch(
spec_input_type=SpecInputType.UNO_VERIFY,
input_ids=candidates,
positions=logical_positions[:, 1:],
out_cache_loc=tail_locs[:, 1:],
prefix_lens=verify_prefix_lens,
seq_lens_cpu=verify_seq_lens_cpu,
seq_lens_sum=verify_seq_lens_sum,
req_pool_indices=batch.req_pool_indices,
)
verify_result, target_top1 = self._run_target_block(
verify_forward_batch,
need_top1=all_greedy,
)
if all_greedy:
output_ids, accept_lens, new_seq_lens, correction = self._accept_and_pack(
candidates=candidates,
target_top1=target_top1,
committed_seq_lens=committed_seq_lens,
)
else:
sampling_result = run_uno_sampling(
candidates=candidates,
next_token_logits=verify_result.logits_output.next_token_logits,
sampling_info=sampling_info,
committed_frontiers=committed_seq_lens,
draft_distribution=draft_distribution,
max_top_k=max_top_k,
uniform_top_k_value=uniform_top_k_value,
)
output_ids = sampling_result.output_ids
accept_lens = sampling_result.accept_lens
new_seq_lens = sampling_result.new_seq_lens
correction = sampling_result.next_seed_tokens
next_draft_input = UnoDraftInput(
bonus_tokens=correction,
new_seq_lens=new_seq_lens,
forward_width=self.forward_width,
)
if on_publish is not None:
on_publish(new_seq_lens)
return GenerationBatchResult(
logits_output=verify_result.logits_output,
next_token_ids=output_ids.reshape(-1),
accept_lens=accept_lens,
next_draft_input=next_draft_input,
speculative_num_draft_tokens=self.forward_width,
speculative_output_stride=self.tail_width,
num_non_draft_tokens_per_req=2,
new_seq_lens=new_seq_lens,
can_run_cuda_graph=verify_result.can_run_cuda_graph,
routed_experts_output=verify_result.routed_experts_output,
indexer_topk_output=verify_result.indexer_topk_output,
)
def forward_batch_generation(
self,
batch: ScheduleBatch,
on_publish=None,
grammar_barrier=None,
) -> GenerationBatchResult:
del grammar_barrier
self._validate_batch(batch)
if batch.forward_mode == ForwardMode.EXTEND:
return self._forward_prefill(batch, on_publish)
if batch.forward_mode == ForwardMode.DECODE:
return self._forward_decode(batch, on_publish)
raise RuntimeError(
f"UNO expected an EXTEND or DECODE batch, got {batch.forward_mode}."
)
def update_weights_from_disk(self, recv_req):
# The scheduler updates the target worker before calling the spec worker.
return True, "UNO has no separate draft weights."
def update_weights_from_ipc(self, recv_req):
# The scheduler updates the target worker before calling the spec worker.
return True, "UNO has no separate draft weights."
def update_weights_from_tensor(self, recv_req):
# This update route selects the spec worker instead of updating both.
return self.target_worker.update_weights_from_tensor(recv_req)
@contextlib.contextmanager
def _bind_uno_draft_runtime(self):
target_attn_backend = self.model_runner.attn_backend
target_graph_runner = self.model_runner.decode_cuda_graph_runner
self.model_runner.attn_backend = self._uno_draft_attn_backend
self.model_runner.decode_cuda_graph_runner = self._uno_draft_cuda_graph_runner
try:
yield
finally:
self.model_runner.attn_backend = target_attn_backend
self.model_runner.decode_cuda_graph_runner = target_graph_runner
def _run_draft_block(
self,
forward_batch: ForwardBatch,
*,
need_top1: bool = True,
):
batch_size = forward_batch.batch_size
self.lora_manager.reset_lora_batch()
backend_context = (
self._bind_uno_draft_runtime()
if self.tree_mode
else contextlib.nullcontext()
)
attn_context = (
forward_context(ForwardContext(attn_backend=self._uno_draft_attn_backend))
if self.tree_mode
else contextlib.nullcontext()
)
# The eager runner plans through model_runner.attn_backend, while
# attention layers execute through ForwardContext. Tree UNO binds both
# to the same private F/1 backend for this draft pass.
with backend_context, attn_context:
if self.num_speculative_proposals == 0:
return self._run_target_block(
forward_batch,
need_top1=need_top1,
)
graph_runner = getattr(
self.model_runner,
"decode_cuda_graph_runner",
None,
)
# If cuda-graph is on, reuse its captured LoRA routing
# and replay it in _run_target_block.
# Else, prepare routing for eager.
if not (
forward_batch.forward_mode.is_cuda_graph()
and graph_runner is not None
and graph_runner.can_run_graph(forward_batch)
):
self.lora_manager.prepare_lora_token_segments(
lora_ids=[None, self.uno_lora_id] * batch_size,
segment_lens=[1, self.num_speculative_proposals] * batch_size,
)
result = self._run_target_block(
forward_batch,
need_top1=need_top1,
)
# Clear LoRA routing before verification
self.lora_manager.reset_lora_batch()
return result
@@ -0,0 +1,121 @@
"""CUDA correctness tests for the fused suffix-attention merge."""
import math
import unittest
import torch
from sglang.kernels.ops.attention.suffix_attention_merge import (
merge_suffix_attention_in_place,
)
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(
est_time=15,
stage="base-b-kernel-unit",
runner_config="1-gpu-large",
)
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
class TestSuffixAttentionMerge(CustomTestCase):
def _case(self, *, num_queries: int, head_dim: int, dtype: torch.dtype):
torch.manual_seed(7)
device = torch.device("cuda")
num_q_heads = 8
num_kv_heads = 2
num_slots = 2 * num_queries + 11
q = torch.randn(num_queries, num_q_heads, head_dim, device=device, dtype=dtype)
k_cache = torch.randn(
num_slots, num_kv_heads, head_dim, device=device, dtype=dtype
)
v_cache = torch.randn_like(k_cache)
page_table = torch.stack(
[
torch.randperm(num_slots, device=device)[:num_queries]
for _ in range(num_queries)
]
).to(torch.int32)
suffix_lengths = (
torch.arange(num_queries, device=device, dtype=torch.int32)
.remainder(num_queries)
.add_(1)
)
prefix = torch.randn_like(q)
prefix_lse = torch.randn(
num_q_heads, num_queries, device=device, dtype=torch.float32
)
scale = 1.0 / math.sqrt(head_dim)
reference = prefix.float().clone()
heads_per_kv = num_q_heads // num_kv_heads
kv_heads = torch.arange(num_q_heads, device=device) // heads_per_kv
for token in range(num_queries):
length = int(suffix_lengths[token])
slots = page_table[token, :length].long()
keys = k_cache[slots][:, kv_heads].float()
values = v_cache[slots][:, kv_heads].float()
scores = torch.einsum("lhd,hd->lh", keys, q[token].float()) * scale
maximum = torch.maximum(prefix_lse[:, token], scores.max(dim=0).values)
prefix_weight = torch.exp(prefix_lse[:, token] - maximum)
suffix_weights = torch.exp(scores - maximum)
reference[token] = (
reference[token] * prefix_weight[:, None]
+ torch.einsum("lh,lhd->hd", suffix_weights, values)
) / (prefix_weight + suffix_weights.sum(dim=0))[:, None]
static_prefix = prefix.clone()
merge_suffix_attention_in_place(
q,
k_cache,
v_cache,
page_table,
suffix_lengths,
static_prefix,
prefix_lse,
scale,
)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
static_prefix.copy_(prefix)
merge_suffix_attention_in_place(
q,
k_cache,
v_cache,
page_table,
suffix_lengths,
static_prefix,
prefix_lse,
scale,
)
graph.replay()
torch.cuda.synchronize()
torch.testing.assert_close(
static_prefix.float(), reference, rtol=2e-2, atol=2e-2
)
def test_representative_shapes(self):
cases = (
(16, 64, torch.float16),
(60, 128, torch.bfloat16),
)
for num_queries, head_dim, dtype in cases:
with self.subTest(
num_queries=num_queries,
head_dim=head_dim,
dtype=dtype,
):
self._case(
num_queries=num_queries,
head_dim=head_dim,
dtype=dtype,
)
if __name__ == "__main__":
unittest.main()
+267
View File
@@ -0,0 +1,267 @@
"""End-to-end CUDA-graph coverage for linear and tree UNO decoding.
The test runs both modes on the same prompts. Linear UNO alternates
LoRA-draft and clean-target variants in one graph runner. Tree UNO uses a
private LoRA-draft runner before native EAGLE tree verification. Besides the
generation contract, short greedy comparisons guard lossless output parity
with autoregressive decoding, and the stochastic comparison guards that tree
search improves TPF over the linear proposal on a small, fixed GSM8K sample.
"""
import os
import unittest
from typing import NamedTuple
import requests
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
register_cuda_ci(
est_time=480,
stage="base-b",
runner_config="1-gpu-large",
)
MODEL = "Qwen/Qwen3-8B"
DEFAULT_UNO_LORA = "s-sahoo/uno-qwen3-8B"
LORA_PATH_ENV = "SGLANG_TEST_UNO_LORA_PATH"
MAX_NEW_TOKENS = 128
# AR decode and UNO verification use different kernel shapes, so compare a
# bounded greedy prefix instead of requiring full-sequence bitwise identity.
PARITY_TOKENS = 32
# One LoRA draft forward plus one clean verification forward.
FORWARDS_PER_UNO_CYCLE = 2
PROMPTS = (
(
"Question: Janet's ducks lay 16 eggs per day. She eats three for "
"breakfast every morning and bakes muffins for her friends every day "
"with four. She sells the remainder at the farmers' market daily for "
"$2 per fresh duck egg. How much in dollars does she make every day "
"at the farmers' market?\nAnswer:"
),
(
"Question: A robe takes 2 bolts of blue fiber and half that much "
"white fiber. How many bolts in total does it take?\nAnswer:"
),
(
"Question: Josh decides to try flipping a house. He buys a house for "
"$80,000 and then puts in $50,000 in repairs. This increased the value "
"of the house by 150%. How much profit did he make?\nAnswer:"
),
)
class _UnoConfig(NamedTuple):
name: str
speculative_num_steps: int
speculative_eagle_topk: int
speculative_num_draft_tokens: int
LINEAR_CONFIG = _UnoConfig(
name="linear",
speculative_num_steps=1,
speculative_eagle_topk=1,
speculative_num_draft_tokens=8, # F = 8
)
TREE_CONFIG = _UnoConfig(
name="tree",
speculative_num_steps=7, # F = 8
speculative_eagle_topk=16,
speculative_num_draft_tokens=8, # Q = 8
)
class TestUnoCudaGraph(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.base_url = DEFAULT_URL_FOR_TEST
cls.adapter_path = os.environ.get(LORA_PATH_ENV, DEFAULT_UNO_LORA)
def _server_args(self, config: _UnoConfig | None) -> list[str]:
args = [
"--dtype",
"bfloat16",
"--attention-backend",
"fa3",
"--max-running-requests",
str(len(PROMPTS)),
"--cuda-graph-max-bs-decode",
str(len(PROMPTS)),
"--mem-fraction-static",
"0.7",
"--page-size",
"1",
"--disable-radix-cache",
"--random-seed",
"17",
]
if config is not None:
args.extend(
[
"--speculative-algorithm",
"UNO",
"--uno-lora-path",
self.adapter_path,
"--speculative-num-steps",
str(config.speculative_num_steps),
"--speculative-eagle-topk",
str(config.speculative_eagle_topk),
"--speculative-num-draft-tokens",
str(config.speculative_num_draft_tokens),
]
)
return args
def _run_ar_reference(self) -> list[list[int]]:
process = None
try:
process = popen_launch_server(
MODEL,
self.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=self._server_args(None),
)
return self._run_greedy_output_ids()
finally:
if process is not None:
kill_process_tree(process.pid)
def _run_config(self, config: _UnoConfig) -> tuple[float, list[list[int]]]:
process = None
try:
process = popen_launch_server(
MODEL,
self.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=self._server_args(config),
)
greedy_output_ids = self._run_greedy_output_ids()
tpf = self._run_generation_contract(config)
return tpf, greedy_output_ids
finally:
if process is not None:
kill_process_tree(process.pid)
def _run_greedy_output_ids(self) -> list[list[int]]:
# A list-valued request can be admitted with different prefill batch
# shapes across server launches. Run each parity prompt at BS1 so the
# AR and UNO comparisons use the same execution shape.
output_ids = []
for prompt in PROMPTS:
response = requests.post(
self.base_url + "/generate",
json={
"text": prompt,
"sampling_params": {
"temperature": 0,
"max_new_tokens": PARITY_TOKENS,
"ignore_eos": True,
},
},
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
)
self.assertEqual(response.status_code, 200, response.text)
result = response.json()
self.assertIn("output_ids", result, result)
self.assertEqual(
len(result["output_ids"]),
PARITY_TOKENS,
f"Wrong greedy output length for prompt {prompt!r}",
)
output_ids.append(result["output_ids"])
return output_ids
def _assert_ar_parity(
self,
mode: str,
actual: list[list[int]],
expected: list[list[int]],
) -> None:
for prompt, actual_ids, expected_ids in zip(PROMPTS, actual, expected):
self.assertEqual(
actual_ids,
expected_ids,
f"{mode} UNO diverged from AR within the first "
f"{PARITY_TOKENS} tokens for prompt {prompt!r}",
)
def _run_generation_contract(self, config: _UnoConfig) -> float:
server_info = requests.get(self.base_url + "/server_info", timeout=30).json()
self.assertEqual(
server_info["speculative_eagle_topk"], config.speculative_eagle_topk
)
self.assertEqual(
server_info["speculative_num_steps"], config.speculative_num_steps
)
self.assertEqual(
server_info["speculative_num_draft_tokens"],
config.speculative_num_draft_tokens,
)
response = requests.post(
self.base_url + "/generate",
json={
"text": PROMPTS,
"sampling_params": {
"temperature": 0.7,
"top_k": 50,
"top_p": 0.95,
"max_new_tokens": MAX_NEW_TOKENS,
"ignore_eos": True,
},
},
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
)
self.assertEqual(response.status_code, 200, response.text)
results = response.json()
self.assertEqual(len(results), len(PROMPTS))
total_completion_tokens = 0
total_verify_ct = 0
for result in results:
self.assertTrue(result["text"].strip())
meta_info = result["meta_info"]
self.assertEqual(meta_info["completion_tokens"], MAX_NEW_TOKENS)
total_completion_tokens += meta_info["completion_tokens"]
total_verify_ct += meta_info.get("spec_verify_ct", 0)
self.assertGreater(
total_verify_ct, 0, f"{config.name} performed no verify steps"
)
total_forwards = FORWARDS_PER_UNO_CYCLE * total_verify_ct
tpf = total_completion_tokens / total_forwards
self.assertGreater(
tpf,
1.5,
f"{config.name} did not advance beyond autoregressive decoding: {tpf=}",
)
return tpf
def test_ar_parity_and_tree_tpf_exceeds_linear(self):
ar_output_ids = self._run_ar_reference()
linear_tpf, linear_output_ids = self._run_config(LINEAR_CONFIG)
self._assert_ar_parity("Linear", linear_output_ids, ar_output_ids)
tree_tpf, tree_output_ids = self._run_config(TREE_CONFIG)
self._assert_ar_parity("Tree", tree_output_ids, ar_output_ids)
print(f"UNO GSM8K sample: {linear_tpf=:.3f}, {tree_tpf=:.3f}")
self.assertGreater(
tree_tpf,
linear_tpf,
f"Tree UNO did not improve TPF: {linear_tpf=:.3f}, {tree_tpf=:.3f}",
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,41 @@
"""Regression test for UNO's base-only LoRA routing fast path."""
import unittest
from types import SimpleNamespace
from sglang.srt.lora.lora_manager import LoRAManager
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class _InactiveSkippingBackend:
skip_inactive_lora_batches = True
def __init__(self):
self.batch_info = object()
self.prepare_called = False
def reset_batch_state(self):
self.batch_info = None
def prepare_lora_batch(self, *args, **kwargs):
self.prepare_called = True
class TestUnoInactiveLoRABatch(CustomTestCase):
def test_all_base_batch_clears_stale_routing_before_graph_metadata(self):
backend = _InactiveSkippingBackend()
manager = LoRAManager.__new__(LoRAManager)
manager.lora_backend = backend
forward_batch = SimpleNamespace(lora_ids=[None], batch_size=1)
manager.prepare_lora_batch(forward_batch)
self.assertIsNone(backend.batch_info)
self.assertFalse(backend.prepare_called)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,236 @@
"""Target-layer validation for UNO's specialized LoRA backend."""
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import torch
from sglang.srt.layers.linear import (
ColumnParallelLinear,
ReplicatedLinear,
RowParallelLinear,
)
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.lora.backend.triton_backend import TritonLoRABackend
from sglang.srt.lora.backend.uno_cublas_backend import UnoCublasLoRABackend
from sglang.srt.lora.lora_manager import LoRAManager
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestUnoLoRATargets(CustomTestCase):
def setUp(self):
self.backend = UnoCublasLoRABackend.__new__(UnoCublasLoRABackend)
self.backend._pending_lora_a = None
self.backend._use_cublas_lora_b = False
@staticmethod
def _model(modules, **attributes):
return SimpleNamespace(
named_modules=lambda: modules,
**attributes,
)
def test_supported_decoder_targets_are_accepted(self):
modules = [
(
"model.layers.0.qkv_proj",
ColumnParallelLinear.__new__(ColumnParallelLinear),
),
(
"model.layers.0.o_proj",
RowParallelLinear.__new__(RowParallelLinear),
),
(
"model.layers.0.fused_qkv_a_proj_with_mqa",
ReplicatedLinear.__new__(ReplicatedLinear),
),
]
self.backend.validate_lora_targets(
base_model=self._model(modules),
target_modules={
"qkv_proj",
"o_proj",
"fused_qkv_a_proj_with_mqa",
},
)
def test_unsupported_targets_are_rejected(self):
cases = {
"unknown decoder layer": (
self._model(
[
(
"model.layers.0.custom_proj",
torch.nn.Linear(2, 2),
)
]
),
{"custom_proj"},
"Linear",
),
"fused MoE": (
self._model(
[
(
"model.layers.0.mlp",
FusedMoE.__new__(FusedMoE),
)
]
),
{"gate_up_proj", "down_proj"},
"FusedMoE",
),
}
for name, (model, targets, expected) in cases.items():
with self.subTest(name=name), self.assertRaisesRegex(ValueError, expected):
self.backend.validate_lora_targets(
base_model=model,
target_modules=targets,
)
def test_nonoverlap_dense_calls_fall_back_to_triton(self):
x = object()
weights = object()
hidden = object()
base_output = object()
pruned_batch_info = object()
expected = object()
with (
patch.object(
TritonLoRABackend,
"run_lora_a_sgemm",
return_value=hidden,
) as run_lora_a,
patch.object(
TritonLoRABackend,
"run_lora_b_sgemm",
return_value=expected,
) as run_lora_b,
):
actual_hidden = self.backend.run_lora_a_sgemm(
x,
weights,
pruned_batch_info=pruned_batch_info,
)
actual = self.backend.run_lora_b_sgemm(
actual_hidden,
weights,
base_output=base_output,
pruned_batch_info=pruned_batch_info,
)
self.assertIs(actual_hidden, hidden)
self.assertIs(actual, expected)
run_lora_a.assert_called_once_with(
x,
weights,
pruned_batch_info,
1,
)
run_lora_b.assert_called_once_with(
hidden,
weights,
base_output,
pruned_batch_info,
)
def test_overlap_launch_selects_cublas(self):
pending = object()
x = object()
weights = object()
hidden = object()
base_output = object()
expected = object()
self.backend._pending_lora_a = pending
self.backend._consume_lora_a_overlap = MagicMock(return_value=hidden)
self.backend._run_lora_b = MagicMock(return_value=expected)
with (
patch.object(TritonLoRABackend, "run_lora_a_sgemm") as run_lora_a,
patch.object(TritonLoRABackend, "run_lora_b_sgemm") as run_lora_b,
):
actual_hidden = self.backend.run_lora_a_sgemm(x, weights)
actual = self.backend.run_lora_b_sgemm(
actual_hidden,
weights,
base_output=base_output,
)
self.assertIs(actual_hidden, hidden)
self.assertIs(actual, expected)
self.backend._consume_lora_a_overlap.assert_called_once_with(pending)
self.backend._run_lora_b.assert_called_once_with(
hidden,
weights,
base_output,
)
self.assertFalse(self.backend._use_cublas_lora_b)
run_lora_a.assert_not_called()
run_lora_b.assert_not_called()
def test_nonoverlap_qkv_call_falls_back_to_triton(self):
expected = object()
args = {
"x": object(),
"qkv_lora_a": object(),
"qkv_lora_b": object(),
"output_offset": object(),
"output_offset_cpu": object(),
"max_qkv_out_dim": 128,
"base_output": object(),
"n_slices": 2,
}
with patch.object(
TritonLoRABackend,
"run_qkv_lora",
return_value=expected,
) as run_qkv_lora:
actual = self.backend.run_qkv_lora(**args)
self.assertIs(actual, expected)
run_qkv_lora.assert_called_once_with(
args["x"],
args["qkv_lora_a"],
args["qkv_lora_b"],
args["output_offset"],
128,
args["base_output"],
2,
)
def test_manager_preflights_targets_before_wrapping(self):
manager = LoRAManager.__new__(LoRAManager)
manager.base_model = object()
manager.lora_backend = MagicMock()
manager._experts_shared_outer_override = None
manager.init_lora_adapters = MagicMock()
manager.init_lora_shapes = MagicMock(
side_effect=lambda **_: setattr(manager, "target_modules", {"qkv_proj"})
)
manager._detect_shared_outer_loras = MagicMock(return_value=False)
manager.init_lora_modules = MagicMock()
manager.init_memory_pool = MagicMock()
manager.update_lora_info = MagicMock()
manager.lora_backend.validate_lora_targets.side_effect = ValueError(
"unsupported target"
)
with self.assertRaisesRegex(ValueError, "unsupported target"):
manager.init_state(max_lora_rank=1, target_modules={"q_proj"})
manager.lora_backend.validate_lora_targets.assert_called_once_with(
base_model=manager.base_model,
target_modules={"qkv_proj"},
)
manager.init_lora_modules.assert_not_called()
if __name__ == "__main__":
unittest.main()
@@ -7,6 +7,7 @@ import torch
from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
@@ -178,17 +179,8 @@ class TestDecodeHiddenStateRetention(CustomTestCase):
second_step = torch.arange(16, dtype=torch.float32).view(8, 2)[4:]
def result(hidden_states):
return SimpleNamespace(
copy_done=None,
auxiliary_host_output=None,
routed_experts_output=None,
indexer_topk_output=None,
return GenerationBatchResult(
logits_output=SimpleNamespace(hidden_states=hidden_states),
next_token_ids=None,
can_run_cuda_graph=False,
num_correct_drafts=0,
num_block_accept_tokens=0,
num_cap_tokens=0,
speculative_num_draft_tokens=4,
)
@@ -10,6 +10,7 @@ from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.runtime_context import get_context
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.test.ci.ci_register import register_cpu_ci
@@ -78,17 +79,9 @@ def _make_processor() -> SchedulerBatchResultProcessor:
def _make_result():
return SimpleNamespace(
copy_done=None,
auxiliary_host_output=None,
routed_experts_output=None,
indexer_topk_output=None,
return GenerationBatchResult(
logits_output=SimpleNamespace(hidden_states=None, customized_info=None),
next_token_ids=[4],
can_run_cuda_graph=False,
num_correct_drafts=0,
num_block_accept_tokens=0,
num_cap_tokens=0,
speculative_num_draft_tokens=0,
)
@@ -14,6 +14,7 @@ from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.sampling.sampling_params import (
REQUEST_REASONING_END_TOKEN_IDS_KEY,
SamplingParams,
@@ -101,16 +102,10 @@ def _make_req(terminate_after: int) -> Req:
def _make_result(num_draft_tokens, accept_lens, flat_tokens):
return SimpleNamespace(
return GenerationBatchResult(
next_token_ids=torch.tensor(flat_tokens, dtype=torch.long),
accept_lens=torch.tensor(accept_lens, dtype=torch.long),
speculative_num_draft_tokens=num_draft_tokens,
num_correct_drafts=None,
num_correct_drafts_per_req_cpu=None,
block_accept_lens=None,
cap_lens=None,
copy_done=None,
grammar_advanced=False,
)
@@ -0,0 +1,67 @@
"""Scheduler containment for unsupported UNO requests."""
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
maybe_stub_sgl_kernel()
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.managers.scheduler import Scheduler
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
class TestSchedulerUnoRequestValidation(CustomTestCase):
def test_invalid_request_is_aborted_before_scheduler_admission(self):
scheduler = Scheduler.__new__(Scheduler)
scheduler.enable_session_radix_cache = False
scheduler.model_config = SimpleNamespace(
hf_eos_token_id={1},
vocab_size=128,
)
scheduler.disaggregation_mode = DisaggregationMode.NULL
scheduler.metrics_reporter = SimpleNamespace(enable_metrics=False)
scheduler.tokenizer = None
scheduler.dllm_config = None
scheduler._maybe_namespace_elastic_radix_cache = MagicMock()
scheduler.spec_algorithm = SimpleNamespace(
is_dflash_family=lambda: False,
is_uno=lambda: True,
)
scheduler.init_req_max_new_tokens = MagicMock()
scheduler._add_request_to_queue = MagicMock()
recv_req = MagicMock(
session_params=None,
session_id=None,
input_embeds=None,
bootstrap_port=1,
)
req = MagicMock()
error = "UNO request is unsupported."
with (
patch(
"sglang.srt.managers.scheduler.BeamCoordinator.request_beam_width",
return_value=1,
),
patch("sglang.srt.managers.scheduler.Req", return_value=req),
patch(
"sglang.srt.managers.scheduler.validate_uno_request",
return_value=error,
) as validate_uno_request,
):
scheduler.handle_generate_request(recv_req)
validate_uno_request.assert_called_once_with(req)
req.set_finish_with_abort.assert_called_once_with(error)
scheduler.init_req_max_new_tokens.assert_called_once_with(req)
scheduler._add_request_to_queue.assert_called_once_with(req)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,74 @@
"""CPU regressions for UNO aggregate token accounting."""
import unittest
from types import SimpleNamespace
from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.managers.scheduler_components.metrics_reporter import (
SchedulerMetricsReporter,
)
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestUnoTokenAccounting(CustomTestCase):
def setUp(self):
self.batch_size = 2
self.result = GenerationBatchResult(
num_correct_drafts=3,
num_non_draft_tokens_per_req=2,
)
def test_generated_token_count_includes_both_non_draft_tokens(self):
self.assertEqual(self.result.get_num_generated_tokens(self.batch_size), 7)
self.assertEqual(
GenerationBatchResult(num_correct_drafts=3).get_num_generated_tokens(
self.batch_size
),
5,
)
def test_spec_metrics_keep_generated_and_draft_counts_separate(self):
reporter = SchedulerMetricsReporter.__new__(SchedulerMetricsReporter)
reporter.spec_num_accept_tokens = 0
reporter.spec_num_correct_drafts = 0
reporter.spec_num_forward_ct = 0
reporter.spec_num_block_accept_tokens = 0
reporter.spec_num_cap_tokens = 0
reporter.update_spec_metrics(
self.batch_size,
self.result.num_correct_drafts,
num_accept_tokens=self.result.get_num_generated_tokens(self.batch_size),
)
self.assertEqual(reporter.spec_num_accept_tokens, 7)
self.assertEqual(reporter.spec_num_correct_drafts, 3)
self.assertEqual(reporter.spec_num_forward_ct, 2)
def test_decode_moment_receives_full_generated_token_count(self):
scheduler = Scheduler.__new__(Scheduler)
scheduler._prev_step = (1, 10.0, False)
scheduler.decode_moment_totals = [0.0] * 6
batch = SimpleNamespace(
forward_mode=SimpleNamespace(
is_extend_without_speculative=lambda: False,
is_decode=lambda: True,
is_target_verify=lambda: False,
),
reqs=[SimpleNamespace(rid="req-0"), SimpleNamespace(rid="req-1")],
forward_iter=2,
launch_ts=10.001,
after_idle_gap=False,
)
scheduler._record_step_counters(batch, self.result)
self.assertEqual(scheduler.decode_moment_totals[5], 7)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,36 @@
"""Unit tests for UNO allocation sizing."""
import unittest
from sglang.srt.mem_cache.allocation_sizing import (
get_alloc_len_per_decode,
get_alloc_reserve_per_decode,
get_req_to_token_extra_context_len,
)
from sglang.srt.runtime_context import get_context, get_parallel
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestUnoAllocationSizing(CustomTestCase):
def test_page_size_one_row_covers_decode_reserve(self):
with (
get_context().override_server_args(
speculative_algorithm="UNO",
speculative_num_draft_tokens=8,
page_size=1,
),
get_parallel().override(attn_dcp_size=1),
):
self.assertEqual(get_alloc_len_per_decode(), 9)
self.assertEqual(get_alloc_reserve_per_decode(), 18)
self.assertGreaterEqual(
get_req_to_token_extra_context_len(),
get_alloc_reserve_per_decode(),
)
if __name__ == "__main__":
unittest.main()
@@ -14,6 +14,7 @@ import unittest
from unittest.mock import patch
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -25,6 +26,7 @@ class TestDraftRunnerSkipsLoRA(CustomTestCase):
runner = ModelRunner.__new__(ModelRunner)
runner.is_draft_worker = is_draft_worker
runner.lora_manager = None
runner.spec_algorithm = SpeculativeAlgorithm.NONE
with patch.object(ModelRunner, "init_lora_manager") as init_lora:
with patch("sglang.srt.model_executor.model_runner.get_lora") as get_lora:
get_lora.return_value.enable_lora = enable_lora
@@ -45,6 +45,10 @@ _DFLASH_DECODE = (
"speculative/dflash_info_v2.py",
"DFlashDraftInputV2.prepare_for_decode",
)
_UNO_DECODE = (
"speculative/uno_info.py",
"UnoDraftInput.prepare_for_decode",
)
_RESOLVE = (
"managers/scheduler_components/batch_result_processor.py",
"SchedulerBatchResultProcessor._resolve_spec_v2_tokens",
@@ -71,6 +75,8 @@ _OWNER_SITES = {
# one of these two owners for each speculative decode iteration.
(*_DFLASH_DECODE, "decode_batch_idx"): 1,
(*_DFLASH_DECODE, "evict"): 1,
(*_UNO_DECODE, "decode_batch_idx"): 1,
(*_UNO_DECODE, "evict"): 1,
(
"mem_cache/allocation.py",
"alloc_for_spec_decode",
@@ -0,0 +1,72 @@
"""CPU contracts for the fused suffix-attention merge dispatch guard."""
import unittest
from types import SimpleNamespace
import torch
from sglang.kernels.ops.attention.suffix_attention_merge import (
can_use_fused_suffix_attention_merge,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestSuffixAttentionMergeDispatch(CustomTestCase):
def _inputs(self):
layer = SimpleNamespace(
head_dim=64,
v_head_dim=64,
is_cross_attention=False,
logit_cap=0.0,
)
q = torch.empty((16, 8 * 64), dtype=torch.bfloat16)
key_cache = torch.empty((8, 16, 2, 64), dtype=torch.bfloat16)
value_cache = torch.empty_like(key_cache)
return layer, q, key_cache, value_cache
def _eligible(self, **overrides):
layer, q, key_cache, value_cache = self._inputs()
arguments = dict(
layer=layer,
q=q,
key_cache=key_cache,
value_cache=value_cache,
extra_kwargs={},
)
arguments.update(overrides)
return can_use_fused_suffix_attention_merge(**arguments)
def test_standard_attention_is_eligible(self):
self.assertTrue(self._eligible())
def test_special_attention_features_fall_back(self):
self.assertFalse(self._eligible(extra_kwargs={"sinks": object()}))
layer, _, _, _ = self._inputs()
layer.is_cross_attention = True
self.assertFalse(self._eligible(layer=layer))
layer, _, _, _ = self._inputs()
layer.logit_cap = 20.0
self.assertFalse(self._eligible(layer=layer))
def test_unsupported_tensor_layout_falls_back(self):
layer, _, _, _ = self._inputs()
layer.v_head_dim = 32
self.assertFalse(self._eligible(layer=layer))
_, q, key_cache, value_cache = self._inputs()
self.assertFalse(
self._eligible(
q=q.float(),
key_cache=key_cache.float(),
value_cache=value_cache.float(),
)
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,93 @@
"""Request-admission validation for UNO speculative decoding."""
import unittest
from types import SimpleNamespace
from sglang.srt.speculative.uno_validation import validate_uno_request
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
def _make_request(**overrides):
sampling_params = SimpleNamespace(
min_p=0.0,
json_schema=None,
regex=None,
ebnf=None,
structural_tag=None,
frequency_penalty=0.0,
presence_penalty=0.0,
repetition_penalty=1.0,
min_new_tokens=0,
logit_bias=None,
)
request = SimpleNamespace(
sampling_params=sampling_params,
grammar=None,
return_logprob=False,
return_hidden_states_mode=SimpleNamespace(need_capture=lambda: False),
custom_logit_processor=None,
lora_id=None,
)
for name, value in overrides.items():
target, field = name.split("__", maxsplit=1)
owner = sampling_params if target == "sampling_params" else request
setattr(owner, field, value)
return request
class TestUnoRequestValidation(CustomTestCase):
def test_supported_request_is_accepted(self):
self.assertIsNone(validate_uno_request(_make_request()))
def test_unsupported_request_features_are_rejected(self):
cases = {
"min_p": ({"sampling_params__min_p": 0.1}, "min_p"),
"grammar": ({"sampling_params__regex": "[0-9]+"}, "grammar"),
"logprobs": ({"request__return_logprob": True}, "logprobs"),
"hidden states": (
{
"request__return_hidden_states_mode": SimpleNamespace(
need_capture=lambda: True
)
},
"return_hidden_states",
),
"frequency penalty": (
{"sampling_params__frequency_penalty": 0.1},
"penalties",
),
"presence penalty": (
{"sampling_params__presence_penalty": 0.1},
"penalties",
),
"repetition penalty": (
{"sampling_params__repetition_penalty": 1.1},
"penalties",
),
"minimum new tokens": (
{"sampling_params__min_new_tokens": 1},
"penalties",
),
"logit bias": (
{"sampling_params__logit_bias": {1: 0.5}},
"logit_bias",
),
"custom processor": (
{"request__custom_logit_processor": "processor"},
"custom logit processors",
),
"public LoRA": ({"request__lora_id": "adapter"}, "LoRA"),
}
for name, (overrides, expected) in cases.items():
with self.subTest(name=name):
error = validate_uno_request(_make_request(**overrides))
self.assertIsNotNone(error)
self.assertIn(expected, error)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,73 @@
"""Startup validation for UNO configuration."""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.arg_groups.speculative_hook import _handle_uno
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestUnoTreeConfig(CustomTestCase):
def test_unsupported_runtime_modes_are_rejected_at_startup(self):
cases = {
"deterministic inference": (
{"enable_deterministic_inference": True},
"enable-deterministic-inference",
),
"strict thinking": (
{"enable_strict_thinking": True},
"enable-strict-thinking",
),
}
for name, (overrides, expected) in cases.items():
with self.subTest(name=name):
values = {
"device": "cuda",
"speculative_draft_model_path": None,
"uno_lora_path": "/tmp/uno-lora",
"enable_deterministic_inference": False,
"enable_strict_thinking": False,
}
values.update(overrides)
server_args = SimpleNamespace(**values)
with (
patch(
"sglang.srt.arg_groups.speculative_hook.resolving_view",
side_effect=lambda args: args,
),
self.assertRaisesRegex(ValueError, expected),
):
_handle_uno(server_args)
def test_parent_list_overflow_is_rejected_at_startup(self):
"""An invalid tree must not survive startup and crash on first decode."""
server_args = SimpleNamespace(
device="cuda",
enable_deterministic_inference=False,
enable_strict_thinking=False,
speculative_draft_model_path=None,
uno_lora_path="/tmp/uno-lora",
speculative_num_draft_tokens=8,
speculative_num_steps=3,
speculative_eagle_topk=2,
)
with (
patch(
"sglang.srt.arg_groups.speculative_hook.resolving_view",
side_effect=lambda args: args,
),
patch("sglang.srt.arg_groups.speculative_hook.declare_resolution"),
self.assertRaisesRegex(ValueError, "parent-list ABI"),
):
_handle_uno(server_args)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,105 @@
"""CPU contracts for UNO tree compact target sampling."""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
from sglang.srt.speculative.eagle_utils import (
_can_use_sparse_uno_tree_target_sampling,
)
from sglang.srt.speculative.uno_utils import sample_uno_tree_target_tokens
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestUnoTreeSparseSampling(CustomTestCase):
def test_sparse_dispatch_guard(self):
sampling_info = SimpleNamespace(
sampling_seed=None,
need_min_p_sampling=False,
)
spec_config = SimpleNamespace(
speculative_use_rejection_sampling=False,
)
with (
patch("sglang.srt.speculative.eagle_utils._is_cuda", True),
patch(
"sglang.srt.speculative.eagle_utils.get_spec",
return_value=spec_config,
),
):
self.assertTrue(
_can_use_sparse_uno_tree_target_sampling(128, sampling_info)
)
self.assertFalse(
_can_use_sparse_uno_tree_target_sampling(None, sampling_info)
)
self.assertFalse(
_can_use_sparse_uno_tree_target_sampling(129, sampling_info)
)
sampling_info.sampling_seed = torch.tensor([1])
self.assertFalse(
_can_use_sparse_uno_tree_target_sampling(128, sampling_info)
)
sampling_info.sampling_seed = None
sampling_info.need_min_p_sampling = True
self.assertFalse(
_can_use_sparse_uno_tree_target_sampling(128, sampling_info)
)
sampling_info.need_min_p_sampling = False
spec_config.speculative_use_rejection_sampling = True
self.assertFalse(
_can_use_sparse_uno_tree_target_sampling(128, sampling_info)
)
def test_targets_are_sampled_from_compact_support(self):
support_ids = torch.tensor(
[
[[10, 11], [20, 21], [30, 31]],
[[40, 41], [50, 51], [60, 61]],
],
dtype=torch.int64,
)
support_probs = torch.full((2, 3, 2), 0.5)
sampled_offsets = torch.tensor(
[[0], [1], [0], [1], [0], [1]],
dtype=torch.long,
)
sampling_info = SimpleNamespace()
next_token_logits = torch.empty((6, 100))
with (
patch(
"sglang.srt.speculative.uno_utils._build_sparse_target_support",
return_value=(support_ids, support_probs),
) as build_support,
patch(
"sglang.srt.speculative.uno_utils.fast_sample",
return_value=(torch.empty((6, 1)), sampled_offsets),
),
):
targets = sample_uno_tree_target_tokens(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
batch_size=2,
verify_width=3,
max_top_k=2,
)
self.assertEqual(targets.tolist(), [[10, 21, 30], [41, 50, 61]])
build_support.assert_called_once_with(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
batch_size=2,
forward_width=3,
max_top_k=2,
)
if __name__ == "__main__":
unittest.main()