48 virtual std::shared_ptr<Op_Base::workbench_t>
create_workbench([[maybe_unused]] uint32_t input_size,
49 [[maybe_unused]]
const std::shared_ptr<param_base_t> params,
52 assert(input_size > 0);
53 assert(
nullptr != params);
57 return std::make_shared<workbench_t>();
60 virtual size_t execute([[maybe_unused]] std::shared_ptr<Op_Base::workbench_t> pwb,
61 const M& input_matrix,
62 const std::shared_ptr<param_base_t> params,
63 const size_t input_begin,
64 const size_t input_end,
65 std::vector<uint32_t>& output)
70 assert(
nullptr != pwb);
71 assert(
nullptr != params);
72 auto layer = std::dynamic_pointer_cast<const params_t>(params);
77 const M input = input_matrix.block(0, 0, input_matrix.rows(), input_end - input_begin);
79 M arc_head = ((layer->m_weight_head * input).colwise() + layer->m_bias_head).transpose();
81 M arc_dep = (layer->m_weight_dep * input).colwise() + layer->m_bias_dep;
87 const V head_bias = arc_head * layer->m_u2;
89 M logits = (arc_head * layer->m_U1) * arc_dep;
90 logits.colwise() += head_bias;
92 for (Eigen::Index i = 0; i < logits.rows(); ++i)
95 logits.row(i).maxCoeff(&idx);
97 assert(idx < std::numeric_limits<uint32_t>::max());
98 output[input_begin + i] = (uint32_t) idx;
104 arborescence<M, uint32_t, typename M::Scalar>(logits, output, input_begin);