* hex-mm: initial support for F32 * F32 -> F32 matmuls * hex-rms-norm: fix src1 stride use in fused rms_norm_mul * hex-ops: clear spad pointers in the ops that clober it This fixes an odd case where fused rms-norm-mul was failing but only in qwen3.5-2B and only at searth op-bath sizes. * hmx-mm: add support for F32 * F32 -> F32 matmul_2d on HMX Decided to use Q4_0 * F32 -> F32 matmul for this. Q4_0 gets dequantized and tiled into F16, and here we quantize and tile F32 into F16. Super simple and pretty efficient. * hmx-mm: route f16 2D matmuls through the same kernel used for all other types * hmx-mm: re-introduce pipelined vs non-pipelined mode that we used to have but is much more generic way This update futher improves matmul performance and at the same time removes most of the redudant logic we had in different paths. * hmx-fa: slighlty improved pipeline simimar to matmul updates * hmx-mm: initial version of MAT_MUL_ID support for HMX * hmx-mm: fixed mxfp4 handling for MUL_MAT_ID * hex-gdn: optimize GATED_DELTA_NET DMA prefetch/double-buff, vectorize everything with HVX, in other words -- the usual :) * hmx-mm: missed one more case where we can use fastmod * hexagon: update DCVS settings for a slight perf bump * hmx-fa: use fastdiv in hmx-flash-attn * hmx-fa: precompute slope values to avoid disrupting the inner loop * hvx-utils/fa: new HVX helpers for powf and logf and using those to speed up FA alibi * hex-ops: fixed a bug in fusion logic that was messing up the order of the src tensors when some srcs are empty * hex-fa: correctly fallback to HVX if we have sinks or the dims are not quite right
48 lines
1.8 KiB
C
48 lines
1.8 KiB
C
#ifndef HVX_FLASH_ATTN_H
|
|
#define HVX_FLASH_ATTN_H
|
|
|
|
#include <math.h>
|
|
#include "hvx-utils.h"
|
|
|
|
// Scalar helper to compute a single ALiBi slope.
|
|
static inline float alibi_slope(uint32_t h, uint32_t n_head_log2, float m0, float m1) {
|
|
return (h < n_head_log2) ? powf(m0, h + 1) : powf(m1, 2 * (h - n_head_log2) + 1);
|
|
}
|
|
|
|
// Vectorized helper to compute 32 ALiBi slopes starting from (kv_head * G).
|
|
static inline HVX_Vector hvx_alibi_slopes(
|
|
uint32_t kv_head,
|
|
uint32_t G,
|
|
uint32_t n_head_log2,
|
|
float m0,
|
|
float m1
|
|
) {
|
|
static const float ramp_32[32] __attribute__((aligned(128))) = {
|
|
0.0f, 1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f, 7.0f,
|
|
8.0f, 9.0f, 10.0f, 11.0f, 12.0f, 13.0f, 14.0f, 15.0f,
|
|
16.0f, 17.0f, 18.0f, 19.0f, 20.0f, 21.0f, 22.0f, 23.0f,
|
|
24.0f, 25.0f, 26.0f, 27.0f, 28.0f, 29.0f, 30.0f, 31.0f
|
|
};
|
|
HVX_Vector v_ramp = hvx_vmem(ramp_32);
|
|
HVX_Vector v_h_base = hvx_vec_splat_f32((float)(kv_head * G));
|
|
HVX_Vector v_h = hvx_vec_add_f32_f32(v_h_base, v_ramp);
|
|
|
|
// Compute exponent_m0: h + 1
|
|
HVX_Vector v_exp_m0 = hvx_vec_add_f32_f32(v_h, hvx_vec_splat_f32(1.0f));
|
|
|
|
// Compute exponent_m1: 2 * (h - n_head_log2) + 1
|
|
HVX_Vector v_n_head_log2 = hvx_vec_splat_f32((float)n_head_log2);
|
|
HVX_Vector v_h_minus = hvx_vec_sub_f32_f32(v_h, v_n_head_log2);
|
|
HVX_Vector v_exp_m1 = hvx_vec_add_f32_f32(hvx_vec_mul_f32_f32(hvx_vec_splat_f32(2.0f), v_h_minus), hvx_vec_splat_f32(1.0f));
|
|
|
|
// Compute powers
|
|
HVX_Vector v_pow_m0 = hvx_vec_pow_const_base_f32(m0, v_exp_m0);
|
|
HVX_Vector v_pow_m1 = hvx_vec_pow_const_base_f32(m1, v_exp_m1);
|
|
|
|
// Select based on h < n_head_log2
|
|
HVX_VectorPred p_cond = Q6_Q_vcmp_gt_VsfVsf(v_n_head_log2, v_h); // v_n_head_log2 > v_h <=> h < n_head_log2
|
|
return Q6_V_vmux_QVV(p_cond, v_pow_m0, v_pow_m1);
|
|
}
|
|
|
|
#endif /* HVX_FLASH_ATTN_H */
|