mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-28 17:07:31 -05:00
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:
+20
-16
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
|
||||
};
|
||||
|
||||
@@ -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
|
||||
//
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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});
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user