devils-agent / baim /policy.py
devildasdf's picture
Optimize CPU feature encoding and add bounded batch inference with measured parity
c7d6933 verified
Raw History Blame Contribute Delete
2.82 kB
"""Learned three-action baseline with explicit uncertainty abstention."""
import json
from pathlib import Path
import re
import torch
from .actions import Action, Decision, Kind
from .features import encode, KINDS
from .model import PointerPolicy
class LearnedPolicy:
def __init__(self, checkpoint, quantized=False, confidence_threshold=.85):
root = Path(checkpoint)
self.model = PointerPolicy.from_pretrained(root).eval()
self.vocab = json.loads((root/'vocab.json').read_text(encoding='utf-8'))
self.temperatures = json.loads((root/'calibration.json').read_text())['temperatures']
self.threshold = confidence_threshold
if quantized:
self.model = torch.ao.quantization.quantize_dynamic(self.model,{torch.nn.Linear},dtype=torch.qint8)
@torch.inference_mode()
def predict(self, goal, state, ticket):
return self.predict_batch([(goal, state, ticket)])[0]
@torch.inference_mode()
def predict_batch(self, requests):
"""Bounded CPU microbatch. Every decision retains its own authority ticket."""
if len(requests) > 64:
raise ValueError('A microbatch may contain at most 64 requests')
if not requests:
return []
rows = [dict(goal=goal,elements=[dict(role=e.role,name=e.name,visible=e.visible,
enabled=e.enabled,sensitive=e.sensitive) for e in state.elements])
for goal,state,_ in requests]
inputs,_,_,maps = encode(rows,self.vocab)
a,t = self.model(*inputs)
ap = (a/self.temperatures[0]).softmax(-1)
tp = (t/self.temperatures[1]).softmax(-1)
return [self._decision(goal,state,ticket,mapping,ap[i],tp[i])
for i,((goal,state,ticket),mapping) in enumerate(zip(requests,maps))]
def _decision(self, goal, state, ticket, mapping, ap, tp):
if not mapping:
return Decision(ticket,Action(Kind.ASK_USER,('No eligible DOM target; a fallback is required.',)))
ai,ti = int(ap.argmax()),int(tp.argmax())
ac,tc = float(ap[ai]),float(tp[ti])
if min(ac,tc) < self.threshold:
return Decision(ticket,Action(Kind.ASK_USER,('The policy is uncertain about this action or target.',)),ac,tc)
kind = Kind(KINDS[ai])
element = state.elements[mapping[ti]]
args = (element.ref,)
if kind in {Kind.TYPE,Kind.SELECT}:
# Literal copying, not inferred website/workflow logic. General span prediction is pending.
literals = re.findall(r'"([^"\n]*)"',goal)
if len(literals) != 1:
return Decision(ticket,Action(Kind.ASK_USER,('This baseline needs one quoted value to copy.',)),ac,tc)
args += (literals[0],)
return Decision(ticket,Action(kind,args),ac,tc)