Coverage for src/crispatt/crispatt.py: 58%
137 statements
« prev ^ index » next coverage.py v7.15.3, created at 2026-09-03 17:23 +1200
« prev ^ index » next coverage.py v7.15.3, created at 2026-09-03 17:23 +1200
1from collections.abc import Iterator, Sequence
2import functools
3import itertools
4import logging
5import typing
6from typing import Self
7import warnings
9import attrs
10import numpy as np
11from wisib import Domain, Floe, Ice, Ocean, Wave
13if typing.TYPE_CHECKING:
14 from wisib import FloeCoupled
16from ._fsd import _FSDHandler
17from ._noise import _generate_noise_on_lengths, _RandomVariableProtocol
18from ._stft import LocalFT
20logger = logging.getLogger(__name__)
23def floe_deflection(
24 floe: FloeCoupled,
25 wave: Wave,
26 x,
27 return_modes: bool = False,
28):
29 x = np.atleast_1d(x)[:, None]
30 coefs_pos, coefs_neg = floe.amplitude_iw_pos, floe.amplitude_iw_neg
31 wavenumbers = floe.ice.wave_numbers[: len(coefs_pos)]
32 _t = (
33 np.real(
34 1j
35 * wavenumbers
36 * (
37 coefs_pos * np.exp(1j * wavenumbers * x)
38 + coefs_neg * np.exp(-1j * wavenumbers * (x - floe.length))
39 )
40 * np.tanh(wavenumbers * floe.ice.dud)
41 )
42 / wave.angular_frequency
43 )
45 if return_modes:
46 return _t
47 return _t.sum(axis=1)
50# NOTE: most of these functions could be made static methods of Ensemble.
53def _compute_edges(
54 lengths_ensemble: np.ndarray, concentration: float, ice_edge: float
55) -> np.ndarray:
56 edges_ensemble = np.zeros_like(lengths_ensemble)
57 edges_ensemble[:, 1:] = np.cumsum(lengths_ensemble[:, :-1], axis=1)
58 edges_ensemble /= concentration
59 edges_ensemble += ice_edge
60 return edges_ensemble
63def _n_floes_from_width(
64 domain_width: float, average_length: float, concentration: float
65) -> int:
66 return np.ceil(domain_width / (average_length / concentration)).astype(int)
69def _init_seed_sequence(
70 seed: None | int | np.random.SeedSequence,
71) -> np.random.SeedSequence:
72 if isinstance(seed, np.random.SeedSequence):
73 seed_sequence = seed
74 else:
75 seed_sequence = np.random.SeedSequence(seed)
76 logger.info(f"Initialising ensemble with {seed_sequence.entropy:#x}.")
77 return seed_sequence
80@attrs.define
81class Ensemble:
82 experiments: list[Experiment]
83 seed_sequence: np.random.SeedSequence
84 fsd_handler: _FSDHandler
86 @property
87 def average_length(self) -> float:
88 return self.fsd_handler.average_length
90 @classmethod
91 def from_width(
92 cls,
93 n_members: int,
94 width: float,
95 *,
96 fsd_handler: _FSDHandler,
97 noise_distribution: _RandomVariableProtocol,
98 wave: Wave,
99 ocean: Ocean,
100 n_coupling_modes: tuple[int, int, int],
101 ices: Ice | Sequence[Ice],
102 ice_edge: float = 0,
103 concentration: float = 0.99,
104 seed: None | np.random.SeedSequence | int = None,
105 ices_weights: np.ndarray | None = None,
106 ) -> Self:
107 """Create an ensemble by specifying the width of the domains.
109 Parameters
110 ----------
111 n_members : int
112 Cardinality of the ensemble.
113 width : float
114 Physical extent of the domains, in m.
115 fsd_handler : _FSDHandler
116 Object parametrising the floe size distribution.
117 noise_distribution : stats._distn_infrastructure.rv_frozen
118 Distribution specifying the noise added to floe lengths.
119 wave : Wave
120 Object holding information on the wave forcing.
121 ocean : Ocean
122 Object holding information on the bearing fluid.
123 n_coupling_modes : tuple[int, int, int]
124 Number of terms kept to compute the scattering kernel and accessory
125 function, and to compute vertical displacement, respectively.
126 ice : Ice
127 Object holding information on the ice characteristics.
128 ice_edge : float, default: 0.0
129 Location of the ice edge, that is, the open--ice-covered ocean
130 boundary, in m.
131 concentration : float, default: 0.99
132 Ice concentration.
133 seed : np.random.SeedSequence or int, optional
134 A seed to initialise the random generator.
135 If a :class:`~numpy.random.SeedSequence` instance, will be used directly.
136 Otherwise, will be used to instantiate such an object.
138 Returns
139 -------
140 Ensemble
141 The initialised ensemble.
143 """
144 seed_sequence = _init_seed_sequence(seed)
145 n_floes = _n_floes_from_width(width, fsd_handler.average_length, concentration)
146 lengths = fsd_handler.compute_lengths(n_floes)
147 lengths_ensemble = cls._add_noise_to_lengths(
148 lengths, n_members, seed_sequence, noise_distribution
149 )
150 edges_ensemble = _compute_edges(lengths_ensemble, concentration, ice_edge)
152 # ices = itertools.repeat(ice)
153 ices_ensemble = cls._generate_ices_ensemble(
154 ices, n_floes, n_members, seed_sequence, ices_weights
155 )
157 experiments = [
158 Experiment.from_objects(
159 wave,
160 ocean,
161 n_coupling_modes,
162 [
163 Floe(left_edge=_e, length=_l, ice=_ice)
164 for _e, _l, _ice in zip(edges, lengths, ices_)
165 ],
166 )
167 for edges, lengths, ices_ in zip(
168 edges_ensemble, lengths_ensemble, ices_ensemble
169 )
170 ]
171 return cls(experiments, seed_sequence, fsd_handler)
173 @classmethod
174 def _generate_ices_ensemble(
175 cls,
176 ices: Ice | Sequence[Ice],
177 n_floes: int,
178 n_members: int,
179 seed_sequence: np.random.SeedSequence,
180 weights: np.ndarray | None,
181 ) -> Iterator[Iterator[Ice]]:
182 if not isinstance(ices, Sequence):
183 ice = ices
184 return itertools.repeat(itertools.repeat(ice, n_floes), n_members)
185 indices = np.arange(len(ices))
186 ice_seed_sequence = seed_sequence.spawn(1)[0]
187 ices_ensemble = (
188 (
189 ices[i]
190 for i in np.random.default_rng(_child_sequence).choice(
191 indices, size=n_floes, replace=True, p=weights
192 )
193 )
194 for _child_sequence in ice_seed_sequence.spawn(n_members)
195 )
196 return ices_ensemble
198 @staticmethod
199 def _add_noise_to_lengths(
200 lengths: np.ndarray,
201 n_members: int,
202 seed_sequence: np.random.SeedSequence,
203 noise_distribution: _RandomVariableProtocol,
204 ) -> np.ndarray:
205 # The last floe is infinite and not included in n_floes.
206 n_floes = lengths.size - 1
207 noise_seed_sequence = seed_sequence.spawn(1)[0]
208 noise = _generate_noise_on_lengths(
209 n_floes, n_members, noise_seed_sequence, noise_distribution
210 )
211 # Promote `lengths` to a 2D array, even if `n_members == 1`.
212 lengths_ensemble = np.repeat(lengths[None], n_members, axis=0)
213 lengths_ensemble[:, :-1] += noise
214 return lengths_ensemble
216 def _get_wavelengths(self) -> np.ndarray:
217 return (
218 2
219 * np.pi
220 / np.real(
221 np.array(
222 [
223 self.experiments[0].domain.floes[0].ice.wave_numbers[2],
224 self.experiments[0].domain.ocean.wave_numbers[0],
225 ]
226 )
227 )
228 )
230 def smallest_support(self) -> tuple[float, float]:
231 left_bounds, right_bounds = zip(
232 *(
233 (_e.domain.floes[0].left_edge, _e.domain.floes[-2].right_edge)
234 for _e in self.experiments
235 )
236 )
237 return max(left_bounds), min(right_bounds)
239 def envelope(
240 self,
241 *,
242 n_samples: int | None = None,
243 resolution: float | None = None,
244 window_length: float | None = None,
245 hop_length: float | None = None,
246 padding_factor: int = 8,
247 ):
248 manager = _EnvelopeManager.from_ensemble(self)
249 x, local_ft = manager.get_local_ft(
250 n_samples=n_samples,
251 resolution=resolution,
252 window_length=window_length,
253 hop_length=hop_length,
254 padding_factor=padding_factor,
255 )
256 envelopes = [
257 local_ft.envelope_from_fft(_e.deflections_from_positions(x))
258 for _e in self.experiments
259 ]
260 return envelopes
263@attrs.define
264class Experiment:
265 domain: Domain
267 @classmethod
268 def from_objects(
269 cls,
270 wave: Wave,
271 ocean: Ocean,
272 n_coupling_modes: tuple[int, int, int],
273 floes: list[Floe],
274 ):
275 domain = Domain(wave, ocean, *n_coupling_modes)
276 domain.add_floes(floes)
277 return cls(domain)
279 @functools.cached_property
280 def true_width(self) -> float:
281 return self.domain.floes[-1].left_edge - self.domain.floes[0].left_edge
283 def deflections_from_positions(self, positions: np.ndarray) -> np.ndarray:
284 indices = self.indices_from_positions(positions)
285 unique_indices, first_matches, counts = np.unique(
286 indices, return_index=True, return_counts=True
287 )
288 deflections = np.hstack(
289 [
290 floe_deflection(
291 self.domain.floes[_idx],
292 self.domain.wave,
293 positions[_first : _first + _count]
294 - self.domain.floes[_idx].left_edge,
295 )
296 for _idx, _first, _count in zip(unique_indices, first_matches, counts)
297 ]
298 )
299 return deflections
301 def indices_from_positions(self, positions: np.ndarray) -> np.ndarray:
302 positions = np.atleast_1d(positions)
303 edges = np.array([_cf.left_edge for _cf in self.domain.floes[:-1]])
304 return np.array([np.nonzero(_m)[0][-1] for _m in positions[:, None] >= edges])
307@attrs.define
308class _EnvelopeManager:
309 support: tuple[float, float]
310 wavelengths: np.ndarray | None = None
311 average_length: float | None = None
313 @classmethod
314 def from_ensemble(cls, ensemble: Ensemble):
315 support = ensemble.smallest_support()
316 wavelengths = ensemble._get_wavelengths()
317 average_length = ensemble.average_length
318 return cls(support, wavelengths, average_length)
320 def _get_x(self, n_samples: int | None, resolution: float | None) -> np.ndarray:
321 if n_samples is not None:
322 if resolution is not None:
323 warnings.warn(
324 "Both 'n_samples' and 'resolution' were provided, "
325 "'resolution' will be ignored.",
326 stacklevel=2,
327 )
328 return self._x_from_n_samples(n_samples)
329 else:
330 if resolution is not None:
331 return self._x_from_resolution(resolution)
332 else:
333 raise TypeError("One of 'n_samples' or 'resolution' is required.")
335 def _x_from_n_samples(self, n_samples: int) -> np.ndarray:
336 support = self.support
337 return np.linspace(*support, n_samples)
339 def _x_from_resolution(self, resolution: float) -> np.ndarray:
340 support = self.support
341 n = np.ceil((support[-1] - support[0]) / resolution).astype(int) + 1
342 return np.linspace(*support, n)
344 def _window_length_heuristic(self):
345 # Default window length equal to a multiple of the propagating wavelength.
346 _f = 4
347 # WARN: will need to be changed when several ice types are allowed.
348 window_length = _f * max(self.wavelengths)
349 return window_length
351 def _hop_length_heuristic(self):
352 # Default hop_length equal to a multiple of the average floe length.
353 _f = 8
354 hop_length = _f * self.average_length
355 return hop_length
357 def get_local_ft(
358 self,
359 window_length: float | None,
360 hop_length: float | None,
361 padding_factor: int,
362 *,
363 n_samples: int | None = None,
364 resolution: float | None = None,
365 ):
366 x = self._get_x(n_samples, resolution)
368 if window_length is None:
369 window_length = self._window_length_heuristic()
370 if hop_length is None:
371 hop_length = self._hop_length_heuristic()
372 return x, LocalFT.from_stft_params(x, window_length, hop_length, padding_factor)