Réduire la variance avec plusieurs générations

Lorsque votre tâche est non déterministe (par exemple, temperature > 0), une seule génération par entrée peut ne pas vous donner une vision fiable. Définissez num_generations=N pour exécuter la tâche N fois par enregistrement. Le SDK agregge automatiquement la moyenne et l'écart-type par enregistrement.

Utilisation

Utilisation

import asyncio
import os

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

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

dataset = [
    {"prompt": "What is the capital of France?", "expected": "Paris"},
    {"prompt": "What is 2+2?", "expected": "4"},
]

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

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

async def main():
    run = await client.evaluation.run(
        evaluation=Evaluation(name="Variance check"),
        dataset=dataset,
        task=task,
        system=System(
            name="small-t09",
            params={"model": "mistral-small-latest", "temperature": 0.9},
        ),
        evaluators=[
            Evaluator(
                name="accuracy",
                description="1 if expected answer is in the output.",
                scorer=scorer,
                goal=Goal.gte(0.5),
            )
        ],
        num_generations=3,
    )
    run.show(level="generations")

asyncio.run(main())
Fonctionnement

Fonctionnement

Avec num_generations=3, le SDK :

  1. exécute la tâche 3 fois par enregistrement d'entrée.
  2. note chaque génération de manière indépendante.
  3. calcule les statistiques par enregistrement (moyenne, écart-type) sur les 3 générations.
  4. calcule les statistiques globales de l'exécution sur tous les enregistrements.

Utilisez run.show(level="generations") pour voir le détail par génération.

Quand l'utiliser

Quand l'utiliser

  • Comparaison de températures : exécutez le même prompt à différentes températures pour mesurer la stabilité.
  • Tâches bruitées : tâches produisant des sorties significativement différentes à chaque appel.
  • Intervalles de confiance : obtenez une idée de la fiabilité de vos scores avant de tirer des conclusions.
Exemple : comparaison de températures

Exemple : comparaison de températures

import asyncio
import os

from mistralai.observability import Evaluation, Evaluator, Goal, Mistral, System, ScorerContext, TaskContext

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

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

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

async def main():
    for temp in [0.0, 0.3, 0.7, 1.0]:
        run = await client.evaluation.run(
            evaluation=Evaluation(name="Temperature Impact"),
            dataset=dataset,
            task=task,
            system=System(
                name=f"small-t{temp}",
                params={"model": "mistral-small-latest", "temperature": temp},
            ),
            evaluators=[Evaluator(name="accuracy", description="1 if expected answer is in the output.", scorer=scorer, goal=Goal.gte(0.5))],
            num_generations=5,
            tags=[f"temperature:{temp}"],
        )

asyncio.run(main())

Dans Studio, chaque exécution a son paramètre System enregistré. Vous pouvez comparer les distributions de scores entre les températures d'un coup d'œil. Voir Paramètres système pour plus de détails.