Instructions to use replicate/flashinfer-draft with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Kernels
How to use replicate/flashinfer-draft with Kernels:
# !pip install kernels from kernels import get_kernel # a version (or an explicit revision) is required; see the "Files and versions" tab for the available ones kernel = get_kernel("replicate/flashinfer-draft", version=1) - Notebooks
- Google Colab
- Kaggle
File size: 3,819 Bytes
57c3a10 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 | /*
* Copyright (c) 2024 by FlashInfer team.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
#ifndef FLASHINFER_ATTENTION_HOPPER_PARAMS_CUH
#define FLASHINFER_ATTENTION_HOPPER_PARAMS_CUH
#include <cuda.h>
#include <vector>
namespace flashinfer {
template <typename DTypeQ_, typename DTypeKV_, typename DTypeO_, typename IdType_ = int32_t>
struct SinglePrefillParams {
using DTypeQ = DTypeQ_;
using DTypeKV = DTypeKV_;
using DTypeO = DTypeO_;
using IdType = IdType_;
// The QKV matrices.
DTypeQ* q_ptr;
DTypeKV* k_ptr;
DTypeKV* v_ptr;
DTypeO* o_ptr;
float* lse_ptr;
struct AdditionalParams {
float logits_soft_cap;
float sm_scale;
float* scale_q;
float* scale_k;
float* scale_v;
} additional_params;
int64_t q_stride_n;
int64_t k_stride_n;
int64_t v_stride_n;
int64_t o_stride_n;
int64_t q_stride_h;
int64_t k_stride_h;
int64_t v_stride_h;
int64_t o_stride_h;
int qo_len;
int kv_len;
int num_qo_heads;
int num_kv_heads;
int group_size;
int window_left;
bool causal;
};
template <typename DTypeQ_, typename DTypeKV_, typename DTypeO_, typename IdType_>
struct BatchPrefillRaggedParams {
using DTypeQ = DTypeQ_;
using DTypeKV = DTypeKV_;
using DTypeO = DTypeO_;
using IdType = IdType_;
// The QKV matrices.
DTypeQ* q_ptr;
DTypeKV* k_ptr;
DTypeKV* v_ptr;
DTypeO* o_ptr;
float* lse_ptr;
IdType* qo_tile_indices;
IdType* qo_indptr;
IdType* kv_indptr;
IdType* qo_lens;
IdType* kv_lens;
IdType* head_indices;
IdType* work_indptr;
IdType* batch_indices;
struct AdditionalParams {
float logits_soft_cap;
float sm_scale;
uint32_t* maybe_prefix_len_ptr;
uint16_t* maybe_token_pos_in_items_ptr;
uint32_t token_pos_in_items_len;
uint16_t* maybe_max_item_len_ptr;
} additional_params;
int64_t q_stride_n;
int64_t k_stride_n;
int64_t v_stride_n;
int64_t o_stride_n;
int64_t q_stride_h;
int64_t k_stride_h;
int64_t v_stride_h;
int64_t o_stride_h;
int64_t nnz_qo;
int64_t nnz_kv;
int num_qo_heads;
int num_kv_heads;
int group_size;
int window_left;
bool causal;
};
template <typename DTypeQ_, typename DTypeKV_, typename DTypeO_, typename IdType_>
struct BatchPrefillPagedParams {
using DTypeQ = DTypeQ_;
using DTypeKV = DTypeKV_;
using DTypeO = DTypeO_;
using IdType = IdType_;
// The QKV matrices.
DTypeQ* q_ptr;
DTypeKV* k_ptr;
DTypeKV* v_ptr;
DTypeO* o_ptr;
float* lse_ptr;
IdType* qo_tile_indices;
IdType* qo_indptr;
IdType* kv_indptr;
IdType* kv_indices;
IdType* qo_lens;
IdType* kv_lens;
IdType* head_indices;
IdType* work_indptr;
IdType* batch_indices;
struct AdditionalParams {
float logits_soft_cap;
float sm_scale;
uint32_t* maybe_prefix_len_ptr;
uint16_t* maybe_token_pos_in_items_ptr;
uint32_t token_pos_in_items_len;
uint16_t* maybe_max_item_len_ptr;
} additional_params;
int64_t q_stride_n;
int64_t k_stride_n;
int64_t v_stride_n;
int64_t o_stride_n;
int64_t q_stride_h;
int64_t k_stride_h;
int64_t v_stride_h;
int64_t o_stride_h;
int64_t nnz_qo;
int num_qo_heads;
int num_kv_heads;
int group_size;
int page_size;
int window_left;
bool causal;
};
} // namespace flashinfer
#endif // FLASHINFER_ATTENTION_HOPPER_PARAMS_CUH
|