[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"
|
title: "Speculative Decoding"
|
||||||
metatags:
|
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
|
## 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-2 Decoding via Frequency-Ranked Speculative Sampling](#eagle-2-decoding-via-frequency-ranked-speculative-sampling)
|
||||||
- [EAGLE-3 Decoding](#eagle-3-decoding)
|
- [EAGLE-3 Decoding](#eagle-3-decoding)
|
||||||
- [Multi Token Prediction](#multi-token-prediction)
|
- [Multi Token Prediction](#multi-token-prediction)
|
||||||
|
- [UNO decoding](#uno-decoding)
|
||||||
- [DFlash Decoding](#dflash-decoding)
|
- [DFlash Decoding](#dflash-decoding)
|
||||||
- [Standalone Speculative Decoding (Small Draft Model)](#standalone-speculative-decoding-small-draft-model)
|
- [Standalone Speculative Decoding (Small Draft Model)](#standalone-speculative-decoding-small-draft-model)
|
||||||
- [Speculative Decoding V2 (Overlap Scheduler)](#speculative-decoding-v2-overlap-scheduler)
|
- [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`.
|
- **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`.
|
- **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).
|
- **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 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`).
|
- **You have a smaller draft LLM**: Use **STANDALONE** (`--speculative-algorithm STANDALONE`).
|
||||||
- **No extra model available**: Use **NGRAM** (`--speculative-algorithm NGRAM`, CUDA-only).
|
- **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)"}}>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>
|
<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>
|
||||||
|
<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>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>DFLASH</td>
|
<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>
|
<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
|
## 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`.
|
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", 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.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.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>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--speculative-draft-model-path</code></td>
|
<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.02)"}}><code>None</code></td>
|
||||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Path to the draft model weights</td>
|
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Path to the draft model weights</td>
|
||||||
</tr>
|
</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>
|
<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", 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>
|
<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", 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.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.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>
|
||||||
<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", 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.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.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>
|
||||||
<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", 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.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.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>
|
||||||
<tr>
|
<tr>
|
||||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--speculative-dflash-block-size</code></td>
|
<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"),
|
("extend_attention", "build_unified_kv_indices"),
|
||||||
("prefill_attention", "context_attention_fwd"),
|
("prefill_attention", "context_attention_fwd"),
|
||||||
("merge_state", "merge_state_triton"),
|
("merge_state", "merge_state_triton"),
|
||||||
|
("suffix_attention_merge", "merge_suffix_attention_in_place"),
|
||||||
("metadata", "get_num_kv_splits_triton"),
|
("metadata", "get_num_kv_splits_triton"),
|
||||||
("metadata", "prepare_swa_spec_page_table_triton"),
|
("metadata", "prepare_swa_spec_page_table_triton"),
|
||||||
("metadata", "normal_decode_set_metadata"),
|
("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:
|
def _target_checkpoint_bundles_dspark_draft(server_args: ServerArgs) -> bool:
|
||||||
from sglang.srt.speculative.dspark_components.dspark_config import (
|
from sglang.srt.speculative.dspark_components.dspark_config import (
|
||||||
checkpoint_bundles_dspark_draft,
|
checkpoint_bundles_dspark_draft,
|
||||||
|
|||||||
@@ -12,6 +12,10 @@ from sglang.kernels.ops.attention.metadata import (
|
|||||||
prepare_swa_spec_page_table_triton,
|
prepare_swa_spec_page_table_triton,
|
||||||
)
|
)
|
||||||
from sglang.kernels.ops.attention.pa_page_table import _build_pa_page_table
|
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.attention.utils import assert_buffer_fits
|
||||||
from sglang.kernels.ops.kvcache.trtllm_mha_page_table import (
|
from sglang.kernels.ops.kvcache.trtllm_mha_page_table import (
|
||||||
build_trtllm_mha_page_table,
|
build_trtllm_mha_page_table,
|
||||||
@@ -1579,14 +1583,36 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
|
|
||||||
if use_cascade_attn:
|
if use_cascade_attn:
|
||||||
o, softmax_lse, *rest = result
|
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(
|
o_expand, softmax_lse_expand, *rest_expand = flash_attn_with_kvcache(
|
||||||
q=q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
|
q=q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim),
|
||||||
# Here metadata_expand.page_table is not divided with page_size.
|
# The suffix table stores physical token slots, so expose
|
||||||
# This is because we loose the fine control of what token to attend,
|
# the paged cache as page-size-one blocks.
|
||||||
# but has to attend to some block completely.
|
|
||||||
k_cache=key_cache.view(-1, 1, layer.tp_k_head_num, layer.head_dim),
|
k_cache=key_cache.view(-1, 1, layer.tp_k_head_num, layer.head_dim),
|
||||||
v_cache=value_cache.view(
|
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,
|
page_table=self.forward_metadata_spec_decode_expand.page_table,
|
||||||
cache_seqlens=self.forward_metadata_spec_decode_expand.cache_seqlens_int32,
|
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.
|
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
|
# Supporting backends implement init_prefill_cuda_graph_batch_info() and
|
||||||
# honor use_prefill_cuda_graph in prepare_lora_batch().
|
# honor use_prefill_cuda_graph in prepare_lora_batch().
|
||||||
supports_prefill_cuda_graph: bool = False
|
supports_prefill_cuda_graph: bool = False
|
||||||
@@ -52,6 +55,14 @@ class BaseLoRABackend(LoRABackendLmHeadMixing):
|
|||||||
self.lm_head_pass_batch_infos = None
|
self.lm_head_pass_batch_infos = None
|
||||||
self._lm_head_pass_idx = 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(
|
def run_lora_a_embedding(
|
||||||
self,
|
self,
|
||||||
input_ids: torch.Tensor,
|
input_ids: torch.Tensor,
|
||||||
@@ -351,6 +362,20 @@ class BaseLoRABackend(LoRABackendLmHeadMixing):
|
|||||||
"""
|
"""
|
||||||
pass
|
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
|
@triton.jit
|
||||||
def _compute_moe_lora_info_kernel(
|
def _compute_moe_lora_info_kernel(
|
||||||
|
|||||||
@@ -44,6 +44,13 @@ def create_torch_native_backend():
|
|||||||
return TorchNativeLoRABackend
|
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")
|
@register_lora_backend("flashinfer")
|
||||||
def create_flashinfer_backend():
|
def create_flashinfer_backend():
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
@@ -362,6 +362,44 @@ class TritonLoRABackend(BaseLoRABackend):
|
|||||||
self._prepare_lm_head_batch_info(forward_batch, weight_indices, batch_info)
|
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(
|
def _prepare_lm_head_batch_info(
|
||||||
self,
|
self,
|
||||||
forward_batch: ForwardBatch,
|
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
|
has LoRA batch metadata. batch_info is None on DP-attention idle
|
||||||
forwards (see LoRAManager.prepare_lora_batch), so idle forwards take
|
forwards (see LoRAManager.prepare_lora_batch), so idle forwards take
|
||||||
the base path."""
|
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):
|
def set_lora_info(self, *args):
|
||||||
pass
|
pass
|
||||||
@@ -482,14 +490,22 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
|
|||||||
)
|
)
|
||||||
return lora_output
|
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):
|
def forward(self, input_: torch.Tensor):
|
||||||
# duplicate the logic in ColumnParallelLinear
|
# 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
|
bias = self.base_layer.bias if not self.base_layer.skip_bias_add else None
|
||||||
output_parallel = self.base_layer.quant_method.apply(
|
output_parallel = self.base_layer.quant_method.apply(
|
||||||
self.base_layer, input_, bias
|
self.base_layer, input_, bias
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.lora_active:
|
if lora_active:
|
||||||
output_parallel = self.apply_lora(output_parallel, input_)
|
output_parallel = self.apply_lora(output_parallel, input_)
|
||||||
|
|
||||||
if self.base_layer.gather_output:
|
if self.base_layer.gather_output:
|
||||||
@@ -596,6 +612,12 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
|
|||||||
)
|
)
|
||||||
return lora_output
|
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):
|
def slice_lora_a_weights(self, A: torch.Tensor):
|
||||||
return A
|
return A
|
||||||
|
|
||||||
@@ -703,6 +725,10 @@ class QKVParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
|
|||||||
|
|
||||||
return lora_output
|
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):
|
def slice_lora_a_weights(self, A: torch.Tensor):
|
||||||
return A
|
return A
|
||||||
|
|
||||||
@@ -773,6 +799,10 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
|
|||||||
)
|
)
|
||||||
return lora_output
|
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):
|
def forward(self, input_: torch.Tensor, skip_all_reduce=False, forward_batch=None):
|
||||||
if self.base_layer.input_is_parallel:
|
if self.base_layer.input_is_parallel:
|
||||||
input_parallel = input_
|
input_parallel = input_
|
||||||
@@ -783,6 +813,10 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
|
|||||||
)
|
)
|
||||||
input_parallel = splitted_input[tp_rank].contiguous()
|
input_parallel = splitted_input[tp_rank].contiguous()
|
||||||
|
|
||||||
|
lora_active = self.lora_active
|
||||||
|
if lora_active:
|
||||||
|
self.start_lora_a_overlap(input_parallel)
|
||||||
|
|
||||||
bias_ = (
|
bias_ = (
|
||||||
None
|
None
|
||||||
if (self.base_layer.tp_rank > 0 or self.base_layer.skip_bias_add)
|
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
|
all_reduce = get_parallel().attn_tp_group.all_reduce
|
||||||
else:
|
else:
|
||||||
all_reduce = tensor_model_parallel_all_reduce
|
all_reduce = tensor_model_parallel_all_reduce
|
||||||
lora_active = self.lora_active
|
|
||||||
if lora_active and should_reduce:
|
if lora_active and should_reduce:
|
||||||
lora_a_output = self.lora_backend.run_lora_a_sgemm(
|
lora_a_output = self.lora_backend.run_lora_a_sgemm(
|
||||||
input_parallel, self.A_buffer
|
input_parallel, self.A_buffer
|
||||||
|
|||||||
@@ -17,7 +17,7 @@
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import re
|
import re
|
||||||
from typing import Dict, Iterable, List, Optional
|
from typing import Dict, Iterable, List, Optional, Sequence
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -431,6 +431,16 @@ class LoRAManager:
|
|||||||
self.lora_backend.reset_batch_state()
|
self.lora_backend.reset_batch_state()
|
||||||
|
|
||||||
def prepare_lora_batch(self, forward_batch: ForwardBatch):
|
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
|
# set up batch info shared by all lora modules
|
||||||
bs = forward_batch.batch_size
|
bs = forward_batch.batch_size
|
||||||
|
|
||||||
@@ -470,6 +480,39 @@ class LoRAManager:
|
|||||||
lora_ranks[wi] > 0 for wi in weight_indices
|
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):
|
def update_lora_info(self):
|
||||||
"""
|
"""
|
||||||
Update all LoRA modules to associate them with the latest memory buffer.
|
Update all LoRA modules to associate them with the latest memory buffer.
|
||||||
@@ -579,6 +622,10 @@ class LoRAManager:
|
|||||||
max_lora_rank=max_lora_rank,
|
max_lora_rank=max_lora_rank,
|
||||||
target_modules=target_modules,
|
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:
|
if self._experts_shared_outer_override is not None:
|
||||||
self.experts_shared_outer_loras = self._experts_shared_outer_override
|
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,
|
get_draft_recurrent_hidden_state_spec_from_config,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
|
from sglang.srt.speculative.uno_validation import validate_uno_request
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
DynamicGradMode,
|
DynamicGradMode,
|
||||||
configure_gc_logger,
|
configure_gc_logger,
|
||||||
@@ -2813,6 +2814,14 @@ class Scheduler(
|
|||||||
self._add_request_to_queue(req)
|
self._add_request_to_queue(req)
|
||||||
return
|
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 (
|
if (
|
||||||
req.return_sampling_mask
|
req.return_sampling_mask
|
||||||
and self.disaggregation_mode != DisaggregationMode.NULL
|
and self.disaggregation_mode != DisaggregationMode.NULL
|
||||||
@@ -4484,7 +4493,7 @@ class Scheduler(
|
|||||||
self.decode_moment_totals,
|
self.decode_moment_totals,
|
||||||
batch_size,
|
batch_size,
|
||||||
step_us,
|
step_us,
|
||||||
batch_size + result.num_correct_drafts,
|
result.get_num_generated_tokens(batch_size),
|
||||||
)
|
)
|
||||||
|
|
||||||
def maybe_send_health_check_signal(self):
|
def maybe_send_health_check_signal(self):
|
||||||
|
|||||||
@@ -78,6 +78,16 @@ if TYPE_CHECKING:
|
|||||||
logger = logging.getLogger(__name__)
|
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)
|
@dataclass(kw_only=True, slots=True, frozen=True)
|
||||||
class SchedulerBatchResultProcessor:
|
class SchedulerBatchResultProcessor:
|
||||||
is_generation: bool
|
is_generation: bool
|
||||||
@@ -711,8 +721,12 @@ class SchedulerBatchResultProcessor:
|
|||||||
|
|
||||||
next_token_ids = result.next_token_ids.tolist()
|
next_token_ids = result.next_token_ids.tolist()
|
||||||
accept_lens = result.accept_lens.tolist()
|
accept_lens = result.accept_lens.tolist()
|
||||||
result.num_correct_drafts = sum(accept_lens) - len(batch.reqs)
|
stride = _get_speculative_output_stride(result)
|
||||||
result.num_correct_drafts_per_req_cpu = [x - 1 for x in accept_lens]
|
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 = (
|
block_accept_lens = (
|
||||||
result.block_accept_lens.tolist()
|
result.block_accept_lens.tolist()
|
||||||
@@ -740,11 +754,6 @@ class SchedulerBatchResultProcessor:
|
|||||||
self.advance_grammar_fsm(result, batch)
|
self.advance_grammar_fsm(result, batch)
|
||||||
|
|
||||||
predict_tokens = []
|
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):
|
for i, req in enumerate(batch.reqs):
|
||||||
accept_tokens = next_token_ids[i * stride : i * stride + accept_lens[i]]
|
accept_tokens = next_token_ids[i * stride : i * stride + accept_lens[i]]
|
||||||
|
|
||||||
@@ -853,8 +862,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
if result.accept_lens is None:
|
if result.accept_lens is None:
|
||||||
return
|
return
|
||||||
accept_lens = result.accept_lens.tolist()
|
accept_lens = result.accept_lens.tolist()
|
||||||
stride = result.speculative_num_draft_tokens
|
stride = _get_speculative_output_stride(result)
|
||||||
assert stride is not None, "spec-v2 result missing speculative_num_draft_tokens"
|
|
||||||
retained = [None] * len(batch.reqs)
|
retained = [None] * len(batch.reqs)
|
||||||
for i, req in enumerate(batch.reqs):
|
for i, req in enumerate(batch.reqs):
|
||||||
if req.grammar is None or req.is_retracted or req.finished():
|
if req.grammar is None or req.is_retracted or req.finished():
|
||||||
@@ -907,11 +915,14 @@ class SchedulerBatchResultProcessor:
|
|||||||
next_token_ids=next_token_ids,
|
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():
|
if not batch.spec_algorithm.is_none():
|
||||||
self.metrics_reporter.update_spec_metrics(
|
self.metrics_reporter.update_spec_metrics(
|
||||||
batch.batch_size(),
|
batch_size,
|
||||||
result.num_correct_drafts,
|
result.num_correct_drafts,
|
||||||
|
num_accept_tokens=num_generated_tokens,
|
||||||
num_block_accept_tokens=result.num_block_accept_tokens,
|
num_block_accept_tokens=result.num_block_accept_tokens,
|
||||||
num_cap_tokens=result.num_cap_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:
|
if req.return_hidden_states and logits_output.hidden_states is not None:
|
||||||
# hidden_states is [bs * stride, hidden_dim], one row per emitted
|
# hidden_states is [bs * stride, hidden_dim], one row per emitted
|
||||||
# token; stride = speculative_num_draft_tokens for spec, 1 for non-spec.
|
# token; speculative workers record their padded row width.
|
||||||
stride = result.speculative_num_draft_tokens or 1
|
stride = _get_speculative_output_stride(result) if is_spec else 1
|
||||||
accept_len = len(next_token_id)
|
accept_len = len(next_token_id)
|
||||||
start = i * stride
|
start = i * stride
|
||||||
self._append_decode_hidden_states(
|
self._append_decode_hidden_states(
|
||||||
@@ -1016,7 +1027,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
self.metrics_reporter.report_decode_stats(
|
self.metrics_reporter.report_decode_stats(
|
||||||
can_run_cuda_graph,
|
can_run_cuda_graph,
|
||||||
running_batch=batch,
|
running_batch=batch,
|
||||||
num_correct_drafts=result.num_correct_drafts,
|
num_generated_tokens=num_generated_tokens,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _normalize_decode_outputs(
|
def _normalize_decode_outputs(
|
||||||
|
|||||||
@@ -175,9 +175,10 @@ class SchedulerMetricsReporter:
|
|||||||
}.get(getattr(self.scheduler, "device", ""), "cuda graph")
|
}.get(getattr(self.scheduler, "device", ""), "cuda graph")
|
||||||
|
|
||||||
# Cumulative spec-decoding counters (reset every decode_log_interval).
|
# Cumulative spec-decoding counters (reset every decode_log_interval).
|
||||||
# Each update adds (num_correct_drafts + bs, bs).
|
# `*_accept_tokens` includes accepted drafts and non-draft output tokens;
|
||||||
# `*_accept_tokens` = drafts + bonus; `*_correct_drafts` = drafts-only.
|
# `*_correct_drafts` counts accepted draft proposals only.
|
||||||
self.spec_num_accept_tokens = 0 # per-log-interval
|
self.spec_num_accept_tokens = 0 # per-log-interval
|
||||||
|
self.spec_num_correct_drafts = 0
|
||||||
self.spec_num_forward_ct = 0
|
self.spec_num_forward_ct = 0
|
||||||
self.spec_total_num_accept_tokens = 0 # lifetime
|
self.spec_total_num_accept_tokens = 0 # lifetime
|
||||||
self.spec_total_num_forward_ct = 0
|
self.spec_total_num_forward_ct = 0
|
||||||
@@ -398,17 +399,16 @@ class SchedulerMetricsReporter:
|
|||||||
self,
|
self,
|
||||||
bs: int,
|
bs: int,
|
||||||
num_correct_drafts: int,
|
num_correct_drafts: int,
|
||||||
|
num_accept_tokens: int,
|
||||||
num_block_accept_tokens: int = 0,
|
num_block_accept_tokens: int = 0,
|
||||||
num_cap_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_forward_ct += bs
|
||||||
self.spec_num_block_accept_tokens += num_block_accept_tokens
|
self.spec_num_block_accept_tokens += num_block_accept_tokens
|
||||||
self.spec_num_cap_tokens += num_cap_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:
|
def _init_estimated_perf_constants(self) -> None:
|
||||||
model_config = self.scheduler.model_config
|
model_config = self.scheduler.model_config
|
||||||
hf_text_config = model_config.hf_text_config
|
hf_text_config = model_config.hf_text_config
|
||||||
@@ -572,6 +572,7 @@ class SchedulerMetricsReporter:
|
|||||||
self.forward_ct_decode = 0
|
self.forward_ct_decode = 0
|
||||||
self.num_generated_tokens = 0
|
self.num_generated_tokens = 0
|
||||||
self.spec_num_accept_tokens = 0
|
self.spec_num_accept_tokens = 0
|
||||||
|
self.spec_num_correct_drafts = 0
|
||||||
self.spec_num_forward_ct = 0
|
self.spec_num_forward_ct = 0
|
||||||
self.spec_total_num_accept_tokens = 0
|
self.spec_total_num_accept_tokens = 0
|
||||||
self.spec_total_num_forward_ct = 0
|
self.spec_total_num_forward_ct = 0
|
||||||
@@ -757,13 +758,13 @@ class SchedulerMetricsReporter:
|
|||||||
self,
|
self,
|
||||||
can_run_cuda_graph: bool,
|
can_run_cuda_graph: bool,
|
||||||
running_batch: ScheduleBatch = None,
|
running_batch: ScheduleBatch = None,
|
||||||
num_correct_drafts: int = 0,
|
num_generated_tokens: int = 0,
|
||||||
):
|
):
|
||||||
batch = running_batch or self.scheduler.running_batch
|
batch = running_batch or self.scheduler.running_batch
|
||||||
|
|
||||||
# Every-iteration work: realtime token counting + status logger
|
# Every-iteration work: realtime token counting + status logger
|
||||||
if self.current_scheduler_metrics_enabled:
|
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(
|
self.metrics_collector.increment_realtime_tokens(
|
||||||
# TODO unify this w/ the bumping logic in `Scheduler.num_generated_tokens` accumulator
|
# TODO unify this w/ the bumping logic in `Scheduler.num_generated_tokens` accumulator
|
||||||
decode_tokens=decode_tokens,
|
decode_tokens=decode_tokens,
|
||||||
@@ -826,7 +827,7 @@ class SchedulerMetricsReporter:
|
|||||||
spec_block_accept_length = 0
|
spec_block_accept_length = 0
|
||||||
else:
|
else:
|
||||||
spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct
|
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:
|
if get_spec().speculative_num_draft_tokens:
|
||||||
draft_per_round = get_spec().speculative_num_draft_tokens - 1
|
draft_per_round = get_spec().speculative_num_draft_tokens - 1
|
||||||
else:
|
else:
|
||||||
@@ -853,7 +854,8 @@ class SchedulerMetricsReporter:
|
|||||||
)
|
)
|
||||||
self.spec_total_num_accept_tokens += self.spec_num_accept_tokens
|
self.spec_total_num_accept_tokens += self.spec_num_accept_tokens
|
||||||
self.spec_total_num_forward_ct += self.spec_num_forward_ct
|
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_block_accept_tokens = 0
|
||||||
self.spec_num_cap_tokens = 0
|
self.spec_num_cap_tokens = 0
|
||||||
msg += f"accept len: {spec_accept_length:.2f}, accept rate: {spec_accept_rate:.2f}, "
|
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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||||
from sglang.srt.sampling.sampling_observer import HostAuxiliaryOutput
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -71,6 +71,12 @@ class GenerationBatchResult:
|
|||||||
delay_sample_func: Optional[callable] = None
|
delay_sample_func: Optional[callable] = None
|
||||||
future_indices: Optional[torch.Tensor] = None
|
future_indices: Optional[torch.Tensor] = None
|
||||||
speculative_num_draft_tokens: Optional[int] = 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
|
# Grammar FSM advance memoization (spec-v2 overlap). advance_grammar_fsm sets
|
||||||
# these once — eagerly via the scheduler's grammar barrier inside verify(), or
|
# 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
|
new_seq_lens: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# relay path: forward stream -> next step forward
|
# 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
|
# 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
|
# 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)."""
|
this rank/split (a non-last PP rank or a non-final prefill split)."""
|
||||||
return isinstance(self.next_token_ids, torch.Tensor)
|
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")
|
@torch.profiler.record_function("copy_result_to_cpu")
|
||||||
def copy_to_cpu(self, return_logprob: bool, return_hidden_states: bool = True):
|
def copy_to_cpu(self, return_logprob: bool, return_hidden_states: bool = True):
|
||||||
"""Copy tensors to CPU in overlap scheduling.
|
"""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
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
|
|
||||||
spec_algo = SpeculativeAlgorithm.from_string(spec.speculative_algorithm)
|
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():
|
if page_size == 1 or spec_topk == 1 or not spec_algo.has_draft_kv():
|
||||||
return max(spec_steps * spec_topk, spec_tokens)
|
return max(spec_steps * spec_topk, spec_tokens)
|
||||||
else:
|
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
|
# FIXME(lsyin): temporary fix for the context length issue under spec decoding
|
||||||
extra = 4 + (max_speculative_num_draft_tokens() or 0)
|
extra = 4 + (max_speculative_num_draft_tokens() or 0)
|
||||||
page_size = get_alloc_page_size()
|
page_size = get_alloc_page_size()
|
||||||
if get_spec().speculative_algorithm is not None and page_size > 1:
|
spec_algorithm = get_spec().speculative_algorithm
|
||||||
# kv_allocated_len is page-aligned (eagle_prepare_for_decode), so near
|
if spec_algorithm is not None:
|
||||||
# the context limit the aligned reserve can overshoot by page_size - 1;
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
# without the headroom the row write silently lands in the neighbor row.
|
|
||||||
|
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)
|
extra = max(extra, get_alloc_reserve_per_decode() + page_size - 1)
|
||||||
return extra
|
return extra
|
||||||
|
|||||||
@@ -777,6 +777,12 @@ class ModelRunner:
|
|||||||
self.apply_torch_tp()
|
self.apply_torch_tp()
|
||||||
|
|
||||||
def maybe_init_lora_manager(self):
|
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.
|
# Adapters apply to the target model only; the draft runs unadapted.
|
||||||
if get_lora().enable_lora and not self.is_draft_worker:
|
if get_lora().enable_lora and not self.is_draft_worker:
|
||||||
self.init_lora_manager()
|
self.init_lora_manager()
|
||||||
@@ -1430,6 +1436,13 @@ class ModelRunner:
|
|||||||
|
|
||||||
Subclasses can override this to install specialized decode graph runners.
|
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 (
|
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
||||||
DecodeCudaGraphRunner,
|
DecodeCudaGraphRunner,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2067,7 +2067,12 @@ class ServerArgs:
|
|||||||
# -------------------------------------------------------------------------
|
# -------------------------------------------------------------------------
|
||||||
speculative_algorithm: A[
|
speculative_algorithm: A[
|
||||||
Optional[str],
|
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"),
|
NS("spec"),
|
||||||
] = None
|
] = None
|
||||||
speculative_draft_model_path: A[
|
speculative_draft_model_path: A[
|
||||||
|
|||||||
@@ -663,11 +663,30 @@ def _verify_coins(
|
|||||||
return coins, coins_for_final_sampling
|
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(
|
def eagle_sample(
|
||||||
verify_input: EagleVerifyInput,
|
verify_input: EagleVerifyInput,
|
||||||
batch: ScheduleBatch,
|
batch: ScheduleBatch,
|
||||||
logits_output: LogitsProcessorOutput,
|
logits_output: LogitsProcessorOutput,
|
||||||
grammar_mask: Optional[GrammarMask] = None,
|
grammar_mask: Optional[GrammarMask] = None,
|
||||||
|
uno_target_max_top_k: Optional[int] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Verify and find accepted tokens based on logits output and batch
|
Verify and find accepted tokens based on logits output and batch
|
||||||
@@ -760,6 +779,42 @@ def eagle_sample(
|
|||||||
# different number of drafts, desynchronize the committed seq_lens, and
|
# different number of drafts, desynchronize the committed seq_lens, and
|
||||||
# deadlock the next TP collective. Broadcast from rank 0 to ensure
|
# deadlock the next TP collective. Broadcast from rank 0 to ensure
|
||||||
# consistency, the same way the sampling branch below does.
|
# consistency, the same way the sampling branch below does.
|
||||||
|
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)
|
||||||
|
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 = (
|
tp_group = (
|
||||||
get_parallel().attn_tp_group
|
get_parallel().attn_tp_group
|
||||||
if is_dp_attention_enabled()
|
if is_dp_attention_enabled()
|
||||||
|
|||||||
@@ -472,6 +472,7 @@ def run_eagle_verify(
|
|||||||
metadata_ready_pre_pad: bool,
|
metadata_ready_pre_pad: bool,
|
||||||
finalize_tree_path: bool,
|
finalize_tree_path: bool,
|
||||||
grammar_barrier=None,
|
grammar_barrier=None,
|
||||||
|
uno_target_max_top_k: Optional[int] = None,
|
||||||
) -> GenerationBatchResult:
|
) -> GenerationBatchResult:
|
||||||
"""Shared verify step: target-verify forward, sampling, acceptance bookkeeping.
|
"""Shared verify step: target-verify forward, sampling, acceptance bookkeeping.
|
||||||
|
|
||||||
@@ -584,7 +585,13 @@ def run_eagle_verify(
|
|||||||
predict,
|
predict,
|
||||||
accept_lens,
|
accept_lens,
|
||||||
accept_index,
|
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
|
new_seq_lens = batch.seq_lens + accept_lens
|
||||||
clear_unaccepted_c128 = getattr(
|
clear_unaccepted_c128 = getattr(
|
||||||
token_to_kv_pool_allocator.get_kvcache(),
|
token_to_kv_pool_allocator.get_kvcache(),
|
||||||
|
|||||||
@@ -37,6 +37,7 @@ class SpeculativeAlgorithm(Enum):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
DFLASH = auto()
|
DFLASH = auto()
|
||||||
|
UNO = auto()
|
||||||
DSPARK = auto()
|
DSPARK = auto()
|
||||||
EAGLE = auto()
|
EAGLE = auto()
|
||||||
EAGLE3 = auto()
|
EAGLE3 = auto()
|
||||||
@@ -114,6 +115,9 @@ class SpeculativeAlgorithm(Enum):
|
|||||||
def is_dflash(self) -> bool:
|
def is_dflash(self) -> bool:
|
||||||
return self == SpeculativeAlgorithm.DFLASH
|
return self == SpeculativeAlgorithm.DFLASH
|
||||||
|
|
||||||
|
def is_uno(self) -> bool:
|
||||||
|
return self == SpeculativeAlgorithm.UNO
|
||||||
|
|
||||||
def is_dspark(self) -> bool:
|
def is_dspark(self) -> bool:
|
||||||
return self == SpeculativeAlgorithm.DSPARK
|
return self == SpeculativeAlgorithm.DSPARK
|
||||||
|
|
||||||
@@ -220,6 +224,7 @@ class SpeculativeAlgorithm(Enum):
|
|||||||
_handle_eagle_family,
|
_handle_eagle_family,
|
||||||
_handle_frozen_kv_mtp,
|
_handle_frozen_kv_mtp,
|
||||||
_handle_ngram,
|
_handle_ngram,
|
||||||
|
_handle_uno,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Validate for every algorithm at startup: the metrics paths read the
|
# Validate for every algorithm at startup: the metrics paths read the
|
||||||
@@ -230,6 +235,8 @@ class SpeculativeAlgorithm(Enum):
|
|||||||
|
|
||||||
if self.is_dflash():
|
if self.is_dflash():
|
||||||
_handle_dflash(server_args)
|
_handle_dflash(server_args)
|
||||||
|
elif self.is_uno():
|
||||||
|
_handle_uno(server_args)
|
||||||
elif self.is_dspark():
|
elif self.is_dspark():
|
||||||
_handle_dspark(server_args)
|
_handle_dspark(server_args)
|
||||||
elif self.is_frozen_kv_mtp():
|
elif self.is_frozen_kv_mtp():
|
||||||
@@ -304,6 +311,11 @@ class SpeculativeAlgorithm(Enum):
|
|||||||
|
|
||||||
return DFlashWorkerV2
|
return DFlashWorkerV2
|
||||||
|
|
||||||
|
if self.is_uno():
|
||||||
|
from sglang.srt.speculative.uno_worker_v2 import UnoWorkerV2
|
||||||
|
|
||||||
|
return UnoWorkerV2
|
||||||
|
|
||||||
if self.is_dspark():
|
if self.is_dspark():
|
||||||
from sglang.srt.speculative.dspark_components.dspark_worker_v2 import (
|
from sglang.srt.speculative.dspark_components.dspark_worker_v2 import (
|
||||||
DSparkWorkerV2,
|
DSparkWorkerV2,
|
||||||
@@ -356,6 +368,9 @@ class SpecInputType(IntEnum):
|
|||||||
DFLASH_DRAFT = auto()
|
DFLASH_DRAFT = auto()
|
||||||
DFLASH_VERIFY = auto()
|
DFLASH_VERIFY = auto()
|
||||||
NGRAM_VERIFY = auto()
|
NGRAM_VERIFY = auto()
|
||||||
|
UNO_STATE = auto()
|
||||||
|
UNO_DRAFT = auto()
|
||||||
|
UNO_VERIFY = auto()
|
||||||
|
|
||||||
|
|
||||||
class SpecInput(ABC):
|
class SpecInput(ABC):
|
||||||
@@ -393,6 +408,7 @@ class SpecInput(ABC):
|
|||||||
SpecInputType.EAGLE_DRAFT_EXTEND,
|
SpecInputType.EAGLE_DRAFT_EXTEND,
|
||||||
SpecInputType.FROZEN_KV_MTP_DRAFT,
|
SpecInputType.FROZEN_KV_MTP_DRAFT,
|
||||||
SpecInputType.DFLASH_DRAFT,
|
SpecInputType.DFLASH_DRAFT,
|
||||||
|
SpecInputType.UNO_DRAFT,
|
||||||
}
|
}
|
||||||
|
|
||||||
def is_verify_input(self) -> bool:
|
def is_verify_input(self) -> bool:
|
||||||
@@ -401,6 +417,7 @@ class SpecInput(ABC):
|
|||||||
SpecInputType.FROZEN_KV_MTP_VERIFY,
|
SpecInputType.FROZEN_KV_MTP_VERIFY,
|
||||||
SpecInputType.DFLASH_VERIFY,
|
SpecInputType.DFLASH_VERIFY,
|
||||||
SpecInputType.NGRAM_VERIFY,
|
SpecInputType.NGRAM_VERIFY,
|
||||||
|
SpecInputType.UNO_VERIFY,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -82,6 +82,9 @@ class CustomSpecAlgo:
|
|||||||
def is_dflash(self) -> bool:
|
def is_dflash(self) -> bool:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
def is_uno(self) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
def is_dspark(self) -> bool:
|
def is_dspark(self) -> bool:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|||||||
@@ -1051,6 +1051,12 @@ def spec_prepare_for_decode(batch: ScheduleBatch) -> None:
|
|||||||
)
|
)
|
||||||
if batch.spec_algorithm.is_dflash_family():
|
if batch.spec_algorithm.is_dflash_family():
|
||||||
batch.spec_info.prepare_for_decode(batch)
|
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:
|
else:
|
||||||
from sglang.srt.speculative.eagle_utils import eagle_prepare_for_decode
|
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 (
|
from sglang.srt.managers.scheduler_components.batch_result_processor import (
|
||||||
SchedulerBatchResultProcessor,
|
SchedulerBatchResultProcessor,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.utils import GenerationBatchResult
|
||||||
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
|
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
|
||||||
from sglang.srt.runtime_context import get_context
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
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:]
|
second_step = torch.arange(16, dtype=torch.float32).view(8, 2)[4:]
|
||||||
|
|
||||||
def result(hidden_states):
|
def result(hidden_states):
|
||||||
return SimpleNamespace(
|
return GenerationBatchResult(
|
||||||
copy_done=None,
|
|
||||||
auxiliary_host_output=None,
|
|
||||||
routed_experts_output=None,
|
|
||||||
indexer_topk_output=None,
|
|
||||||
logits_output=SimpleNamespace(hidden_states=hidden_states),
|
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,
|
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 (
|
from sglang.srt.managers.scheduler_components.batch_result_processor import (
|
||||||
SchedulerBatchResultProcessor,
|
SchedulerBatchResultProcessor,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.utils import GenerationBatchResult
|
||||||
from sglang.srt.runtime_context import get_context
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
@@ -78,17 +79,9 @@ def _make_processor() -> SchedulerBatchResultProcessor:
|
|||||||
|
|
||||||
|
|
||||||
def _make_result():
|
def _make_result():
|
||||||
return SimpleNamespace(
|
return GenerationBatchResult(
|
||||||
copy_done=None,
|
|
||||||
auxiliary_host_output=None,
|
|
||||||
routed_experts_output=None,
|
|
||||||
indexer_topk_output=None,
|
|
||||||
logits_output=SimpleNamespace(hidden_states=None, customized_info=None),
|
logits_output=SimpleNamespace(hidden_states=None, customized_info=None),
|
||||||
next_token_ids=[4],
|
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,
|
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 (
|
from sglang.srt.managers.scheduler_components.batch_result_processor import (
|
||||||
SchedulerBatchResultProcessor,
|
SchedulerBatchResultProcessor,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.utils import GenerationBatchResult
|
||||||
from sglang.srt.sampling.sampling_params import (
|
from sglang.srt.sampling.sampling_params import (
|
||||||
REQUEST_REASONING_END_TOKEN_IDS_KEY,
|
REQUEST_REASONING_END_TOKEN_IDS_KEY,
|
||||||
SamplingParams,
|
SamplingParams,
|
||||||
@@ -101,16 +102,10 @@ def _make_req(terminate_after: int) -> Req:
|
|||||||
|
|
||||||
|
|
||||||
def _make_result(num_draft_tokens, accept_lens, flat_tokens):
|
def _make_result(num_draft_tokens, accept_lens, flat_tokens):
|
||||||
return SimpleNamespace(
|
return GenerationBatchResult(
|
||||||
next_token_ids=torch.tensor(flat_tokens, dtype=torch.long),
|
next_token_ids=torch.tensor(flat_tokens, dtype=torch.long),
|
||||||
accept_lens=torch.tensor(accept_lens, dtype=torch.long),
|
accept_lens=torch.tensor(accept_lens, dtype=torch.long),
|
||||||
speculative_num_draft_tokens=num_draft_tokens,
|
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 unittest.mock import patch
|
||||||
|
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
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.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -25,6 +26,7 @@ class TestDraftRunnerSkipsLoRA(CustomTestCase):
|
|||||||
runner = ModelRunner.__new__(ModelRunner)
|
runner = ModelRunner.__new__(ModelRunner)
|
||||||
runner.is_draft_worker = is_draft_worker
|
runner.is_draft_worker = is_draft_worker
|
||||||
runner.lora_manager = None
|
runner.lora_manager = None
|
||||||
|
runner.spec_algorithm = SpeculativeAlgorithm.NONE
|
||||||
with patch.object(ModelRunner, "init_lora_manager") as init_lora:
|
with patch.object(ModelRunner, "init_lora_manager") as init_lora:
|
||||||
with patch("sglang.srt.model_executor.model_runner.get_lora") as get_lora:
|
with patch("sglang.srt.model_executor.model_runner.get_lora") as get_lora:
|
||||||
get_lora.return_value.enable_lora = enable_lora
|
get_lora.return_value.enable_lora = enable_lora
|
||||||
|
|||||||
@@ -45,6 +45,10 @@ _DFLASH_DECODE = (
|
|||||||
"speculative/dflash_info_v2.py",
|
"speculative/dflash_info_v2.py",
|
||||||
"DFlashDraftInputV2.prepare_for_decode",
|
"DFlashDraftInputV2.prepare_for_decode",
|
||||||
)
|
)
|
||||||
|
_UNO_DECODE = (
|
||||||
|
"speculative/uno_info.py",
|
||||||
|
"UnoDraftInput.prepare_for_decode",
|
||||||
|
)
|
||||||
_RESOLVE = (
|
_RESOLVE = (
|
||||||
"managers/scheduler_components/batch_result_processor.py",
|
"managers/scheduler_components/batch_result_processor.py",
|
||||||
"SchedulerBatchResultProcessor._resolve_spec_v2_tokens",
|
"SchedulerBatchResultProcessor._resolve_spec_v2_tokens",
|
||||||
@@ -71,6 +75,8 @@ _OWNER_SITES = {
|
|||||||
# one of these two owners for each speculative decode iteration.
|
# one of these two owners for each speculative decode iteration.
|
||||||
(*_DFLASH_DECODE, "decode_batch_idx"): 1,
|
(*_DFLASH_DECODE, "decode_batch_idx"): 1,
|
||||||
(*_DFLASH_DECODE, "evict"): 1,
|
(*_DFLASH_DECODE, "evict"): 1,
|
||||||
|
(*_UNO_DECODE, "decode_batch_idx"): 1,
|
||||||
|
(*_UNO_DECODE, "evict"): 1,
|
||||||
(
|
(
|
||||||
"mem_cache/allocation.py",
|
"mem_cache/allocation.py",
|
||||||
"alloc_for_spec_decode",
|
"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