0%
SustainSys AI Academy · Student Guide

PyTorch from Scratch
A Complete Beginner's Guide

A fully systematic, beginner-friendly guide to understanding and using PyTorch — from your very first tensor to a complete trained neural network. No assumed knowledge. No rushed explanations. Every concept explained before it's used.

Python 3.10+ PyTorch 2.x Tensors Autograd nn.Module DataLoader 12 Lessons

What is PyTorch?

PyTorch is a Python library for building AI systems. Specifically, it makes it easy to build and train neural networks — mathematical models that learn patterns from data.

Think of it this way: teaching a computer to recognise cats means showing it thousands of photos, letting it make guesses, telling it how wrong it was, and letting it adjust. PyTorch is the toolkit that makes this entire process efficient and manageable.

Without PyTorch
  • Write complex maths by hand
  • Training takes weeks or longer
  • Gradients calculated manually
  • No GPU support built in
  • Thousands of lines of boilerplate
With PyTorch
  • Maths handled automatically
  • GPU acceleration with one line
  • Gradients computed automatically
  • Clean, readable Python code
  • Used by OpenAI, Meta, Tesla, Google

Four Things People Confuse

Python

The programming language. PyTorch is written in Python — it's a library that runs inside it.

Machine Learning

The broad idea — teaching computers to learn from data. PyTorch is a tool to implement this idea.

Neural Network

A specific ML technique inspired by the brain. PyTorch is especially good at building these.

PyTorch

A software library — pre-written code you use to build and train neural networks efficiently.

Where PyTorch Fits in the AI Workflow

1
Collect Data

Photos, text, numbers — whatever your problem requires.

2
Prepare Data

Clean it, format it, split into train/test sets.

3
Build a Model ← PyTorch lives here

Define the neural network architecture using nn.Module.

4
Train the Model ← Also PyTorch

Feed data, compute loss, backpropagate, update weights. Repeat.

5
Deploy

Save the trained model, load it in an application, run inference.

Real-World Applications

ApplicationWhat It DoesExample
Image RecognitionIdentifies objects in photosSelf-driving car vision systems
Natural LanguageUnderstands and generates textChatGPT-style models
Medical ImagingDetects anomalies in scansTumour detection in MRI
RecommendationsPredicts what you'll likeNetflix, Spotify suggestions
Speech RecognitionConverts audio to textSiri, voice assistants
Generative AICreates images and textDALL-E, Stable Diffusion
💡 Installing PyTorch

Run in your terminal: pip install torch torchvision torchaudio — then verify with import torch; print(torch.__version__). If it prints a version number, you're ready.

Tensors — The Core Data Structure

Everything in PyTorch is built on tensors. A tensor is like a Python list — but supercharged. It can hold millions of numbers, do maths on all of them simultaneously, and run on a GPU.


  Scalar (0D):    7

  Vector (1D):    [1, 2, 3]

  Matrix (2D):    [[1, 2, 3],
                   [4, 5, 6]]

  3D Tensor:      [[[1, 2], [3, 4]],
                    [[5, 6], [7, 8]]]

  4D Tensor:      Used for batches of images → (batch, channels, height, width)
      

Why Tensors Instead of Lists?

Python Lists
  • Slow — processes one number at a time
  • No GPU support
  • No automatic gradient tracking
  • Clunky multi-dimensional handling
PyTorch Tensors
  • Fast — all numbers processed in parallel
  • GPU-accelerated with one line
  • Tracks gradients automatically
  • Natural multi-dimensional operations

Creating Tensors

python
import torch

# From data you provide
scalar = torch.tensor(7)
vector = torch.tensor([1.0, 2.0, 3.0])
matrix = torch.tensor([[1.0, 2.0], [3.0, 4.0]])

# Useful creation functions
zeros  = torch.zeros(3, 4)       # all zeros, shape (3,4)
ones   = torch.ones(2, 3)        # all ones, shape (2,3)
rng    = torch.arange(0, 10, 2) # [0, 2, 4, 6, 8]
rand   = torch.rand(3, 3)       # random 0-1, uniform
randn  = torch.randn(3, 3)      # random, normal distribution

# Inspect a tensor
print(matrix.shape)   # torch.Size([2, 2])
print(matrix.ndim)    # 2
print(matrix.dtype)   # torch.float32

Data Types — Why dtype Matters

Neural networks almost always use float32 (decimal numbers). Feeding integer tensors where float32 is expected is one of the most common beginner errors.

