| """ |
| ArtiAgent Orchestrator — DefectFill Edition (with VLM list selection) |
| |
| The VLM selects object_class from a given list and defect_type from the |
| corresponding per-object-class list. Product description drives the selection. |
| |
| Usage: |
| python artiagent_orchestrator.py \\ |
| --product-desc "VCSEL laser diode with glass lens cap" \\ |
| --image ./clean_chip.png \\ |
| --output-dir ./defect_output \\ |
| --checkpoint-dir "C:/.../checkpoints" \\ |
| --valid-object-classes '["xray_PCB","vcsel"]' \\ |
| --valid-defect-types '{"xray_PCB":["xray_die","bubble"],"vcsel":["scratch"]}' \\ |
| --device cuda |
| """ |
|
|
| import os |
| import sys |
| import json |
| import argparse |
| import uuid |
| import traceback |
| from pathlib import Path |
| from typing import Dict, List, Optional, Tuple |
| from datetime import datetime |
|
|
| import numpy as np |
| import torch |
| from PIL import Image |
|
|
| SCRIPT_DIR = Path(__file__).parent |
| if str(SCRIPT_DIR) not in sys.path: |
| sys.path.insert(0, str(SCRIPT_DIR)) |
|
|
| from pipeline.local_vlm_client import LocalVLMClient |
| from pipeline.prompts import ( |
| plan_defects_for_product, |
| artifact_description, |
| MoneyManager |
| ) |
| from pipeline.gsam_detector import GSAMDetector |
| from pipeline.defectfill_generator import DefectFillGenerator, DefectFillConfig |
| from pipeline.defect_rag import get_rag |
| from pipeline.domain_router import get_router |
|
|
| import cv2 |
|
|
|
|
| def blend_defect_onto_real_image( |
| real_image: np.ndarray, |
| defect_image: Image.Image, |
| defect_mask: np.ndarray, |
| target_bbox: List[int], |
| max_defect_ratio: Optional[float] = None, |
| mask_shape: str = "rectangle" |
| ) -> Tuple[np.ndarray, np.ndarray]: |
| """Injects a DefectFill defect patch onto a real clean factory image.""" |
| defect_np = np.array(defect_image) |
| mask_uint8 = (defect_mask.astype(np.uint8) * 255) if defect_mask.dtype == bool else defect_mask.astype(np.uint8) |
|
|
| ys, xs = np.where(mask_uint8 > 0) |
| if len(ys) == 0 or len(xs) == 0: |
| return real_image.copy(), np.zeros(real_image.shape[:2], dtype=np.uint8) |
|
|
| y1_d, y2_d = ys.min(), ys.max() |
| x1_d, x2_d = xs.min(), xs.max() |
|
|
| defect_patch = defect_np[y1_d:y2_d + 1, x1_d:x2_d + 1] |
| mask_patch = mask_uint8[y1_d:y2_d + 1, x1_d:x2_d + 1] |
|
|
| x1_t, y1_t, x2_t, y2_t = target_bbox |
| target_w = max(1, x2_t - x1_t) |
| target_h = max(1, y2_t - y1_t) |
|
|
| if max_defect_ratio is None or max_defect_ratio >= 1.0: |
| final_w = target_w |
| final_h = target_h |
| else: |
| if max_defect_ratio < 0.2: |
| max_defect_ratio = 0.2 |
| scale_factor = np.sqrt(max_defect_ratio) |
| scaled_w = int(target_w * scale_factor) |
| scaled_h = int(target_h * scale_factor) |
|
|
| patch_h, patch_w = defect_patch.shape[:2] |
| aspect_ratio = patch_w / max(1, patch_h) |
|
|
| if aspect_ratio > 1: |
| final_w = max(15, scaled_w) |
| final_h = max(15, int(final_w / aspect_ratio)) |
| else: |
| final_h = max(15, scaled_h) |
| final_w = max(15, int(final_h * aspect_ratio)) |
|
|
| defect_patch_resized = cv2.resize(defect_patch, (final_w, final_h), interpolation=cv2.INTER_AREA) |
| mask_patch_resized = cv2.resize(mask_patch, (final_w, final_h), interpolation=cv2.INTER_NEAREST) |
|
|
| shape_type = mask_shape.lower().strip() if mask_shape else "rectangle" |
|
|
| if shape_type == "circle": |
| geom_mask = np.zeros((final_h, final_w), dtype=np.uint8) |
| center = (final_w // 2, final_h // 2) |
| radius = max(1, min(final_w, final_h) // 2 - 1) |
| cv2.circle(geom_mask, center, radius, 255, thickness=-1) |
| mask_patch_resized = geom_mask |
|
|
| elif shape_type == "square": |
| geom_mask = np.zeros((final_h, final_w), dtype=np.uint8) |
| side = max(1, min(final_w, final_h) - 2) |
| top_left_x = (final_w - side) // 2 |
| top_left_y = (final_h - side) // 2 |
| cv2.rectangle( |
| geom_mask, |
| (top_left_x, top_left_y), |
| (top_left_x + side, top_left_y + side), |
| 255, |
| thickness=-1 |
| ) |
| mask_patch_resized = geom_mask |
|
|
| elif shape_type == "rectangle": |
| mask_patch_resized = np.full((final_h, final_w), 255, dtype=np.uint8) |
|
|
| center_x = x1_t + target_w // 2 |
| center_y = y1_t + target_h // 2 |
| center = (center_x, center_y) |
|
|
| real_bgr = cv2.cvtColor(real_image, cv2.COLOR_RGB2BGR) |
| patch_bgr = cv2.cvtColor(defect_patch_resized, cv2.COLOR_RGB2BGR) |
|
|
| patch_mean = np.mean(defect_patch_resized) |
| clone_mode = cv2.NORMAL_CLONE if patch_mean < 30 else cv2.MIXED_CLONE |
|
|
| blended_bgr = cv2.seamlessClone( |
| patch_bgr, |
| real_bgr, |
| mask_patch_resized, |
| center, |
| clone_mode |
| ) |
| blended_rgb = cv2.cvtColor(blended_bgr, cv2.COLOR_BGR2RGB) |
|
|
| full_mask = np.zeros(real_image.shape[:2], dtype=np.uint8) |
| top_left_x = max(0, center_x - final_w // 2) |
| top_left_y = max(0, center_y - final_h // 2) |
|
|
| h_end = min(real_image.shape[0], top_left_y + final_h) |
| w_end = min(real_image.shape[1], top_left_x + final_w) |
|
|
| mask_crop_h = h_end - top_left_y |
| mask_crop_w = w_end - top_left_x |
|
|
| if mask_crop_h > 0 and mask_crop_w > 0: |
| full_mask[top_left_y:h_end, top_left_x:w_end] = ( |
| mask_patch_resized[:mask_crop_h, :mask_crop_w] > 128 |
| ).astype(np.uint8) |
|
|
| return blended_rgb, full_mask |
|
|
|
|
| def create_visual_prompt_image(full_image: np.ndarray, bbox: list) -> np.ndarray: |
| """Draws a bright neon bounding box on the full image around the target ROI.""" |
| viz_img = full_image.copy() |
| x1, y1, x2, y2 = bbox |
| cv2.rectangle(viz_img, (x1, y1), (x2, y2), (0, 255, 0), thickness=2) |
| return viz_img |
|
|
|
|
| def resolve_checkpoint_path(checkpoint_dir: str, object_class: str, defect_type: str) -> str: |
| """Resolve DefectFill checkpoint path from object_class + defect_type.""" |
| path = Path(checkpoint_dir) / object_class / defect_type / "checkpoints" / "checkpoint_final.pt" |
| if not path.exists(): |
| alt = Path(checkpoint_dir) / object_class / defect_type / "checkpoint_final.pt" |
| if alt.exists(): |
| return str(alt) |
| raise FileNotFoundError( |
| f"Checkpoint not found for object_class='{object_class}', defect_type='{defect_type}'.\n" |
| f"Tried: {path}\nAlso tried: {alt}" |
| ) |
| return str(path) |
|
|
|
|
| class ArtiAgentOrchestrator: |
| """Agentic orchestrator for directed defect generation with DefectFill.""" |
|
|
| def __init__( |
| self, |
| device='cuda', |
| output_dir='./defect_output', |
| vlm_model='gemma3:12b', |
| checkpoint_dir: str = "", |
| object_class: str = "", |
| defect_type: str = "", |
| valid_object_classes: Optional[List[str]] = None, |
| valid_defect_types: Optional[Dict[str, List[str]]] = None, |
| image_size: int = 512, |
| num_steps: int = 50, |
| guidance_scale: float = 7.0, |
| domain_hint: str = "" |
| ): |
| self.device = device |
| self.output_dir = Path(output_dir) |
| self.output_dir.mkdir(parents=True, exist_ok=True) |
|
|
| self.vlm_client = LocalVLMClient(model=vlm_model) |
| self.money_manager = MoneyManager(model="gpt-4o") |
|
|
| self.gsam_detector = None |
| self.defectfill_generator = None |
|
|
| self.checkpoint_dir = checkpoint_dir |
| self.object_class = object_class |
| self.defect_type = defect_type |
| self.valid_object_classes = valid_object_classes or [] |
| self.valid_defect_types = valid_defect_types or {} |
| self.image_size = image_size |
| self.num_steps = num_steps |
| self.guidance_scale = guidance_scale |
| self.domain_hint = domain_hint |
|
|
| self.rag = get_rag() |
| self.router = get_router() |
|
|
| |
| |
| |
|
|
| def _init_gsam(self): |
| if self.gsam_detector is None or getattr(self.gsam_detector, 'sam_predictor', None) is None: |
| print("[Agent] Initializing GSAM detector...") |
| self.gsam_detector = GSAMDetector( |
| device=self.device, |
| openai_client=self.vlm_client |
| ) |
|
|
| def _init_defectfill(self, object_class: Optional[str] = None, defect_type: Optional[str] = None): |
| obj_cls = object_class or self.object_class |
| dfc_type = defect_type or self.defect_type |
|
|
| if not obj_cls or not dfc_type: |
| raise ValueError("Both object_class and defect_type must be provided.") |
|
|
| ckpt_path = resolve_checkpoint_path(self.checkpoint_dir, obj_cls, dfc_type) |
|
|
| if self.defectfill_generator is None: |
| print(f"[Agent] Initializing DefectFill: {ckpt_path}") |
| config = DefectFillConfig( |
| ckpt_path=ckpt_path, |
| image_size=self.image_size, |
| num_steps=self.num_steps, |
| device=self.device, |
| guidance_scale=self.guidance_scale |
| ) |
| self.defectfill_generator = DefectFillGenerator(config) |
| else: |
| current_ckpt = getattr(self.defectfill_generator, 'ckpt_path', None) |
| if current_ckpt != ckpt_path: |
| print(f"[Agent] Switching checkpoint: {ckpt_path}") |
| self.defectfill_generator.unload_models() |
| config = DefectFillConfig( |
| ckpt_path=ckpt_path, |
| image_size=self.image_size, |
| num_steps=self.num_steps, |
| device=self.device, |
| guidance_scale=self.guidance_scale |
| ) |
| self.defectfill_generator = DefectFillGenerator(config) |
| else: |
| print("[Agent] Reusing DefectFill generator.") |
|
|
| |
| |
| |
|
|
| def plan(self, product_description: str, image: np.ndarray, |
| defect_type: Optional[str] = None, object_class: Optional[str] = None, |
| num_defects: int = 3, domain_hint: Optional[str] = None): |
| """Agent plans defects. VLM selects object_class/defect_type from valid lists if not provided.""" |
| print(f"\n{'='*60}") |
| print("[Agent] Step 1: Planning defects from product description...") |
| print(f"Product: {product_description}") |
| if object_class: |
| print(f"[Agent] User-provided object_class: {object_class}") |
| else: |
| print(f"[Agent] No object_class provided — VLM will select from: {self.valid_object_classes}") |
| if defect_type: |
| print(f"[Agent] User-provided defect_type: {defect_type}") |
| else: |
| print(f"[Agent] No defect_type provided — VLM will select from valid list.") |
|
|
| plan = plan_defects_for_product( |
| self.vlm_client, |
| product_description, |
| image, |
| money_manager=self.money_manager, |
| target_defect_type=defect_type, |
| object_class=object_class, |
| valid_object_classes=self.valid_object_classes if self.valid_object_classes else None, |
| valid_defect_types=self.valid_defect_types if self.valid_defect_types else None, |
| num_defects=num_defects, |
| domain_hint=domain_hint or self.domain_hint |
| ) |
|
|
| if plan is None: |
| raise RuntimeError("Defect planning failed") |
|
|
| print(f"[Agent] Product type: {plan.product_type}") |
| print(f"[Agent] Analysis: {plan.analysis}") |
| print(f"[Agent] Proposed {len(plan.possible_defects)} defects:") |
| for i, d in enumerate(plan.possible_defects, 1): |
| print(f" {i}. [{d.defect_type.upper()}] {d.description}") |
| print(f" Object class: {d.object_class}") |
| print(f" Target: {d.target_entity} / {d.target_subentity or '(whole)'}") |
| print(f" Location: {d.location_hint}") |
| print(f" Coverage: {d.defect_coverage_ratio} | Shape: {d.mask_shape}") |
|
|
| return plan |
|
|
| |
| |
| |
|
|
| def perceive(self, image: np.ndarray, defect_plan): |
| print(f"\n{'='*60}") |
| print("[Agent] Step 2: Directed perception (verification bbox)...") |
|
|
| self._init_gsam() |
|
|
| entity = defect_plan.target_entity |
|
|
| is_solder_target = any(k in entity.lower() for k in ['solder', 'ball', 'bump', 'joint', 'bga', 'die']) |
| is_leg_target = any(k in entity.lower() for k in ['leg', 'pin', 'lead']) |
|
|
| |
| |
| |
| if is_solder_target: |
| print(f"[Agent] Solder-ball target detected. Trying blob-detection fast path...") |
| preds, _, viz = self.gsam_detector.detect_feature_array( |
| image, |
| feature_type="dots", |
| blob_color=0, |
| min_area=5, |
| max_area=300, |
| min_circularity=0.5, |
| min_inertia_ratio=0.1, |
| pad_x=60, |
| pad_y=40, |
| max_y_span=90, |
| min_cluster_size=4, |
| entity_name="solder_ball_array" |
| ) |
| if len(preds) > 0: |
| pred = preds[0] |
| print(f"[Agent] Blob detection succeeded: bbox={preds[0]['bbox']}, " |
| f"area_ratio={pred.get('area_ratio', 0):.4f}") |
| return pred |
|
|
| |
| |
| |
| if is_leg_target: |
| print(f"[Agent] Leg/Pin target detected. Trying blob-detection fast path...") |
| preds, _, viz = self.gsam_detector.detect_feature_array( |
| image, |
| feature_type="lines", |
| blob_color=0, |
| min_area_ratio=0.0005, |
| max_area_ratio=0.2, |
| min_circularity=0.01, |
| min_inertia_ratio=0.4, |
| pad_x=40, |
| pad_y=20, |
| max_y_span=200, |
| min_cluster_size=2, |
| aspect_ratio_range=(1.5, 100.0), |
| vertical_align_threshold=0.15, |
| entity_name="triac_legs", |
| use_adaptive_threshold=True, |
| morph_kernel_size=3 |
| ) |
| if len(preds) > 0: |
| pred = preds[0] |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| print(f"[Agent] Blob detection succeeded: bbox={preds[0]['bbox']}, " |
| f"area_ratio={pred.get('area_ratio', 0):.4f}") |
| return pred |
| print(f"[Agent] Line detection failed, falling back to VLM...") |
|
|
| |
| |
| |
| |
| is_screw_target = ( |
| any(k in entity.lower() for k in ['screw', 'hole', 'fastener', 'mounting']) or |
| any(k in (defect_plan.target_subentity or '').lower() for k in ['screw', 'hole', 'fastener', 'mounting']) or |
| any(k in defect_plan.description.lower() for k in ['screw', 'hole', 'fastener', 'mounting']) or |
| 'screw' in defect_plan.defect_type.lower() |
| ) |
| if is_screw_target: |
| print(f"[Agent] Screw/hole target detected. Trying blob-detection fast path...") |
| |
| |
| |
| if defect_plan.defect_type == "extra_screw": |
| blob_color = 0 |
| elif defect_plan.defect_type == "missing_screw": |
| blob_color = 255 |
| else: |
| blob_color = 0 if "hole" in entity.lower() else 255 |
| preds, _, viz = self.gsam_detector.detect_feature_array( |
| image, |
| feature_type="single_dot", |
| blob_color=blob_color, |
| min_area=10, |
| max_area=200, |
| min_circularity=0.5, |
| min_inertia_ratio=0.1, |
| pad_x=15, |
| pad_y=15, |
| max_y_span=30, |
| min_cluster_size=1, |
| entity_name="screw_target", |
|
|
| location_hint=defect_plan.location_hint |
| ) |
| if len(preds) > 0: |
| best_pred = preds[0] |
| print(f"[Agent] Blob detection succeeded: bbox={best_pred['bbox']}") |
| return best_pred |
|
|
| |
| |
| |
| synonym_map = { |
| "metal can package": ["TO-can", "metal can", "can body", "package body", "metal ring"], |
| "lens cap": ["glass lens cap", "lens", "glass dome", "optical window"], |
| "electrode bars": ["vertical bars", "electrodes", "metal lines"], |
| "die surface": ["die", "chip die", "semiconductor die", "IC die", "black rectangle", "dark square", "central chip"], |
| "solder joint": ["solder ball", "BGA ball", "solder bump", "joint"], |
| "3 pin": ["three metal legs", "vertical metal pins", "metal leads", "power pins"], |
| "3-pin": ["three metal legs", "vertical metal pins", "metal leads", "power pins"], |
| "pin": ["metal leg", "lead", "component terminal"], |
| } |
| search_terms = [entity] + synonym_map.get(entity.lower(), []) |
|
|
| predictions = [] |
| for term in search_terms: |
| |
| is_small_feature = any(k in term.lower() for k in ['hole', 'screw', 'dot', 'pin', 'via']) |
| min_ratio = 0.0001 if is_small_feature else 0.005 |
|
|
| preds, _, viz = self.gsam_detector.detect_parts( |
| image=image, |
| entities=[term], |
| subentities=[defect_plan.target_subentity] if defect_plan.target_subentity else [], |
| entity_subentity_mapping={}, |
| location_hint=defect_plan.location_hint, |
| min_area_ratio=min_ratio, |
| max_area_ratio=0.5, |
| openai_client=self.vlm_client |
| ) |
| if len(preds) > 0: |
| predictions = preds |
| if term != entity: |
| print(f"[Agent] Fallback: detected '{term}' instead of '{entity}'") |
| break |
|
|
| |
| if not predictions and any(k in entity.lower() for k in ['solder', 'ball', 'bump', 'joint']): |
| group_terms = [ |
| "array of solder balls", |
| "BGA ball grid", |
| "group of solder joints", |
| "solder ball array", |
| "BGA array", |
| "ball grid array" |
| ] |
| for term in group_terms: |
| preds, _, viz = self.gsam_detector.detect_parts( |
| image=image, |
| entities=[term], |
| subentities=[], |
| entity_subentity_mapping={}, |
| location_hint=defect_plan.location_hint, |
| min_area_ratio=0.05, |
| max_area_ratio=0.8, |
| openai_client=self.vlm_client |
| ) |
| if len(preds) > 0: |
| predictions = preds |
| print(f"[Agent] Group detection: found '{term}' with {len(preds)} prediction(s)") |
| break |
|
|
| h, w = image.shape[:2] |
|
|
| if not predictions: |
| print(f"[Agent] Warning: No detections for {entity}; using full image") |
|
|
| best_pred = { |
| 'bbox': [0, 0, w, h], |
| 'pred_mask': torch.ones((h, w), dtype=torch.bool) |
| } |
| else: |
| |
| |
| |
| if is_leg_target and len(predictions) >= 2: |
| print(f"[Agent] Merging {len(predictions)} leg detections into group bbox...") |
| all_x1 = [p['bbox'][0] for p in predictions] |
| all_y1 = [p['bbox'][1] for p in predictions] |
| all_x2 = [p['bbox'][2] for p in predictions] |
| all_y2 = [p['bbox'][3] for p in predictions] |
|
|
| pad_x = 15 |
| pad_y = 10 |
|
|
| x1 = max(0, min(all_x1) - pad_x) |
| x2 = min(w, max(all_x2) + pad_x) |
|
|
| |
| y1 = max(0, min(all_y1) - pad_y) |
| y2 = min(h, max(all_y2) + pad_y) |
|
|
| merged_bbox = [x1, y1, x2, y2] |
|
|
| best_pred = { |
| 'bbox': merged_bbox, |
| 'pred_box': torch.tensor(merged_bbox).float(), |
| 'pred_mask': torch.ones((h, w), dtype=torch.bool), |
| 'area_ratio': (x2 - x1) * (y2 - y1) / (h * w), |
| 'entity': entity |
| } |
| print(f"[Agent] Merged leg group bbox: {merged_bbox}") |
| |
| else: |
| best_pred = max(predictions, key=lambda p: p.get('area_ratio', 0)) |
| bbox = best_pred['bbox'] |
|
|
| |
| cx = (bbox[0] + bbox[2]) / 2 |
| cy = (bbox[1] + bbox[3]) / 2 |
| bbox_area = (bbox[2] - bbox[0]) * (bbox[3] - bbox[1]) |
| total_area = h * w |
| |
| on_extreme_edge = (cx < 0.02 * w or cx > 0.98 * w or |
| cy < 0.02 * h or cy > 0.98 * h) |
| too_small = bbox_area < total_area * 0.001 |
| too_large = bbox_area > total_area * 0.35 |
| |
| if (on_extreme_edge and too_small) or too_large: |
| print(f"[Agent] Warning: Bad bbox {bbox}. Using center fallback.") |
| best_pred['bbox'] = [int(0.30 * w), int(0.30 * h), |
| int(0.70 * w), int(0.70 * h)] |
|
|
| |
| bbox = best_pred['bbox'] |
| cx = (bbox[0] + bbox[2]) / 2 |
| cy = (bbox[1] + bbox[3]) / 2 |
| if cx < 0.05 * w or cx > 0.95 * w or cy < 0.05 * h or cy > 0.95 * h: |
| print(f"[Agent] Warning: Bbox {bbox} on extreme edge. Falling back to center.") |
| best_pred['bbox'] = [int(0.20 * w), int(0.20 * h), |
| int(0.80 * w), int(0.80 * h)] |
|
|
| return best_pred |
|
|
| |
| |
| |
|
|
| def prepare_generation_conditions(self, defect_plan) -> Dict: |
| print(f"\n{'='*60}") |
| print("[Agent] Step 3: Preparing DefectFill generation conditions...") |
|
|
| obj_cls = defect_plan.object_class |
| if not obj_cls: |
| raise ValueError("object_class is required for DefectFill prompt standardization.") |
|
|
| defect_class = defect_plan.defect_type |
|
|
| |
| material_hint = "" |
| if defect_class == "missing_screw": |
| material_hint = "by only dark empty threaded hole, black shadowed interior, no screw" |
| elif defect_class == "extra_screw": |
| material_hint = "shiny silver Phillips head screw in empty hole, metallic cross slot, bright reflection" |
| |
| prompt = f"A {obj_cls} with {defect_class} {material_hint}" |
| |
|
|
| print(f"[Agent] Standardized prompt: '{prompt}'") |
| print(f"[Agent] Defect type (checkpoint): {defect_class}") |
| print(f"[Agent] Object class (checkpoint): {obj_cls}") |
|
|
| return { |
| 'prompt': prompt, |
| 'defect_type': defect_class, |
| 'object_class': obj_cls, |
| 'defect_plan': defect_plan |
| } |
|
|
| |
| |
| |
|
|
| def synthesize(self, image_patch: Image.Image, mask_patch: np.ndarray, |
| gen_conditions: Dict) -> Tuple[Image.Image, np.ndarray]: |
| print(f"\n{'='*60}") |
| print("[Agent] Step 4: Synthesizing defect with DefectFill...") |
|
|
| obj_cls = gen_conditions['object_class'] |
| dfc_type = gen_conditions['defect_type'] |
|
|
| self._init_defectfill(object_class=obj_cls, defect_type=dfc_type) |
|
|
| |
| inpainted = self.defectfill_generator.inpaint( |
| image=image_patch, |
| mask=mask_patch, |
| prompt=gen_conditions['prompt'], |
| defect_type=dfc_type, |
| seed=42 |
| ) |
| |
| return inpainted, mask_patch |
|
|
| |
| |
| |
|
|
| def verify(self, original_image: np.ndarray, generated_image: Image.Image, |
| defect_mask: np.ndarray, defect_plan) -> Dict: |
| print(f"\n{'='*60}") |
| print("[Agent] Step 5: Verifying generated defect...") |
|
|
| mask_bool = defect_mask.astype(bool) |
| if mask_bool.sum() == 0: |
| print("[Agent] Warning: Empty mask; verifying full image") |
| y1, x1 = 0, 0 |
| y2, x2 = original_image.shape[0], original_image.shape[1] |
| else: |
| ys, xs = np.where(mask_bool) |
| y1, y2 = ys.min(), ys.max() |
| x1, x2 = xs.min(), xs.max() |
| margin = 64 |
| h, w = original_image.shape[:2] |
| y1 = max(0, y1 - margin) |
| x1 = max(0, x1 - margin) |
| y2 = min(h, y2 + margin) |
| x2 = min(w, x2 + margin) |
|
|
| gen_crop = np.array(generated_image)[y1:y2, x1:x2] |
| orig_crop = original_image[y1:y2, x1:x2] |
|
|
| obj_name = f"a {defect_plan.target_subentity or defect_plan.target_entity}" |
| artifact_type = "addition" |
|
|
| result = artifact_description( |
| self.vlm_client, |
| original_image, |
| orig_crop, |
| gen_crop, |
| obj_name, |
| artifact_type, |
| self.money_manager |
| ) |
|
|
| print(f"[Agent] Verification result: has_artifact={result.has_artifact}") |
| print(f"[Agent] Explanation: {result.explanation}") |
| print(f"[Agent] Label: {result.label}") |
|
|
| return { |
| 'passed': result.has_artifact, |
| 'explanation': result.explanation, |
| 'label': result.label |
| } |
|
|
| |
| |
| |
|
|
| def run(self, product_description: str, image_path: str, |
| caption: Optional[str] = None, num_defects: int = 3, |
| defect_type: Optional[str] = None, |
| object_class: Optional[str] = None) -> Dict: |
| """Run the full agentic pipeline with DefectFill.""" |
| start_time = datetime.now() |
| exp_id = str(uuid.uuid4())[:8] |
|
|
| |
| orig_image = np.array(Image.open(image_path).convert('RGB')) |
| image = np.array(Image.fromarray(orig_image).resize( |
| (self.image_size, self.image_size), Image.LANCZOS |
| )) |
| print(f"[Agent] Loaded image: {orig_image.shape} -> working at {self.image_size}x{self.image_size}") |
|
|
| plan = self.plan( |
| product_description, |
| image, |
| defect_type=defect_type, |
| object_class=object_class, |
| num_defects=num_defects |
| ) |
|
|
| results = [] |
| defects_to_process = plan.possible_defects[:num_defects] |
| print(f"[Agent] Processing top {len(defects_to_process)} of {len(plan.possible_defects)} planned defects") |
|
|
| for i, defect_plan in enumerate(defects_to_process): |
| print(f"\n{'='*80}") |
| print(f"[Agent] Defect {i+1}/{len(defects_to_process)}: [{defect_plan.defect_type.upper()}] {defect_plan.description}") |
| print(f"{'='*80}") |
|
|
| try: |
| prediction = self.perceive(image, defect_plan) |
|
|
| |
| |
| |
| if defect_plan.object_class == "xray_PCB": |
| self._init_gsam() |
| solder_preds, _, _ = self.gsam_detector.detect_feature_array(image) |
| if solder_preds and len(solder_preds) > 0: |
| print(f"[Agent] Forcing xray_PCB type defect to target the 16-dot solder ball array.") |
| prediction = solder_preds[0] |
| else: |
| print(f"[Agent] Warning: Could not detect 16-dot array. Using original detection.") |
| |
|
|
| target_bbox = prediction.get('bbox', [0, 0, image.shape[1], image.shape[0]]) |
|
|
| |
| |
| |
| skip_synthesis = False |
| if defect_plan.defect_type == "extra_screw": |
| self._init_gsam() |
| |
| |
| donor_preds, _, _ = self.gsam_detector.detect_feature_array( |
| image, |
| feature_type="single_dot", |
| blob_color=255, |
| min_area=20, |
| max_area=250, |
| min_circularity=0.5, |
| min_inertia_ratio=0.1, |
| pad_x=20, |
| pad_y=20, |
| min_cluster_size=1, |
| entity_name="donor_screw" |
| ) |
| |
| tx1, ty1, tx2, ty2 = target_bbox |
| tw = max(1, tx2 - tx1) |
| th = max(1, ty2 - ty1) |
| target_area = tw * th |
| |
| |
| donor_bbox = None |
| best_donor_area = 0 |
| for d in donor_preds: |
| dx1, dy1, dx2, dy2 = d['bbox'] |
| |
| if not (dx2 < tx1 or dx1 > tx2 or dy2 < ty1 or dy1 > ty2): |
| continue |
| area = (dx2 - dx1) * (dy2 - dy1) |
| |
|
|
| if area > best_donor_area and area >= target_area * 0.7: |
| best_donor_area = area |
| donor_bbox = [dx1, dy1, dx2, dy2] |
| |
| if donor_bbox is None: |
| print("[Agent] No suitable donor screw found. Falling back to diffusion.") |
| skip_synthesis = False |
| else: |
| dx1, dy1, dx2, dy2 = donor_bbox |
| |
| |
| donor_patch = Image.fromarray(image[dy1:dy2, dx1:dx2]) |
| donor_mask = np.ones((dy2 - dy1, dx2 - dx1), dtype=np.uint8) * 255 |
|
|
| |
| |
| head_scale = 1.4 |
| cx = (tx1 + tx2) // 2 |
| cy = (ty1 + ty2) // 2 |
| half_w = int((tx2 - tx1) * head_scale / 2) |
| half_h = int((ty2 - ty1) * head_scale / 2) |
| expanded_bbox = [ |
| max(0, cx - half_w), |
| max(0, cy - half_h), |
| min(image.shape[1], cx + half_w), |
| min(image.shape[0], cy + half_h) |
| ] |
| |
| |
| blended_rgb, full_mask = blend_defect_onto_real_image( |
| real_image=image, |
| defect_image=donor_patch, |
| defect_mask=donor_mask, |
| target_bbox=expanded_bbox, |
| max_defect_ratio=1.0, |
| mask_shape="circle" |
| ) |
| |
| blended_image = blended_rgb |
| blended_with_green_box = create_visual_prompt_image(blended_image, target_bbox) |
| full_defect_mask = full_mask |
| generated_image = Image.fromarray(blended_image) |
|
|
| skip_synthesis = True |
| print(f"[Agent] extra_screw: Used copy-paste compositing from donor screw at {donor_bbox}") |
|
|
| verification = self.verify(image, Image.fromarray(blended_image), full_defect_mask, defect_plan) |
| |
| |
| |
| |
| |
| if not skip_synthesis: |
| |
| |
| |
|
|
| gen_conditions = self.prepare_generation_conditions(defect_plan) |
|
|
| |
| |
| |
| x1, y1, x2, y2 = target_bbox |
| x1, y1 = max(0, x1), max(0, y1) |
| x2, y2 = min(image.shape[1], x2), min(image.shape[0], y2) |
|
|
| |
| h, w = image.shape[:2] |
| full_mask = np.zeros((h, w), dtype=np.uint8) |
|
|
| |
| mask_shape = "rectangle" |
| coverage_ratio = defect_plan.defect_coverage_ratio |
|
|
| |
| bbox_w = x2 - x1 |
| bbox_h = y2 - y1 |
|
|
| |
| if bbox_w > w * 0.9 and bbox_h > h * 0.9: |
| print("[Agent] Warning: Full-image bbox detected. Using centered quarter region.") |
| bbox_w = int(w * 0.5) |
| bbox_h = int(h * 0.5) |
| x1 = (w - bbox_w) // 2 |
| y1 = (h - bbox_h) // 2 |
| x2 = x1 + bbox_w |
| y2 = y1 + bbox_h |
|
|
| scale_factor = np.sqrt(max(0.05, coverage_ratio)) |
| scaled_w = int(bbox_w * scale_factor) |
| scaled_h = int(bbox_h * scale_factor) |
|
|
| center_x = (x1 + x2) // 2 |
| center_y = (y1 + y2) // 2 |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| if mask_shape == "circle": |
| radius = max(1, min(scaled_w, scaled_h) // 2 - 1) |
| cv2.circle(full_mask, (center_x, center_y), radius, 255, -1) |
| elif mask_shape == "square": |
| side = max(1, min(scaled_w, scaled_h) - 2) |
| top_left_x = center_x - side // 2 |
| top_left_y = center_y - side // 2 |
| cv2.rectangle(full_mask, (top_left_x, top_left_y), (top_left_x + side, top_left_y + side), 255, -1) |
| elif mask_shape == "rectangle": |
| top_left_x = center_x - scaled_w // 2 |
| top_left_y = center_y - scaled_h // 2 |
| |
| entity = defect_plan.target_entity |
| is_leg_target = any(k in entity.lower() for k in ['leg', 'pin', 'lead']) |
| if is_leg_target: |
| |
| |
| |
| top_left_y = y2 - scaled_h |
| top_left_y = max(y1, top_left_y) |
| else: |
| top_left_y = center_y - scaled_h // 2 |
| |
| cv2.rectangle(full_mask, (top_left_x, top_left_y), (top_left_x + scaled_w, top_left_y + scaled_h), 255, -1) |
| else: |
| axes = (max(1, scaled_w // 2), max(1, scaled_h // 2)) |
| cv2.ellipse(full_mask, (center_x, center_y), axes, 0, 0, 360, 255, -1) |
|
|
| |
| y_idx, x_idx = np.where(full_mask > 0) |
| if len(y_idx) > 0: |
| min_y, max_y = np.min(y_idx), np.max(y_idx) |
| min_x, max_x = np.min(x_idx), np.max(x_idx) |
| cy_crop = (min_y + max_y) // 2 |
| cx_crop = (min_x + max_x) // 2 |
| max_dim = max(max_y - min_y, max_x - min_x) |
| else: |
| cy_crop, cx_crop = h // 2, w // 2 |
| max_dim = 0 |
|
|
| padding = 50 |
| crop_size = max(self.image_size, max_dim + padding) |
| |
| |
| |
| |
| |
| |
| |
|
|
| half = crop_size // 2 |
|
|
| x1c = cx_crop - half |
| y1c = cy_crop - half |
| x2c = x1c + crop_size |
| y2c = y1c + crop_size |
|
|
| if x1c < 0: x2c -= x1c; x1c = 0 |
| if y1c < 0: y2c -= y1c; y1c = 0 |
| if x2c > w: x1c -= (x2c - w); x2c = w |
| if y2c > h: y1c -= (y2c - h); y2c = h |
| x1c = max(0, x1c); y1c = max(0, y1c) |
| x2c = min(w, x2c); y2c = min(h, y2c) |
|
|
| crop_img = image[y1c:y2c, x1c:x2c] |
| crop_mask = full_mask[y1c:y2c, x1c:x2c] |
|
|
| if crop_img.shape[0] != self.image_size or crop_img.shape[1] != self.image_size: |
| crop_img = cv2.resize(crop_img, (self.image_size, self.image_size), |
| interpolation=cv2.INTER_AREA) |
| crop_mask = cv2.resize(crop_mask, (self.image_size, self.image_size), |
| interpolation=cv2.INTER_NEAREST) |
|
|
| |
| generated_image, _ = self.synthesize( |
| image_patch=Image.fromarray(crop_img), |
| mask_patch=crop_mask, |
| gen_conditions=gen_conditions |
| ) |
| gen_np = np.array(generated_image) |
|
|
| if defect_plan.defect_type == "missing_screw": |
| gen_np = (gen_np * 0.35).astype(np.uint8) |
| |
| |
| |
| |
| |
|
|
| if defect_plan.defect_type == "extra_screw": |
| |
| gen_np = np.clip(gen_np * 1.2 + 30, 0, 255).astype(np.uint8) |
|
|
| |
| crop_h, crop_w = y2c - y1c, x2c - x1c |
| if gen_np.shape[:2] != (crop_h, crop_w): |
| gen_np = cv2.resize(gen_np, (crop_w, crop_h), interpolation=cv2.INTER_AREA) |
| mask_back = cv2.resize(crop_mask, (crop_w, crop_h), interpolation=cv2.INTER_NEAREST) |
| else: |
| mask_back = crop_mask |
|
|
| |
| orig_crop = image[y1c:y2c, x1c:x2c].copy() |
|
|
| |
| mask_back_uint8 = (mask_back > 0).astype(np.uint8) * 255 |
|
|
| |
| orig_crop_bgr = cv2.cvtColor(orig_crop, cv2.COLOR_RGB2BGR) |
| gen_crop_bgr = cv2.cvtColor(gen_np, cv2.COLOR_RGB2BGR) |
|
|
| |
| patch_mean = np.mean(gen_np) |
| if defect_plan.defect_type == "missing_screw" or defect_plan.defect_type == "extra_screw": |
| clone_mode = cv2.NORMAL_CLONE |
| else: |
| clone_mode = cv2.NORMAL_CLONE if patch_mean < 30 else cv2.MIXED_CLONE |
|
|
| center = (mask_back_uint8.shape[1] // 2, mask_back_uint8.shape[0] // 2) |
| blended_crop_bgr = cv2.seamlessClone( |
| gen_crop_bgr, orig_crop_bgr, mask_back_uint8, center, clone_mode |
| ) |
| blended_crop = cv2.cvtColor(blended_crop_bgr, cv2.COLOR_BGR2RGB) |
|
|
| blended_image = image.copy() |
| blended_image[y1c:y2c, x1c:x2c] = blended_crop |
|
|
| full_defect_mask = np.zeros((h, w), dtype=np.uint8) |
| full_defect_mask[y1c:y2c, x1c:x2c] = mask_back |
|
|
| ys, xs = np.where(full_defect_mask > 0) |
| if len(ys) > 0 and len(xs) > 0: |
| mask_bbox = [int(xs.min()), int(ys.min()), int(xs.max()), int(ys.max())] |
| else: |
| mask_bbox = target_bbox |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| blended_with_green_box = create_visual_prompt_image(blended_image, mask_bbox) |
|
|
| verification = self.verify(image, generated_image, full_defect_mask, defect_plan) |
|
|
| if not verification['passed']: |
| defect_dir = self.output_dir / "failed" / f"{exp_id}_defect_{i}_{defect_plan.defect_type}" |
| else: |
| defect_dir = self.output_dir / f"{exp_id}_defect_{i}_{defect_plan.defect_type}" |
|
|
| defect_dir.mkdir(parents=True, exist_ok=True) |
|
|
| Image.fromarray(image).save(defect_dir / "real_clean_image.png") |
|
|
| if skip_synthesis: |
| Image.fromarray(blended_image).save(defect_dir / "blended_factory_defect.png") |
| else: |
| generated_image.save(defect_dir / "blended_factory_defect.png") |
|
|
| generated_image.save(defect_dir / "raw_defectfill_patch.png") |
|
|
| mask_img = Image.fromarray((full_defect_mask > 0).astype(np.uint8) * 255) |
| mask_img.save(defect_dir / "defect_mask.png") |
|
|
| if not skip_synthesis: |
| gen_conditions_to_save = {k: v for k, v in gen_conditions.items() if k != 'defect_plan'} |
| else: |
| gen_conditions_to_save = { |
| 'prompt': f"A {defect_plan.object_class} with {defect_plan.defect_type} (copy-paste compositing)", |
| 'defect_type': defect_plan.defect_type, |
| 'object_class': defect_plan.object_class |
| } |
|
|
| metadata = { |
| 'experiment_id': exp_id, |
| 'product_description': product_description, |
| 'product_type': plan.product_type, |
| 'object_class': defect_plan.object_class, |
| 'defect_type': defect_plan.defect_type, |
| 'defect_plan': defect_plan.dict() if hasattr(defect_plan, 'dict') else vars(defect_plan), |
| 'generation_conditions': gen_conditions_to_save, |
| 'verification': verification, |
| 'timestamp': datetime.now().isoformat() |
| } |
| with open(defect_dir / "metadata.json", 'w') as f: |
| json.dump(metadata, f, indent=2, default=str) |
|
|
| if not verification['passed']: |
| results.append({ |
| 'defect_type': defect_plan.defect_type, |
| 'object_class': defect_plan.object_class, |
| 'success': False, |
| 'verification_passed': False, |
| 'error': f"VLM verification failed: {verification.get('explanation', 'no explanation')}" |
| }) |
| print(f"[Agent] Verification FAILED for {defect_plan.defect_type}.") |
| else: |
| results.append({ |
| 'defect_type': defect_plan.defect_type, |
| 'object_class': defect_plan.object_class, |
| 'success': True, |
| 'verification_passed': verification['passed'], |
| 'output_dir': str(defect_dir) |
| }) |
| print(f"[Agent] Defect {i+1} complete. Saved to {defect_dir}") |
|
|
| except Exception as e: |
| print(f"[Agent] ERROR processing defect {i+1}: {str(e)}") |
| traceback.print_exc() |
| results.append({ |
| 'defect_type': getattr(defect_plan, 'defect_type', 'unknown'), |
| 'object_class': getattr(defect_plan, 'object_class', 'unknown'), |
| 'success': False, |
| 'error': str(e) |
| }) |
|
|
| if hasattr(self, '_gsam_cache'): |
| keys_to_remove = [k for k in self._gsam_cache if k.startswith(str(image_path) + "::")] |
| for k in keys_to_remove: |
| self._gsam_cache.pop(k, None) |
|
|
| elapsed = (datetime.now() - start_time).total_seconds() |
| print(f"\n{'='*60}") |
| print(f"[Agent] Pipeline complete in {elapsed:.1f}s") |
| print(f"[Agent] Results: {sum(1 for r in results if r['success'])}/{len(results)} succeeded") |
|
|
| return { |
| 'experiment_id': exp_id, |
| 'product_type': plan.product_type, |
| 'results': results, |
| 'output_dir': str(self.output_dir), |
| 'elapsed_time': elapsed |
| } |
|
|
| def cleanup(self): |
| if self.gsam_detector: |
| self.gsam_detector.cleanup() |
| self.gsam_detector = None |
| if self.defectfill_generator: |
| self.defectfill_generator.unload_models() |
| self.defectfill_generator = None |
| print("[Agent] All models cleaned up.") |
|
|
|
|
| |
| |
| |
|
|
| def main(): |
| parser = argparse.ArgumentParser(description='ArtiAgent — DefectFill Edition (with VLM list selection)') |
| parser.add_argument('--product-desc', required=True, help='Product description (drives VLM selection)') |
| parser.add_argument('--image', required=True, help='Path to clean product image') |
|
|
| |
| parser.add_argument('--checkpoint-dir', required=True, |
| help='Root directory containing object_class/defect_type checkpoint subfolders') |
| parser.add_argument('--object-class', default=None, |
| help='Object class (optional; VLM selects from valid list if omitted)') |
| parser.add_argument('--defect-type', default=None, |
| help='Defect type (optional; VLM selects from valid list if omitted)') |
|
|
| |
| parser.add_argument('--valid-object-classes', default=None, |
| help="JSON array of valid object classes, e.g., '[\"xray_PCB\",\"vcsel\"]'" ) |
| parser.add_argument('--valid-defect-types', default=None, |
| help="JSON dict mapping object_class to defect types, e.g., '{\"xray_PCB\":[\"xray_die\",\"bubble\"]}'" ) |
| |
| parser.add_argument('--output-dir', default='./defect_output', help='Output directory') |
| parser.add_argument('--caption', default=None, help='Optional image caption') |
| parser.add_argument('--max-defects', type=int, default=3, help='Max defects to generate') |
| parser.add_argument('--device', default='cuda', help='Device (cuda/cpu)') |
| parser.add_argument('--vlm-model', default='gemma3:12b', help='Local VLM model') |
| parser.add_argument('--image-size', type=int, default=512, help='Generation resolution') |
| parser.add_argument('--num-steps', type=int, default=50, help='Denoising steps') |
| parser.add_argument('--guidance-scale', type=float, default=7.5, help='CFG scale') |
|
|
| args = parser.parse_args() |
|
|
| |
| valid_object_classes = None |
| valid_defect_types = None |
| if args.valid_object_classes: |
| valid_object_classes = json.loads(args.valid_object_classes) |
| if args.valid_defect_types: |
| valid_defect_types = json.loads(args.valid_defect_types) |
|
|
| orchestrator = ArtiAgentOrchestrator( |
| device=args.device, |
| output_dir=args.output_dir, |
| vlm_model=args.vlm_model, |
| checkpoint_dir=args.checkpoint_dir, |
| object_class=args.object_class or "", |
| defect_type=args.defect_type or "", |
| valid_object_classes=valid_object_classes, |
| valid_defect_types=valid_defect_types, |
| image_size=args.image_size, |
| num_steps=args.num_steps, |
| guidance_scale=args.guidance_scale |
| ) |
|
|
| result = orchestrator.run( |
| product_description=args.product_desc, |
| image_path=args.image, |
| caption=args.caption, |
| num_defects=args.num_defects, |
| defect_type=args.defect_type, |
| object_class=args.object_class |
| ) |
|
|
| print(f"\nFinal output saved to: {result['output_dir']}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |