# LiLiCorr speculative-decoding training recipe (arXiv:2608.20530). # # LiLiCorr reuses the DFlash mode/pipeline and adds a reranker over the candidate # lattice the parallel backbone already produces, selected via # dflash_architecture_config.projector_type=lilicorr. The backbone keeps its # top-k candidates per slot; a small transformer scores transitions between them # and serving commits a path greedily, left to right. Trained jointly with the # backbone, so the drafter learns to propose candidates that correlate into longer # accepted sequences. # # The objective is the DFlash block loss plus three weighted terms: # loss = dflash_loss + w_ce*CE + w_margin*hinge + w_pen*penalty # No outer multiplier, so `loss == origin_loss + lilicorr_loss` holds exactly. # Online training is required (the penalty reads the target model's logits). # # This file is the published `base` variant. The `margin` variant differs only in # the composition of the cross-entropy weight — see the two-line override below. # Override fields via an OmegaConf dotlist. # modelopt-schema: modelopt.recipe.config.ModelOptDFlashRecipe metadata: description: LiLiCorr training recipe (DFlash backbone + candidate-lattice reranker). # 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: # Online only: the distractor penalty weights each candidate by the target # model's own logit gap, so there has to be a target model in the process. mode: online data_path: # Jinja chat template with {% generation %} tags for answer_only_loss. chat_template: # maps to TrainingArguments (main.py) training: # --- commonly modified --- output_dir: num_train_epochs: 6 per_device_train_batch_size: 1 gradient_accumulation_steps: 1 learning_rate: 6.0e-4 warmup_ratio: 0.04 training_seq_len: 3072 logging_steps: 50 save_steps: 1000 seed: 42 cp_size: 1 dp_shard_size: 1 disable_tqdm: true # Keep off: eval runs the DFlash backbone only (the reranker is not applied in # pseudo_speculative_generate), so AR here would report the backbone alone and # understate the trained model. Compare via export + a serving benchmark. estimate_ar: false ar_validate_steps: 0 answer_only_loss: true # --- rarely modified --- do_eval: false # Cosine, unlike the sibling recipes' linear: the published variants were trained # under a linear warmup into a cosine decay (specforge's CosineAnnealingWarmupLR), # and the schedule is part of the recipe those numbers came from. lr_scheduler_type: cosine save_strategy: steps weight_decay: 0.0 max_grad_norm: 1.0 dataloader_drop_last: true bf16: true tf32: true remove_unused_columns: false ddp_timeout: 1800 report_to: tensorboard # maps to DFlashConfig (modelopt/torch/speculative/config.py). dflash: dflash_block_size: 16 dflash_num_anchors: 512 dflash_use_torch_compile: false # The reranker's terms are added to the plain weighted cross-entropy. Turning KD # on would replace that base term and change the objective the published # checkpoints were trained under. dflash_self_logit_distillation: false # Static exponential position decay, gamma=7 for block_size=16: early in-block # slots gate acceptance, so they carry more weight. Not 'dpace' — the published # variants were trained on the static decay. dflash_loss_objective: decay dflash_loss_decay_factor: 7.0 # Qwen3 has no native mask token; 151669 is an unused id used by the reference. dflash_mask_token_id: 151669 # Objective composition. Absolute weights, validated all-or-nothing. # `base` (this file): w_ce 0.25, w_margin 0.0 # `margin` : w_ce 0.125, w_margin 0.125 # Both keep w_pen 0.25, so the head's total weight is 0.50 either way and the # variants differ only in how the cross-entropy block is split. # fp32 master weights for the draft. Compute stays bf16 under autocast; this keeps the # master copy -- and therefore AdamW's moments -- in fp32, which a bf16 second moment # cannot represent the updates of. The published results for this recipe were trained # with this on; turning it off changes the optimizer's arithmetic, not just its memory. dflash_fp32_master_weights: true dflash_lilicorr_w_ce: 0.25 dflash_lilicorr_w_margin: 0.0 dflash_lilicorr_w_pen: 0.25 # Hinge width, in units of the log-potential. Unused while w_margin is 0. dflash_lilicorr_margin: 2.0 dflash_architecture_config: num_hidden_layers: 5 # Draft attention/MLP dims — set explicitly (the draft is an independent # Qwen3 model and does NOT inherit these from the base). GQA: 8 KV heads. num_attention_heads: 32 num_key_value_heads: 8 head_dim: 128 intermediate_size: 12288 projector_type: lilicorr # Reranker geometry. Every field is required, never defaulted: # candidate_topk sets the lattice width and the shape of rank_embedding, and # logit_scale/vector_eps change the score without changing any tensor shape — # so a guessed value builds a head that loads cleanly and scores a different # function. K is otherwise free; the method works at any k. # dflash_init_checkpoint restores weights only and reads geometry from here, so # warm-starting reproduces a head only if every field matches the one the # checkpoint was trained with. lilicorr_logit_scale and lilicorr_vector_eps are # the two a mistake would not be caught on, having no effect on any shape. lilicorr_candidate_topk: 8 lilicorr_hidden_size: 1024 lilicorr_factor_dim: 1024 lilicorr_num_layers: 2 lilicorr_num_heads: 8 lilicorr_mlp_ratio: 2.0 # The factors are cosines, so this temperature sets their usable range. lilicorr_logit_scale: 8.0 lilicorr_vector_eps: 1.0e-4 # `data.chat_template` is supplied per run; use the shared # tools/launcher/examples/Qwen/Qwen3-8B/chat_template_train.jinja, as the other # speculative recipes do, so variants stay comparable to each other. # # It differs slightly from the mask the published checkpoints were trained under: the # reference leaves the empty `` preamble out of the assistant span and supervises # `<|im_end|>`, where the shared template does the opposite. The token ids are identical # either way, so this is 6 tokens of supervision per record (4 preamble, 2 end-of-turn) # and nothing else. Noted because it is invisible in the data, not because it is # expected to matter at this scale.