Skip to content

class_sampler

data.batch_sampler.class_sampler

BatchSampler

BatchSampler(dataset: Dataset, batch_size: int = 1, repeat: bool = False, shuffle: bool = False, drop_last: bool = False)

Bases: Registerable, ABC

Baseclass for all batch_sampler classes.

Attributes:

  • dataset (Dataset) –

    dataset to sample from.

  • batch_size (int) –

    batch size.

  • repeat (bool) –

    repeat data when data is not enough for batch size.

  • shuffle (bool) –

    shuffle data.

  • drop_last (bool) –

    drop last data when data is not enough for batch size.

Note
  • batch_size, shuffle, and drop_last should be popped from dataloader config, then passed to the batch sampler config.
  • repeat: for SaigeVision2 api, should be set as True.
    • if number of data is enough for batch size, then repeat is neglected.
  • drop_last: for SaigeVision2 api, should be set as True.
    • to prevent the last batch from being smaller than the specified batch size.
    • if False, even when repeat is set as True, the last batch will be smaller than the specified batch size.
  • for Training, shuffle should be set as True.
    • drop_last and repeat are optional.
  • for Validation, repeat, shuffle, and drop_last should be set as False.

Examples:

cfg_sampler.update(
    {
        "shuffle": cfg_dataloader.pop("shuffle", False),
        "drop_last": cfg_dataloader.pop("drop_last", False),
        "batch_size": cfg_dataloader.pop("batch_size", 1),
    }
)
batch_sampler = build_batch_sampler(dataset=dataset, **cfg_sampler)
dataloader = build_dataloader(
    dataset,
    collate_fn=collate_fn,
    batch_sampler=batch_sampler,
    **cfg_dataloader,
)
Source code in SaigeToolkit/data/batch_sampler/base_sampler.py
def __init__(
    self,
    dataset: data.Dataset,
    batch_size: int = 1,
    repeat: bool = False,
    shuffle: bool = False,
    drop_last: bool = False,
) -> None:
    self.dataset = dataset

    self.batch_size = batch_size
    self.repeat = repeat

    self.shuffle = shuffle
    self.drop_last = drop_last

    self._n_repeat = None
    self._buffer_length = None

registry instance-attribute

registry: Dict[str, BatchSampler]

dataset instance-attribute

dataset = dataset

batch_size instance-attribute

batch_size = batch_size

repeat instance-attribute

repeat = repeat

shuffle instance-attribute

shuffle = shuffle

drop_last instance-attribute

drop_last = drop_last

_n_repeat instance-attribute

_n_repeat = None

_buffer_length instance-attribute

_buffer_length = None

buffer_length property

buffer_length: int

n_repeat property

n_repeat: int

__iter__

__iter__() -> Iterator[List[int]]
Source code in SaigeToolkit/data/batch_sampler/base_sampler.py
def __iter__(self) -> Iterator[List[int]]:
    if self.batch_size > self.buffer_length:
        if self.repeat:
            logger.info("data repeat activated.")
        elif self.drop_last:
            self.drop_last = False
            logger.warning("data too short. drop_last deactivated.")

    while True:
        yield from self._get_processed_buffer()

_get_processed_buffer

_get_processed_buffer() -> deque
Source code in SaigeToolkit/data/batch_sampler/base_sampler.py
def _get_processed_buffer(self) -> deque:
    buffer = self._shuffle_and_repeat_buffer()
    buffer = self._batch_buffer(buffer)
    return buffer

_shuffle_and_repeat_buffer

_shuffle_and_repeat_buffer() -> deque
Source code in SaigeToolkit/data/batch_sampler/base_sampler.py
def _shuffle_and_repeat_buffer(self) -> deque:
    if self.repeat and self.batch_size > self.buffer_length:
        buffer = deque()
        for _ in range(self.n_repeat):
            _buffer = self.generate_buffer()
            if self.shuffle:
                np.random.shuffle(_buffer)
            buffer.extend(_buffer)

    else:
        buffer = self.generate_buffer()
        if self.shuffle:
            np.random.shuffle(buffer)

    return buffer

_batch_buffer

_batch_buffer(buffer: deque) -> deque
Source code in SaigeToolkit/data/batch_sampler/base_sampler.py
def _batch_buffer(self, buffer: deque) -> deque:
    data_length, remainders = divmod(len(buffer), self.batch_size)
    _buffer = deque([[buffer.popleft() for _ in range(self.batch_size)] for _ in range(data_length)])
    if not self.drop_last and remainders > 0:
        _buffer.append([buffer.popleft() for _ in range(remainders)])
    return _buffer

__len__

__len__() -> int
Source code in SaigeToolkit/data/batch_sampler/base_sampler.py
def __len__(self) -> int:
    if self.drop_last:
        length = self.buffer_length // self.batch_size
    else:
        length = ceil(self.buffer_length / self.batch_size)
    return max(1, length)

generate_buffer abstractmethod

generate_buffer() -> deque
Source code in SaigeToolkit/data/batch_sampler/base_sampler.py
@abc.abstractmethod
def generate_buffer(self) -> deque:
    return deque([])

ClassBalancedSampler

ClassBalancedSampler(dataset: Dataset, shuffle: bool = True, **kwargs)

Bases: BatchSampler

Sampler implementation for simple class balancing.

Uses class_index value in dataset file dictionaries. When more than one classes exist in single image, the multiplier of the rarest class within image will be applied to the data.

Raises:

  • IndexError

    Will be raised when class index exceeds MAX_NUM_CLS.

Notes
  • TODO: Weighted balancing.
  • TODO: Cropper applicability.
Source code in SaigeToolkit/data/batch_sampler/class_sampler.py
def __init__(
    self,
    dataset: Dataset,
    shuffle: bool = True,
    **kwargs,
) -> None:
    super().__init__(dataset, shuffle=shuffle, **kwargs)
    self.label_list, self.class_multiplier = self._make_weights_for_balanced_classes()

MAX_NUM_CLS class-attribute instance-attribute

MAX_NUM_CLS = 128

_make_weights_for_balanced_classes

_make_weights_for_balanced_classes() -> Tuple[List[int], List[int]]
Source code in SaigeToolkit/data/batch_sampler/class_sampler.py
def _make_weights_for_balanced_classes(self) -> Tuple[List[int], List[int]]:
    label_count = [0] * self.MAX_NUM_CLS
    label_list = [[] for _ in range(len(self.dataset.files))]
    for idx, file in enumerate(self.dataset.files):
        try:
            for label in file["labels"]:
                class_idx = label["class_index"]
                label_count[class_idx] += 1
                label_list[idx].append(class_idx)
        except KeyError:
            # XXX: For srproj compatibility.
            for label in file["LabelGroup"]["Label"]:
                class_idx = label["ClassIndex"]
                label_count[class_idx] += 1
                label_list[idx].append(class_idx)

    max_label_count = max(label_count)
    label_count = [
        max_label_count if cnt == 0 else cnt for cnt in label_count
    ]  # Multiplier 1 for cls not seen.
    data_multiplier_per_class = [round(max_label_count / i) for i in label_count]

    return label_list, data_multiplier_per_class

generate_buffer

generate_buffer() -> deque
Source code in SaigeToolkit/data/batch_sampler/class_sampler.py
def generate_buffer(self) -> deque:
    buffer_balanced = []
    for i in range(len(self.dataset.files)):
        labels = self.label_list[i]
        multipliers = [self.class_multiplier[i] for i in labels]
        multiplier = max(multipliers, default=1)
        buffer_balanced.extend([i] * multiplier)

    return deque(buffer_balanced)