MCPcopy Create free account
hub / github.com/LargeWorldModel/LWM / __iter__

Method __iter__

lwm/data.py:272–303  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

270 )
271
272 def __iter__(self):
273 chunk_size = self.config.batch_size * self.config.seq_length
274 total_tokens = 0
275 while True:
276 token_buffer = []
277 loss_mask_buffer = []
278 for index, example in enumerate(self._dataset):
279 tokens, loss_masks = self.text_processor(example)
280 token_buffer.extend(tokens)
281 loss_mask_buffer.extend(loss_masks)
282 while len(token_buffer) > chunk_size + 1:
283 total_tokens += chunk_size
284 metrics = {
285 'dataset_example_index': index,
286 'dataset_total_tokens': total_tokens,
287 }
288 batch = {
289 'input_tokens': np.array(token_buffer[:chunk_size], dtype=np.int32).reshape(
290 self.config.batch_size, -1
291 ),
292 'target_tokens': np.array(token_buffer[1:chunk_size + 1], dtype=np.int32).reshape(
293 self.config.batch_size, -1
294 ),
295 'loss_masks': np.array(loss_mask_buffer[1:chunk_size + 1], dtype=np.float32).reshape(
296 self.config.batch_size, -1
297 ),
298 }
299 if self.config.always_start_with_bos:
300 batch['input_tokens'][:, 0] = self.tokenizer.bos_token_id
301 yield batch, metrics
302 token_buffer = token_buffer[chunk_size:]
303 loss_mask_buffer = loss_mask_buffer[chunk_size:]
304
305 def get_state_dict(self):
306 return dict(config=self.config)

Callers

nothing calls this directly

Calls 1

text_processorMethod · 0.95

Tested by

no test coverage detected