LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
deeplima::nets::BiRnnSeq2SeqImpl Class Reference

#include </home/runner/work/lima/lima/deeplima/libs/nn/birnn_seq2seq/birnn_seq2seq.h>

Inheritance diagram for deeplima::nets::BiRnnSeq2SeqImpl:
deeplima::nets::BiRnnClassifierImpl deeplima::nets::StaticGraphImpl deeplima::lemmatization::train::Seq2SeqLemmatizerImpl

Public Types

typedef torch::Tensor tensor_t
 

Public Member Functions

 BiRnnSeq2SeqImpl ()
 
 BiRnnSeq2SeqImpl (DictsHolder &&dicts, const std::vector< embd_descr_t > &encoder_embd_descr, const std::vector< rnn_descr_t > &encoder_rnn_descr, const std::vector< embd_descr_t > &decoder_embd_descr, const std::vector< rnn_descr_t > &decoder_rnn_descr, const std::vector< embd_descr_t > &cat_embd_descr)
 
virtual void load (torch::serialize::InputArchive &archive)
 
virtual void save (torch::serialize::OutputArchive &archive) const
 
void evaluate (const std::vector< std::string > &output_names, const TorchMatrix< int64_t > &input, const TorchMatrix< int64_t > &gold, epoch_stat_t &stat, const torch::Device &device=torch::Device(torch::kCPU))
 
- Public Member Functions inherited from deeplima::nets::BiRnnClassifierImpl
 BiRnnClassifierImpl ()
 
 BiRnnClassifierImpl (DictsHolder &&dicts, const std::vector< embd_descr_t > &embd_descr, const std::vector< rnn_descr_t > &rnn_descr, const std::vector< std::string > &output_names, const std::vector< uint32_t > &classes, float input_dropout_prob)
 
 BiRnnClassifierImpl (DictsHolder &&dicts, const std::vector< embd_descr_t > &embd_descr, const std::string &script)
 
const std::vector< embd_descr_t > & get_embd_descr () const
 
void train (size_t epochs, size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, const TorchMatrix< int64_t > &train_input, const TorchMatrix< int64_t > &train_gold, const TorchMatrix< int64_t > &eval_input, const TorchMatrix< int64_t > &eval_gold, torch::optim::Optimizer &opt, const std::string &model_name="", const torch::Device &device=torch::Device(torch::kCPU))
 
void evaluate (const std::vector< std::string > &output_name, const TorchMatrix< int64_t > &input, const TorchMatrix< int64_t > &gold, epoch_stat_t &stat, const torch::Device &device=torch::Device(torch::kCPU))
 
torch::Tensor predict (const std::string &output_name, const torch::Tensor &input, const torch::Device &device=torch::Device(torch::kCPU))
 
torch::Tensor predict (const std::vector< std::string > &output_names, const torch::Tensor &input, const torch::Device &device=torch::Device(torch::kCPU))
 
void predict (size_t worker_id, const torch::Tensor &inputs, int64_t input_begin, int64_t input_end, int64_t output_begin, int64_t output_end, std::shared_ptr< StdMatrix< uint8_t > > &output, const std::vector< std::string > &outputs_names, const torch::Device &device=torch::Device(torch::kCPU))
 
- Public Member Functions inherited from deeplima::nets::StaticGraphImpl
 StaticGraphImpl (const DictsHolder &dicts, const std::string &script)
 
 StaticGraphImpl (DictsHolder &&dicts, const std::string &script)
 
 StaticGraphImpl ()
 
virtual void set_tags (const std::map< std::string, std::string > &tags)
 
virtual const std::map< std::string, std::string > & get_tags () const
 
virtual bool has_tag (const std::string &key) const
 
virtual void to (torch::Device device, bool non_blocking=false) override
 
std::map< std::string, torch::Tensor > forward (const std::map< std::string, torch::Tensor > &inputs, const std::string &output_name)
 
std::map< std::string, torch::Tensor > forward (const std::map< std::string, torch::Tensor > &inputs, const std::vector< std::string > &output_names)
 
template<class InputIt >
std::map< std::string, torch::Tensor > forward (const std::map< std::string, torch::Tensor > &inputs, InputIt outputs_begin, InputIt outputs_end)
 
virtual void pretty_dump (std::ostream &stream) const
 
const DictsHolder & get_dicts () const
 
const std::string & get_script () const
 
const std::vector< torch::nn::LSTM > & get_layers_lstm () const
 
const std::vector< torch::nn::Linear > & get_layers_linear () const
 
