test-backend-ops perf -o SSM_CONV on an Arc Pro B70, interleaved A/B against master, 6 reps, us/run: ne_a=[515,3328,1,1] ne_b=[4,3328,1,1] n_t=512 97.68 -> 52.95 1.85x ne_a=[937,8192,1,1] ne_b=[4,8192,1,1] n_t=934 516.16 -> 276.13 1.87x ne_a=[4,3328,1,1] ne_b=[4,3328,1,1] n_t=1 2.73 -> 2.71 flat llama-bench on qwen35 27B Q4_K - Medium (48 of its 64 blocks run ssm_conv), -ngl 99 -fa 1 -ctk f16 -ctv f16, interleaved passes of r=3: -b 2048 -ub 2048 pp2048 1045.1 / 1043.5 / 1043.7 -> 1069.5 / 1066.3 / 1065.9 +2.2% -b 2048 -ub 512 pp2048 771.8 / 772.7 -> 785.5 / 786.6 +1.8% -b 2048 -ub 512 tg128 23.81 / 23.88 -> 23.87 / 23.86 flat
137 lines
4.4 KiB
C++
137 lines
4.4 KiB
C++
#include "ssm_conv.hpp"
|
|
#include "common.hpp"
|
|
|
|
#include <cstdio>
|
|
|
|
using namespace sycl;
|
|
|
|
static void kernel_ssm_conv(
|
|
queue &q,
|
|
const float *src_data,
|
|
const float *weights,
|
|
float *dst_data,
|
|
int d_conv,
|
|
int d_inner,
|
|
int n_t,
|
|
int n_s,
|
|
int ncs __attribute__((unused)),
|
|
int src_stride_inner,
|
|
int src_stride_seq,
|
|
int dst_stride_token,
|
|
int dst_stride_seq
|
|
) {
|
|
const size_t total_work = static_cast<size_t>(d_inner) * static_cast<size_t>(n_t) * static_cast<size_t>(n_s);
|
|
const size_t work_group_size = 256;
|
|
const size_t num_work_groups = (total_work + work_group_size - 1) / work_group_size;
|
|
|
|
const range<1> global_range(num_work_groups * work_group_size);
|
|
const range<1> local_range(work_group_size);
|
|
|
|
q.submit([&](handler &h) {
|
|
h.parallel_for(
|
|
nd_range<1>(global_range, local_range),
|
|
[=](nd_item<1> item) {
|
|
const size_t idx = item.get_global_id(0);
|
|
if (idx >= total_work) {
|
|
return;
|
|
}
|
|
|
|
// src has the tokens of one channel contiguous, dst has the channels of one
|
|
// token contiguous, so either the loads or the store must be strided. Indexing
|
|
// token-fastest coalesces the d_conv loads, which measured faster except for
|
|
// short, cache-resident rows.
|
|
const int token = static_cast<int>(idx % n_t);
|
|
const int channel = static_cast<int>((idx / n_t) % d_inner);
|
|
const int seq = static_cast<int>(idx / (static_cast<size_t>(n_t) * static_cast<size_t>(d_inner)));
|
|
|
|
const float *s = src_data
|
|
+ static_cast<size_t>(seq) * static_cast<size_t>(src_stride_seq)
|
|
+ static_cast<size_t>(channel) * static_cast<size_t>(src_stride_inner)
|
|
+ static_cast<size_t>(token);
|
|
|
|
const float *c = weights + static_cast<size_t>(channel) * static_cast<size_t>(d_conv);
|
|
|
|
float sumf = 0.0f;
|
|
for (int i0 = 0; i0 < d_conv; ++i0) {
|
|
sumf += s[i0] * c[i0];
|
|
}
|
|
|
|
const size_t dst_idx =
|
|
static_cast<size_t>(seq) * static_cast<size_t>(dst_stride_seq) +
|
|
static_cast<size_t>(token) * static_cast<size_t>(dst_stride_token) +
|
|
static_cast<size_t>(channel);
|
|
|
|
dst_data[dst_idx] = sumf;
|
|
}
|
|
);
|
|
});
|
|
}
|
|
|
|
inline void ggml_sycl_op_ssm_conv(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
|
|
ggml_tensor * src0 = dst->src[0];
|
|
ggml_tensor * src1 = dst->src[1];
|
|
|
|
GGML_ASSERT(src0->type == GGML_TYPE_F32);
|
|
GGML_ASSERT(src1->type == GGML_TYPE_F32);
|
|
GGML_ASSERT(dst->type == GGML_TYPE_F32);
|
|
|
|
const int d_conv = src1->ne[0];
|
|
const int ncs = src0->ne[0];
|
|
const int d_inner = src0->ne[1];
|
|
const int n_t = dst->ne[1];
|
|
const int n_s = dst->ne[2];
|
|
|
|
GGML_ASSERT(src0->ne[0] == d_conv - 1 + n_t);
|
|
GGML_ASSERT(src0->ne[1] == d_inner);
|
|
GGML_ASSERT(src1->ne[1] == d_inner);
|
|
|
|
GGML_ASSERT(dst->ne[0] == d_inner);
|
|
GGML_ASSERT(dst->ne[1] == n_t);
|
|
GGML_ASSERT(dst->ne[2] == n_s);
|
|
|
|
GGML_ASSERT(src0->nb[0] == sizeof(float));
|
|
GGML_ASSERT(src1->nb[0] == sizeof(float));
|
|
|
|
GGML_ASSERT(src0->nb[1] == src0->ne[0] * sizeof(float));
|
|
|
|
const int src_stride_inner = ncs;
|
|
const int src_stride_seq = ncs * d_inner;
|
|
const int dst_stride_token = d_inner;
|
|
const int dst_stride_seq = d_inner * n_t;
|
|
|
|
try {
|
|
queue *q = ctx.stream();
|
|
|
|
const float *src_data = static_cast<const float *>(src0->data);
|
|
const float *weights = static_cast<const float *>(src1->data);
|
|
float *dst_data = static_cast<float *>(dst->data);
|
|
|
|
GGML_ASSERT(src_data && weights && dst_data);
|
|
|
|
kernel_ssm_conv(
|
|
*q,
|
|
src_data,
|
|
weights,
|
|
dst_data,
|
|
d_conv,
|
|
d_inner,
|
|
n_t,
|
|
n_s,
|
|
ncs,
|
|
src_stride_inner,
|
|
src_stride_seq,
|
|
dst_stride_token,
|
|
dst_stride_seq
|
|
);
|
|
|
|
} catch (const std::exception &e) {
|
|
std::fprintf(stderr, "[SYCL-SSM_CONV] ERROR: %s\n", e.what());
|
|
throw;
|
|
}
|
|
}
|
|
|
|
void ggml_sycl_ssm_conv(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
|
|
scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/2);
|
|
ggml_sycl_op_ssm_conv(ctx, dst);
|
|
}
|