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

Function get_batch

SwissArmyTransformer/examples/deit/finetune_vit_cifar10.py:17–49  ·  view source on GitHub ↗
(data_iterator, args, timers)

Source from the content-addressed store, hash-verified

15from sat.training.deepspeed_training import training_main
16
17def get_batch(data_iterator, args, timers):
18 # Items and their type.
19
20 datatype = torch.int64
21
22 # Broadcast data.
23 timers('data loader').start()
24 if data_iterator is not None:
25 data = next(data_iterator)
26 else:
27 data = None
28 image_data = {"image":data[0]}
29 label_data = {"label":data[1]}
30 timers('data loader').stop()
31 image_data = mpu.broadcast_data(["image"], image_data, torch.float32)
32 label_data = mpu.broadcast_data(["label"], label_data, torch.int64)
33
34 # Unpack.
35 label_data = label_data['label'].long()
36 image_data = image_data['image']
37 batch_size = label_data.size()[0]
38 seq_length = args.pre_len + (args.image_size[0]//args.patch_size)*(args.image_size[1]//args.patch_size) + args.post_len
39 position_ids = torch.zeros(seq_length, device=image_data.device, dtype=torch.long)
40 torch.arange(0, seq_length, out=position_ids[:seq_length])
41 position_ids = position_ids.unsqueeze(0).expand([batch_size, -1])
42 attention_mask = torch.ones((1, 1), device=image_data.device)
43
44 tokens = torch.zeros((batch_size, 1), device=image_data.device, dtype=torch.long)
45 # Convert
46 if args.fp16:
47 attention_mask = attention_mask.half()
48 image_data = image_data.half()
49 return tokens, image_data, label_data, attention_mask, position_ids
50
51
52def forward_step(data_iterator, model, args, timers):

Callers 1

forward_stepFunction · 0.70

Calls 2

startMethod · 0.80
stopMethod · 0.80

Tested by

no test coverage detected