Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
210 changes: 209 additions & 1 deletion caveat/callbacks.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,215 @@
import numpy as np
import torch
from pytorch_lightning import LightningModule, Trainer
from pytorch_lightning.callbacks import Callback


class CollapseMonitor(Callback):
def __init__(self, config: dict):
self.au_threshold = config.get("au_threshold", 0.01)
self.kl_collapse_threshold = config.get("kl_collapse_threshold", 0.1)
self.conditional_threshold = config.get("conditional_threshold", 0.05)
self.check_every_n_epochs = config.get("check_every_n_epochs", 5)
self.warn_au_below = config.get("warn_au_below", 0.5)

# Early stopping: opt-in via collapse_patience (falls back to patience).
# Counts check epochs below conditional_threshold; None disables stopping.
self.stopping_patience = config.get(
"collapse_patience", config.get("patience", None)
)
self._bad_epochs = 0

# Storage across batches within an epoch
self._mus = []
self._log_vars = []
self._conditions = []
self._kl_per_dim = []
# Decoder sensitivity accumulators (scalar per batch)
self._decoder_swap_mse = []
self._decoder_out_var = []

def on_validation_batch_end(
self, trainer, pl_module, outputs, batch, batch_idx
):
"""Collect latent stats for collapse diagnostics."""
x, c = batch
with torch.no_grad():
mu, log_var = pl_module.encode(x, c)

self._mus.append(mu.cpu())
self._log_vars.append(log_var.cpu())
self._conditions.append(c.cpu())

kl = -0.5 * (1 + log_var - mu.pow(2) - log_var.exp())
self._kl_per_dim.append(kl.mean(dim=0).cpu()) # mean over batch

# Decoder sensitivity: decode same z with real vs. shuffled conditions.
if hasattr(pl_module, "label_encoder"):
c_unique = c.unique(dim=0) if c.dim() > 1 else c.unique()
if len(c_unique) < 2:
print(
f" [CollapseMonitor] batch {batch_idx}: only {len(c_unique)} distinct "
f"condition(s) — skipping decoder sensitivity for this batch."
)
else:
if mu.dim() != 2:
raise ValueError(
f"[CollapseMonitor] Expected mu to be 2D [N, latent_dim], "
f"got shape {tuple(mu.shape)}"
)
if mu.shape[0] != c.shape[0]:
raise ValueError(
f"[CollapseMonitor] Batch size mismatch: "
f"mu has {mu.shape[0]} rows, c has {c.shape[0]} rows"
)
perm = self._derangement(len(c))
with torch.no_grad():
out_real = pl_module.decode(mu, labels=c).float()
out_perm = pl_module.decode(mu, labels=c[perm]).float()
self._decoder_swap_mse.append(
(out_real - out_perm).pow(2).mean().item()
)
self._decoder_out_var.append(out_real.var().item())

def on_validation_epoch_end(self, trainer, pl_module):
"""Compute collapse diagnostics and log."""
if trainer.current_epoch % self.check_every_n_epochs != 0:
self._reset()
return

mu_all = torch.cat(self._mus, dim=0) # [N, latent_dim]
log_var_all = torch.cat(self._log_vars, dim=0)
c_all = torch.cat(self._conditions, dim=0)
kl_mean = torch.stack(self._kl_per_dim).mean(dim=0) # [latent_dim]

metrics = {}

# 1. Posterior collapse: active units
au_mask = mu_all.var(dim=0) > self.au_threshold
au_pct = au_mask.float().mean().item()
metrics["collapse/active_units_pct"] = au_pct
metrics["collapse/n_active_dims"] = au_mask.sum().item()

# 2. Posterior collapse: per-dim KL
collapsed_dims = (kl_mean < self.kl_collapse_threshold).sum().item()
metrics["collapse/kl_collapsed_dims"] = collapsed_dims
metrics["collapse/kl_mean"] = kl_mean.mean().item()
metrics["collapse/kl_min"] = kl_mean.min().item()

# 3. Decoder sensitivity: ratio of condition-swap MSE to output variance.
# Accumulated per-batch during on_validation_batch_end.
if self._decoder_swap_mse:
swap_mse = np.mean(self._decoder_swap_mse)
out_var = np.mean(self._decoder_out_var)
dec_sens = swap_mse / (out_var + 1e-8)
else:
dec_sens = float("nan")
metrics["collapse/decoder_sensitivity"] = dec_sens

