Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions docs/cli/train.md
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,8 @@ torchrun --standalone --nproc_per_node=4 scripts/train.py \

- **`--block-size`** (int, default: `8`) Block size for DFlash model.

- **`--use-liger-kernel`** (flag, default: `False`) Use Liger Qwen3 RMSNorm and SwiGLU kernels for DFlash training. Requires the optional `speculators[liger]` extra.

- **`--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, produces block_size-1 tokens).

- **`--max-anchors`** (int, default: `3072`) Maximum anchor positions for DFlash, DSpark, and P-EAGLE training.
Expand Down
13 changes: 12 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@ dependencies = [
]

[project.optional-dependencies]
liger = ["liger-kernel"]
dev = [
# build
"build>=1.5.0",
Expand Down Expand Up @@ -145,7 +146,17 @@ exclude = ["venv", "build", "dist"]
follow_imports = 'silent'

[[tool.mypy.overrides]]
module = ["datasets.*", "transformers.*", "setuptools.*", "setuptools_git_versioning.*", "vllm.*", "triton.*", "hs_connectors.*", "mooncake.*"]
module = [
"datasets.*",
"hs_connectors.*",
"liger_kernel.*",
"mooncake.*",
"transformers.*",
"setuptools.*",
"setuptools_git_versioning.*",
"vllm.*",
"triton.*",
]
ignore_missing_imports=true

[tool.ruff]
Expand Down
34 changes: 28 additions & 6 deletions scripts/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@

from hs_connectors import HiddenStatesBackend
from speculators.model import SpeculatorModel
from speculators.models.dflash.kernels import DFlashKernels, load_liger_dflash_kernels
from speculators.models.eagle3.data import shift_batch
from speculators.models.eagle3.rotary_partial import install_partial_neox_rotary
from speculators.models.mtp.data import shift_batch_mtp
Expand Down Expand Up @@ -365,13 +366,20 @@ def parse_vocab_mappings(args: argparse.Namespace):
return None, None, verifier_config.vocab_size


def _resolve_dflash_kernels(args: argparse.Namespace) -> DFlashKernels | None:
if not getattr(args, "use_liger_kernel", False):
return None
return load_liger_dflash_kernels()


def _build_from_config_only(
model_class: type[SpeculatorModel],
path: str,
t2d: torch.Tensor | None,
d2t: torch.Tensor | None,
verifier_name_or_path: str | None = None,
draft_attn_impl: str | None = None,
dflash_kernels: DFlashKernels | None = None,
) -> SpeculatorModel:
"""Initialize a fresh draft from a saved speculator *config* (no weights).

Expand All @@ -392,7 +400,10 @@ def _build_from_config_only(
and not getattr(speculators_config.verifier, "name_or_path", None)
):
speculators_config.verifier.name_or_path = verifier_name_or_path
model = model_class(config=config)
model_kwargs = {"config": config}
if dflash_kernels is not None:
model_kwargs["dflash_kernels"] = dflash_kernels
model = model_class(**model_kwargs)
if hasattr(model, "load_vocab_mappings"):
model.load_vocab_mappings(t2d, d2t) # type: ignore[attr-defined, operator]
if hasattr(model, "load_verifier_weights"):
Expand Down Expand Up @@ -421,6 +432,8 @@ def build_draft_model(
extracts the native MTP head weights from the verifier, so the decoder-shaping
flags and ``--draft-config`` do not apply.
"""
dflash_kernels = _resolve_dflash_kernels(args)

if args.from_pretrained:
if is_config_only_dir(args.from_pretrained):
logger.info(
Expand All @@ -437,6 +450,7 @@ def build_draft_model(
draft_attn_impl=(
args.draft_attn_impl if args.speculator_type != "mtp" else None
),
dflash_kernels=dflash_kernels,
)
if args.speculator_type != "mtp":
# _attn_implementation is never serialized by HF configs, so re-apply
Expand All @@ -445,12 +459,17 @@ def build_draft_model(
# __init__ resolves its own default ("sdpa") when it is absent.
config = model_class.config_class.from_pretrained(args.from_pretrained)
config.transformer_layer_config._attn_implementation = args.draft_attn_impl
pretrained_kwargs = {
"config": config,
"t2d": t2d,
"d2t": d2t,
"verifier": args.verifier_name_or_path,
}
if dflash_kernels is not None:
pretrained_kwargs["dflash_kernels"] = dflash_kernels
return model_class.from_pretrained(
args.from_pretrained,
config=config,
t2d=t2d,
d2t=d2t,
verifier=args.verifier_name_or_path,
**pretrained_kwargs,
)
return model_class.from_pretrained(
args.from_pretrained,
Expand Down Expand Up @@ -499,11 +518,14 @@ def build_draft_model(
)

args.draft_vocab_size = draft_vocab_size
training_args = vars(args).copy()
if dflash_kernels is not None:
training_args["dflash_kernels"] = dflash_kernels
return model_class.from_training_args(
verifier_config=transformer_layer_config,
t2d=t2d,
d2t=d2t,
**vars(args),
**training_args,
)


Expand Down
29 changes: 17 additions & 12 deletions src/speculators/models/dflash/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,15 +5,13 @@
from torch import nn
from torch.nn.attention.flex_attention import create_block_mask, create_mask
from transformers import PretrainedConfig
from transformers.models.qwen3.modeling_qwen3 import (
Qwen3RMSNorm,
Qwen3RotaryEmbedding,
)
from transformers.models.qwen3.modeling_qwen3 import Qwen3RotaryEmbedding

from speculators.model import DraftVocabMixin, SpeculatorModel
from speculators.models.attention import create_float_mask
from speculators.models.dflash import DFlashSpeculatorConfig
from speculators.models.dflash.attention import create_anchor_block_mask_mod
from speculators.models.dflash.kernels import DEFAULT_DFLASH_KERNELS, DFlashKernels
from speculators.models.dflash.metrics import compute_metrics
from speculators.models.dflash.model_definitions import Qwen3DFlashDecoderLayer
from speculators.models.dflash.utils import (
Expand Down Expand Up @@ -54,6 +52,7 @@ class DFlashDraftModel(DraftVocabMixin, SpeculatorModel):
def __init__(
self,
config: DFlashSpeculatorConfig,
dflash_kernels: DFlashKernels | None = None,
) -> None:
# Forcibly override config settings
if config.transformer_layer_config._attn_implementation is None: # noqa: SLF001
Expand All @@ -70,14 +69,19 @@ def __init__(
)
super().__init__(config=config)
self._init_vocab(config)
kernels = dflash_kernels or DEFAULT_DFLASH_KERNELS

tl_config = config.transformer_layer_config

# Number of draft layers is encoded in transformer_layer_config
num_draft_layers = tl_config.num_hidden_layers
self.layers = nn.ModuleList(
[
Qwen3DFlashDecoderLayer(config.transformer_layer_config, layer_idx) # type: ignore[arg-type]
Qwen3DFlashDecoderLayer(
config.transformer_layer_config, # type: ignore[arg-type]
layer_idx,
kernels,
)
for layer_idx in range(num_draft_layers)
]
)
Expand All @@ -91,9 +95,9 @@ def __init__(
self.uses_full_attn = bool(num_draft_layers - len(self.sliding_window_indices))
self.sliding_window_non_causal = config.sliding_window_non_causal

self.norm = Qwen3RMSNorm(
self.norm = kernels.make_rms_norm(
config.transformer_layer_config.hidden_size,
eps=config.transformer_layer_config.rms_norm_eps, # type: ignore[arg-type]
config.transformer_layer_config.rms_norm_eps, # type: ignore[arg-type]
)
self.rotary_emb = Qwen3RotaryEmbedding(config.transformer_layer_config) # type: ignore[arg-type]

Expand All @@ -102,13 +106,13 @@ def __init__(
config.transformer_layer_config.hidden_size,
bias=False,
)
self.hidden_norm = Qwen3RMSNorm(
self.hidden_norm = kernels.make_rms_norm(
config.transformer_layer_config.hidden_size,
eps=config.transformer_layer_config.rms_norm_eps, # type: ignore[arg-type]
config.transformer_layer_config.rms_norm_eps, # type: ignore[arg-type]
)
self.verifier_norm = Qwen3RMSNorm(
self.verifier_norm = kernels.make_rms_norm(
config.transformer_layer_config.hidden_size,
eps=config.transformer_layer_config.rms_norm_eps, # type: ignore[arg-type]
config.transformer_layer_config.rms_norm_eps, # type: ignore[arg-type]
)
self.verifier_norm.weight.requires_grad = False
self.block_size = config.block_size
Expand Down Expand Up @@ -156,11 +160,12 @@ def from_training_args(
The number of draft layers is encoded in verifier_config.num_hidden_layers,
following the same pattern as EAGLE3.
"""
dflash_kernels = kwargs.pop("dflash_kernels", None)
config = DFlashSpeculatorConfig(
**cls._build_base_config_kwargs("dflash", verifier_config, **kwargs)
)

model = cls(config=config)
model = cls(config=config, dflash_kernels=dflash_kernels)
model.load_vocab_mappings(t2d, d2t)
model.load_verifier_weights()
return model
Expand Down
71 changes: 71 additions & 0 deletions src/speculators/models/dflash/kernels.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,71 @@
"""Explicit module factories for the DFlash Qwen3 backbone."""

from __future__ import annotations

from dataclasses import dataclass
from typing import TYPE_CHECKING, Protocol

import torch
from torch import nn
from transformers.models.qwen3.modeling_qwen3 import Qwen3Config, Qwen3MLP, Qwen3RMSNorm

if TYPE_CHECKING:
from collections.abc import Callable


class RMSNormModule(Protocol):
"""The RMS norm surface the DFlash backbone reads, beyond `nn.Module`."""

@property
def weight(self) -> torch.Tensor: ...

def __call__(self, hidden_states: torch.Tensor) -> torch.Tensor: ...


@dataclass(frozen=True)
class DFlashKernels:
"""Construct the replaceable Qwen3 modules used by a DFlash draft."""

make_rms_norm: Callable[[int, float], RMSNormModule]
make_mlp: Callable[[Qwen3Config], nn.Module]


def _make_qwen3_rms_norm(hidden_size: int, eps: float) -> RMSNormModule:
return Qwen3RMSNorm(hidden_size, eps=eps)


def _make_qwen3_mlp(config: Qwen3Config) -> nn.Module:
return Qwen3MLP(config)


DEFAULT_DFLASH_KERNELS = DFlashKernels(
make_rms_norm=_make_qwen3_rms_norm,
make_mlp=_make_qwen3_mlp,
)


def load_liger_dflash_kernels() -> DFlashKernels:
"""Load Liger lazily and adapt it to the DFlash construction boundary."""
try:
from liger_kernel.transformers import ( # noqa: PLC0415
LigerRMSNorm,
LigerSwiGLUMLP,
)
except ModuleNotFoundError as exc:
if exc.name in {"liger_kernel", "liger_kernel.transformers"}:
raise ImportError(
"--use-liger-kernel requires the optional `speculators[liger]` "
'extra. Install it with `pip install "speculators[liger]"`.'
) from exc
raise

def make_rms_norm(hidden_size: int, eps: float) -> RMSNormModule:
return LigerRMSNorm(hidden_size, eps=eps)

def make_mlp(config: Qwen3Config) -> nn.Module:
return LigerSwiGLUMLP(config)

return DFlashKernels(
make_rms_norm=make_rms_norm,
make_mlp=make_mlp,
)
34 changes: 23 additions & 11 deletions src/speculators/models/dflash/model_definitions.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,12 +8,12 @@
FlashAttentionKwargs,
GradientCheckpointingLayer,
Qwen3Config,
Qwen3MLP,
Qwen3RMSNorm,
eager_attention_forward,
)
from typing_extensions import Unpack

from speculators.models.dflash.kernels import DFlashKernels

if TYPE_CHECKING:
from collections.abc import Callable

Expand Down Expand Up @@ -49,7 +49,7 @@ class Qwen3DFlashAttention(nn.Module):

# Implements the custom attention which injects the target models
# hidden states into the kv cache.
def __init__(self, config: Qwen3Config, layer_idx: int):
def __init__(self, config: Qwen3Config, layer_idx: int, kernels: DFlashKernels):
super().__init__()
self.config = config
self.layer_idx = layer_idx
Expand Down Expand Up @@ -84,8 +84,8 @@ def __init__(self, config: Qwen3Config, layer_idx: int):
config.hidden_size, # type: ignore[arg-type]
bias=config.attention_bias, # type: ignore[arg-type]
)
self.q_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps) # type: ignore[arg-type]
self.k_norm = Qwen3RMSNorm(self.head_dim, eps=config.rms_norm_eps) # type: ignore[arg-type]
self.q_norm = kernels.make_rms_norm(self.head_dim, config.rms_norm_eps) # type: ignore[arg-type]
self.k_norm = kernels.make_rms_norm(self.head_dim, config.rms_norm_eps) # type: ignore[arg-type]
self.sliding_window = (
config.sliding_window
if hasattr(config, "layer_types")
Expand Down Expand Up @@ -156,15 +156,27 @@ def forward(


class Qwen3DFlashDecoderLayer(GradientCheckpointingLayer):
def __init__(self, config: Qwen3Config, layer_idx: int):
def __init__(
self,
config: Qwen3Config,
layer_idx: int,
kernels: DFlashKernels,
):
super().__init__()
self.hidden_size = config.hidden_size
self.self_attn = Qwen3DFlashAttention(config=config, layer_idx=layer_idx)
self.mlp = Qwen3MLP(config)
self.input_layernorm = Qwen3RMSNorm(config.hidden_size, eps=config.rms_norm_eps) # type: ignore[arg-type]
self.post_attention_layernorm = Qwen3RMSNorm(
self.self_attn = Qwen3DFlashAttention(
config=config,
layer_idx=layer_idx,
kernels=kernels,
)
self.mlp = kernels.make_mlp(config)
self.input_layernorm = kernels.make_rms_norm(
config.hidden_size,
config.rms_norm_eps, # type: ignore[arg-type]
)
self.post_attention_layernorm = kernels.make_rms_norm(
config.hidden_size,
eps=config.rms_norm_eps, # type: ignore[arg-type]
config.rms_norm_eps, # type: ignore[arg-type]
)

def forward(
Expand Down
16 changes: 16 additions & 0 deletions src/speculators/train/config/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -436,6 +436,11 @@ class DFlashArgs(_Group):
block_size: int = Field(
default=8, description="Block size for DFlash model (default: 8)."
)
use_liger_kernel: bool = Field(
default=False,
description="Use Liger Qwen3 RMSNorm/SwiGLU kernels for DFlash. Requires the "
"optional `speculators[liger]` extra.",
)
sample_from_anchor: bool | None = Field(
default=None,
description="Sample from the anchor position (all positions predict). "
Expand Down Expand Up @@ -652,6 +657,17 @@ def _resolve_derived_defaults(self) -> "TrainConfig":
self.optimizer.muon_lr = 10 * self.optimizer.lr
return self

@model_validator(mode="after")
def _validate_liger_kernel(self) -> "TrainConfig":
"""The Liger kernels are wired into the DFlash backbone only, so the flag is
rejected outright on any other speculator rather than silently ignored."""
if self.dflash.use_liger_kernel and self.speculator_type != "dflash":
raise ValueError(
"--use-liger-kernel is currently supported only with "
"--speculator-type dflash"
)
return self

@model_validator(mode="after")
def _validate_dpace(self) -> "TrainConfig":
"""D-PACE per-position loss weighting requires CE loss and a smoothing constant
Expand Down
Loading