LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
conllu_file_iterator.h
Go to the documentation of this file.
1// Copyright 2022 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_GRAPH_DP_CONLLU_FILE_ITERATOR_H
7#define DEEPLIMA_LIBS_TASKS_GRAPH_DP_CONLLU_FILE_ITERATOR_H
8
10#include "conllu/treebank.h"
15
16namespace deeplima
17{
18namespace graph_dp
19{
20namespace train
21{
22
24{
25public:
26
28 size_t batch_size,
30 const DictsHolder& morph_tag_dh,
31 std::shared_ptr<FeatureVectorizerBase<>> p_embd,
32 bool add_root = false,
33 std::shared_ptr<const std::map<std::string, int64_t>> deprel2id = nullptr);
34
35 // Target value used in the gold to mark positions that must be ignored by the
36 // loss/accuracy (e.g. the synthetic <ROOT> token, which is not a dependent).
37 // Matches the ignore_index passed to nll_loss in the classifier.
38 static constexpr int64_t IGNORE_INDEX = -100;
39
40 void init();
41
42 std::vector<nets::embd_descr_t> get_embd_descr();
43
44 // Dicts of the embeddable (morph) features, in the SAME order as
45 // get_embd_descr() returns them. Needed to build the model's dict holder
46 // parallel to its embd_descr, so the per-embedding `dict=<i>` indices in the
47 // generated script reference the correct dict (the script numbers dicts by
48 // position in embd_descr, not by position in the full tag dict holder).
50
51 class Iterator : public BatchIterator
52 {
53 public:
54 Iterator(const CoNLLUDataSet& dataset)
55 : m_dataset(dataset),
56 m_batch_size(0),
57 m_current_bucket(0),
58 m_iter_counter(0)
59 {
60 }
61
62 virtual ~Iterator() = default;
63 virtual void set_batch_size(int64_t batch_size);
64 virtual void start_epoch();
65 virtual bool end();
66 virtual const Batch next_batch();
67
68 private:
69 const CoNLLUDataSet& m_dataset;
70 int64_t m_batch_size;
71
72 size_t m_current_bucket;
73 int64_t m_iter_counter;
74
75 friend class CoNLLUDataSet;
76 };
77
78 virtual std::shared_ptr<BatchIterator> get_iterator() const;
79
80private:
81 bool m_add_root;
82 const size_t m_batch_size;
83 const DictsHolder m_morph_tag_dh;
84 // deprel (syntactic relation label) string -> class id; shared between the
85 // train and dev datasets so ids are consistent. May be null (rel column then
86 // filled with IGNORE_INDEX).
87 std::shared_ptr<const std::map<std::string, int64_t>> m_deprel2id;
88 const CoNLLU::Annotation& m_annot;
89 std::vector<std::shared_ptr<FeatureVectorizerBase<>>> m_feat_vectorizers;
90
91 // Input features
97
99
100 std::vector<CoNLLUToTorchMatrix::feature_descr_t> m_feat_descr;
101 std::vector<CoNLLUToTorchMatrix::embeddable_feature_descr_t> m_embd_feat_descr;
102
103 // Vectorized inputs
104 std::vector<size_t> m_bucket_keys;
105 std::map<size_t, CoNLLUToTorchMatrix::vectorization_t> m_input_buckets;
106
107 // Vectorized gold values
108 std::map<size_t, std::shared_ptr<TorchMatrix<int64_t>>> m_gold_buckets;
109
110 //void load_embeddings();
111 void vectorize();
112 size_t vectorize_bucket(size_t len, const std::vector<size_t>& sents, const CoNLLUToTorchMatrix& vectorizer);
113 void vectorize_bucket_gold(const CoNLLU::BoundedWordLevelAdapter& src,
115 uint64_t timepoint);
116};
117
118} // train
119} // graph_dp
120} // deeplima
121
122#endif // DEEPLIMA_LIBS_TASKS_GRAPH_DP_CONLLU_FILE_ITERATOR_H
std::vector< nets::embd_descr_t > get_embd_descr()
virtual std::shared_ptr< BatchIterator > get_iterator() const