move all get_stream in sgl_kernel to c++ to reduce the launch overhead (#12521)
This commit is contained in:
@@ -90,13 +90,13 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
||||
|
||||
m.def(
|
||||
"apply_rope_pos_ids_cos_sin_cache(Tensor q, Tensor k, Tensor! q_rope, Tensor! k_rope, Tensor cos_sin_cache, "
|
||||
"Tensor pos_ids, bool interleave, bool enable_pdl, int cuda_stream, "
|
||||
"Tensor pos_ids, bool interleave, bool enable_pdl, "
|
||||
"Tensor? v, Tensor!? k_buffer, Tensor!? v_buffer, Tensor? kv_cache_loc) -> ()");
|
||||
m.impl("apply_rope_pos_ids_cos_sin_cache", torch::kCUDA, &apply_rope_pos_ids_cos_sin_cache);
|
||||
|
||||
m.def(
|
||||
"downcast_fp8(Tensor k, Tensor v, Tensor k_out, Tensor v_out, Tensor k_scale, Tensor v_scale, Tensor loc, int "
|
||||
"mult, int offset, int cuda_stream) -> ()");
|
||||
"downcast_fp8(Tensor k, Tensor v, Tensor k_out, Tensor v_out, Tensor k_scale, Tensor v_scale, Tensor loc, "
|
||||
"int mult, int offset) -> ()");
|
||||
m.impl("downcast_fp8", torch::kCUDA, &downcast_fp8);
|
||||
|
||||
m.def("copy_to_gpu_no_ce(Tensor input, Tensor! output) -> ()");
|
||||
@@ -303,13 +303,13 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
||||
"Tensor candidates, Tensor retrive_index, Tensor retrive_next_token, Tensor retrive_next_sibling, "
|
||||
"Tensor uniform_samples, Tensor uniform_samples_for_final_sampling, Tensor target_probs, Tensor draft_probs, "
|
||||
"float threshold_single, float threshold_acc, "
|
||||
"bool deterministic, int cuda_stream) -> ()");
|
||||
"bool deterministic) -> ()");
|
||||
m.impl("tree_speculative_sampling_target_only", torch::kCUDA, &tree_speculative_sampling_target_only);
|
||||
|
||||
m.def(
|
||||
"verify_tree_greedy(Tensor! predicts, Tensor! accept_index, Tensor! accept_token_num, "
|
||||
"Tensor candidates, Tensor retrive_index, Tensor retrive_next_token, Tensor retrive_next_sibling, "
|
||||
"Tensor target_predict, int cuda_stream) -> ()");
|
||||
"Tensor target_predict) -> ()");
|
||||
m.impl("verify_tree_greedy", torch::kCUDA, &verify_tree_greedy);
|
||||
|
||||
m.def(
|
||||
@@ -403,8 +403,8 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
||||
* From FlashInfer
|
||||
*/
|
||||
m.def(
|
||||
"bmm_fp8(Tensor A, Tensor B, Tensor! D, Tensor A_scale, Tensor B_scale, Tensor workspace_buffer, int "
|
||||
"cublas_handle, int cuda_stream) -> ()",
|
||||
"bmm_fp8(Tensor A, Tensor B, Tensor! D, Tensor A_scale, Tensor B_scale, Tensor workspace_buffer, "
|
||||
"int cublas_handle) -> ()",
|
||||
{at::Tag::needs_fixed_stride_order});
|
||||
m.impl("bmm_fp8", torch::kCUDA, &bmm_fp8);
|
||||
|
||||
|
||||
@@ -106,7 +106,7 @@ TORCH_LIBRARY_EXPAND(sgl_kernel, m) {
|
||||
m.def(
|
||||
"verify_tree_greedy(Tensor! predicts, Tensor! accept_index, Tensor! accept_token_num, "
|
||||
"Tensor candidates, Tensor retrive_index, Tensor retrive_next_token, Tensor retrive_next_sibling, "
|
||||
"Tensor target_predict, int cuda_stream) -> ()");
|
||||
"Tensor target_predict) -> ()");
|
||||
m.impl("verify_tree_greedy", torch::kCUDA, &verify_tree_greedy);
|
||||
|
||||
m.def(
|
||||
|
||||
@@ -150,14 +150,13 @@ void downcast_fp8(
|
||||
at::Tensor& v_scale,
|
||||
at::Tensor& loc,
|
||||
int64_t mult,
|
||||
int64_t offset,
|
||||
int64_t cuda_stream) {
|
||||
int64_t offset) {
|
||||
CHECK_INPUT(k);
|
||||
CHECK_INPUT(v);
|
||||
CHECK_INPUT(k_out);
|
||||
CHECK_INPUT(v_out);
|
||||
|
||||
cudaStream_t stream = reinterpret_cast<cudaStream_t>(cuda_stream);
|
||||
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
switch (k.scalar_type()) {
|
||||
case at::ScalarType::BFloat16:
|
||||
downcast_fp8_impl<__nv_bfloat16>(k, v, k_out, v_out, k_scale, v_scale, loc, mult, offset, stream);
|
||||
|
||||
@@ -28,7 +28,6 @@ void apply_rope_pos_ids_cos_sin_cache(
|
||||
at::Tensor pos_ids,
|
||||
bool interleave,
|
||||
bool enable_pdl,
|
||||
int64_t cuda_stream,
|
||||
const std::optional<at::Tensor>& v,
|
||||
const std::optional<at::Tensor>& k_buffer,
|
||||
const std::optional<at::Tensor>& v_buffer,
|
||||
@@ -88,7 +87,7 @@ void apply_rope_pos_ids_cos_sin_cache(
|
||||
size_t k_rope_stride_n = k_rope.stride(0);
|
||||
size_t k_rope_stride_h = k_rope.stride(1);
|
||||
|
||||
cudaStream_t stream = reinterpret_cast<cudaStream_t>(cuda_stream);
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
DISPATCH_PYTORCH_DTYPE_TO_CTYPE_FP16(q.scalar_type(), c_type, [&] {
|
||||
// TODO temporarily only use `BatchQKApplyRotaryPosIdsCosSinCacheEnhanced` when save_kv_cache
|
||||
// to avoid changing original code path; but this branch is feature-complete and should switch to this later
|
||||
|
||||
@@ -27,8 +27,7 @@ void bmm_fp8(
|
||||
at::Tensor A_scale,
|
||||
at::Tensor B_scale,
|
||||
at::Tensor workspace_buffer,
|
||||
int64_t cublas_handle,
|
||||
int64_t cuda_stream) {
|
||||
int64_t cublas_handle) {
|
||||
TORCH_CHECK(A.is_cuda(), "A must be a CUDA tensor");
|
||||
TORCH_CHECK(B.is_cuda(), "B must be a CUDA tensor");
|
||||
TORCH_CHECK(D.is_cuda(), "D must be a CUDA tensor");
|
||||
@@ -51,7 +50,7 @@ void bmm_fp8(
|
||||
auto n = B.size(2);
|
||||
|
||||
auto lt_handle = reinterpret_cast<cublasLtHandle_t>(cublas_handle);
|
||||
auto stream = reinterpret_cast<cudaStream_t>(cuda_stream);
|
||||
auto stream = at::cuda::getCurrentCUDAStream();
|
||||
|
||||
auto status = flashinfer::bmm_fp8::bmm_fp8_internal_cublaslt(
|
||||
workspace_buffer.data_ptr(),
|
||||
|
||||
@@ -328,8 +328,7 @@ void verify_tree_greedy(
|
||||
at::Tensor retrive_index,
|
||||
at::Tensor retrive_next_token,
|
||||
at::Tensor retrive_next_sibling,
|
||||
at::Tensor target_predict,
|
||||
int64_t cuda_stream = 0) {
|
||||
at::Tensor target_predict) {
|
||||
CHECK_INPUT(candidates);
|
||||
CHECK_INPUT(retrive_index);
|
||||
CHECK_INPUT(retrive_next_token);
|
||||
@@ -389,7 +388,7 @@ void verify_tree_greedy(
|
||||
throw std::runtime_error("Expected 'target_predict' to be of type long (torch.int64).");
|
||||
}
|
||||
|
||||
cudaStream_t stream = reinterpret_cast<cudaStream_t>(cuda_stream);
|
||||
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
dim3 grid(batch_size);
|
||||
dim3 block(1);
|
||||
|
||||
|
||||
@@ -42,8 +42,7 @@ void tree_speculative_sampling_target_only(
|
||||
at::Tensor draft_probs,
|
||||
double threshold_single,
|
||||
double threshold_acc,
|
||||
bool deterministic = true,
|
||||
int64_t cuda_stream = 0) {
|
||||
bool deterministic = true) {
|
||||
CHECK_INPUT(candidates);
|
||||
CHECK_INPUT(retrive_index);
|
||||
CHECK_INPUT(retrive_next_token);
|
||||
@@ -124,7 +123,7 @@ void tree_speculative_sampling_target_only(
|
||||
CHECK_GE(threshold_acc, 0);
|
||||
CHECK_GE(1, threshold_acc);
|
||||
|
||||
cudaStream_t stream = reinterpret_cast<cudaStream_t>(cuda_stream);
|
||||
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||
cudaError_t status = sampling::TreeSpeculativeSamplingTargetOnly<float, int32_t, int64_t>(
|
||||
static_cast<int32_t*>(predicts.data_ptr()),
|
||||
static_cast<int32_t*>(accept_index.data_ptr()),
|
||||
|
||||
Reference in New Issue
Block a user