9#ifndef FZ_ANS_DIETGPU_ANS_GPUANSCODEC_H
10#define FZ_ANS_DIETGPU_ANS_GPUANSCODEC_H
15#include "utils/DeviceUtils.h"
16#include "utils/StaticUtils.h"
18namespace fz {
namespace ans {
20using ANSStateT = uint32_t;
21using ANSEncodedT = uint16_t;
22using ANSDecodedT = uint8_t;
24struct __align__(16) ANSDecodedTx16 {
28struct __align__(8) ANSDecodedTx8 {
32struct __align__(4) ANSDecodedTx4 {
36constexpr uint32_t kNumSymbols = 1 << (
sizeof(ANSDecodedT) * 8);
37static_assert(kNumSymbols > 1,
"");
40constexpr uint32_t kDefaultBlockSize = 4096;
44constexpr int kANSStateBits =
sizeof(ANSStateT) * 8 - 1;
45constexpr int kANSEncodedBits =
sizeof(ANSEncodedT) * 8;
46constexpr ANSStateT kANSEncodedMask =
47 (ANSStateT(1) << kANSEncodedBits) - ANSStateT(1);
49constexpr ANSStateT kANSStartState = ANSStateT(1)
50 << (kANSStateBits - kANSEncodedBits);
51constexpr ANSStateT kANSMinState = ANSStateT(1)
52 << (kANSStateBits - kANSEncodedBits);
55constexpr uint32_t kANSMagic = 0xd00d;
58constexpr uint32_t kANSVersion = 0x0001;
63constexpr uint32_t kBlockAlignment = 16;
67 ANSStateT warpState[kWarpSize];
70struct __align__(32) ANSCoalescedHeader {
71 static __host__ __device__ uint32_t getCompressedOverhead(
73 constexpr int kAlignment = kBlockAlignment /
sizeof(uint2) == 0
75 : kBlockAlignment / sizeof(uint2);
77 return sizeof(ANSCoalescedHeader) +
79 sizeof(uint16_t) * kNumSymbols +
81 sizeof(ANSWarpState) * numBlocks +
83 sizeof(uint2) * roundUp(numBlocks, kAlignment);
86 __host__ __device__ uint32_t getTotalCompressedSize()
const {
87 return getCompressedOverhead() +
88 getTotalCompressedWords() *
sizeof(ANSEncodedT);
91 __host__ __device__ uint32_t getCompressedOverhead()
const {
92 return getCompressedOverhead(getNumBlocks());
95 __host__ __device__
float getCompressionRatio()
const {
96 return (
float)getTotalCompressedSize() /
97 (float)getTotalUncompressedWords() *
sizeof(ANSDecodedT);
100 __host__ __device__ uint32_t getNumBlocks()
const {
return numBlocks; }
101 __host__ __device__
void setNumBlocks(uint32_t nb) { numBlocks = nb; }
103 __host__ __device__
void setMagicAndVersion() {
104 magicAndVersion = (kANSMagic << 16) | kANSVersion;
107 __host__ __device__
void checkMagicAndVersion()
const {
108 assert((magicAndVersion >> 16) == kANSMagic);
109 assert((magicAndVersion & 0xffffU) == kANSVersion);
112 __host__ __device__ uint32_t getTotalUncompressedWords()
const {
113 return totalUncompressedWords;
115 __host__ __device__
void setTotalUncompressedWords(uint32_t words) {
116 totalUncompressedWords = words;
119 __host__ __device__ uint32_t getTotalCompressedWords()
const {
120 return totalCompressedWords;
122 __host__ __device__
void setTotalCompressedWords(uint32_t words) {
123 totalCompressedWords = words;
126 __host__ __device__ uint32_t getProbBits()
const {
return options & 0xf; }
127 __host__ __device__
void setProbBits(uint32_t bits) {
129 options = (options & 0xfffffff0U) | bits;
132 __host__ __device__
bool getUseChecksum()
const {
return options & 0x10; }
133 __host__ __device__
void setUseChecksum(
bool uc) {
134 options = (options & 0xffffffef) | (uint32_t(uc) << 4);
137 __host__ __device__ uint32_t getChecksum()
const {
return checksum; }
138 __host__ __device__
void setChecksum(uint32_t c) { checksum = c; }
140 __device__ uint16_t* getSymbolProbs() {
141 return (uint16_t*)(
this + 1);
143 __device__
const uint16_t* getSymbolProbs()
const {
144 return (
const uint16_t*)(
this + 1);
147 __device__ ANSWarpState* getWarpStates() {
148 return (ANSWarpState*)(getSymbolProbs() + kNumSymbols);
150 __device__
const ANSWarpState* getWarpStates()
const {
151 return (
const ANSWarpState*)(getSymbolProbs() + kNumSymbols);
154 __device__ uint2* getBlockWords(uint32_t numBlocks) {
155 return (uint2*)(getWarpStates() + numBlocks);
157 __device__
const uint2* getBlockWords(uint32_t numBlocks)
const {
158 return (
const uint2*)(getWarpStates() + numBlocks);
161 __device__ ANSEncodedT* getBlockDataStart(uint32_t numBlocks) {
162 constexpr int kAlignment = kBlockAlignment /
sizeof(uint2) == 0
164 : kBlockAlignment / sizeof(uint2);
165 return (ANSEncodedT*)(getBlockWords(numBlocks) + roundUp(numBlocks, kAlignment));
167 __device__
const ANSEncodedT* getBlockDataStart(uint32_t numBlocks)
const {
168 constexpr int kAlignment = kBlockAlignment /
sizeof(uint2) == 0
170 : kBlockAlignment / sizeof(uint2);
171 return (
const ANSEncodedT*)(getBlockWords(numBlocks) + roundUp(numBlocks, kAlignment));
175 uint32_t magicAndVersion;
177 uint32_t totalUncompressedWords;
178 uint32_t totalCompressedWords;
187static_assert(
sizeof(ANSCoalescedHeader) == 32,
"");
188static_assert(isEvenDivisor(
sizeof(ANSCoalescedHeader),
sizeof(uint4)),
"");
191constexpr int kANSRequiredAlignment = 4;
193constexpr int kANSDefaultProbBits = 10;
196constexpr __host__ __device__ uint32_t
197getRawCompBlockMaxSize(uint32_t uncompressedBlockBytes) {
199 uncompressedBlockBytes + (uncompressedBlockBytes / 4), kBlockAlignment);
202inline uint32_t getMaxBlockSizeUnCoalesced(uint32_t uncompressedBlockBytes) {
203 return sizeof(ANSWarpState) + getRawCompBlockMaxSize(uncompressedBlockBytes);
206inline uint32_t getMaxBlockSizeCoalesced(uint32_t uncompressedBlockBytes) {
207 return getRawCompBlockMaxSize(uncompressedBlockBytes);
210inline uint32_t getMaxCompressedSize(uint32_t uncompressedBytes) {
211 uint32_t blocks = divUp(uncompressedBytes, kDefaultBlockSize);
213 size_t rawSize = ANSCoalescedHeader::getCompressedOverhead(kDefaultBlockSize);
214 rawSize += (size_t)getMaxBlockSizeCoalesced(kDefaultBlockSize) * blocks;
215 rawSize = roundUp(rawSize,
sizeof(uint4));
221 inline __device__ BatchWriter(
void* out)
222 : out_((uint8_t*)out), outBlock_(nullptr) {}
224 inline __device__
void setBlock(uint32_t block) {
225 outBlock_ = out_ + block * kDefaultBlockSize;
228 inline __device__
void write(uint32_t offset, uint8_t sym) {
229 outBlock_[offset] = sym;