Memory Profiling & Optimization¶
Overview¶
Memory profiling is critical for production ML systems:
- Finding bottlenecks: Where does memory go?
- Memory leaks: Long-running inference servers
- Optimization strategies: Reduce memory usage by 50-80%
- Production monitoring: Track memory in real-time
-
Memory Profiling Tools¶
tracemalloc (Built-in)¶
import tracemalloc
import torch
tracemalloc.start()
# Code that might leak
for _ in range(1000):
x = torch.randn(100, 100)
y = x ** 2
z = y.sum()
# Get top memory allocations
current, peak = tracemalloc.get_traced_memory()
print(f"Current: {current / 1e6:.1f}MB; Peak: {peak / 1e6:.1f}MB")
# Get detailed snapshot
snapshot = tracemalloc.take_snapshot()
top_stats = snapshot.statistics('lineno')
for stat in top_stats[:3]:
print(stat)
memory_profiler (Line-by-line)¶
from memory_profiler import profile
import torch
@profile
def train_epoch():
"""Profile memory line by line."""
# This line allocated memory
data = [torch.randn(100, 100) for _ in range(100)]
# This line processes data
results = [x.sum() for x in data]
return results
# Run with
# Shows memory used after each line
objgraph (Object tracking)¶
import objgraph
import torch
# Track object growth
objgraph.show_most_common_types(limit=3)
# Create objects
for _ in range(1000):
x = torch.randn(10, 10)
y = x ** 2
# Show what grew
objgraph.show_most_common_types(limit=3)
# Find leaks
objgraph.show_refs([x], filename='refs.png')
Common Memory Patterns¶
Pattern 1: Accumulating Lists¶
# BAD
results = []
for i in range(10000):
tensor = torch.randn(1000, 1000)
results.append(tensor) # Keeps all tensors in memory!
# Total memory
# GOOD
for i in range(10000):
tensor = torch.randn(1000, 1000)
result = process(tensor)
save_to_disk(result)
# Tensor freed after processing
Pattern 2: Unclosed Resources¶
# BAD
f = open('data.txt')
lines = f.readlines()
# File still open, memory held
# GOOD
with open('data.txt') as f:
lines = f.readlines()
# File automatically closed, memory freed
Pattern 3: Circular References¶
# BAD
class Node:
def __init__(self):
self.ref = self # Circular!
node = Node()
del node # NOT freed due to cycle
# GOOD
import weakref
class Node:
def __init__(self):
self.ref = weakref.ref(self)
node = Node()
del node # Freed immediately
Gradient Checkpointing (50-80% Memory Reduction)¶
import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint
class ModelWithCheckpointing(nn.Module):
def __init__(self):
super().__init__()
self.layer1 = nn.Linear(1000, 1000)
self.layer2 = nn.Linear(1000, 1000)
self.layer3 = nn.Linear(1000, 1000)
def forward(self, x):
# Standard forward (stores all activations)
# x1 = self.layer1(x)
# x2 = self.layer2(x1)
# return self.layer3(x2)
# With checkpointing (recompute instead of storing)
x1 = checkpoint(self.layer1, x)
x2 = checkpoint(self.layer2, x1)
return checkpoint(self.layer3, x2)
model = ModelWithCheckpointing()
# Memory savings
# Speed cost
-
Quantization (75% Memory Reduction)¶
import torch
from torch.quantization import quantize_dynamic
# Original model
model = torch.nn.Sequential(
torch.nn.Linear(1000, 1000),
torch.nn.ReLU(),
torch.nn.Linear(1000, 10)
)
print(f"Original size: {sum(p.numel() for p in model.parameters()) * 4 / 1e6:.1f}MB")
# Quantized model (int8)
quantized_model = quantize_dynamic(
model,
{torch.nn.Linear},
dtype=torch.qint8
)
print(f"Quantized size: {sum(p.numel() for p in quantized_model.parameters()) * 1 / 1e6:.1f}MB")
# Memory reduction
Model Pruning (30-70% Memory Reduction)¶
import torch
import torch.nn.utils.prune as prune
model = torch.nn.Sequential(
torch.nn.Linear(1000, 1000),
torch.nn.ReLU(),
torch.nn.Linear(1000, 10)
)
# Prune 50% of weights in layer 1
prune.l1_unstructured(model[0], name='weight', amount=0.5)
# Make pruning permanent
prune.remove(model[0], 'weight')
# Memory reduction
-
Memory Monitoring Decorators¶
from functools import wraps
import tracemalloc
def memory_monitor(func):
"""Decorator to track memory usage of function."""
@wraps(func)
def wrapper(*args, **kwargs):
tracemalloc.start()
result = func(*args, **kwargs)
current, peak = tracemalloc.get_traced_memory()
print(f"{func.__name__}: Peak {peak / 1e6:.1f}MB")
return result
return wrapper
@memory_monitor
def train_epoch():
"""Function with memory monitoring."""
data = [torch.randn(100, 100) for _ in range(100)]
return sum(x.sum() for x in data)
train_epoch()
# Output
-
Production Monitoring with psutil¶
import psutil
import torch
import time
class MemoryMonitor:
"""Monitor system and process memory in production."""
def __init__(self, alert_threshold_mb=8000):
self.process = psutil.Process()
self.alert_threshold = alert_threshold_mb
def check_memory(self):
"""Check current memory usage."""
info = self.process.memory_info()
rss_mb = info.rss / 1e6 # Resident set size
if rss_mb > self.alert_threshold:
print(f"WARNING: Memory usage {rss_mb:.0f}MB > {self.alert_threshold}MB")
return rss_mb
def periodic_check(self, interval=60):
"""Check memory every N seconds."""
while True:
self.check_memory()
time.sleep(interval)
# Usage in inference server
monitor = MemoryMonitor(alert_threshold_mb=8000)
for request in incoming_requests:
result = model.infer(request)
memory = monitor.check_memory()
print(f"Current memory: {memory:.0f}MB")
-
Optimization Checklist¶
For Training¶
- Use gradient checkpointing for large models
- Use mixed precision (fp16) training
- Enable gradient accumulation (simulate larger batches)
- Use appropriate batch size (not too large)
- Profile with tracemalloc to find leaks
For Inference¶
- Use quantization (int4/int8) for model weights
- Implement batch inference (higher throughput, lower per-token memory)
- Use KV cache for transformers (avoid recomputation)
- Implement token streaming (don't buffer entire output)
- Monitor with psutil in production
General¶
- Use context managers for resources
- Avoid circular references (use weakref if needed)
- Implement proper caching (bounded cache size)
- Profile before optimizing (measure, don't guess)
Memory Reduction Impact Table¶
| Technique | Memory Reduction | Speed Impact | Difficulty | Use Case |
|---|---|---|---|---|
| Gradient Checkpointing | 50-80% | +20-30% slower | Medium | Training large models |
| Quantization | 75% (int4) | +2-4x faster | Low | Inference, storage |
| Model Pruning | 30-70% | -5-20% slower | High | Model optimization |
| Mixed Precision | 50% | +2-3x faster | Low | Training, inference |
| Batch Processing | Neutral | +5-10x | Low | Inference throughput |
| KV Cache | 90% (cache) | +10x | Medium | Transformer inference |
| Token Streaming | 99% (output) | Neutral | Medium | Large output tasks |
Real-World Example: Llama 2 70B¶
Baseline (FP32):
- Model weights: 280GB (70B * 4 bytes)
- Activations: 40GB (batch size 1, seq len 2048)
- Gradients: 280GB (training)
- Total: 600GB (training), 320GB (inference)
With Optimizations:
- Quantization (int4): 70GB (4x reduction)
- KV Cache: 5GB (cached, not 40GB)
- Gradient checkpointing: 30GB (training)
- Total: 105GB (training), 75GB (inference)
Final with all optimizations: 87.5% memory reduction!
-
Related Topics¶
- [01 Reference Counting & Garbage Collection](/05-py3/05-memory-&-performance/(01-reference-counting-garbage-collection/) - Memory management
- [02 Memory Layout & Cache Efficiency](/05-py3/05-memory-&-performance/(02-memory-layout-cache-efficiency/) - Performance tuning
- 03 Global Interpreter Lock (Gil) - Concurrency
- 02 Autograd Implementation - Gradient computation