diff --git a/dev/trainer_rank_check.py b/dev/trainer_rank_check.py index d71784618..b0acf6c42 100644 --- a/dev/trainer_rank_check.py +++ b/dev/trainer_rank_check.py @@ -226,34 +226,10 @@ def _local_outputs( rank: TrainerRank, indexed_requests: Sequence[tuple[int, ForwardInput]], ) -> list[dict[str, object]]: - from art.megatron.lora import use_lora_slot - - requests = [request for _, request in indexed_requests] - plan = rank._plan_flat_forward(requests) - outputs: list[ForwardOutput] = [ - ForwardOutput(None, None, None, None) for _ in requests - ] - sources: list[torch.Tensor] = [torch.empty(0, dtype=torch.long) for _ in requests] - for group in plan.groups: - prepared = rank._prepare_packed_forward(group.packed) - with use_lora_slot(group.slot_ref): - group_outputs = rank._forward_packed(group.items, prepared) - for index, source, output in zip( - group.request_indices, - prepared.source_positions_by_item, - group_outputs, - strict=True, - ): - sources[index] = source - outputs[index] = output + outputs = rank.dp_rank_forward([request for _, request in indexed_requests]) return [ - _output_record(global_index, source, output) - for (global_index, _), source, output in zip( - indexed_requests, - sources, - outputs, - strict=True, - ) + _output_record(index, torch.arange(request.input_tokens.numel()), output) + for (index, request), output in zip(indexed_requests, outputs, strict=True) ] diff --git a/scripts/ci/trainer-rank-gpu-tests.sh b/scripts/ci/trainer-rank-gpu-tests.sh index b496c80b4..690d99db3 100755 --- a/scripts/ci/trainer-rank-gpu-tests.sh +++ b/scripts/ci/trainer-rank-gpu-tests.sh @@ -9,6 +9,7 @@ runtime_python="$( test -x "${runtime_python}" "${runtime_python}" -m pytest --tb=short \ + tests/unit/test_trainer_rank_head_recompute.py \ tests/unit/test_trainer_rank_custom_tensors.py \ tests/integration/megatron/cp_attn/test_attention_packed_vs_flattened.py \ 'tests/integration/megatron/gdn_shared_prefix/test_gdn_cp_packed_correctness.py::test_gdn_cp_packed_sibling_order_matches_cp1_oracle[2]' \ diff --git a/src/art/trainer_rank/__init__.py b/src/art/trainer_rank/__init__.py index 340d71560..fd321a820 100644 --- a/src/art/trainer_rank/__init__.py +++ b/src/art/trainer_rank/__init__.py @@ -253,6 +253,12 @@ def forward_micro_batches( ART moves its packed model inputs and labels internally without mutating the caller-owned `ForwardInput` objects. + Per-position outputs contain the full flattened input sequence in source + order, including with context parallelism. Callers must compute identical + losses on every TP/CP replica; ART routes gradients to owning rows + without multiplying them by the number of replicas. `dp_reduce` combines + only distinct data-parallel batches. + Empty local microbatches are skipped unless `yield_empty=True`. Every rank must use the same setting. When a wave skips ranks, TrainerRank collective methods raise if called from its loop body; fully populated @@ -333,6 +339,9 @@ def dp_rank_forward( ) -> ForwardOutputs: """Forward inputs already local to this data-parallel rank. + Outputs contain full sequences in source order on every TP/CP rank, + with the same loss and reduction contract as `forward_micro_batches`. + Per-input checkpoints and `no_grad` values override the method defaults. `no_grad=None` inherits the ambient PyTorch grad mode; `True` disables grads and `False` enables them. @@ -352,6 +361,7 @@ def dp_reduce( *, op: dist.ReduceOp.RedOpType = dist.ReduceOp.SUM, ) -> None: + """Reduce in place over data-parallel batches, excluding TP/CP replicas.""" super().dp_reduce(tensor, op=op) def optim_step( diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index de15bbdee..0e0045d61 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -184,6 +184,8 @@ class _LocalLoRASlotRef: @dataclass(frozen=True) class ForwardOutput(Generic[LogprobsT, TopKT, LogitsT, HiddenStatesT]): + """Per-request tensors in flattened input order, replicated across TP/CP.""" + target_logprobs: LogprobsT top_k: TopKT logits: LogitsT @@ -614,6 +616,34 @@ def backward( return grad_outputs[0], None +class _GatherContextParallelRows(torch.autograd.Function): + @staticmethod + def forward( + ctx: FunctionCtx, + tensor: torch.Tensor, + positions: torch.Tensor, + length: int, + group: dist.ProcessGroup, + ) -> torch.Tensor: + ctx.save_for_backward(positions) + # Each source row has one owner. Scatter + SUM handles unequal/empty + # shards without padded all-gather buffers, including for full logits. + output = tensor.new_zeros((length, *tensor.shape[1:])) + output.index_copy_(0, positions, tensor) + dist.all_reduce(output, group=group) + return output + + @staticmethod + def backward( + ctx: FunctionCtx, *grad_outputs: torch.Tensor + ) -> tuple[torch.Tensor, None, None, None]: + (positions,) = cast(tuple[torch.Tensor, ...], getattr(ctx, "saved_tensors")) + # The caller's loss is replicated over CP, just as over TP after the + # sequence-parallel gather with tensor_parallel_output_grad=False. + # Route one copy to its owner, without summing CP copies of the loss. + return grad_outputs[0].index_select(0, positions), None, None, None + + class _CustomSlotGraphSentinel(torch.autograd.Function): @staticmethod def forward( @@ -866,6 +896,7 @@ class _PreparedPackedForward: packed_seq_params: "PackedSeqParams | None" positions_by_item: tuple[torch.Tensor, ...] source_positions_by_item: tuple[torch.Tensor, ...] + context_parallel_group: dist.ProcessGroup | None = None type _RowMatch = tuple[torch.Tensor, torch.Tensor, tuple[int, ...]] @@ -3009,10 +3040,11 @@ def dp_reduce( self._guard_forward_collective("dp_reduce") from megatron.core import parallel_state as ps + # Public outputs are CP-replicated; internal shard reductions still include CP. dist.all_reduce( tensor, op=op, - group=ps.get_data_parallel_group(with_context_parallel=True), + group=ps.get_data_parallel_group(with_context_parallel=False), ) def optim_step( @@ -3784,6 +3816,10 @@ def add( for param, grad in zip(params, grads, strict=True): if bool(getattr(param, "allreduce", True)): group = ps.get_data_parallel_group(with_context_parallel=True) + if getattr(param, "_art_custom_checkpoint_param", False): + # Custom heads consume replicated full-sequence outputs; + # average their CP copies while still summing DP batches. + grad.div_(ps.get_context_parallel_world_size()) else: group = ps.get_expert_data_parallel_group() if group is not None and group.size() > 1: @@ -5095,6 +5131,10 @@ def _estimate_required_memory_bytes_from_values( static_compute = max( static_compute, packed_tokens * self._moe_output_bytes_per_token ) + if signature.topology[2] > 1: + # Local head results coexist with full CP outputs during gathering. + # Uneven rank plans can assign all of an item's rows to one rank. + static_compute += output_bytes # A profile learned under lighter sharing (lower logical/packed ratio) # underestimates the per-packed-token footprint of a deeper-shared # plan; scale the trusted estimate up by the ratio gap. @@ -5287,7 +5327,67 @@ def _forward_packed( hidden_by_row = self._gather_sequence_parallel_hidden( self._decoder_hidden(prepared) ) - return self._project_head(items, prepared, hidden_by_row) + outputs = self._project_head(items, prepared, hidden_by_row) + group = prepared.context_parallel_group + if group is None: + return outputs + tensors = [ + tensor + for output in outputs + for tensor in ( + output.target_logprobs, + output.logits, + output.hidden_states, + None if output.top_k is None else output.top_k.logprobs, + None if output.top_k is None else output.top_k.tokens, + ) + if tensor is not None + ] + grad_flags = [False for _ in tensors] + if torch.is_grad_enabled(): + flags = torch.tensor( + [ + tensor.is_floating_point() + and (tensor.requires_grad or hidden_by_row.requires_grad) + for tensor in tensors + ], + device=hidden_by_row.device, + dtype=torch.int32, + ) + dist.all_reduce(flags, op=dist.ReduceOp.MAX, group=group) + grad_flags = flags.tolist() + needs_grad = iter(grad_flags) + + def gather(tensor: torch.Tensor | None) -> torch.Tensor | None: + if tensor is None: + return None + if next(needs_grad) and not tensor.requires_grad: + # Empty shards must still enter decoder backward collectives. + # A frozen decoder with a trainable head only needs a leaf. + tensor = tensor + hidden_by_row.reshape(-1)[:1].sum() * 0.0 + tensor.requires_grad_(True) + return _GatherContextParallelRows.apply(tensor, positions, length, group) + + for index, (item, output, source_positions) in enumerate( + zip(items, outputs, prepared.source_positions_by_item, strict=True) + ): + positions = source_positions.to(device=hidden_by_row.device) + length = int(item.input_ids.numel()) + outputs[index] = replace( + output, + target_logprobs=gather(output.target_logprobs), + logits=gather(output.logits), + hidden_states=gather(output.hidden_states), + top_k=( + TopK( + cast(torch.Tensor, gather(output.top_k.logprobs)), + cast(torch.Tensor, gather(output.top_k.tokens)), + ) + if output.top_k is not None + else None + ), + ) + return outputs def _decoder_hidden( self, @@ -5899,6 +5999,7 @@ def _prepare_context_parallel_forward( packed_seq_params=prepared.packed_seq_params, positions_by_item=tuple(pair[0] for pair in local_position_pairs), source_positions_by_item=tuple(pair[1] for pair in local_position_pairs), + context_parallel_group=ps.get_context_parallel_group(), ) def _topology(self) -> "ParallelTopology": diff --git a/tests/integration/megatron/lora/test_dynamic_lora_slots.py b/tests/integration/megatron/lora/test_dynamic_lora_slots.py index 140e1fef3..cd9d47f60 100644 --- a/tests/integration/megatron/lora/test_dynamic_lora_slots.py +++ b/tests/integration/megatron/lora/test_dynamic_lora_slots.py @@ -272,7 +272,7 @@ def _custom_parameter_reduction_worker( torch.testing.assert_close(parameter, torch.tensor(1.0, device=device)) (parameter * float(rank + 1)).backward() (reduced,) = trainer._reduce_dynamic_grads((parameter,), scale_grads=1.0) - expected = {"dp": 3.0, "tp": 1.5, "cp": 3.0, "tp_cp": 5.0}[topology] + expected = {"dp": 3.0, "tp": 1.5, "cp": 1.5, "tp_cp": 2.5}[topology] torch.testing.assert_close(reduced, torch.tensor(expected, device=device)) finally: if getattr(ps, "model_parallel_is_initialized", lambda: False)(): diff --git a/tests/unit/test_trainer_rank_head_recompute.py b/tests/unit/test_trainer_rank_head_recompute.py index bcf94df97..1a2c8d9a8 100644 --- a/tests/unit/test_trainer_rank_head_recompute.py +++ b/tests/unit/test_trainer_rank_head_recompute.py @@ -1,12 +1,16 @@ from __future__ import annotations from contextlib import nullcontext +from datetime import timedelta from types import SimpleNamespace from unittest.mock import patch import pytest import torch +import torch.distributed as dist +import torch.multiprocessing as mp +from art.megatron.prefix_tree_packing import prefix_tree_pack from art.trainer_rank import ForwardInput, TrainerRank, _impl @@ -25,12 +29,7 @@ def _direct(function, *args, use_reentrant=False, **kwargs): return function(*args, **kwargs) -@pytest.mark.parametrize("top_k", (None, 3, 12)) -@pytest.mark.parametrize("include_logits", (False, True)) -@pytest.mark.parametrize("multi_target", (False, True)) -def test_recomputed_head_preserves_outputs_and_arbitrary_loss_gradients( - monkeypatch, top_k, include_logits, multi_target -): +def _patch_local_head(monkeypatch): monkeypatch.setattr(_impl, "_HEAD_CHUNK_TOKENS", 16) monkeypatch.setattr(_impl, "_language_model", lambda model: model) monkeypatch.setattr(_impl, "_all_reduce_tensor_parallel_max", lambda value: value) @@ -46,6 +45,15 @@ def test_recomputed_head_preserves_outputs_and_arbitrary_loss_gradients( monkeypatch.setattr( TrainerRank, "_gather_tensor_parallel_logits", lambda self, value: value ) + + +@pytest.mark.parametrize("top_k", (None, 3, 12)) +@pytest.mark.parametrize("include_logits", (False, True)) +@pytest.mark.parametrize("multi_target", (False, True)) +def test_recomputed_head_preserves_outputs_and_arbitrary_loss_gradients( + monkeypatch, top_k, include_logits, multi_target +): + _patch_local_head(monkeypatch) generator = torch.Generator().manual_seed(11) weight = torch.randn(65, 7, generator=generator) / 3 hidden = torch.randn(41, 7, generator=generator) @@ -134,3 +142,200 @@ def pack(tensor): assert not any( shape[-1:] == (65,) and len(shape) == 2 for shape, _, _ in recomputed[4] ) + + +@pytest.mark.parametrize("cp_size,dp_size", ((2, 1), (4, 1), (2, 2))) +def test_context_parallel_outputs_match_full_sequence(cp_size, dp_size, tmp_path): + pytest.importorskip("megatron.core") + mp.spawn( + _context_parallel_worker, + args=(cp_size, dp_size, f"file://{tmp_path / 'cp'}", "gloo"), + nprocs=cp_size * dp_size, + join=True, + ) + + +@pytest.mark.parametrize("cp_size", (2, 4)) +def test_context_parallel_outputs_cuda(cp_size, tmp_path): + if not torch.cuda.is_available() or torch.cuda.device_count() < cp_size: + pytest.skip(f"requires {cp_size} CUDA devices") + pytest.importorskip("megatron.core") + mp.spawn( + _context_parallel_worker, + args=(cp_size, 1, f"file://{tmp_path / 'cp'}", "nccl"), + nprocs=cp_size, + join=True, + ) + + +def _context_parallel_worker(rank, cp_size, dp_size, init_method, backend): + from megatron.core import parallel_state as ps + + device = torch.device("cpu" if backend == "gloo" else f"cuda:{rank}") + if device.type == "cuda": + torch.cuda.set_device(device) + dist.init_process_group( + backend, + init_method=init_method, + rank=rank, + world_size=cp_size * dp_size, + timeout=timedelta(seconds=90), + ) + try: + cp_groups = [ + dist.new_group(list(range(dp * cp_size, (dp + 1) * cp_size))) + for dp in range(dp_size) + ] + dp_groups = [ + dist.new_group(list(range(cp, cp_size * dp_size, cp_size))) + for cp in range(cp_size) + ] + cp_rank, dp_rank = rank % cp_size, rank // cp_size + cp_group, dp_group = cp_groups[dp_rank], dp_groups[cp_rank] + with pytest.MonkeyPatch.context() as monkeypatch: + _patch_local_head(monkeypatch) + monkeypatch.setattr(ps, "get_context_parallel_world_size", lambda: cp_size) + monkeypatch.setattr( + ps, + "get_data_parallel_group", + lambda *, with_context_parallel: ( + dist.group.WORLD if with_context_parallel else dp_group + ), + ) + monkeypatch.setattr(ps, "get_tensor_model_parallel_group", lambda **_: None) + for mode in ( + "hidden", + "targets", + "all", + "logits", + "head_only", + "frozen", + "no_grad", + ): + _check_context_parallel_case( + cp_rank, dp_rank, cp_size, dp_size, cp_group, device, mode + ) + finally: + dist.destroy_process_group() + + +def _check_context_parallel_case( + cp_rank, dp_rank, cp_size, dp_size, cp_group, device, mode +): + tokens = [torch.tensor(row) for row in ([1, 2, 3, 4, 5, 6, 7], [1, 2, 3, 8, 9])] + packed = prefix_tree_pack(tokens, max_depth=1) + generator = torch.Generator(device=device).manual_seed(911) + features = torch.randn(9, 3, generator=generator, device=device) + dp_rank / 5 + decoder_weight = torch.randn(3, 5, generator=generator, device=device) + head_weight = torch.randn(11, 5, generator=generator, device=device) + probe_weight = torch.randn(5, 2, generator=generator, device=device) + requests = [] + for index, row in enumerate(tokens): + labels = torch.stack((row, row.roll(1)), dim=1) + labels[::2] = -100 + if index == 1: + labels.fill_(-100) + requests.append( + ForwardInput( + input_tokens=row, + target_tokens=labels if mode not in ("hidden", "logits") else None, + hidden_states=mode not in ("targets", "logits"), + logits=mode not in ("hidden", "targets"), + top_k=3 if mode not in ("hidden", "targets", "logits") else None, + ) + ) + # Unequal, reversed ownership with a shared prefix; CP4 rank 3 is empty. + owners = torch.tensor([0, 1, 1, 2, 0, 1, 1, 2, 0]) % cp_size + rows = torch.nonzero(owners == cp_rank).flatten().flip(0) + dispatched = torch.cat((rows, torch.tensor([-1]))) + pairs = [ + _impl._local_position_pairs(dispatched[None], positions) + for positions in packed.positions_by_sequence + ] + + def run(local): + decoder = torch.nn.Parameter( + decoder_weight.clone(), requires_grad=mode not in ("frozen", "head_only") + ) + head = _Head(head_weight) + head.weight.requires_grad_(mode != "frozen") + probe = torch.nn.Parameter(probe_weight.clone(), requires_grad=mode != "frozen") + trainer = object.__new__(TrainerRank) + trainer._skipped_forward_waves = {} + trainer.runtime = SimpleNamespace( + model=[ + SimpleNamespace( + output_layer=head, + vocab_size=11, + share_embeddings_and_output_weights=False, + _scale_logits=lambda value: value, + ) + ] + ) + trainer._tag_custom_parameters((probe,)) + with torch.set_grad_enabled(mode != "no_grad"): + hidden = features @ decoder + if local: + hidden = torch.cat((hidden[rows], hidden.new_zeros((1, 5)))) + trainer._decoder_hidden = lambda _: hidden + trainer._gather_sequence_parallel_hidden = lambda value: value + outputs = trainer._forward_packed( + [trainer._forward_item(request) for request in requests], + SimpleNamespace( + positions_by_item=( + tuple(pair[0] for pair in pairs) + if local + else packed.positions_by_sequence + ), + source_positions_by_item=( + tuple(pair[1] for pair in pairs) + if local + else tuple(torch.arange(len(row)) for row in tokens) + ), + context_parallel_group=cp_group if local else None, + ), + ) + tensors, terms = [], [] + for output, row in zip(outputs, tokens, strict=True): + assert not hasattr(output, "positions") + values = [output.target_logprobs, output.logits, output.hidden_states] + if output.top_k is not None: + values.append(output.top_k.logprobs) + tensors.append(output.top_k.tokens) + for value in values: + if value is not None: + assert value.shape[0] == len(row) + tensors.append(value) + terms.append(value.mean().square() + value.square().mean()) + if output.hidden_states is not None: + # The original failing caller expression needs no CP mapping. + selected = output.hidden_states[1:][(row[1:] % 2 == 1).to(device)] + terms.append((selected @ probe).square().mean()) + loss = torch.stack(terms).sum() + if mode not in ("frozen", "no_grad"): + loss.backward() + return trainer, tensors, loss.detach(), (decoder, head.weight, probe) + + reference = run(False) + actual = run(True) + for value, expected in zip(actual[1], reference[1], strict=True): + torch.testing.assert_close(value, expected) + assert value.requires_grad == expected.requires_grad + torch.testing.assert_close(actual[2], reference[2]) + expected_loss = reference[2].clone() + dist.all_reduce(expected_loss) + expected_loss /= cp_size + trainer = actual[0] + trainer.dp_reduce(actual[2]) + torch.testing.assert_close(actual[2], expected_loss) + count = torch.tensor(sum(len(row) for row in tokens), device=device) + trainer.dp_reduce(count) + assert count.item() == dp_size * sum(len(row) for row in tokens) + if mode in ("frozen", "no_grad"): + return + reduced = trainer._reduce_dynamic_grads(actual[3], scale_grads=0.5) + for grad, param in zip(reduced, reference[3], strict=True): + expected = torch.zeros_like(param) if param.grad is None else param.grad.clone() + dist.all_reduce(expected) + expected *= 0.5 / cp_size + torch.testing.assert_close(grad, expected, atol=2e-5, rtol=2e-5) diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 03f3b6d3c..e0dabad67 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -86,6 +86,34 @@ def _signature(): return _MemorySignature((1, 1, 1, 1), (1, None), 1, (), True, (True,)) +def test_cp_memory_charges_local_and_gathered_outputs(): + rank = _rank() + signature = replace(_signature(), topology=(1, 1, 4, 1)) + output_bytes = 1 << 30 + estimate = rank._estimate_required_memory_bytes_from_values( + packed_tokens=16, output_bytes=output_bytes, signature=signature + ) + compute = 16 * 2048 * 2 * 14 + assert estimate == int((compute + 2 * output_bytes) * 1.1) + # A warm profile includes gather workspace already; do not add it twice. + rank._update_memory_profile( + SimpleNamespace( + signature=signature, + packed_tokens=16, + output_bytes=output_bytes, + active_logical_tokens=16, + ), + compute + 2 * output_bytes, + retained_bytes=None, + ) + assert ( + rank._estimate_required_memory_bytes_from_values( + packed_tokens=16, output_bytes=output_bytes, signature=signature + ) + == estimate + ) + + def test_supported_constructor_and_original_shape(layer): rank = _rank(layer) assert rank._moe_output_bytes_per_token == (512 + 3 * 2048) * 8 * 2