diff --git a/docs/docs.json b/docs/docs.json index 85fd6e39d..60707fd09 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -936,6 +936,7 @@ "docs/advanced_features/separate_reasoning", "docs/advanced_features/quantization", "docs/advanced_features/quantized_kv_cache", + "docs/advanced_features/reasoning_aware_compression", "docs/advanced_features/dp_dpa_smg_guide", "docs/advanced_features/dcp", "docs/advanced_features/expert_parallelism", diff --git a/docs/docs/advanced_features/reasoning_aware_compression.mdx b/docs/docs/advanced_features/reasoning_aware_compression.mdx new file mode 100644 index 000000000..a7c2e68f9 --- /dev/null +++ b/docs/docs/advanced_features/reasoning_aware_compression.mdx @@ -0,0 +1,105 @@ +--- +title: "Reasoning-Aware Compression" +metatags: + description: "Prune reasoning LLMs without breaking them: calibrate SparseGPT on on-policy chain-of-thought activations collected with SGLang, then serve the sparse checkpoint." +--- + +Pruning a reasoning model with a standard calibration set does more damage than pruning a +conventional LLM — and it can make the model *slower*. Reasoning-Aware Compression (RAC) fixes this +by changing what the pruning solver calibrates on: the model's own chain of thought, generated with +SGLang. + +From [*Reasoning Models Can be Accurately Pruned Via Chain-of-Thought Reconstruction*](https://arxiv.org/abs/2509.12464) +(ICLR 2026). + +## The problem + +One-shot pruning methods such as SparseGPT and Wanda choose which weights to remove by minimizing a +layer-wise reconstruction error against a calibration activation matrix `X`: + +``` +min_{W'} || W X - W' X ||_F^2 s.t. ||W'||_0 <= S +``` + +`X` is conventionally built from **prompt** tokens — a slice of C4, or a set of task prompts. That +is a fair proxy for a typical serving workload, where the prompt dominates the token count. + +Reasoning models invert that ratio. They emit thousands of chain-of-thought tokens per query, so +almost every forward pass the pruned model will ever run is over a token it generated itself. +Calibrating only on prompts optimizes the pruned weights for a distribution the model barely visits. + +The result is not a graceful accuracy decay. The pruned model starts to ramble: it produces longer +chains of thought *and* answers less accurately, so pruning increases end-to-end latency instead of +reducing it. At 50% sparsity on MATH-500, C4-calibrated DeepSeek-R1-Distill-Qwen-7B takes almost six +times as long to evaluate as the dense model it was meant to accelerate. + +## The fix + +RAC samples the dense model's own on-policy rollout during calibration and reconstructs the prompt +and decode activations jointly: + +``` +X_RAC = [ X_prompt , X_decode ] +``` + +The solver is untouched, so this is a drop-in change to any existing SparseGPT or Wanda workflow. + +DeepSeek-R1-Distill-Qwen-7B, MATH-500, SparseGPT at 50% sparsity, 1M calibration tokens: + +| Calibration set | acc@1 | Eval wall clock | +| --- | --- | --- | +| Dense (no pruning) | 0.936 | 23.3 min | +| C4 | 0.744 | 135.0 min | +| Task prompts only | 0.812 | 115.6 min | +| **RAC (prompts + on-policy CoT)** | **0.900** | **35.3 min** | + +Across DeepSeek-R1-Distill-Qwen (1.5B–32B) and Qwen3 (1.7B–14B), the paper reports that RAC keeps up +to 95% of dense accuracy at 50% sparsity, improving on prompt-only calibration by up to 17 points. + +## Using it + +SGLang ships the recipe as a runnable example at +[`examples/usage/reasoning_aware_compression`](https://github.com/sgl-project/sglang/tree/main/examples/usage/reasoning_aware_compression), +in three phases: + +| Phase | Script | What it does | +| --- | --- | --- | +| I | `rac_collect_traces.py` | `sgl.Engine` samples on-policy CoT traces into a calibration set | +| II | `rac_prune.py` | `llm-compressor` runs SparseGPT/Wanda against those activations | +| III | `rac_serve_and_eval.py` | SGLang serves the sparse checkpoint and scores MATH-500 | + +Phase I is the expensive step — the paper's budget is 1M on-policy CoT tokens — and is where SGLang's +batched generation does the work. Phase II delegates the pruning solver to +[`llm-compressor`](https://github.com/vllm-project/llm-compressor), which is **not** an SGLang +dependency; install it separately with `pip install "llmcompressor>=0.12.0"`. + +```bash +cd examples/usage/reasoning_aware_compression + +python rac_collect_traces.py \ + --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ + --dataset open-r1/OpenR1-Math-220k --prompt-column problem \ + --target-tokens 1000000 --output-dir ./rac_traces_math + +python rac_prune.py \ + --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ + --calibration ./rac_traces_math/traces.jsonl \ + --sparsity 0.5 --output-dir ./rac_pruned_50 + +python -m sglang.launch_server --model-path ./rac_pruned_50 +``` + +The example README walks through building the paper's prompt-only baseline from the same prompts so +you can compare calibration strategies head to head. + +## Evaluating a pruned reasoning model + +Accuracy alone will hide the failure mode described above. Always report **mean completion length** +and **wall clock** alongside accuracy when comparing pruned reasoning checkpoints — a model that +scores two points lower while emitting three times the chain of thought is not a good trade. +`rac_serve_and_eval.py` reports all three. + +## Related + +- [Quantization](/docs/advanced_features/quantization) — the other axis of model compression, applied + at serving time. diff --git a/examples/usage/reasoning_aware_compression/README.md b/examples/usage/reasoning_aware_compression/README.md new file mode 100644 index 000000000..e41842985 --- /dev/null +++ b/examples/usage/reasoning_aware_compression/README.md @@ -0,0 +1,173 @@ +# Reasoning-Aware Compression (RAC) + +One-shot pruning of reasoning models, calibrated on the model's own chain of thought. + +Implements the recipe from [*Reasoning Models Can be Accurately Pruned Via Chain-of-Thought +Reconstruction*](https://arxiv.org/abs/2509.12464) (Lucas, Behdin, Wang, Tang, Song, Mazumder; +ICLR 2026). Reference implementation: [RyanLucas3/Reasoning-Aware-Compression](https://github.com/RyanLucas3/Reasoning-Aware-Compression). + +## Why + +Layer-wise one-shot pruning picks weights by minimizing a reconstruction error against a +calibration activation matrix `X`: + +``` +min_{W'} || W X - W' X ||_F^2 s.t. ||W'||_0 <= S +``` + +Every standard pipeline builds `X` from **prompt** tokens — C4 text, or task prompts. That is a +reasonable proxy when `|prompt| >> |output|`. Reasoning models invert the ratio: they emit +thousands of chain-of-thought tokens per query, so nearly all of the forward passes the pruned +model will ever run are over tokens it generated itself. Calibrating on prompts alone leaves the +solver optimizing for a distribution the model barely visits. + +The failure mode this produces is worse than a plain accuracy drop. A poorly calibrated pruned +reasoning model **rambles** — it emits more thinking tokens and still answers less accurately, so +pruning makes it *slower*. From the paper (DeepSeek-R1-Distill-Qwen-7B, MATH-500, SparseGPT at 50% +sparsity, 1M calibration tokens): + +| Calibration set | acc@1 | Eval wall clock | +| --- | --- | --- | +| Dense (no pruning) | 0.936 | 23.3 min | +| C4 | 0.744 | 135.0 min | +| Task prompts only | 0.812 | 115.6 min | +| **RAC (prompts + on-policy CoT)** | **0.900** | **35.3 min** | + +RAC's fix is one line of the algorithm: sample the dense model's own rollout, and calibrate on the +prompt *and* decode activations, + +``` +X_RAC = [ X_prompt , X_decode ] +``` + +The solver is untouched — RAC is a drop-in calibration-set swap for SparseGPT, Wanda, and friends. + +## Why this lives in SGLang + +Collecting the rollout is Phase I of the paper's Algorithm 1, and it is the expensive half: the +paper's budget is 1M on-policy CoT tokens per calibration set. That is batched autoregressive +generation, which is what SGLang does. The pruning solver itself is not an inference-engine +concern, so Phase II delegates to [`llm-compressor`](https://github.com/vllm-project/llm-compressor), +and SGLang serves the result. + +``` +rac_collect_traces.py Phase I sgl.Engine samples on-policy CoT -> traces.jsonl +rac_prune.py Phase II llm-compressor SparseGPT/Wanda -> pruned checkpoint +rac_serve_and_eval.py Phase III sgl.Engine scores MATH-500 -> acc + CoT length + runtime +``` + +## Setup + +Phases I and III need only SGLang. Phase II additionally needs `llm-compressor`, which is **not** an +SGLang dependency: + +```bash +pip install "llmcompressor>=0.12.0" +``` + +Tested against `llmcompressor` 0.12.0. + +## Full run + +Reproduces the paper's DeepSeek-R1-Distill-Qwen-1.5B row at 50% sparsity. The paper runs all +one-shot pruning experiments on a single H100. + +```bash +cd examples/usage/reasoning_aware_compression + +# Phase I -- 1M on-policy CoT tokens (the paper's budget), T_max = 8192, T = 0.6, top_p = 0.95. +python rac_collect_traces.py \ + --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ + --dataset open-r1/OpenR1-Math-220k \ + --prompt-column problem \ + --target-tokens 1000000 \ + --output-dir ./rac_traces_math + +# Phase II -- SparseGPT at 50% unstructured sparsity, calibrated on those traces. +python rac_prune.py \ + --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ + --calibration ./rac_traces_math/traces.jsonl \ + --sparsity 0.5 \ + --output-dir ./rac_pruned_50 + +# Phase III -- accuracy *and* CoT length *and* wall clock. +python rac_serve_and_eval.py --model-path ./rac_pruned_50 --num-problems 500 +``` + +To see what RAC actually buys, build the paper's prompt-only baseline from the same prompts and +compare the two checkpoints directly: + +```bash +python rac_collect_traces.py \ + --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ + --dataset open-r1/OpenR1-Math-220k --prompt-column problem \ + --calibration-mode prompt_only \ + --target-tokens 1000000 \ + --output-dir ./prompt_only_traces_math + +python rac_prune.py \ + --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ + --calibration ./prompt_only_traces_math/traces.jsonl \ + --sparsity 0.5 --output-dir ./prompt_only_pruned_50 + +python rac_serve_and_eval.py \ + --model-path ./prompt_only_pruned_50 ./rac_pruned_50 \ + --num-problems 500 +``` + +`prompt_only` mode skips generation entirely, so it costs nothing but the tokenization pass. + +## Smoke test + +A few minutes on one GPU, to check the plumbing before committing to a 1M-token run: + +```bash +python rac_collect_traces.py --model-path Qwen/Qwen3-0.6B \ + --dataset open-r1/OpenR1-Math-220k --prompt-column problem \ + --target-tokens 20000 --max-new-tokens 1024 --output-dir /tmp/rac_traces +python rac_prune.py --model-path Qwen/Qwen3-0.6B \ + --calibration /tmp/rac_traces/traces.jsonl --sparsity 0.5 --output-dir /tmp/rac_pruned +python rac_serve_and_eval.py --model-path /tmp/rac_pruned --num-problems 50 --max-new-tokens 2048 +``` + +Phase I should report a decode share well above 50% — that gap is the activation mass prompt-only +calibration discards. Phase II should report a realized sparsity within a hair of the target. + +## Models and datasets + +The paper evaluates DeepSeek-R1-Distill-Qwen at 1.5B/7B/14B/32B and Qwen3 at 1.7B/8B/14B, pruned at +20–50% sparsity. Any of them work here; pass `--tp-size` to shard the larger ones. + +Calibration prompts follow the paper: [`open-r1/OpenR1-Math-220k`](https://huggingface.co/datasets/open-r1/OpenR1-Math-220k) +with `--prompt-column problem` for math, and a CodeForces prompt set with `--prompt-column prompt` +for code. `--dataset` also accepts a local `.jsonl` path. + +## Notes + +- **Chat template.** Traces are generated through the model's own chat template with the open-r1 + system prompt, which is what the reference implementation's published traces use. The calibration + distribution *is* the method, so changing `--system-prompt` changes the result. +- **Token ids, not text.** Phase I emits token ids and Phase II consumes them directly, so the + sequence the pruner reconstructs is exactly the sequence the model produced — no + detokenize/retokenize drift. +- **Batch size 1 during calibration.** Padding tokens would enter the layer-wise Hessian as if they + were real activations, which is precisely the contamination RAC exists to avoid. +- **`2:4` masks.** Pass `--mask-structure 2:4` for a semi-structured mask. The paper's headline + results are unstructured (`0:0`). +- **Magnitude pruning** is in the reference implementation but not exposed here: `llm-compressor`'s + magnitude modifier is a gradual, training-time modifier rather than a one-shot solver, and RAC is + a one-shot method. +- **Grading.** `rac_serve_and_eval.py` does lightweight boxed-answer matching, enough to rank + checkpoints. For paper-grade numbers use the `lighteval` harness that the RAC and open-r1 repos + use. + +## Citation + +```bibtex +@inproceedings{lucas2026reasoning, + title = {Reasoning Models Can be Accurately Pruned Via Chain-of-Thought Reconstruction}, + author = {Lucas, Ryan and Behdin, Kayhan and Wang, Zhipeng and Tang, Shao and Song, Qingquan and Mazumder, Rahul}, + booktitle = {International Conference on Learning Representations (ICLR)}, + year = {2026} +} +``` diff --git a/examples/usage/reasoning_aware_compression/rac_collect_traces.py b/examples/usage/reasoning_aware_compression/rac_collect_traces.py new file mode 100755 index 000000000..a9be8f69e --- /dev/null +++ b/examples/usage/reasoning_aware_compression/rac_collect_traces.py @@ -0,0 +1,432 @@ +#!/usr/bin/env python3 +""" +Phase I of Reasoning-Aware Compression (RAC): collect on-policy chain-of-thought +traces with SGLang and write them out as a pruning calibration set. + +RAC ("Reasoning Models Can be Accurately Pruned Via Chain-of-Thought +Reconstruction", ICLR 2026, https://arxiv.org/abs/2509.12464) starts from the +observation that one-shot pruning methods minimize a layer-wise reconstruction +error + + min_{W'} || W X - W' X ||_F^2 s.t. ||W'||_0 <= S + +against a calibration activation matrix X built from *prompt* tokens only. A +reasoning model, however, spends most of its forward passes on tokens it +generated itself (|decode| >> |prompt|), so prompt-only calibration is +distribution-shifted away from what the pruned model will actually run. + +RAC's fix is to build the calibration matrix from the dense model's own rollout: + + X_l^RAC = [ X_l^prompt , X_l^decode ] (paper Eq. 7) + +This script is Phase I of the paper's Algorithm 1 -- sampling that rollout -- +which is the expensive half (the paper uses a 1M token budget). Batched +generation is exactly what SGLang is good at, so it is a much cheaper way to get +there than the Hugging Face `generate` loop used by the reference +implementation. Phase II (the pruning solver) lives in `rac_prune.py`. + +Each output row is one calibration sequence: the chat-templated prompt followed +by the model's own continuation, as token ids. Emitting token ids rather than +text means the sequence fed to the pruner is exactly the sequence the model +produced, with no detokenize/retokenize drift. + +Example (paper's math setup): + + python rac_collect_traces.py \ + --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ + --dataset open-r1/OpenR1-Math-220k \ + --prompt-column problem \ + --output-dir ./rac_traces_math + +To produce the paper's "prompt only" ablation baseline from the same prompts, +re-run with `--calibration-mode prompt_only`. +""" + +import argparse +import json +import os +import time +from typing import Iterator, List, Optional + +import msgspec + +import sglang as sgl +from sglang.srt.utils.hf_transformers_utils import get_tokenizer + +# The system prompt used by open-r1's GRPO recipes, which is what the RAC +# reference implementation generated its published traces with. Keeping it +# identical matters: the calibration distribution is the method. +DEFAULT_SYSTEM_PROMPT = ( + "You are a helpful AI Assistant that provides well-reasoned and detailed " + "responses. You first think about the reasoning process as an internal " + "monologue and then provide the user with the answer. Respond in the " + "following format: \n...\n\n\n...\n" +) + + +class TraceStats(msgspec.Struct, frozen=True): + """What one collection run actually produced.""" + + num_rows: int + num_prompt_tokens: int + num_decode_tokens: int + elapsed_seconds: float + + @property + def num_total_tokens(self) -> int: + return self.num_prompt_tokens + self.num_decode_tokens + + +class TraceManifest(msgspec.Struct, frozen=True): + """Provenance for one calibration set, written next to the traces.""" + + model_path: str + calibration_mode: str + dataset: str + prompt_column: str + system_prompt: Optional[str] + num_rows: int + num_prompt_tokens: int + num_decode_tokens: int + num_total_tokens: int + target_tokens: int + num_generations: int + max_new_tokens: int + temperature: float + top_p: float + seed: int + elapsed_seconds: float + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Collect on-policy CoT calibration traces for RAC pruning.", + formatter_class=argparse.ArgumentDefaultsHelpFormatter, + ) + + parser.add_argument( + "--model-path", + default="deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B", + help="Dense reasoning model to collect traces from.", + ) + parser.add_argument( + "--dataset", + default="open-r1/OpenR1-Math-220k", + help="Hugging Face dataset id, or a path to a local .json/.jsonl file.", + ) + parser.add_argument("--dataset-config-name", default=None) + parser.add_argument("--dataset-split", default="train") + parser.add_argument( + "--prompt-column", + default="problem", + help="Column holding the question. 'problem' for math, 'prompt' for code.", + ) + parser.add_argument( + "--max-prompts", + type=int, + default=None, + help="Cap on prompts read from the dataset. Default: read until the " + "token budget is met.", + ) + parser.add_argument("--output-dir", required=True) + + parser.add_argument( + "--calibration-mode", + choices=["rac", "prompt_only"], + default="rac", + help="'rac' appends on-policy CoT activations (paper Eq. 7). " + "'prompt_only' emits prompts alone, i.e. the paper's ablation baseline.", + ) + parser.add_argument( + "--target-tokens", + type=int, + default=1_000_000, + help="Calibration token budget. The paper uses 1M.", + ) + parser.add_argument( + "--num-generations", + type=int, + default=2, + help="Rollouts sampled per prompt. The paper uses 2.", + ) + parser.add_argument( + "--max-new-tokens", + type=int, + default=8192, + help="T_max, the per-rollout CoT length cap. The paper uses 8192.", + ) + parser.add_argument("--temperature", type=float, default=0.6) + parser.add_argument("--top-p", type=float, default=0.95) + parser.add_argument( + "--system-prompt", + default=DEFAULT_SYSTEM_PROMPT, + help="Pass an empty string to omit the system message.", + ) + parser.add_argument( + "--chunk-size", + type=int, + default=256, + help="Prompts per engine call. Bounds how far past the token budget a " + "run can overshoot.", + ) + parser.add_argument( + "--no-text", + action="store_true", + help="Omit the decoded 'text' field from each row to shrink the file. " + "Token ids are what the pruner actually reads; the text is for humans.", + ) + + parser.add_argument("--tp-size", type=int, default=1) + parser.add_argument("--mem-fraction-static", type=float, default=None) + parser.add_argument("--seed", type=int, default=42) + + return parser.parse_args() + + +def load_rows( + *, + dataset: str, + config_name: Optional[str], + split: str, + prompt_column: str, + max_prompts: Optional[int], +): + """Open the prompt corpus that seeds the rollouts.""" + from datasets import load_dataset + + if os.path.exists(dataset): + rows = load_dataset("json", data_files=dataset, split="train") + else: + rows = load_dataset(dataset, config_name, split=split) + + if prompt_column not in rows.column_names: + raise ValueError( + f"Column '{prompt_column}' not in {dataset}. " + f"Available columns: {rows.column_names}" + ) + if max_prompts is not None: + rows = rows.select(range(min(max_prompts, len(rows)))) + + return rows + + +def build_prompt_token_ids( + *, tokenizer, question: str, system_prompt: str +) -> List[int]: + """Chat-template one question into the token ids the model would see.""" + messages = [] + if system_prompt: + messages.append({"role": "system", "content": system_prompt}) + messages.append({"role": "user", "content": question}) + + return tokenizer.apply_chat_template( + messages, + add_generation_prompt=True, + tokenize=True, + ) + + +def iter_question_chunks( + *, rows, prompt_column: str, chunk_size: int +) -> Iterator[List[str]]: + """Yield questions a chunk at a time. + + Chunking matters beyond batching: the corpus (220k rows for the paper's math + set) is far larger than any token budget needs, so templating and rolling out + lazily means a 1M-token run only touches the prompts it actually uses. + """ + for start in range(0, len(rows), chunk_size): + yield rows[start : start + chunk_size][prompt_column] + + +def rollout( + *, llm, prompt_ids_batch: List[List[int]], sampling_params: dict +) -> List[List[int]]: + """Sample one on-policy continuation per entry (Algorithm 1, decode phase).""" + outputs = llm.generate(input_ids=prompt_ids_batch, sampling_params=sampling_params) + return [output["output_ids"] for output in outputs] + + +def collect_traces( + *, + llm, + tokenizer, + rows, + prompt_column: str, + system_prompt: str, + sampling_params: dict, + calibration_mode: str, + target_tokens: int, + num_generations: int, + chunk_size: int, + emit_text: bool, + trace_path: str, +) -> TraceStats: + """Stream calibration rows to disk until the token budget is met.""" + num_rows = 0 + num_prompt_tokens = 0 + num_decode_tokens = 0 + started_at = time.perf_counter() + + with open(trace_path, "w", encoding="utf-8") as trace_file: + for questions in iter_question_chunks( + rows=rows, prompt_column=prompt_column, chunk_size=chunk_size + ): + chunk = [ + build_prompt_token_ids( + tokenizer=tokenizer, + question=question, + system_prompt=system_prompt, + ) + for question in questions + ] + batch = [ids for ids in chunk for _ in range(num_generations)] + + if calibration_mode == "rac": + decode_ids_batch = rollout( + llm=llm, + prompt_ids_batch=batch, + sampling_params=sampling_params, + ) + else: + decode_ids_batch = [[] for _ in batch] + + for prompt_ids, decode_ids in zip(batch, decode_ids_batch): + input_ids = list(prompt_ids) + list(decode_ids) + row = { + "input_ids": input_ids, + "num_prompt_tokens": len(prompt_ids), + "num_decode_tokens": len(decode_ids), + } + if emit_text: + row["text"] = tokenizer.decode(input_ids) + trace_file.write(json.dumps(row, ensure_ascii=False) + "\n") + + num_rows += 1 + num_prompt_tokens += len(prompt_ids) + num_decode_tokens += len(decode_ids) + + total = num_prompt_tokens + num_decode_tokens + print( + f"[rac] rows={num_rows} " + f"tokens={total}/{target_tokens} " + f"(prompt={num_prompt_tokens} decode={num_decode_tokens})", + flush=True, + ) + if total >= target_tokens: + break + + return TraceStats( + num_rows=num_rows, + num_prompt_tokens=num_prompt_tokens, + num_decode_tokens=num_decode_tokens, + elapsed_seconds=time.perf_counter() - started_at, + ) + + +def report(manifest: TraceManifest) -> None: + """Print the prompt/decode split, which is the paper's core diagnostic.""" + total = manifest.num_total_tokens + decode_share = manifest.num_decode_tokens / total if total else 0.0 + + print("\n=== RAC calibration set ===") + print(f" rows : {manifest.num_rows}") + print(f" prompt tokens : {manifest.num_prompt_tokens}") + print(f" decode tokens : {manifest.num_decode_tokens}") + print(f" total tokens : {total}") + print(f" decode share : {decode_share:.1%}") + print(f" wall clock : {manifest.elapsed_seconds/60:.1f} min") + if manifest.calibration_mode == "rac": + print( + "\nThe decode share is the activation mass that prompt-only " + "calibration throws away." + ) + + +def main() -> None: + args = parse_args() + os.makedirs(args.output_dir, exist_ok=True) + + tokenizer = get_tokenizer(args.model_path) + rows = load_rows( + dataset=args.dataset, + config_name=args.dataset_config_name, + split=args.dataset_split, + prompt_column=args.prompt_column, + max_prompts=args.max_prompts, + ) + print(f"[rac] {len(rows)} prompts available in {args.dataset}") + + sampling_params = { + "temperature": args.temperature, + "top_p": args.top_p, + "max_new_tokens": args.max_new_tokens, + } + engine_kwargs = { + "model_path": args.model_path, + "skip_tokenizer_init": True, + "tp_size": args.tp_size, + "random_seed": args.seed, + } + if args.mem_fraction_static is not None: + engine_kwargs["mem_fraction_static"] = args.mem_fraction_static + + trace_path = os.path.join(args.output_dir, "traces.jsonl") + + # prompt_only needs no rollout, so it needs no engine either. + llm = sgl.Engine(**engine_kwargs) if args.calibration_mode == "rac" else None + try: + stats = collect_traces( + llm=llm, + tokenizer=tokenizer, + rows=rows, + prompt_column=args.prompt_column, + system_prompt=args.system_prompt, + sampling_params=sampling_params, + calibration_mode=args.calibration_mode, + target_tokens=args.target_tokens, + num_generations=args.num_generations, + chunk_size=args.chunk_size, + emit_text=not args.no_text, + trace_path=trace_path, + ) + finally: + if llm is not None: + llm.shutdown() + + manifest = TraceManifest( + model_path=args.model_path, + calibration_mode=args.calibration_mode, + dataset=args.dataset, + prompt_column=args.prompt_column, + system_prompt=args.system_prompt or None, + num_rows=stats.num_rows, + num_prompt_tokens=stats.num_prompt_tokens, + num_decode_tokens=stats.num_decode_tokens, + num_total_tokens=stats.num_total_tokens, + target_tokens=args.target_tokens, + num_generations=args.num_generations, + max_new_tokens=args.max_new_tokens, + temperature=args.temperature, + top_p=args.top_p, + seed=args.seed, + elapsed_seconds=stats.elapsed_seconds, + ) + + manifest_path = os.path.join(args.output_dir, "rac_manifest.json") + with open(manifest_path, "wb") as manifest_file: + manifest_file.write(msgspec.json.format(msgspec.json.encode(manifest))) + + report(manifest) + print(f"\nTraces : {trace_path}") + print(f"Manifest : {manifest_path}") + print("\nNext, prune with these activations:") + print( + f" python rac_prune.py --model-path {args.model_path} " + f"--calibration {trace_path} --sparsity 0.5 --output-dir ./rac_pruned" + ) + + +# sgl.Engine spawns subprocesses, so the entry point must be guarded. +if __name__ == "__main__": + main() diff --git a/examples/usage/reasoning_aware_compression/rac_prune.py b/examples/usage/reasoning_aware_compression/rac_prune.py new file mode 100755 index 000000000..11a2a593d --- /dev/null +++ b/examples/usage/reasoning_aware_compression/rac_prune.py @@ -0,0 +1,232 @@ +#!/usr/bin/env python3 +""" +Phase II of Reasoning-Aware Compression (RAC): one-shot prune a reasoning model +against the on-policy chain-of-thought calibration set built by +`rac_collect_traces.py`. + +RAC (https://arxiv.org/abs/2509.12464, ICLR 2026) does not change the pruning +solver. Its whole contribution is which activations the solver reconstructs: +prompt tokens *plus* the model's own decode tokens (paper Eq. 7) instead of +prompt tokens alone. So this script is deliberately thin -- it hands the RAC +calibration set to `llm-compressor`'s SparseGPT or Wanda implementation and +saves the result as a checkpoint SGLang can serve. + +Requires `llm-compressor`, which is NOT an SGLang dependency: + + pip install "llmcompressor>=0.12.0" + +Example (reproduces the paper's 50%-sparsity math setting): + + python rac_prune.py \ + --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ + --calibration ./rac_traces_math/traces.jsonl \ + --sparsity 0.5 \ + --output-dir ./rac_pruned +""" + +import argparse +import json +import os +import random +from typing import List, Optional + +import torch +from torch.utils.data import DataLoader + +INSTALL_HINT = ( + "This script needs llm-compressor, which SGLang does not depend on.\n" + "Install it with:\n\n" + ' pip install "llmcompressor>=0.12.0"\n' +) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="One-shot prune a reasoning model on RAC calibration traces.", + formatter_class=argparse.ArgumentDefaultsHelpFormatter, + ) + + parser.add_argument("--model-path", required=True, help="Dense model to prune.") + parser.add_argument( + "--calibration", + required=True, + help="traces.jsonl produced by rac_collect_traces.py.", + ) + parser.add_argument("--output-dir", required=True) + + parser.add_argument( + "--method", + choices=["sparsegpt", "wanda"], + default="sparsegpt", + help="Layer-wise solver. The paper's headline results use SparseGPT.", + ) + parser.add_argument( + "--sparsity", + type=float, + default=0.5, + help="Layer-wise sparsity. The paper sweeps 0.2-0.5.", + ) + parser.add_argument( + "--mask-structure", + default="0:0", + help="'0:0' is unstructured (the paper's setting). '2:4' gives a " + "semi-structured mask.", + ) + parser.add_argument( + "--max-seq-length", + type=int, + default=8192, + help="Calibration sequences longer than this are truncated.", + ) + parser.add_argument( + "--num-samples", + type=int, + default=None, + help="Cap on calibration sequences used. Default: use all of them.", + ) + parser.add_argument( + "--pipeline", + default="sequential", + help="llm-compressor calibration pipeline. 'sequential' keeps only one " + "decoder layer's Hessians resident, which is what fits on one GPU.", + ) + parser.add_argument("--dtype", default="bfloat16") + parser.add_argument("--device-map", default="auto") + parser.add_argument("--seed", type=int, default=42) + + return parser.parse_args() + + +def load_calibration_sequences( + *, path: str, max_seq_length: int, num_samples: Optional[int], seed: int +) -> List[List[int]]: + """Read RAC traces back as token id sequences.""" + sequences = [] + with open(path, "r", encoding="utf-8") as trace_file: + for line in trace_file: + input_ids = json.loads(line)["input_ids"] + if input_ids: + sequences.append(input_ids[:max_seq_length]) + + if not sequences: + raise ValueError(f"No calibration sequences found in {path}") + + if num_samples is not None and num_samples < len(sequences): + random.Random(seed).shuffle(sequences) + sequences = sequences[:num_samples] + + return sequences + + +def build_calibration_dataloader(sequences: List[List[int]]) -> DataLoader: + """Wrap token id sequences as batches llm-compressor can calibrate on. + + Batch size is 1 on purpose. Batching sequences of different lengths would + require padding, and pad-token activations would enter the layer-wise + Hessian as if they were real ones -- exactly the calibration contamination + RAC is about avoiding. + """ + + def collate(batch: List[List[int]]) -> dict: + input_ids = torch.tensor(batch[0], dtype=torch.long).unsqueeze(0) + return { + "input_ids": input_ids, + "attention_mask": torch.ones_like(input_ids), + } + + return DataLoader(sequences, batch_size=1, shuffle=False, collate_fn=collate) + + +def build_recipe(*, method: str, sparsity: float, mask_structure: str): + """Instantiate the layer-wise solver. RAC leaves this untouched.""" + try: + from llmcompressor.modifiers.pruning import ( + SparseGPTModifier, + WandaPruningModifier, + ) + except ImportError as exc: + raise ImportError(INSTALL_HINT) from exc + + modifier_cls = SparseGPTModifier if method == "sparsegpt" else WandaPruningModifier + return modifier_cls( + sparsity=sparsity, + mask_structure=mask_structure, + targets=["Linear"], + ignore=["re:.*lm_head"], + ) + + +def measure_sparsity(model) -> float: + """Fraction of zeros across the pruned Linear weights.""" + num_zeros = 0 + num_weights = 0 + for name, module in model.named_modules(): + if not isinstance(module, torch.nn.Linear) or "lm_head" in name: + continue + weight = module.weight + num_zeros += int((weight == 0).sum().item()) + num_weights += weight.numel() + + return num_zeros / num_weights if num_weights else 0.0 + + +def main() -> None: + args = parse_args() + + try: + from llmcompressor import oneshot + except ImportError as exc: + raise ImportError(INSTALL_HINT) from exc + + from transformers import AutoModelForCausalLM, AutoTokenizer + + sequences = load_calibration_sequences( + path=args.calibration, + max_seq_length=args.max_seq_length, + num_samples=args.num_samples, + seed=args.seed, + ) + num_calibration_tokens = sum(len(sequence) for sequence in sequences) + print( + f"[rac] calibrating on {len(sequences)} sequences " + f"({num_calibration_tokens} tokens) from {args.calibration}" + ) + + model = AutoModelForCausalLM.from_pretrained( + args.model_path, + dtype=getattr(torch, args.dtype), + device_map=args.device_map, + ) + tokenizer = AutoTokenizer.from_pretrained(args.model_path) + + print( + f"[rac] pruning to {args.sparsity:.0%} sparsity " + f"with {args.method} (mask_structure={args.mask_structure})" + ) + model = oneshot( + model=model, + processor=tokenizer, + dataset=build_calibration_dataloader(sequences), + recipe=build_recipe( + method=args.method, + sparsity=args.sparsity, + mask_structure=args.mask_structure, + ), + pipeline=args.pipeline, + output_dir=args.output_dir, + ) + + realized = measure_sparsity(model) + print(f"\n[rac] realized sparsity: {realized:.2%} (target {args.sparsity:.2%})") + print(f"[rac] checkpoint written to {os.path.abspath(args.output_dir)}") + print("\nServe it:") + print(f" python -m sglang.launch_server --model-path {args.output_dir}") + print("\nOr score it against the dense model:") + print( + f" python rac_serve_and_eval.py --model-path {args.output_dir} " + "--num-problems 100" + ) + + +if __name__ == "__main__": + main() diff --git a/examples/usage/reasoning_aware_compression/rac_serve_and_eval.py b/examples/usage/reasoning_aware_compression/rac_serve_and_eval.py new file mode 100755 index 000000000..7448cf4db --- /dev/null +++ b/examples/usage/reasoning_aware_compression/rac_serve_and_eval.py @@ -0,0 +1,259 @@ +#!/usr/bin/env python3 +""" +Phase III of the Reasoning-Aware Compression (RAC) recipe: serve one or more +checkpoints with SGLang and score them on MATH-500. + +The point of this script is the *pair* of numbers it reports. The RAC paper +(https://arxiv.org/abs/2509.12464, ICLR 2026) shows that a badly calibrated +pruned reasoning model is not just less accurate -- it also rambles, emitting +far more chain-of-thought tokens for a worse answer, so it is slower than the +dense model it was supposed to speed up (paper Fig. 1, and the runtime columns +of Tables 1-2). Accuracy alone hides that. So every row below reports accuracy +next to mean completion length and wall clock. + +Pass several checkpoints to compare calibration strategies head to head: + + python rac_serve_and_eval.py \ + --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B \ + ./pruned_prompt_only \ + ./pruned_rac \ + --num-problems 100 + +Note on grading: the boxed-answer matching here is intentionally simple, good +enough to rank checkpoints during development. For numbers you would put in a +paper, use the lighteval harness the RAC and open-r1 repos use. +""" + +import argparse +import time +from typing import List, Optional + +import msgspec + +import sglang as sgl +from sglang.srt.utils.hf_transformers_utils import get_tokenizer + +DEFAULT_SYSTEM_PROMPT = ( + "You are a helpful AI Assistant that provides well-reasoned and detailed " + "responses. You first think about the reasoning process as an internal " + "monologue and then provide the user with the answer. Respond in the " + "following format: \n...\n\n\n...\n" +) + + +class EvalResult(msgspec.Struct, frozen=True): + model_path: str + num_problems: int + num_correct: int + mean_completion_tokens: float + elapsed_seconds: float + + @property + def accuracy(self) -> float: + return self.num_correct / self.num_problems if self.num_problems else 0.0 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Score pruned reasoning checkpoints on MATH-500 with SGLang.", + formatter_class=argparse.ArgumentDefaultsHelpFormatter, + ) + + parser.add_argument( + "--model-path", + nargs="+", + required=True, + help="One or more checkpoints. Several are evaluated in sequence and " + "reported side by side.", + ) + parser.add_argument("--dataset", default="HuggingFaceH4/MATH-500") + parser.add_argument("--dataset-split", default="test") + parser.add_argument( + "--num-problems", + type=int, + default=500, + help="Problems to score. The full MATH-500 set is 500.", + ) + parser.add_argument( + "--max-new-tokens", + type=int, + default=8192, + help="Generation budget. The paper evaluates with 32768.", + ) + parser.add_argument("--temperature", type=float, default=0.6) + parser.add_argument("--top-p", type=float, default=0.95) + parser.add_argument("--system-prompt", default=DEFAULT_SYSTEM_PROMPT) + parser.add_argument("--tp-size", type=int, default=1) + parser.add_argument("--mem-fraction-static", type=float, default=None) + parser.add_argument("--seed", type=int, default=42) + + return parser.parse_args() + + +def extract_boxed(text: str) -> Optional[str]: + """Return the content of the last \\boxed{...} in text, brace-matched.""" + marker = "\\boxed{" + start = text.rfind(marker) + if start == -1: + return None + + depth = 0 + for index in range(start + len(marker) - 1, len(text)): + if text[index] == "{": + depth += 1 + elif text[index] == "}": + depth -= 1 + if depth == 0: + return text[start + len(marker) : index] + + return None + + +def normalize_answer(answer: str) -> str: + """Strip the LaTeX noise that makes identical answers compare unequal.""" + normalized = answer.strip().rstrip(".").replace(" ", "") + for token in ("\\left", "\\right", "\\!", "\\,", "$", "\\dfrac", "\\tfrac"): + replacement = "\\frac" if token in ("\\dfrac", "\\tfrac") else "" + normalized = normalized.replace(token, replacement) + + if normalized.startswith("\\text{") and normalized.endswith("}"): + normalized = normalized[len("\\text{") : -1] + + return normalized + + +def is_correct(*, completion: str, reference: str) -> bool: + predicted = extract_boxed(completion) + if predicted is None: + return False + return normalize_answer(predicted) == normalize_answer(reference) + + +def load_problems(*, dataset: str, split: str, limit: int) -> tuple: + from datasets import load_dataset + + rows = load_dataset(dataset, split=split) + rows = rows.select(range(min(limit, len(rows)))) + return [row["problem"] for row in rows], [row["answer"] for row in rows] + + +def build_prompts(*, tokenizer, problems: List[str], system_prompt: str) -> List[str]: + prompts = [] + for problem in problems: + messages = [] + if system_prompt: + messages.append({"role": "system", "content": system_prompt}) + messages.append({"role": "user", "content": problem}) + prompts.append( + tokenizer.apply_chat_template( + messages, + add_generation_prompt=True, + tokenize=False, + ) + ) + return prompts + + +def evaluate( + *, + model_path: str, + problems: List[str], + references: List[str], + sampling_params: dict, + system_prompt: str, + engine_kwargs: dict, +) -> EvalResult: + tokenizer = get_tokenizer(model_path) + prompts = build_prompts( + tokenizer=tokenizer, + problems=problems, + system_prompt=system_prompt, + ) + + llm = sgl.Engine(model_path=model_path, **engine_kwargs) + try: + started_at = time.perf_counter() + outputs = llm.generate(prompts, sampling_params) + elapsed_seconds = time.perf_counter() - started_at + finally: + llm.shutdown() + + num_correct = sum( + is_correct(completion=output["text"], reference=reference) + for output, reference in zip(outputs, references) + ) + total_completion_tokens = sum( + output["meta_info"]["completion_tokens"] for output in outputs + ) + + return EvalResult( + model_path=model_path, + num_problems=len(problems), + num_correct=num_correct, + mean_completion_tokens=total_completion_tokens / len(problems), + elapsed_seconds=elapsed_seconds, + ) + + +def report(results: List[EvalResult]) -> None: + width = max(len(result.model_path) for result in results) + + print("\n=== MATH-500 ===") + header = ( + f"{'model':<{width}} {'acc@1':>7} {'mean CoT tokens':>16} {'wall clock':>12}" + ) + print(header) + print("-" * len(header)) + for result in results: + print( + f"{result.model_path:<{width}} " + f"{result.accuracy:>7.3f} " + f"{result.mean_completion_tokens:>16.0f} " + f"{result.elapsed_seconds/60:>10.1f}m" + ) + + if len(results) > 1: + print( + "\nA pruned model that scores worse *and* emits more CoT tokens is " + "the failure mode RAC targets: calibration drift makes it ramble, " + "so it is both less accurate and slower." + ) + + +def main() -> None: + args = parse_args() + + problems, references = load_problems( + dataset=args.dataset, + split=args.dataset_split, + limit=args.num_problems, + ) + print(f"[rac] scoring {len(problems)} problems from {args.dataset}") + + sampling_params = { + "temperature": args.temperature, + "top_p": args.top_p, + "max_new_tokens": args.max_new_tokens, + } + engine_kwargs = {"tp_size": args.tp_size, "random_seed": args.seed} + if args.mem_fraction_static is not None: + engine_kwargs["mem_fraction_static"] = args.mem_fraction_static + + results = [ + evaluate( + model_path=model_path, + problems=problems, + references=references, + sampling_params=sampling_params, + system_prompt=args.system_prompt, + engine_kwargs=engine_kwargs, + ) + for model_path in args.model_path + ] + + report(results) + + +# sgl.Engine spawns subprocesses, so the entry point must be guarded. +if __name__ == "__main__": + main()