36 auto out = std::make_shared<M>(len, 1);
40 for (
size_t i = 0; i < tokens.size(); i++)
42 const auto& t = tokens[i];
45 out->set(p, 0, segm_tag_t::X);
51 throw std::runtime_error(
"vectorize_gold length should be 0");
57 const bool is_mwt = mwt && t.is_multiword();
63 out->set(p, 0, is_mwt ? segm_tag_t::S_EOS_MWT : segm_tag_t::S_EOS);
67 out->set(p, 0, is_mwt ? segm_tag_t::S_MWT : segm_tag_t::S);
73 out->set(p, 0, segm_tag_t::B);
75 while (p < t.m_pos + t.m_len - 1)
77 out->set(p, 0, segm_tag_t::I);
82 out->set(p, 0, is_mwt ? segm_tag_t::E_EOS_MWT : segm_tag_t::E_EOS);
86 out->set(p, 0, is_mwt ? segm_tag_t::E_MWT : segm_tag_t::E);
103 std::vector<ngram_descr_t> ngram_descr = {
113 uint64_t train_char_counter = 0;
116 auto dicts = dict_builder.process(train_doc.
get_original_text(), 100, train_char_counter);
118 if (train_char_counter >= std::numeric_limits<int64_t>::max())
120 throw std::overflow_error(
"Too much characters in training set.");
124 vectorizer.set_dicts(dicts);
126 (int64_t)train_char_counter);
128 const auto& dev_doc = tb.
get_doc(
"dev");
129 auto dev_input = vectorizer.process(dev_doc.get_original_text(), dev_doc.get_text().size() + 1);
131 auto train_gold = vectorize_gold<TorchMatrix<int64_t>>(tb.
get_annot(
"train"), (int64_t)train_char_counter,
134 auto dev_gold = vectorize_gold<TorchMatrix<int64_t>>(tb.
get_annot(
"dev"), dev_doc.get_text().size() + 1,
137 std::vector<embd_descr_t> embd_descr = { {
"char1gram", 2 }, {
"char2gram", 3 }, {
"char3gram", 4 },
138 {
"class1gram", 2 }, {
"class2gram", 2 }, {
"class3gram", 2 },
139 {
"scriptchange", 1 } };
148 BiRnnClassifierForSegmentation model(std::move(dicts),
156 torch::optim::Adam optimizer(model->parameters(),
159 .betas({params.m_beta_one, params.m_beta_two}));
161 std::string dev =
"cpu";
164 std::ostringstream oss;
165 oss <<
"cuda:" << gpuid;
168 torch::Device device(dev);
170 train_input->to(device);
171 train_gold->to(device);
173 dev_input->to(device);
174 dev_gold->to(device);
181 *(train_input.get()), *(train_gold.get()),
182 *(dev_input.get()), *(dev_gold.get()),
185 catch (
const c10::Error& e)
187 std::cerr <<
"Exception in model training: " << e.what() << std::endl;