MCPcopy Create free account
hub / github.com/amazon-science/mm-cot / load_data_img

Function load_data_img

utils_data.py:34–66  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

32 return problems, qids,
33
34def 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
68class ScienceQADatasetStd(Dataset):
69 """

Callers 1

main.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected