torchrolling 0.1.0__tar.gz
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-0.1.0/.gitignore +12 -0
- torchrolling-0.1.0/CHANGELOG.md +14 -0
- torchrolling-0.1.0/LICENSE +21 -0
- torchrolling-0.1.0/PKG-INFO +269 -0
- torchrolling-0.1.0/README.md +246 -0
- torchrolling-0.1.0/pyproject.toml +101 -0
- torchrolling-0.1.0/src/torchrolling/__init__.py +7 -0
- torchrolling-0.1.0/src/torchrolling/_common.py +117 -0
- torchrolling-0.1.0/src/torchrolling/_ewm.py +309 -0
- torchrolling-0.1.0/src/torchrolling/_rolling.py +522 -0
- torchrolling-0.1.0/src/torchrolling/_triton.py +930 -0
- torchrolling-0.1.0/src/torchrolling/py.typed +0 -0
- torchrolling-0.1.0/tests/conftest.py +68 -0
- torchrolling-0.1.0/tests/test_ewm.py +319 -0
- torchrolling-0.1.0/tests/test_quantile.py +161 -0
- torchrolling-0.1.0/tests/test_rolling.py +231 -0
- torchrolling-0.1.0/tests/test_vs_pandas.py +197 -0
- torchrolling-0.1.0/uv.lock +1620 -0
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
# Changelog
|
|
2
|
+
|
|
3
|
+
## 0.1.0
|
|
4
|
+
|
|
5
|
+
First release.
|
|
6
|
+
|
|
7
|
+
- `rolling(...)`: `count`, `sum`, `mean`, `var`, `std`, `skew`, `kurt`, `min`, `max`,
|
|
8
|
+
`median`, `quantile`, and pairwise `cov` and `corr`, with pandas semantics. O(1) work per
|
|
9
|
+
element for everything except quantiles.
|
|
10
|
+
- `ewm(...)`: `mean`, `var`, `std`, `cov`, `corr` with pandas' `com`/`span`/`halflife`/
|
|
11
|
+
`alpha`, `adjust`, `ignore_na` and `min_periods`.
|
|
12
|
+
- On CUDA, every statistic runs as one fused Triton kernel (Triton 3.2+, torch 2.6+).
|
|
13
|
+
- Autograd everywhere, `torch.compile(fullgraph=True)` support. Sums and moments accumulate
|
|
14
|
+
in the input's precision (at least float32); `acc_dtype` overrides it.
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 Lena Barretta
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|
|
@@ -0,0 +1,269 @@
|
|
|
1
|
+
Metadata-Version: 2.5
|
|
2
|
+
Name: torchrolling
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: Pandas-style rolling and exponentially weighted statistics for PyTorch tensors, fast on GPU.
|
|
5
|
+
Project-URL: Homepage, https://github.com/LenaBarretta/torchrolling
|
|
6
|
+
Project-URL: Issues, https://github.com/LenaBarretta/torchrolling/issues
|
|
7
|
+
Author: Lena Barretta
|
|
8
|
+
License-Expression: MIT
|
|
9
|
+
License-File: LICENSE
|
|
10
|
+
Keywords: ewm,ewma,gpu,moving-window,pandas,pytorch,rolling,time-series,torch,triton
|
|
11
|
+
Classifier: Development Status :: 4 - Beta
|
|
12
|
+
Classifier: Intended Audience :: Developers
|
|
13
|
+
Classifier: Intended Audience :: Science/Research
|
|
14
|
+
Classifier: Operating System :: OS Independent
|
|
15
|
+
Classifier: Programming Language :: Python :: 3
|
|
16
|
+
Classifier: Programming Language :: Python :: 3 :: Only
|
|
17
|
+
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
|
18
|
+
Classifier: Topic :: Scientific/Engineering :: Mathematics
|
|
19
|
+
Classifier: Typing :: Typed
|
|
20
|
+
Requires-Python: >=3.10
|
|
21
|
+
Requires-Dist: torch>=2.0
|
|
22
|
+
Description-Content-Type: text/markdown
|
|
23
|
+
|
|
24
|
+
<p align="center">
|
|
25
|
+
<img src="https://raw.githubusercontent.com/LenaBarretta/torchrolling/main/docs/logo.png" alt="torchrolling logo" width="160">
|
|
26
|
+
</p>
|
|
27
|
+
|
|
28
|
+
# torchrolling
|
|
29
|
+
|
|
30
|
+
[](https://pypi.org/project/torchrolling/)
|
|
31
|
+
[](https://pypi.org/project/torchrolling/)
|
|
32
|
+
[](https://github.com/LenaBarretta/torchrolling/actions/workflows/ci.yml)
|
|
33
|
+
[](LICENSE)
|
|
34
|
+
|
|
35
|
+
**Pandas-style rolling and exponentially weighted statistics for PyTorch tensors, fast on
|
|
36
|
+
GPU.** Only depends on torch.
|
|
37
|
+
|
|
38
|
+
```python
|
|
39
|
+
import torch
|
|
40
|
+
import torchrolling
|
|
41
|
+
|
|
42
|
+
x = torch.tensor([1.0, 2.0, float("nan"), 4.0, 5.0])
|
|
43
|
+
torchrolling.rolling(x, 3, min_periods=2).mean()
|
|
44
|
+
# tensor([nan, 1.5000, 1.5000, 3.0000, 4.5000])
|
|
45
|
+
|
|
46
|
+
prices = torch.randn(512, 10_000, device="cuda").cumsum(-1) # [series, time]
|
|
47
|
+
returns = prices.diff(dim=-1, prepend=prices[..., :1])
|
|
48
|
+
r = torchrolling.rolling(returns, 60)
|
|
49
|
+
features = torch.stack([r.mean(), r.std(), r.skew(), r.median(), r.corr(returns.roll(1, -1))])
|
|
50
|
+
trend = torchrolling.ewm(prices, span=20).mean()
|
|
51
|
+
```
|
|
52
|
+
|
|
53
|
+
## Why
|
|
54
|
+
|
|
55
|
+
`pandas.Series.rolling` works on one column on the CPU. When your series already live in a
|
|
56
|
+
`[batch, time]` tensor on the GPU (features computed on the fly, inside a model, after
|
|
57
|
+
augmentation), going through pandas means a round trip per batch. The usual torch workaround,
|
|
58
|
+
`x.unfold(-1, w, 1).mean(-1)`, does `w` times more work than needed, and for a median or
|
|
59
|
+
quantile it copies a tensor `w` times bigger than `x`. torch has no rolling median and no
|
|
60
|
+
exponential moving average at all.
|
|
61
|
+
|
|
62
|
+
torchrolling computes rolling sums, means, counts, variances, skew, kurtosis, covariances,
|
|
63
|
+
correlations, minima and maxima in O(1) work per element and O(n) memory for any window
|
|
64
|
+
size; medians and quantiles; and exponentially weighted means, variances and correlations
|
|
65
|
+
with a parallel scan. On CUDA every statistic is a single fused Triton kernel. It works on
|
|
66
|
+
any device, supports autograd and `torch.compile`, and gives the same numbers as pandas.
|
|
67
|
+
|
|
68
|
+
## Install
|
|
69
|
+
|
|
70
|
+
```bash
|
|
71
|
+
pip install torchrolling
|
|
72
|
+
```
|
|
73
|
+
|
|
74
|
+
Python 3.10+, torch 2.0+. The CUDA kernels need Triton 3.2+, which comes with torch 2.6+;
|
|
75
|
+
with an older torch, the same statistics run as plain torch operations.
|
|
76
|
+
|
|
77
|
+
## Usage
|
|
78
|
+
|
|
79
|
+
### Rolling windows
|
|
80
|
+
|
|
81
|
+
`torchrolling.rolling(x, window, *, min_periods=None, center=False, dim=-1, acc_dtype=None)`
|
|
82
|
+
returns an object with these methods. Each returns a tensor shaped like `x`.
|
|
83
|
+
|
|
84
|
+
| Method | pandas equivalent |
|
|
85
|
+
| --- | --- |
|
|
86
|
+
| `.count()` | `.rolling(...).count()` |
|
|
87
|
+
| `.sum()` | `.rolling(...).sum()` |
|
|
88
|
+
| `.mean()` | `.rolling(...).mean()` |
|
|
89
|
+
| `.var(ddof=1)` / `.std(ddof=1)` | `.rolling(...).var(ddof=1)` / `.std(ddof=1)` |
|
|
90
|
+
| `.skew()` / `.kurt()` | `.rolling(...).skew()` / `.kurt()` |
|
|
91
|
+
| `.min()` / `.max()` | `.rolling(...).min()` / `.max()` |
|
|
92
|
+
| `.median()` | `.rolling(...).median()` |
|
|
93
|
+
| `.quantile(q, interpolation="linear")` | `.rolling(...).quantile(q, interpolation="linear")` |
|
|
94
|
+
| `.cov(other, ddof=1)` / `.corr(other)` | `.rolling(...).cov(other)` / `.corr(other)` |
|
|
95
|
+
|
|
96
|
+
Call several methods on one object: counts and moments are computed once and shared.
|
|
97
|
+
|
|
98
|
+
### Exponentially weighted windows
|
|
99
|
+
|
|
100
|
+
`torchrolling.ewm(x, com=None, span=None, halflife=None, alpha=None, *, min_periods=0,
|
|
101
|
+
adjust=True, ignore_na=False, dim=-1, acc_dtype=None)` takes exactly one of `com`, `span`,
|
|
102
|
+
`halflife`, `alpha`, as in pandas.
|
|
103
|
+
|
|
104
|
+
| Method | pandas equivalent |
|
|
105
|
+
| --- | --- |
|
|
106
|
+
| `.mean()` | `.ewm(...).mean()` |
|
|
107
|
+
| `.var(bias=False)` / `.std(bias=False)` | `.ewm(...).var(bias=False)` / `.std(bias=False)` |
|
|
108
|
+
| `.cov(other, bias=False)` / `.corr(other)` | `.ewm(...).cov(other)` / `.corr(other)` |
|
|
109
|
+
|
|
110
|
+
### Semantics
|
|
111
|
+
|
|
112
|
+
The test suite checks everything against pandas with property-based tests (hypothesis):
|
|
113
|
+
|
|
114
|
+
- the output has the same length; windows that are not full yet give NaN unless
|
|
115
|
+
`min_periods` allows them;
|
|
116
|
+
- NaN and ±inf are treated as missing and skipped (pandas does the same). `count()` is the
|
|
117
|
+
exception, as in pandas: it counts ±inf;
|
|
118
|
+
- `center=True` centres the window the way pandas does;
|
|
119
|
+
- pairwise statistics (`cov`, `corr`) use only the positions where both series are valid;
|
|
120
|
+
- integer input becomes the default float dtype; float input keeps its dtype;
|
|
121
|
+
- `interpolation` is one of `linear`, `lower`, `higher`, `midpoint`, `nearest`, as in pandas.
|
|
122
|
+
|
|
123
|
+
Where results can differ from pandas:
|
|
124
|
+
|
|
125
|
+
- `ewm(adjust=False)` with missing values and `alpha=0.5` (also `com=1`, `span=3`,
|
|
126
|
+
`halflife=1`): pandas takes a different formula for that case
|
|
127
|
+
([pandas-dev/pandas#66523](https://github.com/pandas-dev/pandas/issues/66523));
|
|
128
|
+
torchrolling uses the documented weights, as pandas does for every other `alpha`;
|
|
129
|
+
- `ewm(...).var()` and `.cov()` with `bias=False`, when the weights have decayed a lot
|
|
130
|
+
(`alpha` close to 1, long runs of missing values): torchrolling computes the bias
|
|
131
|
+
correction without cancellation, so it can differ from pandas from the 8th digit on, and
|
|
132
|
+
gives NaN rather than 0 after a single observation;
|
|
133
|
+
- `corr` is NaN, not ±inf, where one series is constant in the window.
|
|
134
|
+
|
|
135
|
+
### Precision: `acc_dtype`
|
|
136
|
+
|
|
137
|
+
Sums and moments are accumulated in the input's precision, but at least float32: float64
|
|
138
|
+
input is accumulated in float64 and matches pandas to 1e-9 or better; float32, float16 and
|
|
139
|
+
bfloat16 input is accumulated in float32, which is accurate to about 1e-6 relative. Minima,
|
|
140
|
+
maxima and quantiles are exact in the input dtype. To accumulate float32 data in float64,
|
|
141
|
+
pass `acc_dtype=torch.float64`. On consumer and inference GPUs (T4, RTX), float64 is much
|
|
142
|
+
slower than float32.
|
|
143
|
+
|
|
144
|
+
### Gradients
|
|
145
|
+
|
|
146
|
+
Every statistic supports autograd. On CUDA, the fused Triton kernels compute the forward
|
|
147
|
+
pass when no gradient is needed (inference, feature pipelines, `torch.no_grad()`); when one
|
|
148
|
+
is, torchrolling runs the same algorithms as plain torch operations instead, which autograd
|
|
149
|
+
differentiates. Both give the same results.
|
|
150
|
+
|
|
151
|
+
## How it works
|
|
152
|
+
|
|
153
|
+
The series is padded and cut into blocks of exactly `window` elements. Each block is scanned
|
|
154
|
+
forwards (prefix) and backwards (suffix) with `cumsum`, `cummax` or `cummin`. Any window then
|
|
155
|
+
covers the tail of one block and the head of the next, so its value is
|
|
156
|
+
`combine(suffix[start], prefix[end])` (the van Herk / Gil-Werman algorithm). For sums this
|
|
157
|
+
also means only `2 * window` numbers are ever added together, so error does not grow with
|
|
158
|
+
series length.
|
|
159
|
+
|
|
160
|
+
Moments (variance, skew, kurtosis, covariance) come from power sums, which cancel badly when
|
|
161
|
+
taken raw. Each side of a window is therefore shifted by a value that lies inside that same
|
|
162
|
+
side: a head starts at its block's start, so it contains the block's first valid value, and
|
|
163
|
+
a tail contains its block's last one. The shifted sums are as well conditioned as the window
|
|
164
|
+
itself, a constant window gives exact zeros (so its variance is exactly 0, as in pandas), and
|
|
165
|
+
the two sides are merged with the pairwise formulas of Chan, Golub and LeVeque and of Pébay.
|
|
166
|
+
|
|
167
|
+
Quantiles cannot be split into halves; the torch code sorts every window, in chunks so that
|
|
168
|
+
memory stays bounded.
|
|
169
|
+
|
|
170
|
+
On CUDA, each of these is a single fused Triton kernel. For rolling statistics, a program
|
|
171
|
+
loads a run of whole blocks twice, once as is and once shifted by one block and one element,
|
|
172
|
+
so that the tail and the head of every window sit at the same position of two register
|
|
173
|
+
tiles; both scans and the merge happen in registers, in one pass over memory. For quantiles
|
|
174
|
+
with windows above 64, a program sorts the segment its outputs need once, with each value's
|
|
175
|
+
position packed into the sort key, then walks the sorted segment once, counting for every
|
|
176
|
+
window how many of its own values it has passed: O(window) work per output instead of
|
|
177
|
+
O(window log² window). Smaller windows sort each window directly.
|
|
178
|
+
|
|
179
|
+
Exponentially weighted statistics follow pandas' update rule
|
|
180
|
+
`mean_t = (1 - s_t) * mean_{t-1} + s_t * x_t`. The weights `s_t` depend only on where values
|
|
181
|
+
are missing, so they are computed up front, and the mean, the weighted covariance and the
|
|
182
|
+
bias correction become affine recurrences `y_t = a_t * y_{t-1} + b_t`. These are composed in
|
|
183
|
+
parallel in blocks of 64 steps (Hillis-Steele), which is stable because every `a_t` lies in
|
|
184
|
+
[0, 1]. On CUDA, one program per series scans it chunk by chunk and carries the state across
|
|
185
|
+
chunks.
|
|
186
|
+
|
|
187
|
+
## Benchmarks
|
|
188
|
+
|
|
189
|
+
Milliseconds, best of 3 runs, lower is better; the fastest in each row is bold. float32 data,
|
|
190
|
+
measured with [`bench/bench.py`](bench/bench.py). The GPU columns ran on a Tesla T4 (Kaggle,
|
|
191
|
+
torch 2.11, Triton 3.6, via [`notebooks/kaggle_gpu.ipynb`](notebooks/kaggle_gpu.ipynb));
|
|
192
|
+
pandas and polars ran twice, on that Kaggle machine's CPU and on an Apple M3 Max (16 cores;
|
|
193
|
+
pandas 3.0, polars 1.44). Each library gets its data where it wants it (tensors on the GPU,
|
|
194
|
+
a DataFrame for pandas and polars), and only the computation is timed. For EWM, `window` is
|
|
195
|
+
the span. All raw numbers, including float64 accumulation, are in
|
|
196
|
+
[`bench/results/`](bench/results/).
|
|
197
|
+
|
|
198
|
+
**1,000 series × 10,000 points**
|
|
199
|
+
|
|
200
|
+
| | window | torchrolling<br>T4 | torch `unfold`<br>T4 | cuDF<br>T4 | pandas<br>Kaggle CPU | polars<br>Kaggle CPU | pandas<br>M3 Max | polars<br>M3 Max |
|
|
201
|
+
| --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: |
|
|
202
|
+
| mean | 20 | 4.2 | **0.9** | 164 | 247 | 68 | 70 | 10 |
|
|
203
|
+
| | 200 | **4.0** | 6.4 | 265 | 262 | 68 | 71 | 9.6 |
|
|
204
|
+
| | 1000 | **3.1** | 38 | 350 | 238 | 62 | 71 | 9.3 |
|
|
205
|
+
| std | 20 | 5.1 | **1.2** | 291 | 330 | 98 | 138 | 13 |
|
|
206
|
+
| | 200 | **4.4** | 15 | 486 | 356 | 101 | 138 | 15 |
|
|
207
|
+
| | 1000 | **2.3** | 80 | 1,487 | 316 | 85 | 133 | 14 |
|
|
208
|
+
| max | 20 | 3.4 | **1.5** | 154 | 415 | 104 | 149 | 13 |
|
|
209
|
+
| | 200 | **2.0** | 3.2 | 170 | 424 | 110 | 148 | 13 |
|
|
210
|
+
| | 1000 | **1.9** | 30 | 207 | 386 | 97 | 148 | 12 |
|
|
211
|
+
| median | 20 | **18** | 142 | — | 4,721 | 480 | 2,239 | 40 |
|
|
212
|
+
| | 200 | 44 | out of memory | — | 5,998 | 425 | 2,909 | **39** |
|
|
213
|
+
| | 1000 | 81 | out of memory | — | 8,293 | 423 | 3,081 | **37** |
|
|
214
|
+
| corr | 20 | **8.3** | 64 | — | 1,594 | — | 526 | — |
|
|
215
|
+
| | 200 | **8.4** | out of memory | — | 1,462 | — | 527 | — |
|
|
216
|
+
| | 1000 | **7.3** | out of memory | — | 1,416 | — | 512 | — |
|
|
217
|
+
| ewm mean | 20 | **1.1** | — | 948 | 147 | 54 | 49 | 10.0 |
|
|
218
|
+
| | 200 | **1.0** | — | 902 | 137 | 54 | 49 | 9.9 |
|
|
219
|
+
| | 1000 | **1.0** | — | 876 | 137 | 53 | 49 | 11 |
|
|
220
|
+
|
|
221
|
+
**10,000 series × 10,000 points** (GPU libraries only)
|
|
222
|
+
|
|
223
|
+
| | window | torchrolling | torch `unfold` | cuDF |
|
|
224
|
+
| --- | ---: | ---: | ---: | ---: |
|
|
225
|
+
| mean | 20 | 42 | **5.7** | 2,230 |
|
|
226
|
+
| | 200 | **40** | 50 | 3,200 |
|
|
227
|
+
| | 1000 | **30** | 262 | 4,055 |
|
|
228
|
+
| std | 20 | 50 | **14** | 3,574 |
|
|
229
|
+
| | 200 | **40** | 136 | 5,566 |
|
|
230
|
+
| | 1000 | **27** | 804 | 15,515 |
|
|
231
|
+
| max | 20 | 31 | **9.7** | 2,042 |
|
|
232
|
+
| | 200 | **22** | 57 | 2,303 |
|
|
233
|
+
| | 1000 | **19** | 294 | 2,573 |
|
|
234
|
+
| median | 20 | **106** | out of memory | — |
|
|
235
|
+
| | 200 | **221** | out of memory | — |
|
|
236
|
+
| | 1000 | **830** | out of memory | — |
|
|
237
|
+
| corr | 20 | **58** | out of memory | — |
|
|
238
|
+
| | 200 | **50** | out of memory | — |
|
|
239
|
+
| | 1000 | **48** | out of memory | — |
|
|
240
|
+
| ewm mean | 20 | **5.3** | — | 10,136 |
|
|
241
|
+
| | 200 | **4.7** | — | 10,117 |
|
|
242
|
+
| | 1000 | **4.5** | — | 9,648 |
|
|
243
|
+
|
|
244
|
+
`torch unfold` is the usual plain-torch workaround, for example `x.unfold(-1, w, 1).mean(-1)`:
|
|
245
|
+
it has no missing-value handling and no `min_periods`, and for median and corr it copies a
|
|
246
|
+
tensor `window` times the size of the input. On small windows without missing values it is
|
|
247
|
+
the fastest way to get mean, std or max; everywhere else torchrolling is. "—" means the
|
|
248
|
+
library has no such statistic.
|
|
249
|
+
|
|
250
|
+
**Reading these numbers honestly.** Much of the gap to pandas and polars is a GPU against a
|
|
251
|
+
CPU, and it depends on the CPU: on the M3 Max, polars runs 5–7 times faster than on Kaggle's
|
|
252
|
+
CPU. One comparison goes the other way. For rolling median with windows of 200 and 1000,
|
|
253
|
+
polars on the M3 Max takes 39 and 37 ms, against torchrolling's 44 and 81 ms on the T4 (a
|
|
254
|
+
2018 inference GPU). Those polars times leave out what it takes to get data that already
|
|
255
|
+
lives on the GPU into polars: copying it to the CPU and the result back (40 MB each way for
|
|
256
|
+
this table) and building the DataFrame. So where the data lives decides: if it is already
|
|
257
|
+
on the GPU, sending it to the CPU rarely pays off; if it is on the CPU, weigh the copy to the
|
|
258
|
+
GPU, and staying on the CPU with pandas or polars may well be the better choice. cuDF is
|
|
259
|
+
measured at its weakest shape, thousands of short columns, which it processes one by one; on
|
|
260
|
+
a few long series it would be much closer. The first call of each statistic, which compiles
|
|
261
|
+
its Triton kernel (a few seconds, then cached on disk), is not timed either.
|
|
262
|
+
|
|
263
|
+
The kernels' launch configurations were tuned on a T4. Results are the same on every GPU;
|
|
264
|
+
only speed can differ. On a GPU or system where a kernel cannot run, torchrolling warns once
|
|
265
|
+
and computes the same statistics with plain torch operations.
|
|
266
|
+
|
|
267
|
+
## License
|
|
268
|
+
|
|
269
|
+
MIT
|
|
@@ -0,0 +1,246 @@
|
|
|
1
|
+
<p align="center">
|
|
2
|
+
<img src="https://raw.githubusercontent.com/LenaBarretta/torchrolling/main/docs/logo.png" alt="torchrolling logo" width="160">
|
|
3
|
+
</p>
|
|
4
|
+
|
|
5
|
+
# torchrolling
|
|
6
|
+
|
|
7
|
+
[](https://pypi.org/project/torchrolling/)
|
|
8
|
+
[](https://pypi.org/project/torchrolling/)
|
|
9
|
+
[](https://github.com/LenaBarretta/torchrolling/actions/workflows/ci.yml)
|
|
10
|
+
[](LICENSE)
|
|
11
|
+
|
|
12
|
+
**Pandas-style rolling and exponentially weighted statistics for PyTorch tensors, fast on
|
|
13
|
+
GPU.** Only depends on torch.
|
|
14
|
+
|
|
15
|
+
```python
|
|
16
|
+
import torch
|
|
17
|
+
import torchrolling
|
|
18
|
+
|
|
19
|
+
x = torch.tensor([1.0, 2.0, float("nan"), 4.0, 5.0])
|
|
20
|
+
torchrolling.rolling(x, 3, min_periods=2).mean()
|
|
21
|
+
# tensor([nan, 1.5000, 1.5000, 3.0000, 4.5000])
|
|
22
|
+
|
|
23
|
+
prices = torch.randn(512, 10_000, device="cuda").cumsum(-1) # [series, time]
|
|
24
|
+
returns = prices.diff(dim=-1, prepend=prices[..., :1])
|
|
25
|
+
r = torchrolling.rolling(returns, 60)
|
|
26
|
+
features = torch.stack([r.mean(), r.std(), r.skew(), r.median(), r.corr(returns.roll(1, -1))])
|
|
27
|
+
trend = torchrolling.ewm(prices, span=20).mean()
|
|
28
|
+
```
|
|
29
|
+
|
|
30
|
+
## Why
|
|
31
|
+
|
|
32
|
+
`pandas.Series.rolling` works on one column on the CPU. When your series already live in a
|
|
33
|
+
`[batch, time]` tensor on the GPU (features computed on the fly, inside a model, after
|
|
34
|
+
augmentation), going through pandas means a round trip per batch. The usual torch workaround,
|
|
35
|
+
`x.unfold(-1, w, 1).mean(-1)`, does `w` times more work than needed, and for a median or
|
|
36
|
+
quantile it copies a tensor `w` times bigger than `x`. torch has no rolling median and no
|
|
37
|
+
exponential moving average at all.
|
|
38
|
+
|
|
39
|
+
torchrolling computes rolling sums, means, counts, variances, skew, kurtosis, covariances,
|
|
40
|
+
correlations, minima and maxima in O(1) work per element and O(n) memory for any window
|
|
41
|
+
size; medians and quantiles; and exponentially weighted means, variances and correlations
|
|
42
|
+
with a parallel scan. On CUDA every statistic is a single fused Triton kernel. It works on
|
|
43
|
+
any device, supports autograd and `torch.compile`, and gives the same numbers as pandas.
|
|
44
|
+
|
|
45
|
+
## Install
|
|
46
|
+
|
|
47
|
+
```bash
|
|
48
|
+
pip install torchrolling
|
|
49
|
+
```
|
|
50
|
+
|
|
51
|
+
Python 3.10+, torch 2.0+. The CUDA kernels need Triton 3.2+, which comes with torch 2.6+;
|
|
52
|
+
with an older torch, the same statistics run as plain torch operations.
|
|
53
|
+
|
|
54
|
+
## Usage
|
|
55
|
+
|
|
56
|
+
### Rolling windows
|
|
57
|
+
|
|
58
|
+
`torchrolling.rolling(x, window, *, min_periods=None, center=False, dim=-1, acc_dtype=None)`
|
|
59
|
+
returns an object with these methods. Each returns a tensor shaped like `x`.
|
|
60
|
+
|
|
61
|
+
| Method | pandas equivalent |
|
|
62
|
+
| --- | --- |
|
|
63
|
+
| `.count()` | `.rolling(...).count()` |
|
|
64
|
+
| `.sum()` | `.rolling(...).sum()` |
|
|
65
|
+
| `.mean()` | `.rolling(...).mean()` |
|
|
66
|
+
| `.var(ddof=1)` / `.std(ddof=1)` | `.rolling(...).var(ddof=1)` / `.std(ddof=1)` |
|
|
67
|
+
| `.skew()` / `.kurt()` | `.rolling(...).skew()` / `.kurt()` |
|
|
68
|
+
| `.min()` / `.max()` | `.rolling(...).min()` / `.max()` |
|
|
69
|
+
| `.median()` | `.rolling(...).median()` |
|
|
70
|
+
| `.quantile(q, interpolation="linear")` | `.rolling(...).quantile(q, interpolation="linear")` |
|
|
71
|
+
| `.cov(other, ddof=1)` / `.corr(other)` | `.rolling(...).cov(other)` / `.corr(other)` |
|
|
72
|
+
|
|
73
|
+
Call several methods on one object: counts and moments are computed once and shared.
|
|
74
|
+
|
|
75
|
+
### Exponentially weighted windows
|
|
76
|
+
|
|
77
|
+
`torchrolling.ewm(x, com=None, span=None, halflife=None, alpha=None, *, min_periods=0,
|
|
78
|
+
adjust=True, ignore_na=False, dim=-1, acc_dtype=None)` takes exactly one of `com`, `span`,
|
|
79
|
+
`halflife`, `alpha`, as in pandas.
|
|
80
|
+
|
|
81
|
+
| Method | pandas equivalent |
|
|
82
|
+
| --- | --- |
|
|
83
|
+
| `.mean()` | `.ewm(...).mean()` |
|
|
84
|
+
| `.var(bias=False)` / `.std(bias=False)` | `.ewm(...).var(bias=False)` / `.std(bias=False)` |
|
|
85
|
+
| `.cov(other, bias=False)` / `.corr(other)` | `.ewm(...).cov(other)` / `.corr(other)` |
|
|
86
|
+
|
|
87
|
+
### Semantics
|
|
88
|
+
|
|
89
|
+
The test suite checks everything against pandas with property-based tests (hypothesis):
|
|
90
|
+
|
|
91
|
+
- the output has the same length; windows that are not full yet give NaN unless
|
|
92
|
+
`min_periods` allows them;
|
|
93
|
+
- NaN and ±inf are treated as missing and skipped (pandas does the same). `count()` is the
|
|
94
|
+
exception, as in pandas: it counts ±inf;
|
|
95
|
+
- `center=True` centres the window the way pandas does;
|
|
96
|
+
- pairwise statistics (`cov`, `corr`) use only the positions where both series are valid;
|
|
97
|
+
- integer input becomes the default float dtype; float input keeps its dtype;
|
|
98
|
+
- `interpolation` is one of `linear`, `lower`, `higher`, `midpoint`, `nearest`, as in pandas.
|
|
99
|
+
|
|
100
|
+
Where results can differ from pandas:
|
|
101
|
+
|
|
102
|
+
- `ewm(adjust=False)` with missing values and `alpha=0.5` (also `com=1`, `span=3`,
|
|
103
|
+
`halflife=1`): pandas takes a different formula for that case
|
|
104
|
+
([pandas-dev/pandas#66523](https://github.com/pandas-dev/pandas/issues/66523));
|
|
105
|
+
torchrolling uses the documented weights, as pandas does for every other `alpha`;
|
|
106
|
+
- `ewm(...).var()` and `.cov()` with `bias=False`, when the weights have decayed a lot
|
|
107
|
+
(`alpha` close to 1, long runs of missing values): torchrolling computes the bias
|
|
108
|
+
correction without cancellation, so it can differ from pandas from the 8th digit on, and
|
|
109
|
+
gives NaN rather than 0 after a single observation;
|
|
110
|
+
- `corr` is NaN, not ±inf, where one series is constant in the window.
|
|
111
|
+
|
|
112
|
+
### Precision: `acc_dtype`
|
|
113
|
+
|
|
114
|
+
Sums and moments are accumulated in the input's precision, but at least float32: float64
|
|
115
|
+
input is accumulated in float64 and matches pandas to 1e-9 or better; float32, float16 and
|
|
116
|
+
bfloat16 input is accumulated in float32, which is accurate to about 1e-6 relative. Minima,
|
|
117
|
+
maxima and quantiles are exact in the input dtype. To accumulate float32 data in float64,
|
|
118
|
+
pass `acc_dtype=torch.float64`. On consumer and inference GPUs (T4, RTX), float64 is much
|
|
119
|
+
slower than float32.
|
|
120
|
+
|
|
121
|
+
### Gradients
|
|
122
|
+
|
|
123
|
+
Every statistic supports autograd. On CUDA, the fused Triton kernels compute the forward
|
|
124
|
+
pass when no gradient is needed (inference, feature pipelines, `torch.no_grad()`); when one
|
|
125
|
+
is, torchrolling runs the same algorithms as plain torch operations instead, which autograd
|
|
126
|
+
differentiates. Both give the same results.
|
|
127
|
+
|
|
128
|
+
## How it works
|
|
129
|
+
|
|
130
|
+
The series is padded and cut into blocks of exactly `window` elements. Each block is scanned
|
|
131
|
+
forwards (prefix) and backwards (suffix) with `cumsum`, `cummax` or `cummin`. Any window then
|
|
132
|
+
covers the tail of one block and the head of the next, so its value is
|
|
133
|
+
`combine(suffix[start], prefix[end])` (the van Herk / Gil-Werman algorithm). For sums this
|
|
134
|
+
also means only `2 * window` numbers are ever added together, so error does not grow with
|
|
135
|
+
series length.
|
|
136
|
+
|
|
137
|
+
Moments (variance, skew, kurtosis, covariance) come from power sums, which cancel badly when
|
|
138
|
+
taken raw. Each side of a window is therefore shifted by a value that lies inside that same
|
|
139
|
+
side: a head starts at its block's start, so it contains the block's first valid value, and
|
|
140
|
+
a tail contains its block's last one. The shifted sums are as well conditioned as the window
|
|
141
|
+
itself, a constant window gives exact zeros (so its variance is exactly 0, as in pandas), and
|
|
142
|
+
the two sides are merged with the pairwise formulas of Chan, Golub and LeVeque and of Pébay.
|
|
143
|
+
|
|
144
|
+
Quantiles cannot be split into halves; the torch code sorts every window, in chunks so that
|
|
145
|
+
memory stays bounded.
|
|
146
|
+
|
|
147
|
+
On CUDA, each of these is a single fused Triton kernel. For rolling statistics, a program
|
|
148
|
+
loads a run of whole blocks twice, once as is and once shifted by one block and one element,
|
|
149
|
+
so that the tail and the head of every window sit at the same position of two register
|
|
150
|
+
tiles; both scans and the merge happen in registers, in one pass over memory. For quantiles
|
|
151
|
+
with windows above 64, a program sorts the segment its outputs need once, with each value's
|
|
152
|
+
position packed into the sort key, then walks the sorted segment once, counting for every
|
|
153
|
+
window how many of its own values it has passed: O(window) work per output instead of
|
|
154
|
+
O(window log² window). Smaller windows sort each window directly.
|
|
155
|
+
|
|
156
|
+
Exponentially weighted statistics follow pandas' update rule
|
|
157
|
+
`mean_t = (1 - s_t) * mean_{t-1} + s_t * x_t`. The weights `s_t` depend only on where values
|
|
158
|
+
are missing, so they are computed up front, and the mean, the weighted covariance and the
|
|
159
|
+
bias correction become affine recurrences `y_t = a_t * y_{t-1} + b_t`. These are composed in
|
|
160
|
+
parallel in blocks of 64 steps (Hillis-Steele), which is stable because every `a_t` lies in
|
|
161
|
+
[0, 1]. On CUDA, one program per series scans it chunk by chunk and carries the state across
|
|
162
|
+
chunks.
|
|
163
|
+
|
|
164
|
+
## Benchmarks
|
|
165
|
+
|
|
166
|
+
Milliseconds, best of 3 runs, lower is better; the fastest in each row is bold. float32 data,
|
|
167
|
+
measured with [`bench/bench.py`](bench/bench.py). The GPU columns ran on a Tesla T4 (Kaggle,
|
|
168
|
+
torch 2.11, Triton 3.6, via [`notebooks/kaggle_gpu.ipynb`](notebooks/kaggle_gpu.ipynb));
|
|
169
|
+
pandas and polars ran twice, on that Kaggle machine's CPU and on an Apple M3 Max (16 cores;
|
|
170
|
+
pandas 3.0, polars 1.44). Each library gets its data where it wants it (tensors on the GPU,
|
|
171
|
+
a DataFrame for pandas and polars), and only the computation is timed. For EWM, `window` is
|
|
172
|
+
the span. All raw numbers, including float64 accumulation, are in
|
|
173
|
+
[`bench/results/`](bench/results/).
|
|
174
|
+
|
|
175
|
+
**1,000 series × 10,000 points**
|
|
176
|
+
|
|
177
|
+
| | window | torchrolling<br>T4 | torch `unfold`<br>T4 | cuDF<br>T4 | pandas<br>Kaggle CPU | polars<br>Kaggle CPU | pandas<br>M3 Max | polars<br>M3 Max |
|
|
178
|
+
| --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: |
|
|
179
|
+
| mean | 20 | 4.2 | **0.9** | 164 | 247 | 68 | 70 | 10 |
|
|
180
|
+
| | 200 | **4.0** | 6.4 | 265 | 262 | 68 | 71 | 9.6 |
|
|
181
|
+
| | 1000 | **3.1** | 38 | 350 | 238 | 62 | 71 | 9.3 |
|
|
182
|
+
| std | 20 | 5.1 | **1.2** | 291 | 330 | 98 | 138 | 13 |
|
|
183
|
+
| | 200 | **4.4** | 15 | 486 | 356 | 101 | 138 | 15 |
|
|
184
|
+
| | 1000 | **2.3** | 80 | 1,487 | 316 | 85 | 133 | 14 |
|
|
185
|
+
| max | 20 | 3.4 | **1.5** | 154 | 415 | 104 | 149 | 13 |
|
|
186
|
+
| | 200 | **2.0** | 3.2 | 170 | 424 | 110 | 148 | 13 |
|
|
187
|
+
| | 1000 | **1.9** | 30 | 207 | 386 | 97 | 148 | 12 |
|
|
188
|
+
| median | 20 | **18** | 142 | — | 4,721 | 480 | 2,239 | 40 |
|
|
189
|
+
| | 200 | 44 | out of memory | — | 5,998 | 425 | 2,909 | **39** |
|
|
190
|
+
| | 1000 | 81 | out of memory | — | 8,293 | 423 | 3,081 | **37** |
|
|
191
|
+
| corr | 20 | **8.3** | 64 | — | 1,594 | — | 526 | — |
|
|
192
|
+
| | 200 | **8.4** | out of memory | — | 1,462 | — | 527 | — |
|
|
193
|
+
| | 1000 | **7.3** | out of memory | — | 1,416 | — | 512 | — |
|
|
194
|
+
| ewm mean | 20 | **1.1** | — | 948 | 147 | 54 | 49 | 10.0 |
|
|
195
|
+
| | 200 | **1.0** | — | 902 | 137 | 54 | 49 | 9.9 |
|
|
196
|
+
| | 1000 | **1.0** | — | 876 | 137 | 53 | 49 | 11 |
|
|
197
|
+
|
|
198
|
+
**10,000 series × 10,000 points** (GPU libraries only)
|
|
199
|
+
|
|
200
|
+
| | window | torchrolling | torch `unfold` | cuDF |
|
|
201
|
+
| --- | ---: | ---: | ---: | ---: |
|
|
202
|
+
| mean | 20 | 42 | **5.7** | 2,230 |
|
|
203
|
+
| | 200 | **40** | 50 | 3,200 |
|
|
204
|
+
| | 1000 | **30** | 262 | 4,055 |
|
|
205
|
+
| std | 20 | 50 | **14** | 3,574 |
|
|
206
|
+
| | 200 | **40** | 136 | 5,566 |
|
|
207
|
+
| | 1000 | **27** | 804 | 15,515 |
|
|
208
|
+
| max | 20 | 31 | **9.7** | 2,042 |
|
|
209
|
+
| | 200 | **22** | 57 | 2,303 |
|
|
210
|
+
| | 1000 | **19** | 294 | 2,573 |
|
|
211
|
+
| median | 20 | **106** | out of memory | — |
|
|
212
|
+
| | 200 | **221** | out of memory | — |
|
|
213
|
+
| | 1000 | **830** | out of memory | — |
|
|
214
|
+
| corr | 20 | **58** | out of memory | — |
|
|
215
|
+
| | 200 | **50** | out of memory | — |
|
|
216
|
+
| | 1000 | **48** | out of memory | — |
|
|
217
|
+
| ewm mean | 20 | **5.3** | — | 10,136 |
|
|
218
|
+
| | 200 | **4.7** | — | 10,117 |
|
|
219
|
+
| | 1000 | **4.5** | — | 9,648 |
|
|
220
|
+
|
|
221
|
+
`torch unfold` is the usual plain-torch workaround, for example `x.unfold(-1, w, 1).mean(-1)`:
|
|
222
|
+
it has no missing-value handling and no `min_periods`, and for median and corr it copies a
|
|
223
|
+
tensor `window` times the size of the input. On small windows without missing values it is
|
|
224
|
+
the fastest way to get mean, std or max; everywhere else torchrolling is. "—" means the
|
|
225
|
+
library has no such statistic.
|
|
226
|
+
|
|
227
|
+
**Reading these numbers honestly.** Much of the gap to pandas and polars is a GPU against a
|
|
228
|
+
CPU, and it depends on the CPU: on the M3 Max, polars runs 5–7 times faster than on Kaggle's
|
|
229
|
+
CPU. One comparison goes the other way. For rolling median with windows of 200 and 1000,
|
|
230
|
+
polars on the M3 Max takes 39 and 37 ms, against torchrolling's 44 and 81 ms on the T4 (a
|
|
231
|
+
2018 inference GPU). Those polars times leave out what it takes to get data that already
|
|
232
|
+
lives on the GPU into polars: copying it to the CPU and the result back (40 MB each way for
|
|
233
|
+
this table) and building the DataFrame. So where the data lives decides: if it is already
|
|
234
|
+
on the GPU, sending it to the CPU rarely pays off; if it is on the CPU, weigh the copy to the
|
|
235
|
+
GPU, and staying on the CPU with pandas or polars may well be the better choice. cuDF is
|
|
236
|
+
measured at its weakest shape, thousands of short columns, which it processes one by one; on
|
|
237
|
+
a few long series it would be much closer. The first call of each statistic, which compiles
|
|
238
|
+
its Triton kernel (a few seconds, then cached on disk), is not timed either.
|
|
239
|
+
|
|
240
|
+
The kernels' launch configurations were tuned on a T4. Results are the same on every GPU;
|
|
241
|
+
only speed can differ. On a GPU or system where a kernel cannot run, torchrolling warns once
|
|
242
|
+
and computes the same statistics with plain torch operations.
|
|
243
|
+
|
|
244
|
+
## License
|
|
245
|
+
|
|
246
|
+
MIT
|
|
@@ -0,0 +1,101 @@
|
|
|
1
|
+
[build-system]
|
|
2
|
+
requires = ["hatchling>=1.24"]
|
|
3
|
+
build-backend = "hatchling.build"
|
|
4
|
+
|
|
5
|
+
[project]
|
|
6
|
+
name = "torchrolling"
|
|
7
|
+
dynamic = ["version"]
|
|
8
|
+
description = "Pandas-style rolling and exponentially weighted statistics for PyTorch tensors, fast on GPU."
|
|
9
|
+
readme = "README.md"
|
|
10
|
+
license = "MIT"
|
|
11
|
+
license-files = ["LICENSE"]
|
|
12
|
+
authors = [{ name = "Lena Barretta" }]
|
|
13
|
+
requires-python = ">=3.10"
|
|
14
|
+
dependencies = ["torch>=2.0"]
|
|
15
|
+
keywords = [
|
|
16
|
+
"pytorch", "torch", "rolling", "moving-window", "ewm", "ewma", "time-series", "gpu", "pandas",
|
|
17
|
+
"triton",
|
|
18
|
+
]
|
|
19
|
+
classifiers = [
|
|
20
|
+
"Development Status :: 4 - Beta",
|
|
21
|
+
"Intended Audience :: Developers",
|
|
22
|
+
"Intended Audience :: Science/Research",
|
|
23
|
+
"Operating System :: OS Independent",
|
|
24
|
+
"Programming Language :: Python :: 3",
|
|
25
|
+
"Programming Language :: Python :: 3 :: Only",
|
|
26
|
+
"Topic :: Scientific/Engineering :: Mathematics",
|
|
27
|
+
"Topic :: Scientific/Engineering :: Artificial Intelligence",
|
|
28
|
+
"Typing :: Typed",
|
|
29
|
+
]
|
|
30
|
+
|
|
31
|
+
[project.urls]
|
|
32
|
+
Homepage = "https://github.com/LenaBarretta/torchrolling"
|
|
33
|
+
Issues = "https://github.com/LenaBarretta/torchrolling/issues"
|
|
34
|
+
|
|
35
|
+
[dependency-groups]
|
|
36
|
+
dev = [
|
|
37
|
+
"pytest>=8",
|
|
38
|
+
"pytest-cov>=5",
|
|
39
|
+
"ruff>=0.6",
|
|
40
|
+
"mypy>=1.11",
|
|
41
|
+
"hypothesis>=6",
|
|
42
|
+
"numpy>=1.24",
|
|
43
|
+
"pandas>=2.0",
|
|
44
|
+
"pandas-stubs>=2.0",
|
|
45
|
+
"triton>=3.0; sys_platform == 'linux'",
|
|
46
|
+
"mpmath>=1.3",
|
|
47
|
+
]
|
|
48
|
+
|
|
49
|
+
# Develop and test against the small CPU-only torch build.
|
|
50
|
+
[tool.uv.sources]
|
|
51
|
+
torch = [{ index = "pytorch-cpu" }]
|
|
52
|
+
|
|
53
|
+
[[tool.uv.index]]
|
|
54
|
+
name = "pytorch-cpu"
|
|
55
|
+
url = "https://download.pytorch.org/whl/cpu"
|
|
56
|
+
explicit = true
|
|
57
|
+
|
|
58
|
+
# Benchmarks and notebooks live in the repository but not in the package people install.
|
|
59
|
+
# (The wheel only ever contains src/torchrolling; this keeps them out of the sdist too.)
|
|
60
|
+
[tool.hatch.build.targets.sdist]
|
|
61
|
+
exclude = ["bench", "notebooks", "docs", ".github"]
|
|
62
|
+
|
|
63
|
+
[tool.hatch.version]
|
|
64
|
+
path = "src/torchrolling/__init__.py"
|
|
65
|
+
|
|
66
|
+
[tool.pytest.ini_options]
|
|
67
|
+
addopts = "-q --cov=torchrolling --cov-report=term-missing --cov-fail-under=100"
|
|
68
|
+
testpaths = ["tests"]
|
|
69
|
+
# pandas' own reference computations warn on all-missing windows with some numpy versions.
|
|
70
|
+
filterwarnings = ["ignore:All-NaN slice encountered:RuntimeWarning"]
|
|
71
|
+
|
|
72
|
+
[tool.coverage.run]
|
|
73
|
+
# Kernels only import where Triton is installed (Linux). The tests still run them there,
|
|
74
|
+
# on CPU through Triton's interpreter (see tests/conftest.py).
|
|
75
|
+
omit = ["*/torchrolling/_triton.py"]
|
|
76
|
+
|
|
77
|
+
[tool.ruff]
|
|
78
|
+
line-length = 100
|
|
79
|
+
extend-exclude = ["notebooks"]
|
|
80
|
+
target-version = "py310"
|
|
81
|
+
|
|
82
|
+
[tool.ruff.lint]
|
|
83
|
+
select = ["E", "F", "W", "I", "B", "UP", "SIM", "RUF"]
|
|
84
|
+
|
|
85
|
+
[tool.ruff.lint.per-file-ignores]
|
|
86
|
+
# Kernels branch on constexprs, which Triton's compiler resolves from plain comparisons and
|
|
87
|
+
# if blocks; `in` tests and conditional expressions are not reliably supported there.
|
|
88
|
+
"src/torchrolling/_triton.py" = ["SIM108", "SIM109"]
|
|
89
|
+
|
|
90
|
+
[tool.mypy]
|
|
91
|
+
strict = true
|
|
92
|
+
files = ["src", "tests"]
|
|
93
|
+
|
|
94
|
+
[[tool.mypy.overrides]]
|
|
95
|
+
# Triton and mpmath ship no type information.
|
|
96
|
+
module = ["triton", "triton.*", "mpmath"]
|
|
97
|
+
ignore_missing_imports = true
|
|
98
|
+
|
|
99
|
+
[[tool.mypy.overrides]]
|
|
100
|
+
module = ["torchrolling._triton"]
|
|
101
|
+
disallow_untyped_decorators = false
|
|
@@ -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"
|