File size: 18,546 Bytes
a181ec9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 | # train
Trains speculator models using either online or offline hidden states. Supports single-GPU and multi-GPU distributed training.
## Basic Usage
**Single-GPU:**
```bash
speculators train \
--verifier-name-or-path meta-llama/Llama-3.1-8B-Instruct \
--data-path ./training_data \
--save-path ./checkpoints \
--draft-vocab-size 32000 \
--epochs 10
```
**Multi-GPU (DDP):**
```bash
torchrun --standalone --nproc_per_node=4 -m speculators.train \
--verifier-name-or-path meta-llama/Llama-3.1-8B-Instruct \
--data-path ./training_data \
--save-path ./checkpoints \
--draft-vocab-size 32000 \
--epochs 10
```
**Multi-GPU (FSDP sharded):**
```bash
torchrun --standalone --nproc_per_node=4 -m speculators.train \
--verifier-name-or-path meta-llama/Llama-3.1-8B-Instruct \
--data-path ./training_data \
--save-path ./checkpoints \
--draft-vocab-size 32000 \
--epochs 10 \
--fsdp-shard
```
## Arguments
### Model Arguments
- **`--verifier-name-or-path`** (str, required) HuggingFace model ID or local path for the verifier/target model.
- **`--trust-remote-code`** (flag) Allow executing code from HF Hub when loading the verifier's tokenizer.
- **`--speculator-type`** (str, default: `"eagle3"`) Type of speculator model to train. Options: `eagle3`, `dflash`, `dflash2`, `dspark`, `peagle`, `mtp`
- **`--from-pretrained`** (str, default: `""`) Path or HF id of an existing draft checkpoint to load weights from and train — either a previously trained draft or the initialized-but-untrained checkpoint produced by `--dry-run`. May also point to a local directory containing only a `config.json`, in which case a fresh draft is initialized from that full speculator config. Takes precedence over all other model-definition options: it is mutually exclusive with `--draft-config` and the decoder-shaping flags (`--num-layers`, `--draft-arch`, `--draft-hidden-act`, `--sliding-window`, `--full-attention-indices`).
- **`--draft-config`** (str, default: `""`) HF id, directory, or JSON path of a decoder config (`LlamaConfig` for eagle3/peagle, `Qwen3Config` for DFlash-family models) used as the draft `transformer_layer_config`; the rest of the speculator is built from the other CLI args. The draft `hidden_size` must match the verifier (mismatch is not yet supported). If a full speculator config is passed, its nested `transformer_layer_config` is extracted. Mutually exclusive with `--from-pretrained` and with the decoder-shaping flags (`--num-layers`, `--draft-arch`, `--draft-hidden-act`, `--sliding-window`, `--full-attention-indices`).
- **`--dry-run`** (flag) Build the speculator, initialize weights, save a checkpoint to `--save-path`, then exit before training. Useful to validate the config/weights in vLLM before launching a full run; the saved checkpoint can be fed straight back via `--from-pretrained`.
- **`--num-layers`** (int, default: `5` for dflash/dspark/dflash2, `1` otherwise) Number of transformer layers in the draft model.
- **`--draft-arch`** (str, default: `"llama"`) Architecture for the synthesized draft decoder layers. Options: `llama`, `qwen3`. Used by Eagle3 and P-EAGLE, which select the decoder layer class from this value; DFlash-family models always use a Qwen3-style decoder regardless. Both are supported in vLLM for inference, and the target and draft architectures do not have to match.
- **`--draft-hidden-act`** (str, default: `"silu"`) Activation function for draft decoder layers. Setting as `None` will inherit activation function from the verifier model.
### Data Arguments
- **`--data-path`** (str, default: `"./data"`) Path to the processed training data directory.
- **`--on-missing`** (choice: `generate`|`skip`|`warn`|`raise`, default: `generate`) Behavior when cached hidden states are missing:
- `generate`: Generate hidden states on-demand using vLLM endpoint
- `skip`: Skip the sample silently, pads to fill batch.
- `warn`: Skip the sample with a warning, pads to fill batch.
- `raise`: Raise an error
- **`--on-generate`** (choice: `cache`|`delete`, default: `"delete"`) Behavior after generating new hidden states (only applies if `--on-missing=generate`):
- `delete`: Delete hidden states after loading (pure online training)
- `cache`: Store hidden states for reuse in future epochs (hybrid training)
- **`--hidden-states-path`** (str, default: `{data-path}/hidden_states`) Path where cached hidden states files are stored (or will be stored if generating).
- **`--vllm-endpoint`** (str, default: `"http://localhost:8000/v1"`) vLLM endpoint address for generating hidden states on-demand (online training). Ignored if `--on-missing` is not set to `generate`.
- **`--request-timeout`** (float, default: `180.0`) Timeout in seconds for each individual vLLM request.
- **`--max-retries`** (int, default: `3`) Maximum number of retry attempts per vLLM request on failure.
- **`--total-seq-len`** (int, default: `8192`) Maximum total sequence length for training batches. Note: samples will be packed into batches with total combined sequence length `{total-seq-len}`.
### Vocabulary Mapping Arguments
- **`--draft-vocab-size`** (int, default: `None`) Vocabulary size for the draft model. If not specified and no vocab mapping files are provided, uses full verifier vocabulary.
- **`--token-freq-path`** (str, default: `{data-path}/token_freq.pt`) Path to token frequency distribution file. This is used to determine which tokens to include in the reduced draft vocab.
- **`--d2t-path`** (str, default: `None`) Path to draft-to-target vocabulary mapping file (`.npy`). Must be provided with `--t2d-path`.
- **`--t2d-path`** (str, default: `None`) Path to target-to-draft vocabulary mapping file (`.npy`). Must be provided with `--d2t-path`.
- **`--mask-token-id`** (int, default: auto-detect) Token ID to use as mask token (for DFlash). Auto-detected if not provided.
- **`--target-layer-ids`** (int list, default: auto-select) Space-separated list of layer IDs for the auxiliary hidden states. Default: `[2, num_layers//2, num_layers-3]` **If custom layers were specified when launching vLLM, pass the same ids here, excluding the final layer `launch_vllm.py` appends** — that one reaches training separately as the verifier's last hidden states.
### Distributed Training Arguments
- **`--fsdp-shard`** (flag) Shard model parameters across GPUs with FSDP. By default, parameters are fully replicated (DDP-like). Enable this when the model does not fit in a single GPU's memory.
### Training Arguments
- **`--save-path`** (str, default: `"./checkpoints"`) Directory to save model checkpoints.
- **`--epochs`** (int, default: `20`) Number of training epochs.
- **`--lr`** (float, default: `1e-4`) Learning rate.
- **`--train-data-ratio`** (float, default: `0.9`) Ratio of data to use for training, the rest of the provided data will be used for validation.
- **`--no-resume-from-checkpoint`** (flag) Disable automatic checkpoint resumption. Without this flag, this script will automatically load the latest checkpoint in `{save-path}` if one exists.
- **`--logger`** (str, default: `""`) Metric logging backend(s). Options: `trackio`, `wandb`, `tensorboard`, `mlflow` Can specify multiple comma-separated: `--logger tensorboard,wandb`. **Warning:** backend must be pip installed before using.
- **`--log-dir`** (str, default: `"./logs"`) Directory to save training logs. Only applies to some logging backends (e.g. `tensorboard`)
- **`--run-name`** (str, default: `None`) Name for the training run (used by logging backends).
- **`--seed`** (int, default: `42`) Random seed for reproducibility.
- **`--hidden-states-dtype`** (str, default: `"bfloat16"`) Data type for dataloader hidden states and autocast compute. Model master weights are always kept in fp32. Options: `float32` (full precision, for debugging), `bfloat16` (recommended for mixed precision training). Note: `float16` is not supported as it requires gradient scaling to prevent underflow.
- **`--deterministic-cuda`** (flag) Enable deterministic CUDA operations. May impact performance.
- **`--loss-fn`** (str, default: `"ce"` for dflash, `"kl_div"` otherwise) Loss function specification. Pass a name for a single loss (`kl_div`, `rkl`, `jsd`, `ce`, `tv`, `nla`, `lk_hybrid`) or a JSON dict for a weighted combination, e.g. `'{"ce": 0.1, "tv": 0.9}'`. Required to be `ce` when `--per-position-loss-weight dpace` is used.
### Optimizer Arguments
- **`--optimizer`** (str, default: `"muon"`) Optimizer to use. Options: `adamw`, `muon`. The `muon` option applies the Muon optimizer to 2D weight matrices and AdamW to the remaining parameters (norms, biases, embeddings, lm_head).
- **`--weight-decay`** (float, default: `0.01`) Weight decay for the AdamW optimizer (and the AdamW group in muon mode).
- **`--muon-lr`** (float, default: `10*lr`) Learning rate for the Muon (2D weights) group. Only used with `--optimizer muon`. Defaults to 10× the `--lr` value.
- **`--muon-momentum`** (float, default: `0.95`) Momentum for the Muon optimizer. Only used with `--optimizer muon`.
- **`--muon-weight-decay`** (float, default: `0.1`) Weight decay for the Muon optimizer. Only used with `--optimizer muon`.
- **`--muon-ns-steps`** (int, default: `5`) Number of Newton-Schulz steps for Muon. Only used with `--optimizer muon`.
- **`--muon-adjust-lr-fn`** (str, default: `"match_rms_adamw"`) Muon LR adjustment strategy. Options: `original`, `match_rms_adamw`. Only used with `--optimizer muon`.
### Eagle3-Specific Arguments
- **`--norm-before-residual` / `--no-norm-before-residual`** (flag, default: `True`) Toggle normalization before residual connections.
- **`--embed-requires-grad` / `--no-embed-requires-grad`** (flag, default: `False`) Whether to train embedding layer weights.
- **`--norm-before-fc` / `--no-norm-before-fc`** (flag, default: `True` for eagle3, `False` otherwise) Apply a single RMSNorm to the concatenated auxiliary hidden states before the FC projection (gpt-oss style). See `--fc-norm` for the per-layer alternative from the Eagle 3.1 paper.
- **`--fc-norm`** (flag, default: `False`) Apply per-layer RMSNorm to each auxiliary hidden state before concatenation and FC projection (Eagle 3.1 paper approach).
- **`--norm-output` / `--no-norm-output`** (flag, default: `True` for eagle3, `False` otherwise) Feed post-norm hidden states back across TTT steps to stabilize magnitude drift across speculation depths.
- **`--ttt-steps`** (int, default: `3`) Number of test-time training steps
- **`--ttt-step-loss-decay`** (float, default: `1.0`) Loss decay factor for test-time training steps.
### P-EAGLE-Specific Arguments
- **`--num-depths`** (int, default: `8`) Number of parallel prediction depths.
- **`--down-sample-ratio`** (float, default: `0.7`) Geometric decay ratio for COD sampling.
- **`--down-sample-ratio-min`** (float, default: `0.2`) Minimum retention ratio for COD sampling.
### Attention Backend Arguments
- **`--draft-attn-impl`** (str, default: `"simple_flex_attention"`) Attention implementation for draft layers. Options: `simple_flex_attention`, `sdpa`, `eager`. Use `sdpa` or `eager` on hardware where flex attention is unavailable (e.g. Ascend NPU). Applies to Eagle3, P-EAGLE, and DFlash-family models. Not supported for MTP.
### DFlash-Specific Arguments
- **`--block-size`** (int, default: `16` for dflash, `8` otherwise) Block size for DFlash-family models.
- **`--sample-from-anchor`** / **`--no-sample-from-anchor`** (bool, default: algorithm-specific) Whether to sample from the anchor position. `True`: sample from anchor and all mask positions (default for dspark, produces block_size tokens). `False`: anchor is bonus token (default for dflash/dflash2, produces block_size-1 tokens).
- **`--max-anchors`** (int, default: `512`) Maximum anchor positions for DFlash-family and P-EAGLE training.
- **`--dflash-decay-gamma`** (float, default: `4.0`) Decay gamma for DFlash-family loss weighting.
- **`--per-position-loss-weight`** (str, default: `"dpace"` for dflash, `"fixed-exp-decay"` otherwise) Per-position loss weighting scheme. Options: `fixed-exp-decay`, `dpace`. Applies to DFlash-family models. `dpace` requires `--loss-fn ce`.
- **`--dpace-alpha`** (float, default: `0.5`) Confidence smoothing constant for the D-PACE loss. Only used with `--per-position-loss-weight dpace`.
### DFlash2-Specific Arguments
DFlash2 builds on DFlash, so all DFlash-specific arguments apply as well. It defaults to five draft layers, block size 8, and KL loss.
- **`--conv-kernel-size`** (int, default: `2`) Local convolution kernel size.
- **`--conv-group-size`** (int, default: `16`) Channel group size for local convolution.
- **`--selector-rank`** (int, default: `256`) Low-rank dimension of the candidate selector.
- **`--selector-top-k`** (int, default: `16`) Number of candidates retained per position.
- **`--selector-loss-alpha`** (float, default: `1.0`) Weight of the candidate-selector K-way cross-entropy term.
### DSpark-Specific Arguments
DSpark builds on DFlash, so all DFlash-specific arguments apply as well.
- **`--markov-rank`** (int, default: `256`) Low-rank dim of the Markov logit-bias head. `0` disables it.
- **`--markov-head-type`** (str, default: `"vanilla"`) Sequential head variant. Options: `vanilla`, `gated`, `rnn`.
- **`--enable-confidence-head`** / **`--no-enable-confidence-head`** (flag, default: `True`) Attach the per-position acceptance confidence head.
- **`--confidence-head-with-markov`** / **`--no-confidence-head-with-markov`** (flag, default: `True`) Feed the Markov previous-token embedding into the confidence head alongside the backbone hidden state.
- **`--confidence-head-alpha`** (float, default: `1.0`) Weight of the confidence-head BCE term.
### Sliding Window Attention Arguments
All speculator types (except `mtp`) use sliding window attention on all draft layers by default.
- **`--sliding-window`** (int, default: `2048`) Sliding window size for sliding window attention layers.
- **`--full-attention-indices`** (int list, default: none) Space-separated draft layer indices that should use full attention instead of sliding window. Example: `--full-attention-indices 0 2` makes layers 0 and 2 use full attention; the rest use sliding window.
- **`--sliding-window-non-causal`** (flag) Use non-causal (bidirectional) masking within draft blocks for sliding window attention layers. Full attention layers are always bidirectional. Note: vLLM currently doesn't support these models.
### Dataloader Arguments
- **`--num-workers`** (int, default: `12`) Number of dataloader worker processes.
- **`--prefetch-factor`** (int, default: `4`) Number of batches to prefetch per worker.
- **`--noise-std`** (float, default: `0.05`) Standard deviation for noise augmentation on hidden states.
### Checkpoint Arguments
- **`--checkpoint-freq`** (int, default: `1`) Save a checkpoint every N epochs. Must be ≥ 1.
- **`--save-best`** (flag) Save a symbolic link to the checkpoint with the lowest validation loss.
### Learning Rate Scheduler Arguments
- **`--scheduler-type`** (str, default: `"linear"`) Type of learning rate scheduler. Options: `linear`, `cosine`, `none`
- **`--scheduler-warmup-steps`** (int, default: `None`) Number of warmup steps for the scheduler.
- **`--scheduler-warmup-ratio`** (float, default: `None`) Warmup as a fraction of total scheduler steps, in `[0, 1]`. Ignored (with a warning) when `--scheduler-warmup-steps` is also set.
- **`--scheduler-total-steps`** (int, default: `None`) Total number of training steps for the scheduler.
- **`--scheduler-num-cosine-cycles`** (float, default: `0.5`) Number of cosine cycles for cosine scheduler.
## Examples
### Online Training
```bash
# First, start vLLM server
python scripts/launch_vllm.py \
meta-llama/Llama-3.1-8B-Instruct \
-- --port 8000
# Then train with on-demand hidden states generation
speculators train \
--verifier-name-or-path meta-llama/Llama-3.1-8B-Instruct \
--data-path ./training_data \
--vllm-endpoint http://localhost:8000/v1 \
--on-missing generate \
--on-generate delete \
--save-path ./checkpoints \
--draft-vocab-size 32000 \
--epochs 10 \
--lr 3e-5
```
### Offline Training
```bash
# Train using pre-generated hidden states
speculators train \
--verifier-name-or-path meta-llama/Llama-3.1-8B-Instruct \
--data-path ./training_data \
--hidden-states-path ./hidden_states \
--on-missing raise \
--save-path ./checkpoints \
--draft-vocab-size 32000 \
--epochs 10 \
--lr 3e-5
```
### Hybrid Training (Cache on First Epoch)
```bash
speculators train \
--verifier-name-or-path meta-llama/Llama-3.1-8B-Instruct \
--data-path ./training_data \
--hidden-states-path ./hidden_states \
--vllm-endpoint http://localhost:8000/v1 \
--on-missing generate \
--on-generate cache \
--save-path ./checkpoints \
--draft-vocab-size 32000 \
--epochs 10 \
--lr 3e-5
```
### Multi-GPU Training with WandB Logging
```bash
CUDA_VISIBLE_DEVICES=0,1,2,3 torchrun \
--standalone \
--nproc_per_node 4 \
-m speculators.train \
--verifier-name-or-path meta-llama/Llama-3.1-70B-Instruct \
--data-path ./training_data \
--hidden-states-path ./hidden_states \
--save-path ./checkpoints \
--draft-vocab-size 32000 \
--epochs 20 \
--lr 1e-4 \
--logger wandb \
--run-name eagle3-llama-70b \
--scheduler-type cosine \
--scheduler-warmup-steps 100 \
--checkpoint-freq 2 \
--save-best \
--fsdp-shard
```
### Fine-tuning a Pretrained Model
```bash
speculators train \
--verifier-name-or-path meta-llama/Llama-3.1-8B-Instruct \
--from-pretrained ./pretrained_speculator \
--data-path ./new_training_data \
--hidden-states-path ./hidden_states \
--save-path ./finetuned_checkpoints \
--epochs 5 \
--lr 5e-6
```
### Initializing From a Decoder Config (with Dry-Run Validation)
```bash
# Build the speculator from a plain decoder config, initialize weights, save a
# checkpoint, and exit before training so it can be validated in vLLM first.
speculators train \
--verifier-name-or-path Qwen/Qwen3-8B \
--speculator-type dflash \
--draft-config ./qwen3_draft_decoder_config.json \
--draft-vocab-size 32000 \
--save-path ./draft_init \
--dry-run
# After validating ./draft_init in vLLM, train starting from it:
speculators train \
--verifier-name-or-path Qwen/Qwen3-8B \
--speculator-type dflash \
--from-pretrained ./draft_init \
--data-path ./training_data \
--epochs 5 \
--lr 5e-6
```
|