OpenReservoirComputing 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- openreservoircomputing-0.1.0.dist-info/METADATA +168 -0
- openreservoircomputing-0.1.0.dist-info/RECORD +28 -0
- openreservoircomputing-0.1.0.dist-info/WHEEL +5 -0
- openreservoircomputing-0.1.0.dist-info/licenses/LICENSE +201 -0
- openreservoircomputing-0.1.0.dist-info/top_level.txt +1 -0
- orc/__init__.py +29 -0
- orc/classifier/__init__.py +6 -0
- orc/classifier/base.py +6 -0
- orc/classifier/models.py +6 -0
- orc/classifier/train.py +6 -0
- orc/control/__init__.py +15 -0
- orc/control/base.py +362 -0
- orc/control/models.py +139 -0
- orc/control/train.py +90 -0
- orc/data/__init__.py +27 -0
- orc/data/integrators.py +753 -0
- orc/drivers.py +851 -0
- orc/embeddings.py +459 -0
- orc/forecaster/__init__.py +20 -0
- orc/forecaster/base.py +487 -0
- orc/forecaster/models.py +416 -0
- orc/forecaster/train.py +301 -0
- orc/readouts.py +742 -0
- orc/tuning/__init__.py +6 -0
- orc/utils/__init__.py +6 -0
- orc/utils/numerics.py +115 -0
- orc/utils/regressions.py +73 -0
- orc/utils/visualization.py +193 -0
orc/forecaster/base.py
ADDED
|
@@ -0,0 +1,487 @@
|
|
|
1
|
+
"""Defines base classes for Reservoir Computer Forecasters."""
|
|
2
|
+
|
|
3
|
+
from abc import ABC
|
|
4
|
+
|
|
5
|
+
import diffrax
|
|
6
|
+
import equinox as eqx
|
|
7
|
+
import jax
|
|
8
|
+
import jax.numpy as jnp
|
|
9
|
+
from jaxtyping import Array, Float
|
|
10
|
+
|
|
11
|
+
from orc.drivers import DriverBase
|
|
12
|
+
from orc.embeddings import EmbedBase
|
|
13
|
+
from orc.readouts import ReadoutBase
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class RCForecasterBase(eqx.Module, ABC):
|
|
17
|
+
"""Base class for reservoir computer forecasters.
|
|
18
|
+
|
|
19
|
+
Defines the interface for the reservoir computer which includes the driver,
|
|
20
|
+
readout and embedding layers.
|
|
21
|
+
|
|
22
|
+
Attributes
|
|
23
|
+
----------
|
|
24
|
+
driver : DriverBase
|
|
25
|
+
Driver layer of the reservoir computer.
|
|
26
|
+
readout : ReadoutBase
|
|
27
|
+
Readout layer of the reservoir computer.
|
|
28
|
+
embedding : EmbedBase
|
|
29
|
+
Embedding layer of the reservoir computer.
|
|
30
|
+
in_dim : int
|
|
31
|
+
Dimension of the input data.
|
|
32
|
+
out_dim : int
|
|
33
|
+
Dimension of the output data.
|
|
34
|
+
res_dim : int
|
|
35
|
+
Dimension of the reservoir.
|
|
36
|
+
chunks : int
|
|
37
|
+
Number of parallel reservoirs.
|
|
38
|
+
dtype : type
|
|
39
|
+
Data type of the reservoir computer (jnp.float64 is highly recommended).
|
|
40
|
+
seed : int
|
|
41
|
+
Random seed for generating the PRNG key for the reservoir computer.
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
Methods
|
|
45
|
+
-------
|
|
46
|
+
force(in_seq, res_state)
|
|
47
|
+
Teacher forces the reservoir with the input sequence.
|
|
48
|
+
set_readout(readout)
|
|
49
|
+
Replaces the readout layer of the reservoir computer.
|
|
50
|
+
set_embedding(embedding)
|
|
51
|
+
Replaces the embedding layer of the reservoir computer.
|
|
52
|
+
forecast(fcast_len, res_state)
|
|
53
|
+
Forecast from an initial reservoir state.
|
|
54
|
+
forecast_from_IC(fcast_len, spinup_data)
|
|
55
|
+
Forecast from a sequence of spinup data.
|
|
56
|
+
"""
|
|
57
|
+
|
|
58
|
+
driver: DriverBase
|
|
59
|
+
readout: ReadoutBase
|
|
60
|
+
embedding: EmbedBase
|
|
61
|
+
in_dim: int
|
|
62
|
+
out_dim: int
|
|
63
|
+
res_dim: int
|
|
64
|
+
chunks: int = 0
|
|
65
|
+
dtype: Float = jnp.float64
|
|
66
|
+
seed: int = 0
|
|
67
|
+
|
|
68
|
+
def __init__(
|
|
69
|
+
self,
|
|
70
|
+
driver: DriverBase,
|
|
71
|
+
readout: ReadoutBase,
|
|
72
|
+
embedding: EmbedBase,
|
|
73
|
+
chunks: int = 0,
|
|
74
|
+
dtype: Float = jnp.float64,
|
|
75
|
+
seed: int = 0,
|
|
76
|
+
) -> None:
|
|
77
|
+
"""Initialize RCForecaster Base.
|
|
78
|
+
|
|
79
|
+
Parameters
|
|
80
|
+
----------
|
|
81
|
+
driver : DriverBase
|
|
82
|
+
Driver layer of the reservoir computer.
|
|
83
|
+
readout : ReadoutBase
|
|
84
|
+
Readout layer of the reservoir computer.
|
|
85
|
+
embedding : EmbedBase
|
|
86
|
+
Embedding layer of the reservoir computer.
|
|
87
|
+
chunks : int
|
|
88
|
+
Number of parallel reservoirs.
|
|
89
|
+
dtype : type
|
|
90
|
+
Data type of the reservoir computer (jnp.float64 is highly recommended).
|
|
91
|
+
seed : int
|
|
92
|
+
Random seed for generating the PRNG key for the reservoir computer.
|
|
93
|
+
"""
|
|
94
|
+
self.driver = driver
|
|
95
|
+
self.readout = readout
|
|
96
|
+
self.embedding = embedding
|
|
97
|
+
self.in_dim = self.embedding.in_dim
|
|
98
|
+
self.out_dim = self.readout.out_dim
|
|
99
|
+
self.res_dim = self.driver.res_dim
|
|
100
|
+
self.chunks = chunks
|
|
101
|
+
self.dtype = dtype
|
|
102
|
+
self.seed = seed
|
|
103
|
+
|
|
104
|
+
@eqx.filter_jit
|
|
105
|
+
def force(self, in_seq: Array, res_state: Array) -> Array:
|
|
106
|
+
"""Teacher forces the reservoir.
|
|
107
|
+
|
|
108
|
+
Parameters
|
|
109
|
+
----------
|
|
110
|
+
in_seq: Array
|
|
111
|
+
Input sequence to force the reservoir, (shape=(seq_len, data_dim)).
|
|
112
|
+
res_state : Array
|
|
113
|
+
Initial reservoir state, (shape=(chunks, res_dim,)).
|
|
114
|
+
|
|
115
|
+
Returns
|
|
116
|
+
-------
|
|
117
|
+
Array
|
|
118
|
+
Forced reservoir sequence, (shape=(seq_len, chunks, res_dim)).
|
|
119
|
+
"""
|
|
120
|
+
|
|
121
|
+
def scan_fn(state, in_vars):
|
|
122
|
+
proj_vars = self.embedding.embed(in_vars)
|
|
123
|
+
res_state = self.driver.advance(proj_vars, state)
|
|
124
|
+
return (res_state, res_state)
|
|
125
|
+
|
|
126
|
+
_, res_seq = jax.lax.scan(scan_fn, res_state, in_seq)
|
|
127
|
+
return res_seq
|
|
128
|
+
|
|
129
|
+
def __call__(self, in_seq: Array, res_state: Array) -> Array:
|
|
130
|
+
"""Teacher forces the reservoir, wrapper for `force` method.
|
|
131
|
+
|
|
132
|
+
Parameters
|
|
133
|
+
----------
|
|
134
|
+
in_seq: Array
|
|
135
|
+
Input sequence to force the reservoir, (shape=(seq_len, data_dim)).
|
|
136
|
+
res_state : Array
|
|
137
|
+
Initial reservoir state, (shape=(chunks, res_dim,)).
|
|
138
|
+
|
|
139
|
+
Returns
|
|
140
|
+
-------
|
|
141
|
+
Array
|
|
142
|
+
Forced reservoir sequence, (shape=(seq_len, chunks, res_dim)).
|
|
143
|
+
"""
|
|
144
|
+
return self.force(in_seq, res_state)
|
|
145
|
+
|
|
146
|
+
def set_readout(self, readout: ReadoutBase):
|
|
147
|
+
"""Replace readout layer.
|
|
148
|
+
|
|
149
|
+
Parameters
|
|
150
|
+
----------
|
|
151
|
+
readout : ReadoutBase
|
|
152
|
+
New readout layer.
|
|
153
|
+
|
|
154
|
+
Returns
|
|
155
|
+
-------
|
|
156
|
+
RCForecasterBase
|
|
157
|
+
Updated model with new readout layer.
|
|
158
|
+
"""
|
|
159
|
+
|
|
160
|
+
def where(m: RCForecasterBase):
|
|
161
|
+
return m.readout
|
|
162
|
+
|
|
163
|
+
new_model = eqx.tree_at(where, self, readout)
|
|
164
|
+
return new_model
|
|
165
|
+
|
|
166
|
+
def set_embedding(self, embedding: EmbedBase):
|
|
167
|
+
"""Replace embedding layer.
|
|
168
|
+
|
|
169
|
+
Parameters
|
|
170
|
+
----------
|
|
171
|
+
embedding : EmbedBase
|
|
172
|
+
New embedding layer.
|
|
173
|
+
|
|
174
|
+
Returns
|
|
175
|
+
-------
|
|
176
|
+
RCForecasterBase
|
|
177
|
+
Updated model with new embedding layer.
|
|
178
|
+
"""
|
|
179
|
+
|
|
180
|
+
def where(m: RCForecasterBase):
|
|
181
|
+
return m.embedding
|
|
182
|
+
|
|
183
|
+
new_model = eqx.tree_at(where, self, embedding)
|
|
184
|
+
return new_model
|
|
185
|
+
|
|
186
|
+
@eqx.filter_jit
|
|
187
|
+
def forecast(self, fcast_len: int, res_state: Array) -> Array:
|
|
188
|
+
"""Forecast from an initial reservoir state.
|
|
189
|
+
|
|
190
|
+
Parameters
|
|
191
|
+
----------
|
|
192
|
+
fcast_len : int
|
|
193
|
+
Steps to forecast.
|
|
194
|
+
res_state : Array
|
|
195
|
+
Initial reservoir state, (shape=(chunks, res_dim)).
|
|
196
|
+
|
|
197
|
+
Returns
|
|
198
|
+
-------
|
|
199
|
+
Array
|
|
200
|
+
Forecasted states, (shape=(fcast_len, data_dim))
|
|
201
|
+
"""
|
|
202
|
+
|
|
203
|
+
def scan_fn(state, _):
|
|
204
|
+
out_state = self.driver.advance(
|
|
205
|
+
self.embedding.embed(self.readout.readout(state)), state
|
|
206
|
+
)
|
|
207
|
+
return (out_state, self.readout.readout(out_state))
|
|
208
|
+
|
|
209
|
+
_, state_seq = jax.lax.scan(scan_fn, res_state, None, length=fcast_len - 1)
|
|
210
|
+
pre_append_state = self.readout.readout(res_state)
|
|
211
|
+
return jnp.vstack([pre_append_state, state_seq])
|
|
212
|
+
|
|
213
|
+
@eqx.filter_jit
|
|
214
|
+
def forecast_from_IC(self, fcast_len: int, spinup_data: Array) -> Array:
|
|
215
|
+
"""Forecast from a sequence of spinup data.
|
|
216
|
+
|
|
217
|
+
Parameters
|
|
218
|
+
----------
|
|
219
|
+
fcast_len : int
|
|
220
|
+
Steps to forecast.
|
|
221
|
+
spinup_data : Array
|
|
222
|
+
Initial condition sequence, (shape=(seq_len, data_dim)).
|
|
223
|
+
|
|
224
|
+
Returns
|
|
225
|
+
-------
|
|
226
|
+
Array
|
|
227
|
+
Forecasted states, (shape=(fcast_len, data_dim)).
|
|
228
|
+
"""
|
|
229
|
+
if self.chunks > 0:
|
|
230
|
+
res_seq = self.force(
|
|
231
|
+
spinup_data, jnp.zeros((self.chunks, self.res_dim), dtype=self.dtype)
|
|
232
|
+
)
|
|
233
|
+
elif self.chunks == 0:
|
|
234
|
+
res_seq = self.force(
|
|
235
|
+
spinup_data, jnp.zeros((self.res_dim), dtype=self.dtype)
|
|
236
|
+
)
|
|
237
|
+
else:
|
|
238
|
+
raise ValueError(f"chunks must be >= 0, but found chunks = {self.chunks}")
|
|
239
|
+
|
|
240
|
+
return self.forecast(fcast_len, res_seq[-1])
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
class CRCForecasterBase(RCForecasterBase, ABC):
|
|
244
|
+
"""Base class for continuous reservoir computer forecasters.
|
|
245
|
+
|
|
246
|
+
Override the force and forecast methods of RCForecasterBase
|
|
247
|
+
to timestep the RC forward using a continuous time ODE solver.
|
|
248
|
+
|
|
249
|
+
Attributes
|
|
250
|
+
----------
|
|
251
|
+
driver : DriverBase
|
|
252
|
+
Driver layer of the reservoir computer.
|
|
253
|
+
readout : ReadoutBase
|
|
254
|
+
Readout layer of the reservoir computer.
|
|
255
|
+
embedding : EmbedBase
|
|
256
|
+
Embedding layer of the reservoir computer.
|
|
257
|
+
in_dim : int
|
|
258
|
+
Dimension of the input data.
|
|
259
|
+
out_dim : int
|
|
260
|
+
Dimension of the output data.
|
|
261
|
+
res_dim : int
|
|
262
|
+
Dimension of the reservoir.
|
|
263
|
+
chunks : int
|
|
264
|
+
Number of parallel reservoirs.
|
|
265
|
+
dtype : type
|
|
266
|
+
Data type of the reservoir computer (jnp.float64 is highly recommended).
|
|
267
|
+
seed : int
|
|
268
|
+
Random seed for generating the PRNG key for the reservoir computer.
|
|
269
|
+
solver : diffrax.Solver
|
|
270
|
+
ODE solver to use for the reservoir computer.
|
|
271
|
+
stepsize_controller : diffrax.StepsizeController
|
|
272
|
+
Stepsize controller to use for the ODE solver.
|
|
273
|
+
|
|
274
|
+
Methods
|
|
275
|
+
-------
|
|
276
|
+
force(in_seq, res_state)
|
|
277
|
+
Teacher forces the reservoir with the input sequence.
|
|
278
|
+
set_readout(readout)
|
|
279
|
+
Replaces the readout layer of the reservoir computer.
|
|
280
|
+
set_embedding(embedding)
|
|
281
|
+
Replaces the embedding layer of the reservoir computer.
|
|
282
|
+
forecast(fcast_len, res_state)
|
|
283
|
+
Forecast from an initial reservoir state.
|
|
284
|
+
forecast_from_IC(fcast_len, spinup_data)
|
|
285
|
+
Forecast from a sequence of spinup data.
|
|
286
|
+
"""
|
|
287
|
+
|
|
288
|
+
solver: diffrax.AbstractSolver
|
|
289
|
+
stepsize_controller: diffrax.AbstractAdaptiveStepSizeController
|
|
290
|
+
|
|
291
|
+
def __init__(
|
|
292
|
+
self,
|
|
293
|
+
driver: DriverBase,
|
|
294
|
+
readout: ReadoutBase,
|
|
295
|
+
embedding: EmbedBase,
|
|
296
|
+
chunks: int = 0,
|
|
297
|
+
dtype: Float = jnp.float64,
|
|
298
|
+
seed: int = 0,
|
|
299
|
+
solver: diffrax.AbstractSolver = None,
|
|
300
|
+
stepsize_controller: diffrax.AbstractAdaptiveStepSizeController = None,
|
|
301
|
+
):
|
|
302
|
+
"""Initialize the continuous reservoir computer.
|
|
303
|
+
|
|
304
|
+
Parameters
|
|
305
|
+
----------
|
|
306
|
+
driver : DriverBase
|
|
307
|
+
Driver layer of the reservoir computer.
|
|
308
|
+
readout : ReadoutBase
|
|
309
|
+
Readout layer of the reservoir computer.
|
|
310
|
+
embedding : EmbedBase
|
|
311
|
+
Embedding layer of the reservoir computer.
|
|
312
|
+
chunks : int
|
|
313
|
+
Number of parallel reservoirs.
|
|
314
|
+
dtype : type
|
|
315
|
+
Data type of the reservoir computer (jnp.float64 is highly recommended).
|
|
316
|
+
seed : int
|
|
317
|
+
Random seed for generating the PRNG key for the reservoir computer.
|
|
318
|
+
solver : diffrax.AbstractSolver
|
|
319
|
+
ODE solver to use for the reservoir computer.
|
|
320
|
+
stepsize_controller : diffrax.AbstractAdaptiveStepSizeController
|
|
321
|
+
Stepsize controller to use for the ODE solver.
|
|
322
|
+
"""
|
|
323
|
+
super().__init__(driver, readout, embedding, chunks, dtype, seed)
|
|
324
|
+
if solver is None:
|
|
325
|
+
solver = diffrax.Tsit5()
|
|
326
|
+
if stepsize_controller is None:
|
|
327
|
+
stepsize_controller = diffrax.PIDController(
|
|
328
|
+
rtol=1e-3, atol=1e-6, icoeff=1.0
|
|
329
|
+
)
|
|
330
|
+
self.solver = solver
|
|
331
|
+
self.stepsize_controller = stepsize_controller
|
|
332
|
+
|
|
333
|
+
@eqx.filter_jit
|
|
334
|
+
def force(self, in_seq: Array, res_state: Array, ts: Array) -> Array:
|
|
335
|
+
"""
|
|
336
|
+
Teacher forces the reservoir.
|
|
337
|
+
|
|
338
|
+
Parameters
|
|
339
|
+
----------
|
|
340
|
+
in_seq: Array
|
|
341
|
+
Input sequence to force the reservoir, (shape=(seq_len, data_dim)).
|
|
342
|
+
res_state : Array
|
|
343
|
+
Initial reservoir state, (shape=(chunks, res_dim,)).
|
|
344
|
+
ts: Array
|
|
345
|
+
Time steps for the input sequence, (shape=(seq_len,)).
|
|
346
|
+
|
|
347
|
+
Returns
|
|
348
|
+
-------
|
|
349
|
+
Array
|
|
350
|
+
Forced reservoir sequence, (shape=(seq_len, chunks, res_dim)).
|
|
351
|
+
"""
|
|
352
|
+
# form interpolants
|
|
353
|
+
coeffs = diffrax.backward_hermite_coefficients(ts, in_seq)
|
|
354
|
+
in_seq_interp = diffrax.CubicInterpolation(ts, coeffs)
|
|
355
|
+
|
|
356
|
+
# RC forced ODE definition
|
|
357
|
+
@eqx.filter_jit
|
|
358
|
+
def res_ode(t, r, args):
|
|
359
|
+
interp = args
|
|
360
|
+
proj_vars = self.embedding.embed(interp.evaluate(t))
|
|
361
|
+
return self.driver.advance(proj_vars, r)
|
|
362
|
+
|
|
363
|
+
# integrate RC
|
|
364
|
+
dt0 = ts[1] - ts[0]
|
|
365
|
+
ts = ts + dt0 # roll time forward one step for targets
|
|
366
|
+
term = diffrax.ODETerm(res_ode)
|
|
367
|
+
args = in_seq_interp
|
|
368
|
+
save_at = diffrax.SaveAt(ts=ts)
|
|
369
|
+
sol = diffrax.diffeqsolve(
|
|
370
|
+
term,
|
|
371
|
+
t0=0.0,
|
|
372
|
+
t1=ts[-1],
|
|
373
|
+
dt0=dt0,
|
|
374
|
+
y0=res_state,
|
|
375
|
+
solver=self.solver,
|
|
376
|
+
stepsize_controller=self.stepsize_controller,
|
|
377
|
+
args=args,
|
|
378
|
+
saveat=save_at,
|
|
379
|
+
max_steps=None,
|
|
380
|
+
)
|
|
381
|
+
res_seq = sol.ys
|
|
382
|
+
return res_seq
|
|
383
|
+
|
|
384
|
+
def __call__(self, in_seq: Array, res_state: Array, ts: Array) -> Array:
|
|
385
|
+
"""
|
|
386
|
+
Teacher forces the reservoir, wrapper for `force` method.
|
|
387
|
+
|
|
388
|
+
Parameters
|
|
389
|
+
----------
|
|
390
|
+
in_seq: Array
|
|
391
|
+
Input sequence to force the reservoir, (shape=(seq_len, data_dim)).
|
|
392
|
+
res_state : Array
|
|
393
|
+
Initial reservoir state, (shape=(chunks, res_dim,)).
|
|
394
|
+
ts: Array
|
|
395
|
+
Time steps for the input sequence, (shape=(seq_len,)).
|
|
396
|
+
|
|
397
|
+
Returns
|
|
398
|
+
-------
|
|
399
|
+
Array
|
|
400
|
+
Forced reservoir sequence, (shape=(seq_len, chunks, res_dim)).
|
|
401
|
+
"""
|
|
402
|
+
return self.force(in_seq, res_state, ts)
|
|
403
|
+
|
|
404
|
+
@eqx.filter_jit
|
|
405
|
+
def forecast(self, ts: Array, res_state: Array) -> Array:
|
|
406
|
+
"""Forecast from an initial reservoir state.
|
|
407
|
+
|
|
408
|
+
Parameters
|
|
409
|
+
----------
|
|
410
|
+
ts : Array
|
|
411
|
+
Time steps for the forecast, (shape=(fcast_len,)).
|
|
412
|
+
res_state : Array
|
|
413
|
+
Initial reservoir state, (shape=(chunks, res_dim)).
|
|
414
|
+
|
|
415
|
+
Returns
|
|
416
|
+
-------
|
|
417
|
+
Array
|
|
418
|
+
Forecasted states, (shape=(fcast_len, data_dim))
|
|
419
|
+
"""
|
|
420
|
+
|
|
421
|
+
# RC autonomous ODE definition
|
|
422
|
+
@eqx.filter_jit
|
|
423
|
+
def res_ode(t, r, args):
|
|
424
|
+
out_state = self.driver.advance(
|
|
425
|
+
self.embedding.embed(self.readout.readout(r)), r
|
|
426
|
+
)
|
|
427
|
+
return out_state
|
|
428
|
+
|
|
429
|
+
# integrate RC
|
|
430
|
+
dt0 = ts[1] - ts[0]
|
|
431
|
+
term = diffrax.ODETerm(res_ode)
|
|
432
|
+
save_at = diffrax.SaveAt(ts=ts)
|
|
433
|
+
sol = diffrax.diffeqsolve(
|
|
434
|
+
term,
|
|
435
|
+
t0=0.0,
|
|
436
|
+
t1=ts[-1],
|
|
437
|
+
dt0=dt0,
|
|
438
|
+
y0=res_state,
|
|
439
|
+
solver=self.solver,
|
|
440
|
+
stepsize_controller=self.stepsize_controller,
|
|
441
|
+
saveat=save_at,
|
|
442
|
+
max_steps=None,
|
|
443
|
+
)
|
|
444
|
+
res_seq = sol.ys
|
|
445
|
+
return eqx.filter_vmap(self.readout.readout)(res_seq)
|
|
446
|
+
|
|
447
|
+
@eqx.filter_jit
|
|
448
|
+
def forecast_from_IC(
|
|
449
|
+
self, ts: Array, spinup_data: Array, spinup_ts: Array = None
|
|
450
|
+
) -> Array:
|
|
451
|
+
"""Forecast from a sequence of spinup data.
|
|
452
|
+
|
|
453
|
+
Parameters
|
|
454
|
+
----------
|
|
455
|
+
ts : Array
|
|
456
|
+
Time steps for the forecast, (shape=(fcast_len,)).
|
|
457
|
+
spinup_data : Array
|
|
458
|
+
Initial condition sequence, (shape=(seq_len, data_dim)).
|
|
459
|
+
spinup_ts : Array
|
|
460
|
+
Time steps for the spinup data, (shape=(seq_len,)).
|
|
461
|
+
If None, the spinup data is assumed to have the same dt
|
|
462
|
+
as the forecast data. If not None, the spinup data
|
|
463
|
+
Default is None.
|
|
464
|
+
|
|
465
|
+
Returns
|
|
466
|
+
-------
|
|
467
|
+
Array
|
|
468
|
+
Forecasted states, (shape=(fcast_len, data_dim)).
|
|
469
|
+
"""
|
|
470
|
+
if spinup_ts is None:
|
|
471
|
+
dt0 = ts[1] - ts[0]
|
|
472
|
+
spinup_ts = jnp.arange(0.0, spinup_data.shape[0], dtype=self.dtype) * dt0
|
|
473
|
+
|
|
474
|
+
if self.chunks > 0:
|
|
475
|
+
res_seq = self.force(
|
|
476
|
+
spinup_data,
|
|
477
|
+
jnp.zeros((self.chunks, self.res_dim), dtype=self.dtype),
|
|
478
|
+
spinup_ts,
|
|
479
|
+
)
|
|
480
|
+
elif self.chunks == 0:
|
|
481
|
+
res_seq = self.force(
|
|
482
|
+
spinup_data, jnp.zeros((self.res_dim), dtype=self.dtype), spinup_ts
|
|
483
|
+
)
|
|
484
|
+
else:
|
|
485
|
+
raise ValueError(f"chunks must be >= 0, but found chunks = {self.chunks}")
|
|
486
|
+
|
|
487
|
+
return self.forecast(ts, res_seq[-1])
|