File size: 10,023 Bytes
256c9c2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 | # Knowledge: Loss Functions
## Purpose
Canonical implementations of loss functions commonly referenced in papers, with attention to numerical stability, framework-specific differences, and common implementation mistakes.
---
## Cross-Entropy Loss
### Standard cross-entropy (classification)
```python
# PyTorch built-in:
loss_fn = nn.CrossEntropyLoss()
# Expects: predictions (batch, num_classes) as LOGITS (not probabilities)
# targets (batch,) as class indices (not one-hot)
```
**Common mistake:** Passing probabilities (after softmax) instead of logits. PyTorch's `CrossEntropyLoss` applies log-softmax internally for numerical stability.
### Cross-entropy with label smoothing
Papers describe label smoothing differently from what PyTorch does:
**Paper formula (Szegedy et al., 2016):**
For C classes and smoothing ε:
- Target for correct class: `1 - ε`
- Target for each incorrect class: `ε / (C - 1)`
**PyTorch >= 1.10 built-in:**
```python
loss_fn = nn.CrossEntropyLoss(label_smoothing=0.1)
```
PyTorch's implementation distributes ε uniformly: correct class gets `1 - ε + ε/C`, others get `ε/C`. This is slightly different from the paper formula (the paper distributes to incorrect classes only).
**Custom implementation matching the paper exactly:**
```python
def label_smoothed_cross_entropy(logits: torch.Tensor, targets: torch.Tensor,
n_classes: int, smoothing: float = 0.1) -> torch.Tensor:
"""Label smoothing exactly as described in Szegedy et al., 2016.
Args:
logits: (batch, n_classes) — raw logits, NOT probabilities
targets: (batch,) — class indices
n_classes: number of classes
smoothing: label smoothing factor ε
"""
log_probs = F.log_softmax(logits, dim=-1) # (batch, n_classes)
# NLL loss for the true class
nll_loss = -log_probs.gather(dim=-1, index=targets.unsqueeze(-1)).squeeze(-1)
# Smooth loss: uniform over all classes
smooth_loss = -log_probs.mean(dim=-1)
loss = (1.0 - smoothing) * nll_loss + smoothing * smooth_loss
return loss.mean()
```
**The difference matters:** For large vocabularies (NLP), the difference between PyTorch's built-in and the paper formula is negligible. For small class counts (e.g., 10), it can be noticeable.
---
## Contrastive Losses
### NT-Xent / InfoNCE (SimCLR, Oord et al.)
```python
def nt_xent_loss(z_i: torch.Tensor, z_j: torch.Tensor,
temperature: float = 0.5) -> torch.Tensor:
"""Normalized Temperature-scaled Cross Entropy loss.
Used in SimCLR (Chen et al., 2020).
Args:
z_i: (batch, d) — representations of view 1
z_j: (batch, d) — representations of view 2 (positive pairs)
temperature: τ — scaling factor
Returns:
Scalar loss
"""
batch_size = z_i.size(0)
# Normalize representations
z_i = F.normalize(z_i, dim=-1) # (batch, d)
z_j = F.normalize(z_j, dim=-1) # (batch, d)
# Concatenate representations: [z_i; z_j]
z = torch.cat([z_i, z_j], dim=0) # (2*batch, d)
# Compute similarity matrix
sim = torch.matmul(z, z.T) / temperature # (2*batch, 2*batch)
# Mask out self-similarity (diagonal)
mask = ~torch.eye(2 * batch_size, dtype=torch.bool, device=z.device)
sim = sim.masked_fill(~mask, float('-inf'))
# Positive pairs: (i, i+batch) and (i+batch, i)
labels = torch.cat([
torch.arange(batch_size, 2 * batch_size),
torch.arange(0, batch_size)
], dim=0).to(z.device) # (2*batch,)
return F.cross_entropy(sim, labels)
```
**Critical: temperature matters enormously.** SimCLR uses τ=0.5 in the paper, but the appendix shows results are sensitive to this value. CLIP uses τ as a learned parameter initialized to 0.07. Always check what temperature the paper uses.
**Common mistake:** Not normalizing the representations before computing similarity. Without normalization, the cosine similarity becomes a dot product, which is unbounded and destabilizes training.
### Triplet loss
```python
def triplet_loss(anchor: torch.Tensor, positive: torch.Tensor,
negative: torch.Tensor, margin: float = 1.0) -> torch.Tensor:
"""Standard triplet margin loss.
Args:
anchor: (batch, d)
positive: (batch, d)
negative: (batch, d)
margin: minimum desired distance gap
"""
pos_dist = F.pairwise_distance(anchor, positive) # (batch,)
neg_dist = F.pairwise_distance(anchor, negative) # (batch,)
loss = F.relu(pos_dist - neg_dist + margin)
return loss.mean()
```
---
## Diffusion Losses
### DDPM loss (Ho et al., 2020)
The simplified loss from DDPM. This is L_simple from Eq. 14:
```python
def ddpm_loss(model: nn.Module, x_0: torch.Tensor,
noise_schedule: dict) -> torch.Tensor:
"""Denoising Diffusion Probabilistic Models simplified loss.
L_simple = E_{t, x_0, ε}[||ε - ε_θ(x_t, t)||²]
The model predicts the noise ε that was added, not the clean signal.
Args:
model: noise prediction network ε_θ
x_0: (batch, C, H, W) — clean images
noise_schedule: dict with 'betas', 'alphas_cumprod', etc.
"""
batch_size = x_0.size(0)
# Sample random timesteps
t = torch.randint(0, len(noise_schedule['betas']), (batch_size,),
device=x_0.device)
# Sample noise
noise = torch.randn_like(x_0)
# Get noisy image: x_t = sqrt(α̅_t) * x_0 + sqrt(1 - α̅_t) * ε
alpha_cumprod_t = noise_schedule['alphas_cumprod'][t] # (batch,)
alpha_cumprod_t = alpha_cumprod_t.view(-1, 1, 1, 1) # (batch, 1, 1, 1) for broadcasting
x_t = torch.sqrt(alpha_cumprod_t) * x_0 + torch.sqrt(1 - alpha_cumprod_t) * noise
# Predict noise
predicted_noise = model(x_t, t) # (batch, C, H, W)
# MSE loss between true and predicted noise
loss = F.mse_loss(predicted_noise, noise)
return loss
```
**Noise schedule:**
```python
def linear_noise_schedule(timesteps: int, beta_start: float = 1e-4,
beta_end: float = 0.02) -> dict:
"""Linear variance schedule from DDPM.
β_t increases linearly from β_start to β_end over T timesteps.
"""
betas = torch.linspace(beta_start, beta_end, timesteps)
alphas = 1.0 - betas
alphas_cumprod = torch.cumprod(alphas, dim=0)
return {
'betas': betas,
'alphas': alphas,
'alphas_cumprod': alphas_cumprod,
'sqrt_alphas_cumprod': torch.sqrt(alphas_cumprod),
'sqrt_one_minus_alphas_cumprod': torch.sqrt(1, - alphas_cumprod),
}
```
**Common mistakes with diffusion losses:**
1. Indexing: `alphas_cumprod[t]` — make sure `t` is the right shape for broadcasting
2. The DDPM paper uses 1000 timesteps with linear schedule β₁=1e-4 to β_T=0.02
3. The model predicts noise ε, not x₀ (in the simplified loss; other parameterizations exist)
4. `sqrt_one_minus_alphas_cumprod` should NOT go to exactly 0 — check boundary conditions
---
## VAE ELBO
### Evidence Lower Bound
```python
def vae_loss(recon_x: torch.Tensor, x: torch.Tensor,
mu: torch.Tensor, log_var: torch.Tensor,
beta: float = 1.0) -> tuple:
"""VAE loss = Reconstruction loss + β * KL divergence.
ELBO = E_q[log p(x|z)] - β * KL(q(z|x) || p(z))
Args:
recon_x: (batch, ...) — reconstructed output
x: (batch, ...) — original input
mu: (batch, d_latent) — mean of q(z|x)
log_var: (batch, d_latent) — log variance of q(z|x)
beta: weight for KL term (β=1 gives standard VAE)
"""
# Reconstruction loss (pixel-wise)
recon_loss = F.mse_loss(recon_x, x, reduction='sum') / x.size(0)
# Alternative: F.binary_cross_entropy for images in [0,1]
# KL divergence: KL(N(μ, σ²) || N(0, 1))
# = -0.5 * Σ(1 + log(σ²) - μ² - σ²)
kl_loss = -0.5 * torch.sum(1 + log_var - mu.pow(2) - log_var.exp()) / x.size(0)
total_loss = recon_loss + beta * kl_loss
return total_loss, recon_loss, kl_loss
```
**Common mistakes:**
1. **Reconstruction loss type:** Binary cross-entropy for images normalized to [0,1], MSE for others. Papers often don't specify.
2. **Reduction:** `sum` vs `mean` — affects the relative weight of reconstruction vs KL. If using `mean`, the KL term effectively gets more weight relative to reconstruction for high-dimensional data.
3. **β value:** β=1 is standard VAE. β-VAE uses β>1 for more disentangled representations. If the paper says "VAE loss" without specifying β, use β=1 and flag it.
4. **KL divergence sign:** The formula gives a positive KL value. The loss ADDS KL (penalizes deviation from prior). If your KL is negative, the sign is wrong.
---
## Numerical stability patterns
### Log-sum-exp trick
When computing `log(Σ exp(x_i))`:
```python
# WRONG (overflow for large values):
result = torch.log(torch.sum(torch.exp(x)))
# CORRECT:
result = torch.logsumexp(x, dim=-1)
# Internally: max_x + log(Σ exp(x_i - max_x))
```
### Softmax stability
PyTorch's `F.softmax` and `F.log_softmax` are already numerically stable (subtract max internally). But if you implement softmax manually:
```python
# WRONG:
weights = torch.exp(scores) / torch.exp(scores).sum(dim=-1, keepdim=True)
# CORRECT:
scores = scores - scores.max(dim=-1, keepdim=True).values
weights = torch.exp(scores) / torch.exp(scores).sum(dim=-1, keepdim=True)
```
### Epsilon in denominators
When dividing by a value that could be zero (e.g., normalizing):
```python
# Not just any epsilon — match the paper's convention or use a safe default
x_normalized = x / (x.norm(dim=-1, keepdim=True) + 1e-8)
```
### Clamping log probabilities
```python
# When computing log of probabilities that might be 0:
log_probs = torch.log(probs.clamp(min=1e-8))
# Or better: compute in log space from the start using log_softmax
```
|