14using torch::indexing::Slice;
29 shared_ptr<
const std::map<std::string, int64_t>> deprel2id)
30 : m_add_root(add_root),
31 m_batch_size(batch_size),
32 m_morph_tag_dh(morph_tag_dh),
33 m_deprel2id(deprel2id),
35 m_feat_vectorizers({ p_embd }),
36 m_feat_extractor(feat_extractor)
38 for (
auto p_embd : m_feat_vectorizers)
43 for (
size_t i = 0; i < m_morph_tag_dh.size(); ++i)
45 if (morph_tag_dh[i]->size() > 0)
47 m_embd_feat_descr.push_back({
49 m_feat_extractor.
feats()[i],
65 assert(m_feat_descr.size() > 0 || m_embd_feat_descr.size() > 0);
73 for (
const auto& fd : m_embd_feat_descr)
75 dicts.push_back(fd.m_dict);
80void CoNLLUDataSet::vectorize()
82 assert(m_feat_descr.size() > 0 || m_embd_feat_descr.size() > 0);
83 CoNLLUToTorchMatrix vectorizer(m_feat_descr, m_embd_feat_descr, m_feat_extractor);
86 vector<vector<size_t>> sent_by_len;
87 sent_by_len.resize(256);
99 if (len > sent_by_len.size() - 1)
101 sent_by_len.resize(len + 1);
103 sent_by_len[len].push_back(i);
106 m_bucket_keys.reserve(sent_by_len.size());
107 for (
size_t i = 0; i < sent_by_len.size(); i++)
109 const vector<size_t>& sents = sent_by_len[i];
115 size_t k = vectorize_bucket(i, sents, vectorizer);
116 m_bucket_keys.push_back(k);
120size_t CoNLLUDataSet::vectorize_bucket(
size_t len,
const vector<size_t>& sents,
const CoNLLUToTorchMatrix& vectorizer)
122 uint64_t timepoint = 0;
123 uint64_t timepoints_per_sentence = m_add_root ? len + 1 : len;
124 uint64_t total_timepoints = timepoints_per_sentence * sents.size();
125 m_input_buckets[timepoints_per_sentence] = vectorizer.init_dst(total_timepoints);
127 m_gold_buckets[timepoints_per_sentence] = std::make_shared<TorchMatrix<int64_t>>(total_timepoints, 2);
129 CoNLLU::Annotation m_root_annot;
130 stringstream root_line;
131 root_line <<
"1\t<ROOT>\t<ROOT>\tROOT\t_\t_\t0\troot\t0:root\t_\n";
132 m_root_annot.load(root_line);
133 CoNLLU::BoundedWordLevelAdapter root_generator(&m_root_annot, 0, 1);
135 for (
size_t i = 0; i < sents.size(); i++)
137 const CoNLLU::Sentence& sent = m_annot.
get_sentence(sents[i]);
139 CoNLLU::BoundedWordLevelAdapter adapter(&m_annot, sent.get_first_word_idx(), len);
143 vectorizer.process(root_generator, m_input_buckets[timepoints_per_sentence], timepoint);
146 vectorizer.process(adapter, m_input_buckets[timepoints_per_sentence], m_add_root ? timepoint+1 : timepoint);
147 vectorize_bucket_gold(adapter, *(m_gold_buckets[timepoints_per_sentence].get()), timepoint);
149 timepoint += timepoints_per_sentence;
152 return timepoints_per_sentence;
155void CoNLLUDataSet::vectorize_bucket_gold(
const CoNLLU::BoundedWordLevelAdapter& src,
156 TorchMatrix<int64_t>& dst,
160 uint64_t current_timepoint = timepoint;
161 while (src.end() != it)
163 while(!(*it).is_word() && src.end() != it)
173 const CoNLLU::idx_t& head = (*it).head();
174 const std::string& gold_rel = (*it).deprel();
175 assert(head.is_real_word());
177 const CoNLLU::idx_t idx = (*it).idx();
178 if (m_add_root && 1 == idx._first)
189 dst.set(current_timepoint, 0, head._first);
195 auto rel_it = m_deprel2id->find(gold_rel);
196 rel_id = (rel_it != m_deprel2id->end()) ? rel_it->second :
IGNORE_INDEX;
198 dst.set(current_timepoint, 1, rel_id);
201 if (current_timepoint == std::numeric_limits<uint64_t>::max())
203 throw std::overflow_error(
"Too much words in the dataset.");
212 return std::make_shared<Iterator>(*
this);
217 m_batch_size = batch_size;
222 if (0 == m_batch_size)
224 throw std::runtime_error(
"CoNLLUDataSet::Iterator::start_epoch batch size cannot be 0.");
226 m_current_bucket = 0;
232 return m_current_bucket >= m_dataset.m_bucket_keys.size();
237 size_t k = m_dataset.m_bucket_keys[m_current_bucket];
239 auto it_input = m_dataset.m_input_buckets.find(k);
240 assert(m_dataset.m_input_buckets.end() != it_input);
243 auto it_gold = m_dataset.m_gold_buckets.find(k);
244 assert(m_dataset.m_gold_buckets.end() != it_gold);
249 int64_t seq_len = it_input->first;
254 int64_t batch_size = m_batch_size;
255 if (-1 == batch_size)
257 assert(0 == input.first->get_tensor().size(0) % seq_len);
258 batch_size = input.first->get_tensor().size(0) / seq_len;
263 std::random_device r;
264 std::default_random_engine e1(r());
265 int64_t max_start_offset = ( gold_bucket.
get_tensor().size(0) / seq_len ) - batch_size;
266 if (max_start_offset < 0)
268 max_start_offset = 0;
270 std::uniform_int_distribution<int> uniform_dist(0, max_start_offset);
271 int64_t batch_start_offset = uniform_dist(e1);
273 if (0 == batch_start_offset &&
274 seq_len * (batch_start_offset + batch_size) > gold_bucket.
get_tensor().size(0))
284 const torch::Tensor trainable
285 = input.first->get_tensor().index({ Slice(seq_len * batch_start_offset,
286 seq_len * (batch_start_offset + batch_size)),
287 Slice() }).reshape({ batch_size, seq_len, -1 }).transpose(0, 1);
288 const torch::Tensor frozen
289 = input.second->get_tensor().index({ Slice(seq_len * batch_start_offset,
290 seq_len * (batch_start_offset + batch_size)),
291 Slice() }).reshape({ batch_size, seq_len, -1 }).transpose(0, 1);
296 const torch::Tensor gold
297 = gold_bucket.
get_tensor().index({ Slice(seq_len * batch_start_offset,
298 seq_len * (batch_start_offset + batch_size)),
299 Slice() }).reshape({ batch_size, seq_len, -1 });
303 if (m_batch_size >= 0)
307 if (m_iter_counter * m_batch_size > gold_bucket.
get_tensor().size(0) / seq_len)
size_t get_num_sentences() const
const Sentence & get_sentence(size_t idx) const
size_t calc_num_of_words(const Annotation &annot) const
iterator_struct const_iterator
const torch::Tensor & get_tensor() const
virtual void start_epoch()
virtual const Batch next_batch()
virtual void set_batch_size(int64_t batch_size)
std::vector< nets::embd_descr_t > get_embd_descr()
virtual std::shared_ptr< BatchIterator > get_iterator() const
static constexpr int64_t IGNORE_INDEX
CoNLLUDataSet(const CoNLLU::Annotation &annot, size_t batch_size, const ConlluFeatExtractor< CoNLLU::WordLevelAdapter::token_t > &feat_extractor, const DictsHolder &morph_tag_dh, std::shared_ptr< FeatureVectorizerBase<> > p_embd, bool add_root=false, std::shared_ptr< const std::map< std::string, int64_t > > deprel2id=nullptr)
DictsHolder get_embd_feature_dicts() const
std::pair< std::shared_ptr< MatrixInt >, std::shared_ptr< MatrixFloat > > vectorization_t
const std::vector< deeplima::nets::embd_descr_t > get_embd_descr() const