Mel-Band RoFormer — architecture, inference hot path & RTX 3050 feasibility
Written 2026-09-15. The target machine is victoria,
an NVIDIA GeForce RTX 3050 6 GB Laptop GPU (GA107,
sm_86, 6144 MiB VRAM, 144 GB/s) — the same machine behind
the Stable
Audio 3 Small analysis and the Chatterbox
V3 profiling report.
The short version: Mel-Band RoFormer is a 228 M-parameter
axial transformer over a 60-band mel-split complex STFT, of
which 88 % of the parameters and 10 % of the FLOPs sit
in the mask-estimation head. It is not attention-bound:
attention is 16.4 % of total FLOPs and roughly
22 % of wall clock, so the companion Stable Audio 3
conclusion holds here too, despite a long frequency axis.
For melband_deux the measured attention share is
higher — 30.1 % of device
time — and it is the model's largest single cost. The "not
attention-bound" verdict is a statement about FLOPs; the kernel that
implements those FLOPs is where the time goes. It is also
not bandwidth-bound and not VRAM-bound
— it runs at ~26 % of victoria's FP32 peak while using
only about a third of the card's VRAM. The cost model is unusually
clean: over 5–120 s of audio, total latency is
1 ms + 1591 ms × n_chunks with a
19 ms RMS residual.
Two levers were found and measured afterwards, and together they take the same 60 s clip from 20.6 s to 6.7 s at unchanged output: running the model in fp16, and replacing the rotary embedding's stack-based gather. See The two levers, measured.
The headline measured numbers come from demucs_golf/doc/benchreport.md,
a complete benchmark taken on this exact machine on 2026-07-02. This
pass reproduced that benchmark on a fresh run and then
explained it; the prior data stands and is cited as such throughout.
Reproducibility and source basis
| source | path / URL | used for |
|---|---|---|
| Mel-Band RoFormer paper | arXiv:2310.01809 | mel-band projection definition, L=6/L=9 configs, reported SDR and parameter counts, multi-band mask averaging |
| reference implementation | models/melband.py |
every tensor shape, layer count and parameter count below |
| chunked inference path | inference.py
_apply_melband |
chunk size, overlap, fade window, reflect padding, overlap-add |
| loader | separator.py
_load_melband, _lazy_load |
config-to-constructor mapping, effective precision |
| author's architecture walkthrough | doc/model.md,
doc/hot_path.md |
the axial path, and what was deliberately stripped from upstream |
| prior measurements | doc/benchreport.md |
all 2026-07-02 latency, RTF, peak-VRAM and SDR figures |
| SDR methodology | doc/EVAL_SUMMARY.md,
doc/sdr.md |
what the reported SDR values do and do not measure |
| released checkpoints | ckpt/melband_{vocals,instrumental,deux}/ |
config fields, safetensors header tensor shapes |
becruily model cards |
mel-band-roformer-deux, mel-band-roformer-vocals | provenance of the three checkpoints |
| victoria GPU specification | NotebookCheck RTX 3050 6GB Laptop | bus width, bandwidth, core count, clocks |
CONFIRMED means a number read directly out of a shipped config, a line of released source, a checkpoint header, a prior measured report, or a vendor specification, with the source given. DERIVED means this document computed it and the arithmetic is shown inline. MEASURED 2026-09-15 means a fresh run on victoria during this pass. Anything neither confirmed nor derived is marked UNCONFIRMED and appears in Gaps and open questions.
Harness for the fresh measurements:
/home/tensor/code/ml/hummingbird/.venv/bin/python, torch
2.13.0+cu130, CUDA 13.0, driver 580.119.02,
torch.cuda.get_device_capability() reported
(8, 6), victoria's idle VRAM 5804 MiB of 6144. A
byte-identical copy of models/melband.py was used with only
its librosa.filters.mel import replaced by an equivalent
Python implementation of the Slaney mel formula, because librosa is not
installed in that venv. The derived band widths were then checked
against the checkpoint's tensor shapes and match band for band, which
validates the whole shape chain independently of librosa.
Model family and the three local checkpoints
Three checkpoints are present locally, and they are two different architectures, not three sizes of one:
| checkpoint | dim | depth | stems | model.safetensors |
DERIVED params (header) | notes |
|---|---|---|---|---|---|---|
melband_vocals |
384 | 6 | 1 | 456.5 MB | 228,203,172 | best-measured local route |
melband_instrumental |
384 | 6 | 1 | 456.5 MB | 228,203,172 | same architecture, different training target |
melband_deux |
256 | 12 | 2 | 434.7 MB | 217,302,620 | narrower and deeper; fewer params, more compute |
melband_karaoke |
384 | 6 | 2 | 859.4 MB | not present locally | measured 2026-07-02; see below |
All four parameter counts are CONFIRMED two independent ways. The
safetensors headers contain exactly 228,203,172 elements for
melband_vocals and melband_instrumental and
217,302,620 for melband_deux, every tensor
F16. Constructing the model from
config.json yields 228,202,788 and 217,301,852 — the small
deltas are RMSNorm gamma buffers counted differently, and a
live strict=True load of the vocals checkpoint reports
0 missing, 0 unexpected keys. The per-band widths
reconstructed from band_split.to_features.*.1.weight
reproduce the librosa.filters.mel derivation exactly, band
for band.
The paper's own ablation reports 72.2 M and 94.8 M
parameters for its L=6 and L=9 Mel-RoFormer variants (arXiv:2310.01809, Table 1).
The becruily checkpoints are 2.3–3.2×
larger, which is consistent with community fine-tunes trained
at wider dim; they are not the paper's released weights,
and their SDR is not comparable to the paper's MUSDB numbers.
melband_karaoke is a reproducibility gap in the
current checkout. It was measured on 2026-07-02 but the current
MODEL_INDEX dropped it in commit 5f4d6b3
("phase out karaoke variant for deux") and
ckpt/melband_karaoke/ no longer exists locally. Its config
was recovered from a mirror on victoria:
dim=384, depth=6, num_stems=2 — i.e. the vocals
architecture plus a second mask estimator, which is exactly
what its 2× checkpoint size and 1.47× peak VRAM imply. It is used below
only as a controlled measurement of the marginal cost of a second
stem.
Architecture
The forward chain
MelBandRoformer.forward (melband.py:316) is
a fixed nine-stage chain. There is no convolutional U-Net, no
time-domain branch, and no vocoder — unlike HTDemucs,
everything happens in one STFT domain:
waveform (B, 2, T)
│
├─ torch.stft n_fft=2048 hop=441 win=2048 hann center=True
│ → complex (B, 2, F=1025, T')
├─ view_as_real, merge stereo into the freq axis
│ → (B, f·s = 2050, T', 2)
├─ gather by mel band: x = stft_repr[batch_arange, freq_indices]
│ → (B, 1979, T', 2) # 1979 = bins counted with multiplicity
├─ fold the complex axis into channels → (B, T', 2·1979 = 3958)
├─ BandSplit: per-band RMSNorm → Linear(dim_in, dim) → stack
│ → (B, T', 60, dim)
├─ × depth axial blocks:
│ time transformer over T' (60 independent sequences of length T')
│ freq transformer over 60 (T' independent sequences of length 60)
├─ MaskEstimator per stem: per-band MLP → GLU → concat → (B, stems, T', 3958)
├─ unfold complex → complex mask (B, stems, 2050, T')
├─ scatter_add_ back to the 2050-bin axis, then divide by num_bands_per_freq
├─ multiply the original complex STFT
└─ istft → waveform (B, stems, 2, T)
The paper frames the whole design as replacing BS-RoFormer's heuristic non-overlapping band split with a learned mel filter bank: "the Mel-band projection module can be seen as a learnable Mel filter-bank, since its MLP-layers serve as the mechanism to learn the filters" (arXiv:2310.01809, §2).
The mel-band split, exactly
num_bands=60 and dim_freqs_in=1025 are
CONFIRMED in every shipped config.json. The split is built
in melband.py:265-294 and is entirely deterministic, so it
can be reproduced without loading a checkpoint:
librosa.filters.mel(sr=44100, n_fft=2048, n_mels=60)produces a(60, 1025)triangular filter bank.- It is binarised
(
mel_filter_bank > 0), discarding the triangle weights and leaving a 0/1 incidence matrix. This matches the paper's description verbatim. - Two entries are forced to 1:
mel_filter_bank[0][0](DC → band 0) andmel_filter_bank[-1, -1](Nyquist → band 59). Without these the DC bin would belong to no band and theassert freqs_per_band.any(dim=0).all()guard atmelband.py:276would fail. freq_indicesis the per-band list of bin indices in band order, with stereo interleaved (freq_indices * 2 + arange(2)), so the tensor is[band 0 L, band 0 R, band 1 L, band 1 R, …].
DERIVED, and CONFIRMED against the checkpoint tensor shapes:
| quantity | value |
|---|---|
| STFT bins | 1025 |
| mel bands | 60 |
| band-membership entries (ones in the binarised bank) | 1979 |
| bins covered by exactly 1 band | 71 |
| bins covered by exactly 2 bands | 954 |
| bins covered by 0 bands | 0 |
| mean bands per bin | 1.9307 |
| smallest band (band 0, after DC forcing) | 7 bins |
| bands 1–15 | 6 bins |
| largest band (band 59) | 130 bins |
All 60 per-band bin counts, DERIVED and identical to the checkpoint's
band_split input widths:
band 0- 9: 7 6 6 6 6 6 6 6 6 6
band 10-19: 6 6 6 6 6 7 7 7 9 9
band 20-29: 9 10 10 11 13 13 13 15 16 17
band 30-39: 19 20 20 22 24 26 28 29 31 33
band 40-49: 36 39 41 44 47 50 54 57 61 66
band 50-59: 71 76 80 86 93 99 105 113 122 130
The overlap is structural, not an implementation artefact: "the width
of a Mel-band is two times the distance between its center and its
previous Mel-band's center. This makes the second half of a Mel-band
overlap its next Mel-band" — hence 71 bins covered once and 954 covered
twice, and hence the division by num_bands_per_freq after
the scatter-add. Because the bank is binarised, that "overlap &
average" is an exact arithmetic mean of two estimates over 93 % of the
spectrum. Note also the 86× spread in band width (6
bins to 130), which is the mel scale doing its job and is why
BandSplit cannot be a single shared projection.
BandSplit (melband.py:157)
is where the frequency axis becomes the channel axis. Its
dim_inputs are 2 · bins_in_band · 2 channels —
real and imaginary stacked as channels, times stereo — so band 0 is
2·7·2 = 28 wide and band 59 is 2·130·2 = 520
wide, summing to 2 · 1979 · 2 = 7916. Each band gets its
own RMSNorm(dim_in) followed by
Linear(dim_in, 384), and the 60 results are stacked on a
new axis. This is the only place the 60-band structure is
created, and it mixes nothing across bands.
One detail worth recording because it looks like a bug and is not:
RMSNorm here is not a normalisation in the usual sense.
melband.py:43 is
F.normalize(x, dim=-1) * sqrt(dim) * gamma — it divides by
the L2 norm and never subtracts a mean, so it is
scale-invariant but not mean-invariant, and it has no
bias. Every "norm" in this model is that same function.
The two transformers are not the same shape
melband.py:248-254 builds, per block, a time
transformer and a freq transformer, each an
independent Transformer with its own
RotaryEmbedding(dim_head=64). Which axis is attended is
chosen purely by the einops packing:
# time axis: (b t f d) -> (b f t d) -> pack to (* t d) => 60 sequences of length T'
# freq axis: (b f t d) -> (b t f d) -> pack to (* f d) => T' sequences of length 60
Both are called on the same 4-D tensor
(b, t, f, d); only the pack pattern differs. That yields
the single most important structural fact in this report:
| time-axis transformer | freq-axis transformer | |
|---|---|---|
| sequence length fed to attention | T' = 801 at the 8 s chunk | 60 |
| independent sequences per call | 60 (one per band) | T' = 801 (one per frame) |
| attention score matrix per sequence | 801 × 801 | 60 × 60 |
| score elements per layer-instance | 5,132,808 | 28,800 |
| RoPE applied over | frame index | band index |
| rotary dims | dim_head = 64 |
dim_head = 64 |
Per Attention layer (melband.py:109) the
parameters are
to_qkv: Linear(dim, 3·heads·dim_head, bias=False),
to_gates: Linear(dim, heads),
to_out: Linear(heads·dim_head, dim, bias=False). Note that
the gating is per-head
(out * sigmoid(gates)), that to_qkv has
no bias, and that
heads·dim_head = 8·64 = 512 ≠ dim = 384 for the vocals
model — so to_qkv genuinely expands and to_out
genuinely contracts. The model is not
head-dimension-square, which matters for the FLOP counting
below.
801 frames, verified on device. torch.stft with center=True returns 1 + 352800/441 = 801 frames, independent of n_fft, and an instrumented forward on victoria reported 801 tokens. Time-axis length 60 and frequency-axis length 801 are both confirmed rather than assumed.
This is an unusual operating point. In most axial transformers the
long axis has few parallel sequences; here the 801-long
attention runs 60 times in parallel while the 60-long attention runs 801
times in parallel, and both axes carry exactly 60 × 801 =
48,060 tokens. The work is not identical, because attention is
quadratic: 801² × 60 = 3.85e7 score entries per layer
against 60² × 801 = 2.88e6 — a 13.4× asymmetry from
the same token budget.
Parameter accounting
DERIVED from the constructed model, CONFIRMED against the safetensors
header to within the RMSNorm buffers:
| group | params | share |
|---|---|---|
mask_estimators |
201,465,304 | 88.28 % |
layers.*.time.* (6 blocks) |
11,831,088 | 5.18 % |
layers.*.freq.* (6 blocks) |
11,831,088 | 5.18 % |
band_split |
3,070,700 | 1.35 % |
output RMSNorm × 2 |
4,608 | 0.002 % |
total (melband_vocals) |
228,202,788 | 100 % |
The entire axial trunk — both axes, all six blocks — is 23.7 M parameters, 10.4 % of the model. 88 % of the weights live in the mask heads.
That is worth stating carefully because it inverts the usual
intuition. MaskEstimator (melband.py:187)
holds one to_freqs sequential per band, and
mask_estimator_depth=2 (CONFIRMED in all configs) makes
each a three-layer MLP with
dim_hidden = dim * 4 = 1536:
per band, per stem: Linear(dim, 1536) → Tanh → Linear(1536, 1536) → Tanh
→ Linear(1536, dim_in·2) → GLU
The output width is dim_in · 2 and GLU
halves it back to dim_in, which is what produces the odd
[56, 1536] and [1040, 1536] final-layer shapes
— CONFIRMED directly from the checkpoint, alongside
mask_estimators.0.to_freqs.0.0.0.weight = [1536, 384] and
...0.0.2.weight = [1536, 1536]. The 201 M total is that
stack replicated across 60 bands and every stem.
So the model is parameter-heavy in the mask head and
FLOP-heavy in the trunk — the inverse of the usual transformer
picture. For melband_deux the same decomposition gives
189,987,760 in the mask head (87.4 %) with a trunk of 12 blocks at
dim=256: more compute in fewer
parameters. Any optimisation aimed at parameter count attacks the cheap
part; anything aimed at FLOPs attacks the trunk.
Verified inference hot path
_apply_melband (inference.py:191) is the
path the prior benchmark exercised. Step by step:
- Dispatch.
apply_model(inference.py:43) type-checksisinstance(model, MelBandRoformer)and routes here. HTDemucs has its own separate function; the two do not share chunking logic. - Device and mode.
model.to(device),model.eval(). The model moves to CUDA in whatever dtype it was constructed in — fp32, see the precision section. - Chunk arithmetic (CONFIRMED,
inference.py:211-215).
Thechunk_size = int(model.sample_rate * 8) # 352800 samples = exactly 8.00 s overlap = 0.1 # module default stride = int(chunk_size * 0.9) # 317520 samples = 7.20 s fade_size = chunk_size // 10 # 35280 samples = 0.80 s border = chunk_size - stride # 35280 samples = 0.80 soverlapargument is used only to derivestride; it never sets a window length, and passingoverlap=Nonerestores 0.1. - Reflect padding at the borders. If
length > 2 · border, the whole mix is reflect-padded byborder(0.8 s) on each side andlengthis updated to the padded length. This is what makes the first fade-in and last fade-out unnecessary. - Linear fade window.
window = ones(chunk_size), with the first and lastfade_sizesamples replaced bylinspace(0,1)andlinspace(1,0). This is a linear fade, not HTDemucs's triangular transition weight — the Mel-Band path is deliberately simpler. Atoffset == 0the fade-in is forced back to 1, and on the final chunk the fade-out is forced back to 1. - Per-chunk loop. For
offset in range(0, length, stride): slice,_pad_chunkto exactlychunk_size, run the model undertorch.no_grad(), accumulateout += est * windowandcounter += window. _pad_chunk(inference.py:169) pads the tail withmode='reflect'only if the remaining audio exceeds half a chunk, otherwise with zeros. The final short chunk is therefore zero-padded, not reflected.- Overlap-add normalisation.
out = out / counter.clamp(min=1e-8). Because the fades are linear and complementary across the stride, this is a partition of unity away from the borders. - Un-pad. Crop
bordersamples off each end if padding was applied.
The chunk size is not the model's training length.
_inference.dim_t = 1101 is CONFIRMED in all three configs,
and 1101 frames at 100 frames/s is 11.01 s, not 8 s.
The implementation overrides it with a hard-coded 8 s. The architecture
is length-agnostic — RoPE and the transformers care only about sequence
length — so a different chunk is fully supported; the shipped path
simply never exercises one. Nothing in the local commit history explains
the choice.
The final chunk is not the same shape of computation as the
others. It is zero-padded to chunk_size and still
costs a full 801-frame forward pass. For a duration just over a multiple
of 7.2 s, up to 35 % of a chunk's cost is spent on mostly-empty
audio.
Why chunking exists: memory, not speed
The chunked path is a memory mechanism. A single unchunked forward at 16 s costs 3,718 ms against 2 × 1,586 = 3,172 ms for two 8 s chunks, so chunking happens to be 17 % faster here as well — but the reason it is mandatory is the VRAM curve, which is super-linear in sequence length:
| input | frames T' | DERIVED attention score tensor (60 × 8 × T'² × fp32) | MEASURED peak |
|---|---|---|---|
| 1 s | 101 | 40 MiB | 1,076 MiB |
| 4 s | 401 | 616 MiB | 1,383 MiB |
| 8 s (the shipped chunk) | 801 | 2.46 GiB | 1,791 MiB |
| 16 s | 1601 | 9.81 GiB | 2,609 MiB |
The 16 s score tensor exceeds victoria's 6 GiB on its own, and that run only completes because the memory-efficient attention kernel tiles the score matrix rather than materialising it all at once. Chunking is what bounds that allocation. This is the same conclusion the Stable Audio 3 report reached from the other direction ("chunking is a memory mechanism, not a speed one"); here it is both, because the attention is quadratic in a chunk whose length is fixed.
Bottleneck analysis — attention-bound or FFN-bound?
The Stable Audio 3 report found attention at only 12.1 % of FLOPs at that model's 120 s ceiling and told a 3050 owner not to spend effort on attention kernels. The mel-band structure changes the mechanism, and — as the measurements below show — it also changes the verdict. FLOP share and kernel time disagree on this model, so both are derived here.
FLOP shares
DERIVED for melband_vocals (dim=384,
d_inner=512, ff_inner=1536, depth=6)
and melband_deux (dim=256, d_inner=512,
ff_inner=1024, depth=12), per 8 s chunk at
T'=801, 60 bands. MACs, with FLOPs = 2 × MACs:
| term | vocals MACs | deux MACs | share of trunk |
|---|---|---|---|
time to_qkv |
170.08 G | — | 12.2 % |
| time FFN up + down | 340.16 G | — | 24.5 % |
freq to_qkv |
170.08 G | — | 12.2 % |
| freq FFN up + down | 340.16 G | — | 24.5 % |
to_out (both axes) |
113.38 G | — | 8.2 % |
| time attention | 236.52 G | — | 17.0 % |
| freq attention | 17.72 G | — | 1.3 % |
to_gates (both axes) |
1.78 G | — | 0.1 % |
| trunk | 1,389.9 G | 1,720.3 G | 100 % |
| mask head | 161.2 G | 152.0 G | 10.4 % of total |
band_split |
2.4 M | 2.4 M | 0.0002 % |
| total | 1,551.1 G | 1,872.5 G | 3.102 / 3.745 TFLOP |
Four shares are worth carrying forward, and they hold for both checkpoints because depth cancels in every ratio:
- Attention is 16.4 % of total FLOPs (18.3 % of the
trunk) for
vocals, and 29.6 % of the trunk fordeux— the narrower, deeper model is the more attention-dominated one despite having fewer parameters. to_qkvand the FFN GEMMs are 73 % of FLOPs — six equally weighted 170.08 G terms against a 1,389.9 G trunk. Within the trunk the linear layers beat attention 4.46 : 1.- Time versus frequency attention is 13.4 : 1
(236.52 G against 17.72 G), the quadratic asymmetry the sequence lengths
predict. At 60 tokens,
2·60·512 = 61,440MACs per token is ten times smaller than one 589,824-MAC FFN projection. Had the model attended over the raw 1025 bins instead of 60 mel bands that term would be ~292× larger and the model really would be quadratic-bound — the band split is what prevents it. - The mask head is 88 % of parameters but 10.4 % of FLOPs. Weight compression attacks the already-cheap part; anything that helps FLOPs must attack the trunk.
FLOP share is not kernel time, and on this model the two
disagree. MEASURED 2026-09-15 for vocals, one real
8 s chunk with CUDA events around each submodule (instrumented forward
1,605 ms against 1,586 ms uninstrumented, so read the ratios, not the
absolute values):
| module / op | ms | share of attributed time |
|---|---|---|
| time-axis transformer (6 calls) | 873.4 | 54.4 % |
| freq-axis transformer (6 calls) | 600.0 | 37.4 % |
mask_estimators |
102.4 | 6.4 % |
band_split |
6.5 | 0.4 % |
| — of which time attention | 310.3 | 27.4 % |
| — of which all 24 RMSNorm | 101.4 | 9.0 % |
| — of which freq attention | 36.3 | 3.2 % |
Attention takes roughly 30 % of wall clock against 16.4 % of
FLOPs — about 1.8× over-represented — while
the linear layers take ~51 % against 73 %, under-represented by 0.7×.
The FLOP model is a good map of where the arithmetic is and a poor
predictor of where the time goes. The
melband_deux trace below resolves this properly: attention
is 30.1 % of device time and is the model's largest
single cost. The 24 RMSNorm calls move 148 MiB per call at about
35 GB/s achieved — a quarter of the card's bandwidth,
and the only achieved-bandwidth figure in this article.
Where the efficiency actually goes
DERIVED from the same measurements:
| quantity | value |
|---|---|
FLOPs per 8 s chunk (melband_vocals) |
3.102 TFLOP |
| MEASURED forward, 8 s chunk | 1.586 s |
| achieved throughput | 1.96 TFLOP/s |
victoria FP32 vector peak at 1.49 GHz
(2560 × 2 × 1.49e9) |
7.63 TFLOP/s |
| fraction of FP32 peak | ≈ 26 % |
MEASURED streaming floor for the fp32 weights
(871 MiB ÷ 144 GB/s) |
6.3 ms |
| weight-streaming floor as a fraction of chunk time | 0.4 % |
Three-quarters of the card's FP32 capability is idle, and the DRAM system is idle to the point of irrelevance — the weights could be streamed 252 times over in the time one chunk takes. The structural cause the earlier analysis proposed is worth stating, because the measurement answers it directly:
The band axis is split across the batch dimension.
pack([x], '* t d') produces a batch of 60 independent
sequences, so each nn.Linear in the time transformer was
expected to run as 60 GEMMs of 801 × 384
rather than one GEMM of 48,060 × 384 — roughly 1,800
time-axis GEMMs per chunk — and the predicted fix was to flatten the
band axis into the token axis. The trace falsifies both
halves. There are 540 GEMM launches per chunk,
not 1,800 or 7,200, because matmul already collapses the
packed batch into one large GEMM; and the flatten is a measured
0.7–3.5 % regression. See
Device time by kernel
class.
The chunking arithmetic
_apply_melband pays a fixed per-chunk
cost and processes a duration-linear number of
chunks, so total latency should be a function of chunk count and not of
duration. DERIVED from chunk_size = 8 s,
stride = 7.2 s, border = 0.8 s:
n_chunks(d) = len(range(0, L, 317520)) where L = 44100·d, plus 2·35280 of reflect
padding when 44100·d > 70560
| duration | DERIVED chunks | model-audio processed | redundancy vs duration |
|---|---|---|---|
| 5 s | 1 | 8.0 s | 1.600× |
| 30 s | 5 | 40.0 s | 1.333× |
| 60 s | 9 | 72.0 s | 1.200× |
| 120 s | 17 | 136.0 s | 1.133× |
| 600 s | 84 | 672.0 s | 1.120× |
The redundancy converges to 8/7.2 = 1.1111×, so the overlap tax is bounded and modest. Reflect padding does not change the chunk count for any duration above 1.6 s, because the padded length lands in the same stride bin.
MEASURED 2026-09-15: the cost model, fitted
Seven durations, melband_vocals, 44.1 kHz stereo
random-noise input, _apply_melband with defaults, one
warmup excluded, one measured run each:
| duration | chunks | latency | RTF | peak VRAM | ms per chunk |
|---|---|---|---|---|---|
| 5 s | 1 | 1.61 s | 0.3213× | 1.67 GiB | 1,607 |
| 15 s | 3 | 4.79 s | 0.3192× | 1.68 GiB | 1,596 |
| 30 s | 5 | 7.96 s | 0.2652× | 1.70 GiB | 1,591 |
| 45 s | 7 | 11.10 s | 0.2467× | 1.72 GiB | 1,586 |
| 60 s | 9 | 14.31 s | 0.2385× | 1.74 GiB | 1,590 |
| 90 s | 13 | 20.67 s | 0.2296× | 1.78 GiB | 1,590 |
| 120 s | 17 | 27.07 s | 0.2256× | 1.82 GiB | 1,592 |
Least-squares fits over those points:
latency = 1 ms + 1591 ms × n_chunks residual RMS 19 ms
latency = 1186 ms + 217 ms × duration_sec residual RMS 314 ms
The chunk-count model is 16× better than the duration model, and its intercept is 1 ms — statistically zero. Every duration from 5 s to 120 s costs 1,586–1,607 ms per chunk, a spread of 1.3 % across a 17× range in chunk count. Three consequences:
- Cost is exactly linear in duration above the first chunk, and it is linear because it is linear in chunks, not in seconds of audio. The chunk redundancy is a real tax on the naive per-second cost, which is why RTF improves monotonically with duration (0.3213 → 0.2256): a longer request amortises the overlap rather than a fixed overhead.
melband_vocalshas no measurable per-call fixed overhead. This is a genuine difference from the Stable Audio 3 finding of a ~1.92 s constant term; here the constant is 1 ms. The reason is structural:_apply_melbandhas no per-step Python loop, no schedule indexing and notqdm, and model load is excluded from the latency column. The only fixed cost is the transposed STFT/iSTFT and the overlap-add buffers, and it is genuinely negligible.- Peak VRAM grows with chunk count at ~9 MiB per
chunk (1.67 → 1.82 GiB over 16 additional chunks) because
outandcounterare preallocated at full padded length. That slope is DERIVED and easily checked:2 tensors × 1 stem × 2 channels × 4 B × 317,520 samples = 4.84 MiBper chunk, so the measured 9 MiB is within 2× (the remainder is allocator behaviour and fragmentation). It is linear, small, and the reason a 600 s request still fits.
Cross-check against the prior benchmark
The 2026-07-02 run in doc/benchreport.md used the same 60 s
of bad_guy.mp3. Its melband_deux row is
reproduced to four figures in the measured section below (20.444 s
against 20.44 s), so only the comparison across architectures is carried
here:
| model | stems | latency | RTF | peak VRAM | ms/chunk |
|---|---|---|---|---|---|
htdemucs |
4 | 2.71 s | 0.0452x | 1.30 GiB | 301 |
melband_vocals |
1 | 15.66 s | 0.2610x | 1.77 GiB | 1,740 |
melband_deux |
2 | 20.44 s | 0.3407x | 1.75 GiB | 2,271 |
melband_karaoke |
2 | 16.65 s | 0.2775x | 2.60 GiB | 1,850 |
Two things this table is good for. Mel-band is roughly 6x
slower than HTDemucs at comparable reconstruction quality. And
melband_karaoke is the only controlled measurement of the
marginal cost of a second stem: same
dim=384/depth=6 trunk as
melband_vocals, one extra mask estimator, and it costs
+6.3 % latency and +0.8 GiB VRAM against a predicted
+10.4 % of FLOPs. That row cannot be re-measured from this checkout —
commit 5f4d6b3 removed the variant from
MODEL_INDEX — so the number stands as recorded.
Hardware analysis for victoria
The card
NVIDIA GeForce RTX 3050 6 GB Laptop (GA107, sm_86), all
measured on victoria except the silicon data: 6144 MiB VRAM,
5804 MiB usable at idle, 2560 CUDA cores at 1.24–1.49 GHz,
96-bit bus at 12 Gbps = 144 GB/s, 2 MB L2, 7 GB host
RAM. DERIVED FP32 vector peak 2560 × 2 × 1.49e9 =
7.63 TFLOP/s at maximum boost, 6.35 TFLOP/s at base.
The 2 MB L2 against a 913 MB fp32 weight set means zero weight
reuse between forward passes — every chunk re-streams every
weight.
The checkpoint is fp16, the model is fp32
The on-disk model.safetensors is 456.5 MB and
entirely F16 (CONFIRMED from the header:
{'F16': 228203172}). But _load_melband
(separator.py:244) constructs the model in default dtype
and never calls .half(), and
Separator._lazy_load (separator.py:315)
applies its precision cast only on the htdemucs branch:
if self.device.type == "cuda" and self.info["type"] == "htdemucs" and self.precision == 'bf16':
model = load_model(self.name, torch.device("cpu"))
model = model.to(dtype=torch.bfloat16).to(self.device)
else:
model = load_model(self.name, self.device) # mel-band always lands here
So the mel-band path upcasts every weight to fp32 on
load. MEASURED 2026-09-15 with a live model:
dtype = torch.float32 and
weight bytes on GPU 871 MiB, which is 228,202,788 × 4 B =
913 MB = 871 MiB, against the file's 456.5 MB. That is
a 1.91× expansion, and it means the
precision argument is silently ignored for mel-band
models — a trap worth knowing about independently of
performance.
Compute-bound or memory-bound?
Neither, and the distinction matters. VRAM is not the
constraint: peak was 1.74-1.82 GiB measured across 5-120 s
against 5804 MiB usable, about 31 % of the card, of
which 871 MiB is the fp32 weight set. A 4 GB part would run it.
Nor is DRAM: one full stream of the fp32 weights is
871 MiB / 144 GB/s = 6.3 ms against a 1,586 ms forward
pass, so the weights could be streamed 252 times over
in one chunk; a bandwidth-bound workload would sit at 80-95 % of
144 GB/s and the weight stream reaches 0.4 % of it.
Nor is it compute-bound in any useful sense:
3.102 TFLOP per chunk against 1,586 ms is 1.96 TFLOP/s, ~26 %
of the FP32 peak.
The honest characterisation is kernel-efficiency-bound. Three mechanisms, in order of measured size:
- fp32 throughout, which doubles weight and activation traffic and precludes any tensor-core path. This is the one that pays: the fp32 attention kernel is 30.1 % of device time and the same shape is 7.7x faster in fp16.
- The rotary embedding costs 546.8 ms per chunk -
27.5 % of the per-op total - rebuilding a cos/sin table and
materialising a
torch.stackgather twice per attention. - 24 memory-bound
RMSNormcalls move 148 MiB each at about 35 GB/s achieved, a quarter of the card's bandwidth. Together they are 29.2 ms per chunk, 1.3 % of device time - real, but small.
The framing for a follow-up pass: victoria has ~4 GiB of spare VRAM and ~5.7 TFLOP/s it cannot reach, and both follow from the model running fp32 rather than from the hardware.
The flash-attn caveat
The suspicion that the benchmarks ran without flash attention is
correct, but the mechanism is not the broken package and the consequence
is not what it appears. flash_attn=True means
"delegate to F.scaled_dot_product_attention", not
"the flash_attn package is installed"; models/melband.py
never imports that package, so its broken
.so (undefined symbol under torch 2.13) is
irrelevant here. MEASURED 2026-09-15, forced-backend probe at the real
attention shape (b=60, h=8, 801x64, fp32):
forced SDPBackend |
result |
|---|---|
FLASH_ATTENTION |
RuntimeError: No available kernel. Aborting execution. |
EFFICIENT_ATTENTION |
OK |
MATH |
OK |
Every flash-attention kernel in PyTorch requires fp16 or
bf16, and this model is fp32. Note that
torch.backends.cuda.flash_sdp_enabled() still returns
True — the flag is enabled, the kernel refuses the dtype.
"SDPA flash is enabled" and "SDPA flash is being used" are different
claims.
Choosing between the two backends that do run costs nothing.
MEASURED 2026-09-15, full 60 s pipeline, three interleaved repeats: the
default and forced EFFICIENT_ATTENTION land at
14.26 s vs 14.27 s (RTF 0.2377 vs 0.2378), a
0.0 % delta. An earlier unrepeated chunk-level
measurement in the same pass suggested a 20 % advantage for
EFFICIENT_ATTENTION; repeating it showed that to be
run-to-run variance, not a treatment effect.
An earlier draft of this section said those timings ran on the
SDPBackend.MATH path. That was wrong, and the correction is
the finding. Profiling the model's own attention shape shows the
default selection is the memory-efficient cutlass kernel,
fmha_cutlassF_f32_aligned_64x64_rf_sm80; forcing
EFFICIENT_ATTENTION produces the identical kernel, while
MATH produces an entirely different one
(softmax_warp_forward plus 128x128 GEMMs) and
is 1.75x slower at this shape. So the benchmarks were already on the
faster of the two fp32 backends, and the fallback that matters is
fp32 itself: the memory-efficient kernel is the best choice
available at fp32 and still takes 30.1 % of device time
in melband_deux, because it cannot use tensor cores. The
same shape is 7.7x faster in fp16 — see
The two levers, measured. The
earlier verdict that "attention is not the problem" was a statement about
FLOPs and it does not survive contact with the kernel that implements
them.
melband_deux, measured on 60 s of real music
Everything above is derived, or measured on melband_vocals.
This section is melband_deux — dim=256,
depth=12, two stems — measured end to end on victoria with the
first 60 s of bad_guy.mp3 as the input instead of noise. Harness
demucs_golf/scripts/prof_melband.py, torch
2.12.0+cu130, driver 580.119.02, one warm-up run before every
reported number. The sector tree comes from a verbatim copy of
MelBandRoformer.forward with timing ranges spliced in; that copy is
asserted bit-identical to the shipped forward (max abs diff
0.0) before anything is measured.
The July row reproduces exactly
| run | wall s | RTF | × real time | peak VRAM MiB |
|---|---|---|---|---|
| 1 | 20.372 | 0.33953 | 2.95 | 1710 |
| 2 | 20.408 | 0.34013 | 2.94 | 1710 |
| 3 | 20.444 | 0.34073 | 2.93 | 1710 |
| 4 | 20.491 | 0.34152 | 2.93 | 1710 |
| 5 | 20.486 | 0.34143 | 2.93 | 1710 |
Median 20.444 s, RTF 0.34073, spread 0.58 % over five runs, peak VRAM flat at 1710 MiB, and an identical output checksum on all five. The 2026-07-02 row was 20.44 s / 0.3407 / 1.75 GiB. That settles the 9.4 % gap: it is input-dependent, not checkout-dependent. On the same file the two agree to four figures, and the earlier fresh pass had used random noise.
The chunk loop is device-bound and flat
Nine chunks, 2271.6–2279.5 ms each — a 0.35 % spread — with
0.09 ms of host overhead per chunk between the device-time and
wall-time columns. The cost model is 2276 ms × n_chunks with no
intercept worth fitting: the fade windows, reflect padding and overlap-add
normalisation are free, and the final partially-filled chunk costs the same as a
full one.
Where a chunk goes
Per 8 s chunk, CUDA event pairs recorded without synchronising and resolved once at the end, so the numbers are stream time rather than instrumentation artefact. The instrumented forward is 2315.7 ms against 2276 ms uninstrumented, a 1.7 % inflation that the shares absorb.
| sector | calls | ms | share |
|---|---|---|---|
trunk | 1 | 2184.1 | 94.34 % |
mask_head | 2 | 117.3 | 5.07 % |
mask_apply | 1 | 5.10 | 0.22 % |
band_split | 1 | 4.55 | 0.20 % |
istft | 1 | 2.61 | 0.11 % |
gather | 1 | 1.09 | 0.05 % |
stft | 1 | 0.43 | 0.02 % |
The trunk is 94 % of a chunk. Everything outside it — STFT, mel gather, band split, complex-mask scatter, iSTFT — totals 132 ms, and only the mask head is worth naming. The 12 blocks are flat to ±0.1 %: 180.24–180.64 ms each, of which the time axis is 114.4–114.9 ms (63.6 %) and the frequency axis 65.6–65.9 ms (36.4 %). There is no first-block penalty and no hot block once warm.
Device time by kernel class
25,211 device slices over the nine chunks of a full 60 s separation. Total device time is 2,276 ms per chunk against a 2,276 ms wall chunk, so the GPU is fully occupied and nothing is waiting on the host.
| class | calls/chunk | ms/chunk | share |
|---|---|---|---|
attention (fmha_cutlassF_f32) | 24 | 685.6 | 30.13 % |
GEMM (ampere_sgemm) | 540 | 612.6 | 26.92 % |
| elementwise mul | 439 | 314.9 | 13.84 % |
| copy / cat | 179 | 239.6 | 10.53 % |
| elementwise add | 218 | 178.0 | 7.82 % |
| elementwise gelu | 24 | 76.8 | 3.37 % |
| elementwise div | 134 | 58.3 | 2.56 % |
| elementwise neg (RoPE) | 48 | 58.1 | 2.55 % |
reduce NormTwoOps (RMSNorm) | 132 | 29.2 | 1.28 % |
Attention is the largest single class, from 24 launches. The
kernel is fmha_cutlassF_f32_aligned_64x64_rf_sm80, the fp32
memory-efficient SDPA backend — which is what flash_attn=True
selects when the tensors are fp32. Each time-axis call averages
51 ms; the frequency-axis ones about 6 ms.
The "60 small GEMMs" mechanism is wrong at the launch level.
12 blocks × 2 axes × 5 projections × 60 bands is 7,200 GEMMs per chunk in
principle; the trace shows 540. PyTorch's matmul on
the packed (60, 801, 256) tensor is already one large GEMM, and it
sustains 3,827–3,942 GFLOP/s, about half of victoria's fp32
peak. Measured head to head, the shipped packed form beats the proposed
(f t) flatten by 0.7–3.5 % on all four projection
shapes, and a hand-written per-band loop is 1.09–1.84× slower. The flatten
patch — the previous highest-priority item in this report — is a
no-op that costs a little.
Per op, exact
Device time attributed from the trace, so a host stall cannot leak into a
sector. rotary and sdpa together are
54 % of a chunk.
| op | time axis ms | freq axis ms | total ms | share |
|---|---|---|---|---|
sdpa | 610.1 | 70.3 | 680.4 | 34.2 % |
rotary | 275.6 | 271.3 | 546.8 | 27.5 % |
ff_up | 99.4 | 100.0 | 199.5 | 10.0 % |
to_qkv | 96.0 | 96.3 | 192.3 | 9.7 % |
norm | 67.1 | 67.1 | 134.2 | 6.8 % |
ff_down | 60.2 | 59.9 | 120.1 | 6.0 % |
to_out | 29.8 | 30.7 | 60.6 | 3.0 % |
gate_scale | 19.5 | 19.5 | 39.0 | 2.0 % |
to_gates | 7.3 | 7.3 | 14.7 | 0.7 % |
rotary is the finding this report did not previously have.
RotaryEmbedding.rotate_queries_or_keys is called twice per
attention (once for q, once for k) in each of 24 blocks, and each call
rebuilds a cos/sin table with an einsum, a repeat and
two trig calls, then materialises a torch.stack gather over two
strided views inside rotate_half, then runs two full multiplies and
an add over the whole tensor. At 546.8 ms per chunk it is the
second-largest cost in the model and it is pure memory
traffic.
The two levers, measured
Half precision. The attention shape is the reason. Benchmarked directly at the shapes the model issues, 30 repetitions:
| shape | dtype | ms/call | GFLOP/s |
|---|---|---|---|
(60, 8, 801, 64) time axis | fp32 default | 51.32 | 1,536 |
| fp16 / bf16 (flash) | 6.65 | 11,852 | |
| fp32 math | 89.63 | 880 | |
| fp16 math | 109.31 | 721 | |
(801, 8, 60, 64) freq axis | fp32 default | 5.92 | 997 |
| fp16 / bf16 (flash) | 2.30 | 2,572 |
7.7× on the time axis, 2.6× on the frequency axis, and the
fp32 column reproduces the 51 ms the trace measured per call — an independent
cross-check of both. bf16 is not an option for this model:
torch.autocast(bf16) fails at torch.view_as_complex,
which rejects bfloat16 when the complex mask is assembled. fp16 is
the viable half precision, and it is the dtype the checkpoint already ships
in.
A fused rotary. Replacing the stack gather with
a flip view plus a sign multiply is arithmetically exact; compiling
it fuses the three elementwise passes into one. Per call, on the time-axis
shape:
| variant | ms/call | speedup | max abs diff |
|---|---|---|---|
| shipped | 11.49 | 1.00× | — |
| cached cos/sin, same arithmetic | 11.45 | 1.00× | 0 |
cached + flip rotate_half | 8.98 | 1.28× | 0 |
cached + flip + torch.compile | 2.35 | 4.88× | 4.8e-07 |
Caching the tables buys nothing, which localises the cost precisely: it is the elementwise passes over a 94 MiB tensor, not the trig setup.
Together, on the same 60 s clip, three runs each:
| configuration | wall s | RTF | × real time | speedup | vs fp32 | peak VRAM MiB |
|---|---|---|---|---|---|---|
| fp32, as shipped | 20.607 | 0.34345 | 2.91 | 1.00× | — | 1773 |
| fp16 autocast | 8.735 | 0.14558 | 6.87 | 2.36× | 56.0 dB | 1902 |
| fast rotary, eager | 19.672 | 0.32786 | 3.05 | 1.05× | 149.9 dB | 1904 |
| fast rotary, compiled | 16.833 | 0.28055 | 3.56 | 1.22× | 132.7 dB | 1814 |
| fp16 + fast rotary, eager | 8.229 | 0.13715 | 7.29 | 2.50× | 70.5 dB | 1949 |
| fp16 + fast rotary, compiled | 6.685 | 0.11141 | 8.98 | 3.08× | 70.8 dB | 1903 |
The vs fp32 column is SI-SDR between the variant's stems and the
fp32 stems; 70.8 dB is a numerical difference, not an audible one. The best
measured configuration separates 60 s of stereo in 6.69 s, RTF
0.111, on a 6 GB laptop GPU at under 1.9 GiB. Neither lever
requires a code change to the model's mathematics: one is
torch.autocast, the other replaces a gather with a view.
Two things this does not claim. The two levers overlap — both spend the same budget — so their speedups are not additive, and the measured 3.08× is the honest combined figure. And the reconstruction quality of the shipped path is unchanged: fp32 stems reconstruct the mixture at 38.221 dB SI-SDR, against 38.22 dB in the July table, with the fp16 stems at 38.218 dB.
Traces
Perfetto renders of a warm separation on victoria, captured with
torch.profiler and NVTX-labelled per sector. They are the source of
every per-kernel number above. The raw capture is 73 MB and is not published;
the windows below are pre-cropped traces rebased to zero, because Perfetto's
deep-link zoom does not take effect in headless Chromium.
Three consecutive chunks of the nine, GPU stream only. The trunk occupies almost all of each chunk; the mask head and iSTFT are the thin bands at the chunk boundaries.

One chunk, end to end
All tracks. stream 7 carries the kernels; the trunk
range is followed by mask_head and istft, and the
twelve blockNN.time / blockNN.freq pairs are visible
inside it, each fronted by an fmha attention kernel.

The axial trunk
2.08 s of trunk. The twelve blocks alternate time and frequency transformers and are indistinguishable in width — the flatness the hook table reports, visible directly.

One block, block00, at 110 ms: a time transformer and then a
frequency transformer.

The mask head — 117 ms per chunk, the only sector outside the trunk worth naming.

Inside a block
700 µs at the end of the largest kernel in a time block — the 51 ms
fmha_cutlassF_f32 attention call — running into the GEMM that
follows it.

760 µs framing a single ampere_sgemm_128x64_tn. The GEMMs in
this model are large, not numerous: 540 launches per chunk at a 1.6 ms mean,
which is why flattening the band axis buys nothing.

760 µs of the memory-bound band. An elementwise mul runs 590 µs
here; elementwise, copy and add kernels together are a third of all device
time.

The band split, 4.6 ms per chunk, at 760 µs magnification.

What to do next
Rewritten against the measurements above. The previous list was ordered by inference; this one is ordered by measured payback.
- Fix the rotary embedding. 546.8 ms per chunk — 27.5 % of
the per-op total — spent rebuilding a cos/sin table and materialising a
torch.stackgather, twice per attention in 24 blocks. Replacing the gather with aflipview and a sign multiply is bit-exact (max abs diff 0.0), and compiling it is 4.88× on the op and 1.22× end to end. This is the cheapest change in the list and it needs no dtype decision. - Wire up
fp16. 2.36× end to end at 56 dB SI-SDR against the fp32 render.Separatoralready takes aprecisionargument but honours it only on the htdemucs branch, so for mel-band it is inert — that is a small fix, not an experiment.bf16is not an option:torch.view_as_complexrejects it. - Take both. 20.607 s → 6.685 s, RTF 0.343 → 0.111, at 70.8 dB against the fp32 stems. The levers overlap, so 3.08× is the combined figure, not the product of the two.
- Drop the band-axis flatten. It was the highest-priority
item here and it is a measured regression: 0.7–3.5 % slower than the shipped
packed form, because
matmulalready collapses the band batch into one GEMM running at half of fp32 peak. - Deprioritise the RMSNorm fusion. It is 29.2 ms per chunk, 1.28 % of device time — not the 9.0 % the instrumented vocals run suggested. It is still correct and still free, but it is no longer worth scheduling.
- Leave the attention backend alone at fp32. The memory-efficient kernel is already the right choice: forcing the math backend is 1.75× slower at both shapes. The win available in attention is dtype, not kernel selection.
- Still worth doing: the chunk-size sweep, which remains the only untested lever with a plausible double-digit return — per-chunk cost is flat at 2,276 ms, so a longer chunk amortises the per-chunk fixed cost while raising the quadratic attention term, and the optimum is unmeasured.
Gaps and open questions
Still open:
- The chunk-size sweep is untested. Training used
dim_t = 1101(11.01 s) and inference uses 8 s. Per-chunk cost is now known to be flat with a negligible intercept, so the trade is purely a longer chunk against the quadratic attention term — but the optimum has not been measured. - The 144 GB/s bandwidth figure is still a specification. Every streaming floor here rests on it, including the elementwise shares.
- The sustained clock during these runs is unmeasured, so utilisation percentages use the 1.49 GHz maximum boost.
- The per-op table comes from an 8 s (two-chunk) capture, not the 60 s one, because the op-level ranges roughly quadruple the trace size. The per-chunk totals agree with the 60 s class table to within 1 % (attention 680.4 against 685.6 ms), so the shares are sound; treat the absolute per-op milliseconds as good to that margin.
- One band-split detail was never confirmed from a primary
source. This implementation uses binarised
librosa.filters.mel, and that reproduces the checkpoint's tensor shapes exactly — so the shapes are certainly right. Whether the mel band edges match whatbecruilytrained with is not confirmed. This would affect separation quality, not any performance figure here. melband_karaokeis present again but still unmeasured in this tree. Its checkpoint is inckpt/; commit5f4d6b3removed it fromMODEL_INDEX. Restoring the row is a config decision, not an experiment.