{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"},{"sourceId":11932060,"sourceType":"datasetVersion","datasetId":7501747},{"sourceId":327336,"sourceType":"modelInstanceVersion","modelInstanceId":274744,"modelId":295634}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =========================\n# Z-sweep on TRAIN (3 tomos) for Z ∈ {1,10,30}\n# Prints per-tomogram results + total runtime per Z\n# Saves: /kaggle/working/z_sweep_train_full.csv\n#        /kaggle/working/z_sweep_train_table.csv\n#        /kaggle/working/z_sweep_train_summary.csv\n# =========================\n\n# --- Basic setup ---\nimport os, sys, io, time, logging, warnings, re, random\nfrom pathlib import Path\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom PIL import Image\nimport torch\n\n# Paths\nDATA_ROOT  = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025\"\nTRAIN_DIR  = f\"{DATA_ROOT}/train\"\nGT_CSV     = f\"{DATA_ROOT}/train_labels.csv\"\nMODEL_PATH = \"/kaggle/input/mahf-yolo-train/mayolov2f.pt\"\n\n# Params\nSIZE = 1024\nCONFIDENCE_THRESHOLD = 0.8\nWBF_IOU   = 0.5\nNMS3D_IOU = 0.2\nCONCENTRATION = 1.0      # 1.0 = use all slices\nZ_LIST = [30]     # <<< sweep values\n\n# Device\ndevice = \"cuda:0\" if torch.cuda.is_available() else \"cpu\"\nif device.startswith(\"cuda\"):\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cuda.matmul.allow_tf32 = True\n    torch.backends.cudnn.allow_tf32 = True\n\nnp.random.seed(42)\ntorch.manual_seed(42)\nrandom.seed(42)\n\n# --- Pull ultralytics fork used by MHAF-YOLO ---\n!cp -r /kaggle/input/mhafyolo/pytorch/default/1/MHAF-YOLO-main /kaggle/working/\nos.chdir(\"/kaggle/working/MHAF-YOLO-main\")\n\nfrom ultralytics import YOLOv10\nfrom ultralytics.utils import LOGGER\nwarnings.filterwarnings(\"ignore\", category=FutureWarning)\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\nLOGGER.setLevel(logging.ERROR)\n\nclass Quiet:\n    def __enter__(self):\n        self._stdout, self._stderr = sys.stdout, sys.stderr\n        sys.stdout = self._buf_out = io.StringIO()\n        sys.stderr = self._buf_err = io.StringIO()\n        return self\n    def __exit__(self, *args):\n        sys.stdout, sys.stderr = self._stdout, self._stderr\n\n# --- Helpers ---\ndef list_tomos(root_dir: str):\n    return [d for d in os.listdir(root_dir) if os.path.isdir(os.path.join(root_dir, d))]\n\ndef clamp(v, a, b):\n    return max(a, min(b, v))\n\ndef normalize_slice(gray_uint8: np.ndarray) -> np.ndarray:\n    if gray_uint8.dtype != np.uint8:\n        gray_uint8 = gray_uint8.astype(np.uint8)\n    p2, p98 = np.percentile(gray_uint8, 2), np.percentile(gray_uint8, 98)\n    if p98 <= p2 + 1e-6:\n        return gray_uint8.copy()\n    clipped = np.clip(gray_uint8, p2, p98)\n    return (255.0 * (clipped - p2) / (p98 - p2)).astype(np.uint8)\n\ndef read_gray_resized(path: str, size: int) -> np.ndarray:\n    im = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    if im is None:\n        im = np.array(Image.open(path).convert(\"L\"))\n    if (im.shape[1], im.shape[0]) != (size, size):\n        im = cv2.resize(im, (size, size))\n    return im\n\ndef build_25d_rgb(tomo_dir: str, slice_files: list, idx: int, size: int, window: int) -> np.ndarray:\n    n = len(slice_files)\n    g = normalize_slice(read_gray_resized(os.path.join(tomo_dir, slice_files[idx]), size))\n    prev = [read_gray_resized(os.path.join(tomo_dir, slice_files[clamp(idx-d,0,n-1)]), size)\n            for d in range(1, window+1)]\n    nxt  = [read_gray_resized(os.path.join(tomo_dir, slice_files[clamp(idx+d,0,n-1)]), size)\n            for d in range(1, window+1)]\n    r = normalize_slice(np.mean(prev, axis=0).astype(np.uint8)) if prev else g.copy()\n    b = normalize_slice(np.mean(nxt,  axis=0).astype(np.uint8)) if nxt  else g.copy()\n    return np.stack([r, g, b], axis=2)\n\ndef iou_xyxy(a, b) -> float:\n    ax1, ay1, ax2, ay2 = a\n    bx1, by1, bx2, by2 = b\n    inter_x1, inter_y1 = max(ax1, bx1), max(ay1, by1)\n    inter_x2, inter_y2 = min(ax2, bx2), min(ay2, by2)\n    iw, ih = max(0.0, inter_x2 - inter_x1), max(0.0, inter_y2 - inter_y1)\n    inter = iw * ih\n    ua = (ax2-ax1)*(ay2-ay1) + (bx2-bx1)*(by2-by1) - inter\n    return inter/ua if ua > 0 else 0.0\n\ndef weighted_box_fusion_simple(xyxy, conf, iou_thr=WBF_IOU):\n    xyxy = np.asarray(xyxy, dtype=np.float32)\n    conf = np.asarray(conf, dtype=np.float32)\n    if xyxy.size == 0:\n        return []\n    if xyxy.ndim == 1:\n        if xyxy.shape[0] != 4:\n            return []\n        xyxy = xyxy[None, :]\n    if conf.ndim == 0:\n        conf = np.array([float(conf)], dtype=np.float32)\n    conf = conf.reshape(-1)\n    if conf.shape[0] != xyxy.shape[0]:\n        n = min(conf.shape[0], xyxy.shape[0])\n        xyxy, conf = xyxy[:n], conf[:n]\n        if n == 0:\n            return []\n    used = np.zeros(len(xyxy), dtype=bool)\n    fused = []\n    for i in range(len(xyxy)):\n        if used[i]:\n            continue\n        group = [i]; used[i] = True\n        for j in range(i+1, len(xyxy)):\n            if used[j]:\n                continue\n            if iou_xyxy(xyxy[i], xyxy[j]) >= iou_thr:\n                group.append(j); used[j] = True\n        w = conf[group]; w = w / (w.sum() + 1e-9)\n        bb = (xyxy[group] * w[:, None]).sum(axis=0)\n        cf = float(conf[group].mean())\n        fused.append((bb[0], bb[1], bb[2], bb[3], cf))\n    return fused\n\ndef _safe_model_infer(model, image_rgb, img_size, device):\n    if not isinstance(image_rgb, np.ndarray):\n        image_rgb = np.array(image_rgb)\n    if image_rgb.dtype != np.uint8:\n        image_rgb = image_rgb.astype(np.uint8)\n    image_rgb = np.ascontiguousarray(image_rgb)\n    with Quiet():\n        return model(image_rgb, imgsz=img_size, device=device, verbose=False)\n\ndef predict_tta_wbf_single(model, image_rgb: np.ndarray, img_size=SIZE, conf_thres=0.0, wbf_iou=WBF_IOU, device='cuda:0'):\n    H, W = image_rgb.shape[:2]\n    all_boxes, all_confs, all_clss = [], [], []\n    # original\n    res = _safe_model_infer(model, image_rgb, img_size, device)\n    for r in res:\n        if r.boxes is None or len(r.boxes) == 0:\n            continue\n        xyxy = r.boxes.xyxy.cpu().numpy()\n        conf = r.boxes.conf.cpu().numpy()\n        cls  = r.boxes.cls.cpu().numpy().astype(int)\n        all_boxes.append(xyxy); all_confs.append(conf); all_clss.append(cls)\n    # hflip\n    img_f = cv2.flip(image_rgb, 1)\n    resf = _safe_model_infer(model, img_f, img_size, device)\n    for r in resf:\n        if r.boxes is None or len(r.boxes) == 0:\n            continue\n        xyxy = r.boxes.xyxy.cpu().numpy()\n        conf = r.boxes.conf.cpu().numpy()\n        cls  = r.boxes.cls.cpu().numpy().astype(int)\n        inv = xyxy.copy()\n        inv[:, 0] = W - xyxy[:, 2]\n        inv[:, 2] = W - xyxy[:, 0]\n        all_boxes.append(inv); all_confs.append(conf); all_clss.append(cls)\n    if len(all_boxes) == 0:\n        return []\n    boxes_cat = np.concatenate(all_boxes, axis=0)\n    confs_cat = np.concatenate(all_confs, axis=0)\n    clss_cat  = np.concatenate(all_clss, axis=0)\n    keep = confs_cat >= conf_thres\n    boxes_cat, confs_cat, clss_cat = boxes_cat[keep], confs_cat[keep], clss_cat[keep]\n    fused_all = []\n    for c in np.unique(clss_cat):\n        idx = np.where(clss_cat == c)[0]\n        fused = weighted_box_fusion_simple(boxes_cat[idx], confs_cat[idx], iou_thr=wbf_iou)\n        for (x1,y1,x2,y2,cf) in fused:\n            fused_all.append((x1,y1,x2,y2,cf,int(c)))\n    return fused_all\n\ndef perform_3d_nms(detections, iou_threshold=NMS3D_IOU):\n    if not detections:\n        return []\n    dets = sorted(detections, key=lambda d: d['confidence'], reverse=True)\n    final = []\n    box_size = 24\n    dist_thr = box_size * iou_threshold\n    def dist3(a,b):\n        return np.sqrt((a['z']-b['z'])**2 + (a['y']-b['y'])**2 + (a['x']-b['x'])**2)\n    while dets:\n        best = dets.pop(0)\n        final.append(best)\n        dets = [d for d in dets if dist3(d, best) > dist_thr]\n    return final\n\n# return also meta for logging: status, num_slices, elapsed\ndef process_tomogram(root_dir, tomo_id, model, size=SIZE, device=device, window_z=10):\n    t0 = time.time()\n    tomo_dir = os.path.join(root_dir, tomo_id)\n    slice_files_all = sorted([f for f in os.listdir(tomo_dir) if f.lower().endswith('.jpg')])\n    if len(slice_files_all) == 0:\n        elapsed = time.time() - t0\n        return {'tomo_id': tomo_id, 'Motor axis 0': -1, 'Motor axis 1': -1, 'Motor axis 2': -1,\n                '_status': 'NONE', '_slices': 0, '_elapsed': elapsed}\n\n    h0, w0 = cv2.imread(os.path.join(tomo_dir, slice_files_all[0]), cv2.IMREAD_GRAYSCALE).shape[:2]\n    scale_y, scale_x = h0 / float(size), w0 / float(size)\n\n    if CONCENTRATION < 1.0:\n        take = max(1, int(len(slice_files_all) * CONCENTRATION))\n        sel = np.linspace(0, len(slice_files_all)-1, take)\n        slice_idx = np.round(sel).astype(int).tolist()\n    else:\n        slice_idx = list(range(len(slice_files_all)))\n\n    all_dets = []\n    for i in slice_idx:\n        rgb = build_25d_rgb(tomo_dir, slice_files_all, i, size=size, window=window_z)\n        try: znum = int(Path(slice_files_all[i]).stem.split('_')[1])\n        except Exception: znum = i\n        fused = predict_tta_wbf_single(model, rgb, img_size=size, conf_thres=0.0, wbf_iou=WBF_IOU, device=device)\n        for (x1,y1,x2,y2,cf,cl) in fused:\n            if cf >= CONFIDENCE_THRESHOLD:\n                xc = 0.5*(x1+x2); yc = 0.5*(y1+y2)\n                all_dets.append({'z': int(round(znum)),\n                                 'y': int(round(yc * scale_y)),\n                                 'x': int(round(xc * scale_x)),\n                                 'confidence': float(cf)})\n\n    final_dets = perform_3d_nms(all_dets, NMS3D_IOU)\n    status = \"FOUND\" if final_dets else \"NONE\"\n    best = final_dets[0] if final_dets else {'z': -1, 'y': -1, 'x': -1}\n    elapsed = time.time() - t0\n    return {'tomo_id': tomo_id,\n            'Motor axis 0': int(best['z']), 'Motor axis 1': int(best['y']), 'Motor axis 2': int(best['x']),\n            '_status': status, '_slices': len(slice_idx), '_elapsed': elapsed}\n\ndef load_gt_points(gt_csv: str) -> pd.DataFrame:\n    df = pd.read_csv(gt_csv)\n    cols = [c.lower().strip().replace('_',' ') for c in df.columns]\n    df.columns = cols\n    if {'motor axis 0','motor axis 1','motor axis 2','tomo id'}.issubset(df.columns):\n        out = df[['tomo id','motor axis 2','motor axis 1','motor axis 0']].copy()\n        out.columns = ['tomo_id','gt_x','gt_y','gt_z']\n        return out\n    raise ValueError(f\"Unrecognized GT header: {df.columns}\")\n\n# --- Load model once (quiet warmup) ---\nwith Quiet():\n    model = YOLOv10(MODEL_PATH)\nmodel.to(device)\ntry:\n    model.model.float()\nexcept Exception:\n    pass\nwith Quiet():\n    _ = model(np.zeros((SIZE,SIZE,3), np.uint8), imgsz=SIZE, device=device, verbose=False)\n\n# --- Build 3-sample set from TRAIN (intersection of disk & GT) ---\ngt_points = load_gt_points(GT_CSV)\ntrain_ids_disk = set(list_tomos(TRAIN_DIR))\ntrain_ids_gt   = set(gt_points[\"tomo_id\"].unique())\ncandidates     = sorted(list(train_ids_disk & train_ids_gt))\nif len(candidates) == 0:\n    raise RuntimeError(\"No train tomograms found with GT on disk.\")\n\nif len(candidates) < 3:\n    print(f\"WARNING: only {len(candidates)} train tomograms available with GT; using all of them.\")\n    sampled_ids = candidates\nelse:\n    sampled_ids = ['tomo_003acc', 'tomo_00e047', 'tomo_01a877']\n\nprint(\"Selected train tomograms (count={}):\".format(len(sampled_ids)))\nprint(sampled_ids)\n\n# --- Run sweep for each Z and collect predictions (with per-tomo logging) ---\nall_preds = []\noverall_t0 = time.time()\nfor z in Z_LIST:\n    rows = []\n    print(f\"\\n=== Running Z={z} on {len(sampled_ids)} tomograms ===\")\n    z_t0 = time.time()\n\n    for tid in sampled_ids:\n        out = process_tomogram(TRAIN_DIR, tid, model, size=SIZE, device=device, window_z=z)\n        out['Z_setup'] = z\n        rows.append(out)\n\n        # per-tomogram print: pred vs GT + errors\n        gtr = gt_points[gt_points['tomo_id'] == tid]\n        if len(gtr):\n            gx, gy, gz = int(gtr['gt_x'].iloc[0]), int(gtr['gt_y'].iloc[0]), int(gtr['gt_z'].iloc[0])\n            px, py, pz = int(out['Motor axis 2']), int(out['Motor axis 1']), int(out['Motor axis 0'])\n            errx, erry, errz = abs(px-gx), abs(py-gy), abs(pz-gz)\n            l2 = ( (px-gx)**2 + (py-gy)**2 + (pz-gz)**2 )**0.5\n            print(f\"[Z={z}] {tid} — {out['_status']} — {out['_slices']} slices — {out['_elapsed']:.1f}s | \"\n                  f\"pred=(x={px}, y={py}, z={pz})  gt=(x={gx}, y={gy}, z={gz})  \"\n                  f\"|err|=({errx},{erry},{errz})  L2={l2:.1f}\")\n        else:\n            print(f\"[Z={z}] {tid} — {out['_status']} — {out['_slices']} slices — {out['_elapsed']:.1f}s | GT MISSING\")\n\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n    # total runtime for this Z\n    z_elapsed = time.time() - z_t0\n    print(f\"=== Total runtime for Z={z}: {z_elapsed:.1f}s ===\")\n\n    df = pd.DataFrame(rows)\n    all_preds.append(df)\n\noverall_elapsed = time.time() - overall_t0\n\npreds_all = pd.concat(all_preds, ignore_index=True).rename(columns={\n    \"Motor axis 0\":\"pred_z\",\n    \"Motor axis 1\":\"pred_y\",\n    \"Motor axis 2\":\"pred_x\"\n})\n\n# --- Merge with GT & compute errors ---\nmerged = (preds_all.merge(gt_points, on='tomo_id', how='left')\n          .assign(err_x=lambda d: (d['pred_x']-d['gt_x']).abs(),\n                  err_y=lambda d: (d['pred_y']-d['gt_y']).abs(),\n                  err_z=lambda d: (d['pred_z']-d['gt_z']).abs()))\nmerged['err_L2'] = np.sqrt((merged['pred_x']-merged['gt_x'])**2 +\n                           (merged['pred_y']-merged['gt_y'])**2 +\n                           (merged['pred_z']-merged['gt_z'])**2)\n\n# --- Human-readable table: GT & preds per Z per tomo ---\nZ_LIST_SORTED = sorted(merged['Z_setup'].unique().tolist())\nrows_tbl = []\nfor tid, g in merged.groupby('tomo_id'):\n    row = {'tomo_id': tid,\n           'GT(x,y,z)': (int(g['gt_x'].iloc[0]), int(g['gt_y'].iloc[0]), int(g['gt_z'].iloc[0]))}\n    for z in Z_LIST_SORTED:\n        gi = g[g['Z_setup']==z].iloc[0]\n        row[f'Pred@Z={z} (x,y,z)'] = (int(gi['pred_x']), int(gi['pred_y']), int(gi['pred_z']))\n        row[f'|err|@Z={z} (x,y,z)'] = (int(abs(gi['pred_x']-gi['gt_x'])),\n                                       int(abs(gi['pred_y']-gi['gt_y'])),\n                                       int(abs(gi['pred_z']-gi['gt_z'])))\n    rows_tbl.append(row)\nfinal_table = pd.DataFrame(rows_tbl)\n\n# --- Summary per Z ---\nsummary = (merged.groupby('Z_setup')\n           .agg(mean_err_x=('err_x','mean'),\n                mean_err_y=('err_y','mean'),\n                mean_err_z=('err_z','mean'),\n                mean_L2=('err_L2','mean'))\n           .round(3)\n           .reset_index())\n\n# --- Save CSVs ---\nout_dir = \"/kaggle/working\"\nfull_csv    = f\"{out_dir}/z_sweep_train_full.csv\"\ntable_csv   = f\"{out_dir}/z_sweep_train_table.csv\"\nsummary_csv = f\"{out_dir}/z_sweep_train_summary.csv\"\n\nmerged.to_csv(full_csv, index=False)\nfinal_table.to_csv(table_csv, index=False)\nsummary.to_csv(summary_csv, index=False)\n\nprint(\"\\n=== Z-sweep summary (mean absolute error) ===\")\nprint(summary.to_string(index=False))\nprint(\"\\nSaved:\")\nprint(full_csv)\nprint(table_csv)\nprint(summary_csv)\n\nprint(f\"\\n=== Overall runtime (all Z values): {overall_elapsed:.1f}s ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-12T14:07:28.909310Z","iopub.execute_input":"2025-08-12T14:07:28.909694Z","execution_failed":"2025-08-12T14:09:54.030Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2) --- Submission generation section ---\nTEST_DIR = f\"{DATA_ROOT}/test\"\nSUBMISSION_PATH = \"/kaggle/working/submission.csv\"\n\n\ndef generate_submission(test_dir=TEST_DIR, model=model, submission_path=SUBMISSION_PATH, z_for_submission=None):\n    \"\"\"\n    Runs inference on the test set and saves submission.csv\n    \"\"\"\n    if z_for_submission is None:\n        z_for_submission = Z_LIST[0]  # default to first sweep Z value\n    test_tomos = sorted([d for d in os.listdir(test_dir) if os.path.isdir(os.path.join(test_dir, d))])\n    total_tomos = len(test_tomos)\n    print(f\"\\n=== Generating submission on {total_tomos} test tomograms (Z={z_for_submission}) ===\")\n\n    results = []\n    motors_found = 0\n\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n    # Sequential loop \n    for tomo_id in test_tomos:\n        result = process_tomogram(test_dir, tomo_id, model, size=SIZE, device=device, window_z=z_for_submission)\n        result['Z_setup'] = z_for_submission\n        results.append(result)\n\n        has_motor = not pd.isna(result['Motor axis 0']) and result['Motor axis 0'] != -1\n        if has_motor:\n            motors_found += 1\n            print(f\"Motor found in {tomo_id} at position: \"\n                  f\"z={result['Motor axis 0']}, y={result['Motor axis 1']}, x={result['Motor axis 2']}\")\n        else:\n            print(f\"No motor detected in {tomo_id}\")\n\n        print(f\"Current detection rate: {motors_found}/{len(results)} \"\n              f\"({motors_found/len(results)*100:.1f}%)\")\n\n    submission_df = pd.DataFrame(results)\n    submission_df = submission_df[['tomo_id', 'Motor axis 0', 'Motor axis 1', 'Motor axis 2']]\n    submission_df.to_csv(submission_path, index=False)\n\n    print(f\"\\nSubmission complete!\")\n    print(f\"Motors detected: {motors_found}/{total_tomos} ({motors_found/total_tomos*100:.1f}%)\")\n    print(f\"Submission saved to: {submission_path}\")\n    print(\"\\nSubmission preview:\")\n    print(submission_df.head())\n\n    return submission_df","metadata":{"trusted":true,"execution":{"execution_failed":"2025-08-12T14:09:54.031Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission_df = generate_submission(z_for_submission=Z_LIST[0])","metadata":{"trusted":true,"execution":{"execution_failed":"2025-08-12T14:09:54.032Z"}},"outputs":[],"execution_count":null}]}