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>
179 lines
7.0 KiB
Python
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()
|