dtypeMeaningWhen to use
torch.float3232-bit decimalModel weights, inputs — use this by default
torch.float6464-bit decimalHigh-precision scientific computing
torch.int6464-bit integerLabels in classification, indices
torch.boolTrue / FalseMasks and conditions
⚠️ Common Mistake

torch.tensor([1, 2, 3]) creates int64. For AI work, always use decimals: torch.tensor([1.0, 2.0, 3.0]) to get float32. Mixing types causes errors in operations.

Tensor Operations, Reshaping & Indexing

Basic Arithmetic

Operations happen element-wise — PyTorch applies the operation to every number simultaneously, with no loops needed.

python
import torch

a = torch.tensor([1.0, 2.0, 3.0])
b = torch.tensor([4.0, 5.0, 6.0])

print(a + b)   # tensor([5., 7., 9.])
print(a * b)   # tensor([ 4., 10., 18.]) — element-wise
print(a ** 2)  # tensor([1., 4., 9.])
print(a * 10)  # tensor([10., 20., 30.]) — broadcasting

Matrix Multiplication — The Engine of Deep Learning

Every single layer in every neural network uses matrix multiplication. This is the most important operation in all of AI.

python
A = torch.tensor([[1.0, 2.0, 3.0],
                   [4.0, 5.0, 6.0]])  # shape (2, 3)

B = torch.tensor([[7.0, 8.0],
                   [9.0, 10.0],
                   [11.0, 12.0]])  # shape (3, 2)

result = torch.matmul(A, B)
print(result.shape)   # torch.Size([2, 2])
# Rule: (2,3) × (3,2) → inner dims must match → result is (2,2)
💡 The Shape Rule

For torch.matmul(A, B) the inner dimensions must match. (2,3) × (3,2) works — result is (2,2). (2,3) × (2,3) will error. Always check shapes before multiplying.

Reshaping

Same numbers, different arrangement. Like rearranging 12 eggs from one row of 12 into 3 rows of 4.

python
x = torch.arange(12, dtype=torch.float32)
print(x.shape)           # torch.Size([12])

a = x.reshape(3, 4)     # 3 rows, 4 columns
b = x.reshape(4, -1)    # -1 means "figure it out" → (4, 3)
c = x.reshape(-1)       # flatten to 1D

# Add a dimension (needed for batch processing)
v = torch.tensor([1.0, 2.0, 3.0])   # shape (3,)
print(v.unsqueeze(0).shape)          # torch.Size([1, 3])
print(v.unsqueeze(1).shape)          # torch.Size([3, 1])

Indexing & Slicing

python
t = torch.tensor([[1, 2, 3],
                   [4, 5, 6],
                   [7, 8, 9]])

print(t[0, 1])     # tensor(2)       — row 0, col 1
print(t[0])        # tensor([1,2,3]) — entire first row
print(t[:, 0])     # tensor([1,4,7]) — entire first column
print(t[0:2, 1:3]) # rows 0-1, cols 1-2

CPU vs GPU — Device Management

Your computer has two processors that can do maths. The CPU is your main brain — powerful but sequential. The GPU was built for video games but turns out its parallel processing is perfect for AI.

CPU
  • ~8 to 16 powerful cores
  • Good at complex sequential tasks
  • Default for all PyTorch code
  • Fine for learning and small projects
GPU
  • Thousands of tiny parallel cores
  • Processes many numbers at once
  • Training can be 10–100× faster
  • Essential for real AI projects

The Standard Device Pattern

Write this at the top of every PyTorch file. It automatically uses a GPU if available, falls back to CPU if not — without changing any other code.

python — always start your files this way
import torch

# Set device once — use everywhere
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Using: {device}")

# Create tensor directly on device
t = torch.rand(3, 3).to(device)

# Move existing tensor to device
a = torch.tensor([1.0, 2.0])
a = a.to(device)

# Move back to CPU (required before .numpy())
a_cpu = a.cpu().numpy()
⚠️ The Critical Device Rule

Every tensor involved in an operation must be on the same device. Mixing CPU and GPU tensors causes a RuntimeError. Move everything to device before operating on it together.

🌐 No GPU? Use Google Colab

Go to colab.research.google.com, create a notebook, then Runtime → Change runtime type → GPU. You get a free NVIDIA GPU in your browser. torch.cuda.is_available() will return True.

Autograd — How PyTorch Learns Automatically

This is where PyTorch becomes truly magical. The question every beginner skips: how does the model know which way to adjust?

