# DFlash speculative-decoding training recipe. Override fields via OmegaConf dotlist on the CLI. # modelopt-schema: modelopt.recipe.config.ModelOptDFlashRecipe metadata: description: DFlash training recipe (model/data/training/dflash bundled). # maps to ModelArguments (main.py) model: model_name_or_path: trust_remote_code: false use_fake_base_for_offline: false # maps to DataArguments (main.py) data: data_path: offline_data_path: # Jinja chat template with {% generation %} tags for answer_only_loss. # Required when answer_only_loss=true. Set in per-model launcher YAML. # Each model keeps its own beside its launcher example, e.g. # tools/launcher/examples/Qwen/Qwen3-8B/chat_template_train.jinja chat_template: # maps to TrainingArguments (main.py) training: # --- commonly modified --- output_dir: num_train_epochs: 10 per_device_train_batch_size: 1 learning_rate: 6.0e-4 warmup_steps: 100 training_seq_len: 4096 logging_steps: 100 save_steps: 5000 cp_size: 1 dp_shard_size: 1 disable_tqdm: true estimate_ar: false ar_validate_steps: 0 answer_only_loss: true # --- rarely modified --- do_eval: false lr_scheduler_type: linear save_strategy: steps weight_decay: 0.0 dataloader_drop_last: true bf16: true tf32: true remove_unused_columns: false ddp_find_unused_parameters: true ddp_timeout: 1800 report_to: tensorboard # maps to DFlashConfig (modelopt/torch/speculative/config.py). dflash: dflash_block_size: 8 dflash_num_anchors: 512 dflash_use_torch_compile: false dflash_self_logit_distillation: true dflash_loss_decay_factor: 4.0 dflash_architecture_config: num_hidden_layers: 5 # mask_token_id: auto-detected from model vocab (override for specific models) # sliding_window and layer_types are inherited from base model config automatically