35#if defined(CK_TARGET_X86)
47 if (!conv_x || !kernel || !out) {
50 if (kernel_size <= 0 || num_channels <= 0 || num_tokens < 0 || num_seqs <= 0) {
54 const size_t seq_width = (size_t)kernel_size - 1u + (
size_t)num_tokens;
55 const size_t conv_seq_stride = (size_t)num_channels * seq_width;
56 const size_t out_seq_stride = (size_t)num_tokens * (
size_t)num_channels;
58 for (
int seq = 0; seq < num_seqs; ++seq) {
59 const float *conv_seq = conv_x + (size_t)seq * conv_seq_stride;
60 float *out_seq = out + (size_t)seq * out_seq_stride;
62 for (
int tok = 0; tok < num_tokens; ++tok) {
63 float *out_tok = out_seq + (size_t)tok * (
size_t)num_channels;
65 for (
int ch = 0; ch < num_channels; ++ch) {
66 const float *conv_row = conv_seq + (size_t)ch * seq_width + (
size_t)tok;
67 const float *kernel_row = kernel + (size_t)ch * (
size_t)kernel_size;
69 for (
int k = 0; k < kernel_size; ++k) {
70 sumf += conv_row[k] * kernel_row[k];
88 if (!d_out || !conv_x || !kernel || !d_conv_x || !d_kernel) {
91 if (kernel_size <= 0 || num_channels <= 0 || num_tokens < 0 || num_seqs <= 0) {
95 const size_t seq_width = (size_t)kernel_size - 1u + (
size_t)num_tokens;
96 const size_t conv_total = (size_t)num_seqs * (
size_t)num_channels * seq_width;
97 const size_t kernel_total = (size_t)num_channels * (
size_t)kernel_size;
98 const size_t conv_seq_stride = (size_t)num_channels * seq_width;
99 const size_t out_seq_stride = (size_t)num_tokens * (
size_t)num_channels;
101 memset(d_conv_x, 0, conv_total *
sizeof(
float));
102 memset(d_kernel, 0, kernel_total *
sizeof(
float));
104 for (
int seq = 0; seq < num_seqs; ++seq) {
105 const float *d_out_seq = d_out + (size_t)seq * out_seq_stride;
106 const float *conv_seq = conv_x + (size_t)seq * conv_seq_stride;
107 float *d_conv_seq = d_conv_x + (size_t)seq * conv_seq_stride;
109 for (
int tok = 0; tok < num_tokens; ++tok) {
110 const float *d_out_tok = d_out_seq + (size_t)tok * (
size_t)num_channels;
112 for (
int ch = 0; ch < num_channels; ++ch) {
113 const float grad = d_out_tok[ch];
114 const float *conv_row = conv_seq + (size_t)ch * seq_width + (
size_t)tok;
115 float *d_conv_row = d_conv_seq + (size_t)ch * seq_width + (
size_t)tok;
116 const float *kernel_row = kernel + (size_t)ch * (
size_t)kernel_size;
117 float *d_kernel_row = d_kernel + (size_t)ch * (
size_t)kernel_size;
119 for (
int k = 0; k < kernel_size; ++k) {
120 d_kernel_row[k] += grad * conv_row[k];
121 d_conv_row[k] += grad * kernel_row[k];
153} ck_ssm_conv1d_llama_args_t;
157 const ck_ssm_conv1d_llama_args_t *args =
158 (
const ck_ssm_conv1d_llama_args_t *)opaque;
159 const int kernel_size = args->kernel_size;
160 const int num_channels = args->num_channels;
161 const int num_tokens = args->num_tokens;
162 const size_t seq_width = (size_t)kernel_size - 1u + (
size_t)num_tokens;
163 const size_t conv_seq_stride = (size_t)num_channels * seq_width;
164 const size_t out_seq_stride = (size_t)num_tokens * (
size_t)num_channels;
166 for (
int seq = 0; seq < args->num_seqs; ++seq) {
167 const float *conv_seq = args->conv_x + (size_t)seq * conv_seq_stride;
168 float *out_seq = args->out + (size_t)seq * out_seq_stride;
170 for (
int ch = begin; ch <
end; ++ch) {
171 const float *kernel_row =
172 args->kernel + (size_t)ch * (
size_t)kernel_size;
173 const float *conv_channel = conv_seq + (size_t)ch * seq_width;
174 for (
int tok = 0; tok < num_tokens; ++tok) {
175 const float *conv_row =
176 conv_channel + (size_t)tok;
178 for (
int k = 0; k < kernel_size; ++k) {
179 volatile float product = conv_row[k] * kernel_row[k];
180 volatile float next = sum + product;
183 out_seq[(size_t)tok * (
size_t)num_channels + (size_t)ch] = sum;
197 if (!conv_x || !kernel || !out || kernel_size <= 0 || num_channels <= 0 ||
198 num_tokens < 0 || num_seqs <= 0) {
201 ck_ssm_conv1d_llama_args_t args = {
202 conv_x, kernel, out, kernel_size, num_channels, num_tokens, num_seqs,
215 if (!conv_x || !kernel || !out || kernel_size <= 0 || num_channels <= 0 ||
216 num_tokens < 0 || num_seqs <= 0) {
219 ck_ssm_conv1d_llama_args_t args = {
220 conv_x, kernel, out, kernel_size, num_channels, num_tokens, num_seqs,
224 const int active = workers < num_channels ? workers : num_channels;
225 if (active > 1 && num_tokens > 1) {
227 pool, active, 0, num_channels, 32,
241 int begin,
int end,
void *opaque)
243 const ck_ssm_conv1d_llama_args_t *args =
244 (
const ck_ssm_conv1d_llama_args_t *)opaque;
245 const int kernel_size = args->kernel_size;
246 const int num_channels = args->num_channels;
247 const int num_tokens = args->num_tokens;
248 const size_t seq_width = (size_t)kernel_size - 1u + (
size_t)num_tokens;
249 const size_t conv_seq_stride = (size_t)num_channels * seq_width;
250 const size_t out_seq_stride = (size_t)num_tokens * (
size_t)num_channels;
252 for (
int seq = 0; seq < args->num_seqs; ++seq) {
253 const float *conv_seq = args->conv_x + (size_t)seq * conv_seq_stride;
254 float *out_seq = args->out + (size_t)seq * out_seq_stride;
256 for (
int ch = begin; ch <
end; ++ch) {
257 const float *kernel_row =
258 args->kernel + (size_t)ch * (
size_t)kernel_size;
259 const float *conv_channel = conv_seq + (size_t)ch * seq_width;
260 for (
int tok = 0; tok < num_tokens; ++tok) {
261 const float *conv_row = conv_channel + (size_t)tok;
263 for (
int k = 0; k < kernel_size; ++k) {
264 sum = fmaf(conv_row[k], kernel_row[k], sum);
266 out_seq[(size_t)tok * (
size_t)num_channels + (size_t)ch] = sum;
280 if (!conv_x || !kernel || !out || kernel_size <= 0 || num_channels <= 0 ||
281 num_tokens < 0 || num_seqs <= 0) {
284 ck_ssm_conv1d_llama_args_t args = {
285 conv_x, kernel, out, kernel_size, num_channels, num_tokens, num_seqs,
289 const int active = workers < num_channels ? workers : num_channels;
290 if (active > 1 && num_tokens > 1) {
292 pool, active, 0, num_channels, 32,
308 conv_x, kernel, out, kernel_size, num_channels, num_tokens, num_seqs);
310 (size_t)num_seqs * (
size_t)num_tokens * (size_t)num_channels;
311 for (
size_t i = 0; i < count; ++i) {
326 ssm_conv1d_backward_ref(d_out, conv_x, kernel, d_conv_x, d_kernel, kernel_size, num_channels, num_tokens, num_seqs);
static uint16_t float_to_bf16(float f)
static float bf16_to_float(uint16_t v)
Persistent pthread thread pool for CK-Engine inference.
void ck_threadpool_parallel_for_n(ck_threadpool_t *pool, int active_threads, int begin, int end, int grain_size, ck_range_fn_t fn, void *args)
ck_threadpool_t * ck_threadpool_global(void)
int ck_threadpool_n_threads(const ck_threadpool_t *pool)
static void ck_ssm_conv1d_llama_channel_range(int begin, int end, void *opaque)
static void ck_ssm_conv1d_llama_fma_channel_range(int begin, int end, void *opaque)
void ssm_conv1d_forward_llama_production_serial(const float *conv_x, const float *kernel, float *out, int kernel_size, int num_channels, int num_tokens, int num_seqs)
void ssm_conv1d_forward_ref(const float *conv_x, const float *kernel, float *out, int kernel_size, int num_channels, int num_tokens, int num_seqs)
void ssm_conv1d_forward_llama_fma(const float *conv_x, const float *kernel, float *out, int kernel_size, int num_channels, int num_tokens, int num_seqs)
void ssm_conv1d_backward_ref(const float *d_out, const float *conv_x, const float *kernel, float *d_conv_x, float *d_kernel, int kernel_size, int num_channels, int num_tokens, int num_seqs)
void ssm_conv1d_forward_pytorch_bf16_storage(const float *conv_x, const float *kernel, float *out, int kernel_size, int num_channels, int num_tokens, int num_seqs)
void ssm_conv1d_forward(const float *conv_x, const float *kernel, float *out, int kernel_size, int num_channels, int num_tokens, int num_seqs)
void ssm_conv1d_forward_llama_production(const float *conv_x, const float *kernel, float *out, int kernel_size, int num_channels, int num_tokens, int num_seqs)
void ssm_conv1d_backward(const float *d_out, const float *conv_x, const float *kernel, float *d_conv_x, float *d_kernel, int kernel_size, int num_channels, int num_tokens, int num_seqs)