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

Method __init__

codegeex/megatron/data/prompt_dataset.py:180–227  ·  view source on GitHub ↗

Args: name: name of the dataset. data_prefix: prefix of the data. documents: list of document indices. input_ids_indexed_dataset: indexed dataset for prompts. attention_mask_index_dataset: indexed dataset for text. labe

(
        self,
        name,
        data_prefix,
        documents,
        input_ids_indexed_dataset,
        attention_mask_index_dataset,
        labels_indexed_dataset,
        num_samples,
        seq_length,
        seed,
    )

Source from the content-addressed store, hash-verified

178
179class PromptDataset(torch.utils.data.Dataset):
180 def __init__(
181 self,
182 name,
183 data_prefix,
184 documents,
185 input_ids_indexed_dataset,
186 attention_mask_index_dataset,
187 labels_indexed_dataset,
188 num_samples,
189 seq_length,
190 seed,
191 ):
192 """
193 Args:
194 name: name of the dataset.
195 data_prefix: prefix of the data.
196 documents: list of document indices.
197 input_ids_indexed_dataset: indexed dataset for prompts.
198 attention_mask_index_dataset: indexed dataset for text.
199 labels_indexed_dataset: indexed dataset for labels.
200 num_samples: number of samples to draw from the indexed dataset.
201 seq_length: sequence length.
202 seed: seed for random number generator.
203 """
204
205 self.name = name
206 self.input_ids_indexed_dataset = input_ids_indexed_dataset
207 self.attention_mask_index_dataset = attention_mask_index_dataset
208 self.labels_indexed_dataset = labels_indexed_dataset
209 self.seq_length = seq_length
210 self.eod_token = get_tokenizer().eod
211
212 # Checks
213 assert np.min(documents) >= 0
214 assert np.max(documents) < input_ids_indexed_dataset.sizes.shape[0]
215 assert input_ids_indexed_dataset.sizes.shape[0] == attention_mask_index_dataset.sizes.shape[0]
216 assert attention_mask_index_dataset.sizes.shape[0] == labels_indexed_dataset.sizes.shape[0]
217
218 # Build index mappings.
219 self.doc_idx = _build_index_mappings(
220 self.name,
221 data_prefix,
222 documents,
223 self.input_ids_indexed_dataset.sizes,
224 num_samples,
225 seq_length,
226 seed,
227 )
228
229 def __len__(self):
230 # -1 is due to data structure used to retieve the index:

Callers

nothing calls this directly

Calls 2

get_tokenizerFunction · 0.90
_build_index_mappingsFunction · 0.85

Tested by

no test coverage detected