goat / Scripts /eval_ultimate10.py
LightChuan's picture
Upload folder using huggingface_hub
6a5bb7e verified
Raw
History Blame Contribute Delete
6.02 kB
"""Ultimate 10-model WBF evaluation."""
import sys,os,json,gc,numpy as np
from PIL import Image,ImageEnhance
from tqdm import tqdm
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):
if not bl or all(len(b)==0 for b in bl): return np.array([])
ab,as_=[],[]
for bx,sx in zip(bl,sl): ab.extend(bx);as_.extend(sx)
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 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'
mpaths=[
('v6_1',f'{exp}/v6_1_s_refined/weights/best.pt'),
('s111',f'{exp}/v17_seed_111/weights/best.pt'),
('s888',f'{exp}/v18_seed_888/weights/best.pt'),
('s777',f'{exp}/v17_seed_777/weights/best.pt'),
('s333',f'{exp}/v15_seed_333/weights/best.pt'),
('s123',f'{exp}/v12_seed_123/weights/best.pt'),
('s555',f'{exp}/v18_seed_555/weights/best.pt'),
('s42',f'{exp}/v12_seed_42/weights/best.pt'),
('s444',f'{exp}/v18_seed_444/weights/best.pt'),
('M11m',f'{exp}/v16_yolo11m/weights/best.pt'),
]
mpaths=[(n,p) for n,p in mpaths if os.path.exists(p)]
print('{} models: {}'.format(len(mpaths),[n for n,_ in mpaths]))
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]
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)
all_preds={}
for name,path in mpaths:
print('Processing {}...'.format(name))
m=YOLO(path)
ip=[]
for vf in tqdm(vfs,desc=name,leave=False):
img=Image.open(os.path.join(val_dir,vf))
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)
ip.append((bl,sl))
all_preds[name]=ip
del m;gc.collect();torch.cuda.empty_cache()
ksp=[]
for idx in range(len(vfs)):
ab,as_=[],[]
for name,_ in mpaths:
bl,sl=all_preds[name][idx]
for b,s in zip(bl,sl):
if len(b)>0: ab.append(b);as_.append(s)
ksp.append(wbf_fn(ab,as_))
mAP_ks,r75_ks=eval_boxes('KS 10m',ksp)
ksp5=[]
for idx in range(len(vfs)):
ab,as_=[],[]
for name,_ in mpaths[:5]:
bl,sl=all_preds[name][idx]
for b,s in zip(bl,sl):
if len(b)>0: ab.append(b);as_.append(s)
ksp5.append(wbf_fn(ab,as_))
mAP_ks5,r75_ks5=eval_boxes('KS Top5',ksp5)
sep='='*60
print('\n{}'.format(sep))
print('ULTIMATE 10-MODEL')
print(sep)
print('Baseline: {:.4f}'.format(mAP_base))
print('KS Top5: {:.4f} (+{:.4f})'.format(mAP_ks5,mAP_ks5-mAP_base))
print('KS 10m: {:.4f} (+{:.4f})'.format(mAP_ks,mAP_ks-mAP_base))
with open('logs/ultimate10_results.json','w') as f:
json.dump({'baseline':round(mAP_base,4),'ks_top5':round(mAP_ks5,4),'ks_10m':round(mAP_ks,4),'n_models':len(mpaths)},f,indent=2)
print('Saved.')
if __name__=='__main__':
main()