mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-01 05:56:52 +02:00
Compare commits
80
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0bc845d356 | ||
|
|
139997d8e7 | ||
|
|
76a5bc86d1 | ||
|
|
46e17a6352 | ||
|
|
fc07d781e6 | ||
|
|
526c43b8f7 | ||
|
|
1c4729414d | ||
|
|
680a036285 | ||
|
|
66e665c427 | ||
|
|
57b557cb95 | ||
|
|
14ebbd5f2f | ||
|
|
f1ea206218 | ||
|
|
6c7a87f7e5 | ||
|
|
f00a64c147 | ||
|
|
d77dd0806d | ||
|
|
6f767fe960 | ||
|
|
f916130d00 | ||
|
|
03a667aa30 | ||
|
|
c2a9e16068 | ||
|
|
4364bf7232 | ||
|
|
ed7ac35e1e | ||
|
|
0c6a6a7ce5 | ||
|
|
81ef10ea58 | ||
|
|
5262471615 | ||
|
|
4da6337767 | ||
|
|
a97cce86a8 | ||
|
|
136887b665 | ||
|
|
9adc7f420c | ||
|
|
6fd50a4094 | ||
|
|
33c923db1b | ||
|
|
c9064dded7 | ||
|
|
c829670992 | ||
|
|
36d7b08340 | ||
|
|
2ebd9ae621 | ||
|
|
cea74625fa | ||
|
|
da6c28eb13 | ||
|
|
d7fb90e8e2 | ||
|
|
7fb2b082ce | ||
|
|
187664b537 | ||
|
|
85ca3b52c3 | ||
|
|
7ac59a6e3a | ||
|
|
2b129ccfa0 | ||
|
|
95887577ab | ||
|
|
694ec23548 | ||
|
|
6f856c7099 | ||
|
|
fcb3074f2b | ||
|
|
2145525a40 | ||
|
|
81bc6b83f8 | ||
|
|
86a24a182b | ||
|
|
08618ff8e7 | ||
|
|
a1de614ba3 | ||
|
|
965f89794f | ||
|
|
d834d44e64 | ||
|
|
9f70b2cecd | ||
|
|
4e7481175c | ||
|
|
171e8846b4 | ||
|
|
4b1a27fa0e | ||
|
|
fcc891545b | ||
|
|
a25c9865fe | ||
|
|
e85e15cf6d | ||
|
|
b248f4a3c1 | ||
|
|
d81aef1994 | ||
|
|
27b20ba8b1 | ||
|
|
e351231c4f | ||
|
|
5a75f14c0f | ||
|
|
e9f824d8c0 | ||
|
|
d028c697b5 | ||
|
|
66963a8bc7 | ||
|
|
cd74ef6274 | ||
|
|
f9af9be219 | ||
|
|
1ab7e5ad2d | ||
|
|
f805c57a2d | ||
|
|
4de0926596 | ||
|
|
ed319febb1 | ||
|
|
84e76d8a23 | ||
|
|
cdc06426e7 | ||
|
|
bced4595b8 | ||
|
|
a02c7f58c1 | ||
|
|
5cf3a35287 | ||
|
|
07fc586e38 |
@@ -1,4 +1,4 @@
|
||||
ARG ONEAPI_VERSION=2025.3.3-0-devel-ubuntu24.04
|
||||
ARG ONEAPI_VERSION=2026.1.1-devel-ubuntu24.04
|
||||
ARG BUILD_DATE=N/A
|
||||
ARG APP_VERSION=N/A
|
||||
ARG APP_REVISION=N/A
|
||||
@@ -19,7 +19,7 @@ RUN npm ci
|
||||
COPY tools/ui/ ./
|
||||
RUN LLAMA_BUILD_NUMBER="$APP_VERSION" npm run build
|
||||
|
||||
FROM docker.io/intel/deep-learning-essentials:$ONEAPI_VERSION AS build
|
||||
FROM docker.io/intel/oneapi-toolkit:$ONEAPI_VERSION AS build
|
||||
|
||||
ARG GGML_SYCL_F16=ON
|
||||
ARG LEVEL_ZERO_VERSION=1.28.2
|
||||
@@ -59,7 +59,7 @@ RUN mkdir -p /app/full \
|
||||
&& cp requirements.txt /app/full \
|
||||
&& cp .devops/tools.sh /app/full/tools.sh
|
||||
|
||||
FROM docker.io/intel/deep-learning-essentials:$ONEAPI_VERSION AS base
|
||||
FROM docker.io/intel/oneapi-toolkit:$ONEAPI_VERSION AS base
|
||||
|
||||
ARG BUILD_DATE=N/A
|
||||
ARG APP_VERSION=N/A
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
ARG UBUNTU_VERSION=22.04
|
||||
# This needs to generally match the container host's environment.
|
||||
ARG MUSA_VERSION=rc4.3.0
|
||||
# Target the MUSA build image
|
||||
ARG BASE_MUSA_DEV_CONTAINER=docker.io/mthreads/musa:${MUSA_VERSION}-devel-ubuntu${UBUNTU_VERSION}-amd64
|
||||
ARG BASE_MUSA_DEV_CONTAINER=registry.mthreads.com/mcconline/inference/pytorch:2.9.1.post1-py3.10-musa5.2.0-mp31-devel-ubuntu${UBUNTU_VERSION}-amd64
|
||||
|
||||
ARG BASE_MUSA_RUN_CONTAINER=docker.io/mthreads/musa:${MUSA_VERSION}-runtime-ubuntu${UBUNTU_VERSION}-amd64
|
||||
ARG BASE_MUSA_RUN_CONTAINER=${BASE_MUSA_DEV_CONTAINER}
|
||||
|
||||
ARG BUILD_DATE=N/A
|
||||
ARG APP_VERSION=N/A
|
||||
|
||||
@@ -33,6 +33,7 @@ concurrency:
|
||||
env:
|
||||
GGML_NLOOP: 3
|
||||
GGML_N_THREADS: 1
|
||||
GGML_SCHED_DEBUG_REALLOC: 1
|
||||
LLAMA_ARG_LOG_COLORS: 1
|
||||
LLAMA_ARG_LOG_PREFIX: 1
|
||||
LLAMA_ARG_LOG_TIMESTAMPS: 1
|
||||
@@ -98,7 +99,8 @@ jobs:
|
||||
id: cmake_test
|
||||
run: |
|
||||
cd build
|
||||
ctest -L main -E "test-llama-archs" --verbose --timeout 900
|
||||
# ref: https://github.com/ggml-org/llama.cpp/pull/19802#issuecomment-4013704023
|
||||
ctest -L main -E "test-llama-archs|test-save-load-state" --verbose --timeout 900
|
||||
|
||||
macos-latest-x64:
|
||||
runs-on: macos-15-intel
|
||||
|
||||
@@ -37,6 +37,7 @@ concurrency:
|
||||
env:
|
||||
GGML_NLOOP: 3
|
||||
GGML_N_THREADS: 1
|
||||
GGML_SCHED_DEBUG_REALLOC: 1
|
||||
LLAMA_ARG_LOG_COLORS: 1
|
||||
LLAMA_ARG_LOG_PREFIX: 1
|
||||
LLAMA_ARG_LOG_TIMESTAMPS: 1
|
||||
|
||||
@@ -145,7 +145,7 @@ jobs:
|
||||
|
||||
musa:
|
||||
runs-on: ubuntu-22.04
|
||||
container: mthreads/musa:rc4.3.0-devel-ubuntu22.04-amd64
|
||||
container: registry.mthreads.com/mcconline/inference/pytorch:2.9.1.post1-py3.10-musa5.2.0-mp31-devel-ubuntu22.04-amd64
|
||||
|
||||
steps:
|
||||
- name: Clone
|
||||
@@ -156,7 +156,7 @@ jobs:
|
||||
id: depends
|
||||
run: |
|
||||
apt-get update
|
||||
apt-get install -y build-essential git cmake libssl-dev jq
|
||||
apt-get install -y build-essential git cmake libssl-dev jq python3-venv
|
||||
|
||||
- name: ccache
|
||||
uses: ggml-org/[email protected]
|
||||
@@ -178,8 +178,8 @@ jobs:
|
||||
run: |
|
||||
cmake -B build -S . \
|
||||
-DGGML_MUSA=ON \
|
||||
-DMUSA_ARCHITECTURES=21
|
||||
time cmake --build build --config Release -j $(nproc)
|
||||
-DMUSA_ARCHITECTURES=31
|
||||
cmake --build build --config Release -j $(nproc)
|
||||
|
||||
- name: ccache-buckets-save
|
||||
if: ${{ github.event_name == 'push' && github.ref == 'refs/heads/master' }}
|
||||
|
||||
@@ -48,7 +48,7 @@ jobs:
|
||||
|
||||
env:
|
||||
ONEAPI_ROOT: /opt/intel/oneapi/
|
||||
ONEAPI_INSTALLER_VERSION: "2025.3.3"
|
||||
ONEAPI_INSTALLER_VERSION: "2026.1"
|
||||
LEVEL_ZERO_VERSION: "1.33.1"
|
||||
LEVEL_ZERO_UBUNTU_VERSION: "u24.04"
|
||||
|
||||
@@ -63,8 +63,8 @@ jobs:
|
||||
shell: bash
|
||||
run: |
|
||||
cd /tmp
|
||||
wget https://registrationcenter-download.intel.com/akdlm/IRC_NAS/56f7923a-adb8-43f3-8b02-2b60fcac8cab/intel-deep-learning-essentials-2025.3.3.16_offline.sh -O intel-deep-learning-essentials_offline.sh
|
||||
sudo bash intel-deep-learning-essentials_offline.sh -s -a --silent --eula accept
|
||||
wget https://registrationcenter-download.intel.com/akdlm/IRC_NAS/5996e26b-f48a-42b1-8db0-b002ad0bd8d7/intel-oneapi-toolkit-2026.1.1.33_offline.sh -O intel-oneapi-toolkit_offline.sh
|
||||
sudo bash intel-oneapi-toolkit_offline.sh -s -a --silent --eula accept
|
||||
|
||||
- name: Install Level Zero SDK
|
||||
shell: bash
|
||||
@@ -129,11 +129,11 @@ jobs:
|
||||
shell: bash
|
||||
|
||||
env:
|
||||
WINDOWS_BASEKIT_URL: https://registrationcenter-download.intel.com/akdlm/IRC_NAS/b60765d1-2b85-4e85-86b6-cb0e9563a699/intel-deep-learning-essentials-2025.3.3.18_offline.exe
|
||||
WINDOWS_BASEKIT_URL: https://registrationcenter-download.intel.com/akdlm/IRC_NAS/0cb67a0d-67f6-410b-868b-f4a0a17ff0cf/intel-oneapi-toolkit-2026.1.1.32_offline.exe
|
||||
WINDOWS_DPCPP_MKL: intel.oneapi.win.cpp-dpcpp-common:intel.oneapi.win.mkl.devel:intel.oneapi.win.dnnl:intel.oneapi.win.tbb.devel
|
||||
LEVEL_ZERO_SDK_URL: https://github.com/oneapi-src/level-zero/releases/download/v1.33.1/level-zero-win-sdk-1.33.1.zip
|
||||
ONEAPI_ROOT: "C:/Program Files (x86)/Intel/oneAPI"
|
||||
ONEAPI_INSTALLER_VERSION: "2025.3.3"
|
||||
ONEAPI_INSTALLER_VERSION: "2026.1"
|
||||
steps:
|
||||
- name: Clone
|
||||
id: checkout
|
||||
|
||||
@@ -31,6 +31,7 @@ concurrency:
|
||||
env:
|
||||
GGML_NLOOP: 3
|
||||
GGML_N_THREADS: 1
|
||||
GGML_SCHED_DEBUG_REALLOC: 1
|
||||
LLAMA_ARG_LOG_COLORS: 1
|
||||
LLAMA_ARG_LOG_PREFIX: 1
|
||||
LLAMA_ARG_LOG_TIMESTAMPS: 1
|
||||
|
||||
@@ -31,7 +31,7 @@ jobs:
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: "3.11"
|
||||
pip-install: -r requirements/requirements-all.txt ty==0.0.78
|
||||
pip-install: -r requirements/requirements-all.txt ty==0.0.84
|
||||
# - name: Type-check with Pyright
|
||||
# uses: jakebailey/pyright-action@v2
|
||||
# with:
|
||||
|
||||
@@ -1344,11 +1344,11 @@ jobs:
|
||||
shell: bash
|
||||
|
||||
env:
|
||||
WINDOWS_BASEKIT_URL: https://registrationcenter-download.intel.com/akdlm/IRC_NAS/b60765d1-2b85-4e85-86b6-cb0e9563a699/intel-deep-learning-essentials-2025.3.3.18_offline.exe
|
||||
WINDOWS_BASEKIT_URL: https://registrationcenter-download.intel.com/akdlm/IRC_NAS/0cb67a0d-67f6-410b-868b-f4a0a17ff0cf/intel-oneapi-toolkit-2026.1.1.32_offline.exe
|
||||
WINDOWS_DPCPP_MKL: intel.oneapi.win.cpp-dpcpp-common:intel.oneapi.win.mkl.devel:intel.oneapi.win.dnnl:intel.oneapi.win.tbb.devel
|
||||
LEVEL_ZERO_SDK_URL: https://github.com/oneapi-src/level-zero/releases/download/v1.28.2/level-zero-win-sdk-1.28.2.zip
|
||||
LEVEL_ZERO_SDK_URL: https://github.com/oneapi-src/level-zero/releases/download/v1.33.1/level-zero-win-sdk-1.33.1.zip
|
||||
ONEAPI_ROOT: "C:/Program Files (x86)/Intel/oneAPI"
|
||||
ONEAPI_INSTALLER_VERSION: "2025.3.3"
|
||||
ONEAPI_INSTALLER_VERSION: "2026.1"
|
||||
|
||||
steps:
|
||||
- name: Clone
|
||||
@@ -1391,9 +1391,11 @@ jobs:
|
||||
run: |
|
||||
echo "cp oneAPI running time dll files in ${{ env.ONEAPI_ROOT }} to ./build/bin"
|
||||
|
||||
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_sycl_blas.5.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_core.2.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_tbb_thread.2.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_sycl_blas.6.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_core.3.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_def.3.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_avx2.3.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/mkl/latest/bin/mkl_tbb_thread.3.dll" ./build/bin
|
||||
|
||||
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/ur_adapter_level_zero.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/ur_adapter_level_zero_v2.dll" ./build/bin
|
||||
@@ -1408,13 +1410,11 @@ jobs:
|
||||
echo "Level Zero loader DLL not found in oneAPI or SDK; relying on system driver/runtime"
|
||||
fi
|
||||
|
||||
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/sycl8.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/sycl9.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/svml_dispmd.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/libmmd.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/libiomp5md.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/sycl-ls.exe" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/libsycl-fallback-bfloat16.spv" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/compiler/latest/bin/libsycl-native-bfloat16.spv" ./build/bin
|
||||
|
||||
cp "${{ env.ONEAPI_ROOT }}/dnnl/latest/bin/dnnl.dll" ./build/bin
|
||||
cp "${{ env.ONEAPI_ROOT }}/tbb/latest/bin/tbb12.dll" ./build/bin
|
||||
@@ -1454,8 +1454,8 @@ jobs:
|
||||
|
||||
env:
|
||||
ONEAPI_ROOT: /opt/intel/oneapi/
|
||||
ONEAPI_INSTALLER_VERSION: "2025.3.3"
|
||||
LEVEL_ZERO_VERSION: "1.28.2"
|
||||
ONEAPI_INSTALLER_VERSION: "2026.1"
|
||||
LEVEL_ZERO_VERSION: "1.33.1"
|
||||
LEVEL_ZERO_UBUNTU_VERSION: "u24.04"
|
||||
|
||||
steps:
|
||||
@@ -1469,16 +1469,16 @@ jobs:
|
||||
shell: bash
|
||||
run: |
|
||||
cd /tmp
|
||||
wget https://registrationcenter-download.intel.com/akdlm/IRC_NAS/56f7923a-adb8-43f3-8b02-2b60fcac8cab/intel-deep-learning-essentials-2025.3.3.16_offline.sh -O intel-deep-learning-essentials_offline.sh
|
||||
sudo bash intel-deep-learning-essentials_offline.sh -s -a --silent --eula accept
|
||||
wget https://registrationcenter-download.intel.com/akdlm/IRC_NAS/5996e26b-f48a-42b1-8db0-b002ad0bd8d7/intel-oneapi-toolkit-2026.1.1.33_offline.sh -O intel-oneapi-toolkit_offline.sh
|
||||
sudo bash intel-oneapi-toolkit_offline.sh -s -a --silent --eula accept
|
||||
|
||||
- name: Install Level Zero SDK
|
||||
shell: bash
|
||||
run: |
|
||||
cd /tmp
|
||||
wget -q "https://github.com/oneapi-src/level-zero/releases/download/v${LEVEL_ZERO_VERSION}/level-zero_${LEVEL_ZERO_VERSION}%2B${LEVEL_ZERO_UBUNTU_VERSION}_amd64.deb" -O level-zero.deb
|
||||
wget -q "https://github.com/oneapi-src/level-zero/releases/download/v${LEVEL_ZERO_VERSION}/level-zero-devel_${LEVEL_ZERO_VERSION}%2B${LEVEL_ZERO_UBUNTU_VERSION}_amd64.deb" -O level-zero-devel.deb
|
||||
sudo apt-get install -y ./level-zero.deb ./level-zero-devel.deb
|
||||
wget -q "https://github.com/oneapi-src/level-zero/releases/download/v${LEVEL_ZERO_VERSION}/libze1_${LEVEL_ZERO_VERSION}%2B${LEVEL_ZERO_UBUNTU_VERSION}_amd64.deb" -O libze1.deb
|
||||
wget -q "https://github.com/oneapi-src/level-zero/releases/download/v${LEVEL_ZERO_VERSION}/libze-dev_${LEVEL_ZERO_VERSION}%2B${LEVEL_ZERO_UBUNTU_VERSION}_amd64.deb" -O libze-dev.deb
|
||||
sudo apt-get install -y ./libze1.deb ./libze-dev.deb
|
||||
|
||||
- name: Download UI build
|
||||
uses: actions/download-artifact@v7
|
||||
|
||||
@@ -7,6 +7,7 @@ General:
|
||||
- Don't try to build or run the code unless you are explicitly asked to do so
|
||||
- Use the `gh` CLI tool when querying PRs, issues, or other GitHub resources
|
||||
- When [MODEL] is needed, first try to get it from the `PI_MODEL_NAME` env var before asking the user
|
||||
- Never read the `AGENTS.md` file
|
||||
|
||||
Coding:
|
||||
- When in doubt, always refer to the CONTRIBUTING.md file of the project
|
||||
|
||||
+1
-1
@@ -57,7 +57,7 @@
|
||||
/ggml/src/ggml-cann/ @ggml-org/ggml-cann
|
||||
/ggml/src/ggml-common.h @ggerganov
|
||||
/ggml/src/ggml-cpu/ @ggerganov
|
||||
/ggml/src/ggml-cpu/iqp.* @bartowski1182
|
||||
/ggml/src/ggml-cpu/tiled/ @jbooth @bartowski1182
|
||||
/ggml/src/ggml-cpu/spacemit/ @alex-spacemit
|
||||
/ggml/src/ggml-cuda/ @ggml-org/ggml-cuda
|
||||
/ggml/src/ggml-cuda/vendors/hip.h @IMbackK
|
||||
|
||||
+1
-1
@@ -21,7 +21,7 @@ docker run --privileged -it \
|
||||
-v $HOME/llama.cpp/ci-cache:/ci-cache \
|
||||
-v $HOME/llama.cpp/ci-results:/ci-results \
|
||||
-v $PWD:/ws -w /ws \
|
||||
mthreads/musa:rc4.3.0-devel-ubuntu22.04-amd64
|
||||
registry.mthreads.com/mcconline/inference/pytorch:2.9.1.post1-py3.10-musa5.2.0-mp31-devel-ubuntu22.04-amd64
|
||||
```
|
||||
|
||||
Inside the container, execute the following commands:
|
||||
|
||||
@@ -158,8 +158,8 @@ if [ ! -z ${GG_BUILD_WEBGPU} ]; then
|
||||
fi
|
||||
|
||||
if [ ! -z ${GG_BUILD_MUSA} ]; then
|
||||
# Use qy1 by default (MTT S80)
|
||||
MUSA_ARCH=${MUSA_ARCH:-21}
|
||||
# Use ph1 by default (MTT S5000)
|
||||
MUSA_ARCH=${MUSA_ARCH:-31}
|
||||
CMAKE_EXTRA="${CMAKE_EXTRA} -DGGML_MUSA=ON -DMUSA_ARCHITECTURES=${MUSA_ARCH}"
|
||||
fi
|
||||
|
||||
|
||||
+11
-10
@@ -351,7 +351,7 @@ static bool parse_bool_value(const std::string & value) {
|
||||
static std::string get_default_local_path(const std::string & url) {
|
||||
auto f = string_split<std::string>(url, '#').front();
|
||||
f = string_split<std::string>(f, '?').front();
|
||||
return fs_get_cache_file(string_split<std::string>(f, '/').back());
|
||||
return fs_path_to_utf8(fs_get_cache_file(string_split<std::string>(f, '/').back()));
|
||||
}
|
||||
|
||||
static bool spec_types_is_default(const common_params & params) {
|
||||
@@ -2673,16 +2673,17 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
|
||||
params.video_ffmpeg_bin_dir = value;
|
||||
}
|
||||
).set_examples(mmproj_examples).set_env("LLAMA_ARG_VIDEO_FFMPEG_DIR"));
|
||||
if (params.is_gen_docs || llama_supports_rpc()) {
|
||||
add_opt(common_arg(
|
||||
{"--rpc"}, "SERVERS",
|
||||
"comma-separated list of RPC servers (host:port)",
|
||||
[](common_params & params, const std::string & value) {
|
||||
add_rpc_devices(value);
|
||||
GGML_UNUSED(params);
|
||||
add_opt(common_arg(
|
||||
{"--rpc"}, "SERVERS",
|
||||
"comma-separated list of RPC servers (host:port)",
|
||||
[](common_params & params, const std::string & value) {
|
||||
if (!llama_supports_rpc()) {
|
||||
throw std::invalid_argument("RPC not supported in this build");
|
||||
}
|
||||
).set_env("LLAMA_ARG_RPC"));
|
||||
}
|
||||
add_rpc_devices(value);
|
||||
GGML_UNUSED(params);
|
||||
}
|
||||
).set_env("LLAMA_ARG_RPC"));
|
||||
add_opt(common_arg(
|
||||
{"-lm", "--load-mode"}, "MODE",
|
||||
"model loading mode (default: auto)\n"
|
||||
|
||||
+206
-179
@@ -46,11 +46,10 @@
|
||||
#include <io.h>
|
||||
#else
|
||||
#include <sys/ioctl.h>
|
||||
#include <sys/stat.h>
|
||||
#include <unistd.h>
|
||||
#endif
|
||||
|
||||
#if defined(__linux__)
|
||||
#if !defined(_WIN32) && !defined(__APPLE__)
|
||||
#include <sys/types.h>
|
||||
#include <pwd.h>
|
||||
#endif
|
||||
@@ -614,34 +613,6 @@ std::string string_from(const struct llama_context * ctx, const std::vector<llam
|
||||
return buf.str();
|
||||
}
|
||||
|
||||
std::string string_from(const struct llama_context * ctx, const struct llama_batch & batch) {
|
||||
std::stringstream buf;
|
||||
|
||||
buf << "[ ";
|
||||
|
||||
bool first = true;
|
||||
for (int i = 0; i < batch.n_tokens; ++i) {
|
||||
if (!first) {
|
||||
buf << ", ";
|
||||
} else {
|
||||
first = false;
|
||||
}
|
||||
|
||||
auto detokenized = common_token_to_piece(ctx, batch.token[i]);
|
||||
|
||||
buf << "\n" << std::to_string(i)
|
||||
<< ", token '" << detokenized << "'"
|
||||
<< ", pos " << std::to_string(batch.pos[i])
|
||||
<< ", n_seq_id " << std::to_string(batch.n_seq_id[i])
|
||||
<< ", seq_id " << std::to_string(batch.seq_id[i][0])
|
||||
<< ", logits " << std::to_string(batch.logits[i]);
|
||||
}
|
||||
|
||||
buf << " ]";
|
||||
|
||||
return buf.str();
|
||||
}
|
||||
|
||||
void string_process_escapes(std::string & input) {
|
||||
std::size_t input_len = input.length();
|
||||
std::size_t output_idx = 0;
|
||||
@@ -900,7 +871,7 @@ bool fs_validate_filename(const std::string & filename, bool allow_subdirs) {
|
||||
|
||||
|
||||
#ifdef _WIN32
|
||||
static std::wstring utf8_to_wstring(const std::string & str) {
|
||||
std::wstring utf8_to_wstring(const std::string & str) {
|
||||
if (str.empty()) {
|
||||
return std::wstring();
|
||||
}
|
||||
@@ -916,82 +887,29 @@ static std::wstring utf8_to_wstring(const std::string & str) {
|
||||
|
||||
return wstr;
|
||||
}
|
||||
|
||||
std::string wstring_to_utf8(const std::wstring & str) {
|
||||
if (str.empty()) {
|
||||
return std::string();
|
||||
}
|
||||
|
||||
int size = WideCharToMultiByte(CP_UTF8, 0, str.c_str(), (int)str.size(), NULL, 0, NULL, NULL);
|
||||
|
||||
if (size <= 0) {
|
||||
return std::string();
|
||||
}
|
||||
|
||||
std::string utf8(size, 0);
|
||||
WideCharToMultiByte(CP_UTF8, 0, str.c_str(), (int)str.size(), &utf8[0], size, NULL, NULL);
|
||||
|
||||
return utf8;
|
||||
}
|
||||
#endif
|
||||
|
||||
// returns true if successful, false otherwise
|
||||
bool fs_create_directory_with_parents(const std::string & path) {
|
||||
#ifdef _WIN32
|
||||
std::wstring wpath = utf8_to_wstring(path);
|
||||
|
||||
// if the path already exists, check whether it's a directory
|
||||
const DWORD attributes = GetFileAttributesW(wpath.c_str());
|
||||
if ((attributes != INVALID_FILE_ATTRIBUTES) && (attributes & FILE_ATTRIBUTE_DIRECTORY)) {
|
||||
return true;
|
||||
}
|
||||
|
||||
size_t pos_slash = 0;
|
||||
|
||||
// process path from front to back, procedurally creating directories
|
||||
while ((pos_slash = path.find('\\', pos_slash)) != std::string::npos) {
|
||||
const std::wstring subpath = wpath.substr(0, pos_slash);
|
||||
|
||||
pos_slash += 1;
|
||||
|
||||
// skip the drive letter, in some systems it can return an access denied error
|
||||
if (subpath.length() == 2 && subpath[1] == ':') {
|
||||
continue;
|
||||
}
|
||||
|
||||
const bool success = CreateDirectoryW(subpath.c_str(), NULL);
|
||||
|
||||
if (!success) {
|
||||
const DWORD error = GetLastError();
|
||||
|
||||
// if the path already exists, ensure that it's a directory
|
||||
if (error == ERROR_ALREADY_EXISTS) {
|
||||
const DWORD attributes = GetFileAttributesW(subpath.c_str());
|
||||
if (attributes == INVALID_FILE_ATTRIBUTES || !(attributes & FILE_ATTRIBUTE_DIRECTORY)) {
|
||||
return false;
|
||||
}
|
||||
} else {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
#else
|
||||
// if the path already exists, check whether it's a directory
|
||||
struct stat info;
|
||||
if (stat(path.c_str(), &info) == 0) {
|
||||
return S_ISDIR(info.st_mode);
|
||||
}
|
||||
|
||||
size_t pos_slash = 1; // skip leading slashes for directory creation
|
||||
|
||||
// process path from front to back, procedurally creating directories
|
||||
while ((pos_slash = path.find('/', pos_slash)) != std::string::npos) {
|
||||
const std::string subpath = path.substr(0, pos_slash);
|
||||
struct stat info;
|
||||
|
||||
// if the path already exists, ensure that it's a directory
|
||||
if (stat(subpath.c_str(), &info) == 0) {
|
||||
if (!S_ISDIR(info.st_mode)) {
|
||||
return false;
|
||||
}
|
||||
} else {
|
||||
// create parent directories
|
||||
const int ret = mkdir(subpath.c_str(), 0755);
|
||||
if (ret != 0) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
pos_slash += 1;
|
||||
}
|
||||
|
||||
return true;
|
||||
#endif // _WIN32
|
||||
// returns the path as a UTF-8 string, preserving its separators
|
||||
std::string fs_path_to_utf8(const std::filesystem::path & path) {
|
||||
const auto value = path.u8string();
|
||||
return std::string(value.begin(), value.end());
|
||||
}
|
||||
|
||||
bool fs_is_directory(const std::string & path) {
|
||||
@@ -1016,58 +934,52 @@ void common_set_env(const std::string & name, const std::string & value) {
|
||||
#endif
|
||||
}
|
||||
|
||||
std::string fs_get_cache_directory() {
|
||||
std::string cache_directory = "";
|
||||
auto ensure_trailing_slash = [](std::string p) {
|
||||
// Make sure to add trailing slash
|
||||
if (p.empty() || p.back() != DIRECTORY_SEPARATOR) {
|
||||
p += DIRECTORY_SEPARATOR;
|
||||
}
|
||||
return p;
|
||||
};
|
||||
cache_directory = common_get_env("LLAMA_CACHE");
|
||||
std::filesystem::path common_get_path_from_env(const std::string & name) {
|
||||
#if defined(_WIN32)
|
||||
const std::wstring wname = utf8_to_wstring(name);
|
||||
const wchar_t * wvalue = _wgetenv(wname.c_str());
|
||||
return wvalue ? std::filesystem::path(wvalue) : std::filesystem::path();
|
||||
#else
|
||||
const char * value = std::getenv(name.c_str());
|
||||
return value ? std::filesystem::path(value) : std::filesystem::path();
|
||||
#endif
|
||||
}
|
||||
|
||||
std::filesystem::path fs_get_cache_directory() {
|
||||
std::filesystem::path cache_directory = common_get_path_from_env("LLAMA_CACHE");
|
||||
if (!cache_directory.empty()) {
|
||||
return cache_directory;
|
||||
}
|
||||
|
||||
#if defined(_WIN32)
|
||||
cache_directory = common_get_path_from_env("LOCALAPPDATA");
|
||||
if (cache_directory.empty()) {
|
||||
#if defined(__linux__) || defined(__FreeBSD__) || defined(_AIX) || \
|
||||
defined(__OpenBSD__) || defined(__NetBSD__)
|
||||
const std::string xdg_cache_home = common_get_env("XDG_CACHE_HOME");
|
||||
const std::string home = common_get_env("HOME");
|
||||
if (!xdg_cache_home.empty()) {
|
||||
cache_directory = xdg_cache_home;
|
||||
} else if (!home.empty()) {
|
||||
cache_directory = home + "/.cache/";
|
||||
throw std::runtime_error("Failed to find %LOCALAPPDATA% directory");
|
||||
}
|
||||
#elif defined(__APPLE__)
|
||||
cache_directory = common_get_path_from_env("HOME");
|
||||
if (cache_directory.empty()) {
|
||||
throw std::runtime_error("Failed to find $HOME directory");
|
||||
}
|
||||
cache_directory /= "Library/Caches";
|
||||
#else
|
||||
cache_directory = common_get_path_from_env("XDG_CACHE_HOME");
|
||||
if (cache_directory.empty()) {
|
||||
cache_directory = common_get_path_from_env("HOME");
|
||||
if (!cache_directory.empty()) {
|
||||
cache_directory /= ".cache";
|
||||
} else {
|
||||
#if defined(__linux__)
|
||||
/* no $HOME is defined, fallback to getpwuid */
|
||||
struct passwd *pw = getpwuid(getuid());
|
||||
if ((!pw) || (!pw->pw_dir)) {
|
||||
const struct passwd * pw = getpwuid(getuid());
|
||||
if (!pw || !pw->pw_dir || !*pw->pw_dir) {
|
||||
throw std::runtime_error("Failed to find $HOME directory");
|
||||
}
|
||||
|
||||
cache_directory = std::string(pw->pw_dir) + std::string("/.cache/");
|
||||
#else /* defined(__linux__) */
|
||||
throw std::runtime_error("Failed to find $HOME directory");
|
||||
#endif /* defined(__linux__) */
|
||||
cache_directory = pw->pw_dir;
|
||||
cache_directory /= ".cache";
|
||||
}
|
||||
#elif defined(__APPLE__)
|
||||
cache_directory = common_get_env("HOME");
|
||||
if (cache_directory.empty()) {
|
||||
throw std::runtime_error("Failed to find $HOME directory");
|
||||
}
|
||||
cache_directory += "/Library/Caches/";
|
||||
#elif defined(_WIN32)
|
||||
cache_directory = common_get_env("LOCALAPPDATA");
|
||||
if (cache_directory.empty()) {
|
||||
throw std::runtime_error("Failed to find %LOCALAPPDATA% directory");
|
||||
}
|
||||
#elif defined(__EMSCRIPTEN__)
|
||||
GGML_ABORT("not implemented on this platform");
|
||||
#else
|
||||
# error Unknown architecture
|
||||
#endif
|
||||
cache_directory = ensure_trailing_slash(cache_directory);
|
||||
cache_directory += "llama.cpp";
|
||||
}
|
||||
return ensure_trailing_slash(cache_directory);
|
||||
#endif
|
||||
return cache_directory / "llama.cpp";
|
||||
}
|
||||
|
||||
std::string fs_get_config_directory() {
|
||||
@@ -1115,14 +1027,15 @@ std::string fs_get_config_directory() {
|
||||
return ensure_trailing_slash(config_directory);
|
||||
}
|
||||
|
||||
std::string fs_get_cache_file(const std::string & filename) {
|
||||
std::filesystem::path fs_get_cache_file(const std::string & filename) {
|
||||
GGML_ASSERT(filename.find(DIRECTORY_SEPARATOR) == std::string::npos);
|
||||
std::string cache_directory = fs_get_cache_directory();
|
||||
const bool success = fs_create_directory_with_parents(cache_directory);
|
||||
if (!success) {
|
||||
throw std::runtime_error("failed to create cache directory: " + cache_directory);
|
||||
const std::filesystem::path cache_directory = fs_get_cache_directory();
|
||||
std::error_code ec;
|
||||
std::filesystem::create_directories(cache_directory, ec);
|
||||
if (ec) {
|
||||
throw std::runtime_error("failed to create cache directory: " + fs_path_to_utf8(cache_directory));
|
||||
}
|
||||
return cache_directory + filename;
|
||||
return cache_directory / std::filesystem::u8path(filename);
|
||||
}
|
||||
|
||||
std::vector<common_file_info> fs_list(const std::string & path, bool include_directories) {
|
||||
@@ -1527,7 +1440,8 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
|
||||
}
|
||||
|
||||
if (llama_model_has_encoder(model)) {
|
||||
llama_encode(lctx, llama_batch_get_one(tmp.data(), tmp.size()));
|
||||
common_batch batch = common_batch_get_one(lctx, tmp);
|
||||
llama_process(lctx, LLAMA_PROCESS_TYPE_ENCODE, batch.get());
|
||||
llama_token decoder_start_token_id = llama_model_decoder_start_token(model);
|
||||
if (decoder_start_token_id == LLAMA_TOKEN_NULL) {
|
||||
decoder_start_token_id = bos;
|
||||
@@ -1536,7 +1450,9 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
|
||||
tmp.push_back(decoder_start_token_id);
|
||||
}
|
||||
if (llama_model_has_decoder(model)) {
|
||||
llama_decode(lctx, llama_batch_get_one(tmp.data(), std::min(tmp.size(), (size_t) params.n_batch)));
|
||||
tmp.resize(std::min(tmp.size(), (size_t) params.n_batch));
|
||||
common_batch batch = common_batch_get_one(lctx, tmp);
|
||||
llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
|
||||
}
|
||||
llama_memory_clear(llama_get_memory(lctx), true);
|
||||
llama_synchronize(lctx);
|
||||
@@ -1600,9 +1516,13 @@ common_context_seq_rm_type common_context_can_seq_rm(llama_context * ctx) {
|
||||
tmp.push_back(0);
|
||||
tmp.push_back(0);
|
||||
|
||||
int ret = llama_decode(ctx, llama_batch_get_one(tmp.data(), tmp.size()));
|
||||
int ret;
|
||||
{
|
||||
common_batch batch = common_batch_get_one(ctx, tmp);
|
||||
ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
|
||||
}
|
||||
if (ret != 0) {
|
||||
COM_ERR("llama_decode() failed: %d\n", ret);
|
||||
COM_ERR("llama_process() failed: %d\n", ret);
|
||||
res = COMMON_CONTEXT_SEQ_RM_TYPE_NO;
|
||||
goto done;
|
||||
}
|
||||
@@ -2189,29 +2109,138 @@ float lr_opt::get_lr(float epoch) const {
|
||||
}
|
||||
|
||||
bool common_replay_last_token(struct llama_context * ctx, llama_token last_token, int32_t pos) {
|
||||
llama_batch batch = llama_batch_get_one(&last_token, 1);
|
||||
batch.pos = &pos;
|
||||
if (llama_decode(ctx, batch)) {
|
||||
common_batch batch(ctx);
|
||||
batch.add(last_token, pos, 0, true);
|
||||
|
||||
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
|
||||
LOG_ERR("%s: failed to replay last token\n", __func__);
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
llama_batch_ext_ptr common_batch_ext_get_one(llama_context * ctx, const llama_tokens & tokens) {
|
||||
llama_batch_ext_ptr batch(llama_batch_ext_init(ctx));
|
||||
common_batch::common_batch(llama_context * ctx) : batch(llama_batch_ext_init(ctx)) {
|
||||
const auto rope_type = llama_model_rope_type(llama_get_model(ctx));
|
||||
n_pos = rope_type == LLAMA_ROPE_TYPE_MROPE || rope_type == LLAMA_ROPE_TYPE_IMROPE ? GGML_MROPE_SECTIONS : 1;
|
||||
}
|
||||
|
||||
auto mem = llama_get_memory(ctx);
|
||||
llama_pos pos = mem ? llama_memory_seq_pos_max(mem, 0) + 1 : 0;
|
||||
void common_batch::clear() {
|
||||
tokens.clear();
|
||||
llama_batch_ext_clear(batch.get());
|
||||
}
|
||||
|
||||
for (size_t i = 0; i < tokens.size(); ++i) {
|
||||
const int32_t idx = llama_batch_ext_add_token(batch.get(), 0, tokens[i]);
|
||||
llama_batch_ext_set_pos(batch.get(), idx, &pos);
|
||||
pos++;
|
||||
int32_t common_batch::add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output) {
|
||||
const int32_t idx = llama_batch_ext_add_token(batch.get(), seq_id, id);
|
||||
if (idx < 0) {
|
||||
GGML_ABORT("%s: failed to add token %d to the batch (error %d, n_tokens = %d)\n", __func__, id, idx, size());
|
||||
}
|
||||
llama_batch_ext_set_pos(batch.get(), idx, &pos);
|
||||
if (output) {
|
||||
llama_batch_ext_set_output_logits(batch.get(), idx, true);
|
||||
}
|
||||
tokens.push_back({ id, { pos, 0, 0, 0 }, seq_id, output, { nullptr, 0, 0 } });
|
||||
return idx;
|
||||
}
|
||||
|
||||
bool common_batch::set_output(int32_t idx, bool value) {
|
||||
if (idx < 0 || idx >= (int32_t) tokens.size()) {
|
||||
return false;
|
||||
}
|
||||
tokens[idx].output = value;
|
||||
return llama_batch_ext_set_output_logits(batch.get(), idx, value);
|
||||
}
|
||||
|
||||
bool common_batch::set_embd(int32_t idx, llama_embd embd) {
|
||||
if (idx < 0 || idx >= (int32_t) tokens.size()) {
|
||||
return false;
|
||||
}
|
||||
if (!llama_batch_ext_set_embd_token(batch.get(), idx, embd)) {
|
||||
return false;
|
||||
}
|
||||
tokens[idx].embd = embd;
|
||||
return true;
|
||||
}
|
||||
|
||||
int32_t common_batch::add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output) {
|
||||
const int32_t idx = llama_batch_ext_add_embd(batch.get(), seq_id, embd);
|
||||
if (idx < 0) {
|
||||
GGML_ABORT("%s: failed to add embedding to the batch (error %d, n_tokens = %d)\n", __func__, idx, size());
|
||||
}
|
||||
llama_batch_ext_set_pos(batch.get(), idx, pos);
|
||||
if (output) {
|
||||
llama_batch_ext_set_output_logits(batch.get(), idx, true);
|
||||
}
|
||||
token t = { LLAMA_TOKEN_NULL, { 0, 0, 0, 0 }, seq_id, output, embd };
|
||||
for (int32_t j = 0; j < n_pos; ++j) {
|
||||
t.pos[j] = pos[j];
|
||||
}
|
||||
tokens.push_back(t);
|
||||
return idx;
|
||||
}
|
||||
|
||||
common_batch common_batch_from_llama_batch(llama_context * ctx, const llama_batch & batch) {
|
||||
common_batch res(ctx);
|
||||
|
||||
const bool has_token = batch.token != nullptr;
|
||||
const bool has_embd = batch.embd != nullptr;
|
||||
|
||||
const size_t n_embd = llama_model_n_embd_inp(llama_get_model(ctx));
|
||||
|
||||
// positions continue from the memory when none are given
|
||||
auto * mem = llama_get_memory(ctx);
|
||||
std::vector<llama_pos> pos_next(llama_n_seq_max(ctx));
|
||||
for (llama_seq_id s = 0; s < (llama_seq_id) pos_next.size(); ++s) {
|
||||
pos_next[s] = llama_memory_seq_pos_max(mem, s) + 1;
|
||||
}
|
||||
|
||||
if (!tokens.empty()) {
|
||||
llama_batch_ext_set_output_logits(batch.get(), (int32_t) tokens.size() - 1, true);
|
||||
for (int32_t i = 0; i < batch.n_tokens; ++i) {
|
||||
const int32_t n_sid = batch.n_seq_id ? batch.n_seq_id[i] : 1;
|
||||
const llama_seq_id seq_id = batch.seq_id ? batch.seq_id[i][0] : 0;
|
||||
|
||||
llama_pos pos[GGML_MROPE_SECTIONS] = { 0, 0, 0, 0 };
|
||||
if (!batch.pos) {
|
||||
pos[0] = pos_next[seq_id]++;
|
||||
} else if (has_token) {
|
||||
pos[0] = batch.pos[i];
|
||||
} else {
|
||||
// embedding batch: section-major layout pos[j*n_tokens + i]
|
||||
for (int32_t j = 0; j < res.n_pos; ++j) {
|
||||
pos[j] = batch.pos[j * batch.n_tokens + i];
|
||||
}
|
||||
}
|
||||
|
||||
const bool output = batch.logits ? batch.logits[i] != 0 : i == batch.n_tokens - 1;
|
||||
|
||||
const llama_embd embd = { has_embd ? batch.embd + (size_t) i * n_embd : nullptr, 1, n_embd };
|
||||
|
||||
int32_t idx;
|
||||
if (has_token) {
|
||||
idx = res.add(batch.token[i], pos[0], seq_id, output);
|
||||
if (has_embd) {
|
||||
res.set_embd(idx, embd);
|
||||
}
|
||||
} else {
|
||||
idx = res.add_embd(embd, pos, seq_id, output);
|
||||
}
|
||||
|
||||
for (int32_t s = 1; s < n_sid; ++s) {
|
||||
llama_batch_ext_add_seq(res.get(), idx, batch.seq_id[i][s]);
|
||||
}
|
||||
}
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
common_batch common_batch_get_one(llama_context * ctx, const llama_tokens & tokens) {
|
||||
common_batch batch(ctx);
|
||||
|
||||
auto mem = llama_get_memory(ctx);
|
||||
llama_pos pos = llama_memory_seq_pos_max(mem, 0) + 1; // -1 + 1 == 0 when the memory is empty
|
||||
|
||||
for (size_t i = 0; i < tokens.size(); ++i) {
|
||||
const bool output = i == tokens.size() - 1;
|
||||
batch.add(tokens[i], pos, 0, output);
|
||||
pos++;
|
||||
}
|
||||
|
||||
return batch;
|
||||
@@ -2241,7 +2270,7 @@ bool common_prompt_batch_decode(
|
||||
// memory, so we can't just remove the last token from the memory and replay the last token which
|
||||
// is the reason for this logic.
|
||||
llama_tokens prefix_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_tokens_before_last);
|
||||
llama_batch_ext_ptr batch_prefix = common_batch_ext_get_one(ctx, prefix_tokens);
|
||||
common_batch batch_prefix = common_batch_get_one(ctx, prefix_tokens);
|
||||
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_prefix.get())) {
|
||||
COM_ERR("%s", "failed to eval\n");
|
||||
return false;
|
||||
@@ -2251,10 +2280,8 @@ bool common_prompt_batch_decode(
|
||||
llama_state_save_file(ctx, state_path.data(), all_tokens.data(), all_tokens.size());
|
||||
COM_INF("saved session before last token to %s, n_new = %zu\n", state_path.data(), all_tokens.size());
|
||||
|
||||
llama_token last_token = all_tokens.back();
|
||||
llama_batch_ext_ptr batch_last = common_batch_ext_get_one(ctx, { last_token });
|
||||
llama_pos pos = n_past;
|
||||
llama_batch_ext_set_pos(batch_last.get(), 0, &pos);
|
||||
common_batch batch_last(ctx);
|
||||
batch_last.add(all_tokens.back(), n_past, 0, true);
|
||||
|
||||
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_last.get())) {
|
||||
COM_ERR("%s", "failed to eval last token\n");
|
||||
@@ -2263,7 +2290,7 @@ bool common_prompt_batch_decode(
|
||||
n_past++;
|
||||
} else {
|
||||
llama_tokens new_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_new);
|
||||
llama_batch_ext_ptr batch = common_batch_ext_get_one(ctx, new_tokens);
|
||||
common_batch batch = common_batch_get_one(ctx, new_tokens);
|
||||
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
|
||||
COM_ERR("%s", "failed to eval\n");
|
||||
return false;
|
||||
|
||||
+68
-6
@@ -8,6 +8,7 @@
|
||||
#include "ggml.h"
|
||||
#include "llama.h"
|
||||
|
||||
#include <array>
|
||||
#include <list>
|
||||
#include <set>
|
||||
#include <sstream>
|
||||
@@ -16,6 +17,7 @@
|
||||
#include <vector>
|
||||
#include <map>
|
||||
#include <algorithm>
|
||||
#include <filesystem>
|
||||
#include <fstream>
|
||||
|
||||
#if defined(_WIN32) && !defined(_WIN32_WINNT)
|
||||
@@ -808,7 +810,9 @@ static std::vector<T> string_split(const std::string & str, char delim) {
|
||||
while (std::getline(str_stream, token, delim)) {
|
||||
T value;
|
||||
std::istringstream token_stream(token);
|
||||
token_stream >> value;
|
||||
if (!(token_stream >> value)) {
|
||||
throw std::invalid_argument("invalid value: \"" + token + "\"");
|
||||
}
|
||||
values.push_back(value);
|
||||
}
|
||||
return values;
|
||||
@@ -876,10 +880,21 @@ void string_process_escapes(std::string & input);
|
||||
std::string string_from(bool value);
|
||||
std::string string_from(const std::vector<int> & values);
|
||||
std::string string_from(const struct llama_context * ctx, const std::vector<llama_token> & tokens);
|
||||
std::string string_from(const struct llama_context * ctx, const struct llama_batch & batch);
|
||||
|
||||
bool glob_match(const std::string & pattern, const std::string & str);
|
||||
|
||||
//
|
||||
// Unicode utils
|
||||
//
|
||||
|
||||
#ifdef _WIN32
|
||||
std::wstring utf8_to_wstring(const std::string & str);
|
||||
std::string wstring_to_utf8(const std::wstring & str);
|
||||
#endif
|
||||
|
||||
// returns the path as a UTF-8 string, preserving its separators
|
||||
std::string fs_path_to_utf8(const std::filesystem::path & path);
|
||||
|
||||
//
|
||||
// Environment utils
|
||||
//
|
||||
@@ -889,16 +904,18 @@ bool glob_match(const std::string & pattern, const std::string & str);
|
||||
std::string common_get_env(const std::string & name);
|
||||
void common_set_env(const std::string & name, const std::string & value);
|
||||
|
||||
// reads a path from the environment, an unset variable gives an empty path
|
||||
std::filesystem::path common_get_path_from_env(const std::string & name);
|
||||
|
||||
//
|
||||
// Filesystem utils
|
||||
//
|
||||
|
||||
bool fs_validate_filename(const std::string & filename, bool allow_subdirs = false);
|
||||
bool fs_create_directory_with_parents(const std::string & path);
|
||||
bool fs_is_directory(const std::string & path);
|
||||
|
||||
std::string fs_get_cache_directory();
|
||||
std::string fs_get_cache_file(const std::string & filename);
|
||||
std::filesystem::path fs_get_cache_directory();
|
||||
std::filesystem::path fs_get_cache_file(const std::string & filename);
|
||||
std::string fs_get_config_directory();
|
||||
|
||||
struct common_file_info {
|
||||
@@ -1021,9 +1038,54 @@ void common_batch_add(
|
||||
const std::vector<llama_seq_id> & seq_ids,
|
||||
bool logits);
|
||||
|
||||
// wrapper around llama_batch_ext that provide getter functions for downstream code
|
||||
struct common_batch {
|
||||
struct token {
|
||||
llama_token id;
|
||||
std::array<llama_pos, GGML_MROPE_SECTIONS> pos; // only pos[0] is used for text tokens
|
||||
llama_seq_id seq_id;
|
||||
bool output;
|
||||
llama_embd embd; // non-owning view of the data passed to add_embd()/set_embd(), data == NULL if none
|
||||
};
|
||||
|
||||
std::vector<token> tokens; // mirror of the entries, tokens[i] describes batch index i
|
||||
llama_batch_ext_ptr batch;
|
||||
|
||||
int32_t n_pos = 1; // positions per embedding entry, GGML_MROPE_SECTIONS for MROPE/IMROPE
|
||||
|
||||
common_batch() = default;
|
||||
common_batch(struct llama_context * ctx);
|
||||
|
||||
llama_batch_ext * get() const { return batch.get(); }
|
||||
|
||||
// content type of the batch, all entries carry the same combination
|
||||
bool has_token() const { return !tokens.empty() && tokens[0].id != LLAMA_TOKEN_NULL; }
|
||||
bool has_embd () const { return !tokens.empty() && tokens[0].embd.data != nullptr; }
|
||||
|
||||
void clear();
|
||||
|
||||
// returns the batch index (>= 0), aborts if the entry cannot be added (batch full, invalid token or seq id)
|
||||
int32_t add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output);
|
||||
|
||||
bool set_output(int32_t idx, bool value);
|
||||
|
||||
// attach a token embedding to the entry at idx, can only be set once per entry
|
||||
bool set_embd(int32_t idx, llama_embd embd);
|
||||
|
||||
// add an embedding-only entry (no token id), aborts like add() on failure
|
||||
// pos points to n_pos positions
|
||||
int32_t add_embd(llama_embd embd, const llama_pos * pos, llama_seq_id seq_id, bool output);
|
||||
|
||||
int32_t size() const { return (int32_t) tokens.size(); }
|
||||
};
|
||||
|
||||
// create a single-sequence batch from a list of tokens
|
||||
// last token always have output_logits set to true
|
||||
llama_batch_ext_ptr common_batch_ext_get_one(struct llama_context * ctx, const llama_tokens & tokens);
|
||||
common_batch common_batch_get_one(struct llama_context * ctx, const llama_tokens & tokens);
|
||||
|
||||
// convert a legacy llama_batch, applying its defaults: seq 0, positions continue from memory, last token is output
|
||||
// the embd rows are read at the model input width
|
||||
common_batch common_batch_from_llama_batch(struct llama_context * ctx, const llama_batch & batch);
|
||||
|
||||
// decodes a single batch of tokens for a prompt and manages session tokens
|
||||
//
|
||||
|
||||
+2
-3
@@ -1,4 +1,5 @@
|
||||
#include "console.h"
|
||||
#include "common.h"
|
||||
#include "log.h"
|
||||
#include <vector>
|
||||
#include <iostream>
|
||||
@@ -1053,9 +1054,7 @@ namespace console {
|
||||
return false;
|
||||
}
|
||||
|
||||
int size_needed = WideCharToMultiByte(CP_UTF8, 0, &wline[0], (int)wline.size(), NULL, 0, NULL, NULL);
|
||||
line.resize(size_needed);
|
||||
WideCharToMultiByte(CP_UTF8, 0, &wline[0], (int)wline.size(), &line[0], size_needed, NULL, NULL);
|
||||
line = wstring_to_utf8(wline);
|
||||
#else
|
||||
if (!std::getline(std::cin, line)) {
|
||||
// Input stream is bad or EOF received
|
||||
|
||||
+3
-3
@@ -286,7 +286,7 @@ static int common_download_file_single_online(const std::string & url,
|
||||
static const int max_attempts = 3;
|
||||
static const int retry_delay_seconds = 2;
|
||||
|
||||
const bool file_exists = std::filesystem::exists(path);
|
||||
const bool file_exists = std::filesystem::exists(std::filesystem::u8path(path));
|
||||
|
||||
if (file_exists && skip_etag) {
|
||||
LOG_DBG("%s: using cached file: %s\n", __func__, path.c_str());
|
||||
@@ -477,7 +477,7 @@ int common_download_file_single(const std::string & url,
|
||||
return common_download_file_single_online(url, path, online_opts, skip_etag);
|
||||
}
|
||||
|
||||
if (!std::filesystem::exists(path)) {
|
||||
if (!std::filesystem::exists(std::filesystem::u8path(path))) {
|
||||
LOG_ERR("%s: required file is not available in cache (offline mode): %s\n", __func__, path.c_str());
|
||||
return -1;
|
||||
}
|
||||
@@ -943,7 +943,7 @@ std::string common_docker_resolve_model(const std::string & docker) {
|
||||
std::string model_filename = repo;
|
||||
std::replace(model_filename.begin(), model_filename.end(), '/', '_');
|
||||
model_filename += "_" + tag + ".gguf";
|
||||
std::string local_path = fs_get_cache_file(model_filename);
|
||||
std::string local_path = fs_path_to_utf8(fs_get_cache_file(model_filename));
|
||||
|
||||
const std::string blob_url = url_prefix + "/blobs/" + gguf_digest;
|
||||
common_download_opts opts;
|
||||
|
||||
+7
-7
@@ -192,9 +192,9 @@ static void common_params_fit_impl(
|
||||
uint32_t hp_nct = 0; // hparams.n_ctx_train
|
||||
uint32_t hp_nex = 0; // hparams.n_expert
|
||||
|
||||
// size the context for all sequences, but keep minimums and alignment per KV stream
|
||||
const uint32_t n_seq_max = std::max<uint32_t>(1, cparams->n_seq_max);
|
||||
const uint32_t n_streams = cparams->kv_unified ? 1 : n_seq_max;
|
||||
// with non-unified kv, we need to take into account n_streams
|
||||
// for example, if memory can hold more than model's trained context size, we must extend the n_ctx to hold enough n_streams
|
||||
const uint32_t n_streams = cparams->kv_unified ? 1 : std::max<uint32_t>(1, cparams->n_seq_max);
|
||||
const bool n_ctx_auto = cparams->n_ctx == 0;
|
||||
|
||||
dmds_t dmds_extra; // memory of the extra model, laid out on the devices of the main model
|
||||
@@ -264,15 +264,15 @@ static void common_params_fit_impl(
|
||||
dmds_t dmds_full = common_get_device_memory_data_impl(path_model, mparams, cparams, devs, hp_ngl, hp_nct, hp_nex, log_level);
|
||||
|
||||
// saturate instead of overflowing, this also preserves the UINT32_MAX sentinel of n_ctx_min:
|
||||
const uint32_t n_ctx_max = (uint32_t) std::min<uint64_t>(uint64_t(hp_nct) * n_seq_max, UINT32_MAX);
|
||||
const uint32_t n_ctx_max = (uint32_t) std::min<uint64_t>(uint64_t(hp_nct) * n_streams, UINT32_MAX);
|
||||
const uint32_t n_ctx_min_total = (uint32_t) std::min<uint64_t>(uint64_t(n_ctx_min) * n_streams, UINT32_MAX);
|
||||
|
||||
// llama_context would use only hp_nct in total for n_ctx == 0, resolve the context before measuring anything else:
|
||||
if (n_ctx_auto) {
|
||||
cparams->n_ctx = n_ctx_max;
|
||||
if (n_seq_max > 1) {
|
||||
LOG_TRC("%s: context size unset -> using %" PRIu32 " for %" PRIu32 " sequences:\n",
|
||||
__func__, n_ctx_max, n_seq_max);
|
||||
if (n_streams > 1) {
|
||||
LOG_TRC("%s: context size unset and KV cache not unified -> using %" PRIu32 " for %" PRIu32 " sequences:\n",
|
||||
__func__, n_ctx_max, n_streams);
|
||||
dmds_full = common_get_device_memory_data_impl(path_model, mparams, cparams, devs, hp_ngl, hp_nct, hp_nex, log_level);
|
||||
}
|
||||
}
|
||||
|
||||
+28
-25
@@ -44,8 +44,7 @@ static fs::path get_cache_directory() {
|
||||
{HOME_DIR, fs::path(".cache") / "huggingface" / "hub"}
|
||||
};
|
||||
for (const auto & entry : entries) {
|
||||
if (auto * p = std::getenv(entry.var); p && *p) {
|
||||
fs::path base(p);
|
||||
if (fs::path base = common_get_path_from_env(entry.var); !base.empty()) {
|
||||
return entry.path.empty() ? base : base / entry.path;
|
||||
}
|
||||
}
|
||||
@@ -63,12 +62,7 @@ static fs::path get_cache_directory() {
|
||||
}
|
||||
|
||||
std::string get_cache_path() {
|
||||
#if defined(__cpp_lib_char8_t)
|
||||
const std::u8string u8str = get_cache_directory().u8string();
|
||||
return std::string(reinterpret_cast<const char *>(u8str.data()), u8str.size());
|
||||
#else
|
||||
return get_cache_directory().u8string();
|
||||
#endif
|
||||
return fs_path_to_utf8(get_cache_directory());
|
||||
}
|
||||
|
||||
static std::string folder_name_to_repo(const std::string & folder) {
|
||||
@@ -179,7 +173,8 @@ static bool is_valid_subpath(const fs::path & path, const fs::path & subpath) {
|
||||
}
|
||||
|
||||
static void safe_write_file(const fs::path & path, const std::string & data) {
|
||||
fs::path path_tmp = path.string() + ".tmp";
|
||||
fs::path path_tmp = path;
|
||||
path_tmp += ".tmp";
|
||||
|
||||
if (path.has_parent_path()) {
|
||||
fs::create_directories(path.parent_path());
|
||||
@@ -196,7 +191,7 @@ static void safe_write_file(const fs::path & path, const std::string & data) {
|
||||
}
|
||||
if (file.fail() || ec) {
|
||||
fs::remove(path_tmp, ec);
|
||||
throw std::runtime_error("failed to write file: " + path.string());
|
||||
throw std::runtime_error("failed to write file: " + fs_path_to_utf8(path));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -246,6 +241,7 @@ static std::string get_repo_commit(const std::string & repo_id,
|
||||
fs::path refs_path = get_repo_path(repo_id) / "refs";
|
||||
std::string name;
|
||||
std::string commit;
|
||||
fs::path name_path;
|
||||
|
||||
for (const auto & branch : json["branches"]) {
|
||||
if (!branch.is_object() ||
|
||||
@@ -256,24 +252,28 @@ static std::string get_repo_commit(const std::string & repo_id,
|
||||
std::string _name = branch["name"].get<std::string>();
|
||||
std::string _commit = branch["targetCommit"].get<std::string>();
|
||||
|
||||
if (!is_valid_subpath(refs_path, _name)) {
|
||||
LOG_WRN("%s: skip invalid branch: %s\n", __func__, _name.c_str());
|
||||
continue;
|
||||
}
|
||||
if (!is_valid_commit(_commit)) {
|
||||
LOG_WRN("%s: skip invalid commit: %s\n", __func__, _commit.c_str());
|
||||
continue;
|
||||
}
|
||||
const fs::path candidate = fs::u8path(_name);
|
||||
|
||||
if (!is_valid_subpath(refs_path, candidate)) {
|
||||
LOG_WRN("%s: skip invalid branch: %s\n", __func__, _name.c_str());
|
||||
continue;
|
||||
}
|
||||
|
||||
if (_name == "main") {
|
||||
name = _name;
|
||||
commit = _commit;
|
||||
name_path = candidate;
|
||||
break;
|
||||
}
|
||||
|
||||
if (name.empty() || commit.empty()) {
|
||||
name = _name;
|
||||
commit = _commit;
|
||||
name_path = candidate;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -282,7 +282,7 @@ static std::string get_repo_commit(const std::string & repo_id,
|
||||
return {};
|
||||
}
|
||||
|
||||
safe_write_file(refs_path / name, commit);
|
||||
safe_write_file(refs_path / name_path, commit);
|
||||
return commit;
|
||||
|
||||
} catch (const common_json_error & e) {
|
||||
@@ -331,7 +331,9 @@ hf_files get_repo_files(const std::string & repo_id,
|
||||
file.repo_id = repo_id;
|
||||
file.path = item["path"].get<std::string>();
|
||||
|
||||
if (!is_valid_subpath(commit_path, file.path)) {
|
||||
const fs::path subpath = fs::u8path(file.path);
|
||||
|
||||
if (!is_valid_subpath(commit_path, subpath)) {
|
||||
LOG_WRN("%s: skip invalid path: %s\n", __func__, file.path.c_str());
|
||||
continue;
|
||||
}
|
||||
@@ -351,12 +353,12 @@ hf_files get_repo_files(const std::string & repo_id,
|
||||
|
||||
file.url = endpoint + repo_id + "/resolve/" + commit + "/" + file.path;
|
||||
|
||||
fs::path final_path = commit_path / file.path;
|
||||
file.final_path = final_path.string();
|
||||
fs::path final_path = commit_path / subpath;
|
||||
file.final_path = fs_path_to_utf8(final_path);
|
||||
|
||||
if (!file.oid.empty() && !fs::exists(final_path)) {
|
||||
fs::path local_path = blobs_path / file.oid;
|
||||
file.local_path = local_path.string();
|
||||
file.local_path = fs_path_to_utf8(local_path);
|
||||
} else {
|
||||
file.local_path = file.final_path;
|
||||
}
|
||||
@@ -423,7 +425,7 @@ hf_files get_cached_files(const std::string & repo_id) {
|
||||
if (!fs::exists(snapshots_path)) {
|
||||
continue;
|
||||
}
|
||||
std::string _repo_id = folder_name_to_repo(repo.path().filename().string());
|
||||
std::string _repo_id = folder_name_to_repo(fs_path_to_utf8(repo.path().filename()));
|
||||
|
||||
if (!is_valid_repo_id(_repo_id)) {
|
||||
continue;
|
||||
@@ -446,8 +448,9 @@ hf_files get_cached_files(const std::string & repo_id) {
|
||||
if (!path.empty()) {
|
||||
hf_file file;
|
||||
file.repo_id = _repo_id;
|
||||
file.path = path.generic_string();
|
||||
file.local_path = entry.path().string();
|
||||
const auto generic_path = path.generic_u8string();
|
||||
file.path = std::string(generic_path.begin(), generic_path.end());
|
||||
file.local_path = fs_path_to_utf8(entry.path());
|
||||
file.final_path = file.local_path;
|
||||
files.push_back(std::move(file));
|
||||
}
|
||||
@@ -461,8 +464,8 @@ std::string finalize_file(const hf_file & file) {
|
||||
static std::atomic<bool> symlinks_disabled{false};
|
||||
|
||||
std::error_code ec;
|
||||
fs::path local_path(file.local_path);
|
||||
fs::path final_path(file.final_path);
|
||||
fs::path local_path = fs::u8path(file.local_path);
|
||||
fs::path final_path = fs::u8path(file.final_path);
|
||||
|
||||
if (local_path == final_path || fs::exists(final_path, ec)) {
|
||||
return file.final_path;
|
||||
@@ -509,7 +512,7 @@ bool remove_cached_repo(const std::string & repo_id) {
|
||||
std::error_code ec;
|
||||
auto removed = fs::remove_all(repo_path, ec);
|
||||
if (ec) {
|
||||
LOG_ERR("%s: failed to remove repo cache %s: %s\n", __func__, repo_path.string().c_str(), ec.message().c_str());
|
||||
LOG_ERR("%s: failed to remove repo cache %s: %s\n", __func__, fs_path_to_utf8(repo_path).c_str(), ec.message().c_str());
|
||||
return false;
|
||||
}
|
||||
return removed > 0;
|
||||
|
||||
@@ -429,8 +429,15 @@ private:
|
||||
bool negate = false;
|
||||
if (is_identifier("not")) { ++current; negate = true; }
|
||||
auto test_id = parse_primary_expression();
|
||||
// FIXME: tests can also be expressed like this: if x is eq 3
|
||||
if (is(token::open_paren)) test_id = parse_call_expression(std::move(test_id));
|
||||
if (is(token::open_paren)) {
|
||||
test_id = parse_call_expression(std::move(test_id));
|
||||
} else if (is(token::numeric_literal) || is(token::string_literal) || is(token::open_curly_bracket) || is(token::open_square_bracket) ||
|
||||
(is(token::identifier) && !is_identifier("and") && !is_identifier("or") && !is_identifier("else"))) {
|
||||
size_t call_pos = current;
|
||||
statements args;
|
||||
args.push_back(parse_unary_expression());
|
||||
test_id = mk_stmt<call_expression>(call_pos, std::move(test_id), std::move(args));
|
||||
}
|
||||
operand = mk_stmt<test_expression>(start_pos, std::move(operand), negate, std::move(test_id));
|
||||
}
|
||||
return operand;
|
||||
|
||||
+61
-14
@@ -348,6 +348,43 @@ static value default_value(const func_args & args) {
|
||||
return no_value ? args.get_pos(1) : args.get_pos(0);
|
||||
}
|
||||
|
||||
static value toobject(const func_args & args) {
|
||||
auto out = mk_val<value_object>();
|
||||
value iter = args.get_pos(0, mk_val<value_undefined>());
|
||||
bool iter_first = false;
|
||||
if (is_val<value_array>(iter)) {
|
||||
iter_first = true;
|
||||
for (const auto & it : iter->as_array()) {
|
||||
if (is_val<value_array>(it) && it->as_array().size() == 2) {
|
||||
auto tuple = it->as_array();
|
||||
auto key = tuple[0];
|
||||
auto val = tuple[1];
|
||||
JJ_DEBUG("namespace/dict: adding key '%s'", key->as_string().str().c_str());
|
||||
out->insert(key, val);
|
||||
} else {
|
||||
throw raised_exception("namespace/dict() iterable argument must consist of tuples, not " + it->type());
|
||||
}
|
||||
}
|
||||
} else if (is_val<value_object>(iter)) {
|
||||
iter_first = true;
|
||||
for (const auto & pair : iter->as_ordered_object()) {
|
||||
JJ_DEBUG("namespace/dict: adding key '%s'", pair.first->as_string().str().c_str());
|
||||
out->insert(pair.first, pair.second);
|
||||
}
|
||||
}
|
||||
for (const auto & arg : args.get_args()) {
|
||||
if (is_val<value_kwarg>(arg)) {
|
||||
auto kwarg = cast_val<value_kwarg>(arg);
|
||||
JJ_DEBUG("namespace/dict: adding key '%s'", kwarg->key.c_str());
|
||||
out->insert(kwarg->key, kwarg->val);
|
||||
} else if (!iter_first) {
|
||||
throw raised_exception("namespace/dict() arguments must be kwargs, dict and/or iterable of tuples, not " + arg->type());
|
||||
}
|
||||
iter_first = false;
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
const func_builtins & global_builtins() {
|
||||
static const func_builtins builtins = {
|
||||
{"raise_exception", [](const func_args & args) -> value {
|
||||
@@ -355,18 +392,8 @@ const func_builtins & global_builtins() {
|
||||
std::string msg = args.get_pos(0)->as_string().str();
|
||||
throw raised_exception("Jinja Exception: " + msg);
|
||||
}},
|
||||
{"namespace", [](const func_args & args) -> value {
|
||||
auto out = mk_val<value_object>();
|
||||
for (const auto & arg : args.get_args()) {
|
||||
if (!is_val<value_kwarg>(arg)) {
|
||||
throw raised_exception("namespace() arguments must be kwargs");
|
||||
}
|
||||
auto kwarg = cast_val<value_kwarg>(arg);
|
||||
JJ_DEBUG("namespace: adding key '%s'", kwarg->key.c_str());
|
||||
out->insert(kwarg->key, kwarg->val);
|
||||
}
|
||||
return out;
|
||||
}},
|
||||
{"dict", toobject},
|
||||
{"namespace", toobject},
|
||||
{"strftime_now", [](const func_args & args) -> value {
|
||||
args.ensure_vals<value_string>();
|
||||
std::string format = args.get_pos(0)->as_string().str();
|
||||
@@ -515,8 +542,28 @@ const func_builtins & global_builtins() {
|
||||
}},
|
||||
{"test_is_sameas", [](const func_args & args) -> value {
|
||||
// Check if an object points to the same memory address as another object
|
||||
(void)args;
|
||||
throw not_implemented_exception("sameas test not implemented");
|
||||
args.ensure_count(2);
|
||||
auto a = args.get_pos(0);
|
||||
auto b = args.get_pos(1);
|
||||
bool res = false;
|
||||
if (!is_val<value_undefined>(a) && !is_val<value_undefined>(b)) {
|
||||
if (is_val<value_none>(a) && is_val<value_none>(b)) {
|
||||
res = true;
|
||||
} else if (is_val<value_bool>(a) && is_val<value_bool>(b)) {
|
||||
if (a->as_bool() == b->as_bool()) {
|
||||
res = true;
|
||||
}
|
||||
} else if (is_val<value_int>(a) && is_val<value_int>(b)) {
|
||||
const int64_t x = a->as_int();
|
||||
// Allow comparison within small-int cache range
|
||||
if (x >= -5 && x <= 256 && x == b->as_int()) {
|
||||
res = true;
|
||||
}
|
||||
} else if (a == b) {
|
||||
res = true;
|
||||
}
|
||||
}
|
||||
return mk_val<value_bool>(res);
|
||||
}},
|
||||
{"test_is_escaped", [](const func_args & args) -> value {
|
||||
(void)args;
|
||||
|
||||
@@ -43,9 +43,10 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
|
||||
|
||||
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||
|
||||
auto has_tools = inputs.tools.is_array() && !inputs.tools.empty();
|
||||
// Constrained grammar whenever tools are offered.
|
||||
auto include_grammar = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE;
|
||||
auto has_tools = inputs.tools.is_array() && !inputs.tools.empty();
|
||||
auto has_response_format = !inputs.json_schema.is_null() && inputs.json_schema.is_object();
|
||||
// Constrained grammar whenever tools are offered or a response format is requested.
|
||||
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto start = p.rule("start", p.literal("<|start|>assistant"));
|
||||
@@ -65,6 +66,15 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
|
||||
auto final_msg = p.rule("final", recipient + p.literal("<|message|>") +
|
||||
p.content(p.until_one_of({ "<|eot|>", "<|eom|>" })));
|
||||
|
||||
if (has_response_format) {
|
||||
auto response_json = p.content(p.schema(p.json(), "response-format-schema", inputs.json_schema));
|
||||
auto response_format = p.rule("response-format",
|
||||
recipient + p.literal("<|message|>") +
|
||||
((p.literal("```json") + p.space() + response_json + p.space() + p.literal("```")) | response_json));
|
||||
|
||||
return p.zero_or_more(start + analysis) + start + response_format;
|
||||
}
|
||||
|
||||
if (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE) {
|
||||
auto string_value = p.ac(
|
||||
p.tool_arg_string_value(p.until("</atem:parameter>")) + p.tool_arg_close(p.literal("</atem:parameter>")),
|
||||
@@ -124,7 +134,7 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED;
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
+1
-1
@@ -214,7 +214,7 @@ struct common_sampler * common_sampler_init(
|
||||
#ifdef LLAMA_USE_LLGUIDANCE
|
||||
grmr = llama_sampler_init_llg(vocab, "lark", grammar_str.c_str());
|
||||
#else
|
||||
GGML_ABORT("llguidance (cmake -DLLAMA_LLGUIDANCE=ON) is not enabled");
|
||||
throw std::runtime_error("failed to parse grammar: llguidance is not enabled");
|
||||
#endif // LLAMA_USE_LLGUIDANCE
|
||||
} else {
|
||||
std::vector<std::string> trigger_patterns;
|
||||
|
||||
+183
-186
@@ -165,7 +165,7 @@ struct common_speculative_impl {
|
||||
|
||||
virtual void begin(llama_seq_id seq_id, const llama_tokens & prompt) = 0;
|
||||
|
||||
virtual bool process(const llama_batch & batch) = 0;
|
||||
virtual bool process(const common_batch & batch) = 0;
|
||||
|
||||
virtual void draft(common_speculative_draft_params_vec & dparams) = 0;
|
||||
|
||||
@@ -179,7 +179,11 @@ struct common_speculative_impl {
|
||||
struct common_speculative_impl_draft_simple : public common_speculative_impl {
|
||||
common_params_speculative_draft params;
|
||||
|
||||
llama_batch batch;
|
||||
common_batch batch;
|
||||
|
||||
// zero row at the draft input width, stands in for target embeddings the draft cannot read
|
||||
std::vector<float> zeros;
|
||||
bool zeros_warned = false; // the substitution is reported once
|
||||
|
||||
std::vector<common_sampler_ptr> smpls;
|
||||
|
||||
@@ -194,6 +198,8 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
|
||||
throw std::runtime_error("draft-simple requires a draft context");
|
||||
}
|
||||
|
||||
zeros.assign(llama_model_n_embd_inp(llama_get_model(ctx_dft)), 0.0f);
|
||||
|
||||
SPC_TRC("%s", "adding speculative implementation 'draft-simple'\n");
|
||||
SPC_TRC("- n_max=%d, n_min=%d, p_min=%f\n", this->params.n_max, this->params.n_min, this->params.p_min);
|
||||
SPC_TRC("- gpu_layers=%d, cache_k=%s, cache_v=%s, ctx_tgt=%s, ctx_dft=%s, devices=[%s]\n",
|
||||
@@ -204,7 +210,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
|
||||
ctx_dft ? "yes" : "no",
|
||||
common_speculative_get_devices_str(this->params.devices).c_str());
|
||||
|
||||
batch = llama_batch_init(llama_n_batch(ctx_dft), 0, 1);
|
||||
batch = common_batch(ctx_dft);
|
||||
|
||||
// TODO: optimize or pass from outside?
|
||||
// {
|
||||
@@ -249,21 +255,46 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
|
||||
}
|
||||
}
|
||||
|
||||
~common_speculative_impl_draft_simple() override {
|
||||
llama_batch_free(batch);
|
||||
}
|
||||
|
||||
void begin(llama_seq_id /*seq_id*/, const llama_tokens & /*prompt*/) override {
|
||||
// noop
|
||||
}
|
||||
|
||||
bool process(const llama_batch & batch) override {
|
||||
bool process(const common_batch & batch_in) override {
|
||||
auto * ctx_dft = params.ctx_dft;
|
||||
|
||||
llama_batch batch_dft = batch;
|
||||
batch_dft.logits = nullptr;
|
||||
// copy the entries to a batch owned by the draft context, only the last token is output
|
||||
batch.clear();
|
||||
const int32_t n_tokens = batch_in.size();
|
||||
for (int32_t k = 0; k < n_tokens; ++k) {
|
||||
const auto & t = batch_in.tokens[k];
|
||||
const bool output = k == n_tokens - 1;
|
||||
if (t.id != LLAMA_TOKEN_NULL) {
|
||||
const int32_t idx = batch.add(t.id, t.pos[0], t.seq_id, output);
|
||||
if (t.embd.data) {
|
||||
batch.set_embd(idx, t.embd);
|
||||
}
|
||||
} else {
|
||||
// mtmd input is projected by the target encoder, a draft with a different width cannot read it
|
||||
// it gets zeros instead, keeping its positions contiguous
|
||||
// ref: https://github.com/ggml-org/llama.cpp/pull/29385#discussion_r4124743243
|
||||
const size_t n_embd = t.embd.n_rows * t.embd.n_embd;
|
||||
const bool same_width = n_embd == zeros.size();
|
||||
if (!same_width && !zeros_warned) {
|
||||
SPC_WRN("target embeddings of size %zu do not fit the draft input width %zu, "
|
||||
"the draft receives zero rows for them and drafts after multimodal input will be poor\n",
|
||||
n_embd, zeros.size());
|
||||
zeros_warned = true;
|
||||
}
|
||||
const llama_embd embd = same_width ? t.embd : llama_embd{ zeros.data(), 1, zeros.size() };
|
||||
batch.add_embd(embd, t.pos.data(), t.seq_id, output);
|
||||
}
|
||||
}
|
||||
|
||||
const int ret = llama_decode(ctx_dft, batch_dft);
|
||||
if (batch.size() == 0) {
|
||||
return true;
|
||||
}
|
||||
|
||||
const int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
|
||||
|
||||
if (ret != 0) {
|
||||
SPC_ERR("failed to decode draft batch, ret = %d\n", ret);
|
||||
@@ -277,7 +308,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
|
||||
void draft(common_speculative_draft_params_vec & dparams) override {
|
||||
auto & ctx_dft = params.ctx_dft;
|
||||
|
||||
common_batch_clear(batch);
|
||||
batch.clear();
|
||||
|
||||
// keep track of which sequences are still drafting
|
||||
int n_drafting = 0;
|
||||
@@ -294,12 +325,12 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
|
||||
drafting[seq_id] = true;
|
||||
common_sampler_reset(smpls[seq_id].get());
|
||||
|
||||
common_batch_add(batch, dp.id_last, dp.pos0, { seq_id }, true);
|
||||
batch.add(dp.id_last, dp.pos0, seq_id, true);
|
||||
}
|
||||
|
||||
int ret = llama_decode(ctx_dft, batch);
|
||||
int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
|
||||
if (ret != 0) {
|
||||
SPC_ERR("llama_decode returned %d\n", ret);
|
||||
SPC_ERR("llama_process returned %d\n", ret);
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -308,7 +339,7 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
|
||||
while (n_drafting > 0) {
|
||||
int i_batch = 0;
|
||||
|
||||
common_batch_clear(batch);
|
||||
batch.clear();
|
||||
|
||||
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
|
||||
if (!drafting[seq_id]) {
|
||||
@@ -353,17 +384,17 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
|
||||
continue;
|
||||
}
|
||||
|
||||
common_batch_add(batch, id, dp.pos0 + i + 1, { seq_id }, true);
|
||||
batch.add(id, dp.pos0 + i + 1, seq_id, true);
|
||||
}
|
||||
|
||||
if (batch.n_tokens == 0) {
|
||||
if (batch.size() == 0) {
|
||||
break;
|
||||
}
|
||||
|
||||
// evaluate the drafted tokens on the draft model
|
||||
ret = llama_decode(ctx_dft, batch);
|
||||
ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
|
||||
if (ret != 0) {
|
||||
SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
|
||||
SPC_ERR("llama_process[%d] returned %d\n", i, ret);
|
||||
break;
|
||||
}
|
||||
|
||||
@@ -423,7 +454,8 @@ struct common_speculative_impl_draft_simple : public common_speculative_impl {
|
||||
// encoder+decoder on n_accepted+1 rows).
|
||||
struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
common_params_speculative_draft params;
|
||||
llama_batch batch;
|
||||
common_batch batch; // decoder input, (token, g_embd) pairs
|
||||
common_batch batch_enc; // encoder input, built from the extracted target features
|
||||
|
||||
std::vector<common_sampler_ptr> smpls;
|
||||
|
||||
@@ -477,11 +509,8 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
n_embd_enc = (int32_t) target_layer_ids_n * n_embd_tgt;
|
||||
n_layer_tgt = llama_model_n_layer(model_tgt);
|
||||
|
||||
const int32_t n_b = (int32_t) llama_n_batch(ctx_dft);
|
||||
batch = llama_batch_init(/*n_tokens=*/ n_b, /*embd=*/ n_embd_dec, /*n_seq_max=*/ 1);
|
||||
// llama_batch_init allocates only one of token/embd; eagle3 decoder needs both.
|
||||
// TODO: fix, how to call without malloc
|
||||
batch.token = (llama_token *) malloc(sizeof(llama_token) * n_b);
|
||||
batch = common_batch(ctx_dft);
|
||||
batch_enc = common_batch(ctx_dft);
|
||||
|
||||
smpls.resize(n_seq);
|
||||
for (auto & s : smpls) {
|
||||
@@ -543,12 +572,6 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
llama_sampler_free(backend_chains[seq_id]);
|
||||
}
|
||||
backend_chains.clear();
|
||||
|
||||
if (batch.token != nullptr) {
|
||||
free(batch.token);
|
||||
batch.token = nullptr;
|
||||
}
|
||||
llama_batch_free(batch);
|
||||
}
|
||||
|
||||
void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
|
||||
@@ -567,16 +590,16 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
}
|
||||
}
|
||||
|
||||
bool process(const llama_batch & batch_in) override {
|
||||
if (batch_in.n_tokens <= 0) {
|
||||
bool process(const common_batch & batch_in) override {
|
||||
if (batch_in.size() <= 0) {
|
||||
return true;
|
||||
}
|
||||
|
||||
if (batch_in.token == nullptr || batch_in.embd != nullptr) {
|
||||
if (!batch_in.has_token() || batch_in.has_embd()) {
|
||||
return true;
|
||||
}
|
||||
|
||||
const int32_t n_tokens = batch_in.n_tokens;
|
||||
const int32_t n_tokens = batch_in.size();
|
||||
|
||||
// i_batch_beg[seq] / i_batch_end[seq]: inclusive batch indices of this seq's
|
||||
// first/last token in batch_in. Assumes per-seq tokens are contiguous within
|
||||
@@ -584,8 +607,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
std::vector<int32_t> i_batch_beg(n_seq, -1);
|
||||
std::vector<int32_t> i_batch_end(n_seq, -1);
|
||||
for (int k = 0; k < n_tokens; ++k) {
|
||||
GGML_ASSERT(batch_in.n_seq_id[k] == 1);
|
||||
const llama_seq_id seq_id = batch_in.seq_id[k][0];
|
||||
const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
|
||||
if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {
|
||||
continue;
|
||||
}
|
||||
@@ -619,24 +641,23 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
|
||||
g_embd_buf.resize((size_t) n_tokens * n_embd_dec);
|
||||
|
||||
// llama_encode() requires the full encoder batch to fit in n_ubatch.
|
||||
// llama_process() requires the full encoder batch to fit in n_ubatch.
|
||||
// Allow batch > ubatch: eagle3's per-token encoder can be chunked safely.
|
||||
const int32_t n_ubatch_dft = (int32_t) llama_n_ubatch(ctx_dft);
|
||||
for (int32_t i = 0; i < n_tokens; i += n_ubatch_dft) {
|
||||
const int32_t n_chunk = std::min(n_ubatch_dft, n_tokens - i);
|
||||
|
||||
llama_batch enc_batch = {
|
||||
/*.n_tokens =*/ n_chunk,
|
||||
/*.token =*/ nullptr,
|
||||
/*.embd =*/ features_buf.data() + (size_t) i * n_embd_enc,
|
||||
/*.pos =*/ nullptr,
|
||||
/*.n_seq_id =*/ nullptr,
|
||||
/*.seq_id =*/ nullptr,
|
||||
/*.logits =*/ nullptr,
|
||||
};
|
||||
const int32_t rc = llama_encode(ctx_dft, enc_batch);
|
||||
// the per-token encoder does not use positions, generate placeholder ones from the memory state
|
||||
batch_enc.clear();
|
||||
llama_pos pos = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), 0) + 1;
|
||||
for (int32_t j = 0; j < n_chunk; ++j) {
|
||||
batch_enc.add_embd({ features_buf.data() + (size_t) (i + j) * n_embd_enc, 1, (size_t) n_embd_enc }, &pos, 0, true);
|
||||
pos++;
|
||||
}
|
||||
|
||||
const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_ENCODE, batch_enc.get());
|
||||
if (rc != 0) {
|
||||
SPC_ERR("llama_encode(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
|
||||
SPC_ERR("llama_process(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
|
||||
rc, (int) n_chunk, (int) i);
|
||||
return false;
|
||||
}
|
||||
@@ -664,7 +685,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
// deferred boundary, completed by the next process() or draft() call.
|
||||
// (c) refresh deferred state — stash this ubatch's full g_embd into verify_g,
|
||||
// update pending_g_last / pending_pos_last to the last row.
|
||||
common_batch_clear(batch);
|
||||
batch.clear();
|
||||
|
||||
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
|
||||
const int32_t beg = i_batch_beg[seq_id];
|
||||
@@ -679,36 +700,34 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
// 2) pending_pos_last + 1 == pos[beg]
|
||||
// 3) pending_pos_last > dft_pos_max // TODO: is this check needed?
|
||||
const llama_pos pending_pos = pending_pos_last[seq_id];
|
||||
if (pending_pos >= 0 && pending_pos + 1 == batch_in.pos[beg]) {
|
||||
if (pending_pos >= 0 && pending_pos + 1 == batch_in.tokens[beg].pos[0]) {
|
||||
const llama_pos dft_pos_max = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), seq_id);
|
||||
if (pending_pos > dft_pos_max) {
|
||||
common_batch_add(batch, batch_in.token[beg], pending_pos, { seq_id }, /*logits=*/ false);
|
||||
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
|
||||
pending_g_last[seq_id].data(), row_bytes);
|
||||
const int32_t idx = batch.add(batch_in.tokens[beg].id, pending_pos, seq_id, /*output=*/ false);
|
||||
batch.set_embd(idx, { pending_g_last[seq_id].data(), 1, (size_t) n_embd_dec });
|
||||
}
|
||||
}
|
||||
|
||||
for (int32_t k = beg; k < end; ++k) {
|
||||
common_batch_add(batch, batch_in.token[k + 1], batch_in.pos[k], { seq_id }, /*logits=*/ false);
|
||||
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
|
||||
g_embd + (size_t) k * n_embd_dec, row_bytes);
|
||||
const int32_t idx = batch.add(batch_in.tokens[k + 1].id, batch_in.tokens[k].pos[0], seq_id, /*output=*/ false);
|
||||
batch.set_embd(idx, { g_embd + (size_t) k * n_embd_dec, 1, (size_t) n_embd_dec });
|
||||
}
|
||||
|
||||
// refresh deferred state
|
||||
const int32_t n_rows = end - beg + 1;
|
||||
verify_pos_first[seq_id] = batch_in.pos[beg];
|
||||
pending_pos_last[seq_id] = batch_in.pos[end];
|
||||
verify_pos_first[seq_id] = batch_in.tokens[beg].pos[0];
|
||||
pending_pos_last[seq_id] = batch_in.tokens[end].pos[0];
|
||||
verify_g_rows[seq_id] = n_rows;
|
||||
verify_g[seq_id].resize((size_t) n_rows * n_embd_dec, 0.0f);
|
||||
std::memcpy(verify_g[seq_id].data(), g_embd + (size_t) beg * n_embd_dec, row_bytes * n_rows);
|
||||
std::memcpy(pending_g_last[seq_id].data(), g_embd + (size_t) end * n_embd_dec, row_bytes);
|
||||
}
|
||||
|
||||
if (batch.n_tokens > 0) {
|
||||
const int32_t rc = llama_decode(ctx_dft, batch);
|
||||
if (batch.size() > 0) {
|
||||
const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
|
||||
if (rc != 0) {
|
||||
SPC_ERR("llama_decode(ctx_dft) failed rc=%d (n_tokens=%d, ubatch_pos[0]=%d)\n",
|
||||
rc, (int) batch.n_tokens, (int) batch_in.pos[0]);
|
||||
SPC_ERR("llama_process(ctx_dft) failed rc=%d (n_tokens=%d, ubatch_pos[0]=%d)\n",
|
||||
rc, (int) batch.size(), (int) batch_in.tokens[0].pos[0]);
|
||||
return false;
|
||||
}
|
||||
}
|
||||
@@ -719,14 +738,12 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
void draft(common_speculative_draft_params_vec & dparams) override {
|
||||
auto & ctx_dft = params.ctx_dft;
|
||||
|
||||
common_batch_clear(batch);
|
||||
batch.clear();
|
||||
|
||||
// keep track of which sequences are still drafting
|
||||
int n_drafting = 0;
|
||||
std::vector<bool> drafting(n_seq);
|
||||
|
||||
const size_t row_bytes = (size_t) n_embd_dec * sizeof(float);
|
||||
|
||||
// Complete the deferred boundary pair (dp.id_last, pending_g_last) at memory
|
||||
// pos pending_pos_last. dp.id_last is target's freshest sample (= corrected
|
||||
// token after verify, or first generated token after prefill), matching the
|
||||
@@ -747,19 +764,17 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
|
||||
llama_memory_seq_rm(llama_get_memory(ctx_dft), seq_id, pending_pos_last[seq_id], -1);
|
||||
|
||||
common_batch_add(batch, dp.id_last, pending_pos_last[seq_id], { seq_id }, true);
|
||||
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,
|
||||
pending_g_last[seq_id].data(),
|
||||
row_bytes);
|
||||
const int32_t idx = batch.add(dp.id_last, pending_pos_last[seq_id], seq_id, true);
|
||||
batch.set_embd(idx, { pending_g_last[seq_id].data(), 1, (size_t) n_embd_dec });
|
||||
}
|
||||
|
||||
if (batch.n_tokens == 0) {
|
||||
if (batch.size() == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
int ret = llama_decode(ctx_dft, batch);
|
||||
int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
|
||||
if (ret != 0) {
|
||||
SPC_ERR("llama_decode returned %d\n", ret);
|
||||
SPC_ERR("llama_process returned %d\n", ret);
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -768,7 +783,7 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
while (n_drafting > 0) {
|
||||
int i_batch = 0;
|
||||
|
||||
common_batch_clear(batch);
|
||||
batch.clear();
|
||||
|
||||
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
|
||||
if (!drafting[seq_id]) {
|
||||
@@ -814,17 +829,17 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
continue;
|
||||
}
|
||||
|
||||
common_batch_add(batch, id, pending_pos_last[seq_id] + (i + 1), { seq_id }, true);
|
||||
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec, prenorm, row_bytes);
|
||||
const int32_t idx = batch.add(id, pending_pos_last[seq_id] + (i + 1), seq_id, true);
|
||||
batch.set_embd(idx, { prenorm, 1, (size_t) n_embd_dec });
|
||||
}
|
||||
|
||||
if (batch.n_tokens == 0) {
|
||||
if (batch.size() == 0) {
|
||||
break;
|
||||
}
|
||||
|
||||
ret = llama_decode(ctx_dft, batch);
|
||||
ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
|
||||
if (ret != 0) {
|
||||
SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
|
||||
SPC_ERR("llama_process[%d] returned %d\n", i, ret);
|
||||
break;
|
||||
}
|
||||
|
||||
@@ -908,8 +923,10 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {
|
||||
struct common_speculative_impl_draft_dflash : public common_speculative_impl {
|
||||
common_params_speculative_draft params;
|
||||
|
||||
llama_batch batch; // noise tokens
|
||||
llama_batch batch_inject; // target features for KV cache injection
|
||||
common_batch batch; // noise tokens
|
||||
common_batch batch_inject; // target features for KV cache injection
|
||||
|
||||
std::vector<float> features_buf; // [n_chunk, n_embd_enc] gathered target features
|
||||
|
||||
std::vector<common_sampler_ptr> smpls;
|
||||
|
||||
@@ -1005,15 +1022,11 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
|
||||
}
|
||||
this->n_max = this->params.n_max;
|
||||
|
||||
batch = llama_batch_init(llama_n_batch(ctx_dft), 0, n_seq);
|
||||
batch_inject = llama_batch_init(llama_n_ubatch(ctx_dft), n_embd_enc, n_seq);
|
||||
batch = common_batch(ctx_dft);
|
||||
batch_inject = common_batch(ctx_dft);
|
||||
|
||||
// embd batches on an M-RoPE draft need 4 position rows per token
|
||||
// embd batches on an M-RoPE draft carry 4 position rows per token
|
||||
is_mrope = llama_model_rope_type(model_dft) == LLAMA_ROPE_TYPE_MROPE;
|
||||
if (is_mrope) {
|
||||
free(batch_inject.pos);
|
||||
batch_inject.pos = (llama_pos *) malloc(sizeof(llama_pos) * 4 * llama_n_batch(ctx_dft));
|
||||
}
|
||||
|
||||
smpls.resize(n_seq);
|
||||
for (auto & s : smpls) {
|
||||
@@ -1062,9 +1075,6 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
|
||||
llama_sampler_free(backend_chains[seq_id]);
|
||||
}
|
||||
backend_chains.clear();
|
||||
|
||||
llama_batch_free(batch);
|
||||
llama_batch_free(batch_inject);
|
||||
}
|
||||
|
||||
void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
|
||||
@@ -1085,8 +1095,8 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
|
||||
}
|
||||
}
|
||||
|
||||
bool process(const llama_batch & batch_in) override {
|
||||
if (batch_in.n_tokens <= 0) {
|
||||
bool process(const common_batch & batch_in) override {
|
||||
if (batch_in.size() <= 0) {
|
||||
return true;
|
||||
}
|
||||
|
||||
@@ -1094,20 +1104,19 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
|
||||
// produce the target-layer features used to seed the draft KV cache, so
|
||||
// embeddings are injected too, except the pinned ones skipped below.
|
||||
// TODO: revisit after https://github.com/ggml-org/llama.cpp/pull/24669 is merged
|
||||
const bool has_tokens = batch_in.token != nullptr;
|
||||
const bool has_embeddings = batch_in.embd != nullptr;
|
||||
const bool has_tokens = batch_in.has_token();
|
||||
const bool has_embeddings = batch_in.has_embd();
|
||||
if (has_tokens == has_embeddings) {
|
||||
return true;
|
||||
}
|
||||
|
||||
const int32_t n_tokens = batch_in.n_tokens;
|
||||
const int32_t n_tokens = batch_in.size();
|
||||
|
||||
// per-seq inclusive batch range (assumes each seq's tokens are contiguous in the batch)
|
||||
std::vector<int32_t> i_batch_beg(n_seq, -1);
|
||||
std::vector<int32_t> i_batch_end(n_seq, -1);
|
||||
for (int32_t k = 0; k < n_tokens; ++k) {
|
||||
GGML_ASSERT(batch_in.n_seq_id[k] == 1);
|
||||
const llama_seq_id seq_id = batch_in.seq_id[k][0];
|
||||
const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
|
||||
if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {
|
||||
continue;
|
||||
}
|
||||
@@ -1130,7 +1139,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
|
||||
|
||||
// an M-RoPE image pins all its rows to one position, so a windowed draft
|
||||
// cache cannot free cells for it - skip it, the draft can jump over the gap
|
||||
const bool pos_pinned = batch_in.pos[i_batch_beg[seq_id]] == batch_in.pos[i_batch_end[seq_id]];
|
||||
const bool pos_pinned = batch_in.tokens[i_batch_beg[seq_id]].pos[0] == batch_in.tokens[i_batch_end[seq_id]].pos[0];
|
||||
if (has_embeddings && n_rows > 1 && pos_pinned) {
|
||||
continue;
|
||||
}
|
||||
@@ -1140,34 +1149,28 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
|
||||
|
||||
// gather target features per extract layer; the fused decode encodes and
|
||||
// injects them into the K/V cache at the target positions
|
||||
batch_inject.n_tokens = n_chunk;
|
||||
features_buf.resize((size_t) n_chunk * n_embd_enc);
|
||||
for (uint32_t k = 0; k < target_layer_ids_n; ++k) {
|
||||
const float * layer = llama_get_embeddings_layer_inp(ctx_tgt, (uint32_t) target_layer_ids[k]);
|
||||
if (!layer) {
|
||||
GGML_ABORT("DFlash: target layer %d input not extracted.", target_layer_ids[k]);
|
||||
}
|
||||
for (int32_t i = 0; i < n_chunk; ++i) {
|
||||
float * dst = batch_inject.embd + (size_t) i * n_embd_enc + k * (size_t) n_embd_tgt;
|
||||
float * dst = features_buf.data() + (size_t) i * n_embd_enc + k * (size_t) n_embd_tgt;
|
||||
const float * src = layer + (size_t) (i_batch_beg[seq_id] + offset + i) * n_embd_tgt;
|
||||
std::memcpy(dst, src, (size_t) n_embd_tgt * sizeof(float));
|
||||
}
|
||||
}
|
||||
|
||||
batch_inject.clear();
|
||||
for (int32_t i = 0; i < n_chunk; ++i) {
|
||||
const llama_pos p = batch_in.pos[i_batch_beg[seq_id] + offset + i];
|
||||
batch_inject.pos[i] = p;
|
||||
if (is_mrope) {
|
||||
batch_inject.pos[1 * n_chunk + i] = p;
|
||||
batch_inject.pos[2 * n_chunk + i] = p;
|
||||
batch_inject.pos[3 * n_chunk + i] = 0;
|
||||
}
|
||||
batch_inject.n_seq_id[i] = 1;
|
||||
batch_inject.seq_id[i][0] = seq_id;
|
||||
batch_inject.logits[i] = false;
|
||||
const llama_pos p = batch_in.tokens[i_batch_beg[seq_id] + offset + i].pos[0];
|
||||
const llama_pos pos_arr[4] = { p, p, p, 0 };
|
||||
batch_inject.add_embd({ features_buf.data() + (size_t) i * n_embd_enc, 1, (size_t) n_embd_enc }, pos_arr, seq_id, false);
|
||||
}
|
||||
const int32_t rc = llama_decode(ctx_dft, batch_inject);
|
||||
const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch_inject.get());
|
||||
if (rc != 0) {
|
||||
LOG_ERR("%s: llama_decode(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
|
||||
LOG_ERR("%s: llama_process(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",
|
||||
__func__, rc, (int) n_chunk, (int) offset);
|
||||
return false;
|
||||
}
|
||||
@@ -1180,7 +1183,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
|
||||
void draft(common_speculative_draft_params_vec & dparams) override {
|
||||
auto & ctx_dft = params.ctx_dft;
|
||||
|
||||
common_batch_clear(batch);
|
||||
batch.clear();
|
||||
|
||||
// build one batch holding every drafting sequence's noise block into a single decode)
|
||||
// record where each block starts and its size
|
||||
@@ -1200,21 +1203,21 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
|
||||
const int32_t n_draft = params.n_max;
|
||||
|
||||
const int32_t n_block_tokens = n_draft + (is_dspark && sample_from_anchor ? 0 : 1);
|
||||
i_block_beg[seq_id] = batch.n_tokens;
|
||||
i_block_beg[seq_id] = batch.size();
|
||||
n_block [seq_id] = n_block_tokens;
|
||||
for (int32_t i = 0; i < n_block_tokens; ++i) {
|
||||
common_batch_add(batch, i == 0 ? dp.id_last : mask_token_id, n + i, { seq_id }, !is_dflash2);
|
||||
batch.add(i == 0 ? dp.id_last : mask_token_id, n + i, seq_id, !is_dflash2);
|
||||
}
|
||||
}
|
||||
|
||||
if (batch.n_tokens == 0) {
|
||||
if (batch.size() == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
// decode all sequence's noise block in a single batch
|
||||
int ret = llama_decode(ctx_dft, batch);
|
||||
int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
|
||||
if (ret != 0) {
|
||||
LOG_WRN("%s: llama_decode returned %d\n", __func__, ret);
|
||||
LOG_WRN("%s: llama_process returned %d\n", __func__, ret);
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -1328,7 +1331,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl {
|
||||
struct common_speculative_impl_draft_mtp : public common_speculative_impl {
|
||||
common_params_speculative_draft params; // reuses the draft-model params slot (ctx_tgt/ctx_dft)
|
||||
|
||||
llama_batch batch;
|
||||
common_batch batch;
|
||||
|
||||
std::vector<common_sampler_ptr> smpls;
|
||||
|
||||
@@ -1384,11 +1387,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
|
||||
ctx_dft ? "yes" : "no",
|
||||
common_speculative_get_devices_str(this->params.devices).c_str());
|
||||
|
||||
const int32_t n_b = (int32_t) llama_n_batch(ctx_dft);
|
||||
batch = llama_batch_init(/*n_tokens=*/ n_b, /*embd=*/ n_embd, /*n_seq_max=*/ 1);
|
||||
// llama_batch_init allocates only one of token/embd; MTP needs both.
|
||||
// TODO: fix, how to call without malloc
|
||||
batch.token = (llama_token *) malloc(sizeof(llama_token) * n_b);
|
||||
batch = common_batch(ctx_dft);
|
||||
|
||||
smpls.resize(n_seq);
|
||||
for (auto & s : smpls) {
|
||||
@@ -1453,12 +1452,6 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
|
||||
llama_sampler_free(backend_chains[seq_id]);
|
||||
}
|
||||
backend_chains.clear();
|
||||
|
||||
if (batch.token != nullptr) {
|
||||
free(batch.token);
|
||||
batch.token = nullptr;
|
||||
}
|
||||
llama_batch_free(batch);
|
||||
}
|
||||
|
||||
void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {
|
||||
@@ -1473,23 +1466,23 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
|
||||
if (pos_max < N - 1 && !is_mem_shared) {
|
||||
SPC_WRN("ctx_dft pos_max=%d < N-1=%d - "
|
||||
"process() hook may not have run on every prefill ubatch "
|
||||
"(need_embd / logits=1 on every prompt position?). "
|
||||
"(need_embd / output flag on every prompt position?). "
|
||||
"Drafts may degrade.\n",
|
||||
(int) pos_max, N - 1);
|
||||
}
|
||||
}
|
||||
|
||||
bool process(const llama_batch & batch_in) override {
|
||||
if (batch_in.n_tokens <= 0) {
|
||||
bool process(const common_batch & batch_in) override {
|
||||
if (batch_in.size() <= 0) {
|
||||
return true;
|
||||
}
|
||||
|
||||
// TODO: how to make it work with vision tokens?
|
||||
if (batch_in.token == nullptr || batch_in.embd != nullptr) {
|
||||
if (!batch_in.has_token() || batch_in.has_embd()) {
|
||||
return true;
|
||||
}
|
||||
|
||||
const int32_t n_tokens = batch_in.n_tokens;
|
||||
const int32_t n_tokens = batch_in.size();
|
||||
|
||||
// remember the first and last batch index for each sequence
|
||||
std::fill(i_batch_beg.begin(), i_batch_beg.end(), -1);
|
||||
@@ -1497,9 +1490,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
|
||||
|
||||
for (int k = 0; k < n_tokens; ++k) {
|
||||
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
|
||||
GGML_ASSERT(batch_in.n_seq_id[k] == 1);
|
||||
|
||||
if (batch_in.seq_id[k][0] == seq_id) {
|
||||
if (batch_in.tokens[k].seq_id == seq_id) {
|
||||
i_batch_end[seq_id] = k;
|
||||
if (i_batch_beg[seq_id] < 0) {
|
||||
i_batch_beg[seq_id] = k;
|
||||
@@ -1515,33 +1506,26 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
|
||||
|
||||
// if kv is shared with target (e.g Gemma4), then we can skip this catch-up decode
|
||||
if (!is_mem_shared) {
|
||||
common_batch_clear(batch);
|
||||
batch.clear();
|
||||
|
||||
for (int k = 0; k < n_tokens; ++k) {
|
||||
common_batch_add(batch, batch_in.token[k], batch_in.pos[k], { batch_in.seq_id[k][0] }, 0);
|
||||
}
|
||||
|
||||
// shift the tgt embeddings to the right by one position
|
||||
// pair each token with the tgt embedding shifted right by one position, and
|
||||
// the first token of each sequence with the pending embedding from a previous run
|
||||
// assumes that the tokens in the batch are sequential for each sequence
|
||||
// i.e. we cannot have seq_id like this: [0, 0, 0, 1, 1, 0, 1, 1]
|
||||
// ^--- this is a problem
|
||||
// TODO:this is generally true, but would be nice to assert it
|
||||
{
|
||||
const float * h_tgt = llama_get_embeddings_nextn(ctx_tgt);
|
||||
std::memcpy(batch.embd + (size_t) 1 * n_embd, h_tgt, row_bytes * (n_tokens-1));
|
||||
}
|
||||
const float * h_tgt = llama_get_embeddings_nextn(ctx_tgt);
|
||||
|
||||
// fill the pending embeddings from a previous run
|
||||
auto set_h = [&](int idx, const float * h_row) {
|
||||
std::memcpy(batch.embd + (size_t) idx * n_embd, h_row, row_bytes);
|
||||
};
|
||||
for (int k = 0; k < n_tokens; ++k) {
|
||||
const llama_seq_id seq_id = batch_in.tokens[k].seq_id;
|
||||
|
||||
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
|
||||
if (i_batch_beg[seq_id] < 0) {
|
||||
continue;
|
||||
}
|
||||
const int32_t idx = batch.add(batch_in.tokens[k].id, batch_in.tokens[k].pos[0], seq_id, false);
|
||||
|
||||
set_h(i_batch_beg[seq_id], pending_h[seq_id].data());
|
||||
const float * h_row = k == i_batch_beg[seq_id]
|
||||
? pending_h[seq_id].data()
|
||||
: h_tgt + (size_t) (k - 1) * n_embd;
|
||||
|
||||
batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
|
||||
}
|
||||
|
||||
auto * mem_dft = llama_get_memory(ctx_dft);
|
||||
@@ -1554,15 +1538,15 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
|
||||
if (i_batch_beg[seq_id] < 0) {
|
||||
continue;
|
||||
}
|
||||
llama_memory_seq_rm(mem_dft, seq_id, batch_in.pos[i_batch_beg[seq_id]], -1);
|
||||
llama_memory_seq_rm(mem_dft, seq_id, batch_in.tokens[i_batch_beg[seq_id]].pos[0], -1);
|
||||
}
|
||||
llama_set_nextn_layer_offset(ctx_dft, head);
|
||||
}
|
||||
|
||||
const int32_t rc = llama_decode(ctx_dft, batch);
|
||||
const int32_t rc = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
|
||||
if (rc != 0) {
|
||||
SPC_ERR("llama_decode(ctx_dft) head=%d failed rc=%d (pos=%d)\n",
|
||||
head, (int) rc, (int) batch_in.pos[0]);
|
||||
SPC_ERR("llama_process(ctx_dft) head=%d failed rc=%d (pos=%d)\n",
|
||||
head, (int) rc, (int) batch_in.tokens[0].pos[0]);
|
||||
ok = false;
|
||||
break;
|
||||
}
|
||||
@@ -1600,14 +1584,12 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
|
||||
void draft(common_speculative_draft_params_vec & dparams) override {
|
||||
auto & ctx_dft = params.ctx_dft;
|
||||
|
||||
common_batch_clear(batch);
|
||||
batch.clear();
|
||||
|
||||
// keep track of which sequences are still drafting
|
||||
int n_drafting = 0;
|
||||
std::vector<bool> drafting(n_seq);
|
||||
|
||||
const size_t row_bytes = (size_t) n_embd * sizeof(float);
|
||||
|
||||
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
|
||||
auto & dp = dparams[seq_id];
|
||||
|
||||
@@ -1619,10 +1601,10 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
|
||||
drafting[seq_id] = true;
|
||||
common_sampler_reset(smpls[seq_id].get());
|
||||
|
||||
common_batch_add(batch, dp.id_last, dp.pos0, { seq_id }, true);
|
||||
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, pending_h[seq_id].data(), row_bytes);
|
||||
const int32_t idx = batch.add(dp.id_last, dp.pos0, seq_id, true);
|
||||
batch.set_embd(idx, { pending_h[seq_id].data(), 1, (size_t) n_embd });
|
||||
|
||||
i_last[seq_id] = batch.n_tokens - 1;
|
||||
i_last[seq_id] = idx;
|
||||
|
||||
if (chain_heads) {
|
||||
chain_h[seq_id].assign(pending_h[seq_id].begin(), pending_h[seq_id].end());
|
||||
@@ -1648,16 +1630,16 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
|
||||
llama_set_nextn_layer_offset(ctx_dft, i);
|
||||
}
|
||||
|
||||
int ret = llama_decode(ctx_dft, batch);
|
||||
int ret = llama_process(ctx_dft, LLAMA_PROCESS_TYPE_DECODE, batch.get());
|
||||
if (ret != 0) {
|
||||
SPC_ERR("llama_decode[%d] returned %d\n", i, ret);
|
||||
SPC_ERR("llama_process[%d] returned %d\n", i, ret);
|
||||
break;
|
||||
}
|
||||
|
||||
// rebuild the batch for the next step: the growing-KV paths re-add only the
|
||||
// new token (the KV already holds the prefix), while chained heads re-add the
|
||||
// whole prefix at the next head. dropped sequences are simply not re-added.
|
||||
common_batch_clear(batch);
|
||||
batch.clear();
|
||||
|
||||
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {
|
||||
if (!drafting[seq_id]) {
|
||||
@@ -1708,24 +1690,24 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl {
|
||||
const int n_rows = (int) result.size() + 1; // id_last + tokens drafted so far
|
||||
for (int t = 0; t < n_rows; ++t) {
|
||||
const llama_token tok = (t == 0) ? dp.id_last : result[t - 1];
|
||||
common_batch_add(batch, tok, dp.pos0 + t, { seq_id }, t == n_rows - 1);
|
||||
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd,
|
||||
chain_h[seq_id].data() + (size_t) t * n_embd, row_bytes);
|
||||
const int32_t idx = batch.add(tok, dp.pos0 + t, seq_id, t == n_rows - 1);
|
||||
batch.set_embd(idx, { chain_h[seq_id].data() + (size_t) t * n_embd, 1, (size_t) n_embd });
|
||||
i_last[seq_id] = idx;
|
||||
}
|
||||
} else if (is_mem_shared) {
|
||||
// note: with shared memory (e.g. Gemma4 assistants) we use the same position for all draft tokens
|
||||
// ref: https://github.com/huggingface/transformers/blob/effde20942e3f82a1b97449f60b3a48c5ff96145/docs/source/en/model_doc/gemma4_assistant.md?plain=1#L36-L37
|
||||
common_batch_add(batch, id, dp.pos0, { seq_id }, true);
|
||||
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, h_row, row_bytes);
|
||||
const int32_t idx = batch.add(id, dp.pos0, seq_id, true);
|
||||
batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
|
||||
i_last[seq_id] = idx;
|
||||
} else {
|
||||
common_batch_add(batch, id, dp.pos0 + i + 1, { seq_id }, true);
|
||||
std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, h_row, row_bytes);
|
||||
const int32_t idx = batch.add(id, dp.pos0 + i + 1, seq_id, true);
|
||||
batch.set_embd(idx, { h_row, 1, (size_t) n_embd });
|
||||
i_last[seq_id] = idx;
|
||||
}
|
||||
|
||||
i_last[seq_id] = batch.n_tokens - 1;
|
||||
}
|
||||
|
||||
if (batch.n_tokens == 0) {
|
||||
if (batch.size() == 0) {
|
||||
break;
|
||||
}
|
||||
|
||||
@@ -1787,7 +1769,7 @@ struct common_speculative_impl_ngram_simple : public common_speculative_impl {
|
||||
// noop
|
||||
}
|
||||
|
||||
bool process(const llama_batch & /*batch*/) override {
|
||||
bool process(const common_batch & /*batch*/) override {
|
||||
// TODO: implement
|
||||
return true;
|
||||
}
|
||||
@@ -1835,7 +1817,7 @@ struct common_speculative_impl_ngram_map_k : public common_speculative_impl {
|
||||
common_ngram_map_begin(config[seq_id], prompt);
|
||||
}
|
||||
|
||||
bool process(const llama_batch & /*batch*/) override {
|
||||
bool process(const common_batch & /*batch*/) override {
|
||||
// TODO: implement
|
||||
return true;
|
||||
}
|
||||
@@ -1993,7 +1975,7 @@ struct common_speculative_impl_ngram_mod : public common_speculative_impl {
|
||||
sinfo.n_draft_last = result.size();
|
||||
}
|
||||
|
||||
bool process(const llama_batch & /*batch*/) override {
|
||||
bool process(const common_batch & /*batch*/) override {
|
||||
// TODO: implement
|
||||
return true;
|
||||
}
|
||||
@@ -2155,7 +2137,7 @@ struct common_speculative_impl_ngram_cache : public common_speculative_impl {
|
||||
}
|
||||
}
|
||||
|
||||
bool process(const llama_batch & /*batch*/) override {
|
||||
bool process(const common_batch & /*batch*/) override {
|
||||
// TODO: implement
|
||||
return true;
|
||||
}
|
||||
@@ -2181,6 +2163,9 @@ struct common_speculative_impl_ngram_cache : public common_speculative_impl {
|
||||
struct common_speculative {
|
||||
common_speculative_draft_params_vec dparams;
|
||||
|
||||
// the target context, used to convert legacy llama_batch inputs
|
||||
llama_context * ctx_tgt = nullptr;
|
||||
|
||||
// list of implementations to use and their states
|
||||
std::vector<std::unique_ptr<common_speculative_impl>> impls;
|
||||
|
||||
@@ -2726,6 +2711,7 @@ common_speculative * common_speculative_init(common_params_speculative & params,
|
||||
|
||||
common_speculative_ptr result(new common_speculative {
|
||||
/* .dparams = */ common_speculative_draft_params_vec(n_seq),
|
||||
/* .ctx_tgt = */ params.draft.ctx_tgt,
|
||||
/* .impls = */ std::move(impls),
|
||||
/* .impl_last = */ std::vector<common_speculative_impl *>(n_seq, nullptr),
|
||||
/* .synth_probs = */ {},
|
||||
@@ -2789,6 +2775,17 @@ void common_speculative_begin(common_speculative * spec, llama_seq_id seq_id, co
|
||||
}
|
||||
|
||||
bool common_speculative_process(common_speculative * spec, const llama_batch & batch) {
|
||||
if (spec == nullptr) {
|
||||
return true;
|
||||
}
|
||||
|
||||
// ngram-only setups have no target context, they do not read the batch anyway
|
||||
const common_batch tmp = spec->ctx_tgt ? common_batch_from_llama_batch(spec->ctx_tgt, batch) : common_batch();
|
||||
|
||||
return common_speculative_process(spec, tmp);
|
||||
}
|
||||
|
||||
bool common_speculative_process(common_speculative * spec, const common_batch & batch) {
|
||||
bool result = true;
|
||||
|
||||
if (spec == nullptr) {
|
||||
|
||||
@@ -77,6 +77,9 @@ common_speculative_draft_params & common_speculative_get_draft_params(common_spe
|
||||
void common_speculative_begin(common_speculative * spec, llama_seq_id seq_id, const llama_tokens & prompt);
|
||||
|
||||
// process the batch and update the internal state of the speculative context
|
||||
bool common_speculative_process(common_speculative * spec, const common_batch & batch);
|
||||
|
||||
// legacy llama_batch input, converted with common_batch_from_llama_batch()
|
||||
bool common_speculative_process(common_speculative * spec, const llama_batch & batch);
|
||||
|
||||
// generate drafts for the sequences specified with `common_speculative_get_draft_params`
|
||||
|
||||
@@ -170,6 +170,9 @@ class ModelBase:
|
||||
self.dir_model_card = dir_model # overridden in convert_lora_to_gguf.py
|
||||
self._is_nvfp4 = False
|
||||
self._is_mxfp4 = False
|
||||
self._nvfp4_global_algo: str | None = None # checkpoint-wide NVFP4 quant_algo
|
||||
self._nvfp4_layer_algo: dict[str, str | None] = {} # per-layer quant_algo, keyed by HF module path
|
||||
self._prec_a4: dict[str, bool] = {} # gguf tensor name -> can use 4-bit (A4) activations
|
||||
self._fp8_as_q8 = fp8_as_q8
|
||||
self._fp8_dequantized: set[str] = set()
|
||||
|
||||
@@ -664,6 +667,18 @@ class ModelBase:
|
||||
if bias_types:
|
||||
self._fusable_qkv_bias_layers.add(bid)
|
||||
|
||||
def _tag_prec_a4(self, hf_name: str, gguf_name: str) -> None:
|
||||
# W4A16_NVFP4 should not use 4-bit activations
|
||||
name = hf_name.removesuffix(".weight").removesuffix(".bias")
|
||||
algo = self._nvfp4_global_algo
|
||||
while name:
|
||||
if name in self._nvfp4_layer_algo:
|
||||
algo = self._nvfp4_layer_algo[name]
|
||||
break
|
||||
name = name.rpartition(".")[0]
|
||||
if algo == "W4A16_NVFP4":
|
||||
self._prec_a4[gguf_name] = False
|
||||
|
||||
def set_gguf_parameters(self):
|
||||
raise NotImplementedError("set_gguf_parameters() must be implemented in subclasses")
|
||||
|
||||
@@ -837,6 +852,7 @@ class ModelBase:
|
||||
raw, shape = self._nvfp4_pack(weight, scale)
|
||||
logger.info(f"Repacked {new_name} with shape {shape} and quantization NVFP4")
|
||||
self.gguf_writer.add_tensor(new_name, raw, raw_dtype=gguf.GGMLQuantizationType.NVFP4)
|
||||
self._tag_prec_a4(name, new_name)
|
||||
|
||||
self._write_scale_tensor(new_name.replace(".weight", ".scale"), scale2)
|
||||
self._write_scale_tensor(new_name.replace(".weight", ".input_scale"), input_scale)
|
||||
@@ -929,6 +945,7 @@ class ModelBase:
|
||||
new_name = self.map_tensor_name(merged_name)
|
||||
logger.info(f"Repacked {new_name} with shape [{len(experts)}, {shape[0]}, {shape[1]}] and quantization NVFP4")
|
||||
self.gguf_writer.add_tensor(new_name, merged, raw_dtype=gguf.GGMLQuantizationType.NVFP4)
|
||||
self._tag_prec_a4(merged_name, new_name)
|
||||
|
||||
scales.sort(key=lambda x: x[0])
|
||||
self._write_scales_tensor(new_name.replace(".weight", ".scale"), [s[1] for s in scales])
|
||||
@@ -971,6 +988,9 @@ class ModelBase:
|
||||
and bool(quant_groups)
|
||||
and all(g.get("format") == "nvfp4-pack-quantized" for g in quant_groups.values() if isinstance(g, dict))
|
||||
)
|
||||
|
||||
self._nvfp4_global_algo = quant_algo
|
||||
|
||||
if quant_algo != "NVFP4":
|
||||
if nvfp4_compressed_tensors:
|
||||
quant_algo = "NVFP4"
|
||||
@@ -980,6 +1000,22 @@ class ModelBase:
|
||||
self._is_nvfp4 = quant_algo in ("NVFP4", "W4A16_NVFP4")
|
||||
self._is_mxfp4 = quant_method == "mxfp4"
|
||||
|
||||
# Per-tensor NVFP4 precision.
|
||||
self._nvfp4_layer_algo = {}
|
||||
if quant_layers:
|
||||
# store all possible module paths and assert if a quantized layer is not in the model
|
||||
modules: set[str] = set()
|
||||
for name in self.model_tensors:
|
||||
while name := name.rpartition(".")[0]:
|
||||
modules.add(name)
|
||||
|
||||
for layer_name, entry in quant_layers.items():
|
||||
if not isinstance(entry, dict):
|
||||
continue
|
||||
if titem := self.filter_tensors((layer_name, lambda: torch.empty(0))):
|
||||
assert titem[0] in modules, f"quantized_layers entry {layer_name!r} is not in the model tensors"
|
||||
self._nvfp4_layer_algo[titem[0]] = entry.get("quant_algo")
|
||||
|
||||
# NVFP4 weights are repacked and written directly to gguf_writer.
|
||||
# This must run before dequant_model so NVFP4 tensors are removed
|
||||
# from model_tensors, leaving only non-NVFP4 (e.g. FP8) for dequant.
|
||||
@@ -1185,6 +1221,12 @@ class ModelBase:
|
||||
logger.info("Set model quantization version")
|
||||
self.gguf_writer.add_quantization_version(gguf.GGML_QUANT_VERSION)
|
||||
|
||||
if self._prec_a4:
|
||||
names = sorted(self._prec_a4.keys())
|
||||
values = [self._prec_a4[n] for n in names]
|
||||
logger.info(f"Set prec_a4 metadata for {len(names)} tensor(s)")
|
||||
self.gguf_writer.add_tensor_extra_prec_a4(names, values)
|
||||
|
||||
def write_vocab(self):
|
||||
raise NotImplementedError("write_vocab() must be implemented in subclasses")
|
||||
|
||||
|
||||
+1
-1
@@ -341,7 +341,7 @@ class NomicBertModel(BertModel):
|
||||
else:
|
||||
raise ValueError(f"unrecognized parameters: n_positions={npos}, max_trained_positions={mtp}")
|
||||
|
||||
assert self.hparams["activation_function"] == "gelu" if self.is_moe else "swiglu"
|
||||
assert self.hparams["activation_function"] == ("gelu" if self.is_moe else "swiglu")
|
||||
|
||||
# this doesn't do anything in the HF version
|
||||
assert self.hparams["causal"] is False
|
||||
|
||||
+31
-3
@@ -17,6 +17,7 @@ from .base import MmprojModel, ModelBase, TextModel, gguf, logger
|
||||
@ModelBase.example("XiaomiMiMo/MiMo-V2.5")
|
||||
class MimoV2Model(TextModel):
|
||||
model_arch = gguf.MODEL_ARCH.MIMO2
|
||||
supports_mtp_export = True
|
||||
|
||||
# MiMo V2-Flash, V2.5 and V2.5-Pro all ship 3 trained MTP layers under model.mtp.layers.{0,1,2}.
|
||||
# The HF config does not expose the count, so it's hardcoded to match the count found in the safetensors.
|
||||
@@ -25,6 +26,8 @@ class MimoV2Model(TextModel):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
if self.no_mtp:
|
||||
self._n_nextn = 0
|
||||
self.block_count = self.hparams["num_hidden_layers"] + self._n_nextn
|
||||
self.tensor_map = gguf.get_tensor_name_map(self.model_arch, self.block_count)
|
||||
|
||||
@@ -101,7 +104,7 @@ class MimoV2Model(TextModel):
|
||||
qkv_overrides: dict[str, tuple[Callable, Callable, int]] = {}
|
||||
qc = self.hparams.get("quantization_config")
|
||||
if isinstance(qc, dict) and qc.get("quant_method") == "fp8":
|
||||
pat = re.compile(r"^model\.layers\.(\d+)\.self_attn\.qkv_proj\.weight_scale_inv$")
|
||||
pat = re.compile(r"^model\.(mtp\.)?layers\.(\d+)\.self_attn\.qkv_proj\.weight_scale_inv$")
|
||||
for name in list(self.model_tensors.keys()):
|
||||
m = pat.match(name)
|
||||
if not m:
|
||||
@@ -109,10 +112,13 @@ class MimoV2Model(TextModel):
|
||||
weight_name = name.removesuffix("_scale_inv")
|
||||
if weight_name not in self.model_tensors:
|
||||
continue
|
||||
bid = int(m.group(2))
|
||||
if m.group(1) is not None:
|
||||
bid += self.hparams["num_hidden_layers"]
|
||||
qkv_overrides[weight_name] = (
|
||||
self.model_tensors[weight_name],
|
||||
self.model_tensors[name],
|
||||
int(m.group(1)),
|
||||
bid,
|
||||
)
|
||||
|
||||
super().dequant_model()
|
||||
@@ -165,7 +171,8 @@ class MimoV2Model(TextModel):
|
||||
if v_scale is not None:
|
||||
self.gguf_writer.add_attn_value_scale(float(v_scale))
|
||||
|
||||
self.gguf_writer.add_nextn_predict_layers(self._n_nextn)
|
||||
if self._n_nextn > 0:
|
||||
self.gguf_writer.add_nextn_predict_layers(self._n_nextn)
|
||||
|
||||
_MXFP4_EXPERT_RE = re.compile(
|
||||
r"^model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.(gate|up|down)_proj\.weight$"
|
||||
@@ -251,11 +258,32 @@ class MimoV2Model(TextModel):
|
||||
def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:
|
||||
name, gen = item
|
||||
|
||||
is_mtp = name.startswith("model.mtp.layers.")
|
||||
if is_mtp and cls.no_mtp:
|
||||
return None
|
||||
if cls.mtp_only and not is_mtp and name not in (
|
||||
"model.embed_tokens.weight", "model.norm.weight", "lm_head.weight",
|
||||
):
|
||||
return None
|
||||
|
||||
if "attention_sink" in name and not name.endswith(".weight"):
|
||||
name += ".weight"
|
||||
|
||||
return super().filter_tensors((name, gen))
|
||||
|
||||
def prepare_metadata(self, vocab_only: bool):
|
||||
from_dir = self.fname_out.is_dir()
|
||||
super().prepare_metadata(vocab_only=vocab_only)
|
||||
|
||||
if not self.mtp_only or not from_dir:
|
||||
return
|
||||
|
||||
output_type: str = self.ftype.name.partition("_")[2]
|
||||
fname_default: str = gguf.naming_convention(
|
||||
self.metadata.name, self.metadata.basename, self.metadata.finetune,
|
||||
self.metadata.version, size_label=None, output_type=output_type, model_type=None)
|
||||
self.fname_out = self.fname_out.parent / f"mtp-{fname_default}.gguf"
|
||||
|
||||
def modify_tensors(self, data_torch, name, bid):
|
||||
# Remap MTP/NextN tensors to additional layer slots so the standard tensor map handles them.
|
||||
# HF: model.mtp.layers.{i}.foo -> model.layers.{n_layer_text + i}.foo
|
||||
|
||||
@@ -424,12 +424,6 @@ class NemotronHModel(GraniteHybridModel):
|
||||
special_vocab = gguf.SpecialVocab(self.dir_model, load_merges=True)
|
||||
special_vocab.add_to_gguf(self.gguf_writer)
|
||||
|
||||
# The tokenizer _does_ add a BOS token (via post_processor type
|
||||
# TemplateProcessing) but does not set add_bos_token to true in the
|
||||
# config, so we need to explicitly override it here.
|
||||
if not self.is_moe:
|
||||
self.gguf_writer.add_add_bos_token(True)
|
||||
|
||||
_MTP_SPECIAL_RENAMES = {
|
||||
"mtp.layers.0.enorm.weight": "model.layers.{bid}.enorm.weight",
|
||||
"mtp.layers.0.hnorm.weight": "model.layers.{bid}.hnorm.weight",
|
||||
|
||||
@@ -154,6 +154,21 @@ class Plamo2Model(TextModel):
|
||||
class Plamo3Model(TextModel):
|
||||
model_arch = gguf.MODEL_ARCH.PLAMO3
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
# PLaMo-3 builds rope_parameters from flat config keys at runtime; mirror the YaRN settings for GGUF.
|
||||
rope_scaling_factor = self.hparams.get("rope_scaling_factor", 1)
|
||||
if rope_scaling_factor != 1 and "rope_type" not in self.rope_parameters:
|
||||
self.rope_parameters.update({
|
||||
"rope_type": "yarn",
|
||||
"factor": float(rope_scaling_factor),
|
||||
"original_max_position_embeddings": int(self.hparams["initial_context_length"]),
|
||||
"beta_fast": 32.0,
|
||||
"beta_slow": 1.0,
|
||||
"truncate": False,
|
||||
})
|
||||
|
||||
def set_vocab(self):
|
||||
self._set_vocab_plamo()
|
||||
|
||||
|
||||
+12
-4
@@ -52,6 +52,10 @@ The packages for FP32 and FP16 would have different accuracy and performance on
|
||||
|
||||
## News
|
||||
|
||||
- 2026.09
|
||||
- Update the CI build environment for oneAPI 2026.1 (unified oneAPI Toolkit). oneDNN is removed from the Deep Learning Essentials package in 2026.0, so the CI now uses the oneAPI Toolkit installer which still includes oneDNN.
|
||||
- oneAPI 2026.1 improves the SYCL build performance: measured with the same code on Arc B570, prompt processing 1331 vs 434 t/s (3.1x) vs the 2025.3-based release build.
|
||||
|
||||
- 2026.04-05
|
||||
- Optimize mul_mat by reorder feature for data type: Q4_K, Q5_K, Q6_K, Q8_0.
|
||||
- Fused MoE.
|
||||
@@ -257,7 +261,7 @@ Platform #0: Intel(R) OpenCL HD Graphics
|
||||
`-- Device #0: Intel(R) Iris(R) Xe Graphics [0x9a49]
|
||||
```
|
||||
|
||||
2. **Install Intel® oneAPI Base toolkit**
|
||||
2. **Install Intel® oneAPI Toolkit**
|
||||
|
||||
SYCL backend depends on:
|
||||
- Intel® oneAPI DPC++/C++ compiler/running-time.
|
||||
@@ -267,11 +271,11 @@ SYCL backend depends on:
|
||||
|
||||
- **For Intel GPU**
|
||||
|
||||
All above are included in both **Intel® oneAPI Base toolkit** and **Intel® Deep Learning Essentials** packages.
|
||||
With the 2026.0 release, the Intel® oneAPI Base toolkit and the HPC toolkit are combined into the **Intel® oneAPI Toolkit**, and **oneDNN is removed from the Intel® Deep Learning Essentials** package (oneDNN is distributed separately since then). The **Intel® oneAPI Toolkit** includes oneDNN until 2027.0.
|
||||
|
||||
It's recommended to install **Intel® Deep Learning Essentials** which only provides the necessary libraries with less size.
|
||||
It's recommended to install the **Intel® oneAPI Toolkit**.
|
||||
|
||||
The **Intel® oneAPI Base toolkit** and **Intel® Deep Learning Essentials** can be obtained from the official [Intel® oneAPI Base Toolkit](https://www.intel.com/content/www/us/en/developer/tools/oneapi/base-toolkit.html) page.
|
||||
The **Intel® oneAPI Toolkit** can be obtained from the official [Intel® oneAPI Toolkit](https://www.intel.com/content/www/us/en/developer/tools/oneapi/base-toolkit-download.html) page.
|
||||
|
||||
Please follow the instructions for downloading and installing the Toolkit for Linux, and preferably keep the default installation values unchanged, notably the installation path *(`/opt/intel/oneapi` by default)*.
|
||||
|
||||
@@ -281,6 +285,7 @@ Upon a successful installation, SYCL is enabled for the available Intel devices,
|
||||
|
||||
|Verified release|
|
||||
|-|
|
||||
|2026.1 |
|
||||
|2025.3.3 |
|
||||
|2025.2.1|
|
||||
|2025.1|
|
||||
@@ -811,6 +816,9 @@ User can use the device management in [docs/multi-gpu.md](https://github.com/ggm
|
||||
| GGML_SYCL_MKL_FA_DIAG | 0 (default) or 1 | Enable output fingerprinting for MKL flash attention. Dumps the first 64 float output values for the first 6 FA calls with n_kv ≥ 1024, labeled with kernel type (MKL/TILE/VEC) for cross-kernel comparison. |
|
||||
| GGML_SYCL_ENABLE_FUSION | 0 or 1 (default) | Enable fused-kernel dispatch in graph compute. Unsupported types and layouts fall back to the standalone op kernels. See `ggml_sycl_can_fuse()`. |
|
||||
| GGML_SYCL_ENABLE_ESIMD | 0 or 1 (default)| Enable ESIMD kernels when available. |
|
||||
| GGML_SYCL_SPARSE_FA | 0 (default) or 1 | Enable Sparse Flash-attention.|
|
||||
| GGML_SYCL_SPARSE_FA_DEBUG | 0 (default) or 1 | Enable to debug for Sparse Flash-attention.|
|
||||
| GGML_SYCL_SPARSE_FA_MARGIN | [0,..] default:256 | Set the margin value for Sparse Flash-attention.|
|
||||
| ZES_ENABLE_SYSMAN | 0 (default) or 1 | Support to get free memory of GPU by sycl::aspect::ext_intel_free_memory.<br>Recommended to use when --split-mode = layer |
|
||||
| UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS | 0 (default) or 1 | Allow SYCL/Unified Runtime Level Zero device allocations larger than 4 GiB. llama.cpp's direct Level Zero allocation path requests the relaxed maximum-size limit itself when GGML_SYCL_ENABLE_LEVEL_ZERO=1. |
|
||||
| GGML_SYCL_USM_SYSTEM | 0 (default) or 1 | Enable experimental support for [USM system allocations](https://github.khronos.org/SYCL_Reference/iface/usm_basic_concept.html#system-allocations) for large GPU buffers. This requires enough host memory for model weights and caches, an Intel Xe2+ GPU such as BMG or newer and supported on Linux only, with CONFIG_DRM_XE_GPUSVM enabled. |
|
||||
|
||||
@@ -77,6 +77,8 @@
|
||||
|
||||
{ "name": "arm64-android-snapdragon-debug" , "inherits": [ "base", "arm64-android-snapdragon", "debug" ] },
|
||||
{ "name": "arm64-android-snapdragon-release", "inherits": [ "base", "arm64-android-snapdragon", "release" ] },
|
||||
{ "name": "arm64-android-snapdragon-relwithdebinfo", "inherits": [ "arm64-android-snapdragon-release" ],
|
||||
"cacheVariables": { "GGML_HEXAGON_HTP_BUILD_TYPE": "RelWithDebInfo" } },
|
||||
|
||||
{ "name": "arm64-windows-snapdragon-debug" , "inherits": [ "base", "arm64-windows-snapdragon", "debug" ] },
|
||||
{ "name": "arm64-windows-snapdragon-release", "inherits": [ "base", "arm64-windows-snapdragon", "release" ] },
|
||||
|
||||
+17
-2
@@ -181,6 +181,14 @@ cmake -B build -DGGML_CUDA=ON
|
||||
cmake --build build --config Release
|
||||
```
|
||||
|
||||
Note that this also builds the CPU backend by default. On Windows on ARM, MSVC's
|
||||
support for the ARM NEON intrinsics used by the CPU backend may be incomplete, so
|
||||
a CUDA build produced entirely with MSVC might have a slower CPU backend. If CPU
|
||||
performance matters, try following the split build used in our release workflow
|
||||
([.github/workflows/release.yml](../.github/workflows/release.yml)): the CPU backend
|
||||
is built with clang (`cmake/arm64-windows-llvm.cmake`) and the CUDA backend with MSVC
|
||||
(`cmake/arm64-windows-msvc-cuda.cmake`), and the artifacts are merged afterwards.
|
||||
|
||||
### Non-Native Builds
|
||||
|
||||
By default llama.cpp will be built for the hardware that is connected to the system at that time.
|
||||
@@ -282,6 +290,13 @@ Consider setting `CUDA_SCALE_LAUNCH_QUEUES=4x`, which increases the CUDA command
|
||||
Override default, speed-optimized compute types for cuBLAS matrix multiplications.
|
||||
Legal values: `auto`, `f16`, `fp16`, `bf16`, `f32`, `fp32`.
|
||||
|
||||
#### GGML_CUDA_MMQ_PREC
|
||||
|
||||
Override the activation precision that the model requests for NVFP4 and MXFP4 matrix multiplications.
|
||||
Currently supported values: `auto`, `q8`, `q4`.
|
||||
|
||||
NVFP4 and MXFP4 layers marked as W4A16 request 8-bit activations, so on Blackwell those layers run through the W4A8 path instead of the native W4A4 path. Set `q4` to keep the native W4A4 path for faster prompt processing at the cost of accuracy, or `q8` to use the W4A8 path for every layer, `auto` uses per-tensor prec metadata (this is the same behavior as when the environment variable is not set).
|
||||
|
||||
### Unified Memory
|
||||
|
||||
The environment variable `GGML_CUDA_ENABLE_UNIFIED_MEMORY=1` can be used to enable unified memory in Linux. This allows swapping to system RAM instead of crashing when the GPU VRAM is exhausted. In Windows this setting is available in the NVIDIA control panel as `System Memory Fallback`.
|
||||
@@ -323,11 +338,11 @@ cmake --build build --config Release
|
||||
By default, all supported compute capabilities are enabled. To customize this behavior, you can specify the `MUSA_ARCHITECTURES` option in the CMake command:
|
||||
|
||||
```bash
|
||||
cmake -B build -DGGML_MUSA=ON -DMUSA_ARCHITECTURES="21"
|
||||
cmake -B build -DGGML_MUSA=ON -DMUSA_ARCHITECTURES="31"
|
||||
cmake --build build --config Release
|
||||
```
|
||||
|
||||
This configuration enables only compute capability `2.1` (MTT S80) during compilation, which can help reduce compilation time.
|
||||
This configuration enables only compute capability `3.1` (MTT S5000) during compilation, which can help reduce compilation time.
|
||||
|
||||
#### Compilation options
|
||||
|
||||
|
||||
+1
-1
@@ -123,7 +123,7 @@ You may want to pass in some different `ARGS`, depending on the MUSA environment
|
||||
|
||||
The defaults are:
|
||||
|
||||
- `MUSA_VERSION` set to `rc4.3.0`
|
||||
- the base image is the MUSA 5.2.0 image from the Moore Threads registry
|
||||
|
||||
The resulting images, are essentially the same as the non-MUSA images:
|
||||
|
||||
|
||||
+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
@@ -228,7 +228,6 @@ int main(int argc, char ** argv) {
|
||||
common_batch_add(batch_tgt, draft[i], n_past + i, { seq_id }, true);
|
||||
}
|
||||
|
||||
//LOG_DBG("target batch: %s\n", string_from(ctx_tgt, batch_tgt).c_str());
|
||||
|
||||
llama_decode(ctx_tgt, batch_tgt);
|
||||
}
|
||||
|
||||
+1
-1
@@ -5,7 +5,7 @@ project("ggml" C CXX ASM)
|
||||
### GGML Version
|
||||
set(GGML_VERSION_MAJOR 0)
|
||||
set(GGML_VERSION_MINOR 25)
|
||||
set(GGML_VERSION_PATCH 1)
|
||||
set(GGML_VERSION_PATCH 3)
|
||||
set(GGML_VERSION_BASE "${GGML_VERSION_MAJOR}.${GGML_VERSION_MINOR}.${GGML_VERSION_PATCH}")
|
||||
|
||||
list(APPEND CMAKE_MODULE_PATH "${CMAKE_CURRENT_SOURCE_DIR}/cmake/")
|
||||
|
||||
@@ -28,6 +28,7 @@ endfunction()
|
||||
function(ggml_get_system_arch)
|
||||
if (CMAKE_OSX_ARCHITECTURES STREQUAL "arm64" OR
|
||||
CMAKE_GENERATOR_PLATFORM_LWR STREQUAL "arm64" OR
|
||||
(CMAKE_GENERATOR_PLATFORM_LWR STREQUAL "arm64ec" AND MSVC AND NOT CMAKE_C_COMPILER_ID STREQUAL "Clang") OR
|
||||
(NOT CMAKE_OSX_ARCHITECTURES AND NOT CMAKE_GENERATOR_PLATFORM_LWR AND
|
||||
CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm.*|ARM64)$"))
|
||||
set(GGML_SYSTEM_ARCH "ARM" PARENT_SCOPE)
|
||||
|
||||
@@ -31,8 +31,6 @@ function(ggml_add_cpu_backend_variant_impl tag_name)
|
||||
ggml-cpu/ggml-cpu.cpp
|
||||
ggml-cpu/repack.cpp
|
||||
ggml-cpu/repack.h
|
||||
ggml-cpu/iqp.cpp
|
||||
ggml-cpu/iqp.h
|
||||
ggml-cpu/hbm.cpp
|
||||
ggml-cpu/hbm.h
|
||||
ggml-cpu/quants.c
|
||||
@@ -43,6 +41,10 @@ function(ggml_add_cpu_backend_variant_impl tag_name)
|
||||
ggml-cpu/amx/amx.h
|
||||
ggml-cpu/amx/mmq.cpp
|
||||
ggml-cpu/amx/mmq.h
|
||||
ggml-cpu/tiled/tiled.cpp
|
||||
ggml-cpu/tiled/tiled.h
|
||||
ggml-cpu/tiled/tiled-kernel.cpp
|
||||
ggml-cpu/tiled/tiled-kernel.h
|
||||
ggml-cpu/ggml-cpu-impl.h
|
||||
ggml-cpu/common.h
|
||||
ggml-cpu/binary-ops.h
|
||||
@@ -104,54 +106,79 @@ function(ggml_add_cpu_backend_variant_impl tag_name)
|
||||
ggml-cpu/arch/arm/repack.cpp
|
||||
)
|
||||
|
||||
# MSVC proper (not clang-cl) has no -mcpu/-march for ARM: features can only be
|
||||
# selected by probing the host at configure time and forwarding -D__ARM_FEATURE_*
|
||||
if (MSVC AND NOT CMAKE_C_COMPILER_ID STREQUAL "Clang")
|
||||
message(FATAL_ERROR "MSVC is not supported for ARM, use clang")
|
||||
set(GGML_ARM_MSVC ON)
|
||||
else()
|
||||
check_cxx_compiler_flag(-mfp16-format=ieee GGML_COMPILER_SUPPORTS_FP16_FORMAT_I3E)
|
||||
if (NOT "${GGML_COMPILER_SUPPORTS_FP16_FORMAT_I3E}" STREQUAL "")
|
||||
list(APPEND ARCH_FLAGS -mfp16-format=ieee)
|
||||
set(GGML_ARM_MSVC OFF)
|
||||
endif()
|
||||
|
||||
if (GGML_ARM_MSVC AND NOT GGML_NATIVE)
|
||||
message(FATAL_ERROR "MSVC on ARM requires GGML_NATIVE=ON, use clang for GGML_CPU_ARM_ARCH / GGML_CPU_ALL_VARIANTS")
|
||||
else()
|
||||
if (NOT GGML_ARM_MSVC)
|
||||
check_cxx_compiler_flag(-mfp16-format=ieee GGML_COMPILER_SUPPORTS_FP16_FORMAT_I3E)
|
||||
if (NOT "${GGML_COMPILER_SUPPORTS_FP16_FORMAT_I3E}" STREQUAL "")
|
||||
list(APPEND ARCH_FLAGS -mfp16-format=ieee)
|
||||
endif()
|
||||
endif()
|
||||
|
||||
if (GGML_NATIVE)
|
||||
# -mcpu=native does not always enable all the features in some compilers,
|
||||
# so we check for them manually and enable them if available
|
||||
|
||||
execute_process(
|
||||
COMMAND ${CMAKE_C_COMPILER} -mcpu=native -E -v -
|
||||
INPUT_FILE "/dev/null"
|
||||
OUTPUT_QUIET
|
||||
ERROR_VARIABLE ARM_MCPU
|
||||
RESULT_VARIABLE ARM_MCPU_RESULT
|
||||
)
|
||||
if (NOT ARM_MCPU_RESULT)
|
||||
string(REGEX MATCH "-mcpu=[^ ']+" ARM_MCPU_FLAG "${ARM_MCPU}")
|
||||
string(REGEX MATCH "-march=[^ ']+" ARM_MARCH_FLAG "${ARM_MCPU}")
|
||||
# cl.exe cannot be queried this way: no /dev/null, and it rejects the flags.
|
||||
# ARM_NATIVE_FLAG stays empty for MSVC and is never appended below.
|
||||
if (NOT GGML_ARM_MSVC)
|
||||
execute_process(
|
||||
COMMAND ${CMAKE_C_COMPILER} -mcpu=native -E -v -
|
||||
INPUT_FILE "/dev/null"
|
||||
OUTPUT_QUIET
|
||||
ERROR_VARIABLE ARM_MCPU
|
||||
RESULT_VARIABLE ARM_MCPU_RESULT
|
||||
)
|
||||
if (NOT ARM_MCPU_RESULT)
|
||||
string(REGEX MATCH "-mcpu=[^ ']+" ARM_MCPU_FLAG "${ARM_MCPU}")
|
||||
string(REGEX MATCH "-march=[^ ']+" ARM_MARCH_FLAG "${ARM_MCPU}")
|
||||
|
||||
# on some old GCC we need to read -march=
|
||||
if (ARM_MARCH_FLAG AND NOT "${ARM_MARCH_FLAG}" STREQUAL "-march=native")
|
||||
set(ARM_NATIVE_FLAG "${ARM_MARCH_FLAG}")
|
||||
elseif(ARM_MCPU_FLAG AND NOT "${ARM_MCPU_FLAG}" STREQUAL "-mcpu=native")
|
||||
set(ARM_NATIVE_FLAG "${ARM_MCPU_FLAG}")
|
||||
# on some old GCC we need to read -march=
|
||||
if (ARM_MARCH_FLAG AND NOT "${ARM_MARCH_FLAG}" STREQUAL "-march=native")
|
||||
set(ARM_NATIVE_FLAG "${ARM_MARCH_FLAG}")
|
||||
elseif(ARM_MCPU_FLAG AND NOT "${ARM_MCPU_FLAG}" STREQUAL "-mcpu=native")
|
||||
set(ARM_NATIVE_FLAG "${ARM_MCPU_FLAG}")
|
||||
endif()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
if ("${ARM_NATIVE_FLAG}" STREQUAL "")
|
||||
set(ARM_NATIVE_FLAG -mcpu=native)
|
||||
message(WARNING "ARM -march/-mcpu not found, -mcpu=native will be used")
|
||||
else()
|
||||
message(STATUS "ARM detected flags: ${ARM_NATIVE_FLAG}")
|
||||
if ("${ARM_NATIVE_FLAG}" STREQUAL "")
|
||||
set(ARM_NATIVE_FLAG -mcpu=native)
|
||||
message(WARNING "ARM -march/-mcpu not found, -mcpu=native will be used")
|
||||
else()
|
||||
message(STATUS "ARM detected flags: ${ARM_NATIVE_FLAG}")
|
||||
endif()
|
||||
endif()
|
||||
|
||||
include(CheckCXXSourceRuns)
|
||||
|
||||
macro(check_arm_feature tag feature code)
|
||||
set(CMAKE_REQUIRED_FLAGS_SAVE ${CMAKE_REQUIRED_FLAGS})
|
||||
set(CMAKE_REQUIRED_FLAGS "${ARM_NATIVE_FLAG}+${tag}")
|
||||
if (GGML_ARM_MSVC)
|
||||
set(ARM_PROBE_FLAG "")
|
||||
set(ARM_PROBE_NOFLAG "")
|
||||
else()
|
||||
set(ARM_PROBE_FLAG "${ARM_NATIVE_FLAG}+${tag}")
|
||||
set(ARM_PROBE_NOFLAG "${ARM_NATIVE_FLAG}+no${tag}")
|
||||
endif()
|
||||
set(CMAKE_REQUIRED_FLAGS "${ARM_PROBE_FLAG}")
|
||||
check_cxx_source_runs("${code}" GGML_MACHINE_SUPPORTS_${tag})
|
||||
if (GGML_MACHINE_SUPPORTS_${tag})
|
||||
set(ARM_NATIVE_FLAG_FIX "${ARM_NATIVE_FLAG_FIX}+${tag}")
|
||||
if (GGML_ARM_MSVC)
|
||||
# forward runtime probing decisions for MSVC
|
||||
list(APPEND ARCH_FLAGS "-D__ARM_FEATURE_${feature}=1")
|
||||
endif()
|
||||
else()
|
||||
set(CMAKE_REQUIRED_FLAGS "${ARM_NATIVE_FLAG}+no${tag}")
|
||||
set(CMAKE_REQUIRED_FLAGS "${ARM_PROBE_NOFLAG}")
|
||||
check_cxx_source_compiles("int main() { return 0; }" GGML_MACHINE_SUPPORTS_no${tag})
|
||||
if (GGML_MACHINE_SUPPORTS_no${tag})
|
||||
set(ARM_NATIVE_FLAG_FIX "${ARM_NATIVE_FLAG_FIX}+no${tag}")
|
||||
@@ -161,12 +188,24 @@ function(ggml_add_cpu_backend_variant_impl tag_name)
|
||||
set(CMAKE_REQUIRED_FLAGS ${CMAKE_REQUIRED_FLAGS_SAVE})
|
||||
endmacro()
|
||||
|
||||
check_arm_feature(dotprod DOTPROD "#include <arm_neon.h>\nint main() { int8x16_t _a, _b; volatile int32x4_t _s = vdotq_s32(_s, _a, _b); return 0; }")
|
||||
check_arm_feature(i8mm MATMUL_INT8 "#include <arm_neon.h>\nint main() { int8x16_t _a, _b; volatile int32x4_t _s = vmmlaq_s32(_s, _a, _b); return 0; }")
|
||||
check_arm_feature(sve SVE "#include <arm_sve.h>\nint main() { svfloat32_t _a, _b; volatile svfloat32_t _c = svadd_f32_z(svptrue_b8(), _a, _b); return 0; }")
|
||||
if (GGML_ARM_MSVC)
|
||||
# drop volatile
|
||||
check_arm_feature(dotprod DOTPROD "#include <arm_neon.h>\nint main() { int8x16_t _a = vdupq_n_s8(0), _b = vdupq_n_s8(0); int32x4_t _s = vdupq_n_s32(0); _s = vdotq_s32(_s, _a, _b); return 0; }")
|
||||
check_arm_feature(i8mm MATMUL_INT8 "#include <arm_neon.h>\nint main() { int8x16_t _a = vdupq_n_s8(0), _b = vdupq_n_s8(0); int32x4_t _s = vdupq_n_s32(0); _s = vmmlaq_s32(_s, _a, _b); return 0; }")
|
||||
check_arm_feature(sve SVE "#include <arm_sve.h>\nint main() { const svbool_t pg = svptrue_b8(); const svint8_t a = svdup_n_s8(0), b = svdup_n_s8(0); svint32_t s = svdup_n_s32(0); s = svdot_s32(s, a, b); (void) pg; (void) s; return 0; }")
|
||||
# FMA is mandatory on ARMv8-A but MSVC never defines __ARM_FEATURE_FMA,
|
||||
# so probe it like the others and forward the define
|
||||
check_arm_feature(fma FMA "#include <arm_neon.h>\nint main() { float32x4_t _a = vdupq_n_f32(0), _b = vdupq_n_f32(0), _c = vdupq_n_f32(0); _a = vfmaq_f32(_a, _b, _c); return 0; }")
|
||||
else()
|
||||
check_arm_feature(dotprod DOTPROD "#include <arm_neon.h>\nint main() { int8x16_t _a, _b; volatile int32x4_t _s = vdotq_s32(_s, _a, _b); return 0; }")
|
||||
check_arm_feature(i8mm MATMUL_INT8 "#include <arm_neon.h>\nint main() { int8x16_t _a, _b; volatile int32x4_t _s = vmmlaq_s32(_s, _a, _b); return 0; }")
|
||||
check_arm_feature(sve SVE "#include <arm_sve.h>\nint main() { svfloat32_t _a, _b; volatile svfloat32_t _c = svadd_f32_z(svptrue_b8(), _a, _b); return 0; }")
|
||||
endif()
|
||||
check_arm_feature(sme SME "#include <arm_sme.h>\n__arm_locally_streaming int main() { __asm__ volatile(\"smstart; smstop;\"); return 0; }")
|
||||
|
||||
list(APPEND ARCH_FLAGS "${ARM_NATIVE_FLAG}${ARM_NATIVE_FLAG_FIX}")
|
||||
if (NOT GGML_ARM_MSVC)
|
||||
list(APPEND ARCH_FLAGS "${ARM_NATIVE_FLAG}${ARM_NATIVE_FLAG_FIX}")
|
||||
endif()
|
||||
else()
|
||||
if (GGML_CPU_ARM_ARCH)
|
||||
list(APPEND ARCH_FLAGS -march=${GGML_CPU_ARM_ARCH})
|
||||
|
||||
@@ -75,7 +75,7 @@
|
||||
#define ggml_gemm_mxfp4_8x8_q8_0_generic ggml_gemm_mxfp4_8x8_q8_0
|
||||
#define ggml_gemm_q8_0_4x4_q8_0_generic ggml_gemm_q8_0_4x4_q8_0
|
||||
#define ggml_gemm_q8_0_4x8_q8_0_generic ggml_gemm_q8_0_4x8_q8_0
|
||||
#elif defined(__aarch64__) || defined(__arm__) || defined(_M_ARM) || defined(_M_ARM64)
|
||||
#elif defined(__aarch64__) || defined(__arm__) || defined(_M_ARM) || defined(_M_ARM64) || defined(_M_ARM64EC)
|
||||
// repack.cpp
|
||||
#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4
|
||||
#define ggml_quantize_mat_q8_K_4x8_generic ggml_quantize_mat_q8_K_4x8
|
||||
|
||||
@@ -875,23 +875,23 @@ void ggml_vec_dot_nvfp4_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const vo
|
||||
const int8x8_t q8_3_lo = vld1_s8(y[2*ib+1].qs + 16);
|
||||
const int8x8_t q8_3_hi = vld1_s8(y[2*ib+1].qs + 24);
|
||||
|
||||
const int32x4_t sumi = (int32x4_t){
|
||||
const int32x4_t sumi = vld1q_s32(((const int32_t[4]) {
|
||||
vaddvq_s32(ggml_nvfp4_dot8(q4_0_lo, q8_0_lo, q4_0_hi, q8_0_hi)),
|
||||
vaddvq_s32(ggml_nvfp4_dot8(q4_1_lo, q8_1_lo, q4_1_hi, q8_1_hi)),
|
||||
vaddvq_s32(ggml_nvfp4_dot8(q4_2_lo, q8_2_lo, q4_2_hi, q8_2_hi)),
|
||||
vaddvq_s32(ggml_nvfp4_dot8(q4_3_lo, q8_3_lo, q4_3_hi, q8_3_hi)),
|
||||
};
|
||||
}));
|
||||
#endif
|
||||
|
||||
const float dy0 = GGML_CPU_FP16_TO_FP32(y[2*ib].d);
|
||||
const float dy1 = GGML_CPU_FP16_TO_FP32(y[2*ib+1].d);
|
||||
const float32x4_t nvsc = {
|
||||
const float32x4_t nvsc = vld1q_f32(((const float[4]) {
|
||||
GGML_CPU_UE4M3_TO_FP32(x[ib].d[0]),
|
||||
GGML_CPU_UE4M3_TO_FP32(x[ib].d[1]),
|
||||
GGML_CPU_UE4M3_TO_FP32(x[ib].d[2]),
|
||||
GGML_CPU_UE4M3_TO_FP32(x[ib].d[3])
|
||||
};
|
||||
const float32x4_t scales = vmulq_f32(nvsc, (float32x4_t){dy0, dy0, dy1, dy1});
|
||||
}));
|
||||
const float32x4_t scales = vmulq_f32(nvsc, vld1q_f32(((const float[4]) {dy0, dy0, dy1, dy1})));
|
||||
|
||||
acc = vfmaq_f32(acc, vcvtq_f32_s32(sumi), scales);
|
||||
}
|
||||
@@ -2591,7 +2591,7 @@ void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
|
||||
memcpy(scales_mins, x0->scales, 12);
|
||||
const uint32_t mins_0_3 = scales_mins[1] & kmask1;
|
||||
const uint32_t mins_4_7 = ((scales_mins[2] >> 4) & kmask2) | (((scales_mins[1] >> 6) & kmask3) << 4);
|
||||
const uint32x2_t mins = {mins_0_3, mins_4_7};
|
||||
const uint32x2_t mins = vcreate_u32((uint64_t) mins_0_3 | ((uint64_t) mins_4_7 << 32));
|
||||
x0_mins = vreinterpretq_s16_u16(vmovl_u8(vreinterpret_u8_u32(mins)));
|
||||
uint32_t scales[2];
|
||||
scales[0] = scales_mins[0] & kmask1; // scales 0~3
|
||||
@@ -2603,7 +2603,7 @@ void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
|
||||
memcpy(scales_mins, x1->scales, 12);
|
||||
const uint32_t mins_0_3 = scales_mins[1] & kmask1;
|
||||
const uint32_t mins_4_7 = ((scales_mins[2] >> 4) & kmask2) | (((scales_mins[1] >> 6) & kmask3) << 4);
|
||||
const uint32x2_t mins = {mins_0_3, mins_4_7};
|
||||
const uint32x2_t mins = vcreate_u32((uint64_t) mins_0_3 | ((uint64_t) mins_4_7 << 32));
|
||||
x1_mins = vreinterpretq_s16_u16(vmovl_u8(vreinterpret_u8_u32(mins)));
|
||||
uint32_t scales[2];
|
||||
scales[0] = scales_mins[0] & kmask1; // scales 0~3
|
||||
@@ -2611,7 +2611,7 @@ void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
|
||||
memcpy(x1_scales, scales, 8);
|
||||
}
|
||||
|
||||
int32x4_t visum = {0};
|
||||
int32x4_t visum = vdupq_n_s32(0);
|
||||
|
||||
// process 64 data points per iteration, totally 256 data points
|
||||
for (int j = 0; j < QK_K / 64; ++j, qx0 += 32, qx1 += 32, qy0 += 64, qy1 += 64) {
|
||||
@@ -2637,14 +2637,14 @@ void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
|
||||
// process 32 data points (share same block scale) per iteration
|
||||
for (int k = 0; k < 2; ++k) {
|
||||
const int blk = j * 2 + k;
|
||||
const int32x4_t block_scale = {
|
||||
const int32x4_t block_scale = vld1q_s32(((const int32_t[4]) {
|
||||
x0_scales[blk],
|
||||
x0_scales[blk],
|
||||
x1_scales[blk],
|
||||
x1_scales[blk],
|
||||
};
|
||||
}));
|
||||
|
||||
int32x4_t vr = {0};
|
||||
int32x4_t vr = vdupq_n_s32(0);
|
||||
for (int l = 0; l < 2; ++l) {
|
||||
const int idx = k * 2 + l;
|
||||
const int64x2_t vx0_s64 = vreinterpretq_s64_s8(vx0[idx]);
|
||||
@@ -2678,20 +2678,22 @@ void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
|
||||
vmull_s16(vget_high_s16(y0_sums), vget_high_s16(x1_mins))));
|
||||
bias[3] = vaddvq_s32(vaddq_s32(vmull_s16(vget_low_s16(y1_sums), vget_low_s16(x1_mins)),
|
||||
vmull_s16(vget_high_s16(y1_sums), vget_high_s16(x1_mins))));
|
||||
const float32x4_t dmins = {
|
||||
// note: the parentheses around the compound literal are required, vld1q_f32() is a
|
||||
// function-like macro on MSVC and would otherwise see four separate arguments
|
||||
const float32x4_t dmins = vld1q_f32(((const float[4]) {
|
||||
GGML_CPU_FP16_TO_FP32(x0->dmin) * y0->d,
|
||||
GGML_CPU_FP16_TO_FP32(x0->dmin) * y1->d,
|
||||
GGML_CPU_FP16_TO_FP32(x1->dmin) * y0->d,
|
||||
GGML_CPU_FP16_TO_FP32(x1->dmin) * y1->d,
|
||||
};
|
||||
}));
|
||||
vfsum = vmlsq_f32(vfsum, vcvtq_f32_s32(vld1q_s32(bias)), dmins);
|
||||
|
||||
const float32x4_t superblock_scale = {
|
||||
const float32x4_t superblock_scale = vld1q_f32(((const float[4]) {
|
||||
GGML_CPU_FP16_TO_FP32(x0->d) * y0->d,
|
||||
GGML_CPU_FP16_TO_FP32(x0->d) * y1->d,
|
||||
GGML_CPU_FP16_TO_FP32(x1->d) * y0->d,
|
||||
GGML_CPU_FP16_TO_FP32(x1->d) * y1->d,
|
||||
};
|
||||
}));
|
||||
vfsum = vmlaq_f32(vfsum, vcvtq_f32_s32(visum), superblock_scale);
|
||||
}
|
||||
}
|
||||
@@ -3261,12 +3263,8 @@ void ggml_vec_dot_q6_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
|
||||
qy0 += 16;
|
||||
qy1 += 16;
|
||||
|
||||
const int32x4_t block_scale = {
|
||||
x0->scales[blk],
|
||||
x0->scales[blk],
|
||||
x1->scales[blk],
|
||||
x1->scales[blk],
|
||||
};
|
||||
const int32x4_t block_scale =
|
||||
vcombine_s32(vdup_n_s32(x0->scales[blk]), vdup_n_s32(x1->scales[blk]));
|
||||
|
||||
// calculate four results at once with outer product
|
||||
const int8x16_t vx_l = vreinterpretq_s8_s64(vzip1q_s64(vreinterpretq_s64_s8(vx0[k]), vreinterpretq_s64_s8(vx1[k])));
|
||||
@@ -3319,12 +3317,14 @@ void ggml_vec_dot_q6_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
|
||||
|
||||
const int32x4_t vibias = vmulq_n_s32(vld1q_s32(bias), 32);
|
||||
|
||||
const float32x4_t superblock_scale = {
|
||||
// note: the parentheses around the compound literal are required, vld1q_f32() is a
|
||||
// function-like macro on MSVC and would otherwise see four separate arguments
|
||||
const float32x4_t superblock_scale = vld1q_f32(((const float[4]) {
|
||||
GGML_CPU_FP16_TO_FP32(x0->d) * y0->d,
|
||||
GGML_CPU_FP16_TO_FP32(x0->d) * y1->d,
|
||||
GGML_CPU_FP16_TO_FP32(x1->d) * y0->d,
|
||||
GGML_CPU_FP16_TO_FP32(x1->d) * y1->d,
|
||||
};
|
||||
}));
|
||||
|
||||
visum = vsubq_s32(visum, vibias);
|
||||
vfsum = vmlaq_f32(vfsum, vcvtq_f32_s32(visum), superblock_scale);
|
||||
|
||||
@@ -4,7 +4,6 @@
|
||||
#include "ggml-backend-impl.h"
|
||||
#include "ggml-backend.h"
|
||||
#include "traits.h"
|
||||
#include "iqp.h"
|
||||
#include "ggml-cpu-impl.h"
|
||||
#include "ggml-impl.h"
|
||||
#include "quants.h"
|
||||
@@ -15,6 +14,7 @@
|
||||
#include "ops.h"
|
||||
#include "ggml.h"
|
||||
#include "common.h"
|
||||
#include "tiled/tiled.h"
|
||||
|
||||
#if defined(_MSC_VER) || defined(__MINGW32__)
|
||||
#include <malloc.h> // using malloc.h with MSC/MINGW
|
||||
@@ -459,7 +459,7 @@ typedef pthread_mutex_t ggml_mutex_t;
|
||||
|
||||
#define ggml_lock_init(x) UNUSED(x)
|
||||
#define ggml_lock_destroy(x) UNUSED(x)
|
||||
#if defined(__x86_64__) || (defined(_MSC_VER) && defined(_M_AMD64))
|
||||
#if defined(__x86_64__) || (defined(_MSC_VER) && defined(_M_AMD64) && !defined(_M_ARM64EC))
|
||||
#define ggml_lock_lock(x) _mm_pause()
|
||||
#else
|
||||
#define ggml_lock_lock(x) UNUSED(x)
|
||||
@@ -1265,6 +1265,11 @@ void ggml_compute_forward_mul_mat(
|
||||
return;
|
||||
}
|
||||
|
||||
// If tiled is supported, it will execute the full op here and we return
|
||||
if (ggml_compute_forward_mul_mat_tiled(params, dst)) {
|
||||
return;
|
||||
}
|
||||
|
||||
GGML_TENSOR_BINARY_OP_LOCALS
|
||||
|
||||
const int ith = params->ith;
|
||||
@@ -1372,13 +1377,6 @@ UseGgmlGemm1:;
|
||||
|
||||
ggml_barrier(params->threadpool);
|
||||
|
||||
// IQ panel gemm (see iqp.h) - must come after the barrier above, it consumes the q8_K rows
|
||||
// of src1 from the work buffer
|
||||
if (ggml_cpu_iqp_supports_mul_mat(dst) && !params->use_ref) {
|
||||
ggml_compute_forward_mul_mat_iqp(params, dst);
|
||||
return;
|
||||
}
|
||||
|
||||
#if GGML_USE_LLAMAFILE
|
||||
if (src1->type != vec_dot_type) {
|
||||
const void* wdata = (src1->type == vec_dot_type) ? src1->data : params->wdata;
|
||||
@@ -1596,15 +1594,9 @@ static void ggml_compute_forward_mul_mat_id(
|
||||
char (*atomic_current_chunk)[CACHE_LINE_SIZE] = // [n_as]
|
||||
incr_ptr_aligned(&wdata_cur, CACHE_LINE_SIZE * n_as, CACHE_LINE_SIZE);
|
||||
|
||||
// IQ panel gemm (see iqp.h); per expert eligibility is decided below, but the work buffer is
|
||||
// reserved for the whole node (ggml_graph_plan sizes it without params, use_ref only skips the dispatch)
|
||||
const bool iqp = ggml_cpu_iqp_supports_mul_mat_id(dst) && !params->use_ref;
|
||||
|
||||
char * iqp_panels = NULL;
|
||||
|
||||
if (iqp) {
|
||||
iqp_panels = incr_ptr_aligned(&wdata_cur, nth * ggml_cpu_iqp_scratch_size(dst), 64);
|
||||
}
|
||||
// Tiled matmul (see tiled.h); per-thread work buffers, 0 bytes when disabled. The
|
||||
// reservation is unconditional, the per expert eligibility is decided at dispatch time
|
||||
char * tiled_scratch = incr_ptr_aligned(&wdata_cur, ggml_tiled_wdata_size(nth, dst), 64);
|
||||
|
||||
GGML_ASSERT(params->wsize >= (size_t)((char *) wdata_cur - (char *) params->wdata));
|
||||
|
||||
@@ -1677,10 +1669,8 @@ static void ggml_compute_forward_mul_mat_id(
|
||||
continue;
|
||||
}
|
||||
|
||||
if (iqp && ggml_cpu_iqp_mul_mat_id_min_batch(cne1)) {
|
||||
ggml_compute_forward_mul_mat_id_iqp(params, dst, cur_a, cne1, (const int32_t *) &MMID_MATRIX_ROW(cur_a, 0),
|
||||
iqp_panels);
|
||||
|
||||
// tiled takes over if profitable for this expert (see tiled.h)
|
||||
if (ggml_compute_forward_mul_mat_id_tiled(params, dst, cur_a, cne1, (const int32_t *) &MMID_MATRIX_ROW(cur_a, 0), tiled_scratch)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -2887,15 +2877,12 @@ struct ggml_cplan ggml_graph_plan(
|
||||
case GGML_OP_MUL_MAT:
|
||||
{
|
||||
const enum ggml_type vec_dot_type = type_traits_cpu[node->src[0]->type].vec_dot_type;
|
||||
|
||||
if (node->src[1]->type != vec_dot_type) {
|
||||
cur = ggml_row_size(vec_dot_type, ggml_nelements(node->src[1]));
|
||||
}
|
||||
|
||||
// the IQ panel path needs one scratch panel per thread past the q8_K rows
|
||||
if (ggml_cpu_iqp_supports_mul_mat(node)) {
|
||||
cur = GGML_PAD(cur, 64) + n_tasks * ggml_cpu_iqp_scratch_size(node);
|
||||
}
|
||||
// Workspace for tiled (see tiled.h)
|
||||
cur = GGML_PAD(cur, 64);
|
||||
cur += ggml_tiled_wdata_size(n_tasks, node);
|
||||
} break;
|
||||
case GGML_OP_MUL_MAT_ID:
|
||||
{
|
||||
@@ -2915,10 +2902,9 @@ struct ggml_cplan ggml_graph_plan(
|
||||
cur += n_as*ids->ne[0]*ids->ne[1]*sizeof(struct mmid_row_mapping) + sizeof(int64_t);
|
||||
// atomic_current_chunk
|
||||
cur += CACHE_LINE_SIZE*n_as + CACHE_LINE_SIZE;
|
||||
// the IQ panel path needs one scratch panel per thread on top of that
|
||||
if (ggml_cpu_iqp_supports_mul_mat_id(node)) {
|
||||
cur += n_tasks * ggml_cpu_iqp_scratch_size(node) + 64;
|
||||
}
|
||||
// Workspace for tiled (see tiled.h)
|
||||
cur = GGML_PAD(cur, 64);
|
||||
cur += ggml_tiled_wdata_size(n_tasks, node);
|
||||
} break;
|
||||
case GGML_OP_OUT_PROD:
|
||||
{
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,39 +0,0 @@
|
||||
#pragma once
|
||||
|
||||
#include "ggml-cpu-impl.h"
|
||||
#include "ggml.h"
|
||||
|
||||
// GGML internal header
|
||||
|
||||
// batched mul_mat path for the grid based IQ types: decode 8 src0 rows at a time into per thread scratch
|
||||
// (block_iqp_x8, see iqp.cpp) and run an integer gemm over them against all src1 columns
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
// whether cne1 rows of src1 are enough for the decode to pay for itself, per expert, for MUL_MAT_ID
|
||||
bool ggml_cpu_iqp_mul_mat_id_min_batch(int64_t cne1);
|
||||
|
||||
bool ggml_cpu_iqp_supports_mul_mat(const struct ggml_tensor * dst);
|
||||
|
||||
// node level test only - per expert eligibility is decided with ggml_cpu_iqp_mul_mat_id_min_batch
|
||||
bool ggml_cpu_iqp_supports_mul_mat_id(const struct ggml_tensor * dst);
|
||||
|
||||
// per thread panel scratch bytes, padded
|
||||
size_t ggml_cpu_iqp_scratch_size(const struct ggml_tensor * dst);
|
||||
|
||||
// must be called after src1 has been converted to q8_K into params->wdata and the threads have synchronized on it
|
||||
void ggml_compute_forward_mul_mat_iqp(const struct ggml_compute_params * params, struct ggml_tensor * dst);
|
||||
|
||||
// one expert: expert_rows points at its row of the matrix_rows table of (i1, i2) int32 pairs, panels at the base of the per thread panel scratches
|
||||
void ggml_compute_forward_mul_mat_id_iqp(const struct ggml_compute_params * params,
|
||||
struct ggml_tensor * dst,
|
||||
int64_t cur_a,
|
||||
int64_t cne1,
|
||||
const int32_t * expert_rows,
|
||||
void * panels);
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
@@ -228,7 +228,7 @@ template <> inline vfloat32m8_t madd(vbfloat16m4_t a, vbfloat16m4_t b, vfloat32m
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// VECTORIZED HORIZONTAL SUM
|
||||
|
||||
#if defined(__ARM_NEON)
|
||||
#if defined(__ARM_NEON) || defined(_M_ARM64) || defined(_M_ARM64EC)
|
||||
inline float hsum(float32x4_t x) {
|
||||
return vaddvq_f32(x);
|
||||
}
|
||||
|
||||
@@ -9037,6 +9037,11 @@ static void ggml_compute_forward_flash_attn_ext_tiled(
|
||||
simd_gemm(KQ, (const float *)Q_q, K_f32, Q_TILE_SZ, DK, KV_TILE_SZ);
|
||||
ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, scale);
|
||||
|
||||
if (logit_softcap != 0.0f) {
|
||||
ggml_vec_tanh_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, KQ);
|
||||
ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, logit_softcap);
|
||||
}
|
||||
|
||||
// Set padded KQ entries to -inf so softmax gives them zero weight
|
||||
if (kv_tile < KV_TILE_SZ) {
|
||||
for (int tq = 0; tq < Q_TILE_SZ; tq++) {
|
||||
@@ -9046,11 +9051,6 @@ static void ggml_compute_forward_flash_attn_ext_tiled(
|
||||
}
|
||||
}
|
||||
|
||||
if (logit_softcap != 0.0f) {
|
||||
ggml_vec_tanh_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, KQ);
|
||||
ggml_vec_scale_f32(Q_TILE_SZ * KV_TILE_SZ, KQ, logit_softcap);
|
||||
}
|
||||
|
||||
if (mask) {
|
||||
ggml_vec_add_f32(tile_rows * KV_TILE_SZ, KQ, KQ, mask32);
|
||||
}
|
||||
@@ -9320,7 +9320,7 @@ static void ggml_compute_forward_flash_attn_ext_f16(
|
||||
kv_is_f32_or_f16 &&
|
||||
k->type == v->type &&
|
||||
neq1 >= Q_TILE_SZ);
|
||||
#ifdef GGML_SIMD
|
||||
#if defined(GGML_SIMD) && !defined(__x86_64__) && !defined(_M_X64)
|
||||
#if defined(__ARM_FEATURE_SVE)
|
||||
const int64_t f32_epr = svcntw();
|
||||
#else
|
||||
|
||||
@@ -56,6 +56,56 @@ static inline void simd_gemm_ukernel(
|
||||
}
|
||||
}
|
||||
|
||||
template <int RM>
|
||||
static inline void simd_gemm_ukernel_tail(
|
||||
float * GGML_RESTRICT C,
|
||||
const float * GGML_RESTRICT A,
|
||||
const float * GGML_RESTRICT B,
|
||||
int K, int N, int cols)
|
||||
{
|
||||
#if defined(__AVX512F__)
|
||||
const __mmask16 mask = (1u << cols) - 1;
|
||||
__m512 acc[RM];
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
acc[i] = _mm512_maskz_loadu_ps(mask, C + i * N);
|
||||
}
|
||||
for (int64_t kk = 0; kk < K; kk++) {
|
||||
const __m512 b = _mm512_maskz_loadu_ps(mask, B + kk * N);
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
acc[i] = _mm512_mask3_fmadd_ps(_mm512_set1_ps(A[i * K + kk]), b, acc[i], mask);
|
||||
}
|
||||
}
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
_mm512_mask_storeu_ps(C + i * N, mask, acc[i]);
|
||||
}
|
||||
#elif defined(__AVX2__)
|
||||
const __m256i mask = _mm256_cmpgt_epi32(_mm256_set1_epi32(cols), _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7));
|
||||
__m256 acc[RM];
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
acc[i] = _mm256_maskload_ps(C + i * N, mask);
|
||||
}
|
||||
for (int64_t kk = 0; kk < K; kk++) {
|
||||
const __m256 b = _mm256_maskload_ps(B + kk * N, mask);
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
acc[i] = GGML_F32_VEC_FMA(acc[i], b, _mm256_set1_ps(A[i * K + kk]));
|
||||
}
|
||||
}
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
_mm256_maskstore_ps(C + i * N, mask, acc[i]);
|
||||
}
|
||||
#else
|
||||
for (int64_t j = 0; j < cols; j++) {
|
||||
for (int64_t i = 0; i < RM; i++) {
|
||||
float a = C[i * N + j];
|
||||
for (int64_t kk = 0; kk < K; kk++) {
|
||||
a += A[i * K + kk] * B[kk * N + j];
|
||||
}
|
||||
C[i * N + j] = a;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
// C[M x N] += A[M x K] * B[K x N]
|
||||
static void simd_gemm(
|
||||
float * GGML_RESTRICT C,
|
||||
@@ -74,14 +124,8 @@ static void simd_gemm(
|
||||
for (; jj + KN <= N; jj += KN) {
|
||||
simd_gemm_ukernel<GEMM_RM, 1>(C + jj, A, B + jj, K, N);
|
||||
}
|
||||
for (; jj < N; jj++) {
|
||||
for (int64_t i = 0; i < GEMM_RM; i++) {
|
||||
float a = C[i * N + jj];
|
||||
for (int64_t kk = 0; kk < K; kk++) {
|
||||
a += A[i * K + kk] * B[kk * N + jj];
|
||||
}
|
||||
C[i * N + jj] = a;
|
||||
}
|
||||
if (jj < N) {
|
||||
simd_gemm_ukernel_tail<GEMM_RM>(C + jj, A, B + jj, K, N, N - jj);
|
||||
}
|
||||
|
||||
A += GEMM_RM * K;
|
||||
@@ -97,12 +141,8 @@ static void simd_gemm(
|
||||
for (; jj + KN <= N; jj += KN) {
|
||||
simd_gemm_ukernel<1, 1>(C + jj, A, B + jj, K, N);
|
||||
}
|
||||
for (; jj < N; jj++) {
|
||||
float a = C[jj];
|
||||
for (int64_t kk = 0; kk < K; kk++) {
|
||||
a += A[kk] * B[kk * N + jj];
|
||||
}
|
||||
C[jj] = a;
|
||||
if (jj < N) {
|
||||
simd_gemm_ukernel_tail<1>(C + jj, A, B + jj, K, N, N - jj);
|
||||
}
|
||||
|
||||
A += K;
|
||||
|
||||
@@ -0,0 +1,719 @@
|
||||
// mulmat microtile kernels
|
||||
|
||||
#include "tiled-kernel.h"
|
||||
|
||||
#include "ggml.h"
|
||||
|
||||
#include <cstring>
|
||||
|
||||
#if defined(__AVX512VNNI__) || defined(__AVX2__) || defined(__AVX__)
|
||||
#include <immintrin.h>
|
||||
#endif
|
||||
|
||||
|
||||
// Reference implementation, slower than existing vec_dot approach
|
||||
template <int SUBBLK, bool HAS_MIN, int BIAS, int NK>
|
||||
static void tiled_run_microtile_scalar(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
|
||||
constexpr int NB = TILED_TILE_K / SUBBLK; // subblocks per 256-K block
|
||||
constexpr int NS = SUBBLK / 16; // per-16 bsums per subblock
|
||||
|
||||
// num_k/slab: the narrow path holds num_k slabs at row stride num_k*256 (num_k=1, slab=0 = standard)
|
||||
// NK: 1 = standard single-slab (offsets 0, strides compile-time); 0 = runtime
|
||||
// (num_k/slab from the args). The narrow path is memory-bound, so one runtime
|
||||
// version serves all of it.
|
||||
const int nkr = (NK > 0) ? NK : num_k;
|
||||
const int seff = (NK > 0) ? 0 : slab;
|
||||
const int qk_stride = nkr * TILED_TILE_K;
|
||||
const int qk_off = seff * TILED_TILE_K;
|
||||
const int nb_stride = NB * nkr;
|
||||
const int nb_off = seff * NB;
|
||||
const int bs_stride = (nkr == 1) ? TILED_TILE_ROWS : TILED_MICRO;
|
||||
const int bs_off = seff * TILED_MICRO;
|
||||
|
||||
float acc[TILED_MICRO][TILED_MICRO];
|
||||
memset(acc, 0, sizeof(acc));
|
||||
|
||||
// subdots at subblock granularity over the 256-K block, exact integer math
|
||||
for (int s = 0; s < NB; s++) {
|
||||
for (int j = 0; j < TILED_MICRO; j++) {
|
||||
const int br = j0 + j;
|
||||
const int8_t * q1 = &src1.q[br * qk_stride + qk_off + s * SUBBLK];
|
||||
int32_t bsum = 0;
|
||||
for (int u = 0; u < NS; u++) {
|
||||
bsum += src1.bsums[(bs_off + s * NS + u) * bs_stride + br];
|
||||
}
|
||||
|
||||
for (int i = 0; i < TILED_MICRO; i++) {
|
||||
const int ar = i0 + i;
|
||||
const int d_off = seff * TILED_MICRO + ar;
|
||||
const uint8_t * q0 = &src0.q[ar * qk_stride + qk_off + s * SUBBLK];
|
||||
|
||||
int32_t raw = 0;
|
||||
for (int e = 0; e < SUBBLK; e++) {
|
||||
raw += (int32_t) q0[e] * (int32_t) q1[e];
|
||||
}
|
||||
|
||||
// BIAS: subtract BIAS*bsum (src1's per-subblock code sum) from the exact int raw
|
||||
int32_t corr = raw;
|
||||
if constexpr (BIAS != 0) {
|
||||
corr -= BIAS * bsum;
|
||||
}
|
||||
const int32_t scales_raw = (int32_t) src0.scales[ar * nb_stride + nb_off + s] * corr;
|
||||
// d is NOT applied here: it is constant over the s-loop, so we apply it last before write-out
|
||||
if constexpr (HAS_MIN) {
|
||||
const int32_t mins_bsum = (int32_t) src0.mins[ar * nb_stride + nb_off + s] * bsum;
|
||||
acc[i][j] += (float) src0.d[d_off] * (float) scales_raw
|
||||
- (float) src0.dmin[d_off] * (float) mins_bsum;
|
||||
} else {
|
||||
acc[i][j] += (float) src0.d[d_off] * (float) scales_raw;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Apply d and write out to buf
|
||||
for (int i = 0; i < TILED_MICRO; i++) {
|
||||
for (int j = 0; j < TILED_MICRO; j++) {
|
||||
buf[(i0 + i) * buf_stride + (j0 + j)] += src1.d[seff * TILED_MICRO + j0 + j] * acc[i][j];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#if defined(__AVX512VNNI__) && defined(__AVX512VL__) && defined(__AVX512DQ__)
|
||||
|
||||
|
||||
// VNNI microkernel
|
||||
// 8x16 band pass: 8 src0 rows (i0, i0+8) x 16 src1 cols (j0, j0+16)
|
||||
// Math: s1_acc += scales_s*raw_s, s2_acc += mins_s*bsums_s per subblock; the row result
|
||||
// is d0*s1_acc - dmin*s2_acc (f32, exact in range), applied once per row.
|
||||
//
|
||||
// Register pressure: the band pass holds acc16 + s1_acc = 16 zmm for the
|
||||
// 8-row band (s2_acc has no dpbusd dependency, so it is computed after the
|
||||
// band pass and adds no register pressure);
|
||||
template <int SUBBLK, bool HAS_MIN, int BIAS, int NK>
|
||||
static void tiled_run_micro_vnni_8x16(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
|
||||
constexpr int NB = TILED_TILE_K / SUBBLK;
|
||||
constexpr int NS = SUBBLK / 16;
|
||||
constexpr int NG = SUBBLK / 4;
|
||||
|
||||
// band width, see the register-pressure note above
|
||||
constexpr int NUM_ROWS = 8;
|
||||
|
||||
// num_k = K-blocks per row (the narrow path holds num_k slabs at row stride num_k*256 so
|
||||
// each weight row is one long stream); slab = this call's slab. num_k=1, slab=0 = standard.
|
||||
// NK: 1 = standard single-slab (offsets 0, strides compile-time); 0 = runtime
|
||||
// (num_k/slab from the args). The narrow path is memory-bound, so one runtime
|
||||
// version serves all of it.
|
||||
const int nkr = (NK > 0) ? NK : num_k;
|
||||
const int seff = (NK > 0) ? 0 : slab;
|
||||
const int qk_stride = nkr * TILED_TILE_K;
|
||||
const int qk_off = seff * TILED_TILE_K;
|
||||
const int nb_stride = NB * nkr;
|
||||
const int nb_off = seff * NB;
|
||||
const int bs_stride = (nkr == 1) ? TILED_TILE_ROWS : TILED_MICRO;
|
||||
const int bs_off = seff * TILED_MICRO;
|
||||
|
||||
const __m512 d1_vec = _mm512_load_ps(&src1.d[seff * TILED_MICRO + j0]);
|
||||
|
||||
__m512i s1_acc[NUM_ROWS];
|
||||
for (int t = 0; t < NUM_ROWS; t++) { s1_acc[t] = _mm512_setzero_si512(); }
|
||||
__m512i s2_acc[NUM_ROWS];
|
||||
if constexpr (HAS_MIN) {
|
||||
for (int t = 0; t < NUM_ROWS; t++) { s2_acc[t] = _mm512_setzero_si512(); }
|
||||
}
|
||||
|
||||
const int32_t * bsums = src1.bsums; // int32 per-16 sums; NS > 1 combines NS lanes per subblock
|
||||
for (int s = 0; s < NB; s++) {
|
||||
__m512i bsums32 = _mm512_load_si512((const __m512i *) &bsums[(bs_off + s * NS) * bs_stride + j0]);
|
||||
for (int u = 1; u < NS; u++) {
|
||||
bsums32 = _mm512_add_epi32(bsums32, _mm512_load_si512((const __m512i *) &bsums[(bs_off + s * NS + u) * bs_stride + j0]));
|
||||
}
|
||||
|
||||
__m512i bias32 = _mm512_setzero_si512();
|
||||
if constexpr (BIAS != 0) {
|
||||
bias32 = _mm512_mullo_epi32(bsums32, _mm512_set1_epi32(BIAS));
|
||||
}
|
||||
|
||||
__m512i acc16[NUM_ROWS];
|
||||
for (int t = 0; t < NUM_ROWS; t++) { acc16[t] = _mm512_setzero_si512(); }
|
||||
|
||||
#ifdef __GNUC__
|
||||
#pragma GCC unroll 8 // pragma unrolled justified by measuing with/without
|
||||
#endif
|
||||
for (int g = 0; g < NG; g++) {
|
||||
const int kg = s * NG + g;
|
||||
// in-place interleave layout: [kg%16 @ TILED_TILE_K][kg/16 @ 64][row @ 4]
|
||||
const __m512i codes = _mm512_load_si512((const __m512i *) &src1.q[(j0 / TILED_MICRO) * (TILED_MICRO * qk_stride)
|
||||
+ qk_off + (kg % TILED_MICRO) * qk_stride + (kg / TILED_MICRO) * (TILED_MICRO * 4)]);
|
||||
for (int t = 0; t < NUM_ROWS; t++) {
|
||||
const uint32_t u4 = *(const uint32_t *) &src0.q[(i0 + t) * qk_stride + qk_off + kg * 4];
|
||||
const __m512i u4b = _mm512_set1_epi32((int) u4);
|
||||
acc16[t] = _mm512_dpbusd_epi32(acc16[t], u4b, codes);
|
||||
}
|
||||
}
|
||||
|
||||
// int correction: s1_acc += scales*(raw-BIAS*bsums)
|
||||
for (int t = 0; t < NUM_ROWS; t++) {
|
||||
const int ar = i0 + t;
|
||||
__m512i rawi = acc16[t];
|
||||
if constexpr (BIAS != 0) {
|
||||
rawi = _mm512_sub_epi32(rawi, bias32);
|
||||
}
|
||||
s1_acc[t] = _mm512_add_epi32(s1_acc[t],
|
||||
_mm512_mullo_epi32(rawi, _mm512_set1_epi32(src0.scales[ar * nb_stride + nb_off + s])));
|
||||
}
|
||||
}
|
||||
|
||||
// s2_acc += mins*bsums; independent of the dpbusd results, so it runs here instead of in the
|
||||
// band pass: the band pass keeps the register budget for the 8-row band
|
||||
if constexpr (HAS_MIN) {
|
||||
for (int s = 0; s < NB; s++) {
|
||||
__m512i bsums32 = _mm512_load_si512((const __m512i *) &bsums[(bs_off + s * NS) * bs_stride + j0]);
|
||||
for (int u = 1; u < NS; u++) {
|
||||
bsums32 = _mm512_add_epi32(bsums32, _mm512_load_si512((const __m512i *) &bsums[(bs_off + s * NS + u) * bs_stride + j0]));
|
||||
}
|
||||
for (int t = 0; t < NUM_ROWS; t++) {
|
||||
s2_acc[t] = _mm512_add_epi32(s2_acc[t],
|
||||
_mm512_mullo_epi32(bsums32, _mm512_set1_epi32(src0.mins[(i0 + t) * nb_stride + nb_off + s])));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// epilogue: int->float, apply per-row scales, store to buf
|
||||
for (int t = 0; t < NUM_ROWS; t++) {
|
||||
const int ar = i0 + t;
|
||||
const int d_off = seff * TILED_MICRO + ar;
|
||||
__m512 f1 = _mm512_cvtepi32_ps(s1_acc[t]);
|
||||
__m512 result = _mm512_mul_ps(f1, _mm512_set1_ps(src0.d[d_off]));
|
||||
if constexpr (HAS_MIN) {
|
||||
__m512 f2 = _mm512_cvtepi32_ps(s2_acc[t]);
|
||||
result = _mm512_fnmadd_ps(_mm512_set1_ps(src0.dmin[d_off]), f2, result);
|
||||
}
|
||||
float * p = &buf[(i0 + t) * buf_stride + j0];
|
||||
_mm512_store_ps(p, _mm512_add_ps(_mm512_load_ps(p), _mm512_mul_ps(result, d1_vec)));
|
||||
}
|
||||
}
|
||||
|
||||
// 16x16 microtile as two explicit 8x16 band passes
|
||||
template <int SUBBLK, bool HAS_MIN, int BIAS, int NK>
|
||||
static void tiled_run_microtile_vnni(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
|
||||
tiled_run_micro_vnni_8x16<SUBBLK, HAS_MIN, BIAS, NK>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
|
||||
tiled_run_micro_vnni_8x16<SUBBLK, HAS_MIN, BIAS, NK>(src0, src1, i0 + 8, j0, num_k, slab, buf, buf_stride);
|
||||
}
|
||||
|
||||
#endif // __AVX512VNNI__ && __AVX512VL__
|
||||
|
||||
#if defined(__AVX2__)
|
||||
|
||||
// AVX2 kernel.
|
||||
// We're effectively applying the existing vec_dot algorithms to an 8x16 block here.
|
||||
// Different paths based on subblock size as it affects when/where we multiply in scales and apply mins
|
||||
// SUBBLK=32: one 256-bit maddubs per (s, column), a 256-bit set1 scale,
|
||||
// one 8-lane i32 accumulator per column.
|
||||
// SUBBLK=16: two subblocks per 32B load. The 256-bit maddubs product
|
||||
// still gives 16 i16 lanes (0..7 = subblock sp, 8..15 = sp+1), but the
|
||||
// scale is applied per 128-bit half.
|
||||
// With BIAS != 0 the q0 codes are biased (up to 241) and a maddubs i16 pair
|
||||
// overflows, so the weight is debiased (w = q0 - BIAS, small). Two MACs:
|
||||
// ACTBIAS: activation pre-biased (+128) by repack, one maddubs(q1b, w); the +128 is
|
||||
// corrected out in the epilogue (128*sum_s scales[s]*w_bsum[s], precomputed in repack_src0).
|
||||
// Fits only when the debiased weight magnitude Wmax <= 64 (2*255*Wmax < 32767).
|
||||
// else (iq4_xs, Wmax = 127): the signed-int8 sign trick, ax = |w|, sy = q1*sign(w),
|
||||
// one maddubs (2*127*127 < 32767) so it always fits.
|
||||
template <int SUBBLK, bool HAS_MIN, int BIAS, int NK, bool ACTBIAS>
|
||||
static void tiled_run_microtile_avx2(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
|
||||
constexpr int NB = TILED_TILE_K / SUBBLK;
|
||||
constexpr int GROUP = 8; // Process 8 src1 columns concurrently (1 YMM vector)
|
||||
|
||||
static_assert(SUBBLK == 16 || SUBBLK == 32, "unsupported SUBBLK");
|
||||
|
||||
// num_k/slab: the narrow path holds num_k slabs at row stride num_k*256 (num_k=1, slab=0 = standard)
|
||||
// NK: 1 = standard single-slab (offsets 0, strides compile-time); 0 = runtime
|
||||
// (num_k/slab from the args). The narrow path is memory-bound, so one runtime
|
||||
// version serves all of it.
|
||||
const int nkr = (NK > 0) ? NK : num_k;
|
||||
const int seff = (NK > 0) ? 0 : slab;
|
||||
const int qk_stride = nkr * TILED_TILE_K;
|
||||
const int qk_off = seff * TILED_TILE_K;
|
||||
const int nb_stride = NB * nkr;
|
||||
const int nb_off = seff * NB;
|
||||
const int bs_stride = (nkr == 1) ? TILED_TILE_ROWS : TILED_MICRO;
|
||||
const int bs_off = seff * TILED_MICRO;
|
||||
|
||||
// Horizontal sum helper: converts 8x32-bit int lane inside a YMM register to scalar int32
|
||||
auto hsum256_epi32 = [](const __m256i v) -> int32_t {
|
||||
__m128i low = _mm256_castsi256_si128(v);
|
||||
__m128i high = _mm256_extracti128_si256(v, 1);
|
||||
__m128i sum128 = _mm_add_epi32(low, high);
|
||||
sum128 = _mm_add_epi32(sum128, _mm_shuffle_epi32(sum128, _MM_SHUFFLE(2, 3, 0, 1)));
|
||||
sum128 = _mm_add_epi32(sum128, _mm_shuffle_epi32(sum128, _MM_SHUFFLE(1, 0, 3, 2)));
|
||||
return _mm_cvtsi128_si32(sum128);
|
||||
};
|
||||
|
||||
for (int i = 0; i < TILED_MICRO; i++) {
|
||||
const int ar = i0 + i;
|
||||
const int d_off = seff * TILED_MICRO + ar;
|
||||
const float d0 = src0.d[d_off];
|
||||
const float dmin0 = src0.dmin[d_off];
|
||||
const uint8_t * q0 = &src0.q[ar * qk_stride + qk_off];
|
||||
const int32_t * scales_row = &src0.scales[ar * nb_stride + nb_off];
|
||||
const int32_t * mins_row = &src0.mins[ar * nb_stride + nb_off];
|
||||
|
||||
for (int g = 0; g < TILED_MICRO; g += GROUP) {
|
||||
// Pre-calculate weight pointers for this column group
|
||||
const int8_t * q1_ptr[GROUP];
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
q1_ptr[t] = &src1.q[(j0 + g + t) * qk_stride + qk_off];
|
||||
}
|
||||
|
||||
__m256i s2 = _mm256_setzero_si256();
|
||||
|
||||
__m256i acc[GROUP];
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
acc[t] = _mm256_setzero_si256();
|
||||
}
|
||||
|
||||
if constexpr (SUBBLK == 32) {
|
||||
for (int s = 0; s < NB; s++) {
|
||||
// OPTIMIZATION 1: Use aligned loads (_mm256_load_si256)
|
||||
const __m256i q0_32 = _mm256_load_si256((const __m256i *) &q0[s * SUBBLK]);
|
||||
const __m256i scales16 = _mm256_set1_epi16(scales_row[s]);
|
||||
|
||||
if constexpr (BIAS != 0) {
|
||||
const __m256i w = _mm256_sub_epi8(q0_32, _mm256_set1_epi8((int8_t) BIAS)); // debiased weight
|
||||
if constexpr (ACTBIAS) {
|
||||
// activation pre-biased (+128) by repack; q1b unsigned, w signed
|
||||
#pragma GCC unroll 8
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
const __m256i q1b_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][s * SUBBLK]);
|
||||
acc[t] = _mm256_add_epi32(acc[t],
|
||||
_mm256_madd_epi16(scales16, _mm256_maddubs_epi16(q1b_32, w)));
|
||||
}
|
||||
} else {
|
||||
// iq4_xs (Wmax > 64 overflows act-bias): sign trick, both operands <= 127
|
||||
const __m256i ax = _mm256_sign_epi8(w, w);
|
||||
#pragma GCC unroll 8
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
const __m256i q1_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][s * SUBBLK]);
|
||||
acc[t] = _mm256_add_epi32(acc[t],
|
||||
_mm256_madd_epi16(scales16, _mm256_maddubs_epi16(ax, _mm256_sign_epi8(q1_32, w))));
|
||||
}
|
||||
}
|
||||
} else {
|
||||
#pragma GCC unroll 8
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
const __m256i q1_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][s * SUBBLK]);
|
||||
acc[t] = _mm256_add_epi32(acc[t],
|
||||
_mm256_madd_epi16(scales16, _mm256_maddubs_epi16(q0_32, q1_32)));
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (HAS_MIN) {
|
||||
const __m256i bsums_v = _mm256_add_epi32(
|
||||
_mm256_load_si256((const __m256i *) &src1.bsums[(bs_off + s * 2) * bs_stride + j0 + g]),
|
||||
_mm256_load_si256((const __m256i *) &src1.bsums[(bs_off + s * 2 + 1) * bs_stride + j0 + g]));
|
||||
s2 = _mm256_add_epi32(s2, _mm256_mullo_epi32(bsums_v, _mm256_set1_epi32(mins_row[s])));
|
||||
}
|
||||
}
|
||||
} else { // SUBBLK == 16
|
||||
for (int sp = 0; sp < NB; sp += 2) {
|
||||
const __m256i q0_32 = _mm256_load_si256((const __m256i *) &q0[sp * SUBBLK]);
|
||||
const __m256i scalesv = _mm256_set_m128i(_mm_set1_epi16(scales_row[sp + 1]), _mm_set1_epi16(scales_row[sp]));
|
||||
if constexpr (BIAS != 0) {
|
||||
const __m256i w = _mm256_sub_epi8(q0_32, _mm256_set1_epi8((int8_t) BIAS)); // debiased weight
|
||||
if constexpr (ACTBIAS) {
|
||||
// activation pre-biased (+128) by repack; q1b unsigned, w signed
|
||||
#pragma GCC unroll 8
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
const __m256i q1b_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][sp * SUBBLK]);
|
||||
acc[t] = _mm256_add_epi32(acc[t], _mm256_madd_epi16(
|
||||
scalesv, _mm256_maddubs_epi16(q1b_32, w)));
|
||||
}
|
||||
} else {
|
||||
// iq4_xs (Wmax > 64 overflows act-bias): sign trick, both operands <= 127
|
||||
const __m256i ax = _mm256_sign_epi8(w, w);
|
||||
#pragma GCC unroll 8
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
const __m256i q1_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][sp * SUBBLK]);
|
||||
acc[t] = _mm256_add_epi32(acc[t], _mm256_madd_epi16(
|
||||
scalesv, _mm256_maddubs_epi16(ax, _mm256_sign_epi8(q1_32, w))));
|
||||
}
|
||||
}
|
||||
} else {
|
||||
#pragma GCC unroll 8
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
const __m256i q1_32 = _mm256_load_si256((const __m256i *) &q1_ptr[t][sp * SUBBLK]);
|
||||
acc[t] = _mm256_add_epi32(acc[t], _mm256_madd_epi16(
|
||||
scalesv, _mm256_maddubs_epi16(q0_32, q1_32)));
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (HAS_MIN) {
|
||||
const __m256i bsums0_v = _mm256_load_si256((const __m256i *) &src1.bsums[(bs_off + sp) * bs_stride + j0 + g]);
|
||||
const __m256i bsums1_v = _mm256_load_si256((const __m256i *) &src1.bsums[(bs_off + sp + 1) * bs_stride + j0 + g]);
|
||||
s2 = _mm256_add_epi32(s2, _mm256_add_epi32(
|
||||
_mm256_mullo_epi32(bsums0_v, _mm256_set1_epi32(mins_row[sp])),
|
||||
_mm256_mullo_epi32(bsums1_v, _mm256_set1_epi32(mins_row[sp + 1]))));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// OPTIMIZATION 2: Vectorized Epilogue with FMA3 gating
|
||||
// Reduce accumulators into a 256-bit vector across 8 columns
|
||||
alignas(32) int32_t s1_vals[GROUP];
|
||||
#pragma GCC unroll 8
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
s1_vals[t] = hsum256_epi32(acc[t]);
|
||||
}
|
||||
__m256i s1_vec = _mm256_load_si256((const __m256i *) s1_vals);
|
||||
|
||||
if constexpr (BIAS != 0 && ACTBIAS) {
|
||||
// act-bias correction: raw = desired + 128*sum_s scales[s]*w_bsum[s]; corr precomputed in repack_src0
|
||||
s1_vec = _mm256_sub_epi32(s1_vec, _mm256_set1_epi32(mins_row[0]));
|
||||
}
|
||||
|
||||
const __m256 src1_d_vec = _mm256_load_ps(&src1.d[seff * TILED_MICRO + j0 + g]);
|
||||
__m256 res_vec = _mm256_mul_ps(_mm256_set1_ps(d0), _mm256_cvtepi32_ps(s1_vec));
|
||||
|
||||
if constexpr (HAS_MIN) {
|
||||
#if defined(__FMA__)
|
||||
// FMA3 implementation: res_vec = res_vec - (dmin0 * s2)
|
||||
res_vec = _mm256_fnmadd_ps(_mm256_set1_ps(dmin0), _mm256_cvtepi32_ps(s2), res_vec);
|
||||
#else
|
||||
// Non-FMA fallback
|
||||
res_vec = _mm256_sub_ps(res_vec, _mm256_mul_ps(_mm256_set1_ps(dmin0), _mm256_cvtepi32_ps(s2)));
|
||||
#endif
|
||||
}
|
||||
|
||||
float * buf_ptr = &buf[ar * buf_stride + j0 + g];
|
||||
__m256 current_buf = _mm256_loadu_ps(buf_ptr);
|
||||
|
||||
#if defined(__FMA__)
|
||||
// FMA3 accumulation: current_buf + (res_vec * src1_d_vec)
|
||||
__m256 updated_buf = _mm256_fmadd_ps(res_vec, src1_d_vec, current_buf);
|
||||
#else
|
||||
__m256 updated_buf = _mm256_add_ps(current_buf, _mm256_mul_ps(res_vec, src1_d_vec));
|
||||
#endif
|
||||
|
||||
_mm256_storeu_ps(buf_ptr, updated_buf);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#endif // __AVX2__
|
||||
|
||||
#if defined(__AVX__) && !defined(__AVX2__)
|
||||
template <int SUBBLK, bool HAS_MIN, int BIAS, int NK>
|
||||
static void tiled_run_microtile_avx(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
|
||||
constexpr int NB = TILED_TILE_K / SUBBLK;
|
||||
constexpr int NS = SUBBLK / 16;
|
||||
constexpr int GROUP = 4; // src1 columns per group: one acc32 per column
|
||||
|
||||
// num_k/slab: the narrow path holds num_k slabs at row stride num_k*256 (num_k=1, slab=0 = standard)
|
||||
// NK: 1 = standard single-slab (offsets 0, strides compile-time); 0 = runtime
|
||||
// (num_k/slab from the args). The narrow path is memory-bound, so one runtime
|
||||
// version serves all of it.
|
||||
const int nkr = (NK > 0) ? NK : num_k;
|
||||
const int seff = (NK > 0) ? 0 : slab;
|
||||
const int qk_stride = nkr * TILED_TILE_K;
|
||||
const int qk_off = seff * TILED_TILE_K;
|
||||
const int nb_stride = NB * nkr;
|
||||
const int nb_off = seff * NB;
|
||||
const int bs_stride = (nkr == 1) ? TILED_TILE_ROWS : TILED_MICRO;
|
||||
const int bs_off = seff * TILED_MICRO;
|
||||
|
||||
for (int i = 0; i < TILED_MICRO; i++) {
|
||||
const int ar = i0 + i;
|
||||
const int d_off = seff * TILED_MICRO + ar;
|
||||
const float d0 = src0.d[d_off];
|
||||
const float dmin0 = src0.dmin[d_off];
|
||||
const uint8_t * q0 = &src0.q[ar * qk_stride + qk_off];
|
||||
const int32_t * scales_row = &src0.scales[ar * nb_stride + nb_off];
|
||||
const int32_t * mins_row = &src0.mins[ar * nb_stride + nb_off];
|
||||
|
||||
for (int g = 0; g < TILED_MICRO; g += GROUP) {
|
||||
const int8_t * q1g[GROUP];
|
||||
__m128i acc[GROUP];
|
||||
__m128i s2 = _mm_setzero_si128(); // sum_s mins_s * bsums_s
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
q1g[t] = &src1.q[(j0 + g + t) * qk_stride + qk_off];
|
||||
acc[t] = _mm_setzero_si128();
|
||||
}
|
||||
|
||||
for (int s = 0; s < NB; s++) {
|
||||
const __m128i scales16 = _mm_set1_epi16(scales_row[s]);
|
||||
if constexpr (HAS_MIN) {
|
||||
__m128i bsums_v = _mm_setzero_si128(); // per-col bsums as a 4-lane i32 vector
|
||||
for (int u = 0; u < NS; u++) {
|
||||
bsums_v = _mm_add_epi32(bsums_v, _mm_loadu_si128(
|
||||
(const __m128i *) &src1.bsums[(bs_off + s * NS + u) * bs_stride + j0 + g]));
|
||||
}
|
||||
s2 = _mm_add_epi32(s2, _mm_mullo_epi32(bsums_v, _mm_set1_epi32(mins_row[s])));
|
||||
}
|
||||
for (int u = 0; u < NS; u++) {
|
||||
const __m128i a16 = _mm_loadu_si128((const __m128i *) &q0[s * SUBBLK + u * 16]);
|
||||
if constexpr (BIAS != 0) {
|
||||
// debiased weight w = a16 - BIAS; signed dot via the sign trick (see header)
|
||||
const __m128i a16s = _mm_sub_epi8(a16, _mm_set1_epi8((int8_t) BIAS));
|
||||
const __m128i ax = _mm_sign_epi8(a16s, a16s);
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
const __m128i q1_16 = _mm_loadu_si128((const __m128i *) &q1g[t][s * SUBBLK + u * 16]);
|
||||
acc[t] = _mm_add_epi32(acc[t],
|
||||
_mm_madd_epi16(scales16, _mm_maddubs_epi16(ax, _mm_sign_epi8(q1_16, a16s))));
|
||||
}
|
||||
} else {
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
acc[t] = _mm_add_epi32(acc[t], _mm_madd_epi16(scales16,
|
||||
_mm_maddubs_epi16(a16,
|
||||
_mm_loadu_si128((const __m128i *) &q1g[t][s * SUBBLK + u * 16]))));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
int32_t s2_s[GROUP] = { 0, 0, 0, 0 };
|
||||
if constexpr (HAS_MIN) {
|
||||
_mm_storeu_si128((__m128i *) s2_s, s2);
|
||||
}
|
||||
|
||||
for (int t = 0; t < GROUP; t++) {
|
||||
// 4 i32 lanes -> scalar
|
||||
__m128i v = _mm_shuffle_epi32(acc[t], _MM_SHUFFLE(2, 3, 0, 1));
|
||||
acc[t] = _mm_add_epi32(acc[t], v);
|
||||
v = _mm_shuffle_epi32(acc[t], _MM_SHUFFLE(1, 0, 3, 2));
|
||||
acc[t] = _mm_add_epi32(acc[t], v);
|
||||
const int32_t s1 = _mm_cvtsi128_si32(acc[t]);
|
||||
|
||||
float res = d0 * (float) s1;
|
||||
if constexpr (HAS_MIN) {
|
||||
res -= dmin0 * (float) s2_s[t];
|
||||
}
|
||||
buf[ar * buf_stride + j0 + g + t] += src1.d[seff * TILED_MICRO + j0 + g + t] * res;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#endif // __AVX__ && !__AVX2__
|
||||
|
||||
// main microtile entry point. num_k = K-blocks per row (the narrow path reads num_k slabs at row
|
||||
// stride num_k*256 so each weight row is one long stream); slab = this call's slab. num_k=1,
|
||||
// slab=0 is the standard single-slab path.
|
||||
// num_k==1 dispatches to the NK=1 instantiation, whose strides/offsets are compile-time
|
||||
// constants (shifts/scaled-leas). num_k>1 (the narrow path) uses the NK=0 runtime-strided
|
||||
// version; it is memory-bound, so one version serves all of it.
|
||||
template <int SUBBLK, bool HAS_MIN, int BIAS, bool ACTBIAS>
|
||||
void tiled_run_microtile(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride) {
|
||||
const bool standard = (num_k == 1);
|
||||
#if defined(__AVX512VNNI__) && defined(__AVX512VL__) && defined(__AVX512DQ__)
|
||||
if (standard) {
|
||||
tiled_run_microtile_vnni<SUBBLK, HAS_MIN, BIAS, 1>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
|
||||
} else {
|
||||
tiled_run_microtile_vnni<SUBBLK, HAS_MIN, BIAS, 0>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
|
||||
}
|
||||
#elif defined(__AVX2__)
|
||||
if (standard) {
|
||||
tiled_run_microtile_avx2<SUBBLK, HAS_MIN, BIAS, 1, ACTBIAS>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
|
||||
} else {
|
||||
tiled_run_microtile_avx2<SUBBLK, HAS_MIN, BIAS, 0, ACTBIAS>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
|
||||
}
|
||||
#elif defined(__AVX__)
|
||||
if (standard) {
|
||||
tiled_run_microtile_avx<SUBBLK, HAS_MIN, BIAS, 1>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
|
||||
} else {
|
||||
tiled_run_microtile_avx<SUBBLK, HAS_MIN, BIAS, 0>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
|
||||
}
|
||||
#else
|
||||
if (standard) {
|
||||
tiled_run_microtile_scalar<SUBBLK, HAS_MIN, BIAS, 1>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
|
||||
} else {
|
||||
tiled_run_microtile_scalar<SUBBLK, HAS_MIN, BIAS, 0>(src0, src1, i0, j0, num_k, slab, buf, buf_stride);
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
// explicit instantiations for the in-use formats (q4_K and q5_K share the constants)
|
||||
// ACTBIAS selects the act-bias MAC (Wmax <= 64); iq4_xs (Wmax 127) keeps the sign trick
|
||||
// q4_K / q5_K: BIAS = 0
|
||||
template void tiled_run_microtile<32, true, 0, false>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
|
||||
// iq4_xs: sign trick (Wmax 127 overflows act-bias)
|
||||
template void tiled_run_microtile<32, false, 128, false>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
|
||||
// iq2_xxs, iq3_xxs, iq3_s, iq1_s: act-bias (Wmax <= 62)
|
||||
template void tiled_run_microtile<32, false, 128, true>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
|
||||
// iq2_xs, iq2_s, iq1_m: act-bias
|
||||
template void tiled_run_microtile<16, false, 128, true>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
|
||||
// q6_K: act-bias (Wmax 32)
|
||||
template void tiled_run_microtile<16, false, 32, true>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
|
||||
// q3_K: act-bias (Wmax 4)
|
||||
template void tiled_run_microtile<16, false, 4, true>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
|
||||
// q2_K: BIAS = 0
|
||||
template void tiled_run_microtile<16, true, 0, false>(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
|
||||
|
||||
|
||||
#define MIN(a, b) ((a) < (b) ? (a) : (b))
|
||||
|
||||
// block_q8_K in 4-byte words, for the int32 gather indices
|
||||
static_assert(sizeof(block_q8_K) == 292 && offsetof(block_q8_K, qs) == 4,
|
||||
"block_q8_K layout changed, fix the src1 repack");
|
||||
|
||||
// VNNI: interleave the natural [row][k] codes in-place into the group-local
|
||||
// [kg%16][kg/16][row][4] layout. base points at row 0 (row r at base + r*row_stride);
|
||||
// n_tiles 16x16 int32 tiles are transposed to the right (tile t occupies bytes
|
||||
// [t*64, t*64+64) per row). Each 16-row group is self-contained in row_stride*16 bytes.
|
||||
#if defined(__AVX512VNNI__) && defined(__AVX512VL__) && defined(__AVX512DQ__)
|
||||
void tiled_repack_src1(tiled_tile_src1 * src1, int row0, int num_k, bool bias) {
|
||||
GGML_UNUSED(bias);
|
||||
const int n_tiles = num_k * (TILED_TILE_K / (TILED_MICRO * 4));
|
||||
const int row_stride = num_k * TILED_TILE_K;
|
||||
uint8_t * base = (uint8_t *) src1->q + row0 * row_stride;
|
||||
for (int t = 0; t < n_tiles; t++) {
|
||||
uint8_t * cbase = base + t * (TILED_MICRO * 4);
|
||||
__m512i v[16];
|
||||
for (int r = 0; r < 16; r++) {
|
||||
v[r] = _mm512_load_si512((const __m512i *) (cbase + r * row_stride));
|
||||
}
|
||||
|
||||
// 16x16 int32 transpose, 4 butterfly phases
|
||||
// Phase 1: 1-element interleave, pairs (0,1), (2,3), ..., (14,15)
|
||||
{
|
||||
const __m512i idx_a = _mm512_setr_epi32(0, 16, 1, 17, 2, 18, 3, 19, 4, 20, 5, 21, 6, 22, 7, 23);
|
||||
const __m512i idx_b = _mm512_setr_epi32(8, 24, 9, 25, 10, 26, 11, 27, 12, 28, 13, 29, 14, 30, 15, 31);
|
||||
for (int i = 0; i < 16; i += 2) {
|
||||
__m512i a = v[i], b = v[i+1];
|
||||
v[i] = _mm512_permutex2var_epi32(a, idx_a, b);
|
||||
v[i+1] = _mm512_permutex2var_epi32(a, idx_b, b);
|
||||
}
|
||||
}
|
||||
// Phase 2: 2-element interleave, pairs (0,2), (1,3), (4,6), (5,7), ...
|
||||
{
|
||||
const __m512i idx_a = _mm512_setr_epi32(0, 1, 16, 17, 4, 5, 20, 21, 8, 9, 24, 25, 12, 13, 28, 29);
|
||||
const __m512i idx_b = _mm512_setr_epi32(2, 3, 18, 19, 6, 7, 22, 23, 10, 11, 26, 27, 14, 15, 30, 31);
|
||||
for (int i = 0; i < 16; i += 4) {
|
||||
__m512i a = v[i], b = v[i+2];
|
||||
v[i] = _mm512_permutex2var_epi32(a, idx_a, b);
|
||||
v[i+2] = _mm512_permutex2var_epi32(a, idx_b, b);
|
||||
a = v[i+1], b = v[i+3];
|
||||
v[i+1] = _mm512_permutex2var_epi32(a, idx_a, b);
|
||||
v[i+3] = _mm512_permutex2var_epi32(a, idx_b, b);
|
||||
}
|
||||
}
|
||||
// Phase 3: 4-element interleave, pairs (0,4), (1,5), ..., (7,11), (8,12), ...
|
||||
{
|
||||
const __m512i idx_a = _mm512_setr_epi32(0, 1, 2, 3, 16, 17, 18, 19, 4, 5, 6, 7, 20, 21, 22, 23);
|
||||
const __m512i idx_b = _mm512_setr_epi32(8, 9, 10, 11, 24, 25, 26, 27, 12, 13, 14, 15, 28, 29, 30, 31);
|
||||
for (int i = 0; i < 16; i += 8) {
|
||||
for (int j = 0; j < 4; j++) {
|
||||
__m512i a = v[i+j], b = v[i+4+j];
|
||||
v[i+j] = _mm512_permutex2var_epi32(a, idx_a, b);
|
||||
v[i+4+j] = _mm512_permutex2var_epi32(a, idx_b, b);
|
||||
}
|
||||
}
|
||||
}
|
||||
// Phase 4: 8-element interleave, pairs (0,8), (1,9), ..., (7,15)
|
||||
{
|
||||
const __m512i idx_a = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 16, 17, 18, 19, 20, 21, 22, 23);
|
||||
const __m512i idx_b = _mm512_setr_epi32(8, 9, 10, 11, 12, 13, 14, 15, 24, 25, 26, 27, 28, 29, 30, 31);
|
||||
for (int i = 0; i < 8; i++) {
|
||||
__m512i a = v[i], b = v[i+8];
|
||||
v[i] = _mm512_permutex2var_epi32(a, idx_a, b);
|
||||
v[i+8] = _mm512_permutex2var_epi32(a, idx_b, b);
|
||||
}
|
||||
}
|
||||
|
||||
// store: v[g] holds 16 int32s for k-group col_order[g], rows 0..15
|
||||
static const int col_order[16] = {0, 8, 1, 9, 4, 12, 5, 13, 2, 10, 3, 11, 6, 14, 7, 15};
|
||||
for (int g = 0; g < 16; g++) {
|
||||
_mm512_store_si512((void *) (cbase + col_order[g] * row_stride), v[g]);
|
||||
}
|
||||
}
|
||||
}
|
||||
#elif defined(__AVX2__)
|
||||
// act-bias: pre-bias the activation in place (+128) so the MAC's maddubs(q1b, w) sees unsigned bytes;
|
||||
// no transpose on AVX2 (natural layout). 16 rows x n_tiles*64 bytes, row r at base + r*row_stride
|
||||
void tiled_repack_src1(tiled_tile_src1 * src1, int row0, int num_k, bool bias) {
|
||||
if (!bias) {
|
||||
return;
|
||||
}
|
||||
const int row_stride = num_k * TILED_TILE_K;
|
||||
uint8_t * base = (uint8_t *) src1->q + row0 * row_stride;
|
||||
const int nbytes = row_stride; // full row width
|
||||
const __m256i b128 = _mm256_set1_epi8((int8_t) 0x80);
|
||||
for (int r = 0; r < TILED_MICRO; r++) {
|
||||
uint8_t * row = base + r * row_stride;
|
||||
for (int off = 0; off < nbytes; off += 32) {
|
||||
const __m256i v = _mm256_loadu_si256((const __m256i *) (row + off));
|
||||
_mm256_storeu_si256((__m256i *) (row + off), _mm256_xor_si256(v, b128));
|
||||
}
|
||||
}
|
||||
}
|
||||
#else
|
||||
void tiled_repack_src1(tiled_tile_src1 * src1, int row0, int num_k, bool bias) {
|
||||
GGML_UNUSED(src1); GGML_UNUSED(row0); GGML_UNUSED(num_k); GGML_UNUSED(bias);
|
||||
}
|
||||
#endif
|
||||
|
||||
#if defined(__AVX2__)
|
||||
// sum 32 unsigned bytes to int32 (via maddubs -> int16 pairs -> hsum)
|
||||
static int32_t tiled_byte_sum_32(const uint8_t * p) {
|
||||
const __m256i dot = _mm256_maddubs_epi16(_mm256_loadu_si256((const __m256i *) p), _mm256_set1_epi8(1));
|
||||
const __m256i sum = _mm256_madd_epi16(_mm256_set1_epi16(1), dot); // 8 int32
|
||||
// reduce 8 lanes: lo+hi -> 4 lanes, hadd -> 2, hadd -> 1 (a single hadd pair only folds 4 lanes)
|
||||
const __m128i t = _mm_add_epi32(_mm256_castsi256_si128(sum), _mm256_extracti128_si256(sum, 1));
|
||||
const __m128i h = _mm_hadd_epi32(t, t);
|
||||
return _mm_cvtsi128_si32(_mm_hadd_epi32(h, h));
|
||||
}
|
||||
// sum 16 unsigned bytes to int32
|
||||
static int32_t tiled_byte_sum_16(const uint8_t * p) {
|
||||
const __m128i dot = _mm_maddubs_epi16(_mm_loadu_si128((const __m128i *) p), _mm_set1_epi8(1));
|
||||
const __m128i sum = _mm_madd_epi16(_mm_set1_epi16(1), dot); // 4 int32
|
||||
const __m128i h = _mm_hadd_epi32(sum, sum);
|
||||
return _mm_cvtsi128_si32(_mm_hadd_epi32(h, h));
|
||||
}
|
||||
#endif
|
||||
|
||||
// act-bias: precompute the per-weight-row correction corr = 128*sum_s scales[s]*w_bsum[s] where
|
||||
// w_bsum[s] = sum of the debiased weight bytes in subblock s. Stored in mins[r][0] (unused for
|
||||
// HAS_MIN = false, which is exactly the act-bias case). Reused across all act bands and weight groups.
|
||||
// No-op off AVX2 (the act-bias MAC is AVX2-only) and when corr is false (iq4_xs / BIAS = 0).
|
||||
template <int SUBBLK>
|
||||
void tiled_repack_src0(tiled_tile_src0 * tile, int n_rows, int num_k, int BIAS, bool corr) {
|
||||
#if defined(__AVX2__)
|
||||
if (!corr) {
|
||||
return;
|
||||
}
|
||||
constexpr int NB = TILED_TILE_K / SUBBLK;
|
||||
const int nb_stride = NB * num_k;
|
||||
for (int slab = 0; slab < num_k; slab++) {
|
||||
for (int r = 0; r < n_rows; r++) {
|
||||
const uint8_t * q = &tile->q[r * (num_k * TILED_TILE_K) + slab * TILED_TILE_K];
|
||||
const int32_t * scales = &tile->scales[r * nb_stride + slab * NB];
|
||||
int32_t corr_val = 0;
|
||||
for (int s = 0; s < NB; s++) {
|
||||
const int32_t qsum = (SUBBLK == 32) ? tiled_byte_sum_32(&q[s * SUBBLK]) : tiled_byte_sum_16(&q[s * SUBBLK]);
|
||||
corr_val += 128 * scales[s] * (qsum - SUBBLK * BIAS);
|
||||
}
|
||||
tile->mins[r * nb_stride + slab * NB] = corr_val;
|
||||
}
|
||||
}
|
||||
#else
|
||||
GGML_UNUSED(tile); GGML_UNUSED(n_rows); GGML_UNUSED(num_k); GGML_UNUSED(BIAS); GGML_UNUSED(corr);
|
||||
#endif
|
||||
}
|
||||
|
||||
template void tiled_repack_src0<16>(tiled_tile_src0 * tile, int n_rows, int num_k, int BIAS, bool corr);
|
||||
template void tiled_repack_src0<32>(tiled_tile_src0 * tile, int n_rows, int num_k, int BIAS, bool corr);
|
||||
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
#pragma once
|
||||
|
||||
// Tiled matmul kernel API: tile structs, kernel definitions
|
||||
|
||||
// Currently only optimized for x86, new architectures should implement:
|
||||
// tiled_run_microtile: 16x16 microkernel
|
||||
// tiled_repack_src0: Optional repack/recalculation of src0, per macrotile
|
||||
// tiled_repack_src1: Optional repack/recalculation of src1, per microtile-band
|
||||
// bit unpacking routines: tiled_unpk_nib4, tiled_unpk_2bit, tiled_unpk_or
|
||||
// LUT value expansion routines: tiled_lut8, tiled_unpk_sign32, tiled_unpk_tern8
|
||||
|
||||
#define GGML_COMMON_DECL_C
|
||||
#include "ggml-common.h"
|
||||
|
||||
#include <stddef.h>
|
||||
#include <stdint.h>
|
||||
|
||||
#if defined(__AVX2__)
|
||||
#include <immintrin.h>
|
||||
#endif
|
||||
|
||||
#define TILED_TILE_K 256 // one QK_K block
|
||||
#define TILED_TILE_ROWS 256 // max window rows, ragged at edges
|
||||
#define TILED_MICRO 16 // microtile edge (also the bsums code-sum granularity)
|
||||
#define TILED_WS_SLOT (512 * 1024) // per-thread workspace slot, a clean 512KB (multiple of 64B)
|
||||
|
||||
// src0 tile: weight side, shared by all formats.
|
||||
// scales/mins are sized for the max subblock count (SUBBLK=16);
|
||||
// SUBBLK=32 formats index at stride 8 and leave the slack unused.
|
||||
struct tiled_tile_src0 {
|
||||
static constexpr int NB_MAX = TILED_TILE_K / 16; // max subblocks per 256-elem block
|
||||
|
||||
alignas(64) uint8_t q[TILED_TILE_ROWS * TILED_TILE_K]; // unsigned quants, widened to uint8
|
||||
float d[TILED_TILE_ROWS]; // One d from each input block, widened to f32
|
||||
float dmin[TILED_TILE_ROWS]; // dmin from each input block (if applicable), widened to F32
|
||||
int32_t scales[TILED_TILE_ROWS * NB_MAX]; // per-subblock scale, stored as int32_t
|
||||
int32_t mins[TILED_TILE_ROWS * NB_MAX]; // per-subblock min, used when HAS_MIN
|
||||
};
|
||||
|
||||
// src1 tile: built from q8_K (wdata)
|
||||
struct tiled_tile_src1 {
|
||||
// q8 codes, one byte per element. Note for VNNI these are reshaped + transposed to be suitable for dpbusd.
|
||||
alignas(64) int8_t q[TILED_TILE_ROWS * TILED_TILE_K];
|
||||
// per-16 code sums from q8_k (int16), widened to int32 so the kernels load them directly, no per-use cvt
|
||||
alignas(64) int32_t bsums[(TILED_TILE_K / 16) * TILED_TILE_ROWS];
|
||||
// f32 (not f16): q8_k stores fp16, the unpack converts once
|
||||
float d[TILED_TILE_ROWS];
|
||||
};
|
||||
|
||||
// per-thread workspace: all tiled state lives here, allocated in wdata (one slot per thread)
|
||||
struct tiled_ws {
|
||||
tiled_tile_src0 src0;
|
||||
tiled_tile_src1 src1;
|
||||
alignas(64) float acc[TILED_TILE_ROWS * TILED_TILE_ROWS];
|
||||
};
|
||||
|
||||
static_assert(sizeof(tiled_ws) <= TILED_WS_SLOT, "tiled workspace exceeds the 512KB per-thread slot");
|
||||
// the slot base is 64B-aligned and TILED_WS_SLOT is a multiple of 64B, so each per-thread slot
|
||||
// is 64B-aligned; these pin the tile fields and the acc buffer at aligned offsets within a slot
|
||||
static_assert(offsetof(tiled_ws, src0) % 64 == 0, "src0 not 64B-aligned in the workspace");
|
||||
static_assert(offsetof(tiled_ws, src1) % 64 == 0, "src1 not 64B-aligned in the workspace");
|
||||
static_assert(offsetof(tiled_ws, acc) % 64 == 0, "acc not 64B-aligned in the workspace");
|
||||
|
||||
// unpack primitives for reading quants, defined as inline here to keep arch-specific code in kernel.h/.cpp
|
||||
// If this section gets too hairy later, we can break up into separate includes.
|
||||
#if defined(__AVX2__)
|
||||
// packed 4-bit codes -> low nibbles (lo) + high nibbles (hi)
|
||||
inline void tiled_unpk_nib4(const uint8_t * src, uint8_t * lo, uint8_t * hi) {
|
||||
const __m256i v = _mm256_loadu_si256((const __m256i *) src);
|
||||
// mask before the lane shift so bits do not cross byte boundaries
|
||||
_mm256_storeu_si256((__m256i *) lo, _mm256_and_si256(v, _mm256_set1_epi8(0x0F)));
|
||||
_mm256_storeu_si256((__m256i *) hi, _mm256_srli_epi32(_mm256_and_si256(v, _mm256_set1_epi8((int8_t) 0xF0)), 4));
|
||||
}
|
||||
// 2-bit values at bit offset S
|
||||
template <int S> inline void tiled_unpk_2bit(const uint8_t * src, uint8_t * dst) {
|
||||
_mm256_storeu_si256((__m256i *) dst, _mm256_and_si256(
|
||||
_mm256_srli_epi32(_mm256_loadu_si256((const __m256i *) src), S), _mm256_set1_epi8(0x03)));
|
||||
}
|
||||
// OR the M-bit value at bit offset S of src into bit offset D of dst
|
||||
template <int S, int D, int M>
|
||||
inline void tiled_unpk_or(uint8_t * dst, const uint8_t * src) {
|
||||
const __m256i v = _mm256_slli_epi32(_mm256_and_si256(
|
||||
_mm256_srli_epi32(_mm256_loadu_si256((const __m256i *) src), S), _mm256_set1_epi8((uint8_t) M)), D);
|
||||
_mm256_storeu_si256((__m256i *) dst, _mm256_or_si256(_mm256_loadu_si256((const __m256i *) dst), v));
|
||||
}
|
||||
|
||||
|
||||
// Unpacking kernels for IQ quants
|
||||
|
||||
// LUT value expansion for the LUT-based formats (iq4_xs, iq grids): the bit unpackers
|
||||
// above give the indices, these expand 8/16 of them to widened codes in one pass
|
||||
// 16-entry byte LUT: dst[j] = lut[src[j]] (16 bytes)
|
||||
inline void tiled_lut8(const uint8_t * lut, const uint8_t * src, uint8_t * dst) {
|
||||
_mm_storeu_si128((__m128i *) dst, _mm_shuffle_epi8(_mm_loadu_si128((const __m128i *) lut),
|
||||
_mm_loadu_si128((const __m128i *) src)));
|
||||
}
|
||||
|
||||
|
||||
// 32 grid magnitudes (4 x 64-bit groups g0..g3, 8 values each) + 4 sign bytes
|
||||
// (byte l signs values 8*l .. 8*l+7) -> codes stored as (value + 128)
|
||||
inline void tiled_unpk_sign32(uint64_t g0, uint64_t g1, uint64_t g2, uint64_t g3,
|
||||
const uint8_t signs[4], uint8_t * dst32) {
|
||||
const __m256i v = _mm256_set_epi64x((int64_t) g3, (int64_t) g2, (int64_t) g1, (int64_t) g0);
|
||||
const __m256i sv = _mm256_shuffle_epi8(_mm256_set1_epi32((int32_t) (signs[0] | signs[1] << 8 | signs[2] << 16 | signs[3] << 24)),
|
||||
_mm256_setr_epi8(0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1,
|
||||
2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3));
|
||||
const __m256i sel = _mm256_set1_epi64x((int64_t) 0x8040201008040201ULL);
|
||||
// 0xFF in every lane whose sign bit is set; GFNI does the and+compare in one instruction
|
||||
#ifdef __GFNI__
|
||||
const __m256i mask = _mm256_gf2p8affine_epi64_epi8(sel, sv, 0);
|
||||
#else
|
||||
const __m256i mask = _mm256_cmpeq_epi8(_mm256_and_si256(sv, sel), sel);
|
||||
#endif
|
||||
// v ^ mask - mask negates the signed lanes (v elsewhere); + 128 gives the biased code
|
||||
const __m256i sgn = _mm256_sub_epi8(_mm256_xor_si256(v, mask), mask);
|
||||
_mm256_storeu_si256((__m256i *) dst32, _mm256_add_epi8(sgn, _mm256_set1_epi8((int8_t) 128)));
|
||||
}
|
||||
// 8 ternary grid bytes (0 = 0, 1 = +1, 0xFF = -1): dst[j] = 128 + delta + 8 * (int8_t) src[j]
|
||||
inline void tiled_unpk_tern8(const uint8_t * src, int8_t delta, uint8_t * dst) {
|
||||
const __m128i v = _mm_cvtepi8_epi16(_mm_loadl_epi64((const __m128i *) src));
|
||||
const __m128i p = _mm_add_epi16(_mm_slli_epi16(v, 3), _mm_set1_epi16(128 + (int) delta));
|
||||
_mm_storel_epi64((__m128i *) dst, _mm_packus_epi16(p, _mm_setzero_si128()));
|
||||
}
|
||||
#else
|
||||
|
||||
// Scalar definitions for unpackers.
|
||||
|
||||
|
||||
inline void tiled_unpk_nib4(const uint8_t * src, uint8_t * lo, uint8_t * hi) {
|
||||
for (int l = 0; l < 32; l++) { lo[l] = (uint8_t) (src[l] & 0xF); hi[l] = (uint8_t) (src[l] >> 4); }
|
||||
}
|
||||
template <int S>
|
||||
inline void tiled_unpk_2bit(const uint8_t * src, uint8_t * dst) {
|
||||
for (int l = 0; l < 32; l++) { dst[l] = (uint8_t) ((src[l] >> S) & 3); }
|
||||
}
|
||||
template <int S, int D, int M>
|
||||
inline void tiled_unpk_or(uint8_t * dst, const uint8_t * src) {
|
||||
for (int l = 0; l < 32; l++) { dst[l] = (uint8_t) (dst[l] | (((src[l] >> S) & M) << D)); }
|
||||
}
|
||||
inline void tiled_lut8(const uint8_t * lut, const uint8_t * src, uint8_t * dst) {
|
||||
for (int j = 0; j < 16; j++) { dst[j] = lut[src[j]]; }
|
||||
}
|
||||
inline void tiled_unpk_sign32(uint64_t g0, uint64_t g1, uint64_t g2, uint64_t g3,
|
||||
const uint8_t signs[4], uint8_t * dst32) {
|
||||
const uint64_t g[4] = { g0, g1, g2, g3 };
|
||||
for (int l = 0; l < 4; l++) {
|
||||
const uint8_t * v = (const uint8_t *) &g[l];
|
||||
const uint8_t s = signs[l];
|
||||
for (int j = 0; j < 8; j++) {
|
||||
dst32[8 * l + j] = (s & (1 << j)) ? (uint8_t) (128 - v[j]) : (uint8_t) (128 + v[j]);
|
||||
}
|
||||
}
|
||||
}
|
||||
// 8 ternary grid bytes (0 = 0, 1 = +1, 0xFF = -1): dst[j] = 128 + delta + 8 * (int8_t) src[j]
|
||||
inline void tiled_unpk_tern8(const uint8_t * src, int8_t delta, uint8_t * dst) {
|
||||
for (int j = 0; j < 8; j++) { dst[j] = (uint8_t) (128 + (int) delta + 8 * (int8_t) src[j]); }
|
||||
}
|
||||
#endif
|
||||
|
||||
// Accumulate one 16x16 microtile (src0 rows [i0, i0+16), src1 cols [j0, j0+16))
|
||||
// over one 256-K slab held in the tiles into a j-major float buffer
|
||||
// (row width buf_stride): buf[i*buf_stride + j] += partial.
|
||||
// SUBBLK/HAS_MIN/BIAS are the src0 format constants (see tiled_tile_src0).
|
||||
// ACTBIAS (AVX2 only): the activation is pre-biased +128 by tiled_repack_src1,
|
||||
// num_k = K-blocks per row: the tile holds num_k slabs at row stride num_k*256
|
||||
// (Default case is num_k=1, 256x256 tiles, we go to longer num_k to improve memory bandwidth when num_rows is small)
|
||||
template <int SUBBLK, bool HAS_MIN, int BIAS, bool ACTBIAS>
|
||||
void tiled_run_microtile(const tiled_tile_src0 & src0, const tiled_tile_src1 & src1,
|
||||
int i0, int j0, int num_k, int slab, float * buf, int buf_stride);
|
||||
|
||||
// Optional repack, if profitable for the kernel.
|
||||
// Repacks one 16-row band of src1 codes, called by driver as we reach each 16-row band in outer loop
|
||||
void tiled_repack_src1(tiled_tile_src1 * src1, int row0, int num_k, bool bias);
|
||||
|
||||
// Optional repack, if profitable for the kernel
|
||||
// Repack the entire src0 panel in place, called by the driver immediately after dequant
|
||||
template <int SUBBLK>
|
||||
void tiled_repack_src0(tiled_tile_src0 * tile, int n_rows, int num_k, int BIAS, bool corr);
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,31 @@
|
||||
#pragma once
|
||||
|
||||
#include <stddef.h>
|
||||
#include <stdint.h>
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
// Amount of wdata to reserve for tiled workspaces
|
||||
size_t ggml_tiled_wdata_size(int n_tasks, struct ggml_tensor * dst);
|
||||
|
||||
// tiled K-quant matmul; returns true if the op was computed here,
|
||||
// false to fall through to the stock path
|
||||
bool ggml_compute_forward_mul_mat_tiled(const struct ggml_compute_params * params,
|
||||
struct ggml_tensor * dst);
|
||||
|
||||
// MUL_MAT_ID (MoE) path, one expert; returns true if the expert was computed here, per expert
|
||||
// eligibility (type gate, batch floor) is decided inside. expert_rows points at the expert's
|
||||
// row of the matrix_rows table of (expert slot, batch row) int32 pairs; scratch is the
|
||||
// per-thread tiled_ws region reserved in wdata, n_tasks of ggml_tiled_ws_size() bytes
|
||||
bool ggml_compute_forward_mul_mat_id_tiled(const struct ggml_compute_params * params,
|
||||
struct ggml_tensor * dst,
|
||||
int64_t cur_a,
|
||||
int64_t cne1,
|
||||
const int32_t * expert_rows,
|
||||
char * scratch);
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
@@ -1,5 +1,18 @@
|
||||
cmake_minimum_required(VERSION 3.18) # for CMAKE_CUDA_ARCHITECTURES
|
||||
|
||||
# ARM64EC uses the x64 ABI, so it needs the x64 import libraries. The arm64 ones
|
||||
# hold native ARM64 symbols and cannot satisfy EC references at link time.
|
||||
# FindCUDAToolkit picks the library dir from the host arch, not the target, so
|
||||
# anchor its search at lib/x64 here. It derives everything else from CUDA_CUDART.
|
||||
string(TOLOWER "${CMAKE_GENERATOR_PLATFORM}" GGML_CUDA_PLATFORM_LWR)
|
||||
if (MSVC AND GGML_CUDA_PLATFORM_LWR STREQUAL "arm64ec")
|
||||
find_library(CUDA_CUDART
|
||||
NAMES cudart
|
||||
HINTS ${CUDAToolkit_ROOT} ENV CUDA_PATH
|
||||
PATH_SUFFIXES lib/x64
|
||||
)
|
||||
endif()
|
||||
|
||||
find_package(CUDAToolkit)
|
||||
|
||||
if (CUDAToolkit_FOUND)
|
||||
|
||||
@@ -111,9 +111,9 @@
|
||||
#define GGML_CUDA_CC_IS_QY2(cc) (cc >= GGML_CUDA_CC_QY2 && cc < GGML_CUDA_CC_PH1)
|
||||
#define GGML_CUDA_CC_IS_PH1(cc) (cc >= GGML_CUDA_CC_PH1)
|
||||
|
||||
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && CUDART_VERSION >= 11070
|
||||
#if !defined(GGML_USE_HIP) && (defined(GGML_USE_MUSA) || CUDART_VERSION >= 11070)
|
||||
# define GGML_CUDA_USE_CUB
|
||||
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && CUDART_VERSION >= 11070
|
||||
#endif // !defined(GGML_USE_HIP) && (defined(GGML_USE_MUSA) || CUDART_VERSION >= 11070)
|
||||
|
||||
// PDL host-side support (cudaLaunchKernelEx) requires CUDART >= 11.8.
|
||||
// However, this has been bugged in CTK < 12.3 for MSVC builds, see
|
||||
@@ -237,7 +237,7 @@ static const char * cu_get_error_str(CUresult err) {
|
||||
#define CU_CHECK(err) CUDA_CHECK_GEN(err, CUDA_SUCCESS, cu_get_error_str)
|
||||
#endif
|
||||
|
||||
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
#if !defined(GGML_USE_HIP)
|
||||
# define CUDA_SET_SHARED_MEMORY_LIMIT(kernel, nbytes) \
|
||||
do { \
|
||||
static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = { false }; \
|
||||
@@ -252,7 +252,7 @@ static const char * cu_get_error_str(CUresult err) {
|
||||
do { \
|
||||
GGML_UNUSED(nbytes); \
|
||||
} while (0)
|
||||
#endif // !(defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
#endif // !defined(GGML_USE_HIP)
|
||||
|
||||
#if CUDART_VERSION >= 11010 || defined(GGML_USE_MUSA)
|
||||
#define GGML_CUDA_ASSUME(x) __builtin_assume(x)
|
||||
@@ -397,7 +397,7 @@ static constexpr __device__ int ggml_cuda_get_physical_warp_size() {
|
||||
|
||||
// Maximum number of bytes that can be copied in a single instruction.
|
||||
static constexpr __device__ int ggml_cuda_get_max_cpy_bytes() {
|
||||
#ifdef GGML_USE_HIP
|
||||
#if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
|
||||
return 16;
|
||||
#else
|
||||
#if __CUDA_ARCH__ >= GGML_CUDA_CC_VOLTA
|
||||
@@ -405,7 +405,7 @@ static constexpr __device__ int ggml_cuda_get_max_cpy_bytes() {
|
||||
#else
|
||||
return 8;
|
||||
#endif // __CUDA_ARCH__ >= GGML_CUDA_CC_VOLTA
|
||||
#endif // GGML_USE_HIP
|
||||
#endif // defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
|
||||
}
|
||||
|
||||
|
||||
@@ -424,10 +424,6 @@ static __device__ void no_device_code(
|
||||
__trap();
|
||||
|
||||
GGML_UNUSED(no_device_code); // suppress unused function warning
|
||||
|
||||
#if defined(GGML_USE_MUSA)
|
||||
__builtin_unreachable();
|
||||
#endif // defined(GGML_USE_MUSA)
|
||||
}
|
||||
|
||||
#ifdef __CUDA_ARCH__
|
||||
@@ -696,16 +692,11 @@ static __device__ __forceinline__ half2 ggml_cuda_hmax2(const half2 a, const hal
|
||||
|
||||
template<int width = WARP_SIZE>
|
||||
static __device__ __forceinline__ half2 warp_reduce_max(half2 x) {
|
||||
#if !defined(GGML_USE_HIP) && __CUDA_ARCH__ >= GGML_CUDA_CC_PASCAL || defined(GGML_USE_HIP)
|
||||
#pragma unroll
|
||||
for (int offset = width/2; offset > 0; offset >>= 1) {
|
||||
x = ggml_cuda_hmax2(x, __shfl_xor_sync(0xffffffff, x, offset, width));
|
||||
}
|
||||
return x;
|
||||
#else
|
||||
GGML_UNUSED(x);
|
||||
NO_DEVICE_CODE;
|
||||
#endif // !defined(GGML_USE_HIP) && __CUDA_ARCH__ >= GGML_CUDA_CC_PASCAL || defined(GGML_USE_HIP)
|
||||
}
|
||||
|
||||
#if (defined(CUDART_VERSION) && CUDART_VERSION < CUDART_HMASK) || defined(GGML_USE_HIP) || \
|
||||
|
||||
@@ -1875,7 +1875,7 @@ static __global__ void flash_attn_ext_f16(
|
||||
#endif // defined(AMD_WMMA_AVAILABLE)
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE)
|
||||
if (ncols1*ncols2 < 16 || DKQ > 256) {
|
||||
if (ncols1*ncols2 < 16 || (DKQ > 256 && ncols1*ncols2 < 32)) {
|
||||
NO_DEVICE_CODE;
|
||||
return;
|
||||
}
|
||||
@@ -2091,26 +2091,22 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml
|
||||
constexpr bool use_sparse_kernel = false;
|
||||
fattn_kernel = flash_attn_ext_f16<DKQ, DV, ncols1, ncols2, use_logit_softcap, V_is_K_view, use_sparse_kernel>;
|
||||
|
||||
#if !defined(GGML_USE_MUSA)
|
||||
static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = {false};
|
||||
if (!shared_memory_limit_raised[id]) {
|
||||
CUDA_CHECK(cudaFuncSetAttribute(reinterpret_cast<fattn_kernel_ptr_t>(fattn_kernel), cudaFuncAttributeMaxDynamicSharedMemorySize, nbytes_shared_total));
|
||||
shared_memory_limit_raised[id] = true;
|
||||
}
|
||||
#endif // !defined(GGML_USE_MUSA)
|
||||
}
|
||||
} else {
|
||||
constexpr bool use_logit_softcap = true;
|
||||
constexpr bool use_sparse_kernel = false;
|
||||
fattn_kernel = flash_attn_ext_f16<DKQ, DV, ncols1, ncols2, use_logit_softcap, V_is_K_view, use_sparse_kernel>;
|
||||
|
||||
#if !defined(GGML_USE_MUSA)
|
||||
static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = {false};
|
||||
if (!shared_memory_limit_raised[id]) {
|
||||
CUDA_CHECK(cudaFuncSetAttribute(reinterpret_cast<fattn_kernel_ptr_t>(fattn_kernel), cudaFuncAttributeMaxDynamicSharedMemorySize, nbytes_shared_total));
|
||||
shared_memory_limit_raised[id] = true;
|
||||
}
|
||||
#endif // !defined(GGML_USE_MUSA)
|
||||
}
|
||||
|
||||
launch_fattn<DV, ncols1, ncols2>
|
||||
|
||||
@@ -19,41 +19,41 @@
|
||||
} \
|
||||
|
||||
static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_nvidia_fp16(const int DKQ, const int DV, const int ncols) {
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 2, 64, 2, 64, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 4, 128, 2, 64, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 8, 256, 2, 64, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 2, 128, 3, 128, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 4, 128, 2, 128, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 8, 128, 3, 128, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 16, 256, 2, 64, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 40, 40, 32, 256, 2, 64, 40)
|
||||
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 2, 64, 2, 64, 64)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 4, 128, 2, 64, 64)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 8, 256, 2, 64, 64)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 8, 256, 3, 128, 64)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 16, 256, 2, 64, 64)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 64, 64, 32, 256, 2, 64, 64)
|
||||
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 2, 64, 2, 64, 72)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 4, 128, 2, 64, 72)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 2, 128, 2, 64, 72)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 4, 128, 3, 128, 72)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 8, 256, 2, 64, 72)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 16, 256, 2, 64, 72)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 72, 72, 32, 256, 2, 64, 72)
|
||||
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 2, 64, 2, 64, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 2, 128, 2, 64, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 4, 128, 2, 64, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 8, 256, 2, 64, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 16, 256, 2, 64, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 32, 256, 2, 64, 40)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 80, 80, 32, 256, 2, 128, 40)
|
||||
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 2, 64, 2, 64, 48)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 4, 128, 2, 64, 48)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 8, 256, 2, 64, 48)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 16, 256, 2, 64, 48)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 32, 256, 2, 64, 48)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE( 96, 96, 32, 256, 2, 128, 24)
|
||||
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 2, 64, 2, 64, 56)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 2, 128, 3, 64, 56)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 4, 128, 2, 64, 56)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 8, 256, 2, 64, 56)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 16, 256, 2, 64, 56)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 32, 256, 2, 64, 56)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 8, 256, 3, 64, 112)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 16, 256, 2, 32, 112)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(112, 112, 32, 256, 2, 128, 56)
|
||||
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(128, 128, 2, 64, 2, 64, 64)
|
||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(128, 128, 4, 128, 2, 64, 64)
|
||||
|
||||
@@ -681,7 +681,7 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const
|
||||
}
|
||||
|
||||
// AMD MFMA needs a certain minimum batch size to outscale the tile kernel for large head sizes.
|
||||
if ((amd_mfma_available(cc) && Q->ne[0] <= 256) && Q->ne[0] != 40 && Q->ne[0] != 72) {
|
||||
if (amd_mfma_available(cc) && Q->ne[0] != 40 && Q->ne[0] != 72) {
|
||||
if ((Q->ne[0] <= 64 && Q->ne[1] * gqa_ratio_eff > 8)) {
|
||||
return BEST_FATTN_KERNEL_MMA_F16;
|
||||
}
|
||||
@@ -691,6 +691,9 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const
|
||||
if ((Q->ne[0] <= 256 && Q->ne[1] * gqa_ratio_eff > 64)) {
|
||||
return BEST_FATTN_KERNEL_MMA_F16;
|
||||
}
|
||||
if (Q->ne[0] > 256 && gqa_opt_applies && Q->ne[1] * gqa_ratio_eff > 128) {
|
||||
return BEST_FATTN_KERNEL_MMA_F16;
|
||||
}
|
||||
}
|
||||
|
||||
// AMD WMMA is faster than the tile kernel if the wide tiles with high arithmetic intensity can be utilized.
|
||||
|
||||
+40
-13
@@ -1,9 +1,10 @@
|
||||
#include "common.cuh"
|
||||
#include "convert.cuh"
|
||||
#include "fwht.cuh"
|
||||
|
||||
template <int N>
|
||||
template <int N, typename T>
|
||||
__launch_bounds__(4*ggml_cuda_get_physical_warp_size(), 1)
|
||||
__global__ void fwht_cuda(const float * src, float * dst, const int64_t n_rows, const float scale) {
|
||||
__global__ void fwht_cuda(const T * src, float * dst, const int64_t n_rows, const float scale) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
|
||||
const int64_t r = (int64_t) blockIdx.x * blockDim.y + threadIdx.y;
|
||||
@@ -22,7 +23,7 @@ __global__ void fwht_cuda(const float * src, float * dst, const int64_t n_rows,
|
||||
ggml_cuda_pdl_sync();
|
||||
#pragma unroll
|
||||
for (int i = 0; i < el_w; ++i) {
|
||||
reg[i] = src[i * warp_size + lane] * scale;
|
||||
reg[i] = ggml_cuda_cast<float>(src[i * warp_size + lane]) * scale;
|
||||
}
|
||||
|
||||
#pragma unroll
|
||||
@@ -58,15 +59,12 @@ __global__ void fwht_cuda(const float * src, float * dst, const int64_t n_rows,
|
||||
}
|
||||
}
|
||||
|
||||
bool ggml_cuda_op_fwht(ggml_backend_cuda_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
|
||||
GGML_ASSERT(ggml_are_same_shape(src, dst));
|
||||
if (!ggml_is_contiguous(src) || !ggml_is_contiguous(dst)) {
|
||||
return false;
|
||||
}
|
||||
template <typename T>
|
||||
static bool ggml_cuda_op_fwht_impl(ggml_backend_cuda_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
|
||||
const int n = src->ne[0];
|
||||
const int64_t rows = ggml_nrows(src);
|
||||
|
||||
const float * src_d = (const float *) src->data;
|
||||
const T * src_d = (const T *) src->data;
|
||||
float * dst_d = (float *) dst->data;
|
||||
|
||||
const int warp_size = ggml_cuda_info().devices[ggml_cuda_get_device()].warp_size;
|
||||
@@ -84,18 +82,47 @@ bool ggml_cuda_op_fwht(ggml_backend_cuda_context & ctx, const ggml_tensor * src,
|
||||
|
||||
switch (n) {
|
||||
case 64:
|
||||
ggml_cuda_kernel_launch(fwht_cuda<64>, launch_params, src_d, dst_d, rows, scale);
|
||||
ggml_cuda_kernel_launch(fwht_cuda<64, T>, launch_params, src_d, dst_d, rows, scale);
|
||||
return true;
|
||||
case 128:
|
||||
ggml_cuda_kernel_launch(fwht_cuda<128>, launch_params, src_d, dst_d, rows, scale);
|
||||
ggml_cuda_kernel_launch(fwht_cuda<128, T>, launch_params, src_d, dst_d, rows, scale);
|
||||
return true;
|
||||
case 256:
|
||||
ggml_cuda_kernel_launch(fwht_cuda<256>, launch_params, src_d, dst_d, rows, scale);
|
||||
ggml_cuda_kernel_launch(fwht_cuda<256, T>, launch_params, src_d, dst_d, rows, scale);
|
||||
return true;
|
||||
case 512:
|
||||
ggml_cuda_kernel_launch(fwht_cuda<512>, launch_params, src_d, dst_d, rows, scale);
|
||||
ggml_cuda_kernel_launch(fwht_cuda<512, T>, launch_params, src_d, dst_d, rows, scale);
|
||||
return true;
|
||||
default:
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
bool ggml_cuda_op_mul_mat_use_fwht(const struct ggml_tensor * op) {
|
||||
const struct ggml_tensor * a = op->src[0];
|
||||
const struct ggml_tensor * b = op->src[1];
|
||||
|
||||
return op->op == GGML_OP_MUL_MAT && ggml_get_op_params_i32(op, 1) == GGML_HINT_SRC0_IS_HADAMARD &&
|
||||
a->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32 &&
|
||||
(b->type == GGML_TYPE_F32 || b->type == GGML_TYPE_F16) && ggml_is_contiguous(b) && ggml_is_contiguous(op) &&
|
||||
ggml_are_same_shape(b, op);
|
||||
}
|
||||
|
||||
bool ggml_cuda_op_fwht(ggml_backend_cuda_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
|
||||
GGML_ASSERT(ggml_are_same_shape(src, dst));
|
||||
if (!ggml_is_contiguous(src) || !ggml_is_contiguous(dst)) {
|
||||
return false;
|
||||
}
|
||||
if (dst->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
switch (src->type) {
|
||||
case GGML_TYPE_F32:
|
||||
return ggml_cuda_op_fwht_impl<float>(ctx, src, dst);
|
||||
case GGML_TYPE_F16:
|
||||
return ggml_cuda_op_fwht_impl<half>(ctx, src, dst);
|
||||
default:
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
#include "common.cuh"
|
||||
|
||||
bool ggml_cuda_op_mul_mat_use_fwht(const struct ggml_tensor * op);
|
||||
|
||||
// Returns whether the Fast Walsh-Hadamard transform could be used.
|
||||
bool ggml_cuda_op_fwht(ggml_backend_cuda_context & ctx, const ggml_tensor * src, ggml_tensor * dst);
|
||||
|
||||
@@ -310,13 +310,9 @@ static ggml_cuda_device_info ggml_cuda_init() {
|
||||
info.devices[id].smpb = prop.sharedMemPerBlock;
|
||||
info.devices[id].warp_size = prop.warpSize;
|
||||
|
||||
#ifndef GGML_USE_MUSA
|
||||
int supports_coop_launch = 0;
|
||||
CUDA_CHECK(cudaDeviceGetAttribute(&supports_coop_launch, cudaDevAttrCooperativeLaunch, physical_id));
|
||||
info.devices[id].supports_cooperative_launch = !!supports_coop_launch;
|
||||
#else
|
||||
info.devices[id].supports_cooperative_launch = false;
|
||||
#endif // !(GGML_USE_MUSA)
|
||||
|
||||
#if defined(GGML_USE_HIP)
|
||||
info.devices[id].smpbo = prop.sharedMemPerBlock;
|
||||
@@ -337,8 +333,6 @@ static ggml_cuda_device_info ggml_cuda_init() {
|
||||
device_vmm ? "yes" : "no", prop.warpSize,
|
||||
device_vram_mib);
|
||||
#elif defined(GGML_USE_MUSA)
|
||||
// FIXME: Ensure compatibility with varying warp sizes across different MUSA archs.
|
||||
info.devices[id].warp_size = 32;
|
||||
info.devices[id].smpbo = prop.sharedMemPerBlockOptin;
|
||||
info.devices[id].cc = GGML_CUDA_CC_OFFSET_MTHREADS + prop.major * 0x100;
|
||||
info.devices[id].cc += prop.minor * 0x10;
|
||||
@@ -1824,8 +1818,7 @@ static bool ggml_cuda_should_fuse_mul_mat_vec_q(const ggml_tensor * tensor) {
|
||||
static void ggml_cuda_mul_mat(ggml_backend_cuda_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
|
||||
GGML_TENSOR_BINARY_OP_LOCALS
|
||||
|
||||
const int32_t hint = ggml_get_op_params_i32(dst, 1);
|
||||
if (hint == GGML_HINT_SRC0_IS_HADAMARD && ggml_cuda_op_fwht(ctx, src1, dst)) {
|
||||
if (ggml_cuda_op_mul_mat_use_fwht(dst) && ggml_cuda_op_fwht(ctx, src1, dst)) {
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -3316,6 +3309,19 @@ static bool ggml_cuda_can_fuse(const struct ggml_cgraph * cgraph,
|
||||
return true;
|
||||
}
|
||||
|
||||
if (ops.size() == 2 && ops.begin()[0] == GGML_OP_RMS_NORM && ops.begin()[1] == GGML_OP_SCALE) {
|
||||
const ggml_tensor * rms_norm = cgraph->nodes[node_idx];
|
||||
const ggml_tensor * scale = cgraph->nodes[node_idx+1];
|
||||
|
||||
GGML_ASSERT(rms_norm->src[0]->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(rms_norm->type == GGML_TYPE_F32);
|
||||
|
||||
float bias;
|
||||
memcpy(&bias, (const float *) scale->op_params + 1, sizeof(float));
|
||||
|
||||
return bias == 0.0f && scale->type == GGML_TYPE_F32;
|
||||
}
|
||||
|
||||
if (ops.size() == 2 && ops.begin()[0] == GGML_OP_SSM_CONV && ops.begin()[1] == GGML_OP_UNARY
|
||||
&& unary_ops.size() == 1 && unary_ops.begin()[0] == GGML_UNARY_OP_SILU) {
|
||||
const ggml_tensor * ssm_conv = cgraph->nodes[node_idx];
|
||||
@@ -4157,6 +4163,11 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
|
||||
return 1;
|
||||
}
|
||||
|
||||
if (ggml_cuda_can_fuse(cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_SCALE }, {})) {
|
||||
ggml_cuda_op_rms_norm_scale_fused(*cuda_ctx, node, cgraph->nodes[i + 1]);
|
||||
return 1;
|
||||
}
|
||||
|
||||
if (ggml_cuda_can_fuse(cgraph, i, { GGML_OP_SSM_CONV, GGML_OP_ADD, GGML_OP_UNARY }, { GGML_UNARY_OP_SILU })) {
|
||||
ggml_cuda_op_ssm_conv(*cuda_ctx, node, cgraph->nodes[i + 1], cgraph->nodes[i + 2]);
|
||||
return 2;
|
||||
@@ -5200,7 +5211,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
|
||||
if (a->nb[0] != ggml_element_size(a) || b->nb[0] != ggml_element_size(b)) {
|
||||
return false; // TODO this could in principle be implemented though currently there is no use case.
|
||||
}
|
||||
if (b->type == GGML_TYPE_F16 && a->type != GGML_TYPE_F16) {
|
||||
if (b->type == GGML_TYPE_F16 && a->type != GGML_TYPE_F16 && !ggml_cuda_op_mul_mat_use_fwht(op)) {
|
||||
return false;
|
||||
}
|
||||
if (op->op == GGML_OP_MUL_MAT_ID && ggml_get_op_params_i32(op, 3) == GGML_PREC_F32) {
|
||||
@@ -5476,8 +5487,9 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
|
||||
|
||||
if (op->src[3]->ne[0] == 1) {
|
||||
// Mamba2
|
||||
// (kernel only supports (d_state == 128 || d_state == 256) && d_head % 16 == 0)
|
||||
return (op->src[0]->ne[0] == 128 || op->src[0]->ne[0] == 256) && op->src[0]->ne[1] % 16 == 0;
|
||||
// (kernel only supports (d_state == 96 || d_state == 128 || d_state == 256) && d_head % 16 == 0)
|
||||
const int64_t d_state = op->src[0]->ne[0];
|
||||
return (d_state == 96 || d_state == 128 || d_state == 256) && op->src[0]->ne[1] % 16 == 0;
|
||||
} else {
|
||||
if (K > 1) {
|
||||
return false;
|
||||
@@ -5563,12 +5575,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
|
||||
case GGML_OP_RWKV_WKV7:
|
||||
return true;
|
||||
case GGML_OP_GATED_DELTA_NET:
|
||||
//TODO: enable once MUSA compiler is solved https://github.com/ggml-org/llama.cpp/pull/19504#issuecomment-4018634327
|
||||
#ifdef GGML_USE_MUSA
|
||||
return false;
|
||||
#else
|
||||
return true;
|
||||
#endif // GGML_USE_MUSA
|
||||
case GGML_OP_DSV4_HC_COMB:
|
||||
return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 &&
|
||||
op->src[2]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32;
|
||||
|
||||
@@ -1631,16 +1631,16 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_mxfp4_fp4(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q4) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q4);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q4);
|
||||
|
||||
int * x_qs = (int *) x_tile;
|
||||
uint32_t * x_sc = (uint32_t *) (x_qs + 2 * MMQ_TILE_NE_K);
|
||||
|
||||
const int txi = threadIdx.x;
|
||||
|
||||
constexpr int iter_k = ggml_cuda_mmq_get_K_vram(type, J, fallback);
|
||||
constexpr int iter_k = ggml_cuda_mmq_get_K_vram(type, J, fallback, GGML_PREC_Q4);
|
||||
|
||||
constexpr int threads_per_row = iter_k / QK_MXFP4; // each thread processes 1 block
|
||||
constexpr int rows_per_warp = warp_size / threads_per_row;
|
||||
@@ -1670,12 +1670,12 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
}
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_nvfp4(
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_nvfp4(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kb0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, prec_src1);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, prec_src1);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -1729,12 +1729,12 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_nvfp4_nvfp4(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int iter_k = ggml_cuda_mmq_get_K_vram(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q4) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q4);
|
||||
constexpr int iter_k = ggml_cuda_mmq_get_K_vram(type, J, fallback, GGML_PREC_Q4);
|
||||
constexpr int threads_per_row = iter_k / QK_NVFP4; // each thread processes 1 block
|
||||
constexpr int rows_per_warp = warp_size / threads_per_row;
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q4);
|
||||
|
||||
uint32_t * x_u32 = (uint32_t *) x_tile;
|
||||
|
||||
|
||||
@@ -474,7 +474,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
}
|
||||
|
||||
// Used for Q3_K, IQ2_S, and IQ2_XS:
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma(
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
constexpr data_layout input_layout = get_input_data_layout();
|
||||
@@ -482,7 +482,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
typedef tile<16, 4, int, input_layout> tile_B;
|
||||
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, prec_src1);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
@@ -532,7 +532,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
typedef tile< 8, 4, int> tile_B;
|
||||
typedef tile<16, 8, int> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, prec_src1);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
@@ -1180,7 +1180,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
typedef tile<8, 8, int> tile_B;
|
||||
typedef tile<16, 8, float> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q4);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int ntx = rows_per_warp / tile_C::I;
|
||||
constexpr int nfrags = MMQ_TILE_NE_K / tile_A::J;
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
|
||||
#include <cstdint>
|
||||
|
||||
static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
|
||||
static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream, const ggml_prec prec_src1) {
|
||||
switch (args.type_x) {
|
||||
case GGML_TYPE_Q1_0:
|
||||
mul_mat_q_case<GGML_TYPE_Q1_0>(ctx, args, stream);
|
||||
@@ -71,9 +71,18 @@ static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, con
|
||||
break;
|
||||
// -----------------------------------------------------------------------
|
||||
case GGML_TYPE_MXFP4:
|
||||
// src1 at Q4 uses the native FP4 instructions, which are Blackwell-only
|
||||
if (prec_src1 == GGML_PREC_Q4) {
|
||||
mul_mat_q_case<GGML_TYPE_MXFP4, GGML_PREC_Q4>(ctx, args, stream);
|
||||
break;
|
||||
}
|
||||
mul_mat_q_case<GGML_TYPE_MXFP4>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_NVFP4:
|
||||
if (prec_src1 == GGML_PREC_Q4) {
|
||||
mul_mat_q_case<GGML_TYPE_NVFP4, GGML_PREC_Q4>(ctx, args, stream);
|
||||
break;
|
||||
}
|
||||
mul_mat_q_case<GGML_TYPE_NVFP4>(ctx, args, stream);
|
||||
break;
|
||||
default:
|
||||
@@ -82,6 +91,47 @@ static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, con
|
||||
}
|
||||
}
|
||||
|
||||
// overrides the src1 precision requested by the graph, "auto" keeps the requested one
|
||||
static ggml_prec ggml_cuda_mmq_get_prec_env() {
|
||||
const char * env_c = getenv("GGML_CUDA_MMQ_PREC");
|
||||
if (env_c == nullptr) {
|
||||
return GGML_PREC_UNDEFINED;
|
||||
}
|
||||
std::string env_cpp = env_c;
|
||||
for (char & c : env_cpp) {
|
||||
c = std::tolower(c);
|
||||
}
|
||||
if (env_cpp == "q4") {
|
||||
return GGML_PREC_Q4;
|
||||
}
|
||||
if (env_cpp == "q8") {
|
||||
return GGML_PREC_Q8;
|
||||
}
|
||||
if (env_cpp != "auto") {
|
||||
GGML_LOG_WARN("%s: Unknown value for GGML_CUDA_MMQ_PREC: '%s'. Available: 'q4', 'q8', 'auto'.\n", __func__, env_cpp.c_str());
|
||||
}
|
||||
return GGML_PREC_UNDEFINED;
|
||||
}
|
||||
|
||||
// src1 is quantized to Q8_1 unless the FP4 types can use 4-bit activations, in which case they
|
||||
// default to the native W4A4 instructions on Blackwell.
|
||||
static ggml_prec ggml_cuda_mmq_get_prec_src1(const ggml_tensor * src0, const ggml_tensor * dst, const int cc) {
|
||||
static const ggml_prec prec_env = ggml_cuda_mmq_get_prec_env();
|
||||
|
||||
ggml_prec prec = prec_env;
|
||||
if (prec == GGML_PREC_UNDEFINED) {
|
||||
prec = (ggml_prec) ggml_get_op_params_i32(dst, 3);
|
||||
}
|
||||
|
||||
// Q4 only for the FP4 types on Blackwell
|
||||
GGML_ASSERT(prec == GGML_PREC_UNDEFINED || prec == GGML_PREC_Q8 || prec == GGML_PREC_Q4);
|
||||
const bool can_use_q4 = (src0->type == GGML_TYPE_NVFP4 || src0->type == GGML_TYPE_MXFP4) && blackwell_mma_available(cc);
|
||||
if (prec == GGML_PREC_Q8 || !can_use_q4) {
|
||||
return GGML_PREC_Q8;
|
||||
}
|
||||
return GGML_PREC_Q4;
|
||||
}
|
||||
|
||||
void ggml_cuda_mul_mat_q(
|
||||
ggml_backend_cuda_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * ids, ggml_tensor * dst) {
|
||||
GGML_ASSERT( src1->type == GGML_TYPE_F32);
|
||||
@@ -128,7 +178,9 @@ void ggml_cuda_mul_mat_q(
|
||||
|
||||
const bool fallback = ne01 % 128 != 0;
|
||||
|
||||
const bool use_native_fp4 = blackwell_mma_available(cc) && (src0->type == GGML_TYPE_MXFP4 || src0->type == GGML_TYPE_NVFP4);
|
||||
const ggml_prec prec_src1 = ggml_cuda_mmq_get_prec_src1(src0, dst, cc);
|
||||
|
||||
const bool use_native_fp4 = prec_src1 == GGML_PREC_Q4;
|
||||
const size_t y_block_size = use_native_fp4 ? sizeof(block_fp4_mmq) : sizeof(block_q8_1_mmq);
|
||||
const size_t y_values_per_block = use_native_fp4 ? QK_FP4_MMQ : QK8_1_MMQ;
|
||||
|
||||
@@ -172,7 +224,7 @@ void ggml_cuda_mul_mat_q(
|
||||
ne02, ne12, s02, s12, s2,
|
||||
ne03, ne13, s03, s13, s3,
|
||||
ne1, ne1};
|
||||
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream);
|
||||
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream, prec_src1);
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -260,7 +312,7 @@ void ggml_cuda_mul_mat_q(
|
||||
ne03, ne13, s03, s13, s3,
|
||||
ne12, ncols_opt};
|
||||
|
||||
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream);
|
||||
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream, prec_src1);
|
||||
}
|
||||
|
||||
bool ggml_cuda_should_use_mmq(enum ggml_type type, int cc, int64_t ne11, int64_t n_experts) {
|
||||
@@ -389,5 +441,10 @@ bool ggml_cuda_should_use_mmq(enum ggml_type type, int cc, int64_t ne11, int64_t
|
||||
return n_experts > 0;
|
||||
}
|
||||
|
||||
// MUSA: the MMQ kernels compute wrong values on PH1 (MTT S5000).
|
||||
if (cc == GGML_CUDA_CC_PH1) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return (!GGML_CUDA_CC_IS_CDNA(cc)) || ne11 < MMQ_DP4A_MAX_BATCH_SIZE;
|
||||
}
|
||||
|
||||
+120
-106
@@ -227,7 +227,7 @@ struct ggml_cuda_mmq_config {
|
||||
|
||||
#undef CASE
|
||||
|
||||
static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
if (GGML_CUDA_CC_IS_AMD(cc)) {
|
||||
if (GGML_CUDA_CC_IS_GCN(cc)) {
|
||||
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
|
||||
@@ -247,6 +247,10 @@ static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type ty
|
||||
return ggml_cuda_mmq_get_config_rdna2(type, J, fallback);
|
||||
}
|
||||
if (blackwell_mma_available(cc)) {
|
||||
// only src1 at Q4 uses the native FP4 config, higher precisions keep src1 at Q8_1
|
||||
if (prec_src1 != GGML_PREC_Q4 && (type == GGML_TYPE_NVFP4 || type == GGML_TYPE_MXFP4)) {
|
||||
return ggml_cuda_mmq_get_config_ampere(type, J, fallback);
|
||||
}
|
||||
return ggml_cuda_mmq_get_config_blackwell(type, J, fallback);
|
||||
}
|
||||
if (ggml_cuda_highest_compiled_arch(cc) >= GGML_CUDA_CC_VOLTA) {
|
||||
@@ -258,7 +262,7 @@ static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type ty
|
||||
return ggml_cuda_mmq_get_config_pascal_older(type, J, fallback);
|
||||
}
|
||||
|
||||
static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_type type, int J, bool fallback) {
|
||||
static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
#ifdef GGML_USE_HIP
|
||||
#ifdef GCN
|
||||
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
|
||||
@@ -275,6 +279,10 @@ static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_t
|
||||
#endif // CDNA
|
||||
#else
|
||||
#ifdef BLACKWELL_MMA_AVAILABLE
|
||||
// only src1 at Q4 uses the native FP4 config, higher precisions keep src1 at Q8_1
|
||||
if (prec_src1 != GGML_PREC_Q4 && (type == GGML_TYPE_NVFP4 || type == GGML_TYPE_MXFP4)) {
|
||||
return ggml_cuda_mmq_get_config_ampere(type, J, fallback);
|
||||
}
|
||||
return ggml_cuda_mmq_get_config_blackwell(type, J, fallback);
|
||||
#elif __CUDA_ARCH__ >= GGML_CUDA_CC_VOLTA
|
||||
return ggml_cuda_mmq_get_config_ampere(type, J, fallback);
|
||||
@@ -284,79 +292,71 @@ static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_t
|
||||
return ggml_cuda_mmq_get_config_pascal_older(type, J, fallback);
|
||||
#endif // BLACKWELL_MMA_AVAILABLE
|
||||
#endif // GGML_USE_HIP
|
||||
GGML_UNUSED_VARS(type, J, fallback);
|
||||
GGML_UNUSED_VARS(type, J, fallback, prec_src1);
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_type(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).type;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_type(ggml_type type, int J, bool fallback) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback).type;
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_type(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).type;
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_nthreads(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).nthreads;
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_nthreads(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).nthreads;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_nthreads(ggml_type type, int J, bool fallback) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback).nthreads;
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_occupancy(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).occupancy;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_occupancy(ggml_type type, int J, bool fallback) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback).occupancy;
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_occupancy(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).occupancy;
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_I(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).I;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_I(ggml_type type, int J, bool fallback) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback).I;
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_I(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).I;
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_J(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).J;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_J(ggml_type type, int J, bool fallback) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback).J;
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_J(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).J;
|
||||
}
|
||||
|
||||
static __host__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).sram_layout;
|
||||
}
|
||||
|
||||
static constexpr __device__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(ggml_type type, int J, bool fallback) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback).sram_layout;
|
||||
static constexpr __device__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).sram_layout;
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_K_vram(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).K_vram;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_K_vram(ggml_type type, int J, bool fallback) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback).K_vram;
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_K_vram(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).K_vram;
|
||||
}
|
||||
|
||||
static __host__ bool ggml_cuda_mmq_get_stream_k(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).stream_k;
|
||||
}
|
||||
|
||||
static constexpr __device__ bool ggml_cuda_mmq_get_stream_k(ggml_type type, int J, bool fallback) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback).stream_k;
|
||||
static constexpr __device__ bool ggml_cuda_mmq_get_stream_k(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).stream_k;
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_fallback(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).fallback;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_fallback(ggml_type type, int J, bool fallback) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback).fallback;
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_fallback(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).fallback;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------------------------
|
||||
@@ -365,8 +365,8 @@ static __host__ int ggml_cuda_mmq_get_sram_stride(const ggml_type type, const in
|
||||
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, cc));
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_sram_stride(ggml_type type, int J, bool fallback) {
|
||||
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback));
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_sram_stride(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, prec_src1));
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_J_max(const ggml_type type, const bool fallback, const int cc, const int64_t ne11) {
|
||||
@@ -541,9 +541,9 @@ struct ggml_cuda_mmq_util_funcs {
|
||||
vdr(vdr), load_tiles(load_tiles), vec_dot(vec_dot), write_back(write_back) {}
|
||||
};
|
||||
|
||||
template <ggml_type type, int J, bool fallback>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_funcs() {
|
||||
if (!ggml_cuda_mmq_get_config(type, J, fallback).use_mma_data_layout()) {
|
||||
if (!ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).use_mma_data_layout()) {
|
||||
switch (type) {
|
||||
case GGML_TYPE_Q1_0:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
@@ -690,17 +690,23 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
|
||||
#ifdef BLACKWELL_MMA_AVAILABLE
|
||||
switch (type) {
|
||||
case GGML_TYPE_MXFP4:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_mxfp4_fp4<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_fp4_fp4_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
if (prec_src1 == GGML_PREC_Q4) {
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_mxfp4_fp4<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_fp4_fp4_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
}
|
||||
break;
|
||||
case GGML_TYPE_NVFP4:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_nvfp4_nvfp4<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_fp4_fp4_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
if (prec_src1 == GGML_PREC_Q4) {
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_nvfp4_nvfp4<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_fp4_fp4_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
}
|
||||
break;
|
||||
default:
|
||||
break;
|
||||
}
|
||||
@@ -841,37 +847,37 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
|
||||
case GGML_TYPE_NVFP4:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback, prec_src1>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
default:
|
||||
return ggml_cuda_mmq_util_funcs(1, nullptr, nullptr, nullptr);
|
||||
}
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_vdr() {
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback>().vdr;
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().vdr;
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
static constexpr __device__ ggml_cuda_mmq_load_tiles_t ggml_cuda_mmq_get_load_tiles() {
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback>().load_tiles;
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().load_tiles;
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
static constexpr __device__ ggml_cuda_mmq_vec_dot_t ggml_cuda_mmq_get_vec_dot() {
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback>().vec_dot;
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().vec_dot;
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
static constexpr __device__ ggml_cuda_mmq_write_back_t ggml_cuda_mmq_get_write_back() {
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback>().write_back;
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().write_back;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------------------------
|
||||
|
||||
template <ggml_type type, int J, bool fallback, bool fixup>
|
||||
template <ggml_type type, int J, bool fallback, bool fixup, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
static __device__ __forceinline__ void mul_mat_q_process_tile(
|
||||
const char * __restrict__ x, const int offset_x, const int * __restrict__ y,
|
||||
const int * __restrict__ ids_dst, float * __restrict__ dst, float * __restrict__ tmp_fixup,
|
||||
@@ -880,25 +886,27 @@ static __device__ __forceinline__ void mul_mat_q_process_tile(
|
||||
const int tile_x_max_i, const int tile_y_max_j, const int kb0_start, const int kb0_stop) {
|
||||
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1) / warp_size;
|
||||
constexpr int qk = ggml_cuda_type_traits<type>::qk;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr ggml_cuda_mmq_load_tiles_t load_tiles = ggml_cuda_mmq_get_load_tiles<type, J, fallback>();
|
||||
constexpr ggml_cuda_mmq_vec_dot_t vec_dot = ggml_cuda_mmq_get_vec_dot<type, J, fallback>();
|
||||
constexpr ggml_cuda_mmq_write_back_t write_back = ggml_cuda_mmq_get_write_back<type, J, fallback>();
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, prec_src1);
|
||||
constexpr ggml_cuda_mmq_load_tiles_t load_tiles = ggml_cuda_mmq_get_load_tiles<type, J, fallback, prec_src1>();
|
||||
constexpr ggml_cuda_mmq_vec_dot_t vec_dot = ggml_cuda_mmq_get_vec_dot<type, J, fallback, prec_src1>();
|
||||
constexpr ggml_cuda_mmq_write_back_t write_back = ggml_cuda_mmq_get_write_back<type, J, fallback, prec_src1>();
|
||||
|
||||
extern __shared__ int data_mul_mat_q[];
|
||||
int * tile_y = data_mul_mat_q + J;
|
||||
int * tile_x = tile_y + GGML_PAD(J*MMQ_TILE_Y_K, nwarps*warp_size);
|
||||
|
||||
#if defined(BLACKWELL_MMA_AVAILABLE)
|
||||
// FP4 tile stores 8 blocks
|
||||
constexpr int ne_block = (type == GGML_TYPE_MXFP4 || type == GGML_TYPE_NVFP4) ? QK_FP4_MMQ : QK8_1_MMQ;
|
||||
// FP4 tile stores 8 blocks. src1 above Q4 uses the generic
|
||||
// Q8_1 tile layout instead of the packed FP4 tile.
|
||||
constexpr int ne_block = ((type == GGML_TYPE_MXFP4 || type == GGML_TYPE_NVFP4) && prec_src1 == GGML_PREC_Q4) ?
|
||||
QK_FP4_MMQ : QK8_1_MMQ;
|
||||
#else
|
||||
constexpr int ne_block = QK8_1_MMQ;
|
||||
#endif // defined(BLACKWELL_MMA_AVAILABLE)
|
||||
|
||||
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback);
|
||||
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback, prec_src1);
|
||||
constexpr int blocks_per_iter = ITER_K / qk;
|
||||
|
||||
float sum[J*I / (nwarps*warp_size)] = {0.0f};
|
||||
@@ -950,8 +958,8 @@ static __device__ __forceinline__ void mul_mat_q_process_tile(
|
||||
|
||||
// The mul_mat_q kernel implements "stream-k" work partitioning as described in https://arxiv.org/abs/2301.03598
|
||||
|
||||
template <ggml_type type, int J, bool fallback>
|
||||
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback), ggml_cuda_mmq_get_occupancy(type, J, fallback))
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1), ggml_cuda_mmq_get_occupancy(type, J, fallback, prec_src1))
|
||||
static __global__ void mul_mat_q(
|
||||
const char * __restrict__ x, const int * __restrict__ y, const int32_t * __restrict__ ids_dst,
|
||||
const int32_t * __restrict__ expert_bounds, float * __restrict__ dst, float * __restrict__ tmp_fixup,
|
||||
@@ -962,15 +970,15 @@ static __global__ void mul_mat_q(
|
||||
const uint3 ntx) {
|
||||
|
||||
// Skip unused template specializations for faster compilation:
|
||||
if (ggml_cuda_mmq_get_config(type, J, fallback).type == GGML_TYPE_COUNT) {
|
||||
if (ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).type == GGML_TYPE_COUNT) {
|
||||
NO_DEVICE_CODE;
|
||||
return;
|
||||
}
|
||||
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1) / warp_size;
|
||||
constexpr int qk = ggml_cuda_type_traits<type>::qk;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, prec_src1);
|
||||
|
||||
const uint32_t nty = (nrows_x + I - 1) / I; // Number of tiles y
|
||||
|
||||
@@ -990,7 +998,7 @@ static __global__ void mul_mat_q(
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
if constexpr (!ggml_cuda_mmq_get_stream_k(type, J, fallback)) {
|
||||
if constexpr (!ggml_cuda_mmq_get_stream_k(type, J, fallback, prec_src1)) {
|
||||
const uint2 tmp2 = fast_div_modulo(blockIdx.z, nchannels_y);
|
||||
const int wt = tmp2.x;
|
||||
const int zt = tmp2.y;
|
||||
@@ -1053,14 +1061,14 @@ static __global__ void mul_mat_q(
|
||||
const int offset_x = fastdiv(wt, sample_ratio)*stride_sample_x + fastdiv(zt, channel_ratio)*stride_channel_x + it*I*stride_row_x;
|
||||
|
||||
constexpr bool fixup = false;
|
||||
mul_mat_q_process_tile<type, J, fallback, fixup>
|
||||
mul_mat_q_process_tile<type, J, fallback, fixup, prec_src1>
|
||||
(x, offset_x, y + offset_y, ids_dst_shared, dst + offset_dst, tmp_fixup, y_scale_tile,
|
||||
stride_row_x, ncols_y, stride_col_dst,
|
||||
tile_x_max_i, tile_y_max_j, 0, blocks_per_ne00.z);
|
||||
return;
|
||||
}
|
||||
|
||||
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback);
|
||||
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback, prec_src1);
|
||||
constexpr int blocks_per_iter = ITER_K / qk;
|
||||
|
||||
// kbc == k block continuous, current index in continuous ijk space.
|
||||
@@ -1147,7 +1155,7 @@ static __global__ void mul_mat_q(
|
||||
const int offset_x = fastdiv(wt, sample_ratio)*stride_sample_x + fastdiv(zt, channel_ratio)*stride_channel_x + it*I*stride_row_x;
|
||||
|
||||
constexpr bool fixup = false; // All but (potentially) the last iterations write their data to dst rather than the fixup buffer.
|
||||
mul_mat_q_process_tile<type, J, fallback, fixup>
|
||||
mul_mat_q_process_tile<type, J, fallback, fixup, prec_src1>
|
||||
(x, offset_x, y + offset_y, ids_dst_shared, dst + offset_dst, tmp_fixup, y_scale_tile,
|
||||
stride_row_x, ncols_y, stride_col_dst,
|
||||
tile_x_max_i, tile_y_max_j, kb0_start, kb0_stop);
|
||||
@@ -1231,24 +1239,24 @@ static __global__ void mul_mat_q(
|
||||
const int offset_x = fastdiv(wt, sample_ratio)*stride_sample_x + fastdiv(zt, channel_ratio)*stride_channel_x + it*I*stride_row_x;
|
||||
|
||||
constexpr bool fixup = true; // Last index writes its data to fixup buffer to avoid data races with other blocks.
|
||||
mul_mat_q_process_tile<type, J, fallback, fixup>
|
||||
mul_mat_q_process_tile<type, J, fallback, fixup, prec_src1>
|
||||
(x, offset_x, y + offset_y, ids_dst_shared, dst + offset_dst, tmp_fixup, y_scale_tile,
|
||||
stride_row_x, ncols_y, stride_col_dst,
|
||||
tile_x_max_i, tile_y_max_j, kb0_start, kb0_stop);
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback>
|
||||
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback)/2, 1)
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1)/2, 1)
|
||||
static __global__ void mul_mat_q_stream_k_fixup(
|
||||
const int32_t * __restrict__ ids_dst, const int32_t * __restrict__ expert_bounds, float * __restrict__ dst,
|
||||
float * __restrict__ tmp_last_tile, const uint3 blocks_per_ne00, const int nrows_x, const int ncols_dst,
|
||||
const int stride_col_dst, const uint3 nchannels_y, const int stride_channel_dst, const uint3 nsamples_y,
|
||||
const int stride_sample_dst, const uint3 ntx) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = (ggml_cuda_mmq_get_nthreads(type, J, fallback) / 2) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = (ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1) / 2) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, prec_src1);
|
||||
constexpr int qk = ggml_cuda_type_traits<type>::qk;
|
||||
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback);
|
||||
constexpr int ITER_K = ggml_cuda_mmq_get_K_vram(type, J, fallback, prec_src1);
|
||||
constexpr int blocks_per_iter = ITER_K / qk;
|
||||
|
||||
float sum[J / nwarps] = {0.0f};
|
||||
@@ -1392,22 +1400,22 @@ static size_t mmq_get_nbytes_shared(const ggml_cuda_mmq_config & config, const i
|
||||
return nbs_ids + nbs_x + GGML_PAD(nbs_y, config.nthreads*sizeof(int));
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
|
||||
const int id = ggml_cuda_get_device();
|
||||
const int cc = ggml_cuda_info().devices[id].cc;
|
||||
const int nsm = ggml_cuda_info().devices[id].nsm;
|
||||
const int warp_size = ggml_cuda_info().devices[id].warp_size;
|
||||
|
||||
const ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config(type, J, fallback, cc);
|
||||
const ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1);
|
||||
GGML_ASSERT(config.nthreads % warp_size == 0);
|
||||
const int nwarps = config.nthreads / warp_size;
|
||||
const int nbytes_shared = mmq_get_nbytes_shared(config, cc);
|
||||
|
||||
const dim3 block_dims(warp_size, nwarps, 1);
|
||||
|
||||
CUDA_SET_SHARED_MEMORY_LIMIT((mul_mat_q<type, J, false>), nbytes_shared);
|
||||
CUDA_SET_SHARED_MEMORY_LIMIT((mul_mat_q<type, J, true>), nbytes_shared);
|
||||
CUDA_SET_SHARED_MEMORY_LIMIT((mul_mat_q<type, J, false, prec_src1>), nbytes_shared);
|
||||
CUDA_SET_SHARED_MEMORY_LIMIT((mul_mat_q<type, J, true, prec_src1>), nbytes_shared);
|
||||
|
||||
const int nty = (args.nrows_x + config.I - 1) / config.I;
|
||||
const int ntx = (args.ncols_max + config.J - 1) / config.J;
|
||||
@@ -1426,8 +1434,8 @@ static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & a
|
||||
const uint3 channel_ratio_fd = init_fastdiv_values(channel_ratio);
|
||||
const uint3 sample_ratio_fd = init_fastdiv_values(sample_ratio);
|
||||
|
||||
if (!ggml_cuda_mmq_get_stream_k(type, J, fallback, cc)) {
|
||||
mul_mat_q<type, J, fallback><<<block_nums_xy_tiling, block_dims, nbytes_shared, stream>>>
|
||||
if (!config.stream_k) {
|
||||
mul_mat_q<type, J, fallback, prec_src1><<<block_nums_xy_tiling, block_dims, nbytes_shared, stream>>>
|
||||
(args.x, args.y, args.ids_dst, args.expert_bounds, args.dst, nullptr, args.y_scale,
|
||||
blocks_per_ne00_fd, args.nrows_x, args.ncols_dst, args.stride_row_x, args.ncols_y, args.nrows_dst,
|
||||
channel_ratio_fd, nchannels_y_fd, args.stride_channel_x, args.stride_channel_y, args.stride_channel_dst,
|
||||
@@ -1456,7 +1464,7 @@ static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & a
|
||||
const dim3 block_nums_fixup(block_nums_stream_k.x, config.I/warp_size, 1);
|
||||
const dim3 block_dims_fixup(block_dims.x, block_dims.y/2, block_dims.z);
|
||||
|
||||
mul_mat_q<type, J, fallback><<<block_nums_stream_k, block_dims, nbytes_shared, stream>>>
|
||||
mul_mat_q<type, J, fallback, prec_src1><<<block_nums_stream_k, block_dims, nbytes_shared, stream>>>
|
||||
(args.x, args.y, args.ids_dst, args.expert_bounds, args.dst, tmp_fixup.ptr, args.y_scale,
|
||||
blocks_per_ne00_fd, args.nrows_x, args.ncols_dst, args.stride_row_x, args.ncols_y, args.nrows_dst,
|
||||
channel_ratio_fd, nchannels_y_fd, args.stride_channel_x, args.stride_channel_y, args.stride_channel_dst,
|
||||
@@ -1468,13 +1476,13 @@ static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & a
|
||||
}
|
||||
|
||||
CUDA_CHECK(cudaGetLastError());
|
||||
mul_mat_q_stream_k_fixup<type, J, fallback><<<block_nums_fixup, block_dims_fixup, 0, stream>>>
|
||||
mul_mat_q_stream_k_fixup<type, J, fallback, prec_src1><<<block_nums_fixup, block_dims_fixup, 0, stream>>>
|
||||
(args.ids_dst, args.expert_bounds, args.dst, tmp_fixup.ptr, blocks_per_ne00_fd, args.nrows_x, args.ncols_dst,
|
||||
args.nrows_dst, nchannels_y_fd, args.stride_channel_dst, nsamples_y_fd, args.stride_sample_dst,
|
||||
ntx_fd);
|
||||
}
|
||||
|
||||
template <ggml_type type, bool fallback>
|
||||
template <ggml_type type, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
|
||||
const int id = ggml_cuda_get_device();
|
||||
const int cc = ggml_cuda_info().devices[id].cc;
|
||||
@@ -1484,7 +1492,7 @@ void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args,
|
||||
int ntiles_J_best = INT_MAX;
|
||||
|
||||
for (int J = 8; J <= 128 && ntiles_J_best > 1; J += 8) {
|
||||
const ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config(type, J, fallback, cc);
|
||||
const ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1);
|
||||
if (config.type == GGML_TYPE_COUNT) {
|
||||
continue;
|
||||
}
|
||||
@@ -1503,52 +1511,52 @@ void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args,
|
||||
|
||||
switch (J_best) {
|
||||
case 8:
|
||||
launch_mul_mat_q<type, 8, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 8, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 16:
|
||||
launch_mul_mat_q<type, 16, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 16, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 24:
|
||||
launch_mul_mat_q<type, 24, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 24, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 32:
|
||||
launch_mul_mat_q<type, 32, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 32, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 40:
|
||||
launch_mul_mat_q<type, 40, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 40, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 48:
|
||||
launch_mul_mat_q<type, 48, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 48, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 56:
|
||||
launch_mul_mat_q<type, 56, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 56, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 64:
|
||||
launch_mul_mat_q<type, 64, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 64, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 72:
|
||||
launch_mul_mat_q<type, 72, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 72, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 80:
|
||||
launch_mul_mat_q<type, 80, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 80, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 88:
|
||||
launch_mul_mat_q<type, 88, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 88, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 96:
|
||||
launch_mul_mat_q<type, 96, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 96, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 104:
|
||||
launch_mul_mat_q<type, 104, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 104, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 112:
|
||||
launch_mul_mat_q<type, 112, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 112, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 120:
|
||||
launch_mul_mat_q<type, 120, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 120, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
case 128:
|
||||
launch_mul_mat_q<type, 128, fallback>(ctx, args, stream);
|
||||
launch_mul_mat_q<type, 128, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
default:
|
||||
fprintf(stderr, "J_best=%d\n", J_best);
|
||||
@@ -1557,20 +1565,24 @@ void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args,
|
||||
}
|
||||
}
|
||||
|
||||
template <ggml_type type>
|
||||
template <ggml_type type, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
void mul_mat_q_case(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
|
||||
if (args.nrows_x % 128 == 0) {
|
||||
constexpr bool fallback = false;
|
||||
mul_mat_q_switch_J<type, fallback>(ctx, args, stream);
|
||||
mul_mat_q_switch_J<type, fallback, prec_src1>(ctx, args, stream);
|
||||
} else {
|
||||
constexpr bool fallback = true;
|
||||
mul_mat_q_switch_J<type, fallback>(ctx, args, stream);
|
||||
mul_mat_q_switch_J<type, fallback, prec_src1>(ctx, args, stream);
|
||||
}
|
||||
}
|
||||
|
||||
#define DECL_MMQ_CASE(type) \
|
||||
template void mul_mat_q_case<type>(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) \
|
||||
|
||||
// FP4 variant: uses native FP4 MMA instead of keeping src1 at Q8_1.
|
||||
#define DECL_MMQ_CASE_W4A4(type) \
|
||||
template void mul_mat_q_case<type, GGML_PREC_Q4>(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) \
|
||||
|
||||
extern DECL_MMQ_CASE(GGML_TYPE_Q1_0);
|
||||
extern DECL_MMQ_CASE(GGML_TYPE_Q2_0);
|
||||
extern DECL_MMQ_CASE(GGML_TYPE_Q4_0);
|
||||
@@ -1596,6 +1608,8 @@ extern DECL_MMQ_CASE(GGML_TYPE_IQ4_XS);
|
||||
// -----------------------------------------
|
||||
extern DECL_MMQ_CASE(GGML_TYPE_MXFP4);
|
||||
extern DECL_MMQ_CASE(GGML_TYPE_NVFP4);
|
||||
extern DECL_MMQ_CASE_W4A4(GGML_TYPE_MXFP4);
|
||||
extern DECL_MMQ_CASE_W4A4(GGML_TYPE_NVFP4);
|
||||
|
||||
// -------------------------------------------------------------------------------------------------------------------------
|
||||
|
||||
|
||||
+44
-11
@@ -73,7 +73,7 @@ static __global__ void group_norm_f32(const float * x, float * dst, const int gr
|
||||
}
|
||||
}
|
||||
|
||||
template <int block_size, bool do_multiply = false, bool do_add = false>
|
||||
template <int block_size, bool do_multiply = false, bool do_add = false, bool do_scale = false>
|
||||
static __global__ void rms_norm_f32(const float * x,
|
||||
float * dst,
|
||||
const int ncols,
|
||||
@@ -96,7 +96,8 @@ static __global__ void rms_norm_f32(const float * x,
|
||||
const uint3 add_ncols_packed = make_uint3(0, 0, 0),
|
||||
const uint3 add_nrows_packed = make_uint3(0, 0, 0),
|
||||
const uint3 add_nchannels_packed = make_uint3(0, 0, 0),
|
||||
const uint3 add_nsamples_packed = make_uint3(0, 0, 0)) {
|
||||
const uint3 add_nsamples_packed = make_uint3(0, 0, 0),
|
||||
const float scale_out = 1.0f) {
|
||||
ggml_cuda_pdl_lc();
|
||||
const int nrows = gridDim.x;
|
||||
const int nchannels = gridDim.y;
|
||||
@@ -107,6 +108,7 @@ static __global__ void rms_norm_f32(const float * x,
|
||||
const int tid = threadIdx.x;
|
||||
|
||||
static_assert(!do_add || do_multiply, "fusing add is not supported without multiplying");
|
||||
static_assert(!do_scale || !do_multiply, "fusing scale is not supported with multiplying");
|
||||
|
||||
x += sample*stride_sample + channel*stride_channel + row*stride_row;
|
||||
dst += ((sample*nchannels + channel)*nrows + row)*ncols;
|
||||
@@ -148,6 +150,8 @@ static __global__ void rms_norm_f32(const float * x,
|
||||
} else if constexpr (do_multiply) {
|
||||
const int mul_col = fastmodulo(col, mul_ncols_packed);
|
||||
dst[col] = scale * x[col] * mul[mul_col];
|
||||
} else if constexpr (do_scale) {
|
||||
dst[col] = scale_out * (scale * x[col]);
|
||||
} else {
|
||||
dst[col] = scale * x[col];
|
||||
}
|
||||
@@ -301,25 +305,27 @@ static void group_norm_f32_cuda(
|
||||
}
|
||||
}
|
||||
|
||||
template <bool do_scale = false>
|
||||
static void rms_norm_f32_cuda(
|
||||
const float * x, float * dst, const int ncols, const int nrows, const int nchannels, const int nsamples,
|
||||
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps, cudaStream_t stream) {
|
||||
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps, cudaStream_t stream,
|
||||
const float scale_out = 1.0f) {
|
||||
const dim3 blocks_num(nrows, nchannels, nsamples);
|
||||
if (ncols < 1024) {
|
||||
const dim3 block_dims(256, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = {blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
|
||||
ggml_cuda_kernel_launch(rms_norm_f32<256, false>, launch_params,
|
||||
ggml_cuda_kernel_launch(rms_norm_f32<256, false, false, do_scale>, launch_params,
|
||||
x, dst, ncols, stride_row, stride_channel, stride_sample, eps,
|
||||
// underlying cudaLaunchKernelEx does not support default params
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0),
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0));
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), scale_out);
|
||||
} else {
|
||||
const dim3 block_dims(1024, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
|
||||
ggml_cuda_kernel_launch(rms_norm_f32<1024, false>, launch_params, x, dst, ncols, stride_row, stride_channel, stride_sample, eps,
|
||||
ggml_cuda_kernel_launch(rms_norm_f32<1024, false, false, do_scale>, launch_params, x, dst, ncols, stride_row, stride_channel, stride_sample, eps,
|
||||
// underlying cudaLaunchKernelEx does not support default params
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0),
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0));
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), scale_out);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -367,7 +373,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
|
||||
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
|
||||
// underlying cudaLaunchKernelEx does not support default params
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0));
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), 1.0f);
|
||||
} else {
|
||||
const dim3 block_dims(1024, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
|
||||
@@ -375,7 +381,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
|
||||
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
|
||||
// underlying cudaLaunchKernelEx does not support default params
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0));
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), 1.0f);
|
||||
}
|
||||
} else {
|
||||
const uint3 mul_ncols_packed = init_fastdiv_values(mul_ncols);
|
||||
@@ -394,7 +400,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
|
||||
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed, add,
|
||||
add_stride_row, add_stride_channel, add_stride_sample, add_ncols_packed, add_nrows_packed,
|
||||
add_nchannels_packed, add_nsamples_packed);
|
||||
add_nchannels_packed, add_nsamples_packed, 1.0f);
|
||||
} else {
|
||||
const dim3 block_dims(1024, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
|
||||
@@ -402,7 +408,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
|
||||
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed, add,
|
||||
add_stride_row, add_stride_channel, add_stride_sample, add_ncols_packed, add_nrows_packed,
|
||||
add_nchannels_packed, add_nsamples_packed);
|
||||
add_nchannels_packed, add_nsamples_packed, 1.0f);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -499,6 +505,33 @@ void ggml_cuda_op_rms_norm(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
rms_norm_f32_cuda(src0_d, dst_d, ne00, ne01, ne02, ne03, s01, s02, s03, eps, stream);
|
||||
}
|
||||
|
||||
void ggml_cuda_op_rms_norm_scale_fused(ggml_backend_cuda_context & ctx, ggml_tensor * dst, ggml_tensor * scale_tensor) {
|
||||
const ggml_tensor * src0 = dst->src[0];
|
||||
const float * src0_d = (const float *) src0->data;
|
||||
float * dst_d = (float *) scale_tensor->data;
|
||||
cudaStream_t stream = ctx.stream();
|
||||
|
||||
GGML_ASSERT(src0->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(scale_tensor->type == GGML_TYPE_F32);
|
||||
|
||||
GGML_TENSOR_UNARY_OP_LOCALS;
|
||||
|
||||
float eps;
|
||||
memcpy(&eps, dst->op_params, sizeof(float));
|
||||
GGML_ASSERT(eps >= 0.0f);
|
||||
|
||||
float scale;
|
||||
memcpy(&scale, (const float *) scale_tensor->op_params + 0, sizeof(float));
|
||||
|
||||
const size_t ts0 = ggml_type_size(src0->type);
|
||||
GGML_ASSERT(nb00 == ts0);
|
||||
const int64_t s01 = nb01 / ts0;
|
||||
const int64_t s02 = nb02 / ts0;
|
||||
const int64_t s03 = nb03 / ts0;
|
||||
|
||||
rms_norm_f32_cuda<true>(src0_d, dst_d, ne00, ne01, ne02, ne03, s01, s02, s03, eps, stream, scale);
|
||||
}
|
||||
|
||||
void ggml_cuda_op_rms_norm_fused(ggml_backend_cuda_context & ctx, ggml_tensor * dst, ggml_tensor * mul_tensor) {
|
||||
const ggml_tensor * rms_norm_src = (ggml_tensor *) dst->src[0];
|
||||
float eps = 0.0f;
|
||||
|
||||
@@ -8,6 +8,8 @@ void ggml_cuda_op_rms_norm(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
|
||||
|
||||
void ggml_cuda_op_rms_norm_fused(ggml_backend_cuda_context & ctx, ggml_tensor * dst, ggml_tensor * mul_tensor);
|
||||
|
||||
void ggml_cuda_op_rms_norm_scale_fused(ggml_backend_cuda_context & ctx, ggml_tensor * dst, ggml_tensor * scale_tensor);
|
||||
|
||||
void ggml_cuda_op_rms_norm_fused_add(ggml_backend_cuda_context & ctx,
|
||||
ggml_tensor * dst,
|
||||
ggml_tensor * mul_tensor,
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && CUDART_VERSION >= 11070
|
||||
#if !defined(GGML_USE_HIP) && (defined(GGML_USE_MUSA) || CUDART_VERSION >= 11070)
|
||||
#define USE_CUB
|
||||
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && CUDART_VERSION >= 11070
|
||||
#endif // !defined(GGML_USE_HIP) && (defined(GGML_USE_MUSA) || CUDART_VERSION >= 11070)
|
||||
|
||||
#ifdef USE_CUB
|
||||
#include <cub/cub.cuh>
|
||||
@@ -163,13 +163,19 @@ __global__ void __launch_bounds__(d_state, 1)
|
||||
const int lane = threadIdx.x % WARP_SIZE;
|
||||
const int warp_idx = blockIdx.x * c_factor + warp;
|
||||
|
||||
ggml_cuda_pdl_sync();
|
||||
|
||||
// the last block can have unused warps when n_head*d_head is not a multiple of c_factor
|
||||
if (warp_idx >= n_head * d_head) {
|
||||
return;
|
||||
}
|
||||
|
||||
const int head_idx = warp_idx / d_head;
|
||||
const int head_off = (warp_idx % d_head) * sizeof(float);
|
||||
const int seq_idx = blockIdx.y;
|
||||
|
||||
const int group_off = (head_idx / (n_head / n_group)) * d_state * sizeof(float);
|
||||
|
||||
ggml_cuda_pdl_sync();
|
||||
// TODO: refactor strides to be in elements/floats instead of bytes to be cleaner and consistent with the rest of the codebase
|
||||
const float * s0_warp = (const float *) ((const char *) src0 + src6[seq_idx] * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
|
||||
const float * x_warp = (const float *) ((const char *) src1 + (seq_idx * src1_nb3) + (warp_idx * sizeof(float)));
|
||||
@@ -246,7 +252,17 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
|
||||
// NOTE: if you change conditions here, be sure to update the corresponding supports_op condition!
|
||||
if (src3_nb1 == sizeof(float)) {
|
||||
// Mamba-2
|
||||
if (d_state == 128) {
|
||||
if (d_state == 96) {
|
||||
constexpr int threads = 96;
|
||||
constexpr int num_warps = threads/WARP_SIZE;
|
||||
|
||||
const dim3 blocks((n_head * head_dim + (num_warps - 1)) / num_warps, n_seq, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(blocks, threads, 0, stream);
|
||||
ggml_cuda_kernel_launch(ssm_scan_f32_group<96/WARP_SIZE, 96>, launch_params,
|
||||
src0, src1, src2, src3, src4, src5, src6, dst,
|
||||
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
|
||||
} else if (d_state == 128) {
|
||||
constexpr int threads = 128;
|
||||
constexpr int num_warps = threads/WARP_SIZE;
|
||||
|
||||
@@ -267,7 +283,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
|
||||
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
|
||||
} else {
|
||||
GGML_ABORT("doesn't support d_state!=(128 or 256).");
|
||||
GGML_ABORT("doesn't support d_state!=(96, 128 or 256).");
|
||||
}
|
||||
} else {
|
||||
// Mamba-1
|
||||
@@ -342,7 +358,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
|
||||
}
|
||||
}
|
||||
|
||||
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
#if !defined(GGML_USE_HIP)
|
||||
// ============================================================================
|
||||
// SSD (State Space Duality) kernels for Mamba-2 prefill (n_tok > SSM_SSD_MIN_TOKENS)
|
||||
//
|
||||
@@ -821,7 +837,7 @@ void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
GGML_ASSERT(src5->nb[2] <= (size_t)INT_MAX);
|
||||
GGML_ASSERT(src5->nb[3] <= (size_t)INT_MAX);
|
||||
|
||||
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
#if !defined(GGML_USE_HIP)
|
||||
// Mamba-2 with scalar A per head: use SSD matmul path for long sequences.
|
||||
// Requires NVIDIA Turing+ otherwise fallback to scan.
|
||||
const bool is_mamba2 = (src3->nb[1] == sizeof(float));
|
||||
|
||||
@@ -50,6 +50,12 @@ SOURCE_MMQ = """// This file has been autogenerated by generate_cu_files.py, do
|
||||
DECL_MMQ_CASE({type});
|
||||
"""
|
||||
|
||||
TYPES_MMQ_W4A4 = ["GGML_TYPE_MXFP4", "GGML_TYPE_NVFP4"]
|
||||
|
||||
SOURCE_MMQ_W4A4 = """
|
||||
DECL_MMQ_CASE_W4A4({type});
|
||||
"""
|
||||
|
||||
SOURCE_MMF = """// This file has been autogenerated by generate_cu_files.py, do not edit manually.
|
||||
|
||||
#include "../mmf.cuh"
|
||||
@@ -105,6 +111,8 @@ for ncols in [8, 16, 32, 64]:
|
||||
for type in TYPES_MMQ:
|
||||
with open(f"mmq-instance-{get_short_name(type)}.cu", "w") as f:
|
||||
f.write(SOURCE_MMQ.format(type=type))
|
||||
if type in TYPES_MMQ_W4A4:
|
||||
f.write(SOURCE_MMQ_W4A4.format(type=type))
|
||||
|
||||
for type in range(1, 17):
|
||||
with open(f"mmf-instance-ncols_{type}.cu", "w") as f:
|
||||
|
||||
@@ -3,3 +3,5 @@
|
||||
#include "../mmq.cuh"
|
||||
|
||||
DECL_MMQ_CASE(GGML_TYPE_MXFP4);
|
||||
|
||||
DECL_MMQ_CASE_W4A4(GGML_TYPE_MXFP4);
|
||||
|
||||
@@ -3,3 +3,5 @@
|
||||
#include "../mmq.cuh"
|
||||
|
||||
DECL_MMQ_CASE(GGML_TYPE_NVFP4);
|
||||
|
||||
DECL_MMQ_CASE_W4A4(GGML_TYPE_NVFP4);
|
||||
|
||||
@@ -98,7 +98,12 @@ __global__ void topk_moe_cuda(const float * logits,
|
||||
const float clamp_val,
|
||||
const float scale_val,
|
||||
const topk_moe_config config) {
|
||||
#if defined(GGML_USE_MUSA)
|
||||
// MUSA: every warp of a partially filled block must reach the barrier below.
|
||||
const int row = MIN(blockIdx.x * blockDim.y + threadIdx.y, n_rows - 1);
|
||||
#else
|
||||
const int row = blockIdx.x * blockDim.y + threadIdx.y;
|
||||
#endif // defined(GGML_USE_MUSA)
|
||||
if (row >= n_rows) {
|
||||
return;
|
||||
}
|
||||
|
||||
Vendored
+2
-2
@@ -248,11 +248,11 @@
|
||||
typedef __hip_bfloat16 nv_bfloat16;
|
||||
typedef __hip_bfloat162 nv_bfloat162;
|
||||
|
||||
#if HIP_VERSION >= 60200000
|
||||
#if HIP_VERSION >= 60300000
|
||||
#include <hip/hip_fp8.h>
|
||||
typedef __hip_fp8_e4m3 __nv_fp8_e4m3;
|
||||
#define FP8_AVAILABLE
|
||||
#endif // HIP_VERSION >= 60200000
|
||||
#endif // HIP_VERSION >= 60300000
|
||||
|
||||
typedef int8_t int8x4_t __attribute__((ext_vector_type(4)));
|
||||
typedef uint8_t uint8x4_t __attribute__((ext_vector_type(4)));
|
||||
|
||||
Vendored
+5
@@ -44,6 +44,7 @@
|
||||
#define cudaDeviceGetPCIBusId musaDeviceGetPCIBusId
|
||||
#define cudaDeviceProp musaDeviceProp
|
||||
#define cudaDeviceSynchronize musaDeviceSynchronize
|
||||
#define cudaDeviceGetAttribute musaDeviceGetAttribute
|
||||
#define cudaError_t musaError_t
|
||||
#define cudaErrorMemoryAllocation musaErrorMemoryAllocation
|
||||
#define cudaErrorPeerAccessAlreadyEnabled musaErrorPeerAccessAlreadyEnabled
|
||||
@@ -114,6 +115,7 @@
|
||||
#define cuMemRelease muMemRelease
|
||||
#define cuMemSetAccess muMemSetAccess
|
||||
#define cuMemUnmap muMemUnmap
|
||||
#define cudaDevAttrCooperativeLaunch musaDevAttrCooperativeLaunch
|
||||
#define cudaFuncAttributeMaxDynamicSharedMemorySize musaFuncAttributeMaxDynamicSharedMemorySize
|
||||
#define cudaFuncSetAttribute musaFuncSetAttribute
|
||||
#define cudaMemcpy3DPeerParms musaMemcpy3DPeerParms
|
||||
@@ -145,6 +147,9 @@
|
||||
#define cudaStreamCaptureModeRelaxed musaStreamCaptureModeRelaxed
|
||||
#define cudaStreamBeginCapture musaStreamBeginCapture
|
||||
#define cudaStreamEndCapture musaStreamEndCapture
|
||||
#define cudaStreamCaptureStatus musaStreamCaptureStatus
|
||||
#define cudaStreamCaptureStatusNone musaStreamCaptureStatusNone
|
||||
#define cudaStreamIsCapturing musaStreamIsCapturing
|
||||
#define cudaOccupancyMaxActiveBlocksPerMultiprocessor musaOccupancyMaxActiveBlocksPerMultiprocessor
|
||||
|
||||
typedef __mt_bfloat16 nv_bfloat16;
|
||||
|
||||
@@ -79,9 +79,7 @@ static __global__ void rwkv_wkv7_f32(const int B, const int T, const int C, cons
|
||||
float state[head_size];
|
||||
__shared__ float _r[head_size], _w[head_size], _k[head_size], _a[head_size], _b[head_size];
|
||||
|
||||
#ifndef GGML_USE_MUSA
|
||||
#pragma unroll
|
||||
#endif
|
||||
for (int i = 0; i < head_size; i++) {
|
||||
state[i] = s[batch_i * state_size + head_i * head_size * head_size + tid * head_size + i];
|
||||
}
|
||||
|
||||
@@ -23,6 +23,7 @@ include(${HEXAGON_SDK_ROOT}/build/cmake/hexagon_fun.cmake)
|
||||
include(ExternalProject)
|
||||
|
||||
option(GGML_HEXAGON_HTP_DEBUG "ggml-hexagon: enable HTP debug output" OFF)
|
||||
set(GGML_HEXAGON_HTP_BUILD_TYPE "Release" CACHE STRING "ggml-hexagon: HTP skel build type (Release, RelWithDebInfo, Debug)")
|
||||
set(GGML_HEXAGON_HTP_CERT "$ENV{HEXAGON_HTP_CERT}" CACHE PATH "ggml-hexagon: enable HTP library signing using certificate")
|
||||
|
||||
add_library(htp_iface OBJECT
|
||||
@@ -64,7 +65,7 @@ function(build_htp_skel V)
|
||||
SOURCE_DIR ${CMAKE_CURRENT_SOURCE_DIR}/htp BUILD_ALWAYS ON
|
||||
BUILD_BYPRODUCTS ${CMAKE_CURRENT_BINARY_DIR}/libggml-htp-${V}.so
|
||||
CMAKE_ARGS
|
||||
-DCMAKE_BUILD_TYPE=Release
|
||||
-DCMAKE_BUILD_TYPE=${GGML_HEXAGON_HTP_BUILD_TYPE}
|
||||
-DCMAKE_TOOLCHAIN_FILE=${CMAKE_CURRENT_SOURCE_DIR}/htp/cmake-toolchain.cmake
|
||||
-DCMAKE_INSTALL_LIBDIR=${CMAKE_CURRENT_BINARY_DIR}
|
||||
-DHEXAGON_SDK_ROOT=${HEXAGON_SDK_ROOT}
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -267,10 +268,10 @@ static inline bool ggml_hexagon_is_repack_type(enum ggml_type type) {
|
||||
return type == GGML_TYPE_Q4_0 || type == GGML_TYPE_Q4_1 ||
|
||||
type == GGML_TYPE_Q8_0 || type == GGML_TYPE_IQ4_NL ||
|
||||
type == GGML_TYPE_MXFP4 || type == GGML_TYPE_Q6_K ||
|
||||
type == GGML_TYPE_Q4_K;
|
||||
type == GGML_TYPE_Q4_K || type == GGML_TYPE_Q5_K;
|
||||
}
|
||||
|
||||
// Size of one repacked row in the DSP tiled layout. The Q6_K and Q4_K tiles store uncompressed scales/mins,
|
||||
// Size of one repacked row in the DSP tiled layout. The Q6_K, Q5_K and Q4_K tiles store uncompressed scales/mins,
|
||||
// so they are larger than the ggml blocks. For the other repack types the tile has the same size as the ggml blocks.
|
||||
static inline size_t ggml_hexagon_tiled_row_size(enum ggml_type type, int64_t ne0) {
|
||||
if (type == GGML_TYPE_Q6_K) {
|
||||
@@ -279,6 +280,9 @@ static inline size_t ggml_hexagon_tiled_row_size(enum ggml_type type, int64_t ne
|
||||
if (type == GGML_TYPE_Q4_K) {
|
||||
return (size_t) (ne0 / 32) * (HTP_MM_WEIGHT_TILE_SIZE_Q4_1 / 32);
|
||||
}
|
||||
if (type == GGML_TYPE_Q5_K) {
|
||||
return (size_t) (ne0 / 32) * (HTP_MM_WEIGHT_TILE_SIZE_Q5_K / 32);
|
||||
}
|
||||
return ggml_row_size(type, ne0);
|
||||
}
|
||||
|
||||
@@ -365,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,
|
||||
@@ -1742,6 +1753,202 @@ static void repack_tiled_q4_K(void * data, const ggml_tensor * t, size_t offset,
|
||||
GGML_UNUSED(size);
|
||||
}
|
||||
|
||||
// tile layout: see HTP_MM_WEIGHT_TILE_SIZE_Q5_K in htp/matmul-ops.h
|
||||
static void repack_q5_K_tiled(ggml_tensor * t, const void * data, size_t offset, size_t size) {
|
||||
GGML_ASSERT(offset == 0);
|
||||
|
||||
const block_q5_K * src_matrix = (const block_q5_K *) data;
|
||||
int64_t ne0 = t->ne[0];
|
||||
int64_t ne1 = t->ne[1];
|
||||
int64_t ne2 = t->ne[2];
|
||||
int64_t ne3 = t->ne[3];
|
||||
int64_t ne0_padded = hex_round_up(ne0, 32);
|
||||
int64_t ne1_padded = hex_round_up(ne1, 32);
|
||||
|
||||
GGML_ASSERT(ne0 % QK_K == 0);
|
||||
|
||||
const int n_col_tiles = ne1_padded / 32;
|
||||
const int n_k_tiles = ne0_padded / 32;
|
||||
const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_Q5_K;
|
||||
const size_t matrix_size = (size_t) n_col_tiles * n_k_tiles * tile_size;
|
||||
|
||||
const int64_t sb_per_row = ne0 / QK_K;
|
||||
|
||||
for (int i3 = 0; i3 < ne3; i3++) {
|
||||
for (int i2 = 0; i2 < ne2; i2++) {
|
||||
const block_q5_K * src_slice = src_matrix + (i3 * ne2 + i2) * (ne1 * sb_per_row);
|
||||
uint8_t * matrix_dst = (uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size;
|
||||
|
||||
memset(matrix_dst, 0, matrix_size);
|
||||
|
||||
for (int64_t r = 0; r < ne1; r++) {
|
||||
const int ct = (int) (r / 32);
|
||||
const int row = (int) (r % 32);
|
||||
const block_q5_K * src_row = src_slice + r * sb_per_row;
|
||||
|
||||
for (int kt = 0; kt < n_k_tiles; kt++) {
|
||||
const int kt_local = kt % 8;
|
||||
const block_q5_K * b = &src_row[kt / 8];
|
||||
const float d = GGML_FP16_TO_FP32(b->d);
|
||||
const float dmin = GGML_FP16_TO_FP32(b->dmin);
|
||||
|
||||
uint8_t * tile_dst = matrix_dst + ((size_t) ct * n_k_tiles + kt) * tile_size;
|
||||
uint8_t * plane = tile_dst + 640;
|
||||
|
||||
uint8_t sc, m;
|
||||
get_scale_min_k4(kt_local, b->scales, &sc, &m);
|
||||
|
||||
const float D = d * (float) sc;
|
||||
const float M = -dmin * (float) m;
|
||||
|
||||
const uint8_t * qs_sub = b->qs + (kt_local / 2) * 32;
|
||||
const int shift = (kt_local & 1) ? 4 : 0;
|
||||
const uint8_t hbit = (uint8_t) (1 << kt_local);
|
||||
|
||||
for (int cp = 0; cp < 16; cp++) {
|
||||
const uint8_t q0 = (qs_sub[2 * cp + 0] >> shift) & 0x0F;
|
||||
const uint8_t q1 = (qs_sub[2 * cp + 1] >> shift) & 0x0F;
|
||||
tile_dst[cp * 32 + row] = (uint8_t) ((q1 << 4) | q0);
|
||||
|
||||
const int i = cp / 4;
|
||||
const int lane = (cp % 4) * 32 + row;
|
||||
if (b->qh[2 * cp + 0] & hbit) {
|
||||
plane[lane] |= (uint8_t) (1 << (2 * i));
|
||||
}
|
||||
if (b->qh[2 * cp + 1] & hbit) {
|
||||
plane[lane] |= (uint8_t) (1 << (2 * i + 1));
|
||||
}
|
||||
}
|
||||
|
||||
ggml_half * scale_dst = (ggml_half *) (tile_dst + 512);
|
||||
scale_dst[2 * row + 0] = GGML_FP32_TO_FP16(D);
|
||||
scale_dst[2 * row + 1] = GGML_FP32_TO_FP16(M);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
GGML_UNUSED(size);
|
||||
}
|
||||
|
||||
// Reverse of repack_q5_K_tiled. Unpacks quants losslessly and normalizes scales/mins. Read-back only.
|
||||
static void repack_tiled_q5_K(void * data, const ggml_tensor * t, size_t offset, size_t size) {
|
||||
GGML_ASSERT(offset == 0);
|
||||
|
||||
block_q5_K * dst_matrix = (block_q5_K *) data;
|
||||
int64_t ne0 = t->ne[0];
|
||||
int64_t ne1 = t->ne[1];
|
||||
int64_t ne2 = t->ne[2];
|
||||
int64_t ne3 = t->ne[3];
|
||||
int64_t ne0_padded = hex_round_up(ne0, 32);
|
||||
int64_t ne1_padded = hex_round_up(ne1, 32);
|
||||
|
||||
GGML_ASSERT(ne0 % QK_K == 0);
|
||||
|
||||
const int n_col_tiles = ne1_padded / 32;
|
||||
const int n_k_tiles = ne0_padded / 32;
|
||||
const size_t tile_size = HTP_MM_WEIGHT_TILE_SIZE_Q5_K;
|
||||
const size_t matrix_size = (size_t) n_col_tiles * n_k_tiles * tile_size;
|
||||
|
||||
const int64_t sb_per_row = ne0 / QK_K;
|
||||
|
||||
for (int i3 = 0; i3 < ne3; i3++) {
|
||||
for (int i2 = 0; i2 < ne2; i2++) {
|
||||
block_q5_K * dst_slice = dst_matrix + (i3 * ne2 + i2) * (ne1 * sb_per_row);
|
||||
const uint8_t * matrix_src = (const uint8_t *) t->data + (i3 * ne2 + i2) * matrix_size;
|
||||
|
||||
for (int64_t r = 0; r < ne1; r++) {
|
||||
const int ct = (int) (r / 32);
|
||||
const int row = (int) (r % 32);
|
||||
block_q5_K * dst_row = dst_slice + r * sb_per_row;
|
||||
|
||||
for (int64_t sb = 0; sb < sb_per_row; sb++) {
|
||||
block_q5_K * b = &dst_row[sb];
|
||||
memset(b, 0, sizeof(block_q5_K));
|
||||
|
||||
float sub_scales[8];
|
||||
float sub_mins[8];
|
||||
|
||||
for (int kt_local = 0; kt_local < 8; kt_local++) {
|
||||
const int kt = sb * 8 + kt_local;
|
||||
const uint8_t * tile_src = matrix_src + ((size_t) ct * n_k_tiles + kt) * tile_size;
|
||||
const uint8_t * plane = tile_src + 640;
|
||||
const ggml_half * scale_src = (const ggml_half *) (tile_src + 512);
|
||||
|
||||
uint8_t * qs_sub = b->qs + (kt_local / 2) * 32;
|
||||
const int shift = (kt_local & 1) ? 4 : 0;
|
||||
const uint8_t hbit = (uint8_t) (1 << kt_local);
|
||||
|
||||
for (int cp = 0; cp < 16; cp++) {
|
||||
const uint8_t val = tile_src[cp * 32 + row];
|
||||
const uint8_t q0 = val & 0x0F;
|
||||
const uint8_t q1 = val >> 4;
|
||||
qs_sub[2 * cp + 0] |= (uint8_t) (q0 << shift);
|
||||
qs_sub[2 * cp + 1] |= (uint8_t) (q1 << shift);
|
||||
|
||||
const int i = cp / 4;
|
||||
const int lane = (cp % 4) * 32 + row;
|
||||
if (plane[lane] & (1 << (2 * i))) {
|
||||
b->qh[2 * cp + 0] |= hbit;
|
||||
}
|
||||
if (plane[lane] & (1 << (2 * i + 1))) {
|
||||
b->qh[2 * cp + 1] |= hbit;
|
||||
}
|
||||
}
|
||||
|
||||
const float D = GGML_FP16_TO_FP32(scale_src[2 * row + 0]);
|
||||
const float M = GGML_FP16_TO_FP32(scale_src[2 * row + 1]);
|
||||
sub_scales[kt_local] = (D > 0.0f) ? D : 0.0f;
|
||||
sub_mins[kt_local] = (-M > 0.0f) ? -M : 0.0f;
|
||||
}
|
||||
|
||||
float max_scale = 0.0f;
|
||||
float max_min = 0.0f;
|
||||
for (int j = 0; j < 8; j++) {
|
||||
if (sub_scales[j] > max_scale) max_scale = sub_scales[j];
|
||||
if (sub_mins[j] > max_min) max_min = sub_mins[j];
|
||||
}
|
||||
|
||||
float inv_scale = 0.0f;
|
||||
if (max_scale > 0.0f) {
|
||||
b->d = GGML_FP32_TO_FP16(max_scale / 63.0f);
|
||||
const float d_actual = GGML_FP16_TO_FP32(b->d);
|
||||
inv_scale = (d_actual > 0.0f) ? (1.0f / d_actual) : 0.0f;
|
||||
} else {
|
||||
b->d = GGML_FP32_TO_FP16(0.0f);
|
||||
}
|
||||
|
||||
float inv_min = 0.0f;
|
||||
if (max_min > 0.0f) {
|
||||
b->dmin = GGML_FP32_TO_FP16(max_min / 63.0f);
|
||||
const float dmin_actual = GGML_FP16_TO_FP32(b->dmin);
|
||||
inv_min = (dmin_actual > 0.0f) ? (1.0f / dmin_actual) : 0.0f;
|
||||
} else {
|
||||
b->dmin = GGML_FP32_TO_FP16(0.0f);
|
||||
}
|
||||
|
||||
for (int j = 0; j < 8; j++) {
|
||||
uint8_t ls = (uint8_t) roundf(inv_scale * sub_scales[j]);
|
||||
uint8_t lm = (uint8_t) roundf(inv_min * sub_mins[j]);
|
||||
ls = (std::min)((uint8_t) 63, ls);
|
||||
lm = (std::min)((uint8_t) 63, lm);
|
||||
if (j < 4) {
|
||||
b->scales[j] = ls;
|
||||
b->scales[j + 4] = lm;
|
||||
} else {
|
||||
b->scales[j + 4] = (ls & 0xF) | ((lm & 0xF) << 4);
|
||||
b->scales[j - 4] |= ((ls >> 4) << 6);
|
||||
b->scales[j - 0] |= ((lm >> 4) << 6);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
GGML_UNUSED(size);
|
||||
}
|
||||
|
||||
static void repack_tensor_tiled(ggml_tensor * tensor, const void * data, size_t size) {
|
||||
switch (tensor->type) {
|
||||
case GGML_TYPE_Q4_0:
|
||||
@@ -1768,6 +1975,10 @@ static void repack_tensor_tiled(ggml_tensor * tensor, const void * data, size_t
|
||||
repack_mxfp4_tiled(tensor, data, 0, size);
|
||||
break;
|
||||
|
||||
case GGML_TYPE_Q5_K:
|
||||
repack_q5_K_tiled(tensor, data, 0, size);
|
||||
break;
|
||||
|
||||
case GGML_TYPE_Q6_K:
|
||||
repack_q6_K_tiled(tensor, data, 0, size);
|
||||
break;
|
||||
@@ -1856,6 +2067,12 @@ static void ggml_backend_hexagon_buffer_get_tensor(ggml_backend_buffer_t buffer,
|
||||
repack_tiled_q4_K(data, tensor, offset, size);
|
||||
break;
|
||||
|
||||
case GGML_TYPE_Q5_K:
|
||||
GGML_ASSERT(offset == 0);
|
||||
GGML_ASSERT(offset + size <= ggml_nbytes(tensor));
|
||||
repack_tiled_q5_K(data, tensor, offset, size);
|
||||
break;
|
||||
|
||||
case GGML_TYPE_Q8_0:
|
||||
GGML_ASSERT(offset == 0);
|
||||
GGML_ASSERT(offset + size <= ggml_nbytes(tensor));
|
||||
@@ -1989,6 +2206,10 @@ static void ggml_backend_hexagon_buffer_get_tensor_2d(ggml_backend_buffer_t buff
|
||||
repack_tiled_q4_K(temp_buf.data(), tensor, offset, temp_size);
|
||||
break;
|
||||
|
||||
case GGML_TYPE_Q5_K:
|
||||
repack_tiled_q5_K(temp_buf.data(), tensor, offset, temp_size);
|
||||
break;
|
||||
|
||||
case GGML_TYPE_Q8_0:
|
||||
repack_tiled_q8_0(temp_buf.data(), tensor, offset, temp_size);
|
||||
break;
|
||||
@@ -4494,8 +4715,7 @@ static bool ggml_hexagon_precompute_hmx_mm_params(
|
||||
|
||||
if (is_batched_val && wtype == GGML_TYPE_F16 && group_size > 1) {
|
||||
// Try grouped path first
|
||||
const bool use_dma_activation = (src1->nb[1]/sizeof(float) > (size_t)ne00_padded);
|
||||
if (htp_mm_hmx_solve_batched_params(wtype, ne00_padded, ne01_padded, ne11, group_size, use_dma_activation, n_threads, pipeline, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
|
||||
if (htp_mm_hmx_solve_batched_params(wtype, ne00_padded, ne01_padded, ne11, group_size, n_threads, pipeline, src2_size, vtcm_budget, &m_chunk, &n_chunk, &act_threads_selected, &vtcm_size)) {
|
||||
use_grouped = true;
|
||||
}
|
||||
}
|
||||
@@ -4786,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;
|
||||
@@ -4830,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;
|
||||
}
|
||||
|
||||
@@ -4927,49 +5208,51 @@ static void ggml_hexagon_precompute_get_rows_params(
|
||||
const uint32_t ne12 = src1->ne[2];
|
||||
const uint32_t nr = ne10 * ne11 * ne12;
|
||||
|
||||
const size_t nb01 = src0->nb[1];
|
||||
const size_t nb1 = dst->nb[1];
|
||||
const ggml_tensor * src0_base = src0->view_src ? src0->view_src : src0;
|
||||
const auto * extra = src0_base->buffer && ggml_backend_buffer_is_hexagon(src0_base->buffer) ?
|
||||
(const ggml_hexagon_tensor_extra *) src0_base->extra : nullptr;
|
||||
const bool tiled = src0->type == GGML_TYPE_Q4_0 || (extra && (extra->flags & GGML_HEXAGON_TENSOR_REPACK) != 0) ||
|
||||
sess->needs_repack.count(src0_base) || sess->needs_repack.count(src0);
|
||||
|
||||
const bool can_use_dma = (src0->type == dst->type) && (nb01 == nb1);
|
||||
const bool use_dma = can_use_dma && (ne00 >= 2048);
|
||||
|
||||
kparams->use_dma = use_dma ? 1 : 0;
|
||||
|
||||
uint32_t chunks_per_row = 1;
|
||||
uint32_t chunk_size = ne00;
|
||||
uint32_t total_tasks = nr;
|
||||
|
||||
if (use_dma) {
|
||||
kparams->n_threads = (std::min)((uint32_t)sess->n_threads, nr);
|
||||
kparams->tasks_per_thread = (nr + kparams->n_threads - 1) / kparams->n_threads;
|
||||
if (src0->type == dst->type) {
|
||||
kparams->kernel_type = HTP_GET_ROWS_KERNEL_SAMETYPE;
|
||||
} else if (tiled) {
|
||||
kparams->kernel_type = HTP_GET_ROWS_KERNEL_TILED;
|
||||
} else {
|
||||
if (src0->type == GGML_TYPE_F32 && nr < sess->n_threads) {
|
||||
const uint32_t min_chunk_size = 1024;
|
||||
uint32_t max_chunks = ne00 / min_chunk_size;
|
||||
if (max_chunks == 0) {
|
||||
max_chunks = 1;
|
||||
}
|
||||
chunks_per_row = (std::min)((sess->n_threads + nr - 1) / nr, max_chunks);
|
||||
chunk_size = (ne00 + chunks_per_row - 1) / chunks_per_row;
|
||||
total_tasks = nr * chunks_per_row;
|
||||
}
|
||||
kparams->n_threads = (std::min)(total_tasks, (uint32_t)sess->n_threads);
|
||||
kparams->tasks_per_thread = (total_tasks + kparams->n_threads - 1) / kparams->n_threads;
|
||||
kparams->kernel_type = HTP_GET_ROWS_KERNEL_FLAT;
|
||||
}
|
||||
|
||||
const uint32_t chunks_per_row = 1;
|
||||
const uint32_t chunk_size = ne00;
|
||||
const uint32_t total_tasks = nr;
|
||||
|
||||
kparams->n_threads = (std::min)((uint32_t)sess->n_threads, total_tasks);
|
||||
|
||||
struct htp_get_rows_vtcm_layout vtcm_layout = {};
|
||||
while (kparams->n_threads > 0) {
|
||||
htp_get_rows_vtcm_layout_build(&vtcm_layout, kparams->kernel_type, src0->type, ne00, kparams->n_threads);
|
||||
if (vtcm_layout.total_bytes <= sess->vtcm_size) {
|
||||
break;
|
||||
}
|
||||
--kparams->n_threads;
|
||||
}
|
||||
|
||||
if (kparams->n_threads == 0 && total_tasks > 0) {
|
||||
htp_get_rows_vtcm_layout_build(&vtcm_layout, kparams->kernel_type, src0->type, ne00, 1);
|
||||
}
|
||||
|
||||
kparams->vtcm_size = (total_tasks == 0) ? 0 : vtcm_layout.total_bytes;
|
||||
kparams->tasks_per_thread = kparams->n_threads > 0 ? (total_tasks + kparams->n_threads - 1) / kparams->n_threads : 0;
|
||||
|
||||
kparams->chunks_per_row = chunks_per_row;
|
||||
kparams->chunk_size = chunk_size;
|
||||
kparams->total_tasks = total_tasks;
|
||||
|
||||
kparams->div_ne10 = init_fastdiv_values(ne10);
|
||||
kparams->div_ne10_ne11 = init_fastdiv_values(ne10 * ne11);
|
||||
kparams->div_chunks_per_row = init_fastdiv_values(chunks_per_row);
|
||||
kparams->div_ne02 = init_fastdiv_values(ne02);
|
||||
kparams->div_ne03 = init_fastdiv_values(ne03);
|
||||
|
||||
struct htp_get_rows_vtcm_layout vtcm_layout;
|
||||
htp_get_rows_vtcm_layout_build(&vtcm_layout, src0->type, ne00, kparams->n_threads);
|
||||
kparams->vtcm_size = vtcm_layout.total_bytes;
|
||||
kparams->div_ne10 = ne10 > 0 ? init_fastdiv_values(ne10) : fastdiv_values{0, 0};
|
||||
kparams->div_ne10_ne11 = (ne10 * ne11) > 0 ? init_fastdiv_values(ne10 * ne11) : fastdiv_values{0, 0};
|
||||
kparams->div_chunks_per_row = chunks_per_row > 0 ? init_fastdiv_values(chunks_per_row) : fastdiv_values{0, 0};
|
||||
kparams->div_ne02 = ne02 > 0 ? init_fastdiv_values(ne02) : fastdiv_values{0, 0};
|
||||
kparams->div_ne03 = ne03 > 0 ? init_fastdiv_values(ne03) : fastdiv_values{0, 0};
|
||||
}
|
||||
|
||||
static void ggml_hexagon_precompute_set_rows_params(
|
||||
@@ -5267,6 +5550,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
|
||||
@@ -5404,12 +5733,13 @@ static bool ggml_hexagon_supported_mul_mat(const struct ggml_hexagon_session * s
|
||||
case GGML_TYPE_IQ4_NL:
|
||||
case GGML_TYPE_MXFP4:
|
||||
case GGML_TYPE_Q4_K:
|
||||
case GGML_TYPE_Q5_K:
|
||||
case GGML_TYPE_Q6_K:
|
||||
if (!ggml_is_contiguous(src0) || ggml_is_permuted(src0)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (src0->ne[0] % ((src0->type == GGML_TYPE_Q6_K || src0->type == GGML_TYPE_Q4_K) ? QK_K : 32)) {
|
||||
if (src0->ne[0] % ((src0->type == GGML_TYPE_Q6_K || src0->type == GGML_TYPE_Q5_K || src0->type == GGML_TYPE_Q4_K) ? QK_K : 32)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -5487,12 +5817,13 @@ static bool ggml_hexagon_supported_mul_mat_id(const struct ggml_hexagon_session
|
||||
case GGML_TYPE_IQ4_NL:
|
||||
case GGML_TYPE_MXFP4:
|
||||
case GGML_TYPE_Q4_K:
|
||||
case GGML_TYPE_Q5_K:
|
||||
case GGML_TYPE_Q6_K:
|
||||
if (!ggml_is_contiguous(src0) || ggml_is_permuted(src0)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (src0->ne[0] % ((src0->type == GGML_TYPE_Q6_K || src0->type == GGML_TYPE_Q4_K) ? QK_K : 32)) {
|
||||
if (src0->ne[0] % ((src0->type == GGML_TYPE_Q6_K || src0->type == GGML_TYPE_Q5_K || src0->type == GGML_TYPE_Q4_K) ? QK_K : 32)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -5616,7 +5947,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;
|
||||
@@ -5639,6 +5971,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;
|
||||
@@ -5660,6 +6012,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];
|
||||
@@ -5806,19 +6178,36 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session *
|
||||
const struct ggml_tensor * src1 = op->src[1]; // indices
|
||||
const struct ggml_tensor * dst = op;
|
||||
|
||||
if (src0->extra) {
|
||||
const auto * extra = (const ggml_hexagon_tensor_extra *) src0->extra;
|
||||
if (extra->flags & GGML_HEXAGON_TENSOR_REPACK) {
|
||||
if (src0->type == GGML_TYPE_Q4_0 && src0->view_src) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const ggml_tensor * src0_base = src0->view_src ? src0->view_src : src0;
|
||||
bool is_repacked = false;
|
||||
if (src0_base->buffer && ggml_backend_buffer_is_hexagon(src0_base->buffer) && src0_base->extra) {
|
||||
const auto * extra = (const ggml_hexagon_tensor_extra *) src0_base->extra;
|
||||
is_repacked = (extra->flags & GGML_HEXAGON_TENSOR_REPACK) != 0;
|
||||
if (is_repacked && src0->type != GGML_TYPE_Q4_0 && src0->type != GGML_TYPE_Q8_0) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
is_repacked = is_repacked || sess->needs_repack.count(src0_base) || sess->needs_repack.count(src0);
|
||||
|
||||
if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_I32 && src0->ne[0] < 32) {
|
||||
// View offsets use the raw quantized layout and cannot address a tiled allocation.
|
||||
if (src0->view_src && is_repacked) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (src0->type == GGML_TYPE_Q4_0 && src0->buffer && !is_repacked) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (src0->type != dst->type && src0->ne[0] < 32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_F16 &&
|
||||
src0->type != GGML_TYPE_Q8_0 && src0->type != GGML_TYPE_I32) {
|
||||
src0->type != GGML_TYPE_Q4_0 && src0->type != GGML_TYPE_Q8_0 && src0->type != GGML_TYPE_I32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -5826,13 +6215,30 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session *
|
||||
return false;
|
||||
}
|
||||
|
||||
if (src0->type == GGML_TYPE_I32) {
|
||||
if (dst->type != GGML_TYPE_I32) {
|
||||
if (src0->type == dst->type) {
|
||||
if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_I32 && src0->type != GGML_TYPE_F16) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
else if (dst->type != GGML_TYPE_F32) {
|
||||
} else if (src0->type == GGML_TYPE_I32) {
|
||||
return false;
|
||||
} else if (dst->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Empty recurrent-state gathers are skipped at execution; do not split the graph for them.
|
||||
if (ggml_is_empty(op)) {
|
||||
return true;
|
||||
}
|
||||
|
||||
struct htp_get_rows_kernel_params kparams;
|
||||
ggml_hexagon_precompute_get_rows_params(sess, src0, src1, dst, &kparams);
|
||||
if (kparams.n_threads == 0 || (size_t) kparams.vtcm_size > sess->vtcm_size) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Q4_0 has no raw fallback. Mark only accepted tensors for repacking.
|
||||
if (src0->type == GGML_TYPE_Q4_0 && !src0->buffer) {
|
||||
sess->needs_repack.insert(src0);
|
||||
}
|
||||
|
||||
return true;
|
||||
@@ -5844,42 +6250,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;
|
||||
}
|
||||
|
||||
@@ -6187,9 +6587,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;
|
||||
@@ -6226,6 +6628,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;
|
||||
}
|
||||
@@ -6456,6 +6859,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));
|
||||
}
|
||||
@@ -7094,16 +7503,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;
|
||||
@@ -7252,10 +7670,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;
|
||||
@@ -7272,6 +7698,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:
|
||||
@@ -7758,6 +8185,8 @@ static void ggml_hexagon_init(ggml_backend_reg * reg) {
|
||||
"please update hexagon_type to match ggml_type");
|
||||
static_assert((unsigned int) HTP_TYPE_Q4_K == (unsigned int) GGML_TYPE_Q4_K,
|
||||
"please update hexagon_type to match ggml_type");
|
||||
static_assert((unsigned int) HTP_TYPE_Q5_K == (unsigned int) GGML_TYPE_Q5_K,
|
||||
"please update hexagon_type to match ggml_type");
|
||||
static_assert((unsigned int) HTP_TYPE_Q6_K == (unsigned int) GGML_TYPE_Q6_K,
|
||||
"please update hexagon_type to match ggml_type");
|
||||
|
||||
|
||||
@@ -19,6 +19,8 @@
|
||||
#include "htp/ssm-conv.h"
|
||||
#include "htp/gated-delta-net-ops.h"
|
||||
#include "htp/softmax-ops.h"
|
||||
#include "htp/argsort-ops.h"
|
||||
#include "htp/get-rows-ops.h"
|
||||
|
||||
struct htp_opnode {
|
||||
ggml_tensor * node { nullptr };
|
||||
@@ -359,17 +361,31 @@ struct htp_opformat {
|
||||
} else if (node.opcode == HTP_OP_GATED_DELTA_NET) {
|
||||
const auto * kparams = (const struct htp_gdn_kernel_params *) node.kernel_params;
|
||||
const char * path = (kparams->kernel_type == HTP_GDN_KERNEL_HMX_CHUNKED) ? "hmx-chunked" : "hvx-recurrent";
|
||||
snprintf(str, max_size, "%s-%s vtcm %u",
|
||||
path,
|
||||
kparams->kda ? "kda" : "scalar",
|
||||
snprintf(str, max_size, "%s-%s vtcm %u", path, kparams->kda ? "kda" : "scalar",
|
||||
(unsigned int) (kparams->vtcm_size ? kparams->vtcm_size : kparams->vtcm_per_thread * kparams->n_threads));
|
||||
} else if (node.opcode == HTP_OP_MUL || node.opcode == HTP_OP_ADD || node.opcode == HTP_OP_ADD_ID ||
|
||||
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 if (node.opcode == HTP_OP_GET_ROWS) {
|
||||
const auto * kparams = (const struct htp_get_rows_kernel_params *) node.kernel_params;
|
||||
const char * ktype_str = "unknown";
|
||||
switch (kparams->kernel_type) {
|
||||
case HTP_GET_ROWS_KERNEL_SAMETYPE: ktype_str = "sametype"; break;
|
||||
case HTP_GET_ROWS_KERNEL_TILED: ktype_str = "tiled"; break;
|
||||
case HTP_GET_ROWS_KERNEL_FLAT: ktype_str = "flat"; break;
|
||||
}
|
||||
snprintf(str, max_size, "%s%s vtcm %u", ktype_str, kparams->n_threads > 1 ? "-multi" : "", (unsigned int) kparams->vtcm_size);
|
||||
} else {
|
||||
snprintf(str, max_size, "----");
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
void format(const htp_opnode & node) {
|
||||
|
||||
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);
|
||||
|
||||
@@ -20,6 +20,7 @@ struct htp_concat_context {
|
||||
uint32_t nrows;
|
||||
uint32_t elem_start;
|
||||
uint32_t nelems;
|
||||
uint32_t nplanes;
|
||||
struct fastdiv_values div_ne0;
|
||||
struct fastdiv_values div_ne1;
|
||||
struct fastdiv_values div_ne2;
|
||||
@@ -60,39 +61,47 @@ static void concat_2d_f32_transposed(unsigned int nth, unsigned int ith, void *
|
||||
|
||||
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
|
||||
|
||||
for (uint32_t i = start_i; i < end_i; i += block_i) {
|
||||
uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i;
|
||||
for (uint32_t p = 0; p < cctx->nplanes; p++) {
|
||||
const uint32_t i3 = p / dst->ne[2];
|
||||
const uint32_t i2 = p - i3 * dst->ne[2];
|
||||
const dma_addr_t src0_plane = src0->data + i2 * src0->nb[2] + i3 * src0->nb[3];
|
||||
const dma_addr_t src1_plane = src1->data + i2 * src1->nb[2] + i3 * src1->nb[3];
|
||||
const dma_addr_t dst_plane = dst->data + i2 * dst->nb[2] + i3 * dst->nb[3];
|
||||
|
||||
uint32_t src1_width_bytes = current_block_i * sizeof(float);
|
||||
const dma_addr_t src1_addr = src1->data + i * src1->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);
|
||||
for (uint32_t i = start_i; i < end_i; i += block_i) {
|
||||
uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i;
|
||||
|
||||
uint32_t src0_row_bytes = src0_ne0 * sizeof(float);
|
||||
const dma_addr_t src0_addr = src0->data + i * src0->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);
|
||||
uint32_t src1_width_bytes = current_block_i * sizeof(float);
|
||||
const dma_addr_t src1_addr = src1_plane + i * src1->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);
|
||||
|
||||
dma_queue_pop(dma_q); // src1
|
||||
uint32_t src0_row_bytes = src0_ne0 * sizeof(float);
|
||||
const dma_addr_t src0_addr = src0_plane + i * src0->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);
|
||||
|
||||
HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride);
|
||||
dma_queue_pop(dma_q); // src1
|
||||
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
|
||||
for (uint32_t j = 0; j < src1_ne0_padded; j += 32) {
|
||||
#pragma unroll(4)
|
||||
for (uint32_t ii = 0; ii < current_block_i; ii++) {
|
||||
size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(float));
|
||||
Q6_vgather_ARMVw(&vtcm_tmp[ii], rt, mu, vv);
|
||||
uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(float);
|
||||
hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
|
||||
HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride);
|
||||
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
|
||||
for (uint32_t j = 0; j < src1_ne0_padded; j += 32) {
|
||||
#pragma unroll(4)
|
||||
for (uint32_t ii = 0; ii < current_block_i; ii++) {
|
||||
size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(float));
|
||||
Q6_vgather_ARMVw(&vtcm_tmp[ii], rt, mu, vv);
|
||||
uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(float);
|
||||
hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
|
||||
}
|
||||
}
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
|
||||
|
||||
dma_queue_pop(dma_q); // src0
|
||||
|
||||
const dma_addr_t dst_addr = dst_plane + i * dst->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(float), current_block_i);
|
||||
|
||||
dma_queue_pop(dma_q);
|
||||
}
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
|
||||
|
||||
dma_queue_pop(dma_q); // src0
|
||||
|
||||
const dma_addr_t dst_addr = dst->data + i * dst->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(float), current_block_i);
|
||||
|
||||
dma_queue_pop(dma_q);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -131,39 +140,47 @@ static void concat_2d_f16_transposed(unsigned int nth, unsigned int ith, void *
|
||||
|
||||
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
|
||||
|
||||
for (uint32_t i = start_i; i < end_i; i += block_i) {
|
||||
uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i;
|
||||
for (uint32_t p = 0; p < cctx->nplanes; p++) {
|
||||
const uint32_t i3 = p / dst->ne[2];
|
||||
const uint32_t i2 = p - i3 * dst->ne[2];
|
||||
const dma_addr_t src0_plane = src0->data + i2 * src0->nb[2] + i3 * src0->nb[3];
|
||||
const dma_addr_t src1_plane = src1->data + i2 * src1->nb[2] + i3 * src1->nb[3];
|
||||
const dma_addr_t dst_plane = dst->data + i2 * dst->nb[2] + i3 * dst->nb[3];
|
||||
|
||||
uint32_t src1_width_bytes = current_block_i * sizeof(__fp16);
|
||||
const dma_addr_t src1_addr = src1->data + i * src1->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);
|
||||
for (uint32_t i = start_i; i < end_i; i += block_i) {
|
||||
uint32_t current_block_i = (end_i - i < block_i) ? (end_i - i) : block_i;
|
||||
|
||||
uint32_t src0_row_bytes = src0_ne0 * sizeof(__fp16);
|
||||
const dma_addr_t src0_addr = src0->data + i * src0->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);
|
||||
uint32_t src1_width_bytes = current_block_i * sizeof(__fp16);
|
||||
const dma_addr_t src1_addr = src1_plane + i * src1->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(spad1_base, src1_addr), spad1_stride, src1->nb[0], src1_width_bytes, src1_ne0);
|
||||
|
||||
dma_queue_pop(dma_q); // src1
|
||||
uint32_t src0_row_bytes = src0_ne0 * sizeof(__fp16);
|
||||
const dma_addr_t src0_addr = src0_plane + i * src0->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(spad0_base, src0_addr), spad0_row_bytes, src0->nb[1], src0_row_bytes, current_block_i);
|
||||
|
||||
HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride);
|
||||
dma_queue_pop(dma_q); // src1
|
||||
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
|
||||
for (uint32_t j = 0; j < src1_ne0_padded; j += 64) {
|
||||
#pragma unroll(4)
|
||||
for (uint32_t ii = 0; ii < current_block_i; ii++) {
|
||||
size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(__fp16));
|
||||
Q6_vgather_ARMVh(&vtcm_tmp[ii], rt, mu, vv);
|
||||
uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(__fp16);
|
||||
hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
|
||||
HVX_Vector * vtcm_tmp = (HVX_Vector *)(spad1_base + src1_ne0_padded * spad1_stride);
|
||||
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
|
||||
for (uint32_t j = 0; j < src1_ne0_padded; j += 64) {
|
||||
#pragma unroll(4)
|
||||
for (uint32_t ii = 0; ii < current_block_i; ii++) {
|
||||
size_t rt = (size_t)(spad1_base + j * spad1_stride + ii * sizeof(__fp16));
|
||||
Q6_vgather_ARMVh(&vtcm_tmp[ii], rt, mu, vv);
|
||||
uint8_t * dst_ptr = spad0_base + ii * spad0_row_bytes + (src0_ne0 + j) * sizeof(__fp16);
|
||||
hvx_vmemu(dst_ptr) = vtcm_tmp[ii];
|
||||
}
|
||||
}
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
|
||||
|
||||
dma_queue_pop(dma_q); // src0
|
||||
|
||||
const dma_addr_t dst_addr = dst_plane + i * dst->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(__fp16), current_block_i);
|
||||
|
||||
dma_queue_pop(dma_q);
|
||||
}
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
|
||||
|
||||
dma_queue_pop(dma_q); // src0
|
||||
|
||||
const dma_addr_t dst_addr = dst->data + i * dst->nb[1];
|
||||
dma_queue_push(dma_q, dma_make_data(dst_addr, spad0_base), dst->nb[1], spad0_row_bytes, (src0_ne0 + src1_ne0) * sizeof(__fp16), current_block_i);
|
||||
|
||||
dma_queue_pop(dma_q);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -230,6 +247,61 @@ static void concat_generic(unsigned int nth, unsigned int ith, void * data) {
|
||||
}
|
||||
}
|
||||
|
||||
static bool concat_dim1_contiguous_dma(struct htp_ops_context * octx, int dim, uint32_t type_size) {
|
||||
const struct htp_tensor * src0 = octx->src[0];
|
||||
const struct htp_tensor * src1 = octx->src[1];
|
||||
const struct htp_tensor * dst = octx->dst;
|
||||
|
||||
if (dim != 1 || octx->ctx->mdev.count > 1 ||
|
||||
(dst->type != HTP_TYPE_F32 && dst->type != HTP_TYPE_F16 && dst->type != HTP_TYPE_I32) ||
|
||||
src0->type != dst->type || src1->type != dst->type ||
|
||||
src0->ne[0] != dst->ne[0] || src1->ne[0] != dst->ne[0] ||
|
||||
src0->ne[2] != dst->ne[2] || src1->ne[2] != dst->ne[2] ||
|
||||
src0->ne[3] != dst->ne[3] || src1->ne[3] != dst->ne[3] ||
|
||||
dst->ne[1] != src0->ne[1] + src1->ne[1] ||
|
||||
!htp_tensor_is_contiguous(src0, type_size) ||
|
||||
!htp_tensor_is_contiguous(src1, type_size) ||
|
||||
!htp_tensor_is_contiguous(dst, type_size)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const uint32_t src0_row_size = src0->ne[0] * type_size;
|
||||
const uint32_t src1_row_size = src1->ne[0] * type_size;
|
||||
|
||||
// v75+ dma_queue_push() writes a 2D descriptor directly and does not split overflow.
|
||||
#if __HVX_ARCH__ >= 75
|
||||
if (src0_row_size > 0xffffffu || src1_row_size > 0xffffffu ||
|
||||
src0->nb[1] > 0xffffffu || src1->nb[1] > 0xffffffu || dst->nb[1] > 0xffffffu ||
|
||||
src0->ne[1] > UINT16_MAX || src1->ne[1] > UINT16_MAX) {
|
||||
return false;
|
||||
}
|
||||
#endif
|
||||
|
||||
dma_queue * q = octx->ctx->dma[0];
|
||||
|
||||
for (uint32_t i3 = 0; i3 < dst->ne[3]; ++i3) {
|
||||
for (uint32_t i2 = 0; i2 < dst->ne[2]; ++i2) {
|
||||
dma_addr_t dst_addr = dst->data + i3 * dst->nb[3] + i2 * dst->nb[2];
|
||||
dma_addr_t src0_addr = src0->data + i3 * src0->nb[3] + i2 * src0->nb[2];
|
||||
dma_addr_t src1_addr = src1->data + i3 * src1->nb[3] + i2 * src1->nb[2];
|
||||
|
||||
if (!dma_queue_push(q, dma_make_data(dst_addr, src0_addr), dst->nb[1], src0->nb[1], src0_row_size, src0->ne[1])) {
|
||||
dma_queue_flush(q);
|
||||
dma_queue_push(q, dma_make_data(dst_addr, src0_addr), dst->nb[1], src0->nb[1], src0_row_size, src0->ne[1]);
|
||||
}
|
||||
|
||||
dst_addr += src0->ne[1] * dst->nb[1];
|
||||
if (!dma_queue_push(q, dma_make_data(dst_addr, src1_addr), dst->nb[1], src1->nb[1], src1_row_size, src1->ne[1])) {
|
||||
dma_queue_flush(q);
|
||||
dma_queue_push(q, dma_make_data(dst_addr, src1_addr), dst->nb[1], src1->nb[1], src1_row_size, src1->ne[1]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
dma_queue_flush(q);
|
||||
return true;
|
||||
}
|
||||
|
||||
int op_concat(struct htp_ops_context * octx) {
|
||||
const struct htp_tensor * src0 = octx->src[0];
|
||||
const struct htp_tensor * src1 = octx->src[1];
|
||||
@@ -237,12 +309,14 @@ int op_concat(struct htp_ops_context * octx) {
|
||||
|
||||
int dim = octx->op_params[0];
|
||||
|
||||
bool is_2d = dst->ne[2] == 1 && dst->ne[3] == 1;
|
||||
|
||||
const uint32_t type_size = (dst->type == HTP_TYPE_F32 || dst->type == HTP_TYPE_I32) ? 4 : 2;
|
||||
bool is_src1_transposed = (src1->nb[0] > src1->nb[1]);
|
||||
bool is_src0_transposed = (src0->nb[0] > src0->nb[1]);
|
||||
|
||||
if (concat_dim1_contiguous_dma(octx, dim, type_size)) {
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
uint32_t n_threads = octx->n_threads;
|
||||
struct htp_concat_context cctx;
|
||||
cctx.octx = octx;
|
||||
@@ -253,7 +327,9 @@ int op_concat(struct htp_ops_context * octx) {
|
||||
|
||||
void (*worker_func)(unsigned int, unsigned int, void *) = concat_generic;
|
||||
|
||||
if (dim == 0 && is_2d && is_src1_transposed && !is_src0_transposed) {
|
||||
const bool rows_ok = src0->nb[0] == type_size && src1->nb[1] == type_size && dst->nb[0] == type_size;
|
||||
|
||||
if (dim == 0 && is_src1_transposed && !is_src0_transposed && rows_ok) {
|
||||
const uint32_t total_rows = dst->ne[1];
|
||||
const size_t dst_data_row_size = dst->ne[0] * type_size;
|
||||
uint32_t row_start = 0;
|
||||
@@ -272,6 +348,7 @@ int op_concat(struct htp_ops_context * octx) {
|
||||
|
||||
cctx.row_start = row_start;
|
||||
cctx.nrows = nrows;
|
||||
cctx.nplanes = dst->ne[2] * dst->ne[3];
|
||||
|
||||
uint32_t block_i = (type_size == 4) ? 32 : 64;
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -17,6 +17,7 @@
|
||||
#include "htp-tensor.h"
|
||||
#include "hvx-utils.h"
|
||||
#include "hvx-quant.h"
|
||||
#include "matmul-ops.h"
|
||||
#include "get-rows-ops.h"
|
||||
#include "work-queue.h"
|
||||
|
||||
@@ -28,6 +29,9 @@ struct get_rows_context {
|
||||
uint32_t task_start;
|
||||
uint32_t tasks;
|
||||
uint32_t tasks_per_thread;
|
||||
uint32_t tile_size;
|
||||
uint32_t tile_stride;
|
||||
bool index_i32;
|
||||
};
|
||||
|
||||
#define get_rows_preamble \
|
||||
@@ -195,32 +199,191 @@ static void get_rows_thread_##TYPE_NAME##_##IDX_TYPE(unsigned int nth, unsigned
|
||||
dma_queue_flush(dma_q); \
|
||||
}
|
||||
|
||||
#define F32_BYTES(n) ((n) * sizeof(float))
|
||||
#define F16_BYTES(n) ((n) * sizeof(__fp16))
|
||||
#define Q8_0_BYTES(n) (((n) / 32) * sizeof(block_q8_0))
|
||||
|
||||
GET_ROWS_THREAD_DT_FN(f32, F32_BYTES, int32_t, { if (cur_elems > 0) hvx_copy_f32_uu((uint8_t *)dst_spad, (const uint8_t *)src_spad, cur_elems); })
|
||||
GET_ROWS_THREAD_DT_FN(f32, F32_BYTES, int64_t, { if (cur_elems > 0) hvx_copy_f32_uu((uint8_t *)dst_spad, (const uint8_t *)src_spad, cur_elems); })
|
||||
static __attribute__((noinline)) void compute_get_rows_f16(float * dst_spad, const void * src_spad, uint32_t cur_elems) {
|
||||
hvx_dequantize_row_f16_f32(dst_spad, src_spad, cur_elems);
|
||||
}
|
||||
|
||||
GET_ROWS_THREAD_DT_FN(f16, F16_BYTES, int32_t, { hvx_dequantize_row_f16_f32((float *)dst_spad, src_spad, ne00); })
|
||||
GET_ROWS_THREAD_DT_FN(f16, F16_BYTES, int64_t, { hvx_dequantize_row_f16_f32((float *)dst_spad, src_spad, ne00); })
|
||||
static __attribute__((noinline)) void compute_get_rows_q8_0(float * dst_spad, const void * src_spad, uint32_t cur_elems) {
|
||||
hvx_dequantize_row_q8_0_f32(dst_spad, src_spad, cur_elems);
|
||||
}
|
||||
|
||||
GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int32_t, { hvx_dequantize_row_q8_0_f32((float *)dst_spad, src_spad, ne00); })
|
||||
GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int64_t, { hvx_dequantize_row_q8_0_f32((float *)dst_spad, src_spad, ne00); })
|
||||
GET_ROWS_THREAD_DT_FN(f16, F16_BYTES, int32_t, { compute_get_rows_f16((float *)dst_spad, src_spad, cur_elems); })
|
||||
GET_ROWS_THREAD_DT_FN(f16, F16_BYTES, int64_t, { compute_get_rows_f16((float *)dst_spad, src_spad, cur_elems); })
|
||||
|
||||
GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int32_t, { compute_get_rows_q8_0((float *)dst_spad, src_spad, cur_elems); })
|
||||
GET_ROWS_THREAD_DT_FN(q8_0, Q8_0_BYTES, int64_t, { compute_get_rows_q8_0((float *)dst_spad, src_spad, cur_elems); })
|
||||
|
||||
|
||||
static __attribute__((noinline)) void compute_get_rows_tiled(float * dst, const uint8_t * tile, uint32_t row, bool q4) {
|
||||
const HVX_VectorPred first2 = Q6_Q_vsetq_R(2);
|
||||
const HVX_VectorPred first4 = Q6_Q_vsetq_R(4);
|
||||
HVX_Vector vq = Q6_V_vzero();
|
||||
if (q4) {
|
||||
const HVX_VectorPred first1 = Q6_Q_vsetq_R(1);
|
||||
const HVX_VectorPred first3 = Q6_Q_vsetq_R(3);
|
||||
for (int group = 3; group >= 0; --group) {
|
||||
const HVX_Vector v = Q6_V_vror_VR(hvx_vmem(tile + group * VLEN), row);
|
||||
// Four planes contribute bytes at 0, 32, 64 and 96 after rotation.
|
||||
HVX_Vector packed = Q6_V_vmux_QVV(first1, v, Q6_V_vror_VR(v, 31));
|
||||
packed = Q6_V_vmux_QVV(first2, packed, Q6_V_vror_VR(v, 62));
|
||||
packed = Q6_V_vmux_QVV(first3, packed, Q6_V_vror_VR(v, 93));
|
||||
vq = Q6_V_vmux_QVV(first4, packed, Q6_V_vror_VR(vq, VLEN - 4));
|
||||
}
|
||||
const HVX_Vector lo = Q6_V_vand_VV(vq, Q6_Vb_vsplat_R(0x0F));
|
||||
const HVX_Vector hi = Q6_Vub_vlsr_VubR(vq, 4);
|
||||
vq = Q6_V_lo_W(Q6_W_vshuff_VVR(hi, lo, -1));
|
||||
vq = Q6_Vb_vsub_VbVb(vq, Q6_Vb_vsplat_R(8));
|
||||
} else {
|
||||
for (int group = 7; group >= 0; --group) {
|
||||
const HVX_Vector v = Q6_V_vror_VR(hvx_vmem(tile + group * VLEN), 2 * row);
|
||||
// Two planes contribute halfwords at 0 and 64 after rotation.
|
||||
const HVX_Vector packed = Q6_V_vmux_QVV(first2, v, Q6_V_vror_VR(v, 62));
|
||||
vq = Q6_V_vmux_QVV(first4, packed, Q6_V_vror_VR(vq, VLEN - 4));
|
||||
}
|
||||
}
|
||||
const HVX_Vector scales = hvx_vmem(tile + (q4 ? 512 : 1024));
|
||||
const HVX_Vector scale_hf = hvx_vec_repl_f16(Q6_V_vror_VR(scales, 2 * row));
|
||||
const HVX_Vector scale = Q6_V_lo_W(hvx_vec_f16_to_f32(scale_hf));
|
||||
const HVX_VectorPair p16 = Q6_Wh_vunpack_Vb(vq);
|
||||
const HVX_VectorPair p32 = Q6_Ww_vunpack_Vh(Q6_V_lo_W(p16));
|
||||
const HVX_Vector values = hvx_vec_mul_f32_f32(Q6_Vsf_equals_Vw(Q6_V_lo_W(p32)), scale);
|
||||
*(HVX_Vector *) dst = values;
|
||||
}
|
||||
|
||||
struct get_rows_tiled_task {
|
||||
dma_addr_t tile_src_base;
|
||||
dma_addr_t dst_data;
|
||||
uint32_t row;
|
||||
};
|
||||
|
||||
static inline struct get_rows_tiled_task get_rows_tiled_calc_task(
|
||||
const struct htp_ops_context * octx,
|
||||
const struct get_rows_context * grctx,
|
||||
uint32_t i,
|
||||
uint32_t n_k_tiles,
|
||||
uint32_t tile_size
|
||||
) {
|
||||
const struct htp_get_rows_kernel_params * kparams = grctx->kparams;
|
||||
get_rows_preamble;
|
||||
|
||||
const uint32_t i12 = fastdiv(i, &kparams->div_ne10_ne11);
|
||||
const uint32_t rem = i - i12 * ne11 * ne10;
|
||||
const uint32_t i11 = fastdiv(rem, &kparams->div_ne10);
|
||||
const uint32_t i10 = rem - i11 * ne10;
|
||||
const dma_addr_t src1_data = octx->src[1]->data + i10*nb10 + i11*nb11 + i12*nb12;
|
||||
const uint32_t i01 = grctx->index_i32 ? *(const int32_t *)(uintptr_t) src1_data : (uint32_t) *(const int64_t *)(uintptr_t) src1_data;
|
||||
assert(i01 < ne01);
|
||||
|
||||
const uint32_t q02 = fastdiv(i11, &kparams->div_ne02);
|
||||
const uint32_t i02 = i11 - q02 * ne02;
|
||||
const uint32_t q03 = fastdiv(i12, &kparams->div_ne03);
|
||||
const uint32_t i03 = i12 - q03 * ne03;
|
||||
const uint32_t column_tile = i01 / HTP_MM_HMX_TILE_N_ROWS;
|
||||
const uint32_t row = i01 % HTP_MM_HMX_TILE_N_ROWS;
|
||||
const dma_addr_t matrix = octx->src[0]->data + i02*nb02 + i03*nb03;
|
||||
|
||||
struct get_rows_tiled_task task;
|
||||
task.tile_src_base = matrix + (column_tile * n_k_tiles) * tile_size;
|
||||
task.dst_data = octx->dst->data + i10*nb1 + i11*nb2 + i12*nb3;
|
||||
task.row = row;
|
||||
return task;
|
||||
}
|
||||
|
||||
static void get_rows_thread_tiled(unsigned int nth, unsigned int ith, void * data) {
|
||||
struct get_rows_context * grctx = (struct get_rows_context *) data;
|
||||
struct htp_ops_context * octx = grctx->octx;
|
||||
const struct htp_get_rows_kernel_params * kparams = grctx->kparams;
|
||||
get_rows_preamble;
|
||||
|
||||
const uint32_t dr = grctx->tasks_per_thread;
|
||||
const uint32_t ir0 = grctx->task_start + dr * ith;
|
||||
if (ir0 >= grctx->task_start + grctx->tasks) {
|
||||
return;
|
||||
}
|
||||
|
||||
const uint32_t ir1 = MIN(ir0 + dr, grctx->task_start + grctx->tasks);
|
||||
const uint32_t n_k_tiles = ne00 / HTP_MM_HMX_TILE_N_COLS;
|
||||
const struct htp_get_rows_vtcm_layout * vtcm_layout = &grctx->vtcm_layout;
|
||||
uint8_t * src_spad_base = grctx->vtcm_base + vtcm_layout->off_src0 + ith * vtcm_layout->src0_bytes_per_thread;
|
||||
uint8_t * dst_spad_base = grctx->vtcm_base + vtcm_layout->off_dst + ith * vtcm_layout->dst_bytes_per_thread;
|
||||
dma_queue * dma_q = octx->ctx->dma[ith];
|
||||
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
|
||||
|
||||
const uint32_t tile_size = grctx->tile_size;
|
||||
const uint32_t tile_stride = grctx->tile_stride;
|
||||
const uint32_t dst_bytes = ne00 * sizeof(float);
|
||||
const bool is_q4 = (octx->src[0]->type == HTP_TYPE_Q4_0);
|
||||
|
||||
for (uint32_t step = 0, spad_idx = 0; step < ir1 - ir0 && spad_idx < 2; ++step, ++spad_idx) {
|
||||
const uint32_t i = ir0 + step;
|
||||
struct get_rows_tiled_task task = get_rows_tiled_calc_task(octx, grctx, i, n_k_tiles, tile_size);
|
||||
|
||||
// Dummy writeback to prime the queue with dst descriptor
|
||||
dma_queue_push(dma_q,
|
||||
dma_make_data(task.dst_data, dst_spad_base + spad_idx * vtcm_layout->dst_spad_half_size),
|
||||
dst_bytes, vtcm_layout->dst_spad_half_size, dst_bytes, 0);
|
||||
|
||||
// Prefetch row tiles
|
||||
dma_queue_push(dma_q,
|
||||
dma_make_data(src_spad_base + spad_idx * vtcm_layout->src0_spad_half_size, task.tile_src_base),
|
||||
tile_stride, tile_size, tile_size, n_k_tiles);
|
||||
}
|
||||
|
||||
for (uint32_t step = 0; step < ir1 - ir0; ++step) {
|
||||
const uint32_t i = ir0 + step;
|
||||
float * dst_spad = (float *) dma_queue_pop(dma_q).src;
|
||||
uint8_t * src_spad = (uint8_t *) dma_queue_pop(dma_q).dst;
|
||||
|
||||
struct get_rows_tiled_task task = get_rows_tiled_calc_task(octx, grctx, i, n_k_tiles, tile_size);
|
||||
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
|
||||
for (uint32_t k_tile = 0; k_tile < n_k_tiles; ++k_tile) {
|
||||
const uint8_t * tile = src_spad + k_tile * tile_stride;
|
||||
float * dst_block = dst_spad + k_tile * HTP_MM_HMX_TILE_N_COLS;
|
||||
compute_get_rows_tiled(dst_block, tile, task.row, is_q4);
|
||||
}
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) i);
|
||||
|
||||
// Real writeback of dst_spad
|
||||
dma_queue_push(dma_q,
|
||||
dma_make_data(task.dst_data, dst_spad),
|
||||
dst_bytes, vtcm_layout->dst_spad_half_size, dst_bytes, 1);
|
||||
|
||||
const uint32_t next_step = step + 2;
|
||||
if (next_step < ir1 - ir0) {
|
||||
const uint32_t ni = ir0 + next_step;
|
||||
struct get_rows_tiled_task next_task = get_rows_tiled_calc_task(octx, grctx, ni, n_k_tiles, tile_size);
|
||||
dma_queue_push(dma_q,
|
||||
dma_make_data(src_spad, next_task.tile_src_base),
|
||||
tile_stride, tile_size, tile_size, n_k_tiles);
|
||||
}
|
||||
}
|
||||
|
||||
dma_queue_flush(dma_q);
|
||||
}
|
||||
|
||||
int op_get_rows(struct htp_ops_context * octx) {
|
||||
const struct htp_get_rows_kernel_params * kparams = (const struct htp_get_rows_kernel_params *) octx->kernel_params;
|
||||
|
||||
if (octx->src[0]->type != HTP_TYPE_F32 &&
|
||||
octx->src[0]->type != HTP_TYPE_F16 &&
|
||||
octx->src[0]->type != HTP_TYPE_Q8_0 &&
|
||||
octx->src[0]->type != HTP_TYPE_I32) {
|
||||
octx->src[0]->type != HTP_TYPE_F16 &&
|
||||
octx->src[0]->type != HTP_TYPE_Q4_0 &&
|
||||
octx->src[0]->type != HTP_TYPE_Q8_0 &&
|
||||
octx->src[0]->type != HTP_TYPE_I32) {
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
|
||||
if ((octx->src[0]->type == HTP_TYPE_I32 && octx->dst->type != HTP_TYPE_I32) ||
|
||||
(octx->src[0]->type != HTP_TYPE_I32 && octx->dst->type != HTP_TYPE_F32)) {
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
if (kparams->kernel_type == HTP_GET_ROWS_KERNEL_SAMETYPE) {
|
||||
if (octx->src[0]->type != octx->dst->type) {
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
} else {
|
||||
if (octx->dst->type != HTP_TYPE_F32) {
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
}
|
||||
|
||||
if (octx->src[1]->type != HTP_TYPE_I32 && octx->src[1]->type != HTP_TYPE_I64) {
|
||||
@@ -262,33 +425,48 @@ int op_get_rows(struct htp_ops_context * octx) {
|
||||
grctx.vtcm_base = (uint8_t *)octx->ctx->vtcm_base;
|
||||
grctx.task_start = task_start;
|
||||
grctx.tasks = tasks;
|
||||
grctx.tasks_per_thread = fastdiv(tasks + n_threads - 1, &octx->n_threads_div);
|
||||
grctx.tasks_per_thread = octx->ctx->mdev.count == 1 ? kparams->tasks_per_thread : fastdiv(tasks + n_threads - 1, &octx->n_threads_div);
|
||||
grctx.tile_size = octx->src[0]->type == HTP_TYPE_Q4_0 ? HTP_MM_WEIGHT_TILE_SIZE_Q4_0 : HTP_MM_WEIGHT_TILE_SIZE_Q8_0;
|
||||
grctx.tile_stride = (grctx.tile_size + 127) & ~127;
|
||||
grctx.index_i32 = octx->src[1]->type == HTP_TYPE_I32;
|
||||
|
||||
const uint32_t ne00 = octx->src[0]->ne[0];
|
||||
htp_get_rows_vtcm_layout_build(&grctx.vtcm_layout, octx->src[0]->type, ne00, n_threads);
|
||||
htp_get_rows_vtcm_layout_build(&grctx.vtcm_layout, kparams->kernel_type, octx->src[0]->type, ne00, n_threads);
|
||||
|
||||
if (grctx.vtcm_layout.total_bytes > octx->ctx->vtcm_size) {
|
||||
FARF(ERROR, "get-rows: VTCM reservation %zu is too small, needed %zu\n",
|
||||
octx->ctx->vtcm_size, grctx.vtcm_layout.total_bytes);
|
||||
return HTP_STATUS_INVAL_PARAMS;
|
||||
}
|
||||
|
||||
const bool is_i32 = (octx->src[1]->type == HTP_TYPE_I32);
|
||||
|
||||
work_queue_func_t q_func = NULL;
|
||||
if (kparams->use_dma) {
|
||||
q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_st_int32_t : get_rows_thread_st_int64_t);
|
||||
} else {
|
||||
switch (octx->src[0]->type) {
|
||||
case HTP_TYPE_F32: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_f32_int32_t : get_rows_thread_f32_int64_t); break;
|
||||
case HTP_TYPE_F16: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_f16_int32_t : get_rows_thread_f16_int64_t); break;
|
||||
case HTP_TYPE_Q8_0: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_q8_0_int32_t : get_rows_thread_q8_0_int64_t); break;
|
||||
case HTP_TYPE_I32: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_st_int32_t : get_rows_thread_st_int64_t); break;
|
||||
default: return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
switch (kparams->kernel_type) {
|
||||
case HTP_GET_ROWS_KERNEL_SAMETYPE:
|
||||
q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_st_int32_t : get_rows_thread_st_int64_t);
|
||||
break;
|
||||
case HTP_GET_ROWS_KERNEL_TILED:
|
||||
q_func = get_rows_thread_tiled;
|
||||
break;
|
||||
case HTP_GET_ROWS_KERNEL_FLAT:
|
||||
switch (octx->src[0]->type) {
|
||||
case HTP_TYPE_F16: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_f16_int32_t : get_rows_thread_f16_int64_t); break;
|
||||
case HTP_TYPE_Q8_0: q_func = (work_queue_func_t)(is_i32 ? get_rows_thread_q8_0_int32_t : get_rows_thread_q8_0_int64_t); break;
|
||||
default: return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
break;
|
||||
default:
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
|
||||
FARF(HIGH, "get-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu use-dma %d n-threads %d\n",
|
||||
FARF(HIGH, "get-rows: (%ux%ux%ux%u) x (%ux%ux%ux%u) -> (%ux%ux%ux%u) : src0-vtcm-size %zu dst-vtcm-size %zu kernel-type %d n-threads %d\n",
|
||||
octx->src[0]->ne[0], octx->src[0]->ne[1], octx->src[0]->ne[2], octx->src[0]->ne[3],
|
||||
octx->src[1]->ne[0], octx->src[1]->ne[1], octx->src[1]->ne[2], octx->src[1]->ne[3],
|
||||
octx->dst->ne[0], octx->dst->ne[1], octx->dst->ne[2], octx->dst->ne[3],
|
||||
grctx.vtcm_layout.src0_bytes_per_thread * n_threads,
|
||||
grctx.vtcm_layout.dst_bytes_per_thread * n_threads,
|
||||
kparams->use_dma, n_threads);
|
||||
kparams->kernel_type, n_threads);
|
||||
|
||||
work_queue_run(octx->ctx->work_queue, q_func, &grctx, n_threads);
|
||||
return HTP_STATUS_OK;
|
||||
|
||||
@@ -1,11 +1,21 @@
|
||||
#ifndef HTP_GET_ROWS_OPS_H
|
||||
#define HTP_GET_ROWS_OPS_H
|
||||
|
||||
#include <stdbool.h>
|
||||
#include <string.h>
|
||||
|
||||
#include "hex-fastdiv.h"
|
||||
#include "matmul-ops.h"
|
||||
|
||||
enum htp_get_rows_kernel_type {
|
||||
HTP_GET_ROWS_KERNEL_SAMETYPE = 0,
|
||||
HTP_GET_ROWS_KERNEL_TILED,
|
||||
HTP_GET_ROWS_KERNEL_FLAT,
|
||||
};
|
||||
|
||||
struct htp_get_rows_kernel_params {
|
||||
int32_t n_threads;
|
||||
int32_t use_dma;
|
||||
int32_t kernel_type;
|
||||
int32_t chunks_per_row;
|
||||
int32_t chunk_size;
|
||||
int32_t total_tasks;
|
||||
@@ -34,19 +44,37 @@ struct htp_get_rows_vtcm_layout {
|
||||
|
||||
static inline void htp_get_rows_vtcm_layout_build(
|
||||
struct htp_get_rows_vtcm_layout * vtcm_layout,
|
||||
int kernel_type,
|
||||
int type,
|
||||
uint32_t ne00,
|
||||
uint32_t n_threads) {
|
||||
|
||||
if (kernel_type == HTP_GET_ROWS_KERNEL_SAMETYPE) {
|
||||
memset(vtcm_layout, 0, sizeof(*vtcm_layout));
|
||||
return;
|
||||
}
|
||||
|
||||
if (kernel_type == HTP_GET_ROWS_KERNEL_TILED) {
|
||||
const size_t tile_size = type == HTP_TYPE_Q4_0 ? HTP_MM_WEIGHT_TILE_SIZE_Q4_0 : HTP_MM_WEIGHT_TILE_SIZE_Q8_0;
|
||||
const size_t tile_stride = (tile_size + 127) & ~127;
|
||||
const uint32_t n_k_tiles = ne00 / HTP_MM_HMX_TILE_N_COLS;
|
||||
const size_t row_tiles_size = n_k_tiles > 0 ? (n_k_tiles * tile_stride) : tile_stride;
|
||||
vtcm_layout->src0_spad_half_size = (row_tiles_size + 255) & ~255;
|
||||
vtcm_layout->dst_spad_half_size = (ne00 * sizeof(float) + 255) & ~255;
|
||||
vtcm_layout->src0_bytes_per_thread = 2 * vtcm_layout->src0_spad_half_size;
|
||||
vtcm_layout->dst_bytes_per_thread = 2 * vtcm_layout->dst_spad_half_size;
|
||||
vtcm_layout->off_src0 = 0;
|
||||
vtcm_layout->off_dst = vtcm_layout->src0_bytes_per_thread * n_threads;
|
||||
vtcm_layout->total_bytes = vtcm_layout->off_dst + vtcm_layout->dst_bytes_per_thread * n_threads;
|
||||
return;
|
||||
}
|
||||
|
||||
uint32_t src0_row_size = 0;
|
||||
switch (type) {
|
||||
case 0: // HTP_TYPE_F32
|
||||
src0_row_size = ne00 * 4;
|
||||
break;
|
||||
case 1: // HTP_TYPE_F16
|
||||
case HTP_TYPE_F16:
|
||||
src0_row_size = ne00 * 2;
|
||||
break;
|
||||
case 8: // HTP_TYPE_Q8_0
|
||||
case HTP_TYPE_Q8_0:
|
||||
src0_row_size = (ne00 / 32) * 34;
|
||||
break;
|
||||
default:
|
||||
|
||||
@@ -506,6 +506,111 @@ static void dequantize_tiled_weight_to_fp16_task_q8_0(
|
||||
}
|
||||
}
|
||||
|
||||
static void dequantize_tiled_weight_to_fp16_task_q5_k(
|
||||
const tiled_dequantize_state_t *state,
|
||||
uint32_t start_tile, uint32_t end_tile) {
|
||||
|
||||
const HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
|
||||
|
||||
for (uint32_t t = start_tile; t < end_tile; t++) {
|
||||
const uint8_t * tile_src = state->src + t * state->aligned_tile_size;
|
||||
__fp16 * dst_ptr = state->dst + t * HTP_MM_HMX_TILE_N_ELMS;
|
||||
|
||||
HVX_Vector vscale_offset = hvx_vmem(tile_src + 512);
|
||||
HVX_VectorPair dm_deal = Q6_W_vdeal_VVR(vscale_offset, vscale_offset, -2);
|
||||
HVX_Vector vd = Q6_V_lo_W(dm_deal);
|
||||
HVX_Vector vm = Q6_V_hi_W(dm_deal);
|
||||
|
||||
HVX_Vector v_scale_duplicated = Q6_V_lo_W(Q6_W_vshuff_VVR(vd, vd, -2));
|
||||
HVX_Vector v_offset_duplicated = Q6_V_lo_W(Q6_W_vshuff_VVR(vm, vm, -2));
|
||||
|
||||
// Load all 4 groups in parallel
|
||||
HVX_Vector vq0 = hvx_vmem(tile_src + 0 * 128);
|
||||
HVX_Vector vq1 = hvx_vmem(tile_src + 1 * 128);
|
||||
HVX_Vector vq2 = hvx_vmem(tile_src + 2 * 128);
|
||||
HVX_Vector vq3 = hvx_vmem(tile_src + 3 * 128);
|
||||
|
||||
// Nibble extraction
|
||||
HVX_Vector v_lo0 = Q6_V_vand_VV(vq0, mask_h4);
|
||||
HVX_Vector v_hi0 = Q6_Vub_vlsr_VubR(vq0, 4);
|
||||
HVX_Vector v_lo1 = Q6_V_vand_VV(vq1, mask_h4);
|
||||
HVX_Vector v_hi1 = Q6_Vub_vlsr_VubR(vq1, 4);
|
||||
HVX_Vector v_lo2 = Q6_V_vand_VV(vq2, mask_h4);
|
||||
HVX_Vector v_hi2 = Q6_Vub_vlsr_VubR(vq2, 4);
|
||||
HVX_Vector v_lo3 = Q6_V_vand_VV(vq3, mask_h4);
|
||||
HVX_Vector v_hi3 = Q6_Vub_vlsr_VubR(vq3, 4);
|
||||
|
||||
// Q5_K: OR in the 5th bit from the plane
|
||||
HVX_Vector v_plane = hvx_vmem(tile_src + 640);
|
||||
v_lo0 = hvx_q5k_or_hibit(v_lo0, v_plane, 0);
|
||||
v_hi0 = hvx_q5k_or_hibit(v_hi0, v_plane, 1);
|
||||
v_lo1 = hvx_q5k_or_hibit(v_lo1, v_plane, 2);
|
||||
v_hi1 = hvx_q5k_or_hibit(v_hi1, v_plane, 3);
|
||||
v_lo2 = hvx_q5k_or_hibit(v_lo2, v_plane, 4);
|
||||
v_hi2 = hvx_q5k_or_hibit(v_hi2, v_plane, 5);
|
||||
v_lo3 = hvx_q5k_or_hibit(v_lo3, v_plane, 6);
|
||||
v_hi3 = hvx_q5k_or_hibit(v_hi3, v_plane, 7);
|
||||
|
||||
// Shuffling
|
||||
HVX_VectorPair vp_shuf0 = Q6_W_vshuff_VVR(v_hi0, v_lo0, -1);
|
||||
HVX_VectorPair vp_shuf1 = Q6_W_vshuff_VVR(v_hi1, v_lo1, -1);
|
||||
HVX_VectorPair vp_shuf2 = Q6_W_vshuff_VVR(v_hi2, v_lo2, -1);
|
||||
HVX_VectorPair vp_shuf3 = Q6_W_vshuff_VVR(v_hi3, v_lo3, -1);
|
||||
|
||||
// Unpack to 16-bit
|
||||
HVX_VectorPair vp_int16_lo0 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf0));
|
||||
HVX_VectorPair vp_int16_hi0 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf0));
|
||||
HVX_VectorPair vp_int16_lo1 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf1));
|
||||
HVX_VectorPair vp_int16_hi1 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf1));
|
||||
HVX_VectorPair vp_int16_lo2 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf2));
|
||||
HVX_VectorPair vp_int16_hi2 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf2));
|
||||
HVX_VectorPair vp_int16_lo3 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf3));
|
||||
HVX_VectorPair vp_int16_hi3 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf3));
|
||||
|
||||
// Convert, multiply, add offset
|
||||
HVX_Vector v_grp0_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo0)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp0_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo0)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp0_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi0)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp0_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi0)), v_scale_duplicated), v_offset_duplicated));
|
||||
|
||||
HVX_Vector v_grp1_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo1)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp1_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo1)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp1_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi1)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp1_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi1)), v_scale_duplicated), v_offset_duplicated));
|
||||
|
||||
HVX_Vector v_grp2_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo2)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp2_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo2)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp2_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi2)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp2_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi2)), v_scale_duplicated), v_offset_duplicated));
|
||||
|
||||
HVX_Vector v_grp3_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo3)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp3_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo3)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp3_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi3)), v_scale_duplicated), v_offset_duplicated));
|
||||
HVX_Vector v_grp3_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi3)), v_scale_duplicated), v_offset_duplicated));
|
||||
|
||||
// Parallel Stores
|
||||
hvx_vmem(dst_ptr + 0 * 64) = v_grp0_0;
|
||||
hvx_vmem(dst_ptr + 1 * 64) = v_grp0_1;
|
||||
hvx_vmem(dst_ptr + 2 * 64) = v_grp0_2;
|
||||
hvx_vmem(dst_ptr + 3 * 64) = v_grp0_3;
|
||||
|
||||
hvx_vmem(dst_ptr + 4 * 64) = v_grp1_0;
|
||||
hvx_vmem(dst_ptr + 5 * 64) = v_grp1_1;
|
||||
hvx_vmem(dst_ptr + 6 * 64) = v_grp1_2;
|
||||
hvx_vmem(dst_ptr + 7 * 64) = v_grp1_3;
|
||||
|
||||
hvx_vmem(dst_ptr + 8 * 64) = v_grp2_0;
|
||||
hvx_vmem(dst_ptr + 9 * 64) = v_grp2_1;
|
||||
hvx_vmem(dst_ptr + 10 * 64) = v_grp2_2;
|
||||
hvx_vmem(dst_ptr + 11 * 64) = v_grp2_3;
|
||||
|
||||
hvx_vmem(dst_ptr + 12 * 64) = v_grp3_0;
|
||||
hvx_vmem(dst_ptr + 13 * 64) = v_grp3_1;
|
||||
hvx_vmem(dst_ptr + 14 * 64) = v_grp3_2;
|
||||
hvx_vmem(dst_ptr + 15 * 64) = v_grp3_3;
|
||||
}
|
||||
}
|
||||
|
||||
// Q6_K stores 6-bit weights and one fp16 scale per 16 k, see HTP_MM_WEIGHT_TILE_SIZE_Q6_K.
|
||||
// A k-group holds 4 k per row, the HMX tile holds 2, so each group is dealt into two tiles.
|
||||
static void dequantize_tiled_weight_to_fp16_task_q6_k(
|
||||
@@ -914,98 +1019,6 @@ typedef struct {
|
||||
|
||||
// activations : fp32 -> fp16
|
||||
|
||||
static void transfer_activation_chunk_fp32_to_fp16(__fp16 *restrict vtcm_dst, const float *restrict src, uint32_t n_rows, uint32_t k_block, uint32_t k_stride, uint32_t k_valid) {
|
||||
const uint32_t n_rows_padded = hex_align_up(n_rows, HTP_MM_HMX_TILE_N_ROWS);
|
||||
const uint32_t n_rows_tiled = (n_rows / HTP_MM_HMX_TILE_N_ROWS) * HTP_MM_HMX_TILE_N_ROWS;
|
||||
|
||||
uint32_t r = 0;
|
||||
|
||||
#pragma unroll(2)
|
||||
for (r = 0; r < n_rows_tiled; r += 2) {
|
||||
uint32_t r0 = r / HTP_MM_HMX_TILE_N_ROWS; // tile row index
|
||||
uint32_t r1 = r % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
|
||||
|
||||
const float *ptr_in0 = src + (r + 0) * k_stride;
|
||||
const float *ptr_in1 = src + (r + 1) * k_stride;
|
||||
|
||||
uint32_t c = 0;
|
||||
for (; c + 32 <= k_valid; c += 32) {
|
||||
HVX_Vector v0 = *(const HVX_Vector *)(ptr_in0 + c);
|
||||
HVX_Vector v1 = *(const HVX_Vector *)(ptr_in1 + c);
|
||||
HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
|
||||
|
||||
uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
|
||||
uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
|
||||
|
||||
HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
|
||||
tile[r1 / 2] = v_out;
|
||||
}
|
||||
if (c < k_block) {
|
||||
HVX_Vector v0 = *(const HVX_Vector *)(ptr_in0 + c);
|
||||
HVX_Vector v1 = *(const HVX_Vector *)(ptr_in1 + c);
|
||||
|
||||
uint32_t rem = k_valid - c;
|
||||
HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
|
||||
v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
|
||||
v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
|
||||
|
||||
HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
|
||||
|
||||
uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
|
||||
uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
|
||||
|
||||
HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
|
||||
tile[r1 / 2] = v_out;
|
||||
}
|
||||
}
|
||||
|
||||
for (; r < n_rows_padded; r += 2) {
|
||||
uint32_t r0 = r / HTP_MM_HMX_TILE_N_ROWS; // tile row index
|
||||
uint32_t r1 = r % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
|
||||
|
||||
const bool row0_valid = r < n_rows;
|
||||
const bool row1_valid = (r + 1) < n_rows;
|
||||
|
||||
const float *ptr_in0 = row0_valid ? (src + (r + 0) * k_stride) : NULL;
|
||||
const float *ptr_in1 = row1_valid ? (src + (r + 1) * k_stride) : NULL;
|
||||
|
||||
uint32_t c = 0;
|
||||
for (; c + 32 <= k_valid; c += 32) {
|
||||
HVX_Vector v0 = Q6_V_vzero();
|
||||
HVX_Vector v1 = Q6_V_vzero();
|
||||
if (row0_valid) v0 = *(const HVX_Vector *)(ptr_in0 + c);
|
||||
if (row1_valid) v1 = *(const HVX_Vector *)(ptr_in1 + c);
|
||||
|
||||
HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
|
||||
|
||||
uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
|
||||
uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
|
||||
|
||||
HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
|
||||
tile[r1 / 2] = v_out;
|
||||
}
|
||||
if (c < k_block) {
|
||||
HVX_Vector v0 = Q6_V_vzero();
|
||||
HVX_Vector v1 = Q6_V_vzero();
|
||||
if (row0_valid) v0 = *(const HVX_Vector *)(ptr_in0 + c);
|
||||
if (row1_valid) v1 = *(const HVX_Vector *)(ptr_in1 + c);
|
||||
|
||||
uint32_t rem = k_valid - c;
|
||||
HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
|
||||
v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
|
||||
v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
|
||||
|
||||
HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
|
||||
|
||||
uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
|
||||
uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
|
||||
|
||||
HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
|
||||
tile[r1 / 2] = v_out;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
static void transfer_activation_row_pair_fp32_to_fp16(
|
||||
__fp16 *restrict vtcm_dst,
|
||||
const float *restrict row0,
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -23,6 +23,7 @@ enum htp_data_type {
|
||||
HTP_TYPE_Q4_1 = 3,
|
||||
HTP_TYPE_Q8_0 = 8,
|
||||
HTP_TYPE_Q4_K = 12,
|
||||
HTP_TYPE_Q5_K = 13,
|
||||
HTP_TYPE_Q6_K = 14,
|
||||
HTP_TYPE_IQ4_NL = 20,
|
||||
HTP_TYPE_I32 = 26,
|
||||
@@ -68,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,
|
||||
@@ -85,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,
|
||||
@@ -107,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);
|
||||
|
||||
@@ -121,29 +121,43 @@ static inline void quantize_block_f32_q8_0_tiled(float * restrict x, uint8_t * r
|
||||
HVX_Vector * vx = (HVX_Vector *) x;
|
||||
HVX_Vector zero = Q6_V_vzero();
|
||||
|
||||
HVX_Vector vmax0_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[0]));
|
||||
HVX_Vector vmax1_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[1]));
|
||||
HVX_Vector vmax2_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[2]));
|
||||
HVX_Vector vmax3_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[3]));
|
||||
|
||||
HVX_Vector vx0_qf = Q6_Vqf32_vsub_VsfVsf(vx[0], zero);
|
||||
HVX_Vector vx1_qf = Q6_Vqf32_vsub_VsfVsf(vx[1], zero);
|
||||
HVX_Vector vx2_qf = Q6_Vqf32_vsub_VsfVsf(vx[2], zero);
|
||||
HVX_Vector vx3_qf = Q6_Vqf32_vsub_VsfVsf(vx[3], zero);
|
||||
|
||||
HVX_Vector vmax0_qf = Q6_Vqf32_vsub_VsfVsf(vmax0_sf, zero);
|
||||
HVX_Vector vmax1_qf = Q6_Vqf32_vsub_VsfVsf(vmax1_sf, zero);
|
||||
HVX_Vector vmax2_qf = Q6_Vqf32_vsub_VsfVsf(vmax2_sf, zero);
|
||||
HVX_Vector vmax3_qf = Q6_Vqf32_vsub_VsfVsf(vmax3_sf, zero);
|
||||
|
||||
HVX_Vector vmax01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax1_qf, vmax0_qf)));
|
||||
HVX_Vector vmax23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax3_qf, vmax2_qf)));
|
||||
|
||||
HVX_Vector vx01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx1_qf, vx0_qf)));
|
||||
HVX_Vector vx23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx3_qf, vx2_qf)));
|
||||
|
||||
HVX_Vector vmax_hf = hvx_vec_reduce_max_f16(hvx_vec_abs_f16(vx01_hf));
|
||||
vmax_hf = hvx_vec_reduce_max2_f16(hvx_vec_abs_f16(vx23_hf), vmax_hf);
|
||||
HVX_Vector vd01_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax01_hf, Q6_Vh_vsplat_R(0x2008));
|
||||
HVX_Vector vd23_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax23_hf, Q6_Vh_vsplat_R(0x2008));
|
||||
HVX_Vector vd01_hf = Q6_Vhf_equals_Vqf16(vd01_qf16);
|
||||
HVX_Vector vd23_hf = Q6_Vhf_equals_Vqf16(vd23_qf16);
|
||||
|
||||
HVX_Vector vd_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax_hf, Q6_Vh_vsplat_R(0x2008));
|
||||
HVX_Vector vd_hf = Q6_Vhf_equals_Vqf16(vd_qf16);
|
||||
|
||||
HVX_Vector vd_inv_hf = hvx_vec_inverse_f16(vd_hf);
|
||||
vx01_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx01_hf, vd_inv_hf));
|
||||
vx23_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx23_hf, vd_inv_hf));
|
||||
HVX_Vector vd01_inv_hf = hvx_vec_inverse_f16(vd01_hf);
|
||||
HVX_Vector vd23_inv_hf = hvx_vec_inverse_f16(vd23_hf);
|
||||
vx01_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx01_hf, vd01_inv_hf));
|
||||
vx23_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx23_hf, vd23_inv_hf));
|
||||
|
||||
HVX_Vector vx01_i16 = hvx_vec_i16_from_hf_rnd_sat(vx01_hf);
|
||||
HVX_Vector vx23_i16 = hvx_vec_i16_from_hf_rnd_sat(vx23_hf);
|
||||
HVX_Vector vx_i8 = Q6_Vb_vpack_VhVh_sat(vx23_i16, vx01_i16);
|
||||
|
||||
HVX_Vector r_scale = hvx_vec_repl_f16(vd_hf);
|
||||
HVX_VectorPair vp01 = Q6_W_vshuff_VVR(vd01_hf, vd01_hf, -64);
|
||||
HVX_VectorPair vp23 = Q6_W_vshuff_VVR(vd23_hf, vd23_hf, -64);
|
||||
|
||||
static const uint8_t __attribute__((aligned(128))) repl[128] = {
|
||||
0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
||||
@@ -157,8 +171,19 @@ static inline void quantize_block_f32_q8_0_tiled(float * restrict x, uint8_t * r
|
||||
};
|
||||
HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
|
||||
|
||||
#pragma unroll
|
||||
for (int b = 0; b < 4; b++) {
|
||||
HVX_Vector v_act = Q6_V_vror_VR(vx_i8, b * 32);
|
||||
HVX_Vector r_scale;
|
||||
if (b == 0) {
|
||||
r_scale = Q6_V_lo_W(vp01);
|
||||
} else if (b == 1) {
|
||||
r_scale = Q6_V_hi_W(vp01);
|
||||
} else if (b == 2) {
|
||||
r_scale = Q6_V_lo_W(vp23);
|
||||
} else {
|
||||
r_scale = Q6_V_hi_W(vp23);
|
||||
}
|
||||
|
||||
HVX_Vector r0 = Q6_V_vdelta_VV(v_act, v_repl_ctrl);
|
||||
HVX_Vector r1 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 4), v_repl_ctrl);
|
||||
@@ -371,6 +396,78 @@ static inline HVX_VectorPair accum_q8_0_32x2(
|
||||
return Q6_W_vcombine_VV(v_sum1, v_sum0);
|
||||
}
|
||||
|
||||
// Q5_K: OR 0x10 into every lane of v whose flag j is set in the plane (see HTP_MM_WEIGHT_TILE_SIZE_Q5_K)
|
||||
static inline HVX_Vector hvx_q5k_or_hibit(HVX_Vector v, HVX_Vector v_plane, int j) {
|
||||
HVX_VectorPred q = Q6_Q_vand_VR(v_plane, 0x01010101u << j);
|
||||
return Q6_V_vandor_VQR(v, q, 0x10101010);
|
||||
}
|
||||
|
||||
// 5-bit variant: the high bit comes from the plane, see hvx_q5k_or_hibit
|
||||
static inline HVX_VectorPair unpack_and_interleave_5bit_x2(HVX_Vector v_src, HVX_Vector v_plane, int i, HVX_Vector mask_h4) {
|
||||
HVX_Vector v_lo = hvx_q5k_or_hibit(Q6_V_vand_VV(v_src, mask_h4), v_plane, 2 * i);
|
||||
HVX_Vector v_hi = hvx_q5k_or_hibit(Q6_Vub_vlsr_VubR(v_src, 4), v_plane, 2 * i + 1);
|
||||
HVX_VectorPair v01_pair = Q6_W_vshuff_VVR(v_hi, v_lo, -1);
|
||||
HVX_Vector v01_lo = Q6_V_lo_W(v01_pair);
|
||||
HVX_Vector v01_hi = Q6_V_hi_W(v01_pair);
|
||||
|
||||
HVX_Vector v23_lo = Q6_V_valign_VVR(v01_hi, v01_lo, 64);
|
||||
HVX_Vector v_W0 = Q6_V_lo_W(Q6_W_vshuff_VVR(v23_lo, v01_lo, -2));
|
||||
|
||||
HVX_Vector v67_lo = Q6_V_valign_VVR(v01_lo, v01_hi, 64);
|
||||
HVX_Vector v_W1 = Q6_V_lo_W(Q6_W_vshuff_VVR(v67_lo, v01_hi, -2));
|
||||
|
||||
return Q6_W_vcombine_VV(v_W1, v_W0);
|
||||
}
|
||||
|
||||
static inline HVX_Vector accum_5bit_32x1(
|
||||
const HVX_Vector * restrict vptr,
|
||||
const HVX_Vector * restrict v_act,
|
||||
HVX_Vector i8
|
||||
) {
|
||||
HVX_Vector v_sum0 = Q6_V_vzero();
|
||||
HVX_Vector v_sum1 = Q6_V_vzero();
|
||||
HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
|
||||
HVX_Vector v_plane = vptr[5];
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; i++) {
|
||||
HVX_VectorPair v_W_pair = unpack_and_interleave_5bit_x2(vptr[i], v_plane, i, mask_h4);
|
||||
HVX_Vector v_W0 = Q6_Vb_vsub_VbVb(Q6_V_lo_W(v_W_pair), i8);
|
||||
HVX_Vector v_W1 = Q6_Vb_vsub_VbVb(Q6_V_hi_W(v_W_pair), i8);
|
||||
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W0, v_act[i * 2 + 0]);
|
||||
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W1, v_act[i * 2 + 1]);
|
||||
}
|
||||
|
||||
return Q6_Vw_vadd_VwVw(v_sum0, v_sum1);
|
||||
}
|
||||
|
||||
static inline HVX_VectorPair accum_5bit_32x2(
|
||||
const HVX_Vector * restrict vptr,
|
||||
const HVX_Vector * restrict v_act0,
|
||||
const HVX_Vector * restrict v_act1,
|
||||
HVX_Vector i8
|
||||
) {
|
||||
HVX_Vector v_sum0 = Q6_V_vzero();
|
||||
HVX_Vector v_sum1 = Q6_V_vzero();
|
||||
HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
|
||||
HVX_Vector v_plane = vptr[5];
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; i++) {
|
||||
HVX_VectorPair v_W_pair = unpack_and_interleave_5bit_x2(vptr[i], v_plane, i, mask_h4);
|
||||
HVX_Vector v_W0 = Q6_Vb_vsub_VbVb(Q6_V_lo_W(v_W_pair), i8);
|
||||
HVX_Vector v_W1 = Q6_Vb_vsub_VbVb(Q6_V_hi_W(v_W_pair), i8);
|
||||
|
||||
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W0, v_act0[i * 2 + 0]);
|
||||
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W1, v_act0[i * 2 + 1]);
|
||||
|
||||
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W0, v_act1[i * 2 + 0]);
|
||||
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W1, v_act1[i * 2 + 1]);
|
||||
}
|
||||
|
||||
return Q6_W_vcombine_VV(v_sum1, v_sum0);
|
||||
}
|
||||
|
||||
// Q6_K weights are stored unsigned (0..63), see HTP_MM_WEIGHT_TILE_SIZE_Q6_K. Unpack k-group g of a tile to signed bytes (q - 32)
|
||||
static inline HVX_Vector unpack_q6_k_group(const HVX_Vector * restrict vptr, int g, HVX_Vector mask_0f, HVX_Vector mask_03, HVX_Vector i32) {
|
||||
HVX_Vector v_lo = (g & 1) ? Q6_Vub_vlsr_VubR(vptr[g >> 1], 4) : Q6_V_vand_VV(vptr[g >> 1], mask_0f);
|
||||
@@ -689,6 +786,102 @@ static void tiled_vec_dot_q8_0_32x2(const uint32_t n, float * restrict s0, float
|
||||
}
|
||||
}
|
||||
|
||||
static void tiled_vec_dot_q5_k_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
|
||||
const uint8_t * restrict tile_ptr = vx;
|
||||
const uint8_t * restrict y_q = vy;
|
||||
|
||||
HVX_Vector v_sum_float = Q6_V_vzero();
|
||||
|
||||
uint32_t n_k_tiles = n / 32;
|
||||
for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
|
||||
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 768);
|
||||
const HVX_Vector * restrict v_act = (const HVX_Vector *) (y_q + kt * 1280);
|
||||
|
||||
HVX_Vector v_sum = accum_5bit_32x1(vptr, v_act, Q6_V_vzero());
|
||||
HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
|
||||
|
||||
HVX_Vector v_scale_offset = vptr[4];
|
||||
HVX_VectorPair p_deal = Q6_W_vdeal_VVR(v_scale_offset, v_scale_offset, -2);
|
||||
HVX_Vector v_scale = Q6_V_lo_W(p_deal);
|
||||
HVX_Vector v_offset = Q6_V_hi_W(p_deal);
|
||||
|
||||
HVX_Vector v_scale_a = v_act[8];
|
||||
HVX_Vector v_sum_a = v_act[9];
|
||||
|
||||
HVX_Vector v_scale_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a);
|
||||
HVX_Vector v_offset_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a);
|
||||
|
||||
HVX_Vector v_scaled_dot = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
|
||||
HVX_Vector v_sum_scaled = hvx_vec_add_f32_f32(v_scaled_dot, v_offset_comb);
|
||||
|
||||
v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
|
||||
}
|
||||
|
||||
if (sz) {
|
||||
hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
|
||||
} else {
|
||||
hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
|
||||
}
|
||||
}
|
||||
|
||||
static void tiled_vec_dot_q5_k_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
|
||||
const uint8_t * restrict tile_ptr = vx;
|
||||
const uint8_t * restrict y0_q = vy0;
|
||||
const uint8_t * restrict y1_q = vy1;
|
||||
|
||||
HVX_Vector v_sum_float_c0 = Q6_V_vzero();
|
||||
HVX_Vector v_sum_float_c1 = Q6_V_vzero();
|
||||
|
||||
uint32_t n_k_tiles = n / 32;
|
||||
for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
|
||||
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 768);
|
||||
const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1280);
|
||||
const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1280);
|
||||
|
||||
HVX_VectorPair v_sums = accum_5bit_32x2(vptr, v_act0, v_act1, Q6_V_vzero());
|
||||
HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
|
||||
HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
|
||||
|
||||
HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
|
||||
HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
|
||||
|
||||
HVX_Vector v_scale_offset = vptr[4];
|
||||
HVX_VectorPair p_deal = Q6_W_vdeal_VVR(v_scale_offset, v_scale_offset, -2);
|
||||
HVX_Vector v_scale = Q6_V_lo_W(p_deal);
|
||||
HVX_Vector v_offset = Q6_V_hi_W(p_deal);
|
||||
|
||||
HVX_Vector v_scale_a_c0 = v_act0[8];
|
||||
HVX_Vector v_sum_a_c0 = v_act0[9];
|
||||
HVX_Vector v_scale_a_c1 = v_act1[8];
|
||||
HVX_Vector v_sum_a_c1 = v_act1[9];
|
||||
|
||||
HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a_c0);
|
||||
HVX_Vector v_offset_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a_c0);
|
||||
HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a_c1);
|
||||
HVX_Vector v_offset_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a_c1);
|
||||
|
||||
HVX_Vector v_scaled_dot_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
|
||||
HVX_Vector v_sum_scaled_c0 = hvx_vec_add_f32_f32(v_scaled_dot_c0, v_offset_comb_c0);
|
||||
|
||||
HVX_Vector v_scaled_dot_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
|
||||
HVX_Vector v_sum_scaled_c1 = hvx_vec_add_f32_f32(v_scaled_dot_c1, v_offset_comb_c1);
|
||||
|
||||
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
|
||||
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
|
||||
}
|
||||
|
||||
if (sz0) {
|
||||
hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
|
||||
} else {
|
||||
hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
|
||||
}
|
||||
if (sz1) {
|
||||
hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
|
||||
} else {
|
||||
hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
|
||||
}
|
||||
}
|
||||
|
||||
static void tiled_vec_dot_q6_k_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
|
||||
const uint8_t * restrict tile_ptr = vx;
|
||||
const uint8_t * restrict y_q = vy;
|
||||
@@ -941,14 +1134,9 @@ static inline void quantize_f32_q8_0_tiled_kernel(
|
||||
size_t src_row_size,
|
||||
size_t dst_row_size
|
||||
) {
|
||||
const size_t src_row_size_padded = hex_round_up(src_row_size, QK_Q8_0_TILED * sizeof(float));
|
||||
hvx_splat_f32_a(tmp_data, 0.0f, src_row_size_padded / sizeof(float));
|
||||
|
||||
(void) tmp_data;
|
||||
for (uint32_t i = 0; i < nrows; ++i) {
|
||||
hex_l2fetch(src_data, src_row_size, src_row_size, 2);
|
||||
hvx_copy_f32_aa(tmp_data, src_data, ne0);
|
||||
|
||||
quantize_row_f32_q8_0_tiled((float *) tmp_data, dst_data, ne0);
|
||||
quantize_row_f32_q8_0_tiled((float *) src_data, dst_data, ne0);
|
||||
dst_data += dst_row_size;
|
||||
src_data += src_row_size;
|
||||
}
|
||||
@@ -963,14 +1151,9 @@ static inline void quantize_f32_q8_1_tiled_kernel(
|
||||
size_t src_row_size,
|
||||
size_t dst_row_size
|
||||
) {
|
||||
const size_t src_row_size_padded = hex_round_up(src_row_size, QK_Q8_0_TILED * sizeof(float));
|
||||
hvx_splat_f32_a(tmp_data, 0.0f, src_row_size_padded / sizeof(float));
|
||||
|
||||
(void) tmp_data;
|
||||
for (uint32_t i = 0; i < nrows; ++i) {
|
||||
hex_l2fetch(src_data, src_row_size, src_row_size, 2);
|
||||
hvx_copy_f32_aa(tmp_data, src_data, ne0);
|
||||
|
||||
quantize_row_f32_q8_1_tiled((float *) tmp_data, dst_data, ne0);
|
||||
quantize_row_f32_q8_1_tiled((float *) src_data, dst_data, ne0);
|
||||
dst_data += dst_row_size;
|
||||
src_data += src_row_size;
|
||||
}
|
||||
@@ -988,24 +1171,15 @@ static inline void quantize_f32_q8_0_tiled_block_kernel(
|
||||
uint32_t r,
|
||||
uint32_t c
|
||||
) {
|
||||
(void) tmp_data;
|
||||
const uint32_t qk = QK_Q8_0_TILED;
|
||||
const uint32_t nb = (ne0 + qk - 1) / qk;
|
||||
|
||||
for (uint32_t ib = ib_first; ib < ib_last; ++ib) {
|
||||
const uint8_t * restrict src_ptr = (const uint8_t *) src + r * src_row_size + c * qk * sizeof(float);
|
||||
const float * restrict src_ptr = (const float *) ((const uint8_t *) src + r * src_row_size + c * qk * sizeof(float));
|
||||
uint8_t * restrict dst_ptr = dst + r * dst_row_size + c * 4 * 1152;
|
||||
|
||||
hex_l2fetch(src_ptr, qk * sizeof(float), qk * sizeof(float), 1);
|
||||
|
||||
if (c == nb - 1) {
|
||||
uint32_t active_elements = ne0 - c * qk;
|
||||
hvx_splat_f32_a(tmp_data, 0.0f, qk);
|
||||
hvx_copy_f32_aa(tmp_data, src_ptr, active_elements);
|
||||
} else {
|
||||
hvx_copy_f32_aa(tmp_data, src_ptr, qk);
|
||||
}
|
||||
|
||||
quantize_block_f32_q8_0_tiled((float *) tmp_data, dst_ptr);
|
||||
quantize_block_f32_q8_0_tiled((float *) src_ptr, dst_ptr);
|
||||
|
||||
c++;
|
||||
if (c == nb) {
|
||||
@@ -1027,24 +1201,15 @@ static inline void quantize_f32_q8_1_tiled_block_kernel(
|
||||
uint32_t r,
|
||||
uint32_t c
|
||||
) {
|
||||
(void) tmp_data;
|
||||
const uint32_t qk = QK_Q8_0_TILED;
|
||||
const uint32_t nb = (ne0 + qk - 1) / qk;
|
||||
|
||||
for (uint32_t ib = ib_first; ib < ib_last; ++ib) {
|
||||
const uint8_t * restrict src_ptr = (const uint8_t *) src + r * src_row_size + c * qk * sizeof(float);
|
||||
const float * restrict src_ptr = (const float *) ((const uint8_t *) src + r * src_row_size + c * qk * sizeof(float));
|
||||
uint8_t * restrict dst_ptr = dst + r * dst_row_size + c * 4 * 1280;
|
||||
|
||||
hex_l2fetch(src_ptr, qk * sizeof(float), qk * sizeof(float), 1);
|
||||
|
||||
if (c == nb - 1) {
|
||||
uint32_t active_elements = ne0 - c * qk;
|
||||
hvx_splat_f32_a(tmp_data, 0.0f, qk);
|
||||
hvx_copy_f32_aa(tmp_data, src_ptr, active_elements);
|
||||
} else {
|
||||
hvx_copy_f32_aa(tmp_data, src_ptr, qk);
|
||||
}
|
||||
|
||||
quantize_block_f32_q8_1_tiled((float *) tmp_data, dst_ptr);
|
||||
quantize_block_f32_q8_1_tiled((float *) src_ptr, dst_ptr);
|
||||
|
||||
c++;
|
||||
if (c == nb) {
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -25,6 +25,9 @@ extern "C" {
|
||||
#define HTP_MM_WEIGHT_TILE_SIZE_Q8_0 1088
|
||||
#define HTP_MM_WEIGHT_TILE_SIZE_IQ4_NL 576
|
||||
#define HTP_MM_WEIGHT_TILE_SIZE_MXFP4 544
|
||||
// Q5_K: the Q4_1 tile (640) followed by a 128-byte plane with the 5th bit of every quant, transposed so that
|
||||
// plane byte l holds the eight flags of lane l: bit 2i = low nibble of nibble vector i, bit 2i+1 = high nibble
|
||||
#define HTP_MM_WEIGHT_TILE_SIZE_Q5_K 768
|
||||
// Q6_K native 6-bit tile (32 rows x 32 k), vrmpy-ready: byte 4*row+b of a vector holds k = 4*group+b
|
||||
// vectors 0..3: low nibbles, vector i holds group 2i (low nibble) and group 2i+1 (high nibble)
|
||||
// vectors 4..5: high 2 bits, vector m holds groups 4m..4m+3 at bit offsets 0,2,4,6
|
||||
@@ -37,6 +40,7 @@ extern "C" {
|
||||
#define HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q8_0 1152
|
||||
#define HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_IQ4_NL 640
|
||||
#define HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_MXFP4 640
|
||||
#define HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q5_K 768
|
||||
#define HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q6_K 896
|
||||
|
||||
// --- Activation Tiled Block Sizes (including padding) ---
|
||||
@@ -199,6 +203,8 @@ static inline uint32_t htp_mm_get_weight_tile_size(int weight_type) {
|
||||
return HTP_MM_WEIGHT_TILE_SIZE_Q4_1;
|
||||
case HTP_TYPE_Q8_0:
|
||||
return HTP_MM_WEIGHT_TILE_SIZE_Q8_0;
|
||||
case HTP_TYPE_Q5_K:
|
||||
return HTP_MM_WEIGHT_TILE_SIZE_Q5_K;
|
||||
case HTP_TYPE_Q6_K:
|
||||
return HTP_MM_WEIGHT_TILE_SIZE_Q6_K;
|
||||
case HTP_TYPE_MXFP4:
|
||||
@@ -218,6 +224,8 @@ static inline uint32_t htp_mm_get_weight_aligned_tile_size(int weight_type) {
|
||||
return HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q4_1;
|
||||
case HTP_TYPE_Q8_0:
|
||||
return HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q8_0;
|
||||
case HTP_TYPE_Q5_K:
|
||||
return HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q5_K;
|
||||
case HTP_TYPE_Q6_K:
|
||||
return HTP_MM_WEIGHT_ALIGNED_TILE_SIZE_Q6_K;
|
||||
case HTP_TYPE_MXFP4:
|
||||
@@ -227,6 +235,11 @@ static inline uint32_t htp_mm_get_weight_aligned_tile_size(int weight_type) {
|
||||
}
|
||||
}
|
||||
|
||||
// weight types whose tiles carry a per-block offset (x = d * q + m): the activations need block sums (q8_1)
|
||||
static inline bool htp_mm_weight_has_offset(int weight_type) {
|
||||
return weight_type == HTP_TYPE_Q4_1 || weight_type == HTP_TYPE_Q4_K || weight_type == HTP_TYPE_Q5_K;
|
||||
}
|
||||
|
||||
// --- Activation/Row Size Helpers ---
|
||||
static inline size_t htp_mm_q8_0_tiled_row_size(uint32_t ne) {
|
||||
const uint32_t ne_padded = ((ne + 127) / 128) * 128;
|
||||
@@ -248,6 +261,7 @@ static inline size_t htp_mm_get_tiled_row_stride(int weight_type, uint32_t k) {
|
||||
case HTP_TYPE_Q4_1:
|
||||
case HTP_TYPE_Q4_K:
|
||||
case HTP_TYPE_Q8_0:
|
||||
case HTP_TYPE_Q5_K:
|
||||
case HTP_TYPE_Q6_K:
|
||||
case HTP_TYPE_MXFP4:
|
||||
return (size_t) nb * htp_mm_get_weight_tile_size(weight_type);
|
||||
@@ -332,6 +346,7 @@ struct htp_mm_hvx_vtcm_layout {
|
||||
size_t off_src2; // vtcm_src2 (Wq / fused only)
|
||||
size_t off_src3; // vtcm_src3 (Wv / fused only)
|
||||
size_t off_dst; // vtcm_dst (output scratch)
|
||||
size_t off_act_raw; // vtcm_act_raw (raw activation DMA staging)
|
||||
|
||||
// Cached sizes
|
||||
size_t src0_bytes;
|
||||
@@ -339,6 +354,7 @@ struct htp_mm_hvx_vtcm_layout {
|
||||
size_t src2_bytes;
|
||||
size_t src3_bytes;
|
||||
size_t dst_bytes;
|
||||
size_t act_raw_bytes;
|
||||
|
||||
size_t total_bytes;
|
||||
};
|
||||
@@ -351,7 +367,6 @@ static inline void htp_mm_hmx_vtcm_layout_build(
|
||||
size_t mc,
|
||||
size_t nc,
|
||||
uint32_t group_size,
|
||||
bool use_dma_activation,
|
||||
bool pipeline,
|
||||
uint32_t act_threads,
|
||||
uint32_t aligned_tile_size,
|
||||
@@ -366,8 +381,7 @@ static inline void htp_mm_hmx_vtcm_layout_build(
|
||||
const size_t activation_area_size = hex_align_up(group_size * act_head_stride * sizeof(uint16_t), HTP_MM_HMX_TILE_SIZE);
|
||||
const size_t output_area_size = hex_align_up(group_size * mc * nc * sizeof(uint16_t), HTP_MM_HMX_TILE_SIZE);
|
||||
const size_t scratch_area_size = hex_align_up(nc * vec_dot_size, HTP_MM_HMX_TILE_SIZE);
|
||||
const size_t min_f32_size = use_dma_activation
|
||||
? hex_align_up(act_threads * HTP_MM_DMA_ACT_MULTIPLIER * k * sizeof(float), 128) : 0;
|
||||
const size_t min_f32_size = hex_align_up(act_threads * HTP_MM_DMA_ACT_MULTIPLIER * k * sizeof(float), 128);
|
||||
|
||||
// Group A: Permanent activation tiles and scales
|
||||
size_t off_group_a = 0;
|
||||
@@ -388,10 +402,9 @@ static inline void htp_mm_hmx_vtcm_layout_build(
|
||||
|
||||
// Group C: Activation prep temporary buffer (overlaps Group B, starting at off_group_a)
|
||||
const size_t max_f32_size = act_threads * 64 * k * sizeof(float);
|
||||
const size_t act_f32_size = use_dma_activation
|
||||
? hex_align_up(hex_smin(max_f32_size, hex_smax(min_f32_size, group_b_size)), 128) : 0;
|
||||
const size_t act_f32_size = hex_align_up(hex_smin(max_f32_size, hex_smax(min_f32_size, group_b_size)), 128);
|
||||
size_t off_group_c = off_group_a;
|
||||
VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_c, off_act_f32, act_f32_size, use_dma_activation);
|
||||
VTCM_LAYOUT_ALLOC(off_group_c, off_act_f32, act_f32_size);
|
||||
|
||||
const size_t group_c_size = off_group_c - off_group_a;
|
||||
|
||||
@@ -478,20 +491,20 @@ static inline void htp_mm_hvx_vtcm_layout_build(
|
||||
bool is_fused_nx
|
||||
) {
|
||||
(void)src1_row_size;
|
||||
size_t src0_sz = 0;
|
||||
size_t src1_sz = 0;
|
||||
size_t src2_sz = src2_row_size > 0 ? htp_mm_round_up(src2_row_size, 128) : 0;
|
||||
size_t src3_sz = 0;
|
||||
size_t dst_sz = 0;
|
||||
size_t src0_sz = 0;
|
||||
size_t src1_sz = 0;
|
||||
size_t src2_sz = src2_row_size > 0 ? htp_mm_round_up(src2_row_size, 128) : 0;
|
||||
size_t src3_sz = 0;
|
||||
size_t dst_sz = 0;
|
||||
size_t act_raw_sz = 0;
|
||||
|
||||
const bool is_repack = (wtype == HTP_TYPE_Q4_0 || wtype == HTP_TYPE_Q4_1 ||
|
||||
wtype == HTP_TYPE_Q8_0 || wtype == HTP_TYPE_IQ4_NL ||
|
||||
wtype == HTP_TYPE_MXFP4 || wtype == HTP_TYPE_Q6_K ||
|
||||
wtype == HTP_TYPE_Q4_K);
|
||||
wtype == HTP_TYPE_Q4_K || wtype == HTP_TYPE_Q5_K);
|
||||
|
||||
if (is_fused_nx) {
|
||||
const size_t src0_row_size_padded = hex_round_up(src0_row_size, 128);
|
||||
const size_t quant_scratch_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float)) * n_threads;
|
||||
|
||||
size_t weight_sz_per_thread = 0;
|
||||
|
||||
@@ -505,17 +518,19 @@ static inline void htp_mm_hvx_vtcm_layout_build(
|
||||
weight_sz_per_thread = hex_round_up(n_prefetch * src0_row_size_padded, 128);
|
||||
}
|
||||
|
||||
size_t tiled_act_row_size = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
|
||||
size_t tiled_act_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
|
||||
size_t act_sz = hex_round_up(tiled_act_row_size * src1_nrows, 128);
|
||||
size_t raw_row_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
|
||||
|
||||
src0_sz = weight_sz_per_thread * n_threads; // shared single-weight prefetch buffer
|
||||
src1_sz = act_sz; // quantized activation buffer
|
||||
src2_sz = 0;
|
||||
src3_sz = 0;
|
||||
dst_sz = quant_scratch_size;
|
||||
src0_sz = weight_sz_per_thread * n_threads; // shared single-weight prefetch buffer
|
||||
src1_sz = act_sz; // quantized activation buffer
|
||||
src2_sz = 0;
|
||||
src3_sz = 0;
|
||||
dst_sz = 0;
|
||||
act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
|
||||
} else if (is_matmul_id) {
|
||||
const size_t src0_row_size_padded = htp_mm_round_up(src0_row_size, 128);
|
||||
const size_t src1_row_size_tiled = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(ne10)
|
||||
const size_t src1_row_size_tiled = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10)
|
||||
: htp_mm_q8_0_tiled_row_size(ne10);
|
||||
|
||||
size_t src0_sz_per_thread = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256);
|
||||
@@ -529,10 +544,13 @@ static inline void htp_mm_hvx_vtcm_layout_build(
|
||||
src0_sz_per_thread = repacked_vtcm_size;
|
||||
}
|
||||
|
||||
src0_sz = src0_sz_per_thread * n_threads;
|
||||
dst_sz = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float)) * n_threads;
|
||||
src2_sz = 0;
|
||||
src3_sz = 0;
|
||||
size_t raw_row_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
|
||||
|
||||
src0_sz = src0_sz_per_thread * n_threads;
|
||||
dst_sz = 0;
|
||||
src2_sz = 0;
|
||||
src3_sz = 0;
|
||||
act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
|
||||
} else {
|
||||
const size_t src0_row_size_padded = htp_mm_round_up(src0_row_size, 128);
|
||||
const size_t dst_nrows = (src1_nrows > 1) ? 0 : 1;
|
||||
@@ -540,21 +558,23 @@ static inline void htp_mm_hvx_vtcm_layout_build(
|
||||
switch (kernel_type) {
|
||||
case HTP_MM_KERNEL_HVX_F16_F16_VTCM: {
|
||||
size_t f16_src1_row_size = htp_mm_round_up(ne10 * 2, 128);
|
||||
src1_sz = htp_mm_round_up(f16_src1_row_size * src1_nrows, 256);
|
||||
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
|
||||
dst_sz = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
|
||||
src1_sz = htp_mm_round_up(f16_src1_row_size * src1_nrows, 256);
|
||||
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
|
||||
dst_sz = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
|
||||
act_raw_sz = hex_round_up(hex_round_up(ne10 * sizeof(float), 128) * src1_nrows, 128);
|
||||
break;
|
||||
}
|
||||
case HTP_MM_KERNEL_HVX_F32_F32_VTCM: {
|
||||
size_t f32_src1_row_size = htp_mm_round_up(ne10 * 4, 128);
|
||||
src1_sz = htp_mm_round_up(f32_src1_row_size * src1_nrows, 256);
|
||||
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
|
||||
dst_sz = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
|
||||
src1_sz = htp_mm_round_up(f32_src1_row_size * src1_nrows, 256);
|
||||
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256) * n_threads;
|
||||
dst_sz = dst_nrows > 0 ? htp_mm_round_up(dst_row_size, 128) * n_threads : 0;
|
||||
act_raw_sz = 0;
|
||||
break;
|
||||
}
|
||||
case HTP_MM_KERNEL_HVX_QUANT_BLOCK:
|
||||
case HTP_MM_KERNEL_HVX_QUANT_ROW: {
|
||||
size_t q_src1_row_size = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
|
||||
size_t q_src1_row_size = htp_mm_weight_has_offset(wtype) ? htp_mm_q8_1_tiled_row_size(ne10) : htp_mm_q8_0_tiled_row_size(ne10);
|
||||
|
||||
src0_sz = htp_mm_round_up(n_prefetch * src0_row_size_padded, 256);
|
||||
src1_sz = htp_mm_round_up(q_src1_row_size * src1_nrows, 256);
|
||||
@@ -569,10 +589,10 @@ static inline void htp_mm_hvx_vtcm_layout_build(
|
||||
src0_sz = repacked_vtcm_size * n_threads;
|
||||
}
|
||||
|
||||
size_t quant_scratch_size_per_thread = htp_mm_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
|
||||
size_t dst_slice_per_thread = (dst_nrows > 0 && src1_nrows == 1) ? htp_mm_round_up((dst_row_size + n_threads - 1) / n_threads, 128) : 0;
|
||||
size_t dst_size_per_thread = (dst_slice_per_thread > quant_scratch_size_per_thread) ? dst_slice_per_thread : quant_scratch_size_per_thread;
|
||||
dst_sz = dst_size_per_thread * n_threads;
|
||||
dst_sz = dst_slice_per_thread * n_threads;
|
||||
size_t raw_row_size = hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
|
||||
act_raw_sz = hex_round_up(raw_row_size * src1_nrows, 128);
|
||||
break;
|
||||
}
|
||||
default:
|
||||
@@ -580,19 +600,30 @@ static inline void htp_mm_hvx_vtcm_layout_build(
|
||||
}
|
||||
}
|
||||
|
||||
size_t off = 0;
|
||||
VTCM_LAYOUT_ALLOC(off, off_src0, src0_sz);
|
||||
VTCM_LAYOUT_ALLOC(off, off_src1, src1_sz);
|
||||
VTCM_LAYOUT_ALLOC(off, off_src2, src2_sz);
|
||||
VTCM_LAYOUT_ALLOC(off, off_src3, src3_sz);
|
||||
VTCM_LAYOUT_ALLOC(off, off_dst, dst_sz);
|
||||
// Group A: Persistent buffers across chunk compute
|
||||
size_t off_group_a = 0;
|
||||
VTCM_LAYOUT_ALLOC(off_group_a, off_src1, src1_sz);
|
||||
VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src2, src2_sz, src2_sz > 0);
|
||||
VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_a, off_src3, src3_sz, src3_sz > 0);
|
||||
|
||||
L->src0_bytes = src0_sz;
|
||||
L->src1_bytes = src1_sz;
|
||||
L->src2_bytes = src2_sz;
|
||||
L->src3_bytes = src3_sz;
|
||||
L->dst_bytes = dst_sz;
|
||||
L->total_bytes = off;
|
||||
// Group B: Compute-only buffers (starts at off_group_a)
|
||||
size_t off_group_b = off_group_a;
|
||||
VTCM_LAYOUT_ALLOC(off_group_b, off_src0, src0_sz);
|
||||
VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_b, off_dst, dst_sz, dst_sz > 0);
|
||||
const size_t group_b_size = off_group_b - off_group_a;
|
||||
|
||||
// Group C: Raw activation staging buffer (overlaps Group B, starts at off_group_a)
|
||||
size_t off_group_c = off_group_a;
|
||||
VTCM_LAYOUT_ALLOC_OPTIONAL(off_group_c, off_act_raw, act_raw_sz, act_raw_sz > 0);
|
||||
const size_t group_c_size = off_group_c - off_group_a;
|
||||
|
||||
L->src0_bytes = src0_sz;
|
||||
L->src1_bytes = src1_sz;
|
||||
L->src2_bytes = src2_sz;
|
||||
L->src3_bytes = src3_sz;
|
||||
L->dst_bytes = dst_sz;
|
||||
L->act_raw_bytes = act_raw_sz;
|
||||
L->total_bytes = off_group_a + hex_smax(group_b_size, group_c_size);
|
||||
}
|
||||
|
||||
static inline bool htp_mm_hvx_solve_vtcm_params(
|
||||
@@ -630,7 +661,7 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
|
||||
const size_t avail_act = vtcm_budget - fixed_bytes;
|
||||
size_t row_size = 0;
|
||||
if (kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW || kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) {
|
||||
row_size = (wtype == HTP_TYPE_Q4_1 || wtype == HTP_TYPE_Q4_K)
|
||||
row_size = htp_mm_weight_has_offset(wtype)
|
||||
? htp_mm_q8_1_tiled_row_size(ne10)
|
||||
: htp_mm_q8_0_tiled_row_size(ne10);
|
||||
} else if (kernel_type == HTP_MM_KERNEL_HVX_F16_F16_VTCM) {
|
||||
@@ -642,7 +673,12 @@ static inline bool htp_mm_hvx_solve_vtcm_params(
|
||||
return false;
|
||||
}
|
||||
|
||||
uint32_t m_chunk = (uint32_t) (avail_act / row_size);
|
||||
size_t eff_row_size = row_size;
|
||||
if (kernel_type == HTP_MM_KERNEL_HVX_QUANT_ROW || kernel_type == HTP_MM_KERNEL_HVX_QUANT_BLOCK) {
|
||||
eff_row_size += hex_round_up(ne10 * sizeof(float), QK_Q8_0_TILED * sizeof(float));
|
||||
}
|
||||
|
||||
uint32_t m_chunk = (uint32_t) (avail_act / eff_row_size);
|
||||
if (m_chunk > 1) {
|
||||
m_chunk &= ~1U;
|
||||
}
|
||||
@@ -679,15 +715,15 @@ static inline size_t htp_mm_hmx_get_2d_vtcm_size(
|
||||
int wtype, uint32_t k, size_t mc, size_t nc, bool pipeline, uint32_t act_threads, uint32_t aligned_tile_size, size_t src2_size
|
||||
) {
|
||||
struct htp_mm_hmx_vtcm_layout L;
|
||||
htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, wtype, k, mc, nc, 1, false, pipeline, act_threads, aligned_tile_size, src2_size);
|
||||
htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_2D, wtype, k, mc, nc, 1, pipeline, act_threads, aligned_tile_size, src2_size);
|
||||
return L.total_bytes;
|
||||
}
|
||||
|
||||
static inline size_t htp_mm_hmx_get_batched_vtcm_size(
|
||||
int wtype, uint32_t k, size_t mc, size_t nc, uint32_t group_size, bool use_dma_activation, bool pipeline, uint32_t act_threads, size_t src2_size) {
|
||||
int wtype, uint32_t k, size_t mc, size_t nc, uint32_t group_size, bool pipeline, uint32_t act_threads, size_t src2_size) {
|
||||
(void)pipeline;
|
||||
struct htp_mm_hmx_vtcm_layout L;
|
||||
htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, wtype, k, mc, nc, group_size, use_dma_activation, false, act_threads, 0, src2_size);
|
||||
htp_mm_hmx_vtcm_layout_build(&L, HTP_MM_KERNEL_HMX_F16_BATCHED, wtype, k, mc, nc, group_size, false, act_threads, 0, src2_size);
|
||||
return L.total_bytes;
|
||||
}
|
||||
|
||||
@@ -697,7 +733,6 @@ static inline bool htp_mm_hmx_solve_batched_params(
|
||||
uint32_t ne01_padded,
|
||||
uint32_t ne11,
|
||||
uint32_t group_size,
|
||||
bool use_dma_activation,
|
||||
int n_threads,
|
||||
bool pipeline,
|
||||
size_t src2_size,
|
||||
@@ -726,7 +761,7 @@ static inline bool htp_mm_hmx_solve_batched_params(
|
||||
if (htp_mm_hmx_compute_chunks(vtcm_budget, group_overhead, group_size_per_n, group_size_per_m, group_size_per_mn, hex_align_up(ne11, 32), ne01_padded,
|
||||
(size_t) ne01_padded * HTP_MM_HMX_COST_W_DEQUANT, (size_t) ne11 * HTP_MM_HMX_COST_A_CONVERT,
|
||||
&m_chunk_candidate, &n_chunk_candidate, &vtcm_size_candidate) == 0) {
|
||||
size_t exact_size = htp_mm_hmx_get_batched_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, group_size, use_dma_activation, pipeline, act_threads, src2_size);
|
||||
size_t exact_size = htp_mm_hmx_get_batched_vtcm_size(wtype, k, m_chunk_candidate, n_chunk_candidate, group_size, pipeline, act_threads, src2_size);
|
||||
if (exact_size <= vtcm_budget) {
|
||||
size_t mblocks = ((size_t) ne11 + m_chunk_candidate - 1) / m_chunk_candidate;
|
||||
if (mblocks < best_mblocks || (mblocks == best_mblocks && act_threads > best_act_threads)) {
|
||||
|
||||
@@ -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;
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user