Downloads Β· 30 days
0
Kiuyha/PromptM-UNet
PromptM-UNet is a image segmentation model from Kiuyha. Use it for the image segmentation task on the model card, and read the license before you ship it in a product. It is set up for promptm-unet. The card lists the license as apache-2.0.
[](https://github.com/kiuyha/PromptM-UNet) [](https://opensource.org/licenses/Apache-2.0)
Downloads Β· 30 days
0
Access
Public
Updated Sep 6, 2026
Repo size
990 MB
Likes
0
Public
Click a slice to open those files.
.pth1.3 GB Β· 100%
From the Hugging Face model README
PromptM-UNet is an ultra-lightweight, multi-modal 3D medical image segmentation framework that integrates Residual Vision Mamba (SSM) with multi-stage language conditioning. It enables accurate, prompt-guided volumetric CT segmentation while maintaining linear (O(N)) computational complexity and a minimal memory footprint feasible for standard consumer GPUs.
The following results are evaluated on the official TotalSegmentator spleen validation cohort across 30 training epochs (from notebooks/ablation_promptm_unet.ipynb):
[!NOTE] Parameter Accounting & Model Weights: The checkpoint files (
.pth) store only the trainable visual backbone, stage projection layers, squeeze convolutions, and deep supervision heads (~5.03M to ~5.44M parameters). Frozen language encoder weights are decoupled and loaded dynamically:
- CLIP (
openai/clip-vit-base-patch32): 63.17M parameters (Default text encoder used for Ablations 1β3)- BioBERT (
emilyalsentzer/Bio_ClinicalBERT): 108.31M parameters- Sentence-BERT (SBERT) (
sentence-transformers/all-MiniLM-L6-v2): 22.71M parametersThe Total Parameters column indicates the full end-to-end model (Visual Backbone + Fusion Heads + Text Encoder).
Evaluated with default CLIP text encoder (63.17M), and All-Stage Fusion:
| Operation | Trainable Params | Total Params | Peak Epoch | Peak DSC | Peak NSD |
|---|---|---|---|---|---|
| Multiplication | 5.26M | 68.43M | 24 | 0.8619 Β± 0.0028 | 0.7871 Β± 0.0041 |
| Concatenation | 5.44M | 68.60M | 26 | 0.8494 Β± 0.0010 | 0.7528 Β± 0.0021 |
| Hybrid | 5.44M | 68.60M | 28 | 0.8152 Β± 0.0014 | 0.7301 Β± 0.0025 |
Evaluated with Multiplication operation, and default CLIP text encoder (63.17M):
| Fusion Strategy | Trainable Params | Text Encoder | Total Params | Peak Epoch | Peak DSC | Peak NSD | Peak Val VRAM |
|---|---|---|---|---|---|---|---|
| Early Fusion | 5.15M | CLIP (63.17M) | 68.31M | 19 | 0.8622 Β± 0.0012 | 0.7928 Β± 0.0019 | < 4.3 GB |
| All-Stage Fusion | 5.26M | CLIP (63.17M) | 68.43M | 24 | 0.8619 Β± 0.0033 | 0.7871 Β± 0.0047 | < 4.3 GB |
| Late Fusion | 5.03M | CLIP (63.17M) | 68.20M | 12 | 0.8576 Β± 0.0008 | 0.7675 Β± 0.0013 | < 4.3 GB |
Evaluated with Multiplication operation, and Early Fusion:
| Text Encoder | Pretrained Model | Text Dim | Text Encoder Params | Trainable Params | Total Params | Peak Epoch | Peak DSC | Peak NSD |
|---|---|---|---|---|---|---|---|---|
| Sentence-BERT (SBERT) | all-MiniLM-L6-v2 | 384D | 22.71M | 5.11M | 27.83M | 24 | 0.8702 Β± 0.0009 | 0.7896 Β± 0.0019 (0.7938 Β± 0.0020*) |
| CLIP | clip-vit-base-patch32 | 512D | 63.17M | 5.15M | 68.31M | 19 | 0.8622 Β± 0.0012 | 0.7928 Β± 0.0019 |
| BioBERT | Bio_ClinicalBERT | 768D | 108.31M | 5.21M | 113.52M | 19 | 0.8466 Β± 0.0016 | 0.7584 Β± 0.0024 |
*Note: SBERT attained peak DSC of 0.8702 Β± 0.0009 at Epoch 24, and peak NSD of 0.7938 Β± 0.0020 at Epoch 26.
Evaluated with Multiplication operation, Early Fusion and Sentence-BERT:
| Init Type | Trainable Params | Total Params | Peak Epoch | Peak DSC | Peak NSD |
|---|---|---|---|---|---|
| Random Init | 5.11M | 27.83M | 24 | 0.8702 Β± 0.0009 | 0.7896 Β± 0.0019 |
| Pretrained | 5.11M | 27.83M | 25 | 0.8667 Β± 0.0007 | 0.8096 Β± 0.0015 |
# Clone the repository
git clone https://github.com/kiuyha/PromptM-UNet.git
cd PromptM-UNet
# Install dependencies and package
pip install -e .
import torch
import yaml
from promptm_unet.models.PromptMUNet import PromptMUNet
# Load configuration
with open("configs/default.yml", "r") as f:
config = yaml.safe_load(f)
# Instantiate model
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = PromptMUNet(config).to(device)
# Load checkpoint
checkpoint = torch.load("best_model.pth", map_location=device)
model.load_state_dict(checkpoint["model_state_dict"])
model.eval()
# Dummy 3D input: (Batch, Channels, Depth, Height, Width)
dummy_ct = torch.randn(1, 1, 64, 128, 128).to(device)
prompt = ["spleen segmentation in abdominal CT scan"]
# Forward pass
with torch.no_grad():
prediction = model(dummy_ct, prompt) # Output: (1, 1, 64, 128, 128)
probabilities = torch.sigmoid(prediction)
binary_mask = (probabilities > 0.5).cpu().numpy()
print(f"Predicted spleen mask shape: {binary_mask.shape}")
accelerate launch -m promptm_unet.cli train \
--path.raw_data_dir "/path/to/totalsegmentator_raw" \
--path.prepro_data_dir "/path/to/preprocessed_data" \
--training.batch_size 2 \
--training.epochs 30
# Run multi-tier evaluation
python -m promptm_unet.cli test \
--checkpoint "./experiments/checkpoints/best_model.pth"
3D CT Volume (1.5mm / 3.0mm)
β
βΌ
βββββββββββββββββ
β Visual Encoderβ (Residual Vision Mamba - RVM Blocks)
βββββββββ¬ββββββββ
β
βΌ
βββββββββββββββββ ββββββββββββββββββββ
β Bottleneck β ββββββΊ β Text Projection β βββ Frozen CLIP / BioBERT / SBERT
βββββββββ¬ββββββββ ββββββββββββββββββββ (512D / 768D / 384D)
β β²
βΌ β (Multi-Stage Hadamard Product)
βββββββββββββββββ β
β Visual Decoderβ βββββββββββββββββ
βββββββββ¬ββββββββ
β
βΌ
Deep Supervision Heads βββΊ High-Resolution Spleen Binary Mask
This project is licensed under the Apache License 2.0.
If you find this work useful in your research, please cite our repository:
@misc{promptm_unet2026,
title={PromptM-UNet: Efficient Text-Prompted 3D Medical Image Segmentation Using State-Space Model Mamba for Spleen},
author={Shridhara, Ketut and Setyawan, Ivan Andika and Al-Habib, Hasanuddin},
year={2026},
publisher={Universitas Negeri Surabaya},
howpublished={\url{https://github.com/kiuyha/PromptM-UNet}}
}