LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
lstm_beam_decoder.h
Go to the documentation of this file.
1// Copyright 2002-2022 CEA LIST
2// SPDX-FileCopyrightText: 2022 CEA LIST <gael.de-chalendar@cea.fr>
3//
4// SPDX-License-Identifier: MIT
5
6#ifndef DEEPLIMA_SRC_INFERENCE_EIGEN_LSTM_BEAM_DECODER_H
7#define DEEPLIMA_SRC_INFERENCE_EIGEN_LSTM_BEAM_DECODER_H
8
9#include <vector>
10#include <algorithm>
11
12#include "bilstm.h"
13#include "linear.h"
14
15#include "embd_dict.h"
16
17//#define BEAM_DECODER_PROFILE
18#ifdef BEAM_DECODER_PROFILE
19#include <chrono>
20#endif
21
22namespace deeplima
23{
24namespace eigen_impl
25{
26
27template<class M, class V>
33
34template<class M, class V, class T>
36{
37protected:
39 {
40 workbench_t(uint32_t hidden_size,
41 const std::vector<uint32_t>& output_sizes,
42 bool precomputed_input=false)
43 : temp(M::Zero(precomputed_input ? 0 : hidden_size * 4, 1)),
44 out(M::Zero(hidden_size, 1)),
45 lin_out(V::Zero(output_sizes[0])),
46 zero(V::Zero(hidden_size)),
47 temp_indicies(output_sizes[0]),
48 m_precomputed_input(precomputed_input)
49 {
50 }
51
52 virtual ~workbench_t() {}
53
55 M out;
57 V s;
59 const V zero;
60 std::vector<uint32_t> initial_top_classes;
61 std::vector<float> initial_logprob;
62 std::vector<uint32_t> top_classes;
63 std::vector<float> logprob;
64 std::vector<uint32_t> indices;
65 std::vector<uint32_t> temp_indicies;
67 };
68
70 {
71 size_t prev_pos;
72 uint32_t cls;
73
75 };
76
77public:
78
80
81 virtual std::shared_ptr<Op_Base::workbench_t> create_workbench(uint32_t /*input_size*/,
82 const std::shared_ptr<param_base_t> params,
83 bool precomputed_input=false) const override
84 {
85 assert(nullptr != params);
86 const auto& layer = std::dynamic_pointer_cast<const params_t>(params)->lstm;
87 const auto& linear = std::dynamic_pointer_cast<const params_t>(params)->linear;
88
89 std::vector<uint32_t> output_sizes = { static_cast<uint32_t>(linear.weight.rows()) };
90 return std::make_shared<workbench_t>(layer.weight_ih.rows() / 4, output_sizes, precomputed_input);
91 }
92
93 virtual bool supports_precomputing() const
94 {
95 return true;
96 }
97
98 virtual void precompute_inputs(const std::shared_ptr<param_base_t> params, const M& inputs, M& outputs, int64_t first_column)
99 {
100 assert(nullptr != params);
101 const auto& layer = std::dynamic_pointer_cast<const params_t>(params)->lstm;
102
103 const V biases = layer.bias_ih + layer.bias_hh;
104 outputs.block(0, first_column, outputs.rows(), inputs.cols()) = (layer.weight_ih * inputs).colwise() + biases;
105 }
106
107 virtual size_t execute(std::shared_ptr<Op_Base::workbench_t> pwb,
108 const EmbdUInt64Float& embd,
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
117 )
118 {
119#ifdef BEAM_DECODER_PROFILE
120 using clock = std::chrono::system_clock;
121 using ms = std::chrono::duration<double, std::milli>;
122
123 auto before = clock::now();
124#endif
125
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;
130 EmbdUInt64Float::index_t stop_idx = embd.lookup(stop_code);
131 assert(stop_idx != 0);
132
133 auto wb = std::dynamic_pointer_cast<workbench_t>(pwb);
134 M& output = wb->out;
135 V& lin_out = wb->lin_out;
136 std::vector<uint32_t>& temp_indicies = wb->temp_indicies;
137
138 std::vector<std::vector<decoding_timepoint_t>> decoding_log(max_output_len, std::vector<decoding_timepoint_t>(beam_size));
139
140 size_t hidden_size = layer.weight_ih.rows() / 4;
141 V& s = wb->s;
142 V c = initial_state_c;
143 V h = initial_state_h;
144
145 Eigen::Ref<const M> input = embd.get_ref_by_key(start_code);
146
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);
151
152#ifdef BEAM_DECODER_PROFILE
153 auto start_decoding = clock::now();
154#endif
155
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);
158
159#ifdef BEAM_DECODER_PROFILE
160 auto after_first_step = clock::now();
161#endif
162
163 size_t decoding_step = 0;
164 for (size_t i = 0; i < beam_size; ++i)
165 {
166 decoding_log[decoding_step][i].cls = initial_top_classes[i];
167 }
168 decoding_step++;
169
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;
178
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);
185
186#ifdef BEAM_DECODER_PROFILE
187 auto start_next_step = clock::now();
188#endif
189
190 while (decoding_step < max_output_len)
191 {
192 size_t pos = 0;
193 for (size_t i = 0; i < beam_size; ++i)
194 {
195 if (stop_idx == decoding_log[decoding_step - 1][i].cls)
196 {
197 continue;
198 }
199 //embd.get_direct(decoding_log[decoding_step - 1][i].cls, input, 0, 0);
200 Eigen::Ref<const M> input = embd.get_ref_by_idx(decoding_log[decoding_step - 1][i].cls);
201 step_with_decoding(layer, linear, input, output, s, lin_out,
202 hidden_size, beam_size,
203 pos, /*i * beam_size,*/
204 top_classes, temp_indicies, logprob, states_c, states_h, i);
205 for (size_t j = 0; j < beam_size; ++j)
206 {
207 logprob[pos /*i * beam_size*/ + j] += initial_logprob[i];
208 }
209 pos += beam_size;
210 }
211
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];
215 });
216
217 /*
218 for (size_t i = 0; i < beam_size * beam_size; ++i)
219 {
220 size_t idx = indices[i];
221 std::cerr << i << " " << idx << " "
222 << initial_top_classes[idx % beam_size] << " " << top_classes[idx]
223 << " "
224 << initial_logprob[idx % beam_size] << " " << logprob[idx] << std::endl;
225 }
226 */
227
228 for (size_t i = 0; i < beam_size; ++i)
229 {
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);
234 }
235
236 size_t non_eos = 0;
237 for (size_t i = 0; i < std::min(beam_size, pos); ++i)
238 {
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)
242 {
243 non_eos++;
244 }
245 decoding_log[decoding_step][i].prev_pos = idx % beam_size;
246 initial_logprob[i] = logprob[idx];
247 }
248
249 decoding_step++;
250
251 if (0 == non_eos)
252 {
253 break; // all beam_size hypotheses gave EOS
254 }
255
256 //std::cerr << std::endl;
257 }
258
259 /*
260 for (size_t p = 0; p < beam_size; ++p)
261 {
262 for (size_t step = 0; step < decoding_log.size(); ++step)
263 {
264 char32_t ch = embd.decode(decoding_log[step][p].cls);
265 std::cerr << decoding_log[step][p].prev_pos
266 << ":"
267 << decoding_log[step][p].cls
268 << ":"
269 << (char)ch
270 << " ";
271
272 }
273 std::cerr << std::endl;
274 }
275
276 std::cerr << std::endl;
277 */
278
279#ifdef BEAM_DECODER_PROFILE
280 auto start_backtracking = clock::now();
281#endif
282
283 // Backtracking
284 size_t step = 0;
285 size_t pos = 0;
286 while (step < decoding_log.size() && decoding_log[step][pos].cls != 1) step++;
287
288 if (step > 0)
289 {
290 step--;
291 }
292
293 while (true)
294 {
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;
298 if (step == 0) // break out here because step is unsigned int
299 break;
300 step--;
301 }
302 std::reverse(output_seq.begin(), output_seq.end());
303
304#ifdef BEAM_DECODER_PROFILE
305 auto end_backtracking = clock::now();
306
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;
312#endif
313
314 return 0;
315 }
316
317protected:
318
320 const params_lstm_t<M, V>& layer, // [in]
321 const params_linear_t<M, V>& linear, // [in]
322 Eigen::Ref<const M>& input, // [in]
323 M& output, // [temp]
324 V& s, // [temp]
325 V& temp, // [temp]
326 const size_t hidden_size, // [in]
327 const size_t beam_size, // [in]
328 const size_t start_pos, // [in]
329 std::vector<uint32_t>& top_classes, // [out]
330 std::vector<uint32_t> indices, // [temp]
331 std::vector<float>& logprob, // [out]
332 V& c, // [in/out]
333 V& h // [in/out]
334 )
335 {
336 // Forward pass
337 s = input.col(0);
338 s.noalias() += layer.weight_hh * h;
339
340 step_fw(hidden_size, 0, s, c, output);
341 h = output.col(0).topRows(hidden_size);
342
343 linear_and_decoding(linear, output, temp, beam_size, start_pos, top_classes, indices, logprob);
344 }
345
347 const params_lstm_t<M, V>& layer, // [in]
348 const params_linear_t<M, V>& linear, // [in]
349 Eigen::Ref<const M>& input, // [in]
350 M& output, // [temp]
351 V& s, // [in]
352 V& temp, // [temp]
353 const size_t hidden_size, // [in]
354 const size_t beam_size, // [in]
355 const size_t start_pos, // [in]
356 std::vector<uint32_t>& top_classes, // [out]
357 std::vector<uint32_t> indices, // [temp]
358 std::vector<float>& logprob, // [out]
359 M& states_c, // [in/out]
360 M& states_h, // [in/out]
361 const size_t beam_idx // [in]
362 )
363 {
364 // Forward pass
365 s = input.col(0);
366 s.noalias() += layer.weight_hh * states_h.col(beam_idx);
367
368 step_fw(hidden_size, 0, s, states_c, beam_idx, output);
369 states_h.col(beam_idx) = output.col(0).topRows(hidden_size);
370
371 linear_and_decoding(linear, output, temp, beam_size, start_pos, top_classes, indices, logprob);
372 }
373
375 const params_linear_t<M, V>& linear, // [in]
376 const M& output, // [in]
377 V& temp, // [temp]
378 const size_t beam_size, // [in]
379 const size_t start_pos, // [in]
380 std::vector<uint32_t>& top_classes, // [out]
381 std::vector<uint32_t> indices, // [temp]
382 std::vector<float>& logprob // [out]
383 )
384 {
385 // Linear
386 temp = linear.bias;
387 temp.noalias() += linear.weight * output;
388 Eigen::Index idx = 0;
389 typename V::Scalar max_value = temp.maxCoeff(&idx);
390 //Eigen::ArrayXf all_logprob = lin_out.array() - (max_value + Eigen::exp(lin_out.array() - max_value).sum());
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;
393
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];
397 });
398
399 for (size_t i = 0; i < beam_size; ++i)
400 {
401 logprob[start_pos + i] = temp[top_classes[start_pos + i]];
402 /*
403 std::cerr << i << " " << top_classes[start_pos + i]
404 << " " << all_logprob[top_classes[start_pos + i]]
405 << " " << lin_out[top_classes[start_pos + i]] << std::endl;
406 */
407 }
408 }
409
410 #define TANH(x) ( T(2) / ( T(1) + Eigen::exp( T(-2) * (x).array() ) ) - T(1) )
411 #define SIGMOID(x) ( T(1) / ( T(1) + Eigen::exp( - (x).array() ) ) )
412
413 inline void update_c(
414 const size_t hidden_size, // [in] the size of the hidden state
415 const V& s, // [in] the results before gating
416 V& c // [in/out] the cell state
417 )
418 {
419 // c = c * forget_gate = c * 1 / sigmoid(...)
420 c = c.array() / (1 + Eigen::exp( - s.segment(hidden_size, hidden_size).array() ) );
421
422 // c = c + input_gate * update_gate
423 c += ( SIGMOID(s.segment(0, hidden_size)).array() * TANH(s.segment(hidden_size * 2, hidden_size)).array() ).matrix();
424 }
425
426 inline void update_c(
427 const size_t hidden_size, // [in] the size of the hidden state
428 const V& s, // [in] the results before gating
429 M& c, // [in/out] the cell state
430 const size_t beam_idx
431 )
432 {
433 // c = c * forget_gate = c * 1 / sigmoid(...)
434 c.col(beam_idx) = c.col(beam_idx).array() / (1 + Eigen::exp( - s.segment(hidden_size, hidden_size).array() ) );
435
436 // c = c + input_gate * update_gate
437 c.col(beam_idx) += ( SIGMOID(s.segment(0, hidden_size)).array() * TANH(s.segment(hidden_size * 2, hidden_size)).array() ).matrix();
438 }
439
440 inline void step_fw(
441 const size_t hidden_size, // [in] the size of the hidden state
442 const size_t t, // [in] step (the position in output)
443 const V& s, // [in] the results before gating
444 V& c, // [in/out] the cell state
445 M& output // [out] the matrix of output states
446 )
447 {
448 update_c(hidden_size, s, c);
449
450 // output = output_gate * tanh(c)
451 output.col(t).topRows(hidden_size).noalias() = (SIGMOID(s.segment(hidden_size * 3, hidden_size)).array() * TANH(c).array()).matrix();
452 }
453
454 inline void step_fw(
455 const size_t hidden_size, // [in] the size of the hidden state
456 const size_t t, // [in] step (the position in output)
457 const V& s, // [in] the results before gating
458 M& c, // [in/out] the cell state
459 const size_t beam_idx,
460 M& output // [out] the matrix of output states
461 )
462 {
463 update_c(hidden_size, s, c, beam_idx);
464
465 // output = output_gate * tanh(c)
466 output.col(t).topRows(hidden_size).noalias() = (SIGMOID(s.segment(hidden_size * 3, hidden_size)).array() * TANH(c.col(beam_idx)).array()).matrix();
467 }
468};
469
470} // namespace eigen_impl
471} // namespace deeplima
472
473#endif
#define TANH(x)
#define SIGMOID(x)
I lookup(const K key) const
Definition embd_dict.h:108
K decode(const I idx) const
Definition embd_dict.h:102
Eigen::Ref< const M > get_ref_by_key(const K key) const
Definition embd_dict.h:79
Eigen::Ref< const M > get_ref_by_idx(const I idx) const
Definition embd_dict.h:74
void step_with_decoding(const params_lstm_t< M, V > &layer, const params_linear_t< M, V > &linear, Eigen::Ref< const M > &input, M &output, V &s, V &temp, const size_t hidden_size, const size_t beam_size, const size_t start_pos, std::vector< uint32_t > &top_classes, std::vector< uint32_t > indices, std::vector< float > &logprob, V &c, V &h)
void linear_and_decoding(const params_linear_t< M, V > &linear, const M &output, V &temp, const size_t beam_size, const size_t start_pos, std::vector< uint32_t > &top_classes, std::vector< uint32_t > indices, std::vector< float > &logprob)
void update_c(const size_t hidden_size, const V &s, M &c, const size_t beam_idx)
void step_fw(const size_t hidden_size, const size_t t, const V &s, V &c, M &output)
virtual size_t execute(std::shared_ptr< Op_Base::workbench_t > pwb, const EmbdUInt64Float &embd, const V &initial_state_h, const V &initial_state_c, const std::shared_ptr< param_base_t > params, const uint32_t start_code, const uint32_t stop_code, const size_t beam_size, std::vector< uint32_t > &output_seq, const size_t max_output_len)
params_lstm_beam_decoder_t< M, V > params_t
void step_fw(const size_t hidden_size, const size_t t, const V &s, M &c, const size_t beam_idx, M &output)
void update_c(const size_t hidden_size, const V &s, V &c)
virtual void precompute_inputs(const std::shared_ptr< param_base_t > params, const M &inputs, M &outputs, int64_t first_column)
virtual std::shared_ptr< Op_Base::workbench_t > create_workbench(uint32_t, const std::shared_ptr< param_base_t > params, bool precomputed_input=false) const override
void step_with_decoding(const params_lstm_t< M, V > &layer, const params_linear_t< M, V > &linear, Eigen::Ref< const M > &input, M &output, V &s, V &temp, const size_t hidden_size, const size_t beam_size, const size_t start_pos, std::vector< uint32_t > &top_classes, std::vector< uint32_t > indices, std::vector< float > &logprob, M &states_c, M &states_h, const size_t beam_idx)
workbench_t(uint32_t hidden_size, const std::vector< uint32_t > &output_sizes, bool precomputed_input=false)