10#ifndef FZ_ANS_DIETGPU_ANS_GPUANSENCODE_H
11#define FZ_ANS_DIETGPU_ANS_GPUANSENCODE_H
15#include "BatchPrefixSum.h"
16#include "GpuANSCodec.h"
17#include "GpuANSStatistics.h"
20namespace fz {
namespace ans {
22template <
int ProbBits>
23__device__ __forceinline__ uint32_t encodeOne(
27 ANSEncodedT* __restrict__ outWords,
28 const uint4* __restrict__ table) {
29 auto lookup = table[sym];
31 uint32_t pdf = lookup.x;
32 uint32_t cdf = lookup.y;
33 uint32_t div_m1 = lookup.z;
34 uint32_t div_shift = lookup.w;
36 constexpr ANSStateT kStateCheckMul = 1 << (kANSStateBits - ProbBits);
38 ANSStateT maxStateCheck = pdf * kStateCheckMul;
39 bool write = state >= maxStateCheck;
41 auto vote = __ballot_sync(0xffffffff, write);
42 auto prefix = __popc(vote & getLaneMaskLt());
45 outWords[outOffset + prefix] = state & kANSEncodedMask;
46 state >>= kANSEncodedBits;
49 uint32_t t = __umulhi(state, div_m1);
50 uint32_t div = (t + state) >> div_shift;
51 auto mod = state - (div * pdf);
53 constexpr uint32_t kProbBitsMul = 1 << ProbBits;
54 state = div * kProbBitsMul + mod + cdf;
59template <
int ProbBits>
60__device__ __forceinline__ uint32_t encodeOnePartial(
65 ANSEncodedT* __restrict__ outWords,
66 const uint4* __restrict__ table) {
68 auto lookup = table[sym];
70 uint32_t pdf = lookup.x;
71 uint32_t cdf = lookup.y;
72 uint32_t div_m1 = lookup.z;
73 uint32_t div_shift = lookup.w;
75 constexpr ANSStateT kStateCheckMul = 1 << (kANSStateBits - ProbBits);
77 ANSStateT maxStateCheck = pdf * kStateCheckMul;
78 bool write = (state >= maxStateCheck);
80 auto vote = __ballot_sync(0xffffffff, write);
81 auto prefix = __popc(vote & getLaneMaskLt());
84 outWords[outOffset + prefix] = state & kANSEncodedMask;
85 state >>= kANSEncodedBits;
88 uint32_t t = __umulhi(state, div_m1);
89 uint32_t div = (t + state) >> div_shift;
90 auto mod = state - (div * pdf);
92 constexpr uint32_t kProbBitsMul = 1 << ProbBits;
93 state = div * kProbBitsMul + mod + cdf;
98template <
int ProbBits,
int BlockSize>
99__global__
void ansEncodeBatch(
102 uint32_t maxNumCompressedBlocks,
103 uint32_t uncoalescedBlockStride,
104 uint8_t* compressedBlocks_dev,
105 uint32_t* compressedWords_dev,
106 const uint4* table_dev) {
107 uint32_t numBlocks = (inSize_dev + BlockSize - 1) / BlockSize;
108 int tid = threadIdx.x;
109 int grim_warp_numid =
110 __shfl_sync(0xffffffff, (blockIdx.x * blockDim.x + tid) / kWarpSize, 0);
111 int laneId = getLaneId();
113 __shared__ uint4 smemLookup[kNumSymbols];
116 if (tid < kNumSymbols) {
117 smemLookup[tid] = table_dev[tid];
121 uint32_t start = grim_warp_numid * BlockSize;
122 if (start >= inSize_dev) {
126 uint32_t blockSize = min(start + BlockSize, inSize_dev) - start;
128 if (grim_warp_numid >= numBlocks)
131 auto inBlock = in_dev + start;
132 auto outBlock = (ANSWarpState*)(compressedBlocks_dev
133 + grim_warp_numid * uncoalescedBlockStride);
135 assert(isPointerAligned(inBlock, kANSRequiredAlignment));
137 ANSEncodedT* outWords = (ANSEncodedT*)(outBlock + 1);
139 ANSStateT state = kANSStartState;
141 uint32_t inOffset = laneId;
142 uint32_t outOffset = 0;
144 constexpr int kUnroll = 8;
146 uint32_t limit = roundDown(blockSize, kWarpSize * kUnroll);
149 for (; inOffset < limit; inOffset += kWarpSize * kUnroll) {
151 for (
int j = 0; j < kUnroll; ++j) {
153 encodeOne<ProbBits>(state, inBlock[inOffset + j * kWarpSize], outOffset, outWords, smemLookup);
158 if (limit != blockSize) {
159 limit = roundDown(blockSize, kWarpSize);
161 for (; inOffset < limit; inOffset += kWarpSize) {
163 encodeOne<ProbBits>(state, inBlock[inOffset], outOffset, outWords, smemLookup);
165 if (limit != blockSize) {
166 bool valid = inOffset < blockSize;
167 ANSDecodedT sym = valid ? inBlock[inOffset] : ANSDecodedT(0);
168 outOffset += encodeOnePartial<ProbBits>(
169 valid, state, sym, outOffset, outWords, smemLookup);
173 outBlock->warpState[laneId] = state;
176 compressedWords_dev[grim_warp_numid] = outOffset;
180template <
typename A,
int B>
182 typedef uint32_t argument_type;
183 typedef uint32_t result_type;
185 template <
typename T>
186 __host__ __device__ uint32_t operator()(T x)
const {
187 constexpr int kDiv = B /
sizeof(A);
188 constexpr int kSize = kDiv < 1 ? 1 : kDiv;
190 return roundUp(x, T(kSize));
194template <
int Threads>
195__global__
void ansEncodeCoalesceBatch(
196 const uint8_t* __restrict__ compressedBlocks_dev,
197 int uncompressedWords,
198 uint32_t maxNumCompressedBlocks,
199 uint32_t uncoalescedBlockStride,
200 const uint32_t* __restrict__ compressedWords_dev,
201 const uint32_t* __restrict__ compressedWordsPrefix_dev,
202 const uint4* __restrict__ table_dev,
203 uint32_t config_probBits,
205 uint32_t* outSize_dev) {
207 auto numBlocks = divUp(uncompressedWords, kDefaultBlockSize);
209 int block = blockIdx.x;
210 int tid = threadIdx.x;
212 ANSCoalescedHeader* headerOut = (ANSCoalescedHeader*)out_dev;
217 uint32_t totalCompressedWords = 0;
220 totalCompressedWords =
221 compressedWordsPrefix_dev[numBlocks - 1] +
223 compressedWords_dev[numBlocks - 1],
224 kBlockAlignment /
sizeof(ANSEncodedT));
227 ANSCoalescedHeader header;
228 header.setMagicAndVersion();
229 header.setNumBlocks(numBlocks);
230 header.setTotalUncompressedWords(uncompressedWords);
231 header.setTotalCompressedWords(totalCompressedWords);
232 header.setProbBits(config_probBits);
235 *outSize_dev = header.getTotalCompressedSize();
241 auto probsOut = headerOut->getSymbolProbs();
244 for (
int i = tid; i < kNumSymbols; i += Threads) {
245 probsOut[i] = table_dev[i].x;
249 if (block >= numBlocks) {
254 auto uncoalescedBlock = compressedBlocks_dev +
255 block * uncoalescedBlockStride;
258 if (tid < kWarpSize) {
259 auto warpStateIn = (ANSWarpState*)uncoalescedBlock;
261 headerOut->getWarpStates()[block].warpState[tid] =
262 warpStateIn->warpState[tid];
265 auto blockWordsOut = headerOut->getBlockWords(numBlocks);
268 for (
int i = blockIdx.x * Threads + tid; i < numBlocks;
269 i += gridDim.x * Threads) {
270 uint32_t lastBlockWords = uncompressedWords % kDefaultBlockSize;
271 lastBlockWords = lastBlockWords == 0 ? kDefaultBlockSize : lastBlockWords;
273 uint32_t blockWords =
274 (i == numBlocks - 1) ? lastBlockWords : kDefaultBlockSize;
276 blockWordsOut[i] = uint2{
277 (blockWords << 16) | compressedWords_dev[i], compressedWordsPrefix_dev[i]};
281 uint32_t numWords = compressedWords_dev[block];
285 uint32_t limitEnd = divUp(numWords, kBlockAlignment /
sizeof(ANSEncodedT));
287 auto inT = (
const LoadT*)(uncoalescedBlock +
sizeof(ANSWarpState));
289 (LoadT*)(headerOut->getBlockDataStart(numBlocks) + compressedWordsPrefix_dev[block]);
291 for (uint32_t i = tid; i < limitEnd; i += Threads) {