20 const torch::Tensor& input,
21 const vector<torch::Tensor>& input_cat,
22 const torch::Tensor& target,
23 torch::optim::Optimizer& opt,
25 const torch::Device& device)
28 int64_t l = target.sizes()[0];
29 torch::Tensor gold_input = torch::empty({1, target.sizes()[1]},
30 torch::TensorOptions().dtype(torch::kInt64)).fill_(2).to(device);
31 gold_input = torch::cat({gold_input, target.index({Slice(0, l-1), Slice()})}, 0);
35 std::map<string, torch::Tensor> input_map
36 = { {
"enc_chars", input }, {
"dec_chars", gold_input } };
42 input_map[
m_cat_embd_descr[feat_idx].m_name] = torch::squeeze(input_cat[feat_idx], 0).to(device);
44 auto output_map =
forward(input_map, output_names.begin(), output_names.end());
46 for (
size_t i = 0; i < output_names.size(); ++i)
48 const string& task_name = output_names[i];
51 auto output = output_map[task_name].to(device);
56 auto o = output.reshape({ -1, output.size(2) }).
to(device);
58 auto this_task_target = target.reshape({ -1 }).
to(device);
61 torch::Tensor loss_tensor = torch::nn::functional::nll_loss(o, this_task_target);
62 loss_tensor.to(device);
64 double loss_value = loss_tensor.mean().item<
double>();
65 task_stat.
m_loss += loss_value;
67 auto prediction = o.argmax(1);
68 int64_t correct_predictions = prediction.eq(this_task_target).sum().item<int64_t>();
69 task_stat.
m_correct += correct_predictions;
70 task_stat.
m_items += this_task_target.size(0);
72 loss_tensor.backward({},
true);