Skip to content

Classes & Inheritance

Overview

PyTorch's entire architecture is built on object-oriented design: - nn.Module = Base class for all layers - Inheritance = Building custom layers - Multiple inheritance = Mixins for reusable functionality - Cooperative inheritance = Method Resolution Order (MRO)

Understanding OOP is essential to understanding PyTorch internals.


Class Basics

PyTorch Module Pattern

import torch

# PyTorch's base class
class MyModule(torch.nn.Module):
    def __init__(self, input_size, output_size):
        super().__init__()  # Call parent constructor
        self.fc = torch.nn.Linear(input_size, output_size)

    def forward(self, x):
        return self.fc(x)

# Usage
model = MyModule(10, 5)
output = model(torch.randn(3, 10))
print(output.shape)  # torch.Size([3, 5])

Attributes and Methods

class Layer:
    class_attribute = "shared"  # Shared across instances

    def __init__(self, name):
        self.name = name  # Instance attribute (per instance)
        self.weight = None

    def forward(self, x):
        return x * 2

    def backward(self):
        print("Computing gradients...")

# Instance creation
layer1 = Layer("layer1")
layer2 = Layer("layer2")

print(layer1.name)  # "layer1"
print(layer2.name)  # "layer2"
print(layer1.class_attribute)  # "shared" (same for both)
print(layer1.class_attribute is layer2.class_attribute)  # True

Inheritance: Building Layer Hierarchies

Single Inheritance

class Layer(torch.nn.Module):
    """Base layer class."""
    def __init__(self, input_size, output_size):
        super().__init__()
        self.input_size = input_size
        self.output_size = output_size

    def forward(self, x):
        raise NotImplementedError

class Linear(Layer):
    """Linear transformation layer."""
    def __init__(self, input_size, output_size):
        super().__init__(input_size, output_size)
        self.weight = torch.nn.Parameter(
            torch.randn(output_size, input_size) / input_size**0.5
        )

    def forward(self, x):
        return x @ self.weight.T

class ReLU(Layer):
    """Activation function."""
    def forward(self, x):
        return torch.relu(x)

# Usage
linear = Linear(10, 5)
relu = ReLU(10, 5)

Method Override

class BaseModel(torch.nn.Module):
    def forward(self, x):
        print("Base forward")
        return x

    def backward(self):
        print("Base backward")

class CustomModel(BaseModel):
    def forward(self, x):
        # Override parent method
        print("Custom forward")
        return super().forward(x) + 1  # Extend parent behavior

    # backward not overridden, uses parent's

model = CustomModel()
model.forward(torch.zeros(3))  # "Custom forward", "Base forward"
model.backward()  # "Base backward"

Multiple Inheritance: Mixins

Mixin Pattern

Mixins add functionality without being complete classes:

class TimingMixin:
    """Add timing capability to any class."""
    def time_forward(self, x):
        import time
        start = time.time()
        output = self.forward(x)
        elapsed = time.time() - start
        print(f"Forward pass took {elapsed:.3f}s")
        return output

class FreezeableMixin:
    """Add freezing capability."""
    def freeze(self):
        """Freeze all parameters."""
        for param in self.parameters():
            param.requires_grad = False

    def unfreeze(self):
        """Unfreeze all parameters."""
        for param in self.parameters():
            param.requires_grad = True

