diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml new file mode 100644 index 0000000..a66ddcb --- /dev/null +++ b/.github/workflows/publish.yml @@ -0,0 +1,112 @@ +name: Build and publish distributions + +on: + push: + branches: [main] + tags: ["v*"] + pull_request: + branches: [main] + workflow_dispatch: + +permissions: + contents: read + +concurrency: + group: ${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: false + +jobs: + build: + name: Build wheel and source distribution + runs-on: ubuntu-latest + timeout-minutes: 10 + steps: + - uses: actions/checkout@v7 + with: + persist-credentials: false + + - uses: actions/setup-python@v7 + with: + python-version: "3.11" + + - name: Check release version + if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/v') + run: | + python - <<'PY' + import os + import tomllib + from pathlib import Path + + config = tomllib.loads(Path("pyproject.toml").read_text()) + expected = "v" + config["project"]["version"] + actual = os.environ["GITHUB_REF_NAME"] + if actual != expected: + raise SystemExit(f"Tag {actual} does not match package version {expected}") + PY + + - name: Install build tools + run: python -m pip install --upgrade build twine + + - name: Build distributions + run: python -m build + + - name: Check package metadata + run: python -m twine check --strict dist/* + + - name: Check packaged CUDA sources + run: | + python - <<'PY' + from pathlib import Path + import tarfile + import zipfile + + required = { + path.as_posix() + for path in Path("entropack").rglob("*") + if path.suffix in {".cu", ".cuh"} + } + if not required: + raise SystemExit("No CUDA sources found in the checkout") + wheels = list(Path("dist").glob("*.whl")) + sources = list(Path("dist").glob("*.tar.gz")) + if len(wheels) != 1 or len(sources) != 1: + raise SystemExit("Expected one wheel and one source distribution") + with zipfile.ZipFile(wheels[0]) as archive: + wheel_files = set(archive.namelist()) + with tarfile.open(sources[0]) as archive: + source_files = {name.partition("/")[2] for name in archive.getnames()} + for name, files in (("wheel", wheel_files), ("source distribution", source_files)): + missing = required - files + if missing: + raise SystemExit(f"Missing CUDA sources in {name}: {sorted(missing)}") + print(f"Verified {len(required)} CUDA source files in both distributions") + PY + + - name: Upload distributions + uses: actions/upload-artifact@v7 + with: + name: python-distributions + path: dist/* + if-no-files-found: error + retention-days: 14 + + publish: + name: Publish to PyPI + needs: build + if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/v') + runs-on: ubuntu-latest + timeout-minutes: 10 + environment: + name: pypi + url: https://pypi.org/project/entropack/ + permissions: + id-token: write + steps: + - name: Download distributions + uses: actions/download-artifact@v8 + with: + name: python-distributions + path: dist/ + + - name: Publish to PyPI + uses: pypa/gh-action-pypi-publish@release/v1 diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..c71c93c --- /dev/null +++ b/.gitignore @@ -0,0 +1,27 @@ +# Python +__pycache__/ +*.py[cod] +*.egg-info/ +build/ +dist/ +.pytest_cache/ + +# Environments +.venv/ +venv/ + +# Editors / tools +.idea/ +.vscode/ +.qoder/ +.claude/ + +# Benchmark artifacts +*.log + +# Sphinx +docs/_build/ +docs/*/_build/ + +# Local tests +/tests/ diff --git a/LICENSE b/LICENSE index 261eeb9..84a34e9 100644 --- a/LICENSE +++ b/LICENSE @@ -186,7 +186,7 @@ same "printed page" as the copyright notice for easier identification within third-party archives. - Copyright [yyyy] [name of copyright owner] + Copyright [2026] [ModelScope] Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except in compliance with the License. diff --git a/README.md b/README.md new file mode 100644 index 0000000..a5a120a --- /dev/null +++ b/README.md @@ -0,0 +1,140 @@ +# EntroPack + +### General-purpose tensor compression for PyTorch + +EntroPack is a general-purpose tensor compression library for PyTorch. It supports lossless +compression for exact recovery and lossy compression with a target bitrate to balance storage +and reconstruction accuracy. GPU encoding and decoding compress tensors and restore them +in their original shape and dtype. + +[![License](https://img.shields.io/badge/license-Apache_2.0-blue.svg)](LICENSE) +![Python](https://img.shields.io/badge/python-%3E%3D3.10-blue.svg) + +[Documentation](docs/en/index.rst) · [中文](README_zh.md) + +- **Flexible bitrates.** Compress each weight matrix at any non-integer target bitrate, + or preserve every input bit with a lossless scheme. +- **Dtype preservation.** Restore tensors in their input dtype, including BF16, FP16, FP8, + and INT8. +- **PyTorch integration.** Compress and restore tensors through a common API, and save + them with `state_dict`. Compressed linear layers provide an integration for model weights. + +## Installation + +Use Python 3.10 or later and install a CUDA-enabled build of PyTorch 2.10 or later +for your environment. + +### Install from source (recommended) + +```bash +git clone https://github.com/modelscope/entropack.git +cd entropack +pip install -e ".[cuda13]" +``` + +### Install from PyPI + +PyPI releases may lag behind source updates. Install from source for the latest features. + +```bash +pip install "entropack[cuda13]" +``` + +Both installation methods above use CUDA 13 and include the matching CuPy package. +For CUDA 12, replace `cuda13` with `cuda12` in either command. If a compatible CuPy is +already installed, use `pip install -e .` for source installation or `pip install entropack` for PyPI. + +See [Quick start](docs/en/Usage/Quick-start.md) for environment requirements and usage examples. + +## Get started + +### Direct tensor compression + +This example compresses a 2D BF16 tensor at a target of 3.5 bits per element, +then decompresses it to a tensor with the original shape and dtype: + +```python +import torch +import entropack as ep + +tensor = (torch.randn(256, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.LatticeRANSConfig(target_bpp=3.5) + +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) + +print(f"Target: {config.target_bpp:.2f} bits per element") +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +print(restored.shape, restored.dtype) +``` + +`target_bpp` is the requested number of bits per element (bpp). `actual_bpp` reports the stored +rate, including metadata. Targets from 1 to 11 are supported, including non-integer values. + +### Compressed Linear + +`CompressedLinear.from_linear` compresses an existing `torch.nn.Linear`'s weights using +the supplied Config and returns a new Compressed Linear. Call it as `layer(x)` to compute +the output, just as with an ordinary linear layer: + +```python +import torch +import entropack as ep + +linear = torch.nn.Linear(256, 256, dtype=torch.bfloat16, device="cuda") +config = ep.LatticeRANSConfig(target_bpp=4.0) +layer = ep.CompressedLinear.from_linear(linear, config=config) +x = torch.randn(8, 256, dtype=torch.bfloat16, device="cuda") + +with torch.inference_mode(): + output = layer(x) +print(output.shape, f"{layer.compressed_bits:.2f} bits per weight") +``` + +See [Compressed Linear usage](docs/en/Usage/Linear-layers.md) for model replacement and checkpoint examples. + +## Configuration + +Config defines the compression scheme and its settings. It is passed to tensor compression +and decompression functions or to a compressed Linear layer's constructor. + +| Config | Compression and use cases | +| --- | --- | +| `DFloat11Config()` | Specialized lossless compression for BF16 tensors. Decompression restores every input bit, for applications requiring exact recovery. | +| `TileANSConfig()` | Lossless compression for BF16, FP16, FP32, FP8, INT8, and other supported dtypes. The compression ratio depends on the input data distribution. | +| `LatticeRANSConfig(target_bpp=...)` | Lossy compression of 2D floating-point and integer tensors. `target_bpp` specifies the target bits per element, from 1 to 11 including non-integer values, to balance storage size and reconstruction accuracy. | + +See [Compression configuration](docs/en/Usage/Configuration.md) for scheme selection +and the complete parameter reference. + +## Performance + +On one NVIDIA H20, EntroPack compresses the weights of all 276 linear layers in +Z-Image-Turbo's diffusion transformer at a 4 bpp target in **3.5 seconds**. The compressed +weights occupy **4.02 bpp**, with **7.18%** relative L2 reconstruction error. Inference with +these weights takes **544.7 ms** per denoising step, only **7.7%** above the original BF16 +model's 505.7 ms. + +## Documentation + +| Guide | Contents | +| --- | --- | +| [Quick start](docs/en/Usage/Quick-start.md) | Install and run tensor compression and Compressed Linear examples | +| [Compression configuration](docs/en/Usage/Configuration.md) | Choose a scheme and look up supported dtypes and parameters | +| [Tensor compression](docs/en/Usage/Tensor-compression.md) | Encode, decode, inspect storage, move data, and save or load tensors | +| [Compressed Linear usage](docs/en/Usage/Linear-layers.md) | Replace model layers, use low-precision computation, and manage checkpoints | +| [API reference](docs/en/API_Reference/index.md) | Look up functions, classes, and properties | + +Compression principles: [DFloat11](docs/en/Principles/DFloat11.md), [tile-ANS](docs/en/Principles/Tile-ANS.md), +and [EntroPack lattice quantization](docs/en/Principles/Lattice-rANS.md). + +## Acknowledgements + +EntroPack's design is inspired by [DFloat11](https://github.com/LeanModels/DFloat11), +[dahuffman](https://github.com/soxofaan/dahuffman), +[DietGPU](https://github.com/facebookresearch/dietgpu), and +[tile-ANS](https://arxiv.org/abs/2606.15789). + +## License + +[Apache License 2.0](LICENSE). diff --git a/README_zh.md b/README_zh.md new file mode 100644 index 0000000..8d9c1b5 --- /dev/null +++ b/README_zh.md @@ -0,0 +1,131 @@ +# EntroPack + +### 面向 PyTorch 的通用张量压缩 + +EntroPack 是一个面向 PyTorch 的通用张量压缩库,支持完整保留原始数据的无损压缩, +以及通过目标码率控制存储大小与重建精度的有损压缩。EntroPack 提供 GPU 编解码, +将压缩后的张量恢复为原来的形状和数据类型。 + +[![License](https://img.shields.io/badge/license-Apache_2.0-blue.svg)](LICENSE) +![Python](https://img.shields.io/badge/python-%3E%3D3.10-blue.svg) + +[文档](docs/zh/index.rst) · [English](README.md) + +- **灵活设置码率。** 支持每个权重矩阵以任意非整数目标码率压缩,也可以选择逐位保留输入的无损方案。 +- **保留数据类型。** 解压后保留输入的数据类型,包括 BF16、FP16、FP8、INT8 等。 +- **接入 PyTorch。** 通过统一接口压缩和恢复张量,使用 `state_dict` 保存;模型权重还可以通过 Compressed Linear 接入。 + +## 安装 + +需要 Python 3.10 及以上,并先安装与环境匹配的 CUDA 版 PyTorch 2.10 及以上。 + +### 源码安装(推荐) + +```bash +git clone https://github.com/modelscope/entropack.git +cd entropack +pip install -e ".[cuda13]" +``` + +### 从 PyPI 安装 + +PyPI 版本更新可能有所延迟,如需最新功能,推荐从源码安装。 + +```bash +pip install "entropack[cuda13]" +``` + +上述两种安装方式均以 CUDA 13 为例,并包含对应版本的 CuPy。使用 CUDA 12 时, +将命令中的 `cuda13` 改为 `cuda12`。如果已安装匹配的 CuPy,源码安装和 PyPI 安装 +可分别使用 `pip install -e .` 和 `pip install entropack`。 + +环境要求与使用示例见[快速上手](docs/zh/Usage/Quick-start.md)。 + +## 快速开始 + +### 直接压缩张量 + +以下示例将一个二维 BF16 张量以每元素 3.5 bit 为目标压缩, +再解压为相同形状和数据类型的张量: + +```python +import torch +import entropack as ep + +tensor = (torch.randn(256, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.LatticeRANSConfig(target_bpp=3.5) + +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) + +print(f"Target: {config.target_bpp:.2f} bits per element") +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +print(restored.shape, restored.dtype) +``` + +`target_bpp` 表示期望的每元素比特数(bpp),`actual_bpp` 返回包含元数据的实际存储码率。 +目标范围为 1–11,支持非整数值。 + +### 使用 Compressed Linear + +`CompressedLinear.from_linear` 按传入的 Config 压缩现有 `torch.nn.Linear` 的权重, +返回一个新的 Compressed Linear。仍可像普通线性层一样,通过 `layer(x)` 计算输出: + +```python +import torch +import entropack as ep + +linear = torch.nn.Linear(256, 256, dtype=torch.bfloat16, device="cuda") +config = ep.LatticeRANSConfig(target_bpp=4.0) +layer = ep.CompressedLinear.from_linear(linear, config=config) +x = torch.randn(8, 256, dtype=torch.bfloat16, device="cuda") + +with torch.inference_mode(): + output = layer(x) +print(output.shape, f"{layer.compressed_bits:.2f} bits per weight") +``` + +模型中的层替换和检查点操作见 [Compressed Linear 使用指南](docs/zh/Usage/Linear-layers.md)。 + +## Config:压缩配置 + +Config 定义压缩方案及其参数,在调用张量编解码函数或构造 Compressed Linear 时传入。 + +| Config | 压缩方式与适用场景 | +| --- | --- | +| `DFloat11Config()` | 专用于 BF16 张量的无损压缩,解压后逐位恢复输入,适合要求精确恢复的场景。 | +| `TileANSConfig()` | 支持 BF16、FP16、FP32、FP8、INT8 等多种数据类型的无损压缩。压缩比取决于输入的数据分布。 | +| `LatticeRANSConfig(target_bpp=...)` | 支持浮点和整数二维张量的有损压缩。`target_bpp` 指定每元素的目标比特数,范围为 1–11,支持非整数值,用于调整存储大小与重建精度之间的取舍。 | + +方案选择和完整参数见[压缩配置](docs/zh/Usage/Configuration.md)。 + +## 性能 + +在单张 NVIDIA H20 上,EntroPack 以 4 bpp 为目标压缩 Z-Image-Turbo 扩散 Transformer +的 276 个线性层权重,耗时 **3.5 秒**。压缩后的实际存储为 **4.02 bpp**, +权重相对 L2 重建误差为 **7.18%**。使用压缩权重推理时,去噪单步耗时为 **544.7 ms**, +相对原始 BF16 模型的 505.7 ms 仅增加 **7.7%**。 + +## 文档 + +| 指南 | 内容 | +| --- | --- | +| [快速上手](docs/zh/Usage/Quick-start.md) | 安装并运行张量压缩与 Compressed Linear 示例 | +| [压缩配置](docs/zh/Usage/Configuration.md) | 选择方案、查看支持类型与完整参数 | +| [通用张量压缩](docs/zh/Usage/Tensor-compression.md) | 编解码、存储统计、设备迁移和保存加载 | +| [Compressed Linear 使用指南](docs/zh/Usage/Linear-layers.md) | 模型替换、低精度计算和检查点使用 | +| [API 参考](docs/zh/API_Reference/index.md) | 查询函数、类与属性 | + +压缩原理:[DFloat11](docs/zh/Principles/DFloat11.md)、[tile-ANS](docs/zh/Principles/Tile-ANS.md)、 +[EntroPack 格量化](docs/zh/Principles/Lattice-rANS.md)。 + +## 致谢 + +EntroPack 的设计受到 [DFloat11](https://github.com/LeanModels/DFloat11)、 +[dahuffman](https://github.com/soxofaan/dahuffman)、 +[DietGPU](https://github.com/facebookresearch/dietgpu) 和 +[tile-ANS](https://arxiv.org/abs/2606.15789) 的启发。 + +## 许可证 + +[Apache License 2.0](LICENSE)。 diff --git a/docs/assets/entropack-pipeline.png b/docs/assets/entropack-pipeline.png new file mode 100644 index 0000000..ba30f74 Binary files /dev/null and b/docs/assets/entropack-pipeline.png differ diff --git a/docs/en/.readthedocs.yaml b/docs/en/.readthedocs.yaml new file mode 100644 index 0000000..c534012 --- /dev/null +++ b/docs/en/.readthedocs.yaml @@ -0,0 +1,17 @@ +# .readthedocs.yaml +# Read the Docs configuration file +# See https://docs.readthedocs.io/en/stable/config-file/v2.html for details + +version: 2 + +build: + os: ubuntu-22.04 + tools: + python: "3.11" + +sphinx: + configuration: docs/en/conf.py + +python: + install: + - requirements: docs/requirements.txt diff --git a/docs/en/API_Reference/index.md b/docs/en/API_Reference/index.md new file mode 100644 index 0000000..03c0360 --- /dev/null +++ b/docs/en/API_Reference/index.md @@ -0,0 +1,174 @@ +# API reference + +The main functions and classes are available through `import entropack as ep`. +For complete examples, see [General tensor compression](../Usage/Tensor-compression.md) and +[Compressed Linear usage](../Usage/Linear-layers.md). Configuration parameters are listed in +[Compression configuration](../Usage/Configuration.md). + +## Tensor encoding and decoding + +### compress + +```text +compress(tensor: torch.Tensor, config: CompressionConfig) -> CompressedTensor +``` + +Compresses `tensor` using the scheme selected by `config`. Returns a `CompressedTensor` +that decompresses to the same shape and dtype as the input by default. + +Input requirements depend on the scheme: +DFloat11 accepts BF16 tensors, Tile-ANS supports multiple dtypes, and lattice quantization requires +a nonempty two-dimensional tensor with finite values. + +| Parameter | Meaning | +| --- | --- | +| `tensor` | PyTorch tensor to compress | +| `config` | Required. Selects a scheme through `DFloat11Config`, `TileANSConfig`, or `LatticeRANSConfig` | + +Some encoding failures emit a warning explaining the failure and return an uncompressed +container with `compress_method == "raw"`. Invalid configurations and backend dispatch +failures raise errors. + +### decompress + +```text +decompress(compressed: CompressedTensor, config: CompressionConfig) -> torch.Tensor +``` + +Returns a tensor with `compressed.shape`, `compressed.dtype`, and `compressed.device`. +Without an output dtype conversion, lossless schemes restore input values bit for bit; +lossy schemes return an approximate reconstruction. + +| Parameter | Meaning | +| --- | --- | +| `compressed` | Container returned by `compress` or restored from a checkpoint | +| `config` | Configuration for the corresponding scheme; decode settings control reconstruction | + +Changing encoding parameters such as `target_bpp` at decode time does not alter the stored data +or requantize the tensor. + +## CompressedTensor + +A `torch.Tensor` subclass holding the compressed representation of one tensor. Usually returned +by `compress` or restored from a checkpoint with `from_state_dict`. Use `decompress` before +performing numerical operations. + +### Common properties + +| Property | Type | Meaning | +| --- | --- | --- | +| `shape` | `torch.Size` | Original tensor shape | +| `dtype` | `torch.dtype` | Reconstructed tensor dtype | +| `encoded_dtype` | `torch.dtype` | Dtype used for encoding | +| `compress_method` | `str` | Scheme actually used by the container | +| `lossless` | `bool` | Whether the scheme is lossless | +| `actual_bpp` | `float` | Stored bits per element, including metadata | + +`actual_bpp = 8 * storage_nbytes() / math.prod(shape)`. +This measures the compressed representation, not checkpoint file size or runtime memory use. + +### Common methods + +| Method | Returns | Meaning | +| --- | --- | --- | +| `to(...)` | `CompressedTensor` | Changes device or output dtype without recompression; `copy=True` copies storage | +| `storage_nbytes(include_header=True)` | `int` | Total compressed size in bytes; `include_header=False` excludes the container header | +| `state_dict(prefix="")` | `dict[str, torch.Tensor]` | Exports the compressed tensor for saving | +| `CompressedTensor.from_state_dict(state, prefix="")` | `CompressedTensor` | Restores the container from that dictionary without recompression | + +Use the same `prefix` when saving and restoring. The dictionary can be saved with `torch.save` +and loaded with `torch.load(..., weights_only=True)`. Set `map_location` to choose the device +on which it will be restored. +Loading restores the encoded dtype. Call `.to(dtype=...)` afterwards if a different output dtype is needed. + +## CompressedLinear + +A linear layer that uses reconstructed weights for each forward call. Weights are stored in +compressed form; the bias is not compressed. +Requires a CUDA GPU and the matching CuPy package. + +### Creating a layer + +```text +CompressedLinear(in_features, out_features, bias=True, *, + config=None, device=None, dtype=torch.bfloat16) +CompressedLinear.from_linear(linear, **kwargs) -> CompressedLinear +``` + +| Constructor parameter | Meaning | +| --- | --- | +| `in_features` / `out_features` | Input and output feature counts | +| `bias` | Whether to include a bias | +| `config` | Weight compression configuration. The default `None` selects DFloat11 for BF16 and Tile-ANS for other supported dtypes | +| `device` | Bias device when constructing a layer directly | +| `dtype` | Dtype used when compressing weights and initializing the bias | + +`from_linear` returns a new layer with compressed source weights and a copy of the bias. +The source weights must already be loaded and cannot be on the `meta` device. +A typical call is `ep.CompressedLinear.from_linear(linear, config=config)`. +Pass `config` or `dtype` through `kwargs`; `dtype` defaults to the source weight dtype, +and the device is taken from the source layer. + +Calling the constructor directly creates a layer without weight data. Call `compress_weight` +or load a checkpoint before running inference. + +### Common methods and properties + +| Interface | Returns | Meaning | +| --- | --- | --- | +| `compress_weight(weight)` | `None` | Initializes compressed weights with shape `(out_features, in_features)` | +| `dequantize(device=None)` | `torch.Tensor` | Returns dense weights with `weight.dtype`, on the layer's device unless `device` is specified | +| `forward(x)` | `torch.Tensor` | Applies the layer to `x` of shape `(..., in_features)` and returns shape `(..., out_features)` | +| `weight` | `CompressedTensor` | Frozen compressed weight parameter held by the layer | +| `container_dtype` | `torch.dtype` | Dtype of the weights or quantized codes in the compressed container | +| `stored_nbytes` | `int` | Weight storage bytes, including metadata and low-precision quantization scales, excluding bias | +| `compressed_bits` | `float` | `8 * stored_nbytes / (in_features * out_features)` | + +Invoke the forward operation as `layer(x)`. `.weight` is a compressed tensor; use `dequantize()` +when numerical weights are needed. + +Use standard `state_dict()` / `load_state_dict()` calls to save and restore layer state. +Before loading, construct layers with matching classes, shapes, container dtypes, and compression +schemes. The Config object itself is not stored in the checkpoint. `.to(device)` moves the layer; +model dtype conversion does not re-encode its compressed weights. + +## CompressedFP8Linear and CompressedINT8Linear + +Linear layers with FP8 or INT8 weights and activations. They share the constructor arguments, +`from_linear`, storage properties, and checkpoint interfaces of `CompressedLinear`. +For input shape `(..., in_features)`, the output has shape `(..., out_features)` and the input's +dtype and device. + +| Class | Weight and activation format | CUDA GPU requirement | +| --- | --- | --- | +| `CompressedFP8Linear` | FP8 E4M3FN | SM8.9 or later | +| `CompressedINT8Linear` | INT8 | SM8.0 or later | + +The layer class determines the code format. The constructor's `dtype` argument does not change +the FP8 or INT8 format. + +With `config=None`, quantized codes are stored directly. Passing `LatticeRANSConfig` applies +additional lossy compression with `1 <= target_bpp < 8`. `stored_nbytes` includes the +per-row quantization scales needed to reconstruct weights. + +| Method | Returns | Meaning | +| --- | --- | --- | +| `codes(device=None)` | FP8 or INT8 tensor | Restores quantized codes without applying row scales | +| `dequantize(device=None)` | `torch.Tensor` | Returns the layer's initialization dtype, which `from_linear` defaults to the source weight dtype | + +Both methods return tensors on the layer's device unless `device` is specified. +Lossy compression may change the codes from their initial quantized values. + +## Config classes + +Configs control tensor encoding and decoding as well as weight storage in Compressed Linear. +The following classes inherit from `CompressionConfig`: + +| Class | Purpose | +| --- | --- | +| `DFloat11Config` | Lossless BF16 compression | +| `TileANSConfig` | Lossless tiled ANS compression for multiple dtypes | +| `LatticeRANSConfig` | Lossy lattice quantization with bitrate controlled by `target_bpp` | + +See [Compression configuration](../Usage/Configuration.md) for scheme selection, defaults, +and parameter ranges. diff --git a/docs/en/Makefile b/docs/en/Makefile new file mode 100644 index 0000000..4ae5e53 --- /dev/null +++ b/docs/en/Makefile @@ -0,0 +1,14 @@ +# Minimal makefile for Sphinx documentation + +SPHINXOPTS ?= +SPHINXBUILD ?= sphinx-build +SOURCEDIR = . +BUILDDIR = ../_build/en + +help: + @$(SPHINXBUILD) -M help "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O) + +.PHONY: help Makefile + +%: Makefile + @$(SPHINXBUILD) -M $@ "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O) diff --git a/docs/en/Principles/DFloat11.md b/docs/en/Principles/DFloat11.md new file mode 100644 index 0000000..49cb49b --- /dev/null +++ b/docs/en/Principles/DFloat11.md @@ -0,0 +1,52 @@ +# DFloat11 + +DFloat11 compresses BF16 tensors without changing their bits. It exploits the fact that +the exponent values in many tensors are concentrated in a small part of the available +range. Frequent exponents can then be represented with fewer bits, while the sign and +fraction remain unchanged. + +## Encoding + +A BF16 value contains one sign bit, eight exponent bits, and seven fraction bits. +The encoder separates each value into an exponent and a byte containing its sign and +fraction. These bytes are stored directly. The exponent stream is compressed using a +Huffman code constructed from the tensor's exponent frequencies: common exponents receive +short codes and uncommon exponents receive longer ones. + +Huffman codes have variable lengths, so a decoder cannot start at an arbitrary bit and +immediately identify the next symbol. EntroPack records entry positions and symbol counts +for coding regions, allowing different regions to be decoded in parallel. The entry +information and Huffman tables add metadata to the compressed representation. + +## Decoding and storage + +Decoding recovers the exponent sequence through the Huffman tables and combines each +exponent with its stored sign and fraction. Reassembling these fields restores the +original BF16 bits, and the stored shape determines how they form the output tensor. +No numerical quantization or rounding is involved. + +The achieved size depends on the exponent distribution and decoding metadata. A concentrated +distribution offers more compression than a broad one, and metadata has a larger relative +cost for small tensors. `DFloat11Config` does not specify a target bitrate. The name +DFloat11 does not imply that every tensor is stored at exactly 11 bits per element. + +## Usage + +The example compresses a BF16 tensor and verifies that decompression preserves its bits. +See [Config](../Usage/Configuration.md) for coding-region parameters. + +```python +import torch +import entropack as ep + +tensor = (torch.randn(256, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.DFloat11Config() + +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) + +assert restored.shape == tensor.shape +assert restored.dtype == tensor.dtype +assert torch.equal(restored.view(torch.uint8), tensor.view(torch.uint8)) +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +``` diff --git a/docs/en/Principles/Lattice-rANS.md b/docs/en/Principles/Lattice-rANS.md new file mode 100644 index 0000000..c9dd9d9 --- /dev/null +++ b/docs/en/Principles/Lattice-rANS.md @@ -0,0 +1,76 @@ +# EntroPack lattice compression + +EntroPack's lattice scheme combines lossy vector quantization with lossless entropy coding +to compress two-dimensional tensors. `LatticeRANSConfig` accepts target bitrates from 1 to +11 bits per element, including non-integer values. Decompression retains the input dtype, +while the target parameter controls storage rate. + +![EntroPack encoding and decoding pipeline](../../assets/entropack-pipeline.png) + +EntroPack's encoding and decoding pipeline, illustrated with a weight matrix. The upper panel +shows rate search and encoding; the lower panel shows fused GPU decoding and reconstruction. + +## Lattice quantization and integer fields + +Rows can differ substantially in numerical scale. The encoder first normalizes each row +by its root mean square, then groups the normalized values into eight-dimensional vectors. +Each vector is approximated by its nearest point on a scaled E8 lattice, a regular +arrangement of points in eight dimensions. A shared quantization scale controls the spacing +between these points. Finer spacing generally reduces reconstruction error but requires +more bits to describe the selected points. + +E8 has integer-coordinate and half-integer-coordinate subsets, called cosets. EntroPack +represents each point by its coset and eight invertible integer fields, using the lattice's +parity constraint to compact the final coordinate. The probability model conditions each +coordinate field on the coset, capturing differences between the two subsets. Frequently +occurring field values can then be encoded with fewer bits on average. The fields can +later be inverted arithmetically without a reconstruction codebook. + +## Rate selection and refinement + +The encoder searches for a quantization scale using sampled rows. For each candidate scale, +it estimates the coded field size and the metadata needed for decoding. This avoids +repeatedly producing a full compressed stream during the search. Once the scale is selected, +the encoder quantizes the full tensor and fits a reconstruction scale for each row by +least squares. + +Optional per-row rate–distortion refinement compares several resolutions for each row and +allocates them under an estimated storage budget. It alternates candidate selection with +updates to the shared probability model. This adds encoding work and is disabled by default. + +## Encoding and reconstruction + +The selected fields are encoded with rANS in independently decodable tiles. +The representation also stores the probability tables, row scales, and tile metadata. +Decoding recovers the fields, reconstructs lattice points, and applies the row scales in +a fused GPU operation. The result has the input's shape and dtype. + +Quantization and conversion back to the output dtype determine reconstruction error. +Entropy coding itself preserves the selected fields exactly. The achieved bitrate can +differ from the requested target because scale selection uses a size estimate. The +`actual_bpp` property reports the actual stored bytes, including metadata, divided by the +element count and multiplied by eight. + +## Usage + +The example targets 3.5 bits per element and measures relative L2 error against the input. +Search and refinement settings are described in [Config](../Usage/Configuration.md). + +```python +import torch +import entropack as ep + +tensor = (torch.randn(256, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.LatticeRANSConfig(target_bpp=3.5) + +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) + +reference = tensor.float() +relative_l2 = (restored.float() - reference).norm() / reference.norm() +assert restored.shape == tensor.shape +assert restored.dtype == tensor.dtype +print(f"Target: {config.target_bpp:.2f} bits per element") +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +print(f"Relative L2 error: {100 * relative_l2.item():.2f}%") +``` diff --git a/docs/en/Principles/Tile-ANS.md b/docs/en/Principles/Tile-ANS.md new file mode 100644 index 0000000..b42872e --- /dev/null +++ b/docs/en/Principles/Tile-ANS.md @@ -0,0 +1,57 @@ +# Tile-ANS + +Tile-ANS compresses tensors losslessly by encoding their storage bytes. It supports +floating-point and integer tensors, including BF16, FP16, FP32, FP8, and INT8. Because it +works on bit representations rather than numerical approximations, decompression restores +the original values exactly. + +## Byte streams and probability tables + +Different byte positions within a numerical format often have different distributions. +Tile-ANS therefore groups bytes by their position within each element. A two-byte format +produces two streams, while a four-byte format produces four. Each stream collects the +corresponding byte from every tensor element. + +The encoder counts byte frequencies separately for these streams and builds a probability +table for each. A skewed distribution can be encoded compactly because common bytes receive +shorter representations on average. When a stream offers little benefit after accounting +for coding overhead, it is stored directly. A single tensor can therefore contain both +entropy-coded streams and directly stored streams. + +## Tiled encoding and decoding + +Each stream is divided into independently decodable tiles. Entropy-coded tiles use range +asymmetric numeral systems (rANS), which encode symbols through reversible integer-state +updates. Multiple interleaved states allow symbols within a tile to be decoded in parallel, +while separate tiles provide additional parallel work. All tiles of a stream share its +probability table, avoiding a separate table for every tile. + +The decoder uses the same probability tables to reverse the state updates and recover +each coded byte stream. It then combines the decoded and directly stored streams, placing +their bytes back into the original positions within the tensor elements. + +No quantization is performed. Storage depends on the byte distributions and metadata, +so lossless compression does not provide a chosen target bitrate or guarantee a smaller +representation for every input. Larger tiles reduce metadata per element, while smaller +tiles expose more independent decoding tasks. + +## Usage + +The example compresses an FP16 tensor and checks its original bits after decompression. +Tile and probability-table settings are described in [Config](../Usage/Configuration.md). + +```python +import torch +import entropack as ep + +tensor = (torch.randn(256, 256, device="cuda") * 0.02).to(torch.float16) +config = ep.TileANSConfig() + +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) + +assert restored.shape == tensor.shape +assert restored.dtype == tensor.dtype +assert torch.equal(restored.view(torch.uint8), tensor.view(torch.uint8)) +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +``` diff --git a/docs/en/Usage/Configuration.md b/docs/en/Usage/Configuration.md new file mode 100644 index 0000000..56b0fe9 --- /dev/null +++ b/docs/en/Usage/Configuration.md @@ -0,0 +1,100 @@ +# Compression configuration + +A Config selects the compression scheme and its encoding and decoding parameters. +`execution_backend` defaults to `"auto"`, which selects the execution backend automatically. + +## Choose a configuration + +| Config | Compression | Input requirements | +| --- | --- | --- | +| `DFloat11Config()` | Lossless BF16 compression | BF16 tensors | +| `TileANSConfig()` | Lossless compression for multiple dtypes | Supported tensor dtypes | +| `LatticeRANSConfig(target_bpp=...)` | Lossy compression at a target bitrate | Nonempty, finite, two-dimensional tensors | + +Both lossless schemes reproduce every input bit. Their compressed size depends on the tensor's +data distribution. `TileANSConfig` also supports BF16, so either lossless scheme can be used for +that dtype. `LatticeRANSConfig` accepts targets from 1 to 11 bits per element, including +non-integer values. Its actual stored rate is available through `CompressedTensor.actual_bpp`. + +## Supported tensor dtypes + +| Tensor dtype | `DFloat11Config` | `TileANSConfig` | `LatticeRANSConfig` | +|---|---|---|---| +| `float32` | — | lossless | lossy | +| `float16` | — | lossless | lossy | +| `bfloat16` | lossless | lossless | lossy | +| `float8_e4m3fn` | — | lossless | lossy | +| `float8_e4m3fnuz` | — | lossless | lossy | +| `float8_e5m2` | — | lossless | lossy | +| `float8_e5m2fnuz` | — | lossless | lossy | +| `int64` | — | lossless | lossy | +| `int32` | — | lossless | lossy | +| `int16` | — | lossless | lossy | +| `int8` | — | lossless | lossy | +| `uint64` | — | lossless | lossy | +| `uint32` | — | lossless | lossy | +| `uint16` | — | lossless | lossy | +| `uint8` | — | lossless | lossy | +| `bool` | — | lossless | lossy | + +The lossless schemes accept tensors of different shapes, while lattice quantization requires +two-dimensional input. BF16, FP16, and FP8 refer to the corresponding PyTorch dtypes above. +Packed four-bit formats, FP64, and complex dtypes are not supported. + +## Parameter conventions + +Changes to encode settings affect subsequent compression, not existing compressed data. +Decode settings take effect when restoring a tensor. Most settings can retain their defaults. For lossy compression, `target_bpp` controls the +storage rate and a positive `row_rdo_iterations` enables per-row rate–distortion optimization (RDO). + +## CompressionConfig + +These execution settings apply to all three configuration classes. + +| Field | Type | Default | Stage | Meaning | +| --- | --- | --- | --- | --- | +| `execution_backend` | `str` or `None` | `"auto"` | Both | `"auto"` or `None` prefers CUDA and falls back to the PyTorch implementation if unavailable. `"cuda"` requires CUDA. `"eager"` selects the PyTorch fallback for tensor encoding and decoding. | + +## DFloat11Config + +Lossless BF16 compression. + +| Field | Type | Default | Stage | Meaning | +| --- | --- | --- | --- | --- | +| `execution_backend` | `str` or `None` | `"auto"` | Both | Backend selection as above | +| `bytes_per_thread` | Positive `int` or `None` | `16` | Encode | Encoded bytes processed per thread. Affects compression ratio and decoding parallelism. | +| `threads_per_block` | Positive `int` or `None` | `128` | Encode | Threads per block during encoding | + +## TileANSConfig + +Lossless compression of the supported tensor dtypes. + +| Field | Type | Default | Stage | Meaning | +| --- | --- | --- | --- | --- | +| `execution_backend` | `str` or `None` | `"auto"` | Both | Backend selection as above | +| `tile_elements` | `int` in [0, 2^31 − 1] | `0` | Encode | Elements per compressed tile. Affects compression ratio and decoding parallelism. `0` selects automatically. | +| `probability_bits` | `0`, `9`, `10`, `11`, `12` | `0` | Encode | Probability-table precision. `0` selects automatically. | +| `raw_lane_threshold` | `float` in [0, 8] | `7.9` | Encode | Threshold for storing hard-to-compress data directly, measured in estimated encoded bits per input byte. | +| `threads_per_block` | Positive `int` or `None` | `None` | Both | GPU block width. `None` selects automatically. | + +## LatticeRANSConfig + +Lossy compression of finite, two-dimensional tensors. + +| Field | Type | Default | Stage | Meaning | +| --- | --- | --- | --- | --- | +| `execution_backend` | `str` or `None` | `"auto"` | Both | Backend selection as above | +| `target_bpp` | `float` in [1, 11] | `4.0` | Encode | Target bits per input element. Non-integer targets are supported. Inspect `actual_bpp` for the stored rate. | +| `prob_bits` | `int` in [9, 15], `0`, or `None` | `None` | Encode | Probability-table precision. `None` or `0` selects automatically. | +| `tile_elements` | Positive `int` or `None` | `None` | Encode | Elements per compressed tile. Affects compression ratio and decoding parallelism. `None` selects automatically. | +| `row_rdo_iterations` | `int` in [0, 8] | `0` | Encode | Per-row rate–distortion refinement sweeps. `0` disables refinement. More sweeps increase compression time. | +| `row_rdo_candidates` | Positive `int` | `5` | Encode | Number of candidate quantizations per row for RDO. More candidates increase compression time. | +| `scale_search_iterations` | Positive `int` | `12` | Encode | Number of quantization-scale search iterations | +| `scale_search_max_vectors` | Positive `int` | `262144` | Encode | Sample limit for rate search, in vectors of eight elements | +| `threads_per_block` | Positive `int` or `None` | `None` | Decode | GPU block width. `None` selects automatically. | +| `l2_prefetch` | `bool` | `True` | Decode | Enable GPU L2 cache prefetching during decoding | + +## Compression principles + +The encoding and decoding processes are described in [DFloat11](../Principles/DFloat11.md), +[Tile-ANS](../Principles/Tile-ANS.md), and [Lattice-rANS](../Principles/Lattice-rANS.md). diff --git a/docs/en/Usage/Linear-layers.md b/docs/en/Usage/Linear-layers.md new file mode 100644 index 0000000..916684e --- /dev/null +++ b/docs/en/Usage/Linear-layers.md @@ -0,0 +1,163 @@ +# Compressed Linear usage + +Compressed Linear replaces a PyTorch linear layer with one that uses compressed weights. +Call `layer(x)` as usual: the input's last dimension changes from `in_features` to +`out_features`, and the other dimensions stay the same. A [Config](Configuration.md) +selects the compression scheme and its parameters. + +Compressed Linear requires a CUDA GPU and the matching CuPy package. See +[Quick start](Quick-start.md) for installation. + +`CompressedLinear` accepts `DFloat11Config`, `TileANSConfig`, or `LatticeRANSConfig`. +The selected scheme must support the weight dtype. + +## Replace an existing layer + +`from_linear` compresses an existing layer's weights and returns a new layer with a copy +of the original bias. For pretrained models, load the checkpoint before calling this method: + +```python +import torch +import entropack as ep + +linear = torch.nn.Linear(256, 256, dtype=torch.bfloat16, device="cuda") +config = ep.LatticeRANSConfig(target_bpp=4.0) +layer = ep.CompressedLinear.from_linear(linear, config=config) +x = torch.randn(8, 256, dtype=torch.bfloat16, device="cuda") + +with torch.inference_mode(): + output = layer(x) +print(output.shape, f"{layer.compressed_bits:.2f} bits per weight") +``` + +Assign the returned layer to the corresponding module attribute to use it in the model. + +## Replace several layers in a model + +This example replaces ordinary linear layers recursively and keeps the output layer in its +original format. Names in `skip` are module paths, as reported by `named_modules()`. + +```python +import torch +import entropack as ep + +model = torch.nn.Sequential( + torch.nn.Linear(256, 256), + torch.nn.GELU(), + torch.nn.Linear(256, 64), +).to(device="cuda", dtype=torch.bfloat16).eval() +config = ep.LatticeRANSConfig(target_bpp=4.0) + + +def compress_linears(module, config, skip=(), prefix=""): + for name, child in list(module.named_children()): + path = f"{prefix}.{name}" if prefix else name + if path in skip: + continue + if type(child) is torch.nn.Linear: + replacement = ep.CompressedLinear.from_linear(child, config=config) + setattr(module, name, replacement.train(child.training)) + else: + compress_linears(child, config, skip, path) + + +compress_linears(model, config, skip={"2"}) +x = torch.randn(8, 256, device="cuda", dtype=torch.bfloat16) +with torch.inference_mode(): + output = model(x) +print(output.shape, type(model[0]).__name__, type(model[2]).__name__) +``` + +The example selects standard `torch.nn.Linear` layers. Custom linear classes or shared +weights may need model-specific handling. Omit `skip` to compress every ordinary linear +layer. If `CompressedLinear` cannot compress a layer's weights, replacement raises an +error; use `skip` to keep that layer in its original form. + +## Combine compression with FP8 or INT8 computation + +| Layer | Weight format | Computation | +| --- | --- | --- | +| `CompressedLinear` | Input weight dtype | Standard linear operation, using the activation dtype | +| `CompressedFP8Linear` | FP8 E4M3FN codes | FP8 weights and activations, requires a CUDA GPU with SM8.9 or later | +| `CompressedINT8Linear` | INT8 codes | INT8 weights and activations, requires a CUDA GPU with SM8.0 or later | + +For `CompressedFP8Linear` and `CompressedINT8Linear`, use `config=None` for FP8 or INT8 +quantization alone, or pass `LatticeRANSConfig(target_bpp=...)` to apply further lossy +compression to the quantized weights. The target must be at least 1 bpp and less than 8 bpp. + +```python +import torch +import entropack as ep + +linear = torch.nn.Linear(256, 256, dtype=torch.bfloat16, device="cuda") +layer = ep.CompressedINT8Linear.from_linear( + linear, config=ep.LatticeRANSConfig(target_bpp=4.0) +) +x = torch.randn(32, 256, dtype=torch.bfloat16, device="cuda") +with torch.inference_mode(): + output = layer(x) +print(output.shape, layer.container_dtype, f"{layer.compressed_bits:.2f} bits per weight") +``` + +Use `CompressedFP8Linear` in the same pattern for FP8, on supported hardware. +To inspect the weights, `codes()` returns their FP8 or INT8 values, and `dequantize()` +returns their floating-point values after dequantization. + +## Measure storage + +`stored_nbytes` reports the compressed weight size, including metadata and FP8 or INT8 quantization scales. +`compressed_bits` is `8 * stored_nbytes / (in_features * out_features)`. +For multiple layers, sum stored bytes and weight elements before computing the ratio. +Biases are separate from this weight-storage measure. + +This measures weight storage, not peak inference memory. + +## Save and load a model + +Save the model's `state_dict`, then construct a model with the same architecture and +Compressed Linear classes before loading it. The following example compresses two layers +and restores their saved weights into a fresh model: + +```python +from pathlib import Path + +import torch +import entropack as ep + +config = ep.LatticeRANSConfig(target_bpp=4.0) +model = torch.nn.Sequential( + torch.nn.Linear(256, 256), + torch.nn.GELU(), + torch.nn.Linear(256, 64), +).to(device="cuda", dtype=torch.bfloat16) +for index in (0, 2): + model[index] = ep.CompressedLinear.from_linear(model[index], config=config) +model.eval() + +x = torch.randn(8, 256, device="cuda", dtype=torch.bfloat16) +with torch.inference_mode(): + expected = model(x) + +path = Path("compressed_model.pt") +torch.save(model.state_dict(), path) + +restored = torch.nn.Sequential( + ep.CompressedLinear(256, 256, config=config, device="cuda", dtype=torch.bfloat16), + torch.nn.GELU(), + ep.CompressedLinear(256, 64, config=config, device="cuda", dtype=torch.bfloat16), +).eval() +state = torch.load(path, map_location="cuda", weights_only=True) +restored.load_state_dict(state) +with torch.inference_mode(): + actual = restored(x) + +assert torch.allclose(actual, expected) +print(actual.shape) +``` + +The example saves `compressed_model.pt` in the current directory. Change the path as needed. +To load an existing checkpoint, construct the `restored` model, then call `torch.load` and +`load_state_dict`. +Keep the model architecture, layer names and classes, compression configurations, weight +dtypes, and library version with the checkpoint. Ordinary `torch.nn.Linear` layers +cannot load Compressed Linear checkpoints directly. diff --git a/docs/en/Usage/Quick-start.md b/docs/en/Usage/Quick-start.md new file mode 100644 index 0000000..e049e3f --- /dev/null +++ b/docs/en/Usage/Quick-start.md @@ -0,0 +1,84 @@ +# Quick start + +This guide covers installation, tensor compression and decompression, and basic Compressed Linear usage. + +## Installation + +Python 3.10 or later is required. First install a CUDA-enabled build of PyTorch 2.10 or later +for your environment. + +### Install from source (recommended) + +```bash +git clone https://github.com/modelscope/entropack.git +cd entropack +pip install -e ".[cuda13]" +``` + +### Install from PyPI + +PyPI releases may lag behind source updates. Install from source for the latest features. + +```bash +pip install "entropack[cuda13]" +``` + +Both installation methods above use CUDA 13 and include the matching CuPy package. +For CUDA 12, replace `cuda13` with `cuda12` in either command. If a compatible CuPy is +already installed, use `pip install -e .` for source installation or `pip install entropack` for PyPI. + +Select PyTorch's CUDA variant when installing PyTorch. +Optional Triton kernels for INT8 computation can be enabled with `pip install triton`. +See [Compressed Linear usage](Linear-layers.md) for additional FP8 and INT8 hardware requirements. +The first call may be slower while CUDA kernels compile. Measure performance after warm-up. + +## Direct tensor compression + +This example selects lossy compression with `LatticeRANSConfig(target_bpp=3.5)`, compresses +a 2D BF16 tensor at a target of 3.5 bits per element, and restores its original shape and dtype. + +```python +import torch +import entropack as ep + +tensor = (torch.randn(256, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.LatticeRANSConfig(target_bpp=3.5) + +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) + +print(f"Target: {config.target_bpp:.2f} bits per element") +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +print(restored.shape, restored.dtype) +``` + +`target_bpp` is measured in bits per element (bpp) and accepts integer or non-integer +values from 1 to 11. `actual_bpp` includes metadata and can differ from the target, +especially for small tensors. This scheme requires a nonempty 2D input without NaN or infinite values. + +See [Tensor compression](Tensor-compression.md) for reconstruction error, +device transfers, and saving and loading. [Compression configuration](Configuration.md) +covers lossless schemes and the complete parameter reference. + +## Use Compressed Linear + +`CompressedLinear.from_linear` compresses an existing layer's weights and returns a new +layer with a copy of the original bias: + +```python +import torch +import entropack as ep + +linear = torch.nn.Linear(256, 256, dtype=torch.bfloat16, device="cuda") +config = ep.LatticeRANSConfig(target_bpp=4.0) +layer = ep.CompressedLinear.from_linear(linear, config=config) +x = torch.randn(8, 256, dtype=torch.bfloat16, device="cuda") + +with torch.inference_mode(): + output = layer(x) +print(output.shape, f"{layer.compressed_bits:.2f} bits per weight") +``` + +For pretrained models, load the checkpoint before replacing the corresponding layers. +[Compressed Linear usage](Linear-layers.md) +covers replacing multiple layers, low-precision computation, and saving and loading compressed checkpoints. diff --git a/docs/en/Usage/Tensor-compression.md b/docs/en/Usage/Tensor-compression.md new file mode 100644 index 0000000..4fc6a35 --- /dev/null +++ b/docs/en/Usage/Tensor-compression.md @@ -0,0 +1,117 @@ +# Tensor compression + +Use `compress(tensor, config)` to compress a weight or other tensor into a +`CompressedTensor`, then `decompress(compressed, config)` to restore it. +A lossless scheme preserves every input bit; a lossy scheme returns an approximation. + +## Compress and restore + +By default, the decompressed tensor has the same shape, `dtype`, and `device` as the input. +To select a different dtype or device, convert the input with `tensor.to(dtype=..., device=...)` +before compression. + +This example compresses a tensor at a 4 bpp target and measures the relative L2 error +after decompression. Encoding and decoding use a configuration from the same scheme: + +```python +import torch +import entropack as ep + +tensor = (torch.randn(128, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.LatticeRANSConfig(target_bpp=4.0) +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) +relative_error = (restored.float() - tensor.float()).norm() / tensor.float().norm() + +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +print(f"Relative L2 error: {100 * relative_error:.2f}%") +``` + +To change the bitrate, compress the source tensor again with a new `target_bpp`. + +## Inspect stored size + +```python +import torch +import entropack as ep + +tensor = (torch.randn(128, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.LatticeRANSConfig(target_bpp=4.0) +compressed = ep.compress(tensor, config) + +print(compressed.shape, compressed.dtype, compressed.compress_method) +print(f"Stored: {compressed.storage_nbytes()} bytes") +print(f"Rate: {compressed.actual_bpp:.2f} bits per element") +``` + +`storage_nbytes()` reports the compressed size in bytes, including metadata needed for decompression. +`actual_bpp` is `8 * storage_nbytes() / tensor.numel()`. These measure the compressed +result, not the size of a checkpoint file or peak runtime memory. + +The achieved bitrate can differ from `target_bpp`, particularly for small tensors. +Compare the stored rate and reconstruction error when selecting a target. + +`CompressedTensor` also exposes `shape`, `dtype`, `compress_method`, and `lossless`. +The [API reference](../API_Reference/index.md) describes its remaining properties. + +## Move a compressed tensor + +`compressed.to(device)` returns a compressed tensor on the requested device. This example +moves it to CPU for storage, then back to the GPU for decompression: + +```python +import torch +import entropack as ep + +tensor = (torch.randn(128, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.DFloat11Config() +compressed = ep.compress(tensor, config) +cpu_copy = compressed.to("cpu") +gpu_copy = cpu_copy.to("cuda") +restored = ep.decompress(gpu_copy, config) +print(restored.device, restored.dtype) +``` + +Decompressing the object returned by `compressed.to(device)` restores the tensor on that device. +`compressed.to(dtype=...)` changes the decompressed output dtype without recompression. + +## Save and load + +Save a compressed tensor's `state_dict()`, load it with +`torch.load(..., weights_only=True)`, and restore the `CompressedTensor` with +`CompressedTensor.from_state_dict()`: + +```python +from pathlib import Path + +import torch +import entropack as ep + +tensor = (torch.randn(128, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.DFloat11Config() +compressed = ep.compress(tensor, config) + +path = Path("compressed_tensor.pt") +torch.save(compressed.state_dict(), path) +state = torch.load(path, map_location="cuda", weights_only=True) +loaded = ep.CompressedTensor.from_state_dict(state) + +restored = ep.decompress(loaded, config) +assert torch.equal(restored.view(torch.uint8), tensor.view(torch.uint8)) +``` + +The example saves `compressed_tensor.pt` in the current directory. Change the path as needed. +`map_location` selects the device for loading. Keep the scheme configuration and library version +with the checkpoint. + +## Input requirements + +The lattice scheme requires a nonempty, finite, two-dimensional tensor. DFloat11 accepts +BF16, and Tile-ANS accepts the dtypes listed in [Config](Configuration.md). Lossless +schemes accept higher-dimensional tensors directly. For lattice compression, reshape +higher-dimensional data to 2D before compression and restore its outer shape after +decompression. + +If compression emits a warning, check `compress_method`: a value of `"raw"` means the +tensor was stored without compression. Invalid configurations and unsupported backend +selections raise errors. diff --git a/docs/en/conf.py b/docs/en/conf.py new file mode 100644 index 0000000..f7094c0 --- /dev/null +++ b/docs/en/conf.py @@ -0,0 +1,43 @@ +# Configuration file for the Sphinx documentation builder. + +import tomllib +from pathlib import Path + +# -- Project information ----------------------------------------------------- + +project = "entropack" +copyright = "2026, EntroPack Authors" +author = "EntroPack Authors" +html_theme = "sphinx_rtd_theme" +language = "en" + + +def get_version() -> str: + pyproject = Path(__file__).resolve().parents[2] / "pyproject.toml" + with pyproject.open("rb") as handle: + return tomllib.load(handle)["project"]["version"] + + +version = get_version() +release = version + +# -- General configuration --------------------------------------------------- + +extensions = [ + "sphinx_markdown_tables", + "sphinx_copybutton", + "sphinx_rtd_theme", + "sphinx.ext.mathjax", + "myst_parser", +] + +source_suffix = [".rst", ".md"] +root_doc = "index" +exclude_patterns = ["build", "_build"] + +# -- Extension configuration ------------------------------------------------- + +copybutton_prompt_text = r">>> |\.\.\. " +copybutton_prompt_is_regexp = True +intersphinx_mapping = {"https://docs.python.org/": None} +myst_enable_extensions = ["amsmath", "dollarmath", "colon_fence"] diff --git a/docs/en/index.rst b/docs/en/index.rst new file mode 100644 index 0000000..fa7d536 --- /dev/null +++ b/docs/en/index.rst @@ -0,0 +1,27 @@ +EntroPack Documentation +============================================== + +General-purpose tensor compression for PyTorch, with lossless and adjustable lossy modes. + +.. toctree:: + :maxdepth: 2 + :caption: Usage + + Usage/Quick-start + Usage/Configuration + Usage/Tensor-compression + Usage/Linear-layers + +.. toctree:: + :maxdepth: 2 + :caption: API reference + + API_Reference/index + +.. toctree:: + :maxdepth: 2 + :caption: Compression principles + + Principles/DFloat11 + Principles/Tile-ANS + Principles/Lattice-rANS diff --git a/docs/requirements.txt b/docs/requirements.txt new file mode 100644 index 0000000..bfd20f4 --- /dev/null +++ b/docs/requirements.txt @@ -0,0 +1,7 @@ +docutils>=0.16.0 +myst_parser +sphinx>=5.3.0 +sphinx-copybutton +sphinx-rtd-theme +sphinx_markdown_tables +pymdown-extensions diff --git a/docs/zh/.readthedocs.yaml b/docs/zh/.readthedocs.yaml new file mode 100644 index 0000000..fbb7868 --- /dev/null +++ b/docs/zh/.readthedocs.yaml @@ -0,0 +1,17 @@ +# .readthedocs.yaml +# Read the Docs configuration file +# See https://docs.readthedocs.io/en/stable/config-file/v2.html for details + +version: 2 + +build: + os: ubuntu-22.04 + tools: + python: "3.11" + +sphinx: + configuration: docs/zh/conf.py + +python: + install: + - requirements: docs/requirements.txt diff --git a/docs/zh/API_Reference/index.md b/docs/zh/API_Reference/index.md new file mode 100644 index 0000000..4424166 --- /dev/null +++ b/docs/zh/API_Reference/index.md @@ -0,0 +1,154 @@ +# API 参考 + +常用函数与类均可通过 `import entropack as ep` 访问。 +完整用例见[通用张量压缩](../Usage/Tensor-compression.md)和 +[Compressed Linear 使用指南](../Usage/Linear-layers.md),配置参数见[压缩配置](../Usage/Configuration.md)。 + +## 张量编解码 + +### compress + +```text +compress(tensor: torch.Tensor, config: CompressionConfig) -> CompressedTensor +``` + +按 `config` 指定的方案压缩 `tensor`,返回 `CompressedTensor`。 +解压后的张量默认与输入具有相同的形状和数据类型。 + +输入要求由方案决定:DFloat11 接受 BF16 张量,Tile-ANS 支持多种数据类型,格量化要求非空、有限值组成的二维张量。 + +| 参数 | 含义 | +| --- | --- | +| `tensor` | 待压缩的 PyTorch 张量 | +| `config` | 必传,使用 `DFloat11Config`、`TileANSConfig` 或 `LatticeRANSConfig` 选择方案 | + +部分编码失败会发出说明原因的警告,并返回 `compress_method == "raw"` 的未压缩容器。 +无效配置或后端选择失败会直接报错。 + +### decompress + +```text +decompress(compressed: CompressedTensor, config: CompressionConfig) -> torch.Tensor +``` + +返回形状为 `compressed.shape`、数据类型为 `compressed.dtype`、设备为 `compressed.device` 的张量。 +未转换输出类型时,无损方案逐位恢复输入值,有损方案返回近似重建。 + +| 参数 | 含义 | +| --- | --- | +| `compressed` | `compress` 生成或从检查点加载的容器 | +| `config` | 对应压缩方案的配置,解码参数控制恢复过程 | + +解码时修改 `target_bpp` 等编码参数不会改变已保存的数据或重新量化张量。 + +## CompressedTensor + +保存一个张量压缩表示的 `torch.Tensor` 子类。通常由 `compress` 返回,或由 `from_state_dict` 从检查点恢复。 +进行数值计算前,需先使用 `decompress` 解压。 + +### 常用属性 + +| 属性 | 类型 | 含义 | +| --- | --- | --- | +| `shape` | `torch.Size` | 原始张量的形状 | +| `dtype` | `torch.dtype` | 解压后的数据类型 | +| `encoded_dtype` | `torch.dtype` | 编码时的数据类型 | +| `compress_method` | `str` | 容器实际使用的压缩方案 | +| `lossless` | `bool` | 该方案是否无损 | +| `actual_bpp` | `float` | 每元素实际存储比特数,包含元数据 | + +`actual_bpp = 8 * storage_nbytes() / math.prod(shape)`。 +这项指标衡量压缩表示的大小,不等于检查点文件大小或运行时显存占用。 + +### 常用方法 + +| 方法 | 返回值 | 含义 | +| --- | --- | --- | +| `to(...)` | `CompressedTensor` | 改变设备或解压输出类型,不重新压缩;`copy=True` 可复制存储 | +| `storage_nbytes(include_header=True)` | `int` | 压缩结果的总字节数,`include_header=False` 时不计容器头部 | +| `state_dict(prefix="")` | `dict[str, torch.Tensor]` | 将压缩张量导出为可保存的字典 | +| `CompressedTensor.from_state_dict(state, prefix="")` | `CompressedTensor` | 从上述字典恢复容器,不重新压缩 | + +保存与加载的 `prefix` 必须一致。可用 `torch.save` 保存字典,并用 +`torch.load(..., weights_only=True)` 加载,通过 `map_location` 指定恢复后的设备。 +加载后恢复编码时的数据类型;需要其他输出类型时,再调用 `.to(dtype=...)`。 + +## CompressedLinear + +每次前向调用使用重建权重执行线性运算的层。权重以压缩形式保存,偏置不压缩。 +运行需要 CUDA GPU 和对应版本的 CuPy。 + +### 创建层 + +```text +CompressedLinear(in_features, out_features, bias=True, *, + config=None, device=None, dtype=torch.bfloat16) +CompressedLinear.from_linear(linear, **kwargs) -> CompressedLinear +``` + +| 构造参数 | 含义 | +| --- | --- | +| `in_features` / `out_features` | 输入与输出特征数 | +| `bias` | 是否包含偏置 | +| `config` | 权重压缩配置。默认 `None` 为 BF16 选择 DFloat11,为其他支持的数据类型选择 Tile-ANS | +| `device` | 直接构造时偏置所在的设备 | +| `dtype` | 压缩权重和初始化偏置时使用的数据类型 | + +`from_linear` 返回一个新层,压缩源层的权重并复制偏置。源权重必须已加载,不能位于 `meta` 设备。 +常用调用为 `ep.CompressedLinear.from_linear(linear, config=config)`。 +`kwargs` 可指定 `config` 或 `dtype`,其中 `dtype` 默认沿用源权重类型,设备自动沿用源层。 + +直接调用构造函数会创建尚无权重数据的层,需要再调用 `compress_weight` 或加载检查点后才能推理。 + +### 常用方法与属性 + +| 接口 | 返回值 | 含义 | +| --- | --- | --- | +| `compress_weight(weight)` | `None` | 初始化层内压缩权重,形状应为 `(out_features, in_features)` | +| `dequantize(device=None)` | `torch.Tensor` | 返回数据类型为 `weight.dtype` 的稠密权重;未指定 `device` 时位于层所在设备 | +| `forward(x)` | `torch.Tensor` | 对形状为 `(..., in_features)` 的输入 `x` 执行线性运算,返回形状为 `(..., out_features)` 的张量 | +| `weight` | `CompressedTensor` | 层持有的冻结压缩权重参数 | +| `container_dtype` | `torch.dtype` | 压缩容器中权重或量化码的数据类型 | +| `stored_nbytes` | `int` | 权重存储字节数,含元数据和低精度层的量化尺度,不含偏置 | +| `compressed_bits` | `float` | `8 * stored_nbytes / (in_features * out_features)` | + +通过 `layer(x)` 调用前向运算。`.weight` 为压缩张量,需要数值权重时使用 `dequantize()`。 + +使用标准 `state_dict()` / `load_state_dict()` 保存与恢复层状态。 +加载前须创建相同层类型、形状、容器数据类型和压缩方案的层,Config 对象本身不会保存在检查点中。 +`.to(device)` 可迁移层,模型的数据类型转换不会重新编码已压缩的权重。 + +## CompressedFP8Linear 与 CompressedINT8Linear + +权重和激活均使用 FP8 或 INT8 的线性层,沿用 `CompressedLinear` 的构造参数、 +`from_linear`、存储属性和检查点接口。输入形状为 `(..., in_features)` 时, +输出形状为 `(..., out_features)`,数据类型和设备与输入一致。 + +| 类 | 权重与激活格式 | CUDA GPU 要求 | +| --- | --- | --- | +| `CompressedFP8Linear` | FP8 E4M3FN | SM8.9 及以上 | +| `CompressedINT8Linear` | INT8 | SM8.0 及以上 | + +量化码格式由层类决定,构造参数 `dtype` 不改变 FP8 或 INT8 格式。 + +`config=None` 时直接保存量化码。指定 `LatticeRANSConfig` 时进一步进行有损压缩, +目标码率须满足 `1 <= target_bpp < 8`。`stored_nbytes` 包含重建权重所需的逐行量化尺度。 + +| 方法 | 返回值 | 含义 | +| --- | --- | --- | +| `codes(device=None)` | FP8 或 INT8 张量 | 恢复量化码,尚未乘回行尺度 | +| `dequantize(device=None)` | `torch.Tensor` | 返回层初始化时的数据类型,`from_linear` 默认沿用原始权重类型 | + +两种方法未指定 `device` 时,返回张量均位于层所在设备。有损压缩后的量化码可能与初始量化结果不同。 + +## Config 类 + +Config 同时用于张量编解码和 Compressed Linear 的权重存储。以下三个类继承自 `CompressionConfig`: + +| 类 | 用途 | +| --- | --- | +| `DFloat11Config` | BF16 无损压缩 | +| `TileANSConfig` | 多种数据类型的分块 ANS 无损压缩 | +| `LatticeRANSConfig` | 以 `target_bpp` 控制码率的格量化有损压缩 | + +方案选择、默认值和参数范围见[压缩配置](../Usage/Configuration.md)。 diff --git a/docs/zh/Makefile b/docs/zh/Makefile new file mode 100644 index 0000000..45d227d --- /dev/null +++ b/docs/zh/Makefile @@ -0,0 +1,14 @@ +# Minimal makefile for Sphinx documentation + +SPHINXOPTS ?= +SPHINXBUILD ?= sphinx-build +SOURCEDIR = . +BUILDDIR = ../_build/zh + +help: + @$(SPHINXBUILD) -M help "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O) + +.PHONY: help Makefile + +%: Makefile + @$(SPHINXBUILD) -M $@ "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O) diff --git a/docs/zh/Principles/DFloat11.md b/docs/zh/Principles/DFloat11.md new file mode 100644 index 0000000..efe0dff --- /dev/null +++ b/docs/zh/Principles/DFloat11.md @@ -0,0 +1,45 @@ +# DFloat11 + +DFloat11 对 BF16 张量进行无损压缩,完整保留原始位表示。许多张量的指数值集中在较小的范围内, +因此可以用较短的编码表示常见指数,同时原样保存符号位和尾数部分。 + +## 编码 + +一个 BF16 数值包含 1 位符号、8 位指数和 7 位尾数。编码器将它拆成两部分: +指数单独组成符号序列,符号位和尾数则合并成一个字节直接存储。 +编码器统计当前张量的指数频率,再构建 Huffman 编码表。常见指数使用较短的编码, +不常见指数使用较长的编码,从而减少整个指数序列占用的空间。 + +Huffman 编码长度不固定,因此解码器无法从任意一位直接识别下一个符号。 +EntroPack 为编码区域保存起始位置和符号数量,使不同区域可以并行解码。 +这些入口信息和 Huffman 表构成压缩表示中的元数据开销。 + +## 解码与存储大小 + +解码器通过 Huffman 表恢复指数序列,再将每个指数与对应的符号位、尾数重新组合, +得到原始 BF16 位表示,并按照保存的形状组织为输出张量。整个过程不涉及数值量化或舍入。 + +实际存储大小取决于指数分布和解码所需的元数据。指数越集中,通常越容易压缩。 +对于较小的张量,元数据占比也会更高。`DFloat11Config` 不设置目标码率, +DFloat11 这一名称也不意味着所有张量都恰好以每元素 11 bit 存储。 + +## 使用示例 + +以下示例压缩一个 BF16 张量,并检查解压后的位表示是否与输入一致。 +编码区域相关参数见 [Config](../Usage/Configuration.md)。 + +```python +import torch +import entropack as ep + +tensor = (torch.randn(256, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.DFloat11Config() + +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) + +assert restored.shape == tensor.shape +assert restored.dtype == tensor.dtype +assert torch.equal(restored.view(torch.uint8), tensor.view(torch.uint8)) +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +``` diff --git a/docs/zh/Principles/Lattice-rANS.md b/docs/zh/Principles/Lattice-rANS.md new file mode 100644 index 0000000..adc81aa --- /dev/null +++ b/docs/zh/Principles/Lattice-rANS.md @@ -0,0 +1,64 @@ +# EntroPack 格量化压缩 + +EntroPack 的格量化方案结合有损向量量化与无损熵编码,压缩二维张量。 +`LatticeRANSConfig` 支持每元素 1 至 11 bit 的目标,包括非整数码率。 +解压后仍保留输入的数据类型,存储码率则通过目标参数调节。 + +![EntroPack 编码与解码流程](../../assets/entropack-pipeline.png) + +以权重矩阵为例的 EntroPack 编解码流程。上半部分为码率搜索与编码,下半部分为融合 GPU 解码与重建。 + +## 格量化与整数字段 + +不同张量行的数值尺度可能相差较大。编码器首先用每行的均方根归一化该行, +再将归一化后的数值每八个组成一个向量。E8 格是八维空间中按规则排列的一组点, +编码器用缩放后的格中最近的点近似每个向量。共享的量化尺度控制格点之间的间距。 +间距越小,通常重建误差越小,但描述所选格点需要的比特也越多。 + +E8 包含整数坐标与半整数坐标两类格点,对应两个陪集。 +EntroPack 用陪集标记和八个可逆整数字段表示格点,并利用奇偶约束压缩最后一个坐标的表示。 +概率模型根据陪集分别统计各坐标字段的分布,以捕捉两类格点的差异。 +频繁出现的字段值平均可用更少的比特表示。解码时,这些字段可通过算术运算还原为格点,无需重建码本。 + +## 码率选择与精度优化 + +编码器在采样行上搜索量化尺度,对每个候选尺度估计字段的编码大小及解码所需的元数据, +无需在搜索过程中反复生成完整压缩码流。选定尺度后,再量化整个张量, +并通过最小二乘拟合每行的重建尺度。 + +可选的逐行率失真优化会为每行比较多个量化精度,在估计的存储预算内分配候选。 +这一过程交替进行候选选择和共享概率模型更新,需要额外的编码计算,默认关闭。 + +## 编码与重建 + +选定的字段由 rANS 熵编码为可独立解码的 tile。 +压缩表示还保存概率表、行尺度和 tile 定位信息。 +解码时,GPU 在融合操作中恢复字段、重建格点并应用行尺度,输出与输入形状和 dtype 相同的张量。 + +重建误差来自量化以及转换回输出 dtype 时的舍入,熵编码本身完整保留选定的字段。 +由于尺度搜索使用大小估计,实际码率可能与目标有差别。 +`actual_bpp` 按包含元数据的实际存储字节数计算,即字节数乘以八,再除以张量元素数。 + +## 使用示例 + +以下示例以每元素 3.5 bit 为目标压缩张量,并以原始张量为参考计算相对 L2 误差。 +搜索与精度优化参数见 [Config](../Usage/Configuration.md)。 + +```python +import torch +import entropack as ep + +tensor = (torch.randn(256, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.LatticeRANSConfig(target_bpp=3.5) + +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) + +reference = tensor.float() +relative_l2 = (restored.float() - reference).norm() / reference.norm() +assert restored.shape == tensor.shape +assert restored.dtype == tensor.dtype +print(f"Target: {config.target_bpp:.2f} bits per element") +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +print(f"Relative L2 error: {100 * relative_l2.item():.2f}%") +``` diff --git a/docs/zh/Principles/Tile-ANS.md b/docs/zh/Principles/Tile-ANS.md new file mode 100644 index 0000000..0929666 --- /dev/null +++ b/docs/zh/Principles/Tile-ANS.md @@ -0,0 +1,48 @@ +# Tile-ANS + +Tile-ANS 通过编码张量的存储字节实现无损压缩,支持 BF16、FP16、FP32、FP8、INT8 等 +浮点和整数类型。它处理的是数值的位表示,不对数值进行近似,因此解压后能够完整恢复原始数据。 + +## 字节流与概率表 + +同一种数值格式中,不同字节位置的分布往往不同。Tile-ANS 按字节在元素内部的位置, +将张量拆成多个字节流。例如,两字节格式对应两个流,四字节格式对应四个流。 +每个流收集所有元素在对应位置上的字节。 + +编码器分别统计这些流的字节频率,为每个流构建概率表。 +当分布较集中时,常见字节平均使用更短的表示,从而减少存储空间。 +对于计入编码开销后压缩收益仍较小的流,编码器直接存储原始字节。 +因此,同一个张量中可以同时存在熵编码流和直接存储的流。 + +## 分块编码与解码 + +每个流进一步划分为可独立解码的 tile。需要熵编码的 tile 使用范围非对称数字系统 rANS, +通过可逆的整数状态更新编码符号。一个 tile 内部交错使用多个编码状态,支持并行恢复符号。 +不同 tile 之间也可以独立解码。同一字节流的所有 tile 共享概率表,无需逐 tile 存储一份表。 + +解码器使用相同的概率表逆转状态更新,恢复各个经过熵编码的字节流, +再将恢复出的字节流与直接存储的字节流合并,把字节放回元素内的原始位置,重建张量。 + +该过程不涉及量化。实际大小由字节分布和元数据共同决定,因此不能指定一个有损压缩式的目标码率, +也不保证每个输入都能缩小。较大的 tile 可以降低每元素的元数据开销,较小的 tile 则提供更多独立解码任务。 + +## 使用示例 + +以下示例压缩一个 FP16 张量,并检查解压后的位表示是否与输入一致。 +分块大小和概率表参数见 [Config](../Usage/Configuration.md)。 + +```python +import torch +import entropack as ep + +tensor = (torch.randn(256, 256, device="cuda") * 0.02).to(torch.float16) +config = ep.TileANSConfig() + +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) + +assert restored.shape == tensor.shape +assert restored.dtype == tensor.dtype +assert torch.equal(restored.view(torch.uint8), tensor.view(torch.uint8)) +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +``` diff --git a/docs/zh/Usage/Configuration.md b/docs/zh/Usage/Configuration.md new file mode 100644 index 0000000..fe14d54 --- /dev/null +++ b/docs/zh/Usage/Configuration.md @@ -0,0 +1,97 @@ +# 压缩配置 + +Config 选择压缩方案并设置编解码参数。`execution_backend` 默认为 `"auto"`,自动选择计算后端。 + +## 选择配置 + +| Config | 压缩方式 | 输入要求 | +| --- | --- | --- | +| `DFloat11Config()` | BF16 无损压缩 | BF16 张量 | +| `TileANSConfig()` | 多种数据类型的无损压缩 | 支持的数据类型 | +| `LatticeRANSConfig(target_bpp=...)` | 按目标码率进行有损压缩 | 非空、有限值组成的二维张量 | + +两个无损方案均逐位恢复输入,压缩后的大小取决于张量的数据分布。`TileANSConfig` 也支持 BF16, +因此 BF16 张量可以选择其中任一无损方案。`LatticeRANSConfig` 接受每元素 1–11 bit 的目标码率, +支持非整数值,实际存储码率可通过 `CompressedTensor.actual_bpp` 查看。 + +## 支持的张量数据类型 + +| 数据类型 | `DFloat11Config` | `TileANSConfig` | `LatticeRANSConfig` | +|---|---|---|---| +| `float32` | — | 无损 | 有损 | +| `float16` | — | 无损 | 有损 | +| `bfloat16` | 无损 | 无损 | 有损 | +| `float8_e4m3fn` | — | 无损 | 有损 | +| `float8_e4m3fnuz` | — | 无损 | 有损 | +| `float8_e5m2` | — | 无损 | 有损 | +| `float8_e5m2fnuz` | — | 无损 | 有损 | +| `int64` | — | 无损 | 有损 | +| `int32` | — | 无损 | 有损 | +| `int16` | — | 无损 | 有损 | +| `int8` | — | 无损 | 有损 | +| `uint64` | — | 无损 | 有损 | +| `uint32` | — | 无损 | 有损 | +| `uint16` | — | 无损 | 有损 | +| `uint8` | — | 无损 | 有损 | +| `bool` | — | 无损 | 有损 | + +无损方案接受多种形状的张量,格量化要求二维输入。文中的 BF16、FP16、FP8 对应上表列出的 +PyTorch 数据类型。打包的四比特格式、FP64 和复数类型不在支持范围内。 + +## 参数说明 + +修改编码参数只影响后续压缩,不会改变已有压缩结果。解码参数在恢复张量时生效。 +多数设置可保留默认值。 +有损压缩的存储码率由 `target_bpp` 控制,`row_rdo_iterations` 设为正数时启用逐行率失真优化(RDO)。 + +## CompressionConfig + +以下执行参数适用于三个配置类。 + +| 字段 | 类型 | 默认值 | 阶段 | 含义 | +| --- | --- | --- | --- | --- | +| `execution_backend` | `str` 或 `None` | `"auto"` | 编解码 | `"auto"` 或 `None` 优先选择 CUDA,不可用时回退到 PyTorch 实现;`"cuda"` 强制使用 CUDA。`"eager"` 为张量编解码的 PyTorch 后备实现。 | + +## DFloat11Config + +用于 BF16 无损压缩。 + +| 字段 | 类型 | 默认值 | 阶段 | 含义 | +| --- | --- | --- | --- | --- | +| `execution_backend` | `str` 或 `None` | `"auto"` | 编解码 | 后端选择,含义同上 | +| `bytes_per_thread` | 正整数或 `None` | `16` | 编码 | 每个线程处理的编码字节数,影响压缩率和解码并行度 | +| `threads_per_block` | 正整数或 `None` | `128` | 编码 | 编码时每个线程块的线程数 | + +## TileANSConfig + +用于支持的数据类型的无损压缩。 + +| 字段 | 类型 | 默认值 | 阶段 | 含义 | +| --- | --- | --- | --- | --- | +| `execution_backend` | `str` 或 `None` | `"auto"` | 编解码 | 后端选择,含义同上 | +| `tile_elements` | [0, 2^31 − 1] 内整数 | `0` | 编码 | 每个压缩块的元素数,影响压缩率和解码并行度。`0` 自动选择。 | +| `probability_bits` | `0`、`9`、`10`、`11`、`12` | `0` | 编码 | 概率表精度,`0` 自动选择 | +| `raw_lane_threshold` | [0, 8] 内浮点数 | `7.9` | 编码 | 决定何时直接存储难以压缩的数据,单位为每字节的预计编码比特数 | +| `threads_per_block` | 正整数或 `None` | `None` | 编解码 | GPU 线程块宽度,`None` 自动选择 | + +## LatticeRANSConfig + +用于有限值组成的二维张量的有损压缩。 + +| 字段 | 类型 | 默认值 | 阶段 | 含义 | +| --- | --- | --- | --- | --- | +| `execution_backend` | `str` 或 `None` | `"auto"` | 编解码 | 后端选择,含义同上 | +| `target_bpp` | [1, 11] 内浮点数 | `4.0` | 编码 | 每个输入元素的目标比特数,支持非整数。实际码率通过 `actual_bpp` 查看。 | +| `prob_bits` | [9, 15] 内整数、`0` 或 `None` | `None` | 编码 | 概率表精度,`None` 或 `0` 自动选择 | +| `tile_elements` | 正整数或 `None` | `None` | 编码 | 每个压缩块的元素数,影响压缩率和解码并行度。`None` 自动选择。 | +| `row_rdo_iterations` | [0, 8] 内整数 | `0` | 编码 | 逐行率失真优化的轮数,`0` 关闭。更多轮次会增加压缩耗时。 | +| `row_rdo_candidates` | 正整数 | `5` | 编码 | RDO 为每行比较的候选量化结果数,更多候选会增加压缩耗时 | +| `scale_search_iterations` | 正整数 | `12` | 编码 | 量化尺度搜索的迭代次数 | +| `scale_search_max_vectors` | 正整数 | `262144` | 编码 | 码率搜索的采样上限,每个向量包含八个元素 | +| `threads_per_block` | 正整数或 `None` | `None` | 解码 | GPU 线程块宽度,`None` 自动选择 | +| `l2_prefetch` | `bool` | `True` | 解码 | 解码时启用 GPU L2 缓存预取 | + +## 压缩原理 + +各方案的编解码过程分别见 [DFloat11](../Principles/DFloat11.md)、 +[Tile-ANS](../Principles/Tile-ANS.md) 和 [Lattice-rANS](../Principles/Lattice-rANS.md)。 diff --git a/docs/zh/Usage/Linear-layers.md b/docs/zh/Usage/Linear-layers.md new file mode 100644 index 0000000..c6c0f78 --- /dev/null +++ b/docs/zh/Usage/Linear-layers.md @@ -0,0 +1,155 @@ +# Compressed Linear 使用指南 + +Compressed Linear 将 PyTorch 线性层替换为使用压缩权重的层。 +仍通过 `layer(x)` 调用:输入的最后一维从 `in_features` 变为 `out_features`, +其余维度不变。[Config](Configuration.md) 指定压缩方案和参数。 + +Compressed Linear 需要 CUDA GPU 和对应版本的 CuPy,安装方式见[快速上手](Quick-start.md)。 + +`CompressedLinear` 可使用 `DFloat11Config`、`TileANSConfig` 或 `LatticeRANSConfig`, +所选方案需支持权重的数据类型。 + +## 替换一个已有层 + +`from_linear` 压缩现有层的权重并返回一个新层,偏置保持原样复制。 +对于预训练模型,应先加载检查点,再调用该方法: + +```python +import torch +import entropack as ep + +linear = torch.nn.Linear(256, 256, dtype=torch.bfloat16, device="cuda") +config = ep.LatticeRANSConfig(target_bpp=4.0) +layer = ep.CompressedLinear.from_linear(linear, config=config) +x = torch.randn(8, 256, dtype=torch.bfloat16, device="cuda") + +with torch.inference_mode(): + output = layer(x) +print(output.shape, f"{layer.compressed_bits:.2f} bits per weight") +``` + +将返回的层赋给模型中的对应属性后,即可使用压缩版本。 + +## 替换模型中的多个层 + +下面的例子递归替换普通线性层,并让输出层保留原来的格式。 +`skip` 使用 `named_modules()` 中的模块路径。 + +```python +import torch +import entropack as ep + +model = torch.nn.Sequential( + torch.nn.Linear(256, 256), + torch.nn.GELU(), + torch.nn.Linear(256, 64), +).to(device="cuda", dtype=torch.bfloat16).eval() +config = ep.LatticeRANSConfig(target_bpp=4.0) + + +def compress_linears(module, config, skip=(), prefix=""): + for name, child in list(module.named_children()): + path = f"{prefix}.{name}" if prefix else name + if path in skip: + continue + if type(child) is torch.nn.Linear: + replacement = ep.CompressedLinear.from_linear(child, config=config) + setattr(module, name, replacement.train(child.training)) + else: + compress_linears(child, config, skip, path) + + +compress_linears(model, config, skip={"2"}) +x = torch.randn(8, 256, device="cuda", dtype=torch.bfloat16) +with torch.inference_mode(): + output = model(x) +print(output.shape, type(model[0]).__name__, type(model[2]).__name__) +``` + +示例仅选择标准 `torch.nn.Linear`,自定义线性层或共享权重需要结合模型处理。 +省略 `skip` 即可压缩所有普通线性层。如果 `CompressedLinear` 无法压缩某层的权重, +替换时会报错;可通过 `skip` 让该层保留原始格式。 + +## 结合 FP8 或 INT8 计算 + +| 层 | 权重格式 | 计算方式 | +| --- | --- | --- | +| `CompressedLinear` | 输入权重的数据类型 | 使用激活数据类型进行普通线性运算 | +| `CompressedFP8Linear` | FP8 E4M3FN 量化码 | FP8 权重和激活,需 CUDA GPU(SM8.9 及以上) | +| `CompressedINT8Linear` | INT8 量化码 | INT8 权重和激活,需 CUDA GPU(SM8.0 及以上) | + +对于 `CompressedFP8Linear` 和 `CompressedINT8Linear`,`config=None` 仅做 FP8 或 INT8 量化, +传入 `LatticeRANSConfig(target_bpp=...)` 则会对量化后的权重进一步进行有损压缩。 +目标码率需大于等于 1 bpp 且低于 8 bpp。 + +```python +import torch +import entropack as ep + +linear = torch.nn.Linear(256, 256, dtype=torch.bfloat16, device="cuda") +layer = ep.CompressedINT8Linear.from_linear( + linear, config=ep.LatticeRANSConfig(target_bpp=4.0) +) +x = torch.randn(32, 256, dtype=torch.bfloat16, device="cuda") +with torch.inference_mode(): + output = layer(x) +print(output.shape, layer.container_dtype, f"{layer.compressed_bits:.2f} bits per weight") +``` + +在支持的硬件上,FP8 可以按同样方式使用 `CompressedFP8Linear`。 +需要查看权重时,`codes()` 返回 FP8 或 INT8 数值,`dequantize()` 返回反量化后的浮点数值。 + +## 统计存储 + +`stored_nbytes` 统计压缩权重的总字节数,包含元数据及 FP8、INT8 的量化尺度。 +`compressed_bits` 等于 `8 * stored_nbytes / (in_features * out_features)`。 +统计多个层时,应先分别累加字节数与权重元素数,再计算比例。偏置不计入这项权重存储指标。 + +这项指标衡量权重存储大小,不代表推理时的峰值显存。 + +## 保存与加载模型 + +保存模型的 `state_dict` 后,先构造具有相同结构、使用相同 Compressed Linear 类的模型,再加载状态。 +以下示例压缩两个层,并将保存的权重加载到一个新模型中: + +```python +from pathlib import Path + +import torch +import entropack as ep + +config = ep.LatticeRANSConfig(target_bpp=4.0) +model = torch.nn.Sequential( + torch.nn.Linear(256, 256), + torch.nn.GELU(), + torch.nn.Linear(256, 64), +).to(device="cuda", dtype=torch.bfloat16) +for index in (0, 2): + model[index] = ep.CompressedLinear.from_linear(model[index], config=config) +model.eval() + +x = torch.randn(8, 256, device="cuda", dtype=torch.bfloat16) +with torch.inference_mode(): + expected = model(x) + +path = Path("compressed_model.pt") +torch.save(model.state_dict(), path) + +restored = torch.nn.Sequential( + ep.CompressedLinear(256, 256, config=config, device="cuda", dtype=torch.bfloat16), + torch.nn.GELU(), + ep.CompressedLinear(256, 64, config=config, device="cuda", dtype=torch.bfloat16), +).eval() +state = torch.load(path, map_location="cuda", weights_only=True) +restored.load_state_dict(state) +with torch.inference_mode(): + actual = restored(x) + +assert torch.allclose(actual, expected) +print(actual.shape) +``` + +示例将检查点保存到当前目录的 `compressed_model.pt`,可按需修改路径。 +加载已有检查点时,构造 `restored` 模型,再调用 `torch.load` 和 `load_state_dict`。 +应随检查点保留模型结构、层名及类型、压缩配置、权重数据类型和库版本。 +普通 `torch.nn.Linear` 无法直接加载 Compressed Linear 的检查点。 diff --git a/docs/zh/Usage/Quick-start.md b/docs/zh/Usage/Quick-start.md new file mode 100644 index 0000000..2fda2c4 --- /dev/null +++ b/docs/zh/Usage/Quick-start.md @@ -0,0 +1,81 @@ +# 快速上手 + +本页介绍安装、张量编解码和 Compressed Linear 的基本用法。 + +## 安装 + +需要 Python 3.10 及以上,并先安装与环境匹配的 CUDA 版 PyTorch 2.10 及以上。 + +### 源码安装(推荐) + +```bash +git clone https://github.com/modelscope/entropack.git +cd entropack +pip install -e ".[cuda13]" +``` + +### 从 PyPI 安装 + +PyPI 版本更新可能有所延迟,如需最新功能,推荐从源码安装。 + +```bash +pip install "entropack[cuda13]" +``` + +上述两种安装方式均以 CUDA 13 为例,并包含对应版本的 CuPy。使用 CUDA 12 时, +将命令中的 `cuda13` 改为 `cuda12`。如果已安装匹配的 CuPy,源码安装和 PyPI 安装 +可分别使用 `pip install -e .` 和 `pip install entropack`。 + +PyTorch 的 CUDA 版本需在安装 PyTorch 时选定。 +INT8 计算可通过 `pip install triton` 启用可选的 Triton 内核。 +FP8 和 INT8 的额外硬件要求见 [Compressed Linear 使用指南](Linear-layers.md)。 +首次调用需要编译 CUDA 内核,可能比后续调用更慢,性能计时应在预热后进行。 + +## 直接压缩张量 + +以下示例使用 `LatticeRANSConfig(target_bpp=3.5)` 选择有损压缩,将二维 BF16 张量 +压缩到每元素 3.5 bit 的目标码率,再恢复为原来的形状和数据类型。 + +```python +import torch +import entropack as ep + +tensor = (torch.randn(256, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.LatticeRANSConfig(target_bpp=3.5) + +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) + +print(f"Target: {config.target_bpp:.2f} bits per element") +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +print(restored.shape, restored.dtype) +``` + +`target_bpp` 的单位为每元素比特数(bpp),可设为 1–11 范围内的整数或非整数值。 +`actual_bpp` 返回包含元数据的实际存储码率,可能与目标不同,尤其在张量较小时。 +该方案要求输入为非空的二维张量,且不含 NaN 或无穷值。 + +重建误差、设备迁移和保存加载见[通用张量压缩](Tensor-compression.md)。 +无损方案及完整参数见[压缩配置](Configuration.md)。 + +## 使用 Compressed Linear + +`CompressedLinear.from_linear` 压缩现有层的权重并返回一个新层,偏置保持原样复制: + +```python +import torch +import entropack as ep + +linear = torch.nn.Linear(256, 256, dtype=torch.bfloat16, device="cuda") +config = ep.LatticeRANSConfig(target_bpp=4.0) +layer = ep.CompressedLinear.from_linear(linear, config=config) +x = torch.randn(8, 256, dtype=torch.bfloat16, device="cuda") + +with torch.inference_mode(): + output = layer(x) +print(output.shape, f"{layer.compressed_bits:.2f} bits per weight") +``` + +对于预训练模型,应先加载检查点,再将模型中的对应层替换为新层。 +[Compressed Linear 使用指南](Linear-layers.md)介绍多个层的替换、 +低精度计算和压缩检查点的保存加载。 diff --git a/docs/zh/Usage/Tensor-compression.md b/docs/zh/Usage/Tensor-compression.md new file mode 100644 index 0000000..4963fd0 --- /dev/null +++ b/docs/zh/Usage/Tensor-compression.md @@ -0,0 +1,111 @@ +# 通用张量压缩 + +使用 `compress(tensor, config)` 将权重或其他张量压缩为 `CompressedTensor`, +再通过 `decompress(compressed, config)` 恢复张量。 +无损方案逐位还原输入,有损方案返回近似结果。 + +## 压缩与恢复 + +解压后的张量默认与输入张量具有相同的形状、`dtype` 和 `device`。 +如需指定数据类型或设备,在压缩前用 `tensor.to(dtype=..., device=...)` 转换输入即可。 + +以下示例使用 4 bpp 的目标码率压缩张量,并计算解压后的相对 L2 误差。 +压缩与解压使用同一方案的配置: + +```python +import torch +import entropack as ep + +tensor = (torch.randn(128, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.LatticeRANSConfig(target_bpp=4.0) +compressed = ep.compress(tensor, config) +restored = ep.decompress(compressed, config) +relative_error = (restored.float() - tensor.float()).norm() / tensor.float().norm() + +print(f"Stored: {compressed.actual_bpp:.2f} bits per element") +print(f"Relative L2 error: {100 * relative_error:.2f}%") +``` + +需要改变码率时,应使用新的 `target_bpp` 重新压缩源张量。 + +## 查看实际存储 + +```python +import torch +import entropack as ep + +tensor = (torch.randn(128, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.LatticeRANSConfig(target_bpp=4.0) +compressed = ep.compress(tensor, config) + +print(compressed.shape, compressed.dtype, compressed.compress_method) +print(f"Stored: {compressed.storage_nbytes()} bytes") +print(f"Rate: {compressed.actual_bpp:.2f} bits per element") +``` + +`storage_nbytes()` 统计压缩结果的总字节数,包含恢复张量所需的元数据, +`actual_bpp` 等于 `8 * storage_nbytes() / tensor.numel()`。 +它们衡量压缩结果的大小,不代表检查点文件大小或运行时峰值显存。 + +实际码率可能与 `target_bpp` 有所不同,小张量尤其如此。 +选择目标码率时,应同时比较实际存储和重建误差。 + +`CompressedTensor` 还提供 `shape`、`dtype`、`compress_method` 和 `lossless`。 +其他属性见 [API 参考](../API_Reference/index.md)。 + +## 移动压缩张量 + +`compressed.to(device)` 返回位于指定设备的压缩张量。以下示例先将其转存到 CPU, +再移回 GPU 解压: + +```python +import torch +import entropack as ep + +tensor = (torch.randn(128, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.DFloat11Config() +compressed = ep.compress(tensor, config) +cpu_copy = compressed.to("cpu") +gpu_copy = cpu_copy.to("cuda") +restored = ep.decompress(gpu_copy, config) +print(restored.device, restored.dtype) +``` + +对 `compressed.to(device)` 返回的压缩张量解压,得到的张量也位于该设备上。 +`compressed.to(dtype=...)` 可改变解压输出类型,无需重新压缩。 + +## 保存与加载 + +保存压缩张量的 `state_dict()`,用 `torch.load(..., weights_only=True)` 加载后, +通过 `CompressedTensor.from_state_dict()` 恢复 `CompressedTensor`: + +```python +from pathlib import Path + +import torch +import entropack as ep + +tensor = (torch.randn(128, 256, device="cuda") * 0.02).to(torch.bfloat16) +config = ep.DFloat11Config() +compressed = ep.compress(tensor, config) + +path = Path("compressed_tensor.pt") +torch.save(compressed.state_dict(), path) +state = torch.load(path, map_location="cuda", weights_only=True) +loaded = ep.CompressedTensor.from_state_dict(state) + +restored = ep.decompress(loaded, config) +assert torch.equal(restored.view(torch.uint8), tensor.view(torch.uint8)) +``` + +示例将检查点保存到当前目录的 `compressed_tensor.pt`,可按需修改路径。 +`map_location` 指定加载设备。建议随检查点保留对应的配置和库版本。 + +## 输入要求 + +格量化方案要求非空、有限值组成的二维张量。DFloat11 接受 BF16, +Tile-ANS 接受 [Config](Configuration.md) 中列出的数据类型。无损方案可直接接受高维张量。 +使用格量化压缩高维数据时,需先转换为二维布局,并在解压后还原原始形状。 + +压缩时若出现警告,可检查 `compress_method`:值为 `"raw"` 表示张量未经压缩就被保存。 +无效配置或不支持的后端选择会直接报错。 diff --git a/docs/zh/conf.py b/docs/zh/conf.py new file mode 100644 index 0000000..3ec7ef8 --- /dev/null +++ b/docs/zh/conf.py @@ -0,0 +1,43 @@ +# Configuration file for the Sphinx documentation builder. + +import tomllib +from pathlib import Path + +# -- Project information ----------------------------------------------------- + +project = "entropack" +copyright = "2026, EntroPack Authors" +author = "EntroPack Authors" +html_theme = "sphinx_rtd_theme" +language = "zh_CN" + + +def get_version() -> str: + pyproject = Path(__file__).resolve().parents[2] / "pyproject.toml" + with pyproject.open("rb") as handle: + return tomllib.load(handle)["project"]["version"] + + +version = get_version() +release = version + +# -- General configuration --------------------------------------------------- + +extensions = [ + "sphinx_markdown_tables", + "sphinx_copybutton", + "sphinx_rtd_theme", + "sphinx.ext.mathjax", + "myst_parser", +] + +source_suffix = [".rst", ".md"] +root_doc = "index" +exclude_patterns = ["build", "_build"] + +# -- Extension configuration ------------------------------------------------- + +copybutton_prompt_text = r">>> |\.\.\. " +copybutton_prompt_is_regexp = True +intersphinx_mapping = {"https://docs.python.org/": None} +myst_enable_extensions = ["amsmath", "dollarmath", "colon_fence"] diff --git a/docs/zh/index.rst b/docs/zh/index.rst new file mode 100644 index 0000000..50db26e --- /dev/null +++ b/docs/zh/index.rst @@ -0,0 +1,27 @@ +EntroPack 文档 +======================== + +面向 PyTorch 的通用张量压缩,支持无损压缩与目标码率可调的有损压缩。 + +.. toctree:: + :maxdepth: 2 + :caption: 使用文档 + + Usage/Quick-start + Usage/Configuration + Usage/Tensor-compression + Usage/Linear-layers + +.. toctree:: + :maxdepth: 2 + :caption: API 参考 + + API_Reference/index + +.. toctree:: + :maxdepth: 2 + :caption: 压缩原理 + + Principles/DFloat11 + Principles/Tile-ANS + Principles/Lattice-rANS diff --git a/entropack/__init__.py b/entropack/__init__.py new file mode 100644 index 0000000..203b22b --- /dev/null +++ b/entropack/__init__.py @@ -0,0 +1,15 @@ +from importlib import metadata + +try: + __version__ = metadata.version("entropack") +except metadata.PackageNotFoundError: + __version__ = "0+unknown" + +from .compression import CompressedTensor, compress, decompress +from .linear import CompressedFP8Linear, CompressedINT8Linear, CompressedLinear +from .schemes import CompressionConfig, DFloat11Config, LatticeRANSConfig, RawConfig, TileANSConfig + +__all__ = [ + "CompressedFP8Linear", "CompressedINT8Linear", "CompressedLinear", "CompressedTensor", "CompressionConfig", + "DFloat11Config", "LatticeRANSConfig", "RawConfig", "TileANSConfig", "__version__", "compress", "decompress", +] diff --git a/entropack/backends/__init__.py b/entropack/backends/__init__.py new file mode 100644 index 0000000..137e6a9 --- /dev/null +++ b/entropack/backends/__init__.py @@ -0,0 +1,19 @@ +from collections.abc import Callable +from dataclasses import dataclass + +import torch + +from . import cuda + + +@dataclass(frozen=True) +class Backend: + name: str + priority: int = 0 + probe: Callable[[], str | None] | None = None + device: Callable[[], torch.device] | None = None + + +BACKENDS: tuple[Backend, ...] = ( + Backend("cuda", priority=100, probe=cuda.probe, device=cuda.runs_on), Backend("eager", priority=0), +) diff --git a/entropack/backends/cuda/__init__.py b/entropack/backends/cuda/__init__.py new file mode 100644 index 0000000..452877a --- /dev/null +++ b/entropack/backends/cuda/__init__.py @@ -0,0 +1,16 @@ +import torch + + +def probe() -> str | None: + if not torch.cuda.is_available(): + return "torch.cuda.is_available() is False (no NVIDIA device / driver)" + try: + import cupy + except ImportError as e: + return (f"cupy is not installed ({e}); install a matching build, e.g. " + "`pip install cupy-cuda13x` for CUDA 13 or `pip install cupy-cuda12x` for CUDA 12") + return None + + +def runs_on() -> torch.device: + return torch.device("cuda", torch.cuda.current_device()) diff --git a/entropack/backends/cuda/device.py b/entropack/backends/cuda/device.py new file mode 100644 index 0000000..33997a0 --- /dev/null +++ b/entropack/backends/cuda/device.py @@ -0,0 +1,96 @@ +from dataclasses import dataclass + +import cupy +import torch + +from .kernels import device_index + +_FALLBACK_WARP_SIZE = 32 + +_DEFAULT_GRID_WAVES = 8 + +__all__ = ["DeviceCaps", "caps", "resolve_threads", "validate_threads_per_block"] + + +def _shared_optin(index: int, properties) -> int: + optin = getattr(properties, "shared_memory_per_block_optin", 0) + if optin: + return int(optin) + with cupy.cuda.Device(index): + attributes = cupy.cuda.Device(index).attributes + return int(attributes["MaxSharedMemoryPerBlockOptin"]) + + +@dataclass(frozen=True) +class DeviceCaps: + index: int + name: str + compute_capability: tuple[int, int] + sm_count: int + warp_size: int + max_threads_per_block: int + threads_per_sm: int + regs_per_sm: int + shared_per_block: int + shared_optin: int + shared_per_sm: int + l2_bytes: int + + def blocks_per_sm(self, threads_per_block: int, shared_per_block: int = 0, regs_per_thread: int = 0) -> int: + blocks = self.threads_per_sm // threads_per_block + if shared_per_block > 0: + blocks = min(blocks, self.shared_per_sm // shared_per_block) + if regs_per_thread > 0: + blocks = min(blocks, self.regs_per_sm // (regs_per_thread * threads_per_block)) + return max(1, blocks) + + def resident_blocks(self, threads_per_block: int, shared_per_block: int = 0, regs_per_thread: int = 0) -> int: + return self.sm_count * self.blocks_per_sm(threads_per_block, shared_per_block, regs_per_thread) + + def grid(self, wanted: int, threads_per_block: int, shared_per_block: int = 0, waves: int = _DEFAULT_GRID_WAVES) -> int: + limit = self.resident_blocks(threads_per_block, shared_per_block) * waves + return max(1, min(wanted, limit)) + + def shared_limit(self, static_slack: int = 0) -> int: + return max(0, min(self.shared_optin, self.shared_per_sm - static_slack)) + + def threads_per_block(self, wanted: int) -> int: + usable = min(wanted, self.max_threads_per_block) + usable -= usable % self.warp_size + return max(self.warp_size, usable) + + +_caps_cache: dict[int, DeviceCaps] = {} + + +def caps(device=None) -> DeviceCaps: + index = device_index(device) + cached = _caps_cache.get(index) + if cached is not None: + return cached + properties = torch.cuda.get_device_properties(index) + queried = DeviceCaps( + index=index, name=properties.name, compute_capability=(properties.major, properties.minor), + sm_count=properties.multi_processor_count, warp_size=getattr(properties, "warp_size", 0) or _FALLBACK_WARP_SIZE, + max_threads_per_block=properties.max_threads_per_block, threads_per_sm=properties.max_threads_per_multi_processor, + regs_per_sm=properties.regs_per_multiprocessor, shared_per_block=properties.shared_memory_per_block, + shared_optin=_shared_optin(index, properties), shared_per_sm=properties.shared_memory_per_multiprocessor, + l2_bytes=getattr(properties, "L2_cache_size", 0), + ) + _caps_cache[index] = queried + return queried + + +def validate_threads_per_block(caps: DeviceCaps, threads_per_block: int) -> None: + if threads_per_block % caps.warp_size or threads_per_block > caps.max_threads_per_block: + raise ValueError( + f"threads_per_block={threads_per_block} is not launchable on {caps.name}: it must be " + f"a multiple of the {caps.warp_size}-thread warp size and at most {caps.max_threads_per_block}" + ) + + +def resolve_threads(caps: DeviceCaps, requested: int | None, default: int) -> int: + if requested is None: + return caps.threads_per_block(default) + validate_threads_per_block(caps, requested) + return requested diff --git a/entropack/backends/cuda/kernels.py b/entropack/backends/cuda/kernels.py new file mode 100644 index 0000000..6c37a72 --- /dev/null +++ b/entropack/backends/cuda/kernels.py @@ -0,0 +1,70 @@ +from collections.abc import Callable, Hashable, Sequence +from pathlib import Path + +import cupy +import numpy as np +import torch + +_modules: dict[tuple, object] = {} +_kernels: dict[tuple, object] = {} +_streams: dict[tuple, object] = {} + +__all__ = ["KernelLibrary", "device_index", "ensure_dynamic_shared", "external_stream", "pointer"] + + +def ensure_dynamic_shared(kernel, shared_bytes: int) -> None: + if shared_bytes > kernel.max_dynamic_shared_size_bytes: + kernel.max_dynamic_shared_size_bytes = shared_bytes + + +class KernelLibrary: + def __init__( + self, key: str, source: Path, defines: Callable[[int, Hashable], Sequence[str]], + includes: Sequence[Path] = (), kernel_names: Sequence[str] | None = None, + ): + self.key = key + self.source = source + self.defines = defines + self.includes = tuple(includes) + self.kernel_names = None if kernel_names is None else frozenset(kernel_names) + + def kernel(self, device_index: int, variant: Hashable, name: str): + if self.kernel_names is not None and name not in self.kernel_names: + raise KeyError(name) + module_key = (self.key, device_index, variant) + module = _modules.get(module_key) + if module is None: + options = ["--std=c++17"] + options += [f"-D{definition}" for definition in self.defines(device_index, variant)] + options += [f"-I{include}" for include in self.includes] + with cupy.cuda.Device(device_index): + module = cupy.RawModule(code=self.source.read_text(), options=tuple(options)) + _modules[module_key] = module + kernel_key = (self.key, device_index, variant, name) + kernel = _kernels.get(kernel_key) + if kernel is None: + kernel = module.get_function(name) + _kernels[kernel_key] = kernel + return kernel + + +def external_stream(stream: torch.cuda.Stream): + key = (device_index(stream.device), int(stream.cuda_stream)) + wrapper = _streams.get(key) + if wrapper is None: + factory = getattr(cupy.cuda.Stream, "from_external", None) + wrapper = factory(stream) if factory else cupy.cuda.ExternalStream(key[1]) + _streams[key] = wrapper + return wrapper + + +def pointer(tensor: torch.Tensor) -> np.uint64: + return np.uint64(tensor.data_ptr()) + + +def device_index(target) -> int: + if isinstance(target, int): + return target + device = target.device if isinstance(target, torch.Tensor) else target + index = getattr(device, "index", None) + return torch.cuda.current_device() if index is None else index diff --git a/entropack/compression/__init__.py b/entropack/compression/__init__.py new file mode 100644 index 0000000..2affb50 --- /dev/null +++ b/entropack/compression/__init__.py @@ -0,0 +1,4 @@ +from .api import CompressionFallbackWarning, compress, decompress +from .compressed_tensor import CompressedTensor + +__all__ = ["CompressedTensor", "CompressionFallbackWarning", "compress", "decompress"] diff --git a/entropack/compression/api.py b/entropack/compression/api.py new file mode 100644 index 0000000..7de5af3 --- /dev/null +++ b/entropack/compression/api.py @@ -0,0 +1,76 @@ +import logging +import warnings + +import torch + +from ..registry import DispatchError, get_scheme, name_for_config +from ..schemes import CompressionConfig, RawConfig +from ..schemes.config import validate_config +from .compressed_tensor import CompressedTensor + +logger = logging.getLogger("entropack") + + +class CompressionFallbackWarning(RuntimeWarning): + """Warning emitted when compression falls back to storing the tensor uncompressed.""" + + +def compress(tensor: torch.Tensor, config: CompressionConfig) -> CompressedTensor: + """Compress a tensor using the selected scheme. + + Args: + tensor: input values. The container preserves the input shape, dtype, and device. + config: configuration selecting the compression scheme and execution backend. + + Returns: + A :class:`CompressedTensor` containing the encoded data and metadata. + + Encoding failures can return an uncompressed ``raw`` container with a + :class:`CompressionFallbackWarning`. Its header records the requested scheme and + failure reason. Invalid configurations and backend dispatch failures raise errors. + """ + if isinstance(tensor, CompressedTensor): + raise TypeError("compress expects an uncompressed tensor; decompress the container before recompressing") + scheme = get_scheme(name_for_config(config)) + validate_config(config) + try: + packed = scheme.encode(tensor, config) + except DispatchError: + raise + except Exception as error: + logger.warning( + "%s could not encode a %s %s tensor, storing it uncompressed: %s", + scheme.name, tuple(tensor.shape), tensor.dtype, error, + ) + return _compress_raw(tensor, scheme.name, error) + return CompressedTensor(header={"compress_method": scheme.name}, buffers=packed, shape=tuple(tensor.shape), dtype=tensor.dtype) + + +def _compress_raw(tensor: torch.Tensor, requested: str, error: Exception) -> CompressedTensor: + reason = f"{type(error).__name__}: {error}" + warnings.warn( + f"entropack stored a {tuple(tensor.shape)} {tensor.dtype} tensor uncompressed, at " + f"{tensor.element_size() * 8} bits per element: compress_method={requested!r} failed with {reason}", + CompressionFallbackWarning, stacklevel=3, + ) + return CompressedTensor( + header={"compress_method": "raw", "requested": requested, "reason": reason}, + buffers=get_scheme("raw").encode(tensor, RawConfig()), shape=tuple(tensor.shape), dtype=tensor.dtype, + ) + + +def decompress(compressed: CompressedTensor, config: CompressionConfig) -> torch.Tensor: + """Restore a tensor from its compressed representation. + + Args: + compressed: container to decode. Its header identifies the compression scheme. + config: configuration for the same scheme. Decode settings select the backend and + execution options. Encode settings such as the target bitrate do not recompress + the stored data. + + Returns: + A tensor with the container's shape and dtype, on the device holding its buffers. + """ + validate_config(config) + restored = compressed.scheme.decode(compressed.buffers, shape=compressed.shape, dtype=compressed.encoded_dtype, config=config) + return restored.to(compressed.dtype) diff --git a/entropack/compression/compressed_tensor.py b/entropack/compression/compressed_tensor.py new file mode 100644 index 0000000..3806a49 --- /dev/null +++ b/entropack/compression/compressed_tensor.py @@ -0,0 +1,205 @@ +import copy +import json +import math +from collections.abc import Mapping +from typing import Any + +import torch +from torch.utils._python_dispatch import return_and_correct_aliasing + +from ..registry import get_scheme +from ..schemes import Scheme + + +def _parse_dtype(name: str) -> torch.dtype: + if not name: + raise TypeError("CompressedTensor serialized dtype must be a non-empty string") + dtype = getattr(torch, name, None) + if not isinstance(dtype, torch.dtype): + raise ValueError(f"Unsupported serialized torch dtype '{name}'") + return dtype + + +class CompressedTensor(torch.Tensor): + """Frozen compressed tensor with a logical output dtype. + + ``encoded_dtype`` records the codec's input type. ``to(dtype=...)`` changes + the output dtype without casting packed buffers or re-encoding values. + """ + + header: dict[str, Any] + buffers: dict[str, torch.Tensor] + + @staticmethod + def __new__(cls, header, buffers, shape, dtype, *, encoded_dtype=None): + device = next(iter(buffers.values())).device + return torch.Tensor._make_wrapper_subclass(cls, tuple(shape), dtype=dtype, device=device, requires_grad=False) + + def __init__(self, header, buffers, shape, dtype, *, encoded_dtype=None): + self.header = copy.deepcopy(header) + self.buffers = dict(buffers) + self.encoded_dtype = dtype if encoded_dtype is None else encoded_dtype + self.validate() + + def __repr__(self) -> str: + return f"{type(self).__name__}(shape={tuple(self.shape)}, dtype={self.dtype}, device={self.device}, scheme={self.compress_method!r})" + + def __tensor_flatten__(self): + names = tuple(self.buffers) + for name, value in self.buffers.items(): + setattr(self, "_packed_" + name, value) + metadata = (names, copy.deepcopy(self.header), tuple(self.shape), self.dtype, self.encoded_dtype) + return ["_packed_" + name for name in names], metadata + + @classmethod + def __tensor_unflatten__(cls, inner_tensors, metadata, outer_size, outer_stride): + names, header, shape, dtype, encoded_dtype = metadata + return cls( + header=header, buffers={name: inner_tensors["_packed_" + name] for name in names}, + shape=shape, dtype=dtype, encoded_dtype=encoded_dtype, + ) + + def _map_buffers(self, fn, *, dtype=None) -> "CompressedTensor": + return type(self)( + header=self.header, buffers={name: fn(value) for name, value in self.buffers.items()}, + shape=self.shape, dtype=self.dtype if dtype is None else dtype, encoded_dtype=self.encoded_dtype, + ) + + def to(self, *args, copy: bool = False, **kwargs) -> "CompressedTensor": + """Move buffers or change the logical dtype, with an optional keyword-only copy.""" + device, dtype, non_blocking, memory_format = torch._C._nn._parse_to(*args, **kwargs) + result = super().to(device=device, dtype=dtype, non_blocking=non_blocking, copy=copy, memory_format=memory_format) + return result.clone() if copy and result.device == self.device else result + + def __copy__(self) -> "CompressedTensor": + result = self._map_buffers(lambda value: value) + if getattr(self, "_is_param", False): + result._is_param = True + return result + + def __deepcopy__(self, memo) -> "CompressedTensor": + if id(self) in memo: + return memo[id(self)] + result = self._map_buffers(lambda value: copy.deepcopy(value, memo)) + if getattr(self, "_is_param", False): + result._is_param = True + memo[id(self)] = result + return result + + def requires_grad_(self, requires_grad: bool = False): + """Compressed values are frozen; gradients may flow through their consumers.""" + if requires_grad: + raise RuntimeError("CompressedTensor cannot require gradients; decompress it to train dense values") + return torch.Tensor.requires_grad_(self, False) + + @classmethod + def __torch_dispatch__(cls, func, types, args=(), kwargs=None): + kwargs = kwargs or {} + source = args[0] if args else None + aten = torch.ops.aten + + if func in (aten.detach.default, aten.alias.default): + result = source._map_buffers(torch.detach) + return return_and_correct_aliasing(func, args, kwargs, result) + + if func in (aten.clone.default, aten._to_copy.default): + if kwargs.get("memory_format") not in (None, torch.preserve_format, torch.contiguous_format): + raise ValueError("CompressedTensor does not support that memory format") + if func == aten.clone.default: + return source._map_buffers(torch.clone) + if kwargs.get("layout", torch.strided) not in (None, torch.strided): + raise ValueError("CompressedTensor only supports strided storage") + options = {name: value for name, value in kwargs.items() if name in ("device", "non_blocking")} + return source._map_buffers(lambda value: value.to(**options), dtype=kwargs.get("dtype")) + + result = func.decompose(*args, **kwargs) + if result is not NotImplemented: + return result + raise NotImplementedError(f"CompressedTensor does not implement {func}; decompress it before numerical operations") + + @property + def compress_method(self) -> str: + """The scheme name the header carries.""" + return self.header["compress_method"] + + @property + def scheme(self) -> Scheme: + """The codec named by the header.""" + return get_scheme(self.compress_method) + + @property + def lossless(self) -> bool: + """Whether the encoding scheme is lossless, before any output dtype conversion.""" + return self.scheme.lossless + + @property + def actual_bpp(self) -> float: + """Bits per element of :attr:`shape`, serialized header included.""" + return self.storage_nbytes() * 8 / math.prod(self.shape) + + def validate(self) -> None: + """Check the scheme, dtype, and buffer names without inspecting contents.""" + if any(isinstance(value, CompressedTensor) for value in self.buffers.values()): + raise TypeError("CompressedTensor buffers cannot contain another CompressedTensor") + scheme = self.scheme + if not scheme.supports(self.encoded_dtype): + raise ValueError(f"'{scheme.name}' does not support format {self.encoded_dtype}") + if set(self.buffers) != set(scheme.buffer_names): + raise ValueError( + f"CompressedTensor buffers for {self.compress_method} must be " + f"{list(scheme.buffer_names)}, got {sorted(self.buffers)}" + ) + + def _serialized_header_tensor(self) -> torch.Tensor: + metadata = { + "compress_method": self.compress_method, "dtype": str(self.encoded_dtype).removeprefix("torch."), + "shape": list(self.shape), "buffer_names": list(self.buffers), "header": self.header, + } + encoded = json.dumps(metadata, sort_keys=True, separators=(",", ":"), allow_nan=False).encode("utf-8") + return torch.frombuffer(bytearray(encoded), dtype=torch.uint8) + + def storage_nbytes(self, include_header: bool = True) -> int: + """Bytes the buffers occupy, plus the serialized header unless ``include_header`` is false.""" + total = sum(value.numel() * value.element_size() for value in self.buffers.values()) + if include_header: + total += self._serialized_header_tensor().numel() + return total + + def state_dict(self, prefix: str = "") -> dict[str, torch.Tensor]: + """The container as flat tensors under ``prefix``: one 1D uint8 header, then one entry per buffer.""" + self.validate() + state = {f"{prefix}header": self._serialized_header_tensor()} + state.update({f"{prefix}buffers.{name}": value.detach() for name, value in self.buffers.items()}) + return state + + @classmethod + def from_state_dict(cls, state: Mapping[str, torch.Tensor], prefix: str = "") -> "CompressedTensor": + """Restore a container in its encoded dtype from a tensor-only checkpoint.""" + header_key = f"{prefix}header" + if header_key not in state: + raise ValueError(f"CompressedTensor state is missing '{header_key}'") + header_tensor = state[header_key] + if header_tensor.dtype != torch.uint8 or header_tensor.ndim != 1: + raise TypeError("CompressedTensor serialized header must be a 1D uint8 tensor") + try: + metadata = json.loads(bytes(header_tensor.detach().cpu().tolist()).decode("utf-8")) + except (UnicodeDecodeError, json.JSONDecodeError) as error: + raise ValueError("CompressedTensor serialized header is invalid") from error + if not isinstance(metadata, dict): + raise ValueError("CompressedTensor serialized header must contain a JSON object") + + required = {"compress_method", "dtype", "shape", "buffer_names", "header"} + missing = required - metadata.keys() + if missing: + raise ValueError(f"CompressedTensor serialized header is missing fields: {sorted(missing)}") + buffer_names = metadata["buffer_names"] + if not isinstance(buffer_names, list) or any(not isinstance(name, str) or not name for name in buffer_names): + raise TypeError("CompressedTensor buffer_names must be a list of strings") + buffers = {} + for name in buffer_names: + key = f"{prefix}buffers.{name}" + if key not in state: + raise ValueError(f"CompressedTensor state is missing '{key}'") + buffers[name] = state[key] + + return cls(header=metadata["header"], buffers=buffers, shape=tuple(metadata["shape"]), dtype=_parse_dtype(metadata["dtype"])) diff --git a/entropack/linear/__init__.py b/entropack/linear/__init__.py new file mode 100644 index 0000000..30df85e --- /dev/null +++ b/entropack/linear/__init__.py @@ -0,0 +1,5 @@ +from .linear import CompressedFP8Linear, CompressedINT8Linear, CompressedLinear, QuantizedLinear + +__all__ = [ + "CompressedFP8Linear", "CompressedINT8Linear", "CompressedLinear", "QuantizedLinear", +] diff --git a/entropack/linear/linear.py b/entropack/linear/linear.py new file mode 100644 index 0000000..d83f1f0 --- /dev/null +++ b/entropack/linear/linear.py @@ -0,0 +1,444 @@ +import dataclasses +from numbers import Real +from typing import ClassVar + +import torch +from torch.nn import functional as F + +from ..compression import CompressedTensor, compress, decompress +from ..registry import default_config as _default_config, get_scheme, name_for_config, require_dtype +from ..schemes import LatticeRANSConfig, RawConfig +from .utils import capability_of, load_quant_kernels, pad, round_up + +_EPS = torch.finfo(torch.float32).eps +_BACKEND = "cuda" + + +class CompressedLinear(torch.nn.Linear): + """A linear layer with compressed weights, reconstructed during each forward call. + + The frozen ``weight`` is a ``CompressedTensor`` parameter. Each forward + reconstructs the stored weight, casts it to the activation dtype, and applies + ``F.linear``. Checkpoints retain the original packed representation. + CUDA and CuPy are required. + + Args: + in_features: number of input features. + out_features: number of output features. + bias: whether to keep a bias. The bias is not compressed. + config: compression configuration. ``None`` selects a lossless scheme by dtype. + device: device for the bias. Compressed buffers retain the source weight's device. + dtype: container dtype, which must be supported by the selected scheme. + """ + + state_buffer_names: ClassVar[tuple[str, ...]] = () + state_prefix: ClassVar[str] = "weight._entropack." + + def __init__( + self, in_features: int, out_features: int, bias: bool = True, *, config=None, + device: str | torch.device | None = None, dtype: torch.dtype = torch.bfloat16, + ): + with torch.device("meta"): + super().__init__(in_features, out_features, bias=False, dtype=dtype) + self.weight = None + if bias: + self.bias = torch.nn.Parameter(torch.zeros(out_features, dtype=dtype, device=device), requires_grad=False) + self.config = config + self._init_dtype = require_dtype(dtype) + if config is not None: + scheme = get_scheme(name_for_config(config)) + if not scheme.supports(self.container_dtype): + raise ValueError(f"'{scheme.name}' does not support format {self.container_dtype}") + for name in self.state_buffer_names: + self.register_buffer(name, None, persistent=False) + + @property + def _encode_config(self): + if self.config is None: + return _default_config(self.container_dtype, execution_backend=_BACKEND) + return dataclasses.replace(self.config, execution_backend=_BACKEND) + + @property + def _decode_config(self): + if self.config is None: + return self.weight.scheme.make_config({"execution_backend": _BACKEND}) + return dataclasses.replace(self.config, execution_backend=_BACKEND) + + @property + def container_dtype(self) -> torch.dtype: + """The configured compression dtype, unchanged by layer dtype casts.""" + return self._init_dtype + + @property + def qweight(self) -> torch.Tensor: + """A packed weight buffer exposed for quantization-framework compatibility.""" + if self.weight is not None: + for name in self.weight.scheme.buffer_names: + return self.weight.buffers[name] + return self.bias + + @property + def stored_nbytes(self) -> int: + """Stored weight bytes, including the serialized header and any W8A8 scales.""" + total = self.weight.storage_nbytes() + for name in self.state_buffer_names: + buffer = self._buffers.get(name) + if buffer is not None: + total += buffer.numel() * buffer.element_size() + return total + + @property + def compressed_bits(self) -> float: + """Stored bits per weight element, including metadata.""" + rows, cols = self.weight.shape + return self.stored_nbytes * 8 / (rows * cols) + + def compress_weight(self, weight: torch.Tensor) -> None: + """Initialize the layer with compressed ``weight``.""" + self.weight = torch.nn.Parameter(self._compress_tensor(weight.detach()), requires_grad=False) + + def _compress_tensor(self, tensor: torch.Tensor) -> CompressedTensor: + compressed = compress(tensor.to(self.container_dtype), self._encode_config) + if compressed.compress_method == "raw" and "reason" in compressed.header: + raise RuntimeError( + f"{compressed.header['requested']} cannot store a {tuple(tensor.shape)} {self.container_dtype} weight for " + f"{type(self).__name__}: {compressed.header['reason']}" + ) + return compressed + + def _reconstruct(self, device: str | torch.device | None = None) -> torch.Tensor: + return decompress(self.weight.to(device=device), self._decode_config) + + def dequantize(self, device: str | torch.device | None = None) -> torch.Tensor: + """Reconstruct the dense weight on ``device``, defaulting to the weight's device.""" + return self._reconstruct(device) + + @classmethod + def from_linear(cls, linear: torch.nn.Linear, **kwargs) -> "CompressedLinear": + """Compress an existing layer's weight into a new layer of this class. + + Args: + linear: the source layer, with its weight materialized -- so this runs after a + ``load_state_dict``, not on a meta-device skeleton. + **kwargs: passed to the constructor; ``dtype`` defaults to the source weight's. + + Returns: + A layer of this class holding the compressed weight and, when the source had one, a copy + of its bias. + """ + weight = linear.weight + if weight is None or weight.device.type == "meta": + raise ValueError("cannot compress a Linear whose weight is not materialized") + kwargs.setdefault("dtype", weight.dtype) + out = cls(linear.in_features, linear.out_features, bias=linear.bias is not None, device=weight.device, **kwargs) + out.compress_weight(weight.data) + if linear.bias is not None: + out.bias = torch.nn.Parameter(linear.bias.data.clone(), requires_grad=False) + return out + + @property + def device(self) -> torch.device: + """The weight device, or the bias device while the weight is not loaded.""" + return self.weight.device if self.weight is not None else self.bias.device if self.bias is not None else torch.device("meta") + + def _bias_on(self, device: torch.device) -> torch.Tensor | None: + return None if self.bias is None else self.bias.to(device) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Reconstruct the weight on ``x``'s device and apply the layer to ``x``.""" + return F.linear(x, self.dequantize(x.device).to(x.dtype), self._bias_on(x.device)) + + def extra_repr(self) -> str: + scheme = "auto" if self.config is None else name_for_config(self.config) + parts = [f"scheme={scheme}", f"container={str(self.container_dtype).removeprefix('torch.')}"] + if self.weight is not None: + parts.append(f"bits={self.compressed_bits:.3f}") + if self.config is not None: + parts.append(f"config={type(self.config).__name__}") + return ", ".join(parts) + + def _apply(self, fn, recurse=True): + # Preserve auxiliary scale dtypes when Module.to casts the layer. + held = {name: self._buffers.pop(name) for name in self.state_buffer_names if name in self._buffers} + try: + super()._apply(fn, recurse=recurse) + finally: + for name, buffer in held.items(): + if buffer is None: + self._buffers[name] = None + continue + moved = fn(buffer) + self._buffers[name] = moved if moved.dtype == buffer.dtype else buffer.to(device=moved.device) + return self + + def _save_to_state_dict(self, destination: dict, prefix: str, keep_vars: bool) -> None: + super()._save_to_state_dict(destination, prefix, keep_vars) + # Serialize the packed representation instead of the wrapper parameter. + destination.pop(prefix + "weight", None) + if self.weight is None: + return + container_prefix = prefix + self.state_prefix + written = self.weight.state_dict(container_prefix) + for name in self.state_buffer_names: + buffer = self._buffers[name] + if buffer is not None: + written[container_prefix + name] = buffer + for key, value in written.items(): + destination[key] = value if keep_vars else value.detach() + + def _load_from_state_dict(self, state_dict: dict, prefix: str, local_metadata, strict, missing_keys, unexpected_keys, error_msgs) -> None: + container_prefix = prefix + self.state_prefix + header_key = container_prefix + "header" + if header_key in state_dict: + consumed = [header_key] + try: + compressed = CompressedTensor.from_state_dict(state_dict, prefix=container_prefix) + consumed.extend(container_prefix + "buffers." + name for name in compressed.buffers) + assign = local_metadata.get("assign_to_params_buffers", False) + + # Copy loading keeps the target device and existing Parameter; assign=True adopts the checkpoint tensors. + if self.weight is not None: + device = self.weight.device + elif self.bias is not None and self.bias.device.type != "meta": + device = self.bias.device + else: + device = compressed.device + if not assign: + compressed = compressed.to(device=device, copy=True) + loaded_weight = torch.nn.Parameter(compressed, requires_grad=False) + if not assign and self.weight is not None: + torch.utils.swap_tensors(self.weight, loaded_weight) + else: + self.weight = loaded_weight + + # Apply the same copy/assign behavior to auxiliary state, such as FP8/INT8 weight scales. + for name in self.state_buffer_names: + key = container_prefix + name + if key not in state_dict: + raise ValueError(f"compressed state is missing '{key}'") + value = state_dict[key] + if not assign: + if self._buffers[name] is None: + value = value.to(device=device, copy=True) + else: + self._buffers[name].copy_(value) + value = self._buffers[name] + self._buffers[name] = value + consumed.append(key) + except (TypeError, ValueError) as error: + error_msgs.append(f"{prefix[:-1]}: {error}") + + # Leave unrecognized keys for the parent's strict checks. + for key in consumed: + state_dict.pop(key) + elif strict: + missing_keys.append(header_key) + + # Let the parent load bias without expecting a dense weight entry. + weight = self._parameters.pop("weight") + super()._load_from_state_dict(state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs) + self._parameters["weight"] = weight + + +class _W8A8LinearFunction(torch.autograd.Function): + + @staticmethod + def forward(ctx, x, layer): + ctx.layer = layer + ctx.x_shape = x.shape + return layer._forward8(x.detach()) + + @staticmethod + @torch.autograd.function.once_differentiable + def backward(ctx, grad_output): + grad = grad_output.reshape(-1, grad_output.shape[-1]) + weight = ctx.layer.dequantize(grad.device).to(grad.dtype) + return torch.mm(grad, weight).reshape(ctx.x_shape), None + + +class QuantizedLinear(CompressedLinear): + + code_dtype: ClassVar[torch.dtype] + code_max: ClassVar[float] + rounds_to_integer: ClassVar[bool] + code_alignment: ClassVar[int] + min_tokens: ClassVar[int] + min_capability: ClassVar[tuple[int, int]] + state_buffer_names = ("weight_scale",) + + def __init__( + self, in_features: int, out_features: int, bias: bool = True, *, config=None, + device: str | torch.device | None = None, dtype: torch.dtype = torch.bfloat16, + ): + super().__init__(in_features, out_features, bias=bias, config=self._coded(config), device=device, dtype=dtype) + + @property + def container_dtype(self) -> torch.dtype: + return self.code_dtype + + @classmethod + def _coded(cls, config): + if config is None: + return RawConfig() + if isinstance(config, RawConfig): + return config + if not isinstance(config, LatticeRANSConfig): + raise TypeError( + f"{cls.__name__} stores {cls.code_dtype} codes: either uncoded (RawConfig) or lattice-coded " + f"(LatticeRANSConfig), got {type(config).__name__}" + ) + rate = config.target_bpp + if isinstance(rate, bool) or not isinstance(rate, Real) or not 0.0 < float(rate) < 8.0: + raise ValueError( + f"{cls.__name__} stores {cls.code_dtype} codes, 8 bits wide, so a target_bpp " + f"only pays below 8; got {rate!r}. Pass RawConfig to store the codes uncoded." + ) + return config + + def _quantize_rows(self, tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + flat = tensor.reshape(-1, tensor.shape[-1]).float() + scale = (flat.abs().amax(dim=1) / self.code_max).clamp(min=_EPS) + scaled = flat / scale.unsqueeze(1) + if self.rounds_to_integer: + scaled = scaled.round() + return scaled.clamp(-self.code_max, self.code_max).to(self.code_dtype), scale + + def compress_weight(self, weight: torch.Tensor) -> None: + """Quantize the source weight, compress its codes, and store the row scales.""" + codes, scale = self._quantize_rows(weight.detach()) + self.weight = torch.nn.Parameter(self._compress_tensor(codes), requires_grad=False) + self.weight_scale = scale + + def codes(self, device: str | torch.device | None = None) -> torch.Tensor: + """Decode the weight codes in ``code_dtype`` for low-precision matrix multiplication.""" + return decompress(self.weight.to(device=device, dtype=self.code_dtype), self._decode_config) + + def dequantize(self, device: str | torch.device | None = None) -> torch.Tensor: + """Reconstruct dense weights in the dtype used to initialize the layer.""" + codes = self.codes(device) + return (codes.float() * self.weight_scale.to(codes.device).unsqueeze(1)).to(self._init_dtype) + + def _require_8bit_gemm(self, device: torch.device) -> None: + if device.type != "cuda": + raise RuntimeError( + f"{type(self).__name__} needs a CUDA device with an 8-bit tensor core, compute capability " + f"{self.min_capability} or later; got {device}." + ) + capability = capability_of(device) + if capability < self.min_capability: + raise RuntimeError( + f"{type(self).__name__} needs compute capability {self.min_capability} or later and " + f"{torch.cuda.get_device_name(device)} reports {capability}. This format has no path on " + "older hardware." + ) + + def _quantize_activation(self, flat: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + kernels = load_quant_kernels() + if kernels is None: + return self._quantize_rows(flat) + return kernels.quantize_rows(flat, self.code_dtype, self.code_max, self.rounds_to_integer) + + def _gemm_shapes(self, tokens: int) -> tuple[int, int, int]: + return (max(tokens, self.min_tokens), round_up(self.out_features, self.code_alignment), round_up(self.in_features, self.code_alignment)) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Apply the low-precision linear operation. + + The layer recovers its FP8 or INT8 codes, quantizes activations per row, and runs the + corresponding matrix multiplication. Input gradients use the reconstructed numerical + weight. The compressed base weights remain frozen. + """ + self._require_8bit_gemm(x.device) + if x.requires_grad: + return _W8A8LinearFunction.apply(x, self) + return self._forward8(x) + + def _forward8(self, x: torch.Tensor) -> torch.Tensor: + flat = x.reshape(-1, x.shape[-1]) + activation, scale = self._quantize_activation(flat.detach()) + codes = self.codes(x.device) + weight_scale = self.weight_scale.to(codes.device) + tokens = activation.shape[0] + rows, outs, cols = self._gemm_shapes(tokens) + padded_activation = pad(pad(activation, 0, rows), 1, cols) + padded_codes = pad(pad(codes, 0, outs), 1, cols) + out = self._gemm(padded_activation, padded_codes, pad(scale, 0, rows), pad(weight_scale, 0, outs), x.dtype) + return out[:tokens, : self.out_features].reshape(*x.shape[:-1], self.out_features) + + def _bias_padded(self, device: torch.device, columns: int) -> torch.Tensor | None: + bias = self._bias_on(device) + return None if bias is None else pad(bias, 0, columns) + + def _epilogue(self, product: torch.Tensor, activation_scale: torch.Tensor, weight_scale: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: + out = product.to(torch.float32) + out.mul_(activation_scale.unsqueeze(1)) + bias = self._bias_padded(out.device, out.shape[1]) + if bias is None: + return out.mul_(weight_scale).to(out_dtype) + return torch.addcmul(bias.to(torch.float32), out, weight_scale).to(out_dtype) + + def _gemm( + self, activation: torch.Tensor, codes: torch.Tensor, activation_scale: torch.Tensor, + weight_scale: torch.Tensor, out_dtype: torch.dtype, + ) -> torch.Tensor: + raise NotImplementedError + + +class CompressedFP8Linear(QuantizedLinear): + """A linear layer using FP8 E4M3FN weights and activations. + + Requires CUDA compute capability 8.9 or later and uses ``torch._scaled_mm``. + With no compression configuration, quantized codes are stored directly. + :class:`~entropack.LatticeRANSConfig` additionally compresses them at targets from 1 up + to, but excluding, 8 bits per element. Inference decodes the FP8 codes before matrix + multiplication. Stored size also includes metadata and per-row weight scales. + """ + + code_dtype = torch.float8_e4m3fn + code_max = float(torch.finfo(torch.float8_e4m3fn).max) + rounds_to_integer = False + code_alignment = 16 + min_tokens = 1 + min_capability = (8, 9) + + def _gemm(self, activation, codes, activation_scale, weight_scale, out_dtype): + bias = self._bias_padded(codes.device, weight_scale.shape[0]) + fused = bias is None or out_dtype is not torch.float32 + out = torch._scaled_mm( + activation, codes.t(), scale_a=activation_scale.unsqueeze(1), scale_b=weight_scale.unsqueeze(0), + bias=bias if fused else None, out_dtype=out_dtype, + ) + return out if fused else out.add_(bias) + + +class CompressedINT8Linear(QuantizedLinear): + """A linear layer using symmetric INT8 weights and activations. + + Requires CUDA compute capability 8.0 or later. Matrix multiplication uses Triton when + available and ``torch._int_mm`` otherwise. With no compression configuration, + quantized codes are stored directly. :class:`~entropack.LatticeRANSConfig` additionally compresses + them at targets from 1 up to, but excluding, 8 bits per element. Stored size also + includes metadata and per-row weight scales. + """ + + code_dtype = torch.int8 + code_max = 127.0 + rounds_to_integer = True + code_alignment = 8 + min_tokens = 17 + min_capability = (8, 0) + + def _gemm_shapes(self, tokens): + if load_quant_kernels() is not None: + return tokens, self.out_features, self.in_features + return super()._gemm_shapes(tokens) + + def _gemm(self, activation, codes, activation_scale, weight_scale, out_dtype): + bias = self._bias_padded(codes.device, weight_scale.shape[0]) + kernels = load_quant_kernels() + if kernels is not None: + return kernels.int8_gemm(activation, codes, activation_scale, weight_scale, bias, out_dtype) + return self._epilogue(torch._int_mm(activation, codes.t()), activation_scale, weight_scale, out_dtype) + + +__all__ = ["CompressedFP8Linear", "CompressedINT8Linear", "CompressedLinear", "QuantizedLinear"] diff --git a/entropack/linear/quant_kernels.py b/entropack/linear/quant_kernels.py new file mode 100644 index 0000000..a668c8c --- /dev/null +++ b/entropack/linear/quant_kernels.py @@ -0,0 +1,108 @@ + +import torch +import triton +import triton.language as tl +import triton.language.extra.libdevice as libdevice + +_INT8_GEMM_CONFIGS = [ + triton.Config({'BLOCK_M': 64, 'BLOCK_N': 64, 'BLOCK_K': 64, 'GROUP_M': 8}, num_warps=4, num_stages=3), + triton.Config({'BLOCK_M': 64, 'BLOCK_N': 64, 'BLOCK_K': 32, 'GROUP_M': 8}, num_warps=4, num_stages=4), + triton.Config({'BLOCK_M': 64, 'BLOCK_N': 128, 'BLOCK_K': 64, 'GROUP_M': 8}, num_warps=4, num_stages=4), + triton.Config({'BLOCK_M': 128, 'BLOCK_N': 64, 'BLOCK_K': 64, 'GROUP_M': 8}, num_warps=4, num_stages=4), + triton.Config({'BLOCK_M': 128, 'BLOCK_N': 128, 'BLOCK_K': 64, 'GROUP_M': 8}, num_warps=8, num_stages=3), +] + +@triton.autotune(configs=_INT8_GEMM_CONFIGS, key=['M', 'N', 'K']) +@triton.jit +def _int8_gemm_kernel( + A, B, A_SCALE, B_SCALE, BIAS, C, M, N, K, stride_am, stride_bn, + HAS_BIAS: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, + GROUP_M: tl.constexpr, +): + pid = tl.program_id(0) + grid_m = tl.cdiv(M, BLOCK_M) + grid_n = tl.cdiv(N, BLOCK_N) + width = GROUP_M * grid_n + group_id = pid // width + group_size = tl.minimum(grid_m - group_id * GROUP_M, GROUP_M) + pid_m = group_id * GROUP_M + (pid % group_size) + pid_n = (pid % width) // group_size + + rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) + rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) + rk = tl.arange(0, BLOCK_K) + mask_m = rm < M + mask_n = rn < N + + a_ptr = A + rm[:, None] * stride_am + rk[None, :] + b_ptr = B + rn[:, None] * stride_bn + rk[None, :] + acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.int32) + for k in range(0, K, BLOCK_K): + mask_k = (k + rk) < K + a = tl.load(a_ptr, mask=mask_m[:, None] & mask_k[None, :], other=0) + b = tl.load(b_ptr, mask=mask_n[:, None] & mask_k[None, :], other=0) + acc += tl.dot(a, tl.trans(b)) + a_ptr += BLOCK_K + b_ptr += BLOCK_K + + out = acc.to(tl.float32) + out = out * tl.load(A_SCALE + rm, mask=mask_m, other=0.0)[:, None] + out = out * tl.load(B_SCALE + rn, mask=mask_n, other=0.0)[None, :] + if HAS_BIAS: + out += tl.load(BIAS + rn, mask=mask_n, other=0.0).to(tl.float32)[None, :] + tl.store(C + rm[:, None] * N + rn[None, :], out.to(C.dtype.element_ty), + mask=mask_m[:, None] & mask_n[None, :]) + +@triton.jit +def _quantize_rows_kernel( + X, OUT, SCALE, K, stride_xm, CODE_MAX: tl.constexpr, ROUNDS: tl.constexpr, EPS: tl.constexpr, + BLOCK_K: tl.constexpr, +): + row = tl.program_id(0) + base = X + row * stride_xm + amax = tl.zeros((), dtype=tl.float32) + for k in range(0, K, BLOCK_K): + rk = k + tl.arange(0, BLOCK_K) + x = tl.load(base + rk, mask=rk < K, other=0.0).to(tl.float32) + amax = tl.maximum(amax, tl.max(tl.abs(x))) + scale = tl.maximum(amax / CODE_MAX, EPS) + + out_base = OUT + row * K + for k in range(0, K, BLOCK_K): + rk = k + tl.arange(0, BLOCK_K) + x = tl.load(base + rk, mask=rk < K, other=0.0).to(tl.float32) + scaled = x / scale + if ROUNDS: + scaled = libdevice.rint(scaled) + scaled = tl.minimum(tl.maximum(scaled, -CODE_MAX), CODE_MAX) + tl.store(out_base + rk, scaled.to(OUT.dtype.element_ty), mask=rk < K) + tl.store(SCALE + row, scale) + +def int8_gemm( + activation: torch.Tensor, codes: torch.Tensor, activation_scale: torch.Tensor, + weight_scale: torch.Tensor, bias: torch.Tensor | None, out_dtype: torch.dtype, +) -> torch.Tensor: + tokens, inner = activation.shape + outer = codes.shape[0] + out = torch.empty(tokens, outer, dtype=out_dtype, device=activation.device) + grid = lambda meta: (triton.cdiv(tokens, meta['BLOCK_M']) * triton.cdiv(outer, meta['BLOCK_N']),) + _int8_gemm_kernel[grid]( + activation, codes, activation_scale, weight_scale, bias, out, tokens, outer, inner, + activation.stride(0), codes.stride(0), bias is not None, + ) + return out + +def quantize_rows( + tensor: torch.Tensor, code_dtype: torch.dtype, code_max: float, rounds_to_integer: bool, +) -> tuple[torch.Tensor, torch.Tensor]: + flat = tensor.reshape(-1, tensor.shape[-1]) + rows, inner = flat.shape + codes = torch.empty(rows, inner, dtype=code_dtype, device=flat.device) + scale = torch.empty(rows, dtype=torch.float32, device=flat.device) + _quantize_rows_kernel[(rows,)]( + flat, codes, scale, inner, flat.stride(0), code_max, rounds_to_integer, + torch.finfo(torch.float32).eps, BLOCK_K=1024, num_warps=8, + ) + return codes, scale + +__all__ = ["int8_gemm", "quantize_rows"] diff --git a/entropack/linear/utils.py b/entropack/linear/utils.py new file mode 100644 index 0000000..8eb5218 --- /dev/null +++ b/entropack/linear/utils.py @@ -0,0 +1,41 @@ +import torch + +_kernels = None +_looked = False +_device_capabilities: dict[int, tuple[int, int]] = {} + + +def round_up(value: int, multiple: int) -> int: + return -(-value // multiple) * multiple + + +def pad(tensor: torch.Tensor, dim: int, size: int) -> torch.Tensor: + shortfall = size - tensor.shape[dim] + if shortfall <= 0: + return tensor + shape = list(tensor.shape) + shape[dim] = shortfall + return torch.cat([tensor, tensor.new_zeros(shape)], dim=dim) + + +def load_quant_kernels(): + global _kernels, _looked + if not _looked: + _looked = True + try: + from . import quant_kernels + _kernels = quant_kernels + except ImportError: + _kernels = None + return _kernels + + +def capability_of(device: torch.device) -> tuple[int, int]: + index = device.index if device.index is not None else torch.cuda.current_device() + capability = _device_capabilities.get(index) + if capability is None: + capability = torch.cuda.get_device_capability(index) + _device_capabilities[index] = capability + return capability + +__all__ = ["capability_of", "load_quant_kernels", "pad", "round_up"] diff --git a/entropack/registry.py b/entropack/registry.py new file mode 100644 index 0000000..87ba763 --- /dev/null +++ b/entropack/registry.py @@ -0,0 +1,121 @@ +import logging +import threading +from typing import Any + +import torch + +from .backends import BACKENDS, Backend + +logger = logging.getLogger("entropack.registry") + +_backends: dict[str, Backend] = {backend.name: backend for backend in BACKENDS} +_schemes: dict[str, Any] = {} +_reasons: dict[str, str | None] = {} +_lock = threading.Lock() + + +_CONTAINER_DTYPES = ( + torch.float32, torch.float16, torch.bfloat16, + torch.float8_e4m3fn, torch.float8_e4m3fnuz, torch.float8_e5m2, torch.float8_e5m2fnuz, + torch.int64, torch.int32, torch.int16, torch.int8, torch.uint64, torch.uint32, torch.uint16, torch.uint8, torch.bool, +) + + +class DispatchError(RuntimeError): + """``reasons`` maps every backend considered to why it was rejected, so no fallback is silent.""" + + def __init__(self, scheme: str, reasons: dict[str, str]): + self.scheme = scheme + self.reasons = dict(reasons) + detail = "; ".join(f"{name}: {reason}" for name, reason in self.reasons.items()) + super().__init__(f"No backend can handle '{scheme}'" + (f": {detail}" if detail else "")) + + +def register_scheme(scheme) -> Any: + with _lock: + _schemes[scheme.name] = scheme + return scheme + + +def get_scheme(name: str) -> Any: + try: + return _schemes[name] + except KeyError: + raise KeyError(f"Unknown scheme '{name}'; available: {sorted(_schemes)}") from None + + +def all_schemes() -> list: + return sorted(_schemes.values(), key=lambda scheme: (-scheme.priority, scheme.name)) + + +def name_for_config(config) -> str: + for cls in type(config).__mro__: + for scheme in _schemes.values(): + if scheme.config_cls is cls: + return scheme.name + raise TypeError(f"{type(config).__name__} is not the config of any registered scheme") + + +def require_dtype(dtype: torch.dtype) -> torch.dtype: + if dtype not in _CONTAINER_DTYPES: + raise ValueError(f"Unsupported tensor dtype {dtype}; supported: {sorted(_CONTAINER_DTYPES, key=str)}") + return dtype + + +def default_config(dtype, **overrides): + dtype = require_dtype(dtype) + scheme = next(scheme for scheme in all_schemes() if scheme.supports(dtype)) + return scheme.make_config({**scheme.options_for(dtype), **overrides}) + + +def _ordered() -> list[Backend]: + return sorted(_backends.values(), key=lambda backend: backend.priority, reverse=True) + + +def reason(backend: str) -> str | None: + spec = _backends.get(backend) + if spec is None: + return "not registered" + if backend not in _reasons: + with _lock: + if backend not in _reasons: + _reasons[backend] = spec.probe() if spec.probe is not None else None + return _reasons[backend] + + +def _rejection(name: str, scheme, dtype: torch.dtype | None) -> str | None: + if name not in _backends: + return "not registered" + if name not in scheme.lanes: + return f"scheme '{scheme.name}' declares no '{name}' lane" + unavailable = reason(name) + if unavailable is not None: + return unavailable + if dtype is not None and not scheme.supports(dtype): + return f"dtype {dtype} is not supported" + return None + + +def select(scheme, tensor: torch.Tensor | None = None, backend: str | None = None, + *, gate_dtype: bool = True) -> str: + dtype = tensor.dtype if (gate_dtype and tensor is not None) else None + if backend is not None: + rejected = _rejection(backend, scheme, dtype) + if rejected is not None: + raise DispatchError(scheme.name, {backend: rejected}) + logger.debug("Backend %s selected for %s", backend, scheme.name) + return backend + + failures: dict[str, str] = {} + for spec in _ordered(): + rejected = _rejection(spec.name, scheme, dtype) + if rejected is None: + logger.debug("Backend %s selected for %s", spec.name, scheme.name) + return spec.name + failures[spec.name] = rejected + raise DispatchError(scheme.name, failures) + + +def backend_device(backend: str | None) -> torch.device | None: + spec = _backends.get(backend) if backend is not None else None + return spec.device() if spec is not None and spec.device is not None else None diff --git a/entropack/schemes/__init__.py b/entropack/schemes/__init__.py new file mode 100644 index 0000000..a888689 --- /dev/null +++ b/entropack/schemes/__init__.py @@ -0,0 +1,12 @@ +from ..registry import all_schemes +from .base import Scheme +from .config import CompressionConfig, RawConfig +from .dfloat11 import DFloat11Config +from .lattice_rans import LatticeRANSConfig +from .tile_ans import TileANSConfig +from . import dfloat11, lattice_rans, tile_ans + +__all__ = [ + "CompressionConfig", "DFloat11Config", "LatticeRANSConfig", "RawConfig", "Scheme", "TileANSConfig", + "all_schemes", +] diff --git a/entropack/schemes/base.py b/entropack/schemes/base.py new file mode 100644 index 0000000..39c7bce --- /dev/null +++ b/entropack/schemes/base.py @@ -0,0 +1,121 @@ +import importlib +from abc import ABC, abstractmethod +from dataclasses import fields +from typing import Any, get_type_hints + +import torch + +from .config import CompressionConfig, RawConfig, validate_config +from ..registry import DispatchError, backend_device, register_scheme, select + +__all__ = ["RawScheme", "Scheme", "buffers_fingerprint", "cached_parse", "packed_buffers", "register_scheme"] + +_lane_modules: dict[tuple[type, str], Any] = {} +_config_classes: dict[type, type] = {} + + +def buffers_fingerprint(buffers: dict, shape, dtype) -> Any: + return ( + tuple(shape) if shape is not None else None, + dtype, + tuple( + (name, tensor.data_ptr(), tuple(tensor.shape), tuple(tensor.stride()), tensor.dtype, tensor.device) + for name, tensor in sorted(buffers.items()) + ), + ) + + +def packed_buffers(packed: dict, kind: type) -> Any: + return kind(**{name: packed[name] for name in kind._fields}) + + +def cached_parse(layout: torch.Tensor, parse, attribute: str): + cached = getattr(layout, attribute, None) + if cached is None: + cached = parse(layout) + setattr(layout, attribute, cached) + return cached + + +class Scheme(ABC): + name: str + buffer_names: tuple[str, ...] + lossless: bool = True + priority: int = 0 + dtypes: tuple[torch.dtype, ...] | None = None + lanes: dict[str, str] = {} + + def supports(self, dtype: torch.dtype) -> bool: + return self.dtypes is None or dtype in self.dtypes + + def options_for(self, dtype: torch.dtype) -> dict[str, Any]: + return {} + + def lane(self, backend: str) -> Any: + key = (type(self), backend) + module = _lane_modules.get(key) + if module is None: + if backend not in self.lanes: + raise DispatchError(self.name, {backend: f"scheme '{self.name}' declares no '{backend}' lane"}) + try: + module = importlib.import_module(f".{self.lanes[backend]}", package=type(self).__module__) + except Exception as error: + raise DispatchError(self.name, {backend: f"lane module failed to import: {error!r}"}) from error + _lane_modules[key] = module + return module + + def lane_for(self, tensor: torch.Tensor | None = None, backend: str | None = None, *, gate_dtype: bool = True): + name = select(self, tensor, None if backend in (None, "auto") else backend, gate_dtype=gate_dtype) + return self.lane(name), backend_device(name) + + @abstractmethod + def encode(self, weight: Any, config: CompressionConfig) -> dict: + ... + + @abstractmethod + def decode(self, packed: dict, *, shape: tuple[int, ...], dtype: Any, config: CompressionConfig) -> Any: + ... + + @property + def config_cls(self) -> type: + """The config class this scheme's ``encode`` annotation promises.""" + cached = _config_classes.get(type(self)) + if cached is None: + cached = get_type_hints(type(self).encode)["config"] + _config_classes[type(self)] = cached + return cached + + def make_config(self, options: dict) -> Any: + """Build this scheme's config from option names and values, rejecting both unknown and bad ones.""" + unknown = set(options) - {field.name for field in fields(self.config_cls)} + if unknown: + raise TypeError(f"Unknown {self.name} options: {sorted(unknown)}") + config = self.config_cls(**options) + validate_config(config) + return config + + def validate_buffers(self, buffers: dict, shape: tuple[int, ...], dtype: Any) -> None: + ... + + +class RawScheme(Scheme): + name = "raw" + buffer_names = ("data",) + priority = -1 + dtypes = None + lanes = {} + + def encode(self, weight, config: RawConfig) -> dict: + data = weight.detach().contiguous().clone() + return {"data": data} + + def validate_buffers(self, buffers, shape, dtype) -> None: + data = buffers["data"] + if data.dtype != dtype or tuple(data.shape) != tuple(shape): + raise ValueError(f"raw buffer must be {tuple(shape)} {dtype}, got {tuple(data.shape)} {data.dtype}") + + def decode(self, packed: dict, *, shape, dtype, config: RawConfig) -> torch.Tensor: + return packed["data"] + + +register_scheme(RawScheme()) diff --git a/entropack/schemes/checks.py b/entropack/schemes/checks.py new file mode 100644 index 0000000..69cc553 --- /dev/null +++ b/entropack/schemes/checks.py @@ -0,0 +1,33 @@ +from collections.abc import Iterable +from numbers import Real + +import torch + +__all__ = ["prepare_weight"] + +_NO_MINMAX_KERNEL = frozenset({ + torch.float8_e4m3fn, torch.float8_e4m3fnuz, torch.float8_e5m2, torch.float8_e5m2fnuz, +}) + + +def _all_finite(weight: torch.Tensor) -> bool: + if not weight.dtype.is_floating_point: + return True + probe = weight.float() if weight.dtype in _NO_MINMAX_KERNEL else weight + limit = float(torch.finfo(weight.dtype).max) + return float(probe.amin()) >= -limit and float(probe.amax()) <= limit + + +def prepare_weight( + weight: torch.Tensor, *, scheme: str, ndim: int | None = None, dtypes: Iterable[torch.dtype] | None = None, + require_finite: bool = False, +) -> torch.Tensor: + if dtypes is not None and weight.dtype not in frozenset(dtypes): + raise ValueError(f"{scheme} does not support dtype {weight.dtype}") + if ndim is not None and weight.ndim != ndim: + raise ValueError(f"{scheme} encode requires exactly {ndim}D input, got {weight.ndim}D") + if weight.numel() == 0: + raise ValueError(f"{scheme} does not support empty tensors") + if require_finite and not _all_finite(weight): + raise ValueError(f"{scheme} input must contain only finite values") + return weight.contiguous() diff --git a/entropack/schemes/config.py b/entropack/schemes/config.py new file mode 100644 index 0000000..ac24252 --- /dev/null +++ b/entropack/schemes/config.py @@ -0,0 +1,111 @@ +from dataclasses import dataclass, fields +from numbers import Real +from typing import Annotated, get_args, get_origin, get_type_hints + +__all__ = ["CompressionConfig", "OneOf", "Range", "RawConfig", "config_fields", "validate_config"] + + +@dataclass +class CompressionConfig: + """Common execution settings for compression schemes. + + Choose a concrete configuration class to select a scheme. Its fields control compression + or decompression as documented for that scheme.""" + + #: Execution backend: ``"auto"`` or ``None`` selects automatically, or use ``"cuda"`` or ``"eager"``. + execution_backend: str | None = "auto" + + +@dataclass +class RawConfig(CompressionConfig): + """Store tensor values without entropy coding. + + Used explicitly for uncoded FP8 or INT8 weights and as the tensor API's fallback when + encoding fails. Storage still includes container metadata.""" + + +class Range: + """An inclusive numeric bound. ``message`` overrides the generated text.""" + + def __init__(self, low=None, high=None, *, message: str | None = None): + self.low = low + self.high = high + self.message = message + + def holds(self, value) -> bool: + return (self.low is None or value >= self.low) and (self.high is None or value <= self.high) + + def __repr__(self) -> str: + return f"Range({self.low!r}, {self.high!r})" + + def describe(self, name: str) -> str: + if self.message is not None: + return self.message + if self.low is None: + return f"{name} must be <= {self.high}" + if self.high is None: + return f"{name} must be >= {self.low}" + return f"{name} must be in [{self.low}, {self.high}]" + + +class OneOf: + """A closed set of values. ``silent`` holds the ones that ask for automatic selection.""" + + def __init__(self, choices, *, silent=(), message: str | None = None): + self.choices = tuple(choices) + self.silent = frozenset(silent) + self.message = message + + def holds(self, value) -> bool: + return value in self.silent or value in self.choices + + def __repr__(self) -> str: + return f"OneOf({list(self.choices)!r})" + + def describe(self, name: str) -> str: + return self.message or f"{name} must be one of {self.choices}" + + +_hints: dict[type, dict] = {} + + +def config_fields(cls: type) -> dict: + cached = _hints.get(cls) + if cached is None: + hints = get_type_hints(cls, include_extras=True) + cached = {field.name: hints[field.name] for field in fields(cls) if field.name != "execution_backend"} + _hints[cls] = cached + return cached + + +def _check(name: str, value, hint) -> None: + if get_origin(hint) is Annotated: + hint, *markers = get_args(hint) + else: + markers = [] + union = get_args(hint) if get_origin(hint) is not None else (hint,) + kinds = tuple(kind for kind in union if kind is not type(None)) + optional = " or None" if len(kinds) < len(union) else "" + + if value is None: + if not optional: + raise TypeError(f"{name} must be {'an integer' if kinds[0] is int else 'a real number'}") + return + + if kinds[0] is int: + holds, expected = isinstance(value, int) and not isinstance(value, bool), "an integer" + elif kinds[0] is bool: + holds, expected = isinstance(value, bool), "a bool" + else: + holds, expected = isinstance(value, Real) and not isinstance(value, bool), "a real number" + if not holds: + raise TypeError(f"{name} must be {expected}{optional}") + + for marker in markers: + if not marker.holds(value): + raise ValueError(marker.describe(name)) + + +def validate_config(config) -> None: + for name, hint in config_fields(type(config)).items(): + _check(name, getattr(config, name), hint) diff --git a/entropack/schemes/dfloat11/__init__.py b/entropack/schemes/dfloat11/__init__.py new file mode 100644 index 0000000..9c744c1 --- /dev/null +++ b/entropack/schemes/dfloat11/__init__.py @@ -0,0 +1,52 @@ +import torch + +from ..base import Scheme, packed_buffers, register_scheme +from ..checks import prepare_weight +from .eager import get_32bit_codec, get_luts +from .format import PACKED_KEYS, DFloat11Buffers, DFloat11Config, validate_packed + + +class DFloat11Scheme(Scheme): + name = "dfloat11" + buffer_names = PACKED_KEYS + priority = 200 + dtypes = (torch.bfloat16,) + lanes = {"eager": "eager", "cuda": "cuda"} + + def encode(self, weight: torch.Tensor, config: DFloat11Config) -> dict: + weight = prepare_weight(weight, scheme=self.name, dtypes=self.dtypes) + device = weight.device + flat = weight.reshape(-1) + + lane, run_on = self.lane_for(flat, config.execution_backend) + if run_on is not None and device != run_on: + flat = flat.to(run_on) + counter = lane.exponent_counter(flat, config.threads_per_block) + codec, _counter, table = get_32bit_codec(counter) + luts = get_luts(table) + + buffers = lane.encode( + weight=flat, codec=codec, luts=luts, bytes_per_thread=config.bytes_per_thread, + threads_per_block=config.threads_per_block, + ) + + return {key: value.to(device) for key, value in buffers._asdict().items()} + + def validate_buffers(self, buffers, shape, dtype): + validate_packed(buffers, shape) + + def decode(self, packed: dict, *, shape: tuple[int, ...], dtype: torch.dtype, + config: DFloat11Config) -> torch.Tensor: + buffers = packed_buffers(packed, DFloat11Buffers) + source = buffers.layout.device + lane, lane_device = self.lane_for(buffers.layout, config.execution_backend, gate_dtype=False) + if lane_device is not None and source != lane_device: + buffers = DFloat11Buffers._make(value.to(lane_device) for value in buffers) + flat = lane.decode(buffers) + out = flat.reshape(shape) + return out if out.device == source else out.to(source) + + +register_scheme(DFloat11Scheme()) + +__all__ = ["DFloat11Config", "DFloat11Scheme"] diff --git a/entropack/schemes/dfloat11/cuda.py b/entropack/schemes/dfloat11/cuda.py new file mode 100644 index 0000000..7c7c391 --- /dev/null +++ b/entropack/schemes/dfloat11/cuda.py @@ -0,0 +1,223 @@ +from pathlib import Path + +import cupy +import numpy as np +import torch + +from ...backends.cuda import device as _device_caps +from ...backends.cuda.kernels import KernelLibrary +from ...backends.cuda.kernels import device_index as _device_index +from ...backends.cuda.kernels import ensure_dynamic_shared as _ensure_dynamic_shared +from ...backends.cuda.kernels import external_stream as _external_stream +from ...backends.cuda.kernels import pointer as _pointer +from .format import BLOCK_SIZE, MAX_RESIDENT_BLOCKS_PER_SM, DFloat11Buffers, make_layout, max_block_elems, parse_layout_cached + +_CUDA_PATH = Path(__file__).parent / "dfloat11.cu" +_DECODE_KERNEL_NAME = "dfloat11_decode_kernel" +_ENCODE_KERNEL_NAMES = ( + "dfloat11_exponent_histogram_kernel", "dfloat11_split_len_kernel", "dfloat11_pack_kernel", "dfloat11_thread_meta_kernel", + "dfloat11_output_positions_kernel", +) + + +def _compile_defines(device_index: int, threads_per_block: int) -> tuple[str, ...]: + caps = _device_caps.caps(device_index) + min_blocks = max(1, min(MAX_RESIDENT_BLOCKS_PER_SM, caps.threads_per_sm // threads_per_block)) + return (f"DFLOAT11_THREADS_PER_BLOCK={threads_per_block}", f"DFLOAT11_MIN_BLOCKS_PER_SM={min_blocks}") + + +_LIBRARY = KernelLibrary( + key="dfloat11", source=_CUDA_PATH, defines=_compile_defines, kernel_names=(_DECODE_KERNEL_NAME, *_ENCODE_KERNEL_NAMES), +) +_kernel = _LIBRARY.kernel + + +def _encode_kernels(device_index: int, threads_per_block: int) -> dict: + return {name: _kernel(device_index, threads_per_block, name) for name in _ENCODE_KERNEL_NAMES} + + +def _shared_budget(device: torch.device) -> int: + return _device_caps.caps(device).shared_limit() + + +def _kernel_tensor(t: torch.Tensor, stream: torch.cuda.Stream) -> torch.Tensor: + contiguous = t.contiguous() + if contiguous is not t: + contiguous.record_stream(stream) + return contiguous + + +def _cupy_view(t: torch.Tensor): + """Not cached: the DLPack capsule keeps the source tensor alive, so a cache would pin every tensor ever viewed for the life + of the process. + """ + return cupy.from_dlpack(t) + + +def decode(buffers: DFloat11Buffers) -> torch.Tensor: + encoded_exponent, sign_mantissa, luts = buffers.encoded_exponent, buffers.sign_mantissa, buffers.luts + output_positions, thread_meta, layout = buffers.output_positions, buffers.thread_meta, buffers.layout + bytes_per_thread, threads_per_block, max_block_elems = parse_layout_cached(layout) + num_luts = int(luts.shape[0]) + n_bytes = int(encoded_exponent.numel()) + n_elements = int(sign_mantissa.numel()) + + n_threads = (n_bytes + bytes_per_thread - 1) // bytes_per_thread + blocks = (n_threads + threads_per_block - 1) // threads_per_block + + budget = _shared_budget(sign_mantissa.device) + fixed_bytes = threads_per_block * 4 + num_luts * 256 + stage_enc_bytes = threads_per_block * bytes_per_thread + 8 + stage_elems = max_block_elems + if fixed_bytes + stage_enc_bytes + stage_elems + 4 > budget: + stage_elems = 0 + if fixed_bytes + stage_enc_bytes > budget: + stage_enc_bytes = 0 + shared_bytes = fixed_bytes + stage_enc_bytes + stage_elems + (4 if stage_elems else 0) + + out = torch.empty(n_elements, dtype=torch.bfloat16, device=sign_mantissa.device) + if blocks == 0: + return out + + device_index = _device_index(sign_mantissa) + with cupy.cuda.Device(device_index): + kernel = _kernel(device_index, threads_per_block, _DECODE_KERNEL_NAME) + _ensure_dynamic_shared(kernel, shared_bytes) + + torch_stream = torch.cuda.current_stream(sign_mantissa.device) + luts_c = _kernel_tensor(luts, torch_stream) + encoded_c = _kernel_tensor(encoded_exponent, torch_stream) + sign_mantissa_c = _kernel_tensor(sign_mantissa, torch_stream) + output_positions_c = _kernel_tensor(output_positions, torch_stream) + thread_meta_c = _kernel_tensor(thread_meta, torch_stream) + args = ( + _pointer(luts_c), _pointer(encoded_c), _pointer(sign_mantissa_c), _pointer(output_positions_c), + _pointer(thread_meta_c), _pointer(out), np.int32(num_luts), np.int64(n_bytes), np.int64(n_elements), + np.int32(bytes_per_thread), np.int32(stage_enc_bytes), np.int32(stage_elems), + np.int32(1 if (sign_mantissa_c.data_ptr() & 3) == 0 else 0), + ) + with _external_stream(torch_stream): + kernel((blocks,), (threads_per_block,), args, shared_mem=shared_bytes) + return out + + +def exponent_counter(weight: torch.Tensor, threads_per_block: int) -> dict[int, int]: + flat = weight.reshape(-1) + n_elements = int(flat.numel()) + histogram = torch.zeros(256, dtype=torch.int64, device=flat.device) + if n_elements == 0: + return {} + + caps = _device_caps.caps(flat.device) + threads = caps.threads_per_block(BLOCK_SIZE) + blocks = caps.grid((n_elements + threads - 1) // threads, threads) + device_index = _device_index(flat) + torch_stream = torch.cuda.current_stream(flat.device) + stream_ctx = _external_stream(torch_stream) + with cupy.cuda.Device(device_index), stream_ctx: + _kernel(device_index, threads_per_block, "dfloat11_exponent_histogram_kernel")( + (blocks,), (threads,), (_pointer(flat), _pointer(histogram), np.int64(n_elements)), + ) + counts = histogram.cpu().tolist() + return {i: int(count) for i, count in enumerate(counts) if count > 0} + + +def _code_table(codec, device): + code_len = torch.zeros(256, dtype=torch.int32) + code_val = torch.zeros(256, dtype=torch.int32) + for k, (bits, val) in codec._table.items(): + if isinstance(k, int): + code_len[k] = bits + code_val[k] = val + eof_len, eof_val = codec._table[codec._eof] + return (code_len.to(device).contiguous(), code_val.to(device).contiguous(), int(eof_len), int(eof_val)) + + +def encode( + *, weight: torch.Tensor, codec, luts: torch.Tensor, bytes_per_thread: int, threads_per_block: int, +) -> DFloat11Buffers: + device = weight.device + device_index = _device_index(device) + flat = weight.reshape(-1) + n_elements = int(flat.numel()) + if not 5 <= bytes_per_thread <= 255: + raise ValueError( + "dfloat11 bytes_per_thread must be in [5, 255]: the region must exceed " + "the 32-bit maximum code length and its worst-case symbol count must fit " "the 11-bit thread_meta field" + ) + if n_elements >= 1 << 32: + raise ValueError("dfloat11 requires fewer than 2^32 elements per tensor") + code_len_gpu, code_val_gpu, eof_len, eof_val = _code_table(codec, device) + + kernels = _encode_kernels(device_index, threads_per_block) + threads = _device_caps.caps(device).threads_per_block(BLOCK_SIZE) + + exponent = torch.empty(n_elements, dtype=torch.uint8, device=device) + sign_mantissa = torch.empty(n_elements, dtype=torch.uint8, device=device) + len_scratch = torch.empty(n_elements, dtype=torch.uint8, device=device) + pref = torch.empty(n_elements + 1, dtype=torch.int64, device=device) + + blocks_split = (n_elements + threads - 1) // threads + + torch_stream = torch.cuda.current_stream(flat.device) + stream_ctx = _external_stream(torch_stream) + with cupy.cuda.Device(device_index), stream_ctx: + kernels["dfloat11_split_len_kernel"]( + (blocks_split,), + (threads,), + ( + _pointer(flat), _pointer(code_len_gpu), _pointer(exponent), _pointer(sign_mantissa), _pointer(len_scratch), + np.int64(n_elements), + ), + ) + + pref_cp = _cupy_view(pref) + pref_cp[0] = 0 + if n_elements > 0: + cupy.cumsum(_cupy_view(len_scratch), dtype=cupy.int64, out=pref_cp[1:]) + + total_bits = int(pref[n_elements].item()) + + n_bytes = (total_bits + 7) // 8 + region_bits = bytes_per_thread * 8 + block_bits = region_bits * threads_per_block + bytes_per_block = bytes_per_thread * threads_per_block + num_blocks = (n_bytes + bytes_per_block - 1) // bytes_per_block + n_regions = threads_per_block * num_blocks + + encoded = torch.zeros(n_bytes, dtype=torch.uint8, device=device) + thread_meta = torch.empty(n_regions, dtype=torch.uint16, device=device) + output_positions = torch.empty(num_blocks + 1, dtype=torch.uint32, device=device) + + with cupy.cuda.Device(device_index), stream_ctx: + pack_chunks = (n_bytes + 3) // 4 + kernels["dfloat11_pack_kernel"]( + ((pack_chunks + threads - 1) // threads,), + (threads,), + ( + _pointer(pref), _pointer(exponent), _pointer(code_len_gpu), _pointer(code_val_gpu), _pointer(encoded), + np.int64(n_elements), np.int64(n_bytes), np.int64(total_bits), np.int32(eof_len), np.uint32(eof_val), + ), + ) + kernels["dfloat11_thread_meta_kernel"]( + ((n_regions + threads - 1) // threads,), + (threads,), + ( + _pointer(pref), _pointer(thread_meta), np.int64(n_elements), np.int64(total_bits), np.int64(region_bits), + np.int64(n_regions), + ), + ) + kernels["dfloat11_output_positions_kernel"]( + ((num_blocks + 1 + threads - 1) // threads,), (threads,), + (_pointer(pref), _pointer(output_positions), np.int64(n_elements), np.int64(block_bits), np.int64(num_blocks)), + ) + + op_host = torch.empty(num_blocks + 1, dtype=torch.uint32, pin_memory=True) + op_host.copy_(output_positions, non_blocking=True) + torch_stream.synchronize() + + layout = make_layout(bytes_per_thread, threads_per_block, max_block_elems(op_host)) + return DFloat11Buffers( + encoded_exponent=encoded, sign_mantissa=sign_mantissa, luts=luts, output_positions=output_positions, + thread_meta=thread_meta, layout=layout, + ) diff --git a/entropack/schemes/dfloat11/dfloat11.cu b/entropack/schemes/dfloat11/dfloat11.cu new file mode 100644 index 0000000..9bafe8d --- /dev/null +++ b/entropack/schemes/dfloat11/dfloat11.cu @@ -0,0 +1,740 @@ +// EntroPack -- DFloat11 lossless bf16 compression: CUDA kernels. +// +// Decode kernel: +// dfloat11_decode_kernel one thread per bitstream region: walks the multi-level Huffman LUTs to recover the +// exponents, stages them in shared memory, and writes the reconstructed bf16 back +// coalesced. __launch_bounds__ come from -D macros the host derives from the queried +// device. +// +// Encode kernels, in pipeline order: +// dfloat11_exponent_histogram_kernel 256-bin bf16 exponent counts, accumulated in shared memory per block. +// dfloat11_split_len_kernel splits bf16 into exponent and sign/mantissa bytes and looks up each exponent's +// Huffman code length in the same pass. +// dfloat11_pack_kernel writes the MSB-first bitstream, one thread per short output chunk. +// dfloat11_thread_meta_kernel per-region gap and symbol count, as the checkpoint's uint16 thread_meta. +// dfloat11_output_positions_kernel per-block starting element index. +// +// The device helpers below serve both directions: MSB-first bit windows, the multi-level LUT walk, and a block-wide +// exclusive scan. Python orchestration lives in cuda.py; format.py is authoritative for the buffer layout and for the +// decode granularity recorded in each checkpoint. NVRTC has no system include path, so libcudacxx supplies the +// uint8_t/uint32_t/int64_t typedefs that would. +#include + +constexpr int kMetaCountBits = 11; +constexpr uint32_t kMetaCountMask = (1u << kMetaCountBits) - 1u; + +#ifndef DFLOAT11_THREADS_PER_BLOCK +#define DFLOAT11_THREADS_PER_BLOCK 128 +#endif +#ifndef DFLOAT11_MIN_BLOCKS_PER_SM +#define DFLOAT11_MIN_BLOCKS_PER_SM 1 +#endif + + +__device__ __forceinline__ uint32_t read_byte_msb( + const uint8_t* __restrict__ data, int64_t n_bytes, int64_t bit_pos) { + const int64_t byte_idx = bit_pos >> 3; + const uint32_t shift = static_cast(bit_pos & 7); + const uint32_t hi = (byte_idx < n_bytes) ? data[byte_idx] : 0u; + if (shift == 0u) { + return hi; + } + const uint32_t lo = (byte_idx + 1 < n_bytes) ? data[byte_idx + 1] : 0u; + return ((hi << shift) | (lo >> (8u - shift))) & 0xFFu; +} + +__device__ __forceinline__ uint32_t decode_symbol( + const uint8_t* __restrict__ luts, + int num_luts, // == num_levels + 1 + const uint8_t* __restrict__ encoded, + int64_t n_bytes, + int64_t bit_pos, + uint32_t* code_len) { + const int num_levels = num_luts - 1; + const uint32_t ptr_min = + (num_levels > 1) ? static_cast(256 - (num_levels - 1)) : 256u; + const uint8_t* lens_row = luts + static_cast(num_levels) * 256; + + int level = 0; + int hop = 0; + for (;;) { + const uint32_t byte = read_byte_msb(encoded, n_bytes, bit_pos + hop * 8); + const uint32_t entry = luts[static_cast(level) * 256 + byte]; + if (num_levels > 1 && entry >= ptr_min) { + level = 256 - static_cast(entry); // pointer -> child level + ++hop; + } else { + *code_len = lens_row[entry]; // leaf symbol + return entry; + } + } +} + +__device__ __forceinline__ uint16_t make_bf16_bits(uint32_t exponent, uint32_t sign_mantissa) { + return static_cast( + ((sign_mantissa & 0x80u) << 8) | (exponent << 7) | (sign_mantissa & 0x7Fu)); +} + +__device__ __forceinline__ uint64_t pack_bf16x4(uint32_t e32, uint32_t s32) { + const uint32_t lo = ((e32 & 0x01010101u) << 7) | (s32 & 0x7F7F7F7Fu); + const uint32_t hi = (s32 & 0x80808080u) | ((e32 >> 1) & 0x7F7F7F7Fu); + uint32_t w0, w1; + asm("prmt.b32 %0, %1, %2, 0x5140;" : "=r"(w0) : "r"(lo), "r"(hi)); + asm("prmt.b32 %0, %1, %2, 0x7362;" : "=r"(w1) : "r"(lo), "r"(hi)); + return static_cast(w0) | (static_cast(w1) << 32); +} + +// A sliding MSB-first bit window held in registers. `decode_symbol` re-reads the stream for every symbol, and again for each +// LUT hop because a codeword is not byte-aligned. Walking a contiguous run instead keeps the next bits in `buf` and refills a +// byte at a time, so the hot loop touches the stream once per byte consumed. +// +// The `BitWindowLocal` variant stages the block's slice in shared memory, zero-padded and sized for the worst-case lookahead, +// so the generic refill's bounds check is never taken and is dropped, and addressing shrinks to 32 bit. +struct BitWindow { + const uint8_t* __restrict__ data; + int64_t n_bytes; + int64_t byte_pos; // next byte to pull into `buf` + uint64_t buf; + int n_valid; +}; + +__device__ __forceinline__ void bit_window_refill(BitWindow* w) { + while (w->n_valid <= 32) { + const uint32_t byte = (w->byte_pos < w->n_bytes) ? w->data[w->byte_pos] : 0u; + w->buf |= static_cast(byte) << (56 - w->n_valid); + w->n_valid += 8; + ++w->byte_pos; + } +} + +__device__ __forceinline__ BitWindow bit_window_open( + const uint8_t* __restrict__ data, int64_t n_bytes, int64_t bit_pos) { + BitWindow w; + w.data = data; + w.n_bytes = n_bytes; + w.byte_pos = bit_pos >> 3; + w.buf = 0; + w.n_valid = 0; + bit_window_refill(&w); + const int skip = static_cast(bit_pos & 7); // land on the first bit of the codeword + w.buf <<= skip; + w.n_valid -= skip; + bit_window_refill(&w); + return w; +} + +struct LutWalk { + const uint8_t* luts; // staged decode rows + const uint8_t* lens_row; // per-symbol code length row + uint32_t ptr_min; // entries >= ptr_min are pointers, below are leaf symbols + int num_levels; +}; + +__device__ __forceinline__ uint32_t bit_window_decode( + BitWindow* w, + const uint8_t* __restrict__ luts, + int num_luts) { // == num_levels + 1 + const int num_levels = num_luts - 1; + const uint32_t ptr_min = + (num_levels > 1) ? static_cast(256 - (num_levels - 1)) : 256u; + const uint8_t* lens_row = luts + static_cast(num_levels) * 256; + + int level = 0; + int shift = 56; // peek the byte `hop` bytes into the window without consuming it + for (;;) { + const uint32_t byte = static_cast((w->buf >> shift) & 0xFFu); + const uint32_t entry = luts[static_cast(level) * 256 + byte]; + if (num_levels > 1 && entry >= ptr_min) { + level = 256 - static_cast(entry); // pointer -> child level + shift -= 8; + } else { + const int code_len = lens_row[entry]; // length of the whole codeword + w->buf <<= code_len; + w->n_valid -= code_len; + bit_window_refill(w); + return entry; + } + } +} + +struct BitWindowLocal { + const uint8_t* data; + int byte_pos; + uint64_t buf; + int n_valid; +}; + +__device__ __forceinline__ void bwl_refill(BitWindowLocal* w) { + while (w->n_valid <= 32) { + w->buf |= static_cast(w->data[w->byte_pos]) << (56 - w->n_valid); + w->n_valid += 8; + ++w->byte_pos; + } +} + +__device__ __forceinline__ BitWindowLocal bwl_open(const uint8_t* data, int bit_pos) { + BitWindowLocal w; + w.data = data; + w.byte_pos = bit_pos >> 3; + w.buf = 0; + w.n_valid = 0; + bwl_refill(&w); + const int skip = bit_pos & 7; + w.buf <<= skip; + w.n_valid -= skip; + bwl_refill(&w); + return w; +} + +__device__ __forceinline__ uint32_t bwl_decode(BitWindowLocal* w, const LutWalk* ctx) { + int level = 0; + int shift = 56; + for (;;) { + const uint32_t byte = static_cast((w->buf >> shift) & 0xFFu); + const uint32_t entry = ctx->luts[level * 256 + byte]; + if (ctx->num_levels > 1 && entry >= ctx->ptr_min) { + level = 256 - static_cast(entry); + shift -= 8; + } else { + const int code_len = ctx->lens_row[entry]; + w->buf <<= code_len; + w->n_valid -= code_len; + bwl_refill(w); + return entry; + } + } +} + +__device__ __forceinline__ uint32_t bswap32(uint32_t x) { + uint32_t r; + asm("prmt.b32 %0, %1, 0, 0x0123;" : "=r"(r) : "r"(x)); + return r; +} + +// Register-resident bitstream for bytes_per_thread == 16: a thread's whole span, its region plus the worst-case lookahead, is +// preloaded into three big-endian u64 words, so the walk issues no stream loads, only funnel shifts. +struct RegStream { + uint64_t w0, w1, w2; + int bit; +}; + +__device__ __forceinline__ RegStream reg_stream_open(const uint8_t* span16, int gap_bits) { + const uint4 q = *reinterpret_cast(span16); + const uint2 r = *reinterpret_cast(span16 + 16); + RegStream s; + s.w0 = (static_cast(bswap32(q.x)) << 32) | bswap32(q.y); + s.w1 = (static_cast(bswap32(q.z)) << 32) | bswap32(q.w); + s.w2 = (static_cast(bswap32(r.x)) << 32) | bswap32(r.y); + s.bit = gap_bits; + return s; +} + +__device__ __forceinline__ uint32_t rs_peek32(const RegStream* s) { + const int b = s->bit; + const uint64_t v = (s->w0 << b) | (b ? (s->w1 >> (64 - b)) : 0ull); + return static_cast(v >> 32); +} + +__device__ __forceinline__ uint32_t rs_decode(RegStream* s, const LutWalk* ctx) { + uint32_t peek = rs_peek32(s); + int lvl_off = 0; + for (;;) { + const uint32_t byte = peek >> 24; + const uint32_t entry = ctx->luts[lvl_off + byte]; + if (entry >= ctx->ptr_min) { // pointer -> child level (ptr_min == 256: never true) + lvl_off = (256u - entry) << 8; + peek <<= 8; + } else { + const int code_len = ctx->lens_row[entry]; + s->bit += code_len; + if (s->bit >= 64) { // rotate the word window; bit + code_len < 95 < 128, once is enough + s->w0 = s->w1; + s->w1 = s->w2; + s->w2 = 0ull; + s->bit -= 64; + } + return entry; + } + } +} + +// Warp-shuffle scan per warp, then a scan of the warp totals: two barriers instead of a full Hillis-Steele sweep over the +// block. +__device__ __forceinline__ int32_t block_exclusive_scan(int32_t* s_scan, int32_t value) { + if ((blockDim.x & 31u) == 0u && blockDim.x >= 32u) { + const unsigned full = 0xFFFFFFFFu; + const int lane = static_cast(threadIdx.x) & 31; + const int warp_id = static_cast(threadIdx.x) >> 5; + const int nwarps = static_cast(blockDim.x) >> 5; + + int32_t v = value; // inclusive scan within the warp +#pragma unroll + for (int off = 1; off < 32; off <<= 1) { + const int32_t n = __shfl_up_sync(full, v, off); + if (lane >= off) v += n; + } + if (lane == 31) { + s_scan[warp_id] = v; + } + __syncthreads(); + + if (warp_id == 0) { + const unsigned wmask = (nwarps == 32) ? full : ((1u << nwarps) - 1u); + int32_t w = (lane < nwarps) ? s_scan[lane] : 0; +#pragma unroll + for (int off = 1; off < 32; off <<= 1) { + const int32_t n = __shfl_up_sync(wmask, w, off); + if (lane >= off && lane < nwarps) w += n; + } + if (lane < nwarps) { + s_scan[lane] = w; + } + } + __syncthreads(); + + const int32_t prefix = (warp_id > 0) ? s_scan[warp_id - 1] : 0; + return prefix + (v - value); + } + + s_scan[threadIdx.x] = value; + __syncthreads(); + for (unsigned offset = 1; offset < blockDim.x; offset <<= 1) { + const int32_t addend = (threadIdx.x >= offset) ? s_scan[threadIdx.x - offset] : 0; + __syncthreads(); // every read of round `offset` completes before any write + if (threadIdx.x >= offset) { + s_scan[threadIdx.x] += addend; + } + __syncthreads(); + } + return s_scan[threadIdx.x] - value; // inclusive -> exclusive +} + + +extern "C" __global__ void __launch_bounds__(DFLOAT11_THREADS_PER_BLOCK, DFLOAT11_MIN_BLOCKS_PER_SM) dfloat11_decode_kernel( + const uint8_t* __restrict__ luts, + const uint8_t* __restrict__ encoded, + const uint8_t* __restrict__ sign_mantissa, + const uint32_t* __restrict__ output_positions, + const uint16_t* __restrict__ thread_meta, + uint16_t* __restrict__ out_bits, // bf16 bit patterns + int num_luts, + int64_t n_bytes, + int64_t n_elements, + int bytes_per_thread, + int stage_enc_bytes, // 0 = read the bitstream straight from global memory + int stage_elems, // 0 = scatter to out_bits directly instead of staging + int sm_u32_aligned) { // 1 = sign_mantissa pointer is 4-byte aligned (vector write-back) + extern __shared__ __align__(16) uint8_t s_raw[]; + int32_t* const s_scan = reinterpret_cast(s_raw); + uint8_t* const s_luts = s_raw + blockDim.x * sizeof(int32_t); + uint8_t* const s_enc = s_luts + static_cast(num_luts) * 256; + + const int64_t out_base = static_cast(output_positions[blockIdx.x]); + uint8_t* const s_exp = s_enc + stage_enc_bytes + (out_base & 3); + + const int lut_bytes = num_luts * 256; + for (int i = threadIdx.x; i < lut_bytes; i += blockDim.x) { + s_luts[i] = luts[i]; + } + + const int64_t block_byte_base = + static_cast(blockIdx.x) * blockDim.x * bytes_per_thread; + const bool direct_ok = (bytes_per_thread == 16) && (stage_enc_bytes > 0) && + (block_byte_base + stage_enc_bytes <= n_bytes) && + ((reinterpret_cast(encoded + block_byte_base) & 15u) == 0u); + if (stage_enc_bytes > 0 && !direct_ok && + ((reinterpret_cast(encoded + block_byte_base) & 15u) == 0u) && + blockDim.x * 16 >= stage_enc_bytes) { + const uint4* src4 = reinterpret_cast(encoded + block_byte_base); + uint4* dst4 = reinterpret_cast(s_enc); + const int n4 = stage_enc_bytes >> 4; // floor; the tail is handled below + if (static_cast(threadIdx.x) < n4) { + const int64_t src = block_byte_base + (static_cast(threadIdx.x) << 4); + uint4 v; + if (src + 16 <= n_bytes) { + v = src4[threadIdx.x]; + } else { // partial sector: rebuild byte-exactly like the scalar path (zero padding) + uint8_t tmp[16]; +#pragma unroll + for (int k = 0; k < 16; ++k) { + tmp[k] = (src + k < n_bytes) ? encoded[src + k] : 0u; + } + v.x = tmp[0] | (tmp[1] << 8) | (tmp[2] << 16) | (static_cast(tmp[3]) << 24); + v.y = tmp[4] | (tmp[5] << 8) | (tmp[6] << 16) | (static_cast(tmp[7]) << 24); + v.z = tmp[8] | (tmp[9] << 8) | (tmp[10] << 16) | (static_cast(tmp[11]) << 24); + v.w = tmp[12] | (tmp[13] << 8) | (tmp[14] << 16) | (static_cast(tmp[15]) << 24); + } + dst4[threadIdx.x] = v; + } + for (int i = (n4 << 4) + threadIdx.x; i < stage_enc_bytes; i += blockDim.x) { + const int64_t src = block_byte_base + i; + s_enc[i] = (src < n_bytes) ? encoded[src] : 0u; + } + } else if (stage_enc_bytes > 0 && !direct_ok) { + for (int i = threadIdx.x; i < stage_enc_bytes; i += blockDim.x) { + const int64_t src = block_byte_base + i; + s_enc[i] = (src < n_bytes) ? encoded[src] : 0u; + } + } + + const int64_t gt = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const uint32_t meta = thread_meta[gt]; + const int64_t start_bit = + gt * static_cast(bytes_per_thread) * 8 + (meta >> kMetaCountBits); + int32_t count = static_cast(meta & kMetaCountMask); + + __syncthreads(); // s_luts / s_enc are now visible to the whole block + + const int32_t local_base = block_exclusive_scan(s_scan, count); + + if (stage_enc_bytes > 0 && stage_elems > 0) { + if (local_base + count > stage_elems) { // only reachable on a corrupt buffer + count = (local_base < stage_elems) ? (stage_elems - local_base) : 0; + } + + LutWalk ctx; + ctx.luts = s_luts; + ctx.num_levels = num_luts - 1; + ctx.ptr_min = + (ctx.num_levels > 1) ? static_cast(256 - (ctx.num_levels - 1)) : 256u; + ctx.lens_row = s_luts + ctx.num_levels * 256; + + const int cursor = static_cast(start_bit - block_byte_base * 8); + if (bytes_per_thread == 16) { + const uint8_t* const span = + direct_ok ? (encoded + block_byte_base) : s_enc; + RegStream rs = reg_stream_open(span + (static_cast(threadIdx.x) << 4), + static_cast(meta >> kMetaCountBits)); + for (int32_t i = 0; i < count; ++i) { + s_exp[local_base + i] = static_cast(rs_decode(&rs, &ctx)); + } + } else { + BitWindowLocal window = bwl_open(s_enc, cursor); + for (int32_t i = 0; i < count; ++i) { + s_exp[local_base + i] = static_cast(bwl_decode(&window, &ctx)); + } + } + __syncthreads(); + + const int64_t block_elems = + static_cast(output_positions[blockIdx.x + 1]) - out_base; + + if (sm_u32_aligned && block_elems >= 8) { + const int64_t p = (8 - (out_base & 7)) & 7; // first 8-aligned element index + int64_t j = threadIdx.x; + for (; j < block_elems && j < p; j += blockDim.x) { + const int64_t gidx = out_base + j; + if (gidx < n_elements) { + out_bits[gidx] = make_bf16_bits(s_exp[j], sign_mantissa[gidx]); + } + } + const int64_t nvec = (block_elems - p) >> 3; + const bool sm_u64_aligned = ((reinterpret_cast(sign_mantissa) & 7u) == 0u); + const bool in_bounds = out_base + block_elems <= n_elements; + for (int64_t c = threadIdx.x; c < nvec; c += blockDim.x) { + const int64_t v = p + (c << 3); + const int64_t gidx = out_base + v; + const uint32_t e0123 = *reinterpret_cast(s_exp + v); + const uint32_t e4567 = *reinterpret_cast(s_exp + v + 4); + uint32_t s0123, s4567; + if (sm_u64_aligned) { + const uint2 s8 = *reinterpret_cast(sign_mantissa + gidx); + s0123 = s8.x; + s4567 = s8.y; + } else { + s0123 = *reinterpret_cast(sign_mantissa + gidx); + s4567 = *reinterpret_cast(sign_mantissa + gidx + 4); + } + const uint64_t lo4 = pack_bf16x4(e0123, s0123); + const uint64_t hi4 = pack_bf16x4(e4567, s4567); + if (in_bounds) { + uint4 o; + o.x = static_cast(lo4); + o.y = static_cast(lo4 >> 32); + o.z = static_cast(hi4); + o.w = static_cast(hi4 >> 32); + *reinterpret_cast(out_bits + gidx) = o; + } else { // ragged final chunk of the final block +#pragma unroll + for (int k = 0; k < 8; ++k) { + if (gidx + k < n_elements) { + out_bits[gidx + k] = make_bf16_bits(s_exp[v + k], sign_mantissa[gidx + k]); + } + } + } + } + const int64_t tail = (block_elems - p) & 7; + const int64_t vtail = p + (nvec << 3); + if (threadIdx.x < tail) { + const int64_t gidx = out_base + vtail + threadIdx.x; + if (gidx < n_elements) { + out_bits[gidx] = make_bf16_bits(s_exp[vtail + threadIdx.x], sign_mantissa[gidx]); + } + } + return; + } + + for (int64_t j = threadIdx.x; j < block_elems; j += blockDim.x) { + const int64_t gidx = out_base + j; + if (gidx < n_elements) { + out_bits[gidx] = make_bf16_bits(s_exp[j], sign_mantissa[gidx]); + } + } + return; + } + + const uint8_t* const stream = (stage_enc_bytes > 0) ? s_enc : encoded; + const int64_t stream_bytes = (stage_enc_bytes > 0) ? stage_enc_bytes : n_bytes; + const int64_t cursor = (stage_enc_bytes > 0) ? (start_bit - block_byte_base * 8) : start_bit; + + if (stage_elems > 0) { + if (local_base + count > stage_elems) { // only reachable on a corrupt buffer + count = (local_base < stage_elems) ? (stage_elems - local_base) : 0; + } + BitWindow window = bit_window_open(stream, stream_bytes, cursor); + for (int32_t i = 0; i < count; ++i) { + s_exp[local_base + i] = + static_cast(bit_window_decode(&window, s_luts, num_luts)); + } + __syncthreads(); + + const int64_t block_elems = + static_cast(output_positions[blockIdx.x + 1]) - out_base; + for (int64_t j = threadIdx.x; j < block_elems; j += blockDim.x) { + const int64_t gidx = out_base + j; + if (gidx < n_elements) { + out_bits[gidx] = make_bf16_bits(s_exp[j], sign_mantissa[gidx]); + } + } + return; + } + + const int64_t base = out_base + local_base; + BitWindow window = bit_window_open(stream, stream_bytes, cursor); + for (int32_t i = 0; i < count; ++i) { + const uint32_t exponent = bit_window_decode(&window, s_luts, num_luts); + const int64_t gidx = base + i; + if (gidx < n_elements) { + out_bits[gidx] = make_bf16_bits(exponent, sign_mantissa[gidx]); + } + } +} + + +// Each block accumulates in shared uint32 bins and then contributes at most one global uint64 atomic per bin, so the encode +// never materializes a per-element exponent array, which would be larger than the weight being encoded. +extern "C" __global__ void dfloat11_exponent_histogram_kernel( + const uint16_t* __restrict__ bf16_bits, + uint64_t* __restrict__ histogram, + int64_t n_elements) { + __shared__ uint32_t bins[256]; + if (threadIdx.x < 256) { + bins[threadIdx.x] = 0; + } + __syncthreads(); + + int64_t idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t stride = static_cast(gridDim.x) * blockDim.x; + for (; idx < n_elements; idx += stride) { + const uint32_t exponent = (bf16_bits[idx] >> 7) & 0xFFu; + atomicAdd(&bins[exponent], 1u); + } + __syncthreads(); + + if (threadIdx.x < 256 && bins[threadIdx.x] != 0) { + atomicAdd(reinterpret_cast(histogram + threadIdx.x), + static_cast(bins[threadIdx.x])); + } +} + + +extern "C" __global__ void dfloat11_split_len_kernel( + const uint16_t* __restrict__ bf16_bits, + const int32_t* __restrict__ code_len, // 256 entries + uint8_t* __restrict__ exponent, + uint8_t* __restrict__ sign_mantissa, + uint8_t* __restrict__ len_out, + int64_t n_elements) { + const int64_t idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + if (idx >= n_elements) { + return; + } + const uint32_t bits = bf16_bits[idx]; + const uint32_t exp = (bits >> 7) & 0xFFu; + exponent[idx] = static_cast(exp); + sign_mantissa[idx] = + static_cast(((bits >> 8) & 0x80u) | (bits & 0x7Fu)); + len_out[idx] = static_cast(code_len[exp]); +} + + +__device__ __forceinline__ int64_t lower_bound_i64( + const int64_t* __restrict__ a, int64_t m, int64_t key) { + int64_t lo = 0; + int64_t hi = m; + while (lo < hi) { + const int64_t mid = (lo + hi) >> 1; + if (a[mid] < key) { + lo = mid + 1; + } else { + hi = mid; + } + } + return lo; +} + + +__device__ __forceinline__ uint32_t byte_contrib(uint64_t vL, int64_t byte_start, int64_t p) { + const int shift = static_cast(byte_start - p); // in [-7, 31] for overlapping codes + const int sh = 24 - shift; + if (sh >= 0) { + return static_cast((vL >> sh) & 0xFFu); + } + return static_cast((vL << (-sh)) & 0xFFu); +} + + +// Consecutive output bytes walk the same monotone prefix array, so a thread pays one lower_bound per chunk and then advances +// the symbol cursor linearly. The chunk is wide enough that binary searches amortize over several output bytes and narrow +// enough that the cursor advance stays short and the thread count stays high. +constexpr int kPackBytesPerThread = 4; + +extern "C" __global__ void dfloat11_pack_kernel( + const int64_t* __restrict__ pref, // n_elements + 1 exclusive prefix sums of code lengths + const uint8_t* __restrict__ exponent, + const int32_t* __restrict__ code_len, // 256 + const int32_t* __restrict__ code_val, // 256 + uint8_t* __restrict__ encoded, + int64_t n_elements, + int64_t n_bytes, + int64_t total_bits, + int eof_len, + uint32_t eof_val) { + const int64_t chunk = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t first_byte = chunk * kPackBytesPerThread; + if (first_byte >= n_bytes) { + return; + } + const int byte_count = static_cast( + (n_bytes - first_byte < kPackBytesPerThread) ? (n_bytes - first_byte) + : kPackBytesPerThread); + + const int64_t first_bit = first_byte * 8; + int64_t i_start = lower_bound_i64(pref, n_elements + 1, first_bit + 1) - 1; + if (i_start < 0) { + i_start = 0; + } + + uint64_t packed = 0; +#pragma unroll + for (int k = 0; k < kPackBytesPerThread; ++k) { + if (k >= byte_count) { + break; + } + const int64_t j = first_byte + k; + const int64_t byte_start = j * 8; + const int64_t byte_end = byte_start + 8; + + while (i_start < n_elements) { + const uint32_t sym = exponent[i_start]; + const int64_t end = pref[i_start] + code_len[sym]; + if (end > byte_start) { + break; + } + ++i_start; + } + + uint32_t acc = 0; + for (int64_t i = i_start; i < n_elements && pref[i] < byte_end; ++i) { + const int64_t p = pref[i]; + const uint32_t sym = exponent[i]; + const int b = code_len[sym]; + const int64_t end = p + b; + if (end <= byte_start) { + continue; + } + const uint64_t vL = + static_cast(static_cast(code_val[sym])) << (32 - b); + acc |= byte_contrib(vL, byte_start, p); + } + + if (j == n_bytes - 1) { + const int64_t r = total_bits - byte_start; + if (r > 0 && r < 8) { + const uint64_t eL = static_cast(eof_val) << (32 - eof_len); + acc |= byte_contrib(eL, byte_start, total_bits); + } + } + packed |= static_cast(acc) << (k * 8); + } + + if (byte_count == kPackBytesPerThread) { + if constexpr (kPackBytesPerThread == 8) { + *reinterpret_cast(encoded + first_byte) = packed; + } else if constexpr (kPackBytesPerThread == 4) { + *reinterpret_cast(encoded + first_byte) = static_cast(packed); + } else { + *reinterpret_cast(encoded + first_byte) = static_cast(packed); + } + } else { +#pragma unroll + for (int k = 0; k < kPackBytesPerThread; ++k) { + if (k < byte_count) { + encoded[first_byte + k] = static_cast(packed >> (k * 8)); + } + } + } +} + + +extern "C" __global__ void dfloat11_thread_meta_kernel( + const int64_t* __restrict__ pref, + uint16_t* __restrict__ thread_meta, + int64_t n_elements, + int64_t total_bits, + int64_t region_bits, + int64_t n_regions) { + const int64_t t = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + if (t >= n_regions) { + return; + } + const int64_t key_lo = t * region_bits; + if (key_lo > total_bits) { + thread_meta[t] = 0; + return; + } + const int64_t key_hi = key_lo + region_bits; + int64_t lb_lo = lower_bound_i64(pref, n_elements + 1, key_lo); + if (lb_lo > n_elements) { + lb_lo = n_elements; + } + int64_t lb_hi = (key_hi > total_bits) + ? n_elements + : lower_bound_i64(pref, n_elements + 1, key_hi); + if (lb_hi > n_elements) { + lb_hi = n_elements; + } + const uint32_t count = static_cast(lb_hi - lb_lo); + const uint32_t gap = (count > 0) ? static_cast(pref[lb_lo] - key_lo) : 0u; + thread_meta[t] = static_cast((gap << kMetaCountBits) | count); +} + + +extern "C" __global__ void dfloat11_output_positions_kernel( + const int64_t* __restrict__ pref, + uint32_t* __restrict__ output_positions, + int64_t n_elements, + int64_t block_bits, + int64_t num_blocks) { + const int64_t blk = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + if (blk > num_blocks) { + return; + } + if (blk == num_blocks) { + output_positions[blk] = static_cast(n_elements); + return; + } + const int64_t key = blk * block_bits; + const int64_t lb = lower_bound_i64(pref, n_elements + 1, key); + output_positions[blk] = static_cast((lb <= n_elements) ? lb : n_elements); +} diff --git a/entropack/schemes/dfloat11/eager.py b/entropack/schemes/dfloat11/eager.py new file mode 100644 index 0000000..7002261 --- /dev/null +++ b/entropack/schemes/dfloat11/eager.py @@ -0,0 +1,239 @@ +from copy import copy + +import numpy as np +import torch + +from .format import ( + DFloat11Buffers, make_layout, max_block_elems, pack_thread_meta, + reconstruct_bf16_bits, +) + + +def _huffman_codec(): + from dahuffman import HuffmanCodec + + return HuffmanCodec + + +def exponent_counter(weight: torch.Tensor, threads_per_block: int | None = None) -> dict: + W = weight.reshape(-1).view(torch.int16) + exponent_8bits = ((W >> 7) & 0xFF).to(torch.int64) + counts = torch.bincount(exponent_8bits, minlength=256).cpu().tolist() + return {i: int(c) for i, c in enumerate(counts) if c > 0} + + +def get_32bit_codec(counter: dict): + HuffmanCodec = _huffman_codec() + codec = HuffmanCodec.from_frequencies(counter) + table = codec.get_code_table() + max_len = 0 + for _, (length, _) in table.items(): + max_len = max(max_len, length) + + compressed_codec = codec + compressed_counter = counter + + min_k = 2 + freq = np.array(list(counter.values())) + while max_len > 32: + min_indices = np.argpartition(freq, min_k)[:min_k] + min_k += 1 + min_keys = np.array(list(counter.keys()))[min_indices] + + compressed_counter = copy(counter) + for k in min_keys: + compressed_counter[k] = 1 + compressed_codec = HuffmanCodec.from_frequencies(compressed_counter) + table = compressed_codec.get_code_table() + max_len = 0 + for _, (length, _) in table.items(): + max_len = max(max_len, length) + + return compressed_codec, compressed_counter, table + + +def get_luts(table: dict) -> torch.Tensor: + prefixes = [""] + + for key, (bits, val) in table.items(): + if isinstance(key, int): + prefix = bin(val)[2:].rjust(bits, "0")[: ((bits - 1) // 8 * 8)] + if prefix not in prefixes: + prefixes.append(prefix) + + prefixes.sort(key=len) + + luts = np.zeros((len(prefixes), 256), dtype=np.uint8) + + for pi, p in enumerate(prefixes): + bytes_dict = {} + pl = len(p) // 8 + for key, (bits, val) in table.items(): + if isinstance(key, int): + bin_val = bin(val)[2:].rjust(bits, "0") + + if bin_val.startswith(p): + if (bits - 1) // 8 == pl: + dict_key = int(bin_val[(pl * 8) :].ljust(8, "0"), 2) + dict_value = key + else: + dict_key = int(bin_val[(pl * 8) : (pl * 8 + 8)], 2) + dict_value = 256 - prefixes.index(bin_val[: (pl * 8 + 8)]) + + if dict_key in bytes_dict and bytes_dict[dict_key] != dict_value: + raise ValueError(f"Key {dict_key} already exists in {bytes_dict}") + else: + bytes_dict[dict_key] = dict_value + + curr_val = 0 + for i in range(256): + if i in bytes_dict: + curr_val = bytes_dict[i] + luts[pi, i] = curr_val + + lens = np.zeros((1, 256), dtype=np.uint8) + for key, (bits, _val) in table.items(): + if isinstance(key, int): + lens[-1, key] = bits + + return torch.from_numpy(np.concatenate((luts, lens), axis=0)) + + +def _encode_bitstream(data, codec, bytes_per_thread: int, threads_per_block: int): + encoded = [] + + gaps = [] + counts = [] + output_positions = [] + + region_bits = 8 * bytes_per_thread + block_bits = region_bits * threads_per_block + + buffer = 0 + size = 0 + total_size = 0 + element_count = 0 + for s in data: + if total_size // region_bits + 1 > len(gaps): + gaps.append(total_size - total_size // region_bits * region_bits) + counts.append(0) + + if total_size // block_bits + 1 > len(output_positions): + output_positions.append(element_count) + + counts[-1] += 1 + + b, v = codec._table[s] + buffer = (buffer << b) + v + size += b + total_size += b + element_count += 1 + while size >= 8: + byte = buffer >> (size - 8) + encoded.append(byte) + buffer = buffer - (byte << (size - 8)) + size -= 8 + + if size > 0: + if total_size // region_bits + 1 > len(gaps): + gaps.append(0) + counts.append(0) + + if total_size // block_bits + 1 > len(output_positions): + output_positions.append(element_count) + + b, v = codec._table[codec._eof] + buffer = (buffer << b) + v + size += b + if size >= 8: + byte = buffer >> (size - 8) + else: + byte = buffer << (8 - size) + encoded.append(byte) + + output_positions.append(len(data)) + + blocks_per_grid = int(np.ceil(len(encoded) / (threads_per_block * bytes_per_thread))) + n_regions = threads_per_block * blocks_per_grid + gaps.extend([0] * (n_regions - len(gaps))) + counts.extend([0] * (n_regions - len(counts))) + + return ( + np.frombuffer(bytes(encoded), dtype=np.uint8).copy(), np.array(gaps, dtype=np.int64), np.array(counts, dtype=np.int64), + np.array(output_positions, dtype=np.uint32), + ) + + +def encode_weights(weights, codec, bytes_per_thread: int, threads_per_block: int): + W_combined = torch.cat(weights).view(torch.int16) + + exponent_8bits = ((W_combined >> 7) & 0xFF).to(torch.uint8) + other_8bits = ((W_combined >> 8) & 0x80 | (W_combined & 0x7F)).to(torch.uint8) + + encoded, gaps, counts, output_positions = _encode_bitstream( + exponent_8bits.tolist(), codec, bytes_per_thread, threads_per_block + ) + + return ( + torch.from_numpy(encoded), other_8bits, torch.from_numpy(output_positions), pack_thread_meta(gaps, counts), + make_layout(bytes_per_thread, threads_per_block, max_block_elems(output_positions)), + ) + + +def _decode_exponents(luts: np.ndarray, encoded: np.ndarray, n_elements: int) -> np.ndarray: + lut = luts.astype(np.int64) + lens_row = lut[-1] + num_levels = lut.shape[0] - 1 + ptr_min = 256 - (num_levels - 1) if num_levels > 1 else 256 + + bits = np.unpackbits(encoded.astype(np.uint8)) + n_bits = bits.size + + def read_byte(offset: int) -> int: + if offset + 8 <= n_bits: + seg = bits[offset : offset + 8] + else: + seg = np.zeros(8, np.uint8) + avail = n_bits - offset + if avail > 0: + seg[:avail] = bits[offset:] + return int(np.packbits(seg)[0]) + + out = np.empty(n_elements, np.int64) + cursor = 0 + for i in range(n_elements): + level = 0 + hop = 0 + while True: + entry = lut[level][read_byte(cursor + hop * 8)] + if num_levels > 1 and entry >= ptr_min: + level = 256 - entry + hop += 1 + else: + out[i] = entry + cursor += int(lens_row[entry]) + break + return out + + +def decode(buffers: DFloat11Buffers) -> torch.Tensor: + encoded_exponent, sign_mantissa, luts = buffers.encoded_exponent, buffers.sign_mantissa, buffers.luts + n_elements = sign_mantissa.numel() + exponents = _decode_exponents(luts.detach().cpu().numpy(), encoded_exponent.detach().cpu().numpy(), n_elements) + sm = sign_mantissa.detach().cpu().numpy().astype(np.uint8) + bf16_bits = reconstruct_bf16_bits(exponents, sm) + flat = torch.from_numpy(bf16_bits.view(np.int16)).view(torch.bfloat16) + return flat.to(sign_mantissa.device) + + +def encode( + *, weight: torch.Tensor, codec, luts: torch.Tensor, bytes_per_thread: int, threads_per_block: int, +) -> DFloat11Buffers: + flat = weight.reshape(-1).cpu() + encoded_exponent, sign_mantissa, output_positions, thread_meta, layout = encode_weights( + [flat], codec, bytes_per_thread, threads_per_block + ) + return DFloat11Buffers( + encoded_exponent=encoded_exponent, sign_mantissa=sign_mantissa, luts=luts, + output_positions=output_positions, thread_meta=thread_meta, layout=layout, + ) diff --git a/entropack/schemes/dfloat11/format.py b/entropack/schemes/dfloat11/format.py new file mode 100644 index 0000000..db472c3 --- /dev/null +++ b/entropack/schemes/dfloat11/format.py @@ -0,0 +1,146 @@ +import math +from dataclasses import dataclass +from typing import Annotated, NamedTuple + +import numpy as np +import torch + +from ..config import CompressionConfig, Range +from ..base import cached_parse + +BLOCK_SIZE = 256 +MAX_RESIDENT_BLOCKS_PER_SM = 32 + + +@dataclass +class DFloat11Config(CompressionConfig): + """Lossless BF16 compression settings. + + The coding-region settings are stored with the compressed tensor. Decoding reads them + from the representation rather than from the configuration.""" + + #: Bitstream bytes per coding region. Controls metadata overhead and decode parallelism. + bytes_per_thread: Annotated[int | None, Range(1, None)] = 16 + #: Threads per block in the encoded layout. + threads_per_block: Annotated[int | None, Range(1, None)] = 128 + + +META_COUNT_BITS = 11 +META_COUNT_MASK = (1 << META_COUNT_BITS) - 1 + + +class DFloat11Buffers(NamedTuple): + encoded_exponent: torch.Tensor + sign_mantissa: torch.Tensor + luts: torch.Tensor + output_positions: torch.Tensor + thread_meta: torch.Tensor + layout: torch.Tensor + + +PACKED_KEYS = DFloat11Buffers._fields + + +def reconstruct_bf16_bits(exponent: np.ndarray, sign_mantissa: np.ndarray) -> np.ndarray: + exp = exponent.astype(np.uint16) + sm = sign_mantissa.astype(np.uint16) + return (((sm & 0x80) << 8) | (exp << 7) | (sm & 0x7F)).astype(np.uint16) + + +def pack_thread_meta(gaps, counts) -> torch.Tensor: + g = np.asarray(gaps, dtype=np.int64) + c = np.asarray(counts, dtype=np.int64) + if g.size != c.size: + raise ValueError(f"gaps/counts length mismatch: {g.size} vs {c.size}") + if g.size and int(g.max()) > 31: + raise ValueError(f"gap {int(g.max())} exceeds 5 bits; max Huffman code length must be < 32") + if c.size and int(c.max()) > META_COUNT_MASK: + raise ValueError(f"per-thread symbol count {int(c.max())} exceeds {META_COUNT_BITS} bits; lower BYTES_PER_THREAD") + return torch.from_numpy(((g << META_COUNT_BITS) | c).astype(np.uint16)) + + +def make_layout(bytes_per_thread: int, threads_per_block: int, max_block_elems: int): + return torch.tensor([bytes_per_thread, threads_per_block, max_block_elems], dtype=torch.int32) + + +def parse_layout(layout) -> tuple[int, int, int]: + if layout.numel() != 3: + raise ValueError(f"dfloat11 layout must contain 3 int32 values, got {layout.numel()}") + bpt, tpb, max_block_elems = (int(v) for v in layout.detach().cpu().tolist()) + if bpt <= 0 or tpb <= 0 or max_block_elems < 0: + raise ValueError(f"invalid dfloat11 layout values: bpt={bpt}, tpb={tpb}, max_block_elems={max_block_elems}") + return bpt, tpb, max_block_elems + + +def parse_layout_cached(layout) -> tuple[int, int, int]: + return cached_parse(layout, parse_layout, "_dfloat11_layout") + + +def max_block_elems(output_positions) -> int: + op = np.asarray(output_positions, dtype=np.int64) + if op.size < 2: + return int(op[0]) if op.size else 0 + return int(np.diff(op).max()) + + +def validate_packed(buffers: dict, shape=None) -> None: + missing = [key for key in PACKED_KEYS if key not in buffers] + if missing: + raise ValueError(f"dfloat11 packed data is missing buffers: {missing}") + if not all(isinstance(buffers[key], torch.Tensor) for key in PACKED_KEYS): + raise TypeError("dfloat11 packed buffers must be torch.Tensor values") + + expected = { + "encoded_exponent": (torch.uint8, 1), "sign_mantissa": (torch.uint8, 1), + "luts": (torch.uint8, 2), "output_positions": (torch.uint32, 1), + "thread_meta": (torch.uint16, 1), "layout": (torch.int32, 1), + } + for key, (dtype, ndim) in expected.items(): + tensor = buffers[key] + if tensor.dtype != dtype or tensor.ndim != ndim: + raise ValueError( + f"dfloat11 buffer '{key}' must be {ndim}D {dtype}, got shape={tuple(tensor.shape)}, dtype={tensor.dtype}" + ) + + devices = {buffers[key].device for key in PACKED_KEYS} + if len(devices) != 1: + raise ValueError(f"dfloat11 packed buffers must share one device, got {devices}") + if buffers["luts"].shape[0] < 2 or buffers["luts"].shape[1] != 256: + raise ValueError(f"dfloat11 luts must have shape (num_levels + 1, 256), got {tuple(buffers['luts'].shape)}") + + bpt, tpb, max_elems = parse_layout_cached(buffers["layout"]) + n_elements = buffers["sign_mantissa"].numel() + if n_elements == 0: + raise ValueError("dfloat11 does not support empty tensors") + if shape is not None: + normalized_shape = tuple(shape) + if any(not isinstance(dim, int) or dim < 0 for dim in normalized_shape): + raise ValueError(f"invalid dfloat11 tensor shape: {normalized_shape}") + if math.prod(normalized_shape) != n_elements: + raise ValueError( + f"dfloat11 shape {normalized_shape} has {math.prod(normalized_shape)} " + f"elements but sign_mantissa has {n_elements}" + ) + + n_bytes = buffers["encoded_exponent"].numel() + blocks = (n_bytes + bpt * tpb - 1) // (bpt * tpb) + if buffers["thread_meta"].numel() != blocks * tpb: + raise ValueError("dfloat11 thread_meta length does not match layout/bitstream") + if buffers["output_positions"].numel() != blocks + 1: + raise ValueError("dfloat11 output_positions length does not match layout/bitstream") + + positions = buffers["output_positions"].detach().cpu().to(torch.int64) + if positions[0] != 0 or positions[-1] != n_elements: + raise ValueError("dfloat11 output_positions endpoints are invalid") + differences = positions[1:] - positions[:-1] + if (differences < 0).any() or (differences.max() if differences.numel() else 0) != max_elems: + raise ValueError("dfloat11 output_positions are not monotone or mismatch layout") + + metadata = buffers["thread_meta"].detach().cpu().to(torch.int64) + gaps = metadata >> META_COUNT_BITS + counts = metadata & META_COUNT_MASK + if (gaps > 31).any() or counts.sum() != n_elements: + raise ValueError("dfloat11 thread_meta fields are invalid") + lengths = buffers["luts"][-1].detach().cpu() + if int(lengths.max()) > 32: + raise ValueError("dfloat11 Huffman code length exceeds 32 bits") diff --git a/entropack/schemes/lattice_rans/__init__.py b/entropack/schemes/lattice_rans/__init__.py new file mode 100644 index 0000000..0d5228b --- /dev/null +++ b/entropack/schemes/lattice_rans/__init__.py @@ -0,0 +1,60 @@ +import torch + +from ..base import Scheme, packed_buffers, register_scheme +from ..checks import prepare_weight +from .format import ( + LATTICE_DIM, PACKED_KEYS, SUPPORTED_DTYPES, LatticeBuffers, LatticeRANSConfig, recommended_tile_elements, + validate_packed, +) + +__all__ = ["LatticeRANSScheme", "LatticeRANSConfig"] + + +class LatticeRANSScheme(Scheme): + name = "lattice_rans" + buffer_names = PACKED_KEYS + lossless = False + priority = 0 + dtypes = SUPPORTED_DTYPES + lanes = {"eager": "eager", "cuda": "cuda"} + + def encode(self, weight: torch.Tensor, config: LatticeRANSConfig) -> dict: + weight = prepare_weight(weight, scheme=self.name, ndim=2, dtypes=SUPPORTED_DTYPES, require_finite=True) + shape, device = tuple(weight.shape), weight.device + pad = -shape[1] % LATTICE_DIM + if pad: + weight = torch.cat([weight, weight.new_zeros(shape[0], pad)], dim=1) + lane, run_on = self.lane_for(weight, config.execution_backend) + if run_on is not None and device != run_on: + weight = weight.to(run_on) + tile_elements = ( + recommended_tile_elements(float(config.target_bpp)) if config.tile_elements is None + else int(config.tile_elements) + ) + packed = lane.encode( + weight=weight, target_bpp=float(config.target_bpp), + prob_bits=None if config.prob_bits in (None, 0) else int(config.prob_bits), + tile_elements=tile_elements, row_rdo_iterations=config.row_rdo_iterations, + row_rdo_candidates=config.row_rdo_candidates, scale_search_iterations=config.scale_search_iterations, + scale_search_max_vectors=config.scale_search_max_vectors, + ) + return {key: value.to(device) for key, value in packed._asdict().items()} + + def validate_buffers(self, buffers, shape, dtype): + validate_packed(buffers, tuple(shape), dtype) + + def decode(self, packed: dict, *, shape: tuple[int, ...], dtype: torch.dtype, + config: LatticeRANSConfig) -> torch.Tensor: + buffers = packed_buffers(packed, LatticeBuffers) + source = buffers.layout.device + lane, lane_device = self.lane_for(buffers.layout, config.execution_backend, gate_dtype=False) + if lane_device is not None and source != lane_device: + buffers = LatticeBuffers._make(value.to(lane_device) for value in buffers) + out = lane.decode(buffers, shape=shape, dtype=dtype, threads_per_block=config.threads_per_block, + l2_prefetch=config.l2_prefetch) + if shape and tuple(out.shape) != shape: + out = out[: shape[0], : shape[1]] + return out if out.device == source else out.to(source) + + +register_scheme(LatticeRANSScheme()) diff --git a/entropack/schemes/lattice_rans/cuda.py b/entropack/schemes/lattice_rans/cuda.py new file mode 100644 index 0000000..6d4c9ca --- /dev/null +++ b/entropack/schemes/lattice_rans/cuda.py @@ -0,0 +1,670 @@ +from dataclasses import dataclass, replace +from pathlib import Path + +import cupy +import numpy as np +import torch + +from ...backends.cuda import device as _device_caps +from ...backends.cuda.kernels import KernelLibrary +from ...backends.cuda.kernels import device_index as _device_index +from ...backends.cuda.kernels import ensure_dynamic_shared as _ensure_dynamic_shared +from ...backends.cuda.kernels import external_stream as _external_stream +from ...backends.cuda.kernels import pointer as _pointer +from ..tile_ans.format import NUM_STATES +from . import rans, rdo +from . import eager as _eager +from .format import ( + ALPHABET_COARSEN_MARGIN, BITS_PER_BYTE, BLOCK_SIZE, + LATTICE_DIM, META_ALPHABET, META_FREQ_OFFSET, META_N_SYMBOLS, META_SYM_MIN, MIN_SHARED_BLOCKS_PER_SM, + MODEL_LAYOUT_BYTES, NUM_COORD_STREAMS, NUM_STREAMS_FULL, SHARED_STAGING_HEADROOM, STATIC_SHARED, STREAM_META_WIDTH, + LatticeBuffers, check_row_scales_finite, decode_geometry, make_layout, report_alphabet_clamp, resolve_prob_bits, + snap_to_container, subsample_index, vector_tile_elements, +) + +_CUDA_PATH = Path(__file__).parent / "lattice_rans.cu" +_TILE_INCLUDE_PATH = Path(__file__).parent.parent / "tile_ans" +_KERNEL_NAMES = ( + "e8_quantize_fields_kernel", + "e8_quantize_fields_f32_kernel", + "e8_refit_scales_kernel", + "e8_refit_scales_f32_kernel", + "e8_minmax_fields_kernel", + "e8_histogram_kernel", + "e8_rans_encode_vector_kernel", + "e8_compact_kernel", + "e8_decode_vector_shlut8pf_kernel", + "e8_decode_vector_shlut8pf_g_kernel", + "e8_decode_vector_packed32pf_kernel", + "e8_decode_vector_packed32pf_g_kernel", + "e8_decode_vector_packed32_kernel", + "e8_decode_vector_packed32_g_kernel", + "e8_decode_vector_fused_kernel", + "e8_decode_vector_fused_g_kernel", +) +_LIBRARY = KernelLibrary( + key="lattice_rans", source=_CUDA_PATH, + defines=lambda _device, probability_bits: (f"TILE_ANS_PROB_BITS={probability_bits}",), includes=(_TILE_INCLUDE_PATH,), + kernel_names=_KERNEL_NAMES, +) +_kernel = _LIBRARY.kernel + +_MIN_RMS = _eager.MIN_RMS +_INT_MAX = 0x7FFFFFFF +_INT_MIN = -0x7FFFFFFF - 1 +_MM_CMAX = 2 * (NUM_STREAMS_FULL - 1) +_MM_LEN = _MM_CMAX + 1 +_MINMAX_TEMPLATE = np.array( + [(_INT_MAX if (t & 1) == 0 and t != _MM_CMAX else _INT_MIN) for t in range(_MM_LEN)], dtype=np.int32, +) + +_ENCODE_KERNELS = { + torch.bfloat16: ("e8_quantize_fields_kernel", "e8_refit_scales_kernel"), + torch.float32: ("e8_quantize_fields_f32_kernel", "e8_refit_scales_f32_kernel"), +} +_STORE_KINDS = { + torch.float32: 0, torch.float16: 1, torch.float8_e4m3fn: 2, torch.float8_e5m2: 3, torch.int8: 4, torch.int16: 5, + torch.int32: 6, torch.int64: 7, torch.uint8: 8, torch.uint16: 9, torch.uint32: 10, torch.uint64: 11, torch.bool: 12, +} +_SCRATCH_DTYPES = frozenset({torch.float8_e4m3fnuz, torch.float8_e5m2fnuz}) +_GENERIC_DECODE_KERNEL = { + "e8_decode_vector_shlut8pf_kernel": "e8_decode_vector_shlut8pf_g_kernel", + "e8_decode_vector_packed32pf_kernel": "e8_decode_vector_packed32pf_g_kernel", + "e8_decode_vector_packed32_kernel": "e8_decode_vector_packed32_g_kernel", + "e8_decode_vector_fused_kernel": "e8_decode_vector_fused_g_kernel", +} + + +def _grid(caps, total: int, threads: int = BLOCK_SIZE) -> int: + return caps.grid(-(-total // threads), threads) + + +@dataclass(frozen=True) +class _QuantizeRequest: + weight: torch.Tensor + rms: torch.Tensor + scale: float + prob_bits: int + rows: int + cols: int + device: torch.device + device_index: int + + @classmethod + def of(cls, weight: torch.Tensor, rms: torch.Tensor, scale: float, prob_bits: int): + rows, cols = weight.shape + return cls( + weight=weight, rms=rms, scale=scale, prob_bits=prob_bits, rows=rows, cols=cols, device=weight.device, + device_index=_device_index(weight), + ) + + +@dataclass(frozen=True) +class _QuantizeResult: + counts: list + sizes: list + sym_min: list + alphabets: list + fields: torch.Tensor + c_arr: torch.Tensor + minmax: torch.Tensor + scales: torch.Tensor | None = None + row_sse: torch.Tensor | None = None + + +def _quantize_pass(request: _QuantizeRequest) -> _QuantizeResult: + weight, rms = request.weight, request.rms + rows, cols = request.rows, request.cols + device, device_index, prob_bits = request.device, request.device_index, request.prob_bits + V = rows * (cols // LATTICE_DIM) + fields = torch.empty(V * LATTICE_DIM, dtype=torch.int32, device=device) + c_arr = torch.empty(V, dtype=torch.int32, device=device) + minmax = torch.from_numpy(_MINMAX_TEMPLATE.copy()).to(device) + torch_stream = torch.cuda.current_stream(device) + with cupy.cuda.Device(device_index), _external_stream(torch_stream): + _kernel(device_index, prob_bits, _ENCODE_KERNELS[weight.dtype][0])( + (_grid(_device_caps.caps(device_index), V),), + (BLOCK_SIZE,), + ( + _pointer(weight), _pointer(rms), np.float32(request.scale), np.int32(rows), np.int32(cols), + np.int32(cols // LATTICE_DIM), _pointer(fields), _pointer(c_arr), _pointer(minmax), + ), + ) + minmax_np = minmax.cpu().numpy() + counts, sizes, sym_min, alphabets = _counts_from_minmax( + fields, c_arr, minmax, minmax_np, V, device, device_index, prob_bits + ) + return _QuantizeResult( + counts=counts, sizes=sizes, sym_min=sym_min, alphabets=alphabets, fields=fields, c_arr=c_arr, minmax=minmax, + ) + + +def _refit_row_scales(request: _QuantizeRequest, result: _QuantizeResult): + weight, rms = request.weight, request.rms + rows, cols = request.rows, request.cols + device_index, prob_bits = request.device_index, request.prob_bits + scales = torch.empty(rows, dtype=torch.float32, device=weight.device) + row_sse = torch.empty(rows, dtype=torch.float32, device=weight.device) + torch_stream = torch.cuda.current_stream(weight.device) + with cupy.cuda.Device(device_index), _external_stream(torch_stream): + _kernel(device_index, prob_bits, _ENCODE_KERNELS[weight.dtype][1])( + (rows,), + (BLOCK_SIZE,), + ( + _pointer(weight), _pointer(result.fields), _pointer(result.c_arr), _pointer(rms), np.float32(request.scale), + _pointer(scales), _pointer(row_sse), np.int32(rows), np.int32(cols), np.int32(cols // LATTICE_DIM), + ), + ) + return scales, row_sse + + +def _summarize_fields(request: _QuantizeRequest, fields, c_arr) -> _QuantizeResult: + rows, cols = request.rows, request.cols + device, device_index, prob_bits = request.device, request.device_index, request.prob_bits + V = rows * (cols // LATTICE_DIM) + minmax = torch.from_numpy(_MINMAX_TEMPLATE.copy()).to(device) + torch_stream = torch.cuda.current_stream(device) + with cupy.cuda.Device(device_index), _external_stream(torch_stream): + _kernel(device_index, prob_bits, "e8_minmax_fields_kernel")( + (_grid(_device_caps.caps(device_index), V),), (BLOCK_SIZE,), + (_pointer(fields), _pointer(c_arr), np.int64(V), _pointer(minmax)), + ) + minmax_np = minmax.cpu().numpy() + counts, sizes, sym_min, alphabets = _counts_from_minmax( + fields, c_arr, minmax, minmax_np, V, device, device_index, prob_bits + ) + return _QuantizeResult( + counts=counts, sizes=sizes, sym_min=sym_min, alphabets=alphabets, fields=fields, c_arr=c_arr, minmax=minmax, + ) + + +def _counts_from_minmax(fields, c_arr, minmax, minmax_np, V, device, device_index, prob_bits): + alphabets = [int(minmax_np[_MM_CMAX]) + 1] + sym_min = [0] + for stream in range(NUM_COORD_STREAMS): + lo, hi = int(minmax_np[stream * 2]), int(minmax_np[stream * 2 + 1]) + live = lo != _INT_MAX + alphabets.append(hi - lo + 1 if live else 0) + sym_min.append(lo if live else 0) + + bin_off = np.zeros(NUM_STREAMS_FULL, dtype=np.int32) + cursor = 2 + for stream in range(1, NUM_STREAMS_FULL): + bin_off[stream] = cursor + cursor += alphabets[stream] + total_bins = cursor + + bins = torch.zeros(total_bins, dtype=torch.int32, device=device) + bin_off_gpu = torch.from_numpy(bin_off).to(device) + caps = _device_caps.caps(device_index) + staging_bins = (caps.shared_per_block - SHARED_STAGING_HEADROOM) // 4 + shared_bins = total_bins if total_bins <= staging_bins else 0 + torch_stream = torch.cuda.current_stream(device) + with cupy.cuda.Device(device_index), _external_stream(torch_stream): + _kernel(device_index, prob_bits, "e8_histogram_kernel")( + (_grid(caps, V),), + (BLOCK_SIZE,), + ( + _pointer(fields), _pointer(c_arr), _pointer(minmax), _pointer(bin_off_gpu), _pointer(bins), np.int64(V), + np.int32(shared_bins), + ), + shared_mem=shared_bins * 4, + ) + bins_np = bins.cpu().numpy().astype(np.int64) + n0, n1 = int(bins_np[0]), int(bins_np[1]) + sizes = [V] + [n0 if stream % 2 == 0 else n1 for stream in range(NUM_COORD_STREAMS)] + counts = [bins_np[: alphabets[0]].copy()] + for stream in range(1, NUM_STREAMS_FULL): + offset = int(bin_off[stream]) + counts.append(bins_np[offset : offset + alphabets[stream]].copy()) + return counts, sizes, sym_min, alphabets + + +def _analytic_total_bytes(counts, sizes, rows, prob_bits, tile_elements): + return rans.coded_bytes(counts, sizes, prob_bits, tile_elements) + rows * 4 + MODEL_LAYOUT_BYTES + + +def _optimize_rows( + request: _QuantizeRequest, baseline: _QuantizeResult, iterations: int, candidate_count: int, tile_elements: int, +) -> _QuantizeResult: + rows, device, prob_bits = request.rows, request.device, request.prob_bits + ratios = rdo.ratio_ladder(candidate_count) + baseline_index = ratios.index(1.0) + candidates = [] + for index, ratio in enumerate(ratios): + if index == baseline_index: + candidate = baseline + else: + scaled = replace(request, scale=request.scale * ratio) + summary = _quantize_pass(scaled) + scales, row_sse = _refit_row_scales(scaled, summary) + candidate = replace(summary, scales=scales, row_sse=row_sse) + candidates.append(rdo.compact_candidate(candidate)) + + summary, fitted_scales = rdo.optimize_rows( + candidates, rows=rows, cols=request.cols, baseline_index=baseline_index, iterations=iterations, device=device, + summarize=lambda fields, c_arr: _summarize_fields(request, fields, c_arr), + total_bytes=lambda counts, sizes: _analytic_total_bytes(counts, sizes, rows, prob_bits, tile_elements), + ) + return replace(summary, scales=fitted_scales) + + +def _bisect_scale_cuda(work, rms, prob_bits, target_bpp, tile_elements, n_iter=34, max_vectors=262144): + rows, cols = work.shape + device = work.device + vecs_per_row = cols // LATTICE_DIM + sub_rows = max(1, min(rows, max_vectors // vecs_per_row)) + index = subsample_index(rows, cols, sub_rows, device) + w_sub = work if index is None else work.index_select(0, index) + rms_sub = rms if index is None else rms.index_select(0, index) + N_sub = sub_rows * cols + + lo, hi = 0.001, 8.0 + for _ in range(n_iter): + mid = (lo * hi) ** 0.5 + summary = _quantize_pass(_QuantizeRequest.of(w_sub, rms_sub, mid, prob_bits)) + total = _analytic_total_bytes(summary.counts, summary.sizes, sub_rows, prob_bits, tile_elements) + if BITS_PER_BYTE * total / N_sub > target_bpp: + lo = mid + else: + hi = mid + return (lo * hi) ** 0.5 + + +@dataclass(frozen=True) +class _EncodeOptions: + target_bpp: float + prob_bits: int + auto_prob_bits: bool + tile_elements: int + row_rdo_iterations: int + row_rdo_candidates: int + scale_search_iterations: int + scale_search_max_vectors: int + table_size: int + + +def _resolve_options( + target_bpp, prob_bits, tile_elements, row_rdo_iterations, row_rdo_candidates, scale_search_iterations, + scale_search_max_vectors, +) -> _EncodeOptions: + resolved_prob_bits, auto_prob_bits = resolve_prob_bits(prob_bits, target_bpp) + return _EncodeOptions( + target_bpp=float(target_bpp), + prob_bits=resolved_prob_bits, + auto_prob_bits=auto_prob_bits, + tile_elements=int(tile_elements), + row_rdo_iterations=int(row_rdo_iterations), + row_rdo_candidates=int(row_rdo_candidates), + scale_search_iterations=int(scale_search_iterations), + scale_search_max_vectors=int(scale_search_max_vectors), + table_size=1 << resolved_prob_bits, + ) + + +def _resolve_scale(work, rms, options: _EncodeOptions): + """The table never shrinks below its starting precision: a coarser grid normalizes the distributions less exactly, so + shrinking costs rate, while the decode it would buy is not certain -- the decode is priced per tile and per symbol, and + the table size only decides which representation still stages in shared memory. + """ + prob_bits = options.prob_bits + if float(work.abs().amax().item()) == 0.0: + request = _QuantizeRequest.of(work, rms, 1.0, prob_bits) + return options, request, _quantize_pass(request) + clamped = False + while True: + escalate = False + scale = _bisect_scale_cuda( + work, rms, prob_bits, options.target_bpp, options.tile_elements, n_iter=options.scale_search_iterations, + max_vectors=options.scale_search_max_vectors, + ) + while True: + request = _QuantizeRequest.of(work, rms, scale, prob_bits) + summary = _quantize_pass(request) + alpha = max(summary.alphabets) if summary.alphabets else 0 + if alpha <= options.table_size: + break + if options.auto_prob_bits and prob_bits < 15: + prob_bits += 1 + options = replace(options, prob_bits=prob_bits, table_size=1 << prob_bits) + escalate = True + break + clamped = True + scale *= alpha / options.table_size * ALPHABET_COARSEN_MARGIN + if not escalate: + if clamped: + report_alphabet_clamp(scale, alpha, options.table_size) + return options, request, summary + + +def _build_codec_tables(summary: _QuantizeResult, options: _EncodeOptions, device): + """Only alphabet-sized frequencies are stored, never a full ``table_size`` LUT; the decoder rebuilds its slot->symbol + table from them, so a per-layer table stays small. + """ + n_streams = NUM_STREAMS_FULL + table_size = options.table_size + freq_parts = [] + cdf_parts = [] + freq_cursor = 0 + meta = np.zeros((n_streams, STREAM_META_WIDTH), dtype=np.int64) + for table in range(n_streams): + n_s = int(summary.sizes[table]) + alphabet = int(summary.alphabets[table]) if n_s else 0 + if alphabet > table_size: + raise ValueError( + f"E8 coordinate alphabet {alphabet} exceeds rANS table_size {table_size}; " + "raise prob_bits or coarsen the lattice scale" + ) + if alphabet: + frequency = rans.normalize_freq(summary.counts[table], table_size).astype(np.uint16) + cdf = np.zeros(alphabet, dtype=np.uint16) + if alphabet > 1: + cdf[1:] = np.cumsum(frequency.astype(np.int64))[:-1].astype(np.uint16) + freq_parts.append(frequency) + cdf_parts.append(cdf) + meta[table, META_N_SYMBOLS] = n_s + meta[table, META_SYM_MIN] = summary.sym_min[table] + meta[table, META_FREQ_OFFSET] = freq_cursor + meta[table, META_ALPHABET] = alphabet + freq_cursor += alphabet + + freq_tables_np = np.concatenate(freq_parts) if freq_parts else np.empty(0, dtype=np.uint16) + cdfs_np = np.concatenate(cdf_parts) if cdf_parts else np.empty(0, dtype=np.uint16) + return ( + torch.from_numpy(freq_tables_np).to(device), torch.from_numpy(cdfs_np).to(device), + torch.from_numpy(meta.astype(np.int32)).to(device), + ) + + +def _encode_streams( + summary: _QuantizeResult, freq_tables, cdfs_gpu, stream_meta, total_vectors: int, options: _EncodeOptions, device, + device_index: int, +): + tile_vectors = vector_tile_elements(options.tile_elements) + total_tiles = max(1, (total_vectors + tile_vectors - 1) // tile_vectors) + states = torch.empty((total_tiles, NUM_STATES), dtype=torch.uint32, device=device) + word_counts = torch.empty(total_tiles, dtype=torch.uint32, device=device) + scratch_stride = tile_vectors * 9 + scratch = torch.empty(total_tiles * scratch_stride, dtype=torch.uint16, device=device) + warps_per_block = BLOCK_SIZE // _device_caps.caps(device_index).warp_size + blocks_x = max(1, -(-total_tiles // warps_per_block)) + torch_stream = torch.cuda.current_stream(device) + with cupy.cuda.Device(device_index), _external_stream(torch_stream): + _kernel(device_index, options.prob_bits, "e8_rans_encode_vector_kernel")( + (blocks_x,), + (BLOCK_SIZE,), + ( + _pointer(summary.fields), _pointer(summary.c_arr), _pointer(freq_tables), _pointer(cdfs_gpu), + _pointer(stream_meta), _pointer(scratch), _pointer(word_counts), _pointer(states), np.int64(total_vectors), + np.int32(tile_vectors), np.int32(total_tiles), + ), + ) + counts_cp = cupy.from_dlpack(word_counts) + offsets64 = torch.empty(total_tiles + 1, dtype=torch.int64, device=device) + offsets_cp = cupy.from_dlpack(offsets64) + offsets_cp[0] = 0 + cupy.cumsum(counts_cp, dtype=cupy.int64, out=offsets_cp[1:]) + total_words = int(offsets64[-1].item()) + if total_words >= 1 << 32: + raise ValueError("lattice_rans payload exceeds uint32 offset capacity") + offsets = offsets64.to(torch.uint32) + payload = torch.empty(total_words, dtype=torch.uint16, device=device) + with cupy.cuda.Device(device_index), _external_stream(torch_stream): + _kernel(device_index, options.prob_bits, "e8_compact_kernel")( + (total_tiles,), (BLOCK_SIZE,), + (_pointer(scratch), _pointer(offsets), _pointer(payload), np.int32(scratch_stride), np.int32(total_tiles)), + ) + return payload, offsets, states + + +def encode( + weight, *, target_bpp, prob_bits, tile_elements, row_rdo_iterations, row_rdo_candidates, scale_search_iterations, + scale_search_max_vectors, +): + options = _resolve_options(target_bpp, prob_bits, tile_elements, row_rdo_iterations, + row_rdo_candidates, scale_search_iterations, scale_search_max_vectors) + + weight = weight.contiguous() + device = weight.device + device_index = _device_index(weight) + rows, cols = weight.shape + V = rows * (cols // LATTICE_DIM) + + if weight.dtype == torch.bfloat16: + work = weight + xf = weight.float() + else: + work = weight.float().contiguous() + xf = work + rms = xf.square().mean(dim=1, keepdim=True).sqrt().clamp_min(_MIN_RMS) + check_row_scales_finite(rms, weight.dtype) + # For a bf16 container this fp32 view only served the row RMS above. Releasing it before the scale search keeps two + # full-tensor copies from being alive at once. + del xf + + options, request, summary = _resolve_scale(work, rms, options) + scales, row_sse = _refit_row_scales(request, summary) + summary = replace(summary, scales=scales, row_sse=row_sse) + if options.row_rdo_iterations > 0 and float(work.abs().amax().item()) != 0.0: + summary = _optimize_rows(request, summary, options.row_rdo_iterations, + options.row_rdo_candidates, options.tile_elements) + + freq_tables, cdfs_gpu, stream_meta = _build_codec_tables(summary, options, device) + payload, offsets, states = _encode_streams( + summary, freq_tables, cdfs_gpu, stream_meta, V, options, device, device_index) + fitted_scales = summary.scales + del summary + + layout = make_layout(cols, options.prob_bits, options.tile_elements).to(device) + return LatticeBuffers( + payload=payload, offsets=offsets, states=states, stream_meta=stream_meta, freq_tables=freq_tables, + scales=fitted_scales.contiguous(), layout=layout, + ) + + +def _shared_lut_usable(shared_bytes: int | None, caps, threads: int) -> bool: + """A table that consumes most of an SM's shared memory leaves one CTA resident, and this decode hides renormalization + latency with warp count, so it would run slower reading the table from shared memory than from global. The staged table + has to leave room for ``MIN_SHARED_BLOCKS_PER_SM`` resident blocks. + """ + if shared_bytes is None: + return False + staged = shared_bytes + STATIC_SHARED + if staged > caps.shared_limit(STATIC_SHARED): + return False + return caps.blocks_per_sm(threads, staged) >= MIN_SHARED_BLOCKS_PER_SM + + +def _stream_slices(meta_np) -> list[tuple[int, int, int]]: + return [ + (stream, int(row[META_FREQ_OFFSET]), int(row[META_ALPHABET])) for stream, row in enumerate(meta_np) + if int(row[META_ALPHABET]) + ] + + +def _cdf_and_frequency(freq_np, offset: int, alphabet: int): + frequency = freq_np[offset : offset + alphabet].astype(np.int64) + cdf = np.zeros(alphabet, dtype=np.int64) + cdf[1:] = np.cumsum(frequency)[:-1] + return cdf, frequency + + +def _pack_bits(freq_np, slices, table_size: int, n_streams: int) -> np.ndarray | None: + pack_bits = np.zeros(n_streams, dtype=np.int32) + for stream, offset, alphabet in slices: + _cdf, frequency = _cdf_and_frequency(freq_np, offset, alphabet) + stored = np.where(frequency == table_size, 0, frequency) + symbol_bits = max(1, (alphabet - 1).bit_length()) + frequency_bits = max(1, int(stored.max()).bit_length()) + delta_bits = max(1, (int(frequency.max()) - 1).bit_length()) + if symbol_bits + frequency_bits + delta_bits > 32: + return None + pack_bits[stream] = symbol_bits | (frequency_bits << 8) + return pack_bits + + +def _slot_fields(freq_np, offset: int, alphabet: int, table_size: int): + cdf, frequency = _cdf_and_frequency(freq_np, offset, alphabet) + stored = np.where(frequency == table_size, 0, frequency) + return (np.repeat(np.arange(alphabet), frequency), np.repeat(stored, frequency), np.repeat(cdf, frequency)) + + +def _shared_tables(freq_np, slices, table_size: int, n_streams: int): + symbols = np.zeros(n_streams * table_size, dtype=np.uint8) + begin_frequency = np.zeros(int(freq_np.size), dtype=np.uint32) + for stream, offset, alphabet in slices: + cdf, frequency = _cdf_and_frequency(freq_np, offset, alphabet) + begin_frequency[offset : offset + alphabet] = (cdf | (frequency << 16)).astype(np.uint32) + base = stream * table_size + symbols[base : base + table_size] = np.repeat(np.arange(alphabet, dtype=np.uint8), frequency) + return symbols, begin_frequency + + +def _packed_tables(freq_np, slices, table_size: int, n_streams: int, pack_bits): + packed = np.zeros((n_streams, table_size), dtype=np.uint32) + for stream, offset, alphabet in slices: + symbol, frequency, begin = _slot_fields(freq_np, offset, alphabet, table_size) + symbol_bits = int(pack_bits[stream] & 0xFF) + frequency_bits = int((pack_bits[stream] >> 8) & 0xFF) + packed[stream] = ( + symbol.astype(np.uint32) | (frequency.astype(np.uint32) << np.uint32(symbol_bits)) + | ((np.arange(table_size, dtype=np.int64) - begin).astype(np.uint32) << np.uint32(symbol_bits + frequency_bits)) + ) + return packed + + +def _wide_tables(freq_np, slices, table_size: int, n_streams: int): + luts = np.zeros((n_streams, table_size), dtype=np.uint64) + for stream, offset, alphabet in slices: + symbol, frequency, begin = _slot_fields(freq_np, offset, alphabet, table_size) + luts[stream] = ( + (begin.astype(np.uint64) << np.uint64(32)) | (frequency.astype(np.uint64) << np.uint64(16)) + | symbol.astype(np.uint64) + ) + return luts + + +@dataclass(frozen=True) +class _DecodePlan: + representation: str + coset_frequency0: int + error: torch.Tensor + vector_tiles: int + tile_vectors: int + sym_u8: torch.Tensor | None = None + fb_lut: torch.Tensor | None = None + shared_bytes: int = 0 + packed_luts: torch.Tensor | None = None + pack_bits: torch.Tensor | None = None + decode_luts: torch.Tensor | None = None + + +def _decode_plan(layout, stream_meta, freq_tables, info, caps, threads: int) -> _DecodePlan: + device = freq_tables.device + fingerprint = (stream_meta.data_ptr(), freq_tables.data_ptr(), device, threads) + cached = getattr(layout, "_lattice_rans_decode_plan", None) + if cached is not None and cached[0] == fingerprint: + return cached[1] + + meta_np = stream_meta.detach().cpu().numpy().astype(np.int64) + freq_np = freq_tables.detach().cpu().numpy().astype(np.uint16) + prob_bits = info["prob_bits"] + table_size = 1 << prob_bits + n_streams = info["n_streams"] + slices = _stream_slices(meta_np) + alphabets = meta_np[:, META_ALPHABET] + max_alpha = int(alphabets.max()) if alphabets.size else 0 + alphabet_sum = int(alphabets.sum()) + + shared_bytes = n_streams * table_size + alphabet_sum * 4 if max_alpha < 256 else None + fields: dict = {} + if _shared_lut_usable(shared_bytes, caps, threads): + representation = "shared" + symbols, begin_frequency = _shared_tables(freq_np, slices, table_size, n_streams) + fields = { + "sym_u8": torch.from_numpy(symbols).to(device), "fb_lut": torch.from_numpy(begin_frequency).to(device), + "shared_bytes": int(symbols.nbytes + begin_frequency.nbytes), + } + else: + pack_bits = _pack_bits(freq_np, slices, table_size, n_streams) + if pack_bits is not None: + representation = "packed" + packed = _packed_tables(freq_np, slices, table_size, n_streams, pack_bits) + fields = {"packed_luts": torch.from_numpy(packed).to(device), "pack_bits": torch.from_numpy(pack_bits).to(device)} + else: + representation = "wide" + luts = _wide_tables(freq_np, slices, table_size, n_streams) + fields = {"decode_luts": torch.from_numpy(luts).to(device)} + + vectors = info["rows"] * (info["cols"] // LATTICE_DIM) + tile_vectors = vector_tile_elements(info["tile_elements"]) + plan = _DecodePlan( + representation=representation, coset_frequency0=int(freq_np[0]), error=torch.zeros(1, dtype=torch.int32, device=device), + vector_tiles=max(1, (vectors + tile_vectors - 1) // tile_vectors), tile_vectors=tile_vectors, **fields, + ) + layout._lattice_rans_decode_plan = (fingerprint, plan) + return plan + + +def _decode_launch(plan, buffers: LatticeBuffers, output, rows, cols, l2_prefetch): + vecs_per_row = cols // LATTICE_DIM + total_vectors = rows * vecs_per_row + native_bf16 = output.dtype == torch.bfloat16 + store_kind = 0 if native_bf16 else _STORE_KINDS[output.dtype] + if plan.representation == "shared": + kernel_name = "e8_decode_vector_shlut8pf_kernel" + shared = plan.shared_bytes + args = ( + _pointer(buffers.payload), _pointer(buffers.offsets), _pointer(buffers.states), _pointer(plan.sym_u8), + _pointer(plan.fb_lut), np.uint32(plan.coset_frequency0), _pointer(buffers.stream_meta), + _pointer(buffers.scales), _pointer(output), np.int32(vecs_per_row), np.int64(total_vectors), + np.int32(plan.tile_vectors), np.int32(plan.vector_tiles), np.int32(plan.fb_lut.numel()), _pointer(plan.error), + ) + elif plan.representation == "packed": + kernel_name = "e8_decode_vector_packed32pf_kernel" if l2_prefetch else "e8_decode_vector_packed32_kernel" + shared = 0 + args = ( + _pointer(buffers.payload), _pointer(buffers.offsets), _pointer(buffers.states), _pointer(plan.packed_luts), + _pointer(plan.pack_bits), np.uint32(plan.coset_frequency0), _pointer(buffers.stream_meta), + _pointer(buffers.scales), _pointer(output), np.int32(vecs_per_row), np.int64(total_vectors), + np.int32(plan.tile_vectors), np.int32(plan.vector_tiles), _pointer(plan.error), + ) + else: + kernel_name = "e8_decode_vector_fused_kernel" + shared = 0 + args = ( + _pointer(buffers.payload), _pointer(buffers.offsets), _pointer(buffers.states), _pointer(plan.decode_luts), + _pointer(buffers.stream_meta), _pointer(buffers.scales), _pointer(output), np.int32(vecs_per_row), + np.int64(total_vectors), np.int32(plan.tile_vectors), np.int32(plan.vector_tiles), _pointer(plan.error), + ) + if not native_bf16: + kernel_name = _GENERIC_DECODE_KERNEL[kernel_name] + args += (np.int32(store_kind),) + return kernel_name, shared, args + + +def decode(buffers: LatticeBuffers, *, shape, dtype, threads_per_block, l2_prefetch): + layout = buffers.layout + info = decode_geometry(buffers._asdict(), shape) + rows, cols = info["rows"], info["cols"] + prob_bits = info["prob_bits"] + + payload = buffers.payload + device_index = _device_index(payload) + caps = _device_caps.caps(device_index) + threads = _device_caps.resolve_threads(caps, threads_per_block, BLOCK_SIZE) + plan = _decode_plan(layout, buffers.stream_meta, buffers.freq_tables, info, caps, threads) + scratch = dtype in _SCRATCH_DTYPES + out_dtype = torch.float32 if scratch else dtype + output = torch.empty((rows, cols), dtype=out_dtype, device=payload.device) + torch_stream = torch.cuda.current_stream(payload.device) + + kernel_name, shared, args = _decode_launch(plan, buffers, output, rows, cols, l2_prefetch) + launch_blocks = max(1, -(-plan.vector_tiles // (threads // caps.warp_size))) + with cupy.cuda.Device(device_index), _external_stream(torch_stream): + kernel = _kernel(device_index, prob_bits, kernel_name) + _ensure_dynamic_shared(kernel, shared) + kernel((launch_blocks,), (threads,), args, shared_mem=shared) + if scratch: + output = snap_to_container(output, dtype) + return output diff --git a/entropack/schemes/lattice_rans/eager.py b/entropack/schemes/lattice_rans/eager.py new file mode 100644 index 0000000..f1cdaf8 --- /dev/null +++ b/entropack/schemes/lattice_rans/eager.py @@ -0,0 +1,405 @@ +import math +from dataclasses import dataclass + +import numpy as np +import torch + +from . import rans, rdo +from .format import ( + ALPHABET_COARSEN_MARGIN, BITS_PER_BYTE, LATTICE_DIM, + META_ALPHABET, META_FREQ_OFFSET, META_N_SYMBOLS, META_SYM_MIN, MODEL_LAYOUT_BYTES, MODEL_STREAM_META_BYTES, + NUM_COORD_FIELDS, STREAM_META_WIDTH, LatticeBuffers, check_row_scales_finite, decode_geometry, make_layout, + report_alphabet_clamp, resolve_prob_bits, snap_to_container, subsample_index, vector_tile_elements, +) + +MIN_RMS = 1.0e-12 + + +def _nearest_dn(Y: torch.Tensor) -> torch.Tensor: + f = torch.round(Y) + resid = Y - f + par = torch.remainder(f.sum(1), 2.0) + idx = resid.abs().argmax(1, keepdim=True) + sgn = torch.sign(resid.gather(1, idx)) + sgn = torch.where(sgn == 0, torch.ones_like(sgn), sgn) + onehot = torch.zeros_like(f) + onehot.scatter_(1, idx, 1.0) + return f + onehot * sgn * (par != 0).float().unsqueeze(1) + + +def nearest_e8(Y: torch.Tensor) -> torch.Tensor: + """The coset comparison runs in float64 so exact ties, common on bf16 and half-integer grids, resolve as the CUDA kernel + does rather than following fp32 summation order. + """ + c0 = _nearest_dn(Y) + c1 = _nearest_dn(Y - 0.5) + 0.5 + yd = Y.to(torch.float64) + d0 = (yd - c0.to(torch.float64)).square().sum(1, keepdim=True) + d1 = (yd - c1.to(torch.float64)).square().sum(1, keepdim=True) + # ``+ 0.0`` folds -0.0 to +0.0: torch.round of a small negative coordinate is -0.0, and the CUDA kernel rounds in integer + # coordinates so it can never emit one. + return torch.where(d0 <= d1, c0, c1) + 0.0 + + +def point_to_fields(p: torch.Tensor): + doubled = torch.round(2 * p) + c = (doubled.remainder(2.0) != 0).any(dim=1).to(torch.int64) + z = torch.round(p - 0.5 * c.unsqueeze(1)).to(torch.int64) + z0_6 = z[:, :7] + par = z0_6.sum(1).remainder(2) + m = torch.div(z[:, 7] - par, 2, rounding_mode="floor") + return c, z0_6, m + + +def fields_to_point(c: torch.Tensor, z0_6: torch.Tensor, m: torch.Tensor) -> torch.Tensor: + par = z0_6.sum(1).remainder(2) + z7 = 2 * m + par + z = torch.cat([z0_6, z7.unsqueeze(1)], dim=1) + return z.to(torch.float32) + 0.5 * c.to(torch.float32).unsqueeze(1) + + +def _coord_streams(c_np, z0_6_np, m_np): + idx = [np.flatnonzero(c_np == 0), np.flatnonzero(c_np == 1)] + fields = [z0_6_np[:, f] for f in range(NUM_COORD_FIELDS - 1)] + [m_np] + streams = [(c_np.astype(np.int64), 0)] + for f in range(NUM_COORD_FIELDS): + for k in (0, 1): + vals = fields[f][idx[k]].astype(np.int64) + if vals.size == 0: + streams.append((vals, 0)) + else: + mn = int(vals.min()) + streams.append((vals - mn, mn)) + return streams + + +def _field_matrix(z0_6_np, m_np): + return np.column_stack([z0_6_np[:, field] for field in range(NUM_COORD_FIELDS - 1)] + [m_np]) + + +def _stream_stats(streams): + counts, sizes, sym_min, alphabets = [], [], [], [] + for symbols, minimum in streams: + alphabet = int(symbols.max()) + 1 if symbols.size else 0 + counts.append(np.bincount(symbols, minlength=alphabet).astype(np.int64)) + sizes.append(int(symbols.size)) + sym_min.append(int(minimum)) + alphabets.append(alphabet) + return counts, sizes, sym_min, alphabets + + +def _rms_and_X(x: torch.Tensor): + rms = x.square().mean(dim=1, keepdim=True).sqrt().clamp_min(MIN_RMS) + X = (x / rms).reshape(-1, LATTICE_DIM).contiguous() + return rms, X + + +def _quantize_at(X, scale): + c, z0_6, m = point_to_fields(nearest_e8(X / scale)) + return c.cpu().numpy(), z0_6.cpu().numpy(), m.cpu().numpy() + + +def _analytic_bytes_from_counts(counts, sizes, rows, prob_bits, tile_elements): + """The cuda model leaves the stream metadata and the fixed slack out, which moves the scale the bisection settles on, so the + two byte models cannot be unified. + """ + coded = rans.coded_bytes(counts, sizes, prob_bits, tile_elements) + return coded + rows * 4 + MODEL_LAYOUT_BYTES + MODEL_STREAM_META_BYTES + 384 + + +def _analytic_total_bytes(c_np, z0_6_np, m_np, rows, prob_bits, tile_elements): + counts, sizes, _sym_min, _alphabets = _stream_stats(_coord_streams(c_np, z0_6_np, m_np)) + return _analytic_bytes_from_counts(counts, sizes, rows, prob_bits, tile_elements) + + +def _bisect_scale(X, rows, cols, target_bpp, prob_bits, tile_elements, n_iter=34, max_vectors=262144): + vecs_per_row = cols // LATTICE_DIM + sub_rows = max(1, min(rows, max_vectors // vecs_per_row)) + index = subsample_index(rows, cols, sub_rows, X.device) + X_sub = X if index is None else X.reshape(rows, cols).index_select(0, index).reshape(-1, LATTICE_DIM) + N_sub = sub_rows * cols + + def total_at(s: float) -> float: + p = nearest_e8(X_sub / s) + c, z0_6, m = point_to_fields(p) + return _analytic_total_bytes( + c.cpu().numpy(), z0_6.cpu().numpy(), m.cpu().numpy(), sub_rows, prob_bits, tile_elements + ) + + lo, hi = 0.001, 8.0 + for _ in range(n_iter): + mid = math.sqrt(lo * hi) + if BITS_PER_BYTE * total_at(mid) / N_sub > target_bpp: + lo = mid + else: + hi = mid + return math.sqrt(lo * hi) + + +def _max_stream_alpha(quantized) -> int: + alpha = 0 + for symbols, _minimum in _coord_streams(*quantized): + if symbols.size: + alpha = max(alpha, int(symbols.max()) + 1) + return alpha + + +def _resolve_scale(x, X, rows, cols, target_bpp, prob_bits, tile_elements, iterations, max_vectors): + """Auto prob_bits grows a bit at a time, only when the alphabet the chosen scale produces overflows the table, up to the + ceiling; that growth is what lets the real rate track ``target_bpp`` at the high end. It never shrinks below its start: + a coarser grid normalizes the distributions less exactly, so shrinking costs rate. + + Past the ceiling, growing further is self-defeating: a wider table prices the same scale cheaper, so the bisection + refines the scale and widens the alphabet again. The scale is coarsened by the actual overflow instead, and only the + alphabet is re-measured; re-bisecting would undo it. + """ + prob_bits, auto_prob_bits = resolve_prob_bits(prob_bits, target_bpp) + if float(x.abs().amax().item()) == 0.0: + return 1.0, _quantize_at(X, 1.0), prob_bits, True + + table_size = 1 << prob_bits + clamped = False + while True: + escalate = False + scale = _bisect_scale( + X, rows, cols, target_bpp, prob_bits, tile_elements, n_iter=iterations, max_vectors=max_vectors + ) + while True: + quantized = _quantize_at(X, scale) + alpha = _max_stream_alpha(quantized) + if alpha <= table_size: + break + if auto_prob_bits and prob_bits < 15: + prob_bits += 1 + table_size = 1 << prob_bits + escalate = True + break + clamped = True + scale *= alpha / table_size * ALPHABET_COARSEN_MARGIN + if not escalate: + if clamped: + report_alphabet_clamp(scale, alpha, table_size) + return scale, quantized, prob_bits, False + + +@dataclass(frozen=True) +class _Candidate: + counts: list + sizes: list + sym_min: list + alphabets: list + fields: torch.Tensor + c_arr: torch.Tensor + scales: torch.Tensor | None = None + row_sse: torch.Tensor | None = None + + +def _refit_row_scales(x, rms, quantized, scale, device): + """Fitting the scale to the points the quantizer actually chose lowers distortion at no rate cost, so it runs on every + encode, not only under the row-RDO pass. + """ + rows, cols = x.shape + c_np, z0_6_np, m_np = quantized + c = torch.from_numpy(c_np).to(device=device, dtype=torch.int64) + z0_6 = torch.from_numpy(z0_6_np).to(device=device, dtype=torch.int64) + m = torch.from_numpy(m_np).to(device=device, dtype=torch.int64) + levels = fields_to_point(c, z0_6, m).reshape(rows, cols) + numerator = (x * levels).sum(1) + denominator = levels.square().sum(1) + scales = torch.where(denominator > 0, (numerator / denominator).clamp_min(MIN_RMS), scale * rms.squeeze(1)) + energy = x.square().sum(1) + row_sse = torch.where(denominator > 0, (energy - numerator * numerator / denominator).clamp_min(0.0), energy) + return scales, row_sse + + +def _candidate_of(quantized, scales, row_sse, device): + c_np, z_np, m_np = quantized + counts, sizes, sym_min, alphabets = _stream_stats(_coord_streams(c_np, z_np, m_np)) + return _Candidate( + counts, sizes, sym_min, alphabets, torch.from_numpy(_field_matrix(z_np, m_np).reshape(-1)).to(device), + torch.from_numpy(c_np).to(device), scales, row_sse, + ) + + +def _candidate_at(x, rms, X, scale, device): + quantized = _quantize_at(X, scale) + scales, row_sse = _refit_row_scales(x, rms, quantized, scale, device) + return _candidate_of(quantized, scales, row_sse, device) + + +def _summarize_fields(fields, c_arr): + field_matrix = fields.cpu().numpy().reshape(-1, LATTICE_DIM) + c_np = c_arr.cpu().numpy() + counts, sizes, sym_min, alphabets = _stream_stats( + _coord_streams(c_np, field_matrix[:, : NUM_COORD_FIELDS - 1], field_matrix[:, NUM_COORD_FIELDS - 1]) + ) + return _Candidate(counts, sizes, sym_min, alphabets, fields, c_arr) + + +def _optimize_rows(x, rms, X, baseline, scale, prob_bits, tile_elements, iterations, candidate_count): + rows, cols = x.shape + device = x.device + ratios = rdo.ratio_ladder(candidate_count) + baseline_index = ratios.index(1.0) + candidates = [] + for index, ratio in enumerate(ratios): + candidate = baseline if index == baseline_index else _candidate_at(x, rms, X, scale * ratio, device) + candidates.append(rdo.compact_candidate(candidate)) + summary, scales = rdo.optimize_rows( + candidates, rows=rows, cols=cols, baseline_index=baseline_index, iterations=iterations, device=device, + summarize=_summarize_fields, + total_bytes=lambda counts, sizes: _analytic_bytes_from_counts(counts, sizes, rows, prob_bits, tile_elements), + ) + chosen = summary.fields.cpu().numpy().reshape(-1, LATTICE_DIM) + return summary.c_arr.cpu().numpy(), chosen[:, : NUM_COORD_FIELDS - 1], chosen[:, NUM_COORD_FIELDS - 1], scales + + +def _encode_streams(c_np, z_np, m_np, prob_bits, tile_elements): + streams = _coord_streams(c_np, z_np, m_np) + n_streams = len(streams) + + table_size = 1 << prob_bits + meta = np.zeros((n_streams, STREAM_META_WIDTH), dtype=np.int64) + frequencies = [] + freq_parts = [] + freq_cursor = 0 + for table, (symbols, sym_min) in enumerate(streams): + n = symbols.size + if n == 0: + frequency = np.empty(0, dtype=np.uint16) + alphabet = 0 + else: + alphabet = int(symbols.max()) + 1 + counts = np.bincount(symbols, minlength=alphabet).astype(np.int64) + frequency = rans.normalize_freq(counts, table_size).astype(np.uint16) + meta[table, META_N_SYMBOLS] = n + meta[table, META_SYM_MIN] = sym_min + meta[table, META_FREQ_OFFSET] = freq_cursor + meta[table, META_ALPHABET] = alphabet + frequencies.append(frequency) + if alphabet: + freq_parts.append(frequency) + freq_cursor += alphabet + + field_symbols = _field_matrix(z_np, m_np).astype(np.int64) + for field in range(NUM_COORD_FIELDS): + table0 = 1 + 2 * field + table1 = table0 + 1 + field_symbols[c_np == 0, field] -= int(meta[table0, META_SYM_MIN]) + field_symbols[c_np == 1, field] -= int(meta[table1, META_SYM_MIN]) + tile_vectors = vector_tile_elements(tile_elements) + payload, states, offsets = rans.encode_vector_stream(c_np, field_symbols, frequencies, prob_bits, tile_vectors) + freq_tables = np.concatenate(freq_parts) if freq_parts else np.empty(0, dtype=np.uint16) + return payload, states, offsets, meta, freq_tables + + +def _encode_impl( + weight, target_bpp, prob_bits, tile_elements, row_rdo_iterations, row_rdo_candidates, scale_search_iterations, + scale_search_max_vectors, +): + if weight.shape[1] % LATTICE_DIM != 0: + raise ValueError( + "lattice_rans quantizes whole E8 vectors, so the column count must be a multiple of " + f"{LATTICE_DIM}; got cols={weight.shape[1]}" + ) + + device = weight.device + x = weight.float().contiguous() + rows, cols = x.shape + + rms, X = _rms_and_X(x) + check_row_scales_finite(rms, weight.dtype) + + s, quantized, prob_bits, all_zero = _resolve_scale( + x, X, rows, cols, target_bpp, prob_bits, tile_elements, scale_search_iterations, scale_search_max_vectors, + ) + c_np, z_np, m_np = quantized + + scales, row_sse = _refit_row_scales(x, rms, quantized, s, device) + if row_rdo_iterations > 0 and not all_zero: + baseline = _candidate_of(quantized, scales, row_sse, device) + c_np, z_np, m_np, scales = _optimize_rows( + x, rms, X, baseline, s, prob_bits, tile_elements, row_rdo_iterations, row_rdo_candidates + ) + + payload, states, offsets, meta, freq_tables = _encode_streams(c_np, z_np, m_np, prob_bits, tile_elements) + + layout = make_layout(cols, prob_bits, tile_elements) + return LatticeBuffers( + payload=torch.from_numpy(payload.astype(np.uint16)).to(device), + offsets=torch.from_numpy(offsets.astype(np.uint32)).to(device), + states=torch.from_numpy(states.astype(np.uint32)).to(device), + stream_meta=torch.from_numpy(meta.astype(np.int32)).to(device), + freq_tables=torch.from_numpy(freq_tables.astype(np.uint16)).to(device), scales=scales.contiguous(), + layout=layout.to(device), + ) + + +def _host_views(buffers: LatticeBuffers): + return ( + buffers.payload.detach().cpu().numpy().astype(np.uint16), + buffers.offsets.detach().cpu().numpy().astype(np.int64), + buffers.states.detach().cpu().numpy().astype(np.uint32), + buffers.stream_meta.detach().cpu().numpy().astype(np.int64), + buffers.freq_tables.detach().cpu().numpy().astype(np.uint16), + buffers.scales.detach().cpu(), + ) + + +def _stream_frequencies(freq_tables: np.ndarray, meta: np.ndarray) -> list: + return [ + freq_tables[int(row[META_FREQ_OFFSET]) : int(row[META_FREQ_OFFSET] + row[META_ALPHABET])] for row in meta + ] + + +def _restore_field_minima(c_np: np.ndarray, shifted_fields: np.ndarray, meta: np.ndarray) -> np.ndarray: + field_values = shifted_fields.copy() + for field in range(NUM_COORD_FIELDS): + table0 = 1 + 2 * field + table1 = table0 + 1 + field_values[:, field] += np.where(c_np == 0, int(meta[table0, META_SYM_MIN]), int(meta[table1, META_SYM_MIN])) + return field_values + + +def _reconstruct(c_np: np.ndarray, field_values: np.ndarray, scales: torch.Tensor, rows: int, cols: int): + c = torch.from_numpy(c_np).to(torch.int64) + z0_6 = torch.from_numpy(field_values[:, :7]).to(torch.int64) + m = torch.from_numpy(field_values[:, 7]).to(torch.int64) + points = fields_to_point(c, z0_6, m) + scale_vec = scales.repeat_interleave(cols // LATTICE_DIM) + return (points * scale_vec.unsqueeze(1)).reshape(rows, cols) + + +def _decode_impl(buffers: LatticeBuffers, shape, dtype): + info = decode_geometry(buffers._asdict(), shape) + rows, cols = info["rows"], info["cols"] + prob_bits = info["prob_bits"] + tile_vectors = vector_tile_elements(info["tile_elements"]) + payload, offsets, states, meta, freq_tables, scales = _host_views(buffers) + + frequencies = _stream_frequencies(freq_tables, meta) + V = rows * (cols // LATTICE_DIM) + c_np, shifted_fields = rans.decode_vector_stream(payload, offsets, states, frequencies, prob_bits, tile_vectors, V) + + field_values = _restore_field_minima(c_np, shifted_fields, meta) + xhat = _reconstruct(c_np, field_values, scales, rows, cols) + + result = snap_to_container(xhat, dtype) + if buffers.layout.device.type != "cpu": + result = result.to(buffers.layout.device) + return result + + +def encode( + weight, *, target_bpp, prob_bits, tile_elements, row_rdo_iterations, row_rdo_candidates, scale_search_iterations, + scale_search_max_vectors, +): + pb = None if prob_bits in (None, 0) else int(prob_bits) + return _encode_impl( + weight, float(target_bpp), pb, int(tile_elements), int(row_rdo_iterations), int(row_rdo_candidates), + int(scale_search_iterations), int(scale_search_max_vectors), + ) + + +def decode(buffers: LatticeBuffers, *, shape, dtype, threads_per_block, l2_prefetch): + return _decode_impl(buffers, shape, dtype) diff --git a/entropack/schemes/lattice_rans/format.py b/entropack/schemes/lattice_rans/format.py new file mode 100644 index 0000000..2e21e90 --- /dev/null +++ b/entropack/schemes/lattice_rans/format.py @@ -0,0 +1,309 @@ +import logging +from dataclasses import dataclass +from typing import Annotated, NamedTuple + +import torch + +from ..config import CompressionConfig, OneOf, Range +from ..base import buffers_fingerprint, cached_parse +from ..tile_ans.format import NUM_STATES + +logger = logging.getLogger("entropack.lattice_rans") + +BLOCK_SIZE = 256 +SHARED_STAGING_HEADROOM = 8 * 1024 +STATIC_SHARED = 256 +MIN_SHARED_BLOCKS_PER_SM = 2 + +SUPPORTED_PROB_BITS = (9, 10, 11, 12, 13, 14, 15) +#: The margin only saves a re-measure and cannot change the outcome, because every step is re-checked against the real +#: histogram. Both lanes read this one definition so they settle on the same scale. +ALPHABET_COARSEN_MARGIN = 1.05 +SUBSAMPLE_SEED_MULTIPLIER = 1000003 + + +@dataclass +class LatticeRANSConfig(CompressionConfig): + """Target-rate compression settings for finite, two-dimensional tensors. + + The encoder selects a quantization scale using sampled storage estimates, fits row + reconstruction scales, and optionally applies per-row rate–distortion refinement. + Decoding restores the input shape and dtype. Tile size and probability precision are + stored with the compressed representation.""" + + #: Requested bits per input element, from 1 to 11 including non-integer values. Read actual_bpp for the stored rate. + target_bpp: Annotated[float, Range(1.0, 11.0)] = 4.0 + #: Probability-table precision. None or zero selects automatically. + prob_bits: Annotated[int | None, OneOf(SUPPORTED_PROB_BITS, silent=(0,))] = None + #: Elements per tile. Larger tiles reduce per-tile metadata and the number of independent decode tasks. None selects by target rate. + tile_elements: Annotated[int | None, Range(1, None)] = None + #: Per-row rate–distortion allocation sweeps using estimated coding costs. Zero disables this refinement. + row_rdo_iterations: Annotated[int, Range(0, 8)] = 0 + #: Number of candidate quantization scales per row for rate–distortion refinement. + row_rdo_candidates: Annotated[int, Range(1, None)] = 5 + #: Number of bisection steps in quantization-scale selection. + scale_search_iterations: Annotated[int, Range(1, None)] = 12 + #: Sampling budget for scale search, measured in eight-value vectors. + scale_search_max_vectors: Annotated[int, Range(1, None)] = 262144 + #: GPU decode block width. None selects a device-dependent value. + threads_per_block: Annotated[int | None, Range(1, None)] = None + #: Prefetch encoded payload into the GPU L2 cache during decoding. + l2_prefetch: bool = True + + +LATTICE_DIM = 8 +NUM_COORD_FIELDS = 8 +NUM_COORD_STREAMS = 2 * NUM_COORD_FIELDS +NUM_STREAMS_FULL = 1 + NUM_COORD_STREAMS + +SUPPORTED_DTYPES = ( + torch.float32, torch.float16, torch.bfloat16, + torch.float8_e4m3fn, torch.float8_e4m3fnuz, torch.float8_e5m2, torch.float8_e5m2fnuz, + torch.int64, torch.int32, torch.int16, torch.int8, torch.uint64, torch.uint32, torch.uint16, torch.uint8, torch.bool, +) +_INTEGER_DTYPES = frozenset({ + torch.int64, torch.int32, torch.int16, torch.int8, torch.uint64, torch.uint32, torch.uint16, torch.uint8, +}) +_EXPECTED_DTYPES = { + "payload": torch.uint16, "offsets": torch.uint32, "states": torch.uint32, "stream_meta": torch.int32, + "freq_tables": torch.uint16, "scales": torch.float32, "layout": torch.int64, +} + +STREAM_META_WIDTH = 4 +META_N_SYMBOLS = 0 +META_SYM_MIN = 1 +META_FREQ_OFFSET = 2 +META_ALPHABET = 3 + +LAYOUT_LEN = 3 + +BITS_PER_BYTE = 8 +#: Not ``LAYOUT_LEN * 8``: tying it to the format would let a layout change move the scale the bisection converges on, so +#: encoded bytes would shift for reasons unrelated to the rate. +MODEL_LAYOUT_BYTES = 64 +#: Pinned the same way and for the same reason, so narrowing ``stream_meta`` changed no encoded byte. +MODEL_STREAM_META_BYTES = 408 + + +class LatticeBuffers(NamedTuple): + payload: torch.Tensor + offsets: torch.Tensor + states: torch.Tensor + stream_meta: torch.Tensor + freq_tables: torch.Tensor + scales: torch.Tensor + layout: torch.Tensor + + +PACKED_KEYS = LatticeBuffers._fields + + +def recommended_tile_elements(target_bpp: float) -> int: + if target_bpp <= 2.0: + return 32768 + if target_bpp <= 4.0: + return 16384 + if target_bpp <= 7.0: + return 8192 + return 4096 + + +def resolve_prob_bits(prob_bits: int | None, target_bpp: float) -> tuple[int, bool]: + if prob_bits in (None, 0): + return (11 if float(target_bpp) <= 7.0 else 12), True + return int(prob_bits), False + + +def report_alphabet_clamp(scale: float, alphabet: int, table_size: int) -> None: + logger.debug( + "lattice_rans coarsened the lattice scale to %.6g so a coordinate alphabet of %d fits the " + "%d-entry rANS table; actual_bpp for this tensor will fall below target_bpp", scale, alphabet, table_size, + ) + + +def subsample_index(rows: int, cols: int, sub_rows: int, device: torch.device) -> torch.Tensor | None: + """A seeded permutation rather than a prefix or a fixed stride: neither is unbiased against the row order a weight happens + to come in. + """ + if sub_rows >= rows: + return None + generator = torch.Generator().manual_seed(rows * SUBSAMPLE_SEED_MULTIPLIER + cols) + return torch.randperm(rows, generator=generator)[:sub_rows].to(device) + + +def _integer_high_bound(dtype: torch.dtype, like: torch.Tensor) -> float: + edge = torch.tensor(float(torch.iinfo(dtype).max) + 1.0, dtype=like.dtype) + below = torch.nextafter(edge, torch.full_like(edge, float("-inf"))) + return float(below) + + +def check_row_scales_finite(rms: torch.Tensor, dtype: torch.dtype) -> None: + if not bool(torch.isfinite(rms).all()): + raise ValueError( + f"lattice_rans normalizes rows in fp32, and the row RMS of this {dtype} tensor overflowed. " + "The limit is sum(weight**2) per row below ~3.4e38, i.e. |values| below " + "sqrt(3.4e38 / cols) -- about 3e17 for a 4k-wide row, 1.6e18 for a 128-wide one. " + "Rescale the source, or use tile_ans, which codes storage bytes and has no numeric " "range limit." + ) + + +def snap_to_container(values: torch.Tensor, dtype: torch.dtype) -> torch.Tensor: + """Saturation is required: torch's float->``float8_e4m3fn`` cast yields NaN above the format's largest finite value and + float->``float16`` yields inf above its own, and the lattice overshoots the source range regularly. A NaN in a decoded + weight destroys inference. + + Integers round half-to-even via ``torch.round``; the CUDA lane matches that with ``rintf``, since ``llroundf`` rounds + half away from zero and would disagree on exact ties. + """ + if dtype == torch.float32: + return values + if dtype == torch.bool: + return values.round() != 0 + if dtype in _INTEGER_DTYPES: + return values.round().clamp(float(torch.iinfo(dtype).min), _integer_high_bound(dtype, values)).to(dtype) + limit = float(torch.finfo(dtype).max) + return values.clamp(-limit, limit).to(dtype) + + +def vector_tile_elements(symbol_tile_elements: int) -> int: + return max(32, symbol_tile_elements // 9) + + +def make_layout(cols: int, prob_bits: int, tile_elements: int) -> torch.Tensor: + return torch.tensor([cols, prob_bits, tile_elements], dtype=torch.int64) + + +def parse_layout(layout: torch.Tensor) -> dict: + if layout.dtype != torch.int64 or layout.ndim != 1 or layout.numel() != LAYOUT_LEN: + raise ValueError(f"lattice_rans layout must be int64[{LAYOUT_LEN}]") + cols, prob_bits, tile_elements = (int(x) for x in layout.detach().cpu().tolist()) + if cols <= 0 or cols % LATTICE_DIM != 0: + raise ValueError("lattice_rans vector-plane format requires positive cols divisible by 8") + if prob_bits not in SUPPORTED_PROB_BITS: + raise ValueError(f"lattice_rans prob_bits {prob_bits} unsupported") + if tile_elements <= 0: + raise ValueError("lattice_rans tile_elements must be positive") + return {"cols": cols, "prob_bits": prob_bits, "tile_elements": tile_elements} + + +def _check_container(buffers: dict, shape: tuple, dtype: torch.dtype) -> None: + missing = [k for k in PACKED_KEYS if k not in buffers] + if missing: + raise ValueError(f"lattice_rans packed data is missing buffers: {missing}") + if set(buffers) != set(PACKED_KEYS): + extra = sorted(set(buffers) - set(PACKED_KEYS)) + raise ValueError(f"lattice_rans packed data has unexpected buffers: {extra}") + if not all(isinstance(buffers[k], torch.Tensor) for k in PACKED_KEYS): + raise TypeError("lattice_rans packed buffers must be torch.Tensor values") + if dtype not in SUPPORTED_DTYPES: + raise ValueError(f"lattice_rans does not support output dtype {dtype}") + if len(shape) != 2 or any(not isinstance(d, int) or d <= 0 for d in shape): + raise ValueError(f"lattice_rans shape must be non-empty 2D, got {shape}") + + +def _check_buffers(buffers: dict) -> None: + devices = {buffers[k].device for k in PACKED_KEYS} + if len(devices) != 1: + raise ValueError(f"lattice_rans packed buffers must share one device, got {devices}") + for k in PACKED_KEYS: + if not buffers[k].is_contiguous(): + raise ValueError(f"lattice_rans buffer '{k}' must be contiguous") + for k, dt in _EXPECTED_DTYPES.items(): + if buffers[k].dtype != dt: + raise ValueError(f"lattice_rans buffer '{k}' must be {dt}, got {buffers[k].dtype}") + + +def _check_scales(buffers: dict, info: dict) -> None: + scales = buffers["scales"] + if scales.ndim != 1 or scales.numel() != info["rows"]: + raise ValueError(f"lattice_rans scales must be 1D with one entry per row, got {tuple(scales.shape)}") + if not bool((torch.isfinite(scales) & (scales > 0)).all().item()): + raise ValueError("lattice_rans scales must be finite and positive") + + +def _check_tile_geometry(buffers: dict, info: dict) -> None: + n_streams, total_tiles = info["n_streams"], info["total_tiles"] + meta = buffers["stream_meta"] + if meta.ndim != 2 or tuple(meta.shape) != (n_streams, STREAM_META_WIDTH): + raise ValueError(f"lattice_rans stream_meta must be ({n_streams},{STREAM_META_WIDTH}), got {tuple(meta.shape)}") + if tuple(buffers["states"].shape) != (total_tiles, NUM_STATES): + raise ValueError("lattice_rans states shape does not match total_tiles") + if buffers["offsets"].numel() != total_tiles + 1: + raise ValueError("lattice_rans offsets length must be total_tiles+1") + offsets = buffers["offsets"].to(torch.int64) + if int(offsets[0].item()) != 0 or int(offsets[-1].item()) != buffers["payload"].numel(): + raise ValueError("lattice_rans payload offsets endpoints are invalid") + if total_tiles > 0 and bool((offsets[1:] < offsets[:-1]).any().item()): + raise ValueError("lattice_rans payload offsets must be monotone") + + +def _check_stream_metadata(buffers: dict, info: dict) -> None: + n_streams = info["n_streams"] + rows, cols = info["rows"], info["cols"] + table_size = 1 << info["prob_bits"] + meta_cpu = buffers["stream_meta"].detach().cpu().to(torch.int64) + freq_cpu = buffers["freq_tables"].detach().cpu().to(torch.int64) + freq_total = freq_cpu.numel() + running_freq = 0 + symbol_counts = [] + for table in range(n_streams): + row = meta_cpu[table] + n_sym = int(row[META_N_SYMBOLS]) + freq_off = int(row[META_FREQ_OFFSET]) + alphabet = int(row[META_ALPHABET]) + if n_sym < 0 or alphabet < 0 or freq_off < 0: + raise ValueError("lattice_rans table metadata has a negative field") + if freq_off != running_freq or freq_off + alphabet > freq_total: + raise ValueError("lattice_rans frequency-table ranges must be contiguous") + if (n_sym == 0) != (alphabet == 0): + raise ValueError("lattice_rans empty table metadata is inconsistent") + if alphabet: + if alphabet > table_size: + raise ValueError("lattice_rans table alphabet exceeds probability precision") + if int(freq_cpu[freq_off : freq_off + alphabet].sum().item()) != table_size: + raise ValueError("lattice_rans normalized frequencies have an invalid sum") + running_freq += alphabet + symbol_counts.append(n_sym) + if running_freq != freq_total: + raise ValueError("lattice_rans frequency table has trailing entries") + vectors = rows * (cols // LATTICE_DIM) + if symbol_counts[0] != vectors or int(meta_cpu[0, META_SYM_MIN]) != 0: + raise ValueError("lattice_rans coset table metadata is invalid") + n0, n1 = symbol_counts[1], symbol_counts[2] + if n0 + n1 != vectors: + raise ValueError("lattice_rans conditioned table counts do not cover all vectors") + for field in range(NUM_COORD_FIELDS): + if symbol_counts[1 + 2 * field] != n0 or symbol_counts[2 + 2 * field] != n1: + raise ValueError("lattice_rans conditioned table counts disagree across fields") + + +def decode_geometry(buffers: dict, shape: tuple) -> dict: + stored = cached_parse(buffers["layout"], parse_layout, "_lattice_rans_layout") + rows = int(shape[0]) + tile_vectors = vector_tile_elements(stored["tile_elements"]) + vectors = rows * (stored["cols"] // LATTICE_DIM) + return { + **stored, "rows": rows, "n_streams": NUM_STREAMS_FULL, + "total_tiles": (vectors + tile_vectors - 1) // tile_vectors, + } + + +def validate_packed(buffers: dict, shape: tuple, dtype: torch.dtype) -> dict: + _check_container(buffers, shape, dtype) + + cached = getattr(buffers["layout"], "_lattice_rans_validation_cache", None) + fp = buffers_fingerprint(buffers, shape, dtype) + if cached is not None and cached[0] == fp: + return cached[1] + + _check_buffers(buffers) + info = decode_geometry(buffers, shape) + rows, cols = tuple(shape) + if not 0 <= info["cols"] - cols < LATTICE_DIM: + raise ValueError(f"lattice_rans layout columns {info['cols']} do not match {(rows, cols)}") + _check_scales(buffers, info) + _check_tile_geometry(buffers, info) + _check_stream_metadata(buffers, info) + + buffers["layout"]._lattice_rans_validation_cache = (fp, info) + return info diff --git a/entropack/schemes/lattice_rans/lattice_rans.cu b/entropack/schemes/lattice_rans/lattice_rans.cu new file mode 100644 index 0000000..8f6d96d --- /dev/null +++ b/entropack/schemes/lattice_rans/lattice_rans.cu @@ -0,0 +1,1140 @@ +// EntroPack lattice_rans CUDA codec: E8-lattice vector quantization + coset-conditioned rANS. +// +// Encode kernels: +// e8_quantize_fields_kernel/_f32 nearest-E8 (Conway-Sloane: the two cosets D8 and D8+g, each reduced by a parity fix on +// the max-residual coordinate) fused with the point->fields split (coset c, parity-reduced +// coordinates z0..z6, m) and a block-reduced per-stream min/max pass. Bit-exact with +// eager.nearest_e8 / point_to_fields: rintf == torch.round (half-to-even), the squared +// distances use torch's sum(dim=1) tree order, __fmul_rn/__fadd_rn block FMA contraction, +// argmax keeps the lowest index on ties, and the coset tie rule is d0 <= d1. The two +// variants differ only in how the weight is read: bf16 native, fp32 for every other +// container. +// e8_refit_scales_kernel/_f32 least-squares refit of each row scale against the quantized points. +// e8_minmax_fields_kernel per-stream symbol min/max; sizes the alphabets. +// e8_histogram_kernel exact per-stream symbol counts. Bins are staged in shared memory when they fit the +// device budget, and fall back to global atomics otherwise. +// e8_rans_encode_vector_kernel tiled 32-state interleaved rANS encode, one warp per tile, bit-exact with +// tile_ans/encode_cpu._encode_stream: the state machine of device.cuh, 16-bit words, +// ballot-coalesced renormalization emission. +// e8_compact_kernel gathers the per-tile scratch words into one payload. +// +// Decode kernels: one warp decodes one tile and reconstructs straight into the container, with no symbol scratch and no +// prefix sum. The variants differ only in how the slot->symbol table is held, and all produce identical output, so the host +// picks one from the stored alphabet widths and the queried device limits: +// e8_decode_vector_shlut8pf_* slot->symbol (uint8) and begin|freq tables staged in shared memory. Used when every +// alphabet is below 256 and the staged table still leaves the SM its resident-block +// target. +// e8_decode_vector_packed32pf_* symbol|freq|delta packed into 32 bits in global memory, so no shared memory is needed +// and occupancy is unrestricted, plus an L2 prefetch of the descending renorm stream. +// e8_decode_vector_packed32_* the same without the prefetch, reached by decoding with l2_prefetch=False. +// e8_decode_vector_fused_* 64-bit table in global memory; the fallback for alphabets too wide to pack into 32 bits. +// A "_g" suffix marks the generic-container twin, which takes a store kind and writes any supported container instead of +// bf16. + +#include +#include +#include +#include +#include "device.cuh" + +namespace { + +constexpr int kCMax = 32; +constexpr int kMinMaxLen = kCMax + 1; +constexpr int kIntMax = 0x7FFFFFFF; +constexpr int kIntMin = -0x7FFFFFFF - 1; + +constexpr int kMetaStride = 4; +constexpr int kMetaSymMin = 1; +constexpr int kMetaFreqOff = 2; + +__device__ __forceinline__ void set_error(int* error, int code) { + atomicCAS(error, 0, code); +} + +__device__ __forceinline__ int div_i64_i32(int64_t a, int b) { + return static_cast(a / b); +} + +__device__ __forceinline__ void nearest_dn_int( + const float* __restrict__ y, int* __restrict__ out) { + float f[8]; + float r[8]; + int parity = 0; +#pragma unroll + for (int j = 0; j < 8; ++j) { + f[j] = rintf(y[j]); // torch.round == round-half-to-even == rintf + r[j] = __fsub_rn(y[j], f[j]); + parity += static_cast(f[j]); // exact: f is an integer float within int32 range + } + int idx = 0; + float best = -1.0f; +#pragma unroll + for (int j = 0; j < 8; ++j) { + const float a = fabsf(r[j]); + if (a > best) { // strict > keeps the lowest index on ties (torch argmax) + best = a; + idx = j; + } + } +#pragma unroll + for (int j = 0; j < 8; ++j) { + out[j] = static_cast(f[j]); + } + if ((parity & 1) != 0) { + out[idx] += (r[idx] < 0.0f) ? -1 : 1; // torch.sign(0) is replaced by +1 in eager + } +} + +// Accumulated in float64, in torch's sum(dim=1) tree order, so exact ties resolve identically to eager.nearest_e8: +// the parenthesization is the point and must not be flattened. +__device__ __forceinline__ double dist2_f64( + const float* __restrict__ y, const int* __restrict__ z, float offset) { + double s[8]; +#pragma unroll + for (int j = 0; j < 8; ++j) { + const double c = static_cast(z[j]) + static_cast(offset); + const double r = static_cast(y[j]) - c; + s[j] = r * r; + } + return ((s[0] + s[4]) + (s[2] + s[6])) + ((s[1] + s[5]) + (s[3] + s[7])); +} + +__device__ __forceinline__ void block_minmax_init(int* shared) { + for (int t = threadIdx.x; t < kMinMaxLen; t += blockDim.x) { + shared[t] = ((t & 1) == 0 && t != kCMax) ? kIntMax : kIntMin; + } + __syncthreads(); +} + +__device__ __forceinline__ void block_minmax_flush(int* shared, int* __restrict__ minmax) { + __syncthreads(); + for (int t = threadIdx.x; t < kMinMaxLen; t += blockDim.x) { + const int v = shared[t]; + if ((t & 1) == 0 && t != kCMax) { + if (v != kIntMax) atomicMin(&minmax[t], v); + } else { + if (v != kIntMin) atomicMax(&minmax[t], v); + } + } +} + +} // namespace + +__device__ __forceinline__ void e8_load8(const __nv_bfloat16* w, float* xf) { + const uint4 packed = *reinterpret_cast(w); + const unsigned words[4] = {packed.x, packed.y, packed.z, packed.w}; +#pragma unroll + for (int j = 0; j < 8; ++j) { + const unsigned bits = (words[j >> 1] >> ((j & 1) * 16)) & 0xFFFFu; + xf[j] = __bfloat162float(*reinterpret_cast(&bits)); + } +} + +__device__ __forceinline__ void e8_load8(const float* w, float* xf) { + const float4 lo = *reinterpret_cast(w); + const float4 hi = *reinterpret_cast(w + 4); + xf[0] = lo.x; xf[1] = lo.y; xf[2] = lo.z; xf[3] = lo.w; + xf[4] = hi.x; xf[5] = hi.y; xf[6] = hi.z; xf[7] = hi.w; +} + +__device__ __forceinline__ float e8_load_scalar(const __nv_bfloat16* p) { + return __bfloat162float(*p); +} + +__device__ __forceinline__ float e8_load_scalar(const float* p) { + return *p; +} + +template +__device__ __forceinline__ void e8_quantize_fields_body( + const InT* __restrict__ weight, + const float* __restrict__ rms, + float s, + int rows, + int cols, + int vecs_per_row, + int* __restrict__ fields, // [V,8] int32: z0..z6, m + int* __restrict__ c_arr, // [V] int32 + int* __restrict__ minmax) { // kMinMaxLen int32, pre-initialized to kIntMax/kIntMin + __shared__ int sm[kMinMaxLen]; + block_minmax_init(sm); + + const int64_t V = static_cast(rows) * vecs_per_row; + const int64_t stride = static_cast(gridDim.x) * blockDim.x; + + for (int64_t idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + idx < V; idx += stride) { + const int row = div_i64_i32(idx, vecs_per_row); + const int col8 = static_cast(idx - static_cast(row) * vecs_per_row); + const float rmsv = rms[row]; + const InT* w = weight + static_cast(row) * cols + col8 * 8; + float xf[8]; + e8_load8(w, xf); + float y[8]; +#pragma unroll + for (int j = 0; j < 8; ++j) { + y[j] = __fdiv_rn(__fdiv_rn(xf[j], rmsv), s); + } + + int z0[8]; + int z1[8]; + nearest_dn_int(y, z0); + float ym[8]; +#pragma unroll + for (int j = 0; j < 8; ++j) { + ym[j] = __fsub_rn(y[j], 0.5f); + } + nearest_dn_int(ym, z1); + const double d0 = dist2_f64(y, z0, 0.0f); + const double d1 = dist2_f64(y, z1, 0.5f); + const bool pick0 = d0 <= d1; // tie picks the D8 coset, exactly as eager + const int* z = pick0 ? z0 : z1; + const int c = pick0 ? 0 : 1; + + int zv[8]; + int par = 0; +#pragma unroll + for (int j = 0; j < 7; ++j) { + zv[j] = z[j]; + par += zv[j]; + } + par &= 1; // two's-complement &1 == torch.remainder(., 2) + zv[7] = (z[7] - par) >> 1; // m = floor((z7 - par)/2) + + int* fout = fields + idx * 8; +#pragma unroll + for (int j = 0; j < 8; ++j) { + fout[j] = zv[j]; + } + c_arr[idx] = c; + + atomicMax(&sm[kCMax], c); +#pragma unroll + for (int j = 0; j < 8; ++j) { + const int st = j * 2 + c; // coord stream index (field f, coset k) + atomicMin(&sm[st * 2], zv[j]); + atomicMax(&sm[st * 2 + 1], zv[j]); + } + } + block_minmax_flush(sm, minmax); +} + +#define E8_QUANTIZE_WRAPPER(NAME, INT) \ + extern "C" __global__ void NAME( \ + const INT* __restrict__ weight, \ + const float* __restrict__ rms, \ + float s, \ + int rows, \ + int cols, \ + int vecs_per_row, \ + int* __restrict__ fields, \ + int* __restrict__ c_arr, \ + int* __restrict__ minmax) { \ + e8_quantize_fields_body( \ + weight, rms, s, rows, cols, vecs_per_row, fields, c_arr, minmax); \ + } + +E8_QUANTIZE_WRAPPER(e8_quantize_fields_kernel, __nv_bfloat16) +E8_QUANTIZE_WRAPPER(e8_quantize_fields_f32_kernel, float) + +#undef E8_QUANTIZE_WRAPPER + +template +__device__ __forceinline__ void e8_refit_scales_body( + const InT* __restrict__ weight, + const int* __restrict__ fields, // [V,8]: z0..z6, m + const int* __restrict__ c_arr, // [V] + const float* __restrict__ rms, + float initial_s, + float* __restrict__ scales, + float* __restrict__ row_sse, + int rows, + int cols, + int vecs_per_row) { + const int row = blockIdx.x; + if (row >= rows) return; + + float numerator = 0.0f; + float denominator = 0.0f; + float energy = 0.0f; + for (int element = threadIdx.x; element < cols; element += blockDim.x) { + const int vector_in_row = element >> 3; + const int coordinate = element & 7; + const int64_t vector = static_cast(row) * vecs_per_row + vector_in_row; + const int* vector_fields = fields + vector * 8; + const int c = c_arr[vector]; + int z; + if (coordinate < 7) { + z = vector_fields[coordinate]; + } else { + int parity = 0; +#pragma unroll + for (int j = 0; j < 7; ++j) parity += vector_fields[j]; + z = 2 * vector_fields[7] + (parity & 1); + } + const float point = __fadd_rn(static_cast(z), c ? 0.5f : 0.0f); + const float value = e8_load_scalar( + weight + static_cast(row) * cols + element); + numerator = __fadd_rn(numerator, __fmul_rn(value, point)); + denominator = __fadd_rn(denominator, __fmul_rn(point, point)); + energy = __fadd_rn(energy, __fmul_rn(value, value)); + } + + __shared__ float numerator_shared[256]; + __shared__ float denominator_shared[256]; + __shared__ float energy_shared[256]; + numerator_shared[threadIdx.x] = numerator; + denominator_shared[threadIdx.x] = denominator; + energy_shared[threadIdx.x] = energy; + __syncthreads(); + for (int offset = blockDim.x >> 1; offset > 0; offset >>= 1) { + if (threadIdx.x < offset) { + numerator_shared[threadIdx.x] += numerator_shared[threadIdx.x + offset]; + denominator_shared[threadIdx.x] += denominator_shared[threadIdx.x + offset]; + energy_shared[threadIdx.x] += energy_shared[threadIdx.x + offset]; + } + __syncthreads(); + } + if (threadIdx.x == 0) { + const float fallback = __fmul_rn(initial_s, rms[row]); + const float fitted = denominator_shared[0] > 0.0f + ? __fdiv_rn(numerator_shared[0], denominator_shared[0]) + : fallback; + scales[row] = (isfinite(fitted) && fitted > 0.0f) ? fitted : fallback; + row_sse[row] = denominator_shared[0] > 0.0f + ? fmaxf(0.0f, energy_shared[0] - numerator_shared[0] * numerator_shared[0] / + denominator_shared[0]) + : energy_shared[0]; + } +} + +#define E8_REFIT_WRAPPER(NAME, INT) \ + extern "C" __global__ void NAME( \ + const INT* __restrict__ weight, \ + const int* __restrict__ fields, \ + const int* __restrict__ c_arr, \ + const float* __restrict__ rms, \ + float initial_s, \ + float* __restrict__ scales, \ + float* __restrict__ row_sse, \ + int rows, \ + int cols, \ + int vecs_per_row) { \ + e8_refit_scales_body( \ + weight, fields, c_arr, rms, initial_s, scales, row_sse, \ + rows, cols, vecs_per_row); \ + } + +E8_REFIT_WRAPPER(e8_refit_scales_kernel, __nv_bfloat16) +E8_REFIT_WRAPPER(e8_refit_scales_f32_kernel, float) + +#undef E8_REFIT_WRAPPER + +extern "C" __global__ void e8_minmax_fields_kernel( + const int* __restrict__ fields, + const int* __restrict__ c_arr, + int64_t V, + int* __restrict__ minmax) { + __shared__ int sm[kMinMaxLen]; + block_minmax_init(sm); + const int64_t stride = static_cast(gridDim.x) * blockDim.x; + for (int64_t idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + idx < V; idx += stride) { + const int c = c_arr[idx]; + atomicMax(&sm[kCMax], c); + const int* vector_fields = fields + idx * 8; +#pragma unroll + for (int j = 0; j < 8; ++j) { + const int stream = j * 2 + c; + const int value = vector_fields[j]; + atomicMin(&sm[stream * 2], value); + atomicMax(&sm[stream * 2 + 1], value); + } + } + block_minmax_flush(sm, minmax); +} + +extern "C" __global__ void e8_histogram_kernel( + const int* __restrict__ fields, + const int* __restrict__ c_arr, + const int* __restrict__ minmax, + const int* __restrict__ bin_off, // [17]: coset (2 bins), then the 16 coord streams + int* __restrict__ bins, + int64_t V, + int shared_bins) { // >0: stage in dynamic shared memory of this many ints + extern __shared__ int sbins[]; + int* target; + if (shared_bins > 0) { + for (int t = threadIdx.x; t < shared_bins; t += blockDim.x) sbins[t] = 0; + __syncthreads(); + target = sbins; + } else { + target = bins; + } + + const int64_t stride = static_cast(gridDim.x) * blockDim.x; + for (int64_t idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + idx < V; idx += stride) { + const int c = c_arr[idx]; + atomicAdd(&target[bin_off[0] + c], 1); + const int* f = fields + idx * 8; +#pragma unroll + for (int j = 0; j < 8; ++j) { + const int st = j * 2 + c; + atomicAdd(&target[bin_off[1 + st] + (f[j] - minmax[st * 2])], 1); + } + } + if (shared_bins > 0) { + __syncthreads(); + for (int t = threadIdx.x; t < shared_bins; t += blockDim.x) { + if (sbins[t]) atomicAdd(&bins[t], sbins[t]); + } + } +} + +extern "C" __global__ void e8_rans_encode_vector_kernel( + const int* __restrict__ fields, + const int* __restrict__ c_arr, + const uint16_t* __restrict__ freq_tables, + const uint16_t* __restrict__ cdfs, + const int* __restrict__ table_meta, + uint16_t* __restrict__ scratch, + uint32_t* __restrict__ word_counts, + uint32_t* __restrict__ final_states, + int64_t num_vectors, + int tile_vectors, + int num_tiles) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int tile = blockIdx.x * warps_per_block + warp_in_block; + if (tile >= num_tiles) return; + const int lane = threadIdx.x & 31; + const int64_t tile_begin = static_cast(tile) * tile_vectors; + const int tile_count = static_cast( + num_vectors - tile_begin < tile_vectors ? num_vectors - tile_begin : tile_vectors); + uint16_t* tile_scratch = scratch + static_cast(tile) * tile_vectors * 9; + uint32_t state = kStateMin; + uint32_t word_count = 0; + constexpr uint32_t state_check_mul = 1u << (31 - kProbBits); + + for (int base = 0; base < tile_count; base += kNumStates) { + const bool valid = base + lane < tile_count; + const int64_t vector = tile_begin + base + lane; + const int c = valid ? c_arr[vector] : 0; +#pragma unroll + for (int field = 7; field >= 0; --field) { + const int table = 1 + 2 * field + c; + const uint32_t symbol = valid ? static_cast( + fields[vector * 8 + field] - table_meta[table * kMetaStride + kMetaSymMin]) : 0u; + const int freq_off = valid ? table_meta[table * kMetaStride + kMetaFreqOff] : 0; + const uint32_t frequency = valid ? freq_tables[freq_off + symbol] : 1u; + const bool emit = valid && state >= frequency * state_check_mul; + const uint32_t vote = __ballot_sync(0xffffffffu, emit); + const uint32_t prefix = __popc(vote & lane_mask_lt()); + if (emit) { + tile_scratch[word_count + prefix] = static_cast(state); + state >>= 16; + } + word_count += __popc(vote); + if (valid) { + state = (state / frequency) * kTableSize + (state % frequency) + + cdfs[freq_off + symbol]; + } + } + const uint32_t symbol = static_cast(c); + const uint32_t frequency = valid ? freq_tables[symbol] : 1u; + const bool emit = valid && state >= frequency * state_check_mul; + const uint32_t vote = __ballot_sync(0xffffffffu, emit); + const uint32_t prefix = __popc(vote & lane_mask_lt()); + if (emit) { + tile_scratch[word_count + prefix] = static_cast(state); + state >>= 16; + } + word_count += __popc(vote); + if (valid) { + state = (state / frequency) * kTableSize + (state % frequency) + cdfs[symbol]; + } + } + final_states[tile * kNumStates + lane] = state; + if (lane == 0) word_counts[tile] = word_count; +} + +extern "C" __global__ void e8_compact_kernel( + const uint16_t* __restrict__ scratch, + const uint32_t* __restrict__ offsets, + uint16_t* __restrict__ payload, + int tile_elements, + int num_tiles) { + const int tile = blockIdx.x; + if (tile >= num_tiles) return; + const uint32_t begin = offsets[tile]; + const uint32_t count = offsets[tile + 1] - begin; + const uint16_t* source = scratch + static_cast(tile) * tile_elements; + for (uint32_t i = threadIdx.x; i < count; i += blockDim.x) { + payload[begin + i] = source[i]; + } +} + +__device__ __forceinline__ uint32_t e8_rans_decode_symbol( + uint32_t& state, const uint64_t* __restrict__ table) { + const uint32_t slot = state & kStateMask; + const uint64_t entry = __ldg(table + slot); + const uint32_t symbol = static_cast(entry & 0xffffu); + uint32_t frequency = static_cast((entry >> 16) & 0xffffu); + if (frequency == 0) frequency = kTableSize; + const uint32_t cdf = static_cast(entry >> 32); + state = frequency * (state >> kProbBits) + (slot - cdf); + return symbol; +} + +__device__ __forceinline__ uint32_t e8_decode_coset( + uint32_t& state, uint32_t frequency0) { + const uint32_t slot = state & kStateMask; + const uint32_t symbol = slot >= frequency0; + const uint32_t begin = symbol ? frequency0 : 0u; + const uint32_t frequency = symbol ? kTableSize - frequency0 : frequency0; + state = frequency * (state >> kProbBits) + (slot - begin); + return symbol; +} + +__device__ __forceinline__ uint32_t e8_decode_packed32( + uint32_t& state, + const uint32_t* __restrict__ tables, + const int* __restrict__ pack_bits, + int table) { + const uint32_t slot = state & kStateMask; + const uint32_t entry = __ldg(tables + static_cast(table) * kTableSize + slot); + const int bits = pack_bits[table]; + const int symbol_bits = bits & 0xff; + const int frequency_bits = (bits >> 8) & 0xff; + const uint32_t symbol_mask = (1u << symbol_bits) - 1u; + const uint32_t frequency_mask = (1u << frequency_bits) - 1u; + const uint32_t symbol = entry & symbol_mask; + uint32_t frequency = (entry >> symbol_bits) & frequency_mask; + if (frequency == 0) frequency = kTableSize; + const uint32_t delta = entry >> (symbol_bits + frequency_bits); + state = frequency * (state >> kProbBits) + delta; + return symbol; +} + +enum : int { + kStoreF32 = 0, + kStoreF16 = 1, + kStoreF8E4M3 = 2, + kStoreF8E5M2 = 3, + kStoreI8 = 4, + kStoreI16 = 5, + kStoreI32 = 6, + kStoreI64 = 7, + kStoreU8 = 8, + kStoreU16 = 9, + kStoreU32 = 10, + kStoreU64 = 11, + kStoreBool = 12, +}; + +template struct e8_int_range; + +#define E8_INT_RANGE(TYPE, LOW, HIGH_EXCLUSIVE) \ + template <> struct e8_int_range { \ + static constexpr float kLow = LOW; \ + static constexpr float kHighExclusive = HIGH_EXCLUSIVE; \ + } + +E8_INT_RANGE(int8_t, -128.0f, 128.0f); +E8_INT_RANGE(int16_t, -32768.0f, 32768.0f); +E8_INT_RANGE(int32_t, -2147483648.0f, 2147483648.0f); +E8_INT_RANGE(int64_t, -9223372036854775808.0f, 9223372036854775808.0f); +E8_INT_RANGE(uint8_t, 0.0f, 256.0f); +E8_INT_RANGE(uint16_t, 0.0f, 65536.0f); +E8_INT_RANGE(uint32_t, 0.0f, 4294967296.0f); +E8_INT_RANGE(uint64_t, 0.0f, 18446744073709551616.0f); + +#undef E8_INT_RANGE + +template +__device__ __forceinline__ T e8_snap_integer(float value) { + const float high = nextafterf(e8_int_range::kHighExclusive, 0.0f); + const float clamped = fminf(fmaxf(rintf(value), e8_int_range::kLow), high); + return static_cast(clamped); +} + +struct BF16Store { + using pointer = __nv_bfloat16*; + + __device__ static void put8_words( + pointer out, int64_t vector, const int* value, float chalf, float scale, int) { + uint4 packed; + unsigned* words = reinterpret_cast(&packed); +#pragma unroll + for (int pair = 0; pair < 4; ++pair) { + const float p0 = __fadd_rn(static_cast(value[2 * pair]), chalf); + const float p1 = __fadd_rn(static_cast(value[2 * pair + 1]), chalf); + const __nv_bfloat16 lo = __float2bfloat16_rn(__fmul_rn(p0, scale)); + const __nv_bfloat16 hi = __float2bfloat16_rn(__fmul_rn(p1, scale)); + const unsigned lo_bits = reinterpret_cast(&lo)[0]; + const unsigned hi_bits = reinterpret_cast(&hi)[0]; + words[pair] = lo_bits | (hi_bits << 16); + } + *reinterpret_cast(out + vector * 8) = packed; + } + + __device__ static void put8_pair( + pointer out, int64_t vector, const int* value, float chalf, float scale, int) { + uint4 packed; + unsigned* words = reinterpret_cast(&packed); +#pragma unroll + for (int pair = 0; pair < 4; ++pair) { + const float p0 = __fadd_rn(static_cast(value[2 * pair]), chalf); + const float p1 = __fadd_rn(static_cast(value[2 * pair + 1]), chalf); + const __nv_bfloat162 pair_bf = __floats2bfloat162_rn( + __fmul_rn(p0, scale), __fmul_rn(p1, scale)); + words[pair] = *reinterpret_cast(&pair_bf); + } + *reinterpret_cast(out + vector * 8) = packed; + } +}; + +template struct e8_conv; + +template <> struct e8_conv { + using type = uint32_t; + __device__ static uint32_t from(float v) { return __float_as_uint(v); } +}; + +template <> struct e8_conv { + using type = uint16_t; + __device__ static uint16_t from(float v) { + return __half_as_ushort(__float2half_rn(fminf(fmaxf(v, -65504.0f), 65504.0f))); + } +}; + +#define E8_CONV_FP8(K, LIMIT, INTERP) \ + template <> struct e8_conv { \ + using type = __nv_fp8_storage_t; \ + __device__ static __nv_fp8_storage_t from(float v) { \ + return __nv_cvt_float_to_fp8( \ + fminf(fmaxf(v, -LIMIT), LIMIT), __NV_SATFINITE, INTERP); \ + } \ + } + +E8_CONV_FP8(kStoreF8E4M3, 448.0f, __NV_E4M3); +E8_CONV_FP8(kStoreF8E5M2, 57344.0f, __NV_E5M2); + +#undef E8_CONV_FP8 + +#define E8_CONV_INT(K, TYPE) \ + template <> struct e8_conv { \ + using type = TYPE; \ + __device__ static TYPE from(float v) { \ + return e8_snap_integer(v); \ + } \ + } + +E8_CONV_INT(kStoreI8, int8_t); +E8_CONV_INT(kStoreI16, int16_t); +E8_CONV_INT(kStoreI32, int32_t); +E8_CONV_INT(kStoreI64, int64_t); +E8_CONV_INT(kStoreU8, uint8_t); +E8_CONV_INT(kStoreU16, uint16_t); +E8_CONV_INT(kStoreU32, uint32_t); +E8_CONV_INT(kStoreU64, uint64_t); + +#undef E8_CONV_INT + +template <> struct e8_conv { + using type = bool; + __device__ static bool from(float v) { return rintf(v) != 0.0f; } +}; + +// Store count, not conversion cost, dominates here, so the eight converted values are packed into words and written with as +// few wide stores as the container allows. Only 8-byte containers store scalar: one vector of them is wider than a uint4. +template +__device__ __forceinline__ void e8_store8(void* out, int64_t vector, const float* v) { + using T = typename Conv::type; + constexpr int kBytes = 8 * static_cast(sizeof(T)); + char* base = static_cast(out) + vector * kBytes; + if constexpr (sizeof(T) == 8) { +#pragma unroll + for (int j = 0; j < 8; ++j) reinterpret_cast(base)[j] = Conv::from(v[j]); + } else { + constexpr int kWords = kBytes / 4; + constexpr int kPerWord = 4 / static_cast(sizeof(T)); + unsigned words[kWords]; +#pragma unroll + for (int word = 0; word < kWords; ++word) { + unsigned packed = 0; +#pragma unroll + for (int slot = 0; slot < kPerWord; ++slot) { + const unsigned bits = static_cast( + sizeof(T) == 1 ? static_cast(Conv::from(v[word * kPerWord + slot])) + : sizeof(T) == 2 ? static_cast(Conv::from(v[word * kPerWord + slot])) + : static_cast(Conv::from(v[word * kPerWord + slot]))); + packed |= bits << (8 * static_cast(sizeof(T)) * slot); + } + words[word] = packed; + } + if constexpr (kWords == 2) { + *reinterpret_cast(base) = make_uint2(words[0], words[1]); + } else { + *reinterpret_cast(base) = make_uint4(words[0], words[1], words[2], words[3]); + if constexpr (kWords == 8) { + *reinterpret_cast(base + 16) = + make_uint4(words[4], words[5], words[6], words[7]); + } + } + } +} + +struct GenericStore { + using pointer = void*; + + __device__ static void put8( + pointer out, int64_t vector, const int* value, float chalf, float scale, int kind) { + float v[8]; +#pragma unroll + for (int j = 0; j < 8; ++j) { + v[j] = __fmul_rn(__fadd_rn(static_cast(value[j]), chalf), scale); + } + switch (kind) { + case kStoreF32: e8_store8>(out, vector, v); break; + case kStoreF16: e8_store8>(out, vector, v); break; + case kStoreF8E4M3: e8_store8>(out, vector, v); break; + case kStoreF8E5M2: e8_store8>(out, vector, v); break; + case kStoreI8: e8_store8>(out, vector, v); break; + case kStoreI16: e8_store8>(out, vector, v); break; + case kStoreI32: e8_store8>(out, vector, v); break; + case kStoreI64: e8_store8>(out, vector, v); break; + case kStoreU8: e8_store8>(out, vector, v); break; + case kStoreU16: e8_store8>(out, vector, v); break; + case kStoreU32: e8_store8>(out, vector, v); break; + case kStoreU64: e8_store8>(out, vector, v); break; + default: e8_store8>(out, vector, v); break; + } + } + + __device__ static void put8_words( + pointer out, int64_t vector, const int* value, float chalf, float scale, int kind) { + put8(out, vector, value, chalf, scale, kind); + } + + __device__ static void put8_pair( + pointer out, int64_t vector, const int* value, float chalf, float scale, int kind) { + put8(out, vector, value, chalf, scale, kind); + } +}; + +template +__device__ __forceinline__ void e8_decode_vector_fused_body( + const uint16_t* __restrict__ payload, + const uint32_t* __restrict__ offsets, + const uint32_t* __restrict__ states, + const uint64_t* __restrict__ decode_luts, + const int* __restrict__ table_meta, + const float* __restrict__ scales, + typename Store::pointer __restrict__ output, + int store_kind, + int vecs_per_row, + int64_t num_vectors, + int tile_elements, + int num_tiles, + int* __restrict__ error) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int tile = blockIdx.x * warps_per_block + warp_in_block; + if (tile >= num_tiles) return; + const int lane = threadIdx.x & 31; + const int64_t tile_begin = static_cast(tile) * tile_elements; + const int tile_count = static_cast( + num_vectors - tile_begin < tile_elements ? num_vectors - tile_begin : tile_elements); + + const uint16_t* input_begin = payload + offsets[tile]; + const uint16_t* input = payload + offsets[tile + 1]; + uint32_t state = states[tile * kNumStates + lane]; + + const int remainder = tile_count & (kNumStates - 1); + int output_offset = tile_count - remainder; + bool first = true; + while (first || output_offset > 0) { + int valid_lanes; + if (first && remainder) { + valid_lanes = remainder; + } else { + if (first) output_offset = tile_count; + output_offset -= kNumStates; + valid_lanes = kNumStates; + } + first = false; + const bool valid = lane < valid_lanes; + int c = 0; + int value[8]; + if (valid) c = static_cast(e8_rans_decode_symbol(state, decode_luts)); + if (!rans_renormalize_checked(valid, state, input, input_begin)) set_error(error, 2); +#pragma unroll + for (int field = 0; field < 8; ++field) { + if (valid) { + const int table = 1 + 2 * field + c; + const uint64_t* lut = decode_luts + static_cast(table) * kTableSize; + value[field] = static_cast(e8_rans_decode_symbol(state, lut)) + + table_meta[table * kMetaStride + kMetaSymMin]; + } + if (!rans_renormalize_checked(valid, state, input, input_begin)) set_error(error, 2); + } + if (valid) { + int parity = 0; +#pragma unroll + for (int field = 0; field < 7; ++field) parity += value[field]; + value[7] = 2 * value[7] + (parity & 1); + const int64_t vector = tile_begin + output_offset + lane; + const int row = div_i64_i32(vector, vecs_per_row); + Store::put8_words(output, vector, value, c ? 0.5f : 0.0f, scales[row], store_kind); + } + if (output_offset == 0) break; + } + if (input != input_begin || state != kStateMin) set_error(error, 3); +} + +#define E8_FUSED_WRAPPER(NAME) \ + extern "C" __global__ void NAME( \ + const uint16_t* __restrict__ payload, \ + const uint32_t* __restrict__ offsets, \ + const uint32_t* __restrict__ states, \ + const uint64_t* __restrict__ decode_luts, \ + const int* __restrict__ table_meta, \ + const float* __restrict__ scales, \ + __nv_bfloat16* __restrict__ output, \ + int vecs_per_row, \ + int64_t num_vectors, \ + int tile_elements, \ + int num_tiles, \ + int* __restrict__ error) { \ + e8_decode_vector_fused_body( \ + payload, offsets, states, decode_luts, table_meta, scales, \ + output, 0, vecs_per_row, num_vectors, tile_elements, num_tiles, \ + error); \ + } + +#define E8_FUSED_GENERIC_WRAPPER(NAME) \ + extern "C" __global__ void NAME( \ + const uint16_t* __restrict__ payload, \ + const uint32_t* __restrict__ offsets, \ + const uint32_t* __restrict__ states, \ + const uint64_t* __restrict__ decode_luts, \ + const int* __restrict__ table_meta, \ + const float* __restrict__ scales, \ + void* __restrict__ output, \ + int vecs_per_row, \ + int64_t num_vectors, \ + int tile_elements, \ + int num_tiles, \ + int* __restrict__ error, \ + int store_kind) { \ + e8_decode_vector_fused_body( \ + payload, offsets, states, decode_luts, table_meta, scales, \ + output, store_kind, vecs_per_row, num_vectors, tile_elements, \ + num_tiles, error); \ + } + +E8_FUSED_WRAPPER(e8_decode_vector_fused_kernel) +E8_FUSED_GENERIC_WRAPPER(e8_decode_vector_fused_g_kernel) + +#undef E8_FUSED_WRAPPER +#undef E8_FUSED_GENERIC_WRAPPER + +template +__device__ __forceinline__ void e8_decode_packed32_body( + const uint16_t* __restrict__ payload, + const uint32_t* __restrict__ offsets, + const uint32_t* __restrict__ states, + const uint32_t* __restrict__ decode_luts, + const int* __restrict__ pack_bits, + uint32_t coset_frequency0, + const int* __restrict__ table_meta, + const float* __restrict__ scales, + typename Store::pointer __restrict__ output, + int store_kind, + int vecs_per_row, + int64_t num_vectors, + int tile_vectors, + int num_tiles, + int* __restrict__ error) { + __shared__ int pack_bits_shared[17]; + __shared__ int sym_min_shared[17]; + if (threadIdx.x < 17) { + pack_bits_shared[threadIdx.x] = pack_bits[threadIdx.x]; + sym_min_shared[threadIdx.x] = table_meta[threadIdx.x * kMetaStride + kMetaSymMin]; + } + __syncthreads(); + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int tile = blockIdx.x * warps_per_block + warp_in_block; + if (tile >= num_tiles) return; + const int lane = threadIdx.x & 31; + const int64_t tile_begin = static_cast(tile) * tile_vectors; + const int tile_count = static_cast( + num_vectors - tile_begin < tile_vectors ? num_vectors - tile_begin : tile_vectors); + const uint16_t* input_begin = payload + offsets[tile]; + const uint16_t* input = payload + offsets[tile + 1]; + uint32_t state = states[tile * kNumStates + lane]; + const int remainder = tile_count & (kNumStates - 1); + int output_offset = tile_count - remainder; + bool first = true; + while (first || output_offset > 0) { + int valid_lanes; + if (first && remainder) { + valid_lanes = remainder; + } else { + if (first) output_offset = tile_count; + output_offset -= kNumStates; + valid_lanes = kNumStates; + } + first = false; + const bool valid = lane < valid_lanes; + if constexpr (PREF) { + // The renorm rate grows with the coded rate, so at high coded rates the cold-payload latency matters more. + if (lane < 4) { + const char* p = reinterpret_cast(input) - 128 * (lane + 1); + if (p >= reinterpret_cast(payload)) { + asm volatile("prefetch.global.L2 [%0];" ::"l"(p)); + } + } + } + int c = 0; + int value[8]; + if (valid) c = static_cast(e8_decode_coset(state, coset_frequency0)); + if (!rans_renormalize_checked(valid, state, input, input_begin)) set_error(error, 2); +#pragma unroll + for (int field = 0; field < 8; ++field) { + if (valid) { + const int table = 1 + 2 * field + c; + value[field] = static_cast( + e8_decode_packed32(state, decode_luts, pack_bits_shared, table)) + + sym_min_shared[table]; + } + if (!rans_renormalize_checked(valid, state, input, input_begin)) set_error(error, 2); + } + if (valid) { + int parity = 0; +#pragma unroll + for (int field = 0; field < 7; ++field) parity += value[field]; + value[7] = 2 * value[7] + (parity & 1); + const int64_t vector = tile_begin + output_offset + lane; + const int c_row = div_i64_i32(vector, vecs_per_row); + Store::put8_words(output, vector, value, c ? 0.5f : 0.0f, scales[c_row], store_kind); + } + if (output_offset == 0) break; + } + if (input != input_begin || state != kStateMin) set_error(error, 3); +} + +#define E8_PACKED32_WRAPPER(NAME, PREF) \ + extern "C" __global__ void NAME( \ + const uint16_t* __restrict__ payload, \ + const uint32_t* __restrict__ offsets, \ + const uint32_t* __restrict__ states, \ + const uint32_t* __restrict__ decode_luts, \ + const int* __restrict__ pack_bits, \ + uint32_t coset_frequency0, \ + const int* __restrict__ table_meta, \ + const float* __restrict__ scales, \ + __nv_bfloat16* __restrict__ output, \ + int vecs_per_row, \ + int64_t num_vectors, \ + int tile_vectors, \ + int num_tiles, \ + int* __restrict__ error) { \ + e8_decode_packed32_body( \ + payload, offsets, states, decode_luts, pack_bits, coset_frequency0, \ + table_meta, scales, output, 0, vecs_per_row, num_vectors, \ + tile_vectors, num_tiles, error); \ + } + +#define E8_PACKED32_GENERIC_WRAPPER(NAME, PREF) \ + extern "C" __global__ void NAME( \ + const uint16_t* __restrict__ payload, \ + const uint32_t* __restrict__ offsets, \ + const uint32_t* __restrict__ states, \ + const uint32_t* __restrict__ decode_luts, \ + const int* __restrict__ pack_bits, \ + uint32_t coset_frequency0, \ + const int* __restrict__ table_meta, \ + const float* __restrict__ scales, \ + void* __restrict__ output, \ + int vecs_per_row, \ + int64_t num_vectors, \ + int tile_vectors, \ + int num_tiles, \ + int* __restrict__ error, \ + int store_kind) { \ + e8_decode_packed32_body( \ + payload, offsets, states, decode_luts, pack_bits, coset_frequency0, \ + table_meta, scales, output, store_kind, vecs_per_row, \ + num_vectors, tile_vectors, num_tiles, error); \ + } + +E8_PACKED32_WRAPPER(e8_decode_vector_packed32_kernel, false) +E8_PACKED32_WRAPPER(e8_decode_vector_packed32pf_kernel, true) +E8_PACKED32_GENERIC_WRAPPER(e8_decode_vector_packed32_g_kernel, false) +E8_PACKED32_GENERIC_WRAPPER(e8_decode_vector_packed32pf_g_kernel, true) + +#undef E8_PACKED32_WRAPPER +#undef E8_PACKED32_GENERIC_WRAPPER + +// The global-LUT path gathers a 32-lane random table through L1, which a table larger than the cache does not stay in; staging +// two compact tables in dynamic shared memory is faster while the SM still holds enough resident CTAs. +template +__device__ __forceinline__ void e8_decode_shlut_body( + const uint16_t* __restrict__ payload, + const uint32_t* __restrict__ offsets, + const uint32_t* __restrict__ states, + const uint8_t* __restrict__ sym_lut, // [17 * kTableSize] + const uint32_t* __restrict__ fb_lut, // [fb_entries] begin | freq<<16 + uint32_t coset_frequency0, + const int* __restrict__ table_meta, // [17, kMetaStride]: see kMetaFreqOff, kMetaSymMin + const float* __restrict__ scales, + typename Store::pointer __restrict__ output, + int store_kind, + int vecs_per_row, + int64_t num_vectors, + int tile_vectors, + int num_tiles, + int fb_entries, + int* __restrict__ error) { + extern __shared__ unsigned char shmem_raw[]; + uint8_t* sym_sh = reinterpret_cast(shmem_raw); + uint32_t* fb_sh = reinterpret_cast(sym_sh + 17 * kTableSize); + __shared__ int meta_sh[34]; // [0,17): freq_off, [17,34): sym_min + { + const int n4 = 17 * kTableSize / 16; // uint8 slots, staged as uint4 + const uint4* src4 = reinterpret_cast(sym_lut); + uint4* dst4 = reinterpret_cast(sym_sh); + for (int i = threadIdx.x; i < n4; i += blockDim.x) dst4[i] = src4[i]; + for (int i = threadIdx.x; i < fb_entries; i += blockDim.x) fb_sh[i] = fb_lut[i]; + if (threadIdx.x < 17) { + meta_sh[threadIdx.x] = table_meta[threadIdx.x * kMetaStride + kMetaFreqOff]; + meta_sh[threadIdx.x + 17] = table_meta[threadIdx.x * kMetaStride + kMetaSymMin]; + } + } + __syncthreads(); + + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int tile = blockIdx.x * warps_per_block + warp_in_block; + if (tile >= num_tiles) return; + const int lane = threadIdx.x & 31; + const uint32_t mask_ge = lane_mask_ge(); + const int64_t tile_begin = static_cast(tile) * tile_vectors; + const int tile_count = static_cast( + num_vectors - tile_begin < tile_vectors ? num_vectors - tile_begin : tile_vectors); + const uint16_t* input_begin = payload + offsets[tile]; + const uint16_t* input = payload + offsets[tile + 1]; + uint32_t state = states[tile * kNumStates + lane]; + const int remainder = tile_count & (kNumStates - 1); + int output_offset = tile_count - remainder; + bool first = true; + while (first || output_offset > 0) { + int valid_lanes; + if (first && remainder) { + valid_lanes = remainder; + } else { + if (first) output_offset = tile_count; + output_offset -= kNumStates; + valid_lanes = kNumStates; + } + first = false; + const bool valid = lane < valid_lanes; + if (lane < 4) { + const char* p = reinterpret_cast(input) - 128 * (lane + 1); + if (p >= reinterpret_cast(payload)) { + asm volatile("prefetch.global.L2 [%0];" ::"l"(p)); + } + } + int c = 0; + int value[8]; + if (valid) c = static_cast(e8_decode_coset(state, coset_frequency0)); + rans_renormalize_unchecked(valid, state, input, mask_ge); + const int cTS = c * kTableSize; +#pragma unroll + for (int field = 0; field < 8; ++field) { + if (valid) { + const uint32_t slot = state & kStateMask; + const uint32_t sym = static_cast( + sym_sh[(2 * field + 1) * kTableSize + cTS + slot]); + const int table = (2 * field + 1) + c; + const uint32_t fb = fb_sh[meta_sh[table] + sym]; + const uint32_t begin = fb & 0xFFFFu; + const uint32_t freq = fb >> 16; + state = freq * (state >> kProbBits) + (slot - begin); + value[field] = static_cast(sym) + meta_sh[table + 17]; + } + rans_renormalize_unchecked(valid, state, input, mask_ge); + } + if (valid) { + int parity = 0; +#pragma unroll + for (int field = 0; field < 7; ++field) parity += value[field]; + value[7] = 2 * value[7] + (parity & 1); + const int64_t vector = tile_begin + output_offset + lane; + const int c_row = div_i64_i32(vector, vecs_per_row); + const float chalf = c ? 0.5f : 0.0f; + Store::put8_pair(output, vector, value, chalf, scales[c_row], store_kind); + } + if (output_offset == 0) break; + } + if (input != input_begin || state != kStateMin) set_error(error, 3); +} + +#define E8_SHLUT_WRAPPER(NAME) \ + extern "C" __global__ void NAME( \ + const uint16_t* __restrict__ payload, \ + const uint32_t* __restrict__ offsets, \ + const uint32_t* __restrict__ states, \ + const uint8_t* __restrict__ sym_lut, \ + const uint32_t* __restrict__ fb_lut, \ + uint32_t coset_frequency0, \ + const int* __restrict__ table_meta, \ + const float* __restrict__ scales, \ + __nv_bfloat16* __restrict__ output, \ + int vecs_per_row, \ + int64_t num_vectors, \ + int tile_vectors, \ + int num_tiles, \ + int fb_entries, \ + int* __restrict__ error) { \ + e8_decode_shlut_body( \ + payload, offsets, states, sym_lut, fb_lut, coset_frequency0, table_meta, \ + scales, output, 0, vecs_per_row, num_vectors, tile_vectors, \ + num_tiles, fb_entries, error); \ + } + +#define E8_SHLUT_GENERIC_WRAPPER(NAME) \ + extern "C" __global__ void NAME( \ + const uint16_t* __restrict__ payload, \ + const uint32_t* __restrict__ offsets, \ + const uint32_t* __restrict__ states, \ + const uint8_t* __restrict__ sym_lut, \ + const uint32_t* __restrict__ fb_lut, \ + uint32_t coset_frequency0, \ + const int* __restrict__ table_meta, \ + const float* __restrict__ scales, \ + void* __restrict__ output, \ + int vecs_per_row, \ + int64_t num_vectors, \ + int tile_vectors, \ + int num_tiles, \ + int fb_entries, \ + int* __restrict__ error, \ + int store_kind) { \ + e8_decode_shlut_body( \ + payload, offsets, states, sym_lut, fb_lut, coset_frequency0, table_meta, \ + scales, output, store_kind, vecs_per_row, num_vectors, \ + tile_vectors, num_tiles, fb_entries, error); \ + } + +E8_SHLUT_WRAPPER(e8_decode_vector_shlut8pf_kernel) +E8_SHLUT_GENERIC_WRAPPER(e8_decode_vector_shlut8pf_g_kernel) + +#undef E8_SHLUT_WRAPPER +#undef E8_SHLUT_GENERIC_WRAPPER + diff --git a/entropack/schemes/lattice_rans/rans.py b/entropack/schemes/lattice_rans/rans.py new file mode 100644 index 0000000..b13de7f --- /dev/null +++ b/entropack/schemes/lattice_rans/rans.py @@ -0,0 +1,209 @@ +import numpy as np + +from ..tile_ans.eager import normalize_counts, quantized_cross_entropy +from ..tile_ans.format import NUM_STATES, STATE_MIN +from .format import BITS_PER_BYTE, NUM_COORD_FIELDS, vector_tile_elements + +__all__ = [ + "build_codec_tables", "coded_bytes", "decode_vector_stream", "encode_vector_stream", "normalize_freq", + "vector_stream_analytic_bytes", +] + + +def normalize_freq(counts: np.ndarray, table_size: int) -> np.ndarray: + counts = np.asarray(counts, dtype=np.int64) + if counts.size == 0 or int(counts.sum()) == 0: + raise ValueError("cannot build an rANS table from an empty symbol stream") + if counts.size > table_size: + raise ValueError( + f"E8 coordinate alphabet {counts.size} exceeds rANS table_size {table_size}; " + "raise prob_bits or coarsen the lattice scale" + ) + return normalize_counts(counts, table_size) + + +def build_codec_tables(freq: np.ndarray, probability_bits: int): + freq = np.asarray(freq, dtype=np.int64) + table_size = 1 << probability_bits + if freq.size > table_size: + raise ValueError(f"alphabet {freq.size} exceeds table_size {table_size}") + if int(freq.sum()) != table_size: + raise ValueError("normalized frequencies must sum to table_size") + cdf = np.zeros(freq.size, dtype=np.int64) + cdf[1:] = np.cumsum(freq)[:-1] + lut = np.zeros(table_size, dtype=np.uint64) + running = 0 + for symbol in range(freq.size): + value = int(freq[symbol]) + if value == 0: + continue + packed_freq = 0 if value == table_size else value + entry = (np.uint64(running) << np.uint64(32)) | (np.uint64(packed_freq) << np.uint64(16)) | np.uint64(symbol) + lut[running : running + value] = entry + running += value + if running != table_size: + raise AssertionError(f"frequency sum is {running}, expected {table_size}") + return cdf, lut + + +def _decode_group(states, luts, table_ids, words, pointer, output, valid_lanes, probability_bits): + table_size = 1 << probability_bits + reads = [] + for lane in range(valid_lanes): + state = int(states[lane]) + slot = state & (table_size - 1) + entry = int(luts[0 if table_ids is None else int(table_ids[lane])][slot]) + symbol = entry & 0xFFFF + frequency = (entry >> 16) & 0xFFFF + if frequency == 0: + frequency = table_size + cdf = entry >> 32 + output[lane] = symbol + state = frequency * (state >> probability_bits) + (slot - cdf) + states[lane] = state + if state < STATE_MIN: + reads.append(lane) + first_word = pointer - len(reads) + if first_word < 0: + raise ValueError("lattice_rans rANS payload is truncated") + for index, lane in enumerate(reads): + states[lane] = (int(states[lane]) << 16) | int(words[first_word + index]) + return first_word + + +def _decode_tile_group(luts, states, words, pointer, coset_out, field_out, valid_lanes, probability_bits): + pointer = _decode_group(states, luts, None, words, pointer, coset_out, valid_lanes, probability_bits) + for field in range(NUM_COORD_FIELDS): + pointer = _decode_group( + states, [luts[1 + 2 * field], luts[2 + 2 * field]], coset_out, words, pointer, field_out[:, field], valid_lanes, + probability_bits, + ) + return pointer + + +def encode_vector_stream( + cosets: np.ndarray, field_symbols: np.ndarray, frequencies: list[np.ndarray], probability_bits: int, tile_vectors: int, +): + cosets = np.asarray(cosets, dtype=np.int64) + field_symbols = np.asarray(field_symbols, dtype=np.int64) + if field_symbols.shape != (cosets.size, NUM_COORD_FIELDS): + raise ValueError(f"field_symbols must have shape [num_vectors, {NUM_COORD_FIELDS}]") + cdfs = [] + for frequency in frequencies: + frequency = np.asarray(frequency, dtype=np.int64) + if frequency.size: + cdfs.append(build_codec_tables(frequency, probability_bits)[0]) + else: + cdfs.append(np.empty(0, dtype=np.int64)) + table_size = 1 << probability_bits + state_check_shift = 31 - probability_bits + num_tiles = max(1, (cosets.size + tile_vectors - 1) // tile_vectors) + states = np.empty((num_tiles, NUM_STATES), dtype=np.uint32) + parts = [] + offsets = np.zeros(num_tiles + 1, dtype=np.int64) + total = 0 + for tile in range(num_tiles): + begin = tile * tile_vectors + end = min(begin + tile_vectors, cosets.size) + tile_states = np.full(NUM_STATES, STATE_MIN, dtype=np.uint32) + words = [] + for base in range(begin, end, NUM_STATES): + limit = min(NUM_STATES, end - base) + for field in range(7, -1, -1): + for lane in range(limit): + index = base + lane + table = 1 + 2 * field + int(cosets[index]) + symbol = int(field_symbols[index, field]) + frequency = int(frequencies[table][symbol]) + state = int(tile_states[lane]) + if state >= (frequency << state_check_shift): + words.append(state & 0xFFFF) + state >>= 16 + tile_states[lane] = (state // frequency) * table_size + (state % frequency) + int(cdfs[table][symbol]) + for lane in range(limit): + symbol = int(cosets[base + lane]) + frequency = int(frequencies[0][symbol]) + state = int(tile_states[lane]) + if state >= (frequency << state_check_shift): + words.append(state & 0xFFFF) + state >>= 16 + tile_states[lane] = (state // frequency) * table_size + (state % frequency) + int(cdfs[0][symbol]) + tile_words = np.asarray(words, dtype=np.uint16) + states[tile] = tile_states + parts.append(tile_words) + total += tile_words.size + offsets[tile + 1] = total + payload = np.concatenate(parts) if parts else np.empty(0, dtype=np.uint16) + return payload, states, offsets + + +def decode_vector_stream( + words: np.ndarray, offsets: np.ndarray, states: np.ndarray, frequencies: list[np.ndarray], probability_bits: int, + tile_vectors: int, num_vectors: int, +): + luts = [ + np.empty(0, dtype=np.uint64) if np.asarray(frequency).size == 0 + else build_codec_tables(np.asarray(frequency, dtype=np.int64), probability_bits)[1] for frequency in frequencies + ] + cosets = np.empty(num_vectors, dtype=np.int64) + fields = np.empty((num_vectors, NUM_COORD_FIELDS), dtype=np.int64) + num_tiles = max(1, (num_vectors + tile_vectors - 1) // tile_vectors) + for tile in range(num_tiles): + begin = tile * tile_vectors + tile_count = min(tile_vectors, num_vectors - begin) + tile_words = np.asarray(words[offsets[tile] : offsets[tile + 1]], dtype=np.uint16) + tile_states = np.asarray(states[tile], dtype=np.uint32).copy() + pointer = tile_words.size + remainder = tile_count % NUM_STATES + offset = tile_count - remainder + if remainder: + pointer = _decode_tile_group( + luts, tile_states, tile_words, pointer, cosets[begin + offset : begin + tile_count], + fields[begin + offset : begin + tile_count], remainder, probability_bits, + ) + while offset > 0: + offset -= NUM_STATES + pointer = _decode_tile_group( + luts, tile_states, tile_words, pointer, cosets[begin + offset : begin + offset + NUM_STATES], + fields[begin + offset : begin + offset + NUM_STATES], NUM_STATES, probability_bits, + ) + if pointer != 0: + raise ValueError("lattice_rans vector rANS stream contains unread payload words") + return cosets, fields + + +def coded_bytes(counts, sizes, prob_bits: int, tile_elements: int) -> float: + table_size = 1 << prob_bits + counts_by_table = [] + frequencies = [] + overflow_bits = 0.0 + empty_counts = np.empty(0, dtype=np.int64) + empty_freqs = np.empty(0, dtype=np.uint16) + for n_symbols, count in zip(sizes, counts, strict=True): + if n_symbols and count.size <= table_size: + counts_by_table.append(count) + frequencies.append(normalize_freq(count, table_size)) + else: + overflow_bits += n_symbols * prob_bits + counts_by_table.append(empty_counts) + frequencies.append(empty_freqs) + total = vector_stream_analytic_bytes( + counts_by_table, frequencies, prob_bits, vector_tile_elements(tile_elements), sizes[0] + ) + return total + overflow_bits / BITS_PER_BYTE + + +def vector_stream_analytic_bytes( + counts_by_table: list[np.ndarray], frequencies: list[np.ndarray], probability_bits: int, tile_vectors: int, + num_vectors: int, +) -> float: + bits = 0.0 + frequency_bytes = 0 + for counts, frequency in zip(counts_by_table, frequencies, strict=True): + if counts.size: + bits += quantized_cross_entropy(counts.astype(np.int64), frequency.astype(np.int64), probability_bits) + frequency_bytes += frequency.size * 2 + num_tiles = max(1, (num_vectors + tile_vectors - 1) // tile_vectors) + state_residual_bits = num_tiles * NUM_STATES * 16 + payload_words = int(np.ceil(max(0.0, bits - state_residual_bits) / 16.0)) + return payload_words * 2 + num_tiles * NUM_STATES * 4 + (num_tiles + 1) * 4 + frequency_bytes diff --git a/entropack/schemes/lattice_rans/rdo.py b/entropack/schemes/lattice_rans/rdo.py new file mode 100644 index 0000000..b31463a --- /dev/null +++ b/entropack/schemes/lattice_rans/rdo.py @@ -0,0 +1,162 @@ +from dataclasses import replace +from typing import NamedTuple + +import numpy as np +import torch + +from .format import BITS_PER_BYTE, LATTICE_DIM, NUM_COORD_FIELDS + +RATIO_SPAN = (0.70, 1.45) + + +def ratio_ladder(count): + """The scale ratios one refinement pass prices: 1.0, then the geometric midpoint of the widest + log-gap inside RATIO_SPAN, one per further candidate. A K-point ladder is therefore a subset of + every wider one, so widening can only help a row.""" + pts = [1.0] + while len(pts) < count: + bounds = [RATIO_SPAN[0], *pts, RATIO_SPAN[1]] + i = max(range(1, len(bounds)), key=lambda j: bounds[j] / bounds[j - 1]) + pts.insert(i - 1, (bounds[i] * bounds[i - 1]) ** 0.5) + return tuple(pts) + + +class RateTable(NamedTuple): + """Both index tensors are int32, which keeps the per-iteration index temporaries :func:`row_rates` builds small against the + table they index. + """ + + values: torch.Tensor + offsets: torch.Tensor + minima: torch.Tensor + + +def stream_ranges(candidates, n_streams): + minima = [0] * n_streams + maxima = [0] * n_streams + for stream in range(n_streams): + alive = [ + (candidate.sym_min[stream], candidate.sym_min[stream] + candidate.alphabets[stream] - 1) for candidate in candidates + if candidate.alphabets[stream] > 0 + ] + if alive: + minima[stream] = min(item[0] for item in alive) + maxima[stream] = max(item[1] for item in alive) + return minima, maxima + + +def rate_costs(counts, sizes, sym_min, minima, maxima, device, alpha=0.5): + """The streams are concatenated on the host and uploaded as one table; uploading them separately costs one host-to-device + copy per stream and dominates this function. The arithmetic itself is small and runs slower on the device than here. + """ + tables = [] + offsets = [0] + for stream, count in enumerate(counts): + width = maxima[stream] - minima[stream] + 1 + expanded = np.zeros(width, dtype=np.float64) + if count.size: + offset = sym_min[stream] - minima[stream] + expanded[offset : offset + count.size] = count + probability = (expanded + alpha) / (sizes[stream] + alpha * width) + tables.append(-np.log2(probability).astype(np.float32)) + offsets.append(offsets[-1] + width) + return RateTable( + torch.from_numpy(np.concatenate(tables)).to(device), torch.tensor(offsets, dtype=torch.int32, device=device), + torch.tensor(minima, dtype=torch.int32, device=device), + ) + + +def row_rates(rows, cols, candidate, table): + """Gather rather than mask: a boolean mask per stream needs ``nonzero`` to bring the hit count back to the host, and those + per-stream synchronizations leave small layers host-bound. Transposing the field matrix first turns the coordinate reads + from strided into contiguous ones, which matters more the less of the matrix the device cache holds. The accumulation + order is untouched, so the sum is bit-identical. + """ + vectors_per_row = cols // LATTICE_DIM + fields = candidate.fields.reshape(-1, LATTICE_DIM).t().contiguous().to(torch.int32) + coset = candidate.c_arr.to(torch.int32) + rate = table.values[table.offsets[0] + coset] + for field in range(NUM_COORD_FIELDS): + stream = 1 + 2 * field + coset + rate = rate + table.values[table.offsets[stream] + (fields[field] - table.minima[stream])] + return rate.reshape(rows, vectors_per_row).sum(1) + + +def compact_candidate(candidate): + live = [index for index in range(1, len(candidate.sym_min)) if candidate.alphabets[index] > 0] + lows = [candidate.sym_min[index] for index in live] + highs = [candidate.sym_min[index] + candidate.alphabets[index] - 1 for index in live] + low = min(lows, default=0) + high = max(highs, default=0) + if low >= -128 and high <= 127: + field_dtype = torch.int8 + elif low >= -32768 and high <= 32767: + field_dtype = torch.int16 + else: + field_dtype = torch.int32 + return replace(candidate, fields=candidate.fields.to(field_dtype), c_arr=candidate.c_arr.to(torch.uint8)) + + +def select_candidate_rows(rows, cols, candidates, choice): + vectors_per_row = cols // LATTICE_DIM + device = candidates[0].fields.device + fields = torch.empty(candidates[0].fields.numel(), dtype=torch.int32, device=device) + c_arr = torch.empty(candidates[0].c_arr.numel(), dtype=torch.int32, device=device) + fields_rows = fields.reshape(rows, vectors_per_row, LATTICE_DIM) + c_rows = c_arr.reshape(rows, vectors_per_row) + for index, candidate in enumerate(candidates): + selected_rows = torch.nonzero(choice == index, as_tuple=False).flatten() + if selected_rows.numel() == 0: + continue + source = candidate.fields.reshape(rows, vectors_per_row, LATTICE_DIM)[selected_rows] + fields_rows[selected_rows] = source.to(torch.int32) + c_rows[selected_rows] = candidate.c_arr.reshape(rows, vectors_per_row)[selected_rows].to(torch.int32) + scales = torch.stack([candidate.scales for candidate in candidates], dim=1) + row = torch.arange(rows, device=choice.device) + return fields, c_arr, scales[row, choice] + + +def choose_rate_tradeoff(distortion, rates, row, desired_rate): + zero_choice = distortion.argmin(1) + if rates[row, zero_choice].sum() <= desired_rate: + return zero_choice + choice = zero_choice + low = 0.0 + high = max(float(distortion.mean().item()) * 1.0e-4, 1.0e-12) + for _ in range(50): + trial = (distortion + high * rates).argmin(1) + if rates[row, trial].sum() <= desired_rate: + break + high *= 2.0 + for _ in range(20): + middle = high * 0.5 if low == 0.0 else (low * high) ** 0.5 + trial = (distortion + middle * rates).argmin(1) + if rates[row, trial].sum() > desired_rate: + low = middle + else: + high = middle + choice = trial + return choice + + +def optimize_rows(candidates, *, rows, cols, baseline_index, iterations, device, summarize, total_bytes): + minima, maxima = stream_ranges(candidates, len(candidates[0].counts)) + distortion = torch.stack([candidate.row_sse for candidate in candidates], dim=1) + choice = torch.full((rows,), baseline_index, dtype=torch.int64, device=device) + row = torch.arange(rows, device=device) + fields, c_arr, fitted_scales = select_candidate_rows(rows, cols, candidates, choice) + summary = summarize(fields, c_arr) + baseline = candidates[baseline_index] + target_bytes = total_bytes(baseline.counts, baseline.sizes) + + for _ in range(iterations): + table = rate_costs(summary.counts, summary.sizes, summary.sym_min, minima, maxima, device) + rates = torch.stack([row_rates(rows, cols, candidate, table) for candidate in candidates], dim=1) + current_rate = rates[row, choice].sum() + current_bytes = total_bytes(summary.counts, summary.sizes) + desired_rate = current_rate + BITS_PER_BYTE * (target_bytes - current_bytes) + choice = choose_rate_tradeoff(distortion, rates, row, desired_rate) + fields, c_arr, fitted_scales = select_candidate_rows(rows, cols, candidates, choice) + summary = summarize(fields, c_arr) + + return summary, fitted_scales diff --git a/entropack/schemes/tile_ans/__init__.py b/entropack/schemes/tile_ans/__init__.py new file mode 100644 index 0000000..a50af43 --- /dev/null +++ b/entropack/schemes/tile_ans/__init__.py @@ -0,0 +1,55 @@ +import torch + +from ..base import Scheme, packed_buffers, register_scheme +from ..checks import prepare_weight +from .format import OPTIONS_BY_DTYPE, PACKED_KEYS, TileBuffers, TileANSConfig, validate_packed + + +class TileANSScheme(Scheme): + name = "tile_ans" + buffer_names = PACKED_KEYS + priority = 100 + dtypes = tuple(OPTIONS_BY_DTYPE) + lanes = {"eager": "eager", "cuda": "cuda"} + + def options_for(self, dtype): + tile_elements, probability_bits = OPTIONS_BY_DTYPE[dtype] + return { + "tile_elements": tile_elements, "probability_bits": probability_bits, + "raw_lane_threshold": TileANSConfig().raw_lane_threshold, + } + + def encode(self, weight: torch.Tensor, config: TileANSConfig) -> dict: + weight = prepare_weight(weight, scheme=self.name) + tile_elements = config.tile_elements + if tile_elements == 0: + storage_bytes = weight.numel() * weight.element_size() + tile_elements = 4096 if storage_bytes <= 32 * 1024 * 1024 else 8192 + device = weight.device + lane, run_on = self.lane_for(weight, config.execution_backend) + if run_on is not None and device != run_on: + weight = weight.to(run_on) + buffers = lane.encode( + weight=weight, tile_elements=tile_elements, probability_bits=config.probability_bits, + raw_lane_threshold=config.raw_lane_threshold, threads_per_block=config.threads_per_block, + ) + return {key: value.to(device) for key, value in buffers._asdict().items()} + + def validate_buffers(self, buffers, shape, dtype): + validate_packed(buffers, shape, dtype) + + def decode(self, packed: dict, *, shape: tuple[int, ...], dtype: torch.dtype, + config: TileANSConfig) -> torch.Tensor: + buffers = packed_buffers(packed, TileBuffers) + source = buffers.layout.device + lane, lane_device = self.lane_for(buffers.layout, config.execution_backend, gate_dtype=False) + if lane_device is not None and source != lane_device: + buffers = TileBuffers._make(value.to(lane_device) for value in buffers) + flat = lane.decode(buffers, dtype=dtype, threads_per_block=config.threads_per_block) + out = flat.reshape(shape) + return out if out.device == source else out.to(source) + + +register_scheme(TileANSScheme()) + +__all__ = ["TileANSConfig", "TileANSScheme"] diff --git a/entropack/schemes/tile_ans/cuda.py b/entropack/schemes/tile_ans/cuda.py new file mode 100644 index 0000000..25d42dd --- /dev/null +++ b/entropack/schemes/tile_ans/cuda.py @@ -0,0 +1,170 @@ +from pathlib import Path + +import cupy +import numpy as np +import torch + +from ...backends.cuda import device as _device_caps +from ...backends.cuda.kernels import KernelLibrary +from ...backends.cuda.kernels import device_index as _device_index +from ...backends.cuda.kernels import external_stream as _external_stream +from ...backends.cuda.kernels import pointer as _pointer +from .eager import build_tables_from_counts, lane_modes_from_counts, select_coding_options +from .format import ( + AUTO_PROB_BITS, BLOCK_SIZE, ENCODE_TABLE_SHARED_BYTES, HISTOGRAM_MAX_WARPS, HISTOGRAM_MIN_WARPS, + HISTOGRAM_WARP_BUDGET, LANE_ANS, LANE_RAW, NUM_STATES, TileBuffers, make_layout, num_streams, + num_tiles, parse_layout_cached, +) + +_CUDA_PATH = Path(__file__).parent / "tile_ans.cu" +_KERNEL_NAMES = ( + "tile_ans_histogram_kernel", "tile_ans_encode_kernel", "tile_ans_compact_kernel", "tile_ans_decode_raw0_ans1_kernel", + "tile_ans_decode_raw3_ans1_kernel", "tile_ans_decode_all_raw_kernel", "tile_ans_decode_kernel", +) +_LIBRARY = KernelLibrary( + key="tile_ans", source=_CUDA_PATH, defines=lambda _device, probability_bits: (f"TILE_ANS_PROB_BITS={probability_bits}",), + includes=(_CUDA_PATH.parent,), kernel_names=_KERNEL_NAMES, +) +_kernel = _LIBRARY.kernel + + +def _lane_histogram(contiguous, caps, device_index, torch_stream, probability_bits): + num_elements = contiguous.numel() + num_lanes = contiguous.element_size() + histograms = torch.zeros((num_lanes, 256), dtype=torch.int64, device=contiguous.device) + histogram_warps = min( + HISTOGRAM_MAX_WARPS, max(HISTOGRAM_MIN_WARPS, HISTOGRAM_WARP_BUDGET // num_lanes), + ) + histogram_threads = histogram_warps * caps.warp_size + histogram_blocks = caps.grid(-(-num_elements // histogram_threads), histogram_threads) + histogram_shared = histogram_warps * num_lanes * 256 * 4 + with cupy.cuda.Device(device_index), _external_stream(torch_stream): + _kernel(device_index, probability_bits, "tile_ans_histogram_kernel")( + (histogram_blocks,), (histogram_threads,), + (_pointer(contiguous), _pointer(histograms), np.int64(num_elements), np.int32(num_lanes)), + shared_mem=histogram_shared, + ) + return histograms.cpu().numpy() + + +def _lane_tables(counts, probability_bits, raw_lane_threshold, tile_elements, device): + if probability_bits == AUTO_PROB_BITS: + probability_bits, frequencies, cdfs, decode_tables, lane_modes = select_coding_options( + counts, raw_lane_threshold, tile_elements + ) + else: + frequencies, cdfs, decode_tables = build_tables_from_counts(counts, probability_bits) + lane_modes = lane_modes_from_counts(counts, raw_lane_threshold, tile_elements) + return ( + probability_bits, int(np.count_nonzero(lane_modes == LANE_ANS)), torch.from_numpy(frequencies).to(device), + torch.from_numpy(cdfs).to(device), torch.from_numpy(decode_tables).to(device), torch.from_numpy(lane_modes).to(device), + ) + + +def encode( + *, weight: torch.Tensor, tile_elements: int, probability_bits: int, raw_lane_threshold: float, + threads_per_block: int | None, +) -> TileBuffers: + contiguous = weight.contiguous() + num_elements = contiguous.numel() + num_lanes = contiguous.element_size() + streams = num_streams(num_elements, num_lanes, tile_elements) + device = contiguous.device + device_index = _device_index(contiguous) + torch_stream = torch.cuda.current_stream(device) + caps = _device_caps.caps(device) + + counts = _lane_histogram(contiguous, caps, device_index, torch_stream, probability_bits) + ( + probability_bits, num_ans_lanes, frequencies_gpu, cdfs_gpu, decode_tables_gpu, lane_modes_gpu, + ) = _lane_tables(counts, probability_bits, raw_lane_threshold, tile_elements, device) + + tiles = num_tiles(num_elements, tile_elements) + states = torch.empty((tiles * num_ans_lanes, NUM_STATES), dtype=torch.uint32, device=device) + word_counts = torch.empty(streams, dtype=torch.uint32, device=device) + scratch = torch.empty(streams * tile_elements, dtype=torch.uint16, device=device) + threads = _device_caps.resolve_threads(caps, threads_per_block, BLOCK_SIZE) + warps = threads // caps.warp_size + blocks = -(-tiles // warps) + args = ( + _pointer(contiguous), _pointer(frequencies_gpu), _pointer(cdfs_gpu), _pointer(lane_modes_gpu), _pointer(scratch), + _pointer(word_counts), _pointer(states), np.int64(num_elements), np.int32(tile_elements), np.int32(num_lanes), + np.int32(tiles), + ) + + with cupy.cuda.Device(device_index), _external_stream(torch_stream): + _kernel(device_index, probability_bits, "tile_ans_encode_kernel")( + (blocks, num_lanes), (threads,), args, shared_mem=ENCODE_TABLE_SHARED_BYTES, + ) + counts_cp = cupy.from_dlpack(word_counts) + offsets64 = torch.empty(streams + 1, dtype=torch.int64, device=device) + offsets_cp = cupy.from_dlpack(offsets64) + offsets_cp[0] = 0 + cupy.cumsum(counts_cp, dtype=cupy.int64, out=offsets_cp[1:]) + total_words = int(offsets64[-1].item()) + if total_words >= 1 << 32: + raise ValueError("tile_ans payload exceeds uint32 offset capacity") + + offsets = offsets64.to(torch.uint32) + payload = torch.empty(total_words, dtype=torch.uint16, device=device) + with cupy.cuda.Device(device_index), _external_stream(torch_stream): + _kernel(device_index, probability_bits, "tile_ans_compact_kernel")( + (streams,), (threads,), + (_pointer(scratch), _pointer(offsets), _pointer(payload), np.int32(tile_elements), np.int32(streams)), + ) + layout = make_layout(num_elements, num_lanes, tile_elements, probability_bits).to(device) + return TileBuffers( + payload=payload, offsets=offsets, states=states, decode_tables=decode_tables_gpu, lane_modes=lane_modes_gpu, + layout=layout, + ) + + +def decode(buffers: TileBuffers, *, dtype: torch.dtype, threads_per_block: int | None) -> torch.Tensor: + payload, offsets, states = buffers.payload, buffers.offsets, buffers.states + decode_tables, lane_modes, layout = buffers.decode_tables, buffers.lane_modes, buffers.layout + tile_elements, probability_bits, num_lanes, num_elements = parse_layout_cached(layout) + output = torch.empty(num_elements * num_lanes, dtype=torch.uint8, device=payload.device) + device_index = _device_index(payload) + torch_stream = torch.cuda.current_stream(payload.device) + caps = _device_caps.caps(payload.device) + + tiles = num_tiles(num_elements, tile_elements) + lane_mode_values = getattr(layout, "_tile_ans_lane_modes", None) + if lane_mode_values is None: + lane_mode_values = tuple(int(value) for value in lane_modes.detach().cpu().tolist()) + layout._tile_ans_lane_modes = lane_mode_values + + all_raw_writer = all(mode == LANE_RAW for mode in lane_mode_values) + paired_writer = num_lanes == 2 and lane_mode_values == (LANE_RAW, LANE_ANS) + quad_writer = num_lanes == 4 and lane_mode_values == (LANE_RAW, LANE_RAW, LANE_RAW, LANE_ANS) + table_shared = (1 << probability_bits) * 4 + grid_lanes = 1 + block_owns_tile = False + if all_raw_writer: + kernel_name, shared_bytes = "tile_ans_decode_all_raw_kernel", 0 + block_owns_tile = True + elif paired_writer: + # bf16 splits into a raw low byte and an ANS-coded high byte, so the coded lane's table is small enough to stage in + # shared memory: renormalization then drops its bounds check and the lane mask is hoisted out of the symbol loop. + kernel_name = "tile_ans_decode_raw0_ans1_kernel" + shared_bytes = table_shared + elif quad_writer: + kernel_name, shared_bytes = "tile_ans_decode_raw3_ans1_kernel", 0 + else: + kernel_name, shared_bytes = "tile_ans_decode_kernel", table_shared + grid_lanes = num_lanes + + kernel = _kernel(device_index, probability_bits, kernel_name) + args = ( + _pointer(payload), _pointer(offsets), _pointer(states), _pointer(decode_tables), _pointer(lane_modes), _pointer(output), + np.int64(num_elements), np.int32(tile_elements), np.int32(num_lanes), np.int32(tiles), + ) + + def launch_decode(threads: int) -> None: + blocks = tiles if block_owns_tile else -(-tiles // (threads // caps.warp_size)) + grid = (blocks, grid_lanes) if grid_lanes > 1 else (blocks,) + with cupy.cuda.Device(device_index), _external_stream(torch_stream): + kernel(grid, (threads,), args, shared_mem=shared_bytes) + + launch_decode(_device_caps.resolve_threads(caps, threads_per_block, BLOCK_SIZE)) + return output.view(dtype) diff --git a/entropack/schemes/tile_ans/device.cuh b/entropack/schemes/tile_ans/device.cuh new file mode 100644 index 0000000..eb267b6 --- /dev/null +++ b/entropack/schemes/tile_ans/device.cuh @@ -0,0 +1,90 @@ +#pragma once + +#include + +#ifndef TILE_ANS_PROB_BITS +#define TILE_ANS_PROB_BITS 12 +#endif + +constexpr int kProbBits = TILE_ANS_PROB_BITS; +constexpr uint32_t kTableSize = 1u << kProbBits; +constexpr uint32_t kStateMask = kTableSize - 1u; +constexpr uint32_t kStateMin = 1u << 15; +constexpr int kNumStates = 32; + +__device__ __forceinline__ uint32_t lane_mask_ge() { + uint32_t mask; + asm("mov.u32 %0, %%lanemask_ge;" : "=r"(mask)); + return mask; +} + +__device__ __forceinline__ uint32_t lane_mask_lt() { + uint32_t mask; + asm("mov.u32 %0, %%lanemask_lt;" : "=r"(mask)); + return mask; +} + +__device__ __forceinline__ uint32_t rans_decode_symbol( + uint32_t& state, + const uint32_t* __restrict__ table) { + const uint32_t slot = state & kStateMask; + const uint32_t entry = table[slot]; + const uint32_t symbol = entry & 0xFFu; + uint32_t frequency = (entry >> 8) & 0xFFFu; + if (frequency == 0) { + frequency = kTableSize; + } + const uint32_t cdf = entry >> 20; + state = frequency * (state >> kProbBits) + (slot - cdf); + return symbol; +} + +__device__ __forceinline__ bool rans_renormalize_checked( + bool valid, + uint32_t& state, + const uint16_t*& input, + const uint16_t* __restrict__ input_begin) { + const bool read = valid && state < kStateMin; + const uint32_t vote = __ballot_sync(0xFFFFFFFFu, read); + const uint32_t prefix = __popc(vote & lane_mask_ge()); + bool valid_read = true; + if (read) { + const uint16_t* address = input - prefix; + valid_read = address >= input_begin; + const uint32_t word = valid_read ? *address : 0u; + state = (state << 16) | word; + } + input -= __popc(vote); + return valid_read; +} + +__device__ __forceinline__ void rans_renormalize_unchecked( + bool valid, + uint32_t& state, + const uint16_t*& input, + uint32_t mask_ge) { + const bool read = valid && state < kStateMin; + const uint32_t vote = __ballot_sync(0xFFFFFFFFu, read); + if (read) { + const uint32_t prefix = __popc(vote & mask_ge); + state = (state << 16) | input[-static_cast(prefix)]; + } + input -= __popc(vote); +} + +__device__ __forceinline__ void decode_group( + bool valid, + uint32_t& state, + const uint32_t* __restrict__ table, + const uint16_t*& input, + const uint16_t* __restrict__ input_begin, + uint8_t* __restrict__ output, + int64_t output_index, + int64_t num_lanes, + int byte_lane) { + if (valid) { + output[output_index * num_lanes + byte_lane] = + static_cast(rans_decode_symbol(state, table)); + } + (void)rans_renormalize_checked(valid, state, input, input_begin); +} diff --git a/entropack/schemes/tile_ans/eager.py b/entropack/schemes/tile_ans/eager.py new file mode 100644 index 0000000..98e700b --- /dev/null +++ b/entropack/schemes/tile_ans/eager.py @@ -0,0 +1,307 @@ +import heapq + +import numpy as np +import torch + +from .format import ( + AUTO_PROB_BITS, LANE_ANS, LANE_RAW, NUM_STATES, STATE_MIN, SUPPORTED_PROB_BITS, TileBuffers, + make_layout, num_tiles, parse_layout_cached, +) + +_STATE_BITS = 31 +_RENORM_BITS = 16 +_AUTO_RATIO_TOLERANCE = 0.0015 + + +def normalize_counts(counts: np.ndarray, table_size: int) -> np.ndarray: + counts = np.asarray(counts, dtype=np.int64) + present = np.flatnonzero(counts) + if present.size == 0: + raise ValueError("cannot build an rANS codebook from empty input") + + target = counts.astype(np.float64) * (table_size / int(counts.sum())) + frequencies = np.floor(target).astype(np.int64) + frequencies[present] = np.maximum(frequencies[present], 1) + + difference = table_size - int(frequencies.sum()) + if difference > 0: + residual = target - frequencies + queue = [(-float(residual[symbol]), int(symbol)) for symbol in present] + heapq.heapify(queue) + for _ in range(difference): + negative_residual, symbol = heapq.heappop(queue) + frequencies[symbol] += 1 + heapq.heappush(queue, (negative_residual + 1.0, symbol)) + elif difference < 0: + residual = target - frequencies + queue = [(float(residual[symbol]), int(symbol)) for symbol in present if frequencies[symbol] > 1] + if not queue: + raise ValueError("unable to normalize rANS frequencies") + heapq.heapify(queue) + for step in range(-difference): + symbol_residual, symbol = heapq.heappop(queue) + frequencies[symbol] -= 1 + if frequencies[symbol] > 1: + heapq.heappush(queue, (symbol_residual + 1.0, symbol)) + elif not queue and step + 1 < -difference: + raise ValueError("unable to normalize rANS frequencies") + return frequencies.astype(np.uint16) + + +def _build_tables_from_frequencies(frequencies: np.ndarray, probability_bits: int): + frequencies = np.asarray(frequencies, dtype=np.uint16) + table_size = 1 << probability_bits + num_lanes = frequencies.shape[0] + cdfs = np.zeros((num_lanes, 256), dtype=np.uint16) + decode_tables = np.empty((num_lanes, table_size), dtype=np.uint32) + + for lane, freq in enumerate(frequencies): + cdf = np.zeros(256, dtype=np.uint16) + running = 0 + for symbol in range(256): + cdf[symbol] = running + value = int(freq[symbol]) + if value: + packed_frequency = 0 if value == 4096 else value + decode_tables[lane, running : running + value] = (running << 20) | (packed_frequency << 8) | symbol + running += value + if running != table_size: + raise AssertionError(f"normalized frequency sum is {running}, expected {table_size}") + cdfs[lane] = cdf + return frequencies, cdfs, decode_tables + + +def build_tables_from_counts(counts_by_lane: np.ndarray, probability_bits: int): + counts_by_lane = np.asarray(counts_by_lane, dtype=np.int64) + table_size = 1 << probability_bits + frequencies = np.stack([normalize_counts(counts, table_size) for counts in counts_by_lane]) + return _build_tables_from_frequencies(frequencies, probability_bits) + + +def quantized_cross_entropy(counts: np.ndarray, frequencies: np.ndarray, probability_bits: int) -> float: + present = counts > 0 + return float( + np.sum(counts[present].astype(np.float64) * (probability_bits - np.log2(frequencies[present].astype(np.float64)))) + ) + + +def lane_modes_from_counts( + counts_by_lane: np.ndarray, raw_lane_threshold: float, tile_elements: int, frequencies: np.ndarray | None = None, + probability_bits: int | None = None, +) -> np.ndarray: + if not 0.0 <= raw_lane_threshold <= 8.0: + raise ValueError("raw_lane_threshold must be in [0, 8]") + effective_threshold = min(raw_lane_threshold, 8.0 - (NUM_STATES * 32) / tile_elements) + counts_by_lane = np.asarray(counts_by_lane, dtype=np.int64) + modes = np.empty(counts_by_lane.shape[0], dtype=np.uint8) + for lane, counts in enumerate(counts_by_lane): + if frequencies is None: + probabilities = counts[counts > 0].astype(np.float64) / int(counts.sum()) + bits_per_symbol = -np.sum(probabilities * np.log2(probabilities)) + else: + if probability_bits is None: + raise ValueError("probability_bits is required with normalized frequencies") + bits_per_symbol = quantized_cross_entropy(counts, frequencies[lane], probability_bits) / int(counts.sum()) + modes[lane] = LANE_RAW if bits_per_symbol >= effective_threshold else LANE_ANS + return modes + + +def select_coding_options(counts_by_lane: np.ndarray, raw_lane_threshold: float, tile_elements: int): + counts_by_lane = np.asarray(counts_by_lane, dtype=np.int64) + num_elements = int(counts_by_lane[0].sum()) + num_lanes = counts_by_lane.shape[0] + tiles = num_tiles(num_elements, tile_elements) + full_tiles, tail = divmod(num_elements, tile_elements) + raw_words = full_tiles * ((tile_elements + 1) // 2) + (tail + 1) // 2 + original_bytes = num_elements * num_lanes + # Quantizing a table can only add cost, so a lane above the threshold stays above it at every precision; when all of them + # are, only the coarsest table needs pricing. + plain = lane_modes_from_counts(counts_by_lane, raw_lane_threshold, tile_elements) + probability_variants = SUPPORTED_PROB_BITS[:1] if np.all(plain == LANE_RAW) else SUPPORTED_PROB_BITS + + candidates = [] + for probability_bits in probability_variants: + table_size = 1 << probability_bits + frequencies = np.stack([normalize_counts(counts, table_size) for counts in counts_by_lane]) + modes = lane_modes_from_counts(counts_by_lane, raw_lane_threshold, tile_elements, frequencies, probability_bits) + estimated_bytes = (tiles * num_lanes + 1) * 4 + num_lanes * table_size * 4 + num_lanes + 5 * 8 + for lane, mode in enumerate(modes): + if mode == LANE_RAW: + estimated_bytes += raw_words * 2 + else: + cross_entropy_bits = quantized_cross_entropy(counts_by_lane[lane], frequencies[lane], probability_bits) + estimated_bytes += cross_entropy_bits / 8 + estimated_bytes += tiles * NUM_STATES * 4 + candidates.append((estimated_bytes, probability_bits, frequencies, modes)) + + # Among the tables within tolerance of the cheapest, the smallest: a coarser grid costs a little rate and saves a + # proportionally larger decode table. + best_bytes = min(value for value, *_ in candidates) + tolerance = original_bytes * _AUTO_RATIO_TOLERANCE + _, probability_bits, frequencies, modes = min( + (candidate for candidate in candidates if candidate[0] <= best_bytes + tolerance), key=lambda candidate: candidate[1], + ) + frequencies, cdfs, decode_tables = _build_tables_from_frequencies(frequencies, probability_bits) + return probability_bits, frequencies, cdfs, decode_tables, modes + + +def _encode_raw_stream(symbols: np.ndarray) -> np.ndarray: + words = np.zeros((symbols.size + 1) // 2, dtype=np.uint16) + words |= symbols[0::2].astype(np.uint16) + if symbols.size > 1: + words[: symbols[1::2].size] |= symbols[1::2].astype(np.uint16) << 8 + return words + + +def _encode_stream(symbols: np.ndarray, frequencies: np.ndarray, cdfs: np.ndarray, probability_bits: int): + states = np.full(NUM_STATES, STATE_MIN, dtype=np.uint32) + words: list[int] = [] + table_size = 1 << probability_bits + state_check_shift = _STATE_BITS - probability_bits + + for base in range(0, symbols.size, NUM_STATES): + limit = min(NUM_STATES, symbols.size - base) + for lane in range(limit): + symbol = int(symbols[base + lane]) + frequency = int(frequencies[symbol]) + state = int(states[lane]) + if state >= (frequency << state_check_shift): + words.append(state & 0xFFFF) + state >>= _RENORM_BITS + state = (state // frequency) * table_size + (state % frequency) + int(cdfs[symbol]) + states[lane] = state + return states, np.asarray(words, dtype=np.uint16) + + +def encode( + *, weight: torch.Tensor, tile_elements: int, probability_bits: int, raw_lane_threshold: float, **_ignored, +) -> TileBuffers: + if weight.numel() == 0: + raise ValueError("tile_ans does not support empty tensors") + contiguous = weight.detach().contiguous() + num_elements = contiguous.numel() + num_lanes = contiguous.element_size() + raw = contiguous.reshape(-1).view(torch.uint8).cpu().numpy().copy() + lane_bytes = raw.reshape(num_elements, num_lanes) + counts = np.stack([np.bincount(lane_bytes[:, lane], minlength=256) for lane in range(num_lanes)]) + if probability_bits == AUTO_PROB_BITS: + probability_bits, frequencies, cdfs, decode_tables, lane_modes = select_coding_options( + counts, raw_lane_threshold, tile_elements + ) + else: + frequencies, cdfs, decode_tables = build_tables_from_counts(counts, probability_bits) + lane_modes = lane_modes_from_counts(counts, raw_lane_threshold, tile_elements) + + tiles = num_tiles(num_elements, tile_elements) + num_streams = tiles * num_lanes + num_ans_lanes = int((lane_modes == LANE_ANS).sum()) + states = np.empty((tiles * num_ans_lanes, NUM_STATES), dtype=np.uint32) + offsets = np.empty(num_streams + 1, dtype=np.uint32) + offsets[0] = 0 + payload_parts = [] + + stream = 0 + ans_stream = 0 + payload_words = 0 + for lane in range(num_lanes): + for tile in range(tiles): + begin = tile * tile_elements + end = min(begin + tile_elements, num_elements) + if lane_modes[lane] == LANE_RAW: + words = _encode_raw_stream(lane_bytes[begin:end, lane]) + else: + stream_states, words = _encode_stream( + lane_bytes[begin:end, lane], frequencies[lane], cdfs[lane], probability_bits, + ) + states[ans_stream] = stream_states + ans_stream += 1 + if payload_words + words.size >= 1 << 32: + raise ValueError("tile_ans payload exceeds uint32 offset capacity") + payload_parts.append(words) + payload_words += words.size + offsets[stream + 1] = payload_words + stream += 1 + + payload = np.concatenate(payload_parts) if payload_parts else np.empty(0, dtype=np.uint16) + return TileBuffers( + payload=torch.from_numpy(payload), offsets=torch.from_numpy(offsets), states=torch.from_numpy(states), + decode_tables=torch.from_numpy(decode_tables), lane_modes=torch.from_numpy(lane_modes), + layout=make_layout(num_elements, num_lanes, tile_elements, probability_bits), + ) + + +def _decode_group(states, table, words, pointer, output, base, valid_lanes, probability_bits): + table_size = 1 << probability_bits + reads = [] + for lane in range(valid_lanes): + state = int(states[lane]) + slot = state & (table_size - 1) + entry = int(table[slot]) + symbol = entry & 0xFF + frequency = (entry >> 8) & 0xFFF + if frequency == 0: + frequency = table_size + cdf = entry >> 20 + output[base + lane] = symbol + state = frequency * (state >> probability_bits) + (slot - cdf) + states[lane] = state + if state < STATE_MIN: + reads.append(lane) + + first_word = pointer - len(reads) + if first_word < 0: + raise ValueError("tile_ans payload is truncated") + for index, lane in enumerate(reads): + states[lane] = (int(states[lane]) << _RENORM_BITS) | int(words[first_word + index]) + return first_word + + +def decode(buffers: TileBuffers, *, dtype: torch.dtype, **_ignored) -> torch.Tensor: + payload = buffers.payload.detach().cpu().numpy() + offsets = buffers.offsets.detach().cpu().numpy().astype(np.int64) + states = buffers.states.detach().cpu().numpy().astype(np.uint32) + tables = buffers.decode_tables.detach().cpu().numpy().astype(np.uint32) + lane_modes = buffers.lane_modes.detach().cpu().numpy().astype(np.uint8) + tile_elements, probability_bits, num_lanes, num_elements = parse_layout_cached(buffers.layout) + + output = np.empty(num_elements * num_lanes, dtype=np.uint8) + tiles = num_tiles(num_elements, tile_elements) + stream = 0 + ans_stream = 0 + for byte_lane in range(num_lanes): + for tile in range(tiles): + tile_begin = tile * tile_elements + tile_count = min(tile_elements, num_elements - tile_begin) + begin = int(offsets[stream]) + pointer = int(offsets[stream + 1]) + words = payload[begin:pointer] + if lane_modes[byte_lane] == LANE_RAW: + expected_words = (tile_count + 1) // 2 + if words.size != expected_words: + raise ValueError("tile_ans raw lane payload length is invalid") + lane_output = np.empty(tile_count, dtype=np.uint8) + lane_output[0::2] = (words & 0xFF).astype(np.uint8) + if tile_count > 1: + lane_output[1::2] = (words[: tile_count // 2] >> 8).astype(np.uint8) + else: + remainder = tile_count % NUM_STATES + stream_states = states[ans_stream].copy() + ans_stream += 1 + pointer -= begin + lane_output = np.empty(tile_count, dtype=np.uint8) + offset = tile_count - remainder + if remainder: + pointer = _decode_group( + stream_states, tables[byte_lane], words, pointer, lane_output, offset, remainder, probability_bits + ) + while offset > 0: + offset -= NUM_STATES + pointer = _decode_group( + stream_states, tables[byte_lane], words, pointer, lane_output, offset, NUM_STATES, probability_bits + ) + if pointer != 0: + raise ValueError("tile_ans payload contains unread words") + output[tile_begin * num_lanes + byte_lane : (tile_begin + tile_count) * num_lanes : num_lanes] = lane_output + stream += 1 + + return torch.from_numpy(output).view(dtype).to(buffers.payload.device) diff --git a/entropack/schemes/tile_ans/format.py b/entropack/schemes/tile_ans/format.py new file mode 100644 index 0000000..0704b2d --- /dev/null +++ b/entropack/schemes/tile_ans/format.py @@ -0,0 +1,164 @@ +import math +from dataclasses import dataclass +from typing import Annotated, NamedTuple + +import torch + +from ..config import CompressionConfig, OneOf, Range +from ..base import cached_parse + +BLOCK_SIZE = 256 +HISTOGRAM_WARP_BUDGET = 16 +HISTOGRAM_MIN_WARPS = 2 +HISTOGRAM_MAX_WARPS = 8 +ENCODE_TABLE_SHARED_BYTES = 2 * 256 * 2 + +AUTO_PROB_BITS = 0 +AUTO_TILE_ELEMENTS = 0 +SUPPORTED_PROB_BITS = (9, 10, 11, 12) +NUM_STATES = 32 +STATE_MIN = 1 << 15 +LANE_ANS = 0 +LANE_RAW = 1 + +TILE_ELEMENTS_LIMIT = 1 << 31 +RAW_LANE_THRESHOLD_LIMIT = 8.0 +TILE_ELEMENTS_MESSAGE = ( + "tile_elements must be positive, or 0 for automatic selection, and fit the CUDA int32 launch ABI" +) + + +@dataclass +class TileANSConfig(CompressionConfig): + """Lossless compression settings for the supported tensor dtypes. + + Tile size and probability precision are stored with the compressed tensor. A zero value + requests automatic selection. Block width controls GPU execution.""" + + #: Elements per independent tile. Zero selects 4096, or 8192 for tensors larger than 32 MiB. + tile_elements: Annotated[int, Range(0, TILE_ELEMENTS_LIMIT - 1, message=TILE_ELEMENTS_MESSAGE)] = AUTO_TILE_ELEMENTS + #: Probability-table precision. Zero selects using the tensor symbol histogram. + probability_bits: Annotated[int, OneOf(SUPPORTED_PROB_BITS, silent=(AUTO_PROB_BITS,))] = AUTO_PROB_BITS + #: Byte streams at or above this estimated cost in bits per symbol are stored directly. + raw_lane_threshold: Annotated[float, Range(0.0, RAW_LANE_THRESHOLD_LIMIT)] = 7.9 + #: GPU block width for encoding and decoding. None selects a device-dependent value. + threads_per_block: Annotated[int | None, Range(1, None)] = None + + +OPTIONS_BY_DTYPE = { + torch.float32: (8192, 10), torch.float16: (8192, 11), torch.bfloat16: (0, 11), + torch.float8_e4m3fn: (8192, 0), torch.float8_e4m3fnuz: (8192, 0), + torch.float8_e5m2: (8192, 0), torch.float8_e5m2fnuz: (8192, 0), + torch.int64: (8192, 10), torch.int32: (16384, 10), torch.int16: (8192, 9), + torch.int8: (8192, 0), torch.uint64: (8192, 9), torch.uint32: (8192, 10), + torch.uint16: (8192, 9), torch.uint8: (8192, 0), torch.bool: (8192, 9), +} + + +class TileBuffers(NamedTuple): + payload: torch.Tensor + offsets: torch.Tensor + states: torch.Tensor + decode_tables: torch.Tensor + lane_modes: torch.Tensor + layout: torch.Tensor + + +PACKED_KEYS = TileBuffers._fields + + +def num_tiles(num_elements: int, tile_elements: int) -> int: + return -(-num_elements // tile_elements) + + +def num_streams(num_elements: int, num_lanes: int, tile_elements: int) -> int: + return num_tiles(num_elements, tile_elements) * num_lanes + + +def make_layout(num_elements: int, num_lanes: int, tile_elements: int, probability_bits: int): + if probability_bits not in SUPPORTED_PROB_BITS: + raise ValueError(f"probability_bits must be one of {SUPPORTED_PROB_BITS}") + return torch.tensor([tile_elements, probability_bits, num_lanes, num_elements], dtype=torch.int64) + + +def parse_layout(layout: torch.Tensor) -> tuple[int, int, int, int]: + if layout.dtype != torch.int64 or layout.ndim != 1 or layout.numel() != 4: + raise ValueError("tile_ans layout must be int64[4]") + tile_elements, prob_bits, num_lanes, num_elements = (int(value) for value in layout.detach().cpu().tolist()) + if tile_elements <= 0 or prob_bits not in SUPPORTED_PROB_BITS or num_lanes <= 0 or num_elements <= 0: + raise ValueError( + "invalid tile_ans layout values: " f"tile_elements={tile_elements}, prob_bits={prob_bits}, " + f"num_lanes={num_lanes}, num_elements={num_elements}" + ) + return tile_elements, prob_bits, num_lanes, num_elements + + +def parse_layout_cached(layout: torch.Tensor) -> tuple[int, int, int, int]: + return cached_parse(layout, parse_layout, "_tile_ans_layout") + + +def validate_packed(buffers: dict[str, torch.Tensor], shape, dtype: torch.dtype) -> None: + missing = [key for key in PACKED_KEYS if key not in buffers] + if missing: + raise ValueError(f"tile_ans packed data is missing buffers: {missing}") + if not all(isinstance(buffers[key], torch.Tensor) for key in PACKED_KEYS): + raise TypeError("tile_ans packed buffers must be torch.Tensor values") + + expected_layout = { + "payload": (torch.uint16, 1), "offsets": (torch.uint32, 1), "states": (torch.uint32, 2), + "decode_tables": (torch.uint32, 2), "lane_modes": (torch.uint8, 1), "layout": (torch.int64, 1), + } + for key, (expected_dtype, ndim) in expected_layout.items(): + tensor = buffers[key] + if not tensor.is_contiguous(): + raise ValueError(f"tile_ans buffer '{key}' must be contiguous") + if tensor.dtype != expected_dtype or tensor.ndim != ndim: + raise ValueError( + f"tile_ans buffer '{key}' must be {ndim}D {expected_dtype}, " + f"got shape={tuple(tensor.shape)}, dtype={tensor.dtype}" + ) + + devices = {buffers[key].device for key in PACKED_KEYS} + if len(devices) != 1: + raise ValueError(f"tile_ans packed buffers must share one device, got {devices}") + + tile_elements, prob_bits, num_lanes, num_elements = parse_layout(buffers["layout"]) + element_size = torch.empty((), dtype=dtype).element_size() + if num_lanes != element_size: + raise ValueError(f"tile_ans layout has {num_lanes} byte lanes but dtype {dtype} uses {element_size} bytes") + normalized_shape = tuple(shape) + if any(not isinstance(dim, int) or dim < 0 for dim in normalized_shape): + raise ValueError(f"invalid tile_ans tensor shape: {normalized_shape}") + if math.prod(normalized_shape) != num_elements: + raise ValueError(f"tile_ans shape {normalized_shape} does not match {num_elements} elements") + + streams = num_streams(num_elements, num_lanes, tile_elements) + tiles = num_tiles(num_elements, tile_elements) + if tuple(buffers["decode_tables"].shape) != (num_lanes, 1 << prob_bits): + raise ValueError("tile_ans decode_tables shape does not match layout") + if buffers["lane_modes"].numel() != num_lanes: + raise ValueError("tile_ans lane_modes length does not match layout") + lane_modes = buffers["lane_modes"].detach().cpu() + if ((lane_modes != LANE_ANS) & (lane_modes != LANE_RAW)).any(): + raise ValueError("tile_ans lane_modes contains an unknown codec mode") + ans_lanes = int((lane_modes == LANE_ANS).sum()) + if tuple(buffers["states"].shape) != (tiles * ans_lanes, NUM_STATES): + raise ValueError("tile_ans states shape does not match ANS lanes/layout") + if buffers["offsets"].numel() != streams + 1: + raise ValueError("tile_ans offsets length does not match layout") + + offsets = buffers["offsets"].detach().cpu().to(torch.int64) + if offsets[0] != 0 or offsets[-1] != buffers["payload"].numel(): + raise ValueError("tile_ans offsets endpoints are invalid") + if (offsets[1:] < offsets[:-1]).any(): + raise ValueError("tile_ans offsets must be monotone") + + for byte_lane, mode in enumerate(lane_modes.tolist()): + if mode != LANE_RAW: + continue + first_stream = byte_lane * tiles + sizes = offsets[first_stream + 1 : first_stream + tiles + 1] - offsets[first_stream : first_stream + tiles] + raw_words = torch.full_like(sizes, (tile_elements + 1) // 2) + raw_words[-1] = (num_elements - (tiles - 1) * tile_elements + 1) // 2 + if not torch.equal(sizes, raw_words): + raise ValueError("tile_ans raw lane payload length is invalid") diff --git a/entropack/schemes/tile_ans/tile_ans.cu b/entropack/schemes/tile_ans/tile_ans.cu new file mode 100644 index 0000000..400a5ca --- /dev/null +++ b/entropack/schemes/tile_ans/tile_ans.cu @@ -0,0 +1,469 @@ +#include "device.cuh" + +// Encode kernels: +// tile_ans_histogram_kernel per-byte-lane 256-bin symbol counts. Every warp keeps a private histogram in shared +// memory and the block folds them, so a block contributes at most 256 global atomics per +// lane. +// tile_ans_encode_kernel one warp per (tile, byte lane) stream: 32 interleaved rANS states, 16-bit +// renormalization, words written to that stream's scratch region. A lane marked raw is +// packed two bytes per word instead of being coded. +// tile_ans_compact_kernel gathers the per-stream scratch words into one payload. +// +// Decode kernels, selected by the byte-lane pattern stored in the checkpoint. All produce identical output and differ only +// in how many lanes one warp reassembles at once: +// tile_ans_decode_raw0_ans1_kernel two lanes, raw then coded: the bf16 case, where one warp rebuilds both bytes of every +// element and stages the coded lane's table in shared memory. +// tile_ans_decode_raw3_ans1_kernel four lanes, the first three raw: the fp32 case. +// tile_ans_decode_all_raw_kernel every lane raw, i.e. nothing was compressible: one block per tile, a plain unpack. +// tile_ans_decode_kernel the general case: the grid carries the byte lane, coded lanes are decoded by rANS +// through a shared-memory table and raw lanes are unpacked. +// All four share one argument list, so the host builds one tuple and picks a name; the raw-only kernels ignore the state and +// table arguments. +// +// The rANS state machine, the renormalization variants and the shared decode helpers live in device.cuh, which the +// lattice_rans lane compiles against as well. + + +extern "C" __global__ void tile_ans_histogram_kernel( + const uint8_t* __restrict__ input, + uint64_t* __restrict__ histograms, + int64_t num_elements, + int num_lanes) { + extern __shared__ uint32_t warp_bins[]; + const int warp = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int bins_per_warp = num_lanes * 256; + const int total_bins = warps_per_block * bins_per_warp; + for (int index = threadIdx.x; index < total_bins; index += blockDim.x) { + warp_bins[index] = 0; + } + __syncthreads(); + + uint32_t* bins = warp_bins + warp * bins_per_warp; + int64_t element = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t stride = static_cast(gridDim.x) * blockDim.x; + for (; element < num_elements; element += stride) { + const uint8_t* value = input + element * num_lanes; +#pragma unroll + for (int byte_lane = 0; byte_lane < num_lanes; ++byte_lane) { + atomicAdd(&bins[byte_lane * 256 + value[byte_lane]], 1u); + } + } + __syncthreads(); + for (int index = threadIdx.x; index < bins_per_warp; index += blockDim.x) { + uint32_t sum = 0; +#pragma unroll + for (int source_warp = 0; source_warp < warps_per_block; ++source_warp) { + sum += warp_bins[source_warp * bins_per_warp + index]; + } + if (sum) { + atomicAdd( + reinterpret_cast(histograms + index), + static_cast(sum)); + } + } +} + +extern "C" __global__ void tile_ans_encode_kernel( + const uint8_t* __restrict__ input, + const uint16_t* __restrict__ frequencies, + const uint16_t* __restrict__ cdfs, + const uint8_t* __restrict__ lane_modes, + uint16_t* __restrict__ scratch, + uint32_t* __restrict__ word_counts, + uint32_t* __restrict__ final_states, + int64_t num_elements, + int tile_elements, + int num_lanes, + int num_tiles) { + const int warp_in_block = threadIdx.x >> 5; + const int lane = threadIdx.x & 31; + const int warps_per_block = blockDim.x >> 5; + const int tile = blockIdx.x * warps_per_block + warp_in_block; + const int byte_lane = blockIdx.y; + + extern __shared__ uint16_t shared_encode_tables[]; + uint16_t* frequency = shared_encode_tables; + uint16_t* cdf = frequency + 256; + const uint16_t* source_frequency = frequencies + byte_lane * 256; + const uint16_t* source_cdf = cdfs + byte_lane * 256; + for (int index = threadIdx.x; index < 256; index += blockDim.x) { + frequency[index] = source_frequency[index]; + cdf[index] = source_cdf[index]; + } + __syncthreads(); + + if (tile >= num_tiles || byte_lane >= num_lanes) { + return; + } + const int stream = byte_lane * num_tiles + tile; + const int64_t tile_begin = static_cast(tile) * tile_elements; + const int tile_count = static_cast( + (num_elements - tile_begin < tile_elements) + ? (num_elements - tile_begin) + : tile_elements); + + uint16_t* stream_scratch = scratch + static_cast(stream) * tile_elements; + if (lane_modes[byte_lane] != 0) { + const int raw_words = (tile_count + 1) >> 1; + for (int word = lane; word < raw_words; word += kNumStates) { + const int first = word << 1; + const uint32_t low = input[(tile_begin + first) * num_lanes + byte_lane]; + const uint32_t high = (first + 1 < tile_count) + ? input[(tile_begin + first + 1) * num_lanes + byte_lane] + : 0u; + stream_scratch[word] = static_cast(low | (high << 8)); + } + if (lane == 0) { + word_counts[stream] = raw_words; + } + return; + } + + int ans_lane = 0; + for (int prior_lane = 0; prior_lane < byte_lane; ++prior_lane) { + ans_lane += lane_modes[prior_lane] == 0; + } + const int ans_stream = ans_lane * num_tiles + tile; + uint32_t state = kStateMin; + uint32_t word_count = 0; + constexpr uint32_t state_check_mul = 1u << (31 - kProbBits); + + for (int base = 0; base < tile_count; base += kNumStates) { + const bool valid = base + lane < tile_count; + const uint32_t symbol = valid + ? input[(tile_begin + base + lane) * num_lanes + byte_lane] + : 0u; + const uint32_t freq = valid ? frequency[symbol] : 1u; + const bool emit = valid && state >= freq * state_check_mul; + const uint32_t vote = __ballot_sync(0xFFFFFFFFu, emit); + const uint32_t prefix = __popc(vote & lane_mask_lt()); + if (emit) { + stream_scratch[word_count + prefix] = static_cast(state); + state >>= 16; + } + word_count += __popc(vote); + if (valid) { + state = (state / freq) * kTableSize + (state % freq) + cdf[symbol]; + } + } + + final_states[ans_stream * kNumStates + lane] = state; + if (lane == 0) { + word_counts[stream] = word_count; + } +} + +extern "C" __global__ void tile_ans_compact_kernel( + const uint16_t* __restrict__ scratch, + const uint32_t* __restrict__ offsets, + uint16_t* __restrict__ payload, + int tile_elements, + int num_streams) { + const int stream = blockIdx.x; + if (stream >= num_streams) { + return; + } + const uint32_t begin = offsets[stream]; + const uint32_t count = offsets[stream + 1] - begin; + const uint16_t* source = scratch + static_cast(stream) * tile_elements; + for (uint32_t index = threadIdx.x; index < count; index += blockDim.x) { + payload[begin + index] = source[index]; + } +} + +extern "C" __global__ void tile_ans_decode_raw0_ans1_kernel( + const uint16_t* __restrict__ payload, + const uint32_t* __restrict__ offsets, + const uint32_t* __restrict__ states, + const uint32_t* __restrict__ decode_tables, + const uint8_t* __restrict__ lane_modes, + uint8_t* __restrict__ output, + int64_t num_elements, + int tile_elements, + int num_lanes, + int num_tiles) { + extern __shared__ uint32_t shared_table[]; + const uint4* src4 = reinterpret_cast(decode_tables + kTableSize); + uint4* dst4 = reinterpret_cast(shared_table); + for (int i = threadIdx.x; i < kTableSize / 4; i += blockDim.x) dst4[i] = src4[i]; + __syncthreads(); + const uint32_t* __restrict__ table = shared_table; + + const int warp_in_block = threadIdx.x >> 5; + const int lane = threadIdx.x & 31; + const int warps_per_block = blockDim.x >> 5; + const int tile = blockIdx.x * warps_per_block + warp_in_block; + if (tile >= num_tiles || num_lanes != 2) { + return; + } + + const int64_t tile_begin = static_cast(tile) * tile_elements; + const int tile_count = static_cast( + (num_elements - tile_begin < tile_elements) + ? (num_elements - tile_begin) + : tile_elements); + const int raw_stream = tile; + const int ans_stream = num_tiles + tile; + const uint32_t expected_raw_words = (tile_count + 1) >> 1; + if (offsets[raw_stream + 1] - offsets[raw_stream] != expected_raw_words) { + return; + } + const uint16_t* raw_words = payload + offsets[raw_stream]; + const uint16_t* input_begin = payload + offsets[ans_stream]; + const uint16_t* input = payload + offsets[ans_stream + 1]; + uint32_t state = states[tile * kNumStates + lane]; + uint16_t* output16 = reinterpret_cast(output); + const uint32_t mask_ge = lane_mask_ge(); + + const int remainder = tile_count & (kNumStates - 1); + int output_offset = tile_count - remainder; + if (remainder) { + const bool valid = lane < remainder; + const int local_index = output_offset + lane; + if (valid) { + const uint32_t high = rans_decode_symbol(state, table); + const uint32_t packed_low = raw_words[local_index >> 1]; + const uint32_t low = (packed_low >> ((local_index & 1) * 8)) & 0xFFu; + output16[tile_begin + local_index] = static_cast(low | (high << 8)); + } + rans_renormalize_unchecked(valid, state, input, mask_ge); + } + + while (output_offset > 0) { + output_offset -= kNumStates; + const int local_index = output_offset + lane; + const uint32_t high = rans_decode_symbol(state, table); + const uint32_t packed_low = raw_words[local_index >> 1]; + const uint32_t low = (packed_low >> ((local_index & 1) * 8)) & 0xFFu; + output16[tile_begin + local_index] = static_cast(low | (high << 8)); + rans_renormalize_unchecked(true, state, input, mask_ge); + } +} + +__device__ __forceinline__ void decode_group_raw3_ans1( + bool valid, + uint32_t& state, + const uint32_t* __restrict__ table, + const uint16_t*& input, + const uint16_t* __restrict__ input_begin, + const uint16_t* __restrict__ raw0, + const uint16_t* __restrict__ raw1, + const uint16_t* __restrict__ raw2, + uint32_t* __restrict__ output, + int64_t output_index, + int local_index) { + if (valid) { + const uint32_t high = rans_decode_symbol(state, table); + const int word = local_index >> 1; + const int shift = (local_index & 1) * 8; + const uint32_t b0 = (raw0[word] >> shift) & 0xFFu; + const uint32_t b1 = (raw1[word] >> shift) & 0xFFu; + const uint32_t b2 = (raw2[word] >> shift) & 0xFFu; + output[output_index] = b0 | (b1 << 8) | (b2 << 16) | (high << 24); + } + (void)rans_renormalize_checked(valid, state, input, input_begin); +} + +extern "C" __global__ void tile_ans_decode_raw3_ans1_kernel( + const uint16_t* __restrict__ payload, + const uint32_t* __restrict__ offsets, + const uint32_t* __restrict__ states, + const uint32_t* __restrict__ decode_tables, + const uint8_t* __restrict__ lane_modes, + uint8_t* __restrict__ output, + int64_t num_elements, + int tile_elements, + int num_lanes, + int num_tiles) { + const int warp_in_block = threadIdx.x >> 5; + const int lane = threadIdx.x & 31; + const int warps_per_block = blockDim.x >> 5; + const int tile = blockIdx.x * warps_per_block + warp_in_block; + + const uint32_t* table = decode_tables + 3 * kTableSize; + if (tile >= num_tiles || num_lanes != 4) { + return; + } + const int64_t tile_begin = static_cast(tile) * tile_elements; + const int tile_count = static_cast( + (num_elements - tile_begin < tile_elements) + ? (num_elements - tile_begin) + : tile_elements); + + const uint32_t expected_raw_words = (tile_count + 1) >> 1; + if (offsets[tile + 1] - offsets[tile] != expected_raw_words || + offsets[num_tiles + tile + 1] - offsets[num_tiles + tile] != expected_raw_words || + offsets[2 * num_tiles + tile + 1] - offsets[2 * num_tiles + tile] != expected_raw_words) { + return; + } + const uint16_t* raw0 = payload + offsets[tile]; + const uint16_t* raw1 = payload + offsets[num_tiles + tile]; + const uint16_t* raw2 = payload + offsets[2 * num_tiles + tile]; + const int ans_stream = 3 * num_tiles + tile; + const uint16_t* input_begin = payload + offsets[ans_stream]; + const uint16_t* input = payload + offsets[ans_stream + 1]; + uint32_t state = states[tile * kNumStates + lane]; + uint32_t* output32 = reinterpret_cast(output); + + const int remainder = tile_count & (kNumStates - 1); + int output_offset = tile_count - remainder; + if (remainder) { + decode_group_raw3_ans1( + lane < remainder, state, table, input, input_begin, + raw0, raw1, raw2, output32, + tile_begin + output_offset + lane, output_offset + lane); + } + while (output_offset > 0) { + output_offset -= kNumStates; + decode_group_raw3_ans1( + true, state, table, input, input_begin, + raw0, raw1, raw2, output32, + tile_begin + output_offset + lane, output_offset + lane); + } +} + +extern "C" __global__ void tile_ans_decode_all_raw_kernel( + const uint16_t* __restrict__ payload, + const uint32_t* __restrict__ offsets, + const uint32_t* __restrict__ states, + const uint32_t* __restrict__ decode_tables, + const uint8_t* __restrict__ lane_modes, + uint8_t* __restrict__ output, + int64_t num_elements, + int tile_elements, + int num_lanes, + int num_tiles) { + const int tile = blockIdx.x; + if (tile >= num_tiles) { + return; + } + const int64_t tile_begin = static_cast(tile) * tile_elements; + const int tile_count = static_cast( + (num_elements - tile_begin < tile_elements) + ? (num_elements - tile_begin) + : tile_elements); + + for (int element = threadIdx.x; element < tile_count; element += blockDim.x) { + uint64_t value = 0; +#pragma unroll + for (int byte_lane = 0; byte_lane < num_lanes; ++byte_lane) { + const int stream = byte_lane * num_tiles + tile; + const uint32_t expected_words = (tile_count + 1) >> 1; + if (offsets[stream + 1] - offsets[stream] != expected_words) { + return; + } + const uint16_t packed = payload[offsets[stream] + (element >> 1)]; + const uint32_t byte = (packed >> ((element & 1) * 8)) & 0xFFu; + value |= static_cast(byte) << (byte_lane * 8); + } + const int64_t output_index = tile_begin + element; + if (num_lanes == 8) { + reinterpret_cast(output)[output_index] = value; + } else if (num_lanes == 4) { + reinterpret_cast(output)[output_index] = static_cast(value); + } else if (num_lanes == 2) { + reinterpret_cast(output)[output_index] = static_cast(value); + } else { + output[output_index] = static_cast(value); + } + } +} + +extern "C" __global__ void tile_ans_decode_kernel( + const uint16_t* __restrict__ payload, + const uint32_t* __restrict__ offsets, + const uint32_t* __restrict__ states, + const uint32_t* __restrict__ decode_tables, + const uint8_t* __restrict__ lane_modes, + uint8_t* __restrict__ output, + int64_t num_elements, + int tile_elements, + int num_lanes, + int num_tiles) { + const int warp_in_block = threadIdx.x >> 5; + const int lane = threadIdx.x & 31; + const int warps_per_block = blockDim.x >> 5; + const int tile = blockIdx.x * warps_per_block + warp_in_block; + const int byte_lane = blockIdx.y; + + const bool raw_lane = lane_modes[byte_lane] != 0; + extern __shared__ uint32_t shared_decode_tables[]; + uint32_t* table = shared_decode_tables; + const uint32_t* source_table = decode_tables + byte_lane * kTableSize; + if (!raw_lane) { + for (int index = threadIdx.x; index < static_cast(kTableSize); index += blockDim.x) { + table[index] = source_table[index]; + } + } + __syncthreads(); + + if (tile >= num_tiles || byte_lane >= num_lanes) { + return; + } + const int stream = byte_lane * num_tiles + tile; + const int64_t tile_begin = static_cast(tile) * tile_elements; + const int tile_count = static_cast( + (num_elements - tile_begin < tile_elements) + ? (num_elements - tile_begin) + : tile_elements); + + const uint32_t begin_word = offsets[stream]; + const uint32_t end_word = offsets[stream + 1]; + const uint16_t* input_begin = payload + begin_word; + const uint16_t* input = payload + end_word; + + if (raw_lane) { + const int raw_words = (tile_count + 1) >> 1; + if (end_word - begin_word != static_cast(raw_words)) { + return; + } + for (int word = lane; word < raw_words; word += kNumStates) { + const uint32_t packed = input_begin[word]; + const int first = word << 1; + output[(tile_begin + first) * num_lanes + byte_lane] = + static_cast(packed); + if (first + 1 < tile_count) { + output[(tile_begin + first + 1) * num_lanes + byte_lane] = + static_cast(packed >> 8); + } + } + return; + } + + int ans_lane = 0; + for (int prior_lane = 0; prior_lane < byte_lane; ++prior_lane) { + ans_lane += lane_modes[prior_lane] == 0; + } + const int ans_stream = ans_lane * num_tiles + tile; + uint32_t state = states[ans_stream * kNumStates + lane]; + const int remainder = tile_count & (kNumStates - 1); + int output_offset = tile_count - remainder; + if (remainder) { + const bool valid = lane < remainder; + decode_group( + valid, + state, + table, + input, + input_begin, + output, + tile_begin + output_offset + lane, + num_lanes, + byte_lane); + } + + while (output_offset > 0) { + output_offset -= kNumStates; + decode_group( + true, + state, + table, + input, + input_begin, + output, + tile_begin + output_offset + lane, + num_lanes, + byte_lane); + } +} diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..6c31076 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,36 @@ +[build-system] +requires = ["setuptools>=77.0.3", "wheel"] +build-backend = "setuptools.build_meta" + +[project] +name = "entropack" +version = "0.1.0" +description = "EntroPack: general-purpose tensor compression for PyTorch with lossless and rate-controlled lossy schemes and GPU encoding and decoding." +readme = "README.md" +requires-python = ">=3.10" +license = "Apache-2.0" +license-files = ["LICENSE"] +authors = [{ name = "DiffSynth-Studio" }] +# Runtime deps. The CUDA lane additionally needs cupy for the running CUDA major version: +# `pip install entropack[cuda13]` (CUDA 13) or `[cuda12]` (CUDA 12). Without cupy +# the package degrades to the pure-torch eager lane instead of failing. +dependencies = [ + "torch>=2.10", + "numpy>=1.23", + "dahuffman>=0.4", +] + +[project.optional-dependencies] +cuda13 = ["cupy-cuda13x>=14"] +cuda12 = ["cupy-cuda12x>=14"] + +[tool.setuptools.packages.find] +include = ["entropack*"] + +[tool.setuptools.package-data] +"entropack.schemes.dfloat11" = ["*.cu"] +"entropack.schemes.tile_ans" = ["*.cu", "*.cuh"] +"entropack.schemes.lattice_rans" = ["*.cu"] + +[tool.pytest.ini_options] +testpaths = ["tests"]