{"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":11744090,"sourceType":"datasetVersion","datasetId":7372391},{"sourceId":11829003,"sourceType":"datasetVersion","datasetId":7431109}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"First of all, I'd like to thank the organizers of the BYU competition for hosting such a challenging and inspiring event.\n\nThis notebook contains my final training code. Despite working with 3D U-Net for nearly three months, it never quite returned my affection — the final score was nearly 0.\n\nStill, this journey was full of valuable experiments and learning. The code below reflects those efforts.\n\n\nI hope that someday, somewhere, this notebook might help someone — even just a little.\n\nThank you.","metadata":{}},{"cell_type":"code","source":"import gc\nimport os\nimport sys\nimport cv2\nimport time\nimport math\nimport json\nimport glob\nimport random\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nfrom sklearn.model_selection import StratifiedKFold\n\nimport scipy.ndimage\nfrom scipy.stats import mode\nfrom scipy import signal\n\n# torch\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import AdamW, Adam\nfrom torch.cuda.amp import GradScaler, autocast\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import (CosineAnnealingLR,\n                                      OneCycleLR,\n                                      ReduceLROnPlateau\n                                     )\nfrom torch.utils.data import WeightedRandomSampler\nfrom torch.utils.data import Sampler\nfrom torch.utils.data import BatchSampler\nimport torch.utils.checkpoint\nimport torchvision.transforms as T\nfrom scipy.cluster.vq import kmeans, vq\nimport statistics\n\n# type hints\nfrom typing import Dict, List\n\n#cluster\n#import hdbscan\n#from sklearn.cluster import DBSCAN\n\n# Other\nfrom sklearn.metrics import fbeta_score\nfrom collections import OrderedDict\nfrom collections import defaultdict\n\nfrom skimage.feature import peak_local_max\nfrom skimage.morphology import skeletonize, thin\nfrom skimage.morphology import medial_axis\nfrom skimage.morphology import remove_small_objects\n\nimport scipy.ndimage\nfrom scipy.ndimage import center_of_mass\nfrom scipy.ndimage import gaussian_filter\nfrom scipy.ndimage import binary_dilation\nfrom scipy.ndimage import convolve\n\nfrom functools import partial\n\nfrom skimage.transform import resize\nfrom skimage.measure import label\nfrom skimage.measure import regionprops\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\nos.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0\"\nos.environ[\"CUDA_LAUNCH_BLOCKING\"] = \"1\"\nprint(f'PyTorch version : {torch.__version__}')","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.813Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    # General\n    SEED         = 42\n    FOLD_LIST    = [0, 1, 2, 3]\n    N_fold       = len(FOLD_LIST)\n    IMSIZE       = 512\n    TARGET_SLICES= 32\n    MIN_POS      = 3 #number of positive in batch_size\n    \n    # Data path\n    base_dir     = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025\"\n    download_dir = \"/kaggle/input/byu-images-and-masks\"\n    model_dir    = \"/kaggle/input/byu-unet-models\"\n    \n    # Model\n    use_premodel     = False\n    dropout_encoder  = 0.1\n    dropout_decoder  = 0.0\n    z_sampling       = True\n    edge_inp_W       = 0.4 #stage1=0.4\n    custom_output    = False\n    use_droppath     = False\n    output_scaling   = True\n    T                = 2.0\n    use_center_heatm = True\n    \n    # DataSet\n    valid_aug        = False # 1stage=True\n    train_batch_size = 4\n    valid_batch_size = 4\n    train_p          = 0.9 # 1stage=0.6\n    valid_p          = 0.3 # 1stage valid_aug=True\n    USE_CROP         = True\n    CROP_THRESOLD    = 0.9 #or 0.2 1stage=0.6\n    CROP_SIZE        = (384, 384, TARGET_SLICES) # (H, W, D)\n    \"\"\"\n    eg:crop_size\n    (416, 416, 32)\n    (384, 384, 32)\n    (256, 256, 32)：1stage\n    (128, 128, 32)\n    (64,  64,  16)\n    \n    \"\"\"\n    # Dataloder\n    num_workers  = os.cpu_count()\n    \n    # Optimizer\n    use_AdamW    = True\n    lr           = 1e-4\n    weight_decay = 1e-3\n    freezing     = True if use_premodel else False\n    \n    # Scheduler\n    ###onecyclelr\n    max_lr       = 5e-5 # 1stage 3e-4\n    epochs       = 42\n    warmup_epochs= 0.3#2.3 #10%\n    \n    # Training\n    GET_OOF      = True\n    valid_warmup = 0 # if 1stage=epochs else=0\n    print_freq   = 40\n    es_round     = 6\n    use_amp      = True\n    max_grad_norm= 10\n    \n    # Loss\n    \n    #1stage loss\n    bce_weight    = 0.8\n    dice_weight   = 0.15\n    tversky_weight= 0.05\n    \n    # 2stage loss\n    bce_weight2        = 0.10\n    dice_weight2       = 0.10\n    tversky_weight2    = 0.15\n    focal_weight2      = 0.25\n    center_weight      = 0.05\n    center_heat_w      = 0.10\n    z_loss_weight      = 0.25\n    use_W_epoch_min= True\n    weights_total  = (bce_weight2+dice_weight2+tversky_weight2+focal_weight2+center_weight+center_heat_w+z_loss_weight)\n\n    \n    # Target\n    change_target= True\n    USE_MEAN     = False\n    \n    # Debug\n    debug       = True\n    more_log    = True\n    f_check     = False\n    weight_show = False\n    print_freq2 = 20\n    if debug:\n        FOLD_LIST = [0]\n        if more_log:\n            print_freq = 30\n            print_freq2= print_freq*5\n\n    normalized_show = True\n    parameter_check = False\n    s2valid_debug   = False\n    flatness_plot   = True\n    other_plot      = False\n    save_plot       = True\n    show_peak_weight= False\n    # sigmoid (outputs)\n    output_check    = False\n    \ncf = CFG()\nprint(f\"num workers    : {cf.num_workers}\")\nprint(f\"FOLD LIST      : {cf.FOLD_LIST}\")\nprint(f\"Pint Frequency : {cf.print_freq}\")\nprint(f\"Use pretrain M : {cf.use_premodel}\")\nprint(f\"Freezing       : {cf.freezing}\")\nprint(f\"Weights total  : {cf.weights_total}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.814Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Device","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using Device is {device}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.815Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### SEED Everything","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.backends.cudnn.deterministic=True # if T4 False else True (speed up)\n        torch.backends.cudnn.benchmark=False    # if T4  True else False (speed up)\n    print(\"Seed Setting Done!\")\nseed_everything(cf.SEED)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.815Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Make Directory","metadata":{}},{"cell_type":"code","source":"MODEL_DIR = \"./models\"\nif not os.path.exists(MODEL_DIR):\n    os.makedirs(MODEL_DIR)\n    \nHISTORY_DIR = \"./history\"\nif not os.path.exists(HISTORY_DIR):\n    os.makedirs(HISTORY_DIR)\n    \nOOF_DIR = \"./oof\"\nif not os.path.exists(OOF_DIR):\n    os.makedirs(OOF_DIR)\n\nFLATNESS_DIR = \"./flatness\"\nif not os.path.exists(FLATNESS_DIR):\n    os.makedirs(FLATNESS_DIR)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.815Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Competition Metric","metadata":{}},{"cell_type":"code","source":"import sklearn.metrics\n\nclass ParticipantVisibleError(Exception):\n    # If you want an error message to be shown to participants, you must raise the error as a ParticipantVisibleError\n    # All other errors will only be shown to the competition host. This helps prevent unintentional leakage of solution data.\n    pass\n\n\ndef distance_metric(\n    solution    : pd.DataFrame,\n    submission  : pd.DataFrame,\n    thresh_ratio: float,\n    min_radius  : float,\n):\n    coordinate_cols = ['Motor axis 0', 'Motor axis 1', 'Motor axis 2']\n    label_tensor = solution[coordinate_cols].values.reshape(len(solution), -1, len(coordinate_cols))\n    predicted_tensor = submission[coordinate_cols].values.reshape(len(submission), -1, len(coordinate_cols))\n    # Find the minimum euclidean distances between the true and predicted points\n    solution['distance'] = np.linalg.norm(label_tensor - predicted_tensor, axis=2).min(axis=1)\n    # Convert thresholds from angstroms to voxels\n    solution['thresholds']  = solution['Voxel spacing'].apply(lambda x: (min_radius * thresh_ratio) / x)\n    solution['predictions'] = submission['Has motor'].values\n    solution.loc[(solution['distance'] > solution['thresholds']) & (solution['Has motor'] == 1) & (submission['Has motor'] == 1), 'predictions'] = 0\n    return solution['predictions'].values\n\n\ndef score(solution: pd.DataFrame, submission: pd.DataFrame, min_radius: float, beta: float) -> float:\n    \"\"\"\n    Parameters:\n    solution (pd.DataFrame): DataFrame containing ground truth motor positions.\n    submission (pd.DataFrame): DataFrame containing predicted motor positions.\n\n    Returns:\n    float: FBeta score.\n\n    Example\n    --------\n    >>> solution = pd.DataFrame({\n    ...     'tomo_id': [0, 1, 2, 3],\n    ...     'Motor axis 0': [-1, 250, 100, 200],\n    ...     'Motor axis 1': [-1, 250, 100, 200],\n    ...     'Motor axis 2': [-1, 250, 100, 200],\n    ...     'Voxel spacing': [10, 10, 10, 10],\n    ...     'Has motor': [0, 1, 1, 1]\n    ... })\n    >>> submission = pd.DataFrame({\n    ...     'tomo_id': [0, 1, 2, 3],\n    ...     'Motor axis 0': [100, 251, 600, -1],\n    ...     'Motor axis 1': [100, 251, 600, -1],\n    ...     'Motor axis 2': [100, 251, 600, -1]\n    ... })\n    >>> score(solution, submission, 1000, 2)\n    0.3571428571428571\n    \"\"\"\n\n    solution   = solution.sort_values('tomo_id').reset_index(drop=True)\n    submission = submission.sort_values('tomo_id').reset_index(drop=True)\n\n    filename_equiv_array = solution['tomo_id'].eq(submission['tomo_id'], fill_value=0).values\n\n    if np.sum(filename_equiv_array) != len(solution['tomo_id']):\n        raise ValueError('Submitted tomo_id values do not match the sample_submission file')\n\n    submission['Has motor'] = 1\n    # If any columns are missing an axis, it's marked with no motor\n    select = (submission[['Motor axis 0', 'Motor axis 1', 'Motor axis 2']] == -1).any(axis='columns')\n    submission.loc[select, 'Has motor'] = 0\n\n    cols = ['Has motor', 'Motor axis 0', 'Motor axis 1', 'Motor axis 2']\n    assert all(col in submission.columns for col in cols)\n\n    # Calculate a label of 0 or 1 using the 'has motor', and 'motor axis' values\n    predictions = distance_metric(\n        solution,\n        submission,\n        thresh_ratio=1.0,\n        min_radius=min_radius,\n    )\n\n    return sklearn.metrics.fbeta_score(solution['Has motor'].values, predictions, beta=beta)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.816Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Helper","metadata":{}},{"cell_type":"code","source":"def visualize_slice(images, masks, slice_idx, selected_index):\n    fig, ax = plt.subplots(1, 2, figsize=(10, 5))\n    ax[0].imshow(images[slice_idx], cmap=\"gray\")\n    ax[0].set_title(f\"Original Image {selected_index[slice_idx]}\")\n    ax[1].imshow(images[slice_idx], cmap=\"gray\")\n    ax[1].imshow(masks[slice_idx], cmap=\"jet\", alpha=0.7)\n    ax[1].set_title(f\"Mask Overlay {selected_index[slice_idx]}\")\n    plt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_model(check_point):\n    model_name = check_point.split(\"/\")[-1].split(\".\")[0]\n    print(f\"Model : {model_name}\")\n    _model = build_model()\n    state  = torch.load(check_point, map_location=device)\n    _model.load_state_dict(state[\"model\"], strict=False) # strict=False add new func else True\n    _model.eval()\n    \n    del state\n    _=gc.collect()\n    torch.cuda.empty_cache()\n    print(\"Successfully!!!\")\n    print()\n\n    return _model","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.816Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Data","metadata":{}},{"cell_type":"code","source":"if cf.download_dir is None:\n    train = pd.read_csv(os.path.join(cf.base_dir, \"train_labels.csv\"))\n    train.rename(columns={\n        \"Motor axis 0\"        : \"flagellum_z\",\n        \"Motor axis 1\"        : \"flagellum_y\",\n        \"Motor axis 2\"        : \"flagellum_x\",\n        \"Array shape (axis 0)\": \"num_slices\",\n        \"Array shape (axis 1)\": \"height\",\n        \"Array shape (axis 2)\": \"width\",\n        \"Number of motors\"    : \"Motor_count\"\n    }, inplace=True)\n    train[\"Has motor\"] = (train[\"Motor_count\"]>=1).astype(int)\n    \n    display(train[\"Has motor\"].value_counts())\nelse:\n    print(\"Download Train DF\")\n    train = pd.read_csv(os.path.join(cf.download_dir, \"train_data\", \"train_select.csv\"))\n    \n    train.rename(columns={\"Number of motors\"    : \"Motor_count\"}, inplace=True)\n    train[\"Has motor\"] = (train[\"Motor_count\"]>=1).astype(int)\n\ndisplay(train[\"Has motor\"].value_counts())\nsub = pd.read_csv(os.path.join(cf.base_dir, \"sample_submission.csv\"))\nprint(f\"Train Shape      : {train.shape}\")\nprint(f\"Submission Shape : {sub.shape}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if train.loc[(train[\"flagellum_y\"]==-1)&(train[\"flagellum_x\"]==-1)&(train[\"flagellum_z\"] != -1), [\"flagellum_z\", \"flagellum_y\", \"flagellum_x\"]].sum().all():\n    print(f\"Z is Miss!!!\")\n    train.loc[(train[\"flagellum_y\"]==-1)&(train[\"flagellum_x\"]==-1), \"flagellum_z\"] = -1.0\n    display(train[(train[\"flagellum_y\"]==-1)&(train[\"flagellum_x\"]==-1)])","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cf.debug:\n    display(train[train[\"tomo_id\"].isin([\"tomo_0333fa\", \"tomo_033ebe\", \"tomo_05df8a\"])])","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cf.download_dir is None:\n    if cf.USE_MEAN:\n        print(\"Mean\")\n        target_df = train.groupby(\"tomo_id\")[[\"flagellum_z\", \"flagellum_y\", \"flagellum_x\"]].mean().round().reset_index()\n    else:\n        print(\"Mode\")\n        target_df = train.groupby(\"tomo_id\")[[\"flagellum_z\", \"flagellum_y\", \"flagellum_x\"]].agg(lambda x: x.mode().iloc[0]).reset_index()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fixed_cols = [\"tomo_id\", \"num_slices\", \"height\", \"width\", \"Voxel spacing\", \"Motor_count\"]\n\nif cf.download_dir is None:\n    if cf.change_target:\n        print(\"Traget Change\")\n        train = train[~train[fixed_cols].duplicated()]\n        train = train.drop(columns={\"flagellum_z\", \"flagellum_y\", \"flagellum_x\"})\n        train = train.merge(target_df, on=\"tomo_id\", how=\"left\", suffixes=(\"\", \"_change\"))\n        train = train.rename(columns={\"flagellum_z_change\": \"flagellum_z\", \n                                      \"flagellum_y_change\": \"flagellum_y\", \n                                      \"flagellum_x_change\": \"flagellum_x\"})\n        \n        test  = train[train[\"tomo_id\"].isin(sub[\"tomo_id\"])].reset_index(drop=True)\n        train = train[~train[\"tomo_id\"].isin(sub[\"tomo_id\"])].reset_index(drop=True)\n    else:\n        print(\"Nomal\")\n        train   = train[~train[fixed_cols].duplicated()]\n        \ntest    = train[train[\"tomo_id\"].isin(sub[\"tomo_id\"])].reset_index(drop=True)\ntrain   = train[~train[\"tomo_id\"].isin(sub[\"tomo_id\"])].reset_index(drop=True)\n\n\ntrain.shape, test.shape","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_test_df(df):\n    for index, row in tqdm(df.iterrows(), total=len(df), desc=\"make test df\"):\n        row_id = row[\"tomo_id\"]\n        path   = os.path.join(cf.base_dir, \"test\", row_id, \"*\")\n        paths  = sorted(glob.glob(path))\n        imgs   = cv2.imread(paths[0])\n        h, w,_= imgs.shape\n        df.loc[index, \"num_slices\"] = len(paths)\n        df.loc[index, \"height\"]     = h\n        df.loc[index, \"width\"]      = w\n    df[\"num_slices\"] = df[\"num_slices\"].astype(int)\n    df[\"height\"]     = df[\"height\"].astype(int)\n    df[\"width\"]      = df[\"width\"].astype(int)\n    return df","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test = get_test_df(test)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.816Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Fold","metadata":{}},{"cell_type":"code","source":"if cf.download_dir is None:\n    drop_id = train.loc[train[\"Motor_count\"]==5, \"tomo_id\"].values[0]\n    print(drop_id)\n    train = train[~train[\"tomo_id\"].isin([drop_id])].reset_index(drop=True)\n    print(train.shape)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train[\"Motor_count\"].value_counts()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_fold(df, n_splits=cf.N_fold, random_state=cf.SEED):\n    skf = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=random_state)\n    for e, (train_index, valid_index) in enumerate(skf.split(df, df[\"Motor_count\"])):\n        df.loc[valid_index, \"fold\"] = e\n    df[\"fold\"] = df[\"fold\"].astype(int)\n    return df","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = get_fold(train)\ntrain.groupby(\"fold\").size()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cf.debug:\n    display(train.groupby(\"fold\")[\"Motor_count\"].value_counts().sort_index())","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Preprocessing","metadata":{}},{"cell_type":"code","source":"# def normalize_slice(slice_data, lower_percentile=0.1, upper_percentile=99.9):\n#     \"\"\"\n#     パーセンタイルベースで正規化 (0.1%〜99.9%の範囲に収めて0-1スケール)\n#     \"\"\"\n#     lower = np.percentile(slice_data, lower_percentile)\n#     upper = np.percentile(slice_data, upper_percentile)\n    \n#     slice_data = np.clip(slice_data, lower, upper)\n#     slice_data = (slice_data - lower) / (upper - lower + 1e-8)  # ゼロ除算防止\n#     return (slice_data * 255).astype(np.uint8)/255.0\n#     #return slice_data.astype(np.float32)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def normalize_slice(slice_data, \n                    clip_limit=4.0, \n                    tile_grid_size=(16,16)\n                   ):\n    \"\"\"\n    Normalize slice data using CLAHE to enhance low-contrast features like flagellum\n    \"\"\"\n    # 画像の正規化 (0-255)\n    slice_data = (slice_data - np.min(slice_data)) / (np.max(slice_data) - np.min(slice_data)) * 255\n    slice_data = slice_data.astype(np.uint8)\n\n    # CLAHE を適用\n    clahe    = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid_size)\n    enhanced = clahe.apply(slice_data)\n    enhanced = np.clip(enhanced, 0, 255)\n    return enhanced.astype(np.float32) / 255.0","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_resize(imgs, imsize):\n    imgs = [np.array(cv2.resize(img, (imsize, imsize))) for img in imgs]\n    _min, _max = np.min(imgs), np.max(imgs)\n    imgs = (imgs-_min)/(_max-_min+1e-6)\n    imgs = (imgs*255).astype(np.uint8)\n    return imgs\n\ndef img_resize(img_path, imsize):\n    imgs = np.stack([cv2.imread(path, cv2.IMREAD_GRAYSCALE) for path in img_path], axis=0)\n    orig = imgs.shape[0]\n    imgs = [np.array(cv2.resize(img, (imsize, imsize))) for img in imgs]\n    _min, _max = np.min(imgs), np.max(imgs)\n    imgs = (imgs-_min)/(_max-_min+1e-6)\n    imgs = (imgs*255).astype(np.uint8)\n    return imgs, orig\n\ndef sample_slices(volume, selected_indices=None, target_slices=cf.TARGET_SLICES, center_ratio=0.5, mode=\"train\", apply_clahe=True):\n    num_slices, height, width = volume.shape\n\n    if num_slices <= target_slices:\n        # スライス数が足りない場合、補間\n        scale_z   = target_slices / num_slices\n        resampled = scipy.ndimage.zoom(volume, (scale_z, 1, 1), order=1)\n        if mode != \"train\" and selected_indices is None:\n            selected_indices = np.linspace(0, num_slices - 1, target_slices, dtype=int)  # 補間時の仮の対応\n        return resampled, selected_indices\n\n    # 各部分のスライス数を計算\n    n_center = int(target_slices * center_ratio)  # 中央部分\n    n_side   = (target_slices - n_center) // 2    # 両端部分\n\n    # スライスのインデックスを計算\n    start_indices  = np.linspace(0, num_slices // 4, n_side, dtype=int)\n    center_indices = np.linspace(num_slices // 4, 3 * num_slices // 4, n_center, dtype=int)\n    end_indices    = np.linspace(3 * num_slices // 4, num_slices - 1, n_side, dtype=int)\n    \n    selected       = np.unique(np.concatenate([start_indices, center_indices, end_indices]))\n\n    # スライス数が target_slices より少なくなるのを防ぐ\n    while len(selected) < target_slices:\n        selected = np.append(selected, selected[-1])  # 最後のスライスを繰り返す\n    \n    resampled = volume[selected[:target_slices]]\n    if apply_clahe:\n        resampled = np.stack([normalize_slice(im) for im in resampled], axis=0)\n    else:\n        resampled = np.stack(resampled, axis=0).astype(np.float32)\n\n    if mode == \"train\" and selected_indices is not None:\n        selected = selected_indices[selected[:target_slices]]\n    return resampled, selected\n\ndef load_data(tomo_id, IMSIZE=cf.IMSIZE, target_slices=cf.TARGET_SLICES, center_ratio=0.5, mode=\"train\"):\n    if mode == \"train\":\n        img_paths  = sorted(glob.glob(os.path.join(cf.download_dir, \"train_img\", tomo_id, \"*.png\")))\n        mask_paths = sorted(glob.glob(os.path.join(cf.download_dir, \"mask_img\",  tomo_id, \"*.png\")))\n        img_select_ind = np.array([int(os.path.basename(f).split(\"_\")[1].split(\".\")[0]) for f in img_paths])\n        \n        images = np.array([cv2.imread(p, cv2.IMREAD_GRAYSCALE) for p in img_paths])\n        masks  = np.array([cv2.imread(p, cv2.IMREAD_GRAYSCALE) for p in mask_paths])\n        if images.shape[1] != IMSIZE:\n            images = get_resize(images, imsize=IMSIZE)\n            masks  = get_resize(masks,  imsize=IMSIZE)\n            \n        images, img_selected_indices  = sample_slices(images, selected_indices=img_select_ind, target_slices=target_slices, center_ratio=center_ratio)\n        masks,  _ = sample_slices(masks,  target_slices=target_slices, center_ratio=center_ratio, apply_clahe=False)\n        return images, masks, img_selected_indices\n    else:\n        img_paths = sorted(glob.glob(os.path.join(cf.base_dir, \"test\", tomo_id, \"*.jpg\")))\n        images, orignal_slices   = img_resize(img_paths, imsize=IMSIZE)\n        images, selected_indices = sample_slices(images, target_slices=target_slices,  center_ratio=center_ratio, mode=mode)\n        return images, orignal_slices, selected_indices","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### DataSet","metadata":{}},{"cell_type":"code","source":"def compute_padding(images, final_size):\n    \"\"\"\n    images    : (C, D, H, W)\n    final_size: 最終的に揃えたい H, W (ここでは Dはそのままとする)\n    \"\"\"\n    _, D, H, W = images.shape\n\n    pad_H = max(final_size - H, 0)\n    pad_W = max(final_size - W, 0)\n\n    pad_top    = pad_H // 2\n    pad_bottom = pad_H - pad_top\n    pad_left   = pad_W // 2\n    pad_right  = pad_W - pad_left\n\n    # D方向はここではpadしない前提\n    pad_front = 0\n    pad_back  = 0\n\n    return pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def crop_near_motor_center(images, \n                           masks, \n                           labels,\n                           pad_left, \n                           pad_right, \n                           pad_top, \n                           pad_bottom, \n                           pad_front, \n                           pad_back,\n                           crop_size, \n                           final_size,\n                           min_mask_ratio=0.01, \n                           max_tries=5,\n                           target_thersold=0.097\n                          ):\n    _, D, H, W = images.shape\n    z, y, x = labels.numpy()\n    crop_H, crop_W, crop_D = final_size\n\n    label_valid = (x >= target_thersold) and (y >= target_thersold) and (z >= target_thersold)\n\n    # Padding\n    images = F.pad(images, (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back), mode='reflect')\n    masks  = F.pad(masks, (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back), mode='reflect')\n    D_p, H_p, W_p = images.shape[1:]\n\n    half_crop_D = crop_D // 2\n    half_crop_H = crop_H // 2\n    half_crop_W = crop_W // 2\n\n    for _ in range(max_tries):\n        if label_valid:\n            x_p = x + pad_left\n            y_p = y + pad_top\n            z_p = z + pad_front\n            center_x = int(x_p)\n            center_y = int(y_p)\n            center_z = int(z_p)\n        else:\n            center_x = W_p // 2\n            center_y = H_p // 2\n            center_z = D_p // 2\n\n        start_x = max(center_x - half_crop_W, 0)\n        start_y = max(center_y - half_crop_H, 0)\n        start_z = max(center_z - half_crop_D, 0)\n\n        end_x = min(start_x + crop_W, W_p)\n        end_y = min(start_y + crop_H, H_p)\n        end_z = min(start_z + crop_D, D_p)\n\n        if end_x - start_x < crop_W:\n            start_x = max(end_x - crop_W, 0)\n        if end_y - start_y < crop_H:\n            start_y = max(end_y - crop_H, 0)\n        if end_z - start_z < crop_D:\n            start_z = max(end_z - crop_D, 0)\n\n        cropped_images = images[:, start_z:end_z, start_y:end_y, start_x:end_x]\n        cropped_masks  = masks[:,  start_z:end_z, start_y:end_y, start_x:end_x]\n\n        mask_mean = cropped_masks.mean()\n        if mask_mean > min_mask_ratio:\n            break\n        else:\n            # ランダム中心に変更\n            if (W_p > crop_W) and (H_p > crop_H) and (D_p > crop_D):\n                center_x = random.randint(half_crop_W, W_p - half_crop_W)\n                center_y = random.randint(half_crop_H, H_p - half_crop_H)\n                center_z = random.randint(half_crop_D, D_p - half_crop_D)\n            else:\n                center_x = W_p // 2\n                center_y = H_p // 2\n                center_z = D_p // 2\n\n    # Resize\n    cropped_images = F.interpolate(cropped_images.unsqueeze(0), size=(crop_D, crop_H, crop_W), mode=\"trilinear\", align_corners=False).squeeze(0)\n    cropped_masks  = F.interpolate(cropped_masks.unsqueeze(0),  size=(crop_D, crop_H, crop_W), mode=\"trilinear\", align_corners=False).squeeze(0)\n\n    if label_valid:\n        new_x = (x_p - start_x) * (crop_W / crop_W)\n        new_y = (y_p - start_y) * (crop_H / crop_H)\n        new_z = (z_p - start_z) * (crop_D / crop_D)\n        new_labels = torch.tensor([new_z, new_y, new_x], dtype=torch.float32)\n    else:\n        new_labels = torch.tensor([-1, -1, -1], dtype=torch.float32)\n\n    return cropped_images, cropped_masks, new_labels","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, \n                 df, \n                 transform=None, \n                 target_slices=cf.TARGET_SLICES, \n                 center_ratio=0.5, \n                 mode=\"train\"\n                ):\n        self.df   = df\n        self.mode = mode\n        self.target_slices = target_slices\n        self.center_ratio  = center_ratio\n        self.transform     = transform\n\n    def __len__(self):\n        return len(self.df[\"tomo_id\"].unique())\n\n    def __getitem__(self, index):\n        tomo_id = self.df[\"tomo_id\"].unique()[index]\n        row     = self.df[self.df[\"tomo_id\"] == tomo_id].iloc[0]\n        \n        if self.mode == \"train\":\n            # ======== Train mode (Cropあり / maskあり / Augmentationあり) ========\n            original_slices = row[\"num_slices\"]\n            z, y, x         = row[\"flagellum_z\"], row[\"flagellum_y\"], row[\"flagellum_x\"]\n            voxel           = row[\"Voxel spacing\"]\n            H, W            = row[\"height\"], row[\"width\"]\n            has_motor       = row[\"Has motor\"]\n            motor_count     = row[\"Motor_count\"]\n            \n            images, masks, selected_indices = load_data(tomo_id, target_slices=self.target_slices, center_ratio=self.center_ratio)\n            masks = np.clip(masks, 0, 255) / 255.0\n\n            images = torch.tensor(images).unsqueeze(0).float()  # (C, D, H, W)\n            masks  = torch.tensor(masks).unsqueeze(0).float()   # (C, D, H, W)\n            labels = torch.tensor([z, y, x], dtype=torch.float32)\n            selected_indices = torch.tensor(selected_indices, dtype=torch.long)\n            \n            crop_H, crop_W, crop_D = cf.CROP_SIZE\n        \n            pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back = compute_padding(images, final_size=crop_H)\n            \n            # ----- 補正済みラベル（crop/transform前提）-----\n            try:\n                z_index = selected_indices.tolist().index(int(z))\n            except ValueError:\n                z_index = -1  # 無効値\n            \n            y_scaled = y * (images.shape[2] / H)\n            x_scaled = x * (images.shape[3] / W)\n\n            resize_labels = torch.tensor([z_index, y_scaled, x_scaled], dtype=torch.float32)\n        \n            if cf.USE_CROP:\n                if random.random() < cf.CROP_THRESOLD:\n                    images, masks, labels = crop_near_motor_center(images, \n                                                                   masks, \n                                                                   resize_labels,\n                                                                   pad_left=pad_left,\n                                                                   pad_right=pad_right,\n                                                                   pad_top=pad_top,\n                                                                   pad_bottom=pad_bottom,\n                                                                   pad_front=pad_front,\n                                                                   pad_back=pad_back,\n                                                                   crop_size=crop_H,\n                                                                   final_size=cf.CROP_SIZE,\n                                                                  )\n                else:\n                    # cropしなかった場合も、リサイズだけはする\n                    padded_image = F.pad(images, (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back), mode='reflect')\n                    padded_mask  = F.pad(masks,  (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back), mode='reflect')\n\n                    # (batch, channel, depth, height, width) の5D\n                    images = F.interpolate(padded_image.unsqueeze(0), size=(crop_D, crop_H, crop_W), mode=\"trilinear\", align_corners=False).squeeze(0)\n                    masks  = F.interpolate(padded_mask.unsqueeze(0),  size=(crop_D, crop_H, crop_W), mode=\"trilinear\", align_corners=False).squeeze(0)\n                \n\n            if self.transform:\n                images, masks, resize_labels = self.transform(images, \n                                                              masks, \n                                                              resize_labels,\n                                                              pad_left=pad_left, \n                                                              pad_right=pad_right,\n                                                              pad_top=pad_top, \n                                                              pad_bottom=pad_bottom,\n                                                              pad_front=pad_front, \n                                                              pad_back=pad_back,\n                                                              final_size=cf.CROP_SIZE\n                                                             )\n\n            # # 高確度の中心のみで判定\n            # core_ratio = (masks > 0.98).float().mean().item()\n            \n            # # 周辺ぼかし含めた全体の非ゼロ量\n            # extended_ratio = (masks > 0.097).float().mean().item()\n            \n            # if core_ratio > 0.0001:  # 0.001 → 0.0001 に緩和\n            #     motor_count = 1\n            # elif extended_ratio > 0.005:  # 0.01 → 0.005 に緩和\n            #     motor_count = 1\n            # else:\n            #     motor_count = 0\n            \n            if masks.max() > 0.1:  # 最大値が小さすぎなければ motor あり\n                motor_count = 1\n            else:\n                motor_count = 0\n            \n            return {\n                \"tomo_id\"         : tomo_id,\n                \"images\"          : images, \n                \"masks\"           : masks,\n                \"orig_slices\"     : original_slices, \n                \"labels\"          : labels,\n                \"h\"               : H,\n                \"w\"               : W,\n                \"voxel_spacing\"   : voxel,\n                \"has_motor\"       : has_motor,\n                \"Motor_count\"     : motor_count,\n                \"selected_indices\": selected_indices\n            }\n\n        elif self.mode == \"valid\":\n            # ======== Valid mode (Cropなし / maskあり / Augmentationなし) ========\n            original_slices = row[\"num_slices\"]\n            z, y, x         = row[\"flagellum_z\"], row[\"flagellum_y\"], row[\"flagellum_x\"]\n            voxel           = row[\"Voxel spacing\"]\n            H, W            = row[\"height\"], row[\"width\"]\n            has_motor       = row[\"Has motor\"]\n            motor_count     = row[\"Motor_count\"]\n            \n            images, masks, selected_indices = load_data(tomo_id, target_slices=self.target_slices, center_ratio=self.center_ratio)\n            masks = np.clip(masks, 0, 255) / 255.0\n\n            images = torch.tensor(images).unsqueeze(0).float()\n            masks  = torch.tensor(masks).unsqueeze(0).float()\n\n            labels           = torch.tensor([z, y, x], dtype=torch.float32)\n            selected_indices = torch.tensor(selected_indices, dtype=torch.long)\n\n            if cf.USE_CROP:\n                crop_H, crop_W, crop_D = cf.CROP_SIZE\n                pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back = compute_padding(images, final_size=crop_H)\n                \n                padded_image = F.pad(images, (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back), mode='reflect')\n                padded_mask  = F.pad(masks,  (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back), mode='reflect')\n                images = F.interpolate(padded_image.unsqueeze(0), size=(crop_D, crop_H, crop_W), mode=\"trilinear\", align_corners=False).squeeze(0)\n                masks  = F.interpolate(padded_mask.unsqueeze(0),  size=(crop_D, crop_H, crop_W), mode=\"trilinear\", align_corners=False).squeeze(0)\n\n            return {\n                \"tomo_id\"         : tomo_id,\n                \"images\"          : images, \n                \"masks\"           : masks,\n                \"orig_slices\"     : original_slices, \n                \"labels\"          : labels,\n                \"h\"               : H,\n                \"w\"               : W,\n                \"voxel_spacing\"   : voxel,\n                \"has_motor\"       : has_motor,\n                \"Motor_count\"     : motor_count,\n                \"selected_indices\": selected_indices\n            }\n\n        elif self.mode == \"test\":\n            # ======== Test mode (Cropなし / maskなし / Augmentationなし) ========\n            H, W = row[\"height\"], row[\"width\"]\n            images, original_slices, selected_indices = load_data(tomo_id, target_slices=self.target_slices, center_ratio=self.center_ratio, mode=\"test\")\n\n            images = torch.tensor(images).unsqueeze(0).float()\n            selected_indices = torch.tensor(selected_indices, dtype=torch.long)\n\n            if cf.USE_CROP:\n                crop_H, crop_W, crop_D = cf.CROP_SIZE\n                pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back = compute_padding(images, final_size=crop_H)\n                \n                padded_image = F.pad(images, (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back), mode='reflect')\n                images = F.interpolate(padded_image.unsqueeze(0), size=(crop_D, crop_H, crop_W), mode=\"trilinear\", align_corners=False).squeeze(0)\n\n            return {\n                \"tomo_id\"         : tomo_id, \n                \"images\"          : images, \n                \"orig_slices\"     : original_slices,\n                \"h\"               : H,\n                \"w\"               : W,\n                \"selected_indices\": selected_indices\n            }\n\n        else:\n            raise ValueError(f\"Unknown mode: {self.mode}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def augment_3d(images, \n               masks, \n               labels, \n               pad_left,\n               pad_right,\n               pad_top, \n               pad_bottom,\n               pad_front, \n               pad_back,\n               final_size,\n               p=cf.train_p\n              ):\n    \n    _, D, H, W = images.shape\n    z, y, x    = labels.numpy()\n    crop_H, crop_W, crop_D = final_size\n\n    # ラベルが有効か判定\n    label_valid = (x >= 0) and (y >= 0) and (z >= 0)\n\n    # 左右反転（X軸）\n    if random.random() < p:\n        images = torch.flip(images, dims=[3])\n        masks  = torch.flip(masks,  dims=[3])\n        if label_valid:\n            x = W - 1 - x  # flip後のx座標更新\n\n    # 上下反転（Y軸）\n    if random.random() < p:\n        images = torch.flip(images, dims=[2])\n        masks  = torch.flip(masks,  dims=[2])\n        if label_valid:\n            y = H - 1 - y  # flip後のy座標更新\n\n    # 前後反転（Z軸）\n    if random.random() < p:\n        images = torch.flip(images, dims=[1])\n        masks  = torch.flip(masks,  dims=[1])\n        if label_valid:\n            z = D - 1 - z  # flip後のz座標更新\n\n    # 90度回転（X-Y平面）\n    if random.random() < p:\n        images = images.rot90(1, dims=[2, 3])\n        masks  = masks.rot90(1, dims=[2, 3])\n        if label_valid:\n            x, y = y, W - 1 - x  # (x,y)のスワップ＋更新\n\n    # ============ If need ============\n    # # Gaussian Noise\n    if random.random() < p:\n        noise = torch.randn_like(images) * 0.05\n        images = images + noise\n\n    # # Contrast Adjustment\n    if random.random() < p:\n        factor = random.uniform(0.75, 1.25)\n        mean = images.mean()\n        images = (images - mean) * factor + mean\n\n    # # Cutout (ランダムに隠す)\n    if random.random() < p:\n        cutout_size = 32\n        zc = random.randint(0, D - 1)\n        yc = random.randint(0, H - cutout_size)\n        xc = random.randint(0, W - cutout_size)\n        images[:, zc, yc:yc+cutout_size, xc:xc+cutout_size] = 0\n    # ==================================================\n\n    # reflect pad\n    images = F.pad(images, (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back), mode='reflect')\n    masks  = F.pad(masks,  (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back), mode='reflect')\n\n    # Resize\n    images = F.interpolate(images.unsqueeze(0), size=(crop_D, crop_H, crop_W), mode=\"trilinear\", align_corners=False).squeeze(0)\n    masks  = F.interpolate(masks.unsqueeze(0),  size=(crop_D, crop_H, crop_W), mode=\"trilinear\", align_corners=False).squeeze(0)\n\n    # Label scaling\n    if label_valid:\n        y = y * (crop_H / H)\n        x = x * (crop_W / W)\n        labels = torch.tensor([z, y, x], dtype=torch.float32)\n    else:\n        labels = torch.tensor([-1, -1, -1], dtype=torch.float32)\n\n    return images, masks, labels","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### DataLoader","metadata":{}},{"cell_type":"code","source":"class BalancedBatchSampler(Sampler):\n    def __init__(self, df, batch_size, min_pos=cf.MIN_POS):\n        assert batch_size > min_pos, f\"batch_size({batch_size}) must be > min_pos({min_pos})\"\n        self.batch_size = batch_size\n        self.min_pos    = min_pos\n        self.max_pos    = batch_size - 1  # 少なくとも 1つは neg を含むように制限\n\n        self.pos_ids = df[df[\"Motor_count\"] >= 1][\"tomo_id\"].unique().tolist()\n        self.neg_ids = df[df[\"Motor_count\"] == 0][\"tomo_id\"].unique().tolist()\n\n        self.all_ids = df[\"tomo_id\"].unique().tolist()\n        self.pos_indices = [self.all_ids.index(tid) for tid in self.pos_ids if tid in self.all_ids]\n        self.neg_indices = [self.all_ids.index(tid) for tid in self.neg_ids if tid in self.all_ids]\n\n        # conservatively estimate number of batches\n        self.num_batches = len(self.pos_indices) // self.min_pos\n\n    def __len__(self):\n        return self.num_batches\n\n    def __iter__(self):\n        pos = self.pos_indices.copy()\n        neg = self.neg_indices.copy()\n        random.shuffle(pos)\n        random.shuffle(neg)\n\n        pos_i = 0\n        neg_i = 0\n\n        for _ in range(self.num_batches):\n            # ランダムに正例の数を決定\n            num_pos = random.randint(self.min_pos, min(self.max_pos, self.batch_size - 1))\n\n            # pos が足りなければシャッフルして再利用\n            if pos_i + num_pos > len(pos):\n                random.shuffle(pos)\n                pos_i = 0\n            pos_batch = pos[pos_i: pos_i + num_pos]\n            pos_i += num_pos\n\n            # 残りは neg で埋める\n            num_neg = self.batch_size - num_pos\n            if neg_i + num_neg > len(neg):\n                random.shuffle(neg)\n                neg_i = 0\n            neg_batch = neg[neg_i: neg_i + num_neg]\n            neg_i += num_neg\n\n            batch = pos_batch + neg_batch\n            random.shuffle(batch)\n            yield batch","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def custom_collate(batch):\n    batch_dict = {}\n    for key in batch[0]:\n        if isinstance(batch[0][key], torch.Tensor):\n            batch_dict[key] = torch.stack([b[key] for b in batch])\n        elif isinstance(batch[0][key], (int, float)):\n            batch_dict[key] = torch.tensor([b[key] for b in batch])\n        else:\n            batch_dict[key] = [b[key] for b in batch]\n    return batch_dict","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_dataloaders(df, fold):\n    train = df[df[\"fold\"] != fold].reset_index(drop=True)\n    valid = df[df[\"fold\"] == fold].reset_index(drop=True)\n\n    print(f\"Train Length : {len(train):,}\")\n    print(f\"Valid Length : {len(valid):,}\")\n    \n    train_dataset = CustomDataset(train, transform=augment_3d)\n    \n    if cf.valid_aug:\n        print(f\"Valid Augmentation : {cf.valid_aug}\")\n        valid_transform = partial(augment_3d, p=cf.valid_p)\n        valid_dataset   = CustomDataset(valid, transform=valid_transform)\n    else:\n        print(f\"Valid Augmentation : {cf.valid_aug}\")\n        valid_dataset = CustomDataset(valid, mode=\"valid\")\n\n    sampler = BalancedBatchSampler(train, batch_size=cf.train_batch_size)\n\n    dataloaders = dict()\n\n    dataloaders[\"train\"] = DataLoader(train_dataset, \n                                      #batch_size=cf.batch_size, # used sampler batch_size Off\n                                      #shuffle=False,  # if used batch_sampler off\n                                      batch_sampler=sampler,\n                                      num_workers=cf.num_workers,\n                                      pin_memory=True, \n                                      #drop_last=True, # if used batch_sampler off\n                                      collate_fn=custom_collate\n                                     )\n    dataloaders[\"valid\"] = DataLoader(valid_dataset, \n                                      batch_size=cf.valid_batch_size, \n                                      shuffle=False, \n                                      num_workers=cf.num_workers,\n                                      pin_memory=True,\n                                      drop_last=False\n                                     )\n    return dataloaders\n\ndef get_dataloaders_test(df):\n    print(f\"Test Shape : {df.shape}\")\n    test_dataset = CustomDataset(df, mode=\"test\")\n    dataloaders  = dict()\n    \n    dataloaders[\"test\"] = DataLoader(test_dataset, \n                                     batch_size=cf.valid_batch_size, \n                                     shuffle=False, \n                                     num_workers=cf.num_workers,\n                                     pin_memory=True, \n                                     drop_last=False\n                                    )\n    return dataloaders","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Model","metadata":{}},{"cell_type":"code","source":"# --- Encoder ----\nclass DropPath(nn.Module):\n    def __init__(self, drop_prob=0.1):\n        super().__init__()\n        self.drop_prob = drop_prob\n\n    def forward(self, x):\n        if not self.training or self.drop_prob == 0.0:\n            return x\n        keep_prob = 1 - self.drop_prob\n        mask      = torch.rand(x.shape[0], 1, 1, 1, 1, device=x.device) < keep_prob\n        return x * mask / keep_prob\n\n\nclass ConvNextBlock3D(nn.Module):\n    def __init__(self, dim, dropout_rate=0.1, use_droppath=True):\n        super().__init__()\n        self.depthwise = nn.Conv3d(dim, dim, kernel_size=3, padding=1, groups=dim, bias=False)\n        self.norm      = nn.BatchNorm3d(dim)\n        self.pointwise = nn.Sequential(\n            nn.Conv3d(dim, dim // 2, kernel_size=1, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(dim // 2, dim, kernel_size=1, bias=False)\n        )\n        self.dropout      = nn.Dropout3d(p=dropout_rate)\n        self.use_droppath = use_droppath\n        self.droppath     = DropPath(drop_prob=dropout_rate) if use_droppath else nn.Identity()\n\n    def forward(self, x):\n        residual = x\n        x = self.depthwise(x)\n        x = self.norm(x)\n        x = self.pointwise(x)\n        x = self.dropout(x)\n        x = self.droppath(x)\n        return x + residual\n\nclass CBAM3D(nn.Module):\n    def __init__(self, channels, reduction_ratio=cf.TARGET_SLICES//2, kernel_size=7):\n        super().__init__()\n        # Channel Attention\n        self.max_pool   = nn.AdaptiveMaxPool3d(1)\n        self.avg_pool   = nn.AdaptiveAvgPool3d(1)\n        self.shared_mlp = nn.Sequential(\n            nn.Conv3d(channels, channels // reduction_ratio, 1, bias=False),\n            nn.ReLU(),\n            nn.Conv3d(channels // reduction_ratio, channels, 1, bias=False)\n        )\n        self.sigmoid_channel = nn.Sigmoid()\n\n        # Spatial Attention\n        self.conv_spatial    = nn.Conv3d(2, 1, kernel_size=kernel_size, padding=kernel_size // 2, bias=False)\n        self.sigmoid_spatial = nn.Sigmoid()\n\n    def forward(self, x):\n        # Channel attention\n        max_out = self.shared_mlp(self.max_pool(x))\n        avg_out = self.shared_mlp(self.avg_pool(x))\n        channel_att = self.sigmoid_channel(max_out + avg_out)\n        x = x * channel_att\n\n        # Spatial attention\n        max_out, _  = torch.max(x,  dim=1, keepdim=True)\n        avg_out     = torch.mean(x, dim=1, keepdim=True)\n        spatial_att = self.sigmoid_spatial(self.conv_spatial(torch.cat([avg_out, max_out], dim=1)))\n        x = x * spatial_att\n        return x\n\nclass SEBlock3D(nn.Module):\n    def __init__(self, channels, reduction=8):\n        super().__init__()\n        reduced_channels = max(1, channels // reduction)  # 少なくとも1は確保\n        self.avg_pool = nn.AdaptiveAvgPool3d(1)\n        self.fc = nn.Sequential(\n            nn.Linear(channels, reduced_channels, bias=True),\n            nn.ReLU(inplace=True),\n            nn.Linear(reduced_channels, channels, bias=True),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        b, c, d, h, w = x.size()\n        y = self.avg_pool(x).view(b, c)\n        y = self.fc(y).view(b, c, 1, 1, 1)\n        return x * y\n\nclass Conv3DEncoder(nn.Module):\n    def __init__(self, \n                 in_channels=1, \n                 dims=[cf.TARGET_SLICES,\n                       cf.TARGET_SLICES*2, \n                       cf.TARGET_SLICES*4, \n                       cf.TARGET_SLICES*6\n                      ], \n                 dropout_rate=cf.dropout_encoder,\n                 z_sampling=cf.z_sampling\n                ):\n        super().__init__()\n        \n        if z_sampling:\n            self.stem = nn.Conv3d(in_channels, dims[0], kernel_size=3, stride=2, padding=1)\n            self.downsample_layers = nn.ModuleList([\n                nn.Identity(),\n                nn.Conv3d(dims[0], dims[1], kernel_size=3, stride=2, padding=1),\n                nn.Conv3d(dims[1], dims[2], kernel_size=3, stride=2, padding=1),\n                nn.Conv3d(dims[2], dims[3], kernel_size=3, stride=2, padding=1)\n            ])\n        else:\n            self.stem = nn.Conv3d(in_channels, dims[0], kernel_size=(1,3,3), stride=(1,2,2), padding=(0,1,1))\n            self.downsample_layers = nn.ModuleList([\n                nn.Identity(),\n                nn.Conv3d(dims[0], dims[1], kernel_size=(1,3,3), stride=(1,2,2), padding=(0,1,1)),\n                nn.Conv3d(dims[1], dims[2], kernel_size=(1,3,3), stride=(1,2,2), padding=(0,1,1)),\n                nn.Conv3d(dims[2], dims[3], kernel_size=(1,3,3), stride=(1,2,2), padding=(0,1,1))\n            ])\n        self.stages = nn.ModuleList([\n            nn.Sequential(*[ConvNextBlock3D(dim, dropout_rate, use_droppath=cf.use_droppath) for _ in range(2)]) for dim in dims\n        ])\n        \n        #self.norm = nn.BatchNorm3d(dims[-1])\n        self.norm = nn.GroupNorm(num_groups=8, num_channels=dims[-1]) #8~16\n\n        # CBAM(Convolutional Block Attention Module)\n        self.cbam_f1 = CBAM3D(dims[0])\n        self.cbam_f2 = CBAM3D(dims[1])\n        self.cbam_f3 = CBAM3D(dims[2])\n        self.cbam_f4 = CBAM3D(dims[3])\n        #SEBlock 2stage only\n        self.se_f1 = SEBlock3D(dims[0], reduction=4)\n        self.se_f2 = SEBlock3D(dims[1], reduction=8)\n        self.se_f3 = SEBlock3D(dims[2], reduction=4)#or reduction=4\n        self.se_f4 = SEBlock3D(dims[3], reduction=4)\n\n    def forward(self, x):\n        f0 = self.stem(x)\n        f1 = self.stages[0](f0)\n        #f1 = self.se_f1(f1)\n        #f1 = self.cbam_f1(f1)  # CBAM 1stage on\n        \n        f2 = self.stages[1](self.downsample_layers[1](f1))\n        f2 = self.se_f2(f2)\n        #f2 = self.cbam_f2(f2) # CBAM #1satge off\n        \n        f3 = self.stages[2](self.downsample_layers[2](f2))\n        f3 = self.se_f3(f3)\n        #f3 = self.cbam_f3(f3) # CBAM #1stage off\n        \n        f4 = self.norm(self.stages[3](self.downsample_layers[3](f3)))\n        f4 = self.se_f4(f4)\n        #f4 = self.cbam_f4(f4)\n        return [f1, f2, f3, f4]\n\n\n# --- Decoder ---\nclass SEBlock(nn.Module):\n    def __init__(self, in_channels, reduction=32):\n        super().__init__()\n        self.global_avg_pool = nn.AdaptiveAvgPool3d(1)\n        self.fc = nn.Sequential(\n            nn.Linear(in_channels, in_channels // reduction, bias=False),\n            nn.ReLU(),\n            nn.Linear(in_channels // reduction, in_channels, bias=False),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        b, c, _, _, _ = x.size()\n        y = self.global_avg_pool(x).view(b, c)\n        y = self.fc(y).view(b, c, 1, 1, 1)\n        return x * y.expand_as(x)\n\n# class CBAM(nn.Module):\n#     def __init__(self, in_channels, reduction=cf.TARGET_SLICES):\n#         super().__init__()\n#         self.se      = SEBlock(in_channels, reduction)\n#         self.spatial = nn.Conv3d(1, 1, kernel_size=3, padding=1, bias=False)\n\n#     def forward(self, x):\n#         x = self.se(x)\n#         spatial_attn = torch.mean(x, dim=1, keepdim=True)\n#         spatial_attn = torch.sigmoid(self.spatial(spatial_attn))\n#         return x * spatial_attn\n\nclass CBAM(nn.Module):\n    def __init__(self, in_channels, reduction_ratio=16):\n        # reduction 32 or 16\n        super().__init__()\n        \n        # reduction_ratio=16ではなく、チャネルに依存してadaptiveに\n        reduced_channels = max(4, min(64, in_channels // reduction_ratio))\n\n        # Channel Attention using 1x1x1 Conv (instead of Linear)\n        self.channel_attn = nn.Sequential(\n            nn.AdaptiveAvgPool3d(1),  # → (B, C, 1, 1, 1)\n            nn.Conv3d(in_channels, reduced_channels, kernel_size=1, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(reduced_channels, in_channels, kernel_size=1, bias=False),\n            nn.Sigmoid()\n        )\n\n        # Spatial Attention\n        self.spatial_attn = nn.Sequential(\n            nn.Conv3d(1, 1, kernel_size=3, padding=1, bias=False),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        # Channel attention\n        ca = self.channel_attn(x)\n        x  = x * ca\n\n        # Spatial attention (channel mean)\n        sa_input = torch.mean(x, dim=1, keepdim=True)  # → (B, 1, D, H, W)\n        sa = self.spatial_attn(sa_input)\n        x  = x * sa\n        return x\n\n        # Spatial Attention\n        avg_out = torch.mean(x, dim=1, keepdim=True)\n        max_out, _ = torch.max(x, dim=1, keepdim=True)\n        spatial = torch.cat([avg_out, max_out], dim=1)\n        spatial = self.sigmoid_spatial(self.conv_spatial(spatial))\n\n        return x * spatial\n\nclass DepthwiseConv3D(nn.Module):\n    def __init__(self, in_channels):\n        super().__init__()\n        self.depthwise = nn.Conv3d(in_channels, in_channels, kernel_size=3, padding=1, groups=in_channels)\n        self.relu      = nn.ReLU()\n\n    def forward(self, x):\n        return self.relu(self.depthwise(x))\n\nclass RefineBlock(nn.Module):\n    def __init__(self, in_channels):\n        super().__init__()\n        self.conv1 = nn.Conv3d(in_channels, in_channels, kernel_size=3, padding=1)\n    def forward(self, x):\n        return self.conv1(x)\n\nclass Decoder3D(nn.Module):\n    def __init__(self, \n                 decoder_dims=[cf.TARGET_SLICES*6, \n                               cf.TARGET_SLICES*4, \n                               cf.TARGET_SLICES*2, \n                               cf.TARGET_SLICES\n                              ], \n                 out_channels=1, \n                 dropout_rate=cf.dropout_decoder, \n                 apply_cbam_indices=[0, 1],\n                 edge_ch=1,\n                 z_sampling=cf.z_sampling,\n                 check=cf.f_check,\n                 log_interval1=cf.print_freq, \n                 log_interval2=cf.print_freq2,\n                 use_center_heatmap=True\n                ):\n        super().__init__()\n        self.edge_scale_raw = nn.Parameter(torch.tensor(-1.5))\n        self.check          = check\n        self.log_interval1  = log_interval1\n        self.log_interval2  = log_interval2\n        self.iteration      = 0\n\n        def make_cbam_or_none(in_channels, use_cbam):\n            return CBAM(in_channels) if use_cbam else None\n\n        self.attn = nn.ModuleList([\n            make_cbam_or_none(decoder_dims[0], 0 in apply_cbam_indices),\n            make_cbam_or_none(decoder_dims[0], 1 in apply_cbam_indices),\n            make_cbam_or_none(decoder_dims[1], 2 in apply_cbam_indices),\n            make_cbam_or_none(decoder_dims[2]+edge_ch, 3 in apply_cbam_indices),\n        ])\n        if z_sampling:\n            self.up3   = nn.ConvTranspose3d(decoder_dims[0], decoder_dims[0], kernel_size=3, stride=2, padding=1, output_padding=1)\n            self.up2   = nn.ConvTranspose3d(decoder_dims[0], decoder_dims[1], kernel_size=3, stride=2, padding=1, output_padding=1)\n            self.up1   = nn.ConvTranspose3d(decoder_dims[1], decoder_dims[2]+edge_ch, kernel_size=3, stride=2, padding=1, output_padding=1)\n            self.up0   = nn.ConvTranspose3d(decoder_dims[2]+edge_ch, decoder_dims[3], kernel_size=3, stride=2, padding=1, output_padding=1)\n        else:\n            self.up3 = nn.ConvTranspose3d(decoder_dims[0], decoder_dims[0], kernel_size=(1,3,3), stride=(1,2,2), padding=(0,1,1), output_padding=(0,1,1))\n            self.up2 = nn.ConvTranspose3d(decoder_dims[0], decoder_dims[1], kernel_size=(1,3,3), stride=(1,2,2), padding=(0,1,1), output_padding=(0,1,1))\n            self.up1 = nn.ConvTranspose3d(decoder_dims[1], decoder_dims[2]+edge_ch, kernel_size=(1,3,3), stride=(1,2,2), padding=(0,1,1), output_padding=(0,1,1))\n            self.up0 = nn.ConvTranspose3d(decoder_dims[2]+edge_ch, decoder_dims[3], kernel_size=(1,3,3), stride=(1,2,2), padding=(0,1,1), output_padding=(0,1,1))\n\n        \n        self.skip3 = nn.Conv3d(decoder_dims[1], decoder_dims[0], kernel_size=1)\n        self.conv3 = ConvNextBlock3D(decoder_dims[0], dropout_rate, use_droppath=False)\n\n        self.skip2 = nn.Conv3d(decoder_dims[2], decoder_dims[1], kernel_size=3, padding=1)\n        self.conv2 = ConvNextBlock3D(decoder_dims[1], dropout_rate, use_droppath=False)\n\n        self.skip1 = nn.Conv3d(decoder_dims[3], decoder_dims[2]+edge_ch, kernel_size=1)\n        self.conv1 = ConvNextBlock3D(decoder_dims[2]+edge_ch, dropout_rate, use_droppath=False)\n\n        self.conv0 = ConvNextBlock3D(decoder_dims[3], dropout_rate, use_droppath=False)\n\n        self.final_conv = nn.Conv3d(decoder_dims[3], out_channels, kernel_size=1)\n        self.refine     = RefineBlock(in_channels=out_channels)\n\n        if use_center_heatmap:\n            self.center_head = nn.Conv3d(decoder_dims[3], 1, kernel_size=1)\n        else:\n            self.center_head = None\n        \n        # === skip connection enhancement ===\n        self.skip3_enhance = SEBlock(decoder_dims[1])  # for f3\n        #self.skip3_enhance = DepthwiseConv3D(decoder_dims[1]) # for f3\n        self.skip2_enhance = SEBlock(decoder_dims[2])  # for f2\n        #self.skip2_enhance = DepthwiseConv3D(decoder_dims[2])  # for f2\n        self.skip1_enhance = Sobel3D(decoder_dims[3])  # for f1\n        #self.skip1_enhance = SEBlock(decoder_dims[3])\n        #self.skip1_enhance = DepthwiseConv3D(decoder_dims[3])\n\n    def forward(self, features, edge_feature=None):\n        f1, f2, f3, f4 = features\n\n        # enhance skip features\n        f3_enh = self.skip3_enhance(f3)\n        f2_enh = self.skip2_enhance(f2)\n        f1_enh = self.skip1_enhance(f1)\n\n        d3_in = self.attn[0](f4) if self.attn[0] is not None else f4\n        d3 = self.up3(d3_in)\n        d3 = self.conv3(d3 + self.skip3(f3_enh))\n\n        d2_in = self.attn[1](d3) if self.attn[1] is not None else d3\n        d2 = self.up2(d2_in)\n        d2 = self.conv2(d2 + self.skip2(f2_enh))\n\n        d1_in = self.attn[2](d2) if self.attn[2] is not None else d2\n        d1 = self.up1(d1_in)\n        d1 = d1 + self.skip1(f1_enh)\n        # edge_feature upsample & concat or add\n        if edge_feature is not None:\n            if edge_feature.shape[2:] != d1.shape[2:]:\n                edge_feature = F.interpolate(edge_feature, size=d1.shape[2:], mode='trilinear', align_corners=False)\n            \n            # #エッジ特徴をノーマライズ（std=1, mean=0）\n            # edge_feature = (edge_feature - edge_feature.mean()) / (edge_feature.std() + 1e-5)\n\n            # clamp Normalize\n            std          = edge_feature.std().clamp(min=1e-3)\n            edge_feature = (edge_feature - edge_feature.mean()) / std\n\n            #小さくスケーリングしてから加算\n            d1 = d1 + F.softplus(self.edge_scale_raw) * edge_feature\n            #[B, 65, 16, 128, 128])\n            d1 = self.conv1(d1)\n        \n\n        d0_in = self.attn[3](d1) if self.attn[3] is not None else d1\n        d0 = self.up0(d0_in)\n        d0 = self.conv0(d0)\n        \n        if self.check and (self.log_interval2 > 0 and self.iteration % self.log_interval2 == 0):\n            with torch.no_grad():\n                print(f\"D3 Size: {d3.size()}\")\n                describe_tensor(d3, \"D3\")\n                plot_feature_maps_grid(d3, name=\"d3\")\n                print(f\"D2 Size: {d2.size()}\")\n                describe_tensor(d2, \"D2\")\n                plot_feature_maps_grid(d2, name=\"d2\")\n                print(f\"D1 Size: {d1.size()}\")\n                describe_tensor(d1, \"D1\")\n                out_d1 = torch.sigmoid(d1)\n                print()\n                print(\"D1 Sigmoid\")\n                print(f\"min : {out_d1.min().item():.6f}\")\n                print(f\"max : {out_d1.max().item():.6f}\")\n                print(f\"mean: {out_d1.mean().item():6f}\")\n                print()\n                plot_feature_maps_grid(d1, name=\"d1\")\n                \n            del d3, d2, d1\n            _=gc.collect()\n        if cf.debug and self.iteration % self.log_interval1 == 0:\n            print()\n            print(f\"[Decoder3D] edge_scale_raw = {self.edge_scale_raw.item():.4f}\")\n            print()\n        self.iteration +=1\n        \n        x_mask = self.final_conv(d0)\n        #x_mask = self.refine(x)\n        \n        if self.center_head is not None:\n            x_center = self.center_head(d0)\n            return x_mask, x_center\n        else:\n            return x_mask\n\n#------ Edge ------\nclass Sobel3D(nn.Module):\n    def __init__(self, \n                 in_channels, \n                 scale=5,\n                 blur_kernel_size=3, \n                 threshold=0.03\n                ):\n        super().__init__()\n        self.in_channels = in_channels\n        self.scale       = scale\n        self.threshold   = threshold\n\n        # Sobel kernels\n        gx = torch.tensor([\n            [[-1, 0, 1],\n             [-2, 0, 2],\n             [-1, 0, 1]],\n            [[-2, 0, 2],\n             [-4, 0, 4],\n             [-2, 0, 2]],\n            [[-1, 0, 1],\n             [-2, 0, 2],\n             [-1, 0, 1]]\n        ], dtype=torch.float32) / cf.TARGET_SLICES\n\n        gy = torch.tensor([\n            [[-1, -2, -1],\n             [0, 0, 0],\n             [1, 2, 1]],\n            [[-2, -4, -2],\n             [0, 0, 0],\n             [2, 4, 2]],\n            [[-1, -2, -1],\n             [0, 0, 0],\n             [1, 2, 1]]\n        ], dtype=torch.float32) / cf.TARGET_SLICES\n\n        gz = torch.tensor([\n            [[-1, -2, -1],\n             [-2, -4, -2],\n             [-1, -2, -1]],\n            [[0, 0, 0],\n             [0, 0, 0],\n             [0, 0, 0]],\n            [[1, 2, 1],\n             [2, 4, 2],\n             [1, 2, 1]]\n        ], dtype=torch.float32) / cf.TARGET_SLICES\n\n        self.register_buffer(\"kernel_x\", gx[None, None, :, :, :].repeat(in_channels, 1, 1, 1, 1))\n        self.register_buffer(\"kernel_y\", gy[None, None, :, :, :].repeat(in_channels, 1, 1, 1, 1))\n        self.register_buffer(\"kernel_z\", gz[None, None, :, :, :].repeat(in_channels, 1, 1, 1, 1))\n\n        # ノイズ抑制の平均ブラー（固定重み）\n        self.blur = nn.Conv3d(in_channels, \n                              in_channels, \n                              kernel_size=blur_kernel_size,\n                              stride=1, \n                              padding=blur_kernel_size // 2,\n                              groups=in_channels, \n                              bias=False\n                             )\n        self.blur.weight.data.fill_(1.0 / (blur_kernel_size ** 3))\n        self.blur.weight.requires_grad = False\n\n    def forward(self, x):\n        x = self.blur(x)  # ノイズをぼかすことで、細い構造を浮かせる\n        with torch.no_grad():\n            edge_x = F.conv3d(x, self.kernel_x, padding=1, groups=self.in_channels)\n            edge_y = F.conv3d(x, self.kernel_y, padding=1, groups=self.in_channels)\n            edge_z = F.conv3d(x, self.kernel_z, padding=1, groups=self.in_channels)\n    \n            edge = torch.sqrt(torch.clamp(edge_x**2 + edge_y**2 + edge_z**2, min=1e-6))\n        \n            del edge_x, edge_y, edge_z \n            torch.cuda.empty_cache()\n            \n            # チャンネルごとにmin-max正規化\n            min_val = edge.view(edge.size(0), edge.size(1), -1).min(dim=2)[0].view(edge.size(0), edge.size(1), 1, 1, 1)\n            max_val = edge.view(edge.size(0), edge.size(1), -1).max(dim=2)[0].view(edge.size(0), edge.size(1), 1, 1, 1)\n            edge    = (edge - min_val) / (max_val - min_val + 1e-8)\n        return torch.sigmoid((edge - self.threshold) * self.scale).float()\n\n\n# class EdgeEnhancer3D_v2(nn.Module):\n#     def __init__(self, in_channels, scale=0.1):\n#         super().__init__()\n#         self.sobel = Sobel3D(in_channels=in_channels)\n#         self.scale = scale\n\n#     def forward(self, x):\n#         edge = self.sobel(x)\n#         return x + edge * self.scale\n\nclass EdgeEnhancer3D_v2(nn.Module):\n    def __init__(self, \n                 in_channels, \n                 use_learnable_scale=True, \n                 edge_power=0.7, \n                 eps=1e-6, \n                 ratio=0.5 #1stage=0.1\n                ):\n        super().__init__()\n        self.sobel               = Sobel3D(in_channels=in_channels)\n        self.use_learnable_scale = use_learnable_scale\n        self.edge_power = edge_power\n        self.eps        = eps\n        self.ratio      = ratio\n\n        if use_learnable_scale:\n            # チャンネルごとにスケールを学習（初期値は小さく）\n            self.scale = nn.Parameter(torch.ones(1, in_channels, 1, 1, 1) * self.ratio)\n        else:\n            self.scale = 0.1\n\n    def forward(self, x):\n        edge = self.sobel(x)\n\n        # min-max正規化\n        edge_min = edge.amin(dim=[2,3,4], keepdim=True)\n        edge_max = edge.amax(dim=[2,3,4], keepdim=True)\n        edge     = (edge - edge_min) / (edge_max - edge_min + self.eps)\n\n        # 明るいエッジを強調（オプション）\n        if self.edge_power != 1.0:\n            edge = edge.pow(self.edge_power)\n\n        # 重み付き加算\n        if self.use_learnable_scale:\n            out = x + edge * self.scale\n        else:\n            out = x + edge * self.scale\n\n        if cf.parameter_check:\n            # edge の分布チェック\n            print(f\"edge max: {edge.max().item():.6f} edge mean: {edge.mean().item():.6f}\")\n\n        return out\n        \n### Define Model\nclass UNet3D_Attention(nn.Module):\n    def __init__(self, \n                 in_channels=1, \n                 out_channels=1, \n                 dropout_rate=0.1, \n                 eps=1e-6, \n                 edge_channels=1,\n                 edge_input_power=cf.edge_inp_W,\n                 mode=\"2stage\",\n                 gamma=0.2,\n                 output_temp=cf.T,\n                 crop_enhanced=True\n                ):\n        \"\"\"\n        •\tgamma < 1.0 → 強調をさらに強くする（小さな差を大きく拡大）\n    \t•\tgamma > 1.0 → 逆に抑える／鈍化する\n        \"\"\"\n        super().__init__()\n        self.edge_input_p  = edge_input_power\n        self.mode          = mode\n        self.gamma         = gamma\n        self.encoder       = Conv3DEncoder(in_channels, dropout_rate=cf.dropout_encoder)\n        self.decoder       = Decoder3D(out_channels=out_channels, dropout_rate=cf.dropout_decoder, use_center_heatmap=cf.use_center_heatm)\n        self.edge_enhancer = EdgeEnhancer3D_v2(in_channels)\n        self.eps           = eps\n        self.edge_proj = nn.Sequential(\n            nn.Conv3d(1, edge_channels, kernel_size=3, padding=1),\n            nn.GroupNorm(1, edge_channels),\n            nn.Sigmoid(),\n            )\n        self.use_center_heatmap = cf.use_center_heatm\n        self.output_temp   = output_temp\n        self.crop_enhanced = crop_enhanced\n        \n    def forward(self, x):\n        with torch.cuda.amp.autocast(enabled=cf.use_amp):  # Mixed Precision\n            x = x.float()\n    \n            # 1. エッジ抽出\n            x_edge = self.edge_enhancer(x)  # shape: [B, 1, D, H, W]\n            x_edge = x_edge.float()\n            x_edge = x_edge.clamp(0, 1.0)\n            \n            # 2. Convでエッジ特徴変換\n            # SoftClip（例: sigmoid or tanh で緩やかに制限）\n            x_edge_proj = torch.tanh(self.edge_proj(x_edge))  # 値域 [-1, 1]  # shape: [B, C, D, H, W]\n            \n            # 3. x + edge_map\n            if self.mode==\"1stage\":\n                x_input = x + self.edge_input_p * x_edge_proj # 1stage\n            else:\n                # ---- X (input=img) ----\n                x         = x.clamp(0.0, 1.0)  # clamp\n                x_flipped = (1.0 - x).clamp(min=1e-6)\n                x_log     = -torch.log(x_flipped)\n                x_log     = x_log / (x_log.max() + 1e-6)\n                x_log     = x_log ** self.gamma\n                \n                # sigmoidコントラスト強調（0.5中心）\n                \"\"\"\n                sigmoid(5 * (x - 0.5)) は中間値（0.4〜0.6）をかなり鋭く分離するため、強調しすぎる場合は係数5→3などに調整可能\n            \t逆に、ぼやけた構造が残りすぎるなら 5〜8に上げるのもOK\n                \"\"\"\n                #x_log = torch.sigmoid(5 * (x_log - 0.5))\n                \n                # ----- Edge -----\n                # tanh出力を0-1にスケール\n                x_edge_proj_scaled = (x_edge_proj + 1.0) / 2.0  # [-1,1] → [0,1]\n                x_edge_proj_log    = -torch.log(1e-6 + 1.0 - x_edge_proj_scaled)\n                x_edge_proj_log    = x_edge_proj_log / (x_edge_proj_log.max() + 1e-6)\n                x_edge_proj_log    = x_edge_proj_log ** self.gamma\n                \n                # ---- Join ----\n                alpha = self.edge_input_p        # edge強調寄り\n                beta  = 1 - self.edge_input_p   # flipped-logの信号強調\n                \n                x_enhanced = alpha * x_edge_proj_log + beta * x_log\n                x_enhanced = x_enhanced / (x_enhanced.max() + 1e-6)  # scaling\n                \n                if self.crop_enhanced:\n                    # tanh 3~6\n                    x_enhanced = torch.tanh(3 * (x_enhanced - 0.5)) * 0.5 + 0.5\n                    # sigmoid　3~6\n                    #x_enhanced = torch.sigmoid(5 * (x_enhanced - 0.5))\n                    \n                # Final Flapped\n                x_input = 1.0 - x_enhanced\n                # boost\n                #x_input = x_input ** 0.5\n            \n            if x_edge_proj.max().item()>1.0:\n                print(f\"⚠️ X_edge_proj: Max={x_edge_proj.max().item()}\")\n    \n            # 4. Encoder → Decoder\n            features = self.encoder(x_input)\n            if self.use_center_heatmap:\n                logits, center_logits = self.decoder(features, edge_feature=x_edge_proj)\n            else:\n                logits = self.decoder(features, edge_feature=x_edge_proj)\n            \n            # Final Outputs\n            if cf.output_scaling:\n                # soft outputs\n                if self.use_center_heatmap:\n                    output = torch.sigmoid(logits / self.output_temp)\n                    output2= torch.sigmoid(center_logits / self.output_temp)\n                else:\n                    output   = torch.sigmoid(logits / self.output_temp)\n                \n                if cf.debug and cf.output_check:\n                    for t in [1.0, 1.2, 1.5, 2.0, 2.5, 3.0, 3.5]:\n                        out = torch.sigmoid(logits/t)\n                        print(f\"Temp: {t} Max: {out.sum().item():.6f} Min: {out.min().item():.6f} Mean: {out.mean().item():.6f}\")\n                    print()\n            else:\n                output = torch.sigmoid(logits)\n                if self.use_center_heatmap:\n                    output2 = torch.sigmoid(center_logits)\n                \n            if output.max().item()>1.0:\n                print(\"Likely not Sigmoid...\")\n\n        if self.use_center_heatmap:\n            return output, output2, x_input, x_edge_proj\n        else:   \n            return output, x_input, x_edge_proj\n\n\n# Eval Output\nclass CustomOutput(nn.Module):\n    def __init__(self, \n                 threshold=0.1, \n                 mode=\"binary\"  # \"binary\", \"softplus\", \"raw\"\n                ):\n        super().__init__()\n        self.threshold = threshold\n        self.mode      = mode\n\n    def forward(self, x):\n\n        if self.mode == \"binary\":\n            return (x > self.threshold).float()\n\n        elif self.mode == \"softplus\":\n            # スムーズなしきい値処理（ただし曖昧さは残る）\n            return F.softplus(x - self.threshold)\n\n        elif self.mode == \"raw\":\n            # sigmoid 後そのまま返す（可視化など）\n            return x\n\n        else:\n            raise ValueError(f\"Unknown mode: {self.mode}\")\n\nclass CustomOutputWrapper(nn.Module):\n    def __init__(self, model, threshold=0.1, mode=\"binary\"):\n        super().__init__()\n        self.model         = model\n        self.custom_output = CustomOutput(threshold, mode)\n\n    def forward(self, x):\n        out, x_inputs, x_edge = self.model(x)\n        if not self.training:\n            out = self.custom_output(out)\n        return out, x_inputs, x_edge\n\n    def __getattr__(self, name):\n        try:\n            return super().__getattr__(name)\n        except AttributeError:\n            return getattr(self.model, name)\n\n\ndef build_model():\n    model = UNet3D_Attention()\n    if cf.custom_output:\n       model = CustomOutputWrapper(model)\n    return model.to(device)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Loss","metadata":{}},{"cell_type":"code","source":"def show_pred_and_mask_one(prob, target, target_enhanced, tomo_id, motor_count, edge=None, inputs=None, normalized=cf.normalized_show):\n    img       = inputs.squeeze(0).squeeze(0).detach().cpu().numpy()      # [D, H, W]\n    sigm      = prob.squeeze(0).squeeze(0).detach().cpu().numpy()        # [D, H, W]\n    orig_mask = target.squeeze(0).squeeze(0).detach().cpu().numpy()\n    mask_edge = target_enhanced.squeeze(0).squeeze(0).detach().cpu().numpy()\n    edge_map  = edge.squeeze(0).squeeze(0).detach().cpu().numpy() if edge is not None else None\n\n    add_col    = 1  # z-profile用\n    \n    num_cols = (4 if edge_map is not None else 3) + add_col\n    _, H, W    = img.shape\n    motor_count= [mc.item() if torch.is_tensor(mc) else mc for mc in motor_count]\n    z_scores   = sigm.mean(axis=(1, 2))\n    C          = z_scores.argmax()\n    GM         = round(orig_mask[C, ...].max() * 255)\n    sigm_max   = sigm[C, ...].max().item()\n    # スケーリング切り替え\n    if normalized:\n        img       = (img - img.min()) / (img.max() - img.min() + 1e-6)\n        mask_norm = (mask_edge[C, ...] - mask_edge[C, ...].min()) / (mask_edge[C, ...].max() - mask_edge[C, ...].min() + 1e-6)\n        img_show  = img[C, ...] * 255\n        mask_show = mask_norm * 255\n        \n        # if plot out of memory\n        img_show  = resize(img_show,  (128, 128), preserve_range=True)\n        mask_show = resize(mask_show, (128, 128), preserve_range=True)\n    else:\n        img_show  = img[C, ...]\n        mask_show = mask_edge[C, ...]\n\n        img_show  = resize(img_show,  (128, 128), preserve_range=True)\n        mask_show = resize(mask_show, (128, 128), preserve_range=True)\n\n    fig, ax = plt.subplots(1, num_cols, figsize=(num_cols * 3, 3))\n    img_max   = img[C, ...].max().item()\n    \n    mask_max  = mask_edge[C, ...].max().item()\n    mask_mean = mask_edge[C, ...].mean().item()\n    \n    sigm_slice= sigm[C, ...]\n    sigm_max  = sigm_slice.max().item()\n    sigm_mean = sigm_slice.mean().item()\n    sigm_vmax = max(sigm_max, 0.05)  # 最大値が小さすぎると真っ黒になるため\n    # --- 最大反応位置を取得 ---\n    topn=3\n    flat_slice  = sigm_slice.ravel()\n    top_indices = np.argpartition(flat_slice, -topn)[-topn:]\n    top_indices = top_indices[np.argsort(-flat_slice[top_indices])]  # 値の大きい順に並び替え\n    # 各点の (y, x) 座標に変換\n    yx_coords = np.array(np.unravel_index(top_indices, sigm_slice.shape)).T\n    y_coords, x_coords = yx_coords[:, 0], yx_coords[:, 1]\n\n    ax[0].imshow(img_show, cmap=\"gray\", vmin=0, vmax=255 if normalized else None)\n    ax[1].imshow(sigm[C, ...], cmap=\"gray\", vmin=0, vmax=sigm_vmax)\n    # scatter\n    colors = ['red', 'green', 'blue']\n    for j in range(len(x_coords)):\n        ax[1].scatter([x_coords[j]], [y_coords[j]], \n                      c=colors[j % len(colors)], s=30, marker='o', \n                      label=f'Top-{j+1}', alpha=0.4)\n    #ax[1].legend()\n    ax[2].imshow(mask_show, cmap=\"gray\" if normalized else \"hot\", vmin=0, vmax=255 if normalized else None)\n    \n    col_idx = 3\n    if edge_map is not None:\n        edge_show = edge_map[C, ...]\n        edge_show = resize(edge_show, (128, 128), preserve_range=True)\n        ax[col_idx].imshow(edge_show, cmap=\"inferno\")\n        emax = edge_map[C, ...].max().item()\n        emin = edge_map[C, ...].min().item()\n        emean= edge_map[C, ...].mean().item()\n        ax[col_idx].set_title(f\"Edge Map Max: {emax:.3f} Min: {emin:.3f} Mean: {emean:.3f}\", size=8)\n        col_idx += 1\n        \n    ax[col_idx].plot(z_scores)\n    ax[col_idx].set_ylim(0, z_scores.max()*1.1)\n    ax[col_idx].set_xlabel(\"Z\")\n    ax[col_idx].set_ylabel(\"Mean\")\n    ax[col_idx].set_title(\"Z Activation Profile\", size=8)\n\n    ax[0].set_title(f\"Input: Max={img_max:.3f} id={tomo_id}\", size=8)\n    ax[1].set_title(f\"Max={sigm_max:.3f} Mean={sigm_mean:.3f} z={C}  H={H} W={W}\", size=8)\n    ax[2].set_title(f\"Mask:MM={mask_max:.3f} Mean={mask_mean:.3f} Gmax={GM} F={motor_count[0]}\", size=8)\n\n    for j in range(num_cols):\n        ax[j].set_xticks([])\n        ax[j].set_yticks([])\n\n    plt.tight_layout()\n    plt.show()\n    plt.close(fig)  # ★これが最重要\n    \n    del fig, ax\n    _=gc.collect()  ","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_pred_and_mask(probs, targets, target_enhanced, tomo_ids, motor_count, \n                       edges=None, \n                       inputs=None,\n                       target_com=None,\n                       normalized=cf.normalized_show,\n                       use_argmax=False,\n                       show_text=False\n                      ):\n    imgs       = inputs.squeeze(1).detach().cpu().numpy()\n    sigm       = probs.squeeze(1).detach().cpu().numpy()\n    orig_mask  = targets.squeeze(1).detach().cpu().numpy()\n    mask_edge  = target_enhanced.squeeze(1).detach().cpu().numpy()\n    edge_map   = edges.squeeze(1).detach().cpu().numpy() if edges is not None else None\n    target_com = target_com.detach().cpu().numpy() if target_com is not None else None\n\n    add_col    = 1  # z-profile用の追加カラム\n    num_cols   = (4 if edge_map is not None else 3) + add_col\n    batch_size = len(tomo_ids)\n    _, _, H, W = imgs.shape\n    motor_count= [mc.item() if torch.is_tensor(mc) else mc for mc in motor_count]\n\n    fig, ax = plt.subplots(nrows=batch_size, ncols=num_cols, figsize=(num_cols * 3, batch_size * 2))\n    if batch_size == 1:\n        ax = np.expand_dims(ax, 0)\n\n    for i in range(batch_size):\n        z_scores = sigm[i].mean(axis=(1, 2))  # shape: (D,)\n\n        # --- best_z を z軸全体から求める ---\n        threshold   = 0.5\n        best_score  = -1\n        best_coords = None\n        best_z      = 0\n\n        for z in range(sigm.shape[1]):\n            sigm_z      = sigm[i, z]\n            binary_mask = (sigm_z > threshold).astype(np.uint8)\n            labeled     = label(binary_mask, connectivity=2)\n            num         = labeled.max()\n            if num == 0:\n                continue\n            for j in range(1, num + 1):\n                blob  = (labeled == j)\n                score = sigm_z[blob].mean()\n                if score > best_score:\n                    best_score  = score\n                    coords      = np.argwhere(blob)\n                    best_coords = coords\n                    best_z      = z\n\n        C  = best_z  # すべてに best_z を適用\n        GM = round(orig_mask[i, C, ...].max() * 255)\n\n        # 正規化表示\n        if normalized:\n            imgs      = (imgs - imgs.min()) / (imgs.max() - imgs.min() + 1e-6)\n            img_show  = imgs[i, C, ...] * 255\n            mask_norm = (mask_edge[i, C, ...] - mask_edge[i, C, ...].min()) / (mask_edge[i, C, ...].max() - mask_edge[i, C, ...].min() + 1e-6)\n            mask_show = mask_norm * 255\n        else:\n            img_show  = imgs[i, C, ...]\n            mask_show = mask_edge[i, C, ...]\n\n        img_max   = imgs[i, C, ...].max().item()\n        mask_max  = mask_edge[i, C, ...].max().item()\n        mask_mean = mask_edge[i, C, ...].mean().item()\n        sigm_slice= sigm[i, C, ...]\n        sigm_max  = sigm_slice.max().item()\n        sigm_mean = sigm_slice.mean().item()\n        sigm_vmax = max(sigm_max, 0.05)\n\n        # トップスコア座標取得（常に描画）\n        topn=3\n        flat_slice  = sigm_slice.ravel()\n        top_indices = np.argpartition(flat_slice, -topn)[-topn:]\n        top_indices = top_indices[np.argsort(-flat_slice[top_indices])]\n        yx_coords = np.array(np.unravel_index(top_indices, sigm_slice.shape)).T\n        y_coords, x_coords = yx_coords[:, 0], yx_coords[:, 1]\n\n        # --- 描画部 ---\n        ax[i, 0].imshow(img_show, cmap=\"gray\", vmin=0, vmax=255 if normalized else None)\n        ax[i, 1].imshow(sigm[i, C, ...], cmap=\"gray\", vmin=0, vmax=sigm_vmax)\n        ax[i, 2].imshow(mask_show, cmap=\"magma\" if normalized else \"hot\", vmin=0, vmax=255 if normalized else None)\n        if target_com is not None:\n            targ_z, targ_y, targ_x = target_com[i]\n            if abs(round(targ_z) - C) <= 2:  # +- 2 Z slice\n                #print(\"PLOT Target Com\")\n                ax[i, 2].scatter([targ_x], [targ_y],\n                                 c='cyan', s=40, marker='x',\n                                 label='Target COM')\n\n        if use_argmax:\n            colors = ['red', 'green', 'blue']\n            for j in range(len(x_coords)):\n                ax[i, 1].scatter([x_coords[j]], [y_coords[j]], \n                                 c=colors[j % len(colors)], s=30, marker='o', \n                                 label=f'Top-{j+1}', alpha=0.4)\n        else:\n            if best_coords is not None:\n                center_y, center_x = best_coords.mean(axis=0)\n                ax[i, 1].scatter([center_x], [center_y],\n                                 c='red', s=40, marker='o',\n                                 alpha=0.8 if show_text else 0.4,\n                                 label='Blob Center')\n                if show_text:\n                    score_text = f\"{best_score:.2f}\"\n                    ax[i, 1].text(center_x + 2, center_y, score_text,\n                                  color='white', fontsize=8,\n                                  bbox=dict(facecolor='black', alpha=0.5, boxstyle='round,pad=0.2'))\n\n        col_idx = 3\n        if edge_map is not None:\n            ax[i, col_idx].imshow(edge_map[i, C, ...], cmap=\"inferno\", vmin=0, vmax=1)\n            emax = edge_map[i, C, ...].max().item()\n            emin = edge_map[i, C, ...].min().item()\n            emean= edge_map[i, C, ...].mean().item()\n            ax[i, col_idx].set_title(f\"Edge Max: {emax:.3f} Min: {emin:.3f} Mean: {emean:.3f}\", size=8)\n            col_idx += 1\n\n        ax[i, col_idx].plot(z_scores)\n        ax[i, col_idx].set_ylim(0, z_scores.max() * 1.1)\n        ax[i, col_idx].set_title(\"Z Activation Profile\", size=8)\n        ax[i, col_idx].set_xlabel(\"Z\")\n        ax[i, col_idx].set_ylabel(\"Mean\")\n\n        ax[i, 0].set_title(f\"Input: Max={img_max:.3f} id={tomo_ids[i]}\", size=8)\n        ax[i, 1].set_title(f\"Sigmoid Max={sigm_max:.3f} Mean={sigm_mean:.3f} z={C}\", size=8)\n        ax[i, 2].set_title(f\"Mask: MM={mask_max:.3f} Mean={mask_mean:.3f} Gmax={GM} F={motor_count[i]}\", size=8)\n\n        for j in range(num_cols):\n            ax[i, j].set_xticks([])\n            ax[i, j].set_yticks([])\n\n    plt.tight_layout()\n    plt.show(block=False)\n    plt.pause(0.001)  # 少しだけ時間をあげて描画させる\n    plt.close(fig)\n    del fig, ax\n    _ = gc.collect()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_pred_and_mask_com_based(probs, targets, target_enhanced, tomo_ids, motor_count, \n                                 edges=None, \n                                 inputs=None,\n                                 target_com=None,\n                                 pred_com=None,\n                                 normalized=cf.normalized_show,\n                                 show_text=False\n                                ):\n    imgs       = inputs.squeeze(1).detach().cpu().numpy()\n    sigm       = probs.squeeze(1).detach().cpu().numpy()\n    orig_mask  = targets.squeeze(1).detach().cpu().numpy()\n    mask_edge  = target_enhanced.squeeze(1).detach().cpu().numpy()\n    edge_map   = edges.squeeze(1).detach().cpu().numpy() if edges is not None else None\n    target_com = target_com.detach().cpu().numpy() if target_com is not None else None\n    pred_com   = pred_com.detach().cpu().numpy() if pred_com is not None else None\n\n    add_col    = 1\n    num_cols   = (4 if edge_map is not None else 3) + add_col\n    batch_size = len(tomo_ids)\n    _, _, H, W = imgs.shape\n    motor_count= [mc.item() if torch.is_tensor(mc) else mc for mc in motor_count]\n\n    fig, ax = plt.subplots(nrows=batch_size, ncols=num_cols, figsize=(num_cols * 3, batch_size * 2))\n    if batch_size == 1:\n        ax = np.expand_dims(ax, 0)\n\n    for i in range(batch_size):\n        z_scores = sigm[i].mean(axis=(1, 2))  # shape: (D,)\n\n        # --- pred_com から描画sliceを取得 ---\n        if pred_com is not None:\n            C = int(round(pred_com[i, 0]))  # zスライス\n            center_y = pred_com[i, 1]\n            center_x = pred_com[i, 2]\n        else:\n            C = sigm.shape[1] // 2  # fallback\n            center_y, center_x = H//2, W//2\n\n        GM = round(orig_mask[i, C, ...].max() * 255)\n\n        # 正規化表示\n        if normalized:\n            imgs      = (imgs - imgs.min()) / (imgs.max() - imgs.min() + 1e-6)\n            img_show  = imgs[i, C, ...] * 255\n            mask_norm = (mask_edge[i, C, ...] - mask_edge[i, C, ...].min()) / (mask_edge[i, C, ...].max() - mask_edge[i, C, ...].min() + 1e-6)\n            mask_show = mask_norm * 255\n        else:\n            img_show  = imgs[i, C, ...]\n            mask_show = mask_edge[i, C, ...]\n\n        img_max   = imgs[i, C, ...].max().item()\n        mask_max  = mask_edge[i, C, ...].max().item()\n        mask_mean = mask_edge[i, C, ...].mean().item()\n        sigm_slice= sigm[i, C, ...]\n        sigm_max  = sigm_slice.max().item()\n        sigm_mean = sigm_slice.mean().item()\n        sigm_vmax = max(sigm_max, 0.05)\n\n        # --- 描画部 ---\n        ax[i, 0].imshow(img_show, cmap=\"gray\", vmin=0, vmax=255 if normalized else None)\n        ax[i, 1].imshow(sigm[i, C, ...], cmap=\"gray\", vmin=0, vmax=sigm_vmax)\n        ax[i, 2].imshow(mask_show, cmap=\"magma\" if normalized else \"hot\", vmin=0, vmax=255 if normalized else None)\n\n        # target_comがあれば描画\n        if target_com is not None:\n            targ_z, targ_y, targ_x = target_com[i]\n            if abs(round(targ_z) - C) <= 2:\n                ax[i, 2].scatter([targ_x], [targ_y],\n                                 c='cyan', s=40, marker='x',\n                                 label='Target COM')\n\n        # pred_com描画\n        if pred_com is not None and -1 not in pred_com[i]:\n            ax[i, 1].scatter([center_x], [center_y],\n                             c='red', s=40, marker='o',\n                             alpha=0.8 if show_text else 0.4,\n                             label='Pred COM')\n            if show_text:\n                ax[i, 1].text(center_x + 2, center_y, f\"({pred_com[i][0]:.1f},{pred_com[i][1]:.1f},{pred_com[i][2]:.1f})\",\n                              color='white', fontsize=8,\n                              bbox=dict(facecolor='black', alpha=0.5, boxstyle='round,pad=0.2'))\n\n        col_idx = 3\n        if edge_map is not None:\n            ax[i, col_idx].imshow(edge_map[i, C, ...], cmap=\"inferno\", vmin=0, vmax=1)\n            emax = edge_map[i, C, ...].max().item()\n            emin = edge_map[i, C, ...].min().item()\n            emean= edge_map[i, C, ...].mean().item()\n            ax[i, col_idx].set_title(f\"Edge Max: {emax:.3f} Min: {emin:.3f} Mean: {emean:.3f}\", size=8)\n            col_idx += 1\n\n        ax[i, col_idx].plot(z_scores)\n        ax[i, col_idx].set_ylim(0, z_scores.max() * 1.1)\n        ax[i, col_idx].set_title(\"Z Activation Profile\", size=8)\n        ax[i, col_idx].set_xlabel(\"Z\")\n        ax[i, col_idx].set_ylabel(\"Mean\")\n\n        ax[i, 0].set_title(f\"Input: Max={img_max:.3f} id={tomo_ids[i]}\", size=8)\n        ax[i, 1].set_title(f\"Sigmoid Max={sigm_max:.3f} Mean={sigm_mean:.3f} z={C}\", size=8)\n        ax[i, 2].set_title(f\"Mask: MM={mask_max:.3f} Mean={mask_mean:.3f} Gmax={GM} F={motor_count[i]}\", size=8)\n\n        for j in range(num_cols):\n            ax[i, j].set_xticks([])\n            ax[i, j].set_yticks([])\n\n    plt.tight_layout()\n    plt.show(block=False)\n    plt.pause(0.001)\n    plt.close(fig)\n    del fig, ax\n    _ = gc.collect()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_weight_and_prob_slices(weights, probs):\n    \"\"\"\n    weights: (B, 1, D, H, W) torch.Tensor → already converted to numpy\n    probs  : (B, 1, D, H, W) torch.Tensor → already converted to numpy\n    \"\"\"\n    weights = weights.squeeze(1)  # (B, D, H, W)\n    probs   = probs.squeeze(1)    # (B, D, H, W)\n    batch   = weights.shape[0]\n\n    for i in range(batch):\n        weight  = weights[i]  # (D, H, W)\n        prob    = probs[i]    # (D, H, W)\n        D, H, W = weight.shape\n\n        # z方向で最大の位置を取得\n        z_index  = int(np.argmax(weight.max(axis=(1, 2))))\n        min_val  = weight.min()\n        max_val  = weight.max()\n        mean_val = weight.mean()\n\n        if z_index == 0:\n            z_index = D // 2\n\n        # z_index の前後1枚（計3枚）を取得（範囲制限）\n        z_range = [z for z in range(z_index - 1, z_index + 2) if 0 <= z < D]\n\n        # サブプロット数：スライス数 + 2 (hist, percentile)\n        total_plots = len(z_range) + 2\n        fig, axes = plt.subplots(1, total_plots, figsize=(5 * total_plots, 5))\n\n        fig.suptitle(\n            f\"[Batch {i+1}] z_max={z_index} | min={min_val:.3f}, max={max_val:.3f}, mean={mean_val:.3f}\",\n            fontsize=18, fontweight=\"bold\"\n        )\n\n        # --- スライス + 重ね合わせ ---\n        for idx, z in enumerate(z_range):\n            ax = axes[idx]\n\n            if weight[z].shape[0] == cf.IMSIZE:\n                weight_show = resize(weight[z], (128, 128), preserve_range=True)\n                prob_show   = resize(prob[z], (128, 128), preserve_range=True)\n            else:\n                weight_show = weight[z]\n                prob_show   = prob[z]\n            \n            ax.imshow(weight_show, cmap=\"hot\", interpolation=\"nearest\")\n            ax.imshow(prob_show,   cmap=\"Blues\", interpolation=\"nearest\", alpha=0.6)\n            ax.set_title(f\"Z={z}\", fontsize=12)\n            ax.axis(\"off\")\n\n        # --- Histogram ---\n        flat_weights = weight.flatten()\n        ax_hist = axes[-2]\n        ax_hist.hist(flat_weights, bins=100, color=\"orange\", edgecolor=\"black\", alpha=0.7)\n        ax_hist.set_title(\"Histogram\", fontsize=12)\n        ax_hist.set_xlabel(\"Weight\")\n        ax_hist.set_ylabel(\"Count\")\n        ax_hist.grid(True)\n\n        # --- Percentile Plot ---\n        percentiles = np.percentile(flat_weights, np.arange(0, 101, 1))\n        ax_percentile = axes[-1]\n        ax_percentile.plot(np.arange(0, 101, 1), percentiles, color=\"purple\", linewidth=2)\n        ax_percentile.set_title(\"Percentile\", fontsize=12)\n        ax_percentile.set_xlabel(\"Percentile (%)\")\n        ax_percentile.set_ylabel(\"Weight Value\")\n        ax_percentile.grid(True)\n\n        plt.tight_layout()\n        plt.show()\n    plt.close(fig)  # ★これが最重要\n    \n    del fig, axes\n    _=gc.collect()  ","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_peak_z_slice(weights: torch.Tensor, tomo_ids: list[str]):\n    assert weights.ndim == 5 and weights.shape[1] == 1\n    B, _, Z, Y, X = weights.shape\n    weights = weights.squeeze(1).detach().cpu().numpy()\n\n    fig, axes = plt.subplots(1, B, figsize=(4 * B, 4))\n\n    for i in range(B):\n        vol = weights[i]  # [Z, Y, X]\n        # Z方向でmax値を持つインデックスを探す（=gauss ball 中心）\n        peak_z    = np.argmax(vol.max(axis=(1,2)))\n        slice_img = vol[peak_z]\n\n        ax = axes[i] if B > 1 else axes\n        im = ax.imshow(slice_img, cmap='hot')\n        ax.set_title(f\"{tomo_ids[i]} (Z={peak_z})\")\n        plt.colorbar(im, ax=ax)\n\n    plt.tight_layout()\n    plt.show()\n    plt.close(fig)\n    del fig, axes\n    _=gc.collect()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CenterLoss(nn.Module):\n    def __init__(self, center_threshold=0.4):\n        super().__init__()\n        self.center_threshold = center_threshold\n\n    def get_com_with_mask(self, tensor):  # tensor: (B, 1, D, H, W)\n        B, C, D, H, W = tensor.shape\n        if C != 1:\n            raise ValueError(f\"Expected C=1, got {C}\")\n            \n        tensor = tensor[:, 0]  # (B, D, H, W)\n    \n        coords = torch.stack(torch.meshgrid(\n            torch.arange(D, device=tensor.device),\n            torch.arange(H, device=tensor.device),\n            torch.arange(W, device=tensor.device),\n            indexing='ij'\n        ), dim=0).float().unsqueeze(0)  # (1, 3, D, H, W)\n    \n        weighted_coords = tensor.unsqueeze(1) * coords  # (B, 3, D, H, W)\n    \n        sum_tensor = tensor.sum(dim=(1, 2, 3), keepdim=False).clamp(min=1e-6).unsqueeze(1)  # (B, 1)\n    \n        com = weighted_coords.sum(dim=(2, 3, 4)) / sum_tensor  # (B, 3) \n    \n        zero_mask      = (tensor.sum(dim=(1, 2, 3)) == 0)\n        com[zero_mask] = 1.0\n        return com  # (B, 3)\n\n\n    def get_pred_center_with_threshold(self, pred_heatmap):\n        \"\"\"\n        pred_heatmap: Tensor (B, 1, D, H, W) — 出力ヒートマップ\n        threshold   : float — confidenceが低すぎる場合に [-1, -1, -1] を返す\n        \"\"\"\n        B, _, D, H, W = pred_heatmap.shape\n    \n        # 重心ベース（Center of Mass）\n        coords = torch.stack(torch.meshgrid(\n            torch.arange(D, device=pred_heatmap.device),\n            torch.arange(H, device=pred_heatmap.device),\n            torch.arange(W, device=pred_heatmap.device),\n            indexing='ij'\n        ), dim=0).float().unsqueeze(0)  # (1, 3, D, H, W)\n    \n        weighted_coords = pred_heatmap * coords  # (B, 3, D, H, W)\n        sum_tensor = pred_heatmap.sum(dim=(2, 3, 4), keepdim=True) + 1e-6  # (B, 1, 1, 1, 1)\n        com = weighted_coords.sum(dim=(2, 3, 4)) / sum_tensor.view(B, 1)  # (B, 3)\n    \n        # ピークベース（最大値の位置）\n        pred_flat = pred_heatmap.view(B, -1)  # (B, D*H*W)\n        max_vals, max_idxs = pred_flat.max(dim=1)  # (B,), (B,)\n    \n        # ピーク座標に変換\n        z = (max_idxs // (H * W)).long()\n        y = ((max_idxs % (H * W)) // W).long()\n        x = (max_idxs % W).long()\n        peak_coords = torch.stack([z, y, x], dim=1).float()  # (B, 3)\n    \n        # スコアが閾値以下なら [-1, -1, -1]\n        invalid_mask = (max_vals < self.center_threshold).view(B, 1)\n        invalid_value = torch.full_like(com, -1.0)\n    \n        # 重心ベースとピークベースの混合（重心が中央に寄ってしまう対策）\n        # ⇒ 信頼度高い → COM、低め → ピーク優先（柔軟に切り替え可能）\n        final_com = torch.where(invalid_mask.expand_as(com), invalid_value, com)\n    \n        # もしCOMとピークの差が極端に大きい場合 → ピークに置き換える\n        dist = torch.norm(com - peak_coords, dim=1, keepdim=True)\n        replace_mask = (dist > 20.0) & (~invalid_mask)  # ピークとCOMがズレすぎてたら置換\n        final_com = torch.where(replace_mask.expand_as(final_com), peak_coords, final_com)\n    \n        return final_com\n\n    def forward(self, pred, target, motor_count=None):\n        \"\"\"\n        target     : (B, C=1, D, H, W)\n        pred_com   : (B, 3)\n        target_com : (B, 3)\n        \n        \"\"\"\n        pred_com   = self.get_pred_center_with_threshold(pred) #(B, 3)\n        target_com = self.get_com_with_mask(target)\n        \n        dist     = ((pred_com - target_com) ** 2).sum(dim=1).sqrt()\n        max_dist = (pred.shape[2] ** 2 + pred.shape[3] ** 2 + pred.shape[4] ** 2) ** 0.5\n        dist_normalized = dist / max_dist\n\n        if motor_count is not None:\n            motor_mask = (motor_count == 1)\n            dist_normalized = dist_normalized[motor_mask]\n\n        if dist_normalized.numel() == 0:\n            loss = torch.tensor(0.0, device=pred.device, requires_grad=True)\n        else:\n            loss = torch.log1p(dist_normalized * 50.0).mean()\n\n        return loss, pred_com, target_com, dist_normalized\n\n\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=0.5, gamma=1.0, reduction='mean', use_targets_norm=False):\n        \"\"\"\n        alpha    : 正例/負例の重み係数（eg：0.25）\n        gamma    : 難易度調整パラメータ（ef：2.0）\n        reduction: 'mean' or 'none'\n        use_targets_norm: Trueならtargets_norm（ガウス＋エッジ）、Falseなら元のtargets（ガウスのみ\n        \"\"\"\n        super().__init__()\n        self.alpha     = alpha\n        self.gamma     = gamma\n        self.reduction = reduction\n        self.use_targets_norm = use_targets_norm\n\n    def forward(self, probs, targets, targets_norm=None, voxel_weights=None, sample_weights=None):\n        \"\"\"\n        probs         : sigmoid後の出力 (B,1,D,H,W)\n        targets       : ガウス中心のみ (B,1,D,H,W)\n        targets_norm  : ガウス＋エッジマスク (B,1,D,H,W)\n        voxel_weights : voxelごとの重み (B,1,D,H,W)\n        sample_weights: サンプル単位の重み (B,)\n        \n        \"\"\"\n        eps = 1e-7\n        targets_use = targets_norm if self.use_targets_norm else targets\n\n        probs             = probs.clamp(min=eps, max=1 - eps)\n        p_t               = targets_use * probs + (1 - targets_use) * (1 - probs)\n        ce_loss           = - (targets_use * torch.log(probs) + (1 - targets_use) * torch.log(1 - probs))\n        modulating_factor = (1 - p_t) ** self.gamma\n        alpha_weight = targets_use * self.alpha + (1 - targets_use) * (1 - self.alpha)\n        loss         = alpha_weight * modulating_factor * ce_loss  # main focal loss\n\n        if voxel_weights is not None:\n            loss = loss * voxel_weights\n\n        loss = loss.view(probs.size(0), -1).mean(dim=1)  # 各サンプルごと (B,)\n\n        if sample_weights is not None:\n            loss = loss * sample_weights  # sampleごとのweightを掛ける\n\n        if self.reduction == 'mean':\n            return loss.mean()\n        else:\n            return loss\n\nclass DynamicLossWeight:\n    def __init__(self, loss_names, alpha=0.9, eps=1e-6):\n        self.alpha      = alpha\n        self.eps        = eps\n        self.loss_stats = {name: 1.0 for name in loss_names}\n\n    def update(self, current_losses: dict):\n        for name, value in current_losses.items():\n            self.loss_stats[name] = (\n                self.alpha * self.loss_stats[name] + (1 - self.alpha) * value\n            )\n\n    def get_weights(self):\n        inv_mags = {\n            name: 1.0 / (value + self.eps)\n            for name, value in self.loss_stats.items()\n        }\n        total = sum(inv_mags.values())\n        return {name: val / total for name, val in inv_mags.items()}\n        \n        \nclass GaussianBallHybridLoss(nn.Module):\n    def __init__(self, \n                 bce_weight    =0.5, \n                 dice_weight   =0.5,\n                 tversky_weight=0.5,\n                 focal_weight  =0.5,\n                 center_weight =0.5,\n                 center_heat_w =0.5,\n                 z_loss_weight =0.5,\n                 alpha         =0.3, #(0.5~0.7) or 0.3\n                 dice_alpha    =0.7, #1.7 or dice_alpha（α）は 曲線の鋭さを決めるパラメータ。典型的な値は 3〜5、または小さいなら 1.0〜1.5\n                 tversky_alpha =0.2,\n                 tversky_beta  =0.7,#0.7\n                 gamma         =1.33,\n                 base_thresh_ratio=0.005, \n                 min_thresh       =1e-4,\n                 smooth           =1e-6,\n                 clamp_eps        =1e-6,\n                 YX_weight_power  =0.6,\n                 Z_weight_power   =0.8, # 1.2\n                 ZYX_ratio        =1.8, # ~2.0\n                 final_w_ratio    =1.4,\n                 edge_weight      =1.2, \n                 edge_power       =1.5, #(0.8~1.5)\n                 log_interval     =cf.print_freq if cf.debug else 0,\n                 show_weight      =cf.weight_show,\n                 show_peak_weight =cf.show_peak_weight,\n                 total_epochs     =cf.epochs,\n                 mode             =\"1stage\",\n                 final_weight_norm=\"mixed\", #\"mixed\" or \"mean\" or \"amax\" or None\n                 use_W_epoch_min  =cf.use_W_epoch_min,\n                 auto_weight      =True,\n                 center_loss_module=None\n                ):\n        super().__init__()\n        self.auto_weight          = auto_weight\n        self.bce_weight           = bce_weight\n        self.dice_weight          = dice_weight\n        self.tversky_weight       = tversky_weight\n        self.focal_weight         = focal_weight\n        self.center_weight        = center_weight\n        self.center_heatmap_weight= center_heat_w\n        self.z_loss_weight        = z_loss_weight\n        self.alpha             = alpha\n        self.dice_alpha        = dice_alpha\n        self.tversky_alpha     = tversky_alpha\n        self.tversky_beta      = tversky_beta\n        self.gamma             = gamma\n        self.smooth            = smooth\n        self.clamp_eps         = clamp_eps\n        self.base_thresh_ratio = base_thresh_ratio\n        self.min_thresh        = min_thresh\n        self.edge_weight       = edge_weight\n        self.edge_power        = edge_power\n        self.YX_weight_power   = YX_weight_power\n        self.Z_weight_power    = Z_weight_power\n        self.zyx_ratio         = ZYX_ratio\n        self.final_w_ratio     = final_w_ratio\n        self.log_interval      = log_interval\n        self.iteration         = 0\n        self.ema_nonzero_ratio = 0.5\n        self.ema_alpha         = 0.8\n        self.apply_bce_scaling = False\n        self.show_weight       = show_weight\n        self.mode              = mode\n        self.final_weight_norm = final_weight_norm\n        self.total_epochs      = total_epochs\n        self.current_epoch     = 0\n        self.use_W_epoch_min   = use_W_epoch_min\n        self.z_gaussian_weights= False #1stage=False\n        self.use_center_weight = True\n        self.show_peak_weight  = show_peak_weight\n        self.use_edge_enhancement = True\n        self.edge_soft_mode       = True\n        self.edge_power2          = 1.5\n        self.fallback_value       = 0.05 # or 0\n        self.use_center_heatmap   = cf.use_center_heatm\n        self.apply_bce_with_motor_count = True\n        self.neg_ratio        = 1.0\n        self.pos_ratio        = 2.0\n        # --- Focal Loss -----\n        self.FocalLoss        = FocalLoss()\n        # --- Center Loss ----\n        self.criterion_center = center_loss_module\n        # --- EMA log ----\n        self.ema_log = getattr(self, \"ema_log\", {\"nonzero\": [], \"ema\": [], \"scale\": []})\n        # --- Pred_com Base plot ---\n        self.show_pred_com_base= True\n        # --- DynamicLossWeight ---\n        if self.auto_weight:\n            self.loss_tracker = DynamicLossWeight(['bce', \n                                                   'dice', \n                                                   'tversky', \n                                                   'focal',\n                                                   'center', \n                                                   'center_heatmap', \n                                                   'z_loss'\n                                                  ])\n        else:\n            self.bce_weight           = bce_weight\n            self.dice_weight          = dice_weight\n            self.tversky_weight       = tversky_weight\n            self.focal_weight         = focal_weight\n            self.center_weight        = center_weight\n            self.center_heatmap_weight= center_heat_w\n            self.z_loss_weight        = z_loss_weight\n\n    def set_epoch(self, epoch):\n        self.current_epoch = epoch\n\n    def get_clamp_min(self):\n        # eg: 0.05 → 0.01 に直線的に下げる\n        min_start = 0.3#0.05\n        min_end   = 0.1#0.01\n        ratio     = min(self.current_epoch / (self.total_epochs - 1), 1.0)\n        return min_start * (1 - ratio) + min_end * ratio\n\n    # --- Z Weights ---\n    @staticmethod\n    def get_z_importance_weights(targets, \n                                 motor_counts, \n                                 sigma=2.5, \n                                 threshold=10.0, \n                                 fallback_value=0.05\n                                ):\n        \"\"\"\n        targets       : Tensor of shape (B, 1, D, H, W)\n        motor_counts  : Tensor of shape (B,)\n        fallback_value: when motor_count==0, fallback to low flat weight if >0, else 0.\n        \n        \"\"\"\n        B, _, D, H, W = targets.shape\n        \n        z_coords = torch.arange(D, device=targets.device).float().view(1, D)\n        z_grid   = z_coords.view(1, D)  # (1, D)\n    \n        # Z-profile\n        target_z_profile = targets.sum(dim=(3, 4)).squeeze(1)  # (B, D)\n        z_mask = (target_z_profile > threshold).float()        # (B, D)\n    \n        # Motor count mask\n        motor_mask = (motor_counts > 0).float().view(B, 1)     # (B, 1)\n        z_mask = z_mask * motor_mask                           # motor=0なら全て0\n    \n        # 重心計算（z_maskでnoise除去）\n        z_profile = target_z_profile * z_mask + 1e-6\n        z_center  = (z_profile * z_coords).sum(dim=1) / (z_profile.sum(dim=1) + 1e-6)\n        z_center  = z_center.view(B, 1)  # (B, 1)\n    \n        # Gaussian\n        z_weights = torch.exp(-((z_grid - z_center) ** 2) / (2 * sigma ** 2))  # (B, D)\n        z_weights = z_weights * z_mask\n    \n        # Normalization\n        z_weights = z_weights / (z_weights.sum(dim=1, keepdim=True) + 1e-6)\n    \n        # motor_count == 0 の場合 fallback weightを使用（e.g., Soft weight）\n        fallback_mask    = (motor_counts == 0).float().view(B, 1)  # (B, 1)\n        fallback_weights = torch.ones((B, D), device=targets.device) * fallback_value\n        z_weights = z_weights * (1 - fallback_mask) + fallback_weights * fallback_mask\n    \n        return z_weights.view(B, 1, D, 1, 1)\n        \n    @staticmethod\n    def get_z_gaussian_weights(depth, device=\"cpu\", sigma_ratio=0.25):\n        \"\"\"\n        Z方向に中心重視の重みをつけるガウス関数を生成\n        Args:\n            depth (int)                 : z方向のスライス数\n            device (str or torch.device): 配置先\n            sigma_ratio (float)         : sigmaを depth に対する割合で調整（例: 0.25）\n        Returns:\n            Tensor of shape (depth,) with values in [0, 1]\n        \"\"\"\n        z = torch.arange(depth, device=device)\n        center  = (depth - 1) / 2\n        sigma   = depth * sigma_ratio\n        weights = torch.exp(-((z - center) ** 2) / (2 * sigma ** 2))\n        weights = weights / weights.max()  # Normalize to [0, 1]\n        return weights  # shape: (depth,)\n\n\n    # --- Center Heatmap ----\n    @staticmethod\n    def generate_center_heatmap(center_coord, shape=None, sigma=3.0, device='cpu'):\n        \"\"\"\n        center_coord: tensor (B, 3) with values in [0, D), [0, H), [0, W)\n        shape       : (D, H, W) — heatmap output shape\n        sigma       : std of the Gaussian\n        \"\"\"\n        B, _     = center_coord.shape\n        D, H, W  = shape\n        heatmaps = torch.zeros((B, 1, D, H, W), device=device)\n    \n        z = torch.arange(D, device=device).view(1, D, 1, 1).float()\n        y = torch.arange(H, device=device).view(1, 1, H, 1).float()\n        x = torch.arange(W, device=device).view(1, 1, 1, W).float()\n    \n        for b in range(B):\n            cz, cy, cx = center_coord[b]\n            heatmap = torch.exp(-((z - cz)**2 + (y - cy)**2 + (x - cx)**2) / (2 * sigma**2))\n            heatmaps[b, 0] = heatmap / heatmap.max()  # normalize to max=1\n    \n        return heatmaps\n\n    # Custom smooth_l1_loss\n    @staticmethod\n    def custom_smooth_l1_loss(inputs, target, beta=0.1, reduction='mean'):\n        \"\"\"\n        beta = 0.05  # 微細なずれに厳しく（過学習注意）\n        beta = 0.1   # 現状と同等\n        beta = 0.15  # 少しロバストに\n        beta = 0.2   # やや誤差に寛容に（外れ値対策）\n        \n        \"\"\"\n        n    = torch.abs(inputs - target)\n        cond = n < beta\n        loss = torch.where(cond, 0.5 * n ** 2 / beta, n - 0.5 * beta)\n        if reduction == 'mean':\n            return loss.mean()\n        elif reduction == 'sum':\n            return loss.sum()\n        else:\n            return loss\n\n    \n    def forward(self, x_inputs, inputs, targets, tomo_ids, motor_count=None, edge_maps=None, y_center=None, types=\"1stage\"):\n        raw_probs= inputs.float()\n        targets  = targets.float()\n        if self.use_center_heatmap and y_center is not None:\n            raw_center_probs = y_center.float()\n        \n        # === ガウスマスクに edge_map を加える ===\n        eps = 1e-6\n        \n        if self.use_edge_enhancement and edge_maps is not None:\n        \n            edge_maps = edge_maps.float()\n            \n            if self.edge_soft_mode:\n                edge_min = edge_maps.min()\n                edge_max = edge_maps.max()\n                denom    = edge_max - edge_min\n                if denom < self.clamp_eps:\n                    edge_maps = torch.zeros_like(edge_maps)\n                else:\n                    edge_maps = (edge_maps - edge_min) / (denom + self.clamp_eps)\n            else:\n                # Normalize（min-max）\n                edge_maps = (edge_maps - edge_maps.min()) / (edge_maps.max() - edge_maps.min() + self.clamp_eps)\n        \n            # edge強調（ハード／ソフト切り替え）\n            if self.edge_soft_mode:\n                edge_maps = (edge_maps + eps).pow(self.edge_power2)  # thinly\n            else:\n                edge_maps = edge_maps.pow(self.edge_power)\n        \n            # Apply edge weight map\n            if self.mode == \"1stage\":\n                # 背景にも少しエッジを残す\n                edge_weight_map = 0.4 + 0.7 * torch.sigmoid(6 * (targets - 0.05))\n            else:\n                edge_weight_map = 0.1 + 0.9 * torch.sigmoid(12 * (targets - 0.3))\n                #edge_weight_map = 0.1 + 0.9 * torch.sigmoid(20 * (targets - 0.05))\n        \n            # edge × edge_weight\n            edge_maps_filtered = edge_maps * edge_weight_map\n        \n            if self.mode == \"1stage\":\n                # エッジをターゲットに加算し、制限\n                target_enhanced = targets + (self.edge_weight * edge_maps_filtered)\n                target_enhanced = target_enhanced.clamp(0.0, 1.0)\n                target_enhanced = torch.where(target_enhanced > 0.01, target_enhanced, torch.zeros_like(target_enhanced))\n            else:\n                # log・sigmoid変換による強調（2stage）\n                targets_flipped = 1.0 - targets\n                    \n                enhanced_raw    = targets_flipped + (edge_maps_filtered * 2.5)\n                #enhanced_raw    = targets + (edge_maps_filtered * 2.5)\n                enhanced_raw    = torch.clamp(enhanced_raw, min=eps)  # log1p safe\n                target_enhanced = torch.log1p(enhanced_raw)\n                #target_enhanced = torch.sigmoid(12 * (target_enhanced - 0.70)) # eg: 0.70 ~ 0.72\n                target_enhanced = torch.sigmoid(20 * (target_enhanced - 0.78))\n                target_enhanced = (target_enhanced - target_enhanced.min()) / (target_enhanced.max() - target_enhanced.min() + eps)\n        \n        else:\n            # only targets\n            target_enhanced = targets\n\n        # NaN or Inf Check Before Normalize\n        if torch.isnan(target_enhanced).any() or torch.isinf(target_enhanced).any():\n            print(\"⚠️ target_enhanced has NaN or Inf BEFORE normalization!\")\n            print(f\"→ min: {target_enhanced.min().item():.4f}, max: {target_enhanced.max().item():.4f}\")\n            target_enhanced = torch.nan_to_num(target_enhanced, nan=0.0, posinf=1.0, neginf=0.0)\n        \n        # Safe min-max Normalize\n        min_val = target_enhanced.min()\n        max_val = target_enhanced.max()\n        denom   = max_val - min_val\n        \n        if denom < eps:\n            print(\"⚠️ target_enhanced has zero variance! Using zeros.\")\n            target_enhanced = torch.zeros_like(target_enhanced)\n        else:\n            target_enhanced = (target_enhanced - min_val) / (denom + eps)\n\n        # --- Probs Clamp ---\n        probs = raw_probs.clamp(eps, 1.0 - eps).to(dtype=torch.float32)\n        if self.use_center_heatmap:\n            center_probs = raw_center_probs.clamp(eps, 1.0 - eps).to(dtype=torch.float32)\n        \n        # --- 従来通りのしきい値ベースのマスクも生成 ---\n        targets_norm = target_enhanced.clamp(0.0, 1.0)\n        \n        if torch.isnan(targets_norm).any() or torch.isinf(targets_norm).any():\n            print(f\"⚠️ Targets Nan Min: {targets_norm.min().item()}\")\n            print(f\"⚠️ Targets Nan Max: {targets_norm.max().item()}\")\n            \n        targets_bin  = (targets_norm > 0).to(dtype=torch.float32)\n        \n        mask_sum     = targets_bin.sum().item()\n        max_val      = target_enhanced.max().item()\n        \n        #----- Z Y X Weight --------\n        \n        eps = 1e-8\n\n        # ======== YX Weignt =======\n        \n        # 1. ガウスボールが存在する場所だけマスク（≠ 0 の箇所）\n        if self.mode==\"1stage\":\n            gauss_mask = targets.to(dtype=torch.float32)\n        else:\n            #gauss_mask = (targets > 0.01).float() # or 0.05\n            gauss_mask = targets.float()\n\n        if self.mode==\"1stage\":\n            # 2. 通常通りしきい値でweights_maskを作る\n            max_val       = target_enhanced.max().item()\n            threshold     = max(max_val * self.base_thresh_ratio, self.min_thresh)\n            weights_mask  = (target_enhanced >= threshold).to(dtype=torch.float32)\n            \n            # 3. yx方向の平均と正規化\n            yx_mask = weights_mask / (weights_mask.mean(dim=[3, 4], keepdim=True) + eps)\n        else:\n            # threshold： gauss中心に寄せる\n            threshold    = max((target_enhanced * gauss_mask).max().item() * self.base_thresh_ratio, self.min_thresh)\n            weights_mask = (target_enhanced >= threshold).float()\n            \n            yx_mask      = weights_mask / (weights_mask.mean(dim=[3, 4], keepdim=True) + eps)\n            \n        if self.mode==\"1stage\":\n            yx_mask = yx_mask * torch.log1p(gauss_mask * 10)\n        else:\n            #yx_mask = yx_mask * gauss_mask\n            # シャープ化のための weight強調\n            yx_mask = yx_mask * torch.exp(gauss_mask * 1.5)\n        \n        if self.mode==\"1stage\":\n            # Hybrid\n            yx_mask_mean = yx_mask / (torch.mean(yx_mask, dim=[2, 3, 4], keepdim=True) + eps)\n            yx_mask_amax = yx_mask / (torch.amax(yx_mask, dim=[2, 3, 4], keepdim=True) + eps)\n            W=0.3\n            yx_mask = (1 - W) * yx_mask_mean + W * yx_mask_amax\n        else:\n            # Hybrid\n            yx_mask_mean = yx_mask / (torch.mean(yx_mask, dim=[2, 3, 4], keepdim=True) + eps)\n            yx_mask_amax = yx_mask / (torch.amax(yx_mask, dim=[2, 3, 4], keepdim=True) + eps)\n            \n            W=0.4 #0.2 ~ 0.6\n            yx_mask = (1 - W) * yx_mask_mean + W * yx_mask_amax\n            \n            # ★ max or mean normalizeでピークを明確にする\n            #amax=ピークが鋭くなる：点の予測に向いている\n            #yx_mask = yx_mask / (torch.amax(yx_mask, dim=[2, 3, 4], keepdim=True) + eps)\n\n        if self.mode==\"1stage\":\n            yx_mask = yx_mask.pow(self.YX_weight_power)\n        else:\n            # 尖鋭化（→ピークを作る）\n            #yx_mask = yx_mask.clamp(max=30.0)\n            if self.YX_weight_power != 1.0:\n                yx_mask = yx_mask.pow(self.YX_weight_power)\n            yx_mask = yx_mask / (yx_mask.mean(dim=[2, 3, 4], keepdim=True) + eps)\n\n            # --- Apply motor_count mask final block---\n            motor_mask  = (motor_count > 0).to(dtype=torch.float32).view(-1, 1, 1, 1, 1)\n            fallback_yx = torch.ones_like(yx_mask) * self.fallback_value\n            yx_mask     = yx_mask * motor_mask + fallback_yx * (1.0 - motor_mask)\n        #print(f\"Befor clamp yx_mask : Max: {yx_mask.max().item():.6f} Mean: {yx_mask.mean().item():.6f}\")\n        \n        # ==== Z Weight ======\n        if self.z_gaussian_weights:\n            # z_weights to Gaussian\n            z_weights = GaussianBallHybridLoss.get_z_gaussian_weights(depth=targets.shape[2], device=targets.device)  # shape: (Z,)\n            z_weights = z_weights.pow(self.Z_weight_power)\n            z_weights = z_weights.view(1, 1, -1, 1, 1)  # shape: (1,1,Z,1,1) for broadcasting\n        else:\n            # Used 1stage\n            z_weights = GaussianBallHybridLoss.get_z_importance_weights(targets_norm, \n                                                                        motor_count, \n                                                                        fallback_value=self.fallback_value\n                                                                       )\n            z_weights = z_weights.pow(self.Z_weight_power)\n\n        # ==== ALL Weight (Z + YX) =======\n        \n        if self.mode==\"1stage\":\n            weights = z_weights * yx_mask\n        else: \n            weights     = (z_weights * yx_mask) * self.zyx_ratio\n            motor_mask  = (motor_count > 0).view(-1, 1, 1, 1, 1)  # bool型で保持\n            weight_mask = (weights > 0)\n            motor_voxel_mask = motor_mask & weight_mask  # bool型で論理積\n\n            if self.final_weight_norm is not None:\n                # Final Weight min-max normalize or mean normalize or Hyblid\n                if self.final_weight_norm == \"mean\":\n                    # 重みを求めたあとに追加スケーリング\n                    weights = weights / (weights.mean(dim=[2, 3, 4], keepdim=True) + eps)\n                    \n                    # さらに倍率を調整（学習全体に影響）\n                    weights = weights * self.final_w_ratio  # or 0.05, 0.2\n                    \n                elif self.final_weight_norm == \"amax\":\n                    weights = weights / (weights.amax(dim=[2, 3, 4], keepdim=True)+ eps)\n                    \n                elif self.final_weight_norm == \"mixed\":\n                    # Hyblid: mean + amax\n                    # -----------------------------\n                    # motor=1 sample only mean, amax \n                    # -----------------------------\n                    if motor_voxel_mask.sum() > 0:\n                        mask_f = motor_voxel_mask.float()  # need to float\n                        valid_weights = weights * mask_f   # voxel Normalize\n                        W    = 0.4 # or 0.3\n                        denom = mask_f.sum(dim=[2,3,4], keepdim=True) + eps\n                        mean  = valid_weights.sum(dim=[2,3,4], keepdim=True) / denom\n                        amax  = valid_weights.amax(dim=[2,3,4], keepdim=True) + eps\n                    \n                        # Hyblid Normalize\n                        normed_valid = (1 - W) * (valid_weights / (mean + eps)) + W * (valid_weights / amax)\n                        \n                        # motor=1 の mean ReScale\n                        mean_normed  = normed_valid.sum(dim=[2,3,4], keepdim=True) / denom\n                        normed_valid = normed_valid / (mean_normed + eps)\n                        \n                        # Scaling（eg：0.05～1.0の範囲に）\n                        normed_valid = normed_valid * self.final_w_ratio\n                        # Join\n                        # motor_mask（bool）によって条件付き代入\n                        weights = torch.where(motor_mask, normed_valid, self.fallback_value)\n                    else:\n                         # fallback-only motor=0\n                        weights = torch.full_like(weights, fill_value=self.fallback_value)\n                        # if cf.debug:\n                        #     print(\"⚠️ motor_voxel_mask.sum() == 0 → fallback weights only.\")\n                    \n                    \n        # if self.mode==\"2stage\" and (torch.isnan(weights).any() or weights.max() > 10):\n        #     print(f\"⚠️[Weight Debug] Max: {weights.max().item():.6f} Min: {weights.min().item():.6f}\")\n            \n        if self.mode==\"1stage\":\n            weights = weights.clamp(min=0.05, max=1.0)\n            \n        elif self.mode==\"2stage\" and self.use_W_epoch_min:\n            clamp_min = self.get_clamp_min()\n            weights   = weights.clamp(min=clamp_min, max=1.0)\n            \n        else:\n            weights = weights.clamp(min=self.fallback_value, max=1.0) #0.05\n\n        if mask_sum == 0.0 and weights.mean().item() > 0.0:\n            print(f\"⚠️ Mask Sum & Weights Mean Zero ⚠️\")\n            if types == \"2stage_valid\":\n                show_pred_and_mask_one(probs, targets, target_enhanced, tomo_ids, motor_count, edge_maps, inputs=x_inputs)\n            else:\n                show_pred_and_mask(probs, targets, target_enhanced, tomo_ids, motor_count, edge_maps, inputs=x_inputs)\n\n        # --- BCE ---\n        bce_loss_raw = F.binary_cross_entropy(probs, targets_norm, reduction='none')\n        weighted_bce_per_sample = (bce_loss_raw * weights).view(probs.size(0), -1).mean(dim=1)\n\n        # --- logDice ---\n        intersection = (probs * targets_norm * weights).sum(dim=(1, 2, 3, 4))\n        probs_sum    = (probs * weights).sum(dim=(1, 2, 3, 4))\n        target_sum   = (targets_norm * weights).sum(dim=(1, 2, 3, 4))\n        dice_denominator = probs_sum + target_sum + self.smooth\n        _ratio    = (2 * intersection + self.smooth) / (dice_denominator + eps)\n        _ratio    = torch.clamp(_ratio, min=eps, max=1.0)\n        # soft logDice\n        dice_loss = -torch.log(_ratio + eps) * self.dice_alpha\n        \n        # --- Tversky ---\n        tp = (probs * targets_norm * weights).sum(dim=(1, 2, 3, 4))\n        fn = ((1 - probs) * targets_norm * weights).sum(dim=(1, 2, 3, 4))\n        fp = (probs * (1 - targets_norm) * weights).sum(dim=(1, 2, 3, 4))\n        tversky_index = (tp + self.smooth) / (tp + self.tversky_alpha * fn + self.tversky_beta * fp + self.smooth)\n        tversky_loss  = 1 - tversky_index\n\n        \"\"\"\n        •\talpha > beta → FPをより強くペナルティ：予測を保守的に（False Alarm減らす）→ Recall下がりPrecision上がる。\n    \t•\talpha < beta → FNをより強くペナルティ：予測を攻撃的に（漏れを減らす）→ Recall上がりPrecision下がる。\n        \"\"\"\n        \n        # --- 2stage Center loss ---\n        with torch.no_grad():\n            loss_center_per_sample, pred_com, target_com, dist_normalized = self.criterion_center(probs,\n                                                                                                  targets,\n                                                                                                  motor_count\n                                                                                                  )\n        # --- Z Center loss ---\n        # スケール調整 + MAE に変換した例\n        z_pred       = pred_com[:, 0]\n        z_true       = target_com[:, 0]\n        Z_DEPTH      = float(cf.TARGET_SLICES)\n        z_error      = (z_pred - z_true) / Z_DEPTH\n        z_per_sample = GaussianBallHybridLoss.custom_smooth_l1_loss(z_error, torch.zeros_like(z_error), beta=0.1)\n\n            \n        # --- Optional Sample Weight  Moter Count (Flagellum Presence)---\n        sample_weight = torch.ones_like(weighted_bce_per_sample)  # (B,)\n        if motor_count is not None:\n            #sample_weight = (motor_count.float() + 1.0)\n            sample_weight = torch.where(motor_count == 0, self.neg_ratio, self.pos_ratio)# motor=0: 1.0 motor=1: 2(stage=1)\n            sample_weight = torch.clamp(sample_weight, max=3.0)\n            sample_weight = sample_weight.view(-1)\n\n        # BCE / Dice / Tversky など\n        assert sample_weight.ndim == 1 and sample_weight.shape[0] == probs.shape[0], \\\n            f\"Invalid sample_weight shape: {sample_weight.shape}\"\n        \n        # --- BCE ---\n        weighted_bce = (sample_weight * weighted_bce_per_sample).mean()\n        if self.apply_bce_scaling:\n            _non_zero_ratio = (targets_bin.sum() / targets_bin.numel()).item()\n            _scale          = torch.log1p(torch.tensor(1.0 / (_non_zero_ratio + eps), device=probs.device))  # ややマイルド  # 非ゼロ率が小さいとscaleが大きくなる\n            _scale          = min(_scale, 5.0)  # 過度なスケーリングを防ぐ\n            weighted_bce   *= (1.0 + _scale)\n\n        if self.apply_bce_with_motor_count:\n            if motor_count is not None:\n                # flagellumなしの抑制\n                _bce_neg  = 0.5  # 0.2 ~ 0.5\n                neg_mask  = (motor_count == 0).view(-1, 1, 1, 1, 1)\n                bce_neg   = F.binary_cross_entropy(probs, torch.zeros_like(probs), reduction='none')\n                conf_mask = (probs > 0.3).float()  # 弱い出力はignore\n                penalty_mask = neg_mask.expand_as(bce_neg)\n                masked_bce   = (bce_neg * conf_mask)[penalty_mask]\n                \n                if masked_bce.numel() > 0:\n                    penalty_neg = masked_bce.mean()\n                else:\n                    penalty_neg = torch.tensor(0.0, device=probs.device, dtype=probs.dtype, requires_grad=False)\n        \n                weighted_bce += _bce_neg * penalty_neg\n        \n                # false positiveペナルティ\n                B       = probs.shape[0]  # batch size\n                _fp     = 0.35           # 0.3 ~ 0.6\n                fp_thre = 7\n                with torch.no_grad():\n                    max_conf    = probs.view(B, -1).max(dim=1).values\n                    fp_mask     = (motor_count == 0) & (max_conf > 0.4)\n                    fp_strength = max_conf[fp_mask]\n        \n                    if fp_strength.numel() > 0:\n                        fp_penalty = torch.log1p(fp_strength * fp_thre).mean()\n                    else:\n                        fp_penalty = torch.tensor(0.0, device=probs.device, dtype=probs.dtype, requires_grad=False)\n        \n                weighted_bce += _fp * fp_penalty\n\n\n        # --- logDice ---\n        dice_loss_per_sample = dice_loss  # shape: (B,)\n        weighted_dice = (sample_weight * dice_loss_per_sample).mean()\n        \n        # --- Tversky ---\n        tversky_loss_per_sample = tversky_loss  # shape: (B,)\n        weighted_tversky = (sample_weight * tversky_loss_per_sample).mean()\n\n        # --- Center ---\n        use_center  = (motor_count == 1).float()\n        center_mask = use_center * sample_weight\n        weighted_center = (loss_center_per_sample * center_mask).sum() / (center_mask.sum() + 1e-6)\n        #weighted_center = (loss_center_per_sample * sample_weight).mean()\n\n        # --- Center Heatmap ---\n        if self.use_center_heatmap:\n            with torch.no_grad():\n                target_center_heatmap = GaussianBallHybridLoss.generate_center_heatmap(target_com, shape=probs.shape[2:], device=probs.device)  # (B, 1, D, H, W)\n                target_center_heatmap = target_center_heatmap.detach()\n                target_center_heatmap = target_center_heatmap.to(dtype=torch.float32)\n            loss_center_heatmap_per_sample = F.mse_loss(center_probs, target_center_heatmap, reduction='none')  # (B, 1, D, H, W)\n            loss_center_heatmap_per_sample = loss_center_heatmap_per_sample.mean(dim=(1,2,3,4))  # → (B,)\n            weighted_center_heatmap = (sample_weight * loss_center_heatmap_per_sample).mean()\n            \n            del target_center_heatmap\n            del loss_center_heatmap_per_sample\n            _=gc.collect()\n            torch.cuda.empty_cache()\n        else:\n            weighted_center_heatmap = torch.tensor(0.0, device=probs.device)\n\n        # --- Z Loss ---\n        weighted_z_loss = (sample_weight * z_per_sample).mean()\n\n        # --- Focal Loss ---\n        with torch.no_grad():\n            weighted_focal   = self.FocalLoss(probs, \n                                              targets=targets,\n                                              voxel_weights=weights, \n                                              sample_weights=sample_weight\n                                             )\n\n        # --- DynamicWeight ---\n        if self.auto_weight:\n            loss_dict = {\n                'bce'           : weighted_bce.item(),\n                'dice'          : weighted_dice.item(),\n                'tversky'       : weighted_tversky.item(),\n                'focal'         : weighted_focal.item(),\n                'center'        : weighted_center.item(),\n                'center_heatmap': weighted_center_heatmap.item(),\n                'z_loss'        : weighted_z_loss.item(),\n            }\n            self.loss_tracker.update(loss_dict)\n            auto_weights = self.loss_tracker.get_weights()\n        else:\n            auto_weights = {\n                'bce'           : self.bce_weight,\n                'dice'          : self.dice_weight,\n                'tversky'       : self.tversky_weight,\n                'focal'         : self.focal_weight,\n                'center'        : self.center_weight,\n                'center_heatmap': self.center_heatmap_weight,\n                'z_loss'        : self.z_loss_weight,\n            }\n\n        if self.mode==\"2stage\" and self.use_center_weight:\n            total_loss = (\n                auto_weights[\"bce\"]             * weighted_bce +\n                auto_weights[\"dice\"]            * weighted_dice +\n                auto_weights[\"tversky\"]         * weighted_tversky +\n                auto_weights[\"center\"]          * weighted_center +\n                auto_weights[\"center_heatmap\"]  * weighted_center_heatmap +\n                auto_weights[\"z_loss\"]          * weighted_z_loss +\n                auto_weights[\"focal\"]           * weighted_focal\n            )\n        else:\n            total_loss = (\n                self.bce_weight             * weighted_bce +\n                self.dice_weight            * weighted_dice +\n                self.tversky_weight         * weighted_tversky +\n                self.center_heatmap_weight  * weighted_center_heatmap\n            )\n\n        # ---- EMA -----\n        nonzero_ratio = (targets > 0).to(dtype=torch.float32).mean().item()\n        self.ema_nonzero_ratio = (\n            self.ema_alpha * self.ema_nonzero_ratio + (1 - self.ema_alpha) * nonzero_ratio\n        )\n\n        ratio = torch.tensor(max(self.ema_nonzero_ratio, 1e-3), device=probs.device)\n        scale = torch.log1p((1.0 - ratio) * self.alpha + 1e-6)\n        total_loss *= (1.0 + scale)\n        ## EMA plot\n        self.ema_log[\"nonzero\"].append(nonzero_ratio)\n        self.ema_log[\"ema\"].append(self.ema_nonzero_ratio)\n        self.ema_log[\"scale\"].append(scale.item())\n    \n        if total_loss == 0.0:\n            print(f\"⚠️ Total zero : {tomo_ids}\")\n            if types == \"2stage_valid\":\n                show_pred_and_mask_one(probs, targets, target_enhanced, tomo_ids, motor_count, edge_maps, inputs=x_inputs)\n            else:\n                show_pred_and_mask(probs, targets, target_enhanced, tomo_ids, motor_count, edge_maps, inputs=x_inputs)\n\n        if self.training and self.log_interval > 0 and self.iteration % self.log_interval == 0:\n            print(f\"[LossLog] 🦁 Mode: {self.mode}\")\n            print(f\"Pred Sum  : {raw_probs.sum().item():.6f} Max: {raw_probs.max().item():.6f} Min: {raw_probs.min().item():.6f} Mean: {raw_probs.mean().item():.6f}\")\n            print(f\"Mask Shape: {target_enhanced.shape}\")\n            print()\n            if self.mode==\"2stage\" and self.use_center_weight:\n                print(f\"Auto Weight : {self.auto_weight}\")\n                bce_weight            = auto_weights[\"bce\"]\n                dice_weight           = auto_weights[\"dice\"]\n                tversky_weight        = auto_weights[\"tversky\"]\n                focal_weight          = auto_weights[\"focal\"]\n                center_weight         = auto_weights[\"center\"]\n                center_heatmap_weight = auto_weights[\"center_heatmap\"]\n                z_loss_weight         = auto_weights[\"z_loss\"]\n                print(f\"BCE Weight: {bce_weight:.4f} | logDice Weight: {dice_weight:.2f} | Tversky Weight: {tversky_weight:.2f} | Focal Weight: {focal_weight:.2f}| Center Weight: {center_weight:.2f} | Center Heatmap Weight: {center_heatmap_weight:.2f} | Z loss weith: {z_loss_weight:.2f}\")\n            else:\n                print(f\"BCE Weight: {self.bce_weight} | logDice Weight: {self.dice_weight} | Tversky Weight: {self.tversky_weight} | Center Heatmap Weight: {self.center_heatmap_weight}\")\n            \n            print(f\"BCE    : {weighted_bce.item()    :.10f}\")\n            print(f\"logDice: {weighted_dice.item()   :.10f}\")\n            print(f\"Tversky: {weighted_tversky.item():.10f}\")\n            if self.mode==\"2stage\":\n                print(f\"Focal           : {weighted_focal.item()  :.10f}\")\n                print(f\"W Center        : {weighted_center.item() :.10f}\")\n                print(f\"W Center heatmap: {weighted_center_heatmap.item():.10f}\")\n                print(f\"Z Loss          : {weighted_z_loss.item():.10f}\")\n                print(f\"Cneter threshold: {self.criterion_center.center_threshold}\")\n            if self.apply_bce_with_motor_count:\n                print()\n                print(f\"[Penalty] λ_neg={_bce_neg}, penalty_neg={penalty_neg.item():.4f} fp_thresold={fp_thre}\")\n                print(f\"[FP]      λ_fp ={_fp},　count={fp_mask.sum().item()}, avg_max_conf={max_conf[fp_mask].mean().item():.4f}\")\n                print(f\"[BCE Neg] penalty_neg={penalty_neg.item():.4f}, [FP Penalty] fp_penalty={fp_penalty.item():.4f}\")\n                print()\n            print(f\"Total  : {total_loss.item()      :.10f}\")\n            \n            print()\n            print()\n            print(f\"[Debug] targets_norm.min(): {targets_norm.min().item():.4f}\")\n            print(f\"[Debug] targets_norm.max(): {targets_norm.max().item():.4f}\")\n            print(f\"[Debug] Mask nonzero ratio: {nonzero_ratio:.4f}\")\n            print()\n            if self.mode==\"2stage\" and self.use_center_weight:\n                total = (weighted_bce) + (weighted_dice) + (weighted_tversky) + (weighted_center) + (weighted_z_loss)\n                wb    = bce_weight     * weighted_bce\n                wd    = dice_weight    * weighted_dice\n                wt    = tversky_weight * weighted_tversky\n                wc    = center_weight  * weighted_center\n                wch   = center_heatmap_weight * weighted_center_heatmap\n                z_l   = z_loss_weight * weighted_z_loss\n                wf    = focal_weight  * weighted_focal\n                total_raw  = wb + wd + wt + wc + wch + z_l + wf\n                correction = (1.0 + self.alpha * (1.0 - ratio))\n                total_corr = total_raw * correction\n                \n            else:\n                # 比率表示を追加\n                total = (weighted_bce) + (weighted_dice) + (weighted_tversky)\n                wb    = self.bce_weight     * weighted_bce\n                wd    = self.dice_weight    * weighted_dice\n                wt    = self.tversky_weight * weighted_tversky\n                wch   = self.center_heatmap_weight * weighted_center_heatmap\n                total_raw  = wb + wd + wt + wch\n                correction = (1.0 + self.alpha * (1.0 - ratio))\n                total_corr = total_raw * correction\n            \n            if total.detach().item() > 0:\n                if self.mode==\"2stage\" and self.use_center_weight:\n                    print(f\"Before Ratios (not weighted) : Total={total:.3f} BCE={weighted_bce/total:.2%} logDice={weighted_dice/total:.2%} Tversky={weighted_tversky/total:.2%} Focal={weighted_focal/total:.2%} Center={weighted_center/total:.2%} CenterHeat={weighted_center_heatmap/total:.2%} Z_loss={weighted_z_loss/total:.2%}\")\n                    print(f\"[Center] pred_com : \\n{pred_com.detach().cpu().numpy()}\")\n                    print(f\"[Center] targ_com : \\n{target_com.detach().cpu().numpy()}\")\n                    print(f\"[Center] dist_norm: \\n{dist_normalized.detach().cpu().numpy()}\")\n                else:\n                    print(f\"Before Ratios (not weighted) : Total={total:.3f} BCE={weighted_bce/total:.2%} logDice={weighted_dice/total:.2%} Tversky={weighted_tversky/total:.2%} CenterHeat={weighted_center_heatmap/total:.2%}\")\n                print()\n                if self.z_gaussian_weights:\n                    print(\"Use get_z_gaussian_weights  : Z is Fixed\")\n                else:\n                    print(\"Use get_z_importance_weights: Z is dynamic\")\n                print(f\"Z Weight   Mean  : {z_weights.mean().item():.6f} Max: {z_weights.max().item():.6f}   Min: {z_weights.min().item():.6f}   Z weight power : {self.Z_weight_power}\")\n                print(f\"YX Mask    Mean  : {yx_mask.mean().item()  :.6f} Max: {yx_mask.max().item():.6f}   Min: {yx_mask.min().item():.6f}   YX weight power : {self.YX_weight_power}\")\n                print(f\"ZYX Weight Mean  : {weights.mean().item()  :.6f} Max: {weights.max().item():.6f}   Min: {weights.min().item():.6f}   Final Weight ratio : {self.final_w_ratio}\")\n                print(f\"ZYXWeight  Shape : {weights.size()}\")\n                if self.mode==\"2stage\":\n                    #valid領域だけ抽出して統計出す\n                    # if motor_voxel_mask.sum() > 0:\n                    #     mask_bool  = motor_voxel_mask.bool()\n                    #     valid_max  = weights[mask_bool].max()\n                    #     valid_min  = weights[mask_bool].min()\n                    #     valid_mean = weights[mask_bool].mean()\n                    #     print(f\"ZYX Weight (motor_voxel motor=1): Max: {valid_max.item():.4f} Min: {valid_min.item():.4f} Mean: {valid_mean.item():.4f}\")\n                    #     print()\n                        print(f\"Edge power : {self.edge_power2}\")\n                        print(f\"edge_maps_filtered stats: Max: {edge_maps_filtered.max().item():.4f} Min: {edge_maps_filtered.min().item():.4f} Mean: {edge_maps_filtered.mean().item():.4f}\")\n                    # else:\n                    #     print(\"⚠️ motor_voxel_mask has no positive values!\")\n                print()\n                print()\n                if self.mode==\"2stage\" and self.use_center_weight:\n                    print(f\"After Ratios (add weighted) : Total={total_corr:.3f} BCE={wb/total_corr:.2%} logDice={wd/total_corr:.2%} Tversky={wt/total_corr:.2%} Focal={wf/total_corr:.2%} Center={wc/total_corr:.2%} CenterHeat={wch/total_corr:.2%} Z_loss={z_l/total_corr:.2%}\")\n                else:\n                    print(f\"After Ratios (add weighted) : Total={total_corr:.3f} BCE={wb/total_corr:.2%} logDice={wd/total_corr:.2%} Tversky={wt/total_corr:.2%} CenterHeat={wch/total_corr:.2%}\")\n                print()\n                if self.apply_bce_scaling:\n                    print(f\"[BCE Scaling]      : non_zero_ratio={_non_zero_ratio:.4f}, scale={_scale:.2f}\")\n                print(f\"[logDice] raw ratio: {_ratio.mean().item():.4f}\")\n                print(f\"[logDice] loss     : {dice_loss.mean().item():.4f}\")\n                print(f\"[Dice loss scale]  : {(dice_loss.mean() / (self.dice_weight + eps)).item():.4f}\")\n                print()\n                print(f\"Tversky : alpha={self.tversky_alpha} Beta={self.tversky_beta}\")\n                print()\n\n            print(f\"[Debug] sample_weight　　　: neg={self.neg_ratio} pos={self.pos_ratio}\")\n            print(f\"[Debug] Flagellum Presence: {sample_weight.cpu().numpy()}\")\n            print(f\"[Debug] EMA  nonzero ratio: {self.ema_nonzero_ratio:.4f} EMA alpha: {self.ema_alpha} Base alpha: {self.alpha}\")\n            print(f\"[Debug] EMA  : final Scale 1.0 + scale: {(1.0 + scale):.4f} nonzero_ratio={nonzero_ratio:.4f} ratio={ratio.item():.4f} scale={scale.item():.4f}\")\n            print(f\"[Debug] Mask Sum          : {mask_sum}\")\n            print(f\"[Debug] Voxel Weights Sum : {weights.sum().item()}\")\n            print(f\"[Debug] Voxel Weights Min : {weights.min().item():.4f} Max: {weights.max().item():.4f} Mean: {weights.mean().item():.4f}\")\n            print(f\"target_enhanced normed    : min={target_enhanced.min().item():.4f} max={target_enhanced.max().item():.4f} mean={target_enhanced.mean().item():.4f}\")\n            print(f\"yx_mask                   : min={yx_mask.min().item():.4f} max={yx_mask.max().item():.4f} mean={yx_mask.mean().item():.4f}\")\n            if self.mode==\"2stage\" and self.use_W_epoch_min:\n                print(f\"[Epoch {self.current_epoch}] Clamp=VoxelWeights Min: {clamp_min:.5f}\")\n            print()\n            if types == \"2stage_valid\":\n                # weights Show\n                if self.show_weight:\n                    weights_np = weights.detach().cpu().numpy().astype(np.float32)\n                    probs_np   = probs.detach().cpu().numpy().astype(np.float32)\n                    visualize_weight_and_prob_slices(weights_np, probs_np)\n\n                if self.show_peak_weight:\n                    plot_peak_z_slice(weights, tomo_ids)\n                # pred and mask Show\n                show_pred_and_mask_one(probs, targets, target_enhanced, tomo_ids, motor_count, edge_maps, inputs=x_inputs)\n            else:\n                # weights Show\n                if self.show_weight:\n                    weights_np = weights.detach().cpu().numpy().astype(np.float32)\n                    probs_np   = probs.detach().cpu().numpy().astype(np.float32)\n                    visualize_weight_and_prob_slices(weights_np, probs_np)\n\n                if self.show_peak_weight:\n                    plot_peak_z_slice(weights, tomo_ids)\n                    \n                # pred and mask Show\n                try:\n                    print(\"Start plotting\", flush=True)\n                    if self.show_pred_com_base:\n                        show_pred_and_mask_com_based(probs, targets, target_enhanced, tomo_ids, motor_count, edge_maps, inputs=x_inputs, target_com=target_com, pred_com=pred_com)\n                    else:\n                        show_pred_and_mask(probs, targets, target_enhanced, tomo_ids, motor_count, edge_maps, inputs=x_inputs, target_com=target_com)\n                    print(\"Finished plotting\", flush=True)\n                except Exception as e:\n                    print(f\"[PLOT ERROR] {e}\", flush=True)\n\n                print(f\"z_error mean: {z_error.mean().item():.4f}\")\n                print(f\"z_pred  mean: {z_pred.mean().item() :.4f}\")\n                print(f\"z_true  mean: {z_true.mean().item() :.4f}\")\n                print()\n                if cf.other_plot:\n                    fig, axs = plt.subplots(1, 3, figsize=(12, 4))\n                    # absなし\n                    axs[0].hist(z_error.detach().cpu().numpy(), bins=50, range=(-1, 1))\n                    axs[0].set_title(\"Z error distribution (raw)\")\n                    axs[0].set_xlabel(\"z_error (raw)\")\n                    axs[0].grid(True)\n                    \n                    # absあり\n                    axs[1].hist(z_error.abs().detach().cpu().numpy(), bins=50, range=(0, 1))\n                    axs[1].set_title(\"Z error distribution (absolute)\")\n                    axs[1].set_xlabel(\"|z_error|\")\n                    axs[1].grid(True)\n    \n                    # EMA\n                    log = self.ema_log\n                    if len(log[\"nonzero\"]) > 0:\n                        axs[2].plot(log[\"nonzero\"], label=\"Nonzero Ratio (Raw)\")\n                        axs[2].plot(log[\"ema\"], label=\"EMA Nonzero Ratio\", linestyle=\"--\")\n                        axs[2].plot(log[\"scale\"], label=\"Scale Factor (1 + scale)\", linestyle=\":\")\n                        \n                        axs[2].set_xlabel(\"Iteration\")\n                        axs[2].set_ylabel(\"Value\")\n                        axs[2].set_title(\"EMA-based Nonzero Ratio and Loss Scaling\")\n                        axs[2].legend()\n                        axs[2].grid(True)\n                    else:\n                        fig.delaxes(axs[2])\n                        \n                    plt.tight_layout()\n                    plt.show()\n                    plt.close(fig)\n                    del fig, axs\n                    _=gc.collect()\n        self.iteration += 1\n        return total_loss","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Optimizer","metadata":{}},{"cell_type":"code","source":"def print_parameter_status(model):\n    print(f\"{'Parameter':50s} {'Trainable'}\")\n    for name, param in model.named_parameters():\n        print(f\"{name:50s} {param.requires_grad}\")\n\n    print()\n    for name, param in model.named_parameters():\n        if 'scale' in name:\n            print(name, param.requires_grad, param.shape)\n\n    print()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_optimizer(cf, model, lr=cf.lr, wd=cf.weight_decay):\n    if cf.use_premodel and cf.freezing:\n        total_params         = sum(p.numel() for p in model.parameters())\n        all_trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n        print(f\"Total parameters        : {total_params:,}\")\n        print(f\"Trainable parameters    : {all_trainable_params:,}\")\n        print(f\"Non-trainable parameters: {total_params - all_trainable_params:,}\")\n        print()\n        for name, param in model.named_parameters():\n            # decoder freeze\n            if not any(key in name for key in [\"decoder\", \"edge_enhancer\"]): #Unfreezing parameter\n                param.requires_grad = False\n\n            # encoder stage + cbam unfreeze\n            if (\n                name.startswith(\"encoder.stages.1\")\n                or name.startswith(\"encoder.stages.2\")\n                or name.startswith(\"encoder.stages.3\")\n                or name.startswith(\"encoder.downsample_layers.1\")\n                or name.startswith(\"encoder.downsample_layers.2\")\n                or name.startswith(\"encoder.downsample_layers.3\")\n                #or name.startswith(\"encoder.cbam_f1\")\n                #or name.startswith(\"encoder.cbam_f2\")\n                #or name.startswith(\"encoder.cbam_f3\")\n                #or name.startswith(\"encoder.cbam_f4\")\n                #or name.startswith(\"encoder.se_f1\")\n                or name.startswith(\"encoder.se_f2\")\n                or name.startswith(\"encoder.se_f3\")\n                or name.startswith(\"encoder.se_f4\")\n            ):\n                param.requires_grad = True\n\n        # not trainable parameters でも trainingできる\n        #model.decoder.edge_scale_raw.requires_grad = True\n        \n        trainable_params = filter(lambda p: p.requires_grad, model.parameters())\n\n        if cf.use_AdamW:\n            print(f\"Optimizer    : AdamW Freezing\")\n            optimizer = AdamW(params=trainable_params, \n                              lr=lr, \n                              weight_decay=wd\n                             )\n        else:\n            print(f\"Optimizer    : Adam Freezing\")\n            optimizer = torch.optim.Adam(params=trainable_params, \n                                         lr=lr)\n            \n        print(f\"Trainable P  : {sum(p.requires_grad for p in model.parameters())}/ {all_trainable_params}\")\n    \n    else:\n        # normal\n        params = model.parameters()\n        if cf.use_AdamW:\n            print(f\"Optimizer    : AdamW\")\n            optimizer = AdamW(params=params, lr=lr, weight_decay=wd)\n        else:\n            print(f\"Optimizer    : Adam\")\n            optimizer = torch.optim.Adam(params, lr=lr)\n\n    if cf.use_premodel and cf.freezing and cf.parameter_check:\n        print_parameter_status(model)\n        \n    return optimizer","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Scheduler","metadata":{}},{"cell_type":"code","source":"def get_scheduler(cf, optimizer, train_size, max_lr=cf.max_lr, epochs=cf.epochs, warmup_epochs=cf.warmup_epochs):\n    print(f\"Sheduler     : OneCycleLR\")\n    print()\n    scheduler = OneCycleLR(optimizer=optimizer, \n                           max_lr=max_lr,\n                           epochs=epochs,\n                           steps_per_epoch=train_size, \n                           pct_start=warmup_epochs/epochs,\n                           div_factor=25.0,\n                           final_div_factor=1e3,\n                           anneal_strategy=\"cos\"\n                          )\n\n    return scheduler","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### AverageMeter","metadata":{}},{"cell_type":"code","source":"class AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val   = 0\n        self.avg   = 0\n        self.sum   = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val    = val\n        self.sum   += val * n\n        self.count += n\n        self.avg    = self.sum / self.count\n\n\ndef format_time(seconds):\n    if isinstance(seconds, str):\n        seconds = float(seconds)\n    minutes = int(seconds // 60)\n    sec     = seconds % 60\n    return f\"{minutes}m {sec:.1f}s\"","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### History","metadata":{}},{"cell_type":"code","source":"def plot_history(history, model_path=\".\", show=True):\n    epochs = range(1, len(history[\"train_loss\"]) + 1)\n\n    plt.figure()\n    plt.plot(epochs, history[\"train_loss\"], label=\"Training Loss\")\n    plt.plot(epochs, history[\"valid_loss\"], label=\"Validation Loss\")\n    plt.title(\"Loss evolution\")\n    plt.xlabel(\"Epochs\")\n    plt.ylabel(\"Loss\")\n    plt.legend()\n    plt.savefig(os.path.join(model_path, \"loss_evo.png\"))\n    if show:\n        plt.show()\n    plt.close()\n\n    plt.figure()\n    plt.plot(epochs, history[\"lr\"])\n    plt.title(\"Learning Rate evolution\")\n    plt.xlabel(\"Epochs\")\n    plt.ylabel(\"LR\")\n    plt.savefig(os.path.join(model_path, \"lr_evo.png\"))\n    if show:\n        plt.show()\n    plt.close()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.818Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Train Epoch","metadata":{}},{"cell_type":"code","source":"def plot_input_and_mask(image, mask, slice_idx=None, save_path=None):\n    # image: (1, D, H, W)\n    # mask:  (1, D, H, W)\n    \n    image = image.squeeze(0)  # → (D, H, W)\n    mask  = mask.squeeze(0)   # → (D, H, W)\n\n    if slice_idx is None:\n        slice_idx = image.shape[0] // 2  # 中央スライス\n\n    img_slice  = image[slice_idx].cpu().numpy()\n    mask_slice = mask[slice_idx].cpu().numpy()\n\n    fig, axs = plt.subplots(1, 2, figsize=(10, 5))\n    axs[0].imshow(img_slice, cmap=\"gray\")\n    axs[0].imshow(mask_slice, cmap=\"Reds\", alpha=0.5)\n    axs[0].set_title(\"Input Image\")\n    axs[1].imshow(mask_slice, cmap=\"hot\")\n    axs[1].set_title(\"Mask\")\n    for ax in axs:\n        ax.axis(\"off\")\n\n    plt.tight_layout()\n    if save_path:\n        plt.savefig(save_path)\n        print(f\"Saved to {save_path}\")\n    else:\n        plt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_feature_maps_grid(feature, name=\"f\", num_channels=3, slice_idx=None, downsample_to=128):\n    B, C, D, H, W = feature.shape\n\n    if slice_idx is None:\n        slice_idx = D // 2\n\n    num_channels = min(num_channels, C)\n\n    center   = C // 2\n    half     = num_channels // 2\n    start_ch = max(center - half, 0)\n    end_ch   = min(start_ch + num_channels, C)\n    selected_channels = list(range(start_ch, end_ch))\n\n    n_cols = len(selected_channels)\n    n_rows = B\n\n    fig, axes = plt.subplots(n_rows, n_cols, figsize=(n_cols * 2, n_rows * 2))\n\n    if B == 1 or n_cols == 1:\n        axes = np.atleast_2d(axes)\n\n    for b in range(B):\n        for i, c in enumerate(selected_channels):\n            ax = axes[b][i]\n            img = feature[b, c, slice_idx, :, :].detach().cpu()\n\n            # Down Sampling\n            img = F.interpolate(img[None, None, :, :], size=(downsample_to, downsample_to), mode='bilinear', align_corners=False)[0, 0]\n\n            ax.imshow(img, cmap='viridis')\n            ax.axis('off')\n            if b == 0:\n                ax.set_title(f\"ch{c}\", fontsize=10)\n            if i == 0:\n                ax.set_ylabel(f\"B{b}\", fontsize=10)\n\n    fig.suptitle(name, fontsize=14)\n    plt.tight_layout()\n    plt.show()\n    plt.clf()\n    plt.close()\n    del fig, axes\n    _=gc.collect()\n    torch.cuda.empty_cache()\n\ndef describe_tensor(tensor, name=\"tensor\"):\n    t = tensor.detach().cpu()\n    print(f\"{name}: min={t.min():.3f}, max={t.max():.3f}, mean={t.mean():.3f}, std={t.std():.3f}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_epoch(dataloaders, model, criterion, optimizer, epoch, scheduler, device):\n    seed_everything(cf.SEED)\n    model.train()\n    \n    scaler     = None #torch.cuda.amp.GradScaler(enabled=cf.use_amp)\n    train_size = len(dataloaders[\"train\"])\n    losses     = AverageMeter()\n    start      = end = time.time()\n    stream     = tqdm(dataloaders[\"train\"], total=train_size, unit=\"train_batch\", desc=\"Training\")\n    \n    for step, batch in enumerate(stream):\n        img        = batch[\"images\"].to(device)\n        mask       = batch[\"masks\"].to(device)\n        tomo_ids   = batch[\"tomo_id\"]\n        motor_cnt  = batch[\"Motor_count\"]\n        batch_size = mask.size(0)\n            \n        optimizer.zero_grad()\n        if cf.use_center_heatm:\n            y_preds, y_center, x_inputs, x_edge_proj = model(img)\n        else:\n            y_preds, x_inputs, x_edge_proj = model(img)\n        \n        if cf.f_check and (step % cf.print_freq2 == 0 or step == (train_size - 1)):\n            print(\"x_edge_proj\")\n            print(f\"min : {x_edge_proj.min().item() :.6f}\")\n            print(f\"max : {x_edge_proj.max().item() :.6f}\")\n            print(f\"mean: {x_edge_proj.mean().item():.6f}\")\n            print(f\"std : {x_edge_proj.std().item() :.6f}\")\n            # --- モデル実行（中間出力取得） ---\n            with torch.no_grad():\n                feats = model.encoder(img)  # ← Conv3DEncoder の出力\n                f1, f2, f3, f4 = feats\n                # F2\n                upsampled_f2 = F.interpolate(f2, size=f1.shape[2:], mode='trilinear', align_corners=False)\n                print(f\"Upsampled f2 Shape: {upsampled_f2.shape}\")\n                describe_tensor(upsampled_f2, \"f2_upsampled\")\n                plot_feature_maps_grid(upsampled_f2, name=\"f2_upsampled\")\n                # F3\n                upsampled_f3 = F.interpolate(f3, size=f1.shape[2:], mode='trilinear', align_corners=False)\n                print(f\"Upsampled f3 Shape: {upsampled_f3.shape}\")\n                describe_tensor(upsampled_f3, \"f3_upsampled\")\n                plot_feature_maps_grid(upsampled_f3, name=\"f3_upsampled\")\n                \n                for i, feat in enumerate([f1, f2, f3, f4], start=1):\n                    print(f\"f{i} shape: {feat.shape}\")\n                    describe_tensor(feat, f\"f{i}\")\n                    plot_feature_maps_grid(feat, name=f\"f{i}\")\n                del f1, f2, f3, f4, feats\n                _=gc.collect()\n        \n                \n        if mask.sum() == 0 and y_preds.sum() == 0:\n            print(f\"⚠️Non Mask & Non Preds id⚠️ : {tomo_ids}\")\n            loss = torch.tensor(1e-5, requires_grad=True).to(device)      \n        else:\n            loss = criterion(x_inputs, \n                             y_preds, \n                             mask, \n                             tomo_ids, \n                             motor_count=torch.tensor(batch[\"Motor_count\"]).to(device), \n                             edge_maps=x_edge_proj,\n                             y_center= y_center if cf.use_center_heatm else None\n                            )\n            \n        if torch.isnan(loss).any():\n            raise RuntimeError(\"Detected NaN....\")\n\n        losses.update(loss.item(), batch_size)\n        stream.set_postfix(OrderedDict(loss=f\"{loss.item():.6f}\", \n                                       lr=f\"{optimizer.param_groups[0]['lr']:.8f}\"\n                                      )\n                          )\n\n        if scaler is not None:\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm(model.parameters(), max_norm=cf.max_grad_norm or 1e-9)\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            loss.backward()  \n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)  \n            optimizer.step()\n            \n        del img\n        \n        try:\n            scheduler.step()\n        except ZeroDivisionError as e:\n            print(f\"⚠️[Warning] Scheduler step failed: {e}\")\n            print(\"Current tomo_id:\", batch[\"tomo_id\"])\n                \n        end = time.time()\n        if step % cf.print_freq == 0 or step == (train_size -1):\n            #lr     = optimizer.param_groups[0][\"lr\"]\n            lr     = scheduler.get_last_lr()[0] #onecycle\n            \n            elapsed_sec = time.time() - start\n            progress    = float(step + 1) / train_size\n            remain_sec  = elapsed_sec * (1 - progress) / progress\n            print()\n            print(f\"Epoch : [{epoch + 1}] [{step}/{train_size}] Elapsed => {format_time(elapsed_sec)} (remain {format_time(remain_sec)})\")\n            print(f\"Loss  : {losses.avg:.4f} LR : {lr:.10f}\")\n            print()\n            \n    del y_preds, loss, stream\n    if cf.use_center_heatm:\n        del y_center\n    _=gc.collect()\n    torch.cuda.empty_cache()\n    return losses.avg, lr","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Valid Epoch","metadata":{}},{"cell_type":"code","source":"def get_coords_from_high_conf_blob(pred,\n                                   selected_indices,\n                                   height,\n                                   width,\n                                   conf_threshold=0.5,\n                                   use_center_of_mass=True,\n                                   use_crop=True\n                                  ):\n    if isinstance(pred, torch.Tensor):\n        pred = pred.detach().cpu().numpy()\n    pred_np = np.squeeze(pred)\n\n    if pred_np.ndim != 3:\n        raise ValueError(f\"Expected pred with shape [Z, H, W], but got shape {pred_np.shape}\")\n\n    # 1. 閾値処理\n    binary_mask = (pred_np > conf_threshold).astype(np.uint8)\n\n    # Optional cropping\n    if use_crop:\n        coords = np.argwhere(binary_mask)\n        if len(coords) == 0:\n            return -1, -1, -1\n        zmin, ymin, xmin = coords.min(axis=0)\n        zmax, ymax, xmax = coords.max(axis=0) + 1\n        cropped_mask = binary_mask[zmin:zmax, ymin:ymax, xmin:xmax]\n        cropped_pred = pred_np[zmin:zmax, ymin:ymax, xmin:xmax]\n    else:\n        if np.sum(binary_mask) == 0:\n            return -1, -1, -1\n        cropped_mask = binary_mask\n        cropped_pred = pred_np\n        zmin, ymin, xmin = 0, 0, 0\n\n    # 2. ラベリングと blob 選定\n    labeled = label(cropped_mask, connectivity=3)\n    props   = regionprops(labeled, intensity_image=cropped_pred)\n    if len(props) == 0:\n        return -1, -1, -1\n\n    best = max(props, key=lambda r: r.mean_intensity)\n\n    # 3. 座標取得\n    if use_center_of_mass:\n        zc, yc, xc = best.weighted_centroid\n    else:\n        coords = best.coords\n        scores = cropped_pred[tuple(coords.T)]\n        top_idx = np.argmax(scores)\n        zc, yc, xc = coords[top_idx]\n\n    # 4. Crop offset を元に戻す\n    zc += zmin\n    yc += ymin\n    xc += xmin\n\n    # 5. 元画像スケールに変換\n    z_orig = selected_indices[int(round(zc))]\n    y_orig = yc * (height / pred_np.shape[1])\n    x_orig = xc * (width  / pred_np.shape[2])\n\n    return [round(z_orig), round(y_orig), round(x_orig)]","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_coords_from_pred_com(pred_heatmap, selected_indices, height, width, center_loss_module):\n    \"\"\"\n    pred_heatmap    : Tensor of shape (1, 1, D, H, W)\n    selected_indices: list[int] - z軸の元スライスインデックス\n    height, width   : int - 元画像サイズ\n    return          : [z, y, x] あるいは [-1, -1, -1]\n    \"\"\"\n    if pred_heatmap.ndim != 5:\n        raise ValueError(f\"Expected pred_heatmap shape (1, 1, D, H, W), but got {pred_heatmap.shape}\")\n\n    pred_com = center_loss_module.get_pred_center_with_threshold(pred_heatmap)\n\n    if pred_com is None or pred_com.shape != (1, 3):\n        return [-1, -1, -1]\n\n    try:\n        pred_com = pred_com.squeeze(0).tolist()  # (3,)\n    except:\n        return [-1, -1, -1]\n\n    if not isinstance(pred_com, list) or len(pred_com) != 3 or any(p < 0 for p in pred_com):\n        return [-1, -1, -1]\n\n    zc, yc, xc = map(lambda v: int(round(v)), pred_com)\n\n    if zc < 0 or zc >= len(selected_indices):\n        return [-1, -1, -1]\n\n    z_orig = selected_indices[zc]\n    y_orig = yc * (height / pred_heatmap.shape[-2])\n    x_orig = xc * (width  / pred_heatmap.shape[-1])\n\n    return [round(z_orig), round(y_orig), round(x_orig)]","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### IOU","metadata":{}},{"cell_type":"code","source":"def compute_iou(preds    : torch.Tensor,\n                targets  : torch.Tensor,\n                threshold: float = 0.1,# or 0.5\n                eps      : float = 1e-6,\n                binarize_targets: bool = False,\n                reduction       : str = \"mean\"  # \"mean\", \"median\", \"none\"\n               ):\n    \"\"\"\n    Compute IoU between predicted and target masks.\n    \"\"\"\n    preds_bin   = (preds > threshold).float()\n    targets_bin = (targets > threshold).float() if binarize_targets else targets\n\n    intersection = (preds_bin * targets_bin).sum(dim=[1, 2, 3, 4])\n    union = preds_bin.sum(dim=[1, 2, 3, 4]) + targets_bin.sum(dim=[1, 2, 3, 4]) - intersection\n    ious  = (intersection + eps) / (union + eps)\n\n    if reduction == \"mean\":\n        return ious.mean().item()\n    elif reduction == \"median\":\n        return ious.median().item()\n    elif reduction == \"none\":\n        return ious  # shape: (B,)\n    else:\n        raise ValueError(f\"Unknown reduction method: {reduction}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def flatness_plot(positive_flatnesses, negative_flatnesses, save_path=None):\n    plt.figure(figsize=(10, 6))\n    plt.hist(positive_flatnesses, bins=50, alpha=0.6, color=\"blue\", label=\"Positive (has motor)\")\n    plt.hist(negative_flatnesses, bins=50, alpha=0.6, color=\"red\",  label=\"Negative (no motor)\")\n    plt.axvline(np.mean(positive_flatnesses), color=\"blue\", linestyle=\"dashed\", linewidth=1)\n    plt.axvline(np.mean(negative_flatnesses), color=\"red\",  linestyle=\"dashed\", linewidth=1)\n    plt.xlabel(\"Flatness (max - mean of pred)\")\n    plt.ylabel(\"Count\")\n    plt.title(\"Distribution of Flatness in Validation Set\")\n    plt.legend()\n    plt.grid(True)\n    \n    if save_path:\n        print(\"Save Flatness Plot\")\n        plt.savefig(save_path)\n    else:\n        plt.show()\n    plt.close()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def print_grad(grad):\n    print(\"scale grad:\", grad)\n\n# output_temp を予測に応じて自動調整する関数\ndef adjust_output_temp(model, pred_mean, target_mean=0.70, k=0.5, min_temp=1.2, max_temp=2.5):\n    delta = target_mean - pred_mean\n    scale = 1.0 + delta * k  # 例: pred_meanが0.6なら 1 + 0.1*0.5 = 1.05 → 濃くする\n    model.output_temp *= scale\n    model.output_temp = max(min_temp, min(max_temp, model.output_temp))","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def valid_epoch_with_coord(dataloaders, model, criterion, device, fold, epoch,\n                           center_loss_module,\n                           use_pred_com_coords=True\n                          ):\n    seed_everything(cf.SEED)\n    model.eval()\n    if center_loss_module is not None:\n        print(\"🌷 GET Coords Pred_com Base\")\n    \n    valid_size = len(dataloaders[\"valid\"])\n    losses = AverageMeter()\n    stream = tqdm(dataloaders[\"valid\"], total=valid_size, unit=\"valid_batch\", desc=\"Valid\")\n\n    labels        = list()\n    tomo_ids      = list()\n    preds         = list()\n    has_motors    = list()\n    voxel_spacing = list()\n    ious          = list()\n    flatness_list_pos = list()\n    flatness_list_neg = list()\n    all_ypreds        = list()\n    num_neg     = 0\n    motor1_seen = False\n\n    for batch in stream:\n        img        = batch[\"images\"].to(device)\n        mask       = batch[\"masks\"].to(device)\n        batch_size = img.size(0)\n    \n        tomo_id    = batch[\"tomo_id\"]\n        label      = batch[\"labels\"].detach().cpu().numpy()\n        select_ind = batch[\"selected_indices\"].detach().cpu().numpy()\n        h          = batch[\"h\"].cpu().numpy()\n        w          = batch[\"w\"].cpu().numpy()\n        has_motor  = batch[\"has_motor\"].cpu().numpy()\n        voxel_sp   = batch[\"voxel_spacing\"].cpu().numpy()\n    \n        # ↓ motor_count == 1 が1件は含まれているようにチェック（has_motorが代替になっていればOK）\n        motor_count = batch[\"Motor_count\"]  # torch.Tensor\n        has_motor1  = (motor_count == 1).any().item()\n\n        if not motor1_seen and not has_motor1:\n            continue  # skip batch\n        else:\n            motor1_seen = motor1_seen or has_motor1\n\n        with torch.no_grad():\n            if cf.use_center_heatm:\n                y_preds, y_center, x_input, x_edge_proj = model(img)\n            else:\n                y_preds, x_input, x_edge_proj = model(img)\n                \n        all_ypreds.append(y_preds)\n        loss = criterion(x_input, \n                         y_preds, \n                         mask,\n                         tomo_ids=tomo_id,\n                         motor_count=torch.tensor(batch[\"Motor_count\"], dtype=torch.float32).to(device),\n                         edge_maps=x_edge_proj,\n                         y_center=y_center if cf.use_center_heatm else None\n                        )\n\n        if torch.isnan(loss).any():\n            raise RuntimeError(\"Detected NaN...😭\")\n\n        if not motor1_seen:\n            print(\"[Warning] No batch with motor_count == 1 was seen during validation!\")\n        \n        losses.update(loss.item(), batch_size)\n        \n        for i in range(batch_size):\n            y_pred_i     = y_preds[i:i+1]  # shape: (1, C, Z, H, W)\n            mask_i       = mask[i:i+1]\n            x_input_i    = x_input[i:i+1]\n            x_edge_proj_i= x_edge_proj[i:i+1]\n            tomo_i       = tomo_id[i]\n            label_i      = label[i]\n            select_ind_i = select_ind[i]\n            h_i          = h[i]\n            w_i          = w[i]\n            has_motor_i  = has_motor[i]\n            voxel_sp_i   = voxel_sp[i]\n            if cf.use_center_heatm:\n                y_center_i = y_center[i:i+1]\n\n            # Flatness \n            flatness_i = float(y_pred_i.max().item() - y_pred_i.mean().item())\n        \n            if label_i[0] != -1:\n                flatness_list_pos.append(flatness_i)\n            else:\n                flatness_list_neg.append(flatness_i)\n\n            \n            # Get Orignal Coord\n            if use_pred_com_coords:\n                coord = get_coords_from_pred_com(y_pred_i,\n                                                 select_ind_i,\n                                                 h_i,\n                                                 w_i,\n                                                 center_loss_module\n                                                )\n            else:\n                coord = get_coords_from_high_conf_blob(y_pred_i,\n                                                       select_ind_i,\n                                                       h_i, \n                                                       w_i,\n                                                      )\n            if coord == [-1, -1, -1]:\n                num_neg += 1\n            \n            if label_i[0] != -1:\n                iou = compute_iou(y_pred_i, mask_i)\n                ious.append(iou)\n\n\n            if cf.s2valid_debug:\n                print()\n                print(f\"y_preds  Shape : {y_preds.shape}\")\n                print(f\"y_pred_i Shape : {y_pred_i.shape}\")\n                # y_pred_i の統計チェック\n                print(\"max         :\", y_pred_i.max().item())\n                print(\"mean        :\", y_pred_i.mean().item())\n                print(\"sum         :\", y_pred_i.sum().item())\n                print(\"contains inf:\", torch.isinf(y_pred_i).any().item())\n                print(\"contains nan:\", torch.isnan(y_pred_i).any().item())\n\n                print(f\"get_original_coords_from_crop → coord={coord}, label={label_i}\")\n                print(\"→ coord-label distance =\", np.linalg.norm(coord - label_i))\n                print(f\"[Valid] IoU   : {iou:.4f}, loss: {loss.item():.4f}\")\n                print()\n\n            preds.append(coord)\n            tomo_ids.append(tomo_i)\n            labels.append(label_i)\n            has_motors.append(has_motor_i)\n            voxel_spacing.append(voxel_sp_i)\n\n    # データフレーム化 & 評価\n    fixed_col_gt   = [\"tomo_id\", \"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\", \"Voxel spacing\", \"Has motor\"]\n    fixed_col_pred = [\"tomo_id\", \"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"]\n\n    pred_df = pd.DataFrame(preds, columns=[\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"])\n    pred_df[\"tomo_id\"] = tomo_ids\n    pred_df = pred_df[fixed_col_pred]\n\n    gt_df = pd.DataFrame(labels, columns=[\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"])\n    gt_df[\"tomo_id\"]       = tomo_ids\n    gt_df[\"Has motor\"]     = has_motors\n    gt_df[\"Voxel spacing\"] = voxel_spacing\n    gt_df = gt_df[fixed_col_gt]\n\n    #===IOU====\n    mean_iou = np.mean(ious) if len(ious) > 0 else 0.0\n\n    if cf.debug:\n        display(gt_df)\n        display(pred_df)\n        \n    assert len(pred_df) == len(gt_df), \"Mismatch Length...\"\n    valid_score = score(gt_df, pred_df, 1_000, 2)\n    print(f\"Comp Score : {valid_score:.6f}\")\n\n\n    # 各軸ごとのズレ（誤差）を記録\n    errors = list()\n    for i in range(len(gt_df)):\n        if gt_df[\"Has motor\"].iloc[i] == 0:\n            continue  # No motor → 無視（あるいはcontinueせずNaN入れる）\n    \n        gt_coord   = gt_df.iloc[i][  [\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"]].values.astype(float)\n        pred_coord = pred_df.iloc[i][[\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"]].values.astype(float)\n        \n        # 無効な [-1, -1, -1] を除外\n        if (gt_coord == -1).all() or (pred_coord == -1).all():\n            continue\n            \n        diff = pred_coord - gt_coord  # 差分（予測 - 正解）\n        errors.append(diff)\n        \n    if len(errors) == 0:\n        print(\"⚠️ No valid motor predictions to compute errors.\")\n    else:\n        errors = np.stack(errors)  # [N, 3]\n        print(\"Axis-wise error (mean ± std):\")\n        for axis, name in enumerate([\"Z\", \"Y\", \"X\"]):\n            mean = errors[:, axis].mean()\n            std  = errors[:, axis].std()\n            print(f\"{name}-axis: {mean:.2f} ± {std:.2f}\")\n    \n        fig, ax = plt.subplots(1, 3, figsize=(12, 4))\n        for i, axis in enumerate([\"Z\", \"Y\", \"X\"]):\n            ax[i].hist(errors[:, i], bins=20, color=\"steelblue\", alpha=0.7)\n            ax[i].set_title(f\"{axis}-axis error\", size=8)\n            ax[i].axvline(0, color='red', linestyle='--')\n        plt.tight_layout()\n        plt.show()\n        \n        if cf.debug:\n            print(\"Debug Sample:\")\n            for i in range(5):\n                gt   = gt_df.iloc[i][  [\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"]].values\n                pred = pred_df.iloc[i][[\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"]].values\n                diff = pred - gt\n                print(f\"GT:   [{gt[0]:>6.1f}, {gt[1]:>6.1f}, {gt[2]:>6.1f}] \"\n                      f\"PRED: [{pred[0]:>6.1f}, {pred[1]:>6.1f}, {pred[2]:>6.1f}] \"\n                      f\"DIFF: [{diff[0]:>+6.1f}, {diff[1]:>+6.1f}, {diff[2]:>+6.1f}]\")\n        \n        print()\n        print(f\"🎈Validation IoU Score: {mean_iou:.4f}\")\n        print(\"\\n[Flatness Monitor]\")\n        print(f\"Positive (label != -1) flatness avg: {np.mean(flatness_list_pos):.6f}, min: {np.min(flatness_list_pos):.6f}\")\n        print(f\"Negative (label == -1) flatness avg: {np.mean(flatness_list_neg):.6f}, min: {np.min(flatness_list_neg):.6f}\")\n        # scale の値を見る\n        #model.edge_enhancer.scale.register_hook(print_grad)\n        print(f\"[Epoch {epoch}] Edge Scale: {model.edge_enhancer.scale.data.flatten().tolist()}\")\n    \n        # Show（save or plot）\n        if cf.flatness_plot and epoch % 5:\n            save_name = f\"fold{fold}_epoch{epoch}_flatness_hist_valid.png\"\n            if cf.save_plot:\n                SAVE_PATH = os.path.join(FLATNESS_DIR, save_name)\n            else:\n                SAVE_PATH = None\n            flatness_plot(flatness_list_pos, flatness_list_neg, save_path=SAVE_PATH)\n        print()\n    \n    # shape: List of tensors → [B, 1, D, H, W]\n    all_ypreds_tensor = torch.cat(all_ypreds, dim=0)  # (total_B, 1, D, H, W)\n    mean_pred_value   = all_ypreds_tensor.mean().item()\n    adjust_output_temp(model, mean_pred_value)\n    # ログ表示\n    print(f\"[AutoTemp] pred_mean={mean_pred_value:.4f} → output_temp={model.output_temp:.3f}\")\n    print(f\"Flagellum not found in {num_neg}/ {len(dataloaders['valid'].dataset)} samples\")\n    print()\n    if len(errors) > 0:\n        del errors, gt_coord, pred_coord, diff, gt_df, mean_iou\n    del all_ypreds, all_ypreds_tensor, mean_pred_value\n    if cf.flatness_plot:\n        del flatness_list_pos, flatness_list_neg\n\n    if cf.use_center_heatm:\n        del y_center\n        \n    _ = gc.collect()\n    torch.cuda.empty_cache()\n    \n    if cf.GET_OOF:\n        return losses.avg, valid_score, pred_df\n    else:\n        return losses.avg","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def valid_epoch(dataloaders, model, criterion, device):\n    seed_everything(cf.SEED)\n    model.eval()\n    \n    valid_size = len(dataloaders[\"valid\"])\n    losses     = AverageMeter()\n    start      = end = time.time()\n    stream     = tqdm(dataloaders[\"valid\"], total=valid_size, unit=\"valid_batch\", desc=\"Valid\")\n\n    labels        = list()\n    tomo_ids      = list()\n    preds         = list()\n    has_motors    = list()\n    voxel_spacing = list()\n    ious          = list()\n    \n    for step, batch in enumerate(stream):\n        tomo_id   = batch[\"tomo_id\"]\n        label     = batch[\"labels\"].detach().cpu().numpy()\n        select_ind= batch[\"selected_indices\"].detach().cpu().numpy()\n        h         = batch[\"h\"].cpu().numpy()\n        w         = batch[\"w\"].cpu().numpy()\n        has_motor = batch[\"has_motor\"].cpu().numpy()\n        voxel_sp  = batch[\"voxel_spacing\"].cpu().numpy()\n        \n        img       = batch[\"images\"].to(device)\n        mask      = batch[\"masks\"].to(device)\n        batch_size= mask.size(0)\n\n        with torch.no_grad():\n            if cf.use_center_heatm:\n                y_preds, y_center, x_inputs, x_edge_proj = model(img)\n            else:\n                y_preds, x_inputs, x_edge_proj = model(img)\n                \n            iou = compute_iou(y_preds, mask)\n            ious.append(iou)\n            \n        if mask.sum() == 0 and y_preds.sum() == 0:\n            print(f\"⚠️ Non Mask & Non Preds id ⚠️ : {tomo_id}\")\n\n        loss = criterion(x_inputs, \n                         y_preds, \n                         mask, \n                         tomo_id, \n                         motor_count=torch.tensor(batch[\"Motor_count\"], dtype=torch.float32).to(device), \n                         edge_maps=x_edge_proj,\n                         y_center=y_center if cf.use_center_heatm else None\n                        )\n            \n        \n        if torch.isnan(loss).any():\n            raise RuntimeError(\"Detected NaN...😭\")\n\n        losses.update(loss.item(), batch_size)\n        \n    mean_iou = np.mean(ious)\n    print(f\"🎈Validation IoU: {mean_iou:.4f}\")\n\n    del mean_iou\n    _=gc.collect()\n    torch.cuda.empty_cache()\n    return losses.avg","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Train Loop","metadata":{}},{"cell_type":"code","source":"def linear_schedule(epoch, start_epoch, end_epoch, start_value, end_value):\n    ratio = (epoch - start_epoch) / (end_epoch - start_epoch)\n    ratio = min(max(ratio, 0.0), 1.0)\n    return start_value + (end_value - start_value) * ratio","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_loop(df:pd.DataFrame, fold:int):\n    history = {\"train_loss\": list(),\n               \"valid_loss\": list(),\n               \"lr\"        : list()\n              }\n    dataloaders = get_dataloaders(df, fold)\n    if cf.use_premodel:\n        EDGE_SCALE_CHANGE= False\n        SCALE            = 0.5 #or -0.7\n        \n        check_point = glob.glob(os.path.join(cf.model_dir, \"models\", \"*.pth\"))[0]\n        model       = get_model(check_point)\n        if EDGE_SCALE_CHANGE:\n            print(f\"Edge Scale Change: scale= {SCALE}\")\n            # edge_scale_raw の初期値を変更\n            with torch.no_grad():\n                model.decoder.edge_scale_raw.copy_(torch.tensor(SCALE))\n    else:\n        model = build_model()\n        \n    optimizer   = get_optimizer(cf, model)\n    \n    scheduler   = get_scheduler(cf, \n                                optimizer, \n                                len(dataloaders[\"train\"])\n                               )\n\n    center_loss_module = CenterLoss(center_threshold=0.05)\n    \n    criterion1 = GaussianBallHybridLoss(bce_weight=cf.bce_weight,\n                                        dice_weight=cf.dice_weight, \n                                        tversky_weight=cf.tversky_weight,\n                                        tversky_alpha=0.2,\n                                        mode=\"1stage\"\n                                       ).to(device)\n\n    criterion2 = GaussianBallHybridLoss(bce_weight=cf.bce_weight2,\n                                        dice_weight=cf.dice_weight2, \n                                        tversky_weight=cf.tversky_weight2,\n                                        focal_weight=cf.focal_weight2,\n                                        center_weight=cf.center_weight,\n                                        center_heat_w=cf.center_heat_w,\n                                        z_loss_weight=cf.z_loss_weight,\n                                        tversky_alpha=0.3, # or 0.75\n                                        tversky_beta=0.7, # or 0.25\n                                        mode=\"2stage\",\n                                        center_loss_module=center_loss_module\n                                       ).to(device)\n        \n    \n    best_loss        = np.inf\n    not_decreased    = 0\n    epoch_best_score = list()\n    \n\n    print(f\"Use SEED           : {cf.SEED}\")\n    print(f\"Target slices      : {cf.TARGET_SLICES}\")\n    print(f\"Train Batch size   : {cf.train_batch_size}\")\n    print(f\"Valid Batch size   : {cf.valid_batch_size}\")\n    print(f\"Epochs             : {cf.epochs}\")\n    print(f\"Learning rate      : {cf.lr}\")\n    print(f\"Max Learnign rate  : {cf.max_lr}\")\n    print(f\"Early stoping round: {cf.es_round}\")\n    print(f\"Valid Warmup epochs: {cf.valid_warmup}\")\n    print(f\"Use CustomOutput   : {cf.custom_output}\")\n    print(f\"Z_sampling         : {cf.z_sampling}\")\n    print(f\"Use Pretrainedmodel: {cf.use_premodel}\")\n    print(f\"Print Frequency    : {cf.print_freq}\")\n    print(f\"Use Weight epoch   : {cf.use_W_epoch_min}\")\n    print(f\"Freezing           : {cf.freezing}\")\n    print()\n    \n    epochs = cf.epochs\n    for epoch in range(epochs):\n        if cf.use_W_epoch_min:\n            print(\"🤡 Set epoch Reduce ZYX Weight clamp min : Loss Function\")\n            criterion2.set_epoch(epoch)\n        start_time = time.time()\n            \n        if epoch < cf.valid_warmup:\n            print(\"1Stage\")\n            # Train\n            train_avg_loss, lr = train_epoch(dataloaders, \n                                             model, \n                                             criterion1, \n                                             optimizer,\n                                             epoch, \n                                             scheduler, \n                                             device, \n                                            )\n            \n            # Valid\n            print(f\"🌾 Witn 1stage infer : [{epoch+1}] 🌾\")\n            valid_avg_loss = valid_epoch(dataloaders,\n                                         model, \n                                         criterion1,\n                                         device\n                                        )\n\n        else:\n            print(\"2Stage\")\n            # Auto Weight\n            if epoch < 3:\n                criterion2.auto_weight = False\n            else:\n                criterion2.auto_weight = True\n            \n            # Scheduling\n            criterion2.pos_ratio             = linear_schedule(epoch, 0, cf.epochs, 2.0, 1.0)\n            criterion2.center_weight         = linear_schedule(epoch, 0, cf.epochs, 0.05, 0.20)\n            criterion2.center_heatmap_weight = linear_schedule(epoch, 0, cf.epochs, 0.08, 0.22)\n            criterion2.criterion_center.center_threshold = linear_schedule(epoch, 0, cf.epochs, 0.05, 0.27)\n            criterion2.edge_power2           = linear_schedule(epoch, 0, cf.epochs, 1.7, 0.0)\n            \n            \n            train_avg_loss, lr = train_epoch(dataloaders, \n                                             model, \n                                             criterion2, \n                                             optimizer,\n                                             epoch, \n                                             scheduler, \n                                             device,\n                                            )\n        \n            # Valid\n            print(f\"⛳️ With 2stage infer: [{epoch+1}] ⛳️\")\n            if cf.GET_OOF:\n                print(\"🚀Get OOF\")\n                valid_avg_loss, valid_score, oof_df = valid_epoch_with_coord(dataloaders,\n                                                                             model,\n                                                                             criterion2, \n                                                                             device,\n                                                                             fold,\n                                                                             epoch,\n                                                                             center_loss_module\n                                                                             )\n            else:\n                valid_avg_loss = valid_epoch_with_coord(dataloaders,\n                                                        model,\n                                                        criterion2, \n                                                        device,\n                                                        fold,\n                                                        epoch\n                                                        )\n                    \n        elapsed = time.time() - start_time\n        \n        history[\"train_loss\"].append(train_avg_loss)\n        history[\"valid_loss\"].append(valid_avg_loss)\n        history[\"lr\"].append(lr)\n\n        if valid_avg_loss < best_loss:\n            best_loss = valid_avg_loss\n\n            print()\n            # Loss judge\n            print(f\"[{epoch + 1}/{epochs}] Train Loss : {train_avg_loss:.6f} Valid Loss : {valid_avg_loss:.6f} Time : {elapsed:.4f}s\")\n\n            print(\"~~~~~~~~ Improved ~~~~~~\")\n            print()\n            print(f\"Loss      : {best_loss  :.6f}\")\n            if cf.GET_OOF:\n                print(f\"Comp Score: {valid_score:.6f}\")\n            print()\n            print(\"~~~~~~~~~~~~~~~~~~~~~~~~~\")\n            print(\"🌸　Save Model!!!! 🌸\")\n\n            model_save_name = f\"Unet3d_fold{fold+1}.pth\"\n            torch.save({\"model\": model.state_dict()},\n                       os.path.join(MODEL_DIR, model_save_name)\n                      )\n            print()\n            not_decreased = 0\n            if cf.GET_OOF:\n                epoch_best_score.append(valid_score)\n                best_oof = oof_df.copy()\n        else:\n            print(f\"[{epoch + 1}/{epochs}] Train Loss : {train_avg_loss:.6f} Valid Loss : {valid_avg_loss:.6f} Time : {elapsed:.4f}\")\n            print(\"~~~~~~ Not Decreased ~~~~\")\n            print()\n            print(f\"Loss : {valid_avg_loss:.6f}\")\n            if cf.GET_OOF:\n                print(f\"Score: {valid_score:.6f}\")\n            print()\n            print(\"~~~~~~~~~~~~~~~~~~~~~~~~~\")\n            \n            not_decreased  += 1\n            if not_decreased  >= cf.es_round:\n                print(f\"Current Epoch is {epoch+1}. Early Stopping...\")\n                break\n\n    plot_history(history, model_path=HISTORY_DIR)\n    \n    history_path = os.path.join(HISTORY_DIR, f\"history_fold0.json\")\n    with open(history_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(history, f, ensure_ascii=False, indent=4)\n\n    if cf.GET_OOF:\n        epoch_max_score = np.max(epoch_best_score)\n        print()\n        print(f\"Best Score is {epoch_max_score:.6f}\")\n        \n    del model, optimizer, scheduler, dataloaders, history\n    _=gc.collect()\n    torch.cuda.empty_cache()\n    print(\"Done\")\n    \n    if cf.GET_OOF:\n        return best_oof\n    else:\n        return best_loss","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"all_oof = list()\nall_loss= list()\nfor fold in cf.FOLD_LIST:\n    print()\n    print(f\"{'>'.join('~+')*6} Fold {fold+1}/{len(cf.FOLD_LIST)} {'<'.join('+~')*6}\")\n    print()\n    if cf.GET_OOF:\n        oof = train_loop(train, fold)\n        all_oof.append(oof)\n    else:\n        loss = train_loop(train, fold)\n        all_loss.append(loss)\n        \nif cf.GET_OOF: \n    final_oof = pd.concat(all_oof)\n    del all_oof\n    _=gc.collect()","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cf.GET_OOF:\n    final_oof.to_csv(os.path.join(OOF_DIR, \"oof_df.csv\"), index=False)\nelse:\n    print(all_loss)\n    _=gc.collect()\n    torch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.819Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Compute Overall Metric","metadata":{}},{"cell_type":"code","source":"fixed_col_gt = [\"tomo_id\", \"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\", \"Voxel spacing\", \"Has motor\"]","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cf.debug:\n    print(\"Debug...\")\n    gt_df = train[train[\"fold\"]==cf.FOLD_LIST[0]].reset_index(drop=True)\n    gt_df = gt_df.rename(columns={\"flagellum_z\": \"Motor axis 0\", \"flagellum_y\": \"Motor axis 1\", \"flagellum_x\": \"Motor axis 2\"})\n    gt_df = gt_df[fixed_col_gt]\nelse:\n    gt_df = train.copy()\n    gt_df = gt_df.rename(columns={\"flagellum_z\": \"Motor axis 0\", \"flagellum_y\": \"Motor axis 1\", \"flagellum_x\": \"Motor axis 2\"})\n    gt_df = gt_df[fixed_col_gt]\ngt_df.head(10)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cf.GET_OOF:\n    display(final_oof.head(10))","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if cf.GET_OOF:\n    comp_score = score(gt_df, final_oof, 1_000, 2)\n    print(f\"{'+'.join('><')*10}\")\n    print()\n    print(f\"OOF Score is   {comp_score:.6f}\")\n    print()\n    print(f\"{'*'.join('><')*10}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"USE_CURRENT_MODEL=True","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if USE_CURRENT_MODEL:\n    print(\"current model\")\n    model_dir = \"/kaggle/working\"\nelse:\n    print(\"Download model\")\n    model_dir = cf.model_dir","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_model_path(model_dir:str):\n    models = os.listdir(os.path.join(model_dir, \"models\"))\n    model_check_point_list = sorted([os.path.join(model_dir, \"models\", p) for p in models])\n    return model_check_point_list\n\ndef get_models(check_points):\n    model_list = list()\n    for check_point in check_points:\n        model_name = check_point.split(\"/\")[-1].split(\".\")[0]\n        print(f\"Model : {model_name}\")\n        _model = build_model()\n        state  = torch.load(check_point, map_location=device)\n        _model.load_state_dict(state[\"model\"])\n        _model.eval()\n        model_list.append(_model)\n        \n        del state\n        _=gc.collect()\n        torch.cuda.empty_cache()\n        print(\"Successfully!!!\")\n        print()\n\n    return model_list","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"check_points = create_model_path(model_dir)\nlen(check_points)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"check_points","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_list = get_models(check_points)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_coord(output, \n              selected_indices, \n              height, \n              width, \n              threshold=0.6\n             ):\n    if isinstance(output, torch.Tensor):\n        output = output.detach().cpu().numpy()\n            \n    output = output.squeeze(0)\n\n    if np.max(output) < threshold:\n        return -1, -1, -1\n\n    z_pred, y_pred, x_pred = np.unravel_index(np.argmax(output), output.shape)\n    \n    z_orig = selected_indices[z_pred]  \n    y_orig = (y_pred * (height / cf.IMSIZE))\n    x_orig = (x_pred * (width / cf.IMSIZE))\n\n    if round(x_orig) == 0 or round(y_orig) == 0:\n        return -1, -1, -1\n\n    return round(z_orig), round(y_orig), round(x_orig)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def tta_predict(model, image, device, use_center_heatm=cf.use_center_heatm):\n    model.eval()\n    image = image.to(device)\n    preds = []\n\n    # 1. Original\n    preds.append(model(image)[0])\n\n    # 2. Flipped versions\n    for dims in [[2], [3], [2, 3]]:  # Y, X, Y+X flip\n        flipped = torch.flip(image, dims=dims)\n        if use_center_heatm:\n            pred, _, _, _ = model(flipped)\n        else:\n            pred, _, _ = model(flipped)\n            \n        pred = torch.flip(pred, dims=dims)  # reverse flip\n        preds.append(pred)\n\n    # 3. Average predictions\n    avg_pred = torch.stack(preds, dim=0).mean(dim=0)\n    return avg_pred","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def inference(model, dataloaders):\n    \n    preds    = list()\n    tomo_ids = list()\n    fixed_col_pred = [\"tomo_id\", \"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"]\n    stream = tqdm(dataloaders[\"test\"], total=len(dataloaders[\"test\"]), unit=\"test_batch\", desc=\"Infer\")\n    for batch in stream:\n        tomo_id = batch[\"tomo_id\"]\n        img     = batch[\"images\"].to(device)\n        h       = batch[\"h\"].cpu().numpy()\n        w       = batch[\"w\"].cpu().numpy()\n        index   = batch[\"selected_indices\"].detach().cpu().numpy()\n        with torch.no_grad():\n            y_preds = tta_predict(model, img, device)\n        for i in range(y_preds.shape[0]):\n            coord = get_coord(y_preds[i], index[i], h[i], w[i])\n            preds.append(coord)\n            tomo_ids.append(tomo_id[i])\n    pred_df = pd.DataFrame(preds, columns=[\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"])\n    pred_df[\"tomo_id\"] = tomo_ids\n    pred_df = pred_df[fixed_col_pred]\n    return pred_df","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataloaders = get_dataloaders_test(test)\nlen(dataloaders)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_preds = list()\nfor model in model_list:\n    \n    pred_df = inference(model, dataloaders)\n    all_preds.append(pred_df)\nall_preds = pd.concat(all_preds, axis=0)","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"execution_failed":"2025-06-08T00:14:49.820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def aggregate_predictions(ensemble_df):\n    def custom_aggregation(group):\n        z_values = group[\"Motor axis 0\"].values\n        y_values = group[\"Motor axis 1\"].values\n        x_values = group[\"Motor axis 2\"].values\n\n        # -1.0 が含まれている場合は全て -1.0 にする\n        if -1.0 in z_values or -1.0 in y_values or -1.0 in x_values:\n            return pd.Series([-1.0, -1.0, -1.0], index=[\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"])\n\n        # Z は最頻値、Y, X は中央値\n        z_pred = float(mode(z_values)[0])\n        y_pred = np.median(y_values)\n        x_pred = np.median(x_values)\n\n        return pd.Series([z_pred, y_pred, x_pred], index=[\"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"])\n\n    # tomo_id ごとに集約\n    aggregated_df = ensemble_df.groupby(\"tomo_id\").apply(custom_aggregation).reset_index()\n    return aggregated_df\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_pred = aggregate_predictions(all_preds)\n\ndel all_preds\n_=gc.collect()","metadata":{"trusted":true,"_kg_hide-output":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_pred","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"gt_df = test[[\"tomo_id\", \"flagellum_z\", \"flagellum_y\", \"flagellum_x\", \"Voxel spacing\", \"Has motor\"]].copy()\ngt_df.rename(columns={\"flagellum_z\": \"Motor axis 0\", \n                      \"flagellum_y\": \"Motor axis 1\",\n                      \"flagellum_x\": \"Motor axis 2\"\n                     }, inplace=True)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"display(gt_df)\ndisplay(final_pred)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_score = score(gt_df, final_pred, 1_000, 2)\nprint(f\"Final Score : {final_score:.8f}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for b in dataloaders[\"test\"]:\n    img = b[\"images\"].to(device).float()\nimg.shape","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"_=gc.collect()\ntorch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nwith torch.no_grad():\n    if cf.use_center_heatm:\n        out, _, _, _ = model(img)\n    else:\n        out, _, _ = model(img)\nprint(f\"Out Shape: {out.shape}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i in range(out.shape[0]):\n    ou_i = out[i:i+1]\n    break","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ou_i.shape","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Raw output stats:\", out.min().item(), out.max().item())","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"out_orignal = out.squeeze(1)\nout_orig    = out_orignal.detach().cpu().numpy()\nout_orig.shape","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"out0 = out_orig[0]\nout1 = out_orig[1]\nout2 = out_orig[2]","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"id0 = b[\"tomo_id\"][0]\nid1 = b[\"tomo_id\"][1]\nid2 = b[\"tomo_id\"][2]\nid0, id1, id2","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def crop(images, masks):\n    images = torch.tensor(images).unsqueeze(0).float()  # (C, D, H, W)\n    masks  = torch.tensor(masks).unsqueeze(0).float()   \n    crop_H, crop_W, crop_D = cf.CROP_SIZE\n    pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back = compute_padding(images, final_size=crop_H)\n    \n    padded_image = F.pad(images, (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back), mode='reflect')\n    padded_mask  = F.pad(masks,  (pad_left, pad_right, pad_top, pad_bottom, pad_front, pad_back), mode='reflect')\n    images = F.interpolate(padded_image.unsqueeze(0), size=(crop_D, crop_H, crop_W), mode=\"trilinear\", align_corners=False).squeeze(0)\n    masks  = F.interpolate(padded_mask.unsqueeze(0),  size=(crop_D, crop_H, crop_W), mode=\"trilinear\", align_corners=False).squeeze(0)\n    images = images.squeeze(0).detach().cpu().numpy()\n    masks  = masks.squeeze(0).detach().cpu().numpy()\n    return images, masks","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img0, mask0, select_ind0 = load_data(id0)\nimg0, mask0 = crop(img0, mask0)\n\nimg1, mask1, select_ind1 = load_data(id1)\nimg1, mask1 = crop(img1, mask1)\n\nimg2, mask2, select_ind2 = load_data(id2)\nimg2, mask2 = crop(img2, mask2)\n\nlen(select_ind0) == len(select_ind1) == len(select_ind2)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_output(row_id, out, img, mask, steps=3):\n    batch_size = img.shape[0]\n    steps      = steps\n    for i in range(0, batch_size, steps):\n        fig, axes  = plt.subplots(nrows=1, ncols=4, figsize=(14, 6))\n        sig_max = out[i].max()\n        mask_max= mask[i].max()\n        \n        axes[0].imshow(img[i],  cmap=\"viridis\")\n        axes[1].imshow(out[i],  cmap=\"gray\")\n        axes[2].imshow(mask[i], cmap=\"gray\")\n        axes[3].imshow(out[i],  cmap=\"gray\", alpha=0.8)\n        axes[3].imshow(mask[i], cmap=\"gray\", alpha=0.5)\n    \n        axes[0].set_title(f\"Org : Ind: {i} Id:{row_id}\", size=16)\n        axes[1].set_title(f\"Sigm: Max={sig_max:.3f}\", size=16)\n        axes[2].set_title(f\"Mask: Max={mask_max}\",    size=16)\n        axes[3].set_title(f\"Out&Mask\", size=16)\n        for j in range(4):\n            axes[j].set_xticks([])\n            axes[j].set_yticks([])\n        \n        plt.tight_layout()\n        plt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"out1.shape, img1.shape, mask1.shape","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"show_output(id1, out1, img1, mask1)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"show_output(id2, out2, img2, mask2)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"out0_sigmoid = out0\nout1_sigmoid = out1\nout2_sigmoid = out2","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def aggs(output):\n    print(f\"Max : {output.max():.6f}\")\n    print(f\"Min : {output.min():.6f}\")\n    print(f\"Mean: {output.mean():.6f}\")\n    print()\naggs(out0_sigmoid)\naggs(out1_sigmoid)\naggs(out2_sigmoid)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"index = b[\"selected_indices\"].detach().cpu().numpy()\nindex.shape","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"h = b[\"h\"].cpu().numpy()\nw = b[\"w\"].cpu().numpy()\n\nh0 = h[0]\nh1 = h[1]\nh2 = h[2]\n\nw0 = w[0]\nw1 = w[1]\nw2 = w[2]","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"index0 = index[0]\nindex1 = index[1]\nindex2 = index[2]","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"output           = out1_sigmoid\nselected_indices = index1\nheight           = h1\nwidth            = w1","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_slices(volume, step=4):\n    num_slices = volume.shape[0]\n    fig, axes = plt.subplots(1, num_slices // step, figsize=(15, 5))\n    for i, ax in enumerate(axes):\n        idx = i * step\n        ax.imshow(volume[idx], cmap=\"hot\")\n        ax.set_facecolor(\"black\")\n        ax.set_title(f\"Slice {idx}\")\n        ax.axis(\"off\")\n    plt.tight_layout()\n    plt.show()\n\n\nshow_slices(out0_sigmoid)\nshow_slices(out1_sigmoid)\nshow_slices(out2_sigmoid)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-06-08T00:14:49.823Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}