"""Compare recall and semantic reranking without executing source HTML. All candidates here come from the dataset candidate pool. Broad eligibility is an offline diagnostic, not permission to click arbitrary generic live DOM nodes. """ import argparse from collections import Counter import json import math from pathlib import Path import statistics from time import perf_counter import torch from transformers import AutoTokenizer, AutoModelForSequenceClassification from .features import candidates, words from .mind2web_smoke import normalize_task def bm25(goal,elements,limit=80): documents = [Counter(words(e['role']+' '+e['name'])) for e in elements] lengths = [sum(d.values()) for d in documents] avg = sum(lengths)/max(len(lengths),1) df = Counter(token for doc in documents for token in doc) query = set(words(goal)) scores = [] for index,doc in enumerate(documents): e = elements[index] if not e.get('visible',True) or not e.get('enabled',True) or e.get('sensitive',False): continue score = 0.0 for token in query: frequency = doc[token] inverse = math.log(1+(len(documents)-df[token]+.5)/(df[token]+.5)) score += inverse*frequency*2.2/(frequency+1.2*(.25+.75*lengths[index]/max(avg,1))) scores.append((score,index)) return [index for _,index in sorted(scores,key=lambda pair:(-pair[0],pair[1]))[:limit]] def main(): parser = argparse.ArgumentParser() parser.add_argument('--source',required=True) parser.add_argument('--model',required=True) parser.add_argument('--output',default='reports/retrieval-audit.json') args = parser.parse_args() torch.set_num_threads(2) torch.set_num_interop_threads(1) tasks = json.loads(Path(args.source).read_text(encoding='utf-8')) rows = [row for task in tasks for row in normalize_task(task) if row is not None] tokenizer = AutoTokenizer.from_pretrained(args.model,local_files_only=True,trust_remote_code=False) model = AutoModelForSequenceClassification.from_pretrained(args.model,local_files_only=True, trust_remote_code=False).eval() counts = Counter() timings = [] for number,row in enumerate(rows): for limit in [10,20,40,80,200]: counts[f'old_recall_at_{limit}'] += row['target'] in candidates(row['goal'],row['elements'],limit) counts[f'bm25_recall_at_{limit}'] += row['target'] in bm25(row['goal'],row['elements'],limit) selected = bm25(row['goal'],row['elements'],80) counts['bm25_top1'] += selected[0] == row['target'] # Query contains only user task and already completed steps. No current label/value. for use_history in [False,True]: query = row['goal'] if use_history: query += '\nPreviously completed: ' + ' ; '.join(row['history'][-4:]) + '\nNext relevant control:' passages = [row['elements'][i]['role']+' '+row['elements'][i]['name'] for i in selected] start = perf_counter() scores=[] with torch.inference_mode(): for offset in range(0,len(passages),16): batch = passages[offset:offset+16] encoded=tokenizer([query]*len(batch),batch,padding=True,truncation=True, max_length=256,return_tensors='pt') scores.extend(model(**encoded).logits.flatten().tolist()) ranked=[selected[i] for i in sorted(range(len(scores)),key=lambda i:-scores[i])] prefix='semantic_history' if use_history else 'semantic_goal' for limit in [1,5,10,20,40]: counts[f'{prefix}_recall_at_{limit}'] += row['target'] in ranked[:limit] timings.append(dict(history=use_history,ms=(perf_counter()-start)*1000)) if (number+1)%5==0: print(f'{number+1}/{len(rows)} rows evaluated',flush=True) report=dict(samples=len(rows),counts=dict(counts),rates={k:v/len(rows) for k,v in counts.items()}, semantic_rerank_median_ms=statistics.median(t['ms'] for t in timings), model='cross-encoder/ms-marco-MiniLM-L6-v2',revision='233902d25c440f23af6f7d6e94d2946bac0bee0a', threads=2,scope='Mind2Web small training-shard offline diagnostic; no browser execution or production promotion', history='Ground-truth previous actions only; teacher-forced history, not autonomous rollouts', parameter_count=sum(p.numel() for p in model.parameters()),timings=timings) 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!='timings'},indent=2)) if __name__=='__main__': main()