|
| | MoeWna16MarlinGemm (const Tensor a, const Tensor b_q_weight, std::optional< Tensor > b_bias_or_none, const Tensor b_scales, std::optional< Tensor > a_scales, std::optional< Tensor > global_scale, std::optional< Tensor > b_zeros_or_none, std::optional< Tensor > g_idx_or_none, std::optional< Tensor > perm_or_none, const Tensor workspace, const Tensor sorted_token_ids, const Tensor expert_ids, const Tensor num_tokens_past_padded, const Tensor topk_weights, const int64_t moe_block_size, const int64_t top_k, const bool mul_topk_weights, const int64_t b_type_id, const int64_t size_m, const int64_t size_n, const int64_t size_k, const bool is_full_k, const bool use_atomic_add, const bool use_fp32_reduce, const bool is_zp_float, const int64_t thread_k, const int64_t thread_n, const int64_t blocks_per_sm, Tensor out) |
| |
| virtual void | operator() (const Tensor a, const Tensor b_q_weight, std::optional< Tensor > b_bias_or_none, const Tensor b_scales, std::optional< Tensor > a_scales, std::optional< Tensor > global_scale, std::optional< Tensor > b_zeros_or_none, std::optional< Tensor > g_idx_or_none, std::optional< Tensor > perm_or_none, const Tensor workspace, const Tensor sorted_token_ids, const Tensor expert_ids, const Tensor num_tokens_past_padded, const Tensor topk_weights, const int64_t moe_block_size, const int64_t top_k, const bool mul_topk_weights, const int64_t b_type_id, const int64_t size_m, const int64_t size_n, const int64_t size_k, const bool is_full_k, const bool use_atomic_add, const bool use_fp32_reduce, const bool is_zp_float, const int64_t thread_k, const int64_t thread_n, const int64_t blocks_per_sm, Tensor out) const =0 |
| |
| 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 |
| |
| 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) |
| |
|
| void | ValidateCallMetadata (const Tensor a, const Tensor b_q_weight, std::optional< Tensor > b_bias_or_none, const Tensor b_scales, std::optional< Tensor > a_scales, std::optional< Tensor > global_scale, std::optional< Tensor > b_zeros_or_none, std::optional< Tensor > g_idx_or_none, std::optional< Tensor > perm_or_none, const Tensor workspace, const Tensor sorted_token_ids, const Tensor expert_ids, const Tensor num_tokens_past_padded, const Tensor topk_weights, const int64_t moe_block_size, const int64_t top_k, const bool mul_topk_weights, const int64_t b_type_id, const int64_t size_m, const int64_t size_n, const int64_t size_k, const bool is_full_k, const bool use_atomic_add, const bool use_fp32_reduce, const bool is_zp_float, const int64_t thread_k, const int64_t thread_n, const int64_t blocks_per_sm, Tensor out) const |
| |
|
| 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 constexpr Device::Type | device_type_ |
| |
| static constexpr std::size_t | implementation_index_ |
| |