LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
static_graph.h
Go to the documentation of this file.
1// Copyright 2021 CEA LIST
2// SPDX-FileCopyrightText: 2022 CEA LIST <gael.de-chalendar@cea.fr>
3//
4// SPDX-License-Identifier: MIT
5
6#ifndef DEEPLIMA_SRC_STATIC_GRAPH_STATIC_GRAPH_H
7#define DEEPLIMA_SRC_STATIC_GRAPH_STATIC_GRAPH_H
8
9#include <vector>
10#include <functional>
11
12#include <torch/torch.h>
13#include <torch/serialize/archive.h>
14
15#include "dict_base.h"
18// #include "nn/torch_modules/stanza_models_depparse_model.h"
19
20namespace deeplima
21{
22namespace nets
23{
24
31 class StaticGraphImpl : public torch::nn::Module
32 {
33 struct step_descr_t
34 {
35 enum step_type_t
36 {
37 unknown = 0,
38 def = 1,
39 forward = 2,
40 cat = 3,
41 log_softmax = 4,
42 reshape = 5,
43 unbind = 6,
44 unsqueeze = 7,
45 sigmoid = 8,
46 max_step_type
47 };
48
49 step_type_t m_type;
50 std::vector<std::string> m_names;
51 std::map<std::string, std::string> m_args;
52 std::map<std::string, std::vector<int64_t>> m_iargs;
53
54 explicit step_descr_t(step_type_t type = unknown, const std::vector<std::string> &names = {}, const std::map<std::string, std::string> &args = {})
55 : m_type(type), m_names(names), m_args(args)
56 {
57 if (m_type >= max_step_type)
58 {
59 throw std::runtime_error("Error in static graph");
60 }
61 }
62 };
63
64 struct context_t
65 {
66 std::vector<torch::Tensor> m_tensors;
67 explicit context_t(size_t num)
68 {
69 m_tensors.resize(num);
70 }
71 };
72
73 struct op_t
74 {
75 enum op_type_t
76 {
77 unknown = 1,
78 forward = 2,
79 max_op_type
80 };
81
82 op_t(op_type_t type = unknown)
83 : m_type(type)
84 {
85 if (m_type >= max_op_type)
86 {
87 throw std::runtime_error("Error in static graph");
88 }
89 }
90
91 op_type_t m_type;
92 std::vector<size_t> m_outputs; // identifiers of target context's tensors (this op is going to create them)
93 std::vector<size_t> m_inputs; // identifiers of input context's tensors (this op expects they are created)
94 std::function<void(context_t &)> m_fn;
95 std::string m_descr;
96 };
97
98 enum module_type_t
99 {
100 unknown = 0,
101 embedding = 1,
102 lstm = 2,
103 linear = 3,
104 dropout = 4,
105 deep_biaffine_attention_decoder = 5,
106 deep_biaffine_attention_label_decoder = 6,
107 // stanza_models_depparse_model = 7,
108 max_module_type
109 };
110
111 public:
112 StaticGraphImpl(const DictsHolder &dicts, const std::string &script);
113 StaticGraphImpl(DictsHolder &&dicts, const std::string &script);
115
116 virtual void load(torch::serialize::InputArchive &archive) override;
117 virtual void save(torch::serialize::OutputArchive &archive) const override;
118 virtual void set_tags(const std::map<std::string, std::string> &tags);
119 virtual const std::map<std::string, std::string> &get_tags() const
120 {
121 return m_tags;
122 }
123 virtual bool has_tag(const std::string &key) const
124 {
125 return m_tags.end() != m_tags.find(key);
126 }
127
128 virtual void to(torch::Device device, bool non_blocking = false) override;
129
130 std::map<std::string, torch::Tensor> forward(
131 const std::map<std::string, torch::Tensor> &inputs,
132 const std::string &output_name)
133 {
134 std::vector<std::string> output_names = {output_name};
135 return forward(inputs, output_names.cbegin(), output_names.cend());
136 }
137
138 std::map<std::string, torch::Tensor> forward(
139 const std::map<std::string, torch::Tensor> &inputs,
140 const std::vector<std::string> &output_names)
141 {
142 return forward(inputs, output_names.cbegin(), output_names.cend());
143 }
144
145 template <class InputIt>
146 std::map<std::string, torch::Tensor> forward(
147 const std::map<std::string, torch::Tensor> &inputs,
148 InputIt outputs_begin,
149 InputIt outputs_end)
150 {
151 assert(inputs.size() > 0);
152 assert(outputs_begin != outputs_end);
153
154 std::string key = get_exec_plan_key(inputs, outputs_begin, outputs_end);
155 if (m_exec_plans.end() == m_exec_plans.find(key))
156 {
157 prepare_exec_plan(inputs, outputs_begin, outputs_end);
158 }
159
160 const std::vector<size_t> plan = m_exec_plans[key];
161 context_t ctx(m_tensor_name_to_idx.size());
162 for (const auto &kv : inputs)
163 {
164 if (m_tensor_name_to_idx.end() == m_tensor_name_to_idx.find(kv.first))
165 {
166 throw std::runtime_error("Error in static graph");
167 }
168 size_t idx = m_tensor_name_to_idx[kv.first];
169 if (idx >= ctx.m_tensors.size())
170 {
171 throw std::runtime_error("Error in static graph");
172 }
173 ctx.m_tensors[idx] = kv.second;
174 }
175
176 // size_t total_time = 0;
177 for (size_t i = 0; i < plan.size(); i++)
178 {
179 size_t op_idx = plan[i];
180 // chrono::steady_clock::time_point begin = chrono::steady_clock::now();
181 // Call of train function (forward ?)
182 m_ops[op_idx].m_fn(ctx);
183 // chrono::steady_clock::time_point end = chrono::steady_clock::now();
184 // cerr << "\"" << m_ops[op_idx].m_descr << "\" "
185 // << chrono::duration_cast<chrono::milliseconds>(end - begin).count() << "[ms]" << endl;
186 // total_time += chrono::duration_cast<chrono::milliseconds>(end - begin).count();
187 }
188 // cerr << "Total calculations: " << total_time << "\r";
189
190 std::map<std::string, torch::Tensor> rv;
191 for (auto it = outputs_begin; it != outputs_end; ++it)
192 {
193 const std::string &str = *it;
194 if (m_tensor_name_to_idx.end() == m_tensor_name_to_idx.find(str))
195 {
196 throw std::runtime_error("Error in static graph");
197 }
198 size_t idx = m_tensor_name_to_idx[str];
199 if (idx >= ctx.m_tensors.size())
200 {
201 throw std::runtime_error("Error in static graph");
202 }
203 rv[str] = ctx.m_tensors[idx];
204
205 // cout << "Output [" << idx << "].size() == " << ctx.m_tensors[idx].sizes() << endl;
206 }
207
208 return rv;
209 }
210
211 virtual void pretty_dump(std::ostream &stream) const;
212
213 const DictsHolder &get_dicts() const
214 {
215 return m_dicts;
216 }
217
218 protected:
219 virtual void parse_script(const std::string &script);
220 virtual step_descr_t parse_script_line(const std::string &line);
221 virtual std::map<std::string, std::string> parse_options(std::istringstream &ss);
222 virtual std::vector<std::string> parse_list(std::string &str, char sep);
223 virtual std::vector<int64_t> parse_iargs(const std::string &str);
224
225 std::map<std::string, std::vector<size_t>> m_exec_plans;
226
227 template <class InputIt>
229 const std::map<std::string, torch::Tensor> &inputs,
230 InputIt outputs_begin,
231 InputIt outputs_end)
232 {
233 assert(inputs.size() > 0);
234 assert(outputs_begin != outputs_end);
235
236 std::string key = get_exec_plan_key(inputs, outputs_begin, outputs_end);
237 if (m_exec_plans.end() != m_exec_plans.find(key))
238 {
239 throw std::runtime_error("Error in static graph");
240 }
241
242 std::set<size_t> given_inputs_idx;
243 for (const auto &kv : inputs)
244 {
245 size_t in_idx = m_tensor_name_to_idx[kv.first];
246 given_inputs_idx.insert(in_idx);
247 }
248
249 std::vector<size_t> temp;
250 temp.reserve(m_ops.size() * 2);
251
252 std::list<size_t> required;
253 for (auto it = outputs_begin; it != outputs_end; ++it)
254 {
255 const std::string &tensor_name = *it;
256 size_t idx = m_tensor_name_to_idx[tensor_name];
257 required.push_back(idx);
258 }
259
260 while (!required.empty())
261 {
262 size_t out_idx = required.front();
263 required.pop_front();
264 size_t op_idx = m_outidx_to_opidx[out_idx];
265 temp.push_back(op_idx);
266
267 const op_t &op = m_ops[op_idx];
268 for (size_t in_idx : op.m_inputs)
269 {
270 if (given_inputs_idx.end() == given_inputs_idx.find(in_idx))
271 {
272 required.push_back(in_idx);
273 // given_inputs_idx.insert(in_idx);
274 }
275 }
276 }
277
278 std::set<size_t> planned_ops;
279 std::vector<size_t> plan;
280 plan.reserve(temp.size());
281
282 for (int32_t i = temp.size() - 1; i >= 0; i--)
283 {
284 if (planned_ops.end() != planned_ops.find(temp[i]))
285 {
286 continue;
287 }
288 plan.push_back(temp[i]);
289 planned_ops.insert(temp[i]);
290
291 // cout << m_ops[temp[i]].m_descr << endl;
292 }
293
294 m_exec_plans[key] = plan;
295 }
296
297 template <class InputIt>
298 std::string get_exec_plan_key(
299 const std::map<std::string, torch::Tensor> &inputs,
300 InputIt outputs_begin,
301 InputIt outputs_end)
302 {
303 std::string k;
304 for (const auto &kv : inputs)
305 {
306 if (k.size() > 0)
307 {
308 k += " ";
309 }
310 k += kv.first;
311 }
312
313 k += " -> ";
314
315 for (auto it = outputs_begin; it != outputs_end; ++it)
316 {
317 k += *it;
318 }
319
320 return k;
321 }
322
323 template <class T>
324 T get_option(const std::map<std::string, std::string> &opts, const std::string &name)
325 {
326 const auto it = opts.find(name);
327 if (opts.cend() == it)
328 {
329 throw std::runtime_error("Error in static graph");
330 }
331 T v;
332 std::istringstream ss(it->second);
333 ss >> v;
334 std::string s;
335 if (ss >> s)
336 {
337 throw std::runtime_error("Error in static graph");
338 }
339
340 return v;
341 }
342
343 bool get_bool_option(const std::map<std::string, std::string> &opts, const std::string &name)
344 {
345 const auto it = opts.find(name);
346 if (opts.cend() == it)
347 {
348 throw std::runtime_error("Error in static graph");
349 }
350
351 if (it->second == "true")
352 {
353 return true;
354 }
355 else if (it->second == "false")
356 {
357 return false;
358 }
359 throw std::runtime_error("Error in static graph");
360 }
361
362 virtual void create_arg(const std::vector<std::string> &names, const std::map<std::string, std::string> &opts);
363
364 virtual void create_submodule_Embedding(const std::string &name, const std::map<std::string, std::string> &opts);
365 virtual void create_submodule_LSTM(const std::string &name, const std::map<std::string, std::string> &opts);
366 virtual void create_submodule_Linear(const std::string &name, const std::map<std::string, std::string> &opts);
367 virtual void create_submodule_Dropout(const std::string &name, const std::map<std::string, std::string> &opts);
368 virtual void create_submodule_DeepBiaffineAttentionDecoder(const std::string &name, const std::map<std::string, std::string> &opts);
369 virtual void create_submodule_DeepBiaffineAttentionLabelDecoder(const std::string &name, const std::map<std::string, std::string> &opts);
370 // virtual void create_submodule_StanzaDepparseParser(const std::string &name, const std::map<std::string, std::string> &opts);
371
373 std::string m_script;
374 std::map<std::string, std::string> m_tags;
375
376 std::vector<torch::nn::Embedding> m_embedding;
377 std::vector<torch::nn::LSTM> m_lstm;
378 std::vector<torch::nn::Linear> m_linear;
379 std::vector<torch::nn::Dropout> m_dropout;
380 std::vector<deeplima::nets::torch_modules::DeepBiaffineAttentionDecoder> m_deep_biaffine_attention_decoder;
381 std::vector<deeplima::nets::torch_modules::DeepBiaffineAttentionLabelDecoder> m_deep_biaffine_attention_label_decoder;
382 // std::shared_ptr<deeplima::nets::torch_modules::StanzaDepparseParser> m_stanza_depparse_parser;
383
384 public:
385 const std::string &get_script() const
386 {
387 return m_script;
388 }
389
390 // Introspection for Torch -> Eigen converter
391 const std::vector<torch::nn::LSTM> &get_layers_lstm() const
392 {
393 return m_lstm;
394 }
395
396 const std::vector<torch::nn::Linear> &get_layers_linear() const
397 {
398 return m_linear;
399 }
400
401 const std::vector<torch::nn::Embedding> &get_layers_embedding() const
402 {
403 return m_embedding;
404 }
405
406 const std::vector<torch::nn::Dropout> &get_layers_dropout() const
407 {
408 return m_dropout;
409 }
410
411 const std::vector<deeplima::nets::torch_modules::DeepBiaffineAttentionDecoder> &get_layers_deep_biaffine_attn_decoder() const
412 {
414 }
415
416 const std::vector<deeplima::nets::torch_modules::DeepBiaffineAttentionLabelDecoder> &get_layers_deep_biaffine_attn_label_decoder() const
417 {
419 }
420
421 // const deeplima::nets::torch_modules::StanzaDepparseParser &get_layers_stanza_depparse_parser() const
422 // {
423 // return *m_stanza_depparse_parser;
424 // }
425
426 torch::nn::Embedding get_module_by_name(const std::string &name) const
427 {
428 const auto it = m_modules.find(name);
429 if (m_modules.end() == it)
430 {
431 throw std::runtime_error("Unknown module name");
432 }
433 const module_ref_t &mr = it->second;
434 assert(module_type_t::embedding == mr.m_type);
435 return m_embedding[mr.m_idx];
436 }
437
438 const std::string get_module_name(size_t idx, const std::string &type) const
439 {
440 module_type_t t = unknown;
441 if (type == "lstm")
442 {
443 t = lstm;
444 }
445 else if (type == "embedding")
446 {
447 t = embedding;
448 }
449 else if (type == "linear")
450 {
451 t = linear;
452 }
453 else if (type == "dropout")
454 {
455 t = dropout;
456 }
457 else if (type == "deep_biaffine_attention_decoder")
458 {
459 t = deep_biaffine_attention_decoder;
460 }
461 // else if (type == "m_stanza_depparse_parser")
462 // {
463 // t = stanza_models_depparse_model;
464 // }
465 else
466 {
467 throw std::runtime_error("Unknown module type");
468 }
469
470 for (const auto &p : m_modules)
471 {
472 if (idx == p.second.m_idx && t == p.second.m_type)
473 {
474 return p.first;
475 }
476 }
477
478 return std::string("");
479 }
480
481 protected:
482 virtual void init_rnns();
483
485 {
486 module_type_t m_type;
487 size_t m_idx;
488
489 explicit module_ref_t(module_type_t type = module_type_t::unknown, size_t idx = 0)
490 : m_type(type), m_idx(idx)
491 {
492 if (m_type >= module_type_t::max_module_type)
493 {
494 throw std::runtime_error("Error in static graph");
495 }
496 }
497 };
498
499 std::map<std::string, module_ref_t> m_modules;
500 std::set<std::string> m_args;
501 std::map<std::string, size_t> m_tensor_name_to_idx;
502 std::vector<op_t> m_ops;
503 std::map<size_t, size_t> m_outidx_to_opidx; // output tensor's idx -> op idx
504 };
505
506 TORCH_MODULE(StaticGraph);
507
508 }
509}
510
511#endif
Implementation for Torch of the Tensorflow execution graph Used in training only.
virtual std::vector< int64_t > parse_iargs(const std::string &str)
std::vector< torch::nn::LSTM > m_lstm
const std::vector< torch::nn::Dropout > & get_layers_dropout() const
const std::vector< torch::nn::Linear > & get_layers_linear() const
std::vector< torch::nn::Embedding > m_embedding
std::vector< deeplima::nets::torch_modules::DeepBiaffineAttentionDecoder > m_deep_biaffine_attention_decoder
virtual step_descr_t parse_script_line(const std::string &line)
virtual void create_submodule_DeepBiaffineAttentionLabelDecoder(const std::string &name, const std::map< std::string, std::string > &opts)
const std::vector< deeplima::nets::torch_modules::DeepBiaffineAttentionDecoder > & get_layers_deep_biaffine_attn_decoder() const
void prepare_exec_plan(const std::map< std::string, torch::Tensor > &inputs, InputIt outputs_begin, InputIt outputs_end)
std::map< std::string, torch::Tensor > forward(const std::map< std::string, torch::Tensor > &inputs, const std::vector< std::string > &output_names)
virtual void create_submodule_LSTM(const std::string &name, const std::map< std::string, std::string > &opts)
const std::vector< deeplima::nets::torch_modules::DeepBiaffineAttentionLabelDecoder > & get_layers_deep_biaffine_attn_label_decoder() const
const std::string get_module_name(size_t idx, const std::string &type) const
virtual void set_tags(const std::map< std::string, std::string > &tags)
virtual std::vector< std::string > parse_list(std::string &str, char sep)
const std::vector< torch::nn::LSTM > & get_layers_lstm() const
bool get_bool_option(const std::map< std::string, std::string > &opts, const std::string &name)
T get_option(const std::map< std::string, std::string > &opts, const std::string &name)
torch::nn::Embedding get_module_by_name(const std::string &name) const
virtual void create_submodule_Dropout(const std::string &name, const std::map< std::string, std::string > &opts)
std::vector< deeplima::nets::torch_modules::DeepBiaffineAttentionLabelDecoder > m_deep_biaffine_attention_label_decoder
std::vector< torch::nn::Dropout > m_dropout
virtual void create_submodule_Embedding(const std::string &name, const std::map< std::string, std::string > &opts)
virtual void create_submodule_DeepBiaffineAttentionDecoder(const std::string &name, const std::map< std::string, std::string > &opts)
std::map< std::string, module_ref_t > m_modules
virtual bool has_tag(const std::string &key) const
virtual std::map< std::string, std::string > parse_options(std::istringstream &ss)
std::set< std::string > m_args
std::map< std::string, std::string > m_tags
virtual void create_arg(const std::vector< std::string > &names, const std::map< std::string, std::string > &opts)
virtual void pretty_dump(std::ostream &stream) const
std::vector< torch::nn::Linear > m_linear
virtual const std::map< std::string, std::string > & get_tags() const
virtual void to(torch::Device device, bool non_blocking=false) override
std::map< std::string, size_t > m_tensor_name_to_idx
virtual void load(torch::serialize::InputArchive &archive) override
const std::vector< torch::nn::Embedding > & get_layers_embedding() const
std::map< std::string, torch::Tensor > forward(const std::map< std::string, torch::Tensor > &inputs, const std::string &output_name)
virtual void save(torch::serialize::OutputArchive &archive) const override
std::string get_exec_plan_key(const std::map< std::string, torch::Tensor > &inputs, InputIt outputs_begin, InputIt outputs_end)
virtual void create_submodule_Linear(const std::string &name, const std::map< std::string, std::string > &opts)
const std::string & get_script() const
const DictsHolder & get_dicts() const
std::map< std::string, torch::Tensor > forward(const std::map< std::string, torch::Tensor > &inputs, InputIt outputs_begin, InputIt outputs_end)
virtual void parse_script(const std::string &script)
std::map< std::string, std::vector< size_t > > m_exec_plans
std::map< size_t, size_t > m_outidx_to_opidx
TORCH_MODULE(BiRnnSeq2Seq)
module_ref_t(module_type_t type=module_type_t::unknown, size_t idx=0)