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