107 virtual size_t execute(std::shared_ptr<Op_Base::workbench_t> pwb,
109 const V& initial_state_h,
110 const V& initial_state_c,
111 const std::shared_ptr<param_base_t> params,
112 const uint32_t start_code,
113 const uint32_t stop_code,
114 const size_t beam_size,
115 std::vector<uint32_t>& output_seq,
116 const size_t max_output_len
119#ifdef BEAM_DECODER_PROFILE
120 using clock = std::chrono::system_clock;
121 using ms = std::chrono::duration<double, std::milli>;
123 auto before = clock::now();
126 assert(
nullptr != pwb);
127 assert(
nullptr != params);
128 const auto& layer = std::dynamic_pointer_cast<const params_t>(params)->lstm;
129 const auto& linear = std::dynamic_pointer_cast<const params_t>(params)->linear;
131 assert(stop_idx != 0);
133 auto wb = std::dynamic_pointer_cast<workbench_t>(pwb);
135 V& lin_out = wb->lin_out;
136 std::vector<uint32_t>& temp_indicies = wb->temp_indicies;
138 std::vector<std::vector<decoding_timepoint_t>> decoding_log(max_output_len, std::vector<decoding_timepoint_t>(beam_size));
140 size_t hidden_size = layer.weight_ih.rows() / 4;
142 V c = initial_state_c;
143 V h = initial_state_h;
147 std::vector<uint32_t>& initial_top_classes = wb->initial_top_classes;
148 initial_top_classes.resize(beam_size);
149 std::vector<float>& initial_logprob = wb->initial_logprob;
150 initial_logprob.resize(beam_size);
152#ifdef BEAM_DECODER_PROFILE
153 auto start_decoding = clock::now();
156 step_with_decoding(layer, linear, input, output, s, lin_out, hidden_size, beam_size, 0,
157 initial_top_classes, temp_indicies, initial_logprob, c, h);
159#ifdef BEAM_DECODER_PROFILE
160 auto after_first_step = clock::now();
163 size_t decoding_step = 0;
164 for (
size_t i = 0; i < beam_size; ++i)
166 decoding_log[decoding_step][i].cls = initial_top_classes[i];
170 M& states_c = wb->states_c;
171 if ((
size_t)states_c.cols() != beam_size)
172 states_c = M::Zero(hidden_size, beam_size);
173 for (
size_t i = 0; i < beam_size; ++i) states_c.col(i) = c;
174 M& states_h = wb->states_h;
175 if ((
size_t)states_h.cols() != beam_size)
176 states_h = M::Zero(hidden_size, beam_size);
177 for (
size_t i = 0; i < beam_size; ++i) states_h.col(i) = h;
179 std::vector<uint32_t>& top_classes = wb->top_classes;
180 top_classes.resize(beam_size * beam_size);
181 std::vector<float>& logprob = wb->logprob;
182 logprob.resize(beam_size * beam_size);
183 std::vector<uint32_t>& indices = wb->indices;
184 indices.resize(beam_size * beam_size);
186#ifdef BEAM_DECODER_PROFILE
187 auto start_next_step = clock::now();
190 while (decoding_step < max_output_len)
193 for (
size_t i = 0; i < beam_size; ++i)
195 if (stop_idx == decoding_log[decoding_step - 1][i].cls)
200 Eigen::Ref<const M> input = embd.
get_ref_by_idx(decoding_log[decoding_step - 1][i].cls);
202 hidden_size, beam_size,
204 top_classes, temp_indicies, logprob, states_c, states_h, i);
205 for (
size_t j = 0; j < beam_size; ++j)
207 logprob[pos + j] += initial_logprob[i];
212 for (uint32_t i = 0; i < pos; ++i) indices[i] = i;
213 std::sort(indices.begin(), indices.begin() + pos, [&logprob](
size_t a,
size_t b){
214 return logprob[a] > logprob[b];
228 for (
size_t i = 0; i < beam_size; ++i)
230 size_t idx = indices[i];
231 size_t ref = idx % beam_size;
232 states_c.col(i) = states_c.col(ref);
233 states_h.col(i) = states_h.col(ref);
237 for (
size_t i = 0; i < std::min(beam_size, pos); ++i)
239 size_t idx = indices[i];
240 decoding_log[decoding_step][i].cls = top_classes[idx];
241 if (decoding_log[decoding_step][i].cls != stop_idx)
245 decoding_log[decoding_step][i].prev_pos = idx % beam_size;
246 initial_logprob[i] = logprob[idx];
279#ifdef BEAM_DECODER_PROFILE
280 auto start_backtracking = clock::now();
286 while (step < decoding_log.size() && decoding_log[step][pos].cls != 1) step++;
295 char32_t ch = embd.
decode(decoding_log[step][pos].cls);
296 output_seq.push_back(ch);
297 pos = decoding_log[step][pos].prev_pos;
302 std::reverse(output_seq.begin(), output_seq.end());
304#ifdef BEAM_DECODER_PROFILE
305 auto end_backtracking = clock::now();
307 std::cerr <<
" prepare : " << (start_decoding - before).count() << std::endl;
308 std::cerr <<
" first step : " << (after_first_step - start_decoding).count() << std::endl;
309 std::cerr <<
" prepare : " << (start_next_step - after_first_step).count() << std::endl;
310 std::cerr <<
" all steps : " << (start_backtracking - start_next_step).count() << std::endl;
311 std::cerr <<
" backtracking : " << (end_backtracking - start_backtracking).count() << std::endl;
378 const size_t beam_size,
379 const size_t start_pos,
380 std::vector<uint32_t>& top_classes,
381 std::vector<uint32_t> indices,
382 std::vector<float>& logprob
387 temp.noalias() += linear.
weight * output;
388 Eigen::Index idx = 0;
389 typename V::Scalar max_value = temp.maxCoeff(&idx);
391 temp.array() -= (max_value + Eigen::exp(temp.array() - max_value).sum());
392 for (uint32_t i = 0; i < indices.size(); ++i) indices[i] = i;
394 std::partial_sort_copy(indices.begin(), indices.end(), top_classes.begin() + start_pos, top_classes.begin() + start_pos + beam_size,
395 [&temp](uint32_t a, uint32_t b) {
396 return temp[a] > temp[b];
399 for (
size_t i = 0; i < beam_size; ++i)
401 logprob[start_pos + i] = temp[top_classes[start_pos + i]];