LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
segmentation_eigen_inference_impl.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_INCLUDE_SEGMENTATION_SEGMENTATION_EIGEN_INFERENCE_IMPL_H
7#define DEEPLIMA_SRC_INCLUDE_SEGMENTATION_SEGMENTATION_EIGEN_INFERENCE_IMPL_H
8
9#include <eigen3/Eigen/Dense>
10
11#include "bilstm.h"
12#include "bilstm_and_dense.h"
13
14#include "embd_dict.h"
15
18
19namespace deeplima
20{
21namespace segmentation
22{
23namespace eigen_impl
24{
25#ifdef WIN32
26#ifdef SEGM_EXPORTING
27 #define SEGM_EXPORT __declspec(dllexport)
28#else
29 #define SEGM_EXPORT __declspec(dllimport)
30#endif
31#else
32 #define SEGM_EXPORT
33#endif
34
36{
37public:
38 typedef Eigen::MatrixXf Matrix;
39 typedef Eigen::VectorXf Vector;
40 typedef float Scalar;
41 typedef Eigen::MatrixXf tensor_t;
44
45 virtual void load(const std::string& fn)
46 {
47 convert_from_torch(fn);
48 }
49
50 const std::vector<impl::ngram_descr_t>& get_ngram_descr() const
51 {
52 return m_ngram_gescr;
53 }
54
55 virtual void precompute_inputs(
56 const Eigen::MatrixXf& /*inputs*/,
57 Eigen::MatrixXf& /*outputs*/,
58 int64_t /*input_size*/
59 )
60 {
61 throw std::runtime_error("Precomputing isn't supported for segmentation");
62 }
63
64 virtual void predict(
65 size_t worker_id,
66 const Eigen::MatrixXf& inputs,
67 int64_t input_begin,
68 int64_t input_end,
69 int64_t output_begin,
70 int64_t output_end,
71 std::shared_ptr< StdMatrix<uint8_t> >& output,
72 const std::vector<std::string>& /*outputs_names*/
73 )
74 {
75 auto p_op = std::dynamic_pointer_cast<deeplima::eigen_impl::Op_BiLSTM_Dense_ArgMax<Eigen::MatrixXf, Eigen::VectorXf, float>>(Parent::m_ops[0]);
76 assert(Parent::m_wb.size() > 0);
77 assert(worker_id < Parent::m_wb[0].size());
78 p_op->execute(Parent::m_wb[0][worker_id], inputs, Parent::m_params[0], output->m_tensor,
79 input_begin, input_end, output_begin, output_end);
80 }
81
82protected:
83 std::vector<impl::ngram_descr_t> m_ngram_gescr;
84
85 virtual void convert_from_torch(const std::string& fn);
86};
87
88
89} // namespace eigen_impl
90} // namespace segmentation
91} // namespace deeplima
92
93#endif
virtual void precompute_inputs(const Eigen::MatrixXf &, Eigen::MatrixXf &, int64_t)
virtual void predict(size_t worker_id, const Eigen::MatrixXf &inputs, int64_t input_begin, int64_t input_end, int64_t output_begin, int64_t output_end, std::shared_ptr< StdMatrix< uint8_t > > &output, const std::vector< std::string > &)