Skip to content
PrizmalAiPublic

About

Residual-Aware Sparsification and Pruning: Protecting Principal Subspaces During Sparsification or Pruning

Resources

Stars

0 stars

Watchers

0 watching

Forks

Latest commit

 

History

2 Commits

Folders and files

Repository files navigation

RASP: Residual-Aware Sparsification and Pruning

Protecting principal subspaces during sparsification or pruning.

RASP is a training-free, drop-in wrapper for existing pruning and activation-sparsity methods. Instead of masking a projection matrix directly, RASP decomposes each targeted weight W = L_r + R_r via a rank-r truncated SVD, keeps the principal low-rank component L_r dense, and applies the existing keep-mask only to the residual R_r. The dominant singular directions are therefore always evaluated densely, and the destructive effect of masking is confined to the lower-energy residual branch.

Because RASP only changes where a mask is applied, it is complementary to channel-selection rules such as Wanda, Griffin, and TEAL: the same mask can be reused, and its error is bounded by the residual singular tail σ_{r+1}(W). The rank r acts as an explicit error-control knob.

This repository accompanies the paper "Residual-Aware Sparsification and Pruning: Protecting Principal Subspaces During Sparsification or Pruning." Author and affiliation information is withheld for anonymous review.

Highlights

  • Training-free — no fine-tuning or continued pre-training; just an SVD-based re-parameterization applied at load time.
  • Method-agnostic — wraps Wanda (weight/channel pruning), Griffin (contextual activation sparsity), and TEAL (magnitude-based activation sparsity).
  • Rank-controlled error — increasing the retained rank r provably shrinks the part of the projection exposed to pruning.
  • Hardware-efficient decoding — a custom Triton sparse Residual-GEMV (RGEMV) kernel combines a dense low-rank path with a sparse residual path during autoregressive decoding.

Supported models

  • Llama-2 / Llama-3 (instruction-tuned variants used in the paper)
  • Gemma-2 (2B and 9B, instruction-tuned)
  • Mistral
  • Qwen2 / Qwen3 and DeepSeek-R1-distilled Qwen models

Contents

Install

  1. Clone the repository:
git clone <repository-url> rasp
cd rasp
  1. Set up the environment:
conda create -yn rasp python=3.11
conda activate rasp

pip install -e .
  1. (Optional) To calibrate thresholds for your own models or run accuracy evaluations, install the extra dependencies:
pip install -e ".[eval]"

Repository layout

rasp/
├── core/            # RASP algorithm: SVD decomposition, residual-aware adapters, pruning scores
│   ├── adapter.py            # RaspLayer, init_rasp_params
│   ├── adapted_modules.py    # FeedForwardwRASP, eval-time module swap
│   └── pruning_utils.py      # rasp_decompose_weight, rasp_core, pruning score functions
├── model.py         # Sparse model wrappers (Llama/Mistral/Qwen/Gemma) + build_sparse_layers
├── mlp.py           # RASPTealMLP (RASP applied on top of a TEAL base)
├── self_attn.py     # Sparse attention projections
├── eval.py          # Accuracy evaluation entry point (LM Evaluation Harness)
├── grab_acts.py     # Activation-histogram calibration
├── greedyopt.py     # Block-wise greedy per-layer sparsity allocation
└── run_*.sh         # Example configurations (generation / classification / reasoning / tests)
gpt-fast/            # Inference engine with the sparse RGEMV kernel
kernels/             # Triton kernels
utils/               # Shared model / data / sparsity utilities

Calibration

RASP reuses the base method's channel-selection rule, so calibration follows the base method. For a TEAL base, build activation histograms (and optionally per-layer activations) used to select thresholds:

CUDA_VISIBLE_DEVICES=0 python rasp/grab_acts.py \
  --model_name meta-llama/Meta-Llama-3-8B \
  --output_path ./calibration/Llama-3-8B \
  --use_rasp_teal \
  --rasp_rank 256 \
  --rasp_adapter_method rasp \
  --rasp_use_cache

