Download baim/policy.py from devildasdf/devils-agent: direct link, hf CLI and curl.
- Browser
- Download file 2.82 kB
-
https://huggingface.co/devildasdf/devils-agent/resolve/main/baim/policy.py
- Command line
-
hf download hf://devildasdf/devils-agent/baim/policy.py
-
curl -L -o policy.py https://huggingface.co/devildasdf/devils-agent/resolve/main/baim/policy.py
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) | |
| def predict(self, goal, state, ticket): | |
| return self.predict_batch([(goal, state, ticket)])[0] | |
| 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) | |