The answer is gradients — and PyTorch computes them automatically through a system called autograd.

The Blindfolded Hiker Analogy

🏔
The Landscape

Imagine hilly terrain. Your goal is to find the lowest valley. This landscape represents the model's error (loss).

👣
The Slope

Blindfolded, you can feel the slope under your feet. Slope downward left → step left. This slope is the gradient.

✅
The Valley

Keep stepping downhill and you reach the valley — minimum loss, best possible model. PyTorch calculates the slope automatically.

requires_grad, backward(), and Gradients

python
import torch

# Tell PyTorch to track this tensor (weights need this, not inputs)
w = torch.tensor([3.0], requires_grad=True)

# Forward pass — PyTorch secretly builds a "receipt" of every operation
y = w ** 2   # y = w²

# Backward pass — walk back through the receipt, compute gradients
y.backward()

# The gradient: dy/dw = 2w = 2×3 = 6
print(w.grad)   # tensor([6.])

The Three Sacred Lines — Always In This Order

python — the core of every training step
# 1. Zero out old gradients (they accumulate if you don't!)
optimizer.zero_grad()

# 2. Compute gradients by walking backwards through the graph
loss.backward()

# 3. Update weights in the direction that reduces loss
optimizer.step()

torch.no_grad() — For Evaluation

During evaluation you don't need gradients — tracking them wastes memory and time. Disable them with:

python
model.eval()                          # switch to evaluation mode
with torch.no_grad():              # disable gradient tracking
    predictions = model(X_test)   # safe, efficient inference
⚠️ The Accumulation Trap

PyTorch adds gradients to whatever was there before. If you forget optimizer.zero_grad() before each step, gradients pile up and your model learns incorrectly. This is the most common training bug.

Building Neural Networks with torch.nn

Strip away all complexity and a neural network is just: a function that takes numbers in, does a series of mathematical transformations, and produces numbers out.

The Building Blocks

Input

Numbers fed into the network. Images, text, measurements — everything is converted to numbers.

Weights

Learnable numbers that scale inputs. These start random and are adjusted during training.

Bias

An extra learnable number added after weights. Gives the network more flexibility.

Layer

One round of: multiply by weights, add bias. Networks stack multiple layers.

Activation

Applied after each layer. Adds non-linearity — without it, stacking layers does nothing useful.

Output

Final result. A number (regression), a probability (classification), or more complex structure.

nn.Module — The Base Class

Every PyTorch model inherits from nn.Module. It handles storing weights, tracking parameters, and switching between train/eval mode automatically.

python — the standard model pattern
import torch
import torch.nn as nn

class SimpleNetwork(nn.Module):

    def __init__(self):
        super().__init__()          # always required — activates nn.Module machinery

        self.layer1 = nn.Linear(3, 8)   # 3 inputs → 8 neurons
        self.layer2 = nn.Linear(8, 1)   # 8 → 1 output
        self.relu   = nn.ReLU()         # activation: max(0, x)

    def forward(self, x):            # defines how data flows through
        x = self.layer1(x)
        x = self.relu(x)
        x = self.layer2(x)
        return x

model = SimpleNetwork()

# Pass data through
x = torch.tensor([[1.0, 2.0, 3.0]])   # shape (1, 3) — 1 sample, 3 features
print(model(x).shape)                   # torch.Size([1, 1])

nn.Sequential — The Shortcut

python — for straight-line networks
model = nn.Sequential(
    nn.Linear(3, 8),
    nn.ReLU(),
    nn.Linear(8, 1)
)
# Identical result to the class above — less code for simple networks

Common Activation Functions

ActivationFormulaWhen to use
nn.ReLU()max(0, x)Default choice for hidden layers — almost always works
nn.Sigmoid()1 / (1 + e⁻ˣ)Binary classification output — squashes to 0–1
nn.Softmax()eˣ / ΣeˣMulti-class output — probabilities that sum to 1
nn.Tanh()(eˣ - e⁻ˣ) / (eˣ + e⁻ˣ)Squashes to -1 to 1, used in some RNNs

The Training Loop

The training loop is the heart of every PyTorch project. You will write this pattern in every single model you ever train. Memorise its structure.

Loss Functions

Loss FunctionUse CaseWhat it measures
nn.MSELoss()Regression — predict a numberMean squared difference between prediction and target
nn.BCELoss()Binary classification — 2 classesBinary cross entropy — requires Sigmoid before it
nn.CrossEntropyLoss()Multi-class classificationIncludes Softmax internally — most flexible choice

