Skip to content

Commit 147472b

Browse files
committed
Add exact encoded-size oracle for plain/outlier AdaptiveBitpack
Plain/outlier AdaptiveBitpack stages declare their local coder policy; finalize() binds compatible semantics to directly connected AdaptiveLorenzo stages independently of fusion policy. The staged selector now evaluates shared exact size quotes from full/rest/element-0 block statistics instead of a private plain-byte formula.
1 parent 658889f commit 147472b

6 files changed

Lines changed: 338 additions & 35 deletions

File tree

Lines changed: 80 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,80 @@
1+
#pragma once
2+
3+
/**
4+
* @file adaptive_bitpack_oracle.cuh
5+
* @brief Exact, device-inline encoding quotes shared by staged and fused selectors.
6+
*/
7+
8+
#include <cstdint>
9+
10+
namespace fz::adaptive_bitpack_oracle {
11+
12+
struct PlainFixedRateDecision {
13+
uint8_t rate = 0;
14+
};
15+
16+
struct PlainFixedRateQuote {
17+
uint32_t payload_bytes = 0;
18+
PlainFixedRateDecision decision{};
19+
};
20+
21+
struct AdaptiveFixedRateDecision {
22+
uint8_t rate = 0;
23+
uint8_t selector = 0; ///< bit0=outlier; bits1-2=outlier bytes minus one.
24+
};
25+
26+
struct AdaptiveFixedRateQuote {
27+
uint32_t payload_bytes = 0;
28+
AdaptiveFixedRateDecision decision{};
29+
};
30+
31+
/// Exact payload emitted by plain AdaptiveBitpack for one coder unit: an all-zero
32+
/// unit emits nothing; otherwise it emits one sign word and `rate` plane words.
33+
__device__ __forceinline__ PlainFixedRateQuote quotePlainFixedRate(
34+
uint32_t max_magnitude, uint32_t word_bytes)
35+
{
36+
const uint8_t rate = max_magnitude
37+
? static_cast<uint8_t>(32 - __clz(max_magnitude))
38+
: uint8_t{0};
39+
PlainFixedRateQuote q;
40+
q.decision.rate = rate;
41+
q.payload_bytes = rate > 0
42+
? word_bytes * (static_cast<uint32_t>(rate) + 1u)
43+
: 0u;
44+
return q;
45+
}
46+
47+
/// Exact plain-versus-element0-outlier choice used by AdaptiveBitpack's outlier
48+
/// mode. Ties select plain, matching the staged and warp-register encoders.
49+
__device__ __forceinline__ AdaptiveFixedRateQuote quoteAdaptiveFixedRate(
50+
uint32_t max_magnitude, uint32_t max_rest_magnitude,
51+
uint32_t first_magnitude, uint32_t word_bytes)
52+
{
53+
const PlainFixedRateQuote plain = quotePlainFixedRate(max_magnitude, word_bytes);
54+
const uint8_t rest_rate = max_rest_magnitude
55+
? static_cast<uint8_t>(32 - __clz(max_rest_magnitude))
56+
: uint8_t{0};
57+
const uint32_t outlier_bytes = first_magnitude
58+
? static_cast<uint32_t>((32 - __clz(first_magnitude) + 7) / 8)
59+
: 0u;
60+
const uint32_t outlier_payload = outlier_bytes +
61+
(rest_rate > 0
62+
? word_bytes * (static_cast<uint32_t>(rest_rate) + 1u)
63+
: word_bytes);
64+
65+
AdaptiveFixedRateQuote q;
66+
if (plain.payload_bytes <= outlier_payload) {
67+
q.payload_bytes = plain.payload_bytes;
68+
q.decision.rate = plain.decision.rate;
69+
q.decision.selector = 0;
70+
} else {
71+
q.payload_bytes = outlier_payload;
72+
q.decision.rate = rest_rate;
73+
// outlier_payload can beat plain only when outlier_bytes is nonzero.
74+
q.decision.selector = static_cast<uint8_t>(
75+
1u | ((outlier_bytes - 1u) << 1u));
76+
}
77+
return q;
78+
}
79+
80+
} // namespace fz::adaptive_bitpack_oracle

‎modules/coders/adaptive_bitpack/adaptive_bitpack_stage.h‎

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -123,6 +123,25 @@ class AdaptiveBitpackStage : public Stage {
123123
return d;
124124
}
125125

