Yashp2003's picture
download
raw
11.8 kB
"""
fpqa_batch.py — Fast batch evaluation using HF Inference Provider.
Runs 240 items (10 layouts/room × 4 rooms × 6 qtypes) per model.
Uses concurrent requests for speed.
"""
import os, sys, json, time, math, copy, random, re
from collections import defaultdict
from concurrent.futures import ThreadPoolExecutor, as_completed
from huggingface_hub import InferenceClient
# ===== Configuration =====
MODELS = [
"Qwen/Qwen2.5-7B-Instruct",
"Qwen/Qwen3-32B",
"Qwen/Qwen2.5-72B-Instruct",
]
SAMPLE = 5
MAX_WORKERS = 3
QTYPES = ["pair_distance","view_angle","free_space","obstruction","placement","repositioning"]
CATS = {"pair_distance":"Metric","view_angle":"Metric","free_space":"Topology",
"placement":"Topology","obstruction":"Topology","repositioning":"Dynamic"}
LAYOUT_DIR = "/home/buntu1/.cache/huggingface/hub/datasets--OldDelorean--FloorplanQA-Layouts/snapshots/9ffb74c1c158104fa302baea8e18087d31906bcd/layouts"
OUT_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "results")
def log(*a): print("[fpqa]", *a, flush=True)
# ===== Geometry Helpers =====
from shapely.geometry import Polygon, LineString
from shapely.ops import unary_union
def pts(c): return [(p["x"],p["y"]) for p in c] if c else []
def poly(p):
if len(p)<3: return None
try: P=Polygon(p); return P if P.is_valid else None
except: return None
def room_poly(d): return poly(pts(d.get("room_boundary",[])))
def obs_polys(d):
r=[]
for o in d.get("objects",[]):
lab=(o.get("label","") or "").lower()
if any(x in lab for x in ["rug","light","ceiling","vent"]): continue
p=poly(pts(o.get("points",[])))
if p: r.append(p)
return r
# ===== Ground Truth Generators =====
def gen_gt(qt, data, rt):
room=room_poly(data)
if not room or room.is_empty: return None
obs=obs_polys(data)
if qt=="pair_distance":
objs=[(o.get("label",""), poly(pts(o.get("points",[])))) for o in data.get("objects",[])]
objs=[(l,p) for l,p in objs if p]
if len(objs)<2: return None
d=round(objs[0][1].centroid.distance(objs[1][1].centroid),4)
return {"a":objs[0][0],"b":objs[1][0],"answer":d}
elif qt=="view_angle":
objs=[(o.get("label",""), poly(pts(o.get("points",[])))) for o in data.get("objects",[])]
objs=[(l,p) for l,p in objs if p]
if len(objs)<2: return None
ax,ay=objs[0][1].centroid.x,objs[0][1].centroid.y
bx,by=objs[1][1].centroid.x,objs[1][1].centroid.y
angle=round(math.degrees(math.atan2(by-ay,bx-ax)),1)
if angle<0: angle+=360
return {"a":objs[0][0],"b":objs[1][0],"answer":angle}
elif qt=="free_space":
free=room.difference(unary_union(obs)) if obs else room
pct=round(free.area/room.area*100,1) if room.area>0 else 0
return {"answer":pct}
elif qt=="obstruction":
objs=[(o.get("label",""), poly(pts(o.get("points",[])))) for o in data.get("objects",[])]
objs=[(l,p) for l,p in objs if p]
if len(objs)<2: return None
line=LineString([objs[0][1].centroid.coords[0],objs[1][1].centroid.coords[0]])
blocked=[objs[i][0] for i in range(2,len(objs)) if line.intersects(objs[i][1])]
return {"a":objs[0][0],"b":objs[1][0],"answer":blocked if blocked else ["none"]}
elif qt=="placement":
objs=[(o.get("label",""), poly(pts(o.get("points",[])))) for o in data.get("objects",[])]
objs=[(l,p) for l,p in objs if p]
if not objs: return None
wz=room.boundary.buffer(0.5)
return {"answer":{lab:("wall" if wz.contains(p.centroid) else "center") for lab,p in objs}}
elif qt=="repositioning":
objs=[(o.get("label",""), poly(pts(o.get("points",[])))) for o in data.get("objects",[])]
objs=[(l,p) for l,p in objs if p]
if not objs: return None
from shapely.affinity import translate
moved=translate(objs[0][1],xoff=1.0)
others=unary_union([objs[i][1] for i in range(1,len(objs))])
collides=moved.intersects(others) if others else False
return {"a":objs[0][0],"answer":"collision" if collides else "no collision"}
# ===== Scoring =====
def score(qt, gt, txt):
if not txt: return False,"empty"
t=txt.lower().strip(); a=gt["answer"]
if qt in ("pair_distance","view_angle","free_space"):
nums=re.findall(r'[\d.]+',txt)
if not nums: return False,"no_num"
try:
pv=float(nums[-1])
tol={"pair_distance":max(0.02,abs(a)*0.05),"view_angle":5.0,"free_space":max(3.0,abs(a)*0.1)}[qt]
return abs(pv-a)<=tol, f"pred={pv} gt={a}"
except: return False,"parse_err"
elif qt=="obstruction":
if isinstance(a,list) and a==["none"]:
return ("none" in t or "no object" in t or "no obstruction" in t),"none"
found=sum(1 for x in a if x.lower() in t)
return found>=max(1,len(a)//2), f"found={found}/{len(a)}"
elif qt=="placement":
if isinstance(a,dict):
correct=sum(1 for lab,zone in a.items() if zone.lower() in t)
return correct>=max(1,len(a)//2), f"ok={correct}/{len(a)}"
return False,"bad_gt"
elif qt=="repositioning":
return a.lower() in t, f"gt={a}"
return False,"unknown"
# ===== Prompt Builder =====
def make_prompt(qt, gt, data, fmt="json"):
if fmt=="json":
s=json.dumps({k:v for k,v in data.items() if k!="objects"},indent=1,default=str)[:2000]
header=f"Floor Layout (JSON):\n{s}"
else:
s="<layout>\n"
for k,v in data.items():
if k=="objects": continue
s+=f" <{k}>{v}</{k}>\n"
s+="</layout>"
header=f"Floor Layout (XML):\n{s}"
qmap={"pair_distance":f"What is the distance between {gt.get('a','object A')} and {gt.get('b','object B')}?",
"view_angle":f"What is the angle from {gt.get('a','object A')} to {gt.get('b','object B')}?",
"free_space":"What percentage of the room is free space?",
"obstruction":f"What objects obstruct the path from {gt.get('a','object A')} to {gt.get('b','object B')}?",
"placement":"For each object, is it near the wall or in the center?",
"repositioning":f"If you move {gt.get('a','an object')} 1 meter to the right, does it collide with any other object?"}
return f"{header}\n\nQuestion: {qmap.get(qt,'?')}\n\nReply with just the answer."
# ===== Build Benchmark =====
def build_benchmark(sample=SAMPLE, perturb=0.0, fmt="json"):
items=[]
for rt in ["kitchen","living_room","bedroom","hssd"]:
rdir=os.path.join(LAYOUT_DIR,rt)
if not os.path.isdir(rdir): continue
files=sorted([f for f in os.listdir(rdir) if f.endswith(".json")])[:sample]
for fname in files:
data=json.load(open(os.path.join(rdir,fname)))
if perturb>0:
d2=copy.deepcopy(data); rng=random.Random(abs(hash(d2.get("layout_id",0)))+12345)
def j(pts_list):
return [{"x":round(p["x"]+rng.gauss(0,perturb),4),"y":round(p["y"]+rng.gauss(0,perturb),4)} for p in pts_list]
if "room_boundary" in d2: d2["room_boundary"]=j(d2["room_boundary"])
for o in d2.get("objects",[]): o["points"]=j(o["points"])
data=d2
for qt in QTYPES:
gt=gen_gt(qt,data,rt)
if gt: items.append({"rt":rt,"qt":qt,"cat":CATS[qt],"lid":data.get("layout_id"),"gt":gt,"data":data})
return items
# ===== Run Evaluation =====
def run_eval(model, items, fmt="json"):
client = InferenceClient()
sys_msg = "You are a spatial reasoning assistant. Given a structured floor layout, answer the spatial question. Reply with just the answer."
def query_one(it):
user = make_prompt(it["qt"], it["gt"], it["data"], fmt)
msgs = [{"role":"system","content":sys_msg},{"role":"user","content":user}]
for attempt in range(3):
try:
r = client.chat_completion(model=model, messages=msgs, max_tokens=2000, temperature=0.0)
return r.choices[0].message.content
except Exception as e:
if attempt < 2: time.sleep(1 * (attempt + 1))
else: return f"__ERROR__:{e}"
results = []
t0 = time.time()
try:
with ThreadPoolExecutor(max_workers=MAX_WORKERS) as pool:
futures = {pool.submit(query_one, it): i for i, it in enumerate(items)}
for i, fut in enumerate(as_completed(futures)):
it = items[futures[fut]]
txt = fut.result()
ok, reason = score(it["qt"], it["gt"], txt)
results.append({"rt":it["rt"],"qt":it["qt"],"cat":it["cat"],"lid":it["lid"],"ok":bool(ok),"reason":reason})
if (i+1) % 20 == 0:
acc = sum(r["ok"] for r in results)/len(results)
log(f" [{model.split('/')[-1]}] {i+1}/{len(items)} {time.time()-t0:.0f}s acc={acc:.3f}")
except Exception as e:
log(f" INTERRUPTED at {len(results)}/{len(items)}: {e}")
overall = sum(r["ok"] for r in results)/max(1,len(results))
by_cat = defaultdict(lambda:{"c":0,"n":0})
by_rt = defaultdict(lambda:{"c":0,"n":0})
for r in results:
by_cat[r["cat"]]["c"]+=int(r["ok"]); by_cat[r["cat"]]["n"]+=1
by_rt[r["rt"]]["c"]+=int(r["ok"]); by_rt[r["rt"]]["n"]+=1
return {
"model": model, "fmt": fmt, "n_items": len(results),
"overall_acc": round(overall, 4), "elapsed_sec": round(time.time()-t0, 1),
"by_cat": {k:{"acc":round(v["c"]/v["n"],4),"n":v["n"]} for k,v in by_cat.items()},
"by_rt": {k:{"acc":round(v["c"]/v["n"],4),"n":v["n"]} for k,v in by_rt.items()},
}
if __name__ == "__main__":
os.makedirs(OUT_DIR, exist_ok=True)
log("Building benchmark...")
base_items = build_benchmark(sample=SAMPLE, perturb=0.0)
log(f"Benchmark: {len(base_items)} items")
summary = {}
for model in MODELS:
name = model.split("/")[-1]
log(f"\n=== Running {model} (JSON baseline) ===")
r = run_eval(model, base_items, fmt="json")
path = os.path.join(OUT_DIR, f"{name}_json.json")
json.dump(r, open(path, "w"), indent=2)
log(f" acc={r['overall_acc']} elapsed={r['elapsed_sec']}s -> {path}")
for cat in ["Metric","Topology","Dynamic"]:
if cat in r["by_cat"]:
log(f" {cat}: {r['by_cat'][cat]['acc']} n={r['by_cat'][cat]['n']}")
for rt in ["kitchen","living_room","bedroom","hssd"]:
if rt in r["by_rt"]:
log(f" {rt}: {r['by_rt'][rt]['acc']} n={r['by_rt'][rt]['n']}")
summary[model] = r
# XML ablation (Claim 4) — only on 7B
log(f"\n=== XML ABLATION (7B) ===")
r_xml = run_eval("Qwen/Qwen2.5-7B-Instruct", base_items, fmt="xml")
path = os.path.join(OUT_DIR, "Qwen2.5-7B-Instruct_xml.json")
json.dump(r_xml, open(path, "w"), indent=2)
log(f" acc={r_xml['overall_acc']}")
# Perturbation (Claim 5) — only on 7B
log(f"\n=== PERTURBATION (7B) ===")
pert_items = build_benchmark(sample=SAMPLE, perturb=0.1)
r_pert = run_eval("Qwen/Qwen2.5-7B-Instruct", pert_items, fmt="json")
path = os.path.join(OUT_DIR, "Qwen2.5-7B-Instruct_perturb.json")
json.dump(r_pert, open(path, "w"), indent=2)
log(f" acc={r_pert['overall_acc']}")
# Summary
summary_path = os.path.join(OUT_DIR, "summary.json")
summary["xml_7b"] = r_xml
summary["perturb_7b"] = r_pert
json.dump(summary, open(summary_path, "w"), indent=2)
log(f"\n=== FINAL SUMMARY ===")
for k, v in summary.items():
log(f" {k}: acc={v['overall_acc']} elapsed={v['elapsed_sec']}s")
log(f"WROTE {summary_path}")

Xet Storage Details

Size:
11.8 kB
·
Xet hash:
bfc9bae56bab7d48a03a354c8a4d1c9ad6d58d880faf51f87bec8df6de05301f

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.