Files
Jungfraujoch/image_analysis/image_preprocessing/BSLZ4DecodeWarp.h
T
jungfrauandClaude Opus 5 cf2336e523 Decode a compressed frame into shared memory and preprocess it there
The image loop on a 16 Mpx detector is two thirds of the run and both cards are busy for essentially
all of it, so card time removed is wall time removed. Of the six milliseconds a frame costs, two and
a half were spent decompressing it - and not because the card was short of bandwidth. The LZ4 pass
moved 53 GB/s where the strong-pixel flagger, reading the same image and the same bin table, gets
276. It is latency, not bandwidth: the copy loop moves 32 bytes per warp iteration with a syncwarp
after each one, and for a match copy the source and the destination both derive from the same
pointer, so nothing pipelines. The warp spends its time waiting for global memory, one dependent
round trip at a time.

So decode where the waiting is cheap. One CUDA block now owns one bitshuffle block: its first warp
decodes the payload into shared memory, and the whole block then un-transposes and preprocesses out
of shared and writes finished pixels. A shared round trip is tens of cycles rather than hundreds,
and the 72 MB shuffled intermediate never reaches DRAM at all - the pair of kernels moved about 238
MB a frame and the fused one moves 93.

The parser is lifted into a device function that both kernels call over the same bytes, so the
standalone path and the fused one cannot decode a chunk differently. The statistics reduction had to
change with it: 48 bytes of static shared on top of a full bitshuffle block costs a whole resident
block per multiprocessor, so the counts now reduce through a warp shuffle and one integer atomic per
warp. Blocks larger than 16 kB keep the two-kernel path, and the beam stop's own decoder is
untouched.

What this costs is decoder parallelism: a block that holds 16 kB of shared is one of four resident
per multiprocessor on this card, where the old kernel fitted thirty-two warps each decoding on its
own. The trade is favourable here and should be better on the production cards, which have half
again as much shared memory per multiprocessor. Measured on a 16 Mpx rotation set at the production
GPU count, with the indexing work of the next commit: 37.2 s -> 32.9 s, and the merged output is
byte-identical.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_011n8riB6X59oRjkrSHzNPAU
2026-08-23 14:03:40 -04:00

129 lines
6.8 KiB
C++

