From 898392b724548712c6162d7ca8d1645e01492d2e Mon Sep 17 00:00:00 2001 From: Max Krasnyansky Date: Mon, 21 Sep 2026 14:49:52 -0700 Subject: [PATCH] hexagon: new HMX-optimized GATED_DELTA_NET (llama/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 --- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 98 +- ggml/src/ggml-hexagon/htp-opnode.h | 4 +- ggml/src/ggml-hexagon/htp/flash-attn-ops.c | 571 ++++--- ggml/src/ggml-hexagon/htp/flash-attn-ops.h | 5 +- .../ggml-hexagon/htp/gated-delta-net-ops.c | 1324 ++++++++++++++++- .../ggml-hexagon/htp/gated-delta-net-ops.h | 10 +- ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h | 4 +- ggml/src/ggml-hexagon/htp/htp-ops.h | 8 + ggml/src/ggml-hexagon/htp/main.c | 2 +- 9 files changed, 1669 insertions(+), 357 deletions(-) diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index ec5a4aeb6..58806e37f 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -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 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(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(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(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(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; diff --git a/ggml/src/ggml-hexagon/htp-opnode.h b/ggml/src/ggml-hexagon/htp-opnode.h index 0716a8d21..803aa3f5a 100644 --- a/ggml/src/ggml-hexagon/htp-opnode.h +++ b/ggml/src/ggml-hexagon/htp-opnode.h @@ -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 || diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c index 988886082..bfcf7cb0c 100644 --- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c @@ -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; diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h index 2bd232190..22bb8c53d 100644 --- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.h @@ -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; diff --git a/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c index b37313370..1dd828db7 100644 --- a/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c +++ b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c @@ -2,6 +2,7 @@ #include #include #include +#include #include "hvx-base.h" #include "hvx-copy.h" @@ -11,6 +12,10 @@ #include "ggml-common.h" #include "htp-ctx.h" #include "htp-tensor.h" +#include "htp-vtcm.h" +#include "hmx-utils.h" +#include "hmx-fa-kernels.h" +#include "hmx-queue.h" #include "gated-delta-net-ops.h" #ifndef MIN @@ -55,9 +60,8 @@ static inline HVX_Vector gdn_mul_dot_f32(float * restrict dst, const HVX_Vector return hvx_vec_reduce_sum_f32(acc); } -static inline HVX_Vector gdn_mul_scalar_dot_f32(float * restrict dst, float mul, const HVX_Vector * restrict dot, uint32_t n) { +static inline HVX_Vector gdn_mul_scalar_dot_f32(float * restrict dst, HVX_Vector vmul, const HVX_Vector * restrict dot, uint32_t n) { HVX_Vector acc = Q6_V_vzero(); - const HVX_Vector vmul = hvx_vec_splat_f32(mul); const uint32_t epv = 128 / sizeof(float); const uint32_t nvec = n / epv; const uint32_t nloe = n % epv; @@ -589,20 +593,15 @@ static inline void gdn_step_kda_f32( HVX_Vector vk[4]; HVX_Vector vg[4]; - static const float kInf = INFINITY; - static const float kMaxExp = 88.7228f; - const HVX_Vector max_exp = hvx_vec_splat_f32(kMaxExp); - const HVX_Vector inf = hvx_vec_splat_f32(kInf); - for (uint32_t i = 0; i < nvec; ++i) { vq[i] = hvx_vmemu(q_t + i * epv); vk[i] = hvx_vmemu(k_t + i * epv); - vg[i] = hvx_vec_exp_f32_guard(hvx_vmemu(g_t + i * epv), max_exp, inf); + vg[i] = hvx_vec_exp_f32(hvx_vmemu(g_t + i * epv)); } if (nloe) { vq[nvec] = hvx_vmemu(q_t + nvec * epv); vk[nvec] = hvx_vmemu(k_t + nvec * epv); - vg[nvec] = hvx_vec_exp_f32_guard(hvx_vmemu(g_t + nvec * epv), max_exp, inf); + vg[nvec] = hvx_vec_exp_f32(hvx_vmemu(g_t + nvec * epv)); } const HVX_Vector vbeta = hvx_vec_splat_f32(beta_val); @@ -690,9 +689,8 @@ static inline void gdn_step_scalar_f32( vk[nvec] = hvx_vmemu(k_t + nvec * epv); } - const float gate = expf(g_t[0]); - const HVX_Vector vgate = hvx_vec_splat_f32(gate); - const HVX_Vector vbeta = hvx_vec_splat_f32(beta_val); + const HVX_Vector vgate = hvx_vec_exp_f32(hvx_vec_splat_f32(g_t[0])); + const HVX_Vector vbeta = hvx_vec_splat_f32(beta_val); const HVX_Vector vscale = hvx_vec_splat_f32(scale); float delta[8] __attribute__((aligned(128))); @@ -742,7 +740,7 @@ static inline void gdn_step_scalar_f32( } for (; j < S_v; ++j) { float * row = s_work + (uint64_t) j * S_v; - HVX_Vector vsum = gdn_mul_scalar_dot_f32(row, gate, vk, S_v); + HVX_Vector vsum = gdn_mul_scalar_dot_f32(row, vgate, vk, S_v); HVX_Vector vv_t = hvx_vec_splat_f32(v_t[j]); HVX_Vector vdj = hvx_vec_mul_f32_f32(hvx_vec_sub_f32_f32(vv_t, vsum), vbeta); HVX_Vector vres = gdn_add_scaled_dot_f32(row, vk, vdj, vq, S_v); @@ -1022,6 +1020,1255 @@ static void gated_delta_net_f32_tg_thread(unsigned int nth, unsigned int ith, vo dma_queue_flush(dma_q); } +struct htp_gdn_hmx_gemm_task { + const __fp16 * row_tiles; + const __fp16 * col_tiles; + __fp16 * out_tiles; + uint32_t n_row_tiles; + uint32_t n_col_tiles; + uint32_t n_dot_tiles; + uint32_t dot_stride; + uint8_t * hmx_scales; +}; + +static void htp_gdn_hmx_gemm_worker(void * data) { + struct htp_gdn_hmx_gemm_task * task = (struct htp_gdn_hmx_gemm_task *) data; + asm volatile(HMX_SET_BIAS("%0") :: "r"((unsigned int)task->hmx_scales)); + + const size_t dot_stride = task->dot_stride; + for (uint32_t r = 0; r < task->n_row_tiles; ++r) { + const __fp16 * r_tiles = task->row_tiles + r * dot_stride; + const __fp16 * c_tiles = task->col_tiles; + __fp16 * o_tile = task->out_tiles + r * task->n_col_tiles * HMX_FP16_TILE_N_ELMS; + + for (uint32_t c = 0; c < task->n_col_tiles; ++c) { + hmx_fa_qk_dot_tile(r_tiles, c_tiles, o_tile, task->n_dot_tiles); + c_tiles += dot_stride; + o_tile += HMX_FP16_TILE_N_ELMS; + } + } +} + +static inline void htp_gdn_push_hmx_gemm_task( + hmx_queue_t q, + struct htp_gdn_hmx_gemm_task * task, + const __fp16 * row_tiles, + const __fp16 * col_tiles, + __fp16 * out_tiles, + uint32_t n_row_tiles, + uint32_t n_col_tiles, + uint32_t n_dot_tiles, + uint8_t * scales +) { + task->row_tiles = row_tiles; + task->col_tiles = col_tiles; + task->out_tiles = out_tiles; + task->n_row_tiles = n_row_tiles; + task->n_col_tiles = n_col_tiles; + task->n_dot_tiles = n_dot_tiles; + task->dot_stride = n_dot_tiles * HMX_FP16_TILE_N_ELMS; + task->hmx_scales = scales; + + hmx_queue_push(q, hmx_queue_make_desc(htp_gdn_hmx_gemm_worker, task)); +} + +static inline void gdn_unpack_64x64_tiles_to_vectors( + HVX_Vector * restrict rows, + const __fp16 * restrict tiles +) { + const HVX_Vector * t00 = (const HVX_Vector *) (tiles + 0 * HMX_FP16_TILE_N_ELMS); + const HVX_Vector * t01 = (const HVX_Vector *) (tiles + 1 * HMX_FP16_TILE_N_ELMS); + const HVX_Vector * t10 = (const HVX_Vector *) (tiles + 2 * HMX_FP16_TILE_N_ELMS); + const HVX_Vector * t11 = (const HVX_Vector *) (tiles + 3 * HMX_FP16_TILE_N_ELMS); + + for (uint32_t r = 0; r < 16; ++r) { + HVX_VectorPair vp0 = Q6_W_vdeal_VVR(t01[r], t00[r], -2); + rows[2 * r + 0] = Q6_V_lo_W(vp0); + rows[2 * r + 1] = Q6_V_hi_W(vp0); + + HVX_VectorPair vp1 = Q6_W_vdeal_VVR(t11[r], t10[r], -2); + rows[32 + 2 * r + 0] = Q6_V_lo_W(vp1); + rows[32 + 2 * r + 1] = Q6_V_hi_W(vp1); + } +} + +static inline void gdn_pack_64x64_vectors_to_tiles( + __fp16 * restrict tiles, + const HVX_Vector * restrict rows +) { + HVX_Vector * t00 = (HVX_Vector *) (tiles + 0 * HMX_FP16_TILE_N_ELMS); + HVX_Vector * t01 = (HVX_Vector *) (tiles + 1 * HMX_FP16_TILE_N_ELMS); + HVX_Vector * t10 = (HVX_Vector *) (tiles + 2 * HMX_FP16_TILE_N_ELMS); + HVX_Vector * t11 = (HVX_Vector *) (tiles + 3 * HMX_FP16_TILE_N_ELMS); + + for (uint32_t r = 0; r < 16; ++r) { + HVX_VectorPair vp0 = Q6_W_vshuff_VVR(rows[2 * r + 1], rows[2 * r + 0], -2); + t00[r] = Q6_V_lo_W(vp0); + t01[r] = Q6_V_hi_W(vp0); + + HVX_VectorPair vp1 = Q6_W_vshuff_VVR(rows[32 + 2 * r + 1], rows[32 + 2 * r + 0], -2); + t10[r] = Q6_V_lo_W(vp1); + t11[r] = Q6_V_hi_W(vp1); + } +} + +static inline void gdn_unpack_64xS_tiles_to_f32( + float * restrict dst_f32, + const __fp16 * restrict tiles, + uint32_t S_v +) { + const uint32_t n_col_tiles = S_v / 32; + for (uint32_t r0 = 0; r0 < 2; ++r0) { + for (uint32_t d = 0; d < S_v / 64; ++d) { + const HVX_Vector * t0 = (const HVX_Vector *) (tiles + (r0 * n_col_tiles + 2 * d + 0) * HMX_FP16_TILE_N_ELMS); + const HVX_Vector * t1 = (const HVX_Vector *) (tiles + (r0 * n_col_tiles + 2 * d + 1) * HMX_FP16_TILE_N_ELMS); + + for (uint32_t r = 0; r < 16; ++r) { + HVX_VectorPair vp01 = Q6_W_vdeal_VVR(t1[r], t0[r], -2); + HVX_VectorPair p0 = hvx_vec_f16_to_f32(Q6_V_lo_W(vp01)); + HVX_VectorPair p1 = hvx_vec_f16_to_f32(Q6_V_hi_W(vp01)); + + float * out0 = dst_f32 + (r0 * 32 + 2 * r + 0) * S_v + d * 64; + float * out1 = dst_f32 + (r0 * 32 + 2 * r + 1) * S_v + d * 64; + + hvx_vmem(out0 + 0) = Q6_V_lo_W(p0); + hvx_vmem(out0 + 32) = Q6_V_hi_W(p0); + hvx_vmem(out1 + 0) = Q6_V_lo_W(p1); + hvx_vmem(out1 + 32) = Q6_V_hi_W(p1); + } + } + } +} + +static inline void gdn_unpack_64xS_tiles_to_f16( + __fp16 * restrict dst_f16, + const __fp16 * restrict tiles, + uint32_t S_v +) { + const uint32_t n_col_tiles = S_v / 32; + for (uint32_t r0 = 0; r0 < 2; ++r0) { + for (uint32_t d = 0; d < S_v / 64; ++d) { + const HVX_Vector * t0 = (const HVX_Vector *) (tiles + (r0 * n_col_tiles + 2 * d + 0) * HMX_FP16_TILE_N_ELMS); + const HVX_Vector * t1 = (const HVX_Vector *) (tiles + (r0 * n_col_tiles + 2 * d + 1) * HMX_FP16_TILE_N_ELMS); + + for (uint32_t r = 0; r < 16; ++r) { + HVX_VectorPair vp01 = Q6_W_vdeal_VVR(t1[r], t0[r], -2); + __fp16 * out0 = dst_f16 + (r0 * 32 + 2 * r + 0) * S_v + d * 64; + __fp16 * out1 = dst_f16 + (r0 * 32 + 2 * r + 1) * S_v + d * 64; + + hvx_vmem(out0) = Q6_V_lo_W(vp01); + hvx_vmem(out1) = Q6_V_hi_W(vp01); + } + } + } +} + +static inline void gdn_unpack_SxS_tiles_to_f32( + float * restrict dst_f32, + const __fp16 * restrict tiles, + uint32_t S_v +) { + const uint32_t n_tiles = S_v / 32; + for (uint32_t r0 = 0; r0 < n_tiles; ++r0) { + for (uint32_t d = 0; d < S_v / 64; ++d) { + const HVX_Vector * t0 = (const HVX_Vector *) (tiles + (r0 * n_tiles + 2 * d + 0) * HMX_FP16_TILE_N_ELMS); + const HVX_Vector * t1 = (const HVX_Vector *) (tiles + (r0 * n_tiles + 2 * d + 1) * HMX_FP16_TILE_N_ELMS); + + for (uint32_t r = 0; r < 16; ++r) { + HVX_VectorPair vp01 = Q6_W_vdeal_VVR(t1[r], t0[r], -2); + HVX_VectorPair p0 = hvx_vec_f16_to_f32(Q6_V_lo_W(vp01)); + HVX_VectorPair p1 = hvx_vec_f16_to_f32(Q6_V_hi_W(vp01)); + + float * out0 = dst_f32 + (r0 * 32 + 2 * r + 0) * S_v + d * 64; + float * out1 = dst_f32 + (r0 * 32 + 2 * r + 1) * S_v + d * 64; + + hvx_vmem(out0 + 0) = Q6_V_lo_W(p0); + hvx_vmem(out0 + 32) = Q6_V_hi_W(p0); + hvx_vmem(out1 + 0) = Q6_V_lo_W(p1); + hvx_vmem(out1 + 32) = Q6_V_hi_W(p1); + } + } + } +} + +static inline void gdn_f32_to_hmx_row_tiles_and_f16( + __fp16 * restrict dst_tiles, + __fp16 * restrict dst_prime_tiles, + __fp16 * restrict dst_f16, + const float * restrict src, + const __fp16 * restrict scale_per_row, + uint32_t n_rows, + uint32_t n_cols +) { + const uint32_t n_col_tiles = n_cols / 32; + const uint32_t * scale_pairs = (const uint32_t *) scale_per_row; + + for (uint32_t r = 0; r < n_rows; r += 2) { + uint32_t r0 = r / 32; + uint32_t r1 = (r % 32) / 2; + const float * p0 = src + (r + 0) * n_cols; + const float * p1 = src + (r + 1) * n_cols; + + HVX_Vector v_scale; + if (dst_prime_tiles) { + uint32_t scale_pair = scale_pairs ? scale_pairs[r / 2] : 0x3c003c00; + v_scale = Q6_V_vsplat_R(scale_pair); + } + + for (uint32_t c = 0; c < n_col_tiles; c += 2) { + HVX_Vector v0_0 = hvx_vmem(p0 + (c + 0) * 32); + HVX_Vector v1_0 = hvx_vmem(p1 + (c + 0) * 32); + HVX_Vector v0_1 = hvx_vmem(p0 + (c + 1) * 32); + HVX_Vector v1_1 = hvx_vmem(p1 + (c + 1) * 32); + + HVX_Vector vh0 = hvx_vec_f32_to_f16_shuff(v0_0, v1_0); + HVX_Vector vh1 = hvx_vec_f32_to_f16_shuff(v0_1, v1_1); + __fp16 * tile0 = dst_tiles + (r0 * n_col_tiles + c + 0) * HMX_FP16_TILE_N_ELMS; + __fp16 * tile1 = dst_tiles + (r0 * n_col_tiles + c + 1) * HMX_FP16_TILE_N_ELMS; + ((HVX_Vector *) tile0)[r1] = vh0; + ((HVX_Vector *) tile1)[r1] = vh1; + + if (dst_prime_tiles) { + HVX_Vector vh0_s = hvx_vec_mul_f16_f16(vh0, v_scale); + HVX_Vector vh1_s = hvx_vec_mul_f16_f16(vh1, v_scale); + __fp16 * tile0_s = dst_prime_tiles + (r0 * n_col_tiles + c + 0) * HMX_FP16_TILE_N_ELMS; + __fp16 * tile1_s = dst_prime_tiles + (r0 * n_col_tiles + c + 1) * HMX_FP16_TILE_N_ELMS; + ((HVX_Vector *) tile0_s)[r1] = vh0_s; + ((HVX_Vector *) tile1_s)[r1] = vh1_s; + } + + if (dst_f16) { + HVX_VectorPair vp01 = Q6_W_vdeal_VVR(vh1, vh0, -2); + hvx_vmem(dst_f16 + (r + 0) * n_cols + c * 32) = Q6_V_lo_W(vp01); + hvx_vmem(dst_f16 + (r + 1) * n_cols + c * 32) = Q6_V_hi_W(vp01); + } + } + } +} + +static inline void hvx_transpose_32x32_words(HVX_Vector * restrict m, HVX_Vector * restrict tmp) { + for (int i = 0; i < 16; ++i) { + HVX_VectorPair p = Q6_W_vshuff_VVR(m[2*i + 1], m[2*i], -4); + tmp[2*i + 0] = Q6_V_lo_W(p); + tmp[2*i + 1] = Q6_V_hi_W(p); + } + + for (int b = 0; b < 32; b += 4) { + HVX_VectorPair p0 = Q6_W_vshuff_VVR(tmp[b + 2], tmp[b + 0], -8); + HVX_VectorPair p1 = Q6_W_vshuff_VVR(tmp[b + 3], tmp[b + 1], -8); + m[b + 0] = Q6_V_lo_W(p0); m[b + 1] = Q6_V_hi_W(p0); + m[b + 2] = Q6_V_lo_W(p1); m[b + 3] = Q6_V_hi_W(p1); + } + + for (int b = 0; b < 32; b += 8) { + for (int i = 0; i < 4; ++i) { + HVX_VectorPair p = Q6_W_vshuff_VVR(m[b + i + 4], m[b + i], -16); + tmp[b + 2*i + 0] = Q6_V_lo_W(p); + tmp[b + 2*i + 1] = Q6_V_hi_W(p); + } + } + + for (int b = 0; b < 32; b += 16) { + for (int i = 0; i < 8; ++i) { + HVX_VectorPair p = Q6_W_vshuff_VVR(tmp[b + i + 8], tmp[b + i], -32); + m[b + 2*i + 0] = Q6_V_lo_W(p); + m[b + 2*i + 1] = Q6_V_hi_W(p); + } + } + + for (int i = 0; i < 16; ++i) { + HVX_VectorPair p = Q6_W_vshuff_VVR(m[i + 16], m[i], -64); + tmp[2 * i + 0] = Q6_V_lo_W(p); + tmp[2 * i + 1] = Q6_V_hi_W(p); + } + + for (int i = 0; i < 32; ++i) { + m[i] = tmp[i]; + } +} + +static inline void gdn_pack_d_t_row_tiles( + __fp16 * restrict dst_tiles, + const __fp16 * restrict src_d, + uint32_t S_v, + HVX_Vector * restrict m, + HVX_Vector * restrict tmp +) { + for (uint32_t col_half = 0; col_half < S_v / 64; ++col_half) { + uint32_t r0_base = col_half * 2; + for (uint32_t c0 = 0; c0 < 2; ++c0) { + for (uint32_t s_local = 0; s_local < 32; ++s_local) { + uint32_t s = c0 * 32 + s_local; + m[s_local] = hvx_vmem(src_d + s * S_v + col_half * 64); + } + + hvx_transpose_32x32_words(m, tmp); + + uint32_t tile0_idx = (r0_base + 0) * 2 + c0; + uint32_t tile1_idx = (r0_base + 1) * 2 + c0; + HVX_Vector * t0 = (HVX_Vector *)(dst_tiles + tile0_idx * HMX_FP16_TILE_N_ELMS); + HVX_Vector * t1 = (HVX_Vector *)(dst_tiles + tile1_idx * HMX_FP16_TILE_N_ELMS); + + for (uint32_t r = 0; r < 16; ++r) { + t0[r] = m[r]; + t1[r] = m[16 + r]; + } + } + } +} + +static __attribute__((noinline)) void gdn_build_inv_l_blocks( + __fp16 * restrict inv_row_tiles, + const HVX_Vector * restrict rows_kk, + const __fp16 * restrict decay_m, + const float * restrict beta, + __fp16 * restrict l10_tile, + __fp16 * restrict neg_a11_tile +) { + const HVX_Vector v_one_f16 = hvx_vec_splat_f16(1.0f); + const HVX_VectorPred q_mask64 = Q6_Q_vsetq2_R(64); + + uint16_t beta_u16[64] __attribute__((aligned(128))); + uint16_t l00[32][32] __attribute__((aligned(128))); + uint16_t l11[32][32] __attribute__((aligned(128))); + + HVX_Vector * restrict p_l00 = (HVX_Vector *) l00; + HVX_Vector * restrict p_l11 = (HVX_Vector *) l11; + HVX_Vector * restrict p_l10_tile = (HVX_Vector *) l10_tile; + + HVX_Vector * restrict tile00 = (HVX_Vector *) (inv_row_tiles + 0 * HMX_FP16_TILE_N_ELMS); + HVX_Vector * restrict tile01 = (HVX_Vector *) (inv_row_tiles + 1 * HMX_FP16_TILE_N_ELMS); + HVX_Vector * restrict tile11 = (HVX_Vector *) (inv_row_tiles + 3 * HMX_FP16_TILE_N_ELMS); + HVX_Vector * restrict p_neg_a11 = (HVX_Vector *) neg_a11_tile; + + hvx_vmem(beta_u16) = hvx_vec_f32_to_f16(hvx_vmem(beta + 0), hvx_vmem(beta + 32)); + + for (uint32_t r = 0; r < 16; ++r) { + tile01[r] = Q6_V_vzero(); + } + + for (uint32_t r = 0; r < 16; ++r) { + uint32_t t0 = 2 * r; + uint32_t t1 = t0 + 1; + + HVX_Vector v_d0 = hvx_vmem(decay_m + t0 * 64); + HVX_Vector v_d1 = hvx_vmem(decay_m + t1 * 64); + HVX_Vector v_b0 = Q6_Vh_vsplat_R(beta_u16[t0]); + HVX_Vector v_b1 = Q6_Vh_vsplat_R(beta_u16[t1]); + + HVX_Vector r0 = hvx_vec_mul_f16_f16(hvx_vec_mul_f16_f16(rows_kk[t0], v_d0), v_b0); + HVX_Vector r1 = hvx_vec_mul_f16_f16(hvx_vec_mul_f16_f16(rows_kk[t1], v_d1), v_b1); + + p_l00[r] = Q6_V_vmux_QVV(q_mask64, r0, Q6_V_vror_VR(r1, 64)); + } + + for (uint32_t r = 0; r < 16; ++r) { + uint32_t t0 = 32 + 2 * r; + uint32_t t1 = t0 + 1; + + HVX_Vector v_d0 = hvx_vmem(decay_m + t0 * 64); + HVX_Vector v_d1 = hvx_vmem(decay_m + t1 * 64); + HVX_Vector v_b0 = Q6_Vh_vsplat_R(beta_u16[t0]); + HVX_Vector v_b1 = Q6_Vh_vsplat_R(beta_u16[t1]); + + HVX_Vector r0 = hvx_vec_mul_f16_f16(hvx_vec_mul_f16_f16(rows_kk[t0], v_d0), v_b0); + HVX_Vector r1 = hvx_vec_mul_f16_f16(hvx_vec_mul_f16_f16(rows_kk[t1], v_d1), v_b1); + + HVX_VectorPair vp_l10 = Q6_W_vshuff_VVR(r1, r0, -2); + p_l10_tile[r] = Q6_V_lo_W(vp_l10); + p_l11[r] = Q6_V_vmux_QVV(q_mask64, Q6_V_vror_VR(r0, 64), r1); + } + + HVX_Vector a_rows[32]; + for (uint32_t t = 0; t < 32; ++t) { + HVX_Vector v_inv = Q6_V_vzero(); + for (uint32_t k = 0; k < t; ++k) { + HVX_Vector v_lk = Q6_Vh_vsplat_R(l00[t][k]); + v_inv = hvx_vec_sub_f16_f16(v_inv, hvx_vec_mul_f16_f16(v_lk, a_rows[k])); + } + HVX_VectorPred q_diag = (t == 0) ? Q6_Q_vsetq2_R(2) : Q6_Q_and_QQn(Q6_Q_vsetq2_R(2 * (t + 1)), Q6_Q_vsetq2_R(2 * t)); + a_rows[t] = Q6_V_vand_QV(q_mask64, Q6_V_vmux_QVV(q_diag, v_one_f16, v_inv)); + } + + for (uint32_t r = 0; r < 16; ++r) { + HVX_VectorPair vp = Q6_W_vshuff_VVR(a_rows[2 * r + 1], a_rows[2 * r + 0], -2); + tile00[r] = Q6_V_lo_W(vp); + } + + for (uint32_t t = 0; t < 32; ++t) { + HVX_Vector v_inv = Q6_V_vzero(); + for (uint32_t k = 0; k < t; ++k) { + HVX_Vector v_lk = Q6_Vh_vsplat_R(l11[t][k]); + v_inv = hvx_vec_sub_f16_f16(v_inv, hvx_vec_mul_f16_f16(v_lk, a_rows[k])); + } + HVX_VectorPred q_diag = (t == 0) ? Q6_Q_vsetq2_R(2) : Q6_Q_and_QQn(Q6_Q_vsetq2_R(2 * (t + 1)), Q6_Q_vsetq2_R(2 * t)); + a_rows[t] = Q6_V_vand_QV(q_mask64, Q6_V_vmux_QVV(q_diag, v_one_f16, v_inv)); + } + + for (uint32_t r = 0; r < 16; ++r) { + HVX_VectorPair vp = Q6_W_vshuff_VVR(a_rows[2 * r + 1], a_rows[2 * r + 0], -2); + tile11[r] = Q6_V_lo_W(vp); + + HVX_Vector n0 = hvx_vec_sub_f16_f16(Q6_V_vzero(), a_rows[2 * r + 0]); + HVX_Vector n1 = hvx_vec_sub_f16_f16(Q6_V_vzero(), a_rows[2 * r + 1]); + HVX_VectorPair vp_neg = Q6_W_vshuff_VVR(n1, n0, -2); + p_neg_a11[r] = Q6_V_lo_W(vp_neg); + } +} + + +static inline void gdn_dma_push_chunk_inputs( + dma_queue * dma_q, + float * vtcm_q, + float * vtcm_k, + float * vtcm_v, + const struct htp_tensor * q, + const struct htp_tensor * k, + const struct htp_tensor * v, + uint32_t iq3, uint32_t iq1, + uint32_t ik3, uint32_t ik1, + uint32_t iv3, uint32_t iv1, + uint32_t t_chunk, + uint32_t chunk_size, + uint32_t S_v +) { + const dma_addr_t q_dma = q->data + (uint64_t) iq3 * q->nb[3] + (uint64_t) t_chunk * q->nb[2] + (uint64_t) iq1 * q->nb[1]; + const dma_addr_t k_dma = k->data + (uint64_t) ik3 * k->nb[3] + (uint64_t) t_chunk * k->nb[2] + (uint64_t) ik1 * k->nb[1]; + const dma_addr_t v_dma = v->data + (uint64_t) iv3 * v->nb[3] + (uint64_t) t_chunk * v->nb[2] + (uint64_t) iv1 * v->nb[1]; + + dma_queue_push(dma_q, dma_make_data(vtcm_q, q_dma), S_v * sizeof(float), q->nb[2], S_v * sizeof(float), chunk_size); + dma_queue_push(dma_q, dma_make_data(vtcm_k, k_dma), S_v * sizeof(float), k->nb[2], S_v * sizeof(float), chunk_size); + dma_queue_push(dma_q, dma_make_data(vtcm_v, v_dma), S_v * sizeof(float), v->nb[2], S_v * sizeof(float), chunk_size); +} + +static inline void gdn_dma_push_chunk_gb( + dma_queue * dma_q, + float * vtcm_g_raw, + float * vtcm_b_raw, + const struct htp_tensor * g, + const struct htp_tensor * beta, + uint32_t iv3, + uint32_t iv1, + uint32_t t_chunk, + uint32_t chunk_size, + uint32_t n_batch +) { + const dma_addr_t g_dma = g->data + (uint64_t) iv3 * g->nb[3] + (uint64_t) t_chunk * g->nb[2] + (uint64_t) iv1 * g->nb[1]; + const dma_addr_t beta_dma = beta->data + (uint64_t) iv3 * beta->nb[3] + (uint64_t) t_chunk * beta->nb[2] + (uint64_t) iv1 * beta->nb[1]; + const uint32_t row_bytes = n_batch * sizeof(float); + + dma_queue_push(dma_q, dma_make_data(vtcm_g_raw, g_dma), row_bytes, g->nb[2], row_bytes, chunk_size); + dma_queue_push(dma_q, dma_make_data(vtcm_b_raw, beta_dma), row_bytes, beta->nb[2], row_bytes, chunk_size); +} + +static inline void gdn_pack_s_col_tiles( + __fp16 * restrict vtcm_s_col_tiles, + __fp16 * restrict vtcm_s_f16, + const float * restrict vtcm_s_state, + uint32_t S_v +) { + for (uint32_t j = 0; j < S_v; ++j) { + for (uint32_t i = 0; i < S_v; i += 64) { + HVX_Vector v0 = hvx_vmem(vtcm_s_state + j * S_v + i + 0); + HVX_Vector v1 = (i + 32 < S_v) ? hvx_vmem(vtcm_s_state + j * S_v + i + 32) : Q6_V_vzero(); + hvx_vmem(vtcm_s_f16 + j * S_v + i) = hvx_vec_f32_to_f16(v0, v1); + } + } + hmx_interleave_rows_to_tiles(vtcm_s_col_tiles, vtcm_s_f16, S_v, S_v, S_v, 0, S_v); +} + +struct htp_gdn_head_ptrs { + float * s_state; + __fp16 * s_f16; + __fp16 * s_col_tiles; + float * s_update_f32; + __fp16 * s_update_tiles; + + float * q_f32[2]; + float * k_f32[2]; + float * v_f32[2]; + float * g_f32[2]; + float * b_f32[2]; + float * o_f32[2]; + + float * v_inter_f32; + float * o_inter_f32; + float * o_intra_f32; + + __fp16 * k_f16; + __fp16 * v_prime_f16; + __fp16 * delta_f16; + __fp16 * d_f16; + + __fp16 * q_row_tiles; + __fp16 * q_prime_row_tiles; + __fp16 * k_row_tiles; + __fp16 * k_col_tiles; + __fp16 * k_prime_row_tiles; + __fp16 * k_col_tiles_64x128; + __fp16 * kk_tiles; + __fp16 * qk_tiles; + __fp16 * v_inter_tiles; + __fp16 * o_inter_tiles; + __fp16 * inv_row_tiles; + __fp16 * a_row_tiles; + __fp16 * v_prime_col_tiles; + __fp16 * delta_tiles; + __fp16 * delta_col_tiles; + __fp16 * o_intra_tiles; + __fp16 * d_row_tiles; + + __fp16 * gamma; + float * lambda_init; + __fp16 * lambda_init_f16; + __fp16 * decay_m; + __fp16 * decay_a; + + HVX_Vector * rows_kk; + HVX_Vector * rows_qk; + HVX_Vector * rows_inv; + HVX_Vector * rows_a; + + HVX_Vector * vtcm_m; + HVX_Vector * vtcm_tmp; + + uint32_t iv1; + uint32_t iv3; + uint32_t iq1; + uint32_t ik1; + uint32_t iq3; + uint32_t ik3; + dma_addr_t state_in_dma; + dma_addr_t state_out_dma; +}; + +static inline void gdn_init_head_ptrs( + struct htp_gdn_head_ptrs * head, + const struct htp_gdn_hmx_vtcm_layout * L, + uint8_t * vtcm_base, + uint32_t h, + uint32_t base_iv1, + uint32_t iv3, + const struct htp_tensor * q, + const struct htp_tensor * k, + const struct htp_tensor * v, + const struct htp_tensor * state, + const struct htp_tensor * dst, + const struct htp_tensor * dst_cache, + const struct htp_gdn_kernel_params * kparams, + uint32_t S_v, + uint32_t H, + uint32_t n_tokens, + uint32_t chunk_size +) { + const size_t dma_scalar_sz = hex_round_up(chunk_size * sizeof(float), 128); + const size_t decay_sz = 64 * 64 * sizeof(__fp16); + const size_t row_vecs_sz = 64 * 128; + + head->s_state = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_s_state + h * L->state_f32_bytes); + head->s_f16 = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_s_f16 + h * L->state_f16_bytes); + head->s_col_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_s_col_tiles + h * L->state_tiles_bytes); + head->s_update_f32 = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_s_update_f32 + h * L->state_f32_bytes); + head->s_update_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_s_update_tiles + h * L->state_tiles_bytes); + + head->q_f32[0] = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_q_f32[0] + h * L->dma_chunk_bytes); + head->q_f32[1] = L->pipeline ? VTCM_LAYOUT_PTR(float, vtcm_base, L->off_q_f32[1] + h * L->dma_chunk_bytes) : head->q_f32[0]; + head->k_f32[0] = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_k_f32[0] + h * L->dma_chunk_bytes); + head->k_f32[1] = L->pipeline ? VTCM_LAYOUT_PTR(float, vtcm_base, L->off_k_f32[1] + h * L->dma_chunk_bytes) : head->k_f32[0]; + head->v_f32[0] = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_v_f32[0] + h * L->dma_chunk_bytes); + head->v_f32[1] = L->pipeline ? VTCM_LAYOUT_PTR(float, vtcm_base, L->off_v_f32[1] + h * L->dma_chunk_bytes) : head->v_f32[0]; + head->g_f32[0] = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_g_f32[0] + h * dma_scalar_sz); + head->g_f32[1] = L->pipeline ? VTCM_LAYOUT_PTR(float, vtcm_base, L->off_g_f32[1] + h * dma_scalar_sz) : head->g_f32[0]; + head->b_f32[0] = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_b_f32[0] + h * dma_scalar_sz); + head->b_f32[1] = L->pipeline ? VTCM_LAYOUT_PTR(float, vtcm_base, L->off_b_f32[1] + h * dma_scalar_sz) : head->b_f32[0]; + head->o_f32[0] = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_o_f32[0] + h * L->dma_chunk_bytes); + head->o_f32[1] = L->pipeline ? VTCM_LAYOUT_PTR(float, vtcm_base, L->off_o_f32[1] + h * L->dma_chunk_bytes) : head->o_f32[0]; + + head->v_inter_f32 = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_v_inter_f32 + h * L->dma_chunk_bytes); + head->o_inter_f32 = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_o_inter_f32 + h * L->dma_chunk_bytes); + head->o_intra_f32 = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_o_intra_f32 + h * L->dma_chunk_bytes); + + head->k_f16 = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_k_f16 + h * L->act_f16_bytes); + head->v_prime_f16 = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_v_prime_f16 + h * L->act_f16_bytes); + head->delta_f16 = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_delta_f16 + h * L->act_f16_bytes); + head->d_f16 = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_d_f16 + h * L->act_f16_bytes); + + head->q_row_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_q_row_tiles + h * L->tile_64xSv_bytes); + head->q_prime_row_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_q_prime_row_tiles + h * L->tile_64xSv_bytes); + head->k_row_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_k_row_tiles + h * L->tile_64xSv_bytes); + head->k_col_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_k_col_tiles + h * L->tile_64xSv_bytes); + head->k_prime_row_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_k_prime_row_tiles + h * L->tile_64xSv_bytes); + head->k_col_tiles_64x128 = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_k_col_tiles_64x128 + h * L->tile_64xSv_bytes); + head->kk_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_kk_tiles + h * L->tile_64x64_bytes); + head->qk_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_qk_tiles + h * L->tile_64x64_bytes); + head->v_inter_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_v_inter_tiles + h * L->tile_64xSv_bytes); + head->o_inter_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_o_inter_tiles + h * L->tile_64xSv_bytes); + head->inv_row_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_inv_row_tiles + h * L->tile_64x64_bytes); + head->a_row_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_a_row_tiles + h * L->tile_64x64_bytes); + head->v_prime_col_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_v_prime_col_tiles + h * L->tile_64xSv_bytes); + head->delta_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_delta_tiles + h * L->tile_64xSv_bytes); + head->delta_col_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_delta_col_tiles + h * L->tile_64xSv_bytes); + head->o_intra_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_o_intra_tiles + h * L->tile_64xSv_bytes); + head->d_row_tiles = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_d_row_tiles + h * L->tile_64xSv_bytes); + + head->gamma = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_gamma + h * dma_scalar_sz); + head->lambda_init_f16 = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_gamma + h * dma_scalar_sz + 128); + head->lambda_init = VTCM_LAYOUT_PTR(float, vtcm_base, L->off_lambda_init + h * dma_scalar_sz); + head->decay_m = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_decay_m + h * decay_sz); + head->decay_a = VTCM_LAYOUT_PTR(__fp16, vtcm_base, L->off_decay_a + h * decay_sz); + + head->rows_kk = VTCM_LAYOUT_PTR(HVX_Vector, vtcm_base, L->off_rows_kk + h * row_vecs_sz); + head->rows_qk = VTCM_LAYOUT_PTR(HVX_Vector, vtcm_base, L->off_rows_qk + h * row_vecs_sz); + head->rows_inv = VTCM_LAYOUT_PTR(HVX_Vector, vtcm_base, L->off_rows_inv + h * row_vecs_sz); + head->rows_a = VTCM_LAYOUT_PTR(HVX_Vector, vtcm_base, L->off_rows_a + h * row_vecs_sz); + + head->vtcm_m = VTCM_LAYOUT_PTR(HVX_Vector, vtcm_base, L->off_thread_scratch + h * (64 * 128)); + head->vtcm_tmp = head->vtcm_m + 32; + + head->iv1 = base_iv1 + h; + head->iv3 = iv3; + head->iq1 = fastmodulo(head->iv1, q->ne[1], &kparams->div_q1); + head->ik1 = fastmodulo(head->iv1, k->ne[1], &kparams->div_k1); + head->iq3 = fastdiv(head->iv3, &kparams->div_rq3); + head->ik3 = fastdiv(head->iv3, &kparams->div_rk3); + + head->state_in_dma = state->data + + ((uint64_t) head->iv3 * kparams->state_seq_stride + (uint64_t) head->iv1 * S_v * S_v) * sizeof(float); + + head->state_out_dma = dst_cache ? + (dst_cache->data + ((uint64_t) head->iv3 * H + head->iv1) * S_v * S_v * sizeof(float)) : + (dst->data + ((uint64_t) S_v * H * n_tokens * kparams->n_seqs + (uint64_t) (head->iv3 * H + head->iv1) * S_v * S_v) * sizeof(float)); +} + +struct htp_gdn_batch_context { + struct htp_gdn_head_ptrs * heads; + const float * vtcm_g_raw; + const float * vtcm_b_raw; + uint32_t curr_buf; + uint32_t c; + uint32_t n_batch; + uint32_t S_v; + float scale; + struct htp_ops_context * octx; + const struct htp_gdn_kernel_params * kparams; +}; + +static void gdn_hvx_init_state_worker(unsigned int n, unsigned int i, void * data) { + (void) n; + struct htp_gdn_batch_context * bctx = (struct htp_gdn_batch_context *) data; + struct htp_thread_trace * tr = &bctx->octx->ctx->trace[i]; + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, 0); + + struct htp_gdn_head_ptrs * head = &bctx->heads[i]; + gdn_pack_s_col_tiles(head->s_col_tiles, head->s_f16, head->s_state, bctx->S_v); + + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, 0); +} + +static inline __attribute__((unused)) HVX_Vector hvx_clamp_neg20_0(HVX_Vector v, HVX_Vector v_zero, HVX_Vector v_neg20) { + HVX_VectorPred p_gt = Q6_Q_vcmp_gt_VsfVsf(v, v_zero); + v = Q6_V_vmux_QVV(p_gt, v_zero, v); + HVX_VectorPred p_lt = Q6_Q_vcmp_gt_VsfVsf(v_neg20, v); + return Q6_V_vmux_QVV(p_lt, v_neg20, v); +} + +static inline HVX_Vector hvx_prefix_scan_f32(HVX_Vector v, HVX_Vector carry_in) { + const HVX_Vector zero = Q6_V_vzero(); + + v = hvx_vec_add_f32_f32(v, Q6_V_vlalign_VVR(v, zero, 4)); + v = hvx_vec_add_f32_f32(v, Q6_V_vlalign_VVR(v, zero, 8)); + v = hvx_vec_add_f32_f32(v, Q6_V_vlalign_VVR(v, zero, 16)); + v = hvx_vec_add_f32_f32(v, Q6_V_vlalign_VVR(v, zero, 32)); + v = hvx_vec_add_f32_f32(v, Q6_V_vlalign_VVR(v, zero, 64)); + v = hvx_vec_add_f32_f32(v, carry_in); + + return v; +} + +static inline HVX_Vector hvx_splat_last_f32(HVX_Vector v) { + return hvx_vec_repl4(Q6_V_vror_VR(v, 124)); +} + +static void gdn_hvx_phase1a_worker(unsigned int n, unsigned int i, void * data) { + (void) n; + struct htp_gdn_batch_context * bctx = (struct htp_gdn_batch_context *) data; + struct htp_thread_trace * tr = &bctx->octx->ctx->trace[i]; + const uint16_t info = (uint16_t) bctx->c; + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_GDN_PREP, info); + + struct htp_gdn_head_ptrs * head = &bctx->heads[i]; + const uint32_t curr_buf = bctx->curr_buf; + const uint32_t S_v = bctx->S_v; + const uint32_t n_batch = bctx->n_batch; + + if (n_batch == 1) { + hvx_vmem(head->g_f32[curr_buf] + 0) = hvx_vmem(bctx->vtcm_g_raw + 0); + hvx_vmem(head->g_f32[curr_buf] + 32) = hvx_vmem(bctx->vtcm_g_raw + 32); + hvx_vmem(head->b_f32[curr_buf] + 0) = hvx_vmem(bctx->vtcm_b_raw + 0); + hvx_vmem(head->b_f32[curr_buf] + 32) = hvx_vmem(bctx->vtcm_b_raw + 32); + } else { + int32_t offsets[32] __attribute__((aligned(128))); + for (int k = 0; k < 32; ++k) { + offsets[k] = k * n_batch * sizeof(float); + } + HVX_Vector vv = *(const HVX_Vector *) offsets; + const size_t rt_g = (size_t) ((const uint8_t *) bctx->vtcm_g_raw + i * sizeof(float)); + const size_t rt_b = (size_t) ((const uint8_t *) bctx->vtcm_b_raw + i * sizeof(float)); + const size_t mu = 64 * n_batch * sizeof(float); + + Q6_vgather_ARMVw((HVX_Vector *) (head->g_f32[curr_buf] + 0), rt_g, mu, vv); + Q6_vgather_ARMVw((HVX_Vector *) (head->g_f32[curr_buf] + 32), rt_g + 32 * n_batch * sizeof(float), mu, vv); + Q6_vgather_ARMVw((HVX_Vector *) (head->b_f32[curr_buf] + 0), rt_b, mu, vv); + Q6_vgather_ARMVw((HVX_Vector *) (head->b_f32[curr_buf] + 32), rt_b + 32 * n_batch * sizeof(float), mu, vv); + } + + const uint32_t t_chunk = bctx->c * 64; + const uint32_t valid_tokens = hex_smin(64, bctx->kparams->n_tokens - t_chunk); + if (valid_tokens < 64) { + for (uint32_t t = valid_tokens; t < 64; ++t) { + head->g_f32[curr_buf][t] = 0.0f; + head->b_f32[curr_buf][t] = 0.0f; + } + const HVX_Vector vzero = Q6_V_vzero(); + for (uint32_t t = valid_tokens; t < 64; ++t) { + for (uint32_t j = 0; j < S_v; j += 32) { + hvx_vmem(head->q_f32[curr_buf] + t * S_v + j) = vzero; + hvx_vmem(head->k_f32[curr_buf] + t * S_v + j) = vzero; + hvx_vmem(head->v_f32[curr_buf] + t * S_v + j) = vzero; + } + } + } + + const HVX_Vector v_g0 = hvx_vmem(head->g_f32[curr_buf] + 0); + const HVX_Vector v_g1 = hvx_vmem(head->g_f32[curr_buf] + 32); + + HVX_Vector v_gamma0 = hvx_prefix_scan_f32(v_g0, Q6_V_vzero()); + HVX_Vector v_carry = hvx_splat_last_f32(v_gamma0); + HVX_Vector v_gamma1 = hvx_prefix_scan_f32(v_g1, v_carry); + + const HVX_Vector v_zero = Q6_V_vzero(); + const HVX_Vector v_neg20 = hvx_vec_splat_f32(-20.0f); + + hvx_vmem(head->gamma) = hvx_vec_f32_to_f16(v_gamma0, v_gamma1); + + HVX_Vector v_l0 = hvx_vec_exp_f32(hvx_clamp_neg20_0(v_gamma0, v_zero, v_neg20)); + HVX_Vector v_l1 = hvx_vec_exp_f32(hvx_clamp_neg20_0(v_gamma1, v_zero, v_neg20)); + + hvx_vmem(head->lambda_init + 0) = v_l0; + hvx_vmem(head->lambda_init + 32) = v_l1; + hvx_vmem(head->lambda_init_f16) = hvx_vec_f32_to_f16(v_l0, v_l1); + + gdn_f32_to_hmx_row_tiles_and_f16(head->k_row_tiles, head->k_prime_row_tiles, head->k_f16, + head->k_f32[curr_buf], head->lambda_init_f16, 64, S_v); + hmx_interleave_rows_to_tiles(head->k_col_tiles, head->k_f16, 64, S_v, S_v, 0, 64); + + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_GDN_PREP, info); +} + +static void gdn_hvx_phase1b_worker(unsigned int n, unsigned int i, void * data) { + (void) n; + struct htp_gdn_batch_context * bctx = (struct htp_gdn_batch_context *) data; + struct htp_thread_trace * tr = &bctx->octx->ctx->trace[i]; + const uint16_t info = (uint16_t) bctx->c; + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_GDN_PREP, info); + + struct htp_gdn_head_ptrs * head = &bctx->heads[i]; + const uint32_t curr_buf = bctx->curr_buf; + const uint32_t S_v = bctx->S_v; + + gdn_f32_to_hmx_row_tiles_and_f16(head->q_row_tiles, head->q_prime_row_tiles, NULL, + head->q_f32[curr_buf], head->lambda_init_f16, 64, S_v); + + hmx_interleave_cols_to_tiles(head->k_col_tiles_64x128, head->k_f16, 64, S_v, S_v, 2, 0, 64); + + const uint16_t * gamma_u16 = (const uint16_t *) head->gamma; + const HVX_Vector v_gamma = hvx_vmem(head->gamma); + + const HVX_Vector v_zero_f16 = Q6_V_vzero(); + const HVX_Vector v_neg20_f16 = hvx_vec_splat_f16(-20.0f); + const HVX_Vector v_log2e_f16 = hvx_vec_splat_f16(1.4426950408889634f); + const HVX_Vector v_one_f16 = hvx_vec_splat_f16(1.0f); + + hvx_vmem(head->decay_m + 0) = Q6_V_vzero(); + hvx_vmem(head->decay_a + 0) = Q6_V_vand_QV(Q6_Q_vsetq2_R(2), v_one_f16); + + for (uint32_t t = 1; t < 63; t += 2) { + uint32_t t0 = t; + uint32_t t1 = t + 1; + + HVX_Vector v_gamma_t0 = Q6_Vh_vsplat_R(gamma_u16[t0]); + HVX_Vector v_gamma_t1 = Q6_Vh_vsplat_R(gamma_u16[t1]); + + HVX_Vector diff0 = hvx_vec_sub_f16_f16(v_gamma_t0, v_gamma); + HVX_Vector diff1 = hvx_vec_sub_f16_f16(v_gamma_t1, v_gamma); + + HVX_VectorPred p_gt0 = Q6_Q_vcmp_gt_VhfVhf(diff0, v_zero_f16); + HVX_VectorPred p_gt1 = Q6_Q_vcmp_gt_VhfVhf(diff1, v_zero_f16); + + diff0 = Q6_V_vmux_QVV(p_gt0, v_zero_f16, diff0); + diff1 = Q6_V_vmux_QVV(p_gt1, v_zero_f16, diff1); + + diff0 = Q6_Vhf_vmax_VhfVhf(v_neg20_f16, diff0); + diff1 = Q6_Vhf_vmax_VhfVhf(v_neg20_f16, diff1); + + HVX_Vector diff_log2e0 = hvx_vec_mul_f16_f16(diff0, v_log2e_f16); + HVX_Vector diff_log2e1 = hvx_vec_mul_f16_f16(diff1, v_log2e_f16); + + HVX_Vector v_exp0 = hvx_vec_exp2_f16(diff_log2e0); + HVX_Vector v_exp1 = hvx_vec_exp2_f16(diff_log2e1); + + HVX_VectorPred mask_lt0 = Q6_Q_vsetq2_R(2 * t0); + HVX_VectorPred mask_lt1 = Q6_Q_vsetq2_R(2 * t1); + + HVX_Vector v_m0 = Q6_V_vand_QV(mask_lt0, v_exp0); + HVX_Vector v_m1 = Q6_V_vand_QV(mask_lt1, v_exp1); + + HVX_VectorPred mask_le0 = Q6_Q_vsetq2_R(2 * (t0 + 1)); + HVX_VectorPred mask_le1 = Q6_Q_vsetq2_R(2 * (t1 + 1)); + + HVX_VectorPred mask_diag0 = Q6_Q_and_QQn(mask_le0, mask_lt0); + HVX_VectorPred mask_diag1 = Q6_Q_and_QQn(mask_le1, mask_lt1); + + HVX_Vector v_a0 = Q6_V_vmux_QVV(mask_diag0, v_one_f16, v_m0); + HVX_Vector v_a1 = Q6_V_vmux_QVV(mask_diag1, v_one_f16, v_m1); + + hvx_vmem(head->decay_m + t0 * 64) = v_m0; + hvx_vmem(head->decay_a + t0 * 64) = v_a0; + hvx_vmem(head->decay_m + t1 * 64) = v_m1; + hvx_vmem(head->decay_a + t1 * 64) = v_a1; + } + + { + HVX_Vector v_gamma_t = Q6_Vh_vsplat_R(gamma_u16[63]); + HVX_Vector diff = hvx_vec_sub_f16_f16(v_gamma_t, v_gamma); + HVX_VectorPred p_gt = Q6_Q_vcmp_gt_VhfVhf(diff, v_zero_f16); + diff = Q6_V_vmux_QVV(p_gt, v_zero_f16, diff); + diff = Q6_Vhf_vmax_VhfVhf(v_neg20_f16, diff); + + HVX_Vector diff_log2e = hvx_vec_mul_f16_f16(diff, v_log2e_f16); + HVX_Vector v_exp = hvx_vec_exp2_f16(diff_log2e); + + HVX_VectorPred mask_lt_t = Q6_Q_vsetq2_R(2 * 63); + HVX_Vector v_m = Q6_V_vand_QV(mask_lt_t, v_exp); + + HVX_VectorPred mask_le_t = Q6_Q_vcmp_eq_VhVh(v_zero_f16, v_zero_f16); + HVX_VectorPred mask_diag = Q6_Q_and_QQn(mask_le_t, mask_lt_t); + HVX_Vector v_a = Q6_V_vmux_QVV(mask_diag, v_one_f16, v_m); + + hvx_vmem(head->decay_m + 63 * 64) = v_m; + hvx_vmem(head->decay_a + 63 * 64) = v_a; + } + + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_GDN_PREP, info); +} + +static void gdn_hvx_phase2_worker(unsigned int n, unsigned int i, void * data) { + (void) n; + struct htp_gdn_batch_context * bctx = (struct htp_gdn_batch_context *) data; + struct htp_thread_trace * tr = &bctx->octx->ctx->trace[i]; + const uint16_t info = (uint16_t) bctx->c; + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_GDN_SOLVE, info); + + struct htp_gdn_head_ptrs * head = &bctx->heads[i]; + const uint32_t curr_buf = bctx->curr_buf; + + gdn_unpack_64x64_tiles_to_vectors(head->rows_kk, head->kk_tiles); + + gdn_build_inv_l_blocks( + head->inv_row_tiles, + head->rows_kk, + head->decay_m, + head->b_f32[curr_buf], + (__fp16 *) head->vtcm_m, + (__fp16 *) head->vtcm_m + HMX_FP16_TILE_N_ELMS + ); + + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_GDN_SOLVE, info); +} + +static void gdn_hvx_phase3_worker(unsigned int n, unsigned int i, void * data) { + (void) n; + struct htp_gdn_batch_context * bctx = (struct htp_gdn_batch_context *) data; + struct htp_thread_trace * tr = &bctx->octx->ctx->trace[i]; + const uint16_t info = (uint16_t) bctx->c; + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_GDN_V_PREP, info); + + struct htp_gdn_head_ptrs * head = &bctx->heads[i]; + const uint32_t curr_buf = bctx->curr_buf; + const uint32_t S_v = bctx->S_v; + + gdn_unpack_64xS_tiles_to_f32(head->v_inter_f32, head->v_inter_tiles, S_v); + + HVX_VectorAlias local_b[2]; + local_b[0].v = hvx_vmem(head->b_f32[curr_buf] + 0); + local_b[1].v = hvx_vmem(head->b_f32[curr_buf] + 32); + + for (uint32_t t = 0; t < 64; ++t) { + HVX_Vector vb = hvx_vec_splat_f32(local_b[t / 32].fp32[t % 32]); + for (uint32_t j = 0; j < S_v; j += 64) { + HVX_Vector vv0 = hvx_vmem(head->v_f32[curr_buf] + t * S_v + j + 0); + HVX_Vector vv1 = (j + 32 < S_v) ? hvx_vmem(head->v_f32[curr_buf] + t * S_v + j + 32) : Q6_V_vzero(); + HVX_Vector vi0 = hvx_vmem(head->v_inter_f32 + t * S_v + j + 0); + HVX_Vector vi1 = (j + 32 < S_v) ? hvx_vmem(head->v_inter_f32 + t * S_v + j + 32) : Q6_V_vzero(); + + HVX_Vector vp0 = hvx_vec_mul_f32_f32(hvx_vec_sub_f32_f32(vv0, vi0), vb); + HVX_Vector vp1 = hvx_vec_mul_f32_f32(hvx_vec_sub_f32_f32(vv1, vi1), vb); + + hvx_vmem(head->v_prime_f16 + t * S_v + j) = hvx_vec_f32_to_f16(vp0, vp1); + } + } + + hmx_interleave_cols_to_tiles(head->v_prime_col_tiles, head->v_prime_f16, 64, S_v, S_v, 2, 0, 64); + + gdn_unpack_64x64_tiles_to_vectors(head->rows_qk, head->qk_tiles); + for (uint32_t t = 0; t < 64; ++t) { + HVX_Vector v_decay_a = hvx_vmem(head->decay_a + t * 64); + head->rows_a[t] = hvx_vec_mul_f16_f16(head->rows_qk[t], v_decay_a); + } + gdn_pack_64x64_vectors_to_tiles(head->a_row_tiles, head->rows_a); + + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_GDN_V_PREP, info); +} + +static void gdn_hvx_phase4_worker(unsigned int n, unsigned int i, void * data) { + (void) n; + struct htp_gdn_batch_context * bctx = (struct htp_gdn_batch_context *) data; + struct htp_thread_trace * tr = &bctx->octx->ctx->trace[i]; + const uint16_t info = (uint16_t) bctx->c; + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_GDN_D_PREP, info); + + struct htp_gdn_head_ptrs * head = &bctx->heads[i]; + const uint32_t S_v = bctx->S_v; + + gdn_unpack_64xS_tiles_to_f16(head->delta_f16, head->delta_tiles, S_v); + hmx_interleave_cols_to_tiles(head->delta_col_tiles, head->delta_f16, 64, S_v, S_v, 2, 0, 64); + + const uint16_t * decay_last = (const uint16_t *) (head->decay_a + 63 * 64); + const HVX_Vector vzero = Q6_V_vzero(); + + for (uint32_t s = 0; s < 64; ++s) { + HVX_Vector vs = Q6_Vh_vsplat_R(decay_last[s]); + HVX_VectorPred p_zero = Q6_Q_vcmp_eq_VhVh(vs, vzero); + for (uint32_t j = 0; j < S_v; j += 64) { + HVX_Vector vd = hvx_vmem(head->delta_f16 + s * S_v + j); + HVX_Vector prod = hvx_vec_mul_f16_f16(vd, vs); + hvx_vmem(head->d_f16 + s * S_v + j) = Q6_V_vmux_QVV(p_zero, vzero, prod); + } + } + + gdn_pack_d_t_row_tiles(head->d_row_tiles, head->d_f16, S_v, head->vtcm_m, head->vtcm_tmp); + + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_GDN_D_PREP, info); +} + +static void gdn_hvx_phase5_worker(unsigned int n, unsigned int i, void * data) { + (void) n; + struct htp_gdn_batch_context * bctx = (struct htp_gdn_batch_context *) data; + struct htp_thread_trace * tr = &bctx->octx->ctx->trace[i]; + const uint16_t info = (uint16_t) bctx->c; + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_GDN_OUT, info); + + struct htp_gdn_head_ptrs * head = &bctx->heads[i]; + const uint32_t curr_buf = bctx->curr_buf; + const uint32_t S_v = bctx->S_v; + const float scale = bctx->scale; + + gdn_unpack_64xS_tiles_to_f32(head->o_inter_f32, head->o_inter_tiles, S_v); + gdn_unpack_64xS_tiles_to_f32(head->o_intra_f32, head->o_intra_tiles, S_v); + + HVX_Vector vscale = hvx_vec_splat_f32(scale); + for (uint32_t j = 0; j < 64 * S_v / 32; ++j) { + HVX_Vector vi = hvx_vmem(head->o_inter_f32 + j * 32); + HVX_Vector va = hvx_vmem(head->o_intra_f32 + j * 32); + hvx_vmem(head->o_f32[curr_buf] + j * 32) = hvx_vec_mul_f32_f32(hvx_vec_add_f32_f32(vi, va), vscale); + } + + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_GDN_OUT, info); +} + +static void gdn_hvx_phase6_worker(unsigned int n, unsigned int i, void * data) { + (void) n; + struct htp_gdn_batch_context * bctx = (struct htp_gdn_batch_context *) data; + struct htp_thread_trace * tr = &bctx->octx->ctx->trace[i]; + const uint16_t info = (uint16_t) bctx->c; + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_GDN_STATE, info); + + struct htp_gdn_head_ptrs * head = &bctx->heads[i]; + const uint32_t S_v = bctx->S_v; + const uint32_t c = bctx->c; + const uint32_t n_chunks = bctx->kparams->n_chunks; + + gdn_unpack_SxS_tiles_to_f32(head->s_update_f32, head->s_update_tiles, S_v); + + HVX_VectorAlias last_lambda; + last_lambda.v = hvx_vmem(head->lambda_init + 32); + HVX_Vector v_l_final = hvx_vec_splat_f32(last_lambda.fp32[31]); + + for (uint32_t j = 0; j < S_v * S_v / 32; ++j) { + HVX_Vector vs_old = hvx_vmem(head->s_state + j * 32); + HVX_Vector vsu = hvx_vmem(head->s_update_f32 + j * 32); + hvx_vmem(head->s_state + j * 32) = hvx_vec_add_f32_f32(hvx_vec_mul_f32_f32(vs_old, v_l_final), vsu); + } + + if (c + 1 < n_chunks) { + gdn_pack_s_col_tiles(head->s_col_tiles, head->s_f16, head->s_state, S_v); + } + + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_GDN_STATE, info); +} + + +static int gated_delta_net_f32_hmx_chunked( + struct htp_ops_context * octx, + const struct htp_gdn_kernel_params * kparams, + uint32_t row_start, + uint32_t nrows +) { + const struct htp_tensor * q = octx->src[0]; + const struct htp_tensor * k = octx->src[1]; + const struct htp_tensor * v = octx->src[2]; + const struct htp_tensor * g = octx->src[3]; + const struct htp_tensor * beta = octx->src[4]; + const struct htp_tensor * state = octx->src[5]; + const struct htp_tensor * dst = octx->dst; + const struct htp_tensor * dst_cache = octx->dsts[1]; + + const uint32_t S_v = kparams->S_v; + const uint32_t H = kparams->H; + const uint32_t n_tokens = kparams->n_tokens; + const float scale = kparams->scale; + const uint32_t chunk_size = kparams->chunk_size; + const uint32_t n_chunks = kparams->n_chunks; + const uint32_t n_sv_tiles = S_v / 32; + + struct htp_gdn_hmx_vtcm_layout L; + htp_gdn_hmx_vtcm_layout_build(&L, S_v, chunk_size, kparams->n_heads_batch, kparams->n_threads, kparams->pipeline != 0); + + if (L.total_bytes > octx->ctx->vtcm_size) { + return HTP_STATUS_VTCM_TOO_SMALL; + } + + uint8_t * const vtcm_base = (uint8_t *) octx->ctx->vtcm_base; + + float * vtcm_g_raw[2] = { + VTCM_LAYOUT_PTR(float, vtcm_base, L.off_g_raw[0]), + L.pipeline ? VTCM_LAYOUT_PTR(float, vtcm_base, L.off_g_raw[1]) : VTCM_LAYOUT_PTR(float, vtcm_base, L.off_g_raw[0]) + }; + float * vtcm_b_raw[2] = { + VTCM_LAYOUT_PTR(float, vtcm_base, L.off_b_raw[0]), + L.pipeline ? VTCM_LAYOUT_PTR(float, vtcm_base, L.off_b_raw[1]) : VTCM_LAYOUT_PTR(float, vtcm_base, L.off_b_raw[0]) + }; + + uint8_t * vtcm_scales_1 = VTCM_LAYOUT_PTR(uint8_t, vtcm_base, L.off_scales_1); + hmx_init_column_scales(vtcm_scales_1, Q6_V_vsplat_R(0x3c00)); + + hmx_queue_t hmx_q = octx->ctx->hmx_queue; + dma_queue * dma_q = octx->ctx->dma[0]; + work_queue_t wp = octx->ctx->work_queue; + + struct htp_gdn_head_ptrs heads[8]; + struct htp_gdn_hmx_gemm_task gemm_tasks[8][9]; + + uint32_t n_batch = 1; + for (uint32_t r = row_start; r < row_start + nrows; r += n_batch) { + const uint32_t head_in_seq = fastmodulo(r, H, &kparams->div_H); + const uint32_t iv3 = fastdiv(r, &kparams->div_H); + const uint32_t heads_left_in_seq = H - head_in_seq; + const uint32_t heads_left_in_range = (row_start + nrows) - r; + n_batch = hex_smin((uint32_t) kparams->n_heads_batch, hex_smin(heads_left_in_seq, heads_left_in_range)); + + for (uint32_t h = 0; h < n_batch; ++h) { + gdn_init_head_ptrs(&heads[h], &L, vtcm_base, h, head_in_seq, iv3, + q, k, v, state, dst, dst_cache, kparams, S_v, H, n_tokens, chunk_size); + } + + struct htp_gdn_batch_context bctx; + bctx.heads = heads; + bctx.vtcm_g_raw = NULL; + bctx.vtcm_b_raw = NULL; + bctx.curr_buf = 0; + bctx.c = 0; + bctx.n_batch = n_batch; + bctx.S_v = S_v; + bctx.scale = scale; + bctx.octx = octx; + bctx.kparams = kparams; + + for (uint32_t h = 0; h < n_batch; ++h) { + dma_queue_push(dma_q, dma_make_data(heads[h].s_state, heads[h].state_in_dma), + S_v * sizeof(float), S_v * sizeof(float), S_v * sizeof(float), S_v); + } + for (uint32_t h = 0; h < n_batch; ++h) { + dma_queue_pop(dma_q); + } + + if (n_chunks > 0) { + work_queue_run(wp, gdn_hvx_init_state_worker, &bctx, n_batch); + + const uint32_t chunk0_tokens = hex_smin(chunk_size, n_tokens); + for (uint32_t h = 0; h < n_batch; ++h) { + gdn_dma_push_chunk_inputs(dma_q, heads[h].q_f32[0], heads[h].k_f32[0], heads[h].v_f32[0], + q, k, v, heads[h].iq3, heads[h].iq1, heads[h].ik3, heads[h].ik1, + heads[h].iv3, heads[h].iv1, 0, chunk0_tokens, S_v); + } + gdn_dma_push_chunk_gb(dma_q, vtcm_g_raw[0], vtcm_b_raw[0], g, beta, iv3, head_in_seq, 0, chunk0_tokens, n_batch); + } + + for (uint32_t c = 0; c < n_chunks; ++c) { + const uint32_t curr_buf = c & 1; + const uint32_t next_buf = (c + 1) & 1; + const uint32_t t_chunk = c * chunk_size; + + bctx.curr_buf = curr_buf; + bctx.c = c; + bctx.vtcm_g_raw = vtcm_g_raw[curr_buf]; + bctx.vtcm_b_raw = vtcm_b_raw[curr_buf]; + + for (uint32_t h = 0; h < n_batch; ++h) { + dma_queue_pop(dma_q); + dma_queue_pop(dma_q); + dma_queue_pop(dma_q); + } + dma_queue_pop(dma_q); + dma_queue_pop(dma_q); + + if (c + 1 < n_chunks) { + const uint32_t next_t_chunk = (c + 1) * chunk_size; + const uint32_t next_tokens = hex_smin(chunk_size, n_tokens - next_t_chunk); + for (uint32_t h = 0; h < n_batch; ++h) { + gdn_dma_push_chunk_inputs(dma_q, heads[h].q_f32[next_buf], heads[h].k_f32[next_buf], heads[h].v_f32[next_buf], + q, k, v, heads[h].iq3, heads[h].iq1, heads[h].ik3, heads[h].ik1, + heads[h].iv3, heads[h].iv1, next_t_chunk, next_tokens, S_v); + } + gdn_dma_push_chunk_gb(dma_q, vtcm_g_raw[next_buf], vtcm_b_raw[next_buf], + g, beta, iv3, head_in_seq, next_t_chunk, next_tokens, n_batch); + } + + if (c > 0) { + for (uint32_t h = 0; h < n_batch; ++h) { + dma_queue_pop(dma_q); + } + } + + work_queue_run(wp, gdn_hvx_phase1a_worker, &bctx, n_batch); + + for (uint32_t h = 0; h < n_batch; ++h) { + htp_gdn_push_hmx_gemm_task(hmx_q, &gemm_tasks[h][0], heads[h].k_row_tiles, heads[h].k_col_tiles, heads[h].kk_tiles, 2, 2, n_sv_tiles, vtcm_scales_1); + } + for (uint32_t h = 0; h < n_batch; ++h) { + htp_gdn_push_hmx_gemm_task(hmx_q, &gemm_tasks[h][2], heads[h].k_prime_row_tiles, heads[h].s_col_tiles, heads[h].v_inter_tiles, 2, n_sv_tiles, n_sv_tiles, vtcm_scales_1); + } + + work_queue_run(wp, gdn_hvx_phase1b_worker, &bctx, n_batch); + + for (uint32_t h = 0; h < n_batch; ++h) { + hmx_queue_pop(hmx_q); + } + + for (uint32_t h = 0; h < n_batch; ++h) { + htp_gdn_push_hmx_gemm_task(hmx_q, &gemm_tasks[h][1], heads[h].q_row_tiles, heads[h].k_col_tiles, heads[h].qk_tiles, 2, 2, n_sv_tiles, vtcm_scales_1); + } + for (uint32_t h = 0; h < n_batch; ++h) { + htp_gdn_push_hmx_gemm_task(hmx_q, &gemm_tasks[h][3], heads[h].q_prime_row_tiles, heads[h].s_col_tiles, heads[h].o_inter_tiles, 2, n_sv_tiles, n_sv_tiles, vtcm_scales_1); + } + + work_queue_run(wp, gdn_hvx_phase2_worker, &bctx, n_batch); + + for (uint32_t h = 0; h < n_batch; ++h) { + hmx_queue_pop(hmx_q); + } + for (uint32_t h = 0; h < n_batch; ++h) { + hmx_queue_pop(hmx_q); + } + + for (uint32_t h = 0; h < n_batch; ++h) { + htp_gdn_push_hmx_gemm_task( + hmx_q, &gemm_tasks[h][7], + (__fp16 *) heads[h].vtcm_m, + heads[h].inv_row_tiles + 0 * HMX_FP16_TILE_N_ELMS, + (__fp16 *) heads[h].vtcm_tmp, + 1, 1, 1, vtcm_scales_1 + ); + } + for (uint32_t h = 0; h < n_batch; ++h) { + htp_gdn_push_hmx_gemm_task( + hmx_q, &gemm_tasks[h][8], + (__fp16 *) heads[h].vtcm_m + HMX_FP16_TILE_N_ELMS, + (__fp16 *) heads[h].vtcm_tmp, + heads[h].inv_row_tiles + 2 * HMX_FP16_TILE_N_ELMS, + 1, 1, 1, vtcm_scales_1 + ); + } + + work_queue_run(wp, gdn_hvx_phase3_worker, &bctx, n_batch); + + for (uint32_t h = 0; h < n_batch; ++h) { + hmx_queue_pop(hmx_q); + } + for (uint32_t h = 0; h < n_batch; ++h) { + hmx_queue_pop(hmx_q); + } + for (uint32_t h = 0; h < n_batch; ++h) { + hmx_queue_pop(hmx_q); + } + + for (uint32_t h = 0; h < n_batch; ++h) { + htp_gdn_push_hmx_gemm_task(hmx_q, &gemm_tasks[h][4], heads[h].inv_row_tiles, heads[h].v_prime_col_tiles, heads[h].delta_tiles, 2, n_sv_tiles, 2, vtcm_scales_1); + } + + for (uint32_t h = 0; h < n_batch; ++h) { + hmx_queue_pop(hmx_q); + } + + work_queue_run(wp, gdn_hvx_phase4_worker, &bctx, n_batch); + + for (uint32_t h = 0; h < n_batch; ++h) { + htp_gdn_push_hmx_gemm_task(hmx_q, &gemm_tasks[h][5], heads[h].a_row_tiles, heads[h].delta_col_tiles, heads[h].o_intra_tiles, 2, n_sv_tiles, 2, vtcm_scales_1); + } + for (uint32_t h = 0; h < n_batch; ++h) { + htp_gdn_push_hmx_gemm_task(hmx_q, &gemm_tasks[h][6], heads[h].d_row_tiles, heads[h].k_col_tiles_64x128, heads[h].s_update_tiles, n_sv_tiles, n_sv_tiles, 2, vtcm_scales_1); + } + + for (uint32_t h = 0; h < n_batch; ++h) { + hmx_queue_pop(hmx_q); + } + + work_queue_run(wp, gdn_hvx_phase5_worker, &bctx, n_batch); + + const uint32_t valid_tokens = hex_smin(chunk_size, n_tokens - t_chunk); + for (uint32_t h = 0; h < n_batch; ++h) { + const dma_addr_t attn_chunk_dma = dst->data + + ((uint64_t) heads[h].iv3 * n_tokens * H + (uint64_t) t_chunk * H + heads[h].iv1) * S_v * sizeof(float); + dma_queue_push(dma_q, dma_make_data(attn_chunk_dma, heads[h].o_f32[curr_buf]), + dst->nb[1], S_v * sizeof(float), S_v * sizeof(float), valid_tokens); + } + + for (uint32_t h = 0; h < n_batch; ++h) { + hmx_queue_pop(hmx_q); + } + + work_queue_run(wp, gdn_hvx_phase6_worker, &bctx, n_batch); + } + + if (n_chunks > 0) { + for (uint32_t h = 0; h < n_batch; ++h) { + dma_queue_pop(dma_q); + } + } + + for (uint32_t h = 0; h < n_batch; ++h) { + dma_queue_push(dma_q, dma_make_data(heads[h].state_out_dma, heads[h].s_state), + S_v * sizeof(float), S_v * sizeof(float), S_v * sizeof(float), S_v); + } + for (uint32_t h = 0; h < n_batch; ++h) { + dma_queue_pop(dma_q); + } + } + + dma_queue_flush(dma_q); + return HTP_STATUS_OK; +} + int op_gated_delta_net(struct htp_ops_context * octx) { const struct htp_tensor * q = octx->src[0]; const struct htp_tensor * k = octx->src[1]; @@ -1097,11 +2344,34 @@ int op_gated_delta_net(struct htp_ops_context * octx) { kparams_local.K = K; kparams_local.total_rows = total_rows; kparams_local.rows_per_thread = (total_rows + n_threads - 1) / n_threads; - struct htp_gdn_vtcm_layout layout_local; - htp_gdn_vtcm_layout_build(&layout_local, S_v, n_threads); - kparams_local.state_aligned = (uint32_t) layout_local.state_aligned; - kparams_local.vtcm_per_thread = (uint32_t) layout_local.bytes_per_thread; - kparams_local.vtcm_size = (uint32_t) layout_local.total_bytes; + const bool can_use_hmx = (octx->ctx->hmx_enabled) && + (S_v % 64 == 0) && + (n_tokens >= HTP_GDN_MIN_TOKENS) && + (g->ne[0] == 1) && + (K == 1); + + struct htp_gdn_hmx_vtcm_layout hmx_layout_local; + struct htp_gdn_vtcm_layout hvx_layout_local; + uint32_t n_heads_batch = 1; + + if (can_use_hmx && htp_gdn_hmx_solve_layout(&hmx_layout_local, S_v, HTP_GDN_CHUNK_SIZE, total_rows, octx->ctx->vtcm_size, n_threads, true, &n_heads_batch)) { + kparams_local.kernel_type = HTP_GDN_KERNEL_HMX_CHUNKED; + kparams_local.pipeline = hmx_layout_local.pipeline ? 1 : 0; + kparams_local.chunk_size = HTP_GDN_CHUNK_SIZE; + kparams_local.n_chunks = (n_tokens + HTP_GDN_CHUNK_SIZE - 1) / HTP_GDN_CHUNK_SIZE; + kparams_local.n_heads_batch = (uint16_t) n_heads_batch; + kparams_local.vtcm_size = (uint32_t) hmx_layout_local.total_bytes; + kparams_local.state_aligned = (uint32_t) hmx_layout_local.state_f32_bytes; + kparams_local.vtcm_per_thread = (uint32_t) (hmx_layout_local.total_bytes / (n_threads > 0 ? n_threads : 1)); + } else { + htp_gdn_vtcm_layout_build(&hvx_layout_local, S_v, n_threads); + kparams_local.kernel_type = HTP_GDN_KERNEL_HVX_RECURRENT; + kparams_local.pipeline = 0; + kparams_local.n_heads_batch = 1; + kparams_local.state_aligned = (uint32_t) hvx_layout_local.state_aligned; + kparams_local.vtcm_per_thread = (uint32_t) hvx_layout_local.bytes_per_thread; + kparams_local.vtcm_size = (uint32_t) hvx_layout_local.total_bytes; + } kparams_local.kda = (g->ne[0] == S_v) ? 1 : 0; kparams_local.scale = 1.0f / sqrtf((float) S_v); kparams_local.state_seq_stride = (uint32_t) (state->nb[3] / sizeof(float)); @@ -1121,7 +2391,19 @@ int op_gated_delta_net(struct htp_ops_context * octx) { uint32_t row_start = 0; uint32_t nrows = total_rows; - if (octx->op_params[1] != 0) { + if (octx->ctx->mdev.count > 1) { + const bool can_split = htp_tensor_mdev_data_aligned(dst) && + ((dst->nb[1] & (HTP_TENSOR_MDEV_LINE_SIZE - 1)) == 0); + const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition( + total_rows, + can_split ? 1 : 0, + octx->ctx->mdev.idx, + octx->ctx->mdev.count, + &octx->ctx->mdev.count_div + ); + row_start = range.start; + nrows = range.count; + } else if (octx->op_params[1] != 0) { row_start = octx->op_params[1]; nrows = octx->op_params[2]; } @@ -1130,6 +2412,10 @@ int op_gated_delta_net(struct htp_ops_context * octx) { return HTP_STATUS_OK; } + if (kparams->kernel_type == HTP_GDN_KERNEL_HMX_CHUNKED) { + return gated_delta_net_f32_hmx_chunked(octx, kparams, row_start, nrows); + } + const uint32_t n_threads = (nrows < kparams->n_threads) ? nrows : kparams->n_threads; struct htp_gdn_context gctx; diff --git a/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.h b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.h index fd703142e..32fb7d24b 100644 --- a/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.h +++ b/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.h @@ -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; } diff --git a/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h b/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h index 8fd299795..d5fb48ad9 100644 --- a/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +++ b/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h @@ -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, diff --git a/ggml/src/ggml-hexagon/htp/htp-ops.h b/ggml/src/ggml-hexagon/htp/htp-ops.h index 0e63febdd..ee5b92441 100644 --- a/ggml/src/ggml-hexagon/htp/htp-ops.h +++ b/ggml/src/ggml-hexagon/htp/htp-ops.h @@ -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, }; diff --git a/ggml/src/ggml-hexagon/htp/main.c b/ggml/src/ggml-hexagon/htp/main.c index 653c9a250..b4b352b20 100644 --- a/ggml/src/ggml-hexagon/htp/main.c +++ b/ggml/src/ggml-hexagon/htp/main.c @@ -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