Classer les documents avec plus de 500 étiquettes
ai_classify accepte jusqu'à 500 étiquettes par appel. Pour les taxonomies plus vastes, pré-filtrez les étiquettes par document en utilisant la similarité d'intégration, 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 préserve la précision.
NEAREST BY nécessite Databricks Runtime 18 ou une version supérieure, ou Serverless. Databricks Runtime 18 est plus récent que Databricks Runtime 18,0, 18,1 et 18,2.
Avant de commencer
- Un Workspace avec Unity Catalog activé, avec accès à
ai_classify(voir la disponibilité). - Databricks Runtime 18+ ou Serverless (requis pour
NEAREST BY). - Une table Delta de documents à classifier.
- Une table Delta d'étiquettes avec une colonne clé et une colonne de description facultative.
- Un petit ensemble d'évaluation avec des étiquettes de vérité terrain, ou la possibilité d'en créer un (Option B à l'étape 4).
0. Configuration
Définissez vos noms de table, noms de colonne et modèle d'intégration. La fonction d'assistance top_k_labels_json construit l'expression JSON transmise à ai_classify pour les K meilleurs candidats 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. Intégrer les étiquettes
Exécutez-le une fois. Relancez 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. Intégrez 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érez les k étiquettes principales à l'aide de NEAREST BY
NEAREST BY effectue directement une jointure approximative des plus proches voisins — aucune table N×M intermédiaire n'est nécessaire. Pour chaque document, il renvoie les K étiquettes les plus similaires en un seul passage.
# 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éparez un ensemble d'évaluation de vérité terrain
Le K-tuning nécessite un petit ensemble de documents avec des étiquettes correctes connues. Si vous disposez d'une table d'évaluation existante, définissez EVAL_TABLE dans la cellule suivante et passez la cellule d'échantillonnage. Sinon, la deuxième cellule échantillonne les documents que vous pouvez étiqueter manuellement et 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")
5e. Mesurer le rappel@K
Recall@K vérifie si l'étiquette de vérité terrain apparaît dans les meilleurs candidats d'intégration K. Il s'agit d'une métrique de récupération seule — elle n'appelle pas ai_classify et s'exécute instantanément.
Si le rappel est faible pour un K donné, ai_classify ne peut pas renvoyer la bonne réponse car l'étiquette correcte a été exclue de l'ensemble des candidats avant même que la classification ne s'exécute.
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. Mesurer la précision de bout en bout
Pour chaque K, créez l'ensemble d'étiquettes top-K par document d'évaluation, exécutez ai_classify et comparez-le à la vérité terrain.
Cette étape appelle ai_classify et coûte plus cher que la vérification de rappel. Start avec 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 choisir K
Le graphique ci-dessous montre la précision Recall@K et de bout en bout côte à côte. Choisissez le K le plus petit où la précision cesse de s'améliorer — un K plus grand 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
Appliquer le K sélectionné à l'ensemble de votre table 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")