diff --git a/src/art/utils/sft.py b/src/art/utils/sft.py index 776330c9a..b9d571130 100644 --- a/src/art/utils/sft.py +++ b/src/art/utils/sft.py @@ -217,9 +217,6 @@ def create_sft_dataset_iterator( items_per_chunk = batch_size * chunk_size chunks_per_epoch = math.ceil(dataset_size / items_per_chunk) - # Convert initial_step (batch-based) to initial_chunk for skipping - initial_chunk = initial_step // chunk_size - pbar = ( tqdm( initial=initial_step, total=total_batches, desc="SFT Training", unit="step" @@ -234,14 +231,16 @@ def create_sft_dataset_iterator( random.Random(seed + epoch).shuffle(epoch_trajs) for chunk_idx in range(chunks_per_epoch): - global_chunk_idx = epoch * chunks_per_epoch + chunk_idx + chunk_start = chunk_idx * items_per_chunk + chunk_end = min(chunk_start + items_per_chunk, dataset_size) - # Skip chunks before initial_step - if global_chunk_idx < initial_chunk: + # Resume by batch, since the last chunk of an epoch can be shorter. + chunk_start = max( + chunk_start, (initial_step - epoch * batches_per_epoch) * batch_size + ) + if chunk_start >= chunk_end: continue - chunk_start = chunk_idx * items_per_chunk - chunk_end = min(chunk_start + items_per_chunk, dataset_size) chunk_trajs = epoch_trajs[chunk_start:chunk_end] num_batches_in_chunk = math.ceil(len(chunk_trajs) / batch_size) diff --git a/tests/unit/test_sft.py b/tests/unit/test_sft.py index 687c0349c..abcab3e6f 100644 --- a/tests/unit/test_sft.py +++ b/tests/unit/test_sft.py @@ -11,6 +11,7 @@ from art import TrainableModel from art.utils.sft import ( + SFTChunk, create_lr_schedule, create_sft_dataset_iterator, iterate_file, @@ -307,6 +308,58 @@ def test_create_sft_dataset_iterator_initial_step(): assert resumed_chunks[0].config.learning_rate == all_chunks[1].config.learning_rate +@pytest.mark.parametrize("dataset_size", [5, 6]) +@pytest.mark.parametrize("shuffle", [False, True]) +@pytest.mark.parametrize("initial_step", range(8)) +def test_create_sft_dataset_iterator_resume_batches( + dataset_size, shuffle, initial_step +): + # Each epoch has three batches, so its second chunk has only one batch. + trajs = _make_trajectories(dataset_size) + all_chunks = list( + create_sft_dataset_iterator( + trajs, + chunk_size=2, + epochs=2, + batch_size=2, + shuffle=shuffle, + show_progress=False, + ) + ) + resumed_chunks = list( + create_sft_dataset_iterator( + trajs, + chunk_size=2, + epochs=2, + batch_size=2, + shuffle=shuffle, + initial_step=initial_step, + show_progress=False, + ) + ) + + def batches(chunks: list[SFTChunk]): + result = [] + for chunk in chunks: + lrs = chunk.config.learning_rate + assert isinstance(lrs, list) + for offset, lr in zip( + range(0, len(chunk.trajectories), 2), lrs, strict=True + ): + result.append( + ( + chunk.trajectories[offset : offset + 2], + lr, + chunk.step + offset // 2, + chunk.epoch, + chunk.epoch_step + offset // 2, + ) + ) + return result + + assert batches(resumed_chunks) == batches(all_chunks)[initial_step:] + + def test_create_sft_dataset_iterator_deterministic(): """Test that create_sft_dataset_iterator is deterministic with the same seed.""" trajs = _make_trajectories(50)