MCPcopy Create free account
hub / github.com/OpenLMLab/MOSS-RLHF / batchify

Method batchify

ppo/ppo_datahelper.py:337–351  ·  view source on GitHub ↗
(self, batch_samples: List[Dict[str, Any]])

Source from the content-addressed store, hash-verified

335 return output
336
337 def batchify(self, batch_samples: List[Dict[str, Any]]) -> Dict[str, Any]:
338 batch = dict()
339 batch_text_vec = torch.tensor(pad_sequences(
340 [sample['text_vec'] for sample in batch_samples], pad_value=self.tokenizer.pad_token_id, pad_left=False
341 ), dtype=torch.long)
342 loss_mask = torch.tensor(pad_sequences(
343 [sample['loss_mask'] for sample in batch_samples], pad_value=0, pad_left=False
344 ), dtype=torch.bool)
345
346 batch.update({
347 'text_vec': batch_text_vec,
348 'loss_mask': loss_mask
349 })
350
351 return batch
352
353 def batch_generator(self):
354 while True:

Callers 1

final_generatorMethod · 0.45

Calls 2

pad_sequencesFunction · 0.90
updateMethod · 0.80

Tested by

no test coverage detected