SpecDec Bench: PostProcess flag (#759)

## What does this PR do?

**Type of change:** ? Bug Fix: https://nvbugspro.nvidia.com/bug/5795144

**Overview:** ? Pass postprocess flag to handle slicing message.

## Usage
<!-- You can potentially add a usage example below. -->

```python
# Add a code snippet demonstrating how to use this
```

## Testing
<!-- Mention how have you tested your change if applicable. -->

## Before your PR is "*Ready for review*"
<!-- If you haven't finished some of the above items you can still open
`Draft` PR. -->

- **Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)**
and your commits are signed.
- **Is this change backward compatible?**: Yes/No <!--- If No, explain
why. -->
- **Did you write any new necessary tests?**: Yes/No
- **Did you add or update any necessary documentation?**: Yes/No
- **Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?**:
Yes/No <!--- Only for new features, API changes, critical bug fixes or
bw breaking changes. -->

## Additional Information
<!-- E.g. related issue. -->


<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->

## Summary by CodeRabbit

## Release Notes

* **New Features**
* Introduced --postprocess command-line option to select postprocessing
strategy. Users can choose "base" (default, preserves existing behavior)
or "gptoss" (new alternative method) with validation to reject invalid
selections.

<sub>✏️ Tip: You can customize this high-level summary in your review
settings.</sub>

<!-- end of auto-generated comment: release notes by coderabbit.ai -->

Signed-off-by: Izzy Putterman <iputterman@nvidia.com>
This commit is contained in:
Izzy Putterman
2026-01-13 02:24:41 +05:30
committed by GitHub
parent b484efb84e
commit 727da95a91
2 changed files with 23 additions and 3 deletions
+1 -1
View File
@@ -28,7 +28,7 @@ MTBench is available [here](https://huggingface.co/datasets/HuggingFaceH4/mt_ben
Download `nvidia/gpt-oss-120b-Eagle3` to a local directory `/path/to/eagle`.
```bash
python3 run.py --model_dir openai/gpt-oss-120b --tokenizer openai/gpt-oss-120b --draft_model_dir /path/to/eagle --mtbench question.jsonl --tp_size 1 --ep_size 1 --draft_length 3 --output_length 4096 --num_requests 80 --engine TRTLLM --concurrency 1
python3 run.py --model_dir openai/gpt-oss-120b --tokenizer openai/gpt-oss-120b --draft_model_dir /path/to/eagle --mtbench question.jsonl --tp_size 1 --ep_size 1 --draft_length 3 --output_length 4096 --num_requests 80 --engine TRTLLM --concurrency 1 --postprocess gptoss
```
+22 -2
View File
@@ -18,7 +18,13 @@ import asyncio
import yaml
from specdec_bench import datasets, metrics, models, runners
from specdec_bench.utils import decode_chat, encode_chat, get_tokenizer, postprocess_base
from specdec_bench.utils import (
decode_chat,
encode_chat,
get_tokenizer,
postprocess_base,
postprocess_gptoss,
)
engines_available = {
"TRTLLM": models.TRTLLMPYTModel,
@@ -109,7 +115,12 @@ def run_simple(args):
metrics_list.insert(0, metrics.AcceptanceRate())
runner = runners.SimpleRunner(model, metrics=metrics_list)
postprocess = postprocess_base
if args.postprocess == "base":
postprocess = postprocess_base
elif args.postprocess == "gptoss":
postprocess = postprocess_gptoss
else:
raise ValueError(f"Invalid postprocess: {args.postprocess}")
asyncio.run(
run_loop(runner, dataset, tokenizer, args.output_length, postprocess, args.concurrency)
@@ -183,6 +194,15 @@ if __name__ == "__main__":
help="Maximum number of concurrent requests",
)
parser.add_argument("--aa_timing", action="store_true", help="Enable AA timing metric")
parser.add_argument(
"--postprocess",
type=str,
required=False,
default="base",
choices=["base", "gptoss"],
help="Postprocess to use",
)
args = parser.parse_args()
if args.runtime_params is not None: