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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,13 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0

### Added

- **Sequence packing under pipeline parallelism**, for both attention backends. `doc_ids` now reaches every pipeline stage, so `pp > 1` with `data.pack_sequences = true` trains with real block-diagonal attention instead of being rejected. It travels as a schedule **kwarg**: positional args reach only stage 0, but kwargs are handed to every stage, and both are split into micro-batches the same way -- so alignment with the tokens is automatic and costs no extra sends.
- It deliberately does *not* ride the pipe between stages, which was the first attempt. `PipelineStage` treats every inter-stage tensor as a differentiable activation: `_prepare_forward_infra` calls `requires_grad_(True)` on each received buffer with no dtype check, and `_create_grad_send_info` registers a gradient send for every received input "regardless of whether it requires grad" (its own comment). An int64 tensor therefore cannot cross a stage boundary at all, and no dtype adjustment avoids it.
- `kempnerforge/distributed/pipeline_parallel.py`: `PipelineStageModule.forward` takes `doc_ids` and returns hidden states only -- never a tuple -- and builds its own `BlockMask` per stage under `attention_backend="flex"`, since a `BlockMask` is not a tensor and could not be transmitted either.
- `kempnerforge/training/entry.py` / `loop.py`: `pipeline_step` concatenates the micro-batch `doc_ids` and passes them on every rank's own `schedule.step()` call.
- `kempnerforge/config/job.py`: the `pp > 1` + packing rejection is **removed** -- both backends now work under PP.
- Tests: `tests/distributed/test_pp.py` (new). `test_pipeline_output_matches_single_gpu` covers the unpacked pipeline; `test_packed_pipeline_matches_single_gpu[sdpa|flex]` asserts the pipeline matches a single-GPU *packed* forward **and** differs from the unpacked one; and `test_packed_pipeline_training_step[sdpa|flex][1f1b|gpipe]` runs a real step with a loss function and gradients enabled. That last one is not redundant: a schedule built without `loss_fn` leaves `has_backward=False`, which skips the code the first attempt broke, so forward-only coverage reported success on a configuration that could not train.
- Docs: `docs/configuration/validation-rules.md` drops the rule this removes.
- **FlexAttention backend for packed sequences** (`model.attention_backend = "flex"`). Document packing built a dense `(B, 1, S, S)` boolean mask *in every layer* and handed it to `F.scaled_dot_product_attention`; an explicit `attn_mask` is not a FlashAttention-2 shape, so SDPA fell back to the mem-efficient/math kernel and packing gave up throughput to save padding. `"flex"` instead builds one FlexAttention `BlockMask` per forward, shared by every layer, so the mask predicate compiles into the attention kernel and fully-masked blocks are skipped rather than computed. The default `"sdpa"` is bit-identical to before, and the flag only does anything when `data.pack_sequences = true` — unpacked batches take the `is_causal` SDPA fast path under either backend.
- `kempnerforge/model/masking.py` (new): `build_doc_causal_block_mask` (block-diagonal causal; `H=None` so the mask broadcasts over heads, which is what keeps it correct under tensor parallelism's sharded head counts) and `flex_attention_fn`, which compiles `flex_attention` once per process on CUDA and stays eager on CPU.
- `kempnerforge/model/attention.py`: a `block_mask` branch ahead of the dense-mask path. GQA goes through `enable_gqa` instead of `repeat_interleave`, so the flex path never materializes the repeated K/V heads. `capture_attention_weights` raises under flex — the kernel fuses the mask and never forms an attention-weight matrix to capture.
Expand Down
1 change: 0 additions & 1 deletion docs/configuration/validation-rules.md
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,6 @@ File:
unaffected at every length, so existing packed runs are not at risk. Not
reduced further; the bound is empirical. It costs nothing, since a sequence
below the mask block size has no block sparsity to exploit.
- `pp > 1` rejects `data.pack_sequences` (pipeline stages never receive `doc_ids`).
- When `num_experts > 0` (MoE):
- `moe_top_k > 0`
- `moe_top_k ≤ num_experts`
Expand Down
8 changes: 0 additions & 8 deletions kempnerforge/config/job.py
Original file line number Diff line number Diff line change
Expand Up @@ -139,14 +139,6 @@ def __post_init__(self) -> None:
f"got {self.train.seq_len}. Use attention_backend='sdpa' for shorter sequences."
)

# Pipeline stages never receive doc_ids, so packing would silently train
# on cross-document context. See issue #204.
if self.distributed.pp > 1 and self.data.pack_sequences:
raise ValueError(
"Sequence packing + Pipeline Parallelism is not supported. "
"Set data.pack_sequences=false, or train without pipeline parallelism."
)

@property
def is_vlm(self) -> bool:
"""Whether this job builds a ``VLMWrapper`` around the text backbone."""
Expand Down
28 changes: 25 additions & 3 deletions kempnerforge/distributed/pipeline_parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
from kempnerforge.config.schema import ModelConfig
from kempnerforge.model.embedding import OutputHead, TokenEmbedding
from kempnerforge.model.init import init_weights
from kempnerforge.model.masking import build_doc_causal_block_mask
from kempnerforge.model.norm import build_norm
from kempnerforge.model.position import precompute_rope_frequencies
from kempnerforge.model.transformer import TransformerBlock
Expand Down Expand Up @@ -140,7 +141,6 @@ def __init__(
self.num_stages = num_stages
self.is_first = stage_id == 0
self.is_last = stage_id == num_stages - 1

start, end = layer_range

# Token embedding — only on first stage
Expand Down Expand Up @@ -189,12 +189,19 @@ def init_weights_and_freqs(self) -> None:
)
init_weights(self, self.config)

