MCPcopy Create free account
hub / github.com/OpenNMT/CTranslate2 / Padder

Method Padder

src/padder.cc:7–54  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5namespace 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)

Callers

nothing calls this directly

Calls 5

StorageViewClass · 0.85
beginMethod · 0.80
endMethod · 0.80
sizeMethod · 0.45
deviceMethod · 0.45

Tested by

no test coverage detected