const std::vector< torch::nn::Embedding > & get_layers_embedding () const
 
const std::vector< torch::nn::Dropout > & get_layers_dropout () const
 
const std::vector< deeplima::nets::torch_modules::DeepBiaffineAttentionDecoder > & get_layers_deep_biaffine_attn_decoder () const
 
const std::vector< deeplima::nets::torch_modules::DeepBiaffineAttentionLabelDecoder > & get_layers_deep_biaffine_attn_label_decoder () const
 
torch::nn::Embedding get_module_by_name (const std::string &name) const
 
const std::string get_module_name (size_t idx, const std::string &type) const
 

Protected Member Functions

std::string generate_script (const std::vector< embd_descr_t > &encoder_embd_descr, const std::vector< rnn_descr_t > &encoder_rnn_descr, const std::vector< embd_descr_t > &decoder_embd_descr, const std::vector< rnn_descr_t > &decoder_rnn_descr, const std::vector< embd_descr_t > &cat_embd_descr, size_t n_output_classes)
 
void train_batch (const std::vector< std::string > &output_names, const torch::Tensor &input, const std::vector< torch::Tensor > &input_cat, const torch::Tensor &target, torch::optim::Optimizer &opt, epoch_stat_t &stat, const torch::Device &device)
 
- Protected Member Functions inherited from deeplima::nets::BiRnnClassifierImpl
void split_input (const torch::Tensor &src, std::map< std::string, torch::Tensor > &dst, const torch::Device &device)
 
void evaluate (const std::vector< std::string > &output_names, const std::map< std::string, torch::Tensor > &input, const torch::Tensor &target, epoch_stat_t &stat, const torch::Device &device)
 
void train_epoch (size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, const torch::Tensor &input_batches, const torch::Tensor &gold_batches, torch::optim::Optimizer &opt, epoch_stat_t &stat, const torch::Device &device)
 
void train_batch (size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, const torch::Tensor &input, const torch::Tensor &gold, torch::optim::Optimizer &opt, epoch_stat_t &stat, const torch::Device &device)
 
void train_batch (size_t batch_size, size_t seq_len, const std::vector< std::string > &output_names, const std::map< std::string, torch::Tensor > &input, const torch::Tensor &target, torch::optim::Optimizer &opt, epoch_stat_t &stat, const torch::Device &device)
 
- Protected Member Functions inherited from deeplima::nets::StaticGraphImpl
virtual void parse_script (const std::string &script)
 
virtual step_descr_t parse_script_line (const std::string &line)
 
virtual std::map< std::string, std::string > parse_options (std::istringstream &ss)
 
virtual std::vector< std::string > parse_list (std::string &str, char sep)
 
virtual std::vector< int64_t > parse_iargs (const std::string &str)
 
template<class InputIt >
void prepare_exec_plan (const std::map< std::string, torch::Tensor > &inputs, InputIt outputs_begin, InputIt outputs_end)
 
template<class InputIt >
std::string get_exec_plan_key (const std::map< std::string, torch::Tensor > &inputs, InputIt outputs_begin, InputIt outputs_end)
 
template<class T >
T get_option (const std::map< std::string, std::string > &opts, const std::string &name)
 
bool get_bool_option (const std::map< std::string, std::string > &opts, const std::string &name)
 
virtual void create_arg (const std::vector< std::string > &names, const std::map< std::string, std::string > &opts)
 
