Skip to content

About

Histopathologic Cancer Detection using ResNet-18 features and an Attention mechanism

Resources

Stars

0 stars

Watchers

1 watching

Forks

Repository files navigation

Histopathologic Cancer Detection using Attention Mechanism

Deep learning model for binary classification of histopathologic images to detect metastatic cancer in lymph node sections using ResNet18 with attention mechanism

Dataset

Source: HCD-Cropped Dataset on Kaggle

Dataset Characteristics

  • Total Images: 220,025 histopathologic images
  • Format: TIF images (32×32×3 pixels)
  • Classes: Binary classification
    • Class 0: No metastatic tissue
    • Class 1: Metastatic tissue present
  • Image Resolution: 32×32 pixels (upscaled to 224×224 for ResNet18 input)
  • Channels: 3 (RGB)
  • Split: 80% training, 10% validation, 10% test

The dataset contains cropped patches from whole-slide images of lymph node sections, where the task is to identify metastatic tissue

Model Architecture

Attention-Enhanced ResNet18

The model uses a ResNet18 backbone pretrained on ImageNet with a custom attention module for interpretable feature learning

Architecture Details

  • Backbone: ResNet18 (pretrained on ImageNet1K)
  • Feature Extractor: All ResNet18 layers except final FC layer
  • Feature Dimension: 512 → 16 (compressed representation)
  • Attention Mechanism: Tanh-based attention with softmax normalization
  • Dropout: 0.5 for regularization
  • Output: 2 classes (binary classification)
class AttentionModule(nn.Module):
    def __init__(self):
        super(AttentionModule, self).__init__()
        original_model = models.resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)
        self.features = nn.Sequential(*list(original_model.children())[:-1])
        self.fc1 = nn.Linear(512, 16)
        self.dropout = nn.Dropout(0.5)
        self.attention = nn.Sequential(
            nn.Linear(16, 16),
            nn.Tanh(),
            nn.Linear(16, 1),
            nn.Softmax(dim=1)
        )
        self.classifier = nn.Linear(16, 2)

Total Parameters: 11,185,043 (all trainable)

Data Preprocessing

  • Resize: 32×32 → 224×224 pixels
  • Normalization: ImageNet statistics (mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
  • Tensor conversion with automatic augmentation support

Training Configuration

  • Optimizer: Adam (lr=0.001)
  • Loss Function: CrossEntropyLoss
  • Batch Size: 32
  • Epochs: 6
  • Device: CUDA (GPU acceleration)

Performance

Training Progress

Epoch Train Loss Train Acc Val Loss Val Acc
1 0.4271 81.03% 0.4236 81.07%
2 0.3761 83.84% 0.3545 84.75%
3 0.3456 85.43% 0.3393 85.24%
4 0.3235 86.48% 0.3150 86.45%
5 0.3051 87.35% 0.3158 86.66%
6 0.2866 88.21% 0.3242 86.25%

Final Test Performance

  • Test Loss: 0.3130
  • Test Accuracy: 86.71%

Installation

# Clone the repository
git clone https://github.com/yourusername/histopathologic-cancer-detection.git
cd histopathologic-cancer-detection

# Install dependencies
pip install torch torchvision torchinfo pandas pillow rasterio scikit-learn

Usage

Training

# Load dataset
cancer_dataset = CancerDataset(image_dir=image_dir, labels_df=labels_df, transform=transform)

# Split data
train_dataset, val_dataset, test_dataset = torch.utils.data.random_split(
    cancer_dataset, [0.8, 0.1, 0.1]
)

# Initialize model
model = AttentionModule().to(device)
optimizer = optim.Adam(model.parameters(), lr=0.001)
criterion = nn.CrossEntropyLoss()

# Train
for epoch in range(num_epochs):
    train_loss, train_acc = train_model(model, train_loader, device, optimizer, criterion)
    val_loss, val_acc = validate_model(model, val_loader, device, criterion)

Inference

# Load trained weights
model.load_state_dict(torch.load('6-epoch-cancer_detection_weights.pth'))
model.eval()

# Get predictions with attention weights
outputs, attention_weights = model(images, return_attention=True)

Key Features

  • ✅ Attention mechanism for interpretable predictions
  • ✅ Transfer learning with ImageNet-pretrained ResNet18
  • ✅ Efficient training on 220K+ images
  • ✅ 86.71% test accuracy
  • ✅ Attention weight visualization support

Requirements

torch>=2.0.0
torchvision>=0.15.0
torchinfo
pandas
pillow
rasterio
scikit-learn

Acknowledgments

About

Histopathologic Cancer Detection using ResNet-18 features and an Attention mechanism

Resources

Stars

0 stars

Watchers

1 watching

Forks

Releases

Packages

Contributors

Languages