def forward(self, x: torch.Tensor) -> torch.Tensor:
def forward(self, x: torch.Tensor, doc_ids: torch.Tensor | None = None) -> torch.Tensor:
"""Forward pass for this pipeline stage.

Args:
x: For stage 0: token IDs of shape (batch, seq_len).
For other stages: hidden states of shape (batch, seq_len, dim).
doc_ids: Per-token document ids, shape (batch, seq_len), when
sequence packing is on. Supplied as a *schedule kwarg*: the
schedule splits it into micro-batches alongside the tokens and
hands it to every stage directly. It must not travel between
stages as an activation -- ``PipelineStage`` calls
``requires_grad_(True)`` on every received buffer without
checking dtype, which an integer tensor cannot satisfy.

Returns:
For last stage: logits of shape (batch, seq_len, vocab_size).
Expand All @@ -212,9 +219,22 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
cos = self._rope_cos[:seq_len] # type: ignore[reportOptionalSubscript]
sin = self._rope_sin[:seq_len] # type: ignore[reportOptionalSubscript]

# Packed sequences under attention_backend="flex": build the BlockMask
# once here and share it across this stage's layers, mirroring
# Transformer.forward. Each stage builds its own from the doc_ids it
# was handed -- cheap next to attention, and a BlockMask is not a tensor,
# so it could not be transmitted between stages anyway.
block_mask = None
layer_doc_ids = doc_ids
if doc_ids is not None and self.config.attention_backend == "flex":
block_mask = build_doc_causal_block_mask(doc_ids, x.device)
# The BlockMask supersedes doc_ids; clearing it keeps the dense
# SDPA mask branch in Attention.forward unreachable on this path.
layer_doc_ids = None

# Run through assigned layers
for layer in self.layers.values():
x = layer(x, cos, sin)
x = layer(x, cos, sin, doc_ids=layer_doc_ids, block_mask=block_mask)

# Last stage: norm + output head
if self.is_last:
Expand Down Expand Up @@ -295,6 +315,8 @@ def build_pipeline_stage(
device=device,
),
)
# doc_ids is deliberately absent here: it reaches stages as a schedule
# kwarg, which is neither transmitted nor shape-inferred.

return PipelineStage(
submodule=stage_module,
Expand Down
1 change: 0 additions & 1 deletion kempnerforge/training/entry.py
Original file line number Diff line number Diff line change
Expand Up @@ -106,7 +106,6 @@ def build_model(
pp_size = get_pp_size(device_mesh)

tp_enabled_pp = "tp" in device_mesh.mesh_dim_names # type: ignore[reportOperatorIssue]

# apply_{float8,ac,fsdp2} are annotated for Transformer; a stage module
# exposes the same block structure, hence the cast.
if tp_enabled_pp:
Expand Down
15 changes: 11 additions & 4 deletions kempnerforge/training/loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -193,12 +193,14 @@ def pipeline_step(session: TrainingSession, step: int) -> StepResult:

# Collect microbatches into a full batch for the schedule.
# schedule.step() splits along dim 0 into n_microbatches.
input_ids_list, labels_list = [], []
input_ids_list, labels_list, doc_ids_list = [], [], []
for _ in range(tc.grad_accum_steps):
if batches.has_data:
batch = batches.next_batch()
input_ids_list.append(batch["input_ids"].to(device))
labels_list.append(batch["labels"].to(device))
if "doc_ids" in batch:
doc_ids_list.append(batch["doc_ids"].to(device))
else:
input_ids_list.append(
torch.randint(0, mc.vocab_size, (tc.batch_size, tc.seq_len), device=device)
Expand All @@ -209,6 +211,11 @@ def pipeline_step(session: TrainingSession, step: int) -> StepResult:

full_input = torch.cat(input_ids_list, dim=0)
full_labels = torch.cat(labels_list, dim=0)
# Packed runs pass doc_ids as a schedule *kwarg*. Positional args reach only
# stage 0, but kwargs are handed to every stage, and both are split into
# micro-batches the same way -- so alignment with the tokens is automatic and
# doc_ids never becomes an activation the backward pass tries to grade.
pp_kwargs = {"doc_ids": torch.cat(doc_ids_list, dim=0)} if doc_ids_list else {}

# The schedule handles forward/backward for all microbatches.
# First stage needs input; last stage needs target for loss.
Expand All @@ -219,11 +226,11 @@ def pipeline_step(session: TrainingSession, step: int) -> StepResult:
pp_losses: list[torch.Tensor] = []

if is_first:
pipeline.schedule.step(full_input, target=full_labels, losses=pp_losses)
pipeline.schedule.step(full_input, target=full_labels, losses=pp_losses, **pp_kwargs)
elif is_last:
pipeline.schedule.step(target=full_labels, losses=pp_losses)
pipeline.schedule.step(target=full_labels, losses=pp_losses, **pp_kwargs)
else:
pipeline.schedule.step()
pipeline.schedule.step(**pp_kwargs)

# Loss is only meaningful on the last stage
if is_last and pp_losses:
Expand Down
Loading
Loading