{"metadata":{"kernelspec":{"display_name":"Python 3 (ipykernel)","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.14"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":14266465,"datasetId":8751895,"databundleVersionId":15066994},{"sourceType":"kernelVersion","sourceId":290917305,"isSourceIdPinned":false}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false},"papermill":{"default_parameters":{},"duration":29118.985133,"end_time":"2026-02-03T01:11:23.867492","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-02-02T17:06:04.882359","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"3f4f1ab3-8a0c-41d4-9cf7-e7f2db1275b2","cell_type":"markdown","source":"# Bronze Medal - uuNet by ChatGPT5.2 - [Train]\nThis is the notebook to train my ChatGPT5.2 vibe coded bronze medal in Vesuvius Challenge. This notebook was trained offline on 1xA100 80GB VRAM. It took a few days to run. (Then the resultant notebook was uploaded here). The inference notebook is [here][1]. Discussion is [here][2]\n\n[1]: https://www.kaggle.com/code/cdeotte/infer-bronze-medal-uunet-by-chatgpt\n[2]: https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/writeups/bronze-medal-chatgpt-vibe-coding","metadata":{}},{"id":"e2abc496-28a9-422a-b29b-c5712cdb0e30","cell_type":"code","source":"# ============================================================\n# FULL DROP-IN REPLACEMENT (YOUR SCRIPT + (B) + (C))\n#\n# (B) ✅ Decoder closer to host nnUNetv2 plan: n_conv_per_stage_decoder = 1\n#     - Replace decoder \"res_block (2 convs)\" with a single conv block after concat.\n#\n# (C) ✅ Replace clDice with a host-like \"MedialSurfaceRecall-ish\" loss (approx)\n#     - Uses medicai.utils.soft_skeletonize (differentiable) to compute a skeleton\n#     - Optimizes *recall on the skeleton/medial structure* to reduce holes / breaks\n#     - Loss is computed in float32 (stable under mixed_bfloat16)\n#\n# Everything else preserved:\n#   - mixed_bfloat16 compute + fp32 dataloader outputs\n#   - ResEncUNet encoder depth [1,3,4,6,6,6] + InstanceNorm3D\n#   - 4-head deep supervision (aux20/40/80/main); aux losses are DiceCE only\n#   - AdamW + CosineDecay\n# ============================================================\n\nimport os\nos.environ[\"CUDA_VISIBLE_DEVICES\"] = \"7\"\n\nVER = 31\n\nimport glob\nfrom tqdm import tqdm\nimport numpy as np\nimport pandas as pd\nfrom matplotlib import pyplot as plt\n\nimport os, warnings\nos.environ[\"KERAS_BACKEND\"] = \"torch\"\nwarnings.filterwarnings('ignore')\n\nimport keras\nfrom keras import ops\nfrom keras.optimizers.schedules import CosineDecay\n\nimport torch\nimport torch.nn.functional as F\nimport tensorflow as tf\n\n# ✅ mixed precision policy\nkeras.mixed_precision.set_global_policy(\"mixed_bfloat16\")\nprint(\"Keras:\", keras.__version__, \"| Torch:\", torch.__version__, \"| Backend:\", keras.config.backend())\nprint(\"Mixed precision policy:\", keras.mixed_precision.global_policy())\n\n# mainly for 3D or 2D models, transformation, loss, metrics etc\nimport medicai\nfrom medicai.transforms import (\n    Compose,\n    NormalizeIntensity,\n    RandShiftIntensity,\n    RandRotate90,\n    RandFlip,\n    RandCutOut,\n    RandSpatialCrop\n)\nfrom medicai.losses import (\n    SparseDiceCELoss,\n)\nfrom medicai.metrics import SparseDiceMetric\nfrom medicai.utils import SlidingWindowInference\nfrom medicai.utils import soft_skeletonize  # ✅ used for (C)\n\ninput_shape=(160, 160, 160)\nbatch_size=3\nnum_classes=3\n\n# Each tfrecord contains 6 samples, total 786 samples.\nnum_samples = 780\nepochs = 1_000\n\n# ============================================================\n# DATA AUG (UNCHANGED) — RETURNS FP32 (keep)\n# ============================================================\ndef train_transformation(image, label):\n    data = {\"image\": image, \"label\": label}\n    pipeline = Compose([\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        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        RandCutOut(\n            keys=[\"image\", \"label\"],\n            invalid_label=2,\n            mask_size=[input_shape[1]//4, input_shape[2]//4],\n            fill_mode=\"constant\",\n            cutout_mode='volume',\n            prob=0.2,\n            num_cuts=2,\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 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    # keep dataloader outputs fp32 (as requested)\n    image = image[..., None]  # (D,H,W,1)\n    label = label[..., None]  # (D,H,W,1)\n    image = tf.cast(image, tf.float32)\n    label = tf.cast(label, tf.float32)\n    return image, label\n\ndef tfrecord_loader(tfrecord_pattern, batch_size=1, shuffle=True):\n    dataset = tf.data.TFRecordDataset(tf.io.gfile.glob(tfrecord_pattern))\n    dataset = dataset.shuffle(buffer_size=100) if shuffle else dataset\n    dataset = dataset.map(parse_tfrecord_fn, num_parallel_calls=tf.data.AUTOTUNE)\n    dataset = dataset.map(prepare_inputs, num_parallel_calls=tf.data.AUTOTUNE)\n    if shuffle:\n        dataset = dataset.map(train_transformation, num_parallel_calls=tf.data.AUTOTUNE)\n    else:\n        dataset = dataset.map(val_transformation, num_parallel_calls=tf.data.AUTOTUNE)\n    dataset = dataset.batch(batch_size, drop_remainder=shuffle)\n    dataset = dataset.prefetch(tf.data.AUTOTUNE)\n    return dataset\n\nall_tfrec = sorted(\n    glob.glob(\"data/*.tfrec\"),\n    key=lambda x: int(x.split(\"_\")[-1].replace(\".tfrec\", \"\"))\n)\n\nval_idx = -1\nval_patterns = [all_tfrec[val_idx]]\ntrain_patterns = [f for i, f in enumerate(all_tfrec) if i != len(all_tfrec) + val_idx]\n\ntrain_ds = tfrecord_loader(train_patterns, batch_size=batch_size, shuffle=True)\nval_ds   = tfrecord_loader(val_patterns,   batch_size=1,         shuffle=False)\n\n# ============================================================\n# Model: ResEncUNet + InstanceNorm3D + DeepSup4\n# (B) Decoder changed to 1 conv per stage after concat (host-like)\n# ============================================================\nfrom keras import layers\n\nclass InstanceNorm3D(layers.Layer):\n    def __init__(self, eps=1e-5, affine=True, name=None):\n        super().__init__(name=name)\n        self.eps = eps\n        self.affine = affine\n\n    def build(self, input_shape):\n        c = int(input_shape[-1])\n        if self.affine:\n            self.gamma = self.add_weight(name=\"gamma\", shape=(c,), initializer=\"ones\", trainable=True)\n            self.beta  = self.add_weight(name=\"beta\",  shape=(c,), initializer=\"zeros\", trainable=True)\n        else:\n            self.gamma, self.beta = None, None\n        super().build(input_shape)\n\n    def call(self, x):\n        mean = ops.mean(x, axis=(1, 2, 3), keepdims=True)\n        var  = ops.mean(ops.square(x - mean), axis=(1, 2, 3), keepdims=True)\n        xhat = (x - mean) / ops.sqrt(var + self.eps)\n        if self.affine:\n            gamma = ops.reshape(self.gamma, (1, 1, 1, 1, -1))\n            beta  = ops.reshape(self.beta,  (1, 1, 1, 1, -1))\n            xhat = xhat * gamma + beta\n        return xhat\n\ndef IN(x, name: str, eps=1e-5, affine=True):\n    return InstanceNorm3D(eps=eps, affine=affine, name=name)(x)\n\ndef conv3d(x, out_ch, k=3, s=1, name=None):\n    return layers.Conv3D(out_ch, kernel_size=k, strides=s, padding=\"same\", use_bias=True, name=name)(x)\n\ndef res_block_in(x, out_ch, name: str, leaky=0.01, eps=1e-5):\n    in_ch = int(x.shape[-1])\n\n    y = conv3d(x, out_ch, k=3, s=1, name=f\"{name}_conv1\")\n    y = IN(y, name=f\"{name}_in1\", eps=eps, affine=True)\n    y = layers.LeakyReLU(alpha=leaky, name=f\"{name}_lrelu1\")(y)\n\n    y = conv3d(y, out_ch, k=3, s=1, name=f\"{name}_conv2\")\n    y = IN(y, name=f\"{name}_in2\", eps=eps, affine=True)\n\n    if in_ch != out_ch:\n        skip = conv3d(x, out_ch, k=1, s=1, name=f\"{name}_skip\")\n        skip = IN(skip, name=f\"{name}_skip_in\", eps=eps, affine=True)\n    else:\n        skip = x\n\n    y = layers.Add(name=f\"{name}_add\")([y, skip])\n    y = layers.LeakyReLU(alpha=leaky, name=f\"{name}_lrelu2\")(y)\n    return y\n\ndef downsample_in(x, out_ch, name: str, leaky=0.01, eps=1e-5):\n    x = conv3d(x, out_ch, k=3, s=2, name=f\"{name}_downconv\")\n    x = IN(x, name=f\"{name}_down_in\", eps=eps, affine=True)\n    x = layers.LeakyReLU(alpha=leaky, name=f\"{name}_down_lrelu\")(x)\n    return x\n\n# (B) ✅ Single conv decoder block (1 conv per stage)\ndef dec_conv1_in(x, out_ch, name: str, leaky=0.01, eps=1e-5):\n    x = conv3d(x, out_ch, k=3, s=1, name=f\"{name}_conv\")\n    x = IN(x, name=f\"{name}_in\", eps=eps, affine=True)\n    x = layers.LeakyReLU(alpha=leaky, name=f\"{name}_lrelu\")(x)\n    return x\n\ndef up_block_in(x, skip, out_ch, name: str, leaky=0.01, eps=1e-5):\n    x = layers.Conv3DTranspose(out_ch, kernel_size=2, strides=2, padding=\"same\",\n                               use_bias=True, name=f\"{name}_up\")(x)\n    x = IN(x, name=f\"{name}_up_in\", eps=eps, affine=True)\n    x = layers.LeakyReLU(alpha=leaky, name=f\"{name}_up_lrelu\")(x)\n\n    x = layers.Concatenate(axis=-1, name=f\"{name}_cat\")([x, skip])\n\n    # ✅ host-like: 1 conv per decoder stage (instead of 2-conv residual block)\n    x = dec_conv1_in(x, out_ch, name=f\"{name}_dec1\", leaky=leaky, eps=eps)\n    return x\n\ndef build_nnunet3d_resenc_deepsup4(\n    input_shape=(160, 160, 160, 1),\n    num_classes=3,\n    features_per_stage=(32, 64, 128, 256, 320, 320),\n    n_blocks_per_stage=(1, 3, 4, 6, 6, 6),\n    leaky=0.01,\n    classifier_activation=\"softmax\",\n    eps=1e-5,\n):\n    assert len(features_per_stage) == 6\n    assert len(n_blocks_per_stage) == 6\n\n    inputs = keras.Input(shape=input_shape, name=\"image\")\n\n    # Encoder stage 0\n    x = conv3d(inputs, features_per_stage[0], k=3, s=1, name=\"enc0_conv\")\n    x = IN(x, name=\"enc0_in\", eps=eps, affine=True)\n    x = layers.LeakyReLU(alpha=leaky, name=\"enc0_lrelu\")(x)\n    for b in range(n_blocks_per_stage[0]):\n        x = res_block_in(x, features_per_stage[0], name=f\"enc0_res{b}\", leaky=leaky, eps=eps)\n    skip0 = x\n    skips = [skip0]\n\n    # Encoder stages 1..5\n    for s in range(1, 6):\n        x = downsample_in(x, features_per_stage[s], name=f\"enc{s}\", leaky=leaky, eps=eps)\n        for b in range(n_blocks_per_stage[s]):\n            x = res_block_in(x, features_per_stage[s], name=f\"enc{s}_res{b}\", leaky=leaky, eps=eps)\n        skips.append(x)\n\n    # Decoder dec4..dec0\n    feat_dec3 = feat_dec2 = feat_dec1 = None\n    for s in reversed(range(0, 5)):\n        x = up_block_in(x, skips[s], features_per_stage[s], name=f\"dec{s}\", leaky=leaky, eps=eps)\n        if s == 3: feat_dec3 = x  # 20^3\n        if s == 2: feat_dec2 = x  # 40^3\n        if s == 1: feat_dec1 = x  # 80^3\n    feat_dec0 = x                # 160^3\n\n    def head(feat, name):\n        lg = layers.Conv3D(num_classes, kernel_size=1, padding=\"same\", name=f\"{name}_logits\")(feat)\n        return layers.Activation(classifier_activation, name=f\"{name}_prob\")(lg) if classifier_activation else lg\n\n    aux20 = head(feat_dec3, \"aux20\")\n    aux40 = head(feat_dec2, \"aux40\")\n    aux80 = head(feat_dec1, \"aux80\")\n    main  = head(feat_dec0, \"main\")\n\n    model_train = keras.Model(inputs=inputs, outputs=[aux80, aux40, aux20, main],\n                              name=\"nnunet3d_resenc_deepsup4_train_in\")\n    model_infer = keras.Model(inputs=inputs, outputs=main,\n                              name=\"nnunet3d_resenc_deepsup4_infer_in\")\n    return model_train, model_infer\n\nmodel_train, model_infer = build_nnunet3d_resenc_deepsup4(\n    input_shape=input_shape + (1,),\n    num_classes=num_classes,\n    features_per_stage=(32, 64, 128, 256, 320, 320),\n    n_blocks_per_stage=(1, 3, 4, 6, 6, 6),\n    classifier_activation=\"softmax\",\n    eps=1e-5,\n)\n\nprint(\"Train-model params (M):\", model_train.count_params() / 1e6)\nprint(\"Infer-model params (M):\", model_infer.count_params() / 1e6)\n\n# ============================================================\n# LR SCHEDULE + OPTIMIZER (UNCHANGED)\n# ============================================================\nsteps_per_epoch = num_samples // batch_size\ntotal_steps = steps_per_epoch * epochs\nwarmup_steps = int(total_steps * 0.05)\ndecay_steps = max(1, total_steps - warmup_steps)\n\nlr_schedule = CosineDecay(\n    initial_learning_rate=1e-6,\n    decay_steps=decay_steps,\n    warmup_target=min(3e-4, 1e-4 * (batch_size / 2)),\n    warmup_steps=warmup_steps,\n    alpha=0.1,\n)\n\noptim = keras.optimizers.AdamW(\n    learning_rate=lr_schedule,\n    weight_decay=1e-5,\n)\n\n# ============================================================\n# (C) LOSS: DiceCE + MedialSurfaceRecall-ish (approx) on MAIN\n#     Aux heads: DiceCE only\n# ============================================================\ndice_ce_loss_fn = SparseDiceCELoss(\n    from_logits=False,\n    num_classes=num_classes,\n    ignore_class_ids=2,\n)\n\ndef _to_torch(x):\n    return ops.convert_to_tensor(x)\n\ndef downsample_label_nearest_3d(y_true_bdhwc, out_dhw):\n    \"\"\"\n    y_true: (B,D,H,W,1) float32 class ids\n    out_dhw: (D,H,W)\n    \"\"\"\n    y = _to_torch(y_true_bdhwc)\n    y = y.permute(0, 4, 1, 2, 3)            # (B,1,D,H,W)\n    y_ds = F.interpolate(y, size=out_dhw, mode=\"nearest\")\n    y_ds = y_ds.permute(0, 2, 3, 4, 1)      # (B,D,H,W,1)\n    return y_ds\n\n# Deep supervision weights (unchanged)\nW_MAIN  = 1.00\nW_AUX80 = 0.50\nW_AUX40 = 0.25\nW_AUX20 = 0.10\n\n# (C) ✅ MedialSurfaceRecall-ish loss (approx) using soft_skeletonize\n# - Treat target_class_id=1 as \"ink\"\n# - Compute soft skeleton of GT mask and maximize recall under prediction prob\nMSR_TARGET_CLASS_ID = 1\nMSR_WEIGHT = 0.20      # start modest; try 0.10, 0.20, 0.30\nMSR_ITERS  = 15        # soft skeleton iterations; try 10..25\n_EPS = 1e-6\n\ndef _extract_fg_prob(y_prob, class_id: int):\n    # y_prob: (B,D,H,W,C)\n    return y_prob[..., class_id:class_id+1]\n\ndef medial_surface_recall_loss_approx(y_true, y_prob, target_class_id=1, ignore_class_id=2, iters=15):\n    \"\"\"\n    y_true: (B,D,H,W,1) class ids (float32)\n    y_prob: (B,D,H,W,C) probabilities (float/bfloat)\n    Returns: scalar loss in float32\n    \"\"\"\n    y_true_f = ops.cast(y_true, \"float32\")\n    y_prob_f = ops.cast(y_prob, \"float32\")\n\n    # binary GT mask for target class\n    gt = ops.cast(ops.equal(y_true_f, float(target_class_id)), \"float32\")  # (B,D,H,W,1)\n\n    # ignore mask: 1 where valid, 0 where ignore\n    valid = ops.cast(ops.not_equal(y_true_f, float(ignore_class_id)), \"float32\")\n\n    # prediction prob for the same class\n    pr = _extract_fg_prob(y_prob_f, target_class_id)  # (B,D,H,W,1)\n    pr = ops.clip(pr, 0.0, 1.0)\n\n    # mask out ignore voxels\n    gt = gt * valid\n    pr = pr * valid\n\n    # soft skeleton of GT (differentiable)\n    # NOTE: soft_skeletonize expects float tensor in [0,1]\n    skel_gt = soft_skeletonize(gt, iters=iters)\n    skel_gt = ops.clip(skel_gt, 0.0, 1.0)\n\n    # Recall of GT skeleton under prediction: sum(skel_gt * pr) / sum(skel_gt)\n    num = ops.sum(skel_gt * pr)\n    den = ops.sum(skel_gt) + _EPS\n    recall = num / den\n\n    return 1.0 - recall\n\ndef dice_ce_only(y_true, y_pred):\n    y_true_f = ops.cast(y_true, \"float32\")\n    y_pred_f = ops.cast(y_pred, \"float32\")\n    return dice_ce_loss_fn(y_true_f, y_pred_f)\n\ndef combined_loss_main(y_true, y_pred_main_prob):\n    # DiceCE + MSR-ish (approx)\n    y_true_f = ops.cast(y_true, \"float32\")\n    y_pred_f = ops.cast(y_pred_main_prob, \"float32\")\n    loss_dice = dice_ce_loss_fn(y_true_f, y_pred_f)\n    loss_msr  = medial_surface_recall_loss_approx(\n        y_true_f, y_pred_f,\n        target_class_id=MSR_TARGET_CLASS_ID,\n        ignore_class_id=2,\n        iters=MSR_ITERS\n    )\n    return loss_dice + (MSR_WEIGHT * loss_msr)\n\ndef deep_supervision_loss(y_true, y_preds):\n    \"\"\"\n    y_preds: [aux80_prob, aux40_prob, aux20_prob, main_prob]\n    \"\"\"\n    aux80, aux40, aux20, main = y_preds\n\n    s80 = aux80.shape\n    s40 = aux40.shape\n    s20 = aux20.shape\n\n    y80 = downsample_label_nearest_3d(y_true, (int(s80[1]), int(s80[2]), int(s80[3])))\n    y40 = downsample_label_nearest_3d(y_true, (int(s40[1]), int(s40[2]), int(s40[3])))\n    y20 = downsample_label_nearest_3d(y_true, (int(s20[1]), int(s20[2]), int(s20[3])))\n\n    loss_main = combined_loss_main(y_true, main)\n    loss_80   = dice_ce_only(y80, aux80)\n    loss_40   = dice_ce_only(y40, aux40)\n    loss_20   = dice_ce_only(y20, aux20)\n\n    return (W_MAIN * loss_main) + (W_AUX80 * loss_80) + (W_AUX40 * loss_40) + (W_AUX20 * loss_20)\n\n# ============================================================\n# SWI for validation uses ONLY main head (infer model)\n# ============================================================\nswi = SlidingWindowInference(\n    model_infer,\n    num_classes=num_classes,\n    roi_size=input_shape,\n    sw_batch_size=1,\n    overlap=0.5,\n    mode='gaussian',\n)\n\n# ============================================================\n# TRAIN / VALIDATE\n# ============================================================\ndef train_one_epoch(model_train, dataloader, metrics):\n    loop = tqdm(dataloader, desc=\"Training\", leave=False)\n\n    for imgs, labels in loop:\n        y_preds = model_train(imgs)   # [aux80, aux40, aux20, main]\n        loss = deep_supervision_loss(labels, y_preds)\n\n        model_train.zero_grad()\n        trainable_weights = [v for v in model_train.trainable_weights]\n\n        loss.backward()\n        gradients = [v.value.grad for v in trainable_weights]\n\n        with torch.no_grad():\n            optim.apply(gradients, trainable_weights)\n\n        main_pred = y_preds[-1]\n        metrics.update_state(\n            ops.cast(labels, \"float32\"),\n            ops.cast(main_pred, \"float32\"),\n        )\n\n        loss_score = ops.convert_to_numpy(loss)\n        metrics_score = ops.convert_to_numpy(metrics.result())\n        loop.set_postfix(loss=loss_score, dice=metrics_score)\n\n    return loss, metrics\n\ndef validate(model_infer, dataloader, metrics):\n    for x, y in dataloader:\n        output = swi(x)\n        metrics.update_state(\n            ops.cast(y, \"float32\"),\n            ops.cast(output, \"float32\"),\n        )\n    return metrics\n\ndef run_training(train_loader, val_loader, model_train, model_infer, epochs=20):\n    train_metrics = SparseDiceMetric(\n        from_logits=False,\n        num_classes=num_classes,\n        ignore_class_ids=2,\n        name='dice'\n    )\n\n    val_metrics = SparseDiceMetric(\n        from_logits=False,\n        num_classes=num_classes,\n        ignore_class_ids=2,\n        name='val_dice'\n    )\n\n    best_val_dice = 0.0\n\n    for epoch in range(epochs):\n        print(f'Epoch {epoch+1}/{epochs}')\n\n        loss, train_metrics = train_one_epoch(model_train, train_loader, train_metrics)\n\n        train_metrics_score = ops.convert_to_numpy(train_metrics.result())\n        loss_score = ops.convert_to_numpy(loss)\n        train_metrics.reset_state()\n\n        if (epoch + 1) % 5 == 0:\n            val_metrics = validate(model_infer, val_loader, val_metrics)\n            val_metrics_score = ops.convert_to_numpy(val_metrics.result())\n            val_metrics.reset_state()\n\n            print(\n                f'Training - Loss: {loss_score:.4f}, Dice: {train_metrics_score:.4f}'\n                f'\\nValidation - Dice: {val_metrics_score:.4f}\\n'\n            )\n\n            if val_metrics_score > best_val_dice:\n                best_val_dice = val_metrics_score\n                model_train.save_weights(f'model_v{VER}.weights.h5')\n                print(f'Dice score improved: {best_val_dice}. Model saved.')\n        else:\n            print(f'Training - Loss: {loss_score:.4f}, Dice: {train_metrics_score:.4f}\\n')\n\nrun_training(train_ds, val_ds, model_train, model_infer, epochs=epochs)","metadata":{},"outputs":[],"execution_count":null}]}