Download baim/evaluate_policy.py from devildasdf/devils-agent: direct link, hf CLI and curl.
- Browser
- Download file 4.19 kB
-
https://huggingface.co/devildasdf/devils-agent/resolve/main/baim/evaluate_policy.py
- Command line
-
hf download hf://devildasdf/devils-agent/baim/evaluate_policy.py
-
curl -L -o evaluate_policy.py https://huggingface.co/devildasdf/devils-agent/resolve/main/baim/evaluate_policy.py
4.19 kB
| """Single-step browser benchmark with independent fixture outcome checks.""" | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import platform | |
| import statistics | |
| from time import perf_counter | |
| import torch | |
| from playwright.sync_api import sync_playwright | |
| from .authority import Authority | |
| from .browser import Browser | |
| from .policy import LearnedPolicy | |
| from .runtime import Runtime | |
| from .synthetic import load, render | |
| from .baseline import LexicalRoleBaseline | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument('--checkpoint',default='models/v000-mean') | |
| parser.add_argument('--data',default='datasets/synthetic-v1/test.jsonl') | |
| parser.add_argument('--output',default='reports/policy-browser-v000.json') | |
| parser.add_argument('--limit',type=int,default=120) | |
| parser.add_argument('--quantized',action='store_true') | |
| parser.add_argument('--baseline',action='store_true') | |
| args = parser.parse_args() | |
| torch.set_num_threads(2) | |
| torch.set_num_interop_threads(1) | |
| start = perf_counter() | |
| policy = LexicalRoleBaseline() if args.baseline else LearnedPolicy(args.checkpoint,quantized=args.quantized) | |
| load_ms = (perf_counter()-start)*1000 | |
| results = [] | |
| with sync_playwright() as pw: | |
| with pw.chromium.launch(headless=True) as chromium: | |
| context = chromium.new_context(viewport={'width':1100,'height':900}) | |
| browser = Browser(context) | |
| authority = Authority() | |
| for sample in load(args.data)[:args.limit]: | |
| browser.page.set_content(render(sample)) | |
| authority.replace(sample['goal']) | |
| start = perf_counter() | |
| state = browser.observe() | |
| observed_ms = (perf_counter()-start)*1000 | |
| start = perf_counter() | |
| decision = policy.predict(sample['goal'],state,authority.ticket(state)) | |
| policy_ms = (perf_counter()-start)*1000 | |
| runtime = Runtime(browser,authority,lambda *_: True) | |
| outcome = runtime.execute(decision) | |
| expected = sample['target'] | |
| if sample['action'] == 'C': | |
| success = browser.page.evaluate('window.fixtureResult') == expected | |
| else: | |
| # data-index is used only by the evaluator; never shown to the policy. | |
| value = browser.page.locator(f'[data-index="{expected}"]').input_value() | |
| success = value == sample['argument'] | |
| results.append(dict(sample_id=sample['sample_id'],template=sample['template'], | |
| success=success,abstained=decision.action.kind.value=='A', | |
| expected_action=sample['action'],predicted_action=decision.action.kind.value, | |
| action_confidence=decision.action_confidence,target_confidence=decision.target_confidence, | |
| observation_ms=observed_ms,policy_ms=policy_ms,execution_ms=outcome.wall_ms, | |
| status=outcome.status,code=outcome.code)) | |
| if len(results)%20==0: | |
| print(f'{len(results)} fixtures complete',flush=True) | |
| context.close() | |
| def median(key): | |
| return statistics.median(row[key] for row in results) | |
| report = dict(checkpoint='lexical-role-baseline' if args.baseline else args.checkpoint,quantized=args.quantized,platform=platform.platform(), | |
| torch_threads=2,model_load_ms=load_ms,samples=len(results), | |
| success_rate=sum(r['success'] for r in results)/len(results), | |
| abstention_rate=sum(r['abstained'] for r in results)/len(results), | |
| median_observation_ms=median('observation_ms'),median_policy_ms=median('policy_ms'), | |
| median_execution_ms=median('execution_ms'), | |
| scope='single-step generated browser fixtures, held-out layout templates; not arbitrary websites', | |
| target_vps_validated=False,results=results) | |
| Path(args.output).parent.mkdir(parents=True,exist_ok=True) | |
| Path(args.output).write_text(json.dumps(report,indent=2),encoding='utf-8') | |
| print(json.dumps({k:v for k,v in report.items() if k!='results'},indent=2)) | |
| if __name__ == '__main__': | |
| main() | |