Skip to content

Latest commit

History

26 Commits

Folders and files

NameName
Last commit message
Last commit date

Repository files navigation

MPS BitsAndBytes

Real 4-bit and 8-bit quantization for PyTorch on Apple Silicon (M1/M2/M3/M4).

Full bitsandbytes-compatible API with Metal GPU acceleration for running large models on your Mac.

Features

FormatBitsMemory SavingsBest For
NF44-bit~75%LLM weights (normally distributed)
FP44-bit~75%Alternative with better dynamic range
FP8 E4M38-bit~50%Better precision than INT8
INT88-bit~50%General purpose

Plus:

  • Metal GPU kernels - Fused dequant+matmul, no Python overhead
  • Double quantization - Extra ~10% savings on scales
  • 8-bit Optimizers - Adam8bit, AdamW8bit, Lion8bit, SGD8bit
  • Paged Optimizers - CPU offloading for larger models
  • Quantized Embeddings - Embedding4bit, Embedding8bit
  • Sparse Operations - spmm_coo, spmm_coo_int8
  • LLM.int8 - OutlierAwareLinear with col+row quantization
  • HuggingFace compatible - BitsAndBytesConfig API works out of the box
  • QLoRA training - Freeze quantized weights, train LoRA adapters

Installation

pip install mps-bitsandbytes

Or from source:

git clone https://github.com/mpsops/mps-bitsandbytes
cd mps-bitsandbytes
pip install -e .

Quick Start

4-bit Quantization (NF4 - Recommended for LLMs)

importtorchfrommps_bitsandbytesimportLinear4bit, BitsAndBytesConfig, quantize_model# Convert a single layerlinear=torch.nn.Linear(4096, 4096).half().to('mps')
linear_4bit=Linear4bit.from_linear(linear) # NF4 by default# Or use FP4linear_fp4=Linear4bit.from_linear(linear, quant_type='fp4')
# Quantize entire modelconfig=BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.float16,
)
model=quantize_model(your_model, quantization_config=config, device='mps')

8-bit Quantization (FP8 or INT8)

frommps_bitsandbytesimportLinear8bit, LinearFP8# INT8 (traditional)linear_int8=Linear8bit.from_linear(linear)
# FP8 E4M3 (better precision)linear_fp8=LinearFP8.from_linear(linear)

8-bit Optimizers

Memory-efficient optimizers that store momentum/variance in 8-bit:

frommps_bitsandbytesimportAdam8bit, AdamW8bit, Lion8bit, SGD8bit# Drop-in replacement for torch optimizersoptimizer=Adam8bit(model.parameters(), lr=1e-3)
optimizer=AdamW8bit(model.parameters(), lr=1e-3, weight_decay=0.01)
optimizer=Lion8bit(model.parameters(), lr=1e-4)
optimizer=SGD8bit(model.parameters(), lr=0.1, momentum=0.9)

Paged Optimizers

Offload optimizer states to CPU for training larger models:

frommps_bitsandbytesimportPagedAdam, PagedAdamW, PagedLion# States are stored on CPU, copied to GPU during step()optimizer=PagedAdamW(model.parameters(), lr=1e-3, page_to_cpu=True)

Quantized Embeddings

Reduce embedding table memory by 50-75%:

frommps_bitsandbytesimportEmbedding4bit, Embedding8bit, EmbeddingNF4, EmbeddingFP4# Convert existing embeddingembed=torch.nn.Embedding(50000, 4096).half().to('mps')
embed_4bit=Embedding4bit.from_embedding(embed) # NF4 by defaultembed_fp4=EmbeddingFP4.from_embedding(embed) # FP4embed_8bit=Embedding8bit.from_embedding(embed) # INT8

Functional API

frommps_bitsandbytesimport (
# 4-bitquantize_nf4, dequantize_nf4, matmul_nf4,
quantize_fp4, dequantize_fp4, matmul_fp4,
# 8-bitquantize_fp8_e4m3, dequantize_fp8_e4m3, matmul_fp8_e4m3,
quantize_rowwise, dequantize_rowwise, matmul_int8,
# Col+Row INT8 (LLM.int8 style)quantize_colrow, dequantize_colrow, matmul_colrow,
# Double quantizationdouble_quant, dequant_absmax,
# Sparsespmm_coo, spmm_coo_int8, sparse_coo_from_dense, quantize_sparse_coo,
)
# NF4weight=torch.randn(4096, 4096, device='mps', dtype=torch.float16)
packed, absmax=quantize_nf4(weight, block_size=64)
output=matmul_nf4(input, packed, absmax)
# Double quantization (quantize the scales too)absmax_quant, absmax_scales=double_quant(absmax)

