hexagon: support for backend sampler (#29502)

* hex-topk: trying to improve/cleanup the pipeline

* hex-sampling: add STEP op

* hex-sampler: add SUM op

* hex-sampler: update CPY to support sampling cases

* hex-binary: add support for chunking to handle large logits

* hex-argmax: super basic version of ARGMAX

* hex-binary: support for scalars in extended buffers

* hex-binary: fix wrong indexing for dim 1 broadcasts across dim 2 slices

* hex-argsort: fix missing header

* hex-sampler: cleanup dma usage in the sampler related ops, and binary

* hex-build: disable autovectorizer, it is better to use explicit hints for critical loops

* hex-binary: fix perf regression due to is_1d fallback

* hex-ops: update supported ops
This commit is contained in:
Max Krasnyansky
2026-09-26 20:29:36 -07:00
committed by GitHub
parent 95887577ab
commit 2b129ccfa0
21 changed files with 12611 additions and 10816 deletions
+20 -16
View File
@@ -14,16 +14,16 @@ Legend:
| Operation | BLAS | CANN | CPU | CUDA | ET | HTP | MTL | OpenCL | SYCL | Vulkan | WebGPU | ZenDNN | zDNN |
|-----------|------|------|------|------|------|------|------|------|------|------|------|------|------|
| ABS | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ABS | ❌ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ACC | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ADD | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ADD1 | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| ADD_ID | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ARANGE | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ARGMAX | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ARGSORT | ❌ | ✅ | ✅ | ✅ | ❌ | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| ARGMAX | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ARGSORT | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| CEIL | ❌ | ❌ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| CLAMP | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| CLAMP | ❌ | ✅ | ✅ | ✅ | 🟡 | ✅ | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| COL2IM_1D | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| CONCAT | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | 🟡 | ❌ | ❌ |
| CONT | ❌ | 🟡 | ✅ | ✅ | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | 🟡 | ❌ | ❌ |
@@ -55,7 +55,7 @@ Legend:
| GATED_LINEAR_ATTN | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| GEGLU | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| GEGLU_ERF | ❌ | ✅ | ✅ | ✅ | 🟡 | ❌ | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| GEGLU_QUICK | ❌ | ✅ | ✅ | ✅ | 🟡 | ❌ | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| GEGLU_QUICK | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| GELU | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| GELU_ERF | ❌ | ✅ | ✅ | 🟡 | 🟡 | ❌ | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| GELU_QUICK | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
@@ -64,17 +64,21 @@ Legend:
| GROUP_NORM | ❌ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ |
| HARDSIGMOID | ❌ | ✅ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| HARDSWISH | ❌ | ✅ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| IM2COL | ❌ | ✅ | ✅ | ✅ | ✅ | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| IM2COL | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| IM2COL_3D | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| L2_NORM | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ❌ | ✅ | ✅ | 🟡 | ❌ | ❌ |
| LEAKY_RELU | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | 🟡 | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| LEAKY_RELU | ❌ | ✅ | ✅ | ✅ | ❌ | 🟡 | 🟡 | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| LIGHTNING_INDEXER | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ |
| LOG | ❌ | ✅ | ✅ | ✅ | ❌ | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| LOG | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| MEAN | ❌ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | 🟡 | ✅ | ❌ | ❌ | ❌ |
| MUL | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| MUL_MAT | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 |
| MUL_MAT_HADAMARD | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| MUL_MAT_HADAMARD | ❌ | ❌ | ❌ | ❌ | ✅ | 🟡 | ❌ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| MUL_MAT_ID | ❌ | 🟡 | ✅ | ✅ | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | 🟡 | ❌ |
| MUL_MAT_ID_W4A4 | ❌ | ❌ | ❌ | ❌ | ❌ | 🟡 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| MUL_MAT_ID_W4A8 | ❌ | ❌ | ❌ | ❌ | ❌ | 🟡 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| MUL_MAT_W4A4 | ❌ | ❌ | ❌ | ❌ | ❌ | 🟡 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| MUL_MAT_W4A8 | ❌ | ❌ | ❌ | ❌ | ❌ | 🟡 | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| NEG | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| NORM | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | 🟡 | ❌ | ❌ |
| OPT_STEP_ADAMW | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
@@ -85,12 +89,12 @@ Legend:
| POOL_1D | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| POOL_2D | ❌ | 🟡 | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| REGLU | ❌ | ✅ | ✅ | ✅ | 🟡 | ❌ | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| RELU | ❌ | ✅ | ✅ | 🟡 | 🟡 | ❌ | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| RELU | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| REPEAT | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | 🟡 | ❌ | ❌ |
| REPEAT_BACK | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| RMS_NORM | ❌ | ✅ | ✅ | ✅ | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| RMS_NORM_BACK | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ROLL | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ROLL | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ROPE | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ROPE_BACK | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| ROUND | ❌ | ❌ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
@@ -108,20 +112,20 @@ Legend:
| SOFT_MAX | ❌ | 🟡 | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SOFT_MAX_BACK | ❌ | ❌ | 🟡 | 🟡 | ❌ | ❌ | ❌ | ❌ | 🟡 | ✅ | ❌ | ❌ | ❌ |
| SOLVE_TRI | ❌ | ❌ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ❌ | 🟡 | ✅ | ✅ | ❌ | ❌ |
| SQR | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SQRT | ❌ | ✅ | ✅ | ✅ | ❌ | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SQR | ❌ | ✅ | ✅ | ✅ | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SQRT | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SSM_CONV | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SSM_SCAN | ❌ | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ❌ | 🟡 | 🟡 | ✅ | ❌ | ❌ |
| STEP | ❌ | ✅ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| STEP | ❌ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SUB | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SUM | ❌ | 🟡 | ✅ | 🟡 | ❌ | ❌ | 🟡 | ❌ | 🟡 | 🟡 | 🟡 | ❌ | ❌ |
| SUM | ❌ | 🟡 | ✅ | 🟡 | ❌ | 🟡 | 🟡 | ❌ | 🟡 | 🟡 | 🟡 | ❌ | ❌ |
| SUM_ROWS | ❌ | ✅ | ✅ | 🟡 | ❌ | 🟡 | ✅ | 🟡 | 🟡 | ✅ | ✅ | ❌ | ❌ |
| SWIGLU | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| SWIGLU_CLAMP | ❌ | ❌ | ❌ | ❌ | ❌ | 🟡 | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ |
| SWIGLU_OAI | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| TANH | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| TIMESTEP_EMBEDDING | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ |
| TOP_K | ❌ | ❌ | ✅ | ❌ | ❌ | 🟡 | ✅ | ❌ | ✅ | 🟡 | ✅ | ❌ | ❌ |
| TOP_K | ❌ | ❌ | ✅ | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ | 🟡 | ✅ | ❌ | ❌ |
| TRI | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| TRUNC | ❌ | ❌ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| UPSCALE | ❌ | 🟡 | ✅ | ✅ | ❌ | ❌ | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
+10632 -10041
View File
File diff suppressed because it is too large Load Diff
+205 -28
View File
@@ -63,6 +63,7 @@
#include "htp/rope-ops.h"
#include "htp/ssm-conv.h"
#include "htp/gated-delta-net-ops.h"
#include "htp/argsort-ops.h"
#include "htp_iface.h"
#include "htp-drv.h"
@@ -368,6 +369,13 @@ static void ggml_hexagon_precompute_gated_delta_net_params(
struct htp_gdn_kernel_params * kparams
);
static void ggml_hexagon_precompute_sort_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * op,
bool is_top_k,
struct htp_sort_kernel_params * kparams
);
static void ggml_hexagon_precompute_fused_mmnx_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0,
@@ -4998,13 +5006,48 @@ static bool ggml_hexagon_precompute_binary_params(
const bool is_add_id = op == HTP_OP_ADD_ID;
const bool is_scalar = !is_add_id && src1->ne[0] == 1;
const bool is_transposed = src0->nb[1] < src0_row_size || src1->nb[1] < src1_row_size || dst->nb[1] < dst_row_size;
const bool is_row_bcast = !is_add_id && !is_scalar && !is_transposed &&
src1->ne[0] == src0->ne[0] &&
(src0->ne[1] > 1 || src0->ne[2] > 1 || src0->ne[3] > 1) &&
src1->ne[1] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1;
const bool is_same_shape = !is_add_id && !is_scalar && !is_transposed &&
src1->ne[0] == src0->ne[0] &&
(src1->ne[1] == src0->ne[1] || src1->ne[1] == 1) &&
src1->ne[1] == src0->ne[1] &&
(src1->ne[2] == src0->ne[2] || src1->ne[2] == 1) &&
(src1->ne[3] == src0->ne[3] || src1->ne[3] == 1);
const bool is_row_bcast = is_same_shape && src1->ne[1] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1;
const bool is_complex = !is_add_id && !is_scalar && !is_same_shape && (src1->ne[0] == src0->ne[0]);
const bool is_complex = !is_add_id && !is_scalar && !is_same_shape && !is_row_bcast && (src1->ne[0] == src0->ne[0]);
const bool is_contig = ggml_is_contiguous(src0) && ggml_is_contiguous(src1) && ggml_is_contiguous(dst);
const bool is_scalar_broadcast = !is_add_id && (ggml_nelements(src1) == 1);
if (!is_add_id && is_contig && (ggml_are_same_shape(src0, src1) || is_scalar_broadcast)) {
const uint32_t total_elems = (uint32_t) ggml_nelements(src0);
const uint32_t n_threads = sess->n_threads;
const uint32_t max_chunk_elems = 32768 / elem_size;
const uint32_t min_chunk_elems = 256;
const uint32_t target_chunk_elems = hex_round_up((total_elems + (2 * n_threads) - 1) / (2 * n_threads), 32);
const uint32_t chunk_size = (std::min)(max_chunk_elems, (std::max)(target_chunk_elems, min_chunk_elems));
const uint32_t chunk_bytes = hex_round_up(chunk_size * elem_size, 128);
kparams->kernel_type = HTP_BINARY_KERNEL_CHUNKED;
kparams->n_threads = n_threads;
kparams->rows_per_buffer = 1;
kparams->src0_row_size_aligned = src0_row_size_aligned;
kparams->src1_row_size_aligned = is_scalar_broadcast ? 0 : src1_row_size_aligned;
kparams->dst_row_size_aligned = dst_row_size_aligned;
kparams->src1_size = 0;
kparams->chunk_size = chunk_size;
kparams->chunk_bytes = chunk_bytes;
kparams->is_scalar = is_scalar_broadcast ? 1 : 0;
struct htp_binary_vtcm_layout L;
htp_binary_vtcm_layout_build(&L, kparams, sess->vtcm_size);
if (L.total_bytes == 0 || L.total_bytes > sess->vtcm_size) {
return false;
}
kparams->vtcm_size = L.total_bytes;
return true;
}
enum htp_binary_kernel_type kernel_type;
size_t src1_size = 0;
@@ -5042,6 +5085,32 @@ static bool ggml_hexagon_precompute_binary_params(
struct htp_binary_vtcm_layout L;
htp_binary_vtcm_layout_build(&L, kparams, sess->vtcm_size);
if (L.rows_per_buffer == 0 || L.total_bytes > sess->vtcm_size) {
if (!is_add_id && is_contig && (ggml_are_same_shape(src0, src1) || is_scalar_broadcast)) {
const uint32_t total_elems = (uint32_t) ggml_nelements(src0);
const uint32_t n_threads = sess->n_threads;
const uint32_t max_chunk_elems = 32768 / elem_size;
const uint32_t min_chunk_elems = 256;
const uint32_t target_chunk_elems = hex_round_up((total_elems + (2 * n_threads) - 1) / (2 * n_threads), 32);
const uint32_t chunk_size = (std::min)(max_chunk_elems, (std::max)(target_chunk_elems, min_chunk_elems));
const uint32_t chunk_bytes = hex_round_up(chunk_size * elem_size, 128);
kparams->kernel_type = HTP_BINARY_KERNEL_CHUNKED;
kparams->n_threads = n_threads;
kparams->rows_per_buffer = 1;
kparams->src1_row_size_aligned = is_scalar_broadcast ? 0 : src1_row_size_aligned;
kparams->src1_size = 0;
kparams->chunk_size = chunk_size;
kparams->chunk_bytes = chunk_bytes;
kparams->is_scalar = is_scalar_broadcast ? 1 : 0;
htp_binary_vtcm_layout_build(&L, kparams, sess->vtcm_size);
if (L.total_bytes == 0 || L.total_bytes > sess->vtcm_size) {
return false;
}
kparams->vtcm_size = L.total_bytes;
return true;
}
return false;
}
@@ -5479,6 +5548,52 @@ static void ggml_hexagon_precompute_gated_delta_net_params(
if (kparams->n_threads > 0) kparams->div_n_threads = init_fastdiv_values(kparams->n_threads);
}
static void ggml_hexagon_precompute_sort_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * op,
bool is_top_k,
struct htp_sort_kernel_params * kparams
) {
memset(kparams, 0, sizeof(*kparams));
const struct ggml_tensor * src0 = op->src[0];
const struct ggml_tensor * dst = op;
const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const uint32_t ne00 = src0->ne[0];
const uint32_t k = dst->ne[0];
int32_t order = GGML_SORT_ORDER_DESC;
if (!is_top_k) {
order = ((const int32_t *) op->op_params)[0];
}
const uint32_t n_threads_max = sess->n_threads > 0 ? sess->n_threads : 4;
const size_t vtcm_budget = sess->vtcm_size > 0 ? sess->vtcm_size : (8 * 1024 * 1024);
struct htp_sort_vtcm_layout layout;
bool ok = htp_sort_solve_layout(&layout, ne00, total_rows, k, n_threads_max, vtcm_budget, is_top_k);
GGML_ASSERT(ok);
kparams->n_threads = (int32_t) layout.n_threads;
kparams->total_rows = (int32_t) total_rows;
kparams->row_start = 0;
kparams->row_end = (int32_t) total_rows;
kparams->ne00 = (int32_t) ne00;
kparams->k = (int32_t) k;
kparams->order = order;
kparams->is_top_k = is_top_k ? 1 : 0;
kparams->use_dma = 1;
kparams->chunk_elems = (int32_t) layout.chunk_elems;
kparams->n_chunks = (int32_t) layout.n_chunks;
kparams->vtcm_size = (int32_t) layout.total_bytes;
kparams->phase1_slot_size = (int32_t) layout.phase1_slot_size;
kparams->merge_values_off = (int32_t) layout.merge_values_off;
kparams->merge_indices_off = (int32_t) layout.merge_indices_off;
kparams->merge_elems = (int32_t) layout.merge_elems;
kparams->n_slots = (int32_t) layout.n_slots;
}
static void ggml_hexagon_precompute_fused_mmnx_params(
const struct ggml_hexagon_session * sess,
const struct ggml_tensor * src0, // W0
@@ -5830,7 +5945,8 @@ static bool ggml_hexagon_supported_unary(const struct ggml_hexagon_session * ses
case GGML_OP_LOG:
break;
case GGML_OP_UNARY:
if (ggml_get_unary_op(op) != GGML_UNARY_OP_ABS) {
if (ggml_get_unary_op(op) != GGML_UNARY_OP_ABS &&
ggml_get_unary_op(op) != GGML_UNARY_OP_STEP) {
return false;
}
break;
@@ -5853,6 +5969,26 @@ static bool ggml_hexagon_supported_unary(const struct ggml_hexagon_session * ses
GGML_UNUSED(sess);
}
static bool ggml_hexagon_supported_sum(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
const struct ggml_tensor * src0 = op->src[0];
const struct ggml_tensor * dst = op;
if (src0->type != GGML_TYPE_F32) {
return false;
}
if (dst->type != GGML_TYPE_F32) {
return false;
}
if (!ggml_is_contiguous(src0) || !ggml_is_contiguous(dst)) {
return false;
}
return true;
GGML_UNUSED(sess);
}
static bool ggml_hexagon_supported_sum_rows(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
const struct ggml_tensor * src0 = op->src[0];
const struct ggml_tensor * dst = op;
@@ -5874,6 +6010,26 @@ static bool ggml_hexagon_supported_sum_rows(const struct ggml_hexagon_session *
GGML_UNUSED(sess);
}
static bool ggml_hexagon_supported_argmax(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
const struct ggml_tensor * src0 = op->src[0];
const struct ggml_tensor * dst = op;
if (src0->type != GGML_TYPE_F32) {
return false;
}
if (dst->type != GGML_TYPE_I32) {
return false;
}
if (!ggml_is_contiguous(src0) || !ggml_is_contiguous(dst)) {
return false;
}
return true;
GGML_UNUSED(sess);
}
static bool ggml_hexagon_supported_activations(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
const struct ggml_tensor * src0 = op->src[0];
const struct ggml_tensor * src1 = op->src[1];
@@ -6058,42 +6214,36 @@ static bool ggml_hexagon_supported_argsort(const struct ggml_hexagon_session * s
const struct ggml_tensor * src0 = op->src[0]; // values
const struct ggml_tensor * dst = op; // indices
if (src0->type != GGML_TYPE_F32) {
if (src0->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_I32) {
return false;
}
if (dst->type != GGML_TYPE_I32) {
return false;
}
const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const uint32_t n_threads_max = sess->n_threads > 0 ? sess->n_threads : 4;
const size_t vtcm_budget = sess->vtcm_size > 0 ? sess->vtcm_size : (8 * 1024 * 1024);
if (src0->ne[0] > (16*1024)) {
// reject tensors with huge rows for now
struct htp_sort_vtcm_layout layout;
if (!htp_sort_solve_layout(&layout, src0->ne[0], total_rows, dst->ne[0], n_threads_max, vtcm_budget, false)) {
return false;
}
return true;
GGML_UNUSED(sess);
}
static bool ggml_hexagon_supported_top_k(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
const struct ggml_tensor * src0 = op->src[0]; // values
const struct ggml_tensor * dst = op; // indices
if (src0->type != GGML_TYPE_F32) {
if (src0->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_I32) {
return false;
}
if (dst->type != GGML_TYPE_I32) {
return false;
}
const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
const uint32_t n_threads_max = sess->n_threads > 0 ? sess->n_threads : 4;
const size_t vtcm_budget = sess->vtcm_size > 0 ? sess->vtcm_size : (8 * 1024 * 1024);
// Single row uses the threaded chunk+merge path. Multi-row uses one full
// buffer per thread, so it keeps the tighter 64K cap.
const bool single_row = (src0->ne[1] == 1 && src0->ne[2] == 1 && src0->ne[3] == 1);
const int64_t max_ne00 = single_row ? (256*1024) : (64*1024);
if (src0->ne[0] > max_ne00) {
struct htp_sort_vtcm_layout layout;
if (!htp_sort_solve_layout(&layout, src0->ne[0], total_rows, dst->ne[0], n_threads_max, vtcm_budget, true)) {
return false;
}
@@ -6401,9 +6551,11 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) {
case GGML_OP_CONT: return HTP_OP_CPY;
case GGML_OP_GET_ROWS: return HTP_OP_GET_ROWS;
case GGML_OP_SET_ROWS: return HTP_OP_SET_ROWS;
case GGML_OP_SUM: return HTP_OP_SUM;
case GGML_OP_SUM_ROWS: return HTP_OP_SUM_ROWS;
case GGML_OP_ARGSORT: return HTP_OP_ARGSORT;
case GGML_OP_TOP_K: return HTP_OP_TOP_K;
case GGML_OP_ARGMAX: return HTP_OP_ARGMAX;
case GGML_OP_NORM: return HTP_OP_NORM;
case GGML_OP_L2_NORM: return HTP_OP_L2_NORM;
case GGML_OP_RMS_NORM: return HTP_OP_RMS_NORM;
@@ -6440,6 +6592,7 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) {
case GGML_UNARY_OP_TANH: return HTP_OP_UNARY_TANH;
case GGML_UNARY_OP_ABS: return HTP_OP_UNARY_ABS;
case GGML_UNARY_OP_RELU: return HTP_OP_UNARY_RELU;
case GGML_UNARY_OP_STEP: return HTP_OP_UNARY_STEP;
default:
break;
}
@@ -6670,6 +6823,12 @@ static ggml_status ggml_backend_hexagon_graph_compute(ggml_backend_t backend, gg
node.node,
(struct htp_gdn_kernel_params *)node.kernel_params
);
} else if (node.opcode == HTP_OP_ARGSORT || node.opcode == HTP_OP_TOP_K) {
ggml_hexagon_precompute_sort_params(sess,
node.node,
node.opcode == HTP_OP_TOP_K,
(struct htp_sort_kernel_params *) node.kernel_params
);
}
computed_nodes.push_back(std::move(node));
}
@@ -7308,16 +7467,25 @@ static bool ggml_hexagon_supported_cpy(const struct ggml_hexagon_session * sess,
if (dst->type != GGML_TYPE_F32 && dst->type != GGML_TYPE_F16 &&
dst->type != GGML_TYPE_I32) return false;
const bool is_scalar = (ggml_nelements(src0) == 1 && ggml_nelements(dst) == 1);
const bool sametype = (src0->type == dst->type);
const bool transposed = ggml_is_transposed(src0) || ggml_is_transposed(dst);
const bool sameshape = !transposed && ggml_are_same_shape(src0, dst);
const bool transposed = !is_scalar && (ggml_is_transposed(src0) || ggml_is_transposed(dst));
const bool sameshape = is_scalar || (!transposed && ggml_are_same_shape(src0, dst));
// Same-type copies also support I32.
if (src0->type == GGML_TYPE_I32 || dst->type == GGML_TYPE_I32) {
if (!sameshape) return false;
if (sametype) return true;
if ((src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_I32) ||
(src0->type == GGML_TYPE_I32 && dst->type == GGML_TYPE_F32)) {
return true;
}
return false;
}
// can handle any shape and any same-type (pretty slow if reshaping is required)
if (sametype) return true;
// Type conversion is only supported between F32 and F16.
if (src0->type == GGML_TYPE_I32 || dst->type == GGML_TYPE_I32) return false;
// cannot handle re-shaping and type conversion at the same time
if (!sameshape) return false;
return true;
@@ -7466,10 +7634,18 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons
supp = ggml_hexagon_supported_unary(sess, op);
break;
case GGML_OP_SUM:
supp = ggml_hexagon_supported_sum(sess, op);
break;
case GGML_OP_SUM_ROWS:
supp = ggml_hexagon_supported_sum_rows(sess, op);
break;
case GGML_OP_ARGMAX:
supp = ggml_hexagon_supported_argmax(sess, op);
break;
case GGML_OP_SOFT_MAX:
supp = ggml_hexagon_supported_softmax(sess, op);
break;
@@ -7486,6 +7662,7 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons
case GGML_UNARY_OP_GELU:
case GGML_UNARY_OP_GELU_QUICK:
case GGML_UNARY_OP_RELU:
case GGML_UNARY_OP_STEP:
supp = ggml_hexagon_supported_unary(sess, op);
break;
default:
+7
View File
@@ -19,6 +19,7 @@
#include "htp/ssm-conv.h"
#include "htp/gated-delta-net-ops.h"
#include "htp/softmax-ops.h"
#include "htp/argsort-ops.h"
struct htp_opnode {
ggml_tensor * node { nullptr };
@@ -367,6 +368,12 @@ struct htp_opformat {
node.opcode == HTP_OP_SUB || node.opcode == HTP_OP_DIV) {
const auto * kparams = (const struct htp_binary_kernel_params *) node.kernel_params;
snprintf(str, max_size, "vtcm %u", (unsigned int) kparams->vtcm_size);
} else if (node.opcode == HTP_OP_ARGSORT || node.opcode == HTP_OP_TOP_K) {
const auto * kparams = (const struct htp_sort_kernel_params *) node.kernel_params;
snprintf(str, max_size, "%s nth %d nchk %d chk %d vtcm %d",
node.opcode == HTP_OP_TOP_K ? "top_k" : "argsort",
(int) kparams->n_threads, (int) kparams->n_chunks,
(int) kparams->chunk_elems, (int) kparams->vtcm_size);
} else {
snprintf(str, max_size, "----");
}
File diff suppressed because it is too large Load Diff
+178
View File
@@ -0,0 +1,178 @@
#ifndef HTP_ARGSORT_OPS_H
#define HTP_ARGSORT_OPS_H
#include <stdint.h>
#include <stddef.h>
#include <stdbool.h>
#include <string.h>
#include "hex-fastdiv.h"
struct htp_sort_kernel_params {
int32_t n_threads;
int32_t total_rows;
int32_t row_start;
int32_t row_end;
int32_t ne00;
int32_t k;
int32_t order; // GGML_SORT_ORDER_ASC (0) or GGML_SORT_ORDER_DESC (1)
int32_t is_top_k; // 1 if TOP_K, 0 if ARGSORT
int32_t use_dma; // 1 if DMA enabled
int32_t chunk_elems;
int32_t n_chunks;
int32_t vtcm_size;
int32_t phase1_slot_size;
int32_t merge_values_off;
int32_t merge_indices_off;
int32_t merge_elems;
int32_t n_slots;
int32_t pad[15];
};
struct htp_sort_vtcm_layout {
size_t total_bytes;
size_t phase1_slot_size;
size_t merge_values_off;
size_t merge_indices_off;
uint32_t chunk_elems;
uint32_t n_chunks;
uint32_t merge_elems;
uint32_t n_threads;
uint32_t n_slots;
};
static inline bool htp_sort_solve_layout(
struct htp_sort_vtcm_layout * layout,
uint32_t ne00,
uint32_t total_rows,
uint32_t k,
uint32_t n_threads_max,
size_t vtcm_budget,
bool is_top_k) {
memset(layout, 0, sizeof(*layout));
if (total_rows > 1) {
uint32_t n_threads = total_rows < n_threads_max ? total_rows : n_threads_max;
uint32_t n_vec = (ne00 + 31) / 32;
uint32_t n_vec_pow2 = 1;
while (n_vec_pow2 < n_vec) n_vec_pow2 <<= 1;
uint32_t ne00_padded = n_vec_pow2 * 32;
size_t values_size = ((ne00_padded * sizeof(float)) + 127) & ~127;
size_t indices_size = ((ne00_padded * sizeof(int32_t)) + 127) & ~127;
size_t spad_per_slot = ((values_size + indices_size) + 255) & ~255;
uint32_t n_slots = 2;
if (spad_per_slot * 2 > vtcm_budget) {
n_slots = 1;
}
size_t spad_per_thread = spad_per_slot * n_slots;
while (n_threads > 1 && (spad_per_thread * n_threads) > vtcm_budget) {
n_threads--;
}
size_t total_bytes = spad_per_thread * n_threads;
if (total_bytes > vtcm_budget) {
return false;
}
layout->total_bytes = total_bytes;
layout->phase1_slot_size = spad_per_slot;
layout->chunk_elems = ne00_padded;
layout->n_chunks = 1;
layout->n_threads = n_threads;
layout->n_slots = n_slots;
return true;
}
uint32_t n_vec = (ne00 + 31) / 32;
uint32_t n_vec_pow2 = 1;
while (n_vec_pow2 < n_vec) n_vec_pow2 <<= 1;
uint32_t n_chunks = 1;
if (ne00 > 1024) {
while (n_chunks * 2 <= n_threads_max && n_chunks * 2 <= n_vec_pow2) {
n_chunks *= 2;
}
}
uint32_t chunk_n_vec = n_vec_pow2 / n_chunks;
uint32_t chunk_elems = chunk_n_vec * 32;
size_t phase1_values_size = ((chunk_elems * sizeof(float)) + 127) & ~127;
size_t phase1_indices_size = ((chunk_elems * sizeof(int32_t)) + 127) & ~127;
size_t phase1_slot_size = ((phase1_values_size + phase1_indices_size) + 255) & ~255;
size_t phase1_total_size = phase1_slot_size * n_chunks;
size_t merge_values_size = 0;
size_t merge_indices_size = 0;
size_t merge_values_off = phase1_total_size;
size_t merge_indices_off = 0;
uint32_t merge_elems = 0;
if (n_chunks > 1) {
if (is_top_k) {
uint32_t local_k = k < chunk_elems ? k : chunk_elems;
uint32_t total_candidates = n_chunks * local_k;
uint32_t merge_n_vec = (total_candidates + 31) / 32;
uint32_t merge_n_vec_pow2 = 1;
while (merge_n_vec_pow2 < merge_n_vec) merge_n_vec_pow2 <<= 1;
merge_elems = merge_n_vec_pow2 * 32;
} else {
merge_elems = n_vec_pow2 * 32;
}
merge_values_size = ((merge_elems * sizeof(float)) + 127) & ~127;
merge_indices_size = ((merge_elems * sizeof(int32_t)) + 127) & ~127;
merge_indices_off = merge_values_off + merge_values_size;
}
size_t total_bytes = phase1_total_size + merge_values_size + merge_indices_size;
if (total_bytes > vtcm_budget && n_chunks > 1) {
n_chunks = 1;
chunk_elems = n_vec_pow2 * 32;
phase1_values_size = ((chunk_elems * sizeof(float)) + 127) & ~127;
phase1_indices_size = ((chunk_elems * sizeof(int32_t)) + 127) & ~127;
phase1_slot_size = ((phase1_values_size + phase1_indices_size) + 255) & ~255;
phase1_total_size = phase1_slot_size;
merge_values_size = 0;
merge_indices_size = 0;
merge_values_off = phase1_total_size;
merge_indices_off = 0;
merge_elems = 0;
total_bytes = phase1_total_size;
}
if (total_bytes > vtcm_budget) {
return false;
}
uint32_t n_slots = 1;
if (n_chunks == 1 && phase1_slot_size * 2 <= vtcm_budget) {
n_slots = 2;
total_bytes = phase1_slot_size * 2;
}
layout->total_bytes = total_bytes;
layout->phase1_slot_size = phase1_slot_size;
layout->merge_values_off = merge_values_off;
layout->merge_indices_off = merge_indices_off;
layout->chunk_elems = chunk_elems;
layout->n_chunks = n_chunks;
layout->merge_elems = merge_elems;
layout->n_threads = n_chunks;
layout->n_slots = n_slots;
return true;
}
#if defined(__cplusplus)
static_assert(sizeof(struct htp_sort_kernel_params) <= 128, "htp_sort_kernel_params is too large for kernel_params blob");
#else
_Static_assert(sizeof(struct htp_sort_kernel_params) <= 128, "htp_sort_kernel_params is too large for kernel_params blob");
#endif
#endif // HTP_ARGSORT_OPS_H
+277
View File
@@ -917,6 +917,154 @@ static void binary_thread_add_id_f32(unsigned int nth, unsigned int ith, void *
dma_queue_flush(dma_q);
}
static inline void hvx_div_scalar_f32(uint8_t * restrict dst, const uint8_t * restrict src, const float val, const uint32_t num_elems) {
hvx_mul_scalar_f32(dst, src, 1.0f / val, num_elems);
}
static inline void hvx_div_scalar_f16(uint8_t * restrict dst, const uint8_t * restrict src, const _Float16 val, const uint32_t num_elems) {
hvx_div_scalar_f16_aa(dst, src, val, num_elems);
}
typedef void (*compute_binary_chunked_t)(
uint8_t * restrict dst,
const uint8_t * restrict src0,
const uint8_t * restrict src1,
const uint32_t num_elems
);
typedef void (*compute_binary_scalar_chunked_f32_t)(
uint8_t * restrict dst,
const uint8_t * restrict src0,
const float val,
const uint32_t num_elems
);
typedef void (*compute_binary_scalar_chunked_f16_t)(
uint8_t * restrict dst,
const uint8_t * restrict src0,
const _Float16 val,
const uint32_t num_elems
);
struct binary_chunked_context {
struct htp_ops_context * octx;
struct htp_binary_vtcm_layout vtcm_layout;
uint8_t * vtcm_base;
uint32_t elem_start;
uint32_t nelem;
uint32_t chunk_size;
uint32_t chunks_per_thread;
uint32_t total_chunks;
bool is_scalar;
float scalar_f32;
_Float16 scalar_f16;
compute_binary_chunked_t compute;
compute_binary_scalar_chunked_f32_t compute_scalar_f32;
compute_binary_scalar_chunked_f16_t compute_scalar_f16;
};
static void binary_thread_chunked(unsigned int nth, unsigned int ith, void * data) {
(void) nth;
struct binary_chunked_context * ctx = (struct binary_chunked_context *) data;
struct htp_ops_context * octx = ctx->octx;
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * src1 = octx->src[1];
const struct htp_tensor * dst = octx->dst;
const uint32_t start_chunk = ctx->chunks_per_thread * ith;
const uint32_t end_chunk = MIN(start_chunk + ctx->chunks_per_thread, ctx->total_chunks);
if (start_chunk >= end_chunk) {
return;
}
const uint32_t src0_type = src0->type;
const size_t elem_size = (src0_type == HTP_TYPE_F32) ? sizeof(float) : sizeof(_Float16);
const uint32_t chunk_size = ctx->chunk_size;
const size_t chunk_bytes = ctx->vtcm_layout.src0_spad_half_size;
FARF(HIGH, "binary-chunked: %d/%d (%u:%u) chunks %u elems %u",
ith, nth, start_chunk, end_chunk, ctx->total_chunks, ctx->nelem);
const struct htp_binary_vtcm_layout * layout = &ctx->vtcm_layout;
uint8_t * src0_spad_base = VTCM_LAYOUT_PTR(uint8_t, ctx->vtcm_base, layout->off_src0) + (ith * layout->src0_bytes_per_thread);
uint8_t * src1_spad_base = ctx->is_scalar ? NULL : (VTCM_LAYOUT_PTR(uint8_t, ctx->vtcm_base, layout->off_src1) + (ith * layout->src1_bytes_per_thread));
uint8_t * dst_spad_base = VTCM_LAYOUT_PTR(uint8_t, ctx->vtcm_base, layout->off_dst) + (ith * layout->dst_bytes_per_thread);
dma_queue * dma_q = octx->ctx->dma[ith];
uint32_t prefetch_chunk = start_chunk;
int spad_idx = 0;
for (int k = 0; k < 2 && prefetch_chunk < end_chunk; k++) {
const uint32_t c_start = ctx->elem_start + prefetch_chunk * chunk_size;
const uint32_t c_end = MIN(c_start + chunk_size, ctx->elem_start + ctx->nelem);
const uint32_t cur_elems = c_end - c_start;
const uint32_t cur_bytes = cur_elems * elem_size;
dma_addr_t s0_curr = src0->data + (size_t) c_start * elem_size;
dma_addr_t d_curr = dst->data + (size_t) c_start * elem_size;
uint8_t * s0_spad = src0_spad_base + spad_idx * chunk_bytes;
uint8_t * d_spad = dst_spad_base + spad_idx * chunk_bytes;
dma_queue_push(dma_q, dma_make_data(d_curr, d_spad), chunk_bytes, chunk_bytes, cur_bytes, 0);
dma_queue_push(dma_q, dma_make_data(s0_spad, s0_curr), chunk_bytes, chunk_bytes, cur_bytes, 1);
if (!ctx->is_scalar) {
dma_addr_t s1_curr = src1->data + (size_t) c_start * elem_size;
uint8_t * s1_spad = src1_spad_base + spad_idx * chunk_bytes;
dma_queue_push(dma_q, dma_make_data(s1_spad, s1_curr), chunk_bytes, chunk_bytes, cur_bytes, 1);
}
prefetch_chunk++;
spad_idx ^= 1;
}
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
for (uint32_t c = start_chunk; c < end_chunk; c++) {
const uint32_t c_start = ctx->elem_start + c * chunk_size;
const uint32_t c_end = MIN(c_start + chunk_size, ctx->elem_start + ctx->nelem);
const uint32_t cur_elems = c_end - c_start;
const uint32_t cur_bytes = cur_elems * elem_size;
uint8_t * d_spad = (uint8_t *) dma_queue_pop(dma_q).src;
uint8_t * s0_spad = (uint8_t *) dma_queue_pop(dma_q).dst;
uint8_t * s1_spad = ctx->is_scalar ? NULL : (uint8_t *) dma_queue_pop(dma_q).dst;
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) c);
if (ctx->is_scalar) {
if (src0_type == HTP_TYPE_F32) {
ctx->compute_scalar_f32(d_spad, s0_spad, ctx->scalar_f32, cur_elems);
} else {
ctx->compute_scalar_f16(d_spad, s0_spad, ctx->scalar_f16, cur_elems);
}
} else {
ctx->compute(d_spad, s0_spad, s1_spad, cur_elems);
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) c);
dma_addr_t dst_curr = dst->data + (size_t) c_start * elem_size;
dma_queue_push(dma_q, dma_make_data(dst_curr, d_spad), chunk_bytes, chunk_bytes, cur_bytes, 1);
if (prefetch_chunk < end_chunk) {
const uint32_t pc_start = ctx->elem_start + prefetch_chunk * chunk_size;
const uint32_t pc_end = MIN(pc_start + chunk_size, ctx->elem_start + ctx->nelem);
const uint32_t p_elems = pc_end - pc_start;
const uint32_t p_bytes = p_elems * elem_size;
dma_addr_t s0_next = src0->data + (size_t) pc_start * elem_size;
dma_queue_push(dma_q, dma_make_data(s0_spad, s0_next), chunk_bytes, chunk_bytes, p_bytes, 1);
if (!ctx->is_scalar) {
dma_addr_t s1_next = src1->data + (size_t) pc_start * elem_size;
dma_queue_push(dma_q, dma_make_data(s1_spad, s1_next), chunk_bytes, chunk_bytes, p_bytes, 1);
}
prefetch_chunk++;
}
}
dma_queue_flush(dma_q);
}
static int execute_op_binary(struct htp_ops_context * octx) {
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * src1 = octx->src[1];
@@ -932,6 +1080,135 @@ static int execute_op_binary(struct htp_ops_context * octx) {
const size_t src1_row_size = src1->ne[0] * elem_size;
const size_t dst_row_size = dst->ne[0] * elem_size;
if (kparams->kernel_type == HTP_BINARY_KERNEL_CHUNKED) {
const bool is_scalar = kparams->is_scalar || (src1->ne[0] == 1 && src1->ne[1] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1);
compute_binary_chunked_t compute = NULL;
compute_binary_scalar_chunked_f32_t compute_scalar_f32 = NULL;
compute_binary_scalar_chunked_f16_t compute_scalar_f16 = NULL;
if (is_scalar) {
if (src0_type == HTP_TYPE_F32) {
switch (octx->op) {
case HTP_OP_ADD: compute_scalar_f32 = hvx_add_scalar_f32; break;
case HTP_OP_SUB: compute_scalar_f32 = hvx_sub_scalar_f32; break;
case HTP_OP_MUL: compute_scalar_f32 = hvx_mul_scalar_f32; break;
case HTP_OP_DIV: compute_scalar_f32 = hvx_div_scalar_f32; break;
default: break;
}
} else if (src0_type == HTP_TYPE_F16) {
switch (octx->op) {
case HTP_OP_ADD: compute_scalar_f16 = hvx_add_scalar_f16; break;
case HTP_OP_SUB: compute_scalar_f16 = hvx_sub_scalar_f16; break;
case HTP_OP_MUL: compute_scalar_f16 = hvx_mul_scalar_f16; break;
case HTP_OP_DIV: compute_scalar_f16 = hvx_div_scalar_f16; break;
default: break;
}
}
if (!compute_scalar_f32 && !compute_scalar_f16) {
return HTP_STATUS_NO_SUPPORT;
}
} else {
if (src0_type == HTP_TYPE_F32) {
switch (octx->op) {
case HTP_OP_ADD: compute = hvx_add_f32; break;
case HTP_OP_SUB: compute = hvx_sub_f32; break;
case HTP_OP_MUL: compute = hvx_mul_f32; break;
case HTP_OP_DIV: compute = hvx_div_f32; break;
default: break;
}
} else if (src0_type == HTP_TYPE_F16) {
switch (octx->op) {
case HTP_OP_ADD: compute = hvx_add_f16; break;
case HTP_OP_SUB: compute = hvx_sub_f16; break;
case HTP_OP_MUL: compute = hvx_mul_f16; break;
case HTP_OP_DIV: compute = hvx_div_f16; break;
default: break;
}
}
if (!compute) {
return HTP_STATUS_NO_SUPPORT;
}
}
const uint32_t total_elems = (uint32_t) (src0->ne[0] * src0->ne[1] * src0->ne[2] * src0->ne[3]);
if (total_elems == 0) {
return HTP_STATUS_OK;
}
uint32_t elem_start = 0;
uint32_t nelem = total_elems;
const uint32_t elems_per_line = (src0_type == HTP_TYPE_F32) ? 32 : 64;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_mdev_data_aligned(dst) &&
htp_tensor_mdev_data_aligned(src0) &&
(is_scalar || htp_tensor_mdev_data_aligned(src1)) &&
htp_tensor_is_contiguous(dst, elem_size) &&
htp_tensor_is_contiguous(src0, elem_size) &&
(is_scalar || htp_tensor_is_contiguous(src1, elem_size));
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(
total_elems, can_split ? elems_per_line : 0,
octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
elem_start = range.start;
nelem = range.count;
}
if (nelem == 0) {
return HTP_STATUS_OK;
}
if (!htp_ops_context_set_n_threads(octx, kparams->n_threads)) {
return HTP_STATUS_INVAL_PARAMS;
}
struct htp_binary_vtcm_layout vtcm_layout;
htp_binary_vtcm_layout_build(&vtcm_layout, kparams, octx->ctx->vtcm_size);
if (vtcm_layout.total_bytes == 0 || vtcm_layout.total_bytes > octx->ctx->vtcm_size) {
return HTP_STATUS_VTCM_TOO_SMALL;
}
const uint32_t chunk_size = kparams->chunk_size > 0 ? kparams->chunk_size : (32768 / elem_size);
const uint32_t total_chunks = (nelem + chunk_size - 1) / chunk_size;
const uint32_t n_threads = (total_chunks >= 2) ? MIN(octx->n_threads, total_chunks) : 1;
const uint32_t chunks_per_thread = (total_chunks + n_threads - 1) / n_threads;
float scalar_f32 = 0.0f;
_Float16 scalar_f16 = 0;
if (is_scalar) {
uint8_t * vtcm_src1 = VTCM_LAYOUT_PTR(uint8_t, octx->ctx->vtcm_base, vtcm_layout.off_src1);
dma_queue * dma_q = octx->ctx->dma[0];
dma_queue_push(dma_q, dma_make_data(vtcm_src1, src1->data), 128, 0, elem_size, 1);
dma_queue_pop(dma_q);
if (src0_type == HTP_TYPE_F32) {
scalar_f32 = ((const float *) vtcm_src1)[0];
} else {
scalar_f16 = ((const _Float16 *) vtcm_src1)[0];
}
}
struct binary_chunked_context cctx = {
.octx = octx,
.vtcm_layout = vtcm_layout,
.vtcm_base = (uint8_t *) octx->ctx->vtcm_base,
.elem_start = elem_start,
.nelem = nelem,
.chunk_size = chunk_size,
.chunks_per_thread = chunks_per_thread,
.total_chunks = total_chunks,
.is_scalar = is_scalar,
.scalar_f32 = scalar_f32,
.scalar_f16 = scalar_f16,
.compute = compute,
.compute_scalar_f32 = compute_scalar_f32,
.compute_scalar_f16 = compute_scalar_f16,
};
work_queue_run(octx->ctx->work_queue, binary_thread_chunked, &cctx, n_threads);
return HTP_STATUS_OK;
}
uint32_t row_start = 0;
uint32_t nrows = src0_nrows;
+39
View File
@@ -17,6 +17,7 @@ enum htp_binary_kernel_type {
HTP_BINARY_KERNEL_ADD_ID,
HTP_BINARY_KERNEL_COMPLEX,
HTP_BINARY_KERNEL_REPEAT,
HTP_BINARY_KERNEL_CHUNKED,
};
struct htp_binary_kernel_params {
@@ -30,6 +31,10 @@ struct htp_binary_kernel_params {
uint32_t src1_size;
uint32_t vtcm_size;
uint32_t chunk_size;
uint32_t chunk_bytes;
uint32_t is_scalar;
};
#if defined(__cplusplus)
@@ -68,6 +73,40 @@ static inline void htp_binary_vtcm_layout_build(
return;
}
if (kparams->kernel_type == HTP_BINARY_KERNEL_CHUNKED) {
const size_t chunk_bytes = kparams->chunk_bytes;
if (chunk_bytes == 0) {
return;
}
L->src0_bytes_per_thread = 2 * chunk_bytes;
L->src1_bytes_per_thread = kparams->is_scalar ? 0 : (2 * chunk_bytes);
L->dst_bytes_per_thread = 2 * chunk_bytes;
L->src0_spad_half_size = chunk_bytes;
L->src1_spad_half_size = kparams->is_scalar ? 0 : chunk_bytes;
L->dst_spad_half_size = chunk_bytes;
L->rows_per_buffer = 1;
L->src1_size = 0;
const size_t src0_total = n_threads * L->src0_bytes_per_thread;
const size_t src1_total = kparams->is_scalar ? 128 : (n_threads * L->src1_bytes_per_thread);
const size_t dst_total = n_threads * L->dst_bytes_per_thread;
size_t off = 0;
VTCM_LAYOUT_ALLOC(off, off_src0, src0_total);
VTCM_LAYOUT_ALLOC(off, off_src1, src1_total);
VTCM_LAYOUT_ALLOC(off, off_dst, dst_total);
if (off > vtcm_size) {
return;
}
L->total_bytes = off;
return;
}
const size_t spad_row_total = (kparams->kernel_type == HTP_BINARY_KERNEL_SAME_SHAPE)
? 2 * (kparams->src0_row_size_aligned + kparams->src1_row_size_aligned + kparams->dst_row_size_aligned)
: 2 * (kparams->src0_row_size_aligned + kparams->dst_row_size_aligned);
@@ -136,7 +136,7 @@ set(CMAKE_SHARED_LIBRARY_SONAME_C_FLAG "-Wl,-soname,")
set(CMAKE_SHARED_LIBRARY_SONAME_CXX_FLAG "-Wl,-soname,")
# Compiler Options
set(COMMON_FLAGS "${ARCH_FLAGS} -fvectorize -flto -Wall -Werror -fno-zero-initialized-in-bss -G0 -fdata-sections -fpic ${XQF_ARGS}")
set(COMMON_FLAGS "${ARCH_FLAGS} -fno-vectorize -fno-slp-vectorize -flto -Wall -Werror -fno-zero-initialized-in-bss -G0 -fdata-sections -fpic ${XQF_ARGS}")
set(CMAKE_CXX_FLAGS_DEBUG "${COMMON_FLAGS} -O0 -D_DEBUG -g")
set(CMAKE_CXX_FLAGS_RELWITHDEBINFO "${COMMON_FLAGS} -O2 -g")
+120
View File
@@ -301,6 +301,86 @@ static void cpy_thread_f32_f16_sameshape(unsigned int nth, unsigned int ith, voi
}
}
static void cpy_thread_i32_f32_sameshape(unsigned int nth, unsigned int ith, void * data) {
struct htp_copy_context * ct = (struct htp_copy_context *) data;
struct htp_ops_context * octx = ct->octx;
cpy_preamble;
const uint32_t dr = ct->src0_nrows_per_thread;
const uint32_t ir0 = ct->row_start + dr * ith;
const uint32_t ir1 = MIN(ir0 + dr, ct->row_start + ct->nrows);
if (ir0 >= ir1) return;
const uint32_t ne02_ne01 = ne02 * ne01;
uint32_t i03 = fastdiv(ir0, &ct->div_ne02_ne01);
uint32_t rem = ir0 - i03 * ne02_ne01;
uint32_t i02 = fastdiv(rem, &ct->div_ne01);
uint32_t i01 = rem - i02 * ne01;
uint8_t* dst_ptr = (uint8_t*) dst->data + i01*nb1 + i02*nb2 + i03*nb3;
uint8_t* src0_ptr = (uint8_t*) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
for (uint32_t r = ir0; r < ir1; r++) {
hex_l2fetch(src0_ptr, ne00 * sizeof(float), nb01, 2);
const float * restrict src_row = (const float *) src0_ptr;
int32_t * restrict dst_row = (int32_t *) dst_ptr;
for (uint32_t i = 0; i < ne00; i++) {
dst_row[i] = (int32_t) src_row[i];
}
dst_ptr += nb1;
src0_ptr += nb01;
if (++i01 == ne01) {
i01 = 0;
if (++i02 == ne02) {
i02 = 0;
i03++;
}
dst_ptr = (uint8_t*) dst->data + i02*nb2 + i03*nb3;
src0_ptr = (uint8_t*) src0->data + i02*nb02 + i03*nb03;
}
}
}
static void cpy_thread_f32_i32_sameshape(unsigned int nth, unsigned int ith, void * data) {
struct htp_copy_context * ct = (struct htp_copy_context *) data;
struct htp_ops_context * octx = ct->octx;
cpy_preamble;
const uint32_t dr = ct->src0_nrows_per_thread;
const uint32_t ir0 = ct->row_start + dr * ith;
const uint32_t ir1 = MIN(ir0 + dr, ct->row_start + ct->nrows);
if (ir0 >= ir1) return;
const uint32_t ne02_ne01 = ne02 * ne01;
uint32_t i03 = fastdiv(ir0, &ct->div_ne02_ne01);
uint32_t rem = ir0 - i03 * ne02_ne01;
uint32_t i02 = fastdiv(rem, &ct->div_ne01);
uint32_t i01 = rem - i02 * ne01;
uint8_t* dst_ptr = (uint8_t*) dst->data + i01*nb1 + i02*nb2 + i03*nb3;
uint8_t* src0_ptr = (uint8_t*) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
for (uint32_t r = ir0; r < ir1; r++) {
hex_l2fetch(src0_ptr, ne00 * sizeof(int32_t), nb01, 2);
const int32_t * restrict src_row = (const int32_t *) src0_ptr;
float * restrict dst_row = (float *) dst_ptr;
for (uint32_t i = 0; i < ne00; i++) {
dst_row[i] = (float) src_row[i];
}
dst_ptr += nb1;
src0_ptr += nb01;
if (++i01 == ne01) {
i01 = 0;
if (++i02 == ne02) {
i02 = 0;
i03++;
}
dst_ptr = (uint8_t*) dst->data + i02*nb2 + i03*nb3;
src0_ptr = (uint8_t*) src0->data + i02*nb02 + i03*nb03;
}
}
}
static inline void cpy_dma_push_2d_chunked(
dma_queue * dma_q,
dma_addr_t dst,
@@ -366,6 +446,42 @@ static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {
cpy_preamble;
*use_dma = false;
const uint32_t total_elems_src = ne00 * ne01 * ne02 * ne03;
const uint32_t total_elems_dst = ne0 * ne1 * ne2 * ne3;
if (total_elems_src == 1 && total_elems_dst == 1) {
if (octx->ctx->mdev.count > 1 && octx->ctx->mdev.idx > 0) {
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_F32 && dst->type == HTP_TYPE_I32) {
((int32_t *) dst->data)[0] = (int32_t) (((const float *) src0->data)[0]);
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_I32 && dst->type == HTP_TYPE_F32) {
((float *) dst->data)[0] = (float) (((const int32_t *) src0->data)[0]);
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_I32 && dst->type == HTP_TYPE_I32) {
((int32_t *) dst->data)[0] = ((const int32_t *) src0->data)[0];
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_F32 && dst->type == HTP_TYPE_F32) {
((float *) dst->data)[0] = ((const float *) src0->data)[0];
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_F16 && dst->type == HTP_TYPE_F16) {
((__fp16 *) dst->data)[0] = ((const __fp16 *) src0->data)[0];
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_F32 && dst->type == HTP_TYPE_F16) {
((__fp16 *) dst->data)[0] = (__fp16) (((const float *) src0->data)[0]);
return HTP_STATUS_OK;
}
if (src0->type == HTP_TYPE_F16 && dst->type == HTP_TYPE_F32) {
((float *) dst->data)[0] = (float) (((const __fp16 *) src0->data)[0]);
return HTP_STATUS_OK;
}
}
struct htp_copy_context ct;
ct.octx = octx;
@@ -450,6 +566,10 @@ static int exec_cpy(struct htp_ops_context * octx, bool * use_dma) {
copy_fun = cpy_thread_f16_f32_sameshape;
} else if (dst->type == HTP_TYPE_F32 && src0->type == HTP_TYPE_F16) {
copy_fun = cpy_thread_f32_f16_sameshape;
} else if (dst->type == HTP_TYPE_I32 && src0->type == HTP_TYPE_F32) {
copy_fun = cpy_thread_i32_f32_sameshape;
} else if (dst->type == HTP_TYPE_F32 && src0->type == HTP_TYPE_I32) {
copy_fun = cpy_thread_f32_i32_sameshape;
} else {
return HTP_STATUS_NO_SUPPORT;
}
+16
View File
@@ -426,6 +426,22 @@ static inline bool dma_queue_push(dma_queue *q, dma_data ddata, size_t dst_strid
#endif
static inline void dma_sync_read(dma_queue * dma_q, void * dst, dma_addr_t src, size_t bytes) {
const uint32_t b = (uint32_t) bytes;
if (b > 0) {
dma_queue_push(dma_q, dma_make_data(dst, src), b, b, b, 1);
dma_queue_pop(dma_q);
}
}
static inline void dma_sync_write(dma_queue * dma_q, dma_addr_t dst, const void * src, size_t bytes) {
const uint32_t b = (uint32_t) bytes;
if (b > 0) {
dma_queue_push(dma_q, dma_make_data(dst, src), b, b, b, 1);
dma_queue_pop(dma_q);
}
}
#define DMA_CACHE_MAX_SIZE 256U
// Fully assoc LRU cache
+2
View File
@@ -150,7 +150,9 @@ int op_matmul_nx(struct htp_ops_context * octx);
int op_matmul_id_nx(struct htp_ops_context * octx);
int op_binary(struct htp_ops_context * octx);
int op_unary(struct htp_ops_context * octx);
int op_sum(struct htp_ops_context * octx);
int op_sum_rows(struct htp_ops_context * octx);
int op_argmax(struct htp_ops_context * octx);
int op_activations(struct htp_ops_context * octx);
int op_softmax(struct htp_ops_context * octx);
int op_add_id(struct htp_ops_context * octx);
+3
View File
@@ -69,6 +69,7 @@ enum htp_op_code {
HTP_OP_UNARY_ABS,
HTP_OP_UNARY_LOG,
HTP_OP_UNARY_RELU,
HTP_OP_UNARY_STEP,
HTP_OP_GLU_SWIGLU,
HTP_OP_GLU_SWIGLU_OAI,
HTP_OP_GLU_GEGLU,
@@ -86,6 +87,7 @@ enum htp_op_code {
HTP_OP_TOP_K,
HTP_OP_SQR,
HTP_OP_SQRT,
HTP_OP_SUM,
HTP_OP_SUM_ROWS,
HTP_OP_SSM_CONV,
HTP_OP_REPEAT,
@@ -108,6 +110,7 @@ enum htp_op_code {
HTP_OP_GLU_SWIGLU_CLAMP,
HTP_OP_MDEV_GROUP,
HTP_OP_ROLL,
HTP_OP_ARGMAX,
HTP_OP_INVALID
};
+52
View File
@@ -579,6 +579,58 @@ static inline void hvx_abs_f16(uint8_t * restrict dst, const uint8_t * restrict
}
}
//
// Step
//
static inline void hvx_step_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) {
assert((unsigned long) dst % 128 == 0);
assert((unsigned long) src % 128 == 0);
HVX_Vector * restrict vdst = (HVX_Vector *) dst;
HVX_Vector * restrict vsrc = (HVX_Vector *) src;
const uint32_t elem_size = sizeof(float);
const uint32_t epv = 128 / elem_size;
const uint32_t nvec = n / epv;
const uint32_t nloe = n % epv;
uint32_t i = 0;
_Pragma("unroll(4)")
for (; i < nvec; i++) {
vdst[i] = hvx_vec_step_f32(vsrc[i]);
}
if (nloe) {
HVX_Vector v = hvx_vec_step_f32(vsrc[i]);
hvx_vec_store_a((void *) &vdst[i], nloe * elem_size, v);
}
}
static inline void hvx_step_f16_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) {
assert((unsigned long) dst % 128 == 0);
assert((unsigned long) src % 128 == 0);
HVX_Vector * restrict vdst = (HVX_Vector *) dst;
HVX_Vector * restrict vsrc = (HVX_Vector *) src;
const uint32_t elem_size = sizeof(_Float16);
const uint32_t epv = 128 / elem_size;
const uint32_t nvec = n / epv;
const uint32_t nloe = n % epv;
uint32_t i = 0;
_Pragma("unroll(4)")
for (; i < nvec; i++) {
vdst[i] = hvx_vec_step_f16(vsrc[i]);
}
if (nloe) {
HVX_Vector v = hvx_vec_step_f16(vsrc[i]);
hvx_vec_store_a((void *) &vdst[i], nloe * elem_size, v);
}
}
//
// Square
//
+14
View File
@@ -111,6 +111,20 @@ static inline HVX_Vector hvx_vec_neg_f32(HVX_Vector v) {
#endif // __HVX_ARCH__ > 75
}
static inline HVX_Vector hvx_vec_step_f32(HVX_Vector v) {
const HVX_Vector zero = Q6_V_vzero();
const HVX_Vector one = hvx_vec_splat_f32(1.0f);
HVX_VectorPred q = Q6_Q_vcmp_gt_VsfVsf(v, zero);
return Q6_V_vmux_QVV(q, one, zero);
}
static inline HVX_Vector hvx_vec_step_f16(HVX_Vector v) {
const HVX_Vector zero = Q6_V_vzero();
const HVX_Vector one = hvx_vec_splat_f16((_Float16) 1.0f);
HVX_VectorPred q = Q6_Q_vcmp_gt_VhfVhf(v, zero);
return Q6_V_vmux_QVV(q, one, zero);
}
static inline HVX_VectorPred hvx_vec_is_nan_f16(HVX_Vector v) {
const HVX_Vector vnan_exp = Q6_Vh_vsplat_R(0x7C00);
const HVX_Vector vnan_frac = Q6_Vh_vsplat_R(0x7FFF);
+73
View File
@@ -326,6 +326,79 @@ static inline int32_t hvx_reduce_max_i32(const uint8_t * restrict src, const int
}
}
static inline void hvx_argmax_f32(
const float * restrict src,
uint32_t n,
uint32_t offset,
float * out_val,
int32_t * out_idx
) {
if (n == 0) {
*out_val = -INFINITY;
*out_idx = (int32_t) offset;
return;
}
if (n < 32 || !hex_is_aligned((void *) src, 128)) {
float best_val = src[0];
int32_t best_idx = (int32_t) offset;
for (uint32_t i = 1; i < n; i++) {
if (src[i] > best_val) {
best_val = src[i];
best_idx = (int32_t) (offset + i);
}
}
*out_val = best_val;
*out_idx = best_idx;
return;
}
static const int32_t c_lane_idx[32] __attribute__((aligned(128))) = {
0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15,
16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31
};
const HVX_Vector v_init_idx = *(const HVX_Vector *) c_lane_idx;
const HVX_Vector v_step = Q6_V_vsplat_R(32);
HVX_Vector v_cur_idx = Q6_Vw_vadd_VwVw(v_init_idx, Q6_V_vsplat_R((int32_t) offset));
HVX_Vector v_max_val = hvx_vec_splat_f32(-INFINITY);
HVX_Vector v_max_idx = v_cur_idx;
const HVX_Vector * vsrc = (const HVX_Vector *) src;
const uint32_t nvec = n / 32;
for (uint32_t vi = 0; vi < nvec; vi++) {
HVX_Vector v = vsrc[vi];
HVX_VectorPred pred = Q6_Q_vcmp_gt_VsfVsf(v, v_max_val);
v_max_val = Q6_V_vmux_QVV(pred, v, v_max_val);
v_max_idx = Q6_V_vmux_QVV(pred, v_cur_idx, v_max_idx);
v_cur_idx = Q6_Vw_vadd_VwVw(v_cur_idx, v_step);
}
HVX_VectorAlias u_val, u_idx;
u_val.v = v_max_val;
u_idx.v = v_max_idx;
float best_val = u_val.fp32[0];
int32_t best_idx = (int32_t) u_idx.w[0];
for (int i = 1; i < 32; i++) {
if (u_val.fp32[i] > best_val) {
best_val = u_val.fp32[i];
best_idx = (int32_t) u_idx.w[i];
}
}
for (uint32_t i = nvec * 32; i < n; i++) {
if (src[i] > best_val) {
best_val = src[i];
best_idx = (int32_t) (offset + i);
}
}
*out_val = best_val;
*out_idx = best_idx;
}
#undef hvx_reduce_loop_body
#undef HVX_REDUCE_MAX_OP
#undef HVX_REDUCE_SUM_OP
+7
View File
@@ -857,6 +857,7 @@ static int execute_op(struct htp_ops_context * octx) {
case HTP_OP_UNARY_ABS:
case HTP_OP_UNARY_LOG:
case HTP_OP_UNARY_RELU:
case HTP_OP_UNARY_STEP:
case HTP_OP_L2_NORM:
return op_unary(octx);
@@ -882,6 +883,9 @@ static int execute_op(struct htp_ops_context * octx) {
case HTP_OP_GET_ROWS:
return op_get_rows(octx);
case HTP_OP_SUM:
return op_sum(octx);
case HTP_OP_SUM_ROWS:
return op_sum_rows(octx);
@@ -898,6 +902,9 @@ static int execute_op(struct htp_ops_context * octx) {
case HTP_OP_TOP_K:
return op_top_k(octx);
case HTP_OP_ARGMAX:
return op_argmax(octx);
case HTP_OP_SSM_CONV:
return op_ssm_conv(octx);
+260 -4
View File
@@ -83,16 +83,17 @@ static void sum_rows_thread_f32(unsigned int nth, unsigned int ith, void *data)
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start_row);
for (uint32_t ir = 0; ir < n_rows; ir++) {
const float * restrict src_local = src_th + (ir * (src_stride / sizeof(float)));
const float * restrict src_local = (const float *) ((const uint8_t *) src_th + ir * src_stride);
float * restrict dst_local = (float *) ((uint8_t *) dst_th + ir * dst_stride);
if (ir + 1 < n_rows) {
hex_l2fetch(src_local + (src_stride / sizeof(float)), src_stride, src_stride, 1);
hex_l2fetch((const uint8_t *) src_local + src_stride, src_stride, src_stride, 1);
}
if (opt_path) {
dst_th[ir] = hvx_reduce_sum_f32_a((const uint8_t *) src_local, ne00);
*dst_local = hvx_reduce_sum_f32_a((const uint8_t *) src_local, ne00);
} else {
dst_th[ir] = hvx_reduce_sum_f32((const uint8_t *) src_local, ne00);
*dst_local = hvx_reduce_sum_f32((const uint8_t *) src_local, ne00);
}
}
@@ -152,3 +153,258 @@ int op_sum_rows(struct htp_ops_context * octx) {
return HTP_STATUS_OK;
}
struct sum_context {
struct htp_ops_context * octx;
const float * src_data;
float partial_sums[HTP_MAX_NTHREADS];
uint32_t total_elems;
uint32_t elems_per_thread;
};
static void sum_thread_f32(unsigned int nth, unsigned int ith, void * data) {
struct sum_context * sctx = (struct sum_context *) data;
const uint32_t start = sctx->elems_per_thread * ith;
const uint32_t end = MIN(start + sctx->elems_per_thread, sctx->total_elems);
if (start >= end) {
sctx->partial_sums[ith] = 0.0f;
return;
}
const uint32_t n = end - start;
const float * src = sctx->src_data + start;
struct htp_thread_trace * tr = &sctx->octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
hex_l2fetch_block((const void *) src, n * sizeof(float));
sctx->partial_sums[ith] = hvx_reduce_sum_f32((const uint8_t *) src, n);
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
}
int op_sum(struct htp_ops_context * octx) {
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * dst = octx->dst;
if (src0->type != HTP_TYPE_F32) {
return HTP_STATUS_NO_SUPPORT;
}
if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
return HTP_STATUS_NO_SUPPORT;
}
if (octx->ctx->mdev.count > 1 && octx->ctx->mdev.idx > 0) {
return HTP_STATUS_OK;
}
const uint32_t total_elems = (uint32_t) (src0->ne[0] * src0->ne[1] * src0->ne[2] * src0->ne[3]);
if (total_elems == 0) {
((float *) dst->data)[0] = 0.0f;
return HTP_STATUS_OK;
}
const uint32_t n_threads = (total_elems >= 1024) ? MIN(octx->n_threads, HTP_MAX_NTHREADS) : 1;
const uint32_t raw_chunk = (total_elems + n_threads - 1) / n_threads;
const uint32_t elems_per_thread = hex_round_up(raw_chunk, 32);
struct sum_context sctx = {
.octx = octx,
.src_data = (const float *) src0->data,
.total_elems = total_elems,
.elems_per_thread = elems_per_thread,
};
work_queue_run(octx->ctx->work_queue, sum_thread_f32, &sctx, n_threads);
float sum = 0.0f;
for (uint32_t i = 0; i < n_threads; i++) {
sum += sctx.partial_sums[i];
}
((float *) dst->data)[0] = sum;
return HTP_STATUS_OK;
}
static inline void argmax_slice_f32(
const float * restrict src,
uint32_t n,
uint32_t offset,
float * out_val,
int32_t * out_idx
) {
hvx_argmax_f32(src, n, offset, out_val, out_idx);
}
struct argmax_context {
struct htp_ops_context * octx;
const float * src_data;
int32_t * dst_data;
uint32_t ne00;
uint32_t src_stride;
uint32_t dst_stride;
uint32_t row_start;
uint32_t nrows;
uint32_t rows_per_thread;
uint32_t elems_per_thread;
float partial_max[HTP_MAX_NTHREADS];
int32_t partial_idx[HTP_MAX_NTHREADS];
};
static void argmax_thread_single_row(unsigned int nth, unsigned int ith, void * data) {
struct argmax_context * actx = (struct argmax_context *) data;
const uint32_t start = actx->elems_per_thread * ith;
const uint32_t end = MIN(start + actx->elems_per_thread, actx->ne00);
if (start >= end) {
actx->partial_max[ith] = -INFINITY;
actx->partial_idx[ith] = 0;
return;
}
const uint32_t n = end - start;
const float * src = actx->src_data + start;
struct htp_thread_trace * tr = &actx->octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
hex_l2fetch_block((const void *) src, n * sizeof(float));
float max_val;
int32_t max_idx;
argmax_slice_f32(src, n, start, &max_val, &max_idx);
actx->partial_max[ith] = max_val;
actx->partial_idx[ith] = max_idx;
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
}
static void argmax_thread_multi_row(unsigned int nth, unsigned int ith, void * data) {
struct argmax_context * actx = (struct argmax_context *) data;
const uint32_t r0 = actx->row_start + actx->rows_per_thread * ith;
const uint32_t r1 = MIN(r0 + actx->rows_per_thread, actx->row_start + actx->nrows);
if (r0 >= r1) {
return;
}
struct htp_thread_trace * tr = &actx->octx->ctx->trace[ith];
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) r0);
for (uint32_t r = r0; r < r1; r++) {
const float * src_row = (const float *) ((const uint8_t *) actx->src_data + r * actx->src_stride);
int32_t * dst_val = (int32_t *) ((uint8_t *) actx->dst_data + r * actx->dst_stride);
hex_l2fetch_block((const void *) src_row, actx->ne00 * sizeof(float));
float max_val;
int32_t max_idx;
argmax_slice_f32(src_row, actx->ne00, 0, &max_val, &max_idx);
*dst_val = max_idx;
}
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) r0);
}
int op_argmax(struct htp_ops_context * octx) {
const struct htp_tensor * src0 = octx->src[0];
const struct htp_tensor * dst = octx->dst;
if (src0->type != HTP_TYPE_F32 || dst->type != HTP_TYPE_I32) {
return HTP_STATUS_NO_SUPPORT;
}
if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
return HTP_STATUS_NO_SUPPORT;
}
const uint32_t ne00 = src0->ne[0];
const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
if (ne00 == 0 || src0_nrows == 0) {
return HTP_STATUS_OK;
}
if (src0_nrows == 1) {
if (octx->ctx->mdev.count > 1 && octx->ctx->mdev.idx > 0) {
return HTP_STATUS_OK;
}
if (ne00 == 1) {
((int32_t *) dst->data)[0] = 0;
return HTP_STATUS_OK;
}
const uint32_t n_threads = (ne00 >= 1024) ? MIN(octx->n_threads, HTP_MAX_NTHREADS) : 1;
const uint32_t raw_chunk = (ne00 + n_threads - 1) / n_threads;
const uint32_t elems_per_thread = hex_round_up(raw_chunk, 32);
struct argmax_context actx = {
.octx = octx,
.src_data = (const float *) src0->data,
.dst_data = (int32_t *) dst->data,
.ne00 = ne00,
.src_stride = src0->nb[1] > 0 ? (uint32_t) src0->nb[1] : (uint32_t) (ne00 * sizeof(float)),
.dst_stride = dst->nb[0] > 0 ? (uint32_t) dst->nb[0] : (uint32_t) sizeof(int32_t),
.row_start = 0,
.nrows = 1,
.rows_per_thread = 1,
.elems_per_thread = elems_per_thread,
};
work_queue_run(octx->ctx->work_queue, argmax_thread_single_row, &actx, n_threads);
float best_val = actx.partial_max[0];
int32_t best_idx = actx.partial_idx[0];
for (uint32_t i = 1; i < n_threads; i++) {
if (actx.partial_max[i] > best_val) {
best_val = actx.partial_max[i];
best_idx = actx.partial_idx[i];
}
}
((int32_t *) dst->data)[0] = best_idx;
return HTP_STATUS_OK;
}
uint32_t row_start = 0;
uint32_t nrows = src0_nrows;
if (octx->ctx->mdev.count > 1) {
const bool can_split = htp_tensor_mdev_data_aligned(dst) &&
htp_tensor_is_contiguous(dst, sizeof(int32_t));
const uint32_t elems_per_chunk = can_split ? 32 : 0;
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(
src0_nrows, elems_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
row_start = range.start;
nrows = range.count;
}
if (nrows == 0) {
return HTP_STATUS_OK;
}
const uint32_t n_threads = MIN(octx->n_threads, nrows);
const uint32_t rows_per_thread = (nrows + n_threads - 1) / n_threads;
struct argmax_context actx = {
.octx = octx,
.src_data = (const float *) src0->data,
.dst_data = (int32_t *) dst->data,
.ne00 = ne00,
.src_stride = src0->nb[1] > 0 ? (uint32_t) src0->nb[1] : (uint32_t) (ne00 * sizeof(float)),
.dst_stride = dst->nb[0] > 0 ? (uint32_t) dst->nb[0] : (uint32_t) sizeof(int32_t),
.row_start = row_start,
.nrows = nrows,
.rows_per_thread = rows_per_thread,
.elems_per_thread = 0,
};
work_queue_run(octx->ctx->work_queue, argmax_thread_multi_row, &actx, n_threads);
return HTP_STATUS_OK;
}
+38
View File
@@ -409,6 +409,20 @@ static void log_f16(const void * restrict src,
}
}
static void step_f16(const void * restrict src,
void * restrict dst,
const uint32_t num_rows,
const struct htp_unary_context * uctx) {
htp_unary_op_preamble;
for (uint32_t ir = 0; ir < num_rows; ir++) {
const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned);
uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned);
hvx_step_f16_aa((uint8_t *) dst_local, (const uint8_t *) src_local, ne0);
}
}
static void l2_norm_f16(const void * restrict src,
void * restrict dst,
const uint32_t num_rows,
@@ -662,6 +676,20 @@ static void relu_f32(const void * restrict src,
}
}
static void step_f32(const void * restrict src,
void * restrict dst,
const uint32_t num_rows,
const struct htp_unary_context * uctx) {
htp_unary_op_preamble;
for (uint32_t ir = 0; ir < num_rows; ir++) {
const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned);
uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned);
hvx_step_f32_aa(dst_local, src_local, ne0);
}
}
static void log_f32(const void * restrict src,
void * restrict dst,
const uint32_t num_rows,
@@ -767,6 +795,11 @@ static void tile_relu_f32(void * restrict dst, const void * restrict src, uint32
hvx_max_scalar_f32((uint8_t *) dst, (const uint8_t *) src, 0.0f, tw);
}
static void tile_step_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
(void) uctx;
hvx_step_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
}
static void tri_apply_tile_f32(const void * restrict src, void * restrict dst,
uint32_t tile_elems, uint32_t col_start, uint32_t i01,
uint32_t ne0, int32_t ttype) {
@@ -1486,6 +1519,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_UNARY_ABS: op_type = is_f16 ? "abs-f16" : "abs-f32"; break;
case HTP_OP_UNARY_LOG: op_type = is_f16 ? "log-f16" : "log-f32"; break;
case HTP_OP_UNARY_RELU: op_type = "relu-f32"; break;
case HTP_OP_UNARY_STEP: op_type = is_f16 ? "step-f16" : "step-f32"; break;
case HTP_OP_L2_NORM: op_type = is_f16 ? "l2norm-f16" : "l2norm-f32"; break;
case HTP_OP_TRI: op_type = "tri-f32"; break;
default:
@@ -1506,6 +1540,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_L2_NORM:
case HTP_OP_UNARY_ABS:
case HTP_OP_UNARY_LOG:
case HTP_OP_UNARY_STEP:
break;
default:
FARF(ERROR, "unary-%s: not supported for F16\n", op_type);
@@ -1634,6 +1669,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_UNARY_ABS: compute_func = (void *) tile_abs_f32; break;
case HTP_OP_UNARY_LOG: compute_func = (void *) tile_log_f32; break;
case HTP_OP_UNARY_RELU: compute_func = (void *) tile_relu_f32; break;
case HTP_OP_UNARY_STEP: compute_func = (void *) tile_step_f32; break;
case HTP_OP_TRI:
task_func = unary_thread_tiled_tri_f32;
compute_func = (void *) tri_apply_tile_f32;
@@ -1652,6 +1688,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_L2_NORM: compute_func = (void *) l2_norm_f16; break;
case HTP_OP_UNARY_ABS: compute_func = (void *) abs_f16; break;
case HTP_OP_UNARY_LOG: compute_func = (void *) log_f16; break;
case HTP_OP_UNARY_STEP: compute_func = (void *) step_f16; break;
default: break;
}
} else {
@@ -1678,6 +1715,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
case HTP_OP_UNARY_ABS: compute_func = (void *) abs_f32; break;
case HTP_OP_UNARY_LOG: compute_func = (void *) log_f32; break;
case HTP_OP_UNARY_RELU: compute_func = (void *) relu_f32; break;
case HTP_OP_UNARY_STEP: compute_func = (void *) step_f32; break;
case HTP_OP_L2_NORM: compute_func = (void *) l2_norm_f32; break;
case HTP_OP_TRI:
task_func = unary_thread_tri_f32;
+1
View File
@@ -59,6 +59,7 @@ static inline bool htp_op_is_unary(uint32_t opcode) {
case HTP_OP_UNARY_ABS:
case HTP_OP_UNARY_LOG:
case HTP_OP_UNARY_RELU:
case HTP_OP_UNARY_STEP:
case HTP_OP_L2_NORM:
case HTP_OP_TRI:
return true;
+1
View File
@@ -9830,6 +9830,7 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
add_test_bin_bcast(type, {5120, 1, 1, 1}, {1, 256, 1, 1});
add_test_bin_bcast(type, {640, 1, 1, 1}, {1, 1, 1, 1});
add_test_bin_bcast(type, {64, 262144, 1, 1}, {1, 1, 1, 1});
add_test_bin_bcast(type, {128, 1, 8, 1}, {1, 4, 1, 1});
//add_test_bin_bcast(type, {3, 3, 2560, 1280}, {1, 1, 1, 1});
//add_test_bin_bcast(type, {3, 3, 2560, 1280}, {2, 1, 1, 1});
}