{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":114397,"databundleVersionId":13696770,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":282571811,"sourceType":"kernelVersion"}],"dockerImageVersionId":31192,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Clustering the BioTrove Dataset: genus and species clusters with BioCLIP\n\nThis notebook clusters the 49633 images into 746 genus clusters and 4982 species clusters.\n\nIt uses `AgglomerativeClustering` with Euclidean distances and *Ward* linkage, based on BioCLIP embeddings computed in another notebook.\n\nReferences\n- Competition: [Clustering the BioTrove Dataset](https://www.kaggle.com/competitions/biotrove-clustering)\n- [BioCLIP: A Vision Foundation Model for the Tree of Life](https://arxiv.org/pdf/2311.18803)\n- https://imageomics.github.io/pybioclip/","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport pickle\nfrom matplotlib import pyplot as plt\n\nfrom sklearn.cluster import AgglomerativeClustering\nfrom sklearn.manifold import TSNE\n","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-12-12T18:46:04.661288Z","iopub.execute_input":"2025-12-12T18:46:04.661580Z","iopub.status.idle":"2025-12-12T18:46:09.309529Z","shell.execute_reply.started":"2025-12-12T18:46:04.661551Z","shell.execute_reply":"2025-12-12T18:46:09.308357Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Read the metadata (dataframe with columns 'hash_id' and 'family')\nmetadata = pd.read_csv('/kaggle/input/biotrove-clustering/metadata.csv')\nfamilies = np.unique(metadata.family)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-12T18:46:09.311740Z","iopub.execute_input":"2025-12-12T18:46:09.312268Z","iopub.status.idle":"2025-12-12T18:46:09.421085Z","shell.execute_reply.started":"2025-12-12T18:46:09.312241Z","shell.execute_reply":"2025-12-12T18:46:09.420011Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Read the embeddings of the images\nwith open(f'/kaggle/input/cbtd-bioclip-embeddings/embedding.pickle', 'rb') as f:\n    embeddings = pickle.load(f) # array of shape (49633, 768)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-12T18:46:09.422060Z","iopub.execute_input":"2025-12-12T18:46:09.422414Z","iopub.status.idle":"2025-12-12T18:46:10.559348Z","shell.execute_reply.started":"2025-12-12T18:46:09.422389Z","shell.execute_reply":"2025-12-12T18:46:10.558150Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n# Cluster the images to genus clusters\n# Plot the clusters for the first few families\n# The 2d embedding is computed by T-SNE\n# The colors are the clusters of AgglomerativeClustering on the original BioCLIP embedding\ngenus_cluster = []\nn_clusters_list = []\n_, axs = plt.subplots(4, 4, figsize=(15, 15))\naxs = axs.ravel()\nfor i in range(len(families)):\n    family = families[i]\n    family_mask = metadata.family == family\n    embeddings_f = embeddings[family_mask]\n\n    # Guess the number of genera\n    model = AgglomerativeClustering(n_clusters=3,\n                                    metric='euclidean',\n                                    linkage='ward',\n                                    compute_distances=True)\n    model.fit(embeddings_f)\n    diff_distances = np.diff(model.distances_)[:-1]\n    n_clusters = len(model.distances_) - np.argmax(diff_distances)\n    n_clusters = min(n_clusters, len(embeddings_f) // 20) # every genus has at least 20 images\n    if n_clusters >= 7:\n        print(f\"{family:20} {n_clusters} genera\")\n    if n_clusters * 20 > len(embeddings_f):\n        raise ValueError(f\"{family} cannot have {n_clusters} clusters for {len(embeddings_f)} images.\")\n    n_clusters_list.append(n_clusters)\n\n    # Cluster the images for submission\n    model.set_params(n_clusters=n_clusters)\n    clusters = model.fit_predict(embeddings_f)\n\n    # Plot the clusters for the first few families\n    if i < len(axs):\n        tsne = TSNE()\n        embeddings_2d = tsne.fit_transform(embeddings_f)\n        ax = axs[i]\n        ax.scatter(embeddings_2d.T[0], embeddings_2d.T[1], s=8, c=clusters, cmap='brg')\n        ax.set_xticks([])\n        ax.set_yticks([])\n        ax.set_title(family)\n\n    # Format the output for submission\n    clusters = [f\"{family}_g_{c}\" for c in clusters]\n    genus_cluster.append(pd.Series(clusters, index=metadata.hash_id[family_mask]))\n\nplt.tight_layout()\nplt.show()\n\n# Show the distribution of n_clusters; median should be 5\nprint(f\"Median genus clusters: {np.median(n_clusters_list)}\")\nvc = np.unique(n_clusters_list, return_counts=True)\nplt.bar(*vc)\nplt.xlabel('# genus clusters')\nplt.ylabel('count')\nplt.show()\n\ngenus_cluster = pd.concat(genus_cluster) # Series with hash_id as index and genus cluster as value\ngenus_cluster.name = 'genus_cluster'\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-12T18:46:10.560687Z","iopub.execute_input":"2025-12-12T18:46:10.561089Z","iopub.status.idle":"2025-12-12T18:46:45.367122Z","shell.execute_reply.started":"2025-12-12T18:46:10.561041Z","shell.execute_reply":"2025-12-12T18:46:45.366055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cluster the images to species clusters\n\ndef cluster_for_factor(factor):\n    species_cluster = []\n    chs_list, dbs_list, ss_list = [], [], []\n    for family in families:\n        family_mask = metadata.family == family\n        embeddings_f = embeddings[family_mask]\n    \n        model = AgglomerativeClustering(n_clusters=int(len(embeddings_f) * factor + 0.5),\n                                        metric='euclidean',\n                                        linkage='ward')\n        clusters = model.fit_predict(embeddings_f)\n        clusters = [f\"{family}_s_{c}\" for c in clusters]\n        species_cluster.append(pd.Series(clusters, index=metadata.hash_id[family_mask]))\n        # chs = calinski_harabasz_score(embeddings_f, clusters) # higher is better\n        # dbs = davies_bouldin_score(embeddings_f, clusters) # lower is better\n        # ss = silhouette_score(embeddings_f, clusters, metric='cosine') # higher is better\n        # chs_list.append(chs)\n        # dbs_list.append(dbs)\n        # ss_list.append(ss)\n        # print(f\"{family:25} {model.n_clusters_:3}/{family_mask.sum():3} {chs:8.3f}\") # n_clusters/n_images\n\n    # print(f\"{factor=:.3f} chs={np.mean(chs_list):5.2f} (higher is better)\")\n    # print(f\"{factor=:.3f} dbs={np.mean(dbs_list):5.2f} (lower is better)\")\n    # print(f\"{factor=:.3f} ss={np.mean(ss_list):7.5f} (higher is better)\")\n    \n    species_cluster = pd.concat(species_cluster) # Series with hash_id as index and species cluster as value\n    species_cluster.name = 'species_cluster'\n    return species_cluster\n    \nfactor = 0.1\nspecies_cluster = cluster_for_factor(factor)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-12T18:46:45.368420Z","iopub.execute_input":"2025-12-12T18:46:45.368880Z","iopub.status.idle":"2025-12-12T18:46:50.181503Z","shell.execute_reply.started":"2025-12-12T18:46:45.368853Z","shell.execute_reply":"2025-12-12T18:46:50.180487Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Write the submission file\n\nsub = pd.DataFrame({\n    'hash_id': metadata.hash_id,\n    'family_cluster': metadata.family,\n    # 'genus_cluster': metadata.family,\n    # 'species_cluster': np.arange(len(metadata))\n})\nsub = sub.join(genus_cluster, on='hash_id', how='inner', validate='1:1')\nsub = sub.join(species_cluster, on='hash_id', how='inner', validate='1:1')\ndisplay(sub)\nsub.to_csv('submission.csv', index=False)\nprint()\n!head submission.csv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-12T18:46:50.182628Z","iopub.execute_input":"2025-12-12T18:46:50.182893Z","iopub.status.idle":"2025-12-12T18:46:50.712999Z","shell.execute_reply.started":"2025-12-12T18:46:50.182872Z","shell.execute_reply":"2025-12-12T18:46:50.711714Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(np.unique(sub.family_cluster)), len(np.unique(sub.genus_cluster)), len(np.unique(sub.species_cluster))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-12T18:46:50.716644Z","iopub.execute_input":"2025-12-12T18:46:50.716971Z","iopub.status.idle":"2025-12-12T18:46:50.810453Z","shell.execute_reply.started":"2025-12-12T18:46:50.716940Z","shell.execute_reply":"2025-12-12T18:46:50.809252Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}