17 using torch::indexing::Slice;
19 const int64_t batch_size = input.size(1);
22 torch::Tensor input_t = input.transpose(0, 1);
31 torch::Tensor roots = torch::tile(
root, {batch_size, 1, 1});
32 h_head = torch::cat({roots, h_head}, 1);
37 auto append_ones = [](
const torch::Tensor& t) {
38 torch::Tensor ones = torch::ones({t.size(0), t.size(1), 1}, t.options());
39 return torch::cat({t, ones}, 2);
41 torch::Tensor aug_dep = append_ones(h_dep);
42 torch::Tensor aug_head = append_ones(h_head);
46 torch::Tensor logits = torch::einsum(
"bxi,lij,byj->bxyl", {aug_dep,
U, aug_head});