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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions .github/workflows/tests.yml
Original file line number Diff line number Diff line change
Expand Up @@ -42,3 +42,17 @@ jobs:
PY
- name: Regression suite (locked recipes can't silently drift)
run: pytest tests/ -v

routing-cpu:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
- name: Install CPU-only routing test dependencies
run: |
pip install -e . pytest
pip install torch --index-url https://download.pytorch.org/whl/cpu
- name: Core and routing tensor fixtures without vLLM or a GPU
run: pytest tests/ --ignore=tests/studio -v
412 changes: 344 additions & 68 deletions experiments/vllm_decode_routing_hook.py

Large diffs are not rendered by default.

115 changes: 115 additions & 0 deletions experiments/vllm_routing_check.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,115 @@
#!/usr/bin/env python3
"""Check routing capture on fixed synthetic requests, not model quality or speed.

Run once with the hook installed on worker PYTHONPATH and --capture-dir pointing
to the workers' shared dump directory. Run again without the hook and without
--capture-dir; compare output_sha256 for both batches. No request text, generated
text, model paths, hostnames or worker filenames are included in the summary.
"""
import argparse
import hashlib
import json
from pathlib import Path

PROMPTS = [
"Explain why a database transaction needs isolation. Give a concrete example.",
"Write a Python function that merges two sorted lists without changing its inputs.",
"A gardener has a rectangular plot measuring 12 by 8 metres. Explain how to calculate the perimeter and area.",
"Tell a short story about a lighthouse keeper repairing a broken radio.",
]


def summarize_outputs(outputs):
sequences = [list(x.outputs[0].token_ids) for x in outputs]
if any(not ids for ids in sequences):
raise ValueError("expected at least one generated token per request")
return {"requests": len(outputs),
"prompt_tokens": sum(len(x.prompt_token_ids) for x in outputs),
"output_tokens": sum(map(len, sequences)),
# The last generated token is returned, not fed back to the router.
"expected_decode_tokens": sum(len(ids) - 1 for ids in sequences),
"output_sha256": hashlib.sha256(json.dumps(sequences, separators=(",", ":")).encode()).hexdigest()}


def check_capture(directory, workers, expected):
if not workers:
raise ValueError("no worker reports")
reports = []
for worker in workers:
name = worker["file"]
if Path(name).name != name or not name.startswith("routing_") or not name.endswith(".json"):
raise ValueError("expected a routing capture basename")
data = json.loads((Path(directory) / name).read_text())
if worker["errors"] or data["meta"]["errors"]:
raise ValueError("worker reported recording errors")
if not data["layers"] or len(data["layers"]) != worker["layers"]:
raise ValueError("worker has missing routing layers")
phases = {"prefill": expected["prompt_tokens"],
"decode": expected["expected_decode_tokens"], "unknown": 0}
for layer in data["layers"].values():
if any(layer[phase]["tokens"] != count for phase, count in phases.items()):
raise ValueError("per-layer token counts disagree with the requests")
reports.append({"layers": len(data["layers"]), "calls": worker["calls"],
"errors": 0, "tokens_per_layer": phases})
return reports


def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model", required=True)
parser.add_argument("--capture-dir", help="shared worker dump directory; omit for an unhooked control")
parser.add_argument("--output", type=Path, required=True, help="aggregate JSON summary")
parser.add_argument("--moe-backend", help="optional runtime-specific backend, e.g. marlin")
parser.add_argument("--serial", action="store_true", help="one request at a time to control batch shape")
parser.add_argument("--same-engine-control", action="store_true", help="remove the hook and repeat twice in the same engine")
args = parser.parse_args()
if args.same_engine_control and not args.capture_dir:
parser.error("--same-engine-control requires --capture-dir")
# Keep vLLM optional for CPU fixtures and --help.
import vllm
from vllm import LLM, SamplingParams
options = dict(model=args.model, trust_remote_code=False, tensor_parallel_size=1,
enforce_eager=True, max_model_len=2048, max_num_seqs=1 if args.serial else 4,
gpu_memory_utilization=0.55, enable_prefix_caching=False,
enable_chunked_prefill=True, compilation_config={"mode": 0}, seed=0,
kernel_config={"enable_flashinfer_autotune": False})
if args.moe_backend:
options["moe_backend"] = args.moe_backend
if args.capture_dir:
options["worker_extension_cls"] = "vllm_decode_routing_hook.RoutingWorkerExtension"
llm = LLM(**options)
def generate(max_tokens):
params = SamplingParams(temperature=0, max_tokens=max_tokens, ignore_eos=True)
if args.serial:
return [llm.generate([prompt], params)[0] for prompt in PROMPTS]
return llm.generate(PROMPTS, params)
batches = []
for max_tokens in (1, 48):
if args.capture_dir:
named = llm.collective_rpc("pollard_start_capture")
if any(x["named_routers"] == 0 for x in named):
raise ValueError("no named MoE routers in this model")
outputs = generate(max_tokens)
result = summarize_outputs(outputs)
result["max_tokens"] = max_tokens
if args.capture_dir:
result["workers"] = check_capture(args.capture_dir, llm.collective_rpc("pollard_flush_capture"), result)
batches.append(result)
summary = {"schema_version": 1, "vllm": vllm.__version__, "capture": bool(args.capture_dir),
"flashinfer_autotune": False,
"serial": args.serial,
"quality_benchmark": False, "performance_benchmark": False, "batches": batches}
if args.same_engine_control:
summary["hook_removal"] = llm.collective_rpc("pollard_remove_capture")
controls = [[summarize_outputs(generate(n)) for n in (1, 48)] for _ in range(2)]
summary["same_engine_controls"] = controls
summary["control_repeats_match"] = all(x["output_sha256"] == y["output_sha256"]
for x, y in zip(*controls))
summary["capture_matches_controls"] = all(x["output_sha256"] == y["output_sha256"]
for control in controls for x, y in zip(batches, control))
args.output.write_text(json.dumps(summary, indent=2) + "\n")
print(json.dumps(summary))


if __name__ == "__main__":
main()
201 changes: 201 additions & 0 deletions experiments/vllm_routing_profiles.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,201 @@
#!/usr/bin/env python3
"""Compare routing on original synthetic workloads, not quality or throughput.

Uses the routing-check worker extension. Raw snapshots are kept separately from
the aggregate report, without request text or generated text. Synthetic domain
differences do not establish behavior on real traffic or a quantization benefit.
"""
import argparse
import hashlib
import itertools
import json
import math
from pathlib import Path

from vllm_routing_check import check_capture, summarize_outputs

DOMAINS = ("code", "reasoning", "general")
CONTEXTS = {"short": 256, "long": 1536}
KS = (8, 16, 32, 64, 128)


def workload_text(domain, variant, paragraphs):
"""Original, deterministic public fixtures. No copied or private corpus."""
if domain not in DOMAINS or variant not in range(4) or paragraphs < 1:
raise ValueError("invalid workload settings")
sections = []
for i in range(paragraphs):
n = i + variant * 101
if domain == "code":
sections.append(f"def update_{n}(records, limit={n % 19 + 2}):\n"
" kept = [r for r in records if r['active']]\n"
" return sorted(kept, key=lambda r: r['score'])[:limit]\n"
f"# Review function {n} for empty inputs, ties, and mutation.\n")
elif domain == "reasoning":
sections.append(f"Warehouse {n} starts with {n % 29 + 14} crates. "
f"It receives {n % 11 + 3} deliveries of {n % 7 + 2} crates "
f"and sends out {n % 13 + 5} crates. Each crate contains "
"six sealed boxes. Track the remaining crates and boxes, "
"showing each step and stating any assumptions.\n")
else:
sections.append(f"At station {n}, the volunteer archivists catalogued "
"letters and photographs from the old coastal railway. "
"A visitor asked how the weather affected daily journeys. "
"The curator compared passenger diaries with maintenance "
"notes and explained why accounts sometimes disagreed.\n")
instruction = {"code": "Review the code and suggest a safe improvement.",
"reasoning": "Explain the arithmetic for the final warehouse.",
"general": "Summarize the archive account in plain language."}[domain]
return "\n".join(sections) + "\n" + instruction


def workload_tokens(tokenizer, domain, target, variant):
"""Build a token-length fixture without depending on a chat template.

Truncation deliberately makes this a controlled token-stream experiment,
not a scored instruction-following task. Actual encoded IDs are hashed.
"""
if not isinstance(target, int) or isinstance(target, bool) or target < 1:
raise ValueError("target must be a positive token count")
sections = 1
while True:
ids = tokenizer.encode(workload_text(domain, variant, sections), add_special_tokens=False)
if len(ids) >= target:
return ids[:target]
sections *= 2


def normalized(values):
if not values or any(not math.isfinite(v) or v < 0 for v in values):
raise ValueError("expected finite nonnegative routing values")
total = sum(values)
if total <= 0:
raise ValueError("routing histogram is empty")
return [v / total for v in values]


def coverage(values, indices):
p = normalized(values)
return sum(p[i] for i in indices)


def top_indices(values, k):
return sorted(range(len(values)), key=lambda i: (-values[i], i))[:k]


def js_bits(left, right):
if len(left) != len(right):
raise ValueError("expert widths disagree")
p, q = normalized(left), normalized(right)
middle = [(a + b) / 2 for a, b in zip(p, q)]
return sum(0.5 * v * math.log2(v / m) for dist in (p, q)
for v, m in zip(dist, middle) if v > 0)


def describe_layers(data):
results = []
# Router module names are model identifiers, not host filesystem paths.
for ordinal, (name, layer) in enumerate(sorted(data["layers"].items())):
row = {"layer": ordinal, "router_name": name, "experts": layer["experts"]}
for phase in ("prefill", "decode"):
h = layer[phase]
row[phase] = {"tokens": h["tokens"],
"count": list(h["count"]), "mass": list(h["mass"]),
"selection_coverage": {str(k): coverage(h["count"], top_indices(h["count"], k)) for k in KS},
"weight_coverage": {str(k): coverage(h["mass"], top_indices(h["mass"], k)) for k in KS}}
results.append(row)
return results


def compare_layers(source, target):
if source["layers"].keys() != target["layers"].keys():
raise ValueError("capture layer names disagree")
comparisons = []
for ordinal, name in enumerate(sorted(source["layers"])):
a, b = source["layers"][name], target["layers"][name]
if a["experts"] != b["experts"]:
raise ValueError("capture expert widths disagree")
row = {"layer": ordinal}
for phase in ("prefill", "decode"):
row[phase] = {"selection_js_bits": js_bits(a[phase]["count"], b[phase]["count"]),
"source_cache_target_selection_coverage": {
str(k): coverage(b[phase]["count"], top_indices(a[phase]["count"], k)) for k in KS}}
comparisons.append(row)
return comparisons


def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model", required=True)
parser.add_argument("--capture-dir", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--moe-backend")
parser.add_argument("--batched", action="store_true", help="four concurrent requests instead of serial requests")
parser.add_argument("--max-tokens", type=int, default=32)
args = parser.parse_args()
if not 2 <= args.max_tokens <= 128:
parser.error("--max-tokens must be between 2 and 128")
import vllm
from vllm import LLM, SamplingParams
options = dict(model=args.model, trust_remote_code=False, tensor_parallel_size=1,
enforce_eager=True, max_model_len=2048, max_num_seqs=4 if args.batched else 1,
max_num_batched_tokens=1024, gpu_memory_utilization=0.55,
enable_prefix_caching=False, enable_chunked_prefill=True,
compilation_config={"mode": 0}, seed=0,
kernel_config={"enable_flashinfer_autotune": False},
worker_extension_cls="vllm_decode_routing_hook.RoutingWorkerExtension")
if args.moe_backend:
options["moe_backend"] = args.moe_backend
llm = LLM(**options)
tokenizer = llm.get_tokenizer()
params = SamplingParams(temperature=0, max_tokens=args.max_tokens, ignore_eos=True)
profiles, captured = [], {}
requests = {}
for domain, (context, length) in itertools.product(DOMAINS, CONTEXTS.items()):
key = f"{domain}-{context}"
ids = [workload_tokens(tokenizer, domain, length, variant) for variant in range(4)]
requests[key] = [{"prompt_token_ids": x} for x in ids]
def generate(key):
if args.batched:
return llm.generate(requests[key], params)
return [llm.generate([x], params)[0] for x in requests[key]]
raw_dir = args.capture_dir / "snapshots"
raw_dir.mkdir(exist_ok=True)
for key in requests:
workers = llm.collective_rpc("pollard_start_capture")
if any(w["named_routers"] == 0 for w in workers):
raise ValueError("no named routers")
result = summarize_outputs(generate(key))
reports = llm.collective_rpc("pollard_flush_capture")
result["workers"] = check_capture(args.capture_dir, reports, result)
if len(reports) != 1:
raise ValueError("this harness supports one worker; do not merge ranks implicitly")
data = json.loads((args.capture_dir / reports[0]["file"]).read_text())
# reset() deletes the worker's previous snapshot, so preserve each profile first.
(raw_dir / f"{key}.json").write_text(json.dumps(data) + "\n")
captured[key] = data
result.update(profile=key, prompt_ids_sha256=hashlib.sha256(
json.dumps(requests[key], separators=(",", ":")).encode()).hexdigest(),
mixed_router_calls=data["meta"].get("mixed_router_calls"),
layers=describe_layers(data))
profiles.append(result)
print(json.dumps({"completed": key, "prompt_tokens": result["prompt_tokens"],
"decode_tokens": result["expected_decode_tokens"]}), flush=True)
removal = llm.collective_rpc("pollard_remove_capture")
controls = [[summarize_outputs(generate(key)) for key in requests] for _ in range(2)]
summary = {"schema_version": 1, "vllm": vllm.__version__, "synthetic_workloads": True,
"quality_benchmark": False, "performance_benchmark": False,
"serial": not args.batched, "tensor_parallel_size": 1,
"max_num_batched_tokens": 1024, "max_tokens": args.max_tokens,
"flashinfer_autotune": False, "prefix_caching": False,
"profiles": profiles, "hook_removal": removal,
"control_repeats_match": all(a["output_sha256"] == b["output_sha256"] for a, b in zip(*controls)),
"capture_matches_controls": all(a["output_sha256"] == b["output_sha256"] for run in controls for a, b in zip(profiles, run)),
"same_engine_controls": controls,
"comparisons": [{"source": a, "target": b, "layers": compare_layers(captured[a], captured[b])}
for a, b in itertools.permutations(requests, 2)]}
args.output.write_text(json.dumps(summary, indent=2, allow_nan=False) + "\n")


if __name__ == "__main__":
main()
Loading
Loading