Documenten classificeren met meer dan 500 labels

ai_classify accepteert maximaal 500 labels per gesprek. Voor grotere taxonomieën filtert u labels voor elk document vooraf op basis van embeddinggelijkenis en past u vervolgens ai_classify toe op de shortlist van top-K-kandidaten. In deze zelfstudie leert u hoe u het optimale K kunt vinden: het kleinste aantal kandidaten dat de nauwkeurigheid behoudt.

Note

NEAREST BY vereist Databricks Runtime 18 of hoger of serverloos. Databricks Runtime 18 is nieuwer dan Databricks Runtime 18.0, 18.1 en 18.2.

Voordat u begint

  • Een werkruimte met Unity Catalog-functionaliteit met toegang tot ai_classify (zie beschikbaarheid).
  • Databricks Runtime 18+ of Serverloos (vereist voor NEAREST BY).
  • Een Delta-tabel met documenten die moeten worden geclassificeerd.
  • Een Delta-tabel met labels met een sleutelkolom en een optionele beschrijvingskolom.
  • Een kleine evaluatieset met ground-truthlabels, of de mogelijkheid om er een te maken (optie B in stap 4).

0. Configuratie

Stel de tabelnamen, kolomnamen en het insluitmodel in. De helperfunctie top_k_labels_json bouwt de JSON-expressie op die aan ai_classify wordt doorgegeven voor de top-K-kandidaten van elk document.

# -- Your tables --
DOCS_TABLE = "path.to.your_docs_table"                   # table of documents to classify
DOCS_TEXT_COL = "document"                              # column with text to classify
DOCS_ID_COL = None                                      # unique ID column; set to None to auto-generate via md5

LABELS_TABLE = "path.to.your_labels_table"               # table of labels
LABELS_KEY_COL = "label"                                # column with label value
LABELS_DESC_COL = "description"                         # description column; set to None if labels have no descriptions

# -- Embedding model --
EMBEDDING_MODEL = "databricks-qwen3-embedding-0-6b"     # compact model, good default for English text

# -- K values to sweep --
K_VALUES = [10, 20, 50, 100, 200, 500]

# -- Eval set size (if you need to create one) --
EVAL_SAMPLE_SIZE = 100                                   # docs to sample for manual labeling
doc_id_expr = DOCS_ID_COL if DOCS_ID_COL else f"md5({DOCS_TEXT_COL})"
label_embed_text = (
    f"concat({LABELS_KEY_COL}, ': ', {LABELS_DESC_COL})"
    if LABELS_DESC_COL
    else LABELS_KEY_COL
)
def top_k_labels_json(prefix=""):
    """Build a JSON expression for collected labels from NEAREST BY results."""
    col_prefix = f"{prefix}." if prefix else ""
    if LABELS_DESC_COL:
        return f"to_json(map_from_entries(collect_list(struct({col_prefix}{LABELS_KEY_COL}, {col_prefix}{LABELS_DESC_COL}))))"
    else:
        return f"to_json(collect_list({col_prefix}{LABELS_KEY_COL}))"

print(f"Docs table:    {DOCS_TABLE} (text: {DOCS_TEXT_COL}, id: {doc_id_expr})")
print(f"Labels table:  {LABELS_TABLE} (key: {LABELS_KEY_COL}, desc: {LABELS_DESC_COL})")
print(f"Embed text:    {label_embed_text}")
print(f"K sweep:       {K_VALUES}")

1. De labels insluiten

Voer dit eenmaal uit. Voer alleen opnieuw uit wanneer de taxonomie wordt gewijzigd.

spark.sql(f"""
CREATE OR REPLACE TABLE label_embeddings AS
SELECT
  {LABELS_KEY_COL},
  {f'{LABELS_DESC_COL},' if LABELS_DESC_COL else ''}
  cast(
    ai_query('{EMBEDDING_MODEL}', {label_embed_text}) AS ARRAY<FLOAT>
  ) AS embedding
FROM {LABELS_TABLE}
""")

label_count = spark.sql("SELECT count(*) AS n FROM label_embeddings").first()["n"]
print(f"Embedded {label_count} labels")

2. De documenten insluiten

spark.sql(f"""
CREATE OR REPLACE TABLE doc_embeddings AS
SELECT
  {doc_id_expr} AS id,
  {DOCS_TEXT_COL} AS doc_text,
  cast(
    ai_query('{EMBEDDING_MODEL}', {DOCS_TEXT_COL}) AS ARRAY<FLOAT>
  ) AS embedding
FROM {DOCS_TABLE}
""")

