LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
RnnDependencyParser.cpp
Go to the documentation of this file.
1// SPDX-FileCopyrightText: 2022 CEA LIST <gael.de-chalendar@cea.fr>
2//
3// SPDX-License-Identifier: MIT
4
5#include <QtCore/QTemporaryFile>
6#include <QtCore/QRegularExpression>
7#include <QDir>
8
9
14
24
25#include <queue>
26
27#include "RnnDependencyParser.h"
30#include "deeplima/token_type.h"
35
36#include <thread>
37#include <algorithm>
38#include <eigen3/Eigen/Core>
39
40
41#define DEBUG_THIS_FILE true
42
44using namespace Lima::Common::PropertyCode;
45using namespace Lima::Common::MediaticData;
46using namespace Lima::Common::Misc;
48using namespace Lima::Common::AnnotationGraphs;
49using namespace deeplima;
50
51
53
54#if defined(DEBUG_LP) && defined(DEBUG_THIS_FILE)
55#define LOG_MESSAGE(stream, msg) stream << msg;
56#define LOG_MESSAGE_WITH_PROLOG(stream, msg) SALOGINIT; LOG_MESSAGE(stream, msg);
57#else
58 #define LOG_MESSAGE(stream, msg) ;
59 #define LOG_MESSAGE_WITH_PROLOG(stream, msg) ;
60#endif
61
63
65
67{
68public:
71 void init(GroupConfigurationStructure& unitConfiguration);
72 void analyzer(std::shared_ptr<TokenSequenceAnalyzer<>::TokenIterator> ti);
74 MediaId m_language;
75 QString m_data;
76 std::shared_ptr<DependencyParser> m_dependencyParser = nullptr;
77 // std::shared_ptr<TokenSequenceAnalyzer<>> m_sequenceAnalyser = nullptr;
78 std::function<void()> m_load_fn;
79 std::shared_ptr< StringIndex > m_stridx;
81 std::vector< std::map<std::string, std::string> > m_tags;
82 std::vector<QString> m_lemmas;
83 // std::vector<typename DependencyParser::token_with_analysis_t> m_tokens;
84 std::vector<std::string> m_class_names;
85 std::vector< std::vector<std::string> > m_classes;
86 std::vector<uint32_t> m_heads;
87 std::vector<std::string> m_deprels; // predicted deprel per token, parallel to m_heads
88 std::vector<bool> m_isRoot; // true for synthetic <ROOT> tokens (no real vertex)
89 std::vector<std::string> m_relClassNames; // deprel id -> string, from the model
91 bool m_enabled; // false when no parser model is available: process() is a no-op
92};
93
95 THIS_FILE_LOGGING_CATEGORY()),
96 m_stridx(new StringIndex()),
97 m_loaded(false),
98 m_enabled(true)
99{
100}
101
105
110
112{
113 LOG_MESSAGE_WITH_PROLOG(LDEBUG, "RnnDependencyParser::init");
114
115 m_d->m_language = manager->getInitializationParameters().media;
116
117 m_d->init(unitConfiguration);
118
119}
120
122{
124 TimeUtilsController RnnDependencyParserProcessTime("RnnDependencyParser");
125 SALOGINIT;
126 LOG_MESSAGE(LDEBUG, "RnnDependencyParser::process");
127 if (!m_d->m_enabled)
128 {
129 LOG_MESSAGE(LDEBUG, "RnnDependencyParser disabled (no model); skipping.");
130 return SUCCESS_ID;
131 }
132 // Ensure the model is loaded (no-op if already loaded or not lazy-initialized).
133 if (m_d->m_load_fn)
134 {
135 m_d->m_load_fn();
136 }
137 auto tiData = std::dynamic_pointer_cast<TokenIteratorData>(analysis.getData("TokenIterator"));
138 if (tiData == nullptr)
139 {
140 SALOGINIT;
141 LERROR << "Can't Process RnnDependencyParser : missing data 'TokenIterator'";
142 return MISSING_DATA;
143 }
144 auto stridxPtr = tiData->getStringIndex();
145 m_d->m_dependencyParser->setStringIndex(stridxPtr);
146 auto tokenIterator = tiData->getTokenIterator();
147 tokenIterator->reset();
148
149 auto anagraph = std::dynamic_pointer_cast<AnalysisGraph>(analysis.getData("PosGraph"));
150 if (anagraph == nullptr)
151 {
152 LERROR << "no PosGraph ! abort";
153 return MISSING_DATA;
154 }
155
156 auto syntacticData = std::dynamic_pointer_cast<SyntacticAnalysis::SyntacticData>(analysis.getData("SyntacticData"));
157 if (syntacticData == nullptr)
158 {
159 syntacticData = std::make_shared<SyntacticAnalysis::SyntacticData>(anagraph.get(), nullptr);
160 analysis.setData("SyntacticData",syntacticData);
161 }
162 syntacticData->setupDependencyGraph();
163
164 m_d->analyzer(tokenIterator);
165
166 const auto& languageData = static_cast<const Common::MediaticData::LanguageData&>(
168
169 auto sd = std::dynamic_pointer_cast<SegmentationData>(
170 analysis.getData(m_d->m_data.toStdString()));
171 if (sd == nullptr)
172 {
173 LERROR << "RnnDependencyParser: missing segmentation data '" << m_d->m_data << "'";
174 return MISSING_DATA;
175 }
176 LinguisticGraph* posGraph = anagraph->getGraph();
177 const LinguisticGraphVertex lastVertex = anagraph->lastVertex();
178
179 // The parser emits one sentence at a time: a synthetic <ROOT> followed by the
180 // sentence's real tokens, whose heads are sentence-LOCAL (0 == that sentence's
181 // root, 1..k == its tokens in order). Build, per sentence/segment, the ordered
182 // list of real token vertices by walking the PoS graph the same way the CoNLL
183 // dumper does (ordered[0] is the boundary, rendered as HEAD 0 / DEPREL root).
184 std::vector<LinguisticGraphVertex> segmentBegin;
185 std::vector<std::vector<LinguisticGraphVertex>> segmentTokens;
186 for (auto segIt = sd->getSegments().begin(); segIt != sd->getSegments().end();
187 ++segIt)
188 {
189 const LinguisticGraphVertex sentBegin = segIt->getFirstVertex();
190 const LinguisticGraphVertex sentEnd = segIt->getLastVertex();
191 segmentBegin.push_back(sentBegin);
192 segmentTokens.emplace_back();
193 std::vector<LinguisticGraphVertex>& tokens = segmentTokens.back();
194
195 std::queue<LinguisticGraphVertex> toVisit;
196 std::set<LinguisticGraphVertex> visited;
197 toVisit.push(sentBegin);
198 while (!toVisit.empty())
199 {
200 const LinguisticGraphVertex v = toVisit.front();
201 toVisit.pop();
202 if (visited.count(v) > 0)
203 {
204 continue;
205 }
206 visited.insert(v);
207 // Real tokens (skip the sentence boundary and any non-token vertex).
208 if (v != sentBegin && get(vertex_token, *posGraph, v) != nullptr)
209 {
210 tokens.push_back(v);
211 }
212 if (v == sentEnd)
213 {
214 break;
215 }
216 LinguisticGraphOutEdgeIt outIt, outItEnd;
217 for (boost::tie(outIt, outItEnd) = boost::out_edges(v, *posGraph);
218 outIt != outItEnd; ++outIt)
219 {
220 const LinguisticGraphVertex tgt = boost::target(*outIt, *posGraph);
221 if (visited.count(tgt) == 0 && tgt != lastVertex)
222 {
223 toVisit.push(tgt);
224 }
225 }
226 }
227 }
228
229 // Walk the parser output. Each <ROOT> starts a new sentence (matched to the
230 // next segment in order); the following tokens map to that sentence's vertices.
231 int segIdx = -1;
232 size_t localPos = 0; // 1-based position of the current token within its sentence
233 for (size_t i = 0; i < m_d->m_heads.size(); ++i)
234 {
235 if (i < m_d->m_isRoot.size() && m_d->m_isRoot[i])
236 {
237 ++segIdx;
238 localPos = 0;
239 continue;
240 }
241 ++localPos;
242 if (segIdx < 0 || segIdx >= static_cast<int>(segmentTokens.size()))
243 {
244 continue;
245 }
246 const std::vector<LinguisticGraphVertex>& tokens = segmentTokens[segIdx];
247 if (localPos > tokens.size())
248 {
249 continue; // parser/graph token counts disagree for this sentence
250 }
251 const uint32_t head = m_d->m_heads[i];
252 const std::string& deprel = (i < m_d->m_deprels.size())
253 ? m_d->m_deprels[i]
254 : std::string("dep");
256 languageData.getSyntacticRelationId(deprel);
257 const LinguisticGraphVertex src = tokens[localPos - 1];
258 // head == 0 -> this sentence's root boundary (HEAD 0 / DEPREL root);
259 // otherwise the head-th token of the same sentence.
260 const LinguisticGraphVertex dest =
261 (head == 0 || head > tokens.size()) ? segmentBegin[segIdx]
262 : tokens[head - 1];
263 syntacticData->addRelationNoChain(relType, src, dest);
264 }
265 TimeUtils::logElapsedTime("RnnDependencyParser");
266 return SUCCESS_ID;
267}
268
270{
271
272 m_data = QString(getStringParameter(unitConfiguration, "data", 0, "SentenceBoundaries").c_str());
273 QString dependency_parser_model_prefix = getStringParameter(unitConfiguration, "dependency_parser_model_prefix", ConfigurationHelper::REQUIRED | ConfigurationHelper::NOT_EMPTY).c_str();
274 QString tagger_model_prefix = getStringParameter(unitConfiguration, "tagger_model_prefix", ConfigurationHelper::REQUIRED | ConfigurationHelper::NOT_EMPTY).c_str();
275 LOG_MESSAGE_WITH_PROLOG(LDEBUG, "RnnDependencyParserPrivate::init dependency parser model: " << dependency_parser_model_prefix);
276
277 QString lang_str = MediaticData::single().media(m_language).c_str();
278 QString resources_path = MediaticData::single().getResourcesPath().c_str();
279 QString dependency_parser_name = dependency_parser_model_prefix;
280 QString tagger_model_name = tagger_model_prefix;
281
282 std::string udlang;
283 MediaticData::single().getOptionValue("udlang", udlang);
284
285 if (!fix_lang_codes(lang_str, udlang))
286 {
288 "RnnDependencyParserPrivate::init: Can't parse language id " << udlang.c_str(),
290 }
291
292 dependency_parser_name.replace(QString("$udlang"), QString(udlang.c_str()));
293 tagger_model_name.replace(QString("$udlang"), QString(udlang.c_str()));
294
295 auto dependency_parser_file_name = findFileInPaths(resources_path,
296 QString::fromUtf8("/RnnDependencyParser/%1/%2.pt")
297 .arg(lang_str, dependency_parser_name));
298
299 auto tagger_model_file_name = findFileInPaths(resources_path,
300 QString::fromUtf8("/RnnTagger/%1/%2.pt")
301 .arg(lang_str, tagger_model_name));
302 if (dependency_parser_file_name.isEmpty())
303 {
304 // No dependency parser model for this language yet: disable the unit
305 // rather than aborting the whole pipeline. process() becomes a no-op, so
306 // the rest of the pipeline (tagging, lemmatization, dumping) still runs.
307 SALOGINIT;
308 LWARN << "RnnDependencyParserPrivate::init: no dependency parser model found for "
309 << lang_str << " (" << dependency_parser_name
310 << "); dependency parsing disabled.";
311 m_enabled = false;
312 return;
313 }
314
315 if (tagger_model_file_name.isEmpty())
316 {
317 throw InvalidConfiguration("RnnTokensAnalyzerPrivate::init: tagger model file not found.");
318 }
319
320 LOG_MESSAGE(LDEBUG, "RnnDependencyParserPrivate::init call TokenSequenceAnalyzer<>().get_classes_from_fn");
321 TokenSequenceAnalyzer<>().get_classes_from_fn(tagger_model_file_name.toStdString(), m_class_names, m_classes);
322 auto temp_classes_names = m_class_names;
323 auto temp_classes = m_classes;
324 temp_classes_names.erase(temp_classes_names.begin()+1);
325 temp_classes.erase(temp_classes.begin()+1);
326 m_load_fn = [this, dependency_parser_file_name, tagger_model_file_name, temp_classes_names, temp_classes]()
327 {
328 if (m_loaded)
329 {
330 return;
331 }
332
333 // Disable Eigen intra-op (OpenMP) parallelism process-wide. The neural
334 // units already parallelize at the sentence/slot level, so letting Eigen
335 // also fork a parallel region for every small matmul only oversubscribes
336 // the cores. For the parser this is catastrophic: it does dozens of tiny
337 // biaffine products per sentence, and the OpenMP fork/join overhead dwarfs
338 // the arithmetic (measured ~20x slowdown). This is a global Eigen setting,
339 // so it also benefits the tagger/segmenter units in the same process.
340 Eigen::setNbThreads(1);
341
342 // The parser runs a single inference worker: its throughput is bounded by
343 // the serial feature-vectorization/feeding path (fastText), not by the
344 // inference, so extra workers do not help here (unlike the tagger). Keeping
345 // one worker also keeps analyzeText output independent of the core count.
346 m_dependencyParser = std::make_shared<DependencyParser>(dependency_parser_file_name.toStdString(),
348 for (size_t i = 0; i < temp_classes.size(); i++)
349 {
350 m_dependencyParser->set_classes(i, temp_classes_names[i], temp_classes[i]);
351 }
352 m_loaded = true;
353 };
354
355 if (!isInitLazy())
356 {
357 m_load_fn();
358 }
359
360
361}
362
364{
365 // deprel id -> string vocabulary from the model; empty if the model has no
366 // label decoder (then deprel() falls back to "dep").
367 m_relClassNames = m_dependencyParser->get_rel_class_names();
368 m_dependencyParser->register_handler([this](const StringIndex& stridx,
369 const std::vector<typename DependencyParser::token_with_analysis_t>& tokens,
370 std::shared_ptr< StdMatrix<uint32_t> > classes,
371 size_t begin,
372 size_t end)
373 {
374 typename DependencyParser::TokenIterator dti(stridx,
375 tokens,
376 classes,
377 begin,
378 end,
379 m_relClassNames.empty() ? nullptr : &m_relClassNames);
381 });
382 (*m_dependencyParser)(*ti);
383 m_dependencyParser->finalize();
384}
385
387{
388 // The iterator prefixes the stream with a synthetic <ROOT> token; real tokens
389 // carry global head indices (0 == root). Mark the <ROOT>(s) so process() skips
390 // them when assigning relations to actual graph vertices.
391 while (!ti.end())
392 {
393 m_heads.push_back(ti.head());
394 m_deprels.push_back(ti.deprel());
395 m_isRoot.push_back(std::string(ti.form()) == "<ROOT>");
396 ti.next();
397 }
398}
399
400}
This file is the main header file for the data related to annotation graphs.
#define CONFIGURATIONHELPER_LOGGING_INIT(X)
#define LOG_MESSAGE_WITH_PROLOG(stream, msg)
#define LOG_MESSAGE(stream, msg)
#define LWARN
Definition LimaCommon.h:160
#define LIMA_EXCEPTION_SELECT_LOGINIT(X, Y, Z)
This macro writes the message Y to the error stream configured by its first parameter X,...
Definition LimaCommon.h:332
#define LDEBUG
Definition LimaCommon.h:157
#define LERROR
Definition LimaCommon.h:161
A graph structure for linguistic analysis.
@ vertex_token
LinguisticGraph::vertex_descriptor LinguisticGraphVertex
LinguisticGraph::out_edge_iterator LinguisticGraphOutEdgeIt
boost::adjacency_list< boost::vecS, boost::vecS, boost::bidirectionalS, LinguisticVertexProperties > LinguisticGraph
Property to identify the chains in the graph.
#define SALOGINIT
#define RNNDEPENDENCYPARSER_CLASSID
Defines a Factory to create Object of type Base.
Data used for the syntactic analyzis of texts.
Holds all data that pass through the ProcessUnits Analysis data are shared pointers,...
std::shared_ptr< AnalysisData > getData(const QString &id)
return AnalysisData by id
void setData(const QString &id, std::shared_ptr< AnalysisData > data)
set an analysisData with the given id.
Holds linguistic data for one language.
const MediaData & mediaData(MediaId media) const
Manage initialization of InitializableObjects using configuration module and parameters.
const InitializationParameters & getInitializationParameters() const
get Initialization Parameters
Use this exception to signal an error in one of the configuration files.
Definition LimaCommon.h:345
void getStringParameter(Common::XMLConfigurationFiles::GroupConfigurationStructure &unitConfiguration, const std::string &name, std::string &value, int flags=Flags::REQUIRED, std::string default_value="")
void analyzer(std::shared_ptr< TokenSequenceAnalyzer<>::TokenIterator > ti)
LimaStatusCode process(AnalysisContent &analysis) const override
Process on data in analysisContent.
void init(Lima::Common::XMLConfigurationFiles::GroupConfigurationStructure &unitConfiguration, Manager *manager) override
initialize with parameters from configuration file.
static const MediaticData & single()
const singleton accessor
Definition Singleton.h:51
This file contains a class to control log of informations about time, such as logging cumulated time ...
static void logElapsedTime(const std::string &mess, const std::string &taskCategory=std::string(""))
log the number of microseconds since last UpdateCurrentTime
static void updateCurrentTime(const std::string &taskCategory=std::string(""))
store current time for new elapsed time computation
void get_classes_from_fn(const std::string &fn, std::vector< std::string > &classes_names, std::vector< std::vector< std::string > > &classes)
QString findFileInPaths(const QString &paths, const QString &fileName, const QChar &separator)
Find the given file in the given paths.
static SimpleFactory< MediaProcessUnit, RnnDependencyParser > RnnDependencyParserFactory(RNNDEPENDENCYPARSER_CLASSID)
bool fix_lang_codes(QString &lang_str, std::string &udlang)
LimaStatusCode
Definition LimaCommon.h:236
@ SUCCESS_ID
Definition LimaCommon.h:237
@ MISSING_DATA
Definition LimaCommon.h:243