| 32 | return problems, qids, |
| 33 | |
| 34 | def load_data_img(args): |
| 35 | problems = json.load(open(os.path.join(args.data_root, 'scienceqa/problems.json'))) |
| 36 | pid_splits = json.load(open(os.path.join(args.data_root, 'scienceqa/pid_splits.json'))) |
| 37 | captions = json.load(open(args.caption_file))["captions"] |
| 38 | name_maps = json.load(open('data/name_map.json')) |
| 39 | |
| 40 | # check |
| 41 | if args.img_type == "resnet": |
| 42 | image_features = np.load('vision_features/resnet.npy') |
| 43 | image_features = np.expand_dims(image_features, axis=1) |
| 44 | image_features = image_features.repeat(512, axis=1) |
| 45 | elif args.img_type == "clip": |
| 46 | image_features = np.load('vision_features/clip.npy') |
| 47 | elif args.img_type == "detr": |
| 48 | image_features = np.load('vision_features/detr.npy') |
| 49 | elif args.img_type == "vit": |
| 50 | image_features = torch.load("vision_features/vit.pth") |
| 51 | else: |
| 52 | image_features = np.load('vision_features/detr.npy') |
| 53 | print("img_features size: ", image_features.shape) |
| 54 | |
| 55 | for qid in problems: |
| 56 | problems[qid]['caption'] = captions[qid] if qid in captions else "" |
| 57 | |
| 58 | train_qids = pid_splits['%s' % (args.train_split)] |
| 59 | val_qids = pid_splits['%s' % (args.val_split)] |
| 60 | test_qids = pid_splits['%s' % (args.test_split)] |
| 61 | print(f"number of train problems: {len(train_qids)}\n") |
| 62 | print(f"number of val problems: {len(val_qids)}\n") |
| 63 | print(f"number of test problems: {len(test_qids)}\n") |
| 64 | |
| 65 | qids = {'train': train_qids, 'val':val_qids,'test':test_qids} |
| 66 | return problems, qids, name_maps, image_features |
| 67 | |
| 68 | class ScienceQADatasetStd(Dataset): |
| 69 | """ |