doc_count = spark.sql("SELECT count(*) AS n FROM doc_embeddings").first()["n"]
print(f"Embedded {doc_count} documents")

3. De top-K-labels ophalen met NEAREST BY

NEAREST BY voert rechtstreeks een benaderende nearest-neighbor-join uit — zonder tussenliggende N×M-tabel. Voor elk document worden in één doorgang de K meest vergelijkbare labels geretourneerd.

# Preview: top-5 nearest labels for a sample of documents
preview_df = spark.sql(f"""
SELECT
  d.id,
  l.{LABELS_KEY_COL}
  {f', l.{LABELS_DESC_COL}' if LABELS_DESC_COL else ''}
FROM doc_embeddings d
INNER JOIN label_embeddings l
APPROX NEAREST 5 BY SIMILARITY vector_cosine_similarity(d.embedding, l.embedding)
LIMIT 20
""")
preview_df.display()

4. Een basis-waarheidsevaluatieset voorbereiden

Voor K-tuning is een kleine set documenten met bekende juiste labels vereist. Als u al een evaluatietabel hebt, zet u EVAL_TABLE in de volgende cel en slaat u de samplecel over. Zo niet, dan selecteert de tweede cel documenten die u handmatig kunt labelen en opnieuw kunt importeren.

# Option A: point to your existing eval table
# Must have columns: id (matching doc_embeddings.id) and ground_truth_label
EVAL_TABLE = dbutils.widgets.get("eval_table")  # read from notebook widget

if EVAL_TABLE:
    eval_df = spark.table(EVAL_TABLE)
    print(f"Loaded {eval_df.count()} eval examples from {EVAL_TABLE}")
else:
    print("No eval table set — run the next cell to sample documents for labeling.")
# Option B: sample documents for manual labeling
if not EVAL_TABLE:
    sample_df = spark.sql(f"""
        SELECT id, doc_text
        FROM doc_embeddings
        ORDER BY rand()
        LIMIT {EVAL_SAMPLE_SIZE}
    """)
    sample_df.display()
    print(f"\nSampled {EVAL_SAMPLE_SIZE} documents.")
    print("Next steps:")
    print("  1. Export these rows (copy the table above or save to CSV)")
    print("  2. Add a 'ground_truth_label' column and fill in the correct label for each doc")
    print("  3. Re-import as a Delta table and set EVAL_TABLE above")
    print("  4. Re-run cell 4 (Option A) to load it")

5. Meting Recall@K

Recall@K controleert of het label ground-truth wordt weergegeven in de top-K-insluitingskandidaten. Dit is een metriek die alleen gegevens ophaalt — er wordt geen ai_classify aangeroepen en deze wordt direct uitgevoerd.

Als de recall laag is bij een bepaalde K, kan ai_classify onmogelijk het juiste antwoord teruggeven, omdat het juiste label al uit de kandidatenset was uitgesloten voordat de classificatie überhaupt werd uitgevoerd.

assert EVAL_TABLE, "Set EVAL_TABLE in cell 4 before running K-tuning."

spark.sql(f"CREATE OR REPLACE TEMP VIEW eval_set AS SELECT * FROM {EVAL_TABLE}")

recall_results = []
for k in K_VALUES:
    row = spark.sql(f"""
        SELECT
          {k} AS k,
          count(*) AS eval_size,
          sum(CASE WHEN hit THEN 1 ELSE 0 END) AS hits,
          round(sum(CASE WHEN hit THEN 1 ELSE 0 END) / count(*), 4) AS recall_at_k
        FROM (
          SELECT
            e.id,
            array_contains(
              collect_list(l.{LABELS_KEY_COL}),
              e.ground_truth_label
            ) AS hit
          FROM eval_set e
          JOIN doc_embeddings d ON d.id = e.id
          INNER JOIN label_embeddings l
          APPROX NEAREST {k} BY SIMILARITY vector_cosine_similarity(d.embedding, l.embedding)
          GROUP BY e.id, e.ground_truth_label
        )
    """).first()
    recall_results.append(row.asDict())
    print(f"  K={k:>4d}  →  Recall@K = {row['recall_at_k']:.2%}  ({row['hits']}/{row['eval_size']})")

recall_df = spark.createDataFrame(recall_results)
recall_df.display()

6. Meet de nauwkeurigheid van begin tot eind

Bouw voor elke K de top-K-labelset per eval-document, voer ai_classify uit en vergelijk met de referentiegegevens.

Deze stap roept ai_classify aan en kost meer dan de recallcontrole. Start met de K-waarden waarbij de recall al redelijk is.

accuracy_results = []

