73 , actual_output_size_(0)
74 , cached_orig_bytes_(0)
80 void setInverse(
bool inv)
override { is_inverse_ = inv; }
81 bool isInverse()
const override {
return is_inverse_; }
85 void setChunkSize(
size_t bytes) { chunk_size_ =
static_cast<uint32_t
>(bytes); }
86 void setWordSize(
size_t bytes) { word_size_ =
static_cast<uint8_t
>(bytes); }
100 void setMatchLevel(
int level) { match_level_ =
static_cast<uint8_t
>(level); }
101 int getMatchLevel()
const {
return static_cast<int>(match_level_); }
130 bool getSplitMode()
const {
return split_mode_; }
132 size_t getChunkSize()
const {
return chunk_size_; }
147 int getWordSize()
const {
return static_cast<int>(word_size_); }
148 uint32_t getCachedOrigBytes()
const {
return cached_orig_bytes_; }
154 const std::vector<void*>& inputs,
155 const std::vector<void*>& outputs,
156 const std::vector<size_t>& sizes
161 std::string
getName()
const override {
return "GPULZ"; }
162 size_t getNumInputs()
const override {
163 return (is_inverse_ && split_mode_) ? 4 : 1;
165 size_t getNumOutputs()
const override {
166 return (!is_inverse_ && split_mode_) ? 4 : 1;
170 if (!is_inverse_ && split_mode_)
171 return {
"literals",
"lengths",
"offsets",
"meta"};
176 const std::vector<size_t>& input_sizes
181 if (cached_orig_bytes_ > 0)
182 return {
static_cast<size_t>(cached_orig_bytes_)};
183 return {input_sizes.empty() ? 0 : input_sizes[0]};
185 const size_t n_bytes = input_sizes.empty() ? 0 : input_sizes[0];
186 const size_t n_chunks = (n_bytes + chunk_size_ - 1) / chunk_size_;
187 const size_t hdr = 4 + 4 + 8 * n_chunks;
196 const size_t padded = n_chunks * chunk_size_;
199 const size_t block_elems = chunk_size_ / word_size_;
200 const size_t flag_stride = (block_elems + 7) / 8;
204 return {align4(padded),
205 align4(n_chunks * block_elems),
206 align4(n_chunks * block_elems),
207 align4(hdr + n_chunks * flag_stride)};
213 return {align4(padded + hdr)};
216 std::unordered_map<std::string, size_t>
231 const std::vector<size_t>& input_sizes
233 if (input_sizes.empty())
return 0;
234 const size_t in_bytes = input_sizes[0];
235 const size_t n_chunks = (in_bytes + chunk_size_ - 1) / chunk_size_;
236 const size_t block_elems = chunk_size_ / word_size_;
237 const size_t flag_bytes_max = (block_elems + 7) / 8;
241 return split_mode_ ? (in_bytes + n_chunks * flag_bytes_max
242 + 4 * n_chunks *
sizeof(uint32_t))
245 size_t bytes = n_chunks * (
static_cast<size_t>(chunk_size_)
246 + flag_bytes_max + 4 *
sizeof(uint32_t));
247 if (split_mode_) bytes += n_chunks * 5 *
sizeof(uint32_t) + 16;
256 return static_cast<uint8_t
>(DataType::UINT8);
264 size_t output_index, uint8_t* buf,
size_t max_size
267 if (max_size < 14)
return 0;
268 std::memcpy(buf, &chunk_size_,
sizeof(uint32_t));
270 std::memcpy(buf + 5, &cached_orig_bytes_,
sizeof(uint32_t));
271 buf[9] = split_mode_ ? 1u : 0u;
272 std::memcpy(buf + 10, &orig_unpadded_bytes_,
sizeof(uint32_t));
277 if (size >= 4) std::memcpy(&chunk_size_, buf,
sizeof(uint32_t));
278 if (size >= 5) word_size_ = buf[4];
279 if (size >= 9) std::memcpy(&cached_orig_bytes_, buf + 5,
sizeof(uint32_t));
280 if (size >= 10) split_mode_ = (buf[9] != 0);
281 if (size >= 14) std::memcpy(&orig_unpadded_bytes_, buf + 10,
sizeof(uint32_t));
287 saved_chunk_size_ = chunk_size_;
288 saved_word_size_ = word_size_;
289 saved_cached_orig_bytes_ = cached_orig_bytes_;
290 saved_split_mode_ = split_mode_;
291 saved_orig_unpadded_bytes_ = orig_unpadded_bytes_;
294 void restoreState()
override {
295 chunk_size_ = saved_chunk_size_;
296 word_size_ = saved_word_size_;
297 cached_orig_bytes_ = saved_cached_orig_bytes_;
298 split_mode_ = saved_split_mode_;
299 orig_unpadded_bytes_ = saved_orig_unpadded_bytes_;
303 static constexpr size_t align4(
size_t n) {
return (n + 3) & ~size_t(3); }
306 void finishSplitReadback(fz::stream_t stream)
const;
312 void freeForwardScratch(fz::stream_t stream,
bool sync_before_cuda_free);
314 void executeForward(fz::stream_t stream, MemoryPool* pool,
315 const std::vector<void*>& inputs,
316 const std::vector<void*>& outputs,
318 void executeInverse(fz::stream_t stream, MemoryPool* pool,
319 const std::vector<void*>& inputs,
320 const std::vector<void*>& outputs,
324 uint32_t chunk_size_;
325 uint32_t saved_chunk_size_ = 0;
327 uint8_t saved_word_size_ = 0;
328 uint8_t match_level_ = 1;
329 bool split_mode_ =
false;
330 bool saved_split_mode_ =
false;
331 size_t actual_output_size_;
333 size_t actual_split_sizes_[4] = {0, 0, 0, 0};
334 uint32_t cached_orig_bytes_ = 0;
335 uint32_t saved_cached_orig_bytes_ = 0;
340 uint32_t orig_unpadded_bytes_ = 0;
341 uint32_t saved_orig_unpadded_bytes_ = 0;
344 uint8_t* d_data_scratch_ =
nullptr;
345 uint8_t* d_flag_scratch_ =
nullptr;
346 uint32_t* d_flag_size_ =
nullptr;
347 uint32_t* d_data_size_ =
nullptr;
348 uint32_t* d_clean_dev_ =
nullptr;
349 uint32_t* d_dst_off_dev_ =
nullptr;
352 uint32_t* d_lit_off_dev_ =
nullptr;
353 uint32_t* d_tok_off_dev_ =
nullptr;
354 uint32_t* d_meta_off_dev_ =
nullptr;
355 uint32_t* d_lit_cnt_dev_ =
nullptr;
356 uint32_t* d_tok_cnt_dev_ =
nullptr;
357 uint32_t* d_totals_dev_ =
nullptr;
358 mutable uint8_t* split_out_ptr_[4] = {
nullptr,
nullptr,
nullptr,
nullptr};
359 mutable bool split_readback_pending_ =
false;
360 mutable bool tail_readback_pending_ =
false;
361 mutable fz::stream_t tail_readback_stream_ =
nullptr;
362 mutable uint32_t tail_last_index_ = 0;
363 mutable uint8_t* tail_output_ptr_ =
nullptr;
364 size_t scratch_capacity_ = 0;
365 MemoryPool* scratch_pool_owner_ =
nullptr;
366 bool scratch_from_pool_ =
false;
368 std::weak_ptr<const void> scratch_alive_;