26#include <unordered_map>
49 , actual_output_size_(0)
53 void setInverse(
bool inv)
override { is_inverse_ = inv; }
58 if (is_inverse_ || element_width_ != 4 || block_size_ != 16384u)
return {};
59 return FusionSpec{FusionAccess::BlockLocal, block_size_};
65 return FusedOpDecl{FusionStrategy::ChunkCooperative,
"Bitshuffle32",
66 "fused/chunk_fusion/chunk_fusion.cuh", {}};
68 bool isInverse()
const override {
return is_inverse_; }
70 void setBlockSize(
size_t bytes) { block_size_ =
static_cast<uint32_t
>(bytes); }
71 void setElementWidth(
size_t bytes){ element_width_ =
static_cast<uint8_t
>(bytes); }
73 size_t getBlockSize()
const {
return block_size_; }
75 size_t getElementWidth()
const {
return element_width_; }
81 const std::vector<void*>& inputs,
82 const std::vector<void*>& outputs,
83 const std::vector<size_t>& sizes
87 std::string
getName()
const override {
return "Bitshuffle"; }
88 size_t getNumInputs()
const override {
return 1; }
89 size_t getNumOutputs()
const override {
return 1; }
92 const std::vector<size_t>& input_sizes
95 return {input_sizes[0]};
98 std::unordered_map<std::string, size_t>
100 return {{
"output", actual_output_size_}};
103 return (index == 0) ? actual_output_size_ : 0;
107 return static_cast<uint16_t
>(StageType::BITSHUFFLE);
112 return static_cast<uint8_t
>(DataType::UINT8);
118 size_t output_index, uint8_t* buf,
size_t max_size
121 if (max_size < 5)
return 0;
122 std::memcpy(buf, &block_size_,
sizeof(uint32_t));
123 buf[4] = element_width_;
128 if (size >= 4) std::memcpy(&block_size_, buf,
sizeof(uint32_t));
129 if (size >= 5) element_width_ = buf[4];
135 saved_block_size_ = block_size_;
136 saved_element_width_ = element_width_;
137 saved_actual_output_size_ = actual_output_size_;
140 void restoreState()
override {
141 block_size_ = saved_block_size_;
142 element_width_ = saved_element_width_;
143 actual_output_size_ = saved_actual_output_size_;
148 uint32_t block_size_;
149 uint32_t saved_block_size_ = 0;
150 uint8_t element_width_;
151 uint8_t saved_element_width_ = 0;
152 size_t actual_output_size_ = 0;
153 size_t saved_actual_output_size_ = 0;
158 size_t validateConfig()
const {
159 if (element_width_ != 1 && element_width_ != 2 &&
160 element_width_ != 4 && element_width_ != 8)
161 throw std::invalid_argument(
162 "BitshuffleStage: element_width must be 1, 2, 4, or 8");
163 if (block_size_ == 0 || block_size_ % (1024u * element_width_) != 0)
164 throw std::invalid_argument(
165 "BitshuffleStage: block_size must be a positive multiple of "
166 "1024 * element_width (default 16384 satisfies this for all "
167 "supported element widths)");
168 return block_size_ / element_width_;
Definition bitshuffle_stage.h:43
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
Definition bitshuffle_stage.h:91
size_t getMaxHeaderSize(size_t) const override
Definition bitshuffle_stage.h:132
uint8_t getOutputDataType(size_t) const override
Definition bitshuffle_stage.h:110
FusionSpec getFusionSpec() const override
Definition bitshuffle_stage.h:57
size_t getActualOutputSize(int index) const override
Definition bitshuffle_stage.h:102
std::string getName() const override
Definition bitshuffle_stage.h:87
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:134
FusedOpDecl getFusedOp() const override
Chunk-cooperative fixed-length op: 32-bit bitshuffle. Stateless (no params).
Definition bitshuffle_stage.h:63
uint16_t getStageTypeId() const override
Definition bitshuffle_stage.h:106
size_t getRequiredInputAlignment() const override
Definition bitshuffle_stage.h:74
size_t serializeHeader(size_t output_index, uint8_t *buf, size_t max_size) const override
Definition bitshuffle_stage.h:117
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition bitshuffle_stage.h:127
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
Definition bitshuffle_stage.h:99
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:52
Backend-neutral GPU type aliases.