mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-01 05:56:52 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c85b92c69c | ||
|
|
31385c9ceb | ||
|
|
8019dc563b | ||
|
|
86ea01d05e | ||
|
|
18b74ff68e | ||
|
|
c13e04e1dd | ||
|
|
c8cda8b4fe | ||
|
|
6d78fb0727 | ||
|
|
18bbc46b46 | ||
|
|
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 |
@@ -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
|
||||
|
||||
+10
-5
@@ -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/musa_sdk:5.2.0-devel-ubuntu${UBUNTU_VERSION}-s5000
|
||||
|
||||
ARG BASE_MUSA_RUN_CONTAINER=docker.io/mthreads/musa:${MUSA_VERSION}-runtime-ubuntu${UBUNTU_VERSION}-amd64
|
||||
ARG BASE_MUSA_RUN_CONTAINER=registry.mthreads.com/mcconline/musa_sdk:5.2.0-runtime-ubuntu${UBUNTU_VERSION}-s5000
|
||||
|
||||
ARG BUILD_DATE=N/A
|
||||
ARG APP_VERSION=N/A
|
||||
@@ -37,7 +36,10 @@ RUN apt-get update && \
|
||||
python3-pip \
|
||||
git \
|
||||
libssl-dev \
|
||||
libgomp1
|
||||
libgomp1 \
|
||||
musa-mualg-5-2 \
|
||||
musa-muthrust-5-2 \
|
||||
libmthreads-compute
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
@@ -80,13 +82,16 @@ LABEL org.opencontainers.image.created=$BUILD_DATE \
|
||||
org.opencontainers.image.source=$IMAGE_SOURCE
|
||||
|
||||
RUN apt-get update \
|
||||
&& apt-get install -y libgomp1 curl ffmpeg \
|
||||
&& apt-get install -y libgomp1 curl ffmpeg libmthreads-compute \
|
||||
&& apt autoremove -y \
|
||||
&& apt clean -y \
|
||||
&& rm -rf /tmp/* /var/tmp/* \
|
||||
&& find /var/cache/apt/archives /var/lib/apt/lists -not -name lock -type f -delete \
|
||||
&& find /var/cache -type f -delete
|
||||
|
||||
# The MUSA runtime image does not register its library directory
|
||||
RUN echo "/usr/local/musa/lib" > /etc/ld.so.conf.d/musa-runtime.conf && ldconfig
|
||||
|
||||
COPY --from=build /app/lib/ /app
|
||||
|
||||
### Full
|
||||
|
||||
@@ -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: registry.mthreads.com/mcconline/inference/pytorch:2.9.1.post1-py3.10-musa5.2.0-mp31-devel-ubuntu22.04-amd64
|
||||
container: registry.mthreads.com/mcconline/musa_sdk:5.2.0-devel-ubuntu22.04-s5000
|
||||
|
||||
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 python3-venv
|
||||
apt-get install -y build-essential git cmake libssl-dev jq python3-venv musa-mualg-5-2 musa-muthrust-5-2 libmthreads-compute
|
||||
|
||||
- name: ccache
|
||||
uses: ggml-org/[email protected]
|
||||
|
||||
@@ -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
|
||||
|
||||
+2
-2
@@ -21,13 +21,13 @@ 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 \
|
||||
registry.mthreads.com/mcconline/inference/pytorch:2.9.1.post1-py3.10-musa5.2.0-mp31-devel-ubuntu22.04-amd64
|
||||
registry.mthreads.com/mcconline/musa_sdk:5.2.0-devel-ubuntu22.04-s5000
|
||||
```
|
||||
|
||||
Inside the container, execute the following commands:
|
||||
|
||||
```bash
|
||||
apt update -y && apt install -y bc cmake ccache git python3.10-venv time unzip wget
|
||||
apt update -y && apt install -y bc cmake ccache git python3.10-venv time unzip wget musa-mualg-5-2 musa-muthrust-5-2 libmthreads-compute
|
||||
git config --global --add safe.directory /ws
|
||||
GG_BUILD_MUSA=1 bash ./ci/run.sh /ci-results /ci-cache
|
||||
```
|
||||
|
||||
@@ -49,14 +49,6 @@ mkdir -p "$2"
|
||||
OUT=$(realpath "$1")
|
||||
MNT=$(realpath "$2")
|
||||
|
||||
# gpu-rocm self-hosted runner can't upload logs to blob; keep each run's logs in
|
||||
# their own dir keyed by the GitHub run id so an Actions run URL maps to its logs.
|
||||
if [ -n "${GG_BUILD_ROCM}" ] && [ -n "${GITHUB_RUN_ID}" ]; then
|
||||
OUT="$OUT/run-${GITHUB_RUN_ID}-${GITHUB_RUN_ATTEMPT:-1}"
|
||||
mkdir -p "$OUT"
|
||||
echo "ci results dir: $OUT"
|
||||
fi
|
||||
|
||||
rm -f $OUT/*.log
|
||||
|
||||
sd=`dirname $0`
|
||||
|
||||
+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"
|
||||
|
||||
+204
-107
@@ -49,7 +49,7 @@
|
||||
#include <unistd.h>
|
||||
#endif
|
||||
|
||||
#if defined(__linux__)
|
||||
#if !defined(_WIN32) && !defined(__APPLE__)
|
||||
#include <sys/types.h>
|
||||
#include <pwd.h>
|
||||
#endif
|
||||
@@ -613,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;
|
||||
@@ -940,11 +912,27 @@ std::string fs_path_to_utf8(const std::filesystem::path & path) {
|
||||
return std::string(value.begin(), value.end());
|
||||
}
|
||||
|
||||
// returns true if successful, false otherwise
|
||||
bool fs_create_directory_with_parents(const std::string & path) {
|
||||
void fs_write_atomic(const std::filesystem::path & path, const std::string & data) {
|
||||
std::error_code ec;
|
||||
std::filesystem::create_directories(std::filesystem::u8path(path), ec);
|
||||
return !ec;
|
||||
std::filesystem::path path_tmp = path;
|
||||
path_tmp += ".tmp";
|
||||
|
||||
if (path.has_parent_path()) {
|
||||
std::filesystem::create_directories(path.parent_path(), ec);
|
||||
}
|
||||
|
||||
std::ofstream file(path_tmp, std::ios::binary);
|
||||
file << data;
|
||||
file.close();
|
||||
|
||||
if (!file.fail()) {
|
||||
std::filesystem::rename(path_tmp, path, ec);
|
||||
}
|
||||
|
||||
if (file.fail() || ec) {
|
||||
std::filesystem::remove(path_tmp, ec);
|
||||
throw std::runtime_error("failed to write file: " + fs_path_to_utf8(path));
|
||||
}
|
||||
}
|
||||
|
||||
bool fs_is_directory(const std::string & path) {
|
||||
@@ -969,58 +957,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() {
|
||||
@@ -1068,14 +1050,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) {
|
||||
@@ -1480,7 +1463,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;
|
||||
@@ -1489,7 +1473,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);
|
||||
@@ -1553,9 +1539,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;
|
||||
}
|
||||
@@ -2142,29 +2132,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;
|
||||
@@ -2194,7 +2293,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;
|
||||
@@ -2204,10 +2303,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");
|
||||
@@ -2216,7 +2313,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;
|
||||
|
||||
+57
-6
@@ -8,6 +8,7 @@
|
||||
#include "ggml.h"
|
||||
#include "llama.h"
|
||||
|
||||
#include <array>
|
||||
#include <list>
|
||||
#include <set>
|
||||
#include <sstream>
|
||||
@@ -809,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;
|
||||
@@ -877,7 +880,6 @@ 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);
|
||||
|
||||
@@ -902,16 +904,18 @@ std::string fs_path_to_utf8(const std::filesystem::path & path);
|
||||
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 {
|
||||
@@ -925,6 +929,8 @@ std::vector<common_file_info> fs_list(const std::string & path, bool include_dir
|
||||
// fs open, also handle UTF8 on Windows
|
||||
std::ifstream fs_open_ifstream(const std::string & fname, std::ios_base::openmode mode);
|
||||
|
||||
void fs_write_atomic(const std::filesystem::path & path, const std::string & data);
|
||||
|
||||
//
|
||||
// TTY utils
|
||||
//
|
||||
@@ -1034,9 +1040,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
|
||||
//
|
||||
|
||||
+10
-34
@@ -46,39 +46,9 @@
|
||||
// downloader
|
||||
//
|
||||
|
||||
// validate repo name format: owner/repo
|
||||
static void write_file(const std::string & fname, const std::string & content) {
|
||||
const std::string fname_tmp = fname + ".tmp";
|
||||
std::ofstream file(fname_tmp);
|
||||
if (!file) {
|
||||
throw std::runtime_error(string_format("error: failed to open file '%s'\n", fname.c_str()));
|
||||
}
|
||||
|
||||
try {
|
||||
file << content;
|
||||
file.close();
|
||||
|
||||
// Makes write atomic
|
||||
if (rename(fname_tmp.c_str(), fname.c_str()) != 0) {
|
||||
LOG_ERR("%s: unable to rename file: %s to %s\n", __func__, fname_tmp.c_str(), fname.c_str());
|
||||
// If rename fails, try to delete the temporary file
|
||||
if (remove(fname_tmp.c_str()) != 0) {
|
||||
LOG_ERR("%s: unable to delete temporary file: %s\n", __func__, fname_tmp.c_str());
|
||||
}
|
||||
}
|
||||
} catch (...) {
|
||||
// If anything fails, try to delete the temporary file
|
||||
if (remove(fname_tmp.c_str()) != 0) {
|
||||
LOG_ERR("%s: unable to delete temporary file: %s\n", __func__, fname_tmp.c_str());
|
||||
}
|
||||
|
||||
throw std::runtime_error(string_format("error: failed to write file '%s'\n", fname.c_str()));
|
||||
}
|
||||
}
|
||||
|
||||
static void write_etag(const std::string & path, const std::string & etag) {
|
||||
const std::string etag_path = path + ".etag";
|
||||
write_file(etag_path, etag);
|
||||
fs_write_atomic(std::filesystem::u8path(etag_path), etag);
|
||||
LOG_DBG("%s: file etag saved: %s\n", __func__, etag_path.c_str());
|
||||
}
|
||||
|
||||
@@ -274,6 +244,12 @@ static bool common_pull_file(httplib::Client & cli,
|
||||
return false;
|
||||
}
|
||||
|
||||
ofs.close();
|
||||
if (!ofs) {
|
||||
LOG_ERR("%s: error closing file: %s\n", __func__, path_tmp.c_str());
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
@@ -286,7 +262,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 +453,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 +919,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;
|
||||
|
||||
+24
-39
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -173,28 +172,6 @@ static bool is_valid_subpath(const fs::path & path, const fs::path & subpath) {
|
||||
return b_end == b.end();
|
||||
}
|
||||
|
||||
static void safe_write_file(const fs::path & path, const std::string & data) {
|
||||
fs::path path_tmp = path.string() + ".tmp";
|
||||
|
||||
if (path.has_parent_path()) {
|
||||
fs::create_directories(path.parent_path());
|
||||
}
|
||||
|
||||
std::ofstream file(path_tmp);
|
||||
file << data;
|
||||
file.close();
|
||||
|
||||
std::error_code ec;
|
||||
|
||||
if (!file.fail()) {
|
||||
fs::rename(path_tmp, path, ec);
|
||||
}
|
||||
if (file.fail() || ec) {
|
||||
fs::remove(path_tmp, ec);
|
||||
throw std::runtime_error("failed to write file: " + path.string());
|
||||
}
|
||||
}
|
||||
|
||||
static common_json api_get(const std::string & url,
|
||||
const std::string & token) {
|
||||
auto [cli, parts] = common_http_client(url);
|
||||
@@ -241,6 +218,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() ||
|
||||
@@ -251,24 +229,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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -277,7 +259,7 @@ static std::string get_repo_commit(const std::string & repo_id,
|
||||
return {};
|
||||
}
|
||||
|
||||
safe_write_file(refs_path / name, commit);
|
||||
fs_write_atomic(refs_path / name_path, commit);
|
||||
return commit;
|
||||
|
||||
} catch (const common_json_error & e) {
|
||||
@@ -326,7 +308,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;
|
||||
}
|
||||
@@ -346,12 +330,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;
|
||||
}
|
||||
@@ -418,7 +402,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;
|
||||
@@ -441,8 +425,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));
|
||||
}
|
||||
@@ -456,8 +441,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;
|
||||
@@ -504,7 +489,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;
|
||||
|
||||
+39
-12
@@ -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();
|
||||
|
||||
@@ -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`
|
||||
|
||||
+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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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|
|
||||
|
||||
@@ -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.
|
||||
|
||||
+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
@@ -65,8 +65,7 @@ int main(int argc, char ** argv) {
|
||||
}
|
||||
}, nullptr);
|
||||
|
||||
// load dynamic backends
|
||||
ggml_backend_load_all();
|
||||
llama_backend_init();
|
||||
|
||||
// initialize the model
|
||||
llama_model_params model_params = llama_model_default_params();
|
||||
|
||||
@@ -77,9 +77,7 @@ int main(int argc, char ** argv) {
|
||||
}
|
||||
}
|
||||
|
||||
// load dynamic backends
|
||||
|
||||
ggml_backend_load_all();
|
||||
llama_backend_init();
|
||||
|
||||
// initialize the model
|
||||
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
+40
-24
@@ -1372,30 +1372,6 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra
|
||||
const int src_backend_id = sched->hv_tensor_backend_ids[src_id];
|
||||
GGML_ASSERT(src_backend_id != -1); // all inputs should be assigned by now
|
||||
|
||||
if (src->flags & GGML_TENSOR_FLAG_INPUT && sched->n_copies > 1) {
|
||||
if (tensor_id_copy(src_id, src_backend_id, 0) == NULL) {
|
||||
ggml_backend_t backend = sched->backends[src_backend_id];
|
||||
for (int c = 0; c < sched->n_copies; c++) {
|
||||
struct ggml_tensor * tensor_copy;
|
||||
if (c == sched->cur_copy) {
|
||||
tensor_copy = src; // use the original tensor as the current copy
|
||||
} else {
|
||||
tensor_copy = ggml_dup_tensor_layout(sched->ctx, src);
|
||||
ggml_format_name(tensor_copy, "%s#%s#%d", ggml_backend_name(backend), src->name, c);
|
||||
}
|
||||
ggml_set_input(tensor_copy);
|
||||
ggml_set_output(tensor_copy); // prevent ggml-alloc from overwriting the tensor
|
||||
tensor_id_copy(src_id, src_backend_id, c) = tensor_copy;
|
||||
SET_CAUSE(tensor_copy, "4.cpy");
|
||||
}
|
||||
int n_graph_inputs = sched->n_graph_inputs++;
|
||||
if (n_graph_inputs >= sched->graph_inputs_capacity) {
|
||||
ggml_backend_sched_graph_inputs_grow(sched);
|
||||
}
|
||||
sched->graph_inputs[n_graph_inputs] = src;
|
||||
}
|
||||
}
|
||||
|
||||
if (src_backend_id != cur_backend_id && !ggml_backend_sched_buffer_supported(sched, src, cur_backend_id)) {
|
||||
// create a copy of the input in the split's backend
|
||||
if (tensor_id_copy(src_id, cur_backend_id, 0) == NULL) {
|
||||
@@ -1428,6 +1404,46 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra
|
||||
ggml_backend_sched_print_assignments(sched, graph);
|
||||
}
|
||||
|
||||
// pass 6: collect all input tensors into graph_inputs
|
||||
// this includes inputs not consumed by any node (e.g. the embeddings input of a text-only batch) so that
|
||||
// the graph composition does not depend on which inputs are used, which would otherwise cause graph
|
||||
// reallocations when switching between different types of batches [GGML_SCHED_DEBUG_REALLOC]
|
||||
if (sched->n_copies > 1) {
|
||||
for (int i = 0; i < graph->n_leafs; i++) {
|
||||
struct ggml_tensor * leaf = graph->leafs[i];
|
||||
if ((leaf->flags & GGML_TENSOR_FLAG_INPUT) == 0) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const size_t leaf_id = hash_id(leaf);
|
||||
const int leaf_backend_id = tensor_backend_id(leaf);
|
||||
GGML_ASSERT(leaf_backend_id != -1); // all leafs should be assigned by now
|
||||
|
||||
if (tensor_id_copy(leaf_id, leaf_backend_id, 0) == NULL) {
|
||||
ggml_backend_t backend = sched->backends[leaf_backend_id];
|
||||
for (int c = 0; c < sched->n_copies; c++) {
|
||||
struct ggml_tensor * tensor_copy;
|
||||
if (c == sched->cur_copy) {
|
||||
tensor_copy = leaf; // use the original tensor as the current copy
|
||||
} else {
|
||||
tensor_copy = ggml_dup_tensor_layout(sched->ctx, leaf);
|
||||
ggml_format_name(tensor_copy, "%s#%s#%d", ggml_backend_name(backend), leaf->name, c);
|
||||
}
|
||||
ggml_set_input(tensor_copy);
|
||||
ggml_set_output(tensor_copy); // prevent ggml-alloc from overwriting the tensor
|
||||
tensor_id_copy(leaf_id, leaf_backend_id, c) = tensor_copy;
|
||||
SET_CAUSE(tensor_copy, "6.cpy");
|
||||
}
|
||||
}
|
||||
|
||||
int n_graph_inputs = sched->n_graph_inputs++;
|
||||
if (n_graph_inputs >= sched->graph_inputs_capacity) {
|
||||
ggml_backend_sched_graph_inputs_grow(sched);
|
||||
}
|
||||
sched->graph_inputs[n_graph_inputs] = leaf;
|
||||
}
|
||||
}
|
||||
|
||||
// swap node_backend_ids and leaf _backend_ids with prevs
|
||||
{
|
||||
int * tmp = sched->node_backend_ids;
|
||||
|
||||
@@ -106,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}")
|
||||
@@ -163,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);
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -1818,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;
|
||||
}
|
||||
|
||||
@@ -5212,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) {
|
||||
@@ -5488,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;
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -63,6 +63,7 @@
|
||||
#include "htp/rope-ops.h"
|
||||
#include "htp/ssm-conv.h"
|
||||
#include "htp/gated-delta-net-ops.h"
|
||||
#include "htp/argsort-ops.h"
|
||||
#include "htp_iface.h"
|
||||
#include "htp-drv.h"
|
||||
|
||||
@@ -368,6 +369,13 @@ static void ggml_hexagon_precompute_gated_delta_net_params(
|
||||
struct htp_gdn_kernel_params * kparams
|
||||
);
|
||||
|
||||
static void ggml_hexagon_precompute_sort_params(
|
||||
const struct ggml_hexagon_session * sess,
|
||||
const struct ggml_tensor * op,
|
||||
bool is_top_k,
|
||||
struct htp_sort_kernel_params * kparams
|
||||
);
|
||||
|
||||
static void ggml_hexagon_precompute_fused_mmnx_params(
|
||||
const struct ggml_hexagon_session * sess,
|
||||
const struct ggml_tensor * src0,
|
||||
@@ -4998,13 +5006,48 @@ static bool ggml_hexagon_precompute_binary_params(
|
||||
const bool is_add_id = op == HTP_OP_ADD_ID;
|
||||
const bool is_scalar = !is_add_id && src1->ne[0] == 1;
|
||||
const bool is_transposed = src0->nb[1] < src0_row_size || src1->nb[1] < src1_row_size || dst->nb[1] < dst_row_size;
|
||||
const bool is_row_bcast = !is_add_id && !is_scalar && !is_transposed &&
|
||||
src1->ne[0] == src0->ne[0] &&
|
||||
(src0->ne[1] > 1 || src0->ne[2] > 1 || src0->ne[3] > 1) &&
|
||||
src1->ne[1] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1;
|
||||
const bool is_same_shape = !is_add_id && !is_scalar && !is_transposed &&
|
||||
src1->ne[0] == src0->ne[0] &&
|
||||
(src1->ne[1] == src0->ne[1] || src1->ne[1] == 1) &&
|
||||
src1->ne[1] == src0->ne[1] &&
|
||||
(src1->ne[2] == src0->ne[2] || src1->ne[2] == 1) &&
|
||||
(src1->ne[3] == src0->ne[3] || src1->ne[3] == 1);
|
||||
const bool is_row_bcast = is_same_shape && src1->ne[1] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1;
|
||||
const bool is_complex = !is_add_id && !is_scalar && !is_same_shape && (src1->ne[0] == src0->ne[0]);
|
||||
const bool is_complex = !is_add_id && !is_scalar && !is_same_shape && !is_row_bcast && (src1->ne[0] == src0->ne[0]);
|
||||
const bool is_contig = ggml_is_contiguous(src0) && ggml_is_contiguous(src1) && ggml_is_contiguous(dst);
|
||||
const bool is_scalar_broadcast = !is_add_id && (ggml_nelements(src1) == 1);
|
||||
|
||||
if (!is_add_id && is_contig && (ggml_are_same_shape(src0, src1) || is_scalar_broadcast)) {
|
||||
const uint32_t total_elems = (uint32_t) ggml_nelements(src0);
|
||||
const uint32_t n_threads = sess->n_threads;
|
||||
const uint32_t max_chunk_elems = 32768 / elem_size;
|
||||
const uint32_t min_chunk_elems = 256;
|
||||
const uint32_t target_chunk_elems = hex_round_up((total_elems + (2 * n_threads) - 1) / (2 * n_threads), 32);
|
||||
const uint32_t chunk_size = (std::min)(max_chunk_elems, (std::max)(target_chunk_elems, min_chunk_elems));
|
||||
const uint32_t chunk_bytes = hex_round_up(chunk_size * elem_size, 128);
|
||||
|
||||
kparams->kernel_type = HTP_BINARY_KERNEL_CHUNKED;
|
||||
kparams->n_threads = n_threads;
|
||||
kparams->rows_per_buffer = 1;
|
||||
kparams->src0_row_size_aligned = src0_row_size_aligned;
|
||||
kparams->src1_row_size_aligned = is_scalar_broadcast ? 0 : src1_row_size_aligned;
|
||||
kparams->dst_row_size_aligned = dst_row_size_aligned;
|
||||
kparams->src1_size = 0;
|
||||
kparams->chunk_size = chunk_size;
|
||||
kparams->chunk_bytes = chunk_bytes;
|
||||
kparams->is_scalar = is_scalar_broadcast ? 1 : 0;
|
||||
|
||||
struct htp_binary_vtcm_layout L;
|
||||
htp_binary_vtcm_layout_build(&L, kparams, sess->vtcm_size);
|
||||
if (L.total_bytes == 0 || L.total_bytes > sess->vtcm_size) {
|
||||
return false;
|
||||
}
|
||||
|
||||
kparams->vtcm_size = L.total_bytes;
|
||||
return true;
|
||||
}
|
||||
|
||||
enum htp_binary_kernel_type kernel_type;
|
||||
size_t src1_size = 0;
|
||||
@@ -5042,6 +5085,32 @@ static bool ggml_hexagon_precompute_binary_params(
|
||||
struct htp_binary_vtcm_layout L;
|
||||
htp_binary_vtcm_layout_build(&L, kparams, sess->vtcm_size);
|
||||
if (L.rows_per_buffer == 0 || L.total_bytes > sess->vtcm_size) {
|
||||
if (!is_add_id && is_contig && (ggml_are_same_shape(src0, src1) || is_scalar_broadcast)) {
|
||||
const uint32_t total_elems = (uint32_t) ggml_nelements(src0);
|
||||
const uint32_t n_threads = sess->n_threads;
|
||||
const uint32_t max_chunk_elems = 32768 / elem_size;
|
||||
const uint32_t min_chunk_elems = 256;
|
||||
const uint32_t target_chunk_elems = hex_round_up((total_elems + (2 * n_threads) - 1) / (2 * n_threads), 32);
|
||||
const uint32_t chunk_size = (std::min)(max_chunk_elems, (std::max)(target_chunk_elems, min_chunk_elems));
|
||||
const uint32_t chunk_bytes = hex_round_up(chunk_size * elem_size, 128);
|
||||
|
||||
kparams->kernel_type = HTP_BINARY_KERNEL_CHUNKED;
|
||||
kparams->n_threads = n_threads;
|
||||
kparams->rows_per_buffer = 1;
|
||||
kparams->src1_row_size_aligned = is_scalar_broadcast ? 0 : src1_row_size_aligned;
|
||||
kparams->src1_size = 0;
|
||||
kparams->chunk_size = chunk_size;
|
||||
kparams->chunk_bytes = chunk_bytes;
|
||||
kparams->is_scalar = is_scalar_broadcast ? 1 : 0;
|
||||
|
||||
htp_binary_vtcm_layout_build(&L, kparams, sess->vtcm_size);
|
||||
if (L.total_bytes == 0 || L.total_bytes > sess->vtcm_size) {
|
||||
return false;
|
||||
}
|
||||
|
||||
kparams->vtcm_size = L.total_bytes;
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -5139,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(
|
||||
@@ -5479,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
|
||||
@@ -5830,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;
|
||||
@@ -5853,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;
|
||||
@@ -5874,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];
|
||||
@@ -6020,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;
|
||||
}
|
||||
|
||||
@@ -6040,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;
|
||||
@@ -6058,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;
|
||||
}
|
||||
|
||||
@@ -6401,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;
|
||||
@@ -6440,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;
|
||||
}
|
||||
@@ -6670,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));
|
||||
}
|
||||
@@ -7308,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;
|
||||
@@ -7466,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;
|
||||
@@ -7486,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:
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -150,7 +150,9 @@ int op_matmul_nx(struct htp_ops_context * octx);
|
||||
int op_matmul_id_nx(struct htp_ops_context * octx);
|
||||
int op_binary(struct htp_ops_context * octx);
|
||||
int op_unary(struct htp_ops_context * octx);
|
||||
int op_sum(struct htp_ops_context * octx);
|
||||
int op_sum_rows(struct htp_ops_context * octx);
|
||||
int op_argmax(struct htp_ops_context * octx);
|
||||
int op_activations(struct htp_ops_context * octx);
|
||||
int op_softmax(struct htp_ops_context * octx);
|
||||
int op_add_id(struct htp_ops_context * octx);
|
||||
|
||||
@@ -69,6 +69,7 @@ enum htp_op_code {
|
||||
HTP_OP_UNARY_ABS,
|
||||
HTP_OP_UNARY_LOG,
|
||||
HTP_OP_UNARY_RELU,
|
||||
HTP_OP_UNARY_STEP,
|
||||
HTP_OP_GLU_SWIGLU,
|
||||
HTP_OP_GLU_SWIGLU_OAI,
|
||||
HTP_OP_GLU_GEGLU,
|
||||
@@ -86,6 +87,7 @@ enum htp_op_code {
|
||||
HTP_OP_TOP_K,
|
||||
HTP_OP_SQR,
|
||||
HTP_OP_SQRT,
|
||||
HTP_OP_SUM,
|
||||
HTP_OP_SUM_ROWS,
|
||||
HTP_OP_SSM_CONV,
|
||||
HTP_OP_REPEAT,
|
||||
@@ -108,6 +110,7 @@ enum htp_op_code {
|
||||
HTP_OP_GLU_SWIGLU_CLAMP,
|
||||
HTP_OP_MDEV_GROUP,
|
||||
HTP_OP_ROLL,
|
||||
HTP_OP_ARGMAX,
|
||||
|
||||
HTP_OP_INVALID
|
||||
};
|
||||
|
||||
@@ -579,6 +579,58 @@ static inline void hvx_abs_f16(uint8_t * restrict dst, const uint8_t * restrict
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
// Step
|
||||
//
|
||||
|
||||
static inline void hvx_step_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) {
|
||||
assert((unsigned long) dst % 128 == 0);
|
||||
assert((unsigned long) src % 128 == 0);
|
||||
|
||||
HVX_Vector * restrict vdst = (HVX_Vector *) dst;
|
||||
HVX_Vector * restrict vsrc = (HVX_Vector *) src;
|
||||
|
||||
const uint32_t elem_size = sizeof(float);
|
||||
const uint32_t epv = 128 / elem_size;
|
||||
const uint32_t nvec = n / epv;
|
||||
const uint32_t nloe = n % epv;
|
||||
|
||||
uint32_t i = 0;
|
||||
|
||||
_Pragma("unroll(4)")
|
||||
for (; i < nvec; i++) {
|
||||
vdst[i] = hvx_vec_step_f32(vsrc[i]);
|
||||
}
|
||||
if (nloe) {
|
||||
HVX_Vector v = hvx_vec_step_f32(vsrc[i]);
|
||||
hvx_vec_store_a((void *) &vdst[i], nloe * elem_size, v);
|
||||
}
|
||||
}
|
||||
|
||||
static inline void hvx_step_f16_aa(uint8_t * restrict dst, const uint8_t * restrict src, uint32_t n) {
|
||||
assert((unsigned long) dst % 128 == 0);
|
||||
assert((unsigned long) src % 128 == 0);
|
||||
|
||||
HVX_Vector * restrict vdst = (HVX_Vector *) dst;
|
||||
HVX_Vector * restrict vsrc = (HVX_Vector *) src;
|
||||
|
||||
const uint32_t elem_size = sizeof(_Float16);
|
||||
const uint32_t epv = 128 / elem_size;
|
||||
const uint32_t nvec = n / epv;
|
||||
const uint32_t nloe = n % epv;
|
||||
|
||||
uint32_t i = 0;
|
||||
|
||||
_Pragma("unroll(4)")
|
||||
for (; i < nvec; i++) {
|
||||
vdst[i] = hvx_vec_step_f16(vsrc[i]);
|
||||
}
|
||||
if (nloe) {
|
||||
HVX_Vector v = hvx_vec_step_f16(vsrc[i]);
|
||||
hvx_vec_store_a((void *) &vdst[i], nloe * elem_size, v);
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
// Square
|
||||
//
|
||||
|
||||
@@ -111,6 +111,20 @@ static inline HVX_Vector hvx_vec_neg_f32(HVX_Vector v) {
|
||||
#endif // __HVX_ARCH__ > 75
|
||||
}
|
||||
|
||||
static inline HVX_Vector hvx_vec_step_f32(HVX_Vector v) {
|
||||
const HVX_Vector zero = Q6_V_vzero();
|
||||
const HVX_Vector one = hvx_vec_splat_f32(1.0f);
|
||||
HVX_VectorPred q = Q6_Q_vcmp_gt_VsfVsf(v, zero);
|
||||
return Q6_V_vmux_QVV(q, one, zero);
|
||||
}
|
||||
|
||||
static inline HVX_Vector hvx_vec_step_f16(HVX_Vector v) {
|
||||
const HVX_Vector zero = Q6_V_vzero();
|
||||
const HVX_Vector one = hvx_vec_splat_f16((_Float16) 1.0f);
|
||||
HVX_VectorPred q = Q6_Q_vcmp_gt_VhfVhf(v, zero);
|
||||
return Q6_V_vmux_QVV(q, one, zero);
|
||||
}
|
||||
|
||||
static inline HVX_VectorPred hvx_vec_is_nan_f16(HVX_Vector v) {
|
||||
const HVX_Vector vnan_exp = Q6_Vh_vsplat_R(0x7C00);
|
||||
const HVX_Vector vnan_frac = Q6_Vh_vsplat_R(0x7FFF);
|
||||
|
||||
@@ -326,6 +326,79 @@ static inline int32_t hvx_reduce_max_i32(const uint8_t * restrict src, const int
|
||||
}
|
||||
}
|
||||
|
||||
static inline void hvx_argmax_f32(
|
||||
const float * restrict src,
|
||||
uint32_t n,
|
||||
uint32_t offset,
|
||||
float * out_val,
|
||||
int32_t * out_idx
|
||||
) {
|
||||
if (n == 0) {
|
||||
*out_val = -INFINITY;
|
||||
*out_idx = (int32_t) offset;
|
||||
return;
|
||||
}
|
||||
|
||||
if (n < 32 || !hex_is_aligned((void *) src, 128)) {
|
||||
float best_val = src[0];
|
||||
int32_t best_idx = (int32_t) offset;
|
||||
for (uint32_t i = 1; i < n; i++) {
|
||||
if (src[i] > best_val) {
|
||||
best_val = src[i];
|
||||
best_idx = (int32_t) (offset + i);
|
||||
}
|
||||
}
|
||||
*out_val = best_val;
|
||||
*out_idx = best_idx;
|
||||
return;
|
||||
}
|
||||
|
||||
static const int32_t c_lane_idx[32] __attribute__((aligned(128))) = {
|
||||
0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15,
|
||||
16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31
|
||||
};
|
||||
|
||||
const HVX_Vector v_init_idx = *(const HVX_Vector *) c_lane_idx;
|
||||
const HVX_Vector v_step = Q6_V_vsplat_R(32);
|
||||
HVX_Vector v_cur_idx = Q6_Vw_vadd_VwVw(v_init_idx, Q6_V_vsplat_R((int32_t) offset));
|
||||
HVX_Vector v_max_val = hvx_vec_splat_f32(-INFINITY);
|
||||
HVX_Vector v_max_idx = v_cur_idx;
|
||||
|
||||
const HVX_Vector * vsrc = (const HVX_Vector *) src;
|
||||
const uint32_t nvec = n / 32;
|
||||
|
||||
for (uint32_t vi = 0; vi < nvec; vi++) {
|
||||
HVX_Vector v = vsrc[vi];
|
||||
HVX_VectorPred pred = Q6_Q_vcmp_gt_VsfVsf(v, v_max_val);
|
||||
v_max_val = Q6_V_vmux_QVV(pred, v, v_max_val);
|
||||
v_max_idx = Q6_V_vmux_QVV(pred, v_cur_idx, v_max_idx);
|
||||
v_cur_idx = Q6_Vw_vadd_VwVw(v_cur_idx, v_step);
|
||||
}
|
||||
|
||||
HVX_VectorAlias u_val, u_idx;
|
||||
u_val.v = v_max_val;
|
||||
u_idx.v = v_max_idx;
|
||||
|
||||
float best_val = u_val.fp32[0];
|
||||
int32_t best_idx = (int32_t) u_idx.w[0];
|
||||
for (int i = 1; i < 32; i++) {
|
||||
if (u_val.fp32[i] > best_val) {
|
||||
best_val = u_val.fp32[i];
|
||||
best_idx = (int32_t) u_idx.w[i];
|
||||
}
|
||||
}
|
||||
|
||||
for (uint32_t i = nvec * 32; i < n; i++) {
|
||||
if (src[i] > best_val) {
|
||||
best_val = src[i];
|
||||
best_idx = (int32_t) (offset + i);
|
||||
}
|
||||
}
|
||||
|
||||
*out_val = best_val;
|
||||
*out_idx = best_idx;
|
||||
}
|
||||
|
||||
#undef hvx_reduce_loop_body
|
||||
#undef HVX_REDUCE_MAX_OP
|
||||
#undef HVX_REDUCE_SUM_OP
|
||||
|
||||
@@ -857,6 +857,7 @@ static int execute_op(struct htp_ops_context * octx) {
|
||||
case HTP_OP_UNARY_ABS:
|
||||
case HTP_OP_UNARY_LOG:
|
||||
case HTP_OP_UNARY_RELU:
|
||||
case HTP_OP_UNARY_STEP:
|
||||
case HTP_OP_L2_NORM:
|
||||
return op_unary(octx);
|
||||
|
||||
@@ -882,6 +883,9 @@ static int execute_op(struct htp_ops_context * octx) {
|
||||
case HTP_OP_GET_ROWS:
|
||||
return op_get_rows(octx);
|
||||
|
||||
case HTP_OP_SUM:
|
||||
return op_sum(octx);
|
||||
|
||||
case HTP_OP_SUM_ROWS:
|
||||
return op_sum_rows(octx);
|
||||
|
||||
@@ -898,6 +902,9 @@ static int execute_op(struct htp_ops_context * octx) {
|
||||
case HTP_OP_TOP_K:
|
||||
return op_top_k(octx);
|
||||
|
||||
case HTP_OP_ARGMAX:
|
||||
return op_argmax(octx);
|
||||
|
||||
case HTP_OP_SSM_CONV:
|
||||
return op_ssm_conv(octx);
|
||||
|
||||
|
||||
@@ -83,16 +83,17 @@ static void sum_rows_thread_f32(unsigned int nth, unsigned int ith, void *data)
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start_row);
|
||||
|
||||
for (uint32_t ir = 0; ir < n_rows; ir++) {
|
||||
const float * restrict src_local = src_th + (ir * (src_stride / sizeof(float)));
|
||||
const float * restrict src_local = (const float *) ((const uint8_t *) src_th + ir * src_stride);
|
||||
float * restrict dst_local = (float *) ((uint8_t *) dst_th + ir * dst_stride);
|
||||
|
||||
if (ir + 1 < n_rows) {
|
||||
hex_l2fetch(src_local + (src_stride / sizeof(float)), src_stride, src_stride, 1);
|
||||
hex_l2fetch((const uint8_t *) src_local + src_stride, src_stride, src_stride, 1);
|
||||
}
|
||||
|
||||
if (opt_path) {
|
||||
dst_th[ir] = hvx_reduce_sum_f32_a((const uint8_t *) src_local, ne00);
|
||||
*dst_local = hvx_reduce_sum_f32_a((const uint8_t *) src_local, ne00);
|
||||
} else {
|
||||
dst_th[ir] = hvx_reduce_sum_f32((const uint8_t *) src_local, ne00);
|
||||
*dst_local = hvx_reduce_sum_f32((const uint8_t *) src_local, ne00);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -152,3 +153,258 @@ int op_sum_rows(struct htp_ops_context * octx) {
|
||||
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
struct sum_context {
|
||||
struct htp_ops_context * octx;
|
||||
const float * src_data;
|
||||
float partial_sums[HTP_MAX_NTHREADS];
|
||||
uint32_t total_elems;
|
||||
uint32_t elems_per_thread;
|
||||
};
|
||||
|
||||
static void sum_thread_f32(unsigned int nth, unsigned int ith, void * data) {
|
||||
struct sum_context * sctx = (struct sum_context *) data;
|
||||
const uint32_t start = sctx->elems_per_thread * ith;
|
||||
const uint32_t end = MIN(start + sctx->elems_per_thread, sctx->total_elems);
|
||||
|
||||
if (start >= end) {
|
||||
sctx->partial_sums[ith] = 0.0f;
|
||||
return;
|
||||
}
|
||||
|
||||
const uint32_t n = end - start;
|
||||
const float * src = sctx->src_data + start;
|
||||
|
||||
struct htp_thread_trace * tr = &sctx->octx->ctx->trace[ith];
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
|
||||
|
||||
hex_l2fetch_block((const void *) src, n * sizeof(float));
|
||||
|
||||
sctx->partial_sums[ith] = hvx_reduce_sum_f32((const uint8_t *) src, n);
|
||||
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
|
||||
}
|
||||
|
||||
int op_sum(struct htp_ops_context * octx) {
|
||||
const struct htp_tensor * src0 = octx->src[0];
|
||||
const struct htp_tensor * dst = octx->dst;
|
||||
|
||||
if (src0->type != HTP_TYPE_F32) {
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
|
||||
if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
|
||||
if (octx->ctx->mdev.count > 1 && octx->ctx->mdev.idx > 0) {
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
const uint32_t total_elems = (uint32_t) (src0->ne[0] * src0->ne[1] * src0->ne[2] * src0->ne[3]);
|
||||
if (total_elems == 0) {
|
||||
((float *) dst->data)[0] = 0.0f;
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
const uint32_t n_threads = (total_elems >= 1024) ? MIN(octx->n_threads, HTP_MAX_NTHREADS) : 1;
|
||||
const uint32_t raw_chunk = (total_elems + n_threads - 1) / n_threads;
|
||||
const uint32_t elems_per_thread = hex_round_up(raw_chunk, 32);
|
||||
|
||||
struct sum_context sctx = {
|
||||
.octx = octx,
|
||||
.src_data = (const float *) src0->data,
|
||||
.total_elems = total_elems,
|
||||
.elems_per_thread = elems_per_thread,
|
||||
};
|
||||
|
||||
work_queue_run(octx->ctx->work_queue, sum_thread_f32, &sctx, n_threads);
|
||||
|
||||
float sum = 0.0f;
|
||||
for (uint32_t i = 0; i < n_threads; i++) {
|
||||
sum += sctx.partial_sums[i];
|
||||
}
|
||||
((float *) dst->data)[0] = sum;
|
||||
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
static inline void argmax_slice_f32(
|
||||
const float * restrict src,
|
||||
uint32_t n,
|
||||
uint32_t offset,
|
||||
float * out_val,
|
||||
int32_t * out_idx
|
||||
) {
|
||||
hvx_argmax_f32(src, n, offset, out_val, out_idx);
|
||||
}
|
||||
|
||||
struct argmax_context {
|
||||
struct htp_ops_context * octx;
|
||||
const float * src_data;
|
||||
int32_t * dst_data;
|
||||
uint32_t ne00;
|
||||
uint32_t src_stride;
|
||||
uint32_t dst_stride;
|
||||
uint32_t row_start;
|
||||
uint32_t nrows;
|
||||
uint32_t rows_per_thread;
|
||||
uint32_t elems_per_thread;
|
||||
|
||||
float partial_max[HTP_MAX_NTHREADS];
|
||||
int32_t partial_idx[HTP_MAX_NTHREADS];
|
||||
};
|
||||
|
||||
static void argmax_thread_single_row(unsigned int nth, unsigned int ith, void * data) {
|
||||
struct argmax_context * actx = (struct argmax_context *) data;
|
||||
const uint32_t start = actx->elems_per_thread * ith;
|
||||
const uint32_t end = MIN(start + actx->elems_per_thread, actx->ne00);
|
||||
|
||||
if (start >= end) {
|
||||
actx->partial_max[ith] = -INFINITY;
|
||||
actx->partial_idx[ith] = 0;
|
||||
return;
|
||||
}
|
||||
|
||||
const uint32_t n = end - start;
|
||||
const float * src = actx->src_data + start;
|
||||
|
||||
struct htp_thread_trace * tr = &actx->octx->ctx->trace[ith];
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
|
||||
|
||||
hex_l2fetch_block((const void *) src, n * sizeof(float));
|
||||
|
||||
float max_val;
|
||||
int32_t max_idx;
|
||||
argmax_slice_f32(src, n, start, &max_val, &max_idx);
|
||||
|
||||
actx->partial_max[ith] = max_val;
|
||||
actx->partial_idx[ith] = max_idx;
|
||||
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) start);
|
||||
}
|
||||
|
||||
static void argmax_thread_multi_row(unsigned int nth, unsigned int ith, void * data) {
|
||||
struct argmax_context * actx = (struct argmax_context *) data;
|
||||
const uint32_t r0 = actx->row_start + actx->rows_per_thread * ith;
|
||||
const uint32_t r1 = MIN(r0 + actx->rows_per_thread, actx->row_start + actx->nrows);
|
||||
|
||||
if (r0 >= r1) {
|
||||
return;
|
||||
}
|
||||
|
||||
struct htp_thread_trace * tr = &actx->octx->ctx->trace[ith];
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) r0);
|
||||
|
||||
for (uint32_t r = r0; r < r1; r++) {
|
||||
const float * src_row = (const float *) ((const uint8_t *) actx->src_data + r * actx->src_stride);
|
||||
int32_t * dst_val = (int32_t *) ((uint8_t *) actx->dst_data + r * actx->dst_stride);
|
||||
|
||||
hex_l2fetch_block((const void *) src_row, actx->ne00 * sizeof(float));
|
||||
|
||||
float max_val;
|
||||
int32_t max_idx;
|
||||
argmax_slice_f32(src_row, actx->ne00, 0, &max_val, &max_idx);
|
||||
|
||||
*dst_val = max_idx;
|
||||
}
|
||||
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, (uint16_t) r0);
|
||||
}
|
||||
|
||||
int op_argmax(struct htp_ops_context * octx) {
|
||||
const struct htp_tensor * src0 = octx->src[0];
|
||||
const struct htp_tensor * dst = octx->dst;
|
||||
|
||||
if (src0->type != HTP_TYPE_F32 || dst->type != HTP_TYPE_I32) {
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
|
||||
if (htp_tensor_is_extended(src0) || htp_tensor_is_extended(dst)) {
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
|
||||
const uint32_t ne00 = src0->ne[0];
|
||||
const uint32_t src0_nrows = src0->ne[1] * src0->ne[2] * src0->ne[3];
|
||||
|
||||
if (ne00 == 0 || src0_nrows == 0) {
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
if (src0_nrows == 1) {
|
||||
if (octx->ctx->mdev.count > 1 && octx->ctx->mdev.idx > 0) {
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
if (ne00 == 1) {
|
||||
((int32_t *) dst->data)[0] = 0;
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
const uint32_t n_threads = (ne00 >= 1024) ? MIN(octx->n_threads, HTP_MAX_NTHREADS) : 1;
|
||||
const uint32_t raw_chunk = (ne00 + n_threads - 1) / n_threads;
|
||||
const uint32_t elems_per_thread = hex_round_up(raw_chunk, 32);
|
||||
|
||||
struct argmax_context actx = {
|
||||
.octx = octx,
|
||||
.src_data = (const float *) src0->data,
|
||||
.dst_data = (int32_t *) dst->data,
|
||||
.ne00 = ne00,
|
||||
.src_stride = src0->nb[1] > 0 ? (uint32_t) src0->nb[1] : (uint32_t) (ne00 * sizeof(float)),
|
||||
.dst_stride = dst->nb[0] > 0 ? (uint32_t) dst->nb[0] : (uint32_t) sizeof(int32_t),
|
||||
.row_start = 0,
|
||||
.nrows = 1,
|
||||
.rows_per_thread = 1,
|
||||
.elems_per_thread = elems_per_thread,
|
||||
};
|
||||
|
||||
work_queue_run(octx->ctx->work_queue, argmax_thread_single_row, &actx, n_threads);
|
||||
|
||||
float best_val = actx.partial_max[0];
|
||||
int32_t best_idx = actx.partial_idx[0];
|
||||
for (uint32_t i = 1; i < n_threads; i++) {
|
||||
if (actx.partial_max[i] > best_val) {
|
||||
best_val = actx.partial_max[i];
|
||||
best_idx = actx.partial_idx[i];
|
||||
}
|
||||
}
|
||||
((int32_t *) dst->data)[0] = best_idx;
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
uint32_t row_start = 0;
|
||||
uint32_t nrows = src0_nrows;
|
||||
|
||||
if (octx->ctx->mdev.count > 1) {
|
||||
const bool can_split = htp_tensor_mdev_data_aligned(dst) &&
|
||||
htp_tensor_is_contiguous(dst, sizeof(int32_t));
|
||||
const uint32_t elems_per_chunk = can_split ? 32 : 0;
|
||||
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(
|
||||
src0_nrows, elems_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
|
||||
row_start = range.start;
|
||||
nrows = range.count;
|
||||
}
|
||||
|
||||
if (nrows == 0) {
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
const uint32_t n_threads = MIN(octx->n_threads, nrows);
|
||||
const uint32_t rows_per_thread = (nrows + n_threads - 1) / n_threads;
|
||||
|
||||
struct argmax_context actx = {
|
||||
.octx = octx,
|
||||
.src_data = (const float *) src0->data,
|
||||
.dst_data = (int32_t *) dst->data,
|
||||
.ne00 = ne00,
|
||||
.src_stride = src0->nb[1] > 0 ? (uint32_t) src0->nb[1] : (uint32_t) (ne00 * sizeof(float)),
|
||||
.dst_stride = dst->nb[0] > 0 ? (uint32_t) dst->nb[0] : (uint32_t) sizeof(int32_t),
|
||||
.row_start = row_start,
|
||||
.nrows = nrows,
|
||||
.rows_per_thread = rows_per_thread,
|
||||
.elems_per_thread = 0,
|
||||
};
|
||||
|
||||
work_queue_run(octx->ctx->work_queue, argmax_thread_multi_row, &actx, n_threads);
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
@@ -409,6 +409,20 @@ static void log_f16(const void * restrict src,
|
||||
}
|
||||
}
|
||||
|
||||
static void step_f16(const void * restrict src,
|
||||
void * restrict dst,
|
||||
const uint32_t num_rows,
|
||||
const struct htp_unary_context * uctx) {
|
||||
htp_unary_op_preamble;
|
||||
|
||||
for (uint32_t ir = 0; ir < num_rows; ir++) {
|
||||
const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned);
|
||||
uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned);
|
||||
|
||||
hvx_step_f16_aa((uint8_t *) dst_local, (const uint8_t *) src_local, ne0);
|
||||
}
|
||||
}
|
||||
|
||||
static void l2_norm_f16(const void * restrict src,
|
||||
void * restrict dst,
|
||||
const uint32_t num_rows,
|
||||
@@ -662,6 +676,20 @@ static void relu_f32(const void * restrict src,
|
||||
}
|
||||
}
|
||||
|
||||
static void step_f32(const void * restrict src,
|
||||
void * restrict dst,
|
||||
const uint32_t num_rows,
|
||||
const struct htp_unary_context * uctx) {
|
||||
htp_unary_op_preamble;
|
||||
|
||||
for (uint32_t ir = 0; ir < num_rows; ir++) {
|
||||
const uint8_t * restrict src_local = (const uint8_t *)src + (ir * src0_row_size_aligned);
|
||||
uint8_t * restrict dst_local = (uint8_t *)dst + (ir * dst_row_size_aligned);
|
||||
|
||||
hvx_step_f32_aa(dst_local, src_local, ne0);
|
||||
}
|
||||
}
|
||||
|
||||
static void log_f32(const void * restrict src,
|
||||
void * restrict dst,
|
||||
const uint32_t num_rows,
|
||||
@@ -767,6 +795,11 @@ static void tile_relu_f32(void * restrict dst, const void * restrict src, uint32
|
||||
hvx_max_scalar_f32((uint8_t *) dst, (const uint8_t *) src, 0.0f, tw);
|
||||
}
|
||||
|
||||
static void tile_step_f32(void * restrict dst, const void * restrict src, uint32_t tw, const struct htp_unary_context * uctx) {
|
||||
(void) uctx;
|
||||
hvx_step_f32_aa((uint8_t *) dst, (const uint8_t *) src, tw);
|
||||
}
|
||||
|
||||
static void tri_apply_tile_f32(const void * restrict src, void * restrict dst,
|
||||
uint32_t tile_elems, uint32_t col_start, uint32_t i01,
|
||||
uint32_t ne0, int32_t ttype) {
|
||||
@@ -1486,6 +1519,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
|
||||
case HTP_OP_UNARY_ABS: op_type = is_f16 ? "abs-f16" : "abs-f32"; break;
|
||||
case HTP_OP_UNARY_LOG: op_type = is_f16 ? "log-f16" : "log-f32"; break;
|
||||
case HTP_OP_UNARY_RELU: op_type = "relu-f32"; break;
|
||||
case HTP_OP_UNARY_STEP: op_type = is_f16 ? "step-f16" : "step-f32"; break;
|
||||
case HTP_OP_L2_NORM: op_type = is_f16 ? "l2norm-f16" : "l2norm-f32"; break;
|
||||
case HTP_OP_TRI: op_type = "tri-f32"; break;
|
||||
default:
|
||||
@@ -1506,6 +1540,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
|
||||
case HTP_OP_L2_NORM:
|
||||
case HTP_OP_UNARY_ABS:
|
||||
case HTP_OP_UNARY_LOG:
|
||||
case HTP_OP_UNARY_STEP:
|
||||
break;
|
||||
default:
|
||||
FARF(ERROR, "unary-%s: not supported for F16\n", op_type);
|
||||
@@ -1634,6 +1669,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
|
||||
case HTP_OP_UNARY_ABS: compute_func = (void *) tile_abs_f32; break;
|
||||
case HTP_OP_UNARY_LOG: compute_func = (void *) tile_log_f32; break;
|
||||
case HTP_OP_UNARY_RELU: compute_func = (void *) tile_relu_f32; break;
|
||||
case HTP_OP_UNARY_STEP: compute_func = (void *) tile_step_f32; break;
|
||||
case HTP_OP_TRI:
|
||||
task_func = unary_thread_tiled_tri_f32;
|
||||
compute_func = (void *) tri_apply_tile_f32;
|
||||
@@ -1652,6 +1688,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
|
||||
case HTP_OP_L2_NORM: compute_func = (void *) l2_norm_f16; break;
|
||||
case HTP_OP_UNARY_ABS: compute_func = (void *) abs_f16; break;
|
||||
case HTP_OP_UNARY_LOG: compute_func = (void *) log_f16; break;
|
||||
case HTP_OP_UNARY_STEP: compute_func = (void *) step_f16; break;
|
||||
default: break;
|
||||
}
|
||||
} else {
|
||||
@@ -1678,6 +1715,7 @@ static int execute_op_unary(struct htp_ops_context * octx) {
|
||||
case HTP_OP_UNARY_ABS: compute_func = (void *) abs_f32; break;
|
||||
case HTP_OP_UNARY_LOG: compute_func = (void *) log_f32; break;
|
||||
case HTP_OP_UNARY_RELU: compute_func = (void *) relu_f32; break;
|
||||
case HTP_OP_UNARY_STEP: compute_func = (void *) step_f32; break;
|
||||
case HTP_OP_L2_NORM: compute_func = (void *) l2_norm_f32; break;
|
||||
case HTP_OP_TRI:
|
||||
task_func = unary_thread_tri_f32;
|
||||
|
||||
@@ -59,6 +59,7 @@ static inline bool htp_op_is_unary(uint32_t opcode) {
|
||||
case HTP_OP_UNARY_ABS:
|
||||
case HTP_OP_UNARY_LOG:
|
||||
case HTP_OP_UNARY_RELU:
|
||||
case HTP_OP_UNARY_STEP:
|
||||
case HTP_OP_L2_NORM:
|
||||
case HTP_OP_TRI:
|
||||
return true;
|
||||
|
||||
@@ -2397,22 +2397,22 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_pad(ggml_metal_l
|
||||
char base[256];
|
||||
char name[256];
|
||||
|
||||
// note: this is slower
|
||||
//const bool is_c4 = op->src[0]->ne[0] % 4 == 0 && op->ne[0] % 4 == 0;
|
||||
const bool is_c4 = false;
|
||||
const bool circular = ggml_get_op_params_i32(op, 8) != 0;
|
||||
|
||||
snprintf(base, 256, "kernel_pad_%s%s", ggml_type_name(op->src[0]->type), is_c4 ? "_4" : "");
|
||||
snprintf(name, 256, "%s", base);
|
||||
snprintf(base, 256, "kernel_pad_%s", ggml_type_name(op->src[0]->type));
|
||||
snprintf(name, 256, "%s_circular=%d", base, circular);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
if (res.pipeline) {
|
||||
return res;
|
||||
if (!res.pipeline) {
|
||||
ggml_metal_cv_t cv = ggml_metal_cv_init();
|
||||
|
||||
ggml_metal_cv_set_bool(cv, circular, FC_PAD + 0);
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||
|
||||
ggml_metal_cv_free(cv);
|
||||
}
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, nullptr);
|
||||
|
||||
res.c4 = is_c4;
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
|
||||
@@ -1713,13 +1713,6 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
|
||||
case GGML_OP_POOL_2D:
|
||||
return op->src[0]->type == GGML_TYPE_F32;
|
||||
case GGML_OP_PAD:
|
||||
// TODO: add circular padding support for metal, see https://github.com/ggml-org/llama.cpp/pull/16985
|
||||
if (ggml_get_op_params_i32(op, 8) != 0) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return (ggml_get_op_params_i32(op, 0) == 0) && (ggml_get_op_params_i32(op, 2) == 0) &&
|
||||
(ggml_get_op_params_i32(op, 4) == 0) && (ggml_get_op_params_i32(op, 6) == 0);
|
||||
case GGML_OP_PAD_REFLECT_1D:
|
||||
case GGML_OP_TIMESTEP_EMBEDDING:
|
||||
return op->src[0]->type == GGML_TYPE_F32;
|
||||
|
||||
@@ -418,7 +418,8 @@ static bool ggml_metal_fusion_check_topk_moe(
|
||||
const int64_t n_tokens = logits->ne[1];
|
||||
const int64_t n_expert_used = ids->ne[0];
|
||||
|
||||
if (n_expert <= 0 || n_tokens <= 0 || n_expert_used <= 0 || n_expert_used > n_expert ||
|
||||
// note: n_tokens == 0 (no-output batch) must match so that the packing stays shape-independent
|
||||
if (n_expert <= 0 || n_expert_used <= 0 || n_expert_used > n_expert ||
|
||||
n_expert > GGML_METAL_TOPK_MOE_MAX_EXPERTS || n_expert_used > GGML_METAL_TOPK_MOE_MAX_EXPERTS) {
|
||||
return false;
|
||||
}
|
||||
@@ -545,7 +546,8 @@ static bool ggml_metal_fusion_match_moe_reduce(
|
||||
const int64_t n_embd = experts->ne[0];
|
||||
const int64_t n_tokens = experts->ne[2];
|
||||
|
||||
if (n_embd <= 0 || n_tokens <= 0 || experts->ne[1] != n_expert_used || experts->ne[3] != 1 ||
|
||||
// note: n_tokens == 0 (no-output batch) must match so that the packing stays shape-independent
|
||||
if (n_embd <= 0 || experts->ne[1] != n_expert_used || experts->ne[3] != 1 ||
|
||||
weights->ne[0] != 1 || weights->ne[1] != n_expert_used || weights->ne[2] != n_tokens || weights->ne[3] != 1 ||
|
||||
dst->ne[0] != n_embd || dst->ne[1] != n_tokens || dst->ne[2] != 1 || dst->ne[3] != 1) {
|
||||
return false;
|
||||
@@ -1079,16 +1081,17 @@ const ggml_metal_fusion * ggml_metal_fusion_next(
|
||||
// transparent) node sequence that the compute phase uses, so the returned count is the raw index
|
||||
// span from idx to the last matched node (intermediate views are packed along).
|
||||
int ggml_metal_fusion_max(const ggml_cgraph * gf, int idx) {
|
||||
// an empty/view node cannot start a pattern - pack it alone
|
||||
if (ggml_op_is_empty(gf->nodes[idx]->op) || ggml_is_empty(gf->nodes[idx])) {
|
||||
// a view node cannot start a pattern - pack it alone
|
||||
if (ggml_op_is_empty(gf->nodes[idx]->op)) {
|
||||
return 1;
|
||||
}
|
||||
|
||||
// collect the non-empty node indices starting at idx
|
||||
// collect the non-view node indices starting at idx; 0-element tensors are included so
|
||||
// that empty graphs pack like their non-empty counterparts (see ggml_metal_fusion_filter_ops)
|
||||
int idxs[GGML_METAL_FUSION_MAX];
|
||||
int n_idxs = 0;
|
||||
for (int i = idx; i < gf->n_nodes && n_idxs < GGML_METAL_FUSION_MAX; i++) {
|
||||
if (!ggml_op_is_empty(gf->nodes[i]->op) && !ggml_is_empty(gf->nodes[i])) {
|
||||
if (!ggml_op_is_empty(gf->nodes[i]->op)) {
|
||||
idxs[n_idxs++] = i;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -120,6 +120,7 @@
|
||||
#define FC_TOPK_MOE 1800
|
||||
#define FC_MOE_REDUCE 1900
|
||||
#define FC_DSV4_HC 2000
|
||||
#define FC_PAD 2100
|
||||
|
||||
// op-specific constants
|
||||
#define OP_FLASH_ATTN_EXT_NQPSG 8
|
||||
@@ -1120,6 +1121,10 @@ typedef struct {
|
||||
uint64_t nb1;
|
||||
uint64_t nb2;
|
||||
uint64_t nb3;
|
||||
int32_t lp0;
|
||||
int32_t lp1;
|
||||
int32_t lp2;
|
||||
int32_t lp3;
|
||||
} ggml_metal_kargs_pad;
|
||||
|
||||
typedef struct {
|
||||
@@ -1237,7 +1242,7 @@ typedef struct {
|
||||
|
||||
// widths at or above this use the threadgroup FWHT kernel, one row per threadgroup
|
||||
// with GGML_METAL_FWHT_TG_NT threads, instead of one row per simdgroup
|
||||
#define GGML_METAL_FWHT_TG_MIN_N 1024
|
||||
#define GGML_METAL_FWHT_TG_MIN_N 512
|
||||
#define GGML_METAL_FWHT_TG_NT 256
|
||||
|
||||
typedef struct {
|
||||
|
||||
@@ -4999,16 +4999,15 @@ int ggml_metal_op_pad(ggml_metal_op_t ctx, int idx) {
|
||||
/*.nb0 =*/ nb0,
|
||||
/*.nb1 =*/ nb1,
|
||||
/*.nb2 =*/ nb2,
|
||||
/*.nb3 =*/ nb3
|
||||
/*.nb3 =*/ nb3,
|
||||
/*.lp0 =*/ ggml_get_op_params_i32(op, 0),
|
||||
/*.lp1 =*/ ggml_get_op_params_i32(op, 2),
|
||||
/*.lp2 =*/ ggml_get_op_params_i32(op, 4),
|
||||
/*.lp3 =*/ ggml_get_op_params_i32(op, 6),
|
||||
};
|
||||
|
||||
auto pipeline = ggml_metal_library_get_pipeline_pad(lib, op);
|
||||
|
||||
if (pipeline.c4) {
|
||||
args.ne00 = ne00/4;
|
||||
args.ne0 = ne0/4;
|
||||
}
|
||||
|
||||
const int nth_max = MIN(64, ggml_metal_pipeline_max_theads_per_threadgroup(pipeline));
|
||||
const int nth = MIN(args.ne0, nth_max);
|
||||
const int nk0 = (args.ne0 + 1024 - 1)/1024; // note: 1024 is hardcoded in the kernel!
|
||||
|
||||
@@ -114,8 +114,14 @@ kernel void kernel_roll_f32(
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
kernel void kernel_pad_impl(
|
||||
constant bool FC_pad_circular [[function_constant(FC_PAD + 0)]];
|
||||
|
||||
// circular means on a torus, so the coordinates wrap around
|
||||
static inline int32_t wrap_around(int32_t coord, int32_t size) {
|
||||
return (coord + size) % size;
|
||||
}
|
||||
|
||||
kernel void kernel_pad_f32(
|
||||
constant ggml_metal_kargs_pad & args,
|
||||
device const char * src0,
|
||||
device char * dst,
|
||||
@@ -127,12 +133,40 @@ kernel void kernel_pad_impl(
|
||||
const int32_t k0 = tgpig.x/args.ne1;
|
||||
const int32_t i1 = tgpig.x - k0*args.ne1;
|
||||
|
||||
const int32_t i03 = i3;
|
||||
const int32_t i02 = i2;
|
||||
const int32_t i01 = i1;
|
||||
const int32_t ne00 = args.ne00;
|
||||
const int32_t ne01 = args.ne01;
|
||||
const int32_t ne02 = args.ne02;
|
||||
const int32_t ne03 = args.ne03;
|
||||
|
||||
device const T * src0_ptr = (device const T *) (src0 + i03*args.nb03 + i02*args.nb02 + i01*args.nb01);
|
||||
device T * dst_ptr = (device T *) (dst + i3*args.nb3 + i2*args.nb2 + i1*args.nb1);
|
||||
int32_t i01 = i1 - args.lp1;
|
||||
int32_t i02 = i2 - args.lp2;
|
||||
int32_t i03 = i3 - args.lp3;
|
||||
|
||||
if (FC_pad_circular) {
|
||||
i01 = wrap_around(i01, ne01);
|
||||
i02 = wrap_around(i02, ne02);
|
||||
i03 = wrap_around(i03, ne03);
|
||||
}
|
||||
|
||||
device float * dst_ptr = (device float *) (dst + i3*args.nb3 + i2*args.nb2 + i1*args.nb1);
|
||||
|
||||
// the row lies in the padded region, so no source row backs it
|
||||
if (i01 < 0 || i01 >= ne01 ||
|
||||
i02 < 0 || i02 >= ne02 ||
|
||||
i03 < 0 || i03 >= ne03) {
|
||||
for (int32_t l0 = 0; l0 < 1024; l0 += ntg.x) {
|
||||
const int32_t i0 = k0*1024 + tpitg.x + l0;
|
||||
if (i0 >= args.ne0) {
|
||||
break;
|
||||
}
|
||||
|
||||
dst_ptr[i0] = 0.0f;
|
||||
}
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
device const char * src0_row = src0 + i03*args.nb03 + i02*args.nb02 + i01*args.nb01;
|
||||
|
||||
for (int32_t l0 = 0; l0 < 1024; l0 += ntg.x) {
|
||||
const int32_t i0 = k0*1024 + tpitg.x + l0;
|
||||
@@ -140,19 +174,16 @@ kernel void kernel_pad_impl(
|
||||
break;
|
||||
}
|
||||
|
||||
if (i0 < args.ne00 && i1 < args.ne01 && i2 < args.ne02 && i3 < args.ne03) {
|
||||
dst_ptr[i0] = src0_ptr[i0];
|
||||
} else {
|
||||
dst_ptr[i0] = 0.0f;
|
||||
int32_t i00 = i0 - args.lp0;
|
||||
|
||||
if (FC_pad_circular) {
|
||||
i00 = wrap_around(i00, ne00);
|
||||
}
|
||||
|
||||
dst_ptr[i0] = i00 >= 0 && i00 < ne00 ? *((device const float *) (src0_row + i00*args.nb00)) : 0.0f;
|
||||
}
|
||||
}
|
||||
|
||||
typedef decltype(kernel_pad_impl<float>) kernel_pad_t;
|
||||
|
||||
template [[host_name("kernel_pad_f32")]] kernel kernel_pad_t kernel_pad_impl<float>;
|
||||
template [[host_name("kernel_pad_f32_4")]] kernel kernel_pad_t kernel_pad_impl<float4>;
|
||||
|
||||
// TODO: this is slow - optimize
|
||||
kernel void kernel_pad_reflect_1d_f32(
|
||||
constant ggml_metal_kargs_pad_reflect_1d & args,
|
||||
@@ -510,18 +541,18 @@ typedef decltype(kernel_fwht<64, half>) kernel_fwht_f16_t;
|
||||
template [[host_name("kernel_fwht_f32_64")]] kernel kernel_fwht_f32_t kernel_fwht<64, float>;
|
||||
template [[host_name("kernel_fwht_f32_128")]] kernel kernel_fwht_f32_t kernel_fwht<128, float>;
|
||||
template [[host_name("kernel_fwht_f32_256")]] kernel kernel_fwht_f32_t kernel_fwht<256, float>;
|
||||
template [[host_name("kernel_fwht_f32_512")]] kernel kernel_fwht_f32_t kernel_fwht<512, float>;
|
||||
|
||||
template [[host_name("kernel_fwht_f16_64")]] kernel kernel_fwht_f16_t kernel_fwht<64, half>;
|
||||
template [[host_name("kernel_fwht_f16_128")]] kernel kernel_fwht_f16_t kernel_fwht<128, half>;
|
||||
template [[host_name("kernel_fwht_f16_256")]] kernel kernel_fwht_f16_t kernel_fwht<256, half>;
|
||||
template [[host_name("kernel_fwht_f16_512")]] kernel kernel_fwht_f16_t kernel_fwht<512, half>;
|
||||
|
||||
template [[host_name("kernel_fwht_f32_512")]] kernel kernel_fwht_f32_t kernel_fwht_tg<512, GGML_METAL_FWHT_TG_NT, float>;
|
||||
template [[host_name("kernel_fwht_f32_1024")]] kernel kernel_fwht_f32_t kernel_fwht_tg<1024, GGML_METAL_FWHT_TG_NT, float>;
|
||||
template [[host_name("kernel_fwht_f32_2048")]] kernel kernel_fwht_f32_t kernel_fwht_tg<2048, GGML_METAL_FWHT_TG_NT, float>;
|
||||
template [[host_name("kernel_fwht_f32_4096")]] kernel kernel_fwht_f32_t kernel_fwht_tg<4096, GGML_METAL_FWHT_TG_NT, float>;
|
||||
template [[host_name("kernel_fwht_f32_8192")]] kernel kernel_fwht_f32_t kernel_fwht_tg<8192, GGML_METAL_FWHT_TG_NT, float>;
|
||||
|
||||
template [[host_name("kernel_fwht_f16_512")]] kernel kernel_fwht_f16_t kernel_fwht_tg<512, GGML_METAL_FWHT_TG_NT, half>;
|
||||
template [[host_name("kernel_fwht_f16_1024")]] kernel kernel_fwht_f16_t kernel_fwht_tg<1024, GGML_METAL_FWHT_TG_NT, half>;
|
||||
template [[host_name("kernel_fwht_f16_2048")]] kernel kernel_fwht_f16_t kernel_fwht_tg<2048, GGML_METAL_FWHT_TG_NT, half>;
|
||||
template [[host_name("kernel_fwht_f16_4096")]] kernel kernel_fwht_f16_t kernel_fwht_tg<4096, GGML_METAL_FWHT_TG_NT, half>;
|
||||
|
||||
@@ -1186,7 +1186,7 @@ struct ggml_backend_opencl_context {
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
size_t sz;
|
||||
size_t sz = 0;
|
||||
const void * kernel_bin = get_adreno_bin_kernel_func(
|
||||
kernel_name.c_str(), device_name.c_str(), driver_version.c_str(), &sz);
|
||||
if (bin_size) {
|
||||
@@ -3881,7 +3881,7 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
|
||||
backend_ctx->kernel_gemv_noshuffle_q4_0_f32_32b_trans = nullptr;
|
||||
backend_ctx->kernel_gemm_noshuffle_q4_0_f32_32b_trans_ila_a8_bin = nullptr;
|
||||
backend_ctx->kernel_gemm_noshuffle_q4_0_q8_1_dp4a_ila_a8_bin = nullptr;
|
||||
if (backend_ctx->adreno_gen == ADRENO_GPU_GEN::X2E) {
|
||||
{
|
||||
{
|
||||
std::string opts = std::string("-cl-std=") + opencl_c_std +
|
||||
" -cl-mad-enable "
|
||||
@@ -4381,8 +4381,8 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
|
||||
backend_ctx->kernel_gemv_noshuffle_q4_k_f32_32b_trans = nullptr;
|
||||
backend_ctx->kernel_gemm_noshuffle_q4_k_f32_32b_trans_ila_a8_bin = nullptr;
|
||||
backend_ctx->kernel_gemm_noshuffle_q4_k_q8_1_dp4a_ila_a8_bin = nullptr;
|
||||
if (backend_ctx->adreno_gen == ADRENO_GPU_GEN::X2E) {
|
||||
{
|
||||
{
|
||||
if (backend_ctx->has_vector_subgroup_broadcast) {
|
||||
std::string opts = std::string("-cl-std=") + opencl_c_std +
|
||||
" -cl-mad-enable "
|
||||
" -DSIMDGROUP_WIDTH=" +
|
||||
@@ -4430,8 +4430,8 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
|
||||
backend_ctx->kernel_gemv_noshuffle_q6_k_f32_32b_trans = nullptr;
|
||||
backend_ctx->kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8_bin = nullptr;
|
||||
backend_ctx->kernel_gemm_noshuffle_q6_k_q8_1_dp4a_ila_a8_bin = nullptr;
|
||||
if (backend_ctx->adreno_gen == ADRENO_GPU_GEN::X2E) {
|
||||
{
|
||||
{
|
||||
if (backend_ctx->has_vector_subgroup_broadcast) {
|
||||
std::string opts = std::string("-cl-std=") + opencl_c_std +
|
||||
" -cl-mad-enable "
|
||||
" -DSIMDGROUP_WIDTH=" +
|
||||
@@ -4479,8 +4479,8 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
|
||||
backend_ctx->kernel_gemv_noshuffle_q5_k_f32_32b_trans = nullptr;
|
||||
backend_ctx->kernel_gemm_noshuffle_q5_k_f32_32b_trans_ila_a8_bin = nullptr;
|
||||
backend_ctx->kernel_gemm_noshuffle_q5_k_q8_1_dp4a_ila_a8_bin = nullptr;
|
||||
if (backend_ctx->adreno_gen == ADRENO_GPU_GEN::X2E) {
|
||||
{
|
||||
{
|
||||
if (backend_ctx->has_vector_subgroup_broadcast) {
|
||||
std::string opts = std::string("-cl-std=") + opencl_c_std +
|
||||
" -cl-mad-enable "
|
||||
" -DSIMDGROUP_WIDTH=" +
|
||||
|
||||
@@ -9,6 +9,7 @@
|
||||
#include "ggml-quants.h"
|
||||
#include "ggml.h"
|
||||
|
||||
#include <algorithm>
|
||||
#include <atomic>
|
||||
#include <cerrno>
|
||||
#include <climits>
|
||||
@@ -960,6 +961,24 @@ static bool has_view_op_input(const ggml_tensor * op) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// OV slices whole elements per axis, so each stride must be a multiple of the next smaller one
|
||||
// (e.g. a batch stride of m*nb[1] + pad bytes cannot be expressed and would be read wrongly).
|
||||
static bool has_strides_on_element_grid(const ggml_tensor * t) {
|
||||
std::vector<size_t> strides;
|
||||
for (int i = 0; i < GGML_MAX_DIMS; i++) {
|
||||
if (t->ne[i] > 1) {
|
||||
strides.push_back(t->nb[i]);
|
||||
}
|
||||
}
|
||||
std::sort(strides.begin(), strides.end());
|
||||
for (size_t i = 1; i < strides.size(); i++) {
|
||||
if (strides[i - 1] == 0 || strides[i] % strides[i - 1] != 0) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
static bool has_non_contiguous_view_input(const ggml_tensor * op) {
|
||||
for (int i = 0; i < GGML_MAX_SRC; i++) {
|
||||
if (op->src[i] == nullptr) {
|
||||
@@ -1540,6 +1559,9 @@ static ggml_openvino_op_support ggml_backend_openvino_device_supports_op_impl(gg
|
||||
if (supported_types.find(src->type) == supported_types.end()) {
|
||||
return {false, "src[" + std::to_string(i) + "] type " + std::string(ggml_type_name(src->type)) + " is not supported"};
|
||||
}
|
||||
if (!has_strides_on_element_grid(src)) {
|
||||
return {false, "src[" + std::to_string(i) + "] strides are not multiples of each other"};
|
||||
}
|
||||
const bool is_supported_3d_moe_expert =
|
||||
op->op == GGML_OP_MUL_MAT_ID && i == 0 && (src->type == GGML_TYPE_MXFP4 || src->ne[3] == 1);
|
||||
if (ggml_is_quantized(src->type) && src->ne[2] != 1 && !is_supported_3d_moe_expert) {
|
||||
|
||||
@@ -26,6 +26,7 @@
|
||||
// so every SEND posts a whole 128KiB stride over the wire, even when partially filled.
|
||||
// (In testing 128KiB was the best performing among 32, 64, 128, 256)
|
||||
|
||||
//TODO: add mechanism similar to https://github.com/ggml-org/llama.cpp/pull/29440 to prevent idle CPU from spinning
|
||||
static constexpr uint32_t RDMA_SEG_MAGIC = 0x52534547u; // "RSEG"
|
||||
static constexpr int RDMA_NBUF = 16; // ring depth (frames per direction)
|
||||
static constexpr size_t RDMA_FRAME = 4096; // Thunderbolt frame (fixed on Apple)
|
||||
|
||||
@@ -25,6 +25,8 @@
|
||||
#ifdef GGML_RPC_RDMA
|
||||
# include <infiniband/verbs.h>
|
||||
# include <array>
|
||||
# include <cerrno>
|
||||
# include <chrono>
|
||||
# include <time.h>
|
||||
# ifndef _WIN32
|
||||
# include <poll.h>
|
||||
@@ -54,6 +56,8 @@ using rdma_gid_t = std::array<uint8_t, RDMA_GID_SIZE>;
|
||||
#if defined(GGML_RPC_RDMA) && !defined(GGML_RPC_RDMA_APPLE)
|
||||
static constexpr size_t RDMA_CHUNK = 256 * 1024; // 256 KiB per send/recv (fits default 8 MiB memlock)
|
||||
static constexpr int RDMA_RX_DEPTH = 24; // pre-posted recv ring: 24 × 256 KiB = 6 MiB
|
||||
// keep polling the CQ for this long after the last activity, then sleep until the next completion
|
||||
static constexpr auto RDMA_SPIN_TIME = std::chrono::milliseconds(100);
|
||||
|
||||
struct rdma_conn {
|
||||
struct ibv_context * ctx = nullptr;
|
||||
@@ -61,6 +65,9 @@ struct rdma_conn {
|
||||
struct ibv_cq * scq = nullptr; // send completions
|
||||
struct ibv_cq * rcq = nullptr; // recv completions
|
||||
struct ibv_qp * qp = nullptr;
|
||||
struct ibv_comp_channel * ch = nullptr; // CQ events, so an idle connection can sleep instead of spinning
|
||||
|
||||
std::chrono::steady_clock::time_point last_active; // last completion or posted send
|
||||
|
||||
void * tx_buf = nullptr;
|
||||
struct ibv_mr * tx_mr = nullptr;
|
||||
@@ -95,6 +102,7 @@ struct rdma_conn {
|
||||
if (qp) ibv_destroy_qp(qp);
|
||||
if (scq) ibv_destroy_cq(scq);
|
||||
if (rcq) ibv_destroy_cq(rcq);
|
||||
if (ch) ibv_destroy_comp_channel(ch);
|
||||
if (pd) ibv_dealloc_pd(pd);
|
||||
if (ctx) ibv_close_device(ctx);
|
||||
}
|
||||
@@ -142,6 +150,7 @@ struct socket_t::impl {
|
||||
bool tcp_peer_closed();
|
||||
bool rdma_activate(uint32_t remote_qpn, uint32_t remote_psn, const uint8_t * remote_gid);
|
||||
bool rdma_poll(struct ibv_cq * cq, struct ibv_wc * wc);
|
||||
bool rdma_wait_event();
|
||||
|
||||
std::unique_ptr<rdma_conn> rdma;
|
||||
rdma_local_info rdma_local = {};
|
||||
@@ -291,8 +300,10 @@ bool socket_t::impl::rdma_probe() {
|
||||
rdma->pd = ibv_alloc_pd(ibctx);
|
||||
if (!rdma->pd) return false;
|
||||
|
||||
rdma->scq = ibv_create_cq(ibctx, 16, nullptr, nullptr, 0);
|
||||
rdma->rcq = ibv_create_cq(ibctx, RDMA_RX_DEPTH + 4, nullptr, nullptr, 0);
|
||||
// without a completion channel rdma_poll() spins all the time, as before
|
||||
rdma->ch = ibv_create_comp_channel(ibctx);
|
||||
rdma->scq = ibv_create_cq(ibctx, 16, nullptr, rdma->ch, 0);
|
||||
rdma->rcq = ibv_create_cq(ibctx, RDMA_RX_DEPTH + 4, nullptr, rdma->ch, 0);
|
||||
if (!rdma->scq || !rdma->rcq) return false;
|
||||
|
||||
ibv_qp_init_attr qia = {};
|
||||
@@ -394,15 +405,45 @@ bool socket_t::impl::rdma_activate(uint32_t remote_qpn, uint32_t remote_psn, con
|
||||
}
|
||||
}
|
||||
|
||||
rdma->last_active = std::chrono::steady_clock::now();
|
||||
|
||||
GGML_LOG_INFO("RDMA activated: qpn=%u->%u mtu=%d rx_depth=%d\n",
|
||||
rdma_local.qpn, remote_qpn, 128 << rdma_local.path_mtu, RDMA_RX_DEPTH);
|
||||
return true;
|
||||
}
|
||||
|
||||
// Sleep until the completion channel has an event or the TCP peer closes.
|
||||
bool socket_t::impl::rdma_wait_event() {
|
||||
rdma_conn * c = rdma.get();
|
||||
// POLLHUP and POLLERR are always reported, the TCP socket carries no data after the RDMA upgrade
|
||||
struct pollfd pfds[2] = {
|
||||
{ c->ch->fd, POLLIN, 0 },
|
||||
{ fd, POLLRDHUP, 0 },
|
||||
};
|
||||
if (poll(pfds, 2, -1) < 0) {
|
||||
return errno == EINTR;
|
||||
}
|
||||
if (pfds[1].revents & (POLLHUP | POLLERR | POLLRDHUP)) {
|
||||
return false;
|
||||
}
|
||||
if (pfds[0].revents & POLLIN) {
|
||||
struct ibv_cq * ev_cq = nullptr;
|
||||
void * ev_ctx = nullptr;
|
||||
if (ibv_get_cq_event(c->ch, &ev_cq, &ev_ctx) != 0) {
|
||||
return false;
|
||||
}
|
||||
ibv_ack_cq_events(ev_cq, 1);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool socket_t::impl::rdma_poll(struct ibv_cq * cq, struct ibv_wc * wc) {
|
||||
for (uint64_t s = 0; ; s++) {
|
||||
rdma_conn * c = rdma.get();
|
||||
bool armed = false;
|
||||
for (uint64_t s = 1; ; s++) {
|
||||
int n = ibv_poll_cq(cq, 1, wc);
|
||||
if (n > 0) {
|
||||
c->last_active = std::chrono::steady_clock::now();
|
||||
if (wc->status != IBV_WC_SUCCESS) {
|
||||
GGML_LOG_ERROR("RDMA CQ wc error: status=%d (%s) vendor_err=0x%x\n",
|
||||
wc->status, ibv_wc_status_str(wc->status), wc->vendor_err);
|
||||
@@ -410,7 +451,24 @@ bool socket_t::impl::rdma_poll(struct ibv_cq * cq, struct ibv_wc * wc) {
|
||||
return wc->status == IBV_WC_SUCCESS;
|
||||
}
|
||||
if (n < 0) return false;
|
||||
if ((s & 0xFFFFF) == 0 && s > 0) {
|
||||
if (armed) {
|
||||
// armed and still empty: sleep until the next completion
|
||||
if (!rdma_wait_event()) {
|
||||
return false;
|
||||
}
|
||||
armed = false;
|
||||
continue;
|
||||
}
|
||||
// spin while the connection is busy, arm the CQ once it has been idle for RDMA_SPIN_TIME
|
||||
// a completion that arrives before arming raises no event, so poll once more after arming
|
||||
if (c->ch && (s & 0x3FF) == 0 && std::chrono::steady_clock::now() - c->last_active > RDMA_SPIN_TIME) {
|
||||
if (ibv_req_notify_cq(cq, 0) != 0) {
|
||||
return false;
|
||||
}
|
||||
armed = true;
|
||||
continue;
|
||||
}
|
||||
if ((s & 0xFFFFF) == 0) {
|
||||
if (tcp_peer_closed()) {
|
||||
return false;
|
||||
}
|
||||
@@ -444,6 +502,7 @@ bool socket_t::impl::rdma_send(const void * data, size_t size) {
|
||||
}
|
||||
|
||||
if (ibv_post_send(c->qp, &wr, &bad) != 0) return false;
|
||||
c->last_active = std::chrono::steady_clock::now();
|
||||
struct ibv_wc wc;
|
||||
if (!rdma_poll(c->scq, &wc)) return false;
|
||||
|
||||
|
||||
@@ -124,6 +124,107 @@ static void launch_fwht(const float * src, float * dst, const int64_t n_rows, co
|
||||
});
|
||||
}
|
||||
|
||||
// Wide blocks: one row per work-group instead of per sub-group, so each work-item
|
||||
// keeps N/NT values rather than N/WARP_SIZE. Butterflies below the sub-group width
|
||||
// still shuffle; those up to NT go through work-group local memory; the rest stay
|
||||
// in registers.
|
||||
template <int N, int NT>
|
||||
static void fwht_kernel_wide(const float * __restrict__ src,
|
||||
float * __restrict__ dst,
|
||||
const int64_t n_rows,
|
||||
const float scale,
|
||||
const sycl::nd_item<2> & item,
|
||||
float * smem) {
|
||||
const int64_t r = item.get_global_id(0);
|
||||
if (r >= n_rows) {
|
||||
return;
|
||||
}
|
||||
|
||||
src += r * N;
|
||||
dst += r * N;
|
||||
|
||||
constexpr int el_w = N / NT;
|
||||
static_assert(el_w >= 1 && N % NT == 0, "row must be a whole number of work-group widths");
|
||||
|
||||
const int tid = item.get_local_id(1);
|
||||
|
||||
float reg[el_w];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < el_w; ++i) {
|
||||
reg[i] = src[i * NT + tid] * scale;
|
||||
}
|
||||
|
||||
const sycl::sub_group sg = item.get_sub_group();
|
||||
const int lane = sg.get_local_linear_id();
|
||||
|
||||
// Butterflies inside the sub-group, same pattern as the narrow kernel.
|
||||
#pragma unroll
|
||||
for (int h = 1; h < WARP_SIZE; h *= 2) {
|
||||
#pragma unroll
|
||||
for (int j = 0; j < el_w; ++j) {
|
||||
const float val = reg[j];
|
||||
const float val2 = dpct::permute_sub_group_by_xor(sg, val, h, WARP_SIZE);
|
||||
|
||||
reg[j] = (lane & h) == 0 ? val + val2 : val2 - val;
|
||||
}
|
||||
}
|
||||
|
||||
// Butterflies from the sub-group width up to NT: the partner lane is outside
|
||||
// this sub-group, so it goes through work-group local memory instead of a shuffle.
|
||||
for (int h = WARP_SIZE; h < NT; h *= 2) {
|
||||
#pragma unroll
|
||||
for (int j = 0; j < el_w; ++j) {
|
||||
smem[j * NT + tid] = reg[j];
|
||||
}
|
||||
item.barrier(sycl::access::fence_space::local_space);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < el_w; ++j) {
|
||||
const float val = reg[j];
|
||||
const float val2 = smem[j * NT + (tid ^ h)];
|
||||
reg[j] = (tid & h) == 0 ? val + val2 : val2 - val;
|
||||
}
|
||||
item.barrier(sycl::access::fence_space::local_space);
|
||||
}
|
||||
|
||||
// Butterflies across registers: h is a multiple of NT, so the partner of element
|
||||
// i*NT + tid lives in reg[i + h/NT] on the same work-item.
|
||||
for (int h = NT; h < N; h *= 2) {
|
||||
const int step = h / NT;
|
||||
for (int j = 0; j < el_w; j += 2 * step) {
|
||||
for (int k = 0; k < step; ++k) {
|
||||
const float x = reg[j + k];
|
||||
const float y = reg[j + k + step];
|
||||
|
||||
reg[j + k] = x + y;
|
||||
reg[j + k + step] = x - y;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < el_w; ++i) {
|
||||
dst[i * NT + tid] = reg[i];
|
||||
}
|
||||
}
|
||||
|
||||
template <int N, int NT>
|
||||
static void launch_fwht_wide(const float * src,
|
||||
float * dst,
|
||||
const int64_t n_rows,
|
||||
const float scale,
|
||||
dpct::queue_ptr stream) {
|
||||
const sycl::range<2> global(n_rows, NT);
|
||||
const sycl::range<2> local(1, NT);
|
||||
|
||||
stream->submit([&](sycl::handler & cgh) {
|
||||
sycl::local_accessor<float, 1> smem(sycl::range<1>(N), cgh);
|
||||
cgh.parallel_for(sycl::nd_range<2>(global, local),
|
||||
[=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
||||
fwht_kernel_wide<N, NT>(src, dst, n_rows, scale, item, get_pointer(smem));
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
template <int N, int m>
|
||||
static void kronecker_kernel(const float * __restrict__ src,
|
||||
float * __restrict__ dst,
|
||||
@@ -285,6 +386,18 @@ bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src,
|
||||
case 1280:
|
||||
launch_kronecker<1280, 20>(src_d, dst_d, rows, scale, stream);
|
||||
return true;
|
||||
case 1024:
|
||||
launch_fwht_wide<1024, 256>(src_d, dst_d, rows, scale, stream);
|
||||
return true;
|
||||
case 2048:
|
||||
launch_fwht_wide<2048, 256>(src_d, dst_d, rows, scale, stream);
|
||||
return true;
|
||||
case 4096:
|
||||
launch_fwht_wide<4096, 256>(src_d, dst_d, rows, scale, stream);
|
||||
return true;
|
||||
case 8192:
|
||||
launch_fwht_wide<8192, 256>(src_d, dst_d, rows, scale, stream);
|
||||
return true;
|
||||
default:
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -97,7 +97,7 @@ vk_pipeline ggml_vk_get_quantize_pipeline(ggml_backend_vk_context * ctx, ggml_ty
|
||||
void ggml_vk_quantize_q8_1(ggml_backend_vk_context * ctx, vk_context& subctx, const vk_subbuffer & in, const vk_subbuffer & out, uint32_t ne);
|
||||
void ggml_vk_dsv4_hc_comb(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * mixes, const ggml_tensor * scale, const ggml_tensor * base, ggml_tensor * dst);
|
||||
void ggml_vk_dsv4_hc_pre(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * weights, ggml_tensor * dst);
|
||||
void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * residual, const ggml_tensor * post, const ggml_tensor * comb, ggml_tensor * dst);
|
||||
void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * residual, const ggml_tensor * post, const ggml_tensor * comb, ggml_tensor * dst, const ggml_tensor * gate_scale_in = nullptr);
|
||||
void ggml_vk_mul_mat(ggml_backend_vk_context * ctx, vk_context& subctx, const struct ggml_cgraph * cgraph, int node_idx);
|
||||
bool ggml_vk_use_mul_mat_vec_id(const struct ggml_cgraph * cgraph, int node_idx);
|
||||
void ggml_vk_mul_mat_id(ggml_backend_vk_context * ctx, vk_context& subctx, const struct ggml_cgraph * cgraph, int node_idx);
|
||||
@@ -263,9 +263,36 @@ inline void ggml_vk_dispatch_pipeline(ggml_backend_vk_context* ctx, vk_context&
|
||||
GGML_ASSERT(pipeline->parameter_count == descriptor_buffer_infos.size());
|
||||
GGML_ASSERT(pipeline->push_constant_size == push_constant_size(push_constants));
|
||||
|
||||
vk::DescriptorSet& descriptor_set = ctx->descriptor_sets[ctx->descriptor_set_idx++];
|
||||
vk::WriteDescriptorSet write_descriptor_set{ descriptor_set, 0, 0, pipeline->parameter_count, vk::DescriptorType::eStorageBuffer, nullptr, descriptor_buffer_infos.begin() };
|
||||
ctx->device->device.updateDescriptorSets({ write_descriptor_set }, {});
|
||||
const uint32_t descriptor_set_idx = ctx->descriptor_set_idx++;
|
||||
vk::DescriptorSet& descriptor_set = ctx->descriptor_sets[descriptor_set_idx];
|
||||
|
||||
// a new buffer can get the handle of a destroyed one, so drop all cached bindings after any destroy
|
||||
const uint64_t destroy_count = ctx->device->buffer_destroy_count.load(std::memory_order_acquire);
|
||||
if (ctx->descriptor_set_bindings_destroy_count != destroy_count) {
|
||||
for (auto & b : ctx->descriptor_set_bindings) {
|
||||
b.clear();
|
||||
}
|
||||
ctx->descriptor_set_bindings_destroy_count = destroy_count;
|
||||
}
|
||||
|
||||
// skip the write if this set already holds these bindings from the last graph
|
||||
std::vector<vk::DescriptorBufferInfo> & bindings = ctx->descriptor_set_bindings[descriptor_set_idx];
|
||||
bool same = !ctx->device->disable_descriptor_reuse && bindings.size() == descriptor_buffer_infos.size();
|
||||
if (same) {
|
||||
size_t i = 0;
|
||||
for (const vk::DescriptorBufferInfo & info : descriptor_buffer_infos) {
|
||||
const vk::DescriptorBufferInfo & prev = bindings[i++];
|
||||
if (prev.buffer != info.buffer || prev.offset != info.offset || prev.range != info.range) {
|
||||
same = false;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
if (!same) {
|
||||
vk::WriteDescriptorSet write_descriptor_set{ descriptor_set, 0, 0, pipeline->parameter_count, vk::DescriptorType::eStorageBuffer, nullptr, descriptor_buffer_infos.begin() };
|
||||
ctx->device->device.updateDescriptorSets({ write_descriptor_set }, {});
|
||||
bindings.assign(descriptor_buffer_infos.begin(), descriptor_buffer_infos.end());
|
||||
}
|
||||
|
||||
subctx->s->buffer->buf.pushConstants(pipeline->layout, vk::ShaderStageFlagBits::eCompute, 0, push_constant_size(push_constants), push_constant_data(push_constants));
|
||||
subctx->s->buffer->buf.bindPipeline(vk::PipelineBindPoint::eCompute, pipeline->pipeline);
|
||||
|
||||
@@ -205,6 +205,10 @@ struct vk_op_dsv4_hc_post_push_constants {
|
||||
uint32_t p_offset;
|
||||
uint32_t c_offset;
|
||||
uint32_t d_offset;
|
||||
|
||||
uint32_t gate;
|
||||
float gate_scale_in;
|
||||
float gate_scale_out;
|
||||
};
|
||||
|
||||
struct vk_op_count_experts_push_constants {
|
||||
|
||||
@@ -51,6 +51,8 @@ typedef struct VkPhysicalDeviceCooperativeMatrixDecodeVectorFeaturesNV {
|
||||
|
||||
#include <cmath>
|
||||
|
||||
#include <functional>
|
||||
|
||||
#include <iomanip>
|
||||
|
||||
#include <iostream>
|
||||
@@ -553,6 +555,15 @@ static constexpr std::initializer_list<ggml_op> rms_norm_view_set_rows_pattern {
|
||||
|
||||
static constexpr std::initializer_list<ggml_op> rope_view_set_rows_pattern { GGML_OP_ROPE, GGML_OP_VIEW, GGML_OP_SET_ROWS };
|
||||
|
||||
// scale_out*sigmoid(scale_in*x) as the hc_post weights (qwen4exp hc_combine)
|
||||
static constexpr std::initializer_list<ggml_op> hc_post_gate_pattern { GGML_OP_SCALE, GGML_OP_UNARY, GGML_OP_SCALE, GGML_OP_DSV4_HC_POST };
|
||||
|
||||
static constexpr std::initializer_list<std::array<int, 3>> hc_post_gate_edges {
|
||||
{ 1, 0, 0 }, // sigmoid->src[0] == scale
|
||||
{ 2, 0, 1 }, // scale->src[0] == sigmoid
|
||||
{ 3, 2, 2 }, // hc_post->src[2] == scale (post)
|
||||
};
|
||||
|
||||
static constexpr std::initializer_list<std::array<int, 3>> topk_moe_early_softmax_norm_edges {
|
||||
{ 1, 0, 0 }, // reshape->src[0] == softmax
|
||||
{ 2, 0, 0 }, // argsort->src[0] == softmax
|
||||
@@ -1014,6 +1025,8 @@ struct vk_device_struct {
|
||||
ggml_backend_buffer_type buffer_type;
|
||||
|
||||
bool disable_fusion;
|
||||
bool disable_descriptor_reuse;
|
||||
std::atomic<uint64_t> buffer_destroy_count {};
|
||||
bool disable_host_visible_vidmem;
|
||||
bool allow_sysmem_fallback;
|
||||
bool disable_graph_optimize;
|
||||
@@ -1056,6 +1069,8 @@ struct vk_buffer_struct {
|
||||
}
|
||||
VK_LOG_DEBUG("~vk_buffer_struct(" << buffer << ", " << size << ")");
|
||||
|
||||
// bump before destroying, so a thread that sees the buffer gone also sees the new count
|
||||
device->buffer_destroy_count.fetch_add(1, std::memory_order_release);
|
||||
device->device.freeMemory(device_memory);
|
||||
device->device.destroyBuffer(buffer);
|
||||
}
|
||||
@@ -1267,6 +1282,9 @@ struct ggml_backend_vk_context {
|
||||
|
||||
std::vector<vk::DescriptorPool> descriptor_pools;
|
||||
std::vector<vk::DescriptorSet> descriptor_sets;
|
||||
// last bindings written to each set; descriptor_sets is append-only so an index always names the same set
|
||||
std::vector<std::vector<vk::DescriptorBufferInfo>> descriptor_set_bindings;
|
||||
uint64_t descriptor_set_bindings_destroy_count {};
|
||||
uint32_t descriptor_set_idx {};
|
||||
uint32_t pipeline_descriptor_set_requirements {};
|
||||
|
||||
@@ -1284,6 +1302,7 @@ struct ggml_backend_vk_context {
|
||||
bool fused_topk_moe_scale {};
|
||||
// QSA indexer gather+add+top_k fused into one radix-select
|
||||
bool fused_topk_qsa {};
|
||||
bool fused_hc_post_gate {};
|
||||
rms_norm_mode fused_rms_norm_mode {RMS_NORM_COUNT};
|
||||
|
||||
// for GGML_VK_PERF_LOGGER
|
||||
|
||||
@@ -881,6 +881,7 @@ void ggml_pipeline_allocate_descriptor_sets(ggml_backend_vk_context * ctx) {
|
||||
|
||||
pool_idx++;
|
||||
}
|
||||
ctx->descriptor_set_bindings.resize(ctx->descriptor_sets.size());
|
||||
}
|
||||
|
||||
static vk_command_buffer* ggml_vk_create_cmd_buffer(vk_device& device, vk_command_pool& p) {
|
||||
@@ -4885,6 +4886,8 @@ vk_device ggml_vk_get_device(size_t idx) {
|
||||
|
||||
device->disable_fusion = getenv("GGML_VK_DISABLE_FUSION") != nullptr;
|
||||
|
||||
device->disable_descriptor_reuse = getenv("GGML_VK_DISABLE_DESCRIPTOR_REUSE") != nullptr;
|
||||
|
||||
device->add_rms_fusion = !device->disable_fusion &&
|
||||
device->subgroup_arithmetic &&
|
||||
device->vendor_id != VK_VENDOR_ID_INTEL;
|
||||
@@ -6001,6 +6004,11 @@ bool ggml_vk_dim01_contiguous(const ggml_tensor * tensor) {
|
||||
(tensor->ne[3] == 1 || tensor->nb[3] == tensor->nb[2]*tensor->ne[2]);
|
||||
}
|
||||
|
||||
// Batch stride in elements of a tensor read in place.
|
||||
static uint32_t ggml_vk_batch_stride(const ggml_tensor * tensor) {
|
||||
return (uint32_t)(tensor->nb[2] / ggml_type_size(tensor->type) * ggml_blck_size(tensor->type));
|
||||
}
|
||||
|
||||
vk_pipeline ggml_vk_get_cpy_pipeline(ggml_backend_vk_context * ctx, const ggml_tensor * src, const ggml_tensor * dst, ggml_type to) {
|
||||
|
||||
// Choose "contiguous copy" shader if src/dst are contiguous
|
||||
@@ -6488,8 +6496,13 @@ static void ggml_vk_mul_mat_q_f16(ggml_backend_vk_context * ctx, vk_context& sub
|
||||
}
|
||||
}
|
||||
|
||||
uint32_t stride_batch_x = ne00*ne01;
|
||||
uint32_t stride_batch_y = ne10*ne11;
|
||||
// The partial tiles of a quantized A rely on the bound range to read zeros past the last row,
|
||||
// so the range stays exact: the strided extent of a tensor read in place, the staged size otherwise.
|
||||
const uint64_t x_range = qx_needs_dequant ? x_sz : ggml_nbytes(src0);
|
||||
const uint64_t y_range = (qy_needs_dequant || quantize_y) ? y_sz : ggml_nbytes(src1);
|
||||
|
||||
uint32_t stride_batch_x = qx_needs_dequant ? ne00*ne01 : ggml_vk_batch_stride(src0);
|
||||
uint32_t stride_batch_y = (qy_needs_dequant || quantize_y) ? ne10*ne11 : ggml_vk_batch_stride(src1);
|
||||
|
||||
if (!ggml_vk_dim01_contiguous(src0) && !qx_needs_dequant) {
|
||||
stride_batch_x = src0->nb[0] / ggml_type_size(src0->type);
|
||||
@@ -6502,7 +6515,7 @@ static void ggml_vk_mul_mat_q_f16(ggml_backend_vk_context * ctx, vk_context& sub
|
||||
// compute
|
||||
ggml_vk_matmul(
|
||||
ctx, subctx, pipeline,
|
||||
{ d_X, x_buf_offset, x_sz }, { d_Y, y_buf_offset, y_sz },
|
||||
{ d_X, x_buf_offset, x_range }, { d_Y, y_buf_offset, y_range },
|
||||
ggml_vk_subbuffer(ctx, d_D, d_buf_offset), { ctx->prealloc_split_k, 0, d_sz * split_k },
|
||||
ne01, ne11, ne10,
|
||||
ne10, ne10, stride_d, stride_batch_x, stride_batch_y, stride_batch_d,
|
||||
@@ -6770,8 +6783,8 @@ static void ggml_vk_mul_mat_vec_q_f16(ggml_backend_vk_context * ctx, vk_context&
|
||||
}
|
||||
|
||||
// For batch_n, the A matrix is the same for each batch, and B/D use the row stride as the batch stride
|
||||
uint32_t stride_batch_x = batch_n ? 0 : ne00*ne01;
|
||||
uint32_t stride_batch_y = batch_n ? ne10 : (ne10*ne11);
|
||||
uint32_t stride_batch_x = batch_n ? 0 : (qx_needs_dequant ? ne00*ne01 : ggml_vk_batch_stride(src0));
|
||||
uint32_t stride_batch_y = batch_n ? ne10 : ((qy_needs_dequant || quantize_y) ? ne10*ne11 : ggml_vk_batch_stride(src1));
|
||||
uint32_t stride_batch_d = batch_n ? ne20 : (ne20*ne21);
|
||||
|
||||
if (!ggml_vk_dim01_contiguous(src0) && !qx_needs_dequant) {
|
||||
@@ -7157,7 +7170,7 @@ void ggml_vk_dsv4_hc_pre(ggml_backend_vk_context * ctx, vk_context& subctx, cons
|
||||
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { x_buf, w_buf, d_buf }, pc, { n_embd, n_tokens, 1 });
|
||||
}
|
||||
|
||||
void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * residual, const ggml_tensor * post, const ggml_tensor * comb, ggml_tensor * dst) {
|
||||
void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * residual, const ggml_tensor * post, const ggml_tensor * comb, ggml_tensor * dst, const ggml_tensor * gate_scale_in) {
|
||||
VK_LOG_DEBUG("ggml_vk_dsv4_hc_post(" << x << ", " << residual << ", " << post << ", " << comb << ", " << dst << ")");
|
||||
|
||||
vk_pipeline pipeline = comb ? ctx->device->pipeline_dsv4_hc_post_f32 : ctx->device->pipeline_dsv4_hc_post_nocomb_f32;
|
||||
@@ -7170,7 +7183,9 @@ void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, con
|
||||
|
||||
const vk_subbuffer x_buf = ggml_vk_tensor_subbuffer(ctx, x, true);
|
||||
const vk_subbuffer r_buf = ggml_vk_tensor_subbuffer(ctx, residual, true);
|
||||
const vk_subbuffer p_buf = ggml_vk_tensor_subbuffer(ctx, post, true);
|
||||
// with a fused gate, post is scale(sigmoid(scale(p_src))) and the shader applies it to p_src
|
||||
const ggml_tensor * p_src = gate_scale_in ? gate_scale_in->src[0] : post;
|
||||
const vk_subbuffer p_buf = ggml_vk_tensor_subbuffer(ctx, p_src, true);
|
||||
const vk_subbuffer c_buf = comb ? ggml_vk_tensor_subbuffer(ctx, comb, true) : x_buf;
|
||||
const vk_subbuffer d_buf = ggml_vk_tensor_subbuffer(ctx, dst, true);
|
||||
|
||||
@@ -7178,12 +7193,15 @@ void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, con
|
||||
n_embd, n_tokens,
|
||||
ggml_vk_nb_elem(x, 0), ggml_vk_nb_elem(x, 1),
|
||||
ggml_vk_nb_elem(residual, 0), ggml_vk_nb_elem(residual, 1), ggml_vk_nb_elem(residual, 2),
|
||||
ggml_vk_nb_elem(post, 0), ggml_vk_nb_elem(post, 1),
|
||||
ggml_vk_nb_elem(p_src, 0), ggml_vk_nb_elem(p_src, 1),
|
||||
comb ? ggml_vk_nb_elem(comb, 0) : 0, comb ? ggml_vk_nb_elem(comb, 1) : 0, comb ? ggml_vk_nb_elem(comb, 2) : 0,
|
||||
ggml_vk_nb_elem(dst, 0), ggml_vk_nb_elem(dst, 1), ggml_vk_nb_elem(dst, 2),
|
||||
0, 0, 0, 0, 0,
|
||||
gate_scale_in ? 1u : 0u,
|
||||
gate_scale_in ? ggml_get_op_params_f32(gate_scale_in, 0) : 1.0f,
|
||||
gate_scale_in ? ggml_get_op_params_f32(post, 0) : 1.0f,
|
||||
};
|
||||
init_pushconst_tensor_offsets(ctx, pc, x, residual, post, comb, dst);
|
||||
init_pushconst_tensor_offsets(ctx, pc, x, residual, p_src, comb, dst);
|
||||
|
||||
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { x_buf, r_buf, p_buf, c_buf, d_buf }, pc, { n_embd, n_tokens, 1 });
|
||||
}
|
||||
@@ -7601,9 +7619,12 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context&
|
||||
}
|
||||
ggml_vk_sync_buffers(ctx, subctx);
|
||||
|
||||
uint32_t stride_batch_x = ne00*ne01;
|
||||
const uint64_t x_range = qx_needs_dequant ? x_sz : ggml_nbytes(src0);
|
||||
const uint64_t y_range = (qy_needs_dequant || quantize_y) ? y_sz : ggml_nbytes(src1);
|
||||
|
||||
uint32_t stride_batch_x = qx_needs_dequant ? ne00*ne01 : ggml_vk_batch_stride(src0);
|
||||
uint32_t stride_b_y = y_needs_k_padding ? y_staged_row_stride : ne10;
|
||||
uint32_t stride_batch_y = y_needs_k_padding ? y_staged_row_stride * ne11 : ne10*ne11;
|
||||
uint32_t stride_batch_y = y_needs_k_padding ? y_staged_row_stride * ne11 : ((qy_needs_dequant || quantize_y) ? ne10*ne11 : ggml_vk_batch_stride(src1));
|
||||
|
||||
if (!ggml_vk_dim01_contiguous(src0) && !qx_needs_dequant) {
|
||||
stride_batch_x = src0->nb[0] / ggml_type_size(src0->type);
|
||||
@@ -7616,7 +7637,7 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context&
|
||||
// compute
|
||||
ggml_vk_matmul_id(
|
||||
ctx, subctx, pipeline,
|
||||
{ d_X, x_buf_offset, x_sz }, { d_Y, y_buf_offset, y_sz },
|
||||
{ d_X, x_buf_offset, x_range }, { d_Y, y_buf_offset, y_range },
|
||||
{ d_D, d_buf_offset, d_sz }, { d_ids, ids_buf_offset, ids_sz }, expert_count_buf,
|
||||
ne01, ne21, ne10, ne10, stride_b_y, ne01,
|
||||
stride_batch_x, stride_batch_y, ne20*ne21,
|
||||
@@ -7801,7 +7822,8 @@ static void ggml_vk_mul_mat_vec_id_q_f16(ggml_backend_vk_context * ctx, vk_conte
|
||||
}
|
||||
}
|
||||
|
||||
uint32_t stride_batch_y = ne10*ne11;
|
||||
uint32_t stride_batch_x = qx_needs_dequant ? ne00*ne01 : ggml_vk_batch_stride(src0);
|
||||
uint32_t stride_batch_y = (qy_needs_dequant || quantize_y) ? ne10*ne11 : ggml_vk_batch_stride(src1);
|
||||
|
||||
if (!ggml_vk_dim01_contiguous(src1) && !qy_needs_dequant) {
|
||||
stride_batch_y = src1->nb[2] / ggml_type_size(src1->type);
|
||||
@@ -7844,7 +7866,7 @@ static void ggml_vk_mul_mat_vec_id_q_f16(ggml_backend_vk_context * ctx, vk_conte
|
||||
for (uint32_t expert_i1 = 0; expert_i1 < nei1; ++expert_i1) {
|
||||
const vk_mat_vec_id_push_constants pc = {
|
||||
(uint32_t)ne00, (uint32_t)ne10, (uint32_t)ne10, (uint32_t)ne01,
|
||||
(uint32_t)(ne00 * ne01), stride_batch_y, (uint32_t)(ne20 * ne21),
|
||||
stride_batch_x, stride_batch_y, (uint32_t)(ne20 * ne21),
|
||||
fusion_flags,
|
||||
(uint32_t)nei0, (uint32_t)ne11, expert_i1, nbi1
|
||||
};
|
||||
@@ -11172,8 +11194,10 @@ void ggml_vk_argsort(ggml_backend_vk_context * ctx, vk_context& subctx, const gg
|
||||
// Pick the largest workgroup size <= ncolsp2
|
||||
uint32_t pipeline_idx = std::min(ncols_pad_log2, num_argsort_pipelines - 1);
|
||||
|
||||
uint32_t max_wg_log2 = std::min(ctx->device->max_workgroup_size_log2, num_argsort_pipelines - 1);
|
||||
|
||||
// Use the "small" argsort shader if the whole sort can be done by a single workgroup.
|
||||
bool use_small = ncols_pad_log2 <= ctx->device->max_workgroup_size_log2 &&
|
||||
bool use_small = ncols_pad_log2 <= max_wg_log2 &&
|
||||
ctx->device->pipeline_argsort_f32[pipeline_idx] != nullptr;
|
||||
|
||||
vk_pipeline pipeline = use_small ? ctx->device->pipeline_argsort_f32[pipeline_idx]
|
||||
@@ -11208,7 +11232,7 @@ void ggml_vk_argsort(ggml_backend_vk_context * ctx, vk_context& subctx, const gg
|
||||
{
|
||||
vk_op_argsort_push_constants pc2 = pc;
|
||||
pc2.outer_start = 0;
|
||||
pc2.outer_end = std::min(ncols_pad_log2, ctx->device->max_workgroup_size_log2);
|
||||
pc2.outer_end = std::min(ncols_pad_log2, max_wg_log2);
|
||||
pc2.inner_start = 0;
|
||||
pc2.inner_end = 100;
|
||||
ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1);
|
||||
@@ -11217,7 +11241,7 @@ void ggml_vk_argsort(ggml_backend_vk_context * ctx, vk_context& subctx, const gg
|
||||
if (!use_small) {
|
||||
ggml_vk_sync_buffers(ctx, subctx);
|
||||
// Loop over outer/inner passes, synchronizing between each pass.
|
||||
for (uint32_t outer = ctx->device->max_workgroup_size_log2; outer < ncols_pad_log2; ++outer) {
|
||||
for (uint32_t outer = max_wg_log2; outer < ncols_pad_log2; ++outer) {
|
||||
for (uint32_t inner = 0; inner < outer + 1; ++inner) {
|
||||
vk_op_argsort_push_constants pc2 = pc;
|
||||
pc2.outer_start = outer;
|
||||
@@ -12340,7 +12364,12 @@ bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgraph, in
|
||||
|
||||
break;
|
||||
case GGML_OP_SCALE:
|
||||
ggml_vk_scale(ctx, compute_ctx, src0, node);
|
||||
if (ctx->fused_hc_post_gate) {
|
||||
ggml_tensor * hc_post = cgraph->nodes[node_idx + ctx->num_additional_fused_ops];
|
||||
ggml_vk_dsv4_hc_post(ctx, compute_ctx, hc_post->src[0], hc_post->src[1], hc_post->src[2], hc_post->src[3], hc_post, node);
|
||||
} else {
|
||||
ggml_vk_scale(ctx, compute_ctx, src0, node);
|
||||
}
|
||||
|
||||
break;
|
||||
case GGML_OP_SQR:
|
||||
@@ -12820,6 +12849,7 @@ void ggml_vk_cleanup(ggml_backend_vk_context * ctx) {
|
||||
}
|
||||
ctx->descriptor_pools.clear();
|
||||
ctx->descriptor_sets.clear();
|
||||
ctx->descriptor_set_bindings.clear();
|
||||
|
||||
ctx->compute_cmd_pool.destroy(ctx->device->device);
|
||||
if (ctx->device->async_use_transfer_queue) {
|
||||
@@ -13415,6 +13445,19 @@ static bool ggml_vk_can_fuse_unary_mul_pair(const struct ggml_cgraph * cgraph, i
|
||||
ggml_vk_can_fuse_unary_mul(cgraph, node_idx, node_idx + 1);
|
||||
}
|
||||
|
||||
static bool ggml_vk_can_fuse_hc_post_gate(const struct ggml_cgraph * cgraph, int node_idx) {
|
||||
const ggml_tensor * scale_in = cgraph->nodes[node_idx];
|
||||
const ggml_tensor * sigmoid = cgraph->nodes[node_idx + 1];
|
||||
const ggml_tensor * scale_out = cgraph->nodes[node_idx + 2];
|
||||
|
||||
// the shader folds scale -> sigmoid -> scale; a bias on either scale is not handled
|
||||
return ggml_get_unary_op(sigmoid) == GGML_UNARY_OP_SIGMOID &&
|
||||
ggml_get_op_params_f32(scale_in, 1) == 0.0f &&
|
||||
ggml_get_op_params_f32(scale_out, 1) == 0.0f &&
|
||||
scale_in->src[0]->type == GGML_TYPE_F32 &&
|
||||
ggml_are_same_shape(scale_in->src[0], scale_out);
|
||||
}
|
||||
|
||||
bool ggml_vk_can_fuse(const ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx, std::initializer_list<enum ggml_op> ops) {
|
||||
if (ops.size() == 2 && ops.begin()[0] == GGML_OP_UNARY && ops.begin()[1] == GGML_OP_MUL) {
|
||||
return ggml_vk_can_fuse_unary_mul_pair(cgraph, node_idx);
|
||||
@@ -14262,6 +14305,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg
|
||||
ctx->fused_topk_moe_mode = TOPK_MOE_COUNT;
|
||||
ctx->fused_topk_moe_scale = false;
|
||||
ctx->fused_topk_qsa = false;
|
||||
ctx->fused_hc_post_gate = false;
|
||||
ctx->fused_rms_norm_mode = RMS_NORM_COUNT;
|
||||
const char *fusion_string {};
|
||||
if (!ctx->device->disable_fusion) {
|
||||
@@ -14318,6 +14362,13 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg
|
||||
op_srcs_fused_elementwise[0] = false;
|
||||
op_srcs_fused_elementwise[1] = true;
|
||||
op_srcs_fused_elementwise[2] = true;
|
||||
} else if (ggml_can_fuse_subgraph(cgraph, i, hc_post_gate_pattern, { i + 3 }) &&
|
||||
ggml_check_edges(cgraph, i, hc_post_gate_edges) &&
|
||||
ggml_vk_can_fuse_hc_post_gate(cgraph, i)) {
|
||||
ctx->num_additional_fused_ops = hc_post_gate_pattern.size() - 1;
|
||||
ctx->fused_hc_post_gate = true;
|
||||
fusion_string = "HC_POST_GATE";
|
||||
std::fill_n(op_srcs_fused_elementwise, ctx->num_additional_fused_ops + 1, false);
|
||||
} else if (ggml_vk_can_fuse(ctx, cgraph, i, rms_norm_mul_add_mul_pattern)) {
|
||||
ctx->num_additional_fused_ops = 3;
|
||||
ctx->fused_rms_norm_mode = RMS_NORM_MUL_ADD_MUL;
|
||||
@@ -14496,6 +14547,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg
|
||||
ctx->fused_topk_moe_mode = TOPK_MOE_COUNT;
|
||||
ctx->fused_topk_moe_scale = false;
|
||||
ctx->fused_topk_qsa = false;
|
||||
ctx->fused_hc_post_gate = false;
|
||||
ctx->fused_rms_norm_mode = RMS_NORM_COUNT;
|
||||
fusion_string = nullptr;
|
||||
}
|
||||
@@ -14752,6 +14804,11 @@ void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * graph,
|
||||
if (keep_pattern(rope_view_set_rows_pattern)) {
|
||||
continue;
|
||||
}
|
||||
if (match_pattern(hc_post_gate_pattern, first_unused)) {
|
||||
add_pattern_alloc_deps(hc_post_gate_pattern, first_unused + (int) hc_post_gate_pattern.size() - 1);
|
||||
keep_pattern(hc_post_gate_pattern);
|
||||
continue;
|
||||
}
|
||||
|
||||
// First, grab the next unused node.
|
||||
current_set.push_back(first_unused);
|
||||
@@ -14791,7 +14848,8 @@ void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * graph,
|
||||
match_pattern(rms_norm_mul_add_pattern, j) ||
|
||||
match_pattern(rms_norm_mul_rope_view_set_rows_pattern, j) ||
|
||||
match_pattern(rms_norm_view_set_rows_pattern, j) ||
|
||||
match_pattern(rope_view_set_rows_pattern, j)) {
|
||||
match_pattern(rope_view_set_rows_pattern, j) ||
|
||||
match_pattern(hc_post_gate_pattern, j)) {
|
||||
continue;
|
||||
}
|
||||
bool ok = true;
|
||||
@@ -15612,7 +15670,7 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
|
||||
if (device->vulkan_memory_model) {
|
||||
return true;
|
||||
} else {
|
||||
return op->ne[0] <= (1 << device->max_workgroup_size_log2);
|
||||
return op->ne[0] <= (1 << std::min(device->max_workgroup_size_log2, num_argsort_pipelines - 1));
|
||||
}
|
||||
}
|
||||
case GGML_OP_TOP_K:
|
||||
|
||||
@@ -33,6 +33,10 @@ layout(push_constant) uniform parameter
|
||||
uint p_offset;
|
||||
uint c_offset;
|
||||
uint d_offset;
|
||||
|
||||
uint gate; // post = gate_scale_out*sigmoid(gate_scale_in*p)
|
||||
float gate_scale_in;
|
||||
float gate_scale_out;
|
||||
};
|
||||
|
||||
layout(binding = 0, std430) readonly buffer X { float data_x[]; };
|
||||
@@ -51,7 +55,8 @@ void main() {
|
||||
const uint it = gl_WorkGroupID.y;
|
||||
|
||||
if (tid < hc) {
|
||||
post_s[tid] = data_p[p_offset + tid * nbp0 + it * nbp1];
|
||||
const float p = data_p[p_offset + tid * nbp0 + it * nbp1];
|
||||
post_s[tid] = gate != 0 ? (1.0f / (1.0f + exp(-(p * gate_scale_in)))) * gate_scale_out : p;
|
||||
}
|
||||
if (HAS_COMB == 1 && tid < hc * hc) {
|
||||
const uint idst = tid & 3;
|
||||
|
||||
@@ -3774,7 +3774,28 @@ static void ggml_backend_webgpu_buffer_set_tensor(ggml_backend_buffer_t buffer,
|
||||
|
||||
size_t total_offset = ggml_webgpu_tensor_offset(tensor) + offset;
|
||||
|
||||
buf_ctx->global_ctx->queue.WriteBuffer(buf_ctx->buffer, total_offset, data, (size / 4) * 4);
|
||||
// WriteBuffer needs the offset and the size to be multiples of 4.
|
||||
// Write the misaligned head bytes using compute memset, then increment total_offset
|
||||
// and data pointer so that they are 4-aligned.
|
||||
if (total_offset % 4 != 0) {
|
||||
size_t lane = total_offset % 4; // in-word lane the head starts at (the tail below always starts at 0)
|
||||
size_t head = std::min<size_t>(4 - lane, size);
|
||||
|
||||
// Pack head bytes into a uint32_t
|
||||
uint32_t head_val = 0;
|
||||
for (size_t i = 0; i < head; i++) {
|
||||
((uint8_t *) &head_val)[lane + i] = ((const uint8_t *) data)[i];
|
||||
}
|
||||
ggml_backend_webgpu_buffer_memset(buf_ctx->global_ctx, buf_ctx->buffer, head_val, total_offset, head);
|
||||
|
||||
total_offset += head;
|
||||
size -= head;
|
||||
data = (const uint8_t *) data + head;
|
||||
}
|
||||
|
||||
if (size > 0) {
|
||||
buf_ctx->global_ctx->queue.WriteBuffer(buf_ctx->buffer, total_offset, data, (size / 4) * 4);
|
||||
}
|
||||
|
||||
if (size % 4 != 0) {
|
||||
// If size is not a multiple of 4, we need to memset the remaining bytes
|
||||
|
||||
@@ -225,6 +225,11 @@ static enum ggml_status ggml_backend_zdnn_buffer_init_tensor(ggml_backend_buffer
|
||||
return GGML_STATUS_SUCCESS;
|
||||
}
|
||||
|
||||
// reject empty tensors to avoid zDNN crash
|
||||
if (ggml_is_empty(tensor)) {
|
||||
return GGML_STATUS_SUCCESS;
|
||||
}
|
||||
|
||||
ggml_backend_zdnn_buffer_context * ctx = (ggml_backend_zdnn_buffer_context *)buffer->context;
|
||||
|
||||
const int64_t tsize = ggml_nbytes(tensor);
|
||||
|
||||
+12
-11
@@ -14,6 +14,7 @@
|
||||
#include <new>
|
||||
#include <stdexcept>
|
||||
#include <string>
|
||||
#include <unordered_set>
|
||||
#include <vector>
|
||||
|
||||
#define GGUF_MAX_STRING_LENGTH (1024*1024*1024)
|
||||
@@ -550,6 +551,8 @@ static struct gguf_context * gguf_init_from_reader(const struct gguf_reader & gr
|
||||
|
||||
// KV pairs
|
||||
{
|
||||
std::unordered_set<std::string> seen_keys;
|
||||
|
||||
for (int64_t i = 0; ok && i < n_kv; ++i) {
|
||||
std::string key;
|
||||
gguf_type type = gguf_type(-1);
|
||||
@@ -569,11 +572,9 @@ static struct gguf_context * gguf_init_from_reader(const struct gguf_reader & gr
|
||||
GGML_LOG_ERROR("%s: key %" PRIi64 " is empty\n", __func__, i);
|
||||
ok = false;
|
||||
}
|
||||
for (size_t j = 0; ok && j < ctx->kv.size(); ++j) {
|
||||
if (key == ctx->kv[j].key) {
|
||||
GGML_LOG_ERROR("%s: duplicate key '%s' for tensors %zu and %" PRIi64 " \n", __func__, key.c_str(), j, i);
|
||||
ok = false;
|
||||
}
|
||||
if (ok && !seen_keys.insert(key).second) {
|
||||
GGML_LOG_ERROR("%s: duplicate key '%s' for KV pair %" PRIi64 "\n", __func__, key.c_str(), i);
|
||||
ok = false;
|
||||
}
|
||||
if (!ok) {
|
||||
break;
|
||||
@@ -636,6 +637,8 @@ static struct gguf_context * gguf_init_from_reader(const struct gguf_reader & gr
|
||||
}
|
||||
|
||||
// read the tensor info
|
||||
std::unordered_set<std::string> seen_tensor_names;
|
||||
|
||||
for (int64_t i = 0; ok && i < n_tensors; ++i) {
|
||||
struct gguf_tensor_info info;
|
||||
|
||||
@@ -659,12 +662,10 @@ static struct gguf_context * gguf_init_from_reader(const struct gguf_reader & gr
|
||||
ggml_set_name(&info.t, name.c_str());
|
||||
|
||||
// make sure there are no duplicate tensor names
|
||||
for (int64_t j = 0; ok && j < i; ++j) {
|
||||
if (strcmp(info.t.name, ctx->info[j].t.name) == 0) {
|
||||
GGML_LOG_ERROR("%s: duplicate tensor name '%s' for tensors %" PRIi64 " and %" PRIi64 "\n", __func__, info.t.name, j, i);
|
||||
ok = false;
|
||||
break;
|
||||
}
|
||||
if (ok && !seen_tensor_names.insert(name).second) {
|
||||
GGML_LOG_ERROR("%s: duplicate tensor name '%s' for tensor %" PRIi64 "\n", __func__, info.t.name, i);
|
||||
ok = false;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (!ok) {
|
||||
|
||||
+5
-2
@@ -482,14 +482,14 @@ extern "C" {
|
||||
LLAMA_API struct llama_model_quantize_params llama_model_quantize_default_params(void);
|
||||
|
||||
// Initialize the llama + ggml backend
|
||||
// If numa is true, use NUMA optimizations
|
||||
// Call once at the start of the program
|
||||
LLAMA_API void llama_backend_init(void);
|
||||
|
||||
// Call once at the end of the program - currently only used for MPI
|
||||
LLAMA_API void llama_backend_free(void);
|
||||
|
||||
//optional:
|
||||
// Optional: enable numa optimizations
|
||||
// TODO: deprecate and make part of llama_backend_init()
|
||||
LLAMA_API void llama_numa_init(enum ggml_numa_strategy numa);
|
||||
|
||||
// Optional: an auto threadpool gets created in ggml if not passed explicitly
|
||||
@@ -1108,6 +1108,9 @@ extern "C" {
|
||||
// If set to true, the model will only attend to the past tokens
|
||||
LLAMA_API void llama_set_causal_attn(struct llama_context * ctx, bool causal_attn);
|
||||
|
||||
// Returns whether the context is currently using causal attention
|
||||
LLAMA_API bool llama_get_causal_attn(const struct llama_context * ctx);
|
||||
|
||||
// Set whether the model is in warmup mode or not
|
||||
// If true, all model tensors are activated during llama_decode() to load and cache their weights.
|
||||
//
|
||||
|
||||
@@ -64,7 +64,24 @@ def main():
|
||||
'_ZL12rwkv_wkv_f32ILi128EEviiiiPKfS1_S1_S1_S1_S1_Pf',
|
||||
'_ZL9mul_mat_qIL9ggml_type10ELi64ELb1EEvPKcPKiS4_S4_PfS5_PKf15HIP_vector_typeIjLj3EEiiiiiS9_S9_iiiS9_S9_iiiS9_',
|
||||
'_ZL9mul_mat_qIL9ggml_type42ELi128ELb1EEvPKcPKiS4_S4_PfS5_PKf15HIP_vector_typeIjLj3EEiiiiiS9_S9_iiiS9_S9_iiiS9_',
|
||||
'_ZL18flash_attn_ext_f16ILi576ELi512ELi2ELi32ELb0ELb1ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil'
|
||||
'_ZL18flash_attn_ext_f16ILi576ELi512ELi2ELi32ELb0ELb1ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi16ELi4ELb0ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi16ELi4ELb1ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi32ELi2ELb0ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi32ELi2ELb1ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi8ELi8ELb0ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi8ELi8ELb1ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi576ELi512ELi16ELi4ELb0ELb1ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi576ELi512ELi4ELi16ELb0ELb1ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi16ELi2ELb0ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi16ELi2ELb1ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi4ELi8ELb0ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi4ELi8ELb1ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi8ELi4ELb0ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi512ELi512ELi8ELi4ELb1ELb0ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi576ELi512ELi1ELi32ELb0ELb1ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi576ELi512ELi2ELi16ELb0ELb1ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
'_ZL18flash_attn_ext_f16ILi576ELi512ELi8ELi4ELb0ELb1ELb0EEvPKcS1_S1_S1_S1_PKiPfP15HIP_vector_typeIfLj2EEffffjfiS5_IjLj3EEiiiiiiiiiiiliiliiiiil',
|
||||
}
|
||||
|
||||
functions = parse_log_file(log_file)
|
||||
|
||||
@@ -514,7 +514,7 @@ def generate_perfetto_trace(filtered_ops, trace_events, output_path):
|
||||
tm = time_mappers[dev]
|
||||
e['ts_ns'] = tm.cycle_to_ns(e['start_cyc'])
|
||||
dur_ns = tm.dur_cycles_to_ns(e['start_cyc'], e['end_cyc'] - e['start_cyc'])
|
||||
e['dur_ns'] = max(dur_ns, 100)
|
||||
e['dur_ns'] = max(dur_ns, 1)
|
||||
|
||||
# Allocate slots (sub-tracks) to prevent overlaps on same virtual track
|
||||
active_slots = defaultdict(list)
|
||||
|
||||
+1
-1
@@ -722,7 +722,7 @@ static const std::map<llm_tensor, llm_tensor_info> LLM_TENSOR_INFOS = {
|
||||
{LLM_TENSOR_TOKEN_EMBD, {LLM_TENSOR_LAYER_INPUT, GGML_OP_GET_ROWS}},
|
||||
{LLM_TENSOR_POS_EMBD, {LLM_TENSOR_LAYER_INPUT, GGML_OP_GET_ROWS}},
|
||||
{LLM_TENSOR_TOKEN_TYPES, {LLM_TENSOR_LAYER_INPUT, GGML_OP_GET_ROWS}},
|
||||
{LLM_TENSOR_HRM_Z_L_INIT, {LLM_TENSOR_LAYER_INPUT, GGML_OP_ADD}},
|
||||
{LLM_TENSOR_HRM_Z_L_INIT, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_ADD}},
|
||||
{LLM_TENSOR_TOKEN_EMBD_NORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, // do the norms on the first layer (not the input layer)
|
||||
{LLM_TENSOR_OUTPUT, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}},
|
||||
{LLM_TENSOR_CLS, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}},
|
||||
|
||||
+10
-1
@@ -1256,7 +1256,12 @@ void llama_context::set_causal_attn(bool value) {
|
||||
|
||||
cparams.causal_attn = value;
|
||||
|
||||
sched_need_reserve = true;
|
||||
// no scheduler reserve needed because graph shapes must not depend on causal_attn, a flip only rebuilds the graph
|
||||
//sched_need_reserve = true;
|
||||
}
|
||||
|
||||
bool llama_context::get_causal_attn() const {
|
||||
return cparams.causal_attn;
|
||||
}
|
||||
|
||||
void llama_context::set_warmup(bool value) {
|
||||
@@ -3933,6 +3938,10 @@ void llama_set_causal_attn(llama_context * ctx, bool causal_attn) {
|
||||
ctx->set_causal_attn(causal_attn);
|
||||
}
|
||||
|
||||
bool llama_get_causal_attn(const llama_context * ctx) {
|
||||
return ctx->get_causal_attn();
|
||||
}
|
||||
|
||||
void llama_set_warmup(llama_context * ctx, bool warmup) {
|
||||
ctx->set_warmup(warmup);
|
||||
}
|
||||
|
||||
@@ -103,6 +103,8 @@ struct llama_context {
|
||||
const llama_token * get_sampled_candidates_ith(int32_t idx);
|
||||
size_t get_sampled_candidates_count(int32_t idx);
|
||||
|
||||
bool get_causal_attn() const;
|
||||
|
||||
void attach_threadpool(
|
||||
ggml_threadpool_t threadpool,
|
||||
ggml_threadpool_t threadpool_batch);
|
||||
|
||||
+1
-1
@@ -297,7 +297,7 @@ void llm_graph_input_cls::set_input(const llama_ubatch * ubatch) {
|
||||
|
||||
const bool last = (
|
||||
cparams.pooling_type == LLAMA_POOLING_TYPE_LAST ||
|
||||
(cparams.pooling_type == LLAMA_POOLING_TYPE_RANK && (arch == LLM_ARCH_QWEN3 || arch == LLM_ARCH_QWEN3VL)) // qwen3 reranking & embedding models use last token
|
||||
(cparams.pooling_type == LLAMA_POOLING_TYPE_RANK && cparams.causal_attn)
|
||||
);
|
||||
|
||||
for (int i = 0; i < n_tokens; ++i) {
|
||||
|
||||
@@ -277,7 +277,8 @@ void llama_memory_hybrid_idx::set_input_qsa(
|
||||
ggml_tensor * bias,
|
||||
const llama_ubatch * ubatch,
|
||||
uint32_t ratio,
|
||||
bool blk_bias) const {
|
||||
bool blk_bias,
|
||||
bool causal_attn) const {
|
||||
GGML_ASSERT(ratio > 0);
|
||||
GGML_ASSERT(get_mem_idx() != nullptr);
|
||||
|
||||
@@ -545,7 +546,7 @@ void llama_memory_hybrid_idx::set_input_qsa(
|
||||
|
||||
if (blk_bias) {
|
||||
// a block sits wholly inside or outside the tail, so one value covers it
|
||||
// the caller adds the attention mask, which drops empty, foreign and future cells
|
||||
// the caller adds the attention mask, which drops empty, foreign and, when causal, future cells
|
||||
float * cur_blk_bias = dst_bias + i*n_blocks;
|
||||
|
||||
for (int64_t b = 0; b < n_blocks; ++b) {
|
||||
@@ -555,7 +556,7 @@ void llama_memory_hybrid_idx::set_input_qsa(
|
||||
}
|
||||
|
||||
// finite, so it can never meet a -inf and produce a nan
|
||||
cur_blk_bias[b] = bid_idx[b] >= tail_start ? 1e9f : 0.0f;
|
||||
cur_blk_bias[b] = (causal_attn && bid_idx[b] >= tail_start) ? 1e9f : 0.0f;
|
||||
}
|
||||
|
||||
// the spare block holds the unpooled cells, which are the incomplete tail, so
|
||||
@@ -576,7 +577,10 @@ void llama_memory_hybrid_idx::set_input_qsa(
|
||||
if (!cells.is_empty(j) && cells.seq_has(j, seq_id)) {
|
||||
const int64_t idx = ranked ? rank[j] : cells.pos_get(j);
|
||||
|
||||
if (idx <= q) {
|
||||
if (!causal_attn) {
|
||||
// every visible block competes on score and the unpooled cells are always selected
|
||||
v = blk_of[j] < 0 ? 1e9f : 0.0f;
|
||||
} else if (idx <= q) {
|
||||
// finite, so it can never meet a -inf and produce a nan
|
||||
v = idx >= tail_start ? 1e9f : (blk_of[j] < 0 ? -INFINITY : 0.0f);
|
||||
}
|
||||
@@ -676,8 +680,9 @@ void llama_memory_hybrid_idx_context::set_input_qsa(
|
||||
ggml_tensor * bias,
|
||||
const llama_ubatch * ubatch,
|
||||
uint32_t ratio,
|
||||
bool blk_bias) const {
|
||||
bool blk_bias,
|
||||
bool causal_attn) const {
|
||||
GGML_ASSERT(mem != nullptr);
|
||||
|
||||
mem->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias);
|
||||
mem->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias, causal_attn);
|
||||
}
|
||||
|
||||
@@ -12,6 +12,8 @@
|
||||
// llama_memory_hybrid plus a third cache with one indexer key per token, for block-sparse attention (qwen4exp QSA)
|
||||
// the indexer is a side buffer over the attention cells: same size, padding, streams and slots, so cell j is one token in both
|
||||
|
||||
// TODO: this memory module is pending complete reimplementation - do not use for model other than Qwen4
|
||||
|
||||
class llama_memory_hybrid_idx : public llama_memory_hybrid {
|
||||
public:
|
||||
llama_memory_hybrid_idx(
|
||||
@@ -83,9 +85,10 @@ public:
|
||||
// bias F32 [n_kv, n_tokens/ns, ns] -inf where invisible, large where always visible
|
||||
// blk_bias asks for the bias per block instead: [n_blocks, n_tokens/ns, ns]
|
||||
// the caller then adds the attention mask, the only part of the bias that varies within a block
|
||||
// causal_attn selects the rule: causal forces the query's own block on, non-causal lets every visible block compete on score
|
||||
void set_input_qsa(ggml_tensor * cell_blk, ggml_tensor * blk_cells, ggml_tensor * blk_pos,
|
||||
ggml_tensor * bias, const llama_ubatch * ubatch, uint32_t ratio,
|
||||
bool blk_bias) const;
|
||||
bool blk_bias, bool causal_attn) const;
|
||||
|
||||
private:
|
||||
// forget seq_id (all of it if seq_id < 0) in every cache at once, so a failed restore cannot leave the caches out of step
|
||||
@@ -143,7 +146,7 @@ public:
|
||||
|
||||
void set_input_qsa(ggml_tensor * cell_blk, ggml_tensor * blk_cells, ggml_tensor * blk_pos,
|
||||
ggml_tensor * bias, const llama_ubatch * ubatch, uint32_t ratio,
|
||||
bool blk_bias) const;
|
||||
bool blk_bias, bool causal_attn) const;
|
||||
|
||||
private:
|
||||
const llama_memory_hybrid_idx * mem = nullptr;
|
||||
|
||||
@@ -401,6 +401,7 @@ static struct llama_model * llama_model_load_from_file_impl(
|
||||
return nullptr;
|
||||
}
|
||||
}
|
||||
// TODO: remove
|
||||
ggml_time_init();
|
||||
|
||||
if (!params.vocab_only && ggml_backend_reg_count() == 0) {
|
||||
|
||||
+7
-10
@@ -446,19 +446,16 @@ static ggml_tensor * build_dflash2_conv(
|
||||
|
||||
ggml_tensor * weight_all = ggml_add(ctx0, coeff_all, base_side);
|
||||
|
||||
// taps at or past block_size only read the left padding and add nothing
|
||||
const int64_t n_taps = std::min(kernel_size, block_size);
|
||||
|
||||
ggml_tensor * result = nullptr;
|
||||
for (int64_t tap = 0; tap < kernel_size; ++tap) {
|
||||
for (int64_t tap = 0; tap < n_taps; ++tap) {
|
||||
ggml_tensor * values = blocks;
|
||||
if (tap > 0) {
|
||||
ggml_tensor * zeros = ggml_fill(ctx0,
|
||||
ggml_new_tensor_3d(ctx0, hidden->type, hidden_size, std::min(tap, block_size), n_blocks), 0.0f);
|
||||
if (tap < block_size) {
|
||||
ggml_tensor * previous = ggml_view_3d(ctx0, blocks, hidden_size, block_size - tap, n_blocks,
|
||||
blocks->nb[1], blocks->nb[2], 0);
|
||||
values = ggml_concat(ctx0, zeros, previous, 1);
|
||||
} else {
|
||||
values = zeros;
|
||||
}
|
||||
ggml_tensor * previous = ggml_view_3d(ctx0, blocks, hidden_size, block_size - tap, n_blocks,
|
||||
blocks->nb[1], blocks->nb[2], 0);
|
||||
values = ggml_pad_ext(ctx0, previous, 0, 0, tap, 0, 0, 0, 0, 0);
|
||||
}
|
||||
values = ggml_reshape_2d(ctx0, values, hidden_size, n_tokens);
|
||||
|
||||
|
||||
@@ -43,7 +43,7 @@ void llama_model_hrm_text::load_arch_tensors(llama_model_loader &) {
|
||||
output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), { n_embd, n_vocab }, TENSOR_DUPLICATED);
|
||||
}
|
||||
|
||||
hrm_z_l_init = create_tensor(tn(LLM_TENSOR_HRM_Z_L_INIT), { n_embd }, 0);
|
||||
hrm_z_l_init = create_tensor(tn(LLM_TENSOR_HRM_Z_L_INIT, 0), { n_embd }, 0);
|
||||
|
||||
const int lps = hparams.n_hrm_layers_per_stack;
|
||||
|
||||
|
||||
+12
-6
@@ -6,6 +6,9 @@
|
||||
#include <algorithm>
|
||||
#include <cinttypes>
|
||||
|
||||
// [TAG_QWEN4_REIMPLEMENT]
|
||||
// TODO: this graph implementation is pending complete reimplementation - do not use it as a reference
|
||||
|
||||
// bad metadata must be catchable: GGML_ASSERT aborts the whole process
|
||||
static void qwen4exp_require_nonzero(const llama_model_loader & ml, llm_kv kid, uint32_t value) {
|
||||
if (value == 0) {
|
||||
@@ -489,13 +492,13 @@ ggml_tensor * llama_model_qwen4exp::graph::build_norm_gated(
|
||||
// one mean-pooled indexer key scores each block; set_input resolves the cache layout
|
||||
class llama_model_qwen4exp::llm_graph_input_qsa : public llm_graph_input_i {
|
||||
public:
|
||||
llm_graph_input_qsa(const llama_memory_hybrid_idx_context * mctx, uint32_t ratio, bool blk_bias) :
|
||||
mctx(mctx), ratio(ratio), blk_bias(blk_bias) {}
|
||||
llm_graph_input_qsa(const llama_memory_hybrid_idx_context * mctx, uint32_t ratio, bool blk_bias, bool causal_attn) :
|
||||
mctx(mctx), ratio(ratio), blk_bias(blk_bias), causal_attn(causal_attn) {}
|
||||
virtual ~llm_graph_input_qsa() = default;
|
||||
|
||||
void set_input(const llama_ubatch * ubatch) override {
|
||||
mctx->get_idx()->set_input_k_idxs(k_idxs, ubatch);
|
||||
mctx->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias);
|
||||
mctx->set_input_qsa(cell_blk, blk_cells, blk_pos, bias, ubatch, ratio, blk_bias, causal_attn);
|
||||
}
|
||||
|
||||
bool can_reuse(const llm_graph_params & params) override {
|
||||
@@ -537,6 +540,9 @@ public:
|
||||
|
||||
// the per-cell half of the bias is the attention mask, so only the per-block half is uploaded
|
||||
const bool blk_bias;
|
||||
|
||||
// this is fixed for the graph's lifetime, as causal_attn is part of the reuse key (llm_graph_params::allow_reuse)
|
||||
const bool causal_attn;
|
||||
};
|
||||
|
||||
ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k(
|
||||
@@ -564,11 +570,11 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k(
|
||||
|
||||
// only the "which block is visible" half of the bias varies per block
|
||||
// the rest is the visible/not test the attention mask already carries, so upload the per-block half only: 1/ratio of the cells
|
||||
// alibi writes distances instead of a mask and non-causal keeps future cells, so both opt out
|
||||
// alibi writes distances instead of a mask, so it opts out
|
||||
// the mask also holds an mrope rule for the query's own position, but only 2d image positions can differ there
|
||||
const bool blk_bias = kq_mask != nullptr &&
|
||||
kq_mask->ne[0] == n_kv && kq_mask->ne[1] == n_tps && kq_mask->ne[3] == n_stream &&
|
||||
cparams.causal_attn && !hparams.use_alibi;
|
||||
!hparams.use_alibi;
|
||||
|
||||
// nothing above depends on the layer, so the layers sharing a ratio share one input set
|
||||
llm_graph_input_qsa * inp = nullptr;
|
||||
@@ -577,7 +583,7 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_top_k(
|
||||
if (it != qsa_inps.end()) {
|
||||
inp = it->second;
|
||||
} else {
|
||||
auto qsa = std::make_unique<llm_graph_input_qsa>(mctx_hyb, (uint32_t) r, blk_bias);
|
||||
auto qsa = std::make_unique<llm_graph_input_qsa>(mctx_hyb, (uint32_t) r, blk_bias, cparams.causal_attn);
|
||||
|
||||
qsa->k_idxs = mctx_idx->build_input_k_idxs(ctx0, ubatch);
|
||||
qsa->cell_blk = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_kv, n_stream);
|
||||
|
||||
+4
-32
@@ -207,42 +207,13 @@ if (NOT WIN32 OR NOT BUILD_SHARED_LIBS)
|
||||
FIXTURES_SETUP generate-models
|
||||
)
|
||||
|
||||
# Test recurrent-state rollback across all architectures, using the generated dummy models
|
||||
llama_test(
|
||||
test-recurrent-state-rollback
|
||||
LABEL main
|
||||
ARGS -m "${MODEL_DIR}/qwen35-dense.gguf"
|
||||
)
|
||||
set_tests_properties(test-recurrent-state-rollback PROPERTIES
|
||||
FIXTURES_REQUIRED generate-models
|
||||
)
|
||||
|
||||
llama_test(
|
||||
test-recurrent-state-rollback
|
||||
NAME test-recurrent-state-rollback-nemotron-h
|
||||
LABEL main
|
||||
ARGS -m "${MODEL_DIR}/nemotron_h-dense.gguf"
|
||||
)
|
||||
set_tests_properties(test-recurrent-state-rollback-nemotron-h PROPERTIES
|
||||
FIXTURES_REQUIRED generate-models
|
||||
)
|
||||
llama_test(
|
||||
test-recurrent-state-rollback
|
||||
NAME test-recurrent-state-rollback-dsv4
|
||||
LABEL main
|
||||
ARGS -m "${MODEL_DIR}/deepseek4-moe.gguf"
|
||||
)
|
||||
set_tests_properties(test-recurrent-state-rollback-dsv4 PROPERTIES
|
||||
FIXTURES_REQUIRED generate-models
|
||||
)
|
||||
llama_test(
|
||||
test-recurrent-state-rollback
|
||||
NAME test-recurrent-state-rollback-kimi-k3
|
||||
LABEL main
|
||||
ARGS -m "${MODEL_DIR}/kimi-k3-moe.gguf"
|
||||
)
|
||||
set_tests_properties(test-recurrent-state-rollback-kimi-k3 PROPERTIES
|
||||
FIXTURES_REQUIRED generate-models
|
||||
ARGS --models "${MODEL_DIR}"
|
||||
)
|
||||
set_tests_properties(test-recurrent-state-rollback PROPERTIES FIXTURES_REQUIRED generate-models)
|
||||
|
||||
# Test state save/load functionality across all architectures, using the generated dummy models
|
||||
llama_test(
|
||||
@@ -316,6 +287,7 @@ llama_build_and_test(test-backend-sampler.cpp LABEL "model")
|
||||
|
||||
# Test for state restore with fragmented KV cache
|
||||
# Requires a model, uses same args pattern as test-thread-safety
|
||||
# TODO: run on all dummy models
|
||||
llama_build_and_test(test-state-restore-fragmented.cpp LABEL "model" ARGS -m "${MODEL_DEST}")
|
||||
set_tests_properties(test-state-restore-fragmented PROPERTIES FIXTURES_REQUIRED test-download-model)
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user