# 4. Posterior variance health
mean_posterior_var = log_var_all.exp().mean().item()
metrics["collapse/mean_posterior_var"] = mean_posterior_var

pl_module.log_dict(metrics, on_epoch=True)

# 5. Human-readable warnings
self._emit_warnings(trainer, au_pct, collapsed_dims, dec_sens, kl_mean)

# 6. Optional early stopping on persistent decoder insensitivity
if self.stopping_patience is not None and not np.isnan(dec_sens):
if dec_sens < self.conditional_threshold:
self._bad_epochs += 1
if self._bad_epochs >= self.stopping_patience:
print(
f"\n Stopping: decoder sensitivity {dec_sens:.4f} "
f"below {self.conditional_threshold} for "
f"{self.stopping_patience} check epochs."
)
trainer.should_stop = True
else:
self._bad_epochs = 0

self._reset()

def _derangement(self, n):
"""Random permutation with no fixed points (ensures conditions are always swapped)."""
for _ in range(10):
perm = torch.randperm(n)
if not (perm == torch.arange(n)).any():
return perm
# Fallback: cyclic shift always produces a derangement
return torch.roll(torch.arange(n), 1)

def _emit_warnings(
self, trainer, au_pct, collapsed_dims, dec_sens, kl_mean
):
epoch = trainer.current_epoch
issues = []

if au_pct < self.warn_au_below:
issues.append(
f" [POSTERIOR COLLAPSE] Only {au_pct:.1%} of latent dims active. "
f"Consider raising free_bits or reducing beta."
)

if collapsed_dims > 0:
issues.append(
f" [KL COLLAPSE] {collapsed_dims} dims have KL < {self.kl_collapse_threshold}. "
f"Min KL: {kl_mean.min():.4f}"
)

if isinstance(dec_sens, float) and not np.isnan(dec_sens):
if dec_sens < self.conditional_threshold:
issues.append(
f" [DECODER INSENSITIVE] Decoder sensitivity = {dec_sens:.4f}. "
f"Conditions are not influencing output sequences."
)

if issues:
print(f"\n Epoch {epoch} collapse warnings:")
for msg in issues:
print(msg)
# else:
# print(
# f"\n Epoch {epoch}: No collapse detected "
# f"(AU={au_pct:.1%}, cond_sep={cond_sep:.3f})"
# )

def _reset(self):
self._mus.clear()
self._log_vars.clear()
self._conditions.clear()
self._kl_per_dim.clear()
self._decoder_swap_mse.clear()
self._decoder_out_var.clear()


class CyclicalBetaAnnealer(Callback):
def __init__(self, config: dict) -> None:
"""
n_cycles: how many times to repeat the ramp
max_beta: the maximum value of beta to reach at the end of each ramp
ratio: fraction of each cycle spent ramping (rest stays at max_beta)
"""
self.n_cycles = config.get("n_cycles", 4)
self.max_beta = config.get("max_beta", 1.0)
self.ratio = config.get("ratio", 0.5)

def on_train_epoch_start(self, trainer, pl_module):
total_epochs = trainer.max_epochs
cycle_len = total_epochs / self.n_cycles
cycle_pos = trainer.current_epoch % cycle_len
ramp_end = cycle_len * self.ratio

if cycle_pos < ramp_end:
beta = self.max_beta * (cycle_pos / ramp_end)
else:
beta = self.max_beta

pl_module.beta = beta
pl_module.log("beta", beta)


