| """Read-only offline grounding diagnostic; never executes downloaded HTML.""" |
| import argparse |
| from collections import Counter |
| from dataclasses import asdict |
| import hashlib |
| from html.parser import HTMLParser |
| import json |
| from pathlib import Path |
| import random |
| import torch |
| from .features import encode |
| from .policy import LearnedPolicy |
| from .state import Element |
| from .train import logits |
|
|
|
|
| class TextIndex(HTMLParser): |
| def __init__(self): |
| super().__init__(convert_charrefs=True) |
| self.stack = [] |
| self.nodes = {} |
|
|
| def handle_starttag(self,tag,attrs): |
| attrs = dict(attrs) |
| node = dict(tag=tag,attrs=attrs,text=[]) |
| if attrs.get('backend_node_id'): |
| self.nodes[attrs['backend_node_id']] = node |
| if tag not in {'area','base','br','col','embed','hr','img','input','link','meta','param','source','track','wbr'}: |
| self.stack.append(node) |
|
|
| def handle_endtag(self,tag): |
| for index in range(len(self.stack)-1,-1,-1): |
| if self.stack[index]['tag'] == tag: |
| del self.stack[index:] |
| break |
|
|
| def handle_data(self,data): |
| if any(node['tag'] in {'script','style'} for node in self.stack): |
| return |
| for node in self.stack: |
| if sum(map(len,node['text'])) < 512: |
| node['text'].append(data[:512]) |
|
|
|
|
| def normalize_task(task): |
| for step_index, step in enumerate(task['actions']): |
| positives = step['pos_candidates'] |
| if not positives: |
| yield None |
| continue |
| parser = TextIndex() |
| parser.feed(step['cleaned_html']) |
| entries = [(candidate,True) for candidate in positives] + [(candidate,False) for candidate in step['neg_candidates']] |
| rng = random.Random(int(hashlib.sha256(step['action_uid'].encode()).hexdigest(),16)) |
| rng.shuffle(entries) |
| elements, targets = [], [] |
| seen = set() |
| for candidate,positive in entries: |
| ident = str(candidate['backend_node_id']) |
| if ident in seen: |
| continue |
| seen.add(ident) |
| attrs = json.loads(candidate['attributes']) |
| node = parser.nodes.get(ident,{}) |
| html_attrs = node.get('attrs',{}) |
| attrs = {**html_attrs,**attrs} |
| tag = candidate['tag'].lower() |
| role = attrs.get('role') or {'button':'button','a':'link','input':'textbox', |
| 'textarea':'textbox','select':'combobox'}.get(tag,'generic') |
| if tag=='input': |
| role = {'checkbox':'checkbox','radio':'radio','submit':'button','button':'button'}.get(attrs.get('type'),role) |
| if role=='generic' and str(attrs.get('is_clickable','')).lower() in {'true','1'}: |
| role='button' |
| sensitive = attrs.get('type')=='password' or 'cc-' in attrs.get('autocomplete','') |
| name = attrs.get('aria-label') or attrs.get('placeholder') or attrs.get('title') or ' '.join(node.get('text',[])) |
| name = ' '.join(str(name).split())[:512] |
| index = len(elements) |
| elements.append(asdict(Element(f'e{index}',role,'[REDACTED]' if sensitive else name, |
| str(index),enabled='disabled' not in attrs,sensitive=sensitive))) |
| if positive: |
| targets.append(index) |
| if len(targets)!=1: |
| yield None |
| continue |
| operation = {'CLICK':'C','TYPE':'T','SELECT':'O'}.get(step['operation']['op']) |
| if operation is None: |
| yield None |
| continue |
| yield dict(goal=task['confirmed_task'],elements=elements,action=operation,target=targets[0], |
| history=task.get('action_reprs',[])[:step_index]) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument('--source',required=True) |
| parser.add_argument('--checkpoint',default='models/v000-mean') |
| parser.add_argument('--output',default='reports/mind2web-smoke-v000.json') |
| args = parser.parse_args() |
| torch.set_num_threads(2) |
| tasks = json.loads(Path(args.source).read_text(encoding='utf-8')) |
| all_rows = [row for task in tasks for row in normalize_task(task)] |
| rows = [row for row in all_rows if row is not None] |
| policy = LearnedPolicy(args.checkpoint) |
| inputs,actions,targets,_ = encode(rows,policy.vocab) |
| a,t = logits(policy.model,inputs) |
| report = dict(source='osunlp/Mind2Web',revision='6314166657eec4aa0e22c00f8d801e609ce8e80f', |
| file=Path(args.source).name,source_sha256=hashlib.sha256(Path(args.source).read_bytes()).hexdigest(), |
| license='CC-BY-4.0',attribution='Deng et al., Mind2Web: Towards a Generalist Agent for the Web, 2023, arXiv:2306.06070', |
| checkpoint=args.checkpoint,tasks=len(tasks),websites=len({task['website'] for task in tasks}), |
| total_steps=len(all_rows),scorable_steps=len(rows),unscorable_steps=len(all_rows)-len(rows), |
| operations=dict(Counter(row['action'] for row in rows)), |
| candidate_recall=float((targets>=0).float().mean()), |
| action_accuracy=float((a.argmax(-1)==actions).float().mean()), |
| target_accuracy=float((t.argmax(-1)==targets).float().mean()), |
| joint_accuracy=float(((a.argmax(-1)==actions)&(t.argmax(-1)==targets)).float().mean()), |
| browser_execution=False,training_on_source=False, |
| limitations='Small training-shard diagnostic, NOT official held-out benchmark. Approximate HTML names/roles; no history; 24-token goal truncation. No raw page or goal text persisted in report.') |
| 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(report,indent=2)) |
|
|
|
|
| if __name__=='__main__': |
| main() |
|
|