Skip to content
Draft
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
33 changes: 33 additions & 0 deletions docs/sections/basic_usage/training.md
Original file line number Diff line number Diff line change
Expand Up @@ -549,6 +549,39 @@ a complete checkpoint and points `<run_id>-best` at it, even when
`training.save_interval` is zero. `<run_id>-latest` continues to identify the
newest complete checkpoint.

## Sequence packing

Text EAGLE3, DFlash, and DFlash2 can pack variable-length samples within each
microbatch for offline training/evaluation and online disaggregated consumers:

```yaml
training:
attention_backend: flex_attention
sequence_packing: true
```

Packing defaults to `false`. DFlash2 uses `training.strategy: dflash` with a
DFlash2 draft-model config. FlexAttention is required; USP, multimodal inputs,
other algorithms, `compact_teacher`, and `trim_loss_positions` are unsupported.
DFlash/DFlash2 LK and D-PACE are supported; EAGLE3 LK is not.

Packing concatenates the existing microbatch, resets positions per document,
and isolates attention, labels, and teacher shifts. It preserves sample order,
logical batch size, loss normalization, gradient accumulation, and optimizer
schedule. `data.max_length` still limits each original sample; the packed row
can exceed that length.

DFlash/DFlash2 preserve per-document anchor sampling. Invalid padded proposal
slots skip the backbone when host metadata is available; outputs return to the
original batch/anchor/block layout for objectives and metrics. Online packing
happens after the consumer fetches target features, preserving individual
capture requests and sample-ID acknowledgements.

Gains depend on sequence lengths, valid proposal counts, and pipeline overhead.
Batch size one or equal-length samples have no inter-sample context padding to
remove. Compare the same samples, batch size, and accumulation using unpadded
tokens/s; training gains do not imply faster speculative serving.

## Compact offline teacher