class LinearLossScheduler(Callback):
def __init__(self, config: dict) -> None:
self.min_epochs = config.get("min_epochs", 0)
Expand Down Expand Up @@ -58,7 +266,7 @@ def on_train_epoch_start(
elif current_epoch >= e:
pl_module.scheduled_start_weight = 1.0
else:
pl_module.scheduled_end_weight = (current_epoch - s) / (e - s)
pl_module.scheduled_start_weight = (current_epoch - s) / (e - s)

if self.end_schedule is not None:
s, e = self.end_schedule
Expand Down
6 changes: 4 additions & 2 deletions caveat/experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,8 +66,10 @@ def __init__(
print(f"Found teacher forcing ratio: {self.teacher_forcing_ratio}")

# loss function params
self.kld_loss_weight = kwargs.get("kld_weight", 0.001)
print(f"Found KLD weight: {self.kld_loss_weight}")
self.beta = kwargs.get("kld_weight", 0.001)
print(f"Found KLD weight: {self.beta}")
self.free_bits = kwargs.get("free_bits", 0.0)
print(f"Found free bits: {self.free_bits}")
self.activity_loss_weight = kwargs.get("activity_loss_weight", 1.0)
self.duration_loss_weight = kwargs.get("duration_loss_weight", 200.0)
self.start_loss_weight = kwargs.get("start_loss_weight", 0.0)
Expand Down
24 changes: 14 additions & 10 deletions caveat/models/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,12 +121,12 @@ def reparameterize(self, mu: Tensor, logvar: Tensor) -> Tensor:
eps = torch.randn_like(std)
return (eps * std) + mu

def kld(self, mu: Tensor, log_var: Tensor) -> Tensor:
# from https://kvfrans.com/deriving-the-kl/
return torch.mean(
-0.5 * torch.sum(1 + log_var - mu**2 - log_var.exp(), dim=1),
dim=0,
)
def kld(self, mu, log_var):
# Per-dimension KL, then apply free bits floor
kl_per_dim = -0.5 * (1 + log_var - mu.pow(2) - log_var.exp())
# Free bits: only penalise KL above the floor
kl_per_dim = torch.clamp(kl_per_dim, min=self.free_bits)
return kl_per_dim.sum(dim=-1).mean() # mean over batch

def encode(self, input: Tensor, labels: Optional[Tensor]) -> list[Tensor]:
"""Encodes the input by passing through the encoder network.
Expand Down Expand Up @@ -440,12 +440,15 @@ def continuous_loss(

# kld loss
kld_loss = self.kld(mu, log_var)
scheduled_kld_weight = self.kld_loss_weight * self.scheduled_kld_weight
scheduled_kld_weight = self.beta * self.scheduled_kld_weight
w_kld_loss = scheduled_kld_weight * kld_loss

# final loss
loss = w_recons_loss + w_kld_loss

# Active units diagnostic
au = (mu.var(dim=0) > 0.01).float().mean()

return {
"loss": loss,
"KLD": w_kld_loss.detach(),
Expand All @@ -457,6 +460,7 @@ def continuous_loss(
"act_weight": torch.tensor([act_weight]).float(),
"dur_weight": torch.tensor([dur_weight]).float(),
"end_weight": torch.tensor([end_weight]).float(),
"active_units": au,
}

def discretized_loss(
Expand Down Expand Up @@ -487,7 +491,7 @@ def discretized_loss(

# kld loss
kld_loss = self.kld(mu, log_var)
scheduled_kld_weight = self.kld_loss_weight * self.scheduled_kld_weight
scheduled_kld_weight = self.beta * self.scheduled_kld_weight
w_kld_loss = scheduled_kld_weight * kld_loss

# final loss
Expand Down Expand Up @@ -628,7 +632,7 @@ def end_time_seq_loss(

# kld loss
kld_loss = self.kld(mu, log_var)
scheduled_kld_weight = self.kld_loss_weight * self.scheduled_kld_weight
scheduled_kld_weight = self.beta * self.scheduled_kld_weight
w_kld_loss = scheduled_kld_weight * kld_loss

# final loss
Expand Down Expand Up @@ -688,7 +692,7 @@ def combined_seq_loss(

# kld loss
kld_loss = self.kld(mu, log_var)
scheduled_kld_weight = self.kld_loss_weight * self.scheduled_kld_weight
scheduled_kld_weight = self.beta * self.scheduled_kld_weight
w_kld_loss = scheduled_kld_weight * kld_loss

# final loss
Expand Down
2 changes: 1 addition & 1 deletion caveat/models/continuous/cswae_lstm.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,7 +101,7 @@ def loss_function(
kld_loss = self.extra_sliced_wasserstein_distance(
z, prior_z, labels=labels, num_projections=32
)
scheduled_kld_weight = self.kld_loss_weight * self.scheduled_kld_weight
scheduled_kld_weight = self.beta * self.scheduled_kld_weight
w_kld_loss = scheduled_kld_weight * kld_loss

# final loss
Expand Down
2 changes: 1 addition & 1 deletion caveat/models/continuous/cvqvae_lstm.py
Original file line number Diff line number Diff line change
Expand Up @@ -352,7 +352,7 @@ def loss_function(
prior_loss = mu

# regularisation loss
scheduled_kld_weight = self.kld_loss_weight * self.scheduled_kld_weight
scheduled_kld_weight = self.beta * self.scheduled_kld_weight
w_vq_loss = scheduled_kld_weight * log_var

# final loss
Expand Down
2 changes: 1 addition & 1 deletion caveat/models/continuous/swae_lstm.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,7 +59,7 @@ def loss_function(
kld_loss = self.sliced_wasserstein_distance(
z, prior_z, num_projections=128
)
scheduled_kld_weight = self.kld_loss_weight * self.scheduled_kld_weight
scheduled_kld_weight = self.beta * self.scheduled_kld_weight
w_kld_loss = scheduled_kld_weight * kld_loss

# final loss
Expand Down
2 changes: 1 addition & 1 deletion caveat/models/continuous/vae_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,7 @@ def validation_step(self, batch, batch_idx, optimizer_idx=0):
log_var=log_var,
target=y,
weights=y_weights,
kld_weight=self.kld_loss_weight,
kld_weight=self.beta,
duration_weight=self.duration_loss_weight,
optimizer_idx=optimizer_idx,
batch_idx=batch_idx,
Expand Down
2 changes: 1 addition & 1 deletion caveat/models/continuous/vqvae_lstm.py
Original file line number Diff line number Diff line change
Expand Up @@ -288,7 +288,7 @@ def loss_function(
prior_loss = mu

# regularisation loss
scheduled_kld_weight = self.kld_loss_weight * self.scheduled_kld_weight
scheduled_kld_weight = self.beta * self.scheduled_kld_weight
w_vq_loss = scheduled_kld_weight * log_var

# final loss
Expand Down
6 changes: 3 additions & 3 deletions caveat/models/joint_vaes/experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ def training_step(self, batch, batch_idx):
log_var=log_var,
targets=(y, labels),
masks=(y_weights, label_weights),
kld_weight=self.kld_loss_weight,
kld_weight=self.beta,
duration_weight=self.duration_loss_weight,
batch_idx=batch_idx,
)
Expand All @@ -62,7 +62,7 @@ def validation_step(self, batch, batch_idx, optimizer_idx=0):
log_var=log_var,
targets=(y, labels),
masks=(y_weights, label_weights),
kld_weight=self.kld_loss_weight,
kld_weight=self.beta,
duration_weight=self.duration_loss_weight,
optimizer_idx=optimizer_idx,
batch_idx=batch_idx,
Expand Down Expand Up @@ -97,7 +97,7 @@ def test_step(self, batch, batch_idx):
log_var=log_var,
targets=(y, labels),
masks=(y_weights, label_weights),
kld_weight=self.kld_loss_weight,
kld_weight=self.beta,
duration_weight=self.duration_loss_weight,
batch_idx=batch_idx,
)
Expand Down
5 changes: 2 additions & 3 deletions caveat/models/joint_vaes/jvae_continuous.py
Original file line number Diff line number Diff line change
Expand Up @@ -185,7 +185,7 @@ def loss_function(

# kld loss
kld_loss = self.kld(mu, log_var)
scheduled_kld_weight = self.kld_loss_weight * self.scheduled_kld_weight
scheduled_kld_weight = self.beta * self.scheduled_kld_weight
w_kld_loss = scheduled_kld_weight * kld_loss

# final loss
Expand Down Expand Up @@ -224,8 +224,7 @@ def reparameterize(self, mu: Tensor, logvar: Tensor) -> Tensor:
def kld(self, mu: Tensor, log_var: Tensor) -> Tensor:
# from https://kvfrans.com/deriving-the-kl/
return torch.mean(
-0.5 * torch.sum(1 + log_var - mu**2 - log_var.exp(), dim=1),
dim=0,
-0.5 * torch.sum(1 + log_var - mu**2 - log_var.exp(), dim=1), dim=0
)

def predict(self, z: Tensor, device: int, **kwargs) -> Tensor:
Expand Down
Loading
Loading