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>
384424void 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>;
600642template class AdaptiveLorenzoStage <int32_t >;
601643
602644template 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);
604647template 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
607651template void launchAdaptiveLorenzoCompact<int16_t >(
608652 const uint8_t *, const int16_t *, const uint32_t *, uint8_t *, int16_t *, size_t , cudaStream_t);
0 commit comments