File size: 6,307 Bytes
8497c2a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1dc60cf
 
a933e0c
1dc60cf
 
 
a933e0c
 
8497c2a
 
 
 
 
 
 
 
 
 
 
 
a933e0c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1dc60cf
 
 
 
 
a933e0c
 
 
 
 
 
 
 
 
 
 
 
 
ce9d58a
 
 
 
 
a933e0c
 
 
ce9d58a
a933e0c
 
1dc60cf
 
 
 
a933e0c
 
 
 
 
 
 
 
 
ce9d58a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a933e0c
 
 
ce9d58a
 
a933e0c
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
---
language: en
library_name: agiws-neural-quant
license: apache-2.0
tags:
  - quantization
  - int4
  - int8
  - fp8
  - fp4
  - nf4
  - bitnet
  - qat
  - pytorch
---

# AGIWSNeuralQuant

Universal neural network quantization library β€” 53 formats, weights + activations, QAT, residual codebooks, per-head scales.

## What

Quantize ANY PyTorch model (vision encoders, LLMs, TTS, image/video generation) via ONE parameterized `Quantizer`:

- **INT** β€” int2/3/4/6/8, per-tensor/per-channel/per-group/per-head scales, symmetric/asymmetric
- **MX/NV formats** β€” MXFP4/MXFP6/MXFP8 (E8M0 block scale), NVFP4/NVFP6/NVFP8 (FP8 scale), MXINT2/4/6/8
- **FP4/FP6/FP8** β€” E2M1, E3M2/E2M3, E4M3/E5M2 floating-point with LUT-based dequant
- **NF4** β€” QLoRA NormalFloat 4-bit + double quantization
- **GGUF k-quants** β€” q2_k/q3_k/q4_k/q5_k/q6_k/q8_0/q4_0 (super-block layout)
- **BitNet 1.58** β€” ternary {-1,0,+1}, binary {+1,-1}
- **Codebook/VQ** β€” kmeans/fixed codebook, vector quantization, AQLM, residual multi-level codebooks
- **PTQ** β€” GPTQ (Hessian compensation), AWQ (activation-aware), SmoothQuant, LLM.int8 outlier, QuIP (rotation)
- **QAT** β€” learnable latent weights + scale via STE, learnable codebooks, dual-path distillation
- **Pruning** β€” structural (REAP/REAM), unstructured, magnitude-based
- **Activations** β€” ALL formats quantize activations, per-token/per-group/per-channel scale modes
- **per-head** β€” attention-specific scale per head (Q/K/V/O projections)

## Install

```bash
pip install git+https://huggingface.co/agiws/NeuralQuant
```

## Usage

```python
from agiws_neural_quant import quantize_model, save_model, load_model

# Quantize any PyTorch model in-place (53 formats via presets)
quantize_model(model, format="int4")           # 4-bit per-group
quantize_model(model, format="mxfp4")          # MXFP4 (E8M0 block scale)
quantize_model(model, format="nvfp4")           # NVFP4 (FP8 scale, weights+activations)
quantize_model(model, format="ternary")        # BitNet 1.58
quantize_model(model, format="codebook",
               residual_levels=3,              # residual codebooks cascade
               learnable=True)                  # QAT via STE

# per-head for attention projections
quantize_model(model, format="int4",
               per_head_modules=["q_proj","k_proj","v_proj","o_proj"],
               num_heads=32, head_dim=128)

# Save / load quantized model (preserves all format metadata)
save_model(model, "quantized.pt")
model2 = load_model("quantized.pt")

# QAT training
from agiws_neural_quant import UnifiedQATWrapper, strip_latent
qat = UnifiedQATWrapper(quantized_module)
loss = qat.distillation_loss(x)   # MSE(student, teacher.detach())
# ... optimizer.step() ...
strip_latent(quantized_module)     # bake to frozen inference
```

## Structure

