diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 5d61e47ce18..dbd1c941a4a 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -80,6 +80,7 @@ def add_megatron_arguments(parser: argparse.ArgumentParser): parser = _add_msc_args(parser) parser = _add_kitchen_quantization_arguments(parser) parser = _add_sft_args(parser) + parser = _add_varlen_dataset_args(parser) parser = _add_fault_injector_args(parser) @@ -1693,7 +1694,36 @@ def validate_args(args, defaults={}): if args.ckpt_format == "fsdp_dtensor": assert args.use_megatron_fsdp, "--ckpt-format fsdp_dtensor is only tested with Megatron FSDP." - # Scheduler-name and max-seqlen validation live in + # --use-varlen-dataset: independent of --sft. Cannot be combined with --sft + # because they are mutually-exclusive top-level dataset selectors that both + # drive the packed-sequence (THD) path. These stay in validate_args: the + # selectors are CLI-level args, not core config fields. + if args.use_varlen_dataset: + assert not args.sft, ( + "--use-varlen-dataset and --sft are mutually exclusive; both " + "select the packed-sequence dataset family. Pick one." + ) + if args.varlen_sbhd_validation: + assert args.sequence_packing_scheduler is None, ( + "--varlen-sbhd-validation does not use a sequence packing " + "scheduler; drop --sequence-packing-scheduler." + ) + # SBHD validation is a real-data numerical-reference path only; + # MockVarlenDataset does not implement it. + assert not args.mock_data, ( + "--varlen-sbhd-validation is not supported with --mock-data; " + "SBHD validation requires a real dataset." + ) + else: + # VarlenDataset emits one unpacked sample per __getitem__; it + # relies on an upstream packing scheduler to group variable-length + # samples into THD batches. Auto-pick ``dp_balanced`` when the + # user did not request one explicitly. + if args.sequence_packing_scheduler is None: + args.sequence_packing_scheduler = 'dp_balanced' + + # Runs after the varlen auto-select above so it sees the final resolved + # scheduler. Scheduler-name and max-seqlen validation live in # ModelParallelConfig.__post_init__; only the buffer-size check stays here # because seq_length is not a core config field. The None case for # max_seqlen_per_dp_cp_rank is rejected by the config check. @@ -3891,6 +3921,37 @@ def _add_sft_args(parser): 'lognormal_sigma=1.1.') return parser + +def _add_varlen_dataset_args(parser): + group = parser.add_argument_group(title='varlen dataset') + group.add_argument('--use-varlen-dataset', action="store_true", + help='Train with VarlenDataset, a variable-length packed (THD) dataset ' + 'that consumes instruction-tuning data from a HuggingFace Hub repo id, ' + 'a local parquet file, or a local jsonl file. Schema (alpaca / sharegpt ' + '/ openai-messages) is auto-detected from the dataset columns. ' + 'Mutually exclusive with --sft. Auto-picks a sequence packing ' + 'scheduler when none is given: dp_balanced. ' + 'Combine with --mock-data for a synthetic lognormal sequence-length ' + 'distribution; see --varlen-mock-dataset-config-json.') + group.add_argument('--varlen-sbhd-validation', action="store_true", + help='Reference SBHD mode for THD numerical verification. When set, ' + 'VarlenDataset emits SBHD-style samples right-padded to ' + '--seq-length (no cu_seqlens, no packing scheduler), so the run can ' + 'be compared against the THD path to validate correctness. ' + 'Incompatible with --sequence-packing-scheduler.') + group.add_argument('--varlen-mock-dataset-config-json', type=str, default=None, + help='Mock-dataset config for --use-varlen-dataset --mock-data. ' + 'Accepts either an inline JSON literal or a path to a JSON file ' + 'containing the same schema as --sft-mock-dataset-config-json: either ' + '{"mode":"file","path":"/path/to/lengths.csv"}, ' + '{"mode":"distribution","type":"lognormal","min_seq_len":1024,' + '"max_seq_len":2048,"mean_seq_len":1536,"lognormal_sigma":1.1}, or ' + '{"mode":"verification","data_path":"/prefix/of/IndexedDataset"}. ' + 'If not specified, defaults to a lognormal distribution with ' + 'min_seq_len=seq_length//2, max_seq_len=seq_length, ' + 'mean_seq_len=seq_length*3//4, lognormal_sigma=1.1.') + return parser + def _add_logits_distillation_args(parser): group = parser.add_argument_group(title='Logits Distillation') diff --git a/megatron/training/datasets/data_samplers.py b/megatron/training/datasets/data_samplers.py index d51d9c6c8a2..33cbb94b432 100644 --- a/megatron/training/datasets/data_samplers.py +++ b/megatron/training/datasets/data_samplers.py @@ -52,7 +52,7 @@ def build_pretraining_data_loader(dataset, consumed_samples): data_parallel_rank=mpu.get_data_parallel_rank(), data_parallel_size=mpu.get_data_parallel_world_size()) elif args.dataloader_type == 'single': - if args.hybrid_context_parallel: + if args.hybrid_context_parallel and args.sequence_packing_scheduler is None: batch_sampler = HybridCPMegatronPretrainingSampler( total_samples=len(dataset), consumed_samples=consumed_samples, @@ -61,7 +61,8 @@ def build_pretraining_data_loader(dataset, consumed_samples): data_parallel_rank=mpu.get_data_parallel_rank(), data_parallel_size=mpu.get_data_parallel_world_size()) else: - # Megatron sampler + # Megatron sampler. Packing schedulers consume one microbatch at a + # time and form packed global batches themselves. batch_sampler = MegatronPretrainingSampler( total_samples=len(dataset), consumed_samples=consumed_samples, @@ -103,9 +104,17 @@ def close_nvidia_fds(): maybe_worker_init_fn = ( worker_init_fn if args.num_workers > 0 else None ) - # Torch dataloader. - if args.hybrid_context_parallel: - extra_kwargs = {"collate_fn": lambda x: x,} + # Identity collate for VarlenDataset and packing-scheduler paths; they emit + # one variable-length dict per sample, not stack-able by the default + # collate. --varlen-sbhd-validation is excluded: it bypasses packing and + # emits fixed-length [seq_length] samples that the default collate stacks + # normally. + if ( + (args.use_varlen_dataset and not args.varlen_sbhd_validation) + or args.hybrid_context_parallel + or args.sequence_packing_scheduler is not None + ): + extra_kwargs = {"collate_fn": lambda x: x} else: extra_kwargs = {} return torch.utils.data.DataLoader( diff --git a/megatron/training/training.py b/megatron/training/training.py index 53a3b98940f..1491d8bc6de 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -43,7 +43,7 @@ # First-party. from megatron.core._rank_utils import safe_get_rank from megatron.core import mpu, nccl_allocator, tensor_parallel -from megatron.core.datasets.data_schedule import HybridCPDataLoaderWrapper +from megatron.core.datasets.data_schedule import HybridCPDataLoaderWrapper, wrap_data_iterator from megatron.core.distributed import DistributedDataParallel as DDP from megatron.core.distributed import ( DistributedDataParallelConfig, @@ -261,6 +261,7 @@ # never call ``update_*`` so the flag stays ``False`` and no collective fires. _seqlen_stats_in_iteration: Optional[torch.Tensor] = None _seqlen_stats_active: bool = False +_seqlen_stats_are_global: bool = False # Only report memory for first 3 checkpoint saves. num_checkpoints_memory_reported = 0 @@ -728,7 +729,7 @@ def update_seqlen_stats_from_cu_seqlens(cu_seqlens): the all-reduce; BSHD callers that never invoke this function leave the flag at ``False`` and pay zero collective cost. """ - global _seqlen_stats_in_iteration, _seqlen_stats_active + global _seqlen_stats_in_iteration, _seqlen_stats_active, _seqlen_stats_are_global if cu_seqlens is None or cu_seqlens.numel() < 2: return # Pin the accumulator to the current CUDA device when available so the @@ -747,6 +748,21 @@ def update_seqlen_stats_from_cu_seqlens(cu_seqlens): _seqlen_stats_in_iteration[0] += seqlens.sum() _seqlen_stats_in_iteration[1] += (seqlens * seqlens).sum() _seqlen_stats_active = True + _seqlen_stats_are_global = False + + +def set_seqlen_stats_in_iteration(total_real_tokens, seqlen_squared_sum): + """Seed per-iteration THD FLOPs stats that were already computed globally.""" + global _seqlen_stats_in_iteration, _seqlen_stats_active, _seqlen_stats_are_global + if total_real_tokens is None or seqlen_squared_sum is None: + return + if _seqlen_stats_in_iteration is None: + device = torch.device(f'cuda:{torch.cuda.current_device()}') if torch.cuda.is_available() else 'cpu' + _seqlen_stats_in_iteration = torch.zeros(2, dtype=torch.float64, device=device) + _seqlen_stats_in_iteration[0] = float(total_real_tokens) + _seqlen_stats_in_iteration[1] = float(seqlen_squared_sum) + _seqlen_stats_active = True + _seqlen_stats_are_global = True def consume_seqlen_stats_in_iteration() -> Tuple[Optional[float], Optional[float]]: @@ -774,13 +790,15 @@ def consume_seqlen_stats_in_iteration() -> Tuple[Optional[float], Optional[float replicated across TP/CP/PP); the world all-reduce therefore overcounts by a factor of ``TP * CP * PP``, which we divide out. """ - global _seqlen_stats_in_iteration, _seqlen_stats_active + global _seqlen_stats_in_iteration, _seqlen_stats_active, _seqlen_stats_are_global if not _seqlen_stats_active: # BSHD path: never allocated the tensor; tell the caller to use the # closed-form defaults. return None, None t = _seqlen_stats_in_iteration - if torch.distributed.is_initialized() and mpu.model_parallel_is_initialized(): + if _seqlen_stats_are_global: + dedup = 1 + elif torch.distributed.is_initialized() and mpu.model_parallel_is_initialized(): torch.distributed.all_reduce(t) tp_size = max(mpu.get_tensor_model_parallel_world_size(), 1) cp_size = max(mpu.get_context_parallel_world_size(), 1) @@ -796,6 +814,7 @@ def consume_seqlen_stats_in_iteration() -> Tuple[Optional[float], Optional[float # iterations reuse it without reallocating. t.zero_() _seqlen_stats_active = False + _seqlen_stats_are_global = False return total_real_tokens / dedup, seqlen_squared_sum / dedup @@ -3094,6 +3113,20 @@ def train_step(forward_step_func, data_iterator, model, optimizer, opt_param_sch for optim_instance in mxfp8_overlap_optimizers: optim_instance._copy_main_params_to_param_buffer() + if getattr(config, "sequence_packing_scheduler", None) is not None: + ( + data_iterator, + scheduled_num_microbatches, + total_real_tokens_in_batch, + seqlen_squared_sum_in_batch, + ) = wrap_data_iterator(data_iterator, config, get_num_microbatches()) + set_seqlen_stats_in_iteration( + total_real_tokens_in_batch, + seqlen_squared_sum_in_batch, + ) + else: + scheduled_num_microbatches = get_num_microbatches() + # Forward pass. if save_activations_in_this_iteration: enable_activation_logging(model, args.save) @@ -3103,7 +3136,7 @@ def train_step(forward_step_func, data_iterator, model, optimizer, opt_param_sch enable_dgrad_logging(model, args.save) grad_context, forward_only = _forward_backward_grad_context(args) _fb_cm = ( - span_cm("megatron.train.iteration.forward_backward", tracer=_otel_step_tracer, num_microbatches=get_num_microbatches()) + span_cm("megatron.train.iteration.forward_backward", tracer=_otel_step_tracer, num_microbatches=scheduled_num_microbatches) if _otel_sg_enabled('forward_backward') and _otel_step_tracer is not None else nullcontext() ) with grad_context, _fb_cm: @@ -3111,7 +3144,7 @@ def train_step(forward_step_func, data_iterator, model, optimizer, opt_param_sch forward_step_func=forward_step_func, data_iterator=data_iterator, model=model, - num_microbatches=get_num_microbatches(), + num_microbatches=scheduled_num_microbatches, seq_length=args.seq_length, micro_batch_size=args.micro_batch_size, decoder_seq_length=args.decoder_seq_length, @@ -4665,6 +4698,9 @@ def trace_handler(p): # Completely skip iteration if needed. if (iteration + 1) in args.iterations_to_skip: + assert ( + getattr(config, "sequence_packing_scheduler", None) is None + ), "Sequence packing scheduler is not supported in skip iteration mode" # Dummy train_step to fast forward train_data_iterator. dummy_train_step(train_data_iterator) if iteration == start_iteration: @@ -5206,13 +5242,23 @@ def evaluate( # Don't care about timing during evaluation config.timers = None ft_integration.on_eval_step_start() + if getattr(config, "sequence_packing_scheduler", None) is not None: + try: + (packed_data_iterator, scheduled_eval_num_microbatches, _, _) = ( + wrap_data_iterator(data_iterator, config, eval_num_microbatches) + ) + except StopIteration: + break + else: + packed_data_iterator = data_iterator + scheduled_eval_num_microbatches = eval_num_microbatches with _otel_managed_span('evaluate', 'megatron.evaluate.step', **{'megatron.eval_iteration': iteration}): loss_dicts = forward_backward_func( forward_step_func=forward_step_func, - data_iterator=data_iterator, + data_iterator=packed_data_iterator, model=model, - num_microbatches=eval_num_microbatches, + num_microbatches=scheduled_eval_num_microbatches, seq_length=args.seq_length, micro_batch_size=eval_micro_batch_size, decoder_seq_length=args.decoder_seq_length, diff --git a/pretrain_gpt.py b/pretrain_gpt.py index f6e8f5d904c..ab0f495b44e 100644 --- a/pretrain_gpt.py +++ b/pretrain_gpt.py @@ -39,6 +39,7 @@ def _rank0_only_showwarning(message, category, filename, lineno, file=None, line from gpt_builders import gpt_builder from megatron.core import mpu from megatron.core.datasets.blended_megatron_dataset_builder import BlendedMegatronDatasetBuilder +from megatron.core.datasets.data_schedule import get_batch_on_this_rank_for_sequence_packing from megatron.core.datasets.gpt_dataset import GPTDataset, GPTDatasetConfig, MockGPTDataset from megatron.core.enums import ModelType from megatron.core.package_info import __version__ as mcore_version @@ -75,6 +76,7 @@ def _rank0_only_showwarning(message, category, filename, lineno, file=None, line from megatron.training.arguments import core_transformer_config_from_args, parse_and_validate_args from megatron.training.datasets.fim_dataset import GPTFIMDataset, GPTFIMDatasetConfig from megatron.training.datasets.sft_dataset import MockSFTDataset, SFTDataset +from megatron.training.datasets.varlen_dataset import MockVarlenDataset, VarlenDataset from megatron.training.training import update_seqlen_stats_from_cu_seqlens from megatron.training.utils import get_blend_and_blend_per_split, is_first_or_last_pipeline_stage from model_provider import model_provider @@ -113,6 +115,19 @@ def get_batch(data_iterator, vp_stage: Optional[int] = None): args = get_args() config = core_transformer_config_from_args(args) + if args.sequence_packing_scheduler is not None: + return get_batch_on_this_rank_for_sequence_packing( + data_iterator, + vpp_size=config.virtual_pipeline_model_parallel_size, + mtp_on_this_rank=mtp_on_this_rank_func( + layout=config.pipeline_model_parallel_layout, + mtp_num_layers=config.mtp_num_layers, + ignore_virtual=False, + vp_stage=vp_stage, + ), + vp_stage=vp_stage, + ) + cp_size = args.context_parallel_size tp_rank = mpu.get_tensor_model_parallel_rank() is_sft = args.sft @@ -307,44 +322,61 @@ def forward_step(data_iterator, model: GPTModel, return_schedule_plan: bool = Fa timers('batch-generator', log_level=2).start() with stimer(bdata=True): vp_stage = get_attr_wrapped_model(model, "vp_stage") - ( - attention_mask, - cu_seqlens, - cu_seqlens_padded, - hybrid_cp_group, - labels, - local_cp_size, - loss_mask, - max_seqlen, - position_ids, - tokens, - ) = get_batch(data_iterator, vp_stage) - - packed_seq_params = None - if cu_seqlens is not None: - # Squeeze the batch dim: the batch dict keeps cu_seqlens as (1, N) - # for consistency, but PackedSeqParams and TE expect 1-D. - cu_seqlens = cu_seqlens.squeeze(0) - if cu_seqlens_padded is not None: - cu_seqlens_padded = cu_seqlens_padded.squeeze(0) - # Use real (unpadded) cu_seqlens to feed the FLOPs accounting: varlen - # attention only computes work for real tokens within each chunk. - update_seqlen_stats_from_cu_seqlens(cu_seqlens) - cu_seqlens_for_params = ( - cu_seqlens_padded if cu_seqlens_padded is not None else cu_seqlens - ) # TODO(asolergi-nv): Currently there is a bug forcing cu_seqlens to be cu_seqlens_padded - packed_seq_params = PackedSeqParams( - qkv_format="thd", - cu_seqlens_q=cu_seqlens_for_params, - cu_seqlens_kv=cu_seqlens_for_params, - cu_seqlens_q_padded=cu_seqlens_padded, - cu_seqlens_kv_padded=cu_seqlens_padded, - max_seqlen_q=int(max_seqlen.item()), - max_seqlen_kv=int(max_seqlen.item()), - local_cp_size=int(local_cp_size.item()) if local_cp_size is not None else None, - cp_group=hybrid_cp_group, - tokens_per_sample=args.seq_length, - ) + batch = get_batch(data_iterator, vp_stage) + + if len(batch) == 7: + ( + tokens, + labels, + loss_mask, + attention_mask, + position_ids, + packed_seq_params, + padding_mask, + ) = batch + elif len(batch) == 6: + tokens, labels, loss_mask, attention_mask, position_ids, packed_seq_params = batch + padding_mask = None + else: + ( + attention_mask, + cu_seqlens, + cu_seqlens_padded, + hybrid_cp_group, + labels, + local_cp_size, + loss_mask, + max_seqlen, + position_ids, + tokens, + ) = batch + + padding_mask = None + packed_seq_params = None + if cu_seqlens is not None: + # Squeeze the batch dim: the batch dict keeps cu_seqlens as (1, N) + # for consistency, but PackedSeqParams and TE expect 1-D. + cu_seqlens = cu_seqlens.squeeze(0) + if cu_seqlens_padded is not None: + cu_seqlens_padded = cu_seqlens_padded.squeeze(0) + # Use real (unpadded) cu_seqlens to feed the FLOPs accounting: varlen + # attention only computes work for real tokens within each chunk. + update_seqlen_stats_from_cu_seqlens(cu_seqlens) + cu_seqlens_for_params = ( + cu_seqlens_padded if cu_seqlens_padded is not None else cu_seqlens + ) # TODO(asolergi-nv): Currently there is a bug forcing cu_seqlens to be cu_seqlens_padded + packed_seq_params = PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=cu_seqlens_for_params, + cu_seqlens_kv=cu_seqlens_for_params, + cu_seqlens_q_padded=cu_seqlens_padded, + cu_seqlens_kv_padded=cu_seqlens_padded, + max_seqlen_q=int(max_seqlen.item()), + max_seqlen_kv=int(max_seqlen.item()), + local_cp_size=int(local_cp_size.item()) if local_cp_size is not None else None, + cp_group=hybrid_cp_group, + tokens_per_sample=args.seq_length, + ) timers('batch-generator').stop() @@ -354,7 +386,13 @@ def forward_step(data_iterator, model: GPTModel, return_schedule_plan: bool = Fa args.overlap_moe_expert_parallel_comm ), "overlap_moe_expert_parallel_comm must be enabled to return the schedule plan" schedule_plan = model.build_schedule_plan( - tokens, position_ids, attention_mask, labels=labels, loss_mask=loss_mask + tokens, + position_ids, + attention_mask, + labels=labels, + loss_mask=loss_mask, + packed_seq_params=packed_seq_params, + padding_mask=padding_mask, ) return schedule_plan, partial(loss_func, loss_mask, model=model) else: @@ -365,6 +403,7 @@ def forward_step(data_iterator, model: GPTModel, return_schedule_plan: bool = Fa labels=labels, loss_mask=loss_mask, packed_seq_params=packed_seq_params, + padding_mask=padding_mask, ) # [ModelOpt]: model is needed to access ModelOpt distillation losses @@ -429,6 +468,8 @@ def core_gpt_dataset_config_from_args(args: Any) -> GPTDatasetConfig: "hybrid_context_parallel": args.hybrid_context_parallel, "inter_document_masking": args.dataloader_inter_document_masking, "sft_mock_dataset_config_json": args.sft_mock_dataset_config_json, + "varlen_mock_dataset_config_json": args.varlen_mock_dataset_config_json, + "varlen_sbhd_validation": args.varlen_sbhd_validation, } # add FIM args to the config @@ -472,6 +513,17 @@ def train_valid_test_datasets_provider(train_val_test_num_samples, vp_stage=None else: dataset_type = SFTDataset is_packed_sequence = True # SFT always uses packed sequence + elif args.use_varlen_dataset: + # Variable-length packed (THD) dataset, independent of --sft. + # Reuses SFTDataset's THD packing internally but is gated + # by its own top-level flag. + if args.mock_data: + dataset_type = MockVarlenDataset + else: + dataset_type = VarlenDataset + # SBHD validation mode runs the non-packed pipeline; THD mode + # is the packed-sequence path. + is_packed_sequence = not args.varlen_sbhd_validation else: if args.mock_data: dataset_type = MockGPTDataset diff --git a/pretrain_hybrid.py b/pretrain_hybrid.py index 7f70205a51c..5156993a667 100644 --- a/pretrain_hybrid.py +++ b/pretrain_hybrid.py @@ -38,6 +38,7 @@ def _rank0_only_showwarning(message, category, filename, lineno, file=None, line from hybrid_builders import hybrid_builder from megatron.core import mpu from megatron.core.datasets.blended_megatron_dataset_builder import BlendedMegatronDatasetBuilder +from megatron.core.datasets.data_schedule import get_batch_on_this_rank_for_sequence_packing from megatron.core.datasets.gpt_dataset import GPTDataset, GPTDatasetConfig, MockGPTDataset from megatron.core.enums import ModelType from megatron.core.package_info import __version__ as mcore_version @@ -75,6 +76,7 @@ def _rank0_only_showwarning(message, category, filename, lineno, file=None, line ) from megatron.training.arguments import core_transformer_config_from_args, parse_and_validate_args from megatron.training.datasets.sft_dataset import SFTDataset +from megatron.training.datasets.varlen_dataset import MockVarlenDataset, VarlenDataset from megatron.training.training import update_seqlen_stats_from_cu_seqlens from megatron.training.utils import get_blend_and_blend_per_split, is_first_or_last_pipeline_stage from model_provider import model_provider @@ -113,6 +115,19 @@ def get_batch(data_iterator, vp_stage=None): args = get_args() config = core_transformer_config_from_args(args) + if args.sequence_packing_scheduler is not None: + return get_batch_on_this_rank_for_sequence_packing( + data_iterator, + vpp_size=config.virtual_pipeline_model_parallel_size, + mtp_on_this_rank=mtp_on_this_rank_func( + layout=config.pipeline_model_parallel_layout, + mtp_num_layers=config.mtp_num_layers, + ignore_virtual=False, + vp_stage=vp_stage, + ), + vp_stage=vp_stage, + ) + cp_size = args.context_parallel_size tp_rank = mpu.get_tensor_model_parallel_rank() is_sft = args.sft @@ -305,46 +320,65 @@ def forward_step(data_iterator, model: HybridModel): with stimer(bdata=True): vp_stage = get_attr_wrapped_model(model, "vp_stage") - ( - attention_mask, - cu_seqlens, - cu_seqlens_padded, - hybrid_cp_group, - labels, - local_cp_size, - loss_mask, - max_seqlen, - position_ids, - tokens, - ) = get_batch(data_iterator, vp_stage) - - packed_seq_params = None - if cu_seqlens is not None: - # Squeeze the batch dim: the batch dict keeps cu_seqlens as (1, N) - # for consistency, but PackedSeqParams and TE expect 1-D. - cu_seqlens = cu_seqlens.squeeze(0) - if cu_seqlens_padded is not None: - cu_seqlens_padded = cu_seqlens_padded.squeeze(0) - # Use real (unpadded) cu_seqlens to feed the FLOPs accounting: varlen - # attention only computes work for real tokens within each chunk. - update_seqlen_stats_from_cu_seqlens(cu_seqlens) - cu_seqlens_for_params = cu_seqlens_padded if cu_seqlens_padded is not None else cu_seqlens - total_tokens = int(cu_seqlens_for_params[-1].item()) - if args.linear_cp_layout == "contiguous" and args.context_parallel_size > 1: - cu_seqlens_for_params = cu_seqlens - packed_seq_params = PackedSeqParams( - qkv_format="thd", - cu_seqlens_q=cu_seqlens_for_params, - cu_seqlens_kv=cu_seqlens_for_params, - cu_seqlens_q_padded=cu_seqlens_padded, - cu_seqlens_kv_padded=cu_seqlens_padded, - max_seqlen_q=int(max_seqlen.item()), - max_seqlen_kv=int(max_seqlen.item()), - local_cp_size=int(local_cp_size.item()) if local_cp_size is not None else None, - cp_group=hybrid_cp_group, - total_tokens=total_tokens, - tokens_per_sample=args.seq_length, - ) + batch = get_batch(data_iterator, vp_stage) + + if len(batch) == 7: + ( + tokens, + labels, + loss_mask, + attention_mask, + position_ids, + packed_seq_params, + padding_mask, + ) = batch + elif len(batch) == 6: + tokens, labels, loss_mask, attention_mask, position_ids, packed_seq_params = batch + padding_mask = None + else: + ( + attention_mask, + cu_seqlens, + cu_seqlens_padded, + hybrid_cp_group, + labels, + local_cp_size, + loss_mask, + max_seqlen, + position_ids, + tokens, + ) = batch + + padding_mask = None + packed_seq_params = None + if cu_seqlens is not None: + # Squeeze the batch dim: the batch dict keeps cu_seqlens as (1, N) + # for consistency, but PackedSeqParams and TE expect 1-D. + cu_seqlens = cu_seqlens.squeeze(0) + if cu_seqlens_padded is not None: + cu_seqlens_padded = cu_seqlens_padded.squeeze(0) + # Use real (unpadded) cu_seqlens to feed the FLOPs accounting: varlen + # attention only computes work for real tokens within each chunk. + update_seqlen_stats_from_cu_seqlens(cu_seqlens) + cu_seqlens_for_params = ( + cu_seqlens_padded if cu_seqlens_padded is not None else cu_seqlens + ) + total_tokens = int(cu_seqlens_for_params[-1].item()) + if args.linear_cp_layout == "contiguous" and args.context_parallel_size > 1: + cu_seqlens_for_params = cu_seqlens + packed_seq_params = PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=cu_seqlens_for_params, + cu_seqlens_kv=cu_seqlens_for_params, + cu_seqlens_q_padded=cu_seqlens_padded, + cu_seqlens_kv_padded=cu_seqlens_padded, + max_seqlen_q=int(max_seqlen.item()), + max_seqlen_kv=int(max_seqlen.item()), + local_cp_size=int(local_cp_size.item()) if local_cp_size is not None else None, + cp_group=hybrid_cp_group, + total_tokens=total_tokens, + tokens_per_sample=args.seq_length, + ) timers('batch-generator').stop() @@ -356,6 +390,7 @@ def forward_step(data_iterator, model: HybridModel): labels=labels, packed_seq_params=packed_seq_params, loss_mask=loss_mask, + padding_mask=padding_mask, ) # [ModelOpt]: model is needed to access ModelOpt distillation losses @@ -419,6 +454,8 @@ def core_gpt_dataset_config_from_args(args: Any) -> GPTDatasetConfig: sequence_parallel_size=args.tensor_model_parallel_size * args.sequence_parallel, hybrid_context_parallel=args.hybrid_context_parallel, inter_document_masking=args.dataloader_inter_document_masking, + varlen_mock_dataset_config_json=args.varlen_mock_dataset_config_json, + varlen_sbhd_validation=args.varlen_sbhd_validation, ) @@ -435,6 +472,17 @@ def train_valid_test_datasets_provider(train_val_test_num_samples, vp_stage=None if args.sft: dataset_type = SFTDataset is_packed_sequence = True # SFT always uses packed sequence + elif args.use_varlen_dataset: + # Variable-length packed (THD) dataset, independent of --sft. + # Reuses SFTDataset's THD packing internally but is gated + # by its own top-level flag. + if args.mock_data: + dataset_type = MockVarlenDataset + else: + dataset_type = VarlenDataset + # SBHD validation mode runs the non-packed pipeline; THD mode + # is the packed-sequence path. + is_packed_sequence = not args.varlen_sbhd_validation else: if args.mock_data: dataset_type = MockGPTDataset diff --git a/tests/unit_tests/data/test_varlen_dataset.py b/tests/unit_tests/data/test_varlen_dataset.py index ea47bd770f5..368d67da992 100644 --- a/tests/unit_tests/data/test_varlen_dataset.py +++ b/tests/unit_tests/data/test_varlen_dataset.py @@ -585,3 +585,200 @@ def test_mock_getitem_thd_keys_and_pad_fallback(): n = out["tokens"].numel() assert out["labels"].numel() == n and out["loss_mask"].numel() == n assert int(out["padded_seq_len"].item()) % 4 == 0 + + +# ---------------------------------------------------------------------------- +# THD handoff: _unpack_batch contract for VarlenDataset-style samples +# +# VarlenDataset already emits one unpacked sub-sample carrying ``padded_seq_len``, +# so _unpack_batch must short-circuit (no cu_seqlens slicing) and only normalize +# the collate batch dim. SFTDataset-style pre-packed samples (cu_seqlens, no +# padded_seq_len) still take the slicing path. +# ---------------------------------------------------------------------------- + + +def test_unpack_batch_short_circuits_for_varlen_samples(): + from megatron.core.datasets.data_schedule_utils import _unpack_batch + + # Two VarlenDataset-style samples, each already a single sub-sample with a + # leading batch dim (as added by the default collate_fn) and padded_seq_len. + batch = [ + { + "tokens": torch.arange(4, dtype=torch.int64).view(1, 4), + "labels": torch.arange(4, dtype=torch.int64).view(1, 4), + "loss_mask": torch.ones(1, 4), + "position_ids": torch.arange(4, dtype=torch.int64).view(1, 4), + "padded_seq_len": torch.tensor([4], dtype=torch.int32), + }, + { + "tokens": torch.arange(8, dtype=torch.int64).view(1, 8), + "labels": torch.arange(8, dtype=torch.int64).view(1, 8), + "loss_mask": torch.ones(1, 8), + "position_ids": torch.arange(8, dtype=torch.int64).view(1, 8), + "padded_seq_len": torch.tensor([8], dtype=torch.int32), + "original_seq_len": torch.tensor([8], dtype=torch.int32), + }, + ] + out = _unpack_batch(batch) + # Short-circuit: same number of samples (no slicing into sub-samples). + assert len(out) == 2 + # Leading collate batch dim dropped. + assert out[0]["tokens"].shape == (4,) + assert out[1]["tokens"].shape == (8,) + # Missing original_seq_len synthesized from padded_seq_len. + assert "original_seq_len" in out[0] + assert int(out[0]["original_seq_len"].item()) == 4 + # Existing original_seq_len preserved. + assert int(out[1]["original_seq_len"].item()) == 8 + + +def test_unpack_batch_slices_prepacked_cu_seqlens_samples(): + from megatron.core.datasets.data_schedule_utils import _unpack_batch + + # SFTDataset-style pre-packed sample: two sub-sequences [0:3) and [3:5), + # described by cu_seqlens, NO padded_seq_len -> takes the slicing path. + batch = [ + { + "tokens": torch.arange(5, dtype=torch.int64), + "labels": torch.arange(5, dtype=torch.int64), + "loss_mask": torch.ones(5), + "position_ids": torch.arange(5, dtype=torch.int64), + "cu_seqlens": torch.tensor([0, 3, 5], dtype=torch.int32), + } + ] + out = _unpack_batch(batch) + # One packed sample with two sub-sequences -> two unpacked samples. + assert len(out) == 2 + assert out[0]["tokens"].numel() == 3 + assert out[1]["tokens"].numel() == 2 + assert int(out[0]["padded_seq_len"].item()) == 3 + assert int(out[1]["padded_seq_len"].item()) == 2 + + +# ---------------------------------------------------------------------------- +# DataLoader collate selection (distributed; run under torch.distributed.run). +# +# Validates the build_pretraining_data_loader contract for the varlen paths: +# * --varlen-sbhd-validation emits fixed-length [seq_length] samples that the +# DEFAULT collate stacks into a [mbs, seq_length] batch. +# * The THD path (--use-varlen-dataset without SBHD) uses the identity collate +# (variable-length dicts are returned as a list, not stacked). +# ---------------------------------------------------------------------------- + + +def _build_varlen_for_loader(items, config, num_samples): + from megatron.core.datasets.utils import Split + + ds = VarlenDataset.__new__(VarlenDataset) + ds.config = config + ds.dataset = items + ds.indices = np.arange(len(items)) + ds.num_samples = num_samples + ds.index_split = Split.train + return ds + + +def _loader_args(*, use_varlen, sbhd, scheduler, mbs, gbs=None): + return SimpleNamespace( + dataloader_type='single', + micro_batch_size=mbs, + global_batch_size=mbs if gbs is None else gbs, + full_validation=False, + num_workers=0, + use_varlen_dataset=use_varlen, + varlen_sbhd_validation=sbhd, + sequence_packing_scheduler=scheduler, + ) + + +def test_sbhd_validation_dataloader_uses_default_collate(): + from megatron.core import parallel_state + from megatron.training.datasets.data_samplers import build_pretraining_data_loader + from megatron.training.global_vars import destroy_global_vars, set_args + from tests.unit_tests.test_utilities import Utils + + Utils.initialize_model_parallel(1, 1) + try: + tok = _FakeTokenizer(eod=0, pad=7) + seq_len, mbs = 16, 2 + # One global batch needs micro_batch_size * data_parallel_size samples; + # size the dataset off the runtime DP world size so this passes under + # any --nproc-per-node (the CI default is 8 ranks -> dp=8). + dp = parallel_state.get_data_parallel_world_size() + n = mbs * dp * 4 + cfg = _make_config(tok, seq_length=seq_len, sbhd=True) + ds = _build_varlen_for_loader(["hello world"] * n, cfg, num_samples=n) + set_args(_loader_args(use_varlen=True, sbhd=True, scheduler=None, mbs=mbs)) + loader = build_pretraining_data_loader(ds, consumed_samples=0) + batch = next(iter(loader)) + # Default collate stacks fixed-length SBHD samples into a tensor batch. + assert isinstance(batch, dict) + assert batch["tokens"].shape == (mbs, seq_len) + assert batch["labels"].shape == (mbs, seq_len) + assert batch["loss_mask"].shape == (mbs, seq_len) + finally: + destroy_global_vars() + Utils.destroy_model_parallel() + + +def test_thd_dataloader_uses_identity_collate(): + from megatron.core import parallel_state + from megatron.training.datasets.data_samplers import build_pretraining_data_loader + from megatron.training.global_vars import destroy_global_vars, set_args + from tests.unit_tests.test_utilities import Utils + + Utils.initialize_model_parallel(1, 1) + try: + tok = _FakeTokenizer(eod=0, pad=7) + mbs = 2 + dp = parallel_state.get_data_parallel_world_size() + n = mbs * dp * 4 + cfg = _make_config(tok, seq_length=64, sbhd=False) + # Variable-length samples so identity collate is required. + variable = ["a", "abcdef", "xy", "qwerty"] + items = [variable[i % len(variable)] for i in range(n)] + ds = _build_varlen_for_loader(items, cfg, num_samples=n) + set_args(_loader_args(use_varlen=True, sbhd=False, scheduler="dp_balanced", mbs=mbs)) + loader = build_pretraining_data_loader(ds, consumed_samples=0) + batch = next(iter(loader)) + # Identity collate returns the raw list of per-sample dicts (unstacked). + assert isinstance(batch, list) + assert len(batch) == mbs + assert "padded_seq_len" in batch[0] + finally: + destroy_global_vars() + Utils.destroy_model_parallel() + + +def test_packing_scheduler_dataloader_yields_microbatches(): + from megatron.core import parallel_state + from megatron.training.datasets.data_samplers import build_pretraining_data_loader + from megatron.training.global_vars import destroy_global_vars, set_args + from tests.unit_tests.test_utilities import Utils + + Utils.initialize_model_parallel(1, 1) + try: + tok = _FakeTokenizer(eod=0, pad=7) + mbs = 2 + num_microbatches = 3 + dp = parallel_state.get_data_parallel_world_size() + gbs = mbs * dp * num_microbatches + n = gbs * 2 + cfg = _make_config(tok, seq_length=64, dp=dp, cp=1) + variable = ["a", "abcdef", "xy", "qwerty"] + items = [variable[i % len(variable)] for i in range(n)] + ds = _build_varlen_for_loader(items, cfg, num_samples=n) + set_args( + _loader_args(use_varlen=True, sbhd=False, scheduler="dp_balanced", mbs=mbs, gbs=gbs) + ) + loader = build_pretraining_data_loader(ds, consumed_samples=0) + batch = next(iter(loader)) + # The packing scheduler calls next(data_iterator) num_microbatches times; + # each loader step must therefore be one local microbatch, not all + # local samples from the global batch. + assert isinstance(batch, list) + assert len(batch) == mbs + assert "padded_seq_len" in batch[0] + finally: + destroy_global_vars() + Utils.destroy_model_parallel()