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.
@@ -0,0 +1,6 @@
1
+ __pycache__/
2
+ .venv/
3
+ uv.lock
4
+ build/
5
+ dist/
6
+ .pytest_cache/
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.
@@ -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.
@@ -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