LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
deep_biaffine_attn_label_decoder.h
Go to the documentation of this file.
1// Copyright 2026 CEA LIST
2// SPDX-FileCopyrightText: 2026 CEA LIST <gael.de-chalendar@cea.fr>
3//
4// SPDX-License-Identifier: MIT
5
6#ifndef DEEPLIMA_SRC_INFERENCE_EIGEN_DEEP_BIAFFINE_ATTN_LABEL_DECODER_H
7#define DEEPLIMA_SRC_INFERENCE_EIGEN_DEEP_BIAFFINE_ATTN_LABEL_DECODER_H
8
9#include <cmath>
10#include <vector>
11
12#include <eigen3/Eigen/Dense>
13#include "op_base.h"
14
15namespace deeplima
16{
17namespace eigen_impl
18{
19
20// Eigen (CPU-inference) counterpart of
21// nets::torch_modules::DeepBiaffineAttentionLabelDecoder. One sentence at a time.
22template<class M=Eigen::MatrixXf, class V=Eigen::VectorXf>
24{
29 std::vector<M> m_U; // one (hidden+1) x (hidden+1) matrix per label
30 // The m_U matrices concatenated column-wise into a single [(hidden+1), L*(hidden+1)]
31 // matrix, so the per-label products can be computed with one GEMM instead of L
32 // tiny ones. Built once at load time (build_stacked_U); if empty, predict_labels
33 // falls back to the per-label loop.
35 V m_root; // head-side <ROOT> row, used when !m_input_includes_root
37
38 // Concatenate m_U[0..L-1] horizontally into m_U_stacked. Call once after m_U is
39 // populated (e.g. at model conversion time).
41 {
42 if (m_U.empty())
43 {
44 m_U_stacked = M();
45 return;
46 }
47 const Eigen::Index d = m_U.front().rows(); // hidden+1
48 const Eigen::Index L = (Eigen::Index) m_U.size();
49 m_U_stacked.resize(d, L * d);
50 for (Eigen::Index l = 0; l < L; ++l)
51 {
52 m_U_stacked.block(0, l * d, d, d) = m_U[l];
53 }
54 }
55};
56
57template<class M, class V, class T>
59{
60public:
62
63 // Per-label label logits for one sentence.
64 // input: [input_dim, n_tokens] (token features as columns).
65 // returns: a vector of m_U.size() matrices, each [n_dep, n_head], where
66 // n_dep = n_tokens
67 // n_head = n_tokens (+1 when a root row is prepended, i.e. !input_includes_root)
68 std::vector<M> compute_logits(const params_t& p, const M& input) const
69 {
70 // h = elu(W x + b), as [n_tokens, hidden]
71 M h_dep = ((p.m_weight_dep * input).colwise() + p.m_bias_dep).transpose();
72 elu_inplace(h_dep);
73 M h_head = ((p.m_weight_head * input).colwise() + p.m_bias_head).transpose();
74 elu_inplace(h_head);
75
77 {
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;
81 h_head = h_head_r;
82 }
83
84 // Append a constant 1 column (affine augmentation).
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);
89
90 std::vector<M> logits;
91 logits.reserve(p.m_U.size());
92 for (const M& u : p.m_U)
93 {
94 // [n_dep, h+1] (h+1, h+1) (h+1, n_head) -> [n_dep, n_head]
95 logits.push_back(aug_dep * u * aug_head.transpose());
96 }
97 return logits;
98 }
99
100 // For each dependent token, score the labels at its given head and take the
101 // argmax. heads[i] is the head index of token input_begin + i (in the head
102 // space, i.e. already accounting for the root row when applicable).
104 const M& input,
105 const std::vector<uint32_t>& heads,
106 size_t input_begin,
107 std::vector<uint32_t>& output) const
108 {
109 const size_t n_labels = p.m_U.size();
110 if (n_labels == 0)
111 {
112 return;
113 }
114
115 // Reproduce the augmented dep/head representations exactly as compute_logits,
116 // but score only the head each token actually got (already decoded by the arc
117 // decoder). We therefore never materialise the full [n_dep x n_head] logit
118 // matrices; we gather each token's head row and take a row-wise dot product.
119 M h_dep = ((p.m_weight_dep * input).colwise() + p.m_bias_dep).transpose();
120 elu_inplace(h_dep);
121 M h_head = ((p.m_weight_head * input).colwise() + p.m_bias_head).transpose();
122 elu_inplace(h_head);
123
125 {
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;
129 h_head = h_head_r;
130 }
131
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);
136
137 const Eigen::Index n_dep = aug_dep.rows();
138 const Eigen::Index d = aug_dep.cols(); // hidden+1
139
140 // gathered.row(i) = aug_head.row(head_of_token_i). heads[] is already in head
141 // space (root row accounted for), matching aug_head's rows.
142 M gathered(n_dep, d);
143 for (Eigen::Index i = 0; i < n_dep; ++i)
144 {
145 gathered.row(i) = aug_head.row((Eigen::Index) heads[input_begin + i]);
146 }
147
148 // For each label l, score_l(i) = aug_dep.row(i) * U_l * gathered.row(i)^T.
149 // Compute aug_dep * U_l for all labels at once via the pre-stacked U
150 // ([d, L*d]) -> one GEMM producing [n_dep, L*d]; then a row-wise dot with
151 // gathered per label slice. Falls back to per-label GEMMs if U isn't stacked.
152 std::vector<Eigen::Index> best(n_dep, 0);
153 std::vector<T> best_score(n_dep, -std::numeric_limits<T>::infinity());
154
155 if (p.m_U_stacked.cols() == (Eigen::Index) n_labels * d
156 && p.m_U_stacked.rows() == d)
157 {
158 const M projected = aug_dep * p.m_U_stacked; // [n_dep, L*d]
159 for (size_t l = 0; l < n_labels; ++l)
160 {
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)
164 {
165 if (score(i) > best_score[i])
166 {
167 best_score[i] = score(i);
168 best[i] = (Eigen::Index) l;
169 }
170 }
171 }
172 }
173 else
174 {
175 for (size_t l = 0; l < n_labels; ++l)
176 {
177 const M tmp = aug_dep * p.m_U[l]; // [n_dep, d]
178 const V score = (tmp.array() * gathered.array()).rowwise().sum();
179 for (Eigen::Index i = 0; i < n_dep; ++i)
180 {
181 if (score(i) > best_score[i])
182 {
183 best_score[i] = score(i);
184 best[i] = (Eigen::Index) l;
185 }
186 }
187 }
188 }
189
190 for (Eigen::Index i = 0; i < n_dep; ++i)
191 {
192 output[input_begin + i] = (uint32_t) best[i];
193 }
194 }
195
196protected:
197 inline void elu_inplace(M& m) const
198 {
199 for (Eigen::Index r = 0; r < m.rows(); ++r)
200 {
201 for (Eigen::Index c = 0; c < m.cols(); ++c)
202 {
203 if (m(r, c) < 0)
204 {
205 m(r, c) = std::exp(m(r, c)) - 1;
206 }
207 }
208 }
209 }
210};
211
212} // namespace eigen_impl
213} // namespace deeplima
214
215#endif
std::vector< M > compute_logits(const params_t &p, const M &input) const
params_deep_biaffine_attn_label_decoder_t< M, V > params_t
void predict_labels(const params_t &p, const M &input, const std::vector< uint32_t > &heads, size_t input_begin, std::vector< uint32_t > &output) const