167 c10::List<std::string> list_of_class_names;
178 c10::List<std::string> list_of_class_names;
188 c10::List<c10::List<std::string>> list_of_classes;
191 c10::List<std::string> current_class;
192 shared_ptr<StringDict> d = dynamic_pointer_cast<StringDict, DictBase>(
m_input_classes[i]);
193 for (
size_t j = 0; j < d->size(); j++)
195 current_class.push_back(d->get_value(j));
197 list_of_classes.push_back(current_class);
202 c10::List<std::string> list_of_embd_fn;
203 list_of_embd_fn.reserve(1);
208 c10::List<std::string> list_of_rel_classes;
218 const std::vector<std::string>& output_names,
221 torch::optim::Optimizer& opt,
222 double& best_eval_accuracy,
223 const torch::Device& device)
228 double best_eval_loss = std::numeric_limits<double>::max();
229 size_t count_below_best = 0;
231 unsigned int epoch = 0;
234 for (
auto &group : opt.param_groups())
236 if (group.has_options())
238 auto &options =
static_cast<torch::optim::AdamOptions &
>(group.options());
240 lr_copy = options.get_lr();
246 chrono::steady_clock::time_point begin = chrono::steady_clock::now();
256 for (
const std::string& tn : output_names)
258 if (train_stat[tn].m_items > 0)
260 train_stat[tn].m_accuracy = double(train_stat[tn].m_correct) / train_stat[tn].m_items;
261 train_stat[tn].m_loss /= train_stat[tn].m_items;
265 chrono::steady_clock::time_point train_end = chrono::steady_clock::now();
269 for (
const std::string& tn : output_names)
271 if (eval_stat[tn].m_items > 0)
273 eval_stat[tn].m_accuracy = double(eval_stat[tn].m_correct) / eval_stat[tn].m_items;
274 eval_stat[tn].m_loss /= eval_stat[tn].m_items;
278 chrono::steady_clock::time_point eval_end = chrono::steady_clock::now();
280 auto train_duration = std::chrono::duration_cast<std::chrono::milliseconds>(train_end - begin).count();
281 auto eval_duration = std::chrono::duration_cast<std::chrono::milliseconds>(eval_end - train_end).count();
283 for (
const string& task_name : output_names)
285 const auto&
train = train_stat[task_name];
286 const auto& eval = eval_stat[task_name];
287 std::cout <<
"EPOCH " << epoch <<
" | " << task_name <<
" | LR=" << lr_copy
288 <<
" | TRAIN LOSS=" <<
train.m_loss <<
" ACC=" <<
train.m_accuracy
289 <<
" | EVAL LOSS=" << eval.m_loss <<
" ACC=" << eval.m_accuracy;
290 if (2 == eval.m_num_classes)
292 std::cout <<
" P=" << eval.m_precision <<
" R=" << eval.m_recall<<
" F1=" << eval.m_f1;
296 std::cout <<
" CORRECT=" << eval.m_correct;
298 std::cout << std::endl << std::flush;
300 cout <<
"TIME: train=" << train_duration <<
"[ms] eval=" << eval_duration <<
"[ms]" << endl;
302 task_stat_t& main_task_eval = eval_stat[output_names[0]];
304 best_eval_loss = min(best_eval_loss, main_task_eval.
m_loss);
305 if (main_task_eval.
m_accuracy > best_eval_accuracy)
307 best_eval_accuracy = main_task_eval.
m_accuracy;
312 count_below_best = 0;
314 else if (main_task_eval.
m_accuracy < best_eval_accuracy)
316 for (
auto &group : opt.param_groups())
318 if (group.has_options())
320 auto &options =
static_cast<torch::optim::AdamOptions &
>(group.options());
321 options.lr(options.lr() * (0.9));
322 lr_copy = options.lr();
325 if (lr_copy < 0.0000001)
372 const vector<string>& output_names,
373 const torch::Tensor& trainable_input,
374 const torch::Tensor& nontrainable_input,
375 const torch::Tensor& gold,
376 torch::optim::Optimizer& opt,
378 const torch::Device& device)
380 using torch::indexing::Slice;
381 static constexpr int64_t kIgnore = -100;
383 map<string, torch::Tensor> current_batch_inputs;
384 split_input(trainable_input, current_batch_inputs, device);
385 current_batch_inputs[
"raw"] = nontrainable_input.to(device);
387 const int64_t B = gold.size(0);
388 const int64_t D = gold.size(1);
389 auto target = gold.reshape({-1, gold.size(2)}).
to(device);
392 && (std::find(output_names.begin(), output_names.end(), std::string(
"rel")) != output_names.end());
394 std::vector<std::string> requested = {
"arc" };
395 if (train_rel) { requested.push_back(
"rel_logits"); }
398 auto out =
forward(current_batch_inputs, requested.begin(), requested.end());
402 torch::Tensor o = out[
"arc"].reshape({ -1, out[
"arc"].size(2) });
403 torch::Tensor tgt = target.index({ Slice(), Slice(0, 1) }).reshape({ -1 });
404 torch::Tensor loss = torch::nn::functional::nll_loss(
405 o, tgt, torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnore));
406 torch::Tensor valid = tgt.ne(kIgnore);
407 torch::Tensor pred = o.argmax(1);
409 s.
m_loss += loss.sum().item<
double>();
410 s.
m_correct += pred.eq(tgt).logical_and(valid).sum().item<int64_t>();
411 s.
m_items += valid.sum().item<int64_t>();
412 loss.backward({}, train_rel);
418 torch::Tensor rel_logits = out[
"rel_logits"];
419 const int64_t L = rel_logits.size(3);
420 torch::Tensor gold_head = target.index({ Slice(), 0 }).reshape({ B, D });
421 torch::Tensor gh = gold_head.clamp_min(0);
422 torch::Tensor idx = gh.unsqueeze(-1).unsqueeze(-1).expand({ B, D, 1, L });
423 torch::Tensor gathered = rel_logits.gather(2, idx).squeeze(2);
424 torch::Tensor rel_log = torch::log_softmax(gathered, 2).reshape({ -1, L });
425 torch::Tensor tgt = target.index({ Slice(), Slice(1, 2) }).reshape({ -1 });
426 torch::Tensor loss = torch::nn::functional::nll_loss(
427 rel_log, tgt, torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnore));
428 torch::Tensor valid = tgt.ne(kIgnore);
429 torch::Tensor pred = rel_log.argmax(1);
431 s.
m_loss += loss.sum().item<
double>();
432 s.
m_correct += pred.eq(tgt).logical_and(valid).sum().item<int64_t>();
433 s.
m_items += valid.sum().item<int64_t>();
441 shared_ptr<BatchIterator> dataset_iterator,
443 const torch::Device& device)
446 dataset_iterator->set_batch_size(-1);
448 dataset_iterator->start_epoch();
450 while (!dataset_iterator->end())
466 for (
const std::string& task_name : output_names)
468 stat[task_name].m_correct += t[task_name].m_correct;
469 stat[task_name].m_items += t[task_name].m_items;
470 stat[task_name].m_loss += t[task_name].m_loss;
476 const torch::Tensor& trainable_input,
477 const torch::Tensor& nontrainable_input,
478 const torch::Tensor& gold,
480 const torch::Device& device)
482 using torch::indexing::Slice;
483 static constexpr int64_t kIgnore = -100;
485 map<string, torch::Tensor> current_inputs;
486 split_input(trainable_input, current_inputs, device);
487 current_inputs[
"raw"] = nontrainable_input.to(device);
489 const int64_t B = gold.size(0);
490 const int64_t D = gold.size(1);
491 auto target = gold.reshape({-1, gold.size(-1)}).
to(device);
494 && (std::find(output_names.begin(), output_names.end(), std::string(
"rel")) != output_names.end());
496 std::vector<std::string> requested = {
"arc" };
497 if (eval_rel) { requested.push_back(
"rel_logits"); }
499 torch::NoGradGuard no_grad;
500 auto out =
forward(current_inputs, requested.begin(), requested.end());
504 torch::Tensor o = out[
"arc"].reshape({ -1, out[
"arc"].size(2) });
505 torch::Tensor tgt = target.index({ Slice(), Slice(0, 1) }).reshape({ -1 });
506 torch::Tensor loss = torch::nn::functional::nll_loss(
507 o, tgt, torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnore));
508 torch::Tensor valid = tgt.ne(kIgnore);
509 torch::Tensor pred = o.argmax(1);
511 s.
m_loss = loss.sum().item<
double>();
512 s.
m_correct = pred.eq(tgt).logical_and(valid).sum().item<int64_t>();
513 s.
m_items = valid.sum().item<int64_t>();
519 torch::Tensor rel_logits = out[
"rel_logits"];
520 const int64_t L = rel_logits.size(3);
521 torch::Tensor gold_head = target.index({ Slice(), 0 }).reshape({ B, D });
522 torch::Tensor gh = gold_head.clamp_min(0);
523 torch::Tensor idx = gh.unsqueeze(-1).unsqueeze(-1).expand({ B, D, 1, L });
524 torch::Tensor gathered = rel_logits.gather(2, idx).squeeze(2);
525 torch::Tensor rel_log = torch::log_softmax(gathered, 2).reshape({ -1, L });
526 torch::Tensor tgt = target.index({ Slice(), Slice(1, 2) }).reshape({ -1 });
527 torch::Tensor loss = torch::nn::functional::nll_loss(
528 rel_log, tgt, torch::nn::functional::NLLLossFuncOptions().ignore_index(kIgnore));
529 torch::Tensor valid = tgt.ne(kIgnore);
530 torch::Tensor pred = rel_log.argmax(1);
532 s.
m_loss = loss.sum().item<
double>();
533 s.
m_correct = pred.eq(tgt).logical_and(valid).sum().item<int64_t>();
534 s.
m_items = valid.sum().item<int64_t>();