126+
/// Exact local encoded-size policy for an immediately upstream adaptive
127+
/// selector. Unlike getFusedOp(), this semantic declaration is available in
128+
/// plain mode: AdaptiveLorenzo needs the staged coder's exact byte formula
129+
/// even when no fused implementation has been selected.
130+
EncodingOracleDecl getEncodingOracle() const override {
131+
if (is_inverse_ || block_size_ == 0) return {};
132+
EncodingOracleDecl d;
133+
d.kind = outlier_selection_
134+
? EncodingOracleKind::AdaptiveFixedRateBitpack
135+
: EncodingOracleKind::PlainFixedRateBitpack;
136+
d.op_name = outlier_selection_ ? "AdaptiveBitpackCoder" : "PlainBitpackCoder";
137+
d.include_header = "fused/fused_block/warp_fusion.cuh";
138+
d.input_data_type = static_cast<uint8_t>(getElementDataType());
139+
d.unit_elems = block_size_;
140+
d.exact = true;
141+
d.additive = true;
142+
return d;
143+
}
144+
126145
/// Choose the warp coder policy the fused path composes (default
127146
/// "AdaptiveBitpackCoder"). Any policy in warp_fusion.cuh that emits an
128147
/// AdaptiveBitpack-decodable archive works — e.g. "PlainBitpackCoder" for an
@@ -139,6 +158,9 @@ class AdaptiveBitpackStage : public Stage {
139158
num_elements_ = num_elements;
140159
actual_output_size_ = archive_bytes;
141160
}
161+
void setFusedArchiveResult(size_t archive_bytes, size_t orig_bytes) override {
162+
setFusedResult(orig_bytes / sizeof(T), archive_bytes);
163+
}
142164

143165
/// Enable cuSZp2 per-block plain/outlier selection: each block may instead
144166
/// store element 0 as a raw 1..sizeof(T)-byte outlier and pack only the rest,

‎modules/fused/adaptive_lorenzo/adaptive_lorenzo_stage.cu‎

Lines changed: 76 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
*/
2424

2525
#include "fused/adaptive_lorenzo/adaptive_lorenzo_stage.h"
26+
#include "coders/adaptive_bitpack/adaptive_bitpack_oracle.cuh"
2627
#include "stage/stage_registry.h"
2728
#include <cstring>
2829
#include <algorithm>
@@ -56,14 +57,33 @@ __device__ __forceinline__ typename std::make_unsigned<T>::type absU(T v) {
5657
return (v < 0) ? static_cast<U>(~uv + static_cast<U>(1)) : uv;
5758
}
5859

