{"cells":[{"cell_type":"markdown","id":"64c65549","metadata":{},"source":"# exp_20260202_107_transunet3d_baseline debug smoke\n- TPU v5e-8 前提 / Keras Core + JAX バックエンド\n- 2ステップだけ流して JIT/データ周りが動くか確認する最小スモーク\n- 混合精度は bfloat16 固定（v5e は float16 の一部演算が未サポート）"},{"cell_type":"code","execution_count":null,"id":"3abac0c3","metadata":{},"outputs":[],"source":"import os, json, subprocess\nfrom pathlib import Path\n\nos.environ.setdefault(\"KERAS_BACKEND\", \"jax\")\nos.environ.setdefault(\"TF_CPP_MIN_LOG_LEVEL\", \"2\")\n\ndef _pip(args):\n    print(\"[pip]\", \" \".join(args))\n    subprocess.run([\"pip\", \"install\", \"-q\"] + args, check=True)\n\n# オフライン wheel 探索（どちらか存在する方を使う）\nWHEEL_DIR_CANDS = [\n    Path(\"/kaggle/input/vsdetection-packages-offline-installer-only/whls\"),\n    Path(\"/kaggle/input/vesuvius-nnunet3d-offline-wheels-v1\"),\n]\nWHEEL_DIR = next((p for p in WHEEL_DIR_CANDS if p.exists()), None)\nALLOW_NET_INSTALL = os.environ.get(\"ALLOW_NET_INSTALL\", \"0\") == \"1\"\n\ndef _ensure_tifffile():\n    try:\n        import tifffile  # noqa\n        return True\n    except Exception:\n        return False\n\nif WHEEL_DIR is not None and list(WHEEL_DIR.glob(\"*.whl\")):\n    _pip([str(w) for w in sorted(WHEEL_DIR.glob(\"*.whl\"))])\nelse:\n    # TPU 環境では imagecodecs のビルドに失敗しがちなので tifffile だけ入れて fallback\n    if not _ensure_tifffile() and ALLOW_NET_INSTALL:\n        try:\n            _pip([\"keras>=3.3.3\", \"tensorflow>=2.15.0\", \"tifffile\"])\n        except subprocess.CalledProcessError as exc:\n            print(f\"[pip] skip network install (offline?): {exc}\")\n    elif not _ensure_tifffile():\n        print(\"[pip] skip network install (ALLOW_NET_INSTALL=0); assume tifffile bundled in image\")\n\n# JAX TPU\nif os.environ.get(\"KAGGLE_ACCELERATOR\", \"\").lower().startswith(\"tpu\") or os.environ.get(\"COLAB_TPU_ADDR\"):\n    _pip([\"jax[tpu]>=0.4.23\", \"-f\", \"https://storage.googleapis.com/jax-releases/libtpu_releases.html\"])\n\ntry:\n    import medicai  # type: ignore\nexcept Exception:\n    _pip([\"git+https://github.com/innat/medic-ai.git\"])\n    import medicai  # type: ignore\n\ntry:\n    import skimage  # type: ignore\nexcept Exception:\n    _pip([\"scikit-image\"])"},{"cell_type":"code","execution_count":null,"id":"a817ba33","metadata":{},"outputs":[],"source":"import math\nimport warnings\nimport numpy as np\nimport tensorflow as tf\nimport keras\nfrom keras import ops, 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\")\nkeras.mixed_precision.set_global_policy(\"mixed_bfloat16\")\ntry:\n    keras.config.enable_flash_attention(); keras.config.disable_flash_attention()\nexcept TypeError:\n    pass"},{"cell_type":"markdown","id":"8b9d809e","metadata":{},"source":"## 定数（小さく固定）\n- shard: train=[0], val=[1]\n- steps_per_epoch=2 で即終了"},{"cell_type":"code","execution_count":null,"id":"f2b851ea","metadata":{},"outputs":[],"source":"DATA_ROOT = Path(\"/kaggle/input/vesuvius-tfrecords\")\nTRAIN_SHARDS = [0]\nVAL_SHARDS = [1]\nINPUT_SHAPE = (160, 160, 160)\nNUM_CLASSES = 3\nEPOCHS = 1\nSTEPS_PER_EPOCH = 2\nBATCH_SIZE_PER_DEVICE = 1\nSTEP_LOG_EVERY = 1  # 1ステップごとにログ"},{"cell_type":"code","execution_count":null,"id":"770d1217","metadata":{},"outputs":[],"source":"devices = keras.distribution.list_devices()\nif not devices:\n    raise SystemError(\"デバイスが見つかりません（TPU/GPU）。\")\nkeras.distribution.set_distribution(keras.distribution.DataParallel(devices=devices))\ntotal_devices = len(devices)\nbatch_size = BATCH_SIZE_PER_DEVICE * total_devices\nprint(\"devices:\", devices)\nprint(\"global batch size:\", batch_size)"},{"cell_type":"markdown","id":"9f160e88","metadata":{},"source":"## tf.data ローダー（最小構成）"},{"cell_type":"code","execution_count":null,"id":"80e65c4a","metadata":{},"outputs":[],"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 = tf.reshape(image, tf.cast(parsed[\"image_shape\"], tf.int64))\n    label = tf.reshape(label, tf.cast(parsed[\"label_shape\"], tf.int64))\n    image = tf.cast(image[..., None], tf.float32)\n    label = tf.cast(label[..., None], tf.float32)\n    return image, label\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            ScaleIntensityRange(keys=[\"image\"], a_min=0, a_max=255, b_min=0, b_max=1, clip=True),\n        ]\n    )\n    out = pipe(data)\n    return out[\"image\"], out[\"label\"]\n\ndef _add_skel(image, label):\n    tubed = tf.numpy_function(\n        lambda lbl: binary_dilation(skeletonize((lbl[..., 0] == 1)), iterations=1).astype(np.float32)[..., None],\n        [label],\n        tf.float32,\n    )\n    tubed.set_shape(label.shape)\n    return image, tf.concat([tf.cast(label, tf.float32), tubed], axis=-1)\n\ndef build_loader(shards, shuffle=True, batch_size=batch_size):\n    paths = [str(DATA_ROOT / f\"training_shard_{i}.tfrec\") for i in shards]\n    missing = [p for p in paths if not tf.io.gfile.exists(p)]\n    if missing:\n        raise FileNotFoundError(f\"TFRecord missing: {missing}\")\n    ds = tf.data.TFRecordDataset(paths)\n    if shuffle:\n        ds = ds.shuffle(64)\n    ds = ds.map(_parse, num_parallel_calls=tf.data.AUTOTUNE)\n    ds = ds.map(_train_transform if shuffle else _train_transform, num_parallel_calls=tf.data.AUTOTUNE)\n    if shuffle:\n        ds = ds.map(_add_skel, num_parallel_calls=tf.data.AUTOTUNE)\n    ds = ds.batch(batch_size, drop_remainder=True).prefetch(tf.data.AUTOTUNE)\n    return ds\n\ntrain_loader = build_loader(TRAIN_SHARDS, shuffle=True)\nval_loader = build_loader(VAL_SHARDS, shuffle=False)"},{"cell_type":"markdown","id":"4cac5a40","metadata":{},"source":"## データローダ単体の速度チェック\n4バッチだけ取り出して時間を計測する"},{"cell_type":"code","execution_count":null,"id":"9be62e01","metadata":{"lines_to_next_cell":1},"outputs":[],"source":"import time\nt0 = time.time()\nt1 = None\nfor i, b in enumerate(train_loader.take(4)):\n    if i == 0:\n        t1 = time.time()\n    if i == 3:\n        break\nif t1 is None:\n    print(\"data: no batch fetched (check tfrec paths)\")\nelse:\n    print(\"data: first batch sec:\", t1 - t0, \"four batches sec:\", time.time() - t0)"},{"cell_type":"markdown","id":"53ff4851","metadata":{},"source":"## モデルとロス（JAX-safe）"},{"cell_type":"code","execution_count":null,"id":"9caa0a56","metadata":{},"outputs":[],"source":"class SkeletonRecallPlusDiceLoss(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(from_logits=False, num_classes=num_classes, ignore_class_ids=2)\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 ops.zeros_like(y_true_mask)\n        pred_ink = y_pred[..., 1]\n        valid = ops.cast(ops.not_equal(y_true_mask, 2), \"float32\")\n\n        base = self.base_loss(y_true_mask[..., None], y_pred)\n        inter = ops.sum(pred_ink * y_true_skel * valid, axis=(1, 2, 3))\n        skel_sum = ops.sum(y_true_skel * valid, axis=(1, 2, 3))\n        has_skel = ops.cast(ops.greater(skel_sum, 0), \"float32\")\n        skel_loss = ops.mean((1.0 - (inter + 1e-6) / (skel_sum + 1e-6)) * has_skel)\n        gt_bg = ops.cast(ops.equal(y_true_mask, 0), \"float32\") * valid\n        fp_loss = ops.sum(pred_ink * gt_bg) / (ops.sum(gt_bg) + 1e-6)\n        return base + self.w_srec * skel_loss + self.w_fp * fp_loss\n\nclass MaskOnlySparseDiceMetric(SparseDiceMetric):\n    def update_state(self, y_true, y_pred, sample_weight=None):\n        shape = getattr(y_true, \"shape\", None)\n        if shape is not None:\n            rank = len(shape)\n            last = shape[-1] if rank > 0 else None\n            if rank >= 5 and last is not None and last >= 1:\n                y_true = y_true[..., 0:1]\n        return super().update_state(y_true, y_pred, sample_weight)\n\n# 簡易ステップロガー\nclass StepLogger(keras.callbacks.Callback):\n    def __init__(self, every=10):\n        super().__init__()\n        self.every = every\n\n    def on_train_batch_end(self, batch, logs=None):\n        if (batch + 1) % self.every == 0:\n            loss = logs.get(\"loss\")\n            dice = logs.get(\"dice\")\n            dice_val = float(dice) if dice is not None else 0.0\n            print(f\"[train] step {batch+1} loss={loss:.4f} dice={dice_val:.4f}\")\n\nmodel = TransUNet(\n    input_shape=INPUT_SHAPE + (1,),\n    encoder_name=\"seresnext50\",\n    classifier_activation=\"softmax\",\n    num_classes=NUM_CLASSES,\n)\n\nsteps_per_epoch = STEPS_PER_EPOCH\nlr_schedule = optimizers.schedules.CosineDecay(\n    initial_learning_rate=5e-5,\n    decay_steps=steps_per_epoch * EPOCHS,\n    alpha=0.1,\n)\noptimizer = optimizers.AdamW(learning_rate=lr_schedule, weight_decay=1e-5)\nloss_fn = SkeletonRecallPlusDiceLoss(num_classes=NUM_CLASSES)\nmetrics = [MaskOnlySparseDiceMetric(from_logits=False, num_classes=NUM_CLASSES, ignore_class_ids=2, name=\"dice\")]\nmodel.compile(optimizer=optimizer, loss=loss_fn, metrics=metrics)\n\nprint(\"model ready\")"},{"cell_type":"markdown","id":"60c5d812","metadata":{},"source":"## ダミー forward（JIT ウォームアップ）"},{"cell_type":"code","execution_count":null,"id":"15f96b13","metadata":{},"outputs":[],"source":"batch = next(iter(train_loader))\nout = model(batch[0])\nprint(\"forward ok, out shape:\", out.shape)"},{"cell_type":"markdown","id":"846dc45f","metadata":{},"source":"## トレーニング開始前の案内（初回コンパイル待ち用）"},{"cell_type":"code","execution_count":null,"id":"c6a8d79e","metadata":{},"outputs":[],"source":"print(\"starting fit ... first step may take 60-180s on TPU v5e (XLA compile)\")"},{"cell_type":"markdown","id":"6fdc095c","metadata":{},"source":"## 超短縮 fit（2 step × 1 epoch）"},{"cell_type":"code","execution_count":null,"id":"0462e73c","metadata":{},"outputs":[],"source":"history = model.fit(\n    train_loader,\n    epochs=EPOCHS,\n    steps_per_epoch=STEPS_PER_EPOCH,\n    callbacks=[StepLogger(every=STEP_LOG_EVERY)],\n    verbose=1,\n)\nprint(\"fit finished\")"}],"metadata":{"jupytext":{"formats":"py:percent,ipynb"},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"}},"nbformat":4,"nbformat_minor":5}