Udforsk kunst på tværs af kulturer og medier med betingede k-nærmeste naboer

I denne artikel bruger du den betingede k-nærmeste naboer (k-NN) algoritme fra SynapseML til at finde visuelt lignende illustrationer. Du forespørger et datasæt med kunst fra Metropolitan Museum of Art i NYC og filtrerer efter kultur- og mediekategorier.

Forudsætninger

  • Få et Microsoft Fabric-abonnement. Du kan også tilmelde dig en gratis Prøveversion af Microsoft Fabric.

  • Log på Microsoft Fabric.

  • Skift til Fabric ved at bruge experience-switcheren nederst til venstre på din startside.

    Skærmbillede, der viser valget af Fabric i oplevelsesskifter-menuen.

  • Opret en ny notesbog.
  • Vedhæft din notesbog til et lakehouse. I venstre side af notesbogen skal du vælge Tilføj for at tilføje et eksisterende lakehouse eller oprette et nyt.

Importér biblioteker

I den første notebook-celle importeres de nødvendige Python-biblioteker:

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

Alle importer bør gennemføres uden fejl. Hvis du ser ModuleNotFoundError, bekræft at du bruger Fabric runtime 1.2 eller nyere.

Indlæs datasættet

Datasættet er en parketfil, der indeholder kunstmetadata fra Metropolitan Museum of Art. Indlæs det i en Spark DataFrame:

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

Datasættet indeholder cirka 51.000 rækker.

Datasætskema

Tabellen indeholder disse kolonner:

  • id: En unik identifikator for hvert kunstværk (for eksempel, 388395)
  • Titel: Kunstværkets titel gemt i museets database
  • Kunstner: Kunstværkskunstner som gemt i museets database
  • Thumbnail_Url: URL til et JPEG-miniaturebillede af kunstværket
  • Image_Url: Hjemmeside-URL til hele billedværket
  • Kultur: Kulturkategori (for eksempel japansk, amerikansk, italiensk)
  • Klassifikation: Mediumkategori (for eksempel malerier, keramik, glas)
  • Museum_Page: URL-link til kunstværkssiden på museets hjemmeside
  • Norm_Features: Forudberegnet billedindlejringsvektor (bruges til lighedssøgning)
  • Museet: Museet, der huser kunstværket

Definér kategorier og filtrer dataene

Definer de kultur- og mediekategorier, du ønsker at forespørge. Filtrer derefter datasættet, så det kun indeholder kunstværker, der matcher dine valgte kategorier:

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()}")

Outputtet viser et antal på flere tusinde rækker, afhængigt af de valgte kategorier.

Fit betingede k-NN-modeller

Skab to betingede k-NN-modeller – én betinget på mediet (Klassifikation) og én betinget på kultur. Hver model accepterer:

  • En outputkolonne til lagring af matches
  • En feature-kolonne , der indeholder billedindlejringsvektoren
  • En værdikolonne , der angiver, hvad der skal returneres for hvert match (miniature-URL)
  • En labelkolonne , der angiver betingningskategorien
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)
)

Definér matchnings- og visualiseringsmetoder

Definer hjælpefunktioner til at forespørge modellerne og vise resultater.

Funktionen add_matches() anvender en betinget k-NN-model på alle specificerede kategorier og tilføjer en match-kolonne for hver:

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

Og plot_img() funktionerne plot_urls() gengiver forespørgselsresultater som et billedgitter:

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()

Kør forespørgslen og visualiser resultaterne

Definér test_all() funktionen til at orkestrere forespørgsler på begge modeller og generere visualiseringer:

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

Vælg nu eksempler på kunst-ID'er fra det filtrerede datasæt og kør forespørgslen:

# 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="./")

To billedgitter vises i række. Det første gitter viser det originale kunstværk med nærmeste naboer på tværs af kulturer. Det andet gitter viser nærmeste naboer på tværs af medier.

Oprydning

Fjern cachede data og gemte filer, når du er færdig med at udforske:

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

Problem Årsag Opløsning
ModuleNotFoundError: No module named 'synapse.ml' Notesbog bruger ikke Fabric runtime Tjek at din notesbog er tilknyttet et Fabric lakehouse med runtime 1.2+
Py4JJavaError Under spark.read.parquet(...) Netværksforbindelsesproblem Bekræft, at dit arbejdsområde kan nås mmlspark.blob.core.windows.net på port 443
Tomt resultat fra test_all() (0 rækker) Udvalgte ID'er er ikke i det filtrerede datasæt Brug small_df.select("id").show(5) til at vælge gyldige ID'er fra de filtrerede data
HTTPError eller blanke billeder i visualisering Miniature-URL er ikke længere tilgængelig Nogle miniaturebilleder kan blive utilgængelige over tid. Funktionen plot_img viser "Image unavailable" ved mislykkede downloads.
OutOfMemoryError Under modeltilpasning Datasættet er for stort til tilgængelig hukommelse Reducer antallet af kategorier i mediums og cultures lister
Langsom modeltilpasning (>10 minutter) Stort datasæt med mange kategorier Start med færre kategorier (3 hver), og udvid så, når pipelinen virker

Hvordan betinget k-NN fungerer

Den betingede k-NN-model bygger på BallTree-datastrukturen . Et BallTree er et rekursivt binært træ, hvor hver node (eller "kugle") indeholder en partition af de datapunkter, du ønsker at forespørge.

Sådan bygger du et BallTree:

  1. Bestem "kugle"-centret tættest på hvert datapunkt baseret på en specificeret egenskab.
  2. Tildel hvert datapunkt til den nærmeste kugle.
  3. Gentag rekursivt og skab en struktur, der understøtter binære trægennemgange.

Denne struktur muliggør effektive k-nærmeste nabo-opslag ved hver bladnode.