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

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

Inheritance diagram for deeplima::nets::BiRnnClassifierImpl:
deeplima::nets::StaticGraphImpl deeplima::graph_dp::train::BiRnnAndDeepBiaffineAttentionImpl deeplima::nets::BiRnnSeq2SeqImpl deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl deeplima::tagging::train::BiRnnClassifierForNerImpl deeplima::lemmatization::train::Seq2SeqLemmatizerImpl

Public Member Functions

 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
 
virtual void load (torch::serialize::InputArchive &archive)
 
virtual void save (torch::serialize::OutputArchive &archive) 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

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

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_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 39 of file birnn_seq_classifier.h.

Constructor & Destructor Documentation

◆ BiRnnClassifierImpl() [1/3]

deeplima::nets::BiRnnClassifierImpl::BiRnnClassifierImpl ( )
inline

Definition at line 43 of file birnn_seq_classifier.h.

◆ BiRnnClassifierImpl() [2/3]

deeplima::nets::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 
)
inline

Definition at line 45 of file birnn_seq_classifier.h.

◆ BiRnnClassifierImpl() [3/3]

deeplima::nets::BiRnnClassifierImpl::BiRnnClassifierImpl ( DictsHolder &&  dicts,
const std::vector< embd_descr_t > &  embd_descr,
const std::string &  script 
)
inline

Definition at line 57 of file birnn_seq_classifier.h.

Member Function Documentation

◆ evaluate() [1/2]

void deeplima::nets::BiRnnClassifierImpl::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) 
)

Definition at line 131 of file birnn_seq_classifier.cpp.

◆ evaluate() [2/2]

void deeplima::nets::BiRnnClassifierImpl::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 
)
protected

Definition at line 145 of file birnn_seq_classifier.cpp.

◆ generate_script()

string deeplima::nets::BiRnnClassifierImpl::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 
)
staticprotected

Definition at line 17 of file birnn_classifier_script_generator.cpp.

◆ get_embd_descr()

const std::vector< embd_descr_t > & deeplima::nets::BiRnnClassifierImpl::get_embd_descr ( ) const
inline

Definition at line 65 of file birnn_seq_classifier.h.

◆ load()

◆ predict() [1/3]

torch::Tensor deeplima::nets::BiRnnClassifierImpl::predict ( const std::string &  output_name,
const torch::Tensor &  input,
const torch::Device &  device = torch::Device(torch::kCPU) 
)

Definition at line 21 of file birnn_seq_classifier.cpp.

◆ predict() [2/3]

torch::Tensor deeplima::nets::BiRnnClassifierImpl::predict ( const std::vector< std::string > &  output_names,
const torch::Tensor &  input,
const torch::Device &  device = torch::Device(torch::kCPU) 
)

Definition at line 29 of file birnn_seq_classifier.cpp.

◆ predict() [3/3]

void deeplima::nets::BiRnnClassifierImpl::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) 
)

Definition at line 51 of file birnn_seq_classifier.cpp.

◆ save()

◆ split_input()

void deeplima::nets::BiRnnClassifierImpl::split_input ( const torch::Tensor &  src,
std::map< std::string, torch::Tensor > &  dst,
const torch::Device &  device 
)
protected

Definition at line 87 of file birnn_seq_classifier.cpp.

◆ train()

void deeplima::nets::BiRnnClassifierImpl::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) 
)

Definition at line 225 of file birnn_seq_classifier.cpp.

◆ train_batch() [1/2]

void deeplima::nets::BiRnnClassifierImpl::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

Definition at line 366 of file birnn_seq_classifier.cpp.

◆ train_batch() [2/2]

void deeplima::nets::BiRnnClassifierImpl::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 
)
protected

Definition at line 349 of file birnn_seq_classifier.cpp.

◆ train_epoch()

void deeplima::nets::BiRnnClassifierImpl::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 
)
protected

Definition at line 328 of file birnn_seq_classifier.cpp.

Member Data Documentation

◆ m_embd_descr

std::vector<embd_descr_t> deeplima::nets::BiRnnClassifierImpl::m_embd_descr
protected

Definition at line 154 of file birnn_seq_classifier.h.


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