Brain Tumor Classifier in Python - Machine Learning Project

NeuralNineAbout 7 min readJul 29, 2025Watch original
THE SUMMARYAI-generated

Key Concepts

Convolutional Neural Networks (CNNs), PyTorch, Image Classification, Brain Tumor MRI Data, Data Loaders, Transformations, Model Definition, Optimizer, Loss Function, Training Loop, Evaluation, Accuracy, Visualization.

1. Data Preparation and Loading

  • Data Set: The video utilizes the "Brain Tumor MRI Dataset" from Kaggle, containing MRI scans categorized into four classes: "no tumor" (healthy brains), "glioma," "meningioma," and "pituitary tumor."
  • Directory Structure: The dataset is organized into training and testing directories, each containing subdirectories representing the four tumor categories.
  • Data Loading:
    • torchvision.datasets.ImageFolder is used to create datasets from the image folders, treating subdirectories as classes.
    • Two data loaders are created: train_DL for training data and test_DL for testing data.
    • batch_size = 32: Batches of 32 images are used during training.
    • shuffle = True for the training data loader to introduce randomness during training.
    • num_workers = 4 for faster data loading.
    • pin_memory = True to speed up the transfer between CPU and GPU.
  • Transformations (TF): A series of transformations are applied to the images:
    • transforms.Resize((128, 128)): Resizes images to 128x128 pixels.
    • transforms.ToTensor(): Converts images to PyTorch tensors.
    • transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)): Normalizes the tensor data using mean and standard deviation of 0.5 for each channel (RGB). transforms.Compose combines the transformations.

2. Model Definition (CNN)

  • Model Architecture: A convolutional neural network (CNN) is built using nn.Sequential (TensorFlow/Keras-like syntax).
  • Layers:
    • nn.Conv2d: Convolutional layers with 3x3 kernels, stride of 1, and padding of 1. The number of output channels/filters increases progressively (3 -> 32 -> 64 -> 128).
    • nn.ReLU: Rectified Linear Unit activation function to introduce non-linearity.
    • nn.MaxPool2d(2): Max pooling layers with a kernel size of 2 for downsampling and regularization.
    • nn.Flatten(): Flattens the output of the convolutional layers into a 1D tensor.
    • nn.Linear: Fully connected (dense) layers. The first linear layer maps the flattened convolutional features (128 * 16 * 16) to 256 neurons, and the final layer maps 256 neurons to 4 output classes.
    • nn.Dropout(0.5): Dropout layer with a 50% dropout rate for regularization.
  • Example Layer Configuration:
    • nn.Conv2d(3, 32, 3, 1, 1): Input 3 channels (RGB), output 32 channels, kernel size 3x3, stride 1, padding 1.
    • nn.Linear(128 * 16 * 16, 256): Input 128 * 16 * 16 features, output 256 neurons.
    • nn.Linear(256, 4): Input 256 neurons, output 4 classes.
  • Device Placement: The model is moved to the GPU (if available) using model.to(device).

3. Training Process

  • Optimizer: AdamW optimizer (optim.AdamW) is used with a learning rate of 1e-4 (0.0001) and weight decay.
  • Loss Function: Cross-entropy loss (nn.CrossEntropyLoss()) is used for multi-class classification.
  • Training Loop:
    • The model is trained for 25 epochs.
    • For each epoch:
      • The running loss is initialized to zero.
      • For each batch (X, Y) from the training data loader:
        • Optimizer gradients are zeroed using opt.zero_grad().
        • The input batch (X) and labels (Y) are moved to the device (GPU if available).
        • The model's output (logits) is computed by applying the model to the input batch: model(X).
        • The loss is calculated by comparing the model's output to the true labels: loss_fn(model(X), Y).
        • The gradients are computed using backpropagation: loss.backward().
        • The optimizer updates the model's parameters: opt.step().
        • The loss for the batch is added to the running loss.
      • The epoch number and running loss are printed.
  • model.train(): Sets the model to training mode (though the presenter mentions it might not be strictly necessary in this case).

4. Evaluation and Results

  • Evaluation Mode: The model is set to evaluation mode using model.eval().
  • No Gradient Calculation: The evaluation is performed within a with torch.no_grad(): block to disable gradient calculation and save memory.
  • Evaluation Loop:
    • Test loss and the number of correct predictions are initialized to zero.
    • For each batch (X, Y) from the test data loader:
      • The input batch (X) and labels (Y) are moved to the device.
      • The model's output (logits) is computed.
      • The loss is calculated and accumulated.
      • Predictions are obtained by taking the argmax of the logits along dimension 1 (logits.argmax(dim=1)).
      • The number of correct predictions is calculated by comparing the predictions to the true labels and summing the correct matches.
    • The test loss is divided by the length of the test data set to get the average loss.
    • The accuracy is calculated as the percentage of correct predictions.
  • Performance: The model achieves a test accuracy of 96.49% and a test loss of 0.12.
  • Baseline: The presenter notes that with four classes, a random guess would yield approximately 25% accuracy.