Memory Savings

ModelFP16INT8/FP8NF4/FP4
7B params14 GB7 GB3.5 GB
13B params26 GB13 GB6.5 GB
70B params140 GB70 GB35 GB

HuggingFace Integration

fromtransformersimportAutoModelForCausalLMfrommps_bitsandbytesimportBitsAndBytesConfig, quantize_modelmodel=AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-2-7b-hf",
torch_dtype=torch.float16,
)
config=BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.float16,
)
model=quantize_model(model, quantization_config=config, device='mps')

QLoRA Training

frommps_bitsandbytesimportBitsAndBytesConfig, quantize_model, Adam8bitfrompeftimportget_peft_model, LoraConfig# Load in 4-bitconfig=BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_quant_type="nf4")
model=AutoModelForCausalLM.from_pretrained("model_name", torch_dtype=torch.float16)
model=quantize_model(model, quantization_config=config, device='mps')
# Add LoRA adapters (train in fp16 while base stays quantized)lora_config=LoraConfig(r=8, lora_alpha=16, target_modules=["q_proj", "v_proj"])
model=get_peft_model(model, lora_config)
# Use 8-bit optimizer for extra memory savingsoptimizer=Adam8bit(model.parameters(), lr=1e-4)
trainer.train()

API Reference

Linear Modules

ClassFormatUse Case
Linear4bitNF4 or FP4LLM inference, QLoRA
Linear8bitINT8General quantization
LinearFP8FP8 E4M3Better precision 8-bit
OutlierAwareLinearINT8 + FP16LLM.int8 mixed precision
SwitchBackLinearINT8Training with quantized forward

Embedding Modules

ClassFormatMemory Savings
Embedding4bitNF4 (default)~75%
EmbeddingNF4NF4~75%
EmbeddingFP4FP4~75%
Embedding8bitINT8~50%

Optimizers

ClassDescription
Adam8bitAdam with 8-bit states
AdamW8bitAdamW with 8-bit states
Lion8bitLion optimizer with 8-bit momentum
SGD8bitSGD with 8-bit momentum
PagedAdamAdam with CPU offloading
PagedAdamWAdamW with CPU offloading
PagedLionLion with CPU offloading

Functional API

4-bit (NF4/FP4):

  • quantize_nf4(tensor, block_size=64) / quantize_fp4(...)
  • dequantize_nf4(packed, absmax, ...) / dequantize_fp4(...)
  • matmul_nf4(input, weight_packed, weight_absmax, bias=None) / matmul_fp4(...)

8-bit:

  • quantize_fp8_e4m3(tensor) - FP8 quantization
  • quantize_rowwise(tensor) - INT8 row-wise quantization
  • quantize_colrow(tensor) - INT8 col+row quantization (LLM.int8)
  • matmul_fp8_e4m3(...) / matmul_int8(...) / matmul_colrow(...)

Double Quantization:

  • double_quant(absmax, double_quant_block=256) - Quantize scales
  • dequant_absmax(absmax_quant, absmax_scales) - Restore scales

Sparse Operations:

  • sparse_coo_from_dense(tensor) - Convert to COO format
  • spmm_coo(row_idx, col_idx, values, dense, rows, cols) - Sparse matmul
  • spmm_coo_int8(...) - INT8 sparse matmul
  • quantize_sparse_coo(row_idx, col_idx, values) - Quantize sparse values

Utilities:

  • is_available() - Check MPS availability
  • has_native_kernels() - Check Metal kernels loaded
  • get_memory_footprint(model) - Calculate memory usage

Comparison with bitsandbytes

Featurebitsandbytes (CUDA)mps-bitsandbytes
NF4/FP4CUDAMetal
INT8/FP8CUDAMetal
Double quantCUDAMetal
8-bit OptimizersCUDAPure PyTorch
Paged OptimizersCUDAPure PyTorch
Quantized EmbeddingsCUDAPure PyTorch
Sparse matmulCUDAPure PyTorch
LLM.int8 (col+row)CUDAPure PyTorch
PlatformNVIDIAApple Silicon

Demo

# Chat with a quantized LLM
python demo/chat.py

License

MIT

Credits

About

8-bit quantization for PyTorch on Apple Silicon (M1/M2/M3/M4)

Resources

Stars

11 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages