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