6#ifndef DEEPLIMA_LIBS_TASKS_LEMMATIZATION_MODEL_SEQ2SEQ_LEMMATIZER_H
7#define DEEPLIMA_LIBS_TASKS_LEMMATIZATION_MODEL_SEQ2SEQ_LEMMATIZER_H
17namespace lemmatization
33 const std::vector<nets::embd_descr_t>& encoder_embd_descr,
34 const std::vector<nets::rnn_descr_t>& encoder_rnn_descr,
35 const std::vector<nets::embd_descr_t>& decoder_embd_descr,
36 const std::vector<nets::rnn_descr_t>& decoder_rnn_descr,
37 const std::vector<nets::embd_descr_t>& cat_embd_descr,
38 const std::string& fixed_upos)
40 encoder_embd_descr, encoder_rnn_descr,
41 decoder_embd_descr, decoder_rnn_descr,
55 torch::optim::Optimizer& opt,
56 const torch::Device& device = torch::Device(torch::kCPU));
61 const torch::Device& device = torch::Device(torch::kCPU));
68 torch::optim::Optimizer& opt,
69 const torch::Device& device = torch::Device(torch::kCPU));
71 virtual void load(torch::serialize::InputArchive& archive);
72 virtual void save(torch::serialize::OutputArchive& archive)
const;
74 void load(
const std::string& fn)
76 torch::load(*
this, fn);
96 torch::serialize::OutputArchive& archive,
104 torch::serialize::InputArchive& archive,
107 module.load(archive);
morph_model::morph_model_t m_morph_model
Seq2SeqLemmatizerImpl(DictsHolder &&dicts, const morph_model::morph_model_t &lang_morph_model, const std::vector< nets::embd_descr_t > &encoder_embd_descr, const std::vector< nets::rnn_descr_t > &encoder_rnn_descr, const std::vector< nets::embd_descr_t > &decoder_embd_descr, const std::vector< nets::rnn_descr_t > &decoder_rnn_descr, const std::vector< nets::embd_descr_t > &cat_embd_descr, const std::string &fixed_upos)
virtual void save(torch::serialize::OutputArchive &archive) const
void evaluate(const std::vector< TorchMatrix< int64_t > > &input, const std::vector< TorchMatrix< int64_t > > &gold, nets::epoch_stat_t &stat, const torch::Device &device=torch::Device(torch::kCPU))
void load(const std::string &fn)
const std::string & get_fixed_upos() const
const morph_model::morph_model_t & get_morph_model() const
virtual void load(torch::serialize::InputArchive &archive)
nets::epoch_stat_t train_on_subset(const train_params_lemmatization_t ¶ms, const TorchMatrix< int64_t > &train_input, const std::vector< TorchMatrix< int64_t > > &train_input_cat, const TorchMatrix< int64_t > &train_gold, torch::optim::Optimizer &opt, const torch::Device &device=torch::Device(torch::kCPU))
Helper class for morphology data (upos, features) binarization.
torch::serialize::InputArchive & operator>>(torch::serialize::InputArchive &archive, Seq2SeqLemmatizerImpl &module)
TORCH_MODULE(Seq2SeqLemmatizer)
torch::serialize::OutputArchive & operator<<(torch::serialize::OutputArchive &archive, const Seq2SeqLemmatizerImpl &module)
std::map< std::string, task_stat_t > epoch_stat_t