Optimizers

OptimizerCharacterWhen to use
optim.SGDSimple, classic gradient descentLearning fundamentals, some CV tasks
optim.AdamAdaptive, usually faster to convergeDefault choice for beginners — almost always works
optim.AdamWAdam with weight decay regularisationTransformers and large language models

The Complete Training Loop

python — this pattern repeats in every project
import torch
import torch.nn as nn
import torch.optim as optim

model     = nn.Sequential(nn.Linear(3, 8), nn.ReLU(), nn.Linear(8, 1))
loss_fn   = nn.MSELoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

X = torch.rand(100, 3)   # 100 samples, 3 features
y = torch.rand(100, 1)   # 100 targets

for epoch in range(100):

    # ── Training phase ──────────────────────
    model.train()                  # enable training mode

    predictions = model(X)        # forward pass
    loss        = loss_fn(predictions, y)   # measure error

    optimizer.zero_grad()          # clear old gradients ← NEVER SKIP
    loss.backward()                # compute gradients
    optimizer.step()               # update weights

    # ── Evaluation phase ─────────────────────
    model.eval()
    with torch.no_grad():
        test_preds = model(X)

    if (epoch + 1) % 20 == 0:
        print(f"Epoch {epoch+1}: Loss = {loss.item():.4f}")
💡 model.train() vs model.eval()

model.train() enables dropout and batch normalisation during training. model.eval() disables them for consistent predictions. Always set the mode explicitly — even if your current model has no dropout, it's a critical habit for when you add it later.

Dataset & DataLoader

In real AI, datasets have millions of samples — you can't load them all into memory at once. PyTorch solves this with batching: feed the model small chunks at a time.

All Data at Once
  • Memory crash on large datasets
  • One weight update per epoch
  • Slow convergence
  • Only works for tiny datasets
Mini Batches
  • Memory efficient — any dataset size
  • Many updates per epoch
  • Faster, more stable convergence
  • Standard in all real projects

Custom Dataset Class

Every custom Dataset must implement exactly three methods:

python
from torch.utils.data import Dataset, DataLoader

class MyDataset(Dataset):

    def __init__(self, X, y):
        self.X = X
        self.y = y

    def __len__(self):               # DataLoader needs total count
        return len(self.X)

    def __getitem__(self, index):   # DataLoader calls this per sample
        return self.X[index], self.y[index]

# Create dataset and dataloader
dataset    = MyDataset(X_train, y_train)
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)

# Loop through batches
for X_batch, y_batch in dataloader:
    print(X_batch.shape)   # (32, features) — one batch at a time
    break

TensorDataset — The Quick Version

python — when data is already in tensors
from torch.utils.data import TensorDataset

train_ds = TensorDataset(X_train, y_train)
test_ds  = TensorDataset(X_test,  y_test)

train_loader = DataLoader(train_ds, batch_size=32, shuffle=True)
test_loader  = DataLoader(test_ds,  batch_size=32, shuffle=False)
# shuffle=False for test — order doesn't matter, consistency does
💡 Choosing Batch Size

Start with batch_size=32. Too small (4, 8) = noisy gradients, slow. Too large (512+) = needs more memory, can miss good solutions. 32 or 64 works well for almost everything when learning.

Saving & Loading Models

Training a model once and saving it means you never have to retrain. This is essential — deploy it, share it, resume training, or keep the best version found during training.

The State Dict Pattern

PyTorch saves the state dict — a dictionary of all weights and biases. You save numbers, not structure. Loading requires recreating the architecture first, then filling in the saved numbers.

python — save and load weights
import torch

# ── Save ──────────────────────────────────────────────────────
torch.save(model.state_dict(), "model.pth")

# ── Load ──────────────────────────────────────────────────────
loaded_model = MyModelClass()   # recreate same architecture first
loaded_model.load_state_dict(
    torch.load("model.pth", weights_only=True)
)
loaded_model.eval()             # always set eval mode after loading

Full Checkpoint — For Resuming Training

python — save everything needed to resume
# Save checkpoint
torch.save({
    "epoch":      current_epoch,
    "model":      model.state_dict(),
    "optimizer":  optimizer.state_dict(),  # Adam has internal memory too
    "loss":       best_loss
}, "checkpoint.pth")

# Load checkpoint
ckpt = torch.load("checkpoint.pth", weights_only=False)
model.load_state_dict(ckpt["model"])
optimizer.load_state_dict(ckpt["optimizer"])
start_epoch = ckpt["epoch"]

