Step-by-step guides from beginner to advanced topics
Set up your environment and train your first neural network for image classification.
Learn how to prepare and preprocess semiconductor and medical imaging datasets.
Build and train convolutional neural networks with attention mechanisms.
Implement CBAM, SE blocks, and spatial attention for enhanced performance.
Understand and implement Vision Transformers for image classification.
Combine multiple models for robust predictions with uncertainty quantification.
First, set up your Python environment with the required packages.
# Create virtual environment python -m venv nn_env source nn_env/bin/activate # On Windows: nn_env\Scripts\activate # Install PyTorch (adjust for your CUDA version) pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 # Install other dependencies pip install numpy pandas matplotlib seaborn scikit-learn wandb
Load and explore the semiconductor defect dataset.
import torch
from torch.utils.data import DataLoader
from nn_classification.data import SemiconductorDataset
from nn_classification.data.augmentation import get_augmentation_pipeline
# Create dataset
train_dataset = SemiconductorDataset(
root_dir='data/semiconductor',
split='train',
transform=get_augmentation_pipeline('train')
)
# Create dataloader
train_loader = DataLoader(
train_dataset,
batch_size=32,
shuffle=True,
num_workers=4
)
# Check dataset
print(f"Number of samples: {len(train_dataset)}")
print(f"Number of classes: {train_dataset.num_classes}")
print(f"Class names: {train_dataset.classes}")
Initialize a CNN with attention mechanisms.
from nn_classification.models import CNNWithAttention
# Create model
model = CNNWithAttention(
num_classes=8,
in_channels=1, # Grayscale images
attention_type='cbam',
dropout_rate=0.5
)
# Move to GPU if available
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = model.to(device)
# Check model
total_params = sum(p.numel() for p in model.parameters())
trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f"Total parameters: {total_params:,}")
print(f"Trainable parameters: {trainable_params:,}")
Set up training loop with optimizer and loss function.
from nn_classification.training import Trainer
# Configure training
config = {
'learning_rate': 1e-3,
'epochs': 10,
'optimizer': 'adamw',
'scheduler': 'cosine',
'mixed_precision': True,
'save_best': True
}
# Create trainer
trainer = Trainer(
model=model,
train_loader=train_loader,
val_loader=val_loader,
config=config
)
# Start training
trainer.train()
# Evaluate
results = trainer.evaluate(test_loader)
print(f"Test Accuracy: {results['accuracy']:.2%}")
print(f"Test F1 Score: {results['f1_score']:.3f}")
You've successfully trained your first neural network for image classification! Next, try the intermediate tutorials to learn about attention mechanisms and advanced techniques.
Convolutional Block Attention Module (CBAM) combines channel and spatial attention.
from nn_classification.models.attention import CBAM
class ResBlockWithCBAM(nn.Module):
def __init__(self, channels):
super().__init__()
self.conv1 = nn.Conv2d(channels, channels, 3, padding=1)
self.conv2 = nn.Conv2d(channels, channels, 3, padding=1)
self.cbam = CBAM(channels, reduction_ratio=16)
self.relu = nn.ReLU()
def forward(self, x):
residual = x
out = self.relu(self.conv1(x))
out = self.conv2(out)
out = self.cbam(out) # Apply attention
out += residual
return self.relu(out)
Extract and visualize where the model is focusing.
import matplotlib.pyplot as plt
from nn_classification.utils.visualization import plot_attention_maps
# Get attention maps
model.eval()
with torch.no_grad():
output = model(input_image)
attention_maps = model.get_attention_maps()
# Visualize
fig, axes = plt.subplots(2, 3, figsize=(12, 8))
for i, (layer_name, attn_map) in enumerate(attention_maps.items()):
if i >= 6:
break
ax = axes[i // 3, i % 3]
im = ax.imshow(attn_map.squeeze().cpu(), cmap='hot')
ax.set_title(f'{layer_name} Attention')
ax.axis('off')
plt.colorbar(im, ax=ax)
plt.tight_layout()
plt.show()