11#ifndef FZ_ANS_DIETGPU_ANS_GPUANSSTATISTICS_H
12#define FZ_ANS_DIETGPU_ANS_GPUANSSTATISTICS_H
16#include "GpuANSCodec.h"
17#include "utils/DeviceUtils.h"
18#include "utils/PtxUtils.h"
19#include "utils/StaticUtils.h"
25namespace fz {
namespace ans {
31blockSum(
int warpId,
int laneId,
int valForSum,
int* smem) {
32 static_assert(isEvenDivisor(Threads, kWarpSize),
"");
33 constexpr int kWarps = Threads / kWarpSize;
35 auto allSum = warpReduceAllSum(valForSum);
38 smem[warpId] = allSum;
43 int v = laneId < kWarps ? smem[laneId] : 0;
44 v = warpReduceAllSum(v);
61__device__
void normalizeProbabilitiesFromHistogram(
63 const uint32_t* __restrict__ counts,
66 uint4* __restrict__ table) {
69 kNumSymbols == Threads || isEvenDivisor(kNumSymbols, uint32_t(Threads)),
72 constexpr int kNumSymPerThread =
73 kNumSymbols == Threads ? 1 : (kNumSymbols / Threads);
80 constexpr int kWarps = Threads / kWarpSize;
81 uint32_t kProbWeight = 1 << probBits;
82 int tid = threadIdx.x;
83 int warpId = tid / kWarpSize;
84 int laneId = getLaneId();
88 uint32_t qProb[kNumSymPerThread];
93 for (
int i = 0; i < kNumSymPerThread; ++i) {
94 int curSym = i * Threads + tid;
95 uint32_t count = counts[curSym];
98 qProb[i] = kProbWeight * ((float)count / (
float)totalNum);
101 qProb[i] = (count > 0 && qProb[i] == 0) ? 1 : qProb[i];
103 qProbSum += qProb[i];
107 __shared__
int smemSum[kWarps];
108 qProbSum = blockSum<Threads>(warpId, laneId, qProbSum, smemSum);
112 uint32_t sortedPair[kNumSymPerThread];
115 for (
int i = 0; i < kNumSymPerThread; ++i) {
116 int curSym = i * Threads + tid;
117 sortedPair[i] = (qProb[i] << 16) | curSym;
120 using Sort = cub::BlockRadixSort<uint32_t, Threads, kNumSymPerThread>;
121 __shared__
typename Sort::TempStorage smemSort;
122 Sort(smemSort).SortDescending(sortedPair);
124 uint32_t tidSymbol[kNumSymPerThread];
127 for (
int i = 0; i < kNumSymPerThread; ++i) {
128 tidSymbol[i] = sortedPair[i] & 0xffffU;
129 qProb[i] = sortedPair[i] >> 16;
134 int diff = (int)kProbWeight - (
int)qProbSum;
138 int iterToApply = diff < kNumSymbols ? diff : kNumSymbols;
141 for (
int i = 0; i < kNumSymPerThread; ++i) {
142 int curSym = tidSymbol[i];
143 if (curSym < iterToApply) {
150 }
else if (diff < 0) {
157 for (
int i = 0; i < kNumSymPerThread; ++i) {
158 qNumGt1s += (int)(qProb[i] > 1);
161 qNumGt1s = blockSum<Threads>(warpId, laneId, qNumGt1s, smemSum);
164 int iterToApply = diff < qNumGt1s ? diff : qNumGt1s;
165 assert(iterToApply > 0);
166 int startIndex = qNumGt1s - iterToApply;
169 for (
int i = 0; i < kNumSymPerThread; ++i) {
170 int curSym = tid * kNumSymPerThread + i;
171 if (curSym >= startIndex && curSym < qNumGt1s) {
182 __shared__ uint32_t smemPdf[kNumSymbols];
185 for (
int i = 0; i < kNumSymPerThread; ++i) {
186 smemPdf[tidSymbol[i]] = qProb[i];
191 uint32_t symPdf[kNumSymPerThread];
193 for (
int i = 0; i < kNumSymPerThread; ++i) {
194 int curSym = tid * kNumSymPerThread + i;
195 symPdf[i] = smemPdf[curSym];
198 using Scan = cub::BlockScan<uint32_t, Threads>;
199 __shared__
typename Scan::TempStorage smemScan;
201 uint32_t symCdf[kNumSymPerThread];
202 Scan(smemScan).ExclusiveSum(symPdf, symCdf);
206 uint32_t shift[kNumSymPerThread];
207 uint32_t magic[kNumSymPerThread];
210 for (
int i = 0; i < kNumSymPerThread; ++i) {
211 shift[i] = 32 - __clz(symPdf[i] - 1);
212 constexpr uint64_t one = 1;
213 uint64_t magic64 = 0;
216 ((one << 32) * ((one << shift[i]) - symPdf[i])) / symPdf[i] + 1;
217 magic[i] = (uint32_t)magic64;
221 for (
int i = 0; i < kNumSymPerThread; ++i) {
222 int curSym = tid * kNumSymPerThread + i;
223 table[curSym] = uint4{symPdf[i], symCdf[i], magic[i], shift[i]};
227template <
int Threads>
228__global__
void quantizeWeights(
229 const uint32_t* __restrict__ counts,
232 uint4* __restrict__ table) {
233 normalizeProbabilitiesFromHistogram<Threads>(
240inline void ansCalcWeights(
243 const uint32_t* histogram_dev,
245 cudaStream_t stream) {
246 constexpr int kThreads = kNumSymbols;
247 quantizeWeights<kThreads><<<1, kThreads, 0, stream>>>(
248 histogram_dev, inSize_dev, probBits, table_dev);