[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:
co-authored by
drproduck
BBuf
parent
354ed6d66b
commit
2bb25dc18b
@@ -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.
|
||||
@@ -0,0 +1 @@
|
||||
"""Speculative-decoding benchmark utilities."""
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}, "
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user