mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
### What does this PR do?
Type of change: New feature + bug fix
Two related gaps, both hit while enabling EAGLE3 on a checkpoint whose
config nests its text dims.
**1. `config_overrides` for checkpoints whose `text_config` dims don't
propagate.**
Some multimodal checkpoints carry the real text-tower dims only under
`config.text_config`, leaving the parent fields `None`.
`from_pretrained` then builds a text tower with the wrong shape.
`load_vlm_or_llm` gains an optional `config_overrides` dict applied to
*both* the parent config and its `text_config` before instantiation, and
the three entrypoints that load checkpoints — `ar_validate.py`,
`export_hf_checkpoint.py`, `merge_lora.py` — get a `--config_overrides`
passthrough. `main.py` threads it from `ModelArguments`.
**2. `merge_lora.py` could not merge into any VLM base.**
It loaded via `AutoModelForCausalLM`, which cannot load architectures
absent from the CausalLM Auto map — every VLM base failed. It now goes
through `load_vlm_or_llm`, which routes VLMs to
`AutoModelForVision2Seq`/`AutoModelForImageTextToText` and plain LLMs to
`AutoModelForCausalLM` with the same `dtype`/`device_map`, so LLM
behavior is byte-for-byte unchanged.
Also adds an optional `transformers_cosmos3` import so `cosmos3_omni` is
registered with `AutoConfig` before use, and dispatches that
`model_type` to its model class directly — that plugin registers only a
*config*, never a model under `Auto*`, so `AutoModelForCausalLM` raised
`KeyError('cosmos3_omni')` regardless of imports. The import is wrapped
in `contextlib.suppress(ImportError)`, so it is a no-op when the plugin
isn't installed.
### Usage
```bash
# Checkpoint whose real dims live under config.text_config
python examples/speculative_decoding/scripts/ar_validate.py \
--model_path <ckpt> --trust_remote_code \
--config_overrides '{"num_hidden_layers": 36, "intermediate_size": 12288, "num_key_value_heads": 8}'
# Same flag on export and merge
python examples/speculative_decoding/scripts/export_hf_checkpoint.py \
--model_path <ckpt> --export_path <out> --config_overrides '{"num_hidden_layers": 36}'
python examples/speculative_decoding/scripts/merge_lora.py \
--base_model_path <base> --exported_lora_dir <out> --output_path <merged> \
--config_overrides '{"num_hidden_layers": 36}'
```
```python
model = load_vlm_or_llm(path, config_overrides={"num_hidden_layers": 36}) # default None
```
### Testing
Exercised end-to-end on a Cosmos3-Nano (16B, 36-layer text tower) EAGLE3
LoRA run:
- **Training** — the base loads with all 36 text layers and correct
dims; two 4-epoch co-training runs completed (46,816 steps each).
- **Export + merge** — produced `adapter_model.safetensors` and a merged
base. Verified correct by per-layer weight diff: a `start_layer=18` run
changed **exactly** layers 18-35, with layers 0-17 bit-identical to the
base.
- **AR validation** — `--config_overrides` loads the trained checkpoint;
80/80 MT-Bench samples, AR 3.42.
- **Regression check** — `merge_lora` via `load_vlm_or_llm` produces a
base loadable by `lm_eval`; ifeval/arc_challenge/winogrande all ran to
completion.
No local unit-test run: `nvidia-modelopt` isn't installed in my
checkout, so `tests/conftest.py` fails to import. Relying on CI.
### Before your PR is "*Ready for review*"
- Is this change backward compatible?: ✅ — `config_overrides` defaults
to `None`; the `merge_lora` loader swap keeps the same class, dtype and
device_map for plain LLMs.
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: N/A — no new
dependency; `transformers_cosmos3` is an optional import guarded by
`contextlib.suppress`.
- Did you write any new necessary tests?: ❌ — exercising these paths
needs a checkpoint with a nested `text_config`, which the unit suite has
no fixture for. Happy to add one if a reviewer can point me at a small
suitable model.
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
❌ — can add a *Speculative Decoding* entry for the `merge_lora` VLM fix
if you consider it changelog-worthy.
- Did you get Claude approval on this PR?: ❌ — not yet run.
<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit
- **New Features**
- Added JSON-based model configuration overrides across speculative
decoding, training, validation, export, and LoRA workflows.
- Overrides can update primary model and text configuration settings.
- Expanded support for vision-language models and Cosmos3 Omni
checkpoints.
- **Bug Fixes**
- Improved configuration handling for offline loading and
checkpoint-based initialization.
- Restored draft-model precision during checkpoint loading and model
conversion.
- Added validation for malformed, unsupported, and non-finite override
values.
- Standardized configuration override guidance across command-line
workflows.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->
---------
Signed-off-by: Ye Yu <yeyu@nvidia.com>
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
180 lines
6.4 KiB
Python
180 lines
6.4 KiB
Python
# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
|
# SPDX-License-Identifier: Apache-2.0
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
|
|
"""AR validation for speculative decoding models (EAGLE3, DFlash, Medusa).
|
|
|
|
Supports per-category MT-Bench evaluation and online (context-dependent) validation.
|
|
"""
|
|
|
|
import argparse
|
|
from collections import defaultdict
|
|
|
|
from accelerate import Accelerator
|
|
from datasets import load_dataset
|
|
from tqdm import tqdm
|
|
from transformers import AutoTokenizer
|
|
|
|
import modelopt.torch.opt as mto
|
|
from modelopt.torch.speculative.plugins.hf_eagle import HFARValidation
|
|
from modelopt.torch.speculative.utils import (
|
|
CONFIG_OVERRIDES_HELP,
|
|
load_vlm_or_llm,
|
|
parse_config_overrides,
|
|
)
|
|
|
|
mto.enable_huggingface_checkpointing()
|
|
|
|
|
|
def validate_ar(
|
|
model,
|
|
tokenizer,
|
|
ds,
|
|
steps=3,
|
|
osl=20,
|
|
num_samples=80,
|
|
device=None,
|
|
):
|
|
"""Validate acceptance rate on MT-Bench prompts using online validation.
|
|
|
|
Online validation recomputes ground truth after each accepted draft token
|
|
(context-dependent), matching actual speculative decoding behavior.
|
|
|
|
Args:
|
|
model: Speculative decoding model (EAGLE3, DFlash, etc.)
|
|
tokenizer: Tokenizer for the model.
|
|
ds: MT-Bench dataset (HuggingFace dataset with 'prompt' and optional 'category').
|
|
steps: Number of draft tokens per speculative step.
|
|
osl: Output sequence length.
|
|
num_samples: Max number of samples to evaluate.
|
|
device: Device to run on.
|
|
|
|
Returns:
|
|
List of (category, ar) tuples.
|
|
"""
|
|
validator = HFARValidation(model, tokenizer)
|
|
num_samples = min(num_samples, len(ds))
|
|
results = []
|
|
failures = 0
|
|
for i in tqdm(range(num_samples), desc="Validating AR"):
|
|
prompt = ds[i]["prompt"][0]
|
|
category = ds[i].get("category", "unknown")
|
|
if hasattr(tokenizer, "apply_chat_template"):
|
|
chat_messages = [{"role": "user", "content": prompt}]
|
|
prompt = tokenizer.apply_chat_template(
|
|
chat_messages, tokenize=False, add_generation_prompt=True
|
|
)
|
|
input_ids = tokenizer(prompt, return_tensors="pt").input_ids
|
|
if device:
|
|
input_ids = input_ids.to(device)
|
|
|
|
try:
|
|
_, ar = validator.validate_online(osl, input_ids=input_ids, steps=steps)
|
|
results.append((category, ar))
|
|
except Exception as e:
|
|
failures += 1
|
|
print(f" WARNING: sample {i} ({category}) failed: {e}")
|
|
if failures:
|
|
print(f"WARNING: {failures}/{num_samples} samples failed during AR validation")
|
|
return results
|
|
|
|
|
|
def main():
|
|
parser = argparse.ArgumentParser(description="AR validation for speculative decoding models.")
|
|
parser.add_argument("--model_path", type=str, required=True, help="Path to model directory")
|
|
parser.add_argument("--trust_remote_code", action="store_true", help="Trust remote code")
|
|
parser.add_argument("--steps", type=int, default=3, help="Draft tokens per step")
|
|
parser.add_argument("--osl", type=int, default=32, help="Output sequence length")
|
|
parser.add_argument("--num_samples", type=int, default=80, help="Number of samples")
|
|
parser.add_argument("--per_category", action="store_true", help="Report per-category AR")
|
|
parser.add_argument(
|
|
"--ar_lower_bound",
|
|
type=float,
|
|
default=None,
|
|
help="Error if AR is below this threshold.",
|
|
)
|
|
parser.add_argument(
|
|
"--config_overrides",
|
|
type=str,
|
|
default=None,
|
|
help=CONFIG_OVERRIDES_HELP,
|
|
)
|
|
args = parser.parse_args()
|
|
|
|
config_overrides = parse_config_overrides(args.config_overrides)
|
|
|
|
accelerator = Accelerator()
|
|
model = load_vlm_or_llm(
|
|
args.model_path,
|
|
device_map="auto",
|
|
trust_remote_code=args.trust_remote_code,
|
|
config_overrides=config_overrides,
|
|
)
|
|
tokenizer = AutoTokenizer.from_pretrained(
|
|
args.model_path, trust_remote_code=args.trust_remote_code
|
|
)
|
|
model.eval()
|
|
model = accelerator.prepare(model)
|
|
|
|
ds = load_dataset("HuggingFaceH4/mt_bench_prompts")["train"]
|
|
results = validate_ar(
|
|
model,
|
|
tokenizer,
|
|
ds,
|
|
args.steps,
|
|
args.osl,
|
|
args.num_samples,
|
|
accelerator.device,
|
|
)
|
|
|
|
# validate_ar() clamps to len(ds), so report what was actually attempted rather than the
|
|
# requested --num_samples, which can be larger than the dataset. A non-positive count means
|
|
# nothing ran at all -- distinct from "everything ran and failed", so say so separately.
|
|
attempted = min(args.num_samples, len(ds))
|
|
if attempted <= 0:
|
|
raise ValueError(
|
|
f"No samples to validate: --num_samples={args.num_samples} with a dataset of "
|
|
f"{len(ds)} prompts. Pass a positive --num_samples."
|
|
)
|
|
|
|
if not results:
|
|
raise RuntimeError(
|
|
f"AR validation produced no results: all {attempted} samples failed. "
|
|
"See the per-sample WARNING lines above for the underlying error. "
|
|
"Exiting non-zero so this is not mistaken for a successful validation."
|
|
)
|
|
|
|
if accelerator.is_main_process:
|
|
all_ars = [ar for _, ar in results]
|
|
avg_ar = sum(all_ars) / len(all_ars)
|
|
print(f"\n==== AR Validation Results (osl={args.osl}, steps={args.steps}) ====")
|
|
|
|
if args.per_category:
|
|
cat_ars = defaultdict(list)
|
|
for cat, ar in results:
|
|
cat_ars[cat].append(ar)
|
|
for cat in sorted(cat_ars):
|
|
cat_avg = sum(cat_ars[cat]) / len(cat_ars[cat])
|
|
print(f" {cat:>12}: {cat_avg:.4f}")
|
|
|
|
print(f" {'ALL':>12}: {avg_ar:.4f}")
|
|
print(f" Samples: {len(results)}")
|
|
|
|
if args.ar_lower_bound and avg_ar < args.ar_lower_bound:
|
|
raise ValueError(f"AR {avg_ar:.4f} is below lower bound {args.ar_lower_bound}.")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|