78 M h_head_r(h_head.rows() + 1, h_head.cols());
79 h_head_r.row(0) = p.
m_root.transpose();
80 h_head_r.block(1, 0, h_head.rows(), h_head.cols()) = h_head;
85 M aug_dep(h_dep.rows(), h_dep.cols() + 1);
86 aug_dep << h_dep, M::Ones(h_dep.rows(), 1);
87 M aug_head(h_head.rows(), h_head.cols() + 1);
88 aug_head << h_head, M::Ones(h_head.rows(), 1);
90 std::vector<M> logits;
91 logits.reserve(p.
m_U.size());
92 for (
const M& u : p.
m_U)
95 logits.push_back(aug_dep * u * aug_head.transpose());
105 const std::vector<uint32_t>& heads,
107 std::vector<uint32_t>& output)
const
109 const size_t n_labels = p.
m_U.size();
126 M h_head_r(h_head.rows() + 1, h_head.cols());
127 h_head_r.row(0) = p.
m_root.transpose();
128 h_head_r.block(1, 0, h_head.rows(), h_head.cols()) = h_head;
132 M aug_dep(h_dep.rows(), h_dep.cols() + 1);
133 aug_dep << h_dep, M::Ones(h_dep.rows(), 1);
134 M aug_head(h_head.rows(), h_head.cols() + 1);
135 aug_head << h_head, M::Ones(h_head.rows(), 1);
137 const Eigen::Index n_dep = aug_dep.rows();
138 const Eigen::Index d = aug_dep.cols();
142 M gathered(n_dep, d);
143 for (Eigen::Index i = 0; i < n_dep; ++i)
145 gathered.row(i) = aug_head.row((Eigen::Index) heads[input_begin + i]);
152 std::vector<Eigen::Index> best(n_dep, 0);
153 std::vector<T> best_score(n_dep, -std::numeric_limits<T>::infinity());
155 if (p.
m_U_stacked.cols() == (Eigen::Index) n_labels * d
159 for (
size_t l = 0; l < n_labels; ++l)
161 const auto slice = projected.block(0, (Eigen::Index) l * d, n_dep, d);
162 const V score = (slice.array() * gathered.array()).rowwise().sum();
163 for (Eigen::Index i = 0; i < n_dep; ++i)
165 if (score(i) > best_score[i])
167 best_score[i] = score(i);
168 best[i] = (Eigen::Index) l;
175 for (
size_t l = 0; l < n_labels; ++l)
177 const M tmp = aug_dep * p.
m_U[l];
178 const V score = (tmp.array() * gathered.array()).rowwise().sum();
179 for (Eigen::Index i = 0; i < n_dep; ++i)
181 if (score(i) > best_score[i])
183 best_score[i] = score(i);
184 best[i] = (Eigen::Index) l;
190 for (Eigen::Index i = 0; i < n_dep; ++i)
192 output[input_begin + i] = (uint32_t) best[i];