InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
infini::ops::Embedding Class Referenceabstract

#include <embedding.h>

Inheritance diagram for infini::ops::Embedding:
infini::ops::Operator< Embedding > infini::ops::OperatorBase infini::ops::OperatorBase

Public Member Functions

 Embedding (const Tensor input, const Tensor weight, const std::optional< int64_t > padding_idx, const std::optional< double > max_norm, const double norm_type, const bool scale_grad_by_freq, const bool sparse, Tensor out)
 
 Embedding (const Tensor input, const Tensor weight, Tensor out)
 
 Embedding (const Tensor input, const Tensor weight, const int64_t padding_idx, const bool scale_grad_by_freq, const bool sparse, Tensor out)
 
virtual void operator() (const Tensor input, const Tensor weight, const std::optional< int64_t > padding_idx, const std::optional< double > max_norm, const double norm_type, const bool scale_grad_by_freq, const bool sparse, Tensor out) const =0
 
void operator() (const Tensor input, const Tensor weight, Tensor out) const
 
void operator() (const Tensor input, const Tensor weight, const int64_t padding_idx, const bool scale_grad_by_freq, const bool sparse, Tensor out) const
 
- Public Member Functions inherited from infini::ops::Operator< Embedding >
void operator() (const Handle &handle, const Args &... args)
 
void operator() (const Args &... args) const
 
void operator() (const Handle &handle, const Args &... args)
 
void operator() (const Args &... args) const
 
- Public Member Functions inherited from infini::ops::OperatorBase
virtual ~OperatorBase ()=default
 
virtual std::size_t workspace_size_in_bytes () const
 
void set_handle (const Handle &handle)
 
void set_config (const Config &config)
 
void set_stream (void *stream)
 
void set_workspace (void *workspace)
 
void set_workspace_size_in_bytes (std::size_t workspace_size_in_bytes)
 
virtual ~OperatorBase ()=default
 
virtual std::size_t workspace_size_in_bytes () const
 
void set_handle (const Handle &handle)
 
void set_config (const Config &config)
 
void set_stream (void *stream)
 
void set_workspace (void *workspace)
 
void set_workspace_size_in_bytes (std::size_t workspace_size_in_bytes)
 

Static Public Member Functions

template<typename TensorLike >
static auto MakeReturnValue (const TensorLike &input, const TensorLike &weight, const std::optional< int64_t >=std::nullopt, const std::optional< double >=std::nullopt, const double=2.0, const bool=false, const bool=false)
 
- Static Public Member Functions inherited from infini::ops::Operator< Embedding >
static void clear_cache ()
 
static void clear_cache ()
 
static std::unique_ptr< Operator > Make (const Config &config, const Tensor tensor, Args &&... args)
 
static std::unique_ptr< Operator > Make (const Tensor tensor, Args &&... args)
 
static std::unique_ptr< Operator > Make (const Config &config, const std::vector< Tensor > tensors, Args &&... args)
 
static std::unique_ptr< Operator > Make (const std::vector< Tensor > tensors, Args &&... args)
 
static std::unique_ptr< Operator > Make (const Config &config, const Tensor tensor, Args &&... args)
 
static std::unique_ptr< Operator > Make (const Tensor tensor, Args &&... args)
 
static std::unique_ptr< Operator > Make (const Config &config, const std::vector< Tensor > tensors, Args &&... args)
 
static std::unique_ptr< Operator > Make (const std::vector< Tensor > tensors, Args &&... args)
 
static void Call (const Handle &handle, const Config &config, const Args &... args)
 
static void Call (const Tensor tensor, const Args &... args)
 
static auto Call (const TensorLike &tensor, const Args &... args)
 
static void Call (const Handle &handle, const Config &config, const Args &... args)
 
static void Call (const Tensor tensor, const Args &... args)
 
static auto Call (const TensorLike &tensor, const Args &... args)
 
static std::vector< std::size_t > active_implementation_indices (Device::Type dev_type)
 
static std::vector< std::size_t > active_implementation_indices (Device::Type dev_type)
 

Static Protected Member Functions

static Tensor::Size NumIndices (const Tensor::Shape &input_shape)
 

Protected Attributes

Tensor::Shape input_shape_
 
Tensor::Shape weight_shape_
 
Tensor::Shape out_shape_
 
Tensor::Strides input_strides_
 
Tensor::Strides weight_strides_
 
Tensor::Strides out_strides_
 
DataType input_dtype_
 
DataType weight_dtype_
 
DataType out_dtype_
 
Tensor::Size num_indices_ {0}
 
Tensor::Size vocab_size_ {0}
 
Tensor::Size embedding_dim_ {0}
 
