|
20 | 20 | #ifndef PARQUET_UTIL_BIT_STREAM_UTILS_INLINE_H |
21 | 21 | #define PARQUET_UTIL_BIT_STREAM_UTILS_INLINE_H |
22 | 22 |
|
| 23 | +#include <algorithm> |
| 24 | + |
23 | 25 | #include "parquet/util/bit-stream-utils.h" |
| 26 | +#include "parquet/util/bpacking.h" |
24 | 27 |
|
25 | 28 | namespace parquet { |
26 | 29 |
|
@@ -85,35 +88,98 @@ inline bool BitWriter::PutVlqInt(uint32_t v) { |
85 | 88 | return result; |
86 | 89 | } |
87 | 90 |
|
| 91 | +template <typename T> |
| 92 | +inline void GetValue_(int num_bits, T* v, int max_bytes, const uint8_t* buffer, |
| 93 | + int* bit_offset, int* byte_offset, uint64_t* buffered_values) { |
| 94 | + *v = BitUtil::TrailingBits(*buffered_values, *bit_offset + num_bits) >> *bit_offset; |
| 95 | + |
| 96 | + *bit_offset += num_bits; |
| 97 | + if (*bit_offset >= 64) { |
| 98 | + *byte_offset += 8; |
| 99 | + *bit_offset -= 64; |
| 100 | + |
| 101 | + int bytes_remaining = max_bytes - *byte_offset; |
| 102 | + if (LIKELY(bytes_remaining >= 8)) { |
| 103 | + memcpy(buffered_values, buffer + *byte_offset, 8); |
| 104 | + } else { |
| 105 | + memcpy(buffered_values, buffer + *byte_offset, bytes_remaining); |
| 106 | + } |
| 107 | + |
| 108 | + // Read bits of v that crossed into new buffered_values_ |
| 109 | + *v |= BitUtil::TrailingBits(*buffered_values, *bit_offset) |
| 110 | + << (num_bits - *bit_offset); |
| 111 | + DCHECK_LE(*bit_offset, 64); |
| 112 | + } |
| 113 | +} |
| 114 | + |
88 | 115 | template <typename T> |
89 | 116 | inline bool BitReader::GetValue(int num_bits, T* v) { |
| 117 | + return GetBatch(num_bits, v, 1) == 1; |
| 118 | +} |
| 119 | + |
| 120 | +template <typename T> |
| 121 | +inline int BitReader::GetBatch(int num_bits, T* v, int batch_size) { |
90 | 122 | DCHECK(buffer_ != NULL); |
91 | 123 | // TODO: revisit this limit if necessary |
92 | 124 | DCHECK_LE(num_bits, 32); |
93 | 125 | DCHECK_LE(num_bits, static_cast<int>(sizeof(T) * 8)); |
94 | 126 |
|
95 | | - if (UNLIKELY(byte_offset_ * 8 + bit_offset_ + num_bits > max_bytes_ * 8)) return false; |
96 | | - |
97 | | - *v = BitUtil::TrailingBits(buffered_values_, bit_offset_ + num_bits) >> bit_offset_; |
98 | | - |
99 | | - bit_offset_ += num_bits; |
100 | | - if (bit_offset_ >= 64) { |
101 | | - byte_offset_ += 8; |
102 | | - bit_offset_ -= 64; |
| 127 | + int bit_offset = bit_offset_; |
| 128 | + int byte_offset = byte_offset_; |
| 129 | + uint64_t buffered_values = buffered_values_; |
| 130 | + int max_bytes = max_bytes_; |
| 131 | + const uint8_t* buffer = buffer_; |
| 132 | + |
| 133 | + uint64_t needed_bits = num_bits * batch_size; |
| 134 | + uint64_t remaining_bits = (max_bytes - byte_offset) * 8 - bit_offset; |
| 135 | + if (remaining_bits < needed_bits) { batch_size = remaining_bits / num_bits; } |
| 136 | + |
| 137 | + int i = 0; |
| 138 | + if (UNLIKELY(bit_offset != 0)) { |
| 139 | + for (; i < batch_size && bit_offset != 0; ++i) { |
| 140 | + GetValue_(num_bits, &v[i], max_bytes, buffer, &bit_offset, &byte_offset, |
| 141 | + &buffered_values); |
| 142 | + } |
| 143 | + } |
103 | 144 |
|
104 | | - int bytes_remaining = max_bytes_ - byte_offset_; |
105 | | - if (LIKELY(bytes_remaining >= 8)) { |
106 | | - memcpy(&buffered_values_, buffer_ + byte_offset_, 8); |
107 | | - } else { |
108 | | - memcpy(&buffered_values_, buffer_ + byte_offset_, bytes_remaining); |
| 145 | + if (sizeof(T) == 4) { |
| 146 | + int num_unpacked = unpack32(reinterpret_cast<const uint32_t*>(buffer + byte_offset), |
| 147 | + reinterpret_cast<uint32_t*>(v + i), batch_size - i, num_bits); |
| 148 | + i += num_unpacked; |
| 149 | + byte_offset += num_unpacked * num_bits / 8; |
| 150 | + } else { |
| 151 | + const int buffer_size = 1024; |
| 152 | + static uint32_t unpack_buffer[buffer_size]; |
| 153 | + while (i < batch_size) { |
| 154 | + int unpack_size = std::min(buffer_size, batch_size - i); |
| 155 | + int num_unpacked = unpack32(reinterpret_cast<const uint32_t*>(buffer + byte_offset), |
| 156 | + unpack_buffer, unpack_size, num_bits); |
| 157 | + if (num_unpacked == 0) { break; } |
| 158 | + for (int k = 0; k < num_unpacked; ++k) { |
| 159 | + v[i + k] = unpack_buffer[k]; |
| 160 | + } |
| 161 | + i += num_unpacked; |
| 162 | + byte_offset += num_unpacked * num_bits / 8; |
109 | 163 | } |
| 164 | + } |
110 | 165 |
|
111 | | - // Read bits of v that crossed into new buffered_values_ |
112 | | - *v |= BitUtil::TrailingBits(buffered_values_, bit_offset_) |
113 | | - << (num_bits - bit_offset_); |
| 166 | + int bytes_remaining = max_bytes - byte_offset; |
| 167 | + if (bytes_remaining >= 8) { |
| 168 | + memcpy(&buffered_values, buffer + byte_offset, 8); |
| 169 | + } else { |
| 170 | + memcpy(&buffered_values, buffer + byte_offset, bytes_remaining); |
114 | 171 | } |
115 | | - DCHECK_LE(bit_offset_, 64); |
116 | | - return true; |
| 172 | + |
| 173 | + for (; i < batch_size; ++i) { |
| 174 | + GetValue_( |
| 175 | + num_bits, &v[i], max_bytes, buffer, &bit_offset, &byte_offset, &buffered_values); |
| 176 | + } |
| 177 | + |
| 178 | + bit_offset_ = bit_offset; |
| 179 | + byte_offset_ = byte_offset; |
| 180 | + buffered_values_ = buffered_values; |
| 181 | + |
| 182 | + return batch_size; |
117 | 183 | } |
118 | 184 |
|
119 | 185 | template <typename T> |
|
0 commit comments