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

1from collections.abc import Iterator, Sequence 

2import functools 

3import itertools 

4import logging 

5import typing 

6from typing import Self 

7import warnings 

8 

9import attrs 

10import numpy as np 

11from wisib import Domain, Floe, Ice, Ocean, Wave 

12 

13if typing.TYPE_CHECKING: 

14 from wisib import FloeCoupled 

15 

16from ._fsd import _FSDHandler 

17from ._noise import _generate_noise_on_lengths, _RandomVariableProtocol 

18from ._stft import LocalFT 

19 

20logger = logging.getLogger(__name__) 

21 

22 

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 ) 

44 

45 if return_modes: 

46 return _t 

47 return _t.sum(axis=1) 

48 

49 

50# NOTE: most of these functions could be made static methods of Ensemble. 

51 

52 

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 

61 

62 

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) 

67 

68 

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 

78 

79 

80@attrs.define 

81class Ensemble: 

82 experiments: list[Experiment] 

83 seed_sequence: np.random.SeedSequence 

84 fsd_handler: _FSDHandler 

85 

86 @property 

87 def average_length(self) -> float: 

88 return self.fsd_handler.average_length 

89 

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. 

108 

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. 

137 

138 Returns 

139 ------- 

140 Ensemble 

141 The initialised ensemble. 

142 

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) 

151 

152 # ices = itertools.repeat(ice) 

153 ices_ensemble = cls._generate_ices_ensemble( 

154 ices, n_floes, n_members, seed_sequence, ices_weights 

155 ) 

156 

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) 

172 

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 

197 

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 

215 

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 ) 

229 

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) 

238 

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 

261 

262 

263@attrs.define 

264class Experiment: 

265 domain: Domain 

266 

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) 

278 

279 @functools.cached_property 

280 def true_width(self) -> float: 

281 return self.domain.floes[-1].left_edge - self.domain.floes[0].left_edge 

282 

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 

300 

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]) 

305 

306 

307@attrs.define 

308class _EnvelopeManager: 

309 support: tuple[float, float] 

310 wavelengths: np.ndarray | None = None 

311 average_length: float | None = None 

312 

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) 

319 

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.") 

334 

335 def _x_from_n_samples(self, n_samples: int) -> np.ndarray: 

336 support = self.support 

337 return np.linspace(*support, n_samples) 

338 

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) 

343 

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 

350 

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 

356 

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) 

367 

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)