diff --git a/docs/sections/basic_usage/training.md b/docs/sections/basic_usage/training.md index c7c055b63..6d55ecc06 100644 --- a/docs/sections/basic_usage/training.md +++ b/docs/sections/basic_usage/training.md @@ -549,6 +549,39 @@ a complete checkpoint and points `-best` at it, even when `training.save_interval` is zero. `-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 diff --git a/examples/configs/README.md b/examples/configs/README.md index 905c4eb10..3628ff700 100644 --- a/examples/configs/README.md +++ b/examples/configs/README.md @@ -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. | diff --git a/specforge/algorithms/common/dflash_family_model.py b/specforge/algorithms/common/dflash_family_model.py index 719b12232..af8ee339d 100644 --- a/specforge/algorithms/common/dflash_family_model.py +++ b/specforge/algorithms/common/dflash_family_model.py @@ -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 @@ -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.""" @@ -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) @@ -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.""" @@ -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) @@ -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 @@ -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( @@ -700,16 +730,42 @@ 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 @@ -717,8 +773,12 @@ def _forward_draft_blocks( 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 = ( @@ -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" @@ -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, @@ -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.""" @@ -1505,11 +1581,38 @@ 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) --- @@ -1517,6 +1620,12 @@ def forward( 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), @@ -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 @@ -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) diff --git a/specforge/algorithms/common/hidden_states_data.py b/specforge/algorithms/common/hidden_states_data.py index 82d7fd63d..1549d59e3 100644 --- a/specforge/algorithms/common/hidden_states_data.py +++ b/specforge/algorithms/common/hidden_states_data.py @@ -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") @@ -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", @@ -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", diff --git a/specforge/algorithms/common/providers.py b/specforge/algorithms/common/providers.py index d0c22f702..9edc83119 100644 --- a/specforge/algorithms/common/providers.py +++ b/specforge/algorithms/common/providers.py @@ -491,10 +491,15 @@ class OfflineDataProvider: build_normalizer: Factory build_collator: Factory capture_layout: OfflineCaptureLayout | None = None + build_packed_collator: Factory | None = None def __post_init__(self) -> None: _non_empty(self.modality, field_name="modality") _non_empty(self.normalizer_id, field_name="normalizer_id") + if self.build_packed_collator is not None and not callable( + self.build_packed_collator + ): + raise TypeError("build_packed_collator must be callable or None") if self.capture_layout is not None and not isinstance( self.capture_layout, OfflineCaptureLayout, @@ -581,6 +586,7 @@ class ServerStreamingProvider: build_collator: Factory build_input_adapter: Factory | None = None select_layout: Factory | None = None + build_packed_collator: Factory | None = None def __post_init__(self) -> None: _non_empty(self.modality, field_name="modality") @@ -594,6 +600,10 @@ def __post_init__(self) -> None: raise TypeError("layout must be a ServerCaptureLayout") if not callable(self.build_collator): raise TypeError("build_collator must be callable") + if self.build_packed_collator is not None and not callable( + self.build_packed_collator + ): + raise TypeError("build_packed_collator must be callable or None") if self.build_input_adapter is not None and not callable( self.build_input_adapter ): diff --git a/specforge/algorithms/contracts.py b/specforge/algorithms/contracts.py index 4d323a128..e3a8df111 100644 --- a/specforge/algorithms/contracts.py +++ b/specforge/algorithms/contracts.py @@ -244,6 +244,7 @@ class AlgorithmCapabilities: #: ``training.dflash_teacher_metrics=false`` can drop the teacher-only #: final hidden state from the streaming capture. supports_teacher_metrics_opt_out: bool = False + supports_packed_lk_loss: bool = False def __post_init__(self) -> None: attention_backends = _normalized_names( @@ -262,6 +263,7 @@ def __post_init__(self) -> None: "supports_vocab_mapping", "allows_aux_layer_override", "supports_teacher_metrics_opt_out", + "supports_packed_lk_loss", ): if not isinstance(getattr(self, field_name), bool): raise TypeError(f"{field_name} must be a bool") diff --git a/specforge/algorithms/dflash/providers.py b/specforge/algorithms/dflash/providers.py index ff5a824bf..c1a9044f0 100644 --- a/specforge/algorithms/dflash/providers.py +++ b/specforge/algorithms/dflash/providers.py @@ -14,6 +14,7 @@ build_collator, build_offline_normalizer, build_offline_reader, + build_packed_collator, ) from specforge.algorithms.common.providers import ( AlgorithmProviders, @@ -209,6 +210,7 @@ def algorithm_spec() -> AlgorithmSpec: capabilities=AlgorithmCapabilities( attention_backends={"eager", "sdpa", "flex_attention"}, supports_teacher_metrics_opt_out=True, + supports_packed_lk_loss=True, ), ) @@ -260,6 +262,7 @@ def algorithm_providers() -> AlgorithmProviders: build_reader=partial(build_offline_reader, ALGORITHM_NAME), build_normalizer=build_offline_normalizer, build_collator=collator, + build_packed_collator=build_packed_collator, ), ), server_streaming=( @@ -270,6 +273,7 @@ def algorithm_providers() -> AlgorithmProviders: layout=SERVER_CAPTURE_LAYOUT, build_collator=collator, select_layout=select_server_capture_layout, + build_packed_collator=build_packed_collator, ), ), ) diff --git a/specforge/algorithms/eagle3/data.py b/specforge/algorithms/eagle3/data.py index 7521fcfeb..c0979d77b 100644 --- a/specforge/algorithms/eagle3/data.py +++ b/specforge/algorithms/eagle3/data.py @@ -87,17 +87,86 @@ def build_offline_collator(): return DataCollatorWithPadding() +class DataCollatorWithPacking: + """Pack one logical microbatch without changing its samples or loss weight. + + Each sample is a normalized, unpadded text feature with batch dimension one. + The model uses ``sequence_lengths`` for attention and TTT boundaries. Keeping + the original padded denominator makes this an execution optimization, not + a change to the EAGLE3 objective or optimizer schedule. + """ + + def __call__(self, features): + import torch + + if not features: + raise ValueError("cannot pack an empty feature batch") + keys = ("input_ids", "loss_mask", "hidden_state", "target", "attention_mask") + 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 = 3 if key in ("hidden_state", "target") else 2 + if tensor.ndim != ndim or tensor.shape[:2] != (1, length): + raise ValueError(f"packing requires aligned unbatched {key}") + if not bool((feature["attention_mask"] == 1).all()): + raise ValueError("packing requires unpadded samples") + if "position_ids" in feature: + expected = torch.arange(length, device=ids.device).unsqueeze(0) + if not torch.equal(feature["position_ids"], expected): + raise ValueError("packing supports standard text position_ids only") + lengths.append(length) + + batch = {key: torch.cat([f[key] for f in features], dim=1) for key in keys} + batch["position_ids"] = torch.cat( + [torch.arange(n, device=batch["input_ids"].device) for n in lengths] + ).unsqueeze(0) + # These small descriptors stay on the host until the strategy builds + # the device layout; no GPU scalar synchronization is needed. + batch["sequence_lengths"] = torch.tensor(lengths, dtype=torch.long) + batch["loss_denominator"] = torch.tensor( + len(lengths) * max(lengths), dtype=torch.long + ) + return batch + + +def build_packed_collator(): + return DataCollatorWithPacking() + + def build_server_collator(): from specforge.algorithms.common.collation import concatenate_features return concatenate_features +def build_padded_server_collator(): + """Accept ragged, unshifted EAGLE3 features from separate capture requests.""" + from specforge.algorithms.common.collation import pad_and_concatenate_features + + keys = ("input_ids", "attention_mask", "loss_mask", "hidden_state", "target") + return partial( + pad_and_concatenate_features, + sequence_axes={key: 1 for key in keys}, + required_keys=keys, + ) + + __all__ = [ + "DataCollatorWithPacking", "NORMALIZER_ID", "build_offline_collator", "build_offline_normalizer", "build_offline_reader", + "build_packed_collator", + "build_padded_server_collator", "build_server_collator", "normalize_offline_sample", ] diff --git a/specforge/algorithms/eagle3/model.py b/specforge/algorithms/eagle3/model.py index d89cf5d0c..bfcd21172 100644 --- a/specforge/algorithms/eagle3/model.py +++ b/specforge/algorithms/eagle3/model.py @@ -22,7 +22,7 @@ """EAGLE3 training model implementation.""" -from typing import Callable, List, Optional, Tuple +from typing import Callable, List, Optional, Tuple, Union import torch import torch.nn as nn @@ -37,6 +37,7 @@ from specforge.core.lk_loss import compute_acceptance_rate, compute_lk_loss from specforge.core.loss import LogSoftmaxLoss from specforge.modeling.draft import Eagle3DraftModel +from specforge.modeling.packed_sequence import PackedSequenceLayout from specforge.utils import padding @@ -262,6 +263,8 @@ def forward( target_head_weight: Optional[torch.Tensor] = None, compact_teacher_chunk_size: int = DEFAULT_VOCAB_CHUNK_SIZE, trim_loss_positions: bool = False, + sequence_lengths: Optional[Union[torch.Tensor, Tuple[int, ...]]] = None, + loss_denominator: Optional[Union[torch.Tensor, int]] = None, ) -> Tuple[ List[torch.Tensor], List[torch.Tensor], @@ -285,7 +288,42 @@ def forward( states in draft-vocab space and ``target`` is ignored. trim_loss_positions: compute the teacher, draft logits and loss only at supervised positions when the batch/objective supports it. + sequence_lengths: CPU int64 document lengths (or a Python tuple) for a packed single row. + The strategy has already shifted input and target fields per document. + loss_denominator: optional CPU scalar equal to document count times + the longest document, preserving the padded batch's mean loss. """ + packed_layout = None + if sequence_lengths is not None: + if self.attention_backend != "flex_attention": + raise ValueError("sequence packing currently requires flex_attention") + if ( + trim_loss_positions + or target_hidden_for_compact is not None + or self.lk_loss_type is not None + ): + raise ValueError( + "sequence packing does not support trim, compact teacher, or LK loss" + ) + if input_ids.shape[0] != 1 or hidden_states.shape[:2] != input_ids.shape: + raise ValueError("sequence packing requires one concatenated batch row") + packed_layout = PackedSequenceLayout.from_lengths( + sequence_lengths, input_ids.shape[1], hidden_states.device + ) + if loss_denominator is not None: + if isinstance(loss_denominator, torch.Tensor) and ( + loss_denominator.device.type != "cpu" + or loss_denominator.numel() != 1 + ): + raise ValueError("loss_denominator must be a scalar CPU tensor") + if int(loss_denominator) != packed_layout.padded_denominator: + raise ValueError( + "loss_denominator must equal document_count * longest_document" + ) + if position_ids is None: + position_ids = packed_layout.positions.unsqueeze(0) + elif loss_denominator is not None: + raise ValueError("loss_denominator requires sequence_lengths") adapter = self._make_adapter() # Step 1: handle vocab size if target_hidden_for_compact is not None: @@ -385,7 +423,9 @@ def forward( dtype=torch.bool, device=hidden_states.device, ) - if self.attention_backend == "sdpa": + if packed_layout is not None: + attention_mask = packed_layout + elif self.attention_backend == "sdpa": attention_mask = self.draft_model.prepare_decoder_attention_mask( attention_mask=attention_mask, hidden_states=hidden_states, @@ -527,6 +567,16 @@ def forward( position_mask=state.position_mask, loss_mask=state.loss_mask, adapter=adapter, + loss_scale=( + seq_length / packed_layout.padded_denominator + if packed_layout is not None + else 1.0 + ), + full_positions=( + packed_layout.padded_denominator + if packed_layout is not None + else None + ), ) acces.append(acc) acceptance_rates.append(acceptance_rate) @@ -538,9 +588,14 @@ def forward( if not is_last: # Step 5.7: we need to update the loss mask - global_input_ids = padding(global_input_ids, left=False) - position_mask = padding(position_mask, left=False) - loss_mask = padding(loss_mask, left=False) + if packed_layout is None: + global_input_ids = padding(global_input_ids, left=False) + position_mask = padding(position_mask, left=False) + loss_mask = padding(loss_mask, left=False) + else: + global_input_ids = packed_layout.shift_left(global_input_ids) + position_mask = packed_layout.shift_left(position_mask) + loss_mask = packed_layout.shift_left(loss_mask) # Flex attention mask shirnking is handled inside attention module return ( plosses, diff --git a/specforge/algorithms/eagle3/providers.py b/specforge/algorithms/eagle3/providers.py index 835be387c..c2ecea4e8 100644 --- a/specforge/algorithms/eagle3/providers.py +++ b/specforge/algorithms/eagle3/providers.py @@ -32,7 +32,8 @@ build_offline_collator, build_offline_normalizer, build_offline_reader, - build_server_collator, + build_packed_collator, + build_padded_server_collator, ) ALGORITHM_NAME = "eagle3" @@ -211,6 +212,7 @@ def algorithm_providers() -> AlgorithmProviders: build_reader=build_offline_reader, build_normalizer=build_offline_normalizer, build_collator=build_offline_collator, + build_packed_collator=build_packed_collator, ), ), server_streaming=( @@ -227,7 +229,8 @@ def algorithm_providers() -> AlgorithmProviders: ), attention_mask_feature="attention_mask", ), - build_collator=build_server_collator, + build_collator=build_padded_server_collator, + build_packed_collator=build_packed_collator, ), ), vocab_mapping_modes=frozenset({FeatureMode.OFFLINE, FeatureMode.STREAMING}), diff --git a/specforge/application/planning.py b/specforge/application/planning.py index 8298b8420..ac5d25040 100644 --- a/specforge/application/planning.py +++ b/specforge/application/planning.py @@ -78,6 +78,24 @@ def _validate_algorithm_capabilities( ) -> None: capabilities = algorithm.spec.capabilities training = cfg.training + if training.sequence_packing: + provider = ( + algorithm.providers.offline_for(cfg.model.input_modality) + if mode is FeatureMode.OFFLINE + else algorithm.providers.server_streaming_for(cfg.model.input_modality) + ) + if provider.build_packed_collator is None: + raise ValueError( + f"algorithm {algorithm.name!r} does not support training.sequence_packing " + f"for modality {cfg.model.input_modality!r}" + ) + if ( + training.lk_loss_type is not None + and not capabilities.supports_packed_lk_loss + ): + raise ValueError( + f"algorithm {algorithm.name!r} does not support sequence_packing with lk_loss_type" + ) if training.attention_backend not in capabilities.attention_backends: raise ValueError( f"algorithm {algorithm.name!r} does not support attention_backend=" diff --git a/specforge/config/schema.py b/specforge/config/schema.py index 32eed42cf..9685bbe5a 100644 --- a/specforge/config/schema.py +++ b/specforge/config/schema.py @@ -898,6 +898,9 @@ class TrainingConfig(StrictConfigModel): max_steps: Optional[int] = Field(default=None, gt=0) total_steps: Optional[int] = Field(default=None, gt=0) batch_size: int = Field(default=1, gt=0) + #: Concatenate the samples of each microbatch, retaining document boundaries + #: and the original loss normalization. Text EAGLE3/DFlash/DFlash2 + FlexAttention. + sequence_packing: bool = False accumulation_steps: int = Field(default=1, gt=0) fsdp_sharding: Literal["SHARD_GRAD_OP", "FULL_SHARD", "NO_SHARD"] = "SHARD_GRAD_OP" learning_rate: float = Field(default=1e-4, gt=0.0) @@ -991,6 +994,14 @@ class TrainingConfig(StrictConfigModel): @model_validator(mode="after") def _validate_training_shape(self): + if self.sequence_packing: + if self.attention_backend != "flex_attention": + raise ValueError("training.sequence_packing requires flex_attention") + if self.compact_teacher or self.trim_loss_positions: + raise ValueError( + "training.sequence_packing currently requires compact_teacher=false, " + "trim_loss_positions=false" + ) if not 0.0 <= self.dpace_alpha <= 1.0: raise ValueError("training.dpace_alpha must be in [0, 1]") if not 0.0 < self.down_sample_ratio <= 1.0: diff --git a/specforge/launch.py b/specforge/launch.py index 030e5388f..02a798ff4 100644 --- a/specforge/launch.py +++ b/specforge/launch.py @@ -163,10 +163,19 @@ def _offline_io( *, ttt_length: int, use_usp_preprocess: bool, + sequence_packing: bool = False, ): """Resolve the algorithm-owned normalizer and collator for one modality.""" provider = algorithm.providers.offline_for(modality) - return provider.build_collator(), provider.build_normalizer( + if sequence_packing: + if use_usp_preprocess or provider.build_packed_collator is None: + raise ValueError( + "sequence_packing requires a supported non-USP offline provider" + ) + collator = provider.build_packed_collator() + else: + collator = provider.build_collator() + return collator, provider.build_normalizer( max_len, ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, @@ -252,6 +261,7 @@ def _make_offline_eval_data_factory( ttt_length: int, use_usp_preprocess: bool, dataloader_num_workers: int, + sequence_packing: bool = False, ): """Build a fresh re-iterable eval loader over the offline feature path.""" provider = algorithm.providers.offline_for(modality) @@ -261,6 +271,7 @@ def _make_offline_eval_data_factory( max_len, ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, + sequence_packing=sequence_packing, ) eval_run_id = f"{run_id}-eval" refs = provider.build_reader( @@ -293,11 +304,22 @@ def _streaming_collate( algorithm: AlgorithmRegistration, modality: str, collate_fn, + *, + sequence_packing: bool = False, ): """Resolve an algorithm-owned server-streaming collator.""" if collate_fn is not None: + if sequence_packing: + raise ValueError( + "sequence_packing cannot be combined with a custom collate_fn" + ) return collate_fn - return algorithm.providers.server_streaming_for(modality).build_collator() + provider = algorithm.providers.server_streaming_for(modality) + if sequence_packing: + if provider.build_packed_collator is None: + raise ValueError("sequence_packing requires a supported streaming provider") + return provider.build_packed_collator() + return provider.build_collator() def _resolve_metadata_store( @@ -563,6 +585,7 @@ def build_offline_runtime( sp_ulysses_size: int = 1, sp_ring_size: int = 1, use_usp_preprocess: bool = False, + sequence_packing: bool = False, seed: int = 0, logger=None, log_interval: int = 50, @@ -586,6 +609,7 @@ def build_offline_runtime( max_len, ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, + sequence_packing=sequence_packing, ) controller = DataFlowController( run_id, @@ -621,6 +645,7 @@ def refs_for_epoch(epoch): ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, dataloader_num_workers=dataloader_num_workers, + sequence_packing=sequence_packing, ) return _assemble_trainer( algorithm=algorithm, @@ -689,6 +714,7 @@ def build_disagg_offline_runtime( sp_ulysses_size: int = 1, sp_ring_size: int = 1, use_usp_preprocess: bool = False, + sequence_packing: bool = False, seed: int = 0, logger=None, log_interval: int = 50, @@ -711,6 +737,7 @@ def build_disagg_offline_runtime( max_len, ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, + sequence_packing=sequence_packing, ) source_refs = list(refs) @@ -743,6 +770,7 @@ def refs_for_epoch(epoch): ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, dataloader_num_workers=dataloader_num_workers, + sequence_packing=sequence_packing, ) return _assemble_trainer( algorithm=algorithm, @@ -1536,6 +1564,7 @@ def build_disagg_online_consumer( eval_interval: int = 0, eval_data_factory=None, collate_fn=None, + sequence_packing: bool = False, idle_timeout_s: Optional[float] = None, metadata_store: Optional[MetadataStore] = None, metadata_db_path: Optional[str] = None, @@ -1888,7 +1917,9 @@ def stop_distributor_and_drain() -> None: eval_data_factory=eval_data_factory, logger=logger, log_interval=log_interval, - collate_fn=_streaming_collate(algorithm, modality, collate_fn), + collate_fn=_streaming_collate( + algorithm, modality, collate_fn, sequence_packing=sequence_packing + ), strategy_kwargs=strategy_kwargs, per_sample_transform=None, max_checkpoints=max_checkpoints, diff --git a/specforge/modeling/draft/llama3_eagle.py b/specforge/modeling/draft/llama3_eagle.py index 67276ec5c..6fe355c31 100644 --- a/specforge/modeling/draft/llama3_eagle.py +++ b/specforge/modeling/draft/llama3_eagle.py @@ -16,6 +16,10 @@ compile_friendly_flex_attention, generate_eagle3_mask, ) +from specforge.modeling.packed_sequence import ( + PackedSequenceLayout, + generate_packed_eagle3_mask, +) from specforge.utils import print_with_rank from ...distributed import get_sp_ring_group, get_sp_ulysses_group @@ -761,7 +765,12 @@ def forward( self.rope_scaling["mrope_section"], ) else: - cos, sin = self.rotary_emb(query_states, seq_len=q_len + lck) + rope_length = ( + attention_mask.maximum_length + if isinstance(attention_mask, PackedSequenceLayout) + else q_len + ) + cos, sin = self.rotary_emb(query_states, seq_len=rope_length + lck) cos, sin = cos.to(query_states.device), sin.to(query_states.device) # Keep positions ids aligned when padding so the KV cache is unaffected. query_states, key_states = apply_rotary_pos_emb( @@ -780,10 +789,16 @@ def forward( cache_kwargs=cache_kwargs, ) - seq_lengths = attention_mask.sum(dim=-1) - # Shrink the attention mask to align with the padding to the right. - # This is equivalent to the shrinking logic in eagle3.py - seq_lengths -= lck + if isinstance(attention_mask, PackedSequenceLayout): + mask_mod = generate_packed_eagle3_mask(attention_mask, q_len, lck) + else: + seq_lengths = attention_mask.sum(dim=-1) - lck + mask_mod = generate_eagle3_mask( + seq_lengths=seq_lengths, + Q_LEN=q_len, + KV_LEN=key_cache.shape[-2], + lck=lck, + ) # TODO: Remove the usage of uncompiled create_block_mask after # https://github.com/pytorch/pytorch/issues/160018 if q_len <= 128: @@ -794,12 +809,7 @@ def forward( flex_attention_func = compile_friendly_flex_attention block_mask = create_block_mask_func( - mask_mod=generate_eagle3_mask( - seq_lengths=seq_lengths, - Q_LEN=q_len, - KV_LEN=key_cache.shape[-2], - lck=lck, - ), + mask_mod=mask_mod, B=bsz, H=1, # Rely on broadcast Q_LEN=q_len, diff --git a/specforge/modeling/packed_dflash.py b/specforge/modeling/packed_dflash.py new file mode 100644 index 000000000..bda1c8ba2 --- /dev/null +++ b/specforge/modeling/packed_dflash.py @@ -0,0 +1,158 @@ +"""Index metadata for packing DFlash context while preserving sampled blocks.""" + +from dataclasses import dataclass + +import torch + +from specforge.modeling.packed_sequence import PackedSequenceLayout + + +@dataclass(frozen=True) +class PackedDFlashLayout: + tokens: PackedSequenceLayout + lengths: tuple[int, ...] + document_starts: torch.Tensor + padded_indices: torch.Tensor + padded_valid: torch.Tensor + + @classmethod + def from_lengths(cls, lengths, sequence_length, device): + tokens = PackedSequenceLayout.from_lengths(lengths, sequence_length, device) + lengths = torch.as_tensor(lengths, dtype=torch.long) + starts = lengths.cumsum(0) - lengths + columns = torch.arange(tokens.maximum_length) + indices = starts[:, None] + columns[None, :] + return cls( + tokens=tokens, + lengths=tuple(lengths.tolist()), + document_starts=starts.to(device), + padded_indices=indices.clamp(max=sequence_length - 1).to(device), + padded_valid=(columns[None, :] < lengths[:, None]).to(device), + ) + + def padded_loss_mask(self, loss_mask): + """Rebuild only the small sampling mask, preserving baseline RNG shape.""" + return loss_mask[0, self.padded_indices] * self.padded_valid + + def pack_anchors(self, local_anchors): + return (local_anchors + self.document_starts[:, None]).reshape(1, -1) + + def anchor_starts(self, packed_anchors): + return packed_anchors - self.tokens.positions[packed_anchors] + + def anchor_ends(self, packed_anchors): + return ( + self.anchor_starts(packed_anchors) + + self.tokens.document_lengths[packed_anchors] + ) + + def compact_anchor_indices(self, valid_anchor_counts, width): + """Indices of valid prefix slots after the unchanged sorted sampler.""" + if len(valid_anchor_counts) != len(self.lengths) or any( + not isinstance(count, int) or count < 0 for count in valid_anchor_counts + ): + raise ValueError( + "valid_anchor_counts must contain one nonnegative integer per document" + ) + return torch.tensor( + [ + document * width + index + for document, count in enumerate(valid_anchor_counts) + for index in range(min(count, width)) + ], + dtype=torch.long, + device=self.document_starts.device, + ) + + +def create_packed_dflash_block_mask( + anchor_positions, + block_keep_mask, + context_start_positions, + context_length, + proposal_size, + mask_mod, + *, + block_size=128, + sliding_window=None, +): + """Construct sparse tiles without materializing a quadratic token mask. + + A tile's context union is bounded by its earliest start and latest anchor; + its intersection identifies fully allowed context tiles. A conservative + union may include extra partial tiles, whose exact token predicate remains + ``mask_mod``. Draft tiles intersect only the proposals present in a Q tile. + """ + from torch.nn.attention.flex_attention import BlockMask + + q_block, kv_block = ( + (block_size, block_size) if isinstance(block_size, int) else block_size + ) + batch, anchors = anchor_positions.shape + q_length = anchors * proposal_size + kv_length = context_length + q_length + device = anchor_positions.device + q = torch.arange( + ((q_length + q_block - 1) // q_block) * q_block, device=device + ).reshape(-1, q_block) + anchor_index = (q // proposal_size).clamp(max=anchors - 1) + valid = (q < q_length).unsqueeze(0) & block_keep_mask[:, anchor_index] + anchor = anchor_positions[:, anchor_index] + lower = context_start_positions[:, anchor_index] + if sliding_window is not None: + lower = torch.maximum( + lower, anchor + q.remainder(proposal_size) - (sliding_window - 1) + ) + + context_min = torch.where(valid, lower, context_length).amin(-1) + context_max = torch.where(valid, anchor, 0).amax(-1) + intersection_min = torch.where(valid, lower, 0).amax(-1) + intersection_max = torch.where(valid, anchor, context_length).amin(-1) + any_valid = valid.any(-1) + all_valid = valid.all(-1) + + kv_start = ( + torch.arange((kv_length + kv_block - 1) // kv_block, device=device) * kv_block + ) + kv_end = kv_start + kv_block + context_tiles = ( + (kv_start < context_max[..., None]) + & (kv_end > context_min[..., None]) + & (kv_start < context_length) + & (context_min < context_max)[..., None] + ) + full_tiles = ( + all_valid[..., None] + & (kv_start >= intersection_min[..., None]) + & (kv_end <= intersection_max[..., None]) + & (kv_end <= context_length) + ) + draft_min = context_length + (q[:, 0] // proposal_size) * proposal_size + draft_max = context_length + torch.minimum( + (q[:, -1] // proposal_size + 1) * proposal_size, + q.new_tensor(q_length), + ) + draft_tiles = (kv_start < draft_max[:, None]) & (kv_end > draft_min[:, None]) + partial_tiles = any_valid[..., None] & (context_tiles | draft_tiles) & ~full_tiles + + def ordered(mask): + mask = mask.unsqueeze(1) + counts = mask.sum(-1, dtype=torch.int32) + indices = ( + mask.to(torch.int32) + .argsort(dim=-1, descending=True, stable=True) + .to(torch.int32) + ) + return counts, indices + + partial_counts, partial_indices = ordered(partial_tiles) + full_counts, full_indices = ordered(full_tiles) + return BlockMask.from_kv_blocks( + partial_counts, + partial_indices, + full_counts, + full_indices, + BLOCK_SIZE=(q_block, kv_block), + mask_mod=mask_mod, + seq_lengths=(q_length, kv_length), + ) diff --git a/specforge/modeling/packed_sequence.py b/specforge/modeling/packed_sequence.py new file mode 100644 index 000000000..2706f24bb --- /dev/null +++ b/specforge/modeling/packed_sequence.py @@ -0,0 +1,76 @@ +"""Document boundaries for packed, single-row EAGLE3 training batches.""" + +from dataclasses import dataclass + +import torch + + +@dataclass(frozen=True) +class PackedSequenceLayout: + """Token metadata shared by segment shifts and the Flex Attention mask. + + Lengths are small CPU metadata; token metadata lives beside the model's + tensors. Padding remains local to each document during every TTT shift. + """ + + document_ids: torch.Tensor + positions: torch.Tensor + document_lengths: torch.Tensor + padded_denominator: int + maximum_length: int + + @classmethod + def from_lengths(cls, lengths, sequence_length: int, device): + if not isinstance(lengths, torch.Tensor): + lengths = torch.tensor(lengths, dtype=torch.long) + if lengths.device.type != "cpu" or lengths.dtype != torch.long: + raise ValueError("sequence_lengths must be a CPU int64 tensor") + if lengths.ndim != 1 or not lengths.numel() or bool((lengths <= 0).any()): + raise ValueError("sequence_lengths must contain positive document lengths") + if int(lengths.sum()) != sequence_length: + raise ValueError("sequence_lengths must sum to the packed sequence length") + document_ids = torch.repeat_interleave(torch.arange(lengths.numel()), lengths) + starts = lengths.cumsum(0) - lengths + positions = torch.arange(sequence_length) - starts[document_ids] + return cls( + document_ids=document_ids.to(device), + positions=positions.to(device), + document_lengths=lengths[document_ids].to(device), + padded_denominator=int(lengths.numel() * lengths.max()), + maximum_length=int(lengths.max()), + ) + + def shift_left(self, tensor: torch.Tensor) -> torch.Tensor: + """Drop each document's first entry and append one zero to that document.""" + if tensor.shape[:2] != (1, self.positions.numel()): + raise ValueError( + "packed tensors must have shape [1, sum(sequence_lengths), ...]" + ) + shifted = torch.cat((tensor[:, 1:], torch.zeros_like(tensor[:, -1:])), dim=1) + valid = (self.positions + 1 < self.document_lengths).to(tensor.device) + return shifted * valid.reshape(1, -1, *([1] * (tensor.ndim - 2))) + + +def generate_packed_eagle3_mask( + layout: PackedSequenceLayout, query_length: int, depth: int +): + """Match independent EAGLE3 causal prefixes and diagonal TTT cache suffixes.""" + document_ids = layout.document_ids + positions = layout.positions + lengths = layout.document_lengths + + def mask_mod(_b, _h, q_idx, kv_idx): + # Flex evaluates complete tiles, including indices outside the actual Q/KV + # sizes. Clamp metadata reads; the explicit bounds mask removes those cells. + safe_q = q_idx.clamp(max=query_length - 1) + kv_row = kv_idx % query_length + same_document = document_ids[safe_q] == document_ids[kv_row] + valid_query = (q_idx < query_length) & ( + positions[safe_q] < lengths[safe_q] - depth + ) + valid_key = positions[kv_row] < lengths[kv_row] - depth + causal = (kv_idx < query_length) & (q_idx >= kv_idx) + suffix = (kv_idx >= query_length) & (kv_row == q_idx) + return same_document & valid_query & valid_key & (causal | suffix) + + return mask_mod diff --git a/specforge/training/assembly.py b/specforge/training/assembly.py index 07e63405c..26cca714f 100644 --- a/specforge/training/assembly.py +++ b/specforge/training/assembly.py @@ -645,6 +645,7 @@ def build_training_run( max_len=cfg.data.max_length, num_epochs=t.num_epochs, use_usp_preprocess=(t.attention_backend == "usp"), + sequence_packing=t.sequence_packing, seed=t.seed, resume_from=t.resume_from, **_common_launch_kwargs( diff --git a/specforge/training/disaggregated.py b/specforge/training/disaggregated.py index b30eb274c..b91c9d787 100644 --- a/specforge/training/disaggregated.py +++ b/specforge/training/disaggregated.py @@ -559,6 +559,7 @@ def produce() -> int: sp_ulysses_size=cfg.training.sp_ulysses_size, sp_ring_size=cfg.training.sp_ring_size, use_usp_preprocess=(cfg.training.attention_backend == "usp"), + sequence_packing=cfg.training.sequence_packing, seed=cfg.training.seed, dataloader_num_workers=_dataloader_num_workers(cfg, algorithm), profiling_options=_profiling_options(cfg), @@ -820,6 +821,7 @@ def produce() -> int: run_id=cfg.run_id, output_dir=cfg.output_dir, batch_size=cfg.training.batch_size, + sequence_packing=cfg.training.sequence_packing, accumulation_steps=cfg.training.accumulation_steps, max_steps=cfg.training.max_steps, total_steps=total_steps, diff --git a/specforge/training/strategies/base.py b/specforge/training/strategies/base.py index f592ccdc0..0f84981e1 100644 --- a/specforge/training/strategies/base.py +++ b/specforge/training/strategies/base.py @@ -70,7 +70,21 @@ def linear_lambda_base( return max(0.0, min(1.0, lambda_start * (1.0 - progress))) -def _cpu_max_valid_anchors(loss_mask: torch.Tensor) -> Optional[int]: +def _cpu_valid_anchor_counts( + loss_mask: torch.Tensor, sequence_lengths: Tuple[int, ...] +) -> Optional[Tuple[int, ...]]: + """Count packed document anchors on the host before feature H2D copies.""" + if loss_mask.device.type != "cpu": + return None + return tuple( + int(((mask[:-1] > 0.5) & (mask[1:] > 0.5)).sum()) + for mask in loss_mask[0].split(sequence_lengths) + ) + + +def _cpu_max_valid_anchors( + loss_mask: torch.Tensor, sequence_lengths: Optional[Tuple[int, ...]] = None +) -> Optional[int]: """Count the widest valid anchor row without synchronizing the GPU. Online/offline loaders hand strategies CPU integer features; Mooncake @@ -82,6 +96,10 @@ def _cpu_max_valid_anchors(loss_mask: torch.Tensor) -> Optional[int]: """ if loss_mask.device.type != "cpu": return None + if sequence_lengths is not None: + # Preserve the per-document anchor budget, not the sum across a packed row. + counts = _cpu_valid_anchor_counts(loss_mask, sequence_lengths) + return max(counts, default=0) num_candidates = max(loss_mask.shape[1] - 1, 0) valid = (loss_mask[:, :num_candidates] > 0.5) & ( loss_mask[:, 1 : num_candidates + 1] > 0.5 @@ -128,9 +146,9 @@ def _prepare_eagle_target( ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Normalize EAGLE-family teacher features for a training forward. - Online capture already shifts logits and input IDs. Offline capture stores - the target model's final hidden state, so the frozen target head owns the - equivalent shift and projection to full-vocabulary logits. + Raw hidden-state features from offline readers or online server capture + require the frozen target head's shift and full-vocabulary projection. + Preprocessed logits and their aligned input IDs are used as delivered. """ if target_repr == "hidden_state": if target_head is None: @@ -278,7 +296,46 @@ def forward_loss( target_repr = batch.metadata.get("target_repr") compact_kwargs: Dict[str, Any] = {} - if self.compact_teacher: + packed_kwargs: Dict[str, Any] = {} + sequence_lengths = t.get("sequence_lengths") + if sequence_lengths is not None: + if self.compact_teacher or self.trim_loss_positions: + raise ValueError( + "sequence packing does not support compact_teacher or trim_loss_positions" + ) + if target_repr != "hidden_state" or self.target_head is None: + raise ValueError( + "sequence packing requires unshifted hidden_state targets and a target head" + ) + from specforge.modeling.packed_sequence import PackedSequenceLayout + + layout = PackedSequenceLayout.from_lengths( + sequence_lengths, t["input_ids"].shape[1], t["input_ids"].device + ) + # TargetHead's ordinary global shift would import the next document's + # first token into this document's last row. Shift all three fields + # within document boundaries before projecting the frozen teacher. + input_ids = layout.shift_left(t["input_ids"]).to(device, non_blocking=True) + target_hidden = layout.shift_left(t["target"]) + target = self.target_head(target_hidden.to(device, non_blocking=True)) + loss_mask = layout.shift_left(t["loss_mask"])[..., None].to( + device, non_blocking=True + ) + loss_denominator = t.get("loss_denominator") + if loss_denominator is not None and ( + loss_denominator.device.type != "cpu" or loss_denominator.numel() != 1 + ): + raise ValueError("loss_denominator must be a scalar CPU tensor") + # FSDP moves tensor kwargs onto its compute device. Keep the small + # control metadata as Python values so it stays host-side without a + # GPU synchronization when the wrapped model builds its layout. + packed_kwargs = { + "sequence_lengths": tuple(sequence_lengths.tolist()), + "loss_denominator": ( + int(loss_denominator) if loss_denominator is not None else None + ), + } + elif self.compact_teacher: if target_repr != "hidden_state": raise ValueError( "compact teacher is offline-only and requires " @@ -326,6 +383,7 @@ def forward_loss( else None ), trim_loss_positions=self.trim_loss_positions, + **packed_kwargs, **compact_kwargs, ) weights = [self.ploss_decay**i for i in range(len(plosses))] @@ -516,7 +574,33 @@ def forward_loss( t = batch.tensors device = self._device() selector_loss_alpha = self._selector_loss_alpha(ctx) - max_valid_anchors = _cpu_max_valid_anchors(t["loss_mask"]) + sequence_lengths = t.get("sequence_lengths") + if sequence_lengths is not None: + if ( + sequence_lengths.device.type != "cpu" + or sequence_lengths.dtype != torch.long + or sequence_lengths.ndim != 1 + or not sequence_lengths.numel() + or bool((sequence_lengths <= 0).any()) + or t["input_ids"].shape[0] != 1 + or int(sequence_lengths.sum()) != t["input_ids"].shape[1] + ): + raise ValueError( + "sequence_lengths must be positive CPU int64 document lengths for one packed row" + ) + # FSDP moves tensor kwargs to the GPU; small control metadata must + # remain host-side so packing does not introduce a device sync. + sequence_lengths = tuple(sequence_lengths.tolist()) + valid_anchor_counts = ( + _cpu_valid_anchor_counts(t["loss_mask"], sequence_lengths) + if sequence_lengths is not None + else None + ) + max_valid_anchors = ( + max(valid_anchor_counts, default=0) + if valid_anchor_counts is not None + else _cpu_max_valid_anchors(t["loss_mask"]) + ) collect_detailed_metrics = ( ctx.collect_detailed_metrics if ctx is not None else True ) @@ -527,6 +611,10 @@ def forward_loss( "max_valid_anchors": max_valid_anchors, "selector_loss_alpha": selector_loss_alpha, } + if sequence_lengths is not None: + model_inputs["sequence_lengths"] = sequence_lengths + if valid_anchor_counts is not None: + model_inputs["valid_anchor_counts"] = valid_anchor_counts if ctx is not None: model_inputs["collect_detailed_metrics"] = collect_detailed_metrics target_last_hidden_states = t.get("target_last_hidden_states") diff --git a/tests/test_config/test_sequence_packing.py b/tests/test_config/test_sequence_packing.py new file mode 100644 index 000000000..99c6a0add --- /dev/null +++ b/tests/test_config/test_sequence_packing.py @@ -0,0 +1,148 @@ +"""Reject unsupported packing combinations before building any model.""" + +import unittest +from dataclasses import replace + +from pydantic import ValidationError + +from specforge.algorithms.builtin import builtin_algorithm_registry +from specforge.algorithms.eagle3.data import DataCollatorWithPacking +from specforge.application import resolve_run +from specforge.config import Config +from specforge.config.schema import TrainingConfig + + +def _offline_config(**training): + return Config.model_validate( + { + "model": { + "target_model_path": "target", + "draft_model_config": "draft.json", + "vocab_mapping_path": "mapping.pt", + }, + "data": {"hidden_states_path": "features"}, + "training": training, + } + ) + + +class SequencePackingConfigTest(unittest.TestCase): + def test_default_retains_padded_execution(self): + resolved = resolve_run(_offline_config()) + self.assertFalse(resolved.config.training.sequence_packing) + + def test_offline_eagle3_packing_is_available(self): + resolved = resolve_run(_offline_config(sequence_packing=True, batch_size=4)) + self.assertTrue(resolved.config.training.sequence_packing) + provider = resolved.algorithm.providers.offline_for("text") + self.assertIsInstance(provider.build_packed_collator(), DataCollatorWithPacking) + + def test_other_algorithms_reject_packing(self): + for strategy in ("domino", "dspark"): + with ( + self.subTest(strategy=strategy), + self.assertRaisesRegex( + ValueError, "does not support training.sequence_packing" + ), + ): + resolve_run(_offline_config(strategy=strategy, sequence_packing=True)) + + def test_rejects_non_flex_attention(self): + for attention_backend in ("eager", "sdpa", "fa", "usp"): + with ( + self.subTest(backend=attention_backend), + self.assertRaisesRegex( + ValidationError, "sequence_packing requires flex_attention" + ), + ): + TrainingConfig( + sequence_packing=True, attention_backend=attention_backend + ) + + def test_rejects_unimplemented_objective_combinations(self): + for extra in ( + {"compact_teacher": True}, + {"trim_loss_positions": True}, + ): + with ( + self.subTest(extra=extra), + self.assertRaisesRegex( + ValidationError, "sequence_packing currently requires" + ), + ): + TrainingConfig(sequence_packing=True, **extra) + + def test_lk_packing_is_algorithm_specific(self): + for lk_loss_type in ("lambda", "alpha", "tv"): + with self.subTest(lk_loss_type=lk_loss_type): + with self.assertRaisesRegex(ValueError, "LK|lk_loss"): + resolve_run( + _offline_config( + sequence_packing=True, lk_loss_type=lk_loss_type + ) + ) + resolve_run( + _offline_config( + strategy="dflash", + sequence_packing=True, + lk_loss_type=lk_loss_type, + ) + ) + + def test_supports_dflash_offline_and_dflash2_architecture(self): + resolved = resolve_run( + _offline_config(strategy="dflash", sequence_packing=True) + ) + self.assertTrue( + callable( + resolved.algorithm.providers.offline_for("text").build_packed_collator + ) + ) + self.assertIn( + "DFlash2DraftModel", resolved.algorithm.spec.draft.compatible_architectures + ) + + def test_supports_online_text_features(self): + payload = _offline_config(sequence_packing=True).model_dump() + payload["data"] = {"train_data_path": "train.jsonl"} + payload["training"]["max_steps"] = 1 + payload["training"]["role"] = "auto" + payload["deployment"] = { + "mode": "disaggregated", + "disaggregated": { + "control_dir": "outputs/packing-test/control", + "backend": "mooncake", + "server_urls": ["http://127.0.0.1:30000"], + "mooncake_metadata_server": "http://127.0.0.1:35880/metadata", + "mooncake_master_server_addr": "127.0.0.1:35551", + }, + } + for strategy in ("eagle3", "dflash"): + payload["training"]["strategy"] = strategy + with self.subTest(strategy=strategy): + resolved = resolve_run(Config.model_validate(payload)) + provider = resolved.algorithm.providers.server_streaming_for("text") + self.assertTrue(callable(provider.build_packed_collator)) + + def test_provider_rejects_noncallable_packing_factory(self): + provider = ( + builtin_algorithm_registry().resolve("eagle3").providers.offline_for("text") + ) + with self.assertRaisesRegex( + TypeError, "build_packed_collator must be callable" + ): + replace(provider, build_packed_collator=True) + + streaming = ( + builtin_algorithm_registry() + .resolve("eagle3") + .providers.server_streaming_for("text") + ) + with self.assertRaisesRegex( + TypeError, "build_packed_collator must be callable" + ): + replace(streaming, build_packed_collator=True) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_data/test_sequence_packing.py b/tests/test_data/test_sequence_packing.py new file mode 100644 index 000000000..b47c976a4 --- /dev/null +++ b/tests/test_data/test_sequence_packing.py @@ -0,0 +1,239 @@ +"""Packing preserves logical samples while removing only batch padding (CPU).""" + +import tempfile +import unittest +from pathlib import Path + +import torch + +from specforge.algorithms.eagle3.data import ( + DataCollatorWithPacking, + build_offline_normalizer, + build_packed_collator, +) +from specforge.runtime.data_plane.feature_dataloader import FeatureDataLoader +from specforge.runtime.data_plane.feature_store import LocalFeatureStore +from specforge.runtime.data_plane.offline_reader import OfflineManifestReader +from specforge.runtime.data_plane.sample_ref_queue import SampleRefQueue + + +def _feature(length, offset=0): + return { + "input_ids": (torch.arange(length) + offset).unsqueeze(0), + "attention_mask": torch.ones(1, length, dtype=torch.long), + "loss_mask": (torch.arange(length) % 2).unsqueeze(0), + "hidden_state": torch.arange(length * 6).reshape(1, length, 6) + offset, + "target": torch.arange(length * 2).reshape(1, length, 2) + offset, + } + + +class SequencePackingCollatorTest(unittest.TestCase): + def test_concatenates_features_with_document_positions_and_padded_denominator(self): + features = [_feature(2, 10), _feature(5, 20), _feature(1, 30)] + originals = [{key: value.clone() for key, value in f.items()} for f in features] + + batch = build_packed_collator()(features) + + for key in features[0]: + with self.subTest(key=key): + torch.testing.assert_close( + batch[key], torch.cat([f[key] for f in originals], dim=1) + ) + self.assertEqual(batch["input_ids"].shape, (1, 8)) + self.assertEqual(batch["position_ids"].tolist(), [[0, 1, 0, 1, 2, 3, 4, 0]]) + self.assertEqual(batch["sequence_lengths"].tolist(), [2, 5, 1]) + self.assertEqual(batch["loss_denominator"].item(), 3 * 5) + self.assertEqual(batch["sequence_lengths"].dtype, torch.long) + self.assertEqual(batch["sequence_lengths"].device.type, "cpu") + self.assertEqual(batch["loss_denominator"].device.type, "cpu") + for feature, original in zip(features, originals): + self.assertEqual(feature.keys(), original.keys()) + for key in original: + torch.testing.assert_close(feature[key], original[key]) + + def test_accepts_standard_positions_and_single_sample(self): + feature = _feature(3) + feature["position_ids"] = torch.arange(3).unsqueeze(0) + + batch = DataCollatorWithPacking()([feature]) + + self.assertEqual(batch["position_ids"].tolist(), [[0, 1, 2]]) + self.assertEqual(batch["sequence_lengths"].tolist(), [3]) + self.assertEqual(batch["loss_denominator"].item(), 3) + + def test_rejects_empty_batch_and_empty_or_batched_sample(self): + with self.assertRaisesRegex(ValueError, "empty feature batch"): + DataCollatorWithPacking()([]) + for ids in (torch.empty(1, 0), torch.zeros(2, 3), torch.zeros(3)): + feature = _feature(3) + feature["input_ids"] = ids + with ( + self.subTest(shape=ids.shape), + self.assertRaisesRegex(ValueError, "nonempty"), + ): + DataCollatorWithPacking()([feature]) + + def test_rejects_missing_or_misaligned_features(self): + for key in _feature(3): + feature = _feature(3) + del feature[key] + with self.subTest(missing=key), self.assertRaisesRegex(KeyError, key): + DataCollatorWithPacking()([feature]) + for key in ("loss_mask", "attention_mask", "hidden_state", "target"): + for wrong in (torch.zeros(1, 2), torch.zeros(3), torch.zeros(2, 3, 6)): + feature = _feature(3) + feature[key] = wrong + with ( + self.subTest(key=key, shape=wrong.shape), + self.assertRaisesRegex(ValueError, key), + ): + DataCollatorWithPacking()([feature]) + + def test_rejects_padding_and_nonstandard_positions(self): + feature = _feature(3) + feature["attention_mask"][0, -1] = 0 + with self.assertRaisesRegex(ValueError, "unpadded"): + DataCollatorWithPacking()([feature]) + for positions in ( + torch.tensor([[1, 2, 3]]), + torch.tensor([[0, 0, 1]]), + torch.arange(3), + torch.arange(3).reshape(1, 1, 3), + ): + feature = _feature(3) + feature["position_ids"] = positions + with ( + self.subTest(shape=positions.shape), + self.assertRaisesRegex(ValueError, "standard text position_ids"), + ): + DataCollatorWithPacking()([feature]) + + def test_offline_loader_preserves_sample_ids_order_and_partial_batch(self): + with tempfile.TemporaryDirectory() as directory: + for index, length in enumerate((2, 5, 3)): + torch.save( + { + "input_ids": torch.arange(length) + index * 10, + "loss_mask": torch.ones(length, dtype=torch.long), + "hidden_state": torch.full((1, length, 2), float(index)), + "aux_hidden_state": torch.full((1, length, 6), float(index)), + }, + Path(directory) / f"{index:03d}.ckpt", + ) + refs = OfflineManifestReader(directory, run_id="packing-test").read() + loader = FeatureDataLoader( + LocalFeatureStore("packing-test"), + refs=refs, + batch_size=2, + collate_fn=build_packed_collator(), + per_sample_transform=build_offline_normalizer(4), + drop_last=False, + ) + batches = list(loader) + repeated = list(loader) + + expected_ids = [[r.sample_id for r in refs[:2]], [refs[2].sample_id]] + self.assertEqual([b.sample_ids for b in batches], expected_ids) + self.assertEqual([b.sample_ids for b in repeated], expected_ids) + self.assertEqual(batches[0].tensors["sequence_lengths"].tolist(), [2, 4]) + self.assertEqual( + batches[0].tensors["input_ids"].tolist(), [[0, 1, 10, 11, 12, 13]] + ) + self.assertEqual(batches[0].tensors["loss_mask"].tolist(), [[1, 0, 1, 1, 1, 0]]) + self.assertEqual(batches[0].tensors["loss_denominator"].item(), 8) + self.assertEqual(batches[1].tensors["sequence_lengths"].tolist(), [3]) + self.assertEqual(batches[1].tensors["loss_denominator"].item(), 3) + + +class StreamingPackingDataTest(unittest.TestCase): + @staticmethod + def _dflash(length, offset=0, teacher=True): + sample = _feature(length, offset) + result = { + "input_ids": sample["input_ids"], + "loss_mask": sample["loss_mask"], + "hidden_states": sample["hidden_state"], + } + if teacher: + result["target_last_hidden_states"] = sample["target"] + return result + + def test_dflash_preserves_optional_teacher_features_without_loss_denominator(self): + from specforge.algorithms.common.hidden_states_data import build_packed_collator + + for teacher in (False, True): + features = [self._dflash(3, 10, teacher), self._dflash(5, 20, teacher)] + originals = [{k: v.clone() for k, v in f.items()} for f in features] + batch = build_packed_collator()(features) + self.assertEqual(batch["sequence_lengths"].tolist(), [3, 5]) + self.assertNotIn("loss_denominator", batch) + self.assertEqual("target_last_hidden_states" in batch, teacher) + for key in originals[0]: + torch.testing.assert_close( + batch[key], torch.cat([f[key] for f in originals], dim=1) + ) + for actual, original in zip(features, originals): + torch.testing.assert_close(actual[key], original[key]) + + def test_dflash_rejects_inconsistent_teacher_features_and_bad_shapes(self): + from specforge.algorithms.common.hidden_states_data import build_packed_collator + + collate = build_packed_collator() + with self.assertRaises(KeyError): + collate([self._dflash(3, teacher=True), self._dflash(3, teacher=False)]) + for key in self._dflash(3): + sample = self._dflash(3) + sample[key] = sample[key][:, :2] + with self.subTest(key=key), self.assertRaises(ValueError): + collate([sample]) + + def test_online_queue_keeps_logical_batch_size_identity_and_ack_count(self): + from specforge.algorithms.builtin import builtin_algorithm_registry + + for name in ("eagle3", "dflash"): + algorithm = builtin_algorithm_registry().resolve(name) + store = LocalFeatureStore(f"packing-queue-{name}") + refs = [] + for index, length in enumerate((3, 7, 4, 6)): + sample = ( + _feature(length, index * 10) + if name == "eagle3" + else self._dflash(length, index * 10) + ) + refs.append( + store.put( + sample, + sample_id=f"sample-{index}", + metadata={ + "run_id": "queue-test", + "strategy": name, + "target_repr": "hidden_state", + }, + ) + ) + queue = SampleRefQueue() + queue.put(refs) + loader = FeatureDataLoader( + store, + queue, + batch_size=2, + strategy=name, + collate_fn=algorithm.providers.server_streaming_for( + "text" + ).build_packed_collator(), + ) + batches = list(loader) + self.assertEqual( + [batch.sample_ids for batch in batches], + [["sample-0", "sample-1"], ["sample-2", "sample-3"]], + ) + self.assertEqual( + [batch.tensors["sequence_lengths"].tolist() for batch in batches], + [[3, 7], [4, 6]], + ) + self.assertEqual(queue.in_flight(), 0) + self.assertEqual(queue.depth(), 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_modeling/test_dflash_sequence_packing.py b/tests/test_modeling/test_dflash_sequence_packing.py new file mode 100644 index 000000000..86e245da6 --- /dev/null +++ b/tests/test_modeling/test_dflash_sequence_packing.py @@ -0,0 +1,360 @@ +"""DFlash context packing preserves per-document sampling and objectives.""" + +import copy +import unittest + +import torch +from torch import nn + +from specforge.algorithms.common.dflash_family_model import ( + OnlineDFlashModel, + create_dflash_block_mask, + create_dflash_sdpa_mask, +) +from specforge.modeling.packed_dflash import PackedDFlashLayout + + +class PackedDFlashMaskTest(unittest.TestCase): + def test_sparse_tile_metadata_covers_exact_mask_and_full_tiles_are_exact(self): + lengths = (131, 67, 3, 312) + starts = torch.tensor([0, 131, 198, 201]) + local = torch.arange(9)[None, :].expand(4, -1) * 13 + local = torch.minimum(local, torch.tensor(lengths)[:, None] - 1) + anchors = (local + starts[:, None]).reshape(1, -1) + context_starts = starts[:, None].expand(-1, 9).reshape(1, -1) + keep = torch.ones_like(anchors, dtype=torch.bool) + keep[:, 10:13] = False + for proposal in (3, 16, 33): + for window in (None, 7): + for tile_shape in ((128, 128), (256, 128)): + with self.subTest( + proposal=proposal, window=window, tile_shape=tile_shape + ): + arguments = dict( + anchor_positions=anchors, + block_keep_mask=keep, + S=sum(lengths), + block_size=proposal, + device="cpu", + sliding_window=window, + context_start_positions=context_starts, + ) + dense = create_dflash_sdpa_mask(**arguments)[0, 0] + sparse = create_dflash_block_mask( + **arguments, flex_block_size=tile_shape + ) + rows, columns = sparse.kv_indices.shape[-2:] + tiles = torch.zeros(rows, columns, dtype=torch.bool) + full = torch.zeros_like(tiles) + for row in range(rows): + tiles[ + row, + sparse.kv_indices[ + 0, 0, row, : sparse.kv_num_blocks[0, 0, row] + ], + ] = True + full[ + row, + sparse.full_kv_indices[ + 0, 0, row, : sparse.full_kv_num_blocks[0, 0, row] + ], + ] = True + padded = torch.zeros( + rows * tile_shape[0], + columns * tile_shape[1], + dtype=torch.bool, + ) + padded[: dense.shape[0], : dense.shape[1]] = dense + token_tiles = padded.reshape( + rows, tile_shape[0], columns, tile_shape[1] + ).permute(0, 2, 1, 3) + exact_any = token_tiles.any(-1).any(-1) + exact_full = token_tiles.all(-1).all(-1) + self.assertFalse(bool((exact_any & ~(tiles | full)).any())) + self.assertFalse(bool((full & ~exact_full).any())) + + def test_sampling_mask_restores_documents_and_excludes_padding(self): + layout = PackedDFlashLayout.from_lengths((3, 1, 2), 6, "cpu") + packed_mask = torch.tensor([[1, 0, 1, 1, 1, 1]]) + torch.testing.assert_close( + layout.padded_loss_mask(packed_mask), + torch.tensor([[1, 0, 1], [1, 0, 0], [1, 1, 0]]), + ) + + def test_full_and_sliding_masks_never_read_previous_documents(self): + anchors = torch.tensor([[1, 4, 6]]) + starts = torch.tensor([[0, 3, 3]]) + keep = torch.tensor([[True, False, True]]) + for window in (None, 3): + args = dict( + anchor_positions=anchors, + block_keep_mask=keep, + S=8, + block_size=4, + device="cpu", + sliding_window=window, + context_start_positions=starts, + ) + dense = create_dflash_sdpa_mask(**args)[0, 0] + sparse = create_dflash_block_mask(**args) + q = torch.arange(12)[:, None] + k = torch.arange(20)[None, :] + actual = sparse.mask_mod(torch.tensor(0), torch.tensor(0), q, k) + torch.testing.assert_close(actual, dense) + self.assertFalse(bool(dense[8:, :3].any())) + self.assertFalse(bool(dense[4:8].any())) + self.assertFalse(bool(dense[8:, 8:16].any())) + + +@unittest.skipUnless( + torch.cuda.is_available(), "production Flex Attention requires CUDA" +) +class PackedDFlashProductionTest(unittest.TestCase): + def _fixtures(self, *, dflash2, sliding, dtype, loss_type, lk_loss_type=None): + from transformers import Qwen3Config + + from specforge.modeling.draft.dflash import DFlashDraftModel + from specforge.modeling.draft.dflash2 import DFlash2DraftModel + + torch.manual_seed(912) + method = { + "block_size": 4, + "mask_token_id": 63, + "target_layer_ids": [1, 2], + } + if dflash2: + method.update( + conv_group_size=4, conv_kernel_size=3, selector_rank=8, selector_top_k=8 + ) + config = Qwen3Config( + hidden_size=64, + intermediate_size=128, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=16, + num_hidden_layers=2, + num_target_layers=4, + max_position_embeddings=256, + vocab_size=64, + layer_types=( + ["sliding_attention", "full_attention"] + if sliding + else ["full_attention"] * 2 + ), + sliding_window=16 if sliding else None, + use_sliding_window=sliding, + dflash_config=method, + ) + config._attn_implementation = "flex_attention" + draft = ( + (DFlash2DraftModel if dflash2 else DFlashDraftModel)(config) + .cuda() + .to(dtype) + ) + if dflash2: + with torch.no_grad(): + for layer in draft.layers: + layer.attention_conv.kernel_projection.weight.normal_(std=0.01) + layer.mlp_conv.kernel_projection.weight.normal_(std=0.01) + head = nn.Linear(64, 64, bias=False).cuda().to(dtype).requires_grad_(False) + embedding = nn.Embedding(64, 64).cuda().to(dtype).requires_grad_(False) + model = OnlineDFlashModel( + draft, + head, + embedding, + mask_token_id=63, + block_size=4, + num_anchors=8, + attention_backend="flex_attention", + loss_type=loss_type, + lk_loss_type=lk_loss_type, + objective_chunk_blocks=3, + loss_decay_gamma=3.0, + selector_loss_alpha=0.7, + metric_top_k=8, + ) + generator = torch.Generator().manual_seed(555) + features = [] + for length in (131, 67, 3, 1): + loss_mask = torch.ones(1, length) + if length > 10: + loss_mask[:, :4] = 0 + loss_mask[:, 9:12] = 0 + features.append( + { + "input_ids": torch.randint(0, 63, (1, length), generator=generator), + "loss_mask": loss_mask, + "hidden_states": torch.randn( + 1, length, 128, generator=generator + ).to(dtype), + "target_last_hidden_states": torch.randn( + 1, length, 64, generator=generator + ).to(dtype), + } + ) + return model, features + + def _run(self, model, features, *, packed, compact=True): + lengths = [f["input_ids"].shape[1] for f in features] + if packed: + batch = { + key: torch.cat([f[key] for f in features], dim=1).cuda() + for key in features[0] + } + batch["sequence_lengths"] = tuple(lengths) + if compact: + batch["valid_anchor_counts"] = tuple( + int( + ( + (f["loss_mask"][:, :-1] > 0.5) + & (f["loss_mask"][:, 1:] > 0.5) + ).sum() + ) + for f in features + ) + else: + batch = {} + for key in features[0]: + rows = [] + for f, length in zip(features, lengths): + value = f[key] + rows.append( + torch.cat( + ( + value, + value.new_zeros( + (1, max(lengths) - length, *value.shape[2:]) + ), + ), + dim=1, + ) + ) + batch[key] = torch.cat(rows).cuda() + sampled = [] + original_sampler = model._sample_anchor_positions + + def record(*args, **kwargs): + result = original_sampler(*args, **kwargs) + sampled.append(tuple(t.detach().clone() for t in result)) + return result + + model._sample_anchor_positions = record + torch.manual_seed(999) + try: + loss, accuracy, metrics = model(**batch) + loss.backward() + finally: + model._sample_anchor_positions = original_sampler + gradients = { + name: p.grad.detach().clone() + for name, p in model.named_parameters() + if p.grad is not None + } + return (loss.detach(), accuracy.detach(), metrics), gradients, sampled[0] + + def _compare(self, left, right, tol, path=""): + if isinstance(left, dict): + self.assertEqual(left.keys(), right.keys(), path) + for name in left: + self._compare(left[name], right[name], tol, path + "/" + name) + elif isinstance(left, (tuple, list)): + for index, (a, b) in enumerate(zip(left, right)): + self._compare(a, b, tol, path + f"/{index}") + elif isinstance(left, torch.Tensor): + torch.testing.assert_close(left, right, msg=path, **tol) + else: + self.assertEqual(left, right, path) + + def test_production_anchors_losses_metrics_and_every_gradient_match(self): + cases = ( + (False, False, torch.float32, "dflash", None), + (False, True, torch.float32, "dpace", None), + (True, False, torch.float32, "dflash", None), + (True, True, torch.float32, "dpace", None), + (True, False, torch.bfloat16, "dflash", None), + (True, True, torch.bfloat16, "dpace", None), + (True, True, torch.float32, "dpace-cumulative-confidence-only", "lambda"), + (True, False, torch.float32, "dpace-continuation-value-only", "tv"), + ) + for dflash2, sliding, dtype, loss_type, lk_loss_type in cases: + with self.subTest( + dflash2=dflash2, + sliding=sliding, + dtype=dtype, + loss_type=loss_type, + lk_loss_type=lk_loss_type, + ): + model, features = self._fixtures( + dflash2=dflash2, + sliding=sliding, + dtype=dtype, + loss_type=loss_type, + lk_loss_type=lk_loss_type, + ) + packed_model = copy.deepcopy(model) + baseline, baseline_gradients, baseline_anchors = self._run( + model, features, packed=False + ) + packed, packed_gradients, packed_anchors = self._run( + packed_model, features, packed=True + ) + for expected, actual in zip(baseline_anchors, packed_anchors): + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + tol = ( + dict(atol=3e-6, rtol=5e-4) + if dtype == torch.float32 + else dict(atol=7e-4, rtol=5e-2) + ) + self._compare(packed, baseline, tol) + self._compare(packed_gradients, baseline_gradients, tol, "gradients") + self.assertTrue(any("q_proj" in name for name in packed_gradients)) + if dflash2: + self.assertTrue( + any("candidate_selector" in name for name in packed_gradients) + ) + self.assertTrue( + any("attention_conv" in name for name in packed_gradients) + ) + + def test_dflash2_blocks_cannot_read_another_document(self): + model, features = self._fixtures( + dflash2=True, sliding=True, dtype=torch.float32, loss_type="dpace" + ) + changed = copy.deepcopy(features) + changed[0]["input_ids"].fill_(33) + changed[0]["hidden_states"].mul_(100) + captured = [] + handle = model.draft_model.register_forward_hook( + lambda _module, _args, output: captured.append(output.detach().clone()) + ) + try: + self._run(model, features, packed=True) + model.zero_grad(set_to_none=True) + self._run(model, changed, packed=True) + finally: + handle.remove() + self.assertEqual(len(captured), 2) + first_document_blocks = model.num_anchors * model.block_size + torch.testing.assert_close( + captured[0][:, first_document_blocks:], + captured[1][:, first_document_blocks:], + atol=0, + rtol=0, + ) + + def test_missing_host_counts_retains_equivalent_uncompacted_fallback(self): + model, features = self._fixtures( + dflash2=True, sliding=True, dtype=torch.float32, loss_type="dpace" + ) + uncompacted_model = copy.deepcopy(model) + compact, compact_grads, _ = self._run(model, features, packed=True) + uncompacted, uncompacted_grads, _ = self._run( + uncompacted_model, features, packed=True, compact=False + ) + tolerance = dict(atol=3e-6, rtol=5e-4) + self._compare(compact, uncompacted, tolerance) + self._compare(compact_grads, uncompacted_grads, tolerance, "gradients") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_runtime/test_eagle3_sequence_packing.py b/tests/test_runtime/test_eagle3_sequence_packing.py new file mode 100644 index 000000000..6440b5494 --- /dev/null +++ b/tests/test_runtime/test_eagle3_sequence_packing.py @@ -0,0 +1,235 @@ +"""Packed EAGLE3 must retain padded-batch supervision, gradients, and isolation.""" + +import copy +import tempfile +import unittest + +import torch + +from specforge.modeling.packed_sequence import ( + PackedSequenceLayout, + generate_packed_eagle3_mask, +) + + +class PackedSequenceLayoutTest(unittest.TestCase): + def test_repeated_shifts_never_import_the_next_document(self): + layout = PackedSequenceLayout.from_lengths(torch.tensor([3, 2]), 5, "cpu") + values = torch.tensor([[1, 2, 3, 4, 5]]) + for expected in ([2, 3, 0, 5, 0], [3, 0, 0, 0, 0], [0, 0, 0, 0, 0]): + values = layout.shift_left(values) + self.assertEqual(values.tolist(), [expected]) + + def test_mask_matches_independent_documents_at_every_depth(self): + lengths = [5, 2, 4] + layout = PackedSequenceLayout.from_lengths( + torch.tensor(lengths), sum(lengths), "cpu" + ) + size = sum(lengths) + for depth in range(4): + mask = generate_packed_eagle3_mask(layout, size, depth) + for q in range(size + 2): + for k in range(size * (depth + 1)): + row = k % size + valid = q < size + if valid: + valid = ( + layout.document_ids[q] == layout.document_ids[row] + and layout.positions[q] < layout.document_lengths[q] - depth + and layout.positions[row] + < layout.document_lengths[row] - depth + ) + expected = bool( + valid and ((k < size and q >= k) or (k >= size and row == q)) + ) + actual = bool(mask(0, 0, torch.tensor(q), torch.tensor(k))) + self.assertEqual(actual, expected, (depth, q, k)) + + +@unittest.skipUnless( + torch.cuda.is_available(), "production Flex Attention and loss require CUDA" +) +class Eagle3PackedProductionTest(unittest.TestCase): + def test_cpu_layout_can_shift_device_resident_hidden_features(self): + layout = PackedSequenceLayout.from_lengths(torch.tensor([3, 2]), 5, "cpu") + hidden = torch.arange(10, device="cuda").view(1, 5, 2) + torch.testing.assert_close( + layout.shift_left(hidden), + torch.tensor([[[2, 3], [4, 5], [0, 0], [8, 9], [0, 0]]], device="cuda"), + ) + + def _fixtures(self, dtype=torch.float32, rope_scaling=None): + from transformers import LlamaConfig + + from specforge.algorithms.eagle3.model import OnlineEagle3Model + from specforge.modeling.draft.llama3_eagle import LlamaForCausalLMEagle3 + from specforge.modeling.target.target_head import TargetHead + + torch.manual_seed(123) + config = LlamaConfig( + vocab_size=64, + draft_vocab_size=64, + hidden_size=64, + intermediate_size=128, + num_hidden_layers=1, + num_attention_heads=4, + num_key_value_heads=2, + max_position_embeddings=160, + rope_scaling=rope_scaling, + pad_token_id=0, + ) + draft = ( + LlamaForCausalLMEagle3(config, attention_backend="flex_attention") + .cuda() + .to(dtype) + ) + model = OnlineEagle3Model(draft, length=4, attention_backend="flex_attention") + + # Use the real frozen head methods, without downloading any checkpoint. + head = TargetHead.__new__(TargetHead) + torch.nn.Module.__init__(head) + head.fc = torch.nn.Linear(64, 64, bias=False).cuda().to(dtype) + head.freeze_weights() + features = [] + for length in (131, 67, 3): + mask = torch.ones(1, length, dtype=torch.long) + mask[:, -1] = 0 + # A prompt mask and an internal supervision gap exercise mask shifts. + if length > 10: + mask[:, :4] = 0 + mask[:, 9:12] = 0 + features.append( + { + "input_ids": torch.randint(1, 64, (1, length)), + "attention_mask": torch.ones(1, length, dtype=torch.long), + "loss_mask": mask, + "hidden_state": torch.randn(1, length, 192).to(dtype), + "target": torch.randn(1, length, 64).to(dtype), + } + ) + return model, head, features + + def _run(self, model, head, features, packed): + from specforge.algorithms.eagle3.data import DataCollatorWithPacking + from specforge.data.utils import DataCollatorWithPadding + from specforge.runtime.contracts import TrainBatch + from specforge.training.strategies.base import Eagle3TrainStrategy + + collator = DataCollatorWithPacking() if packed else DataCollatorWithPadding() + batch = TrainBatch( + sample_ids=[str(i) for i in range(len(features))], + strategy="eagle3", + tensors=collator(features), + metadata={"target_repr": "hidden_state"}, + ) + output = Eagle3TrainStrategy(model, target_head=head).forward_loss(batch) + output.loss.backward() + gradients = { + name: parameter.grad.clone() + for name, parameter in model.named_parameters() + if parameter.grad is not None + } + return output, gradients + + def test_production_loss_metrics_and_all_gradients_match_padding(self): + cases = ( + (torch.float32, None), + (torch.bfloat16, None), + (torch.float32, {"rope_type": "dynamic", "factor": 2.0}), + ) + for dtype, rope_scaling in cases: + with self.subTest(dtype=dtype, rope_scaling=rope_scaling): + model, head, features = self._fixtures(dtype, rope_scaling) + packed_model = copy.deepcopy(model) + padded, padded_grads = self._run(model, head, features, packed=False) + packed, packed_grads = self._run( + packed_model, head, features, packed=True + ) + if rope_scaling is not None: + # Initialization caches max_position_embeddings + 20 (180). + # The packed total is 201; packing must not extend that cache + # and change dynamic NTK scaling relative to the padded batch. + self.assertEqual( + packed_model.draft_model.midlayer.self_attn.rotary_emb.max_seq_len_cached, + 180, + ) + tol = ( + {"atol": 2e-6, "rtol": 2e-4} + if dtype == torch.float32 + else {"atol": 3e-4, "rtol": 3e-2} + ) + torch.testing.assert_close(packed.loss, padded.loss, **tol) + for key in padded.metrics: + for expected, actual in zip( + padded.metrics[key], packed.metrics[key] + ): + torch.testing.assert_close(actual, expected, **tol) + self.assertEqual(padded_grads.keys(), packed_grads.keys()) + for name in padded_grads: + torch.testing.assert_close( + packed_grads[name], padded_grads[name], msg=name, **tol + ) + + def test_changing_previous_document_cannot_change_later_document_logits(self): + model, head, features = self._fixtures() + changed = copy.deepcopy(features) + changed[0]["input_ids"].fill_(33) + changed[0]["hidden_state"].mul_(100) + changed[0]["target"].neg_() + captured = [] + hook = model.draft_model.lm_head.register_forward_hook( + lambda _module, _inputs, output: captured.append(output.detach().clone()) + ) + try: + self._run(model, head, features, packed=True) + first = captured[:] + captured.clear() + model.zero_grad(set_to_none=True) + self._run(model, head, changed, packed=True) + finally: + hook.remove() + self.assertEqual(len(first), 4) + self.assertEqual(len(captured), 4) + start = features[0]["input_ids"].shape[1] + for before, after in zip(first, captured): + torch.testing.assert_close( + before[:, start:], after[:, start:], atol=0, rtol=0 + ) + + def test_fully_unsupervised_batch_has_zero_loss_and_finite_zero_gradients(self): + model, head, features = self._fixtures() + for feature in features: + feature["loss_mask"].zero_() + result, gradients = self._run(model, head, features, packed=True) + self.assertEqual(float(result.loss.detach()), 0.0) + for name, gradient in gradients.items(): + self.assertTrue(bool(torch.isfinite(gradient).all()), name) + self.assertEqual(float(gradient.abs().max()), 0.0, name) + + def test_fsdp_forward_preserves_host_packing_metadata(self): + import torch.distributed as dist + from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + + if dist.is_initialized(): + self.skipTest("requires its own single-rank process group") + model, head, features = self._fixtures() + with tempfile.TemporaryDirectory() as directory: + dist.init_process_group( + "nccl", + init_method=f"file://{directory}/process_group", + rank=0, + world_size=1, + ) + try: + wrapped = FSDP(model, use_orig_params=True) + result, gradients = self._run(wrapped, head, features, packed=True) + self.assertTrue(bool(torch.isfinite(result.loss))) + self.assertTrue(gradients) + for name, gradient in gradients.items(): + self.assertTrue(bool(torch.isfinite(gradient).all()), name) + finally: + dist.destroy_process_group() + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_runtime/test_online_sequence_packing_gate.py b/tests/test_runtime/test_online_sequence_packing_gate.py new file mode 100644 index 000000000..d758d66d0 --- /dev/null +++ b/tests/test_runtime/test_online_sequence_packing_gate.py @@ -0,0 +1,298 @@ +"""Opt-in live SGLang -> Mooncake -> packed online consumer lifecycle gate. + +Use SPECFORGE_RUN_SERVER_CAPTURE_TESTS=1 and an isolated patched SGLang on +PYTHONPATH. CUDA_VISIBLE_DEVICES chooses the consumer; PACKING_TARGET_GPU +chooses the server. This gate owns and cleans up only its fixture processes. +""" + +import json +import math +import numbers +import os +import shutil +import sqlite3 +import subprocess +import time +from pathlib import Path + +import torch + +from tests.test_runtime import test_server_capture_gate as gate + + +class TestOnlineSequencePackingGate(gate.TestServerCaptureGate): + # Reuse the live fixture; its separate extraction tests remain in their + # original module instead of being duplicated by this derived test case. + test_eagle3_zero_copy_end_to_end = None + test_dflash_capture_same_server = None + test_dflash_capture_without_teacher_metrics_skips_last_hidden = None + + @classmethod + def setUpClass(cls): + gate.PORT = int(os.environ.get("PACKING_SERVER_PORT", "30992")) + target_gpu = os.environ.get( + "PACKING_TARGET_GPU", os.environ.get("CUDA_VISIBLE_DEVICES", "0") + ) + original_gpu = os.environ.get("CUDA_VISIBLE_DEVICES") + os.environ["CUDA_VISIBLE_DEVICES"] = target_gpu + try: + super().setUpClass() + finally: + if original_gpu is None: + os.environ.pop("CUDA_VISIBLE_DEVICES", None) + else: + os.environ["CUDA_VISIBLE_DEVICES"] = original_gpu + + @classmethod + def _ensure_mooncake_master(cls): + rpc_port = os.environ.get("PACKING_MOONCAKE_RPC_PORT", "50192") + metadata_port = os.environ.get("PACKING_MOONCAKE_METADATA_PORT", "8092") + binary = shutil.which("mooncake_master") + if binary is None: + raise RuntimeError("live packing gate requires mooncake_master") + cls.master = subprocess.Popen( + [ + binary, + "--enable-http-metadata-server=true", + f"--rpc_port={rpc_port}", + f"--http_metadata_server_port={metadata_port}", + "--metrics_port=9092", + ], + stdout=open(os.path.join(cls.workdir, "mooncake_master.log"), "w"), + stderr=subprocess.STDOUT, + start_new_session=True, + ) + time.sleep(3) + if cls.master.poll() is not None: + raise RuntimeError("packing mooncake_master exited during startup") + os.environ["MOONCAKE_MASTER_SERVER_ADDR"] = f"127.0.0.1:{rpc_port}" + os.environ["MOONCAKE_METADATA_SERVER"] = ( + f"http://127.0.0.1:{metadata_port}/metadata" + ) + os.environ["MOONCAKE_LOCAL_HOSTNAME"] = "127.0.0.1" + os.environ["MOONCAKE_PROTOCOL"] = "tcp" + + @classmethod + def _cleanup_processes(cls): + super()._cleanup_processes() + destination = os.environ.get("PACKING_ARTIFACT_DIR") + if destination and cls.workdir: + Path(destination).mkdir(parents=True, exist_ok=True) + for source in Path(cls.workdir).glob("*.log"): + shutil.copy2(source, Path(destination) / source.name) + for source in Path(cls.workdir).glob("*.json"): + if source.name.endswith("-result.json"): + shutil.copy2(source, Path(destination) / source.name) + + def _packing_store(self, run_id): + from specforge.runtime.data_plane.mooncake_store import MooncakeFeatureStore + + return MooncakeFeatureStore( + store_id=run_id, + retain_on_release=True, + setup_kwargs={ + "local_hostname": "127.0.0.1", + "metadata_server": os.environ["MOONCAKE_METADATA_SERVER"], + "global_segment_size": 1 << 28, + "local_buffer_size": 1 << 28, + "protocol": "tcp", + "rdma_devices": "", + "master_server_addr": os.environ["MOONCAKE_MASTER_SERVER_ADDR"], + }, + ) + + def _model(self, name, work): + from safetensors.torch import load_file + from torch import nn + from transformers import Qwen3Config + + from specforge.algorithms.common.dflash_family_model import OnlineDFlashModel + from specforge.modeling.draft.dflash import DFlashDraftModel + from specforge.modeling.draft.dflash2 import DFlash2DraftModel + from specforge.modeling.target.target_head import TargetHead + from tests.test_runtime import _fixtures as fx + + if name == "eagle3": + model, _ = fx.build_eagle3(str(work), ttt=3) + return model, TargetHead.from_pretrained(self.target_dir) + config = Qwen3Config( + architectures=[ + "DFlash2DraftModel" if name == "dflash2" else "DFlashDraftModel" + ], + hidden_size=gate.H, + intermediate_size=128, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=16, + num_hidden_layers=1, + num_target_layers=8, + vocab_size=256, + max_position_embeddings=512, + layer_types=["full_attention"], + dflash_config={ + "block_size": 4, + "mask_token_id": 0, + "target_layer_ids": gate.AUX_LAYER_IDS, + "conv_group_size": 4, + "conv_kernel_size": 2, + "selector_rank": 4, + "selector_top_k": 4, + }, + ) + config._attn_implementation = "flex_attention" + draft = (DFlash2DraftModel if name == "dflash2" else DFlashDraftModel)(config) + weights = load_file(os.path.join(self.target_dir, "model.safetensors")) + head = nn.Linear(gate.H, 256, bias=False) + head.weight.data.copy_(weights["lm_head.weight"]) + head.requires_grad_(False) + embeddings = nn.Embedding.from_pretrained( + weights["model.embed_tokens.weight"], freeze=True + ) + model = OnlineDFlashModel( + draft, + head, + embeddings, + mask_token_id=0, + block_size=4, + num_anchors=4, + attention_backend="flex_attention", + ).to(device="cuda", dtype=torch.bfloat16) + return model, None + + def test_live_packed_online_training_and_durable_acks(self): + from specforge.algorithms.builtin import builtin_algorithm_registry + from specforge.algorithms.common.hidden_states_data import ( + PackedHiddenStatesCollator, + ) + from specforge.algorithms.eagle3.data import DataCollatorWithPacking + from specforge.inference.adapters.server_capture import ( + SGLangServerCaptureAdapter, + ) + from specforge.inference.capture import CaptureConfig + from specforge.launch import build_disagg_online_consumer + from specforge.optimizer import BF16Optimizer + from specforge.runtime.contracts import SampleRef + from specforge.runtime.data_plane.streaming_ref_channel import ( + StreamingRefChannel, + ) + from specforge.training.checkpoint import STATE_FILE + from tests.test_runtime import _fixtures as fx + + fx.build_single_rank_distributed(port="29692") + for name in ("eagle3", "dflash", "dflash2"): + with self.subTest(architecture=name): + algorithm_name = "dflash" if name == "dflash2" else name + run_id = f"packing-online-{name}" + work = Path(self.workdir) / name + work.mkdir() + store = self._packing_store(run_id) + try: + adapter = SGLangServerCaptureAdapter( + f"http://127.0.0.1:{gate.PORT}", + store, + run_id=run_id, + algorithm=algorithm_name, + schema=gate._capture_schema(algorithm_name), + ) + required = {"input_ids", "loss_mask"} | ( + {"attention_mask", "hidden_state", "target"} + if name == "eagle3" + else {"hidden_states"} + ) + contract = CaptureConfig.from_strategy( + required_features=required, + aux_hidden_state_layer_ids=tuple(gate.AUX_LAYER_IDS), + target_repr="hidden_state", + target_hidden_size=gate.H, + ) + rows = [ + list(range(3, 3 + length)) for length in (8, 24, 12, 16) * 2 + ] + tasks = self._tasks(rows) + refs = list(adapter.produce_refs(tasks, capture=contract)) + self.assertEqual(len(refs), 8) + self.assertTrue( + all(isinstance(ref, SampleRef) for ref in refs), repr(refs) + ) + channel = StreamingRefChannel(str(work / "refs.jsonl")) + channel.publish_many(refs) + channel.close() + model, target_head = self._model(name, work) + logged = [] + database = str(work / "consumer.sqlite") + trainer = build_disagg_online_consumer( + algorithm=builtin_algorithm_registry().resolve(algorithm_name), + feature_store=store, + channel=channel, + draft_model=model, + target_head=target_head, + optimizer_factory=lambda module: BF16Optimizer( + module, + lr=1e-3, + max_grad_norm=0.5, + warmup_ratio=0.0, + total_steps=2, + ), + run_id=run_id, + output_dir=str(work / "output"), + batch_size=2, + accumulation_steps=2, + max_steps=2, + sequence_packing=True, + save_interval=1, + log_interval=1, + metadata_db_path=database, + async_ack=False, + idle_timeout_s=60, + logger=lambda metrics, step: logged.append( + (dict(metrics), step) + ), + ) + self.assertIsInstance( + trainer._loader.collate_fn, + ( + DataCollatorWithPacking + if name == "eagle3" + else PackedHiddenStatesCollator + ), + ) + self.assertEqual(trainer.fit(), 2) + self.assertEqual(trainer.micro_step, 4) + self.assertEqual(channel.consumer_quantum(), 4) + self.assertTrue(channel.consumer_stopped()) + self.assertIsNone(channel.consumer_failure()) + self.assertEqual([step for _, step in logged], [1, 2]) + for metrics, _ in logged: + for key, value in metrics.items(): + if isinstance(value, numbers.Real): + self.assertTrue(math.isfinite(value), f"{key}={value}") + with sqlite3.connect(database) as connection: + acked = [ + row[0] + for row in connection.execute( + "SELECT sample_id FROM acked ORDER BY sample_id" + ) + ] + self.assertEqual(acked, sorted(ref.sample_id for ref in refs)) + for step in (1, 2): + self.assertTrue( + ( + work / "output" / f"{run_id}-step{step}" / STATE_FILE + ).is_file() + ) + result = { + "architecture": name, + "optimizer_steps": trainer.global_step, + "microsteps": trainer.micro_step, + "acked_samples": len(acked), + "sample_lengths": [len(row) for row in rows], + "consumer_quantum": channel.consumer_quantum(), + "logged_steps": [step for _, step in logged], + } + (Path(self.workdir) / f"{name}-result.json").write_text( + json.dumps(result, indent=2) + ) + del trainer, model, target_head + torch.cuda.empty_cache() + finally: + store.close() diff --git a/tests/test_runtime/test_packing_strategy_metadata.py b/tests/test_runtime/test_packing_strategy_metadata.py new file mode 100644 index 000000000..1095deebb --- /dev/null +++ b/tests/test_runtime/test_packing_strategy_metadata.py @@ -0,0 +1,94 @@ +"""DFlash packing metadata remains host-side across the wrapped model boundary.""" + +import unittest + +import torch +from torch import nn + +from specforge.runtime.contracts import TrainBatch +from specforge.training.strategies.base import DFlashTrainStrategy, StepContext + + +class _RecordingModel(nn.Module): + def __init__(self, device="cpu"): + super().__init__() + self.weight = nn.Parameter(torch.ones((), device=device)) + self.kwargs = None + + def forward(self, **kwargs): + self.kwargs = kwargs + return ( + self.weight * kwargs["hidden_states"].sum(), + self.weight.new_zeros(()), + {}, + ) + + +def _batch(loss_mask, lengths): + return TrainBatch( + sample_ids=[str(i) for i in range(len(lengths))], + strategy="dflash", + tensors={ + "input_ids": torch.zeros_like(loss_mask, dtype=torch.long), + "loss_mask": loss_mask, + "hidden_states": torch.ones(1, sum(lengths), 4, device=loss_mask.device), + "sequence_lengths": torch.tensor(lengths, dtype=torch.long), + }, + metadata={}, + ) + + +class PackingStrategyMetadataTest(unittest.TestCase): + def test_host_counts_respect_document_boundaries_and_stay_python_values(self): + model = _RecordingModel() + mask = torch.tensor([[1, 1, 1, 1, 0, 1, 1, 1, 0]]) + DFlashTrainStrategy(model).forward_loss(_batch(mask, (3, 4, 2)), StepContext()) + self.assertEqual(model.kwargs["sequence_lengths"], (3, 4, 2)) + self.assertEqual(model.kwargs["valid_anchor_counts"], (2, 1, 0)) + self.assertEqual(model.kwargs["max_valid_anchors"], 2) + self.assertIsInstance(model.kwargs["max_valid_anchors"], int) + self.assertTrue( + all(isinstance(n, int) for n in model.kwargs["valid_anchor_counts"]) + ) + + def test_zero_masks_and_cross_document_pair_have_no_anchors(self): + for mask in (torch.zeros(1, 4), torch.tensor([[0, 1, 1, 0]])): + with self.subTest(mask=mask.tolist()): + model = _RecordingModel() + DFlashTrainStrategy(model).forward_loss(_batch(mask, (2, 2))) + self.assertEqual(model.kwargs["valid_anchor_counts"], (0, 0)) + self.assertEqual(model.kwargs["max_valid_anchors"], 0) + + def test_model_rejects_no_anchor_batch_before_attention(self): + from specforge.algorithms.common.dflash_family_model import OnlineDFlashModel + + model = OnlineDFlashModel( + nn.Linear(4, 4), + nn.Linear(4, 8, bias=False), + nn.Embedding(8, 4), + mask_token_id=0, + block_size=2, + num_anchors=2, + attention_backend="flex_attention", + ) + for mask in (torch.zeros(1, 4), torch.tensor([[0, 1, 1, 0]])): + with ( + self.subTest(mask=mask.tolist()), + self.assertRaisesRegex(ValueError, "two consecutive supervised"), + ): + DFlashTrainStrategy(model).forward_loss(_batch(mask, (2, 2))) + + @unittest.skipUnless( + torch.cuda.is_available(), "GPU loss-mask fallback requires CUDA" + ) + def test_gpu_loss_mask_uses_model_fallback_without_host_counts(self): + model = _RecordingModel("cuda") + batch = _batch(torch.tensor([[1, 1, 0, 1, 1]], device="cuda"), (3, 2)) + DFlashTrainStrategy(model).forward_loss(batch) + self.assertEqual(model.kwargs["sequence_lengths"], (3, 2)) + self.assertIsNone(model.kwargs["max_valid_anchors"]) + self.assertNotIn("valid_anchor_counts", model.kwargs) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_runtime/test_sequence_packing_lifecycle.py b/tests/test_runtime/test_sequence_packing_lifecycle.py new file mode 100644 index 000000000..9bf978da1 --- /dev/null +++ b/tests/test_runtime/test_sequence_packing_lifecycle.py @@ -0,0 +1,118 @@ +"""Packed offline loader -> FSDP -> optimizer/eval/checkpoint GPU smoke test.""" + +import math +import numbers +import tempfile +import unittest +from pathlib import Path + +import torch + +from specforge.algorithms.builtin import builtin_algorithm_registry +from specforge.algorithms.eagle3.data import DataCollatorWithPacking +from tests.test_runtime import _fixtures as fx + + +def _write_variable_length_features(directory, lengths): + directory.mkdir() + generator = torch.Generator().manual_seed(17) + for index, length in enumerate(lengths): + torch.save( + { + "input_ids": torch.randint(0, fx.V, (length,), generator=generator), + "loss_mask": torch.ones(length, dtype=torch.long), + "hidden_state": torch.randn( + 1, length, fx.H, generator=generator + ).bfloat16(), + "aux_hidden_state": torch.randn( + 1, length, 3 * fx.H, generator=generator + ).bfloat16(), + }, + directory / f"{index:04d}.ckpt", + ) + return str(directory) + + +@unittest.skipUnless( + torch.cuda.is_available(), "packed trainer lifecycle requires CUDA" +) +class SequencePackingLifecycleTest(unittest.TestCase): + def test_packed_training_preserves_steps_evaluation_and_checkpoints(self): + from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + + from specforge.launch import build_offline_runtime + from specforge.optimizer import BF16Optimizer + from specforge.training.checkpoint import STATE_FILE + + torch.manual_seed(17) + fx.build_single_rank_distributed(port="29687") + logged = [] + with tempfile.TemporaryDirectory(prefix="packed_lifecycle_") as workdir: + work = Path(workdir) + train_path = _write_variable_length_features( + work / "train", (8, 24, 12, 16) * 2 + ) + eval_path = _write_variable_length_features(work / "eval", (8, 24, 12)) + model, target_head = fx.build_eagle3(workdir, ttt=3) + output = work / "output" + + def optimizer_factory(draft_module): + return BF16Optimizer( + draft_module, + lr=1e-3, + max_grad_norm=0.5, + warmup_ratio=0.0, + total_steps=2, + ) + + trainer = build_offline_runtime( + algorithm=builtin_algorithm_registry().resolve("eagle3"), + hidden_states_path=train_path, + eval_hidden_states_path=eval_path, + draft_model=model, + target_head=target_head, + optimizer_factory=optimizer_factory, + run_id="packing-lifecycle", + output_dir=str(output), + ttt_length=3, + max_len=32, + batch_size=2, + sequence_packing=True, + accumulation_steps=2, + num_epochs=1, + max_steps=2, + eval_interval=1, + save_interval=1, + log_interval=1, + logger=lambda metrics, step: logged.append((dict(metrics), step)), + ) + self.assertIsInstance(trainer.core.strategy.trainable_module(), FSDP) + self.assertIsInstance(trainer._loader.collate_fn, DataCollatorWithPacking) + eval_loader = trainer._controller.eval_data_factory() + self.assertIsInstance(eval_loader.collate_fn, DataCollatorWithPacking) + self.assertEqual([len(b.sample_ids) for b in eval_loader], [2, 1]) + + self.assertEqual(trainer.fit(), 2) + self.assertEqual(trainer.global_step, 2) + self.assertEqual(trainer.micro_step, 4) + self.assertEqual(trainer.last_checkpoint_step, 2) + self.assertEqual( + [step for metrics, step in logged if "eval/avg_loss" in metrics], [1, 2] + ) + self.assertTrue(any("loss" in metrics for metrics, _ in logged)) + for metrics, _ in logged: + for name, value in metrics.items(): + if isinstance(value, numbers.Real): + self.assertTrue(math.isfinite(value), f"{name}={value}") + for step in (1, 2): + checkpoint = output / f"packing-lifecycle-step{step}" / STATE_FILE + self.assertTrue(checkpoint.is_file()) + state = torch.load(checkpoint, map_location="cpu", weights_only=False) + self.assertEqual(state["global_step"], step) + self.assertEqual(state["epoch_samples"], 4 * step) + self.assertTrue(state["draft_state_dict"]) + self.assertTrue((output / "packing-lifecycle-latest").is_dir()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_runtime/test_sequence_packing_wiring.py b/tests/test_runtime/test_sequence_packing_wiring.py new file mode 100644 index 000000000..f7f532228 --- /dev/null +++ b/tests/test_runtime/test_sequence_packing_wiring.py @@ -0,0 +1,236 @@ +"""Packing reaches both offline training topologies and offline evaluation.""" + +import tempfile +import unittest +from unittest import mock + +from specforge.algorithms.builtin import builtin_algorithm_registry +from specforge.algorithms.eagle3.data import DataCollatorWithPacking +from specforge.config import Config +from specforge.launch import ( + _make_offline_eval_data_factory, + _offline_io, + _streaming_collate, + build_disagg_offline_runtime, + build_offline_runtime, +) +from specforge.training.assembly import ModelBundle, build_training_run +from specforge.training.disaggregated import _build_offline, _build_online + +ALGORITHM = builtin_algorithm_registry().resolve("eagle3") + + +def _config(): + return Config.model_validate( + { + "model": { + "target_model_path": "target", + "draft_model_config": "draft.json", + "vocab_mapping_path": "mapping.pt", + }, + "data": { + "hidden_states_path": "/train-features", + "eval_hidden_states_path": "/eval-features", + }, + "training": {"sequence_packing": True, "batch_size": 3, "eval_interval": 1}, + } + ) + + +def _bundle(): + return ModelBundle( + model=object(), + draft_model=object(), + draft_config=object(), + target_head=object(), + strategy_kwargs={}, + ) + + +class SequencePackingWiringTest(unittest.TestCase): + def test_io_selects_packing_and_preserves_the_normalizer(self): + packed, normalizer = _offline_io( + ALGORITHM, + "text", + 123, + ttt_length=3, + use_usp_preprocess=False, + sequence_packing=True, + ) + padded, baseline = _offline_io( + ALGORITHM, + "text", + 123, + ttt_length=3, + use_usp_preprocess=False, + ) + self.assertIsInstance(packed, DataCollatorWithPacking) + self.assertNotIsInstance(padded, DataCollatorWithPacking) + self.assertIs(normalizer.func, baseline.func) + self.assertEqual(normalizer.keywords, baseline.keywords) + + def test_direct_io_rejects_usp_and_unsupported_provider(self): + for algorithm, usp in ( + (ALGORITHM, True), + (builtin_algorithm_registry().resolve("domino"), False), + ): + with ( + self.subTest(algorithm=algorithm.name, usp=usp), + self.assertRaisesRegex( + ValueError, "supported non-USP offline provider" + ), + ): + _offline_io( + algorithm, + "text", + 123, + ttt_length=3, + use_usp_preprocess=usp, + sequence_packing=True, + ) + + def test_streaming_selects_the_registered_packing_factory(self): + for name in ("eagle3", "dflash"): + algorithm = builtin_algorithm_registry().resolve(name) + collator = _streaming_collate( + algorithm, "text", None, sequence_packing=True + ) + self.assertTrue(callable(collator)) + with self.assertRaisesRegex(ValueError, "packing"): + _streaming_collate( + builtin_algorithm_registry().resolve("domino"), + "text", + None, + sequence_packing=True, + ) + with self.assertRaisesRegex(ValueError, "collate"): + _streaming_collate( + ALGORITHM, "text", lambda samples: samples, sequence_packing=True + ) + + def test_eval_factory_selects_packing_and_keeps_partial_batches(self): + with tempfile.TemporaryDirectory() as directory: + factory = _make_offline_eval_data_factory( + algorithm=ALGORITHM, + modality="text", + hidden_states_path=directory, + run_id="packing-eval", + batch_size=3, + max_len=123, + ttt_length=3, + use_usp_preprocess=False, + dataloader_num_workers=0, + sequence_packing=True, + ) + first, second = factory(), factory() + self.assertIsNot(first, second) + self.assertIsInstance(first.collate_fn, DataCollatorWithPacking) + self.assertEqual(first.batch_size, 3) + self.assertFalse(first.drop_last) + + def test_both_offline_builders_pack_train_and_eval(self): + with tempfile.TemporaryDirectory() as directory: + for builder in (build_offline_runtime, build_disagg_offline_runtime): + with ( + self.subTest(builder=builder.__name__), + mock.patch("specforge.launch._assemble_trainer") as assemble, + ): + data_args = ( + {"hidden_states_path": directory} + if builder is build_offline_runtime + else {"feature_store": object(), "refs": []} + ) + builder( + algorithm=ALGORITHM, + draft_model=object(), + target_head=object(), + optimizer_factory=object(), + run_id="packing-test", + output_dir=directory, + batch_size=3, + accumulation_steps=2, + eval_hidden_states_path=directory, + sequence_packing=True, + **data_args, + ) + kwargs = assemble.call_args.kwargs + self.assertIsInstance(kwargs["collate_fn"], DataCollatorWithPacking) + self.assertEqual(kwargs["batch_size"], 3) + self.assertEqual(kwargs["accumulation_steps"], 2) + self.assertIsInstance( + kwargs["eval_data_factory"]().collate_fn, + DataCollatorWithPacking, + ) + + def test_unified_offline_assembly_passes_packing(self): + with ( + mock.patch( + "specforge.training.assembly.build_model_bundle", return_value=_bundle() + ), + mock.patch("specforge.launch.build_offline_runtime") as build, + ): + build_training_run(_config(), algorithm=ALGORITHM) + self.assertTrue(build.call_args.kwargs["sequence_packing"]) + self.assertEqual(build.call_args.kwargs["batch_size"], 3) + + def test_disaggregated_consumer_assembly_passes_packing(self): + cfg = _config() + with ( + mock.patch( + "specforge.training.disaggregated._env", return_value="/manifest.json" + ), + mock.patch("specforge.training.disaggregated._wait_for"), + mock.patch("specforge.training.disaggregated._offline_store"), + mock.patch( + "specforge.runtime.data_plane.disagg_ingest.read_ref_manifest", + return_value=[], + ), + mock.patch("specforge.launch.build_disagg_offline_runtime") as build, + ): + _build_offline( + cfg, + algorithm=ALGORITHM, + build_model_bundle=lambda _: _bundle(), + optimizer_factory=lambda _: object(), + logger=None, + ) + self.assertTrue(build.call_args.kwargs["sequence_packing"]) + self.assertEqual(build.call_args.kwargs["batch_size"], 3) + + def test_online_consumer_assembly_passes_packing_and_logical_batch_size(self): + payload = _config().model_dump() + payload["data"] = {"train_data_path": "train.jsonl"} + payload["training"].update(role="consumer", max_steps=2, eval_interval=0) + payload["deployment"] = { + "mode": "disaggregated", + "disaggregated": { + "control_dir": "outputs/packing-test/control", + "backend": "mooncake", + "server_urls": ["http://127.0.0.1:30000"], + "mooncake_metadata_server": "http://127.0.0.1:35880/metadata", + "mooncake_master_server_addr": "127.0.0.1:35551", + }, + } + cfg = Config.model_validate(payload) + with ( + mock.patch( + "specforge.training.disaggregated._env", + return_value="/ref-channel.jsonl", + ), + mock.patch("specforge.training.disaggregated._mooncake_store"), + mock.patch("specforge.launch.build_disagg_online_consumer") as build, + ): + _build_online( + cfg, + algorithm=ALGORITHM, + build_model_bundle=lambda _: _bundle(), + prepare_prompts=lambda *_args, **_kwargs: [], + optimizer_factory=lambda _: object(), + logger=None, + ) + self.assertTrue(build.call_args.kwargs["sequence_packing"]) + self.assertEqual(build.call_args.kwargs["batch_size"], 3) + + +if __name__ == "__main__": + unittest.main()