17 int64_t batch_size = input.size(1);
20 torch::Tensor input_t = input.transpose(0, 1);
23 torch::Tensor arc_head;
26 torch::Tensor roots2 = torch::tile(
root2, { batch_size, 1, 1 });
28 arc_head =
elu(torch::cat({ roots2,
mlp_head(input_t) }, 1));
40 torch::Tensor arc_dep;
43 torch::Tensor roots = torch::tile(
root, { batch_size, 1, 1 });
45 arc_dep =
elu(torch::cat({ roots,
mlp_dep(input_t) }, 1));
55 torch::Tensor W = torch::matmul(arc_head,
U1);
57 torch::Tensor Wx = torch::matmul(W, arc_dep.transpose(1, 2));
61 torch::Tensor b = torch::matmul(arc_head, torch::tile(
u2, { 1, arc_head.size(1) }));
64 torch::Tensor r = torch::add(Wx, b);
69 r = r.index({ torch::indexing::Slice(),
70 torch::indexing::Slice(1, r.size(1)),
71 torch::indexing::Slice()});