// SPDX-FileCopyrightText: 2026 Filip Leonarski, Paul Scherrer Institute // SPDX-License-Identifier: GPL-3.0-only #pragma once #include // 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; }