{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"},{"sourceId":14746760,"sourceType":"datasetVersion","datasetId":8897381},{"sourceId":742034,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":558512,"modelId":571079},{"sourceId":744711,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":558512,"modelId":571079}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install dynamic_network_architectures==0.4.3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T01:09:40.981746Z","iopub.execute_input":"2026-02-08T01:09:40.982349Z","iopub.status.idle":"2026-02-08T01:09:52.136783Z","shell.execute_reply.started":"2026-02-08T01:09:40.982319Z","shell.execute_reply":"2026-02-08T01:09:52.135994Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install tensorrt==10.12.0.36","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T01:09:52.138482Z","iopub.execute_input":"2026-02-08T01:09:52.139084Z","iopub.status.idle":"2026-02-08T01:11:29.864288Z","shell.execute_reply.started":"2026-02-08T01:09:52.139050Z","shell.execute_reply":"2026-02-08T01:11:29.863340Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport tensorrt\nprint(tensorrt.__version__)\nimport tensorrt as trt\nimport os\n\nimport os\nimport sys\nimport subprocess\nfrom pathlib import Path\nfrom collections import OrderedDict\n\nimport torch\nimport torch.nn as nn\nimport pytorch_lightning as pl\nimport tensorrt as trt\nimport os\nimport sys\nimport subprocess\nfrom pathlib import Path\nfrom collections import OrderedDict\n\nimport torch\nimport torch.nn as nn\nimport pytorch_lightning as pl\nimport tensorrt as trt\n\nfrom dynamic_network_architectures.architectures.unet import ResidualEncoderUNet\nfrom dynamic_network_architectures.architectures.primus import PrimusB\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-08T01:11:29.865533Z","iopub.execute_input":"2026-02-08T01:11:29.865875Z","iopub.status.idle":"2026-02-08T01:11:55.750481Z","shell.execute_reply.started":"2026-02-08T01:11:29.865831Z","shell.execute_reply":"2026-02-08T01:11:55.749868Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ROI_SIZE = (160, 160, 160)\n\nOUT_DIR = Path(\"./trt_models\")\nOUT_DIR.mkdir(exist_ok=True, parents=True)\n\nDEVICE = \"cuda\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T01:11:55.751306Z","iopub.execute_input":"2026-02-08T01:11:55.751514Z","iopub.status.idle":"2026-02-08T01:11:55.755693Z","shell.execute_reply.started":"2026-02-08T01:11:55.751492Z","shell.execute_reply":"2026-02-08T01:11:55.755062Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def create_residual_unet(\n#     in_channels,\n#     out_channels=2,\n#     channels=(32, 64, 128, 256, 320, 320),\n#     strides=(1, 2, 2, 2, 2, 2),\n#     n_blocks_per_stage=(1, 3, 4, 6, 6, 6),\n#     deep_supervision=False,\n# ):\n#     \"\"\"Create ResidualEncoderUNet matching nnUNet 3d_fullres configuration.\"\"\"\n#     n_conv_per_stage_decoder = [1] * (len(channels) - 1)\n\n#     return ResidualEncoderUNet(\n#         input_channels=in_channels,\n#         n_stages=len(channels),\n#         features_per_stage=channels,\n#         conv_op=nn.Conv3d,\n#         kernel_sizes=3,\n#         strides=strides,\n#         n_blocks_per_stage=n_blocks_per_stage,\n#         num_classes=out_channels,\n#         n_conv_per_stage_decoder=n_conv_per_stage_decoder,\n#         conv_bias=True,\n#         norm_op=nn.InstanceNorm3d,\n#         norm_op_kwargs={},\n#         dropout_op=None,\n#         nonlin=nn.LeakyReLU,\n#         nonlin_kwargs={\"inplace\": True},\n#         deep_supervision=deep_supervision,\n#     )\n\n\n\n# class SegmentationModule(pl.LightningModule):\n#     \"\"\"\n#     Lightning module matching training code structure.\n#     Uses self.model to match checkpoint keys.\n#     \"\"\"\n#     def __init__(\n#         self,\n#         in_channels=2,\n#         out_channels=2,\n#         channels=(32, 64, 128, 256, 320, 320),\n#         strides=(1, 2, 2, 2, 2, 2),\n#         **kwargs  # Accept extra hparams from checkpoint\n#     ):\n#         super().__init__()\n#         self.save_hyperparameters()\n\n\n#         n_blocks = kwargs.get('n_blocks_per_stage', (1, 3, 4, 6, 6, 6))\n#         use_ds = kwargs.get('use_deep_supervision', False)\n            \n#         self.model = create_residual_unet(\n#             in_channels=in_channels,\n#             out_channels=out_channels,\n#             channels=channels,\n#             strides=strides,\n#             n_blocks_per_stage=n_blocks,\n#             deep_supervision=use_ds\n#         )\n\n#     def forward(self, x):\n#         return self.model(x)\n\ndef create_primus(\n        in_channels=1,\n        out_channels=2,\n        input_shape = 160,\n        patch_embed_size = 8\n):\n    model = PrimusB(in_channels,\n                    out_channels,\n                    (patch_embed_size, patch_embed_size, patch_embed_size),\n                    (input_shape, input_shape, input_shape))\n    return model\n\nclass SegmentationModulePrimus(pl.LightningModule):\n    \"\"\"\n    Lightning module matching training code structure.\n    Uses self.model to match checkpoint keys.\n    \"\"\"\n    def __init__(\n        self,\n        model_type='segformer',\n        in_channels=2,\n        out_channels=2,\n        channels=(32, 64, 128, 256, 320, 320),\n        strides=(1, 2, 2, 2, 2, 2),\n        **kwargs  # Accept extra hparams from checkpoint\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n\n\n        n_blocks = kwargs.get('n_blocks_per_stage', (1, 3, 4, 6, 6, 6))\n        use_ds = kwargs.get('use_deep_supervision', False)\n            \n        self.model = create_primus()\n\n    def forward(self, x):\n        return self.model(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T01:11:55.757940Z","iopub.execute_input":"2026-02-08T01:11:55.758269Z","iopub.status.idle":"2026-02-08T01:11:57.446542Z","shell.execute_reply.started":"2026-02-08T01:11:55.758218Z","shell.execute_reply":"2026-02-08T01:11:57.445688Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# Model Loading - Simple approach using Lightning's load_from_checkpoint\n# ==============================================================================\nMODEL_PATHS = [\n    (\"/kaggle/input/unet3d-1st-stage-vesuvius-challenge/pytorch/default/12/primus_pretrained/primus_pretrained/best-epoch=264-val_loss=0.3657-val_dice=0.5852.ckpt\", \"cuda:0\"),\n    (\"/kaggle/input/unet3d-1st-stage-vesuvius-challenge/pytorch/default/12/primus_pretrained/primus_pretrained/best-epoch=294-val_loss=0.3653-val_dice=0.5917.ckpt\", \"cuda:0\")\n    \n    # (\"./vesuvius-model-zoo/fst_stage_unet/fold-0-best-epoch369-val_loss0.3720-val_dice0.5789.ckpt\", \"cuda:0\"),\n    #(\"/kaggle/input/vesuvius-sergio-models/fold0_new_ema_pretrain/best-epoch369-val_loss0.3720-val_dice0.5789.ckpt\", \"cuda:0\"),\n    #(\"/kaggle/input/vesuvius-sergio-models/fold1_new_ema_pretrain/best-epoch364-val_loss0.3767-val_dice0.5813.ckpt\", \"cuda:0\"),\n    \n    # (\"/kaggle/input/unet3d-1st-stage-vesuvius-challenge/pytorch/default/1/fold-2-best-epoch384-val_loss0.3647-val_dice0.5865.ckpt\", \"cuda:0\"),\n    # (\"/kaggle/input/unet3d-1st-stage-vesuvius-challenge/pytorch/default/1/fold-3-best-epoch349-val_loss0.3471-val_dice0.6062.ckpt\", \"cuda:1\"),\n    #(\"/kaggle/input/unet3d-1st-stage-vesuvius-challenge/pytorch/default/1/fold-4-best-epoch364-val_loss0.3534-val_dice0.6028.ckpt\", \"cuda:0\"),\n]\n\ndef load_models_simple(model_paths):\n    \"\"\"Load models using Lightning's built-in checkpoint loading, handling EMA weights.\"\"\"\n    models = []\n    for (path, device) in model_paths:\n        print(f\"Loading: {path}\")\n\n        # First, load the checkpoint to check for EMA weights\n        checkpoint = torch.load(path, map_location=device)\n        state_dict = checkpoint.get('state_dict', checkpoint)\n\n        # Check if checkpoint has EMA weights\n        has_ema = any(k.startswith('ema.module.') for k in state_dict.keys())\n        use_ema_hparam = checkpoint.get('hyper_parameters', {}).get('use_ema', False)\n\n        # Create model from checkpoint\n        model = SegmentationModulePrimus.load_from_checkpoint(\n            path,\n            map_location=device,\n            pretrained_backbone=False,  # ensure no downloads\n            strict=False\n        )\n\n        # If EMA weights exist, extract and load them into model\n        if has_ema:\n            print(f\"  Found EMA weights in checkpoint, extracting...\")\n            ema_state_dict = {}\n            for k, v in state_dict.items():\n                if k.startswith('ema.module.'):\n                    # Remove 'ema.module.' prefix to get 'model.*' key\n                    new_key = 'model.' + k[len('ema.module.'):]\n                    ema_state_dict[new_key] = v\n\n            # Load EMA weights into the model\n            missing, unexpected = model.load_state_dict(ema_state_dict, strict=False)\n            if missing:\n                # Filter out non-model keys from missing (like loss functions, etc.)\n                model_missing = [k for k in missing if k.startswith('model.')]\n                if model_missing:\n                    print(f\"  Warning: Missing EMA keys: {model_missing[:5]}...\")\n            print(f\"  ✓ Loaded EMA weights into model\")\n\n        model.to(device)\n        model.eval()\n\n        # Print model info\n        model_type = getattr(model.hparams, 'model_type', 'unet')\n        print(f\"  Model type: {model_type}\")\n        if model_type == 'flex_unet':\n            backbone = getattr(model.hparams, 'backbone', 'resnet18')\n            print(f\"  Backbone: {backbone}\")\n        if has_ema or use_ema_hparam:\n            print(f\"  ✓ Using EMA weights for inference\")\n        else:\n            print(f\"  Standard weights (no EMA)\")\n        print(f\"  ✓ Loaded successfully\")\n\n        models.append((Path(path).stem,model))\n    print(f\"\\n✓ Total models loaded: {len(models)}\")\n    return models\n\nmodels_1st_stage = load_models_simple(MODEL_PATHS) if MODEL_PATHS else []\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T01:11:57.447693Z","iopub.execute_input":"2026-02-08T01:11:57.448151Z","iopub.status.idle":"2026-02-08T01:12:28.719559Z","shell.execute_reply.started":"2026-02-08T01:11:57.448109Z","shell.execute_reply":"2026-02-08T01:12:28.718901Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_engine(onnx_file_path, engine_file_path):\n    # 1. Setup Logger and Builder\n    logger = trt.Logger(trt.Logger.INFO)\n    builder = trt.Builder(logger)\n    \n    # 2. Create Network and Parser\n    # 1 << 0 (Explicit Batch) is required for ONNX importers\n    network = builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH))\n    parser = trt.OnnxParser(network, logger)\n    \n    # 3. Create Configuration\n    config = builder.create_builder_config()\n    \n    # Memory pool limit (workspace size). Give it plenty of RAM (e.g., 4GB)\n    config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 8 * (1 << 30))\n    \n    # Enable FP16 (Half Precision) for speed\n    if builder.platform_has_fast_fp16:\n        print(\"Enabling FP16...\")\n        config.set_flag(trt.BuilderFlag.FP16)\n    \n    # 4. Parse ONNX File\n    print(f\"Parsing {onnx_file_path}...\")\n    with open(onnx_file_path, 'rb') as model:\n        if not parser.parse(model.read()):\n            print(\"ERROR: Failed to parse the ONNX file.\")\n            for error in range(parser.num_errors):\n                print(parser.get_error(error))\n            return None\n\n    # 5. Build and Serialize Engine\n    print(\"Building TensorRT Engine... (This may take a few minutes)\")\n    serialized_engine = builder.build_serialized_network(network, config)\n    \n    if serialized_engine:\n        with open(engine_file_path, \"wb\") as f:\n            f.write(serialized_engine)\n        print(f\"Engine saved to {engine_file_path}\")\n    else:\n        print(\"Failed to build engine.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T01:12:28.720627Z","iopub.execute_input":"2026-02-08T01:12:28.721113Z","iopub.status.idle":"2026-02-08T01:12:28.727617Z","shell.execute_reply.started":"2026-02-08T01:12:28.721085Z","shell.execute_reply":"2026-02-08T01:12:28.726922Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"onnx_models = [\n    ('best-epoch=264-val_loss=0.3657-val_dice=0.5852',\n     '/kaggle/input/unet3d-1st-stage-vesuvius-challenge/pytorch/default/13/trt_models/trt_models/primus-best-epoch=264-val_loss=0.3657-val_dice=0.5852.onnx'),\n    ('best-epoch=294-val_loss=0.3653-val_dice=0.5917',\n    '/kaggle/input/unet3d-1st-stage-vesuvius-challenge/pytorch/default/13/trt_models/trt_models/primus-best-epoch=294-val_loss=0.3653-val_dice=0.5917.onnx')\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T01:12:28.728618Z","iopub.execute_input":"2026-02-08T01:12:28.728882Z","iopub.status.idle":"2026-02-08T01:12:28.791203Z","shell.execute_reply.started":"2026-02-08T01:12:28.728857Z","shell.execute_reply":"2026-02-08T01:12:28.790576Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for (name, ONNX_FILE_PATH) in onnx_models:\n    # # 1. Configuration\n    # ONNX_FILE_PATH = f\"{name}.onnx\"\n    # INPUT_SHAPE = (1, 1, 160, 160, 160) # (Batch, Channel, D, H, W)\n    # model.eval()\n    \n    # # Move to CPU for export (usually sufficient and avoids VRAM issues during export)\n    # model.cpu()\n    \n    # # 3. Create Dummy Input\n    # dummy_input = torch.randn(INPUT_SHAPE, device='cpu')\n    \n    # # 4. Export\n    # print(f\"Exporting model to {ONNX_FILE_PATH}...\")\n    # torch.onnx.export(\n    #     model,\n    #     dummy_input,\n    #     ONNX_FILE_PATH,\n    #     export_params=True,\n    #     opset_version=13,          # Opset 13+ is recommended for 3D Convs\n    #     do_constant_folding=True,\n    #     input_names=['input'],\n    #     output_names=['output'],\n    #     # Since your input is fixed, we do NOT set dynamic_axes.\n    #     # This allows TRT to aggressively optimize for this specific size.\n    # )\n    # print(\"Export complete.\")\n    \n    # ONNX_FILE_PATH = f\"primus-{name}.onnx\"\n    ENGINE_FILE_PATH = f\"primus-{name}.engine\"\n    build_engine(ONNX_FILE_PATH, ENGINE_FILE_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-08T01:12:48.680158Z","iopub.execute_input":"2026-02-08T01:12:48.680472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}