Something went wrong. Try again.
This repository has no description
Something went wrong. Try again.
4.1 kB · 105 lines
Python
at main
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106"""Score live Jev against the labelled lines in eval/cases.tsv.
python -m sidecar.evaluate # every case python -m sidecar.evaluate --misses # only print the ones it got wrong
This spends real (tiny) money: one Jev request per case."""import sysimport timefrom concurrent.futures import ThreadPoolExecutor
from .classify import Utterance, classifyfrom .config import REPO_ROOT, load_settingsfrom .da_schema import ADDRESSEE_SPRITE, REFER_TO, SYSTEM_DOESNT_UNDERSTANDfrom .jev_client import JevClient, JevError
CASES_PATH = REPO_ROOT / "eval" / "cases.tsv"SPRITE_NAME = {sprite: name for name, sprite in ADDRESSEE_SPRITE.items()}
def load_cases(path=CASES_PATH): cases = [] for line in path.read_text(encoding="utf-8").splitlines(): if not line.strip() or line.startswith("#"): continue fields = line.split("\t") + ["-"] * 7 text, act, addressee, reference, param, said, contexts = fields[:7] cases.append({"text": text, "act": act, "addressee": addressee, "reference": reference, "param": param, "said": () if said == "-" else (said,), "contexts": () if contexts == "-" else tuple(contexts.split(","))}) return cases
def observed(result): """Flatten a Classification into the same four fields a case labels.""" # "none" (not "-") for act and reference, so a case can demand their absence: # a spurious reference is a real extra act sent to the game. got = {"act": "none", "addressee": "-", "reference": "none", "param": "-"} for act in result.acts: if act.da_id == SYSTEM_DOESNT_UNDERSTAND: continue if act.char_id in SPRITE_NAME: got["addressee"] = SPRITE_NAME[act.char_id] if act.da_id == REFER_TO: got["reference"] = act.param_names[0] else: got["act"] = act.name named = [p for p in act.param_names if p] if named: got["param"] = "/".join(named) return got
def misses(case, got): wrong = [] for key in ("act", "addressee", "reference"): if case[key] != "-" and case[key] != got[key]: wrong.append("%s: wanted %s, got %s" % (key, case[key], got[key])) if case["param"] != "-" and case["param"] not in got["param"].split("/"): wrong.append("param: wanted %s, got %s" % (case["param"], got["param"])) return wrong
def run_case(case, client, settings): started = time.monotonic() try: utterance = Utterance(case["text"], contexts=case["contexts"], recent_dialogue=case["said"]) result = classify(utterance, client, settings) except JevError as error: return None, ["error: %s" % error], 0.0 got = observed(result) return got, misses(case, got), (time.monotonic() - started) * 1000
def main(): only_misses = "--misses" in sys.argv settings = load_settings() client = JevClient.from_settings(settings) cases = load_cases() with ThreadPoolExecutor(max_workers=6) as pool: outcomes = list(pool.map(lambda c: run_case(c, client, settings), cases))
passed, latencies = 0, [] for case, (got, wrong, ms) in zip(cases, outcomes): if not wrong: passed += 1 if got is not None: latencies.append(ms) if wrong or not only_misses: summary = "" if got is None else "%s to=%s ref=%s param=%s" % ( got["act"], got["addressee"], got["reference"], got["param"]) print("%s %-46r %s" % ("ok " if not wrong else "MISS", case["text"], summary)) for problem in wrong: print(" " + problem) latencies.sort() print("\n%d/%d cases fully correct" % (passed, len(cases))) if latencies: print("latency ms: median %d, p90 %d, max %d" % ( latencies[len(latencies) // 2], latencies[int(len(latencies) * 0.9)], latencies[-1])) sys.exit(0 if passed == len(cases) else 1)
if __name__ == "__main__": main()