5. Visualization

  • Random Image Visualization: A random image from the dataset is visualized along with the model's prediction and the true label.
  • Process:
    • A random index is selected.
    • The image and label at that index are retrieved from the dataset.
    • The image is unnormalized using torchvision.transforms.functional.to_pil_image and the inverse of the normalization parameters (mean and std).
    • The image is displayed using matplotlib.pyplot.
    • The model predicts the class of the image, and the prediction is printed along with the true label.
  • Purpose: To visually inspect the model's performance on individual images.

6. Libraries and Tools

  • PyTorch: Deep learning framework.
  • Torchvision: Provides datasets, model architectures, and image transformations.
  • Matplotlib: Used for visualization.
  • UV (Optional): Package manager (alternative to pip).
  • Jupyter Lab (Recommended): Interactive development environment for running code in cells.

7. Technical Terms and Concepts

  • Convolutional Neural Network (CNN): A type of neural network designed for processing data with a grid-like topology, such as images.
  • Kernel/Filter: A small matrix that slides over the input image to extract features.
  • Channels: Color components of an image (e.g., RGB).
  • Stride: The step size of the kernel as it slides over the image.
  • Padding: Adding extra pixels around the border of an image to control the output size of convolutional layers.
  • ReLU (Rectified Linear Unit): An activation function that introduces non-linearity.
  • Max Pooling: A downsampling operation that reduces the spatial dimensions of the feature maps.
  • Flatten: Converting a multi-dimensional tensor into a 1D tensor.
  • Dense/Linear Layer: A fully connected layer where each neuron is connected to all neurons in the previous layer.
  • Dropout: A regularization technique that randomly sets a fraction of neurons to zero during training.
  • Softmax: An activation function that outputs a probability distribution over multiple classes.
  • Logits: The raw, unnormalized output of a neural network before applying a softmax or sigmoid function.
  • Optimizer: An algorithm that updates the model's parameters during training to minimize the loss function (e.g., AdamW).
  • Learning Rate: A hyperparameter that controls the step size of the optimizer.
  • Weight Decay: A regularization technique that penalizes large weights to prevent overfitting.
  • Loss Function: A function that measures the difference between the model's predictions and the true labels (e.g., cross-entropy loss).
  • Epoch: One complete pass through the entire training data set.
  • Batch: A subset of the training data used in one iteration of the training loop.
  • Backpropagation: An algorithm for computing the gradients of the loss function with respect to the model's parameters.
  • Gradient: The rate of change of the loss function with respect to a parameter.
  • Accuracy: The percentage of correctly classified samples.
  • Overfitting: When a model learns the training data too well and performs poorly on unseen data.
  • Regularization: Techniques used to prevent overfitting.
  • Hyperparameter Tuning: The process of finding the optimal values for hyperparameters (e.g., learning rate, batch size, number of layers).
  • Learning Rate Scheduling: Adjusting the learning rate during training to improve performance.
  • Model Checkpointing: Saving the model's parameters at regular intervals during training.

8. Logical Connections and Flow

  • The video starts with the project's goal: to train a CNN to classify brain tumor MRI images.
  • It then explains the data set, its structure, and how to load it using PyTorch's data loading utilities.
  • Data transformations are defined to prepare the images for training.
  • A CNN model is built using nn.Sequential, and its layers are explained.
  • The training process is detailed, including the optimizer, loss function, and training loop.
  • The trained model is evaluated on the test data set to assess its performance.
  • Finally, the video shows how to visualize random images and the model's predictions.
  • The video concludes with suggestions for further improvement, such as hyperparameter tuning and different model architectures.

9. Synthesis/Conclusion

The video provides a practical guide to building and training a CNN in PyTorch for brain tumor classification. It covers all the essential steps, from data preparation and loading to model definition, training, evaluation, and visualization. The model achieves high accuracy (96.49%), demonstrating the effectiveness of CNNs for this type of image classification task. The presenter encourages viewers to experiment with different hyperparameters and architectures to further improve the results. The tutorial bridges the gap between theoretical knowledge and practical application, offering a solid foundation for beginners and intermediate programmers interested in deep learning and medical image analysis.

AI summaries can miss context or contain errors. Check important details against the original video.

Go a little deeper.

Have a question about this video? Load its transcript to open the video chat.