Downloads · 30 days
0
callumtilbury/bubble-distill
bubble-distill is a machine learning model from callumtilbury. Use it for the machine learning task on the model card, and read the license before you ship it in a product.
Cellpose-SAM-FT → Pseudo-labels → TinyBubbleNet
Downloads · 30 days
0
Access
Public
Updated Apr 24, 2026
Repo size
—
Likes
0
Public
Click a slice to open those files.
.py103 KB · 91%
From the Hugging Face model README
Cellpose-SAM-FT → Pseudo-labels → TinyBubbleNet
A 3-stage pipeline for fast, lightweight microbubble sizing and counting via knowledge distillation.
⚠️ IMPORTANT BUG FIX: The original
train_student.pyin this repo usesBCEWithLogitsLosson binary foreground/background masks. This fails catastrophically because microbubble foreground is only ~0.2% of pixels — the model learns to predict ALL background and achieves 99.8% accuracy while detecting zero bubbles. The fixed scripttrain_mse_distill.pyuses MSE distillation on the teacher's raw cell_prob LOGITS instead, which gives gradients on ALL pixels (background pixels have informative negative logits ~-6). Seetrain_mse_distill.pyfor the corrected implementation.
Cellpose-SAM is excellent for cell/bubble segmentation, but at ~300M params (1.1 GB) it's expensive at inference. For lab settings where your slides look similar and you're "just detecting circles", this is massive overkill. You're paying for the ability to also segment dogs, neurons, and a thousand other things — capacity you don't need.
| Model | Params | Size | 256×256 GPU | 256×256 CPU | FPS (GPU) |
|---|---|---|---|---|---|
| Cellpose-SAM | ~300M | 1.1 GB | ~100 ms | seconds | ~10 |
| TinyBubbleNet (base_ch=16) | 389K | 1.5 MB | 3 ms | 45 ms | 337 |
| TinyBubbleNet (base_ch=32) | 1.5M | 5.8 MB | ~5 ms | ~80 ms | ~200 |
~33× faster, ~750× smaller. And when your domain is narrow (similar-looking lab slides), the accuracy loss is minimal because the student only needs to learn one visual distribution.
TinyBubbleNet is a depthwise-separable U-Net (inspired by PicoSAM2) with a 4-channel output:
| Channel | Name | What it encodes |
|---|---|---|
| 0 | dY | Vertical gradient flow (Cellpose-compatible) |
| 1 | dX | Horizontal gradient flow (Cellpose-compatible) |
| 2 | cell_prob | Foreground/background probability |
| 3 | dist_transform | Distance transform (peak = bubble radius) |
Instance masks are reconstructed via Euler integration of the flow field — identical to Cellpose post-processing. This means the student is fully compatible with the Cellpose ecosystem.
The distance transform head is the key addition for sizing: the peak value within each detected instance directly gives you the bubble radius.
train_student.py / losses.py)# BAD: BCE on binary masks
prob_loss = BCEWithLogitsLoss(pred_prob, binary_mask)
With foreground at only ~0.2% of pixels, the model's dominant gradient signal is "predict all background". Even after 300 epochs with "best val loss 0.0008", the model predicts zero bubbles everywhere.
train_mse_distill.py)# GOOD: MSE on teacher's raw logits
prob_loss = MSE(pred_prob_logits, teacher_cell_prob_logits)
The teacher outputs cell_prob as raw logits (range roughly -9 to +5). Every pixel has an informative value — background pixels should reproduce ~-6, foreground pixels should reproduce ~+5. MSE on logits gives strong gradients everywhere, and the student successfully learns to segment bubbles.
| Loss | Target | Gradient on bg pixels? | Result |
|---|---|---|---|
| BCE + binary mask | {0, 1} | NO (bg is "correct" at 0) | Predicts all background |
| MSE + teacher logits | Real numbers (~-6 to +5) | YES (bg must match ~-6) | Learns proper segmentation |
┌─────────────────────┐ ┌──────────────────────────┐ ┌─────────────────────┐
│ Stage 1: Teacher │ │ Stage 2: Distillation │ │ Stage 3: Inference │
│ │ │ │ │ │
│ Cellpose-SAM-FT │────▶│ Train TinyBubbleNet │────▶│ Fast bubble sizing │
│ generates pseudo- │ │ on raw teacher logits │ │ (~3ms/image GPU) │
│ labels on 100s of │ │ (~400 epochs) │ │ │
│ lab images │ │ │ │ │
└─────────────────────┘ └──────────────────────────┘ └─────────────────────┘
pip install cellpose torch torchvision scipy scikit-image huggingface_hub numpy
python generate_pseudolabels.py \
--image_dir /path/to/lab_images/ \
--model_path /path/to/your/cellpose_sam_ft_model \
--output_dir /path/to/pseudolabels/ \
--diameter 30 \
--channels 0 0
This runs your fine-tuned Cellpose-SAM on all images and saves:
.npy)# ✅ FIXED: Uses MSE distillation on teacher logits
python train_mse_distill.py \
--image_dir /path/to/lab_images/ \
--label_dir /path/to/pseudolabels/ \
--output_dir ./checkpoints/ \
--base_ch 16 \
--epochs 400 \
--batch_size 4 \
--lr 1e-3
# ❌ DEPRECATED (has class imbalance bug)
# python train_student.py ...
The fixed script (train_mse_distill.py):
RawPseudoLabelDataset that loads teacher logits directly (no binarization)python inference.py \
--model_path ./checkpoints/best_model.pt \
--image_path /path/to/image.png \
--output_dir ./results/
| File | Description | Status |
|---|---|---|
model.py | TinyBubbleNet architecture (depthwise-separable U-Net) | ✅ |
losses.py | Original distillation loss (BCE+Dice — has bug) | ⚠️ See train_mse_distill.py for fix |
dataset.py | Original dataset (binary masks — has bug) | ⚠️ See train_mse_distill.py for fix |
train_student.py | Original training (BCE-based — has bug) | ⚠️ Deprecated |
train_mse_distill.py | Fixed training with MSE on teacher logits | ✅ Use this! |
generate_pseudolabels.py | Stage 1: Teacher → pseudo-labels | ✅ |
inference.py | Stage 3: Fast inference + bubble measurements | ✅ |
base_ch | Params | Size | GPU Speed | Use Case |
|---|---|---|---|---|
| 16 | 389K | 1.5 MB | 3 ms @ 256² | Default — fast & tiny |
| 32 | 1.5M | 5.8 MB | 5 ms @ 256² | More capacity if needed |
Use --no_depthwise for standard convolutions (more params, possibly better accuracy on complex images).
Why Cellpose flows instead of direct mask prediction? Flows handle overlapping/touching bubbles via convergence — each pixel flows toward its instance center. Direct mask prediction can't separate touching instances.
Why distance transform head? For circles, the DT peak = radius. This gives you sizing "for free" without post-processing the mask.
Why depthwise-separable convs? ~8× fewer params than standard convs. For a narrow domain (your lab slides), this compression is lossless.
Why MSE on logits instead of BCE on masks? See "The Bug and The Fix" section above. BCE on sparse binary masks fails due to extreme class imbalance. MSE on teacher logits gives gradients everywhere.
The student is specialized to your current lab setup. Re-train when:
Re-training is fast: ~30 min for 400 epochs on 50 images with a GPU.