Torch Floating Point
A PyTorch library for custom floating point quantization with autograd support. This library provides efficient implementations of custom floating point formats with automatic differentiation capabilities.
Features
- Custom Floating Point Formats: Support for arbitrary floating point configurations (sign bits, exponent bits, mantissa bits, bias)
- Autograd Support: Full PyTorch autograd integration for training with quantized weights
- CUDA Support: GPU acceleration for both forward and backward passes
- Straight-Through Estimator: Gradient-friendly quantization for training
Installation
From PyPI (Recommended)
pip install torch-floating-point
From Source
git clone https://github.com/SamirMoustafa/torch-floating-point.git
cd torch-floating-point
pip install -e .
Quick Start
import torch
from floating_point import FloatingPoint, Round
# Define a custom 8-bit floating point format (1 sign, 4 exponent, 3 mantissa bits)
fp8 = FloatingPoint(sign_bits=1, exponent_bits=4, mantissa_bits=3, bias=7, bits=8)
# Create a rounding function
rounder = Round(fp8)
# Create a tensor with gradients
x = torch.randn(10, requires_grad=True)
# Quantize the tensor
quantized = rounder(x)
# Use in training (gradients flow through)
loss = quantized.sum()
loss.backward()
print(f"Original: {x}")
print(f"Quantized: {quantized}")
print(f"Gradients: {x.grad}")
Training with Custom Floating Point Weights
import torch
import torch.nn as nn
from floating_point import FloatingPoint, Round
class FloatPointLinear(nn.Module):
def __init__(self, in_features, out_features, fp_config):
super().__init__()
self.weight = nn.Parameter(torch.randn(out_features, in_features))
self.bias = nn.Parameter(torch.randn(out_features))
self.rounder = Round(fp_config)
def forward(self, x):
quantized_weight = self.rounder(self.weight)
return torch.nn.functional.linear(x, quantized_weight, self.bias)
# Define custom floating point format
fp8 = FloatingPoint(sign_bits=1, exponent_bits=4, mantissa_bits=3, bias=7, bits=8)
# Create model with quantized weights
model = FloatPointLinear(10, 5, fp8)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
criterion = nn.MSELoss()
# Create simple data
x = torch.randn(32, 10)
y = torch.randn(32, 5)
# Training loop
for epoch in range(5):
optimizer.zero_grad()
# Forward pass
output = model(x)
loss = criterion(output, y)
# Backward pass
loss.backward()
optimizer.step()
print(f"Epoch {epoch + 1}: Loss = {loss.item():.6f}")
Contributing
- Fork the repository
- Create a feature branch (
git checkout -b feature/amazing-feature) - Install development dependencies (
make setup-dev) - Make your changes
- Run tests (
make test) - Run linting (
make lint) - Commit your changes (
git commit -m 'Add amazing feature') - Push to the branch (
git push origin feature/amazing-feature) - Open a Pull Request
License
This project is licensed under the MIT License - see the LICENSE file for details.
Citation
If you use this library in your research, please cite:
@software{moustafa2025torchfloatingpoint,
title={Torch Floating Point: A PyTorch library for custom floating point quantization},
author={Samir Moustafa},
year={2025},
url={https://github.com/SamirMoustafa/torch-floating-point}
}
Support
- Issues: GitHub Issues
- Discussions: GitHub Discussions
- Email: samir.moustafa.97@gmail.com
Release files for torch-floating-point 0.0.12
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| torch_floating_point-0.0.12.tar.gz | 16.3 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| torch_floating_point-0.0.12-cp310-cp310-manylinux_2_28_x86_64.whl | CPython 3.10 | CPython 3.10 | Linux glibc 2.28+ x86-64 | Details |
Total release size: 3.0 MB
Release files / torch_floating_point-0.0.12.tar.gz
| Download URL | torch_floating_point-0.0.12.tar.gz |
|---|---|
| Size | 16.3 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
d3ed7dfae8ac977169a890caf44f0e5f2afdb76d87b5045897f4a21f33a954a9
|
|
BLAKE2b-256 checksum How to use checksums |
15a8292481e806c75f1de57f2b804bc0752f9f2a0c8dbe19c10ac182f9ea673a
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.10.20
|
Release files / torch_floating_point-0.0.12-cp310-cp310-manylinux_2_28_x86_64.whl
| Download URL | torch_floating_point-0.0.12-cp310-cp310-manylinux_2_28_x86_64.whl |
|---|---|
| Size | 3.0 MB |
| Tags | CPython 3.10 Linux glibc 2.28+ x86-64 |
|
SHA-256 checksum How to use checksums |
1d275184add69bd8bc0cc35f6af8f230e0c4d65cf5cc8d19a77310db3a8ab4b2
|
|
BLAKE2b-256 checksum How to use checksums |
3f147159d9d8bc14a17937021108456e91fabfdef2c1b191faea98f6da287ccf
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.10.20
|