goat / Scripts /eval_prior_wbf.py
LightChuan's picture
Upload folder using huggingface_hub
6a5bb7e verified
Raw
History Blame Contribute Delete
8.23 kB
"""Position-Size Prior WBF: filter implausible boxes using learned priors.
Each camera × position bin has expected goat size. Boxes far from expectation
get downweighted in WBF. Zero training cost.
"""
import sys,os,json,gc,numpy as np
from PIL import Image,ImageEnhance
from tqdm import tqdm
from collections import defaultdict
PROJECT_DIR='/home/user/goat'
os.chdir(PROJECT_DIR);sys.path.insert(0,PROJECT_DIR)
from ultralytics import YOLO
import torch
def iou_fn(b1,b2):
x1,y1=max(b1[0],b2[0]),max(b1[1],b2[1])
x2,y2=min(b1[2],b2[2]),min(b1[3],b2[3])
inter=max(0,x2-x1)*max(0,y2-y1)
a1=(b1[2]-b1[0])*(b1[3]-b1[1]);a2=(b2[2]-b2[0])*(b2[3]-b2[1])
return inter/(a1+a2-inter+1e-8)
def wbf_fn(bl,sl,thr=0.55,prior_weights=None):
if not bl or all(len(b)==0 for b in bl): return np.array([])
ab,as_=[],[]
for i,(bx,sx) in enumerate(zip(bl,sl)):
pw=prior_weights[i] if prior_weights else np.ones(len(bx))
for j in range(len(bx)):
ab.append(bx[j]);as_.append(sx[j]*pw[j])
if not ab: return np.array([])
ab=np.array(ab);as_=np.array(as_)
o=np.argsort(-as_);ab=ab[o];as_=as_[o]
cl,us=[],np.zeros(len(ab),dtype=bool)
for i in range(len(ab)):
if us[i]: continue
c=[(ab[i],as_[i])];us[i]=True
for j in range(i+1,len(ab)):
if us[j]: continue
tw=sum(s for _,s in c)
ct=sum(b*s/tw for b,s in c)
if iou_fn(ct.tolist(),ab[j].tolist())>thr: c.append((ab[j],as_[j]));us[j]=True
cl.append(c)
rb,rs=[],[]
for c in cl:
tw=sum(s for _,s in c)
rb.append(sum(b*s/tw for b,s in c));rs.append(tw)
return np.array(rb)
def pred_fn(model,img,sz,flip=False,bright=1.0):
ia=img
if bright!=1.0: ia=ImageEnhance.Brightness(ia).enhance(bright)
if flip: ia=ia.transpose(Image.FLIP_LEFT_RIGHT)
r=model.predict(ia,imgsz=sz,conf=0.25,iou=0.7,max_det=100,verbose=False)
if not r or len(r[0].boxes)==0: return np.array([]),np.array([])
b=r[0].boxes.xyxy.cpu().numpy();s=r[0].boxes.conf.cpu().numpy()
if flip: w=img.size[0];b[:,[0,2]]=w-b[:,[2,0]]
return b,s
def build_size_prior():
"""Learn per-camera per-position expected goat size."""
img_dir='Data/Detection_dataset/images/train'
lbl_dir='Data/Detection_dataset/labels/train'
# Grid: 10x10 bins per camera
prior=defaultdict(lambda: defaultdict(list))
for f in sorted(os.listdir(img_dir)):
if not f.endswith('.jpg'): continue
cam=f.split('_2025')[0]
lf=f.replace('.jpg','.txt')
with open(os.path.join(lbl_dir,lf)) as fh:
for line in fh:
p=line.strip().split()
if len(p)>=5:
cx,cy,w,h=float(p[1]),float(p[2]),float(p[3]),float(p[4])
bx,by=int(cx*10),int(cy*10)
bx=min(9,max(0,bx));by=min(9,max(0,by))
prior[cam][(bx,by)].append(np.sqrt(w*h))
# Compute mean/std per bin
stats={}
for cam in prior:
stats[cam]={}
for (bx,by),sizes in prior[cam].items():
if len(sizes)>=3:
stats[cam][(bx,by)]={
'mean':np.mean(sizes),'std':np.std(sizes),'n':len(sizes)
}
return stats
def size_prior_weight(box, cam, stats):
"""Return weight [0.5, 1.5] based on how plausible the box size is."""
if cam not in stats: return 1.0
x1,y1,x2,y2=box
cx=(x1+x2)/2/3200;cy=(y1+y2)/2/1800
w=(x2-x1)/3200;h=(y2-y1)/1800
sz=np.sqrt(w*h)
bx,by=int(cx*10),int(cy*10)
bx=min(9,max(0,bx));by=min(9,max(0,by))
if (bx,by) not in stats[cam]: return 1.0
s=stats[cam][(bx,by)]
z=abs(sz-s['mean'])/(s['std']+1e-8)
# z>2 → unlikely size → downweight
if z>3: return 0.3
if z>2: return 0.6
if z<0.5: return 1.2 # very typical size → boost
return 1.0
def main():
val_dir='Data/Detection_dataset/images/val'
lbl_dir='Data/Detection_dataset/labels/val'
vfs=sorted([f for f in os.listdir(val_dir) if f.endswith('.jpg')])
exp='runs/detect/Detection_experiments'
# Build prior
print('Building size priors...')
prior_stats=build_size_prior()
for cam in sorted(prior_stats):
print(f' {cam}: {len(prior_stats[cam])} grid cells')
# Quick test with v6_1
mpaths=[('v6_1',f'{exp}/v6_1_s_refined/weights/best.pt')]
mpaths=[(n,p) for n,p in mpaths if os.path.exists(p)]
iou_thrs=[round(0.5+i*0.05,2) for i in range(10)]
def eval_boxes(name,boxes_per_img):
tp={t:0 for t in iou_thrs};tg=0
for idx,vf in enumerate(vfs):
img=Image.open(os.path.join(val_dir,vf))
gb=[]
lf=vf.replace('.jpg','.txt')
with open(os.path.join(lbl_dir,lf)) as f:
for line in f:
p=line.strip().split()
if len(p)>=5:
cx,cy,w,h=[float(x) for x in p[1:5]]
gb.append([(cx-w/2)*img.size[0],(cy-h/2)*img.size[1],(cx+w/2)*img.size[0],(cy+h/2)*img.size[1]])
tg+=len(gb)
if not gb: continue
for t in iou_thrs:
mt=set()
for pb in boxes_per_img[idx]:
if len(pb)==0: continue
bi,bg=0,-1
for gi,gt in enumerate(gb):
if gi in mt: continue
ii=iou_fn(pb.tolist(),gt)
if ii>bi: bi=ii;bg=gi
if bi>=t and bg>=0: tp[t]+=1;mt.add(bg)
rec=[tp[t]/tg for t in iou_thrs]
mAP=np.mean(rec)
print('{}: mAP50-95={:.4f} IoU@75={:.4f}'.format(name,mAP,rec[5]))
return mAP,rec[5]
# Baseline
m0=YOLO(mpaths[0][1])
bp=[]
for vf in tqdm(vfs,desc='Baseline'):
img=Image.open(os.path.join(val_dir,vf))
b,_=pred_fn(m0,img,1536);bp.append(b)
del m0;gc.collect();torch.cuda.empty_cache()
mAP_base,r75_base=eval_boxes('Baseline',bp)
# Multi-scale WBF without prior
m=YOLO(mpaths[0][1])
mp_no_prior=[]
for vf in tqdm(vfs,desc='MS-WBF'):
img=Image.open(os.path.join(val_dir,vf))
cam=vf.split('_2025')[0]
bl,sl=[],[]
for sz in [1280,1536,1920]:
for fl in [False,True]:
for br in [1.0,1.2]:
b,s=pred_fn(m,img,sz,fl,br)
if len(b)>0: bl.append(b);sl.append(s)
mp_no_prior.append(wbf_fn(bl,sl))
mAP_ms,r75_ms=eval_boxes('MS-WBF (no prior)',mp_no_prior)
# Multi-scale WBF WITH size prior
mp_prior=[]
prior_used=0;total_boxes=0
for vf in tqdm(vfs,desc='PriorWBF'):
img=Image.open(os.path.join(val_dir,vf))
cam=vf.split('_2025')[0]
bl,sl,weights=[],[],[]
for sz in [1280,1536,1920]:
for fl in [False,True]:
for br in [1.0,1.2]:
b,s=pred_fn(m,img,sz,fl,br)
if len(b)>0:
bl.append(b);sl.append(s)
w=np.ones(len(b))
for i in range(len(b)):
pw=size_prior_weight(b[i],cam,prior_stats)
w[i]=pw
if pw!=1.0: prior_used+=1
weights.append(w)
total_boxes+=sum(len(bx) for bx in bl)
mp_prior.append(wbf_fn(bl,sl,prior_weights=weights))
mAP_prior,r75_prior=eval_boxes('PriorWBF',mp_prior)
del m;gc.collect();torch.cuda.empty_cache()
sep='='*60
print('\n{}'.format(sep))
print('POSITION-SIZE PRIOR WBF')
print(sep)
print('Baseline: {:.4f}'.format(mAP_base))
print('MS-WBF: {:.4f} (+{:.4f})'.format(mAP_ms,mAP_ms-mAP_base))
print('PriorWBF: {:.4f} (+{:.4f})'.format(mAP_prior,mAP_prior-mAP_base))
print('Prior applied: {}/{} ({:.1f}%)'.format(prior_used,total_boxes,prior_used/total_boxes*100 if total_boxes>0 else 0))
with open('logs/prior_wbf_results.json','w') as f:
json.dump({'baseline':round(mAP_base,4),'ms_wbf':round(mAP_ms,4),'prior_wbf':round(mAP_prior,4),'delta':round(mAP_prior-mAP_ms,4)},f,indent=2)
print('Saved.')
if __name__=='__main__':
main()