std::optional< int64_t > padding_idx_ {}
 
std::optional< double > max_norm_ {}
 
double norm_type_ {2.0}
 
bool scale_grad_by_freq_ {false}
 
bool sparse_ {false}
 
- Protected Attributes inherited from infini::ops::OperatorBase
std::unique_ptr< Handle > handle_ptr_
 
std::unique_ptr< Config > config_ptr_
 
void * stream_ {nullptr}
 
void * workspace_ {nullptr}
 
std::size_t workspace_size_in_bytes_ {0}
 

Additional Inherited Members

- Static Protected Attributes inherited from infini::ops::Operator< Embedding >
static constexpr Device::Type device_type_
 
static constexpr std::size_t implementation_index_
 

Constructor & Destructor Documentation

◆ Embedding() [1/3]

infini::ops::Embedding::Embedding ( const Tensor  input,
const Tensor  weight,
const std::optional< int64_t >  padding_idx,
const std::optional< double >  max_norm,
const double  norm_type,
const bool  scale_grad_by_freq,
const bool  sparse,
Tensor  out 
)
inline

◆ Embedding() [2/3]

infini::ops::Embedding::Embedding ( const Tensor  input,
const Tensor  weight,
Tensor  out 
)
inline

◆ Embedding() [3/3]

infini::ops::Embedding::Embedding ( const Tensor  input,
const Tensor  weight,
const int64_t  padding_idx,
const bool  scale_grad_by_freq,
const bool  sparse,
Tensor  out 
)
inline
Deprecated:
Use the overload that also accepts max_norm and norm_type instead.

Member Function Documentation

◆ MakeReturnValue()

template<typename TensorLike >
static auto infini::ops::Embedding::MakeReturnValue ( const TensorLike &  input,
const TensorLike &  weight,
const std::optional< int64_t >  = std::nullopt,
const std::optional< double >  = std::nullopt,
const double  = 2.0,
const bool  = false,
const bool  = false 
)
inlinestatic

◆ NumIndices()

static Tensor::Size infini::ops::Embedding::NumIndices ( const Tensor::Shape &  input_shape)
inlinestaticprotected

◆ operator()() [1/3]

void infini::ops::Embedding::operator() ( const Tensor  input,
const Tensor  weight,
const int64_t  padding_idx,
const bool  scale_grad_by_freq,
const bool  sparse,
Tensor  out 
) const
inline
Deprecated:
Use the overload that also accepts max_norm and norm_type instead.

◆ operator()() [2/3]

virtual void infini::ops::Embedding::operator() ( const Tensor  input,
const Tensor  weight,
const std::optional< int64_t >  padding_idx,
const std::optional< double >  max_norm,
const double  norm_type,
const bool  scale_grad_by_freq,
const bool  sparse,
Tensor  out 
) const
pure virtual

◆ operator()() [3/3]

void infini::ops::Embedding::operator() ( const Tensor  input,
const Tensor  weight,
Tensor  out 
) const
inline

Member Data Documentation

◆ embedding_dim_

Tensor::Size infini::ops::Embedding::embedding_dim_ {0}
protected

◆ input_dtype_

DataType infini::ops::Embedding::input_dtype_
protected

◆ input_shape_

Tensor::Shape infini::ops::Embedding::input_shape_
protected

◆ input_strides_

Tensor::Strides infini::ops::Embedding::input_strides_
protected

◆ max_norm_

std::optional<double> infini::ops::Embedding::max_norm_ {}
protected

◆ norm_type_

double infini::ops::Embedding::norm_type_ {2.0}
protected

◆ num_indices_

Tensor::Size infini::ops::Embedding::num_indices_ {0}
protected

◆ out_dtype_

DataType infini::ops::Embedding::out_dtype_
protected

◆ out_shape_

Tensor::Shape infini::ops::Embedding::out_shape_
protected

◆ out_strides_

Tensor::Strides infini::ops::Embedding::out_strides_
protected

◆ padding_idx_

std::optional<int64_t> infini::ops::Embedding::padding_idx_ {}
protected

◆ scale_grad_by_freq_

bool infini::ops::Embedding::scale_grad_by_freq_ {false}
protected

◆ sparse_

bool infini::ops::Embedding::sparse_ {false}
protected

◆ vocab_size_

Tensor::Size infini::ops::Embedding::vocab_size_ {0}
protected

◆ weight_dtype_

DataType infini::ops::Embedding::weight_dtype_
protected

◆ weight_shape_

Tensor::Shape infini::ops::Embedding::weight_shape_
protected

◆ weight_strides_

Tensor::Strides infini::ops::Embedding::weight_strides_
protected

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