virtual void create_submodule_Embedding (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 create_submodule_Linear (const std::string &name, const std::map< std::string, std::string > &opts)
 
virtual void create_submodule_Dropout (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)
 
virtual void create_submodule_DeepBiaffineAttentionLabelDecoder (const std::string &name, const std::map< std::string, std::string > &opts)
 
virtual void init_rnns ()
 

Static Protected Member Functions

template<class T >
static std::vector< T > concatenate (const std::vector< T > &a, const std::vector< T > &b, const std::vector< T > &c={})
 
- Static Protected Member Functions inherited from deeplima::nets::BiRnnClassifierImpl
static std::string generate_script (const std::vector< embd_descr_t > &embd_descr, const std::vector< rnn_descr_t > &rnn_descr, const std::vector< std::string > &output_names, const std::vector< uint32_t > &classes, float input_dropout_prob)
 

Protected Attributes

std::vector< embd_descr_t > m_cat_embd_descr
 
- Protected Attributes inherited from deeplima::nets::BiRnnClassifierImpl
std::vector< embd_descr_t > m_embd_descr
 
- Protected Attributes inherited from deeplima::nets::StaticGraphImpl
std::map< std::string, std::vector< size_t > > m_exec_plans
 
DictsHolder m_dicts
 
std::string m_script
 
std::map< std::string, std::string > m_tags
 
std::vector< torch::nn::Embedding > m_embedding
 
std::vector< torch::nn::LSTM > m_lstm
 
std::vector< torch::nn::Linear > m_linear
 
std::vector< torch::nn::Dropout > m_dropout
 
std::vector< deeplima::nets::torch_modules::DeepBiaffineAttentionDecoder > m_deep_biaffine_attention_decoder
 
std::vector< deeplima::nets::torch_modules::DeepBiaffineAttentionLabelDecoder > m_deep_biaffine_attention_label_decoder
 
std::map< std::string, module_ref_t > m_modules
 
std::set< std::string > m_args
 
std::map< std::string, size_t > m_tensor_name_to_idx
 
std::vector< op_t > m_ops
 
std::map< size_t, size_t > m_outidx_to_opidx
 

Detailed Description

Definition at line 16 of file birnn_seq2seq.h.

Member Typedef Documentation

◆ tensor_t

Definition at line 19 of file birnn_seq2seq.h.

Constructor & Destructor Documentation

◆ BiRnnSeq2SeqImpl() [1/2]

deeplima::nets::BiRnnSeq2SeqImpl::BiRnnSeq2SeqImpl ( )
inline

Definition at line 21 of file birnn_seq2seq.h.

◆ BiRnnSeq2SeqImpl() [2/2]

deeplima::nets::BiRnnSeq2SeqImpl::BiRnnSeq2SeqImpl ( DictsHolder &&  dicts,
const std::vector< embd_descr_t > &  encoder_embd_descr,
const std::vector< rnn_descr_t > &  encoder_rnn_descr,
const std::vector< embd_descr_t > &  decoder_embd_descr,
const std::vector< rnn_descr_t > &  decoder_rnn_descr,
const std::vector< embd_descr_t > &  cat_embd_descr 
)
inline

Definition at line 23 of file birnn_seq2seq.h.

Member Function Documentation

◆ concatenate()

template<class T >
static std::vector< T > deeplima::nets::BiRnnSeq2SeqImpl::concatenate ( const std::vector< T > &  a,
const std::vector< T > &  b,
const std::vector< T > &  c = {} 
)
inlinestaticprotected

Definition at line 63 of file birnn_seq2seq.h.

◆ evaluate()

void deeplima::nets::BiRnnSeq2SeqImpl::evaluate ( const std::vector< std::string > &  output_names,
const TorchMatrix< int64_t > &  input,
const TorchMatrix< int64_t > &  gold,
epoch_stat_t &  stat,
const torch::Device &  device = torch::Device(torch::kCPU) 
)

Definition at line 77 of file birnn_seq2seq.cpp.

◆ generate_script()

string deeplima::nets::BiRnnSeq2SeqImpl::generate_script ( const std::vector< embd_descr_t > &  encoder_embd_descr,
const std::vector< rnn_descr_t > &  encoder_rnn_descr,
const std::vector< embd_descr_t > &  decoder_embd_descr,
const std::vector< rnn_descr_t > &  decoder_rnn_descr,
const std::vector< embd_descr_t > &  cat_embd_descr,
size_t  n_output_classes 
)
protected

Definition at line 19 of file birnn_seq2seq_script_generator.cpp.

◆ load()

void deeplima::nets::BiRnnSeq2SeqImpl::load ( torch::serialize::InputArchive &  archive)
virtual

◆ save()

void deeplima::nets::BiRnnSeq2SeqImpl::save ( torch::serialize::OutputArchive &  archive) const
virtual

◆ train_batch()

void deeplima::nets::BiRnnSeq2SeqImpl::train_batch ( const std::vector< std::string > &  output_names,
const torch::Tensor &  input,
const std::vector< torch::Tensor > &  input_cat,
const torch::Tensor &  target,
torch::optim::Optimizer &  opt,
epoch_stat_t &  stat,
const torch::Device &  device 
)
protected

Definition at line 19 of file birnn_seq2seq.cpp.

Member Data Documentation

◆ m_cat_embd_descr

std::vector<embd_descr_t> deeplima::nets::BiRnnSeq2SeqImpl::m_cat_embd_descr
protected

Definition at line 73 of file birnn_seq2seq.h.


The documentation for this class was generated from the following files: