diff --git a/ci/domino_gpu_smoke.py b/ci/domino_gpu_smoke.py new file mode 100644 index 00000000..06f2f64d --- /dev/null +++ b/ci/domino_gpu_smoke.py @@ -0,0 +1,147 @@ +"""Hardware smoke test for the Domino drafter training backend. + +Drives the real training path on GPU with a real target model (default +Qwen3-4B): it runs the frozen target forward to collect the DFlash-style +multi-layer context hidden states, builds the Domino draft via +``DominoTrainerBackend.build_model``, and runs several optimizer steps through +``compute_loss`` (which invokes the block-drafter forward with the causal GRU +correction head and the dual-logit base-anchor curriculum). + +The draft is cold-started, so the useful signals are: + * ``loss`` / ``final_loss`` (Domino-refined logits CE) trending down, + * ``base_loss`` (backbone-only logits CE) trending down, + * ``accuracy`` (final) and ``base_accuracy`` rising, + * ``lambda_base`` decaying from 1 -> 0 (curriculum handing over to the head). + +Run: + python ci/domino_gpu_smoke.py --target /path/to/target-model --steps 120 +""" + +from __future__ import annotations + +import argparse + +import torch +from omegaconf import OmegaConf +from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer + +PROMPTS = [ + "Explain why the sky appears blue during the day, in a few sentences.", + "Write a short Python function that returns the nth Fibonacci number and explain it.", + "Summarize the water cycle and its main stages in a short paragraph.", + "Describe the main differences between TCP and UDP for a networking student.", +] + + +def _build_batch(target, tokenizer, target_layer_ids, device): + """One packed batch: input_ids, loss_mask, and concatenated context hidden states.""" + id_chunks, mask_chunks, hidden_chunks = [], [], [] + for text in PROMPTS: + messages = [{"role": "user", "content": text}] + prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + enc = tokenizer(prompt, return_tensors="pt").to(device) + with torch.no_grad(): + out = target(input_ids=enc["input_ids"], output_hidden_states=True) + # hidden_states is a tuple of (num_layers + 1) tensors; pick the context layers. + layers = [out.hidden_states[i][0] for i in target_layer_ids] # each [S, H] + id_chunks.append(enc["input_ids"][0]) + mask_chunks.append(torch.ones(enc["input_ids"].size(1), device=device)) + hidden_chunks.append(torch.cat(layers, dim=-1)) # [S, num_ctx*H] + + input_ids = torch.cat(id_chunks).unsqueeze(0) + loss_mask = torch.cat(mask_chunks).unsqueeze(0) + hidden = torch.cat(hidden_chunks).unsqueeze(0).to(torch.bfloat16) + return {"input_ids": input_ids, "loss_mask": loss_mask, "hidden_states": hidden, "attention_mask": torch.ones_like(input_ids)} + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--target", required=True, help="path or HF id of the target causal LM") + parser.add_argument("--steps", type=int, default=120) + parser.add_argument("--lr", type=float, default=1e-4) + parser.add_argument("--num-context-layers", type=int, default=5) + parser.add_argument("--lambda-decay-steps", type=int, default=60) + args = parser.parse_args() + + device = "cuda" + torch.manual_seed(0) + + print(f"[smoke] loading target {args.target}") + tokenizer = AutoTokenizer.from_pretrained(args.target) + target = AutoModelForCausalLM.from_pretrained(args.target, torch_dtype=torch.bfloat16).to(device).eval() + target_cfg = AutoConfig.from_pretrained(args.target) + + from verl_speco.models.dflash import build_target_layer_ids + + target_layers = int(getattr(target_cfg, "num_hidden_layers")) + target_layer_ids = build_target_layer_ids(args.num_context_layers, target_layers) + print(f"[smoke] context layers={target_layer_ids} (of {target_layers})") + + batch = _build_batch(target, tokenizer, target_layer_ids, device) + print(f"[smoke] batch seq_len={batch['input_ids'].size(1)} hidden={batch['hidden_states'].size(-1)}") + + cfg = OmegaConf.create( + { + "rollout": { + "drafter": { + "speculative_algorithm": "DOMINO", + "model_path": "/dev/null/does-not-exist", + "training": { + "domino_block_size": 8, + "domino_num_anchors": 128, + "domino_num_target_layers": args.num_context_layers, + "domino_num_hidden_layers": 1, + "domino_lambda_base_decay_steps": args.lambda_decay_steps, + "lr": args.lr, + }, + } + }, + "model": {"path": args.target}, + } + ) + + from verl_speco.backends.domino_trainer_backend import DominoTrainerBackend + + backend = DominoTrainerBackend(cfg, target_cfg) + model, drafter_cfg = backend.build_model() + model = model.to(device).to(torch.bfloat16).train() + backend.target_lm_head = backend.target_lm_head.to(device).to(torch.bfloat16) + optimizer = backend.setup_optimizer(model, cfg.rollout.drafter.training) + n_params = sum(p.numel() for p in model.parameters() if p.requires_grad) + print( + f"[smoke] block_size={model.block_size} gru={model.draft_model.gru_hidden_dim} " + f"emb_dim={model.draft_model.emb_dim} trainable_params={n_params:,}" + ) + + first = None + for step in range(args.steps): + out = backend.compute_loss(model, batch, 0) + num_tokens = out["local_num_tokens"].clamp_min(1) + loss = out["total_local_ploss"] / num_tokens + optimizer.zero_grad() + loss.backward() + torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) + optimizer.step() + + if step % 10 == 0 or step == args.steps - 1: + d = out["diagnostics"] + fl = float(d["domino_final_loss"]) + bl = float(d["domino_base_loss"]) + acc = float(out["accuracy"]) + bacc = float(d["domino_base_accuracy"]) + lam = float(d["domino_lambda_base"]) + if first is None: + first = (fl, bl, acc) + print( + f"[smoke] step {step:3d} final_loss={fl:.4f} base_loss={bl:.4f} " + f"final_acc={acc:.4f} base_acc={bacc:.4f} lambda_base={lam:.3f}" + ) + + print( + f"[smoke] DONE final_loss {first[0]:.4f}->{fl:.4f} base_loss {first[1]:.4f}->{bl:.4f} " + f"final_acc {first[2]:.4f}->{acc:.4f}" + ) + + +if __name__ == "__main__": + main() diff --git a/tests/integration/test_domino_backend_contract.py b/tests/integration/test_domino_backend_contract.py new file mode 100644 index 00000000..9ae0918d --- /dev/null +++ b/tests/integration/test_domino_backend_contract.py @@ -0,0 +1,173 @@ +"""Contract tests for the Domino drafter backend. + +CPU-light: they exercise the Domino projector modules, the lambda-base +curriculum, the algorithm routing (config, oldlogprob aux layers, vLLM +guardrail), and the block-drafter classification. The full training forward is +validated on GPU by ``ci/domino_gpu_smoke.py``. +""" + +from __future__ import annotations + +import pytest + + +def _tiny_domino_config(): + from verl_speco.models.domino import DominoConfig + + return DominoConfig( + hidden_size=8, + intermediate_size=16, + num_attention_heads=2, + num_key_value_heads=2, + num_hidden_layers=1, + vocab_size=32, + num_target_layers=4, + num_context_layers=2, + target_hidden_size=8, + target_num_hidden_layers=4, + target_layer_ids=[1, 3], + mask_token_id=31, + block_size=4, + num_anchors=8, + emb_dim=6, + gru_hidden_dim=10, + pure_draft_prefix_len=1, + rms_norm_eps=1e-6, + max_position_embeddings=64, + ) + + +def test_domino_model_builds_projector_head() -> None: + pytest.importorskip("torch") + pytest.importorskip("transformers") + from verl_speco.models.domino import DominoDraftModel + + config = _tiny_domino_config() + model = DominoDraftModel(config) + + assert model.projector_type == "domino" + # GRU consumes token embeddings (hidden_size) -> gru_hidden_dim. + assert model.prefix_gru.input_size == config.hidden_size + assert model.prefix_gru.hidden_size == config.gru_hidden_dim + # embed_proj: [hidden + gru_hidden] -> emb_dim -> vocab. + assert model.embed_proj[0].in_features == config.hidden_size + config.gru_hidden_dim + assert model.embed_proj[0].out_features == config.emb_dim + assert model.embed_proj[-1].out_features == config.vocab_size + # It still carries the DFlash backbone token embedding (used by the GRU). + assert model.embed_tokens.num_embeddings == config.vocab_size + + +def test_domino_forward_computes_top5_accuracy() -> None: + """top5_correct_count must actually be reduced, not left at its zero init. + + DFlash and DSpark both compute it; a Domino regression here silently reports + top5_acc=0 forever (base_trainer derives top5_acc from this counter). + """ + pytest.importorskip("torch") + pytest.importorskip("transformers") + import torch + + from verl_speco.backends.domino_trainer_backend import DominoTrainingModel + from verl_speco.models.domino import DominoDraftModel + + torch.manual_seed(0) + config = _tiny_domino_config() + model = DominoTrainingModel( + draft_model=DominoDraftModel(config), + block_size=config.block_size, + num_anchors=config.num_anchors, + pure_draft_prefix_len=config.pure_draft_prefix_len, + ) + + bsz, seq_len = 2, 16 + input_ids = torch.randint(0, config.vocab_size, (bsz, seq_len)) + hidden_states_list = [torch.randn(bsz, seq_len, config.target_hidden_size) for _ in config.target_layer_ids] + loss_mask = torch.ones(bsz, seq_len, dtype=torch.long) + lm_head_weight = torch.randn(config.vocab_size, config.hidden_size) + + _, _, _, _, _, diagnostics = model(input_ids, hidden_states_list, loss_mask, lm_head_weight) + + top1 = float(diagnostics["top1_correct_count"]) + top5 = float(diagnostics["top5_correct_count"]) + quality = float(diagnostics["quality_token_count"]) + + assert quality > 0 + # top-5 is a superset of top-1, and with vocab_size=32 sampling 5 candidates + # over that many tokens must hit at least one target. + assert top5 >= top1 + assert top5 > 0 + + +def test_domino_lambda_base_schedule() -> None: + # get_lambda_base is pure-python, but its module (domino_trainer_backend) + # subclasses the torch-based DFlash backend at import time, so it cannot be + # imported without torch. Skip under the torch-free CPU CI like the siblings. + pytest.importorskip("torch") + pytest.importorskip("transformers") + from verl_speco.backends.domino_trainer_backend import get_lambda_base + + assert get_lambda_base(0, decay_steps=100, lambda_start=1.0) == pytest.approx(1.0) + assert get_lambda_base(50, decay_steps=100, lambda_start=1.0) == pytest.approx(0.5) + assert get_lambda_base(100, decay_steps=100, lambda_start=1.0) == pytest.approx(0.0) + assert get_lambda_base(200, decay_steps=100, lambda_start=1.0) == pytest.approx(0.0) + assert get_lambda_base(25, decay_steps=100, lambda_start=0.4) == pytest.approx(0.3) + + +def test_domino_backend_is_block_drafter_metadata() -> None: + pytest.importorskip("torch") + pytest.importorskip("transformers") + from omegaconf import OmegaConf + + from verl_speco.backends.domino_trainer_backend import DominoTrainerBackend + + backend = DominoTrainerBackend( + OmegaConf.create({"rollout": {"drafter": {"training": {}}}, "model": {"path": "/tmp/none"}}), + OmegaConf.create({}), + ) + assert backend.model_type == "domino" + + +def test_domino_config_from_file_routes_to_domino(tmp_path) -> None: + pytest.importorskip("transformers") + import json + + from verl_speco.models.auto import AutoDraftModelConfig + from verl_speco.models.domino import DominoConfig + + config = _tiny_domino_config().to_dict() + config["architectures"] = ["DominoDraftModel"] + (tmp_path / "config.json").write_text(json.dumps(config), encoding="utf-8") + + loaded = AutoDraftModelConfig.from_file(str(tmp_path / "config.json")) + assert isinstance(loaded, DominoConfig) + assert loaded.architectures == ["DominoDraftModel"] + assert loaded.projector_type == "domino" + + +def test_domino_uses_dflash_aux_layers() -> None: + from verl_speco.integration.oldlogprob_layer_ids import resolve_oldlogprob_aux_layer_ids + + layer_ids = resolve_oldlogprob_aux_layer_ids( + {"speculative_algorithm": "DOMINO", "training": {"domino_num_target_layers": 5}}, + target_num_hidden_layers=36, + ) + # Routes down the DFlash multi-context-layer branch (not the EAGLE default triple). + assert layer_ids is not None + assert len(layer_ids) == 5 + + +def test_domino_rejected_by_vllm_config_builder() -> None: + from verl_speco.integration.vllm_runtime import _speculative_method_from_drafter + + with pytest.raises(ValueError, match="projector sub-mode"): + _speculative_method_from_drafter({"speculative_algorithm": "DOMINO"}) + + +def test_domino_rejected_by_sglang_config_builder() -> None: + from verl_speco.integration.sglang_runtime import _server_args_overrides_from_drafter + + with pytest.raises(ValueError, match="projector sub-mode"): + _server_args_overrides_from_drafter( + {"enable": True, "speculative_algorithm": "DOMINO"}, + supported_fields={"speculative_algorithm"}, + ) diff --git a/verl_speco/backends/domino_trainer_backend.py b/verl_speco/backends/domino_trainer_backend.py new file mode 100644 index 00000000..8cdb4bc8 --- /dev/null +++ b/verl_speco/backends/domino_trainer_backend.py @@ -0,0 +1,406 @@ +"""Domino drafter training backend. + +Logic follows NeMo AutoModel's Domino training wrapper (``dflash/domino_core.py``): +the DFlash parallel block backbone drafts a whole block in one non-causal forward, +and a lightweight *causal* correction head fixes each position's blindness to the +block's earlier (drafted) tokens. A single-layer GRU encodes a causal state from +the block's previous tokens, and a low-rank projection of ``[backbone hidden | GRU +state]`` produces a full-vocabulary logit delta that is added to the parallel base +logits. Training jointly supervises the refined (final) and backbone-only (base) +logits with a base-anchor curriculum ``loss = (1-lambda)*final + lambda*base``, +``lambda`` decaying to 0. + +Domino reuses the DFlash block-drafter plumbing (anchor sampling, noise block, +block attention mask, shifted labels), so ``DominoTrainerBackend`` subclasses +``DFlashTrainerBackend`` exactly like DSpark does; only ``build_model`` and the +training forward differ. The shifted-label alignment (target ``x[a+1:a+1+block]``, +prev ``[x[a], labels[:-1]]``, every position supervised) is the DSpark alignment, +which equals AutoModel's ``shift_label=True`` Domino path. + +Domino is not an engine-level speculative algorithm: engines expose it as a +``projector_type`` sub-mode of DFlash, so the serve method stays ``dflash`` and the +correction head is enabled from the checkpoint's ``dflash_config.projector_type``. +``DOMINO`` is therefore never a valid engine algorithm string; see ``vllm_runtime`` +and ``sglang_runtime`` for the serve-time guardrails that point at ``DFLASH``. +""" + +from __future__ import annotations + +import logging +import os +from copy import deepcopy +from typing import Any + +import torch +import torch.nn.functional as F + +from verl_speco.backends.dflash_trainer_backend import ( + DFlashTrainerBackend, + DFlashTrainingModel, + _create_dflash_mask_mod, +) +from verl_speco.models.dflash.flex_attention import compile_friendly_create_block_mask +from verl_speco.models.domino import DominoConfig, DominoDraftModel +from verl_speco.trainer.checkpoint import log_drafter_checkpoint_step + +logger = logging.getLogger(__name__) +logger.setLevel(os.getenv("VERL_LOGGING_LEVEL", "INFO")) + + +def get_lambda_base(step: int, decay_steps: int, lambda_start: float) -> float: + """Base-anchor curriculum weight: linearly decays lambda_start -> 0 over decay_steps.""" + decay_steps = max(1, int(decay_steps)) + progress = min(max(step, 0) / decay_steps, 1.0) + return max(0.0, min(1.0, lambda_start * (1.0 - progress))) + + +class DominoTrainingModel(DFlashTrainingModel): + """Training wrapper around DominoDraftModel (DFlash backbone + causal head).""" + + def __init__( + self, + draft_model: DominoDraftModel, + block_size: int = 16, + num_anchors: int = 512, + loss_decay_gamma: float = 7.0, + pure_draft_prefix_len: int = 1, + lambda_base_start: float = 1.0, + lambda_base_decay_steps: int = 2000, + ): + super().__init__( + draft_model=draft_model, + block_size=block_size, + num_anchors=num_anchors, + loss_decay_gamma=loss_decay_gamma, + front_position_weight=1.0, + front_position_count=0, + loss_mode="full_vocab", + sampled_ce_negatives=0, + ) + if getattr(draft_model, "projector_type", None) != "domino": + raise ValueError( + "DominoTrainingModel requires a draft model built with projector_type='domino' " + "(so prefix_gru / embed_proj exist)." + ) + self.pure_draft_prefix_len = int(pure_draft_prefix_len) + self.lambda_base_start = float(lambda_base_start) + self.lambda_base_decay_steps = int(lambda_base_decay_steps) + self._forward_count = 0 + + @property + def _suffix_start(self) -> int: + # Shifted labels => position 0 is a real next-token prediction; pure_draft_prefix_len + # leading positions stay backbone-only (uncorrected). + return int(self.pure_draft_prefix_len) + + def _current_lambda_base(self) -> float: + return get_lambda_base(self._forward_count, self.lambda_base_decay_steps, self.lambda_base_start) + + # --- shifted-label anchor sampling / label building (DSpark alignment) --- + def _sample_anchor_positions(self, seq_len: int, loss_mask: torch.Tensor, device: torch.device): + bsz = loss_mask.shape[0] + num_candidates = max(seq_len - 1, 0) + if num_candidates <= 0: + anchors = torch.zeros(bsz, self.num_anchors, dtype=torch.long, device=device) + keep_mask = torch.zeros(bsz, self.num_anchors, dtype=torch.bool, device=device) + return anchors, keep_mask + + valid = (loss_mask[:, :num_candidates] > 0.5) & (loss_mask[:, 1 : num_candidates + 1] > 0.5) + valid_counts = valid.sum(dim=1) + indices = self._cached_arange("domino_anchor_indices", num_candidates, device).unsqueeze(0).expand(bsz, -1) + masked_indices = torch.where(valid, indices, seq_len + 1) + random_vals = torch.rand(bsz, num_candidates, device=device) + random_vals = torch.where(valid, random_vals, 2.0) + take_n = min(self.num_anchors, num_candidates) + _, top_idx = torch.topk(random_vals, k=take_n, dim=1, largest=False, sorted=False) + selected = torch.gather(masked_indices, 1, top_idx).sort(dim=1).values + if take_n < self.num_anchors: + selected = torch.cat( + [selected, torch.zeros(bsz, self.num_anchors - take_n, dtype=torch.long, device=device)], + dim=1, + ) + keep_mask = self._cached_arange( + "domino_anchor_keep", self.num_anchors, device + ).unsqueeze(0) < valid_counts.unsqueeze(1).clamp(max=self.num_anchors) + return torch.where(keep_mask, selected, 0), keep_mask + + def _build_label_tensors(self, *, input_ids, loss_mask, anchor_positions, block_keep_mask): + bsz, seq_len = input_ids.shape + device = input_ids.device + n_blocks = anchor_positions.shape[1] + label_offsets = self._cached_arange("domino_label_offsets", self.block_size, device, view_shape=(1, 1, -1)) + 1 + label_indices = anchor_positions.unsqueeze(-1) + label_offsets + valid_label_mask = label_indices < seq_len + safe_label_indices = label_indices.clamp(max=max(seq_len - 1, 0)) + safe_label_indices = torch.where( + block_keep_mask.unsqueeze(-1), safe_label_indices, torch.zeros_like(safe_label_indices) + ) + target_ids = torch.gather(input_ids.unsqueeze(1).expand(-1, n_blocks, -1), 2, safe_label_indices) + target_loss_mask = torch.gather(loss_mask.unsqueeze(1).expand(-1, n_blocks, -1), 2, safe_label_indices) + eval_mask = valid_label_mask & (target_loss_mask > 0.5) & block_keep_mask.unsqueeze(-1) + eval_mask = eval_mask.to(torch.int32).cumprod(dim=-1).bool() + + anchor_token_ids = torch.gather(input_ids, 1, anchor_positions.clamp(0, max(seq_len - 1, 0))) + prev_token_ids = torch.cat([anchor_token_ids.unsqueeze(-1), target_ids[:, :, :-1]], dim=-1) + return target_ids, prev_token_ids, eval_mask, label_indices + + def forward(self, input_ids, hidden_states_list, loss_mask, lm_head_weight): + bsz, seq_len = input_ids.shape + device = input_ids.device + self._forward_count += 1 + lambda_base = self._current_lambda_base() + + context_feature = self.draft_model.extract_context_feature(hidden_states_list) + anchor_positions, block_keep_mask = self._sample_anchor_positions(seq_len, loss_mask, device) + n_blocks = anchor_positions.shape[1] + noise_embedding = self._create_noise_embed(input_ids, anchor_positions, block_keep_mask) + context_position_ids, draft_position_ids = self._create_position_ids(anchor_positions, seq_len) + draft_len = n_blocks * self.block_size + + block_mask = None + if device.type == "cuda": + block_mask = compile_friendly_create_block_mask( + mask_mod=_create_dflash_mask_mod(anchor_positions, block_keep_mask, seq_len, self.block_size), + B=bsz, + H=None, + Q_LEN=draft_len, + KV_LEN=seq_len + draft_len, + device=device, + ) + + draft_hidden = self.draft_model( + draft_input_ids=None, + context_feature=context_feature, + draft_position_ids=draft_position_ids, + context_position_ids=context_position_ids, + block_mask=block_mask, + noise_embedding=noise_embedding, + ).view(bsz, n_blocks, self.block_size, -1) + + target_ids, prev_token_ids, eval_mask, _ = self._build_label_tensors( + input_ids=input_ids, + loss_mask=loss_mask, + anchor_positions=anchor_positions, + block_keep_mask=block_keep_mask, + ) + + weight_mask = eval_mask.float() + if self.loss_decay_gamma is not None and self.loss_decay_gamma > 0: + positions = self._cached_arange("domino_decay_positions", self.block_size, device, view_shape=(1, 1, -1)) + weight_mask = weight_mask * torch.exp(-positions.float() / float(self.loss_decay_gamma)) + + # Causal GRU state over the block's previous tokens (full block, then gather active). + block_emb = self.draft_model.embed_tokens(prev_token_ids) # [bsz, n, block, H] + gru_out, _ = self.draft_model.prefix_gru(block_emb.reshape(bsz * n_blocks, self.block_size, -1)) + gru_out = gru_out.reshape(bsz, n_blocks, self.block_size, -1) + + pos_in_block = ( + self._cached_arange("domino_pos_in_block", self.block_size, device, view_shape=(1, 1, -1)) + .expand(bsz, n_blocks, -1) + .reshape(-1) + ) + flat_targets = target_ids.reshape(-1) + flat_weights = weight_mask.reshape(-1) + active_mask = flat_weights > 0 + active_hidden = draft_hidden.reshape(-1, draft_hidden.size(-1))[active_mask] + active_gru = gru_out.reshape(-1, gru_out.size(-1))[active_mask] + active_targets = flat_targets[active_mask] + active_weights = flat_weights[active_mask] + active_pos = pos_in_block[active_mask] + suffix_mask = active_pos >= self._suffix_start + + loss_per_token = torch.zeros_like(flat_weights) + sanitized_rows = torch.zeros((), dtype=torch.float32, device=device) + active_final_pred = None + active_base_pred = None + active_final_logits = None + active_top5 = None + if active_targets.numel() == 0: + loss = flat_weights.sum() * 0.0 + final_loss = base_loss = loss.detach() + else: + # Base logits (backbone-only) over every active row. The Domino head only + # perturbs suffix positions, so we compute the correction and the final + # CE on the suffix rows and reuse the base CE elsewhere -- this avoids + # materializing a second full [num_active, vocab] logits tensor. + base_logits = F.linear(active_hidden, lm_head_weight).float() # [num_active, vocab] + base_ce = F.cross_entropy(base_logits, active_targets, reduction="none") + final_ce = base_ce.clone() + if bool(suffix_mask.any()): + suffix_delta = self.draft_model.embed_proj( + torch.cat([active_hidden[suffix_mask], active_gru[suffix_mask]], dim=-1) + ) + active_final_logits = base_logits[suffix_mask] + suffix_delta.float() + final_ce[suffix_mask] = F.cross_entropy( + active_final_logits, active_targets[suffix_mask], reduction="none" + ) + + finite = torch.isfinite(final_ce) & torch.isfinite(base_ce) + sanitized_rows = (~finite).sum().to(dtype=torch.float32) + final_ce = torch.where(finite, final_ce, torch.zeros_like(final_ce)) + base_ce = torch.where(finite, base_ce, torch.zeros_like(base_ce)) + active_loss_weights = active_weights * finite.to(dtype=active_weights.dtype) + den = active_loss_weights.sum().clamp(min=1e-6) + final_loss = (final_ce * active_loss_weights).sum() / den + base_loss = (base_ce * active_loss_weights).sum() / den + loss = (1.0 - lambda_base) * final_loss + lambda_base * base_loss + loss_per_token[active_mask] = final_ce + with torch.no_grad(): + topk = min(5, base_logits.shape[-1]) + active_base_pred = base_logits.argmax(dim=-1) + active_final_pred = active_base_pred.clone() + active_top5 = base_logits.topk(topk, dim=-1).indices + if active_final_logits is not None: + active_final_pred[suffix_mask] = active_final_logits.argmax(dim=-1) + active_top5[suffix_mask] = active_final_logits.topk(topk, dim=-1).indices + + with torch.no_grad(): + flat_eval_mask = eval_mask.reshape(-1) + binary_eval_mask = flat_eval_mask & (flat_weights > 0) + correct = torch.zeros_like(flat_weights, dtype=torch.bool) + base_correct = torch.zeros_like(flat_weights, dtype=torch.bool) + top1_correct = torch.zeros((), dtype=torch.float32, device=device) + top5_correct = torch.zeros((), dtype=torch.float32, device=device) + quality_token_count = torch.zeros((), dtype=torch.float32, device=device) + if active_final_pred is not None and active_targets.numel() > 0: + active_correct = active_final_pred.eq(active_targets) + correct[active_mask] = active_correct + base_correct[active_mask] = active_base_pred.eq(active_targets) + top1_correct = active_correct.float().sum() + top5_correct = active_top5.eq(active_targets.unsqueeze(-1)).any(dim=-1).float().sum() + quality_token_count = active_targets.new_tensor(float(active_targets.numel()), dtype=torch.float32) + + binary_weights = binary_eval_mask.view(bsz, n_blocks, self.block_size).float() + loss_3d = loss_per_token.view(bsz, n_blocks, self.block_size) + correct_3d = correct.view(bsz, n_blocks, self.block_size).float() + count_per_position = binary_weights.sum(dim=(0, 1)).to(torch.float32) + loss_sum_per_position = (loss_3d * binary_weights).sum(dim=(0, 1)) + correct_per_position = correct_3d.sum(dim=(0, 1)) + loss_per_position = loss_sum_per_position / count_per_position.clamp(min=1.0) + acc_per_position = correct_per_position / count_per_position.clamp(min=1.0) + valid_token_count = active_weights.sum().to(dtype=torch.float32) + weighted_token_count = flat_weights.sum().to(dtype=torch.float32) + accuracy = correct.float().sum() / binary_eval_mask.float().sum().clamp(min=1.0) + base_accuracy = base_correct.float().sum() / binary_eval_mask.float().sum().clamp(min=1.0) + + diagnostics = { + "correct_count": correct.float().sum().detach(), + "eval_token_count": binary_eval_mask.float().sum().detach(), + "top1_correct_count": top1_correct.detach(), + "top5_correct_count": top5_correct.detach(), + "quality_token_count": quality_token_count.detach(), + "valid_token_count": valid_token_count.detach(), + "weighted_token_count": weighted_token_count.detach(), + "sanitized_rows": sanitized_rows.detach(), + "masked_rows": (~binary_eval_mask & flat_eval_mask).float().sum().detach(), + "sampled_vocab_size": torch.tensor(float(lm_head_weight.shape[0]), dtype=torch.float32, device=device), + "loss_mode_id": torch.tensor(0.0, dtype=torch.float32, device=device), + "loss_sum_per_position": loss_sum_per_position.detach(), + "correct_per_position": correct_per_position.detach(), + "count_per_position": count_per_position.detach(), + "local_ploss_sum": (loss_per_token * binary_eval_mask.float()).sum().detach(), + # Domino-specific diagnostics. + "domino_final_loss": final_loss.detach() if torch.is_tensor(final_loss) else torch.tensor(0.0, device=device), + "domino_base_loss": base_loss.detach() if torch.is_tensor(base_loss) else torch.tensor(0.0, device=device), + "domino_base_accuracy": base_accuracy.detach(), + "domino_lambda_base": torch.tensor(float(lambda_base), dtype=torch.float32, device=device), + } + return loss, accuracy, loss_per_position, acc_per_position, count_per_position, diagnostics + + +class DominoTrainerBackend(DFlashTrainerBackend): + @property + def model_type(self): + return "domino" + + def _training_value(self, training_cfg, domino_key: str, dflash_key: str, default: Any): + value = training_cfg.get(domino_key, None) + if value is not None: + return value + return training_cfg.get(dflash_key, default) + + def _normalize_dflash_config(self, drafter_config, target_hf_config, normalized_state, spec_model_path): + training_cfg = self.config.rollout.drafter.training + if training_cfg.get("domino_num_target_layers", None) is not None: + if getattr(drafter_config, "num_context_layers", None) is None: + drafter_config.num_context_layers = int(training_cfg["domino_num_target_layers"]) + return super()._normalize_dflash_config(drafter_config, target_hf_config, normalized_state, spec_model_path) + + def _build_fallback_config(self, target_hf_config): + training_cfg = self.config.rollout.drafter.training + target_text_config = getattr(target_hf_config, "text_config", target_hf_config) + hidden_size_cfg = self._training_value(training_cfg, "domino_hidden_size", "dflash_hidden_size", None) + hidden_size = int(hidden_size_cfg if hidden_size_cfg is not None else target_text_config.hidden_size) + num_context_layers = int(self._training_value(training_cfg, "domino_num_target_layers", "dflash_num_target_layers", 5)) + target_num_hidden_layers = int(getattr(target_text_config, "num_hidden_layers", 36)) + mask_token_id_cfg = self._training_value(training_cfg, "domino_mask_token_id", "dflash_mask_token_id", None) + mask_token_id = int(mask_token_id_cfg if mask_token_id_cfg is not None else target_text_config.vocab_size - 1) + target_layer_ids = self._training_value(training_cfg, "domino_target_layer_ids", "dflash_target_layer_ids", None) + if target_layer_ids is None: + from verl_speco.models.dflash import build_target_layer_ids + + target_layer_ids = build_target_layer_ids(num_context_layers, target_num_hidden_layers) + return DominoConfig( + hidden_size=hidden_size, + intermediate_size=int(getattr(target_text_config, "intermediate_size", hidden_size * 4)), + num_hidden_layers=int(self._training_value(training_cfg, "domino_num_hidden_layers", "dflash_num_hidden_layers", 5)), + num_attention_heads=int(getattr(target_text_config, "num_attention_heads")), + num_key_value_heads=int(getattr(target_text_config, "num_key_value_heads", getattr(target_text_config, "num_attention_heads"))), + vocab_size=int(target_text_config.vocab_size), + rms_norm_eps=float(getattr(target_text_config, "rms_norm_eps", 1e-6)), + max_position_embeddings=int(getattr(target_text_config, "max_position_embeddings", 32768)), + rope_theta=float(getattr(target_text_config, "rope_theta", 10000.0)), + num_target_layers=target_num_hidden_layers, + num_context_layers=num_context_layers, + target_hidden_size=int(target_text_config.hidden_size), + target_num_hidden_layers=target_num_hidden_layers, + target_layer_ids=target_layer_ids, + mask_token_id=mask_token_id, + block_size=int(training_cfg.get("domino_block_size", 16)), + num_anchors=int(training_cfg.get("domino_num_anchors", 512)), + loss_decay_gamma=float(training_cfg.get("domino_loss_decay_gamma", 7.0)), + emb_dim=int(training_cfg.get("domino_emb_dim", 256)), + gru_hidden_dim=int(training_cfg.get("domino_gru_hidden_dim", 1024)), + pure_draft_prefix_len=int(training_cfg.get("domino_pure_draft_prefix_len", 1)), + shift_label=bool(training_cfg.get("domino_shift_label", True)), + lambda_base_start=float(training_cfg.get("domino_lambda_base_start", 1.0)), + lambda_base_decay_steps=int(training_cfg.get("domino_lambda_base_decay_steps", 2000)), + architectures=["DominoDraftModel"], + ) + + def build_model(self): + target_model_path = self.config.model.path + spec_model_path = self.config.rollout.drafter.model_path + config_path = os.path.join(spec_model_path, "config.json") if spec_model_path else None + target_hf_config = self._get_target_hf_config() + normalized_state = None + + if config_path and os.path.exists(config_path): + drafter_config = DominoConfig.from_domino_pretrained(spec_model_path) + if spec_model_path and os.path.exists(spec_model_path): + log_drafter_checkpoint_step(logger, spec_model_path, action="Loading Domino drafter weights") + normalized_state = self._normalize_draft_state_dict(self._load_draft_state_dict(spec_model_path)) + else: + drafter_config = self._build_fallback_config(target_hf_config) + + if not isinstance(drafter_config, DominoConfig): + raise TypeError(f"Domino config is not a DominoConfig: {type(drafter_config)}") + drafter_config = self._normalize_dflash_config(drafter_config, target_hf_config, normalized_state, spec_model_path) + + draft_model = DominoDraftModel(deepcopy(drafter_config)) + if spec_model_path and os.path.exists(spec_model_path) and os.path.exists(config_path): + self._load_draft_checkpoint(draft_model, spec_model_path, normalized_state=normalized_state) + draft_model.load_embedding(target_model_path) + draft_model.freeze_embedding() + + self.target_lm_head = self._build_target_lm_head(target_model_path, target_hf_config) + training_cfg = self.config.rollout.drafter.training + return DominoTrainingModel( + draft_model=draft_model, + block_size=int(training_cfg.get("domino_block_size", getattr(drafter_config, "block_size", 16))), + num_anchors=int(training_cfg.get("domino_num_anchors", getattr(drafter_config, "num_anchors", 512))), + loss_decay_gamma=float(training_cfg.get("domino_loss_decay_gamma", getattr(drafter_config, "loss_decay_gamma", 7.0))), + pure_draft_prefix_len=int(training_cfg.get("domino_pure_draft_prefix_len", getattr(drafter_config, "pure_draft_prefix_len", 1))), + lambda_base_start=float(training_cfg.get("domino_lambda_base_start", getattr(drafter_config, "lambda_base_start", 1.0))), + lambda_base_decay_steps=int(training_cfg.get("domino_lambda_base_decay_steps", getattr(drafter_config, "lambda_base_decay_steps", 2000))), + ), drafter_config diff --git a/verl_speco/config/speco_base.yaml b/verl_speco/config/speco_base.yaml index d13359ba..82ee6047 100644 --- a/verl_speco/config/speco_base.yaml +++ b/verl_speco/config/speco_base.yaml @@ -133,6 +133,23 @@ actor_rollout_ref: dspark_debug_log: false dspark_debug_log_first_n: 2 dspark_debug_log_interval: 100 + # Domino (DFlash variant: causal GRU correction head + dual-logit curriculum). + # Unset keys fall back to the matching dflash_* value. + domino_block_size: 16 + domino_num_anchors: 512 + domino_loss_decay_gamma: 7.0 + domino_hidden_size: null + domino_num_target_layers: 5 + domino_num_hidden_layers: 5 + domino_mask_token_id: null + domino_target_layer_ids: null + domino_max_window: 512 + domino_emb_dim: 256 + domino_gru_hidden_dim: 1024 + domino_pure_draft_prefix_len: 1 + domino_shift_label: true + domino_lambda_base_start: 1.0 + domino_lambda_base_decay_steps: 2000 # EAGLE-1 / EAGLE-2 draft training (single-step feature regression + # full-vocab soft-CE distillation against the frozen target head). eagle1_num_hidden_layers: 1 diff --git a/verl_speco/integration/oldlogprob_layer_ids.py b/verl_speco/integration/oldlogprob_layer_ids.py index c51caf4c..d25e64d7 100644 --- a/verl_speco/integration/oldlogprob_layer_ids.py +++ b/verl_speco/integration/oldlogprob_layer_ids.py @@ -58,10 +58,10 @@ def _is_dspark_config(config: Any) -> bool: def _is_dflash_config(drafter_cfg: Any, model_configs: tuple[Any, ...]) -> bool: algorithm = _drafter_algorithm(drafter_cfg) - if algorithm in {"DFLASH", "DSPARK"}: + if algorithm in {"DFLASH", "DSPARK", "DOMINO"}: return True return any( - architecture in {"DFlashDraftModel", "DSparkDraftModel", "Qwen3DSparkModel"} + architecture in {"DFlashDraftModel", "DSparkDraftModel", "Qwen3DSparkModel", "DominoDraftModel", "Qwen3DominoModel"} for config in model_configs for architecture in _config_architectures(config) ) @@ -120,6 +120,7 @@ def _dflash_num_context_layers(drafter_cfg: Any, model_configs: tuple[Any, ...], candidates = [] if is_dspark: candidates.append(_get_nested(training_cfg, ("dspark_num_target_layers",), None)) + candidates.append(_get_nested(training_cfg, ("domino_num_target_layers",), None)) candidates.extend( ( _get_nested(training_cfg, ("dflash_num_target_layers",), None), diff --git a/verl_speco/integration/sglang_runtime.py b/verl_speco/integration/sglang_runtime.py index 604cf7a5..c938e384 100644 --- a/verl_speco/integration/sglang_runtime.py +++ b/verl_speco/integration/sglang_runtime.py @@ -513,6 +513,22 @@ def _server_args_overrides_from_drafter(drafter_cfg: dict[str, Any], supported_f if not bool(drafter_cfg.get("enable")): return {} + algorithm = str(drafter_cfg.get("speculative_algorithm", "") or "").strip().upper() + if algorithm == "DOMINO": + # Domino is a DFlash variant, not an engine-level method: engines expose it as + # "dflash" and enable the causal correction head (prefix_gru + embed_proj) from + # the checkpoint's dflash_config.projector_type="domino" (vllm-project/vllm#48241, + # sgl-project/sglang#31328). DOMINO is never a valid SGLang ServerArgs algorithm, + # so fail loud and point at DFLASH, mirroring + # vllm_runtime._speculative_method_from_drafter. + raise ValueError( + "DOMINO is not an engine-level speculative algorithm; Domino is served as a DFlash " + "projector sub-mode. Set actor_rollout_ref.rollout.drafter.speculative_algorithm=DFLASH " + "for the rollout/serve path; the trained checkpoint's dflash_config.projector_type=domino " + "enables the Domino correction head on engines that support it, keeping DOMINO for " + "drafter training." + ) + rollout_cfg = drafter_cfg.get("rollout") or {} training_cfg = drafter_cfg.get("training") or {} overrides = { diff --git a/verl_speco/integration/vllm_runtime.py b/verl_speco/integration/vllm_runtime.py index 3709fab2..f39c7d19 100644 --- a/verl_speco/integration/vllm_runtime.py +++ b/verl_speco/integration/vllm_runtime.py @@ -400,6 +400,19 @@ def _rollout_config_from_config(config: Any) -> Any: def _speculative_method_from_drafter(drafter_cfg: dict[str, Any]) -> str: algorithm = _drafter_algorithm(drafter_cfg) + if algorithm == "DOMINO": + # Domino is a DFlash variant, not an engine-level method: engines expose it as + # "dflash" and enable the causal correction head (prefix_gru + embed_proj) from + # the checkpoint's dflash_config.projector_type="domino" (vllm-project/vllm#48241, + # sgl-project/sglang#31328). DOMINO is never a valid engine algorithm string, so + # fail loud and point at DFLASH. + raise ValueError( + "DOMINO is not an engine-level speculative algorithm; Domino is served as a DFlash " + "projector sub-mode. Set actor_rollout_ref.rollout.drafter.speculative_algorithm=DFLASH " + "for the rollout/serve path; the trained checkpoint's dflash_config.projector_type=domino " + "enables the Domino correction head on engines that support it, keeping DOMINO for " + "drafter training." + ) if algorithm == "DSPARK": return "dflash" if _is_vllm_ascend_runtime_hint() else "dspark" diff --git a/verl_speco/models/auto.py b/verl_speco/models/auto.py index 7d3bd8d8..9b04293d 100644 --- a/verl_speco/models/auto.py +++ b/verl_speco/models/auto.py @@ -23,6 +23,11 @@ "Qwen3DSparkModel", } +_DOMINO_ARCHITECTURE_ALIASES = { + "DominoDraftModel", + "Qwen3DominoModel", +} + def _normalize_int_list(value): if value is None: @@ -180,7 +185,11 @@ def from_file(cls, config_path: str): architecture = architectures[0] - if architecture not in cls._config_mapping and architecture not in _DSPARK_ARCHITECTURE_ALIASES: + if ( + architecture not in cls._config_mapping + and architecture not in _DSPARK_ARCHITECTURE_ALIASES + and architecture not in _DOMINO_ARCHITECTURE_ALIASES + ): raise ValueError(f"Architecture {architecture} not supported") config_class = cls._config_mapping.get(architecture) @@ -195,6 +204,13 @@ def from_file(cls, config_path: str): config["architectures"] = ["DSparkDraftModel"] if "enable_confidence_head" not in config: config["enable_confidence_head"] = float(config.get("confidence_head_alpha", 0.0)) > 0.0 + elif architecture in _DOMINO_ARCHITECTURE_ALIASES: + from .domino import DominoConfig + + config_class = DominoConfig + config["model_type"] = DominoConfig.model_type + config["architectures"] = ["DominoDraftModel"] + config.setdefault("projector_type", "domino") elif architecture in _EAGLE3_ARCHITECTURE_ALIASES: config = _normalize_eagle3_config_dict(config) diff --git a/verl_speco/models/domino/__init__.py b/verl_speco/models/domino/__init__.py new file mode 100644 index 00000000..39c6a970 --- /dev/null +++ b/verl_speco/models/domino/__init__.py @@ -0,0 +1,4 @@ +from verl_speco.models.domino.configuration_domino import DominoConfig +from verl_speco.models.domino.modeling_domino import DominoDraftModel + +__all__ = ["DominoConfig", "DominoDraftModel"] diff --git a/verl_speco/models/domino/configuration_domino.py b/verl_speco/models/domino/configuration_domino.py new file mode 100644 index 00000000..e737b473 --- /dev/null +++ b/verl_speco/models/domino/configuration_domino.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +import json +import os + +from verl_speco.models.dflash import DFlashConfig + + +class DominoConfig(DFlashConfig): + """Configuration for the Domino draft model. + + Domino uses the same target-hidden-state context backbone as DFlash and adds + a causal correction head (a prefix GRU plus a low-rank ``embed_proj`` that + produces a per-position logit delta). Enabled via ``projector_type='domino'``. + """ + + model_type = "domino" + + def __init__( + self, + *args, + block_size: int = 16, + num_anchors: int = 512, + loss_decay_gamma: float = 7.0, + projector_type: str = "domino", + emb_dim: int = 256, + gru_hidden_dim: int = 1024, + pure_draft_prefix_len: int = 1, + shift_label: bool = True, + lambda_base_start: float = 1.0, + lambda_base_decay_steps: int = 2000, + **kwargs, + ): + architectures = kwargs.pop("architectures", None) + super().__init__(*args, **kwargs) + self.architectures = architectures or ["DominoDraftModel"] + self.block_size = int(block_size) + self.num_anchors = int(num_anchors) + self.loss_decay_gamma = float(loss_decay_gamma) + self.projector_type = str(projector_type) + self.emb_dim = int(emb_dim) + self.gru_hidden_dim = int(gru_hidden_dim) + self.pure_draft_prefix_len = int(pure_draft_prefix_len) + self.shift_label = bool(shift_label) + self.lambda_base_start = float(lambda_base_start) + self.lambda_base_decay_steps = int(lambda_base_decay_steps) + + @classmethod + def from_domino_pretrained(cls, model_path: str): + config_path = os.path.join(model_path, "config.json") + with open(config_path, "r", encoding="utf-8") as f: + config = json.load(f) + + config["model_type"] = cls.model_type + config["architectures"] = ["DominoDraftModel"] + config.setdefault("projector_type", "domino") + return cls.from_dict(config) diff --git a/verl_speco/models/domino/modeling_domino.py b/verl_speco/models/domino/modeling_domino.py new file mode 100644 index 00000000..21f1cae3 --- /dev/null +++ b/verl_speco/models/domino/modeling_domino.py @@ -0,0 +1,49 @@ +from __future__ import annotations + +import torch.nn as nn + +from verl_speco.models.dflash import DFlashDraftModel + +from .configuration_domino import DominoConfig + + +class DominoDraftModel(DFlashDraftModel): + """DFlash block drafter plus the Domino causal correction head. + + Adds two modules on top of the DFlash backbone (both faithful to the + AutoModel ``projector_type='domino'`` build in ``draft_qwen3.py``): + + - ``prefix_gru``: a single-layer GRU over the block's previous-token + embeddings, producing a causal state for each block position. + - ``embed_proj``: a low-rank MLP over ``[backbone hidden | GRU state]`` that + emits a full-vocabulary logit delta added to the parallel base logits. + + The head is applied by the training wrapper (``DominoTrainingModel``), mirroring + how DSpark keeps the Markov head callable but applies the bias in the trainer. + """ + + config_class = DominoConfig + + def __init__(self, config: DominoConfig): + super().__init__(config) + self.projector_type = str(getattr(config, "projector_type", "domino")) + self.pure_draft_prefix_len = int(getattr(config, "pure_draft_prefix_len", 1)) + self.shift_label = bool(getattr(config, "shift_label", True)) + self.emb_dim = int(getattr(config, "emb_dim", 256)) + self.gru_hidden_dim = int(getattr(config, "gru_hidden_dim", 1024)) + + self.prefix_gru = nn.GRU( + input_size=config.hidden_size, + hidden_size=self.gru_hidden_dim, + num_layers=1, + batch_first=True, + bias=False, + ) + self.embed_proj = nn.Sequential( + nn.Linear(config.hidden_size + self.gru_hidden_dim, self.emb_dim, bias=False), + nn.SiLU(), + nn.Linear(self.emb_dim, config.vocab_size, bias=False), + ) + + +__all__ = ["DominoConfig", "DominoDraftModel"] diff --git a/verl_speco/trainer/base_trainer.py b/verl_speco/trainer/base_trainer.py index 5c21b613..fddcf37a 100644 --- a/verl_speco/trainer/base_trainer.py +++ b/verl_speco/trainer/base_trainer.py @@ -529,11 +529,13 @@ def _has_mesh_dim(self, dim_name: str) -> bool: ) def _is_block_drafter_backend(self) -> bool: - return getattr(self.backend, "model_type", None) in {"dflash", "dspark"} + return getattr(self.backend, "model_type", None) in {"dflash", "dspark", "domino"} def _block_drafter_metric_prefix(self) -> str: model_type = str(getattr(self.backend, "model_type", "dflash") or "dflash") - return "dspark" if model_type == "dspark" else "dflash" + if model_type in {"dspark", "domino"}: + return model_type + return "dflash" def _block_drafter_config_value(self, suffix: str, default: Any) -> Any: training_cfg = self.config.rollout.drafter.training @@ -698,7 +700,7 @@ def _build_draft_model(self): # A. 实例化模型(委托给backend) pending_target_weight = self._pending_target_lm_head_weight if ( - getattr(self.backend, "model_type", None) in {"eagle3", "dflash", "dspark"} + getattr(self.backend, "model_type", None) in {"eagle3", "dflash", "dspark", "domino"} and torch.is_tensor(pending_target_weight) and pending_target_weight.dim() == 2 ): diff --git a/verl_speco/workers/speco_worker.py b/verl_speco/workers/speco_worker.py index 82f89f13..4af45725 100644 --- a/verl_speco/workers/speco_worker.py +++ b/verl_speco/workers/speco_worker.py @@ -408,11 +408,15 @@ def init_model(self): from verl_speco.backends.dspark_trainer_backend import DSparkTrainerBackend trainer_backend = DSparkTrainerBackend(self.config, self.config.model) + elif algo == "DOMINO": + from verl_speco.backends.domino_trainer_backend import DominoTrainerBackend + + trainer_backend = DominoTrainerBackend(self.config, self.config.model) else: raise ValueError( "Unsupported drafter algorithm " f"{self.config.rollout.drafter.speculative_algorithm!r}; " - "supported algorithms are EAGLE1, EAGLE2, EAGLE3, DFLASH and DSPARK" + "supported algorithms are EAGLE1, EAGLE2, EAGLE3, DFLASH, DSPARK and DOMINO" ) self.trainer = DrafterBaseTrainer(