MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / get_batch

Function get_batch

SwissArmyTransformer/examples/chatglm/finetune_chatglm.py:50–67  ·  view source on GitHub ↗
(data_iterator, args, timers)

Source from the content-addressed store, hash-verified

48
49from transformers import DataCollatorForSeq2Seq
50def get_batch(data_iterator, args, timers):
51 # Items and their type.
52 keys = ['input_ids', 'labels']
53 datatype = torch.int64
54
55 # Broadcast data.
56 timers('data loader').start()
57 if data_iterator is not None:
58 data = next(data_iterator)
59 else:
60 data = None
61 timers('data loader').stop()
62 data_b = mpu.broadcast_data(keys, data, datatype)
63 # Unpack.
64 tokens = data_b['input_ids'].long()
65 labels = data_b['labels'].long()
66
67 return tokens, labels
68
69
70from torch.nn import CrossEntropyLoss

Callers 2

forward_step_evalFunction · 0.70
forward_stepFunction · 0.70

Calls 2

startMethod · 0.80
stopMethod · 0.80

Tested by

no test coverage detected