41#include <unordered_map>
69 , actual_output_size_(0)
70 , cached_orig_bytes_(0)
72 , d_sizes_dev_(
nullptr)
73 , d_clean_dev_(
nullptr)
74 , d_dst_off_dev_(
nullptr)
75 , scratch_capacity_(0)
81 void setInverse(
bool inv)
override { is_inverse_ = inv; }
82 bool isInverse()
const override {
return is_inverse_; }
89 void setChunkSize(
size_t bytes) { chunk_size_ =
static_cast<uint32_t
>(bytes); }
90 void setWordSize(
size_t bytes) { word_size_ =
static_cast<uint8_t
>(bytes); }
92 size_t getChunkSize()
const {
return chunk_size_; }
94 int getWordSize()
const {
return static_cast<int>(word_size_); }
95 uint32_t getCachedOrigBytes()
const {
return cached_orig_bytes_; }
101 if (is_inverse_ || word_size_ != 1 || chunk_size_ != 16384u)
return {};
102 return FusionSpec{FusionAccess::Cooperative, chunk_size_};
106 return FusedOpDecl{FusionStrategy::ChunkCooperative,
"RARECoder",
107 "fused/chunk_fusion/chunk_fusion.cuh", {}};
112 actual_output_size_ = archive_bytes;
113 cached_orig_bytes_ =
static_cast<uint32_t
>(orig_bytes);
114 tail_readback_pending_ =
false;
124 const std::vector<void*>& inputs,
125 const std::vector<void*>& outputs,
126 const std::vector<size_t>& sizes
131 std::string
getName()
const override {
return "RARE"; }
132 size_t getNumInputs()
const override {
return 1; }
133 size_t getNumOutputs()
const override {
return 1; }
136 const std::vector<size_t>& input_sizes
139 if (cached_orig_bytes_ > 0)
140 return {
static_cast<size_t>(cached_orig_bytes_)};
141 return {input_sizes.empty() ? 0 : input_sizes[0]};
144 const size_t n_bytes = input_sizes.empty() ? 0 : input_sizes[0];
145 const size_t n_chunks = (n_bytes + chunk_size_ - 1) / chunk_size_;
146 const size_t hdr = 4 + 4 + 4 * n_chunks;
153 const size_t worst = n_bytes + hdr;
154 return {(worst + 3) & ~
size_t(3)};
157 std::unordered_map<std::string, size_t>
170 const std::vector<size_t>& input_sizes
172 if (is_inverse_ || input_sizes.empty())
return 0;
173 const size_t in_bytes = input_sizes[0];
174 const size_t n_chunks = (in_bytes + chunk_size_ - 1) / chunk_size_;
175 return n_chunks * (
static_cast<size_t>(chunk_size_) + 3 *
sizeof(uint32_t));
183 return static_cast<uint8_t
>(DataType::UINT8);
188 size_t output_index, uint8_t* buf,
size_t max_size
191 if (max_size < 9)
return 0;
192 std::memcpy(buf, &chunk_size_,
sizeof(uint32_t));
194 std::memcpy(buf + 5, &cached_orig_bytes_,
sizeof(uint32_t));
199 if (size >= 4) std::memcpy(&chunk_size_, buf,
sizeof(uint32_t));
200 if (size >= 5) word_size_ = buf[4];
201 if (size >= 9) std::memcpy(&cached_orig_bytes_, buf + 5,
sizeof(uint32_t));
207 saved_chunk_size_ = chunk_size_;
208 saved_word_size_ = word_size_;
209 saved_cached_orig_bytes_ = cached_orig_bytes_;
212 void restoreState()
override {
213 chunk_size_ = saved_chunk_size_;
214 word_size_ = saved_word_size_;
215 cached_orig_bytes_ = saved_cached_orig_bytes_;
220 uint32_t chunk_size_;
221 uint32_t saved_chunk_size_ = 0;
223 uint8_t saved_word_size_ = 0;
224 size_t actual_output_size_;
225 uint32_t cached_orig_bytes_ = 0;
226 uint32_t saved_cached_orig_bytes_ = 0;
230 uint32_t* d_sizes_dev_;
231 uint32_t* d_clean_dev_;
232 uint32_t* d_dst_off_dev_;
233 mutable bool tail_readback_pending_ =
false;
234 mutable cudaStream_t tail_readback_stream_ =
nullptr;
235 mutable uint32_t tail_last_index_ = 0;
236 mutable uint8_t* tail_output_ptr_ =
nullptr;
237 size_t scratch_capacity_;
238 MemoryPool* scratch_pool_owner_ =
nullptr;
239 bool scratch_from_pool_ =
false;
Definition rare_stage.h:63
void setFusedArchiveResult(size_t archive_bytes, size_t orig_bytes) override
Definition rare_stage.h:116
void setInverse(bool inv) override
Definition rare_stage.h:81
size_t getMaxHeaderSize(size_t) const override
Definition rare_stage.h:204
void saveState() override
Definition rare_stage.h:206
bool isGraphCompatible() const override
Definition rare_stage.h:87
void deserializeHeader(const uint8_t *buf, size_t size) override
Definition rare_stage.h:198
void postStreamSync(cudaStream_t stream) override
std::unordered_map< std::string, size_t > getActualOutputSizesByName() const override
uint8_t getOutputDataType(size_t) const override
Definition rare_stage.h:182
void execute(cudaStream_t stream, MemoryPool *pool, const std::vector< void * > &inputs, const std::vector< void * > &outputs, const std::vector< size_t > &sizes) override
size_t getRequiredInputAlignment() const override
Definition rare_stage.h:93
FusedOpDecl getFusedOp() const override
Definition rare_stage.h:104
size_t serializeHeader(size_t output_index, uint8_t *buf, size_t max_size) const override
Definition rare_stage.h:187
FusionSpec getFusionSpec() const override
Definition rare_stage.h:100
std::vector< size_t > estimateOutputSizes(const std::vector< size_t > &input_sizes) const override
Definition rare_stage.h:135
size_t estimateScratchBytes(const std::vector< size_t > &input_sizes) const override
Definition rare_stage.h:169
uint16_t getStageTypeId() const override
Definition rare_stage.h:178
size_t getActualOutputSize(int index) const override
std::string getName() const override
Definition rare_stage.h:131
void setFusedResult(size_t archive_bytes, size_t orig_bytes)
Definition rare_stage.h:111
@ RARE
Adaptive top-bit matching generalization of RRE (LC component)
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.