{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install faiss-cpu","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile cluster_images.py\nimport os\nimport json\nimport random\nimport numpy as np\nimport faiss\nimport glob\nimport argparse\nimport xml.etree.ElementTree as ET\nfrom PIL import Image\nfrom tqdm import tqdm\nimport torchvision.transforms as transforms\n\n\ndef parse_val_name(filename: str) -> str:\n    \"\"\" Extracts \"00000001\" from \"ILSVRC2012_val_00000001.jpg\". \"\"\"\n    base = os.path.splitext(os.path.basename(filename))[0]\n    return base.replace(\"ILSVRC2012_val_\", \"\")\n\n\ndef parse_train_name(filepath: str) -> str:\n    \"\"\" Extracts \"n01440764_10040\" from \"n01440764/n01440764_10040.jpg\". \"\"\"\n    base = os.path.splitext(filepath)[0]\n    return base.split(\"/\")[-1]\n\n\ndef load_and_flatten(filepath: str) -> np.ndarray:\n    \"\"\" Loads an image, resizes to 256x256, center-crops it to 224x224, and normalizes. \"\"\"\n    imagenet_transform = transforms.Compose([\n        transforms.Resize(256, interpolation=Image.BILINEAR),\n        transforms.CenterCrop(224),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n    ])\n    img = Image.open(filepath).convert(\"RGB\")\n    img = imagenet_transform(img)\n    return img.numpy().flatten()\n\n\ndef extract_class_from_xml(xml_path: str) -> str:\n    \"\"\" Extracts the class ID from an ImageNet validation XML annotation file. \"\"\"\n    try:\n        tree = ET.parse(xml_path)\n        root = tree.getroot()\n        obj = root.find(\"object\")\n        if obj is not None:\n            return obj.find(\"name\").text  # e.g., \"n01751748\"\n    except Exception as e:\n        print(f\"Warning: Could not parse {xml_path} - {e}\")\n    return \"unknown\"\n\n\ndef main():\n    parser = argparse.ArgumentParser()\n    parser.add_argument(\"--train_dir\", type=str, required=True,\n                        help=\"Path to the folder containing 1k subfolders of train images.\")\n    parser.add_argument(\"--val_dir\", type=str, required=True,\n                        help=\"Path to the folder containing 50k val images named ILSVRC2012_val_XXXXXX.jpg.\")\n    parser.add_argument(\"--val_xml_dir\", type=str, required=True,\n                        help=\"Path to the folder containing XML annotations for val images.\")\n    parser.add_argument(\"--output_dir\", type=str, default=\".\",\n                        help=\"Where to save train_grouping_X.json and val_grouping.json.\")\n    parser.add_argument(\"--split_index\", type=int, choices=range(5), required=True,\n                        help=\"Which split of the training data to process (0 to 4).\")\n    parser.add_argument(\"--seed\", type=int, default=42,\n                        help=\"Random seed for picking 10k centroid images.\")\n    args = parser.parse_args()\n\n    random.seed(args.seed)\n    np.random.seed(args.seed)\n\n    train_dir = args.train_dir\n    val_dir = args.val_dir\n    val_xml_dir = args.val_xml_dir\n    output_dir = args.output_dir\n    split_index = args.split_index\n    os.makedirs(output_dir, exist_ok=True)\n\n    # ---------------------------------------------------------------------\n    # 1) GATHER ALL VAL IMAGES\n    # ---------------------------------------------------------------------\n    val_paths = sorted(glob.glob(os.path.join(val_dir, \"ILSVRC2012_val_*.JPEG\")))\n    print(f\"Found {len(val_paths)} validation images.\")\n\n    # 2) RANDOMLY PICK 10k OF THEM AS CENTROIDS\n    centroid_paths = random.sample(val_paths, 10000)\n    centroid_set = set(centroid_paths)\n\n    # 3) LOAD AND FLATTEN THE 10k CENTROID IMAGES\n    centroid_vectors = []\n    for i, cpath in tqdm(enumerate(centroid_paths), total=len(centroid_paths)):\n        centroid_vectors.append(load_and_flatten(cpath))\n        if (i + 1) % 1000 == 0:\n            print(f\"Loaded {i + 1} centroid images.\")\n    centroid_vectors = np.stack(centroid_vectors, axis=0).astype(np.float32)\n\n    # 4) BUILD A FAISS INDEX\n    d = centroid_vectors.shape[1]\n    index = faiss.IndexFlatL2(d)\n    index.add(centroid_vectors)\n\n    # ---------------------------------------------------------------------\n    # 5) PROCESS VALIDATION IMAGES (ALWAYS FULL)\n    # ---------------------------------------------------------------------\n    val_grouping = { str(i): [] for i in range(10000) }\n    remaining_val_paths = [p for p in val_paths if p not in centroid_set]\n\n    print(f\"Remaining val images to cluster: {len(remaining_val_paths)}\")\n\n    batch_size = 64\n    for start_idx in tqdm(range(0, len(remaining_val_paths), batch_size), total=(len(remaining_val_paths) // batch_size + 1)):\n        batch_paths = remaining_val_paths[start_idx:start_idx + batch_size]\n        batch_vecs = np.stack([load_and_flatten(p) for p in batch_paths], axis=0).astype(np.float32)\n        distances, indices = index.search(batch_vecs, 1)\n\n        for i, pth in enumerate(batch_paths):\n            cluster_id = indices[i, 0]\n            val_name = parse_val_name(pth)\n            xml_path = os.path.join(val_xml_dir, f\"ILSVRC2012_val_{val_name}.xml\")\n            class_id = extract_class_from_xml(xml_path)\n            val_grouping[str(cluster_id)].append(f\"{class_id}_{val_name}\")\n\n    # ---------------------------------------------------------------------\n    # 6) PROCESS TRAINING IMAGES (ONLY 1/5 BASED ON SPLIT INDEX)\n    # ---------------------------------------------------------------------\n    train_grouping = { str(i): [] for i in range(10000) }\n    train_paths = sorted(glob.glob(os.path.join(train_dir, \"*/*.JPEG\")))\n\n    total_train = len(train_paths)\n    chunk_size = total_train // 5\n    start_idx = split_index * chunk_size\n    end_idx = total_train if split_index == 4 else (split_index + 1) * chunk_size\n\n    train_paths = train_paths[start_idx:end_idx]  # Assign 1/5th of data to this session\n    print(f\"Processing training images {start_idx} to {end_idx} ({len(train_paths)})\")\n\n    for start_idx in tqdm(range(0, len(train_paths), batch_size), total=(len(train_paths) // batch_size + 1)):\n        batch_slice = train_paths[start_idx: start_idx + batch_size]\n        batch_vecs = np.stack([load_and_flatten(p) for p in batch_slice], axis=0).astype(np.float32)\n        distances, indices = index.search(batch_vecs, 1)\n\n        for i, p in enumerate(batch_slice):\n            cluster_id = indices[i, 0]\n            train_name = parse_train_name(p)\n            train_grouping[str(cluster_id)].append(train_name)\n\n    # ---------------------------------------------------------------------\n    # 7) SAVE JSON FILES\n    # ---------------------------------------------------------------------\n    val_json_path = os.path.join(output_dir, \"val_grouping.json\")\n    train_json_path = os.path.join(output_dir, f\"train_grouping_{split_index}.json\")\n\n    with open(val_json_path, \"w\") as f:\n        json.dump(val_grouping, f)\n    with open(train_json_path, \"w\") as f:\n        json.dump(train_grouping, f)\n\n    print(f\"\\nSaved val_grouping.json and train_grouping_{split_index}.json.\")\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python3 cluster_images.py --train_dir \"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train\" \\\n                          --val_dir \"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/val\" \\\n                          --val_xml_dir \"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Annotations/CLS-LOC/val\" \\\n                          --output_dir \"/kaggle/working\" \\\n                          --split_index 0 \\\n                          --seed 2411","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# END","metadata":{}}]}