NSNet2 under semi-structured sparsity
Six NSNet2 speech-enhancement checkpoints trained under fixed semi-structured sparsity masks — 2:4, 4:8, 1:4, 80% 1×4 blocks, 80% unstructured, and a dense control — for work on sparse-dense MatMul packing and code generation at batch 1. Each ships as a PyTorch checkpoint, an FP32 and a static-int8 ONNX graph, and a numpy export with explicit zeros and masks. No pattern here costs measurable quality, in FP32 or int8.
New: new_design/ adds new-design NSNet2 and ConvFSENet models
(dense and 2:4@c1, widths 64-384) and a finding that changes how to read this page: on a
healthy model, 2:4 buys kernel speed at equal nonzeros, not quality. See
New design below.
Code, training recipe and export tooling:
LarocheC/eco8-neaixt, branch
sparse-masks-rowfusion. See SPARSE_MATMUL_COLLAB.md there for the full
method.
Results
PESQ on the full 824-utterance VoiceBank-DEMAND test set. Every arm was fine-tuned from the same dense baseline on an identical schedule (lr 3e-4, 120 epochs), so the mask is the only variable.
| directory | pattern | sparsity | PESQ |
|---|---|---|---|
dense |
dense (control) | 0% | 2.777 |
2_4 |
2:4 | 50.0% | 2.779 |
4_8 |
4:8 | 50.0% | 2.779 |
1_4 |
1:4 | 75.0% | 2.781 |
unstructured_80 |
unstructured | 80.0% | 2.776 |
1x4_80 |
1×4 blocks | 80.0% | 2.770 |
int8
Static int8 PTQ (QDQ, per-channel symmetric weights, MinMax calibration on 200 utterances), PESQ through onnxruntime on the same test split. Δ is int8 − FP32.
| directory | sparsity | FP32 | int8 | Δ | int8 RTF |
|---|---|---|---|---|---|
dense |
0% | 2.777 | 2.783 | +0.006 | 0.121 |
2_4 |
50% | 2.779 | 2.781 | +0.002 | 0.125 |
4_8 |
50% | 2.779 | 2.790 | +0.011 | 0.122 |
1_4 |
75% | 2.781 | 2.784 | +0.003 | 0.123 |
1x4_80 |
80% | 2.770 | 2.779 | +0.009 | 0.124 |
unstructured_80 |
80% | 2.776 | 2.774 | −0.002 | 0.121 |
Sparsity does not make quantization harder — every Δ is inside the ±0.01 noise band at every sparsity level, and five of six are positive.
The mask survives int8 bit-exactly. Symmetric per-channel weight
quantization maps 0.0 to exactly 0. The N:M arms conform in the int8 graph with
sparsity slightly above target (0.5016 / 0.5011 / 0.7510 — a few small weights
round to zero, which N:M permits), and 1x4_80 holds block support at exactly
0.2000 live against its 0.2000 budget. Check it yourself with
nsnet2/verify_int8_sparsity.py from the repo.
But the sparsity buys no speed today. int8 RTF is 0.121–0.125 across every arm, dense and 80%-sparse alike, and the int8 file is 2.78 MiB regardless — onnxruntime stores the zeros explicitly and multiplies by them like any other weight. 80% of the multiplies are gone mathematically and none of the latency is. Closing that gap is what these checkpoints are for.
None of these patterns costs measurable quality. The spread across all six arms is 0.012 PESQ while the run-to-run variation within a single arm is ~0.010 sd, so they are statistically indistinguishable. Do not read an ordering into the table — 1:4 topping it at 75% sparsity is which validation happened to land last, not a result.
One caveat. Every arm including the dense control sits ~0.07 below the published 200-epoch baseline of 2.845, because these were shortened fine-tunes with a freshly initialised discriminator; the comparison between arms is unaffected since all paid the same penalty, and all six curves were still rising at epoch 120. A full-length run would likely lift every arm.
For reference, magnitude pruning without fine-tuning is far more pessimistic: 2:4 costs 0.378 PESQ and 1×4 at 80% costs 0.656. Almost all of it comes back, so pruning-only numbers are a poor guide to what a pattern actually costs.
Layout
Each directory holds both a runnable checkpoint and a kernel-oriented export:
g_best,config.json— PyTorch checkpoint, loadable with theNSNet2model in the repo above.g_best_fp32.onnx,g_best_int8.onnx— the streaming graph in FP32 and in static int8 (QDQ). The int8 graph preserves the sparsity pattern exactly.weights.npz— per matrix:<name>.weight(float32, dense with explicit zeros),<name>.mask(uint8, 1 = kept),<name>.bias, and golden vectors<name>.ref_x/<name>.ref_ywhereref_y = W @ ref_x + bias.manifest.json— shapes, pattern, grouping axis, achieved sparsity, ragged tail counts, and N at inference vs training.
verify.py at the top level needs only numpy:
python verify.py 2_4 # shapes, mask/weight agreement, pattern conformance,
# and the golden vectors
Conventions
Layout. Every weight is row-major (M, K), used as y = W · x + b with x
of shape (K, N).
N = 1 at deployment. The model runs one 16 ms frame at a time, so each of these is a matrix-vector product. During training N is 256 · T.
Grouping runs along K. For an N:M pattern the groups of M are contiguous
within a row — along the input dimension, contiguous in memory for a row-major
(M, K) array. This matches the NVIDIA 2:4 convention. The masking code
supports grouping along the output dimension too, if a kernel wants that.
Ragged tail. fc_in has K = 257 — 64 groups of 4 plus one leftover column,
left dense — so it measures 49.8% sparse rather than exactly 50%.
manifest.json reports tail_elements per matrix.
GRU gate packing. gru.weight_ih_l* and gru.weight_hh_l* are (3H, K):
PyTorch stacks the r/z/n gates along the output dimension, so each gate is a
contiguous block of rows and a group of 4 along K never straddles a gate
boundary. Each gate submatrix independently satisfies the pattern, so a
1200×400 packs as one matrix or as three 400×400 with identical results.
The matrices
The four GRU matrices are 69% of the weights and run once per frame, so they
dominate. gru.weight_hh_l0 and gru.weight_hh_l1 sit inside the recurrence
and cannot be batched over time even in principle — the strictest N=1 case here.
| matrix | M | K | params |
|---|---|---|---|
fc_in |
400 | 257 | 102,800 |
gru.weight_ih_l0 |
1200 | 400 | 480,000 |
gru.weight_hh_l0 |
1200 | 400 | 480,000 |
gru.weight_ih_l1 |
1200 | 400 | 480,000 |
gru.weight_hh_l1 |
1200 | 400 | 480,000 |
fc1 |
600 | 400 | 240,000 |
fc2 |
600 | 600 | 360,000 |
fc_out |
257 | 600 | 154,200 |
Usage
Kernel work — numpy only, no PyTorch:
import json
import numpy as np
npz = np.load("2_4/weights.npz")
W = npz["gru.weight_hh_l0.weight"] # (1200, 400) float32, explicit zeros
b = npz["gru.weight_hh_l0.bias"] # (1200,)
x = npz["gru.weight_hh_l0.ref_x"] # (400,) float32
assert np.allclose(W @ x + b, npz["gru.weight_hh_l0.ref_y"], atol=1e-4)
Running the model:
import json
import torch
from common.env import AttrDict
from nsnet2.model import NSNet2
h = AttrDict(json.load(open("2_4/config.json")))
model = NSNet2(h)
model.load_state_dict(torch.load("2_4/g_best", map_location="cpu")["generator"])
Reproducing a mask, or training a new one:
python -m nsnet2.train --config configs/ov_2to4.json \
--checkpoint_path cp_ov_2to4 --init_from <dense g_best>
New design (block-design study)
The six checkpoints above are the original NSNet2. A follow-up study asked
whether the kernel's constraints cost anything, found that the answer depends on
a training defect in the original block, fixed it, and re-measured. Everything
from that study is under new_design/: new-design dense models
for NSNet2 and ConvFSENet across a width sweep, their 2:4@c1 versions, and
Row-Fusion hand-off exports. Full write-up: BLOCK_DESIGN.md on branch
block-design of LarocheC/eco8-neaixt.
The story, in order
- The kernel's constraints are free on the original NSNet2. Restricting
2:4 to a codebook of 4 patterns (
2:4@c1:1010 / 0101 / 1001 / 0110, i.e. one weight kept from each adjacent pair) and using square, multiple-of-32 matrices cost nothing measurable, and every mask pattern above was free up to 80% sparsity. - Below the capacity knee the two models disagreed. At the same nonzero count, 2:4 NSNet2 beat a smaller dense model by +0.042 PESQ, whereas 2:4@c1 ConvFSENet lost to it by 0.020 (plain 2:4 tied).
- Mechanism: dead units. 55-72% of NSNet2's
fc_inReLUs died early in training (the input is a raw, all-positive |X|^0.3 magnitude, trained with Adam at lr 3e-3), so its layer inputs were redundant and pruning removed little information: a least-squares refit of the kept weights recovers pruned NSNet2. ConvFSENet's layer inputs are full-rank, so there was nothing free to prune. Refit at pruning time (plus a channel permutation so the mask fits the codebook) matched or beat fine-tuning. - Fix: a better block. Subtract a fixed per-bin input mean and warm up the
learning rate (design A1, NSNet2); for ConvFSENet also put a BatchNorm
after the frontend conv (design B1: conv -> BN -> ReLU). All of it folds
exactly into the plain models at export, so the deployed graph is unchanged.
Result: 0 dead
fc_in/frontend units. NSNet2 gains +0.02-0.025 PESQ dense and about +0.03 after 2:4 at equal nonzeros (adopted at 68/102, +0.040; a near-miss at 192/192, 0.0005 short of the pre-registered +0.03 bar). ConvFSENet gains +0.028 sparse but loses 0.016 dense: no evidence either way. - Pareto front on the new designs: 2:4 sits on the dense front, not above it. At equal nonzeros, 2:4@c1 + permutation + refit minus the dense line is -0.007 ± 0.013 for NSNet2 and -0.018 for ConvFSENet (ties at c128 and c160). So on a healthy model 2:4's value is kernel speed at equal nonzeros, not quality. The old-design sparse models (including the 2.78 checkpoints above) lie below the new dense front.
Caveats. g_best is selected on the same 824-utterance test split it is
scored on (it sits 0.02-0.10 above the last-5 average); mostly one seed per
width (two at NSNet2 192 and ConvFSENet 96); nonzero count stands in for
latency, nothing here was timed.
Pareto table
PESQ on the 824-utterance VoiceBank-DEMAND test split (from
new_design/results/BLOCK_DESIGN_PARETO.csv; 2:4 = the c1_perm_refit
point). NSNet2 width is hidden/fc (square, H = fc); ConvFSENet width is
residual/conv channels (res = C, conv = 2C). Nonzeros count every parameter
(biases included); for NSNet2 the dense count is after removing the few
train-set-dead fc2 units, which is why it is a little below the parameter
count.
| model | width | seed | dense nonzeros | dense PESQ | 2:4 nonzeros | 2:4@c1 PESQ | folders |
|---|---|---|---|---|---|---|---|
| NSNet2 | 64/64 | 1234 | 91,135 | 2.806 | 46,593 | 2.748 | sq64 |
| NSNet2 | 96/96 | 1234 | 179,711 | 2.884 | 91,137 | 2.861 | sq96 |
| NSNet2 | 128/128 | 1234 | 295,415 | 2.852 | 149,243 | 2.819 | sq128 |
| NSNet2 | 192/192 | 1234 | 616,571 | 2.863 | 310,077 | 2.853 | sq192 |
| NSNet2 | 192/192 | 2345 | 613,421 | 2.847 | 308,627 | 2.827 | sq192_s2345 |
| NSNet2 | 256/256 | 1234 | 1,048,559 | 2.853 | 526,837 | 2.862 | sq256 |
| NSNet2 | 384/384 | 1234 | 2,257,505 | 2.909 | 1,131,945 | 2.906 | sq384 |
| ConvFSENet | 64/128 | 1234 | 189,889 | 2.863 | 99,745 | 2.805 | c64 |
| ConvFSENet | 96/192 | 1234 | 395,297 | 2.860 | 204,785 | 2.835 | c96 |
| ConvFSENet | 96/192 | 2345 | 395,297 | 2.851 | 204,785 | 2.849 | c96_s2345 |
| ConvFSENet | 128/256 | 1234 | 674,433 | 2.897 | 346,689 | 2.861 | c128 |
| ConvFSENet | 160/320 | 1234 | 1,027,297 | 2.910 | 525,457 | 2.880 | c160 |
| ConvFSENet | 192/384 | 1234 | 1,453,889 | 2.872 | 741,089 | 2.843 | c192 |
Read the table against the dense front, not row by row: a 2:4 model at
~310k nonzeros should be compared with a dense model of ~310k nonzeros (between
the 128 and 192 rows), not with its own parent. Every PESQ above was
re-measured on the uploaded checkpoints and matches the CSV (see each folder's
info.json).
Folder index
new_design/
nsnet2/dense/sq{64,96,128,192,192_s2345,256,384}/ A1 NSNet2, folded to plain NSNet2
nsnet2/sparse24/sq{...same}/ 2:4@c1 + permutation + LS refit of the above
convfsenet/dense/c{64,96,96_s2345,128,160,192}/ B1 ConvFSENet, folded to plain ConvFSENet
convfsenet/sparse24/c{...same}/ 2:4@c1 + permutation + LS refit of the above
rowfusion_export/{nsnet2_sq192,nsnet2_sq384,convfsenet_c96,convfsenet_c192}/
kernel hand-off: weights.npz + manifest.json
rowfusion_export/verify_c1.py numpy-only checker for those exports
results/ BLOCK_DESIGN_{PARETO,PHASE3,RUNS}.csv, SPARSE_MATMUL_RUNS.csv
dense/<width>/—g_best({"generator": state_dict}),config.jsonandinfo.json. The design options are folded away: A1's input mean is infc_in.bias; B1's input mean and frontend BatchNorm are infrontend.0. The config has those flags removed, so the repo's existing loaders build a plainNSNet2/ ConvFSENet and load the state dict strictly. (warmup_stepsis left in the config; it only affects training.)sparse24/<width>/— the same three files plusmasks.npzandperm.npz. 2:4 models are stored dense with explicit zeros:g_bestis an ordinary dense checkpoint whose pruned weights are exactly 0.0, andmasks.npzholds one uint8 mask (1 = kept, =weight != 0) per pruned matrix, keyed by parameter name (ConvFSENet masks as(C_out, C_in)). Channels are in the natural (unpermuted) order, so these load and run like any plain model. In that order each row keeps one weight from each of a set of matched column pairs;perm.npzholds the channel permutation that makes those pairs adjacent, i.e. turns every mask into a literal2:4@c1mask. The NSNet2 2:4 models also carry the parent's train-set-deadfc2units as all-zero rows/columns (0-16 per model), so they have the plain full width.rowfusion_export/<model>/— the permuted versions, in the hand-off format of the top-level directories:weights.npzwith<name>.weight(row-major(M, K), explicit zeros),<name>.mask,<name>.bias, the 2-bit<name>.pattern_index(one index per group of 4 intomanifest["codebook"]["patterns"]), and golden vectorsref_x/ref_y. Every matrix passesverify_pattern(W, "2:4@c1")(0 violations). The permutation is exact except at the input:manifest["channel_permutation"] ["input_gather"]gives the one gatherx[input_gather]to apply to the 257 input bins; outputs come out in natural bin order. For ConvFSENet the matrices are the 20 1x1 convs (frontend.0,backend.0, andtcm.*.conv1x1/conv1x1_out) flattened toW = weight[:, :, 0], applied per frame asy[:, t] = W @ x[:, t] + b; the depthwise convs and BatchNorms are elementwise and stay ing_best(seemanifest["layout"]andmanifest["not_exported"]). Check an export withpython new_design/rowfusion_export/verify_c1.py new_design/rowfusion_export/nsnet2_sq192(the top-levelverify.pypredates codebook patterns).
Only nsnet2/sparse24/sq192 and sq384 are 2:4@c1-clean in every matrix
after permutation at full width. For the other NSNet2 widths the count of
surviving fc2 units is not a multiple of 4 and the pipeline left the 2-3
leftover fc_out columns dense (a ragged tail), which the codebook check
counts as a violation once the dead units are restored; each info.json
records the per-matrix counts. All ConvFSENet models are clean.
Loading
import json, torch
from huggingface_hub import snapshot_download
from common.env import AttrDict # from the GitHub repo above
from nsnet2.model import NSNet2
from convfsenet.model import build_causal_model
repo = "claroche1/nsnet2-sparse-rowfusion"
root = snapshot_download(repo, allow_patterns=["new_design/nsnet2/sparse24/sq192/*",
"new_design/convfsenet/dense/c96/*"])
d = f"{root}/new_design/nsnet2/sparse24/sq192"
h = AttrDict(json.load(open(f"{d}/config.json")))
ns = NSNet2(h)
ns.load_state_dict(torch.load(f"{d}/g_best", map_location="cpu")["generator"])
d = f"{root}/new_design/convfsenet/dense/c96"
h = AttrDict(json.load(open(f"{d}/config.json")))
cf = build_causal_model(h)
cf.load_state_dict(torch.load(f"{d}/g_best", map_location="cpu")["generator"])
A single file, e.g. a kernel export, needs no PyTorch:
import numpy as np
from huggingface_hub import hf_hub_download
p = hf_hub_download("claroche1/nsnet2-sparse-rowfusion",
"new_design/rowfusion_export/nsnet2_sq384/weights.npz")
npz = np.load(p)
W, idx = npz["gru.weight_hh_l0.weight"], npz["gru.weight_hh_l0.pattern_index"] # (1152, 384), (1152, 96)
Pipeline code (branch block-design): nsnet2/fold.py, convfsenet/fold.py
(folding), nsnet2/sparsity.py (2:4@c1, verify_pattern,
pattern_indices), nsnet2/export_sparse.py and convfsenet/export_sparse.py
(hand-off export).
Related
- claroche1/sparse-nsnet2-checkpoints — the same model under Butterfly / block-diagonal / Monarch structured factorizations (a different kind of sparsity: factorized transforms rather than masked dense matrices), plus the dense 2.845 baseline these were fine-tuned from.
Citation
NSNet2: Braun & Tashev, Towards efficient models for real-time deep noise suppression, ICASSP 2021. Training recipe built on MP-SENet.