diff --git a/3rdparty/composable_kernel b/3rdparty/composable_kernel index af7118e3425..f33252cebe5 160000 --- a/3rdparty/composable_kernel +++ b/3rdparty/composable_kernel @@ -1 +1 @@ -Subproject commit af7118e342580ecd3f71edce7b1d0ba465012ecf +Subproject commit f33252cebe5a52362ec1ee12c124dde7800dda3a diff --git a/csrc/ck_deepgemm/include/deepgemm_common.cuh b/csrc/ck_deepgemm/include/deepgemm_common.cuh index a524e25e320..cd374dd18ad 100644 --- a/csrc/ck_deepgemm/include/deepgemm_common.cuh +++ b/csrc/ck_deepgemm/include/deepgemm_common.cuh @@ -176,7 +176,7 @@ void grouped_flatmm(KernelArguments& args, ck_stream_config& s) auto kargs = Kernel::MakeKernelArgs(args); const dim3 grids = Kernel::GridSize(kargs); - constexpr dim3 blocks = Kernel::BlockSize(); + const dim3 blocks = Kernel::BlockSize(); ck_tile::launch_kernel( s, ck_tile::make_kernel(Kernel{}, grids, blocks, 0, kargs)); diff --git a/csrc/ck_tile_gemm_moe_2stages/include/moe_cktile2stages_common.cuh b/csrc/ck_tile_gemm_moe_2stages/include/moe_cktile2stages_common.cuh index 514f5d96294..03b0f9ef012 100644 --- a/csrc/ck_tile_gemm_moe_2stages/include/moe_cktile2stages_common.cuh +++ b/csrc/ck_tile_gemm_moe_2stages/include/moe_cktile2stages_common.cuh @@ -246,7 +246,7 @@ void moe_gemm(const MoeFlatmmHostArgs& args, const ck_stream_config& s) auto kargs = Kernel::MakeKernelArgs(args); const dim3 grids = Kernel::GridSize(kargs); - constexpr dim3 blocks = Kernel::BlockSize(); + const dim3 blocks = Kernel::BlockSize(); // if(!Kernel::IsSupportedArgument(kargs)) // { diff --git a/csrc/cktile_gemm_a8w8_bpreshuffle/include/gemm_a8w8_bpreshuffle_cktile_common.cuh b/csrc/cktile_gemm_a8w8_bpreshuffle/include/gemm_a8w8_bpreshuffle_cktile_common.cuh index 10d057cf250..a6480622c13 100644 --- a/csrc/cktile_gemm_a8w8_bpreshuffle/include/gemm_a8w8_bpreshuffle_cktile_common.cuh +++ b/csrc/cktile_gemm_a8w8_bpreshuffle/include/gemm_a8w8_bpreshuffle_cktile_common.cuh @@ -144,7 +144,7 @@ float flatmm_calc(const ck_tile::ScaleFlatmmHostArgs& args, auto kargs = Kernel::MakeKernelArgs(args); const dim3 grids = Kernel::GridSize(kargs); - constexpr dim3 blocks = Kernel::BlockSize(); + const dim3 blocks = Kernel::BlockSize(); if(!Kernel::IsSupportedArgument(kargs)) { diff --git a/csrc/cpp_itfs/mha_fwd_batch_prefill.cu b/csrc/cpp_itfs/mha_fwd_batch_prefill.cu index daf298d0832..b7ec63ffe7e 100644 --- a/csrc/cpp_itfs/mha_fwd_batch_prefill.cu +++ b/csrc/cpp_itfs/mha_fwd_batch_prefill.cu @@ -49,7 +49,7 @@ float mha_batch_prefill(mha_batch_prefill_args args, int head_size_q = args.hdim_q; int head_size_v = args.hdim_v; bool has_dropout = args.p_drop > 0.f; - bool has_sink = args.sink_size > 0; + bool has_sink = args.sink_size > 0 || args.sink_ptr != nullptr; // The kUseGlobalLoad decision (>2GB KV cache → use `global_load_lds_*` // instead of SRD `buffer_load_*`) is made per-arm inside the auto-generated diff --git a/csrc/py_itfs_ck/mha_fwd_kernels.cu b/csrc/py_itfs_ck/mha_fwd_kernels.cu index caa071860be..e962cc24415 100644 --- a/csrc/py_itfs_ck/mha_fwd_kernels.cu +++ b/csrc/py_itfs_ck/mha_fwd_kernels.cu @@ -109,7 +109,7 @@ mha_fwd_args get_ck_fmha_fwd_args(bool has_lse, static_cast(bias_type), has_lse, static_cast(qscale_type), - mask.sink > 0, // has_sink + (mask.sink > 0) || (sink_ptr != nullptr), // has_sink: true for streaming-sink window OR a learned per-head sink_ptr (e.g. gpt-oss, sink_size==0) q.data_ptr(), k.data_ptr(), v.data_ptr(), diff --git a/csrc/py_itfs_ck/mha_varlen_fwd_kernels.cu b/csrc/py_itfs_ck/mha_varlen_fwd_kernels.cu index 645c46d57df..3036b15d9ae 100644 --- a/csrc/py_itfs_ck/mha_varlen_fwd_kernels.cu +++ b/csrc/py_itfs_ck/mha_varlen_fwd_kernels.cu @@ -136,7 +136,7 @@ mha_fwd_args get_ck_fmha_varlen_fwd_args(bool has_lse, static_cast(bias_type), has_lse, static_cast(qscale_type), - mask.sink > 0, // hsa_sink + (mask.sink > 0) || (sink_ptr != nullptr), // has_sink: true for streaming-sink window OR a learned per-head sink_ptr (e.g. gpt-oss, sink_size==0) q.data_ptr(), k.data_ptr(), v.data_ptr(), @@ -496,7 +496,7 @@ mha_varlen_fwd( std::string mask_identify = "b:" + std::to_string(window_size_left) + "," + std::to_string(window_size_right) + "," + std::to_string(sink_size); mask = mask_info::decode(mask_identify, max_seqlen_q, max_seqlen_k); // local } - bool has_sink = mask.sink > 0; + bool has_sink = (mask.sink > 0) || sink_ptr.has_value(); CHECK_SHAPE(q, total_q, num_heads, head_size_q); if (!paged_KV) { const int total_k = k.size(0);