"""Attention Probe: refine bbox centers using YOLO's internal features. =================================================================== Idea: YOLO's P3 features contain sub-pixel position information. For each detection, extract the feature vector at the anchor point, and use a lightweight MLP to predict center/size corrections. Unlike BRN (which used raw pixels and failed), this uses YOLO's OWN features which already encode goat-specific spatial information. Zero training - just run inference + apply correction. """ 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 import torch.nn as nn 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) class ProbeCorrector(nn.Module): """Lightweight MLP that predicts bbox correction from feature vectors.""" def __init__(self, in_dim=128, hidden=32): super().__init__() self.net = nn.Sequential( nn.Linear(in_dim, hidden), nn.ReLU(), nn.Linear(hidden, 4), # Δcx, Δcy, Δw, Δh ) self.net[-1].weight.data.zero_() self.net[-1].bias.data.zero_() def forward(self, feats): return self.net(feats) def train_probe(model, img_dir, lbl_dir, n_samples=2000): """Train the probe on training set detections.""" device = next(model.model.parameters()).device probe = ProbeCorrector().to(device) opt = torch.optim.Adam(probe.parameters(), lr=1e-3) # Register hook to capture P3 features before detection head p3_feats = None def hook_fn(module, input, output): nonlocal p3_feats p3_feats = input[0][0] # P3 features: [B, C, H, W] detect = model.model.model[-1] handle = detect.register_forward_hook(hook_fn) img_files = sorted([f for f in os.listdir(img_dir) if f.endswith('.jpg')]) import random; random.seed(42); random.shuffle(img_files) img_files = img_files[:n_samples] samples = [] for f in tqdm(img_files, desc='Train probe'): img = Image.open(os.path.join(img_dir, f)) iw, ih = img.size # GT boxes gt_boxes = [] lf = f.replace('.jpg','.txt') lbl_path = os.path.join(lbl_dir, lf) if os.path.exists(lbl_path): with open(lbl_path) 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]) gt_boxes.append([cx*iw, cy*ih, w*iw, h*ih]) if not gt_boxes: continue # Forward pass p3_feats = None with torch.no_grad(): results = model.predict(img, imgsz=1536, conf=0.25, iou=0.7, max_det=100, verbose=False) if p3_feats is None or not results or len(results[0].boxes) == 0: continue pred_boxes = results[0].boxes.xyxy.cpu().numpy() fmap = p3_feats # [1, C, H, W] C, H, W = fmap.shape[1], fmap.shape[2], fmap.shape[3] # Match predictions to GT for pred in pred_boxes: best_iou, best_gt = 0, None for gt in gt_boxes: gt_xyxy = [gt[0]-gt[2]/2, gt[1]-gt[3]/2, gt[0]+gt[2]/2, gt[1]+gt[3]/2] iou = iou_fn(pred.tolist(), gt_xyxy) if iou > best_iou: best_iou = iou; best_gt = gt if best_iou < 0.5 or best_gt is None: continue # Get feature at predicted center cx_pred = (pred[0]+pred[2])/2 * W / iw cy_pred = (pred[1]+pred[3])/2 * H / ih cx_pred = int(np.clip(cx_pred, 0, W-1)) cy_pred = int(np.clip(cy_pred, 0, H-1)) feat = fmap[0, :, cy_pred, cx_pred].cpu().numpy() # [C] # Target correction (normalized by object size) gt_cx, gt_cy, gt_w, gt_h = best_gt pred_w = pred[2]-pred[0] pred_h = pred[3]-pred[1] norm = max(pred_w, pred_h) + 1e-8 dcx = (gt_cx - (pred[0]+pred[2])/2) / norm dcy = (gt_cy - (pred[1]+pred[3])/2) / norm dw = (gt_w - pred_w) / norm dh = (gt_h - pred_h) / norm samples.append((feat, np.array([dcx, dcy, dw, dh], dtype=np.float32))) if len(samples) < 100: print(f'Only {len(samples)} samples, probe not trained') return None # Train print(f'Training probe on {len(samples)} samples...') X = torch.tensor(np.stack([s[0] for s in samples]), dtype=torch.float32).to(device) Y = torch.tensor(np.stack([s[1] for s in samples]), dtype=torch.float32).to(device) for epoch in range(100): opt.zero_grad() pred = probe(X) loss = nn.functional.l1_loss(pred, Y) loss.backward() opt.step() if epoch % 20 == 0: print(f' ep{epoch}: loss={loss.item():.5f}') handle.remove() return probe def apply_probe(model, probe, img, boxes): """Apply probe corrections to detected boxes.""" if len(boxes) == 0 or probe is None: return boxes device = next(model.model.parameters()).device p3_feats = None def hook_fn(module, input, output): nonlocal p3_feats p3_feats = input[0][0] detect = model.model.model[-1] handle = detect.register_forward_hook(hook_fn) with torch.no_grad(): results = model.predict(img, imgsz=1536, conf=0.25, iou=0.7, max_det=100, verbose=False) handle.remove() if p3_feats is None: return boxes fmap = p3_feats C, H, W = fmap.shape[1], fmap.shape[2], fmap.shape[3] iw, ih = img.size refined = [] for box in boxes: cx_pred = (box[0]+box[2])/2 * W / iw cy_pred = (box[1]+box[3])/2 * H / ih cx_pred = int(np.clip(cx_pred, 0, W-1)) cy_pred = int(np.clip(cy_pred, 0, H-1)) feat = fmap[0, :, cy_pred, cx_pred] # [C] delta = probe(feat.float().unsqueeze(0).to(device)).cpu().numpy()[0] w = box[2]-box[0]; h = box[3]-box[1] norm = max(w, h) + 1e-8 new_cx = (box[0]+box[2])/2 + delta[0]*norm new_cy = (box[1]+box[3])/2 + delta[1]*norm new_w = w + delta[2]*norm new_h = h + delta[3]*norm new_x1 = max(0, new_cx-new_w/2); new_y1 = max(0, new_cy-new_h/2) new_x2 = min(iw, new_cx+new_w/2); new_y2 = min(ih, new_cy+new_h/2) refined.append([new_x1, new_y1, new_x2, new_y2]) return np.array(refined) 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 def main(): val_dir='Data/Detection_dataset/images/val' lbl_dir='Data/Detection_dataset/labels/val' img_dir='Data/Detection_dataset/images/train' train_lbl='Data/Detection_dataset/labels/train' vfs=sorted([f for f in os.listdir(val_dir) if f.endswith('.jpg')]) 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 # Load model model = YOLO('runs/detect/Detection_experiments/v6_1_s_refined/weights/best.pt') # Train probe probe = train_probe(model, img_dir, train_lbl, n_samples=300) # Baseline eval bp=[] for vf in tqdm(vfs,desc='Baseline'): img=Image.open(os.path.join(val_dir,vf)) b=pred_fn(model,img,1536);bp.append(b) mAP_base=eval_boxes('Baseline',bp) # Probe-refined eval rp=[] for vf in tqdm(vfs,desc='Probe'): img=Image.open(os.path.join(val_dir,vf)) b=pred_fn(model,img,1536) if probe is not None and len(b)>0: b=apply_probe(model,probe,img,b) rp.append(b) mAP_probe=eval_boxes('Probe',rp) sep='='*60 print('\n{}'.format(sep)) print('ATTENTION PROBE') print(sep) print('Baseline: {:.4f}'.format(mAP_base)) print('Probe: {:.4f} (+{:.4f})'.format(mAP_probe,mAP_probe-mAP_base)) with open('logs/probe_result.json','w') as f: json.dump({'baseline':round(mAP_base,4),'probe':round(mAP_probe,4)},f) if __name__=='__main__': main()