mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-08 22:10:37 +02:00
hexagon: new HMX-optimized GATED_DELTA_NET (#29199)
* hex-gdn: start putting together HMX support for GDN * hex-gdn: working hmx but not-pipelined and slow for now * hex-gdn: re-write vtcm layout handling and prep for pipelining * hex-gdn: starting to pipeline hmx and dmas * hex-gdn: add hvx threading for most pipeline stages * hex-gdb: add detailed trace events * hex-gdn: vectorize expfs and use aligned hvx reads/writes * hex-gnd: vectorize the rest of expf * hex-gdn: optimize tail processing (pad partial chunks) * hex-gdb: avoid float up/down casts in hot loops * hex-fa: remove float up/down casts from inner loops * hex-gdn: do exp() in f16 to improve HVX utilization * hex-gdn: optimize tiler * hex-hmx: bump hmx-queue to 128 and dispatch all GDN gemms at once * hex-gdn: further pipeline improvements * hex-gdn: optimize gdn prep stage * hex-gdn: yet more tweaks to optimize GND_SOLVE task and pipeline * hex-gdn: improve accuracy and optmize gdn-prep further * hex-gdn: fix rebase conflict * hex-bufs: revert max_bufsize enforcement, it is enough to just enforce max_vmem * hex-scripts: improved inspect script to avoid false alarms in reg spill detector * hex-fa: improve inline softmax with in-reg VKQ32 accum * hex-fa: minor improvement for dma pipeline in hvx kernel * hex-fa: reduce ddr reads by 20-30% during token gen * hex-gdn: proper alignment for hvx vtcm spads
This commit is contained in:
@@ -100,6 +100,7 @@ static bool opt_dma64 = false;
|
||||
|
||||
static int opt_mm_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
|
||||
static int opt_fa_select = 2; // 2 = HMX -> HVX -> CPU, 1 = HVX -> CPU, 0 = CPU (unsupported)
|
||||
static int opt_gdn_select = 2; // 2 = HMX -> HVX, 1 = HVX, 0 = CPU (unsupported)
|
||||
static int opt_ar_select = 2; // 2 = fused ALLREDUCE+ADD (DMA, default), 1 = unfused ALLREDUCE (DMA), 0 = fallback to CPY+FENCE
|
||||
|
||||
// Default PMU events, if profiling with PMU (mode=2) is enabled
|
||||
@@ -182,6 +183,13 @@ static const char * htp_event_name(uint16_t id) {
|
||||
case HTP_TRACE_EVT_HVX_FA_Q_PREP: return "HVX_Q_PREP";
|
||||
case HTP_TRACE_EVT_HVX_FA_K_PREP: return "HVX_K_PREP";
|
||||
case HTP_TRACE_EVT_HVX_FA_V_PREP: return "HVX_V_PREP";
|
||||
case HTP_TRACE_EVT_HVX_GDN_PREP: return "HVX_GDN_PREP";
|
||||
case HTP_TRACE_EVT_HVX_GDN_SOLVE: return "HVX_GDN_SOLVE";
|
||||
case HTP_TRACE_EVT_HVX_GDN_V_PREP: return "HVX_GDN_V_PREP";
|
||||
case HTP_TRACE_EVT_HVX_GDN_D_PREP: return "HVX_GDN_D_PREP";
|
||||
case HTP_TRACE_EVT_HVX_GDN_OUT: return "HVX_GDN_OUT";
|
||||
case HTP_TRACE_EVT_HVX_GDN_STATE: return "HVX_GDN_STATE";
|
||||
case HTP_TRACE_EVT_HVX_GDN_REM: return "HVX_GDN_REM";
|
||||
case HTP_TRACE_EVT_HMX_COMP: return "HMX_COMP";
|
||||
case HTP_TRACE_EVT_L2FLUSH: return "L2FLUSH";
|
||||
case HTP_TRACE_EVT_INIT: return "INIT";
|
||||
@@ -472,7 +480,6 @@ struct ggml_hexagon_session {
|
||||
uint32_t n_hmx = 0;
|
||||
uint64_t vtcm_size = 0;
|
||||
size_t max_vmem = 0;
|
||||
size_t max_bufsize = 0;
|
||||
uint32_t fence_seq = 0;
|
||||
|
||||
std::atomic<uint64_t> batch_req_seq{0};
|
||||
@@ -538,7 +545,6 @@ struct ggml_backend_hexagon_device_context {
|
||||
int dev_id;
|
||||
ggml_hexagon_device_config config;
|
||||
ggml_backend_dev_t dev = nullptr;
|
||||
size_t max_bufsize = 0;
|
||||
|
||||
ggml_backend_buffer_type buffer_type = {};
|
||||
ggml_backend_buffer_type host_buffer_type = {};
|
||||
@@ -554,9 +560,6 @@ struct ggml_backend_hexagon_device_context {
|
||||
ggml_hexagon_session * session() {
|
||||
if (!sess) {
|
||||
sess = std::make_unique<ggml_hexagon_session>(config, dev);
|
||||
if (max_bufsize > sess->max_vmem) {
|
||||
max_bufsize = sess->max_vmem;
|
||||
}
|
||||
}
|
||||
return sess.get();
|
||||
}
|
||||
@@ -2076,11 +2079,6 @@ static const char * ggml_backend_hexagon_buffer_type_name(ggml_backend_buffer_ty
|
||||
static ggml_backend_buffer_t ggml_backend_hexagon_buffer_type_alloc_buffer(
|
||||
ggml_backend_buffer_type_t buffer_type, size_t size) {
|
||||
auto dev_ctx = static_cast<ggml_backend_hexagon_buffer_type_context *>(buffer_type->context)->dev_ctx;
|
||||
if (size > dev_ctx->max_bufsize) {
|
||||
GGML_LOG_ERROR("ggml-hex: %s buffer size %zu exceeds max_bufsize %zu\n",
|
||||
dev_ctx->c_name(), size, dev_ctx->max_bufsize);
|
||||
return nullptr;
|
||||
}
|
||||
auto sess = dev_ctx->session();
|
||||
if (sess && sess->max_vmem && size > sess->max_vmem) {
|
||||
GGML_LOG_ERROR("ggml-hex: %s buffer size %zu exceeds max_vmem %zu\n",
|
||||
@@ -2099,11 +2097,6 @@ static ggml_backend_buffer_t ggml_backend_hexagon_buffer_type_alloc_buffer(
|
||||
static ggml_backend_buffer_t ggml_backend_hexagon_host_buffer_type_alloc_buffer(
|
||||
ggml_backend_buffer_type_t buffer_type, size_t size) {
|
||||
auto dev_ctx = static_cast<ggml_backend_hexagon_buffer_type_context *>(buffer_type->context)->dev_ctx;
|
||||
if (size > dev_ctx->max_bufsize) {
|
||||
GGML_LOG_ERROR("ggml-hex: %s host buffer size %zu exceeds max_bufsize %zu\n",
|
||||
dev_ctx->c_name(), size, dev_ctx->max_bufsize);
|
||||
return nullptr;
|
||||
}
|
||||
auto sess = dev_ctx->session();
|
||||
if (sess && sess->max_vmem && size > sess->max_vmem) {
|
||||
GGML_LOG_ERROR("ggml-hex: %s host buffer size %zu exceeds max_vmem %zu\n",
|
||||
@@ -2138,10 +2131,8 @@ static size_t ggml_backend_hexagon_buffer_type_get_alloc_size(ggml_backend_buffe
|
||||
}
|
||||
|
||||
static size_t ggml_backend_hexagon_buffer_type_get_max_size(ggml_backend_buffer_type_t buft) {
|
||||
auto * context = static_cast<ggml_backend_hexagon_buffer_type_context *>(buft->context);
|
||||
auto dev_ctx = context->dev_ctx;
|
||||
dev_ctx->session();
|
||||
return dev_ctx->max_bufsize;
|
||||
return opt_mbuf;
|
||||
GGML_UNUSED(buft);
|
||||
}
|
||||
|
||||
static bool ggml_backend_hexagon_buffer_type_is_host(ggml_backend_buffer_type_t buft) {
|
||||
@@ -2173,7 +2164,7 @@ static ggml_backend_buffer_type_i ggml_backend_hexagon_host_buffer_type_interfac
|
||||
};
|
||||
|
||||
ggml_backend_hexagon_device_context::ggml_backend_hexagon_device_context(int dev_id, const ggml_hexagon_device_config & config, ggml_backend_dev_t dev)
|
||||
: dev_id(dev_id), config(config), dev(dev), max_bufsize(opt_mbuf) {
|
||||
: dev_id(dev_id), config(config), dev(dev) {
|
||||
buffer_type.device = dev;
|
||||
buffer_type.iface = ggml_backend_hexagon_buffer_type_interface;
|
||||
buffer_type.context = new ggml_backend_hexagon_buffer_type_context(config.name, this);
|
||||
@@ -3927,7 +3918,6 @@ void ggml_hexagon_session::allocate(const ggml_hexagon_device_config & config) n
|
||||
this->valid_handle = true;
|
||||
|
||||
// Query HW info and resolve session options
|
||||
this->max_bufsize = opt_mbuf;
|
||||
{
|
||||
unsigned int hw_n_threads = 0;
|
||||
unsigned int hw_n_hvx = 0;
|
||||
@@ -4340,6 +4330,10 @@ static bool ggml_hexagon_supported_flash_attn_ext(const struct ggml_hexagon_sess
|
||||
}
|
||||
|
||||
static bool ggml_hexagon_supported_gated_delta_net(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
|
||||
if (opt_gdn_select < 1) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const struct ggml_tensor * q = op->src[0];
|
||||
const struct ggml_tensor * k = op->src[1];
|
||||
const struct ggml_tensor * v = op->src[2];
|
||||
@@ -4387,10 +4381,26 @@ static bool ggml_hexagon_supported_gated_delta_net(const struct ggml_hexagon_ses
|
||||
|
||||
const uint32_t total_rows = (uint32_t) (H * n_seqs);
|
||||
const uint32_t n_threads = (std::min)((uint32_t) sess->n_threads, total_rows);
|
||||
struct htp_gdn_vtcm_layout layout;
|
||||
htp_gdn_vtcm_layout_build(&layout, (uint32_t) S_v, n_threads ? n_threads : 1);
|
||||
if (layout.total_bytes > sess->vtcm_size) {
|
||||
return false;
|
||||
|
||||
const bool can_use_hmx = (opt_gdn_select >= 2) &&
|
||||
(sess->n_hmx > 0) &&
|
||||
(S_v % 64 == 0) &&
|
||||
(n_tokens >= HTP_GDN_MIN_TOKENS) &&
|
||||
(g->ne[0] == 1) &&
|
||||
(K == 1);
|
||||
|
||||
if (can_use_hmx) {
|
||||
struct htp_gdn_hmx_vtcm_layout layout;
|
||||
uint32_t n_heads_batch = 0;
|
||||
if (!htp_gdn_hmx_solve_layout(&layout, (uint32_t) S_v, HTP_GDN_CHUNK_SIZE, total_rows, sess->vtcm_size, n_threads, true, &n_heads_batch)) {
|
||||
return false;
|
||||
}
|
||||
} else {
|
||||
struct htp_gdn_vtcm_layout layout;
|
||||
htp_gdn_vtcm_layout_build(&layout, (uint32_t) S_v, n_threads);
|
||||
if (layout.total_bytes > sess->vtcm_size) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
@@ -5206,10 +5216,37 @@ static void ggml_hexagon_precompute_gated_delta_net_params(
|
||||
const uint32_t total_rows = H * n_seqs;
|
||||
const uint32_t n_threads = (std::min)((uint32_t) sess->n_threads, total_rows);
|
||||
|
||||
struct htp_gdn_vtcm_layout layout;
|
||||
htp_gdn_vtcm_layout_build(&layout, S_v, n_threads ? n_threads : 1);
|
||||
const bool can_use_hmx = (opt_gdn_select >= 2) &&
|
||||
(sess->n_hmx > 0) &&
|
||||
(S_v % 64 == 0) &&
|
||||
(n_tokens >= HTP_GDN_MIN_TOKENS) &&
|
||||
(g->ne[0] == 1) &&
|
||||
(K == 1);
|
||||
|
||||
kparams->n_threads = n_threads ? n_threads : 1;
|
||||
struct htp_gdn_hmx_vtcm_layout hmx_layout;
|
||||
struct htp_gdn_vtcm_layout hvx_layout;
|
||||
uint32_t n_heads_batch = 1;
|
||||
|
||||
if (can_use_hmx && htp_gdn_hmx_solve_layout(&hmx_layout, S_v, HTP_GDN_CHUNK_SIZE, total_rows, sess->vtcm_size, n_threads, true, &n_heads_batch)) {
|
||||
kparams->kernel_type = HTP_GDN_KERNEL_HMX_CHUNKED;
|
||||
kparams->pipeline = hmx_layout.pipeline ? 1 : 0;
|
||||
kparams->chunk_size = HTP_GDN_CHUNK_SIZE;
|
||||
kparams->n_chunks = (n_tokens + HTP_GDN_CHUNK_SIZE - 1) / HTP_GDN_CHUNK_SIZE;
|
||||
kparams->n_heads_batch = (uint16_t) n_heads_batch;
|
||||
kparams->vtcm_size = (uint32_t) hmx_layout.total_bytes;
|
||||
kparams->state_aligned = (uint32_t) hmx_layout.state_f32_bytes;
|
||||
kparams->vtcm_per_thread = (uint32_t) (hmx_layout.total_bytes / (n_threads > 0 ? n_threads : 1));
|
||||
} else {
|
||||
htp_gdn_vtcm_layout_build(&hvx_layout, S_v, n_threads);
|
||||
kparams->kernel_type = HTP_GDN_KERNEL_HVX_RECURRENT;
|
||||
kparams->pipeline = 0;
|
||||
kparams->n_heads_batch = 1;
|
||||
kparams->state_aligned = (uint32_t) hvx_layout.state_aligned;
|
||||
kparams->vtcm_per_thread = (uint32_t) hvx_layout.bytes_per_thread;
|
||||
kparams->vtcm_size = (uint32_t) hvx_layout.total_bytes;
|
||||
}
|
||||
|
||||
kparams->n_threads = n_threads;
|
||||
kparams->S_v = S_v;
|
||||
kparams->H = H;
|
||||
kparams->n_tokens = n_tokens;
|
||||
@@ -5218,9 +5255,6 @@ static void ggml_hexagon_precompute_gated_delta_net_params(
|
||||
kparams->total_rows = total_rows;
|
||||
kparams->rows_per_thread = (total_rows + kparams->n_threads - 1) / kparams->n_threads;
|
||||
kparams->kda = (g->ne[0] == S_v) ? 1 : 0;
|
||||
kparams->state_aligned = (uint32_t) layout.state_aligned;
|
||||
kparams->vtcm_per_thread = (uint32_t) layout.bytes_per_thread;
|
||||
kparams->vtcm_size = (uint32_t) layout.total_bytes;
|
||||
kparams->state_seq_stride = (uint32_t) (state->nb[3] / sizeof(float));
|
||||
kparams->state_size_per_snap = S_v * S_v * H * n_seqs;
|
||||
kparams->scale = 1.0f / sqrtf((float) S_v);
|
||||
@@ -7731,6 +7765,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
|
||||
const char * str_nhmx = getenv("GGML_HEXAGON_NHMX");
|
||||
const char * str_mm_select = getenv("GGML_HEXAGON_MM_SELECT");
|
||||
const char * str_fa_select = getenv("GGML_HEXAGON_FA_SELECT");
|
||||
const char * str_gdn_select = getenv("GGML_HEXAGON_GDN_SELECT");
|
||||
const char * str_ar_select = getenv("GGML_HEXAGON_AR_SELECT");
|
||||
const char * str_ndev = getenv("GGML_HEXAGON_NDEV");
|
||||
const char * str_arch = getenv("GGML_HEXAGON_ARCH");
|
||||
@@ -7783,6 +7818,7 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
|
||||
opt_nhmx = str_nhmx ? atoi(str_nhmx) : opt_nhmx;
|
||||
opt_mm_select = str_mm_select ? atoi(str_mm_select) : opt_mm_select;
|
||||
opt_fa_select = str_fa_select ? atoi(str_fa_select) : opt_fa_select;
|
||||
opt_gdn_select = str_gdn_select ? atoi(str_gdn_select) : opt_gdn_select;
|
||||
opt_ar_select = str_ar_select ? atoi(str_ar_select) : opt_ar_select;
|
||||
opt_mbuf = str_mbuf ? strtoul(str_mbuf, NULL, 0) * MiB : opt_mbuf;
|
||||
opt_vmem = str_vmem ? strtoul(str_vmem, NULL, 0) * MiB : opt_vmem;
|
||||
|
||||
@@ -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 ||
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -247,6 +247,7 @@ static inline size_t hmx_fa_compute_vtcm_usage(size_t gqa_factor, size_t DK, siz
|
||||
}
|
||||
|
||||
#define FA_HVX_BLOCK_SIZE 64
|
||||
#define FA_HVX_G_MAX 8
|
||||
|
||||
struct hvx_fa_vtcm_layout {
|
||||
size_t off_q;
|
||||
@@ -275,11 +276,11 @@ static inline void hvx_fa_vtcm_layout_build(struct hvx_fa_vtcm_layout * L,
|
||||
const size_t size_k_row_padded = hex_round_up(DK * sizeof(__fp16), 128);
|
||||
const size_t size_v_row_padded = hex_round_up(DV * sizeof(__fp16), 128);
|
||||
|
||||
const size_t size_q_block = size_q_row_padded * 1;
|
||||
const size_t size_q_block = size_q_row_padded * FA_HVX_G_MAX;
|
||||
const size_t size_k_block = size_k_row_padded * FA_HVX_BLOCK_SIZE;
|
||||
const size_t size_v_block = size_v_row_padded * FA_HVX_BLOCK_SIZE;
|
||||
const size_t size_m_block = hex_round_up(FA_HVX_BLOCK_SIZE * sizeof(__fp16), 128);
|
||||
const size_t size_vkq_acc = hex_round_up(DV * sizeof(float), 128);
|
||||
const size_t size_vkq_acc = hex_round_up(DV * sizeof(float), 128) * FA_HVX_G_MAX;
|
||||
const size_t size_sinks = hex_round_up(n_heads * sizeof(float), 128);
|
||||
|
||||
size_t off = 0;
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -11,6 +11,7 @@
|
||||
|
||||
#define HTP_GDN_MAX_SV 128
|
||||
#define HTP_GDN_CHUNK_SIZE 64
|
||||
#define HTP_GDN_MIN_TOKENS 8
|
||||
|
||||
#ifndef HMX_FP16_TILE_SIZE
|
||||
#define HMX_FP16_TILE_SIZE 2048
|
||||
@@ -132,7 +133,6 @@ struct htp_gdn_hmx_vtcm_layout {
|
||||
size_t off_rows_a;
|
||||
|
||||
size_t off_thread_scratch;
|
||||
size_t off_attn_rem;
|
||||
size_t off_scales_1;
|
||||
|
||||
size_t state_f32_bytes;
|
||||
@@ -192,8 +192,10 @@ static inline void htp_gdn_hmx_vtcm_layout_build(
|
||||
|
||||
VTCM_LAYOUT_ALLOC(off, off_s_state, bh * state_f32_sz);
|
||||
VTCM_LAYOUT_ALLOC(off, off_s_f16, bh * state_f16_sz);
|
||||
off = hex_align_up(off, HMX_FP16_TILE_SIZE);
|
||||
VTCM_LAYOUT_ALLOC(off, off_s_col_tiles, bh * state_tiles_sz);
|
||||
VTCM_LAYOUT_ALLOC(off, off_s_update_f32, bh * state_f32_sz);
|
||||
off = hex_align_up(off, HMX_FP16_TILE_SIZE);
|
||||
VTCM_LAYOUT_ALLOC(off, off_s_update_tiles, bh * state_tiles_sz);
|
||||
|
||||
VTCM_LAYOUT_ALLOC(off, off_q_f32[0], bh * dma_chunk_sz);
|
||||
@@ -222,6 +224,7 @@ static inline void htp_gdn_hmx_vtcm_layout_build(
|
||||
VTCM_LAYOUT_ALLOC(off, off_delta_f16, bh * act_f16_sz);
|
||||
VTCM_LAYOUT_ALLOC(off, off_d_f16, bh * act_f16_sz);
|
||||
|
||||
off = hex_align_up(off, HMX_FP16_TILE_SIZE);
|
||||
VTCM_LAYOUT_ALLOC(off, off_q_row_tiles, bh * tile_64xSv_sz);
|
||||
VTCM_LAYOUT_ALLOC(off, off_q_prime_row_tiles, bh * tile_64xSv_sz);
|
||||
VTCM_LAYOUT_ALLOC(off, off_k_row_tiles, bh * tile_64xSv_sz);
|
||||
@@ -250,9 +253,10 @@ static inline void htp_gdn_hmx_vtcm_layout_build(
|
||||
VTCM_LAYOUT_ALLOC(off, off_rows_a, bh * row_vecs_sz);
|
||||
|
||||
const size_t thread_scratch_sz = 64 * 128;
|
||||
off = hex_align_up(off, HMX_FP16_TILE_SIZE);
|
||||
VTCM_LAYOUT_ALLOC(off, off_thread_scratch, nth * thread_scratch_sz);
|
||||
VTCM_LAYOUT_ALLOC(off, off_attn_rem, nth * (128 * sizeof(float)));
|
||||
VTCM_LAYOUT_ALLOC(off, off_scales_1, 256);
|
||||
off = hex_align_up(off, HMX_FP16_TILE_SIZE);
|
||||
VTCM_LAYOUT_ALLOC(off, off_scales_1, HMX_FP16_TILE_SIZE);
|
||||
|
||||
L->total_bytes = off;
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -36,12 +36,13 @@ import signal
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, NamedTuple, Optional, Tuple
|
||||
from typing import Dict, List, NamedTuple, Optional, Set, Tuple
|
||||
|
||||
# Ignore SIGPIPE to handle pipes (e.g. head, grep) gracefully
|
||||
if hasattr(signal, "SIGPIPE"):
|
||||
signal.signal(signal.SIGPIPE, signal.SIG_DFL)
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
|
||||
logger = logging.getLogger("ggml-hexagon-inspect")
|
||||
|
||||
|
||||
@@ -56,6 +57,37 @@ class InsnInfo(NamedTuple):
|
||||
in_loop: bool
|
||||
|
||||
|
||||
class LoopStats:
|
||||
def __init__(self, loop_type: str, start_addr: int, end_addr: Optional[int] = None, loop_id: int = 0):
|
||||
self.loop_id = loop_id
|
||||
self.loop_type = loop_type # "loop0" or "loop1"
|
||||
self.start_addr = start_addr
|
||||
self.end_addr = end_addr
|
||||
self.packet_count = 0
|
||||
self.insn_count = 0
|
||||
self.vec_insn_count = 0
|
||||
self.vspills_st = 0
|
||||
self.vspills_ld = 0
|
||||
self.sspills_st = 0
|
||||
self.sspills_ld = 0
|
||||
|
||||
@property
|
||||
def vspills_total(self) -> int:
|
||||
return self.vspills_st + self.vspills_ld
|
||||
|
||||
@property
|
||||
def sspills_total(self) -> int:
|
||||
return self.sspills_st + self.sspills_ld
|
||||
|
||||
@property
|
||||
def has_v_roundtrip(self) -> bool:
|
||||
return self.vspills_st > 0 and self.vspills_ld > 0
|
||||
|
||||
@property
|
||||
def vec_density(self) -> float:
|
||||
return (self.vec_insn_count / self.packet_count) if self.packet_count > 0 else 0.0
|
||||
|
||||
|
||||
class FuncStats:
|
||||
def __init__(self, name: str, address: int, size: int):
|
||||
self.name = name
|
||||
@@ -66,14 +98,19 @@ class FuncStats:
|
||||
self.vec_insn_count = 0
|
||||
self.loop_count = 0
|
||||
self.vspills_in_loop = 0
|
||||
self.vspills_in_loop_st = 0
|
||||
self.vspills_in_loop_ld = 0
|
||||
self.vspills_total = 0
|
||||
self.sspills_in_loop = 0
|
||||
self.sspills_in_loop_st = 0
|
||||
self.sspills_in_loop_ld = 0
|
||||
self.sspills_total = 0
|
||||
self.promotions_in_loop = 0
|
||||
self.promotions_total = 0
|
||||
self.promotion_targets: Dict[str, int] = {}
|
||||
self.calls_in_loop = 0
|
||||
self.calls_total = 0
|
||||
self.loops: List[LoopStats] = []
|
||||
self.insns: List[InsnInfo] = []
|
||||
|
||||
|
||||
@@ -90,16 +127,43 @@ RE_INSN_LINE = re.compile(
|
||||
)
|
||||
RE_LOOP0_START = re.compile(r"\bloop0\((0x[0-9a-fA-F]+)")
|
||||
RE_LOOP1_START = re.compile(r"\bloop1\((0x[0-9a-fA-F]+)")
|
||||
RE_VSPILL = re.compile(r"\bvmemu?\s*\(\s*r(?:29|30)\b")
|
||||
RE_SSPILL = re.compile(r"\bmem[bwhd]\s*\(\s*r(?:29|30)\b")
|
||||
RE_VMEM_BASE = re.compile(r"\bvmemu?\s*\(\s*([a-z0-9]+)\b")
|
||||
RE_SMEM_BASE = re.compile(r"\bmem[bwhd](?:_locked|_fifo)?\s*\(\s*([a-z0-9]+)\b")
|
||||
RE_MEM_STORE = re.compile(r"\bv?mem[bwhdu]?(?:_[a-z]+)?\s*\([^)]*\)\s*(\+|-)?=")
|
||||
RE_ADD_OP = re.compile(r"\b(r[0-9]+)\s*=\s*add\s*\(\s*([^,()]+)\s*,\s*([^,()]+)\s*\)")
|
||||
RE_ASSIGN_LHS = re.compile(r"^\s*(?:if\s*\([^)]+\)\s*)?(r[0-9]+)(?::(r[0-9]+))?\s*(?:[+\-*/&|^]?=)")
|
||||
RE_VEC_OP = re.compile(r"\b(v[0-9]+|w[0-9]+|q[0-3]|vmemu?)\b")
|
||||
RE_STORE = re.compile(r"=\s*(?:v[0-9]|r[0-9]|w[0-9]|#)")
|
||||
RE_PROMOTION_CALL = re.compile(
|
||||
r"\b(?:call|jump)\s+(?:0x[0-9a-fA-F]+\s+)?<(__(?:trunc|extend)[a-zA-Z0-9_]+)(?:@plt)?>"
|
||||
)
|
||||
RE_ANY_CALL = re.compile(r"\bcallr?\b")
|
||||
|
||||
|
||||
def is_mem_store(insn: str) -> bool:
|
||||
return bool(RE_MEM_STORE.search(insn))
|
||||
|
||||
|
||||
def update_sp_regs(insn: str, sp_regs: Set[str]) -> None:
|
||||
# Track registers derived from stack frame (r29/r30)
|
||||
m_add = RE_ADD_OP.search(insn)
|
||||
if m_add:
|
||||
dest = m_add.group(1)
|
||||
op1 = m_add.group(2).strip()
|
||||
op2 = m_add.group(3).strip()
|
||||
if op1 in sp_regs or op2 in sp_regs:
|
||||
sp_regs.add(dest)
|
||||
return
|
||||
|
||||
m_assign = RE_ASSIGN_LHS.match(insn.strip())
|
||||
if m_assign:
|
||||
r1 = m_assign.group(1)
|
||||
r2 = m_assign.group(2)
|
||||
if r1 and r1 not in ("r29", "r30"):
|
||||
sp_regs.discard(r1)
|
||||
if r2 and r2 not in ("r29", "r30"):
|
||||
sp_regs.discard(r2)
|
||||
|
||||
|
||||
def get_repo_root() -> Path:
|
||||
# Resolve repository root from script location
|
||||
return Path(__file__).resolve().parent.parent.parent
|
||||
@@ -338,13 +402,15 @@ def parse_disassembly(
|
||||
end_idx = matches[i + 1].start() if i + 1 < len(matches) else len(disasm_text)
|
||||
chunk = disasm_text[start_idx:end_idx]
|
||||
|
||||
# Calculate rough byte size from line addresses
|
||||
stats = FuncStats(name=name, address=addr, size=0)
|
||||
|
||||
loop0_target: Optional[int] = None
|
||||
loop1_target: Optional[int] = None
|
||||
loop0_active = False
|
||||
loop1_active = False
|
||||
current_loop0: Optional[LoopStats] = None
|
||||
current_loop1: Optional[LoopStats] = None
|
||||
sp_regs: Set[str] = {"r29", "r30"}
|
||||
|
||||
first_addr = None
|
||||
last_addr = None
|
||||
@@ -364,6 +430,10 @@ def parse_disassembly(
|
||||
# Track packet count
|
||||
if "{" in asm_chunk:
|
||||
stats.packet_count += 1
|
||||
if current_loop0:
|
||||
current_loop0.packet_count += 1
|
||||
if current_loop1:
|
||||
current_loop1.packet_count += 1
|
||||
|
||||
# Check loop starts
|
||||
m0 = RE_LOOP0_START.search(asm_chunk)
|
||||
@@ -378,8 +448,23 @@ def parse_disassembly(
|
||||
|
||||
if loop0_target is not None and cur_addr >= loop0_target:
|
||||
loop0_active = True
|
||||
if current_loop0 is None:
|
||||
current_loop0 = LoopStats(
|
||||
loop_id=len(stats.loops) + 1,
|
||||
loop_type="loop0",
|
||||
start_addr=loop0_target,
|
||||
end_addr=0,
|
||||
)
|
||||
|
||||
if loop1_target is not None and cur_addr >= loop1_target:
|
||||
loop1_active = True
|
||||
if current_loop1 is None:
|
||||
current_loop1 = LoopStats(
|
||||
loop_id=len(stats.loops) + 1,
|
||||
loop_type="loop1",
|
||||
start_addr=loop1_target,
|
||||
end_addr=0,
|
||||
)
|
||||
|
||||
in_loop = loop0_active or loop1_active
|
||||
|
||||
@@ -388,31 +473,70 @@ def parse_disassembly(
|
||||
sub_insns = [p.strip() for p in cleaned.split(";") if p.strip()]
|
||||
|
||||
for insn in sub_insns:
|
||||
update_sp_regs(insn, sp_regs)
|
||||
|
||||
stats.insn_count += 1
|
||||
if current_loop0:
|
||||
current_loop0.insn_count += 1
|
||||
if current_loop1:
|
||||
current_loop1.insn_count += 1
|
||||
|
||||
is_vec = bool(RE_VEC_OP.search(insn))
|
||||
if is_vec:
|
||||
stats.vec_insn_count += 1
|
||||
if current_loop0:
|
||||
current_loop0.vec_insn_count += 1
|
||||
if current_loop1:
|
||||
current_loop1.vec_insn_count += 1
|
||||
|
||||
is_vspill = bool(RE_VSPILL.search(insn))
|
||||
is_sspill = bool(RE_SSPILL.search(insn))
|
||||
vm = RE_VMEM_BASE.search(insn)
|
||||
is_vspill = bool(vm and vm.group(1) in sp_regs)
|
||||
|
||||
sm = RE_SMEM_BASE.search(insn)
|
||||
is_sspill = bool(sm and sm.group(1) in sp_regs)
|
||||
|
||||
# Identify store vs load
|
||||
is_store = False
|
||||
is_load = False
|
||||
if is_vspill or is_sspill:
|
||||
if RE_STORE.search(insn):
|
||||
is_store = True
|
||||
else:
|
||||
is_load = True
|
||||
is_store = is_mem_store(insn)
|
||||
is_load = not is_store
|
||||
|
||||
if is_vspill:
|
||||
stats.vspills_total += 1
|
||||
if in_loop:
|
||||
stats.vspills_in_loop += 1
|
||||
if is_store:
|
||||
stats.vspills_in_loop_st += 1
|
||||
else:
|
||||
stats.vspills_in_loop_ld += 1
|
||||
if current_loop0:
|
||||
if is_store:
|
||||
current_loop0.vspills_st += 1
|
||||
else:
|
||||
current_loop0.vspills_ld += 1
|
||||
if current_loop1:
|
||||
if is_store:
|
||||
current_loop1.vspills_st += 1
|
||||
else:
|
||||
current_loop1.vspills_ld += 1
|
||||
elif is_sspill:
|
||||
stats.sspills_total += 1
|
||||
if in_loop:
|
||||
stats.sspills_in_loop += 1
|
||||
if is_store:
|
||||
stats.sspills_in_loop_st += 1
|
||||
else:
|
||||
stats.sspills_in_loop_ld += 1
|
||||
if current_loop0:
|
||||
if is_store:
|
||||
current_loop0.sspills_st += 1
|
||||
else:
|
||||
current_loop0.sspills_ld += 1
|
||||
if current_loop1:
|
||||
if is_store:
|
||||
current_loop1.sspills_st += 1
|
||||
else:
|
||||
current_loop1.sspills_ld += 1
|
||||
|
||||
is_call = bool(RE_ANY_CALL.search(insn))
|
||||
prom_m = RE_PROMOTION_CALL.search(insn)
|
||||
@@ -444,9 +568,29 @@ def parse_disassembly(
|
||||
if ":endloop0" in asm_chunk:
|
||||
loop0_active = False
|
||||
loop0_target = None
|
||||
if current_loop0:
|
||||
current_loop0.end_addr = cur_addr
|
||||
stats.loops.append(current_loop0)
|
||||
current_loop0 = None
|
||||
|
||||
if ":endloop1" in asm_chunk:
|
||||
loop1_active = False
|
||||
loop1_target = None
|
||||
if current_loop1:
|
||||
current_loop1.end_addr = cur_addr
|
||||
stats.loops.append(current_loop1)
|
||||
current_loop1 = None
|
||||
|
||||
if current_loop0:
|
||||
current_loop0.end_addr = last_addr or 0
|
||||
stats.loops.append(current_loop0)
|
||||
if current_loop1:
|
||||
current_loop1.end_addr = last_addr or 0
|
||||
stats.loops.append(current_loop1)
|
||||
|
||||
stats.loops.sort(key=lambda lp: lp.start_addr)
|
||||
for idx, loop in enumerate(stats.loops, 1):
|
||||
loop.loop_id = idx
|
||||
|
||||
if first_addr is not None and last_addr is not None:
|
||||
stats.size = (last_addr - first_addr) + 4
|
||||
@@ -463,14 +607,19 @@ def annotate_disasm_line(
|
||||
loop0_active: bool,
|
||||
loop1_active: bool,
|
||||
use_color: bool = True,
|
||||
) -> Tuple[str, Optional[int], Optional[int], bool, bool]:
|
||||
sp_regs: Optional[Set[str]] = None,
|
||||
) -> Tuple[str, Optional[int], Optional[int], bool, bool, bool]:
|
||||
# Annotate disassembly line with spill and loop tags
|
||||
lm = RE_INSN_LINE.match(raw_line)
|
||||
if not lm:
|
||||
return raw_line, loop0_target, loop1_target, loop0_active, loop1_active
|
||||
return raw_line, loop0_target, loop1_target, loop0_active, loop1_active, False
|
||||
|
||||
cur_addr = int(lm.group(1), 16)
|
||||
asm_chunk = lm.group(4)
|
||||
is_event = False
|
||||
|
||||
if sp_regs is None:
|
||||
sp_regs = {"r29", "r30"}
|
||||
|
||||
# Check loop starts
|
||||
m0 = RE_LOOP0_START.search(asm_chunk)
|
||||
@@ -490,39 +639,75 @@ def annotate_disasm_line(
|
||||
tags = []
|
||||
if m0:
|
||||
tags.append("[LOOP0-START]")
|
||||
is_event = True
|
||||
if m1:
|
||||
tags.append("[LOOP1-START]")
|
||||
is_event = True
|
||||
|
||||
if RE_VSPILL.search(asm_chunk):
|
||||
if in_loop:
|
||||
tags.append("[V-SPILL:IN-LOOP]" if not use_color else "\033[1;31m[V-SPILL:IN-LOOP]\033[0m")
|
||||
else:
|
||||
tags.append("[V-SPILL]" if not use_color else "\033[1;33m[V-SPILL]\033[0m")
|
||||
elif RE_SSPILL.search(asm_chunk):
|
||||
if in_loop:
|
||||
tags.append("[S-SPILL:IN-LOOP]" if not use_color else "\033[1;35m[S-SPILL:IN-LOOP]\033[0m")
|
||||
cleaned = re.sub(r"[{}\s]|:endloop[01]", " ", asm_chunk)
|
||||
sub_insns = [p.strip() for p in cleaned.split(";") if p.strip()]
|
||||
|
||||
for insn in sub_insns:
|
||||
update_sp_regs(insn, sp_regs)
|
||||
|
||||
for insn in sub_insns:
|
||||
vm = RE_VMEM_BASE.search(insn)
|
||||
if vm and vm.group(1) in sp_regs:
|
||||
base = vm.group(1)
|
||||
is_st = is_mem_store(insn)
|
||||
op = "STORE" if is_st else "LOAD"
|
||||
tgt = f"({base})" if base not in ("r29", "r30") else ""
|
||||
if in_loop:
|
||||
tag = f"[V-SPILL:{op}{tgt}:IN-LOOP]"
|
||||
tags.append(f"\033[1;31m{tag}\033[0m" if use_color else tag)
|
||||
else:
|
||||
tag = f"[V-SPILL:{op}{tgt}]"
|
||||
tags.append(f"\033[1;33m{tag}\033[0m" if use_color else tag)
|
||||
is_event = True
|
||||
|
||||
sm = RE_SMEM_BASE.search(insn)
|
||||
if sm and sm.group(1) in sp_regs:
|
||||
base = sm.group(1)
|
||||
is_st = is_mem_store(insn)
|
||||
op = "STORE" if is_st else "LOAD"
|
||||
tgt = f"({base})" if base not in ("r29", "r30") else ""
|
||||
if in_loop:
|
||||
tag = f"[S-SPILL:{op}{tgt}:IN-LOOP]"
|
||||
tags.append(f"\033[1;35m{tag}\033[0m" if use_color else tag)
|
||||
else:
|
||||
tag = f"[S-SPILL:{op}{tgt}]"
|
||||
tags.append(f"\033[0;35m{tag}\033[0m" if use_color else tag)
|
||||
is_event = True
|
||||
|
||||
prom_m = RE_PROMOTION_CALL.search(asm_chunk)
|
||||
if prom_m:
|
||||
ptarget = prom_m.group(1)
|
||||
if in_loop:
|
||||
tags.append(f"[PROMOTION:{ptarget}:IN-LOOP]" if not use_color else f"\033[1;31m[PROMOTION:{ptarget}:IN-LOOP]\033[0m")
|
||||
tag = f"[PROMOTION:{ptarget}:IN-LOOP]"
|
||||
tags.append(f"\033[1;31m{tag}\033[0m" if use_color else tag)
|
||||
else:
|
||||
tags.append(f"[PROMOTION:{ptarget}]" if not use_color else f"\033[1;35m[PROMOTION:{ptarget}]\033[0m")
|
||||
tag = f"[PROMOTION:{ptarget}]"
|
||||
tags.append(f"\033[1;35m{tag}\033[0m" if use_color else tag)
|
||||
is_event = True
|
||||
elif RE_ANY_CALL.search(asm_chunk):
|
||||
if in_loop:
|
||||
tags.append("[CALL:IN-LOOP]" if not use_color else "\033[1;31m[CALL:IN-LOOP]\033[0m")
|
||||
tag = "[CALL:IN-LOOP]"
|
||||
tags.append(f"\033[1;31m{tag}\033[0m" if use_color else tag)
|
||||
is_event = True
|
||||
else:
|
||||
tags.append("[CALL]" if not use_color else "\033[1;36m[CALL]\033[0m")
|
||||
tag = "[CALL]"
|
||||
tags.append(f"\033[1;36m{tag}\033[0m" if use_color else tag)
|
||||
|
||||
if ":endloop0" in asm_chunk:
|
||||
tags.append("[LOOP0-END]")
|
||||
loop0_active = False
|
||||
loop0_target = None
|
||||
is_event = True
|
||||
if ":endloop1" in asm_chunk:
|
||||
tags.append("[LOOP1-END]")
|
||||
loop1_active = False
|
||||
loop1_target = None
|
||||
is_event = True
|
||||
|
||||
tag_str = " ".join(tags)
|
||||
if tag_str:
|
||||
@@ -530,7 +715,7 @@ def annotate_disasm_line(
|
||||
else:
|
||||
annotated = raw_line
|
||||
|
||||
return annotated, loop0_target, loop1_target, loop0_active, loop1_active
|
||||
return annotated, loop0_target, loop1_target, loop0_active, loop1_active, is_event
|
||||
|
||||
|
||||
def run_spills(
|
||||
@@ -566,14 +751,15 @@ def run_spills(
|
||||
col_pkts = "Packets"
|
||||
col_insn = "Insns"
|
||||
col_vec = "HVX Ops"
|
||||
col_vloop = "V-Loop"
|
||||
col_vloop = "V-Loop (st/ld)"
|
||||
col_vtot = "V-Tot"
|
||||
col_sloop = "S-Loop"
|
||||
col_sloop = "S-Loop (st/ld)"
|
||||
col_stot = "S-Tot"
|
||||
col_notes = "Notes"
|
||||
|
||||
hdr = (
|
||||
f"{col_addr:<10} | {col_name:<44} | {col_pkts:>7} | {col_insn:>6} | "
|
||||
f"{col_vec:>7} | {col_vloop:>6} | {col_vtot:>5} | {col_sloop:>6} | {col_stot:>5}"
|
||||
f"{col_addr:<10} | {col_name:<40} | {col_pkts:>7} | {col_insn:>6} | "
|
||||
f"{col_vec:>7} | {col_vloop:>14} | {col_vtot:>5} | {col_sloop:>14} | {col_stot:>5} | {col_notes}"
|
||||
)
|
||||
sep = "-" * len(hdr)
|
||||
|
||||
@@ -596,9 +782,11 @@ def run_spills(
|
||||
|
||||
# Check strict criteria
|
||||
if args.strict:
|
||||
if f.vspills_in_loop > args.max_inloop_vspills:
|
||||
inloop_v = f.vspills_in_loop_st if getattr(args, "strict_stores_only", False) else f.vspills_in_loop
|
||||
if inloop_v > args.max_inloop_vspills:
|
||||
lbl = "in-loop vector store spills" if getattr(args, "strict_stores_only", False) else "in-loop vector spills"
|
||||
strict_violations.append(
|
||||
f"{f.name}: {f.vspills_in_loop} in-loop vector spills (max allowed: {args.max_inloop_vspills})"
|
||||
f"{f.name}: {inloop_v} {lbl} (max allowed: {args.max_inloop_vspills})"
|
||||
)
|
||||
if dma_re and dma_re.search(f.name):
|
||||
if f.vec_insn_count > args.max_dma_vec_ops:
|
||||
@@ -606,14 +794,27 @@ def run_spills(
|
||||
f"{f.name}: DMA worker contains {f.vec_insn_count} HVX vector ops (max allowed: {args.max_dma_vec_ops})"
|
||||
)
|
||||
|
||||
# Highlight in-loop vector spills
|
||||
vloop_str = f"{f.vspills_in_loop:>6}"
|
||||
vloop_detail = f"{f.vspills_in_loop} ({f.vspills_in_loop_st}s,{f.vspills_in_loop_ld}l)" if f.vspills_in_loop > 0 else "0"
|
||||
sloop_detail = f"{f.sspills_in_loop} ({f.sspills_in_loop_st}s,{f.sspills_in_loop_ld}l)" if f.sspills_in_loop > 0 else "0"
|
||||
|
||||
notes = ""
|
||||
if f.vspills_in_loop_st > 0 and f.vspills_in_loop_ld > 0:
|
||||
notes = "\033[1;31m[V-ROUNDTRIP!]\033[0m" if use_color else "[V-ROUNDTRIP!]"
|
||||
elif f.vspills_in_loop_st == 0 and f.vspills_in_loop_ld > 0:
|
||||
notes = "v-readonly"
|
||||
|
||||
vloop_str = f"{vloop_detail:>14}"
|
||||
if f.vspills_in_loop > 0 and use_color:
|
||||
vloop_str = f"\033[1;31m{vloop_str}\033[0m"
|
||||
if f.vspills_in_loop_st > 0 and f.vspills_in_loop_ld > 0:
|
||||
vloop_str = f"\033[1;31m{vloop_str}\033[0m"
|
||||
else:
|
||||
vloop_str = f"\033[1;33m{vloop_str}\033[0m"
|
||||
|
||||
sloop_str = f"{sloop_detail:>14}"
|
||||
|
||||
logger.info(
|
||||
f"0x{f.address:08x} | {f.name:<44} | {f.packet_count:>7} | {f.insn_count:>6} | "
|
||||
f"{f.vec_insn_count:>7} | {vloop_str} | {f.vspills_total:>5} | {f.sspills_in_loop:>6} | {f.sspills_total:>5}"
|
||||
f"0x{f.address:08x} | {f.name:<40} | {f.packet_count:>7} | {f.insn_count:>6} | "
|
||||
f"{f.vec_insn_count:>7} | {vloop_str} | {f.vspills_total:>5} | {sloop_str} | {f.sspills_total:>5} | {notes}"
|
||||
)
|
||||
|
||||
logger.info(sep)
|
||||
@@ -744,7 +945,7 @@ def run_disasm(
|
||||
args: argparse.Namespace,
|
||||
) -> int:
|
||||
# Disassemble matching function(s) with annotated loop and spill markers
|
||||
func_pattern = args.disasm
|
||||
func_pattern = args.disasm if args.disasm else (args.func or ".*")
|
||||
logger.info(f"Inspecting library: {lib_path}")
|
||||
logger.info(f"Disassembling functions matching: '{func_pattern}'\n")
|
||||
|
||||
@@ -794,9 +995,11 @@ def run_disasm(
|
||||
logger.info(f"Packets: {func_stats.packet_count} | Instructions: {func_stats.insn_count} | Loops: {func_stats.loop_count}")
|
||||
vec_pct = (func_stats.vec_insn_count / func_stats.insn_count * 100.0) if func_stats.insn_count else 0.0
|
||||
logger.info(f"HVX Ops: {func_stats.vec_insn_count} ({vec_pct:.1f}% of instructions)")
|
||||
vloop_info = f"{func_stats.vspills_in_loop} ({func_stats.vspills_in_loop_st} st, {func_stats.vspills_in_loop_ld} ld)"
|
||||
sloop_info = f"{func_stats.sspills_in_loop} ({func_stats.sspills_in_loop_st} st, {func_stats.sspills_in_loop_ld} ld)"
|
||||
logger.info(
|
||||
f"Spills: Vector in-loop: {func_stats.vspills_in_loop} | Vector total: {func_stats.vspills_total} | "
|
||||
f"Scalar in-loop: {func_stats.sspills_in_loop} | Scalar total: {func_stats.sspills_total}"
|
||||
f"Spills: Vector in-loop: {vloop_info} | Vector total: {func_stats.vspills_total} | "
|
||||
f"Scalar in-loop: {sloop_info} | Scalar total: {func_stats.sspills_total}"
|
||||
)
|
||||
logger.info(
|
||||
f"Calls: Total: {func_stats.calls_total} (in-loop: {func_stats.calls_in_loop}) | "
|
||||
@@ -804,18 +1007,76 @@ def run_disasm(
|
||||
)
|
||||
logger.info(hdr_border)
|
||||
|
||||
# Log annotated disassembly
|
||||
loop0_target: Optional[int] = None
|
||||
loop1_target: Optional[int] = None
|
||||
# Print Loop Breakdown Table if function has loops
|
||||
if func_stats.loops:
|
||||
logger.info(f"\n--- Loops ({len(func_stats.loops)}) " + "-" * 67)
|
||||
loop_hdr = (
|
||||
f"{'#':<3} | {'Type':<5} | {'Address Range':<25} | {'Packets':>7} | "
|
||||
f"{'HVX Ops':>7} | {'Vec/Pkt':>7} | {'V-Spills (st, ld)':>17} | {'S-Spills (st, ld)':>17} | Notes"
|
||||
)
|
||||
logger.info(loop_hdr)
|
||||
logger.info("-" * len(loop_hdr))
|
||||
for loop in func_stats.loops:
|
||||
vspill_str = f"{loop.vspills_total} ({loop.vspills_st}s,{loop.vspills_ld}l)"
|
||||
sspill_str = f"{loop.sspills_total} ({loop.sspills_st}s,{loop.sspills_ld}l)"
|
||||
notes = []
|
||||
if loop.has_v_roundtrip:
|
||||
notes.append("\033[1;31m[V-ROUNDTRIP!]\033[0m" if use_color else "[V-ROUNDTRIP!]")
|
||||
elif loop.vspills_st == 0 and loop.vspills_ld > 0:
|
||||
notes.append("v-readonly")
|
||||
if loop.vec_density >= 1.5:
|
||||
notes.append("\033[1;32mdual-hvx\033[0m" if use_color else "dual-hvx")
|
||||
notes_str = ", ".join(notes)
|
||||
logger.info(
|
||||
f"{loop.loop_id:<3} | {loop.loop_type:<5} | 0x{loop.start_addr:08x} - 0x{loop.end_addr:08x} | "
|
||||
f"{loop.packet_count:>7} | {loop.vec_insn_count:>7} | {loop.vec_density:>7.2f} | "
|
||||
f"{vspill_str:>17} | {sspill_str:>17} | {notes_str}"
|
||||
)
|
||||
logger.info("-" * len(loop_hdr) + "\n")
|
||||
|
||||
# Parse lines and annotations
|
||||
lines = chunk.splitlines()
|
||||
annotated_lines = []
|
||||
is_event_list = []
|
||||
loop0_target = None
|
||||
loop1_target = None
|
||||
loop0_active = False
|
||||
loop1_active = False
|
||||
sp_regs = {"r29", "r30"}
|
||||
|
||||
for line in chunk.splitlines():
|
||||
ann_line, loop0_target, loop1_target, loop0_active, loop1_active = annotate_disasm_line(
|
||||
line, loop0_target, loop1_target, loop0_active, loop1_active, use_color
|
||||
for line in lines:
|
||||
ann_line, loop0_target, loop1_target, loop0_active, loop1_active, is_ev = annotate_disasm_line(
|
||||
line, loop0_target, loop1_target, loop0_active, loop1_active, use_color, sp_regs
|
||||
)
|
||||
logger.info(ann_line)
|
||||
logger.info("")
|
||||
annotated_lines.append(ann_line)
|
||||
is_event_list.append(is_ev)
|
||||
|
||||
# Filter output if --spills-only
|
||||
if getattr(args, "spills_only", False):
|
||||
ctx = args.context if args.context is not None else 2
|
||||
to_show = [False] * len(annotated_lines)
|
||||
for idx, ev in enumerate(is_event_list):
|
||||
if ev:
|
||||
for j in range(max(0, idx - ctx), min(len(annotated_lines), idx + ctx + 1)):
|
||||
to_show[j] = True
|
||||
|
||||
if not any(to_show):
|
||||
logger.info(" (No spills, promotions, or in-loop calls detected in this function)\n")
|
||||
else:
|
||||
in_gap = False
|
||||
for idx, show in enumerate(to_show):
|
||||
if show:
|
||||
in_gap = False
|
||||
logger.info(annotated_lines[idx])
|
||||
else:
|
||||
if not in_gap:
|
||||
logger.info(" ...")
|
||||
in_gap = True
|
||||
logger.info("")
|
||||
else:
|
||||
for ann_line in annotated_lines:
|
||||
logger.info(ann_line)
|
||||
logger.info("")
|
||||
|
||||
return 0
|
||||
|
||||
@@ -964,9 +1225,24 @@ def main():
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disasm",
|
||||
nargs="?",
|
||||
const="",
|
||||
metavar="FUNC",
|
||||
help="Disassemble function symbol or regex pattern with annotated loop and spill markers.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--spills-only",
|
||||
action="store_true",
|
||||
help="In --disasm, only display packets containing spills, promotions, or in-loop calls, with surrounding context.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"-C",
|
||||
"--context",
|
||||
type=int,
|
||||
default=None,
|
||||
metavar="N",
|
||||
help="Number of context packets before and after spills in --disasm --spills-only (default: 2).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--limit",
|
||||
type=int,
|
||||
@@ -983,8 +1259,9 @@ def main():
|
||||
# Filtering & Display
|
||||
parser.add_argument(
|
||||
"--func",
|
||||
"--fn",
|
||||
"-f",
|
||||
help="Regex filter for function names in --spills or --promotions.",
|
||||
help="Regex filter for function names in --spills, --promotions, or --disasm.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--all",
|
||||
@@ -1010,6 +1287,11 @@ def main():
|
||||
default=0,
|
||||
help="Maximum allowed in-loop vector spills in --strict mode (default: 0).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--strict-stores-only",
|
||||
action="store_true",
|
||||
help="In --strict mode, only count vector store spills (st > 0) towards violations, ignoring readonly stack loads.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max-dma-vec-ops",
|
||||
type=int,
|
||||
@@ -1057,7 +1339,7 @@ def main():
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(message)s")
|
||||
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
|
||||
|
||||
repo_root = get_repo_root()
|
||||
|
||||
@@ -1092,7 +1374,7 @@ def main():
|
||||
# Dispatch commands
|
||||
if args.addr2line is not None:
|
||||
sys.exit(run_addr2line(toolchain, lib_path, args))
|
||||
elif args.disasm:
|
||||
elif args.disasm is not None:
|
||||
sys.exit(run_disasm(toolchain, lib_path, args))
|
||||
elif args.promotions:
|
||||
sys.exit(run_promotions(toolchain, lib_path, args))
|
||||
@@ -1102,5 +1384,5 @@ def main():
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
logging.basicConfig(level=logging.INFO, format="%(message)s")
|
||||
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
|
||||
main()
|
||||
|
||||
@@ -54,6 +54,7 @@ def device_matches(record_device, target_device):
|
||||
return False
|
||||
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
|
||||
logger = logging.getLogger("ggml-hexagon-profile")
|
||||
|
||||
|
||||
@@ -648,7 +649,7 @@ def main():
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format='%(message)s')
|
||||
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
|
||||
|
||||
if "pmu" in args.sort and args.pmu_index is None:
|
||||
logger.error(f"Cannot sort by '{args.sort}' without --pmu-index.")
|
||||
|
||||
@@ -10,6 +10,7 @@ import bisect
|
||||
from typing import Any, Dict, List, Optional
|
||||
from collections import defaultdict
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
|
||||
logger = logging.getLogger("ggml-hexagon-trace")
|
||||
|
||||
op_pattern = re.compile(
|
||||
@@ -732,7 +733,7 @@ def main():
|
||||
group.add_argument("--tail", type=int, help="Limit to last N ops")
|
||||
|
||||
args = parser.parse_args()
|
||||
logging.basicConfig(level=logging.INFO, format='%(message)s')
|
||||
logging.basicConfig(level=logging.INFO, format="%(message)s", stream=sys.stdout)
|
||||
|
||||
op_filter_re = None
|
||||
if args.filter:
|
||||
|
||||
@@ -31,6 +31,7 @@ MANAGED_ENV_NAMES = (
|
||||
"GGML_HEXAGON_MBUF",
|
||||
"GGML_HEXAGON_MM_SELECT",
|
||||
"GGML_HEXAGON_FA_SELECT",
|
||||
"GGML_HEXAGON_GDN_SELECT",
|
||||
"GGML_HEXAGON_AR_SELECT",
|
||||
"GGML_HEXAGON_ETM",
|
||||
"GGML_HEXAGON_ARCH",
|
||||
@@ -166,6 +167,7 @@ def main():
|
||||
parser.add_argument("--hex-mbuf", help="Maximum host buffer size limit in MB to allocate (GGML_HEXAGON_MBUF)")
|
||||
parser.add_argument("--hex-mm-select", help="Select MUL_MAT and MUL_MAT_ID kernel (GGML_HEXAGON_MM_SELECT) 2:HMX,1:HVX,0:disable")
|
||||
parser.add_argument("--hex-fa-select", help="Select Flash Attention kernel (GGML_HEXAGON_FA_SELECT) 2:HMX,1:HVX,0:disable")
|
||||
parser.add_argument("--hex-gdn-select", help="Select Gated Delta Net kernel (GGML_HEXAGON_GDN_SELECT) 2:HMX,1:HVX,0:disable")
|
||||
parser.add_argument("--hex-ar-select", help="Select All-Reduce kernel (GGML_HEXAGON_AR_SELECT) 1:enable,0:disable")
|
||||
parser.add_argument("--hex-etm", help="Enable Embedded Trace Macrocell hardware tracing / trace logging (GGML_HEXAGON_ETM)")
|
||||
parser.add_argument("--hex-arch", help="Target Hexagon NPU architecture version override (v73, v75, v79, v81, etc.) (GGML_HEXAGON_ARCH)")
|
||||
@@ -306,6 +308,7 @@ def main():
|
||||
set_env("GGML_HEXAGON_MBUF", args.hex_mbuf)
|
||||
set_env("GGML_HEXAGON_MM_SELECT", args.hex_mm_select)
|
||||
set_env("GGML_HEXAGON_FA_SELECT", args.hex_fa_select)
|
||||
set_env("GGML_HEXAGON_GDN_SELECT", args.hex_gdn_select)
|
||||
set_env("GGML_HEXAGON_AR_SELECT", args.hex_ar_select)
|
||||
set_env("GGML_HEXAGON_ETM", args.hex_etm)
|
||||
set_env("GGML_HEXAGON_ARCH", args.hex_arch)
|
||||
|
||||
Reference in New Issue
Block a user