bhavibhatt commited on
Commit
77e4ff5
·
verified ·
1 Parent(s): 1368c2d

Delete ensemble.py

Browse files
Files changed (1) hide show
  1. ensemble.py +0 -43
ensemble.py DELETED
@@ -1,43 +0,0 @@
1
- import cv2
2
- import numpy as np
3
- from ultralytics import YOLO
4
-
5
- class WeldVisionEnsemble:
6
- CLASS_NAMES={0:'Bad Welding',1:'Crack',2:'Excess Reinforcement',3:'Good Welding',4:'Porosity',5:'Spatters'}
7
- PENALTIES={'Crack':40,'Porosity':15,'Spatters':5,'Excess Reinforcement':20,'Bad Welding':50,'Good Welding':0}
8
- def __init__(self, model_dir, conf=.25, ensemble_iou=.50, imgsz=640):
9
- self.conf=conf; self.ensemble_iou=ensemble_iou; self.imgsz=imgsz
10
- self.base1=YOLO(f'{model_dir}/best.pt'); self.base2=YOLO(f'{model_dir}/best_v0.pt')
11
- self.crack=YOLO(f'{model_dir}/crack_specialist.pt'); self.spatters=YOLO(f'{model_dir}/spatters_specialist.pt')
12
- @staticmethod
13
- def mask_iou(a,b):
14
- a=a.astype(bool); b=b.astype(bool); u=np.logical_or(a,b).sum(); return float(np.logical_and(a,b).sum()/u) if u else 0.0
15
- def extract(self,r,source):
16
- if r.boxes is None or r.masks is None: return []
17
- out=[]
18
- for box,mask in zip(r.boxes.data.cpu().numpy(),r.masks.data.cpu().numpy()):
19
- x1,y1,x2,y2,conf,cls=box; cls=int(cls)
20
- name='Crack' if source=='crack' else 'Spatters' if source=='spatters' else self.CLASS_NAMES.get(cls,str(cls))
21
- out.append({'box':np.array([x1,y1,x2,y2],dtype=np.float32),'conf':float(conf),'class_name':name,'mask':mask.astype(np.float32),'source':source})
22
- return out
23
- def run(self,model,img,source):
24
- return self.extract(model.predict(img,conf=self.conf,imgsz=self.imgsz,verbose=False)[0],source)
25
- def merge(self,preds):
26
- selected=[]
27
- for p in sorted(preds,key=lambda x:x['conf'],reverse=True):
28
- if not any(p['class_name']==q['class_name'] and self.mask_iou(p['mask'],q['mask'])>=self.ensemble_iou for q in selected): selected.append(p)
29
- return selected
30
- def predict(self,image):
31
- image=np.asarray(image)
32
- if image.ndim!=3 or image.shape[2]!=3: raise ValueError('image must be HxWx3 RGB')
33
- bgr=cv2.cvtColor(image,cv2.COLOR_RGB2BGR)
34
- p1=self.run(self.base1,bgr,'base_1'); p2=self.run(self.base2,bgr,'base_2'); pc=self.run(self.crack,bgr,'crack'); ps=self.run(self.spatters,bgr,'spatters')
35
- merged=self.merge(p1+p2+pc+ps); detections=[]; h,w=bgr.shape[:2]
36
- for p in merged:
37
- c=p['class_name']
38
- if c=='Good Welding': continue
39
- sev='HIGH' if self.PENALTIES.get(c,0)>=30 else 'MEDIUM' if self.PENALTIES.get(c,0)>=15 else 'LOW' if self.PENALTIES.get(c,0)>0 else 'NONE'
40
- x1,y1,x2,y2=p['box']; detections.append({'class':c,'confidence':round(p['conf'],4),'severity':sev,'box':[round(float(max(0,min(w,x1))),2),round(float(max(0,min(h,y1))),2),round(float(max(0,min(w,x2))),2),round(float(max(0,min(h,y2))),2)],'source':p['source']})
41
- score=max(0,100-sum(self.PENALTIES.get(d['class'],0) for d in detections)); high=any(d['severity']=='HIGH' for d in detections); med=any(d['severity']=='MEDIUM' for d in detections)
42
- decision='FAIL' if score<70 or high else 'REVIEW' if score<85 or med else 'PASS'
43
- return {'model':'WeldVision-Ensemble','version':'1.0','decision':decision,'score':score,'detections':detections,'model_counts':{'best.pt':len(p1),'best_v0.pt':len(p2),'crack_specialist.pt':len(pc),'spatters_specialist.pt':len(ps),'ensemble':len(merged)},'recommendation':'Significant visible defect detected. Human inspection required.' if decision=='FAIL' else 'Moderate defect or quality concern detected. Human review required.' if decision=='REVIEW' else 'No significant visible defect detected.'}