LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
word_seq_vectorizer.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_LIBS_TASKS_NER_TRAIN_WORD_SEQ_VECTORIZER_H
7#define DEEPLIMA_LIBS_TASKS_NER_TRAIN_WORD_SEQ_VECTORIZER_H
8
9#include <memory>
10
12#include "static_graph/dict.h"
13
16
17namespace deeplima
18{
19namespace nets
20{
21
22template <class DataSet, class StrFeatExtractor, class UIntFeatExtractor, class MatrixInt, class MatrixFloat>
23class WordSeqVectorizerImpl : public vectorizers::WordSeqEmbdVectorizer<DataSet, StrFeatExtractor, UIntFeatExtractor, MatrixFloat>
24{
25public:
27
29 {
30 std::shared_ptr<DictBase> m_dict; // feature extractor
31
32 embeddable_feature_descr_t(typename Parent::feature_type_t type, const std::string& name, int dim, std::shared_ptr<DictBase> dict)
33 : Parent::feature_descr_base_t(type, name, dim), m_dict(dict) {}
34 };
35
36protected:
37 std::vector<embeddable_feature_descr_t> m_embeddable_features;
38
40
41public:
42
43 WordSeqVectorizerImpl(const std::vector<typename Parent::feature_descr_t>& features,
44 const std::vector<embeddable_feature_descr_t>& embeddable_features)
45 : Parent(features),
46 m_embeddable_features(embeddable_features),
48 {
50 }
51
52 WordSeqVectorizerImpl(const std::vector<typename Parent::feature_descr_t>& features,
53 const std::vector<embeddable_feature_descr_t>& embeddable_features,
54 const StrFeatExtractor& str_feat_extractor)
55 : Parent(features, str_feat_extractor),
56 m_embeddable_features(embeddable_features),
58 {
60 }
61
62 const std::vector<deeplima::nets::embd_descr_t> get_embd_descr() const
63 {
64 std::vector<deeplima::nets::embd_descr_t> d;
65 for (const embeddable_feature_descr_t& feat_descr : m_embeddable_features )
66 {
67 assert(feat_descr.m_dim > 0);
68 assert(!feat_descr.m_name.empty());
69 d.emplace_back(feat_descr.m_name, feat_descr.m_dim);
70 }
71 return d;
72 }
73
74 typedef std::pair<std::shared_ptr<MatrixInt>, std::shared_ptr<MatrixFloat>> vectorization_t;
75
76 vectorization_t process(const DataSet& src)
77 {
78 int64_t len = 0;
79
80 typename DataSet::const_iterator i = src.begin();
81 while (src.end() != i)
82 {
83 if ((*i).is_word())
84 {
85 len++;
86 }
87 i++;
88 }
89
90 auto frozen_features = std::make_shared<MatrixFloat>(len, Parent::m_features_size);
91 auto embeddable_features = std::make_shared<MatrixInt>(len, m_embeddable_size);
92 vectorization_t rv(embeddable_features, frozen_features);
93
94 process(src, rv, 0);
95
96 return rv;
97 }
98
99 vectorization_t init_dst(uint64_t len) const
100 {
101 auto frozen_features = std::make_shared<MatrixFloat>(len, Parent::m_features_size);
102 auto embeddable_features = std::make_shared<MatrixInt>(len, m_embeddable_size);
103 vectorization_t rv(embeddable_features, frozen_features);
104
105 return rv;
106 }
107
108 void process(const DataSet& src, vectorization_t dst, uint64_t start) const
109 {
110 std::shared_ptr<MatrixInt> embeddable_features = dst.first;
111 std::shared_ptr<MatrixFloat> frozen_features = dst.second;
112
113 typename DataSet::const_iterator it = src.begin();
114 uint64_t current_timepoint = start;
115 while (src.end() != it)
116 {
117 while(!(*it).is_word() && src.end() != it)
118 {
119 it++;
120 }
121 if (src.end() == it) break;
122
123 vectorize_timepoint(*frozen_features, *embeddable_features, current_timepoint, *it);
124
125 current_timepoint++;
126 if (current_timepoint == std::numeric_limits<uint64_t>::max())
127 {
128 throw std::overflow_error("Too much words in the dataset.");
129 }
130
131 it++;
132 }
133 }
134
135 inline void vectorize_timepoint(MatrixFloat& frozen_features, MatrixInt& embeddable_features,
136 uint64_t timepoint, const typename DataSet::token_t& token) const
137 {
138 Parent::vectorize_timepoint(frozen_features, timepoint, token);
139
140 for (size_t i = 0; i < m_embeddable_features.size(); i++)
141 {
143 switch (feat_descr.m_type)
144 {
146 throw std::runtime_error("Unsupported");
147 break;
149 throw std::runtime_error("Unsupported");
150 break;
152 {
153 size_t ifeat = Parent::m_str_feat_extractor.get_feat_id(feat_descr.m_name);
154 const std::string& feat_val = Parent::m_str_feat_extractor.feat_value(token, ifeat);
155 std::shared_ptr<StringDict> dict
156 = std::dynamic_pointer_cast<StringDict, DictBase>(feat_descr.m_dict);
157 uint64_t idx = dict->get_idx(feat_val);
158 embeddable_features.set(timepoint, i, idx);
159 }
160 break;
161 default:
162 throw std::runtime_error("Unknown argument type");
163 }
164 }
165 }
166};
167
168} // namespace nets
169} // namespace deeplima
170
171#endif
std::pair< std::shared_ptr< MatrixInt >, std::shared_ptr< MatrixFloat > > vectorization_t
const std::vector< deeplima::nets::embd_descr_t > get_embd_descr() const
WordSeqVectorizerImpl(const std::vector< typename Parent::feature_descr_t > &features, const std::vector< embeddable_feature_descr_t > &embeddable_features)
WordSeqVectorizerImpl(const std::vector< typename Parent::feature_descr_t > &features, const std::vector< embeddable_feature_descr_t > &embeddable_features, const StrFeatExtractor &str_feat_extractor)
std::vector< embeddable_feature_descr_t > m_embeddable_features
void process(const DataSet &src, vectorization_t dst, uint64_t start) const
void vectorize_timepoint(MatrixFloat &frozen_features, MatrixInt &embeddable_features, uint64_t timepoint, const typename DataSet::token_t &token) const
vectorizers::WordSeqEmbdVectorizer< DataSet, StrFeatExtractor, UIntFeatExtractor, MatrixFloat > Parent
vectorization_t process(const DataSet &src)
vectorization_t init_dst(uint64_t len) const
void vectorize_timepoint(MatrixFloat &target, uint64_t timepoint, const typename DataSet::token_t &token) const
embeddable_feature_descr_t(typename Parent::feature_type_t type, const std::string &name, int dim, std::shared_ptr< DictBase > dict)