Saving the Best Model During Training

python — best model pattern
best_loss = float("inf")   # start at infinity so any loss beats it

for epoch in range(EPOCHS):
    # ... training code ...

    if avg_loss < best_loss:
        best_loss = avg_loss
        torch.save(model.state_dict(), "best_model.pth")
        print(f"New best model saved (loss={best_loss:.4f})")

Key Functions Reference

Every important PyTorch function you need as a beginner — with purpose and usage at a glance.

Tensor Creation

torch.tensor(data)

Create tensor from Python list or number. Always use floats for AI: [1.0, 2.0].

torch.zeros(n, m)

Tensor of all zeros, shape (n, m). Used to initialise bias values.

torch.ones(n, m)

Tensor of all ones. Used to create masks and test shapes.

torch.arange(start, end, step)

Evenly spaced values. Like Python's range() but returns a tensor.

torch.rand(n, m)

Random values 0–1, uniform distribution. Weight initialisation.

torch.randn(n, m)

Random values, normal distribution centred at 0. More realistic init.

Tensor Inspection & Manipulation

tensor.shape

Returns torch.Size — dimensions of the tensor. Check this constantly.

tensor.dtype

The data type stored. Use float32 for model weights and inputs.

tensor.reshape(a, b)

Change shape without changing data. Use -1 to auto-calculate one dimension.

tensor.unsqueeze(dim)

Add a dimension at position dim. Common for adding batch dimension.

tensor.to(device)

Move tensor to CPU or GPU. Must match device of everything it interacts with.

torch.matmul(A, B)

Matrix multiplication — the engine of neural networks. Inner dims must match.

torch.cat(list, dim)

Concatenate tensors. dim=0 stacks rows, dim=1 stacks columns.

tensor.item()

Extract a single Python number from a scalar tensor. Use when printing losses.

Autograd

requires_grad=True

Tell PyTorch to track this tensor. Apply to weights and biases, not to input data.

loss.backward()

Compute gradients for all tracked tensors by walking back through the computation graph.

tensor.grad

The computed gradient after .backward(). Used by optimizer to update weights.

torch.no_grad()

Context manager — disables gradient tracking. Use during evaluation and inference.

Neural Network Modules

nn.Module

Base class for all models. Inherit from it. Provides parameter tracking and device management.

nn.Linear(in, out)

Fully connected layer. Performs: output = input × weight + bias. Core building block.

nn.ReLU()

Activation function: max(0, x). Default choice between hidden layers.

nn.Sigmoid()

Squashes to 0–1. Use as final layer for binary classification problems.

nn.Sequential

Chains layers in order. Clean shortcut when data flows in a straight line.

nn.MSELoss()

Mean squared error. Loss function for regression — predicting a number.

nn.BCELoss()

Binary cross entropy. Loss for binary classification. Requires Sigmoid before it.

nn.CrossEntropyLoss()

Multi-class classification loss. Includes Softmax internally. Most flexible.

Training Utilities

optim.Adam(params, lr)

Adaptive optimizer. Default choice — almost always works well. Start with lr=0.001.

optim.SGD(params, lr)

Classic gradient descent. Simple and transparent. Good for learning fundamentals.

optimizer.zero_grad()

Clear accumulated gradients. Call before every .backward(). Never skip this.

optimizer.step()

Update all weights using their computed gradients. Call after .backward().

model.train()

Enable training mode — activates dropout, batch norm. Call before training loop.

model.eval()

Enable evaluation mode — disables dropout. Call before testing and inference.

model.parameters()

Returns all learnable parameters. Pass to optimizer so it knows what to update.

model.state_dict()

Returns dictionary of all weights and biases. Pass to torch.save().

Projects — Built Step by Step

Project 1: Linear Regression

The model learns to discover a hidden formula from data alone — without being told what it is.

python — complete linear regression project
import torch, torch.nn as nn, torch.optim as optim

torch.manual_seed(42)

# Data — model doesn't know the formula, must discover it
TRUE_WEIGHT, TRUE_BIAS = 3.0, 1.0
X     = torch.rand(100, 1)
y     = TRUE_WEIGHT * X + TRUE_BIAS + torch.randn(100, 1) * 0.1
split = 80
X_train, y_train = X[:split], y[:split]
X_test,  y_test  = X[split:], y[split:]

