File size: 4,825 Bytes
f1b0176 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 | import argparse
import json
import sys
from pathlib import Path
PROJECT_ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(PROJECT_ROOT))
from tests.fixtures.parser_acceptance import PARSER_ACCEPTANCE_FIXTURES
from training_coach.parser import PARSER_MODEL
from training_coach.parser_ollama import OLLAMA_MODEL, generate_parser_response_ollama
from training_coach.parser_runtime import generate_parser_response
from training_coach.parser import parse_model_response
SEMANTIC_TEXT_FIELDS = {"notes", "follow_up_question", "evidence"}
SEMANTIC_TEXT_LIST_FIELDS = {"follow_up_questions"}
def _tokenize_text(value):
return {
token.strip(".,?!;:()[]{}'\"").lower()
for token in value.split()
if token.strip(".,?!;:()[]{}'\"")
}
def _compare_text_semantics(expected, actual, path):
if not isinstance(actual, str):
return [f"{path}: expected text, got {type(actual).__name__}"]
expected_tokens = _tokenize_text(expected)
actual_tokens = _tokenize_text(actual)
if expected_tokens and expected_tokens.isdisjoint(actual_tokens):
return [f"{path}: expected related text {expected!r}, got {actual!r}"]
return []
def _compare_subset(expected, actual, path=""):
failures = []
if isinstance(expected, dict):
if not isinstance(actual, dict):
return [f"{path or '<root>'}: expected object, got {type(actual).__name__}"]
for key, expected_value in expected.items():
if key not in actual:
failures.append(f"{path}.{key}: missing key")
continue
failures.extend(
_compare_subset(expected_value, actual[key], f"{path}.{key}".strip("."))
)
return failures
if isinstance(expected, list):
if not isinstance(actual, list):
return [f"{path}: expected list, got {type(actual).__name__}"]
if len(actual) < len(expected):
failures.append(f"{path}: expected at least {len(expected)} items, got {len(actual)}")
return failures
field_name = path.split(".")[-1]
if field_name in SEMANTIC_TEXT_LIST_FIELDS:
for expected_item in expected:
if not any(
not _compare_text_semantics(expected_item, actual_item, path)
for actual_item in actual
):
failures.append(f"{path}: expected related text {expected_item!r}")
return failures
for index, expected_item in enumerate(expected):
failures.extend(_compare_subset(expected_item, actual[index], f"{path}[{index}]"))
return failures
field_name = path.split(".")[-1]
if isinstance(expected, str) and field_name in SEMANTIC_TEXT_FIELDS:
return _compare_text_semantics(expected, actual, path)
if expected != actual:
return [f"{path}: expected {expected!r}, got {actual!r}"]
return []
def evaluate_fixture(fixture, model_name, backend):
if backend == "ollama":
response_text = generate_parser_response_ollama(
fixture["raw_text"],
model_name=model_name,
)
else:
response_text = generate_parser_response(fixture["raw_text"], model_name=model_name)
try:
parsed = parse_model_response(response_text)
except ValueError as error:
return response_text, None, [str(error)]
actual = parsed.model_dump(mode="json")
failures = _compare_subset(fixture["expected"], actual)
return response_text, actual, failures
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--backend", choices=["ollama", "transformers"], default="ollama")
parser.add_argument("--model", default=None)
args = parser.parse_args()
model_name = args.model or (OLLAMA_MODEL if args.backend == "ollama" else PARSER_MODEL)
failed = 0
print(f"Evaluating parser model: {model_name}")
print(f"Backend: {args.backend}")
for fixture in PARSER_ACCEPTANCE_FIXTURES:
print(f"\n=== {fixture['name']} ===")
response_text, actual, failures = evaluate_fixture(
fixture,
model_name,
args.backend,
)
if failures:
failed += 1
print("FAIL")
for failure in failures:
print(f"- {failure}")
print("Model response:")
print(response_text)
if actual is not None:
print("Validated JSON:")
print(json.dumps(actual, indent=2))
else:
print("PASS")
total = len(PARSER_ACCEPTANCE_FIXTURES)
passed = total - failed
print(f"\nParser acceptance: {passed}/{total} passed")
raise SystemExit(1 if failed else 0)
if __name__ == "__main__":
main()
|