Download baim/retrieval_audit.py from devildasdf/devils-agent: direct link, hf CLI and curl.
- Browser
- Download file 4.8 kB
-
https://huggingface.co/devildasdf/devils-agent/resolve/main/baim/retrieval_audit.py
- Command line
-
hf download hf://devildasdf/devils-agent/baim/retrieval_audit.py
-
curl -L -o retrieval_audit.py https://huggingface.co/devildasdf/devils-agent/resolve/main/baim/retrieval_audit.py
4.8 kB
| """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() | |