| 98 | } |
| 99 | |
| 100 | void append(const std::vector<argument>& iter_state, |
| 101 | const std::vector<argument>& concatenated_outputs, |
| 102 | const std::vector<int64_t>& scan_output_dirs, |
| 103 | int64_t curr_iter, |
| 104 | int64_t num_iters) const |
| 105 | { |
| 106 | assert(iter_state.size() == concatenated_outputs.size()); |
| 107 | for(auto i : range(iter_state.size())) |
| 108 | { |
| 109 | const auto& iter_stat = iter_state.at(i); |
| 110 | const auto& scan_out = concatenated_outputs.at(i); |
| 111 | |
| 112 | auto dir = scan_output_dirs.empty() ? 0 : scan_output_dirs[i]; |
| 113 | auto idx = (1 - dir) * curr_iter + dir * (num_iters - 1 - curr_iter); |
| 114 | |
| 115 | auto* in_data = iter_stat.data(); |
| 116 | auto* out_data = scan_out.data(); |
| 117 | std::size_t out_size = iter_stat.get_shape().bytes(); |
| 118 | assert((idx + 1) * out_size <= scan_out.get_shape().bytes()); |
| 119 | std::copy(in_data, in_data + out_size, out_data + idx * out_size); |
| 120 | } |
| 121 | } |
| 122 | |
| 123 | void set_zero(context&, const std::vector<argument>& concatenated_outputs, int iter) const |
| 124 | { |