# Combine mixins
class MyModel(TimingMixin, FreezeableMixin, torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.fc = torch.nn.Linear(10, 5)

    def forward(self, x):
        return self.fc(x)

# Usage
model = MyModel()
model.time_forward(torch.randn(3, 10))  # Shows timing
model.freeze()  # Freezes parameters
model.unfreeze()  # Unfreezes

Method Resolution Order (MRO)

Python uses C3 linearization to determine method resolution order:

class A:
    def method(self):
        print("A")

class B(A):
    def method(self):
        print("B")
        super().method()

class C(A):
    def method(self):
        print("C")
        super().method()

class D(B, C):
    pass

# MRO: D → B → C → A
d = D()
d.method()
# Output:
# B
# C
# A

# View MRO
print(D.__mro__)
# (<class 'D'>, <class 'B'>, <class 'C'>, <class 'A'>, <class 'object'>)

Practical: Cooperative Multiple Inheritance

class Base(torch.nn.Module):
    def __init__(self, name):
        super().__init__()
        self.name = name

class Trainable:
    def train_mode(self):
        self.train()

class Serializable:
    def save(self, path):
        torch.save(self.state_dict(), path)

    def load(self, path):
        self.load_state_dict(torch.load(path))

class Model(Base, Trainable, Serializable):
    def __init__(self, name):
        super().__init__(name)
        self.fc = torch.nn.Linear(10, 5)

    def forward(self, x):
        return self.fc(x)

# All mixin methods available
model = Model("my_model")
model.train_mode()
model.save("model.pt")
model.load("model.pt")

Special Methods: Constructor and Destructor

__init__ Constructor

class Tensor:
    def __init__(self, data, requires_grad=False):
        self.data = data
        self.requires_grad = requires_grad
        self.grad = None

    def __repr__(self):
        return f"Tensor({self.data}, requires_grad={self.requires_grad})"

t = Tensor([1, 2, 3], requires_grad=True)
print(t)  # Tensor([1, 2, 3], requires_grad=True)

__del__ Destructor (Use Sparingly)

class GPUBuffer:
    def __init__(self, size):
        self.buffer = torch.cuda.FloatTensor(size)
        print(f"Allocated {size} elements on GPU")

    def __del__(self):
        # Called when object is garbage collected
        print("Freeing GPU memory")
        # Note: Implicit, not guaranteed timing

# Creation
buf = GPUBuffer(1000000)

# Deletion
del buf  # May not immediately free (depends on GC)

Warning: Don't rely on __del__ for critical cleanup. Use context managers instead.


Instance vs Class Attributes

Mutation Gotchas

class Layer:
    # Class attribute (mutable!) - GOTCHA
    weights = []  # Shared across all instances

    def add_weight(self, w):
        self.weights.append(w)

layer1 = Layer()
layer2 = Layer()

layer1.add_weight(1)
print(layer2.weights)  # [1] - SHARED! Bug!

# Correct version
class LayerFixed:
    def __init__(self):
        self.weights = []  # Instance attribute

    def add_weight(self, w):
        self.weights.append(w)

layer1 = LayerFixed()
layer2 = LayerFixed()
layer1.add_weight(1)
print(layer2.weights)  # [] - Correct, separate instances

PyTorch Example: Registering Parameters

# WRONG
class BadModel(torch.nn.Module):
    layers = []  # Shared class attribute!

    def __init__(self):
        super().__init__()
        self.layers.append(torch.nn.Linear(10, 5))

# CORRECT
class GoodModel(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.layers = torch.nn.ModuleList([
            torch.nn.Linear(10, 5),
            torch.nn.Linear(5, 2)
        ])

    def forward(self, x):
        for layer in self.layers:
            x = layer(x)
        return x

Abstract Base Classes

Using abc module

from abc import ABC, abstractmethod

class Optimizer(ABC):
    """Base class for optimizers."""

    @abstractmethod
    def step(self):
        """Perform one optimization step."""
        pass

    @abstractmethod
    def zero_grad(self):
        """Clear gradients."""
        pass

# Can't instantiate abstract class
# opt = Optimizer()  # TypeError

# Must implement abstract methods
class SGD(Optimizer):
    def __init__(self, params, lr=0.01):
        self.params = params
        self.lr = lr

    def step(self):
        for param in self.params:
            param.data -= self.lr * param.grad

    def zero_grad(self):
        for param in self.params:
            param.grad = None

opt = SGD([torch.randn(10, requires_grad=True)])  # OK

Composition vs Inheritance

Inheritance (Is-a relationship)

class Encoder(torch.nn.Module):
    def forward(self, x):
        return x

class LSTM(Encoder):  # "Is-a" Encoder
    def forward(self, x):
        return super().forward(x)

Composition (Has-a relationship)

class Seq2Seq(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.encoder = Encoder()  # "Has-a" Encoder
        self.decoder = Decoder()  # "Has-a" Decoder

    def forward(self, x):
        encoded = self.encoder(x)
        return self.decoder(encoded)

Best Practice: Favor composition over inheritance for flexible architectures.


Properties: Attribute-like Methods

class Model:
    def __init__(self):
        self._train = True

    @property
    def is_training(self):
        """Read-only property."""
        return self._train

    @is_training.setter
    def is_training(self, value):
        """Setter for property."""
        self._train = value

model = Model()
print(model.is_training)  # True (calls getter)
model.is_training = False  # Calls setter

Practical ML Patterns

Pattern 1: Building Custom Layers

class AttentionHead(torch.nn.Module):
    def __init__(self, dim, head_dim):
        super().__init__()
        self.scale = head_dim ** -0.5
        self.qkv = torch.nn.Linear(dim, head_dim * 3)
        self.out = torch.nn.Linear(head_dim, dim)

    def forward(self, x):
        B, N, C = x.shape
        qkv = self.qkv(x)
        q, k, v = qkv.chunk(3, dim=-1)

        attn = (q @ k.T) * self.scale
        attn = attn.softmax(dim=-1)

        return self.out(attn @ v)

Pattern 2: Building Model Stacks

class TransformerBlock(torch.nn.Module):
    def __init__(self, dim, heads):
        super().__init__()
        self.attention = AttentionHead(dim, dim // heads)
        self.norm1 = torch.nn.LayerNorm(dim)
        self.mlp = torch.nn.Sequential(
            torch.nn.Linear(dim, dim * 4),
            torch.nn.GELU(),
            torch.nn.Linear(dim * 4, dim)
        )
        self.norm2 = torch.nn.LayerNorm(dim)

    def forward(self, x):
        x = x + self.attention(self.norm1(x))
        x = x + self.mlp(self.norm2(x))
        return x

class Transformer(torch.nn.Module):
    def __init__(self, depth, dim, heads):
        super().__init__()
        self.blocks = torch.nn.ModuleList([
            TransformerBlock(dim, heads) for _ in range(depth)
        ])

    def forward(self, x):
        for block in self.blocks:
            x = block(x)
        return x

Summary

  • Classes = Objects with attributes and methods
  • Inheritance = Reuse code and extend functionality
  • Mixins = Add capabilities without full inheritance
  • MRO = Method Resolution Order determines which method is called
  • Abstract classes = Define interfaces without implementation
  • Composition = Flexible alternative to inheritance
  • Properties = Attribute-like methods with getters/setters