lab

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:

  1. librosa.filters.mel(sr=44100, n_fft=2048, n_mels=60) produces a (60, 1025) triangular filter bank.
  2. 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.
  3. Two entries are forced to 1: mel_filter_bank[0][0] (DC → band 0) and mel_filter_bank[-1, -1] (Nyquist → band 59). Without these the DC bin would belong to no band and the assert freqs_per_band.any(dim=0).all() guard at melband.py:276 would fail.
  4. freq_indices is 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:

  1. Dispatch. apply_model (inference.py:43) type-checks isinstance(model, MelBandRoformer) and routes here. HTDemucs has its own separate function; the two do not share chunking logic.
  2. 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.
  3. Chunk arithmetic (CONFIRMED, inference.py:211-215).
    chunk_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 s
    The overlap argument is used only to derive stride; it never sets a window length, and passing overlap=None restores 0.1.
  4. Reflect padding at the borders. If length > 2 · border, the whole mix is reflect-padded by border (0.8 s) on each side and length is updated to the padded length. This is what makes the first fade-in and last fade-out unnecessary.
  5. Linear fade window. window = ones(chunk_size), with the first and last fade_size samples replaced by linspace(0,1) and linspace(1,0). This is a linear fade, not HTDemucs's triangular transition weight — the Mel-Band path is deliberately simpler. At offset == 0 the fade-in is forced back to 1, and on the final chunk the fade-out is forced back to 1.
  6. Per-chunk loop. For offset in range(0, length, stride): slice, _pad_chunk to exactly chunk_size, run the model under torch.no_grad(), accumulate out += est * window and counter += window.
  7. _pad_chunk (inference.py:169) pads the tail with mode='reflect' only if the remaining audio exceeds half a chunk, otherwise with zeros. The final short chunk is therefore zero-padded, not reflected.
  8. 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.
  9. Un-pad. Crop border samples 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:

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:

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:

  1. 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.
  2. 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.stack gather twice per attention.
  3. 24 memory-bound RMSNorm calls 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

runwall sRTF× real timepeak VRAM MiB
120.3720.339532.951710
220.4080.340132.941710
320.4440.340732.931710
420.4910.341522.931710
520.4860.341432.931710

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.

sectorcallsmsshare
trunk12184.194.34 %
mask_head2117.35.07 %
mask_apply15.100.22 %
band_split14.550.20 %
istft12.610.11 %
gather11.090.05 %
stft10.430.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.

classcalls/chunkms/chunkshare
attention (fmha_cutlassF_f32)24685.630.13 %
GEMM (ampere_sgemm)540612.626.92 %
elementwise mul439314.913.84 %
copy / cat179239.610.53 %
elementwise add218178.07.82 %
elementwise gelu2476.83.37 %
elementwise div13458.32.56 %
elementwise neg (RoPE)4858.12.55 %
reduce NormTwoOps (RMSNorm)13229.21.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.

optime axis msfreq axis mstotal msshare
sdpa610.170.3680.434.2 %
rotary275.6271.3546.827.5 %
ff_up99.4100.0199.510.0 %
to_qkv96.096.3192.39.7 %
norm67.167.1134.26.8 %
ff_down60.259.9120.16.0 %
to_out29.830.760.63.0 %
gate_scale19.519.539.02.0 %
to_gates7.37.314.70.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:

shapedtypems/callGFLOP/s
(60, 8, 801, 64) time axisfp32 default51.321,536
fp16 / bf16 (flash)6.6511,852
fp32 math89.63880
fp16 math109.31721
(801, 8, 60, 64) freq axisfp32 default5.92997
fp16 / bf16 (flash)2.302,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:

variantms/callspeedupmax abs diff
shipped11.491.00×—
cached cos/sin, same arithmetic11.451.00×0
cached + flip rotate_half8.981.28×0
cached + flip + torch.compile2.354.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:

configurationwall sRTF× real timespeedupvs fp32peak VRAM MiB
fp32, as shipped20.6070.343452.911.00×—1773
fp16 autocast8.7350.145586.872.36×56.0 dB1902
fast rotary, eager19.6720.327863.051.05×149.9 dB1904
fast rotary, compiled16.8330.280553.561.22×132.7 dB1814
fp16 + fast rotary, eager8.2290.137157.292.50×70.5 dB1949
fp16 + fast rotary, compiled6.6850.111418.983.08×70.8 dB1903

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.

Perfetto timeline of three consecutive chunks on the GPU stream, each dominated by the axial trunk

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.

One 8 second chunk from STFT to iSTFT, showing the trunk dominating and the twelve block pairs inside it

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.

The axial trunk showing twelve alternating time and frequency transformer blocks of equal width

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

A single axial block: one time transformer followed by one frequency transformer

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

The mask head sector: per-band MLPs for both stems

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.

700 microseconds at the end of the 51 millisecond fp32 attention kernel, running into the following GEMM

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 microseconds framing a single large sgemm kernel

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.

760 microseconds of the memory bound elementwise band, showing a 590 microsecond multiply

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

760 microseconds inside the band split sector

What to do next

Rewritten against the measurements above. The previous list was ordered by inference; this one is ordered by measured payback.

  1. 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.stack gather, twice per attention in 24 blocks. Replacing the gather with a flip view 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.
  2. Wire up fp16. 2.36× end to end at 56 dB SI-SDR against the fp32 render. Separator already takes a precision argument but honours it only on the htdemucs branch, so for mel-band it is inert — that is a small fix, not an experiment. bf16 is not an option: torch.view_as_complex rejects it.
  3. 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.
  4. 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 matmul already collapses the band batch into one GEMM running at half of fp32 peak.
  5. 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.
  6. 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.
  7. 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: