Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions include/infinicore/ops.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,10 @@
#include "ops/flash_attention.hpp"
#include "ops/fmin.hpp"
#include "ops/fmod.hpp"
#include "ops/fp8_indexer_logits.hpp"
#include "ops/fp8_indexer_quant.hpp"
#include "ops/fp8_mla_rmsnorm_cache.hpp"
#include "ops/fp8_sparse_mla.hpp"
#include "ops/fused_gated_delta_net_gating.hpp"
#include "ops/fused_moe.hpp"
#include "ops/gelu.hpp"
Expand Down Expand Up @@ -68,6 +72,7 @@
#include "ops/rotmg.hpp"
#include "ops/rwkv5_wkv.hpp"
#include "ops/scal.hpp"
#include "ops/select_last_token_hidden.hpp"
#include "ops/sigmoid.hpp"
#include "ops/silu.hpp"
#include "ops/silu_and_mul.hpp"
Expand Down
22 changes: 22 additions & 0 deletions include/infinicore/ops/fp8_indexer_logits.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
#pragma once

#include "../graph/graph.hpp"
#include "common/op.hpp"

namespace infinicore::op {

INFINICORE_GRAPH_OP_CLASS(
Fp8IndexerLogits,
Tensor, const Tensor &, const Tensor &, const Tensor &, const Tensor &,
const Tensor &, const Tensor &);

void fp8_indexer_logits_(
Tensor logits,
const Tensor &q_fp8,
const Tensor &kv_cache,
const Tensor &block_tables,
const Tensor &weights_fp32,
const Tensor &positions,
const Tensor &request_ids);

} // namespace infinicore::op
39 changes: 39 additions & 0 deletions include/infinicore/ops/fp8_indexer_quant.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
#pragma once

#include "../graph/graph.hpp"
#include "common/op.hpp"

namespace infinicore::op {

INFINICORE_GRAPH_OP_CLASS(
Fp8IndexerQuant, Tensor, Tensor, const Tensor &, const Tensor &);

void fp8_indexer_quant_(
Tensor q_fp8,
Tensor weights_fp32,
const Tensor &q,
const Tensor &weights);

INFINICORE_GRAPH_OP_CLASS(
FusedFp8Indexer,
Tensor, Tensor, Tensor,
const Tensor &, const Tensor &, const Tensor &, const Tensor &,
const Tensor &, const Tensor &, const Tensor &,
size_t, double, double);

void fused_fp8_indexer_(
Tensor q_fp8,
Tensor weights_fp32,
Tensor k_cache,
const Tensor &q_raw,
const Tensor &k_weights,
const Tensor &norm_weight,
const Tensor &norm_bias,
const Tensor &positions,
const Tensor &cos_sin_cache,
const Tensor &slot_mapping,
size_t rope_dim,
double eps,
double weights_scale);

} // namespace infinicore::op
37 changes: 37 additions & 0 deletions include/infinicore/ops/fp8_mla_rmsnorm_cache.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
#pragma once

#include "../graph/graph.hpp"
#include "common/op.hpp"

namespace infinicore::op {

INFINICORE_GRAPH_OP_CLASS(
Fp8MlaRmsnormCache,
Tensor,
const Tensor &, const Tensor &, const Tensor &, const Tensor &,
double);

INFINICORE_GRAPH_OP_CLASS(
Fp8MlaRmsnormDualCache,
Tensor, Tensor,
const Tensor &, const Tensor &, const Tensor &, const Tensor &,
double);

void fp8_mla_rmsnorm_cache_(
Tensor cache,
const Tensor &compressed_kv,
const Tensor &norm_weight,
const Tensor &rope,
const Tensor &slot_mapping,
double eps);

void fp8_mla_rmsnorm_dual_cache_(
Tensor cache,
Tensor vendor_cache,
const Tensor &compressed_kv,
const Tensor &norm_weight,
const Tensor &rope,
const Tensor &slot_mapping,
double eps);

} // namespace infinicore::op
20 changes: 20 additions & 0 deletions include/infinicore/ops/fp8_sparse_mla.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,20 @@
#pragma once

#include "../graph/graph.hpp"
#include "common/op.hpp"

namespace infinicore::op {

INFINICORE_GRAPH_OP_CLASS(
Fp8SparseMla,
Tensor, const Tensor &, const Tensor &, const Tensor &, const Tensor &, float);

void fp8_sparse_mla_(
Tensor output,
const Tensor &query,
const Tensor &kv_cache,
const Tensor &indices,
const Tensor &topk_lens,
float scale);

} // namespace infinicore::op
13 changes: 13 additions & 0 deletions include/infinicore/ops/select_last_token_hidden.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,13 @@
#pragma once

#include "../device.hpp"
#include "../graph/graph.hpp"
#include "common/op.hpp"

namespace infinicore::op {

INFINICORE_GRAPH_OP_CLASS(SelectLastTokenHidden, Tensor, const Tensor &, const Tensor &);

void select_last_token_hidden_(Tensor output, const Tensor &hidden_states, const Tensor &input_offsets);

} // namespace infinicore::op
5 changes: 5 additions & 0 deletions include/infiniop.h
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,10 @@
#include "infiniop/ops/floor_divide.h"
#include "infiniop/ops/fmin.h"
#include "infiniop/ops/fmod.h"
#include "infiniop/ops/fp8_indexer_logits.h"
#include "infiniop/ops/fp8_indexer_quant.h"
#include "infiniop/ops/fp8_mla_rmsnorm_cache.h"
#include "infiniop/ops/fp8_sparse_mla.h"
#include "infiniop/ops/fused_gated_delta_net_gating.h"
#include "infiniop/ops/fused_moe.h"
#include "infiniop/ops/gelu.h"
Expand Down Expand Up @@ -127,6 +131,7 @@
#include "infiniop/ops/rwkv5_wkv.h"
#include "infiniop/ops/scal.h"
#include "infiniop/ops/scatter.h"
#include "infiniop/ops/select_last_token_hidden.h"
#include "infiniop/ops/selu.h"
#include "infiniop/ops/sigmoid.h"
#include "infiniop/ops/silu.h"
Expand Down
33 changes: 33 additions & 0 deletions include/infiniop/ops/fp8_indexer_logits.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
#ifndef __INFINIOP_FP8_INDEXER_LOGITS_API_H__
#define __INFINIOP_FP8_INDEXER_LOGITS_API_H__

#include "../operator_descriptor.h"

typedef struct InfiniopDescriptor *infiniopFp8IndexerLogitsDescriptor_t;

__INFINI_C __export infiniStatus_t infiniopCreateFp8IndexerLogitsDescriptor(
infiniopHandle_t handle,
infiniopFp8IndexerLogitsDescriptor_t *desc_ptr,
infiniopTensorDescriptor_t logits_desc,
infiniopTensorDescriptor_t q_fp8_desc,
infiniopTensorDescriptor_t kv_cache_desc,
infiniopTensorDescriptor_t block_tables_desc,
infiniopTensorDescriptor_t weights_fp32_desc,
infiniopTensorDescriptor_t positions_desc,
infiniopTensorDescriptor_t request_ids_desc);

__INFINI_C __export infiniStatus_t infiniopFp8IndexerLogits(
infiniopFp8IndexerLogitsDescriptor_t desc,
void *logits,
const void *q_fp8,
const void *kv_cache,
const void *block_tables,
const void *weights_fp32,
const void *positions,
const void *request_ids,
void *stream);

__INFINI_C __export infiniStatus_t infiniopDestroyFp8IndexerLogitsDescriptor(
infiniopFp8IndexerLogitsDescriptor_t desc);

#endif
64 changes: 64 additions & 0 deletions include/infiniop/ops/fp8_indexer_quant.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
#ifndef __INFINIOP_FP8_INDEXER_QUANT_API_H__
#define __INFINIOP_FP8_INDEXER_QUANT_API_H__

#include "../operator_descriptor.h"

#include <stdint.h>

typedef struct InfiniopDescriptor *infiniopFp8IndexerQuantDescriptor_t;
typedef struct InfiniopDescriptor *infiniopFusedFp8IndexerDescriptor_t;

__INFINI_C __export infiniStatus_t infiniopCreateFp8IndexerQuantDescriptor(
infiniopHandle_t handle,
infiniopFp8IndexerQuantDescriptor_t *desc_ptr,
infiniopTensorDescriptor_t q_fp8_desc,
infiniopTensorDescriptor_t weights_fp32_desc,
infiniopTensorDescriptor_t q_desc,
infiniopTensorDescriptor_t weights_desc);

__INFINI_C __export infiniStatus_t infiniopFp8IndexerQuant(
infiniopFp8IndexerQuantDescriptor_t desc,
void *q_fp8,
void *weights_fp32,
const void *q,
const void *weights,
void *stream);

__INFINI_C __export infiniStatus_t infiniopDestroyFp8IndexerQuantDescriptor(
infiniopFp8IndexerQuantDescriptor_t desc);

__INFINI_C __export infiniStatus_t infiniopCreateFusedFp8IndexerDescriptor(
infiniopHandle_t handle,
infiniopFusedFp8IndexerDescriptor_t *desc_ptr,
infiniopTensorDescriptor_t q_fp8_desc,
infiniopTensorDescriptor_t weights_fp32_desc,
infiniopTensorDescriptor_t k_cache_desc,
infiniopTensorDescriptor_t q_raw_desc,
infiniopTensorDescriptor_t k_weights_desc,
infiniopTensorDescriptor_t norm_weight_desc,
infiniopTensorDescriptor_t norm_bias_desc,
infiniopTensorDescriptor_t positions_desc,
infiniopTensorDescriptor_t cos_sin_cache_desc,
infiniopTensorDescriptor_t slot_mapping_desc,
uint64_t rope_dim,
double eps,
double weights_scale);

__INFINI_C __export infiniStatus_t infiniopFusedFp8Indexer(
infiniopFusedFp8IndexerDescriptor_t desc,
void *q_fp8,
void *weights_fp32,
void *k_cache,
const void *q_raw,
const void *k_weights,
const void *norm_weight,
const void *norm_bias,
const void *positions,
const void *cos_sin_cache,
const void *slot_mapping,
void *stream);

__INFINI_C __export infiniStatus_t infiniopDestroyFusedFp8IndexerDescriptor(
infiniopFusedFp8IndexerDescriptor_t desc);

#endif
32 changes: 32 additions & 0 deletions include/infiniop/ops/fp8_mla_rmsnorm_cache.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
#ifndef __INFINIOP_FP8_MLA_RMSNORM_CACHE_API_H__
#define __INFINIOP_FP8_MLA_RMSNORM_CACHE_API_H__

#include "../operator_descriptor.h"

typedef struct InfiniopDescriptor *infiniopFp8MlaRmsnormCacheDescriptor_t;

__INFINI_C __export infiniStatus_t infiniopCreateFp8MlaRmsnormCacheDescriptor(
infiniopHandle_t handle,
infiniopFp8MlaRmsnormCacheDescriptor_t *desc_ptr,
infiniopTensorDescriptor_t cache_desc,
infiniopTensorDescriptor_t vendor_cache_desc,
infiniopTensorDescriptor_t compressed_kv_desc,
infiniopTensorDescriptor_t norm_weight_desc,
infiniopTensorDescriptor_t rope_desc,
infiniopTensorDescriptor_t slot_mapping_desc,
double eps);

__INFINI_C __export infiniStatus_t infiniopFp8MlaRmsnormCache(
infiniopFp8MlaRmsnormCacheDescriptor_t desc,
void *cache,
void *vendor_cache,
const void *compressed_kv,
const void *norm_weight,
const void *rope,
const void *slot_mapping,
void *stream);

__INFINI_C __export infiniStatus_t infiniopDestroyFp8MlaRmsnormCacheDescriptor(
infiniopFp8MlaRmsnormCacheDescriptor_t desc);

#endif
36 changes: 36 additions & 0 deletions include/infiniop/ops/fp8_sparse_mla.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
#ifndef __INFINIOP_FP8_SPARSE_MLA_API_H__
#define __INFINIOP_FP8_SPARSE_MLA_API_H__

#include "../operator_descriptor.h"

typedef struct InfiniopDescriptor *infiniopFp8SparseMlaDescriptor_t;

__INFINI_C __export infiniStatus_t infiniopCreateFp8SparseMlaDescriptor(
infiniopHandle_t handle,
infiniopFp8SparseMlaDescriptor_t *desc_ptr,
infiniopTensorDescriptor_t output_desc,
infiniopTensorDescriptor_t query_desc,
infiniopTensorDescriptor_t kv_cache_desc,
infiniopTensorDescriptor_t indices_desc,
infiniopTensorDescriptor_t topk_lens_desc,
float scale);

__INFINI_C __export infiniStatus_t infiniopGetFp8SparseMlaWorkspaceSize(
infiniopFp8SparseMlaDescriptor_t desc,
size_t *size);

__INFINI_C __export infiniStatus_t infiniopFp8SparseMla(
infiniopFp8SparseMlaDescriptor_t desc,
void *workspace,
size_t workspace_size,
void *output,
const void *query,
const void *kv_cache,
const void *indices,
const void *topk_lens,
void *stream);

__INFINI_C __export infiniStatus_t infiniopDestroyFp8SparseMlaDescriptor(
infiniopFp8SparseMlaDescriptor_t desc);

#endif
25 changes: 25 additions & 0 deletions include/infiniop/ops/select_last_token_hidden.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
#ifndef __INFINIOP_SELECT_LAST_TOKEN_HIDDEN_API_H__
#define __INFINIOP_SELECT_LAST_TOKEN_HIDDEN_API_H__

#include "../operator_descriptor.h"

typedef struct InfiniopDescriptor *infiniopSelectLastTokenHiddenDescriptor_t;

__INFINI_C __export infiniStatus_t infiniopCreateSelectLastTokenHiddenDescriptor(
infiniopHandle_t handle,
infiniopSelectLastTokenHiddenDescriptor_t *desc_ptr,
infiniopTensorDescriptor_t output_desc,
infiniopTensorDescriptor_t hidden_states_desc,
infiniopTensorDescriptor_t input_offsets_desc);

__INFINI_C __export infiniStatus_t infiniopSelectLastTokenHidden(
infiniopSelectLastTokenHiddenDescriptor_t desc,
void *output,
const void *hidden_states,
const void *input_offsets,
void *stream);

__INFINI_C __export infiniStatus_t infiniopDestroySelectLastTokenHiddenDescriptor(
infiniopSelectLastTokenHiddenDescriptor_t desc);

#endif
2 changes: 2 additions & 0 deletions python/infinicore/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
double,
dtype,
float,
float8,
float16,
float32,
float64,
Expand Down Expand Up @@ -188,6 +189,7 @@
"complex128",
"double",
"float",
"float8",
"float16",
"float32",
"float64",
Expand Down
1 change: 1 addition & 0 deletions python/infinicore/dtype.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@ def __hash__(self):
cdouble = complex128
float16 = dtype(_infinicore.DataType.F16)
half = float16
float8 = dtype(_infinicore.DataType.F8)
bfloat16 = dtype(_infinicore.DataType.BF16)
uint8 = dtype(_infinicore.DataType.U8)
int8 = dtype(_infinicore.DataType.I8)
Expand Down
Loading
Loading