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 the NSNet2 model 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_y where ref_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

  1. 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.
  2. 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).
  3. Mechanism: dead units. 55-72% of NSNet2's fc_in ReLUs 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.
  4. 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.
  5. 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.json and info.json. The design options are folded away: A1's input mean is in fc_in.bias; B1's input mean and frontend BatchNorm are in frontend.0. The config has those flags removed, so the repo's existing loaders build a plain NSNet2 / ConvFSENet and load the state dict strictly. (warmup_steps is left in the config; it only affects training.)
  • sparse24/<width>/ — the same three files plus masks.npz and perm.npz. 2:4 models are stored dense with explicit zeros: g_best is an ordinary dense checkpoint whose pruned weights are exactly 0.0, and masks.npz holds 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.npz holds the channel permutation that makes those pairs adjacent, i.e. turns every mask into a literal 2:4@c1 mask. The NSNet2 2:4 models also carry the parent's train-set-dead fc2 units 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.npz with <name>.weight (row-major (M, K), explicit zeros), <name>.mask, <name>.bias, the 2-bit <name>.pattern_index (one index per group of 4 into manifest["codebook"]["patterns"]), and golden vectors ref_x / ref_y. Every matrix passes verify_pattern(W, "2:4@c1") (0 violations). The permutation is exact except at the input: manifest["channel_permutation"] ["input_gather"] gives the one gather x[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, and tcm.*.conv1x1 / conv1x1_out) flattened to W = weight[:, :, 0], applied per frame as y[:, t] = W @ x[:, t] + b; the depthwise convs and BatchNorms are elementwise and stay in g_best (see manifest["layout"] and manifest["not_exported"]). Check an export with python new_design/rowfusion_export/verify_c1.py new_design/rowfusion_export/nsnet2_sq192 (the top-level verify.py predates 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.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Dataset used to train claroche1/nsnet2-sparse-rowfusion