Relancer les enregistrements en échec

Une évaluation peut échouer partiellement : une erreur API transitoire, une limite de débit ou un bug dans votre scorer. Plutôt que de relancer toute l'évaluation depuis le début, le SDK permet de relancer uniquement les enregistrements en échec et de mettre à jour les résultats sur place.

Pourquoi c'est important

Pourquoi c'est important

Une exécution d'évaluation typique peut impliquer des centaines ou des milliers d'appels à un LLM. Si cinq enregistrements sur 500 échouent, tout relancer fait perdre du temps et de l'argent. retry_failed_records() identifie les enregistrements qui ont échoué, au niveau de la génération ou du scoring, relance uniquement ceux-ci et met à jour l'exécution d'origine afin de conserver vos résultats au même endroit dans Studio.

workflow de base

workflow de base

import asyncio
import os

from mistralai.evaluations import (
    Evaluation, Evaluator, Mistral, ScorerContext, System, TaskContext,
)

client = Mistral(api_key=os.environ["MISTRAL_API_KEY"])

dataset = [
    {"prompt": "Say hello", "expected": "hello"},
    {"prompt": "Say goodbye", "expected": "goodbye"},
]

async def task(ctx: TaskContext):
    response = await client.chat.complete_async(
        model=str(ctx.system.params["model"]),
        messages=[{"role": "user", "content": ctx.input_record["prompt"]}],
    )
    return str(response.choices[0].message.content)

def scorer(ctx: ScorerContext):
    return 1 if ctx.input_record["expected"] in str(ctx.output).lower() else 0

async def main():
    system = System(name="mistral-small", params={"model": "mistral-small-latest"})

    # Step 1: run the evaluation
    run = await client.evaluation.run(
        evaluation=Evaluation(name="My Eval"),
        system=system,
        dataset=dataset,
        task=task,
        evaluators=[Evaluator(name="accuracy", scorer=scorer)],
    )

    # Step 2: if some records failed, retry them
    result = await client.evaluation.retry_failed_records(
        run_id=run.run_id,
        dataset=dataset,
        task=task,
        evaluators=[Evaluator(name="accuracy", scorer=scorer)],
        system=system,
    )
    print(f"Retried: {result.retried_count}, Patched: {result.patched_count}")

asyncio.run(main())
Corriger la tâche avant de relancer

Corriger la tâche avant de relancer

Vous n'avez pas à relancer avec la même tâche. Si les échecs sont dus à un bug dans votre code, corrigez-le et transmettez la version corrigée :

# Fix the task, retry only the failed records
async def fixed_task(ctx: TaskContext):
    response = await client.chat.complete_async(
        model=str(ctx.system.params["model"]),
        messages=[{"role": "user", "content": ctx.input_record["prompt"]}],
    )
    return str(response.choices[0].message.content)

result = await client.evaluation.retry_failed_records(
    run_id=run.run_id,
    dataset=dataset,
    task=fixed_task,
    evaluators=[Evaluator(name="accuracy", scorer=scorer)],
    system=system,
)
Ce qui se passe

Ce qui se passe

  1. Le SDK récupère l'exécution existante et identifie les enregistrements avec le statut "error", correspondant à un échec de génération ou de scoring.
  2. Seuls ces enregistrements sont retraités avec la tâche et les évaluateurs fournis.
  3. Les résultats réussis sont ajoutés à l'exécution d'origine.
  4. Si l'exécution d'origine comportait des run_evaluators, leurs scores sont recalculés avec les données mises à jour.

L'exécution d'origine dans Studio est mise à jour sur place. Pas d'exécutions en double, pas de nettoyage manuel.

Référence API

Référence API

client.evaluation.retry_failed_records(...) prend les paramètres suivants :

ParamètreTypeDescription
run_idstrObligatoire. ID de l'exécution contenant les enregistrements en échec.
datasetSequence[Mapping[str, Any]]Obligatoire. Même ensemble de données que celui utilisé dans l'exécution d'origine.
taskTaskFunctionObligatoire. Fonction de tâche, qui peut être une version corrigée.
evaluatorslist[Evaluator]Obligatoire. Évaluateurs à utiliser pour recalculer les scores.
run_evaluatorslist[RunEvaluator]Évaluateurs au niveau de l'exécution à recalculer après la mise à jour.
num_generationsintNombre de générations par enregistrement (par défaut : 1).
systemSystemParamètres système à injecter dans la tâche.
max_concurrencyintNombre maximal de tâches simultanées (par défaut : 10).
upload_batch_sizeintTaille de batch pour télécharger les résultats mis à jour (par défaut : 10).

Elle renvoie un RetryFailedRecordsResult :

ChampTypeDescription
run_idstrID de l'exécution relancée.
retried_countintNombre d'enregistrements relancés.
patched_countintNombre d'enregistrements mis à jour avec succès.
run_scores_recomputedboolIndique si les scores au niveau de l'exécution ont été recalculés.