Files
yeyu-nvidiaandClaude Opus 4.8 9d0df45849 specdec: config_overrides for nested text_config checkpoints + load VLM-capable bases in merge_lora (#2289)
### 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>
2026-09-10 10:06:35 -07:00

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()