LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
tagging_wrapper.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_TAGGING_IMPL_TAGGING_WRAPPER_H
7#define DEEPLIMA_TAGGING_IMPL_TAGGING_WRAPPER_H
8
9#include <string>
10
11namespace deeplima
12{
13namespace tagging
14{
15namespace impl
16{
17
18// This class does nothing.
19// It only defines the interface of the inference module
20
21template <class M>
23{
24public:
28
29 inline void load(const std::string& fn)
30 {
31 m_impl.load(fn);
32 }
33
34 inline size_t init_new_worker(size_t input_len, bool precomputed_input=false)
35 {
36 return m_impl.init_new_worker(input_len, precomputed_input);
37 }
38
39 inline size_t get_precomputed_dim() const
40 {
41 return m_impl.get_precomputed_dim();
42 }
43
44 virtual void get_classes_from_fn(const std::string& fn, std::vector<std::string>& class_names, std::vector<std::vector<std::string>>& classes){
45 m_impl.convert_classes_from_fn(fn,class_names, classes);
46 }
47
48 inline void precompute_inputs(
49 const typename M::tensor_t& inputs,
50 typename M::tensor_t& outputs,
51 int64_t input_size
52 )
53 {
54 m_impl.precompute_inputs(inputs, outputs, input_size);
55 }
56
57 inline void predict(
58 size_t worker_id,
59 const typename M::tensor_t& inputs,
60 int64_t input_begin,
61 int64_t input_end,
62 int64_t output_begin,
63 int64_t output_end,
64 std::shared_ptr< StdMatrix<uint8_t> >& output,
65 const std::vector<size_t>& /*lengths*/,
66 const std::vector<std::string>& output_names
67 )
68 {
69 m_impl.predict(worker_id, inputs,
70 input_begin, input_end,
71 output_begin, output_end,
72 output,
73 output_names);
74 }
75
76 inline const typename M::uint_dicts_holder_t& get_input_uint_dicts() const
77 {
78 return m_impl.get_input_uint_dicts();
79 }
80
81 inline const typename M::str_dicts_holder_t& get_input_str_dicts() const
82 {
83 return m_impl.get_input_str_dicts();
84 }
85
86 inline const std::vector<std::string>& get_output_str_dicts_names() const
87 {
88 return m_impl.get_output_str_dicts_names();
89 }
90
91 inline const std::vector<std::vector<std::string>>& get_output_str_dicts() const
92 {
93 return m_impl.get_output_str_dicts();
94 }
95
96 inline const std::string& get_embd_fn(size_t idx) const
97 {
98 return m_impl.get_embd_fn(idx);
99 }
100
101protected:
102
104};
105
106} // namespace impl
107} // namespace tagging
108} // namespace deeplima
109
110#endif
virtual void get_classes_from_fn(const std::string &fn, std::vector< std::string > &class_names, std::vector< std::vector< std::string > > &classes)
void predict(size_t worker_id, const typename M::tensor_t &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< size_t > &, const std::vector< std::string > &output_names)
size_t init_new_worker(size_t input_len, bool precomputed_input=false)
const std::string & get_embd_fn(size_t idx) const
const std::vector< std::vector< std::string > > & get_output_str_dicts() const
const M::str_dicts_holder_t & get_input_str_dicts() const
void precompute_inputs(const typename M::tensor_t &inputs, typename M::tensor_t &outputs, int64_t input_size)
const M::uint_dicts_holder_t & get_input_uint_dicts() const
const std::vector< std::string > & get_output_str_dicts_names() const