Restrict (fused) MMVF to 2*ts alignment, and cublas FP8 fo 16-byte

This commit is contained in:
Oliver Simons
2026-10-09 21:11:12 +02:00
parent 8a3884fee8
commit f771a191ef
5 changed files with 41 additions and 23 deletions
+2 -1
View File
@@ -133,7 +133,8 @@ bool ggml_cuda_mul_mat_fp8(
const int cc = ggml_cuda_info().devices[ctx.device].cc; const int cc = ggml_cuda_info().devices[ctx.device].cc;
if (!fp8_mma_hardware_available(cc) || src0->type != GGML_TYPE_F8_E4M3 || src1->type != GGML_TYPE_F32 || if (!fp8_mma_hardware_available(cc) || src0->type != GGML_TYPE_F8_E4M3 || src1->type != GGML_TYPE_F32 ||
dst->type != GGML_TYPE_F32 || !ggml_is_contiguous(dst) || src0->ne[0] % 16 != 0 || src0->ne[1] % 16 != 0 || dst->type != GGML_TYPE_F32 || !ggml_is_contiguous(dst) || src0->ne[0] % 16 != 0 || src0->ne[1] % 16 != 0 ||
src0->nb[0] != sizeof(uint8_t) || src0->nb[1] != (size_t) src0->ne[0] || src1->nb[0] != sizeof(float)) { src0->nb[0] != sizeof(uint8_t) || src0->nb[1] != (size_t) src0->ne[0] || src1->nb[0] != sizeof(float) ||
!ggml_cuda_is_aligned(src0, 16) || !ggml_cuda_is_aligned(dst, 16)) {
return false; return false;
} }
+14 -10
View File
@@ -1800,7 +1800,7 @@ static bool ggml_cuda_should_fuse_mul_mat(const ggml_tensor * ffn_up,
return true; return true;
} }
static bool ggml_cuda_should_fuse_mul_mat_vec_f(const ggml_tensor * tensor) { static bool ggml_cuda_should_fuse_mul_mat_vec_f(const ggml_tensor * tensor, const ggml_tensor * gate = nullptr) {
ggml_tensor * src0 = tensor->src[0]; ggml_tensor * src0 = tensor->src[0];
ggml_tensor * src1 = tensor->src[1]; ggml_tensor * src1 = tensor->src[1];
const ggml_tensor * dst = tensor; const ggml_tensor * dst = tensor;
@@ -1814,7 +1814,11 @@ static bool ggml_cuda_should_fuse_mul_mat_vec_f(const ggml_tensor * tensor) {
const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc; const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc;
const int warp_size = ggml_cuda_info().devices[ggml_cuda_get_device()].warp_size; const int warp_size = ggml_cuda_info().devices[ggml_cuda_get_device()].warp_size;
use_mul_mat_vec_f = use_mul_mat_vec_f && ggml_cuda_should_use_mmvf(src0->type, cc, warp_size, src0->ne, src0->nb, is_mul_mat_id ? src1->ne[2] : src1->ne[1]); use_mul_mat_vec_f = use_mul_mat_vec_f && ggml_cuda_should_use_mmvf(src0, cc, warp_size, is_mul_mat_id ? src1->ne[2] : src1->ne[1]);
if (gate && !ggml_cuda_should_use_mmvf(gate->src[0], cc, warp_size, is_mul_mat_id ? src1->ne[2] : src1->ne[1])) {
return false;
}
//we only support fusion for ncols_dst = 1 //we only support fusion for ncols_dst = 1
if (tensor->op == GGML_OP_MUL_MAT && dst->ne[1] != 1) { if (tensor->op == GGML_OP_MUL_MAT && dst->ne[1] != 1) {
@@ -1930,7 +1934,7 @@ static void ggml_cuda_mul_mat(ggml_backend_cuda_context & ctx, const ggml_tensor
const int cc = ggml_cuda_info().devices[ctx.device].cc; const int cc = ggml_cuda_info().devices[ctx.device].cc;
const int warp_size = ggml_cuda_info().devices[ctx.device].warp_size; const int warp_size = ggml_cuda_info().devices[ctx.device].warp_size;
if (ggml_cuda_should_use_mmvf(src0->type, cc, warp_size, src0->ne, src0->nb, ne11)) { if (ggml_cuda_should_use_mmvf(src0, cc, warp_size, ne11)) {
// The custom vector kernel can be used over batched cuBLAS GEMM. // The custom vector kernel can be used over batched cuBLAS GEMM.
// But this is only faster for GPUs without tensor cores or with a thin src0 matrix (particularly KQV in attention) // But this is only faster for GPUs without tensor cores or with a thin src0 matrix (particularly KQV in attention)
ggml_cuda_mul_mat_vec_f(ctx, src0, src1, nullptr, dst); ggml_cuda_mul_mat_vec_f(ctx, src0, src1, nullptr, dst);
@@ -1940,7 +1944,7 @@ static void ggml_cuda_mul_mat(ggml_backend_cuda_context & ctx, const ggml_tensor
if (ne01 == 1 && ne11 > MMVF_MAX_BATCH_SIZE && ne2 == 1 && ne3 == 1 if (ne01 == 1 && ne11 > MMVF_MAX_BATCH_SIZE && ne2 == 1 && ne3 == 1
&& src0->type == GGML_TYPE_F32 && src0->type == GGML_TYPE_F32
&& ggml_is_contiguous(src0) && ggml_is_contiguous(src1) && ggml_is_contiguous(dst) && ggml_is_contiguous(src0) && ggml_is_contiguous(src1) && ggml_is_contiguous(dst)
&& ggml_cuda_should_use_mmvf(src1->type, cc, warp_size, src1->ne, src1->nb, /*ne11 =*/ 1)) { && ggml_cuda_should_use_mmvf(src1, cc, warp_size, /*ne11 =*/ 1)) {
ggml_tensor dst_vec = *dst; ggml_tensor dst_vec = *dst;
dst_vec.ne[0] = ne11; dst_vec.ne[0] = ne11;
dst_vec.ne[1] = 1; dst_vec.ne[1] = 1;
@@ -1988,7 +1992,7 @@ static bool ggml_cuda_mul_mat_id_needs_sync(const ggml_tensor * dst, const int c
return false; return false;
} }
} else if (src0->type == GGML_TYPE_F8_E4M3 && } else if (src0->type == GGML_TYPE_F8_E4M3 &&
ggml_cuda_should_use_mmvf(src0->type, cc, warp_size, src0->ne, src0->nb, dst->ne[2])) { ggml_cuda_should_use_mmvf(src0, cc, warp_size, dst->ne[2])) {
return false; return false;
} else if (src0->type != GGML_TYPE_F8_E4M3 && GGML_CUDA_CC_IS_AMD(cc)) { } else if (src0->type != GGML_TYPE_F8_E4M3 && GGML_CUDA_CC_IS_AMD(cc)) {
return false; return false;
@@ -2030,7 +2034,7 @@ static void ggml_cuda_mul_mat_id(ggml_backend_cuda_context & ctx, ggml_tensor *
return; return;
} }
} else if (src0->type == GGML_TYPE_F8_E4M3 && } else if (src0->type == GGML_TYPE_F8_E4M3 &&
ggml_cuda_should_use_mmvf(src0->type, cc, warp_size, src0->ne, src0->nb, ne2)) { ggml_cuda_should_use_mmvf(src0, cc, warp_size, ne2)) {
ggml_cuda_mul_mat_vec_f(ctx, src0, src1, ids, dst); ggml_cuda_mul_mat_vec_f(ctx, src0, src1, ids, dst);
return; return;
} else { } else {
@@ -3995,7 +3999,7 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
fusion_data.glu_op = ggml_get_glu_op(glu); fusion_data.glu_op = ggml_get_glu_op(glu);
fusion_data.glu_limit = ggml_get_op_params_f32(glu, 3); fusion_data.glu_limit = ggml_get_op_params_f32(glu, 3);
if (ggml_cuda_should_fuse_mul_mat_vec_f(up_n)) { if (ggml_cuda_should_fuse_mul_mat_vec_f(up_n, gate_n)) {
ggml_cuda_mul_mat_vec_f(*cuda_ctx, src0, src1, ids, cgraph->nodes[glu_idx], &fusion_data); ggml_cuda_mul_mat_vec_f(*cuda_ctx, src0, src1, ids, cgraph->nodes[glu_idx], &fusion_data);
fused_mul_mat_vec = true; fused_mul_mat_vec = true;
fused_node_count = n_ops; fused_node_count = n_ops;
@@ -4096,7 +4100,7 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
fusion_data.glu_op = ggml_get_glu_op(glu); fusion_data.glu_op = ggml_get_glu_op(glu);
fusion_data.glu_limit = ggml_get_op_params_f32(glu, 3); fusion_data.glu_limit = ggml_get_op_params_f32(glu, 3);
if (ggml_cuda_should_fuse_mul_mat_vec_f(up_n)) { if (ggml_cuda_should_fuse_mul_mat_vec_f(up_n, gate_n)) {
ggml_cuda_mul_mat_vec_f(*cuda_ctx, src0, src1, ids, cgraph->nodes[glu_idx], &fusion_data); ggml_cuda_mul_mat_vec_f(*cuda_ctx, src0, src1, ids, cgraph->nodes[glu_idx], &fusion_data);
fused_mul_mat_vec = true; fused_mul_mat_vec = true;
fused_node_count = n_ops; fused_node_count = n_ops;
@@ -4152,7 +4156,7 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
const ggml_tensor * src1 = up_n->src[1]; const ggml_tensor * src1 = up_n->src[1];
const ggml_tensor * ids = up_n->src[2]; const ggml_tensor * ids = up_n->src[2];
if (ggml_cuda_should_fuse_mul_mat_vec_f(up_n)) { if (ggml_cuda_should_fuse_mul_mat_vec_f(up_n, gate_n)) {
ggml_cuda_mm_fusion_args_host fusion_data{}; ggml_cuda_mm_fusion_args_host fusion_data{};
fusion_data.gate = gate_n->src[0]; fusion_data.gate = gate_n->src[0];
fusion_data.x_bias = up_bias_tensor; fusion_data.x_bias = up_bias_tensor;
@@ -4195,7 +4199,7 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
const ggml_tensor * src1 = up->src[1]; const ggml_tensor * src1 = up->src[1];
const ggml_tensor * ids = up->src[2]; const ggml_tensor * ids = up->src[2];
if (ggml_cuda_should_fuse_mul_mat_vec_f(up)) { if (ggml_cuda_should_fuse_mul_mat_vec_f(up, gate)) {
ggml_cuda_mm_fusion_args_host fusion_data{}; ggml_cuda_mm_fusion_args_host fusion_data{};
fusion_data.gate = gate->src[0]; fusion_data.gate = gate->src[0];
fusion_data.glu_op = ggml_get_glu_op(glu); fusion_data.glu_op = ggml_get_glu_op(glu);
+6 -5
View File
@@ -927,7 +927,11 @@ void ggml_cuda_op_mul_mat_vec_f(
GGML_UNUSED_VARS(ctx, src1, dst, src1_ddq_i, src1_ncols, src1_padded_row_size); GGML_UNUSED_VARS(ctx, src1, dst, src1_ddq_i, src1_ncols, src1_padded_row_size);
} }
bool ggml_cuda_should_use_mmvf(enum ggml_type type, int cc, int warp_size, const int64_t * src0_ne, const size_t * src0_nb, int64_t ne11) { bool ggml_cuda_should_use_mmvf(const ggml_tensor * src0, int cc, int warp_size, int64_t ne11) {
const ggml_type type = src0->type;
const int64_t * src0_ne = src0->ne;
const size_t * src0_nb = src0->nb;
if (src0_ne[0] % 2 != 0) { if (src0_ne[0] % 2 != 0) {
return false; return false;
} }
@@ -937,12 +941,9 @@ bool ggml_cuda_should_use_mmvf(enum ggml_type type, int cc, int warp_size, const
return false; return false;
} }
// Pointers not aligned to the size of half2/nv_bfloat162/float2 would result in a crash: if (!ggml_cuda_is_aligned(src0, 2*ts)) {
for (size_t i = 1; i < GGML_MAX_DIMS; ++i) {
if (src0_nb[i] % (2*ts) != 0) {
return false; return false;
} }
}
switch (type) { switch (type) {
case GGML_TYPE_F32: case GGML_TYPE_F32:
+1 -1
View File
@@ -11,4 +11,4 @@ void ggml_cuda_op_mul_mat_vec_f(
const char * src1_ddq_i, float * dst_dd_i, const int64_t row_low, const int64_t row_high, const int64_t src1_ncols, const char * src1_ddq_i, float * dst_dd_i, const int64_t row_low, const int64_t row_high, const int64_t src1_ncols,
const int64_t src1_padded_row_size, cudaStream_t stream); const int64_t src1_padded_row_size, cudaStream_t stream);
bool ggml_cuda_should_use_mmvf(enum ggml_type type, int cc, int warp_size, const int64_t * src0_ne, const size_t * src0_nb, int64_t ne11); bool ggml_cuda_should_use_mmvf(const ggml_tensor * src0, int cc, int warp_size, int64_t ne11);
+17 -5
View File
@@ -5181,9 +5181,10 @@ struct test_mul_mat : public test_case {
const bool src_overlap; // a and b are overlapping views of the same tensor const bool src_overlap; // a and b are overlapping views of the same tensor
const int64_t m_v; // rows of a in memory, the batches of a are strided for m_v > m, no view for m_v == 0 const int64_t m_v; // rows of a in memory, the batches of a are strided for m_v > m, no view for m_v == 0
const int64_t pad; // bytes after the m_v rows of each batch of a, so nb[2] of a is not a multiple of nb[1] const int64_t pad; // bytes after the m_v rows of each batch of a, so nb[2] of a is not a multiple of nb[1]
const size_t offset_a; // byte offset of the a view when k_v > k or m_v > m
std::string vars() override { std::string vars() override {
return VARS_TO_STR13(type_a, type_b, m, n, k, bs, nr, per, k_v, o, src_overlap, m_v, pad); return VARS_TO_STR14(type_a, type_b, m, n, k, bs, nr, per, k_v, o, src_overlap, m_v, pad, offset_a);
} }
double max_nmse_err() override { double max_nmse_err() override {
@@ -5217,8 +5218,8 @@ struct test_mul_mat : public test_case {
std::array<int64_t, 2> bs = {10, 10}, std::array<int64_t, 2> bs = {10, 10},
std::array<int64_t, 2> nr = {2, 2}, std::array<int64_t, 2> nr = {2, 2},
std::array<int64_t, 4> per = {0, 1, 2, 3}, std::array<int64_t, 4> per = {0, 1, 2, 3},
int64_t k_v = 0, uint32_t o = 1, bool src_overlap = false, int64_t m_v = 0, int64_t pad = 0) int64_t k_v = 0, uint32_t o = 1, bool src_overlap = false, int64_t m_v = 0, int64_t pad = 0, size_t offset_a = 0)
: type_a(type_a), type_b(type_b), m(m), n(n), k(k), bs(bs), nr(nr), per(per), k_v(k_v), o(o), src_overlap(src_overlap), m_v(m_v), pad(pad) {} : type_a(type_a), type_b(type_b), m(m), n(n), k(k), bs(bs), nr(nr), per(per), k_v(k_v), o(o), src_overlap(src_overlap), m_v(m_v), pad(pad), offset_a(offset_a) {}
ggml_tensor * build_graph(ggml_context * ctx) override { ggml_tensor * build_graph(ggml_context * ctx) override {
// C^T = A * B^T: (k, m) * (k, n) => (m, n) // C^T = A * B^T: (k, m) * (k, n) => (m, n)
@@ -5280,13 +5281,15 @@ struct test_mul_mat : public test_case {
if (k_v != 0) { if (k_v != 0) {
GGML_ASSERT(k_v > k); GGML_ASSERT(k_v > k);
a = ggml_view_4d(ctx, a, k, m, bs[0], bs[1], a->nb[1], a->nb[2], a->nb[3], 0); GGML_ASSERT(offset_a <= ggml_row_size(type_a, k_v) - ggml_row_size(type_a, k));
a = ggml_view_4d(ctx, a, k, m, bs[0], bs[1], a->nb[1], a->nb[2], a->nb[3], offset_a);
b = ggml_view_4d(ctx, b, k, n, bs[0]*nr[0], bs[1]*nr[1], b->nb[1], b->nb[2], b->nb[3], 0); b = ggml_view_4d(ctx, b, k, n, bs[0]*nr[0], bs[1]*nr[1], b->nb[1], b->nb[2], b->nb[3], 0);
} }
if (m_v != 0) { if (m_v != 0) {
GGML_ASSERT(m_v > m); GGML_ASSERT(m_v > m);
GGML_ASSERT(pad < (int64_t) a->nb[1]); GGML_ASSERT(pad < (int64_t) a->nb[1]);
a = ggml_view_4d(ctx, a, k, m, bs[0], bs[1], a->nb[1], m_v*a->nb[1] + pad, a->nb[3], 0); GGML_ASSERT(k_v != 0 || offset_a <= (m_v - m)*a->nb[1]);
a = ggml_view_4d(ctx, a, k, m, bs[0], bs[1], a->nb[1], m_v*a->nb[1] + pad, a->nb[3], k_v == 0 ? offset_a : 0);
} }
ggml_set_name(a, "a"); ggml_set_name(a, "a");
ggml_set_name(b, "b"); ggml_set_name(b, "b");
@@ -10757,6 +10760,15 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F16, GGML_TYPE_F16, 1, 509, 2051, {1, 1}, {1, 1})); test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F16, GGML_TYPE_F16, 1, 509, 2051, {1, 1}, {1, 1}));
test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F8_E4M3, GGML_TYPE_BF16, 16, 4, 3, {1, 1}, {1, 1})); test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F8_E4M3, GGML_TYPE_BF16, 16, 4, 3, {1, 1}, {1, 1}));
test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F8_E4M3, GGML_TYPE_BF16, 16, 4, 256, {1, 1}, {1, 1})); test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F8_E4M3, GGML_TYPE_BF16, 16, 4, 256, {1, 1}, {1, 1}));
// FP8 views with an odd byte offset must use the scalar fallback instead of MMVF.
for (int n : {1, 2}) {
test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F8_E4M3, GGML_TYPE_F32, 2, n, 64, {1, 1}, {1, 1}, {0, 1, 2, 3}, 66, 1, false, 0, 0, 1));
}
// Native FP8 cuBLASLt requires 16-byte alignment, including each batch pointer.
for (size_t offset_a : {1, 2, 4, 16}) {
test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F8_E4M3, GGML_TYPE_F32, 16, 16, 64, {1, 1}, {1, 1}, {0, 1, 2, 3}, 0, 1, false, 17, 0, offset_a));
}
test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F8_E4M3, GGML_TYPE_F32, 16, 16, 64, {2, 1}, {1, 1}, {0, 1, 2, 3}, 0, 1, false, 17, 4));
test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F32, GGML_TYPE_F32, 31, 509, 2051, {1, 1}, {1, 1})); test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F32, GGML_TYPE_F32, 31, 509, 2051, {1, 1}, {1, 1}));
test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F32, GGML_TYPE_F32, 32, 509, 2112, {1, 1}, {1, 1})); test_cases.emplace_back(new test_mul_mat(GGML_TYPE_F32, GGML_TYPE_F32, 32, 509, 2112, {1, 1}, {1, 1}));