105 std::vector<params_linear_t<M, V>>
linear;
107#ifdef MATMUL_WITH_FACTORIZATION
110 Eigen::PartialPivLU<M> matmul_input;
111 Eigen::PartialPivLU<M> matmul_forget;
112 Eigen::PartialPivLU<M> matmul_update;
113 Eigen::PartialPivLU<M> matmul_output;
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();
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();
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)
157 auto val = src(row, col);
164 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
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
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;
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;
199 std::vector<uint32_t> output_sizes;
200 output_sizes.reserve(linear.size());
201 for (
const auto& p : linear )
203 output_sizes.push_back(p.weight.rows());
205 return std::make_shared<workbench_t>(input_size, layer.fw.weight_ih.rows() / 4, output_sizes, precomputed_input);
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,
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;
254#ifdef MATMUL_WITH_FACTORIZATION
255 const auto& mul_fw = params.mul_fw;
256 const auto& mul_bw = params.mul_bw;
259 auto wb = std::dynamic_pointer_cast<workbench_t>(pwb);
262 const V& zero = wb->zero;
264 size_t hidden_size = layer.fw.weight_ih.rows() / 4;
266 V s = V::Zero(hidden_size * 4);
268 Eigen::Ref<const M> input = input_matrix.block(0, input_begin, input_matrix.rows(), temp.cols());
270 if (wb->m_precomputed_input)
272#ifdef MATMUL_WITH_FACTORIZATION
273 forward_pass(hidden_size, input, layer.fw.weight_hh, mul_fw, s, c, output);
275 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
279 else if constexpr (std::is_floating_point_v<AuxScalar>)
281 forward_pass(hidden_size, input, layer.fw.weight_hh, s, c, output);
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");
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);
295 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
299 else if constexpr (std::is_floating_point_v<AuxScalar>)
301 backward_pass(hidden_size, input, layer.bw.weight_hh, s, c, output);
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");
313 temp.topRows(hidden_size * 4) = (layer.fw.weight_ih * input).colwise() + layer.fw.bias_ih;
315 temp.bottomRows(hidden_size * 4) = (layer.bw.weight_ih * input).colwise() + layer.bw.bias_ih;
317#ifdef MATMUL_WITH_FACTORIZATION
318 forward_pass(hidden_size, temp, layer.fw.weight_hh, mul_fw, s, c, output);
320 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
324 else if constexpr (std::is_floating_point_v<AuxScalar>)
326 forward_pass(hidden_size, temp, layer.fw.weight_hh, s, c, output);
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");
336 c = V::Zero(hidden_size);
338#ifdef MATMUL_WITH_FACTORIZATION
339 backward_pass(hidden_size, temp, layer.bw.weight_hh, mul_bw, s, c, output);
341 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
345 else if constexpr (std::is_floating_point_v<AuxScalar>)
347 backward_pass(hidden_size, temp, layer.bw.weight_hh, s, c, output);
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");
357#ifdef BILSTM_AND_DENSE_PROFILE
358 const auto before_linear = std::chrono::high_resolution_clock::now();
383 M& linear_output = wb->lin_out[0];
384 for (
size_t j = 0; j < final_output.size(); j++)
386 linear_output.noalias() = (linear[j].weight * output).colwise() + linear[j].bias;
389 for (
auto i = output_begin - input_begin; i < output_end - input_begin; i++)
391 Eigen::Index idx = 0;
395 linear_output.col(i).maxCoeff(&idx);
397 assert(idx < std::numeric_limits<uint8_t>::max());
398 final_output[j][input_begin+i] = (uint8_t) idx;
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() <<
" "
437 const size_t hidden_size,
438 Eigen::Ref<const M> input,
440#ifdef MATMUL_WITH_FACTORIZATION
441 const typename params_t::multiplier_t& mul,
448 step_fw(hidden_size, 0, input.col(0).topRows(hidden_size * 4), c, output);
450#ifdef MATMUL_WITH_FACTORIZATION
454#ifdef BILSTM_AND_DENSE_PROFILE
455 std::chrono::duration<double> sum1(0), sum2(0);
457 for (Eigen::Index t = 1; t < input.cols(); t++)
459#ifdef BILSTM_AND_DENSE_PROFILE
460 const auto before1 = std::chrono::high_resolution_clock::now();
463#ifdef MATMUL_WITH_FACTORIZATION
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);
488 s += input.col(t).topRows(hidden_size * 4);
490 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
492 s = input.col(t).topRows(hidden_size * 4);
496 else if constexpr (std::is_floating_point_v<AuxScalar>)
498 s = input.col(t).topRows(hidden_size * 4);
499 s.noalias() += weight_hh * output.block(0, t-1, hidden_size, 1);
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");
511#ifdef BILSTM_AND_DENSE_PROFILE
512 const auto after1 = std::chrono::high_resolution_clock::now();
513 sum1 += after1 - before1;
515 const auto before2 = std::chrono::high_resolution_clock::now();
517 step_fw(hidden_size, t, s, c, output);
519#ifdef BILSTM_AND_DENSE_PROFILE
520 const auto after2 = std::chrono::high_resolution_clock::now();
521 sum2 += after2 - before2;
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()
534 const size_t hidden_size,
535 Eigen::Ref<const M> input,
537#ifdef MATMUL_WITH_FACTORIZATION
538 const typename params_t::multiplier_t& mul,
545 int t = input.cols() - 1;
546 step_bw(hidden_size, t, input.col(t).bottomRows(hidden_size * 4), c, output);
548 for (t = input.cols() - 2; t >= 0; t--)
550 if constexpr (std::is_integral_v<AuxScalar> && std::is_signed_v<AuxScalar>)
552 s = input.col(t).bottomRows(hidden_size * 4);
556 else if constexpr (std::is_floating_point_v<AuxScalar>)
558 s = input.col(t).bottomRows(hidden_size * 4);
559 s.noalias() += weight_hh * output.block(hidden_size, t+1, hidden_size, 1);
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");
566 step_bw(hidden_size, t, s, c, output);