mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-01 05:56:52 +02:00
Compare commits
44
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6379261a77 | ||
|
|
0a743e4e5a | ||
|
|
3cf03257f2 | ||
|
|
b23efaa2ef | ||
|
|
4260903678 | ||
|
|
9a9f939b80 | ||
|
|
f072b10371 | ||
|
|
59657a613a | ||
|
|
e613ef2c81 | ||
|
|
851cb34f21 | ||
|
|
7d4b92bb9b | ||
|
|
1af554f8fc | ||
|
|
eb1e1f495f | ||
|
|
5b59b83f4e | ||
|
|
60b06ab9a9 | ||
|
|
efa28e950e | ||
|
|
59fc5a1ca3 | ||
|
|
b23701f77d | ||
|
|
60081bb2b5 | ||
|
|
2b1847030c | ||
|
|
50631b3d2c | ||
|
|
18a04f09c2 | ||
|
|
ec92815050 | ||
|
|
4fea119de3 | ||
|
|
5b335f413e | ||
|
|
542348a35c | ||
|
|
d663dd3f3a | ||
|
|
44be98f057 | ||
|
|
911f6cdc8a | ||
|
|
bbd488c42a | ||
|
|
dc85f89c7e | ||
|
|
8ed1a55efc | ||
|
|
bb11ebb682 | ||
|
|
f03cf3e9b8 | ||
|
|
bdcbaaf6e7 | ||
|
|
5c53396b89 | ||
|
|
972d2313bc | ||
|
|
c77ae695c9 | ||
|
|
b49650adb3 | ||
|
|
7076180486 | ||
|
|
ebbb185227 | ||
|
|
4ff829ec2e | ||
|
|
f172be756a | ||
|
|
87f9c82f2d |
@@ -1,5 +1,5 @@
|
||||
ARG OPENVINO_VERSION_MAJOR=2026.3.1
|
||||
ARG OPENVINO_VERSION_FULL=2026.3.1.22476.56d9685302d
|
||||
ARG OPENVINO_VERSION_MAJOR=2026.4
|
||||
ARG OPENVINO_VERSION_FULL=2026.4.0.22959.99c81491cc3
|
||||
ARG UBUNTU_VERSION=24.04
|
||||
|
||||
# Intel GPU driver versions. https://github.com/intel/compute-runtime/releases
|
||||
@@ -10,9 +10,9 @@ ARG COMPUTE_RUNTIME_VERSION_FULL=26.31.39395.13-0
|
||||
ARG IGDGMM_VERSION=22.10.0
|
||||
|
||||
# Intel NPU driver versions. https://github.com/intel/linux-npu-driver/releases
|
||||
ARG NPU_DRIVER_VERSION=v1.35.0
|
||||
ARG NPU_DRIVER_FULL=v1.35.0.20260722-29947505341
|
||||
ARG LIBZE1_VERSION=1.28.2-1~24.04~ppa1
|
||||
ARG NPU_DRIVER_VERSION=v1.38.0
|
||||
ARG NPU_DRIVER_FULL=v1.38.0.20260910-34487311128
|
||||
ARG LIBZE1_VERSION=1.32.0-1~24.04~ppa1
|
||||
|
||||
# Optional proxy build arguments
|
||||
ARG http_proxy=
|
||||
@@ -173,7 +173,7 @@ RUN --mount=type=cache,target=/var/cache/intel-npu,sharing=locked \
|
||||
fi; \
|
||||
DEB=/var/cache/intel-npu/libze1_${LIBZE1_VERSION}_amd64.deb; \
|
||||
if [ ! -f "$DEB" ]; then \
|
||||
wget -q -O "$DEB" https://snapshot.ppa.launchpadcontent.net/kobuk-team/intel-graphics/ubuntu/20260606T100000Z/pool/main/l/level-zero-loader/libze1_${LIBZE1_VERSION}_amd64.deb; \
|
||||
wget -q -O "$DEB" https://snapshot.ppa.launchpadcontent.net/kobuk-team/intel-graphics/ubuntu/20260830T100000Z/pool/main/l/level-zero-loader/libze1_${LIBZE1_VERSION}_amd64.deb; \
|
||||
fi; \
|
||||
mkdir /tmp/npu/ && cd /tmp/npu/ && tar -xf "$TGZ" && cp "$DEB" .; \
|
||||
apt-get update; \
|
||||
|
||||
@@ -29,7 +29,7 @@ concurrency:
|
||||
|
||||
jobs:
|
||||
android-ndk-snapdragon:
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: ubuntu-24.04 # previously ubuntu-latest
|
||||
container:
|
||||
image: 'ghcr.io/snapdragon-toolchain/arm64-android:v0.7'
|
||||
defaults:
|
||||
@@ -59,7 +59,7 @@ jobs:
|
||||
path: pkg-snapdragon/llama.cpp
|
||||
|
||||
linux-iot-snapdragon:
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: ubuntu-24.04 # previously ubuntu-latest
|
||||
container:
|
||||
image: 'ghcr.io/snapdragon-toolchain/arm64-linux:v0.7'
|
||||
defaults:
|
||||
|
||||
@@ -33,7 +33,7 @@ env:
|
||||
|
||||
jobs:
|
||||
default:
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: ubuntu-24.04 # previously ubuntu-latest
|
||||
|
||||
steps:
|
||||
- name: Clone
|
||||
@@ -49,7 +49,7 @@ jobs:
|
||||
distribution: zulu
|
||||
|
||||
- name: Setup Android SDK
|
||||
uses: android-actions/setup-android@40fd30fb8d7440372e1316f5d1809ec01dcd3699 # v4.0.1
|
||||
uses: android-actions/setup-android@be39fa834029ff78f1a44aa3bb0819b8fc2bd8fd # v4.0.4
|
||||
with:
|
||||
log-accepted-android-sdk-licenses: false
|
||||
|
||||
@@ -59,7 +59,7 @@ jobs:
|
||||
./gradlew build --no-daemon
|
||||
|
||||
ndk:
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: ubuntu-24.04 # previously ubuntu-latest
|
||||
container:
|
||||
image: 'ghcr.io/snapdragon-toolchain/arm64-android:v0.3'
|
||||
defaults:
|
||||
@@ -93,7 +93,7 @@ jobs:
|
||||
path: pkg-adb/llama.cpp
|
||||
|
||||
arm64:
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: ubuntu-24.04 # previously ubuntu-latest
|
||||
|
||||
env:
|
||||
NDK_VERSION: "29.0.14206865"
|
||||
@@ -123,7 +123,7 @@ jobs:
|
||||
distribution: temurin
|
||||
|
||||
- name: Setup Android SDK
|
||||
uses: android-actions/setup-android@40fd30fb8d7440372e1316f5d1809ec01dcd3699 # v4.0.1
|
||||
uses: android-actions/setup-android@be39fa834029ff78f1a44aa3bb0819b8fc2bd8fd # v4.0.4
|
||||
with:
|
||||
log-accepted-android-sdk-licenses: false
|
||||
|
||||
|
||||
@@ -41,8 +41,8 @@ jobs:
|
||||
|
||||
env:
|
||||
# Sync versions in build-openvino.yml, build-self-hosted.yml, release.yml, build-cache.yml, .devops/openvino.Dockerfile
|
||||
OPENVINO_VERSION_MAJOR: "2026.3.1"
|
||||
OPENVINO_VERSION_FULL: "2026.3.1.22476.56d9685302d"
|
||||
OPENVINO_VERSION_MAJOR: "2026.4"
|
||||
OPENVINO_VERSION_FULL: "2026.4.0.22959.99c81491cc3"
|
||||
|
||||
steps:
|
||||
- name: Clone
|
||||
@@ -69,8 +69,8 @@ jobs:
|
||||
|
||||
env:
|
||||
# Sync versions in build.yml, build-self-hosted.yml, release.yml, build-cache.yml, .devops/openvino.Dockerfile
|
||||
OPENVINO_VERSION_MAJOR: "2026.3.1"
|
||||
OPENVINO_VERSION_FULL: "2026.3.1.22476.56d9685302d"
|
||||
OPENVINO_VERSION_MAJOR: "2026.4"
|
||||
OPENVINO_VERSION_FULL: "2026.4.0.22959.99c81491cc3"
|
||||
|
||||
steps:
|
||||
- name: Clone
|
||||
|
||||
@@ -41,8 +41,8 @@ jobs:
|
||||
|
||||
env:
|
||||
# Sync versions in build-openvino.yml, build-self-hosted.yml, release.yml, build-cache.yml, .devops/openvino.Dockerfile
|
||||
OPENVINO_VERSION_MAJOR: "2026.3.1"
|
||||
OPENVINO_VERSION_FULL: "2026.3.1.22476.56d9685302d"
|
||||
OPENVINO_VERSION_MAJOR: "2026.4"
|
||||
OPENVINO_VERSION_FULL: "2026.4.0.22959.99c81491cc3"
|
||||
|
||||
steps:
|
||||
- name: Clone
|
||||
@@ -96,8 +96,8 @@ jobs:
|
||||
|
||||
env:
|
||||
# Sync versions in build-openvino.yml, build-self-hosted.yml, release.yml, build-cache.yml, .devops/openvino.Dockerfile
|
||||
OPENVINO_VERSION_MAJOR: "2026.3.1"
|
||||
OPENVINO_VERSION_FULL: "2026.3.1.22476.56d9685302d"
|
||||
OPENVINO_VERSION_MAJOR: "2026.4"
|
||||
OPENVINO_VERSION_FULL: "2026.4.0.22959.99c81491cc3"
|
||||
|
||||
steps:
|
||||
- name: Clone
|
||||
|
||||
@@ -412,8 +412,8 @@ jobs:
|
||||
|
||||
env:
|
||||
# Sync versions in build.yml, build-self-hosted.yml, release.yml, build-cache.yml, .devops/openvino.Dockerfile
|
||||
OPENVINO_VERSION_MAJOR: "2026.3.1"
|
||||
OPENVINO_VERSION_FULL: "2026.3.1.22476.56d9685302d"
|
||||
OPENVINO_VERSION_MAJOR: "2026.4"
|
||||
OPENVINO_VERSION_FULL: "2026.4.0.22959.99c81491cc3"
|
||||
|
||||
steps:
|
||||
- name: Clone
|
||||
|
||||
@@ -11,10 +11,12 @@ on:
|
||||
paths:
|
||||
- .github/workflows/copilot-setup-steps.yml
|
||||
|
||||
cache-mode: none
|
||||
|
||||
jobs:
|
||||
# The job MUST be called `copilot-setup-steps` or it will not be picked up by Copilot.
|
||||
copilot-setup-steps:
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: ubuntu-24.04
|
||||
|
||||
# Set the permissions to the lowest permissions possible needed for your steps.
|
||||
# Copilot will be given its own token for its operations.
|
||||
@@ -31,8 +33,8 @@ jobs:
|
||||
- name: ccache
|
||||
uses: ggml-org/[email protected]
|
||||
with:
|
||||
key: copilot-setup-steps
|
||||
evict-old-files: 1d
|
||||
restore: false
|
||||
save: false
|
||||
|
||||
- name: Dependencies
|
||||
id: depends
|
||||
|
||||
@@ -21,7 +21,7 @@ on:
|
||||
jobs:
|
||||
deploy:
|
||||
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: ubuntu-24.04 # previously ubuntu-latest
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
@@ -13,6 +13,16 @@ on:
|
||||
required: true
|
||||
type: boolean
|
||||
default: true
|
||||
skip_apiabi_check:
|
||||
description: 'Skip API/ABI compatibility check'
|
||||
required: false
|
||||
type: boolean
|
||||
default: false
|
||||
apiabi_compare_tag:
|
||||
description: 'Tag to compare against for API/ABI check (default: latest release)'
|
||||
required: false
|
||||
type: string
|
||||
default: ''
|
||||
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
@@ -23,7 +33,7 @@ permissions:
|
||||
|
||||
jobs:
|
||||
make-release:
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: ubuntu-24.04 # previously ubuntu-latest
|
||||
|
||||
steps:
|
||||
- name: Checkout
|
||||
@@ -33,12 +43,18 @@ jobs:
|
||||
ref: ${{ inputs.commit != '' && inputs.commit || github.ref_name }}
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Install API/ABI check tools
|
||||
if: ${{ github.event.inputs.skip_apiabi_check != 'true' }}
|
||||
run: sudo apt-get install -y abi-compliance-checker abigail-tools
|
||||
|
||||
- name: Run release checks
|
||||
id: checks
|
||||
run: bash scripts/make-release-checks.sh ${{ github.event.inputs.dry_run == 'true' && '--dry-run' || '' }}
|
||||
env:
|
||||
GITHUB_REPOSITORY: ${{ github.repository }}
|
||||
RELEASE_BRANCH: ${{ github.ref_name }}
|
||||
SKIP_APIABI_CHECK: ${{ github.event.inputs.skip_apiabi_check }}
|
||||
APIABI_COMPARE_TAG: ${{ github.event.inputs.apiabi_compare_tag }}
|
||||
|
||||
- name: Create release tag
|
||||
if: ${{ github.event.inputs.dry_run == 'false' }}
|
||||
|
||||
@@ -106,6 +106,7 @@ jobs:
|
||||
uses: ggml-org/[email protected]
|
||||
with:
|
||||
key: release-${{ matrix.os }}-${{ matrix.arch }}
|
||||
evict-old-files: 1d
|
||||
|
||||
- name: Build
|
||||
id: cmake_build
|
||||
@@ -190,6 +191,7 @@ jobs:
|
||||
uses: ggml-org/[email protected]
|
||||
with:
|
||||
key: release-${{ matrix.os }}-cpu
|
||||
evict-old-files: 1d
|
||||
|
||||
- name: Build
|
||||
id: cmake_build
|
||||
@@ -275,6 +277,7 @@ jobs:
|
||||
uses: ggml-org/[email protected]
|
||||
with:
|
||||
key: release-${{ matrix.os }}-vulkan
|
||||
evict-old-files: 1d
|
||||
|
||||
- name: Build
|
||||
id: cmake_build
|
||||
@@ -453,7 +456,7 @@ jobs:
|
||||
needs: [check-release, ui-build]
|
||||
if: ${{ needs.check-release.outputs.should_release == 'true' }}
|
||||
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: ubuntu-24.04 # previously ubuntu-latest
|
||||
|
||||
#permissions:
|
||||
# actions: write
|
||||
@@ -481,10 +484,9 @@ jobs:
|
||||
distribution: temurin
|
||||
|
||||
- name: Setup Android SDK
|
||||
uses: android-actions/setup-android@40fd30fb8d7440372e1316f5d1809ec01dcd3699 # v4.0.1
|
||||
uses: android-actions/setup-android@be39fa834029ff78f1a44aa3bb0819b8fc2bd8fd # v4.0.4
|
||||
with:
|
||||
log-accepted-android-sdk-licenses: false
|
||||
packages: 'platform-tools'
|
||||
|
||||
- name: Install NDK
|
||||
run: |
|
||||
@@ -501,6 +503,7 @@ jobs:
|
||||
# uses: ggml-org/[email protected]
|
||||
# with:
|
||||
# key: release-android-arm64
|
||||
# evict-old-files: 1d
|
||||
|
||||
- name: Build
|
||||
id: cmake_build
|
||||
@@ -555,8 +558,8 @@ jobs:
|
||||
|
||||
env:
|
||||
# Sync versions in build-openvino.yml, build-self-hosted.yml, release.yml, build-cache.yml, .devops/openvino.Dockerfile
|
||||
OPENVINO_VERSION_MAJOR: "2026.3.1"
|
||||
OPENVINO_VERSION_FULL: "2026.3.1.22476.56d9685302d"
|
||||
OPENVINO_VERSION_MAJOR: "2026.4"
|
||||
OPENVINO_VERSION_FULL: "2026.4.0.22959.99c81491cc3"
|
||||
|
||||
steps:
|
||||
- name: Set OpenVINO version output
|
||||
@@ -579,6 +582,7 @@ jobs:
|
||||
uses: ggml-org/[email protected]
|
||||
with:
|
||||
key: release-ubuntu-24.04-openvino-release-no-preset-v1
|
||||
evict-old-files: 1d
|
||||
|
||||
- name: Dependencies
|
||||
run: |
|
||||
@@ -669,8 +673,8 @@ jobs:
|
||||
|
||||
env:
|
||||
# Sync versions in build-openvino.yml, build-self-hosted.yml, release.yml, build-cache.yml, .devops/openvino.Dockerfile
|
||||
OPENVINO_VERSION_MAJOR: "2026.3.1"
|
||||
OPENVINO_VERSION_FULL: "2026.3.1.22476.56d9685302d"
|
||||
OPENVINO_VERSION_MAJOR: "2026.4"
|
||||
OPENVINO_VERSION_FULL: "2026.4.0.22959.99c81491cc3"
|
||||
|
||||
steps:
|
||||
- name: Set OpenVINO version output
|
||||
@@ -822,6 +826,7 @@ jobs:
|
||||
uses: ggml-org/[email protected]
|
||||
with:
|
||||
key: release-windows-2025-vs2026-${{ matrix.arch }}-cpu
|
||||
evict-old-files: 1d
|
||||
|
||||
- name: Build
|
||||
shell: cmd
|
||||
@@ -1066,6 +1071,7 @@ jobs:
|
||||
# uses: ggml-org/[email protected]
|
||||
# with:
|
||||
# key: release-windows-2025-${{ matrix.arch }}-${{ matrix.backend }}
|
||||
# evict-old-files: 1d
|
||||
|
||||
- name: Install OpenCL Headers and Libs
|
||||
id: install_opencl
|
||||
@@ -1154,6 +1160,7 @@ jobs:
|
||||
uses: ggml-org/[email protected]
|
||||
with:
|
||||
key: release-windows-2022-${{ matrix.arch }}-cuda-${{ matrix.cuda }}
|
||||
evict-old-files: 1d
|
||||
|
||||
- name: Build
|
||||
id: cmake_build
|
||||
@@ -1250,6 +1257,7 @@ jobs:
|
||||
uses: ggml-org/[email protected]
|
||||
with:
|
||||
key: release-windows-2022-x64-sycl
|
||||
evict-old-files: 1d
|
||||
|
||||
- name: Build
|
||||
id: cmake_build
|
||||
@@ -1368,6 +1376,7 @@ jobs:
|
||||
uses: ggml-org/[email protected]
|
||||
with:
|
||||
key: release-ubuntu-24.04-sycl-${{ matrix.build }}
|
||||
evict-old-files: 1d
|
||||
|
||||
- name: Build
|
||||
id: cmake_build
|
||||
|
||||
@@ -8,7 +8,7 @@ on:
|
||||
jobs:
|
||||
update:
|
||||
name: Update Winget Package
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: ubuntu-24.04 # previously ubuntu-latest
|
||||
if: github.repository_owner == 'ggml-org'
|
||||
|
||||
steps:
|
||||
|
||||
@@ -122,6 +122,8 @@ struct common_params_context {
|
||||
|
||||
// parse input arguments from CLI
|
||||
// if one argument has invalid value, it will automatically display usage of the specific argument (and not the full usage message)
|
||||
// TODO: this function can load ggml backend (by calling llama_support_rpc)
|
||||
// this is a side-effect that should be avoided
|
||||
bool common_params_parse(int argc, char ** argv, common_params & params, llama_example ex, void(*print_usage)(int, char **) = nullptr);
|
||||
|
||||
// load all backends and print the list of available (non-CPU) devices to stdout
|
||||
|
||||
@@ -1133,6 +1133,14 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
return common_chat_params_init_kimi_k3(tmpl, params);
|
||||
}
|
||||
|
||||
// Ling 3.0 / Bailing V3 - <role>X</role> sections with <arg_key>/<arg_value> tagged
|
||||
// tool calls. <role> sections are unique to this family among the tagged-arg templates.
|
||||
if (src.find("<role>ASSISTANT</role>") != std::string::npos &&
|
||||
src.find("<arg_key>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Ling 3.0 (Bailing V3)\n");
|
||||
return common_chat_params_init_ling3(tmpl, params);
|
||||
}
|
||||
|
||||
// Cohere2 MoE / North Code - marker-wrapped format with <|START_TEXT|> content and
|
||||
// <|START_ACTION|> JSON tool calls. <|START_TEXT|> is unique to this template (the older
|
||||
// Command-R templates use <|START_RESPONSE|>).
|
||||
|
||||
+8
-8
@@ -30,8 +30,8 @@ namespace hf_cache {
|
||||
|
||||
namespace fs = std::filesystem;
|
||||
|
||||
static fs::path get_cache_directory() {
|
||||
static const fs::path cache = []() {
|
||||
std::string get_cache_path() {
|
||||
static const std::string cache = []() {
|
||||
struct {
|
||||
const char * var;
|
||||
fs::path path;
|
||||
@@ -46,14 +46,14 @@ static fs::path get_cache_directory() {
|
||||
for (const auto & entry : entries) {
|
||||
if (auto * p = std::getenv(entry.var); p && *p) {
|
||||
fs::path base(p);
|
||||
return entry.path.empty() ? base : base / entry.path;
|
||||
return (entry.path.empty() ? base : base / entry.path).string();
|
||||
}
|
||||
}
|
||||
#ifndef _WIN32
|
||||
const struct passwd * pw = getpwuid(getuid());
|
||||
|
||||
if (pw && pw->pw_dir && *pw->pw_dir) {
|
||||
return fs::path(pw->pw_dir) / ".cache" / "huggingface" / "hub";
|
||||
return (fs::path(pw->pw_dir) / ".cache" / "huggingface" / "hub").string();
|
||||
}
|
||||
#endif
|
||||
throw std::runtime_error("Failed to determine HF cache directory");
|
||||
@@ -80,7 +80,7 @@ static std::string repo_to_folder_name(const std::string & repo_id) {
|
||||
}
|
||||
|
||||
static fs::path get_repo_path(const std::string & repo_id) {
|
||||
return get_cache_directory() / repo_to_folder_name(repo_id);
|
||||
return fs::path(get_cache_path()) / repo_to_folder_name(repo_id);
|
||||
}
|
||||
|
||||
static bool is_hex_char(const char c) {
|
||||
@@ -393,8 +393,8 @@ static std::string get_cached_ref(const fs::path & repo_path) {
|
||||
}
|
||||
|
||||
hf_files get_cached_files(const std::string & repo_id) {
|
||||
fs::path cache_dir = get_cache_directory();
|
||||
if (!fs::exists(cache_dir)) {
|
||||
const fs::path cache_path = get_cache_path();
|
||||
if (!fs::exists(cache_path)) {
|
||||
return {};
|
||||
}
|
||||
|
||||
@@ -405,7 +405,7 @@ hf_files get_cached_files(const std::string & repo_id) {
|
||||
|
||||
hf_files files;
|
||||
|
||||
for (const auto & repo : fs::directory_iterator(cache_dir)) {
|
||||
for (const auto & repo : fs::directory_iterator(cache_path)) {
|
||||
if (!repo.is_directory()) {
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -32,4 +32,7 @@ std::string finalize_file(const hf_file & file);
|
||||
// Remove the entire cached directory for a repo, returns true if removed
|
||||
bool remove_cached_repo(const std::string & repo_id);
|
||||
|
||||
// Returns the HuggingFace hub cache path
|
||||
std::string get_cache_path();
|
||||
|
||||
} // namespace hf_cache
|
||||
|
||||
@@ -321,7 +321,8 @@ static size_t gbnf_escape_length(const std::string & pattern, size_t pos) {
|
||||
case 'x': n_hex = 2; break;
|
||||
case 'u': n_hex = 4; break;
|
||||
case 'U': n_hex = 8; break;
|
||||
case 't': case 'r': case 'n': case '\\': case '"': case '[': case ']':
|
||||
// keep in sync with parse_char() in src/llama-grammar.cpp
|
||||
case 't': case 'r': case 'n': case '\\': case '"': case '[': case ']': case '-':
|
||||
return 2;
|
||||
default:
|
||||
return 0;
|
||||
|
||||
@@ -272,6 +272,10 @@ common_chat_params common_chat_params_init_gemma4(const common_chat_template &
|
||||
/* max = */ inputs.parallel_tool_calls ? -1 : 1
|
||||
));
|
||||
|
||||
if (inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED) {
|
||||
return start + thought + tool_call;
|
||||
}
|
||||
|
||||
auto scan_to_toolcall = p.rule("scan-to-toolcall", p.until("<|tool_call>"));
|
||||
auto content = p.rule("content", p.content(p.until_one_of({"<|channel>", "<channel|>", "<|tool_call>"})));
|
||||
auto message = p.rule("message", thought + content);
|
||||
|
||||
@@ -0,0 +1,194 @@
|
||||
#include "parsers.h"
|
||||
|
||||
// Ling 3.0 / Bailing V3 - <role>X</role> sections with tagged tool calls:
|
||||
// assistant := [<think> ... </think>] [content] {<tool_call>name
|
||||
// <arg_key>k</arg_key>\n<arg_value>v</arg_value> ...</tool_call>}
|
||||
// The generation prompt ends with "<role>ASSISTANT</role>\n<think>", so the model
|
||||
// never emits the opening think tag, and a tool call can arrive before any
|
||||
// </think>. Reasoning therefore terminates at the think close tag or at a tool
|
||||
// call start, like the Qwen3-Coder and Kimi K3 parsers. With thinking off the
|
||||
// template pre-closes the think block instead, and the model emits bare content.
|
||||
common_chat_params common_chat_params_init_ling3(const common_chat_template & tmpl,
|
||||
const autoparser::generation_params & inputs) {
|
||||
common_chat_params data;
|
||||
|
||||
data.prompt = common_chat_template_direct_apply_impl(tmpl, inputs);
|
||||
data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, inputs);
|
||||
data.format = COMMON_CHAT_FORMAT_PEG_NATIVE;
|
||||
data.supports_thinking = true;
|
||||
|
||||
const std::string ROLE = "<role>ASSISTANT</role>";
|
||||
const std::string THINK_START = "<think>";
|
||||
const std::string THINK_END = "</think>";
|
||||
const std::string CALL_START = "<tool_call>";
|
||||
const std::string CALL_END = "</tool_call>";
|
||||
const std::string ARG_KEY = "<arg_key>";
|
||||
const std::string ARG_KEY_END = "</arg_key>";
|
||||
const std::string ARG_VAL = "<arg_value>";
|
||||
const std::string ROLE_END = "<|role_end|>";
|
||||
const std::string ARG_VAL_END = "</arg_value>";
|
||||
|
||||
data.preserved_tokens = {
|
||||
THINK_START, THINK_END, CALL_START, CALL_END,
|
||||
ARG_KEY, ARG_KEY_END, ARG_VAL, ARG_VAL_END, ROLE_END,
|
||||
};
|
||||
|
||||
data.thinking_start_tag = THINK_START;
|
||||
// Support both </think> and <tool_call> as reasoning end sequences: a call
|
||||
// can be emitted before the think block is closed.
|
||||
data.thinking_end_tags = { THINK_END, CALL_START };
|
||||
|
||||
data.message_delimiters = {
|
||||
{ COMMON_CHAT_ROLE_ASSISTANT, "<role>ASSISTANT</role>" },
|
||||
{ COMMON_CHAT_ROLE_USER, "<role>HUMAN</role>" },
|
||||
{ COMMON_CHAT_ROLE_TOOL, "<role>OBSERVATION</role>" },
|
||||
{ COMMON_CHAT_ROLE_SYSTEM, "<role>SYSTEM</role>" },
|
||||
};
|
||||
|
||||
// the model may spell the end-of-turn control token out as text tokens,
|
||||
// which does not stop generation; a literal stop string catches it either
|
||||
// way (as the Laguna patch does for its </assistant> token)
|
||||
data.additional_stops = { ROLE_END };
|
||||
|
||||
if (inputs.has_continuation()) {
|
||||
const auto & msg = inputs.continue_msg;
|
||||
|
||||
data.generation_prompt = ROLE + "\n" + THINK_START + msg.reasoning_content;
|
||||
if (inputs.continue_final_message == COMMON_CHAT_CONTINUATION_CONTENT) {
|
||||
data.generation_prompt += THINK_END + msg.render_content();
|
||||
}
|
||||
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
// The generation prompt pre-opens the think block when thinking is on, so
|
||||
// the opening tag is optional here and reasoning runs until </think> or a
|
||||
// tool call start; with thinking off the template pre-closes the block and
|
||||
// everything the model emits is content.
|
||||
bool think_open = false;
|
||||
if (inputs.has_continuation()) {
|
||||
think_open = inputs.continue_final_message != COMMON_CHAT_CONTINUATION_CONTENT;
|
||||
} else {
|
||||
auto last_open = data.generation_prompt.rfind(THINK_START);
|
||||
auto last_close = data.generation_prompt.rfind(THINK_END);
|
||||
think_open = last_open != std::string::npos &&
|
||||
(last_close == std::string::npos || last_open > last_close);
|
||||
}
|
||||
|
||||
auto has_tools = inputs.tools.is_array() && !inputs.tools.empty();
|
||||
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||
auto include_grammar = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE;
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto end = p.end();
|
||||
|
||||
// the effective parse input is generation_prompt + model output, so the
|
||||
// assistant opener is optionally consumed here
|
||||
auto opener = p.optional(p.literal(ROLE) + p.optional(p.space()));
|
||||
|
||||
// the generation prompt pre-opens the think block, so the opening tag
|
||||
// is optional; a missing close tag does not swallow a tool call
|
||||
auto body_end = think_open ? p.until_one_of({ THINK_END, CALL_START }) : p.until_one_of({ THINK_END });
|
||||
auto think_body = extract_reasoning ? p.reasoning(body_end) : p.content(body_end);
|
||||
|
||||
auto reasoning = p.optional(p.optional(p.literal(THINK_START)) + think_body +
|
||||
p.optional(p.literal(THINK_END)));
|
||||
|
||||
// content between the think block and the first tool call, plus any
|
||||
// trailing text after the last tool call, are plain content
|
||||
auto content = p.optional(p.content(p.until_one_of({ CALL_START })));
|
||||
|
||||
// a trailing end-of-turn token is consumed instead of leaking into content
|
||||
auto tail = p.optional(p.content(p.until(ROLE_END))) + p.optional(p.literal(ROLE_END));
|
||||
|
||||
if (!has_tools || inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_NONE) {
|
||||
return opener + reasoning + tail + end;
|
||||
}
|
||||
|
||||
auto tool_choices = p.choice();
|
||||
auto arg_close = p.tool_arg_close(p.literal(ARG_VAL_END));
|
||||
auto arg_string = p.rule("ling3-arg-string",
|
||||
p.tool_arg_string_value(p.until(ARG_VAL_END)) + arg_close);
|
||||
|
||||
foreach_function(inputs.tools, [&](const json & tool) {
|
||||
const auto & function = tool.at("function");
|
||||
std::string name = function.at("name");
|
||||
|
||||
std::vector<common_peg_parser> required_args;
|
||||
std::vector<common_peg_parser> optional_args;
|
||||
|
||||
// each argument may be preceded by whitespace: the model emits
|
||||
// newlines between arguments, the template history does not
|
||||
foreach_parameter(function, [&](const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) {
|
||||
auto rule_name = "ling3-arg-" + name + "-" + param.name;
|
||||
|
||||
auto types = param.schema->value_types();
|
||||
|
||||
// string arguments are raw text up to the closing tag, other
|
||||
// types parse as JSON per their schema; each alternative
|
||||
// consumes the closing tag itself so a JSON prefix can not
|
||||
// commit the choice before the tag matches
|
||||
auto arg_value = p.eps();
|
||||
if (!types.has(common_chat_schema::TYPE_STRING)) {
|
||||
arg_value = p.tool_arg_json_value(p.schema(p.json(), rule_name + "-schema", doc, *param.schema)) + arg_close;
|
||||
} else if (types.is_only(common_chat_schema::TYPE_STRING)) {
|
||||
arg_value = arg_string;
|
||||
} else {
|
||||
// the parser tries the JSON alternative first to type the value
|
||||
arg_value = p.gbnf(p.atomic(p.tool_arg_json_value(p.schema(p.json(), rule_name + "-schema", doc, *param.schema)) + arg_close) | arg_string,
|
||||
"ling3-arg-string");
|
||||
}
|
||||
|
||||
auto arg = p.rule(rule_name,
|
||||
p.optional(p.space()) +
|
||||
p.tool_arg(p.tool_arg_open(p.literal(ARG_KEY) + p.tool_arg_name(p.literal(param.name)) +
|
||||
p.literal(ARG_KEY_END)) +
|
||||
p.optional(p.space()) + p.literal(ARG_VAL) +
|
||||
arg_value));
|
||||
|
||||
(param.required ? required_args : optional_args).push_back(arg);
|
||||
});
|
||||
|
||||
// required arguments in any order (as Qwen3-Coder does), then
|
||||
// optional ones in any order and number
|
||||
auto args = p.permute("ling3-" + name + "-args", required_args);
|
||||
if (!optional_args.empty()) {
|
||||
args = args + p.zero_or_more(p.choice(optional_args));
|
||||
}
|
||||
|
||||
auto call = p.tool(p.tool_open(p.literal(CALL_START) + p.tool_name(p.literal(name)) +
|
||||
p.optional(p.space())) +
|
||||
p.tool_args(args) +
|
||||
p.tool_close(p.optional(p.space()) + p.literal(CALL_END)));
|
||||
|
||||
tool_choices |= p.rule("ling3-tool-" + name, call);
|
||||
});
|
||||
|
||||
auto calls = inputs.parallel_tool_calls ?
|
||||
tool_choices + p.zero_or_more(p.space() + tool_choices) :
|
||||
tool_choices;
|
||||
|
||||
auto tools_section = p.trigger_rule("ling3-tool-call", calls + p.space() +
|
||||
p.optional(p.content(p.until(ROLE_END))) + p.optional(p.literal(ROLE_END)));
|
||||
|
||||
auto tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED ? tools_section :
|
||||
p.optional(tools_section);
|
||||
|
||||
return opener + reasoning + content + tools + tail + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
{ COMMON_GRAMMAR_TRIGGER_TYPE_WORD, CALL_START },
|
||||
};
|
||||
}
|
||||
|
||||
return data;
|
||||
}
|
||||
@@ -63,6 +63,8 @@ common_chat_params common_chat_params_init_kimi_k2(const common_chat_template &
|
||||
|
||||
common_chat_params common_chat_params_init_kimi_k3(const common_chat_template & tmpl, const autoparser::generation_params & inputs);
|
||||
|
||||
common_chat_params common_chat_params_init_ling3(const common_chat_template & tmpl, const autoparser::generation_params & inputs);
|
||||
|
||||
// tool_list_tokens preserves the LFM2 system tool-list markers; LFM2.5 renders without them
|
||||
common_chat_params common_chat_params_init_lfm2(const common_chat_template & tmpl, const autoparser::generation_params & inputs, bool tool_list_tokens);
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ set(LLAMA_CHAT_PARSERS_SOURCES
|
||||
${CMAKE_CURRENT_LIST_DIR}/gpt-oss.cpp
|
||||
${CMAKE_CURRENT_LIST_DIR}/kimi-k2.cpp
|
||||
${CMAKE_CURRENT_LIST_DIR}/kimi-k3.cpp
|
||||
${CMAKE_CURRENT_LIST_DIR}/ling3.cpp
|
||||
${CMAKE_CURRENT_LIST_DIR}/lfm2.cpp
|
||||
${CMAKE_CURRENT_LIST_DIR}/minicpm5.cpp
|
||||
${CMAKE_CURRENT_LIST_DIR}/minimax-m3.cpp
|
||||
|
||||
@@ -1861,6 +1861,9 @@ class TextModel(ModelBase):
|
||||
if chkhsh == "972da7b59cec44d1f0a490a86c96df53859e486e481563e5dddac155013d87ac":
|
||||
# ref: https://huggingface.co/poolside/Laguna-XS.2
|
||||
res = "laguna"
|
||||
if chkhsh == "653660222fb704f61cbf2b618a8ae6502b7f8b20c980f9a5de07ed78e13319cd":
|
||||
# ref: https://huggingface.co/ufakai/ufakzeka-1
|
||||
res = "ufakzeka"
|
||||
|
||||
if res is None:
|
||||
logger.warning("\n")
|
||||
|
||||
@@ -163,6 +163,7 @@ models = [
|
||||
{"name": "granite-embed-multi-311m", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/ibm-granite/granite-embedding-311m-multilingual-r2", },
|
||||
{"name": "mellum2", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/JetBrains/Mellum2-12B-A2.5B-Base"},
|
||||
{"name": "laguna", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/poolside/Laguna-XS.2", },
|
||||
{"name": "ufakzeka", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/ufakai/ufakzeka-1", },
|
||||
]
|
||||
|
||||
# some models are known to be broken upstream, so we will skip them as exceptions
|
||||
|
||||
+14
-12
@@ -12,6 +12,8 @@ The OpenVINO backend is implemented in `ggml/src/ggml-openvino` and provides a t
|
||||
- Compiles and caches the model for the target device.
|
||||
- Binds GGML tensor memory to OpenVINO inference tensors and runs inference.
|
||||
|
||||
For guidance on contributing to the OpenVINO backend, see the [OpenVINO Backend Contributing Guide](https://github.com/ravi9/llamacpp-ov-dev-guide/blob/main/contributing-llamacpp-ov.md).
|
||||
|
||||
## Contents
|
||||
|
||||
- [Supported Devices](#supported-devices)
|
||||
@@ -96,7 +98,7 @@ Although, the validated models below were tested with `llama-cli` using the `Q4_
|
||||
- **SL** = Stateless (`GGML_OPENVINO_STATEFUL_EXECUTION=0`)
|
||||
- **SF** = Stateful (`GGML_OPENVINO_STATEFUL_EXECUTION=1`)
|
||||
- Note: The NPU operates in stateless mode only.
|
||||
- **Validation system:** Intel® Core™ Ultra 5 238V (Lunar Lake) | 32 GB RAM | Ubuntu 24.04 | Intel OpenCL GPU Driver 26.31.39395.13-0 | Intel NPU Driver 1.35.0.
|
||||
- **Validation system:** Intel® Core™ Ultra 5 238V (Lunar Lake) | 32 GB RAM | Ubuntu 24.04 | Intel Graphics Compiler 2.41.5 | Intel OpenCL GPU Driver 26.31.39395.13-0 | Intel NPU Driver 1.38.0.
|
||||
- See [Known Limitations](#known-limitations) for context on observed failures.
|
||||
|
||||
| Model | CPU (SL / SF) | GPU (SL / SF) | NPU (SL) |
|
||||
@@ -117,9 +119,9 @@ Although, the validated models below were tested with `llama-cli` using the `Q4_
|
||||
| [lmstudio-community/Qwen3.5-9B-Q4_K_M](https://huggingface.co/lmstudio-community/Qwen3.5-9B-GGUF) | ✓ / ✗ | ✓ / ✗ | ✗ |
|
||||
| | | | |
|
||||
| [unsloth/gemma-3-4b-it-Q4_K_M](https://huggingface.co/unsloth/gemma-3-4b-it-GGUF) | ✓ / ✓ | ✓ / ✓ | ✓ |
|
||||
| [bartowski/google_gemma-4-E2B-it-Q4_K_M](https://huggingface.co/bartowski/google_gemma-4-E2B-it-GGUF) | ✓ / ✗ | ✓ / ✗ | ✗ |
|
||||
| [bartowski/google_gemma-4-E4B-it-Q4_K_M](https://huggingface.co/bartowski/google_gemma-4-E4B-it-GGUF) | ✓ / ✗ | ✓ / ✗ | ✓ |
|
||||
| [bartowski/gemma-4-12B-it-Q4_K_M](https://huggingface.co/bartowski/gemma-4-12B-it-GGUF) | ✓ / ✗ | ✓ / ✗ | ✓ |
|
||||
| [bartowski/google_gemma-4-E2B-it-Q4_K_M](https://huggingface.co/bartowski/google_gemma-4-E2B-it-GGUF) | ✓ / ✓ | ✓ / ✓ | ✗ |
|
||||
| [bartowski/google_gemma-4-E4B-it-Q4_K_M](https://huggingface.co/bartowski/google_gemma-4-E4B-it-GGUF) | ✓ / ✓ | ✓ / ✓ | ✓ |
|
||||
| [bartowski/gemma-4-12B-it-Q4_K_M](https://huggingface.co/bartowski/gemma-4-12B-it-GGUF) | ✓ / ✓ | ✓ / ✓ | ✗ |
|
||||
| | | | |
|
||||
| [bartowski/Phi-3-mini-4k-instruct-Q4_K_M](https://huggingface.co/bartowski/Phi-3-mini-4k-instruct-GGUF) | ✓ / ✓ | ✓ / ✓ | ✓ |
|
||||
| [bartowski/Phi-3.5-mini-instruct-Q4_K_M](https://huggingface.co/bartowski/Phi-3.5-mini-instruct-GGUF) | ✓ / ✓ | ✓ / ✓ | ✓ |
|
||||
@@ -132,9 +134,9 @@ Although, the validated models below were tested with `llama-cli` using the `Q4_
|
||||
| [bartowski/DeepSeek-R1-Distill-Llama-8B-Q4_K_M](https://huggingface.co/bartowski/DeepSeek-R1-Distill-Llama-8B-GGUF) | ✓ / ✓ | ✓ / ✓ | ✓ |
|
||||
| [bartowski/DeepSeek-R1-Distill-Qwen-7B-Q4_K_M](https://huggingface.co/bartowski/DeepSeek-R1-Distill-Qwen-7B-GGUF) | ✓ / ✓ | ✓ / ✓ | ✓ |
|
||||
| | | | |
|
||||
| [ibm-granite/granite-4.0-350m-Q4_K_M](https://huggingface.co/ibm-granite/granite-4.0-350m-GGUF) | ✓ / ✓ | ✗ / ✗ | ✓ |
|
||||
| [ibm-granite/granite-4.0-350m-Q4_K_M](https://huggingface.co/ibm-granite/granite-4.0-350m-GGUF) | ✓ / ✓ | ✓ / ✓ | ✓ |
|
||||
| [ibm-granite/granite-4.0-micro-Q4_K_M](https://huggingface.co/ibm-granite/granite-4.0-micro-GGUF) | ✓ / ✓ | ✓ / ✓ | ✓ |
|
||||
| [ibm-granite/granite-4.0-1b-Q4_K_M](https://huggingface.co/ibm-granite/granite-4.0-1b-GGUF) | ✓ / ✓ | ✗ / ✗ | ✗ |
|
||||
| [ibm-granite/granite-4.0-1b-Q4_K_M](https://huggingface.co/ibm-granite/granite-4.0-1b-GGUF) | ✓ / ✓ | ✓ / ✓ | ✗ |
|
||||
| [ibm-research/granite-3.2-8b-instruct-Q4_K_M](https://huggingface.co/ibm-research/granite-3.2-8b-instruct-GGUF) | ✓ / ✓ | ✓ / ✓ | ✓ |
|
||||
| | | | |
|
||||
| [HuggingFaceTB/smollm2-1.7b-instruct-q4_k_m](https://huggingface.co/HuggingFaceTB/SmolLM2-1.7B-Instruct-GGUF) | ✓ / ✓ | ✓ / ✓ | ✓ |
|
||||
@@ -242,8 +244,8 @@ chmod +x build-llamacpp-ov.sh
|
||||
# ============================================
|
||||
set -euo pipefail
|
||||
|
||||
OPENVINO_VERSION_MAJOR="2026.3.1"
|
||||
OPENVINO_VERSION_FULL="2026.3.1.22476.56d9685302d"
|
||||
OPENVINO_VERSION_MAJOR="2026.4"
|
||||
OPENVINO_VERSION_FULL="2026.4.0.22959.99c81491cc3"
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||
OPENVINO_INSTALL_DIR="/opt/intel/openvino_${OPENVINO_VERSION_MAJOR}"
|
||||
@@ -340,7 +342,7 @@ echo " ./build/ReleaseOV/bin/llama-cli -m model.gguf"
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> The script pins OpenVINO `2026.3.1` via the `OPENVINO_VERSION_MAJOR` / `OPENVINO_VERSION_FULL` variables at the top — edit them to track a different release.
|
||||
> The script pins OpenVINO `2026.4` via the `OPENVINO_VERSION_MAJOR` / `OPENVINO_VERSION_FULL` variables at the top — edit them to track a different release.
|
||||
|
||||
</details>
|
||||
|
||||
@@ -370,8 +372,8 @@ REM ============================================
|
||||
REM llama.cpp OpenVINO Build Script (Ninja)
|
||||
REM ============================================
|
||||
|
||||
set "OPENVINO_VERSION_MAJOR=2026.3.1"
|
||||
set "OPENVINO_VERSION_FULL=2026.3.1.22476.56d9685302d"
|
||||
set "OPENVINO_VERSION_MAJOR=2026.4"
|
||||
set "OPENVINO_VERSION_FULL=2026.4.0.22959.99c81491cc3"
|
||||
|
||||
set "SCRIPT_DIR=%~dp0"
|
||||
set "VCPKG_DIR=C:\vcpkg"
|
||||
@@ -550,7 +552,7 @@ endlocal
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> The script pins OpenVINO `2026.3.1` via the `OPENVINO_VERSION_MAJOR` / `OPENVINO_VERSION_FULL` variables at the top — edit them to track a different release. From any new shell, source the matching `setupvars` script via the junction — `call "C:\Intel\openvino\setupvars.bat"` from `cmd`, or `& "C:\Intel\openvino\setupvars.ps1"` from PowerShell. If `winget` cannot register Visual Studio Build Tools on first run, install them once manually and re-run the script from an elevated **Developer Command Prompt for VS 2022**.
|
||||
> The script pins OpenVINO `2026.4` via the `OPENVINO_VERSION_MAJOR` / `OPENVINO_VERSION_FULL` variables at the top — edit them to track a different release. From any new shell, source the matching `setupvars` script via the junction — `call "C:\Intel\openvino\setupvars.bat"` from `cmd`, or `& "C:\Intel\openvino\setupvars.ps1"` from PowerShell. If `winget` cannot register Visual Studio Build Tools on first run, install them once manually and re-run the script from an elevated **Developer Command Prompt for VS 2022**.
|
||||
|
||||
</details>
|
||||
|
||||
|
||||
+1
-1
@@ -121,7 +121,7 @@ Legend:
|
||||
| SWIGLU_OAI | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
|
||||
| TANH | ❌ | ✅ | ✅ | 🟡 | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
|
||||
| TIMESTEP_EMBEDDING | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ |
|
||||
| TOP_K | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ✅ | ❌ | 🟡 | 🟡 | ✅ | ❌ | ❌ |
|
||||
| TOP_K | ❌ | ❌ | ✅ | ❌ | ❌ | 🟡 | ✅ | ❌ | 🟡 | 🟡 | ✅ | ❌ | ❌ |
|
||||
| TRI | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
|
||||
| TRUNC | ❌ | ❌ | ✅ | 🟡 | 🟡 | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
|
||||
| UPSCALE | ❌ | 🟡 | ✅ | ✅ | ❌ | ❌ | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
|
||||
|
||||
+298
-292
@@ -13936,247 +13936,247 @@
|
||||
"HTP0","ARGSORT","type=f32,ne=[2049,2,1,3],order=1","support","1","yes","HTP"
|
||||
"HTP0","ARGSORT","type=f32,ne=[2,8,8192,1],order=1","support","1","yes","HTP"
|
||||
"HTP0","ARGSORT","type=f32,ne=[2048,512,1,1],order=1","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[12,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[13,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[13,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[15,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[15,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[15,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[19,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[19,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[19,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[19,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[27,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[27,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[27,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[27,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[27,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[43,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[43,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[43,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[43,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[43,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[64,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[75,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[64,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[75,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[64,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[75,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[64,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[75,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[64,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[75,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[12,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[13,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[13,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[15,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[15,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[15,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[19,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[19,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[19,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[19,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[27,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[27,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[27,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[27,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[27,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[43,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[43,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[43,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[43,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[43,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[64,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[75,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[64,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[75,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[64,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[75,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[64,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[75,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[64,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[75,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[128,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[139,1,2,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[256,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[267,1,2,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[512,1,1,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[523,1,2,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1035,1,2,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,1,1,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2059,1,2,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4107,1,2,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8203,1,2,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=9999,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16395,1,2,1],k=9999,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32768,1,1,1],k=9999,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[32779,1,2,1],k=9999,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65547,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65547,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65547,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65547,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65547,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65547,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65547,1,2,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65547,1,2,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65536,1,1,1],k=9999,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[65547,1,2,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131083,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131083,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131083,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131083,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131083,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131083,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131083,1,2,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131083,1,2,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131072,1,1,1],k=9999,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[131083,1,2,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262155,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262155,1,2,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262155,1,2,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262155,1,2,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262155,1,2,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=100,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262155,1,2,1],k=100,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=500,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262155,1,2,1],k=500,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=1023,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262155,1,2,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262144,1,1,1],k=9999,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[262155,1,2,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[524288,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[524299,1,2,1],k=1,ties=0","support","0","no","HTP"
|
||||
@@ -14196,99 +14196,105 @@
|
||||
"HTP0","TOP_K","type=f32,ne=[524299,1,2,1],k=1023,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[524288,1,1,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[524299,1,2,1],k=9999,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,1,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,1,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,1,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,1,1,1],k=4,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,1,1,1],k=4,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=4,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,1,1,1],k=4,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,8,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,8,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,8,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,8,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,8,1,1],k=4,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,8,1,1],k=4,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,16,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,16,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,16,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,16,1,1],k=4,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,1,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,1,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,1,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,16,1,1],k=4,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,16,1,1],k=4,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,1,1,1],k=8,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,1,1,1],k=8,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=8,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,1,1,1],k=8,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,8,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,8,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,8,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,8,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,8,1,1],k=8,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,8,1,1],k=8,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,16,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,16,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,16,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,16,1,1],k=8,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,1,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,1,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,1,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,16,1,1],k=8,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,16,1,1],k=8,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,1,1,1],k=16,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,1,1,1],k=16,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=16,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,1,1,1],k=16,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,8,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,8,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,8,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,8,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,8,1,1],k=16,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,8,1,1],k=16,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,16,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,16,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,16,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,16,1,1],k=16,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,1,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,1,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,1,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,16,1,1],k=16,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,16,1,1],k=16,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,1,1,1],k=32,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,1,1,1],k=32,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,1,1,1],k=32,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,1,1,1],k=32,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,8,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,8,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,8,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,8,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,8,1,1],k=32,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,8,1,1],k=32,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[202048,16,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[151936,16,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,16,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,16,1,1],k=32,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,10,10,10],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[60,10,10,10],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1023,2,1,3],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,2,1,3],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1025,2,1,3],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2047,2,1,3],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,2,1,3],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2049,2,1,3],k=1,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,10,10,10],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[60,10,10,10],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1023,2,1,3],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,2,1,3],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1025,2,1,3],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2047,2,1,3],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,2,1,3],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2049,2,1,3],k=2,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,10,10,10],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[60,10,10,10],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1023,2,1,3],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,2,1,3],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1025,2,1,3],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2047,2,1,3],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,2,1,3],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2049,2,1,3],k=3,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,10,10,10],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[60,10,10,10],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1023,2,1,3],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,2,1,3],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1025,2,1,3],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2047,2,1,3],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,2,1,3],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2049,2,1,3],k=7,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,10,10,10],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[60,10,10,10],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1023,2,1,3],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,2,1,3],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1025,2,1,3],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2047,2,1,3],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,2,1,3],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2049,2,1,3],k=15,ties=0","support","0","no","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,16,1,1],k=32,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8193,16,1,1],k=32,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,10,10,10],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[60,10,10,10],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1023,2,1,3],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,2,1,3],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1025,2,1,3],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2047,2,1,3],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,2,1,3],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2049,2,1,3],k=1,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,10,10,10],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[60,10,10,10],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1023,2,1,3],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,2,1,3],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1025,2,1,3],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2047,2,1,3],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,2,1,3],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2049,2,1,3],k=2,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,10,10,10],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[60,10,10,10],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1023,2,1,3],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,2,1,3],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1025,2,1,3],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2047,2,1,3],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,2,1,3],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2049,2,1,3],k=3,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,10,10,10],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[60,10,10,10],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1023,2,1,3],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,2,1,3],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1025,2,1,3],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2047,2,1,3],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,2,1,3],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2049,2,1,3],k=7,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16,10,10,10],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[60,10,10,10],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1023,2,1,3],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,2,1,3],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1025,2,1,3],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[16384,1,1,1],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2047,2,1,3],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,2,1,3],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2049,2,1,3],k=15,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[1024,1,1,1],k=1024,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[2048,2,1,1],k=1024,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[4096,1,1,1],k=2048,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[8192,2,1,1],k=2051,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[33024,1,1,1],k=2051,ties=0","support","1","yes","HTP"
|
||||
"HTP0","TOP_K","type=f32,ne=[33024,4,1,1],k=2051,ties=0","support","1","yes","HTP"
|
||||
"HTP0","UPSCALE","type=f32,ne=[512,512,3,2],scale_factor=2,mode=nearest,transpose=0","support","0","no","HTP"
|
||||
"HTP0","UPSCALE","type=f32,ne=[512,512,3,2],scale_factor=2,mode=nearest,transpose=1","support","0","no","HTP"
|
||||
"HTP0","UPSCALE","type=f32,ne=[2,5,7,11],ne_tgt=[5,7,11,13],mode=nearest","support","0","no","HTP"
|
||||
|
||||
|
Can't render this file because it is too large.
|
@@ -1,5 +1,7 @@
|
||||
set(TARGET llama-convert-llama2c-to-ggml)
|
||||
add_executable(${TARGET} convert-llama2c-to-ggml.cpp)
|
||||
install(TARGETS ${TARGET} RUNTIME)
|
||||
target_link_libraries(${TARGET} PRIVATE llama-common llama ${CMAKE_THREAD_LIBS_INIT})
|
||||
target_compile_features(${TARGET} PRIVATE cxx_std_17)
|
||||
if (GGML_CPU)
|
||||
set(TARGET llama-convert-llama2c-to-ggml)
|
||||
add_executable(${TARGET} convert-llama2c-to-ggml.cpp)
|
||||
install(TARGETS ${TARGET} RUNTIME)
|
||||
target_link_libraries(${TARGET} PRIVATE llama-common llama ${CMAKE_THREAD_LIBS_INIT})
|
||||
target_compile_features(${TARGET} PRIVATE cxx_std_17)
|
||||
endif()
|
||||
|
||||
@@ -11,6 +11,7 @@
|
||||
#include <clocale>
|
||||
#include <cmath>
|
||||
#include <cstdio>
|
||||
#include <random>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
#include <ctime>
|
||||
@@ -156,7 +157,7 @@ static std::vector<std::string> split_string(const std::string& input, char deli
|
||||
int main(int argc, char ** argv) {
|
||||
std::setlocale(LC_NUMERIC, "C");
|
||||
|
||||
srand(1234);
|
||||
std::mt19937 rng(1234);
|
||||
|
||||
common_params params;
|
||||
|
||||
@@ -321,7 +322,7 @@ int main(int argc, char ** argv) {
|
||||
client.t_start_prompt = ggml_time_us();
|
||||
client.t_start_gen = 0;
|
||||
|
||||
client.input = k_prompts[rand() % k_prompts.size()];
|
||||
client.input = k_prompts[rng() % k_prompts.size()];
|
||||
client.response = "";
|
||||
|
||||
// construct the prompt:
|
||||
@@ -334,10 +335,10 @@ int main(int argc, char ** argv) {
|
||||
client.prompt += k_system;
|
||||
}
|
||||
|
||||
const int n_junk_cur = rand() % n_junk;
|
||||
const int n_junk_cur = rng() % n_junk;
|
||||
|
||||
for (int i = 0; i < n_junk_cur; ++i) {
|
||||
const int r = rand() % k_questions.size();
|
||||
const int r = rng() % k_questions.size();
|
||||
client.prompt += "User:\n" + k_questions[r] + "\nAssistant:\n " + k_answers[r] + "\n";
|
||||
}
|
||||
client.prompt += "User:\n" + client.input + "\nAssistant:\n";
|
||||
|
||||
@@ -1630,7 +1630,10 @@ static bool ggml_backend_sched_alloc_splits(ggml_backend_sched_t sched) {
|
||||
ggml_backend_synchronize(sched->backends[i]);
|
||||
}
|
||||
|
||||
ggml_gallocr_reserve_n(sched->galloc, &sched->graph, sched->node_backend_ids, sched->leaf_backend_ids);
|
||||
if (!ggml_gallocr_reserve_n(sched->galloc, &sched->graph, sched->node_backend_ids, sched->leaf_backend_ids)) {
|
||||
GGML_LOG_ERROR("%s: failed to reserve graph buffers\n", __func__);
|
||||
return false;
|
||||
}
|
||||
if (!ggml_gallocr_alloc_graph(sched->galloc, &sched->graph)) {
|
||||
GGML_LOG_ERROR("%s: failed to allocate graph\n", __func__);
|
||||
return false;
|
||||
|
||||
@@ -1329,7 +1329,9 @@ UseGgmlGemm1:;
|
||||
const size_t nbw3 = nbw2*ne12;
|
||||
|
||||
assert(params->wsize >= ne13*nbw3);
|
||||
GGML_ASSERT(src1->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(src1->type == GGML_TYPE_F32 || src1->type == GGML_TYPE_F16);
|
||||
// the F16 path below writes plain floats into wdata, so it needs an F32 vec_dot_type
|
||||
GGML_ASSERT(src1->type == GGML_TYPE_F32 || vec_dot_type == GGML_TYPE_F32);
|
||||
|
||||
#if 0
|
||||
for (int64_t i13 = 0; i13 < ne13; ++i13) {
|
||||
@@ -1348,9 +1350,15 @@ UseGgmlGemm1:;
|
||||
size_t bs = ggml_blck_size(vec_dot_type);
|
||||
int64_t ne10_block_start = (ith * ne10/bs) / nth;
|
||||
int64_t ne10_block_end = ((ith + 1) * ne10/bs) / nth;
|
||||
from_float((float *)((char *) src1->data + i13*nb13 + i12*nb12 + i11*nb11 + ne10_block_start*bs*nb10),
|
||||
(void *) (wdata + i13*nbw3 + i12*nbw2 + i11*nbw1 + ne10_block_start*nbw0),
|
||||
(ne10_block_end - ne10_block_start) * bs);
|
||||
const char * src1_block = (const char *) src1->data + i13*nb13 + i12*nb12 + i11*nb11 + ne10_block_start*bs*nb10;
|
||||
char * dst_block = wdata + i13*nbw3 + i12*nbw2 + i11*nbw1 + ne10_block_start*nbw0;
|
||||
const int64_t n_block = (ne10_block_end - ne10_block_start) * bs;
|
||||
|
||||
if (src1->type == GGML_TYPE_F32) {
|
||||
from_float((const float *) src1_block, dst_block, n_block);
|
||||
} else {
|
||||
ggml_cpu_fp16_to_fp32((const ggml_fp16_t *) src1_block, (float *) dst_block, n_block);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -451,6 +451,10 @@ static bool ggml_backend_cpu_device_supports_op(ggml_backend_dev_t dev, const st
|
||||
op->type != GGML_TYPE_IQ1_S &&
|
||||
op->type != GGML_TYPE_IQ1_M; // missing type_traits.from_float
|
||||
case GGML_OP_MUL_MAT:
|
||||
if (ggml_get_op_params_i32(op, 1) == GGML_HINT_SRC0_IS_HADAMARD &&
|
||||
src0->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32) {
|
||||
return src1->type == GGML_TYPE_F32 || src1->type == GGML_TYPE_F16;
|
||||
}
|
||||
return src1->type == GGML_TYPE_F32 || src1->type == ggml_get_type_traits_cpu(src0->type)->vec_dot_type;
|
||||
case GGML_OP_SOFT_MAX_BACK: {
|
||||
if (op->src[0]->type != GGML_TYPE_F32 || op->src[1]->type != GGML_TYPE_F32) {
|
||||
|
||||
@@ -12015,11 +12015,20 @@ void ggml_compute_forward_opt_step_sgd(const ggml_compute_params * params, ggml_
|
||||
}
|
||||
}
|
||||
|
||||
static void ggml_compute_forward_fwht_f32(const ggml_compute_params * params, ggml_tensor * dst) {
|
||||
static inline float ggml_fwht_load(const float value) {
|
||||
return value;
|
||||
}
|
||||
|
||||
static inline float ggml_fwht_load(const ggml_fp16_t value) {
|
||||
return ggml_fp16_to_fp32(value);
|
||||
}
|
||||
|
||||
template<typename src_t>
|
||||
static void ggml_compute_forward_fwht_impl(const ggml_compute_params * params, ggml_tensor * dst) {
|
||||
const ggml_tensor * src0 = dst->src[0];
|
||||
const ggml_tensor * src1 = dst->src[1];
|
||||
|
||||
GGML_ASSERT(src1->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(src1->type == (std::is_same_v<src_t, float> ? GGML_TYPE_F32 : GGML_TYPE_F16));
|
||||
GGML_ASSERT(dst->type == GGML_TYPE_F32);
|
||||
|
||||
GGML_TENSOR_BINARY_OP_LOCALS
|
||||
@@ -12046,11 +12055,11 @@ static void ggml_compute_forward_fwht_f32(const ggml_compute_params * params, gg
|
||||
const int64_t i12 = (r - i13 * ne11 * ne12) / ne11;
|
||||
const int64_t i11 = r - i13 * ne11 * ne12 - i12 * ne11;
|
||||
|
||||
const float * src_row = (const float *) ((const char *) src1->data + i11 * nb11 + i12 * nb12 + i13 * nb13);
|
||||
const src_t * src_row = (const src_t *) ((const char *) src1->data + i11 * nb11 + i12 * nb12 + i13 * nb13);
|
||||
float * dst_row = (float *) ((char *) dst->data + i11 * nb1 + i12 * nb2 + i13 * nb3);
|
||||
|
||||
for (int64_t j = 0; j < n; j++) {
|
||||
dst_row[j] = src_row[j] * scale;
|
||||
dst_row[j] = ggml_fwht_load(src_row[j]) * scale;
|
||||
}
|
||||
|
||||
// Scalar passes
|
||||
@@ -12097,12 +12106,17 @@ void ggml_compute_forward_fwht(const ggml_compute_params * params, ggml_tensor *
|
||||
switch (src1->type) {
|
||||
case GGML_TYPE_F32:
|
||||
{
|
||||
ggml_compute_forward_fwht_f32(params, dst);
|
||||
ggml_compute_forward_fwht_impl<float>(params, dst);
|
||||
}
|
||||
break;
|
||||
case GGML_TYPE_F16:
|
||||
{
|
||||
ggml_compute_forward_fwht_impl<ggml_fp16_t>(params, dst);
|
||||
}
|
||||
break;
|
||||
default:
|
||||
{
|
||||
GGML_ABORT("fatal error - fwht is F32 only");
|
||||
GGML_ABORT("fatal error - fwht supports F32 and F16 input");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -51,9 +51,12 @@ void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool,
|
||||
cudaStream_t stream) {
|
||||
ggml_cuda_pool_alloc<int> temp_indices_alloc(pool, ncols * nrows);
|
||||
ggml_cuda_pool_alloc<float> temp_keys_alloc(pool, ncols * nrows);
|
||||
// Device*Sort algorithms currently do not allow for in-place sorting/aliasing of input/outputs
|
||||
ggml_cuda_pool_alloc<float> temp_keys_out_alloc(pool, ncols * nrows);
|
||||
|
||||
int * temp_indices = temp_indices_alloc.get();
|
||||
float * temp_keys = temp_keys_alloc.get();
|
||||
float * temp_keys_out = temp_keys_out_alloc.get();
|
||||
|
||||
static const int block_size = 256;
|
||||
const dim3 grid_size((ncols + block_size - 1) / block_size, nrows);
|
||||
@@ -85,18 +88,18 @@ void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool,
|
||||
|
||||
if (order == GGML_SORT_ORDER_ASC) {
|
||||
if (nrows == 1) {
|
||||
CUDA_CHECK(DeviceRadixSort::SortPairs(nullptr, temp_storage_bytes, temp_keys, temp_keys, // keys (in-place)
|
||||
CUDA_CHECK(DeviceRadixSort::SortPairs(nullptr, temp_storage_bytes, temp_keys, temp_keys_out, // keys in, keys out
|
||||
temp_indices, dst, // values (indices)
|
||||
ncols, 0, sizeof(float) * 8, stream));
|
||||
} else if (is_capturing) {
|
||||
CUDA_CHECK(DeviceSegmentedRadixSort::SortPairs(
|
||||
nullptr, temp_storage_bytes, temp_keys, temp_keys, // keys (in-place)
|
||||
nullptr, temp_storage_bytes, temp_keys, temp_keys_out, // keys in, keys out
|
||||
temp_indices, dst, // values (indices)
|
||||
ncols * nrows, nrows, // num items, num segments
|
||||
offset_iterator, offset_iterator + 1, 0, sizeof(float) * 8, stream));
|
||||
} else {
|
||||
CUDA_CHECK(DeviceSegmentedSort::SortPairs(nullptr, temp_storage_bytes, temp_keys,
|
||||
temp_keys, // keys (in-place)
|
||||
temp_keys_out, // keys out
|
||||
temp_indices, dst, // values (indices)
|
||||
ncols * nrows, nrows, // num items, num segments
|
||||
offset_iterator, offset_iterator + 1, stream));
|
||||
@@ -104,15 +107,15 @@ void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool,
|
||||
} else {
|
||||
if (nrows == 1) {
|
||||
CUDA_CHECK(DeviceRadixSort::SortPairsDescending(nullptr, temp_storage_bytes, temp_keys,
|
||||
temp_keys, // keys (in-place)
|
||||
temp_keys_out, // keys out
|
||||
temp_indices, dst, // values (indices)
|
||||
ncols, 0, sizeof(float) * 8, stream));
|
||||
} else if (is_capturing) {
|
||||
CUDA_CHECK(DeviceSegmentedRadixSort::SortPairsDescending(
|
||||
nullptr, temp_storage_bytes, temp_keys, temp_keys, temp_indices, dst, ncols * nrows, nrows,
|
||||
nullptr, temp_storage_bytes, temp_keys, temp_keys_out, temp_indices, dst, ncols * nrows, nrows,
|
||||
offset_iterator, offset_iterator + 1, 0, sizeof(float) * 8, stream));
|
||||
} else {
|
||||
CUDA_CHECK(DeviceSegmentedSort::SortPairsDescending(nullptr, temp_storage_bytes, temp_keys, temp_keys,
|
||||
CUDA_CHECK(DeviceSegmentedSort::SortPairsDescending(nullptr, temp_storage_bytes, temp_keys, temp_keys_out,
|
||||
temp_indices, dst, ncols * nrows, nrows,
|
||||
offset_iterator, offset_iterator + 1, stream));
|
||||
}
|
||||
@@ -124,31 +127,31 @@ void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool,
|
||||
if (order == GGML_SORT_ORDER_ASC) {
|
||||
if (nrows == 1) {
|
||||
CUDA_CHECK(DeviceRadixSort::SortPairs(d_temp_storage, temp_storage_bytes, temp_keys,
|
||||
temp_keys, // keys (in-place)
|
||||
temp_keys_out, // keys out
|
||||
temp_indices, dst, // values (indices)
|
||||
ncols, 0, sizeof(float) * 8, stream));
|
||||
} else if (is_capturing) {
|
||||
CUDA_CHECK(DeviceSegmentedRadixSort::SortPairs(d_temp_storage, temp_storage_bytes, temp_keys, temp_keys,
|
||||
CUDA_CHECK(DeviceSegmentedRadixSort::SortPairs(d_temp_storage, temp_storage_bytes, temp_keys, temp_keys_out,
|
||||
temp_indices, dst, ncols * nrows, nrows, offset_iterator,
|
||||
offset_iterator + 1, 0, sizeof(float) * 8, stream));
|
||||
} else {
|
||||
CUDA_CHECK(DeviceSegmentedSort::SortPairs(d_temp_storage, temp_storage_bytes, temp_keys, temp_keys,
|
||||
CUDA_CHECK(DeviceSegmentedSort::SortPairs(d_temp_storage, temp_storage_bytes, temp_keys, temp_keys_out,
|
||||
temp_indices, dst, ncols * nrows, nrows, offset_iterator,
|
||||
offset_iterator + 1, stream));
|
||||
}
|
||||
} else {
|
||||
if (nrows == 1) {
|
||||
CUDA_CHECK(DeviceRadixSort::SortPairsDescending(d_temp_storage, temp_storage_bytes, temp_keys,
|
||||
temp_keys, // keys (in-place)
|
||||
temp_keys_out, // keys out
|
||||
temp_indices, dst, // values (indices)
|
||||
ncols, 0, sizeof(float) * 8, stream));
|
||||
} else if (is_capturing) {
|
||||
CUDA_CHECK(DeviceSegmentedRadixSort::SortPairsDescending(
|
||||
d_temp_storage, temp_storage_bytes, temp_keys, temp_keys, temp_indices, dst, ncols * nrows, nrows,
|
||||
d_temp_storage, temp_storage_bytes, temp_keys, temp_keys_out, temp_indices, dst, ncols * nrows, nrows,
|
||||
offset_iterator, offset_iterator + 1, 0, sizeof(float) * 8, stream));
|
||||
} else {
|
||||
CUDA_CHECK(DeviceSegmentedSort::SortPairsDescending(d_temp_storage, temp_storage_bytes, temp_keys,
|
||||
temp_keys, temp_indices, dst, ncols * nrows, nrows,
|
||||
temp_keys_out, temp_indices, dst, ncols * nrows, nrows,
|
||||
offset_iterator, offset_iterator + 1, stream));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1580,6 +1580,9 @@ struct ggml_cuda_mm_fusion_args_host {
|
||||
const ggml_tensor * gate_scale = nullptr;
|
||||
ggml_glu_op glu_op;
|
||||
float glu_limit = 0.0f;
|
||||
const ggml_tensor * shared_up = nullptr;
|
||||
const ggml_tensor * shared_gate = nullptr;
|
||||
ggml_tensor * shared_dst = nullptr;
|
||||
};
|
||||
struct ggml_cuda_mm_fusion_args_device {
|
||||
const void * x_bias = nullptr;
|
||||
@@ -1589,6 +1592,9 @@ struct ggml_cuda_mm_fusion_args_device {
|
||||
const void * gate_scale = nullptr;
|
||||
ggml_glu_op glu_op;
|
||||
float glu_limit = 0.0f;
|
||||
const void * shared_up = nullptr;
|
||||
const void * shared_gate = nullptr;
|
||||
float * shared_dst = nullptr;
|
||||
};
|
||||
|
||||
struct ggml_cuda_kernel_launch_params {
|
||||
|
||||
@@ -719,7 +719,7 @@ static __global__ void flash_attn_mask_to_KV_max(
|
||||
}
|
||||
|
||||
void ggml_cuda_flash_attn_ext_compact_mask(
|
||||
const ggml_tensor * mask, int32_t * indices, int32_t n_kv_max, cudaStream_t stream);
|
||||
const ggml_tensor * mask, int32_t * indices, int32_t * counts, int32_t n_queries, int32_t ncols1, int32_t n_kv_max, cudaStream_t stream);
|
||||
|
||||
template<int D, int ncols1, int ncols2> // D == head size
|
||||
__launch_bounds__(D, 1)
|
||||
@@ -1092,14 +1092,18 @@ void launch_fattn(
|
||||
const int ntiles_z_gqa = ((gqa_ratio + ncols2 - 1) / ncols2);
|
||||
const int ntiles_dst = ntiles_x * ntiles_z_gqa * K->ne[2] * Q->ne[3];
|
||||
|
||||
const int32_t n_kv_max = use_sparse ? ggml_get_op_params_i32(KQV, 4) : 0;
|
||||
// sparse: a query tile of ncols1 queries shares one index list, the union of the queries' visible columns
|
||||
int32_t n_kv_max = 0;
|
||||
if (use_sparse) {
|
||||
GGML_ASSERT(mask != nullptr);
|
||||
GGML_ASSERT(n_kv_max > 0);
|
||||
const size_t mask_rows = size_t(mask->ne[1]) * mask->ne[3];
|
||||
const int32_t n_kv_max_query = ggml_get_op_params_i32(KQV, 4);
|
||||
GGML_ASSERT(n_kv_max_query > 0);
|
||||
n_kv_max = std::min<int64_t>(K->ne[1], int64_t(ncols1)*n_kv_max_query);
|
||||
|
||||
KV_max.alloc(size_t(n_kv_max) * mask_rows);
|
||||
ggml_cuda_flash_attn_ext_compact_mask(mask, KV_max.ptr, n_kv_max, main_stream);
|
||||
const size_t n_lists = size_t(ntiles_x) * mask->ne[3];
|
||||
|
||||
KV_max.alloc(size_t(n_kv_max)*n_lists + n_lists);
|
||||
ggml_cuda_flash_attn_ext_compact_mask(mask, KV_max.ptr, KV_max.ptr + size_t(n_kv_max)*n_lists, Q->ne[1], ncols1, n_kv_max, main_stream);
|
||||
}
|
||||
|
||||
// Optional optimization where the mask is scanned to determine whether part of the calculation can be skipped.
|
||||
|
||||
@@ -1760,7 +1760,9 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile(
|
||||
static constexpr __host__ __device__ bool ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(
|
||||
const int DKQ, const int DV, const int ncols1, const int ncols2) {
|
||||
return (DKQ == 512 && DV == 512 && ncols1 == 1 && ncols2 == 8) ||
|
||||
(DKQ == 576 && DV == 512 && ncols1 == 1 && ncols2 == 16);
|
||||
(DKQ == 576 && DV == 512 && ncols1 == 1 && ncols2 == 16) ||
|
||||
(DKQ == 256 && DV == 256 && ncols1 == 1 && ncols2 == 8) ||
|
||||
(DKQ == 256 && DV == 256 && ncols1 == 8 && ncols2 == 8);
|
||||
}
|
||||
|
||||
template<int DKQ, int DV, int ncols1, int ncols2, bool use_logit_softcap, bool V_is_K_view, bool use_sparse>
|
||||
@@ -1794,8 +1796,9 @@ static __global__ void flash_attn_ext_f16(
|
||||
const char * GGML_CUDA_RESTRICT V = V_ptr;
|
||||
const char * GGML_CUDA_RESTRICT mask = mask_ptr;
|
||||
const char * GGML_CUDA_RESTRICT sinks = sinks_ptr;
|
||||
const int * GGML_CUDA_RESTRICT KV_max = use_sparse ? nullptr : KV_max_ptr;
|
||||
// sparse: one index list per (sequence, query tile), the live count of each list follows the lists
|
||||
const int * GGML_CUDA_RESTRICT sparse_indices = use_sparse ? KV_max_ptr : nullptr;
|
||||
const int * GGML_CUDA_RESTRICT KV_max = KV_max_ptr;
|
||||
float * GGML_CUDA_RESTRICT dst = dst_ptr;
|
||||
float2 * GGML_CUDA_RESTRICT dst_meta = dst_meta_ptr;
|
||||
|
||||
@@ -1860,6 +1863,10 @@ static __global__ void flash_attn_ext_f16(
|
||||
const int iter_j = (ne01.z + (ncols1 - 1)) / ncols1;
|
||||
const int iter_z_gqa = (gqa_ratio + (ncols2 - 1)) / ncols2;
|
||||
|
||||
if (use_sparse) {
|
||||
KV_max = KV_max_ptr + int64_t(iter_j)*ne33*ne11;
|
||||
}
|
||||
|
||||
// kbc == k block continuous, current index in continuous ijk space.
|
||||
int kbc = int64_t(blockIdx.x + 0)*(iter_k*iter_j*iter_z_gqa*ne12*ne03) / gridDim.x;
|
||||
const int kbc_stop = int64_t(blockIdx.x + 1)*(iter_k*iter_j*iter_z_gqa*ne12*ne03) / gridDim.x;
|
||||
@@ -1889,11 +1896,13 @@ static __global__ void flash_attn_ext_f16(
|
||||
|
||||
const half2 * V_h2 = V_is_K_view ? K_h2 : (const half2 *) (V + nb23*sequence + nb22*z_KV);
|
||||
const float * sinks_f = sinks ? (const float *) sinks + zt_Q : nullptr;
|
||||
const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*ne31 + jt*ncols1)*ne11 : nullptr;
|
||||
const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*iter_j + jt)*ne11 : nullptr;
|
||||
|
||||
const float slope = ncols2 == 1 ? get_alibi_slope(max_bias, zt_Q, n_head_log2, m0, m1) : 1.0f;
|
||||
|
||||
if (KV_max) {
|
||||
if (use_sparse) {
|
||||
kb0_stop = min(kb0_stop, (KV_max[(sequence % ne33)*iter_j + jt] + nbatch_fa - 1) / nbatch_fa);
|
||||
} else if (KV_max) {
|
||||
kb0_stop = min(kb0_stop, KV_max[sequence*iter_j + jt] / nbatch_fa);
|
||||
}
|
||||
constexpr bool is_fixup = false; // All but (potentially) the last iterations write their data to dst rather than the fixup buffer.
|
||||
@@ -1936,11 +1945,13 @@ static __global__ void flash_attn_ext_f16(
|
||||
|
||||
const half2 * V_h2 = V_is_K_view ? K_h2 : (const half2 *) (V + nb23*sequence + nb22*z_KV);
|
||||
const float * sinks_f = sinks ? (const float *) sinks + zt_Q : nullptr;
|
||||
const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*ne31 + jt*ncols1)*ne11 : nullptr;
|
||||
const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*iter_j + jt)*ne11 : nullptr;
|
||||
|
||||
const float slope = ncols2 == 1 ? get_alibi_slope(max_bias, zt_Q, n_head_log2, m0, m1) : 1.0f;
|
||||
|
||||
if (KV_max) {
|
||||
if (use_sparse) {
|
||||
kb0_stop = min(kb0_stop, (KV_max[(sequence % ne33)*iter_j + jt] + nbatch_fa - 1) / nbatch_fa);
|
||||
} else if (KV_max) {
|
||||
kb0_stop = min(kb0_stop, KV_max[sequence*iter_j + jt] / nbatch_fa);
|
||||
}
|
||||
|
||||
@@ -1963,7 +1974,7 @@ static __global__ void flash_attn_ext_f16(
|
||||
#endif // defined(FLASH_ATTN_AVAILABLE) && (defined(VOLTA_MMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE))
|
||||
}
|
||||
|
||||
bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
|
||||
bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(const int cc, const ggml_tensor * dst, const int ncols1);
|
||||
|
||||
template <int DKQ, int DV, int ncols1, int ncols2>
|
||||
void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
@@ -2016,7 +2027,7 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml
|
||||
constexpr bool use_logit_softcap = false;
|
||||
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
if constexpr (ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(DKQ, DV, ncols1, ncols2)) {
|
||||
if (ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ctx, dst)) {
|
||||
if (ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(cc, dst, ncols1)) {
|
||||
constexpr bool use_sparse_kernel = true;
|
||||
fattn_kernel = flash_attn_ext_f16<DKQ, DV, ncols1, ncols2, use_logit_softcap, V_is_K_view, use_sparse_kernel>;
|
||||
use_sparse = true;
|
||||
|
||||
+33
-17
@@ -6,10 +6,11 @@
|
||||
#include "fattn.cuh"
|
||||
|
||||
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
// one list per group of ncols1 queries: a column is selected if any query of the group can see it
|
||||
__launch_bounds__(256, 1)
|
||||
static __global__ void flash_attn_mask_to_sparse_indices(
|
||||
const half * mask_ptr, int32_t * indices_ptr, const int ne30, const int n_kv_max,
|
||||
const int64_t s31, const int64_t s33) {
|
||||
const half * mask_ptr, int32_t * indices_ptr, int32_t * counts_ptr, const int ne30, const int n_queries,
|
||||
const int ncols1, const int n_kv_max, const int64_t s31, const int64_t s33) {
|
||||
ggml_cuda_pdl_sync();
|
||||
|
||||
constexpr int values_per_lane = 8;
|
||||
@@ -17,10 +18,13 @@ static __global__ void flash_attn_mask_to_sparse_indices(
|
||||
const int warp = tid / WARP_SIZE;
|
||||
const int lane = tid % WARP_SIZE;
|
||||
const int sequence = blockIdx.y;
|
||||
const int query = blockIdx.x;
|
||||
const int group = blockIdx.x;
|
||||
|
||||
const half * mask = mask_ptr + sequence*s33 + query*s31;
|
||||
int32_t * indices = indices_ptr + (int64_t(sequence)*gridDim.x + query)*n_kv_max;
|
||||
const int q0 = group*ncols1;
|
||||
const int q1 = min(q0 + ncols1, n_queries);
|
||||
|
||||
const half * mask = mask_ptr + sequence*s33 + q0*s31;
|
||||
int32_t * indices = indices_ptr + (int64_t(sequence)*gridDim.x + group)*n_kv_max;
|
||||
|
||||
__shared__ int warp_offsets[256/WARP_SIZE];
|
||||
__shared__ int row_count;
|
||||
@@ -37,7 +41,10 @@ static __global__ void flash_attn_mask_to_sparse_indices(
|
||||
#pragma unroll
|
||||
for (int item = 0; item < values_per_lane; ++item) {
|
||||
const int i = i0 + (warp*values_per_lane + item)*WARP_SIZE + lane;
|
||||
const bool selected = i < ne30 && isfinite(__half2float(mask[i]));
|
||||
bool selected = false;
|
||||
for (int q = 0; q < q1 - q0 && !selected; ++q) {
|
||||
selected = i < ne30 && isfinite(__half2float(mask[q*s31 + i]));
|
||||
}
|
||||
selected_warp[item] = __ballot_sync(0xFFFFFFFF, selected);
|
||||
warp_count += __popc(selected_warp[item]);
|
||||
}
|
||||
@@ -78,10 +85,13 @@ static __global__ void flash_attn_mask_to_sparse_indices(
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
const int count = row_count;
|
||||
const int count = min(row_count, n_kv_max);
|
||||
for (int i = count + tid; i < n_kv_max; i += blockDim.x) {
|
||||
indices[i] = -1;
|
||||
}
|
||||
if (tid == 0) {
|
||||
counts_ptr[int64_t(sequence)*gridDim.x + group] = count;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
// the dependent grid reads indices, signal once the row is complete
|
||||
@@ -90,31 +100,30 @@ static __global__ void flash_attn_mask_to_sparse_indices(
|
||||
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
|
||||
void ggml_cuda_flash_attn_ext_compact_mask(
|
||||
const ggml_tensor * mask, int32_t * indices, int32_t n_kv_max, cudaStream_t stream) {
|
||||
const ggml_tensor * mask, int32_t * indices, int32_t * counts, int32_t n_queries, int32_t ncols1, int32_t n_kv_max, cudaStream_t stream) {
|
||||
#if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
|
||||
GGML_UNUSED_VARS(mask, indices, n_kv_max, stream);
|
||||
GGML_UNUSED_VARS(mask, indices, counts, n_queries, ncols1, n_kv_max, stream);
|
||||
GGML_ABORT("sparse flash attention is only supported on NVIDIA CUDA");
|
||||
#else
|
||||
const int64_t s31 = mask->nb[1] / sizeof(half);
|
||||
const int64_t s33 = mask->nb[3] / sizeof(half);
|
||||
const dim3 blocks_num(mask->ne[1], mask->ne[3], 1);
|
||||
const dim3 blocks_num((n_queries + ncols1 - 1)/ncols1, mask->ne[3], 1);
|
||||
const dim3 block_dim(256, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params(blocks_num, block_dim, 0, stream);
|
||||
ggml_cuda_kernel_launch(flash_attn_mask_to_sparse_indices, launch_params,
|
||||
(const half *) mask->data, indices, int(mask->ne[0]), n_kv_max, s31, s33);
|
||||
(const half *) mask->data, indices, counts, int(mask->ne[0]), n_queries, ncols1, n_kv_max, s31, s33);
|
||||
CUDA_CHECK(cudaGetLastError());
|
||||
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
}
|
||||
|
||||
bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(const int cc, const ggml_tensor * dst, const int ncols1) {
|
||||
#if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
|
||||
GGML_UNUSED_VARS(ctx, dst);
|
||||
GGML_UNUSED_VARS(cc, dst, ncols1);
|
||||
return false;
|
||||
#else
|
||||
const ggml_tensor * Q = dst->src[0];
|
||||
const ggml_tensor * K = dst->src[1];
|
||||
const ggml_tensor * mask = dst->src[3];
|
||||
const int cc = ggml_cuda_info().devices[ctx.device].cc;
|
||||
|
||||
float max_bias = 0.0f;
|
||||
float logit_softcap = 0.0f;
|
||||
@@ -122,10 +131,13 @@ bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ggml_backend_cuda_context
|
||||
memcpy(&logit_softcap, (const float *) dst->op_params + 2, sizeof(float));
|
||||
|
||||
const int32_t n_kv_max = ggml_get_op_params_i32(dst, 4);
|
||||
|
||||
const int64_t n_gather = (ncols1 == 1 ? Q->ne[1] : ncols1) * (int64_t) n_kv_max;
|
||||
|
||||
return GGML_CUDA_CC_IS_NVIDIA(cc) && turing_mma_available(cc) &&
|
||||
mask != nullptr && n_kv_max > 0 && max_bias == 0.0f && logit_softcap == 0.0f &&
|
||||
mask->ne[0] == K->ne[1] && mask->ne[1] >= Q->ne[1] && mask->ne[2] == 1 &&
|
||||
K->ne[1] >= std::max<int64_t>(4096, 2LL*n_kv_max);
|
||||
K->ne[1] >= std::max<int64_t>(4096, 2*n_gather);
|
||||
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
}
|
||||
|
||||
@@ -136,7 +148,7 @@ static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ggml_backend_cuda_con
|
||||
|
||||
#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
if constexpr (ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(DKQ, DV, 1, ncols2)) {
|
||||
if (ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(ctx, dst)) {
|
||||
if (ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(cc, dst, 1)) {
|
||||
ggml_cuda_flash_attn_ext_mma_f16_case<DKQ, DV, 1, ncols2>(ctx, dst);
|
||||
return;
|
||||
}
|
||||
@@ -614,7 +626,11 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const
|
||||
if (turing_mma_available(cc) && Q->ne[0] != 40 && Q->ne[0] != 72) {
|
||||
if (can_use_vector_kernel) {
|
||||
if (!ggml_is_quantized(K->type) && !ggml_is_quantized(V->type)) {
|
||||
if (cc >= GGML_CUDA_CC_ADA_LOVELACE && Q->ne[1] == 1 && Q->ne[3] == 1 && !(gqa_ratio > 4 && K->ne[1] >= 8192)) {
|
||||
// the sparse gather exists only in the MMA kernel: (DKQ, DV, 1, 8) with GQA > 4
|
||||
const bool sparse_decode = gqa_opt_applies && gqa_ratio > 4 &&
|
||||
ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse(K->ne[0], V->ne[0], 1, 8) &&
|
||||
ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(cc, dst, 1);
|
||||
if (cc >= GGML_CUDA_CC_ADA_LOVELACE && Q->ne[1] == 1 && Q->ne[3] == 1 && !(gqa_ratio > 4 && K->ne[1] >= 8192) && !sparse_decode) {
|
||||
return BEST_FATTN_KERNEL_VEC;
|
||||
}
|
||||
} else {
|
||||
|
||||
@@ -1820,6 +1820,55 @@ static bool ggml_cuda_should_fuse_mul_mat_vec_q(const ggml_tensor * tensor) {
|
||||
return use_mul_mat_vec_q;
|
||||
}
|
||||
|
||||
static bool ggml_cuda_match_shared_expert(const ggml_cgraph * graph, int routed_idx, int shared_idx) {
|
||||
if (routed_idx + 2 >= graph->n_nodes || shared_idx + 2 >= graph->n_nodes || shared_idx < routed_idx + 3) {
|
||||
return false;
|
||||
}
|
||||
const int nodes[] = { routed_idx, routed_idx + 1, routed_idx + 2, shared_idx, shared_idx + 1, shared_idx + 2 };
|
||||
const ggml_op ops[] = { GGML_OP_MUL_MAT_ID, GGML_OP_MUL_MAT_ID, GGML_OP_GLU,
|
||||
GGML_OP_MUL_MAT, GGML_OP_MUL_MAT, GGML_OP_GLU };
|
||||
const int outputs[] = { routed_idx + 2, shared_idx + 2 };
|
||||
if (!ggml_can_fuse_subgraph_ext(graph, nodes, 6, ops, outputs, 2)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const ggml_tensor * routed = graph->nodes[routed_idx + 2];
|
||||
const ggml_tensor * shared = graph->nodes[shared_idx + 2];
|
||||
const ggml_tensor * gate = routed->src[0];
|
||||
const ggml_tensor * up = routed->src[1];
|
||||
const ggml_tensor * shared_gate = shared->src[0];
|
||||
const ggml_tensor * shared_up = shared->src[1];
|
||||
const auto is_pair = [&](const ggml_tensor * a, const ggml_tensor * b, int idx) {
|
||||
return (a == graph->nodes[idx] && b == graph->nodes[idx + 1]) ||
|
||||
(b == graph->nodes[idx] && a == graph->nodes[idx + 1]);
|
||||
};
|
||||
if (!is_pair(gate, up, routed_idx) || !is_pair(shared_gate, shared_up, shared_idx) ||
|
||||
!ggml_cuda_should_fuse_mul_mat(up, gate, routed) ||
|
||||
!ggml_cuda_should_fuse_mul_mat(shared_up, shared_gate, shared) ||
|
||||
!up->src[0]->buffer ||
|
||||
!ggml_cuda_should_fuse_mul_mat_vec_q(up)) {
|
||||
return false;
|
||||
}
|
||||
const ggml_tensor * input = up->src[1];
|
||||
const ggml_tensor * weight = up->src[0];
|
||||
const ggml_tensor * shared_weight = shared_up->src[0];
|
||||
if (input->op != GGML_OP_RESHAPE || input->src[0] != shared_up->src[1] ||
|
||||
input->ne[1] != 1 || input->ne[3] != 1 || !ggml_is_contiguous(input) ||
|
||||
!ggml_is_contiguous(shared_up->src[1]) || !ggml_is_matrix(shared_up->src[1]) ||
|
||||
weight->type != shared_weight->type || weight->ne[0] != shared_weight->ne[0] ||
|
||||
weight->ne[1] != shared_weight->ne[1] || weight->nb[1] != shared_weight->nb[1] || weight->ne[3] != 1 ||
|
||||
!ggml_is_matrix(shared_weight) || !ggml_is_contiguous(shared_weight) ||
|
||||
!ggml_is_contiguous(shared_gate->src[0]) || !ggml_is_contiguous(routed) || !ggml_is_contiguous(shared)) {
|
||||
return false;
|
||||
}
|
||||
if (shared_weight->op != GGML_OP_NONE || shared_gate->src[0]->op != GGML_OP_NONE ||
|
||||
ggml_get_glu_op(routed) != ggml_get_glu_op(shared) ||
|
||||
ggml_get_op_params_f32(routed, 3) != ggml_get_op_params_f32(shared, 3)) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
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
|
||||
|
||||
@@ -3438,6 +3487,25 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
|
||||
|
||||
ggml_tensor * node = cgraph->nodes[i];
|
||||
|
||||
if (node->op == GGML_OP_MUL_MAT_ID && cuda_ctx->stream_context().concurrent_events.empty() &&
|
||||
ggml_cuda_match_shared_expert(cgraph, i, i + 3)) {
|
||||
const int outputs[] = { i + 2, i + 5 };
|
||||
if (ggml_cuda_check_fusion_memory_ranges(cgraph, i, 6, outputs, 2)) {
|
||||
ggml_tensor * routed = cgraph->nodes[i + 2];
|
||||
ggml_tensor * shared = cgraph->nodes[i + 5];
|
||||
const ggml_tensor * up = routed->src[1];
|
||||
ggml_cuda_mm_fusion_args_host fusion{};
|
||||
fusion.gate = routed->src[0]->src[0];
|
||||
fusion.glu_op = ggml_get_glu_op(routed);
|
||||
fusion.glu_limit = ggml_get_op_params_f32(routed, 3);
|
||||
fusion.shared_up = shared->src[1]->src[0];
|
||||
fusion.shared_gate = shared->src[0]->src[0];
|
||||
fusion.shared_dst = shared;
|
||||
ggml_cuda_mul_mat_vec_q(*cuda_ctx, up->src[0], up->src[1], up->src[2], routed, &fusion);
|
||||
return 5;
|
||||
}
|
||||
}
|
||||
|
||||
if (node->op == GGML_OP_MUL) {
|
||||
ggml_cuda_moe_weighted_reduction_match match;
|
||||
if (ggml_cuda_match_moe_weighted_reduction(cgraph, i, match)) {
|
||||
@@ -4506,6 +4574,27 @@ static void ggml_backend_cuda_graph_optimize(ggml_backend_t backend, ggml_cgraph
|
||||
|
||||
static const bool disable_fusion = getenv("GGML_CUDA_DISABLE_FUSION") != nullptr && std::atoi(getenv("GGML_CUDA_DISABLE_FUSION"));
|
||||
if (!disable_fusion) {
|
||||
ggml_cuda_set_device(cuda_ctx->device);
|
||||
for (int i = 0; i + 5 < cgraph->n_nodes; ++i) {
|
||||
if (cgraph->nodes[i]->op != GGML_OP_MUL_MAT_ID) {
|
||||
continue;
|
||||
}
|
||||
for (int j = i + 3; j + 2 < cgraph->n_nodes; ++j) {
|
||||
if (cgraph->nodes[j]->op == GGML_OP_MUL_MAT_ID && cgraph->nodes[j + 1]->op == GGML_OP_MUL_MAT_ID) {
|
||||
break;
|
||||
}
|
||||
if (cgraph->nodes[j]->op != GGML_OP_MUL_MAT || !ggml_cuda_match_shared_expert(cgraph, i, j)) {
|
||||
continue;
|
||||
}
|
||||
// Group both outputs before allocation so the shared result cannot alias intervening nodes.
|
||||
std::rotate(cgraph->nodes + i + 3, cgraph->nodes + j, cgraph->nodes + j + 3);
|
||||
ggml_tensor * up = cgraph->nodes[i + 2]->src[1];
|
||||
params->add_alloc_dep(params->user_data, up->src[1], cgraph->nodes[i + 5]);
|
||||
params->add_alloc_dep(params->user_data, up->src[2], cgraph->nodes[i + 5]);
|
||||
i += 5;
|
||||
break;
|
||||
}
|
||||
}
|
||||
for (int i = 0; i < cgraph->n_nodes; ++i) {
|
||||
if (cgraph->nodes[i]->op != GGML_OP_MUL) {
|
||||
continue;
|
||||
|
||||
@@ -609,14 +609,19 @@ static __global__ void mul_mat_vec_q(
|
||||
const int blocks_per_row_x = ncols_x / qk;
|
||||
constexpr int blocks_per_iter = vdr * nwarps*warp_size / qi;
|
||||
|
||||
const uint32_t channel_dst = blockIdx.y;
|
||||
const bool shared_expert = has_fusion && fusion.shared_up && blockIdx.y == gridDim.y - 1;
|
||||
const uint32_t channel_dst = shared_expert ? 0 : blockIdx.y;
|
||||
if (shared_expert) {
|
||||
vx = fusion.shared_up;
|
||||
dst = fusion.shared_dst;
|
||||
}
|
||||
|
||||
uint32_t channel_x;
|
||||
uint32_t channel_y;
|
||||
uint32_t sample_dst;
|
||||
|
||||
ggml_cuda_pdl_sync();
|
||||
channel_x = ncols_dst == 1 && ids ? ids[channel_dst] : fastdiv(channel_dst, channel_ratio);
|
||||
channel_x = shared_expert ? 0 : ncols_dst == 1 && ids ? ids[channel_dst] : fastdiv(channel_dst, channel_ratio);
|
||||
channel_y = ncols_dst == 1 && ids ? fastmodulo(channel_dst, nchannels_y) : channel_dst;
|
||||
sample_dst = blockIdx.z;
|
||||
|
||||
@@ -640,7 +645,7 @@ static __global__ void mul_mat_vec_q(
|
||||
use_gate = fusion.gate != nullptr;
|
||||
use_bias = fusion.x_bias != nullptr;
|
||||
use_gate_bias = fusion.gate_bias != nullptr && use_gate;
|
||||
vgate = fusion.gate;
|
||||
vgate = shared_expert ? fusion.shared_gate : fusion.gate;
|
||||
x_bias = (const float *) fusion.x_bias;
|
||||
gate_bias = (const float *) fusion.gate_bias;
|
||||
active_glu = fusion.glu_op;
|
||||
@@ -838,7 +843,7 @@ static __global__ void mul_mat_vec_q_moe(
|
||||
const void * vx_ptr, const void * vy_ptr, const int32_t * ids_ptr, const ggml_cuda_mm_fusion_args_device fusion,
|
||||
float * dst_ptr,
|
||||
const uint32_t ncols_x, const uint3 nchannels_y, const uint32_t nrows_x,
|
||||
const uint32_t stride_row_x, const uint32_t stride_col_y, const uint32_t stride_col_dst,
|
||||
const uint32_t stride_row_x, const uint32_t stride_col_y, uint32_t stride_col_dst,
|
||||
const uint32_t stride_channel_x, const uint32_t stride_channel_y, const uint32_t stride_channel_dst,
|
||||
const uint32_t ncols_dst, const uint32_t ids_stride) {
|
||||
const void * GGML_CUDA_RESTRICT vx = vx_ptr;
|
||||
@@ -853,6 +858,13 @@ static __global__ void mul_mat_vec_q_moe(
|
||||
|
||||
constexpr vec_dot_q_cuda_t vec_dot_q_cuda = get_vec_dot_q_cuda(type);
|
||||
|
||||
const bool shared_expert = has_fusion && fusion.shared_up && blockIdx.y == gridDim.y - 1;
|
||||
if (shared_expert) {
|
||||
vx = fusion.shared_up;
|
||||
dst = fusion.shared_dst;
|
||||
stride_col_dst = nrows_x;
|
||||
}
|
||||
|
||||
// fuse gate, bias, scales, and glu_op into the up projection
|
||||
bool use_gate = false;
|
||||
const void * vgate = nullptr;
|
||||
@@ -865,7 +877,7 @@ static __global__ void mul_mat_vec_q_moe(
|
||||
|
||||
if constexpr (has_fusion) {
|
||||
use_gate = fusion.gate != nullptr;
|
||||
vgate = fusion.gate;
|
||||
vgate = shared_expert ? fusion.shared_gate : fusion.gate;
|
||||
x_bias = (const float *) fusion.x_bias;
|
||||
gate_bias = (const float *) fusion.gate_bias;
|
||||
active_glu = fusion.glu_op;
|
||||
@@ -881,14 +893,14 @@ static __global__ void mul_mat_vec_q_moe(
|
||||
const int blocks_per_row_x = ncols_x / qk;
|
||||
constexpr int blocks_per_iter = vdr * warp_size / qi;
|
||||
|
||||
const uint32_t channel_dst = blockIdx.y;
|
||||
const uint32_t channel_dst = shared_expert ? 0 : blockIdx.y;
|
||||
|
||||
if (token_idx >= ncols_dst) {
|
||||
return;
|
||||
}
|
||||
|
||||
ggml_cuda_pdl_sync();
|
||||
const uint32_t channel_x = ids[channel_dst + token_idx * ids_stride];
|
||||
const uint32_t channel_x = shared_expert ? 0 : ids[channel_dst + token_idx * ids_stride];
|
||||
const uint32_t channel_y = fastmodulo(channel_dst, nchannels_y);
|
||||
|
||||
const block_q8_1 * y = ((const block_q8_1 *) vy) + channel_y*stride_channel_y + token_idx*stride_col_y;
|
||||
@@ -1034,7 +1046,7 @@ static void mul_mat_vec_q_moe_launch(
|
||||
|
||||
constexpr int rows_per_block = 2; // 2 gives best perf based on tuning
|
||||
const int64_t nblocks_rows = (nrows_x + rows_per_block - 1) / rows_per_block;
|
||||
const dim3 block_nums(nblocks_rows, nchannels_dst);
|
||||
const dim3 block_nums(nblocks_rows, nchannels_dst + (fusion.shared_up != nullptr));
|
||||
const dim3 block_dims(warp_size, ncols_dst);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(block_nums, block_dims, 0, stream);
|
||||
|
||||
@@ -1171,7 +1183,7 @@ static void mul_mat_vec_q_switch_ncols_dst(
|
||||
|
||||
constexpr bool c_halve_iters = decltype(halve_iters_tag)::value && c_promoted;
|
||||
|
||||
const std::pair<dim3, dim3> dims = calc_launch_params<type>(c_ncols_dst, nrows_x, nchannels_dst,
|
||||
const std::pair<dim3, dim3> dims = calc_launch_params<type>(c_ncols_dst, nrows_x, nchannels_dst + (fusion.shared_up != nullptr),
|
||||
nsamples_dst, warp_size, table_id, c_small_k, c_halve_iters);
|
||||
mul_mat_vec_q_switch_fusion<type, c_ncols_dst, c_small_k, c_halve_iters>(
|
||||
vx, vy, ids, fusion, dst, ncols_x, nchannels_y_fd, stride_row_x, stride_col_y, stride_col_dst,
|
||||
@@ -1438,6 +1450,22 @@ void ggml_cuda_mul_mat_vec_q(
|
||||
// non-negligible for some models such as gpt-oss-20b
|
||||
GGML_ASSERT((fusion->x_scale == nullptr && fusion->gate_scale == nullptr) || src0->type == GGML_TYPE_NVFP4);
|
||||
|
||||
if (fusion->shared_up) {
|
||||
GGML_ASSERT(ids && fusion->gate && fusion->shared_gate && fusion->shared_dst);
|
||||
GGML_ASSERT(!fusion->x_bias && !fusion->gate_bias && !fusion->x_scale && !fusion->gate_scale);
|
||||
GGML_ASSERT(ne11 == 1 && ne03 == 1 && ne13 == 1);
|
||||
GGML_ASSERT(fusion->shared_up->type == src0->type && fusion->shared_gate->type == src0->type);
|
||||
GGML_ASSERT(ggml_are_same_shape(fusion->shared_up, fusion->shared_gate));
|
||||
GGML_ASSERT(ggml_is_contiguous(fusion->shared_up) && ggml_is_contiguous(fusion->shared_gate));
|
||||
GGML_ASSERT(fusion->shared_up->ne[0] == ne00 && fusion->shared_up->ne[1] == ne01);
|
||||
GGML_ASSERT(fusion->shared_up->nb[1] == nb01 && ggml_is_matrix(fusion->shared_up));
|
||||
GGML_ASSERT(fusion->shared_dst->type == GGML_TYPE_F32 && ggml_is_contiguous(fusion->shared_dst));
|
||||
GGML_ASSERT(fusion->shared_dst->ne[0] == ne0 && fusion->shared_dst->ne[1] == ne2);
|
||||
fusion_local.shared_up = fusion->shared_up->data;
|
||||
fusion_local.shared_gate = fusion->shared_gate->data;
|
||||
fusion_local.shared_dst = (float *) fusion->shared_dst->data;
|
||||
}
|
||||
|
||||
if (fusion->x_bias) {
|
||||
GGML_ASSERT(fusion->x_bias->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(fusion->x_bias->ne[0] == dst->ne[0]);
|
||||
|
||||
@@ -4000,7 +4000,9 @@ static bool ggml_hexagon_flash_attn_is_hmx_eligible(
|
||||
const uint32_t DK = q->ne[0];
|
||||
const uint32_t DV = v->ne[0];
|
||||
|
||||
if (DK % 64 != 0 || DV % 64 != 0) {
|
||||
// Head dims that are not multiples of 64 are handled by internally padding to
|
||||
// DK_pad/DV_pad = round_up(.,64) and zero-filling the tail lanes.
|
||||
if (DK % 8 != 0 || DV % 8 != 0) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -4073,8 +4075,13 @@ static bool ggml_hexagon_precompute_flash_attn_params(
|
||||
// Check HMX eligibility
|
||||
const struct ggml_tensor * sinks = op->src[4];
|
||||
if (ggml_hexagon_flash_attn_is_hmx_eligible(sess, q, k, v, sinks)) {
|
||||
// HMX tiles head_dim in units of 64; when DK/DV are not 64-aligned the kernel
|
||||
// operates on padded dims with zero-filled tail lanes. VTCM budget and chunk-size
|
||||
// are sized for the padded tiles.
|
||||
const uint32_t DK_pad = hex_round_up(DK, 64);
|
||||
const uint32_t DV_pad = hex_round_up(DV, 64);
|
||||
size_t Br = 0, Bc = 0;
|
||||
int ret = hmx_fa_find_chunk_size(&Br, &Bc, G, DK, DV, neq1, nek1, sess->vtcm_size, sess->n_threads, kparams->is_q_fp32 != 0);
|
||||
int ret = hmx_fa_find_chunk_size(&Br, &Bc, G, DK_pad, DV_pad, neq1, nek1, sess->vtcm_size, sess->n_threads, kparams->is_q_fp32 != 0);
|
||||
if (ret == 0) {
|
||||
kparams->kernel_type = HTP_FA_KERNEL_HMX;
|
||||
kparams->Br = Br;
|
||||
@@ -4084,7 +4091,7 @@ static bool ggml_hexagon_precompute_flash_attn_params(
|
||||
|
||||
kparams->u.hmx.g_br = hex_align_up(G * Br, 32);
|
||||
kparams->u.hmx.pipeline = (kparams->n_kv_blocks >= 3 && sess->n_threads >= 2) ? 1 : 0;
|
||||
kparams->vtcm_size = hmx_fa_compute_vtcm_usage(G, DK, DV, Br, Bc, kparams->n_threads, kparams->u.hmx.pipeline != 0, kparams->is_q_fp32 != 0);
|
||||
kparams->vtcm_size = hmx_fa_compute_vtcm_usage(G, DK_pad, DV_pad, Br, Bc, kparams->n_threads, kparams->u.hmx.pipeline != 0, kparams->is_q_fp32 != 0);
|
||||
|
||||
const size_t row_vec_bytes = hex_align_up(Bc * sizeof(uint16_t), 256);
|
||||
kparams->u.hmx.row_buf_stride = row_vec_bytes / 128; // HVX vector is 128 bytes
|
||||
@@ -5332,11 +5339,12 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session *
|
||||
}
|
||||
}
|
||||
|
||||
if (src0->type != GGML_TYPE_F32 && src0->ne[0] < 32) {
|
||||
if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_I32 && src0->ne[0] < 32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_F16 && src0->type != GGML_TYPE_Q8_0) {
|
||||
if (src0->type != GGML_TYPE_F32 && src0->type != GGML_TYPE_F16 &&
|
||||
src0->type != GGML_TYPE_Q8_0 && src0->type != GGML_TYPE_I32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -5344,7 +5352,12 @@ static bool ggml_hexagon_supported_get_rows(const struct ggml_hexagon_session *
|
||||
return false;
|
||||
}
|
||||
|
||||
if (dst->type != GGML_TYPE_F32) {
|
||||
if (src0->type == GGML_TYPE_I32) {
|
||||
if (dst->type != GGML_TYPE_I32) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
else if (dst->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -5375,6 +5388,30 @@ static bool ggml_hexagon_supported_argsort(const struct ggml_hexagon_session * s
|
||||
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) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (dst->type != GGML_TYPE_I32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// 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) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
static bool ggml_hexagon_supported_rope(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];
|
||||
@@ -5523,11 +5560,6 @@ static bool ggml_hexagon_supported_im2col(const struct ggml_hexagon_session * se
|
||||
const struct ggml_tensor * src1 = op->src[1];
|
||||
const struct ggml_tensor * dst = op;
|
||||
|
||||
const bool is_2D = ((const int32_t *) op->op_params)[6] == 1;
|
||||
if (!is_2D) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// For now support F32->F32 and F32->F16 only.
|
||||
if (src1->type != GGML_TYPE_F32 || (dst->type != GGML_TYPE_F16 && dst->type != GGML_TYPE_F32)) {
|
||||
return false;
|
||||
@@ -5537,13 +5569,6 @@ static bool ggml_hexagon_supported_im2col(const struct ggml_hexagon_session * se
|
||||
return false;
|
||||
}
|
||||
|
||||
// For now keep padded OPs on CPU. Will revisit once we expand coverage past patch-embed OPs.
|
||||
const int32_t p0 = ((const int32_t *) op->op_params)[2];
|
||||
const int32_t p1 = ((const int32_t *) op->op_params)[3];
|
||||
if (p0 != 0 || p1 != 0) {
|
||||
return false;
|
||||
}
|
||||
|
||||
GGML_UNUSED(sess);
|
||||
return true;
|
||||
}
|
||||
@@ -5682,6 +5707,7 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) {
|
||||
case GGML_OP_SET_ROWS: return HTP_OP_SET_ROWS;
|
||||
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_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;
|
||||
@@ -5704,6 +5730,7 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) {
|
||||
case GGML_OP_TRI: return HTP_OP_TRI;
|
||||
case GGML_OP_PAD: return HTP_OP_PAD;
|
||||
case GGML_OP_IM2COL: return HTP_OP_IM2COL;
|
||||
case GGML_OP_ROLL: return HTP_OP_ROLL;
|
||||
|
||||
case GGML_OP_UNARY:
|
||||
switch (ggml_get_unary_op(t)) {
|
||||
@@ -5728,6 +5755,7 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) {
|
||||
case GGML_GLU_OP_SWIGLU_OAI: return HTP_OP_GLU_SWIGLU_OAI;
|
||||
case GGML_GLU_OP_SWIGLU_CLAMP: return HTP_OP_GLU_SWIGLU_CLAMP;
|
||||
case GGML_GLU_OP_GEGLU: return HTP_OP_GLU_GEGLU;
|
||||
case GGML_GLU_OP_GEGLU_QUICK: return HTP_OP_GLU_GEGLU_QUICK;
|
||||
default: break;
|
||||
}
|
||||
break;
|
||||
@@ -6636,6 +6664,31 @@ static bool ggml_hexagon_supported_fill(const struct ggml_hexagon_session * sess
|
||||
GGML_UNUSED(sess);
|
||||
}
|
||||
|
||||
static bool ggml_hexagon_supported_roll(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) {
|
||||
GGML_UNUSED(sess);
|
||||
|
||||
const struct ggml_tensor * src0 = op->src[0];
|
||||
const struct ggml_tensor * dst = op;
|
||||
|
||||
if (src0->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!ggml_are_same_shape(src0, dst)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (src0->nb[0] != ggml_type_size(src0->type) || dst->nb[0] != ggml_type_size(dst->type)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!ggml_is_contiguous(dst)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, const struct ggml_tensor * op) {
|
||||
auto dev_ctx = static_cast<ggml_backend_hexagon_device_context *>(dev->context);
|
||||
auto sess = dev_ctx->session();
|
||||
@@ -6723,6 +6776,7 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons
|
||||
case GGML_GLU_OP_SWIGLU_OAI:
|
||||
case GGML_GLU_OP_SWIGLU_CLAMP:
|
||||
case GGML_GLU_OP_GEGLU:
|
||||
case GGML_GLU_OP_GEGLU_QUICK:
|
||||
supp = ggml_hexagon_supported_activations(sess, op);
|
||||
break;
|
||||
default:
|
||||
@@ -6763,6 +6817,10 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons
|
||||
supp = ggml_hexagon_supported_argsort(sess, op);
|
||||
break;
|
||||
|
||||
case GGML_OP_TOP_K:
|
||||
supp = ggml_hexagon_supported_top_k(sess, op);
|
||||
break;
|
||||
|
||||
case GGML_OP_SSM_CONV:
|
||||
supp = ggml_hexagon_supported_ssm_conv(sess, op);
|
||||
break;
|
||||
@@ -6803,6 +6861,10 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons
|
||||
supp = ggml_hexagon_supported_pad(sess, op);
|
||||
break;
|
||||
|
||||
case GGML_OP_ROLL:
|
||||
supp = ggml_hexagon_supported_roll(sess, op);
|
||||
break;
|
||||
|
||||
default:
|
||||
break;
|
||||
}
|
||||
|
||||
@@ -43,6 +43,7 @@ add_library(${HTP_LIB} SHARED
|
||||
pad-ops.c
|
||||
argsort-ops.c
|
||||
im2col-ops.c
|
||||
roll-ops.c
|
||||
allreduce-ops.c
|
||||
)
|
||||
|
||||
|
||||
@@ -312,6 +312,45 @@ static inline void hvx_geglu_f32_aa(uint8_t * restrict dst, const uint8_t * rest
|
||||
}
|
||||
}
|
||||
|
||||
static inline void hvx_geglu_quick_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src0, const uint8_t * restrict src1, uint32_t n) {
|
||||
assert((unsigned long) dst % 128 == 0);
|
||||
assert((unsigned long) src0 % 128 == 0);
|
||||
assert((unsigned long) src1 % 128 == 0);
|
||||
|
||||
HVX_Vector * restrict vdst = (HVX_Vector *) dst;
|
||||
const HVX_Vector * restrict vsrc0 = (const HVX_Vector *) src0;
|
||||
const HVX_Vector * restrict vsrc1 = (const HVX_Vector *) src1;
|
||||
|
||||
const uint32_t epv = 128 / sizeof(float);
|
||||
const uint32_t nvec = n / epv;
|
||||
const uint32_t nloe = n % epv;
|
||||
|
||||
const HVX_Vector v_scale = hvx_vec_splat_f32(1.702f);
|
||||
const HVX_Vector v_one = hvx_vec_splat_f32(1.0f);
|
||||
const HVX_Vector v_max_exp = hvx_vec_splat_f32(87.0f);
|
||||
const HVX_Vector v_min_exp = hvx_vec_splat_f32(-87.0f);
|
||||
|
||||
uint32_t i = 0;
|
||||
|
||||
_Pragma("unroll(4)")
|
||||
for (; i < nvec; i++) {
|
||||
HVX_Vector x = vsrc0[i];
|
||||
HVX_Vector g = vsrc1[i];
|
||||
HVX_Vector scaled_x = hvx_vec_mul_f32_f32(x, v_scale);
|
||||
HVX_Vector sigmoid_x = hvx_vec_fast_sigmoid_f32_guard_2it(scaled_x, v_one, v_max_exp, v_min_exp);
|
||||
vdst[i] = hvx_vec_mul_f32_f32(hvx_vec_mul_f32_f32(x, sigmoid_x), g);
|
||||
}
|
||||
|
||||
if (nloe) {
|
||||
HVX_Vector x = vsrc0[i];
|
||||
HVX_Vector g = vsrc1[i];
|
||||
HVX_Vector scaled_x = hvx_vec_mul_f32_f32(x, v_scale);
|
||||
HVX_Vector sigmoid_x = hvx_vec_fast_sigmoid_f32_guard_2it(scaled_x, v_one, v_max_exp, v_min_exp);
|
||||
HVX_Vector result = hvx_vec_mul_f32_f32(hvx_vec_mul_f32_f32(x, sigmoid_x), g);
|
||||
hvx_vec_store_a((void *) &vdst[i], nloe * sizeof(float), result);
|
||||
}
|
||||
}
|
||||
|
||||
// geglu(x, g) = gelu(x) * g
|
||||
static void geglu_f32(const float * restrict src0,
|
||||
const float * restrict src1,
|
||||
@@ -329,6 +368,23 @@ static void geglu_f32(const float * restrict src0,
|
||||
}
|
||||
}
|
||||
|
||||
// geglu_quick(x, g) = x * sigmoid(1.702 * x) * g
|
||||
static void geglu_quick_f32(const float * restrict src0,
|
||||
const float * restrict src1,
|
||||
float * restrict dst,
|
||||
const uint32_t num_rows,
|
||||
const struct htp_act_context * actx) {
|
||||
htp_glu_op_preamble;
|
||||
|
||||
for (uint32_t ib = 0; ib < num_rows; ib++) {
|
||||
const uint8_t * restrict src0_ptr = (const uint8_t *) src0 + (ib * src0_row_size_aligned);
|
||||
const uint8_t * restrict src1_ptr = (const uint8_t *) src1 + (ib * src1_row_size_aligned);
|
||||
uint8_t * restrict dst_ptr = (uint8_t *) dst + (ib * dst_row_size_aligned);
|
||||
|
||||
hvx_geglu_quick_f32_aa(dst_ptr, src0_ptr, src1_ptr, nc);
|
||||
}
|
||||
}
|
||||
|
||||
#define DEFINE_GLU_PER_THREAD(NAME, OP_STR, CORE_EXPR) \
|
||||
static void glu_##NAME##_f32_per_thread(unsigned int nth, unsigned int ith, void * data) { \
|
||||
struct htp_act_context * actx = (struct htp_act_context *) data; \
|
||||
@@ -433,6 +489,7 @@ DEFINE_GLU_PER_THREAD(swiglu, "swiglu-f32", swiglu_f32(src0_spad, src1_spad, dst
|
||||
DEFINE_GLU_PER_THREAD(swiglu_oai, "swiglu-oai-f32", swiglu_oai_f32(src0_spad, src1_spad, dst_spad, block_size, actx))
|
||||
DEFINE_GLU_PER_THREAD(swiglu_clamp, "swiglu-clamp-f32", swiglu_clamp_f32(src0_spad, src1_spad, dst_spad, block_size, actx))
|
||||
DEFINE_GLU_PER_THREAD(geglu, "geglu-f32", geglu_f32(src0_spad, src1_spad, dst_spad, block_size, actx))
|
||||
DEFINE_GLU_PER_THREAD(geglu_quick, "geglu-quick-f32", geglu_quick_f32(src0_spad, src1_spad, dst_spad, block_size, actx))
|
||||
|
||||
static int execute_op_activations_f32(struct htp_ops_context * octx) {
|
||||
const struct htp_tensor * src0 = octx->src[0];
|
||||
@@ -467,6 +524,11 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
|
||||
act_op_func = (worker_callback_t)glu_geglu_f32_per_thread;
|
||||
op_type = "geglu-f32";
|
||||
break;
|
||||
|
||||
case HTP_OP_GLU_GEGLU_QUICK:
|
||||
act_op_func = (worker_callback_t)glu_geglu_quick_f32_per_thread;
|
||||
op_type = "geglu-quick-f32";
|
||||
break;
|
||||
default:
|
||||
FARF(ERROR, "Unsupported activations Op %u\n", octx->op);
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
@@ -570,7 +632,11 @@ static int execute_op_activations_f32(struct htp_ops_context * octx) {
|
||||
const uint8_t * data_src0 = (const uint8_t *) src0->data;
|
||||
const uint8_t * data_src1 = src1 ? (const uint8_t *) src1->data : NULL;
|
||||
|
||||
if (!src1 && (octx->op == HTP_OP_GLU_SWIGLU || octx->op == HTP_OP_GLU_SWIGLU_OAI || octx->op == HTP_OP_GLU_SWIGLU_CLAMP || octx->op == HTP_OP_GLU_GEGLU)) {
|
||||
if (!src1 && (octx->op == HTP_OP_GLU_SWIGLU ||
|
||||
octx->op == HTP_OP_GLU_SWIGLU_OAI ||
|
||||
octx->op == HTP_OP_GLU_SWIGLU_CLAMP ||
|
||||
octx->op == HTP_OP_GLU_GEGLU ||
|
||||
octx->op == HTP_OP_GLU_GEGLU_QUICK)) {
|
||||
const int32_t swapped = octx->op_params[1];
|
||||
data_src1 = data_src0;
|
||||
actx.src1_row_size = actx.src0_row_size;
|
||||
|
||||
@@ -170,6 +170,21 @@ static void quicksort_values_indices_desc(float * values, int32_t * indices, int
|
||||
if (i < right) quicksort_values_indices_desc(values, indices, i, right);
|
||||
}
|
||||
|
||||
static uint32_t top_k_max_value_index(const float * values, uint32_t n, float * value) {
|
||||
uint32_t index = 0;
|
||||
float max_value = values[0];
|
||||
|
||||
for (uint32_t i = 1; i < n; i++) {
|
||||
if (values[i] > max_value) {
|
||||
max_value = values[i];
|
||||
index = i;
|
||||
}
|
||||
}
|
||||
|
||||
*value = max_value;
|
||||
return index;
|
||||
}
|
||||
|
||||
// LUT for ramp initialization of argsort output (first 32 members)
|
||||
int32_t argosrt_ramp_lut[32] __attribute__((aligned(VLEN))) = {
|
||||
0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15,
|
||||
@@ -303,6 +318,129 @@ static inline void bitonic_sort_generic_hvx(uint8_t * values, uint8_t * indices,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Sorts descending; values pre-padded with -INFINITY. init_indices=true resets
|
||||
// indices to a fresh ramp (normal full-row sort); false leaves caller-supplied
|
||||
// indices in place and only permutes them (used when merging candidates, to
|
||||
// preserve their original global index).
|
||||
|
||||
static void bitonic_sort_vtcm_desc(uint8_t * values, uint8_t * indices, uint32_t n_vec, bool init_indices) {
|
||||
HVX_Vector zero_vec = Q6_V_vzero();
|
||||
HVX_Vector idx_vec = *(HVX_Vector *)argosrt_ramp_lut;
|
||||
|
||||
HVX_VectorPred pred_all_1s = Q6_Q_vcmp_eq_VwVw(zero_vec, zero_vec);
|
||||
HVX_VectorPred pred_all_0s = Q6_Q_not_Q(pred_all_1s);
|
||||
|
||||
if (init_indices) {
|
||||
// Initialize indices ramp (values are already populated by the caller)
|
||||
for (uint32_t v = 0; v < n_vec; v++) {
|
||||
HVX_Vector idx = Q6_Vw_vadd_VwVw(idx_vec, Q6_V_vsplat_R(v * 32));
|
||||
*(HVX_Vector *)(indices + v * 128) = idx;
|
||||
}
|
||||
}
|
||||
|
||||
int M = 5;
|
||||
while ((1u << (M - 5)) < n_vec) M++;
|
||||
|
||||
for (int s = 1; s <= M; s++) {
|
||||
for (int stage_d = s - 1; stage_d >= 0; stage_d--) {
|
||||
int d = 1 << stage_d;
|
||||
if (d >= 32) {
|
||||
uint32_t v_dist = d / 32;
|
||||
for (uint32_t v1 = 0; v1 < n_vec; v1++) {
|
||||
if ((v1 & v_dist) == 0) {
|
||||
uint32_t v2 = v1 + v_dist;
|
||||
bool asc = (s < M) ? ((((v1 * 32) >> s) % 2) == 0) : false;
|
||||
|
||||
HVX_Vector Vv1 = *(HVX_Vector *)(values + v1 * 128);
|
||||
HVX_Vector Iv1 = *(HVX_Vector *)(indices + v1 * 128);
|
||||
HVX_Vector Vv2 = *(HVX_Vector *)(values + v2 * 128);
|
||||
HVX_Vector Iv2 = *(HVX_Vector *)(indices + v2 * 128);
|
||||
|
||||
vec_cas(&Vv1, &Iv1, &Vv2, &Iv2, asc);
|
||||
|
||||
*(HVX_Vector *)(values + v1 * 128) = Vv1;
|
||||
*(HVX_Vector *)(indices + v1 * 128) = Iv1;
|
||||
*(HVX_Vector *)(values + v2 * 128) = Vv2;
|
||||
*(HVX_Vector *)(indices + v2 * 128) = Iv2;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
if (s < 5) {
|
||||
HVX_VectorPred dir_mask = Q6_Q_vcmp_eq_VwVw(Q6_V_vand_VV(idx_vec, Q6_V_vsplat_R(1 << s)), zero_vec);
|
||||
for (uint32_t v = 0; v < n_vec; v++) {
|
||||
HVX_Vector Vv = *(HVX_Vector *)(values + v * 128);
|
||||
HVX_Vector Iv = *(HVX_Vector *)(indices + v * 128);
|
||||
|
||||
bitonic_cas_32(&Vv, &Iv, d, dir_mask, idx_vec, zero_vec);
|
||||
|
||||
*(HVX_Vector *)(values + v * 128) = Vv;
|
||||
*(HVX_Vector *)(indices + v * 128) = Iv;
|
||||
}
|
||||
} else {
|
||||
for (uint32_t v = 0; v < n_vec; v++) {
|
||||
bool asc = (s < M) ? ((((v * 32) >> s) % 2) == 0) : false;
|
||||
HVX_VectorPred dir_mask = asc ? pred_all_1s : pred_all_0s;
|
||||
|
||||
HVX_Vector Vv = *(HVX_Vector *)(values + v * 128);
|
||||
HVX_Vector Iv = *(HVX_Vector *)(indices + v * 128);
|
||||
|
||||
bitonic_cas_32(&Vv, &Iv, d, dir_mask, idx_vec, zero_vec);
|
||||
|
||||
*(HVX_Vector *)(values + v * 128) = Vv;
|
||||
*(HVX_Vector *)(indices + v * 128) = Iv;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
static void top_k_select_tiled(const uint8_t * src, uint32_t n, uint32_t k,
|
||||
float * values_buf, int32_t * indices_buf,
|
||||
float * top_values, int32_t * top_indices) {
|
||||
const uint32_t tile_elems = 1024;
|
||||
uint32_t n_tiles = (n + tile_elems - 1) / tile_elems;
|
||||
uint32_t candidate_count = n_tiles * k;
|
||||
uint32_t merge_n_vec = hmx_ceil_div(candidate_count, 32);
|
||||
uint32_t merge_n_vec_pow2 = 1;
|
||||
while (merge_n_vec_pow2 < merge_n_vec) merge_n_vec_pow2 <<= 1;
|
||||
uint32_t merge_elems = merge_n_vec_pow2 * 32;
|
||||
float * candidate_values = values_buf + tile_elems;
|
||||
int32_t * candidate_indices = indices_buf + tile_elems;
|
||||
uint32_t candidate_pos = 0;
|
||||
|
||||
for (uint32_t offset = 0; offset < n; offset += tile_elems) {
|
||||
uint32_t tile_count = MIN(tile_elems, n - offset);
|
||||
hvx_copy_f32_au((uint8_t *) values_buf, src + offset * sizeof(float), tile_count);
|
||||
if (tile_count < tile_elems) {
|
||||
hvx_splat_f32_u((uint8_t *) (values_buf + tile_count), -INFINITY, tile_elems - tile_count);
|
||||
}
|
||||
|
||||
bitonic_sort_vtcm_desc((uint8_t *) values_buf, (uint8_t *) indices_buf, tile_elems / 32, true);
|
||||
uint32_t tile_k = MIN(k, tile_count);
|
||||
for (uint32_t j = 0; j < tile_k; j++) {
|
||||
candidate_values[candidate_pos] = values_buf[j];
|
||||
candidate_indices[candidate_pos] = indices_buf[j] + (int32_t) offset;
|
||||
candidate_pos++;
|
||||
}
|
||||
}
|
||||
|
||||
if (merge_elems > candidate_pos) {
|
||||
hvx_splat_f32_u((uint8_t *) (candidate_values + candidate_pos), -INFINITY, merge_elems - candidate_pos);
|
||||
for (uint32_t j = candidate_pos; j < merge_elems; j++) {
|
||||
candidate_indices[j] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
bitonic_sort_vtcm_desc((uint8_t *) candidate_values, (uint8_t *) candidate_indices, merge_n_vec_pow2, false);
|
||||
|
||||
for (uint32_t j = 0; j < k; j++) {
|
||||
top_values[j] = candidate_values[j];
|
||||
top_indices[j] = candidate_indices[j];
|
||||
}
|
||||
}
|
||||
|
||||
__attribute__((always_inline))
|
||||
static inline void sort32_f32_hvx(uint8_t * values, uint8_t * indices, enum ggml_sort_order order) {
|
||||
bitonic_sort_generic_hvx(values, indices, 1, order == GGML_SORT_ORDER_ASC);
|
||||
@@ -534,3 +672,420 @@ int op_argsort(struct htp_ops_context * octx) {
|
||||
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
// ggml_compute_forward_top_k
|
||||
//
|
||||
// Reuses ARGSORT's sort kernels. Only the first `k` indices are copied
|
||||
// to dst, and there's no asc/desc param -- always largest-first.
|
||||
|
||||
struct htp_top_k_context {
|
||||
struct htp_ops_context * octx;
|
||||
uint32_t nrows_per_thread;
|
||||
uint32_t row_start;
|
||||
uint32_t row_end;
|
||||
uint8_t * vtcm_base;
|
||||
size_t vtcm_per_thread;
|
||||
uint32_t k;
|
||||
};
|
||||
|
||||
#define HTP_TOP_K_FN(ne00, sort_fn) \
|
||||
static void htp_top_k_f32_##ne00(unsigned int n, unsigned int i, void * data) { \
|
||||
struct htp_top_k_context * actx = (struct htp_top_k_context *)data; \
|
||||
struct htp_ops_context * octx = actx->octx; \
|
||||
const struct htp_tensor * src0 = octx->src[0]; \
|
||||
const struct htp_tensor * dst = octx->dst; \
|
||||
uint8_t * spad = actx->vtcm_base + actx->vtcm_per_thread * i; \
|
||||
uint32_t row_start = actx->row_start; \
|
||||
uint32_t row_end = actx->row_end; \
|
||||
uint32_t rows_per_thread = actx->nrows_per_thread; \
|
||||
uint32_t start_row = row_start + rows_per_thread * i; \
|
||||
uint32_t end_row = MIN(start_row + rows_per_thread, row_end); \
|
||||
size_t values_size = hex_round_up(ne00 * sizeof(float), 128); \
|
||||
float * values_buf = (float *) spad; \
|
||||
int32_t * indices_buf = (int32_t *) (spad + values_size); \
|
||||
uint32_t nb01 = src0->nb[1]; \
|
||||
uint32_t nb1 = dst->nb[1]; \
|
||||
uint32_t k = actx->k; \
|
||||
struct htp_thread_trace * tr = &octx->ctx->trace[i]; \
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, start_row); \
|
||||
for (uint32_t r = start_row; r < end_row; r++) { \
|
||||
uint32_t src_offset = r * nb01; \
|
||||
uint32_t dst_offset = r * nb1; \
|
||||
uint8_t * src_ptr = (uint8_t *) src0->data + src_offset; \
|
||||
uint8_t * dst_ptr = (uint8_t *) dst->data + dst_offset; \
|
||||
hex_l2fetch(src_ptr, ne00 * sizeof(float), ne00 * sizeof(float), 1); \
|
||||
hvx_copy_f32_au((uint8_t*)values_buf, src_ptr, ne00); \
|
||||
sort_fn((uint8_t*)values_buf, (uint8_t*)indices_buf, GGML_SORT_ORDER_DESC); \
|
||||
hvx_copy_f32_ua(dst_ptr, (const uint8_t *) indices_buf, k); \
|
||||
} \
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, start_row); \
|
||||
}
|
||||
|
||||
HTP_TOP_K_FN(32, sort32_f32_hvx)
|
||||
HTP_TOP_K_FN(64, sort64_f32_hvx)
|
||||
HTP_TOP_K_FN(128, sort128_f32_hvx)
|
||||
HTP_TOP_K_FN(256, sort256_f32_hvx)
|
||||
HTP_TOP_K_FN(512, sort512_f32_hvx)
|
||||
HTP_TOP_K_FN(1024, sort1024_f32_hvx)
|
||||
|
||||
static void htp_top_k_f32_fallback(unsigned int n, unsigned int i, void * data) {
|
||||
struct htp_top_k_context * actx = (struct htp_top_k_context *)data;
|
||||
struct htp_ops_context * octx = actx->octx;
|
||||
|
||||
// Unpack context
|
||||
const struct htp_tensor * src0 = octx->src[0];
|
||||
const struct htp_tensor * dst = octx->dst;
|
||||
|
||||
// Scratchpad memory
|
||||
uint8_t * spad = actx->vtcm_base + actx->vtcm_per_thread * i;
|
||||
|
||||
// Dimensions
|
||||
uint32_t ne00 = src0->ne[0];
|
||||
uint32_t nb01 = src0->nb[1];
|
||||
|
||||
uint32_t nb1 = dst->nb[1];
|
||||
|
||||
uint32_t k = actx->k;
|
||||
|
||||
// Rows to process
|
||||
uint32_t row_start = actx->row_start;
|
||||
uint32_t row_end = actx->row_end;
|
||||
uint32_t rows_per_thread = actx->nrows_per_thread;
|
||||
uint32_t start_row = row_start + rows_per_thread * i;
|
||||
uint32_t end_row = MIN(start_row + rows_per_thread, row_end);
|
||||
|
||||
// Pad ne00 to n_vec*32 (n_vec a power of 2) for the bitonic network;
|
||||
// pad with -INFINITY so it never lands in the top-k.
|
||||
uint32_t n_vec = hmx_ceil_div(ne00, 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 = hex_round_up(ne00_padded * sizeof(float), 128);
|
||||
float * values_buf = (float *) spad;
|
||||
int32_t * indices_buf = (int32_t *) (spad + values_size);
|
||||
|
||||
struct htp_thread_trace * tr = &octx->ctx->trace[i];
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, start_row);
|
||||
|
||||
for (uint32_t r = start_row; r < end_row; r++) {
|
||||
uint32_t src_offset = r * nb01;
|
||||
uint32_t dst_offset = r * nb1;
|
||||
|
||||
uint8_t * src_ptr = (uint8_t *) src0->data + src_offset;
|
||||
uint8_t * dst_ptr = (uint8_t *) dst->data + dst_offset;
|
||||
|
||||
hex_l2fetch(src_ptr, ne00 * sizeof(float), ne00 * sizeof(float), 1);
|
||||
|
||||
if (k <= 64 && ne00 > 1024) {
|
||||
float top_values[64];
|
||||
int32_t top_indices[64];
|
||||
top_k_select_tiled(src_ptr, ne00, k, values_buf, indices_buf, top_values, top_indices);
|
||||
memcpy(dst_ptr, top_indices, k * sizeof(int32_t));
|
||||
continue;
|
||||
}
|
||||
|
||||
hvx_copy_f32_au((uint8_t*)values_buf, src_ptr, ne00);
|
||||
|
||||
// Fills the indices ramp itself, so no init needed here.
|
||||
if (ne00_padded > ne00) {
|
||||
hvx_splat_f32_u((uint8_t *)(values_buf + ne00), -INFINITY, ne00_padded - ne00);
|
||||
}
|
||||
bitonic_sort_vtcm_desc((uint8_t*)values_buf, (uint8_t*)indices_buf, n_vec_pow2, true);
|
||||
|
||||
// Copy top-k indices back to DDR
|
||||
hvx_copy_f32_ua(dst_ptr, (const uint8_t *) indices_buf, k);
|
||||
}
|
||||
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, start_row);
|
||||
}
|
||||
|
||||
// Single row (ne01=ne02=ne03=1) + large ne00 would otherwise run on one
|
||||
// HVX thread while the rest sit idle. Split the row into n_chunks
|
||||
// power-of-two chunks, sort each in parallel with bitonic_sort_vtcm_desc,
|
||||
// then merge the n_chunks*local_k winners with one more sort that carries
|
||||
// the global index through instead of re-deriving it.
|
||||
struct htp_top_k_chunk_ctx {
|
||||
struct htp_ops_context * octx;
|
||||
uint8_t * vtcm_base;
|
||||
size_t phase1_slot_size;
|
||||
uint32_t ne00;
|
||||
uint32_t chunk_elems;
|
||||
uint32_t local_k;
|
||||
size_t merge_values_off;
|
||||
size_t merge_indices_off;
|
||||
};
|
||||
|
||||
static void htp_top_k_chunk_job(unsigned int n, unsigned int i, void * data) {
|
||||
struct htp_top_k_chunk_ctx * cctx = (struct htp_top_k_chunk_ctx *) data;
|
||||
struct htp_ops_context * octx = cctx->octx;
|
||||
const struct htp_tensor * src0 = octx->src[0];
|
||||
|
||||
uint32_t chunk_elems = cctx->chunk_elems;
|
||||
uint32_t ne00 = cctx->ne00;
|
||||
uint32_t local_k = cctx->local_k;
|
||||
uint32_t chunk_base = i * chunk_elems;
|
||||
|
||||
uint8_t * spad = cctx->vtcm_base + cctx->phase1_slot_size * i;
|
||||
size_t values_size = hex_round_up(chunk_elems * sizeof(float), 128);
|
||||
float * values_buf = (float *) spad;
|
||||
int32_t * indices_buf = (int32_t *) (spad + values_size);
|
||||
|
||||
uint32_t real_count = (chunk_base < ne00) ? MIN(chunk_elems, ne00 - chunk_base) : 0;
|
||||
|
||||
if (real_count == 0) {
|
||||
float * merge_values = (float *) (cctx->vtcm_base + cctx->merge_values_off);
|
||||
int32_t * merge_indices = (int32_t *) (cctx->vtcm_base + cctx->merge_indices_off);
|
||||
for (uint32_t j = 0; j < local_k; j++) {
|
||||
merge_values[i * local_k + j] = -INFINITY;
|
||||
merge_indices[i * local_k + j] = 0;
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
if (local_k > 1 && local_k <= 64 && chunk_elems > 1024) {
|
||||
uint8_t * src_ptr = (uint8_t *) src0->data + (size_t) chunk_base * sizeof(float);
|
||||
float top_values[64];
|
||||
int32_t top_indices[64];
|
||||
float * merge_values = (float *) (cctx->vtcm_base + cctx->merge_values_off);
|
||||
int32_t * merge_indices = (int32_t *) (cctx->vtcm_base + cctx->merge_indices_off);
|
||||
for (uint32_t j = 0; j < local_k; j++) {
|
||||
top_values[j] = -INFINITY;
|
||||
top_indices[j] = 0;
|
||||
}
|
||||
top_k_select_tiled(src_ptr, real_count, local_k, values_buf, indices_buf, top_values, top_indices);
|
||||
|
||||
for (uint32_t j = 0; j < local_k; j++) {
|
||||
merge_values[i * local_k + j] = top_values[j];
|
||||
merge_indices[i * local_k + j] = top_indices[j] + (int32_t) chunk_base;
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
if (real_count > 0) {
|
||||
uint8_t * src_ptr = (uint8_t *) src0->data + (size_t) chunk_base * sizeof(float);
|
||||
hex_l2fetch(src_ptr, real_count * sizeof(float), real_count * sizeof(float), 1);
|
||||
hvx_copy_f32_au((uint8_t *) values_buf, src_ptr, real_count);
|
||||
}
|
||||
if (chunk_elems > real_count) {
|
||||
hvx_splat_f32_u((uint8_t *) (values_buf + real_count), -INFINITY, chunk_elems - real_count);
|
||||
}
|
||||
|
||||
if (local_k == 1 && cctx->ne00 >= 128*1024) {
|
||||
float max_value;
|
||||
uint32_t max_index = top_k_max_value_index(values_buf, real_count, &max_value);
|
||||
float * merge_values = (float *) (cctx->vtcm_base + cctx->merge_values_off);
|
||||
int32_t * merge_indices = (int32_t *) (cctx->vtcm_base + cctx->merge_indices_off);
|
||||
merge_values[i] = max_value;
|
||||
merge_indices[i] = (int32_t) (max_index + chunk_base);
|
||||
return;
|
||||
}
|
||||
|
||||
// chunk_elems is always a power-of-two multiple of 32
|
||||
bitonic_sort_vtcm_desc((uint8_t *) values_buf, (uint8_t *) indices_buf, chunk_elems / 32, true);
|
||||
|
||||
float * merge_values = (float *) (cctx->vtcm_base + cctx->merge_values_off);
|
||||
int32_t * merge_indices = (int32_t *) (cctx->vtcm_base + cctx->merge_indices_off);
|
||||
|
||||
for (uint32_t j = 0; j < local_k; j++) {
|
||||
merge_values[i * local_k + j] = values_buf[j];
|
||||
merge_indices[i * local_k + j] = indices_buf[j] + (int32_t) chunk_base;
|
||||
}
|
||||
}
|
||||
|
||||
struct htp_top_k_merge_ctx {
|
||||
struct htp_ops_context * octx;
|
||||
uint8_t * vtcm_base;
|
||||
size_t merge_values_off;
|
||||
size_t merge_indices_off;
|
||||
uint32_t merge_elems;
|
||||
uint32_t total_candidates;
|
||||
uint32_t k;
|
||||
};
|
||||
|
||||
static void htp_top_k_merge_job(unsigned int n, unsigned int i, void * data) {
|
||||
struct htp_top_k_merge_ctx * mctx = (struct htp_top_k_merge_ctx *) data;
|
||||
struct htp_ops_context * octx = mctx->octx;
|
||||
const struct htp_tensor * dst = octx->dst;
|
||||
|
||||
float * merge_values = (float *) (mctx->vtcm_base + mctx->merge_values_off);
|
||||
int32_t * merge_indices = (int32_t *) (mctx->vtcm_base + mctx->merge_indices_off);
|
||||
|
||||
if (mctx->merge_elems > mctx->total_candidates) {
|
||||
uint32_t pad = mctx->merge_elems - mctx->total_candidates;
|
||||
hvx_splat_f32_u((uint8_t *) (merge_values + mctx->total_candidates), -INFINITY, pad);
|
||||
for (uint32_t j = mctx->total_candidates; j < mctx->merge_elems; j++) {
|
||||
merge_indices[j] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
// Preserve the global indices computed in phase 1 -- init_indices=false
|
||||
// so they aren't overwritten with a local ramp.
|
||||
bitonic_sort_vtcm_desc((uint8_t *) merge_values, (uint8_t *) merge_indices, mctx->merge_elems / 32, false);
|
||||
|
||||
hvx_copy_f32_ua((uint8_t *) dst->data, (const uint8_t *) merge_indices, mctx->k);
|
||||
}
|
||||
|
||||
static int op_top_k_single_row_threaded(struct htp_ops_context * octx, uint32_t ne00, uint32_t k) {
|
||||
uint32_t n_threads_avail = octx->n_threads;
|
||||
|
||||
uint32_t n_vec = hmx_ceil_div(ne00, 32);
|
||||
uint32_t n_vec_pow2 = 1;
|
||||
while (n_vec_pow2 < n_vec) n_vec_pow2 <<= 1;
|
||||
|
||||
// Largest power-of-two chunk count that both fits the available
|
||||
// threads and evenly divides n_vec_pow2
|
||||
uint32_t n_chunks = 1;
|
||||
while (n_chunks * 2 <= n_threads_avail && 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;
|
||||
uint32_t local_k = MIN(k, chunk_elems);
|
||||
|
||||
uint32_t total_candidates = n_chunks * local_k;
|
||||
uint32_t merge_n_vec = hmx_ceil_div(total_candidates, 32);
|
||||
uint32_t merge_n_vec_pow2 = 1;
|
||||
while (merge_n_vec_pow2 < merge_n_vec) merge_n_vec_pow2 <<= 1;
|
||||
uint32_t merge_elems = merge_n_vec_pow2 * 32;
|
||||
|
||||
size_t phase1_values_size = hex_round_up(chunk_elems * sizeof(float), 128);
|
||||
size_t phase1_indices_size = hex_round_up(chunk_elems * sizeof(int32_t), 128);
|
||||
size_t phase1_slot_size = hex_round_up(phase1_values_size + phase1_indices_size, 256);
|
||||
size_t phase1_total_size = phase1_slot_size * n_chunks;
|
||||
|
||||
size_t merge_values_size = hex_round_up(merge_elems * sizeof(float), 128);
|
||||
size_t merge_indices_size = hex_round_up(merge_elems * sizeof(int32_t), 128);
|
||||
size_t merge_values_off = phase1_total_size;
|
||||
size_t merge_indices_off = merge_values_off + merge_values_size;
|
||||
|
||||
size_t total_vtcm = phase1_total_size + merge_values_size + merge_indices_size;
|
||||
if (octx->ctx->vtcm_size < total_vtcm) {
|
||||
return HTP_STATUS_VTCM_TOO_SMALL;
|
||||
}
|
||||
|
||||
uint8_t * vtcm_base = (uint8_t *) octx->ctx->vtcm_base;
|
||||
|
||||
struct htp_top_k_chunk_ctx cctx;
|
||||
cctx.octx = octx;
|
||||
cctx.vtcm_base = vtcm_base;
|
||||
cctx.phase1_slot_size = phase1_slot_size;
|
||||
cctx.ne00 = ne00;
|
||||
cctx.chunk_elems = chunk_elems;
|
||||
cctx.local_k = local_k;
|
||||
cctx.merge_values_off = merge_values_off;
|
||||
cctx.merge_indices_off = merge_indices_off;
|
||||
|
||||
work_queue_run(octx->ctx->work_queue, htp_top_k_chunk_job, &cctx, n_chunks);
|
||||
|
||||
struct htp_top_k_merge_ctx mctx;
|
||||
mctx.octx = octx;
|
||||
mctx.vtcm_base = vtcm_base;
|
||||
mctx.merge_values_off = merge_values_off;
|
||||
mctx.merge_indices_off = merge_indices_off;
|
||||
mctx.merge_elems = merge_elems;
|
||||
mctx.total_candidates = total_candidates;
|
||||
mctx.k = k;
|
||||
|
||||
work_queue_run(octx->ctx->work_queue, htp_top_k_merge_job, &mctx, 1);
|
||||
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
int op_top_k(struct htp_ops_context * octx) {
|
||||
// Check supported types
|
||||
if (octx->src[0]->type != HTP_TYPE_F32) {
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
|
||||
const struct htp_tensor * src0 = octx->src[0];
|
||||
const struct htp_tensor * dst = octx->dst;
|
||||
|
||||
const uint32_t total_rows = src0->ne[1] * src0->ne[2] * src0->ne[3];
|
||||
const size_t dst_row_size = dst->ne[0] * sizeof(int32_t);
|
||||
|
||||
uint32_t row_start = 0;
|
||||
uint32_t row_end = total_rows;
|
||||
if (octx->ctx->mdev.count > 1) {
|
||||
uint32_t rows_per_chunk = 0;
|
||||
htp_tensor_mdev_rows_per_chunk(dst, sizeof(int32_t), (uint32_t) dst_row_size, &rows_per_chunk);
|
||||
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, rows_per_chunk,
|
||||
octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
|
||||
row_start = range.start;
|
||||
row_end = range.start + range.count;
|
||||
}
|
||||
|
||||
const uint32_t nrows = row_end - row_start;
|
||||
if (nrows == 0) {
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
uint32_t ne00 = src0->ne[0];
|
||||
uint32_t k = dst->ne[0];
|
||||
|
||||
// Single row + large ne00: the per-row dispatch below would run on one
|
||||
// HVX thread while the rest sit idle. Split the row across threads.
|
||||
if (total_rows == 1 && ne00 > 1024) {
|
||||
int status = op_top_k_single_row_threaded(octx, ne00, k);
|
||||
if (status != HTP_STATUS_VTCM_TOO_SMALL) {
|
||||
return status;
|
||||
}
|
||||
// else: fall through to the single-thread path below.
|
||||
}
|
||||
|
||||
const uint32_t n_threads = MIN(nrows, octx->n_threads);
|
||||
|
||||
// Scratchpad layout: values + indices
|
||||
// For bitonic: need padding to power-of-2 size
|
||||
// Allocate for worst case (bitonic with padding)
|
||||
uint32_t n_vec = hmx_ceil_div(ne00, 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 = hex_round_up(ne00_padded * sizeof(float), 128);
|
||||
size_t indices_size = hex_round_up(ne00_padded * sizeof(int32_t), 128);
|
||||
size_t spad_per_thread = values_size + indices_size;
|
||||
|
||||
// Make sure we round up to 256 for alignment requirements
|
||||
spad_per_thread = hex_round_up(spad_per_thread, 256);
|
||||
|
||||
size_t total_spad_size = spad_per_thread * n_threads;
|
||||
|
||||
if (octx->ctx->vtcm_size < total_spad_size) {
|
||||
FARF(ERROR, "top_k: VTCM size too small. Needed %zu, have %zu", total_spad_size, octx->ctx->vtcm_size);
|
||||
return HTP_STATUS_VTCM_TOO_SMALL;
|
||||
}
|
||||
|
||||
FARF(HIGH, "top_k: %ux%ux%ux%u -> %ux%ux%ux%u (0x%x, 0x%x)",
|
||||
octx->src[0]->ne[0], octx->src[0]->ne[1], octx->src[0]->ne[2], octx->src[0]->ne[3],
|
||||
octx->dst->ne[0], octx->dst->ne[1], octx->dst->ne[2], octx->dst->ne[3],
|
||||
octx->src[0]->data, octx->dst->data);
|
||||
|
||||
struct htp_top_k_context actx;
|
||||
const struct fastdiv_values n_threads_div = init_fastdiv_values(n_threads);
|
||||
actx.octx = octx;
|
||||
actx.nrows_per_thread = fastdiv(nrows + n_threads - 1, &n_threads_div);
|
||||
actx.row_start = row_start;
|
||||
actx.row_end = row_end;
|
||||
actx.vtcm_base = (uint8_t *) octx->ctx->vtcm_base;
|
||||
actx.vtcm_per_thread = spad_per_thread;
|
||||
actx.k = k;
|
||||
|
||||
worker_callback_t job_func = htp_top_k_f32_fallback;
|
||||
switch (ne00) {
|
||||
case 1024: job_func = htp_top_k_f32_1024; break;
|
||||
case 512: job_func = htp_top_k_f32_512; break;
|
||||
case 256: job_func = htp_top_k_f32_256; break;
|
||||
case 128: job_func = htp_top_k_f32_128; break;
|
||||
case 64: job_func = htp_top_k_f32_64; break;
|
||||
case 32: job_func = htp_top_k_f32_32; break;
|
||||
default: job_func = htp_top_k_f32_fallback; break;
|
||||
}
|
||||
|
||||
// Run jobs
|
||||
work_queue_run(octx->ctx->work_queue, job_func, &actx, n_threads);
|
||||
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
@@ -108,6 +108,7 @@ struct hmx_fa_context {
|
||||
|
||||
// Dimensions
|
||||
uint32_t DK, DV;
|
||||
uint32_t DK_pad, DV_pad; // head_dim rounded up to 64 for HMX tiling
|
||||
uint32_t n_kv; // kv_len
|
||||
uint32_t n_kv_heads; // number of KV heads
|
||||
uint32_t n_heads; // number of Q heads
|
||||
@@ -652,7 +653,7 @@ static void fa_k_interleave_thread(unsigned int n, unsigned int i, void * data)
|
||||
hvx_dequantize_row_q8_0_f16(row_k, row_k, factx->DK);
|
||||
}
|
||||
}
|
||||
hmx_interleave_rows_to_tiles(factx->vtcm_k_tiles[args->buf_idx], (const __fp16 *) args->curr_k, total_rows, factx->DK,
|
||||
hmx_interleave_rows_to_tiles(factx->vtcm_k_tiles[args->buf_idx], (const __fp16 *) args->curr_k, total_rows, factx->DK_pad,
|
||||
args->src_stride, start, end);
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_K_PREP, (uint16_t) (args->kv_start + start));
|
||||
}
|
||||
@@ -706,7 +707,7 @@ static void fa_v_interleave_thread(unsigned int n, unsigned int i, void * data)
|
||||
hvx_dequantize_row_q8_0_f16(row_v, row_v, factx->DV);
|
||||
}
|
||||
}
|
||||
hmx_interleave_cols_to_tiles(v_tiles_dst, (const __fp16 *) args->v_src, total_rows, factx->DV,
|
||||
hmx_interleave_cols_to_tiles(v_tiles_dst, (const __fp16 *) args->v_src, total_rows, factx->DV_pad,
|
||||
args->src_stride, (uint32_t) args->n_col_tiles, start, end);
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_V_PREP, (uint16_t) (args->kv_start + start));
|
||||
}
|
||||
@@ -832,17 +833,22 @@ static void fa_q_load_thread(unsigned int n, unsigned int i, void * data) {
|
||||
const uint32_t kv_head = args->kv_head;
|
||||
const uint32_t ib3 = args->ib3;
|
||||
|
||||
assert(factx->DK == factx->DV);
|
||||
|
||||
const bool use_q_dma = (factx->vtcm_q_dma != NULL);
|
||||
|
||||
__fp16 * q_tiles = factx->vtcm_q_tiles;
|
||||
const size_t DK_pad = factx->DK_pad;
|
||||
if (use_q_dma) {
|
||||
const size_t g_rows_end = hex_smin(end, n_rows_g);
|
||||
const uint32_t d_limit = factx->is_q_fp32 ? DK / 32 : DK / 64;
|
||||
|
||||
uint8_t * q_flat = (uint8_t *) factx->vtcm_q_dma;
|
||||
if (factx->is_q_fp32) {
|
||||
if (DK_pad != DK) {
|
||||
if (factx->is_q_fp32) {
|
||||
hmx_fa_q_prep_fp32_pad(q_tiles, q_flat, start, end, g_rows_end, DK, DK_pad, G, args->n_rows_q, &factx->div_G, args->q_transposed);
|
||||
} else {
|
||||
hmx_fa_q_prep_fp16_pad(q_tiles, q_flat, start, end, g_rows_end, DK, DK_pad, G, args->n_rows_q, &factx->div_G, args->q_transposed);
|
||||
}
|
||||
} else if (factx->is_q_fp32) {
|
||||
switch (d_limit) {
|
||||
case 2: hmx_fa_q_prep_fp32_d2(q_tiles, q_flat, start, end, g_rows_end, DK, G, args->n_rows_q, &factx->div_G, args->q_transposed); break;
|
||||
case 4: hmx_fa_q_prep_fp32_d4(q_tiles, q_flat, start, end, g_rows_end, DK, G, args->n_rows_q, &factx->div_G, args->q_transposed); break;
|
||||
@@ -858,7 +864,7 @@ static void fa_q_load_thread(unsigned int n, unsigned int i, void * data) {
|
||||
} else {
|
||||
// Fallback: direct-from-DDR/L2 path
|
||||
hmx_fa_q_prep_fallback(q_tiles, q->data, q->nb[1], q->nb[2], q->nb[3],
|
||||
q_start, kv_head, ib3, start, end, n_rows_g, G, DK, factx->is_q_fp32, &factx->div_G);
|
||||
q_start, kv_head, ib3, start, end, n_rows_g, G, DK, DK_pad, factx->is_q_fp32, &factx->div_G);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -952,6 +958,8 @@ static void fa_o_store_thread_f32(unsigned int n, unsigned int i, void * data) {
|
||||
const uint32_t kv_head = args->kv_head;
|
||||
const uint32_t ib3 = args->ib3;
|
||||
|
||||
const size_t DV_pad = factx->DV_pad;
|
||||
|
||||
size_t q_idx = fastdiv(start, &factx->div_G);
|
||||
size_t h_idx = fastmodulo(start, G, &factx->div_G);
|
||||
|
||||
@@ -961,7 +969,7 @@ static void fa_o_store_thread_f32(unsigned int n, unsigned int i, void * data) {
|
||||
|
||||
size_t r0 = r / HMX_FP16_TILE_N_ROWS;
|
||||
size_t r1 = r % HMX_FP16_TILE_N_ROWS;
|
||||
const __fp16 * tile_row_base = o_tile_src + r0 * HMX_FP16_TILE_N_ROWS * DV;
|
||||
const __fp16 * tile_row_base = o_tile_src + r0 * HMX_FP16_TILE_N_ROWS * DV_pad;
|
||||
|
||||
for (uint32_t d = 0; d < DV / 32; ++d) {
|
||||
const HVX_Vector * in_tile = (const HVX_Vector *) (tile_row_base + d * HMX_FP16_TILE_N_ELMS);
|
||||
@@ -972,6 +980,16 @@ static void fa_o_store_thread_f32(unsigned int n, unsigned int i, void * data) {
|
||||
*(HVX_UVector *) (out + d * 32) = Q6_V_hi_W(vp);
|
||||
}
|
||||
}
|
||||
// Ragged tail: DV not a multiple of 32 (e.g. 72 -> last 8 lanes). Partial vector-write
|
||||
// for the remaining (DV % 32) floats.
|
||||
const uint32_t d_tail = DV / 32;
|
||||
const uint32_t rem = DV - d_tail * 32;
|
||||
if (rem) {
|
||||
const HVX_Vector * in_tile = (const HVX_Vector *) (tile_row_base + d_tail * HMX_FP16_TILE_N_ELMS);
|
||||
HVX_VectorPair vp = hvx_vec_f16_to_f32_shuff(in_tile[r1 / 2]);
|
||||
HVX_Vector vd = (r1 % 2 == 0) ? Q6_V_lo_W(vp) : Q6_V_hi_W(vp);
|
||||
hvx_vec_store_u((void *) (out + d_tail * 32), rem * sizeof(float), vd);
|
||||
}
|
||||
|
||||
h_idx++;
|
||||
if (h_idx == G) {
|
||||
@@ -1006,6 +1024,9 @@ static void fa_o_store_thread_f16(unsigned int n, unsigned int i, void * data) {
|
||||
const uint32_t kv_head = args->kv_head;
|
||||
const uint32_t ib3 = args->ib3;
|
||||
|
||||
// O-tiles use the padded head dim (DV_pad); dst holds the real DV lanes.
|
||||
const size_t DV_pad = factx->DV_pad;
|
||||
|
||||
size_t q_idx = fastdiv(start, &factx->div_G);
|
||||
size_t h_idx = fastmodulo(start, G, &factx->div_G);
|
||||
|
||||
@@ -1015,7 +1036,7 @@ static void fa_o_store_thread_f16(unsigned int n, unsigned int i, void * data) {
|
||||
|
||||
size_t r0 = r / HMX_FP16_TILE_N_ROWS;
|
||||
size_t r1 = r % HMX_FP16_TILE_N_ROWS;
|
||||
const __fp16 * tile_row_base = o_tile_src + r0 * HMX_FP16_TILE_N_ROWS * DV;
|
||||
const __fp16 * tile_row_base = o_tile_src + r0 * HMX_FP16_TILE_N_ROWS * DV_pad;
|
||||
|
||||
for (uint32_t d = 0; d < DV / 64; ++d) {
|
||||
const __fp16 * in_dtile = tile_row_base + d * HMX_FP16_TILE_N_ELMS * 2;
|
||||
@@ -1028,6 +1049,17 @@ static void fa_o_store_thread_f16(unsigned int n, unsigned int i, void * data) {
|
||||
*(HVX_UVector *) (out + d * 64) = Q6_V_hi_W(vp);
|
||||
}
|
||||
}
|
||||
// Ragged tail when DV is not a multiple of 64.
|
||||
const uint32_t d_tail = DV / 64;
|
||||
const uint32_t rem = DV - d_tail * 64;
|
||||
if (rem) {
|
||||
const __fp16 * in_dtile = tile_row_base + d_tail * HMX_FP16_TILE_N_ELMS * 2;
|
||||
const HVX_Vector * pv_in0 = ((const HVX_Vector *) in_dtile) + r1 / 2;
|
||||
const HVX_Vector * pv_in1 = pv_in0 + 16;
|
||||
HVX_VectorPair vp = Q6_W_vdeal_VVR(*pv_in1, *pv_in0, -2);
|
||||
HVX_Vector vd = (r1 % 2 == 0) ? Q6_V_lo_W(vp) : Q6_V_hi_W(vp);
|
||||
hvx_vec_store_u((void *) (out + d_tail * 64), rem * sizeof(__fp16), vd);
|
||||
}
|
||||
|
||||
h_idx++;
|
||||
if (h_idx == G) {
|
||||
@@ -1829,8 +1861,11 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
|
||||
const uint32_t DK = neq0;
|
||||
const uint32_t DV = nev0;
|
||||
|
||||
// HMX requires head_dim to be multiple of 32
|
||||
if (DK % 32 != 0 || DV % 32 != 0) {
|
||||
// HMX tiles head_dim in units of 64. head_dim need not be 64- (or 32-) aligned:
|
||||
// we can operate on DK/DV rounded up to 64 with tail lanes [D, D_pad) zero-filled.
|
||||
const uint32_t DK_pad = hex_round_up(DK, 64);
|
||||
const uint32_t DV_pad = hex_round_up(DV, 64);
|
||||
if (DK == 0 || DV == 0) {
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
|
||||
@@ -1847,6 +1882,8 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
|
||||
factx.n_threads = kparams->n_threads;
|
||||
factx.DK = DK;
|
||||
factx.DV = DV;
|
||||
factx.DK_pad = DK_pad;
|
||||
factx.DV_pad = DV_pad;
|
||||
factx.n_kv = nek1;
|
||||
factx.n_kv_heads = n_kv_heads;
|
||||
factx.n_heads = neq2;
|
||||
@@ -1905,16 +1942,18 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
|
||||
|
||||
// ======== VTCM allocation (GQA-aware) ========
|
||||
// K/V row sizes drive the DMA descriptors (not the VTCM layout) and are used
|
||||
// throughout the KV loop below.
|
||||
// throughout the KV loop below. The DMA copies only the real DK/DV columns; the
|
||||
// staging rows are padded to hold DK_pad/DV_pad columns (tail zero-filled below)
|
||||
// so the HMX interleave/tile logic can operate on 64-aligned head dims.
|
||||
const size_t size_k_row = htp_tensor_get_row_size(k->type, DK);
|
||||
const size_t size_v_row = htp_tensor_get_row_size(v->type, DV);
|
||||
const size_t size_k_row_padded = hex_round_up(DK * sizeof(__fp16), 128);
|
||||
const size_t size_v_row_padded = hex_round_up(DV * sizeof(__fp16), 128);
|
||||
const size_t size_k_row_padded = hex_round_up(DK_pad * sizeof(__fp16), 128);
|
||||
const size_t size_v_row_padded = hex_round_up(DV_pad * sizeof(__fp16), 128);
|
||||
|
||||
// Build the VTCM layout once (shared with the host estimator) and place every
|
||||
// scratch buffer at its computed offset.
|
||||
// scratch buffer at its computed offset. Padded head dims size the HMX tiles.
|
||||
struct hmx_fa_vtcm_layout L;
|
||||
hmx_fa_vtcm_layout_build(&L, G, DK, DV, Br, Bc, n_threads, pipeline, factx.is_q_fp32);
|
||||
hmx_fa_vtcm_layout_build(&L, G, DK_pad, DV_pad, Br, Bc, n_threads, pipeline, factx.is_q_fp32);
|
||||
|
||||
if (L.total_bytes > ctx->vtcm_size) {
|
||||
return HTP_STATUS_VTCM_TOO_SMALL;
|
||||
@@ -1961,6 +2000,24 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
|
||||
|
||||
dma_cache_init(&factx.m_cache, (uint8_t *) factx.vtcm_mask_buf, L.m_buf_slot_bytes, HMX_FA_DMA_CACHE_SIZE);
|
||||
|
||||
// Head-dim padding: the K/V DMA staging buffers and the flat-Q buffer are laid out
|
||||
// with padded row strides (size_{k,v,q}_row_padded, covering D_pad columns) but the
|
||||
// DMA only writes the real D columns per row. Zero the whole staging buffers once up
|
||||
// front so tail lanes [D, D_pad) stay zero for all KV blocks. No-op when already aligned.
|
||||
if (DK_pad != DK || DV_pad != DV) {
|
||||
const size_t k_buf_bytes = (size_t) factx.Bc * size_k_row_padded;
|
||||
const size_t v_buf_bytes = (size_t) factx.Bc * size_v_row_padded;
|
||||
hvx_splat_u8_a((char *) factx.vtcm_k_fp16[0], 0, k_buf_bytes);
|
||||
hvx_splat_u8_a((char *) factx.vtcm_k_fp16[1], 0, k_buf_bytes);
|
||||
hvx_splat_u8_a((char *) factx.vtcm_v_fp16[0], 0, v_buf_bytes);
|
||||
hvx_splat_u8_a((char *) factx.vtcm_v_fp16[1], 0, v_buf_bytes);
|
||||
// Flat-Q DMA scratch
|
||||
if (factx.vtcm_q_dma) {
|
||||
const size_t q_dma_bytes = hex_align_up(factx.g_br * DK * (factx.is_q_fp32 ? sizeof(float) : sizeof(__fp16)), 128);
|
||||
hvx_splat_u8_a((char *) factx.vtcm_q_dma, 0, q_dma_bytes);
|
||||
}
|
||||
}
|
||||
|
||||
// ======== Initialize HMX output scales ========
|
||||
hmx_init_column_scales(factx.vtcm_hmx_scales_id, Q6_V_vsplat_R(0x3c00)); // 1.0
|
||||
hmx_init_column_scales(factx.vtcm_hmx_scales_qk, hvx_vec_splat_f16(factx.scale));
|
||||
@@ -2072,7 +2129,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
|
||||
qk_job[0].s_tiles = factx.vtcm_s_tiles[0];
|
||||
qk_job[0].n_row_tiles = n_row_tiles;
|
||||
qk_job[0].n_col_tiles = hmx_ceil_div(kv_rows0, HMX_FP16_TILE_N_COLS);
|
||||
qk_job[0].n_dot_tiles = DK / 32;
|
||||
qk_job[0].n_dot_tiles = DK_pad / 32;
|
||||
qk_job[0].n_tiles_per_bc = n_tiles_per_bc;
|
||||
qk_job[0].hmx_scales = factx.vtcm_hmx_scales_qk;
|
||||
hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_qk_dot_worker, &qk_job[0]));
|
||||
@@ -2116,7 +2173,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
|
||||
hmx_ceil_div(hex_smin(Bc, nek1 - (kv_blk - 1) * Bc), HMX_FP16_TILE_N_COLS);
|
||||
ou_job[prev_buf].n_row_tiles_g_br = n_row_tiles_g_br;
|
||||
ou_job[prev_buf].n_tiles_per_bc = n_tiles_per_bc;
|
||||
ou_job[prev_buf].DV = DV;
|
||||
ou_job[prev_buf].DV = DV_pad;
|
||||
hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job[prev_buf]));
|
||||
}
|
||||
|
||||
@@ -2134,7 +2191,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
|
||||
qk_job[next_buf].s_tiles = factx.vtcm_s_tiles[next_buf];
|
||||
qk_job[next_buf].n_row_tiles = n_row_tiles;
|
||||
qk_job[next_buf].n_col_tiles = hmx_ceil_div(next_rows, HMX_FP16_TILE_N_COLS);
|
||||
qk_job[next_buf].n_dot_tiles = DK / 32;
|
||||
qk_job[next_buf].n_dot_tiles = DK_pad / 32;
|
||||
qk_job[next_buf].n_tiles_per_bc = n_tiles_per_bc;
|
||||
qk_job[next_buf].hmx_scales = factx.vtcm_hmx_scales_qk;
|
||||
hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_qk_dot_worker, &qk_job[next_buf]));
|
||||
@@ -2198,7 +2255,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
|
||||
ou_job[0].n_col_tiles = last_cols;
|
||||
ou_job[0].n_row_tiles_g_br = n_row_tiles_g_br;
|
||||
ou_job[0].n_tiles_per_bc = n_tiles_per_bc;
|
||||
ou_job[0].DV = DV;
|
||||
ou_job[0].DV = DV_pad;
|
||||
hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job[0]));
|
||||
|
||||
// Overlapped: run HVX build diag inv L while HMX is busy executing the update
|
||||
@@ -2246,7 +2303,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
|
||||
qk_job.s_tiles = factx.vtcm_s_tiles[0];
|
||||
qk_job.n_row_tiles = n_row_tiles;
|
||||
qk_job.n_col_tiles = n_col_tiles;
|
||||
qk_job.n_dot_tiles = (size_t) (DK / 32);
|
||||
qk_job.n_dot_tiles = (size_t) (DK_pad / 32);
|
||||
qk_job.n_tiles_per_bc = n_tiles_per_bc;
|
||||
qk_job.hmx_scales = factx.vtcm_hmx_scales_qk;
|
||||
|
||||
@@ -2302,7 +2359,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
|
||||
ou_job.n_col_tiles = n_col_tiles;
|
||||
ou_job.n_row_tiles_g_br = n_row_tiles_g_br;
|
||||
ou_job.n_tiles_per_bc = n_tiles_per_bc;
|
||||
ou_job.DV = DV;
|
||||
ou_job.DV = DV_pad;
|
||||
|
||||
hmx_queue_push(ctx->hmx_queue, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job));
|
||||
if (kv_blk + 1 == factx.n_kv_blocks) {
|
||||
@@ -2380,7 +2437,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
|
||||
on_job.hmx_scales = factx.vtcm_hmx_scales_id;
|
||||
on_job.n_row_tiles = n_row_tiles;
|
||||
on_job.n_row_tiles_g_br = n_row_tiles_g_br;
|
||||
on_job.DV = DV;
|
||||
on_job.DV = DV_pad;
|
||||
hmx_queue_push(ctx->hmx_queue, hmx_queue_make_desc(hmx_fa_o_norm_worker, &on_job));
|
||||
hmx_queue_pop(ctx->hmx_queue);
|
||||
}
|
||||
|
||||
@@ -213,11 +213,13 @@ int op_get_rows(struct htp_ops_context * octx) {
|
||||
|
||||
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_Q8_0 &&
|
||||
octx->src[0]->type != HTP_TYPE_I32) {
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
|
||||
if (octx->dst->type != HTP_TYPE_F32) {
|
||||
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;
|
||||
}
|
||||
|
||||
@@ -275,6 +277,7 @@ int op_get_rows(struct htp_ops_context * octx) {
|
||||
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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -495,12 +495,140 @@ static inline void hmx_fa_q_prep_fp16(
|
||||
}
|
||||
|
||||
|
||||
// Head-dim-padded Q-prep (f32). Used when DK is not a multiple of 64.
|
||||
static inline void hmx_fa_q_prep_fp32_pad(__fp16 * vtcm_q_tiles,
|
||||
const uint8_t * temp_q_vtcm,
|
||||
size_t start,
|
||||
size_t end,
|
||||
size_t g_rows_end,
|
||||
size_t dk_in,
|
||||
size_t dk_out,
|
||||
size_t G,
|
||||
size_t n_rows_q,
|
||||
const struct fastdiv_values * div_G,
|
||||
bool q_transposed) {
|
||||
const uint32_t n_out_tiles = (uint32_t) (dk_out / 32);
|
||||
for (size_t r = start; r < end; r += 2) {
|
||||
size_t r0 = r / HMX_FP16_TILE_N_ROWS;
|
||||
size_t r1 = r % HMX_FP16_TILE_N_ROWS;
|
||||
__fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * dk_out;
|
||||
|
||||
if (r >= g_rows_end) {
|
||||
for (uint32_t d = 0; d < n_out_tiles; ++d) {
|
||||
((HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS))[r1 / 2] = Q6_V_vzero();
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
const size_t q_idx0 = fastdiv(r + 0, div_G);
|
||||
const size_t h_idx0 = fastmodulo(r + 0, G, div_G);
|
||||
const size_t q_idx1 = fastdiv(r + 1, div_G);
|
||||
const size_t h_idx1 = fastmodulo(r + 1, G, div_G);
|
||||
|
||||
const size_t offset0 = q_transposed ? (h_idx0 * n_rows_q + q_idx0) : (q_idx0 * G + h_idx0);
|
||||
const size_t offset1 = q_transposed ? (h_idx1 * n_rows_q + q_idx1) : (q_idx1 * G + h_idx1);
|
||||
|
||||
const HVX_UVector * pv_in0 = (const HVX_UVector *) (temp_q_vtcm + offset0 * dk_in * sizeof(float));
|
||||
const HVX_UVector * pv_in1 = (r + 1 < g_rows_end) ? (const HVX_UVector *) (temp_q_vtcm + offset1 * dk_in * sizeof(float)) : NULL;
|
||||
|
||||
for (uint32_t d = 0; d < n_out_tiles; ++d) {
|
||||
HVX_Vector * out_tile = (HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS);
|
||||
const size_t base_lane = (size_t) d * 32;
|
||||
const size_t real_lanes = (base_lane < dk_in) ? hex_smin(32, dk_in - base_lane) : 0;
|
||||
|
||||
if (real_lanes == 0) {
|
||||
out_tile[r1 / 2] = Q6_V_vzero();
|
||||
continue;
|
||||
}
|
||||
|
||||
HVX_Vector v0 = pv_in0[d];
|
||||
HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero();
|
||||
if (real_lanes < 32) {
|
||||
// Straddle tile: keep the first real_lanes floats, zero the padded tail so
|
||||
// the packed f16 lanes beyond DK are zero.
|
||||
const HVX_VectorPred keep = Q6_Q_vsetq_R((uint32_t) (real_lanes * sizeof(float)));
|
||||
v0 = Q6_V_vmux_QVV(keep, v0, Q6_V_vzero());
|
||||
v1 = Q6_V_vmux_QVV(keep, v1, Q6_V_vzero());
|
||||
}
|
||||
out_tile[r1 / 2] = hvx_vec_f32_to_f16_shuff(v0, v1);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Head-dim-padded Q-prep (f16). Used when DK is not a multiple of 64.
|
||||
static inline void hmx_fa_q_prep_fp16_pad(__fp16 * vtcm_q_tiles,
|
||||
const uint8_t * temp_q_vtcm,
|
||||
size_t start,
|
||||
size_t end,
|
||||
size_t g_rows_end,
|
||||
size_t dk_in,
|
||||
size_t dk_out,
|
||||
size_t G,
|
||||
size_t n_rows_q,
|
||||
const struct fastdiv_values * div_G,
|
||||
bool q_transposed) {
|
||||
const uint32_t n_out_pairs = (uint32_t) (dk_out / 64);
|
||||
for (size_t r = start; r < end; r += 2) {
|
||||
size_t r0 = r / HMX_FP16_TILE_N_ROWS;
|
||||
size_t r1 = r % HMX_FP16_TILE_N_ROWS;
|
||||
__fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * dk_out;
|
||||
|
||||
if (r >= g_rows_end) {
|
||||
for (uint32_t d = 0; d < n_out_pairs; ++d) {
|
||||
__fp16 * out_dtile = out_base + d * HMX_FP16_TILE_N_ELMS * 2;
|
||||
HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
|
||||
HVX_Vector * pv_out1 = pv_out0 + 16;
|
||||
*pv_out0 = Q6_V_vzero();
|
||||
*pv_out1 = Q6_V_vzero();
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
const size_t q_idx0 = fastdiv(r + 0, div_G);
|
||||
const size_t h_idx0 = fastmodulo(r + 0, G, div_G);
|
||||
const size_t q_idx1 = fastdiv(r + 1, div_G);
|
||||
const size_t h_idx1 = fastmodulo(r + 1, G, div_G);
|
||||
|
||||
const size_t offset0 = q_transposed ? (h_idx0 * n_rows_q + q_idx0) : (q_idx0 * G + h_idx0);
|
||||
const size_t offset1 = q_transposed ? (h_idx1 * n_rows_q + q_idx1) : (q_idx1 * G + h_idx1);
|
||||
|
||||
const HVX_UVector * pv_in0 = (const HVX_UVector *) (temp_q_vtcm + offset0 * dk_in * sizeof(__fp16));
|
||||
const HVX_UVector * pv_in1 = (r + 1 < g_rows_end) ? (const HVX_UVector *) (temp_q_vtcm + offset1 * dk_in * sizeof(__fp16)) : NULL;
|
||||
|
||||
for (uint32_t d = 0; d < n_out_pairs; ++d) {
|
||||
__fp16 * out_dtile = out_base + d * HMX_FP16_TILE_N_ELMS * 2;
|
||||
HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
|
||||
HVX_Vector * pv_out1 = pv_out0 + 16;
|
||||
|
||||
const size_t base_lane = (size_t) d * 64;
|
||||
const size_t real_lanes = (base_lane < dk_in) ? hex_smin(64, dk_in - base_lane) : 0;
|
||||
|
||||
if (real_lanes == 0) {
|
||||
*pv_out0 = Q6_V_vzero();
|
||||
*pv_out1 = Q6_V_vzero();
|
||||
continue;
|
||||
}
|
||||
|
||||
HVX_Vector v0 = pv_in0[d];
|
||||
HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero();
|
||||
if (real_lanes < 64) {
|
||||
const HVX_VectorPred keep = Q6_Q_vsetq_R((uint32_t) (real_lanes * sizeof(__fp16)));
|
||||
v0 = Q6_V_vmux_QVV(keep, v0, Q6_V_vzero());
|
||||
v1 = Q6_V_vmux_QVV(keep, v1, Q6_V_vzero());
|
||||
}
|
||||
HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
|
||||
*pv_out0 = Q6_V_lo_W(vp);
|
||||
*pv_out1 = Q6_V_hi_W(vp);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
static inline void hmx_fa_q_prep_fallback(
|
||||
__fp16 * vtcm_q_tiles, uintptr_t q_data,
|
||||
size_t q_nb1, size_t q_nb2, size_t q_nb3,
|
||||
uint32_t q_start, uint32_t kv_head, uint32_t ib3,
|
||||
size_t start, size_t end, size_t n_rows_g,
|
||||
size_t G, size_t DK, bool is_q_fp32,
|
||||
size_t G, size_t dk_in, size_t dk_out, bool is_q_fp32,
|
||||
const struct fastdiv_values * div_G
|
||||
) {
|
||||
for (size_t r = start; r < end; r += 2) {
|
||||
@@ -518,33 +646,55 @@ static inline void hmx_fa_q_prep_fallback(
|
||||
|
||||
size_t r0 = r / HMX_FP16_TILE_N_ROWS;
|
||||
size_t r1 = r % HMX_FP16_TILE_N_ROWS;
|
||||
__fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * DK;
|
||||
__fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * dk_out;
|
||||
|
||||
if (is_q_fp32) {
|
||||
const HVX_UVector * pv_in0 = q_ptr0 ? (const HVX_UVector *) q_ptr0 : NULL;
|
||||
const HVX_UVector * pv_in1 = q_ptr1 ? (const HVX_UVector *) q_ptr1 : NULL;
|
||||
|
||||
for (uint32_t d = 0; d < DK / 32; ++d) {
|
||||
HVX_Vector v0 = pv_in0 ? pv_in0[d] : Q6_V_vzero();
|
||||
HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero();
|
||||
HVX_Vector v_hf = hvx_vec_f32_to_f16_shuff(v0, v1);
|
||||
for (uint32_t d = 0; d < dk_out / 32; ++d) {
|
||||
HVX_Vector * out_tile = (HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS);
|
||||
const size_t base_lane = (size_t) d * 32;
|
||||
const size_t real_lanes = (base_lane < dk_in) ? hex_smin(32, dk_in - base_lane) : 0;
|
||||
|
||||
HVX_Vector * out_tile = (HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS);
|
||||
out_tile[r1 / 2] = v_hf;
|
||||
if (real_lanes == 0) {
|
||||
out_tile[r1 / 2] = Q6_V_vzero();
|
||||
continue;
|
||||
}
|
||||
HVX_Vector v0 = pv_in0 ? pv_in0[d] : Q6_V_vzero();
|
||||
HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero();
|
||||
if (real_lanes < 32) {
|
||||
const HVX_VectorPred keep = Q6_Q_vsetq_R((uint32_t) (real_lanes * sizeof(float)));
|
||||
v0 = Q6_V_vmux_QVV(keep, v0, Q6_V_vzero());
|
||||
v1 = Q6_V_vmux_QVV(keep, v1, Q6_V_vzero());
|
||||
}
|
||||
out_tile[r1 / 2] = hvx_vec_f32_to_f16_shuff(v0, v1);
|
||||
}
|
||||
} else {
|
||||
const HVX_UVector * pv_in0 = q_ptr0 ? (const HVX_UVector *) q_ptr0 : NULL;
|
||||
const HVX_UVector * pv_in1 = q_ptr1 ? (const HVX_UVector *) q_ptr1 : NULL;
|
||||
|
||||
for (uint32_t d = 0; d < DK / 64; ++d) {
|
||||
HVX_Vector v0 = pv_in0 ? pv_in0[d] : Q6_V_vzero();
|
||||
HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero();
|
||||
HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
|
||||
|
||||
for (uint32_t d = 0; d < dk_out / 64; ++d) {
|
||||
__fp16 * out_dtile = out_base + d * HMX_FP16_TILE_N_ELMS * 2;
|
||||
HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
|
||||
HVX_Vector * pv_out1 = pv_out0 + 16;
|
||||
|
||||
const size_t base_lane = (size_t) d * 64;
|
||||
const size_t real_lanes = (base_lane < dk_in) ? hex_smin(64, dk_in - base_lane) : 0;
|
||||
|
||||
if (real_lanes == 0) {
|
||||
*pv_out0 = Q6_V_vzero();
|
||||
*pv_out1 = Q6_V_vzero();
|
||||
continue;
|
||||
}
|
||||
HVX_Vector v0 = pv_in0 ? pv_in0[d] : Q6_V_vzero();
|
||||
HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero();
|
||||
if (real_lanes < 64) {
|
||||
const HVX_VectorPred keep = Q6_Q_vsetq_R((uint32_t) (real_lanes * sizeof(__fp16)));
|
||||
v0 = Q6_V_vmux_QVV(keep, v0, Q6_V_vzero());
|
||||
v1 = Q6_V_vmux_QVV(keep, v1, Q6_V_vzero());
|
||||
}
|
||||
HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
|
||||
*pv_out0 = Q6_V_lo_W(vp);
|
||||
*pv_out1 = Q6_V_hi_W(vp);
|
||||
}
|
||||
|
||||
@@ -165,6 +165,7 @@ int op_get_rows(struct htp_ops_context * octx);
|
||||
int op_cpy(struct htp_ops_context * octx);
|
||||
int op_repeat(struct htp_ops_context * octx);
|
||||
int op_argsort(struct htp_ops_context * octx);
|
||||
int op_top_k(struct htp_ops_context * octx);
|
||||
int op_ssm_conv(struct htp_ops_context * octx);
|
||||
int op_cumsum(struct htp_ops_context * octx);
|
||||
int op_fill(struct htp_ops_context * octx);
|
||||
@@ -175,5 +176,6 @@ int op_gated_delta_net(struct htp_ops_context * octx);
|
||||
int op_pad(struct htp_ops_context * octx);
|
||||
int op_im2col(struct htp_ops_context * octx);
|
||||
int op_allreduce(struct htp_ops_context * octx);
|
||||
int op_roll(struct htp_ops_context * octx);
|
||||
|
||||
#endif /* HTP_CTX_H */
|
||||
|
||||
@@ -71,6 +71,7 @@ enum htp_op_code {
|
||||
HTP_OP_GLU_SWIGLU,
|
||||
HTP_OP_GLU_SWIGLU_OAI,
|
||||
HTP_OP_GLU_GEGLU,
|
||||
HTP_OP_GLU_GEGLU_QUICK,
|
||||
HTP_OP_SOFTMAX,
|
||||
HTP_OP_ADD_ID,
|
||||
HTP_OP_ROPE,
|
||||
@@ -81,6 +82,7 @@ enum htp_op_code {
|
||||
HTP_OP_CPY,
|
||||
HTP_OP_CPY_FENCE,
|
||||
HTP_OP_ARGSORT,
|
||||
HTP_OP_TOP_K,
|
||||
HTP_OP_SQR,
|
||||
HTP_OP_SQRT,
|
||||
HTP_OP_SUM_ROWS,
|
||||
@@ -104,6 +106,7 @@ enum htp_op_code {
|
||||
HTP_OP_ALLREDUCE_ADD,
|
||||
HTP_OP_GLU_SWIGLU_CLAMP,
|
||||
HTP_OP_MDEV_GROUP,
|
||||
HTP_OP_ROLL,
|
||||
|
||||
HTP_OP_INVALID
|
||||
};
|
||||
|
||||
@@ -126,6 +126,7 @@ static inline uint32_t htp_tensor_get_row_size(int type, uint32_t ne00) {
|
||||
case HTP_TYPE_F32: return ne00 * 4;
|
||||
case HTP_TYPE_F16: return ne00 * 2;
|
||||
case HTP_TYPE_Q8_0: return (ne00 / 32) * 34;
|
||||
case HTP_TYPE_I32: return ne00 * 4;
|
||||
default: return 0;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -25,17 +25,20 @@ struct htp_im2col_context {
|
||||
uint32_t npatches; // number of patches assigned to this dev
|
||||
uint32_t npatches_per_thread; // patches = N*OH*OW (pure-DDR kernel)
|
||||
|
||||
uint32_t pe_row_base; // first N*OH row index assigned to this dev (DMA path)
|
||||
uint32_t pe_nrows; // number of N*OH rows assigned to this dev (DMA path)
|
||||
uint32_t pe_rows_per_thread; // N*OH rows per worker
|
||||
uint32_t pe_src_row_bytes; // one output row's source: IC*KH*IW*4, rounded 256
|
||||
uint32_t pe_dst_row_bytes; // one output row's dst: OW*patch_stride*2, rounded 256
|
||||
uint32_t pe_row_base; // first N*OH row index assigned to this dev (DMA path)
|
||||
uint32_t pe_nrows; // number of N*OH rows assigned to this dev (DMA path)
|
||||
uint32_t pe_rows_per_thread; // N*OH rows per worker
|
||||
uint32_t pe_src_row_bytes; // one output row's source: IC*KH*IW*4, rounded 256
|
||||
uint32_t pe_dst_row_bytes; // one output row's dst: OW*patch_stride*2, rounded 256
|
||||
|
||||
// Patch-embed DMA path VTCM ping-pong.
|
||||
uint8_t * pe_vtcm_src; // base of the 2x src buffers region
|
||||
uint8_t * pe_vtcm_dst; // base of the 2x dst buffers region
|
||||
uint32_t pe_src_size_per_thread; // 2 * pe_src_row_bytes
|
||||
uint32_t pe_dst_size_per_thread; // 2 * pe_dst_row_bytes
|
||||
|
||||
uint32_t pe_owb; // output-col block size
|
||||
uint32_t pe_wb; // staged source window width
|
||||
};
|
||||
|
||||
// Per-op VTCM layout for the patch-embed DMA path
|
||||
@@ -59,83 +62,253 @@ static inline void htp_im2col_vtcm_layout_build(struct htp_im2col_vtcm_layout *
|
||||
L->total_bytes = L->off_dst + L->dst_bytes_per_thread * n_threads;
|
||||
}
|
||||
|
||||
#define IM2COL_PATCHEMBED_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG) \
|
||||
static void FNAME(unsigned int nth, unsigned int ith, void * data) { \
|
||||
struct htp_im2col_context * ictx = (struct htp_im2col_context *) data; \
|
||||
struct htp_ops_context * octx = ictx->octx; \
|
||||
struct htp_thread_trace * restrict tr = &octx->ctx->trace[ith]; \
|
||||
const struct htp_tensor * restrict src0 = octx->src[0]; \
|
||||
const struct htp_tensor * restrict src1 = octx->src[1]; \
|
||||
const struct htp_tensor * restrict dst = octx->dst; \
|
||||
const int32_t s0 = octx->op_params[0], s1 = octx->op_params[1]; \
|
||||
const int32_t p0 = octx->op_params[2], p1 = octx->op_params[3]; \
|
||||
const int32_t d0 = octx->op_params[4], d1 = octx->op_params[5]; \
|
||||
const uint32_t N = src1->ne[3], IC = src1->ne[2], IH = src1->ne[1], IW = src1->ne[0]; \
|
||||
const uint32_t KH = src0->ne[1], KW = src0->ne[0]; \
|
||||
const uint32_t OH = dst->ne[2]; \
|
||||
const uint32_t OW = dst->ne[1]; \
|
||||
const uint32_t patch_stride = IC * KH * KW; \
|
||||
const float * restrict src_data = (const float *) src1->data; \
|
||||
DST_CTYPE * restrict dst_data = (DST_CTYPE *) dst->data; \
|
||||
const uint32_t patch_end = ictx->patch_base + ictx->npatches; \
|
||||
const uint32_t patch_start = ictx->patch_base + ictx->npatches_per_thread * ith; \
|
||||
const uint32_t patch_stop = MIN(patch_start + ictx->npatches_per_thread, patch_end);\
|
||||
if (patch_start >= patch_stop) { \
|
||||
return; \
|
||||
} \
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, patch_start); \
|
||||
for (uint32_t p = patch_start; p < patch_stop; p++) { \
|
||||
const uint32_t iow = p % OW; \
|
||||
const uint32_t ioh = (p / OW) % OH; \
|
||||
const uint32_t in = p / (OW * OH); \
|
||||
DST_CTYPE * restrict dst_patch = dst_data + (uint64_t) p * patch_stride; \
|
||||
for (uint32_t iic = 0; iic < IC; iic++) { \
|
||||
const float * restrict src_plane = src_data + ((uint64_t) in * IC + iic) * IH * IW; \
|
||||
for (uint32_t ikh = 0; ikh < KH; ikh++) { \
|
||||
const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1; \
|
||||
DST_CTYPE * restrict out_run = dst_patch + iic * (KH * KW) + ikh * KW; \
|
||||
if (iih < 0 || iih >= (int32_t) IH) { \
|
||||
SPLAT_FN(out_run, 0.0f, KW); \
|
||||
continue; \
|
||||
} \
|
||||
const int32_t iiw0 = (int32_t) iow * s0 - p0; \
|
||||
const float * restrict src_run = src_plane + (uint64_t) iih * IW + iiw0; \
|
||||
if (d0 == 1) { \
|
||||
/* contiguous source run: [lo,hi) is in-bounds, tails are zero pad */ \
|
||||
const int32_t lo = iiw0 < 0 ? -iiw0 : 0; \
|
||||
int32_t hi = (int32_t) IW - iiw0; \
|
||||
if (hi > (int32_t) KW) { \
|
||||
hi = (int32_t) KW; \
|
||||
} \
|
||||
if (hi <= lo) { \
|
||||
SPLAT_FN(out_run, 0.0f, KW); \
|
||||
} else { \
|
||||
if (lo > 0) { \
|
||||
SPLAT_FN(out_run, 0.0f, (uint32_t) lo); \
|
||||
} \
|
||||
COPY_FN((uint8_t *) (out_run + lo), (const uint8_t *) (src_run + lo), \
|
||||
(uint32_t) (hi - lo)); \
|
||||
if (hi < (int32_t) KW) { \
|
||||
SPLAT_FN(out_run + hi, 0.0f, (KW - (uint32_t) hi)); \
|
||||
} \
|
||||
} \
|
||||
continue; \
|
||||
} \
|
||||
for (uint32_t ikw = 0; ikw < KW; ikw++) { \
|
||||
const int32_t iiw = (int32_t) iow * s0 + (int32_t) ikw * d0 - p0; \
|
||||
out_run[ikw] = (iiw < 0 || iiw >= (int32_t) IW) ? \
|
||||
(DST_CTYPE) 0.0f : \
|
||||
(DST_CTYPE) src_plane[(uint64_t) iih * IW + iiw]; \
|
||||
} \
|
||||
} \
|
||||
} \
|
||||
} \
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, patch_start); \
|
||||
#define IM2COL_PATCHEMBED_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG) \
|
||||
static void FNAME(unsigned int nth, unsigned int ith, void * data) { \
|
||||
struct htp_im2col_context * ictx = (struct htp_im2col_context *) data; \
|
||||
struct htp_ops_context * octx = ictx->octx; \
|
||||
struct htp_thread_trace * restrict tr = &octx->ctx->trace[ith]; \
|
||||
const struct htp_tensor * restrict src0 = octx->src[0]; \
|
||||
const struct htp_tensor * restrict src1 = octx->src[1]; \
|
||||
const struct htp_tensor * restrict dst = octx->dst; \
|
||||
const int32_t s0 = octx->op_params[0]; \
|
||||
const int32_t s1 = octx->op_params[1]; \
|
||||
const int32_t p0 = octx->op_params[2]; \
|
||||
const int32_t p1 = octx->op_params[3]; \
|
||||
const int32_t d0 = octx->op_params[4]; \
|
||||
const int32_t d1 = octx->op_params[5]; \
|
||||
const int32_t is_2D = octx->op_params[6] == 1; \
|
||||
const uint32_t N = is_2D ? src1->ne[3] : src1->ne[2]; \
|
||||
const uint32_t IC = is_2D ? src1->ne[2] : src1->ne[1]; \
|
||||
const uint32_t IH = is_2D ? src1->ne[1] : 1; \
|
||||
const uint32_t IW = src1->ne[0]; \
|
||||
const uint32_t KH = is_2D ? src0->ne[1] : 1; \
|
||||
const uint32_t KW = src0->ne[0]; \
|
||||
const uint32_t OH = is_2D ? dst->ne[2] : 1; \
|
||||
const uint32_t OW = dst->ne[1]; \
|
||||
const uint32_t patch_stride = IC * KH * KW; \
|
||||
const float * restrict src_data = (const float *) src1->data; \
|
||||
DST_CTYPE * restrict dst_data = (DST_CTYPE *) dst->data; \
|
||||
const uint32_t patch_end = ictx->patch_base + ictx->npatches; \
|
||||
const uint32_t patch_start = ictx->patch_base + ictx->npatches_per_thread * ith; \
|
||||
const uint32_t patch_stop = MIN(patch_start + ictx->npatches_per_thread, patch_end); \
|
||||
if (patch_start >= patch_stop) { \
|
||||
return; \
|
||||
} \
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, patch_start); \
|
||||
for (uint32_t p = patch_start; p < patch_stop; p++) { \
|
||||
const uint32_t iow = p % OW; \
|
||||
const uint32_t ioh = (p / OW) % OH; \
|
||||
const uint32_t in = p / (OW * OH); \
|
||||
DST_CTYPE * restrict dst_patch = dst_data + (uint64_t) p * patch_stride; \
|
||||
for (uint32_t iic = 0; iic < IC; iic++) { \
|
||||
const float * restrict src_plane = src_data + ((uint64_t) in * IC + iic) * IH * IW; \
|
||||
for (uint32_t ikh = 0; ikh < KH; ikh++) { \
|
||||
const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1; \
|
||||
DST_CTYPE * restrict out_run = dst_patch + iic * (KH * KW) + ikh * KW; \
|
||||
if (iih < 0 || iih >= (int32_t) IH) { \
|
||||
SPLAT_FN(out_run, 0.0f, KW); \
|
||||
continue; \
|
||||
} \
|
||||
const int32_t iiw0 = (int32_t) iow * s0 - p0; \
|
||||
const float * restrict src_run = src_plane + (uint64_t) iih * IW + iiw0; \
|
||||
if (d0 == 1) { \
|
||||
/* contiguous source run: [lo,hi) is in-bounds, tails are zero pad */ \
|
||||
const int32_t lo = iiw0 < 0 ? -iiw0 : 0; \
|
||||
int32_t hi = (int32_t) IW - iiw0; \
|
||||
if (hi > (int32_t) KW) { \
|
||||
hi = (int32_t) KW; \
|
||||
} \
|
||||
if (hi <= lo) { \
|
||||
SPLAT_FN(out_run, 0.0f, KW); \
|
||||
} else { \
|
||||
if (lo > 0) { \
|
||||
SPLAT_FN(out_run, 0.0f, (uint32_t) lo); \
|
||||
} \
|
||||
COPY_FN((uint8_t *) (out_run + lo), (const uint8_t *) (src_run + lo), \
|
||||
(uint32_t) (hi - lo)); \
|
||||
if (hi < (int32_t) KW) { \
|
||||
SPLAT_FN(out_run + hi, 0.0f, (KW - (uint32_t) hi)); \
|
||||
} \
|
||||
} \
|
||||
continue; \
|
||||
} \
|
||||
for (uint32_t ikw = 0; ikw < KW; ikw++) { \
|
||||
const int32_t iiw = (int32_t) iow * s0 + (int32_t) ikw * d0 - p0; \
|
||||
out_run[ikw] = (iiw < 0 || iiw >= (int32_t) IW) ? \
|
||||
(DST_CTYPE) 0.0f : \
|
||||
(DST_CTYPE) src_plane[(uint64_t) iih * IW + iiw]; \
|
||||
} \
|
||||
} \
|
||||
} \
|
||||
} \
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, patch_start); \
|
||||
}
|
||||
|
||||
IM2COL_PATCHEMBED_BODY(im2col_patchembed_thread, __fp16, hvx_copy_f16_f32_uu, hvx_splat_f16_u, sizeof(__fp16), "f32-f16")
|
||||
IM2COL_PATCHEMBED_BODY(im2col_patchembed_f32_thread, float, hvx_copy_f32_uu, hvx_splat_f32_u, sizeof(float), "f32-f32")
|
||||
|
||||
// Software-pipelined 2-deep: while HVX computes block bi from buffer slot
|
||||
// (bi&1), the DMA engine stages block bi+1 into the other slot concurrently.
|
||||
// A single dma_queue_flush per iteration (after issuing the next stage-in and
|
||||
// this block's store-out) waits for both - safe because the ring is strict
|
||||
// FIFO and each buffer slot is only reused after its prior consumer (compute
|
||||
// or store-out) already finished in program order.
|
||||
#define IM2COL_BLOCKED_DMA_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG) \
|
||||
static void FNAME(unsigned int nth, unsigned int ith, void * data) { \
|
||||
struct htp_im2col_context * ictx = (struct htp_im2col_context *) data; \
|
||||
struct htp_ops_context * octx = ictx->octx; \
|
||||
struct htp_thread_trace * restrict tr = &octx->ctx->trace[ith]; \
|
||||
const struct htp_tensor * restrict src1 = octx->src[1]; \
|
||||
const struct htp_tensor * restrict dst = octx->dst; \
|
||||
const int32_t s0 = octx->op_params[0], s1 = octx->op_params[1]; \
|
||||
const int32_t p0 = octx->op_params[2], p1 = octx->op_params[3]; \
|
||||
const int32_t d0 = octx->op_params[4], d1 = octx->op_params[5]; \
|
||||
const int32_t is_2D = octx->op_params[6] == 1; \
|
||||
const uint32_t N = is_2D ? src1->ne[3] : src1->ne[2]; \
|
||||
const uint32_t IC = is_2D ? src1->ne[2] : src1->ne[1]; \
|
||||
const uint32_t IH = is_2D ? src1->ne[1] : 1; \
|
||||
const uint32_t IW = src1->ne[0]; \
|
||||
const uint32_t KH = is_2D ? octx->src[0]->ne[1] : 1; \
|
||||
const uint32_t KW = octx->src[0]->ne[0]; \
|
||||
const uint32_t OH = is_2D ? dst->ne[2] : 1; \
|
||||
const uint32_t OW = dst->ne[1]; \
|
||||
const uint32_t owb = ictx->pe_owb, Wb = ictx->pe_wb; \
|
||||
const uint32_t patch_stride = IC * KH * KW; \
|
||||
const float * restrict src_data = (const float *) src1->data; \
|
||||
DST_CTYPE * restrict dst_data = (DST_CTYPE *) dst->data; \
|
||||
dma_queue * dmaq = octx->ctx->dma[ith]; \
|
||||
uint8_t * srcb_base = ictx->pe_vtcm_src + ith * ictx->pe_src_size_per_thread; \
|
||||
uint8_t * dstb_base = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread; \
|
||||
float * srcb2[2] = { (float *) srcb_base, (float *) (srcb_base + ictx->pe_src_row_bytes) }; \
|
||||
DST_CTYPE * dstb2[2] = { (DST_CTYPE *) dstb_base, (DST_CTYPE *) (dstb_base + ictx->pe_dst_row_bytes) }; \
|
||||
const uint32_t nrows = N * OH; \
|
||||
const uint32_t per_thread = ictx->pe_rows_per_thread; \
|
||||
const uint32_t row_start = per_thread * ith; \
|
||||
const uint32_t row_end = MIN(row_start + per_thread, nrows); \
|
||||
if (row_start >= row_end) \
|
||||
return; \
|
||||
const uint32_t nbpr = (OW + owb - 1) / owb; \
|
||||
const uint32_t nrows_local = row_end - row_start; \
|
||||
const uint32_t total_blocks = nrows_local * nbpr; \
|
||||
for (uint32_t bi = 0; bi < total_blocks; bi++) { \
|
||||
const uint32_t buf = bi & 1u; \
|
||||
float * srcb = srcb2[buf]; \
|
||||
DST_CTYPE * dstb = dstb2[buf]; \
|
||||
const uint32_t r = row_start + bi / nbpr; \
|
||||
const uint32_t in = r / OH; \
|
||||
const uint32_t ioh = r % OH; \
|
||||
const uint32_t c0 = (bi % nbpr) * owb; \
|
||||
const uint32_t nb = MIN(owb, OW - c0); \
|
||||
const int32_t win0 = (int32_t) c0 * s0 - p0; \
|
||||
if (bi == 0) { \
|
||||
/* prologue: stage block 0 and wait - nothing to overlap with yet */ \
|
||||
for (uint32_t ikh = 0; ikh < KH; ikh++) { \
|
||||
const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1; \
|
||||
if (iih < 0 || iih >= (int32_t) IH) \
|
||||
continue; \
|
||||
const int32_t lo = win0 < 0 ? -win0 : 0; \
|
||||
int32_t hi = (int32_t) IW - win0; \
|
||||
if (hi > (int32_t) Wb) \
|
||||
hi = (int32_t) Wb; \
|
||||
if (hi <= lo) \
|
||||
continue; \
|
||||
const uint32_t cpw = (uint32_t) (hi - lo); \
|
||||
float * vdst = srcb + (uint64_t) ikh * Wb + (uint32_t) lo; \
|
||||
const float * vsrc = src_data + ((uint64_t) (in * IC) * IH + iih) * IW + (win0 + lo); \
|
||||
while (!dma_queue_push(dmaq, dma_make_ptr((uint8_t *) vdst, (const uint8_t *) vsrc), \
|
||||
(size_t) KH * Wb * sizeof(float), (size_t) IH * IW * sizeof(float), \
|
||||
cpw * sizeof(float), IC)) { \
|
||||
dma_queue_pop(dmaq); \
|
||||
} \
|
||||
} \
|
||||
dma_queue_flush(dmaq); \
|
||||
} \
|
||||
if (bi + 1 < total_blocks) { \
|
||||
/* prefetch: stage block bi+1 into the other slot; overlaps with this block's compute below */ \
|
||||
const uint32_t nbuf = 1u - buf; \
|
||||
float * nsrcb = srcb2[nbuf]; \
|
||||
const uint32_t nr = row_start + (bi + 1) / nbpr; \
|
||||
const uint32_t nin = nr / OH; \
|
||||
const uint32_t nioh = nr % OH; \
|
||||
const uint32_t nc0 = ((bi + 1) % nbpr) * owb; \
|
||||
const int32_t nwin0 = (int32_t) nc0 * s0 - p0; \
|
||||
for (uint32_t ikh = 0; ikh < KH; ikh++) { \
|
||||
const int32_t iih = (int32_t) nioh * s1 + (int32_t) ikh * d1 - p1; \
|
||||
if (iih < 0 || iih >= (int32_t) IH) \
|
||||
continue; \
|
||||
const int32_t lo = nwin0 < 0 ? -nwin0 : 0; \
|
||||
int32_t hi = (int32_t) IW - nwin0; \
|
||||
if (hi > (int32_t) Wb) \
|
||||
hi = (int32_t) Wb; \
|
||||
if (hi <= lo) \
|
||||
continue; \
|
||||
const uint32_t cpw = (uint32_t) (hi - lo); \
|
||||
float * vdst = nsrcb + (uint64_t) ikh * Wb + (uint32_t) lo; \
|
||||
const float * vsrc = src_data + ((uint64_t) (nin * IC) * IH + iih) * IW + (nwin0 + lo); \
|
||||
while (!dma_queue_push(dmaq, dma_make_ptr((uint8_t *) vdst, (const uint8_t *) vsrc), \
|
||||
(size_t) KH * Wb * sizeof(float), (size_t) IH * IW * sizeof(float), \
|
||||
cpw * sizeof(float), IC)) { \
|
||||
dma_queue_pop(dmaq); \
|
||||
} \
|
||||
} \
|
||||
} \
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, r); \
|
||||
for (uint32_t j = 0; j < nb; j++) { \
|
||||
const uint32_t iow = c0 + j; \
|
||||
DST_CTYPE * dst_patch = dstb + (uint64_t) j * patch_stride; \
|
||||
const int32_t iiw0 = (int32_t) iow * s0 - p0; \
|
||||
for (uint32_t ikh = 0; ikh < KH; ikh++) { \
|
||||
const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1; \
|
||||
const int okh = (iih >= 0 && iih < (int32_t) IH); \
|
||||
for (uint32_t iic = 0; iic < IC; iic++) { \
|
||||
DST_CTYPE * out_run = dst_patch + iic * (KH * KW) + ikh * KW; \
|
||||
if (!okh) { \
|
||||
SPLAT_FN(out_run, 0.0f, KW); \
|
||||
continue; \
|
||||
} \
|
||||
const float * vrow = srcb + ((uint64_t) (iic * KH + ikh)) * Wb; /* col win0 at idx 0*/ \
|
||||
if (d0 == 1) { \
|
||||
/* contiguous run within the staged window: [lo,hi) in-bounds, tails zero pad */ \
|
||||
const int32_t lo = iiw0 < 0 ? -iiw0 : 0; \
|
||||
int32_t hi = (int32_t) IW - iiw0; \
|
||||
if (hi > (int32_t) KW) { \
|
||||
hi = (int32_t) KW; \
|
||||
} \
|
||||
if (hi <= lo) { \
|
||||
SPLAT_FN(out_run, 0.0f, KW); \
|
||||
} else { \
|
||||
if (lo > 0) { \
|
||||
SPLAT_FN(out_run, 0.0f, (uint32_t) lo); \
|
||||
} \
|
||||
COPY_FN((uint8_t *) (out_run + lo), (const uint8_t *) (vrow + (iiw0 + lo - win0)), \
|
||||
(uint32_t) (hi - lo)); \
|
||||
if (hi < (int32_t) KW) { \
|
||||
SPLAT_FN(out_run + hi, 0.0f, (KW - (uint32_t) hi)); \
|
||||
} \
|
||||
} \
|
||||
continue; \
|
||||
} \
|
||||
for (uint32_t ikw = 0; ikw < KW; ikw++) { \
|
||||
const int32_t iiw = iiw0 + (int32_t) ikw * d0; \
|
||||
out_run[ikw] = \
|
||||
(iiw < 0 || iiw >= (int32_t) IW) ? (DST_CTYPE) 0.0f : (DST_CTYPE) vrow[iiw - win0]; \
|
||||
} \
|
||||
} \
|
||||
} \
|
||||
} \
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, r); \
|
||||
DST_CTYPE * ddr = dst_data + ((uint64_t) (in * OH + ioh) * OW + c0) * patch_stride; \
|
||||
dma_queue_push_vtcm_to_ddr(dmaq, dma_make_ptr((uint8_t *) ddr, (uint8_t *) dstb), \
|
||||
nb * patch_stride * (DST_ELEM), nb * patch_stride * (DST_ELEM), 1); \
|
||||
dma_queue_flush(dmaq); \
|
||||
} \
|
||||
}
|
||||
IM2COL_BLOCKED_DMA_BODY(im2col_blocked_dma_thread, __fp16, hvx_copy_f16_f32_uu, hvx_splat_f16_u, sizeof(__fp16), "blk-dma-f16")
|
||||
IM2COL_BLOCKED_DMA_BODY(im2col_blocked_dma_f32_thread, float, hvx_copy_f32_uu, hvx_splat_f32_u, sizeof(float), "blk-dma-f32")
|
||||
|
||||
// Exact-tiling patch-embed DMA fast path (s0==KW, p0=0, d0=1; and 2D s1==KH,
|
||||
// p1=0, d1=1). Intentionally reads no stride/pad/dilation params so the inner
|
||||
// copy stays tight and fully hoisted - do NOT graft the general gather in here.
|
||||
#define IM2COL_PATCHEMBED_DMA_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG) \
|
||||
static void FNAME(unsigned int nth, unsigned int ith, void * data) { \
|
||||
struct htp_im2col_context * ictx = (struct htp_im2col_context *) data; \
|
||||
@@ -143,21 +316,27 @@ IM2COL_PATCHEMBED_BODY(im2col_patchembed_f32_thread, float, hvx_copy_f32_uu, hvx
|
||||
struct htp_thread_trace * restrict tr = &octx->ctx->trace[ith]; \
|
||||
const struct htp_tensor * restrict src1 = octx->src[1]; \
|
||||
const struct htp_tensor * restrict dst = octx->dst; \
|
||||
const uint32_t N = src1->ne[3], IC = src1->ne[2], IH = src1->ne[1], IW = src1->ne[0]; \
|
||||
const uint32_t KH = octx->src[0]->ne[1], KW = octx->src[0]->ne[0]; \
|
||||
const uint32_t OH = dst->ne[2], OW = dst->ne[1]; \
|
||||
const uint32_t patch_stride = IC * KH * KW; \
|
||||
const float * restrict src_data = (const float *) src1->data; \
|
||||
DST_CTYPE * restrict dst_data = (DST_CTYPE *) dst->data; \
|
||||
dma_queue * dmaq = octx->ctx->dma[ith]; \
|
||||
uint8_t * src_base = ictx->pe_vtcm_src + ith * ictx->pe_src_size_per_thread; \
|
||||
uint8_t * dst_base = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread; \
|
||||
float * srcb = (float *) src_base; \
|
||||
DST_CTYPE * dstb = (DST_CTYPE *) dst_base; \
|
||||
const uint32_t row_end_max = ictx->pe_row_base + ictx->pe_nrows; \
|
||||
const uint32_t per_thread = ictx->pe_rows_per_thread; \
|
||||
const uint32_t row_start = ictx->pe_row_base + per_thread * ith; \
|
||||
const uint32_t row_end = MIN(row_start + per_thread, row_end_max); \
|
||||
const int32_t is_2D = octx->op_params[6] == 1; \
|
||||
const uint32_t N = is_2D ? src1->ne[3] : src1->ne[2]; \
|
||||
const uint32_t IC = is_2D ? src1->ne[2] : src1->ne[1]; \
|
||||
const uint32_t IH = is_2D ? src1->ne[1] : 1; \
|
||||
const uint32_t IW = src1->ne[0]; \
|
||||
const uint32_t KH = is_2D ? octx->src[0]->ne[1] : 1; \
|
||||
const uint32_t KW = octx->src[0]->ne[0]; \
|
||||
const uint32_t OH = is_2D ? dst->ne[2] : 1; \
|
||||
const uint32_t OW = dst->ne[1]; \
|
||||
const uint32_t patch_stride = IC * KH * KW; \
|
||||
const float * restrict src_data = (const float *) src1->data; \
|
||||
DST_CTYPE * restrict dst_data = (DST_CTYPE *) dst->data; \
|
||||
dma_queue * dmaq = octx->ctx->dma[ith]; \
|
||||
uint8_t * src_base = ictx->pe_vtcm_src + ith * ictx->pe_src_size_per_thread; \
|
||||
uint8_t * dst_base = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread; \
|
||||
float * srcb = (float *) src_base; \
|
||||
DST_CTYPE * dstb = (DST_CTYPE *) dst_base; \
|
||||
const uint32_t row_end_max = ictx->pe_row_base + ictx->pe_nrows; \
|
||||
const uint32_t per_thread = ictx->pe_rows_per_thread; \
|
||||
const uint32_t row_start = ictx->pe_row_base + per_thread * ith; \
|
||||
const uint32_t row_end = MIN(row_start + per_thread, row_end_max); \
|
||||
if (row_start >= row_end) \
|
||||
return; \
|
||||
for (uint32_t r = row_start; r < row_end; r++) { \
|
||||
@@ -209,21 +388,30 @@ static bool im2col_use_patchembed_dma(const struct htp_ops_context * octx) {
|
||||
const int32_t p0 = octx->op_params[2], p1 = octx->op_params[3];
|
||||
const int32_t d0 = octx->op_params[4], d1 = octx->op_params[5];
|
||||
const int is_2D = octx->op_params[6] == 1;
|
||||
if (!is_2D) {
|
||||
return false;
|
||||
}
|
||||
if (octx->dst->type != HTP_TYPE_F16 && octx->dst->type != HTP_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
const uint32_t KH = octx->src[0]->ne[1], KW = octx->src[0]->ne[0];
|
||||
if (s0 != (int32_t) KW || s1 != (int32_t) KH) {
|
||||
return false; // non-overlapping
|
||||
const uint32_t KH = is_2D ? octx->src[0]->ne[1] : 1;
|
||||
const uint32_t KW = octx->src[0]->ne[0];
|
||||
if (s0 != (int32_t) KW) {
|
||||
return false; // non-overlapping (width)
|
||||
}
|
||||
if (p0 != 0 || p1 != 0) {
|
||||
return false; // no padding
|
||||
if (p0 != 0) {
|
||||
return false; // no padding (width)
|
||||
}
|
||||
if (d0 != 1 || d1 != 1) {
|
||||
return false; // no dilation
|
||||
if (d0 != 1) {
|
||||
return false; // no dilation (width)
|
||||
}
|
||||
if (is_2D) {
|
||||
if (s1 != (int32_t) KH) {
|
||||
return false; // non-overlapping (height)
|
||||
}
|
||||
if (p1 != 0) {
|
||||
return false; // no padding (height)
|
||||
}
|
||||
if (d1 != 1) {
|
||||
return false; // no dilation (height)
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
@@ -233,8 +421,11 @@ static bool im2col_use_patchembed_dma(const struct htp_ops_context * octx) {
|
||||
static bool im2col_patchembed_dma_fits(struct htp_ops_context * octx,
|
||||
struct htp_im2col_context * ictx,
|
||||
uint32_t n_threads) {
|
||||
const uint32_t IC = octx->src[1]->ne[2], IW = octx->src[1]->ne[0];
|
||||
const uint32_t KH = octx->src[0]->ne[1], KW = octx->src[0]->ne[0];
|
||||
const int32_t is_2D = octx->op_params[6] == 1;
|
||||
const uint32_t IC = is_2D ? octx->src[1]->ne[2] : octx->src[1]->ne[1];
|
||||
const uint32_t IW = octx->src[1]->ne[0];
|
||||
const uint32_t KH = is_2D ? octx->src[0]->ne[1] : 1;
|
||||
const uint32_t KW = octx->src[0]->ne[0];
|
||||
const uint32_t OW = octx->dst->ne[1];
|
||||
const uint32_t patch_stride = IC * KH * KW;
|
||||
|
||||
@@ -257,6 +448,45 @@ static bool im2col_patchembed_dma_fits(struct htp_ops_context * octx,
|
||||
return true;
|
||||
}
|
||||
|
||||
// Sizes a per-thread 2x(src,dst) VTCM ping-pong for the blocked general kernel.
|
||||
// Stages Wb=(owb-1)*s0+(KW-1)*d0+1 source cols per (iic,ikh) row and owb patches
|
||||
// of dst. Picks the largest owb that fits; returns false if even owb=1 does not.
|
||||
static bool im2col_blocked_dma_fits(struct htp_ops_context * octx,
|
||||
struct htp_im2col_context * ictx,
|
||||
uint32_t n_threads) {
|
||||
const int32_t is_2D = octx->op_params[6] == 1;
|
||||
const int32_t s0 = octx->op_params[0];
|
||||
const int32_t d0 = octx->op_params[4];
|
||||
const uint32_t IC = is_2D ? octx->src[1]->ne[2] : octx->src[1]->ne[1];
|
||||
const uint32_t KH = is_2D ? octx->src[0]->ne[1] : 1;
|
||||
const uint32_t KW = octx->src[0]->ne[0];
|
||||
const uint32_t OW = octx->dst->ne[1];
|
||||
const uint32_t patch_stride = IC * KH * KW;
|
||||
const uint32_t dst_elem = (octx->dst->type == HTP_TYPE_F16) ? sizeof(__fp16) : sizeof(float);
|
||||
|
||||
for (uint32_t owb = (OW < 256 ? OW : 256); owb >= 1; owb--) {
|
||||
const uint32_t Wb = (owb - 1) * (uint32_t) s0 + (KW - 1) * (uint32_t) d0 + 1;
|
||||
const uint32_t src_row_bytes = hex_round_up(IC * KH * Wb * sizeof(float), 256);
|
||||
const uint32_t dst_row_bytes = hex_round_up(owb * patch_stride * dst_elem, 256);
|
||||
struct htp_im2col_vtcm_layout L;
|
||||
htp_im2col_vtcm_layout_build(&L, src_row_bytes, dst_row_bytes, n_threads);
|
||||
if (L.total_bytes <= octx->ctx->vtcm_size) {
|
||||
uint8_t * const base = octx->ctx->vtcm_base;
|
||||
ictx->pe_owb = owb;
|
||||
ictx->pe_wb = Wb;
|
||||
ictx->pe_src_row_bytes = src_row_bytes;
|
||||
ictx->pe_dst_row_bytes = dst_row_bytes;
|
||||
ictx->pe_vtcm_src = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src);
|
||||
ictx->pe_vtcm_dst = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst);
|
||||
ictx->pe_src_size_per_thread = (uint32_t) L.src_bytes_per_thread;
|
||||
ictx->pe_dst_size_per_thread = (uint32_t) L.dst_bytes_per_thread;
|
||||
return true;
|
||||
}
|
||||
if (owb == 1) break; // avoid unsigned underflow
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
int op_im2col(struct htp_ops_context * octx) {
|
||||
const struct htp_tensor * src1 = octx->src[1];
|
||||
const struct htp_tensor * dst = octx->dst;
|
||||
@@ -270,8 +500,9 @@ int op_im2col(struct htp_ops_context * octx) {
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
const uint32_t N = src1->ne[3];
|
||||
const uint32_t OH = dst->ne[2];
|
||||
const int32_t is_2D = octx->op_params[6] == 1;
|
||||
const uint32_t N = is_2D ? src1->ne[3] : src1->ne[2];
|
||||
const uint32_t OH = is_2D ? dst->ne[2] : 1;
|
||||
const uint32_t OW = dst->ne[1];
|
||||
const uint32_t total_patches = N * OH * OW;
|
||||
const uint32_t total_rows = N * OH;
|
||||
@@ -280,8 +511,11 @@ int op_im2col(struct htp_ops_context * octx) {
|
||||
uint32_t npatches = total_patches;
|
||||
if (octx->ctx->mdev.count > 1) {
|
||||
const uint32_t patch_size = dst->nb[1];
|
||||
const uint32_t patches_per_chunk = (patch_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(patch_size, HEX_L2_LINE_SIZE)) : 1;
|
||||
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_patches, htp_tensor_mdev_data_aligned(dst) ? patches_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
|
||||
const uint32_t patches_per_chunk =
|
||||
(patch_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(patch_size, HEX_L2_LINE_SIZE)) : 1;
|
||||
const struct htp_tensor_mdev_range range =
|
||||
htp_tensor_mdev_partition(total_patches, htp_tensor_mdev_data_aligned(dst) ? patches_per_chunk : 0,
|
||||
octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
|
||||
patch_base = range.start;
|
||||
npatches = range.count;
|
||||
}
|
||||
@@ -290,8 +524,11 @@ int op_im2col(struct htp_ops_context * octx) {
|
||||
uint32_t nrows = total_rows;
|
||||
if (octx->ctx->mdev.count > 1) {
|
||||
const uint32_t row_size = dst->nb[2];
|
||||
const uint32_t rows_per_chunk = (row_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(row_size, HEX_L2_LINE_SIZE)) : 1;
|
||||
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, htp_tensor_mdev_data_aligned(dst) ? rows_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
|
||||
const uint32_t rows_per_chunk =
|
||||
(row_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(row_size, HEX_L2_LINE_SIZE)) : 1;
|
||||
const struct htp_tensor_mdev_range range =
|
||||
htp_tensor_mdev_partition(total_rows, htp_tensor_mdev_data_aligned(dst) ? rows_per_chunk : 0,
|
||||
octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div);
|
||||
row_base = range.start;
|
||||
nrows = range.count;
|
||||
}
|
||||
@@ -309,22 +546,32 @@ int op_im2col(struct htp_ops_context * octx) {
|
||||
ictx.npatches_per_thread = (npatches + n_threads - 1) / n_threads;
|
||||
|
||||
// Clean non-overlapping patch-embed -> DMA kernel (if it fits VTCM);
|
||||
// everything else (padding/dilation/stride edges) -> pure-DDR kernel.
|
||||
if (im2col_use_patchembed_dma(octx) && nrows > 0) {
|
||||
// everything else (padding/dilation/stride edges) -> blocked-staging DMA
|
||||
// kernel; if neither fits VTCM -> pure-DDR kernel.
|
||||
if (nrows > 0) {
|
||||
const uint32_t pth = MIN(octx->n_threads, nrows);
|
||||
if (pth > 0 && im2col_patchembed_dma_fits(octx, &ictx, pth)) {
|
||||
ictx.pe_row_base = row_base;
|
||||
ictx.pe_nrows = nrows;
|
||||
ictx.pe_rows_per_thread = (nrows + pth - 1) / pth;
|
||||
if (dst->type == HTP_TYPE_F16) {
|
||||
work_queue_run(octx->ctx->work_queue, im2col_patchembed_dma_thread, &ictx, pth);
|
||||
} else {
|
||||
work_queue_run(octx->ctx->work_queue, im2col_patchembed_dma_f32_thread, &ictx, pth);
|
||||
if (pth > 0) {
|
||||
ictx.pe_row_base = row_base;
|
||||
ictx.pe_nrows = nrows;
|
||||
const bool exact = im2col_use_patchembed_dma(octx);
|
||||
if (exact && im2col_patchembed_dma_fits(octx, &ictx, pth)) {
|
||||
ictx.pe_rows_per_thread = (nrows + pth - 1) / pth;
|
||||
work_queue_run(octx->ctx->work_queue,
|
||||
dst->type == HTP_TYPE_F16 ? im2col_patchembed_dma_thread
|
||||
: im2col_patchembed_dma_f32_thread, &ictx, pth);
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
if (!exact && im2col_blocked_dma_fits(octx, &ictx, pth)) {
|
||||
ictx.pe_rows_per_thread = (nrows + pth - 1) / pth;
|
||||
work_queue_run(octx->ctx->work_queue,
|
||||
dst->type == HTP_TYPE_F16 ? im2col_blocked_dma_thread
|
||||
: im2col_blocked_dma_f32_thread, &ictx, pth);
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
// else: doesn't fit -> fall through to the pure-DDR kernel below.
|
||||
}
|
||||
// Fall through to pure-DDR.
|
||||
|
||||
|
||||
if (npatches == 0) {
|
||||
return HTP_STATUS_OK;
|
||||
|
||||
@@ -828,6 +828,7 @@ static int execute_op(struct htp_ops_context * octx) {
|
||||
case HTP_OP_GLU_SWIGLU_OAI:
|
||||
case HTP_OP_GLU_SWIGLU_CLAMP:
|
||||
case HTP_OP_GLU_GEGLU:
|
||||
case HTP_OP_GLU_GEGLU_QUICK:
|
||||
return op_activations(octx);
|
||||
|
||||
case HTP_OP_SOFTMAX:
|
||||
@@ -858,6 +859,9 @@ static int execute_op(struct htp_ops_context * octx) {
|
||||
case HTP_OP_ARGSORT:
|
||||
return op_argsort(octx);
|
||||
|
||||
case HTP_OP_TOP_K:
|
||||
return op_top_k(octx);
|
||||
|
||||
case HTP_OP_SSM_CONV:
|
||||
return op_ssm_conv(octx);
|
||||
|
||||
@@ -879,6 +883,9 @@ static int execute_op(struct htp_ops_context * octx) {
|
||||
case HTP_OP_IM2COL:
|
||||
return op_im2col(octx);
|
||||
|
||||
case HTP_OP_ROLL:
|
||||
return op_roll(octx);
|
||||
|
||||
case HTP_OP_CONCAT:
|
||||
return op_concat(octx);
|
||||
|
||||
|
||||
@@ -0,0 +1,316 @@
|
||||
#pragma clang diagnostic ignored "-Wunused-variable"
|
||||
#pragma clang diagnostic ignored "-Wunused-function"
|
||||
#pragma clang diagnostic ignored "-Wunused-but-set-variable"
|
||||
|
||||
#include <HAP_farf.h>
|
||||
#include <HAP_perf.h>
|
||||
|
||||
#include <string.h>
|
||||
|
||||
#include "dma-queue.h"
|
||||
#include "hvx-utils.h"
|
||||
|
||||
#define GGML_COMMON_DECL_C
|
||||
#include "ggml-common.h"
|
||||
#include "htp-ctx.h"
|
||||
#include "hex-common.h"
|
||||
#include "hex-profile.h"
|
||||
#include "htp-ops.h"
|
||||
#include "htp-tensor.h"
|
||||
|
||||
struct htp_roll_context {
|
||||
struct htp_ops_context * octx;
|
||||
|
||||
uint32_t row_start;
|
||||
uint32_t nrows;
|
||||
uint32_t nrows_per_thread;
|
||||
|
||||
struct fastdiv_values div_ne1;
|
||||
struct fastdiv_values div_ne2_ne1;
|
||||
};
|
||||
|
||||
static inline uint32_t htp_roll_wrap(int32_t i, uint32_t ne) {
|
||||
if (i < 0) {
|
||||
return (uint32_t) (i + (int32_t) ne);
|
||||
}
|
||||
if ((uint32_t) i >= ne) {
|
||||
return (uint32_t) i - ne;
|
||||
}
|
||||
return (uint32_t) i;
|
||||
}
|
||||
|
||||
#define htp_roll_preamble \
|
||||
const struct htp_tensor * src0 = octx->src[0]; \
|
||||
const struct htp_tensor * dst = octx->dst; \
|
||||
\
|
||||
const uint32_t ne0 = dst->ne[0]; \
|
||||
const uint32_t ne1 = dst->ne[1]; \
|
||||
const uint32_t ne2 = dst->ne[2]; \
|
||||
const uint32_t ne3 = dst->ne[3]; \
|
||||
\
|
||||
const uint32_t nb01 = src0->nb[1]; \
|
||||
const uint32_t nb02 = src0->nb[2]; \
|
||||
const uint32_t nb03 = src0->nb[3]; \
|
||||
\
|
||||
const uint32_t nb1 = dst->nb[1]; \
|
||||
const uint32_t nb2 = dst->nb[2]; \
|
||||
const uint32_t nb3 = dst->nb[3]; \
|
||||
\
|
||||
const int32_t s0 = octx->op_params[0]; \
|
||||
const int32_t s1 = octx->op_params[1]; \
|
||||
const int32_t s2 = octx->op_params[2]; \
|
||||
const int32_t s3 = octx->op_params[3]; \
|
||||
\
|
||||
const uint32_t i0_src0 = htp_roll_wrap(-s0, ne0); \
|
||||
const uint32_t n0 = ne0 - i0_src0;
|
||||
|
||||
#define htp_roll_dma_preamble dma_queue * q = octx->ctx->dma[0];
|
||||
|
||||
static inline void roll_dma_push(dma_queue * q,
|
||||
uintptr_t dst,
|
||||
uintptr_t src,
|
||||
uint32_t dst_stride,
|
||||
uint32_t src_stride,
|
||||
uint32_t bytes,
|
||||
uint32_t nrows) {
|
||||
if (bytes == 0 || nrows == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
if (!dma_queue_push(q, dma_make_ptr((void *) dst, (const void *) src), dst_stride, src_stride, bytes, nrows)) {
|
||||
dma_queue_flush(q);
|
||||
dma_queue_push(q, dma_make_ptr((void *) dst, (const void *) src),
|
||||
dst_stride, src_stride, bytes, nrows);
|
||||
}
|
||||
}
|
||||
|
||||
static inline void roll_dma_push_rows(dma_queue * q,
|
||||
const struct htp_tensor * dst,
|
||||
const struct htp_tensor * src0,
|
||||
uint32_t dst_row,
|
||||
uint32_t src_row,
|
||||
uint32_t nrows,
|
||||
uint32_t row_size,
|
||||
uint32_t i0_src0) {
|
||||
const uintptr_t dst_base = dst->data + (uintptr_t) dst_row * row_size;
|
||||
const uintptr_t src_base = src0->data + (uintptr_t) src_row * row_size;
|
||||
const uint32_t n0 = src0->ne[0] - i0_src0;
|
||||
|
||||
roll_dma_push(q, dst_base, src_base + (uintptr_t) i0_src0 * sizeof(float),
|
||||
row_size, row_size, n0 * sizeof(float), nrows);
|
||||
roll_dma_push(q, dst_base + (uintptr_t) n0 * sizeof(float), src_base,
|
||||
row_size, row_size, i0_src0 * sizeof(float), nrows);
|
||||
}
|
||||
|
||||
// Same row-wrap split as roll_dma_push_rows, but addressed with explicit byte strides so it
|
||||
// also works for a src0 that is row-contiguous only (e.g. a permuted view) rather than fully packed.
|
||||
static inline void roll_dma_push_range(dma_queue * q,
|
||||
uintptr_t dst_row,
|
||||
uintptr_t src_row,
|
||||
uint32_t dst_stride,
|
||||
uint32_t src_stride,
|
||||
uint32_t nrows,
|
||||
uint32_t i0_src0,
|
||||
uint32_t n0) {
|
||||
roll_dma_push(q, dst_row, src_row + (uintptr_t) i0_src0 * sizeof(float),
|
||||
dst_stride, src_stride, n0 * sizeof(float), nrows);
|
||||
roll_dma_push(q, dst_row + (uintptr_t) n0 * sizeof(float), src_row,
|
||||
dst_stride, src_stride, i0_src0 * sizeof(float), nrows);
|
||||
}
|
||||
|
||||
static int roll_dma_f32_contiguous(struct htp_ops_context * octx) {
|
||||
htp_roll_preamble;
|
||||
htp_roll_dma_preamble;
|
||||
|
||||
const uint32_t row_size = ne0 * sizeof(float);
|
||||
|
||||
if (s1 == 0 && s2 == 0 && s3 == 0) {
|
||||
roll_dma_push_rows(q, dst, src0, 0, 0, ne1 * ne2 * ne3, row_size, i0_src0);
|
||||
dma_queue_flush(q);
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
if (s1 == 0) {
|
||||
const uint32_t i2_src0 = htp_roll_wrap(-s2, ne2);
|
||||
for (uint32_t i3 = 0; i3 < ne3; i3++) {
|
||||
const uint32_t i03 = htp_roll_wrap((int32_t) i3 - s3, ne3);
|
||||
const uint32_t dst_row0 = i3 * ne2 * ne1;
|
||||
const uint32_t src_row0 = (i03 * ne2 + i2_src0) * ne1;
|
||||
const uint32_t n2_first = ne2 - i2_src0;
|
||||
|
||||
roll_dma_push_rows(q, dst, src0, dst_row0, src_row0, n2_first * ne1,
|
||||
row_size, i0_src0);
|
||||
roll_dma_push_rows(q, dst, src0, dst_row0 + n2_first * ne1, i03 * ne2 * ne1,
|
||||
i2_src0 * ne1, row_size, i0_src0);
|
||||
}
|
||||
|
||||
dma_queue_flush(q);
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
const uint32_t i1_src0 = htp_roll_wrap(-s1, ne1);
|
||||
const uint32_t n1_first = ne1 - i1_src0;
|
||||
|
||||
for (uint32_t i3 = 0; i3 < ne3; i3++) {
|
||||
const uint32_t i03 = htp_roll_wrap((int32_t) i3 - s3, ne3);
|
||||
for (uint32_t i2 = 0; i2 < ne2; i2++) {
|
||||
const uint32_t i02 = htp_roll_wrap((int32_t) i2 - s2, ne2);
|
||||
const uint32_t dst_row0 = (i3 * ne2 + i2) * ne1;
|
||||
const uint32_t src_row0 = (i03 * ne2 + i02) * ne1;
|
||||
|
||||
roll_dma_push_rows(q, dst, src0, dst_row0, src_row0 + i1_src0,
|
||||
n1_first, row_size, i0_src0);
|
||||
roll_dma_push_rows(q, dst, src0, dst_row0 + n1_first, src_row0,
|
||||
i1_src0, row_size, i0_src0);
|
||||
}
|
||||
}
|
||||
|
||||
dma_queue_flush(q);
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
// DMA path for a row-contiguous but otherwise arbitrarily strided src0 (e.g. a permuted view).
|
||||
// Same row-wrap split as above, one DMA push per (i2,i3), addressed via the real nb01/nb02/nb03
|
||||
// instead of assuming a packed layout.
|
||||
static int roll_dma_f32_strided(struct htp_ops_context * octx) {
|
||||
htp_roll_preamble;
|
||||
htp_roll_dma_preamble;
|
||||
|
||||
const uint32_t i1_src0 = htp_roll_wrap(-s1, ne1);
|
||||
const uint32_t n1_first = ne1 - i1_src0;
|
||||
|
||||
for (uint32_t i3 = 0; i3 < ne3; i3++) {
|
||||
const uint32_t i03 = htp_roll_wrap((int32_t) i3 - s3, ne3);
|
||||
for (uint32_t i2 = 0; i2 < ne2; i2++) {
|
||||
const uint32_t i02 = htp_roll_wrap((int32_t) i2 - s2, ne2);
|
||||
|
||||
const uintptr_t dst_row0 = dst->data + (uintptr_t) i2 * nb2 + (uintptr_t) i3 * nb3;
|
||||
const uintptr_t src_row0 = src0->data + (uintptr_t) i02 * nb02 + (uintptr_t) i03 * nb03;
|
||||
|
||||
roll_dma_push_range(q, dst_row0, src_row0 + (uintptr_t) i1_src0 * nb01,
|
||||
nb1, nb01, n1_first, i0_src0, n0);
|
||||
roll_dma_push_range(q, dst_row0 + (uintptr_t) n1_first * nb1, src_row0,
|
||||
nb1, nb01, i1_src0, i0_src0, n0);
|
||||
}
|
||||
}
|
||||
|
||||
dma_queue_flush(q);
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
static void roll_thread_f32(unsigned int nth, unsigned int ith, void * data) {
|
||||
struct htp_roll_context * rctx = (struct htp_roll_context *) data;
|
||||
struct htp_ops_context * octx = rctx->octx;
|
||||
|
||||
htp_roll_preamble;
|
||||
|
||||
const uint32_t row_start = rctx->row_start + rctx->nrows_per_thread * ith;
|
||||
const uint32_t row_end = MIN(row_start + rctx->nrows_per_thread, rctx->row_start + rctx->nrows);
|
||||
if (row_start >= row_end) {
|
||||
return;
|
||||
}
|
||||
|
||||
struct htp_thread_trace * tr = &octx->ctx->trace[ith];
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, row_start);
|
||||
|
||||
for (uint32_t row = row_start; row < row_end; row++) {
|
||||
const uint32_t i3 = fastdiv(row, &rctx->div_ne2_ne1);
|
||||
const uint32_t rem = row - i3 * ne2 * ne1;
|
||||
const uint32_t i2 = fastdiv(rem, &rctx->div_ne1);
|
||||
const uint32_t i1 = rem - i2 * ne1;
|
||||
|
||||
const uint32_t i01 = htp_roll_wrap((int32_t) i1 - s1, ne1);
|
||||
const uint32_t i02 = htp_roll_wrap((int32_t) i2 - s2, ne2);
|
||||
const uint32_t i03 = htp_roll_wrap((int32_t) i3 - s3, ne3);
|
||||
|
||||
const uint8_t * src_row = (const uint8_t *) src0->data + i01*nb01 + i02*nb02 + i03*nb03;
|
||||
uint8_t * dst_row = (uint8_t *) dst->data + i1*nb1 + i2*nb2 + i3*nb3;
|
||||
|
||||
hex_l2fetch(src_row + i0_src0 * sizeof(float), n0 * sizeof(float), ne0 * sizeof(float), 1);
|
||||
hvx_copy_uu(dst_row, src_row + i0_src0 * sizeof(float), n0, sizeof(float));
|
||||
|
||||
if (i0_src0 != 0) {
|
||||
hex_l2fetch(src_row, i0_src0 * sizeof(float), ne0 * sizeof(float), 1);
|
||||
hvx_copy_uu(dst_row + n0 * sizeof(float), src_row, i0_src0, sizeof(float));
|
||||
}
|
||||
}
|
||||
|
||||
htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, row_start);
|
||||
|
||||
FARF(HIGH, "roll %d/%d: (%ux%ux%ux%u) rows %u:%u shift=(%d,%d,%d,%d)\n",
|
||||
ith, nth, ne0, ne1, ne2, ne3,
|
||||
row_start, row_end, s0, s1, s2, s3);
|
||||
}
|
||||
|
||||
int execute_op_roll_f32(struct htp_ops_context * octx) {
|
||||
htp_roll_preamble;
|
||||
|
||||
if (src0->type != HTP_TYPE_F32 || dst->type != HTP_TYPE_F32) {
|
||||
FARF(ERROR, "roll: unsupported type %u -> %u\n", src0->type, dst->type);
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
|
||||
if (src0->nb[0] != sizeof(float) || dst->nb[0] != sizeof(float)) {
|
||||
FARF(ERROR, "roll: unsupported nb0 %u -> %u\n", src0->nb[0], dst->nb[0]);
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
|
||||
if (src0->ne[0] != ne0 || src0->ne[1] != ne1 ||
|
||||
src0->ne[2] != ne2 || src0->ne[3] != ne3) {
|
||||
FARF(ERROR, "roll: shape mismatch\n");
|
||||
return HTP_STATUS_INVAL_PARAMS;
|
||||
}
|
||||
|
||||
if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
const uint32_t total_rows = ne1 * ne2 * ne3;
|
||||
const size_t dst_row_size = ne0 * sizeof(float);
|
||||
|
||||
uint32_t row_start = 0;
|
||||
uint32_t nrows = total_rows;
|
||||
|
||||
if (octx->ctx->mdev.count > 1) {
|
||||
uint32_t rows_per_chunk = 0;
|
||||
htp_tensor_mdev_rows_per_chunk(dst, sizeof(float), (uint32_t) dst_row_size, &rows_per_chunk);
|
||||
const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, rows_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;
|
||||
}
|
||||
|
||||
if (octx->ctx->mdev.count <= 1) {
|
||||
if (htp_tensor_is_contiguous(src0, sizeof(float)) && htp_tensor_is_contiguous(dst, sizeof(float))) {
|
||||
return roll_dma_f32_contiguous(octx);
|
||||
}
|
||||
return roll_dma_f32_strided(octx);
|
||||
}
|
||||
|
||||
const uint32_t n_threads = octx->n_threads;
|
||||
struct htp_roll_context rctx = {
|
||||
.octx = octx,
|
||||
.row_start = row_start,
|
||||
.nrows = nrows,
|
||||
.nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div),
|
||||
.div_ne1 = init_fastdiv_values(dst->ne[1]),
|
||||
.div_ne2_ne1 = init_fastdiv_values(dst->ne[2] * dst->ne[1]),
|
||||
};
|
||||
|
||||
work_queue_run(octx->ctx->work_queue, roll_thread_f32, &rctx, n_threads);
|
||||
|
||||
return HTP_STATUS_OK;
|
||||
}
|
||||
|
||||
int op_roll(struct htp_ops_context * octx) {
|
||||
switch (octx->src[0]->type) {
|
||||
case HTP_TYPE_F32:
|
||||
return execute_op_roll_f32(octx);
|
||||
|
||||
default:
|
||||
return HTP_STATUS_NO_SUPPORT;
|
||||
}
|
||||
}
|
||||
@@ -7,6 +7,24 @@
|
||||
|
||||
#include <vector>
|
||||
|
||||
// must stay in sync with the kernel_fwht_<type>_<N> templates in misc.metal
|
||||
static bool ggml_metal_fwht_supported_size(int64_t n) {
|
||||
return n == 64 || n == 128 || n == 256 || n == 512;
|
||||
}
|
||||
|
||||
// the FWHT kernels handle a Hadamard-hinted MUL_MAT only under these conditions. supports_op
|
||||
// and the dispatch must ask the same question: an F16 src1 that is admitted but then falls
|
||||
// through reaches the generic path, which has no F32 src0 by F16 src1 kernel.
|
||||
bool ggml_metal_op_mul_mat_use_fwht(const struct ggml_tensor * op) {
|
||||
return ggml_get_op_params_i32(op, 1) == GGML_HINT_SRC0_IS_HADAMARD &&
|
||||
op->type == GGML_TYPE_F32 &&
|
||||
(op->src[1]->type == GGML_TYPE_F32 || op->src[1]->type == GGML_TYPE_F16) &&
|
||||
ggml_is_contiguous(op->src[1]) &&
|
||||
ggml_is_contiguous(op) &&
|
||||
ggml_are_same_shape(op->src[1], op) &&
|
||||
ggml_metal_fwht_supported_size(op->src[1]->ne[0]);
|
||||
}
|
||||
|
||||
bool ggml_metal_op_mul_mat_use_mm(const struct ggml_tensor * op, bool has_simdgroup_mm) {
|
||||
const int64_t ne00 = op->src[0]->ne[0];
|
||||
const int64_t ne11 = op->src[1]->ne[1];
|
||||
@@ -222,38 +240,63 @@ struct node_info {
|
||||
void add_fused(ggml_tensor * t) {
|
||||
fused.push_back(t);
|
||||
}
|
||||
|
||||
bool is_output(const ggml_tensor * t) const {
|
||||
if (t == node) {
|
||||
return true;
|
||||
}
|
||||
for (const auto * f : fused) {
|
||||
if (t == f) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
};
|
||||
|
||||
static std::vector<int> ggml_metal_graph_optimize_reorder(const std::vector<node_info> & nodes) {
|
||||
// helper to add node src and dst ranges
|
||||
const auto & h_add = [](ggml_mem_ranges_t mrs, const node_info & node) {
|
||||
// only external sources matter: sources produced by the fused group are internal
|
||||
for (int i = 0; i < GGML_MAX_SRC; i++) {
|
||||
if (node.node->src[i]) {
|
||||
if (!ggml_mem_ranges_add_src(mrs, node.node->src[i])) {
|
||||
const ggml_tensor * src = node.node->src[i];
|
||||
if (src && !node.is_output(src)) {
|
||||
if (!ggml_mem_ranges_add_src(mrs, src)) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// keep track of the sources of the fused nodes as well
|
||||
for (const auto * fused : node.fused) {
|
||||
for (int i = 0; i < GGML_MAX_SRC; i++) {
|
||||
if (fused->src[i]) {
|
||||
if (!ggml_mem_ranges_add_src(mrs, fused->src[i])) {
|
||||
const ggml_tensor * src = fused->src[i];
|
||||
if (src && !node.is_output(src)) {
|
||||
if (!ggml_mem_ranges_add_src(mrs, src)) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return ggml_mem_ranges_add_dst(mrs, node.dst());
|
||||
// all fused tensors are produced by the fused kernel
|
||||
if (!ggml_mem_ranges_add_dst(mrs, node.node)) {
|
||||
return false;
|
||||
}
|
||||
for (const auto * fused : node.fused) {
|
||||
if (!ggml_mem_ranges_add_dst(mrs, fused)) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
};
|
||||
|
||||
// helper to check if a node can run concurrently with the existing set of nodes
|
||||
const auto & h_check = [](ggml_mem_ranges_t mrs, const node_info & node) {
|
||||
for (int i = 0; i < GGML_MAX_SRC; i++) {
|
||||
if (node.node->src[i]) {
|
||||
if (!ggml_mem_ranges_check_src(mrs, node.node->src[i])) {
|
||||
const ggml_tensor * src = node.node->src[i];
|
||||
if (src && !node.is_output(src)) {
|
||||
if (!ggml_mem_ranges_check_src(mrs, src)) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
@@ -261,15 +304,25 @@ static std::vector<int> ggml_metal_graph_optimize_reorder(const std::vector<node
|
||||
|
||||
for (const auto * fused : node.fused) {
|
||||
for (int i = 0; i < GGML_MAX_SRC; i++) {
|
||||
if (fused->src[i]) {
|
||||
if (!ggml_mem_ranges_check_src(mrs, fused->src[i])) {
|
||||
const ggml_tensor * src = fused->src[i];
|
||||
if (src && !node.is_output(src)) {
|
||||
if (!ggml_mem_ranges_check_src(mrs, src)) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return ggml_mem_ranges_check_dst(mrs, node.dst());
|
||||
if (!ggml_mem_ranges_check_dst(mrs, node.node)) {
|
||||
return false;
|
||||
}
|
||||
for (const auto * fused : node.fused) {
|
||||
if (!ggml_mem_ranges_check_dst(mrs, fused)) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
};
|
||||
|
||||
// perform reorders only across these types of ops
|
||||
|
||||
@@ -48,6 +48,7 @@ bool ggml_mem_ranges_check(ggml_mem_ranges_t mrs, const struct ggml_tensor * ten
|
||||
void ggml_graph_optimize(struct ggml_cgraph * gf);
|
||||
|
||||
// mat-mat vs mat-vec dispatch; used by both supports_op and ggml_metal_op_mul_mat*
|
||||
bool ggml_metal_op_mul_mat_use_fwht (const struct ggml_tensor * op);
|
||||
bool ggml_metal_op_mul_mat_use_mm (const struct ggml_tensor * op, bool has_simdgroup_mm);
|
||||
bool ggml_metal_op_mul_mat_id_use_mm(const struct ggml_tensor * op, bool has_simdgroup_mm);
|
||||
|
||||
|
||||
@@ -30,7 +30,8 @@ struct ggml_metal {
|
||||
ggml_metal_device_t dev;
|
||||
ggml_metal_library_t lib;
|
||||
|
||||
ggml_metal_event_t ev_cpy; // for async copies
|
||||
ggml_metal_event_t ev_cpy; // for async copies
|
||||
ggml_metal_event_t ev_sync; // destination completion signal
|
||||
|
||||
dispatch_queue_t d_queue;
|
||||
|
||||
@@ -129,7 +130,8 @@ ggml_metal_t ggml_metal_init(ggml_metal_device_t dev) {
|
||||
}
|
||||
}
|
||||
|
||||
res->ev_cpy = ggml_metal_device_event_init(dev);
|
||||
res->ev_cpy = ggml_metal_device_event_init(dev);
|
||||
res->ev_sync = ggml_metal_device_event_init(dev);
|
||||
|
||||
const struct ggml_metal_device_props * props_dev = ggml_metal_device_get_props(dev);
|
||||
|
||||
@@ -240,6 +242,7 @@ void ggml_metal_free(ggml_metal_t ctx) {
|
||||
dispatch_release(ctx->d_queue);
|
||||
|
||||
ggml_metal_device_event_free(ctx->dev, ctx->ev_cpy);
|
||||
ggml_metal_device_event_free(ctx->dev, ctx->ev_sync);
|
||||
|
||||
free(ctx);
|
||||
}
|
||||
@@ -421,10 +424,23 @@ bool ggml_metal_cpy_tensor_async(ggml_metal_t ctx_src, ggml_metal_t ctx_dst, con
|
||||
return false;
|
||||
}
|
||||
|
||||
id<MTLCommandQueue> dst_queue = ggml_metal_device_get_queue(ctx_dst->dev);
|
||||
id<MTLCommandBuffer> sync_cmd_buf = [dst_queue commandBuffer];
|
||||
|
||||
ggml_metal_event_encode_signal(ctx_dst->ev_sync, sync_cmd_buf);
|
||||
|
||||
[sync_cmd_buf commit];
|
||||
|
||||
[ctx_dst->cmd_bufs_ext addObject:sync_cmd_buf];
|
||||
ctx_dst->cmd_buf_last = sync_cmd_buf;
|
||||
|
||||
[sync_cmd_buf retain];
|
||||
|
||||
// queue the copy operation into the Metal context
|
||||
// this will be queued at the end, after any currently ongoing GPU operations
|
||||
id<MTLCommandQueue> queue = ggml_metal_device_get_queue(ctx_src->dev);
|
||||
id<MTLCommandBuffer> cmd_buf = [queue commandBuffer];
|
||||
ggml_metal_event_encode_wait(ctx_dst->ev_sync, cmd_buf);
|
||||
id<MTLBlitCommandEncoder> encoder = [cmd_buf blitCommandEncoder];
|
||||
|
||||
[encoder copyFromBuffer:bid_src.metal
|
||||
|
||||
@@ -496,14 +496,29 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexe
|
||||
return res;
|
||||
}
|
||||
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc(ggml_metal_library_t lib, ggml_op op) {
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc(ggml_metal_library_t lib, const ggml_tensor * op) {
|
||||
const char * name = nullptr;
|
||||
|
||||
switch (op) {
|
||||
case GGML_OP_DSV4_HC_COMB: name = "kernel_dsv4_hc_comb_f32"; break;
|
||||
case GGML_OP_DSV4_HC_PRE: name = "kernel_dsv4_hc_pre_f32"; break;
|
||||
case GGML_OP_DSV4_HC_POST: name = "kernel_dsv4_hc_post_f32"; break;
|
||||
default: GGML_ABORT("fatal error");
|
||||
switch (op->op) {
|
||||
case GGML_OP_DSV4_HC_COMB:
|
||||
name = "kernel_dsv4_hc_comb_f32";
|
||||
break;
|
||||
case GGML_OP_DSV4_HC_PRE:
|
||||
if (ggml_get_op_params_i32(op, 1) != 0) {
|
||||
name = "kernel_dsv4_hc_pre_gated_f32";
|
||||
} else {
|
||||
name = "kernel_dsv4_hc_pre_f32";
|
||||
}
|
||||
break;
|
||||
case GGML_OP_DSV4_HC_POST:
|
||||
if (op->src[3]) {
|
||||
name = "kernel_dsv4_hc_post_f32";
|
||||
} else {
|
||||
name = "kernel_dsv4_hc_post_nocomb_f32";
|
||||
}
|
||||
break;
|
||||
default:
|
||||
GGML_ABORT("fatal error");
|
||||
}
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
@@ -514,7 +529,8 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc(ggml_met
|
||||
return res;
|
||||
}
|
||||
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv(ggml_metal_library_t lib, const ggml_tensor * op) {
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv(
|
||||
ggml_metal_library_t lib, const ggml_tensor * op, int32_t nc, bool use_silu) {
|
||||
GGML_ASSERT(op->src[0]->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(op->src[1]->type == GGML_TYPE_F32);
|
||||
|
||||
@@ -531,17 +547,24 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv(ggml_me
|
||||
}
|
||||
|
||||
snprintf(base, 256, "kernel_ssm_conv_%s_%s%s", ggml_type_name(op->src[0]->type), ggml_type_name(op->src[1]->type), suffix);
|
||||
snprintf(name, 256, "%s", base);
|
||||
snprintf(name, 256, "%s_nc=%d_silu=%d", base, nc, use_silu ? 1 : 0);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
if (!res.pipeline) {
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, nullptr);
|
||||
ggml_metal_cv_t cv = ggml_metal_cv_init();
|
||||
ggml_metal_cv_set_bool(cv, use_silu, FC_SSM_CONV + 1);
|
||||
ggml_metal_cv_set_int32(cv, nc, FC_SSM_CONV + 2);
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||
|
||||
ggml_metal_cv_free(cv);
|
||||
}
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched(ggml_metal_library_t lib, const ggml_tensor * op, int ssm_conv_bs) {
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched(
|
||||
ggml_metal_library_t lib, const ggml_tensor * op, int ssm_conv_bs, int32_t nc, bool use_silu) {
|
||||
GGML_ASSERT(op->src[0]->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(op->src[1]->type == GGML_TYPE_F32);
|
||||
|
||||
@@ -557,13 +580,15 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched
|
||||
}
|
||||
|
||||
snprintf(base, 256, "kernel_ssm_conv_%s_%s_batched%s", ggml_type_name(op->src[0]->type), ggml_type_name(op->src[1]->type), suffix);
|
||||
snprintf(name, 256, "%s_ssm_conv_bs=%d", base, ssm_conv_bs);
|
||||
snprintf(name, 256, "%s_ssm_conv_bs=%d_nc=%d_silu=%d", base, ssm_conv_bs, nc, use_silu ? 1 : 0);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
if (!res.pipeline) {
|
||||
ggml_metal_cv_t cv = ggml_metal_cv_init();
|
||||
|
||||
ggml_metal_cv_set_int16(cv, ssm_conv_bs, FC_SSM_CONV + 0);
|
||||
ggml_metal_cv_set_bool(cv, use_silu, FC_SSM_CONV + 1);
|
||||
ggml_metal_cv_set_int32(cv, nc, FC_SSM_CONV + 2);
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||
|
||||
@@ -1447,11 +1472,11 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_argsort_merge(gg
|
||||
return res;
|
||||
}
|
||||
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_fwht(ggml_metal_library_t lib, int n) {
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_fwht(ggml_metal_library_t lib, int n, ggml_type tsrc) {
|
||||
char base[256];
|
||||
char name[256];
|
||||
|
||||
snprintf(base, 256, "kernel_fwht_f32_%d", n);
|
||||
snprintf(base, 256, "kernel_fwht_%s_%d", ggml_type_name(tsrc), n);
|
||||
snprintf(name, 256, "%s", base);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
@@ -1533,6 +1558,49 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k_merge(ggml
|
||||
return res;
|
||||
}
|
||||
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_topk_moe(
|
||||
ggml_metal_library_t lib, int32_t n_expert, int32_t top_k, bool with_norm) {
|
||||
char base[256];
|
||||
char name[256];
|
||||
|
||||
snprintf(base, 256, "kernel_topk_moe_f32");
|
||||
snprintf(name, 256, "%s_n_expert=%d_top_k=%d_with_norm=%d", base, n_expert, top_k, with_norm ? 1 : 0);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
if (!res.pipeline) {
|
||||
ggml_metal_cv_t cv = ggml_metal_cv_init();
|
||||
ggml_metal_cv_set_bool (cv, with_norm, FC_TOPK_MOE + 0);
|
||||
ggml_metal_cv_set_int32(cv, n_expert, FC_TOPK_MOE + 1);
|
||||
ggml_metal_cv_set_int32(cv, top_k, FC_TOPK_MOE + 2);
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||
|
||||
ggml_metal_cv_free(cv);
|
||||
}
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_moe_reduce(ggml_metal_library_t lib, int32_t n_expert_used) {
|
||||
char base[256];
|
||||
char name[256];
|
||||
|
||||
snprintf(base, 256, "kernel_moe_reduce_f32");
|
||||
snprintf(name, 256, "%s_n_expert_used=%d", base, n_expert_used);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
if (!res.pipeline) {
|
||||
ggml_metal_cv_t cv = ggml_metal_cv_init();
|
||||
ggml_metal_cv_set_int32(cv, n_expert_used, FC_MOE_REDUCE + 0);
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||
|
||||
ggml_metal_cv_free(cv);
|
||||
}
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_pad(
|
||||
ggml_metal_library_t lib,
|
||||
const struct ggml_tensor * op,
|
||||
@@ -1987,7 +2055,48 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_norm(ggml_metal_
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
if (!res.pipeline) {
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, nullptr);
|
||||
ggml_metal_cv_t cv = ggml_metal_cv_init();
|
||||
ggml_metal_cv_set_bool(cv, false, FC_NORM + 0);
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||
|
||||
ggml_metal_cv_free(cv);
|
||||
}
|
||||
|
||||
res.smem = 32*sizeof(float);
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_norm_scale(ggml_metal_library_t lib, const ggml_tensor * op) {
|
||||
assert(op->op == GGML_OP_NORM || op->op == GGML_OP_RMS_NORM);
|
||||
|
||||
GGML_ASSERT(ggml_is_contiguous_rows(op->src[0]));
|
||||
|
||||
char base[256];
|
||||
char name[256];
|
||||
|
||||
const char * suffix = "";
|
||||
if (op->ne[0] % 4 == 0) {
|
||||
suffix = "_4";
|
||||
}
|
||||
|
||||
switch (op->op) {
|
||||
case GGML_OP_NORM: snprintf(base, 256, "kernel_norm_mul_f32%s", suffix); break;
|
||||
case GGML_OP_RMS_NORM: snprintf(base, 256, "kernel_rms_norm_mul_f32%s", suffix); break;
|
||||
default: GGML_ABORT("fatal error");
|
||||
}
|
||||
|
||||
snprintf(name, 256, "%s_use_scale", base);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
if (!res.pipeline) {
|
||||
ggml_metal_cv_t cv = ggml_metal_cv_init();
|
||||
ggml_metal_cv_set_bool(cv, true, FC_NORM + 0);
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||
|
||||
ggml_metal_cv_free(cv);
|
||||
}
|
||||
|
||||
res.smem = 32*sizeof(float);
|
||||
|
||||
@@ -126,9 +126,9 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_cumsum_ad
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_tri (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_soft_max (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc (ggml_metal_library_t lib, enum ggml_op op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched (ggml_metal_library_t lib, const struct ggml_tensor * op, int ssm_conv_bs);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv (ggml_metal_library_t lib, const struct ggml_tensor * op, int32_t nc, bool use_silu);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched (ggml_metal_library_t lib, const struct ggml_tensor * op, int ssm_conv_bs, int32_t nc, bool use_silu);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan (ggml_metal_library_t lib, const struct ggml_tensor * op, bool tail);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan_ssd_mma (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rwkv (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
@@ -145,15 +145,19 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_id
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_argmax (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_argsort (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_argsort_merge (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_fwht (ggml_metal_library_t lib, int n);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_fwht (ggml_metal_library_t lib, int n, enum ggml_type tsrc);
|
||||
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k_radix (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_top_k_merge (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_bin (ggml_metal_library_t lib, const struct ggml_tensor * op, int32_t n_fuse );
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_topk_moe (ggml_metal_library_t lib, int32_t n_expert, int32_t top_k, bool with_norm);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_moe_reduce (ggml_metal_library_t lib, int32_t n_expert_used);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_bin (ggml_metal_library_t lib, const struct ggml_tensor * op, int32_t n_fuse);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_bin_one (ggml_metal_library_t lib, enum ggml_op op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_l2_norm (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_group_norm (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_norm (ggml_metal_library_t lib, const struct ggml_tensor * op, int32_t n_fuse);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_norm_scale (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rope (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_im2col (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_conv_transpose_1d (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
|
||||
@@ -1733,6 +1733,12 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
|
||||
op->src[0]->ne[0] != 576) {
|
||||
return false;
|
||||
}
|
||||
if (op->src[1]->ne[0] == 72 && op->src[1]->ne[0] != op->src[2]->ne[0]) {
|
||||
return false;
|
||||
}
|
||||
if (op->src[1]->ne[0] < op->src[2]->ne[0]) {
|
||||
return false;
|
||||
}
|
||||
if (op->src[1]->type != op->src[2]->type) {
|
||||
return false;
|
||||
}
|
||||
@@ -1802,8 +1808,6 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
|
||||
op->src[1]->type == GGML_TYPE_F32 &&
|
||||
op->type == GGML_TYPE_F32 &&
|
||||
op->src[0]->ne[1] == 4 &&
|
||||
op->src[1]->ne[0] == 4 &&
|
||||
op->src[1]->ne[2] == 1 &&
|
||||
ggml_is_contiguous_rows(op->src[0]) &&
|
||||
ggml_is_contiguous_rows(op->src[1]);
|
||||
case GGML_OP_DSV4_HC_POST:
|
||||
@@ -1811,17 +1815,15 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
|
||||
op->src[0]->type == GGML_TYPE_F32 &&
|
||||
op->src[1]->type == GGML_TYPE_F32 &&
|
||||
op->src[2]->type == GGML_TYPE_F32 &&
|
||||
op->src[3] != NULL &&
|
||||
op->src[3]->type == GGML_TYPE_F32 &&
|
||||
(op->src[3] == NULL || op->src[3]->type == GGML_TYPE_F32) &&
|
||||
op->type == GGML_TYPE_F32 &&
|
||||
op->src[1]->ne[1] == 4 &&
|
||||
op->src[2]->ne[0] == 4 &&
|
||||
op->src[3]->ne[0] == 4 &&
|
||||
op->src[3]->ne[1] == 4 &&
|
||||
(op->src[3] == NULL || (op->src[3]->ne[0] == 4 && op->src[3]->ne[1] == 4)) &&
|
||||
ggml_is_contiguous_rows(op->src[0]) &&
|
||||
ggml_is_contiguous_rows(op->src[1]) &&
|
||||
ggml_is_contiguous_rows(op->src[2]) &&
|
||||
ggml_is_contiguous_rows(op->src[3]);
|
||||
(op->src[3] == NULL || ggml_is_contiguous_rows(op->src[3]));
|
||||
case GGML_OP_SSM_SCAN:
|
||||
return has_simdgroup_reduction;
|
||||
case GGML_OP_SSM_CONV:
|
||||
@@ -1834,6 +1836,12 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
|
||||
case GGML_OP_SOLVE_TRI:
|
||||
return has_simdgroup_reduction && op->src[0]->type == GGML_TYPE_F32;
|
||||
case GGML_OP_MUL_MAT:
|
||||
// the FWHT kernels read an F16 source directly; every other F16 src1 path
|
||||
// still goes through ggml_metal_supports_mul_mat_op
|
||||
if (op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F16 &&
|
||||
ggml_metal_op_mul_mat_use_fwht(op)) {
|
||||
return has_simdgroup_reduction;
|
||||
}
|
||||
return ggml_metal_supports_mul_mat_op(
|
||||
has_simdgroup_reduction, op, true,
|
||||
ggml_metal_op_mul_mat_use_mm(op, has_simdgroup_mm));
|
||||
|
||||
@@ -4,9 +4,38 @@
|
||||
#include "ggml-metal-device.h"
|
||||
|
||||
#include <algorithm>
|
||||
#include <cstddef>
|
||||
#include <cstring>
|
||||
#include <set>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
struct ggml_metal_fusion {
|
||||
ggml_metal_fusion_id id;
|
||||
|
||||
std::vector<ggml_op> ops; // op sequence (fixed length, non-empty nodes)
|
||||
std::vector<ggml_op> ops_all; // full raw op sequence (may include empty RESHAPE/VIEW nodes)
|
||||
std::vector<int> outs; // additional fused output nodes, relative to ops
|
||||
|
||||
// if unsafe: the generic chain/shape + ggml_can_fuse_subgraph checks are skipped and the
|
||||
// check callback below is the sole validator (used for patterns that are not elision chains,
|
||||
// e.g. the gdn + cache-cpy write-through fusion)
|
||||
bool unsafe;
|
||||
|
||||
// extra backend constraints on top of ggml_can_fuse_subgraph
|
||||
// nodes[j] is the j-th node of the pattern; node_idxs[idx + j] is its raw graph index
|
||||
bool (*check)(const struct ggml_metal_fusion * fusion,
|
||||
const struct ggml_tensor * const * nodes,
|
||||
const struct ggml_cgraph * gf,
|
||||
const int * node_idxs,
|
||||
int idx,
|
||||
ggml_metal_fusion_mode mode);
|
||||
};
|
||||
|
||||
ggml_metal_fusion_id ggml_metal_fusion_get_id(const ggml_metal_fusion * fusion) {
|
||||
return fusion->id;
|
||||
}
|
||||
|
||||
// ---- helpers -------------------------------------------------------------
|
||||
|
||||
// true if two tensors live in the same Metal buffer
|
||||
@@ -31,12 +60,30 @@ static bool ggml_metal_fusion_same_buffer(const ggml_tensor * a, const ggml_tens
|
||||
static bool ggml_metal_fusion_check_norm(
|
||||
const ggml_metal_fusion * fusion,
|
||||
const ggml_tensor * const * nodes,
|
||||
const ggml_cgraph * gf,
|
||||
const int * node_idxs,
|
||||
int idx,
|
||||
ggml_metal_fusion_mode mode) {
|
||||
GGML_UNUSED(mode);
|
||||
GGML_UNUSED(gf);
|
||||
GGML_UNUSED(node_idxs);
|
||||
GGML_UNUSED(idx);
|
||||
|
||||
GGML_ASSERT(fusion->n_ops >= 2);
|
||||
GGML_ASSERT(fusion->ops.size() >= 2);
|
||||
|
||||
for (int j = 1; j < fusion->n_ops; j++) {
|
||||
if (fusion->id == GGML_METAL_FUSION_NORM_SCALE) {
|
||||
GGML_ASSERT(fusion->ops.size() == 2);
|
||||
|
||||
const ggml_tensor * scale = nodes[1];
|
||||
if (scale->op != GGML_OP_SCALE || scale->src[0] != nodes[0] || scale->src[1] ||
|
||||
scale->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
for (int j = 1; j < (int) fusion->ops.size(); j++) {
|
||||
// the fused MUL/ADD must read the previous node as src0
|
||||
if (nodes[j]->src[0] != nodes[j - 1]) {
|
||||
return false;
|
||||
@@ -59,15 +106,53 @@ static bool ggml_metal_fusion_check_norm(
|
||||
return true;
|
||||
}
|
||||
|
||||
// SSM_CONV + UNARY (silu)
|
||||
static bool ggml_metal_fusion_check_ssm_conv_silu(
|
||||
const ggml_metal_fusion * fusion,
|
||||
const ggml_tensor * const * nodes,
|
||||
const ggml_cgraph * gf,
|
||||
const int * node_idxs,
|
||||
int idx,
|
||||
ggml_metal_fusion_mode mode) {
|
||||
GGML_UNUSED(fusion);
|
||||
GGML_UNUSED(gf);
|
||||
GGML_UNUSED(node_idxs);
|
||||
GGML_UNUSED(idx);
|
||||
GGML_UNUSED(mode);
|
||||
|
||||
const ggml_tensor * conv = nodes[0];
|
||||
const ggml_tensor * un = nodes[1];
|
||||
|
||||
if (conv->op != GGML_OP_SSM_CONV || un->op != GGML_OP_UNARY || un->src[0] != conv || un->src[1]) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (ggml_get_unary_op(un) != GGML_UNARY_OP_SILU) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (conv->type != GGML_TYPE_F32 || un->type != GGML_TYPE_F32 || !ggml_is_contiguous_rows(un)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
// ADD x N: each ADD reads the previous ADD as src0, and all addends must share layout
|
||||
// (and, in FULL mode, live in the same Metal buffer)
|
||||
static bool ggml_metal_fusion_check_add_chain(
|
||||
const ggml_metal_fusion * fusion,
|
||||
const ggml_tensor * const * nodes,
|
||||
const ggml_cgraph * gf,
|
||||
const int * node_idxs,
|
||||
int idx,
|
||||
ggml_metal_fusion_mode mode) {
|
||||
GGML_ASSERT(fusion->n_ops >= 2);
|
||||
GGML_UNUSED(gf);
|
||||
GGML_UNUSED(node_idxs);
|
||||
GGML_UNUSED(idx);
|
||||
GGML_ASSERT(fusion->ops.size() >= 2);
|
||||
|
||||
for (int j = 1; j < fusion->n_ops; j++) {
|
||||
for (int j = 1; j < (int) fusion->ops.size(); j++) {
|
||||
if (nodes[j]->src[0] != nodes[j - 1]) {
|
||||
return false;
|
||||
}
|
||||
@@ -94,8 +179,14 @@ static bool ggml_metal_fusion_check_add_chain(
|
||||
static bool ggml_metal_fusion_check_gdn_cache(
|
||||
const ggml_metal_fusion * fusion,
|
||||
const ggml_tensor * const * nodes,
|
||||
const ggml_cgraph * gf,
|
||||
const int * node_idxs,
|
||||
int idx,
|
||||
ggml_metal_fusion_mode mode) {
|
||||
GGML_UNUSED(fusion);
|
||||
GGML_UNUSED(gf);
|
||||
GGML_UNUSED(node_idxs);
|
||||
GGML_UNUSED(idx);
|
||||
|
||||
const ggml_tensor * gdn = nodes[0];
|
||||
const ggml_tensor * cpy = nodes[1];
|
||||
@@ -150,9 +241,15 @@ static bool ggml_metal_fusion_check_gdn_cache(
|
||||
static bool ggml_metal_fusion_check_snake(
|
||||
const ggml_metal_fusion * fusion,
|
||||
const ggml_tensor * const * nodes,
|
||||
const ggml_cgraph * gf,
|
||||
const int * node_idxs,
|
||||
int idx,
|
||||
ggml_metal_fusion_mode mode) {
|
||||
GGML_UNUSED(fusion);
|
||||
GGML_UNUSED(mode);
|
||||
GGML_UNUSED(gf);
|
||||
GGML_UNUSED(node_idxs);
|
||||
GGML_UNUSED(idx);
|
||||
|
||||
const ggml_tensor * mul0 = nodes[0];
|
||||
const ggml_tensor * sin_node = nodes[1];
|
||||
@@ -195,42 +292,465 @@ static bool ggml_metal_fusion_check_snake(
|
||||
return types_ok && shape_ok && dim_ok && contig_ok && x_in_add == x;
|
||||
}
|
||||
|
||||
// ---- patterns ------------------------------------------------------------
|
||||
#define GGML_METAL_TOPK_MOE_MAX_EXPERTS 1024
|
||||
|
||||
static const ggml_op ops_norm_mul[] = { GGML_OP_NORM, GGML_OP_MUL };
|
||||
static const ggml_op ops_norm_mul_add[] = { GGML_OP_NORM, GGML_OP_MUL, GGML_OP_ADD };
|
||||
static const ggml_op ops_rms_norm_mul[] = { GGML_OP_RMS_NORM, GGML_OP_MUL };
|
||||
static const ggml_op ops_rms_norm_mul_add[] = { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ADD };
|
||||
|
||||
static const ggml_op ops_add_2[] = { GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const ggml_op ops_add_3[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const ggml_op ops_add_4[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const ggml_op ops_add_5[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const ggml_op ops_add_6[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const ggml_op ops_add_7[] = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const ggml_op ops_snake[] = { GGML_OP_MUL, GGML_OP_SIN, GGML_OP_SQR, GGML_OP_MUL, GGML_OP_ADD };
|
||||
|
||||
static const ggml_op ops_gdn_cache[] = { GGML_OP_GATED_DELTA_NET, GGML_OP_CPY };
|
||||
|
||||
static const ggml_metal_fusion ggml_metal_fusions[] = {
|
||||
{ GGML_METAL_FUSION_NORM_MUL, ops_norm_mul, 2, false, ggml_metal_fusion_check_norm },
|
||||
{ GGML_METAL_FUSION_NORM_MUL_ADD, ops_norm_mul_add, 3, false, ggml_metal_fusion_check_norm },
|
||||
{ GGML_METAL_FUSION_NORM_MUL, ops_rms_norm_mul, 2, false, ggml_metal_fusion_check_norm },
|
||||
{ GGML_METAL_FUSION_NORM_MUL_ADD, ops_rms_norm_mul_add, 3, false, ggml_metal_fusion_check_norm },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_2, 2, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_3, 3, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_4, 4, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_5, 5, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_6, 6, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_7, 7, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_SNAKE, ops_snake, 5, false, ggml_metal_fusion_check_snake },
|
||||
{ GGML_METAL_FUSION_GDN_CACHE, ops_gdn_cache, 2, true, ggml_metal_fusion_check_gdn_cache },
|
||||
// SOFT_MAX + ARGSORT + GET_ROWS (plus optional norm/scale) for MoE routing.
|
||||
// This is a multi-output elision chain: the fused kernel writes both the selected
|
||||
// expert ids and the gathered/normalized routing weights.
|
||||
static const std::vector<ggml_op> ops_topk_moe_all = {
|
||||
GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS
|
||||
};
|
||||
static const std::vector<ggml_op> ops_topk_moe_scale_all = {
|
||||
GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, GGML_OP_SCALE
|
||||
};
|
||||
static const std::vector<ggml_op> ops_topk_moe_norm_all = {
|
||||
GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS,
|
||||
GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_RESHAPE
|
||||
};
|
||||
static const std::vector<ggml_op> ops_topk_moe_norm_scale_all = {
|
||||
GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS,
|
||||
GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_RESHAPE, GGML_OP_SCALE
|
||||
};
|
||||
|
||||
const ggml_metal_fusion * ggml_metal_fusion_all(int * n) {
|
||||
*n = (int) sizeof(ggml_metal_fusions) / sizeof(ggml_metal_fusions[0]);
|
||||
static bool ggml_metal_fusion_check_topk_moe(
|
||||
const ggml_metal_fusion * fusion,
|
||||
const ggml_tensor * const * nodes,
|
||||
const ggml_cgraph * gf,
|
||||
const int * node_idxs,
|
||||
int idx,
|
||||
ggml_metal_fusion_mode mode) {
|
||||
GGML_ASSERT(fusion->ops.size() >= 3);
|
||||
GGML_UNUSED(nodes);
|
||||
|
||||
return ggml_metal_fusions;
|
||||
const int n_ops = (int) fusion->ops.size();
|
||||
|
||||
const bool with_norm = n_ops >= 6;
|
||||
const bool with_scale = n_ops == 4 || n_ops == 7;
|
||||
|
||||
// the fusion table operates on the non-empty node sequence; the raw graph also
|
||||
// contains the RESHAPE/VIEW nodes that the fused kernel elides.
|
||||
const std::vector<ggml_op> & ops_all = fusion->ops_all;
|
||||
|
||||
const int raw_start = node_idxs[idx];
|
||||
int raw_end = node_idxs[idx + n_ops - 1];
|
||||
|
||||
// the norm variant ends with a RESHAPE that the non-empty sequence filters out;
|
||||
// include it so the output use-count check sees the real final routing tensor
|
||||
if (with_norm && !with_scale) {
|
||||
if (raw_end + 1 >= gf->n_nodes) {
|
||||
return false;
|
||||
}
|
||||
const ggml_tensor * trailing_reshape = gf->nodes[raw_end + 1];
|
||||
if (trailing_reshape->op != GGML_OP_RESHAPE || trailing_reshape->src[0] != gf->nodes[raw_end]) {
|
||||
return false;
|
||||
}
|
||||
raw_end++;
|
||||
}
|
||||
|
||||
const int raw_count = raw_end - raw_start + 1;
|
||||
if (raw_count != (int) ops_all.size()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
int raw_idxs[GGML_METAL_FUSION_MAX];
|
||||
for (int i = 0; i < raw_count; ++i) {
|
||||
raw_idxs[i] = raw_start + i;
|
||||
if (gf->nodes[raw_start + i]->op != ops_all[i]) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
const ggml_tensor * softmax = gf->nodes[raw_start];
|
||||
const ggml_tensor * probs_reshaped = gf->nodes[raw_start + 1];
|
||||
const ggml_tensor * argsort = gf->nodes[raw_start + 2];
|
||||
const ggml_tensor * ids = gf->nodes[raw_start + 3];
|
||||
const ggml_tensor * get_rows = gf->nodes[raw_start + 4];
|
||||
const ggml_tensor * out = gf->nodes[raw_end];
|
||||
const ggml_tensor * logits = softmax->src[0];
|
||||
|
||||
// the fused kernel implements plain softmax only
|
||||
float scale = 1.0f;
|
||||
float max_bias = 0.0f;
|
||||
memcpy(&scale, ((const int32_t *) softmax->op_params) + 0, sizeof(scale));
|
||||
memcpy(&max_bias, ((const int32_t *) softmax->op_params) + 1, sizeof(max_bias));
|
||||
if (scale != 1.0f || max_bias != 0.0f || softmax->src[1] || softmax->src[2]) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (logits->type != GGML_TYPE_F32 || softmax->type != GGML_TYPE_F32 ||
|
||||
out->type != GGML_TYPE_F32 || ids->type != GGML_TYPE_I32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const int64_t n_expert = logits->ne[0];
|
||||
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 ||
|
||||
n_expert > GGML_METAL_TOPK_MOE_MAX_EXPERTS || n_expert_used > GGML_METAL_TOPK_MOE_MAX_EXPERTS) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (logits->ne[2] != 1 || logits->ne[3] != 1 ||
|
||||
ids->ne[1] != n_tokens || ids->ne[2] != 1 || ids->ne[3] != 1 ||
|
||||
out->ne[0] != 1 || out->ne[1] != n_expert_used || out->ne[2] != n_tokens || out->ne[3] != 1) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!ggml_is_contiguous(logits) || !ggml_is_contiguous(out) ||
|
||||
ids->nb[0] != ggml_type_size(GGML_TYPE_I32) ||
|
||||
ids->nb[1] != ggml_type_size(GGML_TYPE_I32) * n_expert) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (probs_reshaped->src[0] != softmax || argsort->src[0] != softmax ||
|
||||
ids->src[0] != argsort || get_rows->src[0] != probs_reshaped || get_rows->src[1] != ids) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (with_norm) {
|
||||
const ggml_tensor * weights_reshaped = gf->nodes[raw_start + 5];
|
||||
const ggml_tensor * sum_rows = gf->nodes[raw_start + 6];
|
||||
const ggml_tensor * clamp = gf->nodes[raw_start + 7];
|
||||
const ggml_tensor * div = gf->nodes[raw_start + 8];
|
||||
const ggml_tensor * out_reshaped = gf->nodes[raw_start + 9];
|
||||
|
||||
if (weights_reshaped->src[0] != get_rows || sum_rows->src[0] != weights_reshaped ||
|
||||
clamp->src[0] != sum_rows || div->src[0] != weights_reshaped || div->src[1] != clamp ||
|
||||
out_reshaped->src[0] != div) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (with_scale) {
|
||||
const ggml_tensor * scale_node = gf->nodes[raw_start + 10];
|
||||
if (scale_node->src[0] != out_reshaped) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
} else if (with_scale) {
|
||||
const ggml_tensor * scale_node = gf->nodes[raw_start + 5];
|
||||
if (scale_node->src[0] != get_rows) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
const int outputs[2] = { raw_start + 3, raw_end };
|
||||
if (!ggml_can_fuse_subgraph_ext(gf, raw_idxs, raw_count, ops_all.data(), outputs, 2)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (mode == GGML_METAL_FUSION_FULL) {
|
||||
if (!logits->data || !out->data || !ids->data) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
#define GGML_METAL_MOE_REDUCE_MAX_EXPERTS 8
|
||||
|
||||
struct ggml_metal_moe_reduce_match {
|
||||
const ggml_tensor * experts;
|
||||
const ggml_tensor * weights;
|
||||
const ggml_tensor * dst;
|
||||
int node_count;
|
||||
};
|
||||
|
||||
static bool ggml_metal_fusion_match_moe_reduce(
|
||||
const ggml_cgraph * gf, int node_idx, const std::vector<ggml_op> & ops_all,
|
||||
ggml_metal_moe_reduce_match * match) {
|
||||
if (match == nullptr || node_idx < 0 || node_idx + (int) ops_all.size() > gf->n_nodes) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const ggml_tensor * mul = gf->nodes[node_idx];
|
||||
if (mul->op != GGML_OP_MUL || mul->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// MUL, then one VIEW per expert, then one ADD per additional expert
|
||||
const int raw_count = (int) ops_all.size();
|
||||
const int n_expert_used = raw_count / 2;
|
||||
|
||||
if (n_expert_used < 2 || n_expert_used > GGML_METAL_MOE_REDUCE_MAX_EXPERTS ||
|
||||
raw_count != 2 * n_expert_used) {
|
||||
return false;
|
||||
}
|
||||
|
||||
int n_views = 0;
|
||||
while (node_idx + 1 + n_views < gf->n_nodes &&
|
||||
gf->nodes[node_idx + 1 + n_views]->op == GGML_OP_VIEW) {
|
||||
n_views++;
|
||||
}
|
||||
|
||||
if (n_views != n_expert_used) {
|
||||
return false;
|
||||
}
|
||||
|
||||
for (int i = n_expert_used + 1; i < raw_count; ++i) {
|
||||
if (gf->nodes[node_idx + i]->op != GGML_OP_ADD) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
int raw_idxs[GGML_METAL_FUSION_MAX];
|
||||
for (int i = 0; i < raw_count; ++i) {
|
||||
raw_idxs[i] = node_idx + i;
|
||||
if (gf->nodes[node_idx + i]->op != ops_all[i]) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
const ggml_tensor * experts = mul->src[0];
|
||||
const ggml_tensor * weights = mul->src[1];
|
||||
const ggml_tensor * dst = gf->nodes[node_idx + raw_count - 1];
|
||||
|
||||
if (experts->type != GGML_TYPE_F32 || weights->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
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 ||
|
||||
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;
|
||||
}
|
||||
|
||||
if (!ggml_is_contiguous(experts) || !ggml_is_contiguous(weights) || !ggml_is_contiguous(dst)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
for (int i = 1; i <= n_expert_used; ++i) {
|
||||
const ggml_tensor * view = gf->nodes[node_idx + i];
|
||||
if (view->view_src != mul || view->src[0] != mul ||
|
||||
view->view_offs != (size_t) (i - 1) * mul->nb[1] ||
|
||||
view->ne[0] != n_embd || view->ne[1] != n_tokens ||
|
||||
view->nb[1] != mul->nb[2]) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
const ggml_tensor * prev_add = nullptr;
|
||||
for (int j = 1; j < n_expert_used; ++j) {
|
||||
const ggml_tensor * add = gf->nodes[node_idx + n_expert_used + j];
|
||||
const ggml_tensor * rhs = gf->nodes[node_idx + j + 1];
|
||||
const ggml_tensor * lhs = j == 1 ? gf->nodes[node_idx + 1] : prev_add;
|
||||
if (add->src[0] != lhs || add->src[1] != rhs) {
|
||||
return false;
|
||||
}
|
||||
prev_add = add;
|
||||
}
|
||||
|
||||
const int outputs[1] = { node_idx + raw_count - 1 };
|
||||
if (!ggml_can_fuse_subgraph_ext(gf, raw_idxs, raw_count, ops_all.data(), outputs, 1)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
match->experts = experts;
|
||||
match->weights = weights;
|
||||
match->dst = dst;
|
||||
match->node_count = raw_count;
|
||||
return true;
|
||||
}
|
||||
|
||||
static bool ggml_metal_fusion_check_moe_reduce(
|
||||
const ggml_metal_fusion * fusion,
|
||||
const ggml_tensor * const * nodes,
|
||||
const ggml_cgraph * gf,
|
||||
const int * node_idxs,
|
||||
int idx,
|
||||
ggml_metal_fusion_mode mode) {
|
||||
GGML_UNUSED(nodes);
|
||||
|
||||
ggml_metal_moe_reduce_match match;
|
||||
if (!ggml_metal_fusion_match_moe_reduce(gf, node_idxs[idx], fusion->ops_all, &match)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if ((int) fusion->ops.size() != match.experts->ne[1]) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const int raw_end = node_idxs[idx] + match.node_count - 1;
|
||||
if (node_idxs[idx + (int) fusion->ops.size() - 1] != raw_end) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (mode == GGML_METAL_FUSION_FULL) {
|
||||
if (!match.experts->data || !match.weights->data || !match.dst->data) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
// ---- patterns ------------------------------------------------------------
|
||||
|
||||
static const std::vector<ggml_op> ops_norm_mul = { GGML_OP_NORM, GGML_OP_MUL };
|
||||
static const std::vector<ggml_op> ops_norm_mul_add = { GGML_OP_NORM, GGML_OP_MUL, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_norm_scale = { GGML_OP_NORM, GGML_OP_SCALE };
|
||||
static const std::vector<ggml_op> ops_rms_norm_mul = { GGML_OP_RMS_NORM, GGML_OP_MUL };
|
||||
static const std::vector<ggml_op> ops_rms_norm_mul_add = { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_rms_norm_scale = { GGML_OP_RMS_NORM, GGML_OP_SCALE };
|
||||
|
||||
static const std::vector<ggml_op> ops_add_2 = { GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_add_3 = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_add_4 = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_add_5 = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_add_6 = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_add_7 = { GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_snake = { GGML_OP_MUL, GGML_OP_SIN, GGML_OP_SQR, GGML_OP_MUL, GGML_OP_ADD };
|
||||
|
||||
static const std::vector<ggml_op> ops_gdn_cache = { GGML_OP_GATED_DELTA_NET, GGML_OP_CPY };
|
||||
|
||||
static const std::vector<ggml_op> ops_topk_moe = {
|
||||
GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS
|
||||
};
|
||||
static const std::vector<ggml_op> ops_topk_moe_scale = {
|
||||
GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS, GGML_OP_SCALE
|
||||
};
|
||||
static const std::vector<ggml_op> ops_topk_moe_norm = {
|
||||
GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS,
|
||||
GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV
|
||||
};
|
||||
static const std::vector<ggml_op> ops_topk_moe_norm_scale = {
|
||||
GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS,
|
||||
GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_SCALE
|
||||
};
|
||||
|
||||
static const std::vector<ggml_op> ops_ssm_conv_silu = { GGML_OP_SSM_CONV, GGML_OP_UNARY };
|
||||
|
||||
static const std::vector<ggml_op> ops_moe_reduce_2 = { GGML_OP_MUL, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_moe_reduce_3 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_moe_reduce_4 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_moe_reduce_5 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_moe_reduce_6 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_moe_reduce_7 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const std::vector<ggml_op> ops_moe_reduce_8 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
|
||||
static const std::vector<ggml_op> ops_moe_reduce_all_2 = {
|
||||
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_ADD
|
||||
};
|
||||
static const std::vector<ggml_op> ops_moe_reduce_all_3 = {
|
||||
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_ADD, GGML_OP_ADD
|
||||
};
|
||||
static const std::vector<ggml_op> ops_moe_reduce_all_4 = {
|
||||
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW,
|
||||
GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD
|
||||
};
|
||||
static const std::vector<ggml_op> ops_moe_reduce_all_5 = {
|
||||
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW,
|
||||
GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD
|
||||
};
|
||||
static const std::vector<ggml_op> ops_moe_reduce_all_6 = {
|
||||
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW,
|
||||
GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD
|
||||
};
|
||||
static const std::vector<ggml_op> ops_moe_reduce_all_7 = {
|
||||
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW,
|
||||
GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD
|
||||
};
|
||||
static const std::vector<ggml_op> ops_moe_reduce_all_8 = {
|
||||
GGML_OP_MUL,
|
||||
GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW,
|
||||
GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD
|
||||
};
|
||||
|
||||
static const std::vector<ggml_metal_fusion> ggml_metal_fusions = {
|
||||
{ GGML_METAL_FUSION_NORM_MUL, ops_norm_mul, ops_norm_mul, {}, false, ggml_metal_fusion_check_norm },
|
||||
{ GGML_METAL_FUSION_NORM_MUL_ADD, ops_norm_mul_add, ops_norm_mul_add, {}, false, ggml_metal_fusion_check_norm },
|
||||
{ GGML_METAL_FUSION_NORM_SCALE, ops_norm_scale, ops_norm_scale, {}, false, ggml_metal_fusion_check_norm },
|
||||
{ GGML_METAL_FUSION_NORM_MUL, ops_rms_norm_mul, ops_rms_norm_mul, {}, false, ggml_metal_fusion_check_norm },
|
||||
{ GGML_METAL_FUSION_NORM_MUL_ADD, ops_rms_norm_mul_add, ops_rms_norm_mul_add, {}, false, ggml_metal_fusion_check_norm },
|
||||
{ GGML_METAL_FUSION_NORM_SCALE, ops_rms_norm_scale, ops_rms_norm_scale, {}, false, ggml_metal_fusion_check_norm },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_2, ops_add_2, {}, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_3, ops_add_3, {}, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_4, ops_add_4, {}, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_5, ops_add_5, {}, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_6, ops_add_6, {}, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_7, ops_add_7, {}, false, ggml_metal_fusion_check_add_chain },
|
||||
{ GGML_METAL_FUSION_SNAKE, ops_snake, ops_snake, {}, false, ggml_metal_fusion_check_snake },
|
||||
{ GGML_METAL_FUSION_GDN_CACHE, ops_gdn_cache, ops_gdn_cache, {}, true, ggml_metal_fusion_check_gdn_cache },
|
||||
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe, ops_topk_moe_all, {1}, true, ggml_metal_fusion_check_topk_moe },
|
||||
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_scale, ops_topk_moe_scale_all, {1}, true, ggml_metal_fusion_check_topk_moe },
|
||||
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm, ops_topk_moe_norm_all, {1}, true, ggml_metal_fusion_check_topk_moe },
|
||||
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm_scale, ops_topk_moe_norm_scale_all, {1}, true, ggml_metal_fusion_check_topk_moe },
|
||||
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_2, ops_moe_reduce_all_2, {}, true, ggml_metal_fusion_check_moe_reduce },
|
||||
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_3, ops_moe_reduce_all_3, {}, true, ggml_metal_fusion_check_moe_reduce },
|
||||
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_4, ops_moe_reduce_all_4, {}, true, ggml_metal_fusion_check_moe_reduce },
|
||||
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_5, ops_moe_reduce_all_5, {}, true, ggml_metal_fusion_check_moe_reduce },
|
||||
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_6, ops_moe_reduce_all_6, {}, true, ggml_metal_fusion_check_moe_reduce },
|
||||
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_7, ops_moe_reduce_all_7, {}, true, ggml_metal_fusion_check_moe_reduce },
|
||||
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_8, ops_moe_reduce_all_8, {}, true, ggml_metal_fusion_check_moe_reduce },
|
||||
{ GGML_METAL_FUSION_SSM_CONV_SILU, ops_ssm_conv_silu, ops_ssm_conv_silu, {}, false, ggml_metal_fusion_check_ssm_conv_silu },
|
||||
};
|
||||
|
||||
// ---- alloc deps -----------------------------------------------------------
|
||||
|
||||
static bool ggml_metal_fusion_match_raw_pattern(
|
||||
const ggml_cgraph * gf, int node_idx, const std::vector<ggml_op> & ops) {
|
||||
if (node_idx < 0 || node_idx + (int) ops.size() > gf->n_nodes) {
|
||||
return false;
|
||||
}
|
||||
|
||||
for (int i = 0; i < (int) ops.size(); ++i) {
|
||||
if (gf->nodes[node_idx + i]->op != ops[i]) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
static void ggml_metal_fusion_add_pattern_alloc_deps(
|
||||
void * user_data,
|
||||
void (*add_alloc_dep)(void *, ggml_tensor *, ggml_tensor *),
|
||||
ggml_cgraph * gf,
|
||||
const ggml_metal_fusion * fusion,
|
||||
int node_idx) {
|
||||
const int last_node = node_idx + (int) fusion->ops_all.size() - 1;
|
||||
|
||||
// keep all external inputs alive until the fused output
|
||||
std::set<ggml_tensor *> seen;
|
||||
for (int j = 0; j < (int) fusion->ops_all.size(); ++j) {
|
||||
ggml_tensor * node = gf->nodes[node_idx + j];
|
||||
for (int s = 0; s < GGML_MAX_SRC; ++s) {
|
||||
ggml_tensor * src = node->src[s];
|
||||
if (src && seen.insert(src).second) {
|
||||
add_alloc_dep(user_data, src, gf->nodes[last_node]);
|
||||
}
|
||||
}
|
||||
seen.insert(node);
|
||||
}
|
||||
}
|
||||
|
||||
void ggml_metal_fusion_add_alloc_deps(
|
||||
void * user_data,
|
||||
void (*add_alloc_dep)(void *, ggml_tensor *, ggml_tensor *),
|
||||
ggml_cgraph * gf) {
|
||||
for (int i = 0; i < gf->n_nodes; ++i) {
|
||||
const ggml_metal_fusion * best = nullptr;
|
||||
int best_raw = 0;
|
||||
|
||||
for (const ggml_metal_fusion & fusion : ggml_metal_fusions) {
|
||||
if ((int) fusion.ops_all.size() <= best_raw) {
|
||||
continue;
|
||||
}
|
||||
if (ggml_metal_fusion_match_raw_pattern(gf, i, fusion.ops_all)) {
|
||||
best = &fusion;
|
||||
best_raw = (int) fusion.ops_all.size();
|
||||
}
|
||||
}
|
||||
|
||||
if (best) {
|
||||
ggml_metal_fusion_add_pattern_alloc_deps(user_data, add_alloc_dep, gf, best, i);
|
||||
i += best_raw - 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---- shared fusion info ---------------------------------------------------
|
||||
@@ -239,7 +759,7 @@ static std::string ggml_metal_fusion_label(const ggml_metal_fusion * fusion) {
|
||||
GGML_ASSERT(fusion != nullptr);
|
||||
|
||||
std::string label;
|
||||
for (int j = 0; j < fusion->n_ops; j++) {
|
||||
for (int j = 0; j < (int) fusion->ops.size(); j++) {
|
||||
if (j > 0) {
|
||||
label += '+';
|
||||
}
|
||||
@@ -271,90 +791,77 @@ struct ggml_metal_fusion_info * ggml_metal_fusion_info_init(bool enabled, int de
|
||||
return finfo;
|
||||
}
|
||||
|
||||
void ggml_metal_fusion_info_free(struct ggml_metal_fusion_info * finfo) {
|
||||
void ggml_metal_fusion_info_free(ggml_metal_fusion_info * finfo) {
|
||||
delete finfo;
|
||||
}
|
||||
|
||||
bool ggml_metal_fusion_info_enabled(const struct ggml_metal_fusion_info * finfo) {
|
||||
bool ggml_metal_fusion_info_enabled(const ggml_metal_fusion_info * finfo) {
|
||||
return finfo->enabled;
|
||||
}
|
||||
|
||||
bool ggml_metal_fusion_info_stats(const struct ggml_metal_fusion_info * finfo) {
|
||||
bool ggml_metal_fusion_info_stats(const ggml_metal_fusion_info * finfo) {
|
||||
return finfo->stats;
|
||||
}
|
||||
|
||||
int ggml_metal_fusion_info_debug(const struct ggml_metal_fusion_info * finfo) {
|
||||
int ggml_metal_fusion_info_debug(const ggml_metal_fusion_info * finfo) {
|
||||
return finfo->debug;
|
||||
}
|
||||
|
||||
int ggml_metal_fusion_info_n_fusions(const struct ggml_metal_fusion_info * finfo) {
|
||||
int ggml_metal_fusion_info_n_fusions(const ggml_metal_fusion_info * finfo) {
|
||||
return (int) finfo->labels.size();
|
||||
}
|
||||
|
||||
const char * ggml_metal_fusion_info_label(const struct ggml_metal_fusion_info * finfo, int idx) {
|
||||
const char * ggml_metal_fusion_info_label(const ggml_metal_fusion_info * finfo, int idx) {
|
||||
GGML_ASSERT(idx >= 0 && idx < (int) finfo->labels.size());
|
||||
return finfo->labels[idx].c_str();
|
||||
}
|
||||
|
||||
uint64_t ggml_metal_fusion_info_count(const struct ggml_metal_fusion_info * finfo, int idx) {
|
||||
uint64_t ggml_metal_fusion_info_count(const ggml_metal_fusion_info * finfo, int idx) {
|
||||
GGML_ASSERT(idx >= 0 && idx < (int) finfo->counts.size());
|
||||
return finfo->counts[idx];
|
||||
}
|
||||
|
||||
void ggml_metal_fusion_info_count_fusion(struct ggml_metal_fusion_info * finfo, const struct ggml_metal_fusion * fusion) {
|
||||
void ggml_metal_fusion_info_count_fusion(ggml_metal_fusion_info * finfo, const ggml_metal_fusion * fusion) {
|
||||
if (!finfo->stats || fusion == nullptr) {
|
||||
return;
|
||||
}
|
||||
|
||||
int n = 0;
|
||||
const ggml_metal_fusion * all = ggml_metal_fusion_all(&n);
|
||||
|
||||
int idx = -1;
|
||||
for (int i = 0; i < n; i++) {
|
||||
if (&all[i] == fusion) {
|
||||
idx = i;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if (idx >= 0 && idx < (int) finfo->counts.size()) {
|
||||
const ptrdiff_t idx = fusion - ggml_metal_fusions.data();
|
||||
if (idx >= 0 && idx < (ptrdiff_t) finfo->counts.size()) {
|
||||
finfo->counts[idx]++;
|
||||
}
|
||||
}
|
||||
|
||||
void ggml_metal_fusion_info_set_enabled(struct ggml_metal_fusion_info * finfo, bool enabled) {
|
||||
void ggml_metal_fusion_info_set_enabled(ggml_metal_fusion_info * finfo, bool enabled) {
|
||||
finfo->enabled = enabled;
|
||||
}
|
||||
|
||||
void ggml_metal_fusion_info_labels_init(struct ggml_metal_fusion_info * finfo) {
|
||||
void ggml_metal_fusion_info_labels_init(ggml_metal_fusion_info * finfo) {
|
||||
if (finfo->labels_set) {
|
||||
return;
|
||||
}
|
||||
|
||||
int n = 0;
|
||||
const ggml_metal_fusion * all = ggml_metal_fusion_all(&n);
|
||||
|
||||
finfo->labels.clear();
|
||||
finfo->counts.assign(n, 0);
|
||||
finfo->labels.reserve(n);
|
||||
finfo->counts.assign(ggml_metal_fusions.size(), 0);
|
||||
finfo->labels.reserve(ggml_metal_fusions.size());
|
||||
|
||||
for (int i = 0; i < n; i++) {
|
||||
finfo->labels.emplace_back(ggml_metal_fusion_label(&all[i]));
|
||||
for (const ggml_metal_fusion & fusion : ggml_metal_fusions) {
|
||||
finfo->labels.emplace_back(ggml_metal_fusion_label(&fusion));
|
||||
}
|
||||
|
||||
finfo->labels_set = true;
|
||||
}
|
||||
|
||||
void ggml_metal_fusion_info_stats_init(struct ggml_metal_fusion_info * finfo) {
|
||||
void ggml_metal_fusion_info_stats_init(ggml_metal_fusion_info * finfo) {
|
||||
finfo->stats = true;
|
||||
ggml_metal_fusion_info_labels_init(finfo);
|
||||
}
|
||||
|
||||
void ggml_metal_fusion_info_stats_reset(struct ggml_metal_fusion_info * finfo) {
|
||||
void ggml_metal_fusion_info_stats_reset(ggml_metal_fusion_info * finfo) {
|
||||
std::fill(finfo->counts.begin(), finfo->counts.end(), 0);
|
||||
}
|
||||
|
||||
int ggml_metal_fusion_info_stats_get(const struct ggml_metal_fusion_info * finfo, const char ** labels, uint64_t * counts, int n) {
|
||||
int ggml_metal_fusion_info_stats_get(const ggml_metal_fusion_info * finfo, const char ** labels, uint64_t * counts, int n) {
|
||||
const int n_fusions = (int) finfo->labels.size();
|
||||
|
||||
if (labels == nullptr) {
|
||||
@@ -372,6 +879,97 @@ int ggml_metal_fusion_info_stats_get(const struct ggml_metal_fusion_info * finfo
|
||||
return n_fill;
|
||||
}
|
||||
|
||||
// ---- memory-range checks -------------------------------------------------
|
||||
|
||||
// reject fusions where an external source overlaps any fused output. the fused
|
||||
// kernels elide intermediate nodes, so only sources that are not part of the
|
||||
// fused subgraph can cause read/write races with the output.
|
||||
static bool ggml_metal_fusion_check_memory_ranges(
|
||||
const ggml_metal_fusion * fusion,
|
||||
const ggml_tensor * const * nodes,
|
||||
int node_count) {
|
||||
// some fused kernels write through a tensor that also appears as a source (e.g. the gdn
|
||||
// cache cpy), so a source that is the same memory as the output is not an external read
|
||||
// source
|
||||
auto same_memory = [](const ggml_tensor * a, const ggml_tensor * b) {
|
||||
if (a->data && b->data && a->data == b->data) {
|
||||
return true;
|
||||
}
|
||||
for (const ggml_tensor * v = a; v; v = v->view_src) {
|
||||
if (v == b) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
for (const ggml_tensor * v = b; v; v = v->view_src) {
|
||||
if (v == a) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
};
|
||||
|
||||
auto nodes_overlap = [](const ggml_tensor * a, const ggml_tensor * b) {
|
||||
if (!a || !b || !a->data || !b->data || !a->buffer || !b->buffer) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (a->buffer != b->buffer) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const int64_t a_start = (int64_t) a->data;
|
||||
const int64_t a_end = a_start + ggml_backend_buft_get_alloc_size(a->buffer->buft, a);
|
||||
const int64_t b_start = (int64_t) b->data;
|
||||
const int64_t b_end = b_start + ggml_backend_buft_get_alloc_size(b->buffer->buft, b);
|
||||
|
||||
return (b_start <= a_start && a_start < b_end) ||
|
||||
(a_start <= b_start && b_start < a_end);
|
||||
};
|
||||
|
||||
auto is_intermediate = [](const ggml_tensor * src, const ggml_tensor * const * nodes, int j) {
|
||||
for (int k = 0; k < j; ++k) {
|
||||
if (src == nodes[k]) {
|
||||
return true;
|
||||
}
|
||||
for (const ggml_tensor * view_src = src->view_src; view_src; view_src = view_src->view_src) {
|
||||
if (view_src == nodes[k]) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
}
|
||||
return false;
|
||||
};
|
||||
|
||||
auto check_dst = [&](const ggml_tensor * dst) {
|
||||
for (int j = 0; j < node_count; ++j) {
|
||||
for (int s = 0; s < GGML_MAX_SRC; ++s) {
|
||||
const ggml_tensor * src = nodes[j]->src[s];
|
||||
if (!src || src->op == GGML_OP_NONE || same_memory(src, dst)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
if (nodes_overlap(dst, src) && !is_intermediate(src, nodes, j)) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
}
|
||||
return true;
|
||||
};
|
||||
|
||||
if (!check_dst(nodes[node_count - 1])) {
|
||||
return false;
|
||||
}
|
||||
|
||||
for (int offset : fusion->outs) {
|
||||
GGML_ASSERT(offset >= 0 && offset < node_count);
|
||||
if (!check_dst(nodes[offset])) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
// ---- queries -------------------------------------------------------------
|
||||
|
||||
// find the longest pattern matching the node sequence starting at idx
|
||||
@@ -383,20 +981,17 @@ const ggml_metal_fusion * ggml_metal_fusion_next(
|
||||
int idx,
|
||||
ggml_metal_fusion_mode mode,
|
||||
int * n_out) {
|
||||
int n = 0;
|
||||
const ggml_metal_fusion * all = ggml_metal_fusion_all(&n);
|
||||
|
||||
const ggml_metal_fusion * res = nullptr;
|
||||
int best = 1;
|
||||
|
||||
for (int i = 0; i < n; i++) {
|
||||
const ggml_metal_fusion * fusion = &all[i];
|
||||
for (const ggml_metal_fusion & fusion : ggml_metal_fusions) {
|
||||
const int n_ops = (int) fusion.ops.size();
|
||||
|
||||
// only look for a longer match than the current best
|
||||
if (fusion->n_ops <= best) {
|
||||
if (n_ops <= best) {
|
||||
continue;
|
||||
}
|
||||
if (idx + fusion->n_ops > n_idxs) {
|
||||
if (idx + n_ops > n_idxs) {
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -404,9 +999,9 @@ const ggml_metal_fusion * ggml_metal_fusion_next(
|
||||
|
||||
// the op sequence must match exactly
|
||||
bool ok = true;
|
||||
for (int j = 0; j < fusion->n_ops; j++) {
|
||||
for (int j = 0; j < n_ops; j++) {
|
||||
nodes[j] = gf->nodes[node_idxs[idx + j]];
|
||||
if (nodes[j]->op != fusion->ops[j]) {
|
||||
if (nodes[j]->op != fusion.ops[j]) {
|
||||
ok = false;
|
||||
break;
|
||||
}
|
||||
@@ -415,10 +1010,10 @@ const ggml_metal_fusion * ggml_metal_fusion_next(
|
||||
continue;
|
||||
}
|
||||
|
||||
if (!fusion->unsafe) {
|
||||
if (!fusion.unsafe) {
|
||||
// common element-wise chain constraints: each node reads the previous one,
|
||||
// and all nodes have the same shape
|
||||
for (int j = 1; j < fusion->n_ops && ok; j++) {
|
||||
for (int j = 1; j < n_ops && ok; j++) {
|
||||
if (nodes[j]->src[0] != nodes[j - 1] && nodes[j]->src[1] != nodes[j - 1]) {
|
||||
ok = false;
|
||||
break;
|
||||
@@ -432,24 +1027,37 @@ const ggml_metal_fusion * ggml_metal_fusion_next(
|
||||
continue;
|
||||
}
|
||||
|
||||
// all current fusions are single-output elision chains, so the last node is the only output
|
||||
// TODO: multi-output fusions: store pattern-relative offsets in the table and translate them here
|
||||
int outputs_buf[1];
|
||||
outputs_buf[0] = node_idxs[idx + fusion->n_ops - 1];
|
||||
// primary output is the last node; additional outputs come from fusion.outs
|
||||
int outputs_buf[GGML_METAL_FUSION_MAX];
|
||||
outputs_buf[0] = node_idxs[idx + n_ops - 1];
|
||||
for (size_t i = 0; i < fusion.outs.size(); ++i) {
|
||||
const int out_offset = fusion.outs[i];
|
||||
GGML_ASSERT(out_offset >= 0 && out_offset < n_ops);
|
||||
outputs_buf[i + 1] = node_idxs[idx + out_offset];
|
||||
}
|
||||
|
||||
const int n_outputs = 1 + (int) fusion.outs.size();
|
||||
|
||||
// structural subgraph checks (op sequence, elidable uses, view containment)
|
||||
if (!ggml_can_fuse_subgraph_ext(gf, node_idxs + idx, fusion->n_ops, fusion->ops, outputs_buf, 1)) {
|
||||
if (!ggml_can_fuse_subgraph_ext(gf, node_idxs + idx, n_ops, fusion.ops.data(), outputs_buf, n_outputs)) {
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
// pattern-specific checks (the sole validator for unsafe patterns)
|
||||
if (fusion->check && !fusion->check(fusion, nodes, mode)) {
|
||||
if (fusion.check && !fusion.check(&fusion, nodes, gf, node_idxs, idx, mode)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
best = fusion->n_ops;
|
||||
res = fusion;
|
||||
// the compute phase has allocated tensors and can detect aliasing between
|
||||
// external sources and fused outputs; the optimizer phase cannot do this yet
|
||||
if (mode == GGML_METAL_FUSION_FULL &&
|
||||
!ggml_metal_fusion_check_memory_ranges(&fusion, nodes, n_ops)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
best = n_ops;
|
||||
res = &fusion;
|
||||
}
|
||||
|
||||
*n_out = best;
|
||||
|
||||
@@ -32,33 +32,27 @@ typedef enum ggml_metal_fusion_id {
|
||||
GGML_METAL_FUSION_NONE = 0,
|
||||
GGML_METAL_FUSION_NORM_MUL, // NORM/RMS_NORM + MUL
|
||||
GGML_METAL_FUSION_NORM_MUL_ADD, // NORM/RMS_NORM + MUL + ADD
|
||||
GGML_METAL_FUSION_NORM_SCALE, // NORM/RMS_NORM + SCALE
|
||||
GGML_METAL_FUSION_ADD_CHAIN, // ADD x N (N in [2, 7])
|
||||
GGML_METAL_FUSION_SNAKE, // MUL + SIN + SQR + MUL + ADD
|
||||
GGML_METAL_FUSION_GDN_CACHE, // GATED_DELTA_NET + CPY (write snapshots into the recurrent cache)
|
||||
GGML_METAL_FUSION_TOPK_MOE, // SOFT_MAX + ARGSORT + GET_ROWS + norm/scale (MoE routing)
|
||||
GGML_METAL_FUSION_MOE_REDUCE, // MUL + expert VIEWs + ADD chain (MoE output reduction)
|
||||
GGML_METAL_FUSION_SSM_CONV_SILU, // SSM_CONV + UNARY (silu)
|
||||
} ggml_metal_fusion_id;
|
||||
|
||||
struct ggml_metal_fusion {
|
||||
ggml_metal_fusion_id id;
|
||||
|
||||
const enum ggml_op * ops; // op sequence (fixed length)
|
||||
int n_ops; // number of ops
|
||||
|
||||
// if unsafe: the generic chain/shape + ggml_can_fuse_subgraph checks are skipped and the
|
||||
// check callback below is the sole validator (used for patterns that are not elision chains,
|
||||
// e.g. the gdn + cache-cpy write-through fusion)
|
||||
bool unsafe;
|
||||
|
||||
// extra backend constraints on top of ggml_can_fuse_subgraph
|
||||
// nodes[j] is the j-th node of the pattern
|
||||
bool (*check)(const struct ggml_metal_fusion * fusion,
|
||||
const struct ggml_tensor * const * nodes,
|
||||
ggml_metal_fusion_mode mode);
|
||||
};
|
||||
struct ggml_metal_fusion; // defined in ggml-metal-fusion.cpp
|
||||
|
||||
typedef struct ggml_metal_fusion ggml_metal_fusion;
|
||||
|
||||
// the single table of all fusions supported by the Metal backend
|
||||
const ggml_metal_fusion * ggml_metal_fusion_all(int * n);
|
||||
// access the fusion identifier without exposing the full pattern definition
|
||||
ggml_metal_fusion_id ggml_metal_fusion_get_id(const struct ggml_metal_fusion * fusion);
|
||||
|
||||
// apply any alloc-dependencies required by the fused kernels during graph optimize
|
||||
void ggml_metal_fusion_add_alloc_deps(
|
||||
void * user_data,
|
||||
void (*add_alloc_dep)(void *, struct ggml_tensor *, struct ggml_tensor *),
|
||||
struct ggml_cgraph * gf);
|
||||
|
||||
// ---- shared fusion info ---------------------------------------------------
|
||||
|
||||
|
||||
@@ -116,6 +116,9 @@
|
||||
#define FC_SUM_ROWS 1400
|
||||
#define FC_UPSCALE 1500
|
||||
#define FC_GATED_DELTA_NET 1600
|
||||
#define FC_NORM 1700
|
||||
#define FC_TOPK_MOE 1800
|
||||
#define FC_MOE_REDUCE 1900
|
||||
|
||||
// op-specific constants
|
||||
#define OP_FLASH_ATTN_EXT_NQPSG 8
|
||||
@@ -622,6 +625,7 @@ typedef struct {
|
||||
uint64_t nbf1[3];
|
||||
uint64_t nbf2[3];
|
||||
uint64_t nbf3[3];
|
||||
float scale;
|
||||
} ggml_metal_kargs_norm;
|
||||
|
||||
typedef struct {
|
||||
@@ -910,7 +914,6 @@ typedef struct {
|
||||
uint64_t nb00;
|
||||
uint64_t nb01;
|
||||
uint64_t nb02;
|
||||
int64_t ne10;
|
||||
int64_t ne11;
|
||||
uint64_t nb10;
|
||||
uint64_t nb11;
|
||||
@@ -1231,6 +1234,19 @@ typedef struct {
|
||||
int32_t top_k; // k
|
||||
} ggml_metal_kargs_top_k;
|
||||
|
||||
typedef struct {
|
||||
int32_t ne01; // n_tokens
|
||||
uint64_t nb01; // logits row stride
|
||||
uint64_t nb1_ids; // ids row stride
|
||||
float clamp;
|
||||
float scale;
|
||||
} ggml_metal_kargs_topk_moe;
|
||||
|
||||
typedef struct {
|
||||
int32_t ne00; // n_embd
|
||||
int32_t ne02; // n_tokens
|
||||
} ggml_metal_kargs_moe_reduce;
|
||||
|
||||
typedef struct {
|
||||
int32_t nrows;
|
||||
} ggml_metal_kargs_fwht;
|
||||
@@ -1283,8 +1299,10 @@ typedef struct {
|
||||
uint64_t nb_x2;
|
||||
uint64_t nb_w0;
|
||||
uint64_t nb_w1;
|
||||
uint64_t nb_w2;
|
||||
uint64_t nb_d0;
|
||||
uint64_t nb_d1;
|
||||
float scale;
|
||||
} ggml_metal_kargs_dsv4_hc_pre;
|
||||
|
||||
typedef struct {
|
||||
|
||||
@@ -226,7 +226,16 @@ static int ggml_metal_op_encode_impl(ggml_metal_op_t ctx, int idx) {
|
||||
// otherwise, we add the new ranges to the encoding context and process the node concurrently
|
||||
//
|
||||
{
|
||||
const bool is_concurrent = ggml_metal_op_concurrency_check(ctx, node);
|
||||
bool is_concurrent = ggml_metal_op_concurrency_check(ctx, node);
|
||||
|
||||
if (is_concurrent && ctx->use_fusion()) {
|
||||
int n_fuse = 1;
|
||||
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n_fuse);
|
||||
if (fusion) {
|
||||
// fused kernels write to the last node of the group, not necessarily to the first node's dst
|
||||
is_concurrent = ggml_mem_ranges_check(ctx->mem_ranges, ctx->node(idx + n_fuse - 1));
|
||||
}
|
||||
}
|
||||
|
||||
if (!is_concurrent) {
|
||||
ggml_metal_op_concurrency_reset(ctx);
|
||||
@@ -1405,7 +1414,7 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_tensor * op = ctx->node(idx);
|
||||
|
||||
ggml_metal_encoder_t enc = ctx->enc;
|
||||
auto pipeline = ggml_metal_library_get_pipeline_dsv4_hc(ctx->lib, op->op);
|
||||
auto pipeline = ggml_metal_library_get_pipeline_dsv4_hc(ctx->lib, op);
|
||||
|
||||
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||
|
||||
@@ -1467,8 +1476,10 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) {
|
||||
/*.nb_x2 =*/ x->nb[2],
|
||||
/*.nb_w0 =*/ weights->nb[0],
|
||||
/*.nb_w1 =*/ weights->nb[1],
|
||||
/*.nb_w2 =*/ weights->nb[2],
|
||||
/*.nb_d0 =*/ op->nb[0],
|
||||
/*.nb_d1 =*/ op->nb[1],
|
||||
/*.scale =*/ ggml_get_op_params_f32(op, 0),
|
||||
};
|
||||
|
||||
ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
|
||||
@@ -1491,7 +1502,6 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) {
|
||||
GGML_ASSERT(x->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(residual->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(post->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(comb->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(op->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(residual->ne[1] == 4);
|
||||
|
||||
@@ -1505,9 +1515,9 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) {
|
||||
/*.nb_r2 =*/ residual->nb[2],
|
||||
/*.nb_p0 =*/ post->nb[0],
|
||||
/*.nb_p1 =*/ post->nb[1],
|
||||
/*.nb_c0 =*/ comb->nb[0],
|
||||
/*.nb_c1 =*/ comb->nb[1],
|
||||
/*.nb_c2 =*/ comb->nb[2],
|
||||
/*.nb_c0 =*/ comb ? comb->nb[0] : 0,
|
||||
/*.nb_c1 =*/ comb ? comb->nb[1] : 0,
|
||||
/*.nb_c2 =*/ comb ? comb->nb[2] : 0,
|
||||
/*.nb_d0 =*/ op->nb[0],
|
||||
/*.nb_d1 =*/ op->nb[1],
|
||||
/*.nb_d2 =*/ op->nb[2],
|
||||
@@ -1517,8 +1527,12 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(x), 1);
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(residual), 2);
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(post), 3);
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(comb), 4);
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 5);
|
||||
if (comb) {
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(comb), 4);
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 5);
|
||||
} else {
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 4);
|
||||
}
|
||||
|
||||
const int n_tiles = (args.n_embd + 31)/32;
|
||||
const int nsg = std::min(4, n_tiles);
|
||||
@@ -1535,6 +1549,14 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) {
|
||||
int ggml_metal_op_soft_max(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_tensor * op = ctx->node(idx);
|
||||
|
||||
if (ctx->use_fusion()) {
|
||||
int n = 1;
|
||||
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
|
||||
if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_TOPK_MOE) {
|
||||
return ggml_metal_op_topk_moe(ctx, idx);
|
||||
}
|
||||
}
|
||||
|
||||
ggml_metal_library_t lib = ctx->lib;
|
||||
ggml_metal_encoder_t enc = ctx->enc;
|
||||
|
||||
@@ -1635,6 +1657,20 @@ int ggml_metal_op_ssm_conv(ggml_metal_op_t ctx, int idx) {
|
||||
GGML_TENSOR_LOCALS( int32_t, ne, op, ne);
|
||||
GGML_TENSOR_LOCALS(uint64_t, nb, op, nb);
|
||||
|
||||
int n_fuse = 1;
|
||||
bool use_silu = false;
|
||||
|
||||
if (ctx->use_fusion()) {
|
||||
int n = 1;
|
||||
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
|
||||
if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_SSM_CONV_SILU) {
|
||||
n_fuse = n;
|
||||
use_silu = true;
|
||||
|
||||
ctx->count_fusions(fusion);
|
||||
}
|
||||
}
|
||||
|
||||
ggml_metal_kargs_ssm_conv args = {
|
||||
/*.ne00 =*/ ne00,
|
||||
/*.ne01 =*/ ne01,
|
||||
@@ -1642,7 +1678,6 @@ int ggml_metal_op_ssm_conv(ggml_metal_op_t ctx, int idx) {
|
||||
/*.nb00 =*/ nb00,
|
||||
/*.nb01 =*/ nb01,
|
||||
/*.nb02 =*/ nb02,
|
||||
/*.ne10 =*/ ne10,
|
||||
/*.ne11 =*/ ne11,
|
||||
/*.nb10 =*/ nb10,
|
||||
/*.nb11 =*/ nb11,
|
||||
@@ -1654,6 +1689,8 @@ int ggml_metal_op_ssm_conv(ggml_metal_op_t ctx, int idx) {
|
||||
/*.nb2 =*/ nb2,
|
||||
};
|
||||
|
||||
const ggml_metal_buffer_id bid_dst = ggml_metal_get_buffer_id(n_fuse > 1 ? ctx->node(idx + n_fuse - 1) : op);
|
||||
|
||||
// Use batched kernel for prefill (ne1 > 1) to reduce threadgroup dispatch overhead
|
||||
const bool use_batched = (ne1 > 1);
|
||||
|
||||
@@ -1668,31 +1705,35 @@ int ggml_metal_op_ssm_conv(ggml_metal_op_t ctx, int idx) {
|
||||
else if (ne1 > 4 ) BATCH_SIZE = 8;
|
||||
else BATCH_SIZE = 2;
|
||||
|
||||
auto pipeline = ggml_metal_library_get_pipeline_ssm_conv_batched(lib, op, BATCH_SIZE);
|
||||
auto pipeline = ggml_metal_library_get_pipeline_ssm_conv_batched(lib, op, BATCH_SIZE, (int32_t) ne10, use_silu);
|
||||
|
||||
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||
ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0);
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[0]), 1);
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[1]), 2);
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 3);
|
||||
ggml_metal_encoder_set_buffer(enc, bid_dst, 3);
|
||||
|
||||
// Dispatch: ne01 rows, ceil(ne1/BATCH_SIZE) token batches, ne02 sequences
|
||||
// Each threadgroup has BATCH_SIZE threads, each handling one token
|
||||
const int n_token_batches = (ne1 + BATCH_SIZE - 1) / BATCH_SIZE;
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, ne01, n_token_batches, ne02, BATCH_SIZE, 1, 1);
|
||||
} else {
|
||||
auto pipeline = ggml_metal_library_get_pipeline_ssm_conv(lib, op);
|
||||
auto pipeline = ggml_metal_library_get_pipeline_ssm_conv(lib, op, (int32_t) ne10, use_silu);
|
||||
|
||||
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||
ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0);
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[0]), 1);
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[1]), 2);
|
||||
ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 3);
|
||||
ggml_metal_encoder_set_buffer(enc, bid_dst, 3);
|
||||
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, ne01, ne1, ne02, 1, 1, 1);
|
||||
}
|
||||
|
||||
return 1;
|
||||
if (n_fuse > 1 && ggml_metal_fusion_info_debug(ctx->finfo) > 1) {
|
||||
GGML_LOG_DEBUG("%s: fuse: SSM_CONV + UNARY\n", __func__);
|
||||
}
|
||||
|
||||
return n_fuse;
|
||||
}
|
||||
|
||||
int ggml_metal_op_ssm_scan(ggml_metal_op_t ctx, int idx) {
|
||||
@@ -1899,7 +1940,7 @@ int ggml_metal_op_gated_delta_net(ggml_metal_op_t ctx, int idx) {
|
||||
int n = 1;
|
||||
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
|
||||
|
||||
if (fusion && fusion->id == GGML_METAL_FUSION_GDN_CACHE) {
|
||||
if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_GDN_CACHE) {
|
||||
const ggml_tensor * dst_cache = ctx->node(idx + 1)->src[1]; // cache view
|
||||
|
||||
bid_out = ggml_metal_get_buffer_id(dst_cache);
|
||||
@@ -2278,12 +2319,6 @@ int ggml_metal_op_pool_1d(ggml_metal_op_t ctx, int idx) {
|
||||
return 1;
|
||||
}
|
||||
|
||||
// supported FWHT sizes, must stay in sync with the
|
||||
// kernel_fwht_f32_<N> templates in ggml-metal.metal
|
||||
static bool ggml_metal_fwht_supported_size(int64_t n) {
|
||||
return n == 64 || n == 128 || n == 256 || n == 512;
|
||||
}
|
||||
|
||||
int ggml_metal_op_fwht(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_tensor * op = ctx->node(idx);
|
||||
|
||||
@@ -2299,7 +2334,7 @@ int ggml_metal_op_fwht(ggml_metal_op_t ctx, int idx) {
|
||||
/*.nrows = */ (int32_t) nrows,
|
||||
};
|
||||
|
||||
auto pipeline = ggml_metal_library_get_pipeline_fwht(lib, n);
|
||||
auto pipeline = ggml_metal_library_get_pipeline_fwht(lib, n, src1->type);
|
||||
|
||||
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||
ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0);
|
||||
@@ -2385,17 +2420,8 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_metal_library_t lib = ctx->lib;
|
||||
ggml_metal_encoder_t enc = ctx->enc;
|
||||
|
||||
const int32_t hint = ggml_get_op_params_i32(op, 1);
|
||||
|
||||
if (hint == GGML_HINT_SRC0_IS_HADAMARD) {
|
||||
if (op->src[1]->type == GGML_TYPE_F32 &&
|
||||
op->type == GGML_TYPE_F32 &&
|
||||
ggml_is_contiguous(op->src[1]) &&
|
||||
ggml_is_contiguous(op) &&
|
||||
ggml_are_same_shape(op->src[1], op) &&
|
||||
ggml_metal_fwht_supported_size(op->src[1]->ne[0])) {
|
||||
return ggml_metal_op_fwht(ctx, idx);
|
||||
}
|
||||
if (ggml_metal_op_mul_mat_use_fwht(op)) {
|
||||
return ggml_metal_op_fwht(ctx, idx);
|
||||
}
|
||||
const ggml_metal_device_props * props_dev = ggml_metal_device_get_props(ctx->dev);
|
||||
|
||||
@@ -3816,10 +3842,16 @@ int ggml_metal_op_bin(ggml_metal_op_t ctx, int idx) {
|
||||
n_fuse = n;
|
||||
|
||||
// snake activation autofuse: mul -> sin -> sqr -> mul -> add
|
||||
if (fusion && fusion->id == GGML_METAL_FUSION_SNAKE) {
|
||||
if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_SNAKE) {
|
||||
ctx->count_fusions(fusion);
|
||||
return ggml_metal_op_snake_fused(ctx, idx);
|
||||
}
|
||||
|
||||
// MoE output reduction: experts * weights -> weighted sum
|
||||
if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_MOE_REDUCE) {
|
||||
ctx->count_fusions(fusion);
|
||||
return ggml_metal_op_moe_reduce(ctx, idx);
|
||||
}
|
||||
}
|
||||
|
||||
ggml_tensor * op = ctx->node(idx);
|
||||
@@ -3878,7 +3910,7 @@ int ggml_metal_op_bin(ggml_metal_op_t ctx, int idx) {
|
||||
// c[1] = add(c[0], b[1])
|
||||
// c[2] = add(c[1], b[2])
|
||||
// ...
|
||||
if (use_fusion && fusion && fusion->id == GGML_METAL_FUSION_ADD_CHAIN) {
|
||||
if (use_fusion && fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_ADD_CHAIN) {
|
||||
// the offsets of the fused addends are relative to the start of the src1 buffer
|
||||
for (int i = 1; i < n_fuse; i++) {
|
||||
args.o1[i] = ggml_metal_get_buffer_id(ctx->node(idx + i)->src[1]).offs;
|
||||
@@ -4122,20 +4154,22 @@ int ggml_metal_op_norm(ggml_metal_op_t ctx, int idx) {
|
||||
/*.nbf1 =*/ { nb01 },
|
||||
/*.nbf2 =*/ { nb02 },
|
||||
/*.nbf3 =*/ { nb03 },
|
||||
/*.scale =*/ 1.0f,
|
||||
};
|
||||
|
||||
int n_fuse = 1;
|
||||
bool fused_norm_scale = false;
|
||||
|
||||
ggml_metal_buffer_id bid_fuse[2] = { bid_src0, bid_src0 };
|
||||
|
||||
// d[0] = norm(a)
|
||||
// d[1] = mul(d[0], b)
|
||||
// d[1] = mul(d[0], b) or scale(d[0])
|
||||
// d[2] = add(d[1], c)
|
||||
if (use_fusion) {
|
||||
int n = 1;
|
||||
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
|
||||
|
||||
if (fusion && (fusion->id == GGML_METAL_FUSION_NORM_MUL || fusion->id == GGML_METAL_FUSION_NORM_MUL_ADD)) {
|
||||
if (fusion && (ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_NORM_MUL || ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_NORM_MUL_ADD)) {
|
||||
n_fuse = n;
|
||||
|
||||
ctx->count_fusions(fusion);
|
||||
@@ -4163,6 +4197,20 @@ int ggml_metal_op_norm(ggml_metal_op_t ctx, int idx) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (fusion && ggml_metal_fusion_get_id(fusion) == GGML_METAL_FUSION_NORM_SCALE) {
|
||||
n_fuse = n;
|
||||
fused_norm_scale = true;
|
||||
|
||||
ctx->count_fusions(fusion);
|
||||
|
||||
const ggml_tensor * scale_node = ctx->node(idx + 1);
|
||||
args.scale = ggml_get_op_params_f32(scale_node, 0);
|
||||
|
||||
if (debug_fusion > 1) {
|
||||
GGML_LOG_DEBUG("%s: fuse: %s + SCALE\n", __func__, ggml_op_name(op->op));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (n_fuse > 1) {
|
||||
@@ -4177,7 +4225,9 @@ int ggml_metal_op_norm(ggml_metal_op_t ctx, int idx) {
|
||||
}
|
||||
}
|
||||
|
||||
auto pipeline = ggml_metal_library_get_pipeline_norm(lib, op, n_fuse);
|
||||
auto pipeline = fused_norm_scale ?
|
||||
ggml_metal_library_get_pipeline_norm_scale(lib, op) :
|
||||
ggml_metal_library_get_pipeline_norm(lib, op, n_fuse);
|
||||
|
||||
int nth = 32; // SIMD width
|
||||
|
||||
@@ -5398,6 +5448,110 @@ static void ggml_metal_op_top_k_radix(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, ne01, ne02, ne03, nth, 1, 1);
|
||||
}
|
||||
|
||||
int ggml_metal_op_topk_moe(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_metal_library_t lib = ctx->lib;
|
||||
ggml_metal_encoder_t enc = ctx->enc;
|
||||
|
||||
int n_fuse = 1;
|
||||
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n_fuse);
|
||||
if (!fusion || ggml_metal_fusion_get_id(fusion) != GGML_METAL_FUSION_TOPK_MOE) {
|
||||
return 1;
|
||||
}
|
||||
|
||||
ggml_tensor * softmax = ctx->node(idx);
|
||||
ggml_tensor * logits = softmax->src[0];
|
||||
ggml_tensor * get_rows = ctx->node(idx + 2);
|
||||
ggml_tensor * ids = get_rows->src[1];
|
||||
ggml_tensor * weights = ctx->node(idx + n_fuse - 1);
|
||||
|
||||
const int64_t n_expert = logits->ne[0];
|
||||
const int64_t n_tokens = logits->ne[1];
|
||||
const int64_t n_expert_used = ids->ne[0];
|
||||
|
||||
const bool with_norm = n_fuse >= 6;
|
||||
const bool with_scale = n_fuse == 4 || n_fuse == 7;
|
||||
|
||||
float clamp = -INFINITY;
|
||||
if (with_norm) {
|
||||
ggml_tensor * clamp_node = ctx->node(idx + 4);
|
||||
clamp = ggml_get_op_params_f32(clamp_node, 0);
|
||||
}
|
||||
|
||||
float scale = 1.0f;
|
||||
if (with_scale) {
|
||||
ggml_tensor * scale_node = ctx->node(idx + n_fuse - 1);
|
||||
scale = ggml_get_op_params_f32(scale_node, 0);
|
||||
}
|
||||
|
||||
ggml_metal_kargs_topk_moe args = {
|
||||
/*.ne01 =*/ (int32_t) n_tokens,
|
||||
/*.nb01 =*/ logits->nb[1],
|
||||
/*.nb1_ids =*/ ids->nb[1],
|
||||
/*.clamp =*/ clamp,
|
||||
/*.scale =*/ scale,
|
||||
};
|
||||
|
||||
auto pipeline = ggml_metal_library_get_pipeline_topk_moe(lib, (int32_t) n_expert, (int32_t) n_expert_used, with_norm);
|
||||
|
||||
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||
ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(logits), 1);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(weights), 2);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(ids), 3);
|
||||
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, (uint32_t) n_tokens, 1, 1, 32, 1, 1);
|
||||
|
||||
ctx->count_fusions(fusion);
|
||||
|
||||
if (ggml_metal_fusion_info_debug(ctx->finfo) > 1) {
|
||||
GGML_LOG_DEBUG("%s: fuse: SOFT_MAX + ARGSORT + GET_ROWS\n", __func__);
|
||||
}
|
||||
|
||||
return n_fuse;
|
||||
}
|
||||
|
||||
int ggml_metal_op_moe_reduce(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_metal_library_t lib = ctx->lib;
|
||||
ggml_metal_encoder_t enc = ctx->enc;
|
||||
|
||||
int n_fuse = 1;
|
||||
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n_fuse);
|
||||
if (!fusion || ggml_metal_fusion_get_id(fusion) != GGML_METAL_FUSION_MOE_REDUCE) {
|
||||
return 1;
|
||||
}
|
||||
|
||||
ggml_tensor * mul = ctx->node(idx);
|
||||
ggml_tensor * experts = mul->src[0];
|
||||
ggml_tensor * weights = mul->src[1];
|
||||
ggml_tensor * dst = ctx->node(idx + n_fuse - 1);
|
||||
|
||||
ggml_metal_kargs_moe_reduce args = {
|
||||
/*.ne00 =*/ (int32_t) experts->ne[0],
|
||||
/*.ne02 =*/ (int32_t) experts->ne[2],
|
||||
};
|
||||
|
||||
auto pipeline = ggml_metal_library_get_pipeline_moe_reduce(lib, (int32_t) experts->ne[1]);
|
||||
|
||||
const int nth = std::min(256, ggml_metal_pipeline_max_theads_per_threadgroup(pipeline));
|
||||
const int n_col_tiles = (args.ne00 + nth - 1) / nth;
|
||||
|
||||
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||
ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(experts), 1);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(weights), 2);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(dst), 3);
|
||||
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, (uint32_t) args.ne02, (uint32_t) n_col_tiles, 1, nth, 1, 1);
|
||||
|
||||
ctx->count_fusions(fusion);
|
||||
|
||||
if (ggml_metal_fusion_info_debug(ctx->finfo) > 1) {
|
||||
GGML_LOG_DEBUG("%s: fuse: MOE_REDUCE\n", __func__);
|
||||
}
|
||||
|
||||
return n_fuse;
|
||||
}
|
||||
|
||||
int ggml_metal_op_top_k(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_tensor * op = ctx->node(idx);
|
||||
|
||||
|
||||
@@ -98,6 +98,8 @@ int ggml_metal_op_timestep_embedding(ggml_metal_op_t ctx, int idx);
|
||||
int ggml_metal_op_argmax (ggml_metal_op_t ctx, int idx);
|
||||
int ggml_metal_op_argsort (ggml_metal_op_t ctx, int idx);
|
||||
int ggml_metal_op_top_k (ggml_metal_op_t ctx, int idx);
|
||||
int ggml_metal_op_topk_moe (ggml_metal_op_t ctx, int idx);
|
||||
int ggml_metal_op_moe_reduce (ggml_metal_op_t ctx, int idx);
|
||||
int ggml_metal_op_tri (ggml_metal_op_t ctx, int idx);
|
||||
int ggml_metal_op_opt_step_adamw (ggml_metal_op_t ctx, int idx);
|
||||
int ggml_metal_op_opt_step_sgd (ggml_metal_op_t ctx, int idx);
|
||||
|
||||
@@ -562,7 +562,11 @@ static void ggml_backend_metal_event_wait(ggml_backend_t backend, ggml_backend_e
|
||||
}
|
||||
|
||||
static void ggml_backend_metal_graph_optimize(ggml_backend_t backend, ggml_cgraph * cgraph, ggml_backend_graph_optimize_params * params) {
|
||||
GGML_UNUSED(params);
|
||||
GGML_ASSERT(params && params->add_alloc_dep);
|
||||
|
||||
// keep the MoE weighted-reduction inputs alive until the fused output so the
|
||||
// allocator cannot reuse them while the fused kernel is still reading them
|
||||
ggml_metal_fusion_add_alloc_deps(params->user_data, params->add_alloc_dep, cgraph);
|
||||
|
||||
ggml_metal_t ctx = (ggml_metal_t)backend->context;
|
||||
|
||||
|
||||
@@ -1,5 +1,11 @@
|
||||
#include "common.h"
|
||||
|
||||
constant bool FC_topk_moe_with_norm [[function_constant(FC_TOPK_MOE + 0)]];
|
||||
constant int FC_topk_moe_n_expert [[function_constant(FC_TOPK_MOE + 1)]];
|
||||
constant int FC_topk_moe_top_k [[function_constant(FC_TOPK_MOE + 2)]];
|
||||
|
||||
constant int FC_moe_reduce_n_expert_used [[function_constant(FC_MOE_REDUCE + 0)]];
|
||||
|
||||
// bitonic sort implementation following the CUDA kernels as reference
|
||||
typedef void (argsort_t)(
|
||||
constant ggml_metal_kargs_argsort & args,
|
||||
@@ -335,3 +341,139 @@ kernel void kernel_top_k_f32_i32(
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// fused SOFT_MAX + top-k + GET_ROWS (+ optional norm/scale) for MoE routing.
|
||||
// One SIMDgroup handles one token row; n_expert is limited to 1024 by the host.
|
||||
kernel void kernel_topk_moe_f32(
|
||||
constant ggml_metal_kargs_topk_moe & args,
|
||||
device const char * src0,
|
||||
device float * weights,
|
||||
device int32_t * ids,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort tiisg[[thread_index_in_simdgroup]]) {
|
||||
const int row = (int) tgpig.x;
|
||||
if (row >= args.ne01) {
|
||||
return;
|
||||
}
|
||||
|
||||
const int n_expert = FC_topk_moe_n_expert;
|
||||
const int top_k = FC_topk_moe_top_k;
|
||||
const int lane = (int) tiisg;
|
||||
const int n_per_lane = (n_expert + 31) / 32;
|
||||
|
||||
device const float * logits_row = (device const float *) (src0 + row * args.nb01);
|
||||
device float * weights_row = weights + row * top_k;
|
||||
device int32_t * ids_row = ids + row * (args.nb1_ids / sizeof(int32_t));
|
||||
|
||||
float wt[32];
|
||||
float output_weights[32];
|
||||
FOR_UNROLL (int i = 0; i < 32; ++i) {
|
||||
wt[i] = -INFINITY;
|
||||
output_weights[i] = 0.0f;
|
||||
}
|
||||
|
||||
for (int i = lane; i < n_expert; i += 32) {
|
||||
const float v = logits_row[i];
|
||||
wt[i / 32] = isnan(v) ? -FLT_MAX : v;
|
||||
}
|
||||
|
||||
// softmax over the expert logits
|
||||
float max_val = -INFINITY;
|
||||
FOR_UNROLL (int i = 0; i < n_per_lane; ++i) {
|
||||
max_val = max(max_val, wt[i]);
|
||||
}
|
||||
max_val = simd_max(max_val);
|
||||
|
||||
float sum_val = 0.0f;
|
||||
FOR_UNROLL (int i = 0; i < n_per_lane; ++i) {
|
||||
wt[i] = exp(wt[i] - max_val);
|
||||
sum_val += wt[i];
|
||||
}
|
||||
sum_val = simd_sum(sum_val);
|
||||
|
||||
const float inv_sum = 1.0f / sum_val;
|
||||
FOR_UNROLL (int i = 0; i < n_per_lane; ++i) {
|
||||
wt[i] *= inv_sum;
|
||||
}
|
||||
|
||||
float wt_sum = 0.0f;
|
||||
|
||||
for (int k = 0; k < top_k; ++k) {
|
||||
float best_val = -INFINITY;
|
||||
int best_expert = -1;
|
||||
|
||||
FOR_UNROLL (int i = 0; i < n_per_lane; ++i) {
|
||||
const int expert = lane + i * 32;
|
||||
if (expert < n_expert && (wt[i] > best_val || (wt[i] == best_val && expert < best_expert))) {
|
||||
best_val = wt[i];
|
||||
best_expert = expert;
|
||||
}
|
||||
}
|
||||
|
||||
FOR_UNROLL (int mask = 16; mask > 0; mask >>= 1) {
|
||||
const float val = simd_shuffle_xor(best_val, mask);
|
||||
const int expert = simd_shuffle_xor(best_expert, mask);
|
||||
if (val > best_val || (val == best_val && expert < best_expert)) {
|
||||
best_val = val;
|
||||
best_expert = expert;
|
||||
}
|
||||
}
|
||||
|
||||
if ((best_expert & 31) == lane) {
|
||||
wt[best_expert / 32] = -INFINITY;
|
||||
}
|
||||
|
||||
if ((k & 31) == lane) {
|
||||
output_weights[k / 32] = best_val;
|
||||
}
|
||||
|
||||
if ((best_expert & 31) == lane) {
|
||||
ids_row[k] = best_expert;
|
||||
if (FC_topk_moe_with_norm) {
|
||||
wt_sum += best_val;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (FC_topk_moe_with_norm) {
|
||||
wt_sum = simd_sum(wt_sum);
|
||||
wt_sum = max(wt_sum, args.clamp);
|
||||
const float inv = 1.0f / wt_sum;
|
||||
FOR_UNROLL (int i = 0; i < n_per_lane; ++i) {
|
||||
output_weights[i] *= inv;
|
||||
}
|
||||
}
|
||||
|
||||
FOR_UNROLL (int i = 0; i < n_per_lane; ++i) {
|
||||
const int idx = i * 32 + lane;
|
||||
if (idx < top_k) {
|
||||
weights_row[idx] = output_weights[i] * args.scale;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// fused MoE expert weighting + reduction: weighted = sum(experts[e] * weights[e]).
|
||||
// The host guarantees all tensors are contiguous F32.
|
||||
kernel void kernel_moe_reduce_f32(
|
||||
constant ggml_metal_kargs_moe_reduce & args,
|
||||
device const float * experts,
|
||||
device const float * weights,
|
||||
device float * dst,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort3 tpitg[[thread_position_in_threadgroup]],
|
||||
ushort3 ntg[[threads_per_threadgroup]]) {
|
||||
const int64_t token = tgpig.x;
|
||||
const int64_t col = (int64_t) tgpig.y * ntg.x + tpitg.x;
|
||||
if (token >= args.ne02 || col >= args.ne00) {
|
||||
return;
|
||||
}
|
||||
|
||||
const int n_expert_used = FC_moe_reduce_n_expert_used;
|
||||
|
||||
const int64_t base = token * (int64_t) n_expert_used * args.ne00 + col;
|
||||
float sum = 0.0f;
|
||||
FOR_UNROLL (int e = 0; e < n_expert_used; ++e) {
|
||||
sum += experts[base + e * args.ne00] * weights[token * n_expert_used + e];
|
||||
}
|
||||
dst[token * args.ne00 + col] = sum;
|
||||
}
|
||||
|
||||
@@ -374,10 +374,10 @@ template [[host_name("kernel_snake_f16")]] kernel void kernel_snake<half>(const
|
||||
template [[host_name("kernel_snake_bf16")]] kernel void kernel_snake<bfloat>(constant ggml_metal_kargs_snake &, device const bfloat *, device const float *, device const float *, device bfloat *, uint, uint, uint);
|
||||
#endif
|
||||
|
||||
template<int N>
|
||||
kernel void kernel_fwht_f32(
|
||||
template<int N, typename src_t>
|
||||
kernel void kernel_fwht(
|
||||
constant ggml_metal_kargs_fwht & args,
|
||||
device const float * src,
|
||||
device const src_t * src,
|
||||
device float * dst,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort sgitg[[simdgroup_index_in_threadgroup]],
|
||||
@@ -402,13 +402,13 @@ kernel void kernel_fwht_f32(
|
||||
|
||||
float reg[NE];
|
||||
for (int i = 0; i < NE; i++) {
|
||||
reg[i] = src[i*NW + lane]*scale;
|
||||
reg[i] = float(src[i*NW + lane])*scale;
|
||||
}
|
||||
for (int i = 1; i < NW; i *= 2) {
|
||||
for (int j = 0; j < NE; j++) {
|
||||
const float val = reg[j];
|
||||
const float val2 = simd_shuffle_xor(val, i);
|
||||
reg[j] = (lane & i) == 0 ? val2 + val : val2 - val;
|
||||
reg[j] = val2 - val + 2*((lane & i) == 0)*val;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -429,12 +429,18 @@ kernel void kernel_fwht_f32(
|
||||
}
|
||||
}
|
||||
|
||||
typedef decltype(kernel_fwht_f32<64>) kernel_fwht_t;
|
||||
typedef decltype(kernel_fwht<64, float>) kernel_fwht_f32_t;
|
||||
typedef decltype(kernel_fwht<64, half>) kernel_fwht_f16_t;
|
||||
|
||||
template [[host_name("kernel_fwht_f32_64")]] kernel kernel_fwht_t kernel_fwht_f32<64>;
|
||||
template [[host_name("kernel_fwht_f32_128")]] kernel kernel_fwht_t kernel_fwht_f32<128>;
|
||||
template [[host_name("kernel_fwht_f32_256")]] kernel kernel_fwht_t kernel_fwht_f32<256>;
|
||||
template [[host_name("kernel_fwht_f32_512")]] kernel kernel_fwht_t kernel_fwht_f32<512>;
|
||||
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>;
|
||||
|
||||
kernel void kernel_dsv4_hc_comb_f32(
|
||||
constant ggml_metal_kargs_dsv4_hc_comb & args,
|
||||
@@ -531,7 +537,73 @@ kernel void kernel_dsv4_hc_pre_f32(
|
||||
result = fma(*(device const float *) (xb + ih*args.nb_x1), w[ih], result);
|
||||
}
|
||||
|
||||
*(device float *) (dst + i0*args.nb_d0 + it*args.nb_d1) = result;
|
||||
*(device float *) (dst + i0*args.nb_d0 + it*args.nb_d1) = args.scale*result;
|
||||
}
|
||||
|
||||
kernel void kernel_dsv4_hc_pre_gated_f32(
|
||||
constant ggml_metal_kargs_dsv4_hc_pre & args,
|
||||
device const char * x,
|
||||
device const char * gate,
|
||||
device char * dst,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort tiisg[[thread_index_in_simdgroup]],
|
||||
ushort sgitg[[simdgroup_index_in_threadgroup]],
|
||||
ushort3 ntg[[threads_per_threadgroup]]) {
|
||||
constexpr ushort hc = 4;
|
||||
|
||||
const int it = tgpig.y;
|
||||
const int i0 = ((int) tgpig.x*ntg.y + sgitg)*32 + tiisg;
|
||||
|
||||
if (i0 >= args.n_embd) {
|
||||
return;
|
||||
}
|
||||
|
||||
device const char * xb = x + i0*args.nb_x0 + it*args.nb_x2;
|
||||
device const char * gb = gate + i0*args.nb_w0 + it*args.nb_w2;
|
||||
float result = 0.0f;
|
||||
FOR_UNROLL (ushort ih = 0; ih < hc; ++ih) {
|
||||
const float g = 1.0f/(1.0f + exp(-*(device const float *) (gb + ih*args.nb_w1)));
|
||||
result = fma(*(device const float *) (xb + ih*args.nb_x1), g, result);
|
||||
}
|
||||
|
||||
*(device float *) (dst + i0*args.nb_d0 + it*args.nb_d1) = args.scale*result;
|
||||
}
|
||||
|
||||
kernel void kernel_dsv4_hc_post_nocomb_f32(
|
||||
constant ggml_metal_kargs_dsv4_hc_post & args,
|
||||
device const char * x,
|
||||
device const char * residual,
|
||||
device const char * post,
|
||||
device char * dst,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort tiisg[[thread_index_in_simdgroup]],
|
||||
ushort sgitg[[simdgroup_index_in_threadgroup]],
|
||||
ushort3 ntg[[threads_per_threadgroup]]) {
|
||||
constexpr ushort hc = 4;
|
||||
|
||||
const int it = tgpig.y;
|
||||
const int i0 = ((int) tgpig.x*ntg.y + sgitg)*32 + tiisg;
|
||||
|
||||
float post_lane = 0.0f;
|
||||
if (tiisg < hc) {
|
||||
post_lane = *(device const float *) (post + tiisg*args.nb_p0 + it*args.nb_p1);
|
||||
}
|
||||
|
||||
float post_reg[hc];
|
||||
FOR_UNROLL (ushort idst = 0; idst < hc; ++idst) {
|
||||
post_reg[idst] = simd_shuffle(post_lane, idst);
|
||||
}
|
||||
|
||||
if (i0 >= args.n_embd) {
|
||||
return;
|
||||
}
|
||||
|
||||
const float xv = *(device const float *) (x + i0*args.nb_x0 + it*args.nb_x1);
|
||||
device const char * rb = residual + i0*args.nb_r0 + it*args.nb_r2;
|
||||
FOR_UNROLL (ushort idst = 0; idst < hc; ++idst) {
|
||||
const float rv = *(device const float *) (rb + idst*args.nb_r1);
|
||||
*(device float *) (dst + i0*args.nb_d0 + idst*args.nb_d1 + it*args.nb_d2) = xv*post_reg[idst] + rv;
|
||||
}
|
||||
}
|
||||
|
||||
kernel void kernel_dsv4_hc_post_f32(
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
#include "common.h"
|
||||
|
||||
constant bool FC_norm_use_scale [[function_constant(FC_NORM + 0)]];
|
||||
|
||||
// F == 1 : norm (no fuse)
|
||||
// F == 2 : norm + mul
|
||||
// F == 3 : norm + mul + add
|
||||
@@ -80,7 +82,11 @@ kernel void kernel_norm_fuse_impl(
|
||||
y[i00] = (y[i00]*scale);
|
||||
}
|
||||
if (F == 2) {
|
||||
y[i00] = (y[i00]*scale)*f0[i00];
|
||||
if (FC_norm_use_scale) {
|
||||
y[i00] = (y[i00]*scale) * args.scale;
|
||||
} else {
|
||||
y[i00] = (y[i00]*scale)*f0[i00];
|
||||
}
|
||||
}
|
||||
if (F == 3) {
|
||||
y[i00] = (y[i00]*scale)*f0[i00] + f1[i00];
|
||||
@@ -155,7 +161,11 @@ kernel void kernel_rms_norm_fuse_impl(
|
||||
y[i00] = (x[i00]*scale);
|
||||
}
|
||||
if (F == 2) {
|
||||
y[i00] = (x[i00]*scale)*f0[i00];
|
||||
if (FC_norm_use_scale) {
|
||||
y[i00] = (x[i00]*scale) * args.scale;
|
||||
} else {
|
||||
y[i00] = (x[i00]*scale)*f0[i00];
|
||||
}
|
||||
}
|
||||
if (F == 3) {
|
||||
y[i00] = (x[i00]*scale)*f0[i00] + f1[i00];
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
#include "common.h"
|
||||
|
||||
constant bool FC_ssm_conv_silu [[function_constant(FC_SSM_CONV + 1)]];
|
||||
constant int FC_ssm_conv_nc [[function_constant(FC_SSM_CONV + 2)]];
|
||||
|
||||
// ref: ggml.c:ggml_compute_forward_ssm_conv_f32
|
||||
kernel void kernel_ssm_conv_f32_f32(
|
||||
constant ggml_metal_kargs_ssm_conv & args,
|
||||
@@ -13,7 +16,7 @@ kernel void kernel_ssm_conv_f32_f32(
|
||||
const int64_t i2 = tgpig.y;
|
||||
const int64_t i3 = tgpig.z;
|
||||
|
||||
const int64_t nc = args.ne10;
|
||||
const int64_t nc = FC_ssm_conv_nc;
|
||||
//const int64_t ncs = args.ne00;
|
||||
//const int64_t nr = args.ne01;
|
||||
//const int64_t n_t = args.ne1;
|
||||
@@ -25,11 +28,11 @@ kernel void kernel_ssm_conv_f32_f32(
|
||||
|
||||
float sumf = 0.0f;
|
||||
|
||||
for (int64_t i0 = 0; i0 < nc; ++i0) {
|
||||
FOR_UNROLL (int64_t i0 = 0; i0 < nc; ++i0) {
|
||||
sumf += s[i0] * c[i0];
|
||||
}
|
||||
|
||||
x[0] = sumf;
|
||||
x[0] = FC_ssm_conv_silu ? sumf/(1.0f + exp(-sumf)) : sumf;
|
||||
}
|
||||
|
||||
kernel void kernel_ssm_conv_f32_f32_4(
|
||||
@@ -44,7 +47,7 @@ kernel void kernel_ssm_conv_f32_f32_4(
|
||||
const int64_t i2 = tgpig.y;
|
||||
const int64_t i3 = tgpig.z;
|
||||
|
||||
const int64_t nc = args.ne10;
|
||||
const int64_t nc = FC_ssm_conv_nc;
|
||||
//const int64_t ncs = args.ne00;
|
||||
//const int64_t nr = args.ne01;
|
||||
//const int64_t n_t = args.ne1;
|
||||
@@ -56,11 +59,11 @@ kernel void kernel_ssm_conv_f32_f32_4(
|
||||
|
||||
float sumf = 0.0f;
|
||||
|
||||
for (int64_t i0 = 0; i0 < nc/4; ++i0) {
|
||||
FOR_UNROLL (int64_t i0 = 0; i0 < nc/4; ++i0) {
|
||||
sumf += dot(s[i0], c[i0]);
|
||||
}
|
||||
|
||||
x[0] = sumf;
|
||||
x[0] = FC_ssm_conv_silu ? sumf/(1.0f + exp(-sumf)) : sumf;
|
||||
}
|
||||
|
||||
constant short FC_ssm_conv_bs [[function_constant(FC_SSM_CONV + 0)]];
|
||||
@@ -87,7 +90,7 @@ kernel void kernel_ssm_conv_f32_f32_batched(
|
||||
const int64_t i2_off = tpitg.x;
|
||||
const int64_t i2 = i2_base + i2_off;
|
||||
|
||||
const int64_t nc = args.ne10; // conv kernel size (typically 4)
|
||||
const int64_t nc = FC_ssm_conv_nc; // conv kernel size (typically 4)
|
||||
const int64_t n_t = args.ne1; // number of tokens
|
||||
|
||||
// Bounds check for partial batches at the end
|
||||
@@ -105,11 +108,11 @@ kernel void kernel_ssm_conv_f32_f32_batched(
|
||||
device float * x = (device float *) ((device char *) dst + ir*args.nb0 + i2*args.nb1 + i3*args.nb2);
|
||||
|
||||
float sumf = 0.0f;
|
||||
for (int64_t i0 = 0; i0 < nc; ++i0) {
|
||||
FOR_UNROLL (int64_t i0 = 0; i0 < nc; ++i0) {
|
||||
sumf += s[i0] * c[i0];
|
||||
}
|
||||
|
||||
x[0] = sumf;
|
||||
x[0] = FC_ssm_conv_silu ? sumf/(1.0f + exp(-sumf)) : sumf;
|
||||
}
|
||||
|
||||
kernel void kernel_ssm_conv_f32_f32_batched_4(
|
||||
@@ -132,7 +135,7 @@ kernel void kernel_ssm_conv_f32_f32_batched_4(
|
||||
const int64_t i2_off = tpitg.x;
|
||||
const int64_t i2 = i2_base + i2_off;
|
||||
|
||||
const int64_t nc = args.ne10; // conv kernel size (typically 4)
|
||||
const int64_t nc = FC_ssm_conv_nc; // conv kernel size (typically 4)
|
||||
const int64_t n_t = args.ne1; // number of tokens
|
||||
|
||||
// Bounds check for partial batches at the end
|
||||
@@ -150,11 +153,11 @@ kernel void kernel_ssm_conv_f32_f32_batched_4(
|
||||
device float * x = (device float *) ((device char *) dst + ir*args.nb0 + i2*args.nb1 + i3*args.nb2);
|
||||
|
||||
float sumf = 0.0f;
|
||||
for (int64_t i0 = 0; i0 < nc/4; ++i0) {
|
||||
FOR_UNROLL (int64_t i0 = 0; i0 < nc/4; ++i0) {
|
||||
sumf += dot(s[i0], c[i0]);
|
||||
}
|
||||
|
||||
x[0] = sumf;
|
||||
x[0] = FC_ssm_conv_silu ? sumf/(1.0f + exp(-sumf)) : sumf;
|
||||
}
|
||||
|
||||
// ref: ggml.c:ggml_compute_forward_ssm_scan_f32, Mamba-2 part
|
||||
|
||||
@@ -191,6 +191,7 @@ set(GGML_OPENCL_KERNELS
|
||||
gemv_noshuffle_q6_k_f32_tiled
|
||||
gemm_noshuffle_q6_k_f32
|
||||
gemm_noshuffle_q6_k_f32_tiled
|
||||
gemv_noshuffle_q6_k_f32_32b_trans
|
||||
gemv_noshuffle_q5_k_f32
|
||||
gemm_noshuffle_q5_k_f32
|
||||
mul
|
||||
@@ -232,6 +233,7 @@ set(GGML_OPENCL_KERNELS
|
||||
mul_mm_f16_f32_kq_kqv
|
||||
conv2d
|
||||
conv2d_f16_f32
|
||||
flash_attn_repack
|
||||
flash_attn_pre_f16
|
||||
flash_attn_f32_f16
|
||||
flash_attn_f32_q8_0
|
||||
|
||||
@@ -567,6 +567,16 @@ struct ggml_opencl_fa_kernels {
|
||||
// attempted (variant, (dk, dv))
|
||||
// all attempted FA kernels appear here, but those not registered failed compilation
|
||||
std::set<std::pair<int, std::pair<int, int>>> variant_attempted;
|
||||
|
||||
// FA bin kernels
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
cl_kernel kernel_flash_attn_f32_f16_bin;
|
||||
|
||||
cl_kernel kernel_repack_q_for_wmm;
|
||||
cl_kernel kernel_repack_k_for_wmm;
|
||||
cl_kernel kernel_repack_v_for_wmm;
|
||||
cl_kernel kernel_repack_mask_for_wmm;
|
||||
#endif
|
||||
};
|
||||
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
@@ -1246,6 +1256,8 @@ struct ggml_backend_opencl_context {
|
||||
cl_kernel kernel_gemv_noshuffle_q6_K_f32_mc3; // multi-column (N=3) verify GEMV
|
||||
cl_kernel kernel_gemm_noshuffle_q6_K_f32;
|
||||
cl_kernel kernel_gemm_noshuffle_q6_K_f32_cok;
|
||||
cl_kernel kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8_bin;
|
||||
cl_kernel kernel_gemv_noshuffle_q6_k_f32_32b_trans;
|
||||
cl_kernel kernel_gemv_noshuffle_q5_k_f32;
|
||||
cl_kernel kernel_gemv_noshuffle_q5_k_f32_mc3; // multi-column (N=3) verify GEMV (spec/MTP)
|
||||
cl_kernel kernel_gemm_noshuffle_q5_k_f32;
|
||||
@@ -4367,6 +4379,43 @@ 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;
|
||||
if (backend_ctx->adreno_gen == ADRENO_GPU_GEN::X2E) {
|
||||
{
|
||||
std::string opts = std::string("-cl-std=") + opencl_c_std +
|
||||
" -cl-mad-enable "
|
||||
" -DSIMDGROUP_WIDTH=" +
|
||||
std::to_string(backend_ctx->adreno_wave_size);
|
||||
#ifdef GGML_OPENCL_EMBED_KERNELS
|
||||
const std::string kernel_src {
|
||||
#include "gemv_noshuffle_q6_k_f32_32b_trans.cl.h"
|
||||
};
|
||||
#else
|
||||
const std::string kernel_src = read_file("gemv_noshuffle_q6_k_f32_32b_trans.cl");
|
||||
#endif
|
||||
cl_program prog = build_program_from_source(backend_ctx, kernel_src.c_str(), opts);
|
||||
CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q6_k_f32_32b_trans =
|
||||
clCreateKernel(prog, "kernel_gemv_noshuffle_q6_k_f32_32b_trans", &err), err));
|
||||
CL_CHECK(clReleaseProgram(prog));
|
||||
GGML_LOG_CONT(".");
|
||||
}
|
||||
|
||||
if (use_adreno_bin_kernels(backend_ctx)) {
|
||||
size_t bin_size = 0;
|
||||
const char * kernel_bin = (const char *)backend_ctx->get_adreno_bin_kernel("gemm_noshuffle_q6_k_f32_32b_trans_ila_a8", &bin_size);
|
||||
if (kernel_bin && bin_size > 0) {
|
||||
cl_program bin_prog =
|
||||
build_program_from_binary(backend_ctx->context, backend_ctx->device, kernel_bin, "", bin_size);
|
||||
|
||||
CL_CHECK((backend_ctx->kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8_bin =
|
||||
clCreateKernel(bin_prog, "kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8", &err), err));
|
||||
CL_CHECK(clReleaseProgram(bin_prog));
|
||||
GGML_LOG_CONT(".");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
std::string CL_moe_compile_opts = std::string("-cl-std=") + opencl_c_std +
|
||||
" -cl-mad-enable "
|
||||
" -cl-fast-relaxed-math";
|
||||
@@ -5133,6 +5182,43 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
|
||||
CL_CHECK(clReleaseProgram(prog));
|
||||
GGML_LOG_CONT(".");
|
||||
}
|
||||
|
||||
// repack
|
||||
{
|
||||
#ifdef GGML_OPENCL_EMBED_KERNELS
|
||||
const std::string kernel_src {
|
||||
#include "flash_attn_repack.cl.h"
|
||||
};
|
||||
#else
|
||||
const std::string kernel_src = read_file("flash_attn_repack.cl");
|
||||
#endif
|
||||
cl_program prog =
|
||||
build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts);
|
||||
|
||||
CL_CHECK((backend_ctx->fa.kernel_repack_q_for_wmm = clCreateKernel(prog, "kernel_repack_q_for_wmm", &err), err));
|
||||
CL_CHECK((backend_ctx->fa.kernel_repack_k_for_wmm = clCreateKernel(prog, "kernel_repack_k_for_wmm", &err), err));
|
||||
CL_CHECK((backend_ctx->fa.kernel_repack_v_for_wmm = clCreateKernel(prog, "kernel_repack_v_for_wmm", &err), err));
|
||||
CL_CHECK((backend_ctx->fa.kernel_repack_mask_for_wmm = clCreateKernel(prog, "kernel_repack_mask_for_wmm", &err), err));
|
||||
GGML_LOG_CONT(".");
|
||||
}
|
||||
|
||||
// kernel_flash_attn_f32_f16_bin
|
||||
{
|
||||
size_t bin_size = 0;
|
||||
backend_ctx->fa.kernel_flash_attn_f32_f16_bin = nullptr;
|
||||
|
||||
if (use_adreno_bin_kernels(backend_ctx)) {
|
||||
const char * kernel_bin = (const char *)backend_ctx->get_adreno_bin_kernel("flash_attn_f32_f16_wmm", &bin_size);
|
||||
if (kernel_bin && bin_size > 0) {
|
||||
cl_program prog =
|
||||
build_program_from_binary(backend_ctx->context, backend_ctx->device, kernel_bin, CL_moe_compile_opts, bin_size);
|
||||
|
||||
CL_CHECK((backend_ctx->fa.kernel_flash_attn_f32_f16_bin = clCreateKernel(prog, "flash_attn_f32_f16", &err), err));
|
||||
CL_CHECK(clReleaseProgram(prog));
|
||||
GGML_LOG_CONT(".");
|
||||
}
|
||||
}
|
||||
}
|
||||
#endif // GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
GGML_LOG_CONT("\n");
|
||||
backend_ctx->kernels_loaded = true;
|
||||
@@ -7294,6 +7380,8 @@ struct ggml_tensor_extra_cl_q6_K {
|
||||
cl_mem ql_img = nullptr;
|
||||
// Upper 2 bits of quantized weights.
|
||||
cl_mem qh = nullptr;
|
||||
// Upper 2 bits as image1d_buffer_t
|
||||
cl_mem qh_img = nullptr;
|
||||
// Scales for each block.
|
||||
cl_mem s = nullptr;
|
||||
// Scales for each super block.
|
||||
@@ -7329,6 +7417,10 @@ struct ggml_tensor_extra_cl_q6_K {
|
||||
CL_CHECK(clReleaseMemObject(ql_img));
|
||||
ql_img = nullptr;
|
||||
}
|
||||
if (qh_img != nullptr) {
|
||||
CL_CHECK(clReleaseMemObject(qh_img));
|
||||
qh_img = nullptr;
|
||||
}
|
||||
|
||||
size_ql = 0;
|
||||
size_qh = 0;
|
||||
@@ -8487,6 +8579,28 @@ inline bool use_q4_0_bin_kernels(const ggml_backend_opencl_context *backend_ctx,
|
||||
#endif
|
||||
}
|
||||
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
static bool use_fa_bin_kernels_prefill(const ggml_backend_opencl_context * backend_ctx, const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v) {
|
||||
if (backend_ctx->fa.kernel_flash_attn_f32_f16_bin == nullptr) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const bool is_mixed = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_F16 && v->type == GGML_TYPE_F16;
|
||||
const bool is_q8_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q8_0 && v->type == GGML_TYPE_Q8_0;
|
||||
|
||||
const int n_q = q->ne[1];
|
||||
const int dk = q->ne[0];
|
||||
const int dv = v->ne[0];
|
||||
|
||||
constexpr bool prefill_only = true;
|
||||
|
||||
return (backend_ctx->gpu_family == GPU_FAMILY::ADRENO &&
|
||||
(is_mixed || is_q8_0) && (dk == dv)
|
||||
&& (dk == 64 || dk == 128 || dk == 256 || dk == 512)
|
||||
&& (!prefill_only || n_q != 1));
|
||||
}
|
||||
#endif
|
||||
|
||||
// The flat-GEMV large-m escape is OPT-IN (GGML_OPENCL_FLAT_LARGE_M=1) because it
|
||||
// is SLOWER than the route it replaces, not because it is unsafe. It was first
|
||||
// parked on the theory that it out-of-bounds-writes at vocab-scale shapes; that
|
||||
@@ -8566,6 +8680,21 @@ static inline bool use_flat_gemv_for_large_m_q6_K(const ggml_backend_opencl_cont
|
||||
&& tensor->ne[2] == 1 && tensor->ne[3] == 1;
|
||||
}
|
||||
|
||||
inline bool use_q6_k_bin_kernels(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) {
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
if (!backend_ctx->kernel_gemv_noshuffle_q6_k_f32_32b_trans ||
|
||||
!backend_ctx->kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8_bin) {
|
||||
return false;
|
||||
}
|
||||
return (tensor->ne[0] % 256 == 0) && (tensor->ne[1] % 64 == 0) &&
|
||||
!use_q6k_tiled(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(backend_ctx, tensor);
|
||||
#else
|
||||
GGML_UNUSED(backend_ctx);
|
||||
GGML_UNUSED(tensor);
|
||||
return false;
|
||||
#endif
|
||||
}
|
||||
|
||||
inline bool use_q4_k_bin_kernels(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) {
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
if (!backend_ctx->kernel_gemv_noshuffle_q4_k_f32_32b_trans ||
|
||||
@@ -8930,6 +9059,11 @@ static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_te
|
||||
case GGML_OP_MEAN:
|
||||
return op->src[0]->type == GGML_TYPE_F32;
|
||||
case GGML_OP_FLASH_ATTN_EXT: {
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
if (use_fa_bin_kernels_prefill(backend_ctx, op->src[0], op->src[1], op->src[2])) {
|
||||
return true;
|
||||
}
|
||||
#endif
|
||||
// The E17 compilers segfault while building FA kernels, skip E17 for now
|
||||
if (adreno_e17_compiler_quirks(backend_ctx)) {
|
||||
return false;
|
||||
@@ -11181,18 +11315,39 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
|
||||
cl_int M = tensor->ne[1]; // ne01
|
||||
cl_int K = tensor->ne[0]; // ne00
|
||||
|
||||
// Transpose ql as ushort
|
||||
transpose_2d_as_16b(backend_ctx,
|
||||
extra->ql, extra->ql, size_ql, K/4, M);
|
||||
if (use_q6_k_bin_kernels(backend_ctx, tensor)) {
|
||||
GGML_ASSERT(K % 256 == 0);
|
||||
GGML_ASSERT(M % 64 == 0);
|
||||
|
||||
// Transpose qh as uchar
|
||||
transpose_2d_as_8b(backend_ctx,
|
||||
extra->qh, extra->qh, size_qh, K/4, M);
|
||||
transpose_2d_as_32b(backend_ctx, extra->ql, extra->ql, size_ql, K/8, M);
|
||||
transpose_2d_as_32b(backend_ctx, extra->qh, extra->qh, size_qh, K/16, M);
|
||||
|
||||
// Transpose s as ushort
|
||||
transpose_2d_as_16b(backend_ctx,
|
||||
extra->s, extra->s, size_s, K/16/2, M);
|
||||
cl_image_format wimg_fmt = { CL_R, CL_UNSIGNED_INT32 };
|
||||
cl_image_desc wimg_desc;
|
||||
memset(&wimg_desc, 0, sizeof(wimg_desc));
|
||||
wimg_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
|
||||
wimg_desc.image_width = static_cast<size_t>(ggml_nelements(tensor) / 8);
|
||||
wimg_desc.buffer = extra->ql;
|
||||
CL_CHECK((extra->ql_img = clCreateImage(context, CL_MEM_READ_ONLY, &wimg_fmt, &wimg_desc, NULL, &err), err));
|
||||
|
||||
memset(&wimg_desc, 0, sizeof(wimg_desc));
|
||||
wimg_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
|
||||
wimg_desc.image_width = static_cast<size_t>(ggml_nelements(tensor) / 16);
|
||||
wimg_desc.buffer = extra->qh;
|
||||
CL_CHECK((extra->qh_img = clCreateImage(context, CL_MEM_READ_ONLY, &wimg_fmt, &wimg_desc, NULL, &err), err));
|
||||
} else {
|
||||
// Transpose ql as ushort
|
||||
transpose_2d_as_16b(backend_ctx,
|
||||
extra->ql, extra->ql, size_ql, K/4, M);
|
||||
|
||||
// Transpose qh as uchar
|
||||
transpose_2d_as_8b(backend_ctx,
|
||||
extra->qh, extra->qh, size_qh, K/4, M);
|
||||
|
||||
// Transpose s as ushort
|
||||
transpose_2d_as_16b(backend_ctx,
|
||||
extra->s, extra->s, size_s, K/16/2, M);
|
||||
}
|
||||
// Transpose d as ushort
|
||||
transpose_2d_as_16b(backend_ctx,
|
||||
extra->d, extra->d, size_d, K/256, M);
|
||||
@@ -12317,15 +12472,24 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
|
||||
|
||||
buf_trans_ql.allocate(backend_ctx->context, size_ql);
|
||||
buf_trans_qh.allocate(backend_ctx->context, size_qh);
|
||||
buf_trans_s.allocate(backend_ctx->context, size_s);
|
||||
buf_trans_d.allocate(backend_ctx->context, size_d);
|
||||
buf_unpacked.allocate(backend_ctx->context, ggml_nbytes(tensor));
|
||||
|
||||
// transpose ql, qh, s and d back
|
||||
transpose_2d_as_16b(backend_ctx, extra->ql, buf_trans_ql.buffer, size_ql, M, K/4);
|
||||
transpose_2d_as_8b(backend_ctx, extra->qh, buf_trans_qh.buffer, size_qh, M, K/4);
|
||||
transpose_2d_as_16b(backend_ctx, extra->s, buf_trans_s.buffer, size_s, M, K/16/2);
|
||||
transpose_2d_as_16b(backend_ctx, extra->d, buf_trans_d.buffer, size_d, M, K/256);
|
||||
cl_mem s_buffer;
|
||||
if (use_q6_k_bin_kernels(backend_ctx, tensor)) {
|
||||
transpose_2d_as_32b(backend_ctx, extra->ql, buf_trans_ql.buffer, size_ql, M, K/8);
|
||||
transpose_2d_as_32b(backend_ctx, extra->qh, buf_trans_qh.buffer, size_qh, M, K/16);
|
||||
// s is left row-major, untransposed, for the binary layout.
|
||||
s_buffer = extra->s;
|
||||
} else {
|
||||
// transpose ql, qh, s and d back
|
||||
buf_trans_s.allocate(backend_ctx->context, size_s);
|
||||
transpose_2d_as_16b(backend_ctx, extra->ql, buf_trans_ql.buffer, size_ql, M, K/4);
|
||||
transpose_2d_as_8b(backend_ctx, extra->qh, buf_trans_qh.buffer, size_qh, M, K/4);
|
||||
transpose_2d_as_16b(backend_ctx, extra->s, buf_trans_s.buffer, size_s, M, K/16/2);
|
||||
s_buffer = buf_trans_s.buffer;
|
||||
}
|
||||
transpose_2d_as_16b(backend_ctx, extra->d, buf_trans_d.buffer, size_d, M, K/256);
|
||||
|
||||
// unpack
|
||||
cl_uchar mask = 0xFF;
|
||||
@@ -12333,7 +12497,7 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
|
||||
cl_kernel kernel = backend_ctx->kernel_restore_block_q6_K_noshuffle;
|
||||
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &buf_trans_ql.buffer));
|
||||
CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &buf_trans_qh.buffer));
|
||||
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &buf_trans_s.buffer));
|
||||
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &s_buffer));
|
||||
CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &buf_trans_d.buffer));
|
||||
CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &buf_unpacked.buffer));
|
||||
CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_uchar), &mask));
|
||||
@@ -17108,6 +17272,407 @@ static void ggml_cl_adreno_xmem_attn_run(
|
||||
|
||||
#endif // GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
static void ggml_cl_flash_attn_prefill_bin(ggml_backend_t backend, const ggml_tensor * q, const ggml_tensor * k, ggml_tensor * dst) {
|
||||
const ggml_tensor * v = dst->src[2];
|
||||
const ggml_tensor * mask = dst->src[3];
|
||||
const ggml_tensor * sinks = dst->src[4];
|
||||
GGML_ASSERT(q->extra);
|
||||
GGML_ASSERT(k->extra);
|
||||
GGML_ASSERT(v->extra);
|
||||
GGML_ASSERT(dst->extra);
|
||||
if (mask) {
|
||||
GGML_ASSERT(mask->extra);
|
||||
}
|
||||
if (sinks) {
|
||||
GGML_ASSERT(sinks->extra);
|
||||
}
|
||||
|
||||
ggml_backend_opencl_context *backend_ctx = (ggml_backend_opencl_context *)backend->context;
|
||||
cl_context context = backend_ctx->context;
|
||||
|
||||
const int n_q = q->ne[1];
|
||||
const int n_kv = k->ne[1];
|
||||
const int d_head_q = q->ne[0];
|
||||
const int d_head_v = v->ne[0];
|
||||
const int n_head = q->ne[2];
|
||||
const int n_head_kv = k->ne[2];
|
||||
const int n_batch = q->ne[3];
|
||||
|
||||
const std::pair<int, int> dk_dv = {d_head_q, d_head_v};
|
||||
cl_kernel kernel = backend_ctx->fa.kernel_flash_attn_f32_f16_bin;
|
||||
GGML_ASSERT(kernel != NULL);
|
||||
|
||||
ggml_tensor_extra_cl * extra_q = (ggml_tensor_extra_cl *)q->extra;
|
||||
ggml_tensor_extra_cl * extra_k = (ggml_tensor_extra_cl *)k->extra;
|
||||
ggml_tensor_extra_cl * extra_v = (ggml_tensor_extra_cl *)v->extra;
|
||||
ggml_tensor_extra_cl * extra_o = (ggml_tensor_extra_cl *)dst->extra;
|
||||
ggml_tensor_extra_cl * extra_mask = mask ? (ggml_tensor_extra_cl *)mask->extra : NULL;
|
||||
ggml_tensor_extra_cl * extra_sinks = sinks ? (ggml_tensor_extra_cl *)sinks->extra : NULL;
|
||||
|
||||
cl_ulong offset_q = extra_q->offset + q->view_offs;
|
||||
cl_ulong offset_o = extra_o->offset + dst->view_offs;
|
||||
|
||||
cl_mem mask_buffer = extra_mask ? extra_mask->data_device : NULL;
|
||||
cl_ulong offset_mask = extra_mask ? extra_mask->offset + mask->view_offs : 0;
|
||||
cl_mem sinks_buffer = extra_sinks ? extra_sinks->data_device : NULL;
|
||||
cl_ulong offset_sinks = extra_sinks ? extra_sinks->offset + sinks->view_offs : 0;
|
||||
|
||||
const cl_ulong q_nb1 = q->nb[1];
|
||||
const cl_ulong q_nb2 = q->nb[2];
|
||||
const cl_ulong q_nb3 = q->nb[3];
|
||||
|
||||
cl_mem k_data_device = extra_k->data_device;
|
||||
cl_ulong offset_k = extra_k->offset + k->view_offs;
|
||||
cl_ulong k_nb1 = k->nb[1];
|
||||
cl_ulong k_nb2 = k->nb[2];
|
||||
cl_ulong k_nb3 = k->nb[3];
|
||||
|
||||
cl_mem v_data_device = extra_v->data_device;
|
||||
cl_ulong offset_v = extra_v->offset + v->view_offs;
|
||||
cl_ulong v_nb1 = v->nb[1];
|
||||
cl_ulong v_nb2 = v->nb[2];
|
||||
cl_ulong v_nb3 = v->nb[3];
|
||||
|
||||
const cl_ulong o_nb1 = dst->nb[1];
|
||||
const cl_ulong o_nb2 = dst->nb[2];
|
||||
const cl_ulong o_nb3 = dst->nb[3];
|
||||
|
||||
const cl_ulong mask_nb1 = mask ? mask->nb[1] : 0;
|
||||
const cl_ulong mask_nb2 = mask ? mask->nb[2] : 0;
|
||||
const cl_ulong mask_nb3 = mask ? mask->nb[3] : 0;
|
||||
const int mask_ne2 = mask ? mask->ne[2] : 0;
|
||||
const int mask_ne3 = mask ? mask->ne[3] : 0;
|
||||
|
||||
float * params = (float *)dst->op_params;
|
||||
float scale = params[0];
|
||||
float max_bias = params[1];
|
||||
float logit_softcap = params[2];
|
||||
|
||||
const int is_causal = (mask == NULL && n_q > 1 && n_q == n_kv); // redundant n_q > 1 check ?
|
||||
|
||||
const int n_head_log2_val = n_head > 0 ? 1u << (int)floorf(log2f((float)n_head)) : 0;
|
||||
const float n_head_log2_f = n_head_log2_val > 0 ? (float)n_head_log2_val : 1.0f;
|
||||
const float m0 = powf(2.0f, -(max_bias) / n_head_log2_f);
|
||||
const float m1 = powf(2.0f, -(max_bias / 2.0f) / n_head_log2_f);
|
||||
|
||||
const bool is_q8_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q8_0 && v->type == GGML_TYPE_Q8_0;
|
||||
|
||||
ggml_cl_flash_attn_temp_buffer temp_k;
|
||||
ggml_cl_flash_attn_temp_buffer temp_v;
|
||||
ggml_cl_flash_attn_temp_buffer temp_k_aos;
|
||||
ggml_cl_flash_attn_temp_buffer temp_v_aos;
|
||||
|
||||
if (is_q8_0) {
|
||||
ggml_cl_flash_attn_reconstruct_aos(
|
||||
backend_ctx, k, temp_k_aos, k_data_device, offset_k, k_nb1, k_nb2, k_nb3);
|
||||
|
||||
ggml_cl_flash_attn_reconstruct_aos(
|
||||
backend_ctx, v, temp_v_aos, v_data_device, offset_v, v_nb1, v_nb2, v_nb3);
|
||||
|
||||
bool k_done = ggml_cl_flash_attn_dequant_kv_gpu(
|
||||
backend_ctx, k, GGML_TYPE_F16, k_data_device, offset_k, k_nb1, k_nb2, k_nb3,
|
||||
temp_k, k_data_device, offset_k, k_nb1, k_nb2, k_nb3);
|
||||
|
||||
bool v_done = ggml_cl_flash_attn_dequant_kv_gpu(
|
||||
backend_ctx, v, GGML_TYPE_F16, v_data_device, offset_v, v_nb1, v_nb2, v_nb3,
|
||||
temp_v, v_data_device, offset_v, v_nb1, v_nb2, v_nb3);
|
||||
|
||||
GGML_ASSERT(k_done && v_done);
|
||||
}
|
||||
|
||||
// Allocate input/output memory buffers
|
||||
cl_mem mem_matrixQ;
|
||||
cl_mem mem_matrixK;
|
||||
cl_mem mem_matrixV;
|
||||
cl_mem mem_matrixO;
|
||||
cl_buffer_region region;
|
||||
cl_int err;
|
||||
|
||||
region.origin = offset_q;
|
||||
region.size = ggml_nbytes(q);
|
||||
mem_matrixQ = clCreateSubBuffer(extra_q->data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
region.origin = offset_k;
|
||||
region.size = is_q8_0 ? (size_t) k_nb3 * (size_t) k->ne[3] : ggml_nbytes(k);
|
||||
mem_matrixK = clCreateSubBuffer(k_data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
region.origin = offset_v;
|
||||
region.size = is_q8_0 ? (size_t) v_nb3 * (size_t) v->ne[3] : ggml_nbytes(v);
|
||||
mem_matrixV = clCreateSubBuffer(v_data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
region.origin = offset_o;
|
||||
region.size = ggml_nbytes(dst);
|
||||
mem_matrixO = clCreateSubBuffer(extra_o->data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
cl_image_format img_fmt_1d = { CL_RGBA, CL_FLOAT};
|
||||
cl_image_desc img_desc_1d;
|
||||
|
||||
// use image 1d buffer used as fallback when on mask is applied
|
||||
cl_mem mem_tex_mask_fallback_1dbuf;
|
||||
img_fmt_1d = { CL_RGBA, CL_HALF_FLOAT};
|
||||
memset(&img_desc_1d, 0, sizeof(img_desc_1d));
|
||||
img_desc_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
|
||||
img_desc_1d.image_width = 1;
|
||||
img_desc_1d.buffer = mem_matrixK;
|
||||
mem_tex_mask_fallback_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_1d, &img_desc_1d, NULL, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
cl_mem mem_tex_matrixO_1dbuf;
|
||||
img_fmt_1d = { CL_RGBA, CL_FLOAT};
|
||||
memset(&img_desc_1d, 0, sizeof(img_desc_1d));
|
||||
img_desc_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
|
||||
img_desc_1d.image_width = ggml_nbytes(dst) / 4 / 4;
|
||||
img_desc_1d.buffer = mem_matrixO;
|
||||
mem_tex_matrixO_1dbuf = clCreateImage(context, CL_MEM_WRITE_ONLY, &img_fmt_1d, &img_desc_1d, NULL, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
// The bin kernel requires 2d (or 3d) buffers packed for data loading/multiplication.
|
||||
// These repack kernels launch across all buffers to ensure compatibility
|
||||
cl_mem mem_tex_matrixMask_1dbuf = NULL;
|
||||
cl_mem mem_matrixMask = NULL;
|
||||
cl_mem mem_matrixMask_padded = NULL;
|
||||
cl_ulong mask_nb1_padded = mask_nb1, mask_nb2_padded = mask_nb2, mask_nb3_padded = mask_nb3;
|
||||
if (extra_mask) {
|
||||
// allocate mem_matrixMask w/ new padded size
|
||||
size_t n_kv_padded = GGML_PAD(n_kv, 4);
|
||||
size_t mask_nb_padded = n_kv_padded * sizeof(cl_half) * mask->ne[1] * mask->ne[2] * mask->ne[3];
|
||||
|
||||
// apply offset and create subBuffer for mask
|
||||
region.origin = offset_mask;
|
||||
region.size = ggml_nbytes(mask);
|
||||
mem_matrixMask = clCreateSubBuffer(extra_mask->data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
{
|
||||
// create padded mask to contain all data
|
||||
mem_matrixMask_padded = clCreateBuffer(context, CL_MEM_ALLOC_HOST_PTR, mask_nb_padded, NULL, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
// pass extra_mask->data_device, mem_matrixMask to kernel for copying/padding
|
||||
mask_nb1_padded = (cl_ulong)n_kv_padded * sizeof(cl_half);
|
||||
mask_nb2_padded = mask_nb1_padded * (cl_ulong)mask->ne[1];
|
||||
mask_nb3_padded = mask_nb2_padded * (cl_ulong)mask->ne[2];
|
||||
|
||||
cl_kernel repack_mask = backend_ctx->fa.kernel_repack_mask_for_wmm;
|
||||
CL_CHECK(clSetKernelArg(repack_mask, 0, sizeof(cl_mem), &mem_matrixMask));
|
||||
CL_CHECK(clSetKernelArg(repack_mask, 1, sizeof(cl_ulong), &mask_nb1));
|
||||
CL_CHECK(clSetKernelArg(repack_mask, 2, sizeof(cl_ulong), &mask_nb2));
|
||||
CL_CHECK(clSetKernelArg(repack_mask, 3, sizeof(cl_ulong), &mask_nb3));
|
||||
CL_CHECK(clSetKernelArg(repack_mask, 4, sizeof(int), &mask_ne2));
|
||||
CL_CHECK(clSetKernelArg(repack_mask, 5, sizeof(cl_mem), &mem_matrixMask_padded));
|
||||
CL_CHECK(clSetKernelArg(repack_mask, 6, sizeof(cl_ulong), &mask_nb1_padded));
|
||||
CL_CHECK(clSetKernelArg(repack_mask, 7, sizeof(cl_ulong), &mask_nb2_padded));
|
||||
CL_CHECK(clSetKernelArg(repack_mask, 8, sizeof(cl_ulong), &mask_nb3_padded));
|
||||
|
||||
size_t repack_mask_gws[3] = {(size_t)n_kv, (size_t)mask->ne[1], (size_t)mask_ne2 * (size_t)mask->ne[3]};
|
||||
backend_ctx->enqueue_ndrange_kernel(repack_mask, 3, repack_mask_gws, NULL, dst);
|
||||
}
|
||||
|
||||
// use image 1d buffer for matrix Mask (padded row stride)
|
||||
cl_image_format img_fmt_mask_1d = { CL_RGBA, CL_HALF_FLOAT};
|
||||
cl_image_desc img_desc_mask_1d;
|
||||
memset(&img_desc_mask_1d, 0, sizeof(img_desc_mask_1d));
|
||||
img_desc_mask_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
|
||||
img_desc_mask_1d.image_width = mask_nb_padded / 2 / 4;
|
||||
img_desc_mask_1d.buffer = mem_matrixMask_padded;
|
||||
mem_tex_matrixMask_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_mask_1d, &img_desc_mask_1d, NULL, &err);
|
||||
CL_CHECK(err);
|
||||
}
|
||||
|
||||
// WMM QK uses repacked 3D images.
|
||||
// Q image: rows, heads, packed depth.
|
||||
cl_image_format img_fmt_3d = { CL_RGBA, CL_HALF_FLOAT };
|
||||
cl_image_desc img_desc_3d;
|
||||
|
||||
memset(&img_desc_3d, 0, sizeof(img_desc_3d));
|
||||
img_desc_3d.image_type = CL_MEM_OBJECT_IMAGE3D;
|
||||
img_desc_3d.image_width = (size_t)n_q;
|
||||
img_desc_3d.image_height = (size_t)n_batch * (size_t)n_head;
|
||||
img_desc_3d.image_depth = (size_t)d_head_q / 4;
|
||||
cl_mem img_q_wmm = NULL;
|
||||
img_q_wmm = clCreateImage(context, CL_MEM_READ_WRITE, &img_fmt_3d, &img_desc_3d, NULL, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
{
|
||||
cl_kernel repack_q = backend_ctx->fa.kernel_repack_q_for_wmm;
|
||||
CL_CHECK(clSetKernelArg(repack_q, 0, sizeof(cl_mem), &mem_matrixQ));
|
||||
CL_CHECK(clSetKernelArg(repack_q, 1, sizeof(cl_ulong), &q_nb1));
|
||||
CL_CHECK(clSetKernelArg(repack_q, 2, sizeof(cl_ulong), &q_nb2));
|
||||
CL_CHECK(clSetKernelArg(repack_q, 3, sizeof(cl_ulong), &q_nb3));
|
||||
CL_CHECK(clSetKernelArg(repack_q, 4, sizeof(int), &n_head));
|
||||
CL_CHECK(clSetKernelArg(repack_q, 5, sizeof(cl_mem), &img_q_wmm));
|
||||
|
||||
size_t repack_q_gws[3] = {(size_t)d_head_q / 4, (size_t)n_q, (size_t)n_batch * (size_t)n_head};
|
||||
backend_ctx->enqueue_ndrange_kernel(repack_q, 3, repack_q_gws, NULL, dst);
|
||||
}
|
||||
|
||||
// K image: columns, row groups, KV heads.
|
||||
const size_t n_kv_row4 = ((size_t)n_kv + 3) / 4;
|
||||
|
||||
memset(&img_desc_3d, 0, sizeof(img_desc_3d));
|
||||
img_desc_3d.image_type = CL_MEM_OBJECT_IMAGE3D;
|
||||
img_desc_3d.image_width = (size_t)d_head_q;
|
||||
img_desc_3d.image_height = n_kv_row4;
|
||||
img_desc_3d.image_depth = (size_t)n_batch * (size_t)n_head_kv;
|
||||
cl_mem img_k_wmm = NULL;
|
||||
img_k_wmm = clCreateImage(context, CL_MEM_READ_WRITE, &img_fmt_3d, &img_desc_3d, NULL, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
{
|
||||
cl_kernel repack_k = backend_ctx->fa.kernel_repack_k_for_wmm;
|
||||
CL_CHECK(clSetKernelArg(repack_k, 0, sizeof(cl_mem), &mem_matrixK));
|
||||
CL_CHECK(clSetKernelArg(repack_k, 1, sizeof(cl_ulong), &k_nb1));
|
||||
CL_CHECK(clSetKernelArg(repack_k, 2, sizeof(cl_ulong), &k_nb2));
|
||||
CL_CHECK(clSetKernelArg(repack_k, 3, sizeof(cl_ulong), &k_nb3));
|
||||
CL_CHECK(clSetKernelArg(repack_k, 4, sizeof(int), &n_head_kv));
|
||||
CL_CHECK(clSetKernelArg(repack_k, 5, sizeof(int), &n_kv));
|
||||
CL_CHECK(clSetKernelArg(repack_k, 6, sizeof(cl_mem), &img_k_wmm));
|
||||
|
||||
size_t repack_k_gws[3] = {(size_t)d_head_q, n_kv_row4, (size_t)n_batch * (size_t)n_head_kv};
|
||||
backend_ctx->enqueue_ndrange_kernel(repack_k, 3, repack_k_gws, NULL, dst);
|
||||
}
|
||||
|
||||
// V image: kv-rows (contracted), packed head-dim groups, KV heads.
|
||||
memset(&img_desc_3d, 0, sizeof(img_desc_3d));
|
||||
img_desc_3d.image_type = CL_MEM_OBJECT_IMAGE3D;
|
||||
img_desc_3d.image_width = (size_t)n_kv;
|
||||
img_desc_3d.image_height = (size_t)d_head_v / 4;
|
||||
img_desc_3d.image_depth = (size_t)n_batch * (size_t)n_head_kv;
|
||||
cl_mem img_v_wmm = NULL;
|
||||
img_v_wmm = clCreateImage(context, CL_MEM_READ_WRITE, &img_fmt_3d, &img_desc_3d, NULL, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
{
|
||||
cl_kernel repack_v = backend_ctx->fa.kernel_repack_v_for_wmm;
|
||||
CL_CHECK(clSetKernelArg(repack_v, 0, sizeof(cl_mem), &mem_matrixV));
|
||||
CL_CHECK(clSetKernelArg(repack_v, 1, sizeof(cl_ulong), &v_nb1));
|
||||
CL_CHECK(clSetKernelArg(repack_v, 2, sizeof(cl_ulong), &v_nb2));
|
||||
CL_CHECK(clSetKernelArg(repack_v, 3, sizeof(cl_ulong), &v_nb3));
|
||||
CL_CHECK(clSetKernelArg(repack_v, 4, sizeof(int), &n_head_kv));
|
||||
CL_CHECK(clSetKernelArg(repack_v, 5, sizeof(cl_mem), &img_v_wmm));
|
||||
|
||||
size_t repack_v_gws[3] = {(size_t)d_head_v / 4, (size_t)n_kv, (size_t)n_batch * (size_t)n_head_kv};
|
||||
backend_ctx->enqueue_ndrange_kernel(repack_v, 3, repack_v_gws, NULL, dst);
|
||||
}
|
||||
|
||||
cl_int enable_mask = (extra_mask) ? 1 : 0;
|
||||
mask_buffer = extra_mask ? mem_tex_matrixMask_1dbuf : mem_tex_mask_fallback_1dbuf;
|
||||
|
||||
cl_mem mem_sinksBuf = NULL;
|
||||
cl_mem mem_tex_sinks_1dbuf = NULL;
|
||||
cl_int enable_sinks = (sinks_buffer != NULL) ? 1 : 0;
|
||||
if (enable_sinks) {
|
||||
region.origin = offset_sinks;
|
||||
region.size = ggml_nbytes(sinks);
|
||||
mem_sinksBuf = clCreateSubBuffer(extra_sinks->data_device, CL_MEM_READ_ONLY, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err);
|
||||
CL_CHECK(err);
|
||||
|
||||
cl_image_format img_fmt_sinks_1d = { CL_R, CL_FLOAT };
|
||||
cl_image_desc img_desc_sinks_1d;
|
||||
memset(&img_desc_sinks_1d, 0, sizeof(img_desc_sinks_1d));
|
||||
img_desc_sinks_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
|
||||
img_desc_sinks_1d.image_width = (size_t)n_head;
|
||||
img_desc_sinks_1d.buffer = mem_sinksBuf;
|
||||
mem_tex_sinks_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_sinks_1d, &img_desc_sinks_1d, NULL, &err);
|
||||
CL_CHECK(err);
|
||||
} else {
|
||||
// The image obj cannot be null so we back with buffer of size 1 and use matrixK to back because it always exists
|
||||
cl_image_format img_fmt_sinks_fallback = { CL_R, CL_FLOAT };
|
||||
cl_image_desc img_desc_sinks_fallback;
|
||||
memset(&img_desc_sinks_fallback, 0, sizeof(img_desc_sinks_fallback));
|
||||
img_desc_sinks_fallback.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
|
||||
img_desc_sinks_fallback.image_width = 1;
|
||||
img_desc_sinks_fallback.buffer = mem_matrixK;
|
||||
mem_tex_sinks_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_sinks_fallback, &img_desc_sinks_fallback, NULL, &err);
|
||||
CL_CHECK(err);
|
||||
}
|
||||
|
||||
cl_uint arg = 0;
|
||||
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &mem_tex_matrixO_1dbuf));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float), &scale));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &n_q));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &n_kv));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &is_causal));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &n_head));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &q_nb1));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &q_nb2));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &q_nb3));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &k_nb1));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &k_nb2));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &k_nb3));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &v_nb1));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &v_nb2));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &v_nb3));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &o_nb1));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &o_nb2));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &o_nb3));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float), &max_bias));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float), &m0));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float), &m1));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &n_head_log2_val));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float), &logit_softcap));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &n_head_kv));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &mask_buffer));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &enable_mask));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &mask_nb1_padded));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &mask_nb2_padded));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &mask_nb3_padded));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &mask_ne2));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &mask_ne3));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &mem_tex_sinks_1dbuf));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &enable_sinks));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &img_q_wmm));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &img_k_wmm));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &img_v_wmm));
|
||||
CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &d_head_q));
|
||||
|
||||
size_t global_work_size[3], local_work_size[3];
|
||||
|
||||
const int n_waves_v = d_head_q / 64;
|
||||
|
||||
local_work_size[0] = 64;
|
||||
local_work_size[1] = n_waves_v;
|
||||
local_work_size[2] = 1;
|
||||
|
||||
global_work_size[0] = 64;
|
||||
global_work_size[1] = ((n_q + 64 - 1) / 64) * n_waves_v;
|
||||
global_work_size[2] = n_batch * n_head;
|
||||
|
||||
backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst);
|
||||
|
||||
CL_CHECK(clReleaseMemObject(mem_tex_matrixO_1dbuf));
|
||||
CL_CHECK(clReleaseMemObject(img_q_wmm));
|
||||
CL_CHECK(clReleaseMemObject(img_k_wmm));
|
||||
CL_CHECK(clReleaseMemObject(img_v_wmm));
|
||||
|
||||
if (mem_tex_matrixMask_1dbuf) {
|
||||
CL_CHECK(clReleaseMemObject(mem_tex_matrixMask_1dbuf));
|
||||
}
|
||||
if (mem_matrixMask) {
|
||||
CL_CHECK(clReleaseMemObject(mem_matrixMask));
|
||||
}
|
||||
if (mem_matrixMask_padded) {
|
||||
CL_CHECK(clReleaseMemObject(mem_matrixMask_padded));
|
||||
}
|
||||
if (mem_tex_sinks_1dbuf) {
|
||||
CL_CHECK(clReleaseMemObject(mem_tex_sinks_1dbuf));
|
||||
}
|
||||
if (mem_sinksBuf) {
|
||||
CL_CHECK(clReleaseMemObject(mem_sinksBuf));
|
||||
}
|
||||
CL_CHECK(clReleaseMemObject(mem_matrixQ));
|
||||
CL_CHECK(clReleaseMemObject(mem_matrixK));
|
||||
CL_CHECK(clReleaseMemObject(mem_matrixV));
|
||||
CL_CHECK(clReleaseMemObject(mem_matrixO));
|
||||
}
|
||||
#endif // GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
|
||||
static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, const ggml_tensor * k, ggml_tensor * dst) {
|
||||
const ggml_tensor * v = dst->src[2];
|
||||
const ggml_tensor * mask = dst->src[3];
|
||||
@@ -17163,6 +17728,14 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co
|
||||
const bool is_q8_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q8_0 && v->type == GGML_TYPE_Q8_0;
|
||||
const bool is_q4_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q4_0 && v->type == GGML_TYPE_Q4_0;
|
||||
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
if (use_fa_bin_kernels_prefill(backend_ctx, q, k, v)) {
|
||||
// We support the prefill path of flash attn with a specialized d_head = 64/128/256
|
||||
ggml_cl_flash_attn_prefill_bin(backend, q, k, dst);
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
|
||||
if (is_f16) {
|
||||
ggml_opencl_ensure_fa_variant(backend_ctx, d_head_q, d_head_v, FA_VARIANT_F16);
|
||||
} else if (is_mixed) {
|
||||
@@ -21111,6 +21684,145 @@ static void ggml_cl_mul_mat_q4_k_f32_adreno(ggml_backend_t backend, const ggml_t
|
||||
#endif
|
||||
}
|
||||
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
static void ggml_cl_mul_mat_q6_K_f32_adreno_ila(ggml_backend_t backend, const ggml_tensor * src0,
|
||||
const ggml_tensor * src1, ggml_tensor * dst) {
|
||||
GGML_ASSERT(src0);
|
||||
GGML_ASSERT(src0->extra);
|
||||
GGML_ASSERT(src1);
|
||||
GGML_ASSERT(src1->extra);
|
||||
GGML_ASSERT(dst);
|
||||
GGML_ASSERT(dst->extra);
|
||||
|
||||
ggml_backend_opencl_context *backend_ctx = (ggml_backend_opencl_context *)backend->context;
|
||||
|
||||
ggml_tensor_extra_cl_q6_K * extra0_q6_K = (ggml_tensor_extra_cl_q6_K *)src0->extra;
|
||||
ggml_tensor_extra_cl * extra1 = (ggml_tensor_extra_cl *)src1->extra;
|
||||
ggml_tensor_extra_cl * extrad = (ggml_tensor_extra_cl *)dst->extra;
|
||||
|
||||
cl_ulong offset1 = extra1->offset + src1->view_offs;
|
||||
cl_ulong offsetd = extrad->offset + dst->view_offs;
|
||||
|
||||
const int ne00 = src0->ne[0];
|
||||
const int ne01 = src0->ne[1];
|
||||
|
||||
const int ne1 = dst->ne[1];
|
||||
|
||||
GGML_ASSERT(ne00 % ggml_blck_size(src0->type) == 0);
|
||||
|
||||
cl_context context = backend_ctx->context;
|
||||
cl_kernel kernel;
|
||||
|
||||
cl_int err;
|
||||
cl_buffer_region region;
|
||||
cl_image_format img_fmt;
|
||||
cl_image_desc img_desc;
|
||||
|
||||
const int M = ne01;
|
||||
const int N = ne1;
|
||||
const int K = ne00;
|
||||
|
||||
if (ne1 == 1) {
|
||||
cl_mem b_sub_buf = nullptr;
|
||||
cl_mem b_img = nullptr;
|
||||
|
||||
region.origin = offset1;
|
||||
region.size = (size_t)K * N * sizeof(float);
|
||||
CL_CHECK((b_sub_buf = clCreateSubBuffer(extra1->data_device, 0, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err), err));
|
||||
|
||||
img_fmt = { CL_RGBA, CL_FLOAT };
|
||||
memset(&img_desc, 0, sizeof(img_desc));
|
||||
img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
|
||||
img_desc.image_width = (size_t)K * N / 4;
|
||||
img_desc.buffer = b_sub_buf;
|
||||
CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err));
|
||||
|
||||
kernel = backend_ctx->kernel_gemv_noshuffle_q6_k_f32_32b_trans;
|
||||
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra0_q6_K->ql_img));
|
||||
CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q6_K->qh_img));
|
||||
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra0_q6_K->s));
|
||||
CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra0_q6_K->d));
|
||||
CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &b_img));
|
||||
CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_mem), &extrad->data_device));
|
||||
CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_ulong), &offsetd));
|
||||
CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_int), &ne00));
|
||||
CL_CHECK(clSetKernelArg(kernel, 8, sizeof(cl_int), &ne01));
|
||||
|
||||
size_t local_work_size[3] = { 64, 8, 1 };
|
||||
size_t global_work_size[3] = { (size_t)ne01, 8, 1 };
|
||||
backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst);
|
||||
|
||||
CL_CHECK(clReleaseMemObject(b_img));
|
||||
CL_CHECK(clReleaseMemObject(b_sub_buf));
|
||||
} else {
|
||||
const int gemm_tile_n = 64;
|
||||
int N_pad = CEIL_DIV(N, gemm_tile_n) * gemm_tile_n;
|
||||
|
||||
cl_mem b_sub_buf = nullptr;
|
||||
cl_mem b_padded = nullptr;
|
||||
cl_mem b_buf = nullptr;
|
||||
if (N_pad == N) {
|
||||
region.origin = offset1;
|
||||
region.size = (size_t)K * N * sizeof(float);
|
||||
CL_CHECK((b_sub_buf = clCreateSubBuffer(extra1->data_device, 0, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err), err));
|
||||
b_buf = b_sub_buf;
|
||||
} else {
|
||||
CL_CHECK((b_padded = clCreateBuffer(context, CL_MEM_READ_WRITE, (size_t)K * N_pad * sizeof(float), NULL, &err), err));
|
||||
const float zero = 0.0f;
|
||||
CL_CHECK(clEnqueueFillBuffer(backend_ctx->queue, b_padded, &zero, sizeof(zero), 0, (size_t)K * N_pad * sizeof(float), 0, NULL, NULL));
|
||||
CL_CHECK(clEnqueueCopyBuffer(backend_ctx->queue, extra1->data_device, b_padded, offset1, 0, (size_t)K * N * sizeof(float), 0, NULL, NULL));
|
||||
b_buf = b_padded;
|
||||
}
|
||||
|
||||
img_fmt = { CL_R, CL_FLOAT };
|
||||
memset(&img_desc, 0, sizeof(img_desc));
|
||||
img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
|
||||
img_desc.image_width = (size_t)K * N_pad;
|
||||
img_desc.buffer = b_buf;
|
||||
cl_mem b_img;
|
||||
CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err));
|
||||
|
||||
region.origin = offsetd;
|
||||
region.size = (size_t)M * N * sizeof(float);
|
||||
cl_mem d_sub_buf;
|
||||
CL_CHECK((d_sub_buf = clCreateSubBuffer(extrad->data_device, 0, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err), err));
|
||||
img_fmt = { CL_R, CL_FLOAT };
|
||||
memset(&img_desc, 0, sizeof(img_desc));
|
||||
img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
|
||||
img_desc.image_width = (size_t)M * N;
|
||||
img_desc.buffer = d_sub_buf;
|
||||
cl_mem d_img;
|
||||
CL_CHECK((d_img = clCreateImage(context, CL_MEM_WRITE_ONLY, &img_fmt, &img_desc, NULL, &err), err));
|
||||
|
||||
kernel = backend_ctx->kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8_bin;
|
||||
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra0_q6_K->ql_img));
|
||||
CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q6_K->qh));
|
||||
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra0_q6_K->s));
|
||||
CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra0_q6_K->d));
|
||||
CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &b_img));
|
||||
CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_mem), &d_img));
|
||||
CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_uint), &ne00));
|
||||
CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_uint), &ne01));
|
||||
CL_CHECK(clSetKernelArg(kernel, 8, sizeof(int), &N));
|
||||
|
||||
size_t local_work_size[3] = { 64, 2, 2 };
|
||||
size_t m_tiles = (size_t)CEIL_DIV(M, 64);
|
||||
size_t global_work_size[3] = { 64, m_tiles, (size_t)CEIL_DIV(N_pad, gemm_tile_n) };
|
||||
backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst);
|
||||
|
||||
CL_CHECK(clReleaseMemObject(b_img));
|
||||
if (b_sub_buf) {
|
||||
CL_CHECK(clReleaseMemObject(b_sub_buf));
|
||||
}
|
||||
if (b_padded) {
|
||||
CL_CHECK(clReleaseMemObject(b_padded));
|
||||
}
|
||||
CL_CHECK(clReleaseMemObject(d_img));
|
||||
CL_CHECK(clReleaseMemObject(d_sub_buf));
|
||||
}
|
||||
}
|
||||
#endif // GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
|
||||
static void ggml_cl_mul_mat_q6_K_f32_adreno(ggml_backend_t backend, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
GGML_ASSERT(src0);
|
||||
@@ -21159,6 +21871,20 @@ static void ggml_cl_mul_mat_q6_K_f32_adreno(ggml_backend_t backend, const ggml_t
|
||||
// (the #1 MTP bottleneck; mc3 above can't, it reads the noshuffle layout).
|
||||
const bool use_q6k_tiled_mc = q6k_mc3 && (ne1 == 3) && (ne01 >= 32768) && use_q6k_tiled(backend_ctx, src0);
|
||||
|
||||
const bool use_bin = use_q6_k_bin_kernels(backend_ctx, src0);
|
||||
|
||||
if (use_bin) {
|
||||
if (use_q6k_mc3 || use_q6k_tiled_mc) {
|
||||
static bool warned = false;
|
||||
if (!warned) {
|
||||
GGML_LOG_WARN("ggml_opencl: GGML_OPENCL_Q6K_MC3 is bypassed by Q6_K binary kernels\n");
|
||||
warned = true;
|
||||
}
|
||||
}
|
||||
ggml_cl_mul_mat_q6_K_f32_adreno_ila(backend, src0, src1, dst);
|
||||
return;
|
||||
}
|
||||
|
||||
if (ne1 == 1 || use_q6k_mc3 || use_q6k_tiled_mc) {
|
||||
cl_mem ql_img = nullptr;
|
||||
cl_mem qh_img = nullptr;
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
|
||||
|
||||
__kernel void kernel_repack_mask_for_wmm(
|
||||
const global half* mask_buf,
|
||||
const ulong mask_nb1,
|
||||
const ulong mask_nb2,
|
||||
const ulong mask_nb3,
|
||||
const int mask_ne2,
|
||||
global half* mask_buf_padded,
|
||||
const ulong mask_nb1_padded,
|
||||
const ulong mask_nb2_padded,
|
||||
const ulong mask_nb3_padded
|
||||
) {
|
||||
int col = get_global_id(0); // 0 .. n_kv
|
||||
int row = get_global_id(1); // 0 .. n_q
|
||||
int slice = get_global_id(2); // 0 .. (n_head * n_batch)
|
||||
|
||||
int head_idx = slice % mask_ne2;
|
||||
int batch_idx = slice / mask_ne2;
|
||||
|
||||
ulong src_off = (ulong)batch_idx * mask_nb3 + (ulong)head_idx * mask_nb2 + (ulong)row * mask_nb1;
|
||||
ulong dst_off = (ulong)batch_idx * mask_nb3_padded + (ulong)head_idx * mask_nb2_padded + (ulong)row * mask_nb1_padded;
|
||||
|
||||
mask_buf_padded[dst_off / 2 + col] = mask_buf[src_off / 2 + col];
|
||||
}
|
||||
|
||||
__kernel void kernel_repack_q_for_wmm(
|
||||
const global float* q_buf,
|
||||
const ulong q_nb1,
|
||||
const ulong q_nb2,
|
||||
const ulong q_nb3,
|
||||
const int n_head,
|
||||
__write_only image3d_t img_q_wmm
|
||||
) {
|
||||
int k4 = get_global_id(0);
|
||||
int row = get_global_id(1);
|
||||
int slice = get_global_id(2);
|
||||
int batch_idx = slice / n_head;
|
||||
int head_idx = slice % n_head;
|
||||
|
||||
|
||||
ulong elem_off = (batch_idx * q_nb3 + head_idx * q_nb2 + row * q_nb1) / 4 + (ulong)k4 * 4;
|
||||
float4 v = vload4(elem_off / 4, q_buf);
|
||||
|
||||
write_imageh(img_q_wmm, (int4)(row, slice, k4, 0), convert_half4(v));
|
||||
}
|
||||
|
||||
__kernel void kernel_repack_k_for_wmm(
|
||||
const global half* k_buf,
|
||||
const ulong k_nb1,
|
||||
const ulong k_nb2,
|
||||
const ulong k_nb3,
|
||||
const int n_head_kv,
|
||||
const int n_kv,
|
||||
__write_only image3d_t img_k_wmm
|
||||
) {
|
||||
int kk = get_global_id(0);
|
||||
int row4 = get_global_id(1);
|
||||
int slice = get_global_id(2);
|
||||
int batch_idx = slice / n_head_kv;
|
||||
int head_kv_idx = slice % n_head_kv;
|
||||
|
||||
ulong base = batch_idx * k_nb3 + head_kv_idx * k_nb2;
|
||||
int row0 = row4 * 4;
|
||||
half4 v;
|
||||
v.x = (row0 + 0 < n_kv) ? k_buf[(base + (ulong)(row0 + 0) * k_nb1) / 2 + kk] : (half)0;
|
||||
v.y = (row0 + 1 < n_kv) ? k_buf[(base + (ulong)(row0 + 1) * k_nb1) / 2 + kk] : (half)0;
|
||||
v.z = (row0 + 2 < n_kv) ? k_buf[(base + (ulong)(row0 + 2) * k_nb1) / 2 + kk] : (half)0;
|
||||
v.w = (row0 + 3 < n_kv) ? k_buf[(base + (ulong)(row0 + 3) * k_nb1) / 2 + kk] : (half)0;
|
||||
|
||||
write_imageh(img_k_wmm, (int4)(kk, row4, slice, 0), v);
|
||||
}
|
||||
|
||||
__kernel void kernel_repack_v_for_wmm(
|
||||
const global half* v_buf,
|
||||
const ulong v_nb1,
|
||||
const ulong v_nb2,
|
||||
const ulong v_nb3,
|
||||
const int n_head_kv,
|
||||
__write_only image3d_t img_v_wmm
|
||||
) {
|
||||
int hdim4 = get_global_id(0); // now fastest — walks contiguous memory
|
||||
int row = get_global_id(1);
|
||||
int slice = get_global_id(2);
|
||||
int batch_idx = slice / n_head_kv;
|
||||
int head_kv_idx = slice % n_head_kv;
|
||||
|
||||
ulong row_off = batch_idx * v_nb3 + head_kv_idx * v_nb2 + (ulong)row * v_nb1;
|
||||
half4 v = vload4((row_off / 2 + (ulong)hdim4 * 4) / 4, v_buf);
|
||||
|
||||
write_imageh(img_v_wmm, (int4)(row, hdim4, slice, 0), v);
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
|
||||
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
|
||||
#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable
|
||||
|
||||
#define QK_K 256
|
||||
#define N_SIMDGROUP 8
|
||||
#define SIMDGROUP_WIDTH 64
|
||||
|
||||
static inline float8 q6_k_to_fp32_packed8(ushort2 ql8, ushort qh8, float d_scale) {
|
||||
float8 fp32x8;
|
||||
fp32x8.s0 = ((float)(( ql8.s0 & 0x000F) | ((uint)((qh8 ) & 0x3) << 4)) - 32.f) * d_scale;
|
||||
fp32x8.s1 = ((float)((( ql8.s0 >> 4) & 0x000F) | ((uint)((qh8 >> 2) & 0x3) << 4)) - 32.f) * d_scale;
|
||||
fp32x8.s2 = ((float)((( ql8.s0 >> 8) & 0x000F) | ((uint)((qh8 >> 4) & 0x3) << 4)) - 32.f) * d_scale;
|
||||
fp32x8.s3 = ((float)((( ql8.s0 >> 12)& 0x000F) | ((uint)((qh8 >> 6) & 0x3) << 4)) - 32.f) * d_scale;
|
||||
fp32x8.s4 = ((float)(( ql8.s1 & 0x000F) | ((uint)((qh8 >> 8) & 0x3) << 4)) - 32.f) * d_scale;
|
||||
fp32x8.s5 = ((float)((( ql8.s1 >> 4) & 0x000F) | ((uint)((qh8 >>10) & 0x3) << 4)) - 32.f) * d_scale;
|
||||
fp32x8.s6 = ((float)((( ql8.s1 >> 8) & 0x000F) | ((uint)((qh8 >>12) & 0x3) << 4)) - 32.f) * d_scale;
|
||||
fp32x8.s7 = ((float)((( ql8.s1 >> 12)& 0x000F) | ((uint)((qh8 >>14) & 0x3) << 4)) - 32.f) * d_scale;
|
||||
return fp32x8;
|
||||
}
|
||||
|
||||
__attribute__((qcom_reqd_sub_group_size("half")))
|
||||
__kernel void kernel_gemv_noshuffle_q6_k_f32_32b_trans(
|
||||
__read_only image1d_buffer_t src0_ql,
|
||||
__read_only image1d_buffer_t src0_qh,
|
||||
__global char * src0_s,
|
||||
__global half * src0_d,
|
||||
__read_only image1d_buffer_t src1,
|
||||
__global float * dst,
|
||||
ulong offsetd,
|
||||
int ne00,
|
||||
int ne01
|
||||
) {
|
||||
uint i01 = get_global_id(0);
|
||||
uint sgid = get_local_id(1);
|
||||
uint slid = get_sub_group_local_id();
|
||||
|
||||
int num_superblocks = ne00 / QK_K;
|
||||
int num_subblocks = ne00 / 32; // 2 sub-blocks of 16 processed per iter below
|
||||
int scales_per_row = num_superblocks * 16;
|
||||
|
||||
__private float sum = 0.0f;
|
||||
|
||||
// Loop over 32-element groups (2 sub-blocks of 16 each), N_SIMDGROUP groups per iter.
|
||||
for (uint ib = sgid; ib < num_subblocks; ib += N_SIMDGROUP) {
|
||||
uint sb = ib / 8; // super-block index
|
||||
uint j = ib % 8; // 32-element group within super-block (0..7)
|
||||
|
||||
// Load d for this super-block.
|
||||
half d_val = src0_d[sb * ne01 + i01];
|
||||
|
||||
// Load 2 sub-block scales (int8), one per 16 elements.
|
||||
global const char * sc = src0_s + i01 * scales_per_row + sb * 16;
|
||||
float scale0 = (float)d_val * (float)sc[j * 2];
|
||||
float scale1 = (float)d_val * (float)sc[j * 2 + 1];
|
||||
|
||||
// Load 4 uints of ql (32 elements, 4-bit each = 128 bits), column-major stride ne01.
|
||||
uint ql_base = (ib * 4) * ne01 + i01;
|
||||
uint4 regQL;
|
||||
regQL.s0 = read_imageui(src0_ql, ql_base).x;
|
||||
regQL.s1 = read_imageui(src0_ql, ql_base + ne01).x;
|
||||
regQL.s2 = read_imageui(src0_ql, ql_base + ne01 * 2).x;
|
||||
regQL.s3 = read_imageui(src0_ql, ql_base + ne01 * 3).x;
|
||||
|
||||
// Load 2 uints of qh (32 elements, 2-bit each = 64 bits), column-major stride ne01.
|
||||
uint qh_base = (ib * 2) * ne01 + i01;
|
||||
uint2 regQH;
|
||||
regQH.s0 = read_imageui(src0_qh, qh_base).x;
|
||||
regQH.s1 = read_imageui(src0_qh, qh_base + ne01).x;
|
||||
|
||||
// Load activations: 32 floats = 8 float4s.
|
||||
uint y_offset = ib * 8;
|
||||
|
||||
float4 y_local = (slid < 8) ? read_imagef(src1, (y_offset + slid)) : (float4)0.0f;
|
||||
float4 y0 = sub_group_broadcast(y_local, 0);
|
||||
float4 y1 = sub_group_broadcast(y_local, 1);
|
||||
float4 y2 = sub_group_broadcast(y_local, 2);
|
||||
float4 y3 = sub_group_broadcast(y_local, 3);
|
||||
float4 y4v = sub_group_broadcast(y_local, 4);
|
||||
float4 y5 = sub_group_broadcast(y_local, 5);
|
||||
float4 y6 = sub_group_broadcast(y_local, 6);
|
||||
float4 y7 = sub_group_broadcast(y_local, 7);
|
||||
|
||||
// Dequantize elements 0..7 (scale0).
|
||||
float8 fp32x8 = q6_k_to_fp32_packed8(as_ushort2(regQL.s0), (ushort)(regQH.s0 & 0xFFFF), scale0);
|
||||
|
||||
float4 acc = y0 * fp32x8.lo;
|
||||
acc += y1 * fp32x8.hi;
|
||||
|
||||
// Dequantize elements 8..15 (scale0).
|
||||
fp32x8 = q6_k_to_fp32_packed8(as_ushort2(regQL.s1), (ushort)(regQH.s0 >> 16), scale0);
|
||||
|
||||
acc += y2 * fp32x8.lo;
|
||||
acc += y3 * fp32x8.hi;
|
||||
|
||||
// Dequantize elements 16..23 (scale1).
|
||||
fp32x8 = q6_k_to_fp32_packed8(as_ushort2(regQL.s2), (ushort)(regQH.s1 & 0xFFFF), scale1);
|
||||
|
||||
acc += y4v * fp32x8.lo;
|
||||
acc += y5 * fp32x8.hi;
|
||||
|
||||
// Dequantize elements 24..31 (scale1).
|
||||
fp32x8 = q6_k_to_fp32_packed8(as_ushort2(regQL.s3), (ushort)(regQH.s1 >> 16), scale1);
|
||||
|
||||
acc += y6 * fp32x8.lo;
|
||||
acc += y7 * fp32x8.hi;
|
||||
|
||||
sum += ((acc.s0 + acc.s1) + (acc.s2 + acc.s3));
|
||||
}
|
||||
|
||||
// reduction in local memory, assumes #subgroups=4
|
||||
__local float reduceLM[SIMDGROUP_WIDTH * (N_SIMDGROUP - 1)];
|
||||
if (sgid > 0) {
|
||||
reduceLM[SIMDGROUP_WIDTH * (sgid - 1) + slid] = sum;
|
||||
}
|
||||
barrier(CLK_LOCAL_MEM_FENCE);
|
||||
if (sgid == 0) {
|
||||
for (uint i = 0; i < N_SIMDGROUP - 1; ++i) {
|
||||
sum += reduceLM[SIMDGROUP_WIDTH * i + slid];
|
||||
}
|
||||
}
|
||||
|
||||
// 1 output per thread in subgroup 0
|
||||
if (sgid == 0) {
|
||||
dst = dst + (offsetd >> 2);
|
||||
dst[i01] = sum;
|
||||
}
|
||||
}
|
||||
@@ -245,7 +245,7 @@ void GgmlOvDecoder::set_input_output() {
|
||||
if (src->op == GGML_OP_VIEW) {
|
||||
// Traverse upward through nested VIEW operations
|
||||
std::remove_reference_t<decltype(current_node_info.node_inputs_views[src_name])> view_chain;
|
||||
auto current = src;
|
||||
auto * current = src;
|
||||
|
||||
while (current != nullptr) {
|
||||
auto current_name = get_tensor_ov_name(m_cgraph, current);
|
||||
@@ -612,9 +612,8 @@ std::pair<ModelParams, ComputeParams> GgmlOvDecoder::compute_llm_params(ggml_cgr
|
||||
if (node->src[1]->view_src != nullptr) {
|
||||
if (node->src[3] != nullptr) {
|
||||
return 4; // decoder self-attention
|
||||
} else {
|
||||
return 5; // cross-attention or encoder self-attention
|
||||
};
|
||||
}
|
||||
return 5; // cross-attention or encoder self-attention
|
||||
}
|
||||
break;
|
||||
default:
|
||||
@@ -736,8 +735,7 @@ std::pair<ModelParams, ComputeParams> GgmlOvDecoder::compute_llm_params(ggml_cgr
|
||||
|
||||
bool rope_seen = false;
|
||||
for (int i = 0; i < cgraph->n_nodes; i++) {
|
||||
auto * node = cgraph->nodes[i];
|
||||
std::string name = std::string(node->name);
|
||||
ggml_tensor * node = cgraph->nodes[i];
|
||||
const int attention_pattern_case = get_attention_pattern_case(node);
|
||||
if (attention_pattern_case != -1) {
|
||||
ggml_tensor * cache_k_permute = nullptr;
|
||||
@@ -948,7 +946,6 @@ ov::PartialShape GgmlOvDecoder::get_graph_input_shape(const ggml_tensor * op,
|
||||
if (m_naive) {
|
||||
return input != nullptr ? ov::PartialShape{get_shape(input)} : ov::PartialShape{get_shape(op)};
|
||||
}
|
||||
auto name = std::string(input->name);
|
||||
ov::PartialShape input_shape;
|
||||
|
||||
if (is_inp_tok(input, op) || is_inp_pos(input, op)) {
|
||||
@@ -1474,7 +1471,7 @@ std::shared_ptr<ov::Node> GgmlOvDecoder::create_weight_node(ggml_tensor * tensor
|
||||
void GgmlOvDecoder::dump_cgraph(const ggml_cgraph * cgraph, std::string & filename) {
|
||||
std::ofstream file(filename);
|
||||
if (!file.is_open()) {
|
||||
std::cerr << "Failed to open file" << std::endl;
|
||||
std::cerr << "Failed to open file" << '\n';
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -1580,11 +1577,11 @@ void print_tensor_address_map(const ggml_cgraph * cgraph) {
|
||||
}
|
||||
}
|
||||
for (const auto & pair : address_map) {
|
||||
std::cout << "Address: " << pair.first << std::endl;
|
||||
std::cout << "Address: " << pair.first << '\n';
|
||||
for (const auto & name : pair.second) {
|
||||
std::cout << name << " ; ";
|
||||
}
|
||||
std::cout << std::endl << std::endl;
|
||||
std::cout << "\n\n";
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2226,7 +2223,7 @@ void GgmlOvDecoder::compute_node_dynamic_dims() {
|
||||
std::cout << ", ";
|
||||
}
|
||||
}
|
||||
std::cout << "]" << std::endl;
|
||||
std::cout << "]" << '\n';
|
||||
// print the src name & shape with the dynamic dim for debugging
|
||||
for (int j = 0; j < GGML_MAX_SRC; j++) {
|
||||
ggml_tensor * src = node->src[j];
|
||||
@@ -2245,9 +2242,9 @@ void GgmlOvDecoder::compute_node_dynamic_dims() {
|
||||
std::cout << ", ";
|
||||
}
|
||||
}
|
||||
std::cout << "]" << std::endl;
|
||||
std::cout << "]" << '\n';
|
||||
}
|
||||
std::cout << std::endl;
|
||||
std::cout << '\n';
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -354,41 +354,41 @@ public:
|
||||
|
||||
void update_io(ggml_cgraph * cgraph);
|
||||
|
||||
inline static bool is_inp_tok(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
static bool is_inp_tok(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
return op->op == GGML_OP_GET_ROWS && tensor == op->src[1] && op->src[0]->op == GGML_OP_NONE;
|
||||
}
|
||||
|
||||
inline static bool is_inp_pos(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
static bool is_inp_pos(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
return op->op == GGML_OP_ROPE && tensor == op->src[1];
|
||||
}
|
||||
|
||||
// IMROPE packs 4 stacked position planes (t/h/w/e) into inp_pos, each of length
|
||||
// n_tokens; other modes carry a single position per token.
|
||||
inline static int get_inp_pos_n_planes(const ggml_tensor * op) {
|
||||
static int get_inp_pos_n_planes(const ggml_tensor * op) {
|
||||
return op->op_params[2] == GGML_ROPE_TYPE_IMROPE ? 4 : 1;
|
||||
}
|
||||
|
||||
inline static bool is_inp_emb(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
static bool is_inp_emb(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
return tensor->op == GGML_OP_GET_ROWS && op->op == GGML_OP_RMS_NORM;
|
||||
}
|
||||
|
||||
inline static bool is_inp_mask(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
static bool is_inp_mask(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
return op->op == GGML_OP_CPY || (op->op == GGML_OP_FLASH_ATTN_EXT && tensor == op->src[3]) ||
|
||||
(op->op == GGML_OP_SOFT_MAX && tensor == op->src[1]);
|
||||
}
|
||||
|
||||
inline static bool is_inp_mean(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
static bool is_inp_mean(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
return op->op == GGML_OP_MUL_MAT && tensor == op->src[1] && tensor->op == GGML_OP_NONE &&
|
||||
(tensor->flags & GGML_TENSOR_FLAG_INPUT) && tensor->type == GGML_TYPE_F32 &&
|
||||
op->src[0] != nullptr && op->src[0]->op != GGML_OP_NONE;
|
||||
}
|
||||
|
||||
inline static bool is_rope_freqs_weight(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
static bool is_rope_freqs_weight(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
return op->op == GGML_OP_ROPE && tensor == op->src[2];
|
||||
}
|
||||
|
||||
// also returns true for cache_s and cache_r in SSM/DeltaNet models
|
||||
inline static bool is_kvcache(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
static bool is_kvcache(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
if (tensor == nullptr) {
|
||||
return false;
|
||||
}
|
||||
@@ -396,14 +396,14 @@ public:
|
||||
(op != nullptr && op->op == GGML_OP_SET_ROWS && op->src[2] == tensor);
|
||||
}
|
||||
|
||||
inline static bool is_conv_state_writeback(const ggml_tensor * node) {
|
||||
static bool is_conv_state_writeback(const ggml_tensor * node) {
|
||||
return node->op == GGML_OP_CPY && node->view_src != nullptr && is_kvcache(node->view_src, nullptr) &&
|
||||
node->src[0] != nullptr && node->src[0]->op == GGML_OP_VIEW && node->src[0]->src[0] != nullptr &&
|
||||
node->src[0]->src[0]->op == GGML_OP_CONCAT && node->src[1] != nullptr &&
|
||||
node->src[1]->op == GGML_OP_VIEW && node->src[1]->view_src == node->view_src;
|
||||
}
|
||||
|
||||
inline static bool is_kv_idx(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
static bool is_kv_idx(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
return op->op == GGML_OP_SET_ROWS && op->src[1] == tensor;
|
||||
}
|
||||
|
||||
@@ -411,13 +411,13 @@ public:
|
||||
return m_model_params.swa_mask != nullptr && tensor == m_model_params.swa_mask;
|
||||
}
|
||||
|
||||
inline static bool is_output_idx(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
static bool is_output_idx(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
return op->op == GGML_OP_GET_ROWS && tensor == op->src[1] && op->src[0]->op != GGML_OP_NONE &&
|
||||
op->src[1]->op == GGML_OP_NONE;
|
||||
}
|
||||
|
||||
// the state permutation index input used in SSM/DeltaNet models (inp->s_copy in llama-graph.cpp)
|
||||
inline static bool is_inp_s_copy(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
static bool is_inp_s_copy(const ggml_tensor * tensor, const ggml_tensor * op) {
|
||||
return op->op == GGML_OP_GET_ROWS && tensor == op->src[1] &&
|
||||
op->src[0]->buffer->usage == GGML_BACKEND_BUFFER_USAGE_ANY;
|
||||
}
|
||||
@@ -481,5 +481,3 @@ private:
|
||||
};
|
||||
|
||||
void print_tensor_address_map(const ggml_cgraph * cgraph);
|
||||
|
||||
std::optional<int> extract_layer_from_name(const std::string & name);
|
||||
|
||||
@@ -472,10 +472,6 @@ ggml_openvino_extracted_layout ggml_openvino_get_extracted_layout(const ggml_ten
|
||||
|
||||
switch (tensor->type) {
|
||||
case GGML_TYPE_MXFP4:
|
||||
layout.is_u4 = true;
|
||||
layout.is_symmetric = true;
|
||||
break;
|
||||
|
||||
case GGML_TYPE_Q4_0:
|
||||
layout.is_u4 = true;
|
||||
layout.is_symmetric = true;
|
||||
|
||||
@@ -28,12 +28,7 @@
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
#ifndef _WIN32
|
||||
# include <sys/mman.h>
|
||||
# include <unistd.h>
|
||||
#endif
|
||||
|
||||
#if defined(_WIN32)
|
||||
#ifdef _WIN32
|
||||
# define WIN32_LEAN_AND_MEAN
|
||||
# ifndef NOMINMAX
|
||||
# define NOMINMAX
|
||||
@@ -61,6 +56,7 @@
|
||||
// - CPU repack buffer: tensor->extra stores tensor_traits with repacked data
|
||||
// =====================================================
|
||||
|
||||
namespace {
|
||||
// Buffer context that manages per-tensor allocations (no contiguous buffer for weights)
|
||||
struct ggml_backend_openvino_buffer_context {
|
||||
int device;
|
||||
@@ -199,6 +195,7 @@ struct ggml_backend_openvino_buffer_type_context {
|
||||
int device;
|
||||
std::string name;
|
||||
};
|
||||
} // namespace
|
||||
|
||||
// =====================================================
|
||||
// Host weight-buffer release (GGML_OPENVINO_RELEASE_WEIGHTS)
|
||||
@@ -258,14 +255,16 @@ void ggml_openvino_release_weight_buffers() {
|
||||
for (const auto & b : reg.buffers) {
|
||||
// Align down/up to page boundaries so madvise only drops whole pages
|
||||
// fully owned by this buffer.
|
||||
const long page = sysconf(_SC_PAGESIZE);
|
||||
uintptr_t start = reinterpret_cast<uintptr_t>(b.first);
|
||||
uintptr_t end = start + b.second;
|
||||
uintptr_t astart = (start + page - 1) & ~(uintptr_t) (page - 1);
|
||||
uintptr_t aend = end & ~(uintptr_t) (page - 1);
|
||||
if (aend > astart) {
|
||||
if (madvise(reinterpret_cast<void *>(astart), aend - astart, MADV_DONTNEED) == 0) {
|
||||
total += aend - astart;
|
||||
const size_t page = (size_t) sysconf(_SC_PAGESIZE);
|
||||
const uintptr_t ustart = reinterpret_cast<uintptr_t>(b.first);
|
||||
const size_t offset_to_page = (page - (ustart & (page - 1))) & (page - 1);
|
||||
if (b.second > offset_to_page) {
|
||||
const size_t aligned_len = (b.second - offset_to_page) & ~(page - 1);
|
||||
if (aligned_len > 0) {
|
||||
char * astart = static_cast<char *>(b.first) + offset_to_page;
|
||||
if (madvise(astart, aligned_len, MADV_DONTNEED) == 0) {
|
||||
total += aligned_len;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -876,11 +875,13 @@ GGML_BACKEND_API bool ggml_backend_is_openvino(ggml_backend_t backend) {
|
||||
return backend != NULL && ggml_guid_matches(backend->guid, ggml_backend_openvino_guid());
|
||||
}
|
||||
|
||||
namespace {
|
||||
struct ggml_backend_openvino_device_context {
|
||||
int device;
|
||||
std::string name;
|
||||
std::string description;
|
||||
};
|
||||
}
|
||||
|
||||
static const char * ggml_backend_openvino_device_get_name(ggml_backend_dev_t dev) {
|
||||
ggml_backend_openvino_device_context * ctx = (ggml_backend_openvino_device_context *) dev->context;
|
||||
@@ -1588,9 +1589,11 @@ static const struct ggml_backend_device_i ggml_backend_openvino_device_interface
|
||||
/* .event_synchronize = */ NULL,
|
||||
};
|
||||
|
||||
namespace {
|
||||
struct ggml_backend_openvino_reg_context {
|
||||
std::vector<ggml_backend_dev_t> devices;
|
||||
};
|
||||
}
|
||||
|
||||
static const char * ggml_backend_openvino_reg_get_name(ggml_backend_reg_t reg) {
|
||||
return GGML_OPENVINO_NAME;
|
||||
|
||||
@@ -34,6 +34,15 @@
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
// From <openvino>/src/common/transformations/include/transformations/utils/utils.hpp
|
||||
namespace ov::op::util {
|
||||
// From <openvino>/src/common/transformations/include/transformations/utils/utils.hpp
|
||||
bool get_single_value(const std::shared_ptr<ov::op::v0::Constant> & const_node,
|
||||
float & value,
|
||||
bool check_value_range = true);
|
||||
} // namespace ov::op::util
|
||||
|
||||
namespace {
|
||||
void unpack_32_4(const uint8_t * data, uint8_t * dst) {
|
||||
std::fill_n(dst, 16, 0);
|
||||
for (int j = 0; j < 16; ++j) {
|
||||
@@ -48,11 +57,11 @@ void unpack_32_4(const uint8_t * data, uint8_t * dst) {
|
||||
}
|
||||
}
|
||||
|
||||
static constexpr size_t MXFP4_BLOCK_SIZE = 32;
|
||||
static constexpr size_t MXFP4_BLOCK_QS_SIZE = MXFP4_BLOCK_SIZE / 2;
|
||||
static constexpr size_t MXFP4_BLOCK_BYTES = sizeof(uint8_t) + MXFP4_BLOCK_QS_SIZE;
|
||||
constexpr size_t MXFP4_BLOCK_SIZE = 32;
|
||||
constexpr size_t MXFP4_BLOCK_QS_SIZE = MXFP4_BLOCK_SIZE / 2;
|
||||
constexpr size_t MXFP4_BLOCK_BYTES = sizeof(uint8_t) + MXFP4_BLOCK_QS_SIZE;
|
||||
|
||||
static void pack_32_mxfp4_for_openvino(const uint8_t * data, uint8_t * dst) {
|
||||
void pack_32_mxfp4_for_openvino(const uint8_t * data, uint8_t * dst) {
|
||||
for (int j = 0; j < static_cast<int>(MXFP4_BLOCK_QS_SIZE); j += 2) {
|
||||
const uint8_t v0 = data[j] & 0x0F;
|
||||
const uint8_t v1 = (data[j + 1] & 0x0F) << 4;
|
||||
@@ -419,7 +428,7 @@ void extract_q6_k_data(const ggml_tensor * tensor,
|
||||
}
|
||||
}
|
||||
|
||||
static inline void get_scale_min_k4(int j, const uint8_t * q, uint8_t * d, uint8_t * m) {
|
||||
inline void get_scale_min_k4(int j, const uint8_t * q, uint8_t * d, uint8_t * m) {
|
||||
if (j < 4) {
|
||||
*d = q[j] & 63;
|
||||
*m = q[j + 4] & 63;
|
||||
@@ -514,9 +523,9 @@ void extract_q5_k_data(const ggml_tensor * tensor,
|
||||
ov::Output<ov::Node> make_int8_weights(ov::Tensor & weight,
|
||||
ov::Tensor & scales,
|
||||
ov::Tensor & zp,
|
||||
size_t group_size,
|
||||
bool use_bias,
|
||||
bool for_gather_matmul) {
|
||||
size_t group_size = GGML_QUANTIZATION_GROUP_SIZE,
|
||||
bool use_bias = false,
|
||||
bool for_gather_matmul = false) {
|
||||
ov::Shape orig_shape = weight.get_shape();
|
||||
bool is_signed = (weight.get_element_type() == ov::element::i8); // Symmetric: signed weights, no ZP
|
||||
|
||||
@@ -611,13 +620,24 @@ ov::Output<ov::Node> make_int8_weights(ov::Tensor & weight,
|
||||
return std::make_shared<ov::op::v0::Convert>(result, ov::element::f32);
|
||||
}
|
||||
|
||||
// If for_gather_matmul is true, the weight tensor may be N-D (e.g. 3D MoE expert weights
|
||||
// [n_expert, rows, cols]). The dequantization chain (Convert->[Subtract]->Multiply) is built as
|
||||
// usual but left in f16 (no final Convert to f32) -- ov::pass::MarkDequantization (registered in
|
||||
// translate_session.cpp) marks the chain so it survives model-build-time ConstantFolding -- see
|
||||
// make_int8_weights.cpp/make_int4_weights.cpp. mul_mat_id.cpp constructs ov::op::internal::GatherMatmul
|
||||
// directly from the resulting f16 dequant chain.
|
||||
//
|
||||
// When use_bias is true (explicitly, or implicitly because for_gather_matmul is true), the zp
|
||||
// tensor is expected to hold an exact f16 bias value (rather than a rounded integer zero point);
|
||||
// it is converted in place into an exact zero_point = -bias/scale and consumed via Subtract, not
|
||||
// Add, so the chain still matches OpenVINO's Convert->Subtract->Multiply decompression pattern.
|
||||
// See make_int8_weights for the meaning of for_gather_matmul.
|
||||
ov::Output<ov::Node> make_int4_weights(ov::Tensor & weight,
|
||||
ov::Tensor & scales,
|
||||
ov::Tensor & zp,
|
||||
size_t group_size,
|
||||
bool use_bias,
|
||||
bool for_gather_matmul) {
|
||||
size_t group_size = GGML_QUANTIZATION_GROUP_SIZE,
|
||||
bool use_bias = false,
|
||||
bool for_gather_matmul = false) {
|
||||
ov::Shape orig_weight_shape = weight.get_shape();
|
||||
bool is_signed = (weight.get_element_type() == ov::element::i4); // Symmetric: signed weights, no ZP
|
||||
|
||||
@@ -746,357 +766,6 @@ ov::Output<ov::Node> make_mxfp4_moe_packed_weights(ov::Tensor & weight) {
|
||||
return weights_node;
|
||||
}
|
||||
|
||||
// Extract quantized weights from tensor and create weight subgraph
|
||||
std::shared_ptr<ov::Node> extract_quantized_weights(const ggml_tensor * tensor,
|
||||
const void * data,
|
||||
ov::Tensor & weights,
|
||||
ov::Tensor & scales,
|
||||
ov::Tensor & zp,
|
||||
bool use_bias) {
|
||||
// Create a temporary tensor for extraction functions that read from tensor->data
|
||||
ggml_tensor temp_tensor = *tensor;
|
||||
temp_tensor.data = const_cast<void *>(data);
|
||||
|
||||
if (tensor->type == GGML_TYPE_MXFP4) {
|
||||
extract_mxfp4_data(&temp_tensor, weights, scales);
|
||||
auto result = make_mxfp4_weights(weights, scales).get_node_shared_ptr();
|
||||
result->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
|
||||
// Determine block size based on tensor type
|
||||
int64_t weights_per_block;
|
||||
bool is_u4;
|
||||
switch (tensor->type) {
|
||||
case GGML_TYPE_Q4_0:
|
||||
case GGML_TYPE_Q4_1:
|
||||
case GGML_TYPE_Q4_K:
|
||||
is_u4 = true;
|
||||
weights_per_block = 32;
|
||||
break;
|
||||
case GGML_TYPE_Q8_0:
|
||||
case GGML_TYPE_Q5_1:
|
||||
case GGML_TYPE_Q5_K:
|
||||
is_u4 = false;
|
||||
weights_per_block = 32;
|
||||
break;
|
||||
case GGML_TYPE_Q6_K:
|
||||
is_u4 = false;
|
||||
weights_per_block = 16;
|
||||
break;
|
||||
default:
|
||||
throw std::runtime_error("Unsupported quantized type for extraction: " +
|
||||
std::string(ggml_type_name(tensor->type)));
|
||||
}
|
||||
|
||||
// 3D MoE expert weights (for_gather_matmul) always use the exact f16 zero-point extraction
|
||||
// (see make_int8_weights/make_int4_weights) rather than the rounded integer zero point --
|
||||
// round(min/scale) error is what corrupts Q4_K/Q5_1 experts, and the f16-zp form still fuses
|
||||
// into GatherMatmulCompressed since it stays a Subtract, not an Add.
|
||||
const bool for_gather_matmul = tensor->ne[2] > 1;
|
||||
use_bias = use_bias || for_gather_matmul;
|
||||
|
||||
// Extract quantized data
|
||||
switch (tensor->type) {
|
||||
case GGML_TYPE_Q4_0:
|
||||
extract_q4_0_data(&temp_tensor, weights, scales, zp);
|
||||
break;
|
||||
case GGML_TYPE_Q4_1:
|
||||
extract_q4_1_data(&temp_tensor, weights, scales, zp, use_bias);
|
||||
break;
|
||||
case GGML_TYPE_Q4_K:
|
||||
extract_q4_k_data(&temp_tensor, weights, scales, zp, use_bias);
|
||||
break;
|
||||
case GGML_TYPE_Q5_1:
|
||||
extract_q5_1_data(&temp_tensor, weights, scales, zp, use_bias);
|
||||
break;
|
||||
case GGML_TYPE_Q8_0:
|
||||
extract_q8_0_data(&temp_tensor, weights, scales, zp);
|
||||
break;
|
||||
case GGML_TYPE_Q6_K:
|
||||
extract_q6_k_data(&temp_tensor, weights, scales, zp);
|
||||
break;
|
||||
case GGML_TYPE_Q5_K:
|
||||
extract_q5_k_data(&temp_tensor, weights, scales, zp, use_bias);
|
||||
break;
|
||||
default:
|
||||
throw std::runtime_error("Unsupported quantized type: " + std::string(ggml_type_name(tensor->type)));
|
||||
}
|
||||
|
||||
// Create the OpenVINO weight subgraph. 3D expert weights (MoE) are routed through the
|
||||
// GatherMatmul-oriented path: dequantized in f16, with constant folding disabled on the chain.
|
||||
ov::Output<ov::Node> weight_node;
|
||||
if (is_u4) {
|
||||
weight_node = make_int4_weights(weights, scales, zp, weights_per_block, use_bias, for_gather_matmul);
|
||||
} else {
|
||||
weight_node = make_int8_weights(weights, scales, zp, weights_per_block, use_bias, for_gather_matmul);
|
||||
}
|
||||
|
||||
auto result = weight_node.get_node_shared_ptr();
|
||||
result->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
|
||||
// Requantize weights to target format, writing to provided buffers
|
||||
std::shared_ptr<ov::Node> requantize_to_buffers(const ggml_tensor * tensor,
|
||||
const void * data,
|
||||
ExtraQuantType requant_type,
|
||||
int64_t block_size,
|
||||
ov::Tensor & weights,
|
||||
ov::Tensor & scales,
|
||||
ov::Tensor & zp) {
|
||||
int64_t n_elements = ggml_nelements(tensor);
|
||||
const int64_t ne0 = tensor->ne[0]; // elements per row
|
||||
const int64_t n_rows = n_elements / ne0;
|
||||
const auto * type_traits = ggml_get_type_traits(tensor->type);
|
||||
const size_t src_row_bytes = ggml_row_size(tensor->type, ne0);
|
||||
|
||||
bool is_u4 = (requant_type == ExtraQuantType::Q4_0_C || requant_type == ExtraQuantType::Q4_0_128 ||
|
||||
requant_type == ExtraQuantType::Q4_0_64 || requant_type == ExtraQuantType::Q4_1_64);
|
||||
|
||||
// Streaming dequant (opt-in via GGML_OPENVINO_REDUCE_COMPILE_MEM or
|
||||
// GGML_OPENVINO_MEMORY_OPTIMIZE): instead of
|
||||
// materializing the full n_elements F32 array (e.g. ~1 GB for token_embd), dequantize
|
||||
// a chunk of complete rows into a small scratch and quantize/convert it straight into
|
||||
// the output buffers, capping the transient F32 footprint at CHUNK_ROWS*ne0 floats.
|
||||
//
|
||||
// Only valid (and only used) for the Q8_0_C / Q8_1_C / F16 targets whose block size
|
||||
// divides a row (channel-wise _C uses block_size == ne0) so no target block straddles
|
||||
// a row boundary, and Q8/F16 have no cross-block packing. The u4 (Q4_0) path packs two
|
||||
// weights per byte with running zp ORs that assume a single whole-array call, so it is
|
||||
// never streamed. When the flag is off, behavior is identical to the original
|
||||
// full-materialization path.
|
||||
const bool stream_requant = ggml_openvino_reduce_compile_mem_enabled() && !is_u4 &&
|
||||
!(block_size > 0 && ne0 % block_size != 0);
|
||||
|
||||
if (!stream_requant) {
|
||||
// Full materialization (original behavior): dequantize the whole tensor to F32,
|
||||
// then convert/quantize in one call.
|
||||
std::vector<float> weights_f32(n_elements);
|
||||
type_traits->to_float(data, weights_f32.data(), n_elements);
|
||||
if (requant_type == ExtraQuantType::F16) {
|
||||
ggml_get_type_traits(GGML_TYPE_F16)->from_float_ref(weights_f32.data(), weights.data(), n_elements);
|
||||
auto result = std::make_shared<ov::op::v0::Constant>(weights);
|
||||
result->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
if (requant_type == ExtraQuantType::Q4_1_64) {
|
||||
quantize_q4_1_asym(weights_f32.data(), weights, scales, zp, n_elements, block_size);
|
||||
} else if (is_u4) {
|
||||
quantize_q4_0(weights_f32.data(), weights, scales, zp, n_elements, block_size);
|
||||
} else if (requant_type == ExtraQuantType::Q8_1_C) {
|
||||
quantize_q8_1(weights_f32.data(), weights, scales, zp, n_elements, block_size);
|
||||
} else {
|
||||
quantize_q8_0(weights_f32.data(), weights, scales, zp, n_elements, block_size);
|
||||
}
|
||||
} else {
|
||||
// Streaming path for Q8_0_C / Q8_1_C / F16 (covers token_embd, output.weight,
|
||||
// and per-layer Q6_K/Q5_K requant — the large transient cases).
|
||||
const int64_t CHUNK_ROWS = std::min<int64_t>(n_rows, 256);
|
||||
std::vector<float> scratch(CHUNK_ROWS * ne0);
|
||||
// F16 destination: 2 bytes/element, advanced per chunk by r0*ne0 elements.
|
||||
auto * f16_base = static_cast<uint8_t *>(weights.data());
|
||||
for (int64_t r0 = 0; r0 < n_rows; r0 += CHUNK_ROWS) {
|
||||
const int64_t rows = std::min(CHUNK_ROWS, n_rows - r0);
|
||||
const int64_t elems = rows * ne0;
|
||||
const auto * src = static_cast<const uint8_t *>(data) + r0 * src_row_bytes;
|
||||
type_traits->to_float(src, scratch.data(), elems);
|
||||
|
||||
if (requant_type == ExtraQuantType::F16) {
|
||||
ggml_get_type_traits(GGML_TYPE_F16)
|
||||
->from_float_ref(scratch.data(), f16_base + (r0 * ne0) * sizeof(uint16_t), elems);
|
||||
} else {
|
||||
const int64_t block_offset = (r0 * ne0) / block_size;
|
||||
if (requant_type == ExtraQuantType::Q8_1_C) {
|
||||
quantize_q8_1(scratch.data(), weights, scales, zp, elems, block_size, block_offset);
|
||||
} else {
|
||||
quantize_q8_0(scratch.data(), weights, scales, zp, elems, block_size, block_offset);
|
||||
}
|
||||
}
|
||||
}
|
||||
if (requant_type == ExtraQuantType::F16) {
|
||||
auto result = std::make_shared<ov::op::v0::Constant>(weights);
|
||||
result->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
// Create the OpenVINO weight subgraph
|
||||
ov::Output<ov::Node> weight_node;
|
||||
if (is_u4) {
|
||||
weight_node = make_int4_weights(weights, scales, zp, block_size);
|
||||
} else {
|
||||
weight_node = make_int8_weights(weights, scales, zp, block_size);
|
||||
}
|
||||
|
||||
auto result = weight_node.get_node_shared_ptr();
|
||||
result->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
|
||||
OvWeight process_weight_tensor(const ggml_tensor * tensor, const void * data, void * output_base_ptr, bool use_bias) {
|
||||
GGML_ASSERT(tensor != nullptr);
|
||||
GGML_ASSERT(data != nullptr);
|
||||
|
||||
OvWeight result;
|
||||
|
||||
// Get shape for weights: [rows, cols], or [n_expert, rows, cols] for 3D MoE expert weights.
|
||||
ov::Shape node_shape = (tensor->ne[2] > 1) ?
|
||||
ov::Shape{static_cast<size_t>(tensor->ne[2]), static_cast<size_t>(tensor->ne[1]),
|
||||
static_cast<size_t>(tensor->ne[0])} :
|
||||
ov::Shape{static_cast<size_t>(tensor->ne[1]), static_cast<size_t>(tensor->ne[0])};
|
||||
|
||||
// Handle F16/F32/BF16 weights
|
||||
if (tensor->type == GGML_TYPE_F32 || tensor->type == GGML_TYPE_F16 || tensor->type == GGML_TYPE_BF16) {
|
||||
ov::element::Type element_type;
|
||||
switch (tensor->type) {
|
||||
case GGML_TYPE_F32:
|
||||
element_type = ov::element::f32;
|
||||
break;
|
||||
case GGML_TYPE_F16:
|
||||
element_type = ov::element::f16;
|
||||
break;
|
||||
case GGML_TYPE_BF16:
|
||||
element_type = ov::element::bf16;
|
||||
break;
|
||||
default:
|
||||
OPENVINO_THROW("Unexpected tensor type in F16/F32/BF16 path");
|
||||
}
|
||||
|
||||
if (output_base_ptr && output_base_ptr != data) {
|
||||
// Using external buffer - copy data and create shared-memory constant
|
||||
size_t tensor_bytes = ggml_nbytes(tensor);
|
||||
memcpy(output_base_ptr, data, tensor_bytes);
|
||||
result.weights = ov::Tensor(element_type, node_shape, output_base_ptr);
|
||||
} else {
|
||||
result.weights = ov::Tensor(element_type, node_shape, data);
|
||||
}
|
||||
result.weight_node = std::make_shared<ov::op::v0::Constant>(result.weights);
|
||||
return result;
|
||||
}
|
||||
|
||||
// Handle quantized weights
|
||||
if (!ggml_is_quantized(tensor->type)) {
|
||||
OPENVINO_THROW("Unsupported weight tensor type: ", ggml_type_name(tensor->type));
|
||||
}
|
||||
|
||||
result.layout = ggml_openvino_get_extracted_layout(tensor, use_bias);
|
||||
const auto & layout = result.layout;
|
||||
if (layout.total_size == 0) {
|
||||
OPENVINO_THROW("Unsupported quantized type: ", ggml_type_name(tensor->type));
|
||||
}
|
||||
|
||||
// 3D MoE expert weights (for_gather_matmul) always use the exact f16 zero-point path (see
|
||||
// extract_quantized_weights) -- must be kept in sync with the "use_bias || for_gather_matmul"
|
||||
// check in ggml_openvino_get_extracted_layout, which sizes/offsets the zp slot accordingly.
|
||||
// Requantized tensors (layout.is_requant) are handled by requantize_to_buffers instead, whose
|
||||
// zp sizing/type is unaffected by for_gather_matmul, so they are excluded here.
|
||||
const bool for_gather_matmul = tensor->ne[2] > 1;
|
||||
const bool zp_is_f16 = !layout.is_requant && (use_bias || for_gather_matmul);
|
||||
|
||||
const bool is_3d_mxfp4_moe = tensor->type == GGML_TYPE_MXFP4 && (tensor->ne[2] > 1 || tensor->ne[3] > 1);
|
||||
if (is_3d_mxfp4_moe) {
|
||||
ov::Shape packed_shape = {static_cast<size_t>(tensor->ne[3]),
|
||||
static_cast<size_t>(tensor->ne[2]),
|
||||
static_cast<size_t>(tensor->ne[1]),
|
||||
static_cast<size_t>(tensor->ne[0] / MXFP4_BLOCK_SIZE),
|
||||
MXFP4_BLOCK_BYTES};
|
||||
const size_t tensor_bytes = ggml_nbytes(tensor);
|
||||
if (output_base_ptr) {
|
||||
auto * buf_base = static_cast<uint8_t *>(output_base_ptr);
|
||||
memcpy(buf_base + layout.weights_offset, data, tensor_bytes);
|
||||
result.weights = ov::Tensor(ov::element::u8, packed_shape, buf_base + layout.weights_offset);
|
||||
} else {
|
||||
result.weights = ov::Tensor(ov::element::u8, packed_shape);
|
||||
memcpy(result.weights.data(), data, tensor_bytes);
|
||||
}
|
||||
result.weight_node = make_mxfp4_moe_packed_weights(result.weights).get_node_shared_ptr();
|
||||
result.weight_node->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
|
||||
if (use_bias) {
|
||||
OPENVINO_ASSERT(!layout.is_requant,
|
||||
"use_bias is only used for test-backend-ops, which should not have requantization");
|
||||
// bias node will be created on the fly and not use backend buffer
|
||||
output_base_ptr = nullptr;
|
||||
}
|
||||
|
||||
// F16 requant path - no separate scales/zp needed in result
|
||||
if (layout.is_requant && layout.requant_type.has_value() && layout.requant_type.value() == ExtraQuantType::F16) {
|
||||
if (output_base_ptr) {
|
||||
result.weights = ov::Tensor(ov::element::f16, node_shape,
|
||||
static_cast<uint8_t *>(output_base_ptr) + layout.weights_offset);
|
||||
} else {
|
||||
result.weights = ov::Tensor(ov::element::f16, node_shape);
|
||||
}
|
||||
ov::Tensor dummy_scales, dummy_zp; // Not used for F16
|
||||
result.weight_node =
|
||||
requantize_to_buffers(tensor, data, ExtraQuantType::F16, 0, result.weights, dummy_scales, dummy_zp);
|
||||
return result;
|
||||
}
|
||||
|
||||
// Quantized path (normal extraction or quantized requant)
|
||||
// Create weight/scale/zp tensors - shared between both paths
|
||||
// For symmetric quantization, use signed types (i4/i8) and no ZP tensor
|
||||
ov::element::Type weight_type = tensor->type == GGML_TYPE_MXFP4 ?
|
||||
ov::element::f4e2m1 :
|
||||
(layout.is_symmetric ? (layout.is_u4 ? ov::element::i4 : ov::element::i8) :
|
||||
(layout.is_u4 ? ov::element::u4 : ov::element::u8));
|
||||
ov::Shape scale_shape = node_shape;
|
||||
scale_shape.back() /= layout.weights_per_block;
|
||||
|
||||
if (tensor->type == GGML_TYPE_MXFP4) {
|
||||
if (tensor->ne[2] == 1 && tensor->ne[3] == 1) {
|
||||
node_shape = {static_cast<size_t>(tensor->ne[1]), static_cast<size_t>(tensor->ne[0])};
|
||||
} else {
|
||||
node_shape.clear();
|
||||
for (int i = GGML_MAX_DIMS - 1; i >= 0; --i) {
|
||||
node_shape.push_back(static_cast<size_t>(tensor->ne[i]));
|
||||
}
|
||||
}
|
||||
|
||||
scale_shape = node_shape;
|
||||
scale_shape.back() /= layout.weights_per_block;
|
||||
}
|
||||
|
||||
if (output_base_ptr) {
|
||||
uint8_t * buf_base = static_cast<uint8_t *>(output_base_ptr);
|
||||
result.weights = ov::Tensor(weight_type, node_shape, buf_base + layout.weights_offset);
|
||||
const ov::element::Type scale_type = tensor->type == GGML_TYPE_MXFP4 ? ov::element::f8e8m0 : ov::element::f16;
|
||||
result.scales = ov::Tensor(scale_type, scale_shape, buf_base + layout.scales_offset);
|
||||
if (!layout.is_symmetric) {
|
||||
ov::element::Type zp_type =
|
||||
zp_is_f16 ? ov::element::f16 : (layout.is_u4 ? ov::element::u4 : ov::element::u8);
|
||||
result.zp = ov::Tensor(zp_type, scale_shape, buf_base + layout.zp_offset);
|
||||
}
|
||||
// else: result.zp remains default-constructed (empty) for symmetric
|
||||
} else {
|
||||
result.weights = ov::Tensor(weight_type, node_shape);
|
||||
const ov::element::Type scale_type = tensor->type == GGML_TYPE_MXFP4 ? ov::element::f8e8m0 : ov::element::f16;
|
||||
result.scales = ov::Tensor(scale_type, scale_shape);
|
||||
if (!layout.is_symmetric) {
|
||||
if (zp_is_f16) {
|
||||
result.zp = ov::Tensor(ov::element::f16, scale_shape);
|
||||
} else {
|
||||
ov::element::Type zp_type = layout.is_u4 ? ov::element::u4 : ov::element::u8;
|
||||
result.zp = ov::Tensor(zp_type, scale_shape);
|
||||
}
|
||||
}
|
||||
// else: result.zp remains default-constructed (empty) for symmetric
|
||||
}
|
||||
|
||||
if (layout.is_requant && layout.requant_type.has_value()) {
|
||||
result.weight_node = requantize_to_buffers(tensor, data, layout.requant_type.value(), layout.weights_per_block,
|
||||
result.weights, result.scales, result.zp);
|
||||
} else {
|
||||
result.weight_node =
|
||||
extract_quantized_weights(tensor, data, result.weights, result.scales, result.zp, use_bias);
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
void quantize_q4_0(const float * x,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
@@ -1252,7 +921,7 @@ void quantize_q8_0(const float * x,
|
||||
ov::Tensor & zp_arr,
|
||||
int64_t k,
|
||||
int64_t qk,
|
||||
int64_t block_offset) {
|
||||
int64_t block_offset = 0) {
|
||||
assert(k % qk == 0);
|
||||
const int nb = k / qk;
|
||||
|
||||
@@ -1308,7 +977,7 @@ void quantize_q8_1(const float * x,
|
||||
ov::Tensor & zp_arr,
|
||||
int64_t k,
|
||||
int64_t qk,
|
||||
int64_t block_offset) {
|
||||
int64_t block_offset = 0) {
|
||||
assert(k % qk == 0);
|
||||
const int nb = k / qk;
|
||||
|
||||
@@ -1339,3 +1008,366 @@ void quantize_q8_1(const float * x,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Extract quantized weights from tensor and create weight subgraph
|
||||
// If weights/scales/zp are provided (non-empty), uses them as output buffers
|
||||
// Otherwise allocates new ov::Tensors internally
|
||||
// Returns the weight node (make_int4_weights or make_int8_weights result)
|
||||
std::shared_ptr<ov::Node> extract_quantized_weights(const ggml_tensor * tensor,
|
||||
const void * data, // Source data pointer (may differ from tensor->data)
|
||||
ov::Tensor & weights,
|
||||
ov::Tensor & scales,
|
||||
ov::Tensor & zp,
|
||||
// Use an exact f16 zero point (vs. a rounded integer one); always
|
||||
// used for for_gather_matmul (3D MoE expert) weights regardless of
|
||||
// this flag, and also settable explicitly for test-backend-ops.
|
||||
bool use_bias = false) {
|
||||
// Create a temporary tensor for extraction functions that read from tensor->data
|
||||
ggml_tensor temp_tensor = *tensor;
|
||||
temp_tensor.data = const_cast<void *>(data);
|
||||
|
||||
if (tensor->type == GGML_TYPE_MXFP4) {
|
||||
extract_mxfp4_data(&temp_tensor, weights, scales);
|
||||
auto result = make_mxfp4_weights(weights, scales).get_node_shared_ptr();
|
||||
result->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
|
||||
// Determine block size based on tensor type
|
||||
int64_t weights_per_block;
|
||||
bool is_u4;
|
||||
switch (tensor->type) {
|
||||
case GGML_TYPE_Q4_0:
|
||||
case GGML_TYPE_Q4_1:
|
||||
case GGML_TYPE_Q4_K:
|
||||
is_u4 = true;
|
||||
weights_per_block = 32;
|
||||
break;
|
||||
case GGML_TYPE_Q8_0:
|
||||
case GGML_TYPE_Q5_1:
|
||||
case GGML_TYPE_Q5_K:
|
||||
is_u4 = false;
|
||||
weights_per_block = 32;
|
||||
break;
|
||||
case GGML_TYPE_Q6_K:
|
||||
is_u4 = false;
|
||||
weights_per_block = 16;
|
||||
break;
|
||||
default:
|
||||
throw std::runtime_error("Unsupported quantized type for extraction: " +
|
||||
std::string(ggml_type_name(tensor->type)));
|
||||
}
|
||||
|
||||
// 3D MoE expert weights (for_gather_matmul) always use the exact f16 zero-point extraction
|
||||
// (see make_int8_weights/make_int4_weights) rather than the rounded integer zero point --
|
||||
// round(min/scale) error is what corrupts Q4_K/Q5_1 experts, and the f16-zp form still fuses
|
||||
// into GatherMatmulCompressed since it stays a Subtract, not an Add.
|
||||
const bool for_gather_matmul = tensor->ne[2] > 1;
|
||||
use_bias = use_bias || for_gather_matmul;
|
||||
|
||||
// Extract quantized data
|
||||
switch (tensor->type) {
|
||||
case GGML_TYPE_Q4_0:
|
||||
extract_q4_0_data(&temp_tensor, weights, scales, zp);
|
||||
break;
|
||||
case GGML_TYPE_Q4_1:
|
||||
extract_q4_1_data(&temp_tensor, weights, scales, zp, use_bias);
|
||||
break;
|
||||
case GGML_TYPE_Q4_K:
|
||||
extract_q4_k_data(&temp_tensor, weights, scales, zp, use_bias);
|
||||
break;
|
||||
case GGML_TYPE_Q5_1:
|
||||
extract_q5_1_data(&temp_tensor, weights, scales, zp, use_bias);
|
||||
break;
|
||||
case GGML_TYPE_Q8_0:
|
||||
extract_q8_0_data(&temp_tensor, weights, scales, zp);
|
||||
break;
|
||||
case GGML_TYPE_Q6_K:
|
||||
extract_q6_k_data(&temp_tensor, weights, scales, zp);
|
||||
break;
|
||||
case GGML_TYPE_Q5_K:
|
||||
extract_q5_k_data(&temp_tensor, weights, scales, zp, use_bias);
|
||||
break;
|
||||
default:
|
||||
throw std::runtime_error("Unsupported quantized type: " + std::string(ggml_type_name(tensor->type)));
|
||||
}
|
||||
|
||||
// Create the OpenVINO weight subgraph. 3D expert weights (MoE) are routed through the
|
||||
// GatherMatmul-oriented path: dequantized in f16, with constant folding disabled on the chain.
|
||||
ov::Output<ov::Node> weight_node;
|
||||
if (is_u4) {
|
||||
weight_node = make_int4_weights(weights, scales, zp, weights_per_block, use_bias, for_gather_matmul);
|
||||
} else {
|
||||
weight_node = make_int8_weights(weights, scales, zp, weights_per_block, use_bias, for_gather_matmul);
|
||||
}
|
||||
|
||||
auto result = weight_node.get_node_shared_ptr();
|
||||
result->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
|
||||
// Requantize weights from tensor to target format, writing to provided buffers
|
||||
// For F16 target, only weights buffer is used (scales/zp ignored)
|
||||
// Returns the weight node
|
||||
std::shared_ptr<ov::Node> requantize_to_buffers(const ggml_tensor * tensor,
|
||||
const void * data, // Source data pointer
|
||||
ExtraQuantType requant_type,
|
||||
int64_t block_size,
|
||||
ov::Tensor & weights,
|
||||
ov::Tensor & scales,
|
||||
ov::Tensor & zp) {
|
||||
int64_t n_elements = ggml_nelements(tensor);
|
||||
const int64_t ne0 = tensor->ne[0]; // elements per row
|
||||
const int64_t n_rows = n_elements / ne0;
|
||||
const auto * type_traits = ggml_get_type_traits(tensor->type);
|
||||
const size_t src_row_bytes = ggml_row_size(tensor->type, ne0);
|
||||
|
||||
bool is_u4 = (requant_type == ExtraQuantType::Q4_0_C || requant_type == ExtraQuantType::Q4_0_128 ||
|
||||
requant_type == ExtraQuantType::Q4_0_64 || requant_type == ExtraQuantType::Q4_1_64);
|
||||
|
||||
// Streaming dequant (opt-in via GGML_OPENVINO_REDUCE_COMPILE_MEM or
|
||||
// GGML_OPENVINO_MEMORY_OPTIMIZE): instead of
|
||||
// materializing the full n_elements F32 array (e.g. ~1 GB for token_embd), dequantize
|
||||
// a chunk of complete rows into a small scratch and quantize/convert it straight into
|
||||
// the output buffers, capping the transient F32 footprint at CHUNK_ROWS*ne0 floats.
|
||||
//
|
||||
// Only valid (and only used) for the Q8_0_C / Q8_1_C / F16 targets whose block size
|
||||
// divides a row (channel-wise _C uses block_size == ne0) so no target block straddles
|
||||
// a row boundary, and Q8/F16 have no cross-block packing. The u4 (Q4_0) path packs two
|
||||
// weights per byte with running zp ORs that assume a single whole-array call, so it is
|
||||
// never streamed. When the flag is off, behavior is identical to the original
|
||||
// full-materialization path.
|
||||
const bool stream_requant = ggml_openvino_reduce_compile_mem_enabled() && !is_u4 &&
|
||||
!(block_size > 0 && ne0 % block_size != 0);
|
||||
|
||||
if (!stream_requant) {
|
||||
// Full materialization (original behavior): dequantize the whole tensor to F32,
|
||||
// then convert/quantize in one call.
|
||||
std::vector<float> weights_f32(n_elements);
|
||||
type_traits->to_float(data, weights_f32.data(), n_elements);
|
||||
if (requant_type == ExtraQuantType::F16) {
|
||||
ggml_get_type_traits(GGML_TYPE_F16)->from_float_ref(weights_f32.data(), weights.data(), n_elements);
|
||||
auto result = std::make_shared<ov::op::v0::Constant>(weights);
|
||||
result->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
if (requant_type == ExtraQuantType::Q4_1_64) {
|
||||
quantize_q4_1_asym(weights_f32.data(), weights, scales, zp, n_elements, block_size);
|
||||
} else if (is_u4) {
|
||||
quantize_q4_0(weights_f32.data(), weights, scales, zp, n_elements, block_size);
|
||||
} else if (requant_type == ExtraQuantType::Q8_1_C) {
|
||||
quantize_q8_1(weights_f32.data(), weights, scales, zp, n_elements, block_size);
|
||||
} else {
|
||||
quantize_q8_0(weights_f32.data(), weights, scales, zp, n_elements, block_size);
|
||||
}
|
||||
} else {
|
||||
// Streaming path for Q8_0_C / Q8_1_C / F16 (covers token_embd, output.weight,
|
||||
// and per-layer Q6_K/Q5_K requant — the large transient cases).
|
||||
const int64_t CHUNK_ROWS = std::min<int64_t>(n_rows, 256);
|
||||
std::vector<float> scratch(CHUNK_ROWS * ne0);
|
||||
// F16 destination: 2 bytes/element, advanced per chunk by r0*ne0 elements.
|
||||
auto * f16_base = static_cast<uint8_t *>(weights.data());
|
||||
for (int64_t r0 = 0; r0 < n_rows; r0 += CHUNK_ROWS) {
|
||||
const int64_t rows = std::min(CHUNK_ROWS, n_rows - r0);
|
||||
const int64_t elems = rows * ne0;
|
||||
const auto * src = static_cast<const uint8_t *>(data) + r0 * src_row_bytes;
|
||||
type_traits->to_float(src, scratch.data(), elems);
|
||||
|
||||
if (requant_type == ExtraQuantType::F16) {
|
||||
ggml_get_type_traits(GGML_TYPE_F16)
|
||||
->from_float_ref(scratch.data(), f16_base + (r0 * ne0) * sizeof(uint16_t), elems);
|
||||
} else {
|
||||
const int64_t block_offset = (r0 * ne0) / block_size;
|
||||
if (requant_type == ExtraQuantType::Q8_1_C) {
|
||||
quantize_q8_1(scratch.data(), weights, scales, zp, elems, block_size, block_offset);
|
||||
} else {
|
||||
quantize_q8_0(scratch.data(), weights, scales, zp, elems, block_size, block_offset);
|
||||
}
|
||||
}
|
||||
}
|
||||
if (requant_type == ExtraQuantType::F16) {
|
||||
auto result = std::make_shared<ov::op::v0::Constant>(weights);
|
||||
result->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
// Create the OpenVINO weight subgraph
|
||||
ov::Output<ov::Node> weight_node;
|
||||
if (is_u4) {
|
||||
weight_node = make_int4_weights(weights, scales, zp, block_size);
|
||||
} else {
|
||||
weight_node = make_int8_weights(weights, scales, zp, block_size);
|
||||
}
|
||||
|
||||
auto result = weight_node.get_node_shared_ptr();
|
||||
result->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
} // namespace
|
||||
|
||||
OvWeight process_weight_tensor(const ggml_tensor * tensor, const void * data, void * output_base_ptr, bool use_bias) {
|
||||
GGML_ASSERT(tensor != nullptr);
|
||||
GGML_ASSERT(data != nullptr);
|
||||
|
||||
OvWeight result;
|
||||
|
||||
// Get shape for weights: [rows, cols], or [n_expert, rows, cols] for 3D MoE expert weights.
|
||||
ov::Shape node_shape = (tensor->ne[2] > 1) ?
|
||||
ov::Shape{static_cast<size_t>(tensor->ne[2]), static_cast<size_t>(tensor->ne[1]),
|
||||
static_cast<size_t>(tensor->ne[0])} :
|
||||
ov::Shape{static_cast<size_t>(tensor->ne[1]), static_cast<size_t>(tensor->ne[0])};
|
||||
|
||||
// Handle F16/F32/BF16 weights
|
||||
if (tensor->type == GGML_TYPE_F32 || tensor->type == GGML_TYPE_F16 || tensor->type == GGML_TYPE_BF16) {
|
||||
ov::element::Type element_type;
|
||||
switch (tensor->type) {
|
||||
case GGML_TYPE_F32:
|
||||
element_type = ov::element::f32;
|
||||
break;
|
||||
case GGML_TYPE_F16:
|
||||
element_type = ov::element::f16;
|
||||
break;
|
||||
case GGML_TYPE_BF16:
|
||||
element_type = ov::element::bf16;
|
||||
break;
|
||||
default:
|
||||
OPENVINO_THROW("Unexpected tensor type in F16/F32/BF16 path");
|
||||
}
|
||||
|
||||
if (output_base_ptr && output_base_ptr != data) {
|
||||
// Using external buffer - copy data and create shared-memory constant
|
||||
size_t tensor_bytes = ggml_nbytes(tensor);
|
||||
memcpy(output_base_ptr, data, tensor_bytes);
|
||||
result.weights = ov::Tensor(element_type, node_shape, output_base_ptr);
|
||||
} else {
|
||||
result.weights = ov::Tensor(element_type, node_shape, data);
|
||||
}
|
||||
result.weight_node = std::make_shared<ov::op::v0::Constant>(result.weights);
|
||||
return result;
|
||||
}
|
||||
|
||||
// Handle quantized weights
|
||||
if (!ggml_is_quantized(tensor->type)) {
|
||||
OPENVINO_THROW("Unsupported weight tensor type: ", ggml_type_name(tensor->type));
|
||||
}
|
||||
|
||||
result.layout = ggml_openvino_get_extracted_layout(tensor, use_bias);
|
||||
const auto & layout = result.layout;
|
||||
if (layout.total_size == 0) {
|
||||
OPENVINO_THROW("Unsupported quantized type: ", ggml_type_name(tensor->type));
|
||||
}
|
||||
|
||||
// 3D MoE expert weights (for_gather_matmul) always use the exact f16 zero-point path (see
|
||||
// extract_quantized_weights) -- must be kept in sync with the "use_bias || for_gather_matmul"
|
||||
// check in ggml_openvino_get_extracted_layout, which sizes/offsets the zp slot accordingly.
|
||||
// Requantized tensors (layout.is_requant) are handled by requantize_to_buffers instead, whose
|
||||
// zp sizing/type is unaffected by for_gather_matmul, so they are excluded here.
|
||||
const bool for_gather_matmul = tensor->ne[2] > 1;
|
||||
const bool zp_is_f16 = !layout.is_requant && (use_bias || for_gather_matmul);
|
||||
|
||||
const bool is_3d_mxfp4_moe = tensor->type == GGML_TYPE_MXFP4 && (tensor->ne[2] > 1 || tensor->ne[3] > 1);
|
||||
if (is_3d_mxfp4_moe) {
|
||||
ov::Shape packed_shape = {static_cast<size_t>(tensor->ne[3]),
|
||||
static_cast<size_t>(tensor->ne[2]),
|
||||
static_cast<size_t>(tensor->ne[1]),
|
||||
static_cast<size_t>(tensor->ne[0] / MXFP4_BLOCK_SIZE),
|
||||
MXFP4_BLOCK_BYTES};
|
||||
const size_t tensor_bytes = ggml_nbytes(tensor);
|
||||
if (output_base_ptr) {
|
||||
auto * buf_base = static_cast<uint8_t *>(output_base_ptr);
|
||||
memcpy(buf_base + layout.weights_offset, data, tensor_bytes);
|
||||
result.weights = ov::Tensor(ov::element::u8, packed_shape, buf_base + layout.weights_offset);
|
||||
} else {
|
||||
result.weights = ov::Tensor(ov::element::u8, packed_shape);
|
||||
memcpy(result.weights.data(), data, tensor_bytes);
|
||||
}
|
||||
result.weight_node = make_mxfp4_moe_packed_weights(result.weights).get_node_shared_ptr();
|
||||
result.weight_node->set_friendly_name(tensor->name);
|
||||
return result;
|
||||
}
|
||||
|
||||
if (use_bias) {
|
||||
OPENVINO_ASSERT(!layout.is_requant,
|
||||
"use_bias is only used for test-backend-ops, which should not have requantization");
|
||||
// bias node will be created on the fly and not use backend buffer
|
||||
output_base_ptr = nullptr;
|
||||
}
|
||||
|
||||
// F16 requant path - no separate scales/zp needed in result
|
||||
if (layout.is_requant && layout.requant_type.has_value() && layout.requant_type.value() == ExtraQuantType::F16) {
|
||||
if (output_base_ptr) {
|
||||
result.weights = ov::Tensor(ov::element::f16, node_shape,
|
||||
static_cast<uint8_t *>(output_base_ptr) + layout.weights_offset);
|
||||
} else {
|
||||
result.weights = ov::Tensor(ov::element::f16, node_shape);
|
||||
}
|
||||
// Not used for F16:
|
||||
ov::Tensor dummy_scales;
|
||||
ov::Tensor dummy_zp;
|
||||
result.weight_node =
|
||||
requantize_to_buffers(tensor, data, ExtraQuantType::F16, 0, result.weights, dummy_scales, dummy_zp);
|
||||
return result;
|
||||
}
|
||||
|
||||
// Quantized path (normal extraction or quantized requant)
|
||||
// Create weight/scale/zp tensors - shared between both paths
|
||||
// For symmetric quantization, use signed types (i4/i8) and no ZP tensor
|
||||
ov::element::Type weight_type;
|
||||
if (tensor->type == GGML_TYPE_MXFP4) {
|
||||
weight_type = ov::element::f4e2m1;
|
||||
} else if (layout.is_symmetric) {
|
||||
weight_type = layout.is_u4 ? ov::element::i4 : ov::element::i8;
|
||||
} else {
|
||||
weight_type = layout.is_u4 ? ov::element::u4 : ov::element::u8;
|
||||
}
|
||||
ov::Shape scale_shape = node_shape;
|
||||
scale_shape.back() /= layout.weights_per_block;
|
||||
|
||||
if (tensor->type == GGML_TYPE_MXFP4) {
|
||||
if (tensor->ne[2] == 1 && tensor->ne[3] == 1) {
|
||||
node_shape = {static_cast<size_t>(tensor->ne[1]), static_cast<size_t>(tensor->ne[0])};
|
||||
} else {
|
||||
node_shape.clear();
|
||||
for (int i = GGML_MAX_DIMS - 1; i >= 0; --i) {
|
||||
node_shape.push_back(static_cast<size_t>(tensor->ne[i]));
|
||||
}
|
||||
}
|
||||
|
||||
scale_shape = node_shape;
|
||||
scale_shape.back() /= layout.weights_per_block;
|
||||
}
|
||||
|
||||
const ov::element::Type scale_type = tensor->type == GGML_TYPE_MXFP4 ? ov::element::f8e8m0 : ov::element::f16;
|
||||
ov::element::Type zp_type = layout.is_u4 ? ov::element::u4 : ov::element::u8;
|
||||
if (zp_is_f16) {
|
||||
zp_type = ov::element::f16;
|
||||
}
|
||||
|
||||
if (output_base_ptr) {
|
||||
uint8_t * buf_base = static_cast<uint8_t *>(output_base_ptr);
|
||||
result.weights = ov::Tensor(weight_type, node_shape, buf_base + layout.weights_offset);
|
||||
result.scales = ov::Tensor(scale_type, scale_shape, buf_base + layout.scales_offset);
|
||||
if (!layout.is_symmetric) {
|
||||
result.zp = ov::Tensor(zp_type, scale_shape, buf_base + layout.zp_offset);
|
||||
}
|
||||
// else: result.zp remains default-constructed (empty) for symmetric
|
||||
} else {
|
||||
result.weights = ov::Tensor(weight_type, node_shape);
|
||||
result.scales = ov::Tensor(scale_type, scale_shape);
|
||||
if (!layout.is_symmetric) {
|
||||
result.zp = ov::Tensor(zp_type, scale_shape);
|
||||
}
|
||||
// else: result.zp remains default-constructed (empty) for symmetric
|
||||
}
|
||||
|
||||
if (layout.is_requant && layout.requant_type.has_value()) {
|
||||
result.weight_node = requantize_to_buffers(tensor, data, layout.requant_type.value(), layout.weights_per_block,
|
||||
result.weights, result.scales, result.zp);
|
||||
} else {
|
||||
result.weight_node =
|
||||
extract_quantized_weights(tensor, data, result.weights, result.scales, result.zp, use_bias);
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
@@ -2,112 +2,12 @@
|
||||
#include "ggml-openvino-extra.h" // For ExtraQuantType
|
||||
#include "ggml.h"
|
||||
|
||||
#include <cstdint>
|
||||
#include <openvino/op/constant.hpp>
|
||||
#include <openvino/core/node_output.hpp>
|
||||
#include <openvino/op/constant.hpp>
|
||||
#include <openvino/runtime/tensor.hpp>
|
||||
|
||||
void unpack_32_4(const uint8_t * data, uint8_t * dst);
|
||||
|
||||
void extract_q4_0_data(const ggml_tensor * tensor,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
ov::Tensor & zp_arr);
|
||||
|
||||
void extract_q4_1_data(const ggml_tensor * tensor,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
ov::Tensor & zp_arr,
|
||||
bool use_bias = false);
|
||||
|
||||
void extract_q5_1_data(const ggml_tensor * tensor,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
ov::Tensor & zp_arr,
|
||||
bool use_bias = false);
|
||||
|
||||
void extract_q8_0_data(const ggml_tensor * tensor,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
ov::Tensor & zp_arr);
|
||||
|
||||
void unpack_256_4(const uint8_t * data, uint8_t * dst);
|
||||
|
||||
void extract_q4_k_data(const ggml_tensor * tensor,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
ov::Tensor & zp_arr,
|
||||
bool use_bias = false);
|
||||
|
||||
void extract_q5_k_data(const ggml_tensor * tensor,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
ov::Tensor & zp_arr,
|
||||
bool use_bias = false);
|
||||
|
||||
void extract_q6_k_data(const ggml_tensor * tensor,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
ov::Tensor & zp_arr);
|
||||
|
||||
void extract_mxfp4_data(const ggml_tensor * tensor, ov::Tensor & weights_arr, ov::Tensor & scales_arr);
|
||||
|
||||
static constexpr size_t GGML_QUANTIZATION_GROUP_SIZE = 32;
|
||||
|
||||
// If for_gather_matmul is true, the weight tensor may be N-D (e.g. 3D MoE expert weights
|
||||
// [n_expert, rows, cols]). The dequantization chain (Convert->[Subtract]->Multiply) is built as
|
||||
// usual but left in f16 (no final Convert to f32) -- ov::pass::MarkDequantization (registered in
|
||||
// translate_session.cpp) marks the chain so it survives model-build-time ConstantFolding -- see
|
||||
// make_int8_weights.cpp/make_int4_weights.cpp. mul_mat_id.cpp constructs ov::op::internal::GatherMatmul
|
||||
// directly from the resulting f16 dequant chain.
|
||||
//
|
||||
// When use_bias is true (explicitly, or implicitly because for_gather_matmul is true), the zp
|
||||
// tensor is expected to hold an exact f16 bias value (rather than a rounded integer zero point);
|
||||
// it is converted in place into an exact zero_point = -bias/scale and consumed via Subtract, not
|
||||
// Add, so the chain still matches OpenVINO's Convert->Subtract->Multiply decompression pattern.
|
||||
ov::Output<ov::Node> make_int8_weights(ov::Tensor & weight,
|
||||
ov::Tensor & scales,
|
||||
ov::Tensor & zp,
|
||||
size_t group_size = GGML_QUANTIZATION_GROUP_SIZE,
|
||||
bool use_bias = false,
|
||||
bool for_gather_matmul = false);
|
||||
|
||||
ov::Output<ov::Node> make_int4_weights(ov::Tensor & weight,
|
||||
ov::Tensor & scales,
|
||||
ov::Tensor & zp,
|
||||
size_t group_size = GGML_QUANTIZATION_GROUP_SIZE,
|
||||
bool use_bias = false,
|
||||
bool for_gather_matmul = false);
|
||||
|
||||
ov::Output<ov::Node> make_mxfp4_weights(ov::Tensor & weight, ov::Tensor & scales);
|
||||
|
||||
ov::Output<ov::Node> make_mxfp4_moe_packed_weights(ov::Tensor & weight);
|
||||
|
||||
// Extract quantized weights from tensor and create weight subgraph
|
||||
// If weights/scales/zp are provided (non-empty), uses them as output buffers
|
||||
// Otherwise allocates new ov::Tensors internally
|
||||
// Returns the weight node (make_int4_weights or make_int8_weights result)
|
||||
std::shared_ptr<ov::Node> extract_quantized_weights(
|
||||
const ggml_tensor * tensor,
|
||||
const void * data, // Source data pointer (may differ from tensor->data)
|
||||
ov::Tensor & weights,
|
||||
ov::Tensor & scales,
|
||||
ov::Tensor & zp,
|
||||
bool use_bias = false); // Use an exact f16 zero point (vs. a rounded integer one); always
|
||||
// used for for_gather_matmul (3D MoE expert) weights regardless of
|
||||
// this flag, and also settable explicitly for test-backend-ops.
|
||||
|
||||
// Requantize weights from tensor to target format, writing to provided buffers
|
||||
// For F16 target, only weights buffer is used (scales/zp ignored)
|
||||
// Returns the weight node
|
||||
std::shared_ptr<ov::Node> requantize_to_buffers(const ggml_tensor * tensor,
|
||||
const void * data, // Source data pointer
|
||||
ExtraQuantType requant_type,
|
||||
int64_t block_size,
|
||||
ov::Tensor & weights,
|
||||
ov::Tensor & scales,
|
||||
ov::Tensor & zp);
|
||||
|
||||
inline const char * extra_quant_type_name(ExtraQuantType t) {
|
||||
switch (t) {
|
||||
case ExtraQuantType::F16:
|
||||
@@ -156,41 +56,3 @@ OvWeight process_weight_tensor(
|
||||
// always used for for_gather_matmul (3D MoE expert) weights
|
||||
// regardless of this flag, and also settable explicitly for
|
||||
// test-backend-ops.
|
||||
|
||||
void quantize_q4_0(const float * x,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
ov::Tensor & zp_arr,
|
||||
int64_t k,
|
||||
int64_t qk);
|
||||
void quantize_q8_1(const float * x,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
ov::Tensor & zp_arr,
|
||||
int64_t k,
|
||||
int64_t qk,
|
||||
int64_t block_offset = 0);
|
||||
void quantize_q4_1_asym(const float * x,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
ov::Tensor & zp_arr,
|
||||
int64_t k,
|
||||
int64_t qk);
|
||||
void quantize_q8_0(const float * x,
|
||||
ov::Tensor & weights_arr,
|
||||
ov::Tensor & scales_arr,
|
||||
ov::Tensor & zp_arr,
|
||||
int64_t k,
|
||||
int64_t qk,
|
||||
int64_t block_offset = 0);
|
||||
|
||||
namespace ov {
|
||||
namespace op {
|
||||
namespace util {
|
||||
// From <openvino>/src/common/transformations/include/transformations/utils/utils.hpp
|
||||
bool get_single_value(const std::shared_ptr<ov::op::v0::Constant> & const_node,
|
||||
float & value,
|
||||
bool check_value_range = true);
|
||||
} // namespace util
|
||||
} // namespace op
|
||||
} // namespace ov
|
||||
|
||||
@@ -237,7 +237,8 @@ bool ggml_openvino_model_cache_verify_manifest(const std::string & path,
|
||||
if (!f.is_open()) {
|
||||
return false;
|
||||
}
|
||||
std::string tag, val;
|
||||
std::string tag;
|
||||
std::string val;
|
||||
// header: fingerprint
|
||||
if (!(f >> tag >> val) || tag != "fingerprint" || val != hex64(fingerprint)) {
|
||||
return false;
|
||||
|
||||
@@ -12,7 +12,6 @@ namespace ggml {
|
||||
|
||||
class FrontEnd {
|
||||
public:
|
||||
using Ptr = std::shared_ptr<FrontEnd>;
|
||||
FrontEnd();
|
||||
|
||||
static std::shared_ptr<Model> convert(const InputModel::Ptr & model, bool naive = false);
|
||||
|
||||
@@ -20,7 +20,7 @@ namespace op {
|
||||
static ov::Output<ov::Node> reshape_add_id_input_to_2d(const ov::Output<ov::Node> & input,
|
||||
const ov::PartialShape & input_shape,
|
||||
const std::vector<int> & dims) {
|
||||
const auto actual_shape = input.get_partial_shape();
|
||||
const auto & actual_shape = input.get_partial_shape();
|
||||
if (actual_shape.rank().is_static() && actual_shape.rank().get_length() == 2) {
|
||||
return input;
|
||||
}
|
||||
|
||||
@@ -3,12 +3,9 @@
|
||||
#include "../op_table.h"
|
||||
#include "../utils.h"
|
||||
|
||||
#include <climits>
|
||||
#include <cstdint>
|
||||
#include <memory>
|
||||
#include <openvino/op/reshape.hpp>
|
||||
#include <openvino/op/slice.hpp>
|
||||
#include <vector>
|
||||
|
||||
namespace ov {
|
||||
namespace frontend {
|
||||
|
||||
@@ -195,7 +195,9 @@ OutputVector translate_flash_attn_ext(const NodeContext & context) {
|
||||
auto tile_kv = [&](int64_t n_heads, int64_t n_heads_kv, int64_t hs, ov::Output<Node> kv) {
|
||||
int64_t f = n_heads / n_heads_kv;
|
||||
if (f > 1 && n_heads_kv > 1) {
|
||||
ov::Output<ov::Node> kv_broadcast_shape, kv_unsqueezed, new_kv_shape;
|
||||
ov::Output<ov::Node> kv_broadcast_shape;
|
||||
ov::Output<ov::Node> kv_unsqueezed;
|
||||
ov::Output<ov::Node> new_kv_shape;
|
||||
auto unsqueeze_axes = ov::op::v0::Constant::create(ov::element::i64, Shape{}, {2});
|
||||
kv_unsqueezed = std::make_shared<ov::op::v0::Unsqueeze>(kv, unsqueeze_axes);
|
||||
|
||||
|
||||
@@ -196,7 +196,7 @@ static OutputVector translate_gated_delta_net_ref(const NodeContext & context) {
|
||||
}
|
||||
|
||||
// Merge batch and head dims: [B*H_v, T, S_v]
|
||||
auto merge_bh = [&](ov::Output<ov::Node> x, int64_t last_dim) {
|
||||
auto merge_bh = [&](const ov::Output<ov::Node> & x, int64_t last_dim) {
|
||||
auto shape = ov::op::v0::Constant::create(ov::element::i64, {3}, std::vector<int64_t>{B * H_v, T, last_dim});
|
||||
return std::make_shared<ov::op::v1::Reshape>(x, shape, false);
|
||||
};
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
#include "../node_context.h"
|
||||
#include "../op_table.h"
|
||||
#include "../utils.h"
|
||||
#include "ggml-impl.h"
|
||||
|
||||
#include <cstddef>
|
||||
#include <memory>
|
||||
|
||||
@@ -42,7 +42,7 @@ ov::Output<ov::Node> slice_axis(const ov::Output<ov::Node> & input, int64_t axis
|
||||
|
||||
ov::Output<ov::Node> static_shape_dims_or_shapeof(const ov::Output<ov::Node> & input,
|
||||
const std::vector<int> & dims) {
|
||||
const auto partial_shape = input.get_partial_shape();
|
||||
const auto & partial_shape = input.get_partial_shape();
|
||||
if (partial_shape.is_static()) {
|
||||
std::vector<int64_t> values;
|
||||
values.reserve(dims.size());
|
||||
|
||||
@@ -8,6 +8,7 @@
|
||||
#include <openvino/op/pad.hpp>
|
||||
#include <openvino/op/reshape.hpp>
|
||||
#include <openvino/op/shape_of.hpp>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
namespace ov {
|
||||
@@ -20,7 +21,7 @@ namespace {
|
||||
ov::Output<ov::Node> translate_circular_pad(ov::Output<ov::Node> input,
|
||||
const std::array<int32_t, 8> & pads,
|
||||
const ov::Shape & input_shape) {
|
||||
ov::Output<ov::Node> result = input;
|
||||
ov::Output<ov::Node> result = std::move(input);
|
||||
|
||||
const std::array<int32_t, 4> pads_begin = {pads[6], pads[4], pads[2], pads[0]};
|
||||
const std::array<int32_t, 4> pads_end = {pads[7], pads[5], pads[3], pads[1]};
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
#include "../node_context.h"
|
||||
#include "../op_table.h"
|
||||
#include "../utils.h"
|
||||
#include "ggml.h"
|
||||
|
||||
#include <memory>
|
||||
#include <openvino/op/broadcast.hpp>
|
||||
|
||||
@@ -25,9 +25,7 @@ OutputVector translate_rms_norm(const NodeContext & context) {
|
||||
auto op_case = context.get_op_case();
|
||||
|
||||
ov::Output<ov::Node> input_node;
|
||||
if (op_case == 1) {
|
||||
input_node = process_view_input_new(context, 0);
|
||||
} else if (op_case == 2) {
|
||||
if (op_case == 2) {
|
||||
auto ssm_state_size = context.get_ssm_state_size();
|
||||
// The GDN op packs [attn | new_state] along the row axis; the state occupies the last
|
||||
// ssm_state_size * n_seqs rows. Slice it off (scaling by the active sequence count) to keep
|
||||
|
||||
@@ -7,7 +7,6 @@
|
||||
#include <openvino/op/reshape.hpp>
|
||||
#include <openvino/op/shape_of.hpp>
|
||||
#include <openvino/op/slice.hpp>
|
||||
#include <set>
|
||||
|
||||
namespace ov {
|
||||
namespace frontend {
|
||||
@@ -153,7 +152,8 @@ OutputVector translate_view(const NodeContext & context) {
|
||||
return {input};
|
||||
}
|
||||
|
||||
int64_t src_elems = 1, dst_elems = 1;
|
||||
int64_t src_elems = 1;
|
||||
int64_t dst_elems = 1;
|
||||
for (int64_t i = 0; i < src_shape.rank().get_length(); ++i) {
|
||||
if (src_shape[i].is_dynamic()) {
|
||||
return {input};
|
||||
|
||||
@@ -84,7 +84,7 @@ bool KVStateSeqAxis::run_on_model(const std::shared_ptr<ov::Model> & model) {
|
||||
// Readers still expect seq at dim 1. A reader that is itself the inverse
|
||||
// Transpose wanted seq at dim 2 all along, so drop it; give anything else the
|
||||
// inverse Transpose so its input is unchanged.
|
||||
for (auto & reader : readers) {
|
||||
for (const auto & reader : readers) {
|
||||
auto * node = reader.get_node();
|
||||
if (ov::is_type<ov::op::v6::Assign>(node)) {
|
||||
continue;
|
||||
|
||||
@@ -344,7 +344,7 @@ std::shared_ptr<Model> TranslateSession::translate_graph(const frontend::InputMo
|
||||
}
|
||||
};
|
||||
|
||||
auto node_visitor = [&](std::shared_ptr<GgmlDecoder> decoder, int node_idx) {
|
||||
auto node_visitor = [&](const std::shared_ptr<GgmlDecoder> & decoder, int node_idx) {
|
||||
auto converted_outputs = translate_node(decoder, node_idx);
|
||||
if (converted_outputs.empty()) {
|
||||
return;
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
#include "utils.h"
|
||||
|
||||
#include "ggml-impl.h"
|
||||
|
||||
#include <cmath>
|
||||
#include <cstddef>
|
||||
#include <ctime>
|
||||
@@ -28,13 +26,6 @@ namespace ov {
|
||||
namespace frontend {
|
||||
namespace ggml {
|
||||
|
||||
std::string getCurrentTime() {
|
||||
std::time_t now = std::time(nullptr);
|
||||
char buf[100];
|
||||
std::strftime(buf, sizeof(buf), "%Y-%m-%d %H:%M:%S", std::localtime(&now));
|
||||
return buf;
|
||||
}
|
||||
|
||||
void num_inputs_check(const NodeContext & context, size_t min_inputs, size_t max_inputs) {
|
||||
auto input_size = context.get_input_size();
|
||||
FRONT_END_OP_CONVERSION_CHECK(input_size >= min_inputs, "Got less inputs than expected");
|
||||
@@ -82,7 +73,7 @@ namespace {
|
||||
ov::Output<ov::Node> rope_yarn_ramp_mix(int n_dims, const float corr_dims[2], float ext_factor) {
|
||||
int half_n_dims = n_dims / 2;
|
||||
std::vector<float> dim_ids_vec(half_n_dims);
|
||||
std::iota(dim_ids_vec.begin(), dim_ids_vec.end(), 0);
|
||||
std::iota(dim_ids_vec.begin(), dim_ids_vec.end(), 0.0f);
|
||||
auto dim_ids = ov::op::v0::Constant::create(ov::element::f32, Shape{1, 1, 1, (size_t) half_n_dims}, dim_ids_vec);
|
||||
auto corr_low = ov::op::v0::Constant::create(ov::element::f32, Shape{1, 1, 1, 1}, {corr_dims[0]});
|
||||
auto corr_high = ov::op::v0::Constant::create(ov::element::f32, Shape{1, 1, 1, 1}, {corr_dims[1]});
|
||||
@@ -551,6 +542,7 @@ ov::Output<ov::Node> process_view_input_new(const NodeContext & context, int inp
|
||||
|
||||
if (tail_begin >= 0 && tail_end <= tail_src_elems) {
|
||||
std::vector<int64_t> flat_shape;
|
||||
flat_shape.reserve(slice_dim);
|
||||
for (int i = 0; i < slice_dim; ++i) {
|
||||
flat_shape.push_back(static_cast<int64_t>(view_src_ggml_shape[i]));
|
||||
}
|
||||
|
||||
@@ -14,8 +14,6 @@ namespace ggml {
|
||||
|
||||
std::string getCurrentTime();
|
||||
|
||||
void dump_ov_model(std::shared_ptr<ov::Model> model);
|
||||
|
||||
void num_inputs_check(const NodeContext & context, size_t min_inputs, size_t max_inputs);
|
||||
|
||||
int non_cont_dim(std::vector<size_t> ne, std::vector<size_t> nb);
|
||||
|
||||
+461
-473
File diff suppressed because it is too large
Load Diff
@@ -142,9 +142,6 @@ struct ov_runtime_context {
|
||||
|
||||
enum ggml_status ov_graph_compute(struct ggml_cgraph * cgraph, ggml_backend_t backend);
|
||||
|
||||
enum ggml_status ov_graph_compute_dynamic(struct ggml_cgraph * cgraph, std::shared_ptr<ov_runtime_context> r_ctx);
|
||||
enum ggml_status ov_graph_compute_static(struct ggml_cgraph * cgraph, std::shared_ptr<ov_runtime_context> r_ctx);
|
||||
|
||||
size_t checksum(const void * data, size_t size);
|
||||
|
||||
bool save_ggml_tensor_data_to_txt(const ggml_tensor * tensor, const std::string & file_path);
|
||||
@@ -185,18 +182,6 @@ int64_t get_inp_pos_n_tokens(struct ggml_cgraph * cgraph, const ggml_tensor * in
|
||||
|
||||
bool get_is_prefill(struct ggml_cgraph * cgraph, const ggml_tensor * inp_pos);
|
||||
|
||||
ov::Tensor get_ov_input_tensor(std::shared_ptr<GgmlOvDecoder> ggml_decoder, const std::string & param_name);
|
||||
ov::Tensor get_ov_input_tensor_static_decode(std::shared_ptr<GgmlOvDecoder> ggml_decoder,
|
||||
const std::string & param_name);
|
||||
ov::Tensor get_ov_input_tensor_static_prefill(std::shared_ptr<GgmlOvDecoder> ggml_decoder,
|
||||
const std::string & param_name,
|
||||
int chunk_index);
|
||||
|
||||
ov::Tensor create_ov_output_tensor(std::shared_ptr<GgmlOvDecoder> ggml_decoder,
|
||||
std::shared_ptr<ov::InferRequest> infer_request,
|
||||
int output_index,
|
||||
const ggml_tensor * ggml_tensor);
|
||||
|
||||
bool is_naive(struct ggml_cgraph * cgraph);
|
||||
|
||||
/**
|
||||
@@ -205,9 +190,3 @@ bool is_naive(struct ggml_cgraph * cgraph);
|
||||
* @return true if the graph is identified as split; otherwise false.
|
||||
*/
|
||||
bool is_model_splitted(struct ggml_cgraph * cgraph);
|
||||
|
||||
enum ggml_status naive_compute(struct ggml_cgraph * cgraph,
|
||||
ov::Core & core,
|
||||
const std::string & device,
|
||||
const ov::AnyMap & config,
|
||||
ov_compiled_model_cache & cache);
|
||||
|
||||
@@ -62,6 +62,11 @@ if (Vulkan_FOUND)
|
||||
ggml_add_backend_library(ggml-vulkan
|
||||
ggml-vulkan.cpp
|
||||
../../include/ggml-vulkan.h
|
||||
ggml-vulkan-types.h
|
||||
ggml-vulkan-push-constants.h
|
||||
ggml-vulkan-common.h
|
||||
ggml-vulkan-buffers.cpp
|
||||
ggml-vulkan-debug.cpp
|
||||
)
|
||||
|
||||
set(VULKAN_SHADER_GEN_CMAKE_ARGS "")
|
||||
|
||||
@@ -0,0 +1,783 @@
|
||||
#include "ggml-vulkan-common.h"
|
||||
|
||||
ggml_backend_buffer_type_i ggml_backend_vk_buffer_type_interface = {
|
||||
/* .get_name = */ ggml_backend_vk_buffer_type_name,
|
||||
/* .alloc_buffer = */ ggml_backend_vk_buffer_type_alloc_buffer,
|
||||
/* .get_alignment = */ ggml_backend_vk_buffer_type_get_alignment,
|
||||
/* .get_max_size = */ ggml_backend_vk_buffer_type_get_max_size,
|
||||
/* .get_alloc_size = */ ggml_backend_vk_buffer_type_get_alloc_size,
|
||||
/* .is_host = */ NULL,
|
||||
};
|
||||
|
||||
static std::vector<uint32_t> ggml_vk_find_memory_properties(const vk::PhysicalDeviceMemoryProperties* mem_props, vk::MemoryRequirements* mem_req, vk::MemoryPropertyFlags flags) {
|
||||
std::vector<uint32_t> indices;
|
||||
|
||||
for (uint32_t i = 0; i < mem_props->memoryTypeCount; ++i) {
|
||||
vk::MemoryType memory_type = mem_props->memoryTypes[i];
|
||||
if ((mem_req->memoryTypeBits & ((uint64_t)1 << i)) &&
|
||||
(flags & memory_type.propertyFlags) == flags &&
|
||||
mem_props->memoryHeaps[memory_type.heapIndex].size >= mem_req->size) {
|
||||
indices.push_back(i);
|
||||
}
|
||||
}
|
||||
return indices;
|
||||
}
|
||||
|
||||
static vk_buffer ggml_vk_create_buffer(vk_device& device, size_t size, const std::initializer_list<vk::MemoryPropertyFlags> & req_flags_list,
|
||||
void *import_ptr = nullptr) {
|
||||
VK_LOG_DEBUG("ggml_vk_create_buffer(" << device->name << ", " << size << ", " << to_string(req_flags_list.begin()[0]) << ", " << to_string(req_flags_list.begin()[req_flags_list.size()-1]) << ")");
|
||||
if (size > device->max_buffer_size) {
|
||||
throw vk::OutOfDeviceMemoryError("Requested buffer size exceeds device buffer size limit");
|
||||
}
|
||||
|
||||
vk_buffer buf = std::make_shared<vk_buffer_struct>();
|
||||
|
||||
if (size == 0) {
|
||||
buf->size = 0;
|
||||
return buf;
|
||||
}
|
||||
|
||||
vk::BufferUsageFlags usage_flags = vk::BufferUsageFlagBits::eStorageBuffer | vk::BufferUsageFlagBits::eTransferSrc | vk::BufferUsageFlagBits::eTransferDst;
|
||||
vk::MemoryAllocateFlags mem_flags {};
|
||||
if (device->buffer_device_address) {
|
||||
usage_flags |= vk::BufferUsageFlagBits::eShaderDeviceAddress;
|
||||
mem_flags |= vk::MemoryAllocateFlagBits::eDeviceAddress;
|
||||
}
|
||||
|
||||
vk::BufferCreateInfo buffer_create_info{
|
||||
vk::BufferCreateFlags(),
|
||||
size,
|
||||
usage_flags,
|
||||
vk::SharingMode::eExclusive,
|
||||
0,
|
||||
nullptr,
|
||||
};
|
||||
|
||||
vk::ExternalMemoryBufferCreateInfo external_memory_bci;
|
||||
if (import_ptr) {
|
||||
external_memory_bci.handleTypes = vk::ExternalMemoryHandleTypeFlagBits::eHostAllocationEXT;
|
||||
buffer_create_info.setPNext(&external_memory_bci);
|
||||
}
|
||||
|
||||
buf->buffer = device->device.createBuffer(buffer_create_info);
|
||||
|
||||
vk::MemoryRequirements mem_req = device->device.getBufferMemoryRequirements(buf->buffer);
|
||||
|
||||
vk::PhysicalDeviceMemoryProperties mem_props = device->physical_device.getMemoryProperties();
|
||||
|
||||
const vk::MemoryPriorityAllocateInfoEXT mem_priority_info { 1.0f };
|
||||
|
||||
vk::MemoryAllocateFlagsInfo mem_flags_info { mem_flags };
|
||||
|
||||
if (device->memory_priority) {
|
||||
mem_flags_info.setPNext(&mem_priority_info);
|
||||
}
|
||||
|
||||
if (import_ptr) {
|
||||
vk::MemoryHostPointerPropertiesEXT host_pointer_props;
|
||||
try {
|
||||
host_pointer_props = device->device.getMemoryHostPointerPropertiesEXT(vk::ExternalMemoryHandleTypeFlagBits::eHostAllocationEXT, import_ptr);
|
||||
} catch (vk::SystemError& e) {
|
||||
GGML_LOG_WARN("ggml_vulkan: Failed getMemoryHostPointerPropertiesEXT (%s)\n", e.what());
|
||||
device->device.destroyBuffer(buf->buffer);
|
||||
return {};
|
||||
}
|
||||
vk::PhysicalDeviceMemoryProperties mem_props = device->physical_device.getMemoryProperties();
|
||||
|
||||
uint32_t memory_type_idx;
|
||||
vk::MemoryPropertyFlags property_flags = *req_flags_list.begin();
|
||||
for (memory_type_idx = 0; memory_type_idx < 32; ++memory_type_idx) {
|
||||
if (!(host_pointer_props.memoryTypeBits & (1u << memory_type_idx))) {
|
||||
continue;
|
||||
}
|
||||
if (!(mem_req.memoryTypeBits & (1u << memory_type_idx))) {
|
||||
continue;
|
||||
}
|
||||
|
||||
vk::MemoryType memory_type = mem_props.memoryTypes[memory_type_idx];
|
||||
// check for visible+coherent+cached. Other flags (e.g. devicelocal) are allowed
|
||||
if ((memory_type.propertyFlags & property_flags) == property_flags) {
|
||||
property_flags = memory_type.propertyFlags;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (memory_type_idx == 32) {
|
||||
GGML_LOG_WARN("ggml_vulkan: Memory type for host allocation not found\n");
|
||||
device->device.destroyBuffer(buf->buffer);
|
||||
return {};
|
||||
}
|
||||
|
||||
buf->memory_property_flags = mem_props.memoryTypes[memory_type_idx].propertyFlags;
|
||||
try {
|
||||
vk::ImportMemoryHostPointerInfoEXT import_info;
|
||||
import_info.handleType = vk::ExternalMemoryHandleTypeFlagBits::eHostAllocationEXT;
|
||||
import_info.pHostPointer = import_ptr;
|
||||
import_info.setPNext(&mem_flags_info);
|
||||
buf->device_memory = device->device.allocateMemory({ size, memory_type_idx, &import_info });
|
||||
} catch (const vk::SystemError& e) {
|
||||
}
|
||||
} else {
|
||||
for (auto it = req_flags_list.begin(); it != req_flags_list.end(); it++) {
|
||||
const auto & req_flags = *it;
|
||||
|
||||
const std::vector<uint32_t> memory_type_indices = ggml_vk_find_memory_properties(&mem_props, &mem_req, req_flags);
|
||||
|
||||
if (memory_type_indices.empty()) {
|
||||
continue;
|
||||
}
|
||||
|
||||
bool done = false;
|
||||
|
||||
for (auto mtype_it = memory_type_indices.begin(); mtype_it != memory_type_indices.end(); mtype_it++) {
|
||||
try {
|
||||
buf->device_memory = device->device.allocateMemory({ mem_req.size, *mtype_it, &mem_flags_info });
|
||||
buf->memory_property_flags = mem_props.memoryTypes[*mtype_it].propertyFlags;
|
||||
done = true;
|
||||
break;
|
||||
} catch (const vk::SystemError& e) {
|
||||
// loop and retry
|
||||
// during last attempt throw the exception
|
||||
if (it + 1 == req_flags_list.end() && mtype_it + 1 == memory_type_indices.end()) {
|
||||
device->device.destroyBuffer(buf->buffer);
|
||||
throw e;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (done) {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (!buf->device_memory) {
|
||||
device->device.destroyBuffer(buf->buffer);
|
||||
throw vk::OutOfDeviceMemoryError("No suitable memory type found");
|
||||
}
|
||||
|
||||
buf->ptr = nullptr;
|
||||
|
||||
if (import_ptr) {
|
||||
buf->ptr = import_ptr;
|
||||
} else {
|
||||
if (buf->memory_property_flags & vk::MemoryPropertyFlagBits::eHostVisible) {
|
||||
buf->ptr = device->device.mapMemory(buf->device_memory, 0, VK_WHOLE_SIZE);
|
||||
}
|
||||
}
|
||||
|
||||
device->device.bindBufferMemory(buf->buffer, buf->device_memory, 0);
|
||||
|
||||
buf->device = device;
|
||||
buf->size = size;
|
||||
|
||||
if (device->buffer_device_address) {
|
||||
const vk::BufferDeviceAddressInfo addressInfo(buf->buffer);
|
||||
buf->bda_addr = device->device.getBufferAddress(addressInfo);
|
||||
}
|
||||
|
||||
device->memory_logger->log_allocation(buf, size);
|
||||
|
||||
return buf;
|
||||
}
|
||||
|
||||
vk_buffer ggml_vk_create_buffer_check(vk_device& device, size_t size, vk::MemoryPropertyFlags req_flags, vk::MemoryPropertyFlags fallback_flags) {
|
||||
try {
|
||||
return ggml_vk_create_buffer(device, size, {req_flags, fallback_flags});
|
||||
} catch (const vk::SystemError& e) {
|
||||
std::cerr << "ggml_vulkan: Memory allocation of size " << size << " failed." << std::endl;
|
||||
std::cerr << "ggml_vulkan: " << e.what() << std::endl;
|
||||
throw e;
|
||||
}
|
||||
}
|
||||
|
||||
vk_buffer ggml_vk_create_buffer_device(vk_device& device, size_t size) {
|
||||
vk_buffer buf;
|
||||
try {
|
||||
if (device->prefer_host_memory) {
|
||||
buf = ggml_vk_create_buffer(device, size, {vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent,
|
||||
vk::MemoryPropertyFlagBits::eDeviceLocal});
|
||||
} else if (device->uma) {
|
||||
// On UMA, prefer host-visible memory so direct tensor borrowing works.
|
||||
// If unavailable, fall back to device-local memory.
|
||||
buf = ggml_vk_create_buffer(device, size, {vk::MemoryPropertyFlagBits::eDeviceLocal | vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent,
|
||||
vk::MemoryPropertyFlagBits::eDeviceLocal,
|
||||
vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent});
|
||||
} else if (device->disable_host_visible_vidmem) {
|
||||
if (device->allow_sysmem_fallback) {
|
||||
buf = ggml_vk_create_buffer(device, size, {vk::MemoryPropertyFlagBits::eDeviceLocal,
|
||||
vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent});
|
||||
} else {
|
||||
buf = ggml_vk_create_buffer(device, size, {vk::MemoryPropertyFlagBits::eDeviceLocal});
|
||||
}
|
||||
} else {
|
||||
// use rebar if available, otherwise fallback to device only visible memory
|
||||
if (device->allow_sysmem_fallback) {
|
||||
buf = ggml_vk_create_buffer(device, size, {vk::MemoryPropertyFlagBits::eDeviceLocal | vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent,
|
||||
vk::MemoryPropertyFlagBits::eDeviceLocal,
|
||||
vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent});
|
||||
} else {
|
||||
buf = ggml_vk_create_buffer(device, size, {vk::MemoryPropertyFlagBits::eDeviceLocal | vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent,
|
||||
vk::MemoryPropertyFlagBits::eDeviceLocal});
|
||||
}
|
||||
}
|
||||
} catch (const vk::SystemError& e) {
|
||||
std::cerr << "ggml_vulkan: Device memory allocation of size " << size << " failed." << std::endl;
|
||||
std::cerr << "ggml_vulkan: " << e.what() << std::endl;
|
||||
throw e;
|
||||
}
|
||||
|
||||
return buf;
|
||||
}
|
||||
|
||||
void ggml_vk_destroy_buffer(vk_buffer& buf) {
|
||||
if (buf == nullptr) {
|
||||
return;
|
||||
}
|
||||
|
||||
if (buf->device != nullptr) {
|
||||
buf->device->memory_logger->log_deallocation(buf);
|
||||
}
|
||||
|
||||
buf.reset();
|
||||
}
|
||||
|
||||
void * ggml_vk_host_malloc(vk_device& device, size_t size) {
|
||||
VK_LOG_MEMORY("ggml_vk_host_malloc(" << size << ")");
|
||||
vk_buffer buf = ggml_vk_create_buffer(device, size,
|
||||
{vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent | vk::MemoryPropertyFlagBits::eHostCached,
|
||||
vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent});
|
||||
|
||||
if(!(buf->memory_property_flags & vk::MemoryPropertyFlagBits::eHostVisible)) {
|
||||
fprintf(stderr, "WARNING: failed to allocate %.2f MB of pinned memory\n",
|
||||
size/1024.0/1024.0);
|
||||
device->device.freeMemory(buf->device_memory);
|
||||
device->device.destroyBuffer(buf->buffer);
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
std::lock_guard<std::shared_mutex> guard(device->pinned_memory_mutex);
|
||||
device->pinned_memory.push_back(std::make_tuple(buf->ptr, size, buf));
|
||||
|
||||
return buf->ptr;
|
||||
}
|
||||
|
||||
void ggml_vk_host_free(vk_device& device, void* ptr) {
|
||||
if (ptr == nullptr) {
|
||||
return;
|
||||
}
|
||||
VK_LOG_MEMORY("ggml_vk_host_free(" << ptr << ")");
|
||||
std::lock_guard<std::shared_mutex> guard(device->pinned_memory_mutex);
|
||||
|
||||
vk_buffer buf;
|
||||
size_t index;
|
||||
for (size_t i = 0; i < device->pinned_memory.size(); i++) {
|
||||
const uint8_t* addr = (const uint8_t*) std::get<0>(device->pinned_memory[i]);
|
||||
const uint8_t* endr = addr + std::get<1>(device->pinned_memory[i]);
|
||||
if (ptr >= addr && ptr < endr) {
|
||||
buf = std::get<2>(device->pinned_memory[i]);
|
||||
index = i;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (buf == nullptr) {
|
||||
fprintf(stderr, "WARNING: failed to free pinned memory: memory not in map\n");
|
||||
return;
|
||||
}
|
||||
|
||||
ggml_vk_destroy_buffer(buf);
|
||||
|
||||
device->pinned_memory.erase(device->pinned_memory.begin() + index);
|
||||
}
|
||||
|
||||
void ggml_vk_host_get(const vk_device& device, const void * ptr, vk_buffer& buf, size_t& buf_offset) {
|
||||
std::shared_lock<std::shared_mutex> guard(device->pinned_memory_mutex);
|
||||
buf = nullptr;
|
||||
buf_offset = 0;
|
||||
for (size_t i = 0; i < device->pinned_memory.size(); i++) {
|
||||
const uint8_t* addr = (const uint8_t*) std::get<0>(device->pinned_memory[i]);
|
||||
const uint8_t* endr = addr + std::get<1>(device->pinned_memory[i]);
|
||||
if (ptr >= addr && ptr < endr) {
|
||||
buf = std::get<2>(device->pinned_memory[i]);
|
||||
buf_offset = ((const uint8_t *)ptr) - addr;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void ggml_vk_ensure_sync_staging_buffer(vk_device& device, size_t size) {
|
||||
if (device->sync_staging == nullptr || device->sync_staging->size < size) {
|
||||
VK_LOG_MEMORY("ggml_vk_ensure_sync_staging_buffer(" << size << ")");
|
||||
ggml_vk_destroy_buffer(device->sync_staging);
|
||||
device->sync_staging = ggml_vk_create_buffer_check(device, size,
|
||||
vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent | vk::MemoryPropertyFlagBits::eHostCached,
|
||||
vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent);
|
||||
}
|
||||
}
|
||||
|
||||
void ggml_vk_ensure_sync_staging_buffer(ggml_backend_vk_context * ctx, size_t size) {
|
||||
if (ctx->sync_staging == nullptr || ctx->sync_staging->size < size) {
|
||||
VK_LOG_MEMORY("ggml_vk_ensure_sync_staging_buffer(" << size << ")");
|
||||
ggml_vk_destroy_buffer(ctx->sync_staging);
|
||||
ctx->sync_staging = ggml_vk_create_buffer_check(ctx->device, size,
|
||||
vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent | vk::MemoryPropertyFlagBits::eHostCached,
|
||||
vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent);
|
||||
}
|
||||
}
|
||||
|
||||
static void ggml_vk_buffer_write_nc_async(ggml_backend_vk_context * ctx, vk_context& subctx, vk_buffer& dst, size_t offset, const ggml_tensor * tensor, bool sync_staging = false) {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_write_nc_async(" << tensor << ")");
|
||||
GGML_ASSERT(!ggml_is_contiguous(tensor));
|
||||
// Buffer is already mapped
|
||||
if(dst->memory_property_flags & vk::MemoryPropertyFlagBits::eHostVisible) {
|
||||
std::cerr << "ggml_vulkan: buffer_write_nc_async dst buffer is host_visible. Use synchronous write." << std::endl;
|
||||
GGML_ABORT("fatal error");
|
||||
}
|
||||
// Check if src is pinned memory
|
||||
vk_buffer buf = nullptr;
|
||||
size_t buf_offset = 0;
|
||||
ggml_vk_host_get(ctx->device, tensor->data, buf, buf_offset);
|
||||
|
||||
const uint64_t ne0 = tensor->ne[0];
|
||||
const uint64_t ne1 = tensor->ne[1];
|
||||
const uint64_t ne2 = tensor->ne[2];
|
||||
const uint64_t ne3 = tensor->ne[3];
|
||||
const uint64_t nb0 = tensor->nb[0];
|
||||
const uint64_t nb1 = tensor->nb[1];
|
||||
const uint64_t nb2 = tensor->nb[2];
|
||||
const uint64_t nb3 = tensor->nb[3];
|
||||
const ggml_type type = tensor->type;
|
||||
const uint64_t ts = ggml_type_size(type);
|
||||
const uint64_t bs = ggml_blck_size(type);
|
||||
|
||||
const uint64_t dstnb0 = ts;
|
||||
const uint64_t dstnb1 = dstnb0*(ne0/bs);
|
||||
const uint64_t dstnb2 = dstnb1*ne1;
|
||||
const uint64_t dstnb3 = dstnb2*ne2;
|
||||
|
||||
const uint64_t ne = ggml_nelements(tensor);
|
||||
|
||||
if (buf != nullptr) {
|
||||
// Memory is pinned, use as staging buffer
|
||||
std::vector<vk::BufferCopy> slices;
|
||||
|
||||
for (uint64_t i3 = 0; i3 < ne3; i3++) {
|
||||
for (uint64_t i2 = 0; i2 < ne2; i2++) {
|
||||
// Find longest contiguous slice
|
||||
if (ne1*nb1 == dstnb2) {
|
||||
slices.push_back({ buf_offset + i3*nb3 + i2*nb2, offset + i3*dstnb3 + i2*dstnb2, dstnb2 });
|
||||
} else {
|
||||
for (uint64_t i1 = 0; i1 < ne1; i1++) {
|
||||
if (ne0*nb0/bs == dstnb1) {
|
||||
slices.push_back({ buf_offset + i3*nb3 + i2*nb2 + i1*nb1, offset + i3*dstnb3 + i2*dstnb2 + i1*dstnb1, dstnb1 });
|
||||
} else {
|
||||
const uint64_t s_off = buf_offset + i3*nb3 + i2*nb2 + i1*nb1;
|
||||
const uint64_t d_off = offset + i3*dstnb3 + i2*dstnb2 + i1*dstnb1;
|
||||
for (uint64_t i0 = 0; i0 < ne0; i0++) {
|
||||
slices.push_back({ s_off + i0*nb0, d_off + i0*dstnb0, dstnb0 });
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
ggml_vk_sync_buffers(ctx, subctx);
|
||||
subctx->s->buffer->buf.copyBuffer(buf->buffer, dst->buffer, slices);
|
||||
return;
|
||||
}
|
||||
|
||||
if (!sync_staging) {
|
||||
GGML_ABORT("Asynchronous write to non-pinned memory not supported");
|
||||
}
|
||||
|
||||
// Staging buffer required
|
||||
vk_buffer& staging = ctx->device->sync_staging;
|
||||
const uint64_t copy_size = ts*ne/bs;
|
||||
ggml_vk_ensure_sync_staging_buffer(ctx->device, copy_size);
|
||||
VkBufferCopy buf_copy{ 0, offset, copy_size };
|
||||
|
||||
ggml_vk_sync_buffers(ctx, subctx);
|
||||
vkCmdCopyBuffer(subctx->s->buffer->buf, (VkBuffer)staging->buffer, (VkBuffer)dst->buffer, 1, &buf_copy);
|
||||
|
||||
for (uint64_t i3 = 0; i3 < ne3; i3++) {
|
||||
for (uint64_t i2 = 0; i2 < ne2; i2++) {
|
||||
// Find longest contiguous slice
|
||||
if (ne1*nb1 == dstnb2) {
|
||||
deferred_memcpy((uint8_t *)staging->ptr + i3*dstnb3 + i2*dstnb2, (const uint8_t *) tensor->data + buf_offset + i3*nb3 + i2*nb2, dstnb2, &subctx->in_memcpys);
|
||||
} else {
|
||||
for (uint64_t i1 = 0; i1 < ne1; i1++) {
|
||||
if (ne0*nb0/bs == dstnb1) {
|
||||
deferred_memcpy((uint8_t *)staging->ptr + i3*dstnb3 + i2*dstnb2 + i1*dstnb1, (const uint8_t *) tensor->data + buf_offset + i3*nb3 + i2*nb2 + i1*nb1, dstnb1, &subctx->in_memcpys);
|
||||
} else {
|
||||
const uint64_t s_off = buf_offset + i3*nb3 + i2*nb2 + i1*nb1;
|
||||
const uint64_t d_off = i3*dstnb3 + i2*dstnb2 + i1*dstnb1;
|
||||
for (uint64_t i0 = 0; i0 < ne0; i0++) {
|
||||
deferred_memcpy((uint8_t *)staging->ptr + d_off + i0*dstnb0, (const uint8_t *) tensor->data + s_off + i0*nb0, dstnb0, &subctx->in_memcpys);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
bool ggml_vk_buffer_write_2d_async(vk_context subctx, vk_buffer& dst, size_t offset, const void * src, size_t spitch, size_t dpitch, size_t width, size_t height, bool sync_staging) {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_write_2d_async(" << width << ", " << height << ")");
|
||||
// Check if src is pinned memory
|
||||
vk_buffer buf = nullptr;
|
||||
size_t buf_offset = 0;
|
||||
ggml_vk_host_get(dst->device, src, buf, buf_offset);
|
||||
|
||||
if (buf != nullptr) {
|
||||
// Memory is pinned, use as staging buffer
|
||||
std::vector<vk::BufferCopy> slices(1);
|
||||
if (width == spitch && width == dpitch) {
|
||||
// Only do single write if stride is equal
|
||||
slices[0].srcOffset = buf_offset;
|
||||
slices[0].dstOffset = offset;
|
||||
slices[0].size = width * height;
|
||||
} else {
|
||||
slices.resize(height);
|
||||
for (size_t i = 0; i < height; i++) {
|
||||
slices[i].srcOffset = buf_offset + i * spitch;
|
||||
slices[i].dstOffset = offset + i * dpitch;
|
||||
slices[i].size = width;
|
||||
}
|
||||
}
|
||||
|
||||
ggml_vk_sync_buffers(nullptr, subctx);
|
||||
subctx->s->buffer->buf.copyBuffer(buf->buffer, dst->buffer, slices);
|
||||
return true;
|
||||
}
|
||||
VK_LOG_DEBUG("STAGING");
|
||||
|
||||
if (!sync_staging) {
|
||||
// copy was not handled caller needs to fall back
|
||||
return false;
|
||||
}
|
||||
|
||||
// Staging buffer required
|
||||
const size_t staging_size = width * height;
|
||||
ggml_vk_ensure_sync_staging_buffer(dst->device, staging_size);
|
||||
|
||||
vk_buffer& staging_buffer = dst->device->sync_staging;
|
||||
|
||||
std::vector<vk::BufferCopy> slices(1);
|
||||
if (width == dpitch) {
|
||||
slices[0].srcOffset = 0;
|
||||
slices[0].dstOffset = offset;
|
||||
slices[0].size = staging_size;
|
||||
} else {
|
||||
slices.resize(height);
|
||||
for (size_t i = 0; i < height; i++) {
|
||||
slices[i].srcOffset = i * width;
|
||||
slices[i].dstOffset = offset + i * dpitch;
|
||||
slices[i].size = width;
|
||||
}
|
||||
}
|
||||
|
||||
ggml_vk_sync_buffers(nullptr, subctx);
|
||||
subctx->s->buffer->buf.copyBuffer(staging_buffer->buffer, dst->buffer, slices);
|
||||
|
||||
if (width == spitch) {
|
||||
deferred_memcpy((uint8_t *)staging_buffer->ptr, src, staging_size, &subctx->in_memcpys);
|
||||
} else {
|
||||
for (size_t i = 0; i < height; i++) {
|
||||
deferred_memcpy((uint8_t *)staging_buffer->ptr + i * width, (const uint8_t *) src + i * spitch, width, &subctx->in_memcpys);
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool ggml_vk_buffer_write_async(vk_context subctx, vk_buffer& dst, size_t offset, const void * src, size_t size, bool sync_staging) {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_write_async(" << size << ")");
|
||||
return ggml_vk_buffer_write_2d_async(subctx, dst, offset, src, size, size, size, 1, sync_staging);
|
||||
}
|
||||
|
||||
void ggml_vk_buffer_write_2d(vk_buffer& dst, size_t offset, const void * src, size_t spitch, size_t dpitch, size_t width, size_t height) {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_write_2d(" << width << ", " << height << ")");
|
||||
// Buffer is already mapped
|
||||
if(dst->memory_property_flags & vk::MemoryPropertyFlagBits::eHostVisible) {
|
||||
GGML_ASSERT(dst->memory_property_flags & vk::MemoryPropertyFlagBits::eHostCoherent);
|
||||
|
||||
if (width == spitch && width == dpitch) {
|
||||
memcpy((uint8_t *)dst->ptr + offset, src, width * height);
|
||||
} else {
|
||||
for (size_t i = 0; i < height; i++) {
|
||||
memcpy((uint8_t *)dst->ptr + offset + i * dpitch, (const uint8_t *) src + i * spitch, width);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
std::lock_guard<std::recursive_mutex> guard(dst->device->mutex);
|
||||
|
||||
vk_context subctx = ggml_vk_create_temporary_context(dst->device->transfer_queue->cmd_pool);
|
||||
ggml_vk_ctx_begin(dst->device, subctx);
|
||||
bool ret = ggml_vk_buffer_write_2d_async(subctx, dst, offset, src, spitch, dpitch, width, height, true);
|
||||
GGML_ASSERT(ret);
|
||||
ggml_vk_ctx_end(subctx);
|
||||
|
||||
for (auto& cpy : subctx->in_memcpys) {
|
||||
memcpy(cpy.dst, cpy.src, cpy.n);
|
||||
}
|
||||
|
||||
for (auto& mset : subctx->memsets) {
|
||||
memset(mset.dst, mset.val, mset.n);
|
||||
}
|
||||
|
||||
ggml_vk_submit(subctx, dst->device->fence);
|
||||
VK_CHECK(dst->device->device.waitForFences({ dst->device->fence }, true, UINT64_MAX), "vk_buffer_write_2d waitForFences", dst->device);
|
||||
dst->device->device.resetFences({ dst->device->fence });
|
||||
ggml_vk_queue_command_pools_cleanup(dst->device);
|
||||
}
|
||||
}
|
||||
|
||||
void ggml_vk_buffer_write(vk_buffer& dst, size_t offset, const void * src, size_t size) {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_write(" << size << ")");
|
||||
ggml_vk_buffer_write_2d(dst, offset, src, size, size, size, 1);
|
||||
}
|
||||
|
||||
bool ggml_vk_buffer_read_2d_async(vk_context subctx, vk_buffer& src, size_t offset, void * dst, size_t spitch, size_t dpitch, size_t width, size_t height, bool sync_staging) {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_read_2d_async(offset=" << offset << ", width=" << width << ", height=" << height << ")");
|
||||
GGML_ASSERT(width > 0);
|
||||
GGML_ASSERT(height > 0);
|
||||
GGML_ASSERT(src != nullptr);
|
||||
|
||||
// TODO: staging_offset is not used
|
||||
|
||||
// Check if dst is pinned memory
|
||||
vk_buffer buf = nullptr;
|
||||
size_t buf_offset = 0;
|
||||
ggml_vk_host_get(src->device, dst, buf, buf_offset);
|
||||
|
||||
std::vector<vk::BufferCopy> slices(1);
|
||||
if (width == spitch && width == dpitch) {
|
||||
// Only do single write if stride is equal
|
||||
slices[0].srcOffset = offset;
|
||||
slices[0].dstOffset = buf_offset;
|
||||
slices[0].size = width * height;
|
||||
} else {
|
||||
slices.resize(height);
|
||||
for (size_t i = 0; i < height; i++) {
|
||||
slices[i].srcOffset = offset + i * spitch;
|
||||
slices[i].dstOffset = buf_offset + i * dpitch;
|
||||
slices[i].size = width;
|
||||
}
|
||||
}
|
||||
|
||||
if (buf != nullptr) {
|
||||
// Memory is pinned, use as staging buffer
|
||||
ggml_vk_sync_buffers(nullptr, subctx);
|
||||
subctx->s->buffer->buf.copyBuffer(src->buffer, buf->buffer, slices);
|
||||
|
||||
return true;
|
||||
}
|
||||
VK_LOG_DEBUG("STAGING");
|
||||
|
||||
if (!sync_staging) {
|
||||
// copy was not handled caller needs to fall back
|
||||
return false;
|
||||
}
|
||||
|
||||
// Fall back to staging buffer
|
||||
const size_t staging_size = width * height;
|
||||
ggml_vk_ensure_sync_staging_buffer(src->device, staging_size);
|
||||
|
||||
vk_buffer& staging_buffer = src->device->sync_staging;
|
||||
|
||||
std::vector<vk::BufferCopy> staging_slices(1);
|
||||
if (width == spitch) {
|
||||
staging_slices[0].srcOffset = offset;
|
||||
staging_slices[0].dstOffset = 0;
|
||||
staging_slices[0].size = staging_size;
|
||||
} else {
|
||||
staging_slices.resize(height);
|
||||
for (size_t i = 0; i < height; i++) {
|
||||
staging_slices[i].srcOffset = offset + i * spitch;
|
||||
staging_slices[i].dstOffset = i * width;
|
||||
staging_slices[i].size = width;
|
||||
}
|
||||
}
|
||||
|
||||
ggml_vk_sync_buffers(nullptr, subctx);
|
||||
subctx->s->buffer->buf.copyBuffer(src->buffer, staging_buffer->buffer, staging_slices);
|
||||
|
||||
if (width == dpitch) {
|
||||
deferred_memcpy(dst, staging_buffer->ptr, staging_size, &subctx->out_memcpys);
|
||||
} else {
|
||||
for (size_t i = 0; i < height; i++) {
|
||||
deferred_memcpy((uint8_t *) dst + i * dpitch, (const uint8_t *) staging_buffer->ptr + i * width, width, &subctx->out_memcpys);
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
static bool ggml_vk_buffer_read_async(vk_context subctx, vk_buffer& src, size_t offset, void * dst, size_t size, bool sync_staging = false) {
|
||||
return ggml_vk_buffer_read_2d_async(subctx, src, offset, dst, size, size, size, 1, sync_staging);
|
||||
}
|
||||
|
||||
void ggml_vk_buffer_read_2d(vk_buffer& src, size_t offset, void * dst, size_t spitch, size_t dpitch, size_t width, size_t height) {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_read_2d(" << src->buffer << ", " << offset << ", " << width << ", " << height << ")");
|
||||
|
||||
// If the device is not an UMA device the memory is host-accessible through rebar. While writing
|
||||
// through PCIe is sufficient fast reading back data from PCIe is slower than going through
|
||||
// the HW device to host copy path.
|
||||
if(src->memory_property_flags & vk::MemoryPropertyFlagBits::eHostVisible && src->device->uma) {
|
||||
GGML_ASSERT(src->memory_property_flags & vk::MemoryPropertyFlagBits::eHostCoherent);
|
||||
|
||||
std::lock_guard<std::recursive_mutex> guard(src->device->mutex);
|
||||
vk_context subctx = ggml_vk_create_temporary_context(src->device->compute_queue->cmd_pool);
|
||||
ggml_vk_ctx_begin(src->device, subctx);
|
||||
subctx->s->buffer->buf.pipelineBarrier(
|
||||
vk::PipelineStageFlagBits::eComputeShader | vk::PipelineStageFlagBits::eTransfer,
|
||||
vk::PipelineStageFlagBits::eHost,
|
||||
{},
|
||||
{ { vk::AccessFlagBits::eShaderWrite | vk::AccessFlagBits::eTransferWrite,
|
||||
vk::AccessFlagBits::eHostRead } },
|
||||
{}, {});
|
||||
ggml_vk_ctx_end(subctx);
|
||||
ggml_vk_submit(subctx, src->device->fence);
|
||||
VK_CHECK(src->device->device.waitForFences({ src->device->fence }, true, UINT64_MAX),
|
||||
"vk_buffer_read_2d uma waitForFences", src->device);
|
||||
src->device->device.resetFences({ src->device->fence });
|
||||
ggml_vk_queue_command_pools_cleanup(src->device);
|
||||
|
||||
if (width == spitch && width == dpitch) {
|
||||
memcpy(dst, (const uint8_t *) src->ptr + offset, width * height);
|
||||
} else {
|
||||
for (size_t i = 0; i < height; i++) {
|
||||
memcpy((uint8_t *) dst + i * dpitch, (const uint8_t *) src->ptr + offset + i * spitch, width);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
std::lock_guard<std::recursive_mutex> guard(src->device->mutex);
|
||||
|
||||
vk_context subctx = ggml_vk_create_temporary_context(src->device->transfer_queue->cmd_pool);
|
||||
ggml_vk_ctx_begin(src->device, subctx);
|
||||
bool ret = ggml_vk_buffer_read_2d_async(subctx, src, offset, dst, spitch, dpitch, width, height, true);
|
||||
GGML_ASSERT(ret);
|
||||
ggml_vk_ctx_end(subctx);
|
||||
|
||||
ggml_vk_submit(subctx, src->device->fence);
|
||||
VK_CHECK(src->device->device.waitForFences({ src->device->fence }, true, UINT64_MAX), "vk_buffer_read_2d waitForFences", src->device);
|
||||
src->device->device.resetFences({ src->device->fence });
|
||||
ggml_vk_queue_command_pools_cleanup(src->device);
|
||||
|
||||
for (auto& cpy : subctx->out_memcpys) {
|
||||
memcpy(cpy.dst, cpy.src, cpy.n);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void ggml_vk_buffer_read(vk_buffer& src, size_t offset, void * dst, size_t size) {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_read(" << src->buffer << ", " << offset << ", " << size << ")");
|
||||
ggml_vk_buffer_read_2d(src, offset, dst, size, size, size, 1);
|
||||
}
|
||||
|
||||
void ggml_vk_buffer_copy_async(vk_context& ctx, vk_buffer& dst, size_t dst_offset, vk_buffer& src, size_t src_offset, size_t size) {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_copy_async(" << size << ")");
|
||||
// Make sure both buffers are on same device
|
||||
GGML_ASSERT(src->device == dst->device);
|
||||
|
||||
VkBufferCopy bc{ src_offset, dst_offset, size };
|
||||
|
||||
vkCmdCopyBuffer(ctx->s->buffer->buf, (VkBuffer)src->buffer, (VkBuffer)dst->buffer, 1, &bc);
|
||||
}
|
||||
|
||||
void ggml_vk_buffer_copy(vk_buffer& dst, size_t dst_offset, vk_buffer& src, size_t src_offset, size_t size) {
|
||||
if (src->device == dst->device) {
|
||||
std::lock_guard<std::recursive_mutex> guard(src->device->mutex);
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_copy(SINGLE_DEVICE, " << size << ")");
|
||||
// Copy within the device
|
||||
vk_context subctx = ggml_vk_create_temporary_context(src->device->transfer_queue->cmd_pool);
|
||||
ggml_vk_ctx_begin(src->device, subctx);
|
||||
ggml_vk_buffer_copy_async(subctx, dst, dst_offset, src, src_offset, size);
|
||||
ggml_vk_ctx_end(subctx);
|
||||
ggml_vk_submit(subctx, src->device->fence);
|
||||
VK_CHECK(src->device->device.waitForFences({ src->device->fence }, true, UINT64_MAX), "vk_buffer_copy waitForFences", src->device);
|
||||
src->device->device.resetFences({ src->device->fence });
|
||||
ggml_vk_queue_command_pools_cleanup(src->device);
|
||||
} else {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_copy(MULTI_DEVICE, " << size << ")");
|
||||
// Copy device to device
|
||||
ggml_vk_ensure_sync_staging_buffer(src->device, size);
|
||||
|
||||
// Copy to src staging buffer
|
||||
ggml_vk_buffer_copy(src->device->sync_staging, 0, src, src_offset, size);
|
||||
// Copy to dst buffer
|
||||
ggml_vk_buffer_write(dst, dst_offset, src->device->sync_staging->ptr, size);
|
||||
}
|
||||
}
|
||||
|
||||
void ggml_vk_buffer_memset_async(vk_context& ctx, vk_buffer& dst, size_t offset, uint32_t c, size_t size) {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_memset_async(" << offset << ", " << c << ", " << size << ")");
|
||||
|
||||
if (dst->memory_property_flags & vk::MemoryPropertyFlagBits::eHostVisible &&
|
||||
dst->device->uma) {
|
||||
deferred_memset((uint8_t*)dst->ptr + offset, c, size, &ctx->memsets);
|
||||
return;
|
||||
}
|
||||
|
||||
// Fall back to GPU fillBuffer for non-UMA or non-host-visible buffers
|
||||
ctx->s->buffer->buf.fillBuffer(dst->buffer, offset, size, c);
|
||||
}
|
||||
|
||||
void ggml_vk_buffer_memset(vk_buffer& dst, size_t offset, uint32_t c, size_t size) {
|
||||
VK_LOG_DEBUG("ggml_vk_buffer_memset(" << offset << ", " << c << ", " << size << ")");
|
||||
|
||||
if (dst->memory_property_flags & vk::MemoryPropertyFlagBits::eHostVisible &&
|
||||
dst->device->uma) {
|
||||
memset((uint8_t*)dst->ptr + offset, c, size);
|
||||
return;
|
||||
}
|
||||
|
||||
std::lock_guard<std::recursive_mutex> guard(dst->device->mutex);
|
||||
vk_context subctx = ggml_vk_create_temporary_context(dst->device->transfer_queue->cmd_pool);
|
||||
ggml_vk_ctx_begin(dst->device, subctx);
|
||||
subctx->s->buffer->buf.fillBuffer(dst->buffer, offset, size, c);
|
||||
ggml_vk_ctx_end(subctx);
|
||||
|
||||
ggml_vk_submit(subctx, dst->device->fence);
|
||||
VK_CHECK(dst->device->device.waitForFences({ dst->device->fence }, true, UINT64_MAX), "vk_memset waitForFences", dst->device);
|
||||
dst->device->device.resetFences({ dst->device->fence });
|
||||
ggml_vk_queue_command_pools_cleanup(dst->device);
|
||||
}
|
||||
|
||||
ggml_backend_buffer_i ggml_backend_vk_buffer_interface = {
|
||||
/* .free_buffer = */ ggml_backend_vk_buffer_free_buffer,
|
||||
/* .get_base = */ ggml_backend_vk_buffer_get_base,
|
||||
/* .init_tensor = */ ggml_backend_vk_buffer_init_tensor,
|
||||
/* .memset_tensor = */ ggml_backend_vk_buffer_memset_tensor,
|
||||
/* .set_tensor = */ ggml_backend_vk_buffer_set_tensor,
|
||||
/* .get_tensor = */ ggml_backend_vk_buffer_get_tensor,
|
||||
/* .set_tensor_2d = */ ggml_backend_vk_buffer_set_tensor_2d,
|
||||
/* .get_tensor_2d = */ ggml_backend_vk_buffer_get_tensor_2d,
|
||||
/* .cpy_tensor = */ ggml_backend_vk_buffer_cpy_tensor,
|
||||
/* .clear = */ ggml_backend_vk_buffer_clear,
|
||||
/* .reset = */ NULL,
|
||||
};
|
||||
|
||||
vk_buffer ggml_vk_buffer_from_host_ptr(vk_device & device, void * ptr, size_t size) {
|
||||
if (!device->external_memory_host) {
|
||||
return {};
|
||||
}
|
||||
|
||||
uintptr_t uptr = reinterpret_cast<uintptr_t>(ptr);
|
||||
if (uptr & (device->min_imported_host_pointer_alignment - 1)) {
|
||||
return {};
|
||||
}
|
||||
if (size & (device->min_imported_host_pointer_alignment - 1)) {
|
||||
return {};
|
||||
}
|
||||
|
||||
const vk::MemoryPropertyFlags property_flags = vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent | vk::MemoryPropertyFlagBits::eHostCached;
|
||||
|
||||
vk_buffer buf {};
|
||||
try {
|
||||
buf = ggml_vk_create_buffer(device, size, { property_flags }, ptr);
|
||||
} catch (vk::SystemError& e) {
|
||||
GGML_LOG_WARN("ggml_vulkan: Failed ggml_vk_create_buffer (%s)\n", e.what());
|
||||
}
|
||||
|
||||
return buf;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,282 @@
|
||||
#pragma once
|
||||
#include "ggml-vulkan-push-constants.h"
|
||||
|
||||
// shared globals
|
||||
extern ggml_backend_buffer_type_i ggml_backend_vk_buffer_type_interface;
|
||||
extern bool vk_memory_logger_enabled;
|
||||
extern bool vk_perf_logger_enabled;
|
||||
extern bool vk_perf_logger_concurrent;
|
||||
extern bool vk_enable_sync_logger;
|
||||
extern uint32_t vk_perf_logger_frequency;
|
||||
extern std::string vk_pipeline_stats_filter;
|
||||
extern void * const vk_ptr_base;
|
||||
extern vk_instance_t vk_instance;
|
||||
extern ggml_backend_buffer_i ggml_backend_vk_buffer_interface;
|
||||
|
||||
// instance
|
||||
vk_device ggml_vk_get_device(size_t idx);
|
||||
DispatchLoaderDynamic & ggml_vk_default_dispatcher();
|
||||
void ggml_vk_instance_init();
|
||||
void ggml_vk_init(ggml_backend_vk_context * ctx, size_t idx);
|
||||
int ggml_vk_get_device_count();
|
||||
void ggml_vk_get_device_description(int device, char * description, size_t description_size);
|
||||
bool ggml_vk_instance_layer_settings_available();
|
||||
bool ggml_vk_instance_portability_enumeration_ext_available(const std::vector<vk::ExtensionProperties>& instance_extensions);
|
||||
bool ggml_vk_instance_debug_utils_ext_available(const std::vector<vk::ExtensionProperties> & instance_extensions);
|
||||
bool ggml_vk_device_is_supported(const vk::PhysicalDevice & vkdev);
|
||||
bool ggml_vk_khr_cooperative_matrix_support(const vk::PhysicalDeviceProperties& props, const vk::PhysicalDeviceDriverProperties& driver_props, vk_device_architecture arch);
|
||||
uint32_t ggml_vk_intel_shader_core_count(const vk::PhysicalDevice& vkdev);
|
||||
bool ggml_vk_intel_windows_driver_in_range(uint32_t driver_version, uint32_t lower_major, uint32_t lower_minor, uint32_t upper_major, uint32_t upper_minor);
|
||||
|
||||
// shaders
|
||||
void ggml_vk_destroy_pipeline(vk::Device& device, vk_pipeline& pipeline);
|
||||
vk_fa_tuning_params get_fa_tuning_params(const vk_device& device, uint32_t hsk, uint32_t hsv, uint32_t n_rows, uint32_t n_kv, ggml_type k_type, ggml_type v_type, bool f32acc);
|
||||
vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool aligned, bool f32acc, bool use_mask, bool use_mask_opt, bool use_logit_softcap, bool use_sparse, ggml_type k_type, ggml_type v_type);
|
||||
uint32_t get_subgroup_size(const std::string &pipeline_name, const vk_device_architecture &arch);
|
||||
void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested = nullptr);
|
||||
bool ggml_vk_flash_attn_scalar_shmem_support(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool f32acc, ggml_type k_type, ggml_type v_type);
|
||||
bool ggml_vk_flash_attn_coopmat_shmem_support(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool f32acc, ggml_type k_type = GGML_TYPE_F16, ggml_type v_type = GGML_TYPE_F16);
|
||||
|
||||
// buffers
|
||||
vk_buffer ggml_vk_create_buffer_check(vk_device& device, size_t size, vk::MemoryPropertyFlags req_flags, vk::MemoryPropertyFlags fallback_flags = vk::MemoryPropertyFlags(0));
|
||||
vk_buffer ggml_vk_create_buffer_device(vk_device& device, size_t size);
|
||||
void ggml_vk_destroy_buffer(vk_buffer& buf);
|
||||
void * ggml_vk_host_malloc(vk_device& device, size_t size);
|
||||
void ggml_vk_host_free(vk_device& device, void* ptr);
|
||||
void ggml_vk_host_get(const vk_device& device, const void * ptr, vk_buffer& buf, size_t& buf_offset);
|
||||
void ggml_vk_ensure_sync_staging_buffer(vk_device& device, size_t size);
|
||||
void ggml_vk_ensure_sync_staging_buffer(ggml_backend_vk_context * ctx, size_t size);
|
||||
bool ggml_vk_buffer_write_2d_async(vk_context subctx, vk_buffer& dst, size_t offset, const void * src, size_t spitch, size_t dpitch, size_t width, size_t height, bool sync_staging = false);
|
||||
bool ggml_vk_buffer_write_async(vk_context subctx, vk_buffer& dst, size_t offset, const void * src, size_t size, bool sync_staging = false);
|
||||
void ggml_vk_buffer_write_2d(vk_buffer& dst, size_t offset, const void * src, size_t spitch, size_t dpitch, size_t width, size_t height);
|
||||
void ggml_vk_buffer_write(vk_buffer& dst, size_t offset, const void * src, size_t size);
|
||||
bool ggml_vk_buffer_read_2d_async(vk_context subctx, vk_buffer& src, size_t offset, void * dst, size_t spitch, size_t dpitch, size_t width, size_t height, bool sync_staging = false);
|
||||
void ggml_vk_buffer_read_2d(vk_buffer& src, size_t offset, void * dst, size_t spitch, size_t dpitch, size_t width, size_t height);
|
||||
void ggml_vk_buffer_read(vk_buffer& src, size_t offset, void * dst, size_t size);
|
||||
void ggml_vk_buffer_copy_async(vk_context& ctx, vk_buffer& dst, size_t dst_offset, vk_buffer& src, size_t src_offset, size_t size);
|
||||
void ggml_vk_buffer_copy(vk_buffer& dst, size_t dst_offset, vk_buffer& src, size_t src_offset, size_t size);
|
||||
void ggml_vk_buffer_memset_async(vk_context& ctx, vk_buffer& dst, size_t offset, uint32_t c, size_t size);
|
||||
void ggml_vk_buffer_memset(vk_buffer& dst, size_t offset, uint32_t c, size_t size);
|
||||
vk_buffer ggml_vk_buffer_from_host_ptr(vk_device & device, void * ptr, size_t size);
|
||||
|
||||
// pipelines
|
||||
uint64_t vk_tensor_offset(const ggml_tensor * tensor);
|
||||
uint32_t get_misalign_bytes(const ggml_backend_vk_context * ctx, const ggml_tensor * t);
|
||||
void ggml_vk_wait_for_fence(ggml_backend_vk_context * ctx);
|
||||
void ggml_pipeline_request_descriptor_sets(ggml_backend_vk_context *ctx, vk_pipeline& pipeline, uint32_t n);
|
||||
void ggml_pipeline_allocate_descriptor_sets(ggml_backend_vk_context * ctx);
|
||||
void ggml_vk_submit(vk_context& ctx, vk::Fence fence);
|
||||
uint32_t ggml_vk_find_queue_family_index(std::vector<vk::QueueFamilyProperties>& queue_family_props, const vk::QueueFlags& required, const vk::QueueFlags& avoid, int32_t compute_index, uint32_t min_num_queues);
|
||||
std::unique_ptr<vk_queue> ggml_vk_create_queue(vk_device& device, uint32_t queue_family_index, uint32_t queue_index, vk::PipelineStageFlags&& stage_flags, bool transfer_only);
|
||||
std::unique_ptr<vk_queue> ggml_vk_create_aliased_queue(vk_device& device, const std::unique_ptr<vk_queue>& source);
|
||||
vk_context ggml_vk_create_context(ggml_backend_vk_context * ctx, vk_command_pool& p);
|
||||
vk_context ggml_vk_create_temporary_context(vk_command_pool& p);
|
||||
void ggml_vk_command_pool_cleanup(vk_device& device, vk_command_pool& p);
|
||||
void ggml_vk_queue_command_pools_cleanup(vk_device& device);
|
||||
vk_subbuffer ggml_vk_subbuffer(const ggml_backend_vk_context* ctx, const vk_buffer& buf, size_t offset = 0);
|
||||
void ggml_vk_sync_buffers(ggml_backend_vk_context* ctx, vk_context& subctx);
|
||||
void ggml_vk_set_event(vk_context& ctx, vk::Event& event);
|
||||
void ggml_vk_wait_events(vk_context& ctx, std::vector<vk::Event>&& events);
|
||||
vk_subbuffer ggml_vk_tensor_subbuffer(const ggml_backend_vk_context * ctx, const ggml_tensor * tensor, bool allow_misalign = false);
|
||||
void ggml_vk_cmd_label_begin(vk::CommandBuffer buf, const char * name);
|
||||
void ggml_vk_ctx_end(vk_context& ctx);
|
||||
void ggml_vk_ctx_begin(vk_device& device, vk_context& subctx);
|
||||
vk_context ggml_vk_get_compute_ctx(ggml_backend_vk_context * ctx);
|
||||
vk_context ggml_vk_get_transfer_ctx(ggml_backend_vk_context * ctx);
|
||||
bool ggml_vk_submit_transfer_ctx(ggml_backend_vk_context * ctx);
|
||||
size_t ggml_vk_align_size(size_t width, size_t align);
|
||||
void deferred_memcpy(void * dst, const void * src, size_t size, std::vector<vk_staging_memcpy>* memcpys = nullptr);
|
||||
void deferred_memset(void * dst, uint32_t val, size_t size, std::vector<vk_staging_memset>* memsets = nullptr);
|
||||
|
||||
// matmul
|
||||
vk_pipeline ggml_vk_get_to_fp16(ggml_backend_vk_context * ctx, ggml_type type);
|
||||
void ggml_vk_matmul(ggml_backend_vk_context * ctx, vk_context& subctx, vk_pipeline& pipeline, vk_subbuffer&& a, vk_subbuffer&& b, vk_subbuffer&& d, vk_subbuffer&& split_k_buffer, uint32_t m, uint32_t n, uint32_t k, uint32_t stride_a, uint32_t stride_b, uint32_t stride_d, uint32_t batch_stride_a, uint32_t batch_stride_b, uint32_t batch_stride_d, uint32_t split_k, uint32_t batch, uint32_t ne02, uint32_t ne12, uint32_t broadcast2, uint32_t broadcast3, uint32_t padded_n);
|
||||
bool ggml_vk_dim01_contiguous(const ggml_tensor * tensor);
|
||||
vk_pipeline ggml_vk_get_cpy_pipeline(ggml_backend_vk_context * ctx, const ggml_tensor * src, const ggml_tensor * dst, ggml_type to);
|
||||
vk_pipeline ggml_vk_get_quantize_pipeline(ggml_backend_vk_context * ctx, ggml_type type);
|
||||
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_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);
|
||||
|
||||
// flash-attn
|
||||
void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v, const ggml_tensor * mask, const ggml_tensor * sinks, ggml_tensor * dst);
|
||||
|
||||
// operators
|
||||
void ggml_vk_cpy_to_contiguous(ggml_backend_vk_context * ctx, vk_context& subctx, vk_pipeline pipeline, const ggml_tensor * tensor, const vk_subbuffer & in, const vk_subbuffer & out);
|
||||
bool ggml_vk_can_use_fwht(const ggml_backend_vk_context * ctx, const ggml_tensor * src1, const ggml_tensor * dst);
|
||||
void ggml_vk_fwht(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src, ggml_tensor * dst);
|
||||
void ggml_vk_get_rows(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_get_rows_back(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_acc(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_multi_add(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_cgraph * cgraph, int node_idx);
|
||||
void ggml_vk_add(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_out_prod(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_sub(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_mul(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
int ggml_vk_unary_mul_op_index(ggml_unary_op op);
|
||||
void ggml_vk_unary_mul(ggml_backend_vk_context * ctx, vk_context& subctx, const struct ggml_cgraph * cgraph, int node_idx);
|
||||
void ggml_vk_div(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_add_id(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * src2, ggml_tensor * dst);
|
||||
void ggml_vk_rwkv_wkv6(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst);
|
||||
void ggml_vk_rwkv_wkv7(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst);
|
||||
void ggml_vk_gated_linear_attn(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst);
|
||||
void ggml_vk_lightning_indexer(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst);
|
||||
void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst);
|
||||
void ggml_vk_ssm_scan(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst);
|
||||
void ggml_vk_ssm_conv(ggml_backend_vk_context * ctx, vk_context& subctx, const struct ggml_cgraph * cgraph, int node_idx);
|
||||
void ggml_vk_opt_step_adamw(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst);
|
||||
void ggml_vk_opt_step_sgd(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * src2, ggml_tensor * dst);
|
||||
void ggml_vk_concat(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_upscale(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_scale(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_sqr(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_sqrt(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_add1(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_arange(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst);
|
||||
void ggml_vk_fill(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst);
|
||||
void ggml_vk_sin(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_cos(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_log(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_tri(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_diag(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_clamp(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_pad(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_pad_reflect_1d(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_roll(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_repeat(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_repeat_back(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_cpy(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_set_rows(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_silu_back(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_group_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
uint32_t ggml_vk_rms_partials_size(ggml_backend_vk_context * ctx, const ggml_tensor *node);
|
||||
void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const struct ggml_cgraph * cgraph, int node_idx, float * op_params);
|
||||
void ggml_vk_rms_norm_back(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_l2_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_unary(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_xielu(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_glu(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_diag_mask_inf(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_soft_max(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * src2, ggml_tensor * dst);
|
||||
void ggml_vk_soft_max_back(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_topk_moe(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_cgraph * cgraph, int node_idx);
|
||||
void ggml_vk_rope(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_cgraph * cgraph, int node_idx, bool backprop);
|
||||
void ggml_vk_argsort(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_topk(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_topk_qsa(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_cgraph * cgraph, int node_idx);
|
||||
void ggml_vk_sum(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_sum_rows(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_mean(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_cumsum(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_cross_entropy_loss(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst);
|
||||
void ggml_vk_cross_entropy_loss_back(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst);
|
||||
void ggml_vk_argmax(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_count_equal(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_solve_tri(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_im2col(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_im2col_3d(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_timestep_embedding(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_conv_transpose_1d(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_col2im_1d(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_snake_dispatch_fused(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_cgraph * cgraph, int node_idx);
|
||||
void ggml_vk_pool_1d(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_pool_2d(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
void ggml_vk_conv_2d(ggml_backend_vk_context * ctx, vk_context & subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_conv_3d(ggml_backend_vk_context * ctx, vk_context & subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_conv_2d_dw(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst);
|
||||
void ggml_vk_leaky_relu(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst);
|
||||
|
||||
// graph
|
||||
void ggml_vk_preallocate_buffers(ggml_backend_vk_context * ctx, vk_context subctx);
|
||||
bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgraph, int node_idx, ggml_tensor *node_begin, int node_idx_begin, bool last_node, bool almost_ready, bool submit);
|
||||
void ggml_vk_compute_forward(ggml_backend_vk_context* ctx, ggml_cgraph * cgraph, ggml_tensor* tensor, int tensor_idx, bool almost_ready);
|
||||
void ggml_vk_graph_cleanup(ggml_backend_vk_context * ctx);
|
||||
void ggml_vk_cleanup(ggml_backend_vk_context * ctx);
|
||||
void ggml_vk_synchronize(ggml_backend_vk_context * ctx);
|
||||
bool ggml_vk_is_empty(ggml_tensor * node);
|
||||
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);
|
||||
bool ggml_vk_can_fuse_ssm_conv(const ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx, int num_extra);
|
||||
bool ggml_vk_can_fuse_topk_moe(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx, topk_moe_mode mode);
|
||||
bool ggml_vk_can_fuse_topk_qsa(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx);
|
||||
bool ggml_vk_can_fuse_rope_set_rows(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx);
|
||||
bool ggml_vk_can_fuse_rms_norm_set_rows(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx);
|
||||
bool ggml_vk_can_fuse_snake(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx);
|
||||
bool ggml_vk_tensors_overlap(const ggml_tensor * a, const ggml_tensor * b, bool elementwise);
|
||||
bool ggml_vk_can_fuse_rms_norm_mul_rope(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx);
|
||||
uint32_t ggml_vk_fuse_multi_add(ggml_backend_vk_context * ctx, const struct ggml_cgraph * cgraph, int node_idx);
|
||||
void ggml_vk_graph_optimize(ggml_backend_t backend, struct ggml_cgraph * graph, struct ggml_backend_graph_optimize_params * params);
|
||||
|
||||
// backend
|
||||
bool ggml_backend_buffer_is_vk(ggml_backend_buffer_t buffer);
|
||||
void ggml_backend_vk_buffer_free_buffer(ggml_backend_buffer_t buffer);
|
||||
void * ggml_backend_vk_buffer_get_base(ggml_backend_buffer_t buffer);
|
||||
enum ggml_status ggml_backend_vk_buffer_init_tensor(ggml_backend_buffer_t buffer, ggml_tensor * tensor);
|
||||
void ggml_backend_vk_buffer_memset_tensor(ggml_backend_buffer_t buffer, ggml_tensor * tensor, uint8_t value, size_t offset, size_t size);
|
||||
void ggml_backend_vk_buffer_set_tensor(ggml_backend_buffer_t buffer, ggml_tensor * tensor, const void * data, size_t offset, size_t size);
|
||||
void ggml_backend_vk_buffer_set_tensor_2d(ggml_backend_buffer_t buffer, ggml_tensor * tensor, const void * data, size_t offset, size_t size, size_t n_copies, size_t stride_tensor, size_t stride_data);
|
||||
void ggml_backend_vk_buffer_get_tensor(ggml_backend_buffer_t buffer, const ggml_tensor * tensor, void * data, size_t offset, size_t size);
|
||||
void ggml_backend_vk_buffer_get_tensor_2d(ggml_backend_buffer_t buffer, const ggml_tensor * tensor, void * data, size_t offset, size_t size, size_t n_copies, size_t stride_tensor, size_t stride_data);
|
||||
bool ggml_backend_vk_buffer_cpy_tensor(ggml_backend_buffer_t buffer, const ggml_tensor * src, ggml_tensor * dst);
|
||||
void ggml_backend_vk_buffer_clear(ggml_backend_buffer_t buffer, uint8_t value);
|
||||
const char * ggml_backend_vk_buffer_type_name(ggml_backend_buffer_type_t buft);
|
||||
ggml_backend_buffer_t ggml_backend_vk_buffer_type_alloc_buffer(ggml_backend_buffer_type_t buft, size_t size);
|
||||
size_t ggml_backend_vk_buffer_type_get_alignment(ggml_backend_buffer_type_t buft);
|
||||
size_t ggml_backend_vk_buffer_type_get_max_size(ggml_backend_buffer_type_t buft);
|
||||
size_t ggml_backend_vk_buffer_type_get_alloc_size(ggml_backend_buffer_type_t buft, const ggml_tensor * tensor);
|
||||
void ggml_backend_vk_free(ggml_backend_t backend);
|
||||
ggml_backend_reg_t ggml_backend_vk_reg();
|
||||
|
||||
// debug
|
||||
int64_t ggml_vk_get_op_batch_size(const ggml_tensor * op);
|
||||
|
||||
// ggml-vulkan.cpp (residual)
|
||||
bool ggml_vk_lightning_indexer_k_type_supported(ggml_type type);
|
||||
void ggml_vk_print_device_fault_info(const vk_device& device);
|
||||
uint64_t ggml_vk_get_node_flops(const ggml_tensor * node);
|
||||
void ggml_vk_print_node_list(const ggml_cgraph * cgraph, int start, int end);
|
||||
void ggml_vk_print_device_lost_info(const vk_device& device);
|
||||
size_t ggml_vk_tensor_buffer_offset(const ggml_backend_vk_context * ctx, const ggml_tensor * t);
|
||||
size_t ggml_vk_descriptor_offset(size_t tensor_offset, size_t alignment, size_t type_size);
|
||||
uint32_t ggml_vk_concat_unit_size(ggml_type type);
|
||||
bool ggml_vk_concat_supported(const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * dst);
|
||||
|
||||
template <typename T>
|
||||
inline void ggml_vk_dispatch_pipeline(ggml_backend_vk_context* ctx, vk_context& subctx, vk_pipeline& pipeline, std::initializer_list<vk::DescriptorBufferInfo> const& descriptor_buffer_infos, const T &push_constants, std::array<uint32_t, 3> elements) {
|
||||
const uint32_t wg0 = CEIL_DIV(elements[0], pipeline->wg_denoms[0]);
|
||||
const uint32_t wg1 = CEIL_DIV(elements[1], pipeline->wg_denoms[1]);
|
||||
const uint32_t wg2 = CEIL_DIV(elements[2], pipeline->wg_denoms[2]);
|
||||
VK_LOG_DEBUG("ggml_vk_dispatch_pipeline(" << pipeline->name << ", {";
|
||||
for (auto& buffer : descriptor_buffer_infos) {
|
||||
std::cerr << "(" << buffer.buffer << ", " << buffer.offset << ", " << buffer.range << "), ";
|
||||
}
|
||||
std::cerr << "}, (" << wg0 << "," << wg1 << "," << wg2 << "))");
|
||||
GGML_ASSERT(wg0 <= ctx->device->properties.limits.maxComputeWorkGroupCount[0] &&
|
||||
wg1 <= ctx->device->properties.limits.maxComputeWorkGroupCount[1] &&
|
||||
wg2 <= ctx->device->properties.limits.maxComputeWorkGroupCount[2]);
|
||||
GGML_ASSERT(ctx->descriptor_set_idx < ctx->descriptor_sets.size());
|
||||
GGML_ASSERT(descriptor_buffer_infos.size() <= MAX_PARAMETER_COUNT);
|
||||
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 }, {});
|
||||
|
||||
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);
|
||||
subctx->s->buffer->buf.bindDescriptorSets(vk::PipelineBindPoint::eCompute,
|
||||
pipeline->layout,
|
||||
0,
|
||||
{ descriptor_set },
|
||||
{});
|
||||
{
|
||||
ggml_vk_debug_label dbg(subctx, pipeline->name, wg0, wg1, wg2);
|
||||
subctx->s->buffer->buf.dispatch(wg0, wg1, wg2);
|
||||
}
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user