Esplora l'arte attraverso culture e tecniche con l'algoritmo dei k-nearest neighbors condizionali

In questo articolo si usa l'algoritmo k-nearest neighbors condizionale (k-NN) di SynapseML per trovare grafica visivamente simile. Interroghi un dataset di opere d'arte del Metropolitan Museum of Art di New York, filtrandolo per le categorie relative alla cultura e al materiale.

Prerequisiti

  • Abbonati a Microsoft Fabric. Oppure, registrati per una versione di prova gratuita di Microsoft Fabric.

  • Accedi a Microsoft Fabric.

  • Passare a Fabric usando il selettore di esperienza nell'angolo in basso a sinistra della home page.

    Screenshot che mostra la selezione di

  • Creare un nuovo notebook.
  • Collegare il notebook a un lakehouse. Sul lato sinistro del tuo notebook, seleziona Aggiungi per aggiungere un lakehouse esistente o crearne uno nuovo.

Importare librerie

Nella prima cella del notebook importare le librerie di Python necessarie:

from pyspark.sql.types import BooleanType
from pyspark.sql.functions import lit, array, udf
from synapse.ml.nn import ConditionalKNN
from PIL import Image
from io import BytesIO

import requests
import numpy as np
import matplotlib.pyplot as plt

Tutte le importazioni devono essere completate senza errori. Se viene visualizzato ModuleNotFoundError, verificare di usare Fabric runtime 1.2 o versione successiva.

Caricare il set di dati

Il set di dati è un file parquet contenente i metadati delle opere d'arte del Metropolitan Museum of Art. Caricarlo in un dataframe Spark:

df = spark.read.parquet(
    "wasbs://publicwasb@mmlspark.blob.core.windows.net/met_and_rijks.parquet"
)
display(df.drop("Norm_Features"))

Il set di dati contiene circa 51.000 righe.

Schema del set di dati

La tabella contiene queste colonne:

  • id: identificatore univoco per ogni pezzo d'arte (ad esempio, 388395)
  • Titolo: Titolo del pezzo d'arte archiviato nel database del museo
  • Artista: artista dell'opera come memorizzato nel database del museo
  • Thumbnail_Url: URL di un'anteprima JPEG della parte d'arte
  • Image_Url: URL del sito Web dell'immagine dell'arte completa
  • Cultura: categoria cultura (ad esempio, giapponese, americano, italiano)
  • Classificazione: categoria media (ad esempio , dipinti, ceramica, vetro)
  • Museum_Page: collegamento URL alla pagina del pezzo d'arte nel sito Web del museo
  • Norm_Features: vettore di incorporamento di immagini pre-calcolate (usato per la ricerca di somiglianza)
  • Museo: Il museo che ospita il pezzo d'arte

Definire le categorie e filtrare i dati

Definire le impostazioni cultura e le categorie medie di cui eseguire la query. Filtrare quindi il set di dati in modo da includere solo immagini corrispondenti alle categorie selezionate:

mediums = ["paintings", "glass", "ceramics"]
cultures = ["japanese", "american", "african (general)"]

# For more categories, uncomment the extended lists:
# mediums = ['prints', 'drawings', 'ceramics', 'textiles', 'paintings',
#            'musical instruments', 'glass', 'accessories', 'photographs',
#            'metalwork', 'sculptures', 'weapons', 'stone', 'precious',
#            'paper', 'woodwork', 'leatherwork', 'uncategorized']
# cultures = ['african (general)', 'american', 'ancient american',
#             'ancient asian', 'ancient european', 'ancient middle-eastern',
#             'asian (general)', 'austrian', 'belgian', 'british', 'chinese',
#             'czech', 'dutch', 'egyptian', 'european (general)', 'french',
#             'german', 'greek', 'iranian', 'italian', 'japanese',
#             'latin american', 'middle eastern', 'roman', 'russian',
#             'south asian', 'southeast asian', 'spanish', 'swiss', 'various']

classes = cultures + mediums
medium_set = set(mediums)
culture_set = set(cultures)

small_df = df.where(
    udf(
        lambda medium, culture: (medium in medium_set) or (culture in culture_set),
        BooleanType(),
    )("Classification", "Culture")
)

small_df.cache()
print(f"Filtered dataset row count: {small_df.count()}")

L'output mostra un conteggio di diverse migliaia di righe, a seconda delle categorie selezionate.

Adattare i modelli k-NN condizionali

Crea due modelli k-NN condizionali: uno condizionato dal supporto (Classification) e uno condizionato dalla cultura. Ogni modello accetta:

  • Una colonna di output per memorizzare le corrispondenze
  • Colonna delle funzionalità contenente il vettore di incorporamento dell'immagine
  • Una colonna dei valori che specifica cosa restituire per ogni corrispondenza (URL della miniatura)
  • Una colonna delle etichette che indica la categoria di condizionamento
medium_cknn = (
    ConditionalKNN()
    .setOutputCol("Matches")
    .setFeaturesCol("Norm_Features")
    .setValuesCol("Thumbnail_Url")
    .setLabelCol("Classification")
    .fit(small_df)
)
culture_cknn = (
    ConditionalKNN()
    .setOutputCol("Matches")
    .setFeaturesCol("Norm_Features")
    .setValuesCol("Thumbnail_Url")
    .setLabelCol("Culture")
    .fit(small_df)
)

Definire metodi di corrispondenza e visualizzazione

Definire le funzioni helper per eseguire query sui modelli e visualizzare i risultati.

La add_matches() funzione applica un modello k-NN condizionale in tutte le categorie specificate, aggiungendo una colonna corrispondenze per ognuna:

def add_matches(classes, cknn, df):
    """Apply conditional k-NN for each category label, adding match columns."""
    results = df
    for label in classes:
        results = cknn.transform(
            results.withColumn("conditioner", array(lit(label)))
        ).withColumnRenamed("Matches", "Matches_{}".format(label))
    return results

