mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
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:
@@ -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
|
||||
|
||||
```
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user