pytorch / pytorch/vision

Per element sampling for Mixup/Cutmix

Open
#8,191 1 comment 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
17.9k
Forks
7.3k
Avg merge
1d 15h
Merged PRs (30d)
13

Description

🚀 The feature

Sample different random parameters of Mixup/Cutmix for different elements of the batch to avoid loss instability in large batch setups.

Motivation, pitch

Hello!

My understanding is that only one random sampling of random parameters is done for the entire batch in Mixup/Cutmix, which leads all elements of the batch to have the same augmentation. For example, when using CutMix, all images in the batch end up with the same bounding box at the exact same location.

This is particularly hurting when using very large batch as it leads to unstable training. In one batch, you get lucky, you get easy parameters and the loss is already low, then in the next batch, you get unlucky and the parameters are super hard and the loss is super high. If transform parameters were sampled per element, you would get an averaging effect that mitigates this issue in the case of large batch sizes.

This could be as easy as replacing the various instances of self._dist.sample(()) by self._dist.sample((batch_size,)) and a bit more involved for the bounding boxes, but nothing really outstandingly hard.

Best!
David.

Alternatives

No response

Additional context

No response

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Locate the Mixup/Cutmix implementation and inspect the self._dist.sample(()) calls first. Determine how batch dimensions and CutMix bounding boxes are handled, then make sampling independent per element; done means different elements can receive different parameters and bounding boxes, reducing large-batch loss instability.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
computer-vision, machine-learning
Issue type
Feature
Difficulty
3/5
Estimated time
1-2 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.