# Model
class LinearModel(nn.Module):
    def __init__(self): super().__init__(); self.linear = nn.Linear(1, 1)
    def forward(self, x): return self.linear(x)

model     = LinearModel()
loss_fn   = nn.MSELoss()
optimizer = optim.Adam(model.parameters(), lr=0.01)

# Train
for epoch in range(200):
    model.train()
    loss = loss_fn(model(X_train), y_train)
    optimizer.zero_grad(); loss.backward(); optimizer.step()

# Results — model discovers weight≈3.0, bias≈1.0
for name, param in model.named_parameters():
    print(f"{name}: {param.data.item():.4f}")
# linear.weight: 2.9876  — true was 3.0
# linear.bias:   1.0043  — true was 1.0

Project 2: Binary Classification

Two clusters of points. The model learns to tell them apart.

python — complete binary classification project
import torch, torch.nn as nn, torch.optim as optim

torch.manual_seed(42)

# Two clusters: class 0 near (1,1), class 1 near (3,3)
X = torch.cat([torch.randn(100,2)+1, torch.randn(100,2)+3])
y = torch.cat([torch.zeros(100,1), torch.ones(100,1)])
idx = torch.randperm(200)
X, y = X[idx], y[idx]
X_train, y_train = X[:160], y[:160]
X_test,  y_test  = X[160:], y[160:]

model = nn.Sequential(
    nn.Linear(2, 8), nn.ReLU(),
    nn.Linear(8, 1), nn.Sigmoid()   # outputs probability 0-1
)
loss_fn   = nn.BCELoss()             # binary cross entropy
optimizer = optim.Adam(model.parameters(), lr=0.01)

for epoch in range(200):
    model.train()
    preds = model(X_train)
    loss  = loss_fn(preds, y_train)
    optimizer.zero_grad(); loss.backward(); optimizer.step()

model.eval()
with torch.no_grad():
    test_preds = model(X_test)
    accuracy   = ((test_preds >= 0.5).float() == y_test).float().mean()
    print(f"Test Accuracy: {accuracy.item()*100:.1f}%")   # ~97-99%

Project 3: End-to-End with DataLoader — Student Pass/Fail

The complete, production-grade pipeline. Every best practice applied. This is the template for all real projects.

01

Data Creation with Noise

Generate 1000 students with hours studied and hours slept. Add 5% label noise to simulate real-world messiness.

02

Shuffle & Split

Randomly shuffle with torch.randperm, then split 80% train / 20% test. Never skip shuffling.

03

Custom Dataset & DataLoader

Wrap data in StudentDataset(Dataset). Create train and test DataLoaders with batch_size=32.

04

Model with nn.Module

3-layer classifier: Linear(2,16) → ReLU → Linear(16,16) → ReLU → Linear(16,1) → Sigmoid.

05

Training Loop with Metrics

Track loss and accuracy per epoch across all batches. Save best model whenever test loss improves.

06

Load Best & Run Inference

Load saved weights, set eval mode, predict on new students. Probability → PASS or FAIL.

python — the complete skeleton (all projects follow this structure)
# ── Setup ─────────────────────────────────────────────────────
device = "cuda" if torch.cuda.is_available() else "cpu"
torch.manual_seed(42)

# ── Data → Dataset → DataLoader ───────────────────────────────
train_loader = DataLoader(StudentDataset(X_train, y_train), batch_size=32, shuffle=True)
test_loader  = DataLoader(StudentDataset(X_test,  y_test),  batch_size=32, shuffle=False)

# ── Model + Loss + Optimizer ──────────────────────────────────
model     = StudentClassifier().to(device)
loss_fn   = nn.BCELoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)
best_loss = float("inf")

# ── Training Loop ─────────────────────────────────────────────
for epoch in range(30):
    model.train()
    for xb, yb in train_loader:
        pred = model(xb); loss = loss_fn(pred, yb)
        optimizer.zero_grad(); loss.backward(); optimizer.step()

    model.eval()
    with torch.no_grad():
        test_loss = sum(loss_fn(model(xb), yb) for xb,yb in test_loader) / len(test_loader)

    if test_loss < best_loss:
        best_loss = test_loss
        torch.save(model.state_dict(), "best_model.pth")

# ── Inference ─────────────────────────────────────────────────
model.load_state_dict(torch.load("best_model.pth", weights_only=True))
model.eval()
with torch.no_grad():
    prob = model(new_student_data)
    print("PASS" if prob.item() >= 0.5 else "FAIL")