diff --git a/areal/engine/core/__init__.py b/areal/engine/core/__init__.py index 499d02fd22..1334457ceb 100644 --- a/areal/engine/core/__init__.py +++ b/areal/engine/core/__init__.py @@ -4,12 +4,14 @@ from areal.engine.core.train_engine import ( aggregate_eval_losses, + compute_microbatch_loss_weight, compute_total_loss_weight, reorder_and_pad_outputs, ) __all__ = [ "aggregate_eval_losses", + "compute_microbatch_loss_weight", "compute_total_loss_weight", "reorder_and_pad_outputs", ] diff --git a/areal/engine/core/train_engine.py b/areal/engine/core/train_engine.py index 9bc83c2b77..f416818ad2 100644 --- a/areal/engine/core/train_engine.py +++ b/areal/engine/core/train_engine.py @@ -14,6 +14,7 @@ from areal.infra.platforms import current_platform from areal.utils.data import ( + TRANSPORT_DUMMY_KEY, MicroBatchList, pad_and_stack_tensors_along_first_dim, reorder_list, @@ -21,12 +22,29 @@ ) __all__ = [ + "compute_microbatch_loss_weight", "compute_total_loss_weight", "aggregate_eval_losses", "reorder_and_pad_outputs", ] +def compute_microbatch_loss_weight( + microbatch: dict[str, Any], + loss_weight_fn: Callable[[dict[str, Any]], torch.Tensor], +) -> torch.Tensor: + """Return zero without invoking an objective on transport-only data.""" + if microbatch.get(TRANSPORT_DUMMY_KEY) is not True: + return loss_weight_fn(microbatch) + reference = next( + (value for value in microbatch.values() if isinstance(value, torch.Tensor)), + None, + ) + if reference is None: + raise ValueError("Transport micro-batch does not contain a tensor") + return torch.zeros((), dtype=torch.float32, device=reference.device) + + def compute_total_loss_weight( mb_list: MicroBatchList, loss_weight_fn: Callable[[dict[str, Any]], torch.Tensor], @@ -52,7 +70,9 @@ def compute_total_loss_weight( The total loss weight (scalar tensor) after all_reduce. """ total_weight = ( - torch.stack([loss_weight_fn(mb) for mb in mb_list.mbs]) + torch.stack( + [compute_microbatch_loss_weight(mb, loss_weight_fn) for mb in mb_list.mbs] + ) .sum() .detach() .clone() @@ -138,7 +158,11 @@ def reorder_and_pad_outputs( The processed outputs, padded and stacked along batch dimension. """ res = aggregate_fn(outputs) + semantic_batch_size = len(output_seqlens) + output_seqlens = [*output_seqlens, *([1] * mb_list.transport_dummy_count)] seqlens = [output_seqlens[i] for i in mb_list.forward_indices] unpacked = unpack_sequence(res, lens=seqlens, dim=0) reordered = reorder_list(unpacked, mb_list.backward_indices) + if mb_list.transport_dummy_count: + reordered = reordered[:semantic_batch_size] return pad_and_stack_tensors_along_first_dim(reordered) diff --git a/areal/engine/fsdp_engine.py b/areal/engine/fsdp_engine.py index 2a184c8c39..75643639ab 100644 --- a/areal/engine/fsdp_engine.py +++ b/areal/engine/fsdp_engine.py @@ -60,6 +60,7 @@ from areal.api.io_struct import DeviceRuntimeInfo from areal.engine.core import ( aggregate_eval_losses, + compute_microbatch_loss_weight, compute_total_loss_weight, reorder_and_pad_outputs, ) @@ -782,7 +783,9 @@ def train_batch( input_batched, _ = self._normalize_batch_input(input_) # Step 1: Prepare micro-batches - mb_list = self._prepare_mb_list(input_batched).to(self.device) + mb_list = self._prepare_mb_list(input_batched, allow_transport_padding=True).to( + self.device + ) # Step 2: Compute total loss weight total_loss_weight = compute_total_loss_weight( @@ -822,7 +825,9 @@ def eval_batch( input_batched, _ = self._normalize_batch_input(input_) # Step 1: Prepare micro-batches - mb_list = self._prepare_mb_list(input_batched).to(self.device) + mb_list = self._prepare_mb_list(input_batched, allow_transport_padding=True).to( + self.device + ) # Step 2: Compute total loss weight total_loss_weight = compute_total_loss_weight( @@ -880,7 +885,9 @@ def forward_batch( batch_size = len(output_seqlens) # Step 2: Prepare micro-batches - mb_list = self._prepare_mb_list(input_batched).to(self.device) + mb_list = self._prepare_mb_list(input_batched, allow_transport_padding=True).to( + self.device + ) # Step 3: Forward using process_output_fn callback, collecting results outputs: list[torch.Tensor] = [] @@ -1872,7 +1879,12 @@ def _load_optimizer_state(self, path: str): self.optimizer.load_state_dict(optimizer_state_dict) dist.barrier(group=self.cpu_group) - def _prepare_mb_list(self, input_: dict[str, Any]) -> MicroBatchList: + def _prepare_mb_list( + self, + input_: dict[str, Any], + *, + allow_transport_padding: bool = False, + ) -> MicroBatchList: assert "attention_mask" in input_ and "input_ids" in input_ input_ = input_.copy() @@ -1935,7 +1947,12 @@ def _prepare_mb_list(self, input_: dict[str, Any]) -> MicroBatchList: else: input_ = amend_position_ids(input_) - mb_list = split_padded_tensor_dict_into_mb_list(input_, self.config.mb_spec) + mb_list = split_padded_tensor_dict_into_mb_list( + input_, + self.config.mb_spec, + group=self.data_parallel_group if allow_transport_padding else None, + allow_transport_padding=allow_transport_padding, + ) mb_list.mbs = [pack_tensor_dict(mb) for mb in mb_list.mbs] mb_list = pad_mb_list( mb_list, @@ -2142,7 +2159,7 @@ def _compute_logprobs_and_loss( loss_multiplier: float = 1.0, ) -> torch.Tensor: """Compute logprobs/entropy and return scaled loss.""" - local_weight = loss_weight_fn(ctx.mb_input) + local_weight = compute_microbatch_loss_weight(ctx.mb_input, loss_weight_fn) if local_weight == 0: return logits.mean() * 0.0 diff --git a/areal/engine/megatron_engine.py b/areal/engine/megatron_engine.py index 46ffca073c..55b928adae 100644 --- a/areal/engine/megatron_engine.py +++ b/areal/engine/megatron_engine.py @@ -51,6 +51,7 @@ from areal.api.io_struct import DeviceRuntimeInfo from areal.engine.core import ( aggregate_eval_losses, + compute_microbatch_loss_weight, compute_total_loss_weight, reorder_and_pad_outputs, ) @@ -1008,7 +1009,9 @@ def train_batch( input_batched, _ = self._normalize_batch_input(input_) # Step 1: Prepare micro-batches - mb_list = self._prepare_mb_list(input_batched).to(self.device) + mb_list = self._prepare_mb_list(input_batched, allow_transport_padding=True).to( + self.device + ) # Step 2: Compute total loss weight. # Use DP+CP group: after CP all-gather each rank computes the full-sequence @@ -1069,7 +1072,9 @@ def eval_batch( input_batched, _ = self._normalize_batch_input(input_) # Step 1: Prepare micro-batches - mb_list = self._prepare_mb_list(input_batched).to(self.device) + mb_list = self._prepare_mb_list(input_batched, allow_transport_padding=True).to( + self.device + ) # Step 2: Compute total loss weight (DP+CP, see train_batch comment). total_loss_weight = compute_total_loss_weight( @@ -1128,7 +1133,9 @@ def forward_batch( batch_size = len(output_seqlens) # Step 2: Prepare micro-batches - mb_list = self._prepare_mb_list(input_batched).to(self.device) + mb_list = self._prepare_mb_list(input_batched, allow_transport_padding=True).to( + self.device + ) # Step 3: Forward using Megatron's pipeline function, collecting results outputs: list[torch.Tensor] = [] @@ -2250,7 +2257,12 @@ def _load_model_from_hf(self, path: str) -> None: fp8_direct_convert=self.fp8_direct_convert, ) - def _prepare_mb_list(self, input_: dict[str, Any]) -> MicroBatchList: + def _prepare_mb_list( + self, + input_: dict[str, Any], + *, + allow_transport_padding: bool = False, + ) -> MicroBatchList: assert "attention_mask" in input_ and "input_ids" in input_ # Parallel sizes pp_size = self.parallel_strategy.pipeline_parallel_size @@ -2301,6 +2313,7 @@ def _prepare_mb_list(self, input_: dict[str, Any]) -> MicroBatchList: input_, mb_spec, group=mpu.get_data_parallel_group(), + allow_transport_padding=allow_transport_padding, ) mb_list.mbs = [pack_tensor_dict(mb) for mb in mb_list.mbs] # NOTE: Pad micro-batches to: @@ -2362,7 +2375,7 @@ def _compute_logprobs_and_loss( total_loss_weight: torch.Tensor, loss_multiplier: float = 1.0, ) -> torch.Tensor: - local_weight = loss_weight_fn(inputs) + local_weight = compute_microbatch_loss_weight(inputs, loss_weight_fn) if local_weight == 0: return output.mean() * 0.0 diff --git a/areal/experimental/engine/archon_engine.py b/areal/experimental/engine/archon_engine.py index a7d95fc2bd..ada7701699 100644 --- a/areal/experimental/engine/archon_engine.py +++ b/areal/experimental/engine/archon_engine.py @@ -41,6 +41,7 @@ ) from areal.engine.core.train_engine import ( aggregate_eval_losses, + compute_microbatch_loss_weight, compute_total_loss_weight, reorder_and_pad_outputs, ) @@ -534,7 +535,9 @@ def train_batch( input_batched, _ = self._normalize_batch_input(input_) - mb_list = self._prepare_mb_list(input_batched).to(self.device) + mb_list = self._prepare_mb_list(input_batched, allow_transport_padding=True).to( + self.device + ) total_loss_weight = compute_total_loss_weight( mb_list, loss_weight_fn, self.data_parallel_group @@ -571,7 +574,9 @@ def eval_batch( input_batched, _ = self._normalize_batch_input(input_) - mb_list = self._prepare_mb_list(input_batched).to(self.device) + mb_list = self._prepare_mb_list(input_batched, allow_transport_padding=True).to( + self.device + ) total_loss_weight = compute_total_loss_weight( mb_list, loss_weight_fn, self.data_parallel_group @@ -629,7 +634,9 @@ def forward_batch( assert output_seqlens is not None batch_size = len(output_seqlens) - mb_list = self._prepare_mb_list(input_batched).to(self.device) + mb_list = self._prepare_mb_list(input_batched, allow_transport_padding=True).to( + self.device + ) def process_output( logits: torch.Tensor, ctx_dict: dict[str, Any] @@ -1183,7 +1190,12 @@ def _normalize_batch_input( return concat_batch(input_) return input_, None - def _prepare_mb_list(self, input_: dict[str, Any]) -> MicroBatchList: + def _prepare_mb_list( + self, + input_: dict[str, Any], + *, + allow_transport_padding: bool = False, + ) -> MicroBatchList: assert "attention_mask" in input_ and "input_ids" in input_ input_ = input_.copy() @@ -1211,7 +1223,7 @@ def _prepare_mb_list(self, input_: dict[str, Any]) -> MicroBatchList: stages_per_rank = len(self.pp_stages) num_total_stages = pp_size * stages_per_rank n_seqs = input_["attention_mask"].shape[0] - if n_seqs < num_total_stages: + if n_seqs < num_total_stages and not allow_transport_padding: raise RuntimeError( f"Pipeline parallelism requires at least {num_total_stages} " f"sequences (pp_size={pp_size} * stages_per_rank=" @@ -1227,7 +1239,12 @@ def _prepare_mb_list(self, input_: dict[str, Any]) -> MicroBatchList: else: mb_spec = self.config.mb_spec - mb_list = split_padded_tensor_dict_into_mb_list(input_, mb_spec) + mb_list = split_padded_tensor_dict_into_mb_list( + input_, + mb_spec, + group=self.data_parallel_group if allow_transport_padding else None, + allow_transport_padding=allow_transport_padding, + ) mb_list.mbs = [pack_tensor_dict(mb) for mb in mb_list.mbs] # LCM ensures page-aligned memory and exact CP slicing without extra padding. @@ -1275,7 +1292,7 @@ def _compute_logprobs_and_loss( loss_multiplier: float = 1.0, ) -> torch.Tensor: """Compute logprobs/entropy and return scaled loss.""" - local_weight = loss_weight_fn(ctx.mb_input) + local_weight = compute_microbatch_loss_weight(ctx.mb_input, loss_weight_fn) if local_weight == 0: return logits.mean() * 0.0 diff --git a/areal/models/tree_attn/tree.py b/areal/models/tree_attn/tree.py index 7b655dc7a3..cf766f6d93 100644 --- a/areal/models/tree_attn/tree.py +++ b/areal/models/tree_attn/tree.py @@ -26,7 +26,7 @@ precompute_tree_attention_data, ) from areal.utils import logging, stats_tracker -from areal.utils.data import MicroBatchList +from areal.utils.data import TRANSPORT_DUMMY_KEY, MicroBatchList from areal.utils.perf_tracer import trace_perf, trace_scope logger = logging.getLogger("TreeAttentionCore") @@ -404,6 +404,7 @@ def build_packed_tree_batch( # Build packed outputs for each tree mbs: list[dict[str, Any]] = [] + padded_mbs: list[dict[str, Any]] = [] padding_lengths: list[int] = [] padded_to_lengths: list[int] = [] @@ -448,13 +449,17 @@ def build_packed_tree_batch( non_packable_keys, ) - mb = { + padded_mb = { "input_ids": input_ids, "position_ids": position_ids, "trie_node": trie, **extra_data, } + mb = dict(padded_mb) + if not trie.all_sequence_ids: + mb[TRANSPORT_DUMMY_KEY] = True mbs.append(mb) + padded_mbs.append(padded_mb) padding_lengths.append(padded_size - num_tokens) padded_to_lengths.append(padded_size) @@ -465,7 +470,7 @@ def build_packed_tree_batch( mb_spec=mb_spec, mbs=mbs, group_lens=[num for num in num_tokens_list], - padded_mbs=mbs, + padded_mbs=padded_mbs, padding_lengths=padding_lengths, padded_to_lengths=padded_to_lengths, _max_seqlen=max(padded_to_lengths), diff --git a/areal/utils/data.py b/areal/utils/data.py index 1560742a1a..6b42d10c50 100644 --- a/areal/utils/data.py +++ b/areal/utils/data.py @@ -23,6 +23,8 @@ logger = logging.getLogger("DataUtils") +TRANSPORT_DUMMY_KEY = "_transport_dummy" + def get_batch_size(data: dict[str, Any]) -> int: if not data: @@ -621,6 +623,7 @@ class MicroBatchList: # sequence-level padding information align_to_lengths: list[int] | None = None old_cu_seqlens_list: list[torch.Tensor] | None = None + transport_dummy_count: int = 0 @property def max_seqlen(self) -> int: @@ -691,16 +694,69 @@ def to(self, *args, **kwargs): padded_to_lengths=self.padded_to_lengths, old_cu_seqlens_list=old_cu_seqlens_list, align_to_lengths=self.align_to_lengths, + transport_dummy_count=self.transport_dummy_count, ) DEFAULT_MAX_TOKENS_PER_MB = int(1e12) +def make_transport_dummy(template: dict[str, Any]) -> dict[str, Any]: + """Create one model-valid row for collective participation.""" + batch_size = get_batch_size(template) + if batch_size < 1: + raise ValueError("Cannot create transport padding from an empty batch") + + dummy: dict[str, Any] = {} + for key, value in template.items(): + if is_multi_modal_key(key) and isinstance(value, list): + dummy[key] = [{}] + elif ( + isinstance(value, torch.Tensor) + and value.ndim > 0 + and value.shape[0] == batch_size + ): + dummy[key] = torch.zeros_like(value[:1]) + elif isinstance(value, list) and len(value) == batch_size: + dummy[key] = [copy.deepcopy(value[0])] + else: + dummy[key] = copy.deepcopy(value) + + attention_mask = dummy.get("attention_mask") + if not isinstance(attention_mask, torch.Tensor) or attention_mask.ndim != 2: + raise ValueError("Transport padding requires a 2D attention_mask") + if attention_mask.shape[1] < 1: + raise ValueError("Transport padding requires sequence length >= 1") + attention_mask[:, 0] = 1 + if isinstance(dummy.get("loss_mask"), torch.Tensor): + dummy["loss_mask"].zero_() + return dummy + + +def _pad_batch_to_min_groups( + data: dict[str, Any], + *, + min_groups: int, + granularity: int, +) -> tuple[dict[str, Any], int]: + batch_size = get_batch_size(data) + if batch_size % granularity != 0: + raise RuntimeError( + f"Batch size {batch_size} cannot divide granularity {granularity}." + ) + current_groups = batch_size // granularity + pad_count = max(min_groups - current_groups, 0) * granularity + if pad_count == 0: + return data, 0 + dummies = [make_transport_dummy(data) for _ in range(pad_count)] + return concat_padded_tensors([data, *dummies]), pad_count + + def split_padded_tensor_dict_into_mb_list( data: dict[str, Any], mb_spec: MicroBatchSpec, group: dist.ProcessGroup | None = None, + allow_transport_padding: bool = False, ) -> MicroBatchList: """Split a padded dict of tensors into micro-batches based on the attention mask. @@ -708,6 +764,8 @@ def split_padded_tensor_dict_into_mb_list( data (Dict): Dictionary containing padded tensors. mb_spec (MicroBatchSpec): Specification for micro-batch splitting. group (Optional[dist.ProcessGroup]): Process group for distributed synchronization. + allow_transport_padding: Add model-valid rows when synchronized execution + requires more micro-batches than local semantic data can provide. Returns: MicroBatchList: A structure containing the split micro-batches and metadata. @@ -720,19 +778,54 @@ def split_padded_tensor_dict_into_mb_list( mb_spec, max_tokens_per_mb=DEFAULT_MAX_TOKENS_PER_MB ) granularity = mb_spec.granularity - bs = data["attention_mask"].shape[0] - if bs % granularity != 0: - raise RuntimeError(f"Batch size {bs} cannot divide granularity {granularity}.") - max_seqlen = data["attention_mask"].shape[1] - seq_lens = data["attention_mask"].sum(1).long().cpu().numpy().tolist() - input_lens = ( - data["attention_mask"] - .view(bs // granularity, granularity, -1) - .sum(dim=(1, 2)) - .long() - .cpu() - .numpy() - ) + semantic_batch_size = data["attention_mask"].shape[0] + allocation_spec = mb_spec + transport_dummy_count = 0 + target_n_mbs = max(mb_spec.n_mbs or 1, mb_spec.n_mbs_divisor) + + while True: + if allow_transport_padding: + data, added = _pad_batch_to_min_groups( + data, + min_groups=target_n_mbs, + granularity=granularity, + ) + transport_dummy_count += added + allocation_spec = MicroBatchSpec.new(mb_spec, n_mbs=target_n_mbs) + + bs = data["attention_mask"].shape[0] + if bs % granularity != 0: + raise RuntimeError( + f"Batch size {bs} cannot divide granularity {granularity}." + ) + max_seqlen = data["attention_mask"].shape[1] + seq_lens = data["attention_mask"].sum(1).long().cpu().numpy().tolist() + input_lens = ( + data["attention_mask"] + .view(bs // granularity, granularity, -1) + .sum(dim=(1, 2)) + .long() + .cpu() + .numpy() + ) + if transport_dummy_count: + input_lens[-transport_dummy_count // granularity :] = 0 + + if not allow_transport_padding: + group_indices = allocate_balanced_mbs_synced( + allocation_spec, input_lens, group=group + ) + break + + group_indices = allocate_balanced_mbs(allocation_spec, input_lens) + if not dist.is_initialized(): + break + all_n_mbs: list[int | None] = [None] * dist.get_world_size(group) + dist.all_gather_object(all_n_mbs, len(group_indices), group=group) + synchronized_n_mbs = max(n for n in all_n_mbs if n is not None) + if all(n == synchronized_n_mbs for n in all_n_mbs): + break + target_n_mbs = synchronized_n_mbs # check for multimodal input data multimodal_keys = {key for key in data if is_multi_modal_key(key)} @@ -752,7 +845,6 @@ def split_padded_tensor_dict_into_mb_list( not_to_split[key] = value # split - group_indices = allocate_balanced_mbs_synced(mb_spec, input_lens, group=group) group_indices = [ seqpack.flat2d( [list(range(i * granularity, (i + 1) * granularity)) for i in group_index] @@ -803,16 +895,28 @@ def _split(tensor): results = [] # organize splitted micro batches assert len(mbs) == len(splitted_lens), (len(mbs), len(splitted_lens)) - for i, (mb, lens) in enumerate(zip(mbs, splitted_lens)): - results.append({**mb, **not_to_split}) + for mb, indices in zip(mbs, group_indices, strict=True): + has_transport_dummy = any(index >= semantic_batch_size for index in indices) + is_transport_dummy = has_transport_dummy and all( + index >= semantic_batch_size for index in indices + ) + if has_transport_dummy and not is_transport_dummy: + raise RuntimeError( + "Transport padding must not share a micro-batch with semantic rows" + ) + result = {**mb, **not_to_split} + if is_transport_dummy: + result[TRANSPORT_DUMMY_KEY] = True + results.append(result) return MicroBatchList( data=data, - mb_spec=mb_spec, + mb_spec=allocation_spec, mbs=results, forward_indices=forward_indices, backward_indices=backward_indices.tolist(), group_lens=group_lens, + transport_dummy_count=transport_dummy_count, ) @@ -1033,6 +1137,9 @@ def pad_mb_list( pad_value=pad_value, seq_align_to=seq_align_to, ) + padded_mb = { + key: value for key, value in padded_mb.items() if key != TRANSPORT_DUMMY_KEY + } padded_mb_inputs.append(padded_mb) pad_lengths.append(pad_len) pad_to_lengths.append(pad_to_length) @@ -1338,9 +1445,10 @@ def bcast_mb_list( mb_list.padding_lengths, mb_list.padded_to_lengths, mb_list.align_to_lengths, + mb_list.transport_dummy_count, ] if mb_list - else [None for _ in range(7)] + else [None for _ in range(8)] ) dist.broadcast_object_list(to_broadcast, src=src_rank, group=group) ( @@ -1351,6 +1459,7 @@ def bcast_mb_list( padding_lengths, padded_to_lengths, align_to_lengths, + transport_dummy_count, ) = to_broadcast return MicroBatchList( data=data, @@ -1364,6 +1473,7 @@ def bcast_mb_list( padded_to_lengths=padded_to_lengths, old_cu_seqlens_list=old_cu_seqlens_list, align_to_lengths=align_to_lengths, + transport_dummy_count=transport_dummy_count, ) diff --git a/tests/test_tree_transport.py b/tests/test_tree_transport.py new file mode 100644 index 0000000000..cd5000fcb2 --- /dev/null +++ b/tests/test_tree_transport.py @@ -0,0 +1,50 @@ +import torch +import torch.distributed as dist + +from areal.api.cli_args import MicroBatchSpec +from areal.engine.core.train_engine import compute_microbatch_loss_weight +from areal.models.tree_attn.tree import build_packed_tree_batch +from areal.utils.data import TRANSPORT_DUMMY_KEY + + +def test_tree_transport_dummy_bypasses_objective_weight(monkeypatch): + data = { + "input_ids": torch.arange(4).view(1, 4), + "attention_mask": torch.ones(1, 4, dtype=torch.bool), + "loss_mask": torch.ones(1, 4, dtype=torch.bool), + } + monkeypatch.setattr(dist, "is_initialized", lambda: True) + monkeypatch.setattr(dist, "get_world_size", lambda _group=None: 2) + + def _all_gather(outputs, local_count, group=None): + del group + outputs[0].copy_(local_count) + outputs[1].fill_(2) + + monkeypatch.setattr(dist, "all_gather", _all_gather) + + mb_list = build_packed_tree_batch( + data, + MicroBatchSpec(max_tokens_per_mb=128), + ) + semantic_mb, transport_mb = mb_list.mbs + + assert TRANSPORT_DUMMY_KEY not in semantic_mb + assert transport_mb[TRANSPORT_DUMMY_KEY] is True + assert mb_list.padded_mbs is not None + assert TRANSPORT_DUMMY_KEY not in mb_list.padded_mbs[1] + + callback_called = False + + def _loss_weight(_microbatch): + nonlocal callback_called + callback_called = True + return torch.tensor(1.0) + + torch.testing.assert_close( + compute_microbatch_loss_weight(transport_mb, _loss_weight), + torch.tensor(0.0), + rtol=0.0, + atol=0.0, + ) + assert callback_called is False diff --git a/tests/test_utils.py b/tests/test_utils.py index 0c82b47517..701662fb6c 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -1,10 +1,20 @@ import pytest import torch +import areal.utils.data as data_module from areal.api.cli_args import MicroBatchSpec +from areal.engine.core.train_engine import ( + compute_microbatch_loss_weight, + reorder_and_pad_outputs, +) +from areal.trainer.dpo.dpo_engine import _dpo_loss_weight +from areal.trainer.rw.rw_engine import _rw_loss_weight from areal.utils.data import ( + TRANSPORT_DUMMY_KEY, + MicroBatchList, pack_tensor_dict, pad_and_stack_tensors_along_first_dim, + pad_mb_list, pad_sequences_to_tensors, reorder_list, split_padded_tensor_dict_into_mb_list, @@ -73,3 +83,128 @@ def test_micro_batch_split(mock_padded_data, n_mbs, max_tokens_per_mb, n_mbs_div assert torch.allclose(x, packed_data[key]) y = pad_and_stack_tensors_along_first_dim(xs) assert torch.allclose(mock_padded_data[key], y) + + +def _preference_batch() -> dict[str, torch.Tensor]: + return { + "input_ids": torch.arange(6).view(2, 3), + "attention_mask": torch.ones(2, 3, dtype=torch.bool), + } + + +@pytest.mark.parametrize("loss_weight_fn", [_dpo_loss_weight, _rw_loss_weight]) +def test_transport_padding_bypasses_objective_weight(loss_weight_fn): + mb_list = split_padded_tensor_dict_into_mb_list( + _preference_batch(), + MicroBatchSpec(n_mbs=2, granularity=2), + allow_transport_padding=True, + ) + mb_list.mbs = [pack_tensor_dict(mb) for mb in mb_list.mbs] + semantic_mb, transport_mb = sorted( + mb_list.mbs, key=lambda mb: TRANSPORT_DUMMY_KEY in mb + ) + + # A model-valid preference pair has non-zero objective weight. The transport + # marker, rather than objective-specific fields, is what makes it weightless. + torch.testing.assert_close( + loss_weight_fn(transport_mb), torch.tensor(1.0), rtol=0.0, atol=0.0 + ) + torch.testing.assert_close( + compute_microbatch_loss_weight(semantic_mb, loss_weight_fn), + torch.tensor(1.0), + rtol=0.0, + atol=0.0, + ) + torch.testing.assert_close( + compute_microbatch_loss_weight(transport_mb, loss_weight_fn), + torch.tensor(0.0), + rtol=0.0, + atol=0.0, + ) + + pad_mb_list(mb_list) + assert all( + TRANSPORT_DUMMY_KEY not in padded_mb for padded_mb in mb_list.padded_mbs or [] + ) + + +def test_noop_packed_padding_preserves_semantic_transport_marker(): + transport_mb = pack_tensor_dict( + { + "input_ids": torch.zeros(1, 1, dtype=torch.long), + "attention_mask": torch.ones(1, 1, dtype=torch.bool), + TRANSPORT_DUMMY_KEY: True, + } + ) + mb_list = MicroBatchList( + data=transport_mb, + mb_spec=MicroBatchSpec(max_tokens_per_mb=1), + mbs=[transport_mb], + group_lens=[1], + transport_dummy_count=1, + ) + + pad_mb_list(mb_list, pad_to_maximum=True) + + assert mb_list.padding_lengths == [0] + assert mb_list.mbs[0][TRANSPORT_DUMMY_KEY] is True + assert mb_list.padded_mbs is not None + assert TRANSPORT_DUMMY_KEY not in mb_list.padded_mbs[0] + + callback_called = False + + def _loss_weight(_microbatch): + nonlocal callback_called + callback_called = True + return torch.tensor(1.0) + + torch.testing.assert_close( + compute_microbatch_loss_weight(mb_list.mbs[0], _loss_weight), + torch.tensor(0.0), + rtol=0.0, + atol=0.0, + ) + assert callback_called is False + + +def test_forward_transport_padding_is_removed_from_outputs(): + data = { + "input_ids": torch.arange(3).view(1, 3), + "attention_mask": torch.ones(1, 3, dtype=torch.bool), + } + mb_list = split_padded_tensor_dict_into_mb_list( + data, + MicroBatchSpec(n_mbs=3), + allow_transport_padding=True, + ) + outputs = [mb["input_ids"][mb["attention_mask"]].float() for mb in mb_list.mbs] + + result = reorder_and_pad_outputs(outputs, [3], mb_list) + + torch.testing.assert_close( + result, torch.tensor([[0.0, 1.0, 2.0]]), rtol=0.0, atol=0.0 + ) + + +def test_transport_padding_converges_to_distributed_microbatch_count(monkeypatch): + monkeypatch.setattr(data_module.dist, "is_initialized", lambda: True) + monkeypatch.setattr(data_module.dist, "get_world_size", lambda _group=None: 2) + + def _all_gather_counts(output, local_count, group=None): + del group + output[:] = [3, local_count] + + monkeypatch.setattr(data_module.dist, "all_gather_object", _all_gather_counts) + + mb_list = split_padded_tensor_dict_into_mb_list( + { + "input_ids": torch.arange(3).view(1, 3), + "attention_mask": torch.ones(1, 3, dtype=torch.bool), + }, + MicroBatchSpec(), + allow_transport_padding=True, + ) + + assert len(mb_list.mbs) == 3 + assert mb_list.transport_dummy_count == 2 + assert sum(TRANSPORT_DUMMY_KEY in mb for mb in mb_list.mbs) == 2