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 :
- exécute la tâche 3 fois par enregistrement d'entrée.
- note chaque génération de manière indépendante.
- calcule les statistiques par enregistrement (moyenne, écart-type) sur les 3 générations.
- 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.