From 9411a658bca43e0a0864f8dd0ef8afb7c560790d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 16 Sep 2026 21:16:39 +0000 Subject: [PATCH 01/15] Model retained activations for selective recompute admission --- src/art/trainer_rank/_impl.py | 43 +++- .../test_trainer_rank_recompute_memory.py | 187 ++++++++++++++++++ tests/unit/test_trainer_rank_split.py | 4 +- tests/unit/test_trainer_rank_split_peak.py | 4 +- tests/unit/test_trainer_rank_topology.py | 9 +- 5 files changed, 229 insertions(+), 18 deletions(-) create mode 100644 tests/unit/test_trainer_rank_recompute_memory.py diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 8a0e9f038..6921a5f83 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1357,14 +1357,6 @@ def __init__(self, runtime: TrainingRuntime) -> None: "therefore requires PP=1 with exactly one local model chunk; " f"got pp={pp_size}, chunks={len(runtime.model)}" ) - if getattr(runtime.provider, "recompute_granularity", None) == "selective": - raise TrainerRankRuntimeSupportError( - "TrainerRank memory planning does not support selective recompute; " - "its activation estimate assumes full recompute. Use " - "ART_MEGATRON_RECOMPUTE_GRANULARITY=full with " - "ART_MEGATRON_RECOMPUTE_METHOD=uniform and " - "ART_MEGATRON_RECOMPUTE_NUM_LAYERS=1." - ) # Tensor parallelism is admitted: the vocab-parallel head, sequence- # parallel gather, TP padding of packed batches and sharded LoRA # gradient reduction pre-date the planner, memory checks all-reduce @@ -1388,6 +1380,11 @@ def __init__(self, runtime: TrainingRuntime) -> None: or getattr(runtime.provider, "num_layers", 1) or 1 ) + self._recompute_granularity = getattr( + getattr(metadata_model, "config", None), + "recompute_granularity", + getattr(runtime.provider, "recompute_granularity", None), + ) # Layers that run the gated-delta-net path (Qwen3.5-4B: 24 of 32); the # cost model prices GDN state hand-offs per GDN layer, not per layer. self._gdn_layers = _gdn_layer_count(runtime.model[0]) @@ -5134,6 +5131,36 @@ def _estimate_required_memory_bytes_from_values( * self._param_dtype_size * activation_factor ) + if signature.grad_enabled and self._recompute_granularity != "full": + geometry = self._geometry + ffn_width = max( + geometry.ffn_hidden_size or 4 * self._hidden_size, + geometry.moe_topk * geometry.moe_ffn_hidden_size + + geometry.moe_shared_expert_ffn, + ) + attention_width = max( + self._hidden_size, + geometry.num_attention_heads * geometry.kv_channels, + 2 * geometry.gdn_key_heads * geometry.gdn_key_head_dim + + 2 * geometry.gdn_value_heads * geometry.gdn_value_head_dim, + ) + # Megatron's selective-recompute model retains per-layer attention, + # norm and MLP tensors (arxiv.org/abs/2205.05198). Charge four FFN + # widths for gated MLPs, and routed hidden rows for MoE. Do not take + # TP/CP/EP or optional selective-module discounts without measured + # evidence; the default core_attn checkpoint leaves the MLP live. + layer_features = ( + 9 * attention_width + + 4 * ffn_width + + 2 * self._hidden_size * max(0, geometry.moe_topk - 1) + ) + static_compute = max( + static_compute, + packed_tokens + * self._num_layers + * self._param_dtype_size + * layer_features, + ) # Groups execute sequentially: summed packed rows conservatively bound # this FC2 component, not all workspace or retained graphs. static_compute = max( diff --git a/tests/unit/test_trainer_rank_recompute_memory.py b/tests/unit/test_trainer_rank_recompute_memory.py new file mode 100644 index 000000000..2d8b70f03 --- /dev/null +++ b/tests/unit/test_trainer_rank_recompute_memory.py @@ -0,0 +1,187 @@ +"""Cold admission checks; these CPU estimates are not native GPU measurements.""" + +from dataclasses import replace +from types import SimpleNamespace +from typing import Any, cast + +import pytest +import torch +from torch.utils.checkpoint import checkpoint + +from art.trainer_rank import ForwardInput, TrainerRank, TrainerRankMemoryError +from art.trainer_rank._impl import _MemoryProfile + + +def _rank( + granularity: str | None = "full", + *, + dtype: torch.dtype = torch.bfloat16, + **geometry: Any, +) -> TrainerRank: + provider = SimpleNamespace( + **{ + "hidden_size": 5120, + "ffn_hidden_size": 17408, + "num_layers": 64, + "num_attention_heads": 24, + "num_query_groups": 4, + "kv_channels": 256, + "recompute_granularity": granularity, + **geometry, + } + ) + return TrainerRank( + cast( + Any, + SimpleNamespace( + model=[torch.nn.Linear(1, 1, dtype=dtype)], + provider=provider, + optimizer=None, + model_support_handler=SimpleNamespace( + build_gdn_execution_spec=bool( + geometry.get("linear_num_value_heads") + ), + is_moe=bool(geometry.get("num_moe_experts")), + ), + ), + ) + ) + + +def _plan(rank: TrainerRank, tokens: int = 32710, *, no_grad: bool = False): + request = ForwardInput( + input_tokens=torch.tensor([1, 2]), hidden_states=True, no_grad=no_grad + ) + return replace( + rank._plan_flat_forward([request]), + packed_tokens=tokens, + logical_tokens=tokens, + output_bytes=tokens * rank.hidden_size * rank._param_dtype_size, + ) + + +@pytest.mark.parametrize("tp", (2, 4)) +@pytest.mark.parametrize("granularity", (None, "selective")) +def test_reported_cold_request_is_refused_before_execution( + monkeypatch: pytest.MonkeyPatch, tp: int, granularity: str | None +) -> None: + rank = _rank(granularity) + monkeypatch.setattr(rank, "_topology_key", lambda: (1, tp, 1, 1)) + monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (0, 1)) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: int(119.289e9)) + monkeypatch.setattr( + rank, "_execute_flat_plan", lambda _: pytest.fail("unsafe forward admitted") + ) + assert not rank._memory_check(_plan(rank)).fits + assert rank._memory_check(_plan(rank, tokens=1024)).fits + with pytest.raises(TrainerRankMemoryError): + rank.dp_rank_forward( + [ForwardInput(input_tokens=torch.arange(32710), hidden_states=True)] + ) + + +def test_full_recompute_keeps_existing_estimate() -> None: + rank = _rank() + plan = _plan(rank) + assert rank._memory_check(plan).estimated_required_bytes == int( + (plan.output_bytes + 32710 * 5120 * 2 * 16) * 1.1 + ) + + +@pytest.mark.parametrize("granularity", (None, "selective")) +def test_no_grad_does_not_pay_for_retained_layers(granularity: str | None) -> None: + rank, full = _rank(granularity), _rank() + assert rank._memory_check(_plan(rank, no_grad=True)) == full._memory_check( + _plan(full, no_grad=True) + ) + assert rank._memory_check(_plan(rank)).estimated_required_bytes > ( + rank._memory_check(_plan(rank, no_grad=True)).estimated_required_bytes + ) + + +@pytest.mark.parametrize( + "geometry", + [ + {"num_layers": 128}, + {"ffn_hidden_size": 34816}, + {"num_attention_heads": 48}, + {"ffn_hidden_size": 0}, + { + "linear_num_key_heads": 16, + "linear_key_head_dim": 128, + "linear_num_value_heads": 48, + "linear_value_head_dim": 128, + }, + {"num_moe_experts": 64, "moe_router_topk": 8, "moe_ffn_hidden_size": 8192}, + ], +) +def test_retained_estimate_tracks_model_geometry(geometry: dict[str, Any]) -> None: + rank, base = _rank("selective", **geometry), _rank("selective") + assert rank._memory_check(_plan(rank)).estimated_required_bytes > ( + base._memory_check(_plan(base)).estimated_required_bytes + ) + + +def test_profile_cannot_erase_recompute_floor() -> None: + rank = _rank("selective") + plan = _plan(rank) + cold = rank._memory_check(plan).estimated_required_bytes + rank._memory_profiles[plan.signature] = _MemoryProfile(0.0, plan.packed_tokens) + assert rank._memory_check(plan).estimated_required_bytes == cold + rank._memory_profiles[plan.signature] = _MemoryProfile( + 2 * cold / plan.packed_tokens, plan.packed_tokens + ) + assert rank._memory_check(plan).estimated_required_bytes > cold + + +def test_non_full_estimate_covers_retained_gated_mlp_tensors() -> None: + # Selective core-attention recompute leaves this MLP graph live. Count + # distinct saved activation storage, excluding model parameters/views. + layers = [ + torch.nn.Sequential( + torch.nn.LayerNorm(16), + torch.nn.Linear(16, 256), + torch.nn.GLU(), + torch.nn.Linear(128, 16), + ) + for _ in range(8) + ] + parameters = { + parameter.untyped_storage().data_ptr() + for layer in layers + for parameter in layer.parameters() + } + + def retained(full: bool) -> int: + saved: dict[int, int] = {} + + def pack(tensor: torch.Tensor) -> torch.Tensor: + storage = tensor.untyped_storage() + if storage.data_ptr() not in parameters: + saved[storage.data_ptr()] = storage.nbytes() + return tensor + + with torch.autograd.graph.saved_tensors_hooks(pack, lambda tensor: tensor): + value = torch.zeros(32, 16, requires_grad=True) + for layer in layers: + value = ( + checkpoint(layer, value, use_reentrant=False) + if full + else layer(value) + ) + return sum(saved.values()) + + rank = _rank( + "selective", + dtype=torch.float32, + hidden_size=16, + ffn_hidden_size=128, + num_layers=8, + num_attention_heads=2, + kv_channels=8, + ) + assert ( + retained(True) + < retained(False) + <= rank._memory_check(_plan(rank, tokens=32)).estimated_required_bytes + ) diff --git a/tests/unit/test_trainer_rank_split.py b/tests/unit/test_trainer_rank_split.py index ce2ad6cdb..7b8bb812f 100644 --- a/tests/unit/test_trainer_rank_split.py +++ b/tests/unit/test_trainer_rank_split.py @@ -83,7 +83,9 @@ def _runtime() -> "TrainingRuntime": return SimpleNamespace( model=[_FakeGPT()], optimizer=None, - provider=SimpleNamespace(hidden_size=8, num_layers=4), + provider=SimpleNamespace( + hidden_size=8, num_layers=4, recompute_granularity="full" + ), model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), ) # type: ignore diff --git a/tests/unit/test_trainer_rank_split_peak.py b/tests/unit/test_trainer_rank_split_peak.py index 7eb6cab04..54473e12e 100644 --- a/tests/unit/test_trainer_rank_split_peak.py +++ b/tests/unit/test_trainer_rank_split_peak.py @@ -31,7 +31,9 @@ def _rank(): SimpleNamespace( model=[_Model()], optimizer=None, - provider=SimpleNamespace(hidden_size=8, num_layers=4), + provider=SimpleNamespace( + hidden_size=8, num_layers=4, recompute_granularity="full" + ), model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), ), ) diff --git a/tests/unit/test_trainer_rank_topology.py b/tests/unit/test_trainer_rank_topology.py index e706b36a7..45ec291dd 100644 --- a/tests/unit/test_trainer_rank_topology.py +++ b/tests/unit/test_trainer_rank_topology.py @@ -65,14 +65,7 @@ def test_trainer_rank_recompute_support( ) -> None: runtime = _runtime(tp=tp) runtime.provider.recompute_granularity = granularity - if granularity == "selective": - with pytest.raises( - TrainerRankRuntimeSupportError, - match="selective recompute.*ART_MEGATRON_RECOMPUTE_GRANULARITY=full", - ): - TrainerRank(runtime) - else: - TrainerRank(runtime) + TrainerRank(runtime) def test_trainer_rank_still_refuses_pipeline_parallel_runtimes() -> None: From 8dece8ed32e0d1af747e35de25c012306086e90c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 16 Sep 2026 21:39:08 +0000 Subject: [PATCH 02/15] Calibrate recompute memory estimates on H200 --- dev/trainer_rank_recompute_memory.csv | 55 +++++++ dev/trainer_rank_recompute_memory.md | 106 ++++++++++++ dev/trainer_rank_recompute_memory.py | 181 +++++++++++++++++++++ dev/trainer_rank_recompute_memory.sky.yaml | 51 ++++++ 4 files changed, 393 insertions(+) create mode 100644 dev/trainer_rank_recompute_memory.csv create mode 100644 dev/trainer_rank_recompute_memory.md create mode 100644 dev/trainer_rank_recompute_memory.py create mode 100644 dev/trainer_rank_recompute_memory.sky.yaml diff --git a/dev/trainer_rank_recompute_memory.csv b/dev/trainer_rank_recompute_memory.csv new file mode 100644 index 000000000..cadbf61e3 --- /dev/null +++ b/dev/trainer_rank_recompute_memory.csv @@ -0,0 +1,55 @@ +model,tp,mode,logical_tokens,packed_tokens,estimate_bytes,available_min_bytes,measured_rank_samples,refused_rank_samples,forward_peak_max_bytes,retained_max_bytes,forward_backward_peak_max_bytes,baseline_min_bytes,baseline_max_bytes,finite_loss_all,estimate_covers_forward_all +Qwen/Qwen3-1.7B,2,full,1024,1024,55364812,140941082317,4,0,177251840,133205504,231861248,1781843968,1883032576,True,False +Qwen/Qwen3-1.7B,2,full,4096,4096,221459251,140938985165,4,0,415402496,264382976,587698176,1883032576,1885129728,True,False +Qwen/Qwen3-1.7B,2,full,16384,16384,885837004,140930596557,4,0,1661798912,1057720832,2099318784,1885129728,1893518336,True,False +Qwen/Qwen3-1.7B,2,none,1024,1024,2717489561,140941082317,4,0,1440448000,1427860480,1473997312,1781843968,1883032576,True,True +Qwen/Qwen3-1.7B,2,none,4096,4096,10869958246,140938985165,4,0,5493352960,5443004416,5610778112,1883032576,1885129728,True,True +Qwen/Qwen3-1.7B,2,none,16384,16384,43479832985,140930596557,4,0,21973600768,21772208128,22443298304,1885129728,1893518336,True,True +Qwen/Qwen3-1.7B,2,selective,1024,1024,2717489561,140941082317,4,0,1439530496,1426943488,1535994880,1781843968,1883032576,True,True +Qwen/Qwen3-1.7B,2,selective,4096,4096,10869958246,140938985165,4,0,5489682944,5439334912,5607108608,1883032576,1885129728,True,True +Qwen/Qwen3-1.7B,2,selective,16384,16384,43479832985,140930596557,4,0,21958920704,21757528576,22428618752,1885129728,1893518336,True,True +Qwen/Qwen3.8-27B,1,full,1024,1024,196083712,89386251981,2,0,973403648,759490048,1198017536,54352220160,54452884480,True,False +Qwen/Qwen3.8-27B,1,full,2048,2048,392167424,89384154829,2,0,1812587008,1384759808,2092073472,54452884480,54452884480,True,False +Qwen/Qwen3.8-27B,1,full,3072,3072,588251136,89384154829,2,0,2718884352,2077143552,3134828032,54452884480,54452884480,True,False +Qwen/Qwen3.8-27B,1,full,4096,4096,784334848,89384154829,2,0,3625185792,2769531392,4178635264,54452884480,54452884480,True,False +Qwen/Qwen3.8-27B,1,full,38443,32710,6328148992,89321231565,2,0,28953189376,22178822144,33419714048,54452884480,54452893184,True,False +Qwen/Qwen3.8-27B,1,none,1024,1024,31311108505,89386251981,2,0,17041559040,16955570688,17060460544,54352220160,54452884480,True,True +Qwen/Qwen3.8-27B,1,none,2048,2048,62622217011,89386251981,2,0,33761374720,33589400064,33799116800,54452884480,54452884480,True,True +Qwen/Qwen3.8-27B,1,none,3072,3072,93933325516,89386251981,0,1,,,,,,, +Qwen/Qwen3.8-27B,1,none,4096,4096,125244434022,89386251981,0,1,,,,,,, +Qwen/Qwen3.8-27B,1,none,38443,32710,1000246567936,89386251981,0,1,,,,,,, +Qwen/Qwen3.8-27B,1,selective,1024,1024,31311108505,87997937357,2,0,17043131904,16957146624,17124951040,54352220160,54452884480,True,True +Qwen/Qwen3.8-27B,1,selective,2048,2048,62622217011,87997937357,2,0,33758228992,33586258432,33795975168,54452884480,54452884480,True,True +Qwen/Qwen3.8-27B,1,selective,3072,3072,93933325516,87997937357,0,1,,,,,,, +Qwen/Qwen3.8-27B,1,selective,4096,4096,125244434022,87997937357,0,1,,,,,,, +Qwen/Qwen3.8-27B,1,selective,38443,32710,1000246567936,87997937357,0,1,,,,,,, +Qwen/Qwen3.8-27B,2,full,1024,1024,196083712,115395819213,4,0,526710272,419749376,710319104,27327631360,27428295680,True,False +Qwen/Qwen3.8-27B,2,full,2048,2048,392167424,115393722061,4,0,917103104,703181312,1098500608,27428295680,27428295680,True,False +Qwen/Qwen3.8-27B,2,full,3072,3072,588251136,115393722061,4,0,1375658496,1054775808,1646584832,27428295680,27428295680,True,False +Qwen/Qwen3.8-27B,2,full,4096,4096,784334848,115393722061,4,0,1834217984,1406374400,2194869760,27428295680,27428295680,True,False +Qwen/Qwen3.8-27B,2,full,38443,32710,6328148992,115391616205,4,0,14653886976,11294938624,17587685376,27428295680,27428304384,True,False +Qwen/Qwen3.8-27B,2,none,1024,1024,31311108505,115395819213,4,0,10079153664,10041397760,10146256896,27327631360,27428295680,True,True +Qwen/Qwen3.8-27B,2,none,2048,2048,62622217011,115395819213,4,0,19542929920,19467420160,19677136896,27428295680,27428295680,True,True +Qwen/Qwen3.8-27B,2,none,3072,3072,93933325516,115395819213,4,0,29394684416,29281420800,29595995136,27428295680,27428295680,True,True +Qwen/Qwen3.8-27B,2,none,4096,4096,125244434022,115395819213,0,2,,,,,,, +Qwen/Qwen3.8-27B,2,none,38443,32710,1000246567936,115395819213,0,2,,,,,,, +Qwen/Qwen3.8-27B,2,selective,1024,1024,31311108505,115395819213,4,0,10082561536,10044808704,10149667840,27327631360,27428295680,True,True +Qwen/Qwen3.8-27B,2,selective,2048,2048,62622217011,115395819213,4,0,19541553664,19466048000,19675764736,27428295680,27428295680,True,True +Qwen/Qwen3.8-27B,2,selective,3072,3072,93933325516,115393722061,4,0,29392325120,29279066624,29593640960,27428295680,27428295680,True,True +Qwen/Qwen3.8-27B,2,selective,4096,4096,125244434022,115393722061,0,2,,,,,,, +Qwen/Qwen3.8-27B,2,selective,38443,32710,1000246567936,115393722061,0,2,,,,,,, +Qwen/Qwen3.8-27B,4,full,1024,1024,196083712,128855946957,8,0,308082176,254597632,455984128,13865406464,13966070784,True,False +Qwen/Qwen3.8-27B,4,full,2048,2048,392167424,128855946957,8,0,480371200,373402112,591982080,13966070784,13966070784,True,False +Qwen/Qwen3.8-27B,4,full,3072,3072,588251136,128853849805,8,0,708239872,545689088,873627648,13966070784,13966070784,True,False +Qwen/Qwen3.8-27B,4,full,4096,4096,784334848,128853849805,8,0,939258368,725320192,1161568256,13966070784,13966070784,True,False +Qwen/Qwen3.8-27B,4,full,38443,32712,6328509440,128851743949,8,0,7502990336,5851191296,9831958528,13966070784,13966079488,True,False +Qwen/Qwen3.8-27B,4,none,1024,1024,31311108505,128855946957,8,0,5736394240,5722230272,5827089408,13865406464,13966070784,True,True +Qwen/Qwen3.8-27B,4,none,2048,2048,62622217011,128855946957,8,0,11131212288,11103934976,11313651712,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,none,3072,3072,93933325516,128855946957,8,0,16653462016,16610974208,16925548544,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,none,4096,4096,125244434022,128855946957,8,0,21989044736,21934492160,22353924096,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,none,38443,32712,1000307699916,128855946957,0,4,,,,,,, +Qwen/Qwen3.8-27B,4,selective,1024,1024,31311108505,127666861773,8,0,5732331008,5718170112,5932531200,13865406464,13966070784,True,True +Qwen/Qwen3.8-27B,4,selective,2048,2048,62622217011,127666861773,8,0,11132523008,11105249792,11314966528,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,selective,3072,3072,93933325516,127666861773,8,0,16651233792,16608751104,16923325440,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,selective,4096,4096,125244434022,127666861773,8,0,21987471872,21932925440,22352357376,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,selective,38443,32712,1000307699916,127666861773,0,4,,,,,,, diff --git a/dev/trainer_rank_recompute_memory.md b/dev/trainer_rank_recompute_memory.md new file mode 100644 index 000000000..a9dfe2568 --- /dev/null +++ b/dev/trainer_rank_recompute_memory.md @@ -0,0 +1,106 @@ +# Recompute memory calibration, 2026-09-16 + +First GPU campaign for #915, measuring estimator commit +`22f628d6a8e87d18b28eb9d1d61579520d012d23`. The new selective/no-recompute +estimate covered every executed forward. It remains conservative, particularly +at larger TP sizes. The unchanged full-recompute estimate underestimated every +measured shape; these results do not validate that legacy estimate. + +[CSV evidence](trainer_rank_recompute_memory.csv) contains 54 model/topology/mode/ +shape cells. Each row aggregates both repetitions and all ranks: peaks are maxima, +available memory is the minimum, and refused cells have blank measurement fields. +There were 202 measured rank-samples and 22 refused rank-samples. All measured +losses were finite, and every backward completed without CUDA OOM. + +## Method + +- NVIDIA H200, PyTorch 2.11.0+cu128, CUDA 12.8, bf16, compiled transformer layers. +- Native Qwen3.8-27B (64 layers) at TP1/2/4, and attention-only Qwen3-1.7B at TP2. + DP/CP/PP are 1; active LoRA rank is 1. Models and adapters use random weights. +- Each mode runs in a fresh process: `full/uniform/1`, selective `core_attn`, or + no recompute. Every sample clears the learned memory profile. The first sample + includes any compilation/autotuning required by that shape; the second is warm. +- Requests retain hidden states with gradients. The driver executes the selected + unsplit native plan only when its cold admission check passes, then backpropagates + a mean-square hidden-state loss. No admission bypass, split, or optimizer step. +- Estimates and measured peaks below are **incremental allocated GiB above the + pre-forward baseline**, not total GPU usage. Forward/backward peaks are separate + CSV fields; backward includes the diagnostic loss. Allocator reservation is not + the measured allocation peak. +- GDN lengths: 1,024 / 2,048 / 3,072 / 4,096, plus two synthetic sibling sequences + of 19,221 and 19,222 tokens sharing 5,733 prefix tokens. The siblings reproduce + #913's 38,443 logical / 32,710 packed tokens (32,712 after TP4 padding), not its + original token contents. Attention-only lengths: 1,024 / 4,096 / 16,384. +- TP2 ran locally. SkyPilot jobs `art-915-recompute-tp1-0916:1` and + `art-915-recompute-tp4-0916:1` both succeeded on free `k8s/cks-wb3` capacity. + Both clusters had ten-minute autodown and two-hour pod deadlines, and were + explicitly torn down after evidence retrieval. + +## Results + +Qwen3.8-27B, 2,048 tokens: + +| TP | New estimate | Selective forward peak | No-recompute forward peak | +| -- | --: | --: | --: | +| 1 | 58.321 | 31.440 | 31.443 | +| 2 | 58.321 | 18.199 | 18.201 | +| 4 | 58.321 | 10.368 | 10.367 | + +The long sibling group was refused in both non-full modes at all three TP sizes. +Its estimate is 931.552 GiB at TP1/2 and 931.609 GiB at TP4. These are admission +results, not measured non-full peaks. Smaller false refusals remain possible: +TP1 declined 3,072 and 4,096 tokens; TP2 declined 4,096 tokens. + +Full recompute completed the long sibling group: + +| TP | Legacy estimate | Forward peak | +| -- | --: | --: | +| 1 | 5.894 | 26.965 | +| 2 | 5.894 | 13.647 | +| 4 | 5.894 | 6.988 | + +Qwen3-1.7B, TP2: + +| Tokens | New estimate | Selective forward peak | No-recompute forward peak | +| --: | --: | --: | --: | +| 1,024 | 2.531 | 1.341 | 1.342 | +| 4,096 | 10.123 | 5.113 | 5.116 | +| 16,384 | 40.494 | 20.451 | 20.465 | + +## Remaining work + +Keep #915 in draft. This campaign supports the retained-activation floor for the +tested dense attention/GDN shapes; it is not a general memory bound. A universal +division by TP would already underestimate the attention-only cold 1,024-token +case (2.531 / 2 < 1.341 GiB). A tighter model must separate replicated activations +and fixed workspace from sharded storage, then validate on held-out shapes. + +Full recompute needs a separate correction for retained checkpoint boundaries and +workspace; preserving its old heuristic does not make it calibrated. MoE/EP/CP, +other full-recompute intervals, optional selective modules, deep prefix trees, +mixed gradient groups, other LoRA ranks, and pretrained correctness are untested. +This commit adds measurement evidence without changing the estimator coefficients. + +## Reproduce + +After `INSTALL_VLLM_RUNTIME=false bash src/art/megatron/setup.sh`, run on two GPUs: + +```sh +ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE=2 \ +ART_MEGATRON_CONTEXT_PARALLEL_SIZE=1 \ +ART_MEGATRON_DATA_PARALLEL_SIZE=1 \ +ART_MEGATRON_PIPELINE_MODEL_PARALLEL_SIZE=1 \ +uv run --project megatron_runtime --no-sync python -m torch.distributed.run \ + --standalone --nproc-per-node=2 dev/trainer_rank_recompute_memory.py \ + --mode selective --tokens 1024 2048 3072 4096 --reported-pair \ + --evidence scratch/recompute-memory/tp2-selective.jsonl +``` + +Repeat with `--mode none` and `--mode full`. For the attention-only control, add +`--model Qwen/Qwen3-1.7B --tokens 1024 4096 16384` and omit `--reported-pair`. +The [SkyPilot task](trainer_rank_recompute_memory.sky.yaml) runs all three modes +at TP4; pass the synced commit as `ART_CALIBRATION_SOURCE_SHA` and launch with +`--idle-minutes-to-autostop 10 --down`. Override both GPU count and TP environment +variable for TP1. The driver emits per-rank JSONL and identical rows in job logs. +Driver SHA-256 for this campaign: +`33a7291041913b60232ac2e287edf26c7766b7ab268af034dbcd71df0eae840e`. diff --git a/dev/trainer_rank_recompute_memory.py b/dev/trainer_rank_recompute_memory.py new file mode 100644 index 000000000..ab7ab185f --- /dev/null +++ b/dev/trainer_rank_recompute_memory.py @@ -0,0 +1,181 @@ +"""Measure cold unsplit admission and native CUDA peaks for one recompute mode. + +Run with torchrun; use a fresh process per mode/topology. Random weights keep +the native model geometry and kernels without downloading a checkpoint. This +measures memory, not pretrained-model correctness. Refused plans are recorded +without execution. Each repetition clears the learned memory profile, while +sample 0 includes cold compilation/autotuning. Output is one JSONL row per rank. +""" + +import argparse +from dataclasses import asdict +import gc +import hashlib +import json +import os +from pathlib import Path +import subprocess +import time + +from dotenv import load_dotenv +import torch +import torch.distributed as dist + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--model", default="Qwen/Qwen3.8-27B") + parser.add_argument("--mode", choices=("full", "selective", "none"), required=True) + parser.add_argument("--layers", type=int, default=0) + parser.add_argument("--tokens", type=int, nargs="+", default=[1024, 2048, 4096]) + parser.add_argument("--repeat", type=int, default=2) + parser.add_argument("--reported-pair", action="store_true") + parser.add_argument("--evidence", type=Path, required=True) + args = parser.parse_args() + if args.repeat < 1 or args.layers < 0 or any(n < 1 for n in args.tokens): + parser.error("repeat/tokens must be positive and layers nonnegative") + load_dotenv(".env") + from trainer_rank_support import load_random_checkpoints + + from art.megatron import train + from art.trainer_rank import ForwardInput, TrainerRank + + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + dist.init_process_group("nccl") + + def configure(provider): + if args.layers: + provider.num_layers = args.layers + provider.recompute_granularity = None if args.mode == "none" else args.mode + provider.recompute_method = "uniform" if args.mode == "full" else None + provider.recompute_num_layers = 1 if args.mode == "full" else None + provider.recompute_modules = ["core_attn"] if args.mode == "selective" else [] + + def emit(row): + gathered = [None] * dist.get_world_size() + dist.all_gather_object(gathered, {**facts, **row, "rank": dist.get_rank()}) + if dist.get_rank() == 0: + with args.evidence.open("a") as stream: + for item in gathered: + stream.write(json.dumps(item, default=str, sort_keys=True) + "\n") + print("MEMORY_CALIBRATION " + json.dumps(gathered, default=str), flush=True) + + try: + torch.manual_seed(913) + runtime = train.build_training_runtime( + model_identifier=args.model, + model_initialization="random", + provider_configure=configure, + print_env=False, + ) + for chunk in runtime.model: + chunk.train() + rank = TrainerRank(runtime) + [slot] = load_random_checkpoints( + runtime, rank, 1, base_model=args.model, lora_rank=1 + ) + facts = { + "schema": "art.dev.recompute_memory.v1", + "driver_sha256": hashlib.sha256(Path(__file__).read_bytes()).hexdigest(), + "source_sha": os.environ.get("ART_CALIBRATION_SOURCE_SHA") + or subprocess.check_output(["git", "rev-parse", "HEAD"], text=True).strip(), + "model": args.model, + "initialization": "random", + "lora_rank": 1, + "mode": rank._recompute_granularity, + "recompute_method": runtime.provider.recompute_method, + "recompute_num_layers": runtime.provider.recompute_num_layers, + "recompute_modules": runtime.provider.recompute_modules, + "geometry": rank._geometry.as_dict(), + "layers": rank._num_layers, + "topology_dp_tp_cp_pp": rank._topology_key(), + "dtype": str(next(runtime.model[0].parameters()).dtype), + "device": torch.cuda.get_device_name(), + "device_total_bytes": torch.cuda.get_device_properties(0).total_memory, + "torch": torch.__version__, + "cuda": torch.version.cuda, + "transformer_layers_compiled": runtime.transformer_layers_compiled, + } + args.evidence.parent.mkdir(parents=True, exist_ok=True) + workloads = [([length], 0) for length in args.tokens] + if args.reported_pair: + # Same logical/fully-shared packed counts as #913; synthetic IDs. + workloads.append(([19221, 19222], 5733)) + generator = torch.Generator().manual_seed(913) + for lengths, prefix in workloads: + tokens = [ + torch.randint(100, 10000, (n,), generator=generator) for n in lengths + ] + for item in tokens[1:]: + item[:prefix] = tokens[0][:prefix] + requests = [ + ForwardInput(input_tokens=item, hidden_states=True) for item in tokens + ] + plan = rank._plan_flat_forward(requests, checkpoint=slot) + for sample in range(args.repeat): + rank.zero_grad() + rank._memory_profiles.clear() + gc.collect() + torch.cuda.empty_cache() + dist.barrier() + torch.cuda.synchronize() + check = rank._memory_check(plan) + row = { + "lengths": lengths, + "shared_prefix": prefix, + "logical_tokens": plan.logical_tokens, + "packed_tokens": plan.packed_tokens, + "output_bytes": plan.output_bytes, + "selected_max_depth": plan.selected_max_depth, + "sample": sample, + "admission": asdict(check), + } + if not check.fits: + emit({**row, "status": "refused"}) + break + baseline = torch.cuda.memory_allocated() + reserved = torch.cuda.memory_reserved() + torch.cuda.reset_peak_memory_stats() + started = time.monotonic() + # Direct execution isolates the selected unsplit plan. It uses + # the same native forward as public admission, with no split. + outputs = rank._execute_flat_plan(plan) + torch.cuda.synchronize() + forward_peak = torch.cuda.max_memory_allocated() + retained = torch.cuda.memory_allocated() + forward_seconds = time.monotonic() - started + terms = [ + output.hidden_states.float().square().mean() + for output in outputs + if output.hidden_states is not None + ] + assert len(terms) == len(requests) + loss = torch.stack(terms).sum() + loss.backward() + torch.cuda.synchronize() + emit( + { + **row, + "status": "measured", + "baseline_allocated_bytes": baseline, + "baseline_reserved_bytes": reserved, + "forward_peak_delta_bytes": forward_peak - baseline, + "retained_delta_bytes": retained - baseline, + "forward_backward_peak_delta_bytes": torch.cuda.max_memory_allocated() + - baseline, + "peak_reserved_bytes": torch.cuda.max_memory_reserved(), + "forward_seconds": forward_seconds, + "forward_backward_seconds": time.monotonic() - started, + "finite_loss": bool(torch.isfinite(loss).item()), + "estimate_covers_forward": check.estimated_required_bytes + >= forward_peak - baseline, + } + ) + del outputs, loss, terms + rank.zero_grad() + finally: + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_rank_recompute_memory.sky.yaml b/dev/trainer_rank_recompute_memory.sky.yaml new file mode 100644 index 000000000..4e9b55f15 --- /dev/null +++ b/dev/trainer_rank_recompute_memory.sky.yaml @@ -0,0 +1,51 @@ +# Launch with --infra k8s/cks-wb3 --idle-minutes-to-autostop 10 --down. +# Set ART_CALIBRATION_SOURCE_SHA to the commit synced in workdir. +name: trainer-rank-recompute-memory +workdir: . + +resources: + infra: k8s/cks-wb3 + accelerators: H200:4 + cpus: 32+ + memory: 256+ + image_id: docker:docker.io/bradhiltonnw/art-gpu:latest + +envs: + ART_CALIBRATION_SOURCE_SHA: unset + ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE: "4" + ART_MEGATRON_CONTEXT_PARALLEL_SIZE: "1" + ART_MEGATRON_DATA_PARALLEL_SIZE: "1" + ART_MEGATRON_PIPELINE_MODEL_PARALLEL_SIZE: "1" + OMP_NUM_THREADS: "1" + OPENBLAS_NUM_THREADS: "1" + MKL_NUM_THREADS: "1" + PYTHONUNBUFFERED: "1" + TOKENIZERS_PARALLELISM: "false" + +setup: | + INSTALL_VLLM_RUNTIME=false bash src/art/megatron/setup.sh + +run: | + set -euo pipefail + export PYTHONPATH="$PWD/src:$PWD" + for mode in selective none full; do + timeout --signal=TERM --kill-after=30s 20m \ + megatron_runtime/.venv/bin/python -m torch.distributed.run \ + --standalone --nproc-per-node="$ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE" \ + dev/trainer_rank_recompute_memory.py --mode "$mode" \ + --tokens 1024 2048 3072 4096 --reported-pair \ + --evidence "scratch/recompute-memory/tp$ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE-$mode.jsonl" + done + +config: + kubernetes: + pod_config: + spec: + schedulerName: binpack-scheduler + activeDeadlineSeconds: 7200 + containers: + - name: ray-node + imagePullPolicy: Always + env: + - name: UV_LINK_MODE + value: copy From be21a2599c743d89a62004ba5e9c662c518cd91a Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 16 Sep 2026 21:39:19 +0000 Subject: [PATCH 03/15] Normalize calibration CSV line endings --- dev/trainer_rank_recompute_memory.csv | 110 +++++++++++++------------- 1 file changed, 55 insertions(+), 55 deletions(-) diff --git a/dev/trainer_rank_recompute_memory.csv b/dev/trainer_rank_recompute_memory.csv index cadbf61e3..293ab2665 100644 --- a/dev/trainer_rank_recompute_memory.csv +++ b/dev/trainer_rank_recompute_memory.csv @@ -1,55 +1,55 @@ -model,tp,mode,logical_tokens,packed_tokens,estimate_bytes,available_min_bytes,measured_rank_samples,refused_rank_samples,forward_peak_max_bytes,retained_max_bytes,forward_backward_peak_max_bytes,baseline_min_bytes,baseline_max_bytes,finite_loss_all,estimate_covers_forward_all -Qwen/Qwen3-1.7B,2,full,1024,1024,55364812,140941082317,4,0,177251840,133205504,231861248,1781843968,1883032576,True,False -Qwen/Qwen3-1.7B,2,full,4096,4096,221459251,140938985165,4,0,415402496,264382976,587698176,1883032576,1885129728,True,False -Qwen/Qwen3-1.7B,2,full,16384,16384,885837004,140930596557,4,0,1661798912,1057720832,2099318784,1885129728,1893518336,True,False -Qwen/Qwen3-1.7B,2,none,1024,1024,2717489561,140941082317,4,0,1440448000,1427860480,1473997312,1781843968,1883032576,True,True -Qwen/Qwen3-1.7B,2,none,4096,4096,10869958246,140938985165,4,0,5493352960,5443004416,5610778112,1883032576,1885129728,True,True -Qwen/Qwen3-1.7B,2,none,16384,16384,43479832985,140930596557,4,0,21973600768,21772208128,22443298304,1885129728,1893518336,True,True -Qwen/Qwen3-1.7B,2,selective,1024,1024,2717489561,140941082317,4,0,1439530496,1426943488,1535994880,1781843968,1883032576,True,True -Qwen/Qwen3-1.7B,2,selective,4096,4096,10869958246,140938985165,4,0,5489682944,5439334912,5607108608,1883032576,1885129728,True,True -Qwen/Qwen3-1.7B,2,selective,16384,16384,43479832985,140930596557,4,0,21958920704,21757528576,22428618752,1885129728,1893518336,True,True -Qwen/Qwen3.8-27B,1,full,1024,1024,196083712,89386251981,2,0,973403648,759490048,1198017536,54352220160,54452884480,True,False -Qwen/Qwen3.8-27B,1,full,2048,2048,392167424,89384154829,2,0,1812587008,1384759808,2092073472,54452884480,54452884480,True,False -Qwen/Qwen3.8-27B,1,full,3072,3072,588251136,89384154829,2,0,2718884352,2077143552,3134828032,54452884480,54452884480,True,False -Qwen/Qwen3.8-27B,1,full,4096,4096,784334848,89384154829,2,0,3625185792,2769531392,4178635264,54452884480,54452884480,True,False -Qwen/Qwen3.8-27B,1,full,38443,32710,6328148992,89321231565,2,0,28953189376,22178822144,33419714048,54452884480,54452893184,True,False -Qwen/Qwen3.8-27B,1,none,1024,1024,31311108505,89386251981,2,0,17041559040,16955570688,17060460544,54352220160,54452884480,True,True -Qwen/Qwen3.8-27B,1,none,2048,2048,62622217011,89386251981,2,0,33761374720,33589400064,33799116800,54452884480,54452884480,True,True -Qwen/Qwen3.8-27B,1,none,3072,3072,93933325516,89386251981,0,1,,,,,,, -Qwen/Qwen3.8-27B,1,none,4096,4096,125244434022,89386251981,0,1,,,,,,, -Qwen/Qwen3.8-27B,1,none,38443,32710,1000246567936,89386251981,0,1,,,,,,, -Qwen/Qwen3.8-27B,1,selective,1024,1024,31311108505,87997937357,2,0,17043131904,16957146624,17124951040,54352220160,54452884480,True,True -Qwen/Qwen3.8-27B,1,selective,2048,2048,62622217011,87997937357,2,0,33758228992,33586258432,33795975168,54452884480,54452884480,True,True -Qwen/Qwen3.8-27B,1,selective,3072,3072,93933325516,87997937357,0,1,,,,,,, -Qwen/Qwen3.8-27B,1,selective,4096,4096,125244434022,87997937357,0,1,,,,,,, -Qwen/Qwen3.8-27B,1,selective,38443,32710,1000246567936,87997937357,0,1,,,,,,, -Qwen/Qwen3.8-27B,2,full,1024,1024,196083712,115395819213,4,0,526710272,419749376,710319104,27327631360,27428295680,True,False -Qwen/Qwen3.8-27B,2,full,2048,2048,392167424,115393722061,4,0,917103104,703181312,1098500608,27428295680,27428295680,True,False -Qwen/Qwen3.8-27B,2,full,3072,3072,588251136,115393722061,4,0,1375658496,1054775808,1646584832,27428295680,27428295680,True,False -Qwen/Qwen3.8-27B,2,full,4096,4096,784334848,115393722061,4,0,1834217984,1406374400,2194869760,27428295680,27428295680,True,False -Qwen/Qwen3.8-27B,2,full,38443,32710,6328148992,115391616205,4,0,14653886976,11294938624,17587685376,27428295680,27428304384,True,False -Qwen/Qwen3.8-27B,2,none,1024,1024,31311108505,115395819213,4,0,10079153664,10041397760,10146256896,27327631360,27428295680,True,True -Qwen/Qwen3.8-27B,2,none,2048,2048,62622217011,115395819213,4,0,19542929920,19467420160,19677136896,27428295680,27428295680,True,True -Qwen/Qwen3.8-27B,2,none,3072,3072,93933325516,115395819213,4,0,29394684416,29281420800,29595995136,27428295680,27428295680,True,True -Qwen/Qwen3.8-27B,2,none,4096,4096,125244434022,115395819213,0,2,,,,,,, -Qwen/Qwen3.8-27B,2,none,38443,32710,1000246567936,115395819213,0,2,,,,,,, -Qwen/Qwen3.8-27B,2,selective,1024,1024,31311108505,115395819213,4,0,10082561536,10044808704,10149667840,27327631360,27428295680,True,True -Qwen/Qwen3.8-27B,2,selective,2048,2048,62622217011,115395819213,4,0,19541553664,19466048000,19675764736,27428295680,27428295680,True,True -Qwen/Qwen3.8-27B,2,selective,3072,3072,93933325516,115393722061,4,0,29392325120,29279066624,29593640960,27428295680,27428295680,True,True -Qwen/Qwen3.8-27B,2,selective,4096,4096,125244434022,115393722061,0,2,,,,,,, -Qwen/Qwen3.8-27B,2,selective,38443,32710,1000246567936,115393722061,0,2,,,,,,, -Qwen/Qwen3.8-27B,4,full,1024,1024,196083712,128855946957,8,0,308082176,254597632,455984128,13865406464,13966070784,True,False -Qwen/Qwen3.8-27B,4,full,2048,2048,392167424,128855946957,8,0,480371200,373402112,591982080,13966070784,13966070784,True,False -Qwen/Qwen3.8-27B,4,full,3072,3072,588251136,128853849805,8,0,708239872,545689088,873627648,13966070784,13966070784,True,False -Qwen/Qwen3.8-27B,4,full,4096,4096,784334848,128853849805,8,0,939258368,725320192,1161568256,13966070784,13966070784,True,False -Qwen/Qwen3.8-27B,4,full,38443,32712,6328509440,128851743949,8,0,7502990336,5851191296,9831958528,13966070784,13966079488,True,False -Qwen/Qwen3.8-27B,4,none,1024,1024,31311108505,128855946957,8,0,5736394240,5722230272,5827089408,13865406464,13966070784,True,True -Qwen/Qwen3.8-27B,4,none,2048,2048,62622217011,128855946957,8,0,11131212288,11103934976,11313651712,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,none,3072,3072,93933325516,128855946957,8,0,16653462016,16610974208,16925548544,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,none,4096,4096,125244434022,128855946957,8,0,21989044736,21934492160,22353924096,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,none,38443,32712,1000307699916,128855946957,0,4,,,,,,, -Qwen/Qwen3.8-27B,4,selective,1024,1024,31311108505,127666861773,8,0,5732331008,5718170112,5932531200,13865406464,13966070784,True,True -Qwen/Qwen3.8-27B,4,selective,2048,2048,62622217011,127666861773,8,0,11132523008,11105249792,11314966528,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,selective,3072,3072,93933325516,127666861773,8,0,16651233792,16608751104,16923325440,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,selective,4096,4096,125244434022,127666861773,8,0,21987471872,21932925440,22352357376,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,selective,38443,32712,1000307699916,127666861773,0,4,,,,,,, +model,tp,mode,logical_tokens,packed_tokens,estimate_bytes,available_min_bytes,measured_rank_samples,refused_rank_samples,forward_peak_max_bytes,retained_max_bytes,forward_backward_peak_max_bytes,baseline_min_bytes,baseline_max_bytes,finite_loss_all,estimate_covers_forward_all +Qwen/Qwen3-1.7B,2,full,1024,1024,55364812,140941082317,4,0,177251840,133205504,231861248,1781843968,1883032576,True,False +Qwen/Qwen3-1.7B,2,full,4096,4096,221459251,140938985165,4,0,415402496,264382976,587698176,1883032576,1885129728,True,False +Qwen/Qwen3-1.7B,2,full,16384,16384,885837004,140930596557,4,0,1661798912,1057720832,2099318784,1885129728,1893518336,True,False +Qwen/Qwen3-1.7B,2,none,1024,1024,2717489561,140941082317,4,0,1440448000,1427860480,1473997312,1781843968,1883032576,True,True +Qwen/Qwen3-1.7B,2,none,4096,4096,10869958246,140938985165,4,0,5493352960,5443004416,5610778112,1883032576,1885129728,True,True +Qwen/Qwen3-1.7B,2,none,16384,16384,43479832985,140930596557,4,0,21973600768,21772208128,22443298304,1885129728,1893518336,True,True +Qwen/Qwen3-1.7B,2,selective,1024,1024,2717489561,140941082317,4,0,1439530496,1426943488,1535994880,1781843968,1883032576,True,True +Qwen/Qwen3-1.7B,2,selective,4096,4096,10869958246,140938985165,4,0,5489682944,5439334912,5607108608,1883032576,1885129728,True,True +Qwen/Qwen3-1.7B,2,selective,16384,16384,43479832985,140930596557,4,0,21958920704,21757528576,22428618752,1885129728,1893518336,True,True +Qwen/Qwen3.8-27B,1,full,1024,1024,196083712,89386251981,2,0,973403648,759490048,1198017536,54352220160,54452884480,True,False +Qwen/Qwen3.8-27B,1,full,2048,2048,392167424,89384154829,2,0,1812587008,1384759808,2092073472,54452884480,54452884480,True,False +Qwen/Qwen3.8-27B,1,full,3072,3072,588251136,89384154829,2,0,2718884352,2077143552,3134828032,54452884480,54452884480,True,False +Qwen/Qwen3.8-27B,1,full,4096,4096,784334848,89384154829,2,0,3625185792,2769531392,4178635264,54452884480,54452884480,True,False +Qwen/Qwen3.8-27B,1,full,38443,32710,6328148992,89321231565,2,0,28953189376,22178822144,33419714048,54452884480,54452893184,True,False +Qwen/Qwen3.8-27B,1,none,1024,1024,31311108505,89386251981,2,0,17041559040,16955570688,17060460544,54352220160,54452884480,True,True +Qwen/Qwen3.8-27B,1,none,2048,2048,62622217011,89386251981,2,0,33761374720,33589400064,33799116800,54452884480,54452884480,True,True +Qwen/Qwen3.8-27B,1,none,3072,3072,93933325516,89386251981,0,1,,,,,,, +Qwen/Qwen3.8-27B,1,none,4096,4096,125244434022,89386251981,0,1,,,,,,, +Qwen/Qwen3.8-27B,1,none,38443,32710,1000246567936,89386251981,0,1,,,,,,, +Qwen/Qwen3.8-27B,1,selective,1024,1024,31311108505,87997937357,2,0,17043131904,16957146624,17124951040,54352220160,54452884480,True,True +Qwen/Qwen3.8-27B,1,selective,2048,2048,62622217011,87997937357,2,0,33758228992,33586258432,33795975168,54452884480,54452884480,True,True +Qwen/Qwen3.8-27B,1,selective,3072,3072,93933325516,87997937357,0,1,,,,,,, +Qwen/Qwen3.8-27B,1,selective,4096,4096,125244434022,87997937357,0,1,,,,,,, +Qwen/Qwen3.8-27B,1,selective,38443,32710,1000246567936,87997937357,0,1,,,,,,, +Qwen/Qwen3.8-27B,2,full,1024,1024,196083712,115395819213,4,0,526710272,419749376,710319104,27327631360,27428295680,True,False +Qwen/Qwen3.8-27B,2,full,2048,2048,392167424,115393722061,4,0,917103104,703181312,1098500608,27428295680,27428295680,True,False +Qwen/Qwen3.8-27B,2,full,3072,3072,588251136,115393722061,4,0,1375658496,1054775808,1646584832,27428295680,27428295680,True,False +Qwen/Qwen3.8-27B,2,full,4096,4096,784334848,115393722061,4,0,1834217984,1406374400,2194869760,27428295680,27428295680,True,False +Qwen/Qwen3.8-27B,2,full,38443,32710,6328148992,115391616205,4,0,14653886976,11294938624,17587685376,27428295680,27428304384,True,False +Qwen/Qwen3.8-27B,2,none,1024,1024,31311108505,115395819213,4,0,10079153664,10041397760,10146256896,27327631360,27428295680,True,True +Qwen/Qwen3.8-27B,2,none,2048,2048,62622217011,115395819213,4,0,19542929920,19467420160,19677136896,27428295680,27428295680,True,True +Qwen/Qwen3.8-27B,2,none,3072,3072,93933325516,115395819213,4,0,29394684416,29281420800,29595995136,27428295680,27428295680,True,True +Qwen/Qwen3.8-27B,2,none,4096,4096,125244434022,115395819213,0,2,,,,,,, +Qwen/Qwen3.8-27B,2,none,38443,32710,1000246567936,115395819213,0,2,,,,,,, +Qwen/Qwen3.8-27B,2,selective,1024,1024,31311108505,115395819213,4,0,10082561536,10044808704,10149667840,27327631360,27428295680,True,True +Qwen/Qwen3.8-27B,2,selective,2048,2048,62622217011,115395819213,4,0,19541553664,19466048000,19675764736,27428295680,27428295680,True,True +Qwen/Qwen3.8-27B,2,selective,3072,3072,93933325516,115393722061,4,0,29392325120,29279066624,29593640960,27428295680,27428295680,True,True +Qwen/Qwen3.8-27B,2,selective,4096,4096,125244434022,115393722061,0,2,,,,,,, +Qwen/Qwen3.8-27B,2,selective,38443,32710,1000246567936,115393722061,0,2,,,,,,, +Qwen/Qwen3.8-27B,4,full,1024,1024,196083712,128855946957,8,0,308082176,254597632,455984128,13865406464,13966070784,True,False +Qwen/Qwen3.8-27B,4,full,2048,2048,392167424,128855946957,8,0,480371200,373402112,591982080,13966070784,13966070784,True,False +Qwen/Qwen3.8-27B,4,full,3072,3072,588251136,128853849805,8,0,708239872,545689088,873627648,13966070784,13966070784,True,False +Qwen/Qwen3.8-27B,4,full,4096,4096,784334848,128853849805,8,0,939258368,725320192,1161568256,13966070784,13966070784,True,False +Qwen/Qwen3.8-27B,4,full,38443,32712,6328509440,128851743949,8,0,7502990336,5851191296,9831958528,13966070784,13966079488,True,False +Qwen/Qwen3.8-27B,4,none,1024,1024,31311108505,128855946957,8,0,5736394240,5722230272,5827089408,13865406464,13966070784,True,True +Qwen/Qwen3.8-27B,4,none,2048,2048,62622217011,128855946957,8,0,11131212288,11103934976,11313651712,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,none,3072,3072,93933325516,128855946957,8,0,16653462016,16610974208,16925548544,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,none,4096,4096,125244434022,128855946957,8,0,21989044736,21934492160,22353924096,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,none,38443,32712,1000307699916,128855946957,0,4,,,,,,, +Qwen/Qwen3.8-27B,4,selective,1024,1024,31311108505,127666861773,8,0,5732331008,5718170112,5932531200,13865406464,13966070784,True,True +Qwen/Qwen3.8-27B,4,selective,2048,2048,62622217011,127666861773,8,0,11132523008,11105249792,11314966528,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,selective,3072,3072,93933325516,127666861773,8,0,16651233792,16608751104,16923325440,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,selective,4096,4096,125244434022,127666861773,8,0,21987471872,21932925440,22352357376,13966070784,13966070784,True,True +Qwen/Qwen3.8-27B,4,selective,38443,32712,1000307699916,127666861773,0,4,,,,,,, From e1d00a32f0aeede991ff01e95f6242f18dde69f6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 16 Sep 2026 21:48:46 +0000 Subject: [PATCH 04/15] Account for TP, sequence parallelism and hybrid layer counts --- dev/trainer_rank_recompute_memory.py | 18 ++++- dev/trainer_rank_recompute_memory.sky.yaml | 13 ++-- src/art/trainer_rank/_impl.py | 58 ++++++++++----- .../test_trainer_rank_recompute_memory.py | 72 +++++++++++++++++++ 4 files changed, 138 insertions(+), 23 deletions(-) diff --git a/dev/trainer_rank_recompute_memory.py b/dev/trainer_rank_recompute_memory.py index ab7ab185f..fa2c595e8 100644 --- a/dev/trainer_rank_recompute_memory.py +++ b/dev/trainer_rank_recompute_memory.py @@ -30,10 +30,17 @@ def main() -> None: parser.add_argument("--tokens", type=int, nargs="+", default=[1024, 2048, 4096]) parser.add_argument("--repeat", type=int, default=2) parser.add_argument("--reported-pair", action="store_true") + parser.add_argument( + "--pairs", action="store_true", help="Two sequences at each token length" + ) + parser.add_argument("--prefix-fraction", type=float, default=0.3) + parser.add_argument("--modules", nargs="+", default=["core_attn"]) parser.add_argument("--evidence", type=Path, required=True) args = parser.parse_args() if args.repeat < 1 or args.layers < 0 or any(n < 1 for n in args.tokens): parser.error("repeat/tokens must be positive and layers nonnegative") + if not 0 <= args.prefix_fraction < 1: + parser.error("prefix-fraction must be in [0, 1)") load_dotenv(".env") from trainer_rank_support import load_random_checkpoints @@ -49,7 +56,7 @@ def configure(provider): provider.recompute_granularity = None if args.mode == "none" else args.mode provider.recompute_method = "uniform" if args.mode == "full" else None provider.recompute_num_layers = 1 if args.mode == "full" else None - provider.recompute_modules = ["core_attn"] if args.mode == "selective" else [] + provider.recompute_modules = args.modules if args.mode == "selective" else [] def emit(row): gathered = [None] * dist.get_world_size() @@ -88,6 +95,8 @@ def emit(row): "recompute_modules": runtime.provider.recompute_modules, "geometry": rank._geometry.as_dict(), "layers": rank._num_layers, + "gdn_layers": rank._gdn_layers, + "sequence_parallel": rank._sequence_parallel, "topology_dp_tp_cp_pp": rank._topology_key(), "dtype": str(next(runtime.model[0].parameters()).dtype), "device": torch.cuda.get_device_name(), @@ -97,7 +106,12 @@ def emit(row): "transformer_layers_compiled": runtime.transformer_layers_compiled, } args.evidence.parent.mkdir(parents=True, exist_ok=True) - workloads = [([length], 0) for length in args.tokens] + workloads = [ + ([length, length], int(length * args.prefix_fraction)) + if args.pairs + else ([length], 0) + for length in args.tokens + ] if args.reported_pair: # Same logical/fully-shared packed counts as #913; synthetic IDs. workloads.append(([19221, 19222], 5733)) diff --git a/dev/trainer_rank_recompute_memory.sky.yaml b/dev/trainer_rank_recompute_memory.sky.yaml index 4e9b55f15..ad5b95384 100644 --- a/dev/trainer_rank_recompute_memory.sky.yaml +++ b/dev/trainer_rank_recompute_memory.sky.yaml @@ -21,6 +21,11 @@ envs: MKL_NUM_THREADS: "1" PYTHONUNBUFFERED: "1" TOKENIZERS_PARALLELISM: "false" + MODEL: Qwen/Qwen3.8-27B + MODES: selective none + MODULES: core_attn + CALIBRATION_ARGS: --pairs --tokens 2048 4096 8192 --reported-pair + EVIDENCE_DIR: scratch/recompute-memory-sharded setup: | INSTALL_VLLM_RUNTIME=false bash src/art/megatron/setup.sh @@ -28,13 +33,13 @@ setup: | run: | set -euo pipefail export PYTHONPATH="$PWD/src:$PWD" - for mode in selective none full; do + for mode in ${MODES}; do timeout --signal=TERM --kill-after=30s 20m \ megatron_runtime/.venv/bin/python -m torch.distributed.run \ --standalone --nproc-per-node="$ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE" \ - dev/trainer_rank_recompute_memory.py --mode "$mode" \ - --tokens 1024 2048 3072 4096 --reported-pair \ - --evidence "scratch/recompute-memory/tp$ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE-$mode.jsonl" + dev/trainer_rank_recompute_memory.py --model "$MODEL" --mode "$mode" \ + --modules ${MODULES} ${CALIBRATION_ARGS} \ + --evidence "$EVIDENCE_DIR/tp$ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE-$mode.jsonl" done config: diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 6921a5f83..abd809ca3 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1362,8 +1362,8 @@ def __init__(self, runtime: TrainingRuntime) -> None: # gradient reduction pre-date the planner, memory checks all-reduce # within the TP x CP group, the memory profile is keyed by topology so # TP calibrates itself online, and the fitted layout cost model prices - # TP explicitly. Known limitation: the cold static memory estimate - # ignores sharding (conservative). + # TP explicitly. The cold retained-activation floor also distinguishes + # tensor/sequence-parallel storage from gathered LoRA inputs. self.runtime: TrainingRuntime = runtime self.device: torch.device = next(runtime.model[0].parameters()).device self._param_dtype_size = _dtype_size(next(runtime.model[0].parameters()).dtype) @@ -1385,6 +1385,13 @@ def __init__(self, runtime: TrainingRuntime) -> None: "recompute_granularity", getattr(runtime.provider, "recompute_granularity", None), ) + self._sequence_parallel = bool( + getattr( + getattr(metadata_model, "config", None), + "sequence_parallel", + getattr(runtime.provider, "sequence_parallel", False), + ) + ) # Layers that run the gated-delta-net path (Qwen3.5-4B: 24 of 32); the # cost model prices GDN state hand-offs per GDN layer, not per layer. self._gdn_layers = _gdn_layer_count(runtime.model[0]) @@ -5138,28 +5145,45 @@ def _estimate_required_memory_bytes_from_values( geometry.moe_topk * geometry.moe_ffn_hidden_size + geometry.moe_shared_expert_ffn, ) + hidden = self._hidden_size attention_width = max( - self._hidden_size, - geometry.num_attention_heads * geometry.kv_channels, + hidden, geometry.num_attention_heads * geometry.kv_channels + ) + gdn_width = max( + hidden, 2 * geometry.gdn_key_heads * geometry.gdn_key_head_dim + 2 * geometry.gdn_value_heads * geometry.gdn_value_head_dim, ) - # Megatron's selective-recompute model retains per-layer attention, - # norm and MLP tensors (arxiv.org/abs/2205.05198). Charge four FFN - # widths for gated MLPs, and routed hidden rows for MoE. Do not take - # TP/CP/EP or optional selective-module discounts without measured - # evidence; the default core_attn checkpoint leaves the MLP live. - layer_features = ( - 9 * attention_width - + 4 * ffn_width - + 2 * self._hidden_size * max(0, geometry.moe_topk - 1) + tp = max(1, self._topology_key()[1]) + sp = tp if self._sequence_parallel else 1 + gdn_layers = min(self._num_layers, self._gdn_layers) + # Split Megatron's 9H attention/norm term into 5H of norms/residuals + # (sequence parallel) and projection intermediates (tensor parallel). + # Gated MLPs retain four FFN widths. ART's column-parallel LoRA path + # additionally retains gathered H-wide inputs to attention and MLP, + # even with SP. Price hybrid layers separately, not at the max width. + layer_features = 5 * hidden / sp + 2 * hidden + if geometry.moe_experts: + # Routed/shared expert storage gets no TP/EP discount without + # evidence for expert sharding and imbalanced dispatch. + layer_features += 4 * ffn_width + 2 * hidden * max( + 0, geometry.moe_topk - 1 + ) + else: + layer_features += 4 * ffn_width / tp + retained_features = ( + self._num_layers * layer_features + + ( + (self._num_layers - gdn_layers) * (9 * attention_width - 5 * hidden) + + gdn_layers * (9 * gdn_width - 5 * hidden) + ) + / tp ) + # Optional selective modules retain the conservative undiscounted + # floor until their effective native checkpoint boundaries are tested. static_compute = max( static_compute, - packed_tokens - * self._num_layers - * self._param_dtype_size - * layer_features, + packed_tokens * self._param_dtype_size * retained_features, ) # Groups execute sequentially: summed packed rows conservatively bound # this FC2 component, not all workspace or retained graphs. diff --git a/tests/unit/test_trainer_rank_recompute_memory.py b/tests/unit/test_trainer_rank_recompute_memory.py index 2d8b70f03..28e2d8563 100644 --- a/tests/unit/test_trainer_rank_recompute_memory.py +++ b/tests/unit/test_trainer_rank_recompute_memory.py @@ -134,6 +134,78 @@ def test_profile_cannot_erase_recompute_floor() -> None: assert rank._memory_check(plan).estimated_required_bytes > cold +def _hybrid_rank(monkeypatch: pytest.MonkeyPatch, tp: int) -> TrainerRank: + rank = _rank( + "selective", + sequence_parallel=True, + linear_num_key_heads=16, + linear_key_head_dim=128, + linear_num_value_heads=48, + linear_value_head_dim=128, + ) + rank._gdn_layers = 48 + monkeypatch.setattr(rank, "_topology_key", lambda: (1, tp, 1, 1)) + return rank + + +@pytest.mark.parametrize( + "tp,peak", [(1, 33758228992), (2, 19541553664), (4, 11132523008)] +) +def test_sharded_floor_covers_recorded_native_gdn_peaks(monkeypatch, tp, peak): + # H200, 64-layer Qwen3.8-27B, LoRA r1, SP, cold/warm max, 2,048 tokens. + # Source evidence: dev/trainer_rank_recompute_memory.csv at 22f628d6. + rank = _hybrid_rank(monkeypatch, tp) + assert rank._memory_check(_plan(rank, tokens=2048)).estimated_required_bytes >= peak + + +def test_tp4_admits_eight_k_sibling_pair_and_prices_actual_layer_mix(monkeypatch): + rank = _hybrid_rank(monkeypatch, 4) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 119 * 2**30) + plan = replace( + _plan(rank, tokens=13928), logical_tokens=16384, output_bytes=16384 * 5120 * 2 + ) + check = rank._memory_check(plan) + assert check.fits + rank._gdn_layers = 64 + assert ( + rank._memory_check(plan).estimated_required_bytes + > check.estimated_required_bytes + ) + rank._gdn_layers = 0 + assert ( + rank._memory_check(plan).estimated_required_bytes + < check.estimated_required_bytes + ) + + +def test_sp_discount_excludes_gathered_lora_inputs(monkeypatch): + rank = _hybrid_rank(monkeypatch, 4) + plan = _plan(rank, tokens=2048) + tp4 = rank._memory_check(plan).estimated_required_bytes + rank._sequence_parallel = False + assert rank._memory_check(plan).estimated_required_bytes > tp4 + monkeypatch.setattr(rank, "_topology_key", lambda: (1, 1, 1, 1)) + assert tp4 > rank._memory_check(plan).estimated_required_bytes / 4 + + +def test_gathered_inputs_cover_attention_only_cold_peak(monkeypatch): + # Dividing the old whole estimate by TP misses this 1,439,530,496-byte peak. + rank = _rank( + "selective", + hidden_size=2048, + ffn_hidden_size=6144, + num_layers=28, + num_attention_heads=16, + kv_channels=128, + sequence_parallel=True, + ) + monkeypatch.setattr(rank, "_topology_key", lambda: (1, 2, 1, 1)) + assert ( + rank._memory_check(_plan(rank, tokens=1024)).estimated_required_bytes + >= 1439530496 + ) + + def test_non_full_estimate_covers_retained_gated_mlp_tensors() -> None: # Selective core-attention recompute leaves this MLP graph live. Count # distinct saved activation storage, excluding model parameters/views. From ad4e86b45d680d581cff64361ee7c37a7aac1f07 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 16 Sep 2026 21:51:39 +0000 Subject: [PATCH 05/15] Record expert topology and gradient health during calibration --- dev/trainer_rank_recompute_memory.py | 20 ++++++++++++++++---- 1 file changed, 16 insertions(+), 4 deletions(-) diff --git a/dev/trainer_rank_recompute_memory.py b/dev/trainer_rank_recompute_memory.py index fa2c595e8..868109ada 100644 --- a/dev/trainer_rank_recompute_memory.py +++ b/dev/trainer_rank_recompute_memory.py @@ -98,6 +98,7 @@ def emit(row): "gdn_layers": rank._gdn_layers, "sequence_parallel": rank._sequence_parallel, "topology_dp_tp_cp_pp": rank._topology_key(), + "parallel_shape": asdict(rank._parallel_shape), "dtype": str(next(runtime.model[0].parameters()).dtype), "device": torch.cuda.get_device_name(), "device_total_bytes": torch.cuda.get_device_properties(0).total_memory, @@ -167,6 +168,13 @@ def emit(row): loss = torch.stack(terms).sum() loss.backward() torch.cuda.synchronize() + backward_peak = torch.cuda.max_memory_allocated() + backward_seconds = time.monotonic() - started + gradients = [ + p.grad + for p in rank._checkpoint_slots[slot].params + if p.grad is not None + ] emit( { **row, @@ -175,17 +183,21 @@ def emit(row): "baseline_reserved_bytes": reserved, "forward_peak_delta_bytes": forward_peak - baseline, "retained_delta_bytes": retained - baseline, - "forward_backward_peak_delta_bytes": torch.cuda.max_memory_allocated() - - baseline, + "forward_backward_peak_delta_bytes": backward_peak - baseline, "peak_reserved_bytes": torch.cuda.max_memory_reserved(), "forward_seconds": forward_seconds, - "forward_backward_seconds": time.monotonic() - started, + "forward_backward_seconds": backward_seconds, "finite_loss": bool(torch.isfinite(loss).item()), + "gradient_tensors": len(gradients), + "finite_gradients": bool(gradients) + and all( + bool(torch.isfinite(g).all().item()) for g in gradients + ), "estimate_covers_forward": check.estimated_required_bytes >= forward_peak - baseline, } ) - del outputs, loss, terms + del outputs, loss, terms, gradients rank.zero_grad() finally: dist.destroy_process_group() From d31f9423edf7c1fbdd3e125a08d254ea662377de Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 16 Sep 2026 21:57:16 +0000 Subject: [PATCH 06/15] Measure the minimum-memory admission layout before refusing --- dev/trainer_rank_recompute_memory.py | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/dev/trainer_rank_recompute_memory.py b/dev/trainer_rank_recompute_memory.py index 868109ada..fdfae3971 100644 --- a/dev/trainer_rank_recompute_memory.py +++ b/dev/trainer_rank_recompute_memory.py @@ -2,8 +2,9 @@ Run with torchrun; use a fresh process per mode/topology. Random weights keep the native model geometry and kernels without downloading a checkpoint. This -measures memory, not pretrained-model correctness. Refused plans are recorded -without execution. Each repetition clears the learned memory profile, while +measures memory, not pretrained-model correctness. Like public admission, try +the minimum-memory layout before refusing an unsplit request; never split or +bypass the budget. Each repetition clears the learned memory profile, while sample 0 includes cold compilation/autotuning. Output is one JSONL row per rank. """ @@ -127,6 +128,7 @@ def emit(row): ForwardInput(input_tokens=item, hidden_states=True) for item in tokens ] plan = rank._plan_flat_forward(requests, checkpoint=slot) + memory_minimal = False for sample in range(args.repeat): rank.zero_grad() rank._memory_profiles.clear() @@ -135,6 +137,12 @@ def emit(row): dist.barrier() torch.cuda.synchronize() check = rank._memory_check(plan) + if not check.fits and not memory_minimal: + plan = rank._plan_flat_forward( + requests, checkpoint=slot, memory_minimal=True + ) + memory_minimal = True + check = rank._memory_check(plan) row = { "lengths": lengths, "shared_prefix": prefix, @@ -142,6 +150,7 @@ def emit(row): "packed_tokens": plan.packed_tokens, "output_bytes": plan.output_bytes, "selected_max_depth": plan.selected_max_depth, + "memory_minimal": memory_minimal, "sample": sample, "admission": asdict(check), } From aef12a9e8502d181fbc614198a4bc4b438a34e9f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 16 Sep 2026 22:03:53 +0000 Subject: [PATCH 07/15] Cover eager MLP retention and calibrate GDN storage separately --- src/art/trainer_rank/_impl.py | 12 ++++++---- .../test_trainer_rank_recompute_memory.py | 24 +++++++++++++++---- 2 files changed, 27 insertions(+), 9 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index abd809ca3..3a2de945c 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -5159,23 +5159,27 @@ def _estimate_required_memory_bytes_from_values( gdn_layers = min(self._num_layers, self._gdn_layers) # Split Megatron's 9H attention/norm term into 5H of norms/residuals # (sequence parallel) and projection intermediates (tensor parallel). - # Gated MLPs retain four FFN widths. ART's column-parallel LoRA path + # Gated MLPs also retain unfused gate/up and LoRA sums: four FFN + # widths underpredicted native eager peaks; charge six in both modes + # since torch.compile can fall back to eager. GDN's projection, + # convolution and recurrent streams use a separate seven-width term + # (see dev/trainer_rank_recompute_memory.md). ART's LoRA path # additionally retains gathered H-wide inputs to attention and MLP, # even with SP. Price hybrid layers separately, not at the max width. layer_features = 5 * hidden / sp + 2 * hidden if geometry.moe_experts: # Routed/shared expert storage gets no TP/EP discount without # evidence for expert sharding and imbalanced dispatch. - layer_features += 4 * ffn_width + 2 * hidden * max( + layer_features += 6 * ffn_width + 2 * hidden * max( 0, geometry.moe_topk - 1 ) else: - layer_features += 4 * ffn_width / tp + layer_features += 6 * ffn_width / tp retained_features = ( self._num_layers * layer_features + ( (self._num_layers - gdn_layers) * (9 * attention_width - 5 * hidden) - + gdn_layers * (9 * gdn_width - 5 * hidden) + + gdn_layers * (7 * gdn_width - 5 * hidden) ) / tp ) diff --git a/tests/unit/test_trainer_rank_recompute_memory.py b/tests/unit/test_trainer_rank_recompute_memory.py index 28e2d8563..a8442e3ec 100644 --- a/tests/unit/test_trainer_rank_recompute_memory.py +++ b/tests/unit/test_trainer_rank_recompute_memory.py @@ -188,8 +188,20 @@ def test_sp_discount_excludes_gathered_lora_inputs(monkeypatch): assert tp4 > rank._memory_check(plan).estimated_required_bytes / 4 -def test_gathered_inputs_cover_attention_only_cold_peak(monkeypatch): - # Dividing the old whole estimate by TP misses this 1,439,530,496-byte peak. +@pytest.mark.parametrize( + "packed,logical,peak", + [ + (1742, 2048, 3043651584), + (6964, 8192, 11941280768), + (13928, 16384, 23861987840), + (27854, 32768, 47563188736), + ], +) +def test_gathered_inputs_cover_attention_only_cold_peak( + monkeypatch, packed, logical, peak +): + # Native eager Qwen3-1.7B TP2 paired requests at d31f9423, max across ranks + # and cold/warm repetitions. Four FFN widths missed all four peaks. rank = _rank( "selective", hidden_size=2048, @@ -200,10 +212,12 @@ def test_gathered_inputs_cover_attention_only_cold_peak(monkeypatch): sequence_parallel=True, ) monkeypatch.setattr(rank, "_topology_key", lambda: (1, 2, 1, 1)) - assert ( - rank._memory_check(_plan(rank, tokens=1024)).estimated_required_bytes - >= 1439530496 + plan = replace( + _plan(rank, tokens=packed), + logical_tokens=logical, + output_bytes=logical * 2048 * 2, ) + assert rank._memory_check(plan).estimated_required_bytes >= peak def test_non_full_estimate_covers_retained_gated_mlp_tensors() -> None: From 57f9de2f90ce44d696aaad6a1da4d99d61f17322 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 16 Sep 2026 22:17:29 +0000 Subject: [PATCH 08/15] Record sharded GPU calibration and declare legacy test recompute mode --- dev/trainer_rank_recompute_memory.csv | 106 ++++---- dev/trainer_rank_recompute_memory.md | 249 +++++++++++------- tests/unit/test_trainer_rank_active_memory.py | 4 +- tests/unit/test_trainer_rank_moe_memory.py | 4 +- .../test_trainer_rank_recompute_memory.py | 2 +- 5 files changed, 209 insertions(+), 156 deletions(-) diff --git a/dev/trainer_rank_recompute_memory.csv b/dev/trainer_rank_recompute_memory.csv index 293ab2665..d7e539aeb 100644 --- a/dev/trainer_rank_recompute_memory.csv +++ b/dev/trainer_rank_recompute_memory.csv @@ -1,55 +1,51 @@ -model,tp,mode,logical_tokens,packed_tokens,estimate_bytes,available_min_bytes,measured_rank_samples,refused_rank_samples,forward_peak_max_bytes,retained_max_bytes,forward_backward_peak_max_bytes,baseline_min_bytes,baseline_max_bytes,finite_loss_all,estimate_covers_forward_all -Qwen/Qwen3-1.7B,2,full,1024,1024,55364812,140941082317,4,0,177251840,133205504,231861248,1781843968,1883032576,True,False -Qwen/Qwen3-1.7B,2,full,4096,4096,221459251,140938985165,4,0,415402496,264382976,587698176,1883032576,1885129728,True,False -Qwen/Qwen3-1.7B,2,full,16384,16384,885837004,140930596557,4,0,1661798912,1057720832,2099318784,1885129728,1893518336,True,False -Qwen/Qwen3-1.7B,2,none,1024,1024,2717489561,140941082317,4,0,1440448000,1427860480,1473997312,1781843968,1883032576,True,True -Qwen/Qwen3-1.7B,2,none,4096,4096,10869958246,140938985165,4,0,5493352960,5443004416,5610778112,1883032576,1885129728,True,True -Qwen/Qwen3-1.7B,2,none,16384,16384,43479832985,140930596557,4,0,21973600768,21772208128,22443298304,1885129728,1893518336,True,True -Qwen/Qwen3-1.7B,2,selective,1024,1024,2717489561,140941082317,4,0,1439530496,1426943488,1535994880,1781843968,1883032576,True,True -Qwen/Qwen3-1.7B,2,selective,4096,4096,10869958246,140938985165,4,0,5489682944,5439334912,5607108608,1883032576,1885129728,True,True -Qwen/Qwen3-1.7B,2,selective,16384,16384,43479832985,140930596557,4,0,21958920704,21757528576,22428618752,1885129728,1893518336,True,True -Qwen/Qwen3.8-27B,1,full,1024,1024,196083712,89386251981,2,0,973403648,759490048,1198017536,54352220160,54452884480,True,False -Qwen/Qwen3.8-27B,1,full,2048,2048,392167424,89384154829,2,0,1812587008,1384759808,2092073472,54452884480,54452884480,True,False -Qwen/Qwen3.8-27B,1,full,3072,3072,588251136,89384154829,2,0,2718884352,2077143552,3134828032,54452884480,54452884480,True,False -Qwen/Qwen3.8-27B,1,full,4096,4096,784334848,89384154829,2,0,3625185792,2769531392,4178635264,54452884480,54452884480,True,False -Qwen/Qwen3.8-27B,1,full,38443,32710,6328148992,89321231565,2,0,28953189376,22178822144,33419714048,54452884480,54452893184,True,False -Qwen/Qwen3.8-27B,1,none,1024,1024,31311108505,89386251981,2,0,17041559040,16955570688,17060460544,54352220160,54452884480,True,True -Qwen/Qwen3.8-27B,1,none,2048,2048,62622217011,89386251981,2,0,33761374720,33589400064,33799116800,54452884480,54452884480,True,True -Qwen/Qwen3.8-27B,1,none,3072,3072,93933325516,89386251981,0,1,,,,,,, -Qwen/Qwen3.8-27B,1,none,4096,4096,125244434022,89386251981,0,1,,,,,,, -Qwen/Qwen3.8-27B,1,none,38443,32710,1000246567936,89386251981,0,1,,,,,,, -Qwen/Qwen3.8-27B,1,selective,1024,1024,31311108505,87997937357,2,0,17043131904,16957146624,17124951040,54352220160,54452884480,True,True -Qwen/Qwen3.8-27B,1,selective,2048,2048,62622217011,87997937357,2,0,33758228992,33586258432,33795975168,54452884480,54452884480,True,True -Qwen/Qwen3.8-27B,1,selective,3072,3072,93933325516,87997937357,0,1,,,,,,, -Qwen/Qwen3.8-27B,1,selective,4096,4096,125244434022,87997937357,0,1,,,,,,, -Qwen/Qwen3.8-27B,1,selective,38443,32710,1000246567936,87997937357,0,1,,,,,,, -Qwen/Qwen3.8-27B,2,full,1024,1024,196083712,115395819213,4,0,526710272,419749376,710319104,27327631360,27428295680,True,False -Qwen/Qwen3.8-27B,2,full,2048,2048,392167424,115393722061,4,0,917103104,703181312,1098500608,27428295680,27428295680,True,False -Qwen/Qwen3.8-27B,2,full,3072,3072,588251136,115393722061,4,0,1375658496,1054775808,1646584832,27428295680,27428295680,True,False -Qwen/Qwen3.8-27B,2,full,4096,4096,784334848,115393722061,4,0,1834217984,1406374400,2194869760,27428295680,27428295680,True,False -Qwen/Qwen3.8-27B,2,full,38443,32710,6328148992,115391616205,4,0,14653886976,11294938624,17587685376,27428295680,27428304384,True,False -Qwen/Qwen3.8-27B,2,none,1024,1024,31311108505,115395819213,4,0,10079153664,10041397760,10146256896,27327631360,27428295680,True,True -Qwen/Qwen3.8-27B,2,none,2048,2048,62622217011,115395819213,4,0,19542929920,19467420160,19677136896,27428295680,27428295680,True,True -Qwen/Qwen3.8-27B,2,none,3072,3072,93933325516,115395819213,4,0,29394684416,29281420800,29595995136,27428295680,27428295680,True,True -Qwen/Qwen3.8-27B,2,none,4096,4096,125244434022,115395819213,0,2,,,,,,, -Qwen/Qwen3.8-27B,2,none,38443,32710,1000246567936,115395819213,0,2,,,,,,, -Qwen/Qwen3.8-27B,2,selective,1024,1024,31311108505,115395819213,4,0,10082561536,10044808704,10149667840,27327631360,27428295680,True,True -Qwen/Qwen3.8-27B,2,selective,2048,2048,62622217011,115395819213,4,0,19541553664,19466048000,19675764736,27428295680,27428295680,True,True -Qwen/Qwen3.8-27B,2,selective,3072,3072,93933325516,115393722061,4,0,29392325120,29279066624,29593640960,27428295680,27428295680,True,True -Qwen/Qwen3.8-27B,2,selective,4096,4096,125244434022,115393722061,0,2,,,,,,, -Qwen/Qwen3.8-27B,2,selective,38443,32710,1000246567936,115393722061,0,2,,,,,,, -Qwen/Qwen3.8-27B,4,full,1024,1024,196083712,128855946957,8,0,308082176,254597632,455984128,13865406464,13966070784,True,False -Qwen/Qwen3.8-27B,4,full,2048,2048,392167424,128855946957,8,0,480371200,373402112,591982080,13966070784,13966070784,True,False -Qwen/Qwen3.8-27B,4,full,3072,3072,588251136,128853849805,8,0,708239872,545689088,873627648,13966070784,13966070784,True,False -Qwen/Qwen3.8-27B,4,full,4096,4096,784334848,128853849805,8,0,939258368,725320192,1161568256,13966070784,13966070784,True,False -Qwen/Qwen3.8-27B,4,full,38443,32712,6328509440,128851743949,8,0,7502990336,5851191296,9831958528,13966070784,13966079488,True,False -Qwen/Qwen3.8-27B,4,none,1024,1024,31311108505,128855946957,8,0,5736394240,5722230272,5827089408,13865406464,13966070784,True,True -Qwen/Qwen3.8-27B,4,none,2048,2048,62622217011,128855946957,8,0,11131212288,11103934976,11313651712,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,none,3072,3072,93933325516,128855946957,8,0,16653462016,16610974208,16925548544,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,none,4096,4096,125244434022,128855946957,8,0,21989044736,21934492160,22353924096,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,none,38443,32712,1000307699916,128855946957,0,4,,,,,,, -Qwen/Qwen3.8-27B,4,selective,1024,1024,31311108505,127666861773,8,0,5732331008,5718170112,5932531200,13865406464,13966070784,True,True -Qwen/Qwen3.8-27B,4,selective,2048,2048,62622217011,127666861773,8,0,11132523008,11105249792,11314966528,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,selective,3072,3072,93933325516,127666861773,8,0,16651233792,16608751104,16923325440,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,selective,4096,4096,125244434022,127666861773,8,0,21987471872,21932925440,22352357376,13966070784,13966070784,True,True -Qwen/Qwen3.8-27B,4,selective,38443,32712,1000307699916,127666861773,0,4,,,,,,, +phase,source_sha,driver_sha256,model,layers,gdn_layers,tp,ep,etp,sequence_parallel,compiled,mode,modules,lengths,shared_prefix,logical_tokens,packed_tokens,selected_max_depth,memory_minimal,estimate_bytes,available_min_bytes,measured_rank_samples,refused_rank_samples,forward_peak_max_bytes,cold_forward_peak_max_bytes,warm_forward_peak_max_bytes,retained_max_bytes,forward_backward_peak_max_bytes,baseline_min_bytes,baseline_max_bytes,finite_loss_all,finite_gradients_all,estimate_covers_forward_all,min_estimate_to_forward_ratio +eager_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,1024+1024,307,2048,1742,2,False,2756291788,140940714701,4,0,3043651584,3043651584,2975649792,3036479488,3091428352,1781843968,1883400192,True,True,False,0.9055871580339203 +eager_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,16384+16384,4915,32768,27854,2,False,44072283340,140914919629,4,0,47563188736,47563188736,47548090368,47448521216,48253829632,1894096896,1909195264,True,True,False,0.9266048915396251 +eager_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,4096+4096,1228,8192,6964,2,False,11018859315,140937149133,4,0,11941280768,11941280768,11937715200,11912611840,12113940480,1883400192,1886965760,True,True,False,0.9227535579372788 +eager_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,8192+8192,2457,16384,13928,2,False,22037718630,140930017997,4,0,23861987840,23861987840,23854856704,23804649984,24207305216,1886965760,1894096896,True,True,False,0.9235491518044459 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,1024+1024,307,2048,1742,2,False,3415587225,140940714701,4,0,3043651584,3043651584,2975649792,3036479488,3091428352,1781843968,1883400192,True,True,True,1.1222004657021873 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,16384+16384,4915,32768,27854,2,False,54614197862,140914919629,4,0,47563188736,47563188736,47548090368,47448521216,48253829632,1894096896,1909195264,True,True,True,1.1482450885523403 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,4096+4096,1228,8192,6964,2,False,13654527180,140937149133,4,0,11941280768,11941280768,11937715200,11912611840,12113940480,1883400192,1886965760,True,True,True,1.1434725843304114 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,8192+8192,2457,16384,13928,2,False,27309054361,140930017997,4,0,23861987840,23861987840,23854856704,23804649984,24207305216,1886965760,1894096896,True,True,True,1.1444584811673426 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,none,,1024+1024,307,2048,1742,2,False,3415587225,140940714701,4,0,2406920192,2406920192,2338918400,2386762240,2441711104,1781843968,1883400192,True,True,True,1.4190695796032442 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,none,,16384+16384,4915,32768,27854,2,False,54614197862,140914919629,4,0,37406161408,37406161408,37391063040,37083114496,37888422912,1894096896,1909195264,True,True,True,1.4600321392592757 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,none,,4096+4096,1228,8192,6964,2,False,13654527180,140937149133,4,0,9424284672,9424284672,9420719104,9343731200,9545059840,1883400192,1886965760,True,True,True,1.4488661638764215 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,none,,8192+8192,2457,16384,13928,2,False,27309054361,140930017997,4,0,18819614208,18819614208,18812483072,18658097152,19060752384,1886965760,1894096896,True,True,True,1.4510953337922963 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,selective,core_attn,1024+1024,307,2048,1742,2,False,3415587225,140940714701,4,0,2405357568,2405357568,2337355776,2385200128,2440148992,1781843968,1883400192,True,True,True,1.419991468395272 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,selective,core_attn,16384+16384,4915,32768,27854,2,False,54614197862,140914919629,4,0,37381202432,37381202432,37366104064,37058156032,37863464448,1894096896,1909195264,True,True,True,1.4610069850307377 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,selective,core_attn,4096+4096,1228,8192,6964,2,False,13654527180,140937149133,4,0,9418034176,9418034176,9414468608,9337481216,9538809856,1883400192,1886965760,True,True,True,1.4498277373845028 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,selective,core_attn,8192+8192,2457,16384,13928,2,False,27309054361,140930017997,4,0,18807127552,18807127552,18799996416,18645611008,19048266240,1886965760,1894096896,True,True,True,1.452058762588436 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn,1024+1024,307,2048,2048,1,False,18095066316,124303160525,8,0,4869984768,4869984768,4582460416,4825250304,4876046848,17555830784,17929557504,True,True,True,3.715630988191214 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn,2048+2048,614,4096,4096,1,False,36190132633,124122997965,8,0,8887517184,8887517184,8857254912,8802950144,8903695360,17928215552,18065679872,True,True,True,4.072018302046414 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn,4096+4096,1228,8192,6964,2,False,61535826739,123840261837,8,0,15021402112,15021402112,14957089792,14845153792,15046482432,18065558016,18281307136,True,True,True,4.096543470455496 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn+moe,1024+1024,307,2048,2048,1,False,18095066316,124135796941,8,0,2815510016,2815510016,2317386752,2735823360,2893087232,17555830784,18096921088,True,True,True,6.426923084332583 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn+moe,2048+2048,614,4096,4096,1,False,36190132633,123792823501,8,0,4870201856,4870201856,4568279552,4702040064,4952712704,18094838272,18395854336,True,True,True,7.430930729989028 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn+moe,4096+4096,1228,8192,6964,2,False,61535826739,123262782157,8,0,8278291968,8278291968,7798616576,7982971392,8422111744,18392608256,18856689664,True,True,True,7.433396523928932 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-9B,32,24,2,1,1,True,True,selective,core_attn,3072+3072,1843,6144,4302,2,False,24865721548,133630881485,4,0,15198974976,15198974976,15131242496,15108335616,15410327552,9090471936,9191136256,True,True,True,1.6360130592532927 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-9B,32,24,2,1,1,True,True,selective,core_attn,6144+6144,3686,12288,8602,2,False,49719908761,133630881485,4,0,30070096896,30070096896,30070068224,29888892928,30492874752,9191136256,9191136256,True,True,True,1.6534668622106723 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,selective,core_attn,19221+19222,5733,38443,32712,2,True,282826872627,128853849805,0,4,,,,,,,,,,, +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,selective,core_attn,2048+2048,614,4096,4096,1,False,35405797785,128855946957,8,0,22132520448,22132520448,22065410560,22076925440,22328585728,13865406464,13966070784,True,True,True,1.5997182909278385 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,selective,core_attn,4096+4096,1228,8192,6964,2,False,60210603622,128853849805,8,0,37665510400,37665510400,37665510400,37582912512,38086231040,13966070784,13966070784,True,True,True,1.5985606721527394 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,selective,core_attn,8192+8192,2457,16384,13928,2,False,120421207244,128853849805,8,0,75080549888,75080549888,75080549888,74915458048,75922093056,13966070784,13966070784,True,True,True,1.6038935173441866 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,none,,19221+19222,5733,38443,32712,2,True,282826872627,128853849805,0,4,,,,,,,,,,, +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,none,,2048+2048,614,4096,4096,1,False,35405797785,128855946957,8,0,22094640640,22094640640,22027530752,22040088064,22291748352,13865406464,13966070784,True,True,True,1.6024609027087575 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,none,,4096+4096,1228,8192,6964,2,False,60210603622,128853849805,8,0,37677824512,37677824512,37677824512,37598091264,38101409792,13966070784,13966070784,True,True,True,1.5980382201425547 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,none,,8192+8192,2457,16384,13928,2,False,120421207244,128853849805,8,0,74964444672,74964444672,74964444672,74802887168,75809522176,13966070784,13966070784,True,True,True,1.6063776337021085 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn,19221+19222,5733,38443,32712,2,True,282826872627,128853849805,0,4,,,,,,,,,,, +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn,2048+2048,614,4096,4096,1,False,35405797785,128855946957,8,0,22093067776,22093067776,22025957888,22038521344,22290181632,13865406464,13966070784,True,True,True,1.6025749861439251 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn,4096+4096,1228,8192,6964,2,False,60210603622,128853849805,8,0,37675145728,37675145728,37675145728,37595429376,38098747904,13966070784,13966070784,True,True,True,1.5981518441016076 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn,8192+8192,2457,16384,13928,2,False,120421207244,128853849805,8,0,74959095296,74959095296,74959095296,74797569024,75804204032,13966070784,13966070784,True,True,True,1.606492271131052 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn+mlp,19221+19222,5733,38443,32712,2,True,282826872627,128853849805,0,4,,,,,,,,,,, +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn+mlp,2048+2048,614,4096,4096,1,False,35405797785,128855946957,8,0,12039591424,12039591424,11972481536,11836139008,12136017408,13865406464,13966070784,True,True,True,2.940780674203052 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn+mlp,4096+4096,1228,8192,6964,2,False,60210603622,128853849805,8,0,20407347712,20407347712,20407347712,20073009664,20576328192,13966070784,13966070784,True,True,True,2.9504374831911524 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn+mlp,8192+8192,2457,16384,13928,2,False,120421207244,128853849805,8,0,40674358272,40674358272,40674358272,40005276672,41011911680,13966070784,13966070784,True,True,True,2.9606172625690146 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,none,,16384+16384,4915,32768,27856,2,True,140687035596,135698216653,0,8,,,,,,,,,,, +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,none,,19221+19222,5733,38443,32712,2,True,165211897241,135698216653,0,8,,,,,,,,,,, +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,none,,2048+2048,614,4096,4096,1,False,20678757580,135700313805,16,0,15596174336,15596174336,15531161600,15554143744,15805804032,7018942464,7119606784,True,True,True,1.3258865369482364 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,none,,4096+4096,1228,8192,6968,2,False,35191907942,135698216653,16,0,26428632064,26428632064,26428632064,26357117440,26860435968,7119606784,7119606784,True,True,True,1.3315826508454434 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,none,,8192+8192,2457,16384,13928,2,False,70343517798,135698216653,16,0,52559443456,52559443456,52559443456,52416501248,53423136256,7119606784,7119606784,True,True,True,1.338361161622419 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,selective,core_attn,16384+16384,4915,32768,27856,2,True,140687035596,135698216653,0,8,,,,,,,,,,, +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,selective,core_attn,19221+19222,5733,38443,32712,2,True,165211897241,135698216653,0,8,,,,,,,,,,, +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,selective,core_attn,2048+2048,614,4096,4096,1,False,20678757580,135700313805,16,0,15595388416,15595388416,15530375680,15553363456,15805023744,7018942464,7119606784,True,True,True,1.3259533541841604 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,selective,core_attn,4096+4096,1228,8192,6968,2,False,35191907942,135698216653,16,0,26427289088,26427289088,26427289088,26355792384,26859110912,7119606784,7119606784,True,True,True,1.3316503189114393 +final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,selective,core_attn,8192+8192,2457,16384,13928,2,False,70343517798,135698216653,16,0,52556765184,52556765184,52556765184,52413853696,53420488704,7119606784,7119606784,True,True,True,1.338429363978719 diff --git a/dev/trainer_rank_recompute_memory.md b/dev/trainer_rank_recompute_memory.md index a9dfe2568..d13f6c130 100644 --- a/dev/trainer_rank_recompute_memory.md +++ b/dev/trainer_rank_recompute_memory.md @@ -1,106 +1,159 @@ -# Recompute memory calibration, 2026-09-16 - -First GPU campaign for #915, measuring estimator commit -`22f628d6a8e87d18b28eb9d1d61579520d012d23`. The new selective/no-recompute -estimate covered every executed forward. It remains conservative, particularly -at larger TP sizes. The unchanged full-recompute estimate underestimated every -measured shape; these results do not validate that legacy estimate. - -[CSV evidence](trainer_rank_recompute_memory.csv) contains 54 model/topology/mode/ -shape cells. Each row aggregates both repetitions and all ranks: peaks are maxima, -available memory is the minimum, and refused cells have blank measurement fields. -There were 202 measured rank-samples and 22 refused rank-samples. All measured -losses were finite, and every backward completed without CUDA OOM. - -## Method - -- NVIDIA H200, PyTorch 2.11.0+cu128, CUDA 12.8, bf16, compiled transformer layers. -- Native Qwen3.8-27B (64 layers) at TP1/2/4, and attention-only Qwen3-1.7B at TP2. - DP/CP/PP are 1; active LoRA rank is 1. Models and adapters use random weights. -- Each mode runs in a fresh process: `full/uniform/1`, selective `core_attn`, or - no recompute. Every sample clears the learned memory profile. The first sample - includes any compilation/autotuning required by that shape; the second is warm. -- Requests retain hidden states with gradients. The driver executes the selected - unsplit native plan only when its cold admission check passes, then backpropagates - a mean-square hidden-state loss. No admission bypass, split, or optimizer step. -- Estimates and measured peaks below are **incremental allocated GiB above the - pre-forward baseline**, not total GPU usage. Forward/backward peaks are separate - CSV fields; backward includes the diagnostic loss. Allocator reservation is not - the measured allocation peak. -- GDN lengths: 1,024 / 2,048 / 3,072 / 4,096, plus two synthetic sibling sequences - of 19,221 and 19,222 tokens sharing 5,733 prefix tokens. The siblings reproduce - #913's 38,443 logical / 32,710 packed tokens (32,712 after TP4 padding), not its - original token contents. Attention-only lengths: 1,024 / 4,096 / 16,384. -- TP2 ran locally. SkyPilot jobs `art-915-recompute-tp1-0916:1` and - `art-915-recompute-tp4-0916:1` both succeeded on free `k8s/cks-wb3` capacity. - Both clusters had ten-minute autodown and two-hour pod deadlines, and were - explicitly torn down after evidence retrieval. - -## Results - -Qwen3.8-27B, 2,048 tokens: - -| TP | New estimate | Selective forward peak | No-recompute forward peak | -| -- | --: | --: | --: | -| 1 | 58.321 | 31.440 | 31.443 | -| 2 | 58.321 | 18.199 | 18.201 | -| 4 | 58.321 | 10.368 | 10.367 | - -The long sibling group was refused in both non-full modes at all three TP sizes. -Its estimate is 931.552 GiB at TP1/2 and 931.609 GiB at TP4. These are admission -results, not measured non-full peaks. Smaller false refusals remain possible: -TP1 declined 3,072 and 4,096 tokens; TP2 declined 4,096 tokens. - -Full recompute completed the long sibling group: - -| TP | Legacy estimate | Forward peak | -| -- | --: | --: | -| 1 | 5.894 | 26.965 | -| 2 | 5.894 | 13.647 | -| 4 | 5.894 | 6.988 | - -Qwen3-1.7B, TP2: - -| Tokens | New estimate | Selective forward peak | No-recompute forward peak | -| --: | --: | --: | --: | -| 1,024 | 2.531 | 1.341 | 1.342 | -| 4,096 | 10.123 | 5.113 | 5.116 | -| 16,384 | 40.494 | 20.451 | 20.465 | - -## Remaining work - -Keep #915 in draft. This campaign supports the retained-activation floor for the -tested dense attention/GDN shapes; it is not a general memory bound. A universal -division by TP would already underestimate the attention-only cold 1,024-token -case (2.531 / 2 < 1.341 GiB). A tighter model must separate replicated activations -and fixed workspace from sharded storage, then validate on held-out shapes. - -Full recompute needs a separate correction for retained checkpoint boundaries and -workspace; preserving its old heuristic does not make it calibrated. MoE/EP/CP, -other full-recompute intervals, optional selective modules, deep prefix trees, -mixed gradient groups, other LoRA ranks, and pretrained correctness are untested. -This commit adds measurement evidence without changing the estimator coefficients. - -## Reproduce - -After `INSTALL_VLLM_RUNTIME=false bash src/art/megatron/setup.sh`, run on two GPUs: +# Recompute memory calibration, September 16, 2026 + +The sharded retained-activation estimate admits the requested Qwen3.8-27B +selective TP4 series (two sequences of 2k, 4k, and 8k tokens), bounds measured +forward peaks in compiled and eager execution, and refuses #913's original long +pair before execution. It accounts for the actual 48 GDN / 16 attention layers. + +The [CSV](trainer_rank_recompute_memory.csv) records final estimator commit +`aef12a9e8502d181fbc614198a4bc4b438a34e9f`, plus an earlier eager-execution failure +as a negative control. Each row aggregates all ranks and two repetitions for one +model/topology/mode/shape. Peaks are maxima; available memory is the minimum; +refusals have blank measurement fields. Source and driver hashes are per row. + +## Requested TP4 calibration + +All values are **incremental allocated GiB above the pre-forward baseline**, +not total GPU usage or reserved memory. Model/adaptor weights are random; this +measures the native runtime's memory behavior, not pretrained correctness. + +Qwen3.8-27B, 64 layers, bf16, LoRA rank 1, selective `core_attn`, SP enabled: + +| Tokens per sequence | Packed tokens | Estimate | Compiled forward | Eager forward | Compiled forward + backward | +| --: | --: | --: | --: | --: | --: | +| 2,048 | 4,096 | 32.974 | 20.576 | 20.613 | 20.759 | +| 4,096 | 6,964 | 56.075 | 35.088 | 35.079 | 35.482 | +| 8,192 | 13,928 | 112.151 | 69.811 | 69.924 | 70.598 | + +Each pair shares 30% of its prefix. The planner chose an unshared layout for the +2k pair and shared layouts for 4k/8k; token counts include TP padding. The compiled +run's pre-forward allocation was at most 13.007 GiB, with a usable incremental +budget of at least 120.004 GiB. The 8k prediction also fits the issue's 119.289 GiB +budget. No-recompute peaks were 20.577 / 35.090 / 69.816 GiB under the same estimates. + +The original synthetic sibling geometry (19,221 + 19,222 logical tokens, sharing +5,733 prefix tokens) is refused at TP4 with a 263.403 GiB estimate for 32,712 packed +tokens. This is an admission result, **not a measured long-pair selective peak**. + +## What changed in the estimate + +For gradient-enabled non-full recompute, the retained floor is the packed token +count times dtype bytes times the sum of layer storage. Dense layers retain +`5H/SP + 2H + 6F/TP`, plus `(9A - 5H)/TP` for each attention layer or +`(7G - 5H)/TP` for each GDN layer. Here `H` is hidden width, `F` is FFN width, +`A = max(H, heads * head_dim)`, `G = max(H, 2*key_width + 2*value_width)`, and +`SP` is TP when sequence parallel is enabled, otherwise 1. + +- The norm/residual term follows the SP distinction in + [Megatron's activation model](https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/training/theoretical_memory_usage.py). + Projection and dense MLP storage receive TP discounts. The separate `2H` + remains replicated because ART's LoRA wrappers retain gathered attention and + MLP inputs even with SP (`return_layernorm_output_gathered` and + `_column_parallel_lora_input` in `src/art/megatron/lora.py`). +- Four FFN widths covered compiled attention-only runs but underestimated eager + execution. Six cover the measured eager gate/up activations and LoRA sums. + Both execution modes use six: compilation can fall back to eager. +- GDN uses its own seven-width calibrated envelope for projection, convolution, + and recurrent storage. Charge it only to actual GDN layers, instead of charging + every layer the maximum of attention and GDN widths. This is an empirical + envelope for the measured native paths, not an exact tensor-liveness proof. +- Routed/shared expert FFN and dispatch storage receive no TP/EP/ETP discount. + Optional `mlp`/`moe` recomputation receives no discount in this revision. +- The existing static estimate and MoE FC2 estimate remain floors. Learned + profiles can only raise the estimate; output bytes and the existing 10% safety + margin are still applied. No-grad and full-recompute paths are unchanged. + +## Additional evidence + +Qwen3.8-27B selective TP8, compiled: + +| Tokens per sequence | Estimate | Forward peak | +| --: | --: | --: | +| 2,048 | 19.259 | 14.524 | +| 4,096 | 32.775 | 24.612 | +| 8,192 | 65.513 | 48.947 | + +No-recompute peaks were 14.525 / 24.614 / 48.950 GiB. The final estimate refuses +16k pairs (131.025 GiB) and the original long pair (153.866 GiB) at TP8. Those +selective peaks remain unmeasured; TP sharding does not remove gathered storage, +and the estimate can still over-refuse near capacity. + +The eager attention-only negative control is Qwen3-1.7B TP2, pairs of +1,024 / 4,096 / 8,192 / 16,384 tokens. Before the six-FFN correction, predictions +were 2.567 / 10.262 / 20.524 / 41.046 GiB against observed +2.835 / 11.121 / 22.223 / 44.297 GiB. Re-running the final code produced the same +peaks with estimates 3.181 / 12.717 / 25.434 / 50.863 GiB. The smallest observed +headroom in the final campaign is 12.2%. Recorded peaks are regression witnesses: +restoring the old estimator fails all four, and dropping the gathered-input term +fails three. + +A held-out Qwen3.5-9B model at TP2 used pairs of 3,072 / 6,144 tokens with a 60% +shared prefix, after coefficients were fixed. Estimates 23.158 / 46.305 GiB covered +forward peaks 14.155 / 28.005 GiB and forward/backward peaks 14.352 / 28.399 GiB. + +Adding `mlp` to selective recomputation on the 27B at TP4 reduced the 2k / 4k / +8k forward peaks to 11.213 / 19.006 / 37.881 GiB, under the same estimates. +Qwen3.5-35B-A3B at TP4/EP4/ETP1 used 1k / 2k / 4k pairs: estimates +16.852 / 33.705 / 57.310 GiB covered default selective peaks +4.536 / 8.277 / 13.990 GiB. Adding `moe` reduced them to +2.622 / 4.536 / 7.710 GiB. These runs support keeping the current undiscounted +floor; they do not establish per-module discounts or bounds on pretrained routing +imbalance. Compiled attention-only selective and no-recompute controls also pass. + +The final campaign contains **46 cells: 296 measured rank-samples and 48 refused +rank-samples**. Every measured forward was covered, every backward completed +without CUDA OOM, and every loss and adapter gradient was finite. The CSV also +includes four cells / 16 rank-samples from the earlier eager negative control; +those failed coverage checks are intentionally preserved, not final-code failures. + +## Method, reproduction, and limits + +NVIDIA H200, PyTorch 2.11.0+cu128, CUDA 12.8, bf16, DP/CP/PP=1. Each mode runs in a +fresh process. Every repetition clears learned memory profiles, gradients, and +unused cached allocations. Sample 0 includes any first-execution compilation and +autotuning; sample 1 is warm. Both contribute to reported maxima. + +The driver plans paired hidden-state requests with gradients, tries the +minimum-memory layout if admission refuses, and executes the unsplit native plan +only if it fits. It then backpropagates a mean-square hidden-state loss. There is +no admission bypass, splitting, or optimizer step. Forward/backward peaks include +the diagnostic loss; finite-gradient checks happen after peak/timing collection. +Raw JSONL also records geometry, every rank/sample, timing, and allocator totals. +Driver SHA-256: `42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788`. + +On a configured GPU machine, after +`INSTALL_VLLM_RUNTIME=false bash src/art/megatron/setup.sh`: ```sh -ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE=2 \ +ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE=4 \ ART_MEGATRON_CONTEXT_PARALLEL_SIZE=1 \ ART_MEGATRON_DATA_PARALLEL_SIZE=1 \ ART_MEGATRON_PIPELINE_MODEL_PARALLEL_SIZE=1 \ uv run --project megatron_runtime --no-sync python -m torch.distributed.run \ - --standalone --nproc-per-node=2 dev/trainer_rank_recompute_memory.py \ - --mode selective --tokens 1024 2048 3072 4096 --reported-pair \ - --evidence scratch/recompute-memory/tp2-selective.jsonl + --standalone --nproc-per-node=4 dev/trainer_rank_recompute_memory.py \ + --mode selective --pairs --tokens 2048 4096 8192 --reported-pair \ + --evidence scratch/recompute-memory/tp4-selective.jsonl ``` -Repeat with `--mode none` and `--mode full`. For the attention-only control, add -`--model Qwen/Qwen3-1.7B --tokens 1024 4096 16384` and omit `--reported-pair`. -The [SkyPilot task](trainer_rank_recompute_memory.sky.yaml) runs all three modes -at TP4; pass the synced commit as `ART_CALIBRATION_SOURCE_SHA` and launch with -`--idle-minutes-to-autostop 10 --down`. Override both GPU count and TP environment -variable for TP1. The driver emits per-rank JSONL and identical rows in job logs. -Driver SHA-256 for this campaign: -`33a7291041913b60232ac2e287edf26c7766b7ab268af034dbcd71df0eae840e`. +Use `--mode none`, `ART_DISABLE_MEGATRON_COMPILE=1`, or `--modules core_attn mlp` +for the corresponding controls. Change both process count and TP for TP8. +The [SkyPilot task](trainer_rank_recompute_memory.sky.yaml) defaults to TP4 on +free Kubernetes; pass the synced commit as `ART_CALIBRATION_SOURCE_SHA` and use +`--idle-minutes-to-autostop 15 --down`. + +Final cluster jobs: `art-915-sharded-tp4-0916:4` (selective/none), `:5` (eager), +`:6` (MLP recompute), `:7` (MoE), `:8` (MoE recompute), all on free `k8s/cks-wb3`; +`art-915-sharded-tp8-ext-0916:2` on free `k8s/ext-collab2`. TP2 controls ran locally. +The clusters used 15-minute autodown and two-hour pod deadlines. Raw evidence is +retained locally under `scratch/recompute-memory-final/`. No paid clusters were +used; both clusters were explicitly torn down after successful jobs and evidence +retrieval. + +The [first unsharded campaign](https://github.com/OpenPipe/ART/blob/be21a2599/dev/trainer_rank_recompute_memory.csv) +remains historical evidence. It also found underestimation in the unchanged full +recompute heuristic (for the long group, 5.894 GiB estimated versus 13.647 GiB +observed at TP2). This PR does not validate or fix that legacy path. Larger LoRA +ranks, pretrained routing imbalance, CP, deeper prefix trees, mixed gradient +groups, and other kernels/hardware are not calibrated by these measurements. +Conservative expert/module pricing remains deliberate; these results support +removing the selective guard for the measured paths, not a universal memory bound. diff --git a/tests/unit/test_trainer_rank_active_memory.py b/tests/unit/test_trainer_rank_active_memory.py index 74ced9d63..d9c1d710a 100644 --- a/tests/unit/test_trainer_rank_active_memory.py +++ b/tests/unit/test_trainer_rank_active_memory.py @@ -34,7 +34,9 @@ def _rank(): SimpleNamespace( model=[_Model()], optimizer=None, - provider=SimpleNamespace(hidden_size=8, num_layers=4), + provider=SimpleNamespace( + hidden_size=8, num_layers=4, recompute_granularity="full" + ), model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), ), ) diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index e0dabad67..ccb1ed6a2 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -75,7 +75,9 @@ def _rank(layer=None): SimpleNamespace( model=[model], optimizer=None, - provider=SimpleNamespace(hidden_size=2048, num_layers=40), + provider=SimpleNamespace( + hidden_size=2048, num_layers=40, recompute_granularity="full" + ), model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), ), ) diff --git a/tests/unit/test_trainer_rank_recompute_memory.py b/tests/unit/test_trainer_rank_recompute_memory.py index a8442e3ec..63a6d9784 100644 --- a/tests/unit/test_trainer_rank_recompute_memory.py +++ b/tests/unit/test_trainer_rank_recompute_memory.py @@ -153,7 +153,7 @@ def _hybrid_rank(monkeypatch: pytest.MonkeyPatch, tp: int) -> TrainerRank: ) def test_sharded_floor_covers_recorded_native_gdn_peaks(monkeypatch, tp, peak): # H200, 64-layer Qwen3.8-27B, LoRA r1, SP, cold/warm max, 2,048 tokens. - # Source evidence: dev/trainer_rank_recompute_memory.csv at 22f628d6. + # First unsharded campaign, linked from dev/trainer_rank_recompute_memory.md. rank = _hybrid_rank(monkeypatch, tp) assert rank._memory_check(_plan(rank, tokens=2048)).estimated_required_bytes >= peak From d8358df387648ba7c2cc56a9ad9b16f63487211f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 17 Sep 2026 00:02:25 +0000 Subject: [PATCH 09/15] Attribute eager component retention during GPU calibration --- dev/trainer_rank_recompute_memory.py | 59 ++++++++++++++++++++++++++++ 1 file changed, 59 insertions(+) diff --git a/dev/trainer_rank_recompute_memory.py b/dev/trainer_rank_recompute_memory.py index fdfae3971..4ab82cd2e 100644 --- a/dev/trainer_rank_recompute_memory.py +++ b/dev/trainer_rank_recompute_memory.py @@ -36,6 +36,9 @@ def main() -> None: ) parser.add_argument("--prefix-fraction", type=float, default=0.3) parser.add_argument("--modules", nargs="+", default=["core_attn"]) + parser.add_argument( + "--components", action="store_true", help="Attribute eager layer allocations" + ) parser.add_argument("--evidence", type=Path, required=True) args = parser.parse_args() if args.repeat < 1 or args.layers < 0 or any(n < 1 for n in args.tokens): @@ -78,6 +81,8 @@ def emit(row): ) for chunk in runtime.model: chunk.train() + if args.components and runtime.transformer_layers_compiled: + parser.error("--components requires ART_DISABLE_MEGATRON_COMPILE=1") rank = TrainerRank(runtime) [slot] = load_random_checkpoints( runtime, rank, 1, base_model=args.model, lora_rank=1 @@ -106,6 +111,17 @@ def emit(row): "torch": torch.__version__, "cuda": torch.version.cuda, "transformer_layers_compiled": runtime.transformer_layers_compiled, + "activation_config": { + name: str(getattr(runtime.provider, name, None)) + for name in ( + "bias_activation_fusion", + "use_te_activation_func", + "gated_linear_unit", + "activation_func", + "attention_output_gate", + "qk_layernorm", + ) + }, } args.evidence.parent.mkdir(parents=True, exist_ok=True) workloads = [ @@ -161,9 +177,51 @@ def emit(row): reserved = torch.cuda.memory_reserved() torch.cuda.reset_peak_memory_stats() started = time.monotonic() + components, handles, entries = [], [], {} + if args.components: + from megatron.core.transformer.transformer_layer import ( + TransformerLayer, + ) + + def enter(module, inputs): + entries[id(module)] = torch.cuda.memory_allocated() + + def leave(name): + def record(module, inputs, output): + components.append( + { + "name": name, + "type": type(module).__name__, + "input_shape": list(inputs[0].shape) + if inputs and isinstance(inputs[0], torch.Tensor) + else None, + "retained_delta_bytes": torch.cuda.memory_allocated() + - entries[id(module)], + } + ) + + return record + + for name, layer in runtime.model[0].named_modules(): + if isinstance(layer, TransformerLayer): + for part, module in ( + ("layer", layer), + ("attention", layer.self_attention), + ("mlp", layer.mlp), + ): + handles.extend( + ( + module.register_forward_pre_hook(enter), + module.register_forward_hook( + leave(f"{name}.{part}") + ), + ) + ) # Direct execution isolates the selected unsplit plan. It uses # the same native forward as public admission, with no split. outputs = rank._execute_flat_plan(plan) + for handle in handles: + handle.remove() torch.cuda.synchronize() forward_peak = torch.cuda.max_memory_allocated() retained = torch.cuda.memory_allocated() @@ -188,6 +246,7 @@ def emit(row): { **row, "status": "measured", + "components": components, "baseline_allocated_bytes": baseline, "baseline_reserved_bytes": reserved, "forward_peak_delta_bytes": forward_peak - baseline, From 71055f4639b2d9f247d7f3a3421a7f79cecc2c61 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 17 Sep 2026 00:12:37 +0000 Subject: [PATCH 10/15] Calibrate retained activation components and dense MLP checkpoints --- src/art/trainer_rank/_impl.py | 112 +++++++++++------- .../test_trainer_rank_recompute_memory.py | 5 +- 2 files changed, 71 insertions(+), 46 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 3a2de945c..c2bae1be1 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1380,17 +1380,29 @@ def __init__(self, runtime: TrainingRuntime) -> None: or getattr(runtime.provider, "num_layers", 1) or 1 ) + memory_config = getattr(metadata_model, "config", None) or runtime.provider self._recompute_granularity = getattr( - getattr(metadata_model, "config", None), - "recompute_granularity", - getattr(runtime.provider, "recompute_granularity", None), + memory_config, "recompute_granularity", None + ) + self._recompute_modules = frozenset( + getattr(memory_config, "recompute_modules", ()) or () ) self._sequence_parallel = bool( - getattr( - getattr(metadata_model, "config", None), - "sequence_parallel", - getattr(runtime.provider, "sequence_parallel", False), - ) + getattr(memory_config, "sequence_parallel", False) + ) + self._attention_output_gate = bool( + getattr(memory_config, "attention_output_gate", False) + ) + # Native fused SwiGLU retains gate/up and the output (3F). Eager + # unfused SwiGLU also retains SiLU and offset tensors (5F). Compilation + # may fall back, so only the native fusion setting earns this discount. + self._mlp_activation_factor = ( + 3 + if getattr(memory_config, "bias_activation_fusion", False) + and not getattr(memory_config, "use_te_activation_func", False) + else 5 + + 2 + * (getattr(memory_config, "activation_func_clamp_value", None) is not None) ) # Layers that run the gated-delta-net path (Qwen3.5-4B: 24 of 32); the # cost model prices GDN state hand-offs per GDN layer, not per layer. @@ -5140,51 +5152,61 @@ def _estimate_required_memory_bytes_from_values( ) if signature.grad_enabled and self._recompute_granularity != "full": geometry = self._geometry - ffn_width = max( - geometry.ffn_hidden_size or 4 * self._hidden_size, - geometry.moe_topk * geometry.moe_ffn_hidden_size - + geometry.moe_shared_expert_ffn, - ) hidden = self._hidden_size - attention_width = max( - hidden, geometry.num_attention_heads * geometry.kv_channels - ) - gdn_width = max( - hidden, - 2 * geometry.gdn_key_heads * geometry.gdn_key_head_dim - + 2 * geometry.gdn_value_heads * geometry.gdn_value_head_dim, - ) tp = max(1, self._topology_key()[1]) sp = tp if self._sequence_parallel else 1 - gdn_layers = min(self._num_layers, self._gdn_layers) - # Split Megatron's 9H attention/norm term into 5H of norms/residuals - # (sequence parallel) and projection intermediates (tensor parallel). - # Gated MLPs also retain unfused gate/up and LoRA sums: four FFN - # widths underpredicted native eager peaks; charge six in both modes - # since torch.compile can fall back to eager. GDN's projection, - # convolution and recurrent streams use a separate seven-width term - # (see dev/trainer_rank_recompute_memory.md). ART's LoRA path - # additionally retains gathered H-wide inputs to attention and MLP, - # even with SP. Price hybrid layers separately, not at the max width. - layer_features = 5 * hidden / sp + 2 * hidden - if geometry.moe_experts: - # Routed/shared expert storage gets no TP/EP discount without - # evidence for expert sharding and imbalanced dispatch. - layer_features += 6 * ffn_width + 2 * hidden * max( - 0, geometry.moe_topk - 1 + # Gathered LoRA inputs alias norm output without sequence sharding. + gathered = hidden if sp > 1 else 0 + common = 2 * hidden / sp + gathered + attention_width = ( + geometry.num_attention_heads * geometry.kv_channels or hidden + ) + kv_width = geometry.num_query_groups * geometry.kv_channels or hidden + gated = self._attention_output_gate + attention = ( + common + ((7 if gated else 5) * attention_width + 3 * kv_width) / tp + ) + if 0 < geometry.num_query_groups < tp: + # SelfAttentionLinearQKVLoRA constructs global QKV before + # slicing it when KV groups cannot be partitioned across TP. + attention += ((2 if gated else 1) * attention_width + 2 * kv_width) * ( + 1 - 1 / tp ) - else: - layer_features += 6 * ffn_width / tp - retained_features = ( - self._num_layers * layer_features + gdn = ( + common + ( - (self._num_layers - gdn_layers) * (9 * attention_width - 5 * hidden) - + gdn_layers * (7 * gdn_width - 5 * hidden) + 4 * geometry.gdn_key_heads * geometry.gdn_key_head_dim + + 8 * geometry.gdn_value_heads * geometry.gdn_value_head_dim ) / tp ) - # Optional selective modules retain the conservative undiscounted - # floor until their effective native checkpoint boundaries are tested. + ffn_width = geometry.ffn_hidden_size or 4 * hidden + mlp = common + self._mlp_activation_factor * ffn_width / tp + if geometry.moe_experts: + # Keep the worst-case dispatch envelope: random-weight runs + # cannot establish an EP discount for pretrained routing. + ffn_width = max( + ffn_width, + geometry.moe_topk * geometry.moe_ffn_hidden_size + + geometry.moe_shared_expert_ffn, + ) + mlp = ( + common + 6 * ffn_width + 2 * hidden * max(0, geometry.moe_topk - 1) + ) + gdn_layers = min(self._num_layers, self._gdn_layers) + retained_features = ( + (self._num_layers - gdn_layers) * attention + + gdn_layers * gdn + + self._num_layers * mlp + ) + if ( + self._recompute_granularity == "selective" + and "mlp" in self._recompute_modules + and not geometry.moe_experts + ): + # Checkpointed dense MLPs keep their input; one live MLP still + # needs its full workspace during the forward. + retained_features -= max(0, self._num_layers - 1) * (mlp - hidden / sp) static_compute = max( static_compute, packed_tokens * self._param_dtype_size * retained_features, diff --git a/tests/unit/test_trainer_rank_recompute_memory.py b/tests/unit/test_trainer_rank_recompute_memory.py index 63a6d9784..b2d1a4758 100644 --- a/tests/unit/test_trainer_rank_recompute_memory.py +++ b/tests/unit/test_trainer_rank_recompute_memory.py @@ -68,7 +68,7 @@ def test_reported_cold_request_is_refused_before_execution( rank = _rank(granularity) monkeypatch.setattr(rank, "_topology_key", lambda: (1, tp, 1, 1)) monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (0, 1)) - monkeypatch.setattr(rank, "_available_memory_bytes", lambda: int(119.289e9)) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: int(119.289 * 2**30)) monkeypatch.setattr( rank, "_execute_flat_plan", lambda _: pytest.fail("unsafe forward admitted") ) @@ -138,6 +138,8 @@ def _hybrid_rank(monkeypatch: pytest.MonkeyPatch, tp: int) -> TrainerRank: rank = _rank( "selective", sequence_parallel=True, + bias_activation_fusion=True, + attention_output_gate=True, linear_num_key_heads=16, linear_key_head_dim=128, linear_num_value_heads=48, @@ -208,6 +210,7 @@ def test_gathered_inputs_cover_attention_only_cold_peak( ffn_hidden_size=6144, num_layers=28, num_attention_heads=16, + num_query_groups=8, kv_channels=128, sequence_parallel=True, ) From 4f50366a68be8f9fa8aefacdbaa01e485a2a22e6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 17 Sep 2026 00:15:28 +0000 Subject: [PATCH 11/15] Account for cold workspace and checkpointed MoE retention --- dev/trainer_rank_recompute_memory.py | 20 +++++++ src/art/trainer_rank/_impl.py | 35 +++++++----- .../test_trainer_rank_recompute_memory.py | 54 ++++++++++++++++++- 3 files changed, 95 insertions(+), 14 deletions(-) diff --git a/dev/trainer_rank_recompute_memory.py b/dev/trainer_rank_recompute_memory.py index 4ab82cd2e..bcb64f8b9 100644 --- a/dev/trainer_rank_recompute_memory.py +++ b/dev/trainer_rank_recompute_memory.py @@ -39,6 +39,11 @@ def main() -> None: parser.add_argument( "--components", action="store_true", help="Attribute eager layer allocations" ) + parser.add_argument( + "--concentrate-routing", + action="store_true", + help="Send every token to the first top-k experts", + ) parser.add_argument("--evidence", type=Path, required=True) args = parser.parse_args() if args.repeat < 1 or args.layers < 0 or any(n < 1 for n in args.tokens): @@ -79,6 +84,20 @@ def emit(row): provider_configure=configure, print_env=False, ) + if args.concentrate_routing: + # Stress dispatch imbalance without replacing native routing/experts. + for module in runtime.model[0].modules(): + if type(module).__name__ == "TopKRouter": + + def gating(inputs, original=module.gating): + logits = original(inputs) + selected = ( + torch.arange(logits.shape[-1], device=logits.device) + < runtime.provider.moe_router_topk + ) + return logits + selected * 10000 + + module.gating = gating for chunk in runtime.model: chunk.train() if args.components and runtime.transformer_layers_compiled: @@ -94,6 +113,7 @@ def emit(row): or subprocess.check_output(["git", "rev-parse", "HEAD"], text=True).strip(), "model": args.model, "initialization": "random", + "concentrated_routing": args.concentrate_routing, "lora_rank": 1, "mode": rank._recompute_granularity, "recompute_method": runtime.provider.recompute_method, diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index c2bae1be1..3eaad51aa 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1424,6 +1424,10 @@ def __init__(self, runtime: TrainingRuntime) -> None: ) spec = getattr(runtime, "model_support_spec", None) self._moe_layers = _moe_layer_count(runtime.model[0]) + self._checkpointed_moe_layers = sum( + getattr(module, "moe_layer_recompute", False) is True + for module in runtime.model[0].modules() + ) is_moe = bool( self._moe_layers or getattr(spec, "is_moe", False) @@ -5185,11 +5189,10 @@ def _estimate_required_memory_bytes_from_values( if geometry.moe_experts: # Keep the worst-case dispatch envelope: random-weight runs # cannot establish an EP discount for pretrained routing. - ffn_width = max( - ffn_width, + ffn_width = ( geometry.moe_topk * geometry.moe_ffn_hidden_size - + geometry.moe_shared_expert_ffn, - ) + + geometry.moe_shared_expert_ffn + ) or ffn_width mlp = ( common + 6 * ffn_width + 2 * hidden * max(0, geometry.moe_topk - 1) ) @@ -5199,17 +5202,23 @@ def _estimate_required_memory_bytes_from_values( + gdn_layers * gdn + self._num_layers * mlp ) - if ( - self._recompute_granularity == "selective" - and "mlp" in self._recompute_modules - and not geometry.moe_experts - ): - # Checkpointed dense MLPs keep their input; one live MLP still - # needs its full workspace during the forward. - retained_features -= max(0, self._num_layers - 1) * (mlp - hidden / sp) + if self._recompute_granularity == "selective": + checkpointed = ( + self._checkpointed_moe_layers + if geometry.moe_experts + else self._num_layers + if "mlp" in self._recompute_modules + else 0 + ) + # Checkpoints keep their input (and MoE's external norm); one + # live MLP still needs workspace, including worst-case dispatch. + checkpoint_input = (2 if geometry.moe_experts else 1) * hidden / sp + retained_features -= max(0, checkpointed - 1) * (mlp - checkpoint_input) static_compute = max( static_compute, - packed_tokens * self._param_dtype_size * retained_features, + # Cold eager runs allocate ~58 MiB beyond warm retention for + # native kernel initialization; a slope alone misses short inputs. + 64 * 2**20 + packed_tokens * self._param_dtype_size * retained_features, ) # Groups execute sequentially: summed packed rows conservatively bound # this FC2 component, not all workspace or retained graphs. diff --git a/tests/unit/test_trainer_rank_recompute_memory.py b/tests/unit/test_trainer_rank_recompute_memory.py index b2d1a4758..0c3be3557 100644 --- a/tests/unit/test_trainer_rank_recompute_memory.py +++ b/tests/unit/test_trainer_rank_recompute_memory.py @@ -157,7 +157,59 @@ def test_sharded_floor_covers_recorded_native_gdn_peaks(monkeypatch, tp, peak): # H200, 64-layer Qwen3.8-27B, LoRA r1, SP, cold/warm max, 2,048 tokens. # First unsharded campaign, linked from dev/trainer_rank_recompute_memory.md. rank = _hybrid_rank(monkeypatch, tp) - assert rank._memory_check(_plan(rank, tokens=2048)).estimated_required_bytes >= peak + estimate = rank._memory_check(_plan(rank, tokens=2048)).estimated_required_bytes + assert peak <= estimate <= 1.2 * peak + + +@pytest.mark.parametrize( + "tp,mlp,packed,logical,peak_gib", + [ + (4, False, 4096, 4096, 20.613), + (4, False, 6964, 8192, 35.090), + (4, False, 13928, 16384, 69.924), + (8, False, 4096, 4096, 14.529), + (8, False, 6968, 8192, 24.654), + (8, False, 13928, 16384, 48.950), + (4, True, 4096, 4096, 11.213), + (4, True, 6964, 8192, 19.006), + (4, True, 13928, 16384, 37.881), + ], +) +def test_hybrid_estimate_is_close_to_recorded_peaks( + monkeypatch, tp, mlp, packed, logical, peak_gib +): + # Native H200 cold/warm witnesses in the calibration report. Coverage alone + # would allow the former 60%-high estimate; also check useful admission. + rank = _hybrid_rank(monkeypatch, tp) + rank._recompute_modules = frozenset(("core_attn", "mlp") if mlp else ("core_attn",)) + plan = replace(_plan(rank, tokens=packed), output_bytes=logical * 5120 * 2) + estimate = rank._memory_check(plan).estimated_required_bytes / 2**30 + assert peak_gib <= estimate <= 1.15 * peak_gib + + +def test_native_fusion_discount_is_independent_of_compilation(monkeypatch): + rank = _hybrid_rank(monkeypatch, 4) + plan = _plan(rank, tokens=4096) + fused = rank._memory_check(plan).estimated_required_bytes + rank.runtime.transformer_layers_compiled = False + assert rank._memory_check(plan).estimated_required_bytes == fused + rank._mlp_activation_factor = 5 + assert rank._memory_check(plan).estimated_required_bytes > fused + + +def test_mlp_discount_requires_selective_and_keeps_one_live_workspace(monkeypatch): + rank = _hybrid_rank(monkeypatch, 4) + rank._recompute_granularity = None + plan = _plan(rank) + unrecomputed = rank._memory_check(plan).estimated_required_bytes + rank._recompute_modules = frozenset(("mlp",)) + assert rank._memory_check(plan).estimated_required_bytes == unrecomputed + rank._recompute_granularity = "selective" + assert rank._memory_check(plan).estimated_required_bytes < unrecomputed + rank._num_layers = 1 + checkpointed = rank._memory_check(plan).estimated_required_bytes + rank._recompute_modules = frozenset() + assert rank._memory_check(plan).estimated_required_bytes == checkpointed def test_tp4_admits_eight_k_sibling_pair_and_prices_actual_layer_mix(monkeypatch): From 45297a4af6f1a17211a5a7c1ab02f19dc423f4a8 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 17 Sep 2026 00:15:45 +0000 Subject: [PATCH 12/15] Preserve provider fallback for partial metadata configs --- src/art/trainer_rank/_impl.py | 30 +++++++++++++----------------- 1 file changed, 13 insertions(+), 17 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 3eaad51aa..e3b4b7ad0 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1381,28 +1381,24 @@ def __init__(self, runtime: TrainingRuntime) -> None: or 1 ) memory_config = getattr(metadata_model, "config", None) or runtime.provider - self._recompute_granularity = getattr( - memory_config, "recompute_granularity", None - ) - self._recompute_modules = frozenset( - getattr(memory_config, "recompute_modules", ()) or () - ) - self._sequence_parallel = bool( - getattr(memory_config, "sequence_parallel", False) - ) - self._attention_output_gate = bool( - getattr(memory_config, "attention_output_gate", False) - ) + + def memory_field(name: str, default: Any = None) -> Any: + return getattr( + memory_config, name, getattr(runtime.provider, name, default) + ) + + self._recompute_granularity = memory_field("recompute_granularity", None) + self._recompute_modules = frozenset(memory_field("recompute_modules", ()) or ()) + self._sequence_parallel = bool(memory_field("sequence_parallel", False)) + self._attention_output_gate = bool(memory_field("attention_output_gate", False)) # Native fused SwiGLU retains gate/up and the output (3F). Eager # unfused SwiGLU also retains SiLU and offset tensors (5F). Compilation # may fall back, so only the native fusion setting earns this discount. self._mlp_activation_factor = ( 3 - if getattr(memory_config, "bias_activation_fusion", False) - and not getattr(memory_config, "use_te_activation_func", False) - else 5 - + 2 - * (getattr(memory_config, "activation_func_clamp_value", None) is not None) + if memory_field("bias_activation_fusion", False) + and not memory_field("use_te_activation_func", False) + else 5 + 2 * (memory_field("activation_func_clamp_value", None) is not None) ) # Layers that run the gated-delta-net path (Qwen3.5-4B: 24 of 32); the # cost model prices GDN state hand-offs per GDN layer, not per layer. From 850e20a2c4c60247430b3d2f4e463eaf860f92a7 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 17 Sep 2026 00:26:19 +0000 Subject: [PATCH 13/15] Price retained GDN states per prefix segment --- dev/trainer_rank_recompute_memory.py | 11 +++-- src/art/trainer_rank/_impl.py | 42 ++++++++++++++++++- .../test_trainer_rank_recompute_memory.py | 42 +++++++++++++++++++ 3 files changed, 91 insertions(+), 4 deletions(-) diff --git a/dev/trainer_rank_recompute_memory.py b/dev/trainer_rank_recompute_memory.py index bcb64f8b9..7b1e98245 100644 --- a/dev/trainer_rank_recompute_memory.py +++ b/dev/trainer_rank_recompute_memory.py @@ -34,6 +34,9 @@ def main() -> None: parser.add_argument( "--pairs", action="store_true", help="Two sequences at each token length" ) + parser.add_argument( + "--sequences", type=int, default=0, help="Override sequence count per batch" + ) parser.add_argument("--prefix-fraction", type=float, default=0.3) parser.add_argument("--modules", nargs="+", default=["core_attn"]) parser.add_argument( @@ -48,6 +51,8 @@ def main() -> None: args = parser.parse_args() if args.repeat < 1 or args.layers < 0 or any(n < 1 for n in args.tokens): parser.error("repeat/tokens must be positive and layers nonnegative") + if args.sequences < 0: + parser.error("sequences must be nonnegative") if not 0 <= args.prefix_fraction < 1: parser.error("prefix-fraction must be in [0, 1)") load_dotenv(".env") @@ -144,10 +149,9 @@ def gating(inputs, original=module.gating): }, } args.evidence.parent.mkdir(parents=True, exist_ok=True) + count = args.sequences or (2 if args.pairs else 1) workloads = [ - ([length, length], int(length * args.prefix_fraction)) - if args.pairs - else ([length], 0) + ([length] * count, int(length * args.prefix_fraction) if count > 1 else 0) for length in args.tokens ] if args.reported_pair: @@ -184,6 +188,7 @@ def gating(inputs, original=module.gating): "shared_prefix": prefix, "logical_tokens": plan.logical_tokens, "packed_tokens": plan.packed_tokens, + "grad_segment_count": plan.grad_segment_count, "output_bytes": plan.output_bytes, "selected_max_depth": plan.selected_max_depth, "memory_minimal": memory_minimal, diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index e3b4b7ad0..1eebd9603 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -941,6 +941,12 @@ def active_logical_tokens(self) -> int: # Keep total-input telemetry while pricing only executed requests. return self.logical_tokens - self.inactive_logical_tokens + @property + def grad_segment_count(self) -> int: + return sum( + len(group.packed.segments) for group in self.groups if group.grad_enabled + ) + @property def subforward_count(self) -> int: return 1 @@ -2925,6 +2931,7 @@ def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: output_bytes=plan.output_bytes, signature=plan.signature, logical_tokens=plan.active_logical_tokens, + gdn_segments=plan.grad_segment_count, ) def _subforward_cost( @@ -2934,12 +2941,14 @@ def _subforward_cost( output_bytes: int, signature: _MemorySignature, logical_tokens: int, + gdn_segments: int = 0, ) -> _SubforwardCost: required = self._estimate_required_memory_bytes_from_values( packed_tokens=packed_tokens, output_bytes=output_bytes, signature=signature, logical_tokens=logical_tokens, + gdn_segments=gdn_segments, ) retained = self._retained_memory_bytes( signature, @@ -3937,6 +3946,12 @@ def priced( output_bytes=output_bytes, signature=signature, logical_tokens=logical_tokens, + # A radix tree has fewer than twice as many segments as + # active requests; the exact plan uses its actual count. + gdn_segments=2 + * sum( + _request_mix_key(r) != "inactive" for r in local_requests + ), ) return ( self._memory_check_required(required, sync_across_dp=True), @@ -5078,6 +5093,7 @@ def _memory_check( output_bytes=forward.output_bytes, signature=forward.signature, logical_tokens=forward.active_logical_tokens, + gdn_segments=forward.grad_segment_count, ) return self._memory_check_required(required, sync_across_dp=sync_across_dp) @@ -5139,6 +5155,7 @@ def _estimate_required_memory_bytes_from_values( output_bytes: int, signature: _MemorySignature, logical_tokens: int | None = None, + gdn_segments: int = 0, ) -> int: if packed_tokens <= 0: return output_bytes @@ -5210,11 +5227,34 @@ def _estimate_required_memory_bytes_from_values( # live MLP still needs workspace, including worst-case dispatch. checkpoint_input = (2 if geometry.moe_experts else 1) * hidden / sp retained_features -= max(0, checkpointed - 1) * (mlp - checkpoint_input) + # Each GDN segment can retain an initial and a final recurrent + # state (fp32), plus convolution history. Unlike token activations, + # these do not shrink with segment length. + gdn_state_bytes = ( + 2 + * gdn_segments + * gdn_layers + / tp + * ( + 4 + * geometry.gdn_value_heads + * geometry.gdn_key_head_dim + * geometry.gdn_value_head_dim + + self._param_dtype_size + * ( + 2 * geometry.gdn_key_heads * geometry.gdn_key_head_dim + + geometry.gdn_value_heads * geometry.gdn_value_head_dim + ) + * max(0, geometry.gdn_conv_kernel - 1) + ) + ) static_compute = max( static_compute, # Cold eager runs allocate ~58 MiB beyond warm retention for # native kernel initialization; a slope alone misses short inputs. - 64 * 2**20 + packed_tokens * self._param_dtype_size * retained_features, + 64 * 2**20 + + gdn_state_bytes + + packed_tokens * self._param_dtype_size * retained_features, ) # Groups execute sequentially: summed packed rows conservatively bound # this FC2 component, not all workspace or retained graphs. diff --git a/tests/unit/test_trainer_rank_recompute_memory.py b/tests/unit/test_trainer_rank_recompute_memory.py index 0c3be3557..56f03dd95 100644 --- a/tests/unit/test_trainer_rank_recompute_memory.py +++ b/tests/unit/test_trainer_rank_recompute_memory.py @@ -245,6 +245,7 @@ def test_sp_discount_excludes_gathered_lora_inputs(monkeypatch): @pytest.mark.parametrize( "packed,logical,peak", [ + (436, 512, 829413888), # Cold 256-token pair before adding workspace. (1742, 2048, 3043651584), (6964, 8192, 11941280768), (13928, 16384, 23861987840), @@ -326,3 +327,44 @@ def pack(tensor: torch.Tensor) -> torch.Tensor: < retained(False) <= rank._memory_check(_plan(rank, tokens=32)).estimated_required_bytes ) + + +def test_moe_discount_uses_effective_native_checkpoint_count(monkeypatch): + rank = _rank( + "selective", + sequence_parallel=True, + num_moe_experts=256, + moe_router_topk=8, + moe_ffn_hidden_size=512, + moe_shared_expert_intermediate_size=512, + recompute_modules=["moe"], + ) + monkeypatch.setattr(rank, "_topology_key", lambda: (1, 4, 1, 1)) + plan = _plan(rank, tokens=4096) + # A module list alone cannot guarantee the native MoE checkpoint is active. + undiscounted = rank._memory_check(plan).estimated_required_bytes + rank._checkpointed_moe_layers = rank._num_layers + assert rank._memory_check(plan).estimated_required_bytes < undiscounted + rank._recompute_granularity = None + assert rank._memory_check(plan).estimated_required_bytes == undiscounted + + +def test_short_hybrid_pairs_pay_for_recurrent_states(monkeypatch): + rank = _hybrid_rank(monkeypatch, 4) + requests = [ + ForwardInput(input_tokens=torch.arange(64) + offset, hidden_states=True) + for offset in (0, 100) + ] + plan = rank._plan_flat_forward(requests) + assert plan.grad_segment_count == 2 + estimate = rank._memory_check(plan).estimated_required_bytes + # Cold eager TP4 at 45297a4af missed this peak without segment states. + assert 892974592 <= estimate <= 1.2 * 892974592 + assert rank._plan_cost(plan).required == estimate + rank._memory_profiles[plan.signature] = _MemoryProfile(0, plan.packed_tokens) + assert rank._memory_check(plan).estimated_required_bytes == estimate + inactive = rank._plan_flat_forward( + requests + [ForwardInput(input_tokens=torch.arange(1000))] + ) + assert inactive.grad_segment_count == 2 + assert rank._memory_check(inactive).estimated_required_bytes == estimate From 7760536e7bf10e4054d708383a1333e7b9f0e120 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 17 Sep 2026 00:31:53 +0000 Subject: [PATCH 14/15] Keep expert peers nonempty in concentrated-routing stress runs --- dev/trainer_rank_recompute_memory.py | 23 +++++++++++++++++++---- 1 file changed, 19 insertions(+), 4 deletions(-) diff --git a/dev/trainer_rank_recompute_memory.py b/dev/trainer_rank_recompute_memory.py index 7b1e98245..c888b3766 100644 --- a/dev/trainer_rank_recompute_memory.py +++ b/dev/trainer_rank_recompute_memory.py @@ -45,7 +45,7 @@ def main() -> None: parser.add_argument( "--concentrate-routing", action="store_true", - help="Send every token to the first top-k experts", + help="Concentrate 31/32 of tokens on the first expert rank", ) parser.add_argument("--evidence", type=Path, required=True) args = parser.parse_args() @@ -96,9 +96,24 @@ def emit(row): def gating(inputs, original=module.gating): logits = original(inputs) - selected = ( - torch.arange(logits.shape[-1], device=logits.device) - < runtime.provider.moe_router_topk + experts = logits.shape[-1] + ep = runtime.provider.expert_model_parallel_size + rows = torch.arange( + logits.numel() // experts, device=logits.device + ).reshape(*logits.shape[:-1], 1) + # Fully empty EP peers crash TE's grouped GEMM. Leave + # 1/32 of tokens on peers while stressing near-max load. + owner = ( + torch.where( + rows % 32 == 0, 1 + (rows // 32) % max(1, ep - 1), 0 + ) + if ep > 1 + else torch.zeros_like(rows) + ) + expert = torch.arange(experts, device=logits.device) + selected = (expert >= owner * (experts // ep)) & ( + expert + < owner * (experts // ep) + runtime.provider.moe_router_topk ) return logits + selected * 10000 From b86ad4bdd143e9d0a24e6204957c8ce970c20294 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Thu, 17 Sep 2026 00:40:50 +0000 Subject: [PATCH 15/15] Record tighter GPU calibration and its remaining limits --- dev/trainer_rank_recompute_memory.csv | 125 ++++++---- dev/trainer_rank_recompute_memory.md | 320 ++++++++++++++------------ src/art/trainer_rank/_impl.py | 4 +- 3 files changed, 248 insertions(+), 201 deletions(-) diff --git a/dev/trainer_rank_recompute_memory.csv b/dev/trainer_rank_recompute_memory.csv index d7e539aeb..e9bc8dc5c 100644 --- a/dev/trainer_rank_recompute_memory.csv +++ b/dev/trainer_rank_recompute_memory.csv @@ -1,51 +1,74 @@ -phase,source_sha,driver_sha256,model,layers,gdn_layers,tp,ep,etp,sequence_parallel,compiled,mode,modules,lengths,shared_prefix,logical_tokens,packed_tokens,selected_max_depth,memory_minimal,estimate_bytes,available_min_bytes,measured_rank_samples,refused_rank_samples,forward_peak_max_bytes,cold_forward_peak_max_bytes,warm_forward_peak_max_bytes,retained_max_bytes,forward_backward_peak_max_bytes,baseline_min_bytes,baseline_max_bytes,finite_loss_all,finite_gradients_all,estimate_covers_forward_all,min_estimate_to_forward_ratio -eager_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,1024+1024,307,2048,1742,2,False,2756291788,140940714701,4,0,3043651584,3043651584,2975649792,3036479488,3091428352,1781843968,1883400192,True,True,False,0.9055871580339203 -eager_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,16384+16384,4915,32768,27854,2,False,44072283340,140914919629,4,0,47563188736,47563188736,47548090368,47448521216,48253829632,1894096896,1909195264,True,True,False,0.9266048915396251 -eager_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,4096+4096,1228,8192,6964,2,False,11018859315,140937149133,4,0,11941280768,11941280768,11937715200,11912611840,12113940480,1883400192,1886965760,True,True,False,0.9227535579372788 -eager_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,8192+8192,2457,16384,13928,2,False,22037718630,140930017997,4,0,23861987840,23861987840,23854856704,23804649984,24207305216,1886965760,1894096896,True,True,False,0.9235491518044459 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,1024+1024,307,2048,1742,2,False,3415587225,140940714701,4,0,3043651584,3043651584,2975649792,3036479488,3091428352,1781843968,1883400192,True,True,True,1.1222004657021873 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,16384+16384,4915,32768,27854,2,False,54614197862,140914919629,4,0,47563188736,47563188736,47548090368,47448521216,48253829632,1894096896,1909195264,True,True,True,1.1482450885523403 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,4096+4096,1228,8192,6964,2,False,13654527180,140937149133,4,0,11941280768,11941280768,11937715200,11912611840,12113940480,1883400192,1886965760,True,True,True,1.1434725843304114 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,selective,core_attn,8192+8192,2457,16384,13928,2,False,27309054361,140930017997,4,0,23861987840,23861987840,23854856704,23804649984,24207305216,1886965760,1894096896,True,True,True,1.1444584811673426 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,none,,1024+1024,307,2048,1742,2,False,3415587225,140940714701,4,0,2406920192,2406920192,2338918400,2386762240,2441711104,1781843968,1883400192,True,True,True,1.4190695796032442 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,none,,16384+16384,4915,32768,27854,2,False,54614197862,140914919629,4,0,37406161408,37406161408,37391063040,37083114496,37888422912,1894096896,1909195264,True,True,True,1.4600321392592757 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,none,,4096+4096,1228,8192,6964,2,False,13654527180,140937149133,4,0,9424284672,9424284672,9420719104,9343731200,9545059840,1883400192,1886965760,True,True,True,1.4488661638764215 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,none,,8192+8192,2457,16384,13928,2,False,27309054361,140930017997,4,0,18819614208,18819614208,18812483072,18658097152,19060752384,1886965760,1894096896,True,True,True,1.4510953337922963 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,selective,core_attn,1024+1024,307,2048,1742,2,False,3415587225,140940714701,4,0,2405357568,2405357568,2337355776,2385200128,2440148992,1781843968,1883400192,True,True,True,1.419991468395272 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,selective,core_attn,16384+16384,4915,32768,27854,2,False,54614197862,140914919629,4,0,37381202432,37381202432,37366104064,37058156032,37863464448,1894096896,1909195264,True,True,True,1.4610069850307377 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,selective,core_attn,4096+4096,1228,8192,6964,2,False,13654527180,140937149133,4,0,9418034176,9418034176,9414468608,9337481216,9538809856,1883400192,1886965760,True,True,True,1.4498277373845028 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,selective,core_attn,8192+8192,2457,16384,13928,2,False,27309054361,140930017997,4,0,18807127552,18807127552,18799996416,18645611008,19048266240,1886965760,1894096896,True,True,True,1.452058762588436 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn,1024+1024,307,2048,2048,1,False,18095066316,124303160525,8,0,4869984768,4869984768,4582460416,4825250304,4876046848,17555830784,17929557504,True,True,True,3.715630988191214 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn,2048+2048,614,4096,4096,1,False,36190132633,124122997965,8,0,8887517184,8887517184,8857254912,8802950144,8903695360,17928215552,18065679872,True,True,True,4.072018302046414 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn,4096+4096,1228,8192,6964,2,False,61535826739,123840261837,8,0,15021402112,15021402112,14957089792,14845153792,15046482432,18065558016,18281307136,True,True,True,4.096543470455496 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn+moe,1024+1024,307,2048,2048,1,False,18095066316,124135796941,8,0,2815510016,2815510016,2317386752,2735823360,2893087232,17555830784,18096921088,True,True,True,6.426923084332583 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn+moe,2048+2048,614,4096,4096,1,False,36190132633,123792823501,8,0,4870201856,4870201856,4568279552,4702040064,4952712704,18094838272,18395854336,True,True,True,7.430930729989028 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,selective,core_attn+moe,4096+4096,1228,8192,6964,2,False,61535826739,123262782157,8,0,8278291968,8278291968,7798616576,7982971392,8422111744,18392608256,18856689664,True,True,True,7.433396523928932 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-9B,32,24,2,1,1,True,True,selective,core_attn,3072+3072,1843,6144,4302,2,False,24865721548,133630881485,4,0,15198974976,15198974976,15131242496,15108335616,15410327552,9090471936,9191136256,True,True,True,1.6360130592532927 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.5-9B,32,24,2,1,1,True,True,selective,core_attn,6144+6144,3686,12288,8602,2,False,49719908761,133630881485,4,0,30070096896,30070096896,30070068224,29888892928,30492874752,9191136256,9191136256,True,True,True,1.6534668622106723 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,selective,core_attn,19221+19222,5733,38443,32712,2,True,282826872627,128853849805,0,4,,,,,,,,,,, -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,selective,core_attn,2048+2048,614,4096,4096,1,False,35405797785,128855946957,8,0,22132520448,22132520448,22065410560,22076925440,22328585728,13865406464,13966070784,True,True,True,1.5997182909278385 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,selective,core_attn,4096+4096,1228,8192,6964,2,False,60210603622,128853849805,8,0,37665510400,37665510400,37665510400,37582912512,38086231040,13966070784,13966070784,True,True,True,1.5985606721527394 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,selective,core_attn,8192+8192,2457,16384,13928,2,False,120421207244,128853849805,8,0,75080549888,75080549888,75080549888,74915458048,75922093056,13966070784,13966070784,True,True,True,1.6038935173441866 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,none,,19221+19222,5733,38443,32712,2,True,282826872627,128853849805,0,4,,,,,,,,,,, -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,none,,2048+2048,614,4096,4096,1,False,35405797785,128855946957,8,0,22094640640,22094640640,22027530752,22040088064,22291748352,13865406464,13966070784,True,True,True,1.6024609027087575 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,none,,4096+4096,1228,8192,6964,2,False,60210603622,128853849805,8,0,37677824512,37677824512,37677824512,37598091264,38101409792,13966070784,13966070784,True,True,True,1.5980382201425547 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,none,,8192+8192,2457,16384,13928,2,False,120421207244,128853849805,8,0,74964444672,74964444672,74964444672,74802887168,75809522176,13966070784,13966070784,True,True,True,1.6063776337021085 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn,19221+19222,5733,38443,32712,2,True,282826872627,128853849805,0,4,,,,,,,,,,, -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn,2048+2048,614,4096,4096,1,False,35405797785,128855946957,8,0,22093067776,22093067776,22025957888,22038521344,22290181632,13865406464,13966070784,True,True,True,1.6025749861439251 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn,4096+4096,1228,8192,6964,2,False,60210603622,128853849805,8,0,37675145728,37675145728,37675145728,37595429376,38098747904,13966070784,13966070784,True,True,True,1.5981518441016076 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn,8192+8192,2457,16384,13928,2,False,120421207244,128853849805,8,0,74959095296,74959095296,74959095296,74797569024,75804204032,13966070784,13966070784,True,True,True,1.606492271131052 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn+mlp,19221+19222,5733,38443,32712,2,True,282826872627,128853849805,0,4,,,,,,,,,,, -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn+mlp,2048+2048,614,4096,4096,1,False,35405797785,128855946957,8,0,12039591424,12039591424,11972481536,11836139008,12136017408,13865406464,13966070784,True,True,True,2.940780674203052 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn+mlp,4096+4096,1228,8192,6964,2,False,60210603622,128853849805,8,0,20407347712,20407347712,20407347712,20073009664,20576328192,13966070784,13966070784,True,True,True,2.9504374831911524 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,selective,core_attn+mlp,8192+8192,2457,16384,13928,2,False,120421207244,128853849805,8,0,40674358272,40674358272,40674358272,40005276672,41011911680,13966070784,13966070784,True,True,True,2.9606172625690146 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,none,,16384+16384,4915,32768,27856,2,True,140687035596,135698216653,0,8,,,,,,,,,,, -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,none,,19221+19222,5733,38443,32712,2,True,165211897241,135698216653,0,8,,,,,,,,,,, -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,none,,2048+2048,614,4096,4096,1,False,20678757580,135700313805,16,0,15596174336,15596174336,15531161600,15554143744,15805804032,7018942464,7119606784,True,True,True,1.3258865369482364 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,none,,4096+4096,1228,8192,6968,2,False,35191907942,135698216653,16,0,26428632064,26428632064,26428632064,26357117440,26860435968,7119606784,7119606784,True,True,True,1.3315826508454434 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,none,,8192+8192,2457,16384,13928,2,False,70343517798,135698216653,16,0,52559443456,52559443456,52559443456,52416501248,53423136256,7119606784,7119606784,True,True,True,1.338361161622419 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,selective,core_attn,16384+16384,4915,32768,27856,2,True,140687035596,135698216653,0,8,,,,,,,,,,, -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,selective,core_attn,19221+19222,5733,38443,32712,2,True,165211897241,135698216653,0,8,,,,,,,,,,, -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,selective,core_attn,2048+2048,614,4096,4096,1,False,20678757580,135700313805,16,0,15595388416,15595388416,15530375680,15553363456,15805023744,7018942464,7119606784,True,True,True,1.3259533541841604 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,selective,core_attn,4096+4096,1228,8192,6968,2,False,35191907942,135698216653,16,0,26427289088,26427289088,26427289088,26355792384,26859110912,7119606784,7119606784,True,True,True,1.3316503189114393 -final,aef12a9e8502d181fbc614198a4bc4b438a34e9f,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,selective,core_attn,8192+8192,2457,16384,13928,2,False,70343517798,135698216653,16,0,52556765184,52556765184,52556765184,52413853696,53420488704,7119606784,7119606784,True,True,True,1.338429363978719 +phase,source_sha,driver_sha256,model,layers,gdn_layers,tp,ep,etp,sequence_parallel,compiled,concentrated_routing,mode,modules,lengths,shared_prefix,logical_tokens,packed_tokens,grad_segment_count,selected_max_depth,memory_minimal,estimate_bytes,available_min_bytes,measured_rank_samples,refused_rank_samples,forward_peak_max_bytes,first_forward_peak_max_bytes,repeat_forward_peak_max_bytes,retained_max_bytes,forward_backward_peak_max_bytes,baseline_min_bytes,baseline_max_bytes,finite_loss_all,finite_gradients_all,estimate_covers_forward_all,min_estimate_to_forward_ratio +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,12288+12288,3686,24576,20892,,2,False,124038771507,128845011620,8,0,112387296256,112387296256,112387296256,112146179072,113656130560,13966070784,13966070784,True,True,True,1.103672529183902 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,19221+19222,5733,38443,32712,,2,True,194173605683,128845011620,0,4,,,,,,,,,,, +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,2048+2048,614,4096,4096,,1,False,24369745100,128845011620,8,0,22027530752,22027530752,22027530752,21972978176,22224638464,13966070784,13966070784,True,True,True,1.106331225881382 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,256+256,76,512,512,,1,False,3110810419,128847108772,8,0,3042869760,3042869760,2975759872,3036048896,3114492928,13865406464,13966070784,True,True,True,1.022327823521438 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,4096+4096,1228,8192,6964,,2,False,41395470336,128845011620,8,0,37665950208,37665950208,37665950208,37586216960,38089535488,13966070784,13966070784,True,True,True,1.0990156920880725 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,8192+8192,2457,16384,13928,,2,False,82717120921,128845011620,8,0,74952197632,74952197632,74952197632,74790640128,75797275136,13966070784,13966070784,True,True,True,1.1035983404666023 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,none,,16384+16384,4915,32768,27856,,2,False,115282732646,135698216653,16,0,104854994944,104854994944,104854994944,104569116160,106582384128,7119606784,7119606784,True,True,True,1.0994491269354325 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,none,,19221+19222,5733,38443,32712,,2,False,135366117990,135698207949,16,0,123127285760,123127285760,123127277056,122791571968,125156542976,7119606784,7119615488,True,True,True,1.0993998377732126 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,none,,2048+2048,614,4096,4096,,1,False,17006224998,135700313805,16,0,15596174336,15596174336,15531161600,15554143744,15805804032,7018942464,7119606784,True,True,True,1.0904100346419723 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,none,,4096+4096,1228,8192,6968,,2,False,28892538470,135698216653,16,0,26428632064,26428632064,26428632064,26357117440,26860435968,7119606784,7119606784,True,True,True,1.093228677142024 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,none,,8192+8192,2457,16384,13928,,2,False,57678276198,135698216653,16,0,52559443456,52559443456,52559443456,52416501248,53423136256,7119606784,7119606784,True,True,True,1.0973913041199763 +before_segment_states,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.5-9B,32,24,2,1,1,True,True,False,selective,core_attn,10000+10000,6000,20000,14000,,2,False,53618370150,133630881485,4,0,48752947200,48752947200,48752947200,48458028032,49441070080,9191136256,9191136256,True,True,True,1.0997975143951912 +before_segment_states,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.5-9B,32,24,2,1,1,True,True,False,selective,core_attn,5000+5000,3000,10000,7000,,2,False,26846094950,133632978637,4,0,24515836928,24515836928,24460253184,24369065984,24860588032,9090471936,9191136256,True,True,True,1.0950511307790014 +cold_workspace_negative_control,71055f4639b2d9f247d7f3a3421a7f79cecc2c61,9bfae96d3b04221b7ccac18086bce2bcc13b3544055be38c83db4c4505c3f370,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,256+256,76,512,436,,2,False,813621248,140941383373,4,0,829413888,829413888,768536064,827618816,866536960,1781843968,1882731520,True,True,False,0.9809592771130449 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,1024+1024,307,2048,1742,,2,False,3324583116,140940323533,4,0,2976541696,2976541696,2975649792,2969369600,3019703296,1882899456,1883791360,True,True,True,1.1169281184495794 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,128+128,38,256,218,,2,False,480630374,140941438669,4,0,379880448,379880448,379768832,378982400,385275904,1882564608,1882676224,True,True,True,1.2652148235857614 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,16384+16384,4915,32768,27854,,2,False,52052538982,140914528461,4,0,47563188736,47563188736,47548090368,47448521216,48253829632,1894488064,1909586432,True,True,True,1.0943870746538502 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,256+256,76,512,436,,2,False,887440998,140941215437,4,0,768759296,768759296,768536064,766964224,779549184,1882676224,1882899456,True,True,True,1.154380834960336 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,4096+4096,1228,8192,6964,,2,False,13069429964,140936757965,4,0,11941280768,11941280768,11937715200,11912611840,12113940480,1883791360,1887356928,True,True,True,1.0944747232661334 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,64+64,19,128,110,,2,False,279085875,140941550285,4,0,256186368,256186368,189020160,255733248,290649600,1781843968,1882564608,True,True,True,1.0893861261189355 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,8192+8192,2457,16384,13928,,2,False,26065040179,140929626829,4,0,23861987840,23861987840,23854856704,23804649984,24207305216,1887356928,1894488064,True,True,True,1.0923247616155016 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,False,selective,core_attn,4096+4096,1228,8192,6964,3,2,False,13069429964,140938041037,4,0,9483472896,9483472896,9414886400,9402919936,9604248576,1781843968,1886073856,True,True,True,1.3781269907475042 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,False,selective,core_attn,8192+8192,2457,16384,13928,3,2,False,26065040179,140930909901,4,0,18806726144,18806726144,18799595008,18645209600,19047864832,1886073856,1893204992,True,True,True,1.385942453748956 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn,1500+1500,450,3000,3000,2,1,False,19630804582,124212695757,8,0,14968917504,14968917504,14561442304,14694942208,14968917504,17555830784,17994856448,True,True,True,1.311437822858884 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn,3000+3000,900,6000,5100,3,2,False,33310583398,123993616077,8,0,24725110272,24725110272,24549762048,24262462976,24725110272,17991170048,18165701632,True,True,True,1.3472370004239227 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn,6000+6000,1800,12000,10200,3,2,False,66441104998,123551084237,8,0,49004391936,49004391936,48701675008,48079419904,49004391936,18149577728,18494987264,True,True,True,1.3558193944080041 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn+moe,1500+1500,450,3000,3000,2,1,False,4606881382,123971948237,8,0,4355604992,4355604992,3704471040,3923585536,4501488128,17555830784,18235603968,True,True,True,1.0576903531108819 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn+moe,3000+3000,900,6000,5100,3,2,False,7769913958,123571519693,8,0,6649523200,6649523200,6301277184,5920428544,6840852992,18233154560,18589895168,True,True,True,1.1684918939751952 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn+moe,6000+6000,1800,12000,10200,3,2,False,15359766118,122728326861,8,0,13253387776,13253387776,12513351168,11752620032,13627614208,18589530112,19319841792,True,True,True,1.1589313145891915 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn,1500+1500,450,3000,3000,2,1,False,19630804582,124206687908,8,0,6961144320,6961144320,6634839040,6897091072,6971726336,17555830784,17994123264,True,True,True,2.820054243897647 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn,3000+3000,900,6000,5100,3,2,False,33310583398,123984802980,8,0,11145096192,11145096192,11101520896,11016890880,11164938752,17993316352,18169870848,True,True,True,2.9888107580363896 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn,6000+6000,1800,12000,10200,3,2,False,66441104998,123536855204,8,0,21793972736,21793972736,21733763072,21565764096,21860678144,18163617792,18502475264,True,True,True,3.0485999869243847 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn+moe,1500+1500,450,3000,3000,2,1,False,4606881382,123961172644,8,0,4025781248,4025781248,3373262336,3903570944,4122183680,17555830784,18239638528,True,True,True,1.1443446869570197 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn+moe,3000+3000,900,6000,5100,3,2,False,7769913958,123548129444,8,0,6099939840,6099939840,5727812096,5894723072,6208201728,18236235264,18606544384,True,True,True,1.2737689488426167 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn+moe,6000+6000,1800,12000,10200,3,2,False,15359766118,122723152548,8,0,12123608576,12123608576,11390174208,11691548160,12330070016,18591955456,19316177920,True,True,True,1.2669302231025774 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,10240+10240,5120,20480,15360,3,2,False,28776870707,138374372045,4,0,25853539840,25853539840,25853539840,25068629504,26035840512,4447645696,4447645696,True,True,True,1.1130727507757792 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,128+128,64,256,256,2,1,False,662215065,138374372045,4,0,498778624,498778624,498778624,484490240,501137408,4447645696,4447645696,True,True,True,1.3276733066251052 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,2048+2048,1024,4096,4096,2,1,False,7788272025,138374372045,4,0,6917972480,6917972480,6917972480,6701941248,6959876608,4447645696,4447645696,True,True,True,1.1258026896632003 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,5120+5120,2560,10240,10240,2,1,False,19189963161,138374372045,4,0,17218811392,17218811392,17218811392,16678733312,17323568640,4447645696,4447645696,True,True,True,1.1144766455782082 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,64+64,32,128,128,2,1,False,424679833,138376469197,4,0,336145408,336145408,269035520,328608256,432695808,4346981376,4447645696,True,True,True,1.2633813310934772 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,2,1,1,True,False,False,selective,core_attn,256+256+256+256+256+256+256+256,0,2048,2048,8,1,False,22748594176,115395819213,4,0,20101526016,20101526016,20101526016,20025496064,20101526016,27428295680,27428295680,True,True,True,1.1316849356557825 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,2,1,1,True,False,False,selective,core_attn,512+512+512+512+512+512+512+512,0,4096,4096,8,1,False,44068660838,115395819213,4,0,39580185088,39580185088,39580185088,39428125184,39580185088,27428295680,27428295680,True,True,True,1.1134020909710405 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,2,1,1,True,False,False,selective,core_attn,64+64+64+64+64+64+64+64,0,512,512,8,1,False,6758544179,115395819213,4,0,5649679872,5649679872,5590548992,5628575232,5864041472,27327631360,27428295680,True,True,True,1.1962702900204254 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,12288+12288,3686,24576,20892,3,2,False,124292779212,128845011620,8,0,112506383872,112506383872,112506383872,112259961856,113769913344,13966070784,13966070784,True,True,True,1.1047620138018972 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,2048+2048,614,4096,4096,2,1,False,24539083571,128847108772,8,0,22065410560,22065410560,22065410560,22009815552,22261475840,13966070784,13966070784,True,True,True,1.1121063668529356 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,4096+4096,1228,8192,6964,3,2,False,41649478041,128845011620,8,0,37657602560,37657602560,37657602560,37575004672,38078323200,13966070784,13966070784,True,True,True,1.1060045039946378 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,64+64,19,128,128,2,1,False,1002405888,128847108772,8,0,903788032,903788032,836678144,902051328,938049536,13865406464,13966070784,True,True,True,1.1091161339919138 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,8192+8192,2457,16384,13928,3,2,False,82971128627,128845011620,8,0,75082050048,75082050048,75082050048,74916958208,75923593216,13966070784,13966070784,True,True,True,1.1050727647148222 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,12288+12288,3686,24576,20892,3,2,False,124292779212,128845011620,8,0,112379268096,112379268096,112379268096,112138194432,113648145920,13966070784,13966070784,True,True,True,1.1060116453670341 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,19221+19222,5733,38443,32712,3,2,True,194427613388,128845011620,0,4,,,,,,,,,,, +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,2048+2048,614,4096,4096,2,1,False,24539083571,128845011620,8,0,22025957888,22025957888,22025957888,21971411456,22223071744,13966070784,13966070784,True,True,True,1.1140983604789865 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,256+256,76,512,512,2,1,False,3280148889,128845011620,8,0,2975563264,2975563264,2975563264,2968744960,3000204288,13966070784,13966070784,True,True,True,1.1023623421773767 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,4096+4096,1228,8192,6964,3,2,False,41649478041,128845011620,8,0,37663271424,37663271424,37663271424,37583555072,38086873600,13966070784,13966070784,True,True,True,1.10583803441089 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,64+64,19,128,128,2,1,False,1002405888,128847108772,8,0,900511232,900511232,833401344,898807296,997523456,13865406464,13966070784,True,True,True,1.1131520100795367 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,8192+8192,2457,16384,13928,3,2,False,82971128627,128845011620,8,0,74946848256,74946848256,74946848256,74785321984,75791956992,13966070784,13966070784,True,True,True,1.1070662817413086 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,16384+16384,4915,32768,27856,3,2,False,90497895628,128845011620,8,0,81347673600,81347673600,81347673600,80013564928,82026832896,13966070784,13966070784,True,True,True,1.1124828974580534 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,2048+2048,614,4096,4096,2,1,False,13493803417,128845011620,8,0,11972481536,11972481536,11972481536,11769029120,12035353088,13966070784,13966070784,True,True,True,1.1270682169294264 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,4096+4096,1228,8192,6964,3,2,False,22870344499,128845011620,8,0,20407347712,20407347712,20407347712,20073009664,20576328192,13966070784,13966070784,True,True,True,1.1206916656568604 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,64+64,19,128,128,2,1,False,657240883,128847108772,8,0,548321280,548321280,481211392,539703808,583254016,13865406464,13966070784,True,True,True,1.1986419403602209 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,8192+8192,2457,16384,13928,3,2,False,45412861542,128845011620,8,0,40674358272,40674358272,40674358272,40005276672,41011911680,13966070784,13966070784,True,True,True,1.1164985379317456 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,16384+16384,4915,32768,27856,3,2,False,115409736499,135698216653,16,0,104849291776,104849291776,104849291776,104563468800,106576736768,7119606784,7119606784,True,True,True,1.100720229427599 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,19221+19222,5733,38443,32712,3,2,False,135493121843,135698207949,16,0,123120420864,123120420864,123120412160,122784771584,125149742592,7119606784,7119615488,True,True,True,1.100492679379865 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,2048+2048,614,4096,4096,2,1,False,17090894233,135698216653,16,0,15528278528,15528278528,15526181376,15486253568,15737913856,7119606784,7119606784,True,True,True,1.100630324358386 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,4096+4096,1228,8192,6968,3,2,False,29019542323,135698216653,16,0,26427723264,26427723264,26427723264,26356226560,26859545088,7119606784,7119606784,True,True,True,1.09807197665531 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,64+64,19,128,128,2,1,False,687626649,135700313805,16,0,652222464,652222464,584981504,650909184,748507136,7018942464,7119606784,True,True,True,1.0542823759593782 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,8192+8192,2457,16384,13928,3,2,False,57805280051,135698216653,16,0,52556722176,52556722176,52556722176,52413810688,53420445696,7119606784,7119606784,True,True,True,1.0998646349637982 +segment_state_negative_control,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,64+64,19,128,128,,1,False,833067417,128847108772,8,0,892974592,892974592,836678144,890582528,989954048,13865406464,13966070784,True,True,False,0.9329127888556991 +segment_state_negative_control,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,64+64,19,128,128,,1,False,487902412,128847108772,8,0,548321280,548321280,481211392,539703808,639605248,13865406464,13966070784,True,True,False,0.8898111924454217 +segment_state_negative_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,128+128,64,256,256,,1,False,548890214,138376469197,4,0,563791360,563791360,498778624,549502976,660929536,4346981376,4447645696,True,True,False,0.973569751051169 +unfused_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,1024+1024,307,2048,1742,,2,False,2756291788,140940714701,4,0,3043651584,3043651584,2975649792,3036479488,3091428352,1781843968,1883400192,True,True,False,0.9055871580339203 +unfused_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,16384+16384,4915,32768,27854,,2,False,44072283340,140914919629,4,0,47563188736,47563188736,47548090368,47448521216,48253829632,1894096896,1909195264,True,True,False,0.9266048915396251 +unfused_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,4096+4096,1228,8192,6964,,2,False,11018859315,140937149133,4,0,11941280768,11941280768,11937715200,11912611840,12113940480,1883400192,1886965760,True,True,False,0.9227535579372788 +unfused_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,8192+8192,2457,16384,13928,,2,False,22037718630,140930017997,4,0,23861987840,23861987840,23854856704,23804649984,24207305216,1886965760,1894096896,True,True,False,0.9235491518044459 diff --git a/dev/trainer_rank_recompute_memory.md b/dev/trainer_rank_recompute_memory.md index d13f6c130..4f6ffe830 100644 --- a/dev/trainer_rank_recompute_memory.md +++ b/dev/trainer_rank_recompute_memory.md @@ -1,128 +1,150 @@ -# Recompute memory calibration, September 16, 2026 - -The sharded retained-activation estimate admits the requested Qwen3.8-27B -selective TP4 series (two sequences of 2k, 4k, and 8k tokens), bounds measured -forward peaks in compiled and eager execution, and refuses #913's original long -pair before execution. It accounts for the actual 48 GDN / 16 attention layers. - -The [CSV](trainer_rank_recompute_memory.csv) records final estimator commit -`aef12a9e8502d181fbc614198a4bc4b438a34e9f`, plus an earlier eager-execution failure -as a negative control. Each row aggregates all ranks and two repetitions for one -model/topology/mode/shape. Peaks are maxima; available memory is the minimum; -refusals have blank measurement fields. Source and driver hashes are per row. - -## Requested TP4 calibration - -All values are **incremental allocated GiB above the pre-forward baseline**, -not total GPU usage or reserved memory. Model/adaptor weights are random; this -measures the native runtime's memory behavior, not pretrained correctness. - -Qwen3.8-27B, 64 layers, bf16, LoRA rank 1, selective `core_attn`, SP enabled: - -| Tokens per sequence | Packed tokens | Estimate | Compiled forward | Eager forward | Compiled forward + backward | -| --: | --: | --: | --: | --: | --: | -| 2,048 | 4,096 | 32.974 | 20.576 | 20.613 | 20.759 | -| 4,096 | 6,964 | 56.075 | 35.088 | 35.079 | 35.482 | -| 8,192 | 13,928 | 112.151 | 69.811 | 69.924 | 70.598 | - -Each pair shares 30% of its prefix. The planner chose an unshared layout for the -2k pair and shared layouts for 4k/8k; token counts include TP padding. The compiled -run's pre-forward allocation was at most 13.007 GiB, with a usable incremental -budget of at least 120.004 GiB. The 8k prediction also fits the issue's 119.289 GiB -budget. No-recompute peaks were 20.577 / 35.090 / 69.816 GiB under the same estimates. - -The original synthetic sibling geometry (19,221 + 19,222 logical tokens, sharing -5,733 prefix tokens) is refused at TP4 with a 263.403 GiB estimate for 32,712 packed -tokens. This is an admission result, **not a measured long-pair selective peak**. - -## What changed in the estimate - -For gradient-enabled non-full recompute, the retained floor is the packed token -count times dtype bytes times the sum of layer storage. Dense layers retain -`5H/SP + 2H + 6F/TP`, plus `(9A - 5H)/TP` for each attention layer or -`(7G - 5H)/TP` for each GDN layer. Here `H` is hidden width, `F` is FFN width, -`A = max(H, heads * head_dim)`, `G = max(H, 2*key_width + 2*value_width)`, and -`SP` is TP when sequence parallel is enabled, otherwise 1. - -- The norm/residual term follows the SP distinction in - [Megatron's activation model](https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/training/theoretical_memory_usage.py). - Projection and dense MLP storage receive TP discounts. The separate `2H` - remains replicated because ART's LoRA wrappers retain gathered attention and - MLP inputs even with SP (`return_layernorm_output_gathered` and - `_column_parallel_lora_input` in `src/art/megatron/lora.py`). -- Four FFN widths covered compiled attention-only runs but underestimated eager - execution. Six cover the measured eager gate/up activations and LoRA sums. - Both execution modes use six: compilation can fall back to eager. -- GDN uses its own seven-width calibrated envelope for projection, convolution, - and recurrent storage. Charge it only to actual GDN layers, instead of charging - every layer the maximum of attention and GDN widths. This is an empirical - envelope for the measured native paths, not an exact tensor-liveness proof. -- Routed/shared expert FFN and dispatch storage receive no TP/EP/ETP discount. - Optional `mlp`/`moe` recomputation receives no discount in this revision. -- The existing static estimate and MoE FC2 estimate remain floors. Learned - profiles can only raise the estimate; output bytes and the existing 10% safety - margin are still applied. No-grad and full-recompute paths are unchanged. - -## Additional evidence - -Qwen3.8-27B selective TP8, compiled: - -| Tokens per sequence | Estimate | Forward peak | -| --: | --: | --: | -| 2,048 | 19.259 | 14.524 | -| 4,096 | 32.775 | 24.612 | -| 8,192 | 65.513 | 48.947 | - -No-recompute peaks were 14.525 / 24.614 / 48.950 GiB. The final estimate refuses -16k pairs (131.025 GiB) and the original long pair (153.866 GiB) at TP8. Those -selective peaks remain unmeasured; TP sharding does not remove gathered storage, -and the estimate can still over-refuse near capacity. - -The eager attention-only negative control is Qwen3-1.7B TP2, pairs of -1,024 / 4,096 / 8,192 / 16,384 tokens. Before the six-FFN correction, predictions -were 2.567 / 10.262 / 20.524 / 41.046 GiB against observed -2.835 / 11.121 / 22.223 / 44.297 GiB. Re-running the final code produced the same -peaks with estimates 3.181 / 12.717 / 25.434 / 50.863 GiB. The smallest observed -headroom in the final campaign is 12.2%. Recorded peaks are regression witnesses: -restoring the old estimator fails all four, and dropping the gathered-input term -fails three. - -A held-out Qwen3.5-9B model at TP2 used pairs of 3,072 / 6,144 tokens with a 60% -shared prefix, after coefficients were fixed. Estimates 23.158 / 46.305 GiB covered -forward peaks 14.155 / 28.005 GiB and forward/backward peaks 14.352 / 28.399 GiB. - -Adding `mlp` to selective recomputation on the 27B at TP4 reduced the 2k / 4k / -8k forward peaks to 11.213 / 19.006 / 37.881 GiB, under the same estimates. -Qwen3.5-35B-A3B at TP4/EP4/ETP1 used 1k / 2k / 4k pairs: estimates -16.852 / 33.705 / 57.310 GiB covered default selective peaks -4.536 / 8.277 / 13.990 GiB. Adding `moe` reduced them to -2.622 / 4.536 / 7.710 GiB. These runs support keeping the current undiscounted -floor; they do not establish per-module discounts or bounds on pretrained routing -imbalance. Compiled attention-only selective and no-recompute controls also pass. - -The final campaign contains **46 cells: 296 measured rank-samples and 48 refused -rank-samples**. Every measured forward was covered, every backward completed -without CUDA OOM, and every loss and adapter gradient was finite. The CSV also -includes four cells / 16 rank-samples from the earlier eager negative control; -those failed coverage checks are intentionally preserved, not final-code failures. - -## Method, reproduction, and limits - -NVIDIA H200, PyTorch 2.11.0+cu128, CUDA 12.8, bf16, DP/CP/PP=1. Each mode runs in a -fresh process. Every repetition clears learned memory profiles, gradients, and -unused cached allocations. Sample 0 includes any first-execution compilation and -autotuning; sample 1 is warm. Both contribute to reported maxima. - -The driver plans paired hidden-state requests with gradients, tries the -minimum-memory layout if admission refuses, and executes the unsplit native plan -only if it fits. It then backpropagates a mean-square hidden-state loss. There is -no admission bypass, splitting, or optimizer step. Forward/backward peaks include -the diagnostic loss; finite-gradient checks happen after peak/timing collection. -Raw JSONL also records geometry, every rank/sample, timing, and allocator totals. -Driver SHA-256: `42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788`. - -On a configured GPU machine, after -`INSTALL_VLLM_RUNTIME=false bash src/art/megatron/setup.sh`: +# Recompute memory calibration + +The revised estimate is **10–13% above measured 27B peaks** for the paired +2k-and-larger workloads, including MLP recompute. The previous TP4 estimate was +about 60% high. The new floor admits a 12k TP4 pair and, at TP8, the original +19,221 + 19,222-token pair from #913. Both complete forward and backward. + +All numbers below are **incremental allocated GiB above the pre-forward +baseline**, not total GPU use or reserved memory. The estimate includes the +existing 10% safety margin. Measurements use native H200 execution with bf16, +rank-1 LoRA, random model/adaptor weights, SP enabled, and DP/CP/PP=1. + +## Requested 27B series and admission boundary + +Qwen3.8-27B has 48 GDN and 16 full-attention layers. Pairs share 30% of their +prefix; the planner chooses the layout, including TP padding. + +| Tokens per sequence | TP | Packed tokens | Previous estimate | New estimate | Compiled forward | Eager forward | +| --: | --: | --: | --: | --: | --: | --: | +| 2,048 | 4 | 4,096 | 32.974 | 22.854 | 20.513 | 20.550 | +| 4,096 | 4 | 6,964 | 56.075 | 38.789 | 35.077 | 35.071 | +| 8,192 | 4 | 13,928 | 112.151 | 77.273 | 69.800 | 69.926 | +| 12,288 | 4 | 20,892 | — | 115.757 | 104.661 | 104.780 | +| 2,048 | 8 | 4,096 | 19.259 | 15.917 | 14.462 | — | +| 4,096 | 8 | 6,968 | 32.775 | 27.027 | 24.613 | — | +| 8,192 | 8 | 13,928 | 65.513 | 53.835 | 48.947 | — | +| 16,384 | 8 | 27,856 | 131.025 | 107.484 | 97.649 | — | +| 19,221 + 19,222 | 8 | 32,712 | 153.866 | 126.188 | 114.665 | — | + +The 12k TP4 and original long TP8 shapes were not used to fit the component +coefficients. The latter was previously refused, so its true selective peak +was unknown. TP4 still refuses the original pair, now at 181.075 GiB. The new +8k TP4 estimate fits both the measured budget and #913's 119.289 GiB budget. +The issue's error labels say GB but divide bytes by 1024³. + +Adding `mlp` to selective recomputation at TP4 gives: + +| Tokens per sequence | Estimate | Forward peak | Forward + backward peak | +| --: | --: | --: | --: | +| 2,048 | 12.567 | 11.150 | 11.209 | +| 4,096 | 21.300 | 19.006 | 19.163 | +| 8,192 | 42.294 | 37.881 | 38.195 | +| 16,384 | 84.283 | 75.761 | 76.393 | + +## What is being priced + +Component hooks on eager native layers separated MLP, attention, and GDN +retention. They exposed native activation fusion as the main MLP distinction: +Qwen3.8 uses fused SwiGLU even in eager execution, while Qwen3-1.7B uses the +unfused path. Compilation alone is not a reliable discount because it can fall +back to eager. + +Let H be hidden width, F dense FFN width, Q the full query width, KV the full +key/value width, and K/V the GDN key/value widths. SP is TP when sequence +parallel is enabled, otherwise 1. The common per-component term is +`2H/SP + gathered`, with `gathered = H` when SP>1 and zero otherwise: without +sequence sharding, the LoRA input aliases norm output. + +- Dense MLP storage is `common + cF/TP`: c=3 for native fused activations, + c=5 for unfused SwiGLU, with an additional allowance for unfused clamping. + These count gate/up, activation output, and the unfused SiLU/offset tensors. +- Full attention costs `common + (5Q + 3KV)/TP`, with another `2Q/TP` for + output gating. If KV groups are fewer than TP ranks, ART's replicated-QKV + LoRA path adds the global QKV storage that survives slicing. Omitting that + term underprices TP8. +- GDN costs `common + (4K + 8V)/TP`. Attention and GDN terms are multiplied by + their actual layer counts. These are calibrated envelopes of the measured + native paths, not an exact liveness proof for every kernel. +- Checkpointed dense MLPs retain their inputs. Native MoE checkpoint flags + determine the discounted MoE layer count; its external norm is also priced. + One full MLP workspace remains charged, including worst-case expert dispatch. +- Each gradient-enabled prefix segment gets an allowance for initial/final GDN + recurrent states (fp32) and convolution history. Exact plans use actual + segment counts; cheap admission uses the radix-tree bound, and optimistic + pruning omits this nonnegative term. A separate 64 MiB allowance covers the + observed roughly 58 MiB cold kernel setup cost. + +The existing static heuristic and routed FC2 bound remain floors. Profiles can +only increase the estimate, and output storage and the 10% safety factor are +applied afterward. Full-recompute and no-grad requests keep their existing path. + +## Cross-checks and remaining conservatism + +The [CSV](trainer_rank_recompute_memory.csv) includes short-input failures that +motivated the fixed costs. Without segment states, a 64-token 27B TP4 pair was +estimated at 0.776 GiB against 0.832 GiB observed. The corrected estimate is +0.934 GiB against a fresh eager peak of 0.842 GiB. Eight unshared 64-token +sequences at TP2 also pass: 6.294 estimated versus 5.262 observed. Larger +8-sequence batches and 4B MLP-checkpointed pairs provide additional checks. + +Qwen3.5-4B TP2 with MLP recompute, pairs of 2,048 / 5,120 / 10,240 tokens and a +50% common prefix, yields estimates 7.253 / 17.872 / 26.801 against peaks +6.443 / 16.036 / 24.078 GiB. A 9B TP2 check used previously unmeasured 5k / 10k +lengths with a 60% prefix: 25.002 / 49.936 estimated against 22.832 / 45.405 GiB, +before adding the nonnegative segment-state term. + +The unfused attention-only eager control is covered from 64 through 16,384 +tokens per sequence. Its larger inputs have about 9% headroom. Compiled-only +Qwen3-1.7B peaks are lower: 12.172 / 24.275 estimated against 8.832 / 17.515 GiB +for 4k / 8k pairs. That remaining 38–39% overestimate preserves the eager +fallback allowance; it does not apply to the natively fused 27B MLP. + +MoE cannot assume balanced expert dispatch. Qwen3.5-35B-A3B, TP4/EP4/ETP1: + +| Tokens per sequence | Default estimate | Balanced peak | Concentrated peak | With `moe`: estimate | Balanced peak | Concentrated peak | +| --: | --: | --: | --: | --: | --: | --: | +| 1,500 | 18.283 | 6.483 | 13.941 | 4.290 | 3.749 | 4.056 | +| 3,000 | 31.023 | 10.380 | 23.027 | 7.236 | 5.681 | 6.193 | +| 6,000 | 61.878 | 20.297 | 45.639 | 14.305 | 11.291 | 12.343 | + +The concentrated eager run routes about 31/32 of tokens to the first expert +rank; the balanced run uses the native random-weight router with compilation. +Their difference includes both routing and execution mode. The routed term +uses actual top-k expert plus shared FFN widths, not the unrelated dense FFN +fallback. It still makes no EP/ETP discount and is deliberately loose for +balanced routing. MoE recomputation removes most of that retained storage. + +A fully concentrated stress attempt segfaulted in Transformer Engine's grouped +GEMM on peers receiving zero tokens, before a complete measurement. The revised +stress keeps those peers nonempty. This kernel limitation is not fixed by the +memory estimate and is not counted as a successful calibration cell. + +## Evidence and reproduction + +The final campaign contains **45 cells, 360 measured rank-samples, and 4 refused +rank-samples**. Every measured forward is covered; every backward completes; +all measured losses and adapter gradients are finite. The smallest observed +forward headroom is 5.4%. The CSV also preserves earlier controls and negative +controls as separate phases, with the actual source/driver hash for each row. +It aggregates maximum peaks across every rank and both repetitions, and minimum +available memory. Refusals have no invented observed peak. + +Estimator source: `850e20a2c4c60247430b3d2f4e463eaf860f92a7`. Stress-driver-only +revision: `7760536e7bf10e4054d708383a1333e7b9f0e120`. A later type annotation does +not change the calculation. Native execution uses PyTorch 2.11.0+cu128, CUDA +12.8, and H200s. Component instrumentation was used only for diagnosis; final +measurements have no component hooks. + +Every sample clears learned profiles, gradients, and unused cached allocations. +The first workload in each process includes cold kernel setup; later lengths +can reuse initialized kernels. CSV `first_` and `repeat_` columns refer to the +two executions of each shape, not independent fresh processes. The driver tries +the minimum-memory layout before refusing, executes no split or budget bypass, +and backpropagates a mean-square hidden-state loss. Peak collection precedes +the finite-gradient checks. No optimizer step is included. + +After `INSTALL_VLLM_RUNTIME=false bash src/art/megatron/setup.sh`: ```sh ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE=4 \ @@ -131,29 +153,29 @@ ART_MEGATRON_DATA_PARALLEL_SIZE=1 \ ART_MEGATRON_PIPELINE_MODEL_PARALLEL_SIZE=1 \ uv run --project megatron_runtime --no-sync python -m torch.distributed.run \ --standalone --nproc-per-node=4 dev/trainer_rank_recompute_memory.py \ - --mode selective --pairs --tokens 2048 4096 8192 --reported-pair \ - --evidence scratch/recompute-memory/tp4-selective.jsonl + --mode selective --pairs --tokens 64 256 2048 4096 8192 12288 --reported-pair \ + --evidence scratch/recompute-memory-calibrated/tp4-selective.jsonl ``` -Use `--mode none`, `ART_DISABLE_MEGATRON_COMPILE=1`, or `--modules core_attn mlp` -for the corresponding controls. Change both process count and TP for TP8. -The [SkyPilot task](trainer_rank_recompute_memory.sky.yaml) defaults to TP4 on -free Kubernetes; pass the synced commit as `ART_CALIBRATION_SOURCE_SHA` and use -`--idle-minutes-to-autostop 15 --down`. - -Final cluster jobs: `art-915-sharded-tp4-0916:4` (selective/none), `:5` (eager), -`:6` (MLP recompute), `:7` (MoE), `:8` (MoE recompute), all on free `k8s/cks-wb3`; -`art-915-sharded-tp8-ext-0916:2` on free `k8s/ext-collab2`. TP2 controls ran locally. -The clusters used 15-minute autodown and two-hour pod deadlines. Raw evidence is -retained locally under `scratch/recompute-memory-final/`. No paid clusters were -used; both clusters were explicitly torn down after successful jobs and evidence -retrieval. - -The [first unsharded campaign](https://github.com/OpenPipe/ART/blob/be21a2599/dev/trainer_rank_recompute_memory.csv) -remains historical evidence. It also found underestimation in the unchanged full -recompute heuristic (for the long group, 5.894 GiB estimated versus 13.647 GiB -observed at TP2). This PR does not validate or fix that legacy path. Larger LoRA -ranks, pretrained routing imbalance, CP, deeper prefix trees, mixed gradient -groups, and other kernels/hardware are not calibrated by these measurements. -Conservative expert/module pricing remains deliberate; these results support -removing the selective guard for the measured paths, not a universal memory bound. +Use `ART_DISABLE_MEGATRON_COMPILE=1` for eager, `--modules core_attn mlp` for +dense MLP checkpoints, or `--modules core_attn moe` for MoE checkpoints. +`--sequences 8 --prefix-fraction 0` checks many short unshared sequences. +`--concentrate-routing` enables the expert-imbalance stress. When running four +processes on an eight-GPU host, also set +`ART_MEGATRON_EXPERT_MODEL_PARALLEL_SIZE=4` and +`ART_MEGATRON_EXPERT_TENSOR_PARALLEL_SIZE=1` for the MoE case. + +The [SkyPilot task](trainer_rank_recompute_memory.sky.yaml) uses free Kubernetes. +Final jobs: `art-915-tight-tp4-0917:6,7,8,9` and +`art-915-tight-tp8-0917:5,7`; TP2 controls ran locally. Both clusters had autodown +and two-hour pod deadlines and were explicitly torn down after evidence +retrieval. No paid clusters were used. Raw JSONL is retained locally under +`scratch/recompute-memory-calibrated/`, with diagnostic and negative-control +runs in `scratch/recompute-memory-components/` and `scratch/recompute-memory-tight/`. + +The [previous report](https://github.com/OpenPipe/ART/blob/57f9de2f9/dev/trainer_rank_recompute_memory.md) +records the looser estimate and the earlier full-recompute underestimation. +This work does not validate that legacy path, larger LoRA ranks, CP, pretrained +routing distributions, arbitrary deep prefix trees, or other hardware/kernels. +These measurements support the calibrated native paths, not a universal memory +bound. diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 1eebd9603..8b0603348 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1394,7 +1394,9 @@ def memory_field(name: str, default: Any = None) -> Any: ) self._recompute_granularity = memory_field("recompute_granularity", None) - self._recompute_modules = frozenset(memory_field("recompute_modules", ()) or ()) + self._recompute_modules: frozenset[str] = frozenset( + memory_field("recompute_modules", ()) or () + ) self._sequence_parallel = bool(memory_field("sequence_parallel", False)) self._attention_output_gate = bool(memory_field("attention_output_gate", False)) # Native fused SwiGLU retains gate/up and the output (3F). Eager