6#ifndef DEEPLIMA_SRC_STATIC_GRAPH_STATIC_GRAPH_H
7#define DEEPLIMA_SRC_STATIC_GRAPH_STATIC_GRAPH_H
12#include <torch/torch.h>
13#include <torch/serialize/archive.h>
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;
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)
57 if (m_type >= max_step_type)
59 throw std::runtime_error(
"Error in static graph");
66 std::vector<torch::Tensor> m_tensors;
67 explicit context_t(
size_t num)
69 m_tensors.resize(num);
82 op_t(op_type_t type = unknown)
85 if (m_type >= max_op_type)
87 throw std::runtime_error(
"Error in static graph");
92 std::vector<size_t> m_outputs;
93 std::vector<size_t> m_inputs;
94 std::function<void(context_t &)> m_fn;
105 deep_biaffine_attention_decoder = 5,
106 deep_biaffine_attention_label_decoder = 6,
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
123 virtual bool has_tag(
const std::string &key)
const
128 virtual void to(torch::Device device,
bool non_blocking =
false)
override;
131 const std::map<std::string, torch::Tensor> &inputs,
132 const std::string &output_name)
134 std::vector<std::string> output_names = {output_name};
135 return forward(inputs, output_names.cbegin(), output_names.cend());
139 const std::map<std::string, torch::Tensor> &inputs,
140 const std::vector<std::string> &output_names)
142 return forward(inputs, output_names.cbegin(), output_names.cend());
145 template <
class InputIt>
147 const std::map<std::string, torch::Tensor> &inputs,
148 InputIt outputs_begin,
151 assert(inputs.size() > 0);
152 assert(outputs_begin != outputs_end);
162 for (
const auto &kv : inputs)
166 throw std::runtime_error(
"Error in static graph");
169 if (idx >= ctx.m_tensors.size())
171 throw std::runtime_error(
"Error in static graph");
173 ctx.m_tensors[idx] = kv.second;
177 for (
size_t i = 0; i < plan.size(); i++)
179 size_t op_idx = plan[i];
182 m_ops[op_idx].m_fn(ctx);
190 std::map<std::string, torch::Tensor> rv;
191 for (
auto it = outputs_begin; it != outputs_end; ++it)
193 const std::string &str = *it;
196 throw std::runtime_error(
"Error in static graph");
199 if (idx >= ctx.m_tensors.size())
201 throw std::runtime_error(
"Error in static graph");
203 rv[str] = ctx.m_tensors[idx];
211 virtual void pretty_dump(std::ostream &stream)
const;
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);
227 template <
class InputIt>
229 const std::map<std::string, torch::Tensor> &inputs,
230 InputIt outputs_begin,
233 assert(inputs.size() > 0);
234 assert(outputs_begin != outputs_end);
239 throw std::runtime_error(
"Error in static graph");
242 std::set<size_t> given_inputs_idx;
243 for (
const auto &kv : inputs)
246 given_inputs_idx.insert(in_idx);
249 std::vector<size_t> temp;
250 temp.reserve(
m_ops.size() * 2);
252 std::list<size_t> required;
253 for (
auto it = outputs_begin; it != outputs_end; ++it)
255 const std::string &tensor_name = *it;
257 required.push_back(idx);
260 while (!required.empty())
262 size_t out_idx = required.front();
263 required.pop_front();
265 temp.push_back(op_idx);
267 const op_t &op =
m_ops[op_idx];
268 for (
size_t in_idx : op.m_inputs)
270 if (given_inputs_idx.end() == given_inputs_idx.find(in_idx))
272 required.push_back(in_idx);
278 std::set<size_t> planned_ops;
279 std::vector<size_t> plan;
280 plan.reserve(temp.size());
282 for (int32_t i = temp.size() - 1; i >= 0; i--)
284 if (planned_ops.end() != planned_ops.find(temp[i]))
288 plan.push_back(temp[i]);
289 planned_ops.insert(temp[i]);
297 template <
class InputIt>
299 const std::map<std::string, torch::Tensor> &inputs,
300 InputIt outputs_begin,
304 for (
const auto &kv : inputs)
315 for (
auto it = outputs_begin; it != outputs_end; ++it)
324 T
get_option(
const std::map<std::string, std::string> &opts,
const std::string &name)
326 const auto it = opts.find(name);
327 if (opts.cend() == it)
329 throw std::runtime_error(
"Error in static graph");
332 std::istringstream ss(it->second);
337 throw std::runtime_error(
"Error in static graph");
343 bool get_bool_option(
const std::map<std::string, std::string> &opts,
const std::string &name)
345 const auto it = opts.find(name);
346 if (opts.cend() == it)
348 throw std::runtime_error(
"Error in static graph");
351 if (it->second ==
"true")
355 else if (it->second ==
"false")
359 throw std::runtime_error(
"Error in static graph");
362 virtual void create_arg(
const std::vector<std::string> &names,
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);
374 std::map<std::string, std::string>
m_tags;
431 throw std::runtime_error(
"Unknown module name");
434 assert(module_type_t::embedding == mr.
m_type);
440 module_type_t t = unknown;
445 else if (type ==
"embedding")
449 else if (type ==
"linear")
453 else if (type ==
"dropout")
457 else if (type ==
"deep_biaffine_attention_decoder")
459 t = deep_biaffine_attention_decoder;
467 throw std::runtime_error(
"Unknown module type");
472 if (idx == p.second.m_idx && t == p.second.m_type)
478 return std::string(
"");
489 explicit module_ref_t(module_type_t type = module_type_t::unknown,
size_t idx = 0)
492 if (
m_type >= module_type_t::max_module_type)
494 throw std::runtime_error(
"Error in static graph");
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
std::vector< op_t > m_ops
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)