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