```
src/agiws_neural_quant/
  base.py            β€” QuantizedWeight, QuantizedActivation, QuantizedModule (chunked dequant, dual-path, QAT)
  quantizer.py       β€” unified Quantizer (53 formats as parameter configurations)
  presets.py         β€” FORMAT_PRESETS dict (53 presets) + get_preset(format, **overrides)
  dispatch.py        β€” quantize_model(), save_model(), load_model(), walker
  ssm_patterns.py    β€” SSM/KDA exclusion patterns (Mamba/RWKV/linear-attention aware)
  training_unified.py β€” UnifiedQATWrapper, dual_path_loss, strip_latent
  training/ste.py    β€” STEQuantize, STECodebook (Straight-Through Estimator)
  ternary.py         β€” ternarize_tensor (BitNet 1.58 primitives)
  nf4/               β€” NF4 LUT, quantize/dequantize, double quantization (QLoRA)
  fp4/               β€” FP4 E2M1 LUT, pack/unpack, E8M0 LUT
  fp6/               β€” FP6 E3M2/E2M3 LUT, quantize/dequantize, pack/unpack
  fp8/               β€” FP8 E4M3/E5M2 LUT, quantize/dequantize
  kquant/            β€” GGUF k-quants (super-block layout)
  converters/        β€” universal file-to-file converter (stream, low memory)
    universal.py     β€” convert_model(input, output): safetensors/GGUF/.pt <-> .pt/safetensors/GGUF
    safetensors_io.py β€” stream read/write safetensors (lazy, low memory)
    gguf_reader.py   β€” GGUF binary format parser + k-quant decode
    hf_nvfp4.py      β€” NVFP4 safetensors -> QuantizedWeight
  analysis.py        β€” layer analysis (per-layer quantization suitability)
  cache.py           β€” teacher-cache for distillation
  extract.py         β€” checkpoint extraction (vision encoders, submodules)
tests/              β€” 169 tests (roundtrip, chunked, dual-path, QAT, SSM, residual, per-head, validation, converter)
scripts/            β€” auxiliary scripts
docs/
```

## Principles

- ONE parameterized `Quantizer` β€” all formats are configurations, not separate classes
- Weights AND activations quantized for ALL formats (W8A8, W4A4, NVFP4)
- Chunked dequant for minimal VRAM (adaptive chunk_size at runtime)
- Dual-path: student + teacher QuantizedModule for cross-quantization distillation
- Save/load preserves all format metadata (buffers + meta + quantizer config + dual-path + teacher)
- Supports any nn.Module: Linear, Conv1d/2d/3d, ConvTranspose, Embedding, LayerNorm, Bilinear
- Parameter validation: invalid combinations raise clear errors (not silent fallback)
- SSM-aware: excludes critical projections (q_proj/k_proj/b_proj) from quantization for Mamba/KDA/RWKV

## Converter (file-to-file, low memory)

```python
from agiws_neural_quant import convert_model

# NVFP4 safetensors (GLM-5.2) -> NeuralQuant .pt (preserves quantized layout)
convert_model("glm5.2-nvfp4.safetensors", "glm5.2.pt")

# fp16 safetensors -> NeuralQuant .pt with int4 quantization
convert_model("model.safetensors", "model_int4.pt", quant_format="int4")

# GGUF k-quant -> safetensors (dequantized to fp16)
convert_model("model.gguf", "model_fp16.safetensors", output_format="safetensors")

# NeuralQuant .pt -> safetensors (dequantized)
convert_model("model.pt", "model_out.safetensors", output_format="safetensors")
```

Processes one tensor at a time (peak memory ~2x largest tensor, not the full model).
Supports sharded safetensors (index.json + multiple shards).

## Status

- 53 quantization presets, all validated (roundtrip tested)
- Universal file-to-file converter (safetensors/GGUF/.pt, stream, low memory)
- 169 tests PASS
- Version 0.2.0