Offline text EAGLE3 can project teacher targets in exact vocabulary chunks
Expand Down
1 change: 1 addition & 0 deletions examples/configs/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -277,6 +277,7 @@ Common fields:
| `training.max_steps` | `null` | Positive hard stop in optimizer steps. If it is set while `total_steps` is omitted, it is also the fallback schedule horizon. |
| `training.total_steps` | `null` | Positive optimizer/loss schedule horizon; it does not itself stop an online stream. A finite online disaggregated run may omit both fields: the producer publishes the exact horizon derived from prepared prompts, epochs, DP size, batch size, and accumulation. |
| `training.batch_size` | `1` | Per-rank microbatch size. P-EAGLE and USP require 1. |
| `training.sequence_packing` | `false` | Pack each original microbatch for text EAGLE3, DFlash, or DFlash2, offline or online. Requires FlexAttention; incompatible with compact teacher and loss-position trimming. DFlash/DFlash2 LK objectives are supported; EAGLE3 LK is not. See [sequence packing](../../docs/sections/basic_usage/training.md#sequence-packing). |
| `training.accumulation_steps` | `1` | Positive microbatches per optimizer update. |
| `training.fsdp_sharding` | `SHARD_GRAD_OP` | Trainer FSDP mode: `SHARD_GRAD_OP`, `FULL_SHARD`, or `NO_SHARD`. |
| `training.learning_rate` | `1e-4` | Positive peak learning rate. |
Expand Down
147 changes: 143 additions & 4 deletions specforge/algorithms/common/dflash_family_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,10 @@
from specforge.core.chunking import checkpointed_chunk_reduce
from specforge.modeling.draft.dflash import DFlashDraftModel
from specforge.modeling.draft.flex_attention_backend import flex_attention_backend
from specforge.modeling.packed_dflash import (
PackedDFlashLayout,
create_packed_dflash_block_mask,
)

try:
from torch.nn.attention.flex_attention import BlockMask, create_block_mask
Expand Down Expand Up @@ -252,6 +256,7 @@ def create_dflash_sdpa_mask(
block_size,
device,
sliding_window: Optional[int] = None,
context_start_positions: Optional[torch.Tensor] = None,
):
"""Construct a full or sliding dense boolean DFlash mask."""

Expand All @@ -274,6 +279,11 @@ def create_dflash_sdpa_mask(
)

mask_context = (kv_indices < S) & (kv_indices < anchor_expanded)
if context_start_positions is not None:
context_starts = context_start_positions.view(B, 1, N, 1).repeat_interleave(
block_size, dim=2
)
mask_context = mask_context & (kv_indices >= context_starts)
if sliding_window is not None:
# The current draft token occupies one slot in the window.
context_lower_bound = anchor_expanded + q_block_offsets - (sliding_window - 1)
Expand All @@ -300,6 +310,7 @@ def create_dflash_block_mask(
device: torch.device,
flex_block_size=None,
sliding_window: Optional[int] = None,
context_start_positions: Optional[torch.Tensor] = None,
):
"""Construct a full or sliding Flex Attention mask for DFlash training."""

Expand All @@ -316,6 +327,10 @@ def dflash_mask_mod(b, h, q_idx, kv_idx):
# Strictly less than: matches inference where target_hidden[anchor_pos]
# is not available as context.
mask_context = is_context & (kv_idx < anchor_pos)
if context_start_positions is not None:
mask_context = mask_context & (
kv_idx >= context_start_positions[b, safe_q_block_id]
)
if sliding_window is not None:
# The current draft token occupies one slot in the window.
context_lower_bound = anchor_pos + q_block_offset - (sliding_window - 1)
Expand All @@ -336,6 +351,18 @@ def dflash_mask_mod(b, h, q_idx, kv_idx):
Q_LEN = N * block_size
KV_LEN = S + N * block_size

if context_start_positions is not None:
return create_packed_dflash_block_mask(
anchor_positions,
block_keep_mask,
context_start_positions,
S,
block_size,
dflash_mask_mod,
block_size=flex_block_size if flex_block_size is not None else 128,
sliding_window=sliding_window,
)

kwargs = {}
if flex_block_size is not None:
kwargs["BLOCK_SIZE"] = flex_block_size
Expand Down Expand Up @@ -552,10 +579,13 @@ def _aligned_target_hidden(
self,
target_last_hidden_states: torch.Tensor,
safe_label_indices: torch.Tensor,
minimum_indices: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Gather the frozen target state that predicts each hard label."""

target_pred_indices = (safe_label_indices - 1).clamp(min=0)
if minimum_indices is not None:
target_pred_indices = torch.maximum(target_pred_indices, minimum_indices)
batch_size = target_last_hidden_states.shape[0]
hidden_size = target_last_hidden_states.shape[-1]
gather_indices = target_pred_indices.reshape(batch_size, -1, 1).expand(
Expand Down Expand Up @@ -700,25 +730,55 @@ def _forward_draft_blocks(
hidden_states: torch.Tensor,
loss_mask: torch.Tensor,
max_valid_anchors: Optional[int] = None,
packed_layout: Optional[PackedDFlashLayout] = None,
valid_anchor_counts: Optional[Tuple[int, ...]] = None,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
bsz, seq_len = input_ids.shape
device = input_ids.device

sampling_mask = (
loss_mask
if packed_layout is None
else packed_layout.padded_loss_mask(loss_mask)
)
anchor_positions, block_keep_mask = self._sample_anchor_positions(
seq_len,
loss_mask,
sampling_mask.shape[1],
sampling_mask,
device,
max_valid_anchors=max_valid_anchors,
)
local_anchor_positions = anchor_positions
if packed_layout is not None:
anchor_positions = packed_layout.pack_anchors(anchor_positions)
block_keep_mask = block_keep_mask.reshape(1, -1)
original_anchor_positions = anchor_positions
original_block_keep_mask = block_keep_mask
compact_indices = None
if packed_layout is not None and valid_anchor_counts is not None:
width = local_anchor_positions.shape[1]
compact_indices = packed_layout.compact_anchor_indices(
valid_anchor_counts, width
)
if compact_indices.numel() == 0:
raise ValueError("DFlash packing requires at least one valid anchor")
anchor_positions = anchor_positions.index_select(1, compact_indices)
block_keep_mask = block_keep_mask.index_select(1, compact_indices)
local_anchor_positions = local_anchor_positions.reshape(1, -1).index_select(
1, compact_indices
)

noise_embedding = self._create_noise_embed(
input_ids, anchor_positions, block_keep_mask
)

context_position_ids = (
torch.arange(seq_len, device=device).unsqueeze(0).expand(bsz, -1)
if packed_layout is None
else packed_layout.tokens.positions.unsqueeze(0)
)
draft_position_ids = self._create_position_ids(local_anchor_positions).reshape(
bsz, -1
)
draft_position_ids = self._create_position_ids(anchor_positions)
full_position_ids = torch.cat([context_position_ids, draft_position_ids], dim=1)

mask_builder = (
Expand All @@ -733,6 +793,10 @@ def _forward_draft_blocks(
"block_size": self.block_size,
"device": device,
}
if packed_layout is not None:
mask_args["context_start_positions"] = packed_layout.anchor_starts(
anchor_positions
)
if (
self.attention_backend == "flex_attention"
and flex_attention_backend() == "FLASH"
Expand Down Expand Up @@ -781,7 +845,17 @@ def _forward_draft_blocks(
attention_mask=dflash_attn_mask,
**draft_kwargs,
)
return anchor_positions, block_keep_mask, output_hidden
if compact_indices is not None:
query_rows = (
compact_indices[:, None] * self.block_size
+ torch.arange(self.block_size, device=device)
).reshape(-1)
output_hidden = output_hidden.new_zeros(
1,
original_anchor_positions.shape[1] * self.block_size,
output_hidden.shape[-1],
).index_copy(1, query_rows, output_hidden)
return original_anchor_positions, original_block_keep_mask, output_hidden

def _selector_chunk_terms(
self,
Expand Down Expand Up @@ -1495,6 +1569,8 @@ def forward(
max_valid_anchors: Optional[int] = None,
selector_loss_alpha: Optional[float] = None,
collect_detailed_metrics: bool = True,
sequence_lengths=None,
valid_anchor_counts: Optional[Tuple[int, ...]] = None,
) -> Tuple[torch.Tensor, torch.Tensor, Dict[str, object]]:
"""Parallel block-wise training forward pass; returns
(loss, accuracy, metrics) — same shape as Domino's forward."""
Expand All @@ -1505,18 +1581,51 @@ def forward(
bsz, seq_len = input_ids.shape
device = input_ids.device

packed_layout = None
if sequence_lengths is not None:
if self.attention_backend != "flex_attention":
raise ValueError(
"DFlash sequence packing currently requires flex_attention"
)
if (
bsz != 1
or hidden_states.shape[:2] != input_ids.shape
or loss_mask.shape != input_ids.shape
):
raise ValueError(
"DFlash sequence packing requires aligned single-row inputs"
)
packed_layout = PackedDFlashLayout.from_lengths(
sequence_lengths, seq_len, device
)

block_kwargs = (
{}
if packed_layout is None
else {
"packed_layout": packed_layout,
"valid_anchor_counts": valid_anchor_counts,
}
)
anchor_positions, block_keep_mask, output_hidden = self._forward_draft_blocks(
input_ids=input_ids,
hidden_states=hidden_states,
loss_mask=loss_mask,
max_valid_anchors=max_valid_anchors,
**block_kwargs,
)

# --- Labels: same-position prediction (position k predicts token anchor+k) ---
label_offsets = torch.arange(0, self.block_size, device=device).view(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=seq_len - 1)
if packed_layout is not None:
document_ends = packed_layout.anchor_ends(anchor_positions).unsqueeze(-1)
valid_label_mask = valid_label_mask & (label_indices < document_ends)
# Even masked tail labels and selector predecessors stay inside their
# document; none can import a neighboring document's token IDs.
safe_label_indices = torch.minimum(safe_label_indices, document_ends - 1)

target_ids = torch.gather(
input_ids.unsqueeze(1).expand(-1, anchor_positions.size(1), -1),
Expand Down Expand Up @@ -1554,6 +1663,15 @@ def forward(
self._aligned_target_hidden(
target_last_hidden_states,
safe_label_indices,
**(
{
"minimum_indices": packed_layout.anchor_starts(
anchor_positions
).unsqueeze(-1)
}
if packed_layout is not None
else {}
),
)
if (
target_last_hidden_states is not None
Expand All @@ -1562,6 +1680,27 @@ def forward(
)
else None
)
if packed_layout is not None:
# Preserve the original [documents, anchors, block] reduction axes.
# D-PACE normalizes anchors per document; selector and walk metrics
# also use these axes. Packing changes only the backbone execution.
original_batch = len(packed_layout.lengths)
anchors_per_document = anchor_positions.shape[1] // original_batch
shape = (original_batch, anchors_per_document, self.block_size)
hidden_4d = hidden_4d.reshape(*shape, hidden_4d.shape[-1])
target_ids = target_ids.reshape(shape)
predecessor_ids = predecessor_ids.reshape(shape)
weight_mask = weight_mask.reshape(shape)
if aligned_target_hidden is not None:
aligned_target_hidden = aligned_target_hidden.reshape(
*shape, aligned_target_hidden.shape[-1]
)
anchor_positions = packed_layout.tokens.positions[anchor_positions].reshape(
original_batch, anchors_per_document
)
block_keep_mask = block_keep_mask.reshape(
original_batch, anchors_per_document
)
sequence_anchor_scale = None
if self.loss_type in _DPACE_LOSS_TYPES:
sequence_anchor_scale = self._sequence_anchor_scale(weight_mask)
Expand Down
43 changes: 43 additions & 0 deletions specforge/algorithms/common/hidden_states_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,47 @@ def build_collator():
)


def build_packed_collator():
"""Pack DFlash/DFlash2 context features after each capture is materialized."""
return PackedHiddenStatesCollator()


class PackedHiddenStatesCollator:
def __call__(self, features):
import torch

if not features:
raise ValueError("cannot pack an empty feature batch")
required = ("input_ids", "loss_mask", "hidden_states")
keys = list(required)
teacher_key = "target_last_hidden_states"
present = [teacher_key in feature for feature in features]
if any(present) and not all(present):
raise KeyError(
f"optional feature {teacher_key!r} must be present in every sample or none"
)
if all(present):
keys.append(teacher_key)
lengths = []
for feature in features:
missing = set(keys) - feature.keys()
if missing:
raise KeyError(f"packed sample is missing features: {sorted(missing)}")
ids = feature["input_ids"]
if ids.ndim != 2 or ids.shape[0] != 1 or ids.shape[1] == 0:
raise ValueError("packing requires nonempty [1, length] input_ids")
length = ids.shape[1]
for key in keys:
tensor = feature[key]
ndim = 2 if key in ("input_ids", "loss_mask") else 3
if tensor.ndim != ndim or tensor.shape[:2] != (1, length):
raise ValueError(f"packing requires aligned unbatched {key}")
lengths.append(length)
batch = {key: torch.cat([f[key] for f in features], dim=1) for key in keys}
batch["sequence_lengths"] = torch.tensor(lengths, dtype=torch.long)
return batch


def build_dspark_collator():
return _padded_collator(
("input_ids", "loss_mask", "hidden_states", "target_last_hidden_states")
Expand Down Expand Up @@ -252,6 +293,7 @@ def build_mtp_collator():
"DSPARK_NORMALIZER_ID",
"MTP_NORMALIZER_ID",
"NORMALIZER_ID",
"PackedHiddenStatesCollator",
"build_collator",
"build_dspark_collator",
"build_dspark_offline_normalizer",
Expand All @@ -261,6 +303,7 @@ def build_mtp_collator():
"build_mtp_offline_reader",
"build_offline_normalizer",
"build_offline_reader",
"build_packed_collator",
"normalize_dspark_offline_sample",
"normalize_mtp_offline_sample",
"normalize_offline_sample",
Expand Down
Loading