Coverage for cosmolayer/cosmolayer/cosmolightning.py: 80%
158 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 00:09 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 00:09 +0000
1"""
2.. module:: cosmolayer.cosmolayer.cosmolightning
3 :synopsis: PyTorch Lightning module for batched CosmoLayer training.
4"""
6from __future__ import annotations
8from collections.abc import Sequence
10import numpy as np
11import torch
12from lightning import pytorch as pl
13from numpy.typing import NDArray
14from torch import distributed as td
15from torch.nn import functional as F
16from torchmetrics import MeanAbsoluteError, MeanSquaredError, R2Score
18from .cosmodata import InputsType
19from .layer import CosmoLayer
20from .utils import LossFn, is_loss_function
22EPSILON = 1e-8
25class LogGammaLightningModule(pl.LightningModule):
26 """PyTorch Lightning module for batched training of a learnable
27 :class:`~cosmolayer.CosmoLayer`.
29 This class is the canonical high-level training interface for CosmoLayer.
30 It constructs an internal :class:`~cosmolayer.CosmoLayer` with learnable
31 interaction matrices and defines the optimization, training, validation,
32 test, and prediction logic.
34 The targets are the log-activity coefficients of the components. In order to
35 handle other tasks, the user must subclass :class:`LogGammaLightningModule` and
36 override the :meth:`~LogGammaLightningModule.predict_from_log_gamma` method. For
37 instance:
39 .. code-block:: python
41 from scipy.constants import R
43 class ExcessGibbsLightningModule(LogGammaLightningModule):
44 def predict_from_log_gamma(self, T, x, log_gamma):
45 return (R * T * (x * log_gamma).sum(dim=-1)).unsqueeze(-1)
47 The module is batch-first throughout. All inputs must represent a minibatch
48 of ``b`` datapoints, and the returned predictions must have leading
49 dimension ``b``. Targets must have the same shape as the predictions.
51 Parameters
52 ----------
53 num_segment_types : int
54 Number of COSMO segment types.
55 temperature_exponents : tuple[int, ...]
56 Exponents defining the temperature dependence of the interaction
57 matrices.
58 area_per_segment : float
59 Area associated with one segment.
60 reference_temperature : float, optional
61 Reference temperature used by :class:`CosmoLayer`.
62 Default is ``298.15``.
63 max_iter : int, optional
64 Maximum number of internal fixed-point or iterative solver steps used
65 by :class:`CosmoLayer`. Default is ``100``.
66 learning_rate : float, optional
67 Learning rate for the Adam optimizer. Default is ``1e-3``.
68 weight_decay : float, optional
69 Weight decay for the Adam optimizer. Default is ``0.0``.
70 loss_function : str, optional
71 Loss function used in training, validation, and test steps. Must be a
72 valid loss function from :mod:`torch.nn.functional`.
73 Default is ``"mse_loss"``.
74 initialization : Sequence[NDArray[np.float64]] | int, optional
75 Initialization for the learnable interaction matrices.
77 - If an ``int`` is provided, it is interpreted as the random seed used
78 to sample one matrix per temperature exponent from a standard normal
79 distribution.
80 - If a sequence of NumPy arrays is provided, it must contain exactly
81 one array per temperature exponent, and each array must have shape
82 ``(num_segment_types, num_segment_types)``.
84 Default is ``42``.
86 Examples
87 --------
88 >>> import torch
89 >>> from importlib.resources import files
90 >>> import cosmolayer as cl
91 >>> from cosmolayer import cosmosac
92 >>> model = cosmosac.CosmoSac2010Model
93 >>> module = LogGammaLightningModule(
94 ... num_segment_types=model.num_segment_types,
95 ... temperature_exponents=model.temperature_exponents,
96 ... area_per_segment=model.area_per_segment,
97 ... )
98 >>> solute_path = files("cosmolayer.data") / "NCCO.cosmo"
99 >>> solvent_path = files("cosmolayer.data") / "O.cosmo"
100 >>> datapoint = cosmosac.CosmoSacMixtureDatapoint(
101 ... cosmo_files=[solute_path, solvent_path],
102 ... mole_fractions=[0.2, 0.8],
103 ... temperature=298.15,
104 ... targets=[-0.2, 0.02],
105 ... model=model,
106 ... )
107 >>> single_inputs = datapoint.get_inputs()
108 >>> batched_inputs = tuple(x.unsqueeze(0) for x in single_inputs)
109 >>> preds = module(batched_inputs)
110 >>> preds.shape
111 torch.Size([1, 2])
112 """
114 def __init__( # noqa: PLR0913, PLR0917
115 self,
116 num_segment_types: int,
117 temperature_exponents: Sequence[int],
118 area_per_segment: float,
119 reference_temperature: float = 298.15,
120 max_iter: int = 100,
121 learning_rate: float = 1e-3,
122 weight_decay: float = 0.0,
123 normalize_targets: bool = False,
124 loss_function: str = "mse_loss",
125 initialization: Sequence[NDArray[np.float64]] | int = 42,
126 ) -> None:
127 super().__init__()
129 if num_segment_types <= 0:
130 raise ValueError("num_segment_types must be a positive integer")
131 if len(temperature_exponents) == 0:
132 raise ValueError("temperature_exponents must not be empty")
133 if area_per_segment <= 0.0:
134 raise ValueError("area_per_segment must be positive")
135 if reference_temperature <= 0.0:
136 raise ValueError("reference_temperature must be positive")
137 if max_iter <= 0:
138 raise ValueError("max_iter must be a positive integer")
139 if learning_rate <= 0.0:
140 raise ValueError("learning_rate must be positive")
141 if weight_decay < 0.0:
142 raise ValueError("weight_decay must be non-negative")
143 loss_callable = getattr(F, loss_function, None)
144 if not is_loss_function(loss_callable):
145 raise ValueError(f"Unsupported loss_function '{loss_function}'.")
147 self.save_hyperparameters(ignore=["initialization"])
148 self.normalize_targets = normalize_targets
149 self.learning_rate = learning_rate
150 self.weight_decay = weight_decay
151 self.loss_function: LossFn = loss_callable
153 initial_matrices = self._build_initial_matrices(
154 initialization=initialization,
155 num_segment_types=num_segment_types,
156 num_matrices=len(temperature_exponents),
157 )
159 self.cosmo_layer = CosmoLayer(
160 interaction_matrices=initial_matrices,
161 exponents=temperature_exponents,
162 area_per_segment=area_per_segment,
163 reference_temperature=reference_temperature,
164 max_iter=max_iter,
165 learn_matrices=True,
166 )
168 self.test_mae = MeanAbsoluteError()
169 self.test_rmse = MeanSquaredError(squared=False)
170 self.test_r2 = R2Score()
172 self.register_buffer("target_mean", torch.tensor(0.0))
173 self.register_buffer("target_std", torch.tensor(1.0))
175 @staticmethod
176 def _build_initial_matrices(
177 initialization: Sequence[NDArray[np.float64]] | int,
178 num_segment_types: int,
179 num_matrices: int,
180 ) -> list[NDArray[np.float64]]:
181 """Create and validate the initial interaction matrices."""
182 if isinstance(initialization, int):
183 rng = np.random.default_rng(initialization)
184 return [
185 rng.normal(size=(num_segment_types, num_segment_types))
186 for _ in range(num_matrices)
187 ]
189 matrices = [np.asarray(matrix, dtype=np.float64) for matrix in initialization]
191 if len(matrices) != num_matrices:
192 raise ValueError(
193 "initialization must contain exactly one matrix per temperature "
194 f"exponent: expected {num_matrices}, got {len(matrices)}"
195 )
197 expected_shape = (num_segment_types, num_segment_types)
198 for index, matrix in enumerate(matrices):
199 if matrix.shape != expected_shape:
200 raise ValueError(
201 "Each initialization matrix must have shape "
202 f"{expected_shape}; matrix {index} has shape {matrix.shape}"
203 )
204 if not np.isfinite(matrix).all():
205 raise ValueError(
206 f"Initialization matrix {index} contains non-finite values"
207 )
209 return matrices
211 @staticmethod
212 def _infer_batch_size(predictions: torch.Tensor, targets: torch.Tensor) -> int:
213 """Infer the minibatch size from prediction and target tensors."""
214 if predictions.ndim == 0 or targets.ndim == 0:
215 raise ValueError(
216 "Predictions and targets must be batched tensors with a leading "
217 "batch dimension"
218 )
219 if predictions.shape != targets.shape:
220 raise ValueError(
221 "Predictions and targets must have the same shape; "
222 f"got {predictions.shape} and {targets.shape}"
223 )
224 return int(targets.shape[0])
226 @torch.no_grad()
227 def _compute_target_statistics(self) -> None:
228 trainer = self.trainer
229 datamodule = getattr(trainer, "datamodule", None)
230 train_dl_from_dm = getattr(datamodule, "train_dataloader", None)
231 if callable(train_dl_from_dm):
232 dataloader = train_dl_from_dm()
233 else:
234 dataloader = getattr(trainer, "train_dataloader", None)
236 if dataloader is None:
237 raise ValueError(
238 "Training dataloader is unavailable; cannot normalize targets"
239 )
241 count = torch.tensor(0.0)
242 target_sum: torch.Tensor | None = None
243 target_sumsq: torch.Tensor | None = None
245 for batch in dataloader:
246 _, targets = batch
247 targets = targets.detach()
249 batch_count = torch.tensor(float(targets.shape[0]), device=targets.device)
250 batch_sum = targets.sum(dim=0)
251 batch_sumsq = (targets**2).sum(dim=0)
253 if target_sum is None:
254 count = count.to(targets.device)
255 target_sum = torch.zeros_like(batch_sum)
256 target_sumsq = torch.zeros_like(batch_sumsq)
258 count = count + batch_count
259 target_sum = target_sum + batch_sum
260 target_sumsq = target_sumsq + batch_sumsq
262 if target_sum is None or count.item() == 0:
263 raise ValueError("Training dataloader is empty; cannot normalize targets")
265 if td.is_available() and td.is_initialized():
266 td.all_reduce(count, op=td.ReduceOp.SUM)
267 td.all_reduce(target_sum, op=td.ReduceOp.SUM)
268 td.all_reduce(target_sumsq, op=td.ReduceOp.SUM)
270 if target_sum is None or target_sumsq is None:
271 raise ValueError("Training dataloader is empty; cannot normalize targets")
273 mean = target_sum / count
274 variance = torch.clamp(target_sumsq / count - mean**2, min=0.0)
275 std = torch.sqrt(variance + EPSILON)
277 self.target_mean = mean.to(self.device)
278 self.target_std = std.to(self.device)
280 def forward(self, inputs: InputsType) -> torch.Tensor:
281 """Compute predictions for a minibatch of datapoints.
283 Parameters
284 ----------
285 inputs : InputsType
286 Batched input tuple ``(temperature, mole_fractions, areas, volumes,
287 probabilities)``. All tensors must be batch-first and represent the
288 same minibatch of size ``b``.
290 Returns
291 -------
292 torch.Tensor
293 Batched predictions with leading dimension ``b``.
294 """
295 temperature, mole_fractions, areas, volumes, probabilities = inputs
296 log_gamma: torch.Tensor = self.cosmo_layer(
297 temperature, mole_fractions, areas, volumes, probabilities
298 )
299 return self.predict_from_log_gamma(temperature, mole_fractions, log_gamma)
301 def predict_from_log_gamma(
302 self,
303 T: torch.Tensor,
304 x: torch.Tensor,
305 log_gamma: torch.Tensor,
306 ) -> torch.Tensor:
307 """Convert log-activity coefficients to final predictions.
309 Parameters
310 ----------
311 T : torch.Tensor
312 Temperature in the same units as the reference temperature.
313 Shape: (...,).
314 x : torch.Tensor
315 Mole fractions of the components. Must sum to 1.
316 Shape: (..., num_components).
317 log_gamma : torch.Tensor
318 Logarithms of the activity coefficients.
319 Shape: (..., num_components).
321 Returns
322 -------
323 torch.Tensor
324 Final predictions.
325 """
326 return log_gamma
328 def configure_optimizers(self) -> torch.optim.Optimizer:
329 """Configure the optimizer used during training.
331 Returns
332 -------
333 torch.optim.Optimizer
334 Adam optimizer over all module parameters.
335 """
336 return torch.optim.Adam(
337 self.parameters(),
338 lr=self.learning_rate,
339 weight_decay=self.weight_decay,
340 )
342 def on_fit_start(self) -> None:
343 if self.normalize_targets:
344 self._compute_target_statistics()
346 def training_step(
347 self, batch: tuple[InputsType, torch.Tensor], batch_idx: int
348 ) -> torch.Tensor:
349 """Run one training step on a minibatch.
351 Parameters
352 ----------
353 batch : tuple[InputsType, torch.Tensor]
354 Batched inputs and batched ground-truth targets. Targets must have
355 the same shape as the model predictions, with leading dimension
356 equal to the minibatch size.
357 batch_idx : int
358 Index of the current batch.
360 Returns
361 -------
362 torch.Tensor
363 Training loss for the batch.
364 """
365 inputs, targets = batch
366 predictions = self(inputs)
367 batch_size = self._infer_batch_size(predictions, targets)
368 if self.normalize_targets:
369 target_mean = self.target_mean
370 target_std = self.target_std
371 targets = (targets - target_mean) / target_std
372 predictions = (predictions - target_mean) / target_std
373 loss: torch.Tensor = self.loss_function(predictions, targets)
374 self.log(
375 "train_loss",
376 loss,
377 on_step=False,
378 on_epoch=True,
379 batch_size=batch_size,
380 )
381 return loss
383 def validation_step(
384 self, batch: tuple[InputsType, torch.Tensor], batch_idx: int
385 ) -> torch.Tensor:
386 """Run one validation step on a minibatch.
388 Parameters
389 ----------
390 batch : tuple[InputsType, torch.Tensor]
391 Batched inputs and batched ground-truth targets. Targets must have
392 the same shape as the model predictions, with leading dimension
393 equal to the minibatch size.
394 batch_idx : int
395 Index of the current batch.
397 Returns
398 -------
399 torch.Tensor
400 Validation loss for the batch.
401 """
402 inputs, targets = batch
403 predictions = self(inputs)
404 batch_size = self._infer_batch_size(predictions, targets)
405 if self.normalize_targets:
406 target_mean = self.target_mean
407 target_std = self.target_std
408 targets = (targets - target_mean) / target_std
409 predictions = (predictions - target_mean) / target_std
410 loss: torch.Tensor = self.loss_function(predictions, targets)
411 self.log(
412 "val_loss",
413 loss,
414 on_step=False,
415 on_epoch=True,
416 batch_size=batch_size,
417 prog_bar=True,
418 )
419 return loss
421 def test_step(
422 self, batch: tuple[InputsType, torch.Tensor], batch_idx: int
423 ) -> torch.Tensor:
424 """Run one test step on a minibatch and update regression metrics.
426 Parameters
427 ----------
428 batch : tuple[InputsType, torch.Tensor]
429 Batched inputs and batched ground-truth targets. Targets must have
430 the same shape as the model predictions, with leading dimension
431 equal to the minibatch size.
432 batch_idx : int
433 Index of the current batch.
435 Returns
436 -------
437 torch.Tensor
438 Test loss for the batch.
439 """
440 inputs, targets = batch
441 predictions = self(inputs)
442 batch_size = self._infer_batch_size(predictions, targets)
443 loss_predictions = predictions
444 loss_targets = targets
445 if self.normalize_targets:
446 target_mean = self.target_mean
447 target_std = self.target_std
448 loss_targets = (targets - target_mean) / target_std
449 loss_predictions = (predictions - target_mean) / target_std
450 loss: torch.Tensor = self.loss_function(loss_predictions, loss_targets)
452 self.test_mae.update(predictions, targets)
453 self.test_rmse.update(predictions, targets)
454 self.test_r2.update(predictions, targets)
456 self.log_dict(
457 {
458 "test_loss": loss,
459 "test_mae": self.test_mae,
460 "test_rmse": self.test_rmse,
461 "test_r2": self.test_r2,
462 },
463 on_step=False,
464 on_epoch=True,
465 batch_size=batch_size,
466 )
467 return loss