15using torch::indexing::Slice;
20namespace lemmatization
32 torch::optim::Optimizer& opt,
33 const torch::Device& device)
35 assert(train_input.size() == train_gold.size());
36 assert(eval_input.size() == eval_gold.size());
48 for (
size_t i = 0; i < train_input.size(); ++i)
50 const auto& input = train_input[i];
51 const auto& input_cat = train_input_cat[i];
52 const auto& gold = train_gold[i];
54 assert(input.get_max_feat() == gold.get_max_feat());
55 assert(input.get_max_feat() == input_cat[0].get_max_feat());
56 auto train_subset_stat =
train_on_subset(params, input, input_cat, gold, opt, device);
57 train_stat[
"output"].m_loss = train_subset_stat[
"output"].m_loss;
58 train_stat[
"output"].m_correct += train_subset_stat[
"output"].m_correct;
59 train_stat[
"output"].m_items += train_subset_stat[
"output"].m_items;
61 if (train_stat[
"output"].m_items > 0)
63 train_stat[
"output"].m_accuracy = float(train_stat[
"output"].m_correct) / train_stat[
"output"].m_items;
68 std::cout <<
"EPOCH " << e <<
" | lemmatization"
71 <<
" TRAIN LOSS=" << train_stat[
"output"].m_loss
72 <<
" ACC=" << train_stat[
"output"].m_accuracy
73 <<
" CORRECT=" << train_stat[
"output"].m_correct
74 <<
" TOTAL=" << train_stat[
"output"].m_items
75 << std::endl << std::flush;
90 torch::optim::Optimizer& opt,
91 const torch::Device& device)
96 const auto& input_tensor = train_input.
get_tensor().to(device);
97 const auto& gold_tensor = train_gold.
get_tensor().to(device);
98 int64_t n_samples = input_tensor.sizes()[1];
100 for (int64_t i = 0; i < n_samples; i += params.
m_batch_size)
102 int64_t end_sample = cast_to_signed<decltype(n_samples)>(params.
m_batch_size) > n_samples ? n_samples : i + params.
m_batch_size;
103 assert(i < end_sample);
104 const auto batch_input = input_tensor.index({ Slice(), Slice(i, end_sample) }).
to(device);
105 const auto batch_gold = gold_tensor.index({ Slice(), Slice(i, end_sample) }).
to(device);
107 vector<TorchMatrix<int64_t>::tensor_t> batch_input_cat(train_input_cat.size());
108 for (
size_t feat_idx = 0; feat_idx < train_input_cat.size(); ++feat_idx)
110 const auto& input_cat_tensor = train_input_cat[feat_idx].get_tensor().to(device);
111 batch_input_cat[feat_idx] = input_cat_tensor.index({ Slice(), Slice(i, end_sample) }).
to(device);
114 train_batch({
"output" }, batch_input, batch_input_cat, batch_gold, opt, stat, device);
117 if (stat[
"output"].m_items > 0)
119 stat[
"output"].m_accuracy = float(stat[
"output"].m_correct) / stat[
"output"].m_items;
127 const torch::Device& device)
129 for (
size_t i = 0; i < input.size(); ++i)
135 BiRnnSeq2SeqImpl::evaluate({
"output" }, input_bucket, gold_bucket, stat, device);
141 BiRnnSeq2SeqImpl::load(archive);
143 archive.read(
"morph_model", val);
144 string serialized_morph_model = *(val.toString().get());
147 archive.read(
"fixed_upos", val);
153 BiRnnSeq2SeqImpl::save(archive);
155 archive.write(
"morph_model", serialized_morph_model);
uint64_t get_max_feat() const
const torch::Tensor & get_tensor() const
morph_model::morph_model_t m_morph_model
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 train(const train_params_lemmatization_t ¶ms, const std::vector< TorchMatrix< int64_t > > &train_input, const std::vector< std::vector< TorchMatrix< int64_t > > > &train_input_cat, const std::vector< TorchMatrix< int64_t > > &train_gold, const std::vector< TorchMatrix< int64_t > > &eval_input, const std::vector< std::vector< TorchMatrix< int64_t > > > &eval_input_cat, const std::vector< TorchMatrix< int64_t > > &eval_gold, torch::optim::Optimizer &opt, const torch::Device &device=torch::Device(torch::kCPU))
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.
std::string to_string() const
void train_batch(const std::vector< std::string > &output_names, const torch::Tensor &input, const std::vector< torch::Tensor > &input_cat, const torch::Tensor &target, torch::optim::Optimizer &opt, epoch_stat_t &stat, const torch::Device &device)
virtual void set_tags(const std::map< std::string, std::string > &tags)
virtual void to(torch::Device device, bool non_blocking=false) override
std::map< std::string, task_stat_t > epoch_stat_t
std::string m_output_model_name
std::map< std::string, std::string > m_tags