torchrolling 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.
- torchrolling/__init__.py +7 -0
- torchrolling/_common.py +117 -0
- torchrolling/_ewm.py +309 -0
- torchrolling/_rolling.py +522 -0
- torchrolling/_triton.py +930 -0
- torchrolling/py.typed +0 -0
- torchrolling-0.1.0.dist-info/METADATA +269 -0
- torchrolling-0.1.0.dist-info/RECORD +10 -0
- torchrolling-0.1.0.dist-info/WHEEL +4 -0
- torchrolling-0.1.0.dist-info/licenses/LICENSE +21 -0
torchrolling/__init__.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
"""Pandas-style rolling and exponentially weighted statistics for PyTorch tensors, fast on GPU."""
|
|
2
|
+
|
|
3
|
+
from torchrolling._ewm import Ewm, ewm
|
|
4
|
+
from torchrolling._rolling import Rolling, rolling
|
|
5
|
+
|
|
6
|
+
__all__ = ["Ewm", "Rolling", "ewm", "rolling"]
|
|
7
|
+
__version__ = "0.1.0"
|
torchrolling/_common.py
ADDED
|
@@ -0,0 +1,117 @@
|
|
|
1
|
+
"""Argument checks, dtype rules and the choice of backend, shared by rolling() and ewm()."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import functools
|
|
6
|
+
import warnings
|
|
7
|
+
from collections.abc import Callable
|
|
8
|
+
|
|
9
|
+
import torch
|
|
10
|
+
from torch import Tensor
|
|
11
|
+
|
|
12
|
+
try:
|
|
13
|
+
from torchrolling import _triton
|
|
14
|
+
except ImportError: # pragma: no cover - Triton is only installed on Linux
|
|
15
|
+
_triton = None # type: ignore[assignment]
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def check_int(name: str, value: object, low: int) -> int:
|
|
19
|
+
if isinstance(value, bool) or not isinstance(value, int) or value < low:
|
|
20
|
+
raise ValueError(f"{name} must be an integer >= {low}, got {value!r}")
|
|
21
|
+
return value
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def as_float(x: object, name: str = "x") -> Tensor:
|
|
25
|
+
"""``x`` as a floating tensor: integers become the default float dtype."""
|
|
26
|
+
if not isinstance(x, Tensor):
|
|
27
|
+
raise TypeError(f"{name} must be a torch.Tensor, got {type(x).__name__}")
|
|
28
|
+
if x.is_complex():
|
|
29
|
+
raise TypeError("complex tensors are not supported")
|
|
30
|
+
if x.ndim == 0:
|
|
31
|
+
raise ValueError(f"{name} must have at least one dimension")
|
|
32
|
+
return x if x.is_floating_point() else x.to(torch.get_default_dtype())
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def as_other(other: object, shape: torch.Size) -> Tensor:
|
|
36
|
+
"""The second series of a pairwise statistic, which must be shaped like ``x``."""
|
|
37
|
+
other = as_float(other, "other")
|
|
38
|
+
if other.shape != shape:
|
|
39
|
+
raise ValueError(
|
|
40
|
+
f"other must have the same shape as x {tuple(shape)}, got {tuple(other.shape)}"
|
|
41
|
+
)
|
|
42
|
+
return other
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def accumulator(dtype: torch.dtype, acc_dtype: torch.dtype | None) -> torch.dtype:
|
|
46
|
+
"""dtype for sums and moments: the input's, but at least float32."""
|
|
47
|
+
if acc_dtype is None:
|
|
48
|
+
return torch.float64 if dtype == torch.float64 else torch.float32
|
|
49
|
+
if not isinstance(acc_dtype, torch.dtype) or not acc_dtype.is_floating_point:
|
|
50
|
+
raise TypeError(f"acc_dtype must be a floating torch.dtype, got {acc_dtype!r}")
|
|
51
|
+
return acc_dtype
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
# Tests set this, so that a kernel that fails raises instead of falling back to torch.
|
|
55
|
+
STRICT = False
|
|
56
|
+
# Cleared when a kernel fails on this machine: from then on, everything runs on torch.
|
|
57
|
+
KERNELS_OK = True
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def use_triton(*tensors: Tensor, window: int = 0, limit: str = "") -> bool:
|
|
61
|
+
"""Whether the Triton kernels can compute this: CUDA (or the interpreter, in tests), no
|
|
62
|
+
gradients needed, and ``window`` within the kernel's limit (``_triton.<limit>``)."""
|
|
63
|
+
x = tensors[0]
|
|
64
|
+
return (
|
|
65
|
+
KERNELS_OK
|
|
66
|
+
and _triton is not None
|
|
67
|
+
and x.numel() > 0
|
|
68
|
+
and (not limit or window <= getattr(_triton, limit))
|
|
69
|
+
and ((x.is_cuda and _supported(x.device)) or _triton.INTERPRET)
|
|
70
|
+
and not (torch.is_grad_enabled() and any(t.requires_grad for t in tensors))
|
|
71
|
+
)
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
@functools.cache
|
|
75
|
+
def _supported(device: torch.device) -> bool: # pragma: no cover - needs a GPU
|
|
76
|
+
# Triton compiles for compute capability 7.0 (Volta) and newer.
|
|
77
|
+
return torch.version.hip is not None or torch.cuda.get_device_capability(device) >= (7, 0)
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def kernel_failed(error: Exception) -> None: # pragma: no cover - needs Triton
|
|
81
|
+
"""A kernel cannot run here (an unsupported GPU or Triton version): warn once, and use
|
|
82
|
+
the torch code, which gives the same results, from now on."""
|
|
83
|
+
if STRICT:
|
|
84
|
+
raise error
|
|
85
|
+
global KERNELS_OK
|
|
86
|
+
KERNELS_OK = False
|
|
87
|
+
warnings.warn(
|
|
88
|
+
f"torchrolling's Triton kernels failed on this machine ({type(error).__name__}: "
|
|
89
|
+
f"{str(error).splitlines()[0] if str(error) else ''}); using the slower torch code "
|
|
90
|
+
"from now on. Please report it at https://github.com/LenaBarretta/torchrolling/issues",
|
|
91
|
+
RuntimeWarning,
|
|
92
|
+
stacklevel=4,
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def kernel_dtype(x: Tensor) -> Tensor:
|
|
97
|
+
"""The kernels read float32 or float64."""
|
|
98
|
+
return x if x.dtype in (torch.float32, torch.float64) else x.float()
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def dtype_name(dtype: torch.dtype) -> str:
|
|
102
|
+
return "float64" if dtype == torch.float64 else "float32" # pragma: no cover - needs Triton
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def run_kernel( # pragma: no cover - needs Triton
|
|
106
|
+
launch: Callable[[], Tensor], dtype: torch.dtype, other: Tensor | None, dim: int
|
|
107
|
+
) -> Tensor | None:
|
|
108
|
+
"""Run a Triton kernel and shape its result like the torch code's (``dtype``, promoted
|
|
109
|
+
with ``other``'s, and the time dimension back at ``dim``). None if the kernel failed."""
|
|
110
|
+
try:
|
|
111
|
+
out = launch()
|
|
112
|
+
except Exception as error:
|
|
113
|
+
kernel_failed(error)
|
|
114
|
+
return None
|
|
115
|
+
if other is not None:
|
|
116
|
+
dtype = torch.promote_types(dtype, other.dtype)
|
|
117
|
+
return out.to(dtype).movedim(-1, dim)
|
torchrolling/_ewm.py
ADDED
|
@@ -0,0 +1,309 @@
|
|
|
1
|
+
"""Exponentially weighted statistics with pandas semantics.
|
|
2
|
+
|
|
3
|
+
pandas updates the weighted mean one observation at a time:
|
|
4
|
+
``mean_t = (1 - s_t) * mean_{t-1} + s_t * x_t``, where ``s_t`` is the share of the newest
|
|
5
|
+
weight in the total. The shares depend only on where values are missing, never on the
|
|
6
|
+
values, so they are computed up front, and the mean, the weighted covariance and the bias
|
|
7
|
+
correction all become affine recurrences ``y_t = a_t * y_{t-1} + b_t``. Those are composed
|
|
8
|
+
in parallel (``_affine_scan``), which is stable because every ``a_t`` lies in [0, 1].
|
|
9
|
+
|
|
10
|
+
Values are shifted by each series' first observation before the scan, so a constant
|
|
11
|
+
stretch gives exactly zero variance, as in pandas.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import math
|
|
17
|
+
|
|
18
|
+
import torch
|
|
19
|
+
import torch.nn.functional as F
|
|
20
|
+
from torch import Tensor
|
|
21
|
+
|
|
22
|
+
from torchrolling import _common
|
|
23
|
+
from torchrolling._common import (
|
|
24
|
+
accumulator,
|
|
25
|
+
as_float,
|
|
26
|
+
as_other,
|
|
27
|
+
check_int,
|
|
28
|
+
dtype_name,
|
|
29
|
+
kernel_dtype,
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
try:
|
|
33
|
+
from torchrolling import _triton
|
|
34
|
+
except ImportError: # pragma: no cover - Triton is only installed on Linux
|
|
35
|
+
_triton = None # type: ignore[assignment]
|
|
36
|
+
|
|
37
|
+
# Steps composed directly (in log2 rounds) before moving one level up.
|
|
38
|
+
_SCAN_BLOCK = 64
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _compose(a: Tensor, b: Tensor) -> tuple[Tensor, Tensor]:
|
|
42
|
+
"""Inclusive prefix composition of the maps y -> a * y + b along the last dim."""
|
|
43
|
+
k = 1
|
|
44
|
+
while k < b.shape[-1]:
|
|
45
|
+
b = a * F.pad(b[..., :-k], (k, 0)) + b
|
|
46
|
+
a = a * F.pad(a[..., :-k], (k, 0), value=1.0)
|
|
47
|
+
k *= 2
|
|
48
|
+
return a, b
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _affine_scan(a: Tensor, b: Tensor) -> Tensor:
|
|
52
|
+
"""``y_t = a_t * y_{t-1} + b_t`` along the last dim, starting from ``y_{-1} = 0``.
|
|
53
|
+
|
|
54
|
+
``b`` may have extra leading dims (several series sharing the same ``a``). Blocks of
|
|
55
|
+
``_SCAN_BLOCK`` steps are composed directly; their end values are chained by the same
|
|
56
|
+
scan one level up, so the work is O(n log _SCAN_BLOCK).
|
|
57
|
+
"""
|
|
58
|
+
length = b.shape[-1]
|
|
59
|
+
if length <= _SCAN_BLOCK:
|
|
60
|
+
return _compose(a, b)[1]
|
|
61
|
+
pad = -length % _SCAN_BLOCK
|
|
62
|
+
blocks = (length + pad) // _SCAN_BLOCK
|
|
63
|
+
a = F.pad(a, (0, pad), value=1.0).reshape(*a.shape[:-1], blocks, _SCAN_BLOCK)
|
|
64
|
+
b = F.pad(b, (0, pad)).reshape(*b.shape[:-1], blocks, _SCAN_BLOCK)
|
|
65
|
+
a, b = _compose(a, b)
|
|
66
|
+
ends = _affine_scan(a[..., -1], b[..., -1])
|
|
67
|
+
starts = F.pad(ends[..., :-1], (1, 0))
|
|
68
|
+
return (a * starts.unsqueeze(-1) + b).flatten(-2)[..., :length]
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def _center_of_mass(
|
|
72
|
+
com: float | None, span: float | None, halflife: float | None, alpha: float | None
|
|
73
|
+
) -> float:
|
|
74
|
+
given = {"com": com, "span": span, "halflife": halflife, "alpha": alpha}
|
|
75
|
+
given = {k: v for k, v in given.items() if v is not None}
|
|
76
|
+
if len(given) != 1:
|
|
77
|
+
raise ValueError("pass exactly one of com, span, halflife or alpha")
|
|
78
|
+
((name, value),) = given.items()
|
|
79
|
+
if isinstance(value, bool) or not isinstance(value, (int, float)) or math.isnan(value):
|
|
80
|
+
raise ValueError(f"{name} must be a number, got {value!r}")
|
|
81
|
+
if name == "com":
|
|
82
|
+
if value < 0:
|
|
83
|
+
raise ValueError(f"com must be >= 0, got {value!r}")
|
|
84
|
+
return float(value)
|
|
85
|
+
if name == "span":
|
|
86
|
+
if value < 1:
|
|
87
|
+
raise ValueError(f"span must be >= 1, got {value!r}")
|
|
88
|
+
return (value - 1) / 2
|
|
89
|
+
if name == "halflife":
|
|
90
|
+
if value <= 0:
|
|
91
|
+
raise ValueError(f"halflife must be > 0, got {value!r}")
|
|
92
|
+
return 1 / (1 - math.exp(math.log(0.5) / value)) - 1
|
|
93
|
+
if not 0 < value <= 1:
|
|
94
|
+
raise ValueError(f"alpha must be in (0, 1], got {value!r}")
|
|
95
|
+
return (1 - value) / value
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def ewm(
|
|
99
|
+
x: Tensor,
|
|
100
|
+
com: float | None = None,
|
|
101
|
+
span: float | None = None,
|
|
102
|
+
halflife: float | None = None,
|
|
103
|
+
alpha: float | None = None,
|
|
104
|
+
*,
|
|
105
|
+
min_periods: int = 0,
|
|
106
|
+
adjust: bool = True,
|
|
107
|
+
ignore_na: bool = False,
|
|
108
|
+
dim: int = -1,
|
|
109
|
+
acc_dtype: torch.dtype | None = None,
|
|
110
|
+
) -> Ewm:
|
|
111
|
+
"""Exponentially weighted window over ``dim`` of ``x``, like ``pandas.Series.ewm``.
|
|
112
|
+
|
|
113
|
+
Give exactly one of ``com``, ``span``, ``halflife`` or ``alpha``. Missing values (NaN
|
|
114
|
+
and +-inf) are skipped; ``adjust`` and ``ignore_na`` mean what they mean in pandas. A
|
|
115
|
+
result is NaN until ``min_periods`` (at least 1) values have been seen.
|
|
116
|
+
"""
|
|
117
|
+
return Ewm(
|
|
118
|
+
x,
|
|
119
|
+
com,
|
|
120
|
+
span,
|
|
121
|
+
halflife,
|
|
122
|
+
alpha,
|
|
123
|
+
min_periods=min_periods,
|
|
124
|
+
adjust=adjust,
|
|
125
|
+
ignore_na=ignore_na,
|
|
126
|
+
dim=dim,
|
|
127
|
+
acc_dtype=acc_dtype,
|
|
128
|
+
)
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
class Ewm:
|
|
132
|
+
"""Exponentially weighted statistics. Create it with :func:`ewm`."""
|
|
133
|
+
|
|
134
|
+
def __init__(
|
|
135
|
+
self,
|
|
136
|
+
x: Tensor,
|
|
137
|
+
com: float | None = None,
|
|
138
|
+
span: float | None = None,
|
|
139
|
+
halflife: float | None = None,
|
|
140
|
+
alpha: float | None = None,
|
|
141
|
+
*,
|
|
142
|
+
min_periods: int = 0,
|
|
143
|
+
adjust: bool = True,
|
|
144
|
+
ignore_na: bool = False,
|
|
145
|
+
dim: int = -1,
|
|
146
|
+
acc_dtype: torch.dtype | None = None,
|
|
147
|
+
) -> None:
|
|
148
|
+
x = as_float(x)
|
|
149
|
+
self._alpha = 1 / (1 + _center_of_mass(com, span, halflife, alpha))
|
|
150
|
+
self._min_periods = max(check_int("min_periods", min_periods, 0), 1)
|
|
151
|
+
self._adjust = bool(adjust)
|
|
152
|
+
self._ignore_na = bool(ignore_na)
|
|
153
|
+
self._dtype = x.dtype
|
|
154
|
+
self._dim = dim
|
|
155
|
+
self._shape = x.shape
|
|
156
|
+
self._acc = accumulator(x.dtype, acc_dtype)
|
|
157
|
+
self._orig = x.movedim(dim, -1)
|
|
158
|
+
|
|
159
|
+
def mean(self) -> Tensor:
|
|
160
|
+
"""Exponentially weighted mean."""
|
|
161
|
+
if (out := self._kernel("mean")) is not None: # pragma: no cover - needs Triton
|
|
162
|
+
return out
|
|
163
|
+
x, obs = self._prepare(self._orig)
|
|
164
|
+
if x.numel() == 0:
|
|
165
|
+
return self._finish(x, obs)
|
|
166
|
+
s, kept, nobs = self._shares(obs)
|
|
167
|
+
anchor = x.gather(-1, obs.to(torch.uint8).argmax(-1, keepdim=True))
|
|
168
|
+
mean = _affine_scan(kept, s * (x - anchor)) + anchor
|
|
169
|
+
return self._finish(mean, nobs >= self._min_periods)
|
|
170
|
+
|
|
171
|
+
def var(self, bias: bool = False) -> Tensor:
|
|
172
|
+
"""Exponentially weighted variance (bias-corrected unless ``bias=True``, like pandas)."""
|
|
173
|
+
if (out := self._kernel("var", bias=bias)) is not None: # pragma: no cover - needs Triton
|
|
174
|
+
return out
|
|
175
|
+
return self._cov(None, bias=bias)
|
|
176
|
+
|
|
177
|
+
def std(self, bias: bool = False) -> Tensor:
|
|
178
|
+
"""Exponentially weighted standard deviation."""
|
|
179
|
+
if (out := self._kernel("var", bias=bias, sqrt=True)) is not None: # pragma: no cover
|
|
180
|
+
return out
|
|
181
|
+
return self.var(bias).sqrt()
|
|
182
|
+
|
|
183
|
+
def cov(self, other: Tensor, bias: bool = False) -> Tensor:
|
|
184
|
+
"""Exponentially weighted covariance with ``other`` (same shape as ``x``)."""
|
|
185
|
+
y = as_other(other, self._shape).movedim(self._dim, -1)
|
|
186
|
+
if (out := self._kernel("cov", y, bias=bias)) is not None: # pragma: no cover
|
|
187
|
+
return out
|
|
188
|
+
return self._cov(y, bias=bias)
|
|
189
|
+
|
|
190
|
+
def corr(self, other: Tensor) -> Tensor:
|
|
191
|
+
"""Exponentially weighted correlation with ``other`` (same shape as ``x``).
|
|
192
|
+
|
|
193
|
+
NaN while either series has been constant so far.
|
|
194
|
+
"""
|
|
195
|
+
y = as_other(other, self._shape).movedim(self._dim, -1)
|
|
196
|
+
if (out := self._kernel("corr", y)) is not None: # pragma: no cover - needs Triton
|
|
197
|
+
return out
|
|
198
|
+
dtype = torch.promote_types(self._dtype, y.dtype)
|
|
199
|
+
x, obs = self._prepare(self._orig, y)
|
|
200
|
+
if x.numel() == 0:
|
|
201
|
+
return self._finish(x, obs, dtype)
|
|
202
|
+
y = torch.where(obs, y.to(self._acc), 0.0)
|
|
203
|
+
s, kept, nobs = self._shares(obs)
|
|
204
|
+
cxy, cxx, cyy = self._comoments(s, kept, [x, y], [(0, 1), (0, 0), (1, 1)])
|
|
205
|
+
denom = cxx * cyy
|
|
206
|
+
spread = denom > 0
|
|
207
|
+
corr = (cxy / torch.where(spread, denom, 1.0).sqrt()).clamp(-1, 1)
|
|
208
|
+
return self._finish(corr, (nobs >= self._min_periods) & spread, dtype)
|
|
209
|
+
|
|
210
|
+
def _kernel(
|
|
211
|
+
self, mode: str, other: Tensor | None = None, bias: bool = False, sqrt: bool = False
|
|
212
|
+
) -> Tensor | None:
|
|
213
|
+
"""``mode`` from the Triton kernel, or None where that does not apply."""
|
|
214
|
+
tensors = [self._orig] if other is None else [self._orig, other]
|
|
215
|
+
if self._acc not in (torch.float32, torch.float64) or not _common.use_triton(*tensors):
|
|
216
|
+
return None
|
|
217
|
+
return _common.run_kernel( # pragma: no cover - needs Triton
|
|
218
|
+
lambda: _triton.ewm(
|
|
219
|
+
kernel_dtype(self._orig),
|
|
220
|
+
None if other is None else kernel_dtype(other),
|
|
221
|
+
self._alpha,
|
|
222
|
+
self._min_periods,
|
|
223
|
+
mode,
|
|
224
|
+
bias,
|
|
225
|
+
sqrt,
|
|
226
|
+
self._adjust,
|
|
227
|
+
self._ignore_na,
|
|
228
|
+
dtype_name(self._acc),
|
|
229
|
+
),
|
|
230
|
+
self._dtype,
|
|
231
|
+
other,
|
|
232
|
+
self._dim,
|
|
233
|
+
)
|
|
234
|
+
|
|
235
|
+
def _cov(self, y: Tensor | None, bias: bool) -> Tensor:
|
|
236
|
+
"""Covariance of ``x`` with ``y``, or variance of ``x`` when ``y`` is None."""
|
|
237
|
+
dtype = self._dtype if y is None else torch.promote_types(self._dtype, y.dtype)
|
|
238
|
+
x, obs = self._prepare(self._orig, y)
|
|
239
|
+
if x.numel() == 0:
|
|
240
|
+
return self._finish(x, obs, dtype)
|
|
241
|
+
s, kept, nobs = self._shares(obs)
|
|
242
|
+
if y is None:
|
|
243
|
+
(cov,) = self._comoments(s, kept, [x], [(0, 0)])
|
|
244
|
+
else:
|
|
245
|
+
y = torch.where(obs, y.to(self._acc), 0.0)
|
|
246
|
+
(cov,) = self._comoments(s, kept, [x, y], [(0, 1)])
|
|
247
|
+
ok = nobs >= self._min_periods
|
|
248
|
+
if not bias:
|
|
249
|
+
# pandas multiplies by W^2 / (W^2 - sum of squared weights). With r the ratio
|
|
250
|
+
# (sum of squared weights) / W^2, r_t = kept^2 r_{t-1} + s^2, and since
|
|
251
|
+
# kept + s = 1, q = 1 - r follows q_t = kept^2 q_{t-1} + 2 kept s: no cancellation.
|
|
252
|
+
q = _affine_scan(kept * kept, 2 * kept * s)
|
|
253
|
+
ok = ok & (q > 0)
|
|
254
|
+
cov = cov / torch.where(q > 0, q, 1.0)
|
|
255
|
+
return self._finish(cov, ok, dtype)
|
|
256
|
+
|
|
257
|
+
def _comoments(
|
|
258
|
+
self, s: Tensor, kept: Tensor, series: list[Tensor], pairs: list[tuple[int, int]]
|
|
259
|
+
) -> list[Tensor]:
|
|
260
|
+
"""Weighted (biased) covariances of the given pairs of ``series``.
|
|
261
|
+
|
|
262
|
+
The series are zero where unobserved. pandas' update is
|
|
263
|
+
``c_t = kept * (c_{t-1} + dmean_x * dmean_y) + s * (x - mean_x) * (y - mean_y)``.
|
|
264
|
+
"""
|
|
265
|
+
v = torch.stack(series)
|
|
266
|
+
first = (s > 0).to(torch.uint8).argmax(-1, keepdim=True)
|
|
267
|
+
v = v - v.gather(-1, first.expand(*v.shape[:-1], 1))
|
|
268
|
+
means = _affine_scan(kept, s * v)
|
|
269
|
+
jump = F.pad(means[..., :-1], (1, 0)) - means
|
|
270
|
+
dev = v - means
|
|
271
|
+
b = torch.stack([kept * jump[i] * jump[j] + s * dev[i] * dev[j] for i, j in pairs])
|
|
272
|
+
return list(_affine_scan(kept, b))
|
|
273
|
+
|
|
274
|
+
def _prepare(self, x: Tensor, y: Tensor | None = None) -> tuple[Tensor, Tensor]:
|
|
275
|
+
"""``x`` in the accumulator dtype, zero where it (or ``y``) is missing; and the mask."""
|
|
276
|
+
obs = torch.isfinite(x) if y is None else torch.isfinite(x) & torch.isfinite(y)
|
|
277
|
+
return torch.where(obs, x.to(self._acc), 0.0), obs
|
|
278
|
+
|
|
279
|
+
def _shares(self, obs: Tensor) -> tuple[Tensor, Tensor, Tensor]:
|
|
280
|
+
"""Per step, the weights of the new value (s) and of the old mean (kept), and nobs.
|
|
281
|
+
|
|
282
|
+
Where nothing is observed s = 0 and kept = 1. Both are computed directly rather than
|
|
283
|
+
one as 1 minus the other, which would lose all precision when alpha is close to 1.
|
|
284
|
+
"""
|
|
285
|
+
f = 1 - self._alpha # as pandas computes it, so alpha == 1 gives exactly 0
|
|
286
|
+
o = obs.to(self._acc)
|
|
287
|
+
if self._adjust:
|
|
288
|
+
# The total weight decays every step (only at observations with ignore_na)
|
|
289
|
+
# and gains 1 per observation.
|
|
290
|
+
decay = o * f + (1 - o) if self._ignore_na else torch.full_like(o, f)
|
|
291
|
+
total = _affine_scan(decay, o)
|
|
292
|
+
before = decay * F.pad(total[..., :-1], (1, 0))
|
|
293
|
+
s = o / total.clamp(min=1)
|
|
294
|
+
kept = torch.where(obs, before / total.clamp(min=1), 1.0)
|
|
295
|
+
else:
|
|
296
|
+
# The previous total is normalised to 1, then decays over the steps since the
|
|
297
|
+
# previous observation; the new value gets weight alpha.
|
|
298
|
+
steps = torch.arange(obs.shape[-1], device=obs.device)
|
|
299
|
+
last = torch.where(obs, steps, -1).cummax(-1).values
|
|
300
|
+
previous = F.pad(last[..., :-1], (1, 0), value=-1)
|
|
301
|
+
gap = torch.ones_like(o) if self._ignore_na else (steps - previous).to(self._acc)
|
|
302
|
+
old = torch.where(previous < 0, 0.0, torch.full_like(o, f).pow(gap))
|
|
303
|
+
s = o * self._alpha / (old + self._alpha)
|
|
304
|
+
kept = torch.where(obs, old / (old + self._alpha), 1.0)
|
|
305
|
+
return s, kept, o.cumsum(-1)
|
|
306
|
+
|
|
307
|
+
def _finish(self, values: Tensor, ok: Tensor, dtype: torch.dtype | None = None) -> Tensor:
|
|
308
|
+
out = torch.where(ok, values, math.nan).to(dtype or self._dtype)
|
|
309
|
+
return out.movedim(-1, self._dim)
|