{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpuV5e8","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":14266465,"datasetId":8751895,"databundleVersionId":15066994},{"sourceType":"modelInstanceVersion","sourceId":732880,"databundleVersionId":15477237,"modelInstanceId":516822,"modelId":510647,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":290917305,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":298459500,"isSourceIdPinned":false}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from IPython.display import clear_output\n\n!pip install tensorflow -qU\n\nvar=\"/kaggle/input/vsdetection-packages-offline-installer-only/whls\"\n!pip install \\\n    \"$var\"/keras_nightly-3.12.0.dev2025100703-py3-none-any.whl \\\n    --no-index \\\n    --find-links \"$var\"\n\nclear_output()\n\n!pip install git+https://github.com/innat/medic-ai.git -q\n\nimport os, warnings\nimport numpy as np\n\n# ========== 关键：JAX/TPU 内存优化配置 ==========\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\nos.environ[\"JAX_ARRAY\"] = \"cpu\"\nos.environ[\"TPU_METRICS_SERVER_PORT\"] = \"0\"\n# 内存优化关键配置\nos.environ[\"XLA_PYTHON_CLIENT_MEM_FRACTION\"] = \"0.95\"\nos.environ[\"XLA_PYTHON_CLIENT_PREALLOCATE\"] = \"false\"\nos.environ[\"JAX_DISABLE_JIT\"] = \"false\"\nos.environ[\"JAX_PLATFORM_NAME\"] = \"tpu\"\nwarnings.filterwarnings('ignore')\n\nimport glob\nimport pandas as pd\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\n# mainly for training API\nimport keras\nfrom keras import ops\nfrom keras.optimizers import SGD, AdamW, Muon\nfrom keras.optimizers.schedules import CosineDecay, PolynomialDecay\n\n# only for tf.data API\nimport tensorflow as tf\n\n# mainly for 3D or 2D models, transformation, loss, metrics etc\nimport medicai\nfrom medicai.transforms import (\n    Compose,\n    NormalizeIntensity,\n    ScaleIntensityRange,\n    Resize,\n    RandShiftIntensity,\n    RandRotate90,\n    RandRotate,\n    RandFlip,\n    RandCutOut,\n    RandSpatialCrop\n)\nfrom medicai.layers import ResizingND\nfrom medicai.models import (\n    UNet, SegFormer, TransUNet, SwinUNETR, UPerNet, ConvNeXtV2Tiny, UNETRPlusPlus\n)\nfrom medicai.losses import (\n    SparseDiceCELoss, SparseTverskyLoss, SparseCenterlineDiceLoss\n)\nfrom medicai.metrics import SparseDiceMetric\nfrom medicai.utils import SlidingWindowInference\nfrom medicai.utils import soft_skeletonize\n\n# due to distributed training only\nkeras.config.disable_flash_attention()\n\n# reproducibility\nkeras.utils.set_random_seed(101)\n\n# ========== 适配版：8核TPU分布式配置（兼容无get_distribution的版本） ==========\ndevices = keras.distribution.list_devices()\nuse_devices = devices[:8] if len(devices)>=8 else devices\n# 直接创建分布式策略，不保存原策略（适配你的Keras版本）\ntrain_distribution = keras.distribution.DataParallel(devices=use_devices)\n# 适配：有些版本用experimental_set_distribution\ntry:\n    keras.distribution.set_distribution(train_distribution)\nexcept AttributeError:\n    keras.distribution.experimental_set_distribution(train_distribution)\ntotal_device = len(use_devices)\n\nprint(f'detected devices: {devices}')\nprint(f'used devices: {use_devices}')\nprint(f'total device: {total_device}')\n\nkeras.version(), keras.config.backend(), medicai.version()\n\n# ========== 核心参数：保持8核训练配置 ==========\ninput_shape=(160, 160, 160)\nbatch_size=1 * total_device  # 8核 → batch_size=8（训练用）\ninfer_batch_size = 1  # 推理用1个样本\nnum_classes=3\nnum_samples = 780\nstart_epoch = 50\ntotal_epochs = 250\ncontinue_epochs = total_epochs - start_epoch\nprint(f\"续训配置：从{start_epoch}到{total_epochs}，共续训{continue_epochs}epoch\")\n\n# ========== 数据加载函数（保持不变） ==========\ndef parse_tfrecord_fn(example):\n    feature_description = {\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    parsed_example = tf.io.parse_single_example(example, feature_description)\n    image = tf.io.decode_raw(parsed_example[\"image\"], tf.uint8)\n    label = tf.io.decode_raw(parsed_example[\"label\"], tf.uint8)\n    image_shape = tf.cast(parsed_example[\"image_shape\"], tf.int64)\n    label_shape = tf.cast(parsed_example[\"label_shape\"], tf.int64)\n    image = tf.reshape(image, image_shape)\n    label = tf.reshape(label, label_shape)\n    return image, label\n\ndef prepare_inputs(image, label):\n    # Add channel dimension\n    image = image[..., None] # (D, H, W, 1)\n    label = label[..., None] # (D, H, W, 1)\n\n    # Convert to float32\n    image = tf.cast(image, tf.float32)\n    label = tf.cast(label, tf.float32)\n    return image, label\n\ndef train_transformation(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\n        ## Geometric transformation\n        RandSpatialCrop(\n            keys=[\"image\", \"label\"],\n            roi_size=input_shape,\n            random_center=True,\n            random_size=False,\n            invalid_label=2,         \n            min_valid_ratio=0.5,     \n            max_attempts=10\n        ),\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(\n            keys=[\"image\", \"label\"], \n            prob=0.4, \n            max_k=3, \n            spatial_axes=(0, 1)\n        ),\n        RandRotate(\n            keys=[\"image\", \"label\"], \n            factor=0.2, \n            prob=0.7, \n            fill_mode=\"crop\",\n        ),\n\n        ## Intensiry transformation\n        NormalizeIntensity(\n            keys=[\"image\"], \n            nonzero=True,\n            channel_wise=False\n        ),\n        RandShiftIntensity(\n            keys=[\"image\"], offsets=0.10, prob=0.5\n        ),\n        ## Spatial transformation \n        RandCutOut(\n            keys=[\"image\", \"label\"],\n            invalid_label=2, \n            mask_size=[\n                input_shape[1]//4,\n                input_shape[2]//4\n            ],\n            fill_mode=\"constant\", # \"constant\", \"gaussian\"\n            cutout_mode='volume', # \"slice\", \"volume\"\n            prob=0.8,\n            num_cuts=5,\n        ),\n    ])\n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]\n\ndef val_transformation(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\n        NormalizeIntensity(\n            keys=[\"image\"], \n            nonzero=True,\n            channel_wise=False\n        ),\n    ])\n    result = pipeline(data)\n    return result[\"image\"], result[\"label\"]\n\ndef tfrecord_loader(tfrecord_pattern, batch_size=1, shuffle=True):\n    dataset = tf.data.TFRecordDataset(\n        tf.io.gfile.glob(tfrecord_pattern),\n        num_parallel_reads=tf.data.AUTOTUNE\n    )\n    dataset = dataset.shuffle(buffer_size=400) if shuffle else dataset \n    dataset = dataset.map(\n        parse_tfrecord_fn, num_parallel_calls=tf.data.AUTOTUNE\n    )\n    dataset = dataset.map(\n        prepare_inputs,\n        num_parallel_calls=tf.data.AUTOTUNE\n    )\n    if shuffle:\n        dataset = dataset.map(\n            train_transformation,\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n    else:\n        dataset = dataset.map(\n            val_transformation,\n            num_parallel_calls=tf.data.AUTOTUNE\n        )\n    dataset = dataset.batch(batch_size, drop_remainder=shuffle)\n    dataset = dataset.prefetch(tf.data.AUTOTUNE)\n    return dataset\n\n# ========== 加载数据集 ==========\nall_tfrec = sorted(\n    glob.glob(\"/kaggle/input/vesuvius-tfrecords/*.tfrec\"),\n    key=lambda x: int(x.split(\"_\")[-1].replace(\".tfrec\", \"\"))\n)\n\nval_idx = -2\nval_patterns = all_tfrec[val_idx:]\ntrain_patterns = [\n    f for i, f in enumerate(all_tfrec) if i < len(all_tfrec) + val_idx\n]\n\ntrain_ds = tfrecord_loader(\n    train_patterns, batch_size=batch_size, shuffle=True\n)\nval_ds = tfrecord_loader(\n    val_patterns, batch_size=infer_batch_size, shuffle=False\n)\n\n# ========== 模型构建与权重加载（适配版） ==========\ndef build_model(input_shape, num_classes):\n    \"\"\"独立的模型构建函数，方便创建CPU副本\"\"\"\n    model = TransUNet(\n        input_shape=input_shape,\n        encoder_name='seresnext50',\n        classifier_activation='softmax',\n        num_classes=num_classes,\n    )\n    return model\n\n# 构建TPU训练模型\nmodel = build_model((160, 160, 160, 1), num_classes)\n\n# 加载预训练权重（适配版：移除get_distribution，直接加载）\nweights_path = \"/kaggle/input/notebooks/tonyai007/train-vesuvius-seresnext50-comboloss/final_model.weights.h5\"\nif os.path.exists(weights_path):\n    print(f\"Loading weights from: {weights_path}\")\n    # 适配版：直接加载权重，依赖skip_mismatch处理分片差异\n    model.load_weights(weights_path, skip_mismatch=True)\n    print(\"Weights loaded successfully!\")\nelse:\n    raise FileNotFoundError(f\"Weights file not found: {weights_path}\")\n\n# ========== 学习率调度器 ==========\nsteps_per_epoch = num_samples // batch_size\ntotal_continue_steps = continue_epochs * steps_per_epoch\ncompleted_steps = start_epoch * steps_per_epoch\n\nlr_schedule = CosineDecay(\n    initial_learning_rate=1e-6,\n    decay_steps=total_continue_steps,\n    warmup_target=min(3e-4, 1e-4 * (batch_size / 2)),\n    warmup_steps=0,\n    alpha=0.1,\n)\n\n# ========== 优化器、损失函数、指标 ==========\noptim = keras.optimizers.AdamW(\n    learning_rate=lr_schedule,\n    weight_decay=1e-5,\n)\n\ndice_ce_loss_fn = SparseDiceCELoss(\n    from_logits=False, \n    num_classes=num_classes,\n    ignore_class_ids=2,\n)\ncldice_loss_fn = SparseCenterlineDiceLoss(\n    from_logits=False, \n    num_classes=num_classes,\n    target_class_ids=1,\n    ignore_class_ids=2,\n    iters=50\n)\ncombined_loss_fn = lambda y_true, y_pred: (\n    dice_ce_loss_fn(y_true, y_pred) + cldice_loss_fn(y_true, y_pred)\n)\n\nmetrics = [\n    SparseDiceMetric(\n        from_logits=False, \n        num_classes=num_classes, \n        ignore_class_ids=2,\n        name='dice'\n    ),\n]\n\nmodel.compile(\n    optimizer=optim,\n    loss=combined_loss_fn,\n    metrics=metrics,\n)\n\n# ========== 终极稳定版SWI回调（适配无get_distribution的版本） ==========\nclass UltimateStableSWICallback(keras.callbacks.Callback):\n    def __init__(\n        self,\n        dataset,\n        num_classes,\n        weights_path=\"temp_weights.h5\",\n        infer_weights_path=\"model.weights.h5\",\n        interval=5,\n        overlap=0.3,\n        roi_size=(160,160,160)\n    ):\n        super().__init__()\n        self.dataset = dataset\n        self.num_classes = num_classes\n        self.weights_path = weights_path\n        self.infer_weights_path = infer_weights_path\n        self.interval = interval\n        self.overlap = overlap\n        self.roi_size = roi_size\n        self.best_score = 0.0\n        self.infer_metric = SparseDiceMetric(\n            from_logits=False,\n            ignore_class_ids=2,\n            num_classes=num_classes,\n            name='infer_dice'\n        )\n        \n    def _create_cpu_model(self):\n        \"\"\"创建纯CPU的模型副本（无TPU分片）\"\"\"\n        # 切换到CPU环境\n        os.environ[\"JAX_PLATFORM_NAME\"] = \"cpu\"\n        \n        # 构建CPU模型（无分布式）\n        cpu_model = build_model((160, 160, 160, 1), self.num_classes)\n        # 加载最新权重\n        cpu_model.load_weights(self.weights_path)\n        \n        return cpu_model\n    \n    def _run_cpu_inference(self):\n        \"\"\"在CPU上执行滑动窗口推理\"\"\"\n        # 1. 保存当前TPU模型权重（直接保存，不切换分布式）\n        self.model.save_weights(self.weights_path)\n        \n        # 2. 创建CPU模型并推理\n        try:\n            cpu_model = self._create_cpu_model()\n            \n            # 初始化滑动窗口推理器（CPU版）\n            sw_inferer = SlidingWindowInference(\n                model=cpu_model,\n                num_classes=self.num_classes,\n                roi_size=self.roi_size,\n                overlap=self.overlap,\n                mode='constant',\n                sw_batch_size=1,\n            )\n            \n            self.infer_metric.reset_states()\n            \n            # 逐样本推理\n            for batch_data in self.dataset:\n                x, y = batch_data\n                x_np = x.numpy()\n                y_np = y.numpy()\n                \n                # CPU推理\n                pred_np = sw_inferer(x_np)\n                \n                # 更新指标\n                self.infer_metric.update_state(y_np, pred_np)\n                \n                # 强制清理内存\n                del x, y, x_np, y_np, pred_np\n                import gc\n                gc.collect()\n            \n            score = self.infer_metric.result()\n            score_np = np.asarray(score)\n            \n            return score_np\n            \n        finally:\n            # 清理CPU模型\n            del cpu_model\n            gc.collect()\n            # 恢复TPU环境\n            os.environ[\"JAX_PLATFORM_NAME\"] = \"tpu\"\n        \n    def on_epoch_end(self, epoch, logs=None):\n        if (epoch + 1) % self.interval == 0:\n            print(f\"\\nEpoch {epoch}: Running CPU-based inference (stable mode)...\")\n            try:\n                # 执行CPU推理\n                score = self._run_cpu_inference()\n                print(f\"Epoch {epoch}: Inference Dice Score = {score:.4f}\")\n                \n                # 保存最优模型（直接保存，不切换分布式）\n                if score > self.best_score:\n                    self.best_score = score\n                    self.model.save_weights(self.infer_weights_path)\n                    print(f\"New best score! Model saved to {self.infer_weights_path}\")\n                    \n            except Exception as e:\n                print(f\"Inference warning (non-critical): {str(e)[:200]}\")\n                print(\"Continuing training...\")\n\n# ========== 回调函数 ==========\nswi_callback = UltimateStableSWICallback(\n    dataset=val_ds,\n    num_classes=num_classes,\n    weights_path=\"temp_tpu_weights.h5\",\n    infer_weights_path=\"best_model.weights.h5\",\n    interval=5,\n    overlap=0.3,\n    roi_size=input_shape\n)\n\nearly_stopping = keras.callbacks.EarlyStopping(\n    monitor='dice',\n    patience=80,\n    mode='max',\n    restore_best_weights=True,\n    verbose=1\n)\n\n# ========== 开始训练（8核TPU，适配版） ==========\nprint(f\"Starting training with 8-core TPU (stable mode)\")\nhistory = model.fit(\n    train_ds,\n    initial_epoch=start_epoch,\n    epochs=total_epochs,\n    callbacks=[\n        swi_callback,\n        early_stopping\n    ],\n    verbose=1\n)\n\n# 训练结束后保存最终模型\nprint(\"Training completed. Saving final model...\")\nmodel.save_weights(\"final_model.weights.h5\")\nprint(\"Final model saved successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-17T17:00:24.382611Z","iopub.execute_input":"2026-02-17T17:00:24.382941Z","iopub.status.idle":"2026-02-17T17:37:12.22046Z","shell.execute_reply.started":"2026-02-17T17:00:24.382918Z","shell.execute_reply":"2026-02-17T17:37:12.218819Z"}},"outputs":[],"execution_count":null}]}