Skip to content

perf(data): avoid repeated full-bucket scans in the sampler - #3243

Merged
bghira merged 1 commit into
bghira:mainfrom
hjinnkim:fix/sampler-unseen-index
Oct 1, 2026
Merged

bghira merged 1 commit into
bghira:mainfrom
hjinnkim:fix/sampler-unseen-index

Conversation

@hjinnkim

Copy link
Copy Markdown
Contributor

Summary

Follow-up to #2897. The sampler rebuilds a bucket's unseen list for every batch, making sampling expensive for large buckets. This change maintains a per-bucket index and removes consumed occurrences while preserving duplicate handling and batch order.

Indices are rebuilt when bucket contents change, including same-length VAE-filter refills, and invalidated on state restore or epoch reset. Each lookup returns a copy to preserve iterator behavior. The index uses additional memory proportional to bucket contents; list comparison, copying, and deletion remain linear operations.

Verification

  • tests.test_sampler: 31/31 pass, covering duplicates, partial batches, bucket changes, and resume.
  • CPU replay on upstream f1cb3800c: the same checkpoint state, seed, and batch size produced identical batch and within-batch image order across 200 batches each for ranks 0 and 3.

Mean sampler time per batch, including initial index construction:

Rank Unpatched Patched
0 61.32 ms 6.59 ms
3 66.89 ms 6.55 ms

These measurements cover sampler execution with image validation and conditioning mocked, while production was running on the same host.

…ned index

MultiAspectSampler.__iter__ rebuilt a bucket's unseen list for every batch by
filtering the whole bucket through the occurrence counts. The cost grew with the
bucket (about 0.66 s per batch on a 257k-image rank shard) and, because loading
is synchronous, every other rank waited for it at the first collective. Late in
an epoch only the largest buckets remain, so it hit almost every step.

The iterator now keeps one index per bucket with the unseen occurrences in bucket
order. Marking a filepath seen raises its occurrence count by one, which retires
exactly its earliest unseen occurrence, so each consumed sample is removed from
the index instead of rescanning. The list handed to random.sample is the same as
before, so batches are unchanged; an index is rebuilt when its bucket's contents
differ from the snapshot it was built from (the VAE cache filter pads a split
shard back to its old length in place), when the seen dict is replaced, and after
resume or an epoch reset. _get_unseen_images keeps recomputing for its other
callers.

Claude-Session: https://claude.ai/code/session_01Ai6TFdNTtZTBiGjCsZgyF3
@bghira
bghira merged commit aedbdaf into bghira:main Oct 1, 2026
2 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants