FZGPUModules 2.0
GPU-accelerated modular compression pipelines
Loading...
Searching...
No Matches
bitshuffle_stage.h
Go to the documentation of this file.
1#pragma once
2
19#include "stage/stage.h"
20#include "fzm_format.h"
21#include "backend/types.h"
22#include <cstdint>
23#include <cstring>
24#include <stdexcept>
25#include <string>
26#include <unordered_map>
27#include <vector>
28
29namespace fz {
30
43class BitshuffleStage : public Stage {
44public:
46 : is_inverse_(false)
47 , block_size_(16384)
48 , element_width_(4)
49 , actual_output_size_(0)
50 {}
51
52 // ── Stage control ──────────────────────────────────────────────────────
53 void setInverse(bool inv) override { is_inverse_ = inv; }
54
55 // Fixed-length region-local transform at the chunk granularity.
56 // The fused Bitshuffle32<ChunkBytes> op covers the 4-byte-element shape at
57 // any chunk_size the fusion harness supports — see chunk_geometry.h's
58 // kSupportedChunkBytes.
59 FusionSpec getFusionSpec() const override {
60 if (is_inverse_ || element_width_ != 4 ||
61 (block_size_ != 4096u && block_size_ != 8192u && block_size_ != 16384u)) return {};
62 return FusionSpec{FusionAccess::RegionLocal, block_size_};
63 }
64
66 FusedOpDecl getFusedOp() const override {
67 if (!getFusionSpec().fusable()) return {};
68 return FusedOpDecl{FusionStrategy::ChunkCooperative, "Bitshuffle32",
69 "fused/chunk_fusion/chunk_fusion.cuh", {}};
70 }
71 bool isInverse() const override { return is_inverse_; }
72
73 void setBlockSize(size_t bytes) { block_size_ = static_cast<uint32_t>(bytes); }
74 void setElementWidth(size_t bytes){ element_width_ = static_cast<uint8_t>(bytes); }
75
76 size_t getBlockSize() const { return block_size_; }
77 size_t getRequiredInputAlignment() const override { return block_size_; }
78 size_t getElementWidth() const { return element_width_; }
79
80 // ── Execution ──────────────────────────────────────────────────────────
81 void execute(
82 fz::stream_t stream,
83 MemoryPool* pool,
84 const std::vector<void*>& inputs,
85 const std::vector<void*>& outputs,
86 const std::vector<size_t>& sizes
87 ) override;
88
89 // ── Metadata ───────────────────────────────────────────────────────────
90 std::string getName() const override { return "Bitshuffle"; }
91 size_t getNumInputs() const override { return 1; }
92 size_t getNumOutputs() const override { return 1; }
93
94 std::vector<size_t> estimateOutputSizes(
95 const std::vector<size_t>& input_sizes
96 ) const override {
97 // Size-preserving transform.
98 return {input_sizes[0]};
99 }
100
101 std::unordered_map<std::string, size_t>
102 getActualOutputSizesByName() const override {
103 return {{"output", actual_output_size_}};
104 }
105 size_t getActualOutputSize(int index) const override {
106 return (index == 0) ? actual_output_size_ : 0;
107 }
108
109 uint16_t getStageTypeId() const override {
110 return static_cast<uint16_t>(StageType::BITSHUFFLE);
111 }
112
113 uint8_t getOutputDataType(size_t) const override {
114 // Raw byte stream — report as UINT8.
115 return static_cast<uint8_t>(DataType::UINT8);
116 }
117
118 // ── Serialization ──────────────────────────────────────────────────────
119 // Header: [0..3] block_size (uint32_t LE), [4] element_width (uint8_t)
121 size_t output_index, uint8_t* buf, size_t max_size
122 ) const override {
123 (void)output_index;
124 if (max_size < 5) return 0;
125 std::memcpy(buf, &block_size_, sizeof(uint32_t));
126 buf[4] = element_width_;
127 return 5;
128 }
129
130 void deserializeHeader(const uint8_t* buf, size_t size) override {
131 if (size >= 4) std::memcpy(&block_size_, buf, sizeof(uint32_t));
132 if (size >= 5) element_width_ = buf[4];
133 }
134
135 size_t getMaxHeaderSize(size_t) const override { return 5; }
136
137 void saveState() override {
138 saved_block_size_ = block_size_;
139 saved_element_width_ = element_width_;
140 saved_actual_output_size_ = actual_output_size_;
141 }
142
143 void restoreState() override {
144 block_size_ = saved_block_size_;
145 element_width_ = saved_element_width_;
146 actual_output_size_ = saved_actual_output_size_;
147 }
148
149private:
150 bool is_inverse_;
151 uint32_t block_size_;
152 uint32_t saved_block_size_ = 0;
153 uint8_t element_width_;
154 uint8_t saved_element_width_ = 0;
155 size_t actual_output_size_ = 0;
156 size_t saved_actual_output_size_ = 0;
157
158 // Validate config and return N_chunk (elements per chunk).
159 // block_size must be a multiple of 1024*element_width so that butterfly
160 // kernels always have full warps in every __shfl_xor_sync call.
161 size_t validateConfig() const {
162 if (element_width_ != 1 && element_width_ != 2 &&
163 element_width_ != 4 && element_width_ != 8)
164 throw std::invalid_argument(
165 "BitshuffleStage: element_width must be 1, 2, 4, or 8");
166 if (block_size_ == 0 || block_size_ % (1024u * element_width_) != 0)
167 throw std::invalid_argument(
168 "BitshuffleStage: block_size must be a positive multiple of "
169 "1024 * element_width (default 16384 satisfies this for all "
170 "supported element widths)");
171 return block_size_ / element_width_;
172 }
173};
174
175} // namespace fz
Definition bitshuffle_stage.h:43
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
Definition bitshuffle_stage.h:94
size_t getMaxHeaderSize(size_t) const override
Definition bitshuffle_stage.h:135
uint8_t getOutputDataType(size_t) const override
Definition bitshuffle_stage.h:113
FusionSpec getFusionSpec() const override
Definition bitshuffle_stage.h:59
size_t getActualOutputSize(int index) const override
Definition bitshuffle_stage.h:105
std::string getName() const override
Definition bitshuffle_stage.h:90
void setInverse(bool inv) override
Definition bitshuffle_stage.h:53
void execute(fz::stream_t stream, MemoryPool *pool, const std::vector< void * > &inputs, const std::vector< void * > &outputs, const std::vector< size_t > &sizes) override
void saveState() override
Definition bitshuffle_stage.h:137
FusedOpDecl getFusedOp() const override
Chunk-cooperative fixed-length op: 32-bit bitshuffle. Stateless (no params).
Definition bitshuffle_stage.h:66
uint16_t getStageTypeId() const override
Definition bitshuffle_stage.h:109
size_t getRequiredInputAlignment() const override
Definition bitshuffle_stage.h:77
size_t serializeHeader(size_t output_index, uint8_t *buf, size_t max_size) const override
Definition bitshuffle_stage.h:120
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition bitshuffle_stage.h:130
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
Definition bitshuffle_stage.h:102
Definition mempool.h:82
Definition stage.h:31
FZM binary file format definitions — structs, enums, and helpers.
Definition dag.h:24
Base class interface for all compression stages.
A stage's contribution to a generated fused kernel — the device-op it maps to, where its source lives...
Definition fusion.h:161
A stage's fusion contract. Stages that can participate in a fused kernel override Stage::getFusionSpe...
Definition fusion.h:51
Backend-neutral GPU type aliases.