LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
bilstm_and_dense.h
Go to the documentation of this file.
1// Copyright 2021 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_BILSTM_AND_DENSE_H
7#define DEEPLIMA_SRC_INFERENCE_EIGEN_BILSTM_AND_DENSE_H
8
9#include <vector>
10#include <type_traits>
11
12#include "bilstm.h"
13#include "linear.h"
14
15//#define BILSTM_AND_DENSE_PROFILE
16//#define MATMUL_WITH_FACTORIZATION
17
18#ifdef BILSTM_AND_DENSE_PROFILE
19#include <chrono>
20#endif
21
22namespace deeplima
23{
24namespace eigen_impl
25{
26
46template<typename M, // Basic matrix type
47 typename V, // Basic vector type
48 typename T, // Basic scalar type,
49 typename AuxScalar=float, // Input scalar type:
50 // - float - for normal arithmetics
51 // - int16_t - for fixed point operations where possible
52 uint8_t IFracBits=4, // Fraction bits in input data (used if AuxScalar is integer)
53 uint8_t WFracBits=4 // Fraction bits in weights_hh (used if AuxScalar is integer)
54 >
56{
57protected:
58
59 // For fixed point calculations (used only if AusScalar is integer)
60 static_assert(std::is_signed_v<AuxScalar>);
61 typedef AuxScalar fixed_point_t;
62 using MatrixWeight = typename std::conditional<std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>,
63 Eigen::Matrix<fixed_point_t, Eigen::Dynamic, Eigen::Dynamic, Eigen::RowMajor>,
64 M>::type;
65 using MatrixInput = typename std::conditional<std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>,
66 Eigen::Matrix<fixed_point_t, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor>,
67 M>::type;
68
69 static constexpr fixed_point_t WEIGHT_FRACTION_MULT = fixed_point_t(1) << WFracBits;
70 static constexpr fixed_point_t DATA_FRACTION_MULT = fixed_point_t(1) << IFracBits;
71 static constexpr fixed_point_t WEIGHT_DATA_FRACTION_MULT = fixed_point_t(1) << (WFracBits + IFracBits);
72 // End of fixed point releated definitions
73
75 {
76 workbench_t(uint32_t input_size,
77 uint32_t hidden_size,
78 const std::vector<uint32_t>& output_sizes,
79 bool precomputed_input=false)
80 : temp(M::Zero(precomputed_input ? 0 : hidden_size * 8, input_size)),
81 out(M::Zero(hidden_size * 2, input_size)),
82 zero(V::Zero(hidden_size)),
83 m_precomputed_input(precomputed_input)
84 {
85 lin_out.resize(output_sizes.size());
86 for (size_t i = 0; i < output_sizes.size(); i++)
87 {
88 lin_out[i] = M::Zero(output_sizes[i], input_size);
89 }
90 }
91
92 virtual ~workbench_t() {}
93
95 M out;
96 std::vector<M> lin_out;
97 const V zero;
99 };
100
101public:
103 {
105 std::vector<params_linear_t<M, V>> linear;
106
107#ifdef MATMUL_WITH_FACTORIZATION
108 struct multiplier_t
109 {
110 Eigen::PartialPivLU<M> matmul_input;
111 Eigen::PartialPivLU<M> matmul_forget;
112 Eigen::PartialPivLU<M> matmul_update;
113 Eigen::PartialPivLU<M> matmul_output;
114 };
115
116 multiplier_t mul_fw;
117 multiplier_t mul_bw;
118
119 bool precompute()
120 {
121 // std::cerr << "fw weights size: " << bilstm.fw.weight_hh.rows()
122 // << " x " << bilstm.fw.weight_hh.cols() << std::endl;
123
124 size_t hidden_size = bilstm.fw.weight_hh.cols();
125 // std::cerr << "precompute(fw.input):" << std::endl;
126 // /*
127 mul_fw.matmul_input = bilstm.fw.weight_hh.block(0, 0, hidden_size, hidden_size).inverse().partialPivLu();
128 mul_fw.matmul_forget = bilstm.fw.weight_hh.block(hidden_size, 0, hidden_size, hidden_size).inverse().partialPivLu();
129 mul_fw.matmul_update = bilstm.fw.weight_hh.block(hidden_size*2, 0, hidden_size, hidden_size).inverse().partialPivLu();
130 mul_fw.matmul_output = bilstm.fw.weight_hh.block(hidden_size*3, 0, hidden_size, hidden_size).inverse().partialPivLu();
131 // */
132
133 hidden_size = bilstm.bw.weight_hh.cols();
134 // std::cerr << "precompute(bw.input):" << std::endl;
135 // /*
136 mul_bw.matmul_input = bilstm.bw.weight_hh.block(0, 0, hidden_size, hidden_size).inverse().partialPivLu();
137 mul_bw.matmul_forget = bilstm.bw.weight_hh.block(hidden_size, 0, hidden_size, hidden_size).inverse().partialPivLu();
138 mul_bw.matmul_update = bilstm.bw.weight_hh.block(hidden_size*2, 0, hidden_size, hidden_size).inverse().partialPivLu();
139 mul_bw.matmul_output = bilstm.bw.weight_hh.block(hidden_size*3, 0, hidden_size, hidden_size).inverse().partialPivLu();
140 // */
141
142 // std::cerr << "end of precomputing" << std::endl;
143 return true;
144 }
145#else
146 // These two matrices of weights are used (and filled) only if we are going
147 // to use fixed point multiplication of the recurrent state.
150
151 void convert_matrix(const M& src, MatrixWeight& dst)
152 {
153 dst = MatrixWeight::Zero(src.rows(), src.cols());
154 for (Eigen::Index row = 0; row < src.rows(); ++row)
155 for (Eigen::Index col = 0; col < src.cols(); ++col)
156 {
157 auto val = src(row, col);
158 dst(row, col) = static_cast<fixed_point_t>((val >= 0.0) ? (val * WEIGHT_FRACTION_MULT + T{0.5}) : (val * WEIGHT_FRACTION_MULT - T{0.5}));
159 }
160 }
161
163 {
164 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
165 {
166 // std::cerr << "Converting hh to fixed_point" << std::endl;
167 // std::cerr << "min(fw_weight_hh) = " << bilstm.fw.weight_hh.minCoeff() << " "
168 // << "max(fw_weight_hh) = " << bilstm.fw.weight_hh.maxCoeff() << std::endl;
170 // std::cerr << "min(fw_weight_hh) = " << static_cast<T>(weight_fw_hh_fixed_point.minCoeff()) / WEIGHT_FRACTION_MULT << " "
171 // << "max(fw_weight_hh) = " << static_cast<T>(weight_fw_hh_fixed_point.maxCoeff()) / WEIGHT_FRACTION_MULT << std::endl;
172
173 // std::cerr << "min(bw_weight_hh) = " << bilstm.bw.weight_hh.minCoeff() << " "
174 // << "max(bw_weight_hh) = " << bilstm.bw.weight_hh.maxCoeff() << std::endl;
176 // std::cerr << "min(bw_weight_hh) = " << static_cast<T>(weight_bw_hh_fixed_point.minCoeff()) / /*WEIGHT_FRACTION_MULT << " "
177 // << "max(bw_weight_hh) = " << static_cast<T>(weight_bw_hh_fixed_point.maxCoeff()) / WEIGHT_FRACTION_MULT << std::endl;*/
178 }
179
180 return true;
181 }
182#endif
183};
184
186
187 virtual std::shared_ptr<Op_Base::workbench_t> create_workbench(uint32_t input_size, const std::shared_ptr<param_base_t> params, bool precomputed_input=false) const override
188 {
189 assert(input_size > 0);
190 assert(nullptr != params);
191 const auto& layer = std::dynamic_pointer_cast<const params_t>(params)->bilstm;
192 const auto& linear = std::dynamic_pointer_cast<const params_t>(params)->linear;
193
194#ifdef MATMUL_WITH_FACTORIZATION
195 const auto& mul_fw = std::dynamic_pointer_cast<const params_t>(params)->mul_fw;
196 const auto& mul_bw = std::dynamic_pointer_cast<const params_t>(params)->mul_bw;
197#endif
198
199 std::vector<uint32_t> output_sizes;
200 output_sizes.reserve(linear.size());
201 for ( const auto& p : linear )
202 {
203 output_sizes.push_back(p.weight.rows());
204 }
205 return std::make_shared<workbench_t>(input_size, layer.fw.weight_ih.rows() / 4, output_sizes, precomputed_input);
206 }
207
208 virtual bool supports_precomputing() const
209 {
210 return true;
211 }
212
213 virtual void precompute_inputs(const std::shared_ptr<param_base_t> params, const M& inputs, M& outputs, int64_t first_column)
214 {
215 assert(nullptr != params);
216 const auto& layer = std::dynamic_pointer_cast<const params_t>(params)->bilstm;
217 size_t hidden_size = layer.fw.weight_ih.rows() / 4;
218
219 auto output_block = outputs.block(0, first_column, outputs.rows(), inputs.cols());
220 // Top rows - forward pass
221 const V fw_bias = layer.fw.bias_ih + layer.fw.bias_hh;
222 output_block.topRows(hidden_size * 4).noalias() = (layer.fw.weight_ih * inputs).colwise() + fw_bias;
223
224 /*
225 std::cerr << "min(precomputed inputs) = " << output_block.topRows(hidden_size * 4).minCoeff() << " "
226 << "max(precomputed inputs) = " << output_block.topRows(hidden_size * 4).maxCoeff() << std::endl;
227 */
228
229 // Bottom rows - backward pass
230 const V bw_bias = layer.bw.bias_ih + layer.bw.bias_hh;
231 output_block.bottomRows(hidden_size * 4).noalias() = (layer.bw.weight_ih * inputs).colwise() + bw_bias;
232
233 /*
234 std::cerr << "min(precomputed inputs) = " << output_block.bottomRows(hidden_size * 4).minCoeff() << " "
235 << "max(precomputed inputs) = " << output_block.bottomRows(hidden_size * 4).maxCoeff() << std::endl;
236 */
237 }
238
239 virtual size_t execute(std::shared_ptr<Op_Base::workbench_t> pwb,
240 const M& input_matrix,
241 const std::shared_ptr<param_base_t> pparams,
242 std::vector<std::vector<uint8_t>>& final_output,
243 size_t input_begin,
244 size_t /*input_end*/,
245 size_t output_begin,
246 size_t output_end)
247 {
248 assert(pwb);
249 assert(pparams);
250 const params_t& params = *std::dynamic_pointer_cast<const params_t>(pparams);
251 const auto& layer = params.bilstm;
252 const auto& linear = params.linear;
253
254#ifdef MATMUL_WITH_FACTORIZATION
255 const auto& mul_fw = params.mul_fw;
256 const auto& mul_bw = params.mul_bw;
257#endif
258
259 auto wb = std::dynamic_pointer_cast<workbench_t>(pwb);
260 M& temp = wb->temp;
261 M& output = wb->out;
262 const V& zero = wb->zero;
263
264 size_t hidden_size = layer.fw.weight_ih.rows() / 4;
265 V c = zero;
266 V s = V::Zero(hidden_size * 4);
267
268 Eigen::Ref<const M> input = input_matrix.block(0, input_begin, input_matrix.rows(), temp.cols());
269
270 if (wb->m_precomputed_input)
271 {
272#ifdef MATMUL_WITH_FACTORIZATION
273 forward_pass(hidden_size, input, layer.fw.weight_hh, mul_fw, s, c, output);
274#else
275 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
276 {
277 forward_pass(hidden_size, input, params.weight_fw_hh_fixed_point, s, c, output);
278 }
279 else if constexpr (std::is_floating_point_v<AuxScalar>)
280 {
281 forward_pass(hidden_size, input, layer.fw.weight_hh, s, c, output);
282 }
283 else
284 {
285 static_assert((std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
286 || std::is_floating_point_v<AuxScalar>, "AuxScalar must be a signed integer or floating-point type");
287 }
288#endif
289
290 V c_fw = c;
291 c = V::Zero(hidden_size);
292#ifdef MATMUL_WITH_FACTORIZATION
293 backward_pass(hidden_size, input, layer.bw.weight_hh, mul_bw, s, c, output);
294#else
295 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
296 {
297 backward_pass(hidden_size, input, params.weight_bw_hh_fixed_point, s, c, output);
298 }
299 else if constexpr (std::is_floating_point_v<AuxScalar>)
300 {
301 backward_pass(hidden_size, input, layer.bw.weight_hh, s, c, output);
302 }
303 else
304 {
305 static_assert((std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
306 || std::is_floating_point_v<AuxScalar>, "AuxScalar must be a signed integer or floating-point type");
307 }
308#endif
309 }
310 else
311 {
312 // Top rows - forward pass
313 temp.topRows(hidden_size * 4) = (layer.fw.weight_ih * input).colwise() + layer.fw.bias_ih;
314 // Bottom rows - backward pass
315 temp.bottomRows(hidden_size * 4) = (layer.bw.weight_ih * input).colwise() + layer.bw.bias_ih;
316
317#ifdef MATMUL_WITH_FACTORIZATION
318 forward_pass(hidden_size, temp, layer.fw.weight_hh, mul_fw, s, c, output);
319#else
320 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
321 {
322 forward_pass(hidden_size, temp, params.weight_fw_hh_fixed_point, s, c, output);
323 }
324 else if constexpr (std::is_floating_point_v<AuxScalar>)
325 {
326 forward_pass(hidden_size, temp, layer.fw.weight_hh, s, c, output);
327 }
328 else
329 {
330 static_assert((std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
331 || std::is_floating_point_v<AuxScalar>, "AuxScalar must be a signed integer or floating-point type");
332 }
333#endif
334
335 V c_fw = c;
336 c = V::Zero(hidden_size);
337
338#ifdef MATMUL_WITH_FACTORIZATION
339 backward_pass(hidden_size, temp, layer.bw.weight_hh, mul_bw, s, c, output);
340#else
341 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
342 {
343 backward_pass(hidden_size, temp, params.weight_bw_hh_fixed_point, s, c, output);
344 }
345 else if constexpr (std::is_floating_point_v<AuxScalar>)
346 {
347 backward_pass(hidden_size, temp, layer.bw.weight_hh, s, c, output);
348 }
349 else
350 {
351 static_assert((std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
352 || std::is_floating_point_v<AuxScalar>, "AuxScalar must be a signed integer or floating-point type");
353 }
354#endif
355 }
356
357#ifdef BILSTM_AND_DENSE_PROFILE
358 const auto before_linear = std::chrono::high_resolution_clock::now();
359#endif
360
361/*
362 M lo = linear[0].weight * output.col(0);
363 for (auto i = output_begin - input_begin; i < output_end - input_begin; i++)
364 {
365 for (size_t j = 0; j < final_output.size(); ++j)
366 {
367 //lo = (linear[j].weight * output.col(i)).colwise() + linear[j].bias;
368 Eigen::Index idx = 0;
369 // TODO scalar v value is not used but maxCoeff has side effect. It must be kept. Should we remove the
370 // return value or use it somewhat?
371 // typename M::Scalar v =
372 ((linear[j].weight * output.col(i)) + linear[j].bias).col(0).maxCoeff(&idx);
373
374 //lo.col(0).maxCoeff(&idx);
375 assert(idx >= 0);
376 assert(idx < std::numeric_limits<uint8_t>::max());
377 final_output[j][input_begin+i] = (uint8_t) idx;
378 }
379 }
380*/
381
382 // Linear layer on top of RNN outputs
383 M& linear_output = wb->lin_out[0];
384 for (size_t j = 0; j < final_output.size(); j++)
385 {
386 linear_output.noalias() = (linear[j].weight * output).colwise() + linear[j].bias;
387 //linear_output.colwise() += linear[j].bias;
388
389 for (auto i = output_begin - input_begin; i < output_end - input_begin; i++)
390 {
391 Eigen::Index idx = 0;
392 // TODO scalar v value is not used but maxCoeff has side effect. It must be kept. Should we remove the
393 // return value or use it somewhat?
394 // typename M::Scalar v =
395 linear_output.col(i).maxCoeff(&idx);
396 assert(idx >= 0);
397 assert(idx < std::numeric_limits<uint8_t>::max());
398 final_output[j][input_begin+i] = (uint8_t) idx;
399 }
400 }
401
402#ifdef BILSTM_AND_DENSE_PROFILE
403 const auto after_linear = std::chrono::high_resolution_clock::now();
404 std::cerr << "output.cols()==" << output.cols() << " "
405 << "linear time: " << std::chrono::duration_cast<std::chrono::microseconds>(after_linear - before_linear).count() << " "
406 << std::endl;
407#endif
408
409
410 // Linear layer on top of RNN outputs - alternative way to calculate (not faster)
411 /*for (size_t j = 0; j < final_output.size(); j++)
412 {
413 M& linear_output = wb->lin_out[j];
414 linear_output = (linear[j].weight * output).colwise() + linear[j].bias;
415 }
416
417 for (size_t j = 0; j < final_output.size(); j++)
418 {
419 M& linear_output = wb->lin_out[j];
420
421 for (Eigen::Index i = output_begin - input_begin;
422 i < output_end - input_begin; i++)
423 {
424 Eigen::Index idx = 0;
425 typename M::Scalar v = linear_output.col(i).maxCoeff(&idx);
426 assert(idx >= 0);
427 assert(idx < std::numeric_limits<uint8_t>::max());
428 final_output[j][input_begin+i] = (uint8_t) idx;
429 }
430 }*/
431
432 return 0;
433 }
434
435protected:
436 inline void forward_pass(
437 const size_t hidden_size, // [in] the size of the hidden state
438 Eigen::Ref<const M> input, // [in] the pre-computed inputs
439 const MatrixWeight& weight_hh, // [in] the weights matrix
440#ifdef MATMUL_WITH_FACTORIZATION
441 const typename params_t::multiplier_t& mul,
442#endif
443 V& s, // [temp] preallocated space for results before gating
444 V& c, // [in/out] the cell state
445 M& output // [out] the output states
446 )
447 {
448 step_fw(hidden_size, 0, input.col(0).topRows(hidden_size * 4), c, output);
449
450#ifdef MATMUL_WITH_FACTORIZATION
451 V prev_state;
452#endif
453
454#ifdef BILSTM_AND_DENSE_PROFILE
455 std::chrono::duration<double> sum1(0), sum2(0);
456#endif
457 for (Eigen::Index t = 1; t < input.cols(); t++)
458 {
459#ifdef BILSTM_AND_DENSE_PROFILE
460 const auto before1 = std::chrono::high_resolution_clock::now();
461#endif
462
463#ifdef MATMUL_WITH_FACTORIZATION
464 /*
465 s.segment(0, hidden_size).noalias() = mul.matmul_input.permutationP() * output.block(0, t-1, hidden_size, 1);
466 mul.matmul_input.matrixLU().template triangularView<Eigen::UnitLower>().solveInPlace(s.segment(0, hidden_size));
467 mul.matmul_input.matrixLU().template triangularView<Eigen::Upper>().solveInPlace(s.segment(0, hidden_size));
468
469 s.segment(hidden_size, hidden_size) = mul.matmul_forget.permutationP() * output.block(0, t-1, hidden_size, 1);
470 mul.matmul_forget.matrixLU().template triangularView<Eigen::UnitLower>().solveInPlace(s.segment(hidden_size, hidden_size));
471 mul.matmul_forget.matrixLU().template triangularView<Eigen::Upper>().solveInPlace(s.segment(hidden_size, hidden_size));
472
473 s.segment(hidden_size*2, hidden_size) = mul.matmul_update.permutationP() * output.block(0, t-1, hidden_size, 1);
474 mul.matmul_update.matrixLU().template triangularView<Eigen::UnitLower>().solveInPlace(s.segment(hidden_size*2, hidden_size));
475 mul.matmul_update.matrixLU().template triangularView<Eigen::Upper>().solveInPlace(s.segment(hidden_size*2, hidden_size));
476
477 s.segment(hidden_size*3, hidden_size) = mul.matmul_output.permutationP() * output.block(0, t-1, hidden_size, 1);
478 mul.matmul_output.matrixLU().template triangularView<Eigen::UnitLower>().solveInPlace(s.segment(hidden_size*3, hidden_size));
479 mul.matmul_output.matrixLU().template triangularView<Eigen::Upper>().solveInPlace(s.segment(hidden_size*3, hidden_size));
480 */
481
482 prev_state = output.block(0, t-1, hidden_size, 1);
483 s.segment(0, hidden_size) = mul.matmul_input.solve(prev_state);
484 s.segment(hidden_size, hidden_size) = mul.matmul_forget.solve(prev_state);
485 s.segment(hidden_size*2, hidden_size) = mul.matmul_update.solve(prev_state);
486 s.segment(hidden_size*3, hidden_size) = mul.matmul_output.solve(prev_state);
487
488 s += input.col(t).topRows(hidden_size * 4);
489#else
490 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
491 {
492 s = input.col(t).topRows(hidden_size * 4);
493 MatrixInput arg2 = (output.block(0, t-1, hidden_size, 1) * DATA_FRACTION_MULT).template cast<fixed_point_t>() ;
494 s.noalias() += ((weight_hh * arg2).template cast<T>() / WEIGHT_DATA_FRACTION_MULT);
495 }
496 else if constexpr (std::is_floating_point_v<AuxScalar>)
497 {
498 s = input.col(t).topRows(hidden_size * 4);
499 s.noalias() += weight_hh * output.block(0, t-1, hidden_size, 1);
500
501 //s = weight_hh * output.block(0, t-1, hidden_size, 1);
502 //s += input.col(t).topRows(hidden_size * 4);
503 }
504 else
505 {
506 static_assert((std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
507 || std::is_floating_point_v<AuxScalar>, "AuxScalar must be a signed integer or floating-point type");
508 }
509#endif
510
511#ifdef BILSTM_AND_DENSE_PROFILE
512 const auto after1 = std::chrono::high_resolution_clock::now();
513 sum1 += after1 - before1;
514
515 const auto before2 = std::chrono::high_resolution_clock::now();
516#endif
517 step_fw(hidden_size, t, s, c, output);
518
519#ifdef BILSTM_AND_DENSE_PROFILE
520 const auto after2 = std::chrono::high_resolution_clock::now();
521 sum2 += after2 - before2;
522#endif
523 }
524
525#ifdef BILSTM_AND_DENSE_PROFILE
526 std::cerr << "input.cols()==" << input.cols() << " "
527 << "step_fw avg time: " << std::chrono::duration_cast<std::chrono::microseconds>(sum2).count() << " "
528 << "in+recc avg time: " << std::chrono::duration_cast<std::chrono::microseconds>(sum1).count()
529 << std::endl;
530#endif
531 }
532
533 inline void backward_pass(
534 const size_t hidden_size, // [in] the size of the hidden state
535 Eigen::Ref<const M> input, // [in] the pre-computed inputs
536 const MatrixWeight& weight_hh, // [in] the weights matrix
537#ifdef MATMUL_WITH_FACTORIZATION
538 const typename params_t::multiplier_t& mul,
539#endif
540 V& s, // [temp] preallocated space for results before gating
541 V& c, // [in/out] the cell state
542 M& output // [out] the output states
543 )
544 {
545 int t = input.cols() - 1;
546 step_bw(hidden_size, t, input.col(t).bottomRows(hidden_size * 4), c, output);
547
548 for (t = input.cols() - 2; t >= 0; t--)
549 {
550 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
551 {
552 s = input.col(t).bottomRows(hidden_size * 4);
553 MatrixInput arg2 = (output.block(hidden_size, t+1, hidden_size, 1) * DATA_FRACTION_MULT).template cast<fixed_point_t>() ;
554 s.noalias() += ((weight_hh * arg2).template cast<float>() / WEIGHT_DATA_FRACTION_MULT);
555 }
556 else if constexpr (std::is_floating_point_v<AuxScalar>)
557 {
558 s = input.col(t).bottomRows(hidden_size * 4);
559 s.noalias() += weight_hh * output.block(hidden_size, t+1, hidden_size, 1);
560 }
561 else
562 {
563 static_assert((std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
564 || std::is_floating_point_v<AuxScalar>, "AuxScalar must be a signed integer or floating-point type");
565 }
566 step_bw(hidden_size, t, s, c, output);
567 }
568 }
569
570 #define TANH(x) ( T(2) / ( T(1) + Eigen::exp( T(-2) * (x).array() ) ) - T(1) )
571 #define SIGMOID(x) ( T(1) / ( T(1) + Eigen::exp( - (x).array() ) ) )
572
573 inline void update_c(
574 const size_t hidden_size, // [in] the size of the hidden state
575 const V& s, // [in] the results before gating
576 V& c // [in/out] the cell state
577 )
578 {
579 // c = c * forget_gate = c * 1 / sigmoid(...)
580 c = c.array() / (1 + Eigen::exp( - s.segment(hidden_size, hidden_size).array() ) );
581
582 // c = c + input_gate * update_gate
583 c += ( SIGMOID(s.segment(0, hidden_size)).array() * TANH(s.segment(hidden_size * 2, hidden_size)).array() ).matrix();
584 }
585
586 inline void step_fw(
587 const size_t hidden_size, // [in] the size of the hidden state
588 const size_t t, // [in] step (the position in output)
589 const V& s, // [in] the results before gating
590 V& c, // [in/out] the cell state
591 M& output // [out] the matrix of output states
592 )
593 {
594 update_c(hidden_size, s, c);
595
596 // output = output_gate * tanh(c)
597 output.col(t).topRows(hidden_size).noalias() = (SIGMOID(s.segment(hidden_size * 3, hidden_size)).array() * TANH(c).array()).matrix();
598 }
599
600 inline void step_bw(
601 const size_t hidden_size, // [in] the size of the hidden state
602 const size_t t, // [in] step (the position in output)
603 const V& s, // [in] the results before gating
604 V& c, // [in/out] the cell state
605 M& output // [out] matrix of output states
606 )
607 {
608 update_c(hidden_size, s, c);
609
610 // output = output_gate * tanh(c)
611 output.col(t).bottomRows(hidden_size).noalias() = (SIGMOID(s.segment(hidden_size * 3, hidden_size)).array() * TANH(c).array()).matrix();
612 }
613};
614
615} // namespace eigen_impl
616} // namespace deeplima
617
618#endif
#define TANH(x)
#define SIGMOID(x)
Fusion of several torch modules implemented for inference in Eigen.
typename std::conditional< std::is_integral_v< AuxScalar > &&std::is_signed_v< AuxScalar >, Eigen::Matrix< fixed_point_t, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor >, M >::type MatrixInput
void step_fw(const size_t hidden_size, const size_t t, const V &s, V &c, M &output)
void step_bw(const size_t hidden_size, const size_t t, const V &s, V &c, M &output)
virtual void precompute_inputs(const std::shared_ptr< param_base_t > params, const M &inputs, M &outputs, int64_t first_column)
static constexpr fixed_point_t WEIGHT_DATA_FRACTION_MULT
virtual size_t execute(std::shared_ptr< Op_Base::workbench_t > pwb, const M &input_matrix, const std::shared_ptr< param_base_t > pparams, std::vector< std::vector< uint8_t > > &final_output, size_t input_begin, size_t, size_t output_begin, size_t output_end)
typename std::conditional< std::is_integral_v< AuxScalar > &&std::is_signed_v< AuxScalar >, Eigen::Matrix< fixed_point_t, Eigen::Dynamic, Eigen::Dynamic, Eigen::RowMajor >, M >::type MatrixWeight
static constexpr fixed_point_t WEIGHT_FRACTION_MULT
void update_c(const size_t hidden_size, const V &s, V &c)
void forward_pass(const size_t hidden_size, Eigen::Ref< const M > input, const MatrixWeight &weight_hh, V &s, V &c, M &output)
static constexpr fixed_point_t DATA_FRACTION_MULT
virtual std::shared_ptr< Op_Base::workbench_t > create_workbench(uint32_t input_size, const std::shared_ptr< param_base_t > params, bool precomputed_input=false) const override
void backward_pass(const size_t hidden_size, Eigen::Ref< const M > input, const MatrixWeight &weight_hh, V &s, V &c, M &output)
workbench_t(uint32_t input_size, uint32_t hidden_size, const std::vector< uint32_t > &output_sizes, bool precomputed_input=false)
params_lstm_t< M, V > fw
Definition bilstm.h:51
params_lstm_t< M, V > bw
Definition bilstm.h:52