The output directory contains the model, histograms, and activations used downstream.

Accuracy evaluation

Evaluation uses the LM Evaluation Harness. Select a mode with exactly one of --use_dense, --use_teal, --use_rasp, or --use_rasp_teal.

RASP wrapping a TEAL base at 50% sparsity, rank 256:

CUDA_VISIBLE_DEVICES=0 python rasp/eval.py \
  --use_rasp_teal \
  --model_name meta-llama/Meta-Llama-3-8B \
  --histogram_path ./calibration/Llama-3-8B/histograms \
  --tasks gsm8k \
  --num_fewshot 5 \
  --sparsity 0.5 \
  --rasp_rank 256 \
  --rasp_adapter_method rasp \
  --use_rasp_cache
  • --use_teal — TEAL base method only (add --greedy --greedy_sparsity_path <path>/lookup for block-wise sparsity).
  • --use_rasp — RASP only (residual-aware pruning without a TEAL activation base).
  • --use_dense — dense reference (no sparsification).
  • The single --sparsity value is shared across the TEAL uniform/greedy paths and the RASP residual mask.

See rasp/run_generation.sh, rasp/run_classification.sh, and rasp/run_reasoning.sh for complete task configurations.

Block-wise greedy sparsity

Optionally allocate sparsity per layer with a greedy search:

CUDA_VISIBLE_DEVICES=0 python rasp/greedyopt.py \
  --model_name meta-llama/Meta-Llama-3-8B \
  --model_type Llama-3-8B \
  --teal_path ./calibration/Llama-3-8B \
  --target_sparsity 0.9 \
  --base_step_size 0.05 \
  --last_fraction 0.25

Then pass --greedy --greedy_sparsity_path ./calibration/Llama-3-8B/lookup to eval.py.

Sparse inference

Single-batch decoding is served through the gpt-fast engine, which includes the Triton sparse RGEMV kernel that fuses a dense low-rank path with a sparse residual path.

cd gpt-fast

# Download and convert a checkpoint to gpt-fast format
python scripts/download.py --repo_id meta-llama/Llama-2-7b-hf --path $SAVE_PATH
python scripts/convert_hf_checkpoint.py --checkpoint_dir $SAVE_PATH/meta-llama/Llama-2-7b-hf

# Sparse decoding
CUDA_VISIBLE_DEVICES=0 python generate.py \
    --compile \
    --checkpoint_path $SAVE_PATH/meta-llama/Llama-2-7b-hf/model.pth \
    --hist_path ../models/Llama-2-7B/histograms \
    --sparsity 0.5 \
    --interactive

Remove --interactive to benchmark decoding throughput.

Configuration reference

--rasp_adapter_method — how the projection is re-parameterized:

Value Description
rasp Residual-aware decomposition (the method described in the paper)
rasp2 Two-component residual variant (experimental)
spectral Spectral variant of the residual-aware decomposition
lora / molora Low-rank / mixture-of-low-rank baselines
none No adapter (base method applied directly)

--rasp_prune_function — residual-branch scoring rule: norm, norm_w2, out_energy, hybrid, wanda.

--rasp_rank — retained SVD rank r. Larger r keeps more of the projection dense (smaller pruning error, larger parameter footprint).

Notes and limitations

  • The current sparse inference kernel supports FP16 (Triton does not currently support BF16 atomic_add).
  • Theoretical error guarantees are derived for the linear RASP surrogate; gated MLPs (SiLU + multiplicative gating) introduce nonlinear interactions not fully captured by the analysis.
  • Realized wall-clock speedups depend on efficient kernels and consistent realized sparsity; nominal sparsity alone does not guarantee acceleration.

License

Released under the Apache License 2.0 (see LICENSE).

About

Residual-Aware Sparsification and Pruning: Protecting Principal Subspaces During Sparsification or Pruning

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages