hexagon: new HMX-optimized GATED_DELTA_NET (#29199)

* hex-gdn: start putting together HMX support for GDN

* hex-gdn: working hmx but not-pipelined and slow for now

* hex-gdn: re-write vtcm layout handling and prep for pipelining

* hex-gdn: starting to pipeline hmx and dmas

* hex-gdn: add hvx threading for most pipeline stages

* hex-gdb: add detailed trace events

* hex-gdn: vectorize expfs and use aligned hvx reads/writes

* hex-gnd: vectorize the rest of expf

* hex-gdn: optimize tail processing (pad partial chunks)

* hex-gdb: avoid float up/down casts in hot loops

* hex-fa: remove float up/down casts from inner loops

* hex-gdn: do exp() in f16 to improve HVX utilization

* hex-gdn: optimize tiler

* hex-hmx: bump hmx-queue to 128 and dispatch all GDN gemms at once

* hex-gdn: further pipeline improvements

* hex-gdn: optimize gdn prep stage

* hex-gdn: yet more tweaks to optimize GND_SOLVE task and pipeline

* hex-gdn: improve accuracy and optmize gdn-prep further

* hex-gdn: fix rebase conflict

* hex-bufs: revert max_bufsize enforcement, it is enough to just enforce max_vmem

* hex-scripts: improved inspect script to avoid false alarms in reg spill detector

* hex-fa: improve inline softmax with in-reg VKQ32 accum

* hex-fa: minor improvement for dma pipeline in hvx kernel

* hex-fa: reduce ddr reads by 20-30% during token gen

* hex-gdn: proper alignment for hvx vtcm spads
This commit is contained in:
Max Krasnyansky
2026-09-21 14:49:52 -07:00
committed by GitHub
parent ff0dbb975e
commit 58367713a6
13 changed files with 2011 additions and 412 deletions
+67 -31
View File
@@ -100,6 +100,7 @@ static bool opt_dma64 = false;
static int opt_mm_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
static int opt_fa_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
static int opt_gdn_select = 2; // 2 = HMX -> HVX, 1 = HVX, 0 = CPU (unsupported)
static int opt_ar_select = 2; // 2 = fused ALLREDUCE+ADD (DMA, default), 1 = unfused ALLREDUCE (DMA), 0 = fallback to CPY+FENCE
// Default PMU events, if profiling with PMU (mode=2) is enabled
@@ -182,6 +183,13 @@ static const char * htp_event_name(uint16_t id) {
case HTP_TRACE_EVT_HVX_FA_Q_PREP: return "HVX_Q_PREP";
case HTP_TRACE_EVT_HVX_FA_K_PREP: return "HVX_K_PREP";
case HTP_TRACE_EVT_HVX_FA_V_PREP: return "HVX_V_PREP";
case HTP_TRACE_EVT_HVX_GDN_PREP: return "HVX_GDN_PREP";
case HTP_TRACE_EVT_HVX_GDN_SOLVE: return "HVX_GDN_SOLVE";
case HTP_TRACE_EVT_HVX_GDN_V_PREP: return "HVX_GDN_V_PREP";
case HTP_TRACE_EVT_HVX_GDN_D_PREP: return "HVX_GDN_D_PREP";
case HTP_TRACE_EVT_HVX_GDN_OUT: return "HVX_GDN_OUT";
case HTP_TRACE_EVT_HVX_GDN_STATE: return "HVX_GDN_STATE";
case HTP_TRACE_EVT_HVX_GDN_REM: return "HVX_GDN_REM";
case HTP_TRACE_EVT_HMX_COMP: return "HMX_COMP";
case HTP_TRACE_EVT_L2FLUSH: return "L2FLUSH";
case HTP_TRACE_EVT_INIT: return "INIT";
@@ -472,7 +480,6 @@ struct ggml_hexagon_session {
uint32_t n_hmx = 0;
uint64_t vtcm_size = 0;
size_t max_vmem = 0;
size_t max_bufsize = 0;
uint32_t fence_seq = 0;
std::atomic<uint64_t> batch_req_seq{0};
@@ -538,7 +545,6 @@ struct ggml_backend_hexagon_device_context {
int dev_id;
ggml_hexagon_device_config config;
ggml_backend_dev_t dev = nullptr;
size_t max_bufsize = 0;
ggml_backend_buffer_type buffer_type = {};
ggml_backend_buffer_type host_buffer_type = {};
@@ -554,9 +560,6 @@ struct ggml_backend_hexagon_device_context {
ggml_hexagon_session * session() {
if (!sess) {
sess = std::make_unique<ggml_hexagon_session>(config, dev);
if (max_bufsize > sess->max_vmem) {
max_bufsize = sess->max_vmem;
}
}
return sess.get();
}
@@ -2076,11 +2079,6 @@ static const char * ggml_backend_hexagon_buffer_type_name(ggml_backend_buffer_ty
static ggml_backend_buffer_t ggml_backend_hexagon_buffer_type_alloc_buffer(
ggml_backend_buffer_type_t buffer_type, size_t size) {
auto dev_ctx = static_cast<ggml_backend_hexagon_buffer_type_context *>(buffer_type->context)->dev_ctx;
if (size > dev_ctx->max_bufsize) {
GGML_LOG_ERROR("ggml-hex: %s buffer size %zu exceeds max_bufsize %zu\n",
dev_ctx->c_name(), size, dev_ctx->max_bufsize);
return nullptr;
}
auto sess = dev_ctx->session();
if (sess && sess->max_vmem && size > sess->max_vmem) {
GGML_LOG_ERROR("ggml-hex: %s buffer size %zu exceeds max_vmem %zu\n",
@@ -2099,11 +2097,6 @@ static ggml_backend_buffer_t ggml_backend_hexagon_buffer_type_alloc_buffer(
static ggml_backend_buffer_t ggml_backend_hexagon_host_buffer_type_alloc_buffer(
ggml_backend_buffer_type_t buffer_type, size_t size) {
auto dev_ctx = static_cast<ggml_backend_hexagon_buffer_type_context *>(buffer_type->context)->dev_ctx;
if (size > dev_ctx->max_bufsize) {
GGML_LOG_ERROR("ggml-hex: %s host buffer size %zu exceeds max_bufsize %zu\n",
dev_ctx->c_name(), size, dev_ctx->max_bufsize);
return nullptr;
}
auto sess = dev_ctx->session();
if (sess && sess->max_vmem && size > sess->max_vmem) {
GGML_LOG_ERROR("ggml-hex: %s host buffer size %zu exceeds max_vmem %zu\n",
@@ -2138,10 +2131,8 @@ static size_t ggml_backend_hexagon_buffer_type_get_alloc_size(ggml_backend_buffe
}
static size_t ggml_backend_hexagon_buffer_type_get_max_size(ggml_backend_buffer_type_t buft) {
auto * context = static_cast<ggml_backend_hexagon_buffer_type_context *>(buft->context);
auto dev_ctx = context->dev_ctx;
dev_ctx->session();
return dev_ctx->max_bufsize;
return opt_mbuf;
GGML_UNUSED(buft);
}
static bool ggml_backend_hexagon_buffer_type_is_host(ggml_backend_buffer_type_t buft) {
@@ -2173,7 +2164,7 @@ static ggml_backend_buffer_type_i ggml_backend_hexagon_host_buffer_type_interfac
};
ggml_backend_hexagon_device_context::ggml_backend_hexagon_device_context(int dev_id, const ggml_hexagon_device_config & config, ggml_backend_dev_t dev)
: dev_id(dev_id), config(config), dev(dev), max_bufsize(opt_mbuf) {
: dev_id(dev_id), config(config), dev(dev) {
buffer_type.device = dev;
buffer_type.iface = ggml_backend_hexagon_buffer_type_interface;
buffer_type.context = new ggml_backend_hexagon_buffer_type_context(config.name, this);
@@ -3927,7 +3918,6 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
this->valid_handle = true;
// Query HW info and resolve session options
this->max_bufsize = opt_mbuf;
{
unsigned int hw_n_threads = 0;
unsigned int hw_n_hvx = 0;
@@ -4340,6 +4330,10 @@ static bool ggml_hexagon_supported_flash_attn_ext(const struct ggml_hexagon_sess
}
static bool ggml_hexagon_supported_gated_delta_net(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
if (opt_gdn_select < 1) {
return false;
}
const struct ggml_tensor * q = op->src[0];
const struct ggml_tensor * k = op->src[1];
const struct ggml_tensor * v = op->src[2];
@@ -4387,10 +4381,26 @@ static bool ggml_hexagon_supported_gated_delta_net(const struct ggml_hexagon_ses
const uint32_t total_rows = (uint32_t) (H * n_seqs);
const uint32_t n_threads = (std::min)((uint32_t) sess->n_threads, total_rows);
struct htp_gdn_vtcm_layout layout;
htp_gdn_vtcm_layout_build(&layout, (uint32_t) S_v, n_threads ? n_threads : 1);
if (layout.total_bytes > sess->vtcm_size) {
return false;
const bool can_use_hmx = (opt_gdn_select >= 2) &&
(sess->n_hmx > 0) &&
(S_v % 64 == 0) &&
(n_tokens >= HTP_GDN_MIN_TOKENS) &&
(g->ne[0] == 1) &&
(K == 1);
if (can_use_hmx) {
struct htp_gdn_hmx_vtcm_layout layout;
uint32_t n_heads_batch = 0;
if (!htp_gdn_hmx_solve_layout(&layout, (uint32_t) S_v, HTP_GDN_CHUNK_SIZE, total_rows, sess->vtcm_size, n_threads, true, &n_heads_batch)) {
return false;
}
} else {
struct htp_gdn_vtcm_layout layout;
htp_gdn_vtcm_layout_build(&layout, (uint32_t) S_v, n_threads);
if (layout.total_bytes > sess->vtcm_size) {
return false;
}
}
return true;
@@ -5206,10 +5216,37 @@ static void ggml_hexagon_precompute_gated_delta_net_params(
const uint32_t total_rows = H * n_seqs;
const uint32_t n_threads = (std::min)((uint32_t) sess->n_threads, total_rows);
struct htp_gdn_vtcm_layout layout;
htp_gdn_vtcm_layout_build(&layout, S_v, n_threads ? n_threads : 1);
const bool can_use_hmx = (opt_gdn_select >= 2) &&
(sess->n_hmx > 0) &&
(S_v % 64 == 0) &&
(n_tokens >= HTP_GDN_MIN_TOKENS) &&
(g->ne[0] == 1) &&
(K == 1);
kparams->n_threads = n_threads ? n_threads : 1;
struct htp_gdn_hmx_vtcm_layout hmx_layout;
struct htp_gdn_vtcm_layout hvx_layout;
uint32_t n_heads_batch = 1;
if (can_use_hmx && htp_gdn_hmx_solve_layout(&hmx_layout, S_v, HTP_GDN_CHUNK_SIZE, total_rows, sess->vtcm_size, n_threads, true, &n_heads_batch)) {
kparams->kernel_type = HTP_GDN_KERNEL_HMX_CHUNKED;
kparams->pipeline = hmx_layout.pipeline ? 1 : 0;
kparams->chunk_size = HTP_GDN_CHUNK_SIZE;
kparams->n_chunks = (n_tokens + HTP_GDN_CHUNK_SIZE - 1) / HTP_GDN_CHUNK_SIZE;
kparams->n_heads_batch = (uint16_t) n_heads_batch;
kparams->vtcm_size = (uint32_t) hmx_layout.total_bytes;
kparams->state_aligned = (uint32_t) hmx_layout.state_f32_bytes;
kparams->vtcm_per_thread = (uint32_t) (hmx_layout.total_bytes / (n_threads > 0 ? n_threads : 1));
} else {
htp_gdn_vtcm_layout_build(&hvx_layout, S_v, n_threads);
kparams->kernel_type = HTP_GDN_KERNEL_HVX_RECURRENT;
kparams->pipeline = 0;
kparams->n_heads_batch = 1;
kparams->state_aligned = (uint32_t) hvx_layout.state_aligned;
kparams->vtcm_per_thread = (uint32_t) hvx_layout.bytes_per_thread;
kparams->vtcm_size = (uint32_t) hvx_layout.total_bytes;
}
kparams->n_threads = n_threads;
kparams->S_v = S_v;
kparams->H = H;
kparams->n_tokens = n_tokens;
@@ -5218,9 +5255,6 @@ static void ggml_hexagon_precompute_gated_delta_net_params(
kparams->total_rows = total_rows;
kparams->rows_per_thread = (total_rows + kparams->n_threads - 1) / kparams->n_threads;
kparams->kda = (g->ne[0] == S_v) ? 1 : 0;
kparams->state_aligned = (uint32_t) layout.state_aligned;
kparams->vtcm_per_thread = (uint32_t) layout.bytes_per_thread;
kparams->vtcm_size = (uint32_t) layout.total_bytes;
kparams->state_seq_stride = (uint32_t) (state->nb[3] / sizeof(float));
kparams->state_size_per_snap = S_v * S_v * H * n_seqs;
kparams->scale = 1.0f / sqrtf((float) S_v);
@@ -7731,6 +7765,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
const char * str_nhmx = getenv("GGML_HEXAGON_NHMX");
const char * str_mm_select = getenv("GGML_HEXAGON_MM_SELECT");
const char * str_fa_select = getenv("GGML_HEXAGON_FA_SELECT");
const char * str_gdn_select = getenv("GGML_HEXAGON_GDN_SELECT");
const char * str_ar_select = getenv("GGML_HEXAGON_AR_SELECT");
const char * str_ndev = getenv("GGML_HEXAGON_NDEV");
const char * str_arch = getenv("GGML_HEXAGON_ARCH");
@@ -7783,6 +7818,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
opt_nhmx = str_nhmx ? atoi(str_nhmx) : opt_nhmx;
opt_mm_select = str_mm_select ? atoi(str_mm_select) : opt_mm_select;
opt_fa_select = str_fa_select ? atoi(str_fa_select) : opt_fa_select;
opt_gdn_select = str_gdn_select ? atoi(str_gdn_select) : opt_gdn_select;
opt_ar_select = str_ar_select ? atoi(str_ar_select) : opt_ar_select;
opt_mbuf = str_mbuf ? strtoul(str_mbuf, NULL, 0) * MiB : opt_mbuf;
opt_vmem = str_vmem ? strtoul(str_vmem, NULL, 0) * MiB : opt_vmem;
+3 -1
View File
@@ -358,7 +358,9 @@ struct htp_opformat {
snprintf(str, max_size, "k%d nth %d vtcm %d", (int) kparams->kernel_id, (int) kparams->n_threads, (int) kparams->vtcm_size);
} else if (node.opcode == HTP_OP_GATED_DELTA_NET) {
const auto * kparams = (const struct htp_gdn_kernel_params *) node.kernel_params;
snprintf(str, max_size, "%s vtcm %u",
const char * path = (kparams->kernel_type == HTP_GDN_KERNEL_HMX_CHUNKED) ? "hmx-chunked" : "hvx-recurrent";
snprintf(str, max_size, "%s-%s vtcm %u",
path,
kparams->kda ? "kda" : "scalar",
(unsigned int) (kparams->vtcm_size ? kparams->vtcm_size : kparams->vtcm_per_thread * kparams->n_threads));
} else if (node.opcode == HTP_OP_MUL || node.opcode == HTP_OP_ADD || node.opcode == HTP_OP_ADD_ID ||
+273 -298
View File
@@ -55,6 +55,7 @@ struct htp_fa_context {
float scale;
float max_bias;
bool has_softcap;
__fp16 logit_softcap;
uint32_t n_head_log2;
@@ -103,6 +104,7 @@ struct hmx_fa_context {
// Op parameters
__fp16 scale;
float max_bias;
bool has_softcap;
__fp16 logit_softcap;
uint32_t n_head_log2;
float m0, m1;
@@ -234,7 +236,10 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
dma_cache m_cache;
dma_cache_init(&m_cache, spad_m, factx->size_m_block, HVX_FA_DMA_CACHE_SIZE);
for (uint32_t ir = ir0; ir < ir1; ++ir) {
const size_t size_vkq_acc_single = hex_round_up(DV * sizeof(float), 128);
uint32_t ir = ir0;
while (ir < ir1) {
const uint32_t iq3 = fastdiv(ir, &factx->src0_div21);
const uint32_t iq2 = fastdiv(ir - iq3*neq2*neq1, &factx->src0_div1);
const uint32_t iq1 = (ir - iq3*neq2*neq1 - iq2 * neq1);
@@ -245,6 +250,59 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
const uint32_t iv3 = fastdiv(iq3, &factx->broadcast_rv3);
const uint32_t iv2 = fastdiv(iq2, &factx->broadcast_rv2);
uint32_t G_local = 1;
if (neq1 == 1 && (mask == NULL || mask->ne[2] == 1)) {
while (ir + G_local < ir1 && G_local < FA_HVX_G_MAX) {
const uint32_t next_ir = ir + G_local;
const uint32_t next_iq3 = fastdiv(next_ir, &factx->src0_div21);
const uint32_t next_iq2 = fastdiv(next_ir - next_iq3*neq2*neq1, &factx->src0_div1);
const uint32_t next_iq1 = (next_ir - next_iq3*neq2*neq1 - next_iq2 * neq1);
const uint32_t next_ik3 = fastdiv(next_iq3, &factx->broadcast_rk3);
const uint32_t next_ik2 = fastdiv(next_iq2, &factx->broadcast_rk2);
const uint32_t next_iv3 = fastdiv(next_iq3, &factx->broadcast_rv3);
const uint32_t next_iv2 = fastdiv(next_iq2, &factx->broadcast_rv2);
if (next_ik2 != ik2 || next_ik3 != ik3 || next_iv2 != iv2 || next_iv3 != iv3 || next_iq1 != iq1 || next_iq3 != iq3) {
break;
}
G_local++;
}
}
uint32_t heads[FA_HVX_G_MAX];
HVX_Vector slope_vecs[FA_HVX_G_MAX] __attribute__((aligned(128)));
HVX_Vector S_vec[FA_HVX_G_MAX] __attribute__((aligned(128)));
HVX_Vector M_vec[FA_HVX_G_MAX] __attribute__((aligned(128)));
uint8_t * q_ptrs[FA_HVX_G_MAX];
float * vkq_ptrs[FA_HVX_G_MAX];
for (uint32_t g = 0; g < G_local; ++g) {
const uint32_t r = ir + g;
const uint32_t r_iq3 = fastdiv(r, &factx->src0_div21);
const uint32_t r_iq2 = fastdiv(r - r_iq3*neq2*neq1, &factx->src0_div1);
const uint32_t r_iq1 = (r - r_iq3*neq2*neq1 - r_iq2 * neq1);
heads[g] = r_iq2;
const __fp16 slope = factx->slopes[r_iq2];
slope_vecs[g] = hvx_vec_splat_f16(slope);
S_vec[g] = hvx_vec_splat_f32(0.0f);
M_vec[g] = hvx_vec_splat_f32(HTP_FA_M_INITIAL_VAL);
uint8_t * q_dst = spad_q + g * factx->size_q_row_padded;
q_ptrs[g] = q_dst;
float * vkq_dst = (float *)(spad_a + g * size_vkq_acc_single);
vkq_ptrs[g] = vkq_dst;
hvx_splat_f32_a((uint8_t *) vkq_dst, 0, DV);
// Fetch Q row g
const dma_addr_t q_row_ptr = q->data + r_iq1*nbq1 + r_iq2*nbq2 + r_iq3*nbq3;
dma_queue_push(dma_q, dma_make_data(q_dst, q_row_ptr), factx->size_q_row_padded, nbq1, size_q_row, 1);
}
dma_addr_t mp_base = 0;
if (mask) {
const uint32_t im2 = fastmodulo(iq2, mask->ne[2], &factx->src3_div2);
@@ -252,116 +310,44 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
mp_base = mask->data + iq1*mask->nb[1] + im2*mask->nb[2] + im3*mask->nb[3];
}
// Precalculate next row variables if there is a next row
bool has_next_ir = (ir + 1 < ir1);
uint32_t next_ik2 = 0, next_ik3 = 0, next_iv2 = 0, next_iv3 = 0;
dma_addr_t next_q_row_ptr = 0;
dma_addr_t next_mp_base = 0;
// Prefetch first two blocks
for (uint32_t ib = 0; ib < MIN(factx->n_blocks, 2); ++ib) {
const uint32_t ic_start = ib * FLASH_ATTN_BLOCK_SIZE;
const uint32_t current_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start);
dma_addr_t next_k_src0 = 0;
dma_addr_t next_v_src0 = 0;
dma_addr_t next_m_src0 = 0;
uint32_t next_block_size0 = 0;
// K
const dma_addr_t k_src = k->data + ic_start*nbk1 + ik2*nbk2 + ik3*nbk3;
uint8_t * k_dst = spad_k + (ib % 2) * factx->size_k_block;
dma_queue_push(dma_q, dma_make_data(k_dst, k_src), factx->size_k_row_padded, nbk1, size_k_row, current_block_size);
dma_addr_t next_k_src1 = 0;
dma_addr_t next_v_src1 = 0;
dma_addr_t next_m_src1 = 0;
uint32_t next_block_size1 = 0;
if (has_next_ir) {
const uint32_t next_ir = ir + 1;
const uint32_t next_iq3 = fastdiv(next_ir, &factx->src0_div21);
const uint32_t next_iq2 = fastdiv(next_ir - next_iq3*neq2*neq1, &factx->src0_div1);
const uint32_t next_iq1 = (next_ir - next_iq3*neq2*neq1 - next_iq2 * neq1);
next_ik3 = fastdiv(next_iq3, &factx->broadcast_rk3);
next_ik2 = fastdiv(next_iq2, &factx->broadcast_rk2);
next_iv3 = fastdiv(next_iq3, &factx->broadcast_rv3);
next_iv2 = fastdiv(next_iq2, &factx->broadcast_rv2);
next_q_row_ptr = q->data + next_iq1*nbq1 + next_iq2*nbq2 + next_iq3*nbq3;
// V
const dma_addr_t v_src = v->data + ic_start*nbv1 + iv2*nbv2 + iv3*nbv3;
uint8_t * v_dst = spad_v + (ib % 2) * factx->size_v_block;
dma_queue_push(dma_q, dma_make_data(v_dst, v_src), factx->size_v_row_padded, nbv1, size_v_row, current_block_size);
// Mask
if (mask) {
const uint32_t next_im2 = fastmodulo(next_iq2, mask->ne[2], &factx->src3_div2);
const uint32_t next_im3 = fastmodulo(next_iq3, mask->ne[3], &factx->src3_div3);
next_mp_base = mask->data + next_iq1*mask->nb[1] + next_im2*mask->nb[2] + next_im3*mask->nb[3];
}
// Precalculate next K/V block 0 source pointers
{
const uint32_t ic_start = 0;
next_block_size0 = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start);
next_k_src0 = k->data + ic_start*nbk1 + next_ik2*nbk2 + next_ik3*nbk3;
next_v_src0 = v->data + ic_start*nbv1 + next_iv2*nbv2 + next_iv3*nbv3;
if (mask) {
next_m_src0 = next_mp_base + ic_start * sizeof(__fp16);
}
}
// Precalculate next K/V block 1 source pointers (if n_blocks > 1)
if (factx->n_blocks > 1) {
const uint32_t ic_start = 1 * FLASH_ATTN_BLOCK_SIZE;
next_block_size1 = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start);
next_k_src1 = k->data + ic_start*nbk1 + next_ik2*nbk2 + next_ik3*nbk3;
next_v_src1 = v->data + ic_start*nbv1 + next_iv2*nbv2 + next_iv3*nbv3;
if (mask) {
next_m_src1 = next_mp_base + ic_start * sizeof(__fp16);
}
const dma_addr_t m_src = mp_base + ic_start * sizeof(__fp16);
dma_cache_push(dma_q, &m_cache, m_src, current_block_size * 2, current_block_size * 2, current_block_size * 2, 1);
}
}
if (ir == ir0) {
// Fetch Q row
const dma_addr_t q_row_ptr = q->data + iq1*nbq1 + iq2*nbq2 + iq3*nbq3;
dma_queue_push(dma_q, dma_make_data(spad_q, q_row_ptr), factx->size_q_row_padded, nbq1, size_q_row, 1);
// Prefetch first two blocks
for (uint32_t ib = 0; ib < MIN(factx->n_blocks, 2); ++ib) {
const uint32_t ic_start = ib * FLASH_ATTN_BLOCK_SIZE;
const uint32_t current_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start);
// K
const dma_addr_t k_src = k->data + ic_start*nbk1 + ik2*nbk2 + ik3*nbk3;
uint8_t * k_dst = spad_k + (ib % 2) * factx->size_k_block;
dma_queue_push(dma_q, dma_make_data(k_dst, k_src), factx->size_k_row_padded, nbk1, size_k_row, current_block_size);
// V
const dma_addr_t v_src = v->data + ic_start*nbv1 + iv2*nbv2 + iv3*nbv3;
uint8_t * v_dst = spad_v + (ib % 2) * factx->size_v_block;
dma_queue_push(dma_q, dma_make_data(v_dst, v_src), factx->size_v_row_padded, nbv1, size_v_row, current_block_size);
// Mask
if (mask) {
const dma_addr_t m_src = mp_base + ic_start * sizeof(__fp16);
// Mask is 1D contiguous for this row
dma_cache_push(dma_q, &m_cache, m_src, current_block_size * 2, current_block_size * 2, current_block_size * 2, 1);
}
// Pop all Q rows
for (uint32_t g = 0; g < G_local; ++g) {
uint8_t * q_ptr_vtcm = (void *) dma_queue_pop(dma_q).dst;
if (factx->is_q_fp32) {
hvx_copy_f16_f32_aa(q_ptr_vtcm, q_ptr_vtcm, DK);
}
}
const uint32_t h = iq2; // head index
const __fp16 slope = factx->slopes[h];
HVX_Vector S_vec = hvx_vec_splat_f32(0.0f);
HVX_Vector M_vec = hvx_vec_splat_f32(HTP_FA_M_INITIAL_VAL);
// Clear accumulator
hvx_splat_f32_a(spad_a, 0, DV);
float * VKQ32 = (float *) (spad_a + 0);
uint8_t * q_ptr_vtcm = (void *) dma_queue_pop(dma_q).dst;
if (factx->is_q_fp32) {
hvx_copy_f16_f32_aa(q_ptr_vtcm, q_ptr_vtcm, DK); // inplace convert f32 to f16
}
const HVX_Vector slope_vec = hvx_vec_splat_f16(slope);
const HVX_Vector v_neg_inf = Q6_Vh_vsplat_R(0xfbff);
const HVX_Vector v_cap = (factx->logit_softcap != 0.0f) ? hvx_vec_splat_f16(factx->logit_softcap) : Q6_V_vzero();
const bool has_softcap = factx->has_softcap;
const HVX_Vector v_cap = has_softcap ? hvx_vec_splat_f16(factx->logit_softcap) : Q6_V_vzero();
const HVX_Vector vinf = Q6_Vh_vsplat_R(0xFC00);
const HVX_Vector vmin = Q6_Vh_vsplat_R(0xFBFF);
const HVX_Vector v_log2e = hvx_vec_splat_f16(EXP_LOG2E_F);
const uint32_t stride_v2 = factx->size_v_row_padded * 2;
for (uint32_t ib = 0; ib < factx->n_blocks; ++ib) {
const uint32_t ic_start = ib * FLASH_ATTN_BLOCK_SIZE;
const uint32_t current_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start);
@@ -388,235 +374,222 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_V_PREP, ir);
}
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_QK, ir);
for (uint32_t g = 0; g < G_local; ++g) {
const uint32_t head_ir = ir + g;
uint8_t * q_ptr_vtcm = q_ptrs[g];
float * VKQ32 = vkq_ptrs[g];
const HVX_Vector slope_vec = slope_vecs[g];
// Inner loop processing the block from VTCM
// 1. Compute scores (64 elements FP16)
HVX_Vector scores_f16 = Q6_V_vzero();
if (current_block_size > 0) {
HVX_Vector scores0 = hvx_dot_f16_f16_aa_rx32(q_ptr_vtcm, k_base, factx->size_k_row_padded, DK, factx->scale);
HVX_Vector scores1 = (current_block_size > 32) ? hvx_dot_f16_f16_aa_rx32(q_ptr_vtcm, k_base + 32 * factx->size_k_row_padded, factx->size_k_row_padded, DK, factx->scale) : Q6_V_vzero();
scores_f16 = hvx_vec_f32_to_f16(scores0, scores1);
}
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_QK, head_ir);
// 2. Softcap (in FP16)
if (factx->logit_softcap != 0.0f) {
scores_f16 = hvx_vec_tanh_f16(scores_f16);
scores_f16 = hvx_vec_mul_f16_f16(scores_f16, v_cap);
}
HVX_VectorPred q_tail_keep = Q6_Q_vsetq2_R(current_block_size * sizeof(__fp16));
// 3. Mask (in FP16)
if (mask) {
HVX_Vector m_vals_f16 = *(const HVX_UVector *) m_base;
HVX_VectorPred is_inf = Q6_Q_vcmp_eq_VhVh(m_vals_f16, vinf);
m_vals_f16 = Q6_V_vmux_QVV(is_inf, vmin, m_vals_f16);
HVX_Vector m_scaled = hvx_vec_mul_f16_f16(m_vals_f16, slope_vec);
scores_f16 = Q6_V_vmux_QVV(q_tail_keep, hvx_vec_add_f16_f16(scores_f16, m_scaled), v_neg_inf);
} else {
scores_f16 = Q6_V_vmux_QVV(q_tail_keep, scores_f16, v_neg_inf);
}
// Compute block max in FP16
HVX_Vector v_max_f16 = hvx_vec_reduce_max_f16(scores_f16);
HVX_Vector v_max = Q6_V_lo_W(hvx_vec_f16_to_f32(v_max_f16)); // splat block max in FP32
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_QK, ir);
if (ib + 1 == factx->n_blocks && has_next_ir) {
// Queue next row's Q row!
dma_queue_push(dma_q, dma_make_data(spad_q, next_q_row_ptr), factx->size_q_row_padded, nbq1, size_q_row, 1);
if (factx->n_blocks % 2 == 0) {
// Queue next row's block 0 (into buffer slot 0)
uint8_t * k_dst = spad_k + 0 * factx->size_k_block;
uint8_t * v_dst = spad_v + 0 * factx->size_v_block;
// K (block 0 of next row)
dma_queue_push(dma_q, dma_make_data(k_dst, next_k_src0), factx->size_k_row_padded, nbk1, size_k_row, next_block_size0);
// V (block 0 of next row)
dma_queue_push(dma_q, dma_make_data(v_dst, next_v_src0), factx->size_v_row_padded, nbv1, size_v_row, next_block_size0);
// Mask (block 0 of next row)
if (mask) {
dma_cache_push(dma_q, &m_cache, next_m_src0, next_block_size0 * 2, next_block_size0 * 2, next_block_size0 * 2, 1);
}
HVX_Vector scores_f16 = Q6_V_vzero();
if (current_block_size > 0) {
HVX_Vector scores0 = hvx_dot_f16_f16_aa_rx32(q_ptr_vtcm, k_base, factx->size_k_row_padded, DK, factx->scale);
HVX_Vector scores1 = (current_block_size > 32) ? hvx_dot_f16_f16_aa_rx32(q_ptr_vtcm, k_base + 32 * factx->size_k_row_padded, factx->size_k_row_padded, DK, factx->scale) : Q6_V_vzero();
scores_f16 = hvx_vec_f32_to_f16(scores0, scores1);
}
}
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_SFM, ir);
{
// 4. Online Softmax Update
HVX_Vector M_new_vec = Q6_Vsf_vmax_VsfVsf(v_max, M_vec);
HVX_Vector diff_vec = HVX_OP_SUB_F32(M_vec, M_new_vec);
HVX_Vector diff_f16 = hvx_vec_f32_to_f16(diff_vec, diff_vec);
HVX_Vector diff_base2 = hvx_vec_mul_f16_f16(diff_f16, v_log2e);
HVX_Vector ms_f16 = hvx_vec_exp2_f16(diff_base2);
HVX_Vector ms_vec = Q6_V_lo_W(hvx_vec_f16_to_f32(ms_f16));
M_vec = M_new_vec;
hvx_scale_vec_f32_aa((uint8_t *) VKQ32, (const uint8_t *) VKQ32, DV, ms_vec);
// Compute P = exp2((S - M) * log2(e)) in FP16
HVX_Vector v_m_vec_f16 = hvx_vec_f32_to_f16(M_vec, M_vec);
HVX_Vector v_s_minus_m = Q6_Vqf16_vsub_VhfVhf(scores_f16, v_m_vec_f16);
HVX_Vector v_s_minus_m_base2 = hvx_vec_mul_f16_f16(Q6_Vhf_equals_Vqf16(v_s_minus_m), v_log2e);
HVX_Vector P = hvx_vec_exp2_f16(v_s_minus_m_base2);
P = Q6_V_vmux_QVV(q_tail_keep, P, Q6_V_vzero());
// Convert P to FP32 to update the running sum S_vec
HVX_VectorPair P_pair = hvx_vec_f16_to_f32(P);
HVX_Vector P0 = Q6_V_lo_W(P_pair);
HVX_Vector P1 = Q6_V_hi_W(P_pair);
HVX_Vector p_sum_vec = hvx_vec_reduce_sum_f32(HVX_OP_ADD_F32(P0, P1));
S_vec = HVX_OP_ADD_F32(HVX_OP_MUL_F32(S_vec, ms_vec), p_sum_vec);
// 5. Accumulate V (F16 * F16 -> F32 accumulator)
const uint8_t * v_ptr = v_base;
for (uint32_t j = 0; j < current_block_size; j += 2) {
if (j + 1 == current_block_size) {
HVX_Vector S0 = hvx_vec_repl_f16(Q6_V_vror_VR(P, j * 2));
hvx_mad_f32_f16_aa_vec(VKQ32, v_ptr, S0, DV);
break;
}
HVX_Vector S0 = hvx_vec_repl_f16(Q6_V_vror_VR(P, j * 2));
HVX_Vector S1 = hvx_vec_repl_f16(Q6_V_vror_VR(P, (j + 1) * 2));
hvx_mad_f32_f16_aa_rx2_vec(VKQ32, v_ptr, v_ptr + factx->size_v_row_padded, S0, S1, DV);
v_ptr += stride_v2;
if (has_softcap) {
scores_f16 = hvx_vec_tanh_f16(scores_f16);
scores_f16 = hvx_vec_mul_f16_f16(scores_f16, v_cap);
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_SFM, ir);
// Issue DMA for next+1 block (if exists)
if (ib + 2 < factx->n_blocks) {
const uint32_t next_ib = ib + 2;
const uint32_t next_ic_start = next_ib * FLASH_ATTN_BLOCK_SIZE;
const uint32_t next_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - next_ic_start);
HVX_VectorPred q_tail_keep = Q6_Q_vsetq2_R(current_block_size * sizeof(__fp16));
// K
const dma_addr_t k_src = k->data + next_ic_start*nbk1 + ik2*nbk2 + ik3*nbk3;
dma_queue_push(dma_q, dma_make_data(k_base, k_src), factx->size_k_row_padded, nbk1, size_k_row, next_block_size);
// V
const dma_addr_t v_src = v->data + next_ic_start*nbv1 + iv2*nbv2 + iv3*nbv3;
dma_queue_push(dma_q, dma_make_data(v_base, v_src), factx->size_v_row_padded, nbv1, size_v_row, next_block_size);
// Mask
if (mask) {
const dma_addr_t m_src = mp_base + next_ic_start * sizeof(__fp16);
dma_cache_push(dma_q, &m_cache, m_src, next_block_size * 2, next_block_size * 2, next_block_size * 2, 1);
HVX_Vector m_vals_f16 = *(const HVX_UVector *) m_base;
HVX_VectorPred is_inf = Q6_Q_vcmp_eq_VhVh(m_vals_f16, vinf);
m_vals_f16 = Q6_V_vmux_QVV(is_inf, vmin, m_vals_f16);
HVX_Vector m_scaled = hvx_vec_mul_f16_f16(m_vals_f16, slope_vec);
scores_f16 = Q6_V_vmux_QVV(q_tail_keep, hvx_vec_add_f16_f16(scores_f16, m_scaled), v_neg_inf);
} else {
scores_f16 = Q6_V_vmux_QVV(q_tail_keep, scores_f16, v_neg_inf);
}
}
}
if (has_next_ir) {
if (factx->n_blocks % 2 == 0) {
// Queue next row's block 1 (into buffer slot 1, if n_blocks > 1)
if (factx->n_blocks > 1) {
uint8_t * k_dst = spad_k + 1 * factx->size_k_block;
uint8_t * v_dst = spad_v + 1 * factx->size_v_block;
HVX_Vector v_max_f16 = hvx_vec_reduce_max_f16(scores_f16);
HVX_Vector v_max = Q6_V_lo_W(hvx_vec_f16_to_f32(v_max_f16));
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_QK, head_ir);
// K (block 1 of next row)
dma_queue_push(dma_q, dma_make_data(k_dst, next_k_src1), factx->size_k_row_padded, nbk1, size_k_row, next_block_size1);
// prefetch K for block ib + 2 after last head finished QK
if (g + 1 == G_local && ib + 2 < factx->n_blocks) {
const uint32_t next_ib = ib + 2;
const uint32_t next_ic_start = next_ib * FLASH_ATTN_BLOCK_SIZE;
const uint32_t next_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - next_ic_start);
// V (block 1 of next row)
dma_queue_push(dma_q, dma_make_data(v_dst, next_v_src1), factx->size_v_row_padded, nbv1, size_v_row, next_block_size1);
// Mask (block 1 of next row)
if (mask) {
dma_cache_push(dma_q, &m_cache, next_m_src1, next_block_size1 * 2, next_block_size1 * 2, next_block_size1 * 2, 1);
}
const dma_addr_t k_src = k->data + next_ic_start*nbk1 + ik2*nbk2 + ik3*nbk3;
dma_queue_push(dma_q, dma_make_data(k_base, k_src), factx->size_k_row_padded, nbk1, size_k_row, next_block_size);
}
} else {
// Queue next row's block 0 (into buffer slot 0)
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_SFM, head_ir);
{
uint8_t * k_dst = spad_k + 0 * factx->size_k_block;
uint8_t * v_dst = spad_v + 0 * factx->size_v_block;
HVX_Vector M_new_vec = Q6_Vsf_vmax_VsfVsf(v_max, M_vec[g]);
HVX_Vector diff_vec = HVX_OP_SUB_F32(M_vec[g], M_new_vec);
// K (block 0 of next row)
dma_queue_push(dma_q, dma_make_data(k_dst, next_k_src0), factx->size_k_row_padded, nbk1, size_k_row, next_block_size0);
HVX_Vector diff_f16 = hvx_vec_f32_to_f16(diff_vec, diff_vec);
HVX_Vector diff_base2 = hvx_vec_mul_f16_f16(diff_f16, v_log2e);
HVX_Vector ms_f16 = hvx_vec_exp2_f16(diff_base2);
HVX_Vector ms_vec = Q6_V_lo_W(hvx_vec_f16_to_f32(ms_f16));
// V (block 0 of next row)
dma_queue_push(dma_q, dma_make_data(v_dst, next_v_src0), factx->size_v_row_padded, nbv1, size_v_row, next_block_size0);
M_vec[g] = M_new_vec;
// Mask (block 0 of next row)
if (mask) {
dma_cache_push(dma_q, &m_cache, next_m_src0, next_block_size0 * 2, next_block_size0 * 2, next_block_size0 * 2, 1);
HVX_Vector v_m_vec_f16 = hvx_vec_f32_to_f16(M_vec[g], M_vec[g]);
HVX_Vector v_s_minus_m = Q6_Vqf16_vsub_VhfVhf(scores_f16, v_m_vec_f16);
HVX_Vector v_s_minus_m_base2 = hvx_vec_mul_f16_f16(Q6_Vhf_equals_Vqf16(v_s_minus_m), v_log2e);
HVX_Vector P = hvx_vec_exp2_f16(v_s_minus_m_base2);
P = Q6_V_vmux_QVV(q_tail_keep, P, Q6_V_vzero());
HVX_VectorPair P_pair = hvx_vec_f16_to_f32(P);
HVX_Vector P0 = Q6_V_lo_W(P_pair);
HVX_Vector P1 = Q6_V_hi_W(P_pair);
HVX_Vector p_sum_vec = hvx_vec_reduce_sum_f32(HVX_OP_ADD_F32(P0, P1));
S_vec[g] = HVX_OP_ADD_F32(HVX_OP_MUL_F32(S_vec[g], ms_vec), p_sum_vec);
const uint8_t * v_ptr = v_base;
if (DV == 64) {
HVX_VectorPair vkq0 = *((const HVX_VectorPair *) VKQ32);
vkq0 = Q6_W_vcombine_VV(
HVX_OP_MUL_F32(Q6_V_hi_W(vkq0), ms_vec),
HVX_OP_MUL_F32(Q6_V_lo_W(vkq0), ms_vec)
);
for (uint32_t j = 0; j < current_block_size; j += 2) {
HVX_Vector S0 = hvx_vec_repl_f16(Q6_V_vror_VR(P, j * 2));
const HVX_Vector * vx0 = (const HVX_Vector *) v_ptr;
if (j + 1 == current_block_size) {
vkq0 = hvx_vec_mpyacc_f32_f16(vkq0, Q6_Vh_vshuff_Vh(vx0[0]), S0);
break;
}
HVX_Vector S1 = hvx_vec_repl_f16(Q6_V_vror_VR(P, (j + 1) * 2));
const HVX_Vector * vx1 = (const HVX_Vector *) (v_ptr + factx->size_v_row_padded);
vkq0 = hvx_vec_mpyacc_f32_f16(vkq0, Q6_Vh_vshuff_Vh(vx0[0]), S0);
vkq0 = hvx_vec_mpyacc_f32_f16(vkq0, Q6_Vh_vshuff_Vh(vx1[0]), S1);
v_ptr += stride_v2;
}
*((HVX_VectorPair *) VKQ32) = vkq0;
} else if (DV == 128) {
HVX_VectorPair vkq0 = ((const HVX_VectorPair *) VKQ32)[0];
HVX_VectorPair vkq1 = ((const HVX_VectorPair *) VKQ32)[1];
vkq0 = Q6_W_vcombine_VV(
HVX_OP_MUL_F32(Q6_V_hi_W(vkq0), ms_vec),
HVX_OP_MUL_F32(Q6_V_lo_W(vkq0), ms_vec)
);
vkq1 = Q6_W_vcombine_VV(
HVX_OP_MUL_F32(Q6_V_hi_W(vkq1), ms_vec),
HVX_OP_MUL_F32(Q6_V_lo_W(vkq1), ms_vec)
);
for (uint32_t j = 0; j < current_block_size; j += 2) {
HVX_Vector S0 = hvx_vec_repl_f16(Q6_V_vror_VR(P, j * 2));
const HVX_Vector * vx0 = (const HVX_Vector *) v_ptr;
if (j + 1 == current_block_size) {
vkq0 = hvx_vec_mpyacc_f32_f16(vkq0, Q6_Vh_vshuff_Vh(vx0[0]), S0);
vkq1 = hvx_vec_mpyacc_f32_f16(vkq1, Q6_Vh_vshuff_Vh(vx0[1]), S0);
break;
}
HVX_Vector S1 = hvx_vec_repl_f16(Q6_V_vror_VR(P, (j + 1) * 2));
const HVX_Vector * vx1 = (const HVX_Vector *) (v_ptr + factx->size_v_row_padded);
vkq0 = hvx_vec_mpyacc_f32_f16(vkq0, Q6_Vh_vshuff_Vh(vx0[0]), S0);
vkq0 = hvx_vec_mpyacc_f32_f16(vkq0, Q6_Vh_vshuff_Vh(vx1[0]), S1);
vkq1 = hvx_vec_mpyacc_f32_f16(vkq1, Q6_Vh_vshuff_Vh(vx0[1]), S0);
vkq1 = hvx_vec_mpyacc_f32_f16(vkq1, Q6_Vh_vshuff_Vh(vx1[1]), S1);
v_ptr += stride_v2;
}
((HVX_VectorPair *) VKQ32)[0] = vkq0;
((HVX_VectorPair *) VKQ32)[1] = vkq1;
} else {
hvx_scale_vec_f32_aa((uint8_t *) VKQ32, (const uint8_t *) VKQ32, DV, ms_vec);
for (uint32_t j = 0; j < current_block_size; j += 2) {
if (j + 1 == current_block_size) {
HVX_Vector S0 = hvx_vec_repl_f16(Q6_V_vror_VR(P, j * 2));
hvx_mad_f32_f16_aa_vec(VKQ32, v_ptr, S0, DV);
break;
}
HVX_Vector S0 = hvx_vec_repl_f16(Q6_V_vror_VR(P, j * 2));
HVX_Vector S1 = hvx_vec_repl_f16(Q6_V_vror_VR(P, (j + 1) * 2));
hvx_mad_f32_f16_aa_rx2_vec(VKQ32, v_ptr, v_ptr + factx->size_v_row_padded, S0, S1, DV);
v_ptr += stride_v2;
}
}
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_SFM, head_ir);
// Queue next row's block 1 (into buffer slot 1, if n_blocks > 1)
if (factx->n_blocks > 1) {
uint8_t * k_dst = spad_k + 1 * factx->size_k_block;
uint8_t * v_dst = spad_v + 1 * factx->size_v_block;
// prefetch V and mask for block ib + 2 after last head finished V accumulation
if (g + 1 == G_local && ib + 2 < factx->n_blocks) {
const uint32_t next_ib = ib + 2;
const uint32_t next_ic_start = next_ib * FLASH_ATTN_BLOCK_SIZE;
const uint32_t next_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - next_ic_start);
// K (block 1 of next row)
dma_queue_push(dma_q, dma_make_data(k_dst, next_k_src1), factx->size_k_row_padded, nbk1, size_k_row, next_block_size1);
// V
const dma_addr_t v_src = v->data + next_ic_start*nbv1 + iv2*nbv2 + iv3*nbv3;
dma_queue_push(dma_q, dma_make_data(v_base, v_src), factx->size_v_row_padded, nbv1, size_v_row, next_block_size);
// V (block 1 of next row)
dma_queue_push(dma_q, dma_make_data(v_dst, next_v_src1), factx->size_v_row_padded, nbv1, size_v_row, next_block_size1);
// Mask (block 1 of next row)
// Mask
if (mask) {
dma_cache_push(dma_q, &m_cache, next_m_src1, next_block_size1 * 2, next_block_size1 * 2, next_block_size1 * 2, 1);
const dma_addr_t m_src = mp_base + next_ic_start * sizeof(__fp16);
dma_cache_push(dma_q, &m_cache, m_src, next_block_size * 2, next_block_size * 2, next_block_size * 2, 1);
}
}
} // end for g
} // end for ib
for (uint32_t g = 0; g < G_local; ++g) {
const uint32_t head_ir = ir + g;
const uint32_t h = heads[g];
float * VKQ32 = vkq_ptrs[g];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_O_PROC, head_ir);
float M = hvx_vec_get_f32(M_vec[g]);
float S = hvx_vec_get_f32(S_vec[g]);
if (sinks) {
const float s = factx->spad_sinks[h];
float vs = 1.0f;
if (s > M) {
HVX_Vector diff_vec = hvx_vec_splat_f32(M - s);
HVX_Vector ms_vec = hvx_vec_exp_f32(diff_vec);
hvx_scale_vec_f32_aa((uint8_t *) VKQ32, (const uint8_t *) VKQ32, DV, ms_vec);
float ms = hvx_vec_get_f32(ms_vec);
S = S * ms + vs;
} else {
HVX_Vector diff_vec = hvx_vec_splat_f32(s - M);
vs = hvx_vec_get_f32(hvx_vec_exp_f32(diff_vec));
S += vs;
}
}
}
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_O_PROC, ir);
// sinks
float M = hvx_vec_get_f32(M_vec);
float S = hvx_vec_get_f32(S_vec);
const float S_inv = S == 0.0f ? 0.0f : 1.0f/S;
hvx_scale_f32_aa((uint8_t *) VKQ32, (const uint8_t *) VKQ32, DV, S_inv);
if (sinks) {
const float s = factx->spad_sinks[h];
const uint32_t r_iq3 = fastdiv(head_ir, &factx->src0_div21);
const uint32_t r_iq2 = fastdiv(head_ir - r_iq3*neq2*neq1, &factx->src0_div1);
const uint32_t r_iq1 = (head_ir - r_iq3*neq2*neq1 - r_iq2 * neq1);
float vs = 1.0f;
uint8_t * dst_ptr = (uint8_t *) dst->data + r_iq2 * dst->nb[1] + r_iq1 * dst->nb[2] + r_iq3 * dst->nb[3];
if (s > M) {
HVX_Vector diff_vec = hvx_vec_splat_f32(M - s);
HVX_Vector ms_vec = hvx_vec_exp_f32(diff_vec);
hvx_scale_vec_f32_aa((uint8_t *) VKQ32, (const uint8_t *) VKQ32, DV, ms_vec);
float ms = hvx_vec_get_f32(ms_vec);
S = S * ms + vs;
} else {
HVX_Vector diff_vec = hvx_vec_splat_f32(s - M);
vs = hvx_vec_get_f32(hvx_vec_exp_f32(diff_vec));
S += vs;
if (dst->type == HTP_TYPE_F32) {
hvx_copy_f32_ua(dst_ptr, (uint8_t *) VKQ32, DV);
} else if (dst->type == HTP_TYPE_F16) {
hvx_copy_f16_f32_ua(dst_ptr, (uint8_t *) VKQ32, DV);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_O_PROC, head_ir);
}
const float S_inv = S == 0.0f ? 0.0f : 1.0f/S;
hvx_scale_f32_aa((uint8_t *) VKQ32, (const uint8_t *) VKQ32, DV, S_inv);
// Store result
// dst indices
const uint32_t i1 = iq1;
const uint32_t i2 = iq2;
const uint32_t i3 = iq3;
// dst is permuted: [DV, n_heads, n_tokens, n_seq]
// head stride is nb[1], token stride is nb[2], batch stride is nb[3]
uint8_t * dst_ptr = (uint8_t *) dst->data + i2 * dst->nb[1] + i1 * dst->nb[2] + i3 * dst->nb[3];
if (dst->type == HTP_TYPE_F32) {
hvx_copy_f32_ua(dst_ptr, (uint8_t *) VKQ32, DV);
} else if (dst->type == HTP_TYPE_F16) {
hvx_copy_f16_f32_ua(dst_ptr, (uint8_t *) VKQ32, DV);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_O_PROC, ir);
ir += G_local;
}
}
@@ -1554,7 +1527,7 @@ static void fa_softmax_thread(unsigned int n, unsigned int i, void * data) {
const bool mask_broadcast = factx->mask_broadcast;
const bool is_g1 = (args->G == 1);
const bool has_alibi = args->has_alibi;
const bool has_softcap = (factx->logit_softcap != 0.0f);
const bool has_softcap = factx->has_softcap;
fa_softmax_impl(n, i, data, has_mask, mask_broadcast, is_g1, has_alibi, has_softcap);
}
@@ -1589,9 +1562,9 @@ static void fa_phase_softmax_and_build_d(struct hmx_fa_context * factx,
const size_t n_row_vec_cnt = hmx_ceil_div(sargs->n_rows_g, 64);
worker_callback_t softmax_fn = fa_softmax_thread;
if (sargs->mask == NULL && factx->logit_softcap == 0.0f && !sargs->has_alibi) {
if (sargs->mask == NULL && !factx->has_softcap && !sargs->has_alibi) {
softmax_fn = fa_softmax_thread_nomask;
} else if (sargs->mask != NULL && factx->mask_broadcast && factx->logit_softcap == 0.0f && !sargs->has_alibi) {
} else if (sargs->mask != NULL && factx->mask_broadcast && !factx->has_softcap && !sargs->has_alibi) {
if (sargs->G == 1) {
softmax_fn = fa_softmax_thread_mask_broadcast_g1;
} else {
@@ -1905,13 +1878,14 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
factx.src3_div3 = kparams->src3_div3;
}
if (kparams->logit_softcap == 0.0f) {
factx.has_softcap = (kparams->logit_softcap != 0.0f);
if (!factx.has_softcap) {
factx.scale = (__fp16) (kparams->scale * EXP_LOG2E_F); // log2(e)
} else {
factx.scale = (__fp16) kparams->scale;
}
factx.max_bias = kparams->max_bias;
factx.logit_softcap = (__fp16) (kparams->logit_softcap * EXP_LOG2E_F);
factx.logit_softcap = factx.has_softcap ? (__fp16) (kparams->logit_softcap * EXP_LOG2E_F) : 0;
factx.n_head_log2 = kparams->n_head_log2;
factx.m0 = kparams->m0;
@@ -2513,7 +2487,8 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
factx.scale = kparams->scale;
factx.max_bias = kparams->max_bias;
factx.logit_softcap = (__fp16) kparams->logit_softcap;
factx.has_softcap = (kparams->logit_softcap != 0.0f);
factx.logit_softcap = factx.has_softcap ? (__fp16) kparams->logit_softcap : 0;
factx.n_head_log2 = kparams->n_head_log2;
factx.m0 = kparams->m0;
+3 -2
View File
@@ -247,6 +247,7 @@ static inline size_t hmx_fa_compute_vtcm_usage(size_t gqa_factor, size_t DK, siz
}
#define FA_HVX_BLOCK_SIZE 64
#define FA_HVX_G_MAX 8
struct hvx_fa_vtcm_layout {
size_t off_q;
@@ -275,11 +276,11 @@ static inline void hvx_fa_vtcm_layout_build(struct hvx_fa_vtcm_layout * L,
const size_t size_k_row_padded = hex_round_up(DK * sizeof(__fp16), 128);
const size_t size_v_row_padded = hex_round_up(DV * sizeof(__fp16), 128);
const size_t size_q_block = size_q_row_padded * 1;
const size_t size_q_block = size_q_row_padded * FA_HVX_G_MAX;
const size_t size_k_block = size_k_row_padded * FA_HVX_BLOCK_SIZE;
const size_t size_v_block = size_v_row_padded * FA_HVX_BLOCK_SIZE;
const size_t size_m_block = hex_round_up(FA_HVX_BLOCK_SIZE * sizeof(__fp16), 128);
const size_t size_vkq_acc = hex_round_up(DV * sizeof(float), 128);
const size_t size_vkq_acc = hex_round_up(DV * sizeof(float), 128) * FA_HVX_G_MAX;
const size_t size_sinks = hex_round_up(n_heads * sizeof(float), 128);
size_t off = 0;
File diff suppressed because it is too large Load Diff
@@ -11,6 +11,7 @@
#define HTP_GDN_MAX_SV 128
#define HTP_GDN_CHUNK_SIZE 64
#define HTP_GDN_MIN_TOKENS 8
#ifndef HMX_FP16_TILE_SIZE
#define HMX_FP16_TILE_SIZE 2048
@@ -132,7 +133,6 @@ struct htp_gdn_hmx_vtcm_layout {
size_t off_rows_a;
size_t off_thread_scratch;
size_t off_attn_rem;
size_t off_scales_1;
size_t state_f32_bytes;
@@ -192,8 +192,10 @@ static inline void htp_gdn_hmx_vtcm_layout_build(
VTCM_LAYOUT_ALLOC(off, off_s_state, bh * state_f32_sz);
VTCM_LAYOUT_ALLOC(off, off_s_f16, bh * state_f16_sz);
off = hex_align_up(off, HMX_FP16_TILE_SIZE);
VTCM_LAYOUT_ALLOC(off, off_s_col_tiles, bh * state_tiles_sz);
VTCM_LAYOUT_ALLOC(off, off_s_update_f32, bh * state_f32_sz);
off = hex_align_up(off, HMX_FP16_TILE_SIZE);
VTCM_LAYOUT_ALLOC(off, off_s_update_tiles, bh * state_tiles_sz);
VTCM_LAYOUT_ALLOC(off, off_q_f32[0], bh * dma_chunk_sz);
@@ -222,6 +224,7 @@ static inline void htp_gdn_hmx_vtcm_layout_build(
VTCM_LAYOUT_ALLOC(off, off_delta_f16, bh * act_f16_sz);
VTCM_LAYOUT_ALLOC(off, off_d_f16, bh * act_f16_sz);
off = hex_align_up(off, HMX_FP16_TILE_SIZE);
VTCM_LAYOUT_ALLOC(off, off_q_row_tiles, bh * tile_64xSv_sz);
VTCM_LAYOUT_ALLOC(off, off_q_prime_row_tiles, bh * tile_64xSv_sz);
VTCM_LAYOUT_ALLOC(off, off_k_row_tiles, bh * tile_64xSv_sz);
@@ -250,9 +253,10 @@ static inline void htp_gdn_hmx_vtcm_layout_build(
VTCM_LAYOUT_ALLOC(off, off_rows_a, bh * row_vecs_sz);
const size_t thread_scratch_sz = 64 * 128;
off = hex_align_up(off, HMX_FP16_TILE_SIZE);
VTCM_LAYOUT_ALLOC(off, off_thread_scratch, nth * thread_scratch_sz);
VTCM_LAYOUT_ALLOC(off, off_attn_rem, nth * (128 * sizeof(float)));
VTCM_LAYOUT_ALLOC(off, off_scales_1, 256);
off = hex_align_up(off, HMX_FP16_TILE_SIZE);
VTCM_LAYOUT_ALLOC(off, off_scales_1, HMX_FP16_TILE_SIZE);
L->total_bytes = off;
}
+2 -2
View File
@@ -48,7 +48,7 @@ static const int16_t d_tile_scatter_offsets[64] __attribute__((aligned(128))) =
};
// Inner HMX tile computation kernels
static void hmx_fa_qk_dot_tile(
static inline void hmx_fa_qk_dot_tile(
const __fp16 * row_tiles,
const __fp16 * col_tiles,
__fp16 * out_tile,
@@ -116,7 +116,7 @@ static void hmx_fa_qk_dot_tile(
);
}
static void hmx_fa_o_update_tile(
static inline void hmx_fa_o_update_tile(
const __fp16 * d_diag,
const __fp16 * o_rc,
const __fp16 * p_tile_in,
+8
View File
@@ -204,6 +204,14 @@ enum htp_trace_event_id {
HTP_TRACE_EVT_HVX_FA_K_PREP = 29,
HTP_TRACE_EVT_HVX_FA_V_PREP = 30,
HTP_TRACE_EVT_HVX_GDN_PREP = 31,
HTP_TRACE_EVT_HVX_GDN_SOLVE = 32,
HTP_TRACE_EVT_HVX_GDN_V_PREP = 33,
HTP_TRACE_EVT_HVX_GDN_D_PREP = 34,
HTP_TRACE_EVT_HVX_GDN_OUT = 35,
HTP_TRACE_EVT_HVX_GDN_STATE = 36,
HTP_TRACE_EVT_HVX_GDN_REM = 37,
HTP_TRACE_EVT_HMX_COMP = 40,
};
+1 -1
View File
@@ -36,7 +36,7 @@
#include "allreduce-ops.h"
#include "htp-fence.h"
#define HMX_QUEUE_CAPACITY 16
#define HMX_QUEUE_CAPACITY 128
#define HMX_QUEUE_STACK_SIZE 16384
#define WORK_QUEUE_CAPACITY 16
#define WORK_QUEUE_STACK_SIZE 16384
+335 -53
View File
@@ -36,12 +36,13 @@ import signal
import subprocess
import sys
from pathlib import Path
from typing import Dict, List, NamedTuple, Optional, Tuple
from typing import Dict, List, NamedTuple, Optional, Set, Tuple
# Ignore SIGPIPE to handle pipes (e.g. head, grep) gracefully
if hasattr(signal, "SIGPIPE"):
signal.signal(signal.SIGPIPE, signal.SIG_DFL)
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
logger = logging.getLogger("ggml-hexagon-inspect")
@@ -56,6 +57,37 @@ class InsnInfo(NamedTuple):
in_loop: bool
class LoopStats:
def __init__(self, loop_type: str, start_addr: int, end_addr: Optional[int] = None, loop_id: int = 0):
self.loop_id = loop_id
self.loop_type = loop_type # "loop0" or "loop1"
self.start_addr = start_addr
self.end_addr = end_addr
self.packet_count = 0
self.insn_count = 0
self.vec_insn_count = 0
self.vspills_st = 0
self.vspills_ld = 0
self.sspills_st = 0
self.sspills_ld = 0
@property
def vspills_total(self) -> int:
return self.vspills_st + self.vspills_ld
@property
def sspills_total(self) -> int:
return self.sspills_st + self.sspills_ld
@property
def has_v_roundtrip(self) -> bool:
return self.vspills_st > 0 and self.vspills_ld > 0
@property
def vec_density(self) -> float:
return (self.vec_insn_count / self.packet_count) if self.packet_count > 0 else 0.0
class FuncStats:
def __init__(self, name: str, address: int, size: int):
self.name = name
@@ -66,14 +98,19 @@ class FuncStats:
self.vec_insn_count = 0
self.loop_count = 0
self.vspills_in_loop = 0
self.vspills_in_loop_st = 0
self.vspills_in_loop_ld = 0
self.vspills_total = 0
self.sspills_in_loop = 0
self.sspills_in_loop_st = 0
self.sspills_in_loop_ld = 0
self.sspills_total = 0
self.promotions_in_loop = 0
self.promotions_total = 0
self.promotion_targets: Dict[str, int] = {}
self.calls_in_loop = 0
self.calls_total = 0
self.loops: List[LoopStats] = []
self.insns: List[InsnInfo] = []
@@ -90,16 +127,43 @@ RE_INSN_LINE = re.compile(
)
RE_LOOP0_START = re.compile(r"\bloop0\((0x[0-9a-fA-F]+)")
RE_LOOP1_START = re.compile(r"\bloop1\((0x[0-9a-fA-F]+)")
RE_VSPILL = re.compile(r"\bvmemu?\s*\(\s*r(?:29|30)\b")
RE_SSPILL = re.compile(r"\bmem[bwhd]\s*\(\s*r(?:29|30)\b")
RE_VMEM_BASE = re.compile(r"\bvmemu?\s*\(\s*([a-z0-9]+)\b")
RE_SMEM_BASE = re.compile(r"\bmem[bwhd](?:_locked|_fifo)?\s*\(\s*([a-z0-9]+)\b")
RE_MEM_STORE = re.compile(r"\bv?mem[bwhdu]?(?:_[a-z]+)?\s*\([^)]*\)\s*(\+|-)?=")
RE_ADD_OP = re.compile(r"\b(r[0-9]+)\s*=\s*add\s*\(\s*([^,()]+)\s*,\s*([^,()]+)\s*\)")
RE_ASSIGN_LHS = re.compile(r"^\s*(?:if\s*\([^)]+\)\s*)?(r[0-9]+)(?::(r[0-9]+))?\s*(?:[+\-*/&|^]?=)")
RE_VEC_OP = re.compile(r"\b(v[0-9]+|w[0-9]+|q[0-3]|vmemu?)\b")
RE_STORE = re.compile(r"=\s*(?:v[0-9]|r[0-9]|w[0-9]|#)")
RE_PROMOTION_CALL = re.compile(
r"\b(?:call|jump)\s+(?:0x[0-9a-fA-F]+\s+)?<(__(?:trunc|extend)[a-zA-Z0-9_]+)(?:@plt)?>"
)
RE_ANY_CALL = re.compile(r"\bcallr?\b")
def is_mem_store(insn: str) -> bool:
return bool(RE_MEM_STORE.search(insn))
def update_sp_regs(insn: str, sp_regs: Set[str]) -> None:
# Track registers derived from stack frame (r29/r30)
m_add = RE_ADD_OP.search(insn)
if m_add:
dest = m_add.group(1)
op1 = m_add.group(2).strip()
op2 = m_add.group(3).strip()
if op1 in sp_regs or op2 in sp_regs:
sp_regs.add(dest)
return
m_assign = RE_ASSIGN_LHS.match(insn.strip())
if m_assign:
r1 = m_assign.group(1)
r2 = m_assign.group(2)
if r1 and r1 not in ("r29", "r30"):
sp_regs.discard(r1)
if r2 and r2 not in ("r29", "r30"):
sp_regs.discard(r2)
def get_repo_root() -> Path:
# Resolve repository root from script location
return Path(__file__).resolve().parent.parent.parent
@@ -338,13 +402,15 @@ def parse_disassembly(
end_idx = matches[i + 1].start() if i + 1 < len(matches) else len(disasm_text)
chunk = disasm_text[start_idx:end_idx]
# Calculate rough byte size from line addresses
stats = FuncStats(name=name, address=addr, size=0)
loop0_target: Optional[int] = None
loop1_target: Optional[int] = None
loop0_active = False
loop1_active = False
current_loop0: Optional[LoopStats] = None
current_loop1: Optional[LoopStats] = None
sp_regs: Set[str] = {"r29", "r30"}
first_addr = None
last_addr = None
@@ -364,6 +430,10 @@ def parse_disassembly(
# Track packet count
if "{" in asm_chunk:
stats.packet_count += 1
if current_loop0:
current_loop0.packet_count += 1
if current_loop1:
current_loop1.packet_count += 1
# Check loop starts
m0 = RE_LOOP0_START.search(asm_chunk)
@@ -378,8 +448,23 @@ def parse_disassembly(
if loop0_target is not None and cur_addr >= loop0_target:
loop0_active = True
if current_loop0 is None:
current_loop0 = LoopStats(
loop_id=len(stats.loops) + 1,
loop_type="loop0",
start_addr=loop0_target,
end_addr=0,
)
if loop1_target is not None and cur_addr >= loop1_target:
loop1_active = True
if current_loop1 is None:
current_loop1 = LoopStats(
loop_id=len(stats.loops) + 1,
loop_type="loop1",
start_addr=loop1_target,
end_addr=0,
)
in_loop = loop0_active or loop1_active
@@ -388,31 +473,70 @@ def parse_disassembly(
sub_insns = [p.strip() for p in cleaned.split(";") if p.strip()]
for insn in sub_insns:
update_sp_regs(insn, sp_regs)
stats.insn_count += 1
if current_loop0:
current_loop0.insn_count += 1
if current_loop1:
current_loop1.insn_count += 1
is_vec = bool(RE_VEC_OP.search(insn))
if is_vec:
stats.vec_insn_count += 1
if current_loop0:
current_loop0.vec_insn_count += 1
if current_loop1:
current_loop1.vec_insn_count += 1
is_vspill = bool(RE_VSPILL.search(insn))
is_sspill = bool(RE_SSPILL.search(insn))
vm = RE_VMEM_BASE.search(insn)
is_vspill = bool(vm and vm.group(1) in sp_regs)
sm = RE_SMEM_BASE.search(insn)
is_sspill = bool(sm and sm.group(1) in sp_regs)
# Identify store vs load
is_store = False
is_load = False
if is_vspill or is_sspill:
if RE_STORE.search(insn):
is_store = True
else:
is_load = True
is_store = is_mem_store(insn)
is_load = not is_store
if is_vspill:
stats.vspills_total += 1
if in_loop:
stats.vspills_in_loop += 1
if is_store:
stats.vspills_in_loop_st += 1
else:
stats.vspills_in_loop_ld += 1
if current_loop0:
if is_store:
current_loop0.vspills_st += 1
else:
current_loop0.vspills_ld += 1
if current_loop1:
if is_store:
current_loop1.vspills_st += 1
else:
current_loop1.vspills_ld += 1
elif is_sspill:
stats.sspills_total += 1
if in_loop:
stats.sspills_in_loop += 1
if is_store:
stats.sspills_in_loop_st += 1
else:
stats.sspills_in_loop_ld += 1
if current_loop0:
if is_store:
current_loop0.sspills_st += 1
else:
current_loop0.sspills_ld += 1
if current_loop1:
if is_store:
current_loop1.sspills_st += 1
else:
current_loop1.sspills_ld += 1
is_call = bool(RE_ANY_CALL.search(insn))
prom_m = RE_PROMOTION_CALL.search(insn)
@@ -444,9 +568,29 @@ def parse_disassembly(
if ":endloop0" in asm_chunk:
loop0_active = False
loop0_target = None
if current_loop0:
current_loop0.end_addr = cur_addr
stats.loops.append(current_loop0)
current_loop0 = None
if ":endloop1" in asm_chunk:
loop1_active = False
loop1_target = None
if current_loop1:
current_loop1.end_addr = cur_addr
stats.loops.append(current_loop1)
current_loop1 = None
if current_loop0:
current_loop0.end_addr = last_addr or 0
stats.loops.append(current_loop0)
if current_loop1:
current_loop1.end_addr = last_addr or 0
stats.loops.append(current_loop1)
stats.loops.sort(key=lambda lp: lp.start_addr)
for idx, loop in enumerate(stats.loops, 1):
loop.loop_id = idx
if first_addr is not None and last_addr is not None:
stats.size = (last_addr - first_addr) + 4
@@ -463,14 +607,19 @@ def annotate_disasm_line(
loop0_active: bool,
loop1_active: bool,
use_color: bool = True,
) -> Tuple[str, Optional[int], Optional[int], bool, bool]:
sp_regs: Optional[Set[str]] = None,
) -> Tuple[str, Optional[int], Optional[int], bool, bool, bool]:
# Annotate disassembly line with spill and loop tags
lm = RE_INSN_LINE.match(raw_line)
if not lm:
return raw_line, loop0_target, loop1_target, loop0_active, loop1_active
return raw_line, loop0_target, loop1_target, loop0_active, loop1_active, False
cur_addr = int(lm.group(1), 16)
asm_chunk = lm.group(4)
is_event = False
if sp_regs is None:
sp_regs = {"r29", "r30"}
# Check loop starts
m0 = RE_LOOP0_START.search(asm_chunk)
@@ -490,39 +639,75 @@ def annotate_disasm_line(
tags = []
if m0:
tags.append("[LOOP0-START]")
is_event = True
if m1:
tags.append("[LOOP1-START]")
is_event = True
if RE_VSPILL.search(asm_chunk):
if in_loop:
tags.append("[V-SPILL:IN-LOOP]" if not use_color else "\033[1;31m[V-SPILL:IN-LOOP]\033[0m")
else:
tags.append("[V-SPILL]" if not use_color else "\033[1;33m[V-SPILL]\033[0m")
elif RE_SSPILL.search(asm_chunk):
if in_loop:
tags.append("[S-SPILL:IN-LOOP]" if not use_color else "\033[1;35m[S-SPILL:IN-LOOP]\033[0m")
cleaned = re.sub(r"[{}\s]|:endloop[01]", " ", asm_chunk)
sub_insns = [p.strip() for p in cleaned.split(";") if p.strip()]
for insn in sub_insns:
update_sp_regs(insn, sp_regs)
for insn in sub_insns:
vm = RE_VMEM_BASE.search(insn)
if vm and vm.group(1) in sp_regs:
base = vm.group(1)
is_st = is_mem_store(insn)
op = "STORE" if is_st else "LOAD"
tgt = f"({base})" if base not in ("r29", "r30") else ""
if in_loop:
tag = f"[V-SPILL:{op}{tgt}:IN-LOOP]"
tags.append(f"\033[1;31m{tag}\033[0m" if use_color else tag)
else:
tag = f"[V-SPILL:{op}{tgt}]"
tags.append(f"\033[1;33m{tag}\033[0m" if use_color else tag)
is_event = True
sm = RE_SMEM_BASE.search(insn)
if sm and sm.group(1) in sp_regs:
base = sm.group(1)
is_st = is_mem_store(insn)
op = "STORE" if is_st else "LOAD"
tgt = f"({base})" if base not in ("r29", "r30") else ""
if in_loop:
tag = f"[S-SPILL:{op}{tgt}:IN-LOOP]"
tags.append(f"\033[1;35m{tag}\033[0m" if use_color else tag)
else:
tag = f"[S-SPILL:{op}{tgt}]"
tags.append(f"\033[0;35m{tag}\033[0m" if use_color else tag)
is_event = True
prom_m = RE_PROMOTION_CALL.search(asm_chunk)
if prom_m:
ptarget = prom_m.group(1)
if in_loop:
tags.append(f"[PROMOTION:{ptarget}:IN-LOOP]" if not use_color else f"\033[1;31m[PROMOTION:{ptarget}:IN-LOOP]\033[0m")
tag = f"[PROMOTION:{ptarget}:IN-LOOP]"
tags.append(f"\033[1;31m{tag}\033[0m" if use_color else tag)
else:
tags.append(f"[PROMOTION:{ptarget}]" if not use_color else f"\033[1;35m[PROMOTION:{ptarget}]\033[0m")
tag = f"[PROMOTION:{ptarget}]"
tags.append(f"\033[1;35m{tag}\033[0m" if use_color else tag)
is_event = True
elif RE_ANY_CALL.search(asm_chunk):
if in_loop:
tags.append("[CALL:IN-LOOP]" if not use_color else "\033[1;31m[CALL:IN-LOOP]\033[0m")
tag = "[CALL:IN-LOOP]"
tags.append(f"\033[1;31m{tag}\033[0m" if use_color else tag)
is_event = True
else:
tags.append("[CALL]" if not use_color else "\033[1;36m[CALL]\033[0m")
tag = "[CALL]"
tags.append(f"\033[1;36m{tag}\033[0m" if use_color else tag)
if ":endloop0" in asm_chunk:
tags.append("[LOOP0-END]")
loop0_active = False
loop0_target = None
is_event = True
if ":endloop1" in asm_chunk:
tags.append("[LOOP1-END]")
loop1_active = False
loop1_target = None
is_event = True
tag_str = " ".join(tags)
if tag_str:
@@ -530,7 +715,7 @@ def annotate_disasm_line(
else:
annotated = raw_line
return annotated, loop0_target, loop1_target, loop0_active, loop1_active
return annotated, loop0_target, loop1_target, loop0_active, loop1_active, is_event
def run_spills(
@@ -566,14 +751,15 @@ def run_spills(
col_pkts = "Packets"
col_insn = "Insns"
col_vec = "HVX Ops"
col_vloop = "V-Loop"
col_vloop = "V-Loop (st/ld)"
col_vtot = "V-Tot"
col_sloop = "S-Loop"
col_sloop = "S-Loop (st/ld)"
col_stot = "S-Tot"
col_notes = "Notes"
hdr = (
f"{col_addr:<10} | {col_name:<44} | {col_pkts:>7} | {col_insn:>6} | "
f"{col_vec:>7} | {col_vloop:>6} | {col_vtot:>5} | {col_sloop:>6} | {col_stot:>5}"
f"{col_addr:<10} | {col_name:<40} | {col_pkts:>7} | {col_insn:>6} | "
f"{col_vec:>7} | {col_vloop:>14} | {col_vtot:>5} | {col_sloop:>14} | {col_stot:>5} | {col_notes}"
)
sep = "-" * len(hdr)
@@ -596,9 +782,11 @@ def run_spills(
# Check strict criteria
if args.strict:
if f.vspills_in_loop > args.max_inloop_vspills:
inloop_v = f.vspills_in_loop_st if getattr(args, "strict_stores_only", False) else f.vspills_in_loop
if inloop_v > args.max_inloop_vspills:
lbl = "in-loop vector store spills" if getattr(args, "strict_stores_only", False) else "in-loop vector spills"
strict_violations.append(
f"{f.name}: {f.vspills_in_loop} in-loop vector spills (max allowed: {args.max_inloop_vspills})"
f"{f.name}: {inloop_v} {lbl} (max allowed: {args.max_inloop_vspills})"
)
if dma_re and dma_re.search(f.name):
if f.vec_insn_count > args.max_dma_vec_ops:
@@ -606,14 +794,27 @@ def run_spills(
f"{f.name}: DMA worker contains {f.vec_insn_count} HVX vector ops (max allowed: {args.max_dma_vec_ops})"
)
# Highlight in-loop vector spills
vloop_str = f"{f.vspills_in_loop:>6}"
vloop_detail = f"{f.vspills_in_loop} ({f.vspills_in_loop_st}s,{f.vspills_in_loop_ld}l)" if f.vspills_in_loop > 0 else "0"
sloop_detail = f"{f.sspills_in_loop} ({f.sspills_in_loop_st}s,{f.sspills_in_loop_ld}l)" if f.sspills_in_loop > 0 else "0"
notes = ""
if f.vspills_in_loop_st > 0 and f.vspills_in_loop_ld > 0:
notes = "\033[1;31m[V-ROUNDTRIP!]\033[0m" if use_color else "[V-ROUNDTRIP!]"
elif f.vspills_in_loop_st == 0 and f.vspills_in_loop_ld > 0:
notes = "v-readonly"
vloop_str = f"{vloop_detail:>14}"
if f.vspills_in_loop > 0 and use_color:
vloop_str = f"\033[1;31m{vloop_str}\033[0m"
if f.vspills_in_loop_st > 0 and f.vspills_in_loop_ld > 0:
vloop_str = f"\033[1;31m{vloop_str}\033[0m"
else:
vloop_str = f"\033[1;33m{vloop_str}\033[0m"
sloop_str = f"{sloop_detail:>14}"
logger.info(
f"0x{f.address:08x} | {f.name:<44} | {f.packet_count:>7} | {f.insn_count:>6} | "
f"{f.vec_insn_count:>7} | {vloop_str} | {f.vspills_total:>5} | {f.sspills_in_loop:>6} | {f.sspills_total:>5}"
f"0x{f.address:08x} | {f.name:<40} | {f.packet_count:>7} | {f.insn_count:>6} | "
f"{f.vec_insn_count:>7} | {vloop_str} | {f.vspills_total:>5} | {sloop_str} | {f.sspills_total:>5} | {notes}"
)
logger.info(sep)
@@ -744,7 +945,7 @@ def run_disasm(
args: argparse.Namespace,
) -> int:
# Disassemble matching function(s) with annotated loop and spill markers
func_pattern = args.disasm
func_pattern = args.disasm if args.disasm else (args.func or ".*")
logger.info(f"Inspecting library: {lib_path}")
logger.info(f"Disassembling functions matching: '{func_pattern}'\n")
@@ -794,9 +995,11 @@ def run_disasm(
logger.info(f"Packets: {func_stats.packet_count} | Instructions: {func_stats.insn_count} | Loops: {func_stats.loop_count}")
vec_pct = (func_stats.vec_insn_count / func_stats.insn_count * 100.0) if func_stats.insn_count else 0.0
logger.info(f"HVX Ops: {func_stats.vec_insn_count} ({vec_pct:.1f}% of instructions)")
vloop_info = f"{func_stats.vspills_in_loop} ({func_stats.vspills_in_loop_st} st, {func_stats.vspills_in_loop_ld} ld)"
sloop_info = f"{func_stats.sspills_in_loop} ({func_stats.sspills_in_loop_st} st, {func_stats.sspills_in_loop_ld} ld)"
logger.info(
f"Spills: Vector in-loop: {func_stats.vspills_in_loop} | Vector total: {func_stats.vspills_total} | "
f"Scalar in-loop: {func_stats.sspills_in_loop} | Scalar total: {func_stats.sspills_total}"
f"Spills: Vector in-loop: {vloop_info} | Vector total: {func_stats.vspills_total} | "
f"Scalar in-loop: {sloop_info} | Scalar total: {func_stats.sspills_total}"
)
logger.info(
f"Calls: Total: {func_stats.calls_total} (in-loop: {func_stats.calls_in_loop}) | "
@@ -804,18 +1007,76 @@ def run_disasm(
)
logger.info(hdr_border)
# Log annotated disassembly
loop0_target: Optional[int] = None
loop1_target: Optional[int] = None
# Print Loop Breakdown Table if function has loops
if func_stats.loops:
logger.info(f"\n--- Loops ({len(func_stats.loops)}) " + "-" * 67)
loop_hdr = (
f"{'#':<3} | {'Type':<5} | {'Address Range':<25} | {'Packets':>7} | "
f"{'HVX Ops':>7} | {'Vec/Pkt':>7} | {'V-Spills (st, ld)':>17} | {'S-Spills (st, ld)':>17} | Notes"
)
logger.info(loop_hdr)
logger.info("-" * len(loop_hdr))
for loop in func_stats.loops:
vspill_str = f"{loop.vspills_total} ({loop.vspills_st}s,{loop.vspills_ld}l)"
sspill_str = f"{loop.sspills_total} ({loop.sspills_st}s,{loop.sspills_ld}l)"
notes = []
if loop.has_v_roundtrip:
notes.append("\033[1;31m[V-ROUNDTRIP!]\033[0m" if use_color else "[V-ROUNDTRIP!]")
elif loop.vspills_st == 0 and loop.vspills_ld > 0:
notes.append("v-readonly")
if loop.vec_density >= 1.5:
notes.append("\033[1;32mdual-hvx\033[0m" if use_color else "dual-hvx")
notes_str = ", ".join(notes)
logger.info(
f"{loop.loop_id:<3} | {loop.loop_type:<5} | 0x{loop.start_addr:08x} - 0x{loop.end_addr:08x} | "
f"{loop.packet_count:>7} | {loop.vec_insn_count:>7} | {loop.vec_density:>7.2f} | "
f"{vspill_str:>17} | {sspill_str:>17} | {notes_str}"
)
logger.info("-" * len(loop_hdr) + "\n")
# Parse lines and annotations
lines = chunk.splitlines()
annotated_lines = []
is_event_list = []
loop0_target = None
loop1_target = None
loop0_active = False
loop1_active = False
sp_regs = {"r29", "r30"}
for line in chunk.splitlines():
ann_line, loop0_target, loop1_target, loop0_active, loop1_active = annotate_disasm_line(
line, loop0_target, loop1_target, loop0_active, loop1_active, use_color
for line in lines:
ann_line, loop0_target, loop1_target, loop0_active, loop1_active, is_ev = annotate_disasm_line(
line, loop0_target, loop1_target, loop0_active, loop1_active, use_color, sp_regs
)
logger.info(ann_line)
logger.info("")
annotated_lines.append(ann_line)
is_event_list.append(is_ev)
# Filter output if --spills-only
if getattr(args, "spills_only", False):
ctx = args.context if args.context is not None else 2
to_show = [False] * len(annotated_lines)
for idx, ev in enumerate(is_event_list):
if ev:
for j in range(max(0, idx - ctx), min(len(annotated_lines), idx + ctx + 1)):
to_show[j] = True
if not any(to_show):
logger.info(" (No spills, promotions, or in-loop calls detected in this function)\n")
else:
in_gap = False
for idx, show in enumerate(to_show):
if show:
in_gap = False
logger.info(annotated_lines[idx])
else:
if not in_gap:
logger.info(" ...")
in_gap = True
logger.info("")
else:
for ann_line in annotated_lines:
logger.info(ann_line)
logger.info("")
return 0
@@ -964,9 +1225,24 @@ def main():
)
parser.add_argument(
"--disasm",
nargs="?",
const="",
metavar="FUNC",
help="Disassemble function symbol or regex pattern with annotated loop and spill markers.",
)
parser.add_argument(
"--spills-only",
action="store_true",
help="In --disasm, only display packets containing spills, promotions, or in-loop calls, with surrounding context.",
)
parser.add_argument(
"-C",
"--context",
type=int,
default=None,
metavar="N",
help="Number of context packets before and after spills in --disasm --spills-only (default: 2).",
)
parser.add_argument(
"--limit",
type=int,
@@ -983,8 +1259,9 @@ def main():
# Filtering & Display
parser.add_argument(
"--func",
"--fn",
"-f",
help="Regex filter for function names in --spills or --promotions.",
help="Regex filter for function names in --spills, --promotions, or --disasm.",
)
parser.add_argument(
"--all",
@@ -1010,6 +1287,11 @@ def main():
default=0,
help="Maximum allowed in-loop vector spills in --strict mode (default: 0).",
)
parser.add_argument(
"--strict-stores-only",
action="store_true",
help="In --strict mode, only count vector store spills (st > 0) towards violations, ignoring readonly stack loads.",
)
parser.add_argument(
"--max-dma-vec-ops",
type=int,
@@ -1057,7 +1339,7 @@ def main():
args = parser.parse_args()
logging.basicConfig(level=logging.INFO, format="%(message)s")
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
repo_root = get_repo_root()
@@ -1092,7 +1374,7 @@ def main():
# Dispatch commands
if args.addr2line is not None:
sys.exit(run_addr2line(toolchain, lib_path, args))
elif args.disasm:
elif args.disasm is not None:
sys.exit(run_disasm(toolchain, lib_path, args))
elif args.promotions:
sys.exit(run_promotions(toolchain, lib_path, args))
@@ -1102,5 +1384,5 @@ def main():
if __name__ == "__main__":
logging.basicConfig(level=logging.INFO, format="%(message)s")
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
main()
+2 -1
View File
@@ -54,6 +54,7 @@ def device_matches(record_device, target_device):
return False
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
logger = logging.getLogger("ggml-hexagon-profile")
@@ -648,7 +649,7 @@ def main():
args = parser.parse_args()
logging.basicConfig(level=logging.INFO, format='%(message)s')
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
if "pmu" in args.sort and args.pmu_index is None:
logger.error(f"Cannot sort by '{args.sort}' without --pmu-index.")
+2 -1
View File
@@ -10,6 +10,7 @@ import bisect
from typing import Any, Dict, List, Optional
from collections import defaultdict
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
logger = logging.getLogger("ggml-hexagon-trace")
op_pattern = re.compile(
@@ -732,7 +733,7 @@ def main():
group.add_argument("--tail", type=int, help="Limit to last N ops")
args = parser.parse_args()
logging.basicConfig(level=logging.INFO, format='%(message)s')
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
op_filter_re = None
if args.filter:
+3
View File
@@ -31,6 +31,7 @@ MANAGED_ENV_NAMES = (
"GGML_HEXAGON_MBUF",
"GGML_HEXAGON_MM_SELECT",
"GGML_HEXAGON_FA_SELECT",
"GGML_HEXAGON_GDN_SELECT",
"GGML_HEXAGON_AR_SELECT",
"GGML_HEXAGON_ETM",
"GGML_HEXAGON_ARCH",
@@ -166,6 +167,7 @@ def main():
parser.add_argument("--hex-mbuf", help="Maximum host buffer size limit in MB to allocate (GGML_HEXAGON_MBUF)")
parser.add_argument("--hex-mm-select", help="Select MUL_MAT and MUL_MAT_ID kernel (GGML_HEXAGON_MM_SELECT) 2:HMX,1:HVX,0:disable")
parser.add_argument("--hex-fa-select", help="Select Flash Attention kernel (GGML_HEXAGON_FA_SELECT) 2:HMX,1:HVX,0:disable")
parser.add_argument("--hex-gdn-select", help="Select Gated Delta Net kernel (GGML_HEXAGON_GDN_SELECT) 2:HMX,1:HVX,0:disable")
parser.add_argument("--hex-ar-select", help="Select All-Reduce kernel (GGML_HEXAGON_AR_SELECT) 1:enable,0:disable")
parser.add_argument("--hex-etm", help="Enable Embedded Trace Macrocell hardware tracing / trace logging (GGML_HEXAGON_ETM)")
parser.add_argument("--hex-arch", help="Target Hexagon NPU architecture version override (v73, v75, v79, v81, etc.) (GGML_HEXAGON_ARCH)")
@@ -306,6 +308,7 @@ def main():
set_env("GGML_HEXAGON_MBUF", args.hex_mbuf)
set_env("GGML_HEXAGON_MM_SELECT", args.hex_mm_select)
set_env("GGML_HEXAGON_FA_SELECT", args.hex_fa_select)
set_env("GGML_HEXAGON_GDN_SELECT", args.hex_gdn_select)
set_env("GGML_HEXAGON_AR_SELECT", args.hex_ar_select)
set_env("GGML_HEXAGON_ETM", args.hex_etm)
set_env("GGML_HEXAGON_ARCH", args.hex_arch)