for k in K_VALUES:
    # Get top-K labels per eval doc using NEAREST BY
    spark.sql(f"""
        CREATE OR REPLACE TEMP VIEW eval_top_labels AS
        SELECT
          d.id,
          {top_k_labels_json('l')} AS labels
        FROM eval_set e
        JOIN doc_embeddings d ON d.id = e.id
        INNER JOIN label_embeddings l
        APPROX NEAREST {k} BY SIMILARITY vector_cosine_similarity(d.embedding, l.embedding)
        GROUP BY d.id
    """)

    # Materialize ai_classify first (returns VARIANT, and is non-deterministic so can't go inside aggregate)
    spark.sql(f"""
        CREATE OR REPLACE TEMP VIEW eval_predictions AS
        SELECT
          e.id,
          e.ground_truth_label,
          get_json_object(cast(ai_classify(d.doc_text, t.labels, map('version', '2.0')) as string), '$.response[0]') AS predicted_label
        FROM eval_set e
        JOIN doc_embeddings d ON d.id = e.id
        JOIN eval_top_labels t ON t.id = e.id
    """)

    row = spark.sql(f"""
        SELECT
          {k} AS k,
          count(*) AS eval_size,
          sum(CASE WHEN predicted_label = ground_truth_label THEN 1 ELSE 0 END) AS correct,
          round(
            sum(CASE WHEN predicted_label = ground_truth_label THEN 1 ELSE 0 END) / count(*),
            4
          ) AS accuracy
        FROM eval_predictions
    """).first()

    accuracy_results.append(row.asDict())
    print(f"  K={k:>4d}  →  Accuracy = {row['accuracy']:.2%}  ({row['correct']}/{row['eval_size']})")

accuracy_df = spark.createDataFrame(accuracy_results)
accuracy_df.display()

7. Resultaten vergelijken en K kiezen

In de onderstaande grafiek ziet u Recall@K en end-to-end nauwkeurigheid naast elkaar. Kies het kleinste K waar de nauwkeurigheid stopt met verbeteren - groter K betekent tragere classificatie zonder kwaliteitswinst.

import pandas as pd
import matplotlib.pyplot as plt

recall_pd = pd.DataFrame(recall_results)
accuracy_pd = pd.DataFrame(accuracy_results)
combined = recall_pd.merge(accuracy_pd, on="k", suffixes=("_recall", "_acc"))

fig, ax = plt.subplots(figsize=(10, 5))
ax.plot(combined["k"], combined["recall_at_k"], "o-", label="Recall@K", linewidth=2)
ax.plot(combined["k"], combined["accuracy"], "s--", label="End-to-end accuracy", linewidth=2)
ax.set_xlabel("K (candidate labels per document)")
ax.set_ylabel("Score")
ax.set_title("K-Tuning: Recall@K vs End-to-End Accuracy")
ax.set_ylim(0, 1.05)
ax.set_xticks(combined["k"])
ax.legend()
ax.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()

print("\nFull results:")
print(combined[["k", "recall_at_k", "accuracy"]].to_string(index=False))
# Pick your K based on the chart above
CHOSEN_K = 50  # <-- edit this

chosen_row = combined[combined["k"] == CHOSEN_K].iloc[0]
print(f"Chosen K = {CHOSEN_K}")
print(f"  Recall@K:           {chosen_row['recall_at_k']:.2%}")
print(f"  End-to-end accuracy: {chosen_row['accuracy']:.2%}")

8. Volledige classificatie uitvoeren met gekozen K

Pas de geselecteerde K toe op de hele documenttabel.

spark.sql(f"""
CREATE TABLE IF NOT EXISTS top_labels_per_doc AS
SELECT
  d.id,
  {top_k_labels_json('l')} AS labels
FROM doc_embeddings d
INNER JOIN label_embeddings l
APPROX NEAREST {CHOSEN_K} BY SIMILARITY vector_cosine_similarity(d.embedding, l.embedding)
GROUP BY d.id
""")

print(f"Built top-{CHOSEN_K} label sets for all documents")
result_df = spark.sql(f"""
SELECT
  c.{DOCS_TEXT_COL},
  cast(ai_classify(c.{DOCS_TEXT_COL}, t.labels, map('version', '2.0')) as string) AS classification
FROM {DOCS_TABLE} c
JOIN top_labels_per_doc t ON t.id = {doc_id_expr.replace(DOCS_TEXT_COL, f'c.{DOCS_TEXT_COL}')}
""")

result_df.display()

# Optionally save results
# result_df.write.mode("overwrite").saveAsTable("my_catalog.my_schema.classification_results")