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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
144 changes: 134 additions & 10 deletions bluemath_tk/deeplearning/_base_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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.

Expand All @@ -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):
Expand Down Expand Up @@ -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")
Expand All @@ -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))
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -730,19 +840,26 @@ 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()
optimizer.step()
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
Expand All @@ -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
Expand Down
Loading
Loading