LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
segmentation_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_SRC_TRAIN_SEGMENTATION_WRAPPER_H
7#define DEEPLIMA_SRC_TRAIN_SEGMENTATION_WRAPPER_H
8
9#include <string>
10
12
13namespace deeplima
14{
15namespace segmentation
16{
17namespace impl
18{
19
20// This class does nothing.
21// It only defines the interface of the inference module
22
23template <class M>
25{
26public:
30
31 inline void load(const std::string& fn)
32 {
33 m_impl.load(fn);
34 }
35
36 inline size_t init_new_worker(size_t input_len, bool precomputed_input=false)
37 {
38 return m_impl.init_new_worker(input_len, precomputed_input);
39 }
40
41 inline void predict(
42 size_t worker_id,
43 const typename M::tensor_t& inputs,
44 int64_t input_begin,
45 int64_t input_end,
46 int64_t output_begin,
47 int64_t output_end,
48 std::shared_ptr< StdMatrix<uint8_t> >& output,
49 const std::vector<size_t>& /*lengths*/,
50 const std::vector<std::string>& output_names
51 )
52 {
53 m_impl.predict(worker_id, inputs,
54 input_begin, input_end,
55 output_begin, output_end,
56 output,
57 output_names);
58 }
59
60 inline const std::vector<ngram_descr_t>& get_ngram_descr() const
61 {
62 return m_impl.get_ngram_descr();
63 }
64
65 inline const typename M::dicts_holder_t& get_dicts() const
66 {
67 return m_impl.get_input_uint_dicts();
68 }
69
70 inline const std::vector<std::string>& get_output_str_dicts_names() const
71 {
72 return m_impl.get_output_str_dicts_names();
73 }
74
75 inline const std::vector<std::vector<std::string>>& get_output_str_dicts() const
76 {
77 return m_impl.get_output_str_dicts();
78 }
79
80protected:
81
83};
84
85} // namespace impl
86} // namespace segmenation
87} // namespace deeplima
88
89#endif
const std::vector< std::string > & get_output_str_dicts_names() const
size_t init_new_worker(size_t input_len, bool precomputed_input=false)
const std::vector< std::vector< std::string > > & get_output_str_dicts() const
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)
const std::vector< ngram_descr_t > & get_ngram_descr() const