| """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'] |
| |
| 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() |
|
|