LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
static_graph.cpp
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#include "static_graph.h"
7#include "dict.h"
8
10
11using torch::indexing::Slice;
12using namespace torch;
13using namespace std;
14
15namespace deeplima
16{
17namespace nets
18{
19
20StaticGraphImpl::StaticGraphImpl(const DictsHolder& dicts, const string& script)
21 : m_dicts(dicts),
22 m_script(script)
23{
25 init_rnns();
26}
27
28StaticGraphImpl::StaticGraphImpl(DictsHolder&& dicts, const std::string& script)
29 : m_dicts(std::move(dicts)),
30 m_script(script)
31{
33 init_rnns();
34}
35
36void StaticGraphImpl::load(serialize::InputArchive& archive)
37{
38 c10::IValue val;
39 archive.read("model_script", val);
40 m_script = *(val.toString().get());
41
42 // Load dicts
43 c10::IValue ival_types_of_dicts;
44 archive.read("types_of_dicts", ival_types_of_dicts);
45 const c10::List<c10::IValue>& types_of_dicts = ival_types_of_dicts.toList();
46
47 m_dicts.resize(types_of_dicts.size());
48 for (size_t i = 0; i < m_dicts.size(); i++)
49 {
50 std::ostringstream dict_id;
51 dict_id << "dict_" << i;
52
53 if (!types_of_dicts.get(i).isString())
54 {
55 throw std::runtime_error("Error in static graph");
56 }
57 const std::string& dict_type = types_of_dicts.get(i).toStringRef();
58
59 c10::IValue v;
60 if (archive.try_read(dict_id.str(), v))
61 {
62 if (!v.isList())
63 {
64 throw std::runtime_error("Error in static graph");
65 }
66
67 if (dict_type == UInt64Dict::class_id())
68 {
69 m_dicts[i] = std::make_shared<UInt64Dict>();
70 }
71 else if (dict_type == WstringDict::class_id())
72 {
73 m_dicts[i] = std::make_shared<WstringDict>();
74 }
75 else if (dict_type == StringDict::class_id())
76 {
77 m_dicts[i] = std::make_shared<StringDict>();
78 }
79 else if (dict_type == Char32Dict::class_id())
80 {
81 m_dicts[i] = std::make_shared<Char32Dict>();
82 }
83 else
84 {
85 throw runtime_error("Unknown dict type.");
86 }
87
88 m_dicts[i]->fromIValue(v);
89 }
90 }
91
92 parse_script(m_script); // calls "register_module"
93
94 // Load tags
95 c10::IValue v;
96 if (archive.try_read("tags", v))
97 {
98 if (!v.isGenericDict())
99 {
100 throw runtime_error("\"tags\" must be a dict of strings.");
101 }
102
103 c10::Dict<c10::IValue, c10::IValue> d = v.toGenericDict();
104 for ( const auto &it : d )
105 {
106 if (!it.key().isString())
107 {
108 std::cerr << "ERROR: keys in \"tags\" dict must be strings" << endl;
109 }
110 if (!it.value().isString())
111 {
112 std::cerr << "ERROR: values in \"tags\" dict must be strings" << endl;
113 }
114
115 if (m_tags.end() != m_tags.find(it.key().toStringRef()))
116 {
117 throw std::runtime_error("Duplicated tags in the model");
118 }
119
120 m_tags[it.key().toStringRef()] = it.value().toStringRef();
121 }
122 }
123
124 torch::nn::Module::load(archive);
125}
126
127void StaticGraphImpl::save(serialize::OutputArchive& archive) const
128{
129 archive.write("model_script", m_script);
130
131 torch::nn::Module::save(archive);
132
133 // Save dicts
134 c10::List<std::string> types_of_dicts;
135 for (size_t i = 0; i < m_dicts.size(); i++)
136 {
137 //cerr << i << "\t" << m_dicts[i]->get_class_id() << std::endl;
138 types_of_dicts.push_back(m_dicts[i]->get_class_id());
139 }
140
141 archive.write("types_of_dicts", types_of_dicts);
142
143 for (size_t i = 0; i < m_dicts.size(); i++)
144 {
145 ostringstream s;
146 s << "dict_" << i;
147 archive.write(s.str(), m_dicts[i]->toIValue());
148 }
149
150 // Save tags
151 c10::Dict<string, string> temp_tags;
152 for ( const auto &it : m_tags )
153 {
154 temp_tags.insert(it.first, it.second);
155 }
156
157 archive.write("tags", temp_tags);
158}
159
160void StaticGraphImpl::set_tags(const map<string, string>& tags)
161{
162 for ( const auto &it : tags)
163 {
164 m_tags.insert({ it.first, it.second });
165 }
166}
167
168void StaticGraphImpl::pretty_dump(ostream &stream) const
169{
170 size_t counter = 0;
171 for ( const auto& m : modules(false) )
172 {
173 stream << m->name() << endl;
174 for ( const auto& t : m->parameters() )
175 {
176 int64_t c = 1;
177 for ( int64_t v : t.sizes() )
178 {
179 c *= v;
180 }
181 stream << "\t" << t.name() << " " << c << " " << t.sizes() << endl;
182 counter += c;
183 }
184 }
185 stream << "Total parameters = " << counter << endl;
186}
187
188void StaticGraphImpl::to(torch::Device device, bool non_blocking)
189{
190 for (auto& m : m_embedding) m->to(device);
191 for (auto& m : m_lstm) m->to(device);
192 for (auto& m : m_linear) m->to(device);
193 for (auto& m : m_dropout) m->to(device);
194 for (auto& m : m_deep_biaffine_attention_decoder) m->to(device);
195 for (auto& m : m_deep_biaffine_attention_label_decoder) m->to(device);
196 torch::nn::Module::to(device, non_blocking);
197 // cerr << "StaticGraphImpl::to( " << device << " )" << std::endl;
198}
199
200void StaticGraphImpl::parse_script(const string& script)
201{
202 // cerr << script << endl;
203 stringstream ss(script);
204 string line;
205
206 while (getline(ss, line))
207 {
208 if (line.size() == 0)
209 {
210 continue;
211 }
212 step_descr_t step = parse_script_line(line);
213 //cerr << line << endl;
214
215 if (step.m_type == step_descr_t::def)
216 {
217 continue;
218 }
219
220 // generate op
221 op_t op;
222 op.m_descr = line;
223
224 size_t known_tensors = m_tensor_name_to_idx.size();
225 for ( const string& name : step.m_names )
226 {
227 if (m_tensor_name_to_idx.end() != m_tensor_name_to_idx.find(name))
228 {
229 throw std::runtime_error("Error in static graph");
230 }
231 m_tensor_name_to_idx[name] = known_tensors;
232 op.m_outputs.push_back(m_tensor_name_to_idx[name]);
233 known_tensors += 1;
234 }
235
236 vector<string> input_tensor_names = parse_list(step.m_args["input"], ',');
237 for ( const string& name : input_tensor_names )
238 {
239 auto it = m_tensor_name_to_idx.find(name);
240 if (m_tensor_name_to_idx.end() == it)
241 {
242 throw std::runtime_error("Error in static graph");
243 }
244 op.m_inputs.push_back(it->second);
245 }
246
247 switch (step.m_type)
248 {
249 case step_descr_t::forward:
250 {
251 string module_name = get_option<string>(step.m_args, "module");
252 auto it = m_modules.find(module_name);
253 if (m_modules.end() == it)
254 {
255 throw std::runtime_error("Error in static graph");
256 }
257 module_ref_t module_ref = it->second;
258 switch (module_ref.m_type)
259 {
260 case module_type_t::embedding:
261 {
262 torch::nn::Embedding& m = m_embedding[module_ref.m_idx];
263 vector<size_t>& inputs = op.m_inputs;
264 vector<size_t>& outputs = op.m_outputs;
265
266 if (inputs.size() != 1 || outputs.size() != 1)
267 {
268 throw std::runtime_error("Error in static graph");
269 }
270
271 op.m_fn = [&m, inputs, outputs](context_t& ctx)
272 {
273 auto out = m->forward(ctx.m_tensors[inputs[0]]);
274 //cerr << ctx.m_tensors[inputs[0]] << endl;
275 //cerr << out << endl;
276 ctx.m_tensors[outputs[0]] = out;
277 };
278 }
279 break;
280
281 case module_type_t::lstm:
282 {
283 torch::nn::LSTM& m = m_lstm[module_ref.m_idx];
284 const vector<size_t>& inputs = op.m_inputs;
285 const vector<size_t>& outputs = op.m_outputs;
286
287 if (inputs.size() < 1)
288 {
289 throw std::runtime_error("Error in static graph");
290 }
291
292 op.m_fn = [&m, inputs, outputs](context_t& ctx)
293 {
294 //cerr << ctx.m_tensors[inputs[0]].sizes() << endl;
295 torch::optional<std::tuple<torch::Tensor, torch::Tensor>> h0_and_c0; // (h0, c0)
296 if (inputs.size() > 1)
297 {
298 if (3 == inputs.size())
299 {
300 h0_and_c0 = { ctx.m_tensors[inputs[1]], ctx.m_tensors[inputs[2]] };
301 }
302 else
303 {
304 throw std::runtime_error("Error in static graph");
305 }
306 }
307 auto out = m->forward(ctx.m_tensors[inputs[0]], h0_and_c0);
308 //cerr << ctx.m_tensors[inputs[0]] << endl;
309 if (outputs.size() > 0)
310 {
311 ctx.m_tensors[outputs[0]] = std::get<0>(out);
312 //cerr << ctx.m_tensors[outputs[0]] << endl;
313 if (outputs.size() > 1)
314 {
315 ctx.m_tensors[outputs[1]] = std::get<0>(std::get<1>(out));
316 if (outputs.size() > 2)
317 {
318 ctx.m_tensors[outputs[2]] = std::get<1>(std::get<1>(out));
319 if (outputs.size() > 3)
320 {
321 throw std::runtime_error("Error in static graph");
322 }
323 }
324 }
325 }
326 };
327 }
328 break;
329
330 case module_type_t::linear:
331 {
332 torch::nn::Linear& m = m_linear[module_ref.m_idx];
333 vector<size_t>& inputs = op.m_inputs;
334 vector<size_t>& outputs = op.m_outputs;
335
336 if (inputs.size() != 1 || outputs.size() != 1)
337 {
338 throw std::runtime_error("Error in static graph");
339 }
340
341 op.m_fn = [&m, inputs, outputs](context_t& ctx)
342 {
343 auto out = m->forward(ctx.m_tensors[inputs[0]]);
344 ctx.m_tensors[outputs[0]] = out;
345 };
346 }
347 break;
348
349 case module_type_t::dropout:
350 {
351 torch::nn::Dropout& m = m_dropout[module_ref.m_idx];
352 vector<size_t>& inputs = op.m_inputs;
353 vector<size_t>& outputs = op.m_outputs;
354
355 if (inputs.size() != 1 || outputs.size() != 1)
356 {
357 throw std::runtime_error("Error in static graph");
358 }
359
360 op.m_fn = [&m, inputs, outputs](context_t& ctx)
361 {
362 auto out = m->forward(ctx.m_tensors[inputs[0]]);
363 //cerr << ctx.m_tensors[inputs[0]] << endl;
364 //cerr << out << endl;
365 ctx.m_tensors[outputs[0]] = out;
366 };
367 }
368 break;
369
370 case module_type_t::deep_biaffine_attention_decoder:
371 {
372 torch_modules::DeepBiaffineAttentionDecoder& m
374 vector<size_t>& inputs = op.m_inputs;
375 vector<size_t>& outputs = op.m_outputs;
376
377 if (inputs.size() != 1 || outputs.size() != 1)
378 {
379 throw std::runtime_error("Error in static graph");
380 }
381
382 op.m_fn = [&m, inputs, outputs](context_t& ctx)
383 {
384 auto out = m->forward(ctx.m_tensors[inputs[0]]);
385 ctx.m_tensors[outputs[0]] = out;
386 };
387 }
388 break;
389
390 case module_type_t::deep_biaffine_attention_label_decoder:
391 {
392 torch_modules::DeepBiaffineAttentionLabelDecoder& m
394 vector<size_t>& inputs = op.m_inputs;
395 vector<size_t>& outputs = op.m_outputs;
396
397 if (inputs.size() != 1 || outputs.size() != 1)
398 {
399 throw std::runtime_error("Error in static graph");
400 }
401
402 op.m_fn = [&m, inputs, outputs](context_t& ctx)
403 {
404 auto out = m->forward(ctx.m_tensors[inputs[0]]);
405 ctx.m_tensors[outputs[0]] = out;
406 };
407 }
408 break;
409
410 default:
411 throw std::runtime_error("Error in static graph");
412 }
413 }
414 break;
415
416 case step_descr_t::cat:
417 {
418 vector<size_t>& inputs = op.m_inputs;
419 vector<size_t>& outputs = op.m_outputs;
420 int64_t dim = atoi(step.m_args["dim"].c_str());
421
422 if (inputs.size() == 0 || outputs.size() != 1)
423 {
424 throw std::runtime_error("Error in static graph");
425 }
426
427 op.m_fn = [inputs, outputs, dim](context_t& ctx)
428 {
429 std::vector<torch::Tensor> input_tensors;
430 input_tensors.resize(inputs.size());
431 for (size_t i = 0; i < inputs.size(); i++)
432 {
433 input_tensors[i] = ctx.m_tensors[inputs[i]];
434 //cerr << input_tensors[i].sizes() << std::endl;
435 }
436 auto out = torch::cat(input_tensors, dim);
437 ctx.m_tensors[outputs[0]] = out;
438 };
439 }
440 break;
441
442 case step_descr_t::reshape:
443 {
444 vector<size_t>& inputs = op.m_inputs;
445 vector<size_t>& outputs = op.m_outputs;
446 const auto shape = step.m_iargs["dims"];
447
448 if (inputs.size() == 0 || outputs.size() != 1)
449 {
450 throw std::runtime_error("Error in static graph");
451 }
452
453 op.m_fn = [inputs, outputs, shape](context_t& ctx)
454 {
455 //cerr << "reshape arg: " << ctx.m_tensors[inputs[0]].sizes() << endl;
456 //cerr << "shape="; for (auto v : shape) cerr << v << ", "; cerr << endl;
457 auto out = torch::reshape(ctx.m_tensors[inputs[0]], shape);
458 //cerr << "reshape res: " << out.sizes() << endl;
459 ctx.m_tensors[outputs[0]] = out;
460 };
461 }
462 break;
463
464 case step_descr_t::unbind:
465 {
466 vector<size_t>& inputs = op.m_inputs;
467 vector<size_t>& outputs = op.m_outputs;
468 int64_t dim = atoi(step.m_args["dim"].c_str());
469
470 if (inputs.size() != 1)
471 {
472 throw std::runtime_error("Error in static graph");
473 }
474
475 op.m_fn = [inputs, outputs, dim](context_t& ctx)
476 {
477 //cerr << "unbind arg: " << ctx.m_tensors[inputs[0]].sizes() << endl;
478 auto out = torch::unbind(ctx.m_tensors[inputs[0]], dim);
479 //cerr << "unbind res: " << out.sizes() << endl;
480 if (outputs.size() != out.size())
481 {
482 throw std::runtime_error("Error in static graph");
483 }
484 for (size_t j = 0; j < out.size(); ++j)
485 {
486 ctx.m_tensors[outputs[j]] = out[j];
487 }
488 };
489 }
490 break;
491
492 case step_descr_t::unsqueeze:
493 {
494 vector<size_t>& inputs = op.m_inputs;
495 vector<size_t>& outputs = op.m_outputs;
496 int64_t dim = atoi(step.m_args["dim"].c_str());
497
498 if (inputs.size() != 1 || outputs.size() != 1)
499 {
500 throw std::runtime_error("Error in static graph");
501 }
502
503 op.m_fn = [inputs, outputs, dim](context_t& ctx)
504 {
505 auto out = torch::unsqueeze(ctx.m_tensors[inputs[0]], dim);
506 ctx.m_tensors[outputs[0]] = out;
507 };
508 }
509 break;
510
511 case step_descr_t::log_softmax:
512 {
513 vector<size_t>& inputs = op.m_inputs;
514 vector<size_t>& outputs = op.m_outputs;
515 int64_t dim = 2;
516 if (step.m_iargs.end() != step.m_iargs.find("dim"))
517 {
518 dim = step.m_iargs.find("dim")->second[0];
519 }
520
521 if (inputs.size() != 1 || outputs.size() != 1)
522 {
523 throw std::runtime_error("Error in static graph");
524 }
525
526 op.m_fn = [inputs, outputs, dim](context_t& ctx)
527 {
528 auto in = ctx.m_tensors[inputs[0]];
529 ctx.m_tensors[outputs[0]] = torch::nn::functional::log_softmax(in, dim);
530 };
531 }
532 break;
533
534 case step_descr_t::sigmoid:
535 {
536 vector<size_t>& inputs = op.m_inputs;
537 vector<size_t>& outputs = op.m_outputs;
538
539 if (inputs.size() != 1 || outputs.size() != 1)
540 {
541 throw std::runtime_error("Error in static graph");
542 }
543
544 op.m_fn = [inputs, outputs](context_t& ctx)
545 {
546 ctx.m_tensors[outputs[0]] = ctx.m_tensors[inputs[0]].sigmoid();
547 };
548 }
549 break;
550
551 default:
552 throw std::runtime_error("Error in static graph");
553 }
554
555 m_ops.push_back(op);
556
557 for ( size_t out_idx : op.m_outputs )
558 {
559 if (m_outidx_to_opidx.end() != m_outidx_to_opidx.find(out_idx))
560 {
561 throw std::runtime_error("Error in static graph");
562 }
563 m_outidx_to_opidx[out_idx] = m_ops.size() - 1;
564 }
565 }
566 //cerr << "Done!" << endl;
567}
568
569map<string, string> StaticGraphImpl::parse_options(istringstream& ss)
570{
571 map<string, string> opts;
572
573 string s;
574 while (ss >> s)
575 {
576 string::size_type p = s.find('=');
577 if (string::npos == p)
578 {
579 throw std::runtime_error("Error in static graph");
580 }
581 string k = s.substr(0, p);
582 string v = s.substr(p+1);
583 if (opts.end() != opts.find(k))
584 {
585 throw std::runtime_error("Error in static graph");
586 }
587 opts[k] = v;
588 }
589
590 return opts;
591}
592
593vector<string> StaticGraphImpl::parse_list(string& str, char sep)
594{
595 vector<string> l;
596
597 string::size_type prev = 0;
598 string::size_type next = str.find(sep, prev);
599 while (next != string::npos)
600 {
601 l.push_back(str.substr(prev, next - prev));
602 prev = next + 1;
603 next = str.find(sep, prev);
604 }
605
606 l.push_back(str.substr(prev, next));
607
608 return l;
609}
610
611StaticGraphImpl::step_descr_t StaticGraphImpl::parse_script_line(const std::string& line)
612{
613 istringstream ss(line);
614 StaticGraphImpl::step_descr_t step;
615
616 string names_list;
617 ss >> names_list;
618 step.m_names = parse_list(names_list, ',');
619
620 string type;
621 ss >> type;
622
623 if (type == "=")
624 {
625 ss >> type;
626 }
627
628 if (type == "def")
629 {
630 if (step.m_names.size() != 1)
631 {
632 throw std::runtime_error("Error in static graph");
633 }
634 // Name def Class arg1=Value1 arg2=Value2 ...
635 step.m_type = step_descr_t::step_type_t::def;
636 string cls;
637
638 ss >> cls;
639
640 map<string, string> opts = parse_options(ss);
641
642 if (cls == "Embedding")
643 {
644 create_submodule_Embedding(names_list, opts);
645 }
646 else if (cls == "LSTM")
647 {
648 create_submodule_LSTM(names_list, opts);
649 }
650 else if (cls == "Linear")
651 {
652 create_submodule_Linear(names_list, opts);
653 }
654 else if (cls == "Dropout")
655 {
656 create_submodule_Dropout(names_list, opts);
657 }
658 else if (cls == "DeepBiaffineAttentionDecoder")
659 {
661 }
662 else if (cls == "DeepBiaffineAttentionLabelDecoder")
663 {
665 }
666 // else if (cls == "StanzaDepparseParser")
667 // {
668 // create_submodule_StanzaDepparseParser(names_list, opts);
669 // }
670 else if (cls == "Arg")
671 {
672 create_arg(step.m_names, opts);
673 }
674 else
675 {
676 throw std::runtime_error("Error in static graph");
677 }
678 }
679 else if (type == "cat")
680 {
681 // Name = cat Input1,Input2
682 step.m_type = step_descr_t::cat;
683 step.m_args = parse_options(ss);
684 }
685 else if (type == "forward")
686 {
687 //
688 step.m_type = step_descr_t::forward;
689 step.m_args = parse_options(ss);
690 }
691 else if (type == "log_softmax")
692 {
693 //
694 step.m_type = step_descr_t::log_softmax;
695 step.m_args = parse_options(ss);
696 auto it = step.m_args.find("dim");
697 if (step.m_args.end() != it)
698 {
699 step.m_iargs["dim"] = parse_iargs(it->second);
700 }
701 }
702 else if (type == "sigmoid")
703 {
704 //
705 step.m_type = step_descr_t::sigmoid;
706 step.m_args = parse_options(ss);
707 }
708 else if (type == "reshape")
709 {
710 //
711 step.m_type = step_descr_t::reshape;
712 step.m_args = parse_options(ss);
713 auto it = step.m_args.find("dims");
714 if (step.m_args.end() == it)
715 {
716 throw std::runtime_error("Error in static graph");
717 }
718 step.m_iargs["dims"] = parse_iargs(it->second);
719 assert(step.m_iargs["dims"].size() > 0);
720 }
721 else if (type == "unbind")
722 {
723 // Out1,Out2,... = unbind Input
724 step.m_type = step_descr_t::unbind;
725 step.m_args = parse_options(ss);
726 }
727 else if (type == "unsqueeze")
728 {
729 // Out1,Out2,... = unsqueeze Input
730 step.m_type = step_descr_t::unsqueeze;
731 step.m_args = parse_options(ss);
732 }
733 else
734 {
735 throw std::runtime_error("Error in static graph");
736 }
737
738 return step;
739}
740
741vector<int64_t> StaticGraphImpl::parse_iargs(const string& str)
742{
743 vector<int64_t> rv;
744
745 for (const string& s : deeplima::utils::split(str, ','))
746 rv.push_back(strtol(s.c_str(), nullptr, 10));
747
748 return rv;
749}
750
751void StaticGraphImpl::create_arg(const std::vector<std::string>& names, const std::map<std::string, std::string>& /*opts*/)
752{
753 for ( const string& name : names )
754 {
755 if (m_args.cend() != m_args.find(name))
756 {
757 throw std::runtime_error("Error in static graph");
758 }
759
760 if (m_tensor_name_to_idx.cend() != m_tensor_name_to_idx.find(name))
761 {
762 throw std::runtime_error("Error in static graph");
763 }
764
765 m_args.insert(name);
766 size_t idx = m_tensor_name_to_idx.size();
767 m_tensor_name_to_idx[name] = idx;
768 }
769}
770
771void StaticGraphImpl::create_submodule_Embedding(const std::string& name, const std::map<std::string, std::string>& opts)
772{
773 // Required options. It must throw an exception if they aren't available
774 int64_t dict_idx = get_option<int64_t>(opts, "dict");
775 int64_t dim = get_option<int64_t>(opts, "dim");
776
777 // A feature whose dict is empty for the corpus would otherwise create a
778 // zero-row embedding and crash on any lookup; give it at least one row.
779 int64_t num_embeddings = std::max<int64_t>(m_dicts[dict_idx]->size(), 1);
780 torch::nn::Embedding m(num_embeddings, dim);
781 m_embedding.push_back(m);
782 m_modules[name] = module_ref_t(module_type_t::embedding, m_embedding.size() - 1);
783 // m->pretty_print(cerr);
784 // cerr << endl;
785 // for ( const auto& t : m->parameters())
786 // {
787 // cerr << t.sizes() << endl;
788 //cerr << "itemsize = " << t.type().typeMeta().itemsize() << endl;
789 // }
790 // cerr << endl;
791
792 register_module(name, m);
793}
794
795void StaticGraphImpl::create_submodule_Dropout(const std::string& name, const std::map<std::string, std::string>& opts)
796{
797 // Required options. It must throw an exception if they aren't available
798 float prob = get_option<float>(opts, "prob");
799
800 torch::nn::Dropout m(prob);
801 m_dropout.push_back(m);
802 m_modules[name] = module_ref_t(module_type_t::dropout, m_dropout.size() - 1);
803
804 // m->pretty_print(cerr);
805 // cerr << endl;
806 // for ( const auto& t : m->parameters())
807 // {
808 // cerr << t.sizes() << endl;
809 // }
810 // cerr << endl;
811
812 register_module(name, m);
813}
814
815void StaticGraphImpl::create_submodule_Linear(const std::string& name, const std::map<std::string, std::string>& opts)
816{
817 // Required options. It must throw an exception if they aren't available
818 int64_t input_size = get_option<int64_t>(opts, "input_size");
819 int64_t output_size = get_option<int64_t>(opts, "output_size");
820
821 torch::nn::Linear m(input_size, output_size);
822 m_linear.push_back(m);
823 m_modules[name] = module_ref_t(module_type_t::linear, m_linear.size() - 1);
824
825 // m->pretty_print(cerr);
826 // cerr << endl;
827 // for ( const auto& t : m->parameters())
828 // {
829 // cerr << t.sizes() << endl;
830 // }
831 // cerr << endl;
832
833 register_module(name, m);
834}
835
836void StaticGraphImpl::create_submodule_LSTM(const std::string& name, const std::map<std::string, std::string>& opts)
837{
838 int64_t input_size = get_option<int64_t>(opts, "input_size");
839 int64_t hidden_size = get_option<int64_t>(opts, "hidden_size");
840
841 torch::nn::LSTMOptions lstm_options(input_size, hidden_size);
842 std::set<std::string> consumed_options({ "input_size", "hidden_size" });
843 // input_size=6 hidden_size=4 num_layers=2 batch_first=true bidirectional=true dropout
844 for (const auto& kv: opts)
845 {
846 if (consumed_options.cend() != consumed_options.find(kv.first))
847 {
848 continue;
849 }
850 if (kv.first == "num_layers")
851 {
852 lstm_options.num_layers(get_option<int64_t>(opts, "num_layers"));
853 }
854 else if (kv.first == "batch_first")
855 {
856 lstm_options.batch_first(get_bool_option(opts, "batch_first"));
857 }
858 else if (kv.first == "bidirectional")
859 {
860 lstm_options.bidirectional(get_bool_option(opts, "bidirectional"));
861 }
862 else if (kv.first == "dropout")
863 {
864 lstm_options.dropout(get_option<double>(opts, "dropout"));
865 }
866 else
867 {
868 throw std::runtime_error("Error in static graph");
869 }
870 consumed_options.insert(kv.first);
871 }
872
873 if (consumed_options.size() != opts.size())
874 {
875 throw std::runtime_error("Error in static graph");
876 }
877
878 torch::nn::LSTM m(lstm_options);
879
880 m_lstm.push_back(m);
881 m_modules[name] = module_ref_t(module_type_t::lstm, m_lstm.size() - 1);
882
883 // m->pretty_print(cerr);
884 // cerr << endl;
885
886 register_module(name, m);
887}
888
889void StaticGraphImpl::create_submodule_DeepBiaffineAttentionDecoder(const string& name, const map<string, string>& opts)
890{
891 int64_t input_dim = get_option<int64_t>(opts, "input_dim");
892 int64_t hidden_arc_dim = get_option<int64_t>(opts, "hidden_arc_dim");
893 bool input_includes_root = get_bool_option(opts, "input_includes_root");
894
895 torch_modules::DeepBiaffineAttentionDecoder m(input_dim, hidden_arc_dim, input_includes_root);
897 m_modules[name] = module_ref_t(module_type_t::deep_biaffine_attention_decoder, m_deep_biaffine_attention_decoder.size() - 1);
898 register_module(name, m);
899}
900
901void StaticGraphImpl::create_submodule_DeepBiaffineAttentionLabelDecoder(const string& name, const map<string, string>& opts)
902{
903 int64_t input_dim = get_option<int64_t>(opts, "input_dim");
904 int64_t hidden_dim = get_option<int64_t>(opts, "hidden_dim");
905 int64_t num_labels = get_option<int64_t>(opts, "num_labels");
906 bool input_includes_root = get_bool_option(opts, "input_includes_root");
907
908 torch_modules::DeepBiaffineAttentionLabelDecoder m(input_dim, hidden_dim, num_labels, input_includes_root);
910 m_modules[name] = module_ref_t(module_type_t::deep_biaffine_attention_label_decoder, m_deep_biaffine_attention_label_decoder.size() - 1);
911 register_module(name, m);
912}
913
914// void StaticGraphImpl::create_submodule_StanzaDepparseParser(const string& name, const map<string, string>& opts)
915// {
916// int64_t word_emb_dim = get_option<int64_t>(opts, "word_emb_dim");
917// int64_t tag_emb_dim = get_option<int64_t>(opts, "tag_emb_dim");
918// int64_t hidden_dim = get_option<int64_t>(opts, "hidden_dim");
919// int64_t num_layers = get_option<int64_t>(opts, "num_layers");
920// int64_t deep_biaff_hidden_dim = get_option<int64_t>(opts, "deep_biaff_hidden_dim");
921// int64_t word_dropout = get_option<int64_t>(opts, "word_dropout");
922// float dropout = get_option<int64_t>(opts, "dropout");
923// float rec_dropout = get_option<int64_t>(opts, "rec_dropout");
924// bool linearize = get_option<int64_t>(opts, "linearize");
925// bool dist = get_option<int64_t>(opts, "dist");
926//
927// std::shared_ptr<std::map<std::string, std::vector<std::string>>> vocab;
928// std::shared_ptr<std::vector<std::vector<std::string>>> feats_vocabs;
929//
930// m_stanza_depparse_parser = std::make_shared<torch_modules::StanzaDepparseParser>(
931// word_emb_dim, tag_emb_dim, hidden_dim,
932// num_layers, dropout, rec_dropout,
933// vocab,
934// feats_vocabs,
935// deep_biaff_hidden_dim,
936// linearize,
937// dist,
938// word_dropout);
939// m_modules[name] = module_ref_t(module_type_t::deep_biaffine_attention_decoder, m_deep_biaffine_attention_decoder.size() - 1);
940// register_module(name, *m_stanza_depparse_parser);
941// }
942
944{
945 for (torch::nn::LSTM &m : m_lstm)
946 {
947 OrderedDict<string, Tensor> params = m->named_parameters();
948 OrderedDict<string, Tensor>::Iterator it = params.begin();
949 while (params.end() != it)
950 {
951 OrderedDict<string, Tensor>::Item& item = *it;
952 if (item.key().substr(0, 9) == "weight_hh")
953 {
954 Tensor& t = item.value();
955 assert(t.size(0) % 4 == 0);
956 int width = t.size(0) / 4;
957
958 for (int i = 0; i < 4; i++){
959 auto a = item.value().index({ Slice(width * i, width * (i+1)), Slice() });
960 torch::nn::init::orthogonal_(a);
961 }
962 }
963 it++;
964 }
965 }
966}
967
968} // namespace nets
969} // namespace deeplima
virtual std::vector< int64_t > parse_iargs(const std::string &str)
std::vector< torch::nn::LSTM > m_lstm
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)
virtual void create_submodule_LSTM(const std::string &name, const std::map< std::string, std::string > &opts)
virtual void set_tags(const std::map< std::string, std::string > &tags)
virtual std::vector< std::string > parse_list(std::string &str, char sep)
bool get_bool_option(const std::map< std::string, std::string > &opts, const std::string &name)
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 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 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
virtual void save(torch::serialize::OutputArchive &archive) const override
virtual void create_submodule_Linear(const std::string &name, const std::map< std::string, std::string > &opts)
virtual void parse_script(const std::string &script)
std::map< size_t, size_t > m_outidx_to_opidx
std::vector< std::string > split(const std::string &str, char delim)
STL namespace.