{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":""},"jupytext":{"formats":"py:percent,ipynb"},"kaggle":{"accelerator":"tpuV5e8","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"},{"sourceId":14266465,"sourceType":"datasetVersion","datasetId":8751895}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"465a4b08","cell_type":"markdown","source":"# exp_20260202_107_transunet3d_baseline\nnnUNet が不調なときの「まず動く」3D TransUNet パイプライン。  \n`.prompts/reference/train-transunet-baseline-lb-0-537.ipynb` を簡約し、TPU なしの Kaggle GPU で回る構成にまとめた。  \n- データ: `/kaggle/input/vesuvius-tfrecords/training_shard_*.tfrec` を想定（0-129をtrain, 130をval）  \n- 依存: `medicai`（TransUNet + 3D SWI）、必要ならオフライン wheels で Keras/JAX を差し替え  \n- 出力: `/kaggle/working/model.weights.h5` と `/kaggle/working/history.json`","metadata":{}},{"id":"61088ffe","cell_type":"markdown","source":"## 0. 環境セットアップ\n- GPU前提。TPUは使わない（Keras JAX のままでも可）。  \n- オフライン wheel がある場合だけ拾い、なければPyPIから最小限を入れる。","metadata":{}},{"id":"2585084b","cell_type":"code","source":"import json\nimport os\nimport subprocess\nfrom pathlib import Path\n\nos.environ.setdefault(\"KERAS_BACKEND\", \"tensorflow\")  # TPU/CPUでJAXを使うなら \"jax\" に差し替え\nos.environ.setdefault(\"TF_CPP_MIN_LOG_LEVEL\", \"2\")\n\ndef _pip_install(args):\n    print(f\"[pip] {' '.join(args)}\")\n    subprocess.run([\"pip\", \"install\", \"-q\"] + args, check=True)\n\n# オフライン wheel (任意)\nWHEEL_DIR = Path(\"/kaggle/input/vsdetection-packages-offline-installer-only/whls\")\nif WHEEL_DIR.exists():\n    wheels = [\n        \"keras_nightly-3.12.0.dev2025100703-py3-none-any.whl\",\n        \"tifffile-2025.12.12-py3-none-any.whl\",\n        \"imagecodecs-2025.11.11-cp311-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\",\n    ]\n    _pip_install([str(WHEEL_DIR / w) for w in wheels])\nelse:\n    _pip_install([\"keras>=3.3.3\", \"tensorflow>=2.15.0\", \"imagecodecs\", \"tifffile\"])\n\ntry:\n    import medicai  # type: ignore\nexcept Exception:\n    _pip_install([\"git+https://github.com/innat/medic-ai.git\"])\n    import medicai  # type: ignore\n\n# scikit-image は skeletonize で必須\ntry:\n    import skimage  # type: ignore\nexcept Exception:\n    _pip_install([\"scikit-image\"])","metadata":{},"outputs":[],"execution_count":null},{"id":"f1d079d5","cell_type":"markdown","source":"## 1. 依存と定数","metadata":{}},{"id":"0a129a3e","cell_type":"code","source":"import math\nimport warnings\n\nimport numpy as np\nimport tensorflow as tf\nfrom keras import ops\nfrom keras import optimizers\nfrom medicai.metrics import SparseDiceMetric\nfrom medicai.losses import SparseDiceCELoss\nfrom medicai.models import TransUNet\nfrom medicai.transforms import (\n    Compose,\n    RandFlip,\n    RandRotate90,\n    RandShiftIntensity,\n    RandSpatialCrop,\n    ScaleIntensityRange,\n)\nfrom skimage.morphology import skeletonize\nfrom scipy.ndimage import binary_dilation\n\nwarnings.filterwarnings(\"ignore\")\ntf.keras.mixed_precision.set_global_policy(\"mixed_float16\")","metadata":{},"outputs":[],"execution_count":null},{"id":"40cb2598","cell_type":"markdown","source":"### パスとハイパーパラメータ","metadata":{}},{"id":"c0d8f85e","cell_type":"code","source":"DATA_ROOT = Path(\"/kaggle/input/vesuvius-tfrecords\")\nTRAIN_SHARDS = list(range(0, 130))   # 0-129\nVAL_SHARDS = [130]                   # 130\n\nINPUT_SHAPE = (160, 160, 160)        # D, H, W\nNUM_CLASSES = 3                      # {0: bg, 1: ink, 2: ignore}\nEPOCHS = 60                          # 100だと長いのでまず60で一周\nWARMUP_EPOCHS = 3\nBATCH_SIZE_PER_GPU = 1\nLR = 5e-5\nWEIGHT_DECAY = 1e-5\nOCCLUSION_MAX_BLOCKS = 6\n\nOUT_WEIGHTS = Path(\"/kaggle/working/model.weights.h5\")\nHISTORY_JSON = Path(\"/kaggle/working/history.json\")","metadata":{},"outputs":[],"execution_count":null},{"id":"d1e432a4","cell_type":"markdown","source":"## 2. デバイス & 分散設定","metadata":{}},{"id":"4a4d5775","cell_type":"code","source":"try:\n    resolver = tf.distribute.cluster_resolver.TPUClusterResolver(tpu=\"local\")  # Kaggle/Colab TPU は local でOK\n    tf.config.experimental_connect_to_cluster(resolver)\n    tf.tpu.experimental.initialize_tpu_system(resolver)\n    strategy = tf.distribute.TPUStrategy(resolver)\n    total_devices = strategy.num_replicas_in_sync\n    batch_size = BATCH_SIZE_PER_GPU * total_devices  # per-replica を変えたければ BATCH_SIZE_PER_GPU を調整\n    print(\"detected devices (TPU):\", resolver.cluster_spec().as_dict().get(\"TPU\", []))\n    print(f\"global batch size: {batch_size}\")\nexcept Exception as exc:\n    print(f\"TPU unavailable ({exc}); fallback to GPU/CPU\")\n    devices = tf.config.list_logical_devices(\"GPU\")\n    if devices:\n        strategy = tf.distribute.MirroredStrategy(devices=devices)\n        total_devices = strategy.num_replicas_in_sync\n        batch_size = BATCH_SIZE_PER_GPU * total_devices\n        print(f\"detected GPU devices: {devices}\")\n    else:\n        strategy = tf.distribute.OneDeviceStrategy(device=\"/CPU:0\")\n        total_devices = 1\n        batch_size = BATCH_SIZE_PER_GPU * total_devices\n        print(\"no GPU; using CPU\")\n    print(f\"global batch size: {batch_size}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"b2dbe54f","cell_type":"markdown","source":"## 3. TFRecord ローダー\n- 画像とラベルを uint8 → float32 にデコード  \n- train は RandCrop/Flip/Rotate + occlusion + skeleton channel  \n- val は正規化のみ（skeleton 無し）","metadata":{}},{"id":"ecdf6c65","cell_type":"code","source":"FEATURE_DESC = {\n    \"image\": tf.io.FixedLenFeature([], tf.string),\n    \"label\": tf.io.FixedLenFeature([], tf.string),\n    \"image_shape\": tf.io.FixedLenFeature([3], tf.int64),\n    \"label_shape\": tf.io.FixedLenFeature([3], tf.int64),\n}\n\ndef _parse(example):\n    parsed = tf.io.parse_single_example(example, FEATURE_DESC)\n    image = tf.io.decode_raw(parsed[\"image\"], tf.uint8)\n    label = tf.io.decode_raw(parsed[\"label\"], tf.uint8)\n    image_shape = tf.cast(parsed[\"image_shape\"], tf.int64)\n    label_shape = tf.cast(parsed[\"label_shape\"], tf.int64)\n    image = tf.reshape(image, image_shape)\n    label = tf.reshape(label, label_shape)\n    image = tf.cast(image[..., None], tf.float32)\n    label = tf.cast(label[..., None], tf.float32)\n    return image, label\n\n\ndef random_occlusions(volume, max_blocks=OCCLUSION_MAX_BLOCKS, min_size=2, max_size=8, prob=1.0):\n    def _augment(v):\n        shape = tf.shape(v)\n        D, H, W = shape[0], shape[1], shape[2]\n        d_range = tf.range(D)[:, None, None, None]\n        h_range = tf.range(H)[None, :, None, None]\n        w_range = tf.range(W)[None, None, :, None]\n        mask = tf.ones([D, H, W, 1], v.dtype)\n        for _ in tf.range(max_blocks):\n            bd = tf.random.uniform([], min_size, max_size, dtype=tf.int32)\n            bh = tf.random.uniform([], min_size, max_size, dtype=tf.int32)\n            bw = tf.random.uniform([], min_size, max_size, dtype=tf.int32)\n            d0 = tf.random.uniform([], 0, tf.maximum(D - bd, 1), dtype=tf.int32)\n            h0 = tf.random.uniform([], 0, tf.maximum(H - bh, 1), dtype=tf.int32)\n            w0 = tf.random.uniform([], 0, tf.maximum(W - bw, 1), dtype=tf.int32)\n            d1 = tf.minimum(d0 + bd, D)\n            h1 = tf.minimum(h0 + bh, H)\n            w1 = tf.minimum(w0 + bw, W)\n            block = (\n                (d_range >= d0) & (d_range < d1) &\n                (h_range >= h0) & (h_range < h1) &\n                (w_range >= w0) & (w_range < w1)\n            )\n            mask = mask * (1.0 - tf.cast(block, v.dtype))\n        return v * mask\n\n    return tf.cond(tf.random.uniform([]) < prob, lambda: _augment(volume), lambda: volume)\n\n\ndef _train_transform(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipe = Compose(\n        [\n            RandSpatialCrop(keys=[\"image\", \"label\"], roi_size=INPUT_SHAPE),\n            RandFlip(keys=[\"image\", \"label\"], spatial_axis=[0], prob=0.5),\n            RandFlip(keys=[\"image\", \"label\"], spatial_axis=[1], prob=0.5),\n            RandFlip(keys=[\"image\", \"label\"], spatial_axis=[2], prob=0.5),\n            RandRotate90(keys=[\"image\", \"label\"], prob=0.4, max_k=3, spatial_axes=(0, 1)),\n            RandShiftIntensity(keys=[\"image\"], offsets=0.15, prob=0.5),\n            ScaleIntensityRange(\n                keys=[\"image\"], a_min=0, a_max=255, b_min=0, b_max=1, clip=True\n            ),\n        ]\n    )\n    out = pipe(data)\n    out[\"image\"] = random_occlusions(out[\"image\"])\n    return out[\"image\"], out[\"label\"]\n\n\ndef _val_transform(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipe = Compose(\n        [\n            ScaleIntensityRange(\n                keys=[\"image\"], a_min=0, a_max=255, b_min=0, b_max=1, clip=True\n            ),\n        ]\n    )\n    out = pipe(data)\n    return out[\"image\"], out[\"label\"]\n\n\ndef _generate_tubed_skeleton(label_vol):\n    mask = (label_vol[..., 0] == 1)\n    skel = skeletonize(mask)\n    tubed = binary_dilation(skel, iterations=1)\n    return tubed.astype(np.float32)[..., None]\n\n\ndef _add_skeleton(image, label):\n    tubed = tf.numpy_function(_generate_tubed_skeleton, [label], tf.float32)\n    tubed.set_shape(label.shape)\n    label2 = tf.concat([tf.cast(label, tf.float32), tubed], axis=-1)\n    return image, label2\n\n\ndef build_loader(shards, shuffle=True, batch_size=batch_size):\n    pattern = [str(DATA_ROOT / f\"training_shard_{i}.tfrec\") for i in shards]\n    for p in pattern:\n        if not Path(p).exists():\n            raise FileNotFoundError(f\"missing tfrecord: {p}\")\n    ds = tf.data.TFRecordDataset(pattern)\n    if shuffle:\n        ds = ds.shuffle(buffer_size=256, reshuffle_each_iteration=True)\n    ds = ds.map(_parse, num_parallel_calls=tf.data.AUTOTUNE)\n    ds = ds.map(_train_transform if shuffle else _val_transform, num_parallel_calls=tf.data.AUTOTUNE)\n    if shuffle:\n        ds = ds.map(_add_skeleton, num_parallel_calls=tf.data.AUTOTUNE)\n    ds = ds.batch(batch_size, drop_remainder=True)\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n    return ds\n\n\ntrain_loader = build_loader(TRAIN_SHARDS, shuffle=True)\nval_loader = build_loader(VAL_SHARDS, shuffle=False, batch_size=1 * total_devices)","metadata":{"lines_to_next_cell":1},"outputs":[],"execution_count":null},{"id":"e8685b24","cell_type":"markdown","source":"## 4. モデル & ロス\n- TransUNet (seresnext50 encoder) をそのまま利用  \n- ロスは Dice+CE に Skeleton Recall + FP 抑制を加えたもの","metadata":{}},{"id":"e802daf0","cell_type":"code","source":"class SkeletonRecallPlusDiceLoss(tf.keras.losses.Loss):\n    def __init__(self, num_classes, w_srec=0.5, w_fp=0.3, name=\"skel_recall_fp_loss\"):\n        super().__init__(name=name)\n        self.num_classes = num_classes\n        self.w_srec = w_srec\n        self.w_fp = w_fp\n        self.base_loss = SparseDiceCELoss(\n            from_logits=False, num_classes=num_classes, ignore_class_ids=2\n        )\n\n    def call(self, y_true, y_pred):\n        y_true_mask = y_true[..., 0]\n        y_true_skel = y_true[..., 1] if y_true.shape[-1] > 1 else tf.zeros_like(y_true_mask)\n        pred_ink = y_pred[..., 1]\n        valid = tf.cast(y_true_mask != 2, tf.float32)\n\n        base = self.base_loss(y_true_mask[..., None], y_pred)\n\n        inter = tf.reduce_sum(pred_ink * y_true_skel * valid, axis=(1, 2, 3))\n        skel_sum = tf.reduce_sum(y_true_skel * valid, axis=(1, 2, 3))\n        has_skel = tf.cast(skel_sum > 0, tf.float32)\n        skel_loss = tf.reduce_mean((1.0 - (inter + 1e-6) / (skel_sum + 1e-6)) * has_skel)\n\n        gt_bg = tf.cast(y_true_mask == 0, tf.float32) * valid\n        fp_loss = tf.reduce_sum(pred_ink * gt_bg) / (tf.reduce_sum(gt_bg) + 1e-6)\n\n        return base + self.w_srec * skel_loss + self.w_fp * fp_loss\n\n\nclass MaskOnlySparseDiceMetric(SparseDiceMetric):\n    \"\"\"y_trueが (B,D,H,W,2) のときは channel0 だけを使うラッパー.\"\"\"\n\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        if y_true.shape.rank is not None and y_true.shape.rank >= 5 and y_true.shape[-1] >= 1:\n            y_true = y_true[..., 0:1]\n        return super().update_state(y_true, y_pred, sample_weight)\n\n\nwith strategy.scope():\n    model = TransUNet(\n        input_shape=INPUT_SHAPE + (1,),\n        encoder_name=\"seresnext50\",\n        classifier_activation=\"softmax\",\n        num_classes=NUM_CLASSES,\n    )\n\n    steps_per_epoch = math.floor(780 / batch_size)  # 780 sample count from reference DS\n    total_steps = steps_per_epoch * EPOCHS\n    warmup_steps = steps_per_epoch * WARMUP_EPOCHS\n\n    lr_schedule = optimizers.schedules.CosineDecay(\n        initial_learning_rate=LR,\n        decay_steps=total_steps,\n        alpha=0.1,\n    )\n\n    optimizer = optimizers.AdamW(learning_rate=lr_schedule, weight_decay=WEIGHT_DECAY)\n    loss_fn = SkeletonRecallPlusDiceLoss(num_classes=NUM_CLASSES)\n    metrics = [\n        MaskOnlySparseDiceMetric(\n            from_logits=False, num_classes=NUM_CLASSES, ignore_class_ids=2, name=\"dice\"\n        )\n    ]\n    model.compile(optimizer=optimizer, loss=loss_fn, metrics=metrics)\n\nprint(model.summary())","metadata":{"lines_to_next_cell":1},"outputs":[],"execution_count":null},{"id":"cb18f73e","cell_type":"markdown","source":"## 5. コールバック（簡易val計測 + スナップショット）","metadata":{}},{"id":"9481be44","cell_type":"code","source":"class ValDiceCallback(tf.keras.callbacks.Callback):\n    \"\"\"10epochごとにval_loader全体でDiceを計測する簡易コールバック.\"\"\"\n\n    def __init__(self, interval=10):\n        super().__init__()\n        self.interval = interval\n        self.metric = SparseDiceMetric(\n            from_logits=False, num_classes=NUM_CLASSES, ignore_class_ids=2, name=\"val_dice\"\n        )\n\n    def on_epoch_end(self, epoch, logs=None):\n        if (epoch + 1) % self.interval != 0:\n            return\n        self.metric.reset_state()\n        for batch_x, batch_y in val_loader:\n            preds = self.model(batch_x, training=False)\n            # val_loaderは skeleton チャンネルなしなので y[...,0:1]\n            self.metric.update_state(batch_y[..., 0:1], preds)\n        val_dice = ops.convert_to_numpy(self.metric.result())\n        print(f\"[val] epoch={epoch+1} dice={val_dice:.4f}\")\n        if logs is not None:\n            logs[\"val_dice\"] = val_dice\n        # ベスト更新時のみ保存\n        if not hasattr(self, \"best\") or val_dice > getattr(self, \"best\", -1):\n            self.best = val_dice\n            self.model.save_weights(OUT_WEIGHTS)\n            print(f\"[val] improved -> saved {OUT_WEIGHTS}\")\n\nclass PeriodicSaver(tf.keras.callbacks.Callback):\n    def __init__(self, interval=20, template=\"checkpoint_epoch_{epoch}.weights.h5\"):\n        super().__init__()\n        self.interval = interval\n        self.template = template\n\n    def on_epoch_end(self, epoch, logs=None):\n        if (epoch + 1) % self.interval == 0:\n            path = self.template.format(epoch=epoch + 1)\n            self.model.save_weights(path)\n            print(f\"[snapshot] saved {path}\")","metadata":{"lines_to_next_cell":1},"outputs":[],"execution_count":null},{"id":"90c595fa","cell_type":"markdown","source":"## 6. 学習","metadata":{}},{"id":"00cedf9f","cell_type":"code","source":"history = model.fit(\n    train_loader,\n    epochs=EPOCHS,\n    callbacks=[ValDiceCallback(), PeriodicSaver()],\n    verbose=1,\n)\n\nwith open(HISTORY_JSON, \"w\") as f:\n    json.dump(history.history, f, indent=2)\nmodel.save_weights(OUT_WEIGHTS)\nprint(f\"saved: {OUT_WEIGHTS}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"7190568b","cell_type":"markdown","source":"## 7. 検証用の簡易推論\n最初の val バッチで Dice を確認する。","metadata":{}},{"id":"9325fcc0","cell_type":"code","source":"val_ds = iter(val_loader)\nval_x, val_y = next(val_ds)\nval_pred = model(val_x, training=False)\n\nval_dice_metric = SparseDiceMetric(\n    from_logits=False, num_classes=NUM_CLASSES, ignore_class_ids=2\n)\nval_dice_metric.update_state(val_y[..., 0:1], val_pred)\nval_dice = ops.convert_to_numpy(val_dice_metric.result())\nprint(f\"quick val dice (1 batch): {val_dice:.4f}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"1a63050a","cell_type":"markdown","source":"## 8. 使い方メモ\n- `jupytext --to ipynb notebook.py` でipynb生成して Kaggle に push  \n- データセット slug が異なる場合は `DATA_ROOT` と `kernel-metadata.json` を更新  \n- まず EPOCHS=20 など短縮でスモーク → 60 / 100 に伸ばすのがおすすめ","metadata":{}}]}