59-
__device__ __forceinline__ int bitWidth32(uint32_t x) {
60-
return x ? (32 - __clz(x)) : 0;
61-
}
62-
63-
// AdaptiveBitpackStage's per-block encoded size: nothing when every residual is
64-
// zero, else a sign bitmap plus `r` bit-planes, one 32-bit word each.
65-
__device__ __forceinline__ uint32_t blockCost(int r) {
66-
return (r > 0) ? 4u * (static_cast<uint32_t>(r) + 1u) : 0u;
60+
struct CoderStats {
61+
uint32_t all;
62+
uint32_t rest;
63+
uint32_t first;
64+
};
65+
66+
// Exact policy dispatch shared by staged selection and the future generated
67+
// tile harness. The quote carries the coder decision even though the staged
68+
// predictor needs only its payload length; the downstream staged coder will
69+
// recompute the same decision from the selected residuals.
70+
__device__ __forceinline__ uint32_t blockCost(
71+
CoderStats stats, EncodingOracleKind oracle_kind)
72+
{
73+
switch (oracle_kind) {
74+
case EncodingOracleKind::AdaptiveFixedRateBitpack:
75+
return adaptive_bitpack_oracle::quoteAdaptiveFixedRate(
76+
stats.all, stats.rest, stats.first,
77+
/*word_bytes=*/4u).payload_bytes;
78+
case EncodingOracleKind::PlainFixedRateBitpack:
79+
return adaptive_bitpack_oracle::quotePlainFixedRate(
80+
stats.all, /*word_bytes=*/4u).payload_bytes;
81+
default:
82+
// Standalone/backward-compatible AdaptiveLorenzo uses the legacy
83+
// plain policy. Unsupported downstream policies are not bound.
84+
return adaptive_bitpack_oracle::quotePlainFixedRate(
85+
stats.all, /*word_bytes=*/4u).payload_bytes;
86+
}
6787
}
6888

6989
// Modes are packed 2 bits per tile (4 tiles per byte): at one byte per tile the
@@ -131,14 +151,19 @@ __global__ void adaptive_lorenzo_forward_kernel(
131151
size_t n,
132152
uint32_t tile_size,
133153
bool enable_order2,
134-
bool enable_centering)
154+
bool enable_centering,
155+
EncodingOracleKind oracle_kind)
135156
{
136157
// All shared state is now per-WARP, not per-element: nwarps <= 32 because a
137158
// tile is at most 32 coder blocks of 32. The tile-sized `s` staging buffer and
138159
// the tile-sized `red` reduction buffer are both gone, so this kernel needs no
139160
// dynamic shared memory at all (see the launch site).
140-
__shared__ uint32_t acc1[kMaxBlocksPerTile]; // per coder block, LZ1
141-
__shared__ uint32_t acc2[kMaxBlocksPerTile]; // per coder block, LZ2
161+
__shared__ uint32_t acc1[kMaxBlocksPerTile]; // all magnitudes, LZ1
162+
__shared__ uint32_t acc2[kMaxBlocksPerTile]; // all magnitudes, LZ2
163+
__shared__ uint32_t rest1[kMaxBlocksPerTile]; // lanes 1..31, LZ1
164+
__shared__ uint32_t rest2[kMaxBlocksPerTile]; // lanes 1..31, LZ2
165+
__shared__ uint32_t first1[kMaxBlocksPerTile]; // lane 0 magnitude, LZ1
166+
__shared__ uint32_t first2[kMaxBlocksPerTile]; // lane 0 magnitude, LZ2
142167
__shared__ long long red[kMaxBlocksPerTile]; // per-warp partial sums (mean)
143168
__shared__ T sb_last[kMaxBlocksPerTile]; // v at lane 31 of each warp
144169
__shared__ T sb_prev[kMaxBlocksPerTile]; // v at lane 30 of each warp
@@ -223,29 +248,42 @@ __global__ void adaptive_lorenzo_forward_kernel(
223248
const T q0 = s_q0;
224249

225250
// ---- Per-coder-block magnitudes, uncentered ----
226-
const uint32_t o1 = warpOr(live ? static_cast<uint32_t>(absU<T>(d1)) : 0u);
227-
const uint32_t o2 = warpOr(live ? static_cast<uint32_t>(absU<T>(d2)) : 0u);
228-
if (lane == 0) { acc1[warp] = o1; acc2[warp] = o2; }
251+
const uint32_t m1 = live ? static_cast<uint32_t>(absU<T>(d1)) : 0u;
252+
const uint32_t m2 = live ? static_cast<uint32_t>(absU<T>(d2)) : 0u;
253+
const uint32_t o1 = warpOr(m1);
254+
const uint32_t o2 = warpOr(m2);
255+
const uint32_t r1rest = warpOr(lane > 0u ? m1 : 0u);
256+
const uint32_t r2rest = warpOr(lane > 0u ? m2 : 0u);
257+
if (lane == 0) {
258+
acc1[warp] = o1; rest1[warp] = r1rest; first1[warp] = m1;
259+
acc2[warp] = o2; rest2[warp] = r2rest; first2[warp] = m2;
260+
}
229261

230262
// ---- Centered variants: only coder block 0 can differ ----
231-
uint32_t acc1c0 = 0u, acc2c0 = 0u;
263+
CoderStats c1stats{0u, 0u, 0u}, c2stats{0u, 0u, 0u};
232264
if (enable_centering && warp == 0u) {
233265
const T c0 = static_cast<T>(q0 - mu);
234-
const T r1 = (tid == 0u) ? c0 : d1;
235-
T r2 = d2;
236-
if (tid == 0u) r2 = c0;
237-
else if (tid == 1u) r2 = static_cast<T>(d1 - c0);
238-
acc1c0 = warpOr(live ? static_cast<uint32_t>(absU<T>(r1)) : 0u);
239-
acc2c0 = warpOr(live ? static_cast<uint32_t>(absU<T>(r2)) : 0u);
266+
const T cr1 = (tid == 0u) ? c0 : d1;
267+
T cr2 = d2;
268+
if (tid == 0u) cr2 = c0;
269+
else if (tid == 1u) cr2 = static_cast<T>(d1 - c0);
270+
const uint32_t cm1 = live ? static_cast<uint32_t>(absU<T>(cr1)) : 0u;
271+
const uint32_t cm2 = live ? static_cast<uint32_t>(absU<T>(cr2)) : 0u;
272+
c1stats.all = warpOr(cm1);
273+
c1stats.rest = warpOr(lane > 0u ? cm1 : 0u);
274+
c1stats.first = fz::backend::shfl(cm1, 0, 32);
275+
c2stats.all = warpOr(cm2);
276+
c2stats.rest = warpOr(lane > 0u ? cm2 : 0u);
277+
c2stats.first = fz::backend::shfl(cm2, 0, 32);
240278
}
241279
__syncthreads();
242280

243281
// ---- Cost each variant, pick the cheapest ----
244282
if (tid == 0) {
245283
uint32_t c_lz1 = 0, c_lz2 = 0;
246284
for (unsigned w = 0; w < nwarps; ++w) {
247-
c_lz1 += blockCost(bitWidth32(acc1[w]));
248-
c_lz2 += blockCost(bitWidth32(acc2[w]));
285+
c_lz1 += blockCost(CoderStats{acc1[w], rest1[w], first1[w]}, oracle_kind);
286+
c_lz2 += blockCost(CoderStats{acc2[w], rest2[w], first2[w]}, oracle_kind);
249287
}
250288
uint32_t costs[4];
251289
costs[0] = c_lz1;
@@ -259,11 +297,13 @@ __global__ void adaptive_lorenzo_forward_kernel(
259297
// genuinely saves those bytes. (The 2-bit mode is charged to nobody
260298
// because every tile pays it regardless of what it picks.)
261299
const uint32_t mean_cost = static_cast<uint32_t>(sizeof(T));
262-
costs[2] = c_lz1 - blockCost(bitWidth32(acc1[0]))
263-
+ blockCost(bitWidth32(acc1c0)) + mean_cost;
300+
costs[2] = c_lz1
301+
- blockCost(CoderStats{acc1[0], rest1[0], first1[0]}, oracle_kind)
302+
+ blockCost(c1stats, oracle_kind) + mean_cost;
264303
if (enable_order2)
265-
costs[3] = c_lz2 - blockCost(bitWidth32(acc2[0]))
266-
+ blockCost(bitWidth32(acc2c0)) + mean_cost;
304+
costs[3] = c_lz2
305+
- blockCost(CoderStats{acc2[0], rest2[0], first2[0]}, oracle_kind)
306+
+ blockCost(c2stats, oracle_kind) + mean_cost;
267307
}
268308
uint32_t best = 0;
269309
for (uint32_t i = 1; i < 4; ++i)
@@ -384,15 +424,16 @@ template<typename T>
384424
void launchAdaptiveLorenzoForward(
385425
const T* d_input, T* d_residuals, uint8_t* d_modes_dense, T* d_means_dense,
386426
uint32_t* d_flags, size_t n, uint32_t tile_size,
387-
bool enable_order2, bool enable_centering, cudaStream_t stream)
427+
bool enable_order2, bool enable_centering, EncodingOracleKind oracle_kind,
428+
cudaStream_t stream)
388429
{
389430
if (n == 0) return;
390431
const int grid = static_cast<int>((n + tile_size - 1) / tile_size);
391432
// No dynamic shared memory: the kernel's shared state is now per-warp
392433
// (<= 32 entries each) and lives in static __shared__ arrays.
393434
adaptive_lorenzo_forward_kernel<T><<<grid, tile_size, 0, stream>>>(
394435
d_input, d_residuals, d_modes_dense, d_means_dense, d_flags, n, tile_size,
395-
enable_order2, enable_centering);
436+
enable_order2, enable_centering, oracle_kind);
396437
FZ_CUDA_CHECK(cudaGetLastError());
397438
}
398439

@@ -526,7 +567,8 @@ void AdaptiveLorenzoStage<T>::execute(
526567
launchAdaptiveLorenzoForward<T>(
527568
static_cast<const T*>(inputs[0]), static_cast<T*>(outputs[0]),
528569
d_modes_dense_, d_means_dense_, d_flags_,
529-
n, tile, config_.enable_order2, config_.enable_centering, stream);
570+
n, tile, config_.enable_order2, config_.enable_centering,
571+
getBoundEncodingOracleKind(), stream);
530572

531573
// flags[tiles] = 0 so offsets[tiles] lands on the total centered count.
532574
FZ_CUDA_CHECK(cudaMemsetAsync(d_flags_ + tiles, 0, sizeof(uint32_t), stream));
@@ -600,9 +642,11 @@ template class AdaptiveLorenzoStage<int16_t>;
600642
template class AdaptiveLorenzoStage<int32_t>;
601643

602644
template void launchAdaptiveLorenzoForward<int16_t>(
603-
const int16_t*, int16_t*, uint8_t*, int16_t*, uint32_t*, size_t, uint32_t, bool, bool, cudaStream_t);
645+
const int16_t*, int16_t*, uint8_t*, int16_t*, uint32_t*, size_t, uint32_t,
646+
bool, bool, EncodingOracleKind, cudaStream_t);
604647
template void launchAdaptiveLorenzoForward<int32_t>(
605-
const int32_t*, int32_t*, uint8_t*, int32_t*, uint32_t*, size_t, uint32_t, bool, bool, cudaStream_t);
648+
const int32_t*, int32_t*, uint8_t*, int32_t*, uint32_t*, size_t, uint32_t,
649+
bool, bool, EncodingOracleKind, cudaStream_t);
606650

