From 9d5a7d0cbf221b6f0127d1fc2da7b3ddc1db9b1c Mon Sep 17 00:00:00 2001 From: hw-native-sys-bot Date: Mon, 13 Apr 2026 11:23:55 +0800 Subject: [PATCH 1/4] =?UTF-8?q?Update:=20rename=20upstream=20org=20ChaoWao?= =?UTF-8?q?=20=E2=86=92=20hw-native-sys=20(#530)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Update all references in GitHub workflow skills, issue templates, and shared library docs to reflect the repo transfer. Co-authored-by: wcwxy <26245345+ChaoWao@users.noreply.github.com> --- .claude/lib/github/README.md | 12 ++++++------ .claude/lib/github/common-issues.md | 2 +- .claude/lib/github/detect-permission.md | 4 ++-- .claude/lib/github/fetch-comments.md | 4 +++- .claude/lib/github/setup.md | 4 ++-- .claude/skills/checkout-pr/SKILL.md | 5 ++--- .github/ISSUE_TEMPLATE/config.yml | 2 +- 7 files changed, 17 insertions(+), 16 deletions(-) diff --git a/.claude/lib/github/README.md b/.claude/lib/github/README.md index 9969f41d6b..4588ee9e8f 100644 --- a/.claude/lib/github/README.md +++ b/.claude/lib/github/README.md @@ -8,10 +8,10 @@ Reusable procedures for GitHub PR workflows. | --------- | ----------- | ------- | | [setup](setup.md) | Authenticate and detect repository context | All | | [lookup-pr](lookup-pr.md) | Find PR by number, branch, or list all | All | -| [detect-permission](detect-permission.md) | Check push access to PR | address-pr-comments | -| [commit-and-push](commit-and-push.md) | Squash commits, rebase, and push | github-pr, address-pr-comments | -| [fetch-comments](fetch-comments.md) | Get unresolved PR review comments | address-pr-comments | -| [reply-and-resolve](reply-and-resolve.md) | Reply to and resolve review threads | address-pr-comments | +| [detect-permission](detect-permission.md) | Check push access to PR | fix-pr | +| [commit-and-push](commit-and-push.md) | Squash commits, rebase, and push | github-pr, fix-pr | +| [fetch-comments](fetch-comments.md) | Get unresolved PR review comments | fix-pr | +| [reply-and-resolve](reply-and-resolve.md) | Reply to and resolve review threads | fix-pr | | [branch-naming](branch-naming.md) | Generate branch name from commit | github-pr | | [common-issues](common-issues.md) | Troubleshooting reference | All | @@ -21,9 +21,9 @@ After running `detect-context`, these variables are available: | Variable | Description | Example (owner) | Example (fork) | | -------- | ----------- | --------------- | -------------- | -| `REPO_OWNER` | Origin repo owner | `ChaoWao` | `contributor` | +| `REPO_OWNER` | Origin repo owner | `hw-native-sys` | `contributor` | | `REPO_NAME` | Origin repo name | `simpler` | `simpler` | -| `PR_REPO_OWNER` | PR target owner | `ChaoWao` | `ChaoWao` | +| `PR_REPO_OWNER` | PR target owner | `hw-native-sys` | `hw-native-sys` | | `PR_REPO_NAME` | PR target name | `simpler` | `simpler` | | `DEFAULT_BRANCH` | Base branch name | `main` | `main` | | `BASE_REF` | Full base ref | `origin/main` | `upstream/main` | diff --git a/.claude/lib/github/common-issues.md b/.claude/lib/github/common-issues.md index eb92a4909c..d154db6fe2 100644 --- a/.claude/lib/github/common-issues.md +++ b/.claude/lib/github/common-issues.md @@ -49,7 +49,7 @@ Do NOT use `-f`/`-F` flags to pass GraphQL `$variables`. Bash mangles `$` signs ```bash # BAD — $owner/$repo/$number clash with bash -gh api graphql -f owner="ChaoWao" -f repo="simpler" -F number=276 \ +gh api graphql -f owner="hw-native-sys" -f repo="simpler" -F number=276 \ -f query='query($owner: String!, $repo: String!, $number: Int!) { ... }' # GOOD — inline values, no variables diff --git a/.claude/lib/github/detect-permission.md b/.claude/lib/github/detect-permission.md index 2f9d36e300..f6c2134e21 100644 --- a/.claude/lib/github/detect-permission.md +++ b/.claude/lib/github/detect-permission.md @@ -1,6 +1,6 @@ # Detect Permission -Used by `address-pr-comments` when working on someone else's PR. Determines push access and overrides `PUSH_REMOTE` if needed. +Used by `fix-pr` when working on someone else's PR. Determines push access and overrides `PUSH_REMOTE` if needed. ## Fetch PR Metadata @@ -62,4 +62,4 @@ esac ## Cleanup After Push -Do NOT remove the fork remote — it is reused by upstream tracking for auto-detection in `/github-pr` and `/address-pr-comments`. +Do NOT remove the fork remote — it is reused by upstream tracking for auto-detection in `/github-pr` and `/fix-pr`. diff --git a/.claude/lib/github/fetch-comments.md b/.claude/lib/github/fetch-comments.md index a2c9199f85..1e3477af31 100644 --- a/.claude/lib/github/fetch-comments.md +++ b/.claude/lib/github/fetch-comments.md @@ -33,7 +33,7 @@ query { }' ``` -Replace `OWNER`, `REPO`, `NUMBER` with actual values (e.g., `"ChaoWao"`, `"simpler"`, `276`). +Replace `OWNER`, `REPO`, `NUMBER` with actual values (e.g., `"hw-native-sys"`, `"simpler"`, `276`). Use `--jq` to filter unresolved threads: @@ -43,9 +43,11 @@ gh api graphql -f query='...' \ ``` **Limits:** + - `reviewThreads(first: 100)`: Fetches up to 100 threads. - `comments(first: 50)`: Fetches up to 50 comments per thread. Output: JSON array of unresolved threads. Each thread has: + - `id` — GraphQL node ID (for resolving via mutation) - `comments.nodes[].databaseId` — REST API ID (for replying) diff --git a/.claude/lib/github/setup.md b/.claude/lib/github/setup.md index a34a732e79..6264838b26 100644 --- a/.claude/lib/github/setup.md +++ b/.claude/lib/github/setup.md @@ -17,7 +17,7 @@ Detects repository role, remotes, and current state. Sets standard variables use ### Canonical Repo ```bash -UPSTREAM_OWNER="ChaoWao" +UPSTREAM_OWNER="hw-native-sys" UPSTREAM_NAME="simpler" UPSTREAM_REPO="$UPSTREAM_OWNER/$UPSTREAM_NAME" DEFAULT_BRANCH="main" @@ -90,7 +90,7 @@ fi | `ROLE` | `"owner"` | `"fork"` | | `BASE_REF` | `upstream/main` | `upstream/main` | | `PUSH_REMOTE` | `origin` | `origin` | -| `PR_REPO_OWNER` | `ChaoWao` | `ChaoWao` | +| `PR_REPO_OWNER` | `hw-native-sys` | `hw-native-sys` | | `PR_REPO_NAME` | `simpler` | `simpler` | | `PR_HEAD_PREFIX` | `""` | `"myuser:"` | | `DEFAULT_BRANCH` | `main` | `main` | diff --git a/.claude/skills/checkout-pr/SKILL.md b/.claude/skills/checkout-pr/SKILL.md index 070d4b6255..504e5bdba4 100644 --- a/.claude/skills/checkout-pr/SKILL.md +++ b/.claude/skills/checkout-pr/SKILL.md @@ -62,11 +62,10 @@ Run [checkout-fork-branch](../../lib/github/checkout-fork-branch.md) with `PUSH_ Print summary: -``` -Checked out PR #$PR_NUMBER ($PR_AUTHOR) +```text Remote: $FORK_REMOTE -> git@github.com:$HEAD_REPO_OWNER/$HEAD_REPO_NAME.git Branch: $LOCAL_BRANCH -> $FORK_REMOTE/$HEAD_BRANCH Push: PUSH_REMOTE=$FORK_REMOTE BRANCH_NAME=$LOCAL_BRANCH:$HEAD_BRANCH ``` -Remind user that `/github-pr` and `/address-pr-comments` will pick up the correct push target automatically (via upstream tracking). +Remind user that `/github-pr` and `/fix-pr` will pick up the correct push target automatically (via upstream tracking). diff --git a/.github/ISSUE_TEMPLATE/config.yml b/.github/ISSUE_TEMPLATE/config.yml index a7f9bb83e5..31cba68312 100644 --- a/.github/ISSUE_TEMPLATE/config.yml +++ b/.github/ISSUE_TEMPLATE/config.yml @@ -1,5 +1,5 @@ blank_issues_enabled: false contact_links: - name: General Questions & Discussions - url: https://github.com/ChaoWao/simpler/discussions + url: https://github.com/hw-native-sys/simpler/discussions about: For questions, ideas, or general discussion, please use GitHub Discussions From 6e118b2c005a6133157314eeffb0b32f0bd96052 Mon Sep 17 00:00:00 2001 From: hw-native-sys-bot Date: Mon, 13 Apr 2026 12:45:05 +0800 Subject: [PATCH 2/4] Fix: orch SO platform split, defer dlclose to deinit (#526) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Two independent fixes to orchestration SO handling on AICPU: 1. Orch SO file creation split by platform. mkstemps (libdevice_orch_XXXXXX.so) ensures per-call uniqueness on sim where multiple workers may share a process, but is not always available on AICPU device libc. Added platform interface create_orch_so_file so sim uses mkstemps + fchmod(0755) and onboard uses pid-based naming + open(...,0755) — sufficient since only one runtime runs per device process. 2. Deferred dlclose/unlink from run() to deinit(). Closing the SO handle at the end of run() made it impossible to re-run the orchestrator through repeated calls into the same executor. The handle is kept until deinit, which then unlinks the file. Applied to a2a3 aicpu_build_graph, a2a3 tensormap_and_ringbuffer, and a5 tensormap_and_ringbuffer. Co-authored-by: wcwxy <26245345+ChaoWao@users.noreply.github.com> --- .../platform/include/aicpu/orch_so_file.h | 48 +++++++++++++++++++ .../platform/onboard/aicpu/orch_so_file.cpp | 26 ++++++++++ src/a2a3/platform/sim/aicpu/orch_so_file.cpp | 46 ++++++++++++++++++ .../aicpu/aicpu_executor.cpp | 13 +++-- .../aicpu/aicpu_executor.cpp | 13 +++-- src/a5/platform/include/aicpu/orch_so_file.h | 48 +++++++++++++++++++ .../platform/onboard/aicpu/orch_so_file.cpp | 26 ++++++++++ src/a5/platform/sim/aicpu/orch_so_file.cpp | 46 ++++++++++++++++++ .../aicpu/aicpu_executor.cpp | 13 +++-- 9 files changed, 264 insertions(+), 15 deletions(-) create mode 100644 src/a2a3/platform/include/aicpu/orch_so_file.h create mode 100644 src/a2a3/platform/onboard/aicpu/orch_so_file.cpp create mode 100644 src/a2a3/platform/sim/aicpu/orch_so_file.cpp create mode 100644 src/a5/platform/include/aicpu/orch_so_file.h create mode 100644 src/a5/platform/onboard/aicpu/orch_so_file.cpp create mode 100644 src/a5/platform/sim/aicpu/orch_so_file.cpp diff --git a/src/a2a3/platform/include/aicpu/orch_so_file.h b/src/a2a3/platform/include/aicpu/orch_so_file.h new file mode 100644 index 0000000000..a305ab8fa7 --- /dev/null +++ b/src/a2a3/platform/include/aicpu/orch_so_file.h @@ -0,0 +1,48 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +/** + * @file orch_so_file.h + * @brief Orchestration SO File Creation Interface for AICPU + * + * Creates a writable, executable-mode file under a candidate directory for + * staging the device orchestration shared library prior to dlopen. + * + * Platform Support: + * - a2a3 (onboard): pid-based naming via open() with mode 0755. AICPU device + * libc may not provide mkstemps, and only one runtime runs per device process. + * - a2a3sim (simulation): mkstemps() with fchmod(0755). Multiple sim workers + * can share a process, so names must be unique per call. + */ + +#ifndef PLATFORM_AICPU_ORCH_SO_FILE_H_ +#define PLATFORM_AICPU_ORCH_SO_FILE_H_ + +#include +#include + +/** + * Create a unique orchestration SO file under `dir`. + * + * On success, writes the full chosen path into `out_path` (null-terminated) + * and returns an open writable fd (caller must close). Permissions are set + * so the file is executable (0755) and suitable for dlopen. + * + * On failure (path too long, directory not writable, etc.), returns -1. + * Caller is expected to try the next candidate directory. + * + * @param dir Candidate directory (e.g. "/tmp") + * @param out_path Buffer that receives the full file path on success + * @param out_path_size Size of `out_path` in bytes + * @return Open writable fd on success, -1 on failure + */ +int32_t create_orch_so_file(const char *dir, char *out_path, size_t out_path_size); + +#endif // PLATFORM_AICPU_ORCH_SO_FILE_H_ diff --git a/src/a2a3/platform/onboard/aicpu/orch_so_file.cpp b/src/a2a3/platform/onboard/aicpu/orch_so_file.cpp new file mode 100644 index 0000000000..322cb7dccc --- /dev/null +++ b/src/a2a3/platform/onboard/aicpu/orch_so_file.cpp @@ -0,0 +1,26 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +#include "aicpu/orch_so_file.h" + +#include +#include + +#include + +int32_t create_orch_so_file(const char *dir, char *out_path, size_t out_path_size) { + // Pid-based naming: AICPU device libc may lack mkstemps, and only one + // runtime runs per device process, so pid uniqueness is sufficient. + int32_t written = snprintf(out_path, out_path_size, "%s/libdevice_orch_%d.so", dir, getpid()); + if (written < 0 || static_cast(written) >= out_path_size) { + return -1; + } + return open(out_path, O_WRONLY | O_CREAT | O_TRUNC, 0755); +} diff --git a/src/a2a3/platform/sim/aicpu/orch_so_file.cpp b/src/a2a3/platform/sim/aicpu/orch_so_file.cpp new file mode 100644 index 0000000000..4da92d7de1 --- /dev/null +++ b/src/a2a3/platform/sim/aicpu/orch_so_file.cpp @@ -0,0 +1,46 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// mkstemps is a BSD extension: glibc gates it behind _DEFAULT_SOURCE (hidden +// under strict -std=c++17 which sets __STRICT_ANSI__). macOS libc exposes it +// unconditionally, so this only matters on Linux. +#ifndef _DEFAULT_SOURCE +#define _DEFAULT_SOURCE +#endif + +#include "aicpu/orch_so_file.h" + +#include +#include +#include +#include + +#include + +int32_t create_orch_so_file(const char *dir, char *out_path, size_t out_path_size) { + // mkstemps: multiple sim workers can share a process, so names must be + // unique per call. The "XXXXXX" template is replaced in-place. + int32_t written = snprintf(out_path, out_path_size, "%s/libdevice_orch_XXXXXX.so", dir); + if (written < 0 || static_cast(written) >= out_path_size) { + return -1; + } + + constexpr int32_t kSuffixLen = 3; // strlen(".so") + int32_t fd = mkstemps(out_path, kSuffixLen); + if (fd < 0) { + return -1; + } + if (fchmod(fd, 0755) != 0) { + close(fd); + unlink(out_path); + return -1; + } + return fd; +} diff --git a/src/a2a3/runtime/aicpu_build_graph/aicpu/aicpu_executor.cpp b/src/a2a3/runtime/aicpu_build_graph/aicpu/aicpu_executor.cpp index c6abcd0aee..7830cff210 100644 --- a/src/a2a3/runtime/aicpu_build_graph/aicpu/aicpu_executor.cpp +++ b/src/a2a3/runtime/aicpu_build_graph/aicpu/aicpu_executor.cpp @@ -9,7 +9,6 @@ * ----------------------------------------------------------------------------------------------------------- */ #include -#include #include #include @@ -25,6 +24,7 @@ #include "aicpu/device_log.h" #include "aicpu/device_time.h" +#include "aicpu/orch_so_file.h" #include "pto2_dispatch_payload.h" #include "runtime.h" #include "spin_hint.h" @@ -1609,8 +1609,7 @@ int32_t AicpuExecutor::run(Runtime *runtime) { const int32_t num_candidates = sizeof(candidate_dirs) / sizeof(candidate_dirs[0]); for (int32_t i = 0; i < num_candidates && !file_created; i++) { - snprintf(so_path, sizeof(so_path), "%s/libdevice_orch_%d.so", candidate_dirs[i], getpid()); - int32_t fd = open(so_path, O_WRONLY | O_CREAT | O_TRUNC, 0755); + int32_t fd = create_orch_so_file(candidate_dirs[i], so_path, sizeof(so_path)); if (fd < 0) { DEV_INFO( "Thread %d: Cannot create SO at %s (errno=%d), trying next path", thread_idx, so_path, errno @@ -1973,8 +1972,6 @@ int32_t AicpuExecutor::run(Runtime *runtime) { // Destroy PTO2 runtime and close orchestration SO (moved from orchestrator path) if (!runtime->get_orch_built_on_host() && orch_so_handle_ != nullptr) { pto2_runtime_destroy(rt); - dlclose(orch_so_handle_); - unlink(orch_so_path_); } DEV_ALWAYS("Thread %d: Last thread, marking executor finished", thread_idx); } @@ -2030,6 +2027,12 @@ void AicpuExecutor::deinit(Runtime *runtime) { // Reset orchestration SO state (handle freed by last thread before deinit) orch_func_ = nullptr; orch_args_cached_ = nullptr; + if (orch_so_handle_ != nullptr) { + dlclose(orch_so_handle_); + } + if (orch_so_path_[0] != '\0') { + unlink(orch_so_path_); + } orch_so_handle_ = nullptr; orch_so_path_[0] = '\0'; diff --git a/src/a2a3/runtime/tensormap_and_ringbuffer/aicpu/aicpu_executor.cpp b/src/a2a3/runtime/tensormap_and_ringbuffer/aicpu/aicpu_executor.cpp index c6de08e4eb..08b8ea91ea 100644 --- a/src/a2a3/runtime/tensormap_and_ringbuffer/aicpu/aicpu_executor.cpp +++ b/src/a2a3/runtime/tensormap_and_ringbuffer/aicpu/aicpu_executor.cpp @@ -9,7 +9,6 @@ * ----------------------------------------------------------------------------------------------------------- */ #include -#include #include #include @@ -25,6 +24,7 @@ #include "aicpu/device_log.h" #include "aicpu/device_time.h" +#include "aicpu/orch_so_file.h" #include "pto2_dispatch_payload.h" #include "runtime.h" #include "spin_hint.h" @@ -2253,8 +2253,7 @@ int32_t AicpuExecutor::run(Runtime *runtime) { const int32_t num_candidates = sizeof(candidate_dirs) / sizeof(candidate_dirs[0]); for (int32_t i = 0; i < num_candidates && !file_created; i++) { - snprintf(so_path, sizeof(so_path), "%s/libdevice_orch_%d.so", candidate_dirs[i], getpid()); - int32_t fd = open(so_path, O_WRONLY | O_CREAT | O_TRUNC, 0755); + int32_t fd = create_orch_so_file(candidate_dirs[i], so_path, sizeof(so_path)); if (fd < 0) { DEV_INFO( "Thread %d: Cannot create SO at %s (errno=%d), trying next path", thread_idx, so_path, errno @@ -2671,8 +2670,6 @@ int32_t AicpuExecutor::run(Runtime *runtime) { orch_bind_runtime_(nullptr); } pto2_runtime_destroy(rt); - dlclose(orch_so_handle_); - unlink(orch_so_path_); } } @@ -2726,6 +2723,12 @@ void AicpuExecutor::deinit(Runtime *runtime) { orch_func_ = nullptr; orch_bind_runtime_ = nullptr; orch_args_cached_ = nullptr; + if (orch_so_handle_ != nullptr) { + dlclose(orch_so_handle_); + } + if (orch_so_path_[0] != '\0') { + unlink(orch_so_path_); + } orch_so_handle_ = nullptr; orch_so_path_[0] = '\0'; diff --git a/src/a5/platform/include/aicpu/orch_so_file.h b/src/a5/platform/include/aicpu/orch_so_file.h new file mode 100644 index 0000000000..40bec7411b --- /dev/null +++ b/src/a5/platform/include/aicpu/orch_so_file.h @@ -0,0 +1,48 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +/** + * @file orch_so_file.h + * @brief Orchestration SO File Creation Interface for AICPU + * + * Creates a writable, executable-mode file under a candidate directory for + * staging the device orchestration shared library prior to dlopen. + * + * Platform Support: + * - a5 (onboard): pid-based naming via open() with mode 0755. AICPU device + * libc may not provide mkstemps, and only one runtime runs per device process. + * - a5sim (simulation): mkstemps() with fchmod(0755). Multiple sim workers + * can share a process, so names must be unique per call. + */ + +#ifndef PLATFORM_AICPU_ORCH_SO_FILE_H_ +#define PLATFORM_AICPU_ORCH_SO_FILE_H_ + +#include +#include + +/** + * Create a unique orchestration SO file under `dir`. + * + * On success, writes the full chosen path into `out_path` (null-terminated) + * and returns an open writable fd (caller must close). Permissions are set + * so the file is executable (0755) and suitable for dlopen. + * + * On failure (path too long, directory not writable, etc.), returns -1. + * Caller is expected to try the next candidate directory. + * + * @param dir Candidate directory (e.g. "/tmp") + * @param out_path Buffer that receives the full file path on success + * @param out_path_size Size of `out_path` in bytes + * @return Open writable fd on success, -1 on failure + */ +int32_t create_orch_so_file(const char *dir, char *out_path, size_t out_path_size); + +#endif // PLATFORM_AICPU_ORCH_SO_FILE_H_ diff --git a/src/a5/platform/onboard/aicpu/orch_so_file.cpp b/src/a5/platform/onboard/aicpu/orch_so_file.cpp new file mode 100644 index 0000000000..322cb7dccc --- /dev/null +++ b/src/a5/platform/onboard/aicpu/orch_so_file.cpp @@ -0,0 +1,26 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +#include "aicpu/orch_so_file.h" + +#include +#include + +#include + +int32_t create_orch_so_file(const char *dir, char *out_path, size_t out_path_size) { + // Pid-based naming: AICPU device libc may lack mkstemps, and only one + // runtime runs per device process, so pid uniqueness is sufficient. + int32_t written = snprintf(out_path, out_path_size, "%s/libdevice_orch_%d.so", dir, getpid()); + if (written < 0 || static_cast(written) >= out_path_size) { + return -1; + } + return open(out_path, O_WRONLY | O_CREAT | O_TRUNC, 0755); +} diff --git a/src/a5/platform/sim/aicpu/orch_so_file.cpp b/src/a5/platform/sim/aicpu/orch_so_file.cpp new file mode 100644 index 0000000000..4da92d7de1 --- /dev/null +++ b/src/a5/platform/sim/aicpu/orch_so_file.cpp @@ -0,0 +1,46 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// mkstemps is a BSD extension: glibc gates it behind _DEFAULT_SOURCE (hidden +// under strict -std=c++17 which sets __STRICT_ANSI__). macOS libc exposes it +// unconditionally, so this only matters on Linux. +#ifndef _DEFAULT_SOURCE +#define _DEFAULT_SOURCE +#endif + +#include "aicpu/orch_so_file.h" + +#include +#include +#include +#include + +#include + +int32_t create_orch_so_file(const char *dir, char *out_path, size_t out_path_size) { + // mkstemps: multiple sim workers can share a process, so names must be + // unique per call. The "XXXXXX" template is replaced in-place. + int32_t written = snprintf(out_path, out_path_size, "%s/libdevice_orch_XXXXXX.so", dir); + if (written < 0 || static_cast(written) >= out_path_size) { + return -1; + } + + constexpr int32_t kSuffixLen = 3; // strlen(".so") + int32_t fd = mkstemps(out_path, kSuffixLen); + if (fd < 0) { + return -1; + } + if (fchmod(fd, 0755) != 0) { + close(fd); + unlink(out_path); + return -1; + } + return fd; +} diff --git a/src/a5/runtime/tensormap_and_ringbuffer/aicpu/aicpu_executor.cpp b/src/a5/runtime/tensormap_and_ringbuffer/aicpu/aicpu_executor.cpp index 64d13a3b11..de2c48be25 100644 --- a/src/a5/runtime/tensormap_and_ringbuffer/aicpu/aicpu_executor.cpp +++ b/src/a5/runtime/tensormap_and_ringbuffer/aicpu/aicpu_executor.cpp @@ -9,7 +9,6 @@ * ----------------------------------------------------------------------------------------------------------- */ #include -#include #include #include @@ -25,6 +24,7 @@ #include "aicpu/device_log.h" #include "aicpu/device_time.h" +#include "aicpu/orch_so_file.h" #include "pto2_dispatch_payload.h" #include "runtime.h" #include "spin_hint.h" @@ -2226,8 +2226,7 @@ int32_t AicpuExecutor::run(Runtime *runtime) { const int32_t num_candidates = sizeof(candidate_dirs) / sizeof(candidate_dirs[0]); for (int32_t i = 0; i < num_candidates && !file_created; i++) { - snprintf(so_path, sizeof(so_path), "%s/libdevice_orch_%d.so", candidate_dirs[i], getpid()); - int32_t fd = open(so_path, O_WRONLY | O_CREAT | O_TRUNC, 0755); + int32_t fd = create_orch_so_file(candidate_dirs[i], so_path, sizeof(so_path)); if (fd < 0) { DEV_INFO( "Thread %d: Cannot create SO at %s (errno=%d), trying next path", thread_idx, so_path, errno @@ -2642,8 +2641,6 @@ int32_t AicpuExecutor::run(Runtime *runtime) { orch_bind_runtime_(nullptr); } pto2_runtime_destroy(rt); - dlclose(orch_so_handle_); - unlink(orch_so_path_); } } @@ -2697,6 +2694,12 @@ void AicpuExecutor::deinit(Runtime *runtime) { orch_func_ = nullptr; orch_bind_runtime_ = nullptr; orch_args_cached_ = nullptr; + if (orch_so_handle_ != nullptr) { + dlclose(orch_so_handle_); + } + if (orch_so_path_[0] != '\0') { + unlink(orch_so_path_); + } orch_so_handle_ = nullptr; orch_so_path_[0] = '\0'; From 9735785f640a733ecc8769a5e1195b1e28dd058c Mon Sep 17 00:00:00 2001 From: chenshengxin Date: Thu, 9 Apr 2026 20:32:35 +0800 Subject: [PATCH 3/4] Add: SPMD paged attention example with dual-vector softmax Implements a complete paged attention kernel using SPMD parallelism under the tensormap_and_ringbuffer runtime. The pipeline consists of QK matmul, softmax prepare, PV matmul, and online update stages with dual AIV lanes processing 8-row sub-tiles each for the softmax and accumulation phases. --- .../spmd_paged_attention/golden.py | 79 ++++++ .../kernels/aic/aic_hub.cpp | 29 ++ .../kernels/aic/aic_pv_matmul.cpp | 152 +++++++++++ .../kernels/aic/aic_qk_matmul.cpp | 158 +++++++++++ .../kernels/aiv/aiv_hub.cpp | 27 ++ .../kernels/aiv/aiv_online_update.cpp | 258 ++++++++++++++++++ .../kernels/aiv/aiv_softmax_prepare.cpp | 192 +++++++++++++ .../kernels/kernel_config.py | 85 ++++++ .../spmd_paged_attention_orch.cpp | 239 ++++++++++++++++ 9 files changed, 1219 insertions(+) create mode 100644 examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py create mode 100644 examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_hub.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_hub.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py create mode 100644 examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py new file mode 100644 index 0000000000..3847437a69 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py @@ -0,0 +1,79 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""SPMD Paged Attention Golden - tensormap_and_ringbuffer example (small scale, bfloat16). + +Uses SPMD parallelism: each block handles one (batch, q_tile) position. +Kernels use get_block_idx() to determine their work slice. +""" + +from paged_attention_golden import ( + compute_golden, # noqa: F401 + run_golden_test, +) +from paged_attention_golden import generate_inputs as _generate_inputs + +__outputs__ = ["out"] + +RTOL = 1e-2 +ATOL = 1e-2 + +ALL_CASES = { + "Case1": { + "batch": 1, + "num_heads": 16, + "kv_head_num": 1, + "head_dim": 16, + "block_size": 16, + "context_len": 33, + "max_model_len": 256, + "dtype": "bfloat16", + }, + "Case2": { + "batch": 1, + "num_heads": 16, + "kv_head_num": 1, + "head_dim": 16, + "block_size": 16, + "context_len": 128, + "max_model_len": 256, + "dtype": "bfloat16", + }, + "CaseVarSeq2": { + "batch": 2, + "num_heads": 16, + "kv_head_num": 1, + "head_dim": 16, + "block_size": 16, + "context_len": 33, + "context_lens_list": [33, 17], + "max_model_len": 256, + "dtype": "bfloat16", + }, + "CaseVarSeq4": { + "batch": 4, + "num_heads": 16, + "kv_head_num": 1, + "head_dim": 16, + "block_size": 16, + "context_len": 64, + "context_lens_list": [33, 64, 48, 15], + "max_model_len": 256, + "dtype": "bfloat16", + }, +} + +DEFAULT_CASE = "Case1" + + +def generate_inputs(params: dict) -> list: + return _generate_inputs(params) + + +if __name__ == "__main__": + run_golden_test(ALL_CASES, DEFAULT_CASE, generate_inputs, label="SPMD Paged Attention") diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_hub.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_hub.cpp new file mode 100644 index 0000000000..eb602e0bd6 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_hub.cpp @@ -0,0 +1,29 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// AIC Hub Kernel - No-op stub used as the AIC slot of MIX (AIC+AIV0+AIV1) tasks +// when the real work happens only on the two AIVs (softmax, online update). +// Pairing an idle AIC with two active AIVs forces the scheduler to allocate a +// full cluster, which is what enables the two AIV lanes to run in parallel. + +#include +#include + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) {} diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp new file mode 100644 index 0000000000..f2e2652e2b --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp @@ -0,0 +1,152 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD PV Matmul: pij(M, K) @ vj(K, N) -> oi_new(M, N) +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// Each block computes one 16x16 matmul using paged V cache. +// +// Args: +// args[0] = pij Tensor* (spmd_blocks*Q_TILE, block_size) data_type +// args[1] = value_cache Tensor* (kv_total_rows, head_dim) bf16 +// args[2] = block_table Tensor* (batch, max_blocks_per_req) int32 +// args[3] = context_lens Tensor* (batch,) int32 +// args[4] = oi_new Tensor* (spmd_blocks*Q_TILE, head_dim) float32 [output] +// args[5] = bn scalar: current KV block index +// args[6] = num_heads scalar +// args[7] = head_dim scalar +// args[8] = block_size scalar +// args[9] = max_num_blocks_per_req scalar +// args[10] = q_loop scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +static constexpr int M = 16; +static constexpr int K = 16; +static constexpr int N = 16; + +template +static __aicore__ void pv_matmul_spmd(__gm__ bfloat16_t *pij_addr, __gm__ bfloat16_t *vj_addr, __gm__ float *oi_addr) { + using GlobalA = GlobalTensor, Stride>; + using GlobalB = GlobalTensor, Stride>; + using GlobalOut = GlobalTensor, Stride>; + + GlobalA pijGlobal(pij_addr); + GlobalB vjGlobal(vj_addr); + GlobalOut oiGlobal(oi_addr); + + using TileMatA = Tile; + using TileMatB = Tile; + + using LeftTile = TileLeft; + using RightTile = TileRight; + using AccTile = TileAcc; + + TileMatA aMatTile; + TileMatB bMatTile; + TASSIGN(aMatTile, 0x0); + TASSIGN(bMatTile, 0x20000); + + LeftTile aTile; + RightTile bTile; + AccTile cTile; + TASSIGN(aTile, 0x0); + TASSIGN(bTile, 0x0); + TASSIGN(cTile, 0x0); + + TLOAD(aMatTile, pijGlobal); + TLOAD(bMatTile, vjGlobal); + + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + + TMOV(aTile, aMatTile); + TMOV(bTile, bMatTile); + + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + + TMATMUL(cTile, aTile, bTile); + + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + + TSTORE(oiGlobal, cTile); + + set_flag(PIPE_FIX, PIPE_S, EVENT_ID7); + wait_flag(PIPE_FIX, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *pij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *value_cache_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *block_table_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *oi_new_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + + int64_t bn = static_cast(args[5]); + int64_t num_heads = static_cast(args[6]); + int64_t head_dim = static_cast(args[7]); + int64_t block_size = static_cast(args[8]); + int64_t max_blocks_per_req = static_cast(args[9]); + int64_t q_loop = static_cast(args[10]); + + int32_t block_idx = get_block_idx(args); + int64_t batch_idx = block_idx / q_loop; + + // Check if this batch has data at this KV block + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t bn_this_batch = (cur_seq + block_size - 1) / block_size; + + // Output pointer for this block's oi_new slice + __gm__ float *oi_addr = + reinterpret_cast<__gm__ float *>(oi_new_t->buffer.addr) + oi_new_t->start_offset + block_idx * M * head_dim; + + if (bn >= bn_this_batch) { + for (int i = 0; i < M * static_cast(head_dim); i++) { + oi_addr[i] = 0.0f; + } + return; + } + + // Look up physical block index + __gm__ int32_t *bt_ptr = + reinterpret_cast<__gm__ int32_t *>(block_table_t->buffer.addr) + block_table_t->start_offset; + int64_t phys_block = static_cast(bt_ptr[batch_idx * max_blocks_per_req + bn]); + + // pij offset: block_idx * Q_TILE * block_size + int64_t pij_offset = block_idx * M * block_size; + __gm__ bfloat16_t *pij_addr = + reinterpret_cast<__gm__ bfloat16_t *>(pij_t->buffer.addr) + pij_t->start_offset + pij_offset; + + // Value offset: phys_block * block_size * head_dim + int64_t v_offset = phys_block * block_size * head_dim; + __gm__ bfloat16_t *vj_addr = + reinterpret_cast<__gm__ bfloat16_t *>(value_cache_t->buffer.addr) + value_cache_t->start_offset + v_offset; + + pv_matmul_spmd(pij_addr, vj_addr, oi_addr); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp new file mode 100644 index 0000000000..bdd0844bf8 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp @@ -0,0 +1,158 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD QK Matmul: qi(M, K) @ kj.T(K, N) -> sij(M, N) +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// Each block computes one 16x16 matmul using paged KV. +// +// Args: +// args[0] = query Tensor* (batch*num_heads, head_dim) bf16 +// args[1] = key_cache Tensor* (kv_total_rows, head_dim) bf16 +// args[2] = block_table Tensor* (batch, max_blocks_per_req) int32 +// args[3] = context_lens Tensor* (batch,) int32 +// args[4] = sij Tensor* (spmd_blocks*Q_TILE, block_size) float32 [output] +// args[5] = bn scalar: current KV block index +// args[6] = num_heads scalar +// args[7] = head_dim scalar +// args[8] = block_size scalar +// args[9] = max_num_blocks_per_req scalar +// args[10] = q_loop scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +static constexpr int M = 16; +static constexpr int K = 16; +static constexpr int N = 16; + +template +static __aicore__ void +qk_matmul_spmd(__gm__ bfloat16_t *qi_addr, __gm__ bfloat16_t *kj_addr, __gm__ float *sij_addr) { + using GlobalA = GlobalTensor, Stride>; + using GlobalB = + GlobalTensor, Stride, Layout::DN>; + using GlobalOut = GlobalTensor, Stride>; + + GlobalA qiGlobal(qi_addr); + GlobalB kjGlobal(kj_addr); + GlobalOut sijGlobal(sij_addr); + + using TileMatA = Tile; + using TileMatB = Tile; + + using LeftTile = TileLeft; + using RightTile = TileRight; + using AccTile = TileAcc; + + TileMatA aMatTile; + TileMatB bMatTile; + TASSIGN(aMatTile, 0x0); + TASSIGN(bMatTile, 0x20000); + + LeftTile aTile; + RightTile bTile; + AccTile cTile; + TASSIGN(aTile, 0x0); + TASSIGN(bTile, 0x0); + TASSIGN(cTile, 0x0); + + TLOAD(aMatTile, qiGlobal); + TLOAD(bMatTile, kjGlobal); + + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + + TMOV(aTile, aMatTile); + TMOV(bTile, bMatTile); + + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + + TMATMUL(cTile, aTile, bTile); + + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + + TSTORE(sijGlobal, cTile); + + set_flag(PIPE_FIX, PIPE_S, EVENT_ID7); + wait_flag(PIPE_FIX, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *query_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *key_cache_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *block_table_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *sij_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + + int64_t bn = static_cast(args[5]); + int64_t num_heads = static_cast(args[6]); + int64_t head_dim = static_cast(args[7]); + int64_t block_size = static_cast(args[8]); + int64_t max_blocks_per_req = static_cast(args[9]); + int64_t q_loop = static_cast(args[10]); + + int32_t block_idx = get_block_idx(args); + + // Decode (batch_idx, q_tile_idx) from block_idx + int64_t batch_idx = block_idx / q_loop; + int64_t q_tile_idx = block_idx % q_loop; + + // Check if this batch has data at this KV block + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t bn_this_batch = (cur_seq + block_size - 1) / block_size; + + // Output pointer for this block's sij slice + __gm__ float *sij_addr = + reinterpret_cast<__gm__ float *>(sij_t->buffer.addr) + sij_t->start_offset + block_idx * M * block_size; + + if (bn >= bn_this_batch) { + // No valid KV data for this batch at this bn — zero out sij + for (int i = 0; i < M * static_cast(block_size); i++) { + sij_addr[i] = 0.0f; + } + return; + } + + // Look up physical block index from block_table + __gm__ int32_t *bt_ptr = + reinterpret_cast<__gm__ int32_t *>(block_table_t->buffer.addr) + block_table_t->start_offset; + int64_t phys_block = static_cast(bt_ptr[batch_idx * max_blocks_per_req + bn]); + + // Query offset: (batch_idx * num_heads + q_tile_idx * Q_TILE, 0) + int64_t q_offset = (batch_idx * num_heads + q_tile_idx * M) * head_dim; + __gm__ bfloat16_t *qi_addr = + reinterpret_cast<__gm__ bfloat16_t *>(query_t->buffer.addr) + query_t->start_offset + q_offset; + + // Key offset: (phys_block * block_size, 0) + int64_t k_offset = phys_block * block_size * head_dim; + __gm__ bfloat16_t *kj_addr = + reinterpret_cast<__gm__ bfloat16_t *>(key_cache_t->buffer.addr) + key_cache_t->start_offset + k_offset; + + qk_matmul_spmd(qi_addr, kj_addr, sij_addr); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_hub.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_hub.cpp new file mode 100644 index 0000000000..a42f2790e9 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_hub.cpp @@ -0,0 +1,27 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// AIV Hub Kernel - No-op stub for accumulator tensor allocation. +// The runtime allocates output tensors specified in the Arg; the kernel itself does nothing. + +#include +#include + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) {} diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp new file mode 100644 index 0000000000..c22023fc9f --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp @@ -0,0 +1,258 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD Online Softmax Update + Normalize Kernel (AIV) with dual-vector +// subvector split. +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// The two AIV lanes in a cluster split the Q_TILE=16 rows 8/8 via +// get_sub_block_id(): AIV0 updates rows [0, 8), AIV1 updates rows [8, 16). +// The online softmax update is row-independent, so the two lanes never touch +// the same row of mi/li/oi accumulators or the output buffer. +// +// Scalar layout strategy (same as MPMD version): +// M scalar floats stored contiguously in GM can be loaded as either: +// - ND (kScalarRows, kScalarCols) RowMajor for element-wise ops +// - DN (kAlignedRows, 1) ColMajor for row-broadcast ops (TROWEXPANDMUL/DIV) +// Conversion between layouts uses GM round-trip: ND TSTORE -> DN TLOAD. +// +// Args: +// args[0] = mij Tensor* (spmd_blocks*Q_TILE,) float32 +// args[1] = lij Tensor* (spmd_blocks*Q_TILE,) float32 +// args[2] = oi_new Tensor* (spmd_blocks*Q_TILE, head_dim) float32 +// args[3] = mi_acc Tensor* (spmd_blocks*Q_TILE,) float32 [inout] +// args[4] = li_acc Tensor* (spmd_blocks*Q_TILE,) float32 [inout] +// args[5] = oi_acc Tensor* (spmd_blocks*Q_TILE, head_dim) float32 [inout] +// args[6] = out Tensor* (batch*num_heads, head_dim) float32 [inout] +// args[7] = is_first scalar +// args[8] = is_last scalar +// args[9] = num_heads scalar +// args[10] = head_dim scalar +// args[11] = q_loop scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +static constexpr int QT = 16; // Full Q tile rows (shared between both AIVs) +static constexpr int SUB_QT = 8; // Rows per AIV lane (QT / 2) +static constexpr int HD = 16; // Head dimension + +template +static __aicore__ void online_update_spmd( + __gm__ float *mij_ptr, __gm__ float *lij_ptr, __gm__ float *oi_new_ptr, __gm__ float *mi_ptr, + __gm__ float *li_ptr, __gm__ float *oi_ptr, __gm__ float *dst_ptr, uint64_t is_first, uint64_t is_last +) { + constexpr int kScalarCols = 32 / sizeof(float); + constexpr int kScalarRows = TM / kScalarCols; + constexpr int kAlignedRows = ((TM * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + + using GlobalDataMxN = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + using GlobalScalarND = + GlobalTensor, Stride<1, 1, 1, kScalarCols, 1>>; + using GlobalScalarDN = GlobalTensor, Stride<1, 1, 1, 1, 1>, Layout::DN>; + + GlobalDataMxN oiNewGlobal(oi_new_ptr); + GlobalDataMxN oiGlobal(oi_ptr); + GlobalDataMxN dstGlobal(dst_ptr); + + GlobalScalarND mijGlobalND(mij_ptr); + GlobalScalarND lijGlobalND(lij_ptr); + GlobalScalarND miGlobalND(mi_ptr); + GlobalScalarND liGlobalND(li_ptr); + + GlobalScalarDN mijGlobalDN(mij_ptr); + GlobalScalarDN lijGlobalDN(lij_ptr); + GlobalScalarDN liGlobalDN(li_ptr); + + using TileDataMxN = Tile; + using TileScalarND = + Tile; + using TileScalarDN = Tile; + + constexpr int kDataBytes = TM * TN * sizeof(float); + constexpr int kScalarNDBytes = kScalarRows * kScalarCols * sizeof(float); + constexpr int kScalarDNBytes = kAlignedRows * sizeof(float); + + TileDataMxN oiNewTile; + TileDataMxN oiTile; + TileScalarND mijND, lijND, miND, liND; + TileScalarND miNewND, alphaND, betaND, tmpND; + TileScalarDN alphaDN, betaDN, liDN; + + TASSIGN(oiNewTile, 0); + TASSIGN(oiTile, kDataBytes); + TASSIGN(mijND, 2 * kDataBytes); + TASSIGN(lijND, 2 * kDataBytes + kScalarNDBytes); + TASSIGN(miND, 2 * kDataBytes + 2 * kScalarNDBytes); + TASSIGN(liND, 2 * kDataBytes + 3 * kScalarNDBytes); + TASSIGN(miNewND, 2 * kDataBytes + 4 * kScalarNDBytes); + TASSIGN(alphaND, 2 * kDataBytes + 5 * kScalarNDBytes); + TASSIGN(betaND, 2 * kDataBytes + 6 * kScalarNDBytes); + TASSIGN(tmpND, 2 * kDataBytes + 7 * kScalarNDBytes); + TASSIGN(alphaDN, 2 * kDataBytes + 8 * kScalarNDBytes); + TASSIGN(betaDN, 2 * kDataBytes + 8 * kScalarNDBytes + kScalarDNBytes); + TASSIGN(liDN, 2 * kDataBytes + 8 * kScalarNDBytes + 2 * kScalarDNBytes); + + if (is_first) { + TLOAD(oiNewTile, oiNewGlobal); + TLOAD(mijND, mijGlobalND); + TLOAD(lijND, lijGlobalND); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(miGlobalND, mijND); + TSTORE(liGlobalND, lijND); + TSTORE(oiGlobal, oiNewTile); + + if (is_last) { + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + TLOAD(liDN, liGlobalDN); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + TROWEXPANDDIV(oiNewTile, oiNewTile, liDN); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(dstGlobal, oiNewTile); + } + } else { + TLOAD(oiNewTile, oiNewGlobal); + TLOAD(oiTile, oiGlobal); + TLOAD(mijND, mijGlobalND); + TLOAD(lijND, lijGlobalND); + TLOAD(miND, miGlobalND); + TLOAD(liND, liGlobalND); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + TMAX(miNewND, miND, mijND); + pipe_barrier(PIPE_V); + TSUB(alphaND, miND, miNewND); + pipe_barrier(PIPE_V); + TEXP(alphaND, alphaND); + pipe_barrier(PIPE_V); + TSUB(betaND, mijND, miNewND); + pipe_barrier(PIPE_V); + TEXP(betaND, betaND); + pipe_barrier(PIPE_V); + TMUL(liND, alphaND, liND); + pipe_barrier(PIPE_V); + TMUL(tmpND, betaND, lijND); + pipe_barrier(PIPE_V); + TADD(liND, liND, tmpND); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(miGlobalND, miNewND); + TSTORE(liGlobalND, liND); + TSTORE(mijGlobalND, alphaND); + TSTORE(lijGlobalND, betaND); + + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + TLOAD(alphaDN, mijGlobalDN); + TLOAD(betaDN, lijGlobalDN); + if (is_last) { + TLOAD(liDN, liGlobalDN); + } + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + + TROWEXPANDMUL(oiTile, oiTile, alphaDN); + TROWEXPANDMUL(oiNewTile, oiNewTile, betaDN); + pipe_barrier(PIPE_V); + TADD(oiTile, oiTile, oiNewTile); + + if (is_last) { + pipe_barrier(PIPE_V); + TROWEXPANDDIV(oiTile, oiTile, liDN); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(dstGlobal, oiTile); + } else { + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(oiGlobal, oiTile); + } + } + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + // Safety check: if called with null tensor args (misrouted hub invocation), return. + if (args[0] == 0 || args[1] == 0 || args[2] == 0) { + return; + } + + __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *oi_new_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *mi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *li_acc_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + __gm__ Tensor *oi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[5]); + __gm__ Tensor *out_t = reinterpret_cast<__gm__ Tensor *>(args[6]); + uint64_t is_first = static_cast(args[7]); + uint64_t is_last = static_cast(args[8]); + int64_t num_heads = static_cast(args[9]); + int64_t head_dim = static_cast(args[10]); + int64_t q_loop = static_cast(args[11]); + + int32_t block_idx = get_block_idx(args); + int32_t sub_block_id = get_sub_block_id(args); // 0 = AIV0 (rows 0..7), 1 = AIV1 (rows 8..15) + int64_t batch_idx = block_idx / q_loop; + int64_t q_tile_idx = block_idx % q_loop; + + // Scalar layout: full QT=16 rows pack to kAlignedRowsFull=16 floats per block_idx; + // each AIV lane owns kAlignedRowsSub=8 contiguous floats inside that slab. + constexpr int kAlignedRowsFull = ((QT * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + constexpr int kAlignedRowsSub = ((SUB_QT * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + + int64_t row_offset = sub_block_id * SUB_QT; + + // Accumulator offsets (each AIV lane owns its own 8-row sub-slice within the block_idx slab) + int64_t scalar_offset = block_idx * kAlignedRowsFull + sub_block_id * kAlignedRowsSub; + int64_t data_offset = (block_idx * QT + row_offset) * head_dim; + + __gm__ float *mij_ptr = + reinterpret_cast<__gm__ float *>(mij_t->buffer.addr) + mij_t->start_offset + scalar_offset; + __gm__ float *lij_ptr = + reinterpret_cast<__gm__ float *>(lij_t->buffer.addr) + lij_t->start_offset + scalar_offset; + __gm__ float *oi_new_ptr = + reinterpret_cast<__gm__ float *>(oi_new_t->buffer.addr) + oi_new_t->start_offset + data_offset; + __gm__ float *mi_ptr = + reinterpret_cast<__gm__ float *>(mi_acc_t->buffer.addr) + mi_acc_t->start_offset + scalar_offset; + __gm__ float *li_ptr = + reinterpret_cast<__gm__ float *>(li_acc_t->buffer.addr) + li_acc_t->start_offset + scalar_offset; + __gm__ float *oi_ptr = + reinterpret_cast<__gm__ float *>(oi_acc_t->buffer.addr) + oi_acc_t->start_offset + data_offset; + + // Output offset: (batch_idx * num_heads + q_tile_idx * QT + row_offset, 0) + int64_t out_offset = (batch_idx * num_heads + q_tile_idx * QT + row_offset) * head_dim; + __gm__ float *dst_ptr = reinterpret_cast<__gm__ float *>(out_t->buffer.addr) + out_t->start_offset + out_offset; + + online_update_spmd(mij_ptr, lij_ptr, oi_new_ptr, mi_ptr, li_ptr, oi_ptr, dst_ptr, is_first, is_last); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp new file mode 100644 index 0000000000..85c0187da4 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp @@ -0,0 +1,192 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD Softmax Preparation Kernel (AIV) with partial block masking and +// dual-vector subvector split. +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// The two AIV lanes in a cluster split the Q_TILE=16 rows 8/8 via +// get_sub_block_id(): AIV0 handles rows [0, 8), AIV1 handles rows [8, 16). +// +// Computes (per sub-slice of SUB_M=8 rows): +// sij_masked = pad(sij, valid_len, -inf) +// sij_scale = sij_masked * scale +// mij = row_max(sij_scale) -> (SUB_M, 1) +// pij = exp(sij_scale - mij) -> (SUB_M, N) +// lij = row_sum(pij) -> (SUB_M, 1) +// +// Args: +// args[0] = sij Tensor* (spmd_blocks*Q_TILE, block_size) float32 [input] +// args[1] = context_lens Tensor* (batch,) int32 +// args[2] = pij Tensor* (spmd_blocks*Q_TILE, block_size) bf16 [output] +// args[3] = mij Tensor* (spmd_blocks*Q_TILE,) float32 [output] +// args[4] = lij Tensor* (spmd_blocks*Q_TILE,) float32 [output] +// args[5] = scale_value scalar (as float bits in uint64) +// args[6] = bn scalar: current KV block index +// args[7] = block_size scalar +// args[8] = q_loop scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +static constexpr int M = 16; // Full Q tile rows (shared between both AIVs) +static constexpr int SUB_M = 8; // Rows per AIV lane (M / 2) +static constexpr int N = 16; // block_size + +template +static __aicore__ void softmax_prepare_spmd( + __gm__ float *sij_addr, float scale_value, uint64_t valid_len, __gm__ bfloat16_t *pij_addr, + __gm__ float *mij_addr, __gm__ float *lij_addr +) { + constexpr int kAlignedRows = ((TM * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + + using GlobalDataMxN = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + using GlobalDataMxN_bf16 = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + using GlobalScalarDN = GlobalTensor, Stride<1, 1, 1, 1, 1>, Layout::DN>; + + GlobalDataMxN sijGlobal(sij_addr); + GlobalDataMxN_bf16 pijGlobal(pij_addr); + GlobalScalarDN mijGlobal(mij_addr); + GlobalScalarDN lijGlobal(lij_addr); + + using TileSijDyn = Tile; + using TileSijPad = + Tile; + + using TileVecMxN = Tile; + using TileVecMxN_bf16 = Tile; + using TileScalarDN = Tile; + + TileVecMxN sijTile; + TileSijDyn sijDynTile(static_cast(valid_len)); + TileSijPad sijPadTile; + TileVecMxN pijTile; + TileVecMxN tmpTile; + TileScalarDN maxTile; + TileScalarDN sumTile; + TileVecMxN_bf16 pijBf16Tile; + + TASSIGN(sijTile, 0x0); + TASSIGN(sijDynTile, 0x0); + TASSIGN(sijPadTile, 0x0); + TASSIGN(pijTile, TM * TN * sizeof(float)); + TASSIGN(tmpTile, 2 * TM * TN * sizeof(float)); + TASSIGN(maxTile, 3 * TM * TN * sizeof(float)); + TASSIGN(sumTile, 3 * TM * TN * sizeof(float) + kAlignedRows * sizeof(float)); + TASSIGN(pijBf16Tile, 3 * TM * TN * sizeof(float) + 2 * kAlignedRows * sizeof(float)); + + TLOAD(sijTile, sijGlobal); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + TFILLPAD_INPLACE(sijPadTile, sijDynTile); + pipe_barrier(PIPE_V); + + TMULS(sijTile, sijTile, scale_value); + pipe_barrier(PIPE_V); + TROWMAX(maxTile, sijTile, tmpTile); + pipe_barrier(PIPE_V); + TROWEXPANDSUB(pijTile, sijTile, maxTile); + pipe_barrier(PIPE_V); + TEXP(pijTile, pijTile); + TCVT(pijBf16Tile, pijTile, RoundMode::CAST_ROUND); + TCVT(pijTile, pijBf16Tile, RoundMode::CAST_ROUND); + TROWSUM(sumTile, pijTile, tmpTile); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(mijGlobal, maxTile); + TSTORE(lijGlobal, sumTile); + TSTORE(pijGlobal, pijBf16Tile); + + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *sij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *pij_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + float scale_value = from_u64(static_cast(args[5])); + int64_t bn = static_cast(args[6]); + int64_t block_size = static_cast(args[7]); + int64_t q_loop = static_cast(args[8]); + + int32_t block_idx = get_block_idx(args); + int32_t sub_block_id = get_sub_block_id(args); // 0 = AIV0 (rows 0..7), 1 = AIV1 (rows 8..15) + int64_t batch_idx = block_idx / q_loop; + + // Compute valid_len for this block: how many columns of sij are valid + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t remaining = cur_seq - bn * block_size; + uint64_t valid_len; + if (remaining <= 0) { + valid_len = 0; + } else if (remaining >= block_size) { + valid_len = static_cast(block_size); + } else { + valid_len = static_cast(remaining); + } + + // Row offset for this AIV lane within the block_idx's Q_TILE slice + int64_t row_offset = sub_block_id * SUB_M; + + // Pointers into this block's SUB_M-row sub-slice of the flat tensors + int64_t data_row_offset = block_idx * M + row_offset; + __gm__ float *sij_addr = + reinterpret_cast<__gm__ float *>(sij_t->buffer.addr) + sij_t->start_offset + data_row_offset * block_size; + __gm__ bfloat16_t *pij_addr = + reinterpret_cast<__gm__ bfloat16_t *>(pij_t->buffer.addr) + pij_t->start_offset + data_row_offset * block_size; + + // Scalar layout: full M=16 rows pack to kAlignedRowsFull=16 floats per block_idx; + // each AIV lane owns kAlignedRowsSub=8 contiguous floats inside that slab. + constexpr int kAlignedRowsFull = ((M * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + constexpr int kAlignedRowsSub = ((SUB_M * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + int64_t scalar_offset = block_idx * kAlignedRowsFull + sub_block_id * kAlignedRowsSub; + __gm__ float *mij_addr = + reinterpret_cast<__gm__ float *>(mij_t->buffer.addr) + mij_t->start_offset + scalar_offset; + __gm__ float *lij_addr = + reinterpret_cast<__gm__ float *>(lij_t->buffer.addr) + lij_t->start_offset + scalar_offset; + + if (valid_len == 0) { + // No valid KV data — emit neutral values so online_update is a no-op: + // mij = -1e30 (very negative so beta = exp(mij - mi_new) ≈ 0) + // lij = 0 (no contribution to normalizer) + // pij = 0 (no attention weight) + for (int i = 0; i < kAlignedRowsSub; i++) { + mij_addr[i] = -1e30f; + lij_addr[i] = 0.0f; + } + for (int i = 0; i < SUB_M * static_cast(block_size); i++) { + pij_addr[i] = static_cast(0.0f); + } + return; + } + + softmax_prepare_spmd(sij_addr, scale_value, valid_len, pij_addr, mij_addr, lij_addr); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py new file mode 100644 index 0000000000..f24c94b1f8 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py @@ -0,0 +1,85 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +""" +SPMD Paged Attention Kernel and Orchestration Configuration + +Uses SPMD (block_num) parallelism across batch*q_loop positions. +Each block handles one (batch_idx, q_tile_idx) using get_block_idx(). + +Softmax and online-update run as MIX tasks (AIC idle + AIV0 + AIV1), with the +two AIVs splitting the 16 query rows 8/8 via get_sub_block_id(). + +AIC Kernels (Matrix Multiplication): + - aic_qk_matmul: Q @ K^T (SPMD across batch*q_loop) + - aic_pv_matmul: P @ V (SPMD across batch*q_loop) + - aic_hub: no-op, occupies the AIC slot of softmax/update MIX tasks + +AIV Kernels (Vector Operations): + - aiv_softmax_prepare: scale, rowmax, exp, rowsum on 8-row sub-tile + - aiv_online_update: online softmax accumulation + normalization on 8-row sub-tile + - aiv_hub: no-op, used to allocate persistent accumulators +""" + +from pathlib import Path + +from task_interface import ArgDirection as D # pyright: ignore[reportAttributeAccessIssue] + +_KERNELS_ROOT = Path(__file__).parent + +ORCHESTRATION = { + "source": str(_KERNELS_ROOT / "orchestration" / "spmd_paged_attention_orch.cpp"), + "function_name": "aicpu_orchestration_entry", +} + +KERNELS = [ + # AIC kernels (matrix multiplication using Cube unit) + { + "func_id": 0, + "name": "SPMD_QK", + "source": str(_KERNELS_ROOT / "aic" / "aic_qk_matmul.cpp"), + "core_type": "aic", + }, + { + "func_id": 1, + "name": "SPMD_PV", + "source": str(_KERNELS_ROOT / "aic" / "aic_pv_matmul.cpp"), + "core_type": "aic", + }, + { + "func_id": 2, + "name": "AIC_HUB", + "source": str(_KERNELS_ROOT / "aic" / "aic_hub.cpp"), + "core_type": "aic", + }, + # AIV kernels (vector operations) + { + "func_id": 3, + "name": "SPMD_SF", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_softmax_prepare.cpp"), + "core_type": "aiv", + }, + { + "func_id": 4, + "name": "SPMD_UP", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_online_update.cpp"), + "core_type": "aiv", + }, + { + "func_id": 5, + "name": "AIV_HUB", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_hub.cpp"), + "core_type": "aiv", + }, +] + +RUNTIME_CONFIG = { + "runtime": "tensormap_and_ringbuffer", + "aicpu_thread_num": 4, + "block_dim": 24, +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp new file mode 100644 index 0000000000..22140426de --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp @@ -0,0 +1,239 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +/** + * SPMD Paged Attention Orchestration (dual-vector subvector partitioning) + * + * Uses SPMD parallelism: block_num = batch * q_loop, where each logical + * block handles one (batch_idx, q_tile_idx) position. Kernels use + * get_block_idx() to compute their data offsets. + * + * QK and PV matmuls are AIC-only SPMD tasks. Softmax and online-update are + * submitted as MIX tasks (AIC hub + AIV0 + AIV1) so the two AIV lanes within + * a cluster each process one half of the 16 query rows, using + * get_sub_block_id() to pick their 8-row slice. This mirrors the AscendC + * reference (paged_attention_antiquantkv.h) subvector partitioning strategy. + * + * Memory Layout: + * Query: (batch, num_heads, head_dim) - bfloat16 + * Key/Value: (total_blocks, block_size, kv_head_num, head_dim) - bfloat16 + * Block Table: (batch, max_num_blocks_per_req) - int32 + * Context Lens: (batch,) - int32 + * Output: (batch, num_heads, head_dim) - float32 + * + * Scratch layout (runtime-allocated, indexed by block_idx * Q_TILE): + * sij: (spmd_blocks * Q_TILE, block_size) float32 + * pij: (spmd_blocks * Q_TILE, block_size) data_type + * oi_new: (spmd_blocks * Q_TILE, head_dim) float32 + * mij/lij: (spmd_blocks * Q_TILE,) float32 + * oi_acc/mi_acc/li_acc: persistent accumulators across bn loop + */ + +#include +#include + +#include + +#include "pto_orchestration_api.h" + +#define FUNC_QK_MATMUL 0 +#define FUNC_PV_MATMUL 1 +#define FUNC_AIC_HUB 2 +#define FUNC_SOFTMAX_PREPARE 3 +#define FUNC_ONLINE_UPDATE 4 +#define FUNC_AIV_HUB 5 + +static constexpr uint64_t Q_TILE = 16; + +extern "C" { + +__attribute__((visibility("default"))) PTO2OrchestrationConfig +aicpu_orchestration_config(const ChipStorageTaskArgs &orch_args) { + (void)orch_args; + return PTO2OrchestrationConfig{ + .expected_arg_count = 7, + }; +} + +__attribute__((visibility("default"))) void aicpu_orchestration_entry(const ChipStorageTaskArgs &orch_args) { + // query: shape=[batch, num_heads, head_dim] + uint64_t batch = orch_args.tensor(0).shapes[0]; + uint64_t num_heads = orch_args.tensor(0).shapes[1]; + uint64_t head_dim = orch_args.tensor(0).shapes[2]; + DataType data_type = orch_args.tensor(0).dtype; + + // key_cache: shape=[total_blocks, block_size, kv_head_num, head_dim] + uint64_t block_size = orch_args.tensor(1).shapes[1]; + + // block_table: shape=[batch, max_num_blocks_per_req] + uint64_t max_num_blocks_per_req = orch_args.tensor(3).shapes[1]; + + // scale from scalar arg + uint64_t scale_value = orch_args.scalar(0); + + uint64_t q_loop = (num_heads + Q_TILE - 1) / Q_TILE; + int16_t spmd_block_num = static_cast(batch * q_loop); + + LOG_INFO( + "SPMD PA: batch=%" PRIu64 " heads=%" PRIu64 " hd=%" PRIu64 " bs=%" PRIu64 " q_loop=%" PRIu64 " blocks=%d", + batch, num_heads, head_dim, block_size, q_loop, spmd_block_num + ); + + // Wrap host-provided tensors + void *query_ptr = orch_args.tensor(0).data_as(); + void *kc_ptr = orch_args.tensor(1).data_as(); + void *vc_ptr = orch_args.tensor(2).data_as(); + void *out_ptr = orch_args.tensor(5).data_as(); + + uint64_t total_kv_blocks = orch_args.tensor(1).shapes[0]; + uint64_t kv_total_rows = total_kv_blocks * block_size; + + uint32_t query_shapes[2] = {static_cast(batch * num_heads), static_cast(head_dim)}; + uint32_t kv_shapes[2] = {static_cast(kv_total_rows), static_cast(head_dim)}; + uint32_t out_shapes[2] = {static_cast(batch * num_heads), static_cast(head_dim)}; + + Tensor query = make_tensor_external(query_ptr, query_shapes, 2, data_type); + Tensor key_cache = make_tensor_external(kc_ptr, kv_shapes, 2, data_type); + Tensor value_cache = make_tensor_external(vc_ptr, kv_shapes, 2, data_type); + Tensor out = make_tensor_external(out_ptr, out_shapes, 2, DataType::FLOAT32); + + uint32_t bt_shapes[2] = {static_cast(batch), static_cast(max_num_blocks_per_req)}; + Tensor block_table = + make_tensor_external(orch_args.tensor(3).data_as(), bt_shapes, 2, DataType::INT32, false); + uint32_t cl_shapes[1] = {static_cast(batch)}; + Tensor context_lens = + make_tensor_external(orch_args.tensor(4).data_as(), cl_shapes, 1, DataType::INT32, false); + + // Find max context_len for KV block loop bound + uint64_t max_ctx = 0; + for (uint64_t b = 0; b < batch; b++) { + uint32_t idx[1] = {static_cast(b)}; + uint64_t ctx = static_cast(get_tensor_data(context_lens, 1, idx)); + if (ctx > max_ctx) max_ctx = ctx; + } + uint64_t max_bn = (max_ctx + block_size - 1) / block_size; + + // Scratch tensor create infos (sized for all SPMD blocks) + uint32_t n_rows = static_cast(spmd_block_num) * static_cast(Q_TILE); + uint32_t sij_shapes[2] = {n_rows, static_cast(block_size)}; + uint32_t pij_shapes[2] = {n_rows, static_cast(block_size)}; + uint32_t oi_new_shapes[2] = {n_rows, static_cast(head_dim)}; + uint32_t scalar_shapes[1] = {n_rows}; + + TensorCreateInfo sij_ci(sij_shapes, 2, DataType::FLOAT32); + TensorCreateInfo pij_ci(pij_shapes, 2, data_type); + TensorCreateInfo oi_new_ci(oi_new_shapes, 2, DataType::FLOAT32); + TensorCreateInfo mij_ci(scalar_shapes, 1, DataType::FLOAT32); + TensorCreateInfo lij_ci(scalar_shapes, 1, DataType::FLOAT32); + TensorCreateInfo acc_oi_ci(oi_new_shapes, 2, DataType::FLOAT32); + TensorCreateInfo acc_mi_ci(scalar_shapes, 1, DataType::FLOAT32); + TensorCreateInfo acc_li_ci(scalar_shapes, 1, DataType::FLOAT32); + + PTO2_SCOPE() { + // Allocate persistent accumulators via no-op AIV hub + Arg hub_args; + hub_args.add_output(acc_oi_ci); + hub_args.add_output(acc_mi_ci); + hub_args.add_output(acc_li_ci); + TaskOutputTensors hub_outs = pto2_rt_submit_aiv_task(FUNC_AIV_HUB, hub_args); + const Tensor &oi_acc = hub_outs.get_ref(0); + const Tensor &mi_acc = hub_outs.get_ref(1); + const Tensor &li_acc = hub_outs.get_ref(2); + + for (uint64_t bn = 0; bn < max_bn; bn++) { + uint64_t is_first = (bn == 0) ? 1 : 0; + uint64_t is_last = (bn == max_bn - 1) ? 1 : 0; + + // -- QK Matmul (AIC, SPMD) -- + Arg qk_args; + qk_args.add_input(query); + qk_args.add_input(key_cache); + qk_args.add_input(block_table); + qk_args.add_input(context_lens); + qk_args.add_output(sij_ci); + qk_args.add_scalar(static_cast(bn)); + qk_args.add_scalar(static_cast(num_heads)); + qk_args.add_scalar(static_cast(head_dim)); + qk_args.add_scalar(static_cast(block_size)); + qk_args.add_scalar(static_cast(max_num_blocks_per_req)); + qk_args.add_scalar(static_cast(q_loop)); + qk_args.launch_spec.set_block_num(spmd_block_num); + TaskOutputTensors qk_outs = pto2_rt_submit_aic_task(FUNC_QK_MATMUL, qk_args); + const Tensor &sij = qk_outs.get_ref(0); + + // -- Softmax Prepare (MIX: AIC hub + AIV0 + AIV1, SPMD) -- + // AIV0 processes rows 0..7, AIV1 processes rows 8..15 of the Q_TILE + // slice, discriminated via get_sub_block_id() inside the kernel. + Arg sf_args; + sf_args.add_input(sij); + sf_args.add_input(context_lens); + sf_args.add_output(pij_ci); + sf_args.add_output(mij_ci); + sf_args.add_output(lij_ci); + sf_args.add_scalar(scale_value); + sf_args.add_scalar(static_cast(bn)); + sf_args.add_scalar(static_cast(block_size)); + sf_args.add_scalar(static_cast(q_loop)); + sf_args.launch_spec.set_block_num(spmd_block_num); + MixedKernels sf_mk; + sf_mk.aic_kernel_id = FUNC_AIC_HUB; + sf_mk.aiv0_kernel_id = FUNC_SOFTMAX_PREPARE; + sf_mk.aiv1_kernel_id = FUNC_SOFTMAX_PREPARE; + TaskOutputTensors sf_outs = pto2_rt_submit_task(sf_mk, sf_args); + const Tensor &pij = sf_outs.get_ref(0); + const Tensor &mij = sf_outs.get_ref(1); + const Tensor &lij = sf_outs.get_ref(2); + + // -- PV Matmul (AIC, SPMD) -- + Arg pv_args; + pv_args.add_input(pij); + pv_args.add_input(value_cache); + pv_args.add_input(block_table); + pv_args.add_input(context_lens); + pv_args.add_output(oi_new_ci); + pv_args.add_scalar(static_cast(bn)); + pv_args.add_scalar(static_cast(num_heads)); + pv_args.add_scalar(static_cast(head_dim)); + pv_args.add_scalar(static_cast(block_size)); + pv_args.add_scalar(static_cast(max_num_blocks_per_req)); + pv_args.add_scalar(static_cast(q_loop)); + pv_args.launch_spec.set_block_num(spmd_block_num); + TaskOutputTensors pv_outs = pto2_rt_submit_aic_task(FUNC_PV_MATMUL, pv_args); + const Tensor &oi_new = pv_outs.get_ref(0); + + // -- Online Update (MIX: AIC hub + AIV0 + AIV1, SPMD) -- + // Row-independent online softmax update: AIV0 updates rows 0..7 of + // the Q_TILE accumulator slice, AIV1 updates rows 8..15. + Arg up_args; + up_args.add_input(mij); + up_args.add_input(lij); + up_args.add_input(oi_new); + up_args.add_inout(mi_acc); + up_args.add_inout(li_acc); + up_args.add_inout(oi_acc); + up_args.add_inout(out); + up_args.add_scalar(is_first); + up_args.add_scalar(is_last); + up_args.add_scalar(static_cast(num_heads)); + up_args.add_scalar(static_cast(head_dim)); + up_args.add_scalar(static_cast(q_loop)); + up_args.launch_spec.set_block_num(spmd_block_num); + MixedKernels up_mk; + up_mk.aic_kernel_id = FUNC_AIC_HUB; + up_mk.aiv0_kernel_id = FUNC_ONLINE_UPDATE; + up_mk.aiv1_kernel_id = FUNC_ONLINE_UPDATE; + pto2_rt_submit_task(up_mk, up_args); + } + } + + LOG_INFO("SPMD PA: %" PRIu64 " KV iters x 4 tasks, blocks=%d", max_bn, static_cast(spmd_block_num)); +} + +} // extern "C" From 50b5e3bf0e9ae3bc64ca6104c0c5ce0d2b294403 Mon Sep 17 00:00:00 2001 From: chenshengxin Date: Tue, 14 Apr 2026 09:30:37 +0800 Subject: [PATCH 4/4] modify spmd pa --- .../spmd_paged_attention/golden.py | 48 +- .../kernels/aic/aic_pv_matmul.cpp | 31 +- .../kernels/aic/aic_qk_matmul.cpp | 29 +- .../kernels/aiv/aiv_online_update.cpp | 93 +-- .../kernels/aiv/aiv_softmax_prepare.cpp | 89 ++- .../kernels/kernel_config.py | 9 +- .../spmd_paged_attention_orch.cpp | 204 +++--- .../spmd_paged_attention/golden.py | 67 ++ .../kernels/aic/aic_hub.cpp | 29 + .../kernels/aic/aic_pv_matmul.cpp | 214 +++++++ .../kernels/aic/aic_qk_matmul.cpp | 216 +++++++ .../kernels/aiv/aiv_hub.cpp | 27 + .../kernels/aiv/aiv_online_update.cpp | 257 ++++++++ .../kernels/aiv/aiv_softmax_prepare.cpp | 344 +++++++++++ .../kernels/kernel_config.py | 86 +++ .../spmd_paged_attention_orch.cpp | 248 ++++++++ .../spmd_paged_attention_blk24/golden.py | 67 ++ .../kernels/aic/aic_hub.cpp | 29 + .../kernels/aic/aic_pv_matmul.cpp | 167 +++++ .../kernels/aic/aic_qk_matmul.cpp | 171 ++++++ .../kernels/aiv/aiv_hub.cpp | 27 + .../kernels/aiv/aiv_online_update.cpp | 316 ++++++++++ .../kernels/aiv/aiv_softmax_prepare.cpp | 251 ++++++++ .../kernels/kernel_config.py | 87 +++ .../spmd_paged_attention_orch.cpp | 268 ++++++++ .../golden.py | 68 ++ .../kernels/kernel_config.py | 52 ++ .../kernels/mix/paged_attention_parallel.cpp | 579 ++++++++++++++++++ .../spmd_paged_attention_orch.cpp | 142 +++++ .../spmd_paged_attention_unroll/golden.py | 67 ++ .../kernels/aic/aic_hub.cpp | 29 + .../kernels/aic/aic_pv_matmul.cpp | 214 +++++++ .../kernels/aic/aic_qk_matmul.cpp | 216 +++++++ .../kernels/aiv/aiv_hub.cpp | 27 + .../kernels/aiv/aiv_online_update.cpp | 257 ++++++++ .../kernels/aiv/aiv_softmax_prepare.cpp | 344 +++++++++++ .../kernels/kernel_config.py | 85 +++ .../spmd_paged_attention_orch.cpp | 248 ++++++++ tools/benchmark_rounds.sh | 2 + 39 files changed, 5485 insertions(+), 219 deletions(-) create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_hub.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_hub.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/golden.py create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_hub.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_pv_matmul.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_qk_matmul.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_hub.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_online_update.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_softmax_prepare.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/kernel_config.py create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/orchestration/spmd_paged_attention_orch.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/golden.py create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/kernel_config.py create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/mix/paged_attention_parallel.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/orchestration/spmd_paged_attention_orch.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/golden.py create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_hub.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_pv_matmul.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_qk_matmul.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_hub.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_online_update.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_softmax_prepare.cpp create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/kernel_config.py create mode 100644 tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/orchestration/spmd_paged_attention_orch.cpp diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py index 3847437a69..87df830e0e 100644 --- a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py @@ -25,45 +25,33 @@ ALL_CASES = { "Case1": { - "batch": 1, + "batch": 256, "num_heads": 16, "kv_head_num": 1, - "head_dim": 16, - "block_size": 16, - "context_len": 33, - "max_model_len": 256, + "head_dim": 128, + "block_size": 128, + "context_len": 256, + "max_model_len": 32768, "dtype": "bfloat16", }, "Case2": { - "batch": 1, - "num_heads": 16, - "kv_head_num": 1, - "head_dim": 16, - "block_size": 16, - "context_len": 128, - "max_model_len": 256, - "dtype": "bfloat16", - }, - "CaseVarSeq2": { - "batch": 2, - "num_heads": 16, + "batch": 64, + "num_heads": 64, "kv_head_num": 1, - "head_dim": 16, - "block_size": 16, - "context_len": 33, - "context_lens_list": [33, 17], - "max_model_len": 256, + "head_dim": 128, + "block_size": 64, + "context_len": 8192, + "max_model_len": 32768, "dtype": "bfloat16", }, - "CaseVarSeq4": { - "batch": 4, - "num_heads": 16, + "Case3": { + "batch": 64, + "num_heads": 64, "kv_head_num": 1, - "head_dim": 16, - "block_size": 16, - "context_len": 64, - "context_lens_list": [33, 64, 48, 15], - "max_model_len": 256, + "head_dim": 256, + "block_size": 64, + "context_len": 8192, + "max_model_len": 32768, "dtype": "bfloat16", }, } diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp index f2e2652e2b..6715a1b96b 100644 --- a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp @@ -8,23 +8,25 @@ * See LICENSE in the root of the software repository for the full text of the License. * ----------------------------------------------------------------------------------------------------------- */ -// SPMD PV Matmul: pij(M, K) @ vj(K, N) -> oi_new(M, N) +// SPMD PV Matmul: pij(q_tile, K) @ vj(K, N) -> oi_new(q_tile, N) // // SPMD block_idx encodes (batch_idx, q_tile_idx). -// Each block computes one 16x16 matmul using paged V cache. +// Each block computes one q_tile x K @ K x N matmul using paged V cache. +// q_tile is passed as a runtime scalar and dispatched to the matching template. // // Args: -// args[0] = pij Tensor* (spmd_blocks*Q_TILE, block_size) data_type +// args[0] = pij Tensor* (spmd_blocks*q_tile, block_size) data_type // args[1] = value_cache Tensor* (kv_total_rows, head_dim) bf16 // args[2] = block_table Tensor* (batch, max_blocks_per_req) int32 // args[3] = context_lens Tensor* (batch,) int32 -// args[4] = oi_new Tensor* (spmd_blocks*Q_TILE, head_dim) float32 [output] +// args[4] = oi_new Tensor* (spmd_blocks*q_tile, head_dim) float32 [output] // args[5] = bn scalar: current KV block index // args[6] = num_heads scalar // args[7] = head_dim scalar // args[8] = block_size scalar // args[9] = max_num_blocks_per_req scalar // args[10] = q_loop scalar +// args[11] = q_tile scalar #include #include @@ -43,9 +45,9 @@ using namespace pto; #include "intrinsic.h" -static constexpr int M = 16; -static constexpr int K = 16; -static constexpr int N = 16; +static constexpr int K_128 = 128; +static constexpr int K_64 = 64; +static constexpr int N_128 = 128; template static __aicore__ void pv_matmul_spmd(__gm__ bfloat16_t *pij_addr, __gm__ bfloat16_t *vj_addr, __gm__ float *oi_addr) { @@ -112,6 +114,7 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { int64_t block_size = static_cast(args[8]); int64_t max_blocks_per_req = static_cast(args[9]); int64_t q_loop = static_cast(args[10]); + int64_t q_tile = static_cast(args[11]); int32_t block_idx = get_block_idx(args); int64_t batch_idx = block_idx / q_loop; @@ -124,10 +127,10 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { // Output pointer for this block's oi_new slice __gm__ float *oi_addr = - reinterpret_cast<__gm__ float *>(oi_new_t->buffer.addr) + oi_new_t->start_offset + block_idx * M * head_dim; + reinterpret_cast<__gm__ float *>(oi_new_t->buffer.addr) + oi_new_t->start_offset + block_idx * q_tile * head_dim; if (bn >= bn_this_batch) { - for (int i = 0; i < M * static_cast(head_dim); i++) { + for (int i = 0; i < static_cast(q_tile * head_dim); i++) { oi_addr[i] = 0.0f; } return; @@ -138,8 +141,8 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { reinterpret_cast<__gm__ int32_t *>(block_table_t->buffer.addr) + block_table_t->start_offset; int64_t phys_block = static_cast(bt_ptr[batch_idx * max_blocks_per_req + bn]); - // pij offset: block_idx * Q_TILE * block_size - int64_t pij_offset = block_idx * M * block_size; + // pij offset: block_idx * q_tile * block_size + int64_t pij_offset = block_idx * q_tile * block_size; __gm__ bfloat16_t *pij_addr = reinterpret_cast<__gm__ bfloat16_t *>(pij_t->buffer.addr) + pij_t->start_offset + pij_offset; @@ -148,5 +151,9 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { __gm__ bfloat16_t *vj_addr = reinterpret_cast<__gm__ bfloat16_t *>(value_cache_t->buffer.addr) + value_cache_t->start_offset + v_offset; - pv_matmul_spmd(pij_addr, vj_addr, oi_addr); + if (q_tile == 16) { + pv_matmul_spmd<16, K_128, N_128>(pij_addr, vj_addr, oi_addr); + } else { + pv_matmul_spmd<64, K_64, N_128>(pij_addr, vj_addr, oi_addr); + } } diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp index bdd0844bf8..81ef800c16 100644 --- a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp @@ -8,23 +8,25 @@ * See LICENSE in the root of the software repository for the full text of the License. * ----------------------------------------------------------------------------------------------------------- */ -// SPMD QK Matmul: qi(M, K) @ kj.T(K, N) -> sij(M, N) +// SPMD QK Matmul: qi(q_tile, K) @ kj.T(K, N) -> sij(q_tile, N) // // SPMD block_idx encodes (batch_idx, q_tile_idx). -// Each block computes one 16x16 matmul using paged KV. +// Each block computes one q_tile x K @ K x N matmul using paged KV. +// q_tile is passed as a runtime scalar and dispatched to the matching template. // // Args: // args[0] = query Tensor* (batch*num_heads, head_dim) bf16 // args[1] = key_cache Tensor* (kv_total_rows, head_dim) bf16 // args[2] = block_table Tensor* (batch, max_blocks_per_req) int32 // args[3] = context_lens Tensor* (batch,) int32 -// args[4] = sij Tensor* (spmd_blocks*Q_TILE, block_size) float32 [output] +// args[4] = sij Tensor* (spmd_blocks*q_tile, block_size) float32 [output] // args[5] = bn scalar: current KV block index // args[6] = num_heads scalar // args[7] = head_dim scalar // args[8] = block_size scalar // args[9] = max_num_blocks_per_req scalar // args[10] = q_loop scalar +// args[11] = q_tile scalar #include #include @@ -43,9 +45,9 @@ using namespace pto; #include "intrinsic.h" -static constexpr int M = 16; -static constexpr int K = 16; -static constexpr int N = 16; +static constexpr int K_128 = 128; +static constexpr int N_128 = 128; +static constexpr int N_64 = 64; template static __aicore__ void @@ -114,6 +116,7 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { int64_t block_size = static_cast(args[8]); int64_t max_blocks_per_req = static_cast(args[9]); int64_t q_loop = static_cast(args[10]); + int64_t q_tile = static_cast(args[11]); int32_t block_idx = get_block_idx(args); @@ -129,11 +132,11 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { // Output pointer for this block's sij slice __gm__ float *sij_addr = - reinterpret_cast<__gm__ float *>(sij_t->buffer.addr) + sij_t->start_offset + block_idx * M * block_size; + reinterpret_cast<__gm__ float *>(sij_t->buffer.addr) + sij_t->start_offset + block_idx * q_tile * block_size; if (bn >= bn_this_batch) { // No valid KV data for this batch at this bn — zero out sij - for (int i = 0; i < M * static_cast(block_size); i++) { + for (int i = 0; i < static_cast(q_tile * block_size); i++) { sij_addr[i] = 0.0f; } return; @@ -144,8 +147,8 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { reinterpret_cast<__gm__ int32_t *>(block_table_t->buffer.addr) + block_table_t->start_offset; int64_t phys_block = static_cast(bt_ptr[batch_idx * max_blocks_per_req + bn]); - // Query offset: (batch_idx * num_heads + q_tile_idx * Q_TILE, 0) - int64_t q_offset = (batch_idx * num_heads + q_tile_idx * M) * head_dim; + // Query offset: (batch_idx * num_heads + q_tile_idx * q_tile, 0) + int64_t q_offset = (batch_idx * num_heads + q_tile_idx * q_tile) * head_dim; __gm__ bfloat16_t *qi_addr = reinterpret_cast<__gm__ bfloat16_t *>(query_t->buffer.addr) + query_t->start_offset + q_offset; @@ -154,5 +157,9 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { __gm__ bfloat16_t *kj_addr = reinterpret_cast<__gm__ bfloat16_t *>(key_cache_t->buffer.addr) + key_cache_t->start_offset + k_offset; - qk_matmul_spmd(qi_addr, kj_addr, sij_addr); + if (q_tile == 16) { + qk_matmul_spmd<16, K_128, N_128>(qi_addr, kj_addr, sij_addr); + } else { + qk_matmul_spmd<64, K_128, N_64>(qi_addr, kj_addr, sij_addr); + } } diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp index c22023fc9f..99750ae34b 100644 --- a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp @@ -12,30 +12,36 @@ // subvector split. // // SPMD block_idx encodes (batch_idx, q_tile_idx). -// The two AIV lanes in a cluster split the Q_TILE=16 rows 8/8 via -// get_sub_block_id(): AIV0 updates rows [0, 8), AIV1 updates rows [8, 16). +// The two AIV lanes in a cluster split the q_tile rows via +// get_sub_block_id(): AIV0 updates the first half, AIV1 the second half. +// q_tile is passed as a runtime scalar and dispatched to the matching template. // The online softmax update is row-independent, so the two lanes never touch // the same row of mi/li/oi accumulators or the output buffer. // -// Scalar layout strategy (same as MPMD version): +// Dual-vector subvector split: each AIV lane processes SUB_QT = q_tile/2 rows. +// TROWEXPANDMUL/DIV instructions operate directly on SUB_QT-row tiles. +// Scalar tiles (mi, li, alpha, beta) are stored at aligned width per lane. +// +// Scalar layout strategy: // M scalar floats stored contiguously in GM can be loaded as either: // - ND (kScalarRows, kScalarCols) RowMajor for element-wise ops // - DN (kAlignedRows, 1) ColMajor for row-broadcast ops (TROWEXPANDMUL/DIV) // Conversion between layouts uses GM round-trip: ND TSTORE -> DN TLOAD. // // Args: -// args[0] = mij Tensor* (spmd_blocks*Q_TILE,) float32 -// args[1] = lij Tensor* (spmd_blocks*Q_TILE,) float32 -// args[2] = oi_new Tensor* (spmd_blocks*Q_TILE, head_dim) float32 -// args[3] = mi_acc Tensor* (spmd_blocks*Q_TILE,) float32 [inout] -// args[4] = li_acc Tensor* (spmd_blocks*Q_TILE,) float32 [inout] -// args[5] = oi_acc Tensor* (spmd_blocks*Q_TILE, head_dim) float32 [inout] +// args[0] = mij Tensor* (padded scalar buf) float32 +// args[1] = lij Tensor* (padded scalar buf) float32 +// args[2] = oi_new Tensor* (spmd_blocks*q_tile, head_dim) float32 +// args[3] = mi_acc Tensor* (padded scalar buf) float32 [inout] +// args[4] = li_acc Tensor* (padded scalar buf) float32 [inout] +// args[5] = oi_acc Tensor* (spmd_blocks*q_tile, head_dim) float32 [inout] // args[6] = out Tensor* (batch*num_heads, head_dim) float32 [inout] // args[7] = is_first scalar // args[8] = is_last scalar // args[9] = num_heads scalar // args[10] = head_dim scalar // args[11] = q_loop scalar +// args[12] = q_tile scalar #include #include @@ -54,10 +60,10 @@ using namespace pto; #include "intrinsic.h" -static constexpr int QT = 16; // Full Q tile rows (shared between both AIVs) -static constexpr int SUB_QT = 8; // Rows per AIV lane (QT / 2) -static constexpr int HD = 16; // Head dimension +static constexpr int HD = 128; // Head dimension +// TM = actual number of valid rows this AIV lane owns. +// All TROW* instructions operate on TM-row tiles directly. template static __aicore__ void online_update_spmd( __gm__ float *mij_ptr, __gm__ float *lij_ptr, __gm__ float *oi_new_ptr, __gm__ float *mi_ptr, @@ -67,6 +73,7 @@ static __aicore__ void online_update_spmd( constexpr int kScalarRows = TM / kScalarCols; constexpr int kAlignedRows = ((TM * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + // GM accessors using GlobalDataMxN = GlobalTensor, Stride<1, 1, 1, TN, 1>>; using GlobalScalarND = GlobalTensor, Stride<1, 1, 1, kScalarCols, 1>>; @@ -85,6 +92,7 @@ static __aicore__ void online_update_spmd( GlobalScalarDN lijGlobalDN(lij_ptr); GlobalScalarDN liGlobalDN(li_ptr); + // Compute tiles use TM rows directly using TileDataMxN = Tile; using TileScalarND = Tile; @@ -202,38 +210,26 @@ static __aicore__ void online_update_spmd( wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); } -extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { - // Safety check: if called with null tensor args (misrouted hub invocation), return. - if (args[0] == 0 || args[1] == 0 || args[2] == 0) { - return; - } - - __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); - __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[1]); - __gm__ Tensor *oi_new_t = reinterpret_cast<__gm__ Tensor *>(args[2]); - __gm__ Tensor *mi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[3]); - __gm__ Tensor *li_acc_t = reinterpret_cast<__gm__ Tensor *>(args[4]); - __gm__ Tensor *oi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[5]); - __gm__ Tensor *out_t = reinterpret_cast<__gm__ Tensor *>(args[6]); - uint64_t is_first = static_cast(args[7]); - uint64_t is_last = static_cast(args[8]); - int64_t num_heads = static_cast(args[9]); - int64_t head_dim = static_cast(args[10]); - int64_t q_loop = static_cast(args[11]); +template +static __aicore__ void online_update_entry( + __gm__ Tensor *mij_t, __gm__ Tensor *lij_t, __gm__ Tensor *oi_new_t, __gm__ Tensor *mi_acc_t, + __gm__ Tensor *li_acc_t, __gm__ Tensor *oi_acc_t, __gm__ Tensor *out_t, uint64_t is_first, uint64_t is_last, + int64_t num_heads, int64_t head_dim, int64_t q_loop, __gm__ int64_t *args +) { + constexpr int SUB_QT = QT / 2; int32_t block_idx = get_block_idx(args); - int32_t sub_block_id = get_sub_block_id(args); // 0 = AIV0 (rows 0..7), 1 = AIV1 (rows 8..15) + int32_t sub_block_id = get_sub_block_id(args); int64_t batch_idx = block_idx / q_loop; int64_t q_tile_idx = block_idx % q_loop; - // Scalar layout: full QT=16 rows pack to kAlignedRowsFull=16 floats per block_idx; - // each AIV lane owns kAlignedRowsSub=8 contiguous floats inside that slab. - constexpr int kAlignedRowsFull = ((QT * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + // Scalar layout uses aligned rows per sub-tile (based on SUB_QT) + constexpr int kAlignedRowsFull = 2 * (((SUB_QT * sizeof(float) + 31) / 32) * (32 / sizeof(float))); constexpr int kAlignedRowsSub = ((SUB_QT * sizeof(float) + 31) / 32) * (32 / sizeof(float)); int64_t row_offset = sub_block_id * SUB_QT; - // Accumulator offsets (each AIV lane owns its own 8-row sub-slice within the block_idx slab) + // Accumulator offsets (each AIV lane owns its own sub-slice within the block_idx slab) int64_t scalar_offset = block_idx * kAlignedRowsFull + sub_block_id * kAlignedRowsSub; int64_t data_offset = (block_idx * QT + row_offset) * head_dim; @@ -256,3 +252,30 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { online_update_spmd(mij_ptr, lij_ptr, oi_new_ptr, mi_ptr, li_ptr, oi_ptr, dst_ptr, is_first, is_last); } + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + // Safety check: if called with null tensor args (misrouted hub invocation), return. + if (args[0] == 0 || args[1] == 0 || args[2] == 0) { + return; + } + + __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *oi_new_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *mi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *li_acc_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + __gm__ Tensor *oi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[5]); + __gm__ Tensor *out_t = reinterpret_cast<__gm__ Tensor *>(args[6]); + uint64_t is_first = static_cast(args[7]); + uint64_t is_last = static_cast(args[8]); + int64_t num_heads = static_cast(args[9]); + int64_t head_dim = static_cast(args[10]); + int64_t q_loop = static_cast(args[11]); + int64_t q_tile = static_cast(args[12]); + + if (q_tile == 16) { + online_update_entry<16>(mij_t, lij_t, oi_new_t, mi_acc_t, li_acc_t, oi_acc_t, out_t, is_first, is_last, num_heads, head_dim, q_loop, args); + } else { + online_update_entry<64>(mij_t, lij_t, oi_new_t, mi_acc_t, li_acc_t, oi_acc_t, out_t, is_first, is_last, num_heads, head_dim, q_loop, args); + } +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp index 85c0187da4..284fb70a9b 100644 --- a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp @@ -12,26 +12,32 @@ // dual-vector subvector split. // // SPMD block_idx encodes (batch_idx, q_tile_idx). -// The two AIV lanes in a cluster split the Q_TILE=16 rows 8/8 via -// get_sub_block_id(): AIV0 handles rows [0, 8), AIV1 handles rows [8, 16). +// The two AIV lanes in a cluster split the q_tile rows via +// get_sub_block_id(): AIV0 handles the first half, AIV1 handles the second. +// q_tile is passed as a runtime scalar and dispatched to the matching template. // -// Computes (per sub-slice of SUB_M=8 rows): +// Dual-vector subvector split: each AIV lane processes SUB_M = q_tile/2 rows. +// TROW* instructions (TROWMAX, TROWEXPANDSUB, TROWSUM) operate directly on +// SUB_M-row tiles. Scalar outputs (mij/lij) are stored at aligned width per lane. +// +// Computes (per sub-slice of sub_m rows): // sij_masked = pad(sij, valid_len, -inf) // sij_scale = sij_masked * scale -// mij = row_max(sij_scale) -> (SUB_M, 1) -// pij = exp(sij_scale - mij) -> (SUB_M, N) -// lij = row_sum(pij) -> (SUB_M, 1) +// mij = row_max(sij_scale) -> (sub_m, 1) +// pij = exp(sij_scale - mij) -> (sub_m, N) +// lij = row_sum(pij) -> (sub_m, 1) // // Args: -// args[0] = sij Tensor* (spmd_blocks*Q_TILE, block_size) float32 [input] +// args[0] = sij Tensor* (spmd_blocks*q_tile, block_size) float32 [input] // args[1] = context_lens Tensor* (batch,) int32 -// args[2] = pij Tensor* (spmd_blocks*Q_TILE, block_size) bf16 [output] -// args[3] = mij Tensor* (spmd_blocks*Q_TILE,) float32 [output] -// args[4] = lij Tensor* (spmd_blocks*Q_TILE,) float32 [output] +// args[2] = pij Tensor* (spmd_blocks*q_tile, block_size) bf16 [output] +// args[3] = mij Tensor* (padded scalar buf) float32 [output] +// args[4] = lij Tensor* (padded scalar buf) float32 [output] // args[5] = scale_value scalar (as float bits in uint64) // args[6] = bn scalar: current KV block index // args[7] = block_size scalar // args[8] = q_loop scalar +// args[9] = q_tile scalar #include #include @@ -50,10 +56,12 @@ using namespace pto; #include "intrinsic.h" -static constexpr int M = 16; // Full Q tile rows (shared between both AIVs) -static constexpr int SUB_M = 8; // Rows per AIV lane (M / 2) -static constexpr int N = 16; // block_size +static constexpr int N_128 = 128; // block_size (Case1) +static constexpr int N_64 = 64; // block_size (Case2) +// TM = actual number of valid rows this AIV lane owns. +// All TROW* instructions operate on TM-row tiles directly. +// Scalar outputs (mij/lij) are stored at aligned width. template static __aicore__ void softmax_prepare_spmd( __gm__ float *sij_addr, float scale_value, uint64_t valid_len, __gm__ bfloat16_t *pij_addr, @@ -61,6 +69,7 @@ static __aicore__ void softmax_prepare_spmd( ) { constexpr int kAlignedRows = ((TM * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + // GM accessors using GlobalDataMxN = GlobalTensor, Stride<1, 1, 1, TN, 1>>; using GlobalDataMxN_bf16 = GlobalTensor, Stride<1, 1, 1, TN, 1>>; using GlobalScalarDN = GlobalTensor, Stride<1, 1, 1, 1, 1>, Layout::DN>; @@ -70,6 +79,7 @@ static __aicore__ void softmax_prepare_spmd( GlobalScalarDN mijGlobal(mij_addr); GlobalScalarDN lijGlobal(lij_addr); + // Compute tiles use TM rows directly using TileSijDyn = Tile; using TileSijPad = Tile; @@ -100,6 +110,7 @@ static __aicore__ void softmax_prepare_spmd( set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + // Pad invalid columns [valid_len, N) with -inf TFILLPAD_INPLACE(sijPadTile, sijDynTile); pipe_barrier(PIPE_V); @@ -124,19 +135,15 @@ static __aicore__ void softmax_prepare_spmd( wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); } -extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { - __gm__ Tensor *sij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); - __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[1]); - __gm__ Tensor *pij_t = reinterpret_cast<__gm__ Tensor *>(args[2]); - __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[3]); - __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[4]); - float scale_value = from_u64(static_cast(args[5])); - int64_t bn = static_cast(args[6]); - int64_t block_size = static_cast(args[7]); - int64_t q_loop = static_cast(args[8]); +template +static __aicore__ void softmax_prepare_entry( + __gm__ Tensor *sij_t, __gm__ Tensor *context_lens_t, __gm__ Tensor *pij_t, __gm__ Tensor *mij_t, + __gm__ Tensor *lij_t, float scale_value, int64_t bn, int64_t block_size, int64_t q_loop, __gm__ int64_t *args +) { + constexpr int SUB_M = Q_TILE / 2; int32_t block_idx = get_block_idx(args); - int32_t sub_block_id = get_sub_block_id(args); // 0 = AIV0 (rows 0..7), 1 = AIV1 (rows 8..15) + int32_t sub_block_id = get_sub_block_id(args); int64_t batch_idx = block_idx / q_loop; // Compute valid_len for this block: how many columns of sij are valid @@ -153,19 +160,18 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { valid_len = static_cast(remaining); } - // Row offset for this AIV lane within the block_idx's Q_TILE slice + // Row offset for this AIV lane within the block_idx's q_tile slice int64_t row_offset = sub_block_id * SUB_M; // Pointers into this block's SUB_M-row sub-slice of the flat tensors - int64_t data_row_offset = block_idx * M + row_offset; + int64_t data_row_offset = block_idx * Q_TILE + row_offset; __gm__ float *sij_addr = reinterpret_cast<__gm__ float *>(sij_t->buffer.addr) + sij_t->start_offset + data_row_offset * block_size; __gm__ bfloat16_t *pij_addr = reinterpret_cast<__gm__ bfloat16_t *>(pij_t->buffer.addr) + pij_t->start_offset + data_row_offset * block_size; - // Scalar layout: full M=16 rows pack to kAlignedRowsFull=16 floats per block_idx; - // each AIV lane owns kAlignedRowsSub=8 contiguous floats inside that slab. - constexpr int kAlignedRowsFull = ((M * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + // Scalar layout uses aligned rows per sub-tile (based on SUB_M) + constexpr int kAlignedRowsFull = 2 * (((SUB_M * sizeof(float) + 31) / 32) * (32 / sizeof(float))); constexpr int kAlignedRowsSub = ((SUB_M * sizeof(float) + 31) / 32) * (32 / sizeof(float)); int64_t scalar_offset = block_idx * kAlignedRowsFull + sub_block_id * kAlignedRowsSub; __gm__ float *mij_addr = @@ -174,10 +180,6 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { reinterpret_cast<__gm__ float *>(lij_t->buffer.addr) + lij_t->start_offset + scalar_offset; if (valid_len == 0) { - // No valid KV data — emit neutral values so online_update is a no-op: - // mij = -1e30 (very negative so beta = exp(mij - mi_new) ≈ 0) - // lij = 0 (no contribution to normalizer) - // pij = 0 (no attention weight) for (int i = 0; i < kAlignedRowsSub; i++) { mij_addr[i] = -1e30f; lij_addr[i] = 0.0f; @@ -188,5 +190,24 @@ extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { return; } - softmax_prepare_spmd(sij_addr, scale_value, valid_len, pij_addr, mij_addr, lij_addr); + softmax_prepare_spmd(sij_addr, scale_value, valid_len, pij_addr, mij_addr, lij_addr); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *sij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *pij_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + float scale_value = from_u64(static_cast(args[5])); + int64_t bn = static_cast(args[6]); + int64_t block_size = static_cast(args[7]); + int64_t q_loop = static_cast(args[8]); + int64_t q_tile = static_cast(args[9]); + + if (q_tile == 16) { + softmax_prepare_entry<16, 128>(sij_t, context_lens_t, pij_t, mij_t, lij_t, scale_value, bn, block_size, q_loop, args); + } else { + softmax_prepare_entry<64, 64>(sij_t, context_lens_t, pij_t, mij_t, lij_t, scale_value, bn, block_size, q_loop, args); + } } diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py index f24c94b1f8..f4eb9faf5d 100644 --- a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py @@ -11,9 +11,10 @@ Uses SPMD (block_num) parallelism across batch*q_loop positions. Each block handles one (batch_idx, q_tile_idx) using get_block_idx(). +q_tile adapts to num_heads: q_tile = min(num_heads, MAX_Q_TILE). Softmax and online-update run as MIX tasks (AIC idle + AIV0 + AIV1), with the -two AIVs splitting the 16 query rows 8/8 via get_sub_block_id(). +two AIVs splitting the q_tile query rows via get_sub_block_id(). AIC Kernels (Matrix Multiplication): - aic_qk_matmul: Q @ K^T (SPMD across batch*q_loop) @@ -21,14 +22,14 @@ - aic_hub: no-op, occupies the AIC slot of softmax/update MIX tasks AIV Kernels (Vector Operations): - - aiv_softmax_prepare: scale, rowmax, exp, rowsum on 8-row sub-tile - - aiv_online_update: online softmax accumulation + normalization on 8-row sub-tile + - aiv_softmax_prepare: scale, rowmax, exp, rowsum on sub-tile (q_tile/2 rows) + - aiv_online_update: online softmax accumulation + normalization on sub-tile - aiv_hub: no-op, used to allocate persistent accumulators """ from pathlib import Path -from task_interface import ArgDirection as D # pyright: ignore[reportAttributeAccessIssue] +from simpler.task_interface import ArgDirection as D # pyright: ignore[reportAttributeAccessIssue] _KERNELS_ROOT = Path(__file__).parent diff --git a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp index 22140426de..698b190e5d 100644 --- a/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp +++ b/examples/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp @@ -15,11 +15,13 @@ * block handles one (batch_idx, q_tile_idx) position. Kernels use * get_block_idx() to compute their data offsets. * + * q_tile adapts to num_heads: q_tile = min(num_heads, MAX_Q_TILE). + * When num_heads <= MAX_Q_TILE, q_loop = 1 and each block processes all heads. + * * QK and PV matmuls are AIC-only SPMD tasks. Softmax and online-update are * submitted as MIX tasks (AIC hub + AIV0 + AIV1) so the two AIV lanes within - * a cluster each process one half of the 16 query rows, using - * get_sub_block_id() to pick their 8-row slice. This mirrors the AscendC - * reference (paged_attention_antiquantkv.h) subvector partitioning strategy. + * a cluster each process one half of the q_tile query rows, using + * get_sub_block_id() to pick their sub-slice. * * Memory Layout: * Query: (batch, num_heads, head_dim) - bfloat16 @@ -28,11 +30,11 @@ * Context Lens: (batch,) - int32 * Output: (batch, num_heads, head_dim) - float32 * - * Scratch layout (runtime-allocated, indexed by block_idx * Q_TILE): - * sij: (spmd_blocks * Q_TILE, block_size) float32 - * pij: (spmd_blocks * Q_TILE, block_size) data_type - * oi_new: (spmd_blocks * Q_TILE, head_dim) float32 - * mij/lij: (spmd_blocks * Q_TILE,) float32 + * Scratch layout (runtime-allocated, indexed by block_idx * q_tile): + * sij: (spmd_blocks * q_tile, block_size) float32 + * pij: (spmd_blocks * q_tile, block_size) data_type + * oi_new: (spmd_blocks * q_tile, head_dim) float32 + * mij/lij: (spmd_blocks * q_tile,) float32 * oi_acc/mi_acc/li_acc: persistent accumulators across bn loop */ @@ -50,7 +52,7 @@ #define FUNC_ONLINE_UPDATE 4 #define FUNC_AIV_HUB 5 -static constexpr uint64_t Q_TILE = 16; +static constexpr uint64_t MAX_Q_TILE = 128; extern "C" { @@ -78,12 +80,15 @@ __attribute__((visibility("default"))) void aicpu_orchestration_entry(const Chip // scale from scalar arg uint64_t scale_value = orch_args.scalar(0); - uint64_t q_loop = (num_heads + Q_TILE - 1) / Q_TILE; + // Q_TILE adapts to num_heads: use num_heads directly when it fits, cap at MAX_Q_TILE + uint64_t q_tile = (num_heads <= MAX_Q_TILE) ? num_heads : MAX_Q_TILE; + uint64_t q_loop = (num_heads + q_tile - 1) / q_tile; int16_t spmd_block_num = static_cast(batch * q_loop); LOG_INFO( - "SPMD PA: batch=%" PRIu64 " heads=%" PRIu64 " hd=%" PRIu64 " bs=%" PRIu64 " q_loop=%" PRIu64 " blocks=%d", - batch, num_heads, head_dim, block_size, q_loop, spmd_block_num + "SPMD PA: batch=%" PRIu64 " heads=%" PRIu64 " hd=%" PRIu64 " bs=%" PRIu64 " q_tile=%" PRIu64 + " q_loop=%" PRIu64 " blocks=%d", + batch, num_heads, head_dim, block_size, q_tile, q_loop, spmd_block_num ); // Wrap host-provided tensors @@ -121,11 +126,18 @@ __attribute__((visibility("default"))) void aicpu_orchestration_entry(const Chip uint64_t max_bn = (max_ctx + block_size - 1) / block_size; // Scratch tensor create infos (sized for all SPMD blocks) - uint32_t n_rows = static_cast(spmd_block_num) * static_cast(Q_TILE); + uint32_t n_rows = static_cast(spmd_block_num) * static_cast(q_tile); uint32_t sij_shapes[2] = {n_rows, static_cast(block_size)}; uint32_t pij_shapes[2] = {n_rows, static_cast(block_size)}; uint32_t oi_new_shapes[2] = {n_rows, static_cast(head_dim)}; - uint32_t scalar_shapes[1] = {n_rows}; + + // Scalar buffers (mij, lij, mi_acc, li_acc) use aligned rows per AIV lane + // based on the actual sub-tile height (q_tile / 2). + uint64_t sub_qt = q_tile / 2; + uint64_t aligned_rows_sub = ((sub_qt * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + uint64_t aligned_rows_full = 2 * aligned_rows_sub; + uint32_t scalar_n = static_cast(spmd_block_num) * static_cast(aligned_rows_full); + uint32_t scalar_shapes[1] = {scalar_n}; TensorCreateInfo sij_ci(sij_shapes, 2, DataType::FLOAT32); TensorCreateInfo pij_ci(pij_shapes, 2, data_type); @@ -136,6 +148,10 @@ __attribute__((visibility("default"))) void aicpu_orchestration_entry(const Chip TensorCreateInfo acc_mi_ci(scalar_shapes, 1, DataType::FLOAT32); TensorCreateInfo acc_li_ci(scalar_shapes, 1, DataType::FLOAT32); + // Outer scope holds persistent accumulators that live across all bn iterations. + // Each bn iteration runs in a nested inner scope so that scratch tensors + // (sij, pij, oi_new, mij, lij) are freed when the inner scope exits, + // preventing GM heap ring overflow on large configs. PTO2_SCOPE() { // Allocate persistent accumulators via no-op AIV hub Arg hub_args; @@ -151,85 +167,87 @@ __attribute__((visibility("default"))) void aicpu_orchestration_entry(const Chip uint64_t is_first = (bn == 0) ? 1 : 0; uint64_t is_last = (bn == max_bn - 1) ? 1 : 0; - // -- QK Matmul (AIC, SPMD) -- - Arg qk_args; - qk_args.add_input(query); - qk_args.add_input(key_cache); - qk_args.add_input(block_table); - qk_args.add_input(context_lens); - qk_args.add_output(sij_ci); - qk_args.add_scalar(static_cast(bn)); - qk_args.add_scalar(static_cast(num_heads)); - qk_args.add_scalar(static_cast(head_dim)); - qk_args.add_scalar(static_cast(block_size)); - qk_args.add_scalar(static_cast(max_num_blocks_per_req)); - qk_args.add_scalar(static_cast(q_loop)); - qk_args.launch_spec.set_block_num(spmd_block_num); - TaskOutputTensors qk_outs = pto2_rt_submit_aic_task(FUNC_QK_MATMUL, qk_args); - const Tensor &sij = qk_outs.get_ref(0); - - // -- Softmax Prepare (MIX: AIC hub + AIV0 + AIV1, SPMD) -- - // AIV0 processes rows 0..7, AIV1 processes rows 8..15 of the Q_TILE - // slice, discriminated via get_sub_block_id() inside the kernel. - Arg sf_args; - sf_args.add_input(sij); - sf_args.add_input(context_lens); - sf_args.add_output(pij_ci); - sf_args.add_output(mij_ci); - sf_args.add_output(lij_ci); - sf_args.add_scalar(scale_value); - sf_args.add_scalar(static_cast(bn)); - sf_args.add_scalar(static_cast(block_size)); - sf_args.add_scalar(static_cast(q_loop)); - sf_args.launch_spec.set_block_num(spmd_block_num); - MixedKernels sf_mk; - sf_mk.aic_kernel_id = FUNC_AIC_HUB; - sf_mk.aiv0_kernel_id = FUNC_SOFTMAX_PREPARE; - sf_mk.aiv1_kernel_id = FUNC_SOFTMAX_PREPARE; - TaskOutputTensors sf_outs = pto2_rt_submit_task(sf_mk, sf_args); - const Tensor &pij = sf_outs.get_ref(0); - const Tensor &mij = sf_outs.get_ref(1); - const Tensor &lij = sf_outs.get_ref(2); - - // -- PV Matmul (AIC, SPMD) -- - Arg pv_args; - pv_args.add_input(pij); - pv_args.add_input(value_cache); - pv_args.add_input(block_table); - pv_args.add_input(context_lens); - pv_args.add_output(oi_new_ci); - pv_args.add_scalar(static_cast(bn)); - pv_args.add_scalar(static_cast(num_heads)); - pv_args.add_scalar(static_cast(head_dim)); - pv_args.add_scalar(static_cast(block_size)); - pv_args.add_scalar(static_cast(max_num_blocks_per_req)); - pv_args.add_scalar(static_cast(q_loop)); - pv_args.launch_spec.set_block_num(spmd_block_num); - TaskOutputTensors pv_outs = pto2_rt_submit_aic_task(FUNC_PV_MATMUL, pv_args); - const Tensor &oi_new = pv_outs.get_ref(0); - - // -- Online Update (MIX: AIC hub + AIV0 + AIV1, SPMD) -- - // Row-independent online softmax update: AIV0 updates rows 0..7 of - // the Q_TILE accumulator slice, AIV1 updates rows 8..15. - Arg up_args; - up_args.add_input(mij); - up_args.add_input(lij); - up_args.add_input(oi_new); - up_args.add_inout(mi_acc); - up_args.add_inout(li_acc); - up_args.add_inout(oi_acc); - up_args.add_inout(out); - up_args.add_scalar(is_first); - up_args.add_scalar(is_last); - up_args.add_scalar(static_cast(num_heads)); - up_args.add_scalar(static_cast(head_dim)); - up_args.add_scalar(static_cast(q_loop)); - up_args.launch_spec.set_block_num(spmd_block_num); - MixedKernels up_mk; - up_mk.aic_kernel_id = FUNC_AIC_HUB; - up_mk.aiv0_kernel_id = FUNC_ONLINE_UPDATE; - up_mk.aiv1_kernel_id = FUNC_ONLINE_UPDATE; - pto2_rt_submit_task(up_mk, up_args); + PTO2_SCOPE() { + // -- QK Matmul (AIC, SPMD) -- + Arg qk_args; + qk_args.add_input(query); + qk_args.add_input(key_cache); + qk_args.add_input(block_table); + qk_args.add_input(context_lens); + qk_args.add_output(sij_ci); + qk_args.add_scalar(static_cast(bn)); + qk_args.add_scalar(static_cast(num_heads)); + qk_args.add_scalar(static_cast(head_dim)); + qk_args.add_scalar(static_cast(block_size)); + qk_args.add_scalar(static_cast(max_num_blocks_per_req)); + qk_args.add_scalar(static_cast(q_loop)); + qk_args.add_scalar(static_cast(q_tile)); + qk_args.launch_spec.set_block_num(spmd_block_num); + TaskOutputTensors qk_outs = pto2_rt_submit_aic_task(FUNC_QK_MATMUL, qk_args); + const Tensor &sij = qk_outs.get_ref(0); + + // -- Softmax Prepare (MIX: AIC hub + AIV0 + AIV1, SPMD) -- + Arg sf_args; + sf_args.add_input(sij); + sf_args.add_input(context_lens); + sf_args.add_output(pij_ci); + sf_args.add_output(mij_ci); + sf_args.add_output(lij_ci); + sf_args.add_scalar(scale_value); + sf_args.add_scalar(static_cast(bn)); + sf_args.add_scalar(static_cast(block_size)); + sf_args.add_scalar(static_cast(q_loop)); + sf_args.add_scalar(static_cast(q_tile)); + sf_args.launch_spec.set_block_num(spmd_block_num); + MixedKernels sf_mk; + sf_mk.aic_kernel_id = FUNC_AIC_HUB; + sf_mk.aiv0_kernel_id = FUNC_SOFTMAX_PREPARE; + sf_mk.aiv1_kernel_id = FUNC_SOFTMAX_PREPARE; + TaskOutputTensors sf_outs = pto2_rt_submit_task(sf_mk, sf_args); + const Tensor &pij = sf_outs.get_ref(0); + const Tensor &mij = sf_outs.get_ref(1); + const Tensor &lij = sf_outs.get_ref(2); + + // -- PV Matmul (AIC, SPMD) -- + Arg pv_args; + pv_args.add_input(pij); + pv_args.add_input(value_cache); + pv_args.add_input(block_table); + pv_args.add_input(context_lens); + pv_args.add_output(oi_new_ci); + pv_args.add_scalar(static_cast(bn)); + pv_args.add_scalar(static_cast(num_heads)); + pv_args.add_scalar(static_cast(head_dim)); + pv_args.add_scalar(static_cast(block_size)); + pv_args.add_scalar(static_cast(max_num_blocks_per_req)); + pv_args.add_scalar(static_cast(q_loop)); + pv_args.add_scalar(static_cast(q_tile)); + pv_args.launch_spec.set_block_num(spmd_block_num); + TaskOutputTensors pv_outs = pto2_rt_submit_aic_task(FUNC_PV_MATMUL, pv_args); + const Tensor &oi_new = pv_outs.get_ref(0); + + // -- Online Update (MIX: AIC hub + AIV0 + AIV1, SPMD) -- + Arg up_args; + up_args.add_input(mij); + up_args.add_input(lij); + up_args.add_input(oi_new); + up_args.add_inout(mi_acc); + up_args.add_inout(li_acc); + up_args.add_inout(oi_acc); + up_args.add_inout(out); + up_args.add_scalar(is_first); + up_args.add_scalar(is_last); + up_args.add_scalar(static_cast(num_heads)); + up_args.add_scalar(static_cast(head_dim)); + up_args.add_scalar(static_cast(q_loop)); + up_args.add_scalar(static_cast(q_tile)); + up_args.launch_spec.set_block_num(spmd_block_num); + MixedKernels up_mk; + up_mk.aic_kernel_id = FUNC_AIC_HUB; + up_mk.aiv0_kernel_id = FUNC_ONLINE_UPDATE; + up_mk.aiv1_kernel_id = FUNC_ONLINE_UPDATE; + pto2_rt_submit_task(up_mk, up_args); + } } } diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py new file mode 100644 index 0000000000..8b351bc17f --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/golden.py @@ -0,0 +1,67 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""SPMD Paged Attention Golden - tensormap_and_ringbuffer example (small scale, bfloat16). + +Uses SPMD parallelism: each block handles one (batch, q_tile) position. +Kernels use get_block_idx() to determine their work slice. +""" + +from paged_attention_golden import ( + compute_golden, # noqa: F401 + run_golden_test, +) +from paged_attention_golden import generate_inputs as _generate_inputs + +__outputs__ = ["out"] + +RTOL = 1e-3 +ATOL = 1e-3 + +ALL_CASES = { + "Case1": { + "batch": 256, + "num_heads": 16, + "kv_head_num": 1, + "head_dim": 128, + "block_size": 128, + "context_len": 8192, + "max_model_len": 32768, + "dtype": "bfloat16", + }, + "Case2": { + "batch": 64, + "num_heads": 64, + "kv_head_num": 1, + "head_dim": 128, + "block_size": 64, + "context_len": 8192, + "max_model_len": 32768, + "dtype": "bfloat16", + }, + "Case3": { + "batch": 64, + "num_heads": 64, + "kv_head_num": 1, + "head_dim": 256, + "block_size": 64, + "context_len": 8192, + "max_model_len": 32768, + "dtype": "bfloat16", + }, +} + +DEFAULT_CASE = "Case1" + + +def generate_inputs(params: dict) -> list: + return _generate_inputs(params) + + +if __name__ == "__main__": + run_golden_test(ALL_CASES, DEFAULT_CASE, generate_inputs, label="SPMD Paged Attention") diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_hub.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_hub.cpp new file mode 100644 index 0000000000..eb602e0bd6 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_hub.cpp @@ -0,0 +1,29 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// AIC Hub Kernel - No-op stub used as the AIC slot of MIX (AIC+AIV0+AIV1) tasks +// when the real work happens only on the two AIVs (softmax, online update). +// Pairing an idle AIC with two active AIVs forces the scheduler to allocate a +// full cluster, which is what enables the two AIV lanes to run in parallel. + +#include +#include + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) {} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp new file mode 100644 index 0000000000..bf219cdddc --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_pv_matmul.cpp @@ -0,0 +1,214 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD SplitK PV Matmul: Accumulated P @ V across n_blocks +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// Each SPMD block processes n_blocks using SplitK accumulation: +// Block 0: TMATMUL(C, A, B) — initialize accumulator +// Block i: TMATMUL_ACC(C, C, A, B) — accumulate into same C +// +// Per-block pij: contiguous packed (M, K) tiles in pij_buf +// Per-block vj: value_cache base + block_table lookup +// Single output: oi_new (M, N) fp32 = sum of P_i @ V_i across all blocks +// +// Case1: (16, 128) @ (128, 128) -> (16, 128) +// Case2: (64, 64) @ ( 64, 128) -> (64, 128) +// +// Args: +// args[0] = pij Tensor* (spmd_blocks*Q_TILE, n_blocks*block_size) bf16 +// args[1] = value_cache Tensor* (kv_total_rows, head_dim) bf16 +// args[2] = block_table Tensor* (batch, max_blocks_per_req) int32 +// args[3] = context_lens Tensor* (batch,) int32 +// args[4] = oi_new Tensor* (spmd_blocks*Q_TILE, head_dim) float32 [output] +// args[5] = bn_start scalar: starting KV block index +// args[6] = n_blocks scalar: number of KV blocks to process +// args[7] = num_heads scalar +// args[8] = head_dim scalar +// args[9] = block_size scalar +// args[10] = max_num_blocks_per_req scalar +// args[11] = q_loop scalar + +#include +// NOLINTBEGIN(clang-diagnostic-error,bugprone-reserved-identifier,bugprone-easily-swappable-parameters,modernize-avoid-c-arrays,modernize-use-auto) +#include + +#include "tensor.h" + +// NOLINTNEXTLINE(build/namespaces) +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] // NOLINT(whitespace/braces) +#endif + +#include "intrinsic.h" + +template +static __aicore__ void pv_matmul_n_spmd( + __gm__ bfloat16_t *pij_base, __gm__ bfloat16_t *val_base, __gm__ float *oi_base, uint64_t n_blocks, + __gm__ int32_t *bt, uint64_t bt_offset +) { + using GlobalA = GlobalTensor, Stride>; + using GlobalB = GlobalTensor, Stride>; + using GlobalOut = GlobalTensor, Stride>; + + using TileMatA = Tile; + using TileMatB = Tile; + + using LeftTile = TileLeft; + using RightTile = TileRight; + using AccTile = TileAcc; + + // L1 memory layout: double-buffered A and B tiles + constexpr int kATileBytes = M * K * static_cast(sizeof(bfloat16_t)); + constexpr int kBTileBytes = K * N * static_cast(sizeof(bfloat16_t)); + + TileMatA aMatTile[2]; + TileMatB bMatTile[2]; + TASSIGN(aMatTile[0], 0x0); + TASSIGN(aMatTile[1], kATileBytes); + TASSIGN(bMatTile[0], 2 * kATileBytes); + TASSIGN(bMatTile[1], 2 * kATileBytes + kBTileBytes); + + // L0 memory layout: double-buffered L0A and L0B, single accumulator L0C + LeftTile aTile[2]; + RightTile bTile[2]; + AccTile cTile; + TASSIGN(aTile[0], 0x0); + TASSIGN(aTile[1], kATileBytes); + TASSIGN(bTile[0], 0x0); + TASSIGN(bTile[1], kBTileBytes); + TASSIGN(cTile, 0x0); + + GlobalOut oiGlobal(oi_base); + + // Seed reverse-dependency flags + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + + for (uint64_t i = 0; i < n_blocks; i++) { + int cur = static_cast(i % 2); + GlobalA pijGlobal(pij_base + i * M * K); + GlobalB vjGlobal(val_base + bt[bt_offset + i] * K * N); + + // Stage 1: TLOAD (MTE2: GM -> L1[cur]) + wait_flag(PIPE_MTE1, PIPE_MTE2, (event_t)cur); + TLOAD(aMatTile[cur], pijGlobal); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + TLOAD(bMatTile[cur], vjGlobal); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + + // Stage 2: TMOV (MTE1: L1[cur] -> L0[cur]) + wait_flag(PIPE_M, PIPE_MTE1, (event_t)cur); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + TMOV(aTile[cur], aMatTile[cur]); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + TMOV(bTile[cur], bMatTile[cur]); + set_flag(PIPE_MTE1, PIPE_MTE2, (event_t)cur); + + // Stage 3: TMATMUL (M-pipe: L0A[cur] x L0B[cur] -> L0C) + set_flag(PIPE_MTE1, PIPE_M, (event_t)cur); + wait_flag(PIPE_MTE1, PIPE_M, (event_t)cur); + if (i == 0) { + TMATMUL(cTile, aTile[cur], bTile[cur]); + } else { + TMATMUL_ACC(cTile, cTile, aTile[cur], bTile[cur]); + } + set_flag(PIPE_M, PIPE_MTE1, (event_t)cur); + } + + // Drain outstanding reverse-dependency flags + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + TSTORE(oiGlobal, cTile); + + set_flag(PIPE_FIX, PIPE_S, EVENT_ID7); + wait_flag(PIPE_FIX, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *pij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *value_cache_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *block_table_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *oi_new_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + + int64_t bn_start = static_cast(args[5]); + int64_t n_blocks_total = static_cast(args[6]); + int64_t num_heads = static_cast(args[7]); + int64_t head_dim = static_cast(args[8]); + int64_t block_size = static_cast(args[9]); + int64_t max_blocks_per_req = static_cast(args[10]); + int64_t q_loop = static_cast(args[11]); + + int32_t block_idx = get_block_idx(args); + int64_t batch_idx = block_idx / q_loop; + + int64_t q_tile = (block_size == 128) ? 16 : 64; + + // Check how many KV blocks this batch actually has + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t bn_this_batch = (cur_seq + block_size - 1) / block_size; + + // Clamp n_blocks to valid range for this batch + int64_t valid_blocks = bn_this_batch - bn_start; + if (valid_blocks < 0) valid_blocks = 0; + int64_t n_blocks = (valid_blocks < n_blocks_total) ? valid_blocks : n_blocks_total; + + // Output pointer for this SPMD block's oi_new slice + __gm__ float *oi_base = + reinterpret_cast<__gm__ float *>(oi_new_t->buffer.addr) + oi_new_t->start_offset + block_idx * q_tile * head_dim; + + if (n_blocks <= 0) { + for (int64_t i = 0; i < q_tile * head_dim; i++) { + oi_base[i] = 0.0f; + } + return; + } + + // pij packed tiles: block_idx's region in pij tensor + __gm__ bfloat16_t *pij_base = + reinterpret_cast<__gm__ bfloat16_t *>(pij_t->buffer.addr) + pij_t->start_offset + + block_idx * q_tile * n_blocks_total * block_size; + + // Value cache base + __gm__ bfloat16_t *val_base = + reinterpret_cast<__gm__ bfloat16_t *>(value_cache_t->buffer.addr) + value_cache_t->start_offset; + + // Block table pointer + __gm__ int32_t *bt = + reinterpret_cast<__gm__ int32_t *>(block_table_t->buffer.addr) + block_table_t->start_offset; + uint64_t bt_offset = static_cast(batch_idx * max_blocks_per_req + bn_start); + + if (q_tile == 16) { + pv_matmul_n_spmd<16, 128, 128>( + pij_base, val_base, oi_base, static_cast(n_blocks), bt, bt_offset + ); + } else { + pv_matmul_n_spmd<64, 64, 128>( + pij_base, val_base, oi_base, static_cast(n_blocks), bt, bt_offset + ); + } +} +// NOLINTEND(clang-diagnostic-error,bugprone-reserved-identifier,bugprone-easily-swappable-parameters,modernize-avoid-c-arrays,modernize-use-auto) diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp new file mode 100644 index 0000000000..ee5edf61df --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aic/aic_qk_matmul.cpp @@ -0,0 +1,216 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD Multi-block QK Matmul: qi(M, K) @ kj.T(K, N) -> sij(M, N) for n_blocks +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// Each SPMD block processes n_blocks consecutive KV blocks starting at bn_start, +// using double-buffered L1 B tiles and hoisted qi TLOAD. +// +// Output: packed tiles [tile_0, tile_1, ..., tile_{n-1}] each (M, N) in row-major. +// +// Template: M=q_tile, K=head_dim, N=block_size +// Case1: (16, 128) @ (128, 128).T -> (16, 128) +// Case2: (64, 128) @ (128, 64).T -> (64, 64) +// +// Args: +// args[0] = query Tensor* (batch*num_heads, head_dim) bf16 +// args[1] = key_cache Tensor* (kv_total_rows, head_dim) bf16 +// args[2] = block_table Tensor* (batch, max_blocks_per_req) int32 +// args[3] = context_lens Tensor* (batch,) int32 +// args[4] = sij Tensor* (spmd_blocks*Q_TILE, n_blocks*block_size) float32 [output] +// args[5] = bn_start scalar: starting KV block index +// args[6] = n_blocks scalar: number of KV blocks to process +// args[7] = num_heads scalar +// args[8] = head_dim scalar +// args[9] = block_size scalar +// args[10] = max_num_blocks_per_req scalar +// args[11] = q_loop scalar + +#include +// NOLINTBEGIN(clang-diagnostic-error,bugprone-reserved-identifier,bugprone-easily-swappable-parameters,modernize-use-auto) +#include + +#include "tensor.h" + +// NOLINTNEXTLINE(build/namespaces) +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] // NOLINT(whitespace/braces) +#endif + +#include "intrinsic.h" + +template +static __aicore__ void qk_matmul_n_spmd( + __gm__ bfloat16_t *qi_base, __gm__ bfloat16_t *key_base, __gm__ float *sij_base, uint64_t n_blocks, + __gm__ int32_t *bt, uint64_t bt_offset +) { + using GlobalA = GlobalTensor, Stride>; + using GlobalB = GlobalTensor, Stride, Layout::DN>; + using GlobalOut = GlobalTensor, Stride>; + + using TileMatA = Tile; + using TileMatB = Tile; + + using LeftTile = TileLeft; + using RightTile = TileRight; + using AccTile = TileAcc; + + // Double-buffered L1 B tiles for kj prefetching + constexpr int kBBytes = K * N * static_cast(sizeof(bfloat16_t)); + TileMatA aMatTile; + TileMatB bMatTile_A; + TileMatB bMatTile_B; + TASSIGN(aMatTile, 0x0); + TASSIGN(bMatTile_A, 0x20000); + TASSIGN(bMatTile_B, 0x20000 + kBBytes); + + LeftTile aTile; + RightTile bTile; + AccTile cTile; + TASSIGN(aTile, 0x0); + TASSIGN(bTile, 0x0); + TASSIGN(cTile, 0x0); + + // Hoist qi TLOAD before the loop (qi is constant across all blocks) + GlobalA qiGlobal(qi_base); + TLOAD(aMatTile, qiGlobal); + + // Pre-load first kj into buffer A + GlobalB kjGlobal_0(key_base + bt[bt_offset + 0] * N * K); + TLOAD(bMatTile_A, kjGlobal_0); + + for (uint64_t i = 0; i < n_blocks; i++) { + GlobalOut sijGlobal(sij_base + i * M * N); + + // Wait for current kj TLOAD to complete + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + + // TMOV qi L1->L0A and kj L1->L0B from current buffer + TMOV(aTile, aMatTile); + if (i % 2 == 0) { + TMOV(bTile, bMatTile_A); + } else { + TMOV(bTile, bMatTile_B); + } + + // Prefetch next kj into alternate L1 buffer + if (i + 1 < n_blocks) { + GlobalB kjGlobal_next(key_base + bt[bt_offset + i + 1] * N * K); + if (i % 2 == 0) { + TLOAD(bMatTile_B, kjGlobal_next); + } else { + TLOAD(bMatTile_A, kjGlobal_next); + } + } + + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + + TMATMUL(cTile, aTile, bTile); + + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + + TSTORE(sijGlobal, cTile); + + if (i + 1 < n_blocks) { + pipe_barrier(PIPE_ALL); + } + } + set_flag(PIPE_FIX, PIPE_S, EVENT_ID7); + wait_flag(PIPE_FIX, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *query_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *key_cache_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *block_table_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *sij_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + + int64_t bn_start = static_cast(args[5]); + int64_t n_blocks_total = static_cast(args[6]); + int64_t num_heads = static_cast(args[7]); + int64_t head_dim = static_cast(args[8]); + int64_t block_size = static_cast(args[9]); + int64_t max_blocks_per_req = static_cast(args[10]); + int64_t q_loop = static_cast(args[11]); + + int32_t block_idx = get_block_idx(args); + + // Decode (batch_idx, q_tile_idx) from block_idx + int64_t batch_idx = block_idx / q_loop; + int64_t q_tile_idx = block_idx % q_loop; + + int64_t q_tile = (block_size == 128) ? 16 : 64; + + // Check how many KV blocks this batch actually has + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t bn_this_batch = (cur_seq + block_size - 1) / block_size; + + // Clamp n_blocks to valid range for this batch + int64_t valid_blocks = bn_this_batch - bn_start; + if (valid_blocks < 0) valid_blocks = 0; + int64_t n_blocks = (valid_blocks < n_blocks_total) ? valid_blocks : n_blocks_total; + + // sij packed tile output: block_idx's region starts at block_idx * q_tile * n_blocks_total * block_size + __gm__ float *sij_base = reinterpret_cast<__gm__ float *>(sij_t->buffer.addr) + sij_t->start_offset + + block_idx * q_tile * n_blocks_total * block_size; + + if (n_blocks <= 0) { + for (int64_t i = 0; i < q_tile * n_blocks_total * block_size; i++) { + sij_base[i] = 0.0f; + } + return; + } + + // Zero out trailing invalid tiles + if (n_blocks < n_blocks_total) { + __gm__ float *trail = sij_base + n_blocks * q_tile * block_size; + for (int64_t i = 0; i < (n_blocks_total - n_blocks) * q_tile * block_size; i++) { + trail[i] = 0.0f; + } + } + + // Query offset: (batch_idx * num_heads + q_tile_idx * q_tile, 0) + int64_t q_offset = (batch_idx * num_heads + q_tile_idx * q_tile) * head_dim; + __gm__ bfloat16_t *qi_base = + reinterpret_cast<__gm__ bfloat16_t *>(query_t->buffer.addr) + query_t->start_offset + q_offset; + + // Key cache base + __gm__ bfloat16_t *key_base = + reinterpret_cast<__gm__ bfloat16_t *>(key_cache_t->buffer.addr) + key_cache_t->start_offset; + + // Block table pointer + __gm__ int32_t *bt = + reinterpret_cast<__gm__ int32_t *>(block_table_t->buffer.addr) + block_table_t->start_offset; + uint64_t bt_offset = static_cast(batch_idx * max_blocks_per_req + bn_start); + + if (q_tile == 16) { + qk_matmul_n_spmd<16, 128, 128>( + qi_base, key_base, sij_base, static_cast(n_blocks), bt, bt_offset + ); + } else { + qk_matmul_n_spmd<64, 128, 64>( + qi_base, key_base, sij_base, static_cast(n_blocks), bt, bt_offset + ); + } +} +// NOLINTEND(clang-diagnostic-error,bugprone-reserved-identifier,bugprone-easily-swappable-parameters,modernize-use-auto) diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_hub.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_hub.cpp new file mode 100644 index 0000000000..a42f2790e9 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_hub.cpp @@ -0,0 +1,27 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// AIV Hub Kernel - No-op stub for accumulator tensor allocation. +// The runtime allocates output tensors specified in the Arg; the kernel itself does nothing. + +#include +#include + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) {} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp new file mode 100644 index 0000000000..72069b4c89 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_online_update.cpp @@ -0,0 +1,257 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD Online Softmax Update + Normalize Kernel (AIV) with dual-vector +// subvector split. +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// The two AIV lanes in a cluster split the Q_TILE=16 rows 8/8 via +// get_sub_block_id(): AIV0 updates rows [0, 8), AIV1 updates rows [8, 16). +// The online softmax update is row-independent, so the two lanes never touch +// the same row of mi/li/oi accumulators or the output buffer. +// +// Scalar layout strategy (same as MPMD version): +// M scalar floats stored contiguously in GM can be loaded as either: +// - ND (kScalarRows, kScalarCols) RowMajor for element-wise ops +// - DN (kAlignedRows, 1) ColMajor for row-broadcast ops (TROWEXPANDMUL/DIV) +// Conversion between layouts uses GM round-trip: ND TSTORE -> DN TLOAD. +// +// Args: +// args[0] = mij Tensor* (spmd_blocks*Q_TILE,) float32 +// args[1] = lij Tensor* (spmd_blocks*Q_TILE,) float32 +// args[2] = oi_new Tensor* (spmd_blocks*Q_TILE, head_dim) float32 +// args[3] = mi_acc Tensor* (spmd_blocks*Q_TILE,) float32 [inout] +// args[4] = li_acc Tensor* (spmd_blocks*Q_TILE,) float32 [inout] +// args[5] = oi_acc Tensor* (spmd_blocks*Q_TILE, head_dim) float32 [inout] +// args[6] = out Tensor* (batch*num_heads, head_dim) float32 [inout] +// args[7] = is_first scalar +// args[8] = is_last scalar +// args[9] = num_heads scalar +// args[10] = head_dim scalar +// args[11] = q_loop scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +static constexpr int QT = 16; // Full Q tile rows (shared between both AIVs) +static constexpr int SUB_QT = 8; // Rows per AIV lane (QT / 2) + +template +static __aicore__ void online_update_spmd( + __gm__ float *mij_ptr, __gm__ float *lij_ptr, __gm__ float *oi_new_ptr, __gm__ float *mi_ptr, + __gm__ float *li_ptr, __gm__ float *oi_ptr, __gm__ float *dst_ptr, uint64_t is_first, uint64_t is_last +) { + constexpr int kScalarCols = 32 / sizeof(float); + constexpr int kScalarRows = TM / kScalarCols; + constexpr int kAlignedRows = ((TM * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + + using GlobalDataMxN = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + using GlobalScalarND = + GlobalTensor, Stride<1, 1, 1, kScalarCols, 1>>; + using GlobalScalarDN = GlobalTensor, Stride<1, 1, 1, 1, 1>, Layout::DN>; + + GlobalDataMxN oiNewGlobal(oi_new_ptr); + GlobalDataMxN oiGlobal(oi_ptr); + GlobalDataMxN dstGlobal(dst_ptr); + + GlobalScalarND mijGlobalND(mij_ptr); + GlobalScalarND lijGlobalND(lij_ptr); + GlobalScalarND miGlobalND(mi_ptr); + GlobalScalarND liGlobalND(li_ptr); + + GlobalScalarDN mijGlobalDN(mij_ptr); + GlobalScalarDN lijGlobalDN(lij_ptr); + GlobalScalarDN liGlobalDN(li_ptr); + + using TileDataMxN = Tile; + using TileScalarND = + Tile; + using TileScalarDN = Tile; + + constexpr int kDataBytes = TM * TN * sizeof(float); + constexpr int kScalarNDBytes = kScalarRows * kScalarCols * sizeof(float); + constexpr int kScalarDNBytes = kAlignedRows * sizeof(float); + + TileDataMxN oiNewTile; + TileDataMxN oiTile; + TileScalarND mijND, lijND, miND, liND; + TileScalarND miNewND, alphaND, betaND, tmpND; + TileScalarDN alphaDN, betaDN, liDN; + + TASSIGN(oiNewTile, 0); + TASSIGN(oiTile, kDataBytes); + TASSIGN(mijND, 2 * kDataBytes); + TASSIGN(lijND, 2 * kDataBytes + kScalarNDBytes); + TASSIGN(miND, 2 * kDataBytes + 2 * kScalarNDBytes); + TASSIGN(liND, 2 * kDataBytes + 3 * kScalarNDBytes); + TASSIGN(miNewND, 2 * kDataBytes + 4 * kScalarNDBytes); + TASSIGN(alphaND, 2 * kDataBytes + 5 * kScalarNDBytes); + TASSIGN(betaND, 2 * kDataBytes + 6 * kScalarNDBytes); + TASSIGN(tmpND, 2 * kDataBytes + 7 * kScalarNDBytes); + TASSIGN(alphaDN, 2 * kDataBytes + 8 * kScalarNDBytes); + TASSIGN(betaDN, 2 * kDataBytes + 8 * kScalarNDBytes + kScalarDNBytes); + TASSIGN(liDN, 2 * kDataBytes + 8 * kScalarNDBytes + 2 * kScalarDNBytes); + + if (is_first) { + TLOAD(oiNewTile, oiNewGlobal); + TLOAD(mijND, mijGlobalND); + TLOAD(lijND, lijGlobalND); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(miGlobalND, mijND); + TSTORE(liGlobalND, lijND); + TSTORE(oiGlobal, oiNewTile); + + if (is_last) { + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + TLOAD(liDN, liGlobalDN); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + TROWEXPANDDIV(oiNewTile, oiNewTile, liDN); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(dstGlobal, oiNewTile); + } + } else { + TLOAD(oiNewTile, oiNewGlobal); + TLOAD(oiTile, oiGlobal); + TLOAD(mijND, mijGlobalND); + TLOAD(lijND, lijGlobalND); + TLOAD(miND, miGlobalND); + TLOAD(liND, liGlobalND); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + TMAX(miNewND, miND, mijND); + pipe_barrier(PIPE_V); + TSUB(alphaND, miND, miNewND); + pipe_barrier(PIPE_V); + TEXP(alphaND, alphaND); + pipe_barrier(PIPE_V); + TSUB(betaND, mijND, miNewND); + pipe_barrier(PIPE_V); + TEXP(betaND, betaND); + pipe_barrier(PIPE_V); + TMUL(liND, alphaND, liND); + pipe_barrier(PIPE_V); + TMUL(tmpND, betaND, lijND); + pipe_barrier(PIPE_V); + TADD(liND, liND, tmpND); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(miGlobalND, miNewND); + TSTORE(liGlobalND, liND); + TSTORE(mijGlobalND, alphaND); + TSTORE(lijGlobalND, betaND); + + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + TLOAD(alphaDN, mijGlobalDN); + TLOAD(betaDN, lijGlobalDN); + if (is_last) { + TLOAD(liDN, liGlobalDN); + } + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + + TROWEXPANDMUL(oiTile, oiTile, alphaDN); + TROWEXPANDMUL(oiNewTile, oiNewTile, betaDN); + pipe_barrier(PIPE_V); + TADD(oiTile, oiTile, oiNewTile); + + if (is_last) { + pipe_barrier(PIPE_V); + TROWEXPANDDIV(oiTile, oiTile, liDN); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(dstGlobal, oiTile); + } else { + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(oiGlobal, oiTile); + } + } + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + // Safety check: if called with null tensor args (misrouted hub invocation), return. + if (args[0] == 0 || args[1] == 0 || args[2] == 0) { + return; + } + + __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *oi_new_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *mi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *li_acc_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + __gm__ Tensor *oi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[5]); + __gm__ Tensor *out_t = reinterpret_cast<__gm__ Tensor *>(args[6]); + uint64_t is_first = static_cast(args[7]); + uint64_t is_last = static_cast(args[8]); + int64_t num_heads = static_cast(args[9]); + int64_t head_dim = static_cast(args[10]); + int64_t q_loop = static_cast(args[11]); + + int32_t block_idx = get_block_idx(args); + int32_t sub_block_id = get_sub_block_id(args); // 0 = AIV0 (rows 0..7), 1 = AIV1 (rows 8..15) + int64_t batch_idx = block_idx / q_loop; + int64_t q_tile_idx = block_idx % q_loop; + + // Scalar layout: full QT=16 rows pack to kAlignedRowsFull=16 floats per block_idx; + // each AIV lane owns kAlignedRowsSub=8 contiguous floats inside that slab. + constexpr int kAlignedRowsFull = ((QT * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + constexpr int kAlignedRowsSub = ((SUB_QT * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + + int64_t row_offset = sub_block_id * SUB_QT; + + // Accumulator offsets (each AIV lane owns its own 8-row sub-slice within the block_idx slab) + int64_t scalar_offset = block_idx * kAlignedRowsFull + sub_block_id * kAlignedRowsSub; + int64_t data_offset = (block_idx * QT + row_offset) * head_dim; + + __gm__ float *mij_ptr = + reinterpret_cast<__gm__ float *>(mij_t->buffer.addr) + mij_t->start_offset + scalar_offset; + __gm__ float *lij_ptr = + reinterpret_cast<__gm__ float *>(lij_t->buffer.addr) + lij_t->start_offset + scalar_offset; + __gm__ float *oi_new_ptr = + reinterpret_cast<__gm__ float *>(oi_new_t->buffer.addr) + oi_new_t->start_offset + data_offset; + __gm__ float *mi_ptr = + reinterpret_cast<__gm__ float *>(mi_acc_t->buffer.addr) + mi_acc_t->start_offset + scalar_offset; + __gm__ float *li_ptr = + reinterpret_cast<__gm__ float *>(li_acc_t->buffer.addr) + li_acc_t->start_offset + scalar_offset; + __gm__ float *oi_ptr = + reinterpret_cast<__gm__ float *>(oi_acc_t->buffer.addr) + oi_acc_t->start_offset + data_offset; + + // Output offset: (batch_idx * num_heads + q_tile_idx * QT + row_offset, 0) + int64_t out_offset = (batch_idx * num_heads + q_tile_idx * QT + row_offset) * head_dim; + __gm__ float *dst_ptr = reinterpret_cast<__gm__ float *>(out_t->buffer.addr) + out_t->start_offset + out_offset; + + online_update_spmd(mij_ptr, lij_ptr, oi_new_ptr, mi_ptr, li_ptr, oi_ptr, dst_ptr, is_first, is_last); +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp new file mode 100644 index 0000000000..4c9ccdfbb4 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/aiv/aiv_softmax_prepare.cpp @@ -0,0 +1,344 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD Two-Pass Softmax Kernel (AIV) for n_blocks tiles with dual-vector split +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// The two AIV lanes in a cluster split the Q_TILE rows via get_sub_block_id(): +// AIV0 (sub_block_id=0) handles rows [0, SUB_M) +// AIV1 (sub_block_id=1) handles rows [SUB_M, Q_TILE) +// +// Memory layout: QK kernel writes packed (Q_TILE, block_size) tiles contiguously. +// Within each tile, AIV0 processes the first SUB_M rows and AIV1 the rest. +// Tile stride between consecutive tiles is Q_TILE * block_size elements. +// +// Two-pass softmax (same algorithm as paged_attention_unroll): +// Pass 1: Find global m = scale * max over all blocks of rowmax(S_i) +// Pass 2: Compute P_i = exp(S_i * scale - m) -> bf16, accumulate l = rowsum(P_i) +// +// Case1: SUB_M=8, TN=128 (Q_TILE=16, block_size=128) +// Case2: SUB_M=32, TN=64 (Q_TILE=64, block_size=64) +// +// Args: +// args[0] = sij Tensor* (spmd_blocks*Q_TILE, n_blocks*block_size) float32 +// args[1] = context_lens Tensor* (batch,) int32 +// args[2] = pij Tensor* (spmd_blocks*Q_TILE, n_blocks*block_size) bf16 [output] +// args[3] = mij Tensor* (spmd_blocks*Q_TILE,) float32 [output] +// args[4] = lij Tensor* (spmd_blocks*Q_TILE,) float32 [output] +// args[5] = scale_value scalar (as float bits in uint64) +// args[6] = bn_start scalar: starting KV block index +// args[7] = n_blocks scalar: number of KV blocks in this group +// args[8] = block_size scalar +// args[9] = q_loop scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +template +static __aicore__ void softmax_prepare_n_spmd( + __gm__ float *sij_base, float scale_value, __gm__ bfloat16_t *pij_base, __gm__ float *mij_addr, + __gm__ float *lij_addr, uint64_t n_blocks, uint64_t valid_len_last +) { + constexpr int kAlignedRows = ((TM * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + constexpr int kScalarCols = 32 / sizeof(float); + constexpr int kScalarRows = TM / kScalarCols; + + // --- GlobalTensor types --- + using GlobalDataMxN = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + using GlobalDataMxN_bf16 = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + using GlobalScalarDN = GlobalTensor, Stride<1, 1, 1, 1, 1>, Layout::DN>; + using GlobalScalarND = + GlobalTensor, Stride<1, 1, 1, kScalarCols, 1>>; + + // --- Tile types --- + using TileSijDyn = Tile; + using TileSijPad = + Tile; + using TileVecMxN = Tile; + using TileVecMxN_bf16 = Tile; + using TileScalarDN = Tile; + using TileScalarND = + Tile; + using TileScalarRow = Tile; + + // --- UB memory layout (double-buffered sij) --- + constexpr int kDataBytes = TM * TN * sizeof(float); + constexpr int kScalarDNBytes = kAlignedRows * sizeof(float); + + TileVecMxN sijTile_A; + TileSijPad sijPadTile_A; + TileVecMxN sijTile_B; + TileSijPad sijPadTile_B; + TileVecMxN pijTile; + TileVecMxN tmpTile; + TileVecMxN sumAccTile; + TileScalarDN localMaxDN; + TileScalarDN globalMaxDN; + TileScalarDN sumDN; + TileVecMxN_bf16 pijBf16Tile; + + TileScalarRow localMaxRow; + TileScalarRow globalMaxRow; + TileScalarND globalMaxND; + + TASSIGN(sijTile_A, 0x0); + TASSIGN(sijPadTile_A, 0x0); + TASSIGN(sijTile_B, kDataBytes); + TASSIGN(sijPadTile_B, kDataBytes); + TASSIGN(pijTile, 2 * kDataBytes); + TASSIGN(tmpTile, 3 * kDataBytes); + TASSIGN(sumAccTile, 4 * kDataBytes); + int scalarBase = 5 * kDataBytes; + TASSIGN(localMaxDN, scalarBase); + TASSIGN(localMaxRow, scalarBase); + TASSIGN(globalMaxDN, scalarBase + kScalarDNBytes); + TASSIGN(globalMaxRow, scalarBase + kScalarDNBytes); + TASSIGN(globalMaxND, scalarBase + kScalarDNBytes); + TASSIGN(sumDN, scalarBase + 2 * kScalarDNBytes); + TASSIGN(pijBf16Tile, scalarBase + 3 * kScalarDNBytes); + + GlobalScalarND mijGlobalND(mij_addr); + GlobalScalarDN lijGlobalDN(lij_addr); + + // ======== Pass 1: Find global row max (unscaled) ======== + // Tile stride between consecutive packed tiles = TILE_STRIDE elements + GlobalDataMxN sijGlobal_p1_0(sij_base); + TLOAD(sijTile_A, sijGlobal_p1_0); + + for (uint64_t i = 0; i < n_blocks; i++) { + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + if (i == n_blocks - 1 && valid_len_last < static_cast(TN)) { + TileSijDyn sijDynTile(static_cast(valid_len_last)); + if (i % 2 == 0) { + TASSIGN(sijDynTile, 0x0); + TFILLPAD_INPLACE(sijPadTile_A, sijDynTile); + } else { + TASSIGN(sijDynTile, static_cast(kDataBytes)); + TFILLPAD_INPLACE(sijPadTile_B, sijDynTile); + } + pipe_barrier(PIPE_V); + } + + if (i % 2 == 0) { + TROWMAX(localMaxDN, sijTile_A, tmpTile); + } else { + TROWMAX(localMaxDN, sijTile_B, tmpTile); + } + pipe_barrier(PIPE_V); + + if (i + 1 < n_blocks) { + GlobalDataMxN sijGlobal_next(sij_base + (i + 1) * TILE_STRIDE); + if (i % 2 == 0) { + TLOAD(sijTile_B, sijGlobal_next); + } else { + TLOAD(sijTile_A, sijGlobal_next); + } + } + + TRESHAPE(localMaxRow, localMaxDN); + if (i == 0) { + TMAX(globalMaxRow, localMaxRow, localMaxRow); + } else { + TMAX(globalMaxRow, globalMaxRow, localMaxRow); + } + pipe_barrier(PIPE_V); + } + + TMULS(globalMaxRow, globalMaxRow, scale_value); + pipe_barrier(PIPE_V); + TRESHAPE(globalMaxDN, globalMaxRow); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(mijGlobalND, globalMaxND); + + // ======== Pass 2: Compute softmax ======== + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + + GlobalDataMxN sijGlobal_0(sij_base); + TLOAD(sijTile_A, sijGlobal_0); + + for (uint64_t i = 0; i < n_blocks; i++) { + GlobalDataMxN_bf16 pijGlobal(pij_base + i * TILE_STRIDE); + + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + if (i == n_blocks - 1 && valid_len_last < static_cast(TN)) { + TileSijDyn curSijDyn(static_cast(valid_len_last)); + if (i % 2 == 0) { + TASSIGN(curSijDyn, 0x0); + TFILLPAD_INPLACE(sijPadTile_A, curSijDyn); + } else { + TASSIGN(curSijDyn, static_cast(kDataBytes)); + TFILLPAD_INPLACE(sijPadTile_B, curSijDyn); + } + pipe_barrier(PIPE_V); + } + + if (i % 2 == 0) { + TMULS(sijTile_A, sijTile_A, scale_value); + pipe_barrier(PIPE_V); + TROWEXPANDSUB(pijTile, sijTile_A, globalMaxDN); + } else { + TMULS(sijTile_B, sijTile_B, scale_value); + pipe_barrier(PIPE_V); + TROWEXPANDSUB(pijTile, sijTile_B, globalMaxDN); + } + pipe_barrier(PIPE_V); + TEXP(pijTile, pijTile); + pipe_barrier(PIPE_V); + TCVT(pijBf16Tile, pijTile, RoundMode::CAST_ROUND); + pipe_barrier(PIPE_V); + TCVT(pijTile, pijBf16Tile, RoundMode::CAST_ROUND); + + pipe_barrier(PIPE_V); + if (i == 0) { + TMULS(sumAccTile, pijTile, 1.0f); + } else { + TADD(sumAccTile, sumAccTile, pijTile); + } + + pipe_barrier(PIPE_V); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(pijGlobal, pijBf16Tile); + + if (i + 1 < n_blocks) { + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + GlobalDataMxN sijGlobal_next(sij_base + (i + 1) * TILE_STRIDE); + if (i % 2 == 0) { + TLOAD(sijTile_B, sijGlobal_next); + } else { + TLOAD(sijTile_A, sijGlobal_next); + } + } + } + + pipe_barrier(PIPE_V); + TROWSUM(sumDN, sumAccTile, tmpTile); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(lijGlobalDN, sumDN); + + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *sij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *pij_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + float scale_value = from_u64(static_cast(args[5])); + int64_t bn_start = static_cast(args[6]); + int64_t n_blocks_total = static_cast(args[7]); + int64_t block_size = static_cast(args[8]); + int64_t q_loop = static_cast(args[9]); + + int32_t block_idx = get_block_idx(args); + int32_t sub_block_id = get_sub_block_id(args); // 0 = AIV0, 1 = AIV1 + int64_t batch_idx = block_idx / q_loop; + + int64_t q_tile = (block_size == 128) ? 16 : 64; + int64_t sub_m = q_tile / 2; + + // Compute valid_len for the last block in this group + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t bn_this_batch = (cur_seq + block_size - 1) / block_size; + + // Clamp n_blocks to valid range for this batch + int64_t valid_blocks = bn_this_batch - bn_start; + if (valid_blocks < 0) valid_blocks = 0; + int64_t n_blocks = (valid_blocks < n_blocks_total) ? valid_blocks : n_blocks_total; + + // Compute valid_len for the last valid block + uint64_t valid_len_last; + if (n_blocks <= 0) { + valid_len_last = 0; + } else { + int64_t last_block_seq_start = (bn_start + n_blocks - 1) * block_size; + int64_t remaining = cur_seq - last_block_seq_start; + if (remaining >= block_size) { + valid_len_last = static_cast(block_size); + } else if (remaining > 0) { + valid_len_last = static_cast(remaining); + } else { + valid_len_last = 0; + } + } + + // Packed tile layout: SPMD block block_idx owns a contiguous region of + // n_blocks_total packed (q_tile, block_size) tiles. + // Packed base for this SPMD block: + int64_t packed_base_offset = block_idx * q_tile * n_blocks_total * block_size; + // AIV lane offset within each packed tile (sub_block_id selects the sub_m-row half): + int64_t sub_offset = sub_block_id * sub_m * block_size; + + __gm__ float *sij_base = + reinterpret_cast<__gm__ float *>(sij_t->buffer.addr) + sij_t->start_offset + packed_base_offset + sub_offset; + __gm__ bfloat16_t *pij_base = + reinterpret_cast<__gm__ bfloat16_t *>(pij_t->buffer.addr) + pij_t->start_offset + packed_base_offset + + sub_offset; + + // Scalar layout: full q_tile rows pack to kAlignedRowsFull floats per block_idx; + // each AIV lane owns kAlignedRowsSub contiguous floats inside that slab. + int64_t kAlignedRowsFull = ((q_tile * static_cast(sizeof(float)) + 31) / 32) * (32 / static_cast(sizeof(float))); + int64_t kAlignedRowsSub = ((sub_m * static_cast(sizeof(float)) + 31) / 32) * (32 / static_cast(sizeof(float))); + int64_t scalar_offset = block_idx * kAlignedRowsFull + sub_block_id * kAlignedRowsSub; + __gm__ float *mij_addr = + reinterpret_cast<__gm__ float *>(mij_t->buffer.addr) + mij_t->start_offset + scalar_offset; + __gm__ float *lij_addr = + reinterpret_cast<__gm__ float *>(lij_t->buffer.addr) + lij_t->start_offset + scalar_offset; + + if (n_blocks <= 0) { + // No valid KV data — emit neutral values + for (int64_t i = 0; i < kAlignedRowsSub; i++) { + mij_addr[i] = -1e30f; + lij_addr[i] = 0.0f; + } + return; + } + + // Tile stride = full packed tile size (q_tile * block_size), NOT sub_m * block_size + if (q_tile == 16) { + softmax_prepare_n_spmd<8, 128, 16 * 128>( + sij_base, scale_value, pij_base, mij_addr, lij_addr, + static_cast(n_blocks), valid_len_last + ); + } else { + softmax_prepare_n_spmd<32, 64, 64 * 64>( + sij_base, scale_value, pij_base, mij_addr, lij_addr, + static_cast(n_blocks), valid_len_last + ); + } +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py new file mode 100644 index 0000000000..f4eb9faf5d --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/kernel_config.py @@ -0,0 +1,86 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +""" +SPMD Paged Attention Kernel and Orchestration Configuration + +Uses SPMD (block_num) parallelism across batch*q_loop positions. +Each block handles one (batch_idx, q_tile_idx) using get_block_idx(). +q_tile adapts to num_heads: q_tile = min(num_heads, MAX_Q_TILE). + +Softmax and online-update run as MIX tasks (AIC idle + AIV0 + AIV1), with the +two AIVs splitting the q_tile query rows via get_sub_block_id(). + +AIC Kernels (Matrix Multiplication): + - aic_qk_matmul: Q @ K^T (SPMD across batch*q_loop) + - aic_pv_matmul: P @ V (SPMD across batch*q_loop) + - aic_hub: no-op, occupies the AIC slot of softmax/update MIX tasks + +AIV Kernels (Vector Operations): + - aiv_softmax_prepare: scale, rowmax, exp, rowsum on sub-tile (q_tile/2 rows) + - aiv_online_update: online softmax accumulation + normalization on sub-tile + - aiv_hub: no-op, used to allocate persistent accumulators +""" + +from pathlib import Path + +from simpler.task_interface import ArgDirection as D # pyright: ignore[reportAttributeAccessIssue] + +_KERNELS_ROOT = Path(__file__).parent + +ORCHESTRATION = { + "source": str(_KERNELS_ROOT / "orchestration" / "spmd_paged_attention_orch.cpp"), + "function_name": "aicpu_orchestration_entry", +} + +KERNELS = [ + # AIC kernels (matrix multiplication using Cube unit) + { + "func_id": 0, + "name": "SPMD_QK", + "source": str(_KERNELS_ROOT / "aic" / "aic_qk_matmul.cpp"), + "core_type": "aic", + }, + { + "func_id": 1, + "name": "SPMD_PV", + "source": str(_KERNELS_ROOT / "aic" / "aic_pv_matmul.cpp"), + "core_type": "aic", + }, + { + "func_id": 2, + "name": "AIC_HUB", + "source": str(_KERNELS_ROOT / "aic" / "aic_hub.cpp"), + "core_type": "aic", + }, + # AIV kernels (vector operations) + { + "func_id": 3, + "name": "SPMD_SF", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_softmax_prepare.cpp"), + "core_type": "aiv", + }, + { + "func_id": 4, + "name": "SPMD_UP", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_online_update.cpp"), + "core_type": "aiv", + }, + { + "func_id": 5, + "name": "AIV_HUB", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_hub.cpp"), + "core_type": "aiv", + }, +] + +RUNTIME_CONFIG = { + "runtime": "tensormap_and_ringbuffer", + "aicpu_thread_num": 4, + "block_dim": 24, +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp new file mode 100644 index 0000000000..ef6b67aa81 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention/kernels/orchestration/spmd_paged_attention_orch.cpp @@ -0,0 +1,248 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +/** + * SPMD Paged Attention Orchestration with N_UNROLL block batching + * + * Uses SPMD parallelism: block_num = batch * q_loop, where each logical + * block handles one (batch_idx, q_tile_idx) position. Kernels use + * get_block_idx() to compute their data offsets. + * + * Batches up to N_UNROLL KV blocks per group. Each group submits 4 tasks: + * 1. QK matmul: qi @ K^T for n_blocks → sij (Q_TILE, n_blocks * block_size) + * 2. Softmax: two-pass over sij → pij, mij, lij + * 3. PV matmul: SplitK accumulated P @ V → oi_new (Q_TILE, head_dim) + * 4. Update: online softmax accumulation with group-level mi, li, oi_new + * + * Softmax and online-update run as MIX tasks (AIC hub + AIV0 + AIV1) with + * dual-vector subvector partitioning via get_sub_block_id(). + * + * Memory Layout: + * Query: (batch, num_heads, head_dim) - bfloat16 + * Key/Value: (total_blocks, block_size, kv_head_num, head_dim) - bfloat16 + * Block Table: (batch, max_num_blocks_per_req) - int32 + * Context Lens: (batch,) - int32 + * Output: (batch, num_heads, head_dim) - float32 + */ + +#include +#include + +#include +#include + +#include "pto_orchestration_api.h" + +#define N_UNROLL 32 + +#define FUNC_QK_MATMUL 0 +#define FUNC_PV_MATMUL 1 +#define FUNC_AIC_HUB 2 +#define FUNC_SOFTMAX_PREPARE 3 +#define FUNC_ONLINE_UPDATE 4 +#define FUNC_AIV_HUB 5 + +static constexpr uint64_t Q_TILE = 16; + +extern "C" { + +__attribute__((visibility("default"))) PTO2OrchestrationConfig +aicpu_orchestration_config(const ChipStorageTaskArgs &orch_args) { + (void)orch_args; + return PTO2OrchestrationConfig{ + .expected_arg_count = 7, + }; +} + +__attribute__((visibility("default"))) void aicpu_orchestration_entry(const ChipStorageTaskArgs &orch_args) { + // query: shape=[batch, num_heads, head_dim] + uint64_t batch = orch_args.tensor(0).shapes[0]; + uint64_t num_heads = orch_args.tensor(0).shapes[1]; + uint64_t head_dim = orch_args.tensor(0).shapes[2]; + DataType data_type = orch_args.tensor(0).dtype; + + // key_cache: shape=[total_blocks, block_size, kv_head_num, head_dim] + uint64_t block_size = orch_args.tensor(1).shapes[1]; + + // block_table: shape=[batch, max_num_blocks_per_req] + uint64_t max_num_blocks_per_req = orch_args.tensor(3).shapes[1]; + + // scale from scalar arg + uint64_t scale_value = orch_args.scalar(0); + + uint64_t q_loop = (num_heads + Q_TILE - 1) / Q_TILE; + int16_t spmd_block_num = static_cast(batch * q_loop); + + LOG_INFO( + "SPMD PA Unroll: batch=%" PRIu64 " heads=%" PRIu64 " hd=%" PRIu64 " bs=%" PRIu64 " q_loop=%" PRIu64 + " blocks=%d N_UNROLL=%d", + batch, num_heads, head_dim, block_size, q_loop, spmd_block_num, N_UNROLL + ); + + // Wrap host-provided tensors + void *query_ptr = orch_args.tensor(0).data_as(); + void *kc_ptr = orch_args.tensor(1).data_as(); + void *vc_ptr = orch_args.tensor(2).data_as(); + void *out_ptr = orch_args.tensor(5).data_as(); + + uint64_t total_kv_blocks = orch_args.tensor(1).shapes[0]; + uint64_t kv_total_rows = total_kv_blocks * block_size; + + uint32_t query_shapes[2] = {static_cast(batch * num_heads), static_cast(head_dim)}; + uint32_t kv_shapes[2] = {static_cast(kv_total_rows), static_cast(head_dim)}; + uint32_t out_shapes[2] = {static_cast(batch * num_heads), static_cast(head_dim)}; + + Tensor query = make_tensor_external(query_ptr, query_shapes, 2, data_type); + Tensor key_cache = make_tensor_external(kc_ptr, kv_shapes, 2, data_type); + Tensor value_cache = make_tensor_external(vc_ptr, kv_shapes, 2, data_type); + Tensor out = make_tensor_external(out_ptr, out_shapes, 2, DataType::FLOAT32); + + uint32_t bt_shapes[2] = {static_cast(batch), static_cast(max_num_blocks_per_req)}; + Tensor block_table = + make_tensor_external(orch_args.tensor(3).data_as(), bt_shapes, 2, DataType::INT32, false); + uint32_t cl_shapes[1] = {static_cast(batch)}; + Tensor context_lens = + make_tensor_external(orch_args.tensor(4).data_as(), cl_shapes, 1, DataType::INT32, false); + + // Find max context_len for KV block loop bound + uint64_t max_ctx = 0; + for (uint64_t b = 0; b < batch; b++) { + uint32_t idx[1] = {static_cast(b)}; + uint64_t ctx = static_cast(get_tensor_data(context_lens, 1, idx)); + if (ctx > max_ctx) max_ctx = ctx; + } + uint64_t max_bn = (max_ctx + block_size - 1) / block_size; + + // Accumulator create infos (persistent across the bn loop) + uint32_t n_rows = static_cast(spmd_block_num) * static_cast(Q_TILE); + uint32_t acc_oi_shapes[2] = {n_rows, static_cast(head_dim)}; + uint32_t scalar_shapes[1] = {n_rows}; + TensorCreateInfo acc_oi_ci(acc_oi_shapes, 2, DataType::FLOAT32); + TensorCreateInfo acc_mi_ci(scalar_shapes, 1, DataType::FLOAT32); + TensorCreateInfo acc_li_ci(scalar_shapes, 1, DataType::FLOAT32); + + PTO2_SCOPE() { + // Allocate persistent accumulators via no-op AIV hub + Arg hub_args; + hub_args.add_output(acc_oi_ci); + hub_args.add_output(acc_mi_ci); + hub_args.add_output(acc_li_ci); + TaskOutputTensors hub_outs = pto2_rt_submit_aiv_task(FUNC_AIV_HUB, hub_args); + const Tensor &oi_acc = hub_outs.get_ref(0); + const Tensor &mi_acc = hub_outs.get_ref(1); + const Tensor &li_acc = hub_outs.get_ref(2); + + for (uint64_t bn = 0; bn < max_bn; bn += N_UNROLL) { + uint64_t n_blocks = std::min(static_cast(N_UNROLL), max_bn - bn); + uint64_t is_first = (bn == 0) ? 1 : 0; + uint64_t is_last = (bn + n_blocks >= max_bn) ? 1 : 0; + + // Scratch tensor create infos sized for n_blocks per SPMD block + uint32_t sij_shapes[2] = {n_rows, static_cast(n_blocks * block_size)}; + uint32_t pij_shapes[2] = {n_rows, static_cast(n_blocks * block_size)}; + uint32_t oi_new_shapes[2] = {n_rows, static_cast(head_dim)}; + uint32_t mij_shapes[1] = {n_rows}; + uint32_t lij_shapes[1] = {n_rows}; + + TensorCreateInfo sij_ci(sij_shapes, 2, DataType::FLOAT32); + TensorCreateInfo pij_ci(pij_shapes, 2, data_type); + TensorCreateInfo oi_new_ci(oi_new_shapes, 2, DataType::FLOAT32); + TensorCreateInfo mij_ci(mij_shapes, 1, DataType::FLOAT32); + TensorCreateInfo lij_ci(lij_shapes, 1, DataType::FLOAT32); + + // -- QK Matmul (AIC, SPMD): n_blocks matmuls per SPMD block -- + Arg qk_args; + qk_args.add_input(query); + qk_args.add_input(key_cache); + qk_args.add_input(block_table); + qk_args.add_input(context_lens); + qk_args.add_output(sij_ci); + qk_args.add_scalar(static_cast(bn)); + qk_args.add_scalar(static_cast(n_blocks)); + qk_args.add_scalar(static_cast(num_heads)); + qk_args.add_scalar(static_cast(head_dim)); + qk_args.add_scalar(static_cast(block_size)); + qk_args.add_scalar(static_cast(max_num_blocks_per_req)); + qk_args.add_scalar(static_cast(q_loop)); + qk_args.launch_spec.set_block_num(spmd_block_num); + TaskOutputTensors qk_outs = pto2_rt_submit_aic_task(FUNC_QK_MATMUL, qk_args); + const Tensor &sij = qk_outs.get_ref(0); + + // -- Softmax Prepare (MIX: AIC hub + AIV0 + AIV1, SPMD) -- + // Two-pass softmax over n_blocks tiles per SPMD block + Arg sf_args; + sf_args.add_input(sij); + sf_args.add_input(context_lens); + sf_args.add_output(pij_ci); + sf_args.add_output(mij_ci); + sf_args.add_output(lij_ci); + sf_args.add_scalar(scale_value); + sf_args.add_scalar(static_cast(bn)); + sf_args.add_scalar(static_cast(n_blocks)); + sf_args.add_scalar(static_cast(block_size)); + sf_args.add_scalar(static_cast(q_loop)); + sf_args.launch_spec.set_block_num(spmd_block_num); + MixedKernels sf_mk; + sf_mk.aic_kernel_id = FUNC_AIC_HUB; + sf_mk.aiv0_kernel_id = FUNC_SOFTMAX_PREPARE; + sf_mk.aiv1_kernel_id = FUNC_SOFTMAX_PREPARE; + TaskOutputTensors sf_outs = pto2_rt_submit_task(sf_mk, sf_args); + const Tensor &pij = sf_outs.get_ref(0); + const Tensor &mij = sf_outs.get_ref(1); + const Tensor &lij = sf_outs.get_ref(2); + + // -- PV Matmul (AIC, SPMD): SplitK accumulated across n_blocks -- + Arg pv_args; + pv_args.add_input(pij); + pv_args.add_input(value_cache); + pv_args.add_input(block_table); + pv_args.add_input(context_lens); + pv_args.add_output(oi_new_ci); + pv_args.add_scalar(static_cast(bn)); + pv_args.add_scalar(static_cast(n_blocks)); + pv_args.add_scalar(static_cast(num_heads)); + pv_args.add_scalar(static_cast(head_dim)); + pv_args.add_scalar(static_cast(block_size)); + pv_args.add_scalar(static_cast(max_num_blocks_per_req)); + pv_args.add_scalar(static_cast(q_loop)); + pv_args.launch_spec.set_block_num(spmd_block_num); + TaskOutputTensors pv_outs = pto2_rt_submit_aic_task(FUNC_PV_MATMUL, pv_args); + const Tensor &oi_new = pv_outs.get_ref(0); + + // -- Online Update (MIX: AIC hub + AIV0 + AIV1, SPMD) -- + Arg up_args; + up_args.add_input(mij); + up_args.add_input(lij); + up_args.add_input(oi_new); + up_args.add_inout(mi_acc); + up_args.add_inout(li_acc); + up_args.add_inout(oi_acc); + up_args.add_inout(out); + up_args.add_scalar(is_first); + up_args.add_scalar(is_last); + up_args.add_scalar(static_cast(num_heads)); + up_args.add_scalar(static_cast(head_dim)); + up_args.add_scalar(static_cast(q_loop)); + up_args.launch_spec.set_block_num(spmd_block_num); + MixedKernels up_mk; + up_mk.aic_kernel_id = FUNC_AIC_HUB; + up_mk.aiv0_kernel_id = FUNC_ONLINE_UPDATE; + up_mk.aiv1_kernel_id = FUNC_ONLINE_UPDATE; + pto2_rt_submit_task(up_mk, up_args); + } + } + + uint64_t n_groups = (max_bn + N_UNROLL - 1) / N_UNROLL; + LOG_INFO( + "SPMD PA Unroll: %" PRIu64 " groups x 4 tasks, blocks=%d", n_groups, static_cast(spmd_block_num) + ); +} + +} // extern "C" diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/golden.py b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/golden.py new file mode 100644 index 0000000000..def701c3fc --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/golden.py @@ -0,0 +1,67 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""SPMD Paged Attention Golden - fixed block_num=24 variant (bfloat16). + +Uses SPMD parallelism with a fixed hardware block_num of 24. +Each hardware block strides over batch*q_loop logical work items. +""" + +from paged_attention_golden import ( + compute_golden, # noqa: F401 + run_golden_test, +) +from paged_attention_golden import generate_inputs as _generate_inputs + +__outputs__ = ["out"] + +RTOL = 1e-3 +ATOL = 1e-3 + +ALL_CASES = { + "Case1": { + "batch": 256, + "num_heads": 16, + "kv_head_num": 1, + "head_dim": 128, + "block_size": 128, + "context_len": 8192, + "max_model_len": 32768, + "dtype": "bfloat16", + }, + "Case2": { + "batch": 64, + "num_heads": 64, + "kv_head_num": 1, + "head_dim": 128, + "block_size": 64, + "context_len": 8192, + "max_model_len": 32768, + "dtype": "bfloat16", + }, + "Case3": { + "batch": 64, + "num_heads": 64, + "kv_head_num": 1, + "head_dim": 256, + "block_size": 64, + "context_len": 8192, + "max_model_len": 32768, + "dtype": "bfloat16", + }, +} + +DEFAULT_CASE = "Case1" + + +def generate_inputs(params: dict) -> list: + return _generate_inputs(params) + + +if __name__ == "__main__": + run_golden_test(ALL_CASES, DEFAULT_CASE, generate_inputs, label="SPMD Paged Attention") diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_hub.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_hub.cpp new file mode 100644 index 0000000000..eb602e0bd6 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_hub.cpp @@ -0,0 +1,29 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// AIC Hub Kernel - No-op stub used as the AIC slot of MIX (AIC+AIV0+AIV1) tasks +// when the real work happens only on the two AIVs (softmax, online update). +// Pairing an idle AIC with two active AIVs forces the scheduler to allocate a +// full cluster, which is what enables the two AIV lanes to run in parallel. + +#include +#include + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) {} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_pv_matmul.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_pv_matmul.cpp new file mode 100644 index 0000000000..375cfd005c --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_pv_matmul.cpp @@ -0,0 +1,167 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD PV Matmul: pij(q_tile, K) @ vj(K, N) -> oi_new(q_tile, N) +// +// Hardware block_num is fixed at 24. Each hardware block strides over +// total_blocks logical work items: +// for (idx = hw_block_idx; idx < total_blocks; idx += block_num) +// Each logical block_idx encodes (batch_idx, q_tile_idx). +// q_tile is passed as a runtime scalar and dispatched to the matching template. +// +// Args: +// args[0] = pij Tensor* (total_blocks*q_tile, block_size) data_type +// args[1] = value_cache Tensor* (kv_total_rows, head_dim) bf16 +// args[2] = block_table Tensor* (batch, max_blocks_per_req) int32 +// args[3] = context_lens Tensor* (batch,) int32 +// args[4] = oi_new Tensor* (total_blocks*q_tile, head_dim) float32 [output] +// args[5] = bn scalar: current KV block index +// args[6] = num_heads scalar +// args[7] = head_dim scalar +// args[8] = block_size scalar +// args[9] = max_num_blocks_per_req scalar +// args[10] = q_loop scalar +// args[11] = q_tile scalar +// args[12] = total_blocks scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +static constexpr int K_128 = 128; +static constexpr int K_64 = 64; +static constexpr int N_128 = 128; + +template +static __aicore__ void pv_matmul_spmd(__gm__ bfloat16_t *pij_addr, __gm__ bfloat16_t *vj_addr, __gm__ float *oi_addr) { + using GlobalA = GlobalTensor, Stride>; + using GlobalB = GlobalTensor, Stride>; + using GlobalOut = GlobalTensor, Stride>; + + GlobalA pijGlobal(pij_addr); + GlobalB vjGlobal(vj_addr); + GlobalOut oiGlobal(oi_addr); + + using TileMatA = Tile; + using TileMatB = Tile; + + using LeftTile = TileLeft; + using RightTile = TileRight; + using AccTile = TileAcc; + + TileMatA aMatTile; + TileMatB bMatTile; + TASSIGN(aMatTile, 0x0); + TASSIGN(bMatTile, 0x20000); + + LeftTile aTile; + RightTile bTile; + AccTile cTile; + TASSIGN(aTile, 0x0); + TASSIGN(bTile, 0x0); + TASSIGN(cTile, 0x0); + + TLOAD(aMatTile, pijGlobal); + TLOAD(bMatTile, vjGlobal); + + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + + TMOV(aTile, aMatTile); + TMOV(bTile, bMatTile); + + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + + TMATMUL(cTile, aTile, bTile); + + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + + TSTORE(oiGlobal, cTile); + + set_flag(PIPE_FIX, PIPE_S, EVENT_ID7); + wait_flag(PIPE_FIX, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *pij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *value_cache_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *block_table_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *oi_new_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + + int64_t bn = static_cast(args[5]); + int64_t num_heads = static_cast(args[6]); + int64_t head_dim = static_cast(args[7]); + int64_t block_size = static_cast(args[8]); + int64_t max_blocks_per_req = static_cast(args[9]); + int64_t q_loop = static_cast(args[10]); + int64_t q_tile = static_cast(args[11]); + int64_t total_blocks = static_cast(args[12]); + + int32_t hw_block_idx = get_block_idx(args); + int32_t block_num = get_block_num(args); + + for (int32_t block_idx = hw_block_idx; block_idx < total_blocks; block_idx += block_num) { + int64_t batch_idx = block_idx / q_loop; + + // Check if this batch has data at this KV block + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t bn_this_batch = (cur_seq + block_size - 1) / block_size; + + // Output pointer for this block's oi_new slice + __gm__ float *oi_addr = reinterpret_cast<__gm__ float *>(oi_new_t->buffer.addr) + oi_new_t->start_offset + + block_idx * q_tile * head_dim; + + if (bn >= bn_this_batch) { + for (int i = 0; i < static_cast(q_tile * head_dim); i++) { + oi_addr[i] = 0.0f; + } + continue; + } + + // Look up physical block index + __gm__ int32_t *bt_ptr = + reinterpret_cast<__gm__ int32_t *>(block_table_t->buffer.addr) + block_table_t->start_offset; + int64_t phys_block = static_cast(bt_ptr[batch_idx * max_blocks_per_req + bn]); + + // pij offset: block_idx * q_tile * block_size + int64_t pij_offset = block_idx * q_tile * block_size; + __gm__ bfloat16_t *pij_addr = + reinterpret_cast<__gm__ bfloat16_t *>(pij_t->buffer.addr) + pij_t->start_offset + pij_offset; + + // Value offset: phys_block * block_size * head_dim + int64_t v_offset = phys_block * block_size * head_dim; + __gm__ bfloat16_t *vj_addr = + reinterpret_cast<__gm__ bfloat16_t *>(value_cache_t->buffer.addr) + value_cache_t->start_offset + v_offset; + + if (q_tile == 16) { + pv_matmul_spmd<16, K_128, N_128>(pij_addr, vj_addr, oi_addr); + } else { + pv_matmul_spmd<64, K_64, N_128>(pij_addr, vj_addr, oi_addr); + } + } +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_qk_matmul.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_qk_matmul.cpp new file mode 100644 index 0000000000..2620c51cb7 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aic/aic_qk_matmul.cpp @@ -0,0 +1,171 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD QK Matmul: qi(q_tile, K) @ kj.T(K, N) -> sij(q_tile, N) +// +// Hardware block_num is fixed at 24. Each hardware block strides over +// total_blocks logical work items: +// for (idx = hw_block_idx; idx < total_blocks; idx += block_num) +// Each logical block_idx encodes (batch_idx, q_tile_idx). +// q_tile is passed as a runtime scalar and dispatched to the matching template. +// +// Args: +// args[0] = query Tensor* (batch*num_heads, head_dim) bf16 +// args[1] = key_cache Tensor* (kv_total_rows, head_dim) bf16 +// args[2] = block_table Tensor* (batch, max_blocks_per_req) int32 +// args[3] = context_lens Tensor* (batch,) int32 +// args[4] = sij Tensor* (total_blocks*q_tile, block_size) float32 [output] +// args[5] = bn scalar: current KV block index +// args[6] = num_heads scalar +// args[7] = head_dim scalar +// args[8] = block_size scalar +// args[9] = max_num_blocks_per_req scalar +// args[10] = q_loop scalar +// args[11] = q_tile scalar +// args[12] = total_blocks scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +static constexpr int K_128 = 128; +static constexpr int N_128 = 128; +static constexpr int N_64 = 64; + +template +static __aicore__ void qk_matmul_spmd(__gm__ bfloat16_t *qi_addr, __gm__ bfloat16_t *kj_addr, __gm__ float *sij_addr) { + using GlobalA = GlobalTensor, Stride>; + using GlobalB = + GlobalTensor, Stride, Layout::DN>; + using GlobalOut = GlobalTensor, Stride>; + + GlobalA qiGlobal(qi_addr); + GlobalB kjGlobal(kj_addr); + GlobalOut sijGlobal(sij_addr); + + using TileMatA = Tile; + using TileMatB = Tile; + + using LeftTile = TileLeft; + using RightTile = TileRight; + using AccTile = TileAcc; + + TileMatA aMatTile; + TileMatB bMatTile; + TASSIGN(aMatTile, 0x0); + TASSIGN(bMatTile, 0x20000); + + LeftTile aTile; + RightTile bTile; + AccTile cTile; + TASSIGN(aTile, 0x0); + TASSIGN(bTile, 0x0); + TASSIGN(cTile, 0x0); + + TLOAD(aMatTile, qiGlobal); + TLOAD(bMatTile, kjGlobal); + + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + + TMOV(aTile, aMatTile); + TMOV(bTile, bMatTile); + + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + + TMATMUL(cTile, aTile, bTile); + + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + + TSTORE(sijGlobal, cTile); + + set_flag(PIPE_FIX, PIPE_S, EVENT_ID7); + wait_flag(PIPE_FIX, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *query_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *key_cache_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *block_table_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *sij_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + + int64_t bn = static_cast(args[5]); + int64_t num_heads = static_cast(args[6]); + int64_t head_dim = static_cast(args[7]); + int64_t block_size = static_cast(args[8]); + int64_t max_blocks_per_req = static_cast(args[9]); + int64_t q_loop = static_cast(args[10]); + int64_t q_tile = static_cast(args[11]); + int64_t total_blocks = static_cast(args[12]); + + int32_t hw_block_idx = get_block_idx(args); + int32_t block_num = get_block_num(args); + + for (int32_t block_idx = hw_block_idx; block_idx < total_blocks; block_idx += block_num) { + // Decode (batch_idx, q_tile_idx) from block_idx + int64_t batch_idx = block_idx / q_loop; + int64_t q_tile_idx = block_idx % q_loop; + + // Check if this batch has data at this KV block + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t bn_this_batch = (cur_seq + block_size - 1) / block_size; + + // Output pointer for this block's sij slice + __gm__ float *sij_addr = reinterpret_cast<__gm__ float *>(sij_t->buffer.addr) + sij_t->start_offset + + block_idx * q_tile * block_size; + + if (bn >= bn_this_batch) { + // No valid KV data for this batch at this bn — zero out sij + for (int i = 0; i < static_cast(q_tile * block_size); i++) { + sij_addr[i] = 0.0f; + } + continue; + } + + // Look up physical block index from block_table + __gm__ int32_t *bt_ptr = + reinterpret_cast<__gm__ int32_t *>(block_table_t->buffer.addr) + block_table_t->start_offset; + int64_t phys_block = static_cast(bt_ptr[batch_idx * max_blocks_per_req + bn]); + + // Query offset: (batch_idx * num_heads + q_tile_idx * q_tile, 0) + int64_t q_offset = (batch_idx * num_heads + q_tile_idx * q_tile) * head_dim; + __gm__ bfloat16_t *qi_addr = + reinterpret_cast<__gm__ bfloat16_t *>(query_t->buffer.addr) + query_t->start_offset + q_offset; + + // Key offset: (phys_block * block_size, 0) + int64_t k_offset = phys_block * block_size * head_dim; + __gm__ bfloat16_t *kj_addr = + reinterpret_cast<__gm__ bfloat16_t *>(key_cache_t->buffer.addr) + key_cache_t->start_offset + k_offset; + + if (q_tile == 16) { + qk_matmul_spmd<16, K_128, N_128>(qi_addr, kj_addr, sij_addr); + } else { + qk_matmul_spmd<64, K_128, N_64>(qi_addr, kj_addr, sij_addr); + } + } +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_hub.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_hub.cpp new file mode 100644 index 0000000000..a42f2790e9 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_hub.cpp @@ -0,0 +1,27 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// AIV Hub Kernel - No-op stub for accumulator tensor allocation. +// The runtime allocates output tensors specified in the Arg; the kernel itself does nothing. + +#include +#include + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) {} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_online_update.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_online_update.cpp new file mode 100644 index 0000000000..0b64d6bed8 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_online_update.cpp @@ -0,0 +1,316 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD Online Softmax Update + Normalize Kernel (AIV) with dual-vector +// subvector split. +// +// Hardware block_num is fixed at 24. Each hardware block strides over +// total_blocks logical work items: +// for (idx = hw_block_idx; idx < total_blocks; idx += block_num) +// Each logical block_idx encodes (batch_idx, q_tile_idx). +// The two AIV lanes in a cluster split the q_tile rows via +// get_sub_block_id(): AIV0 updates the first half, AIV1 the second half. +// q_tile is passed as a runtime scalar and dispatched to the matching template. +// The online softmax update is row-independent, so the two lanes never touch +// the same row of mi/li/oi accumulators or the output buffer. +// +// Hardware safety: TROWEXPANDMUL/DIV require a minimum tile height of 16 rows. +// When the sub-tile has fewer rows (e.g., 8), we use 16-row compute tiles with +// pad rows carrying -inf/0 values. Scalar tiles (mi, li, alpha, beta) are also +// padded to 16 elements per lane, stored contiguously in GM so that DN TLOAD +// can read them back with the correct stride. +// +// Scalar layout strategy: +// M scalar floats stored contiguously in GM can be loaded as either: +// - ND (kScalarRows, kScalarCols) RowMajor for element-wise ops +// - DN (kAlignedRows, 1) ColMajor for row-broadcast ops (TROWEXPANDMUL/DIV) +// Conversion between layouts uses GM round-trip: ND TSTORE -> DN TLOAD. +// +// Args: +// args[0] = mij Tensor* (padded scalar buf) float32 +// args[1] = lij Tensor* (padded scalar buf) float32 +// args[2] = oi_new Tensor* (total_blocks*q_tile, head_dim) float32 +// args[3] = mi_acc Tensor* (padded scalar buf) float32 [inout] +// args[4] = li_acc Tensor* (padded scalar buf) float32 [inout] +// args[5] = oi_acc Tensor* (total_blocks*q_tile, head_dim) float32 [inout] +// args[6] = out Tensor* (batch*num_heads, head_dim) float32 [inout] +// args[7] = is_first scalar +// args[8] = is_last scalar +// args[9] = num_heads scalar +// args[10] = head_dim scalar +// args[11] = q_loop scalar +// args[12] = q_tile scalar +// args[13] = total_blocks scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +static constexpr int HD = 128; // Head dimension +static constexpr int MIN_TM = 16; // minimum tile height for TROW* hw safety + +// TM = actual number of valid rows this AIV lane owns. +// All TROW* instructions operate on padded PM(=16)-row tiles to avoid the +// hardware dstRptStride issue at TM=8. Data tiles (oi) and scalar tiles (mi, li) +// are padded accordingly; only TM valid rows are loaded from / stored to GM. +template +static __aicore__ void online_update_spmd( + __gm__ float *mij_ptr, __gm__ float *lij_ptr, __gm__ float *oi_new_ptr, __gm__ float *mi_ptr, __gm__ float *li_ptr, + __gm__ float *oi_ptr, __gm__ float *dst_ptr, uint64_t is_first, uint64_t is_last +) { + constexpr int PM = (TM < MIN_TM) ? MIN_TM : TM; // padded tile height + constexpr int kScalarCols = 32 / sizeof(float); + constexpr int kScalarRows = PM / kScalarCols; + constexpr int kAlignedRows = ((PM * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + + // GM accessors for data: load/store only TM valid rows + using GlobalDataTMxN = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + // GM accessors for scalars: load/store padded kAlignedRows + using GlobalScalarND = + GlobalTensor, Stride<1, 1, 1, kScalarCols, 1>>; + using GlobalScalarDN = GlobalTensor, Stride<1, 1, 1, 1, 1>, Layout::DN>; + + GlobalDataTMxN oiNewGlobal(oi_new_ptr); + GlobalDataTMxN oiGlobal(oi_ptr); + GlobalDataTMxN dstGlobal(dst_ptr); + + GlobalScalarND mijGlobalND(mij_ptr); + GlobalScalarND lijGlobalND(lij_ptr); + GlobalScalarND miGlobalND(mi_ptr); + GlobalScalarND liGlobalND(li_ptr); + + GlobalScalarDN mijGlobalDN(mij_ptr); + GlobalScalarDN lijGlobalDN(lij_ptr); + GlobalScalarDN liGlobalDN(li_ptr); + + // Compute tiles use PM(=16) rows + using TileDataPMxN = Tile; + using TileScalarND = + Tile; + using TileScalarDN = Tile; + + // Load/store tiles (TM rows) aliased to start of padded tiles + using TileLoadTMxN = Tile; + + constexpr int kDataBytes = PM * TN * sizeof(float); + constexpr int kScalarNDBytes = kScalarRows * kScalarCols * sizeof(float); + constexpr int kScalarDNBytes = kAlignedRows * sizeof(float); + + TileDataPMxN oiNewTile; + TileDataPMxN oiTile; + TileLoadTMxN oiNewLoadTile; + TileLoadTMxN oiLoadTile; + TileLoadTMxN oiStoreTile; + TileScalarND mijND, lijND, miND, liND; + TileScalarND miNewND, alphaND, betaND, tmpND; + TileScalarDN alphaDN, betaDN, liDN; + + TASSIGN(oiNewTile, 0); + TASSIGN(oiNewLoadTile, 0); // alias + TASSIGN(oiTile, kDataBytes); + TASSIGN(oiLoadTile, kDataBytes); // alias + TASSIGN(oiStoreTile, kDataBytes); // alias + TASSIGN(mijND, 2 * kDataBytes); + TASSIGN(lijND, 2 * kDataBytes + kScalarNDBytes); + TASSIGN(miND, 2 * kDataBytes + 2 * kScalarNDBytes); + TASSIGN(liND, 2 * kDataBytes + 3 * kScalarNDBytes); + TASSIGN(miNewND, 2 * kDataBytes + 4 * kScalarNDBytes); + TASSIGN(alphaND, 2 * kDataBytes + 5 * kScalarNDBytes); + TASSIGN(betaND, 2 * kDataBytes + 6 * kScalarNDBytes); + TASSIGN(tmpND, 2 * kDataBytes + 7 * kScalarNDBytes); + TASSIGN(alphaDN, 2 * kDataBytes + 8 * kScalarNDBytes); + TASSIGN(betaDN, 2 * kDataBytes + 8 * kScalarNDBytes + kScalarDNBytes); + TASSIGN(liDN, 2 * kDataBytes + 8 * kScalarNDBytes + 2 * kScalarDNBytes); + + if (is_first) { + TLOAD(oiNewLoadTile, oiNewGlobal); + TLOAD(mijND, mijGlobalND); + TLOAD(lijND, lijGlobalND); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(miGlobalND, mijND); + TSTORE(liGlobalND, lijND); + TSTORE(oiGlobal, oiNewLoadTile); + + if (is_last) { + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + TLOAD(liDN, liGlobalDN); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + TROWEXPANDDIV(oiNewTile, oiNewTile, liDN); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(dstGlobal, oiNewLoadTile); + } + } else { + TLOAD(oiNewLoadTile, oiNewGlobal); + TLOAD(oiLoadTile, oiGlobal); + TLOAD(mijND, mijGlobalND); + TLOAD(lijND, lijGlobalND); + TLOAD(miND, miGlobalND); + TLOAD(liND, liGlobalND); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + TMAX(miNewND, miND, mijND); + pipe_barrier(PIPE_V); + TSUB(alphaND, miND, miNewND); + pipe_barrier(PIPE_V); + TEXP(alphaND, alphaND); + pipe_barrier(PIPE_V); + TSUB(betaND, mijND, miNewND); + pipe_barrier(PIPE_V); + TEXP(betaND, betaND); + pipe_barrier(PIPE_V); + TMUL(liND, alphaND, liND); + pipe_barrier(PIPE_V); + TMUL(tmpND, betaND, lijND); + pipe_barrier(PIPE_V); + TADD(liND, liND, tmpND); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(miGlobalND, miNewND); + TSTORE(liGlobalND, liND); + TSTORE(mijGlobalND, alphaND); + TSTORE(lijGlobalND, betaND); + + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + TLOAD(alphaDN, mijGlobalDN); + TLOAD(betaDN, lijGlobalDN); + if (is_last) { + TLOAD(liDN, liGlobalDN); + } + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + + TROWEXPANDMUL(oiTile, oiTile, alphaDN); + TROWEXPANDMUL(oiNewTile, oiNewTile, betaDN); + pipe_barrier(PIPE_V); + TADD(oiTile, oiTile, oiNewTile); + + if (is_last) { + pipe_barrier(PIPE_V); + TROWEXPANDDIV(oiTile, oiTile, liDN); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(dstGlobal, oiStoreTile); + } else { + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(oiGlobal, oiStoreTile); + } + } + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); +} + +template +static __aicore__ void online_update_entry( + __gm__ Tensor *mij_t, __gm__ Tensor *lij_t, __gm__ Tensor *oi_new_t, __gm__ Tensor *mi_acc_t, + __gm__ Tensor *li_acc_t, __gm__ Tensor *oi_acc_t, __gm__ Tensor *out_t, uint64_t is_first, uint64_t is_last, + int64_t num_heads, int64_t head_dim, int64_t q_loop, int64_t total_blocks, __gm__ int64_t *args +) { + constexpr int SUB_QT = QT / 2; + // Padded sub-tile height for hw-safe TROW* ops + constexpr int PAD_SUB_QT = (SUB_QT < MIN_TM) ? MIN_TM : SUB_QT; + + int32_t hw_block_idx = get_block_idx(args); + int32_t block_num = get_block_num(args); + int32_t sub_block_id = get_sub_block_id(args); + + // Scalar layout uses padded kAlignedRows per sub-tile (based on PAD_SUB_QT) + constexpr int kAlignedRowsFull = 2 * (((PAD_SUB_QT * sizeof(float) + 31) / 32) * (32 / sizeof(float))); + constexpr int kAlignedRowsSub = ((PAD_SUB_QT * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + + for (int32_t block_idx = hw_block_idx; block_idx < total_blocks; block_idx += block_num) { + int64_t batch_idx = block_idx / q_loop; + int64_t q_tile_idx = block_idx % q_loop; + + int64_t row_offset = sub_block_id * SUB_QT; + + // Accumulator offsets (each AIV lane owns its own sub-slice within the block_idx slab) + int64_t scalar_offset = block_idx * kAlignedRowsFull + sub_block_id * kAlignedRowsSub; + int64_t data_offset = (block_idx * QT + row_offset) * head_dim; + + __gm__ float *mij_ptr = + reinterpret_cast<__gm__ float *>(mij_t->buffer.addr) + mij_t->start_offset + scalar_offset; + __gm__ float *lij_ptr = + reinterpret_cast<__gm__ float *>(lij_t->buffer.addr) + lij_t->start_offset + scalar_offset; + __gm__ float *oi_new_ptr = + reinterpret_cast<__gm__ float *>(oi_new_t->buffer.addr) + oi_new_t->start_offset + data_offset; + __gm__ float *mi_ptr = + reinterpret_cast<__gm__ float *>(mi_acc_t->buffer.addr) + mi_acc_t->start_offset + scalar_offset; + __gm__ float *li_ptr = + reinterpret_cast<__gm__ float *>(li_acc_t->buffer.addr) + li_acc_t->start_offset + scalar_offset; + __gm__ float *oi_ptr = + reinterpret_cast<__gm__ float *>(oi_acc_t->buffer.addr) + oi_acc_t->start_offset + data_offset; + + // Output offset: (batch_idx * num_heads + q_tile_idx * QT + row_offset, 0) + int64_t out_offset = (batch_idx * num_heads + q_tile_idx * QT + row_offset) * head_dim; + __gm__ float *dst_ptr = reinterpret_cast<__gm__ float *>(out_t->buffer.addr) + out_t->start_offset + out_offset; + + online_update_spmd( + mij_ptr, lij_ptr, oi_new_ptr, mi_ptr, li_ptr, oi_ptr, dst_ptr, is_first, is_last + ); + } +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + // Safety check: if called with null tensor args (misrouted hub invocation), return. + if (args[0] == 0 || args[1] == 0 || args[2] == 0) { + return; + } + + __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *oi_new_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *mi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *li_acc_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + __gm__ Tensor *oi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[5]); + __gm__ Tensor *out_t = reinterpret_cast<__gm__ Tensor *>(args[6]); + uint64_t is_first = static_cast(args[7]); + uint64_t is_last = static_cast(args[8]); + int64_t num_heads = static_cast(args[9]); + int64_t head_dim = static_cast(args[10]); + int64_t q_loop = static_cast(args[11]); + int64_t q_tile = static_cast(args[12]); + int64_t total_blocks = static_cast(args[13]); + + if (q_tile == 16) { + online_update_entry<16>( + mij_t, lij_t, oi_new_t, mi_acc_t, li_acc_t, oi_acc_t, out_t, is_first, is_last, num_heads, head_dim, q_loop, + total_blocks, args + ); + } else { + online_update_entry<64>( + mij_t, lij_t, oi_new_t, mi_acc_t, li_acc_t, oi_acc_t, out_t, is_first, is_last, num_heads, head_dim, q_loop, + total_blocks, args + ); + } +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_softmax_prepare.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_softmax_prepare.cpp new file mode 100644 index 0000000000..57e32d4e15 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/aiv/aiv_softmax_prepare.cpp @@ -0,0 +1,251 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD Softmax Preparation Kernel (AIV) with partial block masking and +// dual-vector subvector split. +// +// Hardware block_num is fixed at 24. Each hardware block strides over +// total_blocks logical work items: +// for (idx = hw_block_idx; idx < total_blocks; idx += block_num) +// Each logical block_idx encodes (batch_idx, q_tile_idx). +// The two AIV lanes in a cluster split the q_tile rows via +// get_sub_block_id(): AIV0 handles the first half, AIV1 handles the second. +// q_tile is passed as a runtime scalar and dispatched to the matching template. +// +// Hardware safety: TROW* instructions (TROWMAX, TROWEXPANDSUB, TROWSUM) require +// a minimum tile height of 16 rows. When the sub-tile has fewer rows (e.g., 8), +// we allocate 16-row compute tiles and let the pad rows carry UB garbage. +// Row-independence of TROW* operations ensures valid rows are unaffected. +// Only the valid rows' results are stored to pij; scalars (mij/lij) are stored +// at padded width (16 per lane) so online_update can load them back consistently. +// +// Computes (per sub-slice of sub_m rows): +// sij_masked = pad(sij, valid_len, -inf) +// sij_scale = sij_masked * scale +// mij = row_max(sij_scale) -> (sub_m, 1) +// pij = exp(sij_scale - mij) -> (sub_m, N) +// lij = row_sum(pij) -> (sub_m, 1) +// +// Args: +// args[0] = sij Tensor* (total_blocks*q_tile, block_size) float32 [input] +// args[1] = context_lens Tensor* (batch,) int32 +// args[2] = pij Tensor* (total_blocks*q_tile, block_size) bf16 [output] +// args[3] = mij Tensor* (padded scalar buf) float32 [output] +// args[4] = lij Tensor* (padded scalar buf) float32 [output] +// args[5] = scale_value scalar (as float bits in uint64) +// args[6] = bn scalar: current KV block index +// args[7] = block_size scalar +// args[8] = q_loop scalar +// args[9] = q_tile scalar +// args[10] = total_blocks scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +static constexpr int N_128 = 128; // block_size (Case1) +static constexpr int N_64 = 64; // block_size (Case2) +static constexpr int MIN_TM = 16; // minimum tile height for TROW* hw safety + +// TM = actual number of valid rows this AIV lane owns (may be < 16). +// All TROW* instructions operate on padded 16-row tiles to avoid the hardware +// dstRptStride issue at TM=8. We load TM rows from GM into the first TM rows of +// the PM-row UB tile. Pad rows [TM,PM) contain UB garbage, but this is safe: +// all TROW* ops (TROWMAX, TROWEXPANDSUB, TROWSUM) are row-independent, so garbage +// in pad rows cannot affect valid-row results. Scalar outputs (mij/lij) are stored +// at padded width (PM per lane); only valid rows' pij data is stored to GM. +template +static __aicore__ void softmax_prepare_spmd( + __gm__ float *sij_addr, float scale_value, uint64_t valid_len, __gm__ bfloat16_t *pij_addr, __gm__ float *mij_addr, + __gm__ float *lij_addr +) { + constexpr int PM = (TM < MIN_TM) ? MIN_TM : TM; // padded tile height + constexpr int kAlignedRows = ((PM * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + + // GM accessors: load/store only TM valid rows + using GlobalDataMxN = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + using GlobalDataMxN_bf16 = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + using GlobalScalarDN = GlobalTensor, Stride<1, 1, 1, 1, 1>, Layout::DN>; + + GlobalDataMxN sijGlobal(sij_addr); + GlobalDataMxN_bf16 pijGlobal(pij_addr); + GlobalScalarDN mijGlobal(mij_addr); + GlobalScalarDN lijGlobal(lij_addr); + + // Compute tiles are PM(=16) rows, matching the proven hardware path + using TileSijDyn = Tile; + using TileSijPad = + Tile; + + using TileVecMxN = Tile; + using TileVecMxN_bf16 = Tile; + using TileScalarDN = Tile; + + // GM load tiles (TM rows) aliased to the same UB offset as the padded tiles + using TileLoadMxN = Tile; + using TileStoreBf16 = Tile; + + TileVecMxN sijTile; + TileSijDyn sijDynTile(static_cast(valid_len)); + TileSijPad sijPadTile; + TileVecMxN pijTile; + TileVecMxN tmpTile; + TileScalarDN maxTile; + TileScalarDN sumTile; + TileVecMxN_bf16 pijBf16Tile; + + // Load-size tiles aliased to the start of the padded tiles + TileLoadMxN sijLoadTile; + TileStoreBf16 pijStoreTile; + + TASSIGN(sijTile, 0x0); + TASSIGN(sijLoadTile, 0x0); // alias: first TM rows of sijTile + TASSIGN(sijDynTile, 0x0); + TASSIGN(sijPadTile, 0x0); + TASSIGN(pijTile, PM * TN * sizeof(float)); + TASSIGN(tmpTile, 2 * PM * TN * sizeof(float)); + TASSIGN(maxTile, 3 * PM * TN * sizeof(float)); + TASSIGN(sumTile, 3 * PM * TN * sizeof(float) + kAlignedRows * sizeof(float)); + TASSIGN(pijBf16Tile, 3 * PM * TN * sizeof(float) + 2 * kAlignedRows * sizeof(float)); + TASSIGN(pijStoreTile, 3 * PM * TN * sizeof(float) + 2 * kAlignedRows * sizeof(float)); // alias + + // Load only TM valid rows into the first TM rows of the PM-row tile. + // Pad rows [TM, PM) contain UB garbage — this is safe because all TROW* + // ops are row-independent: garbage in pad rows cannot affect valid rows. + TLOAD(sijLoadTile, sijGlobal); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + // Pad invalid columns [valid_len, N) with -inf for all PM rows + TFILLPAD_INPLACE(sijPadTile, sijDynTile); + pipe_barrier(PIPE_V); + + TMULS(sijTile, sijTile, scale_value); + pipe_barrier(PIPE_V); + TROWMAX(maxTile, sijTile, tmpTile); + pipe_barrier(PIPE_V); + TROWEXPANDSUB(pijTile, sijTile, maxTile); + pipe_barrier(PIPE_V); + TEXP(pijTile, pijTile); + TCVT(pijBf16Tile, pijTile, RoundMode::CAST_ROUND); + TCVT(pijTile, pijBf16Tile, RoundMode::CAST_ROUND); + TROWSUM(sumTile, pijTile, tmpTile); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(mijGlobal, maxTile); + TSTORE(lijGlobal, sumTile); + TSTORE(pijGlobal, pijStoreTile); // store only TM valid rows + + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); +} + +template +static __aicore__ void softmax_prepare_entry( + __gm__ Tensor *sij_t, __gm__ Tensor *context_lens_t, __gm__ Tensor *pij_t, __gm__ Tensor *mij_t, + __gm__ Tensor *lij_t, float scale_value, int64_t bn, int64_t block_size, int64_t q_loop, int64_t total_blocks, + __gm__ int64_t *args +) { + constexpr int SUB_M = Q_TILE / 2; + // Padded sub-tile height for hw-safe TROW* ops + constexpr int PAD_SUB_M = (SUB_M < MIN_TM) ? MIN_TM : SUB_M; + + int32_t hw_block_idx = get_block_idx(args); + int32_t block_num = get_block_num(args); + int32_t sub_block_id = get_sub_block_id(args); + + for (int32_t block_idx = hw_block_idx; block_idx < total_blocks; block_idx += block_num) { + int64_t batch_idx = block_idx / q_loop; + + // Compute valid_len for this block: how many columns of sij are valid + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t remaining = cur_seq - bn * block_size; + uint64_t valid_len; + if (remaining <= 0) { + valid_len = 0; + } else if (remaining >= block_size) { + valid_len = static_cast(block_size); + } else { + valid_len = static_cast(remaining); + } + + // Row offset for this AIV lane within the block_idx's q_tile slice + int64_t row_offset = sub_block_id * SUB_M; + + // Pointers into this block's SUB_M-row sub-slice of the flat tensors + int64_t data_row_offset = block_idx * Q_TILE + row_offset; + __gm__ float *sij_addr = + reinterpret_cast<__gm__ float *>(sij_t->buffer.addr) + sij_t->start_offset + data_row_offset * block_size; + __gm__ bfloat16_t *pij_addr = reinterpret_cast<__gm__ bfloat16_t *>(pij_t->buffer.addr) + pij_t->start_offset + + data_row_offset * block_size; + + // Scalar layout uses padded kAlignedRows per sub-tile (based on PAD_SUB_M) + constexpr int kAlignedRowsFull = 2 * (((PAD_SUB_M * sizeof(float) + 31) / 32) * (32 / sizeof(float))); + constexpr int kAlignedRowsSub = ((PAD_SUB_M * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + int64_t scalar_offset = block_idx * kAlignedRowsFull + sub_block_id * kAlignedRowsSub; + __gm__ float *mij_addr = + reinterpret_cast<__gm__ float *>(mij_t->buffer.addr) + mij_t->start_offset + scalar_offset; + __gm__ float *lij_addr = + reinterpret_cast<__gm__ float *>(lij_t->buffer.addr) + lij_t->start_offset + scalar_offset; + + if (valid_len == 0) { + for (int i = 0; i < kAlignedRowsSub; i++) { + mij_addr[i] = -1e30f; + lij_addr[i] = 0.0f; + } + for (int i = 0; i < SUB_M * static_cast(block_size); i++) { + pij_addr[i] = static_cast(0.0f); + } + continue; + } + + softmax_prepare_spmd(sij_addr, scale_value, valid_len, pij_addr, mij_addr, lij_addr); + } +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *sij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *pij_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + float scale_value = from_u64(static_cast(args[5])); + int64_t bn = static_cast(args[6]); + int64_t block_size = static_cast(args[7]); + int64_t q_loop = static_cast(args[8]); + int64_t q_tile = static_cast(args[9]); + int64_t total_blocks = static_cast(args[10]); + + if (q_tile == 16) { + softmax_prepare_entry<16, 128>( + sij_t, context_lens_t, pij_t, mij_t, lij_t, scale_value, bn, block_size, q_loop, total_blocks, args + ); + } else { + softmax_prepare_entry<64, 64>( + sij_t, context_lens_t, pij_t, mij_t, lij_t, scale_value, bn, block_size, q_loop, total_blocks, args + ); + } +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/kernel_config.py b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/kernel_config.py new file mode 100644 index 0000000000..3b434725a9 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/kernel_config.py @@ -0,0 +1,87 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +""" +SPMD Paged Attention Kernel and Orchestration Configuration (fixed block_num=24) + +Uses SPMD parallelism with a fixed hardware block_num of 24. +total_logical_blocks = batch * q_loop work items are distributed across +the 24 hardware blocks via stride loops inside each kernel. +q_tile adapts to num_heads: q_tile = min(num_heads, MAX_Q_TILE). + +Softmax and online-update run as MIX tasks (AIC idle + AIV0 + AIV1), with the +two AIVs splitting the q_tile query rows via get_sub_block_id(). + +AIC Kernels (Matrix Multiplication): + - aic_qk_matmul: Q @ K^T (SPMD across batch*q_loop) + - aic_pv_matmul: P @ V (SPMD across batch*q_loop) + - aic_hub: no-op, occupies the AIC slot of softmax/update MIX tasks + +AIV Kernels (Vector Operations): + - aiv_softmax_prepare: scale, rowmax, exp, rowsum on sub-tile (q_tile/2 rows) + - aiv_online_update: online softmax accumulation + normalization on sub-tile + - aiv_hub: no-op, used to allocate persistent accumulators +""" + +from pathlib import Path + +from simpler.task_interface import ArgDirection as D # pyright: ignore[reportAttributeAccessIssue] + +_KERNELS_ROOT = Path(__file__).parent + +ORCHESTRATION = { + "source": str(_KERNELS_ROOT / "orchestration" / "spmd_paged_attention_orch.cpp"), + "function_name": "aicpu_orchestration_entry", +} + +KERNELS = [ + # AIC kernels (matrix multiplication using Cube unit) + { + "func_id": 0, + "name": "SPMD_QK", + "source": str(_KERNELS_ROOT / "aic" / "aic_qk_matmul.cpp"), + "core_type": "aic", + }, + { + "func_id": 1, + "name": "SPMD_PV", + "source": str(_KERNELS_ROOT / "aic" / "aic_pv_matmul.cpp"), + "core_type": "aic", + }, + { + "func_id": 2, + "name": "AIC_HUB", + "source": str(_KERNELS_ROOT / "aic" / "aic_hub.cpp"), + "core_type": "aic", + }, + # AIV kernels (vector operations) + { + "func_id": 3, + "name": "SPMD_SF", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_softmax_prepare.cpp"), + "core_type": "aiv", + }, + { + "func_id": 4, + "name": "SPMD_UP", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_online_update.cpp"), + "core_type": "aiv", + }, + { + "func_id": 5, + "name": "AIV_HUB", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_hub.cpp"), + "core_type": "aiv", + }, +] + +RUNTIME_CONFIG = { + "runtime": "tensormap_and_ringbuffer", + "aicpu_thread_num": 4, + "block_dim": 24, +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/orchestration/spmd_paged_attention_orch.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/orchestration/spmd_paged_attention_orch.cpp new file mode 100644 index 0000000000..425abf4fcb --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_blk24/kernels/orchestration/spmd_paged_attention_orch.cpp @@ -0,0 +1,268 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +/** + * SPMD Paged Attention Orchestration (dual-vector subvector partitioning) + * + * Uses SPMD parallelism with a fixed hardware block_num = 24. + * total_logical_blocks = batch * q_loop logical work items are distributed + * across the 24 hardware blocks via stride loops inside each kernel: + * for (block_idx = hw_idx; block_idx < total_blocks; block_idx += 24) + * + * q_tile adapts to num_heads: q_tile = min(num_heads, MAX_Q_TILE). + * When num_heads <= MAX_Q_TILE, q_loop = 1 and each block processes all heads. + * + * QK and PV matmuls are AIC-only SPMD tasks. Softmax and online-update are + * submitted as MIX tasks (AIC hub + AIV0 + AIV1) so the two AIV lanes within + * a cluster each process one half of the q_tile query rows, using + * get_sub_block_id() to pick their sub-slice. + * + * Memory Layout: + * Query: (batch, num_heads, head_dim) - bfloat16 + * Key/Value: (total_blocks, block_size, kv_head_num, head_dim) - bfloat16 + * Block Table: (batch, max_num_blocks_per_req) - int32 + * Context Lens: (batch,) - int32 + * Output: (batch, num_heads, head_dim) - float32 + * + * Scratch layout (runtime-allocated, indexed by block_idx * q_tile): + * sij: (spmd_blocks * q_tile, block_size) float32 + * pij: (spmd_blocks * q_tile, block_size) data_type + * oi_new: (spmd_blocks * q_tile, head_dim) float32 + * mij/lij: (spmd_blocks * q_tile,) float32 + * oi_acc/mi_acc/li_acc: persistent accumulators across bn loop + */ + +#include +#include + +#include + +#include "pto_orchestration_api.h" + +#define FUNC_QK_MATMUL 0 +#define FUNC_PV_MATMUL 1 +#define FUNC_AIC_HUB 2 +#define FUNC_SOFTMAX_PREPARE 3 +#define FUNC_ONLINE_UPDATE 4 +#define FUNC_AIV_HUB 5 + +static constexpr uint64_t MAX_Q_TILE = 128; + +extern "C" { + +__attribute__((visibility("default"))) PTO2OrchestrationConfig +aicpu_orchestration_config(const ChipStorageTaskArgs &orch_args) { + (void)orch_args; + return PTO2OrchestrationConfig{ + .expected_arg_count = 7, + }; +} + +__attribute__((visibility("default"))) void aicpu_orchestration_entry(const ChipStorageTaskArgs &orch_args) { + // query: shape=[batch, num_heads, head_dim] + uint64_t batch = orch_args.tensor(0).shapes[0]; + uint64_t num_heads = orch_args.tensor(0).shapes[1]; + uint64_t head_dim = orch_args.tensor(0).shapes[2]; + DataType data_type = orch_args.tensor(0).dtype; + + // key_cache: shape=[total_blocks, block_size, kv_head_num, head_dim] + uint64_t block_size = orch_args.tensor(1).shapes[1]; + + // block_table: shape=[batch, max_num_blocks_per_req] + uint64_t max_num_blocks_per_req = orch_args.tensor(3).shapes[1]; + + // scale from scalar arg + uint64_t scale_value = orch_args.scalar(0); + + // Q_TILE adapts to num_heads: use num_heads directly when it fits, cap at MAX_Q_TILE + uint64_t q_tile = (num_heads <= MAX_Q_TILE) ? num_heads : MAX_Q_TILE; + uint64_t q_loop = (num_heads + q_tile - 1) / q_tile; + int64_t total_logical_blocks = static_cast(batch * q_loop); + static constexpr int16_t spmd_block_num = 24; + + LOG_INFO( + "SPMD PA: batch=%" PRIu64 " heads=%" PRIu64 " hd=%" PRIu64 " bs=%" PRIu64 " q_tile=%" PRIu64 " q_loop=%" PRIu64 + " hw_blocks=%d logical_blocks=%" PRId64, + batch, num_heads, head_dim, block_size, q_tile, q_loop, spmd_block_num, total_logical_blocks + ); + + // Wrap host-provided tensors + void *query_ptr = orch_args.tensor(0).data_as(); + void *kc_ptr = orch_args.tensor(1).data_as(); + void *vc_ptr = orch_args.tensor(2).data_as(); + void *out_ptr = orch_args.tensor(5).data_as(); + + uint64_t total_kv_blocks = orch_args.tensor(1).shapes[0]; + uint64_t kv_total_rows = total_kv_blocks * block_size; + + uint32_t query_shapes[2] = {static_cast(batch * num_heads), static_cast(head_dim)}; + uint32_t kv_shapes[2] = {static_cast(kv_total_rows), static_cast(head_dim)}; + uint32_t out_shapes[2] = {static_cast(batch * num_heads), static_cast(head_dim)}; + + Tensor query = make_tensor_external(query_ptr, query_shapes, 2, data_type); + Tensor key_cache = make_tensor_external(kc_ptr, kv_shapes, 2, data_type); + Tensor value_cache = make_tensor_external(vc_ptr, kv_shapes, 2, data_type); + Tensor out = make_tensor_external(out_ptr, out_shapes, 2, DataType::FLOAT32); + + uint32_t bt_shapes[2] = {static_cast(batch), static_cast(max_num_blocks_per_req)}; + Tensor block_table = + make_tensor_external(orch_args.tensor(3).data_as(), bt_shapes, 2, DataType::INT32, false); + uint32_t cl_shapes[1] = {static_cast(batch)}; + Tensor context_lens = + make_tensor_external(orch_args.tensor(4).data_as(), cl_shapes, 1, DataType::INT32, false); + + // Find max context_len for KV block loop bound + uint64_t max_ctx = 0; + for (uint64_t b = 0; b < batch; b++) { + uint32_t idx[1] = {static_cast(b)}; + uint64_t ctx = static_cast(get_tensor_data(context_lens, 1, idx)); + if (ctx > max_ctx) max_ctx = ctx; + } + uint64_t max_bn = (max_ctx + block_size - 1) / block_size; + + // Scratch tensor create infos (sized for all logical blocks) + uint32_t n_rows = static_cast(total_logical_blocks) * static_cast(q_tile); + uint32_t sij_shapes[2] = {n_rows, static_cast(block_size)}; + uint32_t pij_shapes[2] = {n_rows, static_cast(block_size)}; + uint32_t oi_new_shapes[2] = {n_rows, static_cast(head_dim)}; + + // Scalar buffers (mij, lij, mi_acc, li_acc) are padded to MIN_TM=16 rows per + // AIV lane, even when the actual sub-tile is smaller (e.g., 8 rows for q_tile=16). + // This ensures TROW* instructions always operate on 16-row tiles on hardware. + uint64_t sub_qt = q_tile / 2; + uint64_t pad_sub_qt = (sub_qt < 16) ? 16 : sub_qt; + uint64_t aligned_rows_sub = ((pad_sub_qt * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + uint64_t aligned_rows_full = 2 * aligned_rows_sub; + uint32_t scalar_n = static_cast(total_logical_blocks) * static_cast(aligned_rows_full); + uint32_t scalar_shapes[1] = {scalar_n}; + + TensorCreateInfo sij_ci(sij_shapes, 2, DataType::FLOAT32); + TensorCreateInfo pij_ci(pij_shapes, 2, data_type); + TensorCreateInfo oi_new_ci(oi_new_shapes, 2, DataType::FLOAT32); + TensorCreateInfo mij_ci(scalar_shapes, 1, DataType::FLOAT32); + TensorCreateInfo lij_ci(scalar_shapes, 1, DataType::FLOAT32); + TensorCreateInfo acc_oi_ci(oi_new_shapes, 2, DataType::FLOAT32); + TensorCreateInfo acc_mi_ci(scalar_shapes, 1, DataType::FLOAT32); + TensorCreateInfo acc_li_ci(scalar_shapes, 1, DataType::FLOAT32); + + // Outer scope holds persistent accumulators that live across all bn iterations. + // Each bn iteration runs in a nested inner scope so that scratch tensors + // (sij, pij, oi_new, mij, lij) are freed when the inner scope exits, + // preventing GM heap ring overflow on large configs. + PTO2_SCOPE() { + // Allocate persistent accumulators via no-op AIV hub + Arg hub_args; + hub_args.add_output(acc_oi_ci); + hub_args.add_output(acc_mi_ci); + hub_args.add_output(acc_li_ci); + TaskOutputTensors hub_outs = pto2_rt_submit_aiv_task(FUNC_AIV_HUB, hub_args); + const Tensor &oi_acc = hub_outs.get_ref(0); + const Tensor &mi_acc = hub_outs.get_ref(1); + const Tensor &li_acc = hub_outs.get_ref(2); + + for (uint64_t bn = 0; bn < max_bn; bn++) { + uint64_t is_first = (bn == 0) ? 1 : 0; + uint64_t is_last = (bn == max_bn - 1) ? 1 : 0; + + PTO2_SCOPE() { + // -- QK Matmul (AIC, SPMD) -- + Arg qk_args; + qk_args.add_input(query); + qk_args.add_input(key_cache); + qk_args.add_input(block_table); + qk_args.add_input(context_lens); + qk_args.add_output(sij_ci); + qk_args.add_scalar(static_cast(bn)); + qk_args.add_scalar(static_cast(num_heads)); + qk_args.add_scalar(static_cast(head_dim)); + qk_args.add_scalar(static_cast(block_size)); + qk_args.add_scalar(static_cast(max_num_blocks_per_req)); + qk_args.add_scalar(static_cast(q_loop)); + qk_args.add_scalar(static_cast(q_tile)); + qk_args.add_scalar(total_logical_blocks); + qk_args.launch_spec.set_block_num(spmd_block_num); + TaskOutputTensors qk_outs = pto2_rt_submit_aic_task(FUNC_QK_MATMUL, qk_args); + const Tensor &sij = qk_outs.get_ref(0); + + // -- Softmax Prepare (MIX: AIC hub + AIV0 + AIV1, SPMD) -- + Arg sf_args; + sf_args.add_input(sij); + sf_args.add_input(context_lens); + sf_args.add_output(pij_ci); + sf_args.add_output(mij_ci); + sf_args.add_output(lij_ci); + sf_args.add_scalar(scale_value); + sf_args.add_scalar(static_cast(bn)); + sf_args.add_scalar(static_cast(block_size)); + sf_args.add_scalar(static_cast(q_loop)); + sf_args.add_scalar(static_cast(q_tile)); + sf_args.add_scalar(total_logical_blocks); + sf_args.launch_spec.set_block_num(spmd_block_num); + MixedKernels sf_mk; + sf_mk.aic_kernel_id = FUNC_AIC_HUB; + sf_mk.aiv0_kernel_id = FUNC_SOFTMAX_PREPARE; + sf_mk.aiv1_kernel_id = FUNC_SOFTMAX_PREPARE; + TaskOutputTensors sf_outs = pto2_rt_submit_task(sf_mk, sf_args); + const Tensor &pij = sf_outs.get_ref(0); + const Tensor &mij = sf_outs.get_ref(1); + const Tensor &lij = sf_outs.get_ref(2); + + // -- PV Matmul (AIC, SPMD) -- + Arg pv_args; + pv_args.add_input(pij); + pv_args.add_input(value_cache); + pv_args.add_input(block_table); + pv_args.add_input(context_lens); + pv_args.add_output(oi_new_ci); + pv_args.add_scalar(static_cast(bn)); + pv_args.add_scalar(static_cast(num_heads)); + pv_args.add_scalar(static_cast(head_dim)); + pv_args.add_scalar(static_cast(block_size)); + pv_args.add_scalar(static_cast(max_num_blocks_per_req)); + pv_args.add_scalar(static_cast(q_loop)); + pv_args.add_scalar(static_cast(q_tile)); + pv_args.add_scalar(total_logical_blocks); + pv_args.launch_spec.set_block_num(spmd_block_num); + TaskOutputTensors pv_outs = pto2_rt_submit_aic_task(FUNC_PV_MATMUL, pv_args); + const Tensor &oi_new = pv_outs.get_ref(0); + + // -- Online Update (MIX: AIC hub + AIV0 + AIV1, SPMD) -- + Arg up_args; + up_args.add_input(mij); + up_args.add_input(lij); + up_args.add_input(oi_new); + up_args.add_inout(mi_acc); + up_args.add_inout(li_acc); + up_args.add_inout(oi_acc); + up_args.add_inout(out); + up_args.add_scalar(is_first); + up_args.add_scalar(is_last); + up_args.add_scalar(static_cast(num_heads)); + up_args.add_scalar(static_cast(head_dim)); + up_args.add_scalar(static_cast(q_loop)); + up_args.add_scalar(static_cast(q_tile)); + up_args.add_scalar(total_logical_blocks); + up_args.launch_spec.set_block_num(spmd_block_num); + MixedKernels up_mk; + up_mk.aic_kernel_id = FUNC_AIC_HUB; + up_mk.aiv0_kernel_id = FUNC_ONLINE_UPDATE; + up_mk.aiv1_kernel_id = FUNC_ONLINE_UPDATE; + pto2_rt_submit_task(up_mk, up_args); + } + } + } + + LOG_INFO( + "SPMD PA: %" PRIu64 " KV iters x 4 tasks, hw_blocks=%d logical=%" PRId64, max_bn, + static_cast(spmd_block_num), total_logical_blocks + ); +} + +} // extern "C" diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/golden.py b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/golden.py new file mode 100644 index 0000000000..bab2a4affe --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/golden.py @@ -0,0 +1,68 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""SPMD Paged Attention Golden with TPUSH/TPOP (production scale, bfloat16). + +Combined AIC+AIV MixedKernels task per SPMD block. +AIC and AIV communicate via TPUSH/TPOP pipes for sij, pij, and oi_new. +Uses per-block online softmax (FlashAttention style). +""" + +from paged_attention_golden import ( + compute_golden, # noqa: F401 + run_golden_test, +) +from paged_attention_golden import generate_inputs as _generate_inputs + +__outputs__ = ["out"] + +RTOL = 2e-3 +ATOL = 2e-3 + +ALL_CASES = { + "Case0": { + "batch": 1, + "num_heads": 16, + "kv_head_num": 1, + "head_dim": 128, + "block_size": 128, + "context_len": 128, + "max_model_len": 1024, + "dtype": "bfloat16", + }, + "Case1": { + "batch": 256, + "num_heads": 16, + "kv_head_num": 1, + "head_dim": 128, + "block_size": 128, + "context_len": 8192, + "max_model_len": 32768, + "dtype": "bfloat16", + }, + "Case2": { + "batch": 64, + "num_heads": 64, + "kv_head_num": 1, + "head_dim": 128, + "block_size": 64, + "context_len": 8192, + "max_model_len": 32768, + "dtype": "bfloat16", + }, +} + +DEFAULT_CASE = "Case0" + + +def generate_inputs(params: dict) -> list: + return _generate_inputs(params) + + +if __name__ == "__main__": + run_golden_test(ALL_CASES, DEFAULT_CASE, generate_inputs, label="SPMD Paged Attention") diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/kernel_config.py b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/kernel_config.py new file mode 100644 index 0000000000..1893818da4 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/kernel_config.py @@ -0,0 +1,52 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +""" +SPMD Paged Attention with TPUSH/TPOP (Combined AIC+AIV MixedKernels) + +Single MixedKernels task per invocation. AIC handles QK/PV matmul, +AIV handles online softmax and update. Data flows via TPUSH/TPOP pipes: + - sij pipe (C2V): QK scores + - pij pipe (V2C): softmax probabilities + - oi pipe (C2V): PV output + +The same source is compiled twice: once for AIC (__DAV_CUBE__) and +once for AIV (__DAV_VEC__), using if constexpr dispatch. +""" + +from pathlib import Path + +from simpler.task_interface import ArgDirection as D # pyright: ignore[reportAttributeAccessIssue] + +_KERNELS_ROOT = Path(__file__).parent + +ORCHESTRATION = { + "source": str(_KERNELS_ROOT / "orchestration" / "spmd_paged_attention_orch.cpp"), + "function_name": "aicpu_orchestration_entry", +} + +KERNELS = [ + { + "func_id": 0, + "name": "PA_AIC", + "source": str(_KERNELS_ROOT / "mix" / "paged_attention_parallel.cpp"), + "core_type": "aic", + }, + { + "func_id": 1, + "name": "PA_AIV", + "source": str(_KERNELS_ROOT / "mix" / "paged_attention_parallel.cpp"), + "core_type": "aiv", + }, +] + +RUNTIME_CONFIG = { + "runtime": "tensormap_and_ringbuffer", + "aicpu_thread_num": 4, + "block_dim": 24, +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/mix/paged_attention_parallel.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/mix/paged_attention_parallel.cpp new file mode 100644 index 0000000000..9980718706 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/mix/paged_attention_parallel.cpp @@ -0,0 +1,579 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +/** + * Paged Attention Parallel Kernel — Combined AIC + AIV (TPUSH/TPOP) + * + * Single source compiled twice: + * - AIC (cube): __DAV_CUBE__ → QK matmul, PV matmul + * - AIV (vector): __DAV_VEC__ → online softmax, online update + * + * Per-block pipeline (all KV blocks processed in one invocation): + * AIC: QK matmul → TPUSH(sij) → TPOP(pij) → PV matmul → TPUSH(oi_new) + * AIV: TPOP(sij) → online softmax → TPUSH(pij) → TPOP(oi_new) → online update + * + * Three TPUSH/TPOP pipes: + * - sij_pipe (C2V): scores (Q_TILE, block_size) fp32, TILE_UP_DOWN split + * - pij_pipe (V2C): probabilities (Q_TILE, block_size) bf16, TILE_UP_DOWN + * - oi_pipe (C2V): PV output (Q_TILE, head_dim) fp32, TILE_UP_DOWN split + * + * Q_TILE=16, SUB_QT=8. Supports block_size=64|128, head_dim=128. + * + * MixedKernels args: + * args[0] = query Tensor* (batch*num_heads, head_dim) bf16 + * args[1] = key_cache Tensor* (kv_total_rows, head_dim) bf16 + * args[2] = value_cache Tensor* (kv_total_rows, head_dim) bf16 + * args[3] = block_table Tensor* (batch, max_blocks_per_req) int32 + * args[4] = context_lens Tensor* (batch,) int32 + * args[5] = out Tensor* (batch*num_heads, head_dim) float32 [output] + * args[6] = sij_fifo Tensor* GM ring buffer for sij pipe + * args[7] = pij_fifo Tensor* GM ring buffer for pij pipe + * args[8] = oi_fifo Tensor* GM ring buffer for oi_new pipe + * args[9] = scale_value scalar (float bits in uint64) + * args[10] = num_heads scalar + * args[11] = head_dim scalar + * args[12] = block_size scalar + * args[13] = max_num_blocks_per_req scalar + * args[14] = q_loop scalar + */ + +#include +// NOLINTBEGIN(clang-diagnostic-error,bugprone-reserved-identifier,bugprone-easily-swappable-parameters,modernize-use-auto) +#include +#include + +#include "tensor.h" + +using pto::BLayout; +using pto::Direction; +using pto::GlobalTensor; +using pto::Layout; +using pto::PadValue; +using pto::RoundMode; +using pto::Shape; +using pto::SLayout; +using pto::Stride; +using pto::Tile; +using pto::TileAcc; +using pto::TileLeft; +using pto::TileRight; +using pto::TileSplitAxis; +using pto::TileType; +using pto::TPipe; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] // NOLINT(whitespace/braces) +#endif + +#include "intrinsic.h" + +#ifdef __DAV_CUBE__ +constexpr bool DAV_CUBE = true; +#else +constexpr bool DAV_CUBE = false; +#endif + +#ifdef __DAV_VEC__ +constexpr bool DAV_VEC = true; +#else +constexpr bool DAV_VEC = false; +#endif + +static constexpr int Q_TILE = 16; +static constexpr int SUB_QT = Q_TILE / 2; +static constexpr int HEAD_DIM = 128; + +// TPUSH/TPOP pipe flag IDs (each consumes 2 consecutive IDs: data + backpressure) +static constexpr uint16_t SIJ_FLAG_ID = 0; +static constexpr uint16_t PIJ_FLAG_ID = 2; +static constexpr uint16_t OI_FLAG_ID = 4; +static constexpr uint8_t FIFO_DEPTH = 1; + +// GM FIFO slot sizes (max case: block_size=128) +static constexpr uint32_t SIJ_SLOT_SIZE = Q_TILE * 128 * sizeof(float); // 8192 +static constexpr uint32_t PIJ_SLOT_SIZE = Q_TILE * 128 * sizeof(bfloat16_t); // 4096 +static constexpr uint32_t OI_SLOT_SIZE = Q_TILE * HEAD_DIM * sizeof(float); // 8192 + +// Pipe types +using SijPipeT = TPipe; +using PijPipeT = TPipe; +using OiPipeT = TPipe; + +// AIV UB consumer buffer layout (fixed, independent of template params) +// LocalSlotNum=2 by default, so each consumer needs 2 slots +// sij consumer: 2 * SUB_QT * 128 * 4 = 8192 bytes (max, block_size=128) +// oi consumer: 2 * SUB_QT * HEAD_DIM * 4 = 8192 bytes +static constexpr uint32_t SIJ_UB_BASE = 0x0; +static constexpr uint32_t SIJ_UB_SIZE = 2 * SUB_QT * 128 * sizeof(float); // 8192 +static constexpr uint32_t OI_UB_BASE = SIJ_UB_BASE + SIJ_UB_SIZE; // 0x2000 +static constexpr uint32_t OI_UB_SIZE = 2 * SUB_QT * HEAD_DIM * sizeof(float); // 8192 +static constexpr uint32_t WORK_UB_BASE = OI_UB_BASE + OI_UB_SIZE; // 0x4000 + +// AIC L1 consumer buffer for V2C pij pipe +// LocalSlotNum=2, pij consumer: 2 * Q_TILE * 128 * 2 = 8192 bytes +static constexpr uint32_t PIJ_L1_BASE = 0x40000; +static constexpr uint32_t PIJ_L1_SIZE = 2 * Q_TILE * 128 * sizeof(bfloat16_t); // 8192 + +// ============================================================================ +// AIC (Cube) side +// ============================================================================ + +template +static __aicore__ void aic_process_blocks( + __gm__ bfloat16_t *qi_base, __gm__ bfloat16_t *key_base, __gm__ bfloat16_t *val_base, __gm__ int32_t *bt, + uint64_t bt_offset, uint64_t n_blocks, SijPipeT &sij_pipe, PijPipeT &pij_pipe, OiPipeT &oi_pipe +) { + // QK tile types + using GlobalA_QK = GlobalTensor, Stride>; + using GlobalB_QK = GlobalTensor, Stride, Layout::DN>; + using TileMatA_QK = Tile; + using TileMatB_QK = Tile; + using LeftTile_QK = TileLeft; + using RightTile_QK = TileRight; + using AccTile_QK = TileAcc; + + // PV tile types + using GlobalB_PV = GlobalTensor, Stride>; + using TileMatB_PV = Tile; + using PijMatTile = Tile; + using LeftTile_PV = TileLeft; + using RightTile_PV = TileRight; + using AccTile_PV = TileAcc; + + // L1 layout for QK: + // 0x00000: qi (M*K*2) + // 0x20000: kj double-buffer (2 * K*N*2) + constexpr int kQKBBytes = K * N * static_cast(sizeof(bfloat16_t)); + constexpr int kPVBBytes = N * K * static_cast(sizeof(bfloat16_t)); + + TileMatA_QK aMatTile_QK; + TileMatB_QK bMatTile_QK_A, bMatTile_QK_B; + TASSIGN(aMatTile_QK, 0x0); + TASSIGN(bMatTile_QK_A, 0x20000); + TASSIGN(bMatTile_QK_B, 0x20000 + kQKBBytes); + + LeftTile_QK aTile_QK; + RightTile_QK bTile_QK; + AccTile_QK cTile_QK; + TASSIGN(aTile_QK, 0x0); + TASSIGN(bTile_QK, 0x0); + TASSIGN(cTile_QK, 0x0); + + // L1 layout for PV: + // PIJ_L1_BASE: pij from V2C TPOP (auto-assigned, 2 slots of M*N*2) + // PIJ_L1_BASE + PIJ_L1_SIZE: vj double-buffer (2 * N*K*2) + PijMatTile pijMatTile; + TileMatB_PV bMatTile_PV_A, bMatTile_PV_B; + TASSIGN(bMatTile_PV_A, PIJ_L1_BASE + PIJ_L1_SIZE); + TASSIGN(bMatTile_PV_B, PIJ_L1_BASE + PIJ_L1_SIZE + kPVBBytes); + + LeftTile_PV aTile_PV; + RightTile_PV bTile_PV; + AccTile_PV cTile_PV; + TASSIGN(aTile_PV, 0x0); + TASSIGN(bTile_PV, 0x0); + TASSIGN(cTile_PV, 0x0); + + // Hoist qi TLOAD + GlobalA_QK qiGlobal(qi_base); + TLOAD(aMatTile_QK, qiGlobal); + + for (uint64_t i = 0; i < n_blocks; i++) { + // ---- QK Matmul ---- + GlobalB_QK kjGlobal(key_base + bt[bt_offset + i] * N * K); + if (i % 2 == 0) { + TLOAD(bMatTile_QK_A, kjGlobal); + } else { + TLOAD(bMatTile_QK_B, kjGlobal); + } + + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + + TMOV(aTile_QK, aMatTile_QK); + if (i % 2 == 0) { + TMOV(bTile_QK, bMatTile_QK_A); + } else { + TMOV(bTile_QK, bMatTile_QK_B); + } + + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + + TMATMUL(cTile_QK, aTile_QK, bTile_QK); + + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + + // TPUSH sij (C2V): AccTile L0C -> GM -> AIV UB + TPUSH(sij_pipe, cTile_QK); + + // TPOP pij (V2C): AIV UB -> GM -> L1 (auto-assigned from PIJ_L1_BASE) + TPOP(pij_pipe, pijMatTile); + + // ---- PV Matmul ---- + GlobalB_PV vjGlobal(val_base + bt[bt_offset + i] * N * K); + if (i % 2 == 0) { + TLOAD(bMatTile_PV_A, vjGlobal); + } else { + TLOAD(bMatTile_PV_B, vjGlobal); + } + + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + + // pijMatTile address set by TPOP + TMOV(aTile_PV, pijMatTile); + if (i % 2 == 0) { + TMOV(bTile_PV, bMatTile_PV_A); + } else { + TMOV(bTile_PV, bMatTile_PV_B); + } + + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + + TMATMUL(cTile_PV, aTile_PV, bTile_PV); + + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + + // TPUSH oi_new (C2V): AccTile L0C -> GM -> AIV UB + TPUSH(oi_pipe, cTile_PV); + + if (i + 1 < n_blocks) { + pipe_barrier(PIPE_ALL); + } + } + + set_flag(PIPE_FIX, PIPE_S, EVENT_ID7); + wait_flag(PIPE_FIX, PIPE_S, EVENT_ID7); +} + +// ============================================================================ +// AIV (Vector) side +// ============================================================================ + +template +static __aicore__ void aiv_process_blocks( + float scale_value, uint64_t n_blocks, uint64_t valid_len_last, __gm__ float *dst_ptr, SijPipeT &sij_pipe, + PijPipeT &pij_pipe, OiPipeT &oi_pipe +) { + constexpr int HD = HEAD_DIM; + constexpr int kAlignedRows = ((TM * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + constexpr int kScalarCols = 32 / sizeof(float); + constexpr int kScalarRows = TM / kScalarCols; + + using SijVecTile = Tile; + using PijVecBf16Tile = Tile; + using OiVecTile = Tile; + + using TileVecMxN = Tile; + using TileSijDyn = Tile; + using TileSijPad = + Tile; + // DN (ColMajor) for row-broadcast ops: TROWMAX, TROWSUM, TROWEXPAND* + using TileScalarDN = Tile; + // ND (RowMajor) for element-wise ops: TMULS, TMUL, TSUB, TADD, TEXP, TMAX + using TileScalarND = + Tile; + // Row for max operations + using TileScalarRow = Tile; + using TileDataMxHD = Tile; + using GlobalDataMxHD = GlobalTensor, Stride<1, 1, 1, HD, 1>>; + + constexpr int kSijBytes = TM * TN * sizeof(float); + constexpr int kPijBf16Bytes = TM * TN * sizeof(bfloat16_t); + constexpr int kScalarDNBytes = kAlignedRows * sizeof(float); + constexpr int kScalarNDBytes = kScalarRows * kScalarCols * sizeof(float); + + // Working tiles (after FIFO consumer buffers) + SijVecTile sijTile; + TileSijPad sijPadTile; + TileVecMxN pijTile; + TileVecMxN tmpTile; + PijVecBf16Tile pijBf16Tile; + // DN tiles for row-broadcast + TileScalarDN localMaxDN, globalMaxDN; + TileScalarDN alphaDN_dn, llDN, glDN; + // ND tiles for element-wise + TileScalarND gmND, glND, alphaND, llND, dmND, miNewND, mijND; + // Row tiles for max + TileScalarRow localMaxRow, globalMaxRow, gmRow; + OiVecTile oiNewTile; + TileDataMxHD goTile; + + int ub = WORK_UB_BASE; + TASSIGN(pijTile, ub); + ub += kSijBytes; + TASSIGN(pijBf16Tile, ub); + ub += kPijBf16Bytes; + TASSIGN(tmpTile, ub); + ub += kSijBytes; + + int sb = ub; + TASSIGN(localMaxDN, sb); + TASSIGN(localMaxRow, sb); + sb += kScalarDNBytes; + TASSIGN(globalMaxDN, sb); + TASSIGN(globalMaxRow, sb); + sb += kScalarDNBytes; + // gmND/gmRow alias at the same address + TASSIGN(gmND, sb); + TASSIGN(gmRow, sb); + sb += kScalarDNBytes; + // glND/glDN alias + TASSIGN(glND, sb); + TASSIGN(glDN, sb); + sb += kScalarDNBytes; + // alphaND/alphaDN_dn alias + TASSIGN(alphaND, sb); + TASSIGN(alphaDN_dn, sb); + sb += kScalarDNBytes; + // llND/llDN alias + TASSIGN(llND, sb); + TASSIGN(llDN, sb); + sb += kScalarDNBytes; + TASSIGN(dmND, sb); + sb += kScalarNDBytes; + TASSIGN(miNewND, sb); + sb += kScalarNDBytes; + TASSIGN(mijND, sb); + sb += kScalarNDBytes; + + TASSIGN(goTile, sb); + + GlobalDataMxHD dstGlobal(dst_ptr); + + for (uint64_t i = 0; i < n_blocks; i++) { + // ---- TPOP sij from AIC ---- + TPOP(sij_pipe, sijTile); + + // Pad last block if partial + if (i == n_blocks - 1 && valid_len_last < static_cast(TN)) { + int sij_addr = SIJ_UB_BASE + static_cast((i % 2) * TM * TN * sizeof(float)); + TASSIGN(sijPadTile, sij_addr); + TileSijDyn sijDynTile(static_cast(valid_len_last)); + TASSIGN(sijDynTile, sij_addr); + TFILLPAD_INPLACE(sijPadTile, sijDynTile); + pipe_barrier(PIPE_V); + } + + // rowmax (produces DN tile) + TROWMAX(localMaxDN, sijTile, tmpTile); + pipe_barrier(PIPE_V); + // Convert DN -> Row for element-wise max + TRESHAPE(localMaxRow, localMaxDN); + + if (i == 0) { + // First block: initialize globalMax = scale * localMax + TMULS(globalMaxRow, localMaxRow, scale_value); + } else { + // Update: globalMax = max(gm_prev, scale * localMax) + TMULS(localMaxRow, localMaxRow, scale_value); + pipe_barrier(PIPE_V); + TMAX(globalMaxRow, gmRow, localMaxRow); + } + pipe_barrier(PIPE_V); + // Convert Row -> DN for row-broadcast + TRESHAPE(globalMaxDN, globalMaxRow); + + // softmax: exp(sij * scale - globalMax) + TMULS(sijTile, sijTile, scale_value); + pipe_barrier(PIPE_V); + TROWEXPANDSUB(pijTile, sijTile, globalMaxDN); + pipe_barrier(PIPE_V); + TEXP(pijTile, pijTile); + pipe_barrier(PIPE_V); + + // fp32 -> bf16 -> fp32 for PV matmul precision matching + TCVT(pijBf16Tile, pijTile, RoundMode::CAST_ROUND); + pipe_barrier(PIPE_V); + TCVT(pijTile, pijBf16Tile, RoundMode::CAST_ROUND); + pipe_barrier(PIPE_V); + + // rowsum -> DN tile, then reshape to ND for element-wise ops + TROWSUM(llDN, pijTile, tmpTile); + pipe_barrier(PIPE_V); + TRESHAPE(llND, llDN); + pipe_barrier(PIPE_V); + + // ---- TPUSH pij to AIC ---- + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TPUSH(pij_pipe, pijBf16Tile); + + // ---- TPOP oi_new from AIC ---- + TPOP(oi_pipe, oiNewTile); + + // ---- Online update ---- + if (i == 0) { + // First block: go = oi_new, gm = globalMax, gl = ll + TMULS(goTile, oiNewTile, 1.0f); + pipe_barrier(PIPE_V); + // gmRow = globalMaxRow (copy) + TMULS(gmRow, globalMaxRow, 1.0f); + pipe_barrier(PIPE_V); + // glND = llND (copy using ND tiles) + TMULS(glND, llND, 1.0f); + } else { + // Compute alpha = exp(gm_prev - gm_new) using ND tiles + // globalMaxRow contains the new max, gmRow contains the old max + // Convert both to ND for element-wise ops + TRESHAPE(mijND, globalMaxRow); + TRESHAPE(dmND, gmRow); + pipe_barrier(PIPE_V); + TSUB(alphaND, dmND, mijND); + pipe_barrier(PIPE_V); + TEXP(alphaND, alphaND); + pipe_barrier(PIPE_V); + + // go = go * alpha + oi_new (need alphaDN for TROWEXPANDMUL) + TRESHAPE(alphaDN_dn, alphaND); + pipe_barrier(PIPE_V); + TROWEXPANDMUL(goTile, goTile, alphaDN_dn); + pipe_barrier(PIPE_V); + TADD(goTile, goTile, oiNewTile); + pipe_barrier(PIPE_V); + + // gl = gl * alpha + ll (element-wise with ND tiles) + TMUL(glND, glND, alphaND); + pipe_barrier(PIPE_V); + TADD(glND, glND, llND); + pipe_barrier(PIPE_V); + + // Update gm = globalMax + TMULS(gmRow, globalMaxRow, 1.0f); + } + + pipe_barrier(PIPE_V); + } + + // Normalize: go = go / gl (need glDN for TROWEXPANDDIV) + TRESHAPE(glDN, glND); + pipe_barrier(PIPE_V); + TROWEXPANDDIV(goTile, goTile, glDN); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(dstGlobal, goTile); + + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); +} + +// ============================================================================ +// Entry point +// ============================================================================ + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *query_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *key_cache_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *value_cache_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *block_table_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + __gm__ Tensor *out_t = reinterpret_cast<__gm__ Tensor *>(args[5]); + __gm__ Tensor *sij_fifo_t = reinterpret_cast<__gm__ Tensor *>(args[6]); + __gm__ Tensor *pij_fifo_t = reinterpret_cast<__gm__ Tensor *>(args[7]); + __gm__ Tensor *oi_fifo_t = reinterpret_cast<__gm__ Tensor *>(args[8]); + + float scale_value = from_u64(static_cast(args[9])); + int64_t num_heads = static_cast(args[10]); + int64_t head_dim = static_cast(args[11]); + int64_t block_size = static_cast(args[12]); + int64_t max_blocks_per_req = static_cast(args[13]); + int64_t q_loop = static_cast(args[14]); + + int32_t block_idx = get_block_idx(args); + int64_t batch_idx = block_idx / q_loop; + int64_t q_tile_idx = block_idx % q_loop; + + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t n_blocks = (cur_seq + block_size - 1) / block_size; + + uint64_t valid_len_last = 0; + if (n_blocks > 0) { + int64_t last_block_seq = (n_blocks - 1) * block_size; + int64_t remaining = cur_seq - last_block_seq; + valid_len_last = (remaining >= block_size) ? static_cast(block_size) : + (remaining > 0 ? static_cast(remaining) : 0); + } + + // GM FIFO buffer per SPMD block + __gm__ void *sij_fifo_base = reinterpret_cast<__gm__ void *>( + reinterpret_cast<__gm__ uint8_t *>(sij_fifo_t->buffer.addr) + block_idx * SIJ_SLOT_SIZE * FIFO_DEPTH + ); + __gm__ void *pij_fifo_base = reinterpret_cast<__gm__ void *>( + reinterpret_cast<__gm__ uint8_t *>(pij_fifo_t->buffer.addr) + block_idx * PIJ_SLOT_SIZE * FIFO_DEPTH + ); + __gm__ void *oi_fifo_base = reinterpret_cast<__gm__ void *>( + reinterpret_cast<__gm__ uint8_t *>(oi_fifo_t->buffer.addr) + block_idx * OI_SLOT_SIZE * FIFO_DEPTH + ); + + SijPipeT sij_pipe(sij_fifo_base, SIJ_UB_BASE, 0U); + PijPipeT pij_pipe(pij_fifo_base, 0U, PIJ_L1_BASE); + OiPipeT oi_pipe(oi_fifo_base, OI_UB_BASE, 0U); + + if constexpr (DAV_CUBE) { + if (n_blocks <= 0) return; + + int64_t q_offset = (batch_idx * num_heads + q_tile_idx * Q_TILE) * head_dim; + __gm__ bfloat16_t *qi_base = + reinterpret_cast<__gm__ bfloat16_t *>(query_t->buffer.addr) + query_t->start_offset + q_offset; + __gm__ bfloat16_t *key_base = + reinterpret_cast<__gm__ bfloat16_t *>(key_cache_t->buffer.addr) + key_cache_t->start_offset; + __gm__ bfloat16_t *val_base = + reinterpret_cast<__gm__ bfloat16_t *>(value_cache_t->buffer.addr) + value_cache_t->start_offset; + __gm__ int32_t *bt = + reinterpret_cast<__gm__ int32_t *>(block_table_t->buffer.addr) + block_table_t->start_offset; + uint64_t bt_offset = static_cast(batch_idx * max_blocks_per_req); + + if (block_size == 128) { + aic_process_blocks( + qi_base, key_base, val_base, bt, bt_offset, static_cast(n_blocks), sij_pipe, pij_pipe, oi_pipe + ); + } else { + aic_process_blocks( + qi_base, key_base, val_base, bt, bt_offset, static_cast(n_blocks), sij_pipe, pij_pipe, oi_pipe + ); + } + } + + if constexpr (DAV_VEC) { + int32_t sub_block_id = get_sub_block_id(args); + int64_t row_offset = sub_block_id * SUB_QT; + + int64_t out_offset = (batch_idx * num_heads + q_tile_idx * Q_TILE + row_offset) * head_dim; + __gm__ float *dst = reinterpret_cast<__gm__ float *>(out_t->buffer.addr) + out_t->start_offset + out_offset; + + if (n_blocks <= 0) { + for (int64_t j = 0; j < SUB_QT * head_dim; j++) { + dst[j] = 0.0f; + } + return; + } + + if (block_size == 128) { + aiv_process_blocks( + scale_value, static_cast(n_blocks), valid_len_last, dst, sij_pipe, pij_pipe, oi_pipe + ); + } else { + aiv_process_blocks( + scale_value, static_cast(n_blocks), valid_len_last, dst, sij_pipe, pij_pipe, oi_pipe + ); + } + } +} +// NOLINTEND(clang-diagnostic-error,bugprone-reserved-identifier,bugprone-easily-swappable-parameters,modernize-use-auto) diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/orchestration/spmd_paged_attention_orch.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/orchestration/spmd_paged_attention_orch.cpp new file mode 100644 index 0000000000..a7dc700060 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll-tpush-pop/kernels/orchestration/spmd_paged_attention_orch.cpp @@ -0,0 +1,142 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +/** + * SPMD Paged Attention Orchestration with TPUSH/TPOP + * + * Submits a single MixedKernels task that processes ALL KV blocks per + * (batch, q_tile) position. AIC and AIV cooperate via TPUSH/TPOP pipes: + * AIC: QK matmul → TPUSH(sij) → TPOP(pij) → PV matmul → TPUSH(oi_new) + * AIV: TPOP(sij) → online softmax → TPUSH(pij) → TPOP(oi_new) → online update + * + * SPMD parallelism: block_num = batch * q_loop. + * Each SPMD block handles one (batch_idx, q_tile_idx). + * + * GM FIFO buffers for TPUSH/TPOP are allocated as scratch tensors. + * Each SPMD block gets its own FIFO slots. + */ + +#include +#include + +#include + +#include "pto_orchestration_api.h" + +#define FUNC_PA_AIC 0 +#define FUNC_PA_AIV 1 + +static constexpr uint64_t Q_TILE = 16; +static constexpr uint64_t HEAD_DIM = 128; + +// GM FIFO slot sizes (must match kernel's constants) +static constexpr uint32_t SIJ_SLOT_SIZE = Q_TILE * 128 * sizeof(float); // 8192 +static constexpr uint32_t PIJ_SLOT_SIZE = Q_TILE * 128 * sizeof(uint16_t); // 4096 (bf16) +static constexpr uint32_t OI_SLOT_SIZE = Q_TILE * HEAD_DIM * sizeof(float); // 8192 +static constexpr uint32_t FIFO_DEPTH = 1; + +extern "C" { + +__attribute__((visibility("default"))) PTO2OrchestrationConfig +aicpu_orchestration_config(const ChipStorageTaskArgs &orch_args) { + (void)orch_args; + return PTO2OrchestrationConfig{ + .expected_arg_count = 7, + }; +} + +__attribute__((visibility("default"))) void aicpu_orchestration_entry(const ChipStorageTaskArgs &orch_args) { + uint64_t batch = orch_args.tensor(0).shapes[0]; + uint64_t num_heads = orch_args.tensor(0).shapes[1]; + uint64_t head_dim = orch_args.tensor(0).shapes[2]; + DataType data_type = orch_args.tensor(0).dtype; + + uint64_t block_size = orch_args.tensor(1).shapes[1]; + uint64_t max_num_blocks_per_req = orch_args.tensor(3).shapes[1]; + uint64_t scale_value = orch_args.scalar(0); + + uint64_t q_loop = (num_heads + Q_TILE - 1) / Q_TILE; + int16_t spmd_block_num = static_cast(batch * q_loop); + + LOG_INFO( + "SPMD PA TPUSH/TPOP: batch=%" PRIu64 " heads=%" PRIu64 " hd=%" PRIu64 " bs=%" PRIu64 " q_loop=%" PRIu64 + " blocks=%d", + batch, num_heads, head_dim, block_size, q_loop, spmd_block_num + ); + + // Wrap host tensors + void *query_ptr = orch_args.tensor(0).data_as(); + void *kc_ptr = orch_args.tensor(1).data_as(); + void *vc_ptr = orch_args.tensor(2).data_as(); + void *out_ptr = orch_args.tensor(5).data_as(); + + uint64_t total_kv_blocks = orch_args.tensor(1).shapes[0]; + uint64_t kv_total_rows = total_kv_blocks * block_size; + + uint32_t query_shapes[2] = {static_cast(batch * num_heads), static_cast(head_dim)}; + uint32_t kv_shapes[2] = {static_cast(kv_total_rows), static_cast(head_dim)}; + uint32_t out_shapes[2] = {static_cast(batch * num_heads), static_cast(head_dim)}; + + Tensor query = make_tensor_external(query_ptr, query_shapes, 2, data_type); + Tensor key_cache = make_tensor_external(kc_ptr, kv_shapes, 2, data_type); + Tensor value_cache = make_tensor_external(vc_ptr, kv_shapes, 2, data_type); + Tensor out = make_tensor_external(out_ptr, out_shapes, 2, DataType::FLOAT32); + + uint32_t bt_shapes[2] = {static_cast(batch), static_cast(max_num_blocks_per_req)}; + Tensor block_table = + make_tensor_external(orch_args.tensor(3).data_as(), bt_shapes, 2, DataType::INT32, false); + uint32_t cl_shapes[1] = {static_cast(batch)}; + Tensor context_lens = + make_tensor_external(orch_args.tensor(4).data_as(), cl_shapes, 1, DataType::INT32, false); + + // GM FIFO buffers for TPUSH/TPOP (one set of slots per SPMD block) + uint32_t sij_fifo_total = static_cast(spmd_block_num) * SIJ_SLOT_SIZE * FIFO_DEPTH; + uint32_t pij_fifo_total = static_cast(spmd_block_num) * PIJ_SLOT_SIZE * FIFO_DEPTH; + uint32_t oi_fifo_total = static_cast(spmd_block_num) * OI_SLOT_SIZE * FIFO_DEPTH; + + // Allocate as 1D byte tensors (using INT32 for 4-byte alignment, divide by 4) + uint32_t sij_fifo_shapes[1] = {sij_fifo_total / sizeof(int32_t)}; + uint32_t pij_fifo_shapes[1] = {pij_fifo_total / sizeof(int32_t)}; + uint32_t oi_fifo_shapes[1] = {oi_fifo_total / sizeof(int32_t)}; + + TensorCreateInfo sij_fifo_ci(sij_fifo_shapes, 1, DataType::INT32); + TensorCreateInfo pij_fifo_ci(pij_fifo_shapes, 1, DataType::INT32); + TensorCreateInfo oi_fifo_ci(oi_fifo_shapes, 1, DataType::INT32); + + PTO2_SCOPE() { + Arg args; + args.add_input(query); + args.add_input(key_cache); + args.add_input(value_cache); + args.add_input(block_table); + args.add_input(context_lens); + args.add_inout(out); + args.add_output(sij_fifo_ci); + args.add_output(pij_fifo_ci); + args.add_output(oi_fifo_ci); + args.add_scalar(scale_value); + args.add_scalar(static_cast(num_heads)); + args.add_scalar(static_cast(head_dim)); + args.add_scalar(static_cast(block_size)); + args.add_scalar(static_cast(max_num_blocks_per_req)); + args.add_scalar(static_cast(q_loop)); + args.launch_spec.set_block_num(spmd_block_num); + + MixedKernels mk; + mk.aic_kernel_id = FUNC_PA_AIC; + mk.aiv0_kernel_id = FUNC_PA_AIV; + mk.aiv1_kernel_id = FUNC_PA_AIV; + pto2_rt_submit_task(mk, args); + } + + LOG_INFO("SPMD PA TPUSH/TPOP: submitted 1 MixedKernels task, blocks=%d", static_cast(spmd_block_num)); +} + +} // extern "C" diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/golden.py b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/golden.py new file mode 100644 index 0000000000..6c60ad04c7 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/golden.py @@ -0,0 +1,67 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""SPMD Paged Attention Golden - tensormap_and_ringbuffer test (production scale, bfloat16). + +Uses SPMD parallelism: each block handles one (batch, q_tile) position. +Kernels use get_block_idx() to determine their work slice. +""" + +from paged_attention_golden import ( + compute_golden, # noqa: F401 + run_golden_test, +) +from paged_attention_golden import generate_inputs as _generate_inputs + +__outputs__ = ["out"] + +RTOL = 1e-3 +ATOL = 1e-3 + +ALL_CASES = { + "Case1": { + "batch": 256, + "num_heads": 16, + "kv_head_num": 1, + "head_dim": 128, + "block_size": 128, + "context_len": 8192, + "max_model_len": 32768, + "dtype": "bfloat16", + }, + "Case2": { + "batch": 64, + "num_heads": 64, + "kv_head_num": 1, + "head_dim": 128, + "block_size": 64, + "context_len": 8192, + "max_model_len": 32768, + "dtype": "bfloat16", + }, + "Case3": { + "batch": 64, + "num_heads": 64, + "kv_head_num": 1, + "head_dim": 256, + "block_size": 64, + "context_len": 8192, + "max_model_len": 32768, + "dtype": "bfloat16", + }, +} + +DEFAULT_CASE = "Case1" + + +def generate_inputs(params: dict) -> list: + return _generate_inputs(params) + + +if __name__ == "__main__": + run_golden_test(ALL_CASES, DEFAULT_CASE, generate_inputs, label="SPMD Paged Attention") diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_hub.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_hub.cpp new file mode 100644 index 0000000000..eb602e0bd6 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_hub.cpp @@ -0,0 +1,29 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// AIC Hub Kernel - No-op stub used as the AIC slot of MIX (AIC+AIV0+AIV1) tasks +// when the real work happens only on the two AIVs (softmax, online update). +// Pairing an idle AIC with two active AIVs forces the scheduler to allocate a +// full cluster, which is what enables the two AIV lanes to run in parallel. + +#include +#include + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) {} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_pv_matmul.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_pv_matmul.cpp new file mode 100644 index 0000000000..bf219cdddc --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_pv_matmul.cpp @@ -0,0 +1,214 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD SplitK PV Matmul: Accumulated P @ V across n_blocks +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// Each SPMD block processes n_blocks using SplitK accumulation: +// Block 0: TMATMUL(C, A, B) — initialize accumulator +// Block i: TMATMUL_ACC(C, C, A, B) — accumulate into same C +// +// Per-block pij: contiguous packed (M, K) tiles in pij_buf +// Per-block vj: value_cache base + block_table lookup +// Single output: oi_new (M, N) fp32 = sum of P_i @ V_i across all blocks +// +// Case1: (16, 128) @ (128, 128) -> (16, 128) +// Case2: (64, 64) @ ( 64, 128) -> (64, 128) +// +// Args: +// args[0] = pij Tensor* (spmd_blocks*Q_TILE, n_blocks*block_size) bf16 +// args[1] = value_cache Tensor* (kv_total_rows, head_dim) bf16 +// args[2] = block_table Tensor* (batch, max_blocks_per_req) int32 +// args[3] = context_lens Tensor* (batch,) int32 +// args[4] = oi_new Tensor* (spmd_blocks*Q_TILE, head_dim) float32 [output] +// args[5] = bn_start scalar: starting KV block index +// args[6] = n_blocks scalar: number of KV blocks to process +// args[7] = num_heads scalar +// args[8] = head_dim scalar +// args[9] = block_size scalar +// args[10] = max_num_blocks_per_req scalar +// args[11] = q_loop scalar + +#include +// NOLINTBEGIN(clang-diagnostic-error,bugprone-reserved-identifier,bugprone-easily-swappable-parameters,modernize-avoid-c-arrays,modernize-use-auto) +#include + +#include "tensor.h" + +// NOLINTNEXTLINE(build/namespaces) +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] // NOLINT(whitespace/braces) +#endif + +#include "intrinsic.h" + +template +static __aicore__ void pv_matmul_n_spmd( + __gm__ bfloat16_t *pij_base, __gm__ bfloat16_t *val_base, __gm__ float *oi_base, uint64_t n_blocks, + __gm__ int32_t *bt, uint64_t bt_offset +) { + using GlobalA = GlobalTensor, Stride>; + using GlobalB = GlobalTensor, Stride>; + using GlobalOut = GlobalTensor, Stride>; + + using TileMatA = Tile; + using TileMatB = Tile; + + using LeftTile = TileLeft; + using RightTile = TileRight; + using AccTile = TileAcc; + + // L1 memory layout: double-buffered A and B tiles + constexpr int kATileBytes = M * K * static_cast(sizeof(bfloat16_t)); + constexpr int kBTileBytes = K * N * static_cast(sizeof(bfloat16_t)); + + TileMatA aMatTile[2]; + TileMatB bMatTile[2]; + TASSIGN(aMatTile[0], 0x0); + TASSIGN(aMatTile[1], kATileBytes); + TASSIGN(bMatTile[0], 2 * kATileBytes); + TASSIGN(bMatTile[1], 2 * kATileBytes + kBTileBytes); + + // L0 memory layout: double-buffered L0A and L0B, single accumulator L0C + LeftTile aTile[2]; + RightTile bTile[2]; + AccTile cTile; + TASSIGN(aTile[0], 0x0); + TASSIGN(aTile[1], kATileBytes); + TASSIGN(bTile[0], 0x0); + TASSIGN(bTile[1], kBTileBytes); + TASSIGN(cTile, 0x0); + + GlobalOut oiGlobal(oi_base); + + // Seed reverse-dependency flags + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + + for (uint64_t i = 0; i < n_blocks; i++) { + int cur = static_cast(i % 2); + GlobalA pijGlobal(pij_base + i * M * K); + GlobalB vjGlobal(val_base + bt[bt_offset + i] * K * N); + + // Stage 1: TLOAD (MTE2: GM -> L1[cur]) + wait_flag(PIPE_MTE1, PIPE_MTE2, (event_t)cur); + TLOAD(aMatTile[cur], pijGlobal); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + TLOAD(bMatTile[cur], vjGlobal); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + + // Stage 2: TMOV (MTE1: L1[cur] -> L0[cur]) + wait_flag(PIPE_M, PIPE_MTE1, (event_t)cur); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + TMOV(aTile[cur], aMatTile[cur]); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + TMOV(bTile[cur], bMatTile[cur]); + set_flag(PIPE_MTE1, PIPE_MTE2, (event_t)cur); + + // Stage 3: TMATMUL (M-pipe: L0A[cur] x L0B[cur] -> L0C) + set_flag(PIPE_MTE1, PIPE_M, (event_t)cur); + wait_flag(PIPE_MTE1, PIPE_M, (event_t)cur); + if (i == 0) { + TMATMUL(cTile, aTile[cur], bTile[cur]); + } else { + TMATMUL_ACC(cTile, cTile, aTile[cur], bTile[cur]); + } + set_flag(PIPE_M, PIPE_MTE1, (event_t)cur); + } + + // Drain outstanding reverse-dependency flags + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + TSTORE(oiGlobal, cTile); + + set_flag(PIPE_FIX, PIPE_S, EVENT_ID7); + wait_flag(PIPE_FIX, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *pij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *value_cache_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *block_table_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *oi_new_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + + int64_t bn_start = static_cast(args[5]); + int64_t n_blocks_total = static_cast(args[6]); + int64_t num_heads = static_cast(args[7]); + int64_t head_dim = static_cast(args[8]); + int64_t block_size = static_cast(args[9]); + int64_t max_blocks_per_req = static_cast(args[10]); + int64_t q_loop = static_cast(args[11]); + + int32_t block_idx = get_block_idx(args); + int64_t batch_idx = block_idx / q_loop; + + int64_t q_tile = (block_size == 128) ? 16 : 64; + + // Check how many KV blocks this batch actually has + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t bn_this_batch = (cur_seq + block_size - 1) / block_size; + + // Clamp n_blocks to valid range for this batch + int64_t valid_blocks = bn_this_batch - bn_start; + if (valid_blocks < 0) valid_blocks = 0; + int64_t n_blocks = (valid_blocks < n_blocks_total) ? valid_blocks : n_blocks_total; + + // Output pointer for this SPMD block's oi_new slice + __gm__ float *oi_base = + reinterpret_cast<__gm__ float *>(oi_new_t->buffer.addr) + oi_new_t->start_offset + block_idx * q_tile * head_dim; + + if (n_blocks <= 0) { + for (int64_t i = 0; i < q_tile * head_dim; i++) { + oi_base[i] = 0.0f; + } + return; + } + + // pij packed tiles: block_idx's region in pij tensor + __gm__ bfloat16_t *pij_base = + reinterpret_cast<__gm__ bfloat16_t *>(pij_t->buffer.addr) + pij_t->start_offset + + block_idx * q_tile * n_blocks_total * block_size; + + // Value cache base + __gm__ bfloat16_t *val_base = + reinterpret_cast<__gm__ bfloat16_t *>(value_cache_t->buffer.addr) + value_cache_t->start_offset; + + // Block table pointer + __gm__ int32_t *bt = + reinterpret_cast<__gm__ int32_t *>(block_table_t->buffer.addr) + block_table_t->start_offset; + uint64_t bt_offset = static_cast(batch_idx * max_blocks_per_req + bn_start); + + if (q_tile == 16) { + pv_matmul_n_spmd<16, 128, 128>( + pij_base, val_base, oi_base, static_cast(n_blocks), bt, bt_offset + ); + } else { + pv_matmul_n_spmd<64, 64, 128>( + pij_base, val_base, oi_base, static_cast(n_blocks), bt, bt_offset + ); + } +} +// NOLINTEND(clang-diagnostic-error,bugprone-reserved-identifier,bugprone-easily-swappable-parameters,modernize-avoid-c-arrays,modernize-use-auto) diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_qk_matmul.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_qk_matmul.cpp new file mode 100644 index 0000000000..ee5edf61df --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aic/aic_qk_matmul.cpp @@ -0,0 +1,216 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD Multi-block QK Matmul: qi(M, K) @ kj.T(K, N) -> sij(M, N) for n_blocks +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// Each SPMD block processes n_blocks consecutive KV blocks starting at bn_start, +// using double-buffered L1 B tiles and hoisted qi TLOAD. +// +// Output: packed tiles [tile_0, tile_1, ..., tile_{n-1}] each (M, N) in row-major. +// +// Template: M=q_tile, K=head_dim, N=block_size +// Case1: (16, 128) @ (128, 128).T -> (16, 128) +// Case2: (64, 128) @ (128, 64).T -> (64, 64) +// +// Args: +// args[0] = query Tensor* (batch*num_heads, head_dim) bf16 +// args[1] = key_cache Tensor* (kv_total_rows, head_dim) bf16 +// args[2] = block_table Tensor* (batch, max_blocks_per_req) int32 +// args[3] = context_lens Tensor* (batch,) int32 +// args[4] = sij Tensor* (spmd_blocks*Q_TILE, n_blocks*block_size) float32 [output] +// args[5] = bn_start scalar: starting KV block index +// args[6] = n_blocks scalar: number of KV blocks to process +// args[7] = num_heads scalar +// args[8] = head_dim scalar +// args[9] = block_size scalar +// args[10] = max_num_blocks_per_req scalar +// args[11] = q_loop scalar + +#include +// NOLINTBEGIN(clang-diagnostic-error,bugprone-reserved-identifier,bugprone-easily-swappable-parameters,modernize-use-auto) +#include + +#include "tensor.h" + +// NOLINTNEXTLINE(build/namespaces) +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] // NOLINT(whitespace/braces) +#endif + +#include "intrinsic.h" + +template +static __aicore__ void qk_matmul_n_spmd( + __gm__ bfloat16_t *qi_base, __gm__ bfloat16_t *key_base, __gm__ float *sij_base, uint64_t n_blocks, + __gm__ int32_t *bt, uint64_t bt_offset +) { + using GlobalA = GlobalTensor, Stride>; + using GlobalB = GlobalTensor, Stride, Layout::DN>; + using GlobalOut = GlobalTensor, Stride>; + + using TileMatA = Tile; + using TileMatB = Tile; + + using LeftTile = TileLeft; + using RightTile = TileRight; + using AccTile = TileAcc; + + // Double-buffered L1 B tiles for kj prefetching + constexpr int kBBytes = K * N * static_cast(sizeof(bfloat16_t)); + TileMatA aMatTile; + TileMatB bMatTile_A; + TileMatB bMatTile_B; + TASSIGN(aMatTile, 0x0); + TASSIGN(bMatTile_A, 0x20000); + TASSIGN(bMatTile_B, 0x20000 + kBBytes); + + LeftTile aTile; + RightTile bTile; + AccTile cTile; + TASSIGN(aTile, 0x0); + TASSIGN(bTile, 0x0); + TASSIGN(cTile, 0x0); + + // Hoist qi TLOAD before the loop (qi is constant across all blocks) + GlobalA qiGlobal(qi_base); + TLOAD(aMatTile, qiGlobal); + + // Pre-load first kj into buffer A + GlobalB kjGlobal_0(key_base + bt[bt_offset + 0] * N * K); + TLOAD(bMatTile_A, kjGlobal_0); + + for (uint64_t i = 0; i < n_blocks; i++) { + GlobalOut sijGlobal(sij_base + i * M * N); + + // Wait for current kj TLOAD to complete + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + + // TMOV qi L1->L0A and kj L1->L0B from current buffer + TMOV(aTile, aMatTile); + if (i % 2 == 0) { + TMOV(bTile, bMatTile_A); + } else { + TMOV(bTile, bMatTile_B); + } + + // Prefetch next kj into alternate L1 buffer + if (i + 1 < n_blocks) { + GlobalB kjGlobal_next(key_base + bt[bt_offset + i + 1] * N * K); + if (i % 2 == 0) { + TLOAD(bMatTile_B, kjGlobal_next); + } else { + TLOAD(bMatTile_A, kjGlobal_next); + } + } + + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + + TMATMUL(cTile, aTile, bTile); + + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + + TSTORE(sijGlobal, cTile); + + if (i + 1 < n_blocks) { + pipe_barrier(PIPE_ALL); + } + } + set_flag(PIPE_FIX, PIPE_S, EVENT_ID7); + wait_flag(PIPE_FIX, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *query_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *key_cache_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *block_table_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *sij_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + + int64_t bn_start = static_cast(args[5]); + int64_t n_blocks_total = static_cast(args[6]); + int64_t num_heads = static_cast(args[7]); + int64_t head_dim = static_cast(args[8]); + int64_t block_size = static_cast(args[9]); + int64_t max_blocks_per_req = static_cast(args[10]); + int64_t q_loop = static_cast(args[11]); + + int32_t block_idx = get_block_idx(args); + + // Decode (batch_idx, q_tile_idx) from block_idx + int64_t batch_idx = block_idx / q_loop; + int64_t q_tile_idx = block_idx % q_loop; + + int64_t q_tile = (block_size == 128) ? 16 : 64; + + // Check how many KV blocks this batch actually has + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t bn_this_batch = (cur_seq + block_size - 1) / block_size; + + // Clamp n_blocks to valid range for this batch + int64_t valid_blocks = bn_this_batch - bn_start; + if (valid_blocks < 0) valid_blocks = 0; + int64_t n_blocks = (valid_blocks < n_blocks_total) ? valid_blocks : n_blocks_total; + + // sij packed tile output: block_idx's region starts at block_idx * q_tile * n_blocks_total * block_size + __gm__ float *sij_base = reinterpret_cast<__gm__ float *>(sij_t->buffer.addr) + sij_t->start_offset + + block_idx * q_tile * n_blocks_total * block_size; + + if (n_blocks <= 0) { + for (int64_t i = 0; i < q_tile * n_blocks_total * block_size; i++) { + sij_base[i] = 0.0f; + } + return; + } + + // Zero out trailing invalid tiles + if (n_blocks < n_blocks_total) { + __gm__ float *trail = sij_base + n_blocks * q_tile * block_size; + for (int64_t i = 0; i < (n_blocks_total - n_blocks) * q_tile * block_size; i++) { + trail[i] = 0.0f; + } + } + + // Query offset: (batch_idx * num_heads + q_tile_idx * q_tile, 0) + int64_t q_offset = (batch_idx * num_heads + q_tile_idx * q_tile) * head_dim; + __gm__ bfloat16_t *qi_base = + reinterpret_cast<__gm__ bfloat16_t *>(query_t->buffer.addr) + query_t->start_offset + q_offset; + + // Key cache base + __gm__ bfloat16_t *key_base = + reinterpret_cast<__gm__ bfloat16_t *>(key_cache_t->buffer.addr) + key_cache_t->start_offset; + + // Block table pointer + __gm__ int32_t *bt = + reinterpret_cast<__gm__ int32_t *>(block_table_t->buffer.addr) + block_table_t->start_offset; + uint64_t bt_offset = static_cast(batch_idx * max_blocks_per_req + bn_start); + + if (q_tile == 16) { + qk_matmul_n_spmd<16, 128, 128>( + qi_base, key_base, sij_base, static_cast(n_blocks), bt, bt_offset + ); + } else { + qk_matmul_n_spmd<64, 128, 64>( + qi_base, key_base, sij_base, static_cast(n_blocks), bt, bt_offset + ); + } +} +// NOLINTEND(clang-diagnostic-error,bugprone-reserved-identifier,bugprone-easily-swappable-parameters,modernize-use-auto) diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_hub.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_hub.cpp new file mode 100644 index 0000000000..a42f2790e9 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_hub.cpp @@ -0,0 +1,27 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// AIV Hub Kernel - No-op stub for accumulator tensor allocation. +// The runtime allocates output tensors specified in the Arg; the kernel itself does nothing. + +#include +#include + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) {} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_online_update.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_online_update.cpp new file mode 100644 index 0000000000..72069b4c89 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_online_update.cpp @@ -0,0 +1,257 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD Online Softmax Update + Normalize Kernel (AIV) with dual-vector +// subvector split. +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// The two AIV lanes in a cluster split the Q_TILE=16 rows 8/8 via +// get_sub_block_id(): AIV0 updates rows [0, 8), AIV1 updates rows [8, 16). +// The online softmax update is row-independent, so the two lanes never touch +// the same row of mi/li/oi accumulators or the output buffer. +// +// Scalar layout strategy (same as MPMD version): +// M scalar floats stored contiguously in GM can be loaded as either: +// - ND (kScalarRows, kScalarCols) RowMajor for element-wise ops +// - DN (kAlignedRows, 1) ColMajor for row-broadcast ops (TROWEXPANDMUL/DIV) +// Conversion between layouts uses GM round-trip: ND TSTORE -> DN TLOAD. +// +// Args: +// args[0] = mij Tensor* (spmd_blocks*Q_TILE,) float32 +// args[1] = lij Tensor* (spmd_blocks*Q_TILE,) float32 +// args[2] = oi_new Tensor* (spmd_blocks*Q_TILE, head_dim) float32 +// args[3] = mi_acc Tensor* (spmd_blocks*Q_TILE,) float32 [inout] +// args[4] = li_acc Tensor* (spmd_blocks*Q_TILE,) float32 [inout] +// args[5] = oi_acc Tensor* (spmd_blocks*Q_TILE, head_dim) float32 [inout] +// args[6] = out Tensor* (batch*num_heads, head_dim) float32 [inout] +// args[7] = is_first scalar +// args[8] = is_last scalar +// args[9] = num_heads scalar +// args[10] = head_dim scalar +// args[11] = q_loop scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +static constexpr int QT = 16; // Full Q tile rows (shared between both AIVs) +static constexpr int SUB_QT = 8; // Rows per AIV lane (QT / 2) + +template +static __aicore__ void online_update_spmd( + __gm__ float *mij_ptr, __gm__ float *lij_ptr, __gm__ float *oi_new_ptr, __gm__ float *mi_ptr, + __gm__ float *li_ptr, __gm__ float *oi_ptr, __gm__ float *dst_ptr, uint64_t is_first, uint64_t is_last +) { + constexpr int kScalarCols = 32 / sizeof(float); + constexpr int kScalarRows = TM / kScalarCols; + constexpr int kAlignedRows = ((TM * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + + using GlobalDataMxN = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + using GlobalScalarND = + GlobalTensor, Stride<1, 1, 1, kScalarCols, 1>>; + using GlobalScalarDN = GlobalTensor, Stride<1, 1, 1, 1, 1>, Layout::DN>; + + GlobalDataMxN oiNewGlobal(oi_new_ptr); + GlobalDataMxN oiGlobal(oi_ptr); + GlobalDataMxN dstGlobal(dst_ptr); + + GlobalScalarND mijGlobalND(mij_ptr); + GlobalScalarND lijGlobalND(lij_ptr); + GlobalScalarND miGlobalND(mi_ptr); + GlobalScalarND liGlobalND(li_ptr); + + GlobalScalarDN mijGlobalDN(mij_ptr); + GlobalScalarDN lijGlobalDN(lij_ptr); + GlobalScalarDN liGlobalDN(li_ptr); + + using TileDataMxN = Tile; + using TileScalarND = + Tile; + using TileScalarDN = Tile; + + constexpr int kDataBytes = TM * TN * sizeof(float); + constexpr int kScalarNDBytes = kScalarRows * kScalarCols * sizeof(float); + constexpr int kScalarDNBytes = kAlignedRows * sizeof(float); + + TileDataMxN oiNewTile; + TileDataMxN oiTile; + TileScalarND mijND, lijND, miND, liND; + TileScalarND miNewND, alphaND, betaND, tmpND; + TileScalarDN alphaDN, betaDN, liDN; + + TASSIGN(oiNewTile, 0); + TASSIGN(oiTile, kDataBytes); + TASSIGN(mijND, 2 * kDataBytes); + TASSIGN(lijND, 2 * kDataBytes + kScalarNDBytes); + TASSIGN(miND, 2 * kDataBytes + 2 * kScalarNDBytes); + TASSIGN(liND, 2 * kDataBytes + 3 * kScalarNDBytes); + TASSIGN(miNewND, 2 * kDataBytes + 4 * kScalarNDBytes); + TASSIGN(alphaND, 2 * kDataBytes + 5 * kScalarNDBytes); + TASSIGN(betaND, 2 * kDataBytes + 6 * kScalarNDBytes); + TASSIGN(tmpND, 2 * kDataBytes + 7 * kScalarNDBytes); + TASSIGN(alphaDN, 2 * kDataBytes + 8 * kScalarNDBytes); + TASSIGN(betaDN, 2 * kDataBytes + 8 * kScalarNDBytes + kScalarDNBytes); + TASSIGN(liDN, 2 * kDataBytes + 8 * kScalarNDBytes + 2 * kScalarDNBytes); + + if (is_first) { + TLOAD(oiNewTile, oiNewGlobal); + TLOAD(mijND, mijGlobalND); + TLOAD(lijND, lijGlobalND); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(miGlobalND, mijND); + TSTORE(liGlobalND, lijND); + TSTORE(oiGlobal, oiNewTile); + + if (is_last) { + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + TLOAD(liDN, liGlobalDN); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + TROWEXPANDDIV(oiNewTile, oiNewTile, liDN); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(dstGlobal, oiNewTile); + } + } else { + TLOAD(oiNewTile, oiNewGlobal); + TLOAD(oiTile, oiGlobal); + TLOAD(mijND, mijGlobalND); + TLOAD(lijND, lijGlobalND); + TLOAD(miND, miGlobalND); + TLOAD(liND, liGlobalND); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + TMAX(miNewND, miND, mijND); + pipe_barrier(PIPE_V); + TSUB(alphaND, miND, miNewND); + pipe_barrier(PIPE_V); + TEXP(alphaND, alphaND); + pipe_barrier(PIPE_V); + TSUB(betaND, mijND, miNewND); + pipe_barrier(PIPE_V); + TEXP(betaND, betaND); + pipe_barrier(PIPE_V); + TMUL(liND, alphaND, liND); + pipe_barrier(PIPE_V); + TMUL(tmpND, betaND, lijND); + pipe_barrier(PIPE_V); + TADD(liND, liND, tmpND); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(miGlobalND, miNewND); + TSTORE(liGlobalND, liND); + TSTORE(mijGlobalND, alphaND); + TSTORE(lijGlobalND, betaND); + + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + TLOAD(alphaDN, mijGlobalDN); + TLOAD(betaDN, lijGlobalDN); + if (is_last) { + TLOAD(liDN, liGlobalDN); + } + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + + TROWEXPANDMUL(oiTile, oiTile, alphaDN); + TROWEXPANDMUL(oiNewTile, oiNewTile, betaDN); + pipe_barrier(PIPE_V); + TADD(oiTile, oiTile, oiNewTile); + + if (is_last) { + pipe_barrier(PIPE_V); + TROWEXPANDDIV(oiTile, oiTile, liDN); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(dstGlobal, oiTile); + } else { + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(oiGlobal, oiTile); + } + } + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + // Safety check: if called with null tensor args (misrouted hub invocation), return. + if (args[0] == 0 || args[1] == 0 || args[2] == 0) { + return; + } + + __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *oi_new_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *mi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *li_acc_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + __gm__ Tensor *oi_acc_t = reinterpret_cast<__gm__ Tensor *>(args[5]); + __gm__ Tensor *out_t = reinterpret_cast<__gm__ Tensor *>(args[6]); + uint64_t is_first = static_cast(args[7]); + uint64_t is_last = static_cast(args[8]); + int64_t num_heads = static_cast(args[9]); + int64_t head_dim = static_cast(args[10]); + int64_t q_loop = static_cast(args[11]); + + int32_t block_idx = get_block_idx(args); + int32_t sub_block_id = get_sub_block_id(args); // 0 = AIV0 (rows 0..7), 1 = AIV1 (rows 8..15) + int64_t batch_idx = block_idx / q_loop; + int64_t q_tile_idx = block_idx % q_loop; + + // Scalar layout: full QT=16 rows pack to kAlignedRowsFull=16 floats per block_idx; + // each AIV lane owns kAlignedRowsSub=8 contiguous floats inside that slab. + constexpr int kAlignedRowsFull = ((QT * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + constexpr int kAlignedRowsSub = ((SUB_QT * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + + int64_t row_offset = sub_block_id * SUB_QT; + + // Accumulator offsets (each AIV lane owns its own 8-row sub-slice within the block_idx slab) + int64_t scalar_offset = block_idx * kAlignedRowsFull + sub_block_id * kAlignedRowsSub; + int64_t data_offset = (block_idx * QT + row_offset) * head_dim; + + __gm__ float *mij_ptr = + reinterpret_cast<__gm__ float *>(mij_t->buffer.addr) + mij_t->start_offset + scalar_offset; + __gm__ float *lij_ptr = + reinterpret_cast<__gm__ float *>(lij_t->buffer.addr) + lij_t->start_offset + scalar_offset; + __gm__ float *oi_new_ptr = + reinterpret_cast<__gm__ float *>(oi_new_t->buffer.addr) + oi_new_t->start_offset + data_offset; + __gm__ float *mi_ptr = + reinterpret_cast<__gm__ float *>(mi_acc_t->buffer.addr) + mi_acc_t->start_offset + scalar_offset; + __gm__ float *li_ptr = + reinterpret_cast<__gm__ float *>(li_acc_t->buffer.addr) + li_acc_t->start_offset + scalar_offset; + __gm__ float *oi_ptr = + reinterpret_cast<__gm__ float *>(oi_acc_t->buffer.addr) + oi_acc_t->start_offset + data_offset; + + // Output offset: (batch_idx * num_heads + q_tile_idx * QT + row_offset, 0) + int64_t out_offset = (batch_idx * num_heads + q_tile_idx * QT + row_offset) * head_dim; + __gm__ float *dst_ptr = reinterpret_cast<__gm__ float *>(out_t->buffer.addr) + out_t->start_offset + out_offset; + + online_update_spmd(mij_ptr, lij_ptr, oi_new_ptr, mi_ptr, li_ptr, oi_ptr, dst_ptr, is_first, is_last); +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_softmax_prepare.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_softmax_prepare.cpp new file mode 100644 index 0000000000..4c9ccdfbb4 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/aiv/aiv_softmax_prepare.cpp @@ -0,0 +1,344 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// SPMD Two-Pass Softmax Kernel (AIV) for n_blocks tiles with dual-vector split +// +// SPMD block_idx encodes (batch_idx, q_tile_idx). +// The two AIV lanes in a cluster split the Q_TILE rows via get_sub_block_id(): +// AIV0 (sub_block_id=0) handles rows [0, SUB_M) +// AIV1 (sub_block_id=1) handles rows [SUB_M, Q_TILE) +// +// Memory layout: QK kernel writes packed (Q_TILE, block_size) tiles contiguously. +// Within each tile, AIV0 processes the first SUB_M rows and AIV1 the rest. +// Tile stride between consecutive tiles is Q_TILE * block_size elements. +// +// Two-pass softmax (same algorithm as paged_attention_unroll): +// Pass 1: Find global m = scale * max over all blocks of rowmax(S_i) +// Pass 2: Compute P_i = exp(S_i * scale - m) -> bf16, accumulate l = rowsum(P_i) +// +// Case1: SUB_M=8, TN=128 (Q_TILE=16, block_size=128) +// Case2: SUB_M=32, TN=64 (Q_TILE=64, block_size=64) +// +// Args: +// args[0] = sij Tensor* (spmd_blocks*Q_TILE, n_blocks*block_size) float32 +// args[1] = context_lens Tensor* (batch,) int32 +// args[2] = pij Tensor* (spmd_blocks*Q_TILE, n_blocks*block_size) bf16 [output] +// args[3] = mij Tensor* (spmd_blocks*Q_TILE,) float32 [output] +// args[4] = lij Tensor* (spmd_blocks*Q_TILE,) float32 [output] +// args[5] = scale_value scalar (as float bits in uint64) +// args[6] = bn_start scalar: starting KV block index +// args[7] = n_blocks scalar: number of KV blocks in this group +// args[8] = block_size scalar +// args[9] = q_loop scalar + +#include +#include + +#include "tensor.h" + +using namespace pto; + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#define __aicore__ [aicore] +#endif + +#include "intrinsic.h" + +template +static __aicore__ void softmax_prepare_n_spmd( + __gm__ float *sij_base, float scale_value, __gm__ bfloat16_t *pij_base, __gm__ float *mij_addr, + __gm__ float *lij_addr, uint64_t n_blocks, uint64_t valid_len_last +) { + constexpr int kAlignedRows = ((TM * sizeof(float) + 31) / 32) * (32 / sizeof(float)); + constexpr int kScalarCols = 32 / sizeof(float); + constexpr int kScalarRows = TM / kScalarCols; + + // --- GlobalTensor types --- + using GlobalDataMxN = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + using GlobalDataMxN_bf16 = GlobalTensor, Stride<1, 1, 1, TN, 1>>; + using GlobalScalarDN = GlobalTensor, Stride<1, 1, 1, 1, 1>, Layout::DN>; + using GlobalScalarND = + GlobalTensor, Stride<1, 1, 1, kScalarCols, 1>>; + + // --- Tile types --- + using TileSijDyn = Tile; + using TileSijPad = + Tile; + using TileVecMxN = Tile; + using TileVecMxN_bf16 = Tile; + using TileScalarDN = Tile; + using TileScalarND = + Tile; + using TileScalarRow = Tile; + + // --- UB memory layout (double-buffered sij) --- + constexpr int kDataBytes = TM * TN * sizeof(float); + constexpr int kScalarDNBytes = kAlignedRows * sizeof(float); + + TileVecMxN sijTile_A; + TileSijPad sijPadTile_A; + TileVecMxN sijTile_B; + TileSijPad sijPadTile_B; + TileVecMxN pijTile; + TileVecMxN tmpTile; + TileVecMxN sumAccTile; + TileScalarDN localMaxDN; + TileScalarDN globalMaxDN; + TileScalarDN sumDN; + TileVecMxN_bf16 pijBf16Tile; + + TileScalarRow localMaxRow; + TileScalarRow globalMaxRow; + TileScalarND globalMaxND; + + TASSIGN(sijTile_A, 0x0); + TASSIGN(sijPadTile_A, 0x0); + TASSIGN(sijTile_B, kDataBytes); + TASSIGN(sijPadTile_B, kDataBytes); + TASSIGN(pijTile, 2 * kDataBytes); + TASSIGN(tmpTile, 3 * kDataBytes); + TASSIGN(sumAccTile, 4 * kDataBytes); + int scalarBase = 5 * kDataBytes; + TASSIGN(localMaxDN, scalarBase); + TASSIGN(localMaxRow, scalarBase); + TASSIGN(globalMaxDN, scalarBase + kScalarDNBytes); + TASSIGN(globalMaxRow, scalarBase + kScalarDNBytes); + TASSIGN(globalMaxND, scalarBase + kScalarDNBytes); + TASSIGN(sumDN, scalarBase + 2 * kScalarDNBytes); + TASSIGN(pijBf16Tile, scalarBase + 3 * kScalarDNBytes); + + GlobalScalarND mijGlobalND(mij_addr); + GlobalScalarDN lijGlobalDN(lij_addr); + + // ======== Pass 1: Find global row max (unscaled) ======== + // Tile stride between consecutive packed tiles = TILE_STRIDE elements + GlobalDataMxN sijGlobal_p1_0(sij_base); + TLOAD(sijTile_A, sijGlobal_p1_0); + + for (uint64_t i = 0; i < n_blocks; i++) { + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + if (i == n_blocks - 1 && valid_len_last < static_cast(TN)) { + TileSijDyn sijDynTile(static_cast(valid_len_last)); + if (i % 2 == 0) { + TASSIGN(sijDynTile, 0x0); + TFILLPAD_INPLACE(sijPadTile_A, sijDynTile); + } else { + TASSIGN(sijDynTile, static_cast(kDataBytes)); + TFILLPAD_INPLACE(sijPadTile_B, sijDynTile); + } + pipe_barrier(PIPE_V); + } + + if (i % 2 == 0) { + TROWMAX(localMaxDN, sijTile_A, tmpTile); + } else { + TROWMAX(localMaxDN, sijTile_B, tmpTile); + } + pipe_barrier(PIPE_V); + + if (i + 1 < n_blocks) { + GlobalDataMxN sijGlobal_next(sij_base + (i + 1) * TILE_STRIDE); + if (i % 2 == 0) { + TLOAD(sijTile_B, sijGlobal_next); + } else { + TLOAD(sijTile_A, sijGlobal_next); + } + } + + TRESHAPE(localMaxRow, localMaxDN); + if (i == 0) { + TMAX(globalMaxRow, localMaxRow, localMaxRow); + } else { + TMAX(globalMaxRow, globalMaxRow, localMaxRow); + } + pipe_barrier(PIPE_V); + } + + TMULS(globalMaxRow, globalMaxRow, scale_value); + pipe_barrier(PIPE_V); + TRESHAPE(globalMaxDN, globalMaxRow); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(mijGlobalND, globalMaxND); + + // ======== Pass 2: Compute softmax ======== + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + + GlobalDataMxN sijGlobal_0(sij_base); + TLOAD(sijTile_A, sijGlobal_0); + + for (uint64_t i = 0; i < n_blocks; i++) { + GlobalDataMxN_bf16 pijGlobal(pij_base + i * TILE_STRIDE); + + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + + if (i == n_blocks - 1 && valid_len_last < static_cast(TN)) { + TileSijDyn curSijDyn(static_cast(valid_len_last)); + if (i % 2 == 0) { + TASSIGN(curSijDyn, 0x0); + TFILLPAD_INPLACE(sijPadTile_A, curSijDyn); + } else { + TASSIGN(curSijDyn, static_cast(kDataBytes)); + TFILLPAD_INPLACE(sijPadTile_B, curSijDyn); + } + pipe_barrier(PIPE_V); + } + + if (i % 2 == 0) { + TMULS(sijTile_A, sijTile_A, scale_value); + pipe_barrier(PIPE_V); + TROWEXPANDSUB(pijTile, sijTile_A, globalMaxDN); + } else { + TMULS(sijTile_B, sijTile_B, scale_value); + pipe_barrier(PIPE_V); + TROWEXPANDSUB(pijTile, sijTile_B, globalMaxDN); + } + pipe_barrier(PIPE_V); + TEXP(pijTile, pijTile); + pipe_barrier(PIPE_V); + TCVT(pijBf16Tile, pijTile, RoundMode::CAST_ROUND); + pipe_barrier(PIPE_V); + TCVT(pijTile, pijBf16Tile, RoundMode::CAST_ROUND); + + pipe_barrier(PIPE_V); + if (i == 0) { + TMULS(sumAccTile, pijTile, 1.0f); + } else { + TADD(sumAccTile, sumAccTile, pijTile); + } + + pipe_barrier(PIPE_V); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(pijGlobal, pijBf16Tile); + + if (i + 1 < n_blocks) { + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + GlobalDataMxN sijGlobal_next(sij_base + (i + 1) * TILE_STRIDE); + if (i % 2 == 0) { + TLOAD(sijTile_B, sijGlobal_next); + } else { + TLOAD(sijTile_A, sijGlobal_next); + } + } + } + + pipe_barrier(PIPE_V); + TROWSUM(sumDN, sumAccTile, tmpTile); + + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + TSTORE(lijGlobalDN, sumDN); + + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7); +} + +extern "C" __aicore__ void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *sij_t = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ Tensor *context_lens_t = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ Tensor *pij_t = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ Tensor *mij_t = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ Tensor *lij_t = reinterpret_cast<__gm__ Tensor *>(args[4]); + float scale_value = from_u64(static_cast(args[5])); + int64_t bn_start = static_cast(args[6]); + int64_t n_blocks_total = static_cast(args[7]); + int64_t block_size = static_cast(args[8]); + int64_t q_loop = static_cast(args[9]); + + int32_t block_idx = get_block_idx(args); + int32_t sub_block_id = get_sub_block_id(args); // 0 = AIV0, 1 = AIV1 + int64_t batch_idx = block_idx / q_loop; + + int64_t q_tile = (block_size == 128) ? 16 : 64; + int64_t sub_m = q_tile / 2; + + // Compute valid_len for the last block in this group + __gm__ int32_t *ctx_ptr = + reinterpret_cast<__gm__ int32_t *>(context_lens_t->buffer.addr) + context_lens_t->start_offset; + int64_t cur_seq = static_cast(ctx_ptr[batch_idx]); + int64_t bn_this_batch = (cur_seq + block_size - 1) / block_size; + + // Clamp n_blocks to valid range for this batch + int64_t valid_blocks = bn_this_batch - bn_start; + if (valid_blocks < 0) valid_blocks = 0; + int64_t n_blocks = (valid_blocks < n_blocks_total) ? valid_blocks : n_blocks_total; + + // Compute valid_len for the last valid block + uint64_t valid_len_last; + if (n_blocks <= 0) { + valid_len_last = 0; + } else { + int64_t last_block_seq_start = (bn_start + n_blocks - 1) * block_size; + int64_t remaining = cur_seq - last_block_seq_start; + if (remaining >= block_size) { + valid_len_last = static_cast(block_size); + } else if (remaining > 0) { + valid_len_last = static_cast(remaining); + } else { + valid_len_last = 0; + } + } + + // Packed tile layout: SPMD block block_idx owns a contiguous region of + // n_blocks_total packed (q_tile, block_size) tiles. + // Packed base for this SPMD block: + int64_t packed_base_offset = block_idx * q_tile * n_blocks_total * block_size; + // AIV lane offset within each packed tile (sub_block_id selects the sub_m-row half): + int64_t sub_offset = sub_block_id * sub_m * block_size; + + __gm__ float *sij_base = + reinterpret_cast<__gm__ float *>(sij_t->buffer.addr) + sij_t->start_offset + packed_base_offset + sub_offset; + __gm__ bfloat16_t *pij_base = + reinterpret_cast<__gm__ bfloat16_t *>(pij_t->buffer.addr) + pij_t->start_offset + packed_base_offset + + sub_offset; + + // Scalar layout: full q_tile rows pack to kAlignedRowsFull floats per block_idx; + // each AIV lane owns kAlignedRowsSub contiguous floats inside that slab. + int64_t kAlignedRowsFull = ((q_tile * static_cast(sizeof(float)) + 31) / 32) * (32 / static_cast(sizeof(float))); + int64_t kAlignedRowsSub = ((sub_m * static_cast(sizeof(float)) + 31) / 32) * (32 / static_cast(sizeof(float))); + int64_t scalar_offset = block_idx * kAlignedRowsFull + sub_block_id * kAlignedRowsSub; + __gm__ float *mij_addr = + reinterpret_cast<__gm__ float *>(mij_t->buffer.addr) + mij_t->start_offset + scalar_offset; + __gm__ float *lij_addr = + reinterpret_cast<__gm__ float *>(lij_t->buffer.addr) + lij_t->start_offset + scalar_offset; + + if (n_blocks <= 0) { + // No valid KV data — emit neutral values + for (int64_t i = 0; i < kAlignedRowsSub; i++) { + mij_addr[i] = -1e30f; + lij_addr[i] = 0.0f; + } + return; + } + + // Tile stride = full packed tile size (q_tile * block_size), NOT sub_m * block_size + if (q_tile == 16) { + softmax_prepare_n_spmd<8, 128, 16 * 128>( + sij_base, scale_value, pij_base, mij_addr, lij_addr, + static_cast(n_blocks), valid_len_last + ); + } else { + softmax_prepare_n_spmd<32, 64, 64 * 64>( + sij_base, scale_value, pij_base, mij_addr, lij_addr, + static_cast(n_blocks), valid_len_last + ); + } +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/kernel_config.py b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/kernel_config.py new file mode 100644 index 0000000000..ea77a7bfea --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/kernel_config.py @@ -0,0 +1,85 @@ +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +""" +SPMD Paged Attention Kernel and Orchestration Configuration + +Uses SPMD (block_num) parallelism across batch*q_loop positions. +Each block handles one (batch_idx, q_tile_idx) using get_block_idx(). + +Softmax and online-update run as MIX tasks (AIC idle + AIV0 + AIV1), with the +two AIVs splitting the 16 query rows 8/8 via get_sub_block_id(). + +AIC Kernels (Matrix Multiplication): + - aic_qk_matmul: Q @ K^T (SPMD across batch*q_loop) + - aic_pv_matmul: P @ V (SPMD across batch*q_loop) + - aic_hub: no-op, occupies the AIC slot of softmax/update MIX tasks + +AIV Kernels (Vector Operations): + - aiv_softmax_prepare: scale, rowmax, exp, rowsum on 8-row sub-tile + - aiv_online_update: online softmax accumulation + normalization on 8-row sub-tile + - aiv_hub: no-op, used to allocate persistent accumulators +""" + +from pathlib import Path + +from simpler.task_interface import ArgDirection as D # pyright: ignore[reportAttributeAccessIssue] + +_KERNELS_ROOT = Path(__file__).parent + +ORCHESTRATION = { + "source": str(_KERNELS_ROOT / "orchestration" / "spmd_paged_attention_orch.cpp"), + "function_name": "aicpu_orchestration_entry", +} + +KERNELS = [ + # AIC kernels (matrix multiplication using Cube unit) + { + "func_id": 0, + "name": "SPMD_QK", + "source": str(_KERNELS_ROOT / "aic" / "aic_qk_matmul.cpp"), + "core_type": "aic", + }, + { + "func_id": 1, + "name": "SPMD_PV", + "source": str(_KERNELS_ROOT / "aic" / "aic_pv_matmul.cpp"), + "core_type": "aic", + }, + { + "func_id": 2, + "name": "AIC_HUB", + "source": str(_KERNELS_ROOT / "aic" / "aic_hub.cpp"), + "core_type": "aic", + }, + # AIV kernels (vector operations) + { + "func_id": 3, + "name": "SPMD_SF", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_softmax_prepare.cpp"), + "core_type": "aiv", + }, + { + "func_id": 4, + "name": "SPMD_UP", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_online_update.cpp"), + "core_type": "aiv", + }, + { + "func_id": 5, + "name": "AIV_HUB", + "source": str(_KERNELS_ROOT / "aiv" / "aiv_hub.cpp"), + "core_type": "aiv", + }, +] + +RUNTIME_CONFIG = { + "runtime": "tensormap_and_ringbuffer", + "aicpu_thread_num": 4, + "block_dim": 24, +} diff --git a/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/orchestration/spmd_paged_attention_orch.cpp b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/orchestration/spmd_paged_attention_orch.cpp new file mode 100644 index 0000000000..b391a314d5 --- /dev/null +++ b/tests/st/a2a3/tensormap_and_ringbuffer/spmd_paged_attention_unroll/kernels/orchestration/spmd_paged_attention_orch.cpp @@ -0,0 +1,248 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +/** + * SPMD Paged Attention Orchestration with N_UNROLL block batching + * + * Uses SPMD parallelism: block_num = batch * q_loop, where each logical + * block handles one (batch_idx, q_tile_idx) position. Kernels use + * get_block_idx() to compute their data offsets. + * + * Batches up to N_UNROLL KV blocks per group. Each group submits 4 tasks: + * 1. QK matmul: qi @ K^T for n_blocks → sij (Q_TILE, n_blocks * block_size) + * 2. Softmax: two-pass over sij → pij, mij, lij + * 3. PV matmul: SplitK accumulated P @ V → oi_new (Q_TILE, head_dim) + * 4. Update: online softmax accumulation with group-level mi, li, oi_new + * + * Softmax and online-update run as MIX tasks (AIC hub + AIV0 + AIV1) with + * dual-vector subvector partitioning via get_sub_block_id(). + * + * Memory Layout: + * Query: (batch, num_heads, head_dim) - bfloat16 + * Key/Value: (total_blocks, block_size, kv_head_num, head_dim) - bfloat16 + * Block Table: (batch, max_num_blocks_per_req) - int32 + * Context Lens: (batch,) - int32 + * Output: (batch, num_heads, head_dim) - float32 + */ + +#include +#include + +#include +#include + +#include "pto_orchestration_api.h" + +#define N_UNROLL 64 + +#define FUNC_QK_MATMUL 0 +#define FUNC_PV_MATMUL 1 +#define FUNC_AIC_HUB 2 +#define FUNC_SOFTMAX_PREPARE 3 +#define FUNC_ONLINE_UPDATE 4 +#define FUNC_AIV_HUB 5 + +static constexpr uint64_t Q_TILE = 16; + +extern "C" { + +__attribute__((visibility("default"))) PTO2OrchestrationConfig +aicpu_orchestration_config(const ChipStorageTaskArgs &orch_args) { + (void)orch_args; + return PTO2OrchestrationConfig{ + .expected_arg_count = 7, + }; +} + +__attribute__((visibility("default"))) void aicpu_orchestration_entry(const ChipStorageTaskArgs &orch_args) { + // query: shape=[batch, num_heads, head_dim] + uint64_t batch = orch_args.tensor(0).shapes[0]; + uint64_t num_heads = orch_args.tensor(0).shapes[1]; + uint64_t head_dim = orch_args.tensor(0).shapes[2]; + DataType data_type = orch_args.tensor(0).dtype; + + // key_cache: shape=[total_blocks, block_size, kv_head_num, head_dim] + uint64_t block_size = orch_args.tensor(1).shapes[1]; + + // block_table: shape=[batch, max_num_blocks_per_req] + uint64_t max_num_blocks_per_req = orch_args.tensor(3).shapes[1]; + + // scale from scalar arg + uint64_t scale_value = orch_args.scalar(0); + + uint64_t q_loop = (num_heads + Q_TILE - 1) / Q_TILE; + int16_t spmd_block_num = static_cast(batch * q_loop); + + LOG_INFO( + "SPMD PA Unroll: batch=%" PRIu64 " heads=%" PRIu64 " hd=%" PRIu64 " bs=%" PRIu64 " q_loop=%" PRIu64 + " blocks=%d N_UNROLL=%d", + batch, num_heads, head_dim, block_size, q_loop, spmd_block_num, N_UNROLL + ); + + // Wrap host-provided tensors + void *query_ptr = orch_args.tensor(0).data_as(); + void *kc_ptr = orch_args.tensor(1).data_as(); + void *vc_ptr = orch_args.tensor(2).data_as(); + void *out_ptr = orch_args.tensor(5).data_as(); + + uint64_t total_kv_blocks = orch_args.tensor(1).shapes[0]; + uint64_t kv_total_rows = total_kv_blocks * block_size; + + uint32_t query_shapes[2] = {static_cast(batch * num_heads), static_cast(head_dim)}; + uint32_t kv_shapes[2] = {static_cast(kv_total_rows), static_cast(head_dim)}; + uint32_t out_shapes[2] = {static_cast(batch * num_heads), static_cast(head_dim)}; + + Tensor query = make_tensor_external(query_ptr, query_shapes, 2, data_type); + Tensor key_cache = make_tensor_external(kc_ptr, kv_shapes, 2, data_type); + Tensor value_cache = make_tensor_external(vc_ptr, kv_shapes, 2, data_type); + Tensor out = make_tensor_external(out_ptr, out_shapes, 2, DataType::FLOAT32); + + uint32_t bt_shapes[2] = {static_cast(batch), static_cast(max_num_blocks_per_req)}; + Tensor block_table = + make_tensor_external(orch_args.tensor(3).data_as(), bt_shapes, 2, DataType::INT32, false); + uint32_t cl_shapes[1] = {static_cast(batch)}; + Tensor context_lens = + make_tensor_external(orch_args.tensor(4).data_as(), cl_shapes, 1, DataType::INT32, false); + + // Find max context_len for KV block loop bound + uint64_t max_ctx = 0; + for (uint64_t b = 0; b < batch; b++) { + uint32_t idx[1] = {static_cast(b)}; + uint64_t ctx = static_cast(get_tensor_data(context_lens, 1, idx)); + if (ctx > max_ctx) max_ctx = ctx; + } + uint64_t max_bn = (max_ctx + block_size - 1) / block_size; + + // Accumulator create infos (persistent across the bn loop) + uint32_t n_rows = static_cast(spmd_block_num) * static_cast(Q_TILE); + uint32_t acc_oi_shapes[2] = {n_rows, static_cast(head_dim)}; + uint32_t scalar_shapes[1] = {n_rows}; + TensorCreateInfo acc_oi_ci(acc_oi_shapes, 2, DataType::FLOAT32); + TensorCreateInfo acc_mi_ci(scalar_shapes, 1, DataType::FLOAT32); + TensorCreateInfo acc_li_ci(scalar_shapes, 1, DataType::FLOAT32); + + PTO2_SCOPE() { + // Allocate persistent accumulators via no-op AIV hub + Arg hub_args; + hub_args.add_output(acc_oi_ci); + hub_args.add_output(acc_mi_ci); + hub_args.add_output(acc_li_ci); + TaskOutputTensors hub_outs = pto2_rt_submit_aiv_task(FUNC_AIV_HUB, hub_args); + const Tensor &oi_acc = hub_outs.get_ref(0); + const Tensor &mi_acc = hub_outs.get_ref(1); + const Tensor &li_acc = hub_outs.get_ref(2); + + for (uint64_t bn = 0; bn < max_bn; bn += N_UNROLL) { + uint64_t n_blocks = std::min(static_cast(N_UNROLL), max_bn - bn); + uint64_t is_first = (bn == 0) ? 1 : 0; + uint64_t is_last = (bn + n_blocks >= max_bn) ? 1 : 0; + + // Scratch tensor create infos sized for n_blocks per SPMD block + uint32_t sij_shapes[2] = {n_rows, static_cast(n_blocks * block_size)}; + uint32_t pij_shapes[2] = {n_rows, static_cast(n_blocks * block_size)}; + uint32_t oi_new_shapes[2] = {n_rows, static_cast(head_dim)}; + uint32_t mij_shapes[1] = {n_rows}; + uint32_t lij_shapes[1] = {n_rows}; + + TensorCreateInfo sij_ci(sij_shapes, 2, DataType::FLOAT32); + TensorCreateInfo pij_ci(pij_shapes, 2, data_type); + TensorCreateInfo oi_new_ci(oi_new_shapes, 2, DataType::FLOAT32); + TensorCreateInfo mij_ci(mij_shapes, 1, DataType::FLOAT32); + TensorCreateInfo lij_ci(lij_shapes, 1, DataType::FLOAT32); + + // -- QK Matmul (AIC, SPMD): n_blocks matmuls per SPMD block -- + Arg qk_args; + qk_args.add_input(query); + qk_args.add_input(key_cache); + qk_args.add_input(block_table); + qk_args.add_input(context_lens); + qk_args.add_output(sij_ci); + qk_args.add_scalar(static_cast(bn)); + qk_args.add_scalar(static_cast(n_blocks)); + qk_args.add_scalar(static_cast(num_heads)); + qk_args.add_scalar(static_cast(head_dim)); + qk_args.add_scalar(static_cast(block_size)); + qk_args.add_scalar(static_cast(max_num_blocks_per_req)); + qk_args.add_scalar(static_cast(q_loop)); + qk_args.launch_spec.set_block_num(spmd_block_num); + TaskOutputTensors qk_outs = pto2_rt_submit_aic_task(FUNC_QK_MATMUL, qk_args); + const Tensor &sij = qk_outs.get_ref(0); + + // -- Softmax Prepare (MIX: AIC hub + AIV0 + AIV1, SPMD) -- + // Two-pass softmax over n_blocks tiles per SPMD block + Arg sf_args; + sf_args.add_input(sij); + sf_args.add_input(context_lens); + sf_args.add_output(pij_ci); + sf_args.add_output(mij_ci); + sf_args.add_output(lij_ci); + sf_args.add_scalar(scale_value); + sf_args.add_scalar(static_cast(bn)); + sf_args.add_scalar(static_cast(n_blocks)); + sf_args.add_scalar(static_cast(block_size)); + sf_args.add_scalar(static_cast(q_loop)); + sf_args.launch_spec.set_block_num(spmd_block_num); + MixedKernels sf_mk; + sf_mk.aic_kernel_id = FUNC_AIC_HUB; + sf_mk.aiv0_kernel_id = FUNC_SOFTMAX_PREPARE; + sf_mk.aiv1_kernel_id = FUNC_SOFTMAX_PREPARE; + TaskOutputTensors sf_outs = pto2_rt_submit_task(sf_mk, sf_args); + const Tensor &pij = sf_outs.get_ref(0); + const Tensor &mij = sf_outs.get_ref(1); + const Tensor &lij = sf_outs.get_ref(2); + + // -- PV Matmul (AIC, SPMD): SplitK accumulated across n_blocks -- + Arg pv_args; + pv_args.add_input(pij); + pv_args.add_input(value_cache); + pv_args.add_input(block_table); + pv_args.add_input(context_lens); + pv_args.add_output(oi_new_ci); + pv_args.add_scalar(static_cast(bn)); + pv_args.add_scalar(static_cast(n_blocks)); + pv_args.add_scalar(static_cast(num_heads)); + pv_args.add_scalar(static_cast(head_dim)); + pv_args.add_scalar(static_cast(block_size)); + pv_args.add_scalar(static_cast(max_num_blocks_per_req)); + pv_args.add_scalar(static_cast(q_loop)); + pv_args.launch_spec.set_block_num(spmd_block_num); + TaskOutputTensors pv_outs = pto2_rt_submit_aic_task(FUNC_PV_MATMUL, pv_args); + const Tensor &oi_new = pv_outs.get_ref(0); + + // -- Online Update (MIX: AIC hub + AIV0 + AIV1, SPMD) -- + Arg up_args; + up_args.add_input(mij); + up_args.add_input(lij); + up_args.add_input(oi_new); + up_args.add_inout(mi_acc); + up_args.add_inout(li_acc); + up_args.add_inout(oi_acc); + up_args.add_inout(out); + up_args.add_scalar(is_first); + up_args.add_scalar(is_last); + up_args.add_scalar(static_cast(num_heads)); + up_args.add_scalar(static_cast(head_dim)); + up_args.add_scalar(static_cast(q_loop)); + up_args.launch_spec.set_block_num(spmd_block_num); + MixedKernels up_mk; + up_mk.aic_kernel_id = FUNC_AIC_HUB; + up_mk.aiv0_kernel_id = FUNC_ONLINE_UPDATE; + up_mk.aiv1_kernel_id = FUNC_ONLINE_UPDATE; + pto2_rt_submit_task(up_mk, up_args); + } + } + + uint64_t n_groups = (max_bn + N_UNROLL - 1) / N_UNROLL; + LOG_INFO( + "SPMD PA Unroll: %" PRIu64 " groups x 4 tasks, blocks=%d", n_groups, static_cast(spmd_block_num) + ); +} + +} // extern "C" diff --git a/tools/benchmark_rounds.sh b/tools/benchmark_rounds.sh index 6d684ea193..bb02e38e02 100755 --- a/tools/benchmark_rounds.sh +++ b/tools/benchmark_rounds.sh @@ -33,12 +33,14 @@ declare -A TMR_EXAMPLE_CASES=( [benchmark_bgemm]="" [paged_attention_unroll]="Case1,Case2" [batch_paged_attention]="" + [spmd_paged_attention_unroll]="Case1,Case2" ) TMR_EXAMPLE_ORDER=( alternating_matmul_add benchmark_bgemm paged_attention_unroll batch_paged_attention + spmd_paged_attention_unroll ) # --- aicpu_build_graph ---