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 kernel = get_kernel("replicate/flashinfer-draft") - Notebooks
- Google Colab
- Kaggle
File size: 4,466 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 | /*
* Copyright (c) 2023 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_MATH_CUH_
#define FLASHINFER_MATH_CUH_
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cstdint>
namespace flashinfer {
namespace math {
// log2(e)
constexpr float log2e = 1.44269504088896340736f;
constexpr float loge2 = 0.693147180559945309417f;
constexpr float inf = 5e4;
__forceinline__ __device__ half2 uint32_as_half2(uint32_t x) { return *(half2*)&x; }
__forceinline__ __device__ uint32_t half2_as_uint32(half2 x) { return *(uint32_t*)&x; }
/*!
* \brief Wrapper of PTX ex2.approx instruction, which computes 2^x
* \param x input
*/
__forceinline__ __device__ float ptx_exp2(float x) {
float y;
asm volatile("ex2.approx.ftz.f32 %0, %1;" : "=f"(y) : "f"(x));
return y;
}
/*!
* \brief Wrapper of PTX lg2.approx instruction, which computes log2(x)
* \param x input
*/
__forceinline__ __device__ float ptx_log2(float x) {
float y;
asm volatile("lg2.approx.ftz.f32 %0, %1;" : "=f"(y) : "f"(x));
return y;
}
/*!
* \brief Wrapper of PTX ex2.approx.f16x2 instruction, which computes 2^x
* \param x input
*/
__forceinline__ __device__ half2 ptx_exp2(half2 x) {
uint32_t y_u32;
uint32_t x_u32 = half2_as_uint32(x);
asm volatile("ex2.approx.f16x2 %0, %1;" : "=r"(y_u32) : "r"(x_u32));
return uint32_as_half2(y_u32);
}
/*!
* \brief Wrapper of PTX ex2.approx.f16 instruction, which computes 2^x
* \param x input
*/
__forceinline__ __device__ half ptx_exp2(half x) {
ushort y_u16;
asm volatile("ex2.approx.f16 %0, %1;" : "=h"(y_u16) : "h"(__half_as_ushort(x)));
return __ushort_as_half(y_u16);
}
/*!
* \brief Wrapper of PTX rcp.approx instruction, which computes 1/x
* \param x input
*/
__forceinline__ __device__ float ptx_rcp(float x) {
float y;
asm volatile("rcp.approx.ftz.f32 %0, %1;" : "=f"(y) : "f"(x));
return y;
}
/*!
* \brief Wrapper of PTX shfl.sync.bfly instruction, which performs a butterfly shuffle
* between threads in a warp.
* \param x The value in the source lane
* \param lane_mask The mask to perform thread index xor with: y[i] <- x[i ^ delta]
*/
__forceinline__ __device__ float shfl_xor_sync(float x, int lane_mask) {
float y;
asm volatile("shfl.sync.bfly.b32 %0, %1, %2, 0x1f, 0xffffffff;"
: "=f"(y)
: "f"(x), "r"(lane_mask));
return y;
}
/*!
* \brief Wrapper of PTX shfl.sync.bfly instruction on half2, which performs a butterfly
* shuffle between threads in a warp.
* \param x The value in the source lane
* \param lane_mask The mask to perform thread index xor with: y[i] <- x[i ^ lane_mask]
*/
__forceinline__ __device__ half2 shfl_xor_sync(half2 x, int lane_mask) {
return __shfl_xor_sync(0xffffffff, x, lane_mask);
}
/*!
* \brief Wrapper of PTX rsqrt approximation instruction, which computes 1/sqrt(x)
* \param x input
*/
__forceinline__ __device__ float rsqrt(float x) {
float y;
asm volatile("rsqrt.approx.ftz.f32 %0, %1;" : "=f"(y) : "f"(x));
return y;
}
/*!
* \brief Wrapper of PTX tanh.approx.f32 instruction, which computes tanh(x)
* \param x input
*/
__forceinline__ __device__ float tanh(float x) {
float y;
asm volatile("tanh.approx.f32 %0, %1;" : "=f"(y) : "f"(x));
return y;
}
/*!
* \brief Wrapper of PTX tanh.approx.f16x2 instruction, which computes tanh(x)
* \param x input
*/
__forceinline__ __device__ half2 tanh(half2 x) {
uint32_t y_u32;
uint32_t x_u32 = half2_as_uint32(x);
asm volatile("tanh.approx.f16x2 %0, %1;" : "=r"(y_u32) : "r"(x_u32));
return uint32_as_half2(y_u32);
}
/*!
* \brief Wrapper of PTX tanh.approx.f16 instruction, which computes tanh(x)
* \param x input
*/
__forceinline__ __device__ half tanh(half x) {
ushort y_u16;
asm volatile("tanh.approx.f16 %0, %1;" : "=h"(y_u16) : "h"(__half_as_ushort(x)));
return __ushort_as_half(y_u16);
}
} // namespace math
} // namespace flashinfer
#endif // FLASHINFER_MATH_CUH_
|