Skip to content

Commit f53db79

Browse files
authored
Merge pull request sgl-project#731 from modal-projects/dcw02/dflash-correctness
Correctness fixes for DFlash training path
2 parents 5c93516 + edc450a commit f53db79

34 files changed

Lines changed: 1251 additions & 222 deletions

docs/basic_usage/training.md

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -154,6 +154,31 @@ model:
154154
draft_block_size: 8 # DFlash only
155155
```
156156

157+
For DFlash, configure the attention layout in the referenced draft JSON. Each
158+
entry corresponds to one draft layer; sliding layers share one positive window:
159+
160+
```json
161+
{
162+
"num_hidden_layers": 5,
163+
"layer_types": [
164+
"sliding_attention",
165+
"sliding_attention",
166+
"sliding_attention",
167+
"sliding_attention",
168+
"full_attention"
169+
],
170+
"use_sliding_window": true,
171+
"sliding_window": 2048
172+
}
173+
```
174+
175+
Use `"full_attention"` for every entry, `"use_sliding_window": false`, and
176+
`"sliding_window": null` for a full-only draft. The layout length must equal
177+
`num_hidden_layers`. A layer-count override may resize a uniform layout, but a
178+
mixed layout must be edited explicitly in the draft JSON.
179+
180+
The `eager`, `sdpa`, and `flex_attention` backends support both layouts.
181+
157182
Domino and DSpark need their projector/head metadata, so they require an
158183
explicit draft config (or a pretrained warm-start source that contains
159184
`config.json`). The old Domino parser exposed an optional config flag, but its

scripts/prepare_hidden_states.py

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -50,7 +50,7 @@
5050
from concurrent.futures import ThreadPoolExecutor
5151
from dataclasses import dataclass
5252
from pathlib import Path
53-
from typing import Dict, List, Mapping, Optional
53+
from typing import Callable, Dict, List, Mapping, Optional
5454

5555
import torch
5656
import torch.distributed as dist
@@ -92,6 +92,7 @@ class OfflineCapturePlan:
9292
capture_method: str
9393
capture_layers: tuple[int, ...]
9494
layout: OfflineCaptureLayout
95+
loss_mask_filter: Optional[Callable[[object], bool]]
9596

9697

9798
def parse_args():
@@ -341,6 +342,7 @@ def resolve_offline_capture_plan(
341342
capture_method=resolved.capture_method,
342343
capture_layers=resolved.capture_layers,
343344
layout=resolved.layout,
345+
loss_mask_filter=resolved.loss_mask_filter,
344346
)
345347

346348

@@ -826,7 +828,7 @@ def main():
826828
),
827829
num_proc=min(args.build_dataset_num_proc, 32),
828830
)
829-
if args.num_samples is not None:
831+
if args.num_samples is not None and capture_plan.loss_mask_filter is None:
830832
dataset = dataset.select(range(args.num_samples))
831833
# Tokenizer and cache key
832834
tokenizer = load_tokenizer(
@@ -847,6 +849,15 @@ def main():
847849
cache_key=cache_key,
848850
is_preformatted=args.is_preformatted,
849851
num_proc=args.build_dataset_num_proc,
852+
loss_mask_filter=capture_plan.loss_mask_filter,
853+
)
854+
if capture_plan.loss_mask_filter is not None and args.num_samples is not None:
855+
eagle3_dataset = eagle3_dataset.select(
856+
range(min(args.num_samples, len(eagle3_dataset)))
857+
)
858+
if not len(eagle3_dataset):
859+
raise ValueError(
860+
f"no samples satisfy {capture_plan.strategy} training eligibility"
850861
)
851862
print_with_rank(f"Dataset prepared with {len(eagle3_dataset)} samples.")
852863

specforge/algorithms/common/dflash_family_data.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
from functools import partial
66

77
from specforge.algorithms.common.collation import pad_and_concatenate_features
8+
from specforge.data.loss_mask import has_consecutive_supervised_tokens
89

910
NORMALIZER_ID = "dflash_family_offline_v1"
1011
DSPARK_NORMALIZER_ID = "dspark_offline_v1"
@@ -56,6 +57,10 @@ def normalize_offline_sample(raw, max_len: int):
5657
f"loss_mask={loss_mask.shape[1]}, "
5758
f"hidden_states={hidden_states.shape[1]}"
5859
)
60+
if not has_consecutive_supervised_tokens(loss_mask[0]):
61+
raise ValueError(
62+
"offline DFlash-family samples require two consecutive supervised tokens"
63+
)
5964
return {
6065
"input_ids": input_ids,
6166
"loss_mask": loss_mask,

0 commit comments

Comments
 (0)