Remarque
L’accès à cette page nécessite une autorisation. Vous pouvez essayer de vous connecter ou de modifier des répertoires.
L’accès à cette page nécessite une autorisation. Vous pouvez essayer de modifier des répertoires.
ai_classify accepte jusqu’à 500 étiquettes par appel. Pour les taxonomies volumineuses, préfiltrez les étiquettes par document à l'aide de la similarité des embeddings, puis appelez ai_classify sur la liste restreinte des candidats top-K. Ce tutoriel vous montre comment trouver le K optimal , le plus petit nombre de candidats qui conservent la précision.
Note
NEAREST BY nécessite Databricks Runtime 18 ou version ultérieure, ou Serverless. Databricks Runtime 18 est plus récent que Databricks Runtime 18.0, 18.1 et 18.2.
Avant de commencer
- Un espace de travail compatible avec Unity Catalog avec accès à
ai_classify(voir la disponibilité). - Databricks Runtime 18+ ou Serverless (obligatoire pour
NEAREST BY). - Table delta de documents à classifier.
- Table Delta d’étiquettes avec une colonne clé et une colonne de description facultative.
- Un petit ensemble d'évaluation avec des étiquettes de référence, ou la possibilité d'en créer un (option B à l'étape 4).
0. Configuration
Définissez vos noms de table, noms de colonnes et modèle d’incorporation. La fonction d'assistance top_k_labels_json construit l'expression JSON transmise à ai_classify pour les candidats top-K de chaque 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. Incorporer les étiquettes
Exécutez cette opération une seule fois. Réexécutez uniquement lorsque la taxonomie change.
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. Incorporer les documents
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. Récupérer les étiquettes top-K à l'aide de NEAREST BY
NEAREST BY effectue directement une jointure approximative par plus proche voisin, sans nécessiter de table intermédiaire N×M. Pour chaque document, il retourne les étiquettes K les plus similaires en une seule passe.
# 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. Préparer un ensemble d'évaluation de référence
L’ajustement de K nécessite un petit ensemble de documents dont les libellés corrects sont connus. Si vous disposez d’une table d’évaluation existante, définissez EVAL_TABLE la cellule suivante et ignorez la cellule d’échantillonnage. Si ce n’est pas le cas, la deuxième cellule sélectionne un échantillon de documents que vous pouvez étiqueter manuellement puis réimporter.
# 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. Mesurer Recall@K
Recall@K vérifie si l'étiquette de référence figure parmi les candidats top-K issus des embeddings. Il s’agit d’une métrique de récupération uniquement : elle n’appelle ai_classify pas et s’exécute instantanément.
Si le rappel est faible pour une valeur donnée de K, ai_classify ne peut tout simplement pas renvoyer la bonne réponse, car l’étiquette correcte a été exclue de l’ensemble de candidats avant même que la classification ne soit exécutée.
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. Mesurez la précision de bout en bout
Pour chaque valeur de K, construisez l'ensemble des étiquettes top-K par document d'évaluation, exécutez ai_classify, puis comparez le résultat avec les données de référence.
Cette étape fait appel à ai_classify et coûte plus cher que le contrôle de rappel. Commencez par les valeurs K où le rappel est déjà raisonnable.
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. Comparer les résultats et sélectionner K
Le graphique ci-dessous montre Recall@K et la précision de bout en bout côte à côte. Choisissez le plus petit K où la précision cesse de s’améliorer , plus grande K signifie une classification plus lente sans gain de qualité.
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. Exécuter la classification complète avec le K choisi
Appliquez le K sélectionné à l’ensemble de votre tableau de documents.
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")