Files
Model-Optimizer/examples/speculative_decoding/train_eagle3_and_export.sh
T
Benjamin Chislett 4292505512 Refactor: Clean up EAGLE training dataset preparation (#684)
## What does this PR do?

**Type of change:** Refactor

**Overview:** 
- Consolidate input dataset preparation into `make_dataset.py`
- Read dataset mix spec from a YAML file
- - Can now specify how many samples to take from each split
- - Can no longer easily split a dataset into train/test sections. I
don't think this feature was really useful to begin with. Most datasets
can already be separated into train/val/test at the split level, and
those that can't are usually going to be splitted by the training FW
anyways.
- Add support for a few new dataset types, magpie 300k/500k/1M, nemotron
post-training dataset v2.

## Usage
See README for detailed example

## Testing
Ran it locally on all dataset modes, works successfully and output looks
good. Checked shuffling, conversation IDs, and output contents were all
unique and usable.

- **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
- **Did you write any new necessary tests?**: No
- **Did you add or update any necessary documentation?**: Yes
- **Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?**:
No



<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->

## Summary by CodeRabbit

* **Documentation**
* Updated speculative decoding example documentation with new dataset
references and standardized file paths.

* **New Features**
* Introduced configuration-driven dataset preparation supporting
multiple dataset sources with centralized configuration files.

* **Refactor**
* Simplified dataset preparation workflow with unified tooling and
updated default data paths throughout the training pipeline.

<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: Benjamin Chislett <bchislett@nvidia.com>
2026-03-18 09:29:49 -07:00

72 lines
2.1 KiB
Bash
Executable File

#!/bin/bash
# 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.
set -eo pipefail
# Set default values for BASE_MODEL and DATA
BASE_MODEL=meta-llama/Llama-3.2-1B-Instruct
DATA=input_conversations/train.jsonl
# Parse input arguments --base_model and --data
while [[ $# -gt 0 ]]; do
key="$1"
case $key in
--base_model)
BASE_MODEL="$2"
shift; shift
;;
--data)
DATA="$2"
shift; shift
;;
--offline_data)
OFFLINE_DATA_PATH="$2"
shift; shift
;;
*)
echo "Unknown argument: $1"
exit 1
;;
esac
done
if [[ "$OFFLINE_DATA_PATH" != "" ]]; then
OFFLINE_DATA_ARGS="--offline-data $OFFLINE_DATA_PATH"
else
OFFLINE_DATA_ARGS=""
fi
MODEL_BASENAME=$(basename "$BASE_MODEL")
echo "==== [1/3] Training draft model ===="
OUTPUT_DIR=ckpts/${MODEL_BASENAME}-$(date +%Y%m%d_%H%M)
mkdir -p "$(dirname "$OUTPUT_DIR")"
./launch_train.sh --model $BASE_MODEL \
--output_dir $OUTPUT_DIR \
$OFFLINE_DATA_ARGS \
--data $DATA \
--num_epochs 2 \
--eagle_config eagle_config.json
echo "==== [2/3] Evaluating ModelOpt checkpoint on MT-Bench ===="
python scripts/ar_validate.py --model_path $OUTPUT_DIR
echo "==== [3/3] Exporting checkpoint to deployment format ===="
EXPORT_PATH=export/${MODEL_BASENAME}-$(date +%Y%m%d_%H%M)
mkdir -p "$(dirname "$EXPORT_PATH")"
python scripts/export_hf_checkpoint.py --model_path $OUTPUT_DIR --export_path $EXPORT_PATH