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 +12 -0
- flashmel/build.py +80 -0
- flashmel/filters.py +48 -0
- flashmel/kernel.cu +266 -0
- flashmel/runtime.py +65 -0
- flashmel/transform.py +268 -0
- flashmel-0.1.0.dist-info/METADATA +201 -0
- flashmel-0.1.0.dist-info/RECORD +10 -0
- flashmel-0.1.0.dist-info/WHEEL +4 -0
- flashmel-0.1.0.dist-info/licenses/LICENSE +21 -0
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,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.
|