{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":6799,"databundleVersionId":4225553,"sourceType":"competition"},{"sourceId":15444,"sourceType":"datasetVersion","datasetId":11102},{"sourceId":1462296,"sourceType":"datasetVersion","datasetId":857191}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RUN ALL CELLS, the last cell gives a link to the data.pkl ","metadata":{}},{"cell_type":"code","source":"# store everything in dict\ndata = {}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-17T09:28:48.861552Z","iopub.execute_input":"2025-02-17T09:28:48.862061Z","iopub.status.idle":"2025-02-17T09:28:48.911415Z","shell.execute_reply.started":"2025-02-17T09:28:48.862017Z","shell.execute_reply":"2025-02-17T09:28:48.909942Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# PARTI PROMPTS","metadata":{}},{"cell_type":"code","source":"from datasets import load_dataset\n\nparti_ds = load_dataset(\"nateraw/parti-prompts\")\ndata['PARTI'] = {'imgs': None, 'anns': parti_ds['train']['Prompt']}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-17T09:29:51.761305Z","iopub.execute_input":"2025-02-17T09:29:51.761864Z","iopub.status.idle":"2025-02-17T09:29:52.860184Z","shell.execute_reply.started":"2025-02-17T09:29:51.761818Z","shell.execute_reply":"2025-02-17T09:29:52.858875Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# MSCOCO","metadata":{}},{"cell_type":"code","source":"N_IMAGES = 10000\nside_min = 299 # filter min(width, height) > side_min\nside_delta = 130 # filter max |width-height|<side_delta","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-17T09:28:56.653232Z","iopub.execute_input":"2025-02-17T09:28:56.653922Z","iopub.status.idle":"2025-02-17T09:28:56.660717Z","shell.execute_reply.started":"2025-02-17T09:28:56.653878Z","shell.execute_reply":"2025-02-17T09:28:56.658778Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -U \"git+https://github.com/philferriere/cocoapi.git#egg=pycocotools&subdirectory=PythonAPI\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-17T09:28:56.663062Z","iopub.execute_input":"2025-02-17T09:28:56.663478Z","iopub.status.idle":"2025-02-17T09:29:26.83686Z","shell.execute_reply.started":"2025-02-17T09:28:56.663437Z","shell.execute_reply":"2025-02-17T09:29:26.834791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from PIL import Image\nfrom tqdm import tqdm\nimport torchvision.transforms as T\nimport torch\nimport random\nfrom pycocotools.coco import COCO\ncoco_api = COCO('/kaggle/input/coco-2017-dataset/coco2017/annotations/captions_train2017.json')\n\nimg_ids = coco_api.getImgIds()\nimg_jsons = coco_api.loadImgs(img_ids)\n\nimg_jsons = [i for i in img_jsons if min(i['width'], i['height']) >= side_min and abs(i['width'] - i['height']) < side_delta]\n\nrandom.seed(42)\nprint(\"MAX number of images: \", len(img_jsons))\nimg_jsons = random.sample(img_jsons, N_IMAGES)\nimg_ids = [i['id'] for i in img_jsons]\n\nimgs = torch.zeros((N_IMAGES, 3, 299, 299), dtype=torch.uint8)\n\nfor i, img_js in tqdm(enumerate(img_jsons), total=len(img_jsons), desc=\"load images and resize to 299x299\"):\n    img = Image.open(f'/kaggle/input/coco-2017-dataset/coco2017/train2017/{img_js[\"file_name\"]}').convert('RGB').resize((299, 299), Image.LANCZOS)\n    imgs[i] = ((T.ToTensor()(img)*255).to(torch.uint8))\n\nset_img_ids = set(img_ids)\nall_annotations = coco_api.loadAnns(coco_api.getAnnIds())\nanns = {}\nfor ann in all_annotations:\n    if ann['image_id'] in set_img_ids:\n        anns[ann['image_id']] = anns.get(ann['image_id'], []) + [ann['caption']]\n\nanns = [anns[i] for i in img_ids]\n\n\ndata['COCO'] = {'imgs':imgs, 'anns':anns}","metadata":{"execution":{"iopub.status.busy":"2025-02-17T09:30:12.290122Z","iopub.execute_input":"2025-02-17T09:30:12.292105Z","iopub.status.idle":"2025-02-17T09:35:05.366254Z","shell.execute_reply.started":"2025-02-17T09:30:12.292024Z","shell.execute_reply":"2025-02-17T09:35:05.364557Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\ndef process_strings(strings, max_length):\n    sorted_strings = sorted(strings, key=len, reverse=True)\n    result = ''\n    strings_to_check = sorted_strings.copy()\n    popped_strings = []\n\n  # Step 1: Find the first suitable string\n    while strings_to_check:\n        first_row = strings_to_check[0]\n        if len(first_row) <= max_length:\n            result = first_row\n            break\n        else:\n            # Remove strings that are too long\n            popped_strings.append(strings_to_check.pop(0))\n    else:\n        # If all strings are too long, return truncated smallest string\n        smallest_row = sorted_strings[-1]\n        truncated_row = smallest_row[:max_length]\n        return truncated_row\n\n    # Step 2: Add random rows while under max_length\n    random.seed(42)\n    unused_strings = [s for s in strings_to_check if s != result]\n\n    while unused_strings:\n        random_row = random.choice(unused_strings)\n        potential_result = result + '\\n' + random_row\n        if len(potential_result) <= max_length:\n            result = potential_result\n            unused_strings.remove(random_row)\n        else:\n            break\n    return result\n\ndata['COCO']['anns']  = [process_strings(i, 250) for i in data['COCO']['anns']]\ndata['PARTI']['anns'] = [i[:250] for i in data['PARTI']['anns']]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-17T09:36:40.317165Z","iopub.execute_input":"2025-02-17T09:36:40.317694Z","iopub.status.idle":"2025-02-17T09:36:40.534225Z","shell.execute_reply.started":"2025-02-17T09:36:40.317652Z","shell.execute_reply":"2025-02-17T09:36:40.532649Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open('data.pkl', 'wb') as f:\n    import pickle\n    pickle.dump(data, f)\n\n!du -sh data.pkl\n\nfrom IPython.display import FileLink, display\ndisplay(FileLink('data.pkl'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-17T09:37:46.911632Z","iopub.execute_input":"2025-02-17T09:37:46.912149Z","iopub.status.idle":"2025-02-17T09:37:56.843875Z","shell.execute_reply.started":"2025-02-17T09:37:46.912107Z","shell.execute_reply":"2025-02-17T09:37:56.842187Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CIFAR (UNCOMMENT IF NEEDED)","metadata":{"_kg_hide-input":true}},{"cell_type":"code","source":"# !tar xzvf ../input/cifar10-python/cifar-10-python.tar.gz\n\n# from numpy import load\n# p = 'cifar-10-batches-py/data_batch_1'\n\n# def unpickle(file):\n#   import pickle\n#   with open(file, 'rb') as fo:\n#     d = pickle.load(fo, encoding='bytes')\n#   return d\n\n# cifar = unpickle(p)\n# cifar = {k.decode('utf-8'): v for k, v in cifar.items()}\n# cifar['filenames'] = [v.decode('utf-8') for v in cifar['filenames']]\n\n\n# data['CIFAR'] = {'imgs': torch.tensor(cifar['data'].reshape(-1, 3, 32, 32), dtype=torch.uint8), 'anns':[i.split('_s')[0] for i in cifar['filenames']]}","metadata":{"trusted":true,"_kg_hide-output":true,"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ImageNet  (UNCOMMENT IF NEEDED)","metadata":{"_kg_hide-input":true}},{"cell_type":"code","source":"# !pip install xmltodict\n\n# import os \n# anns = []\n# imageNet_root = '/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Annotations/CLS-LOC/val'\n# for fp in tqdm(os.listdir(imageNet_root)):\n#     with open(os.path.join(imageNet_root, fp), 'r', encoding='utf-8') as file:\n#         anns.append(xmltodict.parse(file.read()))\n\n\n# ## plot distr        \n# # %matplotlib inline\n# # import matplotlib.pyplot as plt\n# # plt.scatter(*zip(*[(int(i['annotation']['size']['width']), int(i['annotation']['size']['height'])) for i in anns if int(i['annotation']['size']['width'])>=299 and int(i['annotation']['size']['height'])>=299]))\n# # plt.show()\n\n# anns = [i for i in anns if min(int(i['annotation']['size']['width']), int(i['annotation']['size']['height'])) <= side_min and abs(int(i['annotation']['size']['width']) - int(i['annotation']['size']['height'])) < side_delta]\n\n# random.seed(42)\n# print(\"MAX number of images: \", len(anns))\n# anns = random.sample(anns, N_IMAGES)\n# img_filenames = [i['annotation']['filename'] for i in anns]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-23T13:36:15.394668Z","iopub.execute_input":"2024-11-23T13:36:15.395134Z","iopub.status.idle":"2024-11-23T13:36:15.617279Z","shell.execute_reply.started":"2024-11-23T13:36:15.395094Z","shell.execute_reply":"2024-11-23T13:36:15.615956Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# img_path = '/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/val'\n\n# imgs = torch.zeros((N_IMAGES, 3, 299, 299), dtype=torch.uint8)\n\n# for i, img_filename in tqdm(enumerate(img_filenames), total=len(img_filenames)):\n#     img = Image.open(os.path.join(img_path, img_filename+'.JPEG')).convert('RGB').resize((299, 299), Image.LANCZOS)\n#     imgs[i] = ((T.ToTensor()(img)*255).to(torch.uint8))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-23T13:39:39.150123Z","iopub.execute_input":"2024-11-23T13:39:39.150633Z","iopub.status.idle":"2024-11-23T13:43:17.466895Z","shell.execute_reply.started":"2024-11-23T13:39:39.150593Z","shell.execute_reply":"2024-11-23T13:43:17.465505Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# with open('/kaggle/input/imagenet-object-localization-challenge/LOC_synset_mapping.txt', 'r') as f:\n#     decode_ann = f.readlines()\n#     decode_ann = {i.split(' ', 1)[0]:i.split(' ', 1)[1][:-1] for i in decode_ann}\n\n# imageNet_anns = []\n# for i in anns:\n#     objects = i['annotation']['object']\n#     if isinstance(objects, list) and len(objects)>1:\n#         max_id = np.argmax([(int(obj['bndbox']['xmax']) - int(obj['bndbox']['xmin'])) * (int(obj['bndbox']['ymax']) - int(obj['bndbox']['ymin'])) for obj in objects])\n#         ann = decode_ann[objects[max_id]['name']]\n#     else:\n#         ann = decode_ann[objects['name']]\n#     imageNet_anns.append(ann)\n\n# data['IMAGENET'] = {'imgs': imgs, 'anns': imageNet_anns}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-23T13:51:09.472527Z","iopub.execute_input":"2024-11-23T13:51:09.473Z","iopub.status.idle":"2024-11-23T13:51:09.481782Z","shell.execute_reply.started":"2024-11-23T13:51:09.472936Z","shell.execute_reply":"2024-11-23T13:51:09.480442Z"},"_kg_hide-input":true},"outputs":[],"execution_count":null}]}