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.
208 lines
7.5 KiB
Python
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()
|