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
Install the PyTorch build you will run first (CPU or CUDA). This package compiles an extension against that torch.
From PyPI (Recommended)
pip install torch-floating-point --no-build-isolation
A C++ compiler is required. For CUDA kernels, also install a CUDA toolkit that matches torch.version.cuda (pip's torch wheel does not include nvcc). Set FORCE_CPU=1 to skip CUDA, or FORCE_CUDA=1 to require it.
From Source
git clone https://github.com/SamirMoustafa/torch-floating-point.git
cd torch-floating-point
pip install --no-build-isolation -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}")
Common layouts (OFP8 / MX)
E4M3 and E5M2 are OCP OFP8 encodings (Micikevicius et al., 2022). E2M1 and UE8M0 come from OCP MX (Rouhani et al., 2023). CUDA __nv_* comments below are aliases; decode goldens match cuda_fp4.h / cuda_fp8.h. AMD MI300 FP8 is HIP FNUZ, not OCP. For block-scaled y = (e - z) * s * s_global, use BlockRound. Constructors for FP6, E1M2, UE4M3, FNUZ, INT, and BFP mag are in the docs (not package exports).
from floating_point import FloatingPoint
# E2M1 (__nv_fp4_e2m1; reserved_exponent=False)
fp4_e2m1 = FloatingPoint(sign_bits=1, exponent_bits=2, mantissa_bits=1, bias=1, bits=4, reserved_exponent=False)
# E4M3-FN (__nv_fp8_e4m3): max finite ±448; codes 127/255 are NaN
fp8_e4m3fn = FloatingPoint(
sign_bits=1,
exponent_bits=4,
mantissa_bits=3,
bias=7,
bits=8,
max_mantissa_at_max_exponent=6,
reserved_exponent=False,
)
# E5M2 (__nv_fp8_e5m2)
fp8_e5m2 = FloatingPoint(sign_bits=1, exponent_bits=5, mantissa_bits=2, bias=15, bits=8, reserved_exponent=True)
# UE8M0 (__nv_fp8_e8m0): codes 0..254 → 2^(E-127); 255 → NaN
fp8_e8m0 = FloatingPoint(sign_bits=0, exponent_bits=8, mantissa_bits=0, bias=127, bits=8, reserved_exponent=True)
Block-scaled Round (NVFP4 / MX)
Shared per-block scale: y = (e - z) * s * s_global with e = Round_elem(x / (s * s_global) + z). Absmax mode detaches s (STE on x only); pass scales= for learnable QAT scales with gradients. s_global= is the optional second-level tensor scale (NVFP4).
OCP MX (MXFP8 / MXFP4) uses UE8M0 scales and block_size=32. NVFP4 is NVIDIA-only: E2M1 + E4M3 (UE4M3) scales, block_size=16, plus FP32 s_global (NVIDIA, 2025). NVIDIA block UE8M0 uses ue8m0_ceil; the OCP sample is ocp_floor; AWS Trainium3 is ocp_floor_x2. Element-wise Round(fp8_e8m0) remains nearest. Recipe tables with source URLs: the docs.
from floating_point import BlockFormat, BlockRound
nvfp4 = BlockFormat(fp4_e2m1, fp8_e4m3fn, 16, 6.0, "nearest")
mxfp8 = BlockFormat(fp8_e4m3fn, fp8_e8m0, 32, 448.0, "ue8m0_ceil")
mxfp8_ocp = BlockFormat(fp8_e4m3fn, fp8_e8m0, 32, 448.0, "ocp_floor")
y = BlockRound(nvfp4)(x) # absmax scales, STE on x only
y = BlockRound(nvfp4)(x, scales=learnable_s) # grad into scales
y = BlockRound(nvfp4)(x, s_global=tensor_scale)
y = BlockRound(mxfp8)(x)
y = BlockRound(nvfp4, rounder=MyRound)(x)
Contributing
See CONTRIBUTING.md for setup, tests, and pull-request expectations. Short version:
- Fork the repository
- Create a feature branch (
git checkout -b feature/amazing-feature) - Install development dependencies (
make env) - 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{moustafa2026torchfloatingpoint,
title={Torch Floating Point: a PyTorch library for custom floating-point formats with automatic differentiation},
author={Samir Moustafa},
year={2026},
version={0.0.22},
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.22
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.22.tar.gz | 35.3 kB | Details |
Release files / torch_floating_point-0.0.22.tar.gz
| Download URL | torch_floating_point-0.0.22.tar.gz |
|---|---|
| Size | 35.3 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
c8b0eab77e26810519ed929e73870f4885b7e3e85d9f54e88a379c9c5cf20141
|
|
BLAKE2b-256 checksum How to use checksums |
f8ab3135ff599e1f402054a840909836f9fc2ff552517a30c932bfead3aa3ca1
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.12.14
|