Skip to content

base_sampler

data.batch_sampler.base_sampler

Batch sampler 클래스의 기본이 되는 추상 클래스를 정의한 모듈입니다.

BatchSampler 클래스를 상속받아 새로운 batch sampler를 정의할 수 있습니다. Naive sampler는 가장 기본적인 sampler 구현 예시로, 데이터셋의 인덱스를 순서대로 배치로 묶어 반환합니다.

logger module-attribute

logger = getLogger('SaigeResearch')

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([])

NaiveSampler

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

Bases: BatchSampler

Naive sampler implementation with buffer extension feature added.

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

generate_buffer

generate_buffer() -> deque
Source code in SaigeToolkit/data/batch_sampler/base_sampler.py
def generate_buffer(self) -> deque:
    return deque(np.arange(len(self.dataset)))