Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
147 commits
Select commit Hold shift + click to select a range
f89213a
Add versioned trainer state and logical rank execution
bradhilton Sep 18, 2026
407f887
Use native topology type in graph cache GPU fixtures
bradhilton Sep 18, 2026
3770f95
Model independent forward groups separately in graph lifetime test
bradhilton Sep 18, 2026
00a3f6f
test: declare topology for split graph lifetime fixture
bradhilton Sep 18, 2026
6d62ce2
refactor: remove unused memory placement selector
bradhilton Sep 19, 2026
5475043
Stamp trainer head registration and export observations
bradhilton Sep 19, 2026
0746fca
fix: remap trainer payloads and compact operation retirement
bradhilton Sep 19, 2026
d5b4fb2
Share trainer v1 test and development setup
bradhilton Sep 19, 2026
5a8e491
Fix logical head completion and lookup boundaries
bradhilton Sep 19, 2026
daf6210
Merge commit '5b7acf06d209dffa70327a810d3a9e303411ab77' into stark/re…
bradhilton Sep 19, 2026
4c89eb4
Expose callback release fence for checkpoint entry
bradhilton Sep 19, 2026
04334c3
Separate abandoned batch pull cleanup from iterator close
bradhilton Sep 19, 2026
bbb81f3
Deduplicate live-head test construction
bradhilton Sep 20, 2026
e92db78
Reuse tensor argument container traversal
bradhilton Sep 20, 2026
bcaa2f0
Consolidate trainer-rank test process scaffolding
bradhilton Sep 20, 2026
7365d9a
Inherit duplicate TrainerRank public methods
bradhilton Sep 20, 2026
8f0f35e
Remove obsolete facade type variables and imports
bradhilton Sep 20, 2026
3fc523e
Reuse trainer test process group fixtures
bradhilton Sep 20, 2026
1bcf119
Reuse live parameter test fixtures
bradhilton Sep 20, 2026
0429b71
Reuse custom tensor wrapping for module members
bradhilton Sep 20, 2026
1d9d6ae
Share policy choice checks and preserve public hint bindings
bradhilton Sep 20, 2026
101fd1f
Share argument container traversal with replay restore
bradhilton Sep 20, 2026
6e61553
Use real trainer typing and share CPU gradient stubs
bradhilton Sep 20, 2026
1a0bc40
Share native live-head test registration setup
bradhilton Sep 20, 2026
d99dde4
Share strict gradient reducer installation in validation tests
bradhilton Sep 20, 2026
ee5415e
Share native and live head tensor inventories
bradhilton Sep 20, 2026
fb3be23
test: share trainer snapshot setup in version tests
bradhilton Sep 20, 2026
f197898
test: share CP2 attention layout setup
bradhilton Sep 20, 2026
36b127f
test: finish sharing initial version snapshots
bradhilton Sep 20, 2026
de6a8d9
test: share inline trainer operation capture
bradhilton Sep 20, 2026
30ef515
test: run trainer operation cases directly as async tests
bradhilton Sep 20, 2026
c2fb6f8
test: remove redundant trainer async wrappers
bradhilton Sep 20, 2026
797f463
test: share trainer operation imports
bradhilton Sep 21, 2026
ba8e720
Consolidate initialized trainer test imports
bradhilton Sep 21, 2026
38b53de
Share standard-library imports in trainer tests
bradhilton Sep 21, 2026
5b7e1c6
Consolidate remaining trainer test imports
bradhilton Sep 21, 2026
5ef8489
Share trainer head mutation-name constants
bradhilton Sep 21, 2026
1a13ee6
Coordinate authority buffer snapshot failures across ranks
bradhilton Sep 21, 2026
b73cd1b
Release undelivered trainer output graphs on attachment failure
bradhilton Sep 21, 2026
a75373b
Preserve delivery errors while closing trainer iterators
bradhilton Sep 21, 2026
696e24b
Merge main vocabulary-head lifetime fix into trainer v1
bradhilton Sep 21, 2026
022eeaa
Fix trainer output delivery test typing
bradhilton Sep 21, 2026
595f12c
Keep Tinker client imports isolated from backend initialization
bradhilton Sep 21, 2026
4d9c631
Release logical forward batches before advancing
bradhilton Sep 21, 2026
86f587a
Merge current ART main into trainer v1
bradhilton Sep 26, 2026
3b6f05e
Preserve caller RNG while advancing and replaying trainer graphs
bradhilton Sep 26, 2026
d1273f1
fix(trainer-rank): meter physical cached graph backward
bradhilton Sep 26, 2026
0004fc0
Keep RNG and command failure collectives aligned
bradhilton Sep 26, 2026
b7d53fc
fix(trainer-rank): preserve forward errors during RNG sync
bradhilton Sep 26, 2026
ff1cb78
fix(trainer-rank): drain commands after physical cancellation
bradhilton Sep 26, 2026
667f8be
fix(trainer-rank): coordinate logical head factory failures
bradhilton Sep 26, 2026
9d6304c
Restore split-peak counter fixture recompute assumptions
bradhilton Sep 26, 2026
ea19d62
test(trainer-rank): complete command transport topology fixture
bradhilton Sep 26, 2026
bfa189f
Merge commit 'a632f8d8b5d2405f1eb5813f8d7a16914a1809aa' into stark/tr…
bradhilton Sep 26, 2026
09b64fb
Preserve split-floor oracles and share transport topology fixture
bradhilton Sep 26, 2026
229a7f5
test(trainer-rank): kill and reap surviving test workers
bradhilton Sep 26, 2026
b3b5201
fix(trainer-rank): exclude consumed graphs from HybridEP floor
bradhilton Sep 26, 2026
7e5b34a
test(trainer-rank): reuse head construction helpers
bradhilton Sep 26, 2026
dafa85c
Merge current main optimizer counter validation into trainer v1
bradhilton Sep 26, 2026
b61edd3
fix(trainer-rank): reject fractional custom optimizer counters
bradhilton Sep 26, 2026
77ebaa6
ci(trainer-rank): allow time for public GPU checks
bradhilton Sep 26, 2026
fa036b9
fix(trainer-rank): preserve custom counter writer coverage
bradhilton Sep 26, 2026
1139c04
Merge ART main warm-memory profiles into trainer-v1
bradhilton Sep 26, 2026
229f22c
test(trainer-rank): scope temporary command fault injections
bradhilton Sep 26, 2026
fe20204
refactor(trainer-rank): share iterator dispatch and scope test patches
bradhilton Sep 26, 2026
6eea24b
fix(trainer-rank): preserve output placement errors during cleanup
bradhilton Sep 26, 2026
0c6ea44
fix(trainer-rank): preserve transport failures during cleanup
bradhilton Sep 26, 2026
59f9704
refactor(trainer-rank): inline gradient batch commit helper
bradhilton Sep 26, 2026
a5082b1
fix(trainer-rank): distinguish callback stream exhaustion
bradhilton Sep 26, 2026
cac496c
refactor(trainer-rank): simplify checkpoint and stream internals
bradhilton Sep 26, 2026
64c4c13
Merge ART main 0e0c31b3 into trainer-v1
bradhilton Sep 26, 2026
869925f
Coordinate checkpoint slot initialization failures
bradhilton Sep 26, 2026
7e5dabd
Simplify checkpoint failure injection callbacks
bradhilton Sep 26, 2026
d12ee06
Check checkpoint communicator reuse after failed loads
bradhilton Sep 26, 2026
3b04f05
Reuse fresh adapter configuration in checkpoint tests
bradhilton Sep 26, 2026
7cfa820
Reuse coordinated phase handling for rank-zero checkpoint work
bradhilton Sep 26, 2026
3a159e7
Merge ART main c0b1296d for repeated Tau cancellation cleanup
bradhilton Sep 27, 2026
70c0b33
Share exact checkpoint runtime constructor across trainer tests
bradhilton Sep 27, 2026
2099ee2
Retain explicit Linear type in command transport worker
bradhilton Sep 27, 2026
0debb26
Share lightweight GPT model fixture across trainer tests
bradhilton Sep 27, 2026
924000d
Simplify packed trainer test setup
bradhilton Sep 27, 2026
9f6724f
Remove redundant correction kind state
bradhilton Sep 27, 2026
70588af
Remove unused native module owner state
bradhilton Sep 27, 2026
a4187dc
Simplify inherited checkpoint finalization arbitration
bradhilton Sep 27, 2026
6918176
Merge exact ART main 762c89df into trainer-v1
bradhilton Sep 27, 2026
5d61a6b
Allow cleanup-only abort after checkpoint finish
bradhilton Sep 27, 2026
c57cb1c
Share fixed fake rank constructors in memory tests
bradhilton Sep 28, 2026
9e19ef6
Merge current ART main with fake rank test cleanup
bradhilton Sep 28, 2026
a704674
Merge current ART main into trainer-v1 candidate
bradhilton Sep 28, 2026
e4545fc
Share full recompute config in memory fixtures
bradhilton Sep 28, 2026
788299c
Merge ART main checkpoint-floor reuse and GPU validation reuse
bradhilton Sep 28, 2026
4a3da11
Share real transformer model shells in memory fixtures
bradhilton Sep 28, 2026
067b203
Simplify checkpoint memory test cases and config setup
bradhilton Sep 28, 2026
f89d83b
Consolidate checkpoint profile and TP pricing test scenarios
bradhilton Sep 28, 2026
63b91ca
Normalize packed GPT-OSS exports in the native builder
bradhilton Sep 28, 2026
ef7367f
Merge qualified checkpoint fixture reductions
bradhilton Sep 28, 2026
949f03e
Merge fixed ART main trajectory serialization updates
bradhilton Sep 28, 2026
f347e1d
Merge fixed ART main string subclass interning update
bradhilton Sep 28, 2026
46c9110
Compose native export correction and checkpoint test reuse
bradhilton Sep 28, 2026
9bfdc4e
Remove unused GPT-OSS interleaved padding helper
bradhilton Sep 28, 2026
b745db5
Compose checkpoint export synchronization and accepted private simpli…
bradhilton Sep 28, 2026
3777c0b
Share LoRA publication runtime validation
bradhilton Sep 28, 2026
0fe6f91
Require validated runtime for LoRA publication preparation
bradhilton Sep 28, 2026
e2852ca
Test second export metadata failure on rank one
bradhilton Sep 28, 2026
6222b3c
Merge ART main e65b0dd preserving literal assistant history
bradhilton Sep 28, 2026
8496d79
Qualify joined-reasoning test fixture import
bradhilton Sep 28, 2026
b4da074
Merge ART memory extraction while preserving trainer-v1 admission
bradhilton Sep 28, 2026
bcb3cc8
Remove obsolete thinking workaround and strengthen branch coverage
bradhilton Sep 28, 2026
e80192e
Merge ART planner and checkpoint extractions preserving trainer-v1
bradhilton Sep 28, 2026
74b5d9f
Merge exact ART tokenization update preserving trainer-v1 controls
bradhilton Sep 28, 2026
1a02ed7
Fix runtime annotations on extracted public trainer methods
bradhilton Sep 28, 2026
7ce1cc8
Remove stale imports and repeated extraction notes
bradhilton Sep 28, 2026
7c0d0b9
Preserve literal renderer oracles across explicit override modes
bradhilton Sep 28, 2026
b13b5f5
Respect prompt provenance in the explicit template override control
bradhilton Sep 28, 2026
8711cc5
Share recorded boundary invocation setup in tests
bradhilton Sep 28, 2026
ebf3ccf
Merge main event-loop isolation into trainer-v1
bradhilton Sep 28, 2026
d9e9061
Merge ART main 71af7792 planner input snapshots
bradhilton Sep 28, 2026
620a1f0
Merge main a68fa500 planner retention diagnostics
bradhilton Sep 28, 2026
b5e65a8
Merge current main planner diagnostic bounds into trainer-v1
bradhilton Sep 28, 2026
24bdb52
Release registered forward graphs when correction setup fails
bradhilton Sep 28, 2026
19db288
Merge current main planner diagnostic nesting bounds into trainer-v1
bradhilton Sep 28, 2026
35498de
Compose trainer handoff and callback cleanup fixes
bradhilton Sep 28, 2026
22cd501
Merge bounded grouped planner replay into trainer v1
bradhilton Sep 28, 2026
f002b10
Preserve placed report completeness and failed graph ownership
bradhilton Sep 28, 2026
d758c0a
Release failed native forward capture aliases
bradhilton Sep 29, 2026
4fc438c
Release failed correction capture copies
bradhilton Sep 29, 2026
43b65bb
Release failed forward captures before graph registration
bradhilton Sep 29, 2026
468e2b5
Restore reviewed ART backward-state cleanup from durable patch
bradhilton Sep 29, 2026
e506246
Merge commit 'eac090802533e998d05e3b681aa9401bf33780e3' into HEAD
bradhilton Sep 29, 2026
7a64114
Release saved graph storage on eviction and replay rejection
bradhilton Sep 29, 2026
496a14b
Merge commit 'e42cc67de99906800e63efc46646fa9e41c37254' into HEAD
bradhilton Sep 29, 2026
ec7f2c1
Release correction staging and backward aliases on failure
bradhilton Sep 29, 2026
e156bb7
Release successful preparation results on coordinated failure
bradhilton Sep 29, 2026
179b162
Merge ART main checkpoint snapshot capture at 4ced233c
bradhilton Sep 29, 2026
1f594e7
Simplify backward orchestration and distributed test coordination
bradhilton Sep 29, 2026
e84a857
Narrow filtered checkpoint slot modules for static typing
bradhilton Sep 29, 2026
2b980b2
Release undelivered executor results after peer failure
bradhilton Sep 29, 2026
b5765d1
Release undelivered output aliases across executor failures
bradhilton Sep 29, 2026
aa55abe
Release detached output references after copy failures
bradhilton Sep 29, 2026
3382d36
Merge commit '708c9fe4593290b3fda8c3031b909f8bd30ee554' into HEAD
bradhilton Sep 29, 2026
4530fe0
Bound asynchronous checkpoint captures by host headroom
bradhilton Sep 29, 2026
77658b1
Include persistent buffer gather in snapshot admission
bradhilton Sep 29, 2026
8029716
Narrow command test head types without changing test behavior
bradhilton Sep 29, 2026
9d02f3b
Release partial checkpoint copies from failed capture frames
bradhilton Sep 29, 2026
5c2a530
Match checkpoint capture frames by exact code identity
bradhilton Sep 29, 2026
577f507
Merge commit 'd3809140b20f2ace90b01201ff1ea77c69b79dce' into HEAD
bradhilton Sep 29, 2026
2e1281a
Reuse adapter config factory in checkpoint tests
bradhilton Sep 29, 2026
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
11 changes: 10 additions & 1 deletion .github/workflows/prek.yml
Original file line number Diff line number Diff line change
Expand Up @@ -213,7 +213,7 @@ jobs:
- name: Run Megatron lightweight tests
run: |
megatron_runtime/.venv/bin/python -c "import megatron.core.packed_seq_params"
megatron_runtime/.venv/bin/python -m pytest --nbval --current-env --tb=short \
megatron_runtime/.venv/bin/python -m pytest -v --nbval --current-env --tb=short \
tests/unit/test_megatron_reference_logprobs.py \
tests/unit/test_preprocessing_tokenize.py::test_gemma4_normalizes_json_tool_arguments_for_mapping_template \
tests/unit/test_moe_routing_replay.py \
Expand All @@ -224,10 +224,14 @@ jobs:
tests/unit/test_prefix_tree_grad_parity.py \
tests/unit/test_prefix_tree_packing.py \
tests/unit/test_trainer_rank_handoff_budget.py \
tests/unit/test_trainer_rank_graph_backward_work.py \
tests/unit/test_qwen35_adapter_config.py \
tests/unit/test_trainer_rank_physical_reserve.py \
tests/unit/test_trainer_rank_validation.py \
tests/unit/test_trainer_rank_weird_shapes.py \
tests/unit/test_trainer_rank_slot_graph_lifetime.py \
tests/unit/test_trainer_rank_forward_handoff.py \
tests/unit/test_trainer_rank_commands.py::test_gloo_dp2_tp2_participation_and_gradients \
tests/unit/test_trainer_rank_admission_inputs.py \
tests/unit/test_trainer_rank_checkpoint_memory.py \
tests/unit/test_trainer_rank_profile_warm.py \
Expand All @@ -249,6 +253,7 @@ jobs:
tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \
tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_refuses_after_recovery \
tests/acceptance/trainer_rank_planner \
tests/integration/megatron/lora/test_lora_disk_codecs.py::test_export_preparation_failure_is_collective_and_retryable \
tests/integration/megatron/test_sft_packing.py::test_sft_packing_preserves_training_targets \
tests/integration/megatron/model_support/test_dispatcher_graph_retention.py \
tests/integration/megatron/gdn_shared_prefix/test_gdn_planner_runtime_model.py \
Expand All @@ -264,6 +269,7 @@ jobs:
uv run --no-sync pytest --nbval --current-env --tb=short tests/unit \
--deselect=tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \
--deselect=tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_refuses_after_recovery \
--deselect=tests/unit/test_trainer_rank_commands.py::test_gloo_dp2_tp2_participation_and_gradients \
--ignore=tests/unit/test_megatron_reference_logprobs.py \
--ignore=tests/unit/test_moe_routing_replay.py \
--ignore=tests/unit/test_moe_routing_real_path.py \
Expand All @@ -273,10 +279,13 @@ jobs:
--ignore=tests/unit/test_prefix_tree_grad_parity.py \
--ignore=tests/unit/test_prefix_tree_packing.py \
--ignore=tests/unit/test_trainer_rank_handoff_budget.py \
--ignore=tests/unit/test_trainer_rank_graph_backward_work.py \
--ignore=tests/unit/test_qwen35_adapter_config.py \
--ignore=tests/unit/test_trainer_rank_physical_reserve.py \
--ignore=tests/unit/test_trainer_rank_validation.py \
--ignore=tests/unit/test_trainer_rank_weird_shapes.py \
--ignore=tests/unit/test_trainer_rank_slot_graph_lifetime.py \
--ignore=tests/unit/test_trainer_rank_forward_handoff.py \
--ignore=tests/unit/test_trainer_rank_admission_inputs.py \
--ignore=tests/unit/test_trainer_rank_checkpoint_memory.py \
--ignore=tests/unit/test_trainer_rank_profile_warm.py \
Expand Down
8 changes: 4 additions & 4 deletions dev/trainer_rank.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,17 +70,17 @@ def main(
for step in range(steps):
loss_sum = torch.tensor(0.0, device=rank.device)
token_count = torch.tensor(0.0, device=rank.device)
for micro in rank.forward_micro_batches(inputs, checkpoint=slot):
for micro in rank.forward_batches(inputs, checkpoint=slot):
loss = torch.tensor(0.0, device=rank.device)
for output in micro.outputs:
assert output.target_logprobs is not None
loss = loss - output.target_logprobs.sum()
token_count += output.target_logprobs.numel()
loss.backward()
rank.backward(loss)
loss_sum += loss.detach()

rank.dp_reduce(loss_sum)
rank.dp_reduce(token_count)
rank.reduce(loss_sum)
rank.reduce(token_count)
scale = 1.0 / max(float(token_count.item()), 1.0)
metrics = rank.optim_step(
params=AdamParams(learning_rate=lr),
Expand Down
16 changes: 9 additions & 7 deletions dev/trainer_rank_check.py
Original file line number Diff line number Diff line change
Expand Up @@ -226,7 +226,7 @@ def _local_outputs(
rank: TrainerRank,
indexed_requests: Sequence[tuple[int, ForwardInput]],
) -> list[dict[str, object]]:
outputs = rank.dp_rank_forward([request for _, request in indexed_requests])
outputs = rank.forward([request for _, request in indexed_requests])
return [
_output_record(index, torch.arange(request.input_tokens.numel()), output)
for (index, request), output in zip(indexed_requests, outputs, strict=True)
Expand Down Expand Up @@ -502,7 +502,7 @@ def _head_backward_chunk_parity(
outputs = rank._project_head(items, prepared, candidate)
loss = _output_loss(outputs)
logprob_sums.append(loss.detach())
loss.backward()
rank.backward(loss)
assert candidate.grad is not None
gradients.append(candidate.grad)
finally:
Expand Down Expand Up @@ -584,7 +584,7 @@ def _slot_gradients(
slots: Sequence[str],
) -> dict[str, list[torch.Tensor]]:
rank.zero_grad()
_output_loss(rank.dp_rank_forward(requests)).backward()
rank.backward(_output_loss(rank.forward(requests)))
return {
slot: [
torch.zeros_like(parameter, dtype=torch.float32, device="cpu")
Expand Down Expand Up @@ -665,14 +665,16 @@ def step() -> list[MicroBatchStats]:
rank.zero_grad()
stats: list[MicroBatchStats] = []
if adaptive:
for micro in rank.forward_micro_batches(requests):
_output_loss(cast(Sequence[ForwardOutput], micro.outputs)).backward()
for micro in rank.forward_batches(requests):
rank.backward(
_output_loss(cast(Sequence[ForwardOutput], micro.outputs))
)
stats.append(micro.stats)
else:
outputs = rank.dp_rank_forward(requests[dp_rank::dp_size])
outputs = rank.forward(requests[dp_rank::dp_size])
if workload == "unequal_slots":
_trace_unequal_slots("forward_ready")
_output_loss(outputs).backward()
rank.backward(_output_loss(outputs))
if workload == "unequal_slots":
_trace_unequal_slots("backward_ready")
if optimizer_step:
Expand Down
4 changes: 2 additions & 2 deletions dev/trainer_rank_collective_trace.py
Original file line number Diff line number Diff line change
Expand Up @@ -141,7 +141,7 @@ def logged(*args, **kwargs):
# gdn.layout imported the function by name; rebind it to the logged one.
gdn_layout.all_to_all_single = dist.all_to_all_single

original_forward = TrainerRank.dp_rank_forward
original_forward = TrainerRank.forward

@functools.wraps(original_forward)
def logged_forward(self, *args, **kwargs):
Expand All @@ -166,7 +166,7 @@ def logged_forward(self, *args, **kwargs):
)
return outputs

TrainerRank.dp_rank_forward = logged_forward # type: ignore[method-assign]
TrainerRank.forward = logged_forward # type: ignore[method-assign]


def _start_watchdog(logger: _Logger, log_dir: Path, stall_seconds: float) -> None:
Expand Down
79 changes: 52 additions & 27 deletions dev/trainer_rank_landing_acceptance.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,17 +137,35 @@ def phase_contract() -> None:

problems: list[str] = []
constructor = _public_parameters(trainer_rank.TrainerRank.__init__)
if list(constructor) != ["runtime"]:
if (
list(constructor) != ["runtime", "options"]
or constructor["options"].kind is not inspect.Parameter.KEYWORD_ONLY
or constructor["options"].default is not None
):
problems.append(
f"TrainerRank must accept exactly (runtime); found {sorted(constructor)}"
f"TrainerRank must accept (runtime, *, options=None); found {constructor}"
)
for method_name in ("forward_micro_batches", "dp_rank_forward"):
for method_name in ("forward_batches", "forward"):
parameters = _public_parameters(getattr(trainer_rank.TrainerRank, method_name))
# ``yield_empty`` (PR #864) is a keyword-only flag defaulting to False:
# off, the contract the acceptance suite pins is unchanged.
extra = set(parameters) - {"inputs", "checkpoint", "no_grad", "yield_empty"}
extra = set(parameters) - {
"inputs",
"checkpoint",
"no_grad",
"yield_empty",
"options",
}
if extra:
problems.append(f"{method_name} has extra parameters {sorted(extra)}")
options = parameters.get("options")
if options is None or (
options.kind is not inspect.Parameter.KEYWORD_ONLY
or options.default is not None
):
problems.append(
f"{method_name}: options must be keyword-only and default to None"
)
flag = parameters.get("yield_empty")
if flag is not None and (
flag.kind is not inspect.Parameter.KEYWORD_ONLY or flag.default is not False
Expand Down Expand Up @@ -374,8 +392,8 @@ def phase_measure(cell: str, arm: str, output_jsonl: str, repeat: int) -> None:

# Behavior smoke from the contract: empty inputs are valid zero-work
# calls that must not disturb subsequent planning.
empty = rank.dp_rank_forward([])
assert len(list(empty)) == 0, "dp_rank_forward([]) must return no outputs"
empty = rank.forward([])
assert len(list(empty)) == 0, "forward([]) must return no outputs"

rows: list[dict[str, object]] = []
for sample in range(repeat + 4): # 1 cold + 3 warmup + repeat measured
Expand All @@ -387,8 +405,8 @@ def phase_measure(cell: str, arm: str, output_jsonl: str, repeat: int) -> None:
start.record()
admission_failed = False
try:
outputs = rank.dp_rank_forward(requests)
_output_loss(outputs).backward()
outputs = rank.forward(requests)
rank.backward(_output_loss(outputs))
except TrainerRankMemoryError as error:
admission_failed = True
rows.append(
Expand Down Expand Up @@ -629,7 +647,7 @@ def forward(
) -> tuple[list[torch.Tensor], dict[str, object]]:
torch.cuda.synchronize()
torch.cuda.reset_peak_memory_stats()
outputs = rank.dp_rank_forward(requests(), no_grad=no_grad)
outputs = rank.forward(requests(), no_grad=no_grad)
telemetry = rank.last_forward_telemetry()
logprobs = [
output.target_logprobs.detach().float().clone() for output in outputs
Expand All @@ -638,9 +656,11 @@ def forward(

def combined_backward(info: dict[str, object]) -> None:
outputs = cast(list, info["outputs"])
torch.stack(
[output.target_logprobs.float().sum() for output in outputs]
).sum().backward()
rank.backward(
torch.stack(
[output.target_logprobs.float().sum() for output in outputs]
).sum()
)
torch.cuda.synchronize()

def unsplit_requirement() -> int:
Expand Down Expand Up @@ -684,7 +704,7 @@ def expect_decline(arm: str) -> None:
torch.cuda.synchronize()
before = int(torch.cuda.memory_allocated())
try:
rank.dp_rank_forward(requests())
rank.forward(requests())
except TrainerRankMemoryError as error:
message = str(error).lower()
if "unable to find a feasible split" not in message:
Expand Down Expand Up @@ -773,9 +793,14 @@ def pressured_cap() -> int:
if len(partition) < 2:
_fail("reverse-order arm expected a split plan")
for indices in reversed(partition):
torch.stack(
[outputs[index].target_logprobs.float().sum() for index in indices]
).sum().backward()
rank.backward(
torch.stack(
[
outputs[index].target_logprobs.float().sum()
for index in indices
]
).sum()
)
rank.zero_grad()
del info, outputs

Expand Down Expand Up @@ -972,7 +997,7 @@ def conversion(
# backward through an active LoRA slot, compared against the depth-one arm.
# ``dp2-tp2-waves`` (4 ranks, DP2 x TP2) exercises the global wave planner's
# collectives (world scope) composed with TP execution collectives (pair scope)
# through public ``forward_micro_batches``, including an empty DP slot.
# through public ``forward_batches``, including an empty DP slot.

TP_GATES: dict[str, float] = {
# bf16 kernels reorder reductions across packings (same metric/tolerance as
Expand Down Expand Up @@ -1337,7 +1362,7 @@ def _write_rows(evidence: str | None, rows: list[dict[str, object]], name: str)
def phase_tp2_public(
evidence: str | None, repeat: int, *, tp: int, dump_dir: str | None
) -> None:
"""DP1 x TP{tp} x CP1 public ``dp_rank_forward`` cell (Qwen3.5-4B full model).
"""DP1 x TP{tp} x CP1 public ``forward`` cell (Qwen3.5-4B full model).

Run at ``--tp 2`` (the gate) and at ``--tp 1`` (the control: the identical
cell on one GPU). Structural gates run here; the numerics gates compare
Expand Down Expand Up @@ -1407,9 +1432,9 @@ def run(arm: str) -> dict[str, Any]:
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
outputs = rank.dp_rank_forward(requests, checkpoint=slot)
outputs = rank.forward(requests, checkpoint=slot)
loss = _output_loss(outputs)
loss.backward()
rank.backward(loss)
end.record()
torch.cuda.synchronize()
telemetry = rank.last_forward_telemetry()
Expand Down Expand Up @@ -1759,7 +1784,7 @@ def structured(label: str, profile: dict[str, float]) -> None:


def phase_dp2_tp2_waves(evidence: str | None) -> None:
"""DP2 x TP2 public ``forward_micro_batches`` gate (4 ranks, Qwen3.5-4B).
"""DP2 x TP2 public ``forward_batches`` gate (4 ranks, Qwen3.5-4B).

Arm A (branchy): six hierarchical GRPO groups as top-level items, the
test-only memory cap sized so the stream needs at least two waves, forward
Expand Down Expand Up @@ -1816,7 +1841,7 @@ def stream(arm: str, cap_bytes: int | None) -> dict[str, Any]:
logprobs: dict[int, list[torch.Tensor]] = {}
loss_total = 0.0
depths: list[int] = []
for batch in rank.forward_micro_batches(items, checkpoint=slot):
for batch in rank.forward_batches(items, checkpoint=slot):
seen.extend(int(index) for index in batch.indices)
telemetry = rank.last_forward_telemetry()
depths.append(int(telemetry["selected_max_depth"]))
Expand All @@ -1840,7 +1865,7 @@ def stream(arm: str, cap_bytes: int | None) -> dict[str, Any]:
if flat:
loss = _output_loss(flat)
loss_total += float(loss.detach().float().item())
loss.backward()
rank.backward(loss)
torch.cuda.synchronize()
result = {
"arm": arm,
Expand Down Expand Up @@ -1955,12 +1980,12 @@ def stream(arm: str, cap_bytes: int | None) -> dict[str, Any]:
rank.zero_grad()
waves = 0
local_outputs = 0
for batch in rank.forward_micro_batches([single], checkpoint=slot):
for batch in rank.forward_batches([single], checkpoint=slot):
waves += 1
flat = [output for group in batch.outputs for output in group]
local_outputs += len(flat)
if flat:
_output_loss(flat).backward()
rank.backward(_output_loss(flat))
torch.cuda.synchronize()
counts = _gather_objects((dp_rank, waves, local_outputs))
rows.append({"arm": "dp2-tp2-empty-slot", "per_rank": counts})
Expand Down Expand Up @@ -2625,9 +2650,9 @@ def run(label: str) -> dict[str, object]:
start.record()
failed = 0
try:
outputs = rank.dp_rank_forward(requests, checkpoint=slot)
outputs = rank.forward(requests, checkpoint=slot)
loss = _output_loss(outputs)
loss.backward()
rank.backward(loss)
except TrainerRankMemoryError as error:
failed = 1
message = str(error)
Expand Down
Loading
Loading