// SPDX-FileCopyrightText: 2026 Filip Leonarski, Paul Scherrer Institute <filip.leonarski@psi.ch>
// SPDX-License-Identifier: GPL-3.0-only
#pragma once
#include <cstdint>
// The LZ4 block parser, as a device function, so the standalone decode kernel and the fused
// decode+un-transpose+preprocess kernel run exactly the same code over the same bytes. The two
// differ only in where the output lands - device memory for the first, a shared-memory staging
// buffer for the second - and a generic pointer covers both, so there is one parser and no way for
// the two paths to disagree on what a chunk decodes to.
//
// CUDA only: include it from a .cu, never from a header a .cpp sees.
// One WARP per LZ4 block. Every lane runs the same sequence parser over the same bytes - a
// broadcast read, so no divergence - and the literal and match copies are split across the 32
// lanes so the stores coalesce. One thread per block instead has each thread streaming its own
// 8 kB region, which coalesces not at all and measured 13x slower.
//
// Because the lanes cooperate on the copies, a match can source bytes that OTHER lanes wrote in
// an earlier sequence. Since Volta that needs an explicit __syncwarp() - implicit reconvergence
// is not part of the programming model - so there is one after every copy loop. The full mask is
// correct: the early return and every break test warp-uniform values, so lanes never diverge
// permanently.
//
// Bounds: every read is clamped against iend and every write against oend, so a malformed or
// corrupt payload cannot walk off either buffer. It can still stop early, which leaves the block
// short; that is what the false return reports.
__device__ __forceinline__ bool lz4_decode_block_warp(const uint8_t *ip, const uint8_t *const iend,
uint8_t *const obase, uint8_t *const oend,
int lane) {
uint8_t *op = obase;
bool malformed = false;
while (ip < iend) {
const uint32_t token = *ip++;
uint32_t litlen = token >> 4;
if (litlen == 15) {
// read_variable_length(&ip, iend - RUN_MASK, initial_check=1) in the reference: the
// chain may not start within, nor run into, the last RUN_MASK (15) input bytes. A
// valid stream never does - the literals it counts have to follow it - so a chain
// that reaches there is corruption, and this is the only place it shows up.
if ((size_t)(iend - ip) <= 15) { malformed = true; break; }
uint32_t s;
do {
s = *ip++;
litlen += s;
if ((size_t)(iend - ip) < 15) { malformed = true; break; }
} while (s == 255);
if (malformed) break;
}
if (litlen) {
// Clamped by the INPUT as well as the output: a corrupt litlen must not read past the
// end of this block's payload or write past the end of the block. Clamping keeps the
// kernel in bounds; needing to clamp at all means the stream is not decodable, which
// is what the reference reports as an error, so record it.
if (litlen > (uint32_t)(oend - op) || litlen > (uint32_t)(iend - ip))
malformed = true;
const uint32_t n = min(min(litlen, (uint32_t)(oend - op)), (uint32_t)(iend - ip));
for (uint32_t i = lane; i < n; i += 32) op[i] = ip[i];
__syncwarp();
op += n; ip += litlen;
}
// LZ4's parsing restrictions: an encoder may not leave a match within MFLIMIT (12) bytes
// of the end of the block, nor fewer than 2+1+LASTLITERALS (8) input bytes after a
// literal run that is not the last one. So once either limit is reached this can ONLY be
// the final sequence, and the final sequence must consume the payload exactly. The
// reference applies this whether or not the run was empty, which is why the test sits
// outside the copy - a zero-length literal run near the end is just as illegal.
if ((size_t)(oend - op) < 12 || (size_t)(iend - ip) < 8) {
malformed = (ip != iend) || (op != oend);
break; // necessarily EOF
}
if (iend - ip < 2) break; // last sequence carries literals only
const uint32_t offset = (uint32_t)ip[0] | ((uint32_t)ip[1] << 8);
ip += 2;
uint32_t matchlen = token & 0x0F;
if (matchlen == 15) {
// read_variable_length(&ip, iend - LASTLITERALS + 1, initial_check=0): bounded by the
// last 4 input bytes rather than 15, and with no check before the first read.
uint32_t s;
do {
s = *ip++;
matchlen += s;
if ((size_t)(iend - ip) < 4) { malformed = true; break; }
} while (s == 255);
if (malformed) break;
}
matchlen += 4; // minmatch
if (offset == 0 || offset > (uint32_t)(op - obase)) { malformed = true; break; }
const uint8_t *mp = op - offset;
// A match may reach the end of the block but never past it - the reference treats an
// overrun as an error rather than truncating, and so must this.
if (matchlen > (uint32_t)(oend - op))
malformed = true;
const uint32_t n = min(matchlen, (uint32_t)(oend - op));
if (offset >= matchlen) {
for (uint32_t i = lane; i < n; i += 32) op[i] = mp[i];
} else {
// An overlapping match is a pattern of period `offset`. mp[0..offset-1] all lie
// before op and are already final, so each output byte can be sourced from them
// independently - which keeps this parallel rather than a serial byte loop. Long
// zero runs in sparse detector data arrive here with offset == 1, and a runtime
// modulo is an emulated division, so the two cheap cases are peeled off first.
if (offset == 1) {
const uint8_t v = mp[0];
for (uint32_t i = lane; i < n; i += 32) op[i] = v;
} else if ((offset & (offset - 1)) == 0) {
const uint32_t m = offset - 1;
for (uint32_t i = lane; i < n; i += 32) op[i] = mp[i & m];
} else {
for (uint32_t i = lane; i < n; i += 32) op[i] = mp[i % offset];
}
}
__syncwarp();
op += n;
}
// A block must decode to exactly its declared length AND consume exactly its payload. Both
// are conditions LZ4_decompress_safe reports to the host path, and both are needed: a corrupt
// stream can land on the right output length while leaving input over, or run its input out
// early. Either way the bytes are not the ones that were compressed.
return !malformed && op == oend && ip == iend;
}