HyperPrune-Llama-3.1-8B-2to4

meta-llama/Llama-3.1-8B pruned to 2:4 semi-structured sparsity with HyperPrune (Sun & Sakuma, Learning Semi-Structured Sparsity for LLMs via Shared and Context-Aware Hypernetwork, ICLR 2026, OpenReview).

This is a reproduction run produced at Elastix as part of the BLADE sparsity-method comparison. It is plain sparse bf16/fp16 safetensors and loads with stock transformers:

from transformers import AutoModelForCausalLM, AutoTokenizer
m = AutoModelForCausalLM.from_pretrained("elastix-ai/HyperPrune-Llama-3.1-8B-2to4")
t = AutoTokenizer.from_pretrained("meta-llama/Llama-3.1-8B")

What differs from the paper's own recipe

paper / repo default this checkpoint
calibration corpus allenai/c4 DKYoon/SlimPajama-6B, validation (BLADE's corpus)
pruned modules see below see below

Everything else — hypernet architecture, both training stages, all learning rates, step counts, temperature, prior, row selection — is HyperPrune's own shipped setting.

Configuration

{
  "model": {
    "name_or_path": "meta-llama/Llama-3.1-8B",
    "tokenizer": "meta-llama/Llama-3.1-8B"
  },
  "data": {
    "dataset_name": "slimpajama",
    "num_samples": 128,
    "seq_len": 2048,
    "seed": 42
  },
  "hypernet": {
    "type": "mlp",
    "hidden_dim": 256,
    "emb_dim": 64,
    "use_layer_emb": false,
    "use_comp_emb": false,
    "use_hessian_diag": true
  },
  "training": {
    "sup_steps": 12000,
    "sup_lr": 0.001,
    "ft_lr": 0.0003,
    "ft_nsamples": 4,
    "rows_per_step": 400,
    "cascade_inner_steps": 300,
    "ft_mode": "cascade",
    "tau": 0.5,
    "prior_source": "sparsegpt",
    "wanda_residual_alpha": 2.0,
    "compensated_propagation": true,
    "use_weight_compensation": true,
    "train_on_compensated": true,
    "fixed_rows_count": 200,
    "fixed_rows_pos": "first",
    "dense_layers_list": []
  },
  "output": {
    "save_dir": "/home/ubuntu/hyperprune_work/outputs/hp-llama31_8b-2to4",
    "wanda_dir": "/home/ubuntu/hyperprune_work/outputs/hp-llama31_8b-2to4_ref",
    "preserve_wanda_dir": false
  }
}

Measured

metric value
overall decoder sparsity (check_sparsity) 0.5006 (all 32 decoder layers pruned)
WikiText-2 PPL (HyperPrune eval_ppl.py, seqlen 2048) 15.537
WikiText-2 word PPL (lm-eval-harness, BLADE's protocol) 22.27
training wall-clock 24.2 min
peak GPU during cascade FT 10.88 GB
GPU 1 x NVIDIA RTX PRO 6000 Blackwell (97 GB), CUDA 13.0, torch 2.13.0+cu130

Two things to know before comparing this number to the paper

1. Every decoder layer is pruned here. HyperPrune's own shipped configs set dense_layers_list: [0, 1], leaving 2 layers fully dense and yielding ~46.9 % sparsity rather than 50 %. This checkpoint prunes every layer, matching BLADE's two_four_all spec, so it is a true 2:4 model in the modules BLADE prunes.

2. Only a few percent of this mask was chosen by the hypernet. The shipped recipe sets fixed_rows_count: 200, so the hypernet decides the mask for the first 200 output rows of each projection and every remaining row keeps the SparseGPT prior's mask verbatim. This is HyperPrune's own default, kept here deliberately because the brief was to change nothing but the calibration corpus.

The two perplexity rows are different quantities and are not comparable to each other. The first is token-level PPL over concatenated WikiText-2 at seqlen 2048 (the Wanda/SparseGPT convention). The second is lm-evaluation-harness word_perplexity at max_length=2048, which is BLADE's protocol — pinned empirically by reproducing BLADE's dense LLaMA-2-7B value of 9.19 (measured 9.1915).

On hyperparameter transfer

HyperPrune ships a config for LLaMA-3-8B, not LLaMA-3.1-8B. There is nothing model-specific to carry across: configs/llama2_7b.yaml and configs/llama3_8b.yaml are byte-identical in their hypernet: and training: blocks, and across all four shipped configs the only per-model variation is three OBS-compensation booleans and two budget bumps on the 70B. HyperPrune performs no per-model hyperparameter tuning, so this checkpoint uses the same recipe as every other model in this collection.

Provenance

Produced from HyperPrune commit 6d093d7 with a small set of documented patches (bias-dtype autocast, calibration loader, disk-peak reduction, and — for 4:8 checkpoints — the N:M generalization, which the reference implementation does not ship). See the reproduction report for the full diff.

Evaluation Results

KL Divergence

Dataset Avg KL Total KL Tokens
wikitext2 0.905580 261525.9878 288,794
c4 0.895029 1654360.7671 1,848,388
slimpajama_calib 0.821121 6884702.7358 8,384,512

Downstream Accuracy

Task acc, None stderr, None
arc_challenge 0.3183 0.0136
arc_easy 0.6494 0.0098
hellaswag 0.4245 0.0049
mmlu 0.2990 0.0038
openbookqa 0.2440 0.0192
piqa 0.6942 0.0107
race 0.3742 0.0150
winogrande 0.6535 0.0134

Perplexity (2048-token windows, max_length=2048)

Dataset Word PPL Byte PPL
WikiText-2 22.2693 1.7866
C4 (en) 53.9565 1.9488

BLADE-Eval: lm-eval 0.4.10, torch 2.13.0+cu130, MLflow run

Downloads last month
901
Safetensors
Model size
8B params
Tensor type
F16
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for elastix-ai/HyperPrune-Llama-3.1-8B-2to4

Finetuned
(1480)
this model

Collection including elastix-ai/HyperPrune-Llama-3.1-8B-2to4