FZGPUModules 2.0
GPU-accelerated modular compression pipelines
Loading...
Searching...
No Matches
GpuANSEncode.h
1
10#ifndef FZ_ANS_DIETGPU_ANS_GPUANSENCODE_H
11#define FZ_ANS_DIETGPU_ANS_GPUANSENCODE_H
12
13#pragma once
14
15#include "BatchPrefixSum.h"
16#include "GpuANSCodec.h"
17#include "GpuANSStatistics.h"
18#include <cmath>
19
20namespace fz { namespace ans {
21
22template <int ProbBits>
23__device__ __forceinline__ uint32_t encodeOne(
24 ANSStateT& state,
25 ANSDecodedT sym,
26 uint32_t outOffset,
27 ANSEncodedT* __restrict__ outWords,
28 const uint4* __restrict__ table) {
29 auto lookup = table[sym];
30
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;
35
36 constexpr ANSStateT kStateCheckMul = 1 << (kANSStateBits - ProbBits);
37
38 ANSStateT maxStateCheck = pdf * kStateCheckMul;
39 bool write = state >= maxStateCheck;
40
41 auto vote = __ballot_sync(0xffffffff, write);
42 auto prefix = __popc(vote & getLaneMaskLt());
43
44 if (write) {
45 outWords[outOffset + prefix] = state & kANSEncodedMask;
46 state >>= kANSEncodedBits;
47 }
48
49 uint32_t t = __umulhi(state, div_m1);
50 uint32_t div = (t + state) >> div_shift;
51 auto mod = state - (div * pdf);
52
53 constexpr uint32_t kProbBitsMul = 1 << ProbBits;
54 state = div * kProbBitsMul + mod + cdf;
55
56 return __popc(vote);
57}
58
59template <int ProbBits>
60__device__ __forceinline__ uint32_t encodeOnePartial(
61 bool valid,
62 ANSStateT& state,
63 ANSDecodedT sym,
64 uint32_t outOffset,
65 ANSEncodedT* __restrict__ outWords,
66 const uint4* __restrict__ table) {
67 if (!valid) return 0;
68 auto lookup = table[sym];
69
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;
74
75 constexpr ANSStateT kStateCheckMul = 1 << (kANSStateBits - ProbBits);
76
77 ANSStateT maxStateCheck = pdf * kStateCheckMul;
78 bool write = (state >= maxStateCheck);
79
80 auto vote = __ballot_sync(0xffffffff, write);
81 auto prefix = __popc(vote & getLaneMaskLt());
82
83 if (write) {
84 outWords[outOffset + prefix] = state & kANSEncodedMask;
85 state >>= kANSEncodedBits;
86 }
87
88 uint32_t t = __umulhi(state, div_m1);
89 uint32_t div = (t + state) >> div_shift;
90 auto mod = state - (div * pdf);
91
92 constexpr uint32_t kProbBitsMul = 1 << ProbBits;
93 state = div * kProbBitsMul + mod + cdf;
94
95 return __popc(vote);
96}
97
98template <int ProbBits, int BlockSize>
99__global__ void ansEncodeBatch(
100 uint8_t* in_dev,
101 int inSize_dev,
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();
112
113 __shared__ uint4 smemLookup[kNumSymbols];
114
115 // we always have at least 256 threads
116 if (tid < kNumSymbols) {
117 smemLookup[tid] = table_dev[tid];
118 }
119 __syncthreads();
120
121 uint32_t start = grim_warp_numid * BlockSize;
122 if (start >= inSize_dev) {
123 return;
124 }
125
126 uint32_t blockSize = min(start + BlockSize, inSize_dev) - start;
127
128 if (grim_warp_numid >= numBlocks)
129 return;
130
131 auto inBlock = in_dev + start;
132 auto outBlock = (ANSWarpState*)(compressedBlocks_dev
133 + grim_warp_numid * uncoalescedBlockStride);
134
135 assert(isPointerAligned(inBlock, kANSRequiredAlignment));
136
137 ANSEncodedT* outWords = (ANSEncodedT*)(outBlock + 1);
138
139 ANSStateT state = kANSStartState;
140
141 uint32_t inOffset = laneId;
142 uint32_t outOffset = 0;
143
144 constexpr int kUnroll = 8;
145
146 uint32_t limit = roundDown(blockSize, kWarpSize * kUnroll);
147
148 {
149 for (; inOffset < limit; inOffset += kWarpSize * kUnroll) {
150#pragma unroll
151 for (int j = 0; j < kUnroll; ++j) {
152 outOffset +=
153 encodeOne<ProbBits>(state, inBlock[inOffset + j * kWarpSize], outOffset, outWords, smemLookup);
154 }
155 }
156 }
157
158 if (limit != blockSize) {
159 limit = roundDown(blockSize, kWarpSize);
160
161 for (; inOffset < limit; inOffset += kWarpSize) {
162 outOffset +=
163 encodeOne<ProbBits>(state, inBlock[inOffset], outOffset, outWords, smemLookup);
164 }
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);
170 }
171 }
172 // Write final state at the beginning (aligned addresses)
173 outBlock->warpState[laneId] = state;
174
175 if (laneId == 0) {
176 compressedWords_dev[grim_warp_numid] = outOffset;
177 }
178}
179
180template <typename A, int B>
181struct Align {
182 typedef uint32_t argument_type;
183 typedef uint32_t result_type;
184
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;
189
190 return roundUp(x, T(kSize));
191 }
192};
193
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,
204 uint8_t* out_dev,
205 uint32_t* outSize_dev) {
206
207 auto numBlocks = divUp(uncompressedWords, kDefaultBlockSize);
208
209 int block = blockIdx.x;
210 int tid = threadIdx.x;
211
212 ANSCoalescedHeader* headerOut = (ANSCoalescedHeader*)out_dev;
213
214 // The first block will be responsible for the coalesced header
215 if (block == 0) {
216 if (tid == 0) {
217 uint32_t totalCompressedWords = 0;
218
219 if (numBlocks > 0) {
220 totalCompressedWords =
221 compressedWordsPrefix_dev[numBlocks - 1] +
222 roundUp(
223 compressedWords_dev[numBlocks - 1],
224 kBlockAlignment / sizeof(ANSEncodedT));
225 }
226
227 ANSCoalescedHeader header;
228 header.setMagicAndVersion();
229 header.setNumBlocks(numBlocks);
230 header.setTotalUncompressedWords(uncompressedWords);
231 header.setTotalCompressedWords(totalCompressedWords);
232 header.setProbBits(config_probBits);
233
234 if (outSize_dev) {
235 *outSize_dev = header.getTotalCompressedSize();
236 }
237
238 *headerOut = header;
239 }
240
241 auto probsOut = headerOut->getSymbolProbs();
242
243 #pragma unroll
244 for (int i = tid; i < kNumSymbols; i += Threads) {
245 probsOut[i] = table_dev[i].x;
246 }
247 }
248
249 if (block >= numBlocks) {
250 return;
251 }
252
253 // where our per-warp data lies
254 auto uncoalescedBlock = compressedBlocks_dev +
255 block * uncoalescedBlockStride;
256
257 // Write per-block warp state
258 if (tid < kWarpSize) {
259 auto warpStateIn = (ANSWarpState*)uncoalescedBlock;
260
261 headerOut->getWarpStates()[block].warpState[tid] =
262 warpStateIn->warpState[tid];
263 }
264
265 auto blockWordsOut = headerOut->getBlockWords(numBlocks);
266
267 // Write out per-block word length
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;
272
273 uint32_t blockWords =
274 (i == numBlocks - 1) ? lastBlockWords : kDefaultBlockSize;
275
276 blockWordsOut[i] = uint2{
277 (blockWords << 16) | compressedWords_dev[i], compressedWordsPrefix_dev[i]};
278 }
279
280 // Number of compressed words in this block
281 uint32_t numWords = compressedWords_dev[block];
282
283 using LoadT = uint4;
284
285 uint32_t limitEnd = divUp(numWords, kBlockAlignment / sizeof(ANSEncodedT));
286
287 auto inT = (const LoadT*)(uncoalescedBlock + sizeof(ANSWarpState));
288 auto outT =
289 (LoadT*)(headerOut->getBlockDataStart(numBlocks) + compressedWordsPrefix_dev[block]);
290
291 for (uint32_t i = tid; i < limitEnd; i += Threads) {
292 outT[i] = inT[i];
293 }
294}
295
296}} // namespace fz::ans
297
298#endif // FZ_ANS_DIETGPU_ANS_GPUANSENCODE_H
Definition dag.h:24