flashmel 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
flashmel/__init__.py ADDED
@@ -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__"]
flashmel/build.py ADDED
@@ -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)
flashmel/filters.py ADDED
@@ -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
flashmel/kernel.cu ADDED
@@ -0,0 +1,266 @@
1
+ // Fused torchaudio-compatible MelSpectrogram kernel:
2
+ // framing + signal padding + Hann window + R2C FFT (cuFFTDx) + |X|^power
3
+ // + sparse mel filterbank projection, all in one launch.
4
+ //
5
+ // Compile-time macros injected by build.py:
6
+ // N_FFT 400 or power of 2 in [64, 16384]
7
+ // FPB frames (FFTs) per block
8
+ // W_MAX max mel filter support width
9
+ // N_MELS number of mel bins
10
+ // CENTER 0/1 torch.stft center
11
+ // PAD_MODE 0 reflect, 1 constant(zeros), 2 replicate, 3 circular
12
+ // POWER_MODE 0: |X|^2, 1: |X|, 2: |X|^POWER_VAL (generic)
13
+ // POWER_VAL float literal, e.g. 0.7f (only used when POWER_MODE == 2)
14
+ // IN_I16 0: f32 input, 1: i16 raw PCM input (scaled by 1/32768)
15
+ // OUT_HALF 0: f32 output, 1: fp16 output
16
+ // REAL_MODE 0: real_mode::normal (e.g. 400), 1: real_mode::folded (powers of 2)
17
+ // SM_ARCH e.g. 890
18
+
19
+ #include <cufftdx.hpp>
20
+ #include <cuda_fp16.h>
21
+
22
+ using namespace cufftdx;
23
+
24
+ constexpr unsigned N_FREQS = N_FFT / 2 + 1;
25
+ // Power-tile row stride: odd for every supported n_fft => rows hit distinct banks.
26
+ constexpr unsigned PSTRIDE = N_FREQS;
27
+ static_assert(PSTRIDE % 2 == 1);
28
+
29
+ #if IN_I16
30
+ using in_t = short;
31
+ #else
32
+ using in_t = float;
33
+ #endif
34
+ #if OUT_HALF
35
+ using out_t = __half;
36
+ #else
37
+ using out_t = float;
38
+ #endif
39
+
40
+ #if REAL_MODE == 1
41
+ using fft_real_opts = RealFFTOptions<complex_layout::natural, real_mode::folded>;
42
+ #else
43
+ using fft_real_opts = RealFFTOptions<complex_layout::natural, real_mode::normal>;
44
+ #endif
45
+ using FFT = decltype(Size<N_FFT>() + Precision<float>() + Type<fft_type::r2c>()
46
+ + fft_real_opts() + Direction<fft_direction::forward>()
47
+ + FFTsPerBlock<FPB>() + Block() + SM<SM_ARCH>());
48
+
49
+ using complex_type = typename FFT::value_type;
50
+
51
+ constexpr unsigned NTHR = FFT::block_dim.x * FFT::block_dim.y;
52
+
53
+ // Mel projection lanes per output: when the block has more threads than mel
54
+ // outputs (large n_fft => FPB=1), SPLIT threads cooperate on one dot product
55
+ // and reduce via warp shuffle instead of idling.
56
+ constexpr unsigned OUTPUTS = N_MELS * FPB;
57
+ constexpr unsigned SPLIT = []() consteval {
58
+ // Narrow filters have nothing to split; partial warps would need shfl masks
59
+ // excluding non-existent lanes.
60
+ if (W_MAX <= 16 || NTHR % 32 != 0) return 1u;
61
+ unsigned s = 1;
62
+ while (s < 32 && 2 * s * OUTPUTS <= NTHR) s *= 2;
63
+ return s;
64
+ }();
65
+ constexpr unsigned SMEM_BYTES =
66
+ FFT::shared_memory_size > FPB * PSTRIDE * 4 ? FFT::shared_memory_size : FPB * PSTRIDE * 4;
67
+
68
+ // Fetch sample v of the virtually padded signal: the f32/i16 signal x[0..n) is
69
+ // first zero-padded by `pad` on both sides (length L = n + 2*pad), then extended
70
+ // by PAD_MODE for the centered STFT overhang. Python validates the torch.stft
71
+ // constraints that make a single reflection/wrap sufficient.
72
+ __device__ __forceinline__ float load_sample(const in_t* __restrict__ x, int v, int L, int n,
73
+ int pad) {
74
+ #if PAD_MODE == 0
75
+ v = v < 0 ? -v : v;
76
+ v = v >= L ? 2 * L - 2 - v : v;
77
+ #elif PAD_MODE == 1
78
+ if (v < 0 || v >= L) return 0.f;
79
+ #elif PAD_MODE == 2
80
+ v = v < 0 ? 0 : v;
81
+ v = v >= L ? L - 1 : v;
82
+ #else
83
+ v %= L;
84
+ v = v < 0 ? v + L : v;
85
+ #endif
86
+ v -= pad;
87
+ if (v < 0 || v >= n) return 0.f;
88
+ #if IN_I16
89
+ return (float)x[v] * (1.f / 32768.f);
90
+ #else
91
+ return x[v];
92
+ #endif
93
+ }
94
+
95
+ __device__ __forceinline__ float spec_value(complex_type v) {
96
+ const float p2 = v.real() * v.real() + v.imag() * v.imag();
97
+ #if POWER_MODE == 0
98
+ return p2;
99
+ #elif POWER_MODE == 1
100
+ return sqrtf(p2);
101
+ #else
102
+ return powf(p2, 0.5f * POWER_VAL);
103
+ #endif
104
+ }
105
+
106
+ extern "C" __global__ void flashmel(const in_t* __restrict__ in, // [B, n_samples]
107
+ out_t* __restrict__ out, // [B, N_MELS, n_frames]
108
+ const float* __restrict__ window, // [N_FFT]
109
+ const int* __restrict__ start, // [N_MELS]
110
+ const int* __restrict__ width, // [N_MELS]
111
+ const float* __restrict__ weights, // [N_MELS, W_MAX]
112
+ int n_samples, int pad, int hop, int n_frames) {
113
+ complex_type thread_data[FFT::storage_size];
114
+ extern __shared__ __align__(16) unsigned char smem_raw[];
115
+ float* power = reinterpret_cast<float*>(smem_raw);
116
+
117
+ const unsigned b = blockIdx.y;
118
+ const unsigned frame0 = blockIdx.x * FPB;
119
+ const int L = n_samples + 2 * pad;
120
+ const in_t* x = in + (size_t)b * n_samples;
121
+
122
+ // Framing + padding + window, straight into FFT registers.
123
+ const unsigned frame = frame0 + threadIdx.y;
124
+ const bool valid = frame < (unsigned)n_frames;
125
+ const int base = (int)frame * hop - (CENTER ? (int)(N_FFT / 2) : 0);
126
+ #if REAL_MODE == 1
127
+ // Folded R2C: N_FFT reals enter as N_FFT/2 complex (even, odd) pairs.
128
+ #if IN_I16
129
+ using in2_t = short2;
130
+ #else
131
+ using in2_t = float2;
132
+ #endif
133
+ // Interior frames (no padding in play, pointer pair-aligned — odd batch row
134
+ // lengths shift row bases by 4 bytes) take vector loads: one 2-wide load
135
+ // each for signal and window instead of two scalars.
136
+ const in_t* xf = x + (base - pad);
137
+ if (valid && base >= pad && base + (int)N_FFT <= pad + n_samples &&
138
+ (reinterpret_cast<uintptr_t>(xf) & (sizeof(in2_t) - 1)) == 0) {
139
+ const in2_t* xp = reinterpret_cast<const in2_t*>(xf);
140
+ const float2* wp = reinterpret_cast<const float2*>(window);
141
+ #pragma unroll
142
+ for (unsigned i = 0; i < FFT::input_ept; ++i) {
143
+ const unsigned c = threadIdx.x + i * FFT::stride;
144
+ float re = 0.f, im = 0.f;
145
+ if (c < FFT::input_length) {
146
+ const in2_t s = xp[c];
147
+ const float2 wv = wp[c];
148
+ #if IN_I16
149
+ re = (float)s.x * (1.f / 32768.f) * wv.x;
150
+ im = (float)s.y * (1.f / 32768.f) * wv.y;
151
+ #else
152
+ re = s.x * wv.x;
153
+ im = s.y * wv.y;
154
+ #endif
155
+ }
156
+ thread_data[i] = complex_type(re, im);
157
+ }
158
+ } else {
159
+ #pragma unroll
160
+ for (unsigned i = 0; i < FFT::input_ept; ++i) {
161
+ const unsigned c = threadIdx.x + i * FFT::stride;
162
+ float re = 0.f, im = 0.f;
163
+ if (valid && c < FFT::input_length) {
164
+ re = load_sample(x, base + (int)(2 * c), L, n_samples, pad) * window[2 * c];
165
+ im = load_sample(x, base + (int)(2 * c + 1), L, n_samples, pad) * window[2 * c + 1];
166
+ }
167
+ thread_data[i] = complex_type(re, im);
168
+ }
169
+ }
170
+ #else
171
+ auto* reg = reinterpret_cast<float*>(thread_data);
172
+ // Interior frames (no padding in play) skip the per-sample boundary logic.
173
+ const in_t* xf = x + (base - pad);
174
+ if (valid && base >= pad && base + (int)N_FFT <= pad + n_samples) {
175
+ #pragma unroll
176
+ for (unsigned i = 0; i < FFT::input_ept; ++i) {
177
+ const unsigned idx = threadIdx.x + i * FFT::stride;
178
+ float s = 0.f;
179
+ if (idx < FFT::input_length) {
180
+ #if IN_I16
181
+ s = (float)xf[idx] * (1.f / 32768.f) * window[idx];
182
+ #else
183
+ s = xf[idx] * window[idx];
184
+ #endif
185
+ }
186
+ reg[i] = s;
187
+ }
188
+ } else {
189
+ #pragma unroll
190
+ for (unsigned i = 0; i < FFT::input_ept; ++i) {
191
+ const unsigned idx = threadIdx.x + i * FFT::stride;
192
+ reg[i] = (valid && idx < FFT::input_length)
193
+ ? load_sample(x, base + (int)idx, L, n_samples, pad) * window[idx]
194
+ : 0.f;
195
+ }
196
+ }
197
+ #endif
198
+
199
+ FFT().execute(thread_data, reinterpret_cast<complex_type*>(smem_raw));
200
+ __syncthreads();
201
+
202
+ // FFT workspace is dead; reuse it as the power tile.
203
+ float* prow = power + threadIdx.y * PSTRIDE;
204
+ #pragma unroll
205
+ for (unsigned i = 0; i < FFT::output_ept; ++i) {
206
+ const unsigned idx = threadIdx.x + i * FFT::stride;
207
+ if (idx < N_FREQS) prow[idx] = spec_value(thread_data[i]);
208
+ }
209
+ __syncthreads();
210
+
211
+ // Sparse mel projection; consecutive threads take consecutive frames of the
212
+ // same mel row -> coalesced FPB-wide segments in out.
213
+ const unsigned tid = threadIdx.y * FFT::block_dim.x + threadIdx.x;
214
+ out_t* o = out + (size_t)b * N_MELS * n_frames + frame0;
215
+ if constexpr (SPLIT > 1) {
216
+ const unsigned lane = tid & (SPLIT - 1);
217
+ const unsigned group = tid / SPLIT;
218
+ constexpr unsigned NGROUPS = NTHR / SPLIT;
219
+ #pragma unroll 1
220
+ for (unsigned pb = 0; pb < OUTPUTS; pb += NGROUPS) { // uniform: all lanes shuffle
221
+ const unsigned p = pb + group;
222
+ const unsigned m = p / FPB, f = p % FPB;
223
+ const bool active = p < OUTPUTS && frame0 + f < (unsigned)n_frames;
224
+ float acc = 0.f;
225
+ if (active) {
226
+ const float* w = weights + (size_t)m * W_MAX;
227
+ const float* pr = power + f * PSTRIDE + start[m];
228
+ const int wm = width[m];
229
+ #pragma unroll 4
230
+ for (int j = (int)lane; j < wm; j += SPLIT) acc = fmaf(w[j], pr[j], acc);
231
+ }
232
+ #pragma unroll
233
+ for (unsigned s = SPLIT / 2; s > 0; s >>= 1)
234
+ acc += __shfl_down_sync(0xffffffffu, acc, s, SPLIT);
235
+ if (active && lane == 0) o[(size_t)m * n_frames + f] = (out_t)acc;
236
+ }
237
+ } else {
238
+ #pragma unroll 1
239
+ for (unsigned p = tid; p < OUTPUTS; p += NTHR) {
240
+ const unsigned m = p / FPB, f = p % FPB;
241
+ if (frame0 + f >= (unsigned)n_frames) continue;
242
+ const float* w = weights + (size_t)m * W_MAX;
243
+ const float* pr = power + f * PSTRIDE + start[m];
244
+ float acc = 0.f;
245
+ #if W_MAX <= 16
246
+ // Narrow filters: fixed fully-unrolled loop over the zero-padded row.
247
+ #pragma unroll
248
+ for (int j = 0; j < W_MAX; ++j) acc = fmaf(w[j], pr[j], acc);
249
+ #else
250
+ // Wide filters: most are much narrower than W_MAX; looping to the actual
251
+ // support width cuts shared-memory traffic (the L1/SMEM bottleneck).
252
+ const int wm = width[m];
253
+ #pragma unroll 4
254
+ for (int j = 0; j < wm; ++j) acc = fmaf(w[j], pr[j], acc);
255
+ #endif
256
+ o[(size_t)m * n_frames + f] = (out_t)acc;
257
+ }
258
+ }
259
+ }
260
+
261
+ extern "C" __global__ void flashmel_info(unsigned* o) {
262
+ o[0] = FFT::block_dim.x;
263
+ o[1] = FFT::block_dim.y;
264
+ o[2] = SMEM_BYTES;
265
+ o[3] = FPB;
266
+ }
flashmel/runtime.py ADDED
@@ -0,0 +1,65 @@
1
+ """CUDA driver-API module loading and kernel launches on torch's current stream."""
2
+
3
+ import functools
4
+
5
+ import numpy as np
6
+ import torch
7
+ from cuda.bindings import driver
8
+
9
+ from .build import compile_kernel
10
+
11
+
12
+ def _check(ret):
13
+ err, *rest = ret
14
+ if err != driver.CUresult.CUDA_SUCCESS:
15
+ raise RuntimeError(f"CUDA driver error: {err}")
16
+ return rest[0] if len(rest) == 1 else rest
17
+
18
+
19
+ def _launch_raw(fn, grid, block, smem, stream, *args):
20
+ vals = np.array(args, dtype=np.uint64)
21
+ ptrs = np.array([vals.ctypes.data + 8 * i for i in range(len(args))], dtype=np.uint64)
22
+ _check(driver.cuLaunchKernel(fn, *grid, *block, smem, stream, ptrs.ctypes.data, 0))
23
+
24
+
25
+ class KernelModule:
26
+ """One compiled (cubin) flashmel kernel for a fixed macro configuration.
27
+
28
+ A CUmodule is bound to the CUDA context it was loaded in, so each instance
29
+ is tied to one device: it loads under (and launches on) torch's primary
30
+ context for `device_index`.
31
+ """
32
+
33
+ def __init__(self, macros: dict[str, str], arch: str, device_index: int) -> None:
34
+ cubin = compile_kernel(macros, arch)
35
+ self.device_index = device_index
36
+ with torch.cuda.device(device_index):
37
+ self.mod = _check(driver.cuModuleLoadData(cubin))
38
+ self.fn = _check(driver.cuModuleGetFunction(self.mod, b"flashmel"))
39
+ info_fn = _check(driver.cuModuleGetFunction(self.mod, b"flashmel_info"))
40
+
41
+ info = torch.zeros(4, dtype=torch.uint32, device=f"cuda:{device_index}")
42
+ stream = torch.cuda.current_stream().cuda_stream
43
+ _launch_raw(info_fn, (1, 1, 1), (1, 1, 1), 0, stream, info.data_ptr())
44
+ torch.cuda.synchronize()
45
+ bx, by, self.smem, self.fpb = info.cpu().tolist()
46
+ self.block = (bx, by, 1)
47
+ if self.smem > 48 * 1024:
48
+ _check(
49
+ driver.cuFuncSetAttribute(
50
+ self.fn,
51
+ driver.CUfunction_attribute.CU_FUNC_ATTRIBUTE_MAX_DYNAMIC_SHARED_SIZE_BYTES,
52
+ self.smem,
53
+ )
54
+ )
55
+
56
+ def launch(self, n_frames: int, batch: int, *args) -> None:
57
+ with torch.cuda.device(self.device_index):
58
+ stream = torch.cuda.current_stream().cuda_stream
59
+ grid = ((n_frames + self.fpb - 1) // self.fpb, batch, 1)
60
+ _launch_raw(self.fn, grid, self.block, self.smem, stream, *args)
61
+
62
+
63
+ @functools.lru_cache(maxsize=None)
64
+ def get_module(macro_items: tuple[tuple[str, str], ...], arch: str, device_index: int) -> KernelModule:
65
+ return KernelModule(dict(macro_items), arch, device_index)
flashmel/transform.py ADDED
@@ -0,0 +1,268 @@
1
+ """MelSpectrogram transform mirroring torchaudio.transforms.MelSpectrogram."""
2
+
3
+ import math
4
+ import os
5
+
6
+ import torch
7
+
8
+ from .filters import make_window, pack_fbank
9
+ from .runtime import KernelModule, get_module
10
+
11
+ _VALID_N_FFT = {400} | {1 << k for k in range(6, 15)}
12
+ _PAD_MODES = {"reflect": 0, "constant": 1, "replicate": 2, "circular": 3}
13
+ _MAX_N_MELS = 256
14
+ _MAX_GRID_Y = 65535
15
+
16
+ # Per-n_fft tuning defaults from scripts/bench.py --sweep on RTX 4070 Laptop
17
+ # (overridable via FLASHMEL_FPB).
18
+ _FPB_DEFAULTS = {64: 32, 128: 32, 256: 16, 400: 8, 512: 8, 1024: 4, 2048: 4, 4096: 2, 8192: 1, 16384: 1}
19
+
20
+
21
+ def _resolve_module(
22
+ n_fft: int,
23
+ w_max: int,
24
+ n_mels: int,
25
+ center: bool,
26
+ pad_mode: str,
27
+ power: float,
28
+ in_i16: bool,
29
+ out_half: bool,
30
+ device: torch.device,
31
+ ) -> KernelModule:
32
+ if power == 2.0:
33
+ power_mode = 0
34
+ elif power == 1.0:
35
+ power_mode = 1
36
+ else:
37
+ power_mode = 2
38
+ major, minor = torch.cuda.get_device_capability(device)
39
+ arch = f"{major}{minor}"
40
+ fpb = int(os.environ.get("FLASHMEL_FPB", _FPB_DEFAULTS[n_fft]))
41
+ last_err: Exception | None = None
42
+ while fpb >= 1:
43
+ macros = {
44
+ "N_FFT": str(n_fft),
45
+ "FPB": f"{fpb}u",
46
+ "W_MAX": str(w_max),
47
+ "N_MELS": str(n_mels),
48
+ "CENTER": "1" if center else "0",
49
+ "PAD_MODE": str(_PAD_MODES[pad_mode]),
50
+ "POWER_MODE": str(power_mode),
51
+ "POWER_VAL": f"{power!r}f",
52
+ "IN_I16": "1" if in_i16 else "0",
53
+ "OUT_HALF": "1" if out_half else "0",
54
+ # powers of 2: cuFFTDx folded R2C; 400: normal-mode R2C (see README for
55
+ # the rejected paired-c2c alternative).
56
+ "REAL_MODE": "1" if n_fft != 400 else "0",
57
+ "SM_ARCH": f"{arch}0",
58
+ }
59
+ try:
60
+ return get_module(tuple(sorted(macros.items())), arch, device.index)
61
+ except RuntimeError as e: # invalid cuFFTDx config (e.g. too many threads)
62
+ last_err = e
63
+ fpb //= 2
64
+ raise RuntimeError(f"no valid kernel config found for n_fft={n_fft}") from last_err
65
+
66
+
67
+ def _n_frames(n: int, n_fft: int, hop_length: int, pad: int, center: bool) -> int:
68
+ L = n + 2 * pad
69
+ return L // hop_length + 1 if center else (L - n_fft) // hop_length + 1
70
+
71
+
72
+ # Opaque custom op so torch.compile traces through the transform without graph
73
+ # breaks on the driver-API launch (shapes come from the register_fake impl).
74
+ @torch.library.custom_op("flashmel::mel_forward", mutates_args=(), device_types="cuda")
75
+ def _mel_forward(
76
+ x: torch.Tensor,
77
+ window: torch.Tensor,
78
+ fb_start: torch.Tensor,
79
+ fb_width: torch.Tensor,
80
+ fb_weights: torch.Tensor,
81
+ n_fft: int,
82
+ hop_length: int,
83
+ pad: int,
84
+ center: bool,
85
+ pad_mode: str,
86
+ power: float,
87
+ out_half: bool,
88
+ ) -> torch.Tensor:
89
+ x = x.contiguous()
90
+ batch, n = x.shape
91
+ n_mels, w_max = fb_weights.shape
92
+ n_frames = _n_frames(n, n_fft, hop_length, pad, center)
93
+ out_dtype = torch.float16 if out_half else torch.float32
94
+ out = torch.empty(batch, n_mels, n_frames, device=x.device, dtype=out_dtype)
95
+ if out.numel():
96
+ mod = _resolve_module(
97
+ n_fft, w_max, n_mels, center, pad_mode, power, x.dtype == torch.int16, out_half, x.device
98
+ )
99
+ for b0 in range(0, batch, _MAX_GRID_Y):
100
+ b1 = min(b0 + _MAX_GRID_Y, batch)
101
+ mod.launch(
102
+ n_frames,
103
+ b1 - b0,
104
+ x[b0:b1].data_ptr(),
105
+ out[b0:b1].data_ptr(),
106
+ window.data_ptr(),
107
+ fb_start.data_ptr(),
108
+ fb_width.data_ptr(),
109
+ fb_weights.data_ptr(),
110
+ n,
111
+ pad,
112
+ hop_length,
113
+ n_frames,
114
+ )
115
+ return out
116
+
117
+
118
+ @_mel_forward.register_fake
119
+ def _mel_forward_fake(x, window, fb_start, fb_width, fb_weights, n_fft, hop_length, pad, center, pad_mode, power, out_half):
120
+ n_frames = _n_frames(x.shape[-1], n_fft, hop_length, pad, center)
121
+ out_dtype = torch.float16 if out_half else torch.float32
122
+ return x.new_empty((x.shape[0], fb_weights.shape[0], n_frames), dtype=out_dtype)
123
+
124
+
125
+ class MelSpectrogram(torch.nn.Module):
126
+ """Drop-in fused-CUDA replacement for torchaudio.transforms.MelSpectrogram.
127
+
128
+ Differences from torchaudio:
129
+ - window_fn/wkwargs are not accepted (always periodic Hann);
130
+ - n_fft must be 400 or a power of 2 in [64, 16384]; n_mels <= 256;
131
+ - input must be a CUDA tensor of dtype float32 or int16 (raw PCM, scaled
132
+ by 1/32768);
133
+ - `dtype` selects the output dtype: torch.float32 (default) or torch.float16;
134
+ - inference-only: raises if the input requires grad while grad mode is on.
135
+ """
136
+
137
+ def __init__(
138
+ self,
139
+ sample_rate: int = 16000,
140
+ n_fft: int = 400,
141
+ win_length: int | None = None,
142
+ hop_length: int | None = None,
143
+ f_min: float = 0.0,
144
+ f_max: float | None = None,
145
+ pad: int = 0,
146
+ n_mels: int = 128,
147
+ power: float = 2.0,
148
+ normalized: bool | str = False,
149
+ center: bool = True,
150
+ pad_mode: str = "reflect",
151
+ onesided: bool = True,
152
+ norm: str | None = None,
153
+ mel_scale: str = "htk",
154
+ dtype: torch.dtype = torch.float32,
155
+ ) -> None:
156
+ super().__init__()
157
+ if n_fft not in _VALID_N_FFT:
158
+ raise ValueError(f"n_fft must be 400 or a power of 2 in [64, 16384], got {n_fft}")
159
+ win_length = win_length if win_length is not None else n_fft
160
+ hop_length = hop_length if hop_length is not None else win_length // 2
161
+ if not 0 < win_length <= n_fft:
162
+ raise ValueError(f"expected 0 < win_length <= n_fft, got {win_length}")
163
+ if hop_length <= 0:
164
+ raise ValueError(f"hop_length must be positive, got {hop_length}")
165
+ if not 1 <= n_mels <= _MAX_N_MELS:
166
+ raise ValueError(f"n_mels must be in [1, {_MAX_N_MELS}], got {n_mels}")
167
+ if pad < 0:
168
+ raise ValueError(f"pad must be non-negative, got {pad}")
169
+ power = float(power)
170
+ if not power > 0 or not math.isfinite(power):
171
+ raise ValueError(f"power must be a positive finite float, got {power}")
172
+ if normalized not in (True, False, "window", "frame_length"):
173
+ raise ValueError(f"normalized must be bool, 'window' or 'frame_length', got {normalized!r}")
174
+ if pad_mode not in _PAD_MODES:
175
+ raise ValueError(f"pad_mode must be one of {sorted(_PAD_MODES)}, got {pad_mode!r}")
176
+ if not onesided:
177
+ raise ValueError("onesided=False is not supported (mel scale needs n_fft // 2 + 1 bins)")
178
+ if norm is not None and norm != "slaney":
179
+ raise ValueError(f"norm must be None or 'slaney', got {norm!r}")
180
+ if mel_scale not in ("htk", "slaney"):
181
+ raise ValueError(f"mel_scale must be 'htk' or 'slaney', got {mel_scale!r}")
182
+ if dtype not in (torch.float32, torch.float16):
183
+ raise ValueError(f"dtype must be torch.float32 or torch.float16, got {dtype}")
184
+
185
+ self.sample_rate = sample_rate
186
+ self.n_fft = n_fft
187
+ self.win_length = win_length
188
+ self.hop_length = hop_length
189
+ self.f_min = f_min
190
+ self.f_max = f_max or float(sample_rate // 2) # mirrors torchaudio exactly
191
+ self.pad = pad
192
+ self.n_mels = n_mels
193
+ self.power = power
194
+ self.normalized = normalized
195
+ self.center = center
196
+ self.pad_mode = pad_mode
197
+ self.norm = norm
198
+ self.mel_scale = mel_scale
199
+ self.dtype = dtype
200
+
201
+ n_freqs = n_fft // 2 + 1
202
+ start, width, weights, self.w_max = pack_fbank(
203
+ n_freqs, self.f_min, self.f_max, n_mels, sample_rate, norm, mel_scale
204
+ )
205
+ self.register_buffer("window", make_window(n_fft, win_length, normalized))
206
+ self.register_buffer("fb_start", start)
207
+ self.register_buffer("fb_width", width)
208
+ self.register_buffer("fb_weights", weights)
209
+ self._filters_by_device: dict[int, tuple[torch.Tensor, ...]] = {}
210
+
211
+ def _filters_on(self, device: torch.device) -> tuple[torch.Tensor, ...]:
212
+ """Window + filterbank tensors on `device`, cached per device (the
213
+ registered buffers stay wherever the module was moved to)."""
214
+ if self.window.device == device:
215
+ return self.window, self.fb_start, self.fb_width, self.fb_weights
216
+ cached = self._filters_by_device.get(device.index)
217
+ if cached is None:
218
+ cached = tuple(
219
+ t.to(device) for t in (self.window, self.fb_start, self.fb_width, self.fb_weights)
220
+ )
221
+ self._filters_by_device[device.index] = cached
222
+ return cached
223
+
224
+ def forward(self, waveform: torch.Tensor) -> torch.Tensor:
225
+ if waveform.dtype not in (torch.float32, torch.int16):
226
+ raise TypeError(f"input dtype must be float32 or int16, got {waveform.dtype}")
227
+ if not waveform.is_cuda:
228
+ raise ValueError("input must be a CUDA tensor")
229
+ if waveform.requires_grad and torch.is_grad_enabled():
230
+ raise RuntimeError(
231
+ "flashmel.MelSpectrogram is inference-only and does not support autograd; "
232
+ "run under torch.no_grad() or detach the input"
233
+ )
234
+
235
+ shape = waveform.shape
236
+ n = shape[-1]
237
+ L = n + 2 * self.pad
238
+ half = self.n_fft // 2
239
+ if self.center:
240
+ if self.pad_mode == "reflect" and half >= L:
241
+ raise ValueError(f"reflect padding requires input length > {half}, got {L}")
242
+ if self.pad_mode == "circular" and half > L:
243
+ raise ValueError(f"circular padding requires input length >= {half}, got {L}")
244
+ if L < 1:
245
+ raise ValueError("input is empty")
246
+ n_frames = L // self.hop_length + 1
247
+ else:
248
+ if L < self.n_fft:
249
+ raise ValueError(f"center=False requires input length >= n_fft ({self.n_fft}), got {L}")
250
+ n_frames = (L - self.n_fft) // self.hop_length + 1
251
+
252
+ x = waveform.reshape(-1, n).contiguous()
253
+ window, fb_start, fb_width, fb_weights = self._filters_on(waveform.device)
254
+ out = _mel_forward(
255
+ x,
256
+ window,
257
+ fb_start,
258
+ fb_width,
259
+ fb_weights,
260
+ self.n_fft,
261
+ self.hop_length,
262
+ self.pad,
263
+ self.center,
264
+ self.pad_mode,
265
+ self.power,
266
+ self.dtype == torch.float16,
267
+ )
268
+ return out.reshape(*shape[:-1], self.n_mels, n_frames)
@@ -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,10 @@
1
+ flashmel/__init__.py,sha256=G9wiMGmcgV1liMWuNUlt2zoNHdYg3ER4JIXOlLuL4Jg,381
2
+ flashmel/build.py,sha256=iKKjOh6EcilxNZLtVjmbeKZJhaAMoWxFhKGLsBhSHwU,2823
3
+ flashmel/filters.py,sha256=3A5xwI8qcfNGFutJSpdy6iMW7JZp7Vw5KQ33EQO3aNc,2013
4
+ flashmel/kernel.cu,sha256=S0P6nkqh3lNv2ZCPBK2J4LoclgmFdw-iZ7kPQAqw16E,10390
5
+ flashmel/runtime.py,sha256=DhEabDnfntRvxWS2QBRBmOPgDSF6HdLxyT0eBdInttA,2621
6
+ flashmel/transform.py,sha256=C1l5ypZtrB18ZLvnGYafz_T4eAbEfxptOwDxO2ZTJpE,10613
7
+ flashmel-0.1.0.dist-info/METADATA,sha256=pwmAD7VGv_Djco5MEwFJEruoLnoJquYf5QpFjIb0lUU,9872
8
+ flashmel-0.1.0.dist-info/WHEEL,sha256=W3fkpkm7-wf9vBI5Z-7s0eWkeM-spu78I8Neb98DeEg,87
9
+ flashmel-0.1.0.dist-info/licenses/LICENSE,sha256=abKhAw1lauLcmo-Lv5EUk_Hy9CK0OV0dG1QX_aVOkoY,1064
10
+ flashmel-0.1.0.dist-info/RECORD,,
@@ -0,0 +1,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: hatchling 1.32.4
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
@@ -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.