.. DO NOT EDIT. .. THIS FILE WAS AUTOMATICALLY GENERATED BY SPHINX-GALLERY. .. TO MAKE CHANGES, EDIT THE SOURCE PYTHON FILE: .. "auto_examples/basics/demo_non_parametric_transform_digits.py" .. LINE NUMBERS ARE GIVEN BELOW. .. only:: html .. note:: :class: sphx-glr-download-link-note :ref:`Go to the end ` to download the full example code. .. rst-class:: sphx-glr-example-title .. _sphx_glr_auto_examples_basics_demo_non_parametric_transform_digits.py: UMAP Non-Parametric Transform on Handwritten Digits =================================================== We fit a reference UMAP embedding on one subset of the sklearn handwritten digits dataset and then map unseen samples with :meth:`torchdr.UMAP.transform` instead of refitting on the full dataset. This mimics a deployment workflow where a reference embedding is already available and new samples only arrive with their input features. We use the built-in digits dataset rather than full MNIST so the example stays fully self-contained and lightweight enough for the documentation gallery. By default, we fit the reference embedding on 1,400 points and transform 300 inference-only query points. We report a simple deployment metric: 10-NN label transfer accuracy from the reference embedding to the transformed query points. .. GENERATED FROM PYTHON SOURCE LINES 18-39 .. code-block:: Python # License: BSD 3-Clause License import os import matplotlib.pyplot as plt from sklearn.datasets import load_digits from sklearn.decomposition import PCA from sklearn.model_selection import train_test_split from sklearn.neighbors import KNeighborsClassifier from torchdr import UMAP RANDOM_STATE = 0 N_REFERENCE = int(os.environ.get("TORCHDR_EXAMPLE_N_REFERENCE", "1400")) N_QUERY = int(os.environ.get("TORCHDR_EXAMPLE_N_QUERY", "300")) PCA_DIM = 50 MAX_ITER = int(os.environ.get("TORCHDR_EXAMPLE_MAX_ITER", "150")) .. GENERATED FROM PYTHON SOURCE LINES 40-45 Load handwritten digits and prepare a reference/query split ----------------------------------------------------------- We keep a reference pool that is used during fitting and a disjoint query pool that is embedded later only through :meth:`transform`. .. GENERATED FROM PYTHON SOURCE LINES 45-65 .. code-block:: Python digits = load_digits() X = (digits.data / 16.0).astype("float32") y = digits.target.astype("int64") if N_REFERENCE + N_QUERY > len(X): raise ValueError( f"Requested {N_REFERENCE + N_QUERY} samples but digits only contains {len(X)}." ) X_reference, X_query, y_reference, y_query = train_test_split( X, y, train_size=N_REFERENCE, test_size=N_QUERY, stratify=y, random_state=RANDOM_STATE, ) .. GENERATED FROM PYTHON SOURCE LINES 66-71 Preprocess features and fit the reference embedding --------------------------------------------------- We reduce the input features with PCA before fitting UMAP, then transform the held-out query set with the non-parametric transform path. .. GENERATED FROM PYTHON SOURCE LINES 71-91 .. code-block:: Python pca = PCA(n_components=PCA_DIM) X_reference_pca = pca.fit_transform(X_reference).astype("float32") X_query_pca = pca.transform(X_query).astype("float32") umap = UMAP( n_components=2, n_neighbors=15, max_iter=MAX_ITER, init="pca", optimizer="SGD", backend=None, device="cpu", random_state=RANDOM_STATE, ) Z_reference = umap.fit_transform(X_reference_pca) Z_query = umap.transform(X_query_pca, X_train=X_reference_pca) .. GENERATED FROM PYTHON SOURCE LINES 92-97 Evaluate deployment-time label transfer --------------------------------------- In a production setting, a simple use case is to classify or annotate new points by comparing them to the already embedded reference set. .. GENERATED FROM PYTHON SOURCE LINES 97-107 .. code-block:: Python knn = KNeighborsClassifier(n_neighbors=10) knn.fit(Z_reference, y_reference) label_transfer_accuracy = knn.score(Z_query, y_query) print(f"Reference samples used for fit: {len(X_reference_pca)}") print(f"Query samples used for transform only: {len(X_query_pca)}") print(f"10-NN label transfer accuracy: {label_transfer_accuracy:.3f}") .. rst-class:: sphx-glr-script-out .. code-block:: none Reference samples used for fit: 1400 Query samples used for transform only: 300 10-NN label transfer accuracy: 0.973 .. GENERATED FROM PYTHON SOURCE LINES 108-110 Visualize the reference and transformed query points ---------------------------------------------------- .. GENERATED FROM PYTHON SOURCE LINES 110-152 .. code-block:: Python fig, axes = plt.subplots(1, 2, figsize=(12, 5)) axes[0].scatter( Z_reference[:, 0], Z_reference[:, 1], c=y_reference, cmap="tab10", s=5, alpha=0.55, ) axes[0].set_title(f"UMAP reference embedding\nfit on {len(X_reference_pca)} points") axes[0].set_xticks([]) axes[0].set_yticks([]) axes[1].scatter( Z_reference[:, 0], Z_reference[:, 1], c="lightgray", s=3, alpha=0.12, ) scatter = axes[1].scatter( Z_query[:, 0], Z_query[:, 1], c=y_query, cmap="tab10", s=7, alpha=0.8, ) axes[1].set_title( "UMAP transform on inference-only query points\n" f"transform on {len(X_query_pca)} points, " f"10-NN accuracy = {label_transfer_accuracy:.3f}" ) axes[1].set_xticks([]) axes[1].set_yticks([]) handles, labels = scatter.legend_elements(prop="colors") fig.legend(handles, labels, loc="lower center", ncol=10, frameon=False) plt.subplots_adjust(bottom=0.16, wspace=0.08) plt.show() .. image-sg:: /auto_examples/basics/images/sphx_glr_demo_non_parametric_transform_digits_001.png :alt: UMAP reference embedding fit on 1400 points, UMAP transform on inference-only query points transform on 300 points, 10-NN accuracy = 0.973 :srcset: /auto_examples/basics/images/sphx_glr_demo_non_parametric_transform_digits_001.png :class: sphx-glr-single-img .. rst-class:: sphx-glr-timing **Total running time of the script:** (0 minutes 13.617 seconds) .. _sphx_glr_download_auto_examples_basics_demo_non_parametric_transform_digits.py: .. only:: html .. container:: sphx-glr-footer sphx-glr-footer-example .. container:: sphx-glr-download sphx-glr-download-jupyter :download:`Download Jupyter notebook: demo_non_parametric_transform_digits.ipynb ` .. container:: sphx-glr-download sphx-glr-download-python :download:`Download Python source code: demo_non_parametric_transform_digits.py ` .. container:: sphx-glr-download sphx-glr-download-zip :download:`Download zipped: demo_non_parametric_transform_digits.zip ` .. only:: html .. rst-class:: sphx-glr-signature `Gallery generated by Sphinx-Gallery `_