Architect8999's picture
feat: integrate Galaxy bugbounty checklist, clientside resources, paper2code
256c9c2 verified
|
Raw
History Blame Contribute Delete
10 kB
# 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
```