Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 7 additions & 8 deletions src/art/utils/sft.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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)
Expand Down
53 changes: 53 additions & 0 deletions tests/unit/test_sft.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@

from art import TrainableModel
from art.utils.sft import (
SFTChunk,
create_lr_schedule,
create_sft_dataset_iterator,
iterate_file,
Expand Down Expand Up @@ -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)
Expand Down