diff --git a/caveat/callbacks.py b/caveat/callbacks.py index 711efaf..e83d0b9 100644 --- a/caveat/callbacks.py +++ b/caveat/callbacks.py @@ -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) @@ -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 diff --git a/caveat/experiment.py b/caveat/experiment.py index 3d7b051..13649a3 100644 --- a/caveat/experiment.py +++ b/caveat/experiment.py @@ -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) diff --git a/caveat/models/base.py b/caveat/models/base.py index f9284ea..a143650 100644 --- a/caveat/models/base.py +++ b/caveat/models/base.py @@ -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. @@ -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(), @@ -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( @@ -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 @@ -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 @@ -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 diff --git a/caveat/models/continuous/cswae_lstm.py b/caveat/models/continuous/cswae_lstm.py index d6b2600..c5b7f86 100644 --- a/caveat/models/continuous/cswae_lstm.py +++ b/caveat/models/continuous/cswae_lstm.py @@ -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 diff --git a/caveat/models/continuous/cvqvae_lstm.py b/caveat/models/continuous/cvqvae_lstm.py index fd91dca..d468973 100644 --- a/caveat/models/continuous/cvqvae_lstm.py +++ b/caveat/models/continuous/cvqvae_lstm.py @@ -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 diff --git a/caveat/models/continuous/swae_lstm.py b/caveat/models/continuous/swae_lstm.py index 56ab487..aea1da4 100644 --- a/caveat/models/continuous/swae_lstm.py +++ b/caveat/models/continuous/swae_lstm.py @@ -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 diff --git a/caveat/models/continuous/vae_attention.py b/caveat/models/continuous/vae_attention.py index b4802e3..55bba2b 100644 --- a/caveat/models/continuous/vae_attention.py +++ b/caveat/models/continuous/vae_attention.py @@ -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, diff --git a/caveat/models/continuous/vqvae_lstm.py b/caveat/models/continuous/vqvae_lstm.py index 7c565df..cf9f821 100644 --- a/caveat/models/continuous/vqvae_lstm.py +++ b/caveat/models/continuous/vqvae_lstm.py @@ -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 diff --git a/caveat/models/joint_vaes/experiment.py b/caveat/models/joint_vaes/experiment.py index a8dc22e..9c4dd6d 100644 --- a/caveat/models/joint_vaes/experiment.py +++ b/caveat/models/joint_vaes/experiment.py @@ -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, ) @@ -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, @@ -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, ) diff --git a/caveat/models/joint_vaes/jvae_continuous.py b/caveat/models/joint_vaes/jvae_continuous.py index 1a67ad5..7d185af 100644 --- a/caveat/models/joint_vaes/jvae_continuous.py +++ b/caveat/models/joint_vaes/jvae_continuous.py @@ -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 @@ -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: diff --git a/caveat/models/joint_vaes/jvae_continuous_rerouted.py b/caveat/models/joint_vaes/jvae_continuous_rerouted.py index 2af841d..04f621f 100644 --- a/caveat/models/joint_vaes/jvae_continuous_rerouted.py +++ b/caveat/models/joint_vaes/jvae_continuous_rerouted.py @@ -160,7 +160,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 @@ -196,8 +196,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: diff --git a/caveat/runners.py b/caveat/runners.py index 61309c1..741be53 100644 --- a/caveat/runners.py +++ b/caveat/runners.py @@ -7,17 +7,17 @@ import torch from pandas import DataFrame from pytorch_lightning import LightningModule, Trainer -from pytorch_lightning.callbacks import ( - EarlyStopping, - LearningRateMonitor, - ModelCheckpoint, -) +from pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint from pytorch_lightning.loggers import TensorBoardLogger from torch import Tensor from torch.random import seed as seeder from caveat import cuda_available, data, encoding, label_encoding, models -from caveat.callbacks import LinearLossScheduler +from caveat.callbacks import ( + CollapseMonitor, + CyclicalBetaAnnealer, + LinearLossScheduler, +) from caveat.data.module import DataModule from caveat.encoding import BaseDataset, BaseEncoder from caveat.evaluate import evaluate @@ -584,9 +584,9 @@ def batch_eval_command( version = sorted([d for d in log_dir.iterdir() if d.is_dir()])[-1] outputs_dir = log_dir / version.name schedules_path = outputs_dir / schedules_name - synthetic_schedules_all[ - log_dir.name - ] = data.load_and_validate_schedules(schedules_path) + synthetic_schedules_all[log_dir.name] = ( + data.load_and_validate_schedules(schedules_path) + ) print( f"> Loaded {synthetic_schedules_all[log_dir.name].pid.nunique()} synthetic schedules from {schedules_path}" ) @@ -894,19 +894,17 @@ def generate( print( f"\n======= Sampling {len(population)} new schedules from synthetic attributes =======" ) - ( - synthetic_attributes, - synthetic_schedules, - zs, - ) = generate_from_attributes( - trainer, - attributes=population, - batch_size=batch_size, - latent_dims=latent_dims, - seed=seed, - ckpt_path=ckpt_path, - write_dir=write_dir, - cats=latent_categories, + (synthetic_attributes, synthetic_schedules, zs) = ( + generate_from_attributes( + trainer, + attributes=population, + batch_size=batch_size, + latent_dims=latent_dims, + seed=seed, + ckpt_path=ckpt_path, + write_dir=write_dir, + cats=latent_categories, + ) ) synthetic_attributes = attribute_encoder.decode(synthetic_attributes) synthetic_attributes.to_csv(write_dir / "synthetic_attributes.csv") @@ -1122,27 +1120,39 @@ def load_model(ckpt_path: Path, config: dict) -> LightningModule: def build_trainer(logger: TensorBoardLogger, config: dict) -> Trainer: trainer_config = config.get("trainer_params", {}) - patience = trainer_config.pop("patience", 5) - checkpoint_callback = ModelCheckpoint( - dirpath=Path(logger.log_dir, "checkpoints"), - monitor="val_loss", - save_top_k=2, - save_weights_only=False, - ) - loss_scheduling = trainer_config.pop("loss_scheduling", {}) - custom_loss_scheduler = LinearLossScheduler(loss_scheduling) - return Trainer( - logger=logger, - callbacks=[ + patience = trainer_config.pop("patience", None) + collapse_monitoring_config = trainer_config.pop("collapse_monitoring", {}) + loss_scheduling_config = trainer_config.pop("loss_scheduling", {}) + beta_scheduling_config = trainer_config.pop("beta_scheduling", {}) + + callbacks = [ + ModelCheckpoint( + dirpath=Path(logger.log_dir, "checkpoints"), + monitor="val_loss", + save_top_k=2, + save_weights_only=False, + ) + ] + + if patience is not None: + callbacks.append( EarlyStopping( - monitor="val_loss", patience=patience, stopping_threshold=0.0 - ), - LearningRateMonitor(), - checkpoint_callback, - custom_loss_scheduler, - ], - **trainer_config, - ) + monitor="val_recon_loss", + patience=patience, + stopping_threshold=0.0, + ) + ) + + if collapse_monitoring_config.get("enabled", True): + callbacks.append(CollapseMonitor(**collapse_monitoring_config)) + + if loss_scheduling_config.get("enabled", False): + callbacks.append(LinearLossScheduler(loss_scheduling_config)) + + if beta_scheduling_config.get("enabled", False): + callbacks.append(CyclicalBetaAnnealer(beta_scheduling_config)) + + return Trainer(logger=logger, callbacks=callbacks, **trainer_config) def initiate_logger(save_dir: Union[Path, str], name: str) -> TensorBoardLogger: diff --git a/tests/test_callbacks.py b/tests/test_callbacks.py new file mode 100644 index 0000000..ca81233 --- /dev/null +++ b/tests/test_callbacks.py @@ -0,0 +1,375 @@ +import math +from unittest.mock import MagicMock + +import pytest +import torch + +from caveat.callbacks import CollapseMonitor, CyclicalBetaAnnealer + + +def make_mocks(max_epochs, current_epoch): + trainer = MagicMock() + trainer.max_epochs = max_epochs + trainer.current_epoch = current_epoch + pl_module = MagicMock() + return trainer, pl_module + + +# --- Initialisation --- + + +def test_default_config(): + cb = CyclicalBetaAnnealer({}) + assert cb.n_cycles == 4 + assert cb.max_beta == 1.0 + assert cb.ratio == 0.5 + + +def test_custom_config(): + cb = CyclicalBetaAnnealer({"n_cycles": 2, "max_beta": 0.5, "ratio": 0.75}) + assert cb.n_cycles == 2 + assert cb.max_beta == 0.5 + assert cb.ratio == 0.75 + + +# --- Beta value cases --- +# Config: max_epochs=100, n_cycles=4, ratio=0.5 +# => cycle_len=25, ramp_end=12.5 + + +@pytest.fixture +def callback(): + return CyclicalBetaAnnealer({"n_cycles": 4, "max_beta": 1.0, "ratio": 0.5}) + + +def test_beta_zero_at_cycle_start(callback): + trainer, pl_module = make_mocks(max_epochs=100, current_epoch=0) + callback.on_train_epoch_start(trainer, pl_module) + assert pl_module.beta == 0.0 + + +def test_beta_mid_ramp(callback): + # epoch=6, cycle_pos=6, ramp_end=12.5 => beta = 1.0 * (6 / 12.5) = 0.48 + trainer, pl_module = make_mocks(max_epochs=100, current_epoch=6) + callback.on_train_epoch_start(trainer, pl_module) + assert pl_module.beta == pytest.approx(0.48) + + +def test_beta_near_ramp_end(callback): + # epoch=12, cycle_pos=12, ramp_end=12.5 => beta = 1.0 * (12 / 12.5) = 0.96 + trainer, pl_module = make_mocks(max_epochs=100, current_epoch=12) + callback.on_train_epoch_start(trainer, pl_module) + assert pl_module.beta == pytest.approx(0.96) + + +def test_beta_at_plateau(callback): + # epoch=13, cycle_pos=13 >= ramp_end=12.5 => beta = max_beta = 1.0 + trainer, pl_module = make_mocks(max_epochs=100, current_epoch=13) + callback.on_train_epoch_start(trainer, pl_module) + assert pl_module.beta == 1.0 + + +def test_beta_second_cycle_start(callback): + # epoch=25 => cycle_pos = 25 % 25 = 0 => beta = 0.0 + trainer, pl_module = make_mocks(max_epochs=100, current_epoch=25) + callback.on_train_epoch_start(trainer, pl_module) + assert pl_module.beta == 0.0 + + +def test_beta_last_epoch(callback): + # epoch=99 => cycle_pos = 99 % 25 = 24 >= ramp_end=12.5 => beta = 1.0 + trainer, pl_module = make_mocks(max_epochs=100, current_epoch=99) + callback.on_train_epoch_start(trainer, pl_module) + assert pl_module.beta == 1.0 + + +# --- Side effects --- + + +def test_log_called(callback): + trainer, pl_module = make_mocks(max_epochs=100, current_epoch=6) + callback.on_train_epoch_start(trainer, pl_module) + pl_module.log.assert_called_once_with("beta", pytest.approx(0.48)) + + +def test_pl_module_beta_set(callback): + trainer, pl_module = make_mocks(max_epochs=100, current_epoch=13) + callback.on_train_epoch_start(trainer, pl_module) + assert pl_module.beta == 1.0 + + +# =========================================================================== +# CollapseMonitor +# =========================================================================== + + +def make_collapse_mocks(current_epoch, mu, log_var): + trainer = MagicMock() + trainer.current_epoch = current_epoch + pl_module = MagicMock() + pl_module.encode.return_value = (mu, log_var) + return trainer, pl_module + + +def _populate_buffers(cb, n_batches=2, batch_size=4, latent_dim=3, n_conditions=2): + """Push synthetic data into cb's internal buffers.""" + mu = torch.randn(batch_size, latent_dim) + log_var = torch.zeros(batch_size, latent_dim) + kl = -0.5 * (1 + log_var - mu.pow(2) - log_var.exp()) + cond = torch.zeros(batch_size, 1) + for _ in range(n_batches): + cb._mus.append(mu) + cb._log_vars.append(log_var) + cb._conditions.append(cond) + cb._kl_per_dim.append(kl.mean(dim=0)) + + +# --- Initialisation --- + + +def test_collapse_monitor_defaults(): + cb = CollapseMonitor({}) + assert cb.au_threshold == 0.01 + assert cb.kl_collapse_threshold == 0.1 + assert cb.conditional_threshold == 0.05 + assert cb.check_every_n_epochs == 5 + assert cb.warn_au_below == 0.5 + + +# --- decoder sensitivity (accumulated in on_validation_batch_end) --- + + +def _make_batch_end_pl_module(decode_fn=None): + """Build a pl_module mock with label_encoder and configurable decode.""" + pl_module = MagicMock() + pl_module.label_encoder = MagicMock() + mu = torch.randn(4, 3) + log_var = torch.zeros(4, 3) + pl_module.encode.return_value = (mu, log_var) + if decode_fn is not None: + pl_module.decode.side_effect = decode_fn + else: + pl_module.decode.return_value = torch.zeros(4, 10, 5) + return pl_module + + +def test_decoder_sensitivity_no_label_encoder(): + """Batches with no label_encoder skip accumulation → nan at epoch end.""" + cb = CollapseMonitor({"check_every_n_epochs": 1}) + # batch_end mock: no label_encoder so no decoder probing + pl_batch = MagicMock(spec=["encode"]) + pl_batch.encode.return_value = (torch.randn(4, 3), torch.zeros(4, 3)) + trainer = MagicMock() + trainer.current_epoch = 0 + batch = (torch.randn(4, 10, 2), torch.zeros(4, 1)) + cb.on_validation_batch_end(trainer, pl_batch, None, batch, 0) + # epoch_end mock: regular MagicMock so log_dict works + pl_epoch = MagicMock() + cb.on_validation_epoch_end(trainer, pl_epoch) + logged = pl_epoch.log_dict.call_args[0][0] + assert math.isnan(logged["collapse/decoder_sensitivity"]) + + +def test_decoder_sensitivity_single_condition(capsys): + """All-same conditions in every batch → nan (no shuffle is meaningful) and a warning is printed.""" + cb = CollapseMonitor({"check_every_n_epochs": 1}) + pl_module = _make_batch_end_pl_module() + trainer = MagicMock() + trainer.current_epoch = 0 + batch = (torch.randn(4, 10, 2), torch.zeros(4, 1)) # all condition 0 + cb.on_validation_batch_end(trainer, pl_module, None, batch, 0) + out = capsys.readouterr().out + assert "CollapseMonitor" in out + assert "skipping decoder sensitivity" in out + _populate_buffers(cb, n_batches=0) + cb.on_validation_epoch_end(trainer, pl_module) + logged = pl_module.log_dict.call_args[0][0] + assert math.isnan(logged["collapse/decoder_sensitivity"]) + + +def test_decoder_sensitivity_insensitive_decoder(): + """Decoder returns same output for any condition → swap_mse=0 → dec_sens≈0.""" + cb = CollapseMonitor({"check_every_n_epochs": 1}) + pl_module = _make_batch_end_pl_module() # decode always returns zeros + trainer = MagicMock() + trainer.current_epoch = 0 + c = torch.cat([torch.zeros(2, 1), torch.ones(2, 1)]) + batch = (torch.randn(4, 10, 2), c) + cb.on_validation_batch_end(trainer, pl_module, None, batch, 0) + _populate_buffers(cb, n_batches=0) + cb.on_validation_epoch_end(trainer, pl_module) + logged = pl_module.log_dict.call_args[0][0] + assert logged["collapse/decoder_sensitivity"] == pytest.approx(0.0, abs=1e-6) + + +def test_decoder_sensitivity_responsive_decoder(): + """Decoder output differs by condition → dec_sens > 0.""" + cb = CollapseMonitor({"check_every_n_epochs": 1}) + call_n = {"n": 0} + + def decode_fn(mu, labels): + val = float(call_n["n"]) + call_n["n"] += 1 + return torch.full((len(mu), 10, 5), val) + + pl_module = _make_batch_end_pl_module(decode_fn=decode_fn) + trainer = MagicMock() + trainer.current_epoch = 0 + c = torch.cat([torch.zeros(2, 1), torch.ones(2, 1)]) + batch = (torch.randn(4, 10, 2), c) + cb.on_validation_batch_end(trainer, pl_module, None, batch, 0) + _populate_buffers(cb, n_batches=0) + cb.on_validation_epoch_end(trainer, pl_module) + logged = pl_module.log_dict.call_args[0][0] + assert logged["collapse/decoder_sensitivity"] > 0.0 + + +# --- on_validation_batch_end --- + + +def test_batch_end_accumulates(): + cb = CollapseMonitor({}) + mu = torch.randn(4, 3) + log_var = torch.zeros(4, 3) + trainer, pl_module = make_collapse_mocks(current_epoch=0, mu=mu, log_var=log_var) + batch = (torch.randn(4, 10, 2), torch.zeros(4, 1)) + + cb.on_validation_batch_end(trainer, pl_module, outputs=None, batch=batch, batch_idx=0) + cb.on_validation_batch_end(trainer, pl_module, outputs=None, batch=batch, batch_idx=1) + + assert len(cb._mus) == 2 + assert len(cb._log_vars) == 2 + assert len(cb._conditions) == 2 + assert len(cb._kl_per_dim) == 2 + + +# --- on_validation_epoch_end: skip path --- + + +def test_epoch_end_skips_non_check_epoch(): + cb = CollapseMonitor({"check_every_n_epochs": 5}) + _populate_buffers(cb) + trainer = MagicMock() + trainer.current_epoch = 1 # 1 % 5 != 0 + pl_module = MagicMock() + + cb.on_validation_epoch_end(trainer, pl_module) + + pl_module.log_dict.assert_not_called() + assert len(cb._mus) == 0 # buffers cleared + + +# --- on_validation_epoch_end: check path --- + + +def test_epoch_end_logs_on_check_epoch(): + cb = CollapseMonitor({"check_every_n_epochs": 5}) + _populate_buffers(cb) + trainer = MagicMock() + trainer.current_epoch = 0 # 0 % 5 == 0 + pl_module = MagicMock() + + cb.on_validation_epoch_end(trainer, pl_module) + + pl_module.log_dict.assert_called_once() + logged = pl_module.log_dict.call_args[0][0] + expected_keys = { + "collapse/active_units_pct", + "collapse/n_active_dims", + "collapse/kl_collapsed_dims", + "collapse/kl_mean", + "collapse/kl_min", + "collapse/decoder_sensitivity", + "collapse/mean_posterior_var", + } + assert expected_keys == set(logged.keys()) + + +def test_epoch_end_resets_buffers(): + cb = CollapseMonitor({"check_every_n_epochs": 5}) + _populate_buffers(cb) + trainer = MagicMock() + trainer.current_epoch = 0 + pl_module = MagicMock() + + cb.on_validation_epoch_end(trainer, pl_module) + + assert cb._mus == [] + assert cb._log_vars == [] + assert cb._conditions == [] + assert cb._kl_per_dim == [] + + +# --- Early stopping (integrated) --- + + +def _run_check_epoch(cb, dec_sens_value): + """Run one check epoch with a controlled decoder sensitivity value. + + Injects scalar buffers so mean(swap_mse)/mean(out_var) == dec_sens_value, + or leaves them empty for nan. + """ + trainer = MagicMock() + trainer.current_epoch = 0 # always a check epoch (0 % 5 == 0) + pl_module = MagicMock() + _populate_buffers(cb) + if not math.isnan(dec_sens_value): + cb._decoder_swap_mse = [dec_sens_value] + cb._decoder_out_var = [1.0] + cb.on_validation_epoch_end(trainer, pl_module) + return trainer + + +def test_stopping_disabled_by_default(): + cb = CollapseMonitor({}) + assert cb.stopping_patience is None + + +def test_stopping_patience_from_collapse_patience(): + cb = CollapseMonitor({"collapse_patience": 3}) + assert cb.stopping_patience == 3 + + +def test_stopping_patience_falls_back_to_patience(): + cb = CollapseMonitor({"patience": 7}) + assert cb.stopping_patience == 7 + + +def test_stopping_patience_collapse_patience_takes_priority(): + cb = CollapseMonitor({"collapse_patience": 3, "patience": 99}) + assert cb.stopping_patience == 3 + + +def test_no_stop_when_separation_healthy(): + cb = CollapseMonitor({"collapse_patience": 2, "conditional_threshold": 0.05}) + trainer = _run_check_epoch(cb, dec_sens_value=0.5) + assert trainer.should_stop is not True + assert cb._bad_epochs == 0 + + +def test_bad_epoch_counter_increments(): + cb = CollapseMonitor({"collapse_patience": 3, "conditional_threshold": 0.05}) + _run_check_epoch(cb, dec_sens_value=0.01) + assert cb._bad_epochs == 1 + + +def test_stops_after_patience_exceeded(): + cb = CollapseMonitor({"collapse_patience": 2, "conditional_threshold": 0.05}) + _run_check_epoch(cb, dec_sens_value=0.01) + trainer = _run_check_epoch(cb, dec_sens_value=0.01) + assert trainer.should_stop is True + + +def test_bad_epochs_reset_on_recovery(): + cb = CollapseMonitor({"collapse_patience": 3, "conditional_threshold": 0.05}) + _run_check_epoch(cb, dec_sens_value=0.01) # bad + _run_check_epoch(cb, dec_sens_value=0.5) # recovered + assert cb._bad_epochs == 0 + + +def test_nan_sep_does_not_increment_counter(): + cb = CollapseMonitor({"collapse_patience": 1, "conditional_threshold": 0.05}) + _run_check_epoch(cb, dec_sens_value=float("nan")) + assert cb._bad_epochs == 0 + # trainer.should_stop must not have been set to True + # (verified indirectly: counter stayed at 0) diff --git a/tests/test_callbacks_integration.py b/tests/test_callbacks_integration.py new file mode 100644 index 0000000..130b942 --- /dev/null +++ b/tests/test_callbacks_integration.py @@ -0,0 +1,251 @@ +"""Integration tests for CollapseMonitor using a real CVAEContLSTM model. + +These tests complement the mock-based tests in test_callbacks.py by exercising +the actual encode/predict paths of the model, catching interface mismatches that +mocks cannot detect. +""" +import math +from unittest.mock import MagicMock, patch + +import pytest +import torch + +from caveat.callbacks import CollapseMonitor +from caveat.models.continuous.cvae_lstm import CVAEContLSTM + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +LENGTH = 8 +N_ACTIVITIES = 5 +LATENT_DIM = 4 +HIDDEN_SIZE = 16 +HIDDEN_N = 1 +N_LABEL_COLS = 2 +LABEL_EMBED_SIZES = [3, 4] # 3 categories for label0, 4 for label1 + + +@pytest.fixture +def model(): + """Minimal CVAEContLSTM with small hidden size for fast tests.""" + return CVAEContLSTM( + in_shape=(LENGTH, N_ACTIVITIES), + encodings=N_ACTIVITIES, + labels_size=N_LABEL_COLS, + label_embed_sizes=LABEL_EMBED_SIZES, + sos=0, + latent_dim=LATENT_DIM, + hidden_size=HIDDEN_SIZE, + hidden_n=HIDDEN_N, + dropout=0.0, + ) + + +def make_batch(n=4, varied_conditions=True): + """Build a (x, c) batch compatible with CVAEContLSTM. + + x: [N, LENGTH, 2] — activity index (int as float) + duration + c: [N, N_LABEL_COLS] long — label indices within embed sizes + """ + x = torch.zeros(n, LENGTH, 2) + x[:, :, 0] = torch.randint(0, N_ACTIVITIES, (n, LENGTH)).float() + x[:, :, 1] = torch.rand(n, LENGTH) + + if varied_conditions: + c = torch.stack( + [ + torch.randint(0, LABEL_EMBED_SIZES[0], (n,)), + torch.randint(0, LABEL_EMBED_SIZES[1], (n,)), + ], + dim=1, + ) + # Guarantee at least two distinct rows so the derangement check passes. + c[0] = torch.tensor([0, 0]) + c[1] = torch.tensor([1, 1]) + else: + c = torch.zeros(n, N_LABEL_COLS, dtype=torch.long) + + return x, c + + +def make_trainer(current_epoch=0): + trainer = MagicMock() + trainer.current_epoch = current_epoch + return trainer + + +# --------------------------------------------------------------------------- +# on_validation_batch_end with real model +# --------------------------------------------------------------------------- + + +def test_batch_end_accumulates_mu_log_var(model): + cb = CollapseMonitor({"check_every_n_epochs": 1}) + trainer = make_trainer() + batch = make_batch() + + cb.on_validation_batch_end(trainer, model, None, batch, 0) + + assert len(cb._mus) == 1 + assert len(cb._log_vars) == 1 + assert cb._mus[0].shape == (4, LATENT_DIM) + assert cb._log_vars[0].shape == (4, LATENT_DIM) + + +def test_batch_end_accumulates_across_multiple_batches(model): + cb = CollapseMonitor({"check_every_n_epochs": 1}) + trainer = make_trainer() + + for i in range(3): + cb.on_validation_batch_end(trainer, model, None, make_batch(), i) + + assert len(cb._mus) == 3 + assert len(cb._kl_per_dim) == 3 + + +def test_batch_end_accumulates_decoder_sensitivity_buffers(model): + """With varied conditions the decoder swap buffers should be populated.""" + cb = CollapseMonitor({"check_every_n_epochs": 1}) + trainer = make_trainer() + cb.on_validation_batch_end(trainer, model, None, make_batch(varied_conditions=True), 0) + + assert len(cb._decoder_swap_mse) == 1 + assert len(cb._decoder_out_var) == 1 + assert math.isfinite(cb._decoder_swap_mse[0]) + assert math.isfinite(cb._decoder_out_var[0]) + + +def test_batch_end_skips_sensitivity_when_single_condition(model, capsys): + """All-same conditions → buffers stay empty and a warning is printed.""" + cb = CollapseMonitor({"check_every_n_epochs": 1}) + trainer = make_trainer() + cb.on_validation_batch_end( + trainer, model, None, make_batch(varied_conditions=False), 0 + ) + + assert cb._decoder_swap_mse == [] + assert cb._decoder_out_var == [] + out = capsys.readouterr().out + assert "CollapseMonitor" in out + assert "skipping decoder sensitivity" in out + + +def test_batch_end_raises_on_mu_shape_mismatch(model): + """Batch size mismatch between mu and c raises ValueError.""" + cb = CollapseMonitor({"check_every_n_epochs": 1}) + trainer = make_trainer() + x, c = make_batch(n=4, varied_conditions=True) + # Manually corrupt: override encode to return mu with wrong batch size + import unittest.mock as mock + bad_mu = torch.randn(3, LATENT_DIM) # 3 rows, but c has 4 + bad_log_var = torch.zeros(3, LATENT_DIM) + with mock.patch.object(model, "encode", return_value=(bad_mu, bad_log_var)): + with pytest.raises(ValueError, match="Batch size mismatch"): + cb.on_validation_batch_end(trainer, model, None, (x, c), 0) + + +# --------------------------------------------------------------------------- +# on_validation_epoch_end with real model +# --------------------------------------------------------------------------- + +EXPECTED_METRIC_KEYS = { + "collapse/active_units_pct", + "collapse/n_active_dims", + "collapse/kl_collapsed_dims", + "collapse/kl_mean", + "collapse/kl_min", + "collapse/decoder_sensitivity", + "collapse/mean_posterior_var", +} + + +def _run_epoch(model, cb, current_epoch=0, n_batches=2, varied_conditions=True): + """Populate buffers via batch_end and call epoch_end, returning logged metrics.""" + trainer = make_trainer(current_epoch=current_epoch) + for i in range(n_batches): + cb.on_validation_batch_end( + trainer, model, None, make_batch(varied_conditions=varied_conditions), i + ) + with patch.object(model, "log_dict") as mock_log: + cb.on_validation_epoch_end(trainer, model) + return mock_log + + +def test_epoch_end_logs_all_keys(model): + cb = CollapseMonitor({"check_every_n_epochs": 1}) + mock_log = _run_epoch(model, cb) + + mock_log.assert_called_once() + logged = mock_log.call_args[0][0] + assert set(logged.keys()) == EXPECTED_METRIC_KEYS + + +def test_epoch_end_values_are_finite(model): + cb = CollapseMonitor({"check_every_n_epochs": 1}) + mock_log = _run_epoch(model, cb) + + logged = mock_log.call_args[0][0] + for key, val in logged.items(): + if key != "collapse/decoder_sensitivity": + assert math.isfinite(val), f"{key} = {val} is not finite" + + +def test_epoch_end_decoder_sensitivity_finite_with_varied_conditions(model): + cb = CollapseMonitor({"check_every_n_epochs": 1}) + mock_log = _run_epoch(model, cb, varied_conditions=True) + + logged = mock_log.call_args[0][0] + assert math.isfinite(logged["collapse/decoder_sensitivity"]) + + +def test_epoch_end_decoder_sensitivity_nan_when_no_varied_conditions(model): + cb = CollapseMonitor({"check_every_n_epochs": 1}) + mock_log = _run_epoch(model, cb, varied_conditions=False) + + logged = mock_log.call_args[0][0] + assert math.isnan(logged["collapse/decoder_sensitivity"]) + + +def test_epoch_end_resets_buffers(model): + cb = CollapseMonitor({"check_every_n_epochs": 1}) + _run_epoch(model, cb) + + assert cb._mus == [] + assert cb._log_vars == [] + assert cb._conditions == [] + assert cb._kl_per_dim == [] + assert cb._decoder_swap_mse == [] + assert cb._decoder_out_var == [] + + +def test_epoch_end_skips_non_check_epoch(model): + cb = CollapseMonitor({"check_every_n_epochs": 5}) + trainer = make_trainer(current_epoch=3) + + cb.on_validation_batch_end(trainer, model, None, make_batch(), 0) + with patch.object(model, "log_dict") as mock_log: + cb.on_validation_epoch_end(trainer, model) + + mock_log.assert_not_called() + + +# --------------------------------------------------------------------------- +# Full pipeline: multiple check epochs +# --------------------------------------------------------------------------- + + +def test_full_pipeline_two_check_epochs(model): + """Run two consecutive check epochs and verify metrics logged each time.""" + cb = CollapseMonitor({"check_every_n_epochs": 5}) + log_call_count = 0 + + for epoch in (0, 5): + trainer = make_trainer(current_epoch=epoch) + for i in range(2): + cb.on_validation_batch_end(trainer, model, None, make_batch(), i) + with patch.object(model, "log_dict") as mock_log: + cb.on_validation_epoch_end(trainer, model) + log_call_count += mock_log.call_count + + assert log_call_count == 2 diff --git a/tests/test_models/test_base_kld.py b/tests/test_models/test_base_kld.py new file mode 100644 index 0000000..219b3d9 --- /dev/null +++ b/tests/test_models/test_base_kld.py @@ -0,0 +1,41 @@ +import types + +import pytest +import torch + +from caveat.models.base import Base + + +def make_stub(free_bits): + """Minimal stub with only the attribute Base.kld needs.""" + return types.SimpleNamespace(free_bits=free_bits) + + +def test_kld_zero_free_bits_perfect_posterior(): + # mu=0, log_var=0 → kl_per_dim = -0.5*(1+0-0-1) = 0 + stub = make_stub(free_bits=0.0) + mu = torch.zeros(2, 4) + log_var = torch.zeros(2, 4) + result = Base.kld(stub, mu, log_var) + assert result.item() == pytest.approx(0.0) + + +def test_kld_free_bits_clamps_small_kl(): + # mu=0, log_var=0 → raw kl=0 < free_bits=0.5; clamped to 0.5 per dim + # batch=1, latent_dim=2 → sum=1.0, mean over batch=1 → kld=1.0 + stub = make_stub(free_bits=0.5) + mu = torch.zeros(1, 2) + log_var = torch.zeros(1, 2) + result = Base.kld(stub, mu, log_var) + assert result.item() == pytest.approx(1.0) + + +def test_kld_free_bits_no_clamp_when_kl_exceeds_floor(): + # dim 0: mu=2, log_var=0 → kl = -0.5*(1+0-4-1) = 2.0 > 0.1 (not clamped) + # dim 1: mu=0, log_var=0 → kl = 0.0 < 0.1 (clamped to 0.1) + # batch=1 → kld = 2.0 + 0.1 = 2.1 + stub = make_stub(free_bits=0.1) + mu = torch.tensor([[2.0, 0.0]]) + log_var = torch.tensor([[0.0, 0.0]]) + result = Base.kld(stub, mu, log_var) + assert result.item() == pytest.approx(2.1)