MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / BlendableDataset

Class BlendableDataset

codegeex/megatron/data/blendable_dataset.py:25–69  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

23
24
25class BlendableDataset(torch.utils.data.Dataset):
26 def __init__(self, datasets, weights):
27
28 self.datasets = datasets
29 num_datasets = len(datasets)
30 assert num_datasets == len(weights)
31
32 self.size = 0
33 for dataset in self.datasets:
34 self.size += len(dataset)
35
36 # Normalize weights.
37 weights = np.array(weights, dtype=np.float64)
38 sum_weights = np.sum(weights)
39 assert sum_weights > 0.0
40 weights /= sum_weights
41
42 # Build indecies.
43 start_time = time.time()
44 assert num_datasets < 255
45 self.dataset_index = np.zeros(self.size, dtype=np.uint8)
46 self.dataset_sample_index = np.zeros(self.size, dtype=np.int64)
47
48 from megatron.data import helpers
49
50 helpers.build_blending_indices(
51 self.dataset_index,
52 self.dataset_sample_index,
53 weights,
54 num_datasets,
55 self.size,
56 torch.distributed.get_rank() == 0,
57 )
58 print_rank_0(
59 "> elapsed time for building blendable dataset indices: "
60 "{:.2f} (sec)".format(time.time() - start_time)
61 )
62
63 def __len__(self):
64 return self.size
65
66 def __getitem__(self, idx):
67 dataset_idx = self.dataset_index[idx]
68 sample_idx = self.dataset_sample_index[idx]
69 return self.datasets[dataset_idx][sample_idx]

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected