LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl Class Reference

#include </home/runner/work/lima/lima/deeplima/libs/tasks/segmentation/model/birnn_classifier_for_segmentation.h>

Inheritance diagram for deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl:
deeplima::nets::BiRnnClassifierImpl deeplima::nets::StaticGraphImpl

Public Types

typedef torch::Tensor tensor_t
 
typedef DictsHolder dicts_holder_t
 

Public Member Functions

 BiRnnClassifierForSegmentationImpl ()
 
 BiRnnClassifierForSegmentationImpl (DictsHolder &&dicts, const std::vector< impl::ngram_descr_t > &ngram_descr, const std::vector< nets::embd_descr_t > &embd_descr, const std::vector< nets::rnn_descr_t > &rnn_descr, const std::string &output_name, uint32_t num_classes, float input_dropout_prob)
 
virtual void load (torch::serialize::InputArchive &archive)
 
virtual void save (torch::serialize::OutputArchive &archive) const
 
void load (const std::string &fn)
 
size_t init_new_worker (size_t)
 
const std::vector< impl::ngram_descr_t > & get_ngram_descr () const
 
- 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 Attributes

std::vector< impl::ngram_descr_t > m_ngram_descr
 
size_t m_workers
 
- 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
 

Additional Inherited Members

- 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 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)
 

Detailed Description

Definition at line 19 of file birnn_classifier_for_segmentation.h.

Member Typedef Documentation

◆ dicts_holder_t

◆ tensor_t

Constructor & Destructor Documentation

◆ BiRnnClassifierForSegmentationImpl() [1/2]

deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl::BiRnnClassifierForSegmentationImpl ( )
inline

Definition at line 25 of file birnn_classifier_for_segmentation.h.

◆ BiRnnClassifierForSegmentationImpl() [2/2]

deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl::BiRnnClassifierForSegmentationImpl ( DictsHolder &&  dicts,
const std::vector< impl::ngram_descr_t > &  ngram_descr,
const std::vector< nets::embd_descr_t > &  embd_descr,
const std::vector< nets::rnn_descr_t > &  rnn_descr,
const std::string &  output_name,
uint32_t  num_classes,
float  input_dropout_prob 
)
inline

Definition at line 30 of file birnn_classifier_for_segmentation.h.

Member Function Documentation

◆ get_ngram_descr()

const std::vector< impl::ngram_descr_t > & deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl::get_ngram_descr ( ) const
inline

Definition at line 61 of file birnn_classifier_for_segmentation.h.

◆ init_new_worker()

size_t deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl::init_new_worker ( size_t  )
inline

Definition at line 56 of file birnn_classifier_for_segmentation.h.

◆ load() [1/2]

void deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl::load ( const std::string &  fn)
inline

Definition at line 51 of file birnn_classifier_for_segmentation.h.

◆ load() [2/2]

virtual void deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl::load ( torch::serialize::InputArchive &  archive)
virtual

◆ save()

void deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl::save ( torch::serialize::OutputArchive &  archive) const
virtual

Reimplemented from deeplima::nets::BiRnnClassifierImpl.

Definition at line 50 of file birnn_classifier_for_segmentation.cpp.

Member Data Documentation

◆ m_ngram_descr

std::vector<impl::ngram_descr_t> deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl::m_ngram_descr
protected

Definition at line 67 of file birnn_classifier_for_segmentation.h.

◆ m_workers

size_t deeplima::segmentation::train::BiRnnClassifierForSegmentationImpl::m_workers
protected

Definition at line 68 of file birnn_classifier_for_segmentation.h.


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