607651
template void launchAdaptiveLorenzoCompact<int16_t>(
608652
const uint8_t*, const int16_t*, const uint32_t*, uint8_t*, int16_t*, size_t, cudaStream_t);

‎modules/fused/adaptive_lorenzo/adaptive_lorenzo_stage.h‎

Lines changed: 50 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -122,6 +122,46 @@ class AdaptiveLorenzoStage : public Stage {
122122
return config_.coder_block_size * config_.blocks_per_tile;
123123
}
124124

125+
/// Bind an exact additive fixed-rate policy exposed by a directly connected
126+
/// AdaptiveBitpack stage. Both plain and element-0-outlier selection operate
127+
/// independently on the same 32-element coder units.
128+
bool bindDownstreamEncodingOracle(const EncodingOracleDecl& decl) override {
129+
if (!decl.valid() || !decl.additive ||
130+
(decl.kind != EncodingOracleKind::PlainFixedRateBitpack &&
131+
decl.kind != EncodingOracleKind::AdaptiveFixedRateBitpack) ||
132+
decl.input_data_type != static_cast<uint8_t>(getElementDataType()) ||
133+
decl.unit_elems != config_.coder_block_size) {
134+
return false;
135+
}
136+
bound_oracle_ = decl;
137+
has_bound_oracle_ = true;
138+
return true;
139+
}
140+
141+
bool hasBoundEncodingOracle() const { return has_bound_oracle_; }
142+
EncodingOracleKind getBoundEncodingOracleKind() const {
143+
return has_bound_oracle_ ? bound_oracle_.kind
144+
: EncodingOracleKind::PlainFixedRateBitpack;
145+
}
146+
147+
FusionSpec getFusionSpec() const override {
148+
if (is_inverse_ || !has_bound_oracle_) return {};
149+
return FusionSpec{FusionAccess::TileAdaptive, getTileSize(),
150+
config_.coder_block_size};
151+
}
152+
153+
std::vector<FusedAuxOutputDecl> getFusedAuxOutputs() const override {
154+
if (!getFusionSpec().fusable()) return {};
155+
return {
156+
FusedAuxOutputDecl{1, "modes", FusedAuxSizeKind::FixedBitsPerUnit,
157+
static_cast<uint8_t>(DataType::UINT8), getTileSize(),
158+
2u, 0u},
159+
FusedAuxOutputDecl{2, "means", FusedAuxSizeKind::CompactedElements,
160+
static_cast<uint8_t>(getElementDataType()), getTileSize(),
161+
0u, 1u},
162+
};
163+
}
164+
125165
void execute(
126166
fz::stream_t stream,
127167
MemoryPool* pool,
@@ -185,6 +225,12 @@ class AdaptiveLorenzoStage : public Stage {
185225
? actual_output_sizes_[index] : 0;
186226
}
187227

228+
void setFusedSideOutput(int output_index, size_t bytes) override {
229+
if (actual_output_sizes_.size() < 3) actual_output_sizes_.resize(3, 0);
230+
if (output_index == 1 || output_index == 2)
231+
actual_output_sizes_[static_cast<size_t>(output_index)] = bytes;
232+
}
233+
188234
uint16_t getStageTypeId() const override {
189235
return static_cast<uint16_t>(StageType::ADAPTIVE_LORENZO);
190236
}
@@ -232,6 +278,8 @@ class AdaptiveLorenzoStage : public Stage {
232278

233279
private:
234280
Config config_;
281+
EncodingOracleDecl bound_oracle_;
282+
bool has_bound_oracle_ = false;
235283
bool is_inverse_ = false;
236284
size_t num_elements_ = 0;
237285
std::vector<size_t> actual_output_sizes_{0, 0, 0};
@@ -283,8 +331,8 @@ extern template class AdaptiveLorenzoStage<int32_t>;
283331
template<typename T>
284332
void launchAdaptiveLorenzoForward(
285333
const T* d_input, T* d_residuals, uint8_t* d_modes, T* d_means,
286-
size_t n, uint32_t tile_size, bool enable_order2, bool enable_centering,
287-
fz::stream_t stream);
334+
uint32_t* d_flags, size_t n, uint32_t tile_size, bool enable_order2,
335+
bool enable_centering, EncodingOracleKind oracle_kind, fz::stream_t stream);
288336

289337
/// Inverse: replay each tile's recorded variant.
290338
template<typename T>

0 commit comments

Comments
 (0)