diff --git a/megatron/training/datasets/sft_dataset.py b/megatron/training/datasets/sft_dataset.py index 2cbc4e424eb..b313dafb0ec 100644 --- a/megatron/training/datasets/sft_dataset.py +++ b/megatron/training/datasets/sft_dataset.py @@ -142,19 +142,12 @@ def extend_with_padding(tokens, targets, positions, pad_len): # Handle any necessary truncation if len(pack_tokens) >= pack_length + 1: # +1 here to account for later alignment - truncate_left_not_right = True # TODO(duncan): plumb this switch in - if truncate_left_not_right: # Retain existing eod - max_body = pack_length - pack_tokens = pack_tokens[-max_body:] - pack_targets = pack_targets[-max_body:] - pack_tokens.append(pad) - pack_targets.append(pad) - else: # Truncate right (need to add eod) - max_body = pack_length - 1 - pack_tokens = pack_tokens[:max_body] - pack_targets = pack_targets[:max_body] - pack_tokens.extend([eod, pad]) - pack_targets.extend([eod, pad]) + # Truncate on the right + max_body = pack_length + pack_tokens = pack_tokens[:max_body] + pack_targets = pack_targets[:max_body] + pack_tokens.extend(pad) + pack_targets.extend(pad) pack_positions = pack_positions[:pack_length+1] # Note len({pack_tokens, pack_targets, pack_positions}) should be pack_length + 1 cu_seqlens[-1] = len(pack_tokens) - 1