(data_iterator, args, timers)
| 15 | from sat.training.deepspeed_training import training_main |
| 16 | |
| 17 | def 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 | |
| 52 | def forward_step(data_iterator, model, args, timers): |
no test coverage detected