Skip to content

builder

data.sampler.builder

logger module-attribute

logger = getLogger('SaigeResearch')

SAMPLER_IMPLEMENTATION module-attribute

SAMPLER_IMPLEMENTATION = {'InfiniteRandomSampler': InfiniteRandomSampler}

InfiniteRandomSampler

Bases: RandomSampler

Infinitely samples elements randomly. If without replacement, then sample from a shuffled dataset. If with replacement, then user can specify :attr:num_samples to draw.

Parameters:

  • data_source (Dataset) –

    dataset to sample from

  • replacement (bool) –

    samples are drawn on-demand with replacement if True, default=False

  • num_samples (int) –

    number of samples to draw, default=len(dataset).

  • generator (Generator) –

    Generator used in sampling.

Assume that this DataLoader is used for training with shuffle=True.

We use InfiniteRandomSampler to make the DataLoader to infinitely sample the dataset. This prevents the "drop_last" and the "prefeching" problem across epochs. Be careful that you should stop the training loop by counting the number of steps manually. The DataLoader will never stop by itself.

NOTE: Note that this sampler has infinite length and thus you should be careful when calculating the current epoch. We recommend you to use step//steps_per_epoch to calculate the current epoch.

build_sampler

build_sampler(_target_: Optional[str], data_source: Optional[Sized], **cfg_sampler: dict) -> Optional[Sampler]

Parameters:

  • cfg_sampler (dict, default: {} ) –

    config dict for building sampler

Returns:

  • Optional[Sampler]

    Optional[Sampler]: sampler object instance. if cfg_sampler is None, return None

Notes

Follows the interface of torch.utils.data.Sampler.

Source code in SaigeToolkit/data/sampler/builder.py
def build_sampler(
    _target_: Optional[str], data_source: Optional[Sized], **cfg_sampler: dict
) -> Optional[Sampler]:
    """
    Args:
        cfg_sampler (dict): config dict for building sampler

    Returns:
        Optional[Sampler]: sampler object instance.
            if cfg_sampler is None, return None

    Notes:
        Follows the interface of torch.utils.data.Sampler.
    """
    if _target_ is None:
        return None
    sampler_class = get_sampler_class(_target_)
    logger.info(f"[{'SAMPLER'.center(9)}] {_target_} [params] {cfg_sampler}")
    return sampler_class(data_source=data_source, **cfg_sampler)

get_sampler_class

get_sampler_class(cfg_sampler_name: str) -> Type[Sampler]
Source code in SaigeToolkit/data/sampler/builder.py
def get_sampler_class(cfg_sampler_name: str) -> Type[Sampler]:
    try:
        return SAMPLER_IMPLEMENTATION[cfg_sampler_name]
    except KeyError as e:
        raise RuntimeError(
            f"Sampler {cfg_sampler_name} not implemented. Supported sampler list: {SAMPLER_IMPLEMENTATION.keys()}"
        ) from e
    except Exception as e:
        logger.error(f"Error: {e}")
        raise e