Le funzioni plot_img() e plot_urls() visualizzano i risultati della query in una griglia di immagini:

def plot_img(axis, url, title):
    """Download and display an image from a URL on a matplotlib axis."""
    try:
        response = requests.get(url, timeout=10)
        response.raise_for_status()
        img = Image.open(BytesIO(response.content)).convert("RGB")
        axis.imshow(img, aspect="equal")
    except Exception as e:
        axis.text(0.5, 0.5, "Image\nunavailable", ha="center", va="center", fontsize=6)
    if title is not None:
        axis.set_title(title, fontsize=10)
    axis.axis("off")


def plot_urls(url_arr, titles, filename):
    """Create a grid visualization of artwork thumbnails and save to file."""
    nx, ny = url_arr.shape

    fig, axes = plt.subplots(ny, nx, figsize=(nx * 5, ny * 5), dpi=150)

    # Reshape required for a single-image query
    if len(axes.shape) == 1:
        axes = axes.reshape(1, -1)

    for i in range(nx):
        for j in range(ny):
            if j == 0:
                plot_img(axes[j, i], url_arr[i, j], titles[i])
            else:
                plot_img(axes[j, i], url_arr[i, j], None)

    plt.tight_layout()
    plt.savefig(filename, dpi=150)
    plt.show()

Eseguire la query e visualizzare i risultati

Definire la funzione test_all() per orchestrare l'interrogazione di entrambi i modelli e la generazione di visualizzazioni:

def test_all(data, cknn_medium, cknn_culture, test_ids, root):
    """Query both k-NN models for given art IDs and save visualizations."""
    is_match = udf(lambda obj: obj in test_ids, BooleanType())
    test_df = data.where(is_match("id"))

    test_count = test_df.count()
    if test_count == 0:
        print("Warning: No matching art IDs found. Verify IDs exist in the filtered dataset.")
        return None

    print(f"Querying {test_count} artwork(s)...")

    results_df_medium = add_matches(mediums, cknn_medium, test_df)
    results_df_culture = add_matches(cultures, cknn_culture, results_df_medium)

    results = results_df_culture.collect()

    original_urls = [row["Thumbnail_Url"] for row in results]

    culture_urls = [
        [row["Matches_{}".format(label)][0]["value"] for row in results]
        for label in cultures
    ]
    culture_url_arr = np.array([original_urls] + culture_urls)[:, :]
    plot_urls(culture_url_arr, ["Original"] + cultures, root + "matches_by_culture.png")

    medium_urls = [
        [row["Matches_{}".format(label)][0]["value"] for row in results]
        for label in mediums
    ]
    medium_url_arr = np.array([original_urls] + medium_urls)[:, :]
    plot_urls(medium_url_arr, ["Original"] + mediums, root + "matches_by_medium.png")

    return results_df_culture

Ora, selezionare gli ID delle opere di esempio dal dataset filtrato ed eseguire la query:

# Select 3 sample artwork IDs from the filtered dataset
sample_rows = small_df.select("id").take(3)
selected_ids = {row["id"] for row in sample_rows}
print(f"Selected art IDs: {selected_ids}")

# Run the query and generate visualizations
result_df = test_all(small_df, medium_cknn, culture_cknn, selected_ids, root="./")

Vengono visualizzate due griglie di immagini inline. La prima griglia mostra l'opera d'arte originale con le opere più simili nelle diverse culture. La seconda griglia mostra i vicini più prossimi tra i media.

Cleanup

Rimuovere i dati memorizzati nella cache e i file salvati al termine dell'esplorazione:

small_df.unpersist()
import os
for f in ["./matches_by_culture.png", "./matches_by_medium.png"]:
    if os.path.exists(f):
        os.remove(f)
        print(f"Removed {f}")
print("OK Cleanup complete")

Troubleshooting

Issue Motivo Risoluzione
ModuleNotFoundError: No module named 'synapse.ml' Notebook che non usa il runtime di Fabric Verificare che il notebook sia collegato a un Fabric lakehouse con runtime 1.2+
Py4JJavaError durante spark.read.parquet(...) Problema di connettività di rete Verificare che l'area di lavoro possa raggiungere mmlspark.blob.core.windows.net sulla porta 443
Nessun risultato da test_all() (0 righe) Gli ID selezionati non sono inclusi nel set di dati filtrato Usare small_df.select("id").show(5) per selezionare gli ID validi dai dati filtrati
HTTPError o immagini vuote nella visualizzazione URL anteprima non più accessibile Alcune miniature potrebbero non essere più disponibili con il tempo. La plot_img funzione visualizza "Immagine non disponibile" per i download non riusciti.
OutOfMemoryError durante l'adattamento del modello Set di dati troppo grande per la memoria disponibile Ridurre il numero di categorie negli elenchi mediums e cultures
Adattamento lento del modello (>10 minuti) Set di dati di grandi dimensioni con molte categorie Inizia con meno categorie (3 per volta), poi amplia quando la pipeline funziona

Come funziona la k-NN condizionale

Il modello k-NN condizionale si basa sulla struttura dei dati BallTree . Un BallTree è un albero binario ricorsivo in cui ogni nodo (o "ball") contiene una partizione dei punti su cui si desidera eseguire la query.

Per creare un BallTree:

  1. Determinare il centro della "sfera" più vicino a ciascun punto dati, in base a una caratteristica specificata.
  2. Assegnare ogni punto di dati alla sfera più vicina.
  3. Ripeti ricorsivamente, creando una struttura che supporti gli attraversamenti dell'albero binario.

Questa struttura consente ricerche efficienti dei k vicini più prossimi in ciascun nodo foglia.