30 const torch::Tensor& input,
31 const torch::Device& device)
33 map<string, torch::Tensor> current_inputs;
36 auto output_map =
forward(current_inputs, output_names.begin(), output_names.end());
38 torch::Tensor prediction = torch::zeros({ input.size(0),
long(output_names.size()) },
39 torch::TensorOptions().dtype(torch::kInt64));
40 for (
size_t i = 0; i < output_names.size(); ++i)
42 const string& task_name = output_names[i];
43 torch::Tensor& output = output_map[task_name];
44 torch::Tensor o = output.reshape({-1, output.size(2)});
45 prediction.index({ Slice(), Slice(i, i+1) }) = o.argmax(1);
52 const torch::Tensor& inputs,
58 const std::vector<std::string>& outputs_names,
59 const torch::Device& device)
61 assert(output->size() == outputs_names.size());
63 int64_t start_shift = output_begin - input_begin;
64 assert(start_shift >= 0);
65 int64_t end_shift = input_end - output_end;
66 assert(end_shift >= 0);
68 const torch::Tensor inputs_slice = inputs.index({ Slice(input_begin, input_end), Slice() });
70 map<string, torch::Tensor> current_inputs;
73 auto output_map =
forward(current_inputs, outputs_names.begin(), outputs_names.end());
74 for (
size_t i = 0; i < outputs_names.size(); i++)
76 torch::Tensor& one_task_output = output_map[outputs_names[i]];
77 torch::Tensor o = one_task_output.reshape({-1, one_task_output.size(2)});
78 const torch::Tensor output_tensor = o.argmax(1);
79 torch::TensorAccessor<int64_t, 1> accessor = output_tensor.accessor<int64_t, 1>();
80 for (int64_t p = start_shift; p < output_tensor.size(0) - end_shift; p++)
82 (*output)[i][input_begin + p] = accessor[p];
88 map<string, torch::Tensor>& dst,
89 const torch::Device& device)
96 throw std::runtime_error(
"Duplicates in embd list.");
98 if (src.sizes().size() == 2)
103 = src.index({Slice(), Slice(c, c+1) }).reshape({ (int64_t)src.size(0), 1 }).
to(device);
112 else if (src.sizes().size() == 3)
118 = src.index({Slice(), Slice(), Slice(c, c+1) }).reshape({ (int64_t)src.size(0), -1 }).
to(device);
135 const torch::Device& device)
137 map<string, torch::Tensor> current_inputs;
140 auto target = gold.
get_tensor().reshape({ -1, long(output_names.size()) }).
to(device);
142 evaluate(output_names, current_inputs, target, stat, device);
146 const map<string, torch::Tensor>& input,
147 const torch::Tensor& target,
149 const torch::Device& )
152 auto output_map =
forward(input, output_names.begin(), output_names.end());
154 for (
size_t i = 0; i < output_names.size(); ++i)
156 const string& task_name = output_names[i];
159 torch::Tensor& output = output_map[task_name];
160 torch::Tensor o = output.reshape({-1, output.size(2)});
161 auto this_task_target = target.index({ Slice(), Slice(i, i+1) }).reshape({ -1 });
165 static constexpr int64_t kIgnoreIndex = -100;
166 torch::Tensor loss_tensor = torch::nn::functional::nll_loss(
168 torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnoreIndex));
169 task_stat.
m_loss = loss_tensor.sum().item<
double>();
170 auto valid = this_task_target.ne(kIgnoreIndex);
171 auto prediction = o.argmax(1);
172 task_stat.
m_correct = prediction.eq(this_task_target).logical_and(valid).sum().item<int64_t>();
173 task_stat.
m_items = valid.sum().item<int64_t>();
183 auto gold_positive = prediction.eq(1);
184 auto gold_negative = prediction.eq(0);
185 auto pred_positive = this_task_target.eq(1);
186 auto pred_negative = this_task_target.eq(0);
187 auto tp = torch::logical_and(gold_positive, pred_positive).count_nonzero().item<int64_t>();
190 task_stat.
m_precision = tp / (prediction.eq(1).sum().item<int64_t>() + 0.00001);
191 task_stat.
m_recall = tp / (this_task_target.eq(1).sum().item<int64_t>() + 0.00001);
194 else if (o.size(-1) > 2)
202 const int64_t num_classes = o.size(-1);
204 double sum_p = 0, sum_r = 0, sum_f1 = 0;
205 for (int64_t c = 0; c < num_classes; ++c)
209 auto pred_c = prediction.eq(c).logical_and(valid);
210 auto gold_c = this_task_target.eq(c);
211 auto tp_c = torch::logical_and(pred_c, gold_c).count_nonzero().item<int64_t>();
212 double prec = tp_c / (pred_c.sum().item<int64_t>() + 0.00001);
213 double rec = tp_c / (gold_c.sum().item<int64_t>() + 0.00001);
216 sum_f1 += (2 * prec * rec) / (prec + rec + 0.00001);
219 task_stat.
m_recall = sum_r / num_classes;
220 task_stat.
m_f1 = sum_f1 / num_classes;
228 const vector<string>& output_names,
233 torch::optim::Optimizer& opt,
234 const std::string& model_name,
235 const torch::Device& device)
237 int64_t num_batches = train_input.
size() / seq_len;
238 int64_t num_features = (int64_t)train_input.
get_max_feat();
239 int64_t seq_len_i64 = (int64_t)seq_len;
241 auto aligned_input = train_input.
get_tensor().index({ Slice(0, num_batches * seq_len_i64), Slice() });
242 auto aligned_gold = train_gold.
get_tensor().index({ Slice(0, num_batches * seq_len_i64), Slice() });
244 = aligned_input.reshape({ num_batches, seq_len_i64, num_features }).transpose(0, 1);
246 = aligned_gold.reshape({ num_batches, seq_len_i64 }).transpose(0, 1);
248 double best_eval_accuracy = 0;
249 double best_eval_loss = std::numeric_limits<double>::max();
250 size_t count_below_best = 0;
256 for (
auto& group : opt.param_groups())
258 if (group.has_options())
260 lr_copy =
static_cast<torch::optim::AdamOptions&
>(group.options()).lr();
263 for (
size_t e = 0; e < epochs; e++)
268 train_epoch(batch_size, seq_len, output_names, input_batches, gold_batches, opt, train_stat, device);
270 evaluate(output_names, eval_input, eval_gold, eval_stat, device);
272 for (
const string& task_name : output_names)
276 std::cout <<
"EPOCH " << e <<
" | " << task_name <<
" | LR=" << lr_copy;
277 std::cout <<
" | TRAIN LOSS=" << train.
m_loss <<
" ACC=" << train.
m_accuracy
280 << std::endl << std::flush;
283 task_stat_t& main_task_eval = eval_stat[output_names[0]];
285 best_eval_loss = min(best_eval_loss, main_task_eval.
m_loss);
286 if (main_task_eval.
m_accuracy > best_eval_accuracy)
288 best_eval_accuracy = main_task_eval.
m_accuracy;
289 if (model_name.size() > 0)
291 torch::save(*
this, model_name +
".pt");
293 count_below_best = 0;
300 else if (main_task_eval.
m_accuracy < best_eval_accuracy)
302 for (
auto &group : opt.param_groups())
304 if (group.has_options())
306 auto &options =
static_cast<torch::optim::AdamOptions &
>(group.options());
307 options.lr(options.lr() * (0.9));
308 lr_copy = options.lr();
311 if (lr_copy < 0.0000001)
316 if (main_task_eval.
m_loss > best_eval_loss && count_below_best > 10)
320 if (count_below_best > 50)
330 const vector<string>& output_names,
331 const torch::Tensor& input_batches,
332 const torch::Tensor& gold_batches,
333 torch::optim::Optimizer& opt,
335 const torch::Device& device)
337 for (int64_t b = 0; b < input_batches.size(1); b += batch_size)
339 auto current_batch_size
340 = ((b + (int64_t)batch_size) > input_batches.size(1)) ? input_batches.size(1) - b : batch_size;
342 train_batch(current_batch_size, seq_len, output_names,
343 input_batches.index({Slice(), Slice(b, b + current_batch_size), Slice()}),
344 gold_batches.index({Slice(), Slice(b, b + current_batch_size) }),
351 const vector<string>& output_names,
352 const torch::Tensor& input,
353 const torch::Tensor& gold,
354 torch::optim::Optimizer& opt,
356 const torch::Device& device)
358 map<string, torch::Tensor> current_batch_inputs;
361 auto target = gold.reshape({-1, 1}).
to(device);
363 train_batch(batch_size, seq_len, output_names, current_batch_inputs, target, opt, stat, device);
368 const vector<string>& output_names,
369 const map<string, torch::Tensor>& input,
370 const torch::Tensor& target,
371 torch::optim::Optimizer& opt,
373 const torch::Device& )
377 auto output_map =
forward(input, output_names.begin(), output_names.end());
383 for (
size_t i = 0; i < output_names.size(); ++i)
385 const string& task_name = output_names[i];
388 auto output = output_map[task_name];
390 auto o = output.reshape({ -1, output.size(2) });
391 auto this_task_target = target.index({ Slice(), Slice(i, i+1) }).reshape({ -1 });
397 static constexpr int64_t kIgnoreIndex = -100;
398 torch::Tensor loss_tensor = torch::nn::functional::nll_loss(
400 torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnoreIndex));
401 double loss_value = loss_tensor.sum().item<
double>();
402 task_stat.
m_loss += loss_value;
404 auto valid = this_task_target.ne(kIgnoreIndex);
405 auto prediction = o.argmax(1);
406 int64_t correct_predictions = prediction.eq(this_task_target).logical_and(valid).sum().item<int64_t>();
407 task_stat.
m_correct += correct_predictions;
408 task_stat.
m_items += valid.sum().item<int64_t>();
410 loss_tensor.backward({},
true);
void train_epoch(size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, const torch::Tensor &input_batches, const torch::Tensor &gold_batches, torch::optim::Optimizer &opt, epoch_stat_t &stat, const torch::Device &device)
void train_batch(size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, const torch::Tensor &input, const torch::Tensor &gold, torch::optim::Optimizer &opt, epoch_stat_t &stat, const torch::Device &device)
void train(size_t epochs, size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, const TorchMatrix< int64_t > &train_input, const TorchMatrix< int64_t > &train_gold, const TorchMatrix< int64_t > &eval_input, const TorchMatrix< int64_t > &eval_gold, torch::optim::Optimizer &opt, const std::string &model_name="", const torch::Device &device=torch::Device(torch::kCPU))