Files
qwen3-6-lora/scripts/32_gate2_toolcalls.py
T
aleleba 2d2c45f0fe Fase 5: gate2 - subir max_tokens a 2048 y guardar texto completo de respuestas
Mismo defecto que se encontro y corrigio en gate3 (commit 77e6804): con
--reasoning-parser activo, max_tokens=1024 podia dejar cortar la respuesta a
mitad de razonamiento antes de emitir el tool_call. Desglose por tipo de
fallo entre corridas: Fase 4 BF16 tuvo CERO casos "no_tool_call" (0/200);
ambas corridas NVFP4 (solo-propia y mezclada) tuvieron 4/200 -- la firma
exacta de un modelo cortado a mitad de razonamiento, no de una regresion de
calidad real.

Mejoras permanentes al test:
1. max_tokens: 1024 -> 2048 (configurable via GATE2_MAX_TOKENS). timeout de
   request subido de 120s a 240s.
2. Se guarda content+reasoning+finish_reason completos en cada fila del
   JSON de resultados (antes solo tool_calls/valid/errors), para poder
   auditar con criterio humano cualquier caso que falle o quede sin
   tool_call, sin tener que re-correr el test.
2026-07-30 06:50:27 +00:00

208 lines
7.5 KiB
Python

"""Fase 4 -- Puerta 2: validez de tool-calls contra el parser real de vLLM.
Corre LOCALMENTE (no necesita GPU) contra el endpoint HTTP del contenedor de eval propio
(vllm-eval, docker-compose.eval.yml, puerto 8001 por defecto) ya levantado y respondiendo
en /v1/models.
Para cada prompt de data/holdout_prompts.jsonl (~200, generados por
scripts/31_build_holdout_prompts.py, sin overlap con train/eval): envia una sola llamada a
/v1/chat/completions con las tools reales del MCP correspondiente y
tool_choice="auto". El parseo de tool_calls (`--tool-call-parser=qwen3_coder`,
configurado en docker-compose.eval.yml) lo hace vLLM en el servidor -- este script solo
valida la RESPUESTA ya parseada (nunca re-implementa el parser con una regex propia):
- Si el modelo decide llamar una tool: valida que el nombre exista en el schema del MCP,
que los argumentos parseen como JSON valido, y que las propiedades "required" del
schema esten presentes.
- Si el modelo NO llama ninguna tool: se cuenta aparte (no es un error per se, algunos
prompts pueden resolverse sin tool-call, pero se reporta la tasa).
Reporta: % de prompts con tool_call sintacticamente valido (parseado sin excepcion por
vLLM, arguments=JSON valido, nombre y campos requeridos correctos) por MCP y global.
"""
import argparse
import json
import os
import sys
import time
from collections import defaultdict
from pathlib import Path
import requests
REPO_ROOT = Path(__file__).resolve().parent.parent
HOLDOUT_PATH = REPO_ROOT / "data" / "holdout_prompts.jsonl"
RESULTS_PATH = REPO_ROOT / "data" / os.environ.get("GATE2_RESULTS_FILENAME", "gate2_results.json")
BASE_URL = os.environ.get("VLLM_EVAL_URL", "http://localhost:8001")
MODEL_NAME = os.environ.get("VLLM_EVAL_MODEL", "qwen3.6-35b-a3b-mcp-bf16")
# 1024 dejaba cortar la respuesta a mitad de razonamiento en modelos con
# --reasoning-parser activo antes de emitir el tool_call -- ver hallazgo de
# Fase 5 (misma causa que el fix de gate3, max_tokens 512->2048).
GATE2_MAX_TOKENS = int(os.environ.get("GATE2_MAX_TOKENS", "2048"))
def load_holdout():
examples = []
with open(HOLDOUT_PATH, encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
examples.append(json.loads(line))
return examples
def tool_by_name(tools, name):
for tool in tools:
if tool.get("name") == name or tool.get("function", {}).get("name") == name:
return tool
return None
def to_openai_tools(tools):
openai_tools = []
for tool in tools:
if "function" in tool:
openai_tools.append(tool)
else:
openai_tools.append({
"type": "function",
"function": {
"name": tool["name"],
"description": tool.get("description", ""),
"parameters": tool.get("inputSchema") or tool.get("parameters") or {"type": "object", "properties": {}},
},
})
return openai_tools
def validate_tool_call(tool_call, tools):
name = tool_call["function"]["name"]
raw_args = tool_call["function"]["arguments"]
try:
args = json.loads(raw_args)
except json.JSONDecodeError as e:
return False, f"arguments no es JSON valido: {e}"
tool_def = tool_by_name(tools, name)
if tool_def is None:
return False, f"tool_call a nombre inexistente en el schema del MCP: {name}"
schema = tool_def.get("inputSchema") or tool_def.get("parameters") or {}
required = schema.get("required", [])
missing = [r for r in required if r not in args]
if missing:
return False, f"faltan campos requeridos {missing} en la llamada a {name}"
return True, None
def call_vllm(prompt, tools, timeout=240):
payload = {
"model": MODEL_NAME,
"messages": [{"role": "user", "content": prompt}],
"tools": to_openai_tools(tools),
"tool_choice": "auto",
"max_tokens": GATE2_MAX_TOKENS,
"temperature": 0.0,
}
resp = requests.post(f"{BASE_URL}/v1/chat/completions", json=payload, timeout=timeout)
resp.raise_for_status()
return resp.json()
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--limit", type=int, default=None)
args = parser.parse_args()
examples = load_holdout()
if args.limit:
examples = examples[: args.limit]
print(f"[INFO] {len(examples)} prompts held-out, endpoint={BASE_URL}")
results = []
stats = defaultdict(lambda: {"total": 0, "valid_tool_call": 0, "no_tool_call": 0, "invalid": 0})
t0 = time.time()
for i, ex in enumerate(examples):
mcp = ex["mcp"]
stats[mcp]["total"] += 1
stats["__global__"]["total"] += 1
try:
response = call_vllm(ex["prompt"], ex["tools"])
except Exception as e:
results.append({"mcp": mcp, "prompt": ex["prompt"], "error": str(e)})
stats[mcp]["invalid"] += 1
stats["__global__"]["invalid"] += 1
continue
message = response["choices"][0]["message"]
# Se guarda siempre el texto completo (content + reasoning) para poder auditar
# con criterio humano los casos que fallan o quedan sin tool_call -- antes no se
# guardaba nada de esto, lo que hacia imposible diagnosticar truncamiento.
content = message.get("content") or ""
reasoning = message.get("reasoning") or ""
tool_calls = message.get("tool_calls") or []
if not tool_calls:
stats[mcp]["no_tool_call"] += 1
stats["__global__"]["no_tool_call"] += 1
results.append({
"mcp": mcp,
"prompt": ex["prompt"],
"tool_calls": None,
"valid": None,
"content": content,
"reasoning": reasoning,
"finish_reason": response["choices"][0].get("finish_reason"),
})
continue
all_valid = True
errors = []
for tc in tool_calls:
ok, err = validate_tool_call(tc, ex["tools"])
if not ok:
all_valid = False
errors.append(err)
if all_valid:
stats[mcp]["valid_tool_call"] += 1
stats["__global__"]["valid_tool_call"] += 1
else:
stats[mcp]["invalid"] += 1
stats["__global__"]["invalid"] += 1
results.append({
"mcp": mcp,
"prompt": ex["prompt"],
"tool_calls": [tc["function"]["name"] for tc in tool_calls],
"valid": all_valid,
"errors": errors,
"content": content,
"reasoning": reasoning,
"finish_reason": response["choices"][0].get("finish_reason"),
})
if (i + 1) % 20 == 0:
print(f"[INFO] {i + 1}/{len(examples)} prompts procesados")
dt = time.time() - t0
print(f"\n=== Puerta 2 -- validez de tool-calls (parser real de vLLM) ===")
print(f"[INFO] tiempo total: {dt:.1f}s\n")
for mcp in sorted(stats):
s = stats[mcp]
pct_valid = 100 * s["valid_tool_call"] / s["total"] if s["total"] else 0
print(
f" {mcp:20s} total={s['total']:4d} valid={s['valid_tool_call']:4d} "
f"no_tool_call={s['no_tool_call']:4d} invalid={s['invalid']:4d} "
f"pct_valid={pct_valid:.1f}%"
)
with open(RESULTS_PATH, "w", encoding="utf-8") as f:
json.dump({"stats": stats, "results": results}, f, ensure_ascii=False, indent=2)
print(f"\n[INFO] resultados detallados en {RESULTS_PATH}")
if __name__ == "__main__":
main()