{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"2354f701","cell_type":"markdown","source":"# GroundingSAM Pipeline\n\nForked from [NielsRogge's tutorial](https://github.com/NielsRogge/Transformers-Tutorials/blob/master/Grounding%20DINO/GroundingDINO_with_Segment_Anything.ipynb)","metadata":{"papermill":{"duration":0.007055,"end_time":"2026-05-30T14:55:06.8194","exception":false,"start_time":"2026-05-30T14:55:06.812345","status":"completed"},"tags":[]}},{"id":"05379a09","cell_type":"code","source":"# Set CUDA allocator config BEFORE importing torch so it actually takes effect.\n# Reduces fragmentation, which matters when running SAM3 on a 16 GiB GPU.\n# NOTE: requires kernel restart to take effect after first run.\nimport os\nos.environ.setdefault(\"PYTORCH_CUDA_ALLOC_CONF\", \"expandable_segments:True\")\n\nimport sys, uuid, random, shutil, abc, contextlib, gc, copy, subprocess\nfrom dataclasses import dataclass, field, replace\nfrom pathlib import Path\nfrom typing import Dict, List, Literal, Optional, Tuple, Union\nfrom torchvision.transforms.functional import InterpolationMode\n\nimport cv2\nimport numpy as np\nimport torch\nimport requests\nimport matplotlib.pyplot as plt\nimport plotly.express as px\nimport plotly.graph_objects as go\nimport torchvision.transforms.functional as TF\nimport itertools\nfrom itertools import zip_longest\nfrom PIL import Image\nfrom transformers import pipeline","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"3ed7a352","cell_type":"markdown","source":"## 1. Data Structures","metadata":{"papermill":{"duration":0.004679,"end_time":"2026-05-30T14:55:35.033951","exception":false,"start_time":"2026-05-30T14:55:35.029272","status":"completed"},"tags":[]}},{"id":"af7ee743","cell_type":"code","source":"@dataclass\nclass BoundingBox:\n    xmin: int\n    ymin: int\n    xmax: int\n    ymax: int\n\n    @property\n    def xyxy(self) -> List[float]:\n        return [self.xmin, self.ymin, self.xmax, self.ymax]\n\n    @property\n    def width(self) -> int:\n        return self.xmax - self.xmin\n\n    @property\n    def height(self) -> int:\n        return self.ymax - self.ymin\n\n    @property\n    def area(self) -> int:\n        return self.width * self.height\n\n    @property\n    def center(self) -> Tuple[int, int]:\n        return (self.xmin + self.xmax) // 2, (self.ymin + self.ymax) // 2\n\n\n@dataclass\nclass DetectionResult:\n    score: float\n    label: str\n    box: BoundingBox\n    mask: Optional[np.ndarray] = None\n\n    @classmethod\n    def from_dict(cls, d: Dict) -> \"DetectionResult\":\n        return cls(\n            score=d[\"score\"],\n            label=d[\"label\"],\n            box=BoundingBox(\n                xmin=d[\"box\"][\"xmin\"],\n                ymin=d[\"box\"][\"ymin\"],\n                xmax=d[\"box\"][\"xmax\"],\n                ymax=d[\"box\"][\"ymax\"],\n            ),\n        )\n\n    @property\n    def mask_area(self) -> int:\n        if self.mask is None:\n            return 0\n        return int(np.asarray(self.mask).astype(bool).sum())","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"b5fac5f8","cell_type":"markdown","source":"## 2. Config Dataclasses","metadata":{"papermill":{"duration":0.004351,"end_time":"2026-05-30T14:55:35.063006","exception":false,"start_time":"2026-05-30T14:55:35.058655","status":"completed"},"tags":[]}},{"id":"dbcbd88e","cell_type":"code","source":"@dataclass\nclass DetectorConfig:\n    \"\"\"\n    With SAM3, detection and segmentation are unified.\n    Only `box_threshold` is used (as the SAM3 score threshold for filtering\n    low-confidence detections). Other fields are kept for backward-compat\n    but ignored by GroundingSAMModel.\n    \"\"\"\n    model_id: str = \"facebook/sam3\"    # HF model id\n    box_threshold: float = 0.3         # SAM3 score threshold\n    text_threshold: float = 0.25       # not used\n    batch_size: int = 32               # not used\n\n\n@dataclass\nclass SegmenterConfig:\n    \"\"\"\n    SAM3 from HuggingFace Transformers (facebook/sam3).\n\n    image_size: 1008 is the model-trained resolution (best accuracy).\n                Lower values (560, 336) reduce memory at the cost of accuracy.\n                For T4 / 16 GiB GPUs running tight on memory, try 560 first.\n    mask_threshold: Threshold applied to predicted mask logits before\n                    binarising (HF default is 0.5).\n    \"\"\"\n    model_id: str = \"facebook/sam3\"\n    image_size: int = 1008\n    mask_threshold: float = 0.5\n    polygon_refinement: bool = False\n\n\n@dataclass\nclass CropConfig:\n    \"\"\"\n    policy options\n    --------------\n    sliding_window  Find the square that covers the most mask pixels (integral-image trick).\n    bbox            Crop directly from the detection bounding box, squared up.\n    center_mask     Crop a square centred on the mask centroid.\n    largest_bbox    Pick the detection with the largest bbox, ignore masks.\n\n    size_bias       Controls random size: side = min_win + u^bias * (req_win - min_win).\n                    bias=1 uniform; bias>1 skews toward min_win; bias<1 toward req_win.\n\n    mask_selection  Which detection to use: largest_area | highest_score | first\n    \"\"\"\n    policy: str = \"sliding_window\"\n    min_win: int = 64\n    req_win: int = 224\n    size_bias: float = 1.5\n    mask_selection: str = \"largest_area\"\n\n\n@dataclass\nclass AugmentationConfig:\n    enabled: bool = True\n    rotation: bool = True\n    shearing: bool = True\n    flipping: bool = True\n    brightness: bool = True\n    contrast: bool = True\n    saturation: bool = True\n    hue: bool = True\n    # FIX 3 — ranges ampliati per garantire modifiche visivamente percettibili.\n    # I range precedenti erano troppo stretti (es. hue (-0.1, 0.1) ≈ nessuna variazione).\n    brightness_range: Tuple[float, float] = (0.6, 1.4)\n    contrast_range: Tuple[float, float] = (0.5, 1.5)\n    saturation_range: Tuple[float, float] = (0.5, 1.5)\n    hue_range: Tuple[float, float] = (-0.3, 0.3)\n    shear_range: Tuple[float, float] = (-15, 15)\n\n@dataclass\nclass DatasetConfig:\n    base_path: str = \"/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train\"\n    # Tetto massimo di immagini SORGENTE da ispezionare per concetto.\n    #   None  -> nessun limite: si ispezionano TUTTE le immagini disponibili.\n    #   int   -> tetto esplicito (comportamento legacy con quota bilanciata).\n    # Con None, l'estrazione si ferma SOLO quando:\n    #   (a) sono state ispezionate tutte le immagini di partenza, OPPURE\n    #   (b) si è raggiunto target_images_per_concept.\n    images_per_class: Optional[int] = None\n    target_images_per_concept: int = 120   # target finale per classe di concetto\n    extensions: Tuple[str, ...] = (\".jpg\", \".jpeg\", \".png\", \".bmp\", \".webp\", \".tif\", \".tiff\")\n\n@dataclass\nclass PipelineConfig:\n    detector: DetectorConfig = field(default_factory=DetectorConfig)\n    segmenter: SegmenterConfig = field(default_factory=SegmenterConfig)\n    crop: CropConfig = field(default_factory=CropConfig)\n    augmentation: AugmentationConfig = field(default_factory=AugmentationConfig)\n    dataset: DatasetConfig = field(default_factory=DatasetConfig)\n    device: str = \"auto\"   # \"auto\" -> cuda if available, else cpu\n    crops_path: str = \"/kaggle/working/crops\"\n    augmented_path: str = \"/kaggle/working/augmented_data\"","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"b4785371","cell_type":"markdown","source":"## 3. Low-Level Utils","metadata":{"papermill":{"duration":0.004138,"end_time":"2026-05-30T14:55:35.094927","exception":false,"start_time":"2026-05-30T14:55:35.090789","status":"completed"},"tags":[]}},{"id":"9570c46e","cell_type":"code","source":"def load_image(image_str: str) -> Image.Image:\n    if image_str.startswith(\"http\"):\n        return Image.open(requests.get(image_str, stream=True).raw).convert(\"RGB\")\n    return Image.open(image_str).convert(\"RGB\")\n\n\ndef get_boxes(results: List[DetectionResult]) -> List[List[List[float]]]:\n    \"\"\"Boxes in the nested format expected by the SAM processor.\"\"\"\n    return [[r.box.xyxy for r in results]]\n\n\ndef mask_to_polygon(mask: np.ndarray) -> List[List[int]]:\n    contours, _ = cv2.findContours(mask.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    return max(contours, key=cv2.contourArea).reshape(-1, 2).tolist()\n\n\ndef polygon_to_mask(polygon: List[Tuple[int, int]], image_shape: Tuple[int, int]) -> np.ndarray:\n    mask = np.zeros(image_shape, dtype=np.uint8)\n    cv2.fillPoly(mask, [np.array(polygon, dtype=np.int32)], color=(255,))\n    return mask\n\n\ndef refine_masks(masks: torch.BoolTensor, polygon_refinement: bool = False) -> List[np.ndarray]:\n    masks = masks.cpu().float().permute(0, 2, 3, 1).mean(axis=-1)\n    masks = (masks > 0).int().numpy().astype(np.uint8)\n    masks = list(masks)\n    if polygon_refinement:\n        for i, mask in enumerate(masks):\n            masks[i] = polygon_to_mask(mask_to_polygon(mask), mask.shape)\n    return masks\n\n\ndef resolve_device(device: str) -> str:\n    return (\"cuda\" if torch.cuda.is_available() else \"cpu\") if device == \"auto\" else device\n\n\ndef _resolve_output_path(out_path: str) -> str:\n    if out_path.lower().endswith((\".png\", \".jpg\", \".jpeg\", \".webp\", \".bmp\", \".tiff\")):\n        os.makedirs(os.path.dirname(out_path) or \".\", exist_ok=True)\n        return out_path\n    os.makedirs(out_path, exist_ok=True)\n    return os.path.join(out_path, f\"crop_{uuid.uuid4().hex[:8]}.png\")\n\n\ndef _pick_best_detection(\n    detections: List[DetectionResult], strategy: str\n) -> Optional[DetectionResult]:\n    valid = [d for d in detections if d.mask is not None and d.mask_area > 0]\n    if not valid:\n        return None\n    if strategy == \"highest_score\":\n        return max(valid, key=lambda d: d.score)\n    if strategy == \"first\":\n        return valid[0]\n    return max(valid, key=lambda d: d.mask_area)  # \"largest_area\" (default)","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"5f60a303","cell_type":"markdown","source":"## 4. Crop Policies\n\nAvailable: `sliding_window`, `bbox`, `center_mask`, `largest_bbox`\n\nAll policies share the same interface — swap `CropConfig.policy` with no other changes.","metadata":{"papermill":{"duration":0.004268,"end_time":"2026-05-30T14:55:35.124703","exception":false,"start_time":"2026-05-30T14:55:35.120435","status":"completed"},"tags":[]}},{"id":"d58ad326","cell_type":"code","source":"class CropPolicyBase(abc.ABC):\n    @abc.abstractmethod\n    def crop(self, image: np.ndarray, detection: DetectionResult, config: CropConfig) -> Tuple[int, int, int]:\n        \"\"\"Return (x1, y1, side).\"\"\"\n\n    def execute(\n        self,\n        image: np.ndarray,\n        detections: List[DetectionResult],\n        config: CropConfig,\n        out_path: str,\n    ) -> Tuple[str, Tuple[int, int, int]]:\n        best = _pick_best_detection(detections, config.mask_selection)\n        if best is None:\n            raise ValueError(\"No valid detections with non-empty masks.\")\n        x1, y1, side = self.crop(image, best, config)\n        img_pil = Image.fromarray(image)\n        crop = img_pil.crop((x1, y1, x1 + side, y1 + side))\n        if side != config.req_win:\n            crop = crop.resize((config.req_win, config.req_win), resample=Image.BICUBIC)\n        fn = _resolve_output_path(out_path)\n        crop.save(fn)\n        return fn, (x1, y1, side)\n\n\nclass SlidingWindowCropPolicy(CropPolicyBase):\n    \"\"\"Integral-image sliding window that maximises mask coverage.\"\"\"\n\n    @staticmethod\n    def _get_window(mask: np.ndarray, win: int) -> Tuple[int, int, int]:\n        m = np.asarray(mask).astype(np.uint8)\n        H, W = m.shape\n        win = min(win, H, W)\n        ii = np.pad(m, ((1, 0), (1, 0)), mode=\"constant\").cumsum(0).cumsum(1)\n        sums = ii[win:, win:] - ii[:-win, win:] - ii[win:, :-win] + ii[:-win, :-win]\n        maxv = int(sums.max())\n        ys, xs = np.where(sums == maxv)\n        k = np.random.randint(len(xs))\n        return int(xs[k]), int(ys[k]), maxv\n\n    def crop(self, image, detection, config):\n        H, W = image.shape[:2]\n        upper = min(config.req_win, H, W)\n        lower = min(config.min_win, upper)\n        t = random.random() ** config.size_bias\n        side = int(round(lower + t * (upper - lower)))\n        x1, y1, _ = self._get_window(detection.mask, side)\n        return x1, y1, side\n\n\nclass BBoxCropPolicy(CropPolicyBase):\n    \"\"\"Crop the detection's bounding box, squared up and clamped.\"\"\"\n\n    def crop(self, image, detection, config):\n        H, W = image.shape[:2]\n        box = detection.box\n        cx, cy = box.center\n        side = max(min(max(box.width, box.height), H, W), config.min_win)\n        x1 = max(0, min(cx - side // 2, W - side))\n        y1 = max(0, min(cy - side // 2, H - side))\n        return x1, y1, side\n\n\nclass CenterMaskCropPolicy(CropPolicyBase):\n    \"\"\"Crop a square centred on the mask centroid (falls back to bbox centre).\"\"\"\n\n    def crop(self, image, detection, config):\n        H, W = image.shape[:2]\n        side = min(config.req_win, H, W)\n        if detection.mask is not None and detection.mask_area > 0:\n            ys, xs = np.where(np.asarray(detection.mask).astype(bool))\n            cx, cy = int(xs.mean()), int(ys.mean())\n        else:\n            cx, cy = detection.box.center\n        x1 = max(0, min(cx - side // 2, W - side))\n        y1 = max(0, min(cy - side // 2, H - side))\n        return x1, y1, side\n\n\nclass LargestBBoxCropPolicy(CropPolicyBase):\n    \"\"\"Pick detection with the largest bbox area; masks not required.\"\"\"\n\n    def execute(self, image, detections, config, out_path):\n        if not detections:\n            raise ValueError(\"No detections provided.\")\n        best = max(detections, key=lambda d: d.box.area)\n        x1, y1, side = self.crop(image, best, config)\n        img_pil = Image.fromarray(image)\n        crop = img_pil.crop((x1, y1, x1 + side, y1 + side))\n        if side != config.req_win:\n            crop = crop.resize((config.req_win, config.req_win), resample=Image.BICUBIC)\n        fn = _resolve_output_path(out_path)\n        crop.save(fn)\n        return fn, (x1, y1, side)\n\n    def crop(self, image, detection, config):\n        H, W = image.shape[:2]\n        cx, cy = detection.box.center\n        side = max(min(max(detection.box.width, detection.box.height), H, W), config.min_win)\n        x1 = max(0, min(cx - side // 2, W - side))\n        y1 = max(0, min(cy - side // 2, H - side))\n        return x1, y1, side\n\n\n_CROP_POLICY_REGISTRY: Dict[str, type] = {\n    \"sliding_window\": SlidingWindowCropPolicy,\n    \"bbox\": BBoxCropPolicy,\n    \"center_mask\": CenterMaskCropPolicy,\n    \"largest_bbox\": LargestBBoxCropPolicy,\n}\n\n\nclass Cropper:\n    \"\"\"Delegates cropping to the policy named in CropConfig.policy.\"\"\"\n\n    def __init__(self, config: CropConfig):\n        self.config = config\n        cls = _CROP_POLICY_REGISTRY.get(config.policy)\n        if cls is None:\n            raise ValueError(f\"Unknown crop policy '{config.policy}'. Available: {list(_CROP_POLICY_REGISTRY)}\")\n        self._policy: CropPolicyBase = cls()\n\n    def crop(self, image: np.ndarray, detections: List[DetectionResult], out_path: str):\n        return self._policy.execute(image, detections, self.config, out_path)","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"6c96dec2","cell_type":"markdown","source":"## 5. Augmenter","metadata":{"papermill":{"duration":0.004742,"end_time":"2026-05-30T14:55:35.167172","exception":false,"start_time":"2026-05-30T14:55:35.16243","status":"completed"},"tags":[]}},{"id":"701df89d","cell_type":"code","source":"class Augmenter:\n    def __init__(self, config):\n        self.cfg = config\n\n    def augment_one(self, img: Image.Image) -> Image.Image:\n        original_size = img.size  # (width, height)\n        # FIX 3 — snapshot per il safety-check finale\n        original_array = np.array(img)\n\n        # ------------------------\n        # 1. ROTATION (restricted)\n        # ------------------------\n        if self.cfg.rotation:\n            angle = random.choice([90, 180, -90, -180])\n            img = TF.rotate(img, angle)\n\n        # ------------------------\n        # 2. SHEARING\n        # ------------------------\n        if self.cfg.shearing:\n            shear_x = random.uniform(*self.cfg.shear_range)\n            img = TF.affine(\n                img,\n                angle=0.0,\n                translate=[0, 0],\n                scale=1.0,\n                shear=[shear_x, 0.0],\n                interpolation=InterpolationMode.BILINEAR\n            )\n\n        # ------------------------\n        # 3. CROP + RESIZE\n        # ------------------------\n        width, height = img.size\n        crop_h = int(height * 0.75)\n        crop_w = int(width * 0.75)\n\n        img = TF.center_crop(img, (crop_h, crop_w))\n        img = TF.resize(img, (original_size[1], original_size[0]))\n\n        # ------------------------\n        # 4. FLIPPING\n        # ------------------------\n        if self.cfg.flipping:\n            flip = random.choice([\"h\", \"v\", \"hv\", \"none\"])\n            if flip in (\"h\", \"hv\"):\n                img = TF.hflip(img)\n            if flip in (\"v\", \"hv\"):\n                img = TF.vflip(img)\n\n        # ------------------------\n        # 5. COLOR AUGMENTATION\n        # ------------------------\n\n        if self.cfg.brightness:\n            factor = random.uniform(*self.cfg.brightness_range)\n            img = TF.adjust_brightness(img, factor)\n\n        if self.cfg.contrast:\n            factor = random.uniform(*self.cfg.contrast_range)\n            img = TF.adjust_contrast(img, factor)\n\n        if self.cfg.saturation:\n            factor = random.uniform(*self.cfg.saturation_range)\n            img = TF.adjust_saturation(img, factor)\n\n        if self.cfg.hue:\n            # NOTE: hue expects range roughly [-0.5, 0.5]\n            delta = random.uniform(*self.cfg.hue_range)\n            img = TF.adjust_hue(img, delta)\n\n        # FIX 3 — safety-check: se per qualche motivo l'output è pixel-identico\n        # all'input (es. rotation=180 + flip \"none\" + color factors ≈ 1.0),\n        # forziamo almeno un flip orizzontale per garantire la differenza.\n        if np.array_equal(original_array, np.array(img)):\n            img = TF.hflip(img)\n\n        return img","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"78077193","cell_type":"markdown","source":"# 6. SAM3 Setup","metadata":{"papermill":{"duration":0.004888,"end_time":"2026-05-30T14:55:35.198968","exception":false,"start_time":"2026-05-30T14:55:35.19408","status":"completed"},"tags":[]}},{"id":"fc87463f","cell_type":"code","source":"# Install SAM3 from facebookresearch/sam3 (the official Meta repo).\n# We use the source install because the HuggingFace Transformers wrapper for\n# SAM3 in late 2025 / early 2026 has an unresolved checkpoint key-mapping bug\n# that ships the text encoder with random weights — every text prompt returns\n# garbage scores. The source path uses the original checkpoint format directly\n# and avoids that issue entirely.\n#\n# Pre-flight checklist:\n#  - Accept the license at https://huggingface.co/facebook/sam3\n#  - Add an HF token to Kaggle Secrets under the key \"HF_TOKEN\"\n#\n# This cell is idempotent: safe to re-run, will skip steps already done.\n\nSAM3_DIR = \"/kaggle/working/sam3\"\n\n# Step 1: clone the repo if it's not already there (Kaggle's /kaggle/working\n# is persistent across kernel restarts within a session, but you should re-run\n# this cell after restarting in case site-packages got wiped).\nif not os.path.isdir(SAM3_DIR):\n    print(f\"Cloning sam3 to {SAM3_DIR}...\")\n    subprocess.run(\n        [\"git\", \"clone\", \"--depth=1\",\n         \"https://github.com/facebookresearch/sam3.git\", SAM3_DIR],\n        check=True,\n    )\nelse:\n    print(f\"sam3 already cloned at {SAM3_DIR}\")\n\n# Step 2: editable install. --no-deps avoids pip trying to install PyTorch\n# from scratch (Kaggle already ships a working version) and --ignore-requires-python\n# lets us run on Python 3.11 (sam3 advertises 3.12+ but the code works on 3.11).\nsubprocess.run(\n    [sys.executable, \"-m\", \"pip\", \"install\", \"--quiet\", \"--no-deps\",\n     \"--ignore-requires-python\", \"-e\", SAM3_DIR],\n    check=True,\n)\n\n# Step 3: install the small runtime deps sam3 actually uses at import time.\nsubprocess.run(\n    [sys.executable, \"-m\", \"pip\", \"install\", \"--quiet\",\n     \"ftfy\", \"iopath>=0.1.10\", \"decord\", \"ipycanvas\"],\n    check=True,\n)\n\n# Step 4: belt-and-braces sys.path injection, so that even if the editable\n# install registration was lost across a kernel restart, the imports below\n# still succeed.\nif SAM3_DIR not in sys.path:\n    sys.path.insert(0, SAM3_DIR)\n\n# Step 5: import and verify.\nfrom sam3.model_builder import build_sam3_image_model\nfrom sam3.model.sam3_image_processor import Sam3Processor\nimport sam3 as _sam3_pkg\nprint(f\"sam3 imported from: {_sam3_pkg.__file__}\")\n\n# Step 6: HuggingFace auth so the gated checkpoint can be downloaded.\nfrom huggingface_hub import login\nfrom kaggle_secrets import UserSecretsClient\n\nsecrets = UserSecretsClient()\nhf_token = secrets.get_secret(\"HF_TOKEN\")\nlogin(token=hf_token)\n\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM total: {torch.cuda.get_device_properties(0).total_memory / 1024**3:.2f} GiB\")","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"d7e262ae","cell_type":"markdown","source":"## 6. GroundingSAM Model","metadata":{"papermill":{"duration":0.004698,"end_time":"2026-05-30T14:55:57.183024","exception":false,"start_time":"2026-05-30T14:55:57.178326","status":"completed"},"tags":[]}},{"id":"770170a7","cell_type":"code","source":"class GroundingSAMModel:\n    \"\"\"\n    SAM3-based Promptable Concept Segmentation using the official\n    facebookresearch/sam3 source distribution.\n\n    Why this path (not HF Transformers):\n        The HF Transformers wrapper for SAM3 in late 2025 had a checkpoint key\n        mismatch that left the text encoder with random weights — every text\n        prompt produced garbage scores. The source install avoids that entirely.\n\n    Memory strategy on T4 (16 GiB):\n        1. Cast the model to float16 with `.half()` after build. Weights drop\n           from ~3.4 GiB (fp32) to ~1.7 GiB (fp16); activations of the ViT-L at\n           1008x1008 also halve.\n        2. Wrap inference in `torch.amp.autocast(bfloat16)` so any internal\n           ops that need a wider dtype still work.\n        3. Call `set_image` once per text prompt — required because the SAM3\n           processor mutates `state[\"backbone_out\"]` and `set_text_prompt`\n           refuses to run a second time on the same state. We tested keeping\n           state across prompts; the source code explicitly disallows it.\n        4. Free CUDA cache after every label and every image to avoid the\n           per-iteration ViT activations accumulating.\n\n    The class keeps the original detect / segment / process interface so the\n    rest of the pipeline (Cropper, Pipeline.run, etc.) works unchanged.\n    \"\"\"\n\n    def __init__(self, config: PipelineConfig):\n        self.cfg = config\n        self.device = resolve_device(config.device)\n\n        # build_sam3_image_model downloads the checkpoint from HuggingFace\n        # (facebook/sam3) on first call and caches it. Requires HF login.\n        sam3_model = build_sam3_image_model(device=self.device)\n\n        # Convert to float16 on CUDA: ~halves both weights and activations.\n        # On CPU we stay at fp32 because fp16 CPU ops are slow.\n        if self.device == \"cuda\":\n            sam3_model = sam3_model\n            self.dtype = torch.float16\n        else:\n            self.dtype = torch.float32\n\n        sam3_model.eval()\n        self.processor = Sam3Processor(sam3_model)\n\n    def _autocast_ctx(self):\n        \"\"\"Use autocast to handle any internal ops that need mixed precision.\"\"\"\n        if self.device == \"cuda\":\n            return torch.amp.autocast(device_type=\"cuda\", dtype=self.dtype)\n        return contextlib.nullcontext()\n\n    def _free_gpu(self) -> None:\n        if self.device == \"cuda\":\n            gc.collect()\n            torch.cuda.empty_cache()\n\n    @staticmethod\n    def _normalise_label(label: str) -> str:\n        return label.rstrip(\".\").strip()\n\n    @staticmethod\n    def _to_2d_mask(mask: np.ndarray, target_hw: Tuple[int, int]) -> np.ndarray:\n        m = np.asarray(mask)\n        while m.ndim > 2:\n            m = m.squeeze(0)\n        m = (m > 0).astype(np.uint8)\n        H, W = target_hw\n        if m.shape != (H, W):\n            m = cv2.resize(m, (W, H), interpolation=cv2.INTER_NEAREST)\n        return m\n\n    def _run_sam3(\n        self, image: Image.Image, label: str\n    ) -> Optional[Tuple[np.ndarray, np.ndarray, np.ndarray]]:\n        \"\"\"\n        Run one SAM3 forward pass. Returns (boxes, masks, scores) as numpy.\n        Returns None on error. Always frees GPU memory before returning.\n        \"\"\"\n        state = None\n        result = None\n        try:\n            with torch.inference_mode(), self._autocast_ctx():\n                state = self.processor.set_image(image)\n                result = self.processor.set_text_prompt(state=state, prompt=label)\n                boxes  = result[\"boxes\"].detach().float().cpu().numpy()\n                masks  = result[\"masks\"].detach().float().cpu().numpy()\n                scores = result[\"scores\"].detach().float().cpu().numpy()\n            return boxes, masks, scores\n        except torch.cuda.OutOfMemoryError as e:\n            print(f\"  [SAM3 OOM] '{label}': {e}\")\n            return None\n        except Exception as e:\n            print(f\"  [SAM3 prompt error] '{label}': {type(e).__name__}: {e}\")\n            return None\n        finally:\n            del state, result\n            self._free_gpu()\n\n    def detect(\n        self,\n        image: Image.Image,\n        labels: List[str],\n        threshold: Optional[float] = None,\n    ) -> List[DetectionResult]:\n        threshold = threshold if threshold is not None else self.cfg.detector.box_threshold\n        img_W, img_H = image.size\n        detections: List[DetectionResult] = []\n\n        for raw_label in labels:\n            label = self._normalise_label(raw_label)\n            out = self._run_sam3(image, label)\n            if out is None:\n                continue\n            boxes, masks, scores = out\n\n            for box, mask, score in zip(boxes, masks, scores):\n                if float(score) < threshold:\n                    continue\n                x1 = max(0, min(int(round(float(box[0]))), img_W - 1))\n                y1 = max(0, min(int(round(float(box[1]))), img_H - 1))\n                x2 = max(0, min(int(round(float(box[2]))), img_W))\n                y2 = max(0, min(int(round(float(box[3]))), img_H))\n                if x2 <= x1 or y2 <= y1:\n                    continue\n\n                m_2d = self._to_2d_mask(mask, (img_H, img_W))\n                if self.cfg.segmenter.polygon_refinement and m_2d.sum() > 0:\n                    m_2d = polygon_to_mask(mask_to_polygon(m_2d), m_2d.shape)\n\n                detections.append(DetectionResult(\n                    score=float(score),\n                    label=label,\n                    box=BoundingBox(xmin=x1, ymin=y1, xmax=x2, ymax=y2),\n                    mask=m_2d,\n                ))\n\n        return detections\n\n    def segment(\n        self,\n        image: Image.Image,\n        detections: List[DetectionResult],\n    ) -> List[DetectionResult]:\n        \"\"\"No-op for SAM3: detect() already returns masks.\"\"\"\n        return detections\n\n    def process(\n        self,\n        image: Union[Image.Image, str],\n        labels: List[str],\n        threshold: Optional[float] = None,\n    ) -> Tuple[np.ndarray, List[DetectionResult]]:\n        if isinstance(image, str):\n            image = load_image(image)\n        elif not isinstance(image, Image.Image):\n            raise TypeError(f\"Expected PIL.Image or str, got {type(image)}\")\n        detections = self.detect(image, labels, threshold)\n        detections = self.segment(image, detections)\n        return np.array(image), detections","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"37e9e0aa","cell_type":"markdown","source":"## 7. Dataset Collector","metadata":{"papermill":{"duration":0.004414,"end_time":"2026-05-30T14:55:57.220258","exception":false,"start_time":"2026-05-30T14:55:57.215844","status":"completed"},"tags":[]}},{"id":"6ecf0dff","cell_type":"code","source":"class DatasetCollector:\n    \"\"\"\n    Collects image paths from an ImageNet-style tree.\n    class_mapping: {concept: [synset_ids]}\n\n    Quando un concetto è mappato a più synset (es. case \"correlated\" o\n    \"random\"), le immagini vengono prese **da tutte le classi** in modo\n    bilanciato:\n      - quota per-synset equa (budget // N, con resto sui primi synset);\n      - redistribuzione del deficit: se un synset ha meno immagini della\n        sua quota, le mancanti vengono prese dai synset con surplus;\n      - interleave round-robin sull\\'output, così che il loop di\n        estrazione tocchi tutte le classi anche se non scorre l\\'intera\n        lista.\n    \"\"\"\n\n    def __init__(self, config: DatasetConfig, class_mapping: Dict[str, List[str]]):\n        self.cfg = config\n        self.class_mapping = class_mapping\n\n    def collect(self) -> Dict[str, List[str]]:\n        ext_lower = tuple(e.lower() for e in self.cfg.extensions)\n        result: Dict[str, List[str]] = {}\n\n        for concept, synsets in self.class_mapping.items():\n            if not synsets:\n                result[concept] = []\n                print(f\"[DatasetCollector] {concept}: 0 images (no synsets).\")\n                continue\n\n            # 1) Raccogli tutte le immagini disponibili PER synset\n            per_synset: List[List[str]] = []\n            for synset in synsets:\n                input_dir = os.path.join(self.cfg.base_path, synset)\n                if not os.path.isdir(input_dir):\n                    print(f\"[DatasetCollector] Warning: {input_dir} not found.\")\n                    per_synset.append([])\n                    continue\n                \n                found: List[str] = []\n                for dp, _, fns in os.walk(input_dir):\n                    # Rimuoviamo il sorted() e raccogliamo i path\n                    for fn in fns: \n                        if fn.lower().endswith(ext_lower):\n                            found.append(os.path.join(dp, fn))\n                \n                # --- LA MODIFICA CHIAVE ---\n                # Mescoliamo le immagini all'interno del singolo synset\n                random.shuffle(found) \n                per_synset.append(found)\n\n            # 2) Selezione delle immagini per synset.\n            budget = self.cfg.images_per_class\n            n = len(synsets)\n\n            if budget is None:\n                # NESSUN tetto: prendiamo TUTTE le immagini disponibili per ogni\n                # synset. In questo modo il loop di estrazione a valle può ispezionare\n                # l'intero pool e fermarsi SOLO quando (a) tutte le immagini di\n                # partenza sono state ispezionate, oppure (b) è stato raggiunto il\n                # target_images_per_concept. Nessuna immagine sorgente viene scartata\n                # a priori.\n                taken: List[List[str]] = per_synset\n            else:\n                # Tetto esplicito (comportamento legacy): quota equa per synset...\n                base_q, extra = divmod(budget, n)\n                quotas = [base_q + (1 if i < extra else 0) for i in range(n)]\n\n                taken = [paths[:q] for paths, q in zip(per_synset, quotas)]\n                surplus: List[List[str]] = [paths[q:] for paths, q in zip(per_synset, quotas)]\n\n                # ...con redistribuzione del deficit ai synset con surplus (round-robin).\n                deficit = budget - sum(len(t) for t in taken)\n                i = 0\n                safety = budget * n + 1\n                while deficit > 0 and any(surplus) and safety > 0:\n                    if surplus[i % n]:\n                        taken[i % n].append(surplus[i % n].pop(0))\n                        deficit -= 1\n                    i += 1\n                    safety -= 1\n\n            # 3) Interleave round-robin tra i synset: garantisce varietà e\n            # bilanciamento nell'ordine di processing, così che — anche se il loop\n            # di estrazione si ferma a target raggiunto — le immagini provengano\n            # da tutte le classi e non solo dalle prime.\n            interleaved: List[str] = []\n            for tup in zip_longest(*taken, fillvalue=None):\n                for p in tup:\n                    if p is not None:\n                        interleaved.append(p)\n\n            result[concept] = interleaved\n            breakdown = \", \".join(f\"{s}={len(t)}\" for s, t in zip(synsets, taken))\n            print(f\"[DatasetCollector] {concept}: {len(interleaved)} images ({breakdown})\")\n\n        return result\n","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"d718144c","cell_type":"markdown","source":"## 8. Main Pipeline","metadata":{"papermill":{"duration":0.00461,"end_time":"2026-05-30T14:55:57.252097","exception":false,"start_time":"2026-05-30T14:55:57.247487","status":"completed"},"tags":[]}},{"id":"da7ca7d7","cell_type":"code","source":"class GroundingSAMPipeline:\n    \"\"\"\n    End-to-end pipeline: load -> detect (SAM3) -> crop -> augment.\n\n    The expensive SAM3 model is loaded once in __init__ and reused across all\n    experiments. Crop policy and augmentation toggle can be overridden per-call\n    via `run_experiment()`, so multiple datasets (ablation studies) can be\n    produced from one pipeline instance without reloading the model.\n\n    Usage\n    -----\n        pipe = GroundingSAMPipeline(config)\n\n        # Vanilla run with config defaults:\n        pipe.run_experiment(\"vanilla\", CLASS_MAPPING_BASE, LABELS_BASE)\n\n        # Override crop policy:\n        pipe.run_experiment(\"crop_bbox\", CLASS_MAPPING_BASE, LABELS_BASE,\n                            crop_policy=\"bbox\")\n\n        # Turn augmentation on/off:\n        pipe.run_experiment(\"no_aug\", ..., augmentation_enabled=False)\n    \"\"\"\n\n    def __init__(self, config: PipelineConfig):\n        self.config = config\n        self.model = GroundingSAMModel(config)\n        # No pre-built cropper/augmenter; rebuilt per experiment if overridden.\n        self._default_cropper = Cropper(config.crop)\n        self._default_augmenter = Augmenter(config.augmentation)\n\n    # ---------- helpers ----------\n\n    def _reset_output_dirs(self, *dirs: str) -> None:\n        for d in dirs:\n            if os.path.exists(d):\n                shutil.rmtree(d)\n            os.makedirs(d, exist_ok=True)\n\n    def _build_cropper(self, crop_policy: Optional[str]) -> \"Cropper\":\n        if crop_policy is None or crop_policy == self.config.crop.policy:\n            return self._default_cropper\n        crop_cfg = replace(self.config.crop, policy=crop_policy)\n        return Cropper(crop_cfg)\n\n    def _build_augmenter(self, augmentation_enabled: Optional[bool]) -> \"Augmenter\":\n        if augmentation_enabled is None or augmentation_enabled == self.config.augmentation.enabled:\n            return self._default_augmenter\n        aug_cfg = replace(self.config.augmentation, enabled=augmentation_enabled)\n        return Augmenter(aug_cfg)\n\n    def _process_image(\n        self,\n        image_path: str,\n        labels: List[str],\n        out_dir: str,\n        cropper: \"Cropper\",\n        threshold: Optional[float] = None,\n    ) -> Optional[Tuple[str, Tuple[int, int, int]]]:\n        try:\n            image_array, detections = self.model.process(image_path, labels, threshold)\n        except Exception as e:\n            print(f\"  [model error] {image_path}: {e}\")\n            return None\n        if not detections:\n            return None\n        try:\n            return cropper.crop(image_array, detections, out_dir)\n        except ValueError as e:\n            print(f\"  [crop error] {image_path}: {e}\")\n            return None\n\n    # ---------- main entry point ----------\n\n    def run_experiment(\n        self,\n        name: str,\n        class_mapping: Dict[str, List[str]],\n        labels: Dict[str, List[str]],\n        crop_policy: Optional[str] = None,\n        augmentation_enabled: Optional[bool] = None,\n        threshold: Optional[float] = None,\n    ) -> Dict[str, List[str]]:\n        \"\"\"\n        Run the pipeline as one named experiment. Outputs go to:\n          - {crops_path}/{name}/{concept}/...\n          - {augmented_path}/{name}/{concept}/...\n\n        Parameters\n        ----------\n        name : str\n            Subdirectory name and label for this run. Use it to keep\n            ablation outputs separate (e.g., \"crop_bbox\", \"no_aug\", \"prompt_v2\").\n        class_mapping : Dict[str, List[str]]\n            concept -> [ImageNet wnid, ...]. Defines which source images to load.\n        labels : Dict[str, List[str]]\n            concept -> [text prompts to feed SAM3]. Multiple prompts per concept\n            are tried independently; their detections are merged.\n        crop_policy : Optional[str]\n            If given, overrides config.crop.policy for this run only.\n            One of \"sliding_window\", \"bbox\", \"center_mask\", \"largest_bbox\".\n        augmentation_enabled : Optional[bool]\n            If given, overrides config.augmentation.enabled for this run only.\n        threshold : Optional[float]\n            SAM3 score threshold. If None, uses config.detector.box_threshold.\n\n        Returns\n        -------\n        Dict[str, List[str]] : concept -> list of final image paths\n            (crops + augmentations combined, capped at target_images_per_concept).\n        \"\"\"\n        cropper   = self._build_cropper(crop_policy)\n        augmenter = self._build_augmenter(augmentation_enabled)\n        aug_on    = augmenter.cfg.enabled\n\n        collector = DatasetCollector(self.config.dataset, class_mapping)\n        image_paths = collector.collect()\n        target = self.config.dataset.target_images_per_concept\n\n        print(f\"\\n=== Experiment: {name} ===\")\n        print(f\"    crop_policy={cropper.config.policy}  augmentation={aug_on}  \"\n              f\"threshold={threshold if threshold is not None else self.config.detector.box_threshold}\")\n\n        all_crops: Dict[str, List[str]] = {}\n\n        for concept, paths in image_paths.items():\n            concept_labels = labels.get(concept, [concept])\n            crop_dir = os.path.join(self.config.crops_path,     name, concept)\n            augm_dir = os.path.join(self.config.augmented_path, name, concept)\n            self._reset_output_dirs(crop_dir, augm_dir)\n\n            # Fase 1: estrai i crop dalle immagini sorgente.\n            # Ci si ferma SOLO quando una delle due condizioni è vera:\n            #   (a) si è raggiunto il target (len(saved) >= target), OPPURE\n            #   (b) sono state ispezionate TUTTE le immagini di partenza\n            #       (il ciclo `for` su `paths` si esaurisce).\n            saved: List[str] = []\n            for img_path in paths:\n                if len(saved) >= target:\n                    break\n                result = self._process_image(img_path, concept_labels, crop_dir, cropper, threshold)\n                if result:\n                    saved.append(result[0])\n\n            n_crops = len(saved)\n            if n_crops == 0:\n                print(f\"[{concept}] No crops produced — skipping augmentation.\")\n                continue\n\n            # Phase 2: augment to fill up to target, if augmentation enabled\n            n_needed = target - n_crops\n            augmented: List[str] = []\n            if n_needed > 0 and aug_on:\n                cycle = itertools.cycle(saved)\n                while len(augmented) < n_needed:\n                    src = next(cycle)\n                    img = Image.open(src).convert(\"RGB\")\n                    out_img = augmenter.augment_one(img)\n                    out_fn = os.path.join(augm_dir, f\"aug_{uuid.uuid4().hex[:8]}.png\")\n                    out_img.save(out_fn)\n                    augmented.append(out_fn)\n\n            # Phase 3: trim if over target (rare)\n            final = saved + augmented\n            if len(final) > target:\n                final = random.sample(final, target)\n            all_crops[concept] = final\n\n            print(f\"[{concept}] final dataset: {len(final)} images \"\n                  f\"({n_crops} crops + {len(augmented)} augmented).\")\n        return all_crops\n\n    # ---------- backward-compat wrapper ----------\n\n    def run(self, class_mapping, labels, crop_type, threshold=None):\n        \"\"\"Backward-compat alias for run_experiment(name=crop_type, ...).\"\"\"\n        return self.run_experiment(\n            name=crop_type,\n            class_mapping=class_mapping,\n            labels=labels,\n            threshold=threshold,\n        )\n\n    def process_image(self, image_path, labels, out_dir, threshold=None):\n        \"\"\"Backward-compat alias — uses the default cropper.\"\"\"\n        return self._process_image(image_path, labels, out_dir, self._default_cropper, threshold)","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"93a50db1","cell_type":"markdown","source":"## 9. Visualisation Utils","metadata":{"papermill":{"duration":0.004656,"end_time":"2026-05-30T14:55:57.288925","exception":false,"start_time":"2026-05-30T14:55:57.284269","status":"completed"},"tags":[]}},{"id":"71b2aa15","cell_type":"code","source":"def annotate(image: Union[Image.Image, np.ndarray], detections: List[DetectionResult]) -> np.ndarray:\n    img = np.array(image) if isinstance(image, Image.Image) else image\n    img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)\n    for det in detections:\n        color = np.random.randint(0, 256, size=3).tolist()\n        b = det.box\n        cv2.rectangle(img, (b.xmin, b.ymin), (b.xmax, b.ymax), color, 2)\n        cv2.putText(img, f\"{det.label}: {det.score:.2f}\", (b.xmin, b.ymin - 10),\n                    cv2.FONT_HERSHEY_SIMPLEX, 0.5, color, 2)\n        if det.mask is not None:\n            contours, _ = cv2.findContours((det.mask * 255).astype(np.uint8),\n                                           cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n            cv2.drawContours(img, contours, -1, color, 2)\n    return cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n\ndef plot_detections(\n    image: Union[Image.Image, np.ndarray],\n    detections: List[DetectionResult],\n    save_name: Optional[str] = None,\n) -> None:\n    plt.imshow(annotate(image, detections))\n    plt.axis(\"off\")\n    if save_name:\n        plt.savefig(save_name, bbox_inches=\"tight\")\n    plt.show()\n\n\ndef plot_images(rows: int, cols: int, directory: str) -> None:\n    \"\"\"Fixed: uses os.path.join instead of string concatenation.\"\"\"\n    exts = (\".png\", \".jpg\", \".jpeg\", \".webp\", \".bmp\")\n    files = [f for f in os.listdir(directory) if f.lower().endswith(exts)]\n    if not files:\n        print(f\"No images in {directory}\")\n        return\n    fig = plt.figure(figsize=(cols * 3, rows * 3))\n    for i in range(rows * cols):\n        ax = fig.add_subplot(rows, cols, i + 1)\n        ax.imshow(Image.open(os.path.join(directory, random.choice(files))))\n        ax.axis(\"off\")\n        ax.set_title(f\"Image {i+1}\", fontsize=8)\n    plt.tight_layout()\n    plt.show()\n\n\ndef plot_detections_plotly(\n    image: np.ndarray,\n    detections: List[DetectionResult],\n    class_colors: Optional[Dict[int, str]] = None,\n) -> None:\n    named_colors = [\n        \"aqua\",\"blue\",\"blueviolet\",\"brown\",\"cadetblue\",\"chartreuse\",\"chocolate\",\n        \"coral\",\"cornflowerblue\",\"crimson\",\"cyan\",\"darkblue\",\"darkcyan\",\"darkgoldenrod\",\n        \"darkgreen\",\"darkorange\",\"darkorchid\",\"darkred\",\"darkseagreen\",\"darkturquoise\",\n        \"darkviolet\",\"deeppink\",\"deepskyblue\",\"dodgerblue\",\"firebrick\",\"forestgreen\",\n        \"fuchsia\",\"gold\",\"goldenrod\",\"green\",\"hotpink\",\"indianred\",\"indigo\",\"lawngreen\",\n        \"limegreen\",\"magenta\",\"maroon\",\"mediumblue\",\"mediumorchid\",\"mediumpurple\",\n        \"mediumseagreen\",\"mediumturquoise\",\"mediumvioletred\",\"midnightblue\",\"navy\",\n        \"olive\",\"orange\",\"orangered\",\"orchid\",\"peru\",\"pink\",\"plum\",\"purple\",\"red\",\n        \"rosybrown\",\"royalblue\",\"saddlebrown\",\"salmon\",\"seagreen\",\"sienna\",\"skyblue\",\n        \"slateblue\",\"springgreen\",\"steelblue\",\"tan\",\"teal\",\"tomato\",\"turquoise\",\n        \"violet\",\"yellowgreen\",\n    ]\n    if class_colors is None:\n        sample = random.sample(named_colors, min(len(detections), len(named_colors)))\n        class_colors = {i: c for i, c in enumerate(sample)}\n\n    fig = px.imshow(image)\n    shapes = []\n    for idx, det in enumerate(detections):\n        color = class_colors.get(idx, \"red\")\n        if det.mask is not None:\n            polygon = mask_to_polygon(det.mask)\n            fig.add_trace(go.Scatter(\n                x=[p[0] for p in polygon] + [polygon[0][0]],\n                y=[p[1] for p in polygon] + [polygon[0][1]],\n                mode=\"lines\", line=dict(color=color, width=2),\n                fill=\"toself\", name=f\"{det.label}: {det.score:.2f}\",\n            ))\n        xmin, ymin, xmax, ymax = det.box.xyxy\n        shapes.append(dict(type=\"rect\", xref=\"x\", yref=\"y\",\n                           x0=xmin, y0=ymin, x1=xmax, y1=ymax,\n                           line=dict(color=color)))\n\n    buttons = (\n        [dict(label=\"None\", method=\"relayout\", args=[\"shapes\", []])]\n        + [dict(label=f\"Det {i+1}\", method=\"relayout\", args=[\"shapes\", [s]]) for i, s in enumerate(shapes)]\n        + [dict(label=\"All\", method=\"relayout\", args=[\"shapes\", shapes])]\n    )\n    fig.update_layout(\n        xaxis=dict(visible=False), yaxis=dict(visible=False), showlegend=True,\n        updatemenus=[dict(type=\"buttons\", direction=\"up\", buttons=buttons)],\n        legend=dict(orientation=\"h\", yanchor=\"bottom\", y=1.02, xanchor=\"right\", x=1),\n    )\n    fig.show()","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"8d3917ac","cell_type":"markdown","source":"## 10. Example Usage","metadata":{"papermill":{"duration":0.004619,"end_time":"2026-05-30T14:55:57.325181","exception":false,"start_time":"2026-05-30T14:55:57.320562","status":"completed"},"tags":[]}},{"id":"ba60846e","cell_type":"code","source":"config = PipelineConfig(\n    # FIX 1 — threshold abbassato da 0.4 a 0.25.\n    # 0.4 era troppo alto per prompt di texture astratta (es. \"striped pattern\"):\n    # SAM3 assegna spesso score 0.25-0.39 a texture che non sono oggetti discreti,\n    # producendo pochissimi crop. 0.25 è il floor consigliato per questo dominio.\n    detector=DetectorConfig(box_threshold=0.25),\n    segmenter=SegmenterConfig(polygon_refinement=False),\n    crop=CropConfig(\n        policy=\"sliding_window\",   # default; overridden per-experiment below\n        min_win=64,\n        req_win=224,\n        size_bias=1.5,\n        mask_selection=\"highest_score\",\n    ),\n    augmentation=AugmentationConfig(\n        enabled=True,\n        rotation=True, shearing=True, flipping=True,\n        brightness=True, contrast=True, saturation=True, hue=True,\n        # FIX 3 — ranges ampliati rispetto ai valori precedenti:\n        #   brightness (0.8,1.2) -> (0.6,1.4)  |  contrast (0.7,1.3) -> (0.5,1.5)\n        #   saturation (0.7,1.3) -> (0.5,1.5)  |  hue (-0.1,0.1) -> (-0.3,0.3)\n        #   shear (-10,10) -> (-15,15)\n        brightness_range=(0.6, 1.4),\n        contrast_range=(0.5, 1.5),\n        saturation_range=(0.5, 1.5),\n        hue_range=(-0.3, 0.3),\n        shear_range=(-15, 15),\n    ),\n    dataset=DatasetConfig(\n        base_path=\"/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train\",\n        images_per_class=None,        # nessun tetto: ispeziona TUTTE le immagini sorgente\n        target_images_per_concept=120,\n    ),\n    device=\"auto\",\n    crops_path=\"/kaggle/working/crops\",\n    augmented_path=\"/kaggle/working/augmented_data\",\n)","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"f3dc1328","cell_type":"code","source":"pipe = GroundingSAMPipeline(config)","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"a4380178","cell_type":"code","source":"# Mappatura gerarchica completa: concept -> case -> [synsets]\nSYNSET_MAPPING = {\n    \"striped\": {\n        \"base\": [\"n02391049\"], # zebra\n        \"correlated\": [\"n02391049\", \"n02129604\"], # zebra, tiger\n        \"random\": [\"n02134084\", \"n03942813\", \"n03584254\", \"n03481172\"] # ice bear, ping-pong ball, iPod, hammer\n    },\n    \"dotted\": {\n        \"base\": [\"n02110341\"], # dalmatian\n        \"correlated\": [\"n02110341\", \"n02165456\"], # dalmatian, ladybug\n        \"random\": [\"n02445715\", \"n03930313\", \"n03271574\", \"n04536866\"] # skunk, picket fence, electric fan, violin\n    },\n    \"chequered\": {\n        \"base\": [\"n06785654\"], # crossword puzzle\n        \"correlated\": [\"n06785654\", \"n04033995\"], # crossword puzzle, quilt\n        \"random\": [\"n09468604\", \"n01910747\", \"n04069434\", \"n03792782\"] # valley, jellyfish, reflex camera, mountain bike\n    },\n    \"wood\": {\n        \"base\": [\"n03930313\"], # picket fence, lumberyard\n        \"correlated\": [\"n04597913\", \"n03891251\"], # wooden spoon, park bench\n        \"random\": [\"n04311004\", \"n04589890\", \"n01910747\", \"n09472597\"] # steel arch bridge, window screen, jellyfish, volcano\n    },\n    \"water\": {\n        \"base\": [\"n09332890\", \"n03388043\"], # lakeside, fountain\n        \"correlated\": [\"n09421951\", \"n09332890\", \"n03388043\"], # coral reef, sandbar\n        \"random\": [\"n03347037\", \"n03661043\", \"n04154565\"] # fire screen, library, screwdriver\n    },\n    \"braided\": {\n        \"base\": [\"n03482405\"], # hamper\n        \"correlated\": [\"n03482405\", \"n03627232\", \"n07695742\"], # knot, pretzel\n        \"random\": [\"n03982430\", \"n02088466\", \"n07715103\"] # billiard table, bloodhound, space shuttle, cauliflower\n    },\n    \"bubbly\": {\n        \"base\": [\"n09229709\"], # bubble\n        \"correlated\": [\"n09229709\", \"n02823750\", \"n09428293\"], # bubble, beer glass, seashore\n        \"random\": [\"n04326547\", \"n04482393\", \"n02280649\"] # stone wall, tricycle, cabbage butterfly\n    },\n    # Concetti ortogonali inseriti direttamente come \"case\" aggiuntivo\n    \"fibrous\": {\n        \"orthogonal\": [\"n07802026\", \"n04584207\", \"n02906734\"] # hay, wig, broom\n    },\n    \"veined\": {\n        \"orthogonal\": [\"n07714571\", \"n01917289\"] # cabbage, brain coral\n    }\n}\n\n# Fusione di LABELS_BASE e LABELS_ORTHO\nLABELS_BASE = {\n    \"striped\":   [\"striped pattern\", \"stripes\", \"striped texture\"],\n    \"dotted\":    [\"dotted pattern\", \"dots\", \"dotted texture\", \"spotted pattern\", \"spots\", \"polka dots\", \"spotted texture\"],\n    \"chequered\": [\"checkered pattern\",\"checkerboard\", \"chequered pattern\"],\n    \"wood\":      [\"wooden material\", \"wood\", \"wooden texture\", \"wooden object\"],\n    \"water\":     [\"water element\", \"water\", \"water texture\"],\n    \"braided\":   [\"braided pattern\", \"braided\", \"braided texture\", \"braided object\"],\n    \"bubbly\":    [\"bubbly pattern\", \"bubbles\", \"bubbly texture\", \"bubbly object\"],\n    \"fibrous\":   [\"fibrous pattern\", \"fibers\", \"fibrous texture\", \"fibrous object\"],\n    \"veined\":    [\"veined pattern\", \"veins\", \"veined texture\", \"veined object\", \"marble texture\", \"leaf veins\", \"veined marble\",\n                    \"organic vein pattern\"]\n}\n\n# Prompt variations espanse per coprire tutte le classi ed evitare KeyError nel loop\nPROMPT_VARIANTS = {\n    \"vanilla_prompt_variations\": {\n        \"striped\": [\"a striped surface\"],\n        \"dotted\": [\"a dotted surface\"],\n        \"chequered\": [\"a chequered surface\"],\n        \"wood\": [\"a wooden surface\"],\n        \"water\": [\"a body of water\"],\n        \"braided\": [\"a braided surface\"],\n        \"bubbly\": [\"a bubbly surface\"],\n        \"fibrous\": [\"a fibrous surface\"],\n        \"veined\": [\"a veined surface\"]\n    },\n    \"augmentation_prompt_variations\": {\n        \"striped\": [\"stripes only\"],\n        \"dotted\": [\"dots only\"],\n        \"chequered\": [\"checkerboard only\"],\n        \"wood\": [\"wood only\"],\n        \"water\": [\"water only\"],\n        \"braided\": [\"braids only\"],\n        \"bubbly\": [\"bubbles only\"],\n        \"fibrous\": [\"fibers only\"],\n        \"veined\": [\"veins only\"]\n    }\n}","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"c3bdf7bb","cell_type":"markdown","source":"## 11. Ablation Experiments\n\nThe pipeline is built once and run many times with different module settings.\nEach experiment is saved to its own subdirectory under `crops_path` / `augmented_path`,\nso runs do not clobber each other.\n\n**Available knobs on `pipe.run_experiment(...)`:**\n\n- `crop_policy`: one of `\"sliding_window\"`, `\"bbox\"`, `\"center_mask\"`, `\"largest_bbox\"`\n- `augmentation_enabled`: `True` or `False`\n- `labels`: any `Dict[concept, List[str]]` you like for prompt variations\n- `threshold`: SAM3 score threshold (default from `config.detector.box_threshold`)\n","metadata":{"papermill":{"duration":0.005149,"end_time":"2026-05-30T14:56:20.097051","exception":false,"start_time":"2026-05-30T14:56:20.091902","status":"completed"},"tags":[]}},{"id":"f4c84972","cell_type":"markdown","source":"### 11. Full Experiment","metadata":{"papermill":{"duration":0.004932,"end_time":"2026-05-30T14:56:20.106763","exception":false,"start_time":"2026-05-30T14:56:20.101831","status":"completed"},"tags":[]}},{"id":"e6cb57a2","cell_type":"code","source":"import os\nimport uuid\nimport itertools\nfrom itertools import zip_longest\nfrom dataclasses import replace\nfrom PIL import Image\n\ndef generate_dual_dataset(pipe: GroundingSAMPipeline, target_imgs: int = 120, mode: str = \"vanilla\"):\n    \"\"\"\n    Genera dataset differenziati:\n    - mode=\"vanilla\": solo crop puri (fino a target_imgs).\n    - mode=\"pure_augmentation\": genera target_imgs immagini tutte trasformate partendo dai crop puri.\n    \"\"\"\n    POLICIES = [\"sliding_window\", \"bbox\", \"center_mask\", \"largest_bbox\"]\n    TARGET_CONCEPTS = [\"striped\", \"dotted\", \"chequered\", \"wood\", \"water\", \"braided\", \"bubbly\", \"fibrous\", \"veined\"]\n\n    croppers = {p: Cropper(replace(pipe.config.crop, policy=p)) for p in POLICIES}\n    # Forziamo l'abilitazione dell'augmenter per il secondo dataset\n    augmenter = Augmenter(replace(pipe.config.augmentation, enabled=True))\n\n    for concept in TARGET_CONCEPTS:\n        cases = SYNSET_MAPPING[concept] #\n        labels = LABELS_BASE[concept]   #\n\n        for case, synsets in cases.items():\n            print(f\"\\n>>> Mode: {mode} | Concept: {concept} / {case} <<<\")\n\n            collector = DatasetCollector(pipe.config.dataset, {concept: synsets})\n            image_paths = collector.collect().get(concept, [])\n\n            saved_crops = {p: [] for p in POLICIES}\n\n            # --- FASE 1: Estrazione dei \"Semi\" (Vanilla Crops) ---\n            # FIX 1 — threshold esplicito a 0.25 invece del default della config\n            # (che potrebbe essere più alto se modificato dall'utente).\n            # FIX 2 — in mode=\"pure_augmentation\" la Fase 1 viene saltata se i\n            # vanilla crop esistono già su disco (evita di perdere immagini che\n            # SAM3 trovava nella run vanilla ma non trova in questa passata).\n            _vanilla_exist = all(\n                os.path.isdir(os.path.join(pipe.config.crops_path, \"vanilla\", f\"vanilla_{pol}\", concept, case))\n                and bool(os.listdir(os.path.join(pipe.config.crops_path, \"vanilla\", f\"vanilla_{pol}\", concept, case)))\n                for pol in POLICIES\n            )\n            _skip_extraction = (mode == \"pure_augmentation\" and _vanilla_exist)\n\n            # L'estrazione scorre TUTTE le immagini sorgente (image_paths contiene\n            # ora l'intero pool, vedi DatasetCollector con images_per_class=None) e si\n            # ferma SOLO quando: (a) ogni policy ha raggiunto target_imgs, oppure\n            # (b) le immagini di partenza sono finite.\n            for img_path in ([] if _skip_extraction else image_paths):\n                if all(len(saved_crops[p]) >= target_imgs for p in POLICIES):\n                    break\n\n                try:\n                    image_array, detections = pipe.model.process(img_path, labels, threshold=0.25)\n                except Exception:\n                    continue\n\n                if not detections:\n                    continue\n\n                for pol in POLICIES:\n                    if len(saved_crops[pol]) >= target_imgs:\n                        continue\n\n                    # Cartella di destinazione per i crop puri\n                    vanilla_dir = os.path.join(pipe.config.crops_path, \"vanilla\", f\"vanilla_{pol}\", concept, case)\n                    os.makedirs(vanilla_dir, exist_ok=True)\n\n                    try:\n                        crop_fn, _ = croppers[pol].crop(image_array, detections, vanilla_dir)\n                        saved_crops[pol].append(crop_fn)\n                    except ValueError:\n                        pass\n\n            # --- FASE 2: Gestione Output in base al Mode ---\n            for pol in POLICIES:\n                n_crops = len(saved_crops[pol])\n\n                if mode == \"pure_augmentation\":\n                    # Creiamo un dataset di target_imgs immagini interamente aumentate\n                    augm_dir = os.path.join(pipe.config.crops_path, \"augmented\", f\"augmented_{pol}\", concept, case)\n                    os.makedirs(augm_dir, exist_ok=True)\n\n                    # FIX 2 — leggi i vanilla già su disco se l'estrazione è stata saltata\n                    # oppure se questa run non ha trovato crop nuovi.\n                    # Questo garantisce che \"pure_augmentation\" abbia sempre semi\n                    # anche quando SAM3 non trova nulla nella seconda passata.\n                    _exts = (\".png\", \".jpg\", \".jpeg\", \".webp\", \".bmp\")\n                    vanilla_dir = os.path.join(pipe.config.crops_path, \"vanilla\", f\"vanilla_{pol}\", concept, case)\n                    if n_crops == 0 and os.path.isdir(vanilla_dir):\n                        disk_seeds = [\n                            os.path.join(vanilla_dir, f)\n                            for f in os.listdir(vanilla_dir)\n                            if f.lower().endswith(_exts)\n                        ]\n                        print(f\"  [{pol}] Nessun crop nella run corrente; uso {len(disk_seeds)} semi da disco.\")\n                    else:\n                        disk_seeds = saved_crops[pol]\n\n                    if disk_seeds:\n                        print(f\"  [{pol}] Genero {target_imgs} immagini aumentate da {len(disk_seeds)} semi.\")\n                        cycle = itertools.cycle(disk_seeds)\n                        for _ in range(target_imgs):\n                            src = next(cycle)\n                            img = Image.open(src).convert(\"RGB\")\n                            aug_img = augmenter.augment_one(img)\n                            aug_fn = os.path.join(augm_dir, f\"aug_{uuid.uuid4().hex[:8]}.png\")\n                            aug_img.save(aug_fn)\n                    else:\n                        print(f\"  [{pol}] Impossibile aumentare: 0 semi trovati (né in run né su disco).\")\n\n                else:\n                    # In modalità vanilla abbiamo già salvato i file in Fase 1\n                    print(f\"  [{pol}] Dataset vanilla completato: {n_crops} immagini.\")\n\n# Esecuzione per generare entrambi i dataset\nprint(\"--- GENERAZIONE DATASET VANILLA ---\")\ngenerate_dual_dataset(pipe, target_imgs=120, mode=\"vanilla\")\n\nprint(\"\\n--- GENERAZIONE DATASET AUGMENTED ---\")\ngenerate_dual_dataset(pipe, target_imgs=120, mode=\"pure_augmentation\")","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"9e48f09c","cell_type":"code","source":"print(\"\\n=== Inizio Ablation sui Prompt ===\")\nTARGET_CONCEPTS = [\"striped\", \"dotted\", \"chequered\", \"wood\", \"water\", \"braided\", \"bubbly\", \"fibrous\", \"veined\"] # Limite imposto\n\nfor variant_name, prompts in PROMPT_VARIANTS.items():\n    # Determina se la variante richiede augmentation dal nome\n    # is_augmented = \"augmentation\" in variant_name\n    \n    for concept in TARGET_CONCEPTS:\n        cases = SYNSET_MAPPING[concept]\n        \n        for case, synsets in cases.items():\n            \n            # Unendo concept e case inganniamo run_experiment facendogli creare \n            # l'albero di cartelle corretto: .../variant_name/concept/case/\n            pseudo_concept = f\"{concept}/{case}\"\n            \n            pipe.run_experiment(\n                name=f\"prompt_variations/{variant_name}\",\n                class_mapping={pseudo_concept: synsets},\n                labels={pseudo_concept: prompts.get(concept, [concept])},\n                crop_policy=\"sliding_window\", # Mantieni una policy fissa per valutare i prompt\n                augmentation_enabled=False\n            )","metadata":{"tags":[]},"outputs":[],"execution_count":null},{"id":"5cd7436e","cell_type":"markdown","source":"### 11.5 Visualise any experiment's outputs","metadata":{"papermill":{"duration":0.013754,"end_time":"2026-05-31T02:14:46.056974","exception":false,"start_time":"2026-05-31T02:14:46.04322","status":"completed"},"tags":[]}},{"id":"fd566c50","cell_type":"markdown","source":"### Nota\nLa cella finale di sola **visualizzazione** (sezione *11.5*) non è inclusa: il file caricato era stato troncato al limite di 50 MB a causa delle immagini incorporate negli output, quindi il suo sorgente non era recuperabile. Non incide sulla logica di estrazione modificata qui.","metadata":{}}]}