PTWM

Quickstart

Install PTWM, compress and decompress a PyTorch tensor, and verify a bit-exact round-trip.

This page installs the package, compresses a tensor, decompresses it, and verifies a bit-exact round-trip.

Install

pip install ptwm

Requires Python ≥ 3.12 and PyTorch ≥ 2.6. For a specific PyTorch build (CPU-only, a particular CUDA version), install it from the matching PyTorch index before or alongside ptwm.

Round-trip a tensor

import torch
from ptwm import Compressor, Decompressor, CompressionConfig, Format

config = CompressionConfig(input_format=Format.TORCH)
compressor = Compressor(config)
decompressor = Decompressor()

tensor = torch.randn(1024, 1024, dtype=torch.bfloat16)
compressed = compressor.compress(tensor)
restored = decompressor.decompress(compressed)

assert torch.equal(tensor, restored)

Compress a safetensors file from the command line

ptwm compress model.safetensors
# produces model.safetensors.ptwm
ptwm decompress model.safetensors.ptwm
# restores model.safetensors

Load a compressed model into 🤗 transformers

from ptwm import patch_safetensors
patch_safetensors()

from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained("path/to/model")
# PTWM decompresses .ptwm shards on the fly

Next steps