diff --git a/bluemath_tk/deeplearning/_base_model.py b/bluemath_tk/deeplearning/_base_model.py index 65cafed..852b04d 100644 --- a/bluemath_tk/deeplearning/_base_model.py +++ b/bluemath_tk/deeplearning/_base_model.py @@ -9,6 +9,7 @@ from tqdm import tqdm from ..core.models import BlueMathModel +from .latent_structure import validate_latent_structure_mode from .metrics import _validate_eps from .metrics import evaluate_reconstruction as evaluate_reconstruction_metric from .metrics import reconstruction_error as reconstruction_error_metric @@ -32,7 +33,15 @@ class BaseDeepLearningModel(BlueMathModel): """ @abstractmethod - def __init__(self, device: str | torch.device | None = None, **kwargs): + def __init__( + self, + device: str | torch.device | None = None, + latent_structure: str = "none", + latent_orthogonality_weight: float = 1e-2, + latent_decorrelation_weight: float = 1e-2, + latent_ordering_probability: float = 0.5, + **kwargs, + ): """ Initialize the base deep learning model. @@ -48,6 +57,32 @@ def __init__(self, device: str | torch.device | None = None, **kwargs): super().__init__(**kwargs) + self.latent_structure = validate_latent_structure_mode(latent_structure) + + for name, value in ( + ("latent_orthogonality_weight", latent_orthogonality_weight), + ("latent_decorrelation_weight", latent_decorrelation_weight), + ): + if ( + not isinstance(value, Real) + or isinstance(value, (bool, np.bool_)) + or not np.isfinite(float(value)) + or float(value) < 0.0 + ): + raise ValueError(f"{name} must be a finite non-negative number.") + setattr(self, name, float(value)) + + if ( + not isinstance(latent_ordering_probability, Real) + or isinstance(latent_ordering_probability, (bool, np.bool_)) + or not np.isfinite(float(latent_ordering_probability)) + or not 0.0 <= float(latent_ordering_probability) <= 1.0 + ): + raise ValueError( + "latent_ordering_probability must be finite and between 0 and 1." + ) + self.latent_ordering_probability = float(latent_ordering_probability) + if device is None: self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") elif isinstance(device, str): @@ -555,6 +590,36 @@ def _stage_checkpoint(self, checkpoint: dict): staged = object.__new__(self.__class__) staged.__dict__ = copy.deepcopy(self.__dict__) checkpoint_config_names = [] + init_config = checkpoint.get("init_config", {}) + if not isinstance(init_config, dict): + raise TypeError("Checkpoint init_config must be a dictionary.") + + latent_config_names = ( + "latent_structure", + "latent_orthogonality_weight", + "latent_decorrelation_weight", + "latent_ordering_probability", + ) + if staged.model is not None: + mismatches = [] + for name in latent_config_names: + if name not in init_config or not hasattr(staged, name): + continue + current = getattr(staged, name) + checkpoint_value = init_config[name] + if current != checkpoint_value: + mismatches.append( + f"{name}: current={current!r}, " + f"checkpoint={checkpoint_value!r}" + ) + if mismatches: + details = "\n - ".join(mismatches) + raise ValueError( + "Checkpoint latent configuration does not match the " + "already-built model. Load into an unbuilt compatible " + "instance or use from_pytorch_model().\n - " + f"{details}" + ) if staged.model is None: build_input_shape = checkpoint.get("build_input_shape") @@ -564,9 +629,6 @@ def _stage_checkpoint(self, checkpoint: dict): "Build the model manually before loading it." ) - init_config = checkpoint.get("init_config", {}) - if not isinstance(init_config, dict): - raise TypeError("Checkpoint init_config must be a dictionary.") for name, value in init_config.items(): if hasattr(staged, name): setattr(staged, name, copy.deepcopy(value)) @@ -627,6 +689,54 @@ def _loss_to_sample_total( return value return value * batch_sample_count + def _model_regularization_losses(self) -> dict[str, torch.Tensor]: + """Collect scalar losses from BlueMath latent regularizer modules.""" + if self.model is None: + return {} + + collected: dict[str, torch.Tensor] = {} + for module_name, module in self.model.named_modules(): + if not getattr(module, "_bluemath_latent_regularizer", False): + continue + for loss_name, loss in module.regularization_losses().items(): + if not isinstance(loss, torch.Tensor) or loss.ndim != 0: + raise TypeError( + f"Regularization loss {loss_name!r} from " + f"{module_name!r} must be a scalar tensor." + ) + key = f"{module_name}.{loss_name}" if module_name else loss_name + collected[key] = loss + return collected + + def _model_regularization_loss( + self, + reference: torch.Tensor, + ) -> torch.Tensor: + """Return the current sum of structured-latent losses.""" + total = torch.zeros( + (), + dtype=reference.dtype, + device=reference.device, + ) + for loss in self._model_regularization_losses().values(): + total = total + loss.to( + dtype=reference.dtype, + device=reference.device, + ) + return total + + def latent_diagnostics( + self, + X: np.ndarray, + batch_size: int = 64, + verbose: int = 0, + ) -> dict: + """Return diagnostics for encoded latent scores.""" + from .latent_structure import compute_latent_diagnostics + + latent = self.encode(X, batch_size=batch_size, verbose=verbose) + return compute_latent_diagnostics(latent) + def fit( self, X: np.ndarray, @@ -730,8 +840,12 @@ def fit( self._require_matching_output_shape(output, batch_y, "Training") self._require_finite_tensor(output, "Training output") self._require_finite_buffers() - loss = criterion(output, batch_y) - self._require_scalar_loss(loss) + reconstruction_loss = criterion(output, batch_y) + self._require_scalar_loss(reconstruction_loss) + regularization_loss = self._model_regularization_loss( + reconstruction_loss + ) + loss = reconstruction_loss + regularization_loss self._require_finite_loss(loss, "Training") loss.backward() self._require_finite_gradients() @@ -739,10 +853,13 @@ def fit( self._require_finite_parameters() train_total += self._loss_to_sample_total( - loss, + reconstruction_loss, current_batch_size, criterion, ) + train_total += ( + float(regularization_loss.item()) * current_batch_size + ) train_sample_count += current_batch_size train_loss = train_total / train_sample_count @@ -763,14 +880,21 @@ def fit( self._require_matching_output_shape(output, batch_y, "Validation") self._require_finite_tensor(output, "Validation output") self._require_finite_parameters() - loss = criterion(output, batch_y) - self._require_scalar_loss(loss) + reconstruction_loss = criterion(output, batch_y) + self._require_scalar_loss(reconstruction_loss) + regularization_loss = self._model_regularization_loss( + reconstruction_loss + ) + loss = reconstruction_loss + regularization_loss self._require_finite_loss(loss, "Validation") validation_total += self._loss_to_sample_total( - loss, + reconstruction_loss, current_batch_size, criterion, ) + validation_total += ( + float(regularization_loss.item()) * current_batch_size + ) validation_sample_count += current_batch_size validation_loss = validation_total / validation_sample_count diff --git a/bluemath_tk/deeplearning/autoencoders.py b/bluemath_tk/deeplearning/autoencoders.py index 4640c3e..c178c96 100644 --- a/bluemath_tk/deeplearning/autoencoders.py +++ b/bluemath_tk/deeplearning/autoencoders.py @@ -43,6 +43,7 @@ from tqdm import tqdm from ._base_model import BaseDeepLearningModel +from .latent_structure import StructuredLatentLinear from .layers import ( LatentDecorr, LinearSelfAttention, @@ -147,6 +148,11 @@ def __init__( k: int = 20, hidden_dims: Optional[list] = None, device: Optional[torch.device] = None, + *, + latent_structure: str = "none", + latent_orthogonality_weight: float = 1e-2, + latent_decorrelation_weight: float = 1e-2, + latent_ordering_probability: float = 0.5, **kwargs, ): self.k = _validate_positive_integer("k", k) @@ -155,7 +161,14 @@ def __init__( self.hidden_dims = _validate_positive_integer_sequence( "hidden_dims", hidden_dims ) - super().__init__(device=device, **kwargs) + super().__init__( + device=device, + latent_structure=latent_structure, + latent_orthogonality_weight=latent_orthogonality_weight, + latent_decorrelation_weight=latent_decorrelation_weight, + latent_ordering_probability=latent_ordering_probability, + **kwargs, + ) def _build_model(self, input_shape: Tuple, **kwargs) -> nn.Module: """Build the standard fully-connected autoencoder model.""" @@ -168,6 +181,11 @@ def _build_model(self, input_shape: Tuple, **kwargs) -> nn.Module: sample_shape = tuple(input_shape[1:]) n_features = int(np.prod(sample_shape)) + latent_structure = self.latent_structure + latent_orthogonality_weight = self.latent_orthogonality_weight + latent_decorrelation_weight = self.latent_decorrelation_weight + latent_ordering_probability = self.latent_ordering_probability + class StandardAutoencoderModel(nn.Module): def __init__(self, n_features, hidden_dims, k, sample_shape): super().__init__() @@ -181,7 +199,16 @@ def __init__(self, n_features, hidden_dims, k, sample_shape): encoder_layers.append(nn.BatchNorm1d(dim)) encoder_layers.append(nn.ReLU()) prev_dim = dim - encoder_layers.append(nn.Linear(prev_dim, k)) + encoder_layers.append( + StructuredLatentLinear( + prev_dim, + k, + mode=latent_structure, + orthogonality_weight=latent_orthogonality_weight, + decorrelation_weight=latent_decorrelation_weight, + ordering_probability=latent_ordering_probability, + ) + ) self.encoder = nn.Sequential(*encoder_layers) # Decoder @@ -618,13 +645,25 @@ def __init__( k: int = 20, hidden: Tuple[int, int] = (256, 128), device: Optional[torch.device] = None, + *, + latent_structure: str = "none", + latent_orthogonality_weight: float = 1e-2, + latent_decorrelation_weight: float = 1e-2, + latent_ordering_probability: float = 0.5, **kwargs, ): self.k = _validate_positive_integer("k", k) self.hidden = tuple( _validate_positive_integer_sequence("hidden", hidden, expected_length=2) ) - super().__init__(device=device, **kwargs) + super().__init__( + device=device, + latent_structure=latent_structure, + latent_orthogonality_weight=latent_orthogonality_weight, + latent_decorrelation_weight=latent_decorrelation_weight, + latent_ordering_probability=latent_ordering_probability, + **kwargs, + ) def _build_model(self, input_shape: Tuple, **kwargs) -> nn.Module: """Build the LSTM autoencoder model.""" @@ -637,6 +676,11 @@ def _build_model(self, input_shape: Tuple, **kwargs) -> nn.Module: n_features = input_shape[-1] seq_len = input_shape[1] # Infer from input shape + latent_structure = self.latent_structure + latent_orthogonality_weight = self.latent_orthogonality_weight + latent_decorrelation_weight = self.latent_decorrelation_weight + latent_ordering_probability = self.latent_ordering_probability + class LSTMAutoencoderModel(nn.Module): def __init__(self, seq_len, n_features, hidden, k): super().__init__() @@ -646,7 +690,14 @@ def __init__(self, seq_len, n_features, hidden, k): # Encoder self.lstm1 = nn.LSTM(n_features, hidden[0], batch_first=True) self.lstm2 = nn.LSTM(hidden[0], hidden[1], batch_first=True) - self.latent = nn.Linear(hidden[1], k) + self.latent = StructuredLatentLinear( + hidden[1], + k, + mode=latent_structure, + orthogonality_weight=latent_orthogonality_weight, + decorrelation_weight=latent_decorrelation_weight, + ordering_probability=latent_ordering_probability, + ) # Decoder self.latent_to_seq = nn.Linear(k, hidden[1]) @@ -735,10 +786,22 @@ def __init__( self, k: int = 20, device: Optional[torch.device] = None, + *, + latent_structure: str = "none", + latent_orthogonality_weight: float = 1e-2, + latent_decorrelation_weight: float = 1e-2, + latent_ordering_probability: float = 0.5, **kwargs, ): self.k = _validate_positive_integer("k", k) - super().__init__(device=device, **kwargs) + super().__init__( + device=device, + latent_structure=latent_structure, + latent_orthogonality_weight=latent_orthogonality_weight, + latent_decorrelation_weight=latent_decorrelation_weight, + latent_ordering_probability=latent_ordering_probability, + **kwargs, + ) def _build_model(self, input_shape: Tuple, **kwargs) -> nn.Module: """Build the CNN autoencoder model.""" @@ -759,6 +822,11 @@ def _build_model(self, input_shape: Tuple, **kwargs) -> nn.Module: pad_h = (4 - (H % 4)) % 4 pad_w = (4 - (W % 4)) % 4 + latent_structure = self.latent_structure + latent_orthogonality_weight = self.latent_orthogonality_weight + latent_decorrelation_weight = self.latent_decorrelation_weight + latent_ordering_probability = self.latent_ordering_probability + class CNNAutoencoderModel(nn.Module): def __init__(self, H, W, C, k, pad_h, pad_w): super().__init__() @@ -789,7 +857,14 @@ def __init__(self, H, W, C, k, pad_h, pad_w): self.flat_size = H_enc * W_enc * 64 self.fc1 = nn.Linear(self.flat_size, 256) - self.fc2 = nn.Linear(256, k) + self.fc2 = StructuredLatentLinear( + 256, + k, + mode=latent_structure, + orthogonality_weight=latent_orthogonality_weight, + decorrelation_weight=latent_decorrelation_weight, + ordering_probability=latent_ordering_probability, + ) # Decoder self.fc3 = nn.Linear(k, 256) @@ -947,6 +1022,11 @@ def __init__( depth_dec: int = 2, heads: int = 4, device: Optional[torch.device] = None, + *, + latent_structure: str = "none", + latent_orthogonality_weight: float = 1e-2, + latent_decorrelation_weight: float = 1e-2, + latent_ordering_probability: float = 0.5, **kwargs, ): self.k = _validate_positive_integer("k", k) @@ -959,7 +1039,14 @@ def __init__( self.heads = _validate_positive_integer("heads", heads) if self.d_model % self.heads != 0: raise ValueError("d_model must be divisible by heads.") - super().__init__(device=device, **kwargs) + super().__init__( + device=device, + latent_structure=latent_structure, + latent_orthogonality_weight=latent_orthogonality_weight, + latent_decorrelation_weight=latent_decorrelation_weight, + latent_ordering_probability=latent_ordering_probability, + **kwargs, + ) def _build_model(self, input_shape: Tuple, **kwargs) -> nn.Module: """Build the ViT autoencoder model.""" @@ -983,6 +1070,11 @@ def _build_model(self, input_shape: Tuple, **kwargs) -> nn.Module: N = Hp * Wp Pdim = self.patch_size * self.patch_size * C + latent_structure = self.latent_structure + latent_orthogonality_weight = self.latent_orthogonality_weight + latent_decorrelation_weight = self.latent_decorrelation_weight + latent_ordering_probability = self.latent_ordering_probability + class ViTAutoencoderModel(nn.Module): def __init__( self, @@ -1030,7 +1122,14 @@ def __init__( # Global bottleneck (latent k) self.global_pool = nn.AdaptiveAvgPool1d(1) - self.latent_k = nn.Linear(d_model, k) + self.latent_k = StructuredLatentLinear( + d_model, + k, + mode=latent_structure, + orthogonality_weight=latent_orthogonality_weight, + decorrelation_weight=latent_decorrelation_weight, + ordering_probability=latent_ordering_probability, + ) # Project back to token space for decoding self.dec_seed = nn.Linear(k, N * d_model) @@ -1200,6 +1299,11 @@ def __init__( self, k: int = 20, device: str | torch.device | None = None, + *, + latent_structure: str = "none", + latent_orthogonality_weight: float = 1e-2, + latent_decorrelation_weight: float = 1e-2, + latent_ordering_probability: float = 0.5, **kwargs, ): if "reconstruction_mode" in kwargs: @@ -1208,7 +1312,14 @@ def __init__( "ConvLSTMAutoencoder always reconstructs the full sequence." ) self.k = _validate_positive_integer("k", k) - super().__init__(device=device, **kwargs) + super().__init__( + device=device, + latent_structure=latent_structure, + latent_orthogonality_weight=latent_orthogonality_weight, + latent_decorrelation_weight=latent_decorrelation_weight, + latent_ordering_probability=latent_ordering_probability, + **kwargs, + ) def fit( self, @@ -1282,6 +1393,11 @@ def _build_model(self, input_shape: tuple, **kwargs) -> nn.Module: pad_w = (-width) % 4 latent_dim = self.k + latent_structure = self.latent_structure + latent_orthogonality_weight = self.latent_orthogonality_weight + latent_decorrelation_weight = self.latent_decorrelation_weight + latent_ordering_probability = self.latent_ordering_probability + class ConvLSTMAutoencoderModel(nn.Module): def __init__(self): super().__init__() @@ -1320,7 +1436,14 @@ def __init__(self): encoded_h = (height + pad_h) // 4 encoded_w = (width + pad_w) // 4 self.flat_size = encoded_h * encoded_w * 64 - self.latent = nn.Linear(self.flat_size, latent_dim) + self.latent = StructuredLatentLinear( + self.flat_size, + latent_dim, + mode=latent_structure, + orthogonality_weight=latent_orthogonality_weight, + decorrelation_weight=latent_decorrelation_weight, + ordering_probability=latent_ordering_probability, + ) self.fc_dec = nn.Linear(latent_dim, self.flat_size) self.upsample1 = nn.Upsample( @@ -1472,6 +1595,11 @@ def __init__( n_layers: int = 2, efficient_attention: str | None = "linear", device: str | torch.device | None = None, + *, + latent_structure: str = "none", + latent_orthogonality_weight: float = 1e-2, + latent_decorrelation_weight: float = 1e-2, + latent_ordering_probability: float = 0.5, **kwargs, ): if "reconstruction_mode" in kwargs: @@ -1491,7 +1619,14 @@ def __init__( if efficient_attention not in {"linear", None}: raise ValueError("efficient_attention must be 'linear' or None.") self.efficient_attention = efficient_attention - super().__init__(device=device, **kwargs) + super().__init__( + device=device, + latent_structure=latent_structure, + latent_orthogonality_weight=latent_orthogonality_weight, + latent_decorrelation_weight=latent_decorrelation_weight, + latent_ordering_probability=latent_ordering_probability, + **kwargs, + ) def fit( self, @@ -1571,6 +1706,11 @@ def _build_model(self, input_shape: tuple, **kwargs) -> nn.Module: n_layers = self.n_layers efficient_attention = self.efficient_attention + latent_structure = self.latent_structure + latent_orthogonality_weight = self.latent_orthogonality_weight + latent_decorrelation_weight = self.latent_decorrelation_weight + latent_ordering_probability = self.latent_ordering_probability + class HybridAutoencoderModel(nn.Module): def __init__(self): super().__init__() @@ -1611,7 +1751,14 @@ def __init__(self): self.transformer_blocks = self._make_blocks() self.global_pool_time = nn.AdaptiveAvgPool1d(1) - self.latent = nn.Linear(d_model, latent_dim) + self.latent = StructuredLatentLinear( + d_model, + latent_dim, + mode=latent_structure, + orthogonality_weight=latent_orthogonality_weight, + decorrelation_weight=latent_decorrelation_weight, + ordering_probability=latent_ordering_probability, + ) encoded_h = (height + pad_h) // 4 encoded_w = (width + pad_w) // 4 diff --git a/bluemath_tk/deeplearning/latent_structure.py b/bluemath_tk/deeplearning/latent_structure.py new file mode 100644 index 0000000..a64101c --- /dev/null +++ b/bluemath_tk/deeplearning/latent_structure.py @@ -0,0 +1,491 @@ +"""Reusable latent-structure regularization for BlueMath_tk autoencoders. + +The feature is deliberately architecture-agnostic. Existing encoder projection +layers remain unchanged; this module operates on the resulting latent scores +and, optionally, on the projection weight tensor. + +Modes +----- +none + Exact pass-through. No regularization and no ordered masking. +orthogonal + Penalize non-orthogonality of the final latent projection and latent + cross-correlation. +pca_like + Same penalties as ``orthogonal`` plus ordered nested latent dropout during + training. Earlier latent coordinates are therefore required to be useful + more often than later coordinates. + +This does NOT make a nonlinear autoencoder equivalent to PCA. The purpose is +to impose PCA-like geometric structure while preserving the nonlinear encoder +and decoder. +""" + +from __future__ import annotations + +import math +from typing import Any + +import numpy as np +import torch +import torch.nn as nn + +LATENT_STRUCTURE_MODES = ("none", "orthogonal", "pca_like") + + +def _validate_nonnegative_finite(name: str, value: float) -> float: + if ( + not isinstance(value, (int, float)) + or isinstance(value, bool) + or not math.isfinite(float(value)) + or float(value) < 0.0 + ): + raise ValueError(f"{name} must be a finite non-negative number.") + return float(value) + + +def _validate_probability(name: str, value: float) -> float: + value = _validate_nonnegative_finite(name, value) + if value > 1.0: + raise ValueError(f"{name} must be between 0 and 1.") + return value + + +def validate_latent_structure_mode(mode: str) -> str: + """Validate and normalize a latent-structure mode.""" + if not isinstance(mode, str): + raise TypeError("latent_structure must be a string.") + normalized = mode.strip().lower() + if normalized not in LATENT_STRUCTURE_MODES: + allowed = ", ".join(repr(item) for item in LATENT_STRUCTURE_MODES) + raise ValueError(f"latent_structure must be one of: {allowed}.") + return normalized + + +class LatentStructureRegularizer(nn.Module): + """Pass-through latent regularizer with optional ordered masking. + + Parameters + ---------- + k : int + Latent dimension. + mode : {"none", "orthogonal", "pca_like"} + Structural mode. + orthogonality_weight : float + Weight applied to the normalized projection orthogonality penalty. + decorrelation_weight : float + Weight applied to the off-diagonal latent correlation penalty. + ordering_probability : float + In ``pca_like`` mode and training mode, probability that each sample is + reconstructed from a random prefix of its latent vector. A value of + ``0.5`` means roughly half of samples use a shortened prefix while the + remainder retain all ``k`` coordinates. + eps : float + Dimensionless tolerance used to identify collapsed coordinates after rescaling. + + Notes + ----- + This module owns no trainable parameters or persistent buffers. Therefore, + adding it does not change the default model state_dict when mode="none", + and it can be introduced without invalidating legacy parameter tensors. + + The projection orthogonality penalty is + + mean((W W^T - I)^2) + + and is only defined when the latent projection weight is supplied. + + The latent decorrelation penalty is the mean squared off-diagonal correlation + of the *unmasked* latent scores. Ordered dropout is applied only after the + regularization terms have been computed. + """ + + _bluemath_latent_regularizer = True + + def __init__( + self, + k: int, + mode: str = "none", + orthogonality_weight: float = 1e-2, + decorrelation_weight: float = 1e-2, + ordering_probability: float = 0.5, + eps: float = 1e-8, + ): + super().__init__() + if not isinstance(k, int) or isinstance(k, bool) or k < 1: + raise ValueError("k must be a positive integer.") + self.k = k + self.mode = validate_latent_structure_mode(mode) + self.orthogonality_weight = _validate_nonnegative_finite( + "orthogonality_weight", orthogonality_weight + ) + self.decorrelation_weight = _validate_nonnegative_finite( + "decorrelation_weight", decorrelation_weight + ) + self.ordering_probability = _validate_probability( + "ordering_probability", ordering_probability + ) + self.eps = _validate_nonnegative_finite("eps", eps) + if self.eps == 0.0: + raise ValueError("eps must be greater than zero.") + + self._last_losses: dict[str, torch.Tensor] = {} + + @property + def enabled(self) -> bool: + """Return whether latent-structure regularization is enabled.""" + return self.mode != "none" + + def _validate_latent(self, z: torch.Tensor) -> None: + if not isinstance(z, torch.Tensor): + raise TypeError("Latent scores must be a PyTorch tensor.") + if z.ndim != 2: + raise ValueError( + "Latent scores must have shape " + f"(batch, {self.k}); got {tuple(z.shape)}." + ) + if z.shape[1] != self.k: + raise ValueError( + f"Latent scores must have width {self.k}; got {z.shape[1]}." + ) + if not torch.isfinite(z).all(): + raise FloatingPointError("Latent scores are not finite.") + + def _orthogonality_loss( + self, + projection_weight: torch.Tensor | None, + z: torch.Tensor, + ) -> torch.Tensor | None: + if self.orthogonality_weight == 0.0: + return None + if projection_weight is None: + raise ValueError( + "projection_weight is required when latent orthogonality is enabled." + ) + if not isinstance(projection_weight, torch.Tensor): + raise TypeError("projection_weight must be a PyTorch tensor.") + if projection_weight.ndim != 2: + raise ValueError("projection_weight must be a 2D tensor.") + if projection_weight.shape[0] != self.k: + raise ValueError( + "projection_weight must have shape " + f"({self.k}, in_features); got {tuple(projection_weight.shape)}." + ) + if projection_weight.shape[1] < self.k: + raise ValueError( + "Exact row orthogonality requires latent input dimension >= k; " + f"got in_features={projection_weight.shape[1]} and k={self.k}." + ) + if not torch.isfinite(projection_weight).all(): + raise FloatingPointError("projection_weight must be finite.") + + gram = projection_weight @ projection_weight.transpose(0, 1) + identity = torch.eye( + self.k, + dtype=gram.dtype, + device=gram.device, + ) + raw = torch.mean((gram - identity) ** 2) + return raw.to(dtype=z.dtype) * self.orthogonality_weight + + def _decorrelation_loss(self, z: torch.Tensor) -> torch.Tensor | None: + """Return a scale-invariant latent-correlation penalty. + + Noncollapsed coordinates are centered, independently rescaled, and + normalized to unit Euclidean norm. Their dot products are Pearson + correlations, so the loss is invariant to any representable nonzero + per-coordinate scaling. + + Pearson correlation is undefined for an exactly collapsed coordinate. + To prevent collapse from reducing the objective, every off-diagonal + pair involving such a coordinate receives the maximal squared- + correlation cost of 1.0. + """ + if ( + self.decorrelation_weight == 0.0 + or z.shape[0] < 2 + or self.k < 2 + ): + return None + + centered = z - z.mean(dim=0, keepdim=True) + coordinate_scale = torch.amax(torch.abs(centered), dim=0) + noncollapsed = coordinate_scale > 0 + safe_scale = torch.where( + noncollapsed, + coordinate_scale, + torch.ones_like(coordinate_scale), + ) + scaled = centered / safe_scale + + coordinate_norm = torch.linalg.vector_norm(scaled, dim=0) + active = noncollapsed & (coordinate_norm > self.eps) + safe_norm = torch.where( + active, + coordinate_norm, + torch.ones_like(coordinate_norm), + ) + normalized = scaled / safe_norm + + correlation = normalized.transpose(0, 1) @ normalized + squared_correlation = torch.clamp(correlation**2, max=1.0) + pair_active = active.unsqueeze(1) & active.unsqueeze(0) + guarded_squared_correlation = torch.where( + pair_active, + squared_correlation, + torch.ones_like(squared_correlation), + ) + + off_diagonal_mask = ~torch.eye( + self.k, + dtype=torch.bool, + device=z.device, + ) + raw = torch.mean( + guarded_squared_correlation[off_diagonal_mask] + ) + return raw * self.decorrelation_weight + + def _compute_losses( + self, + z: torch.Tensor, + projection_weight: torch.Tensor | None, + ) -> dict[str, torch.Tensor]: + if not self.enabled: + return {} + + losses: dict[str, torch.Tensor] = {} + orthogonality = self._orthogonality_loss(projection_weight, z) + if orthogonality is not None: + losses["latent_orthogonality"] = orthogonality + + decorrelation = self._decorrelation_loss(z) + if decorrelation is not None: + losses["latent_decorrelation"] = decorrelation + + return losses + + def apply_ordering(self, z: torch.Tensor) -> torch.Tensor: + """Apply per-sample nested latent dropout in training mode only.""" + self._validate_latent(z) + if ( + self.mode != "pca_like" + or not self.training + or self.ordering_probability == 0.0 + or self.k == 1 + ): + return z + + batch_size = z.shape[0] + shorten = ( + torch.rand(batch_size, device=z.device) < self.ordering_probability + ) + # Prefix lengths 1..k-1 for shortened samples. Non-shortened samples + # keep all k coordinates. + short_prefixes = torch.randint( + low=1, + high=self.k, + size=(batch_size,), + device=z.device, + ) + full_prefixes = torch.full( + (batch_size,), + self.k, + dtype=torch.long, + device=z.device, + ) + prefix_lengths = torch.where(shorten, short_prefixes, full_prefixes) + coordinates = torch.arange(self.k, device=z.device).unsqueeze(0) + mask = coordinates < prefix_lengths.unsqueeze(1) + return z * mask.to(dtype=z.dtype) + + def forward( + self, + z: torch.Tensor, + projection_weight: torch.Tensor | None = None, + *, + apply_ordering: bool = True, + ) -> torch.Tensor: + """Register current regularization losses and return structured scores.""" + self._validate_latent(z) + if not self.enabled: + self._last_losses = {} + return z + + self._last_losses = self._compute_losses(z, projection_weight) + if apply_ordering: + return self.apply_ordering(z) + return z + + def regularization_losses(self) -> dict[str, torch.Tensor]: + """Return losses from the most recent forward pass.""" + return dict(self._last_losses) + + def extra_repr(self) -> str: + """Return a concise module configuration representation.""" + return ( + f"k={self.k}, mode={self.mode!r}, " + f"orthogonality_weight={self.orthogonality_weight}, " + f"decorrelation_weight={self.decorrelation_weight}, " + f"ordering_probability={self.ordering_probability}" + ) + + + +class StructuredLatentLinear(nn.Linear): + """Linear projection that keeps legacy weight/bias state-dict names.""" + + def __init__( + self, + in_features: int, + out_features: int, + bias: bool = True, + *, + mode: str = "none", + orthogonality_weight: float = 1e-2, + decorrelation_weight: float = 1e-2, + ordering_probability: float = 0.5, + apply_ordering_in_forward: bool = True, + ): + super().__init__(in_features, out_features, bias=bias) + self.latent_regularizer = LatentStructureRegularizer( + k=out_features, + mode=mode, + orthogonality_weight=orthogonality_weight, + decorrelation_weight=decorrelation_weight, + ordering_probability=ordering_probability, + ) + self.apply_ordering_in_forward = bool(apply_ordering_in_forward) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Project inputs and apply the configured latent structure.""" + z = super().forward(x) + return self.latent_regularizer( + z, + projection_weight=self.weight, + apply_ordering=self.apply_ordering_in_forward, + ) + + def apply_ordering(self, z: torch.Tensor) -> torch.Tensor: + """Apply the configured training-only ordered prefix mask.""" + return self.latent_regularizer.apply_ordering(z) + + +def projection_orthogonality_error(weight: np.ndarray) -> float: + """Return normalized Frobenius error of row orthogonality.""" + array = np.asarray(weight, dtype=np.float64) + if array.ndim != 2: + raise ValueError("weight must be a 2D array.") + k, in_features = array.shape + if k < 1 or in_features < 1: + raise ValueError("weight dimensions must be positive.") + gram = array @ array.T + identity = np.eye(k, dtype=np.float64) + return float(np.sqrt(np.mean((gram - identity) ** 2))) + + +def compute_latent_diagnostics( + z: np.ndarray, + *, + variance_order_tolerance: float = 1e-12, + eps: float = 1e-12, +) -> dict[str, Any]: + """Compute interpretation diagnostics for a fitted latent representation. + + The function is intentionally independent of PCA. It reports whether latent + coordinates are decorrelated and whether variance is monotonically ordered. + Correlation summaries use only pairs for which both coordinates have + nonzero centered norm; exactly collapsed coordinates are reported + separately. + + Returns JSON-compatible values. + """ + array = np.asarray(z, dtype=np.float64) + if array.ndim != 2: + raise ValueError("z must have shape (n_samples, k).") + if array.shape[0] < 2: + raise ValueError("At least two samples are required.") + if array.shape[1] < 1: + raise ValueError("At least one latent dimension is required.") + if not np.isfinite(array).all(): + raise ValueError("z must contain only finite values.") + if variance_order_tolerance < 0 or not np.isfinite(variance_order_tolerance): + raise ValueError("variance_order_tolerance must be finite and non-negative.") + if eps <= 0 or not np.isfinite(eps): + raise ValueError("eps must be finite and positive.") + + centered = array - array.mean(axis=0, keepdims=True) + covariance = centered.T @ centered / float(array.shape[0] - 1) + variance = np.diag(covariance).copy() + + coordinate_scale = np.max(np.abs(centered), axis=0) + noncollapsed = coordinate_scale > 0.0 + safe_scale = np.where(noncollapsed, coordinate_scale, 1.0) + scaled = centered / safe_scale + + coordinate_norm = np.linalg.norm(scaled, axis=0) + active = noncollapsed & (coordinate_norm > eps) + safe_norm = np.where(active, coordinate_norm, 1.0) + normalized = scaled / safe_norm + correlation = np.clip(normalized.T @ normalized, -1.0, 1.0) + + if array.shape[1] == 1: + mean_abs_offdiag = 0.0 + max_abs_offdiag = 0.0 + else: + off_diagonal = ~np.eye(array.shape[1], dtype=bool) + defined_pairs = off_diagonal & np.outer(active, active) + if np.any(defined_pairs): + values = np.abs(correlation[defined_pairs]) + mean_abs_offdiag = float(values.mean()) + max_abs_offdiag = float(values.max()) + else: + mean_abs_offdiag = 0.0 + max_abs_offdiag = 0.0 + + variance_total = float(variance.sum()) + if variance_total > 0.0: + variance_fraction = variance / variance_total + else: + variance_fraction = np.zeros_like(variance) + + ordering_violations = int( + np.sum( + variance[1:] + > variance[:-1] + float(variance_order_tolerance) + ) + ) + + return { + "n_samples": int(array.shape[0]), + "k": int(array.shape[1]), + "mean": array.mean(axis=0).tolist(), + "variance": variance.tolist(), + "variance_fraction": variance_fraction.tolist(), + "cumulative_variance_fraction": np.cumsum(variance_fraction).tolist(), + "mean_abs_offdiag_correlation": mean_abs_offdiag, + "max_abs_offdiag_correlation": max_abs_offdiag, + "n_collapsed_coordinates": int(np.sum(~active)), + "collapsed_coordinates": np.flatnonzero(~active).astype(int).tolist(), + "variance_ordering_violations": ordering_violations, + "variance_monotonic_nonincreasing": ordering_violations == 0, + } + + +def prefix_latent(z: np.ndarray, n_components: int) -> np.ndarray: + """Zero all coordinates after ``n_components`` for prefix-reconstruction tests.""" + array = np.asarray(z) + if array.ndim != 2: + raise ValueError("z must have shape (n_samples, k).") + if ( + not isinstance(n_components, int) + or isinstance(n_components, bool) + or n_components < 1 + or n_components > array.shape[1] + ): + raise ValueError( + f"n_components must be an integer between 1 and {array.shape[1]}." + ) + result = array.copy() + result[:, n_components:] = 0 + return result diff --git a/bluemath_tk/deeplearning/spatiotemporal_autoencoders.py b/bluemath_tk/deeplearning/spatiotemporal_autoencoders.py index a0a9e0e..87f3174 100644 --- a/bluemath_tk/deeplearning/spatiotemporal_autoencoders.py +++ b/bluemath_tk/deeplearning/spatiotemporal_autoencoders.py @@ -8,6 +8,7 @@ import torch.nn.functional as functional from ._base_model import BaseDeepLearningModel +from .latent_structure import StructuredLatentLinear from .layers import ConvLSTM @@ -107,6 +108,11 @@ def __init__( n_heads: int = 4, n_layers: int = 2, device: str | torch.device | None = None, + *, + latent_structure: str = "none", + latent_orthogonality_weight: float = 1e-2, + latent_decorrelation_weight: float = 1e-2, + latent_ordering_probability: float = 0.5, **kwargs, ): if not isinstance(k, int) or isinstance(k, bool) or k < 1: @@ -140,7 +146,14 @@ def __init__( self.d_model = d_model self.n_heads = n_heads self.n_layers = n_layers - super().__init__(device=device, **kwargs) + super().__init__( + device=device, + latent_structure=latent_structure, + latent_orthogonality_weight=latent_orthogonality_weight, + latent_decorrelation_weight=latent_decorrelation_weight, + latent_ordering_probability=latent_ordering_probability, + **kwargs, + ) def fit( self, @@ -217,6 +230,11 @@ def _build_model(self, input_shape: tuple, **kwargs) -> nn.Module: n_heads = self.n_heads n_layers = self.n_layers + latent_structure = self.latent_structure + latent_orthogonality_weight = self.latent_orthogonality_weight + latent_decorrelation_weight = self.latent_decorrelation_weight + latent_ordering_probability = self.latent_ordering_probability + class SpatialTokenModel(nn.Module): def __init__(self): super().__init__() @@ -268,7 +286,14 @@ def __init__(self): ] ) self.latent_norm = nn.LayerNorm(d_model) - self.latent = nn.Linear(d_model, latent_dim) + self.latent = StructuredLatentLinear( + d_model, + latent_dim, + mode=latent_structure, + orthogonality_weight=latent_orthogonality_weight, + decorrelation_weight=latent_decorrelation_weight, + ordering_probability=latent_ordering_probability, + ) self.latent_to_tokens = nn.Linear(latent_dim, d_model) self.decoder_time_query = nn.Parameter( diff --git a/bluemath_tk/deeplearning/variational_autoencoders.py b/bluemath_tk/deeplearning/variational_autoencoders.py index b58bfa4..dfa7a3b 100644 --- a/bluemath_tk/deeplearning/variational_autoencoders.py +++ b/bluemath_tk/deeplearning/variational_autoencoders.py @@ -12,6 +12,7 @@ from tqdm import tqdm from ._base_model import BaseDeepLearningModel +from .latent_structure import StructuredLatentLinear class VariationalAutoencoder(BaseDeepLearningModel): @@ -52,6 +53,11 @@ def __init__( beta: float = 1.0, validation_mc_samples: int = 4, device: str | torch.device | None = None, + *, + latent_structure: str = "none", + latent_orthogonality_weight: float = 1e-2, + latent_decorrelation_weight: float = 1e-2, + latent_ordering_probability: float = 0.5, **kwargs, ): if not isinstance(k, int) or isinstance(k, bool) or k < 1: @@ -83,7 +89,14 @@ def __init__( self.hidden_dims = list(hidden_dims) self.beta = float(beta) self.validation_mc_samples = validation_mc_samples - super().__init__(device=device, **kwargs) + super().__init__( + device=device, + latent_structure=latent_structure, + latent_orthogonality_weight=latent_orthogonality_weight, + latent_decorrelation_weight=latent_decorrelation_weight, + latent_ordering_probability=latent_ordering_probability, + **kwargs, + ) def _build_model(self, input_shape: tuple, **kwargs) -> nn.Module: """Build the encoder, posterior parameterization, and decoder.""" @@ -99,6 +112,11 @@ def _build_model(self, input_shape: tuple, **kwargs) -> nn.Module: hidden_dims = tuple(self.hidden_dims) latent_dim = self.k + latent_structure = self.latent_structure + latent_orthogonality_weight = self.latent_orthogonality_weight + latent_decorrelation_weight = self.latent_decorrelation_weight + latent_ordering_probability = self.latent_ordering_probability + class VariationalAutoencoderModel(nn.Module): def __init__(self): super().__init__() @@ -117,7 +135,15 @@ def __init__(self): ) previous_dim = hidden_dim self.encoder = nn.Sequential(*encoder_layers) - self.mu_layer = nn.Linear(previous_dim, latent_dim) + self.mu_layer = StructuredLatentLinear( + previous_dim, + latent_dim, + mode=latent_structure, + orthogonality_weight=latent_orthogonality_weight, + decorrelation_weight=latent_decorrelation_weight, + ordering_probability=latent_ordering_probability, + apply_ordering_in_forward=False, + ) self.variance_layer = nn.Linear(previous_dim, latent_dim) decoder_layers: list[nn.Module] = [] @@ -206,7 +232,11 @@ def forward( mu, log_var = self.encode_distribution_forward(x) if stochastic is None: stochastic = self.training - z = self.reparameterize(mu, log_var) if stochastic else mu + if stochastic: + z = self.reparameterize(mu, log_var) + z = self.mu_layer.apply_ordering(z) + else: + z = mu return self.decode_forward(z) return VariationalAutoencoderModel() @@ -428,6 +458,12 @@ def _run_vae_epoch( reconstruction_losses = [] for _ in range(stochastic_samples): z = self.model.reparameterize(mu, log_var) + latent_projection = getattr(self.model, "mu_layer", None) + if ( + latent_projection is not None + and hasattr(latent_projection, "apply_ordering") + ): + z = latent_projection.apply_ordering(z) reconstruction = self.model.decode_forward(z) self._require_matching_output_shape( reconstruction, batch_y, "VAE reconstruction" @@ -444,7 +480,14 @@ def _run_vae_epoch( mean_reconstruction_loss = torch.stack(reconstruction_losses).mean() kl_loss = self.model.kl_divergence(mu, log_var) self._require_finite_loss(kl_loss, "VAE KL") - loss = mean_reconstruction_loss + self.beta * kl_loss + regularization_loss = self._model_regularization_loss( + mean_reconstruction_loss + ) + loss = ( + mean_reconstruction_loss + + self.beta * kl_loss + + regularization_loss + ) self._require_finite_loss(loss, "VAE total") deterministic_loss = None diff --git a/tests/deeplearning/test_latent_structure.py b/tests/deeplearning/test_latent_structure.py new file mode 100644 index 0000000..2565ab8 --- /dev/null +++ b/tests/deeplearning/test_latent_structure.py @@ -0,0 +1,376 @@ +"""Unit tests for PCA-like latent-structure utilities.""" + +import numpy as np +import pytest +import torch + +from bluemath_tk.deeplearning.latent_structure import ( + LatentStructureRegularizer, + compute_latent_diagnostics, + prefix_latent, + projection_orthogonality_error, + validate_latent_structure_mode, +) + + +def test_validate_latent_structure_mode(): + assert validate_latent_structure_mode("NONE") == "none" + assert validate_latent_structure_mode("orthogonal") == "orthogonal" + assert validate_latent_structure_mode("pca_like") == "pca_like" + with pytest.raises(ValueError): + validate_latent_structure_mode("pca") + + +def test_none_mode_is_exact_passthrough_and_has_no_state(): + z = torch.randn(5, 3, requires_grad=True) + regularizer = LatentStructureRegularizer(k=3, mode="none") + output = regularizer(z) + assert output is z + assert regularizer.regularization_losses() == {} + assert regularizer.state_dict() == {} + + +def test_identity_projection_has_zero_orthogonality_penalty(): + z = torch.randn(8, 3) + weight = torch.eye(3, requires_grad=True) + regularizer = LatentStructureRegularizer( + k=3, + mode="orthogonal", + orthogonality_weight=1.0, + decorrelation_weight=0.0, + ) + regularizer.train() + regularizer(z, projection_weight=weight) + losses = regularizer.regularization_losses() + assert torch.allclose( + losses["latent_orthogonality"], + torch.zeros((), dtype=losses["latent_orthogonality"].dtype), + atol=1e-7, + ) + + +def test_projection_requires_enough_input_features_for_row_orthogonality(): + z = torch.randn(5, 4) + weight = torch.randn(4, 3) + regularizer = LatentStructureRegularizer( + k=4, + mode="orthogonal", + orthogonality_weight=1.0, + decorrelation_weight=0.0, + ) + with pytest.raises(ValueError, match="in_features=3"): + regularizer(z, projection_weight=weight) + + +def test_decorrelation_penalty_is_differentiable(): + torch.manual_seed(7) + z = torch.randn(16, 4, requires_grad=True) + weight = torch.eye(4, requires_grad=True) + regularizer = LatentStructureRegularizer( + k=4, + mode="orthogonal", + orthogonality_weight=0.5, + decorrelation_weight=0.5, + ) + regularizer(z, projection_weight=weight) + loss = sum(regularizer.regularization_losses().values()) + loss.backward() + assert z.grad is not None + assert torch.isfinite(z.grad).all() + assert weight.grad is not None + assert torch.isfinite(weight.grad).all() + + +def test_pca_like_ordering_only_masks_during_training(): + z = torch.ones(128, 6) + weight = torch.randn(6, 8) + regularizer = LatentStructureRegularizer( + k=6, + mode="pca_like", + orthogonality_weight=0.0, + decorrelation_weight=0.0, + ordering_probability=1.0, + ) + + torch.manual_seed(11) + regularizer.train() + masked = regularizer(z, projection_weight=weight) + assert torch.all(masked[:, 0] == 1) + assert torch.any(masked[:, 1:] == 0) + + regularizer.eval() + unmasked = regularizer(z, projection_weight=weight) + assert torch.equal(unmasked, z) + + +def test_compute_latent_diagnostics_detects_uncorrelated_ordered_scores(): + # Orthogonal columns with decreasing variance. + z = np.array( + [ + [-3.0, -1.0], + [-1.0, 3.0], + [ 1.0, -3.0], + [ 3.0, 1.0], + ], + dtype=float, + ) + diagnostics = compute_latent_diagnostics(z) + assert diagnostics["k"] == 2 + assert diagnostics["mean_abs_offdiag_correlation"] < 1e-12 + + +def test_prefix_latent(): + z = np.arange(12, dtype=float).reshape(3, 4) + prefixed = prefix_latent(z, 2) + np.testing.assert_array_equal(prefixed[:, :2], z[:, :2]) + np.testing.assert_array_equal(prefixed[:, 2:], 0.0) + + +def test_projection_orthogonality_error_identity(): + assert projection_orthogonality_error(np.eye(4)) == pytest.approx(0.0) + +def test_decorrelation_penalty_is_scale_invariant(): + """Per-coordinate rescaling must not change the correlation penalty.""" + z = torch.tensor( + [ + [-2.0, -1.0, 0.5], + [-1.0, 0.2, 1.0], + [0.5, 0.7, 1.8], + [1.5, 1.2, 2.4], + [2.0, 1.8, 3.2], + ], + dtype=torch.float64, + ) + scales = torch.tensor([0.1, 7.0, 2.5], dtype=torch.float64) + + regularizer = LatentStructureRegularizer( + k=3, + mode="orthogonal", + orthogonality_weight=0.0, + decorrelation_weight=1.0, + ) + + regularizer(z, projection_weight=None) + base = regularizer.regularization_losses()["latent_decorrelation"] + + regularizer(z * scales, projection_weight=None) + scaled = regularizer.regularization_losses()["latent_decorrelation"] + + assert torch.allclose(base, scaled, atol=1e-10, rtol=1e-10) + + +def test_decorrelation_penalty_detects_redundant_latent_coordinates(): + """Highly correlated latent coordinates should receive a clear penalty.""" + x = torch.linspace(-2.0, 2.0, 32) + z = torch.stack([x, 2.0 * x, -0.5 * x], dim=1) + + regularizer = LatentStructureRegularizer( + k=3, + mode="orthogonal", + orthogonality_weight=0.0, + decorrelation_weight=1.0, + ) + regularizer(z, projection_weight=None) + loss = regularizer.regularization_losses()["latent_decorrelation"] + + assert loss.item() > 0.95 + + +def test_decorrelation_penalty_near_zero_for_orthogonal_scores(): + """Pairwise uncorrelated score columns should have negligible penalty.""" + z = torch.tensor( + [ + [-1.0, -1.0], + [-1.0, 1.0], + [1.0, -1.0], + [1.0, 1.0], + ], + dtype=torch.float64, + ) + + regularizer = LatentStructureRegularizer( + k=2, + mode="orthogonal", + orthogonality_weight=0.0, + decorrelation_weight=1.0, + ) + regularizer(z, projection_weight=None) + loss = regularizer.regularization_losses()["latent_decorrelation"] + + assert loss.item() < 1e-12 + +def test_decorrelation_penalty_remains_scale_invariant_across_old_eps_floor(): + """Extreme nonzero rescaling must not weaken the correlation objective.""" + z = torch.tensor( + [ + [-2.0, -1.0, 0.5], + [-1.0, 0.2, 1.0], + [0.5, 0.7, 1.8], + [1.5, 1.2, 2.4], + [2.0, 1.8, 3.2], + ], + dtype=torch.float32, + ) + scales = torch.tensor([1e-20, 1e20, 1e-10], dtype=torch.float32) + + regularizer = LatentStructureRegularizer( + k=3, + mode="orthogonal", + orthogonality_weight=0.0, + decorrelation_weight=1.0, + ) + + regularizer(z, projection_weight=None) + base = regularizer.regularization_losses()["latent_decorrelation"] + + regularizer(z * scales, projection_weight=None) + scaled = regularizer.regularization_losses()["latent_decorrelation"] + + assert torch.allclose(base, scaled, atol=1e-5, rtol=1e-5) + + +def test_decorrelation_penalty_does_not_reward_collapsed_coordinate(): + """Exact collapse must not reduce a correlated representation's loss.""" + x = torch.linspace(-2.0, 2.0, 32) + correlated = torch.stack([x, 2.0 * x, -0.5 * x], dim=1) + collapsed = correlated.clone() + collapsed[:, 1] = 0.0 + + regularizer = LatentStructureRegularizer( + k=3, + mode="orthogonal", + orthogonality_weight=0.0, + decorrelation_weight=1.0, + ) + + regularizer(correlated, projection_weight=None) + correlated_loss = regularizer.regularization_losses()[ + "latent_decorrelation" + ] + + regularizer(collapsed, projection_weight=None) + collapsed_loss = regularizer.regularization_losses()[ + "latent_decorrelation" + ] + + assert correlated_loss.item() == pytest.approx(1.0, abs=1e-6) + assert collapsed_loss.item() >= correlated_loss.item() - 1e-6 + + +def test_structured_linear_preserves_linear_state_keys_and_parameter_count(): + """Structured projection must remain checkpoint-compatible with nn.Linear.""" + from bluemath_tk.deeplearning.latent_structure import StructuredLatentLinear + + plain = torch.nn.Linear(7, 3) + structured = StructuredLatentLinear(7, 3, mode="none") + + assert list(structured.state_dict()) == list(plain.state_dict()) + assert sum(p.numel() for p in structured.parameters()) == sum( + p.numel() for p in plain.parameters() + ) + + +def test_ordering_boundary_probabilities_are_well_defined(): + """p=0 and k=1 are exact pass-through cases in training mode.""" + z = torch.randn(16, 3) + zero_probability = LatentStructureRegularizer( + k=3, + mode="pca_like", + orthogonality_weight=0.0, + decorrelation_weight=0.0, + ordering_probability=0.0, + ) + zero_probability.train() + assert torch.equal(zero_probability(z), z) + + z_single = torch.randn(16, 1) + one_dimension = LatentStructureRegularizer( + k=1, + mode="pca_like", + orthogonality_weight=0.0, + decorrelation_weight=0.0, + ordering_probability=1.0, + ) + one_dimension.train() + assert torch.equal(one_dimension(z_single), z_single) + +def test_near_collapsed_coordinate_has_finite_nonzero_gradient(): + """A tiny but nonzero active coordinate must still receive a gradient.""" + z = torch.tensor( + [ + [-2.0, 0.3e-8, 1.0], + [-1.0, -0.7e-8, 0.2], + [0.0, 1.2e-8, -0.5], + [1.0, 0.1e-8, 1.1], + [2.0, 1.8e-8, -1.2], + [3.0, -0.4e-8, 0.7], + ], + dtype=torch.float64, + requires_grad=True, + ) + regularizer = LatentStructureRegularizer( + k=3, + mode="orthogonal", + orthogonality_weight=0.0, + decorrelation_weight=1.0, + ) + + regularizer(z, projection_weight=None) + loss = regularizer.regularization_losses()["latent_decorrelation"] + loss.backward() + + assert z.grad is not None + assert torch.isfinite(z.grad).all() + assert torch.sum(torch.abs(z.grad[:, 1])).item() > 0.0 + + +def test_latent_diagnostics_correlation_is_scale_invariant(): + """Public Pearson diagnostics must survive extreme coordinate rescaling.""" + z = np.array( + [ + [-2.0, -1.0, 0.5], + [-1.0, 0.2, 1.0], + [0.5, 0.7, 1.8], + [1.5, 1.2, 2.4], + [2.0, 1.8, 3.2], + ], + dtype=np.float64, + ) + scales = np.array([1e-100, 1e100, 1e-50], dtype=np.float64) + + base = compute_latent_diagnostics(z) + scaled = compute_latent_diagnostics(z * scales) + + assert scaled["mean_abs_offdiag_correlation"] == pytest.approx( + base["mean_abs_offdiag_correlation"], + rel=1e-12, + abs=1e-12, + ) + assert scaled["max_abs_offdiag_correlation"] == pytest.approx( + base["max_abs_offdiag_correlation"], + rel=1e-12, + abs=1e-12, + ) + assert scaled["n_collapsed_coordinates"] == 0 + assert scaled["collapsed_coordinates"] == [] + + +def test_latent_diagnostics_reports_exactly_collapsed_coordinates(): + """Undefined Pearson coordinates must be explicit in public diagnostics.""" + z = np.array( + [ + [-2.0, 5.0, 1.0], + [-1.0, 5.0, 0.5], + [0.0, 5.0, -0.5], + [1.0, 5.0, -1.0], + [2.0, 5.0, 0.2], + ], + dtype=np.float64, + ) + + diagnostics = compute_latent_diagnostics(z) + + assert diagnostics["n_collapsed_coordinates"] == 1 + assert diagnostics["collapsed_coordinates"] == [1] + assert np.isfinite(diagnostics["mean_abs_offdiag_correlation"]) + assert np.isfinite(diagnostics["max_abs_offdiag_correlation"]) diff --git a/tests/deeplearning/test_pca_like_latent_integration.py b/tests/deeplearning/test_pca_like_latent_integration.py new file mode 100644 index 0000000..6b63169 --- /dev/null +++ b/tests/deeplearning/test_pca_like_latent_integration.py @@ -0,0 +1,302 @@ +"""Acceptance tests for the PCA-like latent feature. + +These tests are intended to be added together with the integration edits +described in docs/PCA_LIKE_LATENT_FEATURE_IMPLEMENTATION.md. + +Before those edits they will fail because the public constructors do not yet +accept ``latent_structure``. +""" + +import numpy as np +import pytest + +from bluemath_tk.deeplearning.autoencoders import ( + CNNAutoencoder, + ConvLSTMAutoencoder, + HybridConvLSTMTransformerAutoencoder, + LSTMAutoencoder, + SpatialTokenConvLSTMTransformerAutoencoder, + StandardAutoencoder, + VariationalAutoencoder, + VisionTransformerAutoencoder, +) + + +def _cases(): + return [ + ( + lambda: StandardAutoencoder( + k=3, + hidden_dims=[12, 8], + latent_structure="pca_like", + device="cpu", + ), + np.random.default_rng(1).normal(size=(12, 10)).astype("float32"), + ), + ( + lambda: LSTMAutoencoder( + k=3, + hidden=(8, 6), + latent_structure="pca_like", + device="cpu", + ), + np.random.default_rng(2).normal(size=(12, 3, 4)).astype("float32"), + ), + ( + lambda: CNNAutoencoder( + k=3, + latent_structure="pca_like", + device="cpu", + ), + np.random.default_rng(3).normal(size=(10, 1, 8, 8)).astype("float32"), + ), + ( + lambda: VisionTransformerAutoencoder( + k=3, + patch_size=4, + d_model=16, + depth_enc=1, + depth_dec=1, + heads=4, + latent_structure="pca_like", + device="cpu", + ), + np.random.default_rng(4).normal(size=(10, 1, 8, 8)).astype("float32"), + ), + ( + lambda: ConvLSTMAutoencoder( + k=3, + latent_structure="pca_like", + device="cpu", + ), + np.random.default_rng(5).normal(size=(12, 2, 1, 8, 8)).astype("float32"), + ), + ( + lambda: HybridConvLSTMTransformerAutoencoder( + k=3, + d_model=16, + n_heads=4, + n_layers=1, + latent_structure="pca_like", + device="cpu", + ), + np.random.default_rng(6).normal(size=(12, 2, 1, 8, 8)).astype("float32"), + ), + ( + lambda: SpatialTokenConvLSTMTransformerAutoencoder( + k=3, + spatial_pool_size=(1, 1), + d_model=16, + n_heads=4, + n_layers=1, + latent_structure="pca_like", + device="cpu", + ), + np.random.default_rng(7).normal(size=(12, 2, 1, 8, 8)).astype("float32"), + ), + ( + lambda: VariationalAutoencoder( + k=3, + hidden_dims=[12, 8], + beta=0.01, + validation_mc_samples=1, + latent_structure="pca_like", + device="cpu", + ), + np.random.default_rng(8).normal(size=(12, 10)).astype("float32"), + ), + ] + + +@pytest.mark.parametrize(("factory", "X"), _cases()) +def test_pca_like_mode_fits_predicts_and_encodes(factory, X): + np.random.seed(17) + model = factory() + model.fit( + X[:8], + epochs=1, + batch_size=4, + patience=1, + verbose=0, + validation_data=(X[8:], None), + ) + reconstruction = model.predict(X[8:], batch_size=4, verbose=0) + latent = model.encode(X[8:], batch_size=4, verbose=0) + assert reconstruction.shape == X[8:].shape + assert latent.shape == (len(X[8:]), model.k) + assert np.isfinite(reconstruction).all() + assert np.isfinite(latent).all() + + +@pytest.mark.parametrize(("factory", "X"), _cases()) +def test_default_none_mode_remains_available(factory, X): + configured = factory() + default_model = type(configured)(k=configured.k, device="cpu") + assert default_model.latent_structure == "none" + + +def test_invalid_mode_rejected_early(): + with pytest.raises(ValueError): + StandardAutoencoder(k=2, latent_structure="not-a-mode") + +def test_legacy_positional_device_calls_remain_valid(): + """New latent options must not consume the historical device position.""" + models = [ + StandardAutoencoder(3, [8], "cpu"), + LSTMAutoencoder(3, (8, 4), "cpu"), + CNNAutoencoder(3, "cpu"), + VisionTransformerAutoencoder(3, 4, 16, 1, 1, 4, "cpu"), + ConvLSTMAutoencoder(3, "cpu"), + HybridConvLSTMTransformerAutoencoder( + 3, + 16, + 4, + 1, + "linear", + "cpu", + ), + SpatialTokenConvLSTMTransformerAutoencoder( + 3, + (2, 2), + 16, + 4, + 1, + "cpu", + ), + VariationalAutoencoder(3, [8], 1.0, 1, "cpu"), + ] + + for model in models: + assert model.device.type == "cpu" + assert model.latent_structure == "none" + + +def test_new_latent_constructor_options_are_keyword_only(): + """All public latent options should be keyword-only compatibility additions.""" + import inspect + + classes = ( + StandardAutoencoder, + LSTMAutoencoder, + CNNAutoencoder, + VisionTransformerAutoencoder, + ConvLSTMAutoencoder, + HybridConvLSTMTransformerAutoencoder, + SpatialTokenConvLSTMTransformerAutoencoder, + VariationalAutoencoder, + ) + latent_names = ( + "latent_structure", + "latent_orthogonality_weight", + "latent_decorrelation_weight", + "latent_ordering_probability", + ) + + for cls in classes: + parameters = inspect.signature(cls.__init__).parameters + assert parameters["device"].kind is inspect.Parameter.POSITIONAL_OR_KEYWORD + for name in latent_names: + assert parameters[name].kind is inspect.Parameter.KEYWORD_ONLY + + +def test_vae_direct_training_forward_applies_ordered_prefix_mask(): + """The ordinary VAE model forward must match custom-fit ordering semantics.""" + import torch + + outer = VariationalAutoencoder( + k=3, + hidden_dims=[8], + beta=0.0, + validation_mc_samples=1, + latent_structure="pca_like", + latent_orthogonality_weight=0.0, + latent_decorrelation_weight=0.0, + latent_ordering_probability=1.0, + device="cpu", + ) + inner = outer._build_model((4, 6)) + inner.train() + + captured = {} + original_decode = inner.decode_forward + + def capture_decode(z): + captured["z"] = z.detach().clone() + return original_decode(z) + + inner.decode_forward = capture_decode + x = torch.randn(4, 6) + + torch.manual_seed(123) + _ = inner(x) + + assert "z" in captured + assert torch.all(captured["z"][:, -1] == 0) + + mu, _ = inner.encode_distribution_forward(x) + encoded = inner.encode_forward(x) + assert torch.equal(mu, encoded) + assert torch.any(mu[:, -1] != 0) + + +def test_built_model_rejects_latent_checkpoint_config_mismatch(tmp_path): + """A built receiver must not silently retain a conflicting latent mode.""" + source = StandardAutoencoder( + k=2, + hidden_dims=[4], + latent_structure="pca_like", + latent_orthogonality_weight=1.0, + latent_decorrelation_weight=0.1, + latent_ordering_probability=0.5, + device="cpu", + ) + build_shape = (4, 6) + source._build_input_shape = build_shape + source.model = source._build_model(build_shape).to(source.device) + source.is_fitted = True + + checkpoint = tmp_path / "structured.pt" + source.save_pytorch_model(checkpoint) + + receiver = StandardAutoencoder( + k=2, + hidden_dims=[4], + latent_structure="none", + device="cpu", + ) + receiver._build_input_shape = build_shape + receiver.model = receiver._build_model(build_shape).to(receiver.device) + + with pytest.raises(ValueError, match="latent configuration"): + receiver.load_pytorch_model(checkpoint, weights_only=False) + + +def test_structured_checkpoint_restores_config_when_rebuilt(tmp_path): + """Self-describing loading must restore the saved latent configuration.""" + source = StandardAutoencoder( + k=2, + hidden_dims=[4], + latent_structure="pca_like", + latent_orthogonality_weight=1.0, + latent_decorrelation_weight=0.1, + latent_ordering_probability=0.5, + device="cpu", + ) + build_shape = (4, 6) + source._build_input_shape = build_shape + source.model = source._build_model(build_shape).to(source.device) + source.is_fitted = True + + checkpoint = tmp_path / "structured-roundtrip.pt" + source.save_pytorch_model(checkpoint) + + loaded = StandardAutoencoder.from_pytorch_model( + checkpoint, + device="cpu", + weights_only=False, + ) + + assert loaded.latent_structure == "pca_like" + assert loaded.latent_orthogonality_weight == pytest.approx(1.0) + assert loaded.latent_decorrelation_weight == pytest.approx(0.1) + assert loaded.latent_ordering_probability == pytest.approx(0.5)