21#pragma GCC diagnostic push
22#pragma GCC diagnostic ignored "-Wmaybe-uninitialized"
29 const float *cos_cache,
30 const float *sin_cache,
39 head_dim, aligned_head_dim, pos_offset, head_dim, scratch);
43 const float *cos_cache,
44 const float *sin_cache,
55 size_t total = (size_t)num_heads * (
size_t)num_tokens * (size_t)aligned_head_dim;
59 num_heads, num_tokens, head_dim, aligned_head_dim, pos_offset, rotary_dim);
69 const float *cos_cache,
70 const float *sin_cache,
79 if (!scratch_d_out || !scratch_d_x)
return;
81 size_t total = (size_t)num_heads * (
size_t)num_tokens * (size_t)aligned_head_dim;
84 rope_backward(scratch_d_out, scratch_d_x, cos_cache, sin_cache,
85 num_heads, num_tokens, head_dim, aligned_head_dim, pos_offset);
96 const float *cos_cache,
97 const float *sin_cache,
102 int aligned_head_dim,
108 num_tokens, head_dim, aligned_head_dim, pos_offset,
109 head_dim, scratch_q, scratch_k);
114 const float *cos_cache,
115 const float *sin_cache,
120 int aligned_head_dim,
126 if (!q || !k)
return;
129 num_heads, num_tokens, head_dim, aligned_head_dim, pos_offset,
130 rotary_dim, scratch_q);
132 num_kv_heads, num_tokens, head_dim, aligned_head_dim, pos_offset,
133 rotary_dim, scratch_k);
140 const uint16_t *d_k_out,
143 const float *cos_cache,
144 const float *sin_cache,
149 int aligned_head_dim,
151 float *scratch_dq_out,
153 float *scratch_dk_out,
156 if (!d_q_out || !d_k_out || !d_q || !d_k)
return;
159 num_heads, num_tokens, head_dim, aligned_head_dim, pos_offset,
160 scratch_dq_out, scratch_dq);
162 num_kv_heads, num_tokens, head_dim, aligned_head_dim, pos_offset,
163 scratch_dk_out, scratch_dk);
166#pragma GCC diagnostic pop
static void float_tensor_to_bf16(const float *src, uint16_t *dst, size_t count)
static void bf16_tensor_to_float(const uint16_t *src, float *dst, size_t count)
void rope_forward_with_rotary_dim(float *x, const float *cos_cache, const float *sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim)
void rope_backward(const float *d_out, float *d_x, const float *cos_cache, const float *sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset)
void rope_backward_bf16(const uint16_t *d_out, uint16_t *d_x, const float *cos_cache, const float *sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, float *scratch_d_out, float *scratch_d_x)
void rope_backward_qk_bf16(const uint16_t *d_q_out, const uint16_t *d_k_out, uint16_t *d_q, uint16_t *d_k, const float *cos_cache, const float *sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, float *scratch_dq_out, float *scratch_dq, float *scratch_dk_out, float *scratch_dk)
void rope_forward_bf16(uint16_t *x, const float *cos_cache, const float *sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, float *scratch)
void rope_forward_bf16_with_rotary_dim(uint16_t *x, const float *cos_cache, const float *sin_cache, int num_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim, float *scratch)
void rope_forward_qk_bf16(uint16_t *q, uint16_t *k, const float *cos_cache, const float *sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, float *scratch_q, float *scratch_k)
void rope_forward_qk_bf16_with_rotary_dim(uint16_t *q, uint16_t *k, const float *cos_cache, const float *sin_cache, int num_heads, int num_kv_heads, int num_tokens, int head_dim, int aligned_head_dim, int pos_offset, int rotary_dim, float *scratch_q, float *scratch_k)