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

179 lines
7.0 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.
"""Merge LoRA weights from an exported EAGLE checkpoint into the base model and save.
Usage:
python merge_lora.py \
--base_model_path /path/to/original/base/model \
--exported_lora_dir /path/to/exported/eagle/checkpoint \
--output_path /path/to/merged/output
The exported checkpoint (from export_hf_checkpoint.py) contains
adapter_model.safetensors and adapter_config.json in standard peft format.
This script loads the original base model, applies the trained LoRA adapters,
merges them into the base weights, and saves the fused model + tokenizer.
"""
import argparse
from pathlib import Path
from safetensors.torch import load_file
from transformers import AutoTokenizer
from modelopt.torch.speculative.utils import (
CONFIG_OVERRIDES_HELP,
load_vlm_or_llm,
parse_config_overrides,
)
def parse_args():
parser = argparse.ArgumentParser(
description="Merge LoRA weights from an exported EAGLE checkpoint into the base model."
)
parser.add_argument(
"--base_model_path",
type=str,
required=True,
help="Path to the original base model (HF model name or local path).",
)
parser.add_argument(
"--exported_lora_dir",
type=str,
required=True,
help="Path to the exported EAGLE checkpoint containing adapter_model.safetensors.",
)
parser.add_argument(
"--output_path",
type=str,
required=True,
help="Directory to save the merged (fused) base model.",
)
parser.add_argument(
"--trust_remote_code",
action="store_true",
help="Allow loading models that define custom code on the HF Hub. Off by default.",
)
parser.add_argument(
"--config_overrides",
type=str,
default=None,
help=CONFIG_OVERRIDES_HELP,
)
return parser.parse_args()
def main():
args = parse_args()
config_overrides = parse_config_overrides(args.config_overrides)
lora_dir = Path(args.exported_lora_dir)
# Verify exported files exist (standard peft naming)
config_path = lora_dir / "adapter_config.json"
weights_path = lora_dir / "adapter_model.safetensors"
if not config_path.exists() or not weights_path.exists():
raise FileNotFoundError(
f"Expected adapter_config.json and adapter_model.safetensors "
f"in {lora_dir}. Run export_hf_checkpoint.py first."
)
lora_sd = load_file(weights_path)
print(f"Loaded {len(lora_sd)} LoRA tensors from {lora_dir}")
print(f" Sample keys: {list(lora_sd.keys())[:4]}")
# Load the original base model.
#
# Use load_vlm_or_llm rather than AutoModelForCausalLM directly: it falls back to
# AutoModelForCausalLM for plain LLMs (same dtype/device_map, so unchanged behavior), but also
# handles VLMs and registers/loads architectures the Auto* maps don't cover. Cosmos3 is the
# motivating case -- the transformers-cosmos3 plugin registers only the `cosmos3_omni` config,
# never a model under Auto*, so AutoModelForCausalLM raises KeyError('cosmos3_omni') no matter
# what is imported.
print(f"Loading base model from {args.base_model_path}...")
model = load_vlm_or_llm(
args.base_model_path,
dtype="auto",
device_map="cpu",
trust_remote_code=args.trust_remote_code,
config_overrides=config_overrides,
)
tokenizer = AutoTokenizer.from_pretrained(
args.base_model_path, trust_remote_code=args.trust_remote_code
)
# Load LoRA adapter into the base model (export dir uses standard peft naming)
print("Loading LoRA adapter via PeftModel.from_pretrained...")
from peft import PeftModel
model = PeftModel.from_pretrained(model, str(lora_dir))
print(" PeftModel loaded successfully")
# Debug: check adapter file keys vs model keys and values
adapter_keys = set(lora_sd.keys())
model_lora_keys = {k for k in model.state_dict() if ".lora_A." in k or ".lora_B." in k}
print(f" Adapter file keys (first 4): {sorted(adapter_keys)[:4]}")
print(f" Model LoRA keys (first 4): {sorted(model_lora_keys)[:4]}")
# Check if exported lora_B values are actually non-zero
for k, v in lora_sd.items():
if ".lora_B." in k:
print(f" Exported {k}: shape={v.shape}, norm={v.norm().item():.6f}")
break
# Verify lora_B weights are non-zero (B is init'd to zero, so non-zero means loaded)
lora_b_norms = [v.norm().item() for k, v in model.state_dict().items() if ".lora_B." in k]
if not lora_b_norms or all(n == 0 for n in lora_b_norms):
raise RuntimeError("LoRA-B weights are all zero — adapter loading failed.")
print(
f" Verified: {len(lora_b_norms)} LoRA-B matrices "
f"(mean norm={sum(lora_b_norms) / len(lora_b_norms):.4f})"
)
# Merge LoRA into base weights and remove adapter wrappers
model = model.merge_and_unload()
print("LoRA merged successfully.")
# Save
print(f"Saving merged model to {args.output_path}...")
model.save_pretrained(args.output_path)
tokenizer.save_pretrained(args.output_path)
# Restore the original base model's config.json. save_pretrained() with newer
# transformers (>=5.x) rewrites config fields (e.g. rope_theta → rope_parameters,
# torch_dtype → dtype) which can confuse downstream engines like TRT-LLM or vLLM.
# Since LoRA only changes weights — not architecture — the original config is correct.
import shutil
# ...but only when the loaded config still matches the base's. With --config_overrides the
# weights were built from corrected dims, so copying the uncorrected base config back would
# leave config.json disagreeing with model.safetensors and force every downstream reader to
# re-supply the same overrides.
base_config = Path(args.base_model_path) / "config.json"
output_config = Path(args.output_path) / "config.json"
if config_overrides:
print(
" Keeping the saved config.json (config_overrides were applied, so the original "
"base config would not match the merged weights)"
)
elif base_config.exists():
shutil.copy2(str(base_config), str(output_config))
print(f" Restored original config.json from {base_config}")
print(f"Done! Merged model saved to {args.output_path}")
if __name__ == "__main__":
main()