brain-tumor-detect / README.md
aadium's picture
Update README.md
d244641 verified
|
Raw
History Blame Contribute Delete
2.56 kB
metadata
license: mit
library_name: torchvision
tags:
  - medical
  - mri
  - object-detection
  - brain-tumor
  - faster-rcnn
  - pytorch
datasets:
  - neuron-m1yxd/brain-tumor-ppo4z
metrics:
  - map
pipeline_tag: object-detection

Brain Tumor Object Detection (Faster R-CNN)

An end-to-end object detection model using Faster R-CNN (ResNet-50-FPN) trained to detect, localize, and classify brain tumors from MRI scans.

  • Model Architecture: Faster R-CNN with ResNet-50-FPN backbone
  • Framework: PyTorch & TorchVision
  • Input: COCO-formatted brain MRI images
  • Classes (3): glioma, meningioma, pituitary

Dataset

The model was trained using the Brain Tumor Dataset on Roboflow Universe:

  • Classes: glioma (1), meningioma (2), pituitary (3)
  • Format: COCO JSON annotations
  • Augmentations: Horizontal Flip, Random Brightness/Contrast, Color Jitter via Albumentations

Quickstart: Python Inference

You can run inference using standard torchvision and PIL:

import torch
import torchvision
from torchvision.transforms import functional as F
from PIL import Image

# 1. Load Model Architecture
NUM_CLASSES = 4  # Background + 3 tumor classes
model = torchvision.models.detection.fasterrcnn_resnet50_fpn(
    weights=None, 
    min_size=800, 
    max_size=1333
)
in_features = model.roi_heads.box_predictor.cls_score.in_features
model.roi_heads.box_predictor = torchvision.models.detection.faster_rcnn.FastRCNNPredictor(in_features, NUM_CLASSES)

# 2. Download and Load Checkpoint from Hugging Face Hub
from huggingface_hub import hf_hub_download

checkpoint_path = hf_hub_download(repo_id="YOUR_HF_USERNAME/YOUR_MODEL_REPO", filename="fasterrcnn_best.pth")
state_dict = torch.load(checkpoint_path, map_location="cpu", weights_only=True)
model.load_state_dict(state_dict)

device = "cuda" if torch.cuda.is_available() else "cpu"
model.to(device)
model.eval()

# 3. Predict on an Image
CLASS_NAMES = {1: "glioma", 2: "meningioma", 3: "pituitary"}
img = Image.open("sample_mri.jpg").convert("RGB")
img_tensor = F.to_tensor(img).unsqueeze(0).to(device)

THRESHOLD = 0.6
with torch.no_grad():
    prediction = model(img_tensor)[0]

for i in range(len(prediction["boxes"])):
    score = prediction["scores"][i].item()
    if score > THRESHOLD:
        box = prediction["boxes"][i].cpu().numpy()
        label_id = prediction["labels"][i].item()
        print(f"Detected {CLASS_NAMES.get(label_id)} with {score:.2f} confidence at {box}")