flashmel 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.
- flashmel-0.1.0/.gitignore +6 -0
- flashmel-0.1.0/LICENSE +21 -0
- flashmel-0.1.0/PKG-INFO +201 -0
- flashmel-0.1.0/PROJECT.md +18 -0
- flashmel-0.1.0/README.md +167 -0
- flashmel-0.1.0/flashmel/__init__.py +12 -0
- flashmel-0.1.0/flashmel/build.py +80 -0
- flashmel-0.1.0/flashmel/filters.py +48 -0
- flashmel-0.1.0/flashmel/kernel.cu +266 -0
- flashmel-0.1.0/flashmel/runtime.py +65 -0
- flashmel-0.1.0/flashmel/transform.py +268 -0
- flashmel-0.1.0/pyproject.toml +59 -0
- flashmel-0.1.0/scripts/bench.py +127 -0
- flashmel-0.1.0/scripts/table.py +57 -0
- flashmel-0.1.0/scripts/test_correctness.py +190 -0
flashmel-0.1.0/LICENSE
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 AlumKal
|
|
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.
|
flashmel-0.1.0/PKG-INFO
ADDED
|
@@ -0,0 +1,201 @@
|
|
|
1
|
+
Metadata-Version: 2.5
|
|
2
|
+
Name: flashmel
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: Fused CUDA mel spectrogram (cuFFTDx) — a fast, torchaudio-compatible MelSpectrogram in a single kernel launch.
|
|
5
|
+
Project-URL: Homepage, https://github.com/AlumKal/flashmel
|
|
6
|
+
Project-URL: Repository, https://github.com/AlumKal/flashmel
|
|
7
|
+
Project-URL: Issues, https://github.com/AlumKal/flashmel/issues
|
|
8
|
+
Author-email: AlumKal <alumkal-pub@outlook.com>
|
|
9
|
+
License-Expression: MIT
|
|
10
|
+
License-File: LICENSE
|
|
11
|
+
Keywords: audio,cuda,cufftdx,gpu,mel,spectrogram,stft,torchaudio
|
|
12
|
+
Classifier: Development Status :: 4 - Beta
|
|
13
|
+
Classifier: Environment :: GPU :: NVIDIA CUDA
|
|
14
|
+
Classifier: Intended Audience :: Developers
|
|
15
|
+
Classifier: Intended Audience :: Science/Research
|
|
16
|
+
Classifier: Operating System :: POSIX :: Linux
|
|
17
|
+
Classifier: Programming Language :: Python :: 3
|
|
18
|
+
Classifier: Programming Language :: Python :: 3.10
|
|
19
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
20
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
21
|
+
Classifier: Programming Language :: Python :: 3.13
|
|
22
|
+
Classifier: Topic :: Multimedia :: Sound/Audio :: Analysis
|
|
23
|
+
Classifier: Topic :: Scientific/Engineering
|
|
24
|
+
Requires-Python: >=3.10
|
|
25
|
+
Requires-Dist: cuda-python<14,>=13.0
|
|
26
|
+
Requires-Dist: numpy>=1.24
|
|
27
|
+
Requires-Dist: nvidia-cuda-cccl<14,>=13.0
|
|
28
|
+
Requires-Dist: nvidia-cuda-nvrtc<14,>=13.0
|
|
29
|
+
Requires-Dist: nvidia-cuda-runtime<14,>=13.0
|
|
30
|
+
Requires-Dist: nvidia-mathdx>=25.6.0
|
|
31
|
+
Provides-Extra: dev
|
|
32
|
+
Requires-Dist: pytest>=9.0.3; extra == 'dev'
|
|
33
|
+
Description-Content-Type: text/markdown
|
|
34
|
+
|
|
35
|
+
# flashmel
|
|
36
|
+
|
|
37
|
+
Fused CUDA mel spectrogram: framing + padding + Hann window + R2C FFT (cuFFTDx)
|
|
38
|
+
+ `|X|^power` + sparse mel filterbank projection in a **single kernel launch**,
|
|
39
|
+
with a drop-in `torchaudio.transforms.MelSpectrogram`-compatible Python wrapper.
|
|
40
|
+
|
|
41
|
+
3–14× faster than torchaudio on an RTX 4070 Laptop, bit-matching torchaudio at
|
|
42
|
+
`atol=rtol=1e-4` (fp32).
|
|
43
|
+
|
|
44
|
+
## Installation
|
|
45
|
+
|
|
46
|
+
flashmel needs Python ≥ 3.10, an NVIDIA GPU, a recent NVIDIA driver, and a CUDA
|
|
47
|
+
build of **torch** (and **torchaudio**) — install those first, the way your
|
|
48
|
+
environment needs them, e.g.:
|
|
49
|
+
|
|
50
|
+
```sh
|
|
51
|
+
pip install torch torchaudio --index-url https://download.pytorch.org/whl/cu130
|
|
52
|
+
```
|
|
53
|
+
|
|
54
|
+
flashmel deliberately does **not** depend on torch/torchaudio: the CUDA wheel
|
|
55
|
+
stack is environment-specific (driver version, CUDA version, sometimes a private
|
|
56
|
+
mirror), and pinning it here would only fight your install. Then:
|
|
57
|
+
|
|
58
|
+
```sh
|
|
59
|
+
pip install flashmel
|
|
60
|
+
```
|
|
61
|
+
|
|
62
|
+
This pulls flashmel's own deps — the cuFFTDx/CUDA headers and NVRTC
|
|
63
|
+
(`nvidia-mathdx`, `nvidia-cuda-cccl`, `nvidia-cuda-runtime`, `nvidia-cuda-nvrtc`),
|
|
64
|
+
the driver bindings (`cuda-python`), and `numpy`. No system CUDA toolkit or
|
|
65
|
+
`nvcc` is required. The kernel is compiled on first use per configuration (~2 s)
|
|
66
|
+
and cached on disk (see below).
|
|
67
|
+
|
|
68
|
+
## Usage
|
|
69
|
+
|
|
70
|
+
```python
|
|
71
|
+
import torch
|
|
72
|
+
from flashmel import MelSpectrogram
|
|
73
|
+
|
|
74
|
+
mel = MelSpectrogram(n_fft=400, hop_length=160, n_mels=128) # torchaudio args
|
|
75
|
+
x = torch.randn(8, 480000, device="cuda") # f32 or i16 CUDA tensor
|
|
76
|
+
y = mel(x) # [8, 128, 3001]
|
|
77
|
+
```
|
|
78
|
+
|
|
79
|
+
Supports every `torchaudio.transforms.MelSpectrogram` argument
|
|
80
|
+
(`sample_rate, n_fft, win_length, hop_length, f_min, f_max, pad, n_mels, power,
|
|
81
|
+
normalized, center, pad_mode, onesided, norm, mel_scale`) except
|
|
82
|
+
`window_fn`/`wkwargs` (always periodic Hann). Additionally:
|
|
83
|
+
|
|
84
|
+
- **input dtype**: `float32`, or `int16` raw PCM (scaled by 1/32768) — dispatched
|
|
85
|
+
from the input tensor's dtype;
|
|
86
|
+
- **output dtype**: `dtype=torch.float32` (default) or `torch.float16`;
|
|
87
|
+
- input shape `(..., time)`, output `(..., n_mels, n_frames)`, like torchaudio.
|
|
88
|
+
|
|
89
|
+
Constraints: `n_fft` must be 400 or a power of 2 in [64, 16384]; `n_mels ≤ 256`;
|
|
90
|
+
`onesided=True` only (torchaudio's own `MelScale` can't consume two-sided
|
|
91
|
+
spectrograms either); input must be on CUDA.
|
|
92
|
+
|
|
93
|
+
## How it works
|
|
94
|
+
|
|
95
|
+
- **NVRTC-compiled** `flashmel/kernel.cu`, one cubin per parameter configuration,
|
|
96
|
+
disk-cached by content hash in the user cache dir (`~/.cache/flashmel`,
|
|
97
|
+
overridable via `FLASHMEL_CACHE_DIR`). No build step; first call per config
|
|
98
|
+
compiles in ~2 s.
|
|
99
|
+
- Frames are loaded straight into cuFFTDx registers with the signal padding
|
|
100
|
+
(`pad`, `center`, all four `pad_mode`s) applied as index arithmetic — the
|
|
101
|
+
padded/framed signal is never materialized.
|
|
102
|
+
- The STFT `normalized` modes are folded into the window vector on the host.
|
|
103
|
+
- Power-of-2 sizes use cuFFTDx `real_mode::folded` R2C (an N/2-point complex
|
|
104
|
+
FFT), halving shared memory and butterflies; 400 uses `real_mode::normal`.
|
|
105
|
+
- The FFT's shared-memory workspace is reused as the per-frame power tile
|
|
106
|
+
(odd row stride ⇒ bank-conflict-free), from which the sparse mel projection
|
|
107
|
+
reads. Mel filters are triangular ⇒ contiguous support, packed as
|
|
108
|
+
`start[m] / width[m] / weights[m, W_MAX]`; the inner loop runs to each
|
|
109
|
+
filter's actual width, which cut L1/SMEM traffic ~40% for wide-filter
|
|
110
|
+
configs (1024: 0.28→0.23 ms, 4096: 0.57→0.41 ms, 16384: 1.35→0.94 ms).
|
|
111
|
+
- When the block has more threads than mel outputs (large n_fft ⇒ FPB=1),
|
|
112
|
+
SPLIT lanes cooperate per dot product with a warp-shuffle reduce instead of
|
|
113
|
+
idling (16384: another 12%).
|
|
114
|
+
- Folded-mode interior frames (no padding in play) load signal and window as
|
|
115
|
+
2-wide vectors, halving input-phase load instructions (64: −30%, 1024/4096:
|
|
116
|
+
−13%; the LSU pipe, not DRAM, is the binding resource at these sizes).
|
|
117
|
+
- Normal-mode (400) interior frames take a direct-load path that skips the
|
|
118
|
+
per-sample reflect/boundary arithmetic entirely (whisper case: 1.08→0.54 ms,
|
|
119
|
+
2.0×, now at copy peak).
|
|
120
|
+
- Launches via the CUDA driver API on torch's current stream; fp16 output is
|
|
121
|
+
cast at the final store (all accumulation in fp32).
|
|
122
|
+
- The launch is wrapped in a `torch.library` custom op (`flashmel::mel_forward`
|
|
123
|
+
with a fake impl), so `torch.compile(..., fullgraph=True)` traces through it.
|
|
124
|
+
- Kernel modules are cached per (config, arch, device) — CUmodules are
|
|
125
|
+
context-bound — and filter/window buffers are cached per device, so multiple
|
|
126
|
+
GPUs work concurrently. Inference-only: requires-grad input under grad mode
|
|
127
|
+
raises (no backward is implemented).
|
|
128
|
+
|
|
129
|
+
## Performance (RTX 4070 Laptop, sm_89, measured copy peak ≈ 198 GB/s)
|
|
130
|
+
|
|
131
|
+
`uv run python scripts/bench.py` (CUDA-event timing, median of 50; "BW" counts
|
|
132
|
+
essential bytes = input + output only):
|
|
133
|
+
|
|
134
|
+
| case | config | batch | ours | torchaudio | speedup | essential BW |
|
|
135
|
+
|---|---|---|---|---|---|---|
|
|
136
|
+
| whisper | 400/160, 128 mel, 30 s | 32 | 0.54 ms | 6.34 ms | 11.7× | 205 GB/s (103%) |
|
|
137
|
+
| whisper i16→fp16 | 400/160, 128 mel | 32 | 0.53 ms | 7.28 ms | 13.8× | 105 GB/s |
|
|
138
|
+
| speech | 1024/256, 80 mel | 32 | 0.20 ms | 3.15 ms | 16.2× | 138 GB/s (70%) |
|
|
139
|
+
| small | 64/16, 40 mel | 64 | 0.051 ms | 0.29 ms | 5.6× | 280 GB/s (141%) |
|
|
140
|
+
| music | 4096/1024, 256 mel, 44.1k | 16 | 0.35 ms | 4.54 ms | 12.9× | 100 GB/s (51%) |
|
|
141
|
+
| huge | 16384/4096, 256 mel, 44.1k | 16 | 0.92 ms | 4.61 ms | 5.0× | 33 GB/s |
|
|
142
|
+
|
|
143
|
+
Full-spectrum sweep (hop = n_fft/4, 128 mel, B=32, T=160000): 64–512 run at
|
|
144
|
+
97–120% of copy peak, 1024 at 82%, then declining as the FFT itself dominates.
|
|
145
|
+
Numbers drift ±15% with laptop thermals.
|
|
146
|
+
|
|
147
|
+
### Where each case actually sits (Nsight Compute)
|
|
148
|
+
|
|
149
|
+
- **small–512**: DRAM-saturated. Effective BW exceeds copy peak because
|
|
150
|
+
overlapped frame reads hit L2 (essential-byte metric counts them once).
|
|
151
|
+
- **speech (1024) / music (4096)**: L1/SMEM-throughput bound at 76–81% SOL
|
|
152
|
+
with DRAM at 30–40% — the remaining L1 traffic is cuFFTDx's internal
|
|
153
|
+
register exchanges, not under our control (vectorizing our own smem reads is
|
|
154
|
+
blocked by the odd power-tile stride that keeps them bank-conflict-free).
|
|
155
|
+
- **whisper (400)**: DRAM-saturated since the interior direct-load path
|
|
156
|
+
(DRAM 77% / L1 80% / SM 74% SOL — 77% of theoretical DRAM ≈ the measured
|
|
157
|
+
copy peak; previously latency-bound at DRAM 48% / SM 47% with the reflect
|
|
158
|
+
arithmetic on every sample). Occupancy stays at 41%, capped by the FFT
|
|
159
|
+
workspace (12.8 KB per 4-frame block ⇒ 7 blocks/SM; normal-mode R2C is the
|
|
160
|
+
only cuFFTDx option for non-power-of-2), but that no longer binds. A paired
|
|
161
|
+
c2c trick (two frames per complex FFT via conjugate symmetry) halves FFT
|
|
162
|
+
work but measured ~1.5× *slower*: the full Z-spectrum round trip through
|
|
163
|
+
shared memory costs more than the saved butterflies. Implemented, A/B'd,
|
|
164
|
+
removed.
|
|
165
|
+
- **huge (8192/16384)**: occupancy-bound (33–64 KB workspace ⇒ 1–2 blocks/SM)
|
|
166
|
+
and barrier-serialized inside the block FFT. Inherent to a single-kernel
|
|
167
|
+
block FFT at these sizes; going faster would mean a multi-kernel FFT (and
|
|
168
|
+
materializing the spectrum in HBM) or fp16 twiddles (breaking the 1e-4
|
|
169
|
+
contract).
|
|
170
|
+
|
|
171
|
+
### Tuning
|
|
172
|
+
|
|
173
|
+
Per-`n_fft` frames-per-block defaults live in `flashmel/transform.py`
|
|
174
|
+
(`_FPB_DEFAULTS`), found with `scripts/bench.py --sweep`. Override per process
|
|
175
|
+
with `FLASHMEL_FPB`.
|
|
176
|
+
|
|
177
|
+
## Development
|
|
178
|
+
|
|
179
|
+
```sh
|
|
180
|
+
uv sync --extra dev # flashmel deps + pytest
|
|
181
|
+
uv pip install torch torchaudio --index-url https://download.pytorch.org/whl/cu130
|
|
182
|
+
uv run --no-sync pytest scripts/test_correctness.py # 87 tests vs torchaudio (needs CUDA)
|
|
183
|
+
uv run --no-sync python scripts/bench.py [--sweep]
|
|
184
|
+
```
|
|
185
|
+
|
|
186
|
+
torch/torchaudio aren't in the lockfile (see above), so install them into the
|
|
187
|
+
venv separately and pass `--no-sync` to `uv run` so it doesn't prune them.
|
|
188
|
+
|
|
189
|
+
Correctness thresholds per PROJECT.md: `atol=rtol=1e-4` (fp32 out),
|
|
190
|
+
`1e-2` (fp16 out, compared against the fp16-quantized reference; values beyond
|
|
191
|
+
fp16 range legitimately saturate to inf on both sides).
|
|
192
|
+
|
|
193
|
+
Layout: `flashmel/` (package: `transform.py` wrapper, `filters.py` host-side
|
|
194
|
+
window/filterbank packing, `build.py` NVRTC + cache, `runtime.py` driver-API
|
|
195
|
+
launch, `kernel.cu`), `scripts/` (tests + bench). The cubin cache lives in
|
|
196
|
+
`~/.cache/flashmel` (`FLASHMEL_CACHE_DIR` to override). cuFFTDx headers come from
|
|
197
|
+
the `nvidia-mathdx` wheel.
|
|
198
|
+
|
|
199
|
+
## License
|
|
200
|
+
|
|
201
|
+
MIT — see [LICENSE](LICENSE).
|
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
Your goal is to create a fused mel spectrogram kernel in cuFFTDx, and provide a python wrapper.
|
|
2
|
+
|
|
3
|
+
You must support all arguments of <https://docs.pytorch.org/audio/stable/generated/torchaudio.transforms.MelSpectrogram.html>, except:
|
|
4
|
+
|
|
5
|
+
- `window_fn` and `wkwargs`: always use the default
|
|
6
|
+
- `n_fft`: assume the value is either 400 or power-of-2 between 2^6 and 2^14
|
|
7
|
+
- `n_mels`: assume a maximum of 256
|
|
8
|
+
|
|
9
|
+
And additionally support:
|
|
10
|
+
|
|
11
|
+
- input dtype: i16 (raw pcm) or f32
|
|
12
|
+
- output dtype: fp16 or fp32
|
|
13
|
+
|
|
14
|
+
Write proper scripts for correctness test and timing before writing the kernel. Use `atol=1e-4, rtol=1e-4` (100x for fp16) as correctness threshold and cuda events for timing.
|
|
15
|
+
|
|
16
|
+
Keep optimizing until your kernel saturates on memory bandwidth (likely) or compute.
|
|
17
|
+
|
|
18
|
+
The project is empty now - decide a good directory structure yourself. See ../mel-test for a relevant example.
|
flashmel-0.1.0/README.md
ADDED
|
@@ -0,0 +1,167 @@
|
|
|
1
|
+
# flashmel
|
|
2
|
+
|
|
3
|
+
Fused CUDA mel spectrogram: framing + padding + Hann window + R2C FFT (cuFFTDx)
|
|
4
|
+
+ `|X|^power` + sparse mel filterbank projection in a **single kernel launch**,
|
|
5
|
+
with a drop-in `torchaudio.transforms.MelSpectrogram`-compatible Python wrapper.
|
|
6
|
+
|
|
7
|
+
3–14× faster than torchaudio on an RTX 4070 Laptop, bit-matching torchaudio at
|
|
8
|
+
`atol=rtol=1e-4` (fp32).
|
|
9
|
+
|
|
10
|
+
## Installation
|
|
11
|
+
|
|
12
|
+
flashmel needs Python ≥ 3.10, an NVIDIA GPU, a recent NVIDIA driver, and a CUDA
|
|
13
|
+
build of **torch** (and **torchaudio**) — install those first, the way your
|
|
14
|
+
environment needs them, e.g.:
|
|
15
|
+
|
|
16
|
+
```sh
|
|
17
|
+
pip install torch torchaudio --index-url https://download.pytorch.org/whl/cu130
|
|
18
|
+
```
|
|
19
|
+
|
|
20
|
+
flashmel deliberately does **not** depend on torch/torchaudio: the CUDA wheel
|
|
21
|
+
stack is environment-specific (driver version, CUDA version, sometimes a private
|
|
22
|
+
mirror), and pinning it here would only fight your install. Then:
|
|
23
|
+
|
|
24
|
+
```sh
|
|
25
|
+
pip install flashmel
|
|
26
|
+
```
|
|
27
|
+
|
|
28
|
+
This pulls flashmel's own deps — the cuFFTDx/CUDA headers and NVRTC
|
|
29
|
+
(`nvidia-mathdx`, `nvidia-cuda-cccl`, `nvidia-cuda-runtime`, `nvidia-cuda-nvrtc`),
|
|
30
|
+
the driver bindings (`cuda-python`), and `numpy`. No system CUDA toolkit or
|
|
31
|
+
`nvcc` is required. The kernel is compiled on first use per configuration (~2 s)
|
|
32
|
+
and cached on disk (see below).
|
|
33
|
+
|
|
34
|
+
## Usage
|
|
35
|
+
|
|
36
|
+
```python
|
|
37
|
+
import torch
|
|
38
|
+
from flashmel import MelSpectrogram
|
|
39
|
+
|
|
40
|
+
mel = MelSpectrogram(n_fft=400, hop_length=160, n_mels=128) # torchaudio args
|
|
41
|
+
x = torch.randn(8, 480000, device="cuda") # f32 or i16 CUDA tensor
|
|
42
|
+
y = mel(x) # [8, 128, 3001]
|
|
43
|
+
```
|
|
44
|
+
|
|
45
|
+
Supports every `torchaudio.transforms.MelSpectrogram` argument
|
|
46
|
+
(`sample_rate, n_fft, win_length, hop_length, f_min, f_max, pad, n_mels, power,
|
|
47
|
+
normalized, center, pad_mode, onesided, norm, mel_scale`) except
|
|
48
|
+
`window_fn`/`wkwargs` (always periodic Hann). Additionally:
|
|
49
|
+
|
|
50
|
+
- **input dtype**: `float32`, or `int16` raw PCM (scaled by 1/32768) — dispatched
|
|
51
|
+
from the input tensor's dtype;
|
|
52
|
+
- **output dtype**: `dtype=torch.float32` (default) or `torch.float16`;
|
|
53
|
+
- input shape `(..., time)`, output `(..., n_mels, n_frames)`, like torchaudio.
|
|
54
|
+
|
|
55
|
+
Constraints: `n_fft` must be 400 or a power of 2 in [64, 16384]; `n_mels ≤ 256`;
|
|
56
|
+
`onesided=True` only (torchaudio's own `MelScale` can't consume two-sided
|
|
57
|
+
spectrograms either); input must be on CUDA.
|
|
58
|
+
|
|
59
|
+
## How it works
|
|
60
|
+
|
|
61
|
+
- **NVRTC-compiled** `flashmel/kernel.cu`, one cubin per parameter configuration,
|
|
62
|
+
disk-cached by content hash in the user cache dir (`~/.cache/flashmel`,
|
|
63
|
+
overridable via `FLASHMEL_CACHE_DIR`). No build step; first call per config
|
|
64
|
+
compiles in ~2 s.
|
|
65
|
+
- Frames are loaded straight into cuFFTDx registers with the signal padding
|
|
66
|
+
(`pad`, `center`, all four `pad_mode`s) applied as index arithmetic — the
|
|
67
|
+
padded/framed signal is never materialized.
|
|
68
|
+
- The STFT `normalized` modes are folded into the window vector on the host.
|
|
69
|
+
- Power-of-2 sizes use cuFFTDx `real_mode::folded` R2C (an N/2-point complex
|
|
70
|
+
FFT), halving shared memory and butterflies; 400 uses `real_mode::normal`.
|
|
71
|
+
- The FFT's shared-memory workspace is reused as the per-frame power tile
|
|
72
|
+
(odd row stride ⇒ bank-conflict-free), from which the sparse mel projection
|
|
73
|
+
reads. Mel filters are triangular ⇒ contiguous support, packed as
|
|
74
|
+
`start[m] / width[m] / weights[m, W_MAX]`; the inner loop runs to each
|
|
75
|
+
filter's actual width, which cut L1/SMEM traffic ~40% for wide-filter
|
|
76
|
+
configs (1024: 0.28→0.23 ms, 4096: 0.57→0.41 ms, 16384: 1.35→0.94 ms).
|
|
77
|
+
- When the block has more threads than mel outputs (large n_fft ⇒ FPB=1),
|
|
78
|
+
SPLIT lanes cooperate per dot product with a warp-shuffle reduce instead of
|
|
79
|
+
idling (16384: another 12%).
|
|
80
|
+
- Folded-mode interior frames (no padding in play) load signal and window as
|
|
81
|
+
2-wide vectors, halving input-phase load instructions (64: −30%, 1024/4096:
|
|
82
|
+
−13%; the LSU pipe, not DRAM, is the binding resource at these sizes).
|
|
83
|
+
- Normal-mode (400) interior frames take a direct-load path that skips the
|
|
84
|
+
per-sample reflect/boundary arithmetic entirely (whisper case: 1.08→0.54 ms,
|
|
85
|
+
2.0×, now at copy peak).
|
|
86
|
+
- Launches via the CUDA driver API on torch's current stream; fp16 output is
|
|
87
|
+
cast at the final store (all accumulation in fp32).
|
|
88
|
+
- The launch is wrapped in a `torch.library` custom op (`flashmel::mel_forward`
|
|
89
|
+
with a fake impl), so `torch.compile(..., fullgraph=True)` traces through it.
|
|
90
|
+
- Kernel modules are cached per (config, arch, device) — CUmodules are
|
|
91
|
+
context-bound — and filter/window buffers are cached per device, so multiple
|
|
92
|
+
GPUs work concurrently. Inference-only: requires-grad input under grad mode
|
|
93
|
+
raises (no backward is implemented).
|
|
94
|
+
|
|
95
|
+
## Performance (RTX 4070 Laptop, sm_89, measured copy peak ≈ 198 GB/s)
|
|
96
|
+
|
|
97
|
+
`uv run python scripts/bench.py` (CUDA-event timing, median of 50; "BW" counts
|
|
98
|
+
essential bytes = input + output only):
|
|
99
|
+
|
|
100
|
+
| case | config | batch | ours | torchaudio | speedup | essential BW |
|
|
101
|
+
|---|---|---|---|---|---|---|
|
|
102
|
+
| whisper | 400/160, 128 mel, 30 s | 32 | 0.54 ms | 6.34 ms | 11.7× | 205 GB/s (103%) |
|
|
103
|
+
| whisper i16→fp16 | 400/160, 128 mel | 32 | 0.53 ms | 7.28 ms | 13.8× | 105 GB/s |
|
|
104
|
+
| speech | 1024/256, 80 mel | 32 | 0.20 ms | 3.15 ms | 16.2× | 138 GB/s (70%) |
|
|
105
|
+
| small | 64/16, 40 mel | 64 | 0.051 ms | 0.29 ms | 5.6× | 280 GB/s (141%) |
|
|
106
|
+
| music | 4096/1024, 256 mel, 44.1k | 16 | 0.35 ms | 4.54 ms | 12.9× | 100 GB/s (51%) |
|
|
107
|
+
| huge | 16384/4096, 256 mel, 44.1k | 16 | 0.92 ms | 4.61 ms | 5.0× | 33 GB/s |
|
|
108
|
+
|
|
109
|
+
Full-spectrum sweep (hop = n_fft/4, 128 mel, B=32, T=160000): 64–512 run at
|
|
110
|
+
97–120% of copy peak, 1024 at 82%, then declining as the FFT itself dominates.
|
|
111
|
+
Numbers drift ±15% with laptop thermals.
|
|
112
|
+
|
|
113
|
+
### Where each case actually sits (Nsight Compute)
|
|
114
|
+
|
|
115
|
+
- **small–512**: DRAM-saturated. Effective BW exceeds copy peak because
|
|
116
|
+
overlapped frame reads hit L2 (essential-byte metric counts them once).
|
|
117
|
+
- **speech (1024) / music (4096)**: L1/SMEM-throughput bound at 76–81% SOL
|
|
118
|
+
with DRAM at 30–40% — the remaining L1 traffic is cuFFTDx's internal
|
|
119
|
+
register exchanges, not under our control (vectorizing our own smem reads is
|
|
120
|
+
blocked by the odd power-tile stride that keeps them bank-conflict-free).
|
|
121
|
+
- **whisper (400)**: DRAM-saturated since the interior direct-load path
|
|
122
|
+
(DRAM 77% / L1 80% / SM 74% SOL — 77% of theoretical DRAM ≈ the measured
|
|
123
|
+
copy peak; previously latency-bound at DRAM 48% / SM 47% with the reflect
|
|
124
|
+
arithmetic on every sample). Occupancy stays at 41%, capped by the FFT
|
|
125
|
+
workspace (12.8 KB per 4-frame block ⇒ 7 blocks/SM; normal-mode R2C is the
|
|
126
|
+
only cuFFTDx option for non-power-of-2), but that no longer binds. A paired
|
|
127
|
+
c2c trick (two frames per complex FFT via conjugate symmetry) halves FFT
|
|
128
|
+
work but measured ~1.5× *slower*: the full Z-spectrum round trip through
|
|
129
|
+
shared memory costs more than the saved butterflies. Implemented, A/B'd,
|
|
130
|
+
removed.
|
|
131
|
+
- **huge (8192/16384)**: occupancy-bound (33–64 KB workspace ⇒ 1–2 blocks/SM)
|
|
132
|
+
and barrier-serialized inside the block FFT. Inherent to a single-kernel
|
|
133
|
+
block FFT at these sizes; going faster would mean a multi-kernel FFT (and
|
|
134
|
+
materializing the spectrum in HBM) or fp16 twiddles (breaking the 1e-4
|
|
135
|
+
contract).
|
|
136
|
+
|
|
137
|
+
### Tuning
|
|
138
|
+
|
|
139
|
+
Per-`n_fft` frames-per-block defaults live in `flashmel/transform.py`
|
|
140
|
+
(`_FPB_DEFAULTS`), found with `scripts/bench.py --sweep`. Override per process
|
|
141
|
+
with `FLASHMEL_FPB`.
|
|
142
|
+
|
|
143
|
+
## Development
|
|
144
|
+
|
|
145
|
+
```sh
|
|
146
|
+
uv sync --extra dev # flashmel deps + pytest
|
|
147
|
+
uv pip install torch torchaudio --index-url https://download.pytorch.org/whl/cu130
|
|
148
|
+
uv run --no-sync pytest scripts/test_correctness.py # 87 tests vs torchaudio (needs CUDA)
|
|
149
|
+
uv run --no-sync python scripts/bench.py [--sweep]
|
|
150
|
+
```
|
|
151
|
+
|
|
152
|
+
torch/torchaudio aren't in the lockfile (see above), so install them into the
|
|
153
|
+
venv separately and pass `--no-sync` to `uv run` so it doesn't prune them.
|
|
154
|
+
|
|
155
|
+
Correctness thresholds per PROJECT.md: `atol=rtol=1e-4` (fp32 out),
|
|
156
|
+
`1e-2` (fp16 out, compared against the fp16-quantized reference; values beyond
|
|
157
|
+
fp16 range legitimately saturate to inf on both sides).
|
|
158
|
+
|
|
159
|
+
Layout: `flashmel/` (package: `transform.py` wrapper, `filters.py` host-side
|
|
160
|
+
window/filterbank packing, `build.py` NVRTC + cache, `runtime.py` driver-API
|
|
161
|
+
launch, `kernel.cu`), `scripts/` (tests + bench). The cubin cache lives in
|
|
162
|
+
`~/.cache/flashmel` (`FLASHMEL_CACHE_DIR` to override). cuFFTDx headers come from
|
|
163
|
+
the `nvidia-mathdx` wheel.
|
|
164
|
+
|
|
165
|
+
## License
|
|
166
|
+
|
|
167
|
+
MIT — see [LICENSE](LICENSE).
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
"""flashmel: fused CUDA mel spectrogram (cuFFTDx), torchaudio-compatible."""
|
|
2
|
+
|
|
3
|
+
from importlib.metadata import PackageNotFoundError, version
|
|
4
|
+
|
|
5
|
+
from .transform import MelSpectrogram
|
|
6
|
+
|
|
7
|
+
try:
|
|
8
|
+
__version__ = version("flashmel")
|
|
9
|
+
except PackageNotFoundError: # not installed (e.g. running from a source tree)
|
|
10
|
+
__version__ = "0.0.0+unknown"
|
|
11
|
+
|
|
12
|
+
__all__ = ["MelSpectrogram", "__version__"]
|
|
@@ -0,0 +1,80 @@
|
|
|
1
|
+
"""NVRTC compilation of kernel.cu -> cubin, cached on disk by content hash."""
|
|
2
|
+
|
|
3
|
+
import hashlib
|
|
4
|
+
import importlib.util
|
|
5
|
+
import os
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
|
|
8
|
+
from cuda.bindings import nvrtc
|
|
9
|
+
|
|
10
|
+
_KERNEL_SRC = Path(__file__).with_name("kernel.cu")
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def _cache_dir() -> Path:
|
|
14
|
+
"""Per-user, env-overridable cubin cache (never inside site-packages)."""
|
|
15
|
+
if env := os.environ.get("FLASHMEL_CACHE_DIR"):
|
|
16
|
+
return Path(env)
|
|
17
|
+
base = os.environ.get("XDG_CACHE_HOME") or (Path.home() / ".cache")
|
|
18
|
+
return Path(base) / "flashmel"
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def _nvidia_root() -> Path:
|
|
22
|
+
# "nvidia" is a namespace package: pick the location that actually has mathdx
|
|
23
|
+
# (stray "nvidia" directories elsewhere on sys.path also join the namespace).
|
|
24
|
+
spec = importlib.util.find_spec("nvidia")
|
|
25
|
+
locations = [Path(p) for p in spec.submodule_search_locations]
|
|
26
|
+
for p in locations:
|
|
27
|
+
if (p / "mathdx" / "include").is_dir():
|
|
28
|
+
return p
|
|
29
|
+
raise RuntimeError(f"nvidia-mathdx include dir not found in {locations}")
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def include_dirs() -> list[Path]:
|
|
33
|
+
nv = _nvidia_root()
|
|
34
|
+
return [
|
|
35
|
+
nv / "mathdx" / "include",
|
|
36
|
+
nv / "mathdx" / "external" / "cutlass" / "include",
|
|
37
|
+
nv / "cu13" / "include",
|
|
38
|
+
nv / "cu13" / "include" / "cccl",
|
|
39
|
+
]
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _check(ret):
|
|
43
|
+
err, *rest = ret
|
|
44
|
+
if err != nvrtc.nvrtcResult.NVRTC_SUCCESS:
|
|
45
|
+
raise RuntimeError(f"NVRTC error: {err}")
|
|
46
|
+
return rest[0] if len(rest) == 1 else rest
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def compile_kernel(macros: dict[str, str], arch: str) -> bytes:
|
|
50
|
+
"""Compile kernel.cu with the given -D macros for sm_{arch}; returns cubin bytes."""
|
|
51
|
+
src = _KERNEL_SRC.read_text()
|
|
52
|
+
opts = [
|
|
53
|
+
"--std=c++20",
|
|
54
|
+
f"--gpu-architecture=sm_{arch}",
|
|
55
|
+
"--device-as-default-execution-space",
|
|
56
|
+
"--use_fast_math",
|
|
57
|
+
*(f"-D{k}={v}" for k, v in sorted(macros.items())),
|
|
58
|
+
*(f"--include-path={d}" for d in include_dirs()),
|
|
59
|
+
]
|
|
60
|
+
key = hashlib.sha256("\0".join([src, *opts]).encode()).hexdigest()[:16]
|
|
61
|
+
cache_dir = _cache_dir()
|
|
62
|
+
cached = cache_dir / f"flashmel_{key}.cubin"
|
|
63
|
+
if cached.exists():
|
|
64
|
+
return cached.read_bytes()
|
|
65
|
+
|
|
66
|
+
prog = _check(nvrtc.nvrtcCreateProgram(src.encode(), b"kernel.cu", 0, [], []))
|
|
67
|
+
try:
|
|
68
|
+
res = nvrtc.nvrtcCompileProgram(prog, len(opts), [o.encode() for o in opts])[0]
|
|
69
|
+
log = bytearray(_check(nvrtc.nvrtcGetProgramLogSize(prog)))
|
|
70
|
+
_check(nvrtc.nvrtcGetProgramLog(prog, log))
|
|
71
|
+
if res != nvrtc.nvrtcResult.NVRTC_SUCCESS:
|
|
72
|
+
raise RuntimeError(f"NVRTC compilation failed:\n{log.decode()}")
|
|
73
|
+
cubin = bytearray(_check(nvrtc.nvrtcGetCUBINSize(prog)))
|
|
74
|
+
_check(nvrtc.nvrtcGetCUBIN(prog, cubin))
|
|
75
|
+
finally:
|
|
76
|
+
nvrtc.nvrtcDestroyProgram(prog)
|
|
77
|
+
|
|
78
|
+
cache_dir.mkdir(parents=True, exist_ok=True)
|
|
79
|
+
cached.write_bytes(bytes(cubin))
|
|
80
|
+
return bytes(cubin)
|
|
@@ -0,0 +1,48 @@
|
|
|
1
|
+
"""Window construction and sparse mel filterbank packing (CPU, fp32)."""
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
from torchaudio.functional import melscale_fbanks
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def make_window(n_fft: int, win_length: int, normalized) -> torch.Tensor:
|
|
8
|
+
"""Hann window with STFT normalization folded in, zero-padded centered to n_fft."""
|
|
9
|
+
w = torch.hann_window(win_length, periodic=True, dtype=torch.float32)
|
|
10
|
+
if normalized is True or normalized == "window":
|
|
11
|
+
w = w / w.pow(2).sum().sqrt()
|
|
12
|
+
elif normalized == "frame_length":
|
|
13
|
+
w = w / win_length**0.5
|
|
14
|
+
if win_length < n_fft:
|
|
15
|
+
left = (n_fft - win_length) // 2
|
|
16
|
+
w = torch.nn.functional.pad(w, (left, n_fft - win_length - left))
|
|
17
|
+
return w.contiguous()
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def pack_fbank(
|
|
21
|
+
n_freqs: int,
|
|
22
|
+
f_min: float,
|
|
23
|
+
f_max: float,
|
|
24
|
+
n_mels: int,
|
|
25
|
+
sample_rate: int,
|
|
26
|
+
norm,
|
|
27
|
+
mel_scale: str,
|
|
28
|
+
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int]:
|
|
29
|
+
"""Pack the (triangular => contiguous-support) mel filterbank as sparse rows.
|
|
30
|
+
|
|
31
|
+
Returns (start [n_mels] i32, width [n_mels] i32, weights [n_mels, w_max] f32,
|
|
32
|
+
w_max) with start[m] clamped so start[m] + w_max <= n_freqs; weights are the
|
|
33
|
+
exact melscale_fbanks values (zero outside each filter's support).
|
|
34
|
+
"""
|
|
35
|
+
fb = melscale_fbanks(n_freqs, f_min, f_max, n_mels, sample_rate, norm, mel_scale)
|
|
36
|
+
fb = fb.T.contiguous().to(torch.float32) # (n_mels, n_freqs)
|
|
37
|
+
start = torch.zeros(n_mels, dtype=torch.int64)
|
|
38
|
+
width = torch.zeros(n_mels, dtype=torch.int64)
|
|
39
|
+
for m in range(n_mels):
|
|
40
|
+
nz = fb[m].nonzero().flatten()
|
|
41
|
+
if nz.numel():
|
|
42
|
+
start[m] = nz[0]
|
|
43
|
+
width[m] = nz[-1] - nz[0] + 1
|
|
44
|
+
w_max = max(int(width.max()), 1)
|
|
45
|
+
clamped = start.clamp(0, n_freqs - w_max)
|
|
46
|
+
width = width + start - clamped # support end relative to the clamped start
|
|
47
|
+
weights = torch.stack([fb[m, clamped[m] : clamped[m] + w_max] for m in range(n_mels)])
|
|
48
|
+
return clamped.to(torch.int32).contiguous(), width.to(torch.int32).contiguous(), weights.contiguous(), w_max
|