MCPcopy Create free account
hub / github.com/K-Quant/HiDy / create_data_loader

Function create_data_loader

tax_policy/utils.py:30–50  ·  view source on GitHub ↗

Create dataloader. Args: dataset(obj:`paddle.io.Dataset`): Dataset instance. mode(obj:`str`, optional, defaults to obj:`train`): If mode is 'train', it will shuffle the dataset randomly. batch_size(obj:`int`, optional, defaults to 1): The sample number of a mini-batc

(dataset, mode="train", batch_size=1, trans_fn=None)

Source from the content-addressed store, hash-verified

28
29
30def create_data_loader(dataset, mode="train", batch_size=1, trans_fn=None):
31 """
32 Create dataloader.
33 Args:
34 dataset(obj:`paddle.io.Dataset`): Dataset instance.
35 mode(obj:`str`, optional, defaults to obj:`train`): If mode is 'train', it will shuffle the dataset randomly.
36 batch_size(obj:`int`, optional, defaults to 1): The sample number of a mini-batch.
37 trans_fn(obj:`callable`, optional, defaults to `None`): function to convert a data sample to input ids, etc.
38 Returns:
39 dataloader(obj:`paddle.io.DataLoader`): The dataloader which generates batches.
40 """
41 if trans_fn:
42 dataset = dataset.map(trans_fn)
43
44 shuffle = True if mode == "train" else False
45 if mode == "train":
46 sampler = paddle.io.DistributedBatchSampler(dataset=dataset, batch_size=batch_size, shuffle=shuffle)
47 else:
48 sampler = paddle.io.BatchSampler(dataset=dataset, batch_size=batch_size, shuffle=shuffle)
49 dataloader = paddle.io.DataLoader(dataset, batch_sampler=sampler, return_list=True)
50 return dataloader
51
52
53def map_offset(ori_offset, offset_mapping):

Callers 1

do_evalFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected