| 5 | namespace ctranslate2 { |
| 6 | |
| 7 | Padder::Padder(const StorageView& lengths, |
| 8 | const dim_t max_time, |
| 9 | const dim_t pad_batch_to_multiple) |
| 10 | : _batch_size(lengths.size()) { |
| 11 | const std::vector<int32_t> lengths_vec = lengths.to_vector<int32_t>(); |
| 12 | if (max_time < 0) |
| 13 | _max_time = *std::max_element(lengths_vec.begin(), lengths_vec.end()); |
| 14 | else |
| 15 | _max_time = max_time; |
| 16 | const bool has_padding = std::any_of(lengths_vec.begin(), |
| 17 | lengths_vec.end(), |
| 18 | [this](const int32_t length) { |
| 19 | return length != _max_time; |
| 20 | }); |
| 21 | if (!has_padding) |
| 22 | return; |
| 23 | |
| 24 | const dim_t max_size = _max_time * _batch_size; |
| 25 | std::vector<int32_t> padded_to_flat; |
| 26 | std::vector<int32_t> flat_to_padded; |
| 27 | padded_to_flat.reserve(max_size); |
| 28 | flat_to_padded.reserve(max_size); |
| 29 | |
| 30 | dim_t padded_offset = 0; |
| 31 | dim_t flat_offset = 0; |
| 32 | |
| 33 | for (dim_t i = 0; i < _batch_size; ++i) { |
| 34 | const dim_t length = lengths_vec[i]; |
| 35 | for (dim_t t = 0; t < length; ++t) { |
| 36 | padded_to_flat.push_back(padded_offset + t); |
| 37 | flat_to_padded.push_back(flat_offset + t); |
| 38 | } |
| 39 | for (dim_t t = length; t < _max_time; ++t) { |
| 40 | flat_to_padded.push_back(flat_offset + length - 1); |
| 41 | } |
| 42 | padded_offset += _max_time; |
| 43 | flat_offset += length; |
| 44 | } |
| 45 | |
| 46 | while (padded_to_flat.size() % pad_batch_to_multiple != 0) { |
| 47 | padded_to_flat.push_back(padded_to_flat.back()); |
| 48 | ++flat_offset; |
| 49 | } |
| 50 | |
| 51 | const Device device = lengths.device(); |
| 52 | _padded_to_flat = StorageView({flat_offset}, padded_to_flat, device); |
| 53 | _flat_to_padded = StorageView({padded_offset}, flat_to_padded, device); |
| 54 | } |
| 55 | |
| 56 | void Padder::remove_padding(StorageView& x) const { |
| 57 | if (!_padded_to_flat) |
nothing calls this directly
no test coverage detected