diff --git a/.github/script/sweep_common.py b/.github/script/sweep_common.py index 238001f9a..68114f82b 100644 --- a/.github/script/sweep_common.py +++ b/.github/script/sweep_common.py @@ -100,10 +100,16 @@ def build_parser(description: str) -> argparse.ArgumentParser: parser.add_argument( "--platform", default="windows" if IS_WINDOWS else "linux", - # The -hrx variants are the same OS running an HRX build. Nothing in - # the sweep behaves differently; the label exists so the two builds' - # results do not land in the same filename or the same summary row. - choices=["linux", "windows", "linux-hrx", "windows-hrx"], + # Build variants use separate labels so their results do not land in + # the same filename or summary row. + choices=[ + "linux", + "windows", + "linux-hrx", + "windows-hrx", + "linux-models-from-source", + "windows-models-from-source", + ], help="Label recorded in the output filename (default: autodetected).", ) parser.add_argument("--host", default="127.0.0.1", help="Server bind address.") diff --git a/.github/script/sweep_summary.py b/.github/script/sweep_summary.py index 0d5176120..26fa3883e 100644 --- a/.github/script/sweep_summary.py +++ b/.github/script/sweep_summary.py @@ -35,6 +35,16 @@ def main(argv: list[str] | None = None) -> int: parser.add_argument( "--windows-hrx-result", default="", help="Job result for the Windows HRX sweep." ) + parser.add_argument( + "--linux-models-from-source-result", + default="", + help="Job result for the Linux models-from-source sweep.", + ) + parser.add_argument( + "--windows-models-from-source-result", + default="", + help="Job result for the Windows models-from-source sweep.", + ) args = parser.parse_args(argv) out: list[str] = ["# Model Sweep Results", ""] @@ -46,6 +56,8 @@ def main(argv: list[str] | None = None) -> int: ("Linux (HRX)", args.linux_hrx_result), ("Windows", args.windows_result), ("Windows (HRX)", args.windows_hrx_result), + ("Linux (models from source)", args.linux_models_from_source_result), + ("Windows (models from source)", args.windows_models_from_source_result), ] if any(result for _, result in overall): out += ["| Platform | Overall |", "| --- | --- |"] diff --git a/.github/workflows/debian-portable.yml b/.github/workflows/debian-portable.yml index b3751cf77..c7716d8af 100644 --- a/.github/workflows/debian-portable.yml +++ b/.github/workflows/debian-portable.yml @@ -97,6 +97,52 @@ jobs: retention-days: 7 if-no-files-found: error + - name: Build models from source + run: | + cd src + cmake --preset linux-portable \ + -B build-models-from-source \ + -DFLM_BUILD_GEMMA4E=ON + cmake --build build-models-from-source --parallel 2 + + - name: Create models-from-source package + env: + VERSION: ${{ steps.get_version.outputs.version }} + run: | + cd src/build-models-from-source + rm -rf install-models-from-source + DESTDIR=$PWD/install-models-from-source cmake --install . + + package_root=install-models-from-source/opt/fastflowlm + mv "${package_root}/flm" "${package_root}/flm-real" + mv "${package_root}/flm-wrapper.sh" "${package_root}/flm" + chmod +x "${package_root}/flm" + strip --strip-debug "${package_root}/flm-real" + + staged_engine="${package_root}/lib/libgemma4e_npu.so" + test -f "${staged_engine}" + if cmp -s "$GITHUB_WORKSPACE/src/lib/xrt/libgemma4e_npu.so" "${staged_engine}"; then + echo "The staged Gemma 4 engine matches the prebuilt library." >&2 + exit 1 + fi + readelf -n "${staged_engine}" | grep -A1 'Build ID' + tar czf "$GITHUB_WORKSPACE/models-from-source_fastflowlm_${VERSION}_linux.tar.gz" \ + -C "${package_root}" \ + flm \ + flm-real \ + lib/ \ + xclbins/ \ + model_list.json \ + model_info.json + + - name: Upload models-from-source build artifact + uses: actions/upload-artifact@v4 + with: + name: models-from-source-portable + path: models-from-source_fastflowlm_${{ steps.get_version.outputs.version }}_linux.tar.gz + retention-days: 7 + if-no-files-found: error + - name: Collect FFmpeg logs if: always() shell: bash diff --git a/.github/workflows/model-sweep.yml b/.github/workflows/model-sweep.yml index 6fe02954a..e90dd14e4 100644 --- a/.github/workflows/model-sweep.yml +++ b/.github/workflows/model-sweep.yml @@ -17,12 +17,10 @@ name: Model Sweep # matrix job each, so a failure in one modality does not hide the others and # "Re-run failed jobs" restarts only what actually broke. # -# Each platform is swept twice: once against the default portable build and -# once against the HRX one, which is the same source built with FLM_USE_HRX=ON -# and a different set of op libraries underneath. Only the artifact differs, so -# the HRX jobs are their non-HRX counterparts with a different download step -# and a different --platform label. The HRX jobs are commit-sweep only, because -# releases do not currently ship an HRX asset. +# Each platform is swept against the default portable build and the HRX build. +# Commit sweeps also run Gemma 4 against separate Linux and Windows packages +# that build model engines from source. Releases do not ship these additional +# artifacts. on: pull_request: @@ -92,11 +90,10 @@ jobs: # what stops a sweep from silently picking up a stale artifact built from an # earlier push on the same branch. # - # Readiness is per build, not per workflow run: each of the four sweeps has - # its own [_hrx]_ready flag, decided by whether that one artifact - # exists. So an HRX build that fails costs you the HRX sweeps and nothing - # else, and vice versa. The build workflow already reports its own failure, - # so there is no reason to fail twice or to hold back a build that is fine. + # Readiness is per build, not per workflow run. Each artifact has its own + # readiness flag. A failed optional build does not block the other sweeps. + # The build workflow reports its failure, so this workflow skips the missing + # artifact. resolve-build: runs-on: ubuntu-latest timeout-minutes: 90 @@ -107,7 +104,9 @@ jobs: linux_run_id: ${{ steps.resolve.outputs.linux_run_id }} windows_run_id: ${{ steps.resolve.outputs.windows_run_id }} linux_ready: ${{ steps.resolve.outputs.linux_ready }} + linux_models_from_source_ready: ${{ steps.resolve.outputs.linux_models_from_source_ready }} windows_ready: ${{ steps.resolve.outputs.windows_ready }} + windows_models_from_source_ready: ${{ steps.resolve.outputs.windows_models_from_source_ready }} linux_hrx_ready: ${{ steps.resolve.outputs.linux_hrx_ready }} windows_hrx_ready: ${{ steps.resolve.outputs.windows_hrx_ready }} @@ -212,6 +211,7 @@ jobs: wait_for_build() { local workflow="$1" platform="$2" default_pattern="$3" hrx_pattern="$4" + local source_pattern="${5:-}" source_flag="${6:-}" source_label="${7:-}" local run_id="" status artifacts while :; do @@ -254,14 +254,22 @@ jobs: "${platform}_ready" "${platform} portable" "${run_id}" check_artifact "${artifacts}" "${hrx_pattern}" \ "${platform}_hrx_ready" "${platform} HRX portable" "${run_id}" + if [ -n "${source_pattern}" ]; then + check_artifact "${artifacts}" "${source_pattern}" \ + "${source_flag}" "${source_label}" "${run_id}" + fi } # Neither platform gates the other, so a failure here leaves that # platform not-ready rather than aborting the step. wait_for_build debian-portable.yml linux \ - '^portable$' '^HRX-portable$' || echo "Linux build unavailable." + '^portable$' '^HRX-portable$' \ + '^models-from-source-portable$' 'linux_models_from_source_ready' \ + 'Linux models-from-source portable' || echo "Linux build unavailable." wait_for_build windows-build.yml windows \ '^fastflowlm-windows-[0-9a-f]+$' '^HRX-fastflowlm-windows-[0-9a-f]+$' \ + '^models-from-source-fastflowlm-windows-[0-9a-f]+$' \ + 'windows_models_from_source_ready' 'Windows models-from-source portable' \ || echo "Windows build unavailable." sweep-linux: @@ -352,7 +360,12 @@ jobs: set -euo pipefail python3 -m venv .venv .venv/bin/python -m pip install --upgrade pip - .venv/bin/python -m pip install openai + for attempt in 1 2 3; do + .venv/bin/python -m pip install --no-cache-dir openai && break + if [ "${attempt}" -eq 3 ]; then + exit 1 + fi + done # Each script starts and stops its own flm server, picks its own models, # and exits non-zero if any model failed. @@ -378,6 +391,83 @@ jobs: if-no-files-found: warn retention-days: 30 + sweep-linux-models-from-source: + needs: resolve-build + if: needs.resolve-build.outputs.linux_models_from_source_ready == 'true' + runs-on: [self-hosted, linux, x64, npu, sweep] + timeout-minutes: 60 + + strategy: + fail-fast: false + matrix: + task: [llm, vision] + + env: + FLM_MODEL_PATH: /scratch + + steps: + - name: Checkout sweep scripts + uses: actions/checkout@v4 + with: + sparse-checkout: .github/script + filter: blob:none + + - name: Download Linux models-from-source build + uses: dawidd6/action-download-artifact@v25 + with: + github_token: ${{ secrets.GITHUB_TOKEN }} + run_id: ${{ needs.resolve-build.outputs.linux_run_id }} + name: models-from-source-portable + path: build-artifact + if_no_artifact_found: fail + + - name: Unpack Linux models-from-source build + run: | + set -euo pipefail + mkdir -p flm-linux-models-from-source + tarball=$(ls build-artifact/models-from-source_fastflowlm_*_linux.tar.gz) + echo "Using tarball: ${tarball}" + tar xzf "${tarball}" -C flm-linux-models-from-source + chmod +x flm-linux-models-from-source/flm flm-linux-models-from-source/flm-real + echo "FLM_BIN=${GITHUB_WORKSPACE}/flm-linux-models-from-source/flm" >> "$GITHUB_ENV" + + - name: Verify FLM binary runs + run: | + set -euo pipefail + "${FLM_BIN}" version + + - name: Set up Python virtualenv + run: | + set -euo pipefail + python3 -m venv .venv + .venv/bin/python -m pip install --upgrade pip + for attempt in 1 2 3; do + .venv/bin/python -m pip install --no-cache-dir openai && break + if [ "${attempt}" -eq 3 ]; then + exit 1 + fi + done + + - name: Run Gemma 4 ${{ matrix.task }} sweep + run: | + set -euo pipefail + .venv/bin/python .github/script/sweep_${{ matrix.task }}.py \ + --flm-bin "${FLM_BIN}" \ + --platform linux-models-from-source \ + --models gemma4-it:e2b gemma4-it:e4b \ + --gen-lim "${GEN_LIM}" \ + --request-timeout "${REQUEST_TIMEOUT}" \ + --output-dir sweep-results + + - name: Upload Gemma 4 ${{ matrix.task }} results + if: always() + uses: actions/upload-artifact@v4 + with: + name: model-sweep-linux-models-from-source-${{ matrix.task }}-${{ github.run_number }} + path: sweep-results/ + if-no-files-found: warn + retention-days: 30 + # The Linux sweep again, against the HRX build. "Build Linux Portable" # produces both from the same run, so this shares linux_run_id with # sweep-linux but has its own readiness flag: whichever of the two builds @@ -446,7 +536,12 @@ jobs: set -euo pipefail python3 -m venv .venv .venv/bin/python -m pip install --upgrade pip - .venv/bin/python -m pip install openai + for attempt in 1 2 3; do + .venv/bin/python -m pip install --no-cache-dir openai && break + if [ "${attempt}" -eq 3 ]; then + exit 1 + fi + done # --platform is only a label, but it is the one that keeps these results # apart from the default build's in the filenames and the summary table. @@ -582,6 +677,84 @@ jobs: if-no-files-found: warn retention-days: 30 + sweep-windows-models-from-source: + needs: resolve-build + if: needs.resolve-build.outputs.windows_models_from_source_ready == 'true' + runs-on: [self-hosted, windows, x64, npu, sweep] + timeout-minutes: 60 + + strategy: + fail-fast: false + matrix: + task: [llm, vision] + + steps: + - name: Checkout sweep scripts + uses: actions/checkout@v4 + with: + sparse-checkout: .github/script + filter: blob:none + + - name: Resolve model path + shell: powershell + run: | + $modelPath = Join-Path $env:USERPROFILE ".flm" + New-Item -ItemType Directory -Force $modelPath | Out-Null + "FLM_MODEL_PATH=$modelPath" | Out-File -FilePath $env:GITHUB_ENV -Encoding utf8 -Append + + - name: Download Windows models-from-source build + uses: dawidd6/action-download-artifact@v25 + with: + github_token: ${{ secrets.GITHUB_TOKEN }} + run_id: ${{ needs.resolve-build.outputs.windows_run_id }} + name: models-from-source-fastflowlm-windows-${{ github.sha }} + path: flm-windows-models-from-source + if_no_artifact_found: fail + + - name: Verify FLM binary runs + shell: powershell + run: | + $ErrorActionPreference = "Stop" + $flm = Join-Path $env:GITHUB_WORKSPACE "flm-windows-models-from-source\flm.exe" + if (-not (Test-Path $flm)) { + throw "flm.exe not found at $flm" + } + "FLM_BIN=$flm" | Out-File -FilePath $env:GITHUB_ENV -Encoding utf8 -Append + & $flm version + + - name: Set up Python virtualenv + shell: powershell + run: | + $ErrorActionPreference = "Stop" + python -m venv .venv + .\.venv\Scripts\python.exe -m pip install --upgrade pip + .\.venv\Scripts\python.exe -m pip install openai + + - name: Run models-from-source ${{ matrix.task }} sweep + shell: powershell + run: | + $ErrorActionPreference = "Stop" + $cmdArgs = @( + ".github\script\sweep_${{ matrix.task }}.py", + "--flm-bin", $env:FLM_BIN, + "--platform", "windows-models-from-source", + "--models", "gemma4-it:e2b", "gemma4-it:e4b", + "--gen-lim", $env:GEN_LIM, + "--request-timeout", $env:REQUEST_TIMEOUT, + "--output-dir", "sweep-results" + ) + & ".\.venv\Scripts\python.exe" $cmdArgs + if ($LASTEXITCODE -ne 0) { exit $LASTEXITCODE } + + - name: Upload models-from-source ${{ matrix.task }} results + if: always() + uses: actions/upload-artifact@v4 + with: + name: model-sweep-windows-models-from-source-${{ matrix.task }}-${{ github.run_number }} + path: sweep-results/ + if-no-files-found: warn + retention-days: 30 + # The Windows sweep again, against the HRX build. "Build Windows Packages" # produces both from the same run, so this shares windows_run_id with # sweep-windows but has its own readiness flag. @@ -683,7 +856,7 @@ jobs: # Reporting only, no NPU needed, so this one stays on a GitHub-hosted runner. summary: runs-on: ubuntu-latest - needs: [sweep-linux, sweep-windows, sweep-linux-hrx, sweep-windows-hrx] + needs: [sweep-linux, sweep-linux-models-from-source, sweep-windows, sweep-windows-models-from-source, sweep-linux-hrx, sweep-windows-hrx] if: always() steps: # Only the sweep scripts and their test assets, never the sources: this @@ -714,4 +887,6 @@ jobs: --windows-result "${{ needs.sweep-windows.result }}" \ --linux-hrx-result "${{ needs.sweep-linux-hrx.result }}" \ --windows-hrx-result "${{ needs.sweep-windows-hrx.result }}" \ + --linux-models-from-source-result "${{ needs.sweep-linux-models-from-source.result }}" \ + --windows-models-from-source-result "${{ needs.sweep-windows-models-from-source.result }}" \ >> "$GITHUB_STEP_SUMMARY" diff --git a/.github/workflows/windows-build.yml b/.github/workflows/windows-build.yml index 6b3f84bf2..5657a62fe 100644 --- a/.github/workflows/windows-build.yml +++ b/.github/workflows/windows-build.yml @@ -67,6 +67,54 @@ jobs: path: src/dist/ if-no-files-found: error + - name: Configure models-from-source build + shell: powershell + run: | + cmake --preset windows-vs18 ` + -B build-models-from-source ` + -DFLM_BUILD_GEMMA4E=ON + + - name: Build models from source + shell: powershell + run: | + cmake --build build-models-from-source --config Release --parallel 2 + + - name: Prepare models-from-source package files + shell: powershell + run: | + $ErrorActionPreference = "Stop" + $dist = "dist-models-from-source" + + if (Test-Path $dist) { + Remove-Item $dist -Recurse -Force + } + + New-Item -ItemType Directory -Force $dist | Out-Null + Copy-Item -Path .\build-models-from-source\flm.exe -Destination $dist\ -Force + Copy-Item -Path .\xclbins -Destination $dist\ -Recurse -Force + Copy-Item -Path .\model_list.json -Destination $dist\ -Force + Copy-Item -Path .\model_info.json -Destination $dist\ -Force + Copy-Item -Path .\lib\xrt\*.dll -Destination $dist\ -Force + Copy-Item -Path .\lib\*.dll -Destination $dist\ -Force + Copy-Item -Path .\build-models-from-source\engines\*.dll -Destination $dist\ -Force + + $sourceEngine = Get-FileHash .\build-models-from-source\engines\gemma4e_npu.dll -Algorithm SHA256 + $prebuiltEngine = Get-FileHash .\lib\xrt\gemma4e_npu.dll -Algorithm SHA256 + $stagedEngine = Get-FileHash $dist\gemma4e_npu.dll -Algorithm SHA256 + if ($sourceEngine.Hash -eq $prebuiltEngine.Hash) { + throw "The source-built Gemma 4 DLL matches the prebuilt DLL." + } + if ($sourceEngine.Hash -ne $stagedEngine.Hash) { + throw "The staged Gemma 4 DLL does not match the source-built DLL." + } + + - name: Upload models-from-source build artifact + uses: actions/upload-artifact@v4 + with: + name: models-from-source-fastflowlm-windows-${{ github.sha }} + path: src/dist-models-from-source/ + if-no-files-found: error + build-windows-msi: runs-on: [self-hosted, windows, x64, msi] diff --git a/README.md b/README.md index 505aff7ef..78d555173 100644 --- a/README.md +++ b/README.md @@ -170,12 +170,23 @@ More details on the exact procedure, with dependencies to be installed, for Linu This will configure the build to install to `/opt/fastflowlm`. + To build the Gemma 4 engine from source instead of using the prebuilt engine: + + ```bash + cmake --preset linux-default -DFLM_BUILD_GEMMA4E=ON + cmake --build build + ``` + + CMake caches this option. To switch back to the prebuilt engine, reconfigure with `-DFLM_BUILD_GEMMA4E=OFF` or use a fresh build directory. + - **For Windows (in a developer command prompt):** ```bash cmake --preset windows-default ``` + Add `-DFLM_BUILD_GEMMA4E=ON` to build the Gemma 4 engine from source. + 3. **Build the project:** ```bash diff --git a/docs/docs/install_lin.md b/docs/docs/install_lin.md index 397c1c622..62ae23706 100644 --- a/docs/docs/install_lin.md +++ b/docs/docs/install_lin.md @@ -179,6 +179,26 @@ If `flm validate` passes but `flm run` fails with `No such device with index '0' cmake --install --preset linux-default ``` +#### Optional Gemma 4 Source Build + +By default, FLM links the prebuilt Gemma 4 engine from `src/lib//`. To build the engine from `src/detail/` on Linux: + +```sh +cd src +cmake --preset linux-default -DFLM_BUILD_GEMMA4E=ON +cmake --build build +``` + +CMake caches `FLM_BUILD_GEMMA4E`. To switch back to the prebuilt engine, reconfigure with `-DFLM_BUILD_GEMMA4E=OFF` or use a fresh build directory. + +| Option | Default | Effect | +|---|---:|---| +| `FLM_BUILD_GEMMA4E` | `OFF` | Build `gemma4e_npu` from `src/detail/` instead of using the prebuilt engine. | +| `FLM_ENGINE_NATIVE_ARCH` | `OFF` | Add `-march=native`. This produces host-specific binaries that should not be redistributed. | +| `FLM_ENGINE_VERBOSE` | `0` | Set the logging level for the source-built engine. | +| `FLM_ENGINE_DEBUG_LEVEL` | `0` | Set the debug level for the source-built engine. | +| `FLM_OVERRIDE_FLAGS` | empty | Add compile flags for operator overrides. See `src/include/flm_override.hpp`. | + #### Advanced Build Options **Static Build with Bundled XRT/XDNA** diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index 48910cc06..77da38ed6 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -26,7 +26,7 @@ set(CMAKE_ERROR_DEPRECATED ON) # Set output directories -set(CMAKE_RUNTIME_OUTPUT_DIRECTORY ${CMAKE_SOURCE_DIR}/build/) +set(CMAKE_RUNTIME_OUTPUT_DIRECTORY ${CMAKE_BINARY_DIR}) # Force output directories to be absolute @@ -110,6 +110,12 @@ set(FLM_ENGINE_LIB_DIR "${CMAKE_SOURCE_DIR}/lib/${FLM_RUNTIME_NAME}") # The flash engine libs (qwen3vl_flash, gemma4e_flash) are always shipped in # src/lib/${FLM_RUNTIME_NAME} for every backend, so their model families are # always compiled in -- no probe/gate/macro needed. +# Gemma4e can be built from source here; the other engines stay prebuilt. +option(FLM_BUILD_GEMMA4E "Build the Gemma4e engine from src/detail instead of using the prebuilt" OFF) +option(FLM_ENGINE_NATIVE_ARCH "Build the from-source engine with -march=native (host-specific, not redistributable)" OFF) +set(FLM_ENGINE_VERBOSE 0 CACHE STRING "VERBOSE level for the from-source engine") +set(FLM_ENGINE_DEBUG_LEVEL 0 CACHE STRING "DEBUG_LEVEL for the from-source engine") +set(FLM_OVERRIDE_FLAGS "" CACHE STRING "Extra compile flags for the from-source engine (see include/flm_override.hpp)") # ——————————————————————————————————————————————— # NPU runtime discovery. @@ -487,6 +493,12 @@ if(MSVC) target_link_libraries(flm PUBLIC ${STATIC_LIBS}) endif() +# Defines the gemma4e_npu target, which the link list below then resolves to +# instead of the prebuilt of the same name. +if(FLM_BUILD_GEMMA4E) + add_subdirectory(detail) +endif() + set(FLM_ENGINE_LINK_LIBS q4_npu_eXpress llama_npu @@ -567,6 +579,17 @@ else() endif() endif() +# The source-built engine sits in the build tree rather than lib/, and +# has to precede it: a prebuilt of the same name is still sitting there. +if(FLM_BUILD_GEMMA4E AND NOT WIN32) + get_target_property(_flm_build_rpath flm BUILD_RPATH) + if(NOT _flm_build_rpath) + set(_flm_build_rpath "") + endif() + set_target_properties(flm PROPERTIES + BUILD_RPATH "${CMAKE_BINARY_DIR}/engines;${_flm_build_rpath}") +endif() + if(WIN32 AND VCPKG_TOOLCHAIN) # Local/managed vcpkg: link the imported targets from the CONFIG packages # found above (versioned import-lib names resolved automatically). @@ -723,7 +746,15 @@ else() endif() file(GLOB so_libs "${FLM_ENGINE_LIB_DIR}/*.so*") - install(FILES ${so_libs} DESTINATION "${FLM_ENGINE_LIB_DESTINATION}") + # Installed from the glob unless a source build supersedes one of them. The + # name still has to reach _flm_engine_names below, so filter a copy. + set(_flm_prebuilt_libs ${so_libs}) + if(FLM_BUILD_GEMMA4E) + list(FILTER _flm_prebuilt_libs EXCLUDE REGEX "/libgemma4e_npu\\.so[^/]*$") + set_target_properties(gemma4e_npu PROPERTIES INSTALL_RPATH "${FLM_ENGINE_INSTALL_RPATH}") + install(TARGETS gemma4e_npu LIBRARY DESTINATION "${FLM_ENGINE_LIB_DESTINATION}") + endif() + install(FILES ${_flm_prebuilt_libs} DESTINATION "${FLM_ENGINE_LIB_DESTINATION}") set_target_properties(flm PROPERTIES INSTALL_RPATH "${FLM_FLM_INSTALL_RPATH}") # Engine .so file names, used below to keep them out of the flm dependency diff --git a/src/detail/CMakeLists.txt b/src/detail/CMakeLists.txt new file mode 100644 index 000000000..16870c5b5 --- /dev/null +++ b/src/detail/CMakeLists.txt @@ -0,0 +1,141 @@ +# Builds the Gemma4e engine from source, in place of the prebuilt +# lib//libgemma4e_npu.so that ships for every other model. + +set(GEMMA4E_ENGINE_DIR "${CMAKE_CURRENT_SOURCE_DIR}") + +file(GLOB GEMMA4E_SOURCES "${GEMMA4E_ENGINE_DIR}/gemma4e_npu/*.cpp") +# Every engine compiles its own copy of these four; they also ship as separate +# libraries for flm itself to link. +file(GLOB GEMMA4E_SHARED_SOURCES + "${GEMMA4E_ENGINE_DIR}/dequant/*.cpp" + "${GEMMA4E_ENGINE_DIR}/gemm/*.cpp" + "${GEMMA4E_ENGINE_DIR}/lm_head/*.cpp" +) +list(APPEND GEMMA4E_SHARED_SOURCES "${GEMMA4E_ENGINE_DIR}/vision_common/norm.cpp") + +add_library(gemma4e_npu SHARED ${GEMMA4E_SOURCES} ${GEMMA4E_SHARED_SOURCES}) + +if(MSVC) + set_target_properties(gemma4e_npu PROPERTIES + MSVC_RUNTIME_LIBRARY "MultiThreaded$<$:Debug>" + WINDOWS_EXPORT_ALL_SYMBOLS ON + ) +endif() + +target_include_directories(gemma4e_npu BEFORE PRIVATE + "${CMAKE_CURRENT_SOURCE_DIR}/../include" +) +if(WIN32 AND NOT VCPKG_TOOLCHAIN) + target_include_directories(gemma4e_npu PRIVATE + C:/dev/boost_1_88_0 + C:/dev/vcpkg/installed/x64-windows/include/ + ) +endif() +# Whichever XRT the parent discovered: pkg-config, a portable build's from-source +# fetch, or the /opt/xilinx fallback. HRX carries its own headers on hrx::hrx. +if(NOT FLM_USE_HRX) + if(NOT WIN32 AND XRT_FOUND) + target_include_directories(gemma4e_npu PRIVATE ${XRT_INCLUDE_DIRS}) + else() + target_include_directories(gemma4e_npu PRIVATE ${XRT_INCLUDE_DIR}) + endif() +endif() + +target_compile_features(gemma4e_npu PRIVATE cxx_std_20) +target_compile_definitions(gemma4e_npu PRIVATE + VERBOSE=${FLM_ENGINE_VERBOSE} + DEBUG_LEVEL=${FLM_ENGINE_DEBUG_LEVEL} +) +if(FLM_USE_HRX) + target_compile_definitions(gemma4e_npu PRIVATE FLM_USE_HRX=1) +endif() +if(WIN32) + target_compile_definitions(gemma4e_npu PRIVATE + WIN32_LEAN_AND_MEAN + NOMINMAX + _CRT_SECURE_NO_WARNINGS + _CRT_NONSTDC_NO_DEPRECATE + __WINDOWS__ + ) +endif() + +if(NOT MSVC) + target_compile_options(gemma4e_npu PRIVATE + -O3 -Wall -fmax-errors=1 + -mavx512f -mavx512vl -mavx512bw -mavx512dq -mfma + -ffast-math + ) + # -march=native bakes the build host's ISA into the library, which a package + # then carries onto machines that fault on it. The internal Makefile uses it. + if(FLM_ENGINE_NATIVE_ARCH) + target_compile_options(gemma4e_npu PRIVATE -march=native) + endif() +endif() + +# Point FLM_OVERRIDES at a header that redefines FLM_OVERRIDE to dispatch an +# operator elsewhere; see include/flm_override.hpp. Empty leaves the build +# byte-identical to an unannotated one. +if(FLM_OVERRIDE_FLAGS) + separate_arguments(_gemma4e_override_flags NATIVE_COMMAND "${FLM_OVERRIDE_FLAGS}") + target_compile_options(gemma4e_npu PRIVATE ${_gemma4e_override_flags}) +endif() + +find_package(OpenMP REQUIRED) +target_link_libraries(gemma4e_npu PRIVATE OpenMP::OpenMP_CXX) + +target_link_directories(gemma4e_npu PRIVATE "${FLM_ENGINE_LIB_DIR}") +if(FLM_USE_HRX) + target_link_libraries(gemma4e_npu PRIVATE hrx::hrx) +else() + if(NOT WIN32 AND XRT_FOUND) + target_link_directories(gemma4e_npu PRIVATE ${XRT_LIBRARY_DIRS}) + else() + target_link_directories(gemma4e_npu PRIVATE ${XRT_LIB_DIR}) + endif() + if(XRT_BUILT_FROM_SOURCE AND TARGET xrt_coreutil) + target_link_libraries(gemma4e_npu PRIVATE xrt_coreutil) + elseif(NOT WIN32 AND XRT_FOUND) + target_link_libraries(gemma4e_npu PRIVATE ${XRT_LIBRARIES}) + else() + target_link_libraries(gemma4e_npu PRIVATE xrt_coreutil) + endif() + if(WIN32) + target_link_libraries(gemma4e_npu PRIVATE aiebu_static) + else() + target_link_libraries(gemma4e_npu PRIVATE aiebu) + endif() +endif() + +if(WIN32) + # MSVC resolves every symbol when it links a DLL. q4nx and mha stay prebuilt, + # so name them; their import libs sit in FLM_ENGINE_LIB_DIR. + target_link_libraries(gemma4e_npu PRIVATE q4_npu_eXpress mha) +endif() + +if(NOT WIN32) + # Each engine .so carries its own copy of the shared runtime (npu_sequence, + # npu_app_manager, weight_desc_t, buffer, ...) at default visibility. flm + # loads every engine at once, so without this flag ELF interposition binds + # references -- including calls from inside this .so -- to whichever engine + # comes first in DT_NEEDED, running another model's implementation. + # Functions only: typeinfo and vtables stay interposable so dynamic_cast + # across the boundary keeps working. + target_link_options(gemma4e_npu PRIVATE -Wl,-Bsymbolic-functions) + # q4nx and mha stay prebuilt; flm resolves them when it links the engine. + target_link_options(gemma4e_npu PRIVATE -Wl,--allow-shlib-undefined) +endif() + +# Built into the build tree, leaving lib//libgemma4e_npu.so untouched +# so FLM_BUILD_GEMMA4E=OFF still has a prebuilt to fall back on. The parent adds +# this directory to flm's BUILD_RPATH and installs this target in place of the +# prebuilt. +set_target_properties(gemma4e_npu PROPERTIES + LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/engines" + RUNTIME_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}/engines" +) +foreach(_cfg RELEASE DEBUG RELWITHDEBINFO MINSIZEREL) + set_target_properties(gemma4e_npu PROPERTIES + LIBRARY_OUTPUT_DIRECTORY_${_cfg} "${CMAKE_BINARY_DIR}/engines" + RUNTIME_OUTPUT_DIRECTORY_${_cfg} "${CMAKE_BINARY_DIR}/engines" + ) +endforeach() diff --git a/src/detail/dequant/dequant.cpp b/src/detail/dequant/dequant.cpp new file mode 100644 index 000000000..1a2d51c2d --- /dev/null +++ b/src/detail/dequant/dequant.cpp @@ -0,0 +1,330 @@ +#include "dequant_detail.hpp" + +// constructors +Dequant::Impl::Impl(LM_Config& config) : config(config){ + // Initialize any dequant-specific configuration here +} + +Dequant::Impl::~Impl() = default; + +// methods +/// @brief generate the dequant sequence +/// @param seq: the sequence +/// @param D_in: input dimension of the projection weight +/// @param D_out: output dimension of the projection weight +/// @param weight_offset: the weight offset in byte +/// @param mode: dequant output mode +void Dequant::Impl::generate_dequant_q80_packed_in_q4nx_seq(npu_sequence* seq_ptr, const u32 D_in, const u32 D_out, const u32 weight_offset, dequant_output_mode_t output_mode){ + std::cout << "generate_dequant_q80_packed_in_q4nx_seq, D_in: " << D_in << ", D_out: " << D_out << ", weight_offset: " << weight_offset << std::endl; + if (D_in % k_tile_q4 != 0) { + std::cerr << "D_in % k_tile_q4 != 0" << std::endl; + exit(1); + } + + int bd_wait_counter[8] = {0, 0, 0, 0, 0, 0, 0, 0}; + // although each data block is in mxk block, but the data block could be reorder in col-stride on block view + /* + For example, quant_block_col_stride = 2 means + + //This is the logical view of the data block, each block of m_tile_q4 x k_tile_q4 + [block0, block1, ...... blockD, + blockD+1, blockD+2, ...... + ] + + But in memory order, the data block is arrange as block0, blockD+1, block1, blockD+2 .... + + */ + + if(D_in % desired_k_dequant != 0){ + std::cerr << "D_in % desired_k_dequant != 0" << std::endl; + exit(1); + } + + const uint32_t blocks_per_row = D_in / k_tile_q4 * 2; + std::cout << "blocks_per_row: " << blocks_per_row << std::endl; + + if(D_out % desired_m_dequant != 0 ){ + std::cerr << "D_out % desired_m_dequant != 0" << std::endl; + exit(1); + } + + const int quant_in_per_column = (desired_m_dequant / m_tile_q4) * blocks_per_row * block_size_in_byte_q4_1; + const int total_column_rounds = D_out / (desired_m_dequant); + + const int row_per_round = desired_m_dequant * total_cols; + // down rounds, go though D_out + const int down_rounds = (D_out + row_per_round - 1) / row_per_round; + + npu_sequence& seq = *seq_ptr; + seq.clear_cmds(); + + uint32_t input_offset = weight_offset; + + if(output_mode == dequant_output_mode_t::GATE_MATRIX){ + input_offset += (gate_up_m_interleave_size / m_tile_q4) * blocks_per_row * block_size_in_byte_q4_1; + } + uint32_t gate_up_interleave_counter= 0; + + // first, the dequant of down + for(int i = 0; i < down_rounds; i++){ + for(int col = 0; col < 8; col++){ + uint32_t bd_offset = (i % 2) * 8; + uint32_t round_offset = i * 8 + col; + if(round_offset < total_column_rounds){ + seq.npu_dma_memcpy_nd( + sizeof(char), + qw_in_arg_idx, + MM2S, + IT[col], + (npu_bd_id)(0+bd_offset), + it_channel_0, + {0, 0, 0, input_offset}, + //NOTE: this for now only work if desired_m_dequant == quant_block_col_stride*m_tile_q4 + { + blocks_per_row, + (desired_m_dequant / m_tile_q4) / quant_block_col_stride, + quant_block_col_stride * block_size_in_byte_q4_1 / 512, + 512 + }, + { + quant_block_col_stride * block_size_in_byte_q4_1, + quant_block_col_stride * block_size_in_byte_q4_1 * blocks_per_row, + 512, + 1 + }, + -1 ,0, false + ); + + if(output_mode == dequant_output_mode_t::NORMAL_DEQUANT){ + std::cout << "Use normal output!" << std::endl; + input_offset += quant_in_per_column; + } + else{ + gate_up_interleave_counter++; + input_offset += quant_in_per_column; + if(gate_up_interleave_counter == (gate_up_m_interleave_size / desired_m_dequant) ){ + gate_up_interleave_counter = 0; + input_offset += (gate_up_m_interleave_size / m_tile_q4) * blocks_per_row * block_size_in_byte_q4_1; + } + } + + // Each port receive 2*Q4NX_ROWx D_Q4NX_BLOCK_PER_ROW*Q4NX_COL + uint32_t output_offset_0 = round_offset * desired_m_dequant * D_in; + + seq.npu_dma_memcpy_nd( + sizeof(uint16_t),//bf16 outpout + w_out_arg_idx, + S2MM, + IT[col], + (npu_bd_id)(1+bd_offset), + it_channel_0, + {0, 0, 0, output_offset_0}, + { + (uint32_t)D_in/desired_k_dequant, + desired_k_dequant/k_tile_q4, + desired_m_dequant, + k_tile_q4 + }, + { + desired_m_dequant * desired_k_dequant, + k_tile_q4, + desired_k_dequant, + 1 + }, + -1, 0, true, + aggressive_cache + ); + bd_wait_counter[col]++; + } + } + // note: for now + for(int col = 0; col < 8; col++){ + if(bd_wait_counter[col] == 2){ + seq.npu_dma_wait(IT[col], S2MM, it_channel_0); + bd_wait_counter[col]--; + } + } + } + + for(int col = 0; col < 8; col++){ + while(bd_wait_counter[col] != 0){ + seq.npu_dma_wait(IT[col], S2MM, it_channel_0); + bd_wait_counter[col]--; + } + } + seq.cmds2seq(); +} + +/// @brief generate the dequant sequence +/// @param seq: the sequence +/// @param D_in: input dimension of the projection weight +/// @param D_out: output dimension of the projection weight +/// @param weight_offset: the weight offset in byte +/// @param mode: dequant output mode +void Dequant::Impl::generate_dequant_q4_1_seq(npu_sequence* seq_ptr, const u32 D_in, const u32 D_out, const u32 weight_offset, dequant_output_mode_t output_mode){ + if (D_in % k_tile_q4 != 0) { + std::cerr << "D_in % k_tile_q4 != 0" << std::endl; + exit(1); + } + + int bd_wait_counter[8] = {0, 0, 0, 0, 0, 0, 0, 0}; + // although each data block is in mxk block, but the data block could be reorder in col-stride on block view + /* + For example, quant_block_col_stride = 2 means + + //This is the logical view of the data block, each block of m_tile_q4 x k_tile_q4 + [block0, block1, ...... blockD, + blockD+1, blockD+2, ...... + ] + + But in memory order, the data block is arrange as block0, blockD+1, block1, blockD+2 .... + + */ + + if(D_in % desired_k_dequant != 0){ + std::cerr << "D_in % desired_k_dequant != 0" << std::endl; + exit(1); + } + + const uint32_t blocks_per_row = D_in / k_tile_q4; + + if(D_out % desired_m_dequant != 0 ){ + std::cerr << "D_out % desired_m_dequant != 0" << std::endl; + exit(1); + } + + const int quant_in_per_column = (desired_m_dequant / m_tile_q4) * blocks_per_row * block_size_in_byte_q4_1; + const int total_column_rounds = D_out / (desired_m_dequant); + + const int row_per_round = desired_m_dequant * total_cols; + // down rounds, go though D_out + const int down_rounds = (D_out + row_per_round - 1) / row_per_round; + + npu_sequence& seq = *seq_ptr; + seq.clear_cmds(); + + uint32_t input_offset = weight_offset; + + if(output_mode == dequant_output_mode_t::GATE_MATRIX){ + input_offset += (gate_up_m_interleave_size / m_tile_q4) * blocks_per_row * block_size_in_byte_q4_1; + } + uint32_t gate_up_interleave_counter= 0; + + // first, the dequant of down + for(int i = 0; i < down_rounds; i++){ + for(int col = 0; col < 8; col++){ + uint32_t bd_offset = (i % 2) * 8; + uint32_t round_offset = i * 8 + col; + if(round_offset < total_column_rounds){ + + seq.npu_dma_memcpy_nd( + sizeof(char), + qw_in_arg_idx, + MM2S, + IT[col], + (npu_bd_id)(0+bd_offset), + it_channel_0, + {0, 0, 0, input_offset}, + //NOTE: this for now only work if desired_m_dequant == quant_block_col_stride*m_tile_q4 + { + blocks_per_row, + (desired_m_dequant / m_tile_q4) / quant_block_col_stride, + quant_block_col_stride * block_size_in_byte_q4_1 / 512, + 512 + }, + { + quant_block_col_stride * block_size_in_byte_q4_1, + quant_block_col_stride * block_size_in_byte_q4_1 * blocks_per_row, + 512, + 1 + }, + -1 ,0, false + ); + + if(output_mode == dequant_output_mode_t::NORMAL_DEQUANT){ + input_offset += quant_in_per_column; + } + else{ + gate_up_interleave_counter++; + input_offset += quant_in_per_column; + if(gate_up_interleave_counter == (gate_up_m_interleave_size / desired_m_dequant) ){ + gate_up_interleave_counter = 0; + input_offset += (gate_up_m_interleave_size / m_tile_q4) * blocks_per_row * block_size_in_byte_q4_1; + } + } + + // Each port receive 2*Q4NX_ROWx D_Q4NX_BLOCK_PER_ROW*Q4NX_COL + uint32_t output_offset_0 = round_offset * desired_m_dequant * D_in; + + seq.npu_dma_memcpy_nd( + sizeof(uint16_t),//bf16 outpout + w_out_arg_idx, + S2MM, + IT[col], + (npu_bd_id)(1+bd_offset), + it_channel_0, + {0, 0, 0, output_offset_0}, + { + (uint32_t)D_in/desired_k_dequant, + desired_k_dequant/k_tile_q4, + desired_m_dequant, + k_tile_q4 + }, + { + desired_m_dequant * desired_k_dequant, + k_tile_q4, + desired_k_dequant, + 1 + }, + -1, 0, true, + aggressive_cache + ); + bd_wait_counter[col]++; + } + } + // note: for now + for(int col = 0; col < 8; col++){ + if(bd_wait_counter[col] == 2){ + seq.npu_dma_wait(IT[col], S2MM, it_channel_0); + bd_wait_counter[col]--; + } + } + } + + for(int col = 0; col < 8; col++){ + while(bd_wait_counter[col] != 0){ + seq.npu_dma_wait(IT[col], S2MM, it_channel_0); + bd_wait_counter[col]--; + } + } + seq.cmds2seq(); +} + +// wrappers +Dequant::Dequant(LM_Config& config) : _impl(new Impl(config)){} +Dequant::~Dequant(){ + delete _impl; +} +void Dequant::generate_dequant_q80_packed_in_q4nx_seq(npu_sequence* seq, const u32 D_in, const u32 D_out, const u32 weight_offset, int mode){ + _impl->generate_dequant_q80_packed_in_q4nx_seq(seq, D_in, D_out, weight_offset, (Dequant::Impl::dequant_output_mode_t)mode); +} + +void Dequant::generate_dequant_q4_1_seq(npu_sequence* seq, const u32 D_in, const u32 D_out, const u32 weight_offset, int mode){ + _impl->generate_dequant_q4_1_seq(seq, D_in, D_out, weight_offset, (Dequant::Impl::dequant_output_mode_t)mode); +} + +void Dequant::reorder_cpy( + u8 *dst, buffer &src, + quant_block_t quant_block_type, + const int quant_matrix_row, + const int quant_matrix_col, + const int vertical_blocks , + const int vetrical_block_interleave_byte_size + +){ + + _impl->reorder_cpy( + dst, src, quant_block_type, quant_matrix_row, quant_matrix_col, + vertical_blocks, vetrical_block_interleave_byte_size + ); +} diff --git a/src/detail/dequant/dequant_detail.hpp b/src/detail/dequant/dequant_detail.hpp new file mode 100644 index 000000000..cdef9f1c0 --- /dev/null +++ b/src/detail/dequant/dequant_detail.hpp @@ -0,0 +1,175 @@ +#pragma once +#include "modules/dequant.hpp" + +struct Dequant::Impl{ +private: + static constexpr npu_tiles IT[] = {IT0, IT1, IT2, IT3, IT4, IT5, IT6, IT7}; + + static constexpr u32 total_cols = 8; + static constexpr u32 total_rows = 4; + + static constexpr int w_out_arg_idx = 0; + static constexpr int qw_in_arg_idx = 1; + + static constexpr int m_tile_q4 = 32; + static constexpr int k_tile_q4 = 256; + + static constexpr uint32_t block_size_in_byte_q4_0 = ((m_tile_q4 * k_tile_q4 * 4.5) / 8.0); + static constexpr uint32_t block_size_in_byte_q4_1 = m_tile_q4 * k_tile_q4 * 5 / 8; + + static constexpr int m_tile_q8 = 32; + static constexpr int k_tile_q8 = 256; + + static constexpr uint32_t block_size_in_byte_q8_0 = (m_tile_q8 * k_tile_q8 * (8.5) )/8.0; + static constexpr uint32_t block_size_in_byte_q8_1 = (m_tile_q8 * k_tile_q8 * (9) )/8.0; + static constexpr int quant_block_col_stride = 2; + static constexpr int quant_block_interleave_byte_size = 512; + + static constexpr int desired_k_dequant = 512; + static constexpr int desired_m_dequant = 128; + + static constexpr int glu_slice = 1024; + static constexpr int gate_up_m_interleave_size = glu_slice / 2; + +public: + /// @brief dequant output mode, as the quantized weight of UP and GATE are interleaved in memory, now we want to seperate them. + /// @note NORMAL_DEQUANT: normal dequant output + /// @note UP_MATRIX: up projection matrix output + /// @note GATE_MATRIX: gate projection matrix output + typedef enum: int{ + NORMAL_DEQUANT = 0, + UP_MATRIX = 1, + GATE_MATRIX = 2 + } dequant_output_mode_t; + + Impl(){} + Impl(LM_Config& config); + ~Impl(); + /// @brief generate the dequant sequence + /// @param seq: the sequence + /// @param D_in: input dimension of the projection weight + /// @param D_out: output dimension of the projection weight + /// @param weight_offset: the weight offset in byte + /// @param mode: dequant output mode + void generate_dequant_q4_1_seq(npu_sequence* seq, const u32 D_in, const u32 D_out, const u32 weight_offset, dequant_output_mode_t mode); + void generate_dequant_q80_packed_in_q4nx_seq(npu_sequence* seq, const u32 D_in, const u32 D_out, const u32 weight_offset, dequant_output_mode_t mode); + LM_Config config; + + void reorder_cpy(u8 *dst, buffer &src, + Dequant::quant_block_t quant_block_type, + const int quant_matrix_row, + const int quant_matrix_col, + const int vertical_blocks, + const int vetrical_block_interleave_byte_size) + { + int a_block_size = 0; + int block_col_size = 0; + int block_row_size = 0; + switch(quant_block_type){ + case Q4_1: + a_block_size = block_size_in_byte_q4_1; + block_col_size = k_tile_q4; + block_row_size = m_tile_q4; + break; + case Q8_0: + a_block_size= block_size_in_byte_q8_0; + block_col_size = k_tile_q8; + block_row_size = m_tile_q8; + break; + default: + std::cerr << "Unsupport type for now"; + exit(-1); + break; + } + + assert( quant_matrix_col % block_col_size== 0); + const int blocks_per_row = quant_matrix_col / block_col_size; + + assert(quant_matrix_row%(block_row_size* vertical_blocks) == 0 ); + + const int rows = src.size() / a_block_size / blocks_per_row; + + u8 *dst_ptr = dst; + std::vector src_ptr(vertical_blocks); + for (int i = 0; i < vertical_blocks; i++) + { + src_ptr[i] = src.data() + i * a_block_size * blocks_per_row; + } + for (int r = 0; r < rows; r += vertical_blocks) + { + for (int c = 0; c < blocks_per_row; c++) + { + for (int i = 0; i < vertical_blocks; i++) + { + memcpy(dst_ptr, src_ptr[i], a_block_size); + dst_ptr += a_block_size; + src_ptr[i] += a_block_size; + } + } + for (int i = 0; i < vertical_blocks; i++) + { + src_ptr[i] += (vertical_blocks - 1) * a_block_size * blocks_per_row; + if (src_ptr[i] + a_block_size * blocks_per_row > src.end()) + { + src_ptr[i] = src.data(); // useless padding + } + } + } + + // At this step, the blocks are now reorder with vertical blocks + + // For example + // IF previous are row-block order + /* + [A, B, C + D, E, F] + Where blocks are ordered as A, B, C, D, E, F + + With the vertical_blocks =2, + blocks are reorder as A, D, B, E, C, F + + */ + + if(vetrical_block_interleave_byte_size <=0){ + return ; // no need this step + } + // Apply vetrical_block_interleave_byte_size reorder + + // From example above, now A, D blocks are continousy in memory at block level + + // However, we want to do a byte-block level mixing + /** + For example, If A, D block are Block size of 2K and vetrical_block_interleave_byte_size = 1024 + + In memory, the data are layout as A(1-1204) A(1025-2048), D(1-1024), D(1025-2048) + + After the block level reorder, we have + A(1-1204), D(1-1024), A(1025-2048), D(1025-2048) + + */ + size_t num_data_block = (quant_matrix_row/block_row_size) * (quant_matrix_col/block_col_size); + std::vector temp_buffer(vertical_blocks*a_block_size ); + + size_t num_byte_data_chunk = (a_block_size) / vetrical_block_interleave_byte_size; + assert(a_block_size % vetrical_block_interleave_byte_size == 0); + + for(int i = 0; i < num_data_block; i+= vertical_blocks){ + uint8_t* cur_ptr = dst + i*a_block_size; + memcpy( temp_buffer.data(), cur_ptr, temp_buffer.size() ); + + uint8_t* chunk_dst_ptr = cur_ptr; + + for(int byte_chunk_idx = 0; byte_chunk_idx +#include // Required for std::max +#include + +// // constructors +Gemm::Impl::Impl() { + // mapping of valid shimtile index for A -> index offset + valid_A_MT_shimtile_index[0] = 0; + valid_A_MT_shimtile_index[2] = 1; + valid_A_MT_shimtile_index[4] = 2; + valid_A_MT_shimtile_index[6] = 3; +} + +Gemm::Impl::~Impl() = default; + +/// \brief Generate the sequence +/// \param seq the npu sequence +/// \param M the M dimension +/// \param K the K dimension +/// \param N the N dimension +/// \param weight_offset the weight offset +/// \param ADD_BIAS whether to add bias +/// \param OUTPUT_MODE the output activation mode +/// \param bias_offset the bias offset +void Gemm::Impl::generate_seq( + npu_sequence *seq, + uint32_t M, + uint32_t K, + uint32_t N, + const uint32_t weight_offset, + bool ADD_BIAS, + Activation_Type_t OUTPUT_MODE, + const uint32_t bias_offset, + uint32_t C_const_offset +){ + uint32_t A_const_offset = 0; + uint32_t B_const_offset = weight_offset; + + const int K_div_k = K/k; + + // some sanity checks + if (M % (m * total_rows) != 0) { + std::cerr << "GEMM M size not aligned with total npu rows"<< std::endl; + exit(1); + } + if (K % k != 0) { + std::cerr << "GEMM K size not aligned with k"<< std::endl; + exit(1); + } + if (N % n != 0) { + std::cerr << "GEMM N size not aligned with n"<< std::endl; + exit(1); + } + + seq->clear_cmds(); + + std::vector list_C_shim_queue; // int counter of how many DMA_Wait for C + std::vector list_A_shim_queue; // int counter of how many DMA_Wait for A + std::vector list_B_shim_queue; // int counter of how many DMA_Wait for B + for(size_t i = 0; i < shimtile_size; i++){ + list_C_shim_queue.push_back(0); + list_A_shim_queue.push_back(0); + list_B_shim_queue.push_back(0); + } + + // first, setup the rtp buffer and the rtp locks + for(size_t row_idx = 0; row_idx < total_rows; row_idx++){ + for(size_t col_idx = 0; col_idx< total_cols; col_idx++){ + auto CT_tile = get_tile(row_idx + 2, col_idx); + // set RTP value + seq->rtp_write(CT_tile, CT_rtp_address, K_div_k); + seq->rtp_write(CT_tile, CT_rtp_address + 4, M); + seq->rtp_write(CT_tile, CT_rtp_address + 8, N); + if(ADD_BIAS){ + seq->rtp_write( CT_tile, CT_rtp_address + 12, 1 ); + }else{ + seq->rtp_write( CT_tile, CT_rtp_address + 12, 0 ); + } + seq->rtp_write( CT_tile, CT_rtp_address + 16, OUTPUT_MODE ); // OUTPUT MODE + // set RTP lock, enable running + seq->rtp_write(CT_tile, CT_lock_address_base + 16 * CT_rtp_sync_lock_id, 1); // set lock to 1 + } + } + + generate_runtime_sequence( + seq, + A_const_offset, B_const_offset, C_const_offset, bias_offset, + M, N, K, + list_A_shim_queue, list_B_shim_queue, + list_C_shim_queue, + IS_B_ROW_MAJOR, ENABLE_AXI4, B_in_K_N_block_col_major_order, + ADD_BIAS + ); + + int max_C_remain = 0; + for( auto li: list_C_shim_queue){ + max_C_remain = std::max(max_C_remain, li); + } + + for(size_t k = 0; k< max_C_remain; k++){ + for(size_t shim_index = 0; shim_index < total_cols; shim_index++){ + + if(list_A_shim_queue.at(shim_index) > 0){ + seq->npu_dma_wait( + shim_tiles[shim_index], + MM2S, + it_channel_0 + ); + list_A_shim_queue.at(shim_index)--; + } + if(list_B_shim_queue.at(shim_index) > 0){ + seq->npu_dma_wait( + shim_tiles[shim_index], + MM2S, + it_channel_1 + ); + list_B_shim_queue.at(shim_index)--; + } + if(list_C_shim_queue.at(shim_index) > 0){ + seq->npu_dma_wait( + shim_tiles[shim_index], + S2MM, + it_channel_0 + + ); + list_C_shim_queue.at(shim_index)--; + } + } + } + + seq->cmds2seq(); +} + +template +void Gemm::Impl::generate_runtime_sequence( + npu_sequence* seq, + uint32_t A_const_offset, uint32_t B_const_offset, uint32_t C_const_offset, uint32_t Bias_const_offset, + uint32_t M_size, uint32_t N_size, uint32_t K_size, + std::vector &list_A_shim_queue, + std::vector &list_B_shim_queue, + std::vector &list_C_shim_queue, + bool IS_B_ROW_MAJOR, + bool ENABLE_AXI4, + bool B_in_K_N_block_col_major_order, + bool ADD_BIAS +){ + uint32_t M_div_num_row_m = M_size/(m * total_rows); + uint32_t N_div_num_col_n = N_size/(n * total_cols); + + uint32_t N_div_num_col_n_remainder_blocks = (N_size % (n * total_cols)) / n; + + std::vector list_A_BD_pingpong_flag(shimtile_size, 0); + std::vector list_B_BD_pingpong_flag(shimtile_size, 0); + std::vector list_C_BD_pingpong_flag(shimtile_size, 0); + + uint32_t col_block_range = N_div_num_col_n; + if (N_div_num_col_n_remainder_blocks!= 0){ + col_block_range += 1; + } + + for(uint32_t mega_block_col_idx = 0; mega_block_col_idx < col_block_range; mega_block_col_idx++){ + for (uint32_t mega_block_row_idx = 0; mega_block_row_idx < M_div_num_row_m; mega_block_row_idx++) { + for (uint32_t shim_index = 0; shim_index < shimtile_size; shim_index++) { + bool SEND_ADD_BIAS = false; + + if (mega_block_row_idx == 0 &&ADD_BIAS){ + SEND_ADD_BIAS = true; + } + + if (N_div_num_col_n_remainder_blocks!= 0 && mega_block_col_idx == N_div_num_col_n){ + if (shim_index < N_div_num_col_n_remainder_blocks){ + generate_shimtile_sequence_per_k_block( + seq, + shim_index, + mega_block_row_idx, mega_block_col_idx, + M_size, K_size, N_size, + A_const_offset, B_const_offset, C_const_offset, Bias_const_offset, + list_A_shim_queue, list_B_shim_queue, list_C_shim_queue, + list_A_BD_pingpong_flag, list_B_BD_pingpong_flag, list_C_BD_pingpong_flag, + IS_B_ROW_MAJOR, ENABLE_AXI4, + B_in_K_N_block_col_major_order, + true, + ADD_BIAS, SEND_ADD_BIAS + ); + } + + else{ + generate_shimtile_sequence_per_k_block( + seq, + shim_index, + mega_block_row_idx, mega_block_col_idx, + M_size, K_size, N_size, + A_const_offset, B_const_offset, C_const_offset, Bias_const_offset, + list_A_shim_queue, list_B_shim_queue, list_C_shim_queue, + list_A_BD_pingpong_flag, list_B_BD_pingpong_flag, list_C_BD_pingpong_flag, + IS_B_ROW_MAJOR, ENABLE_AXI4, + B_in_K_N_block_col_major_order, + false, + ADD_BIAS, SEND_ADD_BIAS + ); + } + } + else{ + generate_shimtile_sequence_per_k_block( + seq, + shim_index, + mega_block_row_idx, mega_block_col_idx, + M_size, K_size, N_size, + A_const_offset, B_const_offset, C_const_offset, Bias_const_offset, + list_A_shim_queue, list_B_shim_queue, list_C_shim_queue, + list_A_BD_pingpong_flag, list_B_BD_pingpong_flag, list_C_BD_pingpong_flag, + IS_B_ROW_MAJOR, ENABLE_AXI4, + B_in_K_N_block_col_major_order, + true, + ADD_BIAS, SEND_ADD_BIAS + ); + } + } + } + } +} + +template +void Gemm::Impl::generate_shimtile_sequence_per_k_block( + npu_sequence*seq, + uint32_t shim_index, + uint32_t mega_block_row_idx, uint32_t mega_block_col_idx, + uint32_t M_size, uint32_t K_size, uint32_t N_size, + uint32_t A_const_offset, uint32_t B_const_offset, uint32_t C_const_offset, uint32_t Bias_const_offset, + std::vector &list_A_shim_queue, std::vector &list_B_shim_queue, std::vector &list_C_shim_queue, + std::vector &list_A_bd_pingpong_flag, std::vector &list_B_bd_pingpong_flag, std::vector &list_C_bd_pingpong_flag, + bool IS_B_ROW_MAJOR, bool ENABLE_AXI4, + bool B_in_K_N_block_col_major_order, + bool VALID_COLUMN, + bool ADD_BIAS, bool SEND_BIAS +){ + + if(B_in_K_N_block_col_major_order){ + assert(IS_B_ROW_MAJOR == false); // on valid for B in col major order + if (IS_B_ROW_MAJOR){ + std::cerr << "Error: When B_in_K_N_block_col_major_order is set to true, IS_B_ROW_MAJOR cannot be true." << std::endl; + exit(1); + } + } + + // When B_in_K_N_block_col_major_order is set to true, it mean + // B is col-major order && + // B is rearrange into kxn blocks, where blocks are in col-major. Moreover, the data in each blocks is + // also in col-major order. + + // Basically,B as a col-major matrix goes through + // stride: [N_size/n,K_size/k ,n, k] + // offset: [K_size*n,k ,K_SIZE, 1] + + uint32_t K_div_k = K_size/k; + + npu_tiles cur_shimtile = shim_tiles[shim_index]; + + if (valid_A_MT_shimtile_index.contains(shim_index) && valid_A_MT_shimtile_index[shim_index] < total_rows){ + + if (list_A_shim_queue.at(shim_index) == 2) { + seq->npu_dma_wait( + cur_shimtile, MM2S, it_channel_0 + ); + list_A_shim_queue.at(shim_index)--; + } + + uint32_t A_offset = mega_block_row_idx * (total_rows * m) * K_size; + A_offset += valid_A_MT_shimtile_index[shim_index] * (m * K_size); + npu_bd_id A_bd_id; + if (list_A_bd_pingpong_flag.at(shim_index) ==0){ + A_bd_id = bd_0; + list_A_bd_pingpong_flag.at(shim_index) =1; + }else{ + A_bd_id = bd_1; + list_A_bd_pingpong_flag.at(shim_index) =0; + } + + seq->npu_dma_memcpy_nd( + sizeof(T_in), // bfloat16 + Arg_A, + MM2S, + cur_shimtile, + A_bd_id, + it_channel_0, + {0,0,0,A_offset+ A_const_offset}, + {1, K_div_k, m,k}, + {0, k, K_size, 1}, + -1, 0, true, + ENABLE_AXI4 ? aggressive_cache : normal_cache + ); + + list_A_shim_queue.at(shim_index)++; + } + + if((shim_index < total_cols) && VALID_COLUMN){ + if(SEND_BIAS){ + if (list_B_shim_queue[shim_index] == 2){ + + seq->npu_dma_wait( + cur_shimtile, MM2S, it_channel_1 + ); + list_B_shim_queue[shim_index] -= 1; + } + uint32_t _BIAS_DATA_OFFSET = Bias_const_offset + mega_block_col_idx * (total_cols * n) + shim_index *n; + seq->npu_dma_memcpy_nd( + sizeof(T_in), + Arg_Bias, + MM2S, + cur_shimtile, + npu_bd_id(bd_6), //reserved for sending bias + it_channel_1, + {0,0,0,_BIAS_DATA_OFFSET}, + {1, 1,1, k*n}, + {0, 0, 0, 1}, + -1, 0, true, + ENABLE_AXI4 ? aggressive_cache : normal_cache + ); + list_B_shim_queue[shim_index]++; + } + + npu_bd_id b_bd_id; + if (list_B_shim_queue[shim_index] == 2){ + seq->npu_dma_wait( + cur_shimtile, MM2S, it_channel_1 + ); + list_B_shim_queue[shim_index] -= 1; + } + + if (list_B_bd_pingpong_flag[shim_index] == 0) { + b_bd_id = bd_2; + list_B_bd_pingpong_flag[shim_index] = 1; + } else { + b_bd_id = bd_3; + list_B_bd_pingpong_flag[shim_index] = 0; + } + + if (IS_B_ROW_MAJOR){ + uint32_t B_offset = mega_block_col_idx* (total_cols) * n; + B_offset += shim_index * n; + seq->npu_dma_memcpy_nd( + sizeof(T_in), + Arg_B, + MM2S, + cur_shimtile, + b_bd_id, + it_channel_1, + {0,0,0,B_offset+ B_const_offset }, + {1, K_div_k, k, n}, + {0, k*N_size, N_size, 1}, + -1, 0, true, + ENABLE_AXI4 ? aggressive_cache : normal_cache + ); + }else{ + uint32_t B_offset = mega_block_col_idx * (total_cols * n) * K_size; + B_offset += shim_index * n * K_size; + if(B_in_K_N_block_col_major_order){ + seq->npu_dma_memcpy_nd( + sizeof(T_in), + Arg_B, + MM2S, + cur_shimtile, + b_bd_id, + it_channel_1, + {0,0,0,B_offset+ B_const_offset }, + {1, 1,1, K_div_k* n*k}, + {0, 0, 0, 1}, + -1, 0, true, + ENABLE_AXI4 ? aggressive_cache : normal_cache + ); + }else{ + seq->npu_dma_memcpy_nd( + sizeof(T_in), + Arg_B, + MM2S, + cur_shimtile, + b_bd_id, + it_channel_1, + {0,0,0,B_offset+ B_const_offset }, + {1, K_div_k, n, k}, + {0, k, K_size, 1}, + -1, 0, true, + ENABLE_AXI4 ? aggressive_cache : normal_cache + ); + } + } + list_B_shim_queue[shim_index]++; + } + + uint32_t C_offset = mega_block_col_idx * n * total_cols; + C_offset += mega_block_row_idx * m * total_rows * N_size; + C_offset += shim_index * n; + + if (shim_index < total_cols && VALID_COLUMN){ + + if (list_C_shim_queue.at(shim_index) == 2) { + seq->npu_dma_wait( + cur_shimtile, S2MM, it_channel_0 + ); + list_C_shim_queue.at(shim_index)--; + } + + npu_bd_id c_bd_id; + if(list_C_bd_pingpong_flag.at(shim_index) ==0){ + c_bd_id = bd_14; + list_C_bd_pingpong_flag.at(shim_index) =1; + }else{ + c_bd_id = bd_15; + list_C_bd_pingpong_flag.at(shim_index) =0; + } + + seq->npu_dma_memcpy_nd( + sizeof(T_out), + Arg_C, + S2MM, + cur_shimtile, + c_bd_id, + it_channel_0, + {0,0,0, C_offset+ C_const_offset}, + {1,1,4*m, n}, + {0,0,N_size, 1}, + -1, 0, true, + normal_cache + ); + list_C_shim_queue.at(shim_index)++; + } +} + +// wrappers +Gemm::Gemm(LM_Config& config) : _impl(new Impl()){} +Gemm::~Gemm(){ + delete _impl; +} + +uint32_t Gemm::get_m() const{ + return _impl->m; +} +uint32_t Gemm::get_k() const{ + return _impl->k; +} +uint32_t Gemm::get_n() const{ + return _impl->n; +} + +void Gemm::generate_seq(npu_sequence* seq, const uint32_t M, const uint32_t K, const uint32_t N, const uint32_t weight_offset, bool ADD_BIAS, Activation_Type_t OUTPUT_MODE, const uint32_t bias_offset){ + _impl->generate_seq(seq, M, K, N, weight_offset, ADD_BIAS, OUTPUT_MODE, bias_offset, 0); +} + +void Gemm::generate_seq(npu_sequence* seq, const uint32_t M, const uint32_t K, const uint32_t N, const uint32_t weight_offset, bool ADD_BIAS, Activation_Type_t OUTPUT_MODE, const uint32_t bias_offset, + const uint32_t output_offset +){ + _impl->generate_seq(seq, M, K, N, weight_offset, ADD_BIAS, OUTPUT_MODE, bias_offset, output_offset); +} diff --git a/src/detail/gemm/gemm_detail.hpp b/src/detail/gemm/gemm_detail.hpp new file mode 100644 index 000000000..01f5aab11 --- /dev/null +++ b/src/detail/gemm/gemm_detail.hpp @@ -0,0 +1,94 @@ +#ifndef __gemm_detail__ +#define __gemm_detail__ +#include "modules/gemm.hpp" + +struct Gemm::Impl +{ +private: + static constexpr int Arg_C = 0; + static constexpr int Arg_A = 1; + static constexpr int Arg_B = 2; + static constexpr int Arg_Bias = 3; + static constexpr int CT_lock_address_base = 0x000001F000; + + static constexpr int mm_y_group_id = 3; + static constexpr int mm_x_group_id = 4; + static constexpr int mm_w_group_id = 5; + static constexpr npu_tiles shim_tiles[] = {IT0, IT1, IT2, IT3, IT4, IT5, IT6, IT7}; + + static constexpr uint32_t total_cols = 8; + static constexpr uint32_t total_rows = 4; + + static constexpr uint32_t shimtile_size = total_cols > total_rows ? total_cols : total_rows; // max of the two + + static constexpr uint32_t CT_rtp_address = 4096; // for 128 + + static constexpr int CT_rtp_sync_lock_id = 10; // for now hard coded + + std::map valid_A_MT_shimtile_index; + +public: + static constexpr uint32_t m = 64; + static constexpr uint32_t k = 512; + static constexpr uint32_t n = 128; // for now + static constexpr uint32_t r = 8; + static constexpr uint32_t s = 8; + static constexpr uint32_t t = 8; + bool IS_B_ROW_MAJOR = false; // for language model, B is always in col-major order + bool B_in_K_N_block_col_major_order = true; // for language model, B is default in kxn block col-major order + bool ENABLE_AXI4 = true; // for language model, B is always in kxn block col-major order + + Impl(); + ~Impl(); + + /// \brief Generate the sequence + /// \param seq the npu sequence + /// \param M the M dimension + /// \param K the K dimension + /// \param N the N dimension + /// \param weight_offset the weight offset + /// \param ADD_BIAS whether to add bias + /// \param OUTPUT_MODE the output activation mode + /// \param bias_offset the bias offset + void generate_seq( + npu_sequence *seq, + uint32_t M, + uint32_t K, + uint32_t N, + const uint32_t weight_offset, + bool ADD_BIAS, + Activation_Type_t OUTPUT_MODE, + const uint32_t bias_offset, + uint32_t C_const_offset + ); + + template + void generate_runtime_sequence( + npu_sequence* seq, + uint32_t A_const_offset, uint32_t B_const_offset, uint32_t C_const_offset, uint32_t Bias_const_offset, + uint32_t M_size, uint32_t N_size, uint32_t K_size, + std::vector &list_A_shim_queue, + std::vector &list_B_shim_queue, + std::vector &list_C_shim_queue, + bool IS_B_ROW_MAJOR, + bool ENABLE_AXI4, + bool B_in_K_N_block_col_major_order, + bool ADD_BIAS + ); + + template + void generate_shimtile_sequence_per_k_block( + npu_sequence*seq, + uint32_t shim_index, + uint32_t mega_block_row_idx, uint32_t mega_block_col_idx, + uint32_t M_size, uint32_t K_size, uint32_t N_size, + uint32_t A_const_offset, uint32_t B_const_offset, uint32_t C_const_offset, uint32_t Bias_const_offset, + std::vector &list_A_shim_queue, std::vector &list_B_shim_queue, std::vector &list_C_shim_queue, + std::vector &list_A_bd_pingpong_flag, std::vector &list_B_bd_pingpong_flag, std::vector &list_C_bd_pingpong_flag, + bool IS_B_ROW_MAJOR, bool ENABLE_AXI4, + bool B_in_K_N_block_col_major_order, + bool VALID_COLUMN, + bool ADD_BIAS, bool SEND_BIAS + ); +}; +#endif diff --git a/src/detail/gemma4e_npu/avx512_util.hpp b/src/detail/gemma4e_npu/avx512_util.hpp new file mode 100644 index 000000000..60a18ac57 --- /dev/null +++ b/src/detail/gemma4e_npu/avx512_util.hpp @@ -0,0 +1,296 @@ +#pragma once +#include +#include "typedef.hpp" +#include +#include +#include + +// Maximum number of threads for prefill/encode SIMD operations +constexpr int max_prefill_threads = 1; + +/** + * @brief Helper function to load 16 bfloat16 values and convert to __m512 (float). + */ +inline __m512 load_bfloat16_to_m512(const bf16* ptr) { + // Load 16 bfloat16 values (32 bytes) into a __m256i + __m256i bf16_data = _mm256_loadu_si256(reinterpret_cast(ptr)); + + // Convert bfloat16 to float by shifting left 16 bits (bfloat16 is upper 16 bits of float) + __m512i shifted = _mm512_cvtepu16_epi32(bf16_data); + shifted = _mm512_slli_epi32(shifted, 16); + + return _mm512_castsi512_ps(shifted); +} + +/** + * @brief Helper function to store __m512 (float) as 16 bfloat16 values. + * Uses truncation (no rounding). + */ +inline void store_m512_to_bfloat16(bf16* ptr, __m512 data) { + // Convert float to bfloat16 by extracting upper 16 bits + __m512i int_data = _mm512_castps_si512(data); + __m512i shifted = _mm512_srli_epi32(int_data, 16); + __m256i bf16_data = _mm512_cvtepi32_epi16(shifted); + + _mm256_storeu_si256(reinterpret_cast<__m256i*>(ptr), bf16_data); +} + +/** + * @brief Helper function to store __m512 (float) as 16 bfloat16 values with rounding. + * Uses round-to-nearest-even for better accuracy. + */ +inline void store_m512_to_bfloat16_rne(bf16* ptr, __m512 data) { + // Convert float to bfloat16 with rounding to nearest even + __m512i int_data = _mm512_castps_si512(data); + + // Add 0x7FFF for round-to-nearest-even + __m512i rounding = _mm512_set1_epi32(0x7FFF); + __m512i rounded = _mm512_add_epi32(int_data, rounding); + + // Shift right by 16 to get bf16 in lower 16 bits + __m512i shifted = _mm512_srli_epi32(rounded, 16); + + // Pack to 16-bit values + __m256i bf16_data = _mm512_cvtepi32_epi16(shifted); + + _mm256_storeu_si256(reinterpret_cast<__m256i*>(ptr), bf16_data); +} + +// Fast, corrected AVX-512 exp approximation (single-precision). +// Notes: +// - Input x is clamped to [-88, 88] to avoid overflow/underflow. +// - Uses range reduction x = n*ln2 + r, where n is rounded to nearest int. +// - Uses a degree-5 polynomial for exp(r) evaluated with Horner + FMAs. +// - Constructs 2^n by writing the biased exponent field; the biased exponent +// is clamped to [0,255] as a safety measure. +// +// This is an approximation (not fully IEEE-754 accurate for all cases). +inline __m512 _mm512_exp_ps_corrected(__m512 x) { + // clamp x to a reasonable range to avoid overflow/underflow + const __m512 max_val = _mm512_set1_ps(88.0f); + const __m512 min_val = _mm512_set1_ps(-88.0f); + x = _mm512_min_ps(x, max_val); + x = _mm512_max_ps(x, min_val); + + // constants: 1/ln2 and split ln2 = ln2_hi + ln2_lo for extra precision + const __m512 ln2_inv = _mm512_set1_ps(1.44269504088896341f); // 1/ln(2) + const __m512 ln2_hi = _mm512_set1_ps(0.6931471824645996f); // hi part + const __m512 ln2_lo = _mm512_set1_ps(1.9082149292705877e-10f);// lo part + + // compute fx = x * (1/ln2) + __m512 fx = _mm512_mul_ps(x, ln2_inv); + + // round to nearest integer (using rounding intrinsic), storing integer-valued floats + fx = _mm512_roundscale_ps(fx, _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC); + + // convert to int32 (safe since fx holds integer values after rounding) + __m512i emm0 = _mm512_cvttps_epi32(fx); + + // convert back to float for range-reduction arithmetic + __m512 n_ps = _mm512_cvtepi32_ps(emm0); + + // r = x - n * ln2 (use fnmadd to compute c - a*b robustly) + // first r1 = x - n*ln2_hi + __m512 r = _mm512_fnmadd_ps(n_ps, ln2_hi, x); // r = x - n*ln2_hi + // then r = r - n*ln2_lo + r = _mm512_fnmadd_ps(n_ps, ln2_lo, r); // r = x - n*(ln2_hi + ln2_lo) + + // polynomial coefficients for exp(r) ~ 1 + r + r^2/2 + r^3/6 + r^4/24 + r^5/120 + const __m512 c5 = _mm512_set1_ps(0.008333333333333333f); // 1/120 + const __m512 c4 = _mm512_set1_ps(0.041666666666666664f); // 1/24 + const __m512 c3 = _mm512_set1_ps(0.16666666666666666f); // 1/6 + const __m512 c2 = _mm512_set1_ps(0.5f); // 1/2 + const __m512 c1 = _mm512_set1_ps(1.0f); + const __m512 one = _mm512_set1_ps(1.0f); + + // Horner evaluation using FMA: (((c5*r + c4)*r + c3)*r + c2)*r + c1 ; then final *r + 1 + __m512 y = _mm512_fmadd_ps(c5, r, c4); + y = _mm512_fmadd_ps(y, r, c3); + y = _mm512_fmadd_ps(y, r, c2); + y = _mm512_fmadd_ps(y, r, c1); + y = _mm512_fmadd_ps(y, r, one); // y now approximates exp(r) + + // Build 2^n by inserting biased exponent into float bits: + // biased = n + 127 + __m512i biased = _mm512_add_epi32(emm0, _mm512_set1_epi32(127)); + + // clamp biased exponent to [0,255] to avoid invalid bit patterns + biased = _mm512_max_epi32(biased, _mm512_set1_epi32(0)); + biased = _mm512_min_epi32(biased, _mm512_set1_epi32(255)); + + // shift into exponent position (bits 23..30) and reinterpret as float + biased = _mm512_slli_epi32(biased, 23); + __m512 pow2n = _mm512_castsi512_ps(biased); + + // final result: exp(x) ≈ exp(r) * 2^n + return _mm512_mul_ps(y, pow2n); +} + +// Fast AVX-512 log approximation (single-precision). +// Input x must be strictly positive. +inline __m512 _mm512_log_ps_approx(__m512 x) { + const __m512i inv_mant_mask = _mm512_set1_epi32(~0x7f800000); + const __m512i min_norm_pos = _mm512_set1_epi32(0x00800000); + const __m512i exponent_mask = _mm512_set1_epi32(0x7f800000); + const __m512 one = _mm512_set1_ps(1.0f); + + // Extract exponent + __m512i vx = _mm512_castps_si512(x); + __m512i emm0 = _mm512_srli_epi32(vx, 23); + emm0 = _mm512_sub_epi32(emm0, _mm512_set1_epi32(127)); + __m512 e = _mm512_cvtepi32_ps(emm0); + + // Extract mantissa and force exponent to 0 (which means range [1.0, 2.0)) + __m512i m_bits = _mm512_and_si512(vx, inv_mant_mask); + m_bits = _mm512_or_si512(m_bits, _mm512_set1_epi32(0x3f800000)); + __m512 m = _mm512_castsi512_ps(m_bits); + + // Map m from [1, 2) to a symmetric range using p = (m - 1) / (m + 1) + __m512 p1 = _mm512_sub_ps(m, one); + __m512 p2 = _mm512_add_ps(m, one); + __m512 p = _mm512_div_ps(p1, p2); + __m512 p_sq = _mm512_mul_ps(p, p); + + // Evaluate Taylor series for log((1+p)/(1-p)) = 2 * (p + p^3/3 + p^5/5 + p^7/7) + const __m512 c7 = _mm512_set1_ps(2.0f / 7.0f); + const __m512 c5 = _mm512_set1_ps(2.0f / 5.0f); + const __m512 c3 = _mm512_set1_ps(2.0f / 3.0f); + const __m512 c1 = _mm512_set1_ps(2.0f); + + __m512 res = _mm512_fmadd_ps(c7, p_sq, c5); + res = _mm512_fmadd_ps(res, p_sq, c3); + res = _mm512_fmadd_ps(res, p_sq, c1); + res = _mm512_mul_ps(res, p); + + // log(x) = res + e * ln(2) + const __m512 ln2 = _mm512_set1_ps(0.6931471805599453f); + return _mm512_fmadd_ps(e, ln2, res); +} + +// AVX-512 GELU tanh-based, now using the corrected exp function +inline __m512 gelu_tanh_avx512_simd(__m512 gate_vec_fp32) { + // ---- GELU(gate) with tanh approximation: 0.5 * gate * (1 + tanh(sqrt(2/pi) * (gate + 0.044715 * gate^3))) ---- + const __m512 half = _mm512_set1_ps(0.5f); + const __m512 one = _mm512_set1_ps(1.0f); + const __m512 sqrt_2_pi = _mm512_set1_ps(0.7978845608f); // sqrt(2/pi) + const __m512 coeff = _mm512_set1_ps(0.044715f); + + // Compute gate^3 + __m512 gate_squared = _mm512_mul_ps(gate_vec_fp32, gate_vec_fp32); + __m512 gate_cubed = _mm512_mul_ps(gate_squared, gate_vec_fp32); + + // Compute gate + 0.044715 * gate^3 + __m512 inner_term = _mm512_fmadd_ps(coeff, gate_cubed, gate_vec_fp32); + + // Compute sqrt(2/pi) * (gate + 0.044715 * gate^3) + __m512 scaled_term = _mm512_mul_ps(sqrt_2_pi, inner_term); + + // Compute tanh using the corrected exp function: tanh(x) ≈ (exp(x) - exp(-x)) / (exp(x) + exp(-x)) + __m512 exp_pos = _mm512_exp_ps_corrected(scaled_term); + __m512 exp_neg = _mm512_exp_ps_corrected(_mm512_sub_ps(_mm512_setzero_ps(), scaled_term)); + + __m512 numerator = _mm512_sub_ps(exp_pos, exp_neg); + __m512 denominator = _mm512_add_ps(exp_pos, exp_neg); + __m512 tanh_approx = _mm512_div_ps(numerator, denominator); + + // Compute 1 + tanh(...) + __m512 one_plus_tanh = _mm512_add_ps(one, tanh_approx); + + // Compute 0.5 * gate * (1 + tanh(...)) + __m512 gelu = _mm512_mul_ps(half, _mm512_mul_ps(gate_vec_fp32, one_plus_tanh)); + return gelu; +} + +// AVX-512 sigmoid: 1 / (1 + exp(-x)) +inline __m512 sigmoid_avx512(__m512 x) { + const __m512 one = _mm512_set1_ps(1.0f); + __m512 neg_x = _mm512_sub_ps(_mm512_setzero_ps(), x); + __m512 exp_neg_x = _mm512_exp_ps_corrected(neg_x); + return _mm512_div_ps(one, _mm512_add_ps(one, exp_neg_x)); +} + +// AVX-512 SiLU (Swish): x * sigmoid(x) = x / (1 + exp(-x)) +inline __m512 silu_avx512(__m512 x) { + return _mm512_mul_ps(x, sigmoid_avx512(x)); +} + +// Vectorized gaussian function for 16 floats using AVX-512 exp approximation +inline __m512 gaussian_avx512(__m512 x, __m512 sigma) { + const __m512 one = _mm512_set1_ps(1.0f); + const __m512 two = _mm512_set1_ps(2.0f); + const __m512 zero = _mm512_setzero_ps(); + + // Check if sigma <= 0 + __mmask16 mask_zero_sigma = _mm512_cmp_ps_mask(sigma, zero, _CMP_LE_OQ); + + // Compute exp(-(x*x)/(2*sigma*sigma)) using fast AVX-512 approximation + __m512 x_sq = _mm512_mul_ps(x, x); + __m512 sigma_sq = _mm512_mul_ps(sigma, sigma); + __m512 two_sigma_sq = _mm512_mul_ps(two, sigma_sq); + + // Compute -(x*x)/(2*sigma*sigma) + __m512 neg_x_sq_over_2sigma_sq = _mm512_div_ps(_mm512_sub_ps(zero, x_sq), two_sigma_sq); + + // Apply fast exponential + __m512 exp_result = _mm512_exp_ps_corrected(neg_x_sq_over_2sigma_sq); + + // Return 1.0 if sigma <= 0, otherwise exp result + return _mm512_mask_blend_ps(mask_zero_sigma, exp_result, one); +} + +// Fast conversion from uint8 to float with normalization +inline void convert_uint8_to_float_avx512(const uint8_t* src, float* dst, size_t count) { + const size_t simd_count = count & ~15; // Process in chunks of 16 + + for (size_t i = 0; i < simd_count; i += 16) { + // Load 16 uint8 values + __m128i u8_vec = _mm_loadu_si128(reinterpret_cast(src + i)); + + // Convert to 32-bit integers + __m512i i32_vec = _mm512_cvtepu8_epi32(u8_vec); + + // Convert to float + __m512 f32_vec = _mm512_cvtepi32_ps(i32_vec); + + // Store result + _mm512_storeu_ps(dst + i, f32_vec); + } + + // Handle remaining elements + for (size_t i = simd_count; i < count; ++i) { + dst[i] = static_cast(src[i]); + } +} + +// Fast conversion from float to uint8 with clamping +inline void convert_float_to_uint8_avx512(const float* src, uint8_t* dst, size_t count) { + const __m512 zero = _mm512_setzero_ps(); + const __m512 max_val = _mm512_set1_ps(255.0f); + const size_t simd_count = count & ~15; // Process in chunks of 16 + + for (size_t i = 0; i < simd_count; i += 16) { + // Load 16 float values + __m512 f32_vec = _mm512_loadu_ps(src + i); + + // Round to nearest integer + f32_vec = _mm512_roundscale_ps(f32_vec, _MM_FROUND_TO_NEAREST_INT); + + // Clamp to [0, 255] + f32_vec = _mm512_max_ps(f32_vec, zero); + f32_vec = _mm512_min_ps(f32_vec, max_val); + + // Convert to 32-bit integers + __m512i i32_vec = _mm512_cvtps_epi32(f32_vec); + + // Pack to uint8 (with saturation) + __m128i u8_vec = _mm512_cvtusepi32_epi8(i32_vec); + + // Store result + _mm_storeu_si128(reinterpret_cast<__m128i*>(dst + i), u8_vec); + } + + // Handle remaining elements + for (size_t i = simd_count; i < count; ++i) { + dst[i] = static_cast(std::clamp(std::round(src[i]), 0.0f, 255.0f)); + } +} diff --git a/src/detail/gemma4e_npu/conv1d_prefill.hpp b/src/detail/gemma4e_npu/conv1d_prefill.hpp new file mode 100644 index 000000000..2ecbe5044 --- /dev/null +++ b/src/detail/gemma4e_npu/conv1d_prefill.hpp @@ -0,0 +1,230 @@ +#pragma once +#include +#include // Required for std::max +#include +#include "npu_utils/npu_instr_utils.hpp" +#include "tensor_utils/q4_npu_eXpress.hpp" +void conv1d_prefill( + npu_sequence* seq, + const uint32_t L_OUT, + bf16 min_value, + bf16 max_value, + const uint32_t external_x_offset, // in bf16 + const uint32_t external_o_offset, + const uint32_t d +){ + + auto round_up_to_multiple = [](int x, int multiple) -> int { + if (multiple == 0) { + return x; // Cannot divide by zero + } + // This uses integer division to achieve the rounding + return ((x + multiple - 1) / multiple) * multiple; + }; + + npu_tiles IT[8] = {IT0, IT1, IT2, IT3, IT4, IT5, IT6, IT7}; + seq->clear_cmds(); + + const int l_address = 49664; + const int round_address = 8704; + const int min_address = 27136; + const int max_address = 35840; + constexpr int CT_lock_address_base = 0x000001F000; + constexpr int CT_rtp_sync_lock_id = 6; + + assert(d == 1024); + + const int num_col = 8; + float max_float = (float)max_value; + float min_float = (float)min_value; + + const int l_column = 256; + int ROUND = L_OUT / (l_column * num_col); + int remaining = L_OUT % (l_column * num_col); + int needed_col = 0; + int l_left_last_col = 0; + + if (remaining > 0){ + // calculate how many columns are needed for the remaining part + needed_col = remaining / l_column; + if (remaining % l_column != 0) { + // calculate how many rows are needed for the remaining part in the last column + l_left_last_col = remaining - needed_col * l_column; + } + } + + if(ROUND > 0){ + int num_col = 8; + for (int row = 2; row < 6; row++){ + for (int col = 0; col < num_col; col++){ + npu_tiles tile = get_tile(row, col); + seq->rtp_write(tile, l_address, l_column); + seq->rtp_write(tile, round_address, ROUND); + seq->rtp_write(tile, max_address, *(uint32_t*)(&max_float)); + seq->rtp_write(tile, min_address, *(uint32_t*)(&min_float)); + seq->rtp_write(tile, CT_lock_address_base + 16 * CT_rtp_sync_lock_id, 1); // set lock to 1 + } + } + for (int col = 0; col < num_col; col++){ + // send w + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[col], + (npu_bd_id)(1), it_channel_1, + {0, 0, 0, (uint32_t)0}, + {1, 1, (uint32_t)1, (uint32_t)(5 * d)}, + {0, 0, (uint32_t)0, (uint32_t)1}, + -1, 0, false + ); + } + for (int round = 0; round < ROUND; round++){ + int bd_offset = (round % 2) * 8; + for (int col = 0; col < num_col; col++){ + // send x + uint32_t x_offset = external_x_offset + round * l_column * num_col * d + col * l_column * d; + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[col], + (npu_bd_id)(bd_offset + 2), it_channel_0, + {0, 0, 0, x_offset}, + {1, 1, 1, (uint32_t)((l_column + 4) * d)}, + {0, 0, 0, (uint32_t)1}, + -1, 0, false + ); + // receive o + uint32_t o_offset = external_o_offset + round * l_column * num_col * d + col * l_column * d; + seq->npu_dma_memcpy_nd( + 2, 0, + S2MM, IT[col], + (npu_bd_id)(bd_offset + 0), it_channel_0, + {0, 0, 0, o_offset}, + {1, 1, 1, (uint32_t)(l_column * d)}, + {0, 0, 0, (uint32_t)1}, + -1, 0, true + ); + } + if (round > 0){ + for (int col = 0; col < num_col; col++){ + seq->npu_dma_wait( + IT[col], + S2MM, + it_channel_0 + ); + } + } + } + for (int col = 0; col < num_col; col++){ + seq->npu_dma_wait( + IT[col], + S2MM, + it_channel_0 + ); + } + } + + if(remaining > 0){ + int num_col = needed_col; + + for (int row = 2; row < 6; row++){ + for (int col = 0; col < num_col; col++){ + npu_tiles tile = get_tile(row, col); + seq->rtp_write(tile, l_address, l_column); + seq->rtp_write(tile, round_address, 1); + seq->rtp_write(tile, max_address, *(uint32_t*)(&max_float)); + seq->rtp_write(tile, min_address, *(uint32_t*)(&min_float)); + seq->rtp_write(tile, CT_lock_address_base + 16 * CT_rtp_sync_lock_id, 1); // set lock to 1 + } + } + for (int col = 0; col < num_col; col++){ + // send w + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[col], + (npu_bd_id)(1), it_channel_1, + {0, 0, 0, (uint32_t)0}, + {1, 1, (uint32_t)1, (uint32_t)(5 * d)}, + {0, 0, (uint32_t)0, (uint32_t)1}, + -1, 0, false + ); + } + for (int col = 0; col < num_col; col++){ + // send x + uint32_t x_offset = external_x_offset + ROUND * l_column * 8 * d + col * l_column * d; + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[col], + (npu_bd_id)(2), it_channel_0, + {0, 0, 0, x_offset}, + {1, 1, 1, (uint32_t)((l_column + 4) * d)}, + {0, 0, 0, (uint32_t)1}, + -1, 0, false + ); + // receive o + uint32_t o_offset = external_o_offset + ROUND * l_column * 8 * d + col * l_column * d; + seq->npu_dma_memcpy_nd( + 2, 0, + S2MM, IT[col], + (npu_bd_id)(0), it_channel_0, + {0, 0, 0, o_offset}, + {1, 1, 1, (uint32_t)(l_column * d)}, + {0, 0, 0, (uint32_t)1}, + -1, 0, true + ); + seq->npu_dma_wait( + IT[col], + S2MM, + it_channel_0 + ); + } + + if (remaining % l_column != 0){ + for (int row = 2; row < 6; row++){ + npu_tiles tile = get_tile(row, num_col); + seq->rtp_write(tile, l_address, l_left_last_col); + seq->rtp_write(tile, round_address, 1); + seq->rtp_write(tile, max_address, *(uint32_t*)(&max_float)); + seq->rtp_write(tile, min_address, *(uint32_t*)(&min_float)); + seq->rtp_write(tile, CT_lock_address_base + 16 * CT_rtp_sync_lock_id, 1); // set lock to 1 + } + // send w + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[num_col], + (npu_bd_id)(1), it_channel_1, + {0, 0, 0, (uint32_t)0}, + {1, 1, (uint32_t)1, (uint32_t)(5 * d)}, + {0, 0, (uint32_t)0, (uint32_t)1}, + -1, 0, false + ); + // send x + uint32_t x_offset = external_x_offset + ROUND * l_column * 8 * d + num_col * l_column * d; + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[num_col], + (npu_bd_id)(2), it_channel_0, + {0, 0, 0, x_offset}, + {1, 1, 1, (uint32_t)((l_left_last_col + 4) * d)}, + {0, 0, 0, (uint32_t)1}, + -1, 0, false + ); + // receive o + uint32_t o_offset = external_o_offset +ROUND * l_column * 8 * d + num_col * l_column * d; + seq->npu_dma_memcpy_nd( + 2, 0, + S2MM, IT[num_col], + (npu_bd_id)(0), it_channel_0, + {0, 0, 0, o_offset}, + {1, 1, 1, (uint32_t)(l_left_last_col * d)}, + {0, 0, 0, (uint32_t)1}, + -1, 0, true + ); + seq->npu_dma_wait( + IT[num_col], + S2MM, + it_channel_0 + ); + } + } + + seq->cmds2seq(); +} diff --git a/src/detail/gemma4e_npu/embedding_q8_0.hpp b/src/detail/gemma4e_npu/embedding_q8_0.hpp new file mode 100644 index 000000000..5f1565a69 --- /dev/null +++ b/src/detail/gemma4e_npu/embedding_q8_0.hpp @@ -0,0 +1,85 @@ +#ifndef __EMBEDDING_Q8_0_HPP__ +#define __EMBEDDING_Q8_0_HPP__ +#include "buffer.hpp" +#include "tensor_utils/q4_npu_eXpress.hpp" +#include "tensor_2d.hpp" +#include + +class embedding_q8_0{ +private: + static constexpr int Q8_0_GROUP_SIZE = 32; + buffer scale; + buffer qweight; + size_t vocabe_size; + size_t dim; + tensor_2d tensor_qweight; + tensor_2d tensor_scale; + + buffer out_buffer; + + public: + embedding_q8_0(size_t vocab_size, size_t dim) { + this->vocabe_size = vocab_size; + this->dim = dim; + this->scale = buffer(vocab_size * dim / Q8_0_GROUP_SIZE); + this->qweight = buffer(vocab_size * dim); + this->out_buffer = buffer(dim); + this->tensor_qweight = tensor_2d(qweight, dim); + this->tensor_scale = tensor_2d(scale, dim / Q8_0_GROUP_SIZE); + } + + void init_weights(Q4NX& q4nx, const std::string& weight_name){ + q4nx.load_weights(this->scale, weight_name + ".weight.scale"); + q4nx.load_weights(this->qweight, weight_name + ".weight"); + } + + inline void dequant_row_avx512(const int8_t* __restrict qw, const float* __restrict sc, bf16* __restrict dst) { + const size_t num_groups = dim / Q8_0_GROUP_SIZE; + for (size_t i = 0; i < num_groups; i++) { + const int8_t* group_ptr = qw + i * Q8_0_GROUP_SIZE; + bf16* out_ptr = dst + i * Q8_0_GROUP_SIZE; + __m512 scale_vec = _mm512_set1_ps(sc[i]); + + // First 16 elements + __m128i q8_lo = _mm_loadu_si128((const __m128i*)group_ptr); + __m512i q32_lo = _mm512_cvtepi8_epi32(q8_lo); + __m512 f32_lo = _mm512_cvtepi32_ps(q32_lo); + f32_lo = _mm512_mul_ps(f32_lo, scale_vec); + __m512i bf16_lo = _mm512_srli_epi32(_mm512_castps_si512( + _mm512_add_ps(f32_lo, _mm512_castsi512_ps( + _mm512_add_epi32(_mm512_set1_epi32(0x7FFF), + _mm512_and_si512(_mm512_srli_epi32(_mm512_castps_si512(f32_lo), 16), + _mm512_set1_epi32(1)))))), 16); + __m256i out_lo = _mm512_cvtepi32_epi16(bf16_lo); + _mm256_storeu_si256((__m256i*)out_ptr, out_lo); + + // Next 16 elements + __m128i q8_hi = _mm_loadu_si128((const __m128i*)(group_ptr + 16)); + __m512i q32_hi = _mm512_cvtepi8_epi32(q8_hi); + __m512 f32_hi = _mm512_cvtepi32_ps(q32_hi); + f32_hi = _mm512_mul_ps(f32_hi, scale_vec); + __m512i bf16_hi = _mm512_srli_epi32(_mm512_castps_si512( + _mm512_add_ps(f32_hi, _mm512_castsi512_ps( + _mm512_add_epi32(_mm512_set1_epi32(0x7FFF), + _mm512_and_si512(_mm512_srli_epi32(_mm512_castps_si512(f32_hi), 16), + _mm512_set1_epi32(1)))))), 16); + __m256i out_hi = _mm512_cvtepi32_epi16(bf16_hi); + _mm256_storeu_si256((__m256i*)(out_ptr + 16), out_hi); + } + } + + buffer forward(int idx){ + buffer qweight_row = tensor_qweight[idx]; + buffer scale_row = tensor_scale[idx]; + dequant_row_avx512(qweight_row.data(), scale_row.data(), out_buffer.data()); + return out_buffer; + } + + void forward(int idx, buffer& out){ + buffer qweight_row = tensor_qweight[idx]; + buffer scale_row = tensor_scale[idx]; + dequant_row_avx512(qweight_row.data(), scale_row.data(), out.data()); + } +}; + +#endif diff --git a/src/detail/gemma4e_npu/gemma4e_audio.cpp b/src/detail/gemma4e_npu/gemma4e_audio.cpp new file mode 100644 index 000000000..5677164d9 --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_audio.cpp @@ -0,0 +1,2281 @@ +#include "flm_override.hpp" +#include "gemma4e_audio.hpp" + +#include +#include +#include +#ifdef _WIN32 +#include +#endif +#include "utils/debug_utils.hpp" +#include "utils/error_measure.hpp" + +#include "gemma4e_vision_prefill_helper.hpp" + +#include "vision/norm.hpp" +#include "mmRuntimeSequence.hpp" +#include "rot_pos_emb.hpp" + +#include "gemma4e_audio_attention.hpp" +#include "conv1d_prefill.hpp" +#include "utils/utils.hpp" +#include +// #define DEBUG_PRINT_ENCODE_ERROR_METRICS 1 + +Gemma4e_AudioEncoder::~Gemma4e_AudioEncoder() {} + +Gemma4e_AudioEncoder::Gemma4e_AudioEncoder(LM_Config config, npu_xclbin_manager *npu_instance, gemma4e_npu* parent_npu_ptr) + : config(config), npu(npu_instance), model_path(config.model_path), parent_npu_ptr(parent_npu_ptr) +{ + + // load parameters from json file + + { + MM_tile_M = config.sub("audio_config").value("Audio_MM_TILE_M", -1); + MM_tile_K = config.sub("audio_config").value("Audio_MM_TILE_K", -1); + MM_tile_N = config.sub("audio_config").value("Audio_MM_TILE_N", -1); + + seq_len_pad_requirement_for_MM = MM_ROW_SIZE*MM_tile_M; + assert( MM_tile_K % MM_tile_N == 0); + + Gemma4E_Audio_residual_weight = config.sub("audio_config").value("Gemma4E_Audio_residual_weight", 0.0); + assert(this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE % this->parent_npu_ptr->Gemma4E_Audio_num_attention_heads == 0); + Gemma4E_Audio_attention_head_dim = this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE / this->parent_npu_ptr->Gemma4E_Audio_num_attention_heads; + } + + Gemma4E_Audio_q_scale = (1/std::sqrt(Gemma4E_Audio_attention_head_dim)) / std::log(2); + Gemma4E_Audio_k_scale = std::log(1 + std::numbers::e) / std::log(2); + + Gemma4E_Audio_padded_requirement_for_conv1d = this->parent_npu_ptr->Gemma4E_Audio_conv1d_kernel_size\ + - this->parent_npu_ptr->Gemma4E_Audio_conv1d_stride; + DEBUG_BLOCK(1, + std::cout << "Audio q scale: " << Gemma4E_Audio_q_scale << ", k_scale: " << Gemma4E_Audio_k_scale << std::endl; + std::cout << "Gemma4E_Audio_padded_requirement_for_conv1d : " << Gemma4E_Audio_padded_requirement_for_conv1d << std::endl; + ) + + Padded_GEMMA4E_Audio_HIDDEN_SIZE = round_up_to_multiple(this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, MM_tile_K); + Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE = round_up_to_multiple(this->parent_npu_ptr->Gemma4E_Audio_INTERMEDIATE_SIZE, MM_tile_K); + Padded_Gemma4E_Audio_Multimodal_Output_SIZE = round_up_to_multiple(this->parent_npu_ptr->Gemma4E_Audio_Multimodal_Output_SIZE, MM_tile_K); + Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE = round_up_to_multiple( + this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE*2, MM_tile_K + ); + assert(Padded_GEMMA4E_Audio_HIDDEN_SIZE %Gemma4E_Audio_attention_head_dim == 0); + Padded_Gemma4E_Audio_num_attention_heads = Padded_GEMMA4E_Audio_HIDDEN_SIZE / Gemma4E_Audio_attention_head_dim; + assert(Padded_Gemma4E_Audio_num_attention_heads >= this->parent_npu_ptr->Gemma4E_Audio_num_attention_heads); + + this->proj = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "vision_mm.xclbin")); + this->proj_high_precision = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "vision_mm_high_precision.xclbin")); + this->conv1d = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "audio_conv1d.xclbin")); + + this->sub_sampleConvProjection_app = this->proj_high_precision->create_app(); + this->q_proj_app = this->proj->create_app(); + this->k_proj_app = this->proj->create_app(); + this->k_relative_proj_app = this->proj->create_app(); + this->v_proj_app = this->proj->create_app(); + this->o_proj_app = this->proj->create_app(); + this->ffn_down_proj_app = this->proj->create_app(); + this->ffn_up_proj_app = this->proj->create_app(); + this->conv1d_start_proj_app = this->proj->create_app(); + this->conv1d_end_proj_app = this->proj->create_app(); + this->conv1d_app = this->conv1d->create_app(); + audio_pre_encode_proj_app = this->proj->create_app(); + audio_to_language_proj_app = this->proj->create_app(); + + this->audio_attn_k_rel_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->audio_attn_k_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->audio_attn_v_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->audio_attn_q_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->audio_attn_o_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->audio_conv1d_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->audio_conv_pw_1_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->audio_conv_pw_2_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->audio_ffn_down_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->audio_ffn_up_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->audio_ffn_down_1_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->audio_ffn_up_1_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->attn_post_norm_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->attn_pre_norm_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->conv_norm_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->norm_conv_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->ffn_norm_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->ffn_norm_1_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->ffn_post_norm_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->ffn_post_norm_1_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->norm2_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); + this->per_dim_scale_with_softplus_weight.resize(this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers); +} + +void Gemma4e_AudioEncoder::conv1d_layer( + + int layer_idx, + int seq_len, int seq_len_padded, + std::vector &seq_len_per_audio, std::vector &start_seq_len_index_per_audio, + buffer &conv1d_start_proj_input, buffer &conv1d_start_proj_output, + buffer &conv1d_input, buffer &conv1d_output, + buffer &conv1d_end_proj_input, buffer &conv1d_end_proj_output, + + bf16 conv1d_start_input_min, bf16 conv1d_start_input_max, + bf16 conv1d_start_output_min, bf16 conv1d_start_output_max, + + bf16 conv1d_end_input_min, bf16 conv1d_end_input_max, + bf16 conv1d_end_output_min, bf16 conv1d_end_output_max, + + gemma4e_audio_payload_t* audio_payload, + SafeTensors *reference_safetensor +){ + + memcpy(residual.data(), hidden_state.data(), seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16)); + + // compare the norm weigths + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Comparing conv1d_layer " << layer_idx << " norm weights..." << std::endl; + buffer ref_per_layer_norm_weight; + reference_safetensor->load_weights( + ref_per_layer_norm_weight, + "Gemma4AudioLightConv1d_"+std::to_string(layer_idx)+ "_pre_layer_norm_weight" + ); + print_error_metrics( + this->conv_norm_weight[layer_idx].data(), ref_per_layer_norm_weight.data(), + 1, + ref_per_layer_norm_weight.size(), 1, + ref_per_layer_norm_weight.size(), 1 + ); + + std::cout << "Comparing conv1d_layer " << layer_idx << " conv norm weights..." << std::endl; + buffer ref_conv_norm_weight; + reference_safetensor->load_weights( + ref_conv_norm_weight, + "Gemma4AudioLightConv1d_"+std::to_string(layer_idx)+ "_conv_norm_weight" + ); + print_error_metrics( + this->norm_conv_weight[layer_idx].data(), ref_conv_norm_weight.data(), + 1, + ref_conv_norm_weight.size(), 1, + ref_conv_norm_weight.size(), 1 + ); + } + #endif + + simd_rms_norm( + hidden_state.data(), this->conv_norm_weight[layer_idx].data(), conv1d_start_proj_input.data(), + seq_len, this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, 1e-6f + ); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Comparing conv1d_layer " << layer_idx << " pre-attention norm output..." << std::endl; + buffer ref_Gemma4AudioLightConv1d_layer_idx_hidden_states_after_pre_layer_norm; + reference_safetensor->load_weights( + ref_Gemma4AudioLightConv1d_layer_idx_hidden_states_after_pre_layer_norm, + "Gemma4AudioLightConv1d_"+std::to_string(layer_idx)+ "_hidden_states_after_pre_layer_norm" + ); + size_t ref_offset_per_audio = ref_Gemma4AudioLightConv1d_layer_idx_hidden_states_after_pre_layer_norm.size() / audio_payload->num_audios; + for(int i = 0; i < audio_payload->num_audios; i++){ + print_error_metrics( + conv1d_start_proj_input.data() + start_seq_len_index_per_audio[i]*Padded_GEMMA4E_Audio_HIDDEN_SIZE, + ref_Gemma4AudioLightConv1d_layer_idx_hidden_states_after_pre_layer_norm.data() + i*ref_offset_per_audio, + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif + + simd_clamp( + conv1d_start_proj_input.data(), conv1d_start_proj_input.data(), + conv1d_start_input_min, conv1d_start_input_max, + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + { + + generate_mm_sequence( + *this->conv1d_start_proj_app.seq(), + seq_len_padded, this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE, + MM_tile_M, MM_tile_K, MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, //no bias, no activation + 1, conv1d_start_output_min, conv1d_start_output_max, // enable clamp in output + ENABLE_QKV_REORDER, 0// since we don't need it anymore + + ); + } + DEBUG_BLOCK(1, + std::cout << "DEBUG: Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE: " << Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE << std::endl; + ) + conv1d_start_proj_input.sync_to_device(); + audio_conv_pw_1_weight[layer_idx].sync_to_device(); + FLM_OVERRIDE(audio_conv1d_start_proj, conv1d_start_proj_app(conv1d_start_proj_input, audio_conv_pw_1_weight[layer_idx], conv1d_start_proj_output), layer_idx); + conv1d_start_proj_output.sync_from_device(); + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Comparing conv1d_layer " << layer_idx << " conv1d linear output..." << std::endl; + buffer ref_conv1d_start_proj_output; + reference_safetensor->load_weights( + ref_conv1d_start_proj_output, + "Gemma4AudioLightConv1d_" + std::to_string(layer_idx) +"_hidden_states_after_linear_start" + ); + size_t ref_offset_per_audio = ref_conv1d_start_proj_output.size() / audio_payload->num_audios; + for(int i = 0; i < audio_payload->num_audios; i++){ + print_error_metrics( + conv1d_start_proj_output.data() + start_seq_len_index_per_audio[i]*Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE, + ref_conv1d_start_proj_output.data() + i*ref_offset_per_audio, + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE * 2, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE + ); + } + } + #endif + + assert(Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE == (Padded_GEMMA4E_Audio_HIDDEN_SIZE*2) ); + + for(int i = 0, seq_len_offset = 0; i < audio_payload->num_audios; i++){ + + seq_len_offset += Gemma4E_Audio_padded_requirement_for_conv1d; + simd_glu( + conv1d_start_proj_output.data() + start_seq_len_index_per_audio[i] * Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE, + conv1d_input.data() + (seq_len_offset * Padded_GEMMA4E_Audio_HIDDEN_SIZE), + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + + seq_len_offset += (seq_len_per_audio[i] ); + } + conv1d_start_proj_output.sync_to_device(); + + // compare with error metrics + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Comparing conv1d_layer " << layer_idx << " conv1d GLU output..." << std::endl; + buffer ref_conv1d_glu_output; + reference_safetensor->load_weights( + ref_conv1d_glu_output, + "Gemma4AudioLightConv1d_" + std::to_string(layer_idx) +"_hidden_states_after_glu" + ); + + size_t ref_offset_per_audio = ref_conv1d_glu_output.size() / audio_payload->num_audios; + + for(int i = 0, seq_len_offset = 0; i < audio_payload->num_audios; i++){ + seq_len_offset += Gemma4E_Audio_padded_requirement_for_conv1d; + + print_error_metrics( + conv1d_input.data() + seq_len_offset*Padded_GEMMA4E_Audio_HIDDEN_SIZE, + ref_conv1d_glu_output.data() + i*ref_offset_per_audio, + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + + seq_len_offset += seq_len_per_audio[i]; + } + } + #endif + + // conv1d + for(int i = 0, seq_len_offset = 0; i < audio_payload->num_audios; i++){ + + conv1d_prefill( + this->conv1d_app.seq(), + seq_len_per_audio[i], + -1.18e30f, 3.38e30f, // no clamping, + seq_len_offset * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + start_seq_len_index_per_audio[i] *Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + + // scalar_conv1d( + // this->parent_npu_ptr->Gemma4E_Audio_conv1d_kernel_size, + // this->parent_npu_ptr->Gemma4E_Audio_conv1d_stride, + // conv1d_input.data() + seq_len_offset * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + // audio_conv1d_weight[layer_idx].data(), + // conv1d_output.data() + start_seq_len_index_per_audio[i] *Padded_GEMMA4E_Audio_HIDDEN_SIZE, + // seq_len_per_audio[i], + // Padded_GEMMA4E_Audio_HIDDEN_SIZE + // ); + + seq_len_offset += (Gemma4E_Audio_padded_requirement_for_conv1d + seq_len_per_audio[i]); + + conv1d_input.sync_to_device(); + audio_conv1d_weight[layer_idx].sync_to_device(); + FLM_OVERRIDE(audio_conv1d, conv1d_app(conv1d_output, conv1d_input,audio_conv1d_weight[layer_idx] ), layer_idx); + conv1d_output.sync_from_device(); + } + + // utils::print_matrix( + // audio_conv1d_weight[layer_idx], Padded_GEMMA4E_Audio_HIDDEN_SIZE + + // ); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Comparing conv1d_layer " << layer_idx << " conv1d output..." << std::endl; + buffer ref_conv1d_output; + reference_safetensor->load_weights( + ref_conv1d_output, + "Gemma4AudioLightConv1d_" + std::to_string(layer_idx) +"_hidden_states_after_depthwise_conv1d" + ); + + size_t ref_offset_per_audio = ref_conv1d_output.size() / audio_payload->num_audios; + for(int i = 0; i < audio_payload->num_audios; i++){ + print_error_metrics( + conv1d_output.data() + start_seq_len_index_per_audio[i]*Padded_GEMMA4E_Audio_HIDDEN_SIZE, + ref_conv1d_output.data() + i*ref_offset_per_audio, + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif + + simd_rms_norm( + conv1d_output.data(), this->norm_conv_weight[layer_idx].data(), conv1d_output.data(), + seq_len, this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, 1e-6f + ); + + simd_silu( + conv1d_output.data(), conv1d_end_proj_input.data(), + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + conv1d_output.sync_to_device(); + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer ref_hidden_states_after_conv_act; + reference_safetensor->load_weights( + ref_hidden_states_after_conv_act, + "Gemma4AudioLightConv1d_" + std::to_string(layer_idx) +"_hidden_states_after_conv_act" + ); + + size_t ref_offset_per_audio = ref_hidden_states_after_conv_act.size() / audio_payload->num_audios; + + for(int i = 0; i < audio_payload->num_audios; i++){ + std::cout << "Comparing conv1d_layer " << layer_idx << " conv1d activation output for audio " << i << std::endl; + print_error_metrics( + conv1d_end_proj_input.data() + start_seq_len_index_per_audio[i]*Padded_GEMMA4E_Audio_HIDDEN_SIZE, + ref_hidden_states_after_conv_act.data() + i*ref_offset_per_audio, + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif + + simd_clamp( + conv1d_end_proj_input.data(), conv1d_end_proj_input.data(), + conv1d_end_input_min, conv1d_end_input_max, + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + + { + generate_mm_sequence( + *this->conv1d_end_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, Padded_GEMMA4E_Audio_HIDDEN_SIZE, + MM_tile_M, MM_tile_K, MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, //no bias, no activation + 1, conv1d_end_output_min,conv1d_end_output_max, // enable clamp in output + ENABLE_QKV_REORDER, 0// since we don't need it anymore + + ); + } + + conv1d_end_proj_input.sync_to_device(); + audio_conv_pw_2_weight[layer_idx].sync_to_device(); + FLM_OVERRIDE(audio_conv1d_end_proj, conv1d_end_proj_app(conv1d_end_proj_input, audio_conv_pw_2_weight[layer_idx], conv1d_end_proj_output), layer_idx); + conv1d_end_proj_output.sync_from_device(); + + simd_add(conv1d_end_proj_output.data(), residual.data(), hidden_state.data(), + seq_len*Padded_GEMMA4E_Audio_HIDDEN_SIZE); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer ref_conv1d_final_output; + reference_safetensor->load_weights( + ref_conv1d_final_output, + "Gemma4AudioLightConv1d_" + std::to_string(layer_idx) +"_hidden_states_after_residual" + ); + + size_t ref_offset_per_audio = ref_conv1d_final_output.size() / audio_payload->num_audios; + for(int i = 0; i < audio_payload->num_audios; i++){ + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_audio[i]*Padded_GEMMA4E_Audio_HIDDEN_SIZE, + ref_conv1d_final_output.data() + i*ref_offset_per_audio, + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif +} + +void Gemma4e_AudioEncoder::ffn_layer( + + buffer &ffn_up_proj_input, + buffer &ffn_up_proj_output_down_input, + buffer &ffn_down_proj_output, + + buffer &cur_ffn_norm_weight, // ffn_norm or ffn_norm_1 + buffer &cur_ffn_post_norm_weight,// ffn_post_norm_weight or ffn_post_norm_1_weight + buffer &cur_ffn_up_weight, // audio_ffn_up_weight or audio_ffn_up_1_weight + buffer &cur_ffn_down_weight, // audio_ffn_down_weight + int seq_len, int seq_len_padded, + + bf16 cur_audio_ffn_up_input_min, bf16 cur_audio_ffn_up_input_max, + bf16 cur_audio_ffn_up_output_min, bf16 cur_audio_ffn_up_output_max, + bf16 cur_audio_ffn_down_input_min, bf16 cur_audio_ffn_down_input_max, + bf16 cur_audio_ffn_down_output_min, bf16 cur_audio_ffn_down_output_max + +){ + + memcpy(residual.data(), hidden_state.data(), seq_len * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16)); + + generate_mm_sequence( + *this->ffn_up_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE, + MM_tile_M, MM_tile_K, MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 2, //no bias, silu activation + 1, cur_audio_ffn_up_output_min,cur_audio_ffn_up_output_max, // with clamp + ENABLE_QKV_REORDER, 0// since we don't need it anymore + + ); + + simd_rms_norm( + hidden_state.data(), + cur_ffn_norm_weight.data(), + ffn_up_proj_input.data(), + seq_len, + this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 1e-6f + ); + simd_clamp(ffn_up_proj_input.data(), + ffn_up_proj_input.data(), + cur_audio_ffn_up_input_min, cur_audio_ffn_up_input_max, + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + + //ffn_up_proj_app(ffn_up_proj_input, cur_ffn_up_weight, ffn_up_proj_output_down_input); + auto ffn_up_proj_run = FLM_OVERRIDE(audio_ffn_up_proj, ffn_up_proj_app.create_run( + ffn_up_proj_input,cur_ffn_up_weight, ffn_up_proj_output_down_input + )); + ffn_up_proj_input.sync_to_device(); + cur_ffn_up_weight.sync_to_device(); + ffn_up_proj_run.start(); + + generate_mm_sequence( + *this->ffn_down_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE, Padded_GEMMA4E_Audio_HIDDEN_SIZE, + MM_tile_M, MM_tile_K, MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, //no bias, no activation + 1, cur_audio_ffn_down_output_min,cur_audio_ffn_down_output_max, // with clamp + ENABLE_QKV_REORDER, 0// since we don't need it anymore + + ); + + ffn_up_proj_run.wait(); + ffn_up_proj_output_down_input.sync_from_device(); + + simd_clamp( + ffn_up_proj_output_down_input.data(), + ffn_up_proj_output_down_input.data(), + cur_audio_ffn_down_input_min, cur_audio_ffn_down_input_max, + seq_len * Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE + ); + + ffn_up_proj_output_down_input.sync_to_device(); + cur_ffn_down_weight.sync_to_device(); + FLM_OVERRIDE(audio_ffn_down_proj, ffn_down_proj_app(ffn_up_proj_output_down_input,cur_ffn_down_weight, ffn_down_proj_output )); + ffn_down_proj_output.sync_from_device(); + + // perform a post_layer_nrom + simd_rms_norm( + ffn_down_proj_output.data(), + cur_ffn_post_norm_weight.data(), + ffn_down_proj_output.data(), + seq_len, + this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 1e-6f + ); + + simd_add( + ffn_down_proj_output.data(), residual.data(), hidden_state.data(), + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + ffn_down_proj_output.sync_to_device(); +} + +void Gemma4e_AudioEncoder::init_weights(SafeTensors &q4nx){ + DEBUG_BLOCK(1, + std::cout << "Initializing Gemma4e_AudioEncoder weights from model path: " << model_path << std::endl; + ) + + q4nx.load_weights(this->audio_subsample_conv2d_weight_0,"model.audio.subsample.conv_layer0.weight"); + q4nx.load_weights(this->audio_subsample_conv2d_norm_weight_0,"model.audio.subsample.conv_layer0.norm.weight"); + q4nx.load_weights(this->audio_subsample_conv2d_weight_1,"model.audio.subsample.conv_layer1.weight"); + q4nx.load_weights(this->audio_subsample_conv2d_norm_weight_1,"model.audio.subsample.conv_layer1.norm.weight"); + { + buffer temp_buffer; + q4nx.load_weights(temp_buffer, "model.audio.encode_input_projection.weight"); + this->audio_embedding_projection_weight = this->sub_sampleConvProjection_app.create_bo_buffer(temp_buffer.size()); + memcpy( + this->audio_embedding_projection_weight.data(), + temp_buffer.data(), + temp_buffer.size() * sizeof(bf16) + ); + } + + audio_pre_encode_weight = this->audio_pre_encode_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_HIDDEN_SIZE * Padded_Gemma4E_Audio_Multimodal_Output_SIZE + Padded_Gemma4E_Audio_Multimodal_Output_SIZE + + ); + DEBUG_BLOCK(1, + std::cout <<"Padded_Gemma4E_Audio_Multimodal_Output_SIZE: " << Padded_Gemma4E_Audio_Multimodal_Output_SIZE << std::endl; + ) + q4nx.load_weights( + this->audio_pre_encode_weight, + "model.audio.pre_encoder.bias", 0 // the bias + ); + q4nx.load_weights( + this->audio_pre_encode_weight, + "model.audio.pre_encoder.weight", + sizeof(bf16) * Padded_Gemma4E_Audio_Multimodal_Output_SIZE // the weight + ); + + audio_to_language_projection_weight = this->audio_to_language_proj_app.create_bo_buffer( + Padded_Gemma4E_Audio_Multimodal_Output_SIZE * parent_npu_ptr->Gemma4E_Audio_language_projection_output_size + ); + assert( + parent_npu_ptr->Gemma4E_Audio_language_projection_output_size% MM_tile_N == 0 + ); + q4nx.load_weights( + audio_to_language_projection_weight, + "model.audio.embedding_projection.weight" + ); + + for(int layer_id =0; layer_id < this->parent_npu_ptr->Gemma4E_Audio_num_attention_layers; layer_id ++){ + + this->audio_attn_k_weight[layer_id] = this->k_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_HIDDEN_SIZE*Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + q4nx.load_weights( + this->audio_attn_k_weight[layer_id], + "model.audio." + std::to_string(layer_id)+ ".attn_k_proj.weight" + ); + + buffer k_input_max; + q4nx.load_weights(k_input_max, + "model.audio."+std::to_string(layer_id)+ ".attn_k_proj.input_max"); + assert(k_input_max.size() == 1); + this->audio_k_input_max.push_back(k_input_max[0]); + + buffer k_input_min; + q4nx.load_weights(k_input_min, + "model.audio."+std::to_string(layer_id)+ ".attn_k_proj.input_min"); + assert(k_input_min.size() == 1); + this->audio_k_input_min.push_back(k_input_min[0]); + + buffer k_output_max; + q4nx.load_weights(k_output_max, + "model.audio."+std::to_string(layer_id)+ ".attn_k_proj.output_max"); + assert(k_output_max.size() == 1); + this->audio_k_output_max.push_back(k_output_max[0]); + + buffer k_output_min; + q4nx.load_weights(k_output_min, + "model.audio."+std::to_string(layer_id)+ ".attn_k_proj.output_min"); + assert(k_output_min.size() == 1); + this->audio_k_output_min.push_back(k_output_min[0]); + + this->audio_attn_k_rel_weight[layer_id] = this->k_relative_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_HIDDEN_SIZE*Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + q4nx.load_weights( + this->audio_attn_k_rel_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".attn_k_proj.rel_weight" + ); + + this->audio_attn_o_weight[layer_id] = this->o_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_HIDDEN_SIZE*Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + q4nx.load_weights( + this->audio_attn_o_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".attn_out_proj.weight" + ); + + buffer o_input_max; + q4nx.load_weights(o_input_max, + "model.audio."+std::to_string(layer_id)+ ".attn_out_proj.input_max"); + assert(o_input_max.size() == 1); + + buffer o_input_min; + q4nx.load_weights(o_input_min, + "model.audio."+std::to_string(layer_id)+ ".attn_out_proj.input_min"); + assert(o_input_min.size() == 1); + this->audio_o_input_min.push_back(o_input_min[0]); + + this->audio_o_input_max.push_back(o_input_max[0]); + buffer o_output_max; + q4nx.load_weights(o_output_max, + "model.audio."+std::to_string(layer_id)+ ".attn_out_proj.output_max"); + assert(o_output_max.size() == 1); + + this->audio_o_output_max.push_back(o_output_max[0]); + buffer o_output_min; + q4nx.load_weights(o_output_min, + "model.audio."+std::to_string(layer_id)+ ".attn_out_proj.output_min"); + assert(o_output_min.size() == 1); + this->audio_o_output_min.push_back(o_output_min[0]); + + q4nx.load_weights( + this->attn_post_norm_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".attn_post_norm.weight" + ); + q4nx.load_weights( + this->attn_pre_norm_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".attn_pre_norm.weight" + ); + + this->audio_attn_q_weight[layer_id] = this->q_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_HIDDEN_SIZE*Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + q4nx.load_weights( + this->audio_attn_q_weight[layer_id], + "model.audio."+ std::to_string(layer_id)+".attn_q_proj.weight" + ); + buffer q_input_max; + q4nx.load_weights(q_input_max, + "model.audio."+std::to_string(layer_id)+ ".attn_q_proj.input_max"); + assert(q_input_max.size() == 1); + this->audio_q_input_max.push_back(q_input_max[0]); + + buffer q_input_min; + q4nx.load_weights(q_input_min, + "model.audio."+std::to_string(layer_id)+ ".attn_q_proj.input_min"); + assert(q_input_min.size() == 1); + this->audio_q_input_min.push_back(q_input_min[0]); + + buffer q_output_max; + q4nx.load_weights(q_output_max, + "model.audio."+std::to_string(layer_id)+ ".attn_q_proj.output_max"); + assert(q_output_max.size() == 1); + this->audio_q_output_max.push_back(q_output_max[0]); + + buffer q_output_min; + q4nx.load_weights(q_output_min, + "model.audio."+std::to_string(layer_id)+ ".attn_q_proj.output_min"); + assert(q_output_min.size() == 1); + this->audio_q_output_min.push_back(q_output_min[0]); + + this->audio_attn_v_weight[layer_id] = this->v_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_HIDDEN_SIZE*Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + q4nx.load_weights( + this->audio_attn_v_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".attn_v_proj.weight"); + buffer v_input_max; + q4nx.load_weights(v_input_max, + "model.audio."+std::to_string(layer_id)+ ".attn_v_proj.input_max"); + assert(v_input_max.size() == 1); + this->audio_v_input_max.push_back(v_input_max[0]); + + buffer v_input_min; + q4nx.load_weights(v_input_min, + "model.audio."+std::to_string(layer_id)+ ".attn_v_proj.input_min"); + assert(v_input_min.size() == 1); + this->audio_v_input_min.push_back(v_input_min[0]); + + buffer v_output_max; + q4nx.load_weights(v_output_max, + "model.audio."+std::to_string(layer_id)+ ".attn_v_proj.output_max"); + assert(v_output_max.size() == 1); + this->audio_v_output_max.push_back(v_output_max[0]); + + buffer v_output_min; + q4nx.load_weights(v_output_min, + "model.audio."+std::to_string(layer_id)+ ".attn_v_proj.output_min"); + assert(v_output_min.size() == 1); + this->audio_v_output_min.push_back(v_output_min[0]); + + this->audio_conv1d_weight[layer_id] = this->conv1d_app.create_bo_buffer( + this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE * this->parent_npu_ptr->Gemma4E_Audio_conv1d_kernel_size + ); + q4nx.load_weights( + this->audio_conv1d_weight[layer_id], + "model.audio." +std::to_string(layer_id) + ".conv_dw.weight" + ); + + q4nx.load_weights( + this->conv_norm_weight[layer_id], + "model.audio."+std::to_string(layer_id) +".conv_norm.weight" + ); + q4nx.load_weights( + this->norm_conv_weight[layer_id], + "model.audio."+std::to_string(layer_id) +".norm_conv.weight" + ); + + // + this->audio_conv_pw_1_weight[layer_id] = this->conv1d_start_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_HIDDEN_SIZE * Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE + ); + q4nx.load_weights( + this->audio_conv_pw_1_weight[layer_id], + "model.audio." +std::to_string(layer_id) + ".conv_pw_1.weight" + ); + buffer conv_pw_1_input_max; + q4nx.load_weights(conv_pw_1_input_max, + "model.audio."+std::to_string(layer_id)+ ".conv_pw_1.input_max"); + assert(conv_pw_1_input_max.size() == 1); + this->audio_conv_pw1_input_max.push_back(conv_pw_1_input_max[0]); + + buffer conv_pw_1_input_min; + q4nx.load_weights(conv_pw_1_input_min, + "model.audio."+std::to_string(layer_id)+ ".conv_pw_1.input_min"); + assert(conv_pw_1_input_min.size() == 1); + this->audio_conv_pw1_input_min.push_back(conv_pw_1_input_min[0]); + + buffer conv_pw_1_output_max; + q4nx.load_weights(conv_pw_1_output_max, + "model.audio."+std::to_string(layer_id)+ ".conv_pw_1.output_max"); + assert(conv_pw_1_output_max.size() == 1); + this->audio_conv_pw1_output_max.push_back(conv_pw_1_output_max[0]); + + buffer conv_pw_1_output_min; + q4nx.load_weights(conv_pw_1_output_min, + "model.audio."+std::to_string(layer_id)+ ".conv_pw_1.output_min"); + assert(conv_pw_1_output_min.size() == 1); + this->audio_conv_pw1_output_min.push_back(conv_pw_1_output_min[0]); + + this->audio_conv_pw_2_weight[layer_id] = this->conv1d_end_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_HIDDEN_SIZE * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + q4nx.load_weights( + this->audio_conv_pw_2_weight[layer_id], + "model.audio." +std::to_string(layer_id) + ".conv_pw_2.weight" + ); + buffer conv_pw_2_input_max; + q4nx.load_weights(conv_pw_2_input_max, + "model.audio."+std::to_string(layer_id)+ ".conv_pw_2.input_max"); + assert(conv_pw_2_input_max.size() == 1); + this->audio_conv_pw2_input_max.push_back(conv_pw_2_input_max[0]); + + buffer conv_pw_2_input_min; + q4nx.load_weights(conv_pw_2_input_min, + "model.audio."+std::to_string(layer_id)+ ".conv_pw_2.input_min"); + assert(conv_pw_2_input_min.size() == 1); + this->audio_conv_pw2_input_min.push_back(conv_pw_2_input_min[0]); + + buffer conv_pw_2_output_max; + q4nx.load_weights(conv_pw_2_output_max, + "model.audio."+std::to_string(layer_id)+ ".conv_pw_2.output_max"); + assert(conv_pw_2_output_max.size() == 1); + this->audio_conv_pw2_output_max.push_back(conv_pw_2_output_max[0]); + + buffer conv_pw_2_output_min; + q4nx.load_weights(conv_pw_2_output_min, + "model.audio."+std::to_string(layer_id)+ ".conv_pw_2.output_min"); + assert(conv_pw_2_output_min.size() == 1); + this->audio_conv_pw2_output_min.push_back(conv_pw_2_output_min[0]); + + this->audio_ffn_down_weight[layer_id] = this->ffn_down_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_HIDDEN_SIZE * Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE + ); + q4nx.load_weights( + this->audio_ffn_down_weight[layer_id], + "model.audio." +std::to_string(layer_id) + ".ffn.down_proj.weight" + ); + buffer ffn_down_input_max; + q4nx.load_weights(ffn_down_input_max, + "model.audio."+std::to_string(layer_id)+ ".ffn.down_proj.input_max"); + assert(ffn_down_input_max.size() == 1); + this->audio_ffn_down_input_max.push_back(ffn_down_input_max[0]); + + buffer ffn_down_input_min; + q4nx.load_weights(ffn_down_input_min, + "model.audio."+std::to_string(layer_id)+ ".ffn.down_proj.input_min"); + assert(ffn_down_input_min.size() == 1); + this->audio_ffn_down_input_min.push_back(ffn_down_input_min[0]); + + buffer ffn_down_output_max; + q4nx.load_weights(ffn_down_output_max, + "model.audio."+std::to_string(layer_id)+ ".ffn.down_proj.output_max"); + assert(ffn_down_output_max.size() == 1); + this->audio_ffn_down_output_max.push_back(ffn_down_output_max[0]); + + buffer ffn_down_output_min; + q4nx.load_weights(ffn_down_output_min, + "model.audio."+std::to_string(layer_id)+ ".ffn.down_proj.output_min"); + assert(ffn_down_output_min.size() == 1); + this->audio_ffn_down_output_min.push_back(ffn_down_output_min[0]); + + this->audio_ffn_down_1_weight[layer_id] = this->ffn_down_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_HIDDEN_SIZE * Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE + ); + q4nx.load_weights( + this->audio_ffn_down_1_weight[layer_id], + "model.audio." +std::to_string(layer_id) + ".ffn.down_proj_1.weight" + ); + buffer ffn_down_1_input_max; + q4nx.load_weights(ffn_down_1_input_max, + "model.audio."+std::to_string(layer_id)+ ".ffn.down_proj_1.input_max"); + assert(ffn_down_1_input_max.size() == 1); + this->audio_ffn_down_1_input_max.push_back(ffn_down_1_input_max[0]); + + buffer ffn_down_1_input_min; + q4nx.load_weights(ffn_down_1_input_min, + "model.audio."+std::to_string(layer_id)+ ".ffn.down_proj_1.input_min"); + assert(ffn_down_1_input_min.size() == 1); + this->audio_ffn_down_1_input_min.push_back(ffn_down_1_input_min[0]); + + buffer ffn_down_1_output_max; + q4nx.load_weights(ffn_down_1_output_max, + "model.audio."+std::to_string(layer_id)+ ".ffn.down_proj_1.output_max"); + assert(ffn_down_1_output_max.size() == 1); + this->audio_ffn_down_1_output_max.push_back(ffn_down_1_output_max[0]); + + buffer ffn_down_1_output_min; + q4nx.load_weights(ffn_down_1_output_min, + "model.audio."+std::to_string(layer_id)+ ".ffn.down_proj_1.output_min"); + assert(ffn_down_1_output_min.size() == 1); + this->audio_ffn_down_1_output_min.push_back(ffn_down_1_output_min[0]); + + q4nx.load_weights( + this->ffn_norm_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".ffn_norm.weight" + ); + q4nx.load_weights( + this->ffn_norm_1_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".ffn_norm_1.weight" + ); + q4nx.load_weights( + this->ffn_post_norm_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".ffn_post_norm.weight" + ); + //NOTE: Optimization in multiple ffn_post_norm_weight with this.Gemma4E_Audio_residual_weight + for(int i = 0; i < this->ffn_post_norm_weight[layer_id].size(); i++){ + this->ffn_post_norm_weight[layer_id][i] = this->ffn_post_norm_weight[layer_id][i] * this->Gemma4E_Audio_residual_weight; + } + + q4nx.load_weights( + this->ffn_post_norm_1_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".ffn_post_norm_1.weight" + ); + //NOTE: Optimization in multiple ffn_post_norm_1_weight with this.Gemma4E_Audio_residual_weight + for(int i = 0; i < this->ffn_post_norm_1_weight[layer_id].size(); i++){ + this->ffn_post_norm_1_weight[layer_id][i] = this->ffn_post_norm_1_weight[layer_id][i] * this->Gemma4E_Audio_residual_weight; + } + + this->audio_ffn_up_weight[layer_id] = this->ffn_up_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + q4nx.load_weights( + this->audio_ffn_up_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".ffn.up_proj.weight" + ); + + buffer ffn_up_input_max; + q4nx.load_weights(ffn_up_input_max, + "model.audio."+std::to_string(layer_id)+ ".ffn.up_proj.input_max"); + assert(ffn_up_input_max.size() == 1); + this->audio_ffn_up_input_max.push_back(ffn_up_input_max[0]); + + buffer ffn_up_input_min; + q4nx.load_weights(ffn_up_input_min, + "model.audio."+std::to_string(layer_id)+ ".ffn.up_proj.input_min"); + assert(ffn_up_input_min.size() == 1); + this->audio_ffn_up_input_min.push_back(ffn_up_input_min[0]); + + buffer ffn_up_output_max; + q4nx.load_weights(ffn_up_output_max, + "model.audio."+std::to_string(layer_id)+ ".ffn.up_proj.output_max"); + assert(ffn_up_output_max.size() == 1); + this->audio_ffn_up_output_max.push_back(ffn_up_output_max[0]); + + buffer ffn_up_output_min; + q4nx.load_weights(ffn_up_output_min, + "model.audio."+std::to_string(layer_id)+ ".ffn.up_proj.output_min"); + assert(ffn_up_output_min.size() == 1); + this->audio_ffn_up_output_min.push_back(ffn_up_output_min[0]); + + this->audio_ffn_up_1_weight[layer_id] = this->ffn_up_proj_app.create_bo_buffer( + Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + q4nx.load_weights( + this->audio_ffn_up_1_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".ffn.up_proj_1.weight" + ); + buffer ffn_up_1_input_max; + q4nx.load_weights(ffn_up_1_input_max, + "model.audio."+std::to_string(layer_id)+ ".ffn.up_proj_1.input_max"); + assert(ffn_up_1_input_max.size() == 1); + this->audio_ffn_up_1_input_max.push_back(ffn_up_1_input_max[0]); + + buffer ffn_up_1_input_min; + q4nx.load_weights(ffn_up_1_input_min, + "model.audio."+std::to_string(layer_id)+ ".ffn.up_proj_1.input_min"); + assert(ffn_up_1_input_min.size() == 1); + this->audio_ffn_up_1_input_min.push_back(ffn_up_1_input_min[0]); + + buffer ffn_up_1_output_max; + q4nx.load_weights(ffn_up_1_output_max, + "model.audio."+std::to_string(layer_id)+ ".ffn.up_proj_1.output_max"); + assert(ffn_up_1_output_max.size() == 1); + this->audio_ffn_up_1_output_max.push_back(ffn_up_1_output_max[0]); + + buffer ffn_up_1_output_min; + q4nx.load_weights(ffn_up_1_output_min, + "model.audio."+std::to_string(layer_id)+ ".ffn.up_proj_1.output_min"); + assert(ffn_up_1_output_min.size() == 1); + this->audio_ffn_up_1_output_min.push_back(ffn_up_1_output_min[0]); + + q4nx.load_weights( + this->norm2_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".ln2.weight" + ); + + q4nx.load_weights( + this->per_dim_scale_with_softplus_weight[layer_id], + "model.audio."+std::to_string(layer_id)+".pre_dim_scale.weight" + ); + // Any optimization, multiple per_dim_scale with self.q_scale at load weights + for(int i = 0; i < this->per_dim_scale_with_softplus_weight[layer_id].size(); i++){ + this->per_dim_scale_with_softplus_weight[layer_id][i] = this->per_dim_scale_with_softplus_weight[layer_id][i] * Gemma4E_Audio_q_scale; + } + } +} + +std::vector Gemma4e_AudioEncoder::encode(void* audio_payload_ptr){ + + DEBUG_BLOCK(1, + std::cout << "Gemma4e_AudioEncoder::encode called with audio_payload_ptr: " << audio_payload_ptr << std::endl; + ) + + //DEBUG + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + SafeTensors reference_tensors( + this->model_path + "/audio_reference_data.safetensors" + ); + + #endif + + gemma4e_audio_payload_t* audio_payload = static_cast(audio_payload_ptr); + assert(audio_payload != nullptr); + + /* + // compare the weights + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "DEBUG: Start computing error metrics for audio encoder weights" << std::endl; + buffer reference_conv_weight_0; + buffer reference_conv_weight_1; + reference_tensors.load_weights(reference_conv_weight_0,"2D_convolution_weights_0"); + reference_tensors.load_weights(reference_conv_weight_1,"2D_convolution_weights_1"); + + assert(audio_subsample_conv2d_weight_0.size() == reference_conv_weight_0.size()); + assert(audio_subsample_conv2d_weight_1.size() == reference_conv_weight_1.size()); + print_error_metrics( + this->audio_subsample_conv2d_weight_0.data(), reference_conv_weight_0.data(), + 1, + this->audio_subsample_conv2d_weight_0.size(),1, + this->audio_subsample_conv2d_weight_0.size(),1 + ); + print_error_metrics( + this->audio_subsample_conv2d_weight_1.data(), reference_conv_weight_1.data(), + 1, + this->audio_subsample_conv2d_weight_1.size(), 1, + this->audio_subsample_conv2d_weight_1.size(), 1 + ); + } + #endif + */ + + // now, compare with Gemma4AudioSubSampleConvProjection_hidden_states_before_conv + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "DEBUG: Start computing error metrics for audio encoder input" << std::endl; + buffer reference_projection_input; + reference_tensors.load_weights(reference_projection_input,"Gemma4AudioSubSampleConvProjection_hidden_states_before_conv"); + size_t reference_projection_input_num_elements = reference_projection_input.size() / audio_payload->num_audios ; + + for(int i = 0; i < audio_payload->num_audios; i++){ + print_error_metrics( + audio_payload->mel_spectrograms[i].data(), + reference_projection_input.data() + i* reference_projection_input_num_elements , + 1, + audio_payload->mel_spectrogram_frames_per_audio[i], audio_payload->mel_spectrogram_bins_per_audio[i], + audio_payload->mel_spectrogram_frames_per_audio[i], audio_payload->mel_spectrogram_bins_per_audio[i] + ); + } + + //TODO: FIXME: remove it later + for(int i = 0; i < audio_payload->num_audios; i++){ + + for(int l = 0; l< audio_payload->mel_spectrogram_frames_per_audio[i]*audio_payload->mel_spectrogram_bins_per_audio[i]; l++ ){ + audio_payload->mel_spectrograms[i][l] = (bf16)(reference_projection_input[i* reference_projection_input_num_elements + l]); + } + } + } + #endif + + // sanity checks, ensure to be the same + int initial_audio_bins = audio_payload->mel_spectrogram_bins_per_audio[0]; + for(int i = 1; i < audio_payload->num_audios; i++){ + assert(initial_audio_bins == audio_payload->mel_spectrogram_bins_per_audio[i]); + } + + // recall the equation for calculation conv2d output as the following + //H_out = (H_in + 2*padding - K) / stride + 1 + //W_out = (W_in + 2*padding - K) / stride + 1 + auto calc_conv2d_out = [](int h_in, int padding, int k, int stride) -> int { + return (h_in + 2 * padding - k) / stride + 1; + }; + + std::vector> subSample_conv_layer_0_res( audio_payload->num_audios); + //does the first SubSample convlution operation + int audio_bin_after_conv = calc_conv2d_out( + initial_audio_bins, this->parent_npu_ptr->Gemma4e_Audio_conv2d_Padding, + this->parent_npu_ptr->Gemma4E_Audio_conv2d_kernel_size, this->parent_npu_ptr->Gemma4E_Audio_conv2d_Stride + ); + std::vector audio_frames_after_conv2d_0(audio_payload->num_audios); + + { + + for(int i = 0; i num_audios; i++){ + audio_frames_after_conv2d_0[i] = calc_conv2d_out( + audio_payload->mel_spectrogram_frames_per_audio[i], this->parent_npu_ptr->Gemma4e_Audio_conv2d_Padding, + this->parent_npu_ptr->Gemma4E_Audio_conv2d_kernel_size, this->parent_npu_ptr->Gemma4E_Audio_conv2d_Stride + ); + std::vector conv_out( + this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_0 * audio_frames_after_conv2d_0[i] * audio_bin_after_conv + ); + simd_conv2d( + audio_payload->mel_spectrograms[i].data(), + this->audio_subsample_conv2d_weight_0.data(), + conv_out.data(), + 1, + audio_payload->mel_spectrogram_frames_per_audio[i], audio_payload->mel_spectrogram_bins_per_audio[i], + this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_0, + this->parent_npu_ptr->Gemma4E_Audio_conv2d_kernel_size, + this->parent_npu_ptr->Gemma4E_Audio_conv2d_Stride, + this->parent_npu_ptr->Gemma4e_Audio_conv2d_Padding + + ); + // scalar_conv2d( + // audio_payload->mel_spectrograms[i].data(), + // this->audio_subsample_conv2d_weight_0.data(), + // conv_out.data(), + // 1, + // audio_payload->mel_spectrogram_frames_per_audio[i], audio_payload->mel_spectrogram_bins_per_audio[i], + // this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_0, + // this->parent_npu_ptr->Gemma4E_Audio_conv2d_kernel_size, + // this->parent_npu_ptr->Gemma4E_Audio_conv2d_Stride, + // this->parent_npu_ptr->Gemma4e_Audio_conv2d_Padding + // ); + subSample_conv_layer_0_res[i] = std::move(conv_out); + } + + /* + // #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + // { + + // size_t reference_projection_input_num_elements_per_chanel = reference_projection_input_num_elements / this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_0; + + // print_error_metrics( + // subSample_conv_layer_0_res[i].data() + c* audio_frames_after_conv2d_0[i]* audio_bin_after_conv, + // Gemma4AudioSubSampleConvProjectionLayer_0_hidden_states_after_conv.data() + i* reference_projection_input_num_elements + c* reference_projection_input_num_elements_per_chanel, + // 1, + // audio_frames_after_conv2d_0[i], audio_bin_after_conv, + // audio_frames_after_conv2d_0[i], audio_bin_after_conv + // ); + + // } + + // } + + // // //TODO: FIXME: + // // for(int i = 0; i < audio_payload->num_audios; i++){ + // // for(int c = 0; c parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_0; c++ ){ + // // memcpy( + // // subSample_conv_layer_0_res[i].data() + c* audio_frames_after_conv2d_0[i]* audio_bin_after_conv, + // // Gemma4AudioSubSampleConvProjectionLayer_0_hidden_states_after_conv.data() + i* reference_projection_input_num_elements + c* reference_projection_input_num_elements_per_chanel, + // // audio_frames_after_conv2d_0[i]* audio_bin_after_conv * sizeof(bf16) + // // ); + // // } + + // // } + + // } + + // #endif + + */ + + // now, perform the reorder with layernorm + // Python reference does: act(norm(hidden_states.permute(0,2,3,1)).permute(0,3,1,2)) + // Instead of permuting NCHW→NHWC, norming, permuting back, we apply LayerNorm + // directly over the channel dimension (stride = H*W) in NCHW layout, then ReLU. + + for(int i = 0; i < audio_payload->num_audios; i++){ + int H = audio_frames_after_conv2d_0[i]; + int W = audio_bin_after_conv; + int HW = H * W; + bf16* data = subSample_conv_layer_0_res[i].data(); + + // LayerNorm over C channels at each (h,w), then ReLU, all in NCHW layout + layernorm_relu_nchw(data, this->audio_subsample_conv2d_norm_weight_0.data(), + this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_0,HW, 1e-6f); + } + + // #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + // { + + // size_t reference_projection_input_num_elements_per_chanel = reference_projection_input_num_elements / this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_0; + + // print_error_metrics( + // subSample_conv_layer_0_res[i].data() + c* audio_frames_after_conv2d_0[i]* audio_bin_after_conv, + // Gemma4AudioSubSampleConvProjectionLayer_0_hidden_states_after_act.data() + i* reference_projection_input_num_elements + c* reference_projection_input_num_elements_per_chanel, + // 1, + // audio_frames_after_conv2d_0[i], audio_bin_after_conv, + // audio_frames_after_conv2d_0[i], audio_bin_after_conv + // ); + + // } + + // } + + // } + // #endif + } + + //subSample_conv_layer_0_res is [num_audio, this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_0, + // audio_frames_after_conv2d_0[i], audio_bin_after_conv ] + + // now, the second subSample_conv_layer_1 + int audio_bin_after_conv_1 = calc_conv2d_out( + audio_bin_after_conv, this->parent_npu_ptr->Gemma4e_Audio_conv2d_Padding, + this->parent_npu_ptr->Gemma4E_Audio_conv2d_kernel_size, this->parent_npu_ptr->Gemma4E_Audio_conv2d_Stride + ); + std::vector audio_frames_after_conv2d_1( audio_payload->num_audios); + std::vector> subSample_conv_layer_1_res( audio_payload->num_audios); + { + + for(int i = 0; i < audio_payload->num_audios; i++){ + audio_frames_after_conv2d_1[i] = calc_conv2d_out( + audio_frames_after_conv2d_0[i], this->parent_npu_ptr->Gemma4e_Audio_conv2d_Padding, + this->parent_npu_ptr->Gemma4E_Audio_conv2d_kernel_size, this->parent_npu_ptr->Gemma4E_Audio_conv2d_Stride + ); + + std::vector conv_out( + this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1 * audio_frames_after_conv2d_1[i] * audio_bin_after_conv_1 + ); + simd_conv2d( + subSample_conv_layer_0_res[i].data(), + this->audio_subsample_conv2d_weight_1.data(), + conv_out.data(), + this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_0, + audio_frames_after_conv2d_0[i], audio_bin_after_conv, + this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1, + this->parent_npu_ptr->Gemma4E_Audio_conv2d_kernel_size, + this->parent_npu_ptr->Gemma4E_Audio_conv2d_Stride, + this->parent_npu_ptr->Gemma4e_Audio_conv2d_Padding + ); + + // scalar_conv2d( + // subSample_conv_layer_0_res[i].data(), + // this->audio_subsample_conv2d_weight_1.data(), + // conv_out.data(), + // this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_0, + // audio_frames_after_conv2d_0[i], audio_bin_after_conv, + // this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1, + // this->parent_npu_ptr->Gemma4E_Audio_conv2d_kernel_size, + // this->parent_npu_ptr->Gemma4E_Audio_conv2d_Stride, + // this->parent_npu_ptr->Gemma4e_Audio_conv2d_Padding + // ); + subSample_conv_layer_1_res[i] = std::move(conv_out); + } + + // #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + // { + // buffer Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_conv; + // reference_tensors.load_weights(Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_conv,"Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_conv"); + // size_t reference_projection_input_num_elements = Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_conv.size() / audio_payload->num_audios ; + // std::cout << "DEBUG: Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_conv" << std::endl; + + // size_t reference_projection_input_num_elements_per_chanel = reference_projection_input_num_elements / this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1; + // for(int i = 0; i < audio_payload->num_audios; i++){ + // for(int c = 0; c parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1; c++ ){ + // print_error_metrics( + // subSample_conv_layer_1_res[i].data() + c* audio_frames_after_conv2d_1[i]* audio_bin_after_conv_1, + // Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_conv.data() + i* reference_projection_input_num_elements + c* reference_projection_input_num_elements_per_chanel, + // 1, + // audio_frames_after_conv2d_1[i], audio_bin_after_conv_1, + // audio_frames_after_conv2d_1[i], audio_bin_after_conv_1 + // ); + // } + // } + + // } + // #endif + + for(int i = 0; i < audio_payload->num_audios; i++){ + int H = audio_frames_after_conv2d_1[i]; + int W = audio_bin_after_conv_1; + int HW = H * W; + bf16* data = subSample_conv_layer_1_res[i].data(); + + // LayerNorm over C channels at each (h,w), then ReLU, all in NCHW layout + layernorm_relu_nchw(data, this->audio_subsample_conv2d_norm_weight_1.data(), + this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1,HW, 1e-6f); + } + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_act; + reference_tensors.load_weights(Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_act,"Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_act"); + size_t reference_projection_input_num_elements = Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_act.size() / audio_payload->num_audios ; + std::cout << "DEBUG: Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_act" << std::endl; + + size_t reference_projection_input_num_elements_per_chanel = reference_projection_input_num_elements / this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1; + + for(int i = 0; i < audio_payload->num_audios; i++){ + for(int c = 0; c parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1; c++ ){ + print_error_metrics( + subSample_conv_layer_1_res[i].data() + c* audio_frames_after_conv2d_1[i]* audio_bin_after_conv_1, + Gemma4AudioSubSampleConvProjectionLayer_1_hidden_states_after_act.data() + i* reference_projection_input_num_elements + c* reference_projection_input_num_elements_per_chanel, + 1, + audio_frames_after_conv2d_1[i], audio_bin_after_conv_1, + audio_frames_after_conv2d_1[i], audio_bin_after_conv_1 + ); + } + } + } + #endif + } + + // now, subSample_conv_layer_1_res is shape of [ num_audio, this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1, audio_frames_after_conv2d_1[i], audio_bin_after_conv_1 ] + + // reorder subSample_conv_layer_1_res.permute(0, 2, 3, 1).contiguous().reshape(batch_size, seq_len, -1) + // aka seq_len = audio_frames_after_conv2d_1[i] + // the hidden_Size is audio_bin_after_conv_1 *this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1 + // the reorder output to audio_embedding_projection_input, which is [num_audio, seq_len_per_audio, Padded_Gemma4E_Audio_Multimodal_Output_SIZE ] + // NOTE: Padded_Gemma4E_Audio_Multimodal_Output_SIZE>= audio_bin_after_conv_1*this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1 == Gemma4E_Audio_HIDDEN_SIZE + + std::vector seq_len_per_audio(audio_payload->num_audios); + std::vector start_seq_len_index_per_audio(audio_payload->num_audios); + int seq_len = 0; + int seq_len_of_last_audio = 0; + + for(int i = 0; i < audio_payload->num_audios; i++){ + seq_len_per_audio[i] = audio_frames_after_conv2d_1[i] ; + + seq_len += seq_len_per_audio[i]; + seq_len_of_last_audio = seq_len_per_audio[i]; + + start_seq_len_index_per_audio[i] = seq_len - seq_len_per_audio[i]; + } + int seq_len_padded = 0; + + //TODO: FIXME: padding for conv1d + + seq_len_padded = round_up_to_multiple(seq_len, seq_len_pad_requirement_for_MM); + + assert(this->parent_npu_ptr->Gemma4E_Audio_Multimodal_Output_SIZE % MM_tile_K == 0); + assert(this->parent_npu_ptr->Gemma4E_Audio_Multimodal_Output_SIZE % MM_tile_N == 0); + assert(this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE == audio_bin_after_conv_1*this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1); + assert(this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_0/4 * this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1 == + this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + for(int i = 0; i < audio_payload->num_audios; i++) { + std::cout << "DEBUG: seq_len_per_audio[" << i << "]: " << seq_len_per_audio[i] << std::endl; + std::cout << "DEBUG: start_seq_len_index_per_audio[" << i << "]: " << start_seq_len_index_per_audio[i] << std::endl; + } + std::cout << "DEBUG: seq_len: " << seq_len << std::endl; + std::cout << "DEBUG: seq_len_padded: " << seq_len_padded << std::endl; + } + + #endif + + buffer audio_embedding_projection_input = this->sub_sampleConvProjection_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(audio_embedding_projection_input.data() + seq_len * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + audio_embedding_projection_input.sync_to_device(); + + buffer audio_embedding_projection_output = this->sub_sampleConvProjection_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(audio_embedding_projection_output.data() + seq_len * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + audio_embedding_projection_output.sync_to_device(); + + buffer ffn_up_proj_input = this->ffn_up_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(ffn_up_proj_input.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + ffn_up_proj_input.sync_to_device(); + + buffer ffn_up_proj_output_down_input = this->ffn_up_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE + ); + memset(ffn_up_proj_output_down_input.data() + seq_len * Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE * sizeof(bf16) + ); + ffn_up_proj_output_down_input.sync_to_device(); + + buffer ffn_down_proj_output = this->ffn_down_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(ffn_down_proj_output.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + ffn_down_proj_output.sync_to_device(); + + buffer q_proj_input = this->q_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(q_proj_input.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + q_proj_input.sync_to_device(); + + buffer q_proj_output = this->q_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(q_proj_output.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + q_proj_output.sync_to_device(); + + buffer k_proj_input = this->k_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(k_proj_input.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + k_proj_input.sync_to_device(); + + buffer k_proj_output = this->k_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(k_proj_output.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + k_proj_output.sync_to_device(); + + buffer v_proj_input = this->v_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(v_proj_input.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + v_proj_input.sync_to_device(); + + buffer v_proj_output = this->v_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(v_proj_output.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + v_proj_output.sync_to_device(); + + buffer o_output_proj_input = this->o_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(o_output_proj_input.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + o_output_proj_input.sync_to_device(); + + buffer o_output_proj_output = this->o_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(o_output_proj_output.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + o_output_proj_output.sync_to_device(); + + buffer conv1d_start_proj_input = this->conv1d_start_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(conv1d_start_proj_input.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + conv1d_start_proj_input.sync_to_device(); + + buffer conv1d_start_proj_output = this->conv1d_start_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE + ); + memset(conv1d_start_proj_output.data() + seq_len * Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE * sizeof(bf16) + ); + conv1d_start_proj_output.sync_to_device(); + + buffer audio_conv1d_input = this->conv1d_app.create_bo_buffer( + (seq_len_padded + audio_payload->num_audios* this->Gemma4E_Audio_padded_requirement_for_conv1d) * Padded_GEMMA4E_Audio_HIDDEN_SIZE + + ); + memset(audio_conv1d_input.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len + audio_payload->num_audios* this->Gemma4E_Audio_padded_requirement_for_conv1d) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + audio_conv1d_input.sync_to_device(); + + buffer audio_conv1d_output = this->conv1d_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + + ); + memset(audio_conv1d_output.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + audio_conv1d_output.sync_to_device(); + + buffer conv1d_end_proj_input = this->conv1d_end_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(conv1d_end_proj_input.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + conv1d_end_proj_input.sync_to_device(); + + buffer conv1d_end_proj_output = this->conv1d_end_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(conv1d_end_proj_output.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + conv1d_end_proj_output.sync_to_device(); + + buffer audio_pre_encode_input = this->audio_pre_encode_proj_app.create_bo_buffer( + seq_len_padded * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + memset(audio_pre_encode_input.data() + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 0, (seq_len_padded - seq_len) * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + ); + audio_pre_encode_input.sync_to_device(); + + buffer audio_pre_encode_output = this->audio_pre_encode_proj_app.create_bo_buffer( + seq_len_padded * Padded_Gemma4E_Audio_Multimodal_Output_SIZE + ); + memset(audio_pre_encode_output.data() + seq_len * Padded_Gemma4E_Audio_Multimodal_Output_SIZE, + 0, (seq_len_padded - seq_len) * Padded_Gemma4E_Audio_Multimodal_Output_SIZE * sizeof(bf16) + ); + audio_pre_encode_output.sync_to_device(); + + buffer audio_to_language_project_input = this->audio_to_language_proj_app.create_bo_buffer( + seq_len_padded * Padded_Gemma4E_Audio_Multimodal_Output_SIZE + ); + memset(audio_to_language_project_input.data() + seq_len * Padded_Gemma4E_Audio_Multimodal_Output_SIZE, + 0, (seq_len_padded - seq_len) * Padded_Gemma4E_Audio_Multimodal_Output_SIZE * sizeof(bf16) + ); + audio_to_language_project_input.sync_to_device(); + + buffer audio_to_language_project_output = this->audio_to_language_proj_app.create_bo_buffer( + seq_len_padded * parent_npu_ptr->Gemma4E_Audio_language_projection_output_size + ); + memset(audio_to_language_project_output.data() + seq_len * parent_npu_ptr->Gemma4E_Audio_language_projection_output_size, + 0, (seq_len_padded - seq_len) * parent_npu_ptr->Gemma4E_Audio_language_projection_output_size * sizeof(bf16) + ); + audio_to_language_project_output.sync_to_device(); + assert(parent_npu_ptr->Gemma4E_Audio_language_projection_output_size % MM_tile_N == 0); + + { + generate_mm_sequence( + *this->sub_sampleConvProjection_app.seq(), + seq_len_padded, this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, + MM_tile_M, MM_tile_K, MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, //no bias, no activation + 0, -10000.0, 1000000.0, // do not clamp on output + ENABLE_QKV_REORDER, 0// since we don't need it anymore + + ); + + generate_mm_sequence( + *this->audio_pre_encode_proj_app.seq(), + seq_len_padded, this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, this->Padded_Gemma4E_Audio_Multimodal_Output_SIZE, + MM_tile_M, MM_tile_K, MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + true, 0, //no bias, no activation + 0, -10000.0, 1000000.0, // do not clamp on output + ENABLE_QKV_REORDER, 0// since we don't need it anymore + + ); + generate_mm_sequence( + *this->audio_to_language_proj_app.seq(), + seq_len_padded, Padded_Gemma4E_Audio_Multimodal_Output_SIZE, parent_npu_ptr->Gemma4E_Audio_language_projection_output_size, + MM_tile_M, MM_tile_K, MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, //no bias, no activation + 0, -10000.0, 1000000.0, // do not clamp on output + ENABLE_QKV_REORDER, 0// since we don't need it anymore + + ); + } + + memset(audio_embedding_projection_input.data(), 0, audio_embedding_projection_input.size() * sizeof(bf16)); + + // reorder of subSample_conv_layer_1_res.permute(0, 2, 3, 1).contiguous().reshape(batch_size, seq_len, -1) + for(int i = 0; i < audio_payload->num_audios; i++){ + + for(int h = 0; h < audio_frames_after_conv2d_1[i]; h++){ + for(int w = 0; w < audio_bin_after_conv_1; w++){ + for(int c = 0; c < this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1; c++){ + size_t input_index = c* audio_frames_after_conv2d_1[i]* audio_bin_after_conv_1 + h* audio_bin_after_conv_1 + w; + + size_t output_index = start_seq_len_index_per_audio[i] * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE + h* this->Padded_GEMMA4E_Audio_HIDDEN_SIZE\ + + w* this->parent_npu_ptr->Gemma4E_Audio_subsampling_conv_channels_1 + c; + audio_embedding_projection_input[output_index] = subSample_conv_layer_1_res[i][input_index]; + } + } + } + } + // subSample_conv_layer_1_res is now shape of [ seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE ] + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioSubSampleConvProjection_hidden_states_before_linear; + reference_tensors.load_weights(Gemma4AudioSubSampleConvProjection_hidden_states_before_linear,"Gemma4AudioSubSampleConvProjection_hidden_states_before_linear"); + size_t reference_projection_input_num_elements = Gemma4AudioSubSampleConvProjection_hidden_states_before_linear.size() / audio_payload->num_audios ; + std::cout << "DEBUG: Gemma4AudioSubSampleConvProjection_hidden_states_before_linear" << std::endl; + + // for(int i = 0; i < audio_payload->num_audios; i++){ + // print_error_metrics( + // audio_embedding_projection_input.data() + start_seq_len_index_per_audio[i] * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, + // Gemma4AudioSubSampleConvProjection_hidden_states_before_linear.data() + i* reference_projection_input_num_elements , + // 1, + // seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + // seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + // ); + // } + // //TODO:FIXME: remove it later + + // memcpy( + // seq_len_per_audio[i]* Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + } + #endif + + // debugt + assert(audio_embedding_projection_weight.size() == this->Padded_GEMMA4E_Audio_HIDDEN_SIZE* this->Padded_GEMMA4E_Audio_HIDDEN_SIZE); + DEBUG_BLOCK(1, + std::cout << "Padded_GEMMA4E_Audio_HIDDEN_SIZE : " << this->Padded_GEMMA4E_Audio_HIDDEN_SIZE << std::endl; + ) + + audio_embedding_projection_input.sync_to_device(); + this->audio_embedding_projection_weight.sync_to_device(); + FLM_OVERRIDE(audio_sub_sample_proj, sub_sampleConvProjection_app( audio_embedding_projection_input, audio_embedding_projection_weight, audio_embedding_projection_output)); + audio_embedding_projection_output.sync_from_device(); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioSubSampleConvProjection_hidden_states_after_linear; + reference_tensors.load_weights(Gemma4AudioSubSampleConvProjection_hidden_states_after_linear,"Gemma4AudioSubSampleConvProjection_hidden_states_after_linear"); + size_t reference_projection_input_num_elements = Gemma4AudioSubSampleConvProjection_hidden_states_after_linear.size() / audio_payload->num_audios ; + std::cout << "DEBUG: Gemma4AudioSubSampleConvProjection_hidden_states_after_linear" << std::endl; + + // for(int i = 0; i < audio_payload->num_audios; i++){ + // print_error_metrics( + // audio_embedding_projection_output.data() + start_seq_len_index_per_audio[i] * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, + // Gemma4AudioSubSampleConvProjection_hidden_states_after_linear.data() + i* reference_projection_input_num_elements , + // 1, + // seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + // seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + // ); + // } + } + #endif + + std::vector position_embedding(13 *this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE ); + generate_gemma4_audio_rotary_pos_emb( + this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + this->parent_npu_ptr->Gemma4E_Audio_attention_chunk_size, + this->parent_npu_ptr->Gemma4E_Audio_attention_context_left, + this->parent_npu_ptr->Gemma4E_Audio_attention_context_right, + position_embedding + ); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer ref_position_embeddings; + reference_tensors.load_weights(ref_position_embeddings,"position_embeddings"); + + std::cout << "DEBUG: position_embeddings" << std::endl; + + print_error_metrics( + position_embedding.data(), ref_position_embeddings.data(), + 1, + position_embedding.size(), 1, + position_embedding.size(), 1 + ); + } + #endif + + std::vector> audio_sliding_window_attention_mask(audio_payload->num_audios); + for(int i = 0; i < audio_payload->num_audios; i++){ + create_sliding_window_attention_mask( + seq_len_per_audio[i], + this->parent_npu_ptr->Gemma4E_Audio_attention_context_left-1, + this->parent_npu_ptr->Gemma4E_Audio_attention_context_right, + audio_sliding_window_attention_mask[i] + + ); + } + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer attention_mask_before_block_5d; + reference_tensors.load_weights(attention_mask_before_block_5d,"attention_mask_before_block_5d"); + + size_t reference_attention_mask_num_elements = attention_mask_before_block_5d.size() / audio_payload->num_audios ; + size_t reference_attention_mask_length = std::sqrt(reference_attention_mask_num_elements); + + std::cout << "DEBUG: attention_mask_before_block_5d" << std::endl; + + //NOTE: because the reference attention mask is >= each seq_len_per_audio[i] + // For ease of comparison, we create same mask size and load audio_sliding_windo_attention_mask to it + std::vector expanded_attention_mask(attention_mask_before_block_5d.size(), 0); + + for(int i = 0; i < audio_payload->num_audios; i++){ + int* mask_start_ptr = expanded_attention_mask.data() + i* reference_attention_mask_num_elements; + // Fill valid rows (0 to seq_len_per_audio[i]-1) from the C++ sliding window mask + for(int l = 0; l < seq_len_per_audio[i]; l++){ + memcpy( + mask_start_ptr + l* reference_attention_mask_length, + audio_sliding_window_attention_mask[i].data() + l* seq_len_per_audio[i], + seq_len_per_audio[i] * sizeof(int) + ); + } + // Fill padding rows (seq_len_per_audio[i] to reference_attention_mask_length-1) + // Python's create_bidirectional_mask doesn't mask queries, only keys. + // So padding rows still have 1s for valid columns within the sliding window. + int sliding_window_left = this->parent_npu_ptr->Gemma4E_Audio_attention_context_left - 1; + int sliding_window_right = this->parent_npu_ptr->Gemma4E_Audio_attention_context_right; + for(int l = seq_len_per_audio[i]; l < (int)reference_attention_mask_length; l++){ + for(int c = 0; c < seq_len_per_audio[i]; c++){ + int dist = l - c; + bool left_mask = (dist >= 0) && (dist < sliding_window_left); + bool right_mask = (dist < 0) && (-dist < sliding_window_right); + if(left_mask || right_mask){ + mask_start_ptr[l * reference_attention_mask_length + c] = 1; + } + } + } + } + print_error_metrics( + expanded_attention_mask.data(), attention_mask_before_block_5d.data(), + 1, + expanded_attention_mask.size(), 1, + expanded_attention_mask.size(), 1 + ); + } + #endif + + // now, convert to block attention mask + // [batch_Size, num_blocks, chunk_size, context_size] + std::vector> block_attention_mask_per_audio(audio_payload->num_audios); + std::vector num_blocks_per_audio(audio_payload->num_audios); + std::vector context_size_per_audio(audio_payload->num_audios); + for(int i = 0; i < audio_payload->num_audios; i++){ + convert_mask_to_blocked( + audio_sliding_window_attention_mask[i], + this->parent_npu_ptr->Gemma4E_Audio_attention_chunk_size, + this->parent_npu_ptr->Gemma4E_Audio_attention_context_left, + this->parent_npu_ptr->Gemma4E_Audio_attention_context_right, + block_attention_mask_per_audio[i], + num_blocks_per_audio[i], context_size_per_audio[i] + ); + } + + hidden_state.resize(seq_len_padded * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, 0.0f); + residual.resize(seq_len_padded * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, 0.0f); + + memcpy(hidden_state.data(), audio_embedding_projection_output.data(), seq_len * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16)); + for(int layer_id = 0; layer_id Gemma4E_Audio_num_attention_layers; layer_id++){ + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + //Gemma4AudioLayer_{layer_idx}_hidden_states_before_ffw1 + buffer Gemma4AudioLayer_hidden_State_before_ffw1; + reference_tensors.load_weights(Gemma4AudioLayer_hidden_State_before_ffw1,"Gemma4AudioLayer_"+std::to_string(layer_id)+"_hidden_states_before_ffw1"); + size_t reference_ffn_input_num_elements = Gemma4AudioLayer_hidden_State_before_ffw1.size() / audio_payload->num_audios ; + std::cout << "DEBUG: Gemma4AudioLayer_hidden_State_before_ffw1" << std::endl; + + for(int i = 0; i < audio_payload->num_audios; i++){ + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_audio[i] * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Gemma4AudioLayer_hidden_State_before_ffw1.data() + i* reference_ffn_input_num_elements , + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + + // // TODO: FIXME: remove later + + // for(int i = 0; i < audio_payload->num_audios; i++){ + // memcpy( + // hidden_state.data() + start_seq_len_index_per_audio[i] * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, + // Gemma4AudioLayer_hidden_State_before_ffw1.data() + i* reference_ffn_input_num_elements , + // seq_len_per_audio[i] * Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16) + // ); + // } + // memset(hidden_state.data() + seq_len * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, 0, + // (seq_len_padded - seq_len)* this->Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16)); + } + #endif + + ffn_layer( + ffn_up_proj_input, ffn_up_proj_output_down_input, ffn_down_proj_output, + ffn_norm_weight[layer_id], + ffn_post_norm_weight[layer_id], + audio_ffn_up_weight[layer_id], audio_ffn_down_weight[layer_id], + seq_len, seq_len_padded, + + audio_ffn_up_input_min[layer_id], audio_ffn_up_input_max[layer_id], + audio_ffn_up_output_min[layer_id], audio_ffn_up_output_max[layer_id], + audio_ffn_down_input_min[layer_id], audio_ffn_down_input_max[layer_id], + audio_ffn_down_output_min[layer_id], audio_ffn_down_output_max[layer_id] + + ); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioLayer_hidden_states_after_ffn1; + reference_tensors.load_weights(Gemma4AudioLayer_hidden_states_after_ffn1, + "Gemma4AudioLayer_"+std::to_string(layer_id)+"_hidden_states_after_ffw1"); + size_t reference_attn_output_num_elements = + Gemma4AudioLayer_hidden_states_after_ffn1.size() / audio_payload->num_audios; + std::cout << "DEBUG: Gemma4AudioLayer " +std::to_string(layer_id)+ "hidden_states_after_ffw1" << std::endl; + + for (int i = 0; i < audio_payload->num_audios; i++) { + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_audio[i] * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Gemma4AudioLayer_hidden_states_after_ffn1.data() + i * reference_attn_output_num_elements, + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif + + memcpy(residual.data(), hidden_state.data(), seq_len * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16)); + + simd_rms_norm( + hidden_state.data(), + this->attn_pre_norm_weight[layer_id].data(), + hidden_state.data(), + seq_len, + this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + + memcpy(q_proj_input.data(), hidden_state.data(), seq_len * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16)); + simd_clamp( + q_proj_input.data(), + q_proj_input.data(), + audio_q_input_min[layer_id], audio_q_input_max[layer_id], + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + + { + generate_mm_sequence( + *this->q_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, Padded_GEMMA4E_Audio_HIDDEN_SIZE, + MM_tile_M, MM_tile_K, MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, //no bias, no activation + 1, audio_q_output_min[layer_id], audio_q_output_max[layer_id], // do not clamp on output + ENABLE_QKV_REORDER, 0// since we don't need it anymore + + ); + } + + auto q_proj_run = FLM_OVERRIDE(audio_q_proj, q_proj_app.create_run( + q_proj_input, this->audio_attn_q_weight[layer_id], q_proj_output + ), layer_id); + + q_proj_input.sync_to_device(); + this->audio_attn_q_weight[layer_id].sync_to_device(); + q_proj_run.start(); + + // setup for K + generate_mm_sequence( + *this->k_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, Padded_GEMMA4E_Audio_HIDDEN_SIZE, + MM_tile_M, MM_tile_K, MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, //no bias, no activation + 1, audio_k_output_min[layer_id], audio_k_output_max[layer_id], // do not clamp on output + ENABLE_QKV_REORDER, 0// since we don't need it anymore + + ); + memcpy(k_proj_input.data(), hidden_state.data(), seq_len * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16)); + simd_clamp( + k_proj_input.data(), + k_proj_input.data(), + audio_k_input_min[layer_id], audio_k_input_max[layer_id], + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + k_proj_input.sync_to_device(); + + auto k_proj_run = FLM_OVERRIDE(audio_k_proj, k_proj_app.create_run( + k_proj_input, this->audio_attn_k_weight[layer_id], k_proj_output + ), layer_id); + + q_proj_run.wait(); + q_proj_output.sync_from_device(); + + k_proj_input.sync_to_device(); + this->audio_attn_k_weight[layer_id].sync_to_device(); + k_proj_run.start(); + + generate_mm_sequence( + *this->v_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, Padded_GEMMA4E_Audio_HIDDEN_SIZE, + MM_tile_M, MM_tile_K, MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, //no bias, no activation + 1, audio_v_output_min[layer_id], audio_v_output_max[layer_id], // do not clamp on output + ENABLE_QKV_REORDER, 0// since we don't need it anymore + + ); + memcpy(v_proj_input.data(), hidden_state.data(), seq_len * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE * sizeof(bf16)); + simd_clamp( + v_proj_input.data(), + v_proj_input.data(), + audio_v_input_min[layer_id], audio_v_input_max[layer_id], + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + v_proj_input.sync_to_device(); + auto v_proj_run = FLM_OVERRIDE(audio_v_proj, v_proj_app.create_run( + v_proj_input, this->audio_attn_v_weight[layer_id], v_proj_output + ), layer_id); + k_proj_run.wait(); + k_proj_output.sync_from_device(); + + v_proj_input.sync_to_device(); + this->audio_attn_v_weight[layer_id].sync_to_device(); + v_proj_run.start(); + // At this point, q, k proj_output is shape of [num_audio, seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE] + // But can also be viewed as [num_audio, seq_len_per_audio[i], Padded_Gemma4E_Audio_num_attention_heads, Gemma4E_Audio_attention_head_dim] + for(int b = 0; b < audio_payload->num_audios; b++){ + + for(int s= 0; s< seq_len_per_audio[b]; s++){ + + for(int h_idx = 0; h_idx < parent_npu_ptr->Gemma4E_Audio_num_attention_heads; h_idx++){ + + size_t offset = (start_seq_len_index_per_audio[b]+s) * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE \ + + h_idx * Gemma4E_Audio_attention_head_dim; + simd_mul( + q_proj_output.data() + offset, + per_dim_scale_with_softplus_weight[layer_id].data(), + q_proj_output.data() + offset, + Gemma4E_Audio_attention_head_dim + ); + } + } + } + q_proj_output.sync_to_device(); + simd_mul( + k_proj_output.data(), + this->Gemma4E_Audio_k_scale, + k_proj_output.data(), + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + k_proj_output.sync_to_device(); + + // setup o_projection + generate_mm_sequence( + *this->o_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, Padded_GEMMA4E_Audio_HIDDEN_SIZE, + MM_tile_M, MM_tile_K, MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, //no bias, no activation + 1, audio_o_output_min[layer_id], audio_o_output_max[layer_id], // do not clamp on output + ENABLE_QKV_REORDER, 0// since we don't need it anymore + + ); + v_proj_run.wait(); + v_proj_output.sync_from_device(); + + // now, compare q, k, v with + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioAttention_query_states_after_scaling; + buffer Gemma4AudioAttention_key_states_after_scaling; + buffer Gemma4AudioAttention_value_states_after_scaling; + + reference_tensors.load_weights(Gemma4AudioAttention_query_states_after_scaling,"Gemma4AudioAttention_"+std::to_string(layer_id)+"_query_states_after_scaling"); + reference_tensors.load_weights(Gemma4AudioAttention_key_states_after_scaling,"Gemma4AudioAttention_"+std::to_string(layer_id)+"_key_states_after_scaling"); + reference_tensors.load_weights(Gemma4AudioAttention_value_states_after_scaling,"Gemma4AudioAttention_"+std::to_string(layer_id)+"_value_states_after_scaling"); + + size_t reference_qkv_num_elements = Gemma4AudioAttention_query_states_after_scaling.size() / audio_payload->num_audios ; + std::cout << "DEBUG: Gemma4AudioAttention_query_states_after_scaling" << std::endl; + + for(int i = 0; i < audio_payload->num_audios; i++){ + print_error_metrics( + q_proj_output.data() + start_seq_len_index_per_audio[i] * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Gemma4AudioAttention_query_states_after_scaling.data() + i* reference_qkv_num_elements , + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + print_error_metrics( + k_proj_output.data() + start_seq_len_index_per_audio[i] * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Gemma4AudioAttention_key_states_after_scaling.data() + i* reference_qkv_num_elements , + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + print_error_metrics( + v_proj_output.data() + start_seq_len_index_per_audio[i] * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Gemma4AudioAttention_value_states_after_scaling.data() + i* reference_qkv_num_elements , + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif + + int hidden_size = this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE; + int num_positions = (int)(position_embedding.size() / hidden_size); + + memset(o_output_proj_input.data(), 0, o_output_proj_input.size() * sizeof(bf16)); + + compute_audio_self_attention( + q_proj_output.data(), + k_proj_output.data(), + v_proj_output.data(), + this->audio_attn_k_rel_weight[layer_id].data(), + position_embedding.data(), + block_attention_mask_per_audio, + num_blocks_per_audio, + context_size_per_audio, + seq_len_per_audio, + start_seq_len_index_per_audio, + o_output_proj_input.data(), + audio_payload->num_audios, + this->parent_npu_ptr->Gemma4E_Audio_attention_chunk_size, + this->parent_npu_ptr->Gemma4E_Audio_attention_context_left, + this->parent_npu_ptr->Gemma4E_Audio_attention_context_right, + this->parent_npu_ptr->Gemma4E_Audio_num_attention_heads, + Gemma4E_Audio_attention_head_dim, + hidden_size, + Padded_GEMMA4E_Audio_HIDDEN_SIZE, + num_positions, + parent_npu_ptr->Gemma4E_Audio_attention_softcap, + -1e9f // invalid_logits_value + ); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioAttention_attn_output_before_post; + reference_tensors.load_weights(Gemma4AudioAttention_attn_output_before_post, + "Gemma4AudioAttention_"+std::to_string(layer_id)+"_attn_output_before_post"); + size_t reference_attn_output_num_elements = + Gemma4AudioAttention_attn_output_before_post.size() / audio_payload->num_audios; + std::cout << "DEBUG: Gemma4AudioAttention_attn_output_before_post" << std::endl; + + for (int i = 0; i < audio_payload->num_audios; i++) { + print_error_metrics( + o_output_proj_input.data() + start_seq_len_index_per_audio[i] * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Gemma4AudioAttention_attn_output_before_post.data() + i * reference_attn_output_num_elements, + 1, + seq_len_per_audio[i], hidden_size, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif + + // o_proj + simd_clamp( + o_output_proj_input.data(), + o_output_proj_input.data(), + audio_o_input_min[layer_id], audio_o_input_max[layer_id], + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + o_output_proj_input.sync_to_device(); + this->audio_attn_o_weight[layer_id].sync_to_device(); + FLM_OVERRIDE(audio_o_proj, o_proj_app(o_output_proj_input, this->audio_attn_o_weight[layer_id], o_output_proj_output), layer_id); + o_output_proj_output.sync_from_device(); + + simd_rms_norm( + o_output_proj_output.data(), + this->attn_post_norm_weight[layer_id].data(), + hidden_state.data(), + seq_len, + this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 1e-6f + ); + + simd_add(hidden_state.data(), residual.data(), hidden_state.data(), + seq_len * Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + + // now compare it + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioAttention_hidden_states_before_conv1d; + reference_tensors.load_weights(Gemma4AudioAttention_hidden_states_before_conv1d, + "Gemma4AudioLayer_"+std::to_string(layer_id)+"_hidden_states_before_conv1d"); + size_t reference_attn_output_num_elements = + Gemma4AudioAttention_hidden_states_before_conv1d.size() / audio_payload->num_audios; + std::cout << "DEBUG: Gemma4AudioLayer__hidden_states_before_conv1d" << std::endl; + + for (int i = 0; i < audio_payload->num_audios; i++) { + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_audio[i] * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Gemma4AudioAttention_hidden_states_before_conv1d.data() + i * reference_attn_output_num_elements, + 1, + seq_len_per_audio[i], hidden_size, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif + conv1d_layer( + layer_id, + seq_len, seq_len_padded, + seq_len_per_audio, start_seq_len_index_per_audio, + conv1d_start_proj_input, conv1d_start_proj_output, + audio_conv1d_input, audio_conv1d_output, + conv1d_end_proj_input, conv1d_end_proj_output, + this->audio_conv_pw1_input_min[layer_id], this->audio_conv_pw1_input_max[layer_id], + this->audio_conv_pw1_output_min[layer_id], this->audio_conv_pw1_output_max[layer_id], + this->audio_conv_pw2_input_min[layer_id], this->audio_conv_pw2_input_max[layer_id], + this->audio_conv_pw2_output_min[layer_id], this->audio_conv_pw2_output_max[layer_id], + audio_payload, + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + &reference_tensors + #else + nullptr + #endif + ); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioLayer_hidden_states_after_conv1d; + reference_tensors.load_weights(Gemma4AudioLayer_hidden_states_after_conv1d, + "Gemma4AudioLayer_"+std::to_string(layer_id)+"_hidden_states_after_conv1d"); + size_t reference_attn_output_num_elements = + Gemma4AudioLayer_hidden_states_after_conv1d.size() / audio_payload->num_audios; + std::cout << "DEBUG: Gemma4AudioLayer_hidden_states_after_conv1d" << std::endl; + + for (int i = 0; i < audio_payload->num_audios; i++) { + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_audio[i] * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Gemma4AudioLayer_hidden_states_after_conv1d.data() + i * reference_attn_output_num_elements, + 1, + seq_len_per_audio[i], hidden_size, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif + + ffn_layer( + + ffn_up_proj_input, ffn_up_proj_output_down_input, ffn_down_proj_output, + ffn_norm_1_weight[layer_id], + ffn_post_norm_1_weight[layer_id], + audio_ffn_up_1_weight[layer_id], audio_ffn_down_1_weight[layer_id], + seq_len, seq_len_padded, + + audio_ffn_up_1_input_min[layer_id], audio_ffn_up_1_input_max[layer_id], + audio_ffn_up_1_output_min[layer_id], audio_ffn_up_1_output_max[layer_id], + audio_ffn_down_1_input_min[layer_id], audio_ffn_down_1_input_max[layer_id], + audio_ffn_down_1_output_min[layer_id], audio_ffn_down_1_output_max[layer_id] + + ); + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioLayer_hidden_states_after_ffn1; + reference_tensors.load_weights(Gemma4AudioLayer_hidden_states_after_ffn1, + "Gemma4AudioLayer_"+std::to_string(layer_id)+"_hidden_states_after_ffw2"); + size_t reference_attn_output_num_elements = + Gemma4AudioLayer_hidden_states_after_ffn1.size() / audio_payload->num_audios; + std::cout << "DEBUG: Gemma4AudioLayer_hidden_states_after_ffw2" << std::endl; + + for (int i = 0; i < audio_payload->num_audios; i++) { + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_audio[i] * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Gemma4AudioLayer_hidden_states_after_ffn1.data() + i * reference_attn_output_num_elements, + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif + + simd_rms_norm( + hidden_state.data(), norm2_weight[layer_id].data(), hidden_state.data(), + seq_len, this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, seq_len_padded, Padded_GEMMA4E_Audio_HIDDEN_SIZE, + 1e-6f + ); + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioLayer_hidden_states_after_norm_out; + reference_tensors.load_weights(Gemma4AudioLayer_hidden_states_after_norm_out, + "Gemma4AudioLayer_"+std::to_string(layer_id)+"_hidden_states_after_norm_out"); + size_t reference_attn_output_num_elements = + Gemma4AudioLayer_hidden_states_after_norm_out.size() / audio_payload->num_audios; + std::cout << "DEBUG: Gemma4AudioLayer" + std::to_string(layer_id)+ "_hidden_states_after_norm_out" << std::endl; + + for (int i = 0; i < audio_payload->num_audios; i++) { + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_audio[i] * Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Gemma4AudioLayer_hidden_states_after_norm_out.data() + i * reference_attn_output_num_elements, + 1, + seq_len_per_audio[i], hidden_size, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif + } + + memcpy(audio_pre_encode_input.data(), hidden_state.data(), + seq_len* this->Padded_GEMMA4E_Audio_HIDDEN_SIZE* sizeof(bf16) + ); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioPreEncodeProjection_hidden_states_before_linear; + reference_tensors.load_weights(Gemma4AudioPreEncodeProjection_hidden_states_before_linear,"hidden_states_before_output_proj"); + size_t reference_projection_input_num_elements = Gemma4AudioPreEncodeProjection_hidden_states_before_linear.size() / audio_payload->num_audios ; + std::cout << "DEBUG: hidden_states_before_output_proj" << std::endl; + + for(int i = 0; i < audio_payload->num_audios; i++){ + print_error_metrics( + audio_pre_encode_input.data() + start_seq_len_index_per_audio[i] * this->Padded_GEMMA4E_Audio_HIDDEN_SIZE, + Gemma4AudioPreEncodeProjection_hidden_states_before_linear.data() + i* reference_projection_input_num_elements , + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_HIDDEN_SIZE, + seq_len_per_audio[i], Padded_GEMMA4E_Audio_HIDDEN_SIZE + ); + } + } + #endif + + audio_pre_encode_input.sync_to_device(); + this->audio_pre_encode_weight.sync_to_device(); + FLM_OVERRIDE(audio_pre_encode_proj, audio_pre_encode_proj_app( audio_pre_encode_input, audio_pre_encode_weight, audio_pre_encode_output)); + audio_pre_encode_output.sync_from_device(); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer Gemma4AudioPreEncodeProjection_output; + reference_tensors.load_weights(Gemma4AudioPreEncodeProjection_output,"audio_model_output_proj_result"); + size_t reference_projection_input_num_elements = Gemma4AudioPreEncodeProjection_output.size() / audio_payload->num_audios ; + std::cout << "DEBUG: audio_model_output_proj_result" << std::endl; + + for(int i = 0; i < audio_payload->num_audios; i++){ + print_error_metrics( + audio_pre_encode_output.data() + start_seq_len_index_per_audio[i] * this->Padded_Gemma4E_Audio_Multimodal_Output_SIZE, + Gemma4AudioPreEncodeProjection_output.data() + i* reference_projection_input_num_elements , + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_Multimodal_Output_SIZE, + seq_len_per_audio[i], Padded_Gemma4E_Audio_Multimodal_Output_SIZE + ); + } + } + #endif + + simd_rms_norm( + audio_pre_encode_output.data(), + audio_to_language_project_input.data(), + seq_len, this->parent_npu_ptr->Gemma4E_Audio_Multimodal_Output_SIZE, seq_len_padded, Padded_Gemma4E_Audio_Multimodal_Output_SIZE, + 1e-6f + ); + audio_pre_encode_output.sync_to_device(); + + audio_to_language_project_input.sync_to_device(); + this->audio_to_language_projection_weight.sync_to_device(); + FLM_OVERRIDE(audio_to_language_proj, audio_to_language_proj_app(audio_to_language_project_input, this->audio_to_language_projection_weight, audio_to_language_project_output)); + audio_to_language_project_output.sync_from_device(); + + #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + { + buffer audio_to_language_projection_output_ref; + reference_tensors.load_weights(audio_to_language_projection_output_ref,"audio_final_embs_after_embed_audio"); + size_t reference_projection_input_num_elements = audio_to_language_projection_output_ref.size() / audio_payload->num_audios ; + std::cout << "DEBUG: audio_final_embs_after_embed_audio" << std::endl; + + for(int i = 0; i < audio_payload->num_audios; i++){ + print_error_metrics( + audio_to_language_project_output.data() + start_seq_len_index_per_audio[i] * this->parent_npu_ptr->Gemma4E_Audio_language_projection_output_size, + audio_to_language_projection_output_ref.data() + i* reference_projection_input_num_elements , + 1, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_language_projection_output_size, + seq_len_per_audio[i], this->parent_npu_ptr->Gemma4E_Audio_language_projection_output_size + ); + } + } + #endif + + std::vector dummy_output(audio_to_language_project_output.size()); // return half of the hidden size as dummy output + memcpy(dummy_output.data(), audio_to_language_project_output.data(), audio_to_language_project_output.size() * sizeof(bf16)); + return dummy_output; +} + diff --git a/src/detail/gemma4e_npu/gemma4e_audio.hpp b/src/detail/gemma4e_npu/gemma4e_audio.hpp new file mode 100644 index 000000000..e7f3dfdb4 --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_audio.hpp @@ -0,0 +1,199 @@ +#pragma once +#include "typedef.hpp" +#include +#include +#include "tensor_utils/q4_npu_eXpress.hpp" +#include "models/gemma4e/flm/aie2p/gemma4e_npu.hpp" + +#include "vision/norm.hpp" + +class Gemma4e_AudioEncoder{ + public: + ~Gemma4e_AudioEncoder(); + void init_weights(SafeTensors &q4nx); + Gemma4e_AudioEncoder(LM_Config config, npu_xclbin_manager *npu_instance, gemma4e_npu* parent_npu_ptr); + + void ffn_layer( + + buffer &ffn_up_proj_input, + buffer &ffn_up_proj_output_down_input, + buffer &ffn_down_proj_output, + + buffer &cur_ffn_norm_weight, // ffn_norm or ffn_norm_1 + buffer &cur_ffn_post_norm_weight,// ffn_post_norm_weight or ffn_post_norm_1_weight + buffer &cur_ffn_up_weight, // audio_ffn_up_weight or audio_ffn_up_1_weight + buffer &cur_ffn_down_weight, // audio_ffn_down_weight + int seq_len, int seq_len_padded, + + bf16 cur_audio_ffn_up_input_min, bf16 cur_audio_ffn_up_input_max, + bf16 cur_audio_ffn_up_output_min, bf16 cur_audio_ffn_up_output_max, + bf16 cur_audio_ffn_down_input_min, bf16 cur_audio_ffn_down_input_max, + bf16 cur_audio_ffn_down_output_min, bf16 cur_audio_ffn_down_output_max + + ); + void conv1d_layer( + int layer_idx, + int seq_len, int seq_len_padded, + std::vector &seq_len_per_audio, std::vector &start_seq_len_index_per_audio, + buffer &conv1d_start_proj_input, buffer &conv1d_start_proj_output, + buffer &conv1d_input, buffer &conv1d_output, + buffer &conv1d_end_proj_input, buffer &conv1d_end_proj_output, + + bf16 conv1d_start_input_min, bf16 conv1d_start_input_max, + bf16 conv1d_start_output_min, bf16 conv1d_start_output_max, + bf16 conv1d_end_input_min, bf16 conv1d_end_input_max, + bf16 conv1d_end_output_min, bf16 conv1d_end_output_max, + + gemma4e_audio_payload_t* audio_payload, + SafeTensors *reference_safetensor + ); + + std::vector encode( void* audio_payload_ptr); + + LM_Config config; + npu_xclbin_manager *npu; + gemma4e_npu* parent_npu_ptr; + + uint32_t MM_tile_M; + uint32_t MM_tile_K; + uint32_t MM_tile_N; + + float Gemma4E_Audio_residual_weight; + float Gemma4E_Audio_q_scale; + float Gemma4E_Audio_k_scale; + int Gemma4E_Audio_attention_head_dim; + int Gemma4E_Audio_padded_requirement_for_conv1d; + + uint32_t seq_len_pad_requirement_for_MM; + uint32_t MM_ROW_SIZE = 4; + uint32_t MM_COL_SIZE = 8; + + unsigned int Padded_GEMMA4E_Audio_HIDDEN_SIZE; + unsigned int Padded_GEMMA4E_Audio_MLP_INTERMEDIATE_SIZE; + unsigned int Padded_Gemma4E_Audio_Multimodal_Output_SIZE; + unsigned int Padded_GEMMA4E_Audio_Conv1d_Linear_OUTPUT_SIZE; + unsigned int Padded_Gemma4E_Audio_num_attention_heads; + + inline int round_up_to_multiple (int x, int multiple) + { + return ((x + multiple - 1) / multiple) * multiple; + }; + + // no reorder for MM + bool ENABLE_QKV_REORDER = false; + + // for MM runtime sequence + uint32_t rtp_address = 4096; // offset right after stack size + uint32_t rtp_sync_lock_id = 10; // the rtp sync lock + bool ENABLE_AXI4 = true; + bool IS_B_ROW_MAJOR = false; + + std::string model_path; + npu_app_manager *proj; + npu_app_manager *proj_high_precision; + npu_app_manager *conv1d; + + npu_app conv1d_app; + + npu_app sub_sampleConvProjection_app; + npu_app q_proj_app; + npu_app k_proj_app; + npu_app k_relative_proj_app; + npu_app v_proj_app; + npu_app o_proj_app; + + npu_app ffn_down_proj_app; // share for both first and second FFN in the layer + npu_app ffn_up_proj_app; + npu_app conv1d_start_proj_app; + npu_app conv1d_end_proj_app; + npu_app audio_pre_encode_proj_app; + npu_app audio_to_language_proj_app; + + std::vector residual; + std::vector hidden_state; + + std::vector q_blocks; + std::vector k_blocks; + std::vector v_blocks; + + // the weights + buffer audio_embedding_projection_weight; + + buffer audio_subsample_conv2d_weight_0; + buffer audio_subsample_conv2d_norm_weight_0; + buffer audio_subsample_conv2d_weight_1; + buffer audio_subsample_conv2d_norm_weight_1; + + buffer audio_pre_encode_weight; + buffer audio_to_language_projection_weight; + // weights for each layer + std::vector> audio_attn_k_weight; + std::vector> audio_attn_k_rel_weight; + std::vector> audio_attn_v_weight; + std::vector> audio_attn_q_weight; + std::vector> audio_attn_o_weight; + std::vector> audio_conv1d_weight; + + std::vector> audio_conv_pw_1_weight; + std::vector> audio_conv_pw_2_weight; + + std::vector> audio_ffn_down_weight; + std::vector> audio_ffn_up_weight; + std::vector> audio_ffn_down_1_weight; + std::vector> audio_ffn_up_1_weight; + + std::vector> attn_post_norm_weight; + std::vector> attn_pre_norm_weight; + std::vector> conv_norm_weight; + std::vector> norm_conv_weight; + std::vector> ffn_norm_weight; + std::vector> ffn_norm_1_weight; + std::vector> ffn_post_norm_weight; + std::vector> ffn_post_norm_1_weight; + std::vector> norm2_weight; + std::vector> per_dim_scale_with_softplus_weight; + + std::vector audio_k_input_min; + std::vector audio_k_input_max; + std::vector audio_v_input_min; + std::vector audio_v_input_max; + std::vector audio_q_input_min; + std::vector audio_q_input_max; + std::vector audio_o_input_min; + std::vector audio_o_input_max; + std::vector audio_conv_pw1_input_min; + std::vector audio_conv_pw1_input_max; + std::vector audio_conv_pw2_input_min; + std::vector audio_conv_pw2_input_max; + std::vector audio_ffn_down_input_min; + std::vector audio_ffn_down_input_max; + std::vector audio_ffn_up_input_min; + std::vector audio_ffn_up_input_max; + std::vector audio_ffn_down_1_input_min; + std::vector audio_ffn_down_1_input_max; + std::vector audio_ffn_up_1_input_min; + std::vector audio_ffn_up_1_input_max; + + std::vector audio_k_output_min; + std::vector audio_k_output_max; + std::vector audio_v_output_min; + std::vector audio_v_output_max; + std::vector audio_q_output_min; + std::vector audio_q_output_max; + std::vector audio_o_output_min; + std::vector audio_o_output_max; + std::vector audio_conv1d_output_min; + std::vector audio_conv1d_output_max; + std::vector audio_conv_pw1_output_min; + std::vector audio_conv_pw1_output_max; + std::vector audio_conv_pw2_output_min; + std::vector audio_conv_pw2_output_max; + std::vector audio_ffn_down_output_min; + std::vector audio_ffn_down_output_max; + std::vector audio_ffn_up_output_min; + std::vector audio_ffn_up_output_max; + std::vector audio_ffn_down_1_output_min; + std::vector audio_ffn_down_1_output_max; + std::vector audio_ffn_up_1_output_min; + std::vector audio_ffn_up_1_output_max; +}; diff --git a/src/detail/gemma4e_npu/gemma4e_audio_attention.cpp b/src/detail/gemma4e_npu/gemma4e_audio_attention.cpp new file mode 100644 index 000000000..21b912d35 --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_audio_attention.cpp @@ -0,0 +1,440 @@ +#include "gemma4e_audio_attention.hpp" +#include "gemma4e_vision_prefill_helper.hpp" +#include "avx512_util.hpp" +#include +#include +#include +#include +#include +#include + +void create_sliding_window_attention_mask( + int seq_len, + int sliding_window_left, + int sliding_window_right, + std::vector &mask + +){ + mask.assign(seq_len * seq_len, 0); + for (int q_idx = 0; q_idx < seq_len; ++q_idx) { + for (int kv_idx = 0; kv_idx < seq_len; ++kv_idx) { + int dist = q_idx - kv_idx; + bool left_mask = (dist >= 0) && (dist < sliding_window_left); + bool right_mask = (dist < 0) && (-dist < sliding_window_right); + + if (left_mask || right_mask) { + mask[q_idx * seq_len + kv_idx] = 1; + } else { + mask[q_idx * seq_len + kv_idx] = 0; + } + } + } +} +void convert_mask_to_blocked( + std::vector &input_mask, // [seq_len, seq_len] + int chunk_size, + int attention_context_left, + int attention_context_right, + std::vector &blocked_mask, // [num_chunks, chunk_size, context_size], + int &num_blocks, + int &context_size +) { + int seq_len = std::round(std::sqrt(input_mask.size())); + int max_past_horizon = attention_context_left - 1; + int max_future_horizon = attention_context_right; + + num_blocks = (seq_len + chunk_size - 1) / chunk_size; + context_size = chunk_size + max_past_horizon + max_future_horizon; + + blocked_mask.assign(num_blocks * chunk_size * context_size, 0); + + for (int b = 0; b < num_blocks; ++b) { + for (int c = 0; c < chunk_size; ++c) { + int q_idx = b * chunk_size + c; + for (int ctx = 0; ctx < context_size; ++ctx) { + int kv_idx = b * chunk_size + ctx - max_past_horizon; + int out_idx = b * chunk_size * context_size + c * context_size + ctx; + + if (q_idx >= 0 && q_idx < seq_len && kv_idx >= 0 && kv_idx < seq_len) { + blocked_mask[out_idx] = input_mask[q_idx * seq_len + kv_idx]; + } else { + blocked_mask[out_idx] = 0; + } + } + } + } +} + +void convert_to_block( + const bf16* input, + bf16* output, + int seq_len, + int chunk_size, + int row_stride, + int num_blocks +) { + // Output is expected to be zero-initialized by caller. + // The reshape from [num_blocks*chunk_size, row_stride] to [num_blocks, chunk_size, row_stride] + // is a no-op in row-major memory. We just copy valid rows and leave padding as zeros. + int total_output_rows = num_blocks * chunk_size; + int rows_to_copy = std::min(seq_len, total_output_rows); + memcpy(output, input, (size_t)rows_to_copy * row_stride * sizeof(bf16)); +} + +void extract_block_context( + const bf16* input, + bf16* output, + int seq_len, + int chunk_size, + int max_past_horizon, + int max_future_horizon, + int row_stride, + int num_blocks, + int context_size +) { + // 1. Create padded buffer: [max_past_horizon + seq_len + max_future_horizon + chunk_size - 1, row_stride] + // Left padding: max_past_horizon zeros, Right padding: max_future_horizon + chunk_size - 1 zeros + int padded_len = max_past_horizon + seq_len + max_future_horizon + chunk_size - 1; + std::vector padded((size_t)padded_len * row_stride, bf16(0)); + + // Copy input into padded buffer at offset max_past_horizon + memcpy( + padded.data() + (size_t)max_past_horizon * row_stride, + input, + (size_t)seq_len * row_stride * sizeof(bf16) + ); + + // 2. Extract overlapping windows (unfold): for block b, copy context_size rows starting at b*chunk_size + for (int b = 0; b < num_blocks; b++) { + int src_start = b * chunk_size; + memcpy( + output + (size_t)(b * context_size) * row_stride, + padded.data() + (size_t)src_start * row_stride, + (size_t)context_size * row_stride * sizeof(bf16) + ); + } +} + +// ============================================================================ +// AVX512 dot product of bf16 vectors (returns float32) +// ============================================================================ +static inline float avx512_dot_bf16(const bf16* a, const bf16* b, int len) { + __m512 acc0 = _mm512_setzero_ps(); + __m512 acc1 = _mm512_setzero_ps(); + int k = 0; + for (; k + 32 <= len; k += 32) { + acc0 = _mm512_fmadd_ps(load_bfloat16_to_m512(a + k), + load_bfloat16_to_m512(b + k), acc0); + acc1 = _mm512_fmadd_ps(load_bfloat16_to_m512(a + k + 16), + load_bfloat16_to_m512(b + k + 16), acc1); + } + __m512 acc = _mm512_add_ps(acc0, acc1); + for (; k + 16 <= len; k += 16) { + acc = _mm512_fmadd_ps(load_bfloat16_to_m512(a + k), + load_bfloat16_to_m512(b + k), acc); + } + float sum = _mm512_reduce_add_ps(acc); + for (; k < len; k++) { + sum += (float)a[k] * (float)b[k]; + } + return sum; +} + +// ============================================================================ +// _rel_shift for one block: [chunk_size, position_length] -> [chunk_size, context_size] +// Python equivalent: +// x = F.pad(x, (0, context_size + 1 - position_length)) +// x = x.view(..., block_size * (context_size + 1)) +// x = x[..., : block_size * context_size] +// x = x.view(..., block_size, context_size) +// ============================================================================ +static void rel_shift_block( + const float* input, // [chunk_size, position_length] + float* output, // [chunk_size, context_size] + int chunk_size, + int position_length, + int context_size +) { + int padded_col = context_size + 1; + int padded_size = chunk_size * padded_col; + // Use stack buffer for small sizes (typical: 12*25=300 floats = 1.2KB), heap fallback otherwise + float stack_buf[512]; + float* padded = (padded_size <= 512) ? stack_buf : new float[padded_size]; + memset(padded, 0, padded_size * sizeof(float)); + for (int c = 0; c < chunk_size; c++) { + memcpy(&padded[c * padded_col], &input[c * position_length], + position_length * sizeof(float)); + } + memcpy(output, padded, chunk_size * context_size * sizeof(float)); + if (padded_size > 512) delete[] padded; +} + +void compute_audio_blocked_attention( + const bf16* q_blocked, + const bf16* k_blocked, + const bf16* v_blocked, + const bf16* relative_key_states, + const int* block_attention_mask, + bf16* output, + int seq_len, + int num_blocks, + int chunk_size, + int context_size, + int num_positions, + int num_heads, + int head_dim, + int hidden_size, + int padded_hidden_size, + float softcap, + float invalid_logits_value +) { + memset(output, 0, (size_t)seq_len * padded_hidden_size * sizeof(bf16)); + + int total_q_rows = num_blocks * chunk_size; + float inv_softcap = 1.0f / softcap; + + // Pre-allocate per-thread buffers to avoid repeated heap allocation inside OMP loop + std::vector> thread_matrix_bd(num_heads, std::vector(total_q_rows * num_positions)); + std::vector> thread_shifted_bd(num_heads, std::vector(total_q_rows * context_size)); + std::vector> thread_attn_weights(num_heads, std::vector(chunk_size * context_size)); + + // Parallelize across heads — each head writes to non-overlapping columns + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) + for (int h = 0; h < num_heads; h++) { + int h_offset = h * head_dim; + int hd_vecs = head_dim / 16; + + float* matrix_bd = thread_matrix_bd[h].data(); + float* shifted_bd = thread_shifted_bd[h].data(); + float* attn_weights = thread_attn_weights[h].data(); + + // ---- Step 1: matrix_bd = queries_flat_h @ rel_K_h^T ---- + // queries_flat_h: [total_q_rows, head_dim] (head slice of q_blocked) + // rel_K_h: [num_positions, head_dim] (head slice of relative_key_states) + // matrix_bd: [total_q_rows, num_positions] + for (int i = 0; i < total_q_rows; i++) { + const bf16* q_row = q_blocked + (size_t)i * padded_hidden_size + h_offset; + for (int p = 0; p < num_positions; p++) { + const bf16* rk_row = relative_key_states + (size_t)p * hidden_size + h_offset; + matrix_bd[i * num_positions + p] = avx512_dot_bf16(q_row, rk_row, head_dim); + } + } + + // ---- Step 2: Reshape to [num_blocks, chunk_size, num_positions] then _rel_shift per block ---- + // _rel_shift: [chunk_size, num_positions] -> [chunk_size, context_size] + for (int blk = 0; blk < num_blocks; blk++) { + rel_shift_block( + &matrix_bd[blk * chunk_size * num_positions], + &shifted_bd[blk * chunk_size * context_size], + chunk_size, + num_positions, + context_size + ); + } + + // ---- Step 3: Per-block attention (matrix_ac + bd, softcap, mask, softmax, attn@V) ---- + for (int blk = 0; blk < num_blocks; blk++) { + // 3a. Compute combined logits: matrix_ac + shifted_bd, apply softcap + mask + for (int c = 0; c < chunk_size; c++) { + const bf16* q_row = q_blocked + + (size_t)(blk * chunk_size + c) * padded_hidden_size + h_offset; + + for (int ctx = 0; ctx < context_size; ctx++) { + const bf16* k_row = k_blocked + + (size_t)(blk * context_size + ctx) * padded_hidden_size + h_offset; + + // matrix_ac = Q_row · K_row + float w = avx512_dot_bf16(q_row, k_row, head_dim); + + // + shifted_bd + w += shifted_bd[(blk * chunk_size + c) * context_size + ctx]; + + // Softcap: tanh(w / softcap) * softcap + w = std::tanh(w * inv_softcap) * softcap; + + // Mask + int mask_idx = blk * chunk_size * context_size + c * context_size + ctx; + if (!block_attention_mask[mask_idx]) { + w = invalid_logits_value; + } + + attn_weights[c * context_size + ctx] = w; + } + } + + // 3b. Softmax over context_size dimension (per chunk row) + for (int c = 0; c < chunk_size; c++) { + float* row = &attn_weights[c * context_size]; + + // Find max for numerical stability + __m512 max_vec = _mm512_set1_ps(-1e30f); + int ctx = 0; + for (; ctx + 16 <= context_size; ctx += 16) { + max_vec = _mm512_max_ps(max_vec, _mm512_loadu_ps(row + ctx)); + } + float max_val = _mm512_reduce_max_ps(max_vec); + for (; ctx < context_size; ctx++) { + max_val = std::max(max_val, row[ctx]); + } + + // Exp and sum + __m512 sum_vec = _mm512_setzero_ps(); + __m512 max_broadcast = _mm512_set1_ps(max_val); + ctx = 0; + for (; ctx + 16 <= context_size; ctx += 16) { + __m512 val = _mm512_loadu_ps(row + ctx); + __m512 exp_val = _mm512_exp_ps_corrected(_mm512_sub_ps(val, max_broadcast)); + _mm512_storeu_ps(row + ctx, exp_val); + sum_vec = _mm512_add_ps(sum_vec, exp_val); + } + float sum_exp = _mm512_reduce_add_ps(sum_vec); + for (; ctx < context_size; ctx++) { + row[ctx] = std::exp(row[ctx] - max_val); + sum_exp += row[ctx]; + } + + // Normalize (guard against division by zero for fully-masked rows) + float inv_sum = (sum_exp > 0.0f) ? (1.0f / sum_exp) : 0.0f; + __m512 inv_sum_vec = _mm512_set1_ps(inv_sum); + ctx = 0; + for (; ctx + 16 <= context_size; ctx += 16) { + _mm512_storeu_ps(row + ctx, + _mm512_mul_ps(_mm512_loadu_ps(row + ctx), inv_sum_vec)); + } + for (; ctx < context_size; ctx++) { + row[ctx] *= inv_sum; + } + } + + // 3c. attn_output = attn_weights @ V_block (write to output) + for (int c = 0; c < chunk_size; c++) { + int out_seq_idx = blk * chunk_size + c; + if (out_seq_idx >= seq_len) break; + + bf16* out_ptr = output + (size_t)out_seq_idx * padded_hidden_size + h_offset; + + // Accumulate over context_size using AVX512 across head_dim + __m512 acc[16]; // supports up to head_dim=256 + for (int v = 0; v < hd_vecs; v++) acc[v] = _mm512_setzero_ps(); + + for (int ctx = 0; ctx < context_size; ctx++) { + __m512 w_broadcast = _mm512_set1_ps(attn_weights[c * context_size + ctx]); + const bf16* v_row = v_blocked + + (size_t)(blk * context_size + ctx) * padded_hidden_size + h_offset; + + for (int v = 0; v < hd_vecs; v++) { + __m512 v_vec = load_bfloat16_to_m512(v_row + v * 16); + acc[v] = _mm512_fmadd_ps(w_broadcast, v_vec, acc[v]); + } + } + + // Store bf16 result + for (int v = 0; v < hd_vecs; v++) { + store_m512_to_bfloat16(out_ptr + v * 16, acc[v]); + } + + // Scalar tail (head_dim not multiple of 16) + int tail_start = hd_vecs * 16; + for (int d = tail_start; d < head_dim; d++) { + float sum = 0.0f; + for (int ctx = 0; ctx < context_size; ctx++) { + float v_val = (float)v_blocked[ + (size_t)(blk * context_size + ctx) * padded_hidden_size + h_offset + d]; + sum += attn_weights[c * context_size + ctx] * v_val; + } + out_ptr[d] = (bf16)sum; + } + } + } + } +} + +void compute_audio_self_attention( + const bf16* q_proj_output, + const bf16* k_proj_output, + const bf16* v_proj_output, + const bf16* k_rel_weight, + const bf16* position_embedding, + const std::vector>& block_attention_mask_per_audio, + const std::vector& num_blocks_per_audio, + const std::vector& context_size_per_audio, + const std::vector& seq_len_per_audio, + const std::vector& start_seq_len_index_per_audio, + bf16* output, + int num_audios, + int chunk_size, + int context_left, + int context_right, + int num_heads, + int head_dim, + int hidden_size, + int padded_hidden_size, + int num_positions, + float softcap, + float invalid_logits_value +) { + int max_past_horizon = context_left - 1; + int max_future_horizon = context_right; + + // ---- Block Q/K/V per audio ---- + std::vector> blocked_q_per_audio(num_audios); + std::vector> blocked_k_per_audio(num_audios); + std::vector> blocked_v_per_audio(num_audios); + + for (int i = 0; i < num_audios; i++) { + int nb = num_blocks_per_audio[i]; + int cs = context_size_per_audio[i]; + + blocked_q_per_audio[i].assign((size_t)nb * chunk_size * padded_hidden_size, bf16(0)); + convert_to_block( + q_proj_output + start_seq_len_index_per_audio[i] * padded_hidden_size, + blocked_q_per_audio[i].data(), + seq_len_per_audio[i], chunk_size, padded_hidden_size, nb + ); + + blocked_k_per_audio[i].assign((size_t)nb * cs * padded_hidden_size, bf16(0)); + extract_block_context( + k_proj_output + start_seq_len_index_per_audio[i] * padded_hidden_size, + blocked_k_per_audio[i].data(), + seq_len_per_audio[i], chunk_size, max_past_horizon, max_future_horizon, + padded_hidden_size, nb, cs + ); + + blocked_v_per_audio[i].assign((size_t)nb * cs * padded_hidden_size, bf16(0)); + extract_block_context( + v_proj_output + start_seq_len_index_per_audio[i] * padded_hidden_size, + blocked_v_per_audio[i].data(), + seq_len_per_audio[i], chunk_size, max_past_horizon, max_future_horizon, + padded_hidden_size, nb, cs + ); + } + + // ---- Compute relative_key_states = position_embedding @ k_rel_weight^T ---- + std::vector relative_key_states(num_positions * hidden_size, bf16(0)); + simd_gemm_abt_bf16( + position_embedding, + k_rel_weight, + relative_key_states.data(), + num_positions, hidden_size, hidden_size, + hidden_size, padded_hidden_size, hidden_size + ); + + // ---- Compute blocked attention per audio ---- + for (int i = 0; i < num_audios; i++) { + int nb = num_blocks_per_audio[i]; + int cs = context_size_per_audio[i]; + + compute_audio_blocked_attention( + blocked_q_per_audio[i].data(), + blocked_k_per_audio[i].data(), + blocked_v_per_audio[i].data(), + relative_key_states.data(), + block_attention_mask_per_audio[i].data(), + output + start_seq_len_index_per_audio[i] * padded_hidden_size, + seq_len_per_audio[i], + nb, chunk_size, cs, num_positions, + num_heads, head_dim, hidden_size, padded_hidden_size, + softcap, invalid_logits_value + ); + } +} diff --git a/src/detail/gemma4e_npu/gemma4e_audio_attention.hpp b/src/detail/gemma4e_npu/gemma4e_audio_attention.hpp new file mode 100644 index 000000000..e8263fd24 --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_audio_attention.hpp @@ -0,0 +1,111 @@ +#pragma once + +#include +#include "typedef.hpp" +#include "buffer.hpp" + +void create_sliding_window_attention_mask( + int seq_len, + int sliding_window_left, + int sliding_window_right, + std::vector &mask + +); + +void convert_mask_to_blocked( + std::vector &input_mask, // [seq_len, seq_len] + int chunk_size, + int attention_context_left, + int attention_context_right, + std::vector &blocked_mask, // [num_chunks, chunk_size, context_size], + int &num_chunks, + int &context_size +); + +// Splits [seq_len, row_stride] into [num_blocks * chunk_size, row_stride] with zero-padding +// Equivalent to Python: F.pad(hidden_states, (0,0,0,0,0,pad)).reshape(batch, num_blocks, chunk_size, num_heads, head_dim) +void convert_to_block( + const bf16* input, // [seq_len, row_stride] + bf16* output, // [num_blocks * chunk_size, row_stride], caller must pre-allocate & zero-init + int seq_len, + int chunk_size, + int row_stride, // Padded hidden size (num_heads_padded * head_dim) + int num_blocks // = ceil(seq_len / chunk_size) +); + +// Extracts overlapping context windows for blocked attention +// Equivalent to Python: F.pad(...).unfold(1, context_size, chunk_size).movedim(-1, 2) +// Output shape per audio: [num_blocks, context_size, num_heads_padded, head_dim] +void extract_block_context( + const bf16* input, // [seq_len, row_stride] + bf16* output, // [num_blocks * context_size, row_stride], caller must pre-allocate & zero-init + int seq_len, + int chunk_size, + int max_past_horizon, // attention_context_left - 1 + int max_future_horizon, // attention_context_right + int row_stride, // Padded hidden size + int num_blocks, + int context_size // = chunk_size + max_past_horizon + max_future_horizon +); + +// Full blocked audio attention computation for a single audio sample. +// Implements the Python forward() from after Q/K/V scaling through attn_output (before post projection). +// +// q_blocked: [num_blocks * chunk_size, padded_hidden_size] (from convert_to_block) +// k_blocked: [num_blocks * context_size, padded_hidden_size] (from extract_block_context) +// v_blocked: [num_blocks * context_size, padded_hidden_size] (from extract_block_context) +// relative_key_states: [num_positions, hidden_size] dense row-major (= position_embedding @ k_rel_weight^T) +// block_attention_mask: [num_blocks * chunk_size * context_size] (1=attend, 0=mask) +// output: [seq_len, padded_hidden_size] row-major (caller pre-allocates) +void compute_audio_blocked_attention( + const bf16* q_blocked, + const bf16* k_blocked, + const bf16* v_blocked, + const bf16* relative_key_states, + const int* block_attention_mask, + bf16* output, + int seq_len, + int num_blocks, + int chunk_size, + int context_size, + int num_positions, + int num_heads, + int head_dim, + int hidden_size, + int padded_hidden_size, + float softcap, + float invalid_logits_value +); + +// Top-level audio self-attention: blocking + relative key computation + blocked attention + output assembly. +// Replaces the inline code in encode() that does convert_to_block, extract_block_context, +// simd_gemm_abt_bf16 for relative_key_states, and compute_audio_blocked_attention per audio. +// +// q_proj_output / k_proj_output / v_proj_output: [seq_len_padded, padded_hidden_size] packed for all audios +// k_rel_weight: [padded_hidden_size, padded_hidden_size] row-major +// position_embedding: [num_positions, hidden_size] dense +// output: [seq_len_padded, padded_hidden_size] pre-zeroed by caller +void compute_audio_self_attention( + const bf16* q_proj_output, + const bf16* k_proj_output, + const bf16* v_proj_output, + const bf16* k_rel_weight, + const bf16* position_embedding, + const std::vector>& block_attention_mask_per_audio, + const std::vector& num_blocks_per_audio, + const std::vector& context_size_per_audio, + const std::vector& seq_len_per_audio, + const std::vector& start_seq_len_index_per_audio, + bf16* output, + int num_audios, + int chunk_size, + int context_left, + int context_right, + int num_heads, + int head_dim, + int hidden_size, + int padded_hidden_size, + int num_positions, + float softcap, + float invalid_logits_value +); diff --git a/src/detail/gemma4e_npu/gemma4e_cpu_functions.hpp b/src/detail/gemma4e_npu/gemma4e_cpu_functions.hpp new file mode 100644 index 000000000..add5bd1c8 --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_cpu_functions.hpp @@ -0,0 +1,374 @@ +#ifndef __GEMMA4E_CPU_FUNCTIONS_HPP__ +#define __GEMMA4E_CPU_FUNCTIONS_HPP__ +#include +#include +#include +#include "typedef.hpp" +#include "buffer.hpp" +#include "models/gemma4e/flm/aie2p/gemma4e_npu.hpp" +#include "avx512_util.hpp" + +/// @brief Host-side batched kernels used by the gemma4e prefill path. +/// @note These used to be private members of gemma4e_npu::Impl; they are free +/// functions here so the prefill blocks can call them without a handle +/// on the model. Every one of them walks whole rows of a padded batch, +/// hence the L_offset_* arguments: row 0 of the buffer is the first +/// chunk-padding row, and the first live token sits at L_offset. +namespace gemma4e_cpu_func { + +/// @brief Threads the batched kernels fan out over. +static constexpr int MAX_PREFILL_THREAD = 4; + +/// @brief Rope frequency tables, one per attention flavour. +/// @note Both are zero-padded up to the head dim the kernels stride over. +inline constexpr f32 inv_freq_global[] = { + 1.0000e+00, 9.4746e-01, 8.9769e-01, 8.5053e-01, 8.0584e-01, 7.6351e-01, 7.2339e-01, 6.8539e-01, + 6.4938e-01, 6.1527e-01, 5.8294e-01, 5.5232e-01, 5.2330e-01, 4.9581e-01, 4.6976e-01, 4.4508e-01, + 4.2170e-01, 3.9954e-01, 3.7855e-01, 3.5866e-01, 3.3982e-01, 3.2197e-01, 3.0505e-01, 2.8903e-01, + 2.7384e-01, 2.5946e-01, 2.4582e-01, 2.3291e-01, 2.2067e-01, 2.0908e-01, 1.9810e-01, 1.8769e-01, + 1.7783e-01, 1.6849e-01, 1.5963e-01, 1.5125e-01, 1.4330e-01, 1.3577e-01, 1.2864e-01, 1.2188e-01, + 1.1548e-01, 1.0941e-01, 1.0366e-01, 9.8217e-02, 9.3057e-02, 8.8168e-02, 8.3536e-02, 7.9148e-02, + 7.4989e-02, 7.1050e-02, 6.7317e-02, 6.3780e-02, 6.0430e-02, 5.7255e-02, 5.4247e-02, 5.1397e-02, + 4.8697e-02, 4.6138e-02, 4.3714e-02, 4.1418e-02, 3.9242e-02, 3.7180e-02, 3.5227e-02, 3.3376e-02, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, + 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00 +}; +inline constexpr f32 inv_freq_swa[] = { + 1.0000e+00, 9.3057e-01, 8.6596e-01, 8.0584e-01, 7.4989e-01, 6.9783e-01, 6.4938e-01, 6.0430e-01, + 5.6234e-01, 5.2330e-01, 4.8697e-01, 4.5316e-01, 4.2170e-01, 3.9242e-01, 3.6517e-01, 3.3982e-01, + 3.1623e-01, 2.9427e-01, 2.7384e-01, 2.5483e-01, 2.3714e-01, 2.2067e-01, 2.0535e-01, 1.9110e-01, + 1.7783e-01, 1.6548e-01, 1.5399e-01, 1.4330e-01, 1.3335e-01, 1.2409e-01, 1.1548e-01, 1.0746e-01, + 1.0000e-01, 9.3057e-02, 8.6596e-02, 8.0584e-02, 7.4989e-02, 6.9783e-02, 6.4938e-02, 6.0430e-02, + 5.6234e-02, 5.2330e-02, 4.8697e-02, 4.5316e-02, 4.2170e-02, 3.9242e-02, 3.6517e-02, 3.3982e-02, + 3.1623e-02, 2.9427e-02, 2.7384e-02, 2.5483e-02, 2.3714e-02, 2.2067e-02, 2.0535e-02, 1.9110e-02, + 1.7783e-02, 1.6548e-02, 1.5399e-02, 1.4330e-02, 1.3335e-02, 1.2409e-02, 1.1548e-02, 1.0746e-02, + 1.0000e-02, 9.3057e-03, 8.6596e-03, 8.0584e-03, 7.4989e-03, 6.9783e-03, 6.4938e-03, 6.0430e-03, + 5.6234e-03, 5.2330e-03, 4.8697e-03, 4.5316e-03, 4.2170e-03, 3.9242e-03, 3.6517e-03, 3.3982e-03, + 3.1623e-03, 2.9427e-03, 2.7384e-03, 2.5483e-03, 2.3714e-03, 2.2067e-03, 2.0535e-03, 1.9110e-03, + 1.7783e-03, 1.6548e-03, 1.5399e-03, 1.4330e-03, 1.3335e-03, 1.2409e-03, 1.1548e-03, 1.0746e-03, + 1.0000e-03, 9.3057e-04, 8.6596e-04, 8.0584e-04, 7.4989e-04, 6.9783e-04, 6.4938e-04, 6.0430e-04, + 5.6234e-04, 5.2330e-04, 4.8697e-04, 4.5316e-04, 4.2170e-04, 3.9242e-04, 3.6517e-04, 3.3982e-04, + 3.1623e-04, 2.9427e-04, 2.7384e-04, 2.5483e-04, 2.3714e-04, 2.2067e-04, 2.0535e-04, 1.9110e-04, + 1.7783e-04, 1.6548e-04, 1.5399e-04, 1.4330e-04, 1.3335e-04, 1.2409e-04, 1.1548e-04, 1.0746e-04 +}; + +/// @brief y[l] = rms_norm(x[l]) * w, over L rows of a D-wide batch. +/// @param w_ptr may be nullptr, which normalizes without a learned weight. +inline void _rms_norm_batch(bf16* y, bf16* x, const bf16* w_ptr, int norm_D, int D, int L, int L_offset_dest, int L_offset_input){ + + bf16* x_base = x + (size_t)L_offset_input * D; + bf16* y_base = y + (size_t)L_offset_dest * D; + const int simd_width = 16; + + #pragma omp parallel for num_threads(MAX_PREFILL_THREAD) schedule(static) + for (int l = 0; l < L; l++) { + bf16* x_row = x_base + (size_t)l * D; + bf16* y_row = y_base + (size_t)l * D; + for (int d = 0; d < D / norm_D; d++){ + // Step 1: sum of squares over norm_D elements + bf16* x_row_inner = x_row + d * norm_D; + bf16* y_row_inner = y_row + d * norm_D; + __m512 sum_xx_vec = _mm512_setzero_ps(); + int i = 0; + for (; i + simd_width <= norm_D; i += simd_width) { + __m256i bf16_vals = _mm256_loadu_si256((const __m256i*)(x_row_inner + i)); + __m512 fp32_vals = bf16o_fp32_512(bf16_vals); + sum_xx_vec = _mm512_fmadd_ps(fp32_vals, fp32_vals, sum_xx_vec); + } + f32 sum_xx = _mm512_reduce_add_ps(sum_xx_vec); + for (; i < norm_D; i++) { f32 v = static_cast(x_row_inner[i]); sum_xx += v * v; } + + f32 inv_rms_x = 1.0f / sqrtf(sum_xx / (f32)norm_D + 1e-6f); + + // Step 2: y[d] = w[d] * x[d] * inv_rms_x + __m512 inv_rms_vec = _mm512_set1_ps(inv_rms_x); + i = 0; + for (; i + simd_width <= norm_D; i += simd_width) { + __m256i bf16_x = _mm256_loadu_si256((const __m256i*)(x_row_inner + i)); + if (w_ptr != nullptr){ + __m256i bf16_w = _mm256_loadu_si256((const __m256i*)(w_ptr + i)); + __m512 fp32_x = bf16o_fp32_512(bf16_x); + __m512 fp32_w = bf16o_fp32_512(bf16_w); + __m512 y_vec = _mm512_mul_ps(_mm512_mul_ps(fp32_w, fp32_x), inv_rms_vec); + _mm256_storeu_si256((__m256i*)(y_row_inner + i), f32o_bf16_512(y_vec)); + } + else{ + __m512 fp32_x = bf16o_fp32_512(bf16_x); + __m512 y_vec = _mm512_mul_ps(fp32_x, inv_rms_vec); + _mm256_storeu_si256((__m256i*)(y_row_inner + i), f32o_bf16_512(y_vec)); + } + } + for (; i < norm_D; i++) { + y_row_inner[i] = static_cast(static_cast(w_ptr[i]) * static_cast(x_row_inner[i]) * inv_rms_x); + } + } + } +} + +/// @brief x[l] *= scale, over L rows of a D-wide batch. +inline void _elementwise_scale_batch(bf16* x, float scale, int D, int L, int L_offset){ + bf16* dest_base = x + (size_t)L_offset * D; + int total_elements = L * D; + int i = 0; + + __m512 scale_vec = _mm512_set1_ps(scale); + for (; i + 15 < total_elements; i += 16) { + __m256i bf16_vals = _mm256_loadu_si256((const __m256i*)(dest_base + i)); + __m512 fp32_vals = bf16o_fp32_512(bf16_vals); + __m512 result = _mm512_mul_ps(fp32_vals, scale_vec); + _mm256_storeu_si256((__m256i*)(dest_base + i), f32o_bf16_512(result)); + } + for (; i < total_elements; i++){ + dest_base[i] = static_cast(static_cast(dest_base[i]) * scale); + } +} + +/// @brief Applies rope to q or k and the per-head rms norm, over L_effective rows. +/// @note The frequency table is chosen by layer type: swa layers use a shorter +/// wavelength set than global ones. _DH is the head dimension of this +/// layer type, which the caller reads off the model descriptor. +inline void _rope_rms_batch(bf16* x, int D, int L_offset, int L_begin, int L_effective, bf16* rms_weight, gemma4e_layer_type_t layer_type, int _DH){ + const float* inv_freq = is_swa_layer(layer_type) ? inv_freq_swa : inv_freq_global; + const int heads = D / _DH; + bf16* x_base = x + (size_t)L_offset * D; + #pragma omp parallel for num_threads(MAX_PREFILL_THREAD) schedule(static) + for (int ll = 0; ll < L_effective; ll++){ + // Compute thread-local sin/cos values + std::vector local_cos(_DH / 2); + std::vector local_sin(_DH / 2); + for (int j = 0; j < _DH / 2; j++){ + float angle = inv_freq[j] * (L_begin + ll); + local_cos[j] = cosf(angle); + local_sin[j] = sinf(angle); + } + + bf16* x_token = x_base + ll * D; + for (int h = 0; h < heads; h++){ + bf16* x_head = x_token + h * _DH; + // inline rms_norm on single head + { + const int simd_width = 16; + __m512 sum_xx_vec = _mm512_setzero_ps(); + int i = 0; + for (; i + simd_width <= (int)_DH; i += simd_width) { + __m256i bf16_vals = _mm256_loadu_si256((const __m256i*)(x_head + i)); + __m512 fp32_vals = bf16o_fp32_512(bf16_vals); + sum_xx_vec = _mm512_fmadd_ps(fp32_vals, fp32_vals, sum_xx_vec); + } + f32 sum_xx = _mm512_reduce_add_ps(sum_xx_vec); + for (; i < (int)_DH; i++) { f32 v = static_cast(x_head[i]); sum_xx += v * v; } + f32 inv_rms = 1.0f / sqrtf(sum_xx / (f32)_DH + 1e-6f); + __m512 inv_rms_vec = _mm512_set1_ps(inv_rms); + i = 0; + for (; i + simd_width <= (int)_DH; i += simd_width) { + __m256i bf16_x = _mm256_loadu_si256((const __m256i*)(x_head + i)); + __m256i bf16_w = _mm256_loadu_si256((const __m256i*)(rms_weight + i)); + __m512 y_vec = _mm512_mul_ps(_mm512_mul_ps(bf16o_fp32_512(bf16_x), bf16o_fp32_512(bf16_w)), inv_rms_vec); + _mm256_storeu_si256((__m256i*)(x_head + i), f32o_bf16_512(y_vec)); + } + for (; i < (int)_DH; i++) { + x_head[i] = static_cast(static_cast(rms_weight[i]) * static_cast(x_head[i]) * inv_rms); + } + } + // apply rope rotation + bf16* x_left = x_head; + bf16* x_right = x_head + _DH / 2; + int i = 0, simd = 16; + for (; i + simd <= _DH / 2; i += simd) { + __m256i left_bf16 = _mm256_loadu_si256((__m256i*)(x_left + i)); + __m256i right_bf16 = _mm256_loadu_si256((__m256i*)(x_right + i)); + __m512 Lv = bf16o_fp32_512(left_bf16); + __m512 Rv = bf16o_fp32_512(right_bf16); + __m512 C = _mm512_loadu_ps(local_cos.data() + i); + __m512 S = _mm512_loadu_ps(local_sin.data() + i); + __m512 newL = _mm512_sub_ps(_mm512_mul_ps(Lv, C), _mm512_mul_ps(Rv, S)); + __m512 newR = _mm512_add_ps(_mm512_mul_ps(Lv, S), _mm512_mul_ps(Rv, C)); + _mm256_storeu_si256((__m256i*)(x_left + i), f32o_bf16_512(newL)); + _mm256_storeu_si256((__m256i*)(x_right + i), f32o_bf16_512(newR)); + } + } + } +} + +/// @brief dest = x + residual, over L rows of a D-wide batch. +inline void _residual_add_batch(bf16* dest, bf16* x, bf16* residual, int D, int L, int L_offset_dest, int L_offset_x, int L_offset_r){ + bf16* dest_base = dest + (size_t)L_offset_dest * D; + const bf16* src_base = x + (size_t)L_offset_x * D; + const bf16* res_base = residual + (size_t)L_offset_r * D; + + const int simd_width = 16; + + #pragma omp parallel for num_threads(MAX_PREFILL_THREAD) schedule(static) + for (int l = 0; l < L; l++) { + bf16* dest_row = dest_base + (size_t)l * D; + const bf16* src_row = src_base + (size_t)l * D; + const bf16* res_row = res_base + (size_t)l * D; + + int i = 0; + for (; i + simd_width <= D; i += simd_width) { + __m256i bf16_src = _mm256_loadu_si256((const __m256i*)(src_row + i)); + __m256i bf16_res = _mm256_loadu_si256((const __m256i*)(res_row + i)); + __m512 result = _mm512_add_ps(bf16o_fp32_512(bf16_src), bf16o_fp32_512(bf16_res)); + _mm256_storeu_si256((__m256i*)(dest_row + i), f32o_bf16_512(result)); + } + for (; i < D; i++) { + dest_row[i] = static_cast(static_cast(src_row[i]) + static_cast(res_row[i])); + } + } +} + +/// @brief dest = a * b, over L rows of a D-wide batch. +inline void _elementwise_mul_batch(bf16* dest, bf16* a,bf16* b, int D, int L, int L_offset_dest, int L_offset_a, int L_offset_b){ + + bf16* dest_base = dest + (size_t)L_offset_dest * D; + const bf16* a_base = a + (size_t)L_offset_a * D; + const bf16* b_base = b + (size_t)L_offset_b * D; + + const int simd_width = 16; + const int unroll = simd_width * 4; + + #pragma omp parallel for num_threads(MAX_PREFILL_THREAD) schedule(static) + for (int l = 0; l < L; l++) { + bf16* dest_row = dest_base + (size_t)l * D; + const bf16* a_row = a_base + (size_t)l * D; + const bf16* b_row = b_base + (size_t)l * D; + + int i = 0; + // 4-way unrolled loop for maximum throughput + for (; i + unroll <= D; i += unroll){ + __m256i bf16_vals_a0 = _mm256_loadu_si256((const __m256i*)(a_row + i)); + __m256i bf16_vals_a1 = _mm256_loadu_si256((const __m256i*)(a_row + i + simd_width)); + __m256i bf16_vals_a2 = _mm256_loadu_si256((const __m256i*)(a_row + i + simd_width * 2)); + __m256i bf16_vals_a3 = _mm256_loadu_si256((const __m256i*)(a_row + i + simd_width * 3)); + + __m256i bf16_vals_b0 = _mm256_loadu_si256((const __m256i*)(b_row + i)); + __m256i bf16_vals_b1 = _mm256_loadu_si256((const __m256i*)(b_row + i + simd_width)); + __m256i bf16_vals_b2 = _mm256_loadu_si256((const __m256i*)(b_row + i + simd_width * 2)); + __m256i bf16_vals_b3 = _mm256_loadu_si256((const __m256i*)(b_row + i + simd_width * 3)); + + __m512 fp32_vals_a0 = bf16o_fp32_512(bf16_vals_a0); + __m512 fp32_vals_a1 = bf16o_fp32_512(bf16_vals_a1); + __m512 fp32_vals_a2 = bf16o_fp32_512(bf16_vals_a2); + __m512 fp32_vals_a3 = bf16o_fp32_512(bf16_vals_a3); + + __m512 fp32_vals_b0 = bf16o_fp32_512(bf16_vals_b0); + __m512 fp32_vals_b1 = bf16o_fp32_512(bf16_vals_b1); + __m512 fp32_vals_b2 = bf16o_fp32_512(bf16_vals_b2); + __m512 fp32_vals_b3 = bf16o_fp32_512(bf16_vals_b3); + + __m512 result_vec0 = _mm512_mul_ps(fp32_vals_a0, fp32_vals_b0); + __m512 result_vec1 = _mm512_mul_ps(fp32_vals_a1, fp32_vals_b1); + __m512 result_vec2 = _mm512_mul_ps(fp32_vals_a2, fp32_vals_b2); + __m512 result_vec3 = _mm512_mul_ps(fp32_vals_a3, fp32_vals_b3); + + _mm256_storeu_si256((__m256i*)(dest_row + i), f32o_bf16_512(result_vec0)); + _mm256_storeu_si256((__m256i*)(dest_row + i + simd_width), f32o_bf16_512(result_vec1)); + _mm256_storeu_si256((__m256i*)(dest_row + i + simd_width * 2), f32o_bf16_512(result_vec2)); + _mm256_storeu_si256((__m256i*)(dest_row + i + simd_width * 3), f32o_bf16_512(result_vec3)); + } + + // Handle remaining SIMD-width chunks + for (; i + simd_width <= D; i += simd_width){ + __m256i bf16_vals_a = _mm256_loadu_si256((const __m256i*)(a_row + i)); + __m256i bf16_vals_b = _mm256_loadu_si256((const __m256i*)(b_row + i)); + __m512 result_vec = _mm512_mul_ps(bf16o_fp32_512(bf16_vals_a), bf16o_fp32_512(bf16_vals_b)); + _mm256_storeu_si256((__m256i*)(dest_row + i), f32o_bf16_512(result_vec)); + } + + // Scalar tail + for (; i < D; i++){ + dest_row[i] = static_cast(static_cast(a_row[i]) * static_cast(b_row[i])); + } + } +} + +/// @brief dest = gate * b, where b is a PLI_D-wide slice of a D-wide row. +inline void _elementwise_mul_batch(bf16* dest, bf16* gate,bf16* b, int PLI_D, int D, int L, int L_offset_dest, int L_offset_a, int L_offset_b){ + + bf16* dest_base = dest + (size_t)L_offset_dest * PLI_D; + const bf16* a_base = gate + (size_t)L_offset_a * PLI_D; + const bf16* b_base = b + (size_t)L_offset_b * D; + + const int simd_width = 16; + const int unroll = simd_width * 4; + + #pragma omp parallel for num_threads(MAX_PREFILL_THREAD) schedule(static) + for (int l = 0; l < L; l++) { + bf16* dest_row = dest_base + (size_t)l * PLI_D; + const bf16* a_row = a_base + (size_t)l * PLI_D; + const bf16* b_row = b_base + (size_t)l * D; + + int i = 0; + // 4-way unrolled loop for maximum throughput + for (; i + unroll <= PLI_D; i += unroll){ + __m256i bf16_vals_a0 = _mm256_loadu_si256((const __m256i*)(a_row + i)); + __m256i bf16_vals_a1 = _mm256_loadu_si256((const __m256i*)(a_row + i + simd_width)); + __m256i bf16_vals_a2 = _mm256_loadu_si256((const __m256i*)(a_row + i + simd_width * 2)); + __m256i bf16_vals_a3 = _mm256_loadu_si256((const __m256i*)(a_row + i + simd_width * 3)); + + __m256i bf16_vals_b0 = _mm256_loadu_si256((const __m256i*)(b_row + i)); + __m256i bf16_vals_b1 = _mm256_loadu_si256((const __m256i*)(b_row + i + simd_width)); + __m256i bf16_vals_b2 = _mm256_loadu_si256((const __m256i*)(b_row + i + simd_width * 2)); + __m256i bf16_vals_b3 = _mm256_loadu_si256((const __m256i*)(b_row + i + simd_width * 3)); + + __m512 fp32_vals_a0 = bf16o_fp32_512(bf16_vals_a0); + __m512 fp32_vals_a1 = bf16o_fp32_512(bf16_vals_a1); + __m512 fp32_vals_a2 = bf16o_fp32_512(bf16_vals_a2); + __m512 fp32_vals_a3 = bf16o_fp32_512(bf16_vals_a3); + + __m512 fp32_vals_b0 = bf16o_fp32_512(bf16_vals_b0); + __m512 fp32_vals_b1 = bf16o_fp32_512(bf16_vals_b1); + __m512 fp32_vals_b2 = bf16o_fp32_512(bf16_vals_b2); + __m512 fp32_vals_b3 = bf16o_fp32_512(bf16_vals_b3); + + __m512 result_vec0 = _mm512_mul_ps(fp32_vals_a0, fp32_vals_b0); + __m512 result_vec1 = _mm512_mul_ps(fp32_vals_a1, fp32_vals_b1); + __m512 result_vec2 = _mm512_mul_ps(fp32_vals_a2, fp32_vals_b2); + __m512 result_vec3 = _mm512_mul_ps(fp32_vals_a3, fp32_vals_b3); + + _mm256_storeu_si256((__m256i*)(dest_row + i), f32o_bf16_512(result_vec0)); + _mm256_storeu_si256((__m256i*)(dest_row + i + simd_width), f32o_bf16_512(result_vec1)); + _mm256_storeu_si256((__m256i*)(dest_row + i + simd_width * 2), f32o_bf16_512(result_vec2)); + _mm256_storeu_si256((__m256i*)(dest_row + i + simd_width * 3), f32o_bf16_512(result_vec3)); + } + + // Handle remaining SIMD-width chunks + for (; i + simd_width <= PLI_D; i += simd_width){ + __m256i bf16_vals_a = _mm256_loadu_si256((const __m256i*)(a_row + i)); + __m256i bf16_vals_b = _mm256_loadu_si256((const __m256i*)(b_row + i)); + __m512 result_vec = _mm512_mul_ps(bf16o_fp32_512(bf16_vals_a), bf16o_fp32_512(bf16_vals_b)); + _mm256_storeu_si256((__m256i*)(dest_row + i), f32o_bf16_512(result_vec)); + } + + // Scalar tail + for (; i < PLI_D; i++){ + dest_row[i] = static_cast(static_cast(a_row[i]) * static_cast(b_row[i])); + } + } +} + +} // namespace gemma4e_cpu_func + +#endif // __GEMMA4E_CPU_FUNCTIONS_HPP__ diff --git a/src/detail/gemma4e_npu/gemma4e_image.cpp b/src/detail/gemma4e_npu/gemma4e_image.cpp new file mode 100644 index 000000000..31f3b3710 --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_image.cpp @@ -0,0 +1,1854 @@ +#include "flm_override.hpp" +#include "gemma4e_image.hpp" + +#include +#include +#include +#ifdef _WIN32 +#include +#endif +#include "utils/debug_utils.hpp" +#include "utils/error_measure.hpp" + +#include "gemma4e_vision_prefill_helper.hpp" + +#include "vision/norm.hpp" +#include "mmRuntimeSequence.hpp" +#include "rot_pos_emb.hpp" +#include "seq_gen.hpp" + +#define DEBUG_PRINT_ENCODE_TIME_DETAIL (DEBUG_LEVEL >= 1) +#define DEBUG_PRINT_ENCODE_ERROR_METRICS (DEBUG_LEVEL >= 1) + +Gemma4e_ImageEncoder::~Gemma4e_ImageEncoder() +{ +} + +Gemma4e_ImageEncoder::Gemma4e_ImageEncoder(LM_Config config, npu_xclbin_manager *npu_instance, gemma4e_npu* parent_npu_ptr) + : config(config), npu(npu_instance), model_path(config.model_path), parent_npu_ptr(parent_npu_ptr) +{ + + //TODO: FIXME:hi + + // load parameters from json file + + { + + MM_tile_M = config.sub("vision_config").value("VISION_MM_TILE_M", -1); + MM_tile_K = config.sub("vision_config").value("VISION_MM_TILE_K", -1); + MM_tile_N = config.sub("vision_config").value("VISION_MM_TILE_N", -1); + + seq_len_pad_requirement_for_MM = MM_ROW_SIZE*MM_tile_M; + assert( MM_tile_K % MM_tile_N == 0); // we need this for the way we generate the sequence, because k and N is interchangable at MLP (gate-down) + DEBUG_BLOCK(1, + std::cout << "MM_tile_M: " << MM_tile_M << ", MM_tile_K: " << MM_tile_K << ", MM_tile_N: " << MM_tile_N << std::endl; + ) + } + + Padded_GEMMA4E_VISION_HIDDEN_SIZE = round_up_to_multiple(this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, MM_tile_K); + Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE = round_up_to_multiple(this->parent_npu_ptr->GEMMA4E_VISION_INTERMEDIATE_SIZE, MM_tile_K); + Padded_GEMMA4E_VISION_OUT_HIDDEN_SIZE = round_up_to_multiple(this->parent_npu_ptr->GEMMA4E_VISION_IMAGE_OUTPUT_SIZE, MM_tile_K); + + // #if DEBUG_PRINT_ENCODE_ERROR_METRICS + // // print all the parameters for debug + // std::cout << "QWEN3_5_VISION_EMBED_DIM: " << QWEN3_5_VISION_EMBED_DIM << std::endl; + // std::cout << "QWEN3_5_VISION_NUM_HEADS: " << QWEN3_5_VISION_NUM_HEADS << std::endl; + // std::cout << "QWEN3_5_VISION_HEAD_DIM: " << QWEN3_5_VISION_HEAD_DIM << std::endl; + // std::cout << "QWEN3_5_VISION_HIDDEN_SIZE: " << _QWEN3_5_VISION_HIDDEN_SIZE << std::endl; + // std::cout << "QWEN3_5_VISION_MLP_INTERMEDIATE_SIZE: " << _QWEN3_5_VISION_MLP_INTERMEDIATE_SIZE << std::endl; + // std::cout << "QWEN3_5_VISION_NUM_POSITION_EMBEDDINGS: " << QWEN3_5_VISION_NUM_POSITION_EMBEDDINGS << std::endl; + // std::cout << "QWEN3_5_VISION_NUM_LAYERS: " << QWEN3_5_VISION_NUM_LAYERS << std::endl; + // std::cout << "QWEN3_5_VISION_LAYER_NORM_EPSILON: " << QWEN3_5_VISION_LAYER_NORM_EPSILON << std::endl; + // std::cout << "QWEN3_5_MERGER_HIDDEN_SIZE: " << _QWEN3_5_MERGER_HIDDEN_SIZE << std::endl; + // std::cout << "QWEN3_5_VISION_OUT_HIDDEN_SIZE: " << _QWEN3_5_VISION_OUT_HIDDEN_SIZE << std::endl; + + // // print the padded parameters for debug + + // #endif + + this->fla = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "vision_attn.xclbin")); + this->proj = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "vision_mm.xclbin")); + this->proj_high_precision = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "vision_mm_high_precision.xclbin")); + + this->flash_attention_app = this->fla->create_app(); + this->patch_embedder_posisiton_embedding_dim_0_app = this->proj->create_app();//this->proj->create_app(); //TODO: this might not needeD? + this->patch_embedder_posisiton_embedding_dim_1_app = this->proj->create_app();//this->proj->create_app(); + this->patch_embedding_app = this->proj_high_precision->create_app();//this->proj->create_app(); + this->q_proj_app = this->proj->create_app(); + this->k_proj_app = this->proj->create_app(); + this->v_proj_app = this->proj->create_app(); + this->o_proj_app = this->proj->create_app(); + this->gate_proj_app = this->proj->create_app(); + this->up_proj_app = this->proj->create_app(); + this->down_proj_app = this->proj->create_app(); + + this->vision_to_language_input_projection_app = this->proj->create_app(); + + // // attempting to read the info that vision model needs from config.sub("vision_config") + + q_proj_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + k_proj_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + v_proj_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + o_proj_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + gate_proj_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + up_proj_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + down_proj_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + q_norm_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + k_norm_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + post_o_norm_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + post_ffn_norm_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + layer_norm_1_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); + layer_norm_2_weight.resize(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS); +} +void Gemma4e_ImageEncoder::init_weights(SafeTensors &q4nx){ + DEBUG_BLOCK(1, + std::cout << "HIT: init_weights of Gemma4e_ImageEncoder is called. Loading weights from " << this->model_path << std::endl; + ); + + { + + buffer temp_buffer; + + q4nx.load_weights(temp_buffer,"model.vision.patch_embedder.position_embedding_table" ); + this->patch_embedder_position_embedding_table = this->patch_embedder_posisiton_embedding_dim_0_app.create_bo_buffer( + temp_buffer.size() + ); + memcpy(this->patch_embedder_position_embedding_table.data(), temp_buffer.data(), temp_buffer.size()*sizeof(bf16)); + } + { + buffer temp_buffer; + q4nx.load_weights(temp_buffer,"model.vision.patch_embd.weight" ); + this->patch_embd_weight = this->patch_embedding_app.create_bo_buffer( + temp_buffer.size() + ); + memcpy(this->patch_embd_weight.data(), temp_buffer.data(), temp_buffer.size()*sizeof(bf16)); + } + { + buffer temp_buffer; + q4nx.load_weights(temp_buffer,"model.vision.embedding_projection.weight" ); + this->vision_to_language_input_projection_weight = this->vision_to_language_input_projection_app.create_bo_buffer( + temp_buffer.size() + ); + memcpy(this->vision_to_language_input_projection_weight.data(), temp_buffer.data(), temp_buffer.size()*sizeof(bf16)); + } + + for(int layer_id=0; layer_id < this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS; layer_id++){ + + this->q_proj_weight[layer_id] = this->q_proj_app.create_bo_buffer( + Padded_GEMMA4E_VISION_HIDDEN_SIZE*Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + q4nx.load_weights(this->q_proj_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".vision_attn.q_proj.weight" + ); + + this->k_proj_weight[layer_id] = this->k_proj_app.create_bo_buffer( + Padded_GEMMA4E_VISION_HIDDEN_SIZE*Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + q4nx.load_weights(this->k_proj_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".vision_attn.k_proj.weight" + ); + this->v_proj_weight[layer_id] = this->v_proj_app.create_bo_buffer( + Padded_GEMMA4E_VISION_HIDDEN_SIZE*Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + q4nx.load_weights(this->v_proj_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".vision_attn.v_proj.weight" + ); + this->o_proj_weight[layer_id] = this->o_proj_app.create_bo_buffer( + Padded_GEMMA4E_VISION_HIDDEN_SIZE*Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + q4nx.load_weights(this->o_proj_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".vision_attn.out_proj.weight" + ); + this->gate_proj_weight[layer_id] = this->gate_proj_app.create_bo_buffer( + Padded_GEMMA4E_VISION_HIDDEN_SIZE*Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE + ); + q4nx.load_weights(this->gate_proj_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".ffn.gate_proj.weight" + ); + this->up_proj_weight[layer_id] = this->up_proj_app.create_bo_buffer( + Padded_GEMMA4E_VISION_HIDDEN_SIZE*Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE + ); + q4nx.load_weights(this->up_proj_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".ffn.up_proj.weight" + ); + this->down_proj_weight[layer_id] = this->down_proj_app.create_bo_buffer( + Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE*Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + q4nx.load_weights(this->down_proj_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".ffn.down_proj.weight" + ); + + // load all norm weights + q4nx.load_weights( + this->q_norm_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".vision_attn.q_norm.weight" + ); + q4nx.load_weights( + this->k_norm_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".vision_attn.k_norm.weight" + ); + q4nx.load_weights( + this->post_o_norm_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".post_attn_norm.weight" + ); + q4nx.load_weights( + this->post_ffn_norm_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".ffn_post_norm.weight" + ); + q4nx.load_weights( + this->layer_norm_1_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".norm1.weight" + ); + q4nx.load_weights( + this->layer_norm_2_weight[layer_id], + "model.vision."+std::to_string(layer_id)+".norm2.weight" + ); + + // now, we load the min max scalar + buffer q_input_min; + q4nx.load_weights(q_input_min, + "model.vision."+std::to_string(layer_id)+".vision_attn.q_input_min" + ); + assert(q_input_min.size() == 1); + this->input_q_min.push_back(q_input_min[0]); + + buffer q_input_max; + q4nx.load_weights(q_input_max, + "model.vision."+std::to_string(layer_id)+".vision_attn.q_input_max" + ); + assert(q_input_max.size() == 1); + this->input_q_max.push_back(q_input_max[0]); + + buffer q_output_min; + q4nx.load_weights(q_output_min, + "model.vision."+std::to_string(layer_id)+".vision_attn.q_output_min" + ); + assert(q_output_min.size() == 1); + this->output_q_min.push_back(q_output_min[0]); + + buffer q_output_max; + q4nx.load_weights(q_output_max, + "model.vision."+std::to_string(layer_id)+".vision_attn.q_output_max" + ); + assert(q_output_max.size() == 1); + this->output_q_max.push_back(q_output_max[0]); + + // k min/max + buffer k_input_min; + q4nx.load_weights(k_input_min, + "model.vision."+std::to_string(layer_id)+".vision_attn.k_input_min" + ); + assert(k_input_min.size() == 1); + this->input_k_min.push_back(k_input_min[0]); + + buffer k_input_max; + q4nx.load_weights(k_input_max, + "model.vision."+std::to_string(layer_id)+".vision_attn.k_input_max" + ); + assert(k_input_max.size() == 1); + this->input_k_max.push_back(k_input_max[0]); + + buffer k_output_min; + q4nx.load_weights(k_output_min, + "model.vision."+std::to_string(layer_id)+".vision_attn.k_output_min" + ); + assert(k_output_min.size() == 1); + this->output_k_min.push_back(k_output_min[0]); + + buffer k_output_max; + q4nx.load_weights(k_output_max, + "model.vision."+std::to_string(layer_id)+".vision_attn.k_output_max" + ); + assert(k_output_max.size() == 1); + this->output_k_max.push_back(k_output_max[0]); + + // v min/max + buffer v_input_min; + q4nx.load_weights(v_input_min, + "model.vision."+std::to_string(layer_id)+".vision_attn.v_input_min" + ); + assert(v_input_min.size() == 1); + this->input_v_min.push_back(v_input_min[0]); + + buffer v_input_max; + q4nx.load_weights(v_input_max, + "model.vision."+std::to_string(layer_id)+".vision_attn.v_input_max" + ); + assert(v_input_max.size() == 1); + this->input_v_max.push_back(v_input_max[0]); + + buffer v_output_min; + q4nx.load_weights(v_output_min, + "model.vision."+std::to_string(layer_id)+".vision_attn.v_output_min" + ); + assert(v_output_min.size() == 1); + this->output_v_min.push_back(v_output_min[0]); + + buffer v_output_max; + q4nx.load_weights(v_output_max, + "model.vision."+std::to_string(layer_id)+".vision_attn.v_output_max" + ); + assert(v_output_max.size() == 1); + this->output_v_max.push_back(v_output_max[0]); + + // o (attn_out) min/max + buffer o_input_min; + q4nx.load_weights(o_input_min, + "model.vision."+std::to_string(layer_id)+".attn_out_input_min" + ); + assert(o_input_min.size() == 1); + this->input_o_min.push_back(o_input_min[0]); + + buffer o_input_max; + q4nx.load_weights(o_input_max, + "model.vision."+std::to_string(layer_id)+".attn_out_input_max" + ); + assert(o_input_max.size() == 1); + this->input_o_max.push_back(o_input_max[0]); + + buffer o_output_min; + q4nx.load_weights(o_output_min, + "model.vision."+std::to_string(layer_id)+".attn_out_output_min" + ); + assert(o_output_min.size() == 1); + this->output_o_min.push_back(o_output_min[0]); + + buffer o_output_max; + q4nx.load_weights(o_output_max, + "model.vision."+std::to_string(layer_id)+".attn_out_output_max" + ); + assert(o_output_max.size() == 1); + this->output_o_max.push_back(o_output_max[0]); + + // gate min/max + buffer gate_input_min; + q4nx.load_weights(gate_input_min, + "model.vision."+std::to_string(layer_id)+".ffn.gate_input_min" + ); + assert(gate_input_min.size() == 1); + this->input_gate_min.push_back(gate_input_min[0]); + + buffer gate_input_max; + q4nx.load_weights(gate_input_max, + "model.vision."+std::to_string(layer_id)+".ffn.gate_input_max" + ); + assert(gate_input_max.size() == 1); + this->input_gate_max.push_back(gate_input_max[0]); + + buffer gate_output_min; + q4nx.load_weights(gate_output_min, + "model.vision."+std::to_string(layer_id)+".ffn.gate_output_min" + ); + assert(gate_output_min.size() == 1); + this->output_gate_min.push_back(gate_output_min[0]); + + buffer gate_output_max; + q4nx.load_weights(gate_output_max, + "model.vision."+std::to_string(layer_id)+".ffn.gate_output_max" + ); + assert(gate_output_max.size() == 1); + this->output_gate_max.push_back(gate_output_max[0]); + + // up min/max + buffer up_input_min; + q4nx.load_weights(up_input_min, + "model.vision."+std::to_string(layer_id)+".ffn.up_input_min" + ); + assert(up_input_min.size() == 1); + this->input_up_min.push_back(up_input_min[0]); + + buffer up_input_max; + q4nx.load_weights(up_input_max, + "model.vision."+std::to_string(layer_id)+".ffn.up_input_max" + ); + assert(up_input_max.size() == 1); + this->input_up_max.push_back(up_input_max[0]); + + buffer up_output_min; + q4nx.load_weights(up_output_min, + "model.vision."+std::to_string(layer_id)+".ffn.up_output_min" + ); + assert(up_output_min.size() == 1); + this->output_up_min.push_back(up_output_min[0]); + + buffer up_output_max; + q4nx.load_weights(up_output_max, + "model.vision."+std::to_string(layer_id)+".ffn.up_output_max" + ); + assert(up_output_max.size() == 1); + this->output_up_max.push_back(up_output_max[0]); + + // down min/max + buffer down_input_min; + q4nx.load_weights(down_input_min, + "model.vision."+std::to_string(layer_id)+".ffn.down_input_min" + ); + assert(down_input_min.size() == 1); + this->input_down_min.push_back(down_input_min[0]); + + buffer down_input_max; + q4nx.load_weights(down_input_max, + "model.vision."+std::to_string(layer_id)+".ffn.down_input_max" + ); + assert(down_input_max.size() == 1); + this->input_down_max.push_back(down_input_max[0]); + + buffer down_output_min; + q4nx.load_weights(down_output_min, + "model.vision."+std::to_string(layer_id)+".ffn.down_output_min" + ); + assert(down_output_min.size() == 1); + this->output_down_min.push_back(down_output_min[0]); + + buffer down_output_max; + q4nx.load_weights(down_output_max, + "model.vision."+std::to_string(layer_id)+".ffn.down_output_max" + ); + assert(down_output_max.size() == 1); + this->output_down_max.push_back(down_output_max[0]); + } + + DEBUG_BLOCK(1, + std::cout << "[DBG] Gemma4e_ImageEncoder::init_weights finished" << std::endl; + ); +} + +std::vector Gemma4e_ImageEncoder::encode( void* image_payload_ptr) +{ + + DEBUG_BLOCK(1, + std::cout << "HIT: Gemma4e_ImageEncoder::encode is called with image_payload_ptr: " << image_payload_ptr << std::endl; + ); + // // + //DEBUG + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + SafeTensors reference_tensor(this->model_path + "/vision_reference_data.safetensors"); + std::cout << "reference_Tensor path is " << this->model_path + "/vision_reference_data.safetensors" << std::endl; + std::cout << "debug Gemma4e_ImageEncoder::encode called with image_payload_ptr: " << image_payload_ptr << std::endl; + #endif + // auto encoder_start_time = std::chrono::high_resolution_clock::now(); + + gemma4e_image_payload_t* image_payload = (gemma4e_image_payload_t*)image_payload_ptr; + + // first, compare the pixel_values with pre_Gemma4VisionPatchEmbedder_pixel_values + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + std::cout << "Error after patch embedding comparison:" << std::endl; + buffer pre_Gemma4VisionPatchEmbedder_pixel_values; + reference_tensor.load_weights( + pre_Gemma4VisionPatchEmbedder_pixel_values, + "pre_Gemma4VisionPatchEmbedder_pixel_values" + ); + std::cout << "Comparing pixel values with reference..." << std::endl; + uint32_t offset = 0; + for(int i = 0; i < image_payload->num_images; i++){ + uint32_t cur_size = image_payload->image_patch__element_per_patch[i].first*image_payload->image_patch__element_per_patch[i].second; + print_error_metrics( + image_payload->pixel_values[i].data(), + pre_Gemma4VisionPatchEmbedder_pixel_values.data() + offset, + 1, + cur_size, 1, + cur_size, 1 + ); + offset += cur_size; + } + + #endif + // std::cout << "Finished comparing pixel values with reference." << std::endl; + + // TODO: lets use avx512 for it later + + // for each (image_payload->pixel_values) pixel_values = 2 * (pixel_values - 0.5) + + // sanity checkst + for(int i = 0; i < image_payload->num_images; i++){ + assert( image_payload->image_patch__element_per_patch[i].second == this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE); + } + + std::vector seq_len_per_image; // unpadded seq_len per image + std::vector start_seq_len_index_per_image; // the start index in the sequence for each image1 + + int seq_len = 0; + int seq_len_of_last_image = 0; + for(auto image_valid_patch: image_payload->valid_patch_size_per_image ){ + seq_len += image_valid_patch; + seq_len_of_last_image = image_valid_patch; + seq_len_per_image.push_back(image_valid_patch); + + start_seq_len_index_per_image.push_back(seq_len - image_valid_patch); // cumulative offset before this image + } + + int seq_len_padded = round_up_to_multiple(seq_len_of_last_image , vision_L_padded_requirement_for_attention) - seq_len_of_last_image; + seq_len_padded += seq_len; + seq_len_padded = round_up_to_multiple(seq_len_padded, seq_len_pad_requirement_for_MM); + + assert(this->parent_npu_ptr->GEMMA4E_VISION_IMAGE_OUTPUT_SIZE % MM_tile_K == 0); + assert(this->parent_npu_ptr->GEMMA4E_VISION_IMAGE_OUTPUT_SIZE % MM_tile_N == 0); + + DEBUG_BLOCK(1, + std::cout <<"seq_len_pad_requirement_for_MM is " << seq_len_pad_requirement_for_MM << std::endl; + ) + buffer patch_emb_input = patch_embedding_app.create_bo_buffer( + seq_len_padded *this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + buffer patch_emb_output = patch_embedding_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + buffer q_projection_input= q_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + buffer k_projection_input = k_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + buffer v_projection_input = v_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + buffer q_projection_output= q_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + buffer k_projection_output = k_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + buffer v_projection_output = v_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + buffer attention_output = this->flash_attention_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + buffer o_projection_output = o_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + buffer gate_input = gate_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + buffer up_input = up_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + buffer gate_output = gate_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE + ); + buffer up_output = up_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE + ); + buffer down_output = down_proj_app.create_bo_buffer( + seq_len_padded * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + //[2, seq_len_padded, GEMMA4E_POSITION_EMBEDDING_SIZE] + buffer one_shot_buffer = patch_embedder_posisiton_embedding_dim_0_app.create_bo_buffer( + seq_len_padded* 2* this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE + ); + //[2, seq_len_padded, PADDED_GEMMA4E_VISION_HIDDEN_SIZE] + buffer position_embedding_table_output = patch_embedder_posisiton_embedding_dim_0_app.create_bo_buffer( + seq_len_padded* 2* Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + // // memset the buffer to zero + + { + + uint32_t ADD_BIAS = false;// no bias at all for all the mm + generate_mm_sequence(*this->patch_embedding_app.seq(), + seq_len_padded, Padded_GEMMA4E_VISION_HIDDEN_SIZE,Padded_GEMMA4E_VISION_HIDDEN_SIZE, + MM_tile_M,MM_tile_K,MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + ADD_BIAS, 0, + 0, -10000.0, 1000000.0, // do not clamp on output + ENABLE_QKV_REORDER, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + + generate_mm_sequence(*this->patch_embedder_posisiton_embedding_dim_0_app.seq(), + seq_len_padded, this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE ,Padded_GEMMA4E_VISION_HIDDEN_SIZE, + MM_tile_M,MM_tile_K,MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + ADD_BIAS, 0, + 0, -10000.0, 1000000.0, // do not clamp on output + ENABLE_QKV_REORDER, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + + generate_mm_sequence(*this->patch_embedder_posisiton_embedding_dim_1_app.seq(), + seq_len_padded, this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE ,Padded_GEMMA4E_VISION_HIDDEN_SIZE, + MM_tile_M,MM_tile_K,MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + + seq_len_padded* this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE, + this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE* Padded_GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_padded* Padded_GEMMA4E_VISION_HIDDEN_SIZE, + + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + ADD_BIAS, 0, + 0, -10000.0, 1000000.0, // do not clamp on output + ENABLE_QKV_REORDER, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + + gen_mha_vision_attention( + + this->flash_attention_app.seq(), + seq_len_per_image, + vision_L_padded_requirement_for_attention, + vision_S_padded_requirement_for_attention, + vision_num_of_columns, + vision_num_of_rows, + vision_CU_mode, + vision_LQ_per_CT, + vision_LK_per_CT, + vision_LQ_internal, + vision_LK_internal, + ENABLE_QKV_REORDER, + this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + this->Padded_GEMMA4E_VISION_HIDDEN_SIZE, + this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM, + this->parent_npu_ptr->GEMMA4E_VISION_NUM_ATTENTION_HEADS + + ); + } + + // now, apply the pixel_values = 2 * (pixel_values - 0.5) to all valye in patch_emb_input + std::vector temp_pixel_values(patch_emb_input.size()); + memset(temp_pixel_values.data(), 0, temp_pixel_values.size() * sizeof(bf16)); + for(int i = 0, seq_len_offset = 0; i < image_payload->num_images; i++){ + + bf16* raw_pixel_values_ptr = image_payload->pixel_values[i].data(); + bf16* patch_emb_input_ptr = (bf16*)temp_pixel_values.data() + (seq_len_offset* Padded_GEMMA4E_VISION_HIDDEN_SIZE); + + for(int l = 0; l < seq_len_per_image[i]; l++){ + for(int d = 0; d < this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; d++){ + float pixel_value = (float)raw_pixel_values_ptr[l*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE + d]; + pixel_value = 2.0f * (pixel_value - 0.5f); + patch_emb_input_ptr[l*Padded_GEMMA4E_VISION_HIDDEN_SIZE + d] = (bf16)pixel_value; + } + } + seq_len_offset += seq_len_per_image[i]; + } + memcpy(patch_emb_input.data(), temp_pixel_values.data(), temp_pixel_values.size()*sizeof(bf16)); + + patch_emb_input.sync_to_device(); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + std::cout << "Error after scaled vision hidden_size :" << std::endl; + buffer scaled_Gemma4VisionPatchEmbedder_pixel_values; + reference_tensor.load_weights( + scaled_Gemma4VisionPatchEmbedder_pixel_values, + "scaled_Gemma4VisionPatchEmbedder_pixel_values" + ); + + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + patch_emb_input.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + scaled_Gemma4VisionPatchEmbedder_pixel_values.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + image_payload->image_patch__element_per_patch[i].first, Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*image_payload->image_patch__element_per_patch[i].second; + } + + // //TODO: FIXME: + + // for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + // // memcpy( + // // patch_emb_input.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + // // scaled_Gemma4VisionPatchEmbedder_pixel_values.data() + ref_offset, + + // // seq_len_per_image[i]* this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE * sizeof(bf16) + + // } + + #endif + DEBUG_BLOCK(1, + std::cout << "model.vision.patch_embd.weight size: " << patch_embd_weight.size() << std::endl; + ) + + #if DEBUG_PRINT_ENCODE_TIME_DETAIL + auto patch_emb_start_time = std::chrono::high_resolution_clock::now(); + #endif + patch_emb_input.sync_to_device(); + patch_embd_weight.sync_to_device(); + FLM_OVERRIDE(vision_patch_embed, patch_embedding_app(patch_emb_input,patch_embd_weight, patch_emb_output )); + patch_emb_output.sync_from_device(); + #if DEBUG_PRINT_ENCODE_TIME_DETAIL + auto patch_emb_end_time = std::chrono::high_resolution_clock::now(); + std::chrono::duration patch_emb_duration = patch_emb_end_time - patch_emb_start_time; + std::cout << "Time taken for patch embedding: " << patch_emb_duration.count() << " ms" << std::endl; + #endif + // #if DEBUG_PRINT_ENCODE_ERROR_METRICS + + // std::cout << "Error for Gemma4VisionPatchEmbedder_hidden_states:" << std::endl; + // buffer Gemma4VisionPatchEmbedder_hidden_states; + // reference_tensor.load_weights( + // Gemma4VisionPatchEmbedder_hidden_states, + // "Gemma4VisionPatchEmbedder_hidden_states" + // ); + + // for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + // print_error_metrics( + // patch_emb_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + // Gemma4VisionPatchEmbedder_hidden_states.data() + ref_offset, + // 1, + // seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + // image_payload->image_patch__element_per_patch[i].first, Padded_GEMMA4E_VISION_HIDDEN_SIZE + + // // //TODO: FIXME: + // // for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + // // memcpy( + // // patch_emb_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + // // Gemma4VisionPatchEmbedder_hidden_states.data() + ref_offset, + + // // seq_len_per_image[i] * this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE * sizeof(bf16) + + // #endif + + { + // now the _position_embeddings + + ///torch.Size([3, 2, 2520, 10240]) + // first, zero out all one_shot_buffer [2, seq_len_padded, GEMMA4E_POSITION_EMBEDDING_SIZE] + memset(one_shot_buffer.data(), 0, one_shot_buffer.size()*sizeof(bf16)); + + bf16* one_shot_buffer_x_base = (bf16*)one_shot_buffer.data(); + bf16* one_shot_buffer_y_base = one_shot_buffer_x_base + seq_len_padded* this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE; + + for(int i = 0; i < image_payload->num_images; i++){ + bf16* x_ptr = one_shot_buffer_x_base + start_seq_len_index_per_image[i] * this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE; + bf16* y_ptr = one_shot_buffer_y_base + start_seq_len_index_per_image[i] * this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE; + + for(int s = 0; s< seq_len_per_image[i]; s++){ + auto x_val = image_payload->image_grid_pairs_per_image[i][s*2]; + auto y_val = image_payload->image_grid_pairs_per_image[i][s*2 + 1]; + + x_ptr[s * this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE + x_val] = 1.0f; + y_ptr[s * this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE + y_val] = 1.0f; + } + } + } + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + + bf16* one_shot_buffer_x_ptr = (bf16*)one_shot_buffer.data(); + bf16* one_shot_buffer_y_ptr = one_shot_buffer_x_ptr + seq_len_padded* this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE; + + std::cout << "Error for Gemma4VisionPatchEmbedder_one_hot_positions:" << std::endl; + buffer Gemma4VisionPatchEmbedder_one_hot_positions; + reference_tensor.load_weights( + Gemma4VisionPatchEmbedder_one_hot_positions, // shape of [num_image, 2, seq_len_per_image, GEMMA4E_POSITION_EMBEDDING_SIZE] + "Gemma4VisionPatchEmbedder_one_hot_positions" + ); + + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + one_shot_buffer_x_ptr + start_seq_len_index_per_image[i]*this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE, + Gemma4VisionPatchEmbedder_one_hot_positions.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE, + image_payload->image_patch__element_per_patch[i].first, this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE + + ); + + print_error_metrics( + one_shot_buffer_y_ptr + start_seq_len_index_per_image[i]*this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE, + Gemma4VisionPatchEmbedder_one_hot_positions.data() + ref_offset + image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE , + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE, + image_payload->image_patch__element_per_patch[i].first, this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE + + ); + + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_POSITION_EMBEDDING_SIZE*2; // times 2 is for x and y; + } + } + + #endif + + // now, we do the + one_shot_buffer.sync_to_device(); + patch_embedder_position_embedding_table.sync_to_device(); + FLM_OVERRIDE(vision_pos_embed_dim0, patch_embedder_posisiton_embedding_dim_0_app(one_shot_buffer, patch_embedder_position_embedding_table, position_embedding_table_output)); + position_embedding_table_output.sync_from_device(); + + // dim2 + one_shot_buffer.sync_to_device(); + patch_embedder_position_embedding_table.sync_to_device(); + FLM_OVERRIDE(vision_pos_embed_dim1, patch_embedder_posisiton_embedding_dim_1_app(one_shot_buffer, patch_embedder_position_embedding_table, position_embedding_table_output)); + position_embedding_table_output.sync_from_device(); + + // now, compare with the python reference + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + + bf16* position_embedding_table_output_x_ptr = (bf16*)position_embedding_table_output.data(); + bf16* position_embedding_table_output_y_ptr = position_embedding_table_output_x_ptr + seq_len_padded* Padded_GEMMA4E_VISION_HIDDEN_SIZE; + + std::cout << "Error Gemma4VisionPatchEmbedder_position_embeddings_before_sum:" << std::endl; + buffer Gemma4VisionPatchEmbedder_position_embeddings_before_sum; + reference_tensor.load_weights( + Gemma4VisionPatchEmbedder_position_embeddings_before_sum, // shape of [num_image, 2, seq_len_per_image, GEMMA4E_VISION_HIDDEN_SIZE] (unpadded) + "Gemma4VisionPatchEmbedder_position_embeddings_before_sum" + ); + + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + const int patches = image_payload->image_patch__element_per_patch[i].first; + const int ref_hidden = this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // Python ref is unpadded [seq_len, 768] + + print_error_metrics( + position_embedding_table_output_x_ptr + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + Gemma4VisionPatchEmbedder_position_embeddings_before_sum.data() + ref_offset, + 1, + seq_len_per_image[i], ref_hidden, // cols to compare = ref row stride (unpadded) + patches, Padded_GEMMA4E_VISION_HIDDEN_SIZE // cpp row stride (padded) + ); + + print_error_metrics( + position_embedding_table_output_y_ptr + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + Gemma4VisionPatchEmbedder_position_embeddings_before_sum.data() + ref_offset + patches * ref_hidden, + 1, + seq_len_per_image[i], ref_hidden, // cols to compare = ref row stride (unpadded) + patches, Padded_GEMMA4E_VISION_HIDDEN_SIZE // cpp row stride (padded) + ); + + ref_offset += patches * ref_hidden * 2; // times 2 for x and y, use unpadded ref stride + } + } + #endif + + hidden_state.resize(seq_len_padded * Padded_GEMMA4E_VISION_HIDDEN_SIZE); + memset(hidden_state.data(), 0, hidden_state.size() * sizeof(bf16)); + + // now, we sum the two dim + //TODO: use avx512 for it later + + { + + bf16* position_embedding_table_output_x_ptr = (bf16*)position_embedding_table_output.data(); + bf16* position_embedding_table_output_y_ptr = position_embedding_table_output_x_ptr + seq_len_padded* Padded_GEMMA4E_VISION_HIDDEN_SIZE; + for(int i = 0; i < seq_len* Padded_GEMMA4E_VISION_HIDDEN_SIZE; i++){ + hidden_state[i] = position_embedding_table_output_x_ptr[i] + position_embedding_table_output_y_ptr[i] + patch_emb_output[i]; + } + } + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for vision_inputs_embeds_after_patch_embedder:" << std::endl; + buffer vision_inputs_embeds_after_patch_embedder; + reference_tensor.load_weights( + vision_inputs_embeds_after_patch_embedder, // shape of [seq_len_padded, GEMMA4E_VISION_HIDDEN_SIZE] + "vision_inputs_embeds_after_patch_embedder" + ); + + auto ref_hidden_state = this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // Python ref is unpadded [seq_len, 768] + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + vision_inputs_embeds_after_patch_embedder.data() + ref_offset, + 1, + seq_len_per_image[i],ref_hidden_state, + image_payload->image_patch__element_per_patch[i].first, Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*ref_hidden_state; // use unpadded ref stride + } + + // //TODO: fixme + // //TODO: FIXME: + // for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + // memcpy( + // hidden_state.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + // vision_inputs_embeds_after_patch_embedder.data() + ref_offset, + + // seq_len_per_image[i] * Padded_GEMMA4E_VISION_HIDDEN_SIZE * sizeof(bf16) + + // ); + // ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + // } + } + + #endif + + // generate the rope, + std::vector cos_emb; + std::vector sin_emb; + + generate_gemma4_vision_rotary_pos_emb( + image_payload->image_grid_pairs_per_image, + seq_len_per_image, + start_seq_len_index_per_image, + seq_len_padded, + (int)this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM, + this->parent_npu_ptr->GEMMA4E_ROPE_THETA, + 1.0f, // attention_scaling = 1.0 + cos_emb, + sin_emb + ); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + const int head_dim = (int)this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM; + + buffer ref_cos, ref_sin; + reference_tensor.load_weights(ref_cos, "Gemma4VisionEncoder_position_embeddings_cos"); + reference_tensor.load_weights(ref_sin, "Gemma4VisionEncoder_position_embeddings_sin"); + + // ref shape is [num_images, ref_seq_per_image, head_dim] — derive per-image stride from total size + const int ref_seq_per_image = (int)ref_cos.size() / ((int)image_payload->num_images * head_dim); + std::cout << "ref_seq_per_image" << ref_seq_per_image << std::endl; + + std::cout << "Error for Gemma4VisionEncoder_position_embeddings_cos:" << std::endl; + for (int i = 0; i < (int)image_payload->num_images; i++) { + print_error_metrics( + cos_emb.data() + start_seq_len_index_per_image[i] * head_dim, + ref_cos.data() + i * ref_seq_per_image * head_dim, + 1, + seq_len_per_image[i], head_dim, // rows/cols to compare; ref row stride = head_dim + seq_len_per_image[i], head_dim // cpp cos_emb row stride = head_dim (no padding) + ); + } + + std::cout << "Error for Gemma4VisionEncoder_position_embeddings_sin:" << std::endl; + for (int i = 0; i < (int)image_payload->num_images; i++) { + print_error_metrics( + sin_emb.data() + start_seq_len_index_per_image[i] * head_dim, + ref_sin.data() + i * ref_seq_per_image * head_dim, + 1, + seq_len_per_image[i], head_dim, + seq_len_per_image[i], head_dim + ); + } + } + #endif + + residual_buffer.resize(seq_len_padded * Padded_GEMMA4E_VISION_HIDDEN_SIZE); + memset(residual_buffer.data(), 0, residual_buffer.size() * sizeof(bf16)); + //TODO: FIXME: + for(int layer_idx = 0; layer_idx < this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS; layer_idx++){ + memcpy(residual_buffer.data(), hidden_state.data(), seq_len* Padded_GEMMA4E_VISION_HIDDEN_SIZE * sizeof(bf16)); + + // first, comparew with f"Gemma4VisionEncoderLayer_{layer_idx}_initial_hidden_states"] + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionEncoderLayer_" << layer_idx << "_initial_hidden_states:" << std::endl; + buffer ref_initial_hidden_states; + reference_tensor.load_weights( + ref_initial_hidden_states, + "Gemma4VisionEncoderLayer_" + std::to_string(layer_idx) + "_initial_hidden_states" + ); + + // for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + // print_error_metrics( + // hidden_state.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + // ref_initial_hidden_states.data() + ref_offset, + // 1, + // seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + // seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + // ); + // ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + // } + } + #endif + + // do the first RMS norm on hidden_states + + simd_rms_norm( + hidden_state.data(), + this->layer_norm_1_weight[layer_idx].data(), + hidden_state.data(), + seq_len, this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_padded, this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + // compare with "Gemma4VisionEncoderLayer_{layer_idx}_initial_hidden_states" in the reference for the result + // #if DEBUG_PRINT_ENCODE_ERROR_METRICS + // { + // std::cout << "Error for Gemma4VisionEncoderLayer_" << layer_idx << "_post_input_layernorm_hidden_states:" << std::endl; + // buffer ref_post_input_layernorm_hidden_states; + // reference_tensor.load_weights( + // ref_post_input_layernorm_hidden_states, + // "Gemma4VisionEncoderLayer_" + std::to_string(layer_idx) + "_post_input_layernorm_hidden_states" + // ); + // } + // #endif + // #ifdef DEBUG_PRINT_ENCODE_ERROR_METRICS + // { + // std::cout << "Error for Gemma4VisionEncoderLayer_" << layer_idx << "_post_input_layernorm_hidden_states:" << std::endl; + // buffer ref_post_input_layernorm_hidden_states; + // reference_tensor.load_weights( + // ref_post_input_layernorm_hidden_states, + // "Gemma4VisionEncoderLayer_" + std::to_string(layer_idx) + "_post_input_layernorm_hidden_states" + // ); + + // for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + // print_error_metrics( + // hidden_state.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + // ref_post_input_layernorm_hidden_states.data() + ref_offset, + // 1, + // seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + // seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + // ); + // ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + // } + // } + // #endif + + generate_mm_sequence(*this->q_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_VISION_HIDDEN_SIZE ,Padded_GEMMA4E_VISION_HIDDEN_SIZE, + MM_tile_M,MM_tile_K,MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, + 1, (float)this->output_q_min[layer_idx],(float) this->output_q_max[layer_idx], // clamp output to quantization range for q_proj + ENABLE_QKV_REORDER, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + + simd_clamp( + hidden_state.data(), + q_projection_input.data(), + this->input_q_min[layer_idx], this->input_q_max[layer_idx], + seq_len * Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + memset(q_projection_input.data() + seq_len * Padded_GEMMA4E_VISION_HIDDEN_SIZE, 0, (seq_len_padded - seq_len)* Padded_GEMMA4E_VISION_HIDDEN_SIZE * sizeof(bf16)); + q_projection_input.sync_to_device(); + + auto q_proj_run = FLM_OVERRIDE(vision_q_proj, this->q_proj_app.create_run( + q_projection_input, q_proj_weight[layer_idx], q_projection_output + ), layer_idx); + + q_projection_input.sync_to_device(); + this->q_proj_weight[layer_idx].sync_to_device(); + q_proj_run.start(); + + generate_mm_sequence(*this->k_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_VISION_HIDDEN_SIZE ,Padded_GEMMA4E_VISION_HIDDEN_SIZE, + MM_tile_M,MM_tile_K,MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, + 1, (float)this->output_k_min[layer_idx],(float) this->output_k_max[layer_idx], // clamp output to quantization range for k_proj + ENABLE_QKV_REORDER, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + simd_clamp( + hidden_state.data(), + k_projection_input.data(), + this->input_k_min[layer_idx], this->input_k_max[layer_idx], + seq_len * Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + memset(k_projection_input.data() + seq_len * Padded_GEMMA4E_VISION_HIDDEN_SIZE, 0, (seq_len_padded - seq_len)* Padded_GEMMA4E_VISION_HIDDEN_SIZE * sizeof(bf16)); + k_projection_input.sync_to_device(); + + q_proj_run.wait(); + q_projection_output.sync_from_device(); + + k_projection_input.sync_to_device(); + this->k_proj_weight[layer_idx].sync_to_device(); + auto k_proj_run = FLM_OVERRIDE(vision_k_proj, this->k_proj_app.create_run( + k_projection_input, k_proj_weight[layer_idx], k_projection_output + ), layer_idx); + k_proj_run.start(); + + generate_mm_sequence(*this->v_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_VISION_HIDDEN_SIZE ,Padded_GEMMA4E_VISION_HIDDEN_SIZE, + MM_tile_M,MM_tile_K,MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, + 1, (float)this->output_v_min[layer_idx],(float) this->output_v_max[layer_idx], // clamp output to quantization range for v_proj + ENABLE_QKV_REORDER, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + simd_clamp( + hidden_state.data(), + v_projection_input.data(), + this->input_v_min[layer_idx], this->input_v_max[layer_idx], + seq_len * Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + memset(v_projection_input.data() + seq_len * Padded_GEMMA4E_VISION_HIDDEN_SIZE, 0, (seq_len_padded - seq_len)* Padded_GEMMA4E_VISION_HIDDEN_SIZE * sizeof(bf16)); + v_projection_input.sync_to_device(); + + k_proj_run.wait(); + k_projection_output.sync_from_device(); + + v_projection_input.sync_to_device(); + this->v_proj_weight[layer_idx].sync_to_device(); + auto v_proj_run = FLM_OVERRIDE(vision_v_proj, this->v_proj_app.create_run( + v_projection_input, v_proj_weight[layer_idx], v_projection_output + ), layer_idx); + v_proj_run.start(); + // apply norm for q, and k + simd_rms_norm( + q_projection_output.data(), + this->q_norm_weight[layer_idx].data(), + q_projection_output.data(), + seq_len * (this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE / this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM),this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM, + seq_len * (this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE / this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM),this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + apply_multidimensional_rope( + q_projection_output.data(), cos_emb.data(), sin_emb.data(), + seq_len, this->parent_npu_ptr->GEMMA4E_VISION_NUM_ATTENTION_HEADS, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM); + q_projection_output.sync_to_device(); + simd_rms_norm( + k_projection_output.data(), + this->k_norm_weight[layer_idx].data(), + k_projection_output.data(), + seq_len * (this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE / this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM),this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM, + seq_len * (this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE / this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM),this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + + apply_multidimensional_rope( + k_projection_output.data(), cos_emb.data(), sin_emb.data(), + seq_len, this->parent_npu_ptr->GEMMA4E_VISION_NUM_ATTENTION_HEADS, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM); + k_projection_output.sync_to_device(); + + v_proj_run.wait(); + v_projection_output.sync_from_device(); + + // apply norm for v + simd_rms_norm( + v_projection_output.data(), + v_projection_output.data(), + seq_len * (this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE / this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM),this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM, + seq_len * (this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE / this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM),this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + v_projection_output.sync_to_device(); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionAttention_" + std::to_string(layer_idx)+ "_query_states_after_rope:" << std::endl; + buffer ref_q_projection_output; + reference_tensor.load_weights( + ref_q_projection_output, + "Gemma4VisionAttention_" + std::to_string(layer_idx) + "_query_states_after_rope" + ); + + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + q_projection_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + ref_q_projection_output.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + } + std::cout << "Error for Gemma4VisionAttention_" + std::to_string(layer_idx)+ "_key_states_after_rope:" << std::endl; + buffer ref_k_projection_output; + reference_tensor.load_weights( + ref_k_projection_output, + "Gemma4VisionAttention_" + std::to_string(layer_idx) + "_key_states_after_rope" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + k_projection_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + ref_k_projection_output.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + } + + std::cout << "Error for Gemma4VisionAttention_" + std::to_string(layer_idx)+ "_value_states_after_norm:" << std::endl; + buffer ref_v_projection_output; + reference_tensor.load_weights( + ref_v_projection_output, + "Gemma4VisionAttention_" + std::to_string(layer_idx) + "_value_states_after_norm" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + v_projection_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + ref_v_projection_output.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + } + } + #endif + + auto attention_run = FLM_OVERRIDE(vision_attn_core, this->flash_attention_app.create_run( + attention_output, q_projection_output, k_projection_output, v_projection_output + )); + q_projection_output.sync_to_device(); + k_projection_output.sync_to_device(); + v_projection_output.sync_to_device(); + attention_run.start(); + + generate_mm_sequence(*this->o_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_VISION_HIDDEN_SIZE ,Padded_GEMMA4E_VISION_HIDDEN_SIZE, + MM_tile_M,MM_tile_K,MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, + 1, (float)this->output_o_min[layer_idx],(float) this->output_o_max[layer_idx], // clamp output to quantization range for v_proj + ENABLE_QKV_REORDER, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + + attention_run.wait(); + attention_output.sync_from_device(); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionAttention_" + std::to_string(layer_idx)+ "_attention_output:" << std::endl; + buffer ref_attention_output; + reference_tensor.load_weights( + ref_attention_output, + "Gemma4VisionAttention_" + std::to_string(layer_idx) +"_attn_output_before_o_proj" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + attention_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + ref_attention_output.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + } + + // //TOOD: FIXME: remove later + // for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + // memcpy( + // attention_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + // ref_attention_output.data() + ref_offset, + // seq_len_per_image[i]* this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE * sizeof(bf16) + + // ); + // ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + // } + } + + #endif + + simd_clamp( + attention_output.data(), + attention_output.data(), + this->input_o_min[layer_idx], this->input_o_max[layer_idx], + seq_len * Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + attention_output.sync_to_device(); + this->o_proj_weight[layer_idx].sync_to_device(); + FLM_OVERRIDE(vision_o_proj, o_proj_app(attention_output, this->o_proj_weight[layer_idx], o_projection_output ), layer_idx); + o_projection_output.sync_from_device(); + + simd_rms_norm( + o_projection_output.data(), + this->post_o_norm_weight[layer_idx].data(), + o_projection_output.data(), + seq_len, this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_padded, this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionEncoderLayer_" + std::to_string(layer_idx)+ "_post_attention_layernorm_hidden_states:" << std::endl; + buffer ref__post_attention_layernorm_hidden_states; + reference_tensor.load_weights( + ref__post_attention_layernorm_hidden_states, + "Gemma4VisionEncoderLayer_" + std::to_string(layer_idx) + "_post_attention_layernorm_hidden_states" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + o_projection_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + ref__post_attention_layernorm_hidden_states.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + } + } + #endif + + o_projection_output.sync_to_device(); + + simd_add( + o_projection_output.data(), + residual_buffer.data(), + hidden_state.data(), + seq_len * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + // copy to residual buffer + memcpy(residual_buffer.data(), hidden_state.data(), seq_len* Padded_GEMMA4E_VISION_HIDDEN_SIZE * sizeof(bf16)); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionEncoderLayer_" + std::to_string(layer_idx)+ "_hidden_states_after_attention_residual:" << std::endl; + buffer ref__hidden_states_after_attention_residual; + reference_tensor.load_weights( + ref__hidden_states_after_attention_residual, + "Gemma4VisionEncoderLayer_" + std::to_string(layer_idx) + "_hidden_states_after_attention_residual" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + ref__hidden_states_after_attention_residual.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + } + } + #endif + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for pre mlp norm weight" << std::endl; + //Gemma4VisionEncoderLayer_{layer_idx}_pre_feedforward_layernorm_weights + buffer ref_pre_feedforward_layernorm_weights; + reference_tensor.load_weights( + ref_pre_feedforward_layernorm_weights, + "Gemma4VisionEncoderLayer_" + std::to_string(layer_idx) + "_pre_feedforward_layernorm_weights" + ); + print_error_metrics( + this->layer_norm_2_weight[layer_idx].data(), + ref_pre_feedforward_layernorm_weights.data(), + 1, + 1, this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + 1, this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE + ); + } + #endif + + simd_rms_norm( + hidden_state.data(), + this->layer_norm_2_weight[layer_idx].data(), + hidden_state.data(), + seq_len, this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_padded, this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionEncoderLayer_" + std::to_string(layer_idx)+ "_post_pre_feedforward_layernorm_hidden_states:" << std::endl; + buffer ref_post_pre_feedforward_layernorm_hidden_states; + reference_tensor.load_weights( + ref_post_pre_feedforward_layernorm_hidden_states, + "Gemma4VisionEncoderLayer_" + std::to_string(layer_idx) + "_post_pre_feedforward_layernorm_hidden_states" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + ref_post_pre_feedforward_layernorm_hidden_states.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + } + } + #endif + + simd_clamp( + hidden_state.data(), + gate_input.data(), + this->input_gate_min[layer_idx], this->input_gate_max[layer_idx], + seq_len * Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + gate_input.sync_to_device(); + + generate_mm_sequence(*this->gate_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_VISION_HIDDEN_SIZE ,Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE, + MM_tile_M,MM_tile_K,MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 1, // gelu for gate + 1, (float)this->output_gate_min[layer_idx],(float) this->output_gate_max[layer_idx], // clamp output to quantization range for v_proj + ENABLE_QKV_REORDER, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + + gate_input.sync_to_device(); + this->gate_proj_weight[layer_idx].sync_to_device(); + auto gate_proj_run = FLM_OVERRIDE(vision_gate_proj, this->gate_proj_app.create_run( + gate_input, gate_proj_weight[layer_idx], gate_output + ), layer_idx); + gate_proj_run.start(); + + simd_clamp( + hidden_state.data(), + up_input.data(), + this->input_up_min[layer_idx], this->input_up_max[layer_idx], + seq_len * Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + up_input.sync_to_device(); + generate_mm_sequence(*this->up_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_VISION_HIDDEN_SIZE ,Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE, + MM_tile_M,MM_tile_K,MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, + 1, (float)this->output_up_min[layer_idx],(float) this->output_up_max[layer_idx], // clamp output to quantization range for v_proj + ENABLE_QKV_REORDER, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + + gate_proj_run.wait(); + gate_output.sync_from_device(); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionEncoderLayer_" + std::to_string(layer_idx)+ "_gate_proj_act:" << std::endl; + buffer ref_gate_proj_output; + reference_tensor.load_weights( + ref_gate_proj_output, + "Gemma4VisionMLP_layer_" + std::to_string(layer_idx) + "_gate_proj_act" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + gate_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE, + ref_gate_proj_output.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_INTERMEDIATE_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_INTERMEDIATE_SIZE; // use unpadded ref stride + } + } + #endif + + up_input.sync_to_device(); + this->up_proj_weight[layer_idx].sync_to_device(); + auto up_proj_run = FLM_OVERRIDE(vision_up_proj, this->up_proj_app.create_run( + up_input, up_proj_weight[layer_idx], up_output + ), layer_idx); + up_proj_run.start(); + + generate_mm_sequence(*this->down_proj_app.seq(), + seq_len_padded, Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE, Padded_GEMMA4E_VISION_HIDDEN_SIZE , + MM_tile_M,MM_tile_K,MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, + 1, (float)this->output_down_min[layer_idx],(float) this->output_down_max[layer_idx], // clamp output to quantization range for v_proj + ENABLE_QKV_REORDER, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + + up_proj_run.wait(); + up_output.sync_from_device(); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionEncoderLayer_" + std::to_string(layer_idx)+ "_up_proj_output:" << std::endl; + buffer ref_up_proj_output; + reference_tensor.load_weights( + ref_up_proj_output, + "Gemma4VisionMLP_layer_" + std::to_string(layer_idx) + "_up_proj" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + up_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE, + ref_up_proj_output.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_INTERMEDIATE_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_INTERMEDIATE_SIZE; // use unpadded ref stride + } + } + #endif + + simd_mul( + gate_output.data(), + up_output.data(), + gate_output.data(), // write back to gate_output buffer to save memory + seq_len * Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE + ); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionEncoderLayer_" + std::to_string(layer_idx)+ "_after_act:" << std::endl; + buffer ref_post_gate_mul_hidden_states; + reference_tensor.load_weights( + ref_post_gate_mul_hidden_states, + "Gemma4VisionMLP_layer_" + std::to_string(layer_idx) + "_after_act" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + gate_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE, + ref_post_gate_mul_hidden_states.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_INTERMEDIATE_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_INTERMEDIATE_SIZE; // use unpadded ref stride + } + } + #endif + + simd_clamp( + gate_output.data(), + gate_output.data(), + this->input_down_min[layer_idx], this->input_down_max[layer_idx], + seq_len * Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + gate_output.sync_to_device(); + this->down_proj_weight[layer_idx].sync_to_device(); + FLM_OVERRIDE(vision_down_proj, this->down_proj_app(gate_output, this->down_proj_weight[layer_idx], down_output), layer_idx); + down_output.sync_from_device(); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionEncoderLayer_" + std::to_string(layer_idx)+ "_post_mlp_hidden_states:" << std::endl; + buffer ref_post_mlp_hidden_states; + reference_tensor.load_weights( + ref_post_mlp_hidden_states, + "Gemma4VisionEncoderLayer_" + std::to_string(layer_idx) + "_post_mlp_hidden_states" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + down_output.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + ref_post_mlp_hidden_states.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + } + } + #endif + + simd_rms_norm( + down_output.data(), + this->post_ffn_norm_weight[layer_idx].data(), + hidden_state.data(), + seq_len, this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_padded, this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionEncoderLayer_" + std::to_string(layer_idx)+ "_post_feedforward_layernorm_hidden_states:" << std::endl; + buffer ref_post_ffn_layernorm_hidden_states; + reference_tensor.load_weights( + ref_post_ffn_layernorm_hidden_states, + "Gemma4VisionEncoderLayer_" + std::to_string(layer_idx) + "_post_feedforward_layernorm_hidden_states" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + ref_post_ffn_layernorm_hidden_states.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + } + } + #endif + + simd_add( + hidden_state.data(), + residual_buffer.data(), + hidden_state.data(), + seq_len * this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "Error for Gemma4VisionEncoderLayer_" + std::to_string(layer_idx)+ "_final_hidden_states:" << std::endl; + buffer ref__hidden_states_after_ffn_residual; + reference_tensor.load_weights( + ref__hidden_states_after_ffn_residual, + "Gemma4VisionEncoderLayer_" + std::to_string(layer_idx) + "_final_hidden_states" + ); + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + hidden_state.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + ref__hidden_states_after_ffn_residual.data() + ref_offset, + 1, + seq_len_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + seq_len_per_image[i], Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + } + } + #endif + } + + // //TODO: FIXME: DEBUG + // { + + // buffer ref__hidden_states_after_ffn_residual; + // reference_tensor.load_weights( + // ref__hidden_states_after_ffn_residual, + // "Gemma4VisionEncoderLayer_" + std::to_string(this->parent_npu_ptr->GEMMA4E_VISION_NUM_HIDDEN_LAYERS-1) + "_final_hidden_states" + // ); + // for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + // memcpy( + // hidden_state.data() + start_seq_len_index_per_image[i]*Padded_GEMMA4E_VISION_HIDDEN_SIZE, + // ref__hidden_states_after_ffn_residual.data() + ref_offset, + // seq_len_per_image[i] * this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE * sizeof(bf16) + // ); + // ref_offset+= image_payload->image_patch__element_per_patch[i].first*this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; // use unpadded ref stride + // } + + // } + + // // debug, print all the content in seq_len_per_image + + // the vision pooler stage + std::vector k_per_image; + std::vector k_squared_per_image; + for(int i = 0; i < image_payload->num_images; i++){ + + int k = std::sqrt( seq_len_per_image[i]/ image_payload->num_soft_tokens_per_image[i] ); + k_per_image.push_back(k); + k_squared_per_image.push_back(k*k); + } + + std::vector max_x; + + for(int i = 0; i < image_payload->num_images; i++){ + int cur_max_x = -1; + for(int j = 0; j < seq_len_per_image[i]; j++){ + + auto cur_x = image_payload->image_grid_pairs_per_image[i][2*j]; + if(cur_x > cur_max_x){ + cur_max_x = cur_x; + } + } + max_x.push_back(cur_max_x + 1); + } + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + + std::cout << "compare max x" < Gemma4VisionPooler_max_x; + reference_tensor.load_weights( + Gemma4VisionPooler_max_x, + "Gemma4VisionPooler_max_x" + ); + for(int i = 0; i < image_payload->num_images; i++){ + if(max_x[i] != Gemma4VisionPooler_max_x.data()[i]){ + std::cout << "max_x mismatch for image " << i << ": " << max_x[i] << " vs ref " << Gemma4VisionPooler_max_x.data()[i] << std::endl; + } + } + } + #endif + + // now, we divide every image's image_grid_pairs_per_image by k (floor division), matching Python: kernel_idxs = floor(pos / k) + std::vector> kernel_idx(image_payload->num_images); + for(int i = 0; i < image_payload->num_images; i++){ + + for(int j = 0; j < seq_len_per_image[i]; j++){ + image_payload->image_grid_pairs_per_image[i][2*j] /= (float)k_per_image[i]; //NOTE: store back to int, same as round_down in floor mode + image_payload->image_grid_pairs_per_image[i][2*j + 1] /= (float)k_per_image[i]; + + kernel_idx.at(i).push_back( + image_payload->image_grid_pairs_per_image[i][2*j] + (max_x[i] /k_per_image[i] ) * image_payload->image_grid_pairs_per_image[i][2*j+1] + ); + } + } + + // now, compare the error of kernel_idx with "Gemma4VisionPooler_kernel_idxs" + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "compare kernel idx" < Gemma4VisionPooler_kernel_idxs; + reference_tensor.load_weights( + Gemma4VisionPooler_kernel_idxs, + "Gemma4VisionPooler_kernel_idxs" + ); + + int ref_per_image = Gemma4VisionPooler_kernel_idxs.size() / image_payload->num_images; + for(int i = 0, ref_offset=0; i < image_payload->num_images; i++){ + for(int j = 0; j < seq_len_per_image[i]; j++){ + if(kernel_idx[i][j] != Gemma4VisionPooler_kernel_idxs.data()[ref_offset + j]){ + std::cout << "kernel_idx mismatch for image " << i << " token " << j << ": " << kernel_idx[i][j] << " vs ref " << Gemma4VisionPooler_kernel_idxs.data()[ref_offset + j] << std::endl; + } + } + ref_offset+=ref_per_image; + } + } + #endif + + //NOTE: num_soft_token_per_image is in image_payload->num_soft_tokens_per_image[i]; + // L_per_image is in seq_len_per_image[i] + std::vector pooling_output; + //hidden_state is a row-major buffer + // at this point, the hidden_state is in shape of [num_image,seq_len_per_image[], Padded_GEMMA4E_VISION_HIDDEN_SIZE], we will do pooling for each image separately, and the pooling weight will be generated based on kernel_idx and k_squared_per_image (which is the number of tokens in each pooling region), and the pooling weight will be shape of [num_image, seq_len_per_image, num_soft_token_per_image] + //NOTE: seq_len_per_image[] means this varies from image to image, and num_soft_token_per_image is in image_payload->num_soft_tokens_per_image[i] + + //NOTE: the python code generates a matrix that is just too sparse, we use scatter-add instead + // output = weights.transpose(1, 2) @ hidden_states.float() + // equivalently: output[t, h] = (1/k²) * Σ hidden_states[l, h] for all l where kernel_idx[l] == t + + float root_hidden_size = std::sqrt((float)this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE); + for(int i = 0; i < image_payload->num_images; i++){ + const int cur_seq_len = seq_len_per_image[i]; + const int cur_soft_tokens = image_payload->num_soft_tokens_per_image[i]; + const int HIDDEN_SIZE = this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE; + const float scale = 1.0f / (float)k_squared_per_image[i]; + + // accumulate in float for precision + std::vector accum(cur_soft_tokens * HIDDEN_SIZE, 0.0f); + + bf16* hidden_base = hidden_state.data() + start_seq_len_index_per_image[i] * Padded_GEMMA4E_VISION_HIDDEN_SIZE; + for(int l = 0; l < cur_seq_len; l++){ + const int t = kernel_idx[i][l]; + bf16* src = hidden_base + l * Padded_GEMMA4E_VISION_HIDDEN_SIZE; + float* dst = accum.data() + t * HIDDEN_SIZE; + for(int h = 0; h < HIDDEN_SIZE; h++){ + dst[h] += (float)src[h] * scale; + } + } + + // convert back to bf16 + int cur_pooling_output_size = pooling_output.size(); + pooling_output.resize(cur_pooling_output_size + cur_soft_tokens * Padded_GEMMA4E_VISION_HIDDEN_SIZE, bf16(0)); + for(int t = 0; t < cur_soft_tokens; t++){ + for(int h = 0; h < HIDDEN_SIZE; h++){ + // pooling_output[i][t * Padded_GEMMA4E_VISION_HIDDEN_SIZE + h] = bf16(accum[t * HIDDEN_SIZE + h]); + pooling_output[cur_pooling_output_size + t * Padded_GEMMA4E_VISION_HIDDEN_SIZE + h] = bf16(accum[t * HIDDEN_SIZE + h] * root_hidden_size ); // add a scaling factor to prevent overflow, matching the implementation in python code + } + } + } + + size_t vision_output_token_size = pooling_output.size() / this->Padded_GEMMA4E_VISION_HIDDEN_SIZE; + size_t padded_vision_output_token_size = round_up_to_multiple(vision_output_token_size, seq_len_pad_requirement_for_MM); + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "compare Gemma4VisionPooler_hidden_states_after_scaling" << std::endl; + buffer ref_pooling_output; + reference_tensor.load_weights( + ref_pooling_output, + "Gemma4VisionPooler_hidden_states_after_scaling" + ); + + int ref_per_image = ref_pooling_output.size() / image_payload->num_images; + for(int i = 0,pool_offset=0, ref_offset=0; i < image_payload->num_images; i++){ + print_error_metrics( + pooling_output.data() + pool_offset, + ref_pooling_output.data() + ref_offset, + 1, + image_payload->num_soft_tokens_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + image_payload->num_soft_tokens_per_image[i], this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE + ); + ref_offset += ref_per_image; + pool_offset += image_payload->num_soft_tokens_per_image[i] * Padded_GEMMA4E_VISION_HIDDEN_SIZE; + } + } + #endif + + buffer language_embed_input = this->vision_to_language_input_projection_app.create_bo_buffer( + padded_vision_output_token_size* Padded_GEMMA4E_VISION_HIDDEN_SIZE + ); + memset(language_embed_input.data(), 0, padded_vision_output_token_size* Padded_GEMMA4E_VISION_HIDDEN_SIZE * sizeof(bf16)); // zero padding for padded tokens + + buffer language_embed_output = this->vision_to_language_input_projection_app.create_bo_buffer( + padded_vision_output_token_size * this->parent_npu_ptr->GEMMA4E_VISION_IMAGE_OUTPUT_SIZE + ); + + simd_rms_norm( + pooling_output.data(), + language_embed_input.data(), + vision_output_token_size, this->parent_npu_ptr->GEMMA4E_VISION_HIDDEN_SIZE, + padded_vision_output_token_size, this->Padded_GEMMA4E_VISION_HIDDEN_SIZE + + ); + + generate_mm_sequence(*this->vision_to_language_input_projection_app.seq(), + padded_vision_output_token_size, Padded_GEMMA4E_VISION_HIDDEN_SIZE,this->parent_npu_ptr->GEMMA4E_VISION_IMAGE_OUTPUT_SIZE, //TODO: FIXME: + MM_tile_M,MM_tile_K,MM_tile_N, + 8,8,8, + rtp_address, rtp_sync_lock_id, + MM_ROW_SIZE,MM_COL_SIZE, + 0,0,0, + IS_B_ROW_MAJOR, ENABLE_AXI4, true, + false, 0, /// no biase + 0, -10000.0, 1000000.0, // do not clamp on output + ENABLE_QKV_REORDER, this->parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM + ); + + language_embed_input.sync_to_device(); + this->vision_to_language_input_projection_weight.sync_to_device(); + FLM_OVERRIDE(vision_to_language_proj, this->vision_to_language_input_projection_app(language_embed_input, this->vision_to_language_input_projection_weight, language_embed_output)); + language_embed_output.sync_from_device(); + + #if DEBUG_PRINT_ENCODE_ERROR_METRICS + { + std::cout << "compare vision to language projection output" << std::endl; + buffer ref_vision_to_language_projection_output; + reference_tensor.load_weights( + ref_vision_to_language_projection_output, + "vision_final_embs_after_embed_vision" + ); + std::cout << "finished laoding ref vision to language projection output" << std::endl; + print_error_metrics( + language_embed_output.data(), + ref_vision_to_language_projection_output.data(), + 1, + vision_output_token_size, this->parent_npu_ptr->GEMMA4E_VISION_IMAGE_OUTPUT_SIZE, + padded_vision_output_token_size, this->parent_npu_ptr->GEMMA4E_VISION_IMAGE_OUTPUT_SIZE + + ); + } + #endif + + ///TODO: consider change the output type interface + std::vector final_res(padded_vision_output_token_size *this->parent_npu_ptr->GEMMA4E_VISION_IMAGE_OUTPUT_SIZE ); + memcpy( + final_res.data(), + language_embed_output.data(), + padded_vision_output_token_size * this->parent_npu_ptr->GEMMA4E_VISION_IMAGE_OUTPUT_SIZE * sizeof(bf16) + ); + + return final_res; +} diff --git a/src/detail/gemma4e_npu/gemma4e_image.hpp b/src/detail/gemma4e_npu/gemma4e_image.hpp new file mode 100644 index 000000000..6cb40fe7e --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_image.hpp @@ -0,0 +1,136 @@ +#pragma once +#include "typedef.hpp" +#include +#include +#include "tensor_utils/q4_npu_eXpress.hpp" +#include "models/gemma4e/flm/aie2p/gemma4e_npu.hpp" + +#include "vision/norm.hpp" + +class Gemma4e_ImageEncoder{ + + public: + + ~Gemma4e_ImageEncoder(); + void init_weights(SafeTensors &q4nx); + Gemma4e_ImageEncoder(LM_Config config, npu_xclbin_manager *npu_instance, gemma4e_npu* parent_npu_ptr); + + std::vector encode( void* image_payload_ptr); + + LM_Config config; + npu_xclbin_manager *npu; + gemma4e_npu* parent_npu_ptr; + + unsigned int Padded_GEMMA4E_VISION_HIDDEN_SIZE; + unsigned int Padded_GEMMA4E_VISION_MLP_INTERMEDIATE_SIZE; + unsigned int Padded_GEMMA4E_VISION_OUT_HIDDEN_SIZE; + + uint32_t MM_tile_M; + uint32_t MM_tile_K; + uint32_t MM_tile_N; + + uint32_t seq_len_pad_requirement_for_MM; + uint32_t MM_ROW_SIZE = 4; + uint32_t MM_COL_SIZE = 8; + + // parameter for vision attention kernel + uint32_t vision_num_of_columns = 8; + uint32_t vision_num_of_rows = 4; + uint32_t vision_CU_mode = 2; + uint32_t vision_LQ_per_CT = 32; + uint32_t vision_LK_per_CT = 512; + uint32_t vision_LQ_internal=32; + uint32_t vision_LK_internal=32; + uint32_t vision_L_padded_requirement_for_attention = 512; // padding for L_seq in vision attention + uint32_t vision_S_padded_requirement_for_attention = vision_LK_per_CT; // padding for S_seq in vision attention + //uint32_t vision_L_padded_requirement_for_attention = 32; + + // IF true, reorder qkv from L_Seq x (3*QWEN3_5_VISION_NUM_HEADS*QWEN3_VISION_HEAD_DIM) row major -> + // [ 3, QWEN3_5_VISION_NUM_HEADS, L_Seq ,QWEN3_VISION_HEAD_DIM] + bool ENABLE_QKV_REORDER = false; + // for MM runtime sequence + uint32_t rtp_address = 4096; // offset right after stack size + uint32_t rtp_sync_lock_id = 10; // the rtp sync lock + bool ENABLE_AXI4 = true; + bool IS_B_ROW_MAJOR = false; + + // debug ptr for now + std::string model_path; + + npu_app_manager* fla; + npu_app_manager* proj; + npu_app_manager* proj_high_precision; + + // define the necessary bitstream no + npu_app flash_attention_app; + npu_app patch_embedder_posisiton_embedding_dim_0_app; + npu_app patch_embedder_posisiton_embedding_dim_1_app; + npu_app patch_embedding_app; + npu_app q_proj_app; + npu_app k_proj_app; + npu_app v_proj_app; + npu_app o_proj_app; + npu_app gate_proj_app; + npu_app up_proj_app; + npu_app down_proj_app; + + npu_app vision_to_language_input_projection_app; + + // The weights + + buffer patch_embedder_position_embedding_table; + buffer patch_embd_weight; + buffer vision_to_language_input_projection_weight; // project into language hidden space, + + std::vector> q_proj_weight; + std::vector> k_proj_weight; + std::vector> v_proj_weight; + std::vector> o_proj_weight; + std::vector> gate_proj_weight; + std::vector> up_proj_weight; + std::vector> down_proj_weight; + + std::vector hidden_state; + std::vector residual_buffer; + std::vector> q_norm_weight; + std::vector> k_norm_weight; + std::vector> post_o_norm_weight; + std::vector> post_ffn_norm_weight; + std::vector> layer_norm_1_weight; + std::vector> layer_norm_2_weight; + + std::vector input_q_max; + std::vector input_q_min; + std::vector input_k_max; + std::vector input_k_min; + std::vector input_v_max; + std::vector input_v_min; + std::vector input_o_max; + std::vector input_o_min; + std::vector input_gate_max; + std::vector input_gate_min; + std::vector input_up_max; + std::vector input_up_min; + std::vector input_down_max; + std::vector input_down_min; + + std::vector output_q_max; + std::vector output_q_min; + std::vector output_k_max; + std::vector output_k_min; + std::vector output_v_max; + std::vector output_v_min; + std::vector output_o_max; + std::vector output_o_min; + std::vector output_gate_max; + std::vector output_gate_min; + std::vector output_up_max; + std::vector output_up_min; + std::vector output_down_max; + std::vector output_down_min; + + inline int round_up_to_multiple (int x, int multiple) + { + return ((x + multiple - 1) / multiple) * multiple; + }; +}; diff --git a/src/detail/gemma4e_npu/gemma4e_npu.cpp b/src/detail/gemma4e_npu/gemma4e_npu.cpp new file mode 100644 index 000000000..20f87d42d --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_npu.cpp @@ -0,0 +1,860 @@ +#include +#include "flm_override.hpp" +#include "gemma4e_npu_detail.hpp" +#include "metrices.hpp" +#include "utils/error_measure.hpp" +#include "mmRuntimeSequence.hpp" + +gemma4e_npu::Impl::Impl(LM_Config config, npu_xclbin_manager *npu_instance, gemma4e_npu* parent_ptr, int MAX_L ) + : config(config), npu(npu_instance), parent_ptr(parent_ptr){ + FLM_OVERRIDE(engine_init, (void)0, this->npu, this->config); + is_vlm = config.get("is_vlm", false); + is_audio = config.get("is_audio", false); + + current_context_length = 0; + + this->MAX_L = std::max(MAX_L, 4096); + // make MAX_L a power of 2 + float log2_MAX_L = std::log2(this->MAX_L); + if (log2_MAX_L != std::floor(log2_MAX_L)){ + this->MAX_L = 1 << ((int)std::floor(log2_MAX_L) + 1); + } + MAX_L = this->MAX_L; // synchronize MAX_L + is_preload_launched = false; + + + try{ + this->desc.build(this->config); + + D = desc.D; + vocab_size = desc.vocab_size; + vocab_size_padded = desc.vocab_size_padded; + DH = desc.DH; + DQ = desc.DQ; + DK = desc.DK; + DV = desc.DV; + SWA_DH = desc.SWA_DH; + SWA_DQ = desc.SWA_DQ; + SWA_DK = desc.SWA_DK; + SWA_DV = desc.SWA_DV; + num_hidden_layers = desc.num_hidden_layers; + SLIDING_LENGTH = desc.SLIDING_LENGTH; + PLI_D = desc.PLI_D; + num_kv_shared_layers = desc.num_kv_shared_layers; + final_logit_softcapping = desc.final_logit_softcapping; + non_skip_layers = desc.non_skip_layers; + INTERMEDIATE_SIZE = desc.INTERMEDIATE_SIZE; + enable_double_wide_mlp = desc.enable_double_wide_mlp; + global_layer_period = desc.global_layer_period; + layer_types = desc.layer_types; + + if (this->is_global_layer_idx(non_skip_layers - 1)){ + last_global_kv_cache_layer_idx = non_skip_layers - 1; + last_swa_kv_cache_layer_idx = non_skip_layers - 2; + } + else { + last_swa_kv_cache_layer_idx = non_skip_layers - 1; + last_global_kv_cache_layer_idx = (last_swa_kv_cache_layer_idx / global_layer_period) * global_layer_period - 1; + } + DEBUG_BLOCK(1, + header_print_g("info", "Gemma4e NPU config:"); + std::cout << "\tD: " << D << std::endl; + std::cout << "\tPLI_D: " << PLI_D << std::endl; + std::cout << "\tDH: " << DH << std::endl; + std::cout << "\tDQ: " << DQ << std::endl; + std::cout << "\tDK: " << DK << std::endl; + std::cout << "\tDV: " << DV << std::endl; + std::cout << "\tSWA_DH: " << SWA_DH << std::endl; + std::cout << "\tSWA_DQ: " << SWA_DQ << std::endl; + std::cout << "\tSWA_DK: " << SWA_DK << std::endl; + std::cout << "\tSWA_DV: " << SWA_DV << std::endl; + std::cout << "\tSLIDING_LENGTH: " << SLIDING_LENGTH << std::endl; + std::cout << "\tINTERMEDIATE_SIZE: " << INTERMEDIATE_SIZE << std::endl; + std::cout << "\tFINAL_LOGIT_SOFTCAPPING: " << final_logit_softcapping << std::endl; + std::cout << "\tvocab_size: " << vocab_size << " (" << vocab_size_padded << " padded)" << std::endl; + std::cout << "\tnum_hidden_layers: " << num_hidden_layers << std::endl; + std::cout << "\tnum_kv_shared_layers: " << num_kv_shared_layers << std::endl; + std::cout << "\tnon_skip_layers: " << non_skip_layers << std::endl; + std::cout << "\tglobal_layer_period: " << global_layer_period << std::endl; + std::cout << "\tlast_swa_kv_cache_layer_idx: " << last_swa_kv_cache_layer_idx << std::endl; + std::cout << "\tlast_global_kv_cache_layer_idx: " << last_global_kv_cache_layer_idx << std::endl; + + ) + DEBUG_BLOCK(2, + std::cout << "\tlayer_types: \n"; + for (int i = 0; i < num_hidden_layers; i++){ + std::string type_str; + switch (layer_types[i]){ + case e_gemma4e_swa_layer: + type_str = "SWA"; + break; + case e_gemma4e_global_layer: + type_str = "Global"; + break; + case e_gemma4e_swa_layer_skip: + type_str = "SWA_Skip"; + break; + case e_gemma4e_global_layer_skip: + type_str = "Global_Skip"; + break; + default: + type_str = "Unknown"; + } + std::cout << "\t\t layer " << i << ": " << type_str << " " << std::endl; + } + std::cout << std::endl; + ) + } + catch (std::exception& e){ + header_print_r("ERROR", "Failed to parse model config: " << e.what()); + throw e; + } + + gemma4e_seq_gen_parameters_t seq_gen_params = { + .D = D, + .DH = DH, + .DQ = DQ, + .DK = DK, + .DV = DV, + .SWA_DH = SWA_DH, + .SWA_DQ = SWA_DQ, + .SWA_DK = SWA_DK, + .SWA_DV = SWA_DV, + .PLI_D = PLI_D, + .INTERMEDIATE_SIZE = INTERMEDIATE_SIZE, + .NUM_ATTENTION_HEADS = (int)config.get("num_attention_heads"), + .NUM_KEY_VALUE_HEADS = (int)config.get("num_key_value_heads"), + .SLIDING_WINDOW_SIZE = SLIDING_LENGTH, + .VOCAB_SIZE_PADDED = vocab_size_padded, + .enable_double_wide_mlp = enable_double_wide_mlp + }; + this->sequence = std::make_unique(seq_gen_params, MAX_L); + // Every sequence below addresses weights through the descriptors, so bind the + // description before generating any of them. + this->sequence->set_desc(&this->desc); + + DEBUG_BLOCK(1, + header_print_g("info", "VLM Enabled: " << (is_vlm ? "Yes" : "No")); + header_print_g("info", "Audio Model Enabled: " << (is_audio ? "Yes" : "No")); + ); + + if (is_vlm){ + this->gemma4e_image_encoder = std::make_unique(config, npu, this->parent_ptr); + } + if(is_audio){ + this->gemma4e_audio_encoder = std::make_unique(config, npu, this->parent_ptr); + } + // NOTE: order matter. The per layer input apps live on the image encoder's + // manager and must be created before layer.xclbin is registered, so this block + // is built here and handed to the prefill context at the end of the ctor. + std::unique_ptr pli_block = + std::make_unique(&this->desc, nullptr, this->gemma4e_image_encoder.get()); + layer_app_manager = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "layer.xclbin")); + + this->global_layer = layer_app_manager->create_app(); + this->swa_layer = layer_app_manager->create_app(); + this->global_skip_layer = layer_app_manager->create_app(); + this->swa_skip_layer = layer_app_manager->create_app(); + this->layer_pre_load = layer_app_manager->create_app(); + + if (!this->npu->is_preemption_enabled()){ + this->layers_run = layer_app_manager->create_runlist(); + } + + lm_head_app_manager = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "lm_head.xclbin")); + this->lm_head = lm_head_app_manager->create_app(); + + // allocate all buffers + rms_weights.resize(num_hidden_layers); + proj_weights.resize(num_hidden_layers); + pli_gate_up_weights.resize(num_hidden_layers); + kv_caches.resize(non_skip_layers); // only allocate kv_cache for non-skip layers + rope_rms_weights.resize(num_hidden_layers); + layer_scale = buffer(num_hidden_layers); + + this->x = layer_app_manager->create_bo_buffer((3 * D + 8191) / 8192 * 8192); // Iterative Hidden States (D) + Final RMS NORM (D) + INITIAL EMBEDDING (D) + + for (int layer_idx = 0; layer_idx < num_hidden_layers; layer_idx++){ + gemma4e_layer_type_t type = layer_types[layer_idx]; + size_t proj_buffer_size = desc.get_proj_weights_byte_size(type); + proj_weights[layer_idx] = layer_app_manager->create_bo_buffer(proj_buffer_size); // allocate 32MB for each layer, which is enough for current model scale. For larger model scale, we may need to dynamically load weights + this->proj_weights[layer_idx].memset((uint8_t)0); + this->proj_weights[layer_idx].sync_to_device(); + this->pli_gate_up_weights[layer_idx] = npu->create_bo_buffer(PLI_D * D * 2); // gate and up projection weights for per layer input, which will be added to the input of each layer after the first layer + DEBUG_BLOCK(2, + header_print("info", "Buffer sizes (in bytes): proj_weights=" + std::to_string(proj_weights[layer_idx].size())); + ) + } + + for (int layer_idx = 0; layer_idx < num_hidden_layers; layer_idx++){ + rms_weights[layer_idx] = layer_app_manager->create_bo_buffer(desc.get_rms_elems(layer_types[layer_idx])); // input_layer_norm (D) + post_attn_layer_norm (D) + pre_ffn_layer_norm (D) + post_ffn_layer_norm (D) for each layer from host. + } + for (int layer_idx = 0; layer_idx < num_hidden_layers; layer_idx++){ + gemma4e_layer_type_t type = layer_types[layer_idx]; + rope_rms_weights[layer_idx] = layer_app_manager->create_bo_buffer(desc.get_rope_rms_elems(type)); // COS/SIN (_DH) + Q_NORM (_DH) + K_NORM (_DH) + PLI_EMBED (PLI_D) + PLI_NORM (PLI_D) + POST_PLI_NORM (D) + layer_scale + } + + for (int layer_idx = 0; layer_idx < non_skip_layers; layer_idx++){ + gemma4e_layer_type_t type = layer_types[layer_idx]; + kv_caches[layer_idx] = layer_app_manager->create_bo_buffer(desc.get_kv_cache_size(type, MAX_L)); + kv_caches[layer_idx].memset((bf16)0); + kv_caches[layer_idx].sync_to_device(); + DEBUG_BLOCK(2, + header_print("info", "Buffer sizes (in bytes): kv_cache for layer " + std::to_string(layer_idx) + " = " + std::to_string(kv_caches[layer_idx].size())); + ) + } + is_checkpoint_valid = false; // checkpoint is not valid until we load kv cache to device after resizing buffers + pli_down_weights = layer_app_manager->create_bo_buffer(num_hidden_layers * PLI_D * D); // down projection weights for per layer input, which will be added to the input of each layer after the first layer + + lm_head_weights = lm_head_app_manager->create_bo_buffer(desc.get_lm_head_w_size()); + logits = lm_head_app_manager->create_bo_buffer(vocab_size_padded); + logits_valid = buffer(logits.data(), vocab_size); + + this->embedding = std::make_unique(vocab_size, D); + this->pli_embedding = std::make_unique(vocab_size, PLI_D * num_hidden_layers); // per layer input embedding, which will be added to the input of each layer after the first layer + + // Everything below this point belongs to the prefill path: its own xclbins, + // sequences, dequantized weights and scratch buffers. + this->prefill_ctx = std::make_unique( + npu, &this->desc, this->config, this->sequence.get(), std::move(pli_block), MAX_L + ); + + this->sequence->gen_lm_head_seq(this->lm_head.seq(), final_logit_softcapping); + this->lm_head_run = FLM_OVERRIDE(lm_head, + this->lm_head.create_run(this->logits, this->lm_head_weights, this->x)); + + // empty seq for preload xclbin + npu_sequence* pre_load_seq = this->layer_pre_load.seq(); + pre_load_seq->clear_cmds(); + pre_load_seq->cmds2seq(); + + // this->sequence->gen_layer_seq(this->linear_layer.seq(), 0, false); +} + +void gemma4e_npu::Impl::_set_rope_rms_weights(int idx){ + + buffer cos_buf(DH / 2); + buffer sin_buf(DH / 2); + buffer swa_cos_buf(SWA_DH / 2); + buffer swa_sin_buf(SWA_DH / 2); + + for (int j = 0; j < DH / 2; j++){ + cos_buf[j] = (bf16)cos(gemma4e_cpu_func::inv_freq_global[j] * idx); + sin_buf[j] = (bf16)sin(gemma4e_cpu_func::inv_freq_global[j] * idx); + } + for (int j = 0; j < SWA_DH / 2; j++){ + swa_cos_buf[j] = (bf16)cos(gemma4e_cpu_func::inv_freq_swa[j] * idx); + swa_sin_buf[j] = (bf16)sin(gemma4e_cpu_func::inv_freq_swa[j] * idx); + } + + for (int i = 0; i < num_hidden_layers; i++){ + bf16* w_rope_ptr = this->rope_rms_weights[i].data(); + if (is_swa_layer(layer_types[i])){ + memcpy(w_rope_ptr, swa_cos_buf.data(), SWA_DH * sizeof(bf16) / 2); + memcpy(w_rope_ptr + SWA_DH / 2, swa_sin_buf.data(), SWA_DH * sizeof(bf16) / 2); + } + else { + memcpy(w_rope_ptr, cos_buf.data(), DH * sizeof(bf16) / 2); + memcpy(w_rope_ptr + DH / 2, sin_buf.data(), DH * sizeof(bf16) / 2); + } + this->rope_rms_weights[i].sync_to_device(); + } +} + +void gemma4e_npu::Impl::set_context_length(int L){ + this->current_context_length = L; + DEBUG_BLOCK(2, + header_print_r("info", "Setting context length to " + std::to_string(L) + " for all layers in the sequence"); + ) + this->sequence->gen_layer_seq(this->global_layer.seq(), L + 1, e_gemma4e_global_layer); + this->sequence->gen_layer_seq(this->swa_layer.seq(), L + 1, e_gemma4e_swa_layer); + this->sequence->gen_layer_seq(this->global_skip_layer.seq(), L + 1, e_gemma4e_global_layer_skip); + this->sequence->gen_layer_seq(this->swa_skip_layer.seq(), L + 1, e_gemma4e_swa_layer_skip); + + if (!this->npu->is_preemption_enabled()){ + this->layers_run.reset(); + + DEBUG_BLOCK(2, + header_print_r("info", "Create runs for all layers in the sequence with the new context length"); + ) + for (uint32_t i = 0; i < num_hidden_layers; i++){ + + DEBUG_BLOCK(2, + header_print_r("info", "Generate run for layer " + std::to_string(i) + " of type " + std::to_string(layer_types[i])); + ) + switch(layer_types[i]){ + case e_gemma4e_global_layer: + this->layers_run.add(FLM_OVERRIDE(global_layer_run, + this->global_layer.create_run(this->x, this->proj_weights[i], this->rms_weights[i], this->rope_rms_weights[i], this->kv_caches[i]), layer_types[i], i, L, this->MAX_L)); + break; + case e_gemma4e_swa_layer: + this->layers_run.add(FLM_OVERRIDE(swa_layer_run, + this->swa_layer.create_run(this->x, this->proj_weights[i], this->rms_weights[i], this->rope_rms_weights[i], this->kv_caches[i]), layer_types[i], i, L, this->MAX_L)); + break; + case e_gemma4e_global_layer_skip: + this->layers_run.add(FLM_OVERRIDE(global_skip_layer_run, + this->global_skip_layer.create_run(this->x, this->proj_weights[i], this->rms_weights[i], this->rope_rms_weights[i], this->kv_caches[last_global_kv_cache_layer_idx]), layer_types[i], i, L, this->MAX_L)); + break; + case e_gemma4e_swa_layer_skip: + this->layers_run.add(FLM_OVERRIDE(swa_skip_layer_run, + this->swa_skip_layer.create_run(this->x, this->proj_weights[i], this->rms_weights[i], this->rope_rms_weights[i], this->kv_caches[last_swa_kv_cache_layer_idx]), layer_types[i], i, L, this->MAX_L)); + break; + } + } + } + + this->_set_rope_rms_weights(L); +} + +buffer gemma4e_npu::Impl::forward(int ids){ + if (is_preload_launched){ + this->pre_load_run.wait(); + is_preload_launched = false; + } + + _process_embedding(ids); + if (!this->npu->is_preemption_enabled()){ + this->layers_run.execute(); + this->layers_run.wait(); + } + else{ + for (uint32_t i = 0; i < num_hidden_layers; i++){ + switch(layer_types[i]){ + case e_gemma4e_global_layer: + FLM_OVERRIDE(global_layer, + this->global_layer(this->x, this->proj_weights[i], this->rms_weights[i], this->rope_rms_weights[i], this->kv_caches[i]), layer_types[i], i, this->current_context_length, this->MAX_L); + break; + case e_gemma4e_swa_layer: + FLM_OVERRIDE(swa_layer, + this->swa_layer(this->x, this->proj_weights[i], this->rms_weights[i], this->rope_rms_weights[i], this->kv_caches[i]), layer_types[i], i, this->current_context_length, this->MAX_L); + break; + case e_gemma4e_global_layer_skip: + FLM_OVERRIDE(global_skip_layer, + this->global_skip_layer(this->x, this->proj_weights[i], this->rms_weights[i], this->rope_rms_weights[i], this->kv_caches[last_global_kv_cache_layer_idx]), layer_types[i], i, this->current_context_length, this->MAX_L); + break; + case e_gemma4e_swa_layer_skip: + FLM_OVERRIDE(swa_skip_layer, + this->swa_skip_layer(this->x, this->proj_weights[i], this->rms_weights[i], this->rope_rms_weights[i], this->kv_caches[last_swa_kv_cache_layer_idx]), layer_types[i], i, this->current_context_length, this->MAX_L); + break; + } + + DEBUG_BLOCK(2, + header_print("info", "Finished layer " + std::to_string(i) + " of type " + std::to_string(layer_types[i]) + " for the current token"); + this->x.sync_from_device(); + ) + } + } + DEBUG_BLOCK(2, + std::cout << std::endl; + header_print("info", "Finished executing all layers for the current token:" + std::to_string(ids)); + this->x.sync_from_device(); + buffer valid_x = buffer(this->x.data(), D); + utils::print_matrix(valid_x, D); + this->x.sync_to_device(); + ) + + this->lm_head_run.start(); + this->set_context_length(this->current_context_length + 1); + this->lm_head_run.wait(); + this->logits.sync_from_device(); + DEBUG_BLOCK(2, + header_print("info", "Finished LM head run for the current token " + std::to_string(current_context_length)); + utils::print_matrix(this->logits_valid, vocab_size); + ) + + this->pre_load_run = FLM_OVERRIDE(layer_pre_load, this->layer_pre_load.create_run()); // create run for the preload xclbin, which has an empty sequence and will be used to preload the next layer's weights in the background while the current token is being processed + this->pre_load_run.start(); + is_preload_launched = true; + return this->logits_valid; +} + +buffer gemma4e_npu::Impl::prefill(std::vector& ids, void* payload){ + if (is_preload_launched){ + this->pre_load_run.wait(); + is_preload_launched = false; + } + DEBUG_BLOCK(1, + std::cout << "DEBUG: Entering prefill, input length: " << ids.size() << std::endl; + ) +#ifdef MVPREFILL + header_print("warning", "Prefilling with Decoder, Slow!"); + return this->_prefill_with_mv(ids); +#else + return this->_prefill_with_mm(ids, payload); +#endif +} + +buffer gemma4e_npu::Impl::_prefill_with_mv(std::vector& ids, void* payload){ + DEBUG_BLOCK(1, + header_print("info", "Entering prefill with MV, input length: " + std::to_string(ids.size())); + ) + for (int i = 0; i < ids.size(); i++){ + _process_embedding(ids[i]); + for (uint32_t l = 0; l < num_hidden_layers; l++){ + + DEBUG_BLOCK(2, + header_print_r("info", "Running layer " + std::to_string(l) + " of type " + std::to_string(layer_types[l]) + " for token " + std::to_string(ids[i])); + ) + // int test_layer = 15; + switch(layer_types[l]){ + case e_gemma4e_global_layer: + FLM_OVERRIDE(global_layer_mv, this->global_layer(this->x, this->proj_weights[l], this->rms_weights[l], this->rope_rms_weights[l], this->kv_caches[l]), layer_types[l], l); + break; + case e_gemma4e_swa_layer: + FLM_OVERRIDE(swa_layer_mv, this->swa_layer(this->x, this->proj_weights[l], this->rms_weights[l], this->rope_rms_weights[l], this->kv_caches[l]), layer_types[l], l); + break; + case e_gemma4e_global_layer_skip: + FLM_OVERRIDE(global_skip_layer_mv, this->global_skip_layer(this->x, this->proj_weights[l], this->rms_weights[l], this->rope_rms_weights[l], this->kv_caches[last_global_kv_cache_layer_idx]), layer_types[l], l); + break; + case e_gemma4e_swa_layer_skip: + FLM_OVERRIDE(swa_skip_layer_mv, this->swa_skip_layer(this->x, this->proj_weights[l], this->rms_weights[l], this->rope_rms_weights[l], this->kv_caches[last_swa_kv_cache_layer_idx]), layer_types[l], l); + break; + } + DEBUG_BLOCK(2, + header_print("info", "Finished layer " + std::to_string(l) + " of type " + std::to_string(layer_types[l]) + " for the current token"); + this->x.sync_from_device(); + this->x.sync_to_device(); + buffer valid_x = buffer(this->x.data(), D); + utils::print_matrix(valid_x, 256); + ) + } + + DEBUG_BLOCK(2, exit(0); ) + this->set_context_length(this->current_context_length + 1); + } + + DEBUG_BLOCK(1, + header_print("info", "Finished all layers"); + this->x.sync_from_device(); + this->x.sync_to_device(); + buffer valid_x = buffer(this->x.data(), D); + utils::print_matrix(valid_x, 256); + ) + this->lm_head_weights.sync_to_device(); + this->lm_head_run.start(); + this->lm_head_run.wait(); + this->logits.sync_from_device(); + + return logits_valid; +} + +buffer gemma4e_npu::Impl::_prefill_with_mm(std::vector& ids, void* payload){ + int L_in = ids.size(); + + std::unique_ptr reference; + buffer input_ids; + DEBUG_BLOCK(2, + std::cout << "DEBUG: Entering prefill with mm, input length: " << ids.size() << std::endl; + std::string reference_path = utils::path_join(config.model_path, "gemma4_ref.safetensors"); + reference = std::make_unique(reference_path); + reference->load_weights(input_ids, "input_ids"); + L_in = input_ids.size(); + ) + + // Sizes the batch and readies every block; sequences and buffers that already + // match this geometry are kept as they are. + const gemma4e_prefill_shape s = this->prefill_ctx->setup(L_in, this->current_context_length); + gemma4e_common_buffers& bufs = this->prefill_ctx->bufs; + + DEBUG_BLOCK(1, + std::cout << "DEBUG: L_begin: " << s.L_begin << ", L_end: " << s.L_end << std::endl; + std::cout << "DEBUG: L_begin_chunked: " << s.L_begin_chunked << ", L_end_chunked: " << s.L_end_chunked << ", L_padded: " << s.L_padded << std::endl; + std::cout << "DEBUG: L_effective: " << s.L_effective << std::endl; + std::cout << "DEBUG: sliding_l_begin: " << s.sliding_l_begin << std::endl; + ) + + for (int i = 0; i < s.L_effective; i++){ + DEBUG_BLOCK(2, + this->pli_embedding->forward(input_ids[i], this->prefill_ctx->pli_embed_row(i)); + this->embedding->forward(input_ids[i], this->prefill_ctx->residual_row(i)); + continue; + ) + if ((ids[i] == image_token_id) || (ids[i] == audio_token_id)){ + this->pli_embedding->forward(0, this->prefill_ctx->pli_embed_row(i)); + continue; + } + else{ + this->embedding->forward(ids[i], this->prefill_ctx->residual_row(i)); + this->pli_embedding->forward(ids[i], this->prefill_ctx->pli_embed_row(i)); + } + } + + if (payload != nullptr){ + std::vector image_embedding; + std::vector audio_embedding; + + gemma4e_multi_modal_payload_t* multi_modal_payload_ptr = (gemma4e_multi_modal_payload_t*)payload; + + if(is_audio && multi_modal_payload_ptr->audio_payload.num_audios > 0){ + audio_embedding = this->gemma4e_audio_encoder->encode(&multi_modal_payload_ptr->audio_payload); + } + if(is_vlm && multi_modal_payload_ptr->image_payload.num_images > 0){ + image_embedding = this->gemma4e_image_encoder->encode(&multi_modal_payload_ptr->image_payload); + } + + bf16* image_token_ptr = image_embedding.data(); + bf16* audio_token_ptr = audio_embedding.data(); + + for (int i = 0; i < ids.size(); i++) { + if (is_vlm && (ids[i] == image_token_id)){ + // copy image embedding to pli_embed_buffer + memcpy(this->prefill_ctx->residual_row(i).data(), image_token_ptr, D * sizeof(bf16)); + image_token_ptr += D; + } + else if (is_audio && (ids[i] == audio_token_id)){ + // copy audio embedding to pli_embed_buffer + memcpy(this->prefill_ctx->residual_row(i).data(), audio_token_ptr, D * sizeof(bf16)); + audio_token_ptr += D; + } + } + } + + DEBUG_BLOCK(2, + header_print("info", "Input embedding for the first " + std::to_string(s.L_effective) + " tokens:"); + buffer embedding_valid = buffer(bufs.pli_embed_buffer.data() + s.L_offset * PLI_D * num_hidden_layers, s.L_effective * PLI_D * num_hidden_layers); + utils::print_matrix(embedding_valid, PLI_D * num_hidden_layers); + ) + + DEBUG_BLOCK(2, + buffer embedding_ref; + reference->load_weights(embedding_ref, "input_embeds"); + buffer embedding_valid = buffer(bufs.residual_buffer.data() + s.L_offset * D, s.L_effective * D); + buffer embedding_valid_ref = buffer(embedding_ref.data() + s.L_offset * D, s.L_effective * D); + print_error_metrics(get_error_metrics(embedding_valid, embedding_valid_ref), "Input Embedding Error: "); + ) + + // fold the token embeddings into the per layer input stream, once for the batch + buffer pli_input_norm(this->rope_rms_weights[0].data() + desc.get_pli_norm_offset(layer_types[0]), PLI_D); + this->prefill_ctx->pli->pre_pass(s, this->pli_down_weights, pli_input_norm); + + for (uint32_t layer_idx = 0; layer_idx < num_hidden_layers; layer_idx++){ + gemma4e_layer_type_t type = layer_types[layer_idx]; + + DEBUG_BLOCK(2, + if (layer_idx > 0){ + buffer residual_overide; + reference->load_weights(residual_overide, "layer_" + std::to_string(layer_idx - 1)); + memcpy(bufs.residual_buffer.data() + s.L_offset * D, residual_overide.data() + s.L_offset * D, s.L_effective * D * sizeof(bf16)); + } + ) + + this->prefill_ctx->forward( + layer_idx, type, s, + this->proj_weights[layer_idx], + this->rms_weights[layer_idx], + this->rope_rms_weights[layer_idx], + this->pli_gate_up_weights[layer_idx], + this->kv_caches[is_skip_layer(type) + ? (is_swa_layer(type) ? last_swa_kv_cache_layer_idx : last_global_kv_cache_layer_idx) + : (int)layer_idx], + this->layer_scale[layer_idx], + reference.get() + ); + } + + DEBUG_BLOCK(2, + reference.reset(); + header_print_r("info", "DEBUG EXIT"); + exit(0); + ) + buffer predict = this->prefill_ctx->residual_row(s.L_effective - 1); + this->lm_head_weights.sync_to_device(); + get_logits(predict); + this->set_context_length(this->current_context_length + s.L_effective); + + this->pre_load_run = FLM_OVERRIDE(layer_pre_load, this->layer_pre_load.create_run()); + this->pre_load_run.start(); + is_preload_launched = true; + return logits_valid; +} + +buffer gemma4e_npu::Impl::get_k_cache(int layer_idx, int idx){ + this->kv_caches[layer_idx].sync_from_device(); + buffer k_cache(DK); + bf16* k_cache_ptr = this->kv_caches[layer_idx].data(); + uint32_t offset = idx * DK; + memcpy(static_cast(k_cache.data()), static_cast(k_cache_ptr + offset), DK * sizeof(bf16)); + + return k_cache; +} + +buffer gemma4e_npu::Impl::get_v_cache(int layer_idx, int idx){ + this->kv_caches[layer_idx].sync_from_device(); + buffer v_cache(DV); + bf16* v_cache_ptr = this->kv_caches[layer_idx].data(); + uint32_t offset = idx * DV; + memcpy(static_cast(v_cache.data()), static_cast(v_cache_ptr + offset), DV * sizeof(bf16)); + return v_cache; +} + +buffer gemma4e_npu::Impl::get_logits(buffer& predict){ + memcpy(this->x.data(), predict.data(), D * sizeof(bf16)); + this->x.sync_to_device(); + this->lm_head_run.start(); + this->lm_head_run.wait(); + this->logits.sync_from_device(); + return this->logits_valid; +} + +void gemma4e_npu::Impl::clear_context(){ + this->set_context_length(0); + for (int i = 0; i < non_skip_layers; i++){ + this->kv_caches[i].sync_from_device(); + memset(this->kv_caches[i].data(), 0, this->kv_caches[i].size() * sizeof(bf16)); + this->kv_caches[i].sync_to_device(); + } +} + +int gemma4e_npu::Impl::get_current_context_length(){ + return this->current_context_length; +} + +void gemma4e_npu::Impl::update_max_length(uint32_t MAX_L){ + if (MAX_L <= this->MAX_L){ + header_print("FLM", "New length is shorter than the current length, no need to update!"); + return; // no need to update + } + this->MAX_L = MAX_L; + size_t kv_cache_size = MAX_L * (DK + DV); + this->sequence->set_max_length(MAX_L); + + this->prefill_ctx->set_max_length(MAX_L); + + for (uint32_t i = 0; i < non_skip_layers; i++){ + if (!is_swa_layer(layer_types[i])){ + this->kv_caches[i] = buffer(kv_cache_size); + memset(this->kv_caches[i].data(), 0, kv_cache_size * sizeof(bf16)); + this->kv_caches[i].sync_to_device(); + } + else{ + memset(this->kv_caches[i].data(), 0, this->kv_caches[i].size() * sizeof(bf16)); + this->kv_caches[i].sync_to_device(); + } + } + is_checkpoint_valid = false; // force reload checkpoint to clear kv cache on device + this->set_context_length(0); // clear +} + +void gemma4e_npu::Impl::load_weights(Q4NX& q4nx){ + + this->embedding->init_weights(q4nx, "model.embed_tokens"); + this->pli_embedding->init_weights(q4nx, "model.per_layer_token_embd"); + + // Shared by every layer, so it is read once and handed to each of them. + buffer per_layer_norm; + q4nx.load_weights(per_layer_norm, "model.per_layer_proj_norm.weight"); + + for (int layer_idx = 0; layer_idx < num_hidden_layers; layer_idx++){ + this->desc.load_layer_weights(layer_idx, q4nx, + this->proj_weights[layer_idx], + this->rms_weights[layer_idx], + this->rope_rms_weights[layer_idx], + this->pli_gate_up_weights[layer_idx], + per_layer_norm, + this->layer_scale[layer_idx]); + this->pli_gate_up_weights[layer_idx].sync_to_device(); + this->proj_weights[layer_idx].sync_to_device(); + this->rms_weights[layer_idx].sync_to_device(); + this->rope_rms_weights[layer_idx].sync_to_device(); + FLM_OVERRIDE(layer_weights_loaded, (void)0, &this->desc, layer_idx); + DEBUG_BLOCK(1, + header_print("info", "Finished loading weights for layer " + std::to_string(layer_idx)); + ) + } + + // model.norm.weight lives in the second D-sized slot of x, behind the hidden state. + this->desc.load_head_weights(q4nx, this->lm_head_weights, this->x.data() + D, this->pli_down_weights); + + this->pli_down_weights.sync_to_device(); + this->lm_head_weights.sync_to_device(); + FLM_OVERRIDE(head_weights_loaded, (void)0, &this->desc); + + DEBUG_BLOCK(1, + header_print("info", "Finished loading all layer weights"); + ) + + DEBUG_BLOCK(1, + header_print("info", "Finished updating sequence buffer offsets based on loaded weights"); + ) + + this->clear_context(); + + DEBUG_BLOCK(1, + header_print("info", "Finished clearing context after loading weights"); + ) + + // read qkv weights + if(is_vlm){ + DEBUG_BLOCK(1, + std::cout << "[DBG] gemma4e_npu::Impl::load_weights: entering VLM branch, path=" + << config.get("vision_model_weight", "") << std::endl; + ) + { + SafeTensors vision_weights(config.get("vision_model_weight", "")); + DEBUG_BLOCK(1, + std::cout << "[DBG] gemma4e_npu::Impl::load_weights: SafeTensors constructed" << std::endl; + ) + this->gemma4e_image_encoder->init_weights(vision_weights); + DEBUG_BLOCK(1, + std::cout << "[DBG] gemma4e_npu::Impl::load_weights: init_weights returned" << std::endl; + ) + } + DEBUG_BLOCK(1, + std::cout << "[DBG] gemma4e_npu::Impl::load_weights: SafeTensors out of scope" << std::endl; + ) + } + if(is_audio){ + SafeTensors audio_weights(config.get("audio_model_weight", "")); + this->gemma4e_audio_encoder->init_weights(audio_weights); + } +} + +int gemma4e_npu::Impl::checkpoint(){ + if (!is_checkpoint_valid){ + _allocate_checkpoint_buffers(); + } // use lazy allocation + header_print_r("FLM", "Creating checkpoint at context length " + std::to_string(this->current_context_length)); + checkpoint_context_length = this->current_context_length; + for (int i = 0; i < non_skip_layers; i++){ + if (!is_global_layer_idx(i)){ + this->kv_caches[i].sync_from_device(); + memcpy(this->kv_checkpoint[i].data(), this->kv_caches[i].data(), this->kv_caches[i].size() * sizeof(bf16)); + this->kv_caches[i].sync_to_device(); + } + } + is_checkpoint_valid = true; + return this->checkpoint_context_length; +} + +int gemma4e_npu::Impl::restore(){ + if (!is_checkpoint_valid){ + header_print("FLM", "No valid checkpoint found, cannot restore context!"); + return -1; + } + header_print_r("FLM", "Restoring checkpoint at context length " + std::to_string(checkpoint_context_length)); + this->current_context_length = checkpoint_context_length; + set_context_length(this->current_context_length); + for (int i = 0; i < non_skip_layers; i++){ + this->kv_caches[i].sync_from_device(); + if (!is_global_layer_idx(i)){ + this->kv_caches[i].sync_from_device(); + memcpy(this->kv_caches[i].data(), this->kv_checkpoint[i].data(), this->kv_caches[i].size() * sizeof(bf16)); + } + else { + int offset = (size_t)current_context_length * (DK); + // first half is K, second half is V, and they are stored contiguously in kv_cache + memset(this->kv_caches[i].data() + offset, 0, (this->kv_caches[i].size() / 2 - offset) * sizeof(bf16)); // zero out the part that is not restored for global layers, since we only restore the sliding part for global layers + memset(this->kv_caches[i].data() + this->kv_caches[i].size() / 2 + offset, 0, (this->kv_caches[i].size() / 2 - offset) * sizeof(bf16)); // zero out the second half for swa + } + this->kv_caches[i].sync_to_device(); + } + return this->current_context_length; +} + +void gemma4e_npu::Impl::_allocate_checkpoint_buffers(){ + // allocate kv cache checkpoint for sliding layers + this->kv_checkpoint.clear(); + this->kv_checkpoint.resize(non_skip_layers); + for (int i = 0; i < non_skip_layers; i++){ + if(!is_global_layer_idx(i)){ + this->kv_checkpoint[i] = buffer(this->kv_caches[i].size()); + } + } + is_checkpoint_valid = false; +} + +///@brief destructor of qwen3vl_npu +gemma4e_npu::Impl::~Impl(){ + // a preload run may still be in flight; it must complete before the + // xrt objects it references are torn down + if (is_preload_launched){ + try{ + this->pre_load_run.wait(); + } + catch (const std::exception& e){ + header_print("FLM", std::string("Failed to wait for the pending preload run: ") + e.what()); + } + is_preload_launched = false; + } +} +/* +qwen3vl_npu::Impl::~Impl(){ + for (int i = 0; i < config.get("num_hidden_layers"); i++){ + this->kv_caches[i].free(); + this->proj_weights[i].free(); + this->rms_weights[i].free(); + } + this->kv_caches.clear(); + this->rms_weights.clear(); + this->proj_weights.clear(); + this->rope_weights.release(); + this->dequantized_qkv_weights.release(); + this->dequantized_o_weights.release(); + this->x.release(); + this->lm_head.reset(); +} +*/ +// =============================================== +// Externals for qwen3vl_npu +// =============================================== +gemma4e_npu::gemma4e_npu(LM_Config config, npu_xclbin_manager *npu_instance, int MAX_L){ + + this->load_vision_preprocess_parameters(config); + this->load_audio_preprocess_parameters(config); + this->_impl = new Impl(config, npu_instance, this, MAX_L); +} + +buffer gemma4e_npu::forward(int ids){ + return this->_impl->forward(ids); +} + +buffer gemma4e_npu::prefill(std::vector& ids, void* payload){ + try { + buffer result = this->_impl->prefill(ids, payload); + return result; + } + catch (const std::runtime_error& e) { + header_print("FLM", e.what()); + throw; + } +} + +void gemma4e_npu::set_context_length(int L){ + std::cout << "Setting context length is not supported for qwen3vl_npu!" << std::endl; +} + +void gemma4e_npu::load_weights(Q4NX& q4nx){ + this->_impl->load_weights(q4nx); +} + +void gemma4e_npu::update_max_length(uint32_t MAX_L){ + this->_impl->update_max_length(MAX_L); +} + +void gemma4e_npu::clear_context(){ + this->_impl->clear_context(); +} + +int gemma4e_npu::checkpoint(){ + return this->_impl->checkpoint(); +} + +int gemma4e_npu::restore(){ + return this->_impl->restore(); +} + +buffer gemma4e_npu::get_k_cache(int layer_idx, int idx){ + return this->_impl->get_k_cache(layer_idx, idx); +} + +buffer gemma4e_npu::get_v_cache(int layer_idx, int idx){ + return this->_impl->get_v_cache(layer_idx, idx); +} + +int gemma4e_npu::get_current_context_length(){ + return this->_impl->get_current_context_length(); +} + +gemma4e_npu::~gemma4e_npu(){ + delete this->_impl; +} diff --git a/src/detail/gemma4e_npu/gemma4e_npu_def.hpp b/src/detail/gemma4e_npu/gemma4e_npu_def.hpp new file mode 100644 index 000000000..9c29309aa --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_npu_def.hpp @@ -0,0 +1,558 @@ +/// \file gemma4e_npu_def.hpp +/// \brief Gemma4e text backbone config + weight descriptors. +/// \author FastFlowLM Team +/// \note Owns everything that describes the model on disk and in device memory: +/// the config-derived dimensions, the per-layer weight descriptors (and +/// therefore the buffer layout the kernels read), and the weight loading +/// itself. Mirrors llama_desc / gemma4_12b_desc. +/// \note Gemma4e is a hybrid model in two independent ways: +/// - attention is sliding-window (SWA) except every `global_layer_period`-th +/// layer, which is full (global) attention; the two differ in head_dim. +/// - the last `num_kv_shared_layers` layers ("skip" layers) reuse the KV +/// cache of an earlier layer, so they carry no k/v projection and (when +/// `use_double_wide_mlp`) a twice-as-wide MLP. +/// That gives four distinct buffer layouts, indexed by gemma4e_layer_type_t: +/// [0]=swa, [1]=global, [2]=swa_skip, [3]=global_skip. +#pragma once +#include +#include +#include +#include "lm_config.hpp" +#include "weight_desc.hpp" +#include "tensor_utils/q4_npu_eXpress.hpp" +#include "models/gemma4e/flm/aie2p/gemma4e_npu.hpp" + +/// minimum tail padding (in bf16 elements) of the rope/rms buffer, kept in sync +/// with gemma4e_npu_sequence.hpp. +#ifndef MIN_BF16_PAD +#define MIN_BF16_PAD 32 +#endif + +/// \brief Descriptors of every weight of one decoder layer. +/// \note The order in which these are handed to the weight_container in +/// gemma4e_desc::_build_layer() IS the on-device layout, and both the +/// sequence generator (gemma4e_npu_sequence::gen_layer_seq) and the +/// dequant sequences depend on it: q/k/v contiguous, then o, then the +/// interleaved up/gate block, then down, then the three bf16 PLI +/// projections. +struct gemma4e_layer_weight_def { + // ---- projections, quantized, live in the per-layer proj buffer ---- + weight_desc_t attn_q; + weight_desc_t attn_k; //!< absent on skip layers, see has_kv + weight_desc_t attn_v; //!< absent on skip layers, see has_kv + weight_desc_t attn_output; + weight_desc_t ffn_up_gate; //!< up and gate interleaved, see UP_GATE_OUTDIM_REORDER + weight_desc_t ffn_down; + + // ---- per-layer-input projections, bf16, same buffer, after the quantized block ---- + weight_desc_t pli_down_proj; + weight_desc_t pli_gate_proj; + weight_desc_t pli_up_proj; + + // ---- rms norms, bf16, live in the per-layer rms buffer ---- + weight_desc_t input_layer_norm; + weight_desc_t post_attention_norm; + weight_desc_t pre_feedforward_norm; + weight_desc_t post_feedforward_norm; + + // ---- rope/rms buffer, bf16: [cos|sin, q_norm, k_norm, pli_embed, pli_norm, post_pli_norm, scale, pad] ---- + weight_desc_t rope_cos_sin; //!< not a checkpoint tensor: host-computed per position + weight_desc_t attn_q_norm; + weight_desc_t attn_k_norm; + weight_desc_t pli_embed; //!< scratch slot the PLI path writes into, not a checkpoint tensor + weight_desc_t pli_norm; //!< model.per_layer_proj_norm.weight, shared by every layer + weight_desc_t post_pli_norm; + weight_desc_t layer_output_scale; + + // ---- fused / aliased views, never registered in a container ---- + // These carry no storage of their own: they name a span that already exists so + // that the code moving it around can be handed a descriptor instead of an + // offset + two dimensions. Their offsets are copied from the real descriptor + // they alias once it has been placed. + weight_desc_t attn_qkv; //!< q(|k|v) as one block, what _move_weights ships + weight_desc_t ffn_up; //!< the up half of ffn_up_gate, what the dequant kernel picks out + weight_desc_t ffn_gate; //!< the gate half of ffn_up_gate + + bool has_kv = true; //!< false on skip layers (they share an earlier layer's cache) + + /// \brief The projection weights, in device-buffer order. + std::vector proj_all() { + std::vector out = {&attn_q}; + if (has_kv) { out.push_back(&attn_k); out.push_back(&attn_v); } + out.push_back(&attn_output); + out.push_back(&ffn_up_gate); + out.push_back(&ffn_down); + out.push_back(&pli_down_proj); + out.push_back(&pli_gate_proj); + out.push_back(&pli_up_proj); + return out; + } + + /// \brief The four layer norms, in device-buffer order. + std::vector rms_all() { + return {&input_layer_norm, &post_attention_norm, &pre_feedforward_norm, &post_feedforward_norm}; + } + + /// \brief The rope/rms buffer entries, in device-buffer order. + std::vector rope_all() { + return {&rope_cos_sin, &attn_q_norm, &attn_k_norm, + &pli_embed, &pli_norm, &post_pli_norm, &layer_output_scale}; + } + + /// \brief Every descriptor of the layer that owns storage. + std::vector all() { + std::vector out = proj_all(); + for (weight_desc_t* w : rms_all()) out.push_back(w); + for (weight_desc_t* w : rope_all()) out.push_back(w); + return out; + } +}; + +/// \brief Gemma4e text backbone description (config.json + model.q4nx). +struct gemma4e_desc { + /// up/gate are interleaved in slices of this many output rows + static constexpr int UP_GATE_OUTDIM_REORDER = 512; + /// \note Switching the body to another 4-bit dtype is a one-line change here: + /// every size, offset and DMA length below is derived from it. + static constexpr flm_dtype_t PROJ_DTYPE = flm_q41; + static constexpr flm_dtype_t LM_HEAD_DTYPE = flm_q41; + static constexpr flm_dtype_t PLI_DTYPE = flm_bf16; + + /// gemma4e moves weights in 32x256 hardware blocks (unlike the 16x256 of gemma4-12b). + static constexpr int QXNX_M = QXNX_ROW_BLOCK_SIZE; + static constexpr int QXNX_K = QXNX_COL_BLOCK_SIZE; + + /// \brief Bytes of one 32x256 hardware block of a quantized dtype. + static size_t block_bytes(flm_dtype_t dtype) { + return get_quantization_byte_size((size_t)QXNX_M * QXNX_K, dtype); + } + + // ---- dimensions (config.json) ---- + uint32_t num_hidden_layers = 0; + uint32_t D = 0; //!< hidden_size + uint32_t INTERMEDIATE_SIZE = 0; + uint32_t PLI_D = 0; //!< hidden_size_per_layer_input + uint32_t num_attention_heads = 0; + uint32_t num_kv_heads = 0; + uint32_t DH = 0; //!< global_head_dim + uint32_t DQ = 0; + uint32_t DK = 0; + uint32_t DV = 0; + uint32_t SWA_DH = 0; //!< head_dim + uint32_t SWA_DQ = 0; + uint32_t SWA_DK = 0; + uint32_t SWA_DV = 0; + uint32_t SLIDING_LENGTH = 0; + uint32_t vocab_size = 0; + uint32_t vocab_size_padded = 0; + uint32_t num_kv_shared_layers = 0; + uint32_t non_skip_layers = 0; + uint32_t global_layer_period = 0; + bool enable_double_wide_mlp = false; + f32 final_logit_softcapping = 0.0f; + + std::vector layer_types; + + // ---- weight descriptors ---- + /// one representative layout per layer kind, indexed by gemma4e_layer_type_t + gemma4e_layer_weight_def layer_defs[4]; + weight_desc_t final_norm; + weight_desc_t lm_head; + + // ---- aggregated buffer byte sizes, indexed by gemma4e_layer_type_t ---- + size_t proj_weights_byte_size[4] = {0, 0, 0, 0}; + size_t rms_weights_byte_size[4] = {0, 0, 0, 0}; + size_t rope_rms_byte_size[4] = {0, 0, 0, 0}; //!< payload only, without MIN_BF16_PAD + + gemma4e_desc() {} + gemma4e_desc(LM_Config config) { build(config); } + + /// \brief Whether layer `layer_idx` uses full (global) attention. + inline bool is_global_layer_idx(int layer_idx) const { + return ((layer_idx + 1) % (int)global_layer_period) == 0; + } + + /// \brief The representative descriptor of a layer kind. + inline gemma4e_layer_weight_def& weight_desc(gemma4e_layer_type_t type) { + return layer_defs[int(type)]; + } + /// \brief The representative descriptor of a layer, by index. + inline gemma4e_layer_weight_def& weight_desc_of(int layer_idx) { + return layer_defs[int(layer_types[layer_idx])]; + } + + // ---- layer-kind-specific dimensions ---- + inline uint32_t get_DH(gemma4e_layer_type_t t) const { return is_global_layer(t) ? DH : SWA_DH; } + inline uint32_t get_DQ(gemma4e_layer_type_t t) const { return is_global_layer(t) ? DQ : SWA_DQ; } + inline uint32_t get_DK(gemma4e_layer_type_t t) const { return is_global_layer(t) ? DK : SWA_DK; } + inline uint32_t get_DV(gemma4e_layer_type_t t) const { return is_global_layer(t) ? DV : SWA_DV; } + /// \brief MLP width of a layer kind: skip layers are twice as wide when enabled. + inline uint32_t get_intermediate_size(gemma4e_layer_type_t t) const { + return (is_skip_layer(t) && enable_double_wide_mlp) ? INTERMEDIATE_SIZE * 2 : INTERMEDIATE_SIZE; + } + /// \brief MLP width of the widest layer kind, which sizes the dequant buffers. + inline uint32_t get_max_intermediate_size() const { + return enable_double_wide_mlp ? INTERMEDIATE_SIZE * 2 : INTERMEDIATE_SIZE; + } + + // ---- buffer sizes ---- + inline size_t get_proj_weights_byte_size(gemma4e_layer_type_t t) const { return proj_weights_byte_size[int(t)]; } + /// \brief rms buffer size, in bf16 elements (input | post attn | pre ffn | post ffn). + inline size_t get_rms_elems(gemma4e_layer_type_t t) const { return rms_weights_byte_size[int(t)] / sizeof(bf16); } + /// \brief rope/rms buffer size, in bf16 elements, padding included. + /// \note The layer scale is the last entry, and MIN_BF16_PAD elements are kept + /// behind it so that the kernel never reads past the buffer. + inline size_t get_rope_rms_elems(gemma4e_layer_type_t t) { + return rope_rms_byte_size[int(t)] / sizeof(bf16) + MIN_BF16_PAD; + } + /// \brief kv cache size of one layer, in bf16 elements. + /// \note Sliding layers only ever keep `sliding_window` positions. + inline size_t get_kv_cache_size(gemma4e_layer_type_t t, uint32_t MAX_L) const { + return is_global_layer(t) ? (size_t)MAX_L * (DK + DV) + : (size_t)SLIDING_LENGTH * (SWA_DK + SWA_DV); + } + inline size_t get_lm_head_w_size() { return lm_head.get_size(); } + + // ---- rope/rms buffer offsets, in bf16 elements ---- + inline uint32_t rope_elem_offset(weight_desc_t& w) const { return (uint32_t)(w.offset / sizeof(bf16)); } + inline uint32_t get_q_norm_offset(gemma4e_layer_type_t t) { return rope_elem_offset(weight_desc(t).attn_q_norm); } + inline uint32_t get_k_norm_offset(gemma4e_layer_type_t t) { return rope_elem_offset(weight_desc(t).attn_k_norm); } + inline uint32_t get_pli_embed_offset(gemma4e_layer_type_t t) { return rope_elem_offset(weight_desc(t).pli_embed); } + inline uint32_t get_pli_norm_offset(gemma4e_layer_type_t t) { return rope_elem_offset(weight_desc(t).pli_norm); } + inline uint32_t get_post_pli_norm_offset(gemma4e_layer_type_t t) { return rope_elem_offset(weight_desc(t).post_pli_norm); } + inline uint32_t get_layer_scale_offset(gemma4e_layer_type_t t){ return rope_elem_offset(weight_desc(t).layer_output_scale); } + + /// \brief Parse config.json and lay out every weight. + inline void build(LM_Config& config) { + const nlohmann::json& jc = config._json_config; + + uint32_t head_dim = 0, global_head_dim = 0; + JSON_GET(num_hidden_layers, jc, "num_hidden_layers", 0, uint32_t); + JSON_GET(D, jc, "hidden_size", 0, uint32_t); + JSON_GET(INTERMEDIATE_SIZE, jc, "intermediate_size", 0, uint32_t); + JSON_GET(num_attention_heads, jc, "num_attention_heads", 0, uint32_t); + JSON_GET(num_kv_heads, jc, "num_key_value_heads", 0, uint32_t); + JSON_GET(head_dim, jc, "head_dim", 0, uint32_t); + JSON_GET(global_head_dim, jc, "global_head_dim", 0, uint32_t); + JSON_GET(SLIDING_LENGTH, jc, "sliding_window", 0, uint32_t); + JSON_GET(PLI_D, jc, "hidden_size_per_layer_input", 0, uint32_t); + JSON_GET(vocab_size, jc, "vocab_size", 0, uint32_t); + JSON_GET(num_kv_shared_layers, jc, "num_kv_shared_layers", 0, uint32_t); + JSON_GET(final_logit_softcapping, jc, "final_logit_softcapping", 0.0f, f32); + JSON_GET(enable_double_wide_mlp, jc, "use_double_wide_mlp", false, bool); + + DH = global_head_dim; + DQ = DH * num_attention_heads; + DK = DH * num_kv_heads; + DV = DK; + SWA_DH = head_dim; + SWA_DQ = SWA_DH * num_attention_heads; + SWA_DK = SWA_DH * num_kv_heads; + SWA_DV = SWA_DK; + vocab_size_padded = (vocab_size + 1024 - 1) / 1024 * 1024; + non_skip_layers = num_hidden_layers - num_kv_shared_layers; + + // The global-layer period is not in config.json; it follows from the model size. + if (D == 1536) global_layer_period = 5; // E2B + else if (D == 2560) global_layer_period = 6; // E4B + else throw std::runtime_error("gemma4e_desc: unsupported hidden size " + std::to_string(D)); + + layer_types.resize(num_hidden_layers); + for (uint32_t i = 0; i < num_hidden_layers; i++) { + int type = 0; + if (is_global_layer_idx(i)) type |= GEMMA4E_IS_GLOBAL_MASK; + if (i >= non_skip_layers) type |= GEMMA4E_IS_SKIP_MASK; + layer_types[i] = static_cast(type); + } + + DEBUG_BLOCK(1, + std::cout << "================ gemma4e_desc ================" << std::endl; + std::cout << std::left << std::setw(28) << "num_hidden_layers" << " = " << num_hidden_layers << std::endl; + std::cout << std::left << std::setw(28) << "D (hidden_size)" << " = " << D << std::endl; + std::cout << std::left << std::setw(28) << "PLI_D" << " = " << PLI_D << std::endl; + std::cout << std::left << std::setw(28) << "INTERMEDIATE_SIZE" << " = " << INTERMEDIATE_SIZE << std::endl; + std::cout << std::left << std::setw(28) << "DH / DQ / DK / DV" << " = " << DH << " / " << DQ << " / " << DK << " / " << DV << std::endl; + std::cout << std::left << std::setw(28) << "SWA DH / DQ / DK / DV" << " = " << SWA_DH << " / " << SWA_DQ << " / " << SWA_DK << " / " << SWA_DV << std::endl; + std::cout << std::left << std::setw(28) << "SLIDING_LENGTH" << " = " << SLIDING_LENGTH << std::endl; + std::cout << std::left << std::setw(28) << "global_layer_period" << " = " << global_layer_period << std::endl; + std::cout << std::left << std::setw(28) << "num_kv_shared_layers" << " = " << num_kv_shared_layers << std::endl; + std::cout << std::left << std::setw(28) << "use_double_wide_mlp" << " = " << enable_double_wide_mlp << std::endl; + std::cout << std::left << std::setw(28) << "final_logit_softcapping" << " = " << final_logit_softcapping << std::endl; + std::cout << std::left << std::setw(28) << "vocab_size" << " = " << vocab_size << " (" << vocab_size_padded << " padded)" << std::endl; + std::cout << "=============================================" << std::endl; + ) + + final_norm = weight_desc_t(flm_bf16, {(int64_t)D}, "model.norm.weight"); + lm_head = weight_desc_t(LM_HEAD_DTYPE, {(int64_t)D, (int64_t)vocab_size_padded}, "lm_head.weight"); + // Neither shares a buffer with anything, so they never go through a + // weight_container and would never get `added` set. Mark them registered + // here; their offset inside their own buffer is 0 by construction. + final_norm.indp(); + lm_head.indp(); + + for (int t = 0; t < 4; t++) _build_layer(static_cast(t)); + } + + /// \brief Copy a quantized weight into the device buffer in NPU block order. + /// \param dst destination inside the per-layer projection buffer + /// \param src the weight as stored in the q4nx file + /// \param col the input dimension of the weight + /// \param dtype the quantized dtype of the weight + /// \param vertical_blocks how many 32-row bands are interleaved + void reorder_cpy(u8* dst, buffer& src, const int col, flm_dtype_t dtype, + const int vertical_blocks = 2) + { + assert(is_quantize(dtype)); + const size_t a_block_size = block_bytes(dtype); + const int blocks_per_row = col / QXNX_K; + const int rows = src.size() / a_block_size / blocks_per_row; + + u8* dst_ptr = dst; + std::vector src_ptr(vertical_blocks); + for (int i = 0; i < vertical_blocks; i++) + src_ptr[i] = src.data() + i * a_block_size * blocks_per_row; + + for (int r = 0; r < rows; r += vertical_blocks) + { + for (int c = 0; c < blocks_per_row; c++) + { + for (int i = 0; i < vertical_blocks; i++) + { + memcpy(dst_ptr, src_ptr[i], a_block_size); + dst_ptr += a_block_size; + src_ptr[i] += a_block_size; + } + } + for (int i = 0; i < vertical_blocks; i++) + { + src_ptr[i] += (vertical_blocks - 1) * a_block_size * blocks_per_row; + if (src_ptr[i] + a_block_size * blocks_per_row > src.end()) + src_ptr[i] = src.data(); // useless padding + } + } + } + + /// \brief Load one decoder layer's weights from the checkpoint into its device buffers. + /// \param layer_idx which layer to load + /// \param q4nx the opened checkpoint + /// \param proj_buffer the layer's projection buffer (quantized block + bf16 PLI block) + /// \param rms_buffer the layer's four layer norms + /// \param rope_rms_buffer the layer's rope/norm/scale buffer + /// \param pli_gate_up_buffer the prefill-shaped copies of the PLI gate/up projections + /// \param per_layer_norm model.per_layer_proj_norm.weight, shared by every layer + /// \param layer_scale_out receives the layer output scale as a float + /// \note Every destination is addressed through its descriptor's offset, so the + /// physical layout lives in _build_layer() alone. + void load_layer_weights(int layer_idx, Q4NX& q4nx, + buffer& proj_buffer, + buffer& rms_buffer, + buffer& rope_rms_buffer, + buffer& pli_gate_up_buffer, + buffer& per_layer_norm, + float& layer_scale_out) + { + const gemma4e_layer_type_t type = layer_types[layer_idx]; + gemma4e_layer_weight_def& L = weight_desc(type); + u8* base = proj_buffer.data(); + + // ---- q (k, v) ---- + { + buffer w; + q4nx.load_weights(w, L.attn_q.format_name(layer_idx)); + reorder_cpy(base + L.attn_q.offset, w, D, L.attn_q.dtype); + L.attn_q.load(); + } + if (L.has_kv) { + for (weight_desc_t* desc : {&L.attn_k, &L.attn_v}) { + buffer w; + q4nx.load_weights(w, desc->format_name(layer_idx)); + reorder_cpy(base + desc->offset, w, D, desc->dtype); + desc->load(); + } + } + // ---- o ---- + { + buffer w; + q4nx.load_weights(w, L.attn_output.format_name(layer_idx)); + reorder_cpy(base + L.attn_output.offset, w, get_DQ(type), L.attn_output.dtype); + L.attn_output.load(); + } + // ---- up / gate, interleaved in slices of UP_GATE_OUTDIM_REORDER output rows ---- + { + buffer w_up, w_gate; + q4nx.load_weights(w_up, L.ffn_up.format_name(layer_idx)); + q4nx.load_weights(w_gate, L.ffn_gate.format_name(layer_idx)); + + const size_t chunk_size = get_quantization_byte_size((size_t)UP_GATE_OUTDIM_REORDER * D, + L.ffn_up_gate.dtype); + const size_t half_size = L.ffn_up_gate.get_size() / 2; + const size_t phases = half_size / chunk_size; + assert(phases * chunk_size == half_size && "up/gate must split into whole slices"); + + tensor_2d up_tensor(w_up, chunk_size, 0); + tensor_2d gate_tensor(w_gate, chunk_size, 0); + u8* w_ptr = base + L.ffn_up_gate.offset; + for (size_t i = 0; i < phases; i++) { + LOG_VERBOSE(1, "Copying up and gate weights to buffer, phase " << i + 1 << "/" << phases); + reorder_cpy(w_ptr, up_tensor[i], D, L.ffn_up_gate.dtype); + w_ptr += chunk_size; + reorder_cpy(w_ptr, gate_tensor[i], D, L.ffn_up_gate.dtype); + w_ptr += chunk_size; + } + L.ffn_up_gate.load(); + L.ffn_up.load(); + L.ffn_gate.load(); + } + // ---- down ---- + { + buffer w; + q4nx.load_weights(w, L.ffn_down.format_name(layer_idx)); + reorder_cpy(base + L.ffn_down.offset, w, get_intermediate_size(type), L.ffn_down.dtype); + L.ffn_down.load(); + } + // ---- per-layer-input projections, plain bf16 ---- + for (weight_desc_t* desc : {&L.pli_down_proj, &L.pli_gate_proj, &L.pli_up_proj}) { + buffer w; + q4nx.load_weights(w, desc->format_name(layer_idx)); + memcpy(base + desc->offset, w.data(), desc->get_size()); + desc->load(); + } + + // ---- the four layer norms ---- + for (weight_desc_t* desc : L.rms_all()) { + buffer w; + q4nx.load_weights(w, desc->format_name(layer_idx)); + memcpy(desc->locate_myself(rms_buffer).data(), w.data(), desc->get_size()); + desc->load(); + } + + // ---- rope/rms buffer ---- + // rope_cos_sin and pli_embed are scratch slots the runtime fills per position, + // so only the four checkpoint-backed entries are read here. + for (weight_desc_t* desc : {&L.attn_q_norm, &L.attn_k_norm, &L.post_pli_norm}) { + buffer w; + q4nx.load_weights(w, desc->format_name(layer_idx)); + memcpy(desc->locate_myself(rope_rms_buffer).data(), w.data(), desc->get_size()); + desc->load(); + } + memcpy(L.pli_norm.locate_myself(rope_rms_buffer).data(), + per_layer_norm.data(), L.pli_norm.get_size()); + L.pli_norm.load(); + { + buffer w; + q4nx.load_weights(w, L.layer_output_scale.format_name(layer_idx)); + layer_scale_out = (float)w[0]; + // The kernel broadcasts the scale, but it still reads a full vector's + // worth behind it, so zero the padding that follows. + bf16* scale_ptr = L.layer_output_scale.locate_myself(rope_rms_buffer).data(); + memset(scale_ptr, 0, (1 + MIN_BF16_PAD) * sizeof(bf16)); + scale_ptr[0] = w[0]; + L.layer_output_scale.load(); + } + + // ---- prefill-shaped copies of the PLI gate/up projections ---- + { + buffer w_gate, w_up; + q4nx.load_weights(w_gate, "model.layers." + std::to_string(layer_idx) + ".inp_gate.weight_prefill"); + q4nx.load_weights(w_up, "model.layers." + std::to_string(layer_idx) + ".per_layer_projection.weight_prefill"); + memcpy(pli_gate_up_buffer.data(), w_gate.data(), (size_t)PLI_D * D * sizeof(bf16)); + memcpy(pli_gate_up_buffer.data() + (size_t)PLI_D * D, w_up.data(), (size_t)PLI_D * D * sizeof(bf16)); + } + } + + /// \brief Load the weights that do not belong to any single layer. + /// \param q4nx the opened checkpoint + /// \param lm_head_buffer destination of the (reordered) lm head + /// \param final_norm_dst destination of model.norm.weight + /// \param pli_down_buffer destination of the prefill-shaped per-layer down projection + void load_head_weights(Q4NX& q4nx, + buffer& lm_head_buffer, + bf16* final_norm_dst, + buffer& pli_down_buffer) + { + { + buffer w; + q4nx.load_weights(w, final_norm.name); + memcpy(final_norm_dst, w.data(), final_norm.get_size()); + final_norm.load(); + } + { + buffer w; + q4nx.load_weights(w, lm_head.name); + // The lm head spans four columns of cores rather than two. + reorder_cpy(lm_head_buffer.data(), w, D, lm_head.dtype, 4); + lm_head.load(); + } + { + buffer w; + q4nx.load_weights(w, "model.per_layer_model_proj.weight_prefill"); + memcpy(pli_down_buffer.data(), w.data(), + (size_t)num_hidden_layers * D * PLI_D * sizeof(bf16)); + } + } + +private: + /// \brief Lay out one layer kind's three buffers and name every tensor. + /// \param type the layer kind whose layout is being built + /// \note The add_weight() call order is the physical layout; see the note on + /// gemma4e_layer_weight_def. + inline void _build_layer(gemma4e_layer_type_t type) { + gemma4e_layer_weight_def& L = layer_defs[int(type)]; + + const int64_t _DH = get_DH(type); + const int64_t _DQ = get_DQ(type); + const int64_t _DK = get_DK(type); + const int64_t _DV = get_DV(type); + const int64_t _IS = get_intermediate_size(type); + const int64_t d = D; + const int64_t pli = PLI_D; + + L.has_kv = !is_skip_layer(type); + + L.attn_q = weight_desc_t(PROJ_DTYPE, {d, _DQ}, "model.layers.%d.self_attn.q_proj.weight"); + L.attn_k = weight_desc_t(PROJ_DTYPE, {d, _DK}, "model.layers.%d.self_attn.k_proj.weight"); + L.attn_v = weight_desc_t(PROJ_DTYPE, {d, _DV}, "model.layers.%d.self_attn.v_proj.weight"); + L.attn_output = weight_desc_t(PROJ_DTYPE, {_DQ, d}, "model.layers.%d.self_attn.o_proj.weight"); + L.ffn_up_gate = weight_desc_t(PROJ_DTYPE, {d, 2 * _IS}, "model.layers.%d.mlp.{up,gate}_proj.weight"); + L.ffn_down = weight_desc_t(PROJ_DTYPE, {_IS, d}, "model.layers.%d.mlp.down_proj.weight"); + L.pli_down_proj = weight_desc_t(PLI_DTYPE, {pli, d}, "model.per_layer_model_proj.weight_layer%d"); + L.pli_gate_proj = weight_desc_t(PLI_DTYPE, {pli, d}, "model.layers.%d.inp_gate.weight"); + L.pli_up_proj = weight_desc_t(PLI_DTYPE, {pli, d}, "model.layers.%d.per_layer_projection.weight"); + + L.input_layer_norm = weight_desc_t(flm_bf16, {d}, "model.layers.%d.input_layernorm.weight"); + L.post_attention_norm = weight_desc_t(flm_bf16, {d}, "model.layers.%d.post_attention_layernorm.weight"); + L.pre_feedforward_norm = weight_desc_t(flm_bf16, {d}, "model.layers.%d.pre_feedforward_layernorm.weight"); + L.post_feedforward_norm = weight_desc_t(flm_bf16, {d}, "model.layers.%d.post_feedforward_layernorm.weight"); + + L.rope_cos_sin = weight_desc_t(flm_bf16, {_DH}, ""); + L.attn_q_norm = weight_desc_t(flm_bf16, {_DH}, "model.layers.%d.self_attn.q_norm.weight"); + L.attn_k_norm = weight_desc_t(flm_bf16, {_DH}, "model.layers.%d.self_attn.k_norm.weight"); + L.pli_embed = weight_desc_t(flm_bf16, {pli}, ""); + L.pli_norm = weight_desc_t(flm_bf16, {pli}, "model.per_layer_proj_norm.weight"); + L.post_pli_norm = weight_desc_t(flm_bf16, {d}, "model.layers.%d.post_layernorm.weight"); + L.layer_output_scale = weight_desc_t(flm_bf16, {1}, "model.layers.%d.layer_output_scale.weight"); + + weight_container proj, rms, rope; + for (weight_desc_t* w : L.proj_all()) proj.add_weight(*w); + for (weight_desc_t* w : L.rms_all()) rms.add_weight(*w); + for (weight_desc_t* w : L.rope_all()) rope.add_weight(*w); + + proj_weights_byte_size[int(type)] = proj.get_size(); + rms_weights_byte_size[int(type)] = rms.get_size(); + rope_rms_byte_size[int(type)] = rope.get_size(); + + // Aliased views over regions that are already placed above. q/k/v are + // contiguous and are shipped and dequantized as one block; up/gate share + // one interleaved region that the dequant kernel splits by output mode. + L.attn_qkv = weight_desc_t(PROJ_DTYPE, {d, L.has_kv ? (_DQ + _DK + _DV) : _DQ}, + ""); + L.attn_qkv.offset = L.attn_q.offset; + L.attn_qkv.indp(); + + L.ffn_up = weight_desc_t(PROJ_DTYPE, {d, _IS}, "model.layers.%d.mlp.up_proj.weight"); + L.ffn_gate = weight_desc_t(PROJ_DTYPE, {d, _IS}, "model.layers.%d.mlp.gate_proj.weight"); + L.ffn_up.offset = L.ffn_gate.offset = L.ffn_up_gate.offset; + L.ffn_up.indp(); + L.ffn_gate.indp(); + } +}; diff --git a/src/detail/gemma4e_npu/gemma4e_npu_detail.hpp b/src/detail/gemma4e_npu/gemma4e_npu_detail.hpp new file mode 100644 index 000000000..fafe92632 --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_npu_detail.hpp @@ -0,0 +1,181 @@ +#pragma once +#include "models/gemma4e/flm/aie2p/gemma4e_npu.hpp" +#include "gemma4e_npu_sequence.hpp" +#include "gemma4e_image.hpp" +#include "modules/gemm.hpp" +#include "reorder_cpy.hpp" +#include "avx512_util.hpp" +#include "gemma4e_audio.hpp" +#include "embedding_q8_0.hpp" +#include "gemma4e_prefill.hpp" + +struct gemma4e_npu::Impl{ + /// @brief constexprs + + // some model specific template variables + static constexpr int boi_token_id = 255999; // begin of image token id + static constexpr int image_token_id = 258880; // image token id + static constexpr int eoi_token_id = 258882; // end of image token id + + static constexpr int boa_token_id = 256000; // begin of audio token id + static constexpr int audio_token_id = 258881; // audio token id + static constexpr int eoa_token_id = 258883; // end of audio token id + /// @brief constexprs + //TODO: FIXME: + + static constexpr int LC = 16; + + size_t V_offset; + + int vocab_size; + int vocab_size_padded; + bool is_preload_launched; + + LM_Config config; + npu_xclbin_manager *npu; + + npu_app_manager* layer_app_manager; + npu_app_manager* lm_head_app_manager; + + npu_app global_layer; + npu_app swa_layer; + npu_app global_skip_layer; + npu_app swa_skip_layer; + npu_app layer_pre_load; + flm_rt::run pre_load_run; + + // per layer input + npu_app per_layer_input_down_proj; + npu_app per_layer_input_up_proj; + + npu_app lm_head; + + std::unique_ptr sequence; + std::unique_ptr gemma4e_image_encoder; + std::unique_ptr gemma4e_audio_encoder; + + gemma4e_npu* parent_ptr; + + int D; + int DH; + int DQ; + int DK; + int DV; + int SWA_DH; + int SWA_DQ; + int SWA_DK; + int SWA_DV; + + int MAX_L; + int HIDDEN_SIZE; + int INTERMEDIATE_SIZE; + int PLI_D; + int SLIDING_LENGTH; + + int sliding_layer_interval; + int num_hidden_layers; + int num_kv_shared_layers; + int non_skip_layers; + int global_layer_period; + int last_swa_kv_cache_layer_idx; + int last_global_kv_cache_layer_idx; + float final_logit_softcapping; + + bool enable_double_wide_mlp; + + /// \brief Config-derived dimensions and every weight descriptor of the model. + /// \note Owns the buffer layout: the dequant sequences, the decode sequence + /// generator and load_weights all address weights through it, so the + /// quantization type is switchable from gemma4e_desc::PROJ_DTYPE alone. + gemma4e_desc desc; + std::vector layer_types; + + // size + std::vector> rms_weights; + std::vector> rope_rms_weights; + std::vector> proj_weights; + std::vector> kv_caches; + std::vector> pli_gate_up_weights; + + std::vector> kv_checkpoint; + + buffer pli_down_weights; + buffer layer_scale; + /// Everything the prefill path needs: its own xclbins, sequences and buffers. + std::unique_ptr prefill_ctx; + + buffer x; + std::unique_ptr embedding; + std::unique_ptr pli_embedding; + + buffer logits; + buffer logits_valid; + buffer lm_head_weights; + + flm_rt::runlist layers_run; + flm_rt::run lm_head_run; + flm_rt::runlist dequant_all; + int current_context_length; + + int checkpoint_context_length; + + /// @brief initialize the qwen_npu + /// @param config + /// @param npu_instance + Impl(LM_Config config, npu_xclbin_manager *npu_instance, gemma4e_npu* parent_ptr, int MAX_L = 4096); + ~Impl(); // waits for any in-flight preload run before tearing down + + /// @brief forward the qwen_npu + buffer forward(int ids); + buffer prefill(std::vector& ids, void* payload = nullptr); + buffer _prefill_with_mv(std::vector& ids, void* payload = nullptr); + buffer _prefill_with_mm(std::vector& ids, void* payload = nullptr); + void _load_attn_layer_weights(Q4NX& q4nx, int layer_idx); + void _load_linear_layer_weights(Q4NX& q4nx, int layer_idx); + + void set_context_length(int L); + void load_weights(Q4NX& q4nx); + void update_max_length(uint32_t MAX_L); + void clear_context(); + + bool is_checkpoint_valid; + bool is_vlm; + bool is_audio; + void _allocate_checkpoint_buffers(); + int checkpoint(); + int restore(); + + buffer get_k_cache(int layer_idx, int idx); + buffer get_v_cache(int layer_idx, int idx); + buffer get_logits(buffer& x); + + int get_current_context_length(); + + void _set_rope_rms_weights(int idx); + + inline bool is_global_layer_idx(int layer_idx) {return (layer_idx % global_layer_period) == (global_layer_period - 1); } + inline void _process_embedding(int idx){ + buffer embedding_output = this->embedding->forward(idx); + DEBUG_BLOCK(2, + header_print("debug", "Embedding output for token idx " + std::to_string(idx)); + utils::print_matrix(embedding_output, D); + ) + memcpy(this->x.data(), embedding_output.data(), D * sizeof(bf16)); + memcpy(this->x.data() + D * 2, embedding_output.data(), D * sizeof(bf16)); // copy to the second half for swa + this->x.sync_to_device(); + buffer pli_embedding_output = this->pli_embedding->forward(idx); + + DEBUG_BLOCK(2, + header_print("debug", "PLI Embedding output for token idx " + std::to_string(idx)); + utils::print_matrix(pli_embedding_output, PLI_D); + ) + + bf16* p_pli = pli_embedding_output.data(); + for (int i = 0; i < num_hidden_layers; i++) { + memcpy(this->rope_rms_weights[i].data() + desc.get_pli_embed_offset(layer_types[i]), p_pli, PLI_D * sizeof(bf16)); + p_pli += PLI_D; + this->rope_rms_weights[i].sync_to_device(); + } + } + +}; diff --git a/src/detail/gemma4e_npu/gemma4e_npu_sequence.cpp b/src/detail/gemma4e_npu/gemma4e_npu_sequence.cpp new file mode 100644 index 000000000..b6da27dd4 --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_npu_sequence.cpp @@ -0,0 +1,1279 @@ +#include "gemma4e_npu_sequence.hpp" + +gemma4e_npu_sequence::gemma4e_npu_sequence(gemma4e_seq_gen_parameters_t params, uint32_t MAX_L){ + D = params.D; + DH = params.DH; + DQ = params.DQ; + DK = params.DK; + DV = params.DV; + SWA_DH = params.SWA_DH; + SWA_DQ = params.SWA_DQ; + SWA_DK = params.SWA_DK; + SWA_DV = params.SWA_DV; + PLI_D = params.PLI_D; + num_attn_heads = params.NUM_ATTENTION_HEADS; + num_kv_heads = params.NUM_KEY_VALUE_HEADS; + num_kv_per_round = num_kv_heads; + INTERMEDIATE_SIZE = params.INTERMEDIATE_SIZE; + VOCAB_SIZE = params.VOCAB_SIZE_PADDED; + SLIDING_LENGTH = params.SLIDING_WINDOW_SIZE; + enable_double_wide_mlp = params.enable_double_wide_mlp; + + this->MAX_L = MAX_L; + + if (INTERMEDIATE_SIZE == 6144){ // e2b + rtp_addresses = e2b_rtp_addresses; + } + else if (INTERMEDIATE_SIZE == 10240){ + rtp_addresses = e4b_rtp_addresses; + } + else{ + std::cerr << "Unsupported intermediate size: " << INTERMEDIATE_SIZE << std::endl; + throw std::runtime_error("SEQ Unsupported intermediate size"); + } + + DEBUG_BLOCK(1, + header_print_g("info", "Sequence Gen Params: "); + std::cout << "\tD: " << D << std::endl; + std::cout << "\tDH: " << DH << std::endl; + std::cout << "\tDQ: " << DQ << std::endl; + std::cout << "\tDK: " << DK << std::endl; + std::cout << "\tDV: " << DV << std::endl; + std::cout << "\tSWA_DH: " << SWA_DH << std::endl; + std::cout << "\tSWA_DQ: " << SWA_DQ << std::endl; + std::cout << "\tSWA_DK: " << SWA_DK << std::endl; + std::cout << "\tSWA_DV: " << SWA_DV << std::endl; + std::cout << "\tPLI_D: " << PLI_D << std::endl; + std::cout << "\tMAX_L: " << MAX_L << std::endl; + std::cout << "\tINTERMEDIATE_SIZE: " << INTERMEDIATE_SIZE << std::endl; + std::cout << "\tNUM_ATTENTION_HEADS: " << num_attn_heads << std::endl; + std::cout << "\tNUM_KEY_VALUE_HEADS: " << num_kv_heads << std::endl; + std::cout << "\tSLIDING_WINDOW_SIZE: " << SLIDING_LENGTH << std::endl; + std::cout << "\tVOCAB_SIZE_PADDED: " << VOCAB_SIZE << std::endl; + ) + DEBUG_BLOCK(1, + header_print_g("info", "RTP ADDRESS BOOK: "); + std::cout << "\tl_qk_address: " << rtp_addresses.l_qk_address << std::endl; + std::cout << "\tl_kv_address: " << rtp_addresses.l_kv_address << std::endl; + std::cout << "\tswa_l_qk_address: " << rtp_addresses.swa_l_qk_address << std::endl; + std::cout << "\tswa_l_kv_address: " << rtp_addresses.swa_l_kv_address << std::endl; + std::cout << "\tproj_swa_address: " << rtp_addresses.proj_swa_address << std::endl; + std::cout << "\tproj_skip_address: " << rtp_addresses.proj_skip_address << std::endl; + std::cout << "\trms_swa_address: " << rtp_addresses.rms_swa_address << std::endl; + std::cout << "\trms_skip_address: " << rtp_addresses.rms_skip_address << std::endl; + std::cout << "\trope_skip_kv_address: " << rtp_addresses.rope_skip_kv_address << std::endl; + std::cout << "\tswa_rope_skip_kv_address: " << rtp_addresses.swa_rope_skip_kv_address << std::endl; + ) +} + +void gemma4e_npu_sequence::set_max_length(const uint32_t MAX_L){ + this->MAX_L = MAX_L; +} + +void gemma4e_npu_sequence::gen_layer_seq(npu_sequence* seq, const uint32_t L, gemma4e_layer_type_t layer_type){ + constexpr size_t CT_lock_address_base = 0x000001F000; + LAYER_SPECIFIC_DIM(layer_type) + DEBUG_BLOCK(2, + header_print_g("info", "Generating sequence for layer type " + std::to_string(layer_type) + " with L = " + std::to_string(L)); + header_print_g("info", "Layer specific dimensions: "); + std::cout << "\t_ D: " << D << std::endl; + std::cout << "\t_ DH: " << _DH << std::endl; + std::cout << "\t_ DQ: " << _DQ << std::endl; + std::cout << "\t_ DK: " << _DK << std::endl; + std::cout << "\t_ DV: " << _DV << std::endl; + std::cout << "\t_ INTERMEDIATE_SIZE: " << _INTERMEDIATE_SIZE << std::endl; + std::cout << "\t_ QKV_OFFSET: " << weight_elem_offset(layer_weights(layer_type).attn_qkv) << std::endl; + std::cout << "\t_ O_OFFSET: " << weight_elem_offset(layer_weights(layer_type).attn_output) << std::endl; + std::cout << "\t_ UP_OFFSET: " << weight_elem_offset(layer_weights(layer_type).ffn_up_gate) << std::endl; + std::cout << "\t_ DOWN_OFFSET: " << weight_elem_offset(layer_weights(layer_type).ffn_down) << std::endl; + std::cout << "\t_ PLI_DOWN_OFFSET: " << weight_elem_offset(layer_weights(layer_type).pli_down_proj) << std::endl; + std::cout << "\t_ PLI_GATE_OFFSET: " << weight_elem_offset(layer_weights(layer_type).pli_gate_proj) << std::endl; + std::cout << "\t_ PLI_UP_OFFSET: " << weight_elem_offset(layer_weights(layer_type).pli_up_proj) << std::endl; + std::cout << "\t_ IS_SWA: " << is_swa_layer(layer_type) << std::endl; + std::cout << "\t_ IS_SKIP: " << is_skip_layer(layer_type) << std::endl; + ) + int proj_rtp_lock_id = 6; + int rms_rtp_lock_id = 6; + int glu_rtp_lock_id = 6; + seq->clear_cmds(); + int L_local = L; + if (is_swa_layer(layer_type)) { + if (L > SLIDING_LENGTH){ + L_local = SLIDING_LENGTH; + } + } + + for (int i = 0; i < 16; i++){ + seq->rtp_write(proj_tiles[i], rtp_addresses.proj_swa_address, is_swa_layer(layer_type) ? 1 : 0); + seq->rtp_write(proj_tiles[i], rtp_addresses.proj_skip_address, is_skip_layer(layer_type) ? 1 : 0); + seq->rtp_write(proj_tiles[i], CT_lock_address_base + 16 * proj_rtp_lock_id, 1); // set lock to 1 + } + seq->rtp_write(attn_qk_tile, rtp_addresses.l_qk_address, L_local); + seq->rtp_write(attn_kv_tile, rtp_addresses.l_kv_address, L_local); + seq->rtp_write(swa_attn_qk_tile, rtp_addresses.swa_l_qk_address, L_local); + seq->rtp_write(swa_attn_kv_tile, rtp_addresses.swa_l_kv_address, L_local); + seq->rtp_write(rms_tile, rtp_addresses.rms_swa_address, is_swa_layer(layer_type) ? 1 : 0); + seq->rtp_write(rms_tile, rtp_addresses.rms_skip_address, is_skip_layer(layer_type) ? 1 : 0); + seq->rtp_write(rope_ct, rtp_addresses.rope_skip_kv_address, is_skip_layer(layer_type) ? 1 : 0); + seq->rtp_write(swa_rope_ct, rtp_addresses.swa_rope_skip_kv_address, is_skip_layer(layer_type) ? 1 : 0); + seq->rtp_write(glu_tile, rtp_addresses.glu_skip_address, (is_skip_layer(layer_type) && enable_double_wide_mlp) ? 1 : 0); + seq->rtp_write(rms_tile, CT_lock_address_base + 16 * rms_rtp_lock_id, 1); // set lock to 1 + seq->rtp_write(glu_tile, CT_lock_address_base + 16 * glu_rtp_lock_id, 1); // set lock to 1 + + _send_hidden_states(seq); + _send_rms_weights(seq); + _send_rope_rms_weights(seq, layer_type); + + // receive y + + seq->npu_dma_memcpy_nd( + sizeof(bf16), + x_arg_id, + S2MM, + xr_tile, + bd_15, + it_channel_0, + {0, 0, 0, 0}, + {1, 1, 1, (uint32_t)D}, + {0, 0, 0, 1}, + -1, 0, true + ); + gemma4e_layer_weight_def& W = layer_weights(layer_type); + // attn_qkv already spans q alone on skip layers, which carry no k/v projection. + _move_weights(seq, W.attn_qkv); + if (!is_skip_layer(layer_type)){ + _receive_kv_cache(seq, L, layer_type); + } + + _gen_pli_path_seq(seq, layer_type); + _move_kv_cache(seq, L, layer_type); + _move_weights(seq, W.attn_output); + _move_weights(seq, W.ffn_up_gate); + _move_weights(seq, W.ffn_down); + + seq->npu_dma_memcpy_nd( + sizeof(bf16), + proj_arg_id, + MM2S, + IT5, + bd_14, + it_channel_0, + {0, 0, 0, weight_elem_offset(layer_weights(layer_type).pli_gate_proj)}, + {1, 1, 1, (uint32_t)(PLI_D * D)}, + {0, 0, 0, 1}, + -1, 0, true, aggressive_cache + ); + + seq->npu_dma_wait( + IT5, + MM2S, + it_channel_0 + ); + + seq->npu_dma_memcpy_nd( + sizeof(bf16), + proj_arg_id, + MM2S, + IT5, + bd_15, + it_channel_1, + {0, 0, 0, weight_elem_offset(layer_weights(layer_type).pli_up_proj)}, + {1, 1, 1, (uint32_t)(PLI_D * D)}, + {0, 0, 0, 1}, + -1, 0, true, aggressive_cache + ); + + seq->npu_dma_wait( + IT5, + MM2S, + it_channel_1 + ); + // wait for receiving y + seq->npu_dma_wait( + xr_tile, + S2MM, + it_channel_0 + ); + seq->cmds2seq(); +} + +void gemma4e_npu_sequence::gen_lm_head_seq(npu_sequence* seq, float final_scale){ + static constexpr int y_arg_id = 0; + static constexpr int w_arg_id = 1; + static constexpr int x_arg_id = 2; + static constexpr int M_PER_ROUND = 32 * 32; + static constexpr int M = 32; + static constexpr int M_PER_COL = M * 4; + static constexpr int COLS = 8; + static npu_tiles ITs[] = {IT0, IT1, IT2, IT3, IT4, IT5, IT6, IT7}; + const uint32_t final_scale_address = this->rtp_addresses.lm_head_final_tune_address; + assert(VOCAB_SIZE % M_PER_ROUND == 0); + int rounds = VOCAB_SIZE / M_PER_ROUND; + size_t TOTAL_W_SIZE = size_t(D) * size_t(VOCAB_SIZE) * 5 / 8 / 2; + size_t WEIGHTS_PER_IT = TOTAL_W_SIZE / std::size(ITs); + size_t W_PER_COL = (M_PER_COL * D * 5 / 8 / 2); + size_t W_PER_ROUND = W_PER_COL * COLS; + + size_t w_offset = WEIGHTS_PER_IT; + size_t y_offset = VOCAB_SIZE / std::size(ITs); + uint32_t scale_int_view = *((uint32_t*)(&final_scale)); + seq->clear_cmds(); + for (int row = 0; row < 4; row++){ + for (int col = 0; col < 8; col++){ + npu_tiles tile = get_tile(row + 2, col); + seq->rtp_write(tile, final_scale_address, scale_int_view); + } + } + seq->npu_dma_memcpy_nd( + sizeof(bf16), + x_arg_id, + MM2S, + ITs[0], + npu_bd_id(bd_0), + it_channel_0, + {0, 0, 0, 0}, + {1, 1, 1, (uint32_t)D * 2}, + {0, 0, 0, 1}, + -1, 0, false, aggressive_cache + ); + + for (int r = 0; r < rounds; r++) { + int bd_offset = (r % 2) * 8; + npu_bd_id bd_y = npu_bd_id(bd_1 + bd_offset); + npu_bd_id bd_w = npu_bd_id(bd_2 + bd_offset); + for (size_t col = 0; col < std::size(ITs); col++){ + size_t w_col_offset = r * W_PER_ROUND + col * W_PER_COL; + uint32_t y_offset = r * M_PER_ROUND + col * M_PER_COL; + seq->npu_dma_memcpy_nd( + sizeof(bf16), + y_arg_id, + S2MM, + ITs[col], + bd_y, + it_channel_0, + {0, 0, 0, (uint32_t)(y_offset)}, + {1, 1, 1, (uint32_t)M_PER_COL}, + {0, 0, 0, 1}, + -1, 0, true, aggressive_cache + ); + seq->npu_dma_memcpy_nd( + sizeof(bf16), + w_arg_id, + MM2S, + ITs[col], + bd_w, + it_channel_1, + {0, 0, 0, (uint32_t)(w_col_offset)}, + {1, 1, 1, (uint32_t)W_PER_COL}, + {0, 0, 0, 1}, + -1, 0, false, aggressive_cache + ); + if (r > 0){ + seq->npu_dma_wait( + ITs[col], + S2MM, + it_channel_0 + ); + } + } + } + for (size_t col = 0; col < std::size(ITs); col++){ + seq->npu_dma_wait( + ITs[col], + S2MM, + it_channel_0 + ); + } + seq->cmds2seq(); +} + +void gemma4e_npu_sequence::_send_hidden_states(npu_sequence* seq){ + // send x + seq->npu_dma_memcpy_nd( + sizeof(bf16), + x_arg_id, + MM2S, + xr_tile, + bd_0, + it_channel_0, + {0, 0, 0, 0}, + {1, 1, 1, (uint32_t)D}, + {0, 0, 0, 1}, + 0, 0, true, aggressive_cache + ); + seq->npu_dma_wait( + xr_tile, + MM2S, + it_channel_0 + ); +} + +void gemma4e_npu_sequence::_send_rms_weights(npu_sequence* seq){ + seq->npu_dma_memcpy_nd( + sizeof(bf16), + rms_arg_id, + MM2S, + xr_tile, + bd_1, + it_channel_1, + {0, 0, 0, 0}, + {1, 1, 1, (uint32_t)D * 4}, + {0, 0, 0, 1}, + 0, 0, true, aggressive_cache + ); + + seq->npu_dma_wait( + xr_tile, + MM2S, + it_channel_1 + ); +} + +void gemma4e_npu_sequence::_send_rope_rms_weights(npu_sequence* seq, gemma4e_layer_type_t layer_type){ + LAYER_SPECIFIC_DIM(layer_type) + if (is_swa_layer(layer_type)){ + seq->npu_dma_memcpy_nd( + sizeof(bf16), + rope_rms_arg_id, + MM2S, + xr_tile, + bd_4, + it_channel_1, + {0, 0, 0, 0}, + {1, 1, 1, (uint32_t)(_DH * 3)}, + {0, 0, 0, 1}, + 1, 0, true, no_cache + ); + seq->npu_dma_wait( + xr_tile, + MM2S, + it_channel_1 + ); + } + else { + seq->npu_dma_memcpy_nd( + sizeof(bf16), + rope_rms_arg_id, + MM2S, + xr_tile, + bd_4, + it_channel_0, + {0, 0, 0, 0}, + {1, 1, 1, (uint32_t)(_DH * 3)}, + {0, 0, 0, 1}, + 1, 0, true, no_cache + ); + seq->npu_dma_wait( + xr_tile, + MM2S, + it_channel_0 + ); + } +} + +void gemma4e_npu_sequence::_receive_kv_cache(npu_sequence* seq, const int L, gemma4e_layer_type_t layer_type){ + LAYER_SPECIFIC_DIM(layer_type) + DEBUG_BLOCK(2, + header_print_g("info", "Moving KV cache for layer type " + std::to_string(layer_type) + " with L = " + std::to_string(L)); + std::cout << "\t_ L: " << L << std::endl; + std::cout << "\t_ _DK: " << _DK << std::endl; + std::cout << "\t_ _DV: " << _DV << std::endl; + std::cout << "\t_ MAX_L: " << MAX_L << std::endl; + ) + int L_local = L - 1; + if (is_swa_layer(layer_type)) { + L_local = L_local % SLIDING_LENGTH; + } + uint32_t kv_cache_size; + if (is_swa_layer(layer_type)){ + kv_cache_size = SLIDING_LENGTH * (_DK + _DV); + } + else { + kv_cache_size = (_DK + _DV) * MAX_L; + } + uint32_t v_offset = kv_cache_size / 2; + uint32_t L_offset = L_local * _DK; + + npu_it_channel receiving_channel = is_swa_layer(layer_type)? it_channel_1 : it_channel_0; + + seq->npu_dma_memcpy_nd( + sizeof(bf16), + kv_cache_arg_id, + S2MM, + attn_tile, + bd_0, + receiving_channel, + {0, 0, 0, (uint32_t)L_offset}, + {1, 1, 1, (uint32_t)_DK}, + {0, 0, 0, 1}, + -1, 0, true + ); + + seq->npu_dma_wait( + attn_tile, + S2MM, + receiving_channel + ); + seq->npu_dma_memcpy_nd( + sizeof(bf16), + kv_cache_arg_id, + S2MM, + attn_tile, + bd_1, + receiving_channel, + {0, 0, 0, (uint32_t)(L_offset + v_offset)}, + {1, 1, 1, (uint32_t)_DV}, + {0, 0, 0, 1}, + -1, 0, true + ); + seq->npu_dma_wait( + attn_tile, + S2MM, + receiving_channel + ); +} + +void gemma4e_npu_sequence::_move_kv_cache(npu_sequence* seq, const size_t L, gemma4e_layer_type_t layer_type){ + LAYER_SPECIFIC_DIM(layer_type) + DEBUG_BLOCK(2, + header_print_g("info", "Moving KV cache for layer type " + std::to_string(layer_type) + " with L = " + std::to_string(L)); + std::cout << "\t_ L: " << L << std::endl; + std::cout << "\t_ _DK: " << _DK << std::endl; + std::cout << "\t_ _DV: " << _DV << std::endl; + std::cout << "\t_ MAX_L: " << MAX_L << std::endl; + ) + uint32_t kv_cache_size; + if (is_swa_layer(layer_type)){ + kv_cache_size = SLIDING_LENGTH * (_DK + _DV); + } + else { + kv_cache_size = (_DK + _DV) * MAX_L; + } + uint32_t v_offset = kv_cache_size / 2; + int pkt_id = is_swa_layer(layer_type) ? 13 : 12; + + if (is_swa_layer(layer_type)) { + if (L % SLIDING_LENGTH == 0){ // corner, from 0 to end + // move the oldest block to the new block + uint32_t data2move = SLIDING_LENGTH * _DK; + seq->npu_dma_memcpy_nd( + sizeof(bf16), + kv_cache_arg_id, + MM2S, + attn_tile, + bd_8, + it_channel_0, + {0, 0, 0, 0}, + {1, 1, 1, data2move}, + {0, 0, 0, 1}, + pkt_id, 0, false, aggressive_cache + ); + + seq->npu_dma_memcpy_nd( + sizeof(bf16), + kv_cache_arg_id, + MM2S, + attn_tile, + bd_9, + it_channel_1, + {0, 0, 0, (uint32_t)(v_offset)}, + {1, 1, 1, data2move}, + {0, 0, 0, 1}, + pkt_id, 0, false, aggressive_cache + ); + return ; + } + else if (L > SLIDING_LENGTH) { // dual phase + uint32_t L_begin = L % SLIDING_LENGTH; + uint32_t L_phase_1 = SLIDING_LENGTH - L_begin; + uint32_t L_phase_2 = L_begin; + uint32_t offset = L_begin * _DK; + uint32_t data2move_1 = L_phase_1 * _DK; + uint32_t data2move_2 = L_phase_2 * _DK; + + seq->npu_dma_memcpy_nd( + sizeof(bf16), + kv_cache_arg_id, + MM2S, + attn_tile, + bd_8, + it_channel_0, + {0, 0, 0, offset}, + {1, 1, 1, data2move_1}, + {0, 0, 0, 1}, + pkt_id, 0, false, aggressive_cache + ); + + seq->npu_dma_memcpy_nd( + sizeof(bf16), + kv_cache_arg_id, + MM2S, + attn_tile, + bd_9, + it_channel_1, + {0, 0, 0, (uint32_t)(v_offset + offset)}, + {1, 1, 1, data2move_1}, + {0, 0, 0, 1}, + pkt_id, 0, false, aggressive_cache + ); + // phase 2 + seq->npu_dma_memcpy_nd( + sizeof(bf16), + kv_cache_arg_id, + MM2S, + attn_tile, + bd_10, + it_channel_0, + {0, 0, 0, 0}, + {1, 1, 1, data2move_2}, + {0, 0, 0, 1}, + pkt_id, 0, false, aggressive_cache + ); + + seq->npu_dma_memcpy_nd( + sizeof(bf16), + kv_cache_arg_id, + MM2S, + attn_tile, + bd_11, + it_channel_1, + {0, 0, 0, (uint32_t)(v_offset)}, + {1, 1, 1, data2move_2}, + {0, 0, 0, 1}, + pkt_id, 0, false, aggressive_cache + ); + return; + } + } + DEBUG_BLOCK(2, + std::cout << "Single phase move for layer type (Fall back path) " << layer_type << std::endl; + ) + // fallback path, single phase, from zero to L + const int L_padded = (L + L_CHUNK - 1) / L_CHUNK * L_CHUNK; + const uint32_t data2move = L_padded * _DK; + seq->npu_dma_memcpy_nd( + sizeof(bf16), + kv_cache_arg_id, + MM2S, + attn_tile, + bd_8, + it_channel_0, + {0, 0, 0, 0}, + {1, 1, 1, data2move}, + {0, 0, 0, 1}, + pkt_id, 0, true, aggressive_cache + ); + + seq->npu_dma_memcpy_nd( + sizeof(bf16), + kv_cache_arg_id, + MM2S, + attn_tile, + bd_9, + it_channel_1, + {0, 0, 0, (uint32_t)(v_offset)}, + {1, 1, 1, data2move}, + {0, 0, 0, 1}, + pkt_id, 0, true, aggressive_cache + ); + + seq->npu_dma_wait( + attn_tile, + MM2S, + it_channel_0 + ); + seq->npu_dma_wait( + attn_tile, + MM2S, + it_channel_1 + ); +} + +/// \brief Stream one quantized weight from DDR into the mvm cores. +/// \param seq the sequence to append the DMAs to +/// \param weight the weight to move; its shape is {input dim, output dim} and its +/// offset is a byte offset into the layer's projection buffer +/// \note The proj port addresses DDR in bf16 elements, so both the descriptor's +/// byte offset and the block size are halved here. Nothing about this +/// function is tied to a particular 4-bit layout any more: switch +/// gemma4e_desc::PROJ_DTYPE and the block size follows. +void gemma4e_npu_sequence::_move_weights(npu_sequence* seq, weight_desc_t& weight){ + assert(weight.added && "weight must be placed in its buffer before it can be moved"); + assert(is_quantize(weight.dtype) && "_move_weights moves quantized projections only"); + const int m = QXNX_ROW_BLOCK_SIZE; + const int k = QXNX_COL_BLOCK_SIZE; + const uint32_t columns = 4; + const size_t Din = (size_t)weight.shape[0]; + const size_t Dout = (size_t)weight.shape[1]; + const uint32_t a_block_size = (uint32_t)(get_quantization_byte_size((size_t)m * k, weight.dtype) / sizeof(bf16)); + const uint32_t w_offset = weight_elem_offset(weight); + const uint32_t blocks_per_row = Din / k; + const uint32_t cores = columns * 4; + assert(Dout / m / cores > 0); + assert(Dout % (m * cores) == 0); + for (size_t round = 0; round < Dout / m / cores; round++){ + uint32_t bd_offset = (round % 2) * 8; + for (uint32_t col = 0; col < columns; col++){ + seq->npu_dma_memcpy_nd( + sizeof(bf16), + proj_arg_id, + MM2S, + mvm_tiles[col], + npu_bd_id(bd_1 + bd_offset), + it_channel_0, + {0, 0, 0, (uint32_t)((round * cores + col * 4) * a_block_size * blocks_per_row + (uint32_t)w_offset)}, + {1, 1, 1, 2 * blocks_per_row * a_block_size}, + {0, 0, 0, 1}, + -1, 0, true, aggressive_cache + ); + seq->npu_dma_memcpy_nd( + sizeof(bf16), + proj_arg_id, + MM2S, + mvm_tiles[col], + npu_bd_id(bd_2 + bd_offset), + it_channel_1, + {0, 0, 0, (uint32_t)((round * cores + col * 4 + 2) * a_block_size * blocks_per_row + (uint32_t)w_offset)}, + {1, 1, 1, 2 * blocks_per_row * a_block_size}, + {0, 0, 0, 1}, + -1, 0, true, aggressive_cache + ); + } + if (round > 0){ + for (uint32_t col = 0; col < columns; col++){ + seq->npu_dma_wait( + mvm_tiles[col], + MM2S, + it_channel_0 + ); + seq->npu_dma_wait( + mvm_tiles[col], + MM2S, + it_channel_1 + ); + } + } + } + for (uint32_t col = 0; col < columns; col++){ + seq->npu_dma_wait( + mvm_tiles[col], + MM2S, + it_channel_0 + ); + seq->npu_dma_wait( + mvm_tiles[col], + MM2S, + it_channel_1 + ); + } +} + +void gemma4e_npu_sequence::_gen_pli_path_seq(npu_sequence* seq, gemma4e_layer_type_t layer_type){ + LAYER_SPECIFIC_DIM(layer_type) + DEBUG_BLOCK(2, + header_print_g("info", "Generating PLI path sequence for layer type " + std::to_string(layer_type)); + std::cout << "\t_ D: " << D << std::endl; + std::cout << "\t_ DH: " << _DH << std::endl; + std::cout << "\t_ PLI_D: " << PLI_D << std::endl; + std::cout << "\t_ pli_down_proj_offset: " << weight_elem_offset(layer_weights(layer_type).pli_down_proj) << std::endl; + std::cout << "\t_ pli_gate_proj_offset: " << weight_elem_offset(layer_weights(layer_type).pli_gate_proj) << std::endl; + std::cout << "\t_ pli_up_proj_offset: " << weight_elem_offset(layer_weights(layer_type).pli_up_proj) << std::endl; + ) + // move per layer input weights + // send x, x, w + seq->npu_dma_memcpy_nd( + sizeof(bf16), + rope_rms_arg_id, + MM2S, + IT4, + bd_11, + it_channel_0, + {0, 0, 0, (uint32_t)(_DH * 3)}, + {1, 1, 1, (uint32_t)(PLI_D * 2 + D + MIN_BF16_PAD)}, + {0, 0, 0, 1}, + -1, 0, true, no_cache + ); + seq->npu_dma_wait( + IT4, + MM2S, + it_channel_0 + ); + + seq->npu_dma_memcpy_nd( + sizeof(bf16), + x_arg_id, + MM2S, + IT4, + bd_12, + it_channel_0, + {0, 0, 0, (uint32_t)D * 2}, + {1, 1, 1, (uint32_t)D}, + {0, 0, 0, 1}, + -1, 0, false, aggressive_cache + ); + + seq->npu_dma_memcpy_nd( + sizeof(bf16), + proj_arg_id, + MM2S, + IT4, + bd_13, + it_channel_1, + {0, 0, 0, weight_elem_offset(layer_weights(layer_type).pli_down_proj)}, + {1, 1, 1, (uint32_t)(PLI_D * D)}, + {0, 0, 0, 1}, + -1, 0, true, aggressive_cache + ); + + seq->npu_dma_wait( + IT4, + MM2S, + it_channel_1 + ); +} + +/// \brief Build the sequence that dequantizes one weight into a bf16 buffer. +/// \param seq_ptr the sequence to fill +/// \param weight the weight to dequantize; shape is {input dim, output dim} and +/// offset is a byte offset into the layer's projection buffer +/// \param output_mode which half of an interleaved up/gate region to emit, or +/// NORMAL_DEQUANT for a weight that is not interleaved +void gemma4e_npu_sequence::generate_dequant_seq(npu_sequence* seq_ptr, weight_desc_t& weight, dequant_output_mode_t output_mode){ + assert(weight.added && "weight must be placed in its buffer before it can be dequantized"); + assert(is_quantize(weight.dtype) && "only quantized weights need dequantizing"); + const u32 D_in = (u32)weight.shape[0]; + const u32 D_out = (u32)weight.shape[1]; + const u32 weight_offset = (u32)weight.offset; + static constexpr npu_tiles IT[] = {IT0, IT1, IT2, IT3, IT4, IT5, IT6, IT7}; + + static constexpr u32 total_cols = 8; + static constexpr u32 total_rows = 4; + + static constexpr int w_out_arg_idx = 0; + static constexpr int qw_in_arg_idx = 1; + + static constexpr int m_tile_q4 = 32; + static constexpr int k_tile_q4 = 256; + + const uint32_t block_size_in_byte_q4 = (uint32_t)get_quantization_byte_size((size_t)m_tile_q4 * k_tile_q4, weight.dtype); + + static constexpr int m_tile_q8 = 32; + static constexpr int k_tile_q8 = 128; + static constexpr uint32_t block_size_in_byte_q8 = m_tile_q8 * k_tile_q8 * (10)/8; + + static constexpr int quant_block_col_stride = 2; + + static constexpr int desired_k_dequant = 512; + static constexpr int desired_m_dequant = 128; + + static constexpr int glu_slice = 1024; + static constexpr int gate_up_m_interleave_size = glu_slice / 2; + if (D_in % k_tile_q4 != 0) { + std::cerr << "D_in % k_tile_q4 != 0" << std::endl; + exit(1); + } + int bd_wait_counter[8] = {0, 0, 0, 0, 0, 0, 0, 0}; + // although each data block is in mxk block, but the data block could be reorder in col-stride on block view + /* + For example, quant_block_col_stride = 2 means + + //This is the logical view of the data block, each block of m_tile_q4 x k_tile_q4 + [block0, block1, ...... blockD, + blockD+1, blockD+2, ...... + ] + + But in memory order, the data block is arrange as block0, blockD+1, block1, blockD+2 .... + + */ + + if(D_in % desired_k_dequant != 0){ + std::cerr << "D_in % desired_k_dequant != 0" << std::endl; + exit(1); + } + + const uint32_t blocks_per_row = D_in / k_tile_q4; + + if(D_out % desired_m_dequant != 0 ){ + std::cerr << "D_out % desired_m_dequant != 0" << std::endl; + exit(1); + } + + const int quant_in_per_column = (desired_m_dequant / m_tile_q4) * blocks_per_row * block_size_in_byte_q4; + const int total_column_rounds = D_out / (desired_m_dequant); + + const int row_per_round = desired_m_dequant * total_cols; + // down rounds, go though D_out + const int down_rounds = (D_out + row_per_round - 1) / row_per_round; + + npu_sequence& seq = *seq_ptr; + seq.clear_cmds(); + + for(int row = 0; row < 4; row++){ + for (int col = 0; col < 8; col++){ + npu_tiles tile = get_tile(row + 2, col); + seq_ptr->rtp_write(tile, dequant_rtp_address, 0); + } + } + uint32_t input_offset = weight_offset; + + if(output_mode == dequant_output_mode_t::GATE_MATRIX){ + input_offset += (gate_up_m_interleave_size / m_tile_q4) * blocks_per_row * block_size_in_byte_q4; + } + uint32_t gate_up_interleave_counter= 0; + + // first, the dequant of down + for(int i = 0; i < down_rounds; i++){ + for(int col = 0; col < 8; col++){ + uint32_t bd_offset = (i % 2) * 8; + uint32_t round_offset = i * 8 + col; + if(round_offset < total_column_rounds){ + seq.npu_dma_memcpy_nd( + sizeof(char), + qw_in_arg_idx, + MM2S, + IT[col], + (npu_bd_id)(0+bd_offset), + it_channel_0, + {0, 0, 0, input_offset}, + //NOTE: this for now only work if desired_m_dequant == quant_block_col_stride*m_tile_q4 + { + blocks_per_row, + (desired_m_dequant / m_tile_q4) / quant_block_col_stride, + quant_block_col_stride * block_size_in_byte_q4 / 512, + 512 + }, + { + quant_block_col_stride * block_size_in_byte_q4, + quant_block_col_stride * block_size_in_byte_q4 * blocks_per_row, + 512, + 1 + }, + -1 ,0, false + ); + + if(output_mode == dequant_output_mode_t::NORMAL_DEQUANT){ + input_offset += quant_in_per_column; + } + else{ + gate_up_interleave_counter++; + input_offset += quant_in_per_column; + if(gate_up_interleave_counter == (gate_up_m_interleave_size / desired_m_dequant) ){ + gate_up_interleave_counter = 0; + input_offset += (gate_up_m_interleave_size / m_tile_q4) * blocks_per_row * block_size_in_byte_q4; + } + } + + // Each port receive 2*Q4NX_ROWx D_Q4NX_BLOCK_PER_ROW*Q4NX_COL + uint32_t output_offset_0 = round_offset * desired_m_dequant * D_in; + + seq.npu_dma_memcpy_nd( + sizeof(uint16_t),//bf16 outpout + w_out_arg_idx, + S2MM, + IT[col], + (npu_bd_id)(1+bd_offset), + it_channel_0, + {0, 0, 0, output_offset_0}, + { + (uint32_t)D_in/desired_k_dequant, + desired_k_dequant/k_tile_q4, + desired_m_dequant, + k_tile_q4 + }, + { + desired_m_dequant * desired_k_dequant, + k_tile_q4, + desired_k_dequant, + 1 + }, + -1, 0, true, + aggressive_cache + ); + bd_wait_counter[col]++; + } + } + // note: for now + for(int col = 0; col < 8; col++){ + if(bd_wait_counter[col] == 2){ + seq.npu_dma_wait(IT[col], S2MM, it_channel_0); + bd_wait_counter[col]--; + } + } + } + + for(int col = 0; col < 8; col++){ + while(bd_wait_counter[col] != 0){ + seq.npu_dma_wait(IT[col], S2MM, it_channel_0); + bd_wait_counter[col]--; + } + } + seq.cmds2seq(); +} + +void gemma4e_npu_sequence::gen_mha_engine_seq( + npu_sequence* seq, + const uint32_t L_begin, + const uint32_t L_end +){ + const int lc = 8; // local chunk size, each CU works on lc*8 of l at a time. This is determined by the hardware design. + + npu_tiles IT[2][4] = {{IT0, IT1, IT2, IT3}, {IT4, IT5, IT6, IT7}}; + assert(L_begin % (lc * 16) == 0); + assert(L_end % (lc * 16) == 0); + int using_window_size = L_end; + const int Heads = num_attn_heads; + const int GQA = num_attn_heads / num_kv_heads; + const int Num_of_d_in_KV_cache = num_kv_heads; + const int DH = this->DH; + const uint32_t KV_CACHE_SIZE = (DK + DV) * MAX_L; + seq->clear_cmds(); + + const int l_begin_mha_address = 59520; + const int l_end_mha_address = 11520; + const int window_size_address = 59552; + for (int row = 2; row < 6; row++){ + for (int col = 0; col < 8; col++){ + npu_tiles tile = get_tile(row, col); + seq->rtp_write(tile, l_begin_mha_address, L_begin); + seq->rtp_write(tile, l_end_mha_address, L_end); + seq->rtp_write(tile, window_size_address, using_window_size); + } + } + int all_data_size = L_end - L_begin; + const int data_per_round = lc * 16; + const int down_rounds = (all_data_size + data_per_round - 1) / data_per_round; + int num_cu = 2; + for (int head = 0; head < Heads / num_cu; head++){ + for (int round = 0; round < down_rounds; round++){ + int Lq_current = L_begin + round * data_per_round; + int kv_begin = ((Lq_current - using_window_size) > 0) ? (Lq_current - using_window_size) : 0; + int kv_length = Lq_current - kv_begin + data_per_round; + + int bd_offset = (round % 2) * 8; + for (int cu = 0; cu < num_cu; cu++){ + int head_offset = head * num_cu + cu; + for (int col = 0; col < 4; col++){ + // receive y + size_t y_offset = head_offset * DH + (round * data_per_round + col * lc * 4) * DH * Heads; + seq->npu_dma_memcpy_nd( + 2, 0, + S2MM, IT[cu][col], + (npu_bd_id)(bd_offset + 0), it_channel_0, + {0, 0, 0, (uint32_t)y_offset}, + {1, 1, (uint32_t)4 * lc, (uint32_t)(DH)}, + {0, 0, (uint32_t)DH * Heads, 1}, + -1, 0, true + ); + } + // send q + size_t q_offset = head_offset * DH + round * data_per_round * DH * Heads; + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[cu][0], + (npu_bd_id)(bd_offset + 1), it_channel_0, + {0, 0, 0, (uint32_t)q_offset}, + {1, 1, (uint32_t)lc * 4, (uint32_t)(DH)}, + {0, 0, (uint32_t)DH * Heads, 1}, + -1, 0, false + ); + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[cu][0], + (npu_bd_id)(bd_offset + 2), it_channel_1, + {0, 0, 0, (uint32_t)(q_offset + lc * 4 * DH * Heads)}, + {1, 1, (uint32_t)lc * 4, (uint32_t)(DH)}, + {0, 0, (uint32_t)DH * Heads, 1}, + -1, 0, false + ); + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[cu][3], + (npu_bd_id)(bd_offset + 3), it_channel_0, + {0, 0, 0, (uint32_t)(q_offset + lc * 8 * DH * Heads)}, + {1, 1, (uint32_t)lc * 4, (uint32_t)(DH)}, + {0, 0, (uint32_t)DH * Heads, 1}, + -1, 0, false + ); + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[cu][3], + (npu_bd_id)(bd_offset + 4), it_channel_1, + {0, 0, 0, (uint32_t)(q_offset + lc * 12 * DH * Heads)}, + {1, 1, (uint32_t)lc * 4, (uint32_t)(DH)}, + {0, 0, (uint32_t)DH * Heads, 1}, + -1, 0, false + ); + } // cu + int kv_head_offset = head / (GQA / num_cu); + int kv_chunk_offset = kv_head_offset / Num_of_d_in_KV_cache; + int kv_head_offset_in_chunk = kv_head_offset % Num_of_d_in_KV_cache; + size_t k_offset; + size_t v_offset; + if (kv_length <= 128 * 1024){ + k_offset = kv_chunk_offset * MAX_L * DH * Num_of_d_in_KV_cache + kv_head_offset_in_chunk * DH + kv_begin * DH * Num_of_d_in_KV_cache; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[0][2], + (npu_bd_id)(bd_offset + 5), it_channel_0, + {0, 0, 0, (uint32_t)k_offset}, + {1, (uint32_t)kv_length / 128, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + v_offset = k_offset + KV_CACHE_SIZE / 2; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[0][2], + (npu_bd_id)(bd_offset + 6), it_channel_1, + {0, 0, 0, (uint32_t)v_offset}, + {1, (uint32_t)kv_length / 128, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + } + else{ + k_offset = kv_chunk_offset * MAX_L * DH * Num_of_d_in_KV_cache + kv_head_offset_in_chunk * DH + kv_begin * DH * Num_of_d_in_KV_cache; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[0][2], + (npu_bd_id)(bd_offset + 5), it_channel_0, + {0, 0, 0, (uint32_t)k_offset}, + {1, (uint32_t)1024, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + v_offset = k_offset + KV_CACHE_SIZE / 2; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[0][2], + (npu_bd_id)(bd_offset + 6), it_channel_1, + {0, 0, 0, (uint32_t)v_offset}, + {1, (uint32_t)1024, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + int remaining_kv_length = kv_length - 1024 * 128; + k_offset = kv_chunk_offset * MAX_L * DH * Num_of_d_in_KV_cache + kv_head_offset_in_chunk * DH + kv_begin * DH * Num_of_d_in_KV_cache + 1024 * 128 * DH * Num_of_d_in_KV_cache; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[0][2], + (npu_bd_id)(bd_offset + 5), it_channel_0, + {0, 0, 0, (uint32_t)k_offset}, + {1, (uint32_t)remaining_kv_length / 128, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + v_offset = k_offset + KV_CACHE_SIZE / 2 + 1024 * 128 * DH * Num_of_d_in_KV_cache; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[0][2], + (npu_bd_id)(bd_offset + 6), it_channel_1, + {0, 0, 0, (uint32_t)v_offset}, + {1, (uint32_t)remaining_kv_length / 128, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + } + if (round > 0){ + for (int cu = 0; cu < num_cu; cu++){ + for (int col = 0; col < 4; col++){ + seq->npu_dma_wait( + IT[cu][col], + S2MM, + it_channel_0 + ); + } // col + } // cu + }// round > 0 + }// round + for (int cu = 0; cu < num_cu; cu++){ + for (int col = 0; col < 4; col++){ + seq->npu_dma_wait( + IT[cu][col], + S2MM, + it_channel_0 + ); + } // col + } // cu + } // head + seq->cmds2seq(); +} + +void gemma4e_npu_sequence::gen_swa_engine_seq( + npu_sequence* seq, + const uint32_t L_begin, + const uint32_t L_end +){ + const int lc = 16; // local chunk size, each CU works on lc*8 of l at a time. This is determined by the hardware design. + npu_tiles IT[4][2] = {{IT0, IT1}, {IT2, IT3}, {IT4, IT5}, {IT6, IT7}}; + assert(L_begin % (lc * 8) == 0); + assert(L_end % (lc * 8) == 0); + int using_window_size = SLIDING_LENGTH; + const int Heads = num_attn_heads; + const int GQA = num_attn_heads / num_kv_heads; + const int Num_of_d_in_KV_cache = num_kv_heads; + const int DH = this->SWA_DH; + const uint32_t KV_CACHE_SIZE = (SWA_DK + SWA_DV) * MAX_L; + seq->clear_cmds(); + + const int l_begin_mha_address = 61568; + const int l_end_mha_address = 11904; + const int window_size_address = 61600; + for (int row = 2; row < 6; row++){ + for (int col = 0; col < 8; col++){ + npu_tiles tile = get_tile(row, col); + seq->rtp_write(tile, l_begin_mha_address, L_begin); + seq->rtp_write(tile, l_end_mha_address, L_end); + seq->rtp_write(tile, window_size_address, using_window_size); + } + } + // each CU has 8 CTs and works on 1 head. + int all_data_size = L_end - L_begin; + // each round, each CU works on 1 head and lc * 8 of l, total_cols / 2 is corresponding to the number of CUs + const int data_per_round = lc * 8; + // zero padding the last round if necessary + const int down_rounds = (all_data_size + data_per_round - 1) / data_per_round; + + int num_cu = 4; + for (int head = 0; head < Heads / num_cu; head++){ + for (int round = 0; round < down_rounds; round++){ + int Lq_current = L_begin + round * data_per_round; + int kv_begin = ((Lq_current - using_window_size) > 0) ? (Lq_current - using_window_size) : 0; + int kv_length = Lq_current - kv_begin + data_per_round; + + int bd_offset = (round % 2) * 8; + + for (int cu = 0; cu < num_cu; cu++){ + int head_offset = head * num_cu + cu; + + for (int col = 0; col < 2; col++){ + // receive y + size_t y_offset = head_offset * DH + (round * data_per_round + col * lc * 4) * DH * Heads; + seq->npu_dma_memcpy_nd( + 2, 0, + S2MM, IT[cu][col], + (npu_bd_id)(bd_offset + 0), it_channel_0, + {0, 0, 0, (uint32_t)y_offset}, + {1, 1, (uint32_t)4 * lc, (uint32_t)(DH)}, + {0, 0, (uint32_t)DH * Heads, 1}, + -1, 0, true + ); + } + // send q + size_t q_offset = head_offset * DH + round * data_per_round * DH * Heads; + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[cu][0], + (npu_bd_id)(bd_offset + 1), it_channel_0, + {0, 0, 0, (uint32_t)q_offset}, + {1, 1, (uint32_t)lc * 4, (uint32_t)(DH)}, + {0, 0, (uint32_t)DH * Heads, 1}, + -1, 0, false + ); + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[cu][0], + (npu_bd_id)(bd_offset + 2), it_channel_1, + {0, 0, 0, (uint32_t)(q_offset + lc * 4 * DH * Heads)}, + {1, 1, (uint32_t)lc * 4, (uint32_t)(DH)}, + {0, 0, (uint32_t)DH * Heads, 1}, + -1, 0, false + ); + } // cu + + int kv_head_offset = head / (GQA / num_cu); + int kv_chunk_offset = kv_head_offset / Num_of_d_in_KV_cache; + int kv_head_offset_in_chunk = kv_head_offset % Num_of_d_in_KV_cache; + + size_t k_offset; + size_t v_offset; + if (kv_length <= 128 * 1024){ + k_offset = kv_chunk_offset * MAX_L * DH * Num_of_d_in_KV_cache + kv_head_offset_in_chunk * DH + kv_begin * DH * Num_of_d_in_KV_cache; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[1][1], + (npu_bd_id)(bd_offset + 3), it_channel_0, + {0, 0, 0, (uint32_t)k_offset}, + {1, (uint32_t)kv_length / 128, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + v_offset = kv_chunk_offset * MAX_L * DH * Num_of_d_in_KV_cache + kv_head_offset_in_chunk * DH + kv_begin * DH * Num_of_d_in_KV_cache + KV_CACHE_SIZE / 2; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[1][1], + (npu_bd_id)(bd_offset + 4), it_channel_1, + {0, 0, 0, (uint32_t)v_offset}, + {1, (uint32_t)kv_length / 128, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + } else{ + k_offset = kv_chunk_offset * MAX_L * DH * Num_of_d_in_KV_cache + kv_head_offset_in_chunk * DH + kv_begin * DH * Num_of_d_in_KV_cache; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[1][1], + (npu_bd_id)(bd_offset + 3), it_channel_0, + {0, 0, 0, (uint32_t)k_offset}, + {1, (uint32_t)1024, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + v_offset = kv_chunk_offset * MAX_L * DH * Num_of_d_in_KV_cache + kv_head_offset_in_chunk * DH + kv_begin * DH * Num_of_d_in_KV_cache + KV_CACHE_SIZE / 2; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[1][1], + (npu_bd_id)(bd_offset + 4), it_channel_1, + {0, 0, 0, (uint32_t)v_offset}, + {1, (uint32_t)1024, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + int remaining_kv_length = kv_length - 1024 * 128; + k_offset = kv_chunk_offset * MAX_L * DH * Num_of_d_in_KV_cache + kv_head_offset_in_chunk * DH + kv_begin * DH * Num_of_d_in_KV_cache + 1024 * 128 * DH * Num_of_d_in_KV_cache; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[1][1], + (npu_bd_id)(bd_offset + 5), it_channel_0, + {0, 0, 0, (uint32_t)k_offset}, + {1, (uint32_t)remaining_kv_length / 128, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + v_offset = kv_chunk_offset * MAX_L * DH * Num_of_d_in_KV_cache + kv_head_offset_in_chunk * DH + kv_begin * DH * Num_of_d_in_KV_cache + KV_CACHE_SIZE / 2 + 1024 * 128 * DH * Num_of_d_in_KV_cache; + seq->npu_dma_memcpy_nd( + 2, 2, + MM2S, IT[1][1], + (npu_bd_id)(bd_offset + 6), it_channel_1, + {0, 0, 0, (uint32_t)v_offset}, + {1, (uint32_t)remaining_kv_length / 128, (uint32_t)128, (uint32_t)(DH)}, + {0, (uint32_t)128 * DH * Num_of_d_in_KV_cache, (uint32_t)DH * Num_of_d_in_KV_cache, 1}, + -1, 0, false + ); + } + if (round > 0){ + for (int cu = 0; cu < num_cu; cu++){ + for (int col = 0; col < 2; col++){ + seq->npu_dma_wait( + IT[cu][col], + S2MM, + it_channel_0 + ); + } // col + } // cu + }// round > 0 + }// round + for (int cu = 0; cu < num_cu; cu++){ + for (int col = 0; col < 2; col++){ + seq->npu_dma_wait( + IT[cu][col], + S2MM, + it_channel_0 + ); + } // col + } // cu + } // head + seq->cmds2seq(); +} + +gemma4e_npu_sequence::~gemma4e_npu_sequence() = default; diff --git a/src/detail/gemma4e_npu/gemma4e_npu_sequence.hpp b/src/detail/gemma4e_npu/gemma4e_npu_sequence.hpp new file mode 100644 index 000000000..0097fc969 --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_npu_sequence.hpp @@ -0,0 +1,216 @@ +#pragma once + +#include "npu_utils/npu_instr_utils.hpp" +#include "lm_config.hpp" +#include "tensor_utils/q4_npu_eXpress.hpp" +#include "models/gemma4e/flm/aie2p/gemma4e_npu.hpp" +#include "weight_desc.hpp" +#include "gemma4e_npu_def.hpp" + +#define MIN_BF16_PAD 32 // minimum padding in number of bf16 elements to avoid NPU OOM when processing long sequences, which is determined empirically + +#define LAYER_SPECIFIC_DIM(type) \ + int _DH = is_global_layer(type) ? DH : SWA_DH; \ + int _DQ = is_global_layer(type) ? DQ : SWA_DQ; \ + int _DK = is_global_layer(type) ? DK : SWA_DK; \ + int _DV = is_global_layer(type) ? DV : SWA_DV; \ + int _INTERMEDIATE_SIZE = (is_skip_layer(type) && enable_double_wide_mlp) ? INTERMEDIATE_SIZE * 2 : INTERMEDIATE_SIZE; + +typedef struct { + int D; + int DH; + int DQ; + int DK; + int DV; + int SWA_DH; + int SWA_DQ; + int SWA_DK; + int SWA_DV; + int PLI_D; + int INTERMEDIATE_SIZE; + int NUM_ATTENTION_HEADS; + int NUM_KEY_VALUE_HEADS; + int SLIDING_WINDOW_SIZE; + int VOCAB_SIZE_PADDED; + bool enable_double_wide_mlp; +} gemma4e_seq_gen_parameters_t; + +struct gemma4e_npu_sequence{ + typedef enum: int{ + NORMAL_DEQUANT = 0, + UP_MATRIX = 1, + GATE_MATRIX = 2 + } dequant_output_mode_t; + + typedef struct { + // for decoding layer + uint32_t l_qk_address; + uint32_t l_kv_address; + uint32_t swa_l_qk_address; + uint32_t swa_l_kv_address; + uint32_t proj_swa_address; + uint32_t proj_skip_address; + uint32_t rms_swa_address; + uint32_t rms_skip_address; + uint32_t rope_skip_kv_address; + uint32_t swa_rope_skip_kv_address; + uint32_t glu_skip_address; + uint32_t lm_head_final_tune_address; + // for others + uint32_t place_holder_0; + } rtp_address_book_t; + + static constexpr rtp_address_book_t e2b_rtp_addresses = { + .l_qk_address = 57344, + .l_kv_address = 14976, + .swa_l_qk_address = 9216, + .swa_l_kv_address = 40960, + .proj_swa_address = 33280, + .proj_skip_address = 49664, + .rms_swa_address = 52224, + .rms_skip_address = 25600, + .rope_skip_kv_address = 33792, + .swa_rope_skip_kv_address = 33280, + .glu_skip_address = 34816, + .lm_head_final_tune_address = 49152, + .place_holder_0 = 0x20050625 + }; + + static constexpr rtp_address_book_t e4b_rtp_addresses = { + .l_qk_address = 57344, + .l_kv_address = 14976, + .swa_l_qk_address = 53248, + .swa_l_kv_address = 57664, + .proj_swa_address = 33280, + .proj_skip_address = 49664, + .rms_swa_address = 55328, + .rms_skip_address = 55392, + .rope_skip_kv_address = 34816, + .swa_rope_skip_kv_address = 33792, + .glu_skip_address = 30720, + .lm_head_final_tune_address = 10240, + .place_holder_0 = 0x20050625 + }; + /// @brief constexprs + static constexpr npu_tiles xr_tile = IT3; + static constexpr npu_tiles out_tile = IT4; + static constexpr npu_tiles attn_tile = IT2; + static constexpr npu_tiles mvm_tiles[4] = {IT0, IT1, IT6, IT7}; + static constexpr npu_tiles rms_tile = CT03; + static constexpr npu_tiles glu_tile = CT13; + static constexpr npu_tiles rope_ct = CT23; + static constexpr npu_tiles swa_rope_ct = CT33; + static constexpr npu_tiles attn_qk_tile = CT02; + static constexpr npu_tiles attn_kv_tile = CT12; + static constexpr npu_tiles swa_attn_qk_tile = CT22; + static constexpr npu_tiles swa_attn_kv_tile = CT32; + static constexpr npu_tiles proj_tiles[] = {CT00, CT10, CT20, CT30, + CT01, CT11, CT21, CT31, + CT06, CT16, CT26, CT36, + CT07, CT17, CT27, CT37 + }; + static constexpr npu_tiles pli_tile = CT05; + + static constexpr int x_arg_id = 0; + static constexpr int proj_arg_id = 1; + static constexpr int rms_arg_id = 2; + static constexpr int rope_rms_arg_id = 3; + static constexpr int kv_cache_arg_id = 4; + + static constexpr int L_CHUNK = 16; + + /// @brief dequant kernel RTP offset, identical for E2B and E4B. + static constexpr uint32_t dequant_rtp_address = 54272; + + /// \brief The model description; owns every weight descriptor this generator addresses. + /// \note Not owned. Set by set_desc() before any sequence is generated. + gemma4e_desc* desc = nullptr; + rtp_address_book_t rtp_addresses; + + uint32_t D; + uint32_t DH; + uint32_t DQ; + uint32_t DK; + uint32_t DV; + uint32_t SWA_DH; + uint32_t SWA_DQ; + uint32_t SWA_DK; + uint32_t SWA_DV; + + uint32_t MAX_L; + uint32_t HIDDEN_SIZE; + uint32_t INTERMEDIATE_SIZE; + uint32_t PLI_D; + uint32_t SLIDING_LENGTH; + + bool enable_double_wide_mlp; + + int num_attn_heads; + int num_kv_heads; + int num_kv_per_round; + + uint32_t rms_addr; + + int VOCAB_SIZE; + + gemma4e_npu_sequence(){} + gemma4e_npu_sequence(gemma4e_seq_gen_parameters_t params, uint32_t MAX_L); + ~gemma4e_npu_sequence(); + /// \brief Human-readable name of a layer kind, for debug output. + static const char* layer_type_name(gemma4e_layer_type_t t) { + switch (t) { + case e_gemma4e_swa_layer: return "swa"; + case e_gemma4e_global_layer: return "global"; + case e_gemma4e_swa_layer_skip: return "swa_skip"; + case e_gemma4e_global_layer_skip: return "global_skip"; + default: return "unknown"; + } + } + + /// \brief Point the generator at the model description. + /// \param desc the description whose weight descriptors name every DMA source + void set_desc(gemma4e_desc* desc) { + this->desc = desc; + DEBUG_BLOCK(2, + header_print("info", "Sequence generator bound to gemma4e_desc; proj layout (bytes):"); + for (int t = 0; t < 4; t++){ + gemma4e_layer_weight_def& L = desc->weight_desc(static_cast(t)); + std::cout << " " << layer_type_name(static_cast(t)) + << ": qkv=" << L.attn_qkv.offset + << ", o=" << L.attn_output.offset + << ", upgate=" << L.ffn_up_gate.offset + << ", down=" << L.ffn_down.offset + << ", pli_down=" << L.pli_down_proj.offset + << ", pli_gate=" << L.pli_gate_proj.offset + << ", pli_up=" << L.pli_up_proj.offset << std::endl; + } + ) + } + + /// \brief The weight layout of a layer kind. + gemma4e_layer_weight_def& layer_weights(gemma4e_layer_type_t layer_type) { + assert(desc != nullptr && "set_desc() must run before any sequence is generated"); + return desc->weight_desc(layer_type); + } + + /// \brief A weight's start, in bf16 elements, which is how the DMA ports address it. + static uint32_t weight_elem_offset(weight_desc_t& weight) { + return (uint32_t)(weight.offset / sizeof(bf16)); + } + + void _send_hidden_states(npu_sequence* seq); + void _send_rms_weights(npu_sequence* seq); + void _send_rope_rms_weights(npu_sequence* seq, gemma4e_layer_type_t layer_type); + void _receive_kv_cache(npu_sequence* seq, const int L, gemma4e_layer_type_t layer_type); + void _move_kv_cache(npu_sequence* seq, const size_t L, gemma4e_layer_type_t layer_type); + void _move_weights(npu_sequence* seq, weight_desc_t& weight); + void _gen_pli_path_seq(npu_sequence* seq_ptr, gemma4e_layer_type_t layer_type); + void gen_lm_head_seq(npu_sequence* seq, float final_scale); + + void gen_layer_seq(npu_sequence* seq, const uint32_t L, gemma4e_layer_type_t layer_type); + void gen_mha_engine_seq(npu_sequence* seq, const uint32_t L_begin, const uint32_t L_end); + void gen_swa_engine_seq(npu_sequence* seq, const uint32_t L_begin, const uint32_t L_end); + void generate_dequant_seq(npu_sequence* seq_ptr, weight_desc_t& weight, dequant_output_mode_t output_mode); + + void set_max_length(const uint32_t MAX_L); +}; diff --git a/src/detail/gemma4e_npu/gemma4e_prefill.cpp b/src/detail/gemma4e_npu/gemma4e_prefill.cpp new file mode 100644 index 000000000..fe7056fa4 --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_prefill.cpp @@ -0,0 +1,584 @@ +#include "flm_override.hpp" +#include "gemma4e_prefill.hpp" +#include "metrices.hpp" + +// --------------------------------------------------------------------------- +// attention block +// --------------------------------------------------------------------------- + +void gemma4e_attn_block_prefill_context::_forward_swa( + const gemma4e_prefill_shape& s, + gemma4e_layer_type_t type, + buffer& qkv_weights, + buffer& o_weights, + buffer& kv_cache, + buffer& q_norm, + buffer& k_norm, + SafeTensors* reference +){ + const bool is_skip = is_skip_layer(type); + const int D = desc->D; + const int SWA_DQ = desc->SWA_DQ; + const int SWA_DK = desc->SWA_DK; + const int SWA_DV = desc->SWA_DV; + + FLM_OVERRIDE(q_swa_proj, q_swa_proj(bufs->q_buffer, bufs->hidden_state_buffer, qkv_weights), this->desc, type, s); + bufs->q_buffer.sync_from_device(); + + DEBUG_BLOCK(2, + buffer q_ref; + reference->load_weights(q_ref, "q_proj"); + buffer q_valid = buffer(bufs->q_buffer.data(), s.L_effective * SWA_DQ); + buffer q_valid_ref = buffer(q_ref.data(), s.L_effective * SWA_DQ); + print_error_metrics(get_error_metrics(q_valid, q_valid_ref), "SWA Q Projection Error: "); + ) + if (is_skip){ + // a skip layer carries no k/v projection: it reuses the cache the last + // non-skip layer of its kind filled. + gemma4e_cpu_func::_rope_rms_batch(bufs->q_buffer.data(), SWA_DQ, s.L_offset, s.L_begin, s.L_effective, q_norm.data(), type, desc->get_DH(type)); + bufs->q_buffer.sync_to_device(); + } + else { + auto run_k = FLM_OVERRIDE(k_swa_proj, + k_swa_proj.create_run(bufs->k_buffer, bufs->hidden_state_buffer, qkv_weights), this->desc, type, s); + run_k.start(); + + gemma4e_cpu_func::_rope_rms_batch(bufs->q_buffer.data(), SWA_DQ, s.L_offset, s.L_begin, s.L_effective, q_norm.data(), type, desc->get_DH(type)); + bufs->q_buffer.sync_to_device(); + + run_k.wait(); + bufs->k_buffer.sync_from_device(); + auto run_v = FLM_OVERRIDE(v_swa_proj, + v_swa_proj.create_run(bufs->v_buffer, bufs->hidden_state_buffer, qkv_weights), this->desc, type, s); + run_v.start(); + gemma4e_cpu_func::_rope_rms_batch(bufs->k_buffer.data(), SWA_DK, s.L_offset, s.L_begin, s.L_effective, k_norm.data(), type, desc->get_DH(type)); + bufs->k_buffer.sync_to_device(); + run_v.wait(); + + bufs->v_buffer.sync_from_device(); + gemma4e_cpu_func::_rms_norm_batch(bufs->v_buffer.data(), bufs->v_buffer.data(), nullptr, SWA_DV, SWA_DV, s.L_effective, s.L_offset, s.L_offset); + bufs->v_buffer.sync_to_device(); + } + + DEBUG_BLOCK(2, + buffer q_ref; + reference->load_weights(q_ref, "q_embed"); + buffer q_valid = buffer(bufs->q_buffer.data(), s.L_effective * SWA_DQ); + buffer q_valid_ref = buffer(q_ref.data(), s.L_effective * SWA_DQ); + print_error_metrics(get_error_metrics(q_valid, q_valid_ref), "SWA Q Projection Error: "); + if (!is_skip){ + buffer k_ref; + reference->load_weights(k_ref, "k_embed"); + buffer k_valid = buffer(bufs->k_buffer.data(), s.L_effective * SWA_DK); + buffer k_valid_ref = buffer(k_ref.data(), s.L_effective * SWA_DK); + print_error_metrics(get_error_metrics(k_valid, k_valid_ref), "SWA K Projection Error: "); + + buffer v_ref; + reference->load_weights(v_ref, "v_norm"); + buffer v_valid = buffer(bufs->v_buffer.data(), s.L_effective * SWA_DV); + buffer v_valid_ref = buffer(v_ref.data(), s.L_effective * SWA_DV); + print_error_metrics(get_error_metrics(v_valid, v_valid_ref), "SWA V Projection Error: "); + } + ) + if (!is_skip){ + sync_sliding_kv_cache(bufs->k_buffer, bufs->v_buffer, kv_cache_sliding_prefill, kv_cache, s.L_offset, s.L_begin, s.L_effective, s.sliding_l_begin, s.L_end_chunked, SWA_DK, SWA_DV); + } + else{ + kv_cache_sliding_prefill.sync_from_device(); + kv_cache_sliding_prefill.sync_to_device(); + } + FLM_OVERRIDE(swa_attn_core, + this->swa_engine(bufs->attn_out_buffer, bufs->q_buffer, kv_cache_sliding_prefill), type, s, this->MAX_L); + bufs->attn_out_buffer.sync_from_device(); + DEBUG_BLOCK(2, + buffer attn_out_ref; + reference->load_weights(attn_out_ref, "attention_output"); + buffer attn_out_valid = buffer(bufs->attn_out_buffer.data(), s.L_effective * SWA_DQ); + buffer attn_out_valid_ref = buffer(attn_out_ref.data(), s.L_effective * SWA_DQ); + print_error_metrics(get_error_metrics(attn_out_valid, attn_out_valid_ref), "SWA Attention Output Error: "); + ) + + FLM_OVERRIDE(o_swa_proj, + this->o_swa_proj(bufs->hidden_state_buffer, bufs->attn_out_buffer, o_weights), this->desc, type, s); + bufs->hidden_state_buffer.sync_from_device(); + + DEBUG_BLOCK(2, + buffer o_proj_ref; + reference->load_weights(o_proj_ref, "after_o_proj"); + buffer o_proj_valid = buffer(bufs->hidden_state_buffer.data(), s.L_effective * D); + buffer o_proj_valid_ref = buffer(o_proj_ref.data(), s.L_effective * D); + print_error_metrics(get_error_metrics(o_proj_valid, o_proj_valid_ref), "SWA O Projection Error: "); + ) +} + +void gemma4e_attn_block_prefill_context::_forward_global( + const gemma4e_prefill_shape& s, + gemma4e_layer_type_t type, + buffer& qkv_weights, + buffer& o_weights, + buffer& kv_cache, + buffer& q_norm, + buffer& k_norm, + SafeTensors* reference +){ + const bool is_skip = is_skip_layer(type); + const int D = desc->D; + const int DQ = desc->DQ; + const int DK = desc->DK; + const int DV = desc->DV; + + FLM_OVERRIDE(q_global_proj, q_global_proj(bufs->q_buffer, bufs->hidden_state_buffer, qkv_weights), this->desc, type, s); + bufs->q_buffer.sync_from_device(); + + DEBUG_BLOCK(2, + buffer q_ref; + reference->load_weights(q_ref, "q_proj"); + buffer q_valid = buffer(bufs->q_buffer.data(), s.L_effective * DQ); + buffer q_valid_ref = buffer(q_ref.data(), s.L_effective * DQ); + print_error_metrics(get_error_metrics(q_valid, q_valid_ref), "Q Projection Error: "); + ) + if (is_skip){ + gemma4e_cpu_func::_rope_rms_batch(bufs->q_buffer.data(), DQ, s.L_offset, s.L_begin, s.L_effective, q_norm.data(), type, desc->get_DH(type)); + bufs->q_buffer.sync_to_device(); + } + else { + auto run_k = FLM_OVERRIDE(k_global_proj, + k_global_proj.create_run(bufs->k_buffer, bufs->hidden_state_buffer, qkv_weights), this->desc, type, s); + run_k.start(); + gemma4e_cpu_func::_rope_rms_batch(bufs->q_buffer.data(), DQ, s.L_offset, s.L_begin, s.L_effective, q_norm.data(), type, desc->get_DH(type)); + bufs->q_buffer.sync_to_device(); + run_k.wait(); + bufs->k_buffer.sync_from_device(); + + auto run_v = FLM_OVERRIDE(v_global_proj, + v_global_proj.create_run(bufs->v_buffer, bufs->hidden_state_buffer, qkv_weights), this->desc, type, s); + run_v.start(); + gemma4e_cpu_func::_rope_rms_batch(bufs->k_buffer.data(), DK, s.L_offset, s.L_begin, s.L_effective, k_norm.data(), type, desc->get_DH(type)); + bufs->k_buffer.sync_to_device(); + run_v.wait(); + + bufs->v_buffer.sync_from_device(); + gemma4e_cpu_func::_rms_norm_batch(bufs->v_buffer.data(), bufs->v_buffer.data(), nullptr, DV, DV, s.L_effective, s.L_offset, s.L_offset); + bufs->v_buffer.sync_to_device(); + } + + DEBUG_BLOCK(2, + buffer q_ref; + reference->load_weights(q_ref, "q_embed"); + buffer q_valid = buffer(bufs->q_buffer.data(), s.L_effective * DQ); + buffer q_valid_ref = buffer(q_ref.data(), s.L_effective * DQ); + print_error_metrics(get_error_metrics(q_valid, q_valid_ref), "Q Projection Error: "); + + buffer k_ref; + reference->load_weights(k_ref, "k_embed"); + buffer k_valid = buffer(bufs->k_buffer.data(), s.L_effective * DK); + buffer k_valid_ref = buffer(k_ref.data(), s.L_effective * DK); + print_error_metrics(get_error_metrics(k_valid, k_valid_ref), "K Projection Error: "); + + buffer v_ref; + reference->load_weights(v_ref, "v_norm"); + buffer v_valid = buffer(bufs->v_buffer.data(), s.L_effective * DV); + buffer v_valid_ref = buffer(v_ref.data(), s.L_effective * DV); + print_error_metrics(get_error_metrics(v_valid, v_valid_ref), "V Projection Error: "); + ) + if (!is_skip){ + sync_kv_cache(bufs->k_buffer, bufs->v_buffer, kv_cache_global_prefill, kv_cache, s.L_offset, s.L_begin, s.L_effective, DK, DV); + } + else { + kv_cache_global_prefill.sync_from_device(); + kv_cache_global_prefill.sync_to_device(); + } + FLM_OVERRIDE(global_attn_core, + this->mha_engine(bufs->attn_out_buffer, bufs->q_buffer, kv_cache_global_prefill), type, s, this->MAX_L); + DEBUG_BLOCK(2, + header_print("info", "Attn done!"); + ) + bufs->attn_out_buffer.sync_from_device(); + DEBUG_BLOCK(2, + buffer attn_out_ref; + reference->load_weights(attn_out_ref, "attention_output"); + buffer attn_out_valid = buffer(bufs->attn_out_buffer.data(), s.L_effective * DQ); + buffer attn_out_valid_ref = buffer(attn_out_ref.data(), s.L_effective * DQ); + print_error_metrics(get_error_metrics(attn_out_valid, attn_out_valid_ref), "Attention Output Error: "); + ) + + FLM_OVERRIDE(o_global_proj, + this->o_global_proj(bufs->hidden_state_buffer, bufs->attn_out_buffer, o_weights), this->desc, type, s); + bufs->hidden_state_buffer.sync_from_device(); + + DEBUG_BLOCK(2, + buffer o_proj_ref; + reference->load_weights(o_proj_ref, "after_o_proj"); + buffer o_proj_valid = buffer(bufs->hidden_state_buffer.data(), s.L_effective * D); + buffer o_proj_valid_ref = buffer(o_proj_ref.data(), s.L_effective * D); + print_error_metrics(get_error_metrics(o_proj_valid, o_proj_valid_ref), "O Projection Error: "); + ) +} + +void gemma4e_attn_block_prefill_context::sync_kv_cache( + buffer& buffer_k, + buffer& buffer_v, + buffer& prefill_cache, + buffer& decoding_cache, + int L_offset, + int L_begin, + int L_effective, + int DK, int DV +){ + bf16* prefill_k_cache_ptr = prefill_cache.data(); + bf16* prefill_v_cache_ptr = prefill_cache.data() + (size_t)MAX_L * DK; + bf16* decoding_k_cache_ptr = decoding_cache.data(); + bf16* decoding_v_cache_ptr = decoding_cache.data() + (size_t)MAX_L * DK; + // copy the existing cache up to L_begin from decoding_cache to prefill_cache, and then copy the new k cache from buffer_k and v cache from buffer_v to prefill_cache at the appropriate location, then sync the prefill_cache to device for the attention run + decoding_cache.sync_from_device(); + prefill_cache.sync_from_device(); + memcpy(prefill_k_cache_ptr, decoding_k_cache_ptr, L_begin * DK * sizeof(bf16)); // copy the existing cache up to L_begin + memcpy(prefill_v_cache_ptr, decoding_v_cache_ptr, L_begin * DV * sizeof(bf16)); // copy the existing cache up to L_begin + + prefill_k_cache_ptr += L_begin * DK; + prefill_v_cache_ptr += L_begin * DV; + decoding_k_cache_ptr += L_begin * DK; + decoding_v_cache_ptr += L_begin * DV; + + // copy the new k cache and v cache to the appropriate location in prefill_cache and decoding_cache + bf16* new_k_cache_ptr = buffer_k.data() + L_offset * DK; + bf16* new_v_cache_ptr = buffer_v.data() + L_offset * DV; + + memcpy(prefill_k_cache_ptr, new_k_cache_ptr, L_effective * DK * sizeof(bf16)); + memcpy(prefill_v_cache_ptr, new_v_cache_ptr, L_effective * DV * sizeof(bf16)); + memcpy(decoding_k_cache_ptr, new_k_cache_ptr, L_effective * DK * sizeof(bf16)); + memcpy(decoding_v_cache_ptr, new_v_cache_ptr, L_effective * DV * sizeof(bf16)); + + decoding_cache.sync_to_device(); + prefill_cache.sync_to_device(); +} + +void gemma4e_attn_block_prefill_context::sync_sliding_kv_cache( + buffer& buffer_k, + buffer& buffer_v, + buffer& prefill_cache, + buffer& decoding_cache, + int L_offset, + int L_begin, + int L_effective, + int sliding_l_begin, + int L_end_chunked, + int DK, int DV +){ + const int SLIDING_LENGTH = desc->SLIDING_LENGTH; + auto linear2ring = [sliding = SLIDING_LENGTH](int l) { return l % sliding; }; + int in_memory = (L_begin - SLIDING_LENGTH) > 0 ? L_begin - SLIDING_LENGTH : 0; + int l_idx = in_memory % SLIDING_LENGTH; + int local_offset = in_memory - sliding_l_begin; + int local_length = L_end_chunked - sliding_l_begin; + DEBUG_BLOCK(2, + std::cout << "DEBUG: L_begin: " << L_begin << ", L_end_chunked: " << L_end_chunked << ", sliding_l_begin: " << sliding_l_begin << std::endl; + std::cout << "DEBUG: in_memory: " << in_memory << std::endl; + std::cout << "DEBUG: local_offset: " << local_offset << ", local_length: " << local_length << std::endl; + ) + // | sliding_l_begin -> in_memory | in_memory -> L_begin | L_begin -> L_end | L_end -> L_end_chunked | + // | Part 1 | Part 2 | Part 3 | Part 4 | + + int part_1_length = in_memory - sliding_l_begin; + int part_2_length = L_begin - in_memory; + int part_3_length = L_effective; + int part_4_length = L_end_chunked - L_begin - L_effective; + + bf16* prefill_k_cache_ptr = prefill_cache.data() + sliding_l_begin * DK; + bf16* prefill_v_cache_ptr = prefill_cache.data() + sliding_l_begin * DV + (size_t)MAX_L * DK; + bf16* decoding_k_cache_ptr = decoding_cache.data(); + bf16* decoding_v_cache_ptr = decoding_cache.data() + SLIDING_LENGTH * DK; + + // copy the existing cache up to L_begin from decoding_cache to prefill_cache, and then copy the new k cache from buffer_k and v cache from buffer_v to prefill_cache at the appropriate location, then sync the prefill_cache to device for the attention run + decoding_cache.sync_from_device(); + prefill_cache.sync_from_device(); + + // part 1: padded, no actual data exists + if (part_1_length > 0){ + memset(prefill_k_cache_ptr, 0, part_1_length * DK * sizeof(bf16)); + memset(prefill_v_cache_ptr, 0, part_1_length * DV * sizeof(bf16)); + prefill_k_cache_ptr += part_1_length * DK; + prefill_v_cache_ptr += part_1_length * DV; + } + + // part 2: copy from the existing cache in decoding_cache, but we need to copy in the order of the ring buffer + for (int i = 0; i < part_2_length; i++){ + int ring_idx = linear2ring(in_memory + i); + memcpy(prefill_k_cache_ptr, decoding_k_cache_ptr + ring_idx * DK, DK * sizeof(bf16)); + memcpy(prefill_v_cache_ptr, decoding_v_cache_ptr + ring_idx * DV, DV * sizeof(bf16)); + prefill_k_cache_ptr += DK; + prefill_v_cache_ptr += DV; + } + + // copy the new k cache and v cache to the appropriate location in prefill_cache and decoding_cache + bf16* new_k_cache_ptr = buffer_k.data() + L_offset * DK; + bf16* new_v_cache_ptr = buffer_v.data() + L_offset * DV; + //part 3: copy the new k cache and v cache from buffer_k and buffer_v to prefill_cache and decoding_cache, but we also need to copy in the order of the ring buffer + for (int i = 0; i < part_3_length; i++){ + int ring_idx = linear2ring(L_begin + i); + memcpy(prefill_k_cache_ptr, new_k_cache_ptr + i * DK, DK * sizeof(bf16)); + memcpy(prefill_v_cache_ptr, new_v_cache_ptr + i * DV, DV * sizeof(bf16)); + memcpy(decoding_k_cache_ptr + ring_idx * DK, new_k_cache_ptr + i * DK, DK * sizeof(bf16)); + memcpy(decoding_v_cache_ptr + ring_idx * DV, new_v_cache_ptr + i * DV, DV * sizeof(bf16)); + prefill_k_cache_ptr += DK; + prefill_v_cache_ptr += DV; + } + + // // part 4: padded, no actual data exists + // memset(prefill_k_cache_ptr, 0, part_4_length * DK * sizeof(bf16)); + // memset(prefill_v_cache_ptr, 0, part_4_length * DV * sizeof(bf16)); + decoding_cache.sync_to_device(); + prefill_cache.sync_to_device(); +} + +// --------------------------------------------------------------------------- +// mlp block +// --------------------------------------------------------------------------- + +void gemma4e_mlp_prefill_context::forward( + const gemma4e_prefill_shape& s, + bool double_wide, + buffer& gate_weights, + buffer& up_weights, + buffer& down_weights, + SafeTensors* reference +){ + // a double-wide layer runs the same three gemms over twice the mlp width + npu_app& gate = double_wide ? this->gate_skip_proj : this->gate_proj; + npu_app& up = double_wide ? this->up_skip_proj : this->up_proj; + npu_app& down = double_wide ? this->down_skip_proj : this->down_proj; + const int I = double_wide ? desc->INTERMEDIATE_SIZE * 2 : desc->INTERMEDIATE_SIZE; + + bufs->hidden_state_buffer.sync_to_device(); + FLM_OVERRIDE(gate_proj, gate(bufs->gate_buffer, bufs->hidden_state_buffer, gate_weights), this->desc, double_wide, s); + bufs->gate_buffer.sync_from_device(); + FLM_OVERRIDE(up_proj, up(bufs->up_buffer, bufs->hidden_state_buffer, up_weights), this->desc, double_wide, s); + bufs->up_buffer.sync_from_device(); + + gemma4e_cpu_func::_elementwise_mul_batch(bufs->hid_buffer.data(), bufs->gate_buffer.data(), bufs->up_buffer.data(), I, s.L_effective, s.L_offset, s.L_offset, s.L_offset); + + bufs->hid_buffer.sync_to_device(); + + FLM_OVERRIDE(down_proj, down(bufs->hidden_state_buffer, bufs->hid_buffer, down_weights), this->desc, double_wide, s); + bufs->hidden_state_buffer.sync_from_device(); + + DEBUG_BLOCK(2, + buffer down_ref; + reference->load_weights(down_ref, "mlp_output"); + buffer down_valid = buffer(bufs->hidden_state_buffer.data(), s.L_effective * desc->D); + buffer down_valid_ref = buffer(down_ref.data(), s.L_effective * desc->D); + print_error_metrics(get_error_metrics(down_valid, down_valid_ref), "MLP Error: "); + ) +} + +// --------------------------------------------------------------------------- +// per layer input path +// --------------------------------------------------------------------------- + +void gemma4e_pli_prefill_context::setup(const gemma4e_prefill_shape& s){ + if (L_padded_512_old == s.L_padded_512) { + return; + } + const int D = desc->D; + const int PLI_D = desc->PLI_D; + const int num_hidden_layers = desc->num_hidden_layers; + Gemma4e_ImageEncoder* enc = this->image_encoder; + + generate_mm_sequence(*this->pli_down_proj.seq(), + s.L_padded_512, D, num_hidden_layers * PLI_D, + enc->MM_tile_M, enc->MM_tile_K, enc->MM_tile_N, + 8,8,8, + enc->rtp_address, enc->rtp_sync_lock_id, + enc->MM_ROW_SIZE, enc->MM_COL_SIZE, + 0,0,0, + enc->IS_B_ROW_MAJOR, enc->ENABLE_AXI4, true, + false, 0,// no activation + 0, -10000.0, 1000000.0, // do not clamp on output + false, 0 // no need to reorder it + ); + + generate_mm_sequence(*this->pli_gate_proj.seq(), + s.L_padded_512, D, PLI_D, + enc->MM_tile_M, enc->MM_tile_K, enc->MM_tile_N, + 8,8,8, + enc->rtp_address, enc->rtp_sync_lock_id, + enc->MM_ROW_SIZE, enc->MM_COL_SIZE, + 0,0,0, + enc->IS_B_ROW_MAJOR, enc->ENABLE_AXI4, true, + false, 1,// no activation + 0, -10000.0, 1000000.0, // do not clamp on output + false, 0 // no need to reorder it + ); + + generate_mm_sequence(*this->pli_up_proj.seq(), + s.L_padded_512, PLI_D, D, + enc->MM_tile_M, enc->MM_tile_K, enc->MM_tile_N, + 8,8,8, + enc->rtp_address, enc->rtp_sync_lock_id, + enc->MM_ROW_SIZE, enc->MM_COL_SIZE, + 0,PLI_D * D,0, + enc->IS_B_ROW_MAJOR, enc->ENABLE_AXI4, true, + false, 0,// no activation + 0, -10000.0, 1000000.0, // do not clamp on output + false, 0 // no need to reorder it + ); + L_padded_512_old = s.L_padded_512; +} + +void gemma4e_pli_prefill_context::pre_pass( + const gemma4e_prefill_shape& s, + buffer& pli_down_weights, + buffer& pli_input_norm +){ + const int D = desc->D; + const int PLI_D = desc->PLI_D; + const int num_hidden_layers = desc->num_hidden_layers; + + // the down projection reads the token embeddings, which start life in the residual + memcpy(bufs->hidden_state_buffer.data() + (size_t)s.L_offset * D, bufs->residual_buffer.data() + (size_t)s.L_offset * D, (size_t)s.L_effective * D * sizeof(bf16)); + bufs->hidden_state_buffer.sync_to_device(); + FLM_OVERRIDE(pli_down_proj, this->pli_down_proj(bufs->hidden_state_buffer, pli_down_weights, bufs->pli_down_buffer), this->desc, s); + bufs->pli_down_buffer.sync_from_device(); + + gemma4e_cpu_func::_elementwise_scale_batch(bufs->pli_down_buffer.data(), 1.0 / sqrtf((float)D), PLI_D * num_hidden_layers, s.L_effective, s.L_offset); + gemma4e_cpu_func::_rms_norm_batch(bufs->pli_down_buffer.data(), bufs->pli_down_buffer.data(), pli_input_norm.data(), PLI_D, num_hidden_layers * PLI_D, s.L_effective, s.L_offset, s.L_offset); + gemma4e_cpu_func::_residual_add_batch(bufs->pli_embed_buffer.data(), bufs->pli_down_buffer.data(), bufs->pli_embed_buffer.data(), num_hidden_layers * PLI_D, s.L_effective, s.L_offset, s.L_offset, s.L_offset); + + gemma4e_cpu_func::_elementwise_scale_batch(bufs->pli_embed_buffer.data(), 1.0 / sqrtf((float)2.0), PLI_D * num_hidden_layers, s.L_effective, s.L_offset); +} + +void gemma4e_pli_prefill_context::layer_pass( + const gemma4e_prefill_shape& s, + int layer_idx, + buffer& pli_gate_up_weights, + buffer& pli_final_norm +){ + const int D = desc->D; + const int PLI_D = desc->PLI_D; + const int num_hidden_layers = desc->num_hidden_layers; + + FLM_OVERRIDE(pli_gate_proj, this->pli_gate_proj(bufs->hidden_state_buffer, pli_gate_up_weights, bufs->pli_gate_buffer), this->desc, layer_idx, s); + bufs->pli_gate_buffer.sync_from_device(); + + // this layer's slice of the per-layer embeddings gates the projection + gemma4e_cpu_func::_elementwise_mul_batch(bufs->pli_hid_buffer.data(), bufs->pli_gate_buffer.data(), bufs->pli_embed_buffer.data() + (size_t)layer_idx * PLI_D, PLI_D, num_hidden_layers * PLI_D, s.L_effective, s.L_offset, s.L_offset, s.L_offset); + + bufs->pli_hid_buffer.sync_to_device(); + FLM_OVERRIDE(pli_up_proj, this->pli_up_proj(bufs->pli_hid_buffer, pli_gate_up_weights, bufs->hidden_state_buffer), this->desc, layer_idx, s); + bufs->hidden_state_buffer.sync_from_device(); + + gemma4e_cpu_func::_rms_norm_batch(bufs->hidden_state_buffer.data(), bufs->hidden_state_buffer.data(), pli_final_norm.data(), D, D, s.L_effective, s.L_offset, s.L_offset); + bufs->hidden_state_buffer.sync_to_device(); +} + +// --------------------------------------------------------------------------- +// one decoder layer +// --------------------------------------------------------------------------- + +void gemma4e_prefill_context::forward( + int layer_idx, + gemma4e_layer_type_t type, + const gemma4e_prefill_shape& s, + buffer& proj_weights, + buffer& rms_weights, + buffer& rope_rms_weights, + buffer& pli_gate_up_weights, + buffer& kv_cache, + float layer_scale, + SafeTensors* reference +){ + const int D = desc->D; + const bool is_skip = is_skip_layer(type); + + FLM_OVERRIDE(prefill_layer_begin, (void)0, this->desc, layer_idx, type, s); + + this->dequant_block->run(type, proj_weights); + + DEBUG_BLOCK(1, + header_print_r("info", "Running layer " + std::to_string(layer_idx) + " of type " + std::to_string(type)); + ) + + // every norm weight of this layer, located through the descriptor + buffer input_layernorm_weight(rms_weights.data(), D); + buffer post_attention_layernorm_weight(rms_weights.data() + D, D); + buffer pre_feedforward_layernorm_weight(rms_weights.data() + D * 2, D); + buffer post_feedforward_layernorm_weight(rms_weights.data() + D * 3, D); + buffer pli_final_norm(rope_rms_weights.data() + desc->get_post_pli_norm_offset(type), D); + buffer q_norm(rope_rms_weights.data() + desc->get_q_norm_offset(type), desc->get_DH(type)); + buffer k_norm(rope_rms_weights.data() + desc->get_k_norm_offset(type), desc->get_DH(type)); + + // input layernorm + gemma4e_cpu_func::_rms_norm_batch(bufs.hidden_state_buffer.data(), bufs.residual_buffer.data(), input_layernorm_weight.data(), D, D, s.L_effective, s.L_offset, s.L_offset); + bufs.hidden_state_buffer.sync_to_device(); + + DEBUG_BLOCK(2, + buffer norm_ref; + reference->load_weights(norm_ref, "input_layernorm_output"); + buffer norm_valid = buffer(bufs.hidden_state_buffer.data(), s.L_effective * D); + buffer norm_valid_ref = buffer(norm_ref.data(), s.L_effective * D); + print_error_metrics(get_error_metrics(norm_valid, norm_valid_ref), "Input LayerNorm Error: "); + ) + + this->attn_block->forward(s, type, + this->dequant_block->qkv_weights, this->dequant_block->o_weights, + kv_cache, q_norm, k_norm, reference); + + gemma4e_cpu_func::_rms_norm_batch(bufs.hidden_state_buffer.data(), bufs.hidden_state_buffer.data(), post_attention_layernorm_weight.data(), D, D, s.L_effective, s.L_offset, s.L_offset); + bufs.attn_out_buffer.sync_to_device(); + + DEBUG_BLOCK(2, + buffer post_attn_norm_ref; + reference->load_weights(post_attn_norm_ref, "post_attention_norm_output"); + buffer post_attn_norm_valid = buffer(bufs.hidden_state_buffer.data(), s.L_effective * D); + buffer post_attn_norm_valid_ref = buffer(post_attn_norm_ref.data(), s.L_effective * D); + print_error_metrics(get_error_metrics(post_attn_norm_valid, post_attn_norm_valid_ref), "Post Attention LayerNorm Error: "); + ) + + gemma4e_cpu_func::_residual_add_batch(bufs.residual_buffer.data(), bufs.hidden_state_buffer.data(), bufs.residual_buffer.data(), D, s.L_effective, s.L_offset, s.L_offset, s.L_offset); + + gemma4e_cpu_func::_rms_norm_batch(bufs.hidden_state_buffer.data(), bufs.residual_buffer.data(), pre_feedforward_layernorm_weight.data(), D, D, s.L_effective, s.L_offset, s.L_offset); + bufs.hidden_state_buffer.sync_to_device(); + + DEBUG_BLOCK(2, + buffer pre_ffn_norm_ref; + reference->load_weights(pre_ffn_norm_ref, "pre_ffn_norm_output"); + buffer pre_ffn_norm_valid = buffer(bufs.hidden_state_buffer.data(), s.L_effective * D); + buffer pre_ffn_norm_valid_ref = buffer(pre_ffn_norm_ref.data(), s.L_effective * D); + print_error_metrics(get_error_metrics(pre_ffn_norm_valid, pre_ffn_norm_valid_ref), "Pre-FFN LayerNorm Error: "); + ) + + this->mlp->forward(s, is_skip && desc->enable_double_wide_mlp, + this->dequant_block->gate_weights, this->dequant_block->up_weights, + this->dequant_block->down_weights, reference); + + gemma4e_cpu_func::_rms_norm_batch(bufs.hidden_state_buffer.data(), bufs.hidden_state_buffer.data(), post_feedforward_layernorm_weight.data(), D, D, s.L_effective, s.L_offset, s.L_offset); + bufs.hidden_state_buffer.sync_to_device(); + + DEBUG_BLOCK(2, + buffer post_ffn_norm_ref; + reference->load_weights(post_ffn_norm_ref, "post_ffn_norm_output"); + buffer post_ffn_norm_valid = buffer(bufs.hidden_state_buffer.data(), s.L_effective * D); + buffer post_ffn_norm_valid_ref = buffer(post_ffn_norm_ref.data(), s.L_effective * D); + print_error_metrics(get_error_metrics(post_ffn_norm_valid, post_ffn_norm_valid_ref), "Post-FFN LayerNorm Error: "); + ) + + gemma4e_cpu_func::_residual_add_batch(bufs.residual_buffer.data(), bufs.hidden_state_buffer.data(), bufs.residual_buffer.data(), D, s.L_effective, s.L_offset, s.L_offset, s.L_offset); + + // per layer input: the gate reads the layer output, which lives in the residual + memcpy(bufs.hidden_state_buffer.data() + (size_t)s.L_offset * D, bufs.residual_buffer.data() + (size_t)s.L_offset * D, (size_t)s.L_effective * D * sizeof(bf16)); + bufs.hidden_state_buffer.sync_to_device(); + + this->pli->layer_pass(s, layer_idx, pli_gate_up_weights, pli_final_norm); + + gemma4e_cpu_func::_residual_add_batch(bufs.residual_buffer.data(), bufs.hidden_state_buffer.data(), bufs.residual_buffer.data(), D, s.L_effective, s.L_offset, s.L_offset, s.L_offset); + + gemma4e_cpu_func::_elementwise_scale_batch(bufs.residual_buffer.data(), layer_scale, D, s.L_effective, s.L_offset); + + DEBUG_BLOCK(2, + buffer output_ref; + reference->load_weights(output_ref, "layer_" + std::to_string(layer_idx)); + buffer output_valid = buffer(bufs.residual_buffer.data(), s.L_effective * D); + buffer output_valid_ref = buffer(output_ref.data(), s.L_effective * D); + print_error_metrics(get_error_metrics(output_valid, output_valid_ref), "Main path Error: "); + ) +} diff --git a/src/detail/gemma4e_npu/gemma4e_prefill.hpp b/src/detail/gemma4e_npu/gemma4e_prefill.hpp new file mode 100644 index 000000000..9f42f196a --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_prefill.hpp @@ -0,0 +1,521 @@ +#include "flm_override.hpp" +#ifndef __GEMMA4E_PREFILL_HPP__ +#define __GEMMA4E_PREFILL_HPP__ +#include +#include "modules/gemm.hpp" +#include "tensor_2d.hpp" +#include "gemma4e_npu_def.hpp" +#include "gemma4e_npu_sequence.hpp" +#include "gemma4e_cpu_functions.hpp" +#include "gemma4e_image.hpp" +#include "mmRuntimeSequence.hpp" +#include "utils/error_measure.hpp" + +/// @brief Row geometry of one prefill call. +/// +/// The attention kernel consumes whole chunks, so the batch is padded in front +/// (L_offset rows that belong to already-cached tokens) and at the back (up to +/// L_padded rows), and every gemm runs over the full L_padded rows. The per +/// layer input path runs over a coarser 512-row grid, hence L_padded_512. +struct gemma4e_prefill_shape { + int L_in = 0; ///< tokens handed to this call + int L_begin = 0; ///< context length before this call + int L_end = 0; ///< context length after this call + int L_effective = 0; ///< rows that carry a real token, == L_in + int L_offset = 0; ///< rows of chunk padding in front of the first new token + int L_begin_chunked = 0; ///< chunk-aligned start fed to the attention kernel + int L_end_chunked = 0; ///< chunk-aligned end fed to the attention kernel + int L_padded = 0; ///< rows every gemm runs over + int L_padded_512 = 0; ///< L_padded rounded up for the per-layer-input gemms + int sliding_l_begin = 0; ///< first position the sliding window still covers +}; + +/// @brief Buffers shared by every block of a prefill layer. +/// +/// Allocated once per batch geometry and reused across prefill calls; only the +/// residual and the per-layer-input embeddings need clearing, since every other +/// buffer is fully overwritten by the kernel that reads it. +struct gemma4e_common_buffers { + buffer residual_buffer; ///< host: the running residual stream + buffer hidden_state_buffer; ///< device: block input and block output + buffer q_buffer; + buffer k_buffer; + buffer v_buffer; + buffer attn_out_buffer; + buffer gate_buffer; + buffer up_buffer; + buffer hid_buffer; + + // per layer input path + buffer pli_embed_buffer; + buffer pli_down_buffer; + buffer pli_gate_buffer; + buffer pli_hid_buffer; + + tensor_2d tensor_residual; + tensor_2d tensor_pli_embed; + + gemma4e_common_buffers() {} + + /// @brief (Re)allocates every buffer for a batch geometry. + void allocate(npu_xclbin_manager* npu, gemma4e_desc* desc, const gemma4e_prefill_shape& s) { + const size_t D = desc->D; + const size_t I = desc->INTERMEDIATE_SIZE; + const size_t PLI = (size_t)desc->PLI_D * desc->num_hidden_layers; + residual_buffer = buffer((size_t)s.L_padded_512 * D); + hidden_state_buffer = npu->create_bo_buffer((size_t)s.L_padded_512 * D); + q_buffer = npu->create_bo_buffer((size_t)s.L_padded * desc->DQ); + k_buffer = npu->create_bo_buffer((size_t)s.L_padded * desc->DK); + v_buffer = npu->create_bo_buffer((size_t)s.L_padded * desc->DV); + attn_out_buffer = npu->create_bo_buffer((size_t)s.L_padded_512 * desc->DQ); + // the widest mlp a layer can ask for, so skip layers share these buffers + gate_buffer = npu->create_bo_buffer((size_t)s.L_padded * I * 2); + up_buffer = npu->create_bo_buffer((size_t)s.L_padded * I * 2); + hid_buffer = npu->create_bo_buffer((size_t)s.L_padded * I * 2); + pli_embed_buffer = npu->create_bo_buffer((size_t)s.L_padded_512 * PLI); + pli_down_buffer = npu->create_bo_buffer((size_t)s.L_padded_512 * PLI); + pli_gate_buffer = npu->create_bo_buffer((size_t)s.L_padded_512 * desc->PLI_D); + pli_hid_buffer = npu->create_bo_buffer((size_t)s.L_padded_512 * desc->PLI_D); + } + + /// @brief Clears the padding rows and re-points the row views at this batch. + void reset(gemma4e_desc* desc, const gemma4e_prefill_shape& s) { + // rows outside [L_offset, L_offset + L_in) are padding; they still take + // part in every gemm, so leave no stale values in them. + memset(residual_buffer.data(), 0, residual_buffer.size() * sizeof(bf16)); + memset(pli_embed_buffer.data(), 0, pli_embed_buffer.size() * sizeof(bf16)); + tensor_residual.assign(residual_buffer, desc->D, s.L_offset); + tensor_pli_embed.assign(pli_embed_buffer, desc->PLI_D * desc->num_hidden_layers, s.L_offset); + } +}; + +/// @brief Dequantizes one layer's projections into the bf16 buffers the prefill +/// gemms consume. +/// +/// One sequence per (layer kind, weight): which bytes each one reads is entirely +/// the descriptor's business, so changing the quantization type needs no edit +/// here. The five matrices of a layer are launched back to back. +struct gemma4e_dequant_prefill_context { + /// index into the per-layer-kind app table + enum matrix_t { QKV = 0, O = 1, GATE = 2, UP = 3, DOWN = 4, NUM_MATRICES = 5 }; + + gemma4e_desc* desc; + npu_app apps[4][NUM_MATRICES]; ///< [gemma4e_layer_type_t][matrix_t] + + buffer qkv_weights; + buffer o_weights; + buffer gate_weights; + buffer up_weights; + buffer down_weights; + + gemma4e_dequant_prefill_context( + gemma4e_desc* desc, + gemma4e_npu_sequence* seq_gen, + npu_app_manager* dequant_app_manager + ) : desc(desc) + { + const uint32_t D = desc->D; + const uint32_t I = desc->get_max_intermediate_size(); + this->qkv_weights = dequant_app_manager->create_bo_buffer((size_t)D * (desc->DQ + desc->DK + desc->DV)); + this->o_weights = dequant_app_manager->create_bo_buffer((size_t)desc->DQ * D); + this->gate_weights = dequant_app_manager->create_bo_buffer((size_t)D * I); + this->up_weights = dequant_app_manager->create_bo_buffer((size_t)D * I); + this->down_weights = dequant_app_manager->create_bo_buffer((size_t)I * D); + + for (int t = 0; t < 4; t++) { + gemma4e_layer_type_t type = static_cast(t); + gemma4e_layer_weight_def& W = desc->weight_desc(type); + for (int m = 0; m < NUM_MATRICES; m++) { + apps[t][m] = dequant_app_manager->create_app(); + } + // up and gate share one interleaved band and are told apart by the mode + seq_gen->generate_dequant_seq(apps[t][QKV].seq(), W.attn_qkv, gemma4e_npu_sequence::NORMAL_DEQUANT); + seq_gen->generate_dequant_seq(apps[t][O].seq(), W.attn_output, gemma4e_npu_sequence::NORMAL_DEQUANT); + seq_gen->generate_dequant_seq(apps[t][GATE].seq(), W.ffn_gate, gemma4e_npu_sequence::GATE_MATRIX); + seq_gen->generate_dequant_seq(apps[t][UP].seq(), W.ffn_up, gemma4e_npu_sequence::UP_MATRIX); + seq_gen->generate_dequant_seq(apps[t][DOWN].seq(), W.ffn_down, gemma4e_npu_sequence::NORMAL_DEQUANT); + } + } + + /// @brief Dequantizes every projection of one layer, in place. + void run(gemma4e_layer_type_t type, buffer& proj_weights) { + npu_app* a = apps[int(type)]; + FLM_OVERRIDE(dequant_qkv, a[QKV](this->qkv_weights, proj_weights), this->desc, type); + FLM_OVERRIDE(dequant_o, a[O](this->o_weights, proj_weights), this->desc, type); + FLM_OVERRIDE(dequant_up, a[UP](this->up_weights, proj_weights), this->desc, type); + FLM_OVERRIDE(dequant_gate, a[GATE](this->gate_weights, proj_weights), this->desc, type); + FLM_OVERRIDE(dequant_down, a[DOWN](this->down_weights, proj_weights), this->desc, type); + } +}; + +/// @brief Attention half of one prefill layer: q/k/v projections, rope, kv cache +/// fill, attention, output projection. +/// +/// Sliding and global layers run different kernels over differently shaped +/// caches, so both sets of apps live here and forward() dispatches on the layer +/// kind. Skip layers reuse the previous layer's cache and project q only. +struct gemma4e_attn_block_prefill_context { + gemma4e_desc* desc; + Gemm* gemm_seq_gen; + gemma4e_npu_sequence* seq_gen; + gemma4e_common_buffers* bufs; + + int L_padded_old = -1; + int L_begin_chunked_old = -1; + int L_end_chunked_old = -1; + uint32_t MAX_L = 0; + + npu_app q_swa_proj, k_swa_proj, v_swa_proj, o_swa_proj; + npu_app q_global_proj, k_global_proj, v_global_proj, o_global_proj; + npu_app mha_engine; + npu_app swa_engine; + + /// caches the attention kernels read: contiguous, and rebuilt every layer + buffer kv_cache_sliding_prefill; + buffer kv_cache_global_prefill; + + npu_app_manager* mha_app_manager; + npu_app_manager* swa_app_manager; + + gemma4e_attn_block_prefill_context( + gemma4e_desc* desc, + Gemm* gemm_seq_gen, + gemma4e_npu_sequence* seq_gen, + gemma4e_common_buffers* bufs, + npu_app_manager* gemm_app_manager, + npu_app_manager* mha_app_manager, + npu_app_manager* swa_app_manager + ) : desc(desc), gemm_seq_gen(gemm_seq_gen), seq_gen(seq_gen), bufs(bufs), + mha_app_manager(mha_app_manager), swa_app_manager(swa_app_manager) + { + this->q_global_proj = gemm_app_manager->create_app(); + this->k_global_proj = gemm_app_manager->create_app(); + this->v_global_proj = gemm_app_manager->create_app(); + this->o_global_proj = gemm_app_manager->create_app(); + this->q_swa_proj = gemm_app_manager->create_app(); + this->k_swa_proj = gemm_app_manager->create_app(); + this->v_swa_proj = gemm_app_manager->create_app(); + this->o_swa_proj = gemm_app_manager->create_app(); + this->mha_engine = mha_app_manager->create_app(); + this->swa_engine = swa_app_manager->create_app(); + } + + /// @brief Sizes the prefill-side caches, which span the whole context. + void set_max_length(uint32_t MAX_L) { + this->MAX_L = MAX_L; + this->kv_cache_sliding_prefill = swa_app_manager->create_bo_buffer((size_t)MAX_L * (desc->SWA_DK + desc->SWA_DV)); + this->kv_cache_global_prefill = mha_app_manager->create_bo_buffer((size_t)MAX_L * (desc->DK + desc->DV)); + this->kv_cache_sliding_prefill.memset((bf16)0); + this->kv_cache_global_prefill.memset((bf16)0); + this->kv_cache_sliding_prefill.sync_to_device(); + this->kv_cache_global_prefill.sync_to_device(); + // the cached sequences address these buffers, so force a regeneration + this->L_padded_old = -1; + this->L_begin_chunked_old = -1; + this->L_end_chunked_old = -1; + } + + /// @brief Regenerates the projection and attention sequences, if the geometry moved. + void setup(const gemma4e_prefill_shape& s) { + const uint32_t D = desc->D; + if (L_padded_old != s.L_padded) { + gemm_seq_gen->generate_seq(this->q_swa_proj.seq(), s.L_padded, D, desc->SWA_DQ, 0, + false, Gemm::NO_Activation, 0); + gemm_seq_gen->generate_seq(this->k_swa_proj.seq(), s.L_padded, D, desc->SWA_DK, D * desc->SWA_DQ, + false, Gemm::NO_Activation, 0); + gemm_seq_gen->generate_seq(this->v_swa_proj.seq(), s.L_padded, D, desc->SWA_DV, D * (desc->SWA_DQ + desc->SWA_DK), + false, Gemm::NO_Activation, 0); + gemm_seq_gen->generate_seq(this->o_swa_proj.seq(), s.L_padded, desc->SWA_DQ, D, 0, + false, Gemm::NO_Activation, 0); + + gemm_seq_gen->generate_seq(this->q_global_proj.seq(), s.L_padded, D, desc->DQ, 0, + false, Gemm::NO_Activation, 0); + gemm_seq_gen->generate_seq(this->k_global_proj.seq(), s.L_padded, D, desc->DK, D * desc->DQ, + false, Gemm::NO_Activation, 0); + gemm_seq_gen->generate_seq(this->v_global_proj.seq(), s.L_padded, D, desc->DV, D * (desc->DQ + desc->DK), + false, Gemm::NO_Activation, 0); + gemm_seq_gen->generate_seq(this->o_global_proj.seq(), s.L_padded, desc->DQ, D, 0, + false, Gemm::NO_Activation, 0); + L_padded_old = s.L_padded; + } + if (L_begin_chunked_old != s.L_begin_chunked || L_end_chunked_old != s.L_end_chunked) { + seq_gen->gen_mha_engine_seq(this->mha_engine.seq(), s.L_begin_chunked, s.L_end_chunked); + seq_gen->gen_swa_engine_seq(this->swa_engine.seq(), s.L_begin_chunked, s.L_end_chunked); + L_begin_chunked_old = s.L_begin_chunked; + L_end_chunked_old = s.L_end_chunked; + } + } + + /// @brief Runs the attention block of one layer, hidden_state_buffer in and out. + void forward( + const gemma4e_prefill_shape& s, + gemma4e_layer_type_t type, + buffer& qkv_weights, + buffer& o_weights, + buffer& kv_cache, + buffer& q_norm, + buffer& k_norm, + SafeTensors* reference + ) { + if (is_swa_layer(type)) { + _forward_swa(s, type, qkv_weights, o_weights, kv_cache, q_norm, k_norm, reference); + } + else { + _forward_global(s, type, qkv_weights, o_weights, kv_cache, q_norm, k_norm, reference); + } + } + + /// @brief Rebuilds the global prefill cache: everything up to L_begin from the + /// decode cache, then this batch's new rows, written to both caches. + void sync_kv_cache( + buffer& buffer_k, + buffer& buffer_v, + buffer& prefill_cache, + buffer& decoding_cache, + int L_offset, + int L_begin, + int L_effective, + int DK, int DV + ); + + /// @brief Rebuilds the sliding prefill cache, unrolling the decode ring buffer + /// into the linear order the swa kernel expects. + void sync_sliding_kv_cache( + buffer& buffer_k, + buffer& buffer_v, + buffer& prefill_cache, + buffer& decoding_cache, + int L_offset, + int L_begin, + int L_effective, + int sliding_l_begin, + int L_end_chunked, + int DK, int DV + ); + +private: + void _forward_swa(const gemma4e_prefill_shape& s, gemma4e_layer_type_t type, + buffer& qkv_weights, buffer& o_weights, buffer& kv_cache, + buffer& q_norm, buffer& k_norm, SafeTensors* reference); + void _forward_global(const gemma4e_prefill_shape& s, gemma4e_layer_type_t type, + buffer& qkv_weights, buffer& o_weights, buffer& kv_cache, + buffer& q_norm, buffer& k_norm, SafeTensors* reference); +}; + +/// @brief Feed-forward half of one prefill layer. +/// +/// Skip layers run a double-wide mlp when the model enables it, which is a +/// different sequence over the same buffers, so both widths are kept ready. +struct gemma4e_mlp_prefill_context { + gemma4e_desc* desc; + Gemm* gemm_seq_gen; + gemma4e_common_buffers* bufs; + + int L_padded_old = -1; + + npu_app gate_proj, up_proj, down_proj; + npu_app gate_skip_proj, up_skip_proj, down_skip_proj; + + gemma4e_mlp_prefill_context( + gemma4e_desc* desc, + Gemm* gemm_seq_gen, + gemma4e_common_buffers* bufs, + npu_app_manager* gemm_app_manager + ) : desc(desc), gemm_seq_gen(gemm_seq_gen), bufs(bufs) + { + this->gate_proj = gemm_app_manager->create_app(); + this->up_proj = gemm_app_manager->create_app(); + this->down_proj = gemm_app_manager->create_app(); + this->gate_skip_proj = gemm_app_manager->create_app(); + this->up_skip_proj = gemm_app_manager->create_app(); + this->down_skip_proj = gemm_app_manager->create_app(); + } + + void setup(const gemma4e_prefill_shape& s) { + if (L_padded_old == s.L_padded) { + return; + } + const uint32_t D = desc->D; + const uint32_t I = desc->INTERMEDIATE_SIZE; + gemm_seq_gen->generate_seq(this->gate_proj.seq(), s.L_padded, D, I, 0, + false, Gemm::GeLU, 0); + gemm_seq_gen->generate_seq(this->up_proj.seq(), s.L_padded, D, I, 0, + false, Gemm::NO_Activation, 0); + gemm_seq_gen->generate_seq(this->down_proj.seq(), s.L_padded, I, D, 0, + false, Gemm::NO_Activation, 0); + + gemm_seq_gen->generate_seq(this->gate_skip_proj.seq(), s.L_padded, D, I * 2, 0, + false, Gemm::GeLU, 0); + gemm_seq_gen->generate_seq(this->up_skip_proj.seq(), s.L_padded, D, I * 2, 0, + false, Gemm::NO_Activation, 0); + gemm_seq_gen->generate_seq(this->down_skip_proj.seq(), s.L_padded, I * 2, D, 0, + false, Gemm::NO_Activation, 0); + L_padded_old = s.L_padded; + } + + /// @brief Runs the mlp of one layer, hidden_state_buffer in and out. + void forward( + const gemma4e_prefill_shape& s, + bool double_wide, + buffer& gate_weights, + buffer& up_weights, + buffer& down_weights, + SafeTensors* reference + ); +}; + +/// @brief The per-layer-input path: a model-wide down projection run once per +/// batch, and a gate/up pair run at the end of every layer. +struct gemma4e_pli_prefill_context { + gemma4e_desc* desc; + gemma4e_common_buffers* bufs; + Gemma4e_ImageEncoder* image_encoder; ///< owns the mm tiling the pli gemms reuse + + int L_padded_512_old = -1; + + npu_app pli_down_proj, pli_gate_proj, pli_up_proj; + + gemma4e_pli_prefill_context( + gemma4e_desc* desc, + gemma4e_common_buffers* bufs, + Gemma4e_ImageEncoder* image_encoder + ) : desc(desc), bufs(bufs), image_encoder(image_encoder) + { + this->pli_down_proj = image_encoder->proj->create_app(); + this->pli_gate_proj = image_encoder->proj->create_app(); + this->pli_up_proj = image_encoder->proj->create_app(); + } + + void setup(const gemma4e_prefill_shape& s); + + /// @brief Projects the token embeddings down into the per-layer embeddings and + /// folds them into the per-layer-input stream. Runs once per batch. + void pre_pass(const gemma4e_prefill_shape& s, buffer& pli_down_weights, buffer& pli_input_norm); + + /// @brief Gates this layer's per-layer embedding and adds it back to the + /// residual, leaving the result in hidden_state_buffer. + void layer_pass(const gemma4e_prefill_shape& s, int layer_idx, + buffer& pli_gate_up_weights, buffer& pli_final_norm); +}; + +/// @brief One prefill pass over the model, one layer at a time. +/// +/// Owns the prefill xclbins, the sequence generators and every scratch buffer +/// the prefill path needs; the model only hands it weights and a kv cache. The +/// generated sequences are cached on the batch geometry, so a prefill that +/// repeats a shape regenerates nothing. +struct gemma4e_prefill_context { + gemma4e_desc* desc; + npu_xclbin_manager* npu; + gemma4e_common_buffers bufs; + + std::unique_ptr gemm_seq_gen; + std::unique_ptr dequant_block; + std::unique_ptr attn_block; + std::unique_ptr mlp; + std::unique_ptr pli; + + npu_app_manager* gemm_app_manager; + npu_app_manager* dequant_app_manager; + npu_app_manager* mha_app_manager; + npu_app_manager* swa_app_manager; + + /// @brief the attention kernel consumes whole chunks of this many rows + static constexpr int L_chunk = 16 * 8; + /// @brief the attention kernel refuses batches shorter than this + static constexpr int L_MIN = 256; + + /// @note The shared buffers are sized from both L_padded_512 (the residual-side + /// tensors) and L_padded (the q/k/v/gate/up scratch), and the two can move + /// independently, so reallocate when either changes. + int L_padded_512_old = -1; + int L_padded_old = -1; + + gemma4e_prefill_context( + npu_xclbin_manager* npu, + gemma4e_desc* desc, + LM_Config& config, + gemma4e_npu_sequence* seq_gen, + std::unique_ptr pli, + uint32_t MAX_L + ) : desc(desc), npu(npu) + { + this->gemm_app_manager = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "mm.xclbin")); + this->dequant_app_manager = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "dequant.xclbin")); + this->mha_app_manager = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "attn.xclbin")); + this->swa_app_manager = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "swa.xclbin")); + + this->gemm_seq_gen = std::make_unique(config); + + this->attn_block = std::make_unique( + desc, gemm_seq_gen.get(), seq_gen, &bufs, + gemm_app_manager, mha_app_manager, swa_app_manager + ); + this->mlp = std::make_unique( + desc, gemm_seq_gen.get(), &bufs, gemm_app_manager + ); + this->dequant_block = std::make_unique( + desc, seq_gen, dequant_app_manager + ); + // the per layer input apps must already exist by now: they are created on + // the image encoder's manager, which has to happen before layer.xclbin is + // registered, so the model builds this block and hands it over. + this->pli = std::move(pli); + this->pli->bufs = &bufs; + + this->attn_block->set_max_length(MAX_L); + } + + void set_max_length(uint32_t MAX_L) { this->attn_block->set_max_length(MAX_L); } + + /// @brief Computes the row geometry of a prefill call and readies every block. + /// @param L_in number of tokens to prefill + /// @param L_begin context length already in the kv cache + gemma4e_prefill_shape setup(int L_in, int L_begin) { + gemma4e_prefill_shape s; + s.L_in = L_in; + s.L_effective = L_in; + s.L_begin = L_begin; + s.L_end = L_begin + L_in; + s.L_begin_chunked = (L_begin / L_chunk) * L_chunk; + s.L_offset = L_begin - s.L_begin_chunked; + s.L_end_chunked = s.L_begin_chunked + (L_in + s.L_offset + L_MIN - 1) / L_MIN * L_MIN; + s.L_padded = s.L_end_chunked - s.L_begin_chunked; + s.L_padded_512 = ((s.L_padded + 511) / 512) * 512; + s.sliding_l_begin = (s.L_begin_chunked - (int)desc->SLIDING_LENGTH) > 0 + ? (s.L_begin_chunked - (int)desc->SLIDING_LENGTH) : 0; + + if (L_padded_512_old != s.L_padded_512 || L_padded_old != s.L_padded) { + bufs.allocate(npu, desc, s); + L_padded_512_old = s.L_padded_512; + L_padded_old = s.L_padded; + } + bufs.reset(desc, s); + + this->attn_block->setup(s); + this->mlp->setup(s); + this->pli->setup(s); + return s; + } + + /// @brief Row `i` of the residual stream, where the token embeddings are written. + buffer& residual_row(int i) { return bufs.tensor_residual[i]; } + /// @brief Row `i` of the per-layer-input embeddings. + buffer& pli_embed_row(int i) { return bufs.tensor_pli_embed[i]; } + + /// @brief Runs one decoder layer over the whole batch, in place on the residual. + void forward( + int layer_idx, + gemma4e_layer_type_t type, + const gemma4e_prefill_shape& s, + buffer& proj_weights, + buffer& rms_weights, + buffer& rope_rms_weights, + buffer& pli_gate_up_weights, + buffer& kv_cache, + float layer_scale, + SafeTensors* reference + ); +}; + +#endif // __GEMMA4E_PREFILL_HPP__ diff --git a/src/detail/gemma4e_npu/gemma4e_vision_prefill_helper.cpp b/src/detail/gemma4e_npu/gemma4e_vision_prefill_helper.cpp new file mode 100644 index 000000000..1008c5246 --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_vision_prefill_helper.cpp @@ -0,0 +1,1320 @@ +#include "gemma4e_vision_prefill_helper.hpp" + +#include "avx512_util.hpp" +#include +#include +#include +#include // std::bit_cast (C++20+) +#include +#include // OpenMP for multi-threading + +void simd_add( + bf16* input1, + bf16* input2, + bf16* output, + size_t size +){ + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 4; // Process 64 elements per iteration + constexpr size_t CHUNK_SIZE = SIMD_WIDTH * UNROLL_FACTOR; // 64 elements + + // Use OpenMP only if size is large enough to benefit from parallelization + // Threshold: 8 chunks (512 elements) per thread minimum to avoid overhead + constexpr size_t min_elements_per_thread = CHUNK_SIZE * 8; + + const bool use_parallel = size >= (min_elements_per_thread * 2); + + // Calculate number of full chunks for parallel processing + const size_t num_chunks = size / CHUNK_SIZE; + const int signed_num_chunks = static_cast(num_chunks); + + // Process full chunks with OpenMP (chunked distribution for cache locality) + #pragma omp parallel for num_threads(max_prefill_threads) if(use_parallel) schedule(static) + for (int chunk_idx = 0; chunk_idx < signed_num_chunks; ++chunk_idx) { + size_t i = static_cast(chunk_idx) * CHUNK_SIZE; + + // Prefetch next cache lines + _mm_prefetch(reinterpret_cast(input1 + i + 64), _MM_HINT_T0); + _mm_prefetch(reinterpret_cast(input2 + i + 64), _MM_HINT_T0); + + // Load all inputs first (better for out-of-order execution) + __m256i a_bh_vec0 = _mm256_loadu_si256(reinterpret_cast(input1 + i)); + __m256i b_bh_vec0 = _mm256_loadu_si256(reinterpret_cast(input2 + i)); + + __m256i a_bh_vec1 = _mm256_loadu_si256(reinterpret_cast(input1 + i + 16)); + __m256i b_bh_vec1 = _mm256_loadu_si256(reinterpret_cast(input2 + i + 16)); + + __m256i a_bh_vec2 = _mm256_loadu_si256(reinterpret_cast(input1 + i + 32)); + __m256i b_bh_vec2 = _mm256_loadu_si256(reinterpret_cast(input2 + i + 32)); + + __m256i a_bh_vec3 = _mm256_loadu_si256(reinterpret_cast(input1 + i + 48)); + __m256i b_bh_vec3 = _mm256_loadu_si256(reinterpret_cast(input2 + i + 48)); + + // Convert to 32-bit and shift (interleaved for better pipeline usage) + __m512i a_shifted0 = _mm512_slli_epi32(_mm512_cvtepu16_epi32(a_bh_vec0), 16); + __m512i b_shifted0 = _mm512_slli_epi32(_mm512_cvtepu16_epi32(b_bh_vec0), 16); + + __m512i a_shifted1 = _mm512_slli_epi32(_mm512_cvtepu16_epi32(a_bh_vec1), 16); + __m512i b_shifted1 = _mm512_slli_epi32(_mm512_cvtepu16_epi32(b_bh_vec1), 16); + + __m512i a_shifted2 = _mm512_slli_epi32(_mm512_cvtepu16_epi32(a_bh_vec2), 16); + __m512i b_shifted2 = _mm512_slli_epi32(_mm512_cvtepu16_epi32(b_bh_vec2), 16); + + __m512i a_shifted3 = _mm512_slli_epi32(_mm512_cvtepu16_epi32(a_bh_vec3), 16); + __m512i b_shifted3 = _mm512_slli_epi32(_mm512_cvtepu16_epi32(b_bh_vec3), 16); + + // Add as floats + __m512 sum0 = _mm512_add_ps(_mm512_castsi512_ps(a_shifted0), _mm512_castsi512_ps(b_shifted0)); + __m512 sum1 = _mm512_add_ps(_mm512_castsi512_ps(a_shifted1), _mm512_castsi512_ps(b_shifted1)); + __m512 sum2 = _mm512_add_ps(_mm512_castsi512_ps(a_shifted2), _mm512_castsi512_ps(b_shifted2)); + __m512 sum3 = _mm512_add_ps(_mm512_castsi512_ps(a_shifted3), _mm512_castsi512_ps(b_shifted3)); + + // Convert back to bfloat16 and store with round-to-nearest-even + store_m512_to_bfloat16_rne(output + i, sum0); + store_m512_to_bfloat16_rne(output + i + 16, sum1); + store_m512_to_bfloat16_rne(output + i + 32, sum2); + store_m512_to_bfloat16_rne(output + i + 48, sum3); + } + + // Process remaining elements (remainder after full chunks) + size_t i = num_chunks * CHUNK_SIZE; + + // Process remaining 16-element chunks + for (; i + SIMD_WIDTH <= size; i += SIMD_WIDTH) { + __m256i a_bh_vec = _mm256_loadu_si256(reinterpret_cast(input1 + i)); + __m256i b_bh_vec = _mm256_loadu_si256(reinterpret_cast(input2 + i)); + + __m512i a_shifted = _mm512_slli_epi32(_mm512_cvtepu16_epi32(a_bh_vec), 16); + __m512i b_shifted = _mm512_slli_epi32(_mm512_cvtepu16_epi32(b_bh_vec), 16); + + __m512 sum = _mm512_add_ps(_mm512_castsi512_ps(a_shifted), _mm512_castsi512_ps(b_shifted)); + + store_m512_to_bfloat16_rne(output + i, sum); + } + + // Scalar tail + for (; i < size; ++i) { + output[i] = static_cast(static_cast(input1[i]) + static_cast(input2[i])); + } +} + +void transpose_2d( + const bf16* input, + bf16* output, + size_t rows, + size_t cols +){ + int signed_rows = static_cast(rows); + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) if(rows >= 8) + for(int r = 0; r < signed_rows; r++){ + for(int c = 0; c < static_cast(cols); c++){ + output[c * rows + r] = input[r * cols + c]; + } + } +} + +void simd_add( + const float* input1, + const bf16* input2, + float* output, + size_t size +) { + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 4; + constexpr size_t CHUNK_SIZE = SIMD_WIDTH * UNROLL_FACTOR; + constexpr size_t min_elements_per_thread = CHUNK_SIZE * 8; + + const bool use_parallel = size >= (min_elements_per_thread * 2); + const size_t num_chunks = size / CHUNK_SIZE; + const int signed_num_chunks = static_cast(num_chunks); + + #pragma omp parallel for num_threads(max_prefill_threads) if(use_parallel) schedule(static) + for (int chunk_idx = 0; chunk_idx < signed_num_chunks; ++chunk_idx) { + size_t base = static_cast(chunk_idx) * CHUNK_SIZE; + + _mm_prefetch(reinterpret_cast(input1 + base + 64), _MM_HINT_T0); + _mm_prefetch(reinterpret_cast(input2 + base + 64), _MM_HINT_T0); + + for (size_t j = 0; j < CHUNK_SIZE; j += SIMD_WIDTH) { + size_t i = base + j; + __m512 a_ps_vec = _mm512_loadu_ps(input1 + i); + __m512 b_ps_vec = load_bfloat16_to_m512(input2 + i); + _mm512_storeu_ps(output + i, _mm512_add_ps(a_ps_vec, b_ps_vec)); + } + } + + // Remainder after full chunks + size_t i = num_chunks * CHUNK_SIZE; + for (; i + SIMD_WIDTH <= size; i += SIMD_WIDTH) { + __m512 a_ps_vec = _mm512_loadu_ps(input1 + i); + __m512 b_ps_vec = load_bfloat16_to_m512(input2 + i); + _mm512_storeu_ps(output + i, _mm512_add_ps(a_ps_vec, b_ps_vec)); + } + + // Scalar tail + for (; i < size; ++i) { + output[i] = input1[i] + static_cast(input2[i]); + } +} + +void simd_add( + const float* input1, + const bf16* input2, + bf16* output, + size_t size +) { + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 4; + constexpr size_t CHUNK_SIZE = SIMD_WIDTH * UNROLL_FACTOR; + constexpr size_t min_elements_per_thread = CHUNK_SIZE * 8; + + const bool use_parallel = size >= (min_elements_per_thread * 2); + const size_t num_chunks = size / CHUNK_SIZE; + const int signed_num_chunks = static_cast(num_chunks); + + #pragma omp parallel for num_threads(max_prefill_threads) if(use_parallel) schedule(static) + for (int chunk_idx = 0; chunk_idx < signed_num_chunks; ++chunk_idx) { + size_t base = static_cast(chunk_idx) * CHUNK_SIZE; + + _mm_prefetch(reinterpret_cast(input1 + base + 64), _MM_HINT_T0); + _mm_prefetch(reinterpret_cast(input2 + base + 64), _MM_HINT_T0); + + __m512 a_vec0 = _mm512_loadu_ps(input1 + base); + __m512 b_vec0 = load_bfloat16_to_m512(input2 + base); + __m512 a_vec1 = _mm512_loadu_ps(input1 + base + 16); + __m512 b_vec1 = load_bfloat16_to_m512(input2 + base + 16); + __m512 a_vec2 = _mm512_loadu_ps(input1 + base + 32); + __m512 b_vec2 = load_bfloat16_to_m512(input2 + base + 32); + __m512 a_vec3 = _mm512_loadu_ps(input1 + base + 48); + __m512 b_vec3 = load_bfloat16_to_m512(input2 + base + 48); + + store_m512_to_bfloat16_rne(output + base, _mm512_add_ps(a_vec0, b_vec0)); + store_m512_to_bfloat16_rne(output + base + 16, _mm512_add_ps(a_vec1, b_vec1)); + store_m512_to_bfloat16_rne(output + base + 32, _mm512_add_ps(a_vec2, b_vec2)); + store_m512_to_bfloat16_rne(output + base + 48, _mm512_add_ps(a_vec3, b_vec3)); + } + + size_t i = num_chunks * CHUNK_SIZE; + for (; i + SIMD_WIDTH <= size; i += SIMD_WIDTH) { + __m512 a_vec = _mm512_loadu_ps(input1 + i); + __m512 b_vec = load_bfloat16_to_m512(input2 + i); + store_m512_to_bfloat16_rne(output + i, _mm512_add_ps(a_vec, b_vec)); + } + + for (; i < size; ++i) { + float sum = input1[i] + static_cast(input2[i]); + output[i] = static_cast(sum); + } +} + +void simd_add( + const bf16* input1, + const bf16* input2, + float* output, + size_t size +) { + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 4; + constexpr size_t CHUNK_SIZE = SIMD_WIDTH * UNROLL_FACTOR; + constexpr size_t min_elements_per_thread = CHUNK_SIZE * 8; + + const bool use_parallel = size >= (min_elements_per_thread * 2); + const size_t num_chunks = size / CHUNK_SIZE; + const int signed_num_chunks = static_cast(num_chunks); + + #pragma omp parallel for num_threads(max_prefill_threads) if(use_parallel) schedule(static) + for (int chunk_idx = 0; chunk_idx < signed_num_chunks; ++chunk_idx) { + size_t base = static_cast(chunk_idx) * CHUNK_SIZE; + + _mm_prefetch(reinterpret_cast(input1 + base + 64), _MM_HINT_T0); + _mm_prefetch(reinterpret_cast(input2 + base + 64), _MM_HINT_T0); + + __m512 a_vec0 = load_bfloat16_to_m512(input1 + base); + __m512 b_vec0 = load_bfloat16_to_m512(input2 + base); + __m512 a_vec1 = load_bfloat16_to_m512(input1 + base + 16); + __m512 b_vec1 = load_bfloat16_to_m512(input2 + base + 16); + __m512 a_vec2 = load_bfloat16_to_m512(input1 + base + 32); + __m512 b_vec2 = load_bfloat16_to_m512(input2 + base + 32); + __m512 a_vec3 = load_bfloat16_to_m512(input1 + base + 48); + __m512 b_vec3 = load_bfloat16_to_m512(input2 + base + 48); + + _mm512_storeu_ps(output + base, _mm512_add_ps(a_vec0, b_vec0)); + _mm512_storeu_ps(output + base + 16, _mm512_add_ps(a_vec1, b_vec1)); + _mm512_storeu_ps(output + base + 32, _mm512_add_ps(a_vec2, b_vec2)); + _mm512_storeu_ps(output + base + 48, _mm512_add_ps(a_vec3, b_vec3)); + } + + size_t i = num_chunks * CHUNK_SIZE; + for (; i + SIMD_WIDTH <= size; i += SIMD_WIDTH) { + __m512 a_vec = load_bfloat16_to_m512(input1 + i); + __m512 b_vec = load_bfloat16_to_m512(input2 + i); + _mm512_storeu_ps(output + i, _mm512_add_ps(a_vec, b_vec)); + } + + for (; i < size; ++i) { + float sum = static_cast(input1[i]) + static_cast(input2[i]); + output[i] = sum; + } +} + +// Optimized AVX-512 version that combines bias addition and GELU activation +// This version processes the entire array in one pass, reducing memory traffic +// Note: bias is broadcast across all sequence positions (hidden_dim sized, repeated for seq_len) +// Optimizations: +// - OpenMP parallelization (4 threads) for sequence positions +// - Prefetching for cache optimization +// - 2x loop unrolling for better ILP (instruction-level parallelism) +void simd_bias_add_gelu( + bf16* input, + const bf16* bias, + bf16* output, + size_t total_size, + size_t hidden_dim +){ + size_t seq_len = total_size / hidden_dim; + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 2; + + // Use signed integer for OpenMP compatibility + int signed_seq_len = static_cast(seq_len); + + // Parallelize over sequence positions with OpenMP (max 4 threads) + // Each thread processes different sequence positions independently + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) + for (int seq = 0; seq < signed_seq_len; ++seq) { + size_t seq_offset = static_cast(seq) * hidden_dim; + size_t i = 0; + + // Process 32 elements at a time (2x unrolled for better ILP) + for (; i + SIMD_WIDTH * UNROLL_FACTOR <= hidden_dim; i += SIMD_WIDTH * UNROLL_FACTOR) { + // Prefetch next cache lines for input, bias, and output + _mm_prefetch(reinterpret_cast(input + seq_offset + i + 64), _MM_HINT_T0); + _mm_prefetch(reinterpret_cast(bias + i + 64), _MM_HINT_T0); + + // First iteration (elements i to i+15) + __m512 input_vec1 = load_bfloat16_to_m512(input + seq_offset + i); + __m512 bias_vec1 = load_bfloat16_to_m512(bias + i); + + // Second iteration (elements i+16 to i+31) - load in parallel to hide latency + __m512 input_vec2 = load_bfloat16_to_m512(input + seq_offset + i + SIMD_WIDTH); + __m512 bias_vec2 = load_bfloat16_to_m512(bias + i + SIMD_WIDTH); + + // Compute first iteration + __m512 sum_vec1 = _mm512_add_ps(input_vec1, bias_vec1); + + // Compute second iteration (parallel to first GELU computation) + __m512 sum_vec2 = _mm512_add_ps(input_vec2, bias_vec2); + + // Apply GELU activation (computationally expensive, so interleave) + __m512 gelu_vec1 = gelu_tanh_avx512_simd(sum_vec1); + __m512 gelu_vec2 = gelu_tanh_avx512_simd(sum_vec2); + + // Store results + store_m512_to_bfloat16_rne(output + seq_offset + i, gelu_vec1); + store_m512_to_bfloat16_rne(output + seq_offset + i + SIMD_WIDTH, gelu_vec2); + } + + // Process remaining 16-element chunks + for (; i + SIMD_WIDTH <= hidden_dim; i += SIMD_WIDTH) { + __m512 input_vec = load_bfloat16_to_m512(input + seq_offset + i); + __m512 bias_vec = load_bfloat16_to_m512(bias + i); + + __m512 sum_vec = _mm512_add_ps(input_vec, bias_vec); + __m512 gelu_vec = gelu_tanh_avx512_simd(sum_vec); + + store_m512_to_bfloat16_rne(output + seq_offset + i, gelu_vec); + } + + // Handle remaining elements with scalar loop + constexpr float sqrt_2_over_pi = 0.7978845608f; // √(2/π) + constexpr float coeff = 0.044715f; + + for (; i < hidden_dim; ++i) { + float x = static_cast(input[seq_offset + i]); + float b = static_cast(bias[i]); + float sum = x + b; + + // GELU approximation (tanh-based) + float x_cubed = sum * sum * sum; + float inner = sqrt_2_over_pi * (sum + coeff * x_cubed); + float tanh_val = std::tanh(inner); + float gelu = 0.5f * sum * (1.0f + tanh_val); + + output[seq_offset + i] = static_cast(gelu); + } + } +} + +void gelu_bfloat16_ref( + const bf16* input, + bf16* output, + size_t size +) { + size_t i = 0; + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 2; // Process 32 elements per iteration + + // Process 32 elements at a time (2x unrolled for GELU's computational intensity) + for (; i + SIMD_WIDTH * UNROLL_FACTOR <= size; i += SIMD_WIDTH * UNROLL_FACTOR) { + // Prefetch next cache line + _mm_prefetch(reinterpret_cast(input + i + 64), _MM_HINT_T0); + + // Load both sets of data + __m512 input_vec1 = load_bfloat16_to_m512(input + i); + __m512 input_vec2 = load_bfloat16_to_m512(input + i + SIMD_WIDTH); + + // Apply GELU activation to both + __m512 gelu_vec1 = gelu_tanh_avx512_simd(input_vec1); + __m512 gelu_vec2 = gelu_tanh_avx512_simd(input_vec2); + + // Store results + store_m512_to_bfloat16_rne(output + i, gelu_vec1); + store_m512_to_bfloat16_rne(output + i + SIMD_WIDTH, gelu_vec2); + } + + // Process remaining 16-element chunks + for (; i + SIMD_WIDTH <= size; i += SIMD_WIDTH) { + __m512 input_vec = load_bfloat16_to_m512(input + i); + __m512 gelu_vec = gelu_tanh_avx512_simd(input_vec); + store_m512_to_bfloat16_rne(output + i, gelu_vec); + } + + // Constants for the GELU approximation (tanh-based) + constexpr double sqrt_2_over_pi = 0.7978845608028654; // √(2/π) + constexpr double coeff = 0.044715; + + // Scalar tail + for (; i < size; ++i) { + double x = static_cast(input[i]); + double x_cubed = x * x * x; + double inner = sqrt_2_over_pi * (x + coeff * x_cubed); + double tanh_val = std::tanh(inner); + double gelu = 0.5 * x * (1.0 + tanh_val); + output[i] = static_cast(static_cast(gelu)); + } +} + +// RMS Norm without scale weights +// Input layout: [seq_len_padded x X_padded], processes seq_len rows, X elements per row +// Stride between rows is X_padded +void simd_rms_norm( + const bf16* input, + bf16* output, + size_t seq_len, + size_t X, + size_t seq_len_padded, + size_t X_padded, + float eps +) { + constexpr size_t SIMD_WIDTH = 16; + const float inv_X = 1.0f / static_cast(X); + + int signed_seq_len = static_cast(seq_len); + + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) if(seq_len >= 8) + for (int row = 0; row < signed_seq_len; ++row) { + const bf16* row_in = input + static_cast(row) * X_padded; + bf16* row_out = output + static_cast(row) * X_padded; + + // === Pass 1: compute sum of squares === + __m512 acc0 = _mm512_setzero_ps(); + __m512 acc1 = _mm512_setzero_ps(); + size_t i = 0; + + // 2x unrolled SIMD loop + for (; i + SIMD_WIDTH * 2 <= X; i += SIMD_WIDTH * 2) { + __m512 v0 = load_bfloat16_to_m512(row_in + i); + __m512 v1 = load_bfloat16_to_m512(row_in + i + SIMD_WIDTH); + acc0 = _mm512_fmadd_ps(v0, v0, acc0); + acc1 = _mm512_fmadd_ps(v1, v1, acc1); + } + for (; i + SIMD_WIDTH <= X; i += SIMD_WIDTH) { + __m512 v = load_bfloat16_to_m512(row_in + i); + acc0 = _mm512_fmadd_ps(v, v, acc0); + } + + float sum_sq = _mm512_reduce_add_ps(_mm512_add_ps(acc0, acc1)); + + // Scalar tail for sum of squares + for (; i < X; ++i) { + float val = static_cast(row_in[i]); + sum_sq += val * val; + } + + // mean_sq = sum_sq / X + eps, then rsqrt + float mean_sq = sum_sq * inv_X + eps; + float scale = 1.0f / std::sqrt(mean_sq); + __m512 scale_vec = _mm512_set1_ps(scale); + + // === Pass 2: multiply input by scale === + i = 0; + for (; i + SIMD_WIDTH * 2 <= X; i += SIMD_WIDTH * 2) { + __m512 v0 = load_bfloat16_to_m512(row_in + i); + __m512 v1 = load_bfloat16_to_m512(row_in + i + SIMD_WIDTH); + store_m512_to_bfloat16_rne(row_out + i, _mm512_mul_ps(v0, scale_vec)); + store_m512_to_bfloat16_rne(row_out + i + SIMD_WIDTH, _mm512_mul_ps(v1, scale_vec)); + } + for (; i + SIMD_WIDTH <= X; i += SIMD_WIDTH) { + __m512 v = load_bfloat16_to_m512(row_in + i); + store_m512_to_bfloat16_rne(row_out + i, _mm512_mul_ps(v, scale_vec)); + } + + // Scalar tail + for (; i < X; ++i) { + float val = static_cast(row_in[i]); + row_out[i] = static_cast(val * scale); + } + } +} + +// RMS Norm with scale weights (Gemma4RMSNorm with with_scale=True) +// Input layout: [seq_len_padded x X_padded], processes seq_len rows, X elements per row +// norm_weight: bf16 pointer of size [X], broadcast across all rows +void simd_rms_norm( + const bf16* input, + const bf16* norm_weight, + bf16* output, + size_t seq_len, + size_t X, + size_t seq_len_padded, + size_t X_padded, + float eps +) { + constexpr size_t SIMD_WIDTH = 16; + const float inv_X = 1.0f / static_cast(X); + + int signed_seq_len = static_cast(seq_len); + + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) if(seq_len >= 8) + for (int row = 0; row < signed_seq_len; ++row) { + const bf16* row_in = input + static_cast(row) * X_padded; + bf16* row_out = output + static_cast(row) * X_padded; + + // === Pass 1: compute sum of squares === + __m512 acc0 = _mm512_setzero_ps(); + __m512 acc1 = _mm512_setzero_ps(); + size_t i = 0; + + // 2x unrolled SIMD loop + for (; i + SIMD_WIDTH * 2 <= X; i += SIMD_WIDTH * 2) { + __m512 v0 = load_bfloat16_to_m512(row_in + i); + __m512 v1 = load_bfloat16_to_m512(row_in + i + SIMD_WIDTH); + acc0 = _mm512_fmadd_ps(v0, v0, acc0); + acc1 = _mm512_fmadd_ps(v1, v1, acc1); + } + for (; i + SIMD_WIDTH <= X; i += SIMD_WIDTH) { + __m512 v = load_bfloat16_to_m512(row_in + i); + acc0 = _mm512_fmadd_ps(v, v, acc0); + } + + float sum_sq = _mm512_reduce_add_ps(_mm512_add_ps(acc0, acc1)); + + // Scalar tail for sum of squares + for (; i < X; ++i) { + float val = static_cast(row_in[i]); + sum_sq += val * val; + } + + // mean_sq = sum_sq / X + eps, then rsqrt + float mean_sq = sum_sq * inv_X + eps; + float scale = 1.0f / std::sqrt(mean_sq); + __m512 scale_vec = _mm512_set1_ps(scale); + + // === Pass 2: multiply input by scale and norm_weight === + i = 0; + for (; i + SIMD_WIDTH * 2 <= X; i += SIMD_WIDTH * 2) { + __m512 v0 = load_bfloat16_to_m512(row_in + i); + __m512 w0 = load_bfloat16_to_m512(norm_weight + i); + __m512 v1 = load_bfloat16_to_m512(row_in + i + SIMD_WIDTH); + __m512 w1 = load_bfloat16_to_m512(norm_weight + i + SIMD_WIDTH); + store_m512_to_bfloat16_rne(row_out + i, _mm512_mul_ps(_mm512_mul_ps(v0, scale_vec), w0)); + store_m512_to_bfloat16_rne(row_out + i + SIMD_WIDTH, _mm512_mul_ps(_mm512_mul_ps(v1, scale_vec), w1)); + } + for (; i + SIMD_WIDTH <= X; i += SIMD_WIDTH) { + __m512 v = load_bfloat16_to_m512(row_in + i); + __m512 w = load_bfloat16_to_m512(norm_weight + i); + store_m512_to_bfloat16_rne(row_out + i, _mm512_mul_ps(_mm512_mul_ps(v, scale_vec), w)); + } + + // Scalar tail + for (; i < X; ++i) { + float val = static_cast(row_in[i]); + row_out[i] = static_cast(val * scale * static_cast(norm_weight[i])); + } + } +} + +// AVX-512 clamp for bfloat16: output[i] = clamp(input[i], min_val, max_val) +// Uses 4x unrolling for better ILP and OpenMP parallelization for large arrays. +void simd_clamp( + const bf16* input, + bf16* output, + bf16 min_val, + bf16 max_val, + size_t size +) { + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 4; + constexpr size_t CHUNK_SIZE = SIMD_WIDTH * UNROLL_FACTOR; // 64 elements + + const __m512 vmin = _mm512_set1_ps(static_cast(min_val)); + const __m512 vmax = _mm512_set1_ps(static_cast(max_val)); + + constexpr size_t min_elements_per_thread = CHUNK_SIZE * 8; + const bool use_parallel = size >= (min_elements_per_thread * 2); + const size_t num_chunks = size / CHUNK_SIZE; + const int signed_num_chunks = static_cast(num_chunks); + + #pragma omp parallel for num_threads(max_prefill_threads) if(use_parallel) schedule(static) + for (int chunk_idx = 0; chunk_idx < signed_num_chunks; ++chunk_idx) { + size_t i = static_cast(chunk_idx) * CHUNK_SIZE; + + _mm_prefetch(reinterpret_cast(input + i + 64), _MM_HINT_T0); + + __m512 v0 = load_bfloat16_to_m512(input + i); + __m512 v1 = load_bfloat16_to_m512(input + i + 16); + __m512 v2 = load_bfloat16_to_m512(input + i + 32); + __m512 v3 = load_bfloat16_to_m512(input + i + 48); + + v0 = _mm512_min_ps(_mm512_max_ps(v0, vmin), vmax); + v1 = _mm512_min_ps(_mm512_max_ps(v1, vmin), vmax); + v2 = _mm512_min_ps(_mm512_max_ps(v2, vmin), vmax); + v3 = _mm512_min_ps(_mm512_max_ps(v3, vmin), vmax); + + store_m512_to_bfloat16_rne(output + i, v0); + store_m512_to_bfloat16_rne(output + i + 16, v1); + store_m512_to_bfloat16_rne(output + i + 32, v2); + store_m512_to_bfloat16_rne(output + i + 48, v3); + } + + // Remaining 16-element chunks + size_t i = num_chunks * CHUNK_SIZE; + for (; i + SIMD_WIDTH <= size; i += SIMD_WIDTH) { + __m512 v = load_bfloat16_to_m512(input + i); + v = _mm512_min_ps(_mm512_max_ps(v, vmin), vmax); + store_m512_to_bfloat16_rne(output + i, v); + } + + // Scalar tail + for (; i < size; ++i) { + float val = static_cast(input[i]); + float fmin = static_cast(min_val); + float fmax = static_cast(max_val); + if (val < fmin) val = fmin; + if (val > fmax) val = fmax; + output[i] = static_cast(val); + } +} + +void simd_mul( + const bf16* input1, + const bf16* input2, + bf16* output, + size_t size +) { + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 4; + constexpr size_t CHUNK_SIZE = SIMD_WIDTH * UNROLL_FACTOR; + constexpr size_t min_elements_per_thread = CHUNK_SIZE * 8; + constexpr int max_threads = max_prefill_threads; + + const bool use_parallel = size >= (min_elements_per_thread * 2); + const size_t num_chunks = size / CHUNK_SIZE; + const int signed_num_chunks = static_cast(num_chunks); + + #pragma omp parallel for num_threads(max_threads) if(use_parallel) schedule(static) + for (int chunk_idx = 0; chunk_idx < signed_num_chunks; ++chunk_idx) { + const size_t i = static_cast(chunk_idx) * CHUNK_SIZE; + + _mm_prefetch(reinterpret_cast(input1 + i + 64), _MM_HINT_T0); + _mm_prefetch(reinterpret_cast(input2 + i + 64), _MM_HINT_T0); + + const __m512 a_vec0 = load_bfloat16_to_m512(input1 + i); + const __m512 b_vec0 = load_bfloat16_to_m512(input2 + i); + const __m512 a_vec1 = load_bfloat16_to_m512(input1 + i + 16); + const __m512 b_vec1 = load_bfloat16_to_m512(input2 + i + 16); + const __m512 a_vec2 = load_bfloat16_to_m512(input1 + i + 32); + const __m512 b_vec2 = load_bfloat16_to_m512(input2 + i + 32); + const __m512 a_vec3 = load_bfloat16_to_m512(input1 + i + 48); + const __m512 b_vec3 = load_bfloat16_to_m512(input2 + i + 48); + + store_m512_to_bfloat16_rne(output + i, _mm512_mul_ps(a_vec0, b_vec0)); + store_m512_to_bfloat16_rne(output + i + 16, _mm512_mul_ps(a_vec1, b_vec1)); + store_m512_to_bfloat16_rne(output + i + 32, _mm512_mul_ps(a_vec2, b_vec2)); + store_m512_to_bfloat16_rne(output + i + 48, _mm512_mul_ps(a_vec3, b_vec3)); + } + + size_t i = num_chunks * CHUNK_SIZE; + + for (; i + SIMD_WIDTH <= size; i += SIMD_WIDTH) { + const __m512 a_vec = load_bfloat16_to_m512(input1 + i); + const __m512 b_vec = load_bfloat16_to_m512(input2 + i); + store_m512_to_bfloat16_rne(output + i, _mm512_mul_ps(a_vec, b_vec)); + } + + for (; i < size; ++i) { + output[i] = static_cast(static_cast(input1[i]) * static_cast(input2[i])); + } +} + +void simd_mul( + const bf16* input1, + bf16 input2_scalar, + bf16* output, + size_t size +) { + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 4; + constexpr size_t CHUNK_SIZE = SIMD_WIDTH * UNROLL_FACTOR; + constexpr size_t min_elements_per_thread = CHUNK_SIZE * 8; + constexpr int max_threads = max_prefill_threads; + + // Broadcast scalar to all 16 lanes once + const __m512 scalar_vec = _mm512_set1_ps(static_cast(input2_scalar)); + + const bool use_parallel = size >= (min_elements_per_thread * 2); + const size_t num_chunks = size / CHUNK_SIZE; + const int signed_num_chunks = static_cast(num_chunks); + + #pragma omp parallel for num_threads(max_threads) if(use_parallel) schedule(static) + for (int chunk_idx = 0; chunk_idx < signed_num_chunks; ++chunk_idx) { + const size_t i = static_cast(chunk_idx) * CHUNK_SIZE; + + _mm_prefetch(reinterpret_cast(input1 + i + 64), _MM_HINT_T0); + + const __m512 a_vec0 = load_bfloat16_to_m512(input1 + i); + const __m512 a_vec1 = load_bfloat16_to_m512(input1 + i + 16); + const __m512 a_vec2 = load_bfloat16_to_m512(input1 + i + 32); + const __m512 a_vec3 = load_bfloat16_to_m512(input1 + i + 48); + + store_m512_to_bfloat16_rne(output + i, _mm512_mul_ps(a_vec0, scalar_vec)); + store_m512_to_bfloat16_rne(output + i + 16, _mm512_mul_ps(a_vec1, scalar_vec)); + store_m512_to_bfloat16_rne(output + i + 32, _mm512_mul_ps(a_vec2, scalar_vec)); + store_m512_to_bfloat16_rne(output + i + 48, _mm512_mul_ps(a_vec3, scalar_vec)); + } + + size_t i = num_chunks * CHUNK_SIZE; + + for (; i + SIMD_WIDTH <= size; i += SIMD_WIDTH) { + const __m512 a_vec = load_bfloat16_to_m512(input1 + i); + store_m512_to_bfloat16_rne(output + i, _mm512_mul_ps(a_vec, scalar_vec)); + } + + const float scalar_f = static_cast(input2_scalar); + for (; i < size; ++i) { + output[i] = static_cast(static_cast(input1[i]) * scalar_f); + } +} + +// ============================================================================ +// simd_relu: AVX-512 ReLU for bfloat16 +// output[i] = max(input[i], 0) +// ============================================================================ +void simd_relu( + const bf16* input, + bf16* output, + size_t size +) { + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 4; + constexpr size_t CHUNK_SIZE = SIMD_WIDTH * UNROLL_FACTOR; // 64 elements + + const __m512 zero = _mm512_setzero_ps(); + + constexpr size_t min_elements_per_thread = CHUNK_SIZE * 8; + const bool use_parallel = size >= (min_elements_per_thread * 2); + const size_t num_chunks = size / CHUNK_SIZE; + const int signed_num_chunks = static_cast(num_chunks); + + #pragma omp parallel for num_threads(max_prefill_threads) if(use_parallel) schedule(static) + for (int chunk_idx = 0; chunk_idx < signed_num_chunks; ++chunk_idx) { + size_t i = static_cast(chunk_idx) * CHUNK_SIZE; + + _mm_prefetch(reinterpret_cast(input + i + 64), _MM_HINT_T0); + + __m512 v0 = load_bfloat16_to_m512(input + i); + __m512 v1 = load_bfloat16_to_m512(input + i + 16); + __m512 v2 = load_bfloat16_to_m512(input + i + 32); + __m512 v3 = load_bfloat16_to_m512(input + i + 48); + + store_m512_to_bfloat16_rne(output + i, _mm512_max_ps(v0, zero)); + store_m512_to_bfloat16_rne(output + i + 16, _mm512_max_ps(v1, zero)); + store_m512_to_bfloat16_rne(output + i + 32, _mm512_max_ps(v2, zero)); + store_m512_to_bfloat16_rne(output + i + 48, _mm512_max_ps(v3, zero)); + } + + size_t i = num_chunks * CHUNK_SIZE; + for (; i + SIMD_WIDTH <= size; i += SIMD_WIDTH) { + __m512 v = load_bfloat16_to_m512(input + i); + store_m512_to_bfloat16_rne(output + i, _mm512_max_ps(v, zero)); + } + + // Scalar tail + for (; i < size; ++i) { + float val = static_cast(input[i]); + output[i] = static_cast(val > 0.0f ? val : 0.0f); + } +} + +// ============================================================================ +// simd_silu: AVX-512 SiLU (Swish) for bfloat16 +// output[i] = input[i] * sigmoid(input[i]) +// ============================================================================ +void simd_silu( + const bf16* input, + bf16* output, + size_t size +) { + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 4; + constexpr size_t CHUNK_SIZE = SIMD_WIDTH * UNROLL_FACTOR; // 64 elements + + constexpr size_t min_elements_per_thread = CHUNK_SIZE * 8; + const bool use_parallel = size >= (min_elements_per_thread * 2); + const size_t num_chunks = size / CHUNK_SIZE; + const int signed_num_chunks = static_cast(num_chunks); + + #pragma omp parallel for num_threads(max_prefill_threads) if(use_parallel) schedule(static) + for (int chunk_idx = 0; chunk_idx < signed_num_chunks; ++chunk_idx) { + size_t i = static_cast(chunk_idx) * CHUNK_SIZE; + + _mm_prefetch(reinterpret_cast(input + i + 64), _MM_HINT_T0); + + __m512 v0 = load_bfloat16_to_m512(input + i); + __m512 v1 = load_bfloat16_to_m512(input + i + 16); + __m512 v2 = load_bfloat16_to_m512(input + i + 32); + __m512 v3 = load_bfloat16_to_m512(input + i + 48); + + store_m512_to_bfloat16_rne(output + i, silu_avx512(v0)); + store_m512_to_bfloat16_rne(output + i + 16, silu_avx512(v1)); + store_m512_to_bfloat16_rne(output + i + 32, silu_avx512(v2)); + store_m512_to_bfloat16_rne(output + i + 48, silu_avx512(v3)); + } + + size_t i = num_chunks * CHUNK_SIZE; + for (; i + SIMD_WIDTH <= size; i += SIMD_WIDTH) { + __m512 v = load_bfloat16_to_m512(input + i); + store_m512_to_bfloat16_rne(output + i, silu_avx512(v)); + } + + // Scalar tail + for (; i < size; ++i) { + float val = static_cast(input[i]); + float sig = 1.0f / (1.0f + std::exp(-val)); + output[i] = static_cast(val * sig); + } +} + +// ============================================================================ +// simd_layernorm: AVX-512 Layer Normalization for bfloat16 +// input shape: [seq_len, D_padded] in row major, only D columns valid +// output shape: [seq_len, D_padded] +// weights: [D] (per-element scale, applied after normalization) +// Formula per row: output = ((input - mean) / sqrt(var + eps)) * weights +// ============================================================================ +void simd_layernorm( + bf16* input, + bf16* output, + const bf16* weights, + int D, + int D_padded, + int seq_len, + float eps +) { + constexpr size_t SIMD_WIDTH = 16; + const float inv_D = 1.0f / static_cast(D); + + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) if(seq_len >= 8) + for (int row = 0; row < seq_len; ++row) { + bf16* row_in = input + static_cast(row) * D_padded; + bf16* row_out = output + static_cast(row) * D_padded; + + // === Pass 1: compute mean === + __m512 sum_acc0 = _mm512_setzero_ps(); + __m512 sum_acc1 = _mm512_setzero_ps(); + size_t i = 0; + + for (; i + SIMD_WIDTH * 2 <= static_cast(D); i += SIMD_WIDTH * 2) { + sum_acc0 = _mm512_add_ps(sum_acc0, load_bfloat16_to_m512(row_in + i)); + sum_acc1 = _mm512_add_ps(sum_acc1, load_bfloat16_to_m512(row_in + i + SIMD_WIDTH)); + } + for (; i + SIMD_WIDTH <= static_cast(D); i += SIMD_WIDTH) { + sum_acc0 = _mm512_add_ps(sum_acc0, load_bfloat16_to_m512(row_in + i)); + } + + float sum = _mm512_reduce_add_ps(_mm512_add_ps(sum_acc0, sum_acc1)); + for (; i < static_cast(D); ++i) { + sum += static_cast(row_in[i]); + } + + float mean = sum * inv_D; + __m512 mean_vec = _mm512_set1_ps(mean); + + // === Pass 2: compute variance === + __m512 var_acc0 = _mm512_setzero_ps(); + __m512 var_acc1 = _mm512_setzero_ps(); + i = 0; + + for (; i + SIMD_WIDTH * 2 <= static_cast(D); i += SIMD_WIDTH * 2) { + __m512 diff0 = _mm512_sub_ps(load_bfloat16_to_m512(row_in + i), mean_vec); + __m512 diff1 = _mm512_sub_ps(load_bfloat16_to_m512(row_in + i + SIMD_WIDTH), mean_vec); + var_acc0 = _mm512_fmadd_ps(diff0, diff0, var_acc0); + var_acc1 = _mm512_fmadd_ps(diff1, diff1, var_acc1); + } + for (; i + SIMD_WIDTH <= static_cast(D); i += SIMD_WIDTH) { + __m512 diff = _mm512_sub_ps(load_bfloat16_to_m512(row_in + i), mean_vec); + var_acc0 = _mm512_fmadd_ps(diff, diff, var_acc0); + } + + float var_sum = _mm512_reduce_add_ps(_mm512_add_ps(var_acc0, var_acc1)); + for (; i < static_cast(D); ++i) { + float diff = static_cast(row_in[i]) - mean; + var_sum += diff * diff; + } + + float variance = var_sum * inv_D; + float inv_std = 1.0f / std::sqrt(variance + eps); + __m512 inv_std_vec = _mm512_set1_ps(inv_std); + + // === Pass 3: normalize and apply weights === + i = 0; + for (; i + SIMD_WIDTH * 2 <= static_cast(D); i += SIMD_WIDTH * 2) { + __m512 x0 = load_bfloat16_to_m512(row_in + i); + __m512 x1 = load_bfloat16_to_m512(row_in + i + SIMD_WIDTH); + __m512 w0 = load_bfloat16_to_m512(weights + i); + __m512 w1 = load_bfloat16_to_m512(weights + i + SIMD_WIDTH); + + __m512 norm0 = _mm512_mul_ps(_mm512_sub_ps(x0, mean_vec), inv_std_vec); + __m512 norm1 = _mm512_mul_ps(_mm512_sub_ps(x1, mean_vec), inv_std_vec); + + store_m512_to_bfloat16_rne(row_out + i, _mm512_mul_ps(norm0, w0)); + store_m512_to_bfloat16_rne(row_out + i + SIMD_WIDTH, _mm512_mul_ps(norm1, w1)); + } + for (; i + SIMD_WIDTH <= static_cast(D); i += SIMD_WIDTH) { + __m512 x = load_bfloat16_to_m512(row_in + i); + __m512 w = load_bfloat16_to_m512(weights + i); + __m512 norm = _mm512_mul_ps(_mm512_sub_ps(x, mean_vec), inv_std_vec); + store_m512_to_bfloat16_rne(row_out + i, _mm512_mul_ps(norm, w)); + } + + // Scalar tail + for (; i < static_cast(D); ++i) { + float val = (static_cast(row_in[i]) - mean) * inv_std; + row_out[i] = static_cast(val * static_cast(weights[i])); + } + } +} + +// ============================================================================ +// scalar_conv2d: Reference scalar 2D convolution for bfloat16 (verification) +// Same interface and semantics as simd_conv2d but purely scalar. +// Accumulates in float32 for precision. +// ============================================================================ +void scalar_conv2d( + const bf16* input, + const bf16* kernel, + bf16* output, + int C_in, int H_in, int W_in, + int C_out, int K, int stride, int padding +) { + const int H_out = (H_in + 2 * padding - K) / stride + 1; + const int W_out = (W_in + 2 * padding - K) / stride + 1; + const int kernel_size = C_in * K * K; + + for (int oc = 0; oc < C_out; ++oc) { + const bf16* oc_kernel = kernel + static_cast(oc) * kernel_size; + bf16* oc_output = output + static_cast(oc) * H_out * W_out; + + for (int oh = 0; oh < H_out; ++oh) { + for (int ow = 0; ow < W_out; ++ow) { + float acc = 0.0f; + int k_idx = 0; + + for (int ic = 0; ic < C_in; ++ic) { + const bf16* ic_input = input + static_cast(ic) * H_in * W_in; + + for (int kh = 0; kh < K; ++kh) { + int ih = oh * stride - padding + kh; + + if (ih < 0 || ih >= H_in) { + k_idx += K; + continue; + } + + const bf16* input_row = ic_input + ih * W_in; + + for (int kw = 0; kw < K; ++kw, ++k_idx) { + int iw = ow * stride - padding + kw; + + if (iw < 0 || iw >= W_in) { + continue; + } + + acc += static_cast(input_row[iw]) * static_cast(oc_kernel[k_idx]); + } + } + } + + oc_output[oh * W_out + ow] = static_cast(acc); + } + } + } +} + +// ============================================================================ +// simd_conv2d: AVX-512 2D convolution for bfloat16 with OpenMP +// input: [C_in, H_in, W_in] in CHW layout +// kernel: [C_out, C_in, K, K] +// output: [C_out, H_out, W_out] +// H_out = (H_in + 2*padding - K) / stride + 1 +// W_out = (W_in + 2*padding - K) / stride + 1 +// +// Strategy: im2col-style accumulation with cache-friendly access patterns. +// - Outer loop over output channels (parallelized with OpenMP) +// - For each output position, accumulate over C_in * K * K with SIMD +// ============================================================================ +void simd_conv2d( + bf16* input, + const bf16* kernel, + bf16* output, + int C_in, int H_in, int W_in, + int C_out, int K, int stride, int padding +) { + const int H_out = (H_in + 2 * padding - K) / stride + 1; + const int W_out = (W_in + 2 * padding - K) / stride + 1; + const int kernel_size = C_in * K * K; // elements per output channel kernel + + constexpr size_t SIMD_WIDTH = 16; + + // Parallelize over output channels — each thread works on independent output channels + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) if(C_out >= max_prefill_threads) + for (int oc = 0; oc < C_out; ++oc) { + const bf16* oc_kernel = kernel + static_cast(oc) * kernel_size; + bf16* oc_output = output + static_cast(oc) * H_out * W_out; + + for (int oh = 0; oh < H_out; ++oh) { + // --- SIMD path: vectorize across output width (ow) --- + // For each (ic, kh, kw), broadcast the kernel weight and accumulate + // over SIMD_WIDTH output positions simultaneously. + int ow = 0; + for (; ow + static_cast(SIMD_WIDTH) <= W_out; ow += static_cast(SIMD_WIDTH)) { + __m512 acc_vec = _mm512_setzero_ps(); + int k_idx = 0; + + for (int ic = 0; ic < C_in; ++ic) { + const bf16* ic_input = input + static_cast(ic) * H_in * W_in; + + for (int kh = 0; kh < K; ++kh) { + int ih = oh * stride - padding + kh; + + if (ih < 0 || ih >= H_in) { + k_idx += K; + continue; + } + + const bf16* input_row = ic_input + ih * W_in; + + for (int kw = 0; kw < K; ++kw, ++k_idx) { + __m512 k_vec = _mm512_set1_ps(static_cast(oc_kernel[k_idx])); + + // Gather SIMD_WIDTH input values for consecutive ow positions + // iw = (ow + lane) * stride - padding + kw + int iw_base = ow * stride - padding + kw; + + if (stride == 1) { + // Contiguous access: load directly from input_row + if (iw_base >= 0 && iw_base + static_cast(SIMD_WIDTH) <= W_in) { + // All elements in bounds — fast path + __m512 in_vec = load_bfloat16_to_m512(input_row + iw_base); + acc_vec = _mm512_fmadd_ps(in_vec, k_vec, acc_vec); + } else { + // Some elements may be out of bounds — zero-padded + alignas(64) float tmp[16] = {0}; + for (int lane = 0; lane < static_cast(SIMD_WIDTH); ++lane) { + int iw = iw_base + lane; + if (iw >= 0 && iw < W_in) { + tmp[lane] = static_cast(input_row[iw]); + } + } + __m512 in_vec = _mm512_load_ps(tmp); + acc_vec = _mm512_fmadd_ps(in_vec, k_vec, acc_vec); + } + } else { + // Strided access — gather element by element + alignas(64) float tmp[16] = {0}; + for (int lane = 0; lane < static_cast(SIMD_WIDTH); ++lane) { + int iw = (ow + lane) * stride - padding + kw; + if (iw >= 0 && iw < W_in) { + tmp[lane] = static_cast(input_row[iw]); + } + } + __m512 in_vec = _mm512_load_ps(tmp); + acc_vec = _mm512_fmadd_ps(in_vec, k_vec, acc_vec); + } + } + } + } + + // Store SIMD_WIDTH output values + store_m512_to_bfloat16_rne(oc_output + oh * W_out + ow, acc_vec); + } + + // --- Scalar tail for remaining ow positions --- + for (; ow < W_out; ++ow) { + float acc_scalar = 0.0f; + int k_idx = 0; + + for (int ic = 0; ic < C_in; ++ic) { + const bf16* ic_input = input + static_cast(ic) * H_in * W_in; + + for (int kh = 0; kh < K; ++kh) { + int ih = oh * stride - padding + kh; + + if (ih < 0 || ih >= H_in) { + k_idx += K; + continue; + } + + const bf16* input_row = ic_input + ih * W_in; + + for (int kw = 0; kw < K; ++kw, ++k_idx) { + int iw = ow * stride - padding + kw; + + if (iw < 0 || iw >= W_in) { + continue; + } + + float in_val = static_cast(input_row[iw]); + float k_val = static_cast(oc_kernel[k_idx]); + acc_scalar += in_val * k_val; + } + } + } + + oc_output[oh * W_out + ow] = static_cast(acc_scalar); + } + } + } +} + +// ============================================================================ +// layernorm_relu_nchw: Applies LayerNorm over C dimension + ReLU activation +// in an NCHW layout. Strides over channels by HW. +// Uses float32 accumulation to match PyTorch's LayerNorm internal precision. +// ============================================================================ +void layernorm_relu_nchw(bf16* data, const bf16* norm_weight, int C, int HW, float eps) { + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) + for(int hw = 0; hw < HW; hw++){ + float mean = 0.0f; + for(int c = 0; c < C; c++){ + mean += static_cast(data[c * HW + hw]); + } + mean /= C; + + float var = 0.0f; + for(int c = 0; c < C; c++){ + float diff = static_cast(data[c * HW + hw]) - mean; + var += diff * diff; + } + var /= C; + + float inv_std = 1.0f / std::sqrt(var + eps); + + for(int c = 0; c < C; c++){ + float val = static_cast(data[c * HW + hw]); + float normed = (val - mean) * inv_std * static_cast(norm_weight[c]); + data[c * HW + hw] = static_cast(std::max(0.0f, normed)); + } + } +} + +// ============================================================================ +// layernorm_gelu_nchw: Applies LayerNorm over C dimension + GELU (tanh approx) +// ============================================================================ +void layernorm_gelu_nchw(bf16* data, const bf16* norm_weight, int C, int HW, float eps) { + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) + for(int hw = 0; hw < HW; hw++){ + double mean = 0.0; + for(int c = 0; c < C; c++){ + mean += static_cast(data[c * HW + hw]); + } + mean /= C; + + double var = 0.0; + for(int c = 0; c < C; c++){ + double diff = static_cast(data[c * HW + hw]) - mean; + var += diff * diff; + } + var /= C; + + float inv_std = 1.0f / std::sqrt(static_cast(var) + eps); + + for(int c = 0; c < C; c++){ + float val = static_cast(data[c * HW + hw]); + float normed = (val - static_cast(mean)) * inv_std * static_cast(norm_weight[c]); + + // GELU tanh approx + float x_cubed = normed * normed * normed; + float inner = 0.7978845608f * (normed + 0.044715f * x_cubed); + float gelu_val = 0.5f * normed * (1.0f + std::tanh(inner)); + + data[c * HW + hw] = static_cast(gelu_val); + } + } +} + +// ============================================================================ +// rmsnorm_gelu_nchw: Applies RMSNorm over C dimension + GELU (tanh approx) +// ============================================================================ +void rmsnorm_gelu_nchw(bf16* data, const bf16* norm_weight, int C, int HW, float eps) { + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) + for(int hw = 0; hw < HW; hw++){ + double sq_sum = 0.0; + for(int c = 0; c < C; c++){ + float val = static_cast(data[c * HW + hw]); + sq_sum += val * val; + } + float var = sq_sum / C; + float inv_std = 1.0f / std::sqrt(var + eps); + + for(int c = 0; c < C; c++){ + float val = static_cast(data[c * HW + hw]); + float normed = val * inv_std * static_cast(norm_weight[c]); + + // GELU tanh approx + float x_cubed = normed * normed * normed; + float inner = 0.7978845608f * (normed + 0.044715f * x_cubed); + float gelu_val = 0.5f * normed * (1.0f + std::tanh(inner)); + + data[c * HW + hw] = static_cast(gelu_val); + } + } +} + +void simd_gemm_abt_bf16( + const bf16* A, const bf16* B, bf16* C, + int M, int N, int K, + int lda, int ldb, int ldc +) { + // C[i][j] = sum_k A[i*lda + k] * B[j*ldb + k] + constexpr size_t SIMD_WIDTH = 16; + + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) if(M * N > 64) + for (int i = 0; i < M; i++) { + const bf16* a_row = A + (size_t)i * lda; + for (int j = 0; j < N; j++) { + const bf16* b_row = B + (size_t)j * ldb; + + __m512 acc0 = _mm512_setzero_ps(); + __m512 acc1 = _mm512_setzero_ps(); + int k = 0; + for (; k + 2 * (int)SIMD_WIDTH <= K; k += 2 * SIMD_WIDTH) { + acc0 = _mm512_fmadd_ps(load_bfloat16_to_m512(a_row + k), + load_bfloat16_to_m512(b_row + k), acc0); + acc1 = _mm512_fmadd_ps(load_bfloat16_to_m512(a_row + k + SIMD_WIDTH), + load_bfloat16_to_m512(b_row + k + SIMD_WIDTH), acc1); + } + __m512 acc = _mm512_add_ps(acc0, acc1); + for (; k + (int)SIMD_WIDTH <= K; k += SIMD_WIDTH) { + acc = _mm512_fmadd_ps(load_bfloat16_to_m512(a_row + k), + load_bfloat16_to_m512(b_row + k), acc); + } + float sum = _mm512_reduce_add_ps(acc); + for (; k < K; k++) { + sum += (float)a_row[k] * (float)b_row[k]; + } + C[(size_t)i * ldc + j] = (bf16)sum; + } + } +} + +// ============================================================================ +// simd_glu: Gated Linear Unit for bfloat16 +// input: [seq_len, 2, hidden_dim] — first half is value, second half is gate +// output: [seq_len, hidden_dim] +// Formula: output[s, :] = input[s, 0, :] * sigmoid(input[s, 1, :]) +// Fused single-pass: load both halves, sigmoid the gate, multiply, store. +// ============================================================================ +void simd_glu( + const bf16* input, + bf16* output, + size_t seq_len, + size_t hidden_dim +) { + constexpr size_t SIMD_WIDTH = 16; + constexpr size_t UNROLL_FACTOR = 4; + constexpr size_t CHUNK_SIZE = SIMD_WIDTH * UNROLL_FACTOR; // 64 elements + const size_t stride = 2 * hidden_dim; // row stride in input + + int signed_seq_len = static_cast(seq_len); + + #pragma omp parallel for num_threads(max_prefill_threads) schedule(static) if(seq_len >= 4) + for (int s = 0; s < signed_seq_len; ++s) { + const bf16* val_ptr = input + (size_t)s * stride; // first half + const bf16* gate_ptr = input + (size_t)s * stride + hidden_dim; // second half + bf16* out_ptr = output + (size_t)s * hidden_dim; + + size_t d = 0; + + // Main loop: 4x unrolled SIMD + for (; d + CHUNK_SIZE <= hidden_dim; d += CHUNK_SIZE) { + _mm_prefetch(reinterpret_cast(val_ptr + d + 64), _MM_HINT_T0); + _mm_prefetch(reinterpret_cast(gate_ptr + d + 64), _MM_HINT_T0); + + __m512 v0 = load_bfloat16_to_m512(val_ptr + d); + __m512 g0 = sigmoid_avx512(load_bfloat16_to_m512(gate_ptr + d)); + + __m512 v1 = load_bfloat16_to_m512(val_ptr + d + 16); + __m512 g1 = sigmoid_avx512(load_bfloat16_to_m512(gate_ptr + d + 16)); + + __m512 v2 = load_bfloat16_to_m512(val_ptr + d + 32); + __m512 g2 = sigmoid_avx512(load_bfloat16_to_m512(gate_ptr + d + 32)); + + __m512 v3 = load_bfloat16_to_m512(val_ptr + d + 48); + __m512 g3 = sigmoid_avx512(load_bfloat16_to_m512(gate_ptr + d + 48)); + + store_m512_to_bfloat16_rne(out_ptr + d, _mm512_mul_ps(v0, g0)); + store_m512_to_bfloat16_rne(out_ptr + d + 16, _mm512_mul_ps(v1, g1)); + store_m512_to_bfloat16_rne(out_ptr + d + 32, _mm512_mul_ps(v2, g2)); + store_m512_to_bfloat16_rne(out_ptr + d + 48, _mm512_mul_ps(v3, g3)); + } + + // Remaining 16-element chunks + for (; d + SIMD_WIDTH <= hidden_dim; d += SIMD_WIDTH) { + __m512 v = load_bfloat16_to_m512(val_ptr + d); + __m512 g = sigmoid_avx512(load_bfloat16_to_m512(gate_ptr + d)); + store_m512_to_bfloat16_rne(out_ptr + d, _mm512_mul_ps(v, g)); + } + + // Scalar tail + for (; d < hidden_dim; ++d) { + float v = static_cast(val_ptr[d]); + float g = static_cast(gate_ptr[d]); + float sig = 1.0f / (1.0f + std::exp(-g)); + out_ptr[d] = static_cast(v * sig); + } + } +} + +void scalar_conv1d( + int conv_kernel_size, + int conv_stride, + const bf16* input, // [seq_len + (conv_kernel_size - conv_stride), hidden_dim] + const bf16* kernel_weights, // [conv_kernel_size, hidden_dim] + bf16* output, // [seq_len, hidden_dim] + int seq_len, + int hidden_dim +){ + for (int o = 0; o < seq_len; o++) { + bf16* out_ptr = output + o * hidden_dim; + for (int d = 0; d < hidden_dim; d++) { + float acc = 0.0f; + for (int k = 0; k < conv_kernel_size; k++) { + int in_idx = (o * conv_stride + k) * hidden_dim + d; + int w_idx = k * hidden_dim + d; + acc += static_cast(input[in_idx]) * static_cast(kernel_weights[w_idx]); + } + out_ptr[d] = static_cast(acc); + } + } +} diff --git a/src/detail/gemma4e_npu/gemma4e_vision_prefill_helper.hpp b/src/detail/gemma4e_npu/gemma4e_vision_prefill_helper.hpp new file mode 100644 index 000000000..29f441f4f --- /dev/null +++ b/src/detail/gemma4e_npu/gemma4e_vision_prefill_helper.hpp @@ -0,0 +1,173 @@ +#include +#include +#include +#include +#include +#include +#include "typedef.hpp" +#pragma once + +void simd_conv2d( + bf16* input, + const bf16* kernel, + bf16* output, + int C_in, int H_in, int W_in, + int C_out, int K, int stride, int padding +); +void scalar_conv2d( + const bf16* input, + const bf16* kernel, + bf16* output, + int C_in, int H_in, int W_in, + int C_out, int K, int stride, int padding +); + +void simd_layernorm( + bf16* input, // shape of [seq_len, D_padded] in row major ordr, but only D is valud + bf16* output, + const bf16* weights, + int D, + int D_padded, + int seq_len, + float eps = 1e-6f +); + +// NCHW channel-wise operations (acting across C dimension, strided by HW) +void layernorm_relu_nchw(bf16* data, const bf16* norm_weight, int C, int HW, float eps = 1e-6f); +void layernorm_gelu_nchw(bf16* data, const bf16* norm_weight, int C, int HW, float eps = 1e-6f); +void rmsnorm_gelu_nchw(bf16* data, const bf16* norm_weight, int C, int HW, float eps = 1e-6f); + +void simd_relu( + const bf16* input, + bf16* output, + size_t size +); + +void simd_silu( + const bf16* input, + bf16* output, + size_t size +); + +void simd_add( + bf16* input1, + bf16* input2, + bf16* output, + size_t size +); + +// Optimized AVX-512 version that combines bias addition and GELU activation +// bias is broadcast across sequence positions (size = hidden_dim, not total_size) +void simd_bias_add_gelu( + bf16* input, + const bf16* bias, + bf16* output, + size_t total_size, + size_t hidden_dim +); + +void gelu_bfloat16_ref( + const bf16* input, + bf16* output, + size_t size +); + +// GELU activation function with tanh approximation (same as PytorchGELUTanh) +inline float gelu_tanh(float x) { + return 0.5f * x * (1 + std::tanh(std::sqrt(2.0f / 3) * (x + 0.044715f * x * x * x))); +} + +// RMS Norm without scale weights +// Input layout: [seq_len_padded x X_padded], only processes seq_len rows and X cols per row +// Formula: output = input * rsqrt(mean(input^2) + eps) +void simd_rms_norm( + const bf16* input, + bf16* output, + size_t seq_len, + size_t X, + size_t seq_len_padded, + size_t X_padded, + float eps = 1e-6f +); + +// RMS Norm with scale weights (Gemma4RMSNorm with with_scale=True) +// Input layout: [seq_len_padded x X_padded], only processes seq_len rows and X cols per row +// norm_weight: bf16 pointer of size [X] +// Formula: output = (input * rsqrt(mean(input^2) + eps)) * weight +void simd_rms_norm( + const bf16* input, + const bf16* norm_weight, + bf16* output, + size_t seq_len, + size_t X, + size_t seq_len_padded, + size_t X_padded, + float eps = 1e-6f +); + +void simd_clamp( + const bf16* input, + bf16* output, + bf16 min_val, + bf16 max_val, + size_t size +); + +void simd_mul( + const bf16* input1, + const bf16* input2, + bf16* output, + size_t size +); + +void simd_mul( + const bf16* input1, + bf16 input2_scalar, + bf16* output, + size_t size +); + +void transpose_2d( + const bf16* input, + bf16* output, + size_t rows, + size_t cols +); + +// General bf16 matmul: C = A @ B^T with float32 accumulation +// A: [M, K] row-major with stride lda +// B: [N, K] row-major with stride ldb (B^T gives [K, N]) +// C: [M, N] row-major with stride ldc +void simd_gemm_abt_bf16( + const bf16* A, const bf16* B, bf16* C, + int M, int N, int K, + int lda, int ldb, int ldc +); + +void simd_glu( + const bf16* input, //[seq_len, 2,hidden_dim] + bf16* output, // [seq_len, hidden_dim] + size_t seq_len, + size_t hidden_dim +); + +void conv1d( + int conv_kernel_size, + int conv_stride, + const bf16* input, /// [seq_len+ (conv_kernel_size-conv_stride), hidden_dim] + const bf16* kernel_weights, // [conv_kernel_size, hidden_dim] + bf16* output, // [seq_len, hidden_dim] + int seq_len, + int hidden_dim + +); + +void scalar_conv1d( + int conv_kernel_size, + int conv_stride, + const bf16* input, /// [seq_len+ (conv_kernel_size-conv_stride), hidden_dim] + const bf16* kernel_weights, // [conv_kernel_size, hidden_dim] + bf16* output, // [seq_len, hidden_dim] + int seq_len, + int hidden_dim +); diff --git a/src/detail/gemma4e_npu/mmRuntimeSequence.hpp b/src/detail/gemma4e_npu/mmRuntimeSequence.hpp new file mode 100644 index 000000000..8be70830c --- /dev/null +++ b/src/detail/gemma4e_npu/mmRuntimeSequence.hpp @@ -0,0 +1,550 @@ +#ifndef __MM_SEQUENCE_HPP__ +#define __MM_SEQUENCE_HPP__ +#include +#include // Required for std::max +#include +#include "npu_utils/npu_instr_utils.hpp" + +template +void generate_shimtile_sequence_per_k_block( + uint32_t shim_index, uint32_t total_npu_row, uint32_t total_npu_col, + uint32_t mega_block_row_idx, uint32_t mega_block_col_idx, + uint32_t M_size, uint32_t K_size, uint32_t N_size, + uint32_t m, uint32_t k, uint32_t n, + uint32_t Arg_A, uint32_t Arg_B, uint32_t Arg_C, + uint32_t A_const_offset, uint32_t B_const_offset, uint32_t C_const_offset, + + std::vector &list_A_shim_queue, std::vector &list_B_shim_queue, std::vector &list_C_shim_queue, + std::vector &list_A_bd_pingpong_flag, std::vector &list_B_bd_pingpong_flag, std::vector &list_C_bd_pingpong_flag, + + bool IS_B_ROW_MAJOR, bool ENABLE_AXI4, + bool B_in_K_N_block_col_major_order, + bool VALID_COLUMN, + bool ADD_BIAS, bool SEND_BIAS, + std::map& valid_A_MT_shimtile_index, + npu_sequence& seq, + std::vector & shimtile_list, + + bool REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH, + uint32_t DH + +){ + + if(B_in_K_N_block_col_major_order){ + assert(IS_B_ROW_MAJOR== false); // on valid for B in col major order + } + // When B_in_K_N_block_col_major_order is set to true, it mean + // B is col-major order && + // B is rearrange into kxn blocks, where blocks are in col-major. Moreover, the data in each blocks is + // also in col-major order. + + // Basically,B as a col-major matrix goes through + // stride: [N_size/n,K_size/k ,n, k] + // offset: [K_size*n,k ,K_SIZE, 1] + + auto AXI_FLAG = aggressive_cache; + ; + uint32_t K_div_k = K_size/k; + + npu_tiles cur_shimtile = shimtile_list.at(shim_index); + + if (valid_A_MT_shimtile_index.contains(shim_index) && valid_A_MT_shimtile_index[shim_index] < total_npu_row){ + + if (list_A_shim_queue.at(shim_index) == 2) { + seq.npu_dma_wait( + cur_shimtile, MM2S, it_channel_0 + ); + list_A_shim_queue.at(shim_index)--; + } + + uint32_t A_offset = mega_block_row_idx *(total_npu_row*m) *K_size; + A_offset += valid_A_MT_shimtile_index[shim_index]*(m*K_size); + npu_bd_id A_bd_id; + if (list_A_bd_pingpong_flag.at(shim_index) ==0){ + A_bd_id = bd_0; + list_A_bd_pingpong_flag.at(shim_index) =1; + }else{ + A_bd_id = bd_1; + list_A_bd_pingpong_flag.at(shim_index) =0; + } + + seq.npu_dma_memcpy_nd( + sizeof(T_in), // bfloat16 + Arg_A, + MM2S, + cur_shimtile, + A_bd_id, + it_channel_0, + {0,0,0,A_offset+ A_const_offset}, + {1, K_div_k, m,k}, + {0, k, K_size, 1}, + -1, 0, true, + // ENABLE_AXI4 ? AXI_FLAG: normal_cache + aggressive_cache + ); + list_A_shim_queue.at(shim_index)++; + } + + if(shim_index < total_npu_col && VALID_COLUMN){ + + if(SEND_BIAS){ + if (list_B_shim_queue[shim_index] == 2){ + + seq.npu_dma_wait( + cur_shimtile, MM2S, it_channel_1 + ); + list_B_shim_queue[shim_index] -= 1; + } + uint32_t _BIAS_DATA_OFFSET = mega_block_col_idx * (total_npu_col*n) + shim_index *n; + seq.npu_dma_memcpy_nd( + sizeof(T_in), + Arg_B, + MM2S, + cur_shimtile, + npu_bd_id(bd_6), //reserved for sending bias + it_channel_1, + {0,0,0,_BIAS_DATA_OFFSET}, + {1, 1,1, k*n}, + {0, 0, 0, 1}, + -1, 0, true, + // ENABLE_AXI4 ? AXI_FLAG: normal_cache + aggressive_cache + ); + list_B_shim_queue[shim_index]++; + } + + uint32_t BIAS_OFFSET = 0; + if (ADD_BIAS){ + BIAS_OFFSET = N_size; + } + + npu_bd_id b_bd_id; + if (list_B_shim_queue[shim_index] == 2){ + + seq.npu_dma_wait( + cur_shimtile, MM2S, it_channel_1 + ); + list_B_shim_queue[shim_index] -= 1; + } + + if (list_B_bd_pingpong_flag[shim_index] == 0) { + b_bd_id = bd_2; + list_B_bd_pingpong_flag[shim_index] = 1; + } else { + b_bd_id = bd_3; + list_B_bd_pingpong_flag[shim_index] = 0; + } + + if (IS_B_ROW_MAJOR){ + uint32_t B_offset = mega_block_col_idx* (total_npu_col) * n; + B_offset += shim_index * n; + seq.npu_dma_memcpy_nd( + sizeof(T_in), + Arg_B, + MM2S, + cur_shimtile, + b_bd_id, + it_channel_1, + {0,0,0,B_offset+ B_const_offset + BIAS_OFFSET}, + {1, K_div_k, k, n}, + {0, k*N_size, N_size, 1}, + -1, 0, true, + // ENABLE_AXI4 ? AXI_FLAG: normal_cache + aggressive_cache + ); + }else{ + uint32_t B_offset = mega_block_col_idx*(total_npu_col*n)*K_size; + B_offset += shim_index * n*K_size; + if(B_in_K_N_block_col_major_order){ + seq.npu_dma_memcpy_nd( + sizeof(T_in), + Arg_B, + MM2S, + cur_shimtile, + b_bd_id, + it_channel_1, + {0,0,0,B_offset+ B_const_offset + BIAS_OFFSET}, + {1, 1,1, K_div_k* n*k}, + {0, 0, 0, 1}, + -1, 0, true, + // ENABLE_AXI4 ? AXI_FLAG: normal_cache + aggressive_cache + ); + }else{ + seq.npu_dma_memcpy_nd( + sizeof(T_in), + Arg_B, + MM2S, + cur_shimtile, + b_bd_id, + it_channel_1, + {0,0,0,B_offset+ B_const_offset + BIAS_OFFSET}, + {1, K_div_k, n, k}, + {0, k, K_size, 1}, + -1, 0, true, + // ENABLE_AXI4 ? AXI_FLAG: normal_cache + aggressive_cache + ); + } + } + list_B_shim_queue[shim_index]++; + } + + if (shim_index < total_npu_col && VALID_COLUMN){ + + if (list_C_shim_queue.at(shim_index) == 2) { + seq.npu_dma_wait( + cur_shimtile, S2MM, it_channel_0 + ); + list_C_shim_queue.at(shim_index)--; + } + + npu_bd_id c_bd_id; + if(list_C_bd_pingpong_flag.at(shim_index) ==0){ + c_bd_id = bd_14; + list_C_bd_pingpong_flag.at(shim_index) =1; + }else{ + c_bd_id = bd_15; + list_C_bd_pingpong_flag.at(shim_index) =0; + } + + if(REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH){ + // Reorder from [M, N] row-major to [N/DH, M, DH] row-major + // where M = L_Seq and N = 3*NUM_HEADS*DH + /// debug + //std::cerr << "DH is " << DH << std::endl; + if( N_size % DH != 0 ){ + std::cerr << "MM: N_size % DH != 0 " << std::endl; + exit(-1); + } + + if( (total_npu_col * n)%DH != 0){ // for now + std::cerr << "MM: (total_npu_col * n)%DH != 0" < +void generate_runtime_sequence( + uint32_t Arg_A, uint32_t Arg_B, uint32_t Arg_C, + uint32_t A_const_offset, uint32_t B_const_offset, uint32_t C_const_offset, + uint32_t M_size, uint32_t N_size, uint32_t K_size, + uint32_t m, uint32_t n, uint32_t k, + uint32_t total_npu_row, uint32_t total_npu_col, + std::vector &list_A_shim_queue, + std::vector &list_B_shim_queue, + std::vector &list_C_shim_queue, + + bool IS_B_ROW_MAJOR, bool ENABLE_AXI4, bool B_in_K_N_block_col_major_order, + bool ADD_BIAS, + std::map& valid_A_MT_shimtile_index, + npu_sequence& seq, + std::vector &shim_tiles, + bool REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH, + uint32_t DH +){ + + uint32_t M_div_num_row_m = M_size/(m*total_npu_row); + uint32_t N_div_num_col_n = N_size/(n*total_npu_col); + + uint32_t N_div_num_col_n_remainder_blocks = (N_size % (n*total_npu_col))/ n; + + std::vector list_A_BD_pingpong_flag(std::max(total_npu_row, total_npu_col), 0); + std::vector list_B_BD_pingpong_flag(std::max(total_npu_row, total_npu_col), 0); + std::vector list_C_BD_pingpong_flag(std::max(total_npu_row, total_npu_col), 0); + + uint32_t col_block_range = N_div_num_col_n; + if (N_div_num_col_n_remainder_blocks!= 0){ + col_block_range += 1; + } + + for(uint32_t mega_block_col_idx=0; mega_block_col_idx( + shim_index, + total_npu_row, total_npu_col, + mega_block_row_idx, mega_block_col_idx, + M_size, K_size, N_size, + m, k, n, + Arg_A, Arg_B, Arg_C, + A_const_offset, B_const_offset, C_const_offset, + list_A_shim_queue, list_B_shim_queue, list_C_shim_queue, + list_A_BD_pingpong_flag, list_B_BD_pingpong_flag, list_C_BD_pingpong_flag, + IS_B_ROW_MAJOR, ENABLE_AXI4, + B_in_K_N_block_col_major_order, + true, + ADD_BIAS, SEND_ADD_BIAS, + valid_A_MT_shimtile_index, + seq, shim_tiles, + REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH, + DH + + ); + } + + else{ + generate_shimtile_sequence_per_k_block( + shim_index, + total_npu_row, total_npu_col, + mega_block_row_idx, mega_block_col_idx, + M_size, K_size, N_size, + m, k, n, + Arg_A, Arg_B, Arg_C, + A_const_offset, B_const_offset, C_const_offset, + list_A_shim_queue, list_B_shim_queue, list_C_shim_queue, + list_A_BD_pingpong_flag, list_B_BD_pingpong_flag, list_C_BD_pingpong_flag, + IS_B_ROW_MAJOR, ENABLE_AXI4, + B_in_K_N_block_col_major_order, + false, + ADD_BIAS, SEND_ADD_BIAS, + valid_A_MT_shimtile_index, + seq, shim_tiles, + REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH, + DH + ); + } + } + else{ + generate_shimtile_sequence_per_k_block( + shim_index, + total_npu_row, total_npu_col, + mega_block_row_idx, mega_block_col_idx, + M_size, K_size, N_size, + m, k, n, + Arg_A, Arg_B, Arg_C, + A_const_offset, B_const_offset, C_const_offset, + list_A_shim_queue, list_B_shim_queue, list_C_shim_queue, + list_A_BD_pingpong_flag, list_B_BD_pingpong_flag, list_C_BD_pingpong_flag, + IS_B_ROW_MAJOR, ENABLE_AXI4, + B_in_K_N_block_col_major_order, + true, + ADD_BIAS, SEND_ADD_BIAS, + valid_A_MT_shimtile_index, + seq, shim_tiles, + REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH, + DH + ); + } + } + } + } +} + +template +void generate_mm_sequence(npu_sequence &seq, uint32_t M, uint32_t K, uint32_t N, + uint32_t m, uint32_t k, uint32_t n, + uint32_t r, uint32_t s, uint32_t t, + uint32_t CT_rtp_address, uint32_t CT_rtp_sync_lock_id, + uint32_t total_row, uint32_t total_col, + uint32_t A_const_offset, uint32_t B_const_offset, uint32_t C_const_offset, + bool IS_B_ROW_MAJOR, bool ENABLE_AXI4, bool B_in_K_N_block_col_major_order, + bool ADD_BIAS, int OUTPUT_MODE, + int OUTPUT_CLAMP, + float output_clamp_min, float output_clamp_max, + bool REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH, // if false, output is in MXN row major + // if true, output is in NUM_DH x M x DH row major + uint32_t DH +){ + + constexpr int CT_lock_address_base = 0x000001F000; + const int Arg_A = 0; + const int Arg_B = 1; + const int Arg_C = 2; + + const int K_div_k = K/k; + + if( M%(m*total_row) != 0){ + std::cerr << "Error: M size not multiple of m * total_row"<< std::endl; + exit(1); + } + if( K%k != 0){ + std::cerr << "Error: K size not multiple of k"<< std::endl; + exit(1); + } + if( N%n != 0){ + std::cerr << "Error: N size not multiple of n"<< std::endl; + exit(1); + } + if(total_row != 4){ + std::cerr << "Error: total_row greater than 4 not supported"<< std::endl; + exit(1); + } + if(total_col !=8){ + std::cerr << "Error: total_col greater than 8 not supported"<< std::endl; + exit(1); + } + if(REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH){ + if( (N % DH) !=0 || DH%n !=0 || DH index offset + std::map valid_A_MT_shimtile_index; + valid_A_MT_shimtile_index[0] = 0; + valid_A_MT_shimtile_index[2] = 1; + valid_A_MT_shimtile_index[4] = 2; + valid_A_MT_shimtile_index[6] = 3; + + //create list of tiles + uint32_t shimtile_size = std::max(total_col, total_row); + std::vector shim_tiles; + + std::vector list_C_shim_queue; // int counter of how many DMA_Wait for C + std::vector list_A_shim_queue; // int counter of how many DMA_Wait for A + std::vector list_B_shim_queue; // int counter of how many DMA_Wait for B + for(size_t i = 0; i < shimtile_size; i++){ + shim_tiles.push_back((get_tile(0, i))); + list_C_shim_queue.push_back(0); + list_A_shim_queue.push_back(0); + list_B_shim_queue.push_back(0); + } + + // first, setup the rtp buffer and the rtp locks + + for(size_t row_idx = 0; row_idx < total_row; row_idx++){ + for(size_t col_idx = 0; col_idx< total_col; col_idx++){ + auto CT_tile = get_tile(row_idx+2, col_idx); + // set RTP value + seq.rtp_write( CT_tile, CT_rtp_address, K_div_k ); + seq.rtp_write( CT_tile, CT_rtp_address+4, M ); + seq.rtp_write( CT_tile, CT_rtp_address+8, N ); + if(ADD_BIAS){ + seq.rtp_write( CT_tile, CT_rtp_address+12, 1 ); + }else{ + seq.rtp_write( CT_tile, CT_rtp_address+12, 0 ); + } + seq.rtp_write( CT_tile, CT_rtp_address+16, OUTPUT_MODE ); // OUTPUT MODE + seq.rtp_write( CT_tile, CT_rtp_address+20, OUTPUT_CLAMP); // OUTPUT CLAMP 0 means no clamp, 1 means clamp + + int32_t output_min_int, output_max_int; + std::memcpy(&output_min_int, &output_clamp_min, sizeof(int32_t)); + std::memcpy(&output_max_int, &output_clamp_max, sizeof(int32_t)); + + seq.rtp_write( CT_tile, CT_rtp_address+24, output_min_int ); // output clamp min value in float32 + seq.rtp_write( CT_tile, CT_rtp_address+28, output_max_int ); // output clamp max value in float32 + // set RTP lock + seq.rtp_write(CT_tile, CT_lock_address_base+16*(CT_rtp_sync_lock_id), 1); // set lock to 1 + } + } + + generate_runtime_sequence( + Arg_A, Arg_B, Arg_C, + A_const_offset, B_const_offset, C_const_offset, + M, N, K, m,n,k, + total_row, total_col, + list_A_shim_queue, list_B_shim_queue, + list_C_shim_queue, + IS_B_ROW_MAJOR, ENABLE_AXI4, B_in_K_N_block_col_major_order, + ADD_BIAS, + valid_A_MT_shimtile_index, seq, shim_tiles, + REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH, + DH + ); + + int max_C_remain = 0; + for( auto li: list_C_shim_queue){ + + max_C_remain = std::max(max_C_remain, li); + } + + for(size_t k = 0; k< max_C_remain; k++){ + for(size_t shim_index = 0; shim_index < total_col; shim_index++){ + + if(list_A_shim_queue.at(shim_index) > 0){ + seq.npu_dma_wait( + shim_tiles.at(shim_index), + MM2S, + it_channel_0 + ); + list_A_shim_queue.at(shim_index)--; + } + if(list_B_shim_queue.at(shim_index) > 0){ + seq.npu_dma_wait( + shim_tiles.at(shim_index), + MM2S, + it_channel_1 + ); + list_B_shim_queue.at(shim_index)--; + } + if(list_C_shim_queue.at(shim_index) > 0){ + seq.npu_dma_wait( + shim_tiles.at(shim_index), + S2MM, + it_channel_0 + + ); + list_C_shim_queue.at(shim_index)--; + } + } + } + + seq.cmds2seq(); +} + +#endif diff --git a/src/detail/gemma4e_npu/reorder_cpy.hpp b/src/detail/gemma4e_npu/reorder_cpy.hpp new file mode 100644 index 000000000..f1456bedb --- /dev/null +++ b/src/detail/gemma4e_npu/reorder_cpy.hpp @@ -0,0 +1,37 @@ +#pragma once +#include "typedef.hpp" + +inline void reorder_cpy(u8 *dst, buffer &src, const int col, const int vertical_blocks = 2) +{ + const int a_block_size = 32 * 256 * 5 / 8; + const int blocks_per_row = col / 256; + + const int rows = src.size() / a_block_size / blocks_per_row; + + u8 *dst_ptr = dst; + std::vector src_ptr(vertical_blocks); + for (int i = 0; i < vertical_blocks; i++) + { + src_ptr[i] = src.data() + i * a_block_size * blocks_per_row; + } + for (int r = 0; r < rows; r += vertical_blocks) + { + for (int c = 0; c < blocks_per_row; c++) + { + for (int i = 0; i < vertical_blocks; i++) + { + memcpy(dst_ptr, src_ptr[i], a_block_size); + dst_ptr += a_block_size; + src_ptr[i] += a_block_size; + } + } + for (int i = 0; i < vertical_blocks; i++) + { + src_ptr[i] += (vertical_blocks - 1) * a_block_size * blocks_per_row; + if (src_ptr[i] + a_block_size * blocks_per_row > src.end()) + { + src_ptr[i] = src.data(); // useless padding + } + } + } +} diff --git a/src/detail/gemma4e_npu/rot_pos_emb.cpp b/src/detail/gemma4e_npu/rot_pos_emb.cpp new file mode 100644 index 000000000..4045b08ae --- /dev/null +++ b/src/detail/gemma4e_npu/rot_pos_emb.cpp @@ -0,0 +1,219 @@ +#include "rot_pos_emb.hpp" +#include +#include +#include +#include +#include +#include +#include +#include "avx512_util.hpp" +#include +#include + +void generate_gemma4_audio_rotary_pos_emb( + int hidden_size, + int attention_chunk_size, + int attention_context_left, + int attention_context_right, + + std::vector& position_embedding // shape of [13, hidden_size] +){ + + int context_size = attention_chunk_size + (attention_context_left-1) + attention_context_right; + + float min_timescale = 1.0; + float max_timescale = 10000.0; + + int num_timescales = hidden_size / 2; + float log_timescale_increment = std::log( + + max_timescale/min_timescale + + )/ std::max( num_timescales - 1, 1) ; + + std::vector inv_timescales(num_timescales); + + for(int i = 0; i < num_timescales; i++){ + inv_timescales[i] = min_timescale * std::exp(i * -log_timescale_increment); + } + + for(int i = 0; i <= 12; i++){ + int pos_val = 12 - i; + + for(int j = 0; j < num_timescales; j++){ + float angle = pos_val * inv_timescales[j]; + + position_embedding[i *hidden_size + j] = bf16(std::sin(angle)); + position_embedding[i *hidden_size + j + num_timescales] = bf16(std::cos(angle)); + } + } +} + +/** + * Apply rotary position embeddings to Q and K in-place within mm_res buffer. + * New Layout: [3 * num_heads, seq_len, head_dim] + * * @param mm_res_ptr Pointer to buffer of shape [3 * num_heads, seq_len, head_dim] + * Structure: [Block Q (Heads 0..H-1)] | [Block K (Heads 0..H-1)] | [Block V...] + * @param cos_emb_ptr Pointer to cosine embeddings of shape [seq_len, head_dim] (float) + * @param sin_emb_ptr Pointer to sine embeddings of shape [seq_len, head_dim] (float) + * @param seq_len The L_Seq dimension size + * @param hidden_size Hidden dimension size (total across all heads) + * @param num_heads Number of attention heads + */ +void generate_gemma4_vision_rotary_pos_emb( + const std::vector>& grid_pairs_per_image, + const std::vector& seq_len_per_image, + const std::vector& start_seq_len_index_per_image, + int seq_len_padded, + int head_dim, + float theta, + float scale, + std::vector& cos_emb, + std::vector& sin_emb +) { + const int spatial_dim = head_dim / 2; // 32 for head_dim=64 + const int inv_freq_len = spatial_dim / 2; // 16 + + // inv_freq[j] = 1.0 / (theta ^ (j*2 / spatial_dim)) + std::vector inv_freq(inv_freq_len); + for (int j = 0; j < inv_freq_len; j++) { + inv_freq[j] = 1.0f / std::pow(theta, (float)(j * 2) / (float)spatial_dim); + } + + cos_emb.assign(seq_len_padded * head_dim, bf16(0.0f)); + sin_emb.assign(seq_len_padded * head_dim, bf16(0.0f)); + + for (int img = 0; img < (int)grid_pairs_per_image.size(); img++) { + const int num_patches = seq_len_per_image[img]; + const int start = start_seq_len_index_per_image[img]; + const auto& grid_pairs = grid_pairs_per_image[img]; + int compact_patch_idx = 0; + + for (int pair_idx = 0; pair_idx + 1 < (int)grid_pairs.size() && compact_patch_idx < num_patches; pair_idx += 2) { + const int x_val = grid_pairs[pair_idx]; + const int y_val = grid_pairs[pair_idx + 1]; + + // Gemma4 position ids are padded with (-1, -1). The encoder state is compacted to + // valid patches only, so rotary embeddings must compact the valid coordinates too. + if (x_val < 0 || y_val < 0) { + continue; + } + + const int global_s = start + compact_patch_idx; + + bf16* cos_row = cos_emb.data() + global_s * head_dim; + bf16* sin_row = sin_emb.data() + global_s * head_dim; + + for (int j = 0; j < inv_freq_len; j++) { + const float x_freq = x_val * inv_freq[j]; + const float y_freq = y_val * inv_freq[j]; + + // x-axis: channels [0..inv_freq_len-1] and [inv_freq_len..spatial_dim-1] (duplication) + cos_row[j] = bf16(std::cos(x_freq) * scale); + cos_row[j + inv_freq_len]= bf16(std::cos(x_freq) * scale); + sin_row[j] = bf16(std::sin(x_freq) * scale); + sin_row[j + inv_freq_len]= bf16(std::sin(x_freq) * scale); + + // y-axis: channels [spatial_dim..spatial_dim+inv_freq_len-1] and [spatial_dim+inv_freq_len..head_dim-1] + cos_row[spatial_dim + j] = bf16(std::cos(y_freq) * scale); + cos_row[spatial_dim + j + inv_freq_len]= bf16(std::cos(y_freq) * scale); + sin_row[spatial_dim + j] = bf16(std::sin(y_freq) * scale); + sin_row[spatial_dim + j + inv_freq_len]= bf16(std::sin(y_freq) * scale); + } + + compact_patch_idx++; + } + + assert(compact_patch_idx == num_patches); + } +} + +void apply_multidimensional_rope( + bf16* qkv, + const bf16* cos_emb, + const bf16* sin_emb, + int seq_len, + int num_head, + int head_dim, + int ndim +) { + // Python: num_rotated_channels_per_dim = 2 * (head_dim // (2 * ndim)) + const int spatial_dim = 2 * (head_dim / (2 * ndim)); // channels per spatial dim (32 for head_dim=64, ndim=2) + const int quarter_dim = spatial_dim / 2; // rotate_half midpoint (16 for spatial_dim=32) + +#ifdef __AVX512F__ + // Fast path: AVX-512, processes 16 bf16 values per register. + // Requires quarter_dim == 16 (head_dim == 64). + if (quarter_dim == 16) { + for (int s = 0; s < seq_len; s++) { + const bf16* cos_row = cos_emb + s * head_dim; + const bf16* sin_row = sin_emb + s * head_dim; + + // Load cos/sin for x-spatial dim (cos[0..15], duplicated at [16..31]) + const __m512 cos_x = load_bfloat16_to_m512(cos_row); + const __m512 sin_x = load_bfloat16_to_m512(sin_row); + // Load cos/sin for y-spatial dim (cos[32..47], duplicated at [48..63]) + const __m512 cos_y = load_bfloat16_to_m512(cos_row + spatial_dim); + const __m512 sin_y = load_bfloat16_to_m512(sin_row + spatial_dim); + + bf16* row = qkv + s * num_head * head_dim; + for (int h = 0; h < num_head; h++) { + bf16* x = row + h * head_dim; + + // --- x-spatial half: channels [0..spatial_dim-1] --- + const __m512 x_lo = load_bfloat16_to_m512(x); + const __m512 x_hi = load_bfloat16_to_m512(x + quarter_dim); + // out_lo = x_lo * cos_x - x_hi * sin_x + const __m512 out_lo = _mm512_fmsub_ps(x_lo, cos_x, _mm512_mul_ps(x_hi, sin_x)); + // out_hi = x_hi * cos_x + x_lo * sin_x + const __m512 out_hi = _mm512_fmadd_ps(x_hi, cos_x, _mm512_mul_ps(x_lo, sin_x)); + store_m512_to_bfloat16_rne(x, out_lo); + store_m512_to_bfloat16_rne(x + quarter_dim, out_hi); + + // --- y-spatial half: channels [spatial_dim..head_dim-1] --- + bf16* y = x + spatial_dim; + const __m512 y_lo = load_bfloat16_to_m512(y); + const __m512 y_hi = load_bfloat16_to_m512(y + quarter_dim); + // out_lo = y_lo * cos_y - y_hi * sin_y + const __m512 y_out_lo = _mm512_fmsub_ps(y_lo, cos_y, _mm512_mul_ps(y_hi, sin_y)); + // out_hi = y_hi * cos_y + y_lo * sin_y + const __m512 y_out_hi = _mm512_fmadd_ps(y_hi, cos_y, _mm512_mul_ps(y_lo, sin_y)); + store_m512_to_bfloat16_rne(y, y_out_lo); + store_m512_to_bfloat16_rne(y + quarter_dim, y_out_hi); + } + } + return; + } +#endif // __AVX512F__ + + // Scalar fallback: works for any head_dim divisible by 4. + for (int s = 0; s < seq_len; s++) { + const bf16* cos_row = cos_emb + s * head_dim; + const bf16* sin_row = sin_emb + s * head_dim; + + bf16* row = qkv + s * num_head * head_dim; + for (int h = 0; h < num_head; h++) { + bf16* x = row + h * head_dim; + + // x-spatial half [0..spatial_dim-1] + for (int j = 0; j < quarter_dim; j++) { + const float x_lo = float(x[j]); + const float x_hi = float(x[j + quarter_dim]); + const float c = float(cos_row[j]); + const float s_ = float(sin_row[j]); + x[j] = bf16(x_lo * c - x_hi * s_); + x[j + quarter_dim]= bf16(x_hi * c + x_lo * s_); + } + + // y-spatial half [spatial_dim..head_dim-1] + for (int j = 0; j < quarter_dim; j++) { + const float y_lo = float(x[spatial_dim + j]); + const float y_hi = float(x[spatial_dim + j + quarter_dim]); + const float c = float(cos_row[spatial_dim + j]); + const float s_ = float(sin_row[spatial_dim + j]); + x[spatial_dim + j] = bf16(y_lo * c - y_hi * s_); + x[spatial_dim + j + quarter_dim]= bf16(y_hi * c + y_lo * s_); + } + } + } +} diff --git a/src/detail/gemma4e_npu/rot_pos_emb.hpp b/src/detail/gemma4e_npu/rot_pos_emb.hpp new file mode 100644 index 000000000..7acc4c69d --- /dev/null +++ b/src/detail/gemma4e_npu/rot_pos_emb.hpp @@ -0,0 +1,82 @@ +#pragma once +#include +#include +#include "typedef.hpp" +#include +#include + +void generate_gemma4_audio_rotary_pos_emb( + int hidden_size, + int attention_chunk_size, + int attention_context_left, + int attention_context_right, + + std::vector& position_embedding // shape of [13, hidden_size] +); + +/** + * Generate Gemma4 vision rotary position embeddings (cos and sin). + * + * Python equivalent: Gemma4VisionRotaryEmbedding.forward(hidden_states, pixel_position_ids) + * + * For each patch s with (x_val, y_val): + * spatial_dim = head_dim / 2 + * inv_freq[j] = 1 / (theta ^ (j*2 / spatial_dim)) for j in [0, spatial_dim/2) + * emb_x[j] = x_val * inv_freq[j] (duplicated: positions [0..spatial_dim-1]) + * emb_y[j] = y_val * inv_freq[j] (duplicated: positions [spatial_dim..head_dim-1]) + * cos_row = [cos(emb_x)*scale, cos(emb_y)*scale] of length head_dim + * + * Output shape: [seq_len_padded, head_dim] (padded rows are zero) + * + * @param grid_pairs_per_image per-image flat (x,y) pairs [img][s*2] + * @param seq_len_per_image unpadded patch count per image + * @param start_seq_len_index_per_image cumulative start offset per image + * @param seq_len_padded padded total sequence length + * @param head_dim GEMMA4E_VISION_HEAD_DIM (e.g. 64) + * @param theta GEMMA4E_ROPE_THETA (e.g. 100.0) + * @param scale attention_scaling (typically 1.0) + * @param cos_emb output [seq_len_padded * head_dim] (resized, bf16) + * @param sin_emb output [seq_len_padded * head_dim] (resized, bf16) + */ +void generate_gemma4_vision_rotary_pos_emb( + const std::vector>& grid_pairs_per_image, + const std::vector& seq_len_per_image, + const std::vector& start_seq_len_index_per_image, + int seq_len_padded, + int head_dim, + float theta, + float scale, + std::vector& cos_emb, + std::vector& sin_emb +); + +/** + * Apply multidimensional rotary position embeddings (RoPE) in-place. + * + * C++ equivalent of Gemma4's apply_multidimensional_rope with ndim=2, followed + * by apply_rotary_pos_emb (rotate_half variant). + * + * The head_dim is split into two spatial halves: + * x-spatial: channels [0 .. spatial_dim-1] (spatial_dim = head_dim/2) + * y-spatial: channels [spatial_dim .. head_dim-1] + * Standard RoPE is applied independently to each half: + * out_lo = x_lo * cos - x_hi * sin + * out_hi = x_hi * cos + x_lo * sin + * where _lo/_hi refer to the lower and upper quarter of each spatial half. + * + * @param qkv Pointer to buffer of shape [seq_len, num_head, head_dim] (bf16, in-place) + * @param cos_emb Pointer to cosine embeddings of shape [seq_len_padded, head_dim] (bf16) + * @param sin_emb Pointer to sine embeddings of shape [seq_len_padded, head_dim] (bf16) + * @param seq_len Number of valid (non-padded) sequence positions to process + * @param num_head Number of attention heads + * @param head_dim Head dimension; must be divisible by 4 (64 for Gemma4e vision) + */ +void apply_multidimensional_rope( + bf16* qkv, + const bf16* cos_emb, + const bf16* sin_emb, + int seq_len, + int num_head, + int head_dim, + int ndim = 2 +); diff --git a/src/detail/gemma4e_npu/seq_gen.hpp b/src/detail/gemma4e_npu/seq_gen.hpp new file mode 100644 index 000000000..18929a437 --- /dev/null +++ b/src/detail/gemma4e_npu/seq_gen.hpp @@ -0,0 +1,327 @@ +#pragma once +#include "npu_utils/npu_instr_utils.hpp" +#include "npu_sequences/image_attention_sequence.hpp" + +void _gen_mha_engine_seq(npu_sequence* seq, const uint32_t L_padded, const uint32_t L, + const int QKV_buffer_offset_bf, + const int O_buffer_offset_bf +){ + constexpr int DQ = 64 * 16; + constexpr int DK = 64 * 16; + constexpr int DV = DK; + constexpr int DH = 64; + int K_OFFSET = DQ; + int V_OFFSET = DQ + DK; + int D_TOTAL = DQ + DK + DV; + constexpr int total_cols = 8; + constexpr int total_rows = 4; + constexpr npu_tiles IT[8] = {IT0, IT1, IT2, IT3, IT4, IT5, IT6, IT7}; + + // each function call init a new list + int k_ping_pong_flag[8] = {0,0,0,0,0,0,0,0}; + int v_ping_pong_flag[8] = {0,0,0,0,0,0,0,0}; + int k_bd_queue[8] = {0,0,0,0,0,0,0,0}; + int v_bd_queue[8] = {0,0,0,0,0,0,0,0}; + + constexpr int lc = 32; + assert(L_padded % (32) == 0); + constexpr int l_qk_mha_address = 21760; + constexpr int l_kv_mha_address = 28672; + // seq->clear_cmds(); + + for (int row = 2; row < total_rows + 2; row++){ + for (int col = 0; col < 4; col++){ + npu_tiles tile_qk = get_tile(row, col * 2); + npu_tiles tile_kv = get_tile(row, col * 2 + 1); + seq->rtp_write(tile_qk, l_qk_mha_address, L); // 32 is lk + seq->rtp_write(tile_kv, l_kv_mha_address, L); // 32 is lk + } + } + + for (int round = 0; round < L_padded / 32; round++){ + int bd_offset = (round % 2) * 8; + for (int col = 0; col < 4; col++){ + // receive y + size_t y_offset = round * lc * DQ + col * 4 * DH + O_buffer_offset_bf; // y does not have history + seq->npu_dma_memcpy_nd( + 2, 0, + S2MM, IT[col * 2 + 1], + (npu_bd_id)(bd_offset + 0), it_channel_0, + {0, 0, 0, (uint32_t)y_offset}, + {1, 1, (uint32_t)lc, (uint32_t)DH * 4}, + {0, 0, (uint32_t)DQ, 1}, + -1, 0, true + ); + // send q + size_t q_offset = round * lc * D_TOTAL + col * 4 * DH +QKV_buffer_offset_bf; // q does not have history + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[col * 2], + (npu_bd_id)(bd_offset + 1), it_channel_0, + {0, 0, 0, (uint32_t)q_offset}, + {1, 1, (uint32_t)lc, (uint32_t)DH * 4}, + {0, 0, (uint32_t)D_TOTAL, 1}, + -1, 0, false + ); + + uint32_t L_padded_div_32 = L_padded / 32; + uint32_t max_L_per_chunk = 512; + // send k + size_t k_offset = col * 4 * DH + K_OFFSET + QKV_buffer_offset_bf; + size_t v_offset = col * 4 * DH + V_OFFSET + QKV_buffer_offset_bf; + for(int chunk_idx = 0; chunk_idx < L_padded_div_32; chunk_idx += max_L_per_chunk){ + + uint32_t cur_chunk_seqlen = std::min(max_L_per_chunk, L_padded_div_32 - chunk_idx); + int k_shim_id = col*2; + int k_bd_offset ; + if(k_ping_pong_flag[k_shim_id] == 0){ + k_bd_offset = 0; + k_ping_pong_flag[k_shim_id] = 1; + }else{ + k_bd_offset = 8; + k_ping_pong_flag[k_shim_id] = 0; + } + if(k_bd_queue[k_shim_id] == 2){ + // wait k + seq->npu_dma_wait( + IT[k_shim_id], + MM2S, + it_channel_1 + ); + k_bd_queue[k_shim_id]--; + } + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[k_shim_id], + (npu_bd_id)(k_bd_offset + 2), it_channel_1, + {0, 0, 0, (uint32_t)k_offset + chunk_idx * D_TOTAL * 32}, //row major of q, k, v per row + {1, cur_chunk_seqlen, (uint32_t)32, (uint32_t)DH * 4}, + {0, (uint32_t)32 * D_TOTAL, (uint32_t)D_TOTAL, 1}, + -1, 0, true, + aggressive_cache + ); + k_bd_queue[k_shim_id]++; + + int v_shim_id = col * 2 + 1; + int v_bd_offset ; + if(v_ping_pong_flag[v_shim_id] == 0){ + v_bd_offset = 0; + v_ping_pong_flag[v_shim_id] = 1; + }else{ + v_bd_offset = 8; + v_ping_pong_flag[v_shim_id] = 0; + } + if(v_bd_queue[v_shim_id] == 2){ + // wait v + seq->npu_dma_wait( + IT[v_shim_id], + MM2S, + it_channel_1 + ); + v_bd_queue[v_shim_id]--; + } + seq->npu_dma_memcpy_nd( + 2, 1, + MM2S, IT[v_shim_id], + (npu_bd_id)(v_bd_offset + 3), it_channel_1, + {0, 0, 0, (uint32_t)v_offset + chunk_idx * D_TOTAL * 32}, + {1, cur_chunk_seqlen, (uint32_t)32, (uint32_t)DH * 4}, + {0, (uint32_t)32 * D_TOTAL, (uint32_t)D_TOTAL, 1}, + -1, 0, true, + aggressive_cache + ); + v_bd_queue[v_shim_id]++; + } + } // col loop + + if (round > 0){ + for (int col = 0; col < 4; col++){ + seq->npu_dma_wait( + IT[col * 2 + 1], + S2MM, + it_channel_0 + ); + } + } + } + + int max_k_queue_remaing = -1; //should be same with v + for(int col = 0; col < 8; col++){ + if( k_bd_queue[col] > max_k_queue_remaing){ + max_k_queue_remaing = k_bd_queue[col]; + } + } + + for(int queue_size = 0; queue_size 0){ + seq->npu_dma_wait( + IT[col], + MM2S, + it_channel_1 + ); + k_bd_queue[col]--; + } + if(v_bd_queue[col] > 0){ + seq->npu_dma_wait( + IT[col], + MM2S, + it_channel_1 + ); + v_bd_queue[col]--; + } + } + } + + for (int col = 0; col < 4; col++){ + seq->npu_dma_wait( + IT[col * 2 + 1], + S2MM, + it_channel_0 + ); + } + // seq->cmds2seq(); +} + +//deprecated +void gen_mha_main( + + npu_sequence* seq, + std::vector> image_grid_thw, + int pad_requirement_for_attention, + int QWEN3_5_VISION_HIDDEN_SIZE + +){ + + // support of multiple batch mha + + seq->clear_cmds(); + seq->npu_preemption(0); + + auto round_up_to_multiple_lambda = [](int x, int multiple) -> int { + return ((x + multiple - 1) / multiple) * multiple; + }; + + int cur_seq_len = 0; + for(int b = 0; b < image_grid_thw.size(); b++){ + + int start_seq_len = cur_seq_len; + int end_seq_len = 1; + for(int v: image_grid_thw[b]){ + end_seq_len *= v; + } + + _gen_mha_engine_seq( + seq, round_up_to_multiple_lambda(end_seq_len, pad_requirement_for_attention), + end_seq_len, + start_seq_len * (3*QWEN3_5_VISION_HIDDEN_SIZE), //q,k,v offser + start_seq_len* QWEN3_5_VISION_HIDDEN_SIZE// offset + + ); + cur_seq_len += end_seq_len; + } + + seq->cmds2seq(); +} + +void gen_mha_vision_attention( + + npu_sequence* seq, + std::vector &seq_len_per_image, + int vision_L_padded_requirement_for_attention, + int vision_S_padded_requirement_for_attention, + uint32_t vision_num_of_columns, + uint32_t vision_num_of_rows, + uint32_t vision_CU_mode, + uint32_t vision_LQ_per_CT, + uint32_t vision_LK_per_CT, + uint32_t vision_LQ_internal, + uint32_t vision_LK_internal, + bool REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH, + int VISION_HIDDEN_SIZE, + int Padded_VISION_HIDDEN_SIZE, + int VISION_HEAD_DIM, + int VISION_NUM_ATTENTION_HEADS +){ + + if(vision_S_padded_requirement_for_attention >vision_L_padded_requirement_for_attention ){ + std::cerr << "vision S padded requirement greater than L padded requirement"<< std::endl; + exit(1); + } + + // support of multiple batch mha + + auto round_up_to_multiple_lambda = [](int x, int multiple) -> int { + return ((x + multiple - 1) / multiple) * multiple; + }; + + std::vector L_seq_list; + std::vector S_seq_list; + std::vector S_seq_padded_list; + + std::vector Q_batch_offset_list; + std::vector K_batch_offset_list; + std::vector V_batch_offset_list; + std::vector O_batch_offset_list; + + int cur_seq_len = 0; + for(int b = 0; b < seq_len_per_image.size(); b++){ + + int start_seq_len = cur_seq_len; + int end_seq_len = seq_len_per_image[b]; + + uint32_t l_seq_padded = round_up_to_multiple_lambda(end_seq_len, vision_L_padded_requirement_for_attention);// same for L, and S + uint32_t s_seq_padded = round_up_to_multiple_lambda(end_seq_len, vision_S_padded_requirement_for_attention); + L_seq_list.push_back(l_seq_padded); + S_seq_list.push_back(end_seq_len); + S_seq_padded_list.push_back(s_seq_padded); + + // but since + if(!REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH){ + Q_batch_offset_list.push_back(start_seq_len * (Padded_VISION_HIDDEN_SIZE) ); + K_batch_offset_list.push_back(start_seq_len * Padded_VISION_HIDDEN_SIZE); + V_batch_offset_list.push_back(start_seq_len * Padded_VISION_HIDDEN_SIZE); + O_batch_offset_list.push_back(start_seq_len * Padded_VISION_HIDDEN_SIZE); + }else{ + std::cout << "Currently the code only support REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH = false, please set it to false" << std::endl; + exit(-1); + // // The buffer is now in [3*QWEN3_VISION_NUM_HEAD, batch, L_Seq_per_batch, parent_npu_ptr->GEMMA4E_VISION_HEAD_DIM] + } + + cur_seq_len += end_seq_len; + } + + // because q, k, v buffer all in on buffer after MM + constexpr int ARG_Q = 1; + constexpr int ARG_K = 2; + constexpr int ARG_V = 3; + constexpr int ARG_O = 0; + + AttentionConfig config = { + (uint32_t)VISION_HEAD_DIM, (uint32_t)VISION_NUM_ATTENTION_HEADS, + vision_LQ_per_CT, vision_LK_per_CT, + vision_LQ_internal, vision_LK_internal, + vision_num_of_columns, vision_num_of_rows, + vision_CU_mode, + (uint32_t)seq_len_per_image.size(), + + ARG_Q, ARG_K, ARG_V, ARG_O, + + Padded_VISION_HIDDEN_SIZE, Padded_VISION_HIDDEN_SIZE, Padded_VISION_HIDDEN_SIZE, + Padded_VISION_HIDDEN_SIZE, + REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH, + 0,0,0 // since this only vailid if REORDER_OUTPUT_FROM_M_N_TO_NUM_DH_M_DH == true, and currently we only support false, so set it to -1 to avoid misuse + }; + + BatchMetadata batch_data = { + L_seq_list, S_seq_list, S_seq_padded_list, + Q_batch_offset_list, K_batch_offset_list, V_batch_offset_list, O_batch_offset_list + }; + + setup_SHM_configuration( + *seq, + config, + batch_data, + 1.0f// NOTE: very special for gemma4e + ); +} diff --git a/src/detail/lm_head/lm_head.cpp b/src/detail/lm_head/lm_head.cpp new file mode 100644 index 000000000..e811cd54b --- /dev/null +++ b/src/detail/lm_head/lm_head.cpp @@ -0,0 +1,191 @@ +#include "lm_head_detail.hpp" +#include "utils/utils.hpp" + +LMHead::Impl::Impl(LM_Config config, npu_xclbin_manager *npu) : config(config), npu(npu){ + vocab_size = config.get("vocab_size"); + hidden_size = config.get("hidden_size"); + vocab_size_padded = (vocab_size + padding_size - 1) / padding_size * padding_size; + chunk_size = vocab_size_padded / split_apps; + if (hidden_size % 256 != 0){ + hidden_size = (hidden_size / 256 + 1) * 256; + } + blocks_per_row = hidden_size / 256; + + lm_head_app_manager = npu->register_xclbin(utils::path_join(config.exec_path, "xclbins", config.model_name, "lm_head.xclbin")); + assert(blocks_per_row <= 64); + // assert(a_block_size * blocks_per_row <= 1024); + apps.resize(split_apps); + for (int i = 0; i < split_apps; i++){ + apps[i] = lm_head_app_manager->create_app(); + } + if (!this->npu->is_preemption_enabled()){ + lm_head_run = lm_head_app_manager->create_runlist(); + } + + lm_head_x = apps[0].create_bo_buffer(hidden_size); + lm_head_y = apps[0].create_bo_buffer(vocab_size_padded); + lm_head_y.memset((bf16)0.0f); + lm_head_y.sync_to_device(); + + lm_head_weights.resize(split_apps); + for (int i = 0; i < split_apps; i++){ + lm_head_weights[i] = apps[i].create_bo_buffer(chunk_size * hidden_size * 5 / 8); + } + + _generate_seq(); + if (!this->npu->is_preemption_enabled()){ + for (int i = 0; i < split_apps; i++){ + lm_head_run.add(apps[i].create_run(lm_head_y, lm_head_weights[i], lm_head_x)); + } + } + + logits = buffer(lm_head_y.data(), vocab_size); + x_exposed = buffer(this->lm_head_x.data(), hidden_size); +} + +void LMHead::Impl::_generate_seq(){ + for (int i = 0; i < split_apps; i++){ + npu_sequence& seq = *this->apps[i].seq(); + + seq.npu_preemption(0); + assert(chunk_size / m / cores > 0); + assert(chunk_size % (m * cores) == 0); + assert(chunk_size % (columns * 4 * m) == 0); + // x_in + seq.npu_dma_memcpy_nd( + sizeof(bf16), lm_head_x_arg_idx, + MM2S, lm_head_tiles[0], bd_0, it_channel_0, + {0, 0, 0, 0}, + {1, 1, 1, (uint32_t)hidden_size}, + {0, 0, 0, 1}, + -1, 0, false + ); + + // y_out + for (uint32_t col = 0; col < columns; col++){ + uint32_t y_offset = 4 * m * col; + seq.npu_dma_memcpy_nd( + sizeof(bf16), lm_head_y_arg_idx, + S2MM, lm_head_tiles[col], bd_2, it_channel_0, + {0, 0, 0, y_offset + i * chunk_size}, + {1, 1, chunk_size / columns / 4 / m, 4 * m}, + {0, 0, columns * 4 * m, 1}, + -1, 0, true + ); + } + + for (uint32_t round = 0; round < chunk_size / m / cores; round++){ + uint32_t bd_offset = (round % 2) * 8; + for (uint32_t col = 0; col < columns; col++){ + seq.npu_dma_memcpy_nd( + sizeof(uint32_t), + lm_head_w_arg_idx, + MM2S, + lm_head_tiles[col], + npu_bd_id(bd_1 + bd_offset), + it_channel_1, + {0, 0, 0, (round * cores + col * 4) * a_block_size * blocks_per_row}, + {blocks_per_row, 4, 8, a_block_size / 8}, + {a_block_size, a_block_size * blocks_per_row, a_block_size / 8, 1}, + -1, 0, true + ); + } + if (round > 0){ + for (uint32_t col = 0; col < columns; col++){ + seq.npu_dma_wait( + lm_head_tiles[col], + MM2S, + it_channel_1 + ); + } + } + } + for (uint32_t col = 0; col < columns; col++){ + seq.npu_dma_wait( + lm_head_tiles[col], + MM2S, + it_channel_1 + ); + } + for (uint32_t col = 0; col < columns; col++){ + seq.npu_dma_wait( + lm_head_tiles[col], + S2MM, + it_channel_0 + ); + } + seq.cmds2seq(); + apps[i].update_ctrl_seq(); + } +} + +void LMHead::Impl::load_weights(Q4NX& q4nx){ + buffer w_lm_head_w; + q4nx.load_weights(w_lm_head_w, "lm_head.weight"); + size_t w_size = lm_head_weights[0].size(); + u8* ptr = w_lm_head_w.data(); + for (int i = 0; i < split_apps - 1; i++){ + memcpy(lm_head_weights[i].data(), ptr, w_size); + ptr += w_size; + } + lm_head_weights[split_apps - 1].memset(0); + size_t remaining_size = w_lm_head_w.size() - w_size * (split_apps - 1); + memcpy(lm_head_weights[split_apps - 1].data(), ptr, remaining_size); + + for (int i = 0; i < split_apps; i++){ + lm_head_weights[i].sync_to_device(); + } +} + +void LMHead::Impl::execute(){ + this->lm_head_x.sync_to_device(); + if (!this->npu->is_preemption_enabled()){ + this->lm_head_run.execute(); + } + else{ + for (int i = 0; i < split_apps - 1; i++){ + apps[i](lm_head_y, lm_head_weights[i], lm_head_x); + } + this->final_run = apps[split_apps - 1].create_run(lm_head_y, lm_head_weights[split_apps - 1], lm_head_x); + this->final_run.start(); + } +} + +buffer LMHead::Impl::wait(){ + if (!this->npu->is_preemption_enabled()){ + this->lm_head_run.wait(); + } + else{ + this->final_run.wait(); + } + this->lm_head_y.sync_from_device(); + return this->logits; +} + +LMHead::Impl::~Impl() = default; + +// =============================================== +// Externals for LMHead +// =============================================== +LMHead::LMHead(LM_Config config, npu_xclbin_manager *npu){ + this->_impl = new Impl(config, npu); +} + +void LMHead::load_weights(Q4NX& q4nx){ + this->_impl->load_weights(q4nx); +} + +void LMHead::execute(){ + this->_impl->execute(); +} + +buffer LMHead::wait(){ + return this->_impl->wait(); +} + +buffer LMHead::x_exposed(){ + return this->_impl->x_exposed; +} +LMHead::~LMHead(){ + delete this->_impl; +} diff --git a/src/detail/lm_head/lm_head_detail.hpp b/src/detail/lm_head/lm_head_detail.hpp new file mode 100644 index 000000000..77ba526b1 --- /dev/null +++ b/src/detail/lm_head/lm_head_detail.hpp @@ -0,0 +1,45 @@ +#pragma once +#include "modules/lm_head.hpp" + +struct LMHead::Impl{ +public: + // Impl(){} + Impl(LM_Config config, npu_xclbin_manager *npu); + ~Impl(); + void load_weights(Q4NX& q4nx); + void execute(); + buffer wait(); + buffer x_exposed; + + static constexpr npu_tiles lm_head_tiles[] = {IT0, IT1, IT2, IT3, IT4, IT5, IT6, IT7}; + + static constexpr int lm_head_y_arg_idx = 0; + static constexpr int lm_head_w_arg_idx = 1; + static constexpr int lm_head_x_arg_idx = 2; + + static constexpr uint32_t columns = 8; + static constexpr uint32_t m = 32; + static constexpr int split_apps = 4; + static constexpr int padding_size = 4096; + static constexpr uint32_t a_block_size = 32 * 256 * 5 / 8 / 4; + static constexpr uint32_t cores = columns * 4; + LM_Config config; + npu_xclbin_manager *npu; + npu_app_manager* lm_head_app_manager; + + std::vector apps; + flm_rt::run final_run; + flm_rt::runlist lm_head_run; + std::vector> lm_head_weights; + buffer lm_head_y; + buffer lm_head_x; + uint32_t blocks_per_row; + + uint32_t vocab_size; + uint32_t vocab_size_padded; + uint32_t hidden_size; + uint32_t chunk_size; + buffer logits; + + void _generate_seq(); +}; diff --git a/src/detail/vision_common/norm.cpp b/src/detail/vision_common/norm.cpp new file mode 100644 index 000000000..6d17010e6 --- /dev/null +++ b/src/detail/vision_common/norm.cpp @@ -0,0 +1,184 @@ +#include +#include +#include +#include +#include +#include "typedef.hpp" +#include "utils/avx512_util.hpp" + +/** + * @brief AVX-512 optimized Layer Normalization with high precision accumulation. + * + * This implementation combines AVX-512 vectorization with numerical accuracy: + * - Uses AVX-512 for parallel processing (16 floats at a time) + * - Double precision for mean and variance accumulation + * - Handles bfloat16 input/output with float32 intermediate calculations + * + * @param N The size of the hidden dimension. + * @param hidden_states The input vector (bfloat16_t). + * @param weight The scale parameter (bfloat16_t). + * @param bias The bias parameter (bfloat16_t), can be nullptr if no bias. + * @param eps The epsilon value (float) to prevent division by zero (typically 1e-6). + * @param output The output vector (bfloat16_t). + */ +void layernorm_high_precision( + size_t N, + bf16* hidden_states, + bf16* weight, + bf16* bias, + float eps, + bf16* output) +{ + constexpr size_t SIMD_WIDTH = 16; // AVX-512 processes 16 floats at once + const size_t vec_count = N / SIMD_WIDTH; + const size_t remainder = N % SIMD_WIDTH; + + // Step 1: Calculate mean using float precision for accumulation + __m512 sum_vec = _mm512_setzero_ps(); // 16 floats at a time + float sum = 0.0f; + + // Vectorized sum calculation + for (size_t i = 0; i < vec_count; ++i) { + size_t idx = i * SIMD_WIDTH; + + // Load 16 bfloat16 values and convert to float + __m512 vals = load_bfloat16_to_m512(hidden_states + idx); + + // Accumulate in float precision + sum_vec = _mm512_add_ps(sum_vec, vals); + } + + // Reduce vector sum to scalar + sum = _mm512_reduce_add_ps(sum_vec); + + // Handle remainder elements + for (size_t i = vec_count * SIMD_WIDTH; i < N; ++i) { + sum += static_cast(hidden_states[i]); + } + + float mean = sum / static_cast(N); + __m512 mean_vec = _mm512_set1_ps(mean); + + // Step 2: Calculate variance using float precision + __m512 var_vec = _mm512_setzero_ps(); + float var_sum = 0.0f; + + for (size_t i = 0; i < vec_count; ++i) { + size_t idx = i * SIMD_WIDTH; + + // Load and convert to float + __m512 vals = load_bfloat16_to_m512(hidden_states + idx); + + // Calculate (val - mean)^2 + __m512 diff = _mm512_sub_ps(vals, mean_vec); + __m512 diff_sq = _mm512_mul_ps(diff, diff); + + // Accumulate in float precision + var_vec = _mm512_add_ps(var_vec, diff_sq); + } + + var_sum = _mm512_reduce_add_ps(var_vec); + + // Handle remainder + for (size_t i = vec_count * SIMD_WIDTH; i < N; ++i) { + float val = static_cast(hidden_states[i]); + float diff = val - mean; + var_sum += diff * diff; + } + + float variance = var_sum / static_cast(N); + + // Step 3: Calculate 1/sqrt(variance + eps) + float std_inv = 1.0f / std::sqrt(variance + eps); + __m512 std_inv_vec = _mm512_set1_ps(std_inv); + + // Step 4: Normalize, scale, and shift + for (size_t i = 0; i < vec_count; ++i) { + size_t idx = i * SIMD_WIDTH; + + // Load input + __m512 vals = load_bfloat16_to_m512(hidden_states + idx); + + // Normalize: (x - mean) / std + __m512 normalized = _mm512_sub_ps(vals, mean_vec); + normalized = _mm512_mul_ps(normalized, std_inv_vec); + + // Load weight and scale + __m512 w = load_bfloat16_to_m512(weight + idx); + __m512 result = _mm512_mul_ps(normalized, w); + + // Add bias if provided + if (bias != nullptr) { + __m512 b = load_bfloat16_to_m512(bias + idx); + result = _mm512_add_ps(result, b); + } + + // Convert back to bfloat16 and store + store_m512_to_bfloat16_rne(output + idx, result); + } + + // Handle remainder elements + for (size_t i = vec_count * SIMD_WIDTH; i < N; ++i) { + float val = static_cast(hidden_states[i]); + float normalized = (val - mean) * std_inv; + float w = static_cast(weight[i]); + float result = normalized * w; + + if (bias != nullptr) { + float b = static_cast(bias[i]); + result += b; + } + + output[i] = static_cast(result); + } +} + +/** + * @brief Parallel Layer Normalization across multiple sequences with OpenMP. + * + * Applies layer normalization to multiple sequences in parallel using up to 4 threads. + * Uses chunked distribution (static scheduling) for better cache locality. + * Automatically disables parallelization for small sequence counts. + * + * @param seq_len Number of sequences to normalize. + * @param hidden_dim The size of the hidden dimension for each sequence. + * @param hidden_states_base Pointer to base of input sequences (bfloat16_t). + * @param weight The scale parameter (bfloat16_t). + * @param bias The bias parameter (bfloat16_t), can be nullptr if no bias. + * @param eps The epsilon value (float) to prevent division by zero. + * @param output_base Pointer to base of output sequences (bfloat16_t). + */ +void layernorm_parallel( + int seq_len, + size_t hidden_dim, + size_t hidden_dim_padded, + bf16* hidden_states_base, + bf16* weight, + bf16* bias, + float eps, + bf16* output_base) +{ + // Use OpenMP only if seq_len is large enough to benefit from parallelization + // Threshold: 8 sequences per thread minimum to avoid overhead + const int min_seq_per_thread = 8; + const int max_threads = 4; + const bool use_parallel = seq_len >= (min_seq_per_thread * 2); + + // Process each sequence in parallel with chunked distribution (static scheduling) + // This ensures thread-0 gets sequences [0, chunk_size), thread-1 gets [chunk_size, 2*chunk_size), etc. + // Better for cache locality than round-robin distribution + #pragma omp parallel for num_threads(max_threads) if(use_parallel) schedule(static) + for (int s = 0; s < seq_len; ++s) { + bf16* input_seq = hidden_states_base + s * hidden_dim_padded; + bf16* output_seq = output_base + s * hidden_dim_padded; + + layernorm_high_precision( + hidden_dim, + input_seq, + weight, + bias, + eps, + output_seq + ); + } +} diff --git a/src/include/flm_override.hpp b/src/include/flm_override.hpp new file mode 100644 index 000000000..2b0dc4dfe --- /dev/null +++ b/src/include/flm_override.hpp @@ -0,0 +1,30 @@ +#ifndef FLM_OVERRIDE_HPP +#define FLM_OVERRIDE_HPP + +/// \file flm_override.hpp +/// \brief Compile-time override points for a model's operator dispatch. +/// +/// `FLM_OVERRIDE(name, expr, ...)` expands to `expr`. A build that points +/// `FLM_OVERRIDES` at a header may redefine it to dispatch on `name`; every +/// other build emits the code it would have emitted without the annotation. +/// +/// g++ -DFLM_OVERRIDES='"iron_overrides.h"' ... +/// +/// `name` is a token, not a string, so it costs nothing. The trailing +/// arguments carry values an override needs and the call itself does not. The +/// default expansion discards them unevaluated, so they must be free of side +/// effects, and a value computed only to be passed here belongs inside the +/// annotation too, or `-Wall` reports it unused. +/// +/// A statement-position hook that only hands an override some context writes +/// `(void)0` as its expression. + +#ifdef FLM_OVERRIDES +#include FLM_OVERRIDES +#endif + +#ifndef FLM_OVERRIDE +#define FLM_OVERRIDE(name, expr, ...) (expr) +#endif + +#endif // FLM_OVERRIDE_HPP diff --git a/src/include/metrices.hpp b/src/include/metrices.hpp index f46469de7..9eb9fe270 100644 --- a/src/include/metrices.hpp +++ b/src/include/metrices.hpp @@ -89,9 +89,67 @@ inline error_metrics get_error_metrics(buffer& y, buffer& y_ref){ return metrics; } +inline error_metrics get_error_metrics(buffer& y, buffer& y_ref){ + assert(y.size() == y_ref.size()); + error_metrics metrics; + const int simd_width = 8; + + __m256 dot_product = _mm256_setzero_ps(); + __m256 y_square_sum = _mm256_setzero_ps(); + __m256 y_ref_square_sum = _mm256_setzero_ps(); + __m256 error_square_sum = _mm256_setzero_ps(); + __m256 abs_error_sum = _mm256_setzero_ps(); + __m256 abs_y_ref_sum = _mm256_setzero_ps(); + for (int i = 0; i < y.size(); i += simd_width){ + __m256 y_vec = _mm256_loadu_ps(y.data() + i); + __m256 y_ref_vec = _mm256_loadu_ps(y_ref.data() + i); + __m256 error_vec = _mm256_sub_ps(y_vec, y_ref_vec); + y_square_sum = _mm256_add_ps(y_square_sum, _mm256_mul_ps(y_vec, y_vec)); + y_ref_square_sum = _mm256_add_ps(y_ref_square_sum, _mm256_mul_ps(y_ref_vec, y_ref_vec)); + dot_product = _mm256_add_ps(dot_product, _mm256_mul_ps(y_vec, y_ref_vec)); + error_square_sum = _mm256_add_ps(error_square_sum, _mm256_mul_ps(error_vec, error_vec)); + abs_error_sum = _mm256_add_ps(abs_error_sum, _mm256_abs_ps(error_vec)); + abs_y_ref_sum = _mm256_add_ps(abs_y_ref_sum, _mm256_abs_ps(y_ref_vec)); + } + f32 temp[simd_width]; + _mm256_storeu_ps(temp, dot_product); + f32 dot_product_sum = temp[0] + temp[1] + temp[2] + temp[3] + + temp[4] + temp[5] + temp[6] + temp[7]; + _mm256_storeu_ps(temp, y_square_sum); + f32 y_square_sum_sum = temp[0] + temp[1] + temp[2] + temp[3] + + temp[4] + temp[5] + temp[6] + temp[7]; + _mm256_storeu_ps(temp, y_ref_square_sum); + f32 y_ref_square_sum_sum = temp[0] + temp[1] + temp[2] + temp[3] + + temp[4] + temp[5] + temp[6] + temp[7]; + _mm256_storeu_ps(temp, error_square_sum); + f32 error_square_sum_sum = temp[0] + temp[1] + temp[2] + temp[3] + + temp[4] + temp[5] + temp[6] + temp[7]; + _mm256_storeu_ps(temp, abs_error_sum); + f32 abs_error_sum_sum = temp[0] + temp[1] + temp[2] + temp[3] + + temp[4] + temp[5] + temp[6] + temp[7]; + _mm256_storeu_ps(temp, abs_y_ref_sum); + f32 abs_y_ref_sum_sum = temp[0] + temp[1] + temp[2] + temp[3] + + temp[4] + temp[5] + temp[6] + temp[7]; + + // Cosine Similarity + float cosine_similarity = dot_product_sum / (sqrt(y_square_sum_sum) * sqrt(y_ref_square_sum_sum)); + // Relative L1 + float relative_l1 = abs_error_sum_sum / abs_y_ref_sum_sum; + // RMSE + float rmse = sqrt(error_square_sum_sum / y.size()); + // Relative L2 + float relative_l2 = sqrt(error_square_sum_sum / y_ref_square_sum_sum); + metrics.CosineSimilarity = cosine_similarity; + metrics.RelativeL1 = relative_l1; + metrics.RMSE = rmse; + metrics.RelativeL2 = relative_l2; + return metrics; +} + /// \brief print the error metrics /// \param metrics the error metrics -inline void print_error_metrics(error_metrics metrics){ +inline void print_error_metrics(error_metrics metrics, std::string name = "Error Metrics"){ + header_print("info", name); header_print("info", "Cosine Similarity: " << metrics.CosineSimilarity); header_print("info", "Relative L1 : " << metrics.RelativeL1); header_print("info", "RMSE : " << metrics.RMSE); diff --git a/src/include/modules/dequant.hpp b/src/include/modules/dequant.hpp index 43596d08a..ba1d2982b 100644 --- a/src/include/modules/dequant.hpp +++ b/src/include/modules/dequant.hpp @@ -1,37 +1,45 @@ -/// \file dequant.hpp -/// \brief dequant class -/// \author FastFlowLM Team -/// \date 2025-06-24 -/// \version 0.9.10 -/// \note This is a header file for the dequant class -#pragma once -#include "lm_config.hpp" -#include "npu_utils/npu_instr_utils.hpp" - -/// \brief dequant class -/// \note This is a class for the dequant layer -class Dequant{ -public: - Dequant(){} - - /// \brief Constructor - /// \param config the configuration - /// \param xclbin_name the xclbin name - /// \param npu the npu manager - Dequant(LM_Config& config); - ~Dequant(); - - /// @brief generate the dequant sequence - /// @param seq: the sequence - /// @param D_in: input dimension of the projection weight - /// @param D_out: output dimension of the projection weight - /// @param weight_offset: the weight offset in byte - /// @param mode: dequant output mode - void generate_dequant_q4_1_seq(npu_sequence* seq, const uint32_t D_in, const uint32_t D_out, const uint32_t weight_offset, int mode); - void generate_dequant_q80_packed_in_q4nx_seq(npu_sequence* seq, const uint32_t D_in, const uint32_t D_out, const uint32_t weight_offset, int mode); -private: - struct Impl; - Impl* _impl; - -}; - +/// \file dequant.hpp +/// \brief dequant class +/// \author FastFlowLM Team +/// \date 2025-06-24 +/// \version 0.9.10 +/// \note This is a header file for the dequant class +#pragma once +#include "lm_config.hpp" +#include "npu_utils/npu_instr_utils.hpp" + +/// \brief dequant class +/// \note This is a class for the dequant layer +class Dequant{ +public: + Dequant(){} + + /// \brief Constructor + /// \param config the configuration + /// \param xclbin_name the xclbin name + /// \param npu the npu manager + Dequant(LM_Config& config); + ~Dequant(); + + typedef enum: int { + Q4_1 = 0, + Q8_0 = 1, + Q4_0 = 2 + } quant_block_t; + + void reorder_cpy(u8 *dst, buffer &src, + quant_block_t quant_block_type, + const int quant_matrix_row, + const int quant_matrix_col, + const int vertical_blocks=2, + const int vetrical_block_interleave_byte_size=-1); + void generate_dequant_q4_1_seq(npu_sequence* seq, const uint32_t D_in, const uint32_t D_out, const uint32_t weight_offset, int mode); + /// \brief same sequence as q4_1, for the 4.5 bpw (4608 byte) block + void generate_dequant_q4_0_seq(npu_sequence* seq, const uint32_t D_in, const uint32_t D_out, const uint32_t weight_offset, int mode); + void generate_dequant_q80_packed_in_q4nx_seq(npu_sequence* seq, const uint32_t D_in, const uint32_t D_out, const uint32_t weight_offset, int mode); +private: + struct Impl; + Impl* _impl; + +}; + diff --git a/src/include/npu_sequences/image_attention_sequence.hpp b/src/include/npu_sequences/image_attention_sequence.hpp new file mode 100644 index 000000000..299b57a98 --- /dev/null +++ b/src/include/npu_sequences/image_attention_sequence.hpp @@ -0,0 +1,1458 @@ +#ifndef __Attention_RUNTIME_SEQUENCE_HPP__ +#define __Attention_RUNTIME_SEQUENCE_HPP__ + + +//direct copy from gemma3 + +#include +#include // Required for std::max +#include +#include "npu_utils/npu_instr_utils.hpp" +#include "tensor_utils/q4_npu_eXpress.hpp" + + + + + +inline void update_ping_pong_flag(int &flag){ + if(flag ==0){ + flag =1; + }else{ + flag =0; + } +} + + +//ALL SHIMTILE agree on common bds for Q, K, V O that it is using +//for Q, used bd 0, 8 +//for k, used bd 1, 9 +//for V, used bd 2, 10 +//for O, used bd 3, 11 +inline int get_Q_bd_id(int ping_pong){ + if(ping_pong ==0){ + return 0; + }else{ + return 8; + } +} + +inline int get_K_bd_id(int ping_pong){ + if(ping_pong ==0){ + return 1; + }else{ + return 9; + } +} +inline int get_V_bd_id(int ping_pong){ + if(ping_pong ==0){ + return 2; + }else{ + return 10; + } +} +inline int get_O_bd_id(int ping_pong){ + if(ping_pong ==0){ + return 3; + }else{ + return 11; + } +} + + +inline void wait_DMA_queue_if_full( + npu_sequence &seq, + int cur_shim_idx, std::vector &SHIM_queue_counter, + npu_tiles cur_shimtile, dma_direction channel_direction, npu_it_channel it_channel, + int max_queue_size = 2 +){ + + //NOTE: it will not wait if max_queue_size is not achieve + if(SHIM_queue_counter.at(cur_shim_idx) > max_queue_size){ + std::cerr <<"SHIM DMA queue overflow"<< std::endl; + exit(1); + } + + if(SHIM_queue_counter.at(cur_shim_idx) == max_queue_size){ + // need to wait + seq.npu_dma_wait( + cur_shimtile, channel_direction, it_channel + ); + SHIM_queue_counter[cur_shim_idx]--; + } +} + + + +inline void force_wait_DMA_queue( + npu_sequence &seq, + int cur_shim_idx, std::vector &SHIM_queue_counter, + npu_tiles cur_shimtile, dma_direction channel_direction, npu_it_channel it_channel + +){ + + if (SHIM_queue_counter.at(cur_shim_idx) >0){ + // need to wait + seq.npu_dma_wait( + cur_shimtile, channel_direction, it_channel + ); + SHIM_queue_counter[cur_shim_idx]--; + } + +} + + +inline void increment_DMA_queue( + + npu_sequence &seq, + int cur_shim_idx, std::vector &SHIM_queue_counter, + int max_queue_size = 2 +){ + + if(SHIM_queue_counter.at(cur_shim_idx) < max_queue_size){ + SHIM_queue_counter[cur_shim_idx]++; + }else{ + std::cerr <<"SHIM DMA queue overflow"<< std::endl; + exit(1); + } + +} + + + + + + + +template +void process_SHM_Q_chunk_Q_repeat( + npu_sequence &seq, + int NUM_OF_CT_per_column, + int MT_KV_repeat, + int DH, + int NUM_OF_DH, + int head_idx_start, // current DH index that is being processed + + int SHIM_LQ_start_row, + int SHIM_S_SEQ_PADDED, + + std::vector> &CU_column, + int LQ_CU_chunk_size, + int LK_CU_chunk_size, + + int LQ_per_CT, + int LK_per_CT, + + std::vector &SHIM_Q_queue_counter, + std::vector &SHIM_K_queue_counter, + std::vector &SHIM_V_queue_counter, + std::vector &SHIM_O_queue_counter, + + std::vector &SHIM_Q_ping_pong, + std::vector &SHIM_K_ping_pong, + std::vector &SHIM_V_ping_pong, + std::vector &SHIM_O_ping_pong, + + uint32_t Q_PADDDED_D, uint32_t K_PADDED_D, uint32_t V_PADDED_D, uint32_t O_PADDDED_D, + uint32_t Q_offset, uint32_t K_offset, uint32_t V_offset, uint32_t O_offset, + + int ARG_Q, int ARG_K, int ARG_V, int ARG_O, + npu_it_channel it_for_send_Q, npu_it_channel it_for_send_KV, + npu_it_channel it_for_recv_O, + std::vector &shim_tiles +){ + + + + + uint32_t LQ_per_column = LQ_CU_chunk_size/(CU_column[0].size()); + + bool sanity_check = true; + + sanity_check &= SHIM_S_SEQ_PADDED%LK_CU_chunk_size ==0; + sanity_check &= MT_KV_repeat >1; + sanity_check &= LK_CU_chunk_size == LK_per_CT; + sanity_check &= LQ_CU_chunk_size % CU_column[0].size() ==0; + + sanity_check &= NUM_OF_CT_per_column*LQ_per_CT == LQ_per_column; + if(!sanity_check){ + std::cerr << "process_SHM_Q_chunk_Q_repeat sanity check failed"<< std::endl; + exit(1); + } + + + for(int cu_idx = 0; cu_idx = NUM_OF_DH){ + break; + } + + + uint32_t SHIM_O_offset = SHIM_LQ_start_row * (O_PADDDED_D) + DH_index * DH +O_offset; + + for(int idx = 0; idx < cur_CU_shimtile_cols.size(); idx++){ + + auto shim_col_idx = cur_CU_shimtile_cols.at(idx); + + + // now, receive the O + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_O_queue_counter, + shim_tiles.at(shim_col_idx), + S2MM, it_for_recv_O + ); + seq.npu_dma_memcpy_nd( + sizeof(T_out), + ARG_O, + S2MM, + shim_tiles.at(shim_col_idx), + static_cast(get_O_bd_id(SHIM_O_ping_pong.at(shim_col_idx))), + it_for_recv_O, + {0,0,0, (uint32_t)SHIM_O_offset + LQ_per_column*idx*(O_PADDDED_D)}, + {1,1, (uint32_t)LQ_per_column, (uint32_t)DH}, + {0,0, (uint32_t)O_PADDDED_D, 1}, + -1, 0, true + ); + update_ping_pong_flag(SHIM_O_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_O_queue_counter + ); + + + + + } + + } + + + for(int kv_chunk_idx = 0; kv_chunk_idx < SHIM_S_SEQ_PADDED/LK_CU_chunk_size; kv_chunk_idx++){ + + + + for(int cu_idx = 0; cu_idx = NUM_OF_DH){ + break; + } + + + + + + + uint32_t Shim_Q_offset = SHIM_LQ_start_row * (Q_PADDDED_D) + DH_index * DH +Q_offset; + + uint32_t SHIM_K_offset = DH_index * DH + K_offset; + uint32_t SHIM_V_offset = DH_index * DH + V_offset; + + for(int idx = 0; idx < cur_CU_shimtile_cols.size(); idx++){ + + auto shim_col_idx = cur_CU_shimtile_cols.at(idx); + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_Q_queue_counter, + shim_tiles.at(shim_col_idx), + MM2S, it_for_send_Q + ); + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_Q, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_Q_bd_id(SHIM_Q_ping_pong.at(shim_col_idx))), + it_for_send_Q, + {0,0,0, (uint32_t)Shim_Q_offset + LQ_per_column*idx*(Q_PADDDED_D)}, + {1,1, (uint32_t)LQ_per_column, (uint32_t)DH}, + {0,0, (uint32_t)Q_PADDDED_D, 1}, + -1, 0, true + + ); + update_ping_pong_flag(SHIM_Q_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_Q_queue_counter + ); + } + + + //only the first shimtile column in cur_CU_shimtile_cols will be used to send KV data + { + + uint32_t shim_col_idx = cur_CU_shimtile_cols.at(0); + + //NOTE: NO need to wait for k,since we only issue the v + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_V_queue_counter, + shim_tiles.at(shim_col_idx), + MM2S, it_for_send_KV + ); + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_K, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_K_bd_id(SHIM_K_ping_pong.at(shim_col_idx))), + it_for_send_KV, + {0,0,0, (uint32_t)SHIM_K_offset + kv_chunk_idx*LK_CU_chunk_size*(K_PADDED_D)}, + {1,1, (uint32_t)LK_CU_chunk_size, (uint32_t)DH}, + {0,0, (uint32_t)K_PADDED_D, 1}, + -1, 0, false, + aggressive_cache + + ); + update_ping_pong_flag(SHIM_K_ping_pong.at(shim_col_idx)); + + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_V, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_V_bd_id(SHIM_V_ping_pong.at(shim_col_idx))), + it_for_send_KV, + {0,0,0, (uint32_t)SHIM_V_offset + kv_chunk_idx*LK_CU_chunk_size*(V_PADDED_D)}, + {1,1, (uint32_t)LK_CU_chunk_size, (uint32_t)DH}, + {0,0, (uint32_t)V_PADDED_D, 1}, + -1, 0, true, + aggressive_cache + ); + update_ping_pong_flag(SHIM_V_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_V_queue_counter + ); + + } + + + + } + + + + + + } + + + + + + + + + +} + + + + + +template +void process_SHM_Q_chunk_Q_repeat_reorder_NUM_DH_L_DH( + npu_sequence &seq, + int NUM_OF_CT_per_column, + int MT_KV_repeat, + int DH, + int NUM_OF_DH, + int DH_index_start, // current DH index that is being processed + + int SHIM_LQ_start_row, + int SHIM_S_SEQ_PADDED, + + std::vector> &CU_column, + int LQ_CU_chunk_size, + int LK_CU_chunk_size, + + int LQ_per_CT, + int LK_per_CT, + + std::vector &SHIM_Q_queue_counter, + std::vector &SHIM_K_queue_counter, + std::vector &SHIM_V_queue_counter, + std::vector &SHIM_O_queue_counter, + + std::vector &SHIM_Q_ping_pong, + std::vector &SHIM_K_ping_pong, + std::vector &SHIM_V_ping_pong, + std::vector &SHIM_O_ping_pong, + + uint32_t O_PADDDED_D, + //[QWEN3_VISION_NUM_HEAD, batch, L_Seq_per_batch, QWEN3_VISION_HEAD_DIM] + uint32_t Q_total_seq_len_per_head, // Total sequence length per head for Q in layout[NUM_HEAD, batch, L_Seq_per_batch, HEAD_DIM]. Can be L_seq_padded*batch_size (uniform padding) or L_seq*(batch_size-1)+L_seq_padded (only last batch padded) + uint32_t K_total_seq_len_per_head, // Total sequence length per head for K in layout [NUM_HEAD, batch, L_Seq_per_batch, HEAD_DIM]. Can be L_seq_padded*batch_size (uniform padding) or L_seq*(batch_size-1)+L_seq_padded (only last batch padded) + uint32_t V_total_seq_len_per_head, // Total sequence length per head for V in layout [NUM_HEAD, batch, L_Seq_per_batch, HEAD_DIM]. Can be L_seq_padded*batch_size (uniform padding) or L_seq*(batch_size-1)+L_seq_padded (only last batch padded) + uint32_t Q_offset, uint32_t K_offset, uint32_t V_offset, uint32_t O_offset, + + int ARG_Q, int ARG_K, int ARG_V, int ARG_O, + npu_it_channel it_for_send_Q, npu_it_channel it_for_send_KV, + npu_it_channel it_for_recv_O, + std::vector &shim_tiles +){ + + + uint32_t LQ_per_column = LQ_CU_chunk_size/(CU_column[0].size()); + + bool sanity_check = true; + + sanity_check &= SHIM_S_SEQ_PADDED%LK_CU_chunk_size ==0; + sanity_check &= MT_KV_repeat >1; + sanity_check &= LK_CU_chunk_size == LK_per_CT; + sanity_check &= LQ_CU_chunk_size % CU_column[0].size() ==0; + + sanity_check &= NUM_OF_CT_per_column*LQ_per_CT == LQ_per_column; + if(!sanity_check){ + std::cerr << "process_SHM_Q_chunk_Q_repeat sanity check failed"<< std::endl; + exit(1); + } + + + + for(int cu_idx = 0; cu_idx = NUM_OF_DH){ + break; + } + + for(int idx = 0; idx < cur_CU_shimtile_cols.size(); idx++){ + + auto shim_col_idx = cur_CU_shimtile_cols.at(idx); + + uint32_t SHIM_O_offset = SHIM_LQ_start_row * (O_PADDDED_D) + DH_index * DH +O_offset; + // now, receive the O + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_O_queue_counter, + shim_tiles.at(shim_col_idx), + S2MM, it_for_recv_O + ); + seq.npu_dma_memcpy_nd( + sizeof(T_out), + ARG_O, + S2MM, + shim_tiles.at(shim_col_idx), + static_cast(get_O_bd_id(SHIM_O_ping_pong.at(shim_col_idx))), + it_for_recv_O, + {0,0,0, (uint32_t)SHIM_O_offset + LQ_per_column*idx*(O_PADDDED_D)}, + {1,1, (uint32_t)LQ_per_column, (uint32_t)DH}, + {0,0, (uint32_t)O_PADDDED_D, 1}, + -1, 0, true + ); + update_ping_pong_flag(SHIM_O_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_O_queue_counter + ); + + + } + + } + + + + + + for(int kv_chunk_idx = 0; kv_chunk_idx < SHIM_S_SEQ_PADDED/LK_CU_chunk_size; kv_chunk_idx++){ + + + + for(int cu_idx = 0; cu_idx = NUM_OF_DH){ + break; + } + + + uint32_t Shim_Q_offset = SHIM_LQ_start_row* DH + ((Q_total_seq_len_per_head * DH) * DH_index) + Q_offset; + + uint32_t SHIM_K_offset = DH_index * (DH*K_total_seq_len_per_head) + K_offset; + uint32_t SHIM_V_offset = DH_index * (DH*V_total_seq_len_per_head) + V_offset; + + + for(int idx = 0; idx < cur_CU_shimtile_cols.size(); idx++){ + + auto shim_col_idx = cur_CU_shimtile_cols.at(idx); + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_Q_queue_counter, + shim_tiles.at(shim_col_idx), + MM2S, it_for_send_Q + ); + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_Q, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_Q_bd_id(SHIM_Q_ping_pong.at(shim_col_idx))), + it_for_send_Q, + {0,0,0, (uint32_t)Shim_Q_offset + LQ_per_column*idx*(DH)}, + // {1,1, (uint32_t)LQ_per_column, (uint32_t)DH}, + // {0,0, (uint32_t)DH, 1}, + {1, 1, 1, (uint32_t)LQ_per_column*DH}, + {0, 0, 0, 1}, + -1, 0, true + + ); + update_ping_pong_flag(SHIM_Q_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_Q_queue_counter + ); + } + + + //only the first shimtile column in cur_CU_shimtile_cols will be used to send KV data + { + + uint32_t shim_col_idx = cur_CU_shimtile_cols.at(0); + + //NOTE: NO need to wait for k,since we only issue the v + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_V_queue_counter, + shim_tiles.at(shim_col_idx), + MM2S, it_for_send_KV + ); + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_K, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_K_bd_id(SHIM_K_ping_pong.at(shim_col_idx))), + it_for_send_KV, + {0,0,0, (uint32_t)SHIM_K_offset + kv_chunk_idx*LK_CU_chunk_size*(DH)}, + //{1,1, (uint32_t)LK_CU_chunk_size, (uint32_t)DH}, + //{0,0, (uint32_t)DH, 1}, + {1, 1, 1, (uint32_t)LK_CU_chunk_size*DH}, + {0, 0, 0, 1}, + -1, 0, false, + aggressive_cache + + ); + update_ping_pong_flag(SHIM_K_ping_pong.at(shim_col_idx)); + + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_V, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_V_bd_id(SHIM_V_ping_pong.at(shim_col_idx))), + it_for_send_KV, + {0,0,0, (uint32_t)SHIM_V_offset + kv_chunk_idx*LK_CU_chunk_size*(DH)}, + // {1,1, (uint32_t)LK_CU_chunk_size, (uint32_t)DH}, + // {0,0, (uint32_t)DH, 1}, + {1, 1, 1, (uint32_t)LK_CU_chunk_size*DH}, + {0, 0, 0, 1}, + -1, 0, true, + aggressive_cache + ); + update_ping_pong_flag(SHIM_V_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_V_queue_counter + ); + + } + + + } + + + + + } + +} + +template +void process_SHM_Q_chunk_no_Q_repeat( + + npu_sequence &seq, + int NUM_OF_CT_per_column, + int MT_KV_repeat, + int DH, + int NUM_OF_DH, + int DH_index_start, // current DH index that is being processed + + int SHIM_LQ_start_row, + int SHIM_S_SEQ_PADDED, + + std::vector> &CU_column, + int LQ_CU_chunk_size, + int LK_CU_chunk_size, + + int LQ_per_CT, + int LK_per_CT, + + std::vector &SHIM_Q_queue_counter, + std::vector &SHIM_K_queue_counter, + std::vector &SHIM_V_queue_counter, + std::vector &SHIM_O_queue_counter, + + std::vector &SHIM_Q_ping_pong, + std::vector &SHIM_K_ping_pong, + std::vector &SHIM_V_ping_pong, + std::vector &SHIM_O_ping_pong, + + + uint32_t Q_PADDDED_D, uint32_t K_PADDED_D, uint32_t V_PADDED_D, uint32_t O_PADDDED_D, + uint32_t Q_offset, uint32_t K_offset, uint32_t V_offset, uint32_t O_offset, + int ARG_Q, int ARG_K, int ARG_V, int ARG_O, + npu_it_channel it_for_send_Q, npu_it_channel it_for_send_KV, + npu_it_channel it_for_recv_O, + std::vector &shim_tiles + +){ + if(LQ_CU_chunk_size % LQ_per_CT !=0){ + std::cerr << "process_SHM_Q_chunk_no_Q_repeat LQ_CU_chunk_size % LQ_per_CT !=0"<< std::endl; + exit(1); + } + + + + uint32_t LQ_per_column = LQ_CU_chunk_size/(CU_column[0].size()); + bool sanity_check = true; + sanity_check &= SHIM_S_SEQ_PADDED%LK_CU_chunk_size ==0; + sanity_check &= MT_KV_repeat == 1; + sanity_check &= LK_CU_chunk_size == LK_per_CT; + sanity_check &= NUM_OF_CT_per_column*LQ_per_CT == LQ_per_column; + if(!sanity_check){ + std::cerr << "process_SHM_Q_chunk_no_Q_repeat sanity check failed"<< std::endl; + exit(1); + } + + + + + + for(int cu_idx = 0; cu_idx < CU_column.size(); cu_idx++){ + + auto cur_CU_shimtile_cols = CU_column.at(cu_idx); + int DH_index = DH_index_start + cu_idx; + if(DH_index >= NUM_OF_DH){ + break;// done + } + + + uint32_t Shim_Q_offset = SHIM_LQ_start_row * (Q_PADDDED_D) + DH_index * DH + Q_offset; + uint32_t SHIM_O_offset = SHIM_LQ_start_row * (O_PADDDED_D) + DH_index * DH + O_offset; + + + + for(int idx = 0; idx < cur_CU_shimtile_cols.size(); idx++){ + + auto shim_col_idx = cur_CU_shimtile_cols.at(idx); + + + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_Q_queue_counter, + shim_tiles.at(shim_col_idx), + MM2S, it_for_send_Q + ); + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_Q, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_Q_bd_id(SHIM_Q_ping_pong.at(shim_col_idx))), + it_for_send_Q, + {0,0,0, (uint32_t)Shim_Q_offset + LQ_per_column*idx*(Q_PADDDED_D)}, + {1,1, (uint32_t)LQ_per_column, (uint32_t)DH}, + {0,0, (uint32_t)Q_PADDDED_D, 1}, + -1, 0, true + + ); + update_ping_pong_flag(SHIM_Q_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_Q_queue_counter + ); + + + + + // now, receive the O + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_O_queue_counter, + shim_tiles.at(shim_col_idx), + S2MM, it_for_recv_O + ); + seq.npu_dma_memcpy_nd( + sizeof(T_out), + ARG_O, + S2MM, + shim_tiles.at(shim_col_idx), + static_cast(get_O_bd_id(SHIM_O_ping_pong.at(shim_col_idx))), + it_for_recv_O, + {0,0,0, (uint32_t)SHIM_O_offset + LQ_per_column*idx*(O_PADDDED_D)}, + {1,1, (uint32_t)LQ_per_column, (uint32_t)DH}, + {0,0, (uint32_t)O_PADDDED_D, 1}, + -1, 0, true + ); + update_ping_pong_flag(SHIM_O_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_O_queue_counter + ); + + + } + + + } + + + for(uint32_t kv_chunk_idx = 0; kv_chunk_idx < SHIM_S_SEQ_PADDED/LK_CU_chunk_size; kv_chunk_idx++){ + + + for(int cu_idx = 0; cu_idx < CU_column.size(); cu_idx++){ + + auto cur_CU_shimtile_cols = CU_column.at(cu_idx); + int DH_index = DH_index_start + cu_idx; + if(DH_index >= NUM_OF_DH){ + break;// done + } + + uint32_t SHIM_K_offset = DH_index * DH + K_offset; + uint32_t SHIM_V_offset = DH_index * DH + V_offset; + // only the first shimtile column in cur_CU_shimtile_cols will be used to send KV data + uint32_t shim_col_idx = cur_CU_shimtile_cols.at(0); + + + //NOTE: NO need to wait for k,since we only issue the v + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_V_queue_counter, + shim_tiles.at(shim_col_idx), + MM2S, it_for_send_KV + ); + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_K, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_K_bd_id(SHIM_K_ping_pong.at(shim_col_idx))), + it_for_send_KV, + {0,0,0, (uint32_t)SHIM_K_offset + kv_chunk_idx*LK_CU_chunk_size*(K_PADDED_D)}, + {1,1, (uint32_t)LK_CU_chunk_size, (uint32_t)DH}, + {0,0, (uint32_t)K_PADDED_D, 1}, + -1, 0, false, + aggressive_cache + ); + update_ping_pong_flag(SHIM_K_ping_pong.at(shim_col_idx)); + + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_V, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_V_bd_id(SHIM_V_ping_pong.at(shim_col_idx))), + it_for_send_KV, + {0,0,0, (uint32_t)SHIM_V_offset + kv_chunk_idx*LK_CU_chunk_size*(V_PADDED_D)}, + {1,1, (uint32_t)LK_CU_chunk_size, (uint32_t)DH}, + {0,0, (uint32_t)V_PADDED_D, 1}, + -1, 0, true, + aggressive_cache + ); + update_ping_pong_flag(SHIM_V_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_V_queue_counter + ); + + } + + + } + +} + + + + + + + + + + + + + + + + +template +void process_SHM_Q_chunk_no_Q_repeat_qkv_reorder_NUM_DH_L_DH( + + npu_sequence &seq, + int NUM_OF_CT_per_column, + int MT_KV_repeat, + int DH, + int NUM_OF_DH, + int head_idx_start, // current DH index that is being processed + + int SHIM_LQ_start_row, + int SHIM_S_SEQ_PADDED, + + std::vector> &CU_column, + int LQ_CU_chunk_size, + int LK_CU_chunk_size, + + int LQ_per_CT, + int LK_per_CT, + + std::vector &SHIM_Q_queue_counter, + std::vector &SHIM_K_queue_counter, + std::vector &SHIM_V_queue_counter, + std::vector &SHIM_O_queue_counter, + + std::vector &SHIM_Q_ping_pong, + std::vector &SHIM_K_ping_pong, + std::vector &SHIM_V_ping_pong, + std::vector &SHIM_O_ping_pong, + + + uint32_t O_PADDDED_D, + //[QWEN3_VISION_NUM_HEAD, batch, L_Seq_per_batch, QWEN3_VISION_HEAD_DIM] + uint32_t Q_total_seq_len_per_head, // Total sequence length per head for Q in layout[NUM_HEAD, batch, L_Seq_per_batch, HEAD_DIM]. Can be L_seq_padded*batch_size (uniform padding) or L_seq*(batch_size-1)+L_seq_padded (only last batch padded) + uint32_t K_total_seq_len_per_head, // Total sequence length per head for K in layout [NUM_HEAD, batch, L_Seq_per_batch, HEAD_DIM]. Can be L_seq_padded*batch_size (uniform padding) or L_seq*(batch_size-1)+L_seq_padded (only last batch padded) + uint32_t V_total_seq_len_per_head, // Total sequence length per head for V in layout [NUM_HEAD, batch, L_Seq_per_batch, HEAD_DIM]. Can be L_seq_padded*batch_size (uniform padding) or L_seq*(batch_size-1)+L_seq_padded (only last batch padded) + + uint32_t Q_offset, uint32_t K_offset, uint32_t V_offset, uint32_t O_offset, + + int ARG_Q, int ARG_K, int ARG_V, int ARG_O, + npu_it_channel it_for_send_Q, npu_it_channel it_for_send_KV, + npu_it_channel it_for_recv_O, + std::vector &shim_tiles + +){ + if(LQ_CU_chunk_size % LQ_per_CT !=0){ + std::cerr << "process_SHM_Q_chunk_no_Q_repeat LQ_CU_chunk_size % LQ_per_CT !=0"<< std::endl; + exit(1); + } + uint32_t LQ_per_column = LQ_CU_chunk_size/(CU_column[0].size()); + bool sanity_check = true; + sanity_check &= SHIM_S_SEQ_PADDED%LK_CU_chunk_size ==0; + sanity_check &= MT_KV_repeat == 1; + sanity_check &= LK_CU_chunk_size == LK_per_CT; + sanity_check &= NUM_OF_CT_per_column*LQ_per_CT == LQ_per_column; + if(!sanity_check){ + std::cerr << "process_SHM_Q_chunk_no_Q_repeat sanity check failed"<< std::endl; + exit(1); + } + + + + + + for(int cu_idx = 0; cu_idx < CU_column.size(); cu_idx++){ + + auto cur_CU_shimtile_cols = CU_column.at(cu_idx); + int DH_index = head_idx_start + cu_idx; + if(DH_index >= NUM_OF_DH){ + break;// done + } + + uint32_t Shim_Q_offset = SHIM_LQ_start_row* DH + ((Q_total_seq_len_per_head * DH) * DH_index) + Q_offset; + uint32_t SHIM_O_offset = SHIM_LQ_start_row * (O_PADDDED_D) + DH_index * DH + O_offset; + + + for(int idx = 0; idx < cur_CU_shimtile_cols.size(); idx++){ + + auto shim_col_idx = cur_CU_shimtile_cols.at(idx); + + + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_Q_queue_counter, + shim_tiles.at(shim_col_idx), + MM2S, it_for_send_Q + ); + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_Q, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_Q_bd_id(SHIM_Q_ping_pong.at(shim_col_idx))), + it_for_send_Q, + {0,0,0, (uint32_t) Shim_Q_offset + idx*(LQ_per_column*DH)}, + // {1,1, (uint32_t)LQ_per_column, (uint32_t)DH}, + // {0,0, (uint32_t)DH, 1}, + {1, 1, 1, (uint32_t)LQ_per_column*DH}, + {0, 0, 0, 1}, + -1, 0, true + + ); + update_ping_pong_flag(SHIM_Q_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_Q_queue_counter + ); + + + + + // now, receive the O + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_O_queue_counter, + shim_tiles.at(shim_col_idx), + S2MM, it_for_recv_O + ); + seq.npu_dma_memcpy_nd( + sizeof(T_out), + ARG_O, + S2MM, + shim_tiles.at(shim_col_idx), + static_cast(get_O_bd_id(SHIM_O_ping_pong.at(shim_col_idx))), + it_for_recv_O, + {0,0,0, (uint32_t)SHIM_O_offset + LQ_per_column*idx*(O_PADDDED_D)}, + {1,1, (uint32_t)LQ_per_column, (uint32_t)DH}, + {0,0, (uint32_t)O_PADDDED_D, 1}, + -1, 0, true + ); + update_ping_pong_flag(SHIM_O_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_O_queue_counter + ); + + + } + + + } + + + + for(uint32_t kv_chunk_idx = 0; kv_chunk_idx < SHIM_S_SEQ_PADDED/LK_CU_chunk_size; kv_chunk_idx++){ + + + for(int cu_idx = 0; cu_idx < CU_column.size(); cu_idx++){ + + auto cur_CU_shimtile_cols = CU_column.at(cu_idx); + int DH_index = head_idx_start + cu_idx; + if(DH_index >= NUM_OF_DH){ + break;// done + } + + uint32_t SHIM_K_offset = DH_index * (DH*K_total_seq_len_per_head) + K_offset; + uint32_t SHIM_V_offset = DH_index * (DH*V_total_seq_len_per_head) + V_offset; + + + // only the first shimtile column in cur_CU_shimtile_cols will be used to send KV data + uint32_t shim_col_idx = cur_CU_shimtile_cols.at(0); + + //NOTE: NO need to wait for k,since we only issue the v + wait_DMA_queue_if_full( + seq, shim_col_idx, SHIM_V_queue_counter, + shim_tiles.at(shim_col_idx), + MM2S, it_for_send_KV + ); + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_K, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_K_bd_id(SHIM_K_ping_pong.at(shim_col_idx))), + it_for_send_KV, + {0,0,0, (uint32_t)SHIM_K_offset + kv_chunk_idx*LK_CU_chunk_size*(DH )}, + //{1,1, (uint32_t)LK_CU_chunk_size, (uint32_t)DH}, + //{0,0, (uint32_t)DH, 1}, + {1,1, 1,(uint32_t)LK_CU_chunk_size *DH}, + {0,0, 0, 1}, + -1, 0, false, + aggressive_cache + ); + update_ping_pong_flag(SHIM_K_ping_pong.at(shim_col_idx)); + + + seq.npu_dma_memcpy_nd( + sizeof(T_in), + ARG_V, + MM2S, + shim_tiles.at(shim_col_idx), + static_cast(get_V_bd_id(SHIM_V_ping_pong.at(shim_col_idx))), + it_for_send_KV, + {0,0,0, (uint32_t)SHIM_V_offset + kv_chunk_idx*LK_CU_chunk_size*(DH)}, + // {1,1, (uint32_t)LK_CU_chunk_size, (uint32_t)DH}, + // {0,0, (uint32_t)DH, 1}, + {1,1, 1,(uint32_t)LK_CU_chunk_size *DH}, + {0,0, 0, 1}, + -1, 0, true, + aggressive_cache + ); + update_ping_pong_flag(SHIM_V_ping_pong.at(shim_col_idx)); + increment_DMA_queue( + seq, shim_col_idx, SHIM_V_queue_counter + ); + } + + + + } + +} + + + + + + + + + + +struct AttentionConfig { + uint32_t DH; + uint32_t NUM_OF_DH; + uint32_t LQ_per_CT; + uint32_t LK_per_CT; + uint32_t lq_internal; + uint32_t lk_lv_internal; + uint32_t NUM_OF_COLUMNS; + uint32_t NUM_OF_CT_PER_COLUMN; + uint32_t CU_mode; + uint32_t num_of_batches; + + int ARG_Q, ARG_K, ARG_V, ARG_O; + + + // Q_PADDED_D, K_PADDED_D, V_PADDED_D is singificant only if QKV_reordered == false + // since QKV_reordered= false mean + // Q, K, V buffer is view as [batch, L_seq, NUM_DH*DH] + uint32_t Q_padded_D; + uint32_t K_padded_D; + uint32_t V_padded_D; + uint32_t O_PADDDED_D; + + bool QKV_reordered; + // parameter below is only siginificatn when QKV_reordered = true + // when QKV_reordered == true, it means + // Q, K, V buffer is viewed as [NUM_DH, batch, L_seq, DH] + uint32_t Q_total_seq_len_per_head; + uint32_t K_total_seq_len_per_head; + uint32_t V_total_seq_len_per_head; +}; + +struct BatchMetadata { + const std::vector& SHIM_L_SEQ_list; + const std::vector& SHIM_S_SEQ_list; + const std::vector& SHIM_S_SEQ_PADDED_list; + + const std::vector& Q_batch_offset_list; + const std::vector& K_batch_offset_list; + const std::vector& V_batch_offset_list; + const std::vector& O_batch_offset_list; +}; + +template +void setup_SHM_configuration( + npu_sequence &seq, + const AttentionConfig& config, + const BatchMetadata& batch_data, + float attention_scalar +){ + uint32_t DH = config.DH; + uint32_t NUM_OF_DH = config.NUM_OF_DH; + uint32_t LQ_per_CT = config.LQ_per_CT; + uint32_t LK_per_CT = config.LK_per_CT; + uint32_t lq_internal = config.lq_internal; + uint32_t lk_lv_internal = config.lk_lv_internal; + uint32_t NUM_OF_COLUMNS = config.NUM_OF_COLUMNS; + uint32_t NUM_OF_CT_PER_COLUMN = config.NUM_OF_CT_PER_COLUMN; + uint32_t CU_mode = config.CU_mode; + uint32_t num_of_batches = config.num_of_batches; + int ARG_Q = config.ARG_Q; + int ARG_K = config.ARG_K; + int ARG_V = config.ARG_V; + int ARG_O = config.ARG_O; + uint32_t Q_padded_D = config.Q_padded_D; + uint32_t K_padded_D = config.K_padded_D; + uint32_t V_padded_D = config.V_padded_D; + uint32_t O_PADDDED_D = config.O_PADDDED_D; + bool QKV_reordered = config.QKV_reordered; + uint32_t Q_total_seq_len_per_head = config.Q_total_seq_len_per_head; + uint32_t K_total_seq_len_per_head = config.K_total_seq_len_per_head; + uint32_t V_total_seq_len_per_head = config.V_total_seq_len_per_head; + + const std::vector& SHIM_L_SEQ_list = batch_data.SHIM_L_SEQ_list; + const std::vector& SHIM_S_SEQ_list = batch_data.SHIM_S_SEQ_list; + const std::vector& SHIM_S_SEQ_PADDED_list = batch_data.SHIM_S_SEQ_PADDED_list; + const std::vector& Q_batch_offset_list = batch_data.Q_batch_offset_list; + const std::vector& K_batch_offset_list = batch_data.K_batch_offset_list; + const std::vector& V_batch_offset_list = batch_data.V_batch_offset_list; + const std::vector& O_batch_offset_list = batch_data.O_batch_offset_list; + + if(Q_padded_D < DH*NUM_OF_DH || K_padded_D < DH*NUM_OF_DH || V_padded_D < DH*NUM_OF_DH || O_PADDDED_D < DH*NUM_OF_DH) { + std::cerr << "padded D less than required D"<< std::endl; + exit(1); + } + + if(SHIM_L_SEQ_list.size() != num_of_batches || + SHIM_S_SEQ_list.size() != num_of_batches || + SHIM_S_SEQ_PADDED_list.size() != num_of_batches){ + std::cerr << "SHIM_L/S_SEQ_LIST size not equal to num_of_batches"<< std::endl; + exit(1); + } + + + assert(LQ_per_CT % lq_internal ==0); + int MT_KV_repeat = LQ_per_CT/lq_internal; + + + std::vector shim_tiles; + // initialize the shimtiles + for(int i = 0; i < NUM_OF_COLUMNS; i++){ + shim_tiles.push_back(get_tile(0, i)); + } + + constexpr int CT_lock_address_base = 0x000001F000; + + constexpr npu_it_channel it_for_send_Q = it_channel_0; + constexpr npu_it_channel it_for_send_KV = it_channel_1; + constexpr npu_it_channel it_for_recv_O = it_channel_0; + + constexpr int CT_rtp_lock_id = 10; + constexpr int CT_rtp_address = 4096; + + // parameter sanity checks + if(NUM_OF_COLUMNS != 8){ + std::cerr << "Currently only support 8 columns" << std::endl; + exit(1); + } + if(NUM_OF_CT_PER_COLUMN !=4){ + std::cerr << "Currently only support 4 CT per column" << std::endl; + exit(1); + } + // different CU_mode result in different shim->kv broadcast pattern + + std::vector> CU_column; + if (CU_mode == 0){ + // 1-1 + CU_column = { {0}, {1}, {2}, {3}, {4}, {5}, {6}, {7} }; // each column + }else if(CU_mode == 1){ + // 1-8 + CU_column = { {0,1,2,3,4,5,6,7} }; // all columns + }else if(CU_mode == 2){ + // 1-4 + CU_column = { {0,1,2,3}, {4,5,6,7} }; // each column + }else if(CU_mode == 3){ + // 1-2 + CU_column = { {0,1}, {2,3}, {4,5}, {6,7} }; // each column + }else{ + std::cerr << "CU_mode not supported" << std::endl; + exit(1); + } + + + seq.clear_cmds(); + seq.npu_preemption(0); + + int LQ_per_CU = CU_column[0].size() *(LQ_per_CT * NUM_OF_CT_PER_COLUMN); + + + for(int batch_idx =0; batch_idx < num_of_batches; batch_idx++){ + if(SHIM_L_SEQ_list[batch_idx] % LQ_per_CU !=0){ + std::cerr << "SHIM_L_SEQ must be multiple of LQ_per_CU" << std::endl; + exit(1); + } + } + + std::vector SHIM_Q_queue_counter(NUM_OF_COLUMNS, 0); // counter for each column + std::vector SHIM_K_queue_counter(NUM_OF_COLUMNS, 0); // counter for each column + std::vector SHIM_V_queue_counter(NUM_OF_COLUMNS, 0); // counter for each column + std::vector SHIM_O_queue_counter(NUM_OF_COLUMNS, 0); // counter for each column + + + std::vector SHIM_Q_ping_pong(NUM_OF_COLUMNS, 0); // ping pong flag for each column + std::vector SHIM_K_ping_pong(NUM_OF_COLUMNS, 0); // ping pong flag for each column + std::vector SHIM_V_ping_pong(NUM_OF_COLUMNS, 0); // ping pong flag for each column + std::vector SHIM_O_ping_pong(NUM_OF_COLUMNS, 0); // ping pong flag for each column + + + + + + + for(uint32_t batch_idx = 0; batch_idx < num_of_batches; batch_idx++){ + + + + // # calculate how many SHIM_L chunks we have in the column + std::vector LQ_chunk_per_CT_in_column(NUM_OF_COLUMNS, 0);// init to 0 + + int NUM_OF_CU_UNITS = CU_column.size(); + + for (int head_idx_start = 0; head_idx_start < NUM_OF_DH; head_idx_start+= NUM_OF_CU_UNITS) { + for (int SHIM_L_chunk_idx = 0; SHIM_L_chunk_idx < SHIM_L_SEQ_list[batch_idx] / LQ_per_CU; SHIM_L_chunk_idx++) { + for(int cu_idx = 0; cu_idx < CU_column.size(); cu_idx++){ + int DH_index = head_idx_start + cu_idx; + if(DH_index >= NUM_OF_DH){ + break; + } + for(auto col_idx : CU_column[cu_idx]){ + LQ_chunk_per_CT_in_column.at(col_idx)++; + } + } + } + } + + + + //NOTEL for now, we only give SHIM_S_SEQ + for(auto cur_CU_shimtile_cols : CU_column){ + + for (auto shim_col_idx: cur_CU_shimtile_cols){ + for(int row_idx= 0; row_idx < NUM_OF_CT_PER_COLUMN; row_idx++){ + + auto CT_tile = get_tile(row_idx+2, shim_col_idx); + seq.rtp_write(CT_tile, CT_rtp_address, SHIM_S_SEQ_list[batch_idx]); // SHIM_S_SEQ + seq.rtp_write(CT_tile, CT_rtp_address+4, LQ_chunk_per_CT_in_column[shim_col_idx] * LQ_per_CT); // how many LQ chunks per CT in this column + + uint32_t attention_scalar_bits; + memcpy(&attention_scalar_bits, &attention_scalar, sizeof(uint32_t)); + seq.rtp_write(CT_tile, CT_rtp_address+8, attention_scalar_bits); + + seq.rtp_write(CT_tile, CT_lock_address_base+16*(CT_rtp_lock_id), 1); // set lock to 1 + } + + + } + } + + + + // uint32_t Q_offset = batch_idx * SHIM_L_SEQ_list[batch_idx]* Q_padded_D + Q_external_data_offset; + // uint32_t K_offset = batch_idx * SHIM_S_SEQ_PADDED_list[batch_idx] * K_padded_D+ K_external_data_offset; + // uint32_t V_offset = batch_idx * SHIM_S_SEQ_PADDED_list[batch_idx] * V_padded_D + V_external_data_offset; + // uint32_t O_offset = batch_idx * SHIM_L_SEQ_list[batch_idx] * O_PADDDED_D + O_external_data_offset; + + uint32_t Q_offset = Q_batch_offset_list[batch_idx];; + uint32_t K_offset = K_batch_offset_list[batch_idx]; + uint32_t V_offset = V_batch_offset_list[batch_idx]; + uint32_t O_offset = O_batch_offset_list[batch_idx]; + + + + for (int head_idx_start = 0; head_idx_start < NUM_OF_DH; head_idx_start+= NUM_OF_CU_UNITS) + { + + for (int SHIM_L_chunk_idx = 0; SHIM_L_chunk_idx < SHIM_L_SEQ_list[batch_idx] / LQ_per_CU; SHIM_L_chunk_idx++) + { + + if (MT_KV_repeat == 1) + { + if (!QKV_reordered) + { + process_SHM_Q_chunk_no_Q_repeat( + seq, + NUM_OF_CT_PER_COLUMN, + MT_KV_repeat, + DH, + NUM_OF_DH, + head_idx_start, + SHIM_L_chunk_idx * LQ_per_CU, + SHIM_S_SEQ_PADDED_list[batch_idx], + CU_column, + LQ_per_CU, + LK_per_CT, // same as LK_CU_chunk_size + + LQ_per_CT, + LK_per_CT, + SHIM_Q_queue_counter, + SHIM_K_queue_counter, + SHIM_V_queue_counter, + SHIM_O_queue_counter, + SHIM_Q_ping_pong, + SHIM_K_ping_pong, + SHIM_V_ping_pong, + SHIM_O_ping_pong, + Q_padded_D, K_padded_D, V_padded_D, O_PADDDED_D, + Q_offset, K_offset, V_offset, O_offset, + ARG_Q, ARG_K, ARG_V, ARG_O, + it_for_send_Q, it_for_send_KV, + it_for_recv_O, + shim_tiles); + } + else + { + process_SHM_Q_chunk_no_Q_repeat_qkv_reorder_NUM_DH_L_DH( + seq, + NUM_OF_CT_PER_COLUMN, + MT_KV_repeat, + DH, + NUM_OF_DH, + head_idx_start, + SHIM_L_chunk_idx * LQ_per_CU, + SHIM_S_SEQ_PADDED_list[batch_idx], + CU_column, + LQ_per_CU, + LK_per_CT, // same as LK_CU_chunk_size + + LQ_per_CT, + LK_per_CT, + SHIM_Q_queue_counter, + SHIM_K_queue_counter, + SHIM_V_queue_counter, + SHIM_O_queue_counter, + SHIM_Q_ping_pong, + SHIM_K_ping_pong, + SHIM_V_ping_pong, + SHIM_O_ping_pong, + O_PADDDED_D, + Q_total_seq_len_per_head, K_total_seq_len_per_head, V_total_seq_len_per_head, + Q_offset, K_offset, V_offset, O_offset, + ARG_Q, ARG_K, ARG_V, ARG_O, + it_for_send_Q, it_for_send_KV, + it_for_recv_O, + shim_tiles); + } + } + else + { + if (!QKV_reordered) + { + process_SHM_Q_chunk_Q_repeat( + seq, + NUM_OF_CT_PER_COLUMN, + MT_KV_repeat, + DH, + NUM_OF_DH, + head_idx_start, + SHIM_L_chunk_idx * LQ_per_CU, + SHIM_S_SEQ_PADDED_list[batch_idx], + CU_column, + LQ_per_CU, + LK_per_CT, // same as LK_CU_chunk_size + + LQ_per_CT, + LK_per_CT, + SHIM_Q_queue_counter, + SHIM_K_queue_counter, + SHIM_V_queue_counter, + SHIM_O_queue_counter, + SHIM_Q_ping_pong, + SHIM_K_ping_pong, + SHIM_V_ping_pong, + SHIM_O_ping_pong, + Q_padded_D, K_padded_D, V_padded_D, O_PADDDED_D, + Q_offset, K_offset, V_offset, O_offset, + ARG_Q, ARG_K, ARG_V, ARG_O, + it_for_send_Q, it_for_send_KV, + it_for_recv_O, + shim_tiles); + } + else + { + process_SHM_Q_chunk_Q_repeat_reorder_NUM_DH_L_DH( + seq, + NUM_OF_CT_PER_COLUMN, + MT_KV_repeat, + DH, + NUM_OF_DH, + head_idx_start, + SHIM_L_chunk_idx * LQ_per_CU, + SHIM_S_SEQ_PADDED_list[batch_idx], + CU_column, + LQ_per_CU, + LK_per_CT, // same as LK_CU_chunk_size + + LQ_per_CT, + LK_per_CT, + SHIM_Q_queue_counter, + SHIM_K_queue_counter, + SHIM_V_queue_counter, + SHIM_O_queue_counter, + SHIM_Q_ping_pong, + SHIM_K_ping_pong, + SHIM_V_ping_pong, + SHIM_O_ping_pong, + O_PADDDED_D, + Q_total_seq_len_per_head, K_total_seq_len_per_head, V_total_seq_len_per_head, + Q_offset, K_offset, V_offset, O_offset, + ARG_Q, ARG_K, ARG_V, ARG_O, + it_for_send_Q, it_for_send_KV, + it_for_recv_O, + shim_tiles + + ); + } + } + + + } + } + + + // clear and wait for all remaining + bool is_all_flushed = false; + + while(!is_all_flushed){ + is_all_flushed = true; // init to True + + // start with q + for(int idx = 0; idx < NUM_OF_COLUMNS; idx++){ + if(SHIM_Q_queue_counter.at(idx) > 0){ + is_all_flushed = false; + force_wait_DMA_queue( + seq, idx, SHIM_Q_queue_counter, + shim_tiles.at(idx), + MM2S, it_for_send_Q + ); + } + } + + // THEN K, V + for(int idx = 0; idx < NUM_OF_COLUMNS; idx++){ + // assert SHIM_K_queue_counter[idx] == 0 + assert(SHIM_K_queue_counter.at(idx) == 0); + + if(SHIM_V_queue_counter.at(idx) > 0){ + is_all_flushed = false; + force_wait_DMA_queue( + seq, idx, SHIM_V_queue_counter, + shim_tiles.at(idx), + MM2S, it_for_send_KV + ); + } + } + } + + is_all_flushed = false; + while(!is_all_flushed){ + is_all_flushed = true; // init to True + + for(int idx = 0; idx < NUM_OF_COLUMNS; idx++){ + assert(SHIM_O_queue_counter.at(idx) <= 2); + if(SHIM_O_queue_counter.at(idx) > 0){ + is_all_flushed = false; + force_wait_DMA_queue( + seq, idx, SHIM_O_queue_counter, + shim_tiles.at(idx), + S2MM, it_for_recv_O + ); + } + } + } + + + } + + + + + + + + + + + + seq.cmds2seq(); + +} + + + + +#endif \ No newline at end of file diff --git a/src/include/npu_utils/flm_runtime.hpp b/src/include/npu_utils/flm_runtime.hpp new file mode 100644 index 000000000..92355654e --- /dev/null +++ b/src/include/npu_utils/flm_runtime.hpp @@ -0,0 +1,17 @@ +/// \file flm_runtime.hpp +/// \brief Neutral NPU-runtime namespace alias (flm_rt) selecting the active +/// backend at build time. FLM_USE_HRX=ON maps flm_rt to the hrx C++ +/// shim (over libhrx); otherwise flm_rt maps to XRT. Engine sources use +/// flm_rt:: so a single tree compiles against either runtime. +#pragma once + +#if defined(FLM_USE_HRX) +#include "hrx_cpp/hrx_cpp.hpp" +namespace flm_rt = hrx; +#else +#include "xrt/xrt_device.h" +#include "xrt/xrt_kernel.h" +#include "xrt/xrt_bo.h" +#include "xrt/experimental/xrt_kernel.h" +namespace flm_rt = xrt; +#endif diff --git a/src/include/npu_utils/instr_utils/npu_cmd.hpp b/src/include/npu_utils/instr_utils/npu_cmd.hpp index 66a8f9d25..d1252396e 100644 --- a/src/include/npu_utils/instr_utils/npu_cmd.hpp +++ b/src/include/npu_utils/instr_utils/npu_cmd.hpp @@ -1,3 +1,4 @@ +#if defined(FLM_USE_HRX) /// \file npu_cmd.hpp /// \brief npu command /// \author FastFlowLM Team, Alfred @@ -105,4 +106,115 @@ struct npu_cmd{ +#endif +#else +/// \file npu_cmd.hpp +/// \brief npu command +/// \author FastFlowLM Team, Alfred +/// \date 2025-09-09 +/// \note This is a class for the npu command, it is a virtual class for all npu commands +#ifndef __NPU_CMD_HPP__ +#define __NPU_CMD_HPP__ + +#include +#include +#include +#include +#include +#include +#include +#include +#include "buffer.hpp" +#include "utils/debug_utils.hpp" +#include "xrt/xrt_bo.h" + +const int INSTR_PRINT_WIDTH = 80; + +typedef enum : uint32_t { + XAIE_IO_WRITE, + XAIE_IO_BLOCKWRITE, + XAIE_IO_BLOCKSET, + XAIE_IO_MASKWRITE, + XAIE_IO_MASKPOLL, + XAIE_IO_NOOP, + XAIE_IO_PREEMPT, + XAIE_IO_MASKPOLL_BUSY, + XAIE_IO_LOADPDI, + XAIE_IO_LOAD_PM_START, + XAIE_IO_CREATE_SCRATCHPAD, + XAIE_IO_UPDATE_STATE_TABLE, + XAIE_IO_UPDATE_REG, + XAIE_IO_UPDATE_SCRATCH, + XAIE_CONFIG_SHIMDMA_BD, + XAIE_CONFIG_SHIMDMA_DMABUF_BD, + XAIE_IO_CUSTOM_OP_BEGIN = 1U<<7U, // 0x80 + XAIE_IO_CUSTOM_OP_TCT = XAIE_IO_CUSTOM_OP_BEGIN, // still 0x80 + XAIE_IO_CUSTOM_OP_DDR_PATCH, // Previously this was XAIE_IO_CUSTOM_OP_BEGIN + 1, 0x81 + XAIE_IO_CUSTOM_OP_READ_REGS, // Previously this was XAIE_IO_CUSTOM_OP_BEGIN + 2 + XAIE_IO_CUSTOM_OP_RECORD_TIMER, // Previously this was XAIE_IO_CUSTOM_OP_BEGIN + 3 + XAIE_IO_CUSTOM_OP_MERGE_SYNC, // Previously this was XAIE_IO_CUSTOM_OP_BEGIN + 4 + XAIE_IO_CUSTOM_OP_NEXT, + XAIE_IO_LOAD_PM_END_INTERNAL = 200, + XAIE_IO_CUSTOM_OP_MAX = UCHAR_MAX, +} op_headers; + + +typedef enum{ + e_npu_cmd_ddr, + e_npu_cmd_issue_token, + e_npu_cmd_wait, + e_npu_cmd_write_dma, + e_npu_cmd_write +} npu_cmd_type; + +typedef enum { + S2MM, + MM2S +} dma_direction; + +typedef enum :uint32_t { + no_cache = 0x00, + normal_cache = 0x02, + aggressive_cache = 0x0e +} cache_flag_t; + +inline void instr_print(int line_number, uint32_t word, std::string msg){ + if (line_number == -1){ // -1 for the case when one line has multiple messages + MSG_BOX_LINE(INSTR_PRINT_WIDTH, std::dec << std::setw(7) << " | " << std::setw(11) << " | " << msg); + } + else{ + MSG_BOX_LINE(INSTR_PRINT_WIDTH, std::dec << std::setw(4) << line_number << " | " << std::hex << std::setfill('0') << std::setw(8) << word << " | " << msg); + } +} + +///@brief npu command +///@note This is an interface (pure virtual class) for all npu commands +///@note The class is used to print the command, convert the command to the npu format and dump the command to the buffer +///@warning The class is not used directly, but is used as a base class for all npu commands +struct npu_cmd{ + ///@brief print the command + ///@param bd the buffer to dump the command + ///@param line_number the line number of the command + ///@param op_count the operation count of the command + ///@return the number of lines of the command + virtual int print_cmd(uint32_t *bd, int line_number, int op_count) = 0; + + ///@brief convert the command to the npu format + ///@param npu_seq the npu sequence to dump the command + ///@note The function will convert the command to the npu format and dump the command to the buffer + virtual void to_npu(std::vector& npu_seq) = 0; + + // virtual npu_cmd_type get_type() = 0; + ///@brief dump the command to the buffer + ///@param bd the buffer to dump the command + virtual void dump_cmd(uint32_t *bd) = 0; + + ///@brief get the number of lines of the command + ///@return the number of lines of the command + virtual int get_op_lines() = 0; +}; + + + +#endif #endif diff --git a/src/include/npu_utils/npu_instr_utils.hpp b/src/include/npu_utils/npu_instr_utils.hpp index 7c5a9111a..f0f805cdc 100644 --- a/src/include/npu_utils/npu_instr_utils.hpp +++ b/src/include/npu_utils/npu_instr_utils.hpp @@ -143,6 +143,18 @@ class npu_sequence{ ///@note The function will read the npu sequence from the file and parse it ///@note If the file is not found, the function will throw an error ///@warning If the from_file is false, the function will not check if the filename is valid, and the npu sequence is empty + ///@brief Take the npu sequence from words already in memory. + ///@param words the sequence, one uint32_t per instruction word + ///@note For a sequence a generator produces per dispatch, where + /// writing it to a file first would be the only other option. + void from_vector(const std::vector& words){ + this->npu_seq = words; + // Without these dump() would rebuild npu_seq from the (empty) cmds + // list, and the app would keep the kernel it already built. + this->is_valid = true; + this->instr_version++; + } + void from_file(std::string filename, bool is_binary = true){ std::ifstream instr_file(filename, std::ios::binary); if (!instr_file.is_open()){ diff --git a/src/include/tensor_2d.hpp b/src/include/tensor_2d.hpp index 7edae1395..c8195e445 100644 --- a/src/include/tensor_2d.hpp +++ b/src/include/tensor_2d.hpp @@ -18,6 +18,8 @@ class tensor_2d{ uint32_t offset; std::vector> temp; public: + tensor_2d() : D(0), offset(0) {buf = nullptr;} + /// \brief constructor /// \param D the dimension of the tensor /// \param offset the offset of the tensor @@ -32,6 +34,14 @@ class tensor_2d{ /// \brief assign the buffer to the tensor_2d /// \param buf the buffer to assign + /// \brief assign the buffer, setting its row width and offset + /// \param buf the buffer to assign + void assign(buffer &buf, uint32_t D, uint32_t offset = 0){ + this->D = D; + this->offset = offset; + assign(buf); + } + void assign(buffer &buf){ this->buf = &buf; temp.resize(buf.size() / D); diff --git a/src/include/tensor_utils/safe_tensors.hpp b/src/include/tensor_utils/safe_tensors.hpp index 128086976..ecf104634 100644 --- a/src/include/tensor_utils/safe_tensors.hpp +++ b/src/include/tensor_utils/safe_tensors.hpp @@ -1,74 +1,88 @@ -/// \file safe_tensors.hpp -/// \brief SafeTensors class -/// \author FastFlowLM Team -/// \date 2025-06-24 -/// \version 0.9.24 -/// \note This class is used to load weights from a safe-tensors file. -#pragma once - -#include "typedef.hpp" -#include "buffer.hpp" -#include "nlohmann/json.hpp" -#include -#include - -/// \brief Tensor metadata -typedef struct { - std::string name; - std::vector shape; - std::string dtype; - std::vector offsets; - size_t size; - size_t byte_size; -} tensor_metadata; - -/// \brief SafeTensors class -class SafeTensors{ -private: - std::string model_path; - std::ifstream file; - nlohmann::json metadata; - size_t _get_data_size(tensor_metadata tensor_meta); - void _load_tensors(); - void _open_file(); - -protected: - std::vector tensors_data; - -public: - /// \brief Constructor - /// \param model_path the model path - SafeTensors(){ - this->model_path = ""; - } - - /// \brief Constructor - /// \param model_path the model path - SafeTensors(const std::string& model_path); - - /// \brief Destructor - ~SafeTensors(); - - /// \brief Load the weights - /// \param weight_buffer the weight buffer - /// \param weights_name the weights name - /// \return the weights name - std::string load_weights(bytes& weight_buffer, std::string weights_name); - - /// \brief Get the tensor metadata - /// \param tensor_name the tensor name - /// \return the tensor metadata - tensor_metadata get_tensor_metadata(std::string tensor_name); - - /// \brief Write the safetensors - /// \param output_path the output path - void write_safetensors(std::string output_path); - - /// \brief Get the metadata - /// \return the metadata - nlohmann::json get_metadata(); - - /// \brief Switch the model - /// \param new_model_path the new model path - void switch_model(std::string new_model_path); -}; +/// \file safe_tensors.hpp +/// \brief SafeTensors class +/// \author FastFlowLM Team +/// \date 2025-06-24 +/// \version 0.9.10 +/// \note This class is used to load weights from a safe-tensors file. +#pragma once + +#include "typedef.hpp" +#include "buffer.hpp" +#include "nlohmann/json.hpp" +#include +#include + +/// \brief Tensor metadata +typedef struct { + std::string name; + std::vector shape; + std::string dtype; + std::vector offsets; + size_t size; + size_t byte_size; +} tensor_metadata; + +/// \brief SafeTensors class +class SafeTensors{ +private: + std::string model_path; + std::ifstream file; + nlohmann::json metadata; + size_t _get_data_size(tensor_metadata tensor_meta); + void _load_tensors(); + void _open_file(); + +protected: + std::vector tensors_data; + +public: + /// \brief Constructor + /// \param model_path the model path + SafeTensors(){ + this->model_path = ""; + } + + /// \brief Constructor + /// \param model_path the model path + SafeTensors(const std::string& model_path); + + /// \brief Destructor + ~SafeTensors(); + + /// \brief Load the weights + /// \param weight_buffer the weight buffer + /// \param weights_name the weights name + /// \return the weights name + std::string load_weights(bytes& weight_buffer, std::string weights_name); + + /// \brief Load the weights with offset + /// \param weight_buffer the weight buffer + /// \param weights_name the weights name + /// \param byte_offset the byte offset in the weight buffer + /// \return the weights name + std::string load_weights(bytes& weight_buffer, std::string weights_name, size_t byte_offset); + + /// \brief Check whether a tensor is present in the file + /// \param tensor_name the tensor name + /// \return true if the tensor exists + /// \note get_tensor_metadata exits on a missing tensor, so use this to probe + /// for optional weights. + bool has_tensor(std::string tensor_name); + + /// \brief Get the tensor metadata + /// \param tensor_name the tensor name + /// \return the tensor metadata + tensor_metadata get_tensor_metadata(std::string tensor_name); + + /// \brief Write the safetensors + /// \param output_path the output path + void write_safetensors(std::string output_path); + + /// \brief Get the metadata + /// \return the metadata + nlohmann::json get_metadata(); + + /// \brief Switch the model + /// \param new_model_path the new model path + void switch_model(std::string new_model_path); +}; diff --git a/src/include/typedef.hpp b/src/include/typedef.hpp index 60a312a6f..532242d31 100644 --- a/src/include/typedef.hpp +++ b/src/include/typedef.hpp @@ -16,6 +16,8 @@ #ifdef _WIN32 #include #endif +// bf16 has to be constructible from float for the engine sources. +#define BIOVAULT_BFLOAT16_CONVERTING_CONSTRUCTORS #include "biovault_bfloat16.h" typedef float f32; diff --git a/src/include/utils/avx512_util.hpp b/src/include/utils/avx512_util.hpp new file mode 100644 index 000000000..d3fd964f9 --- /dev/null +++ b/src/include/utils/avx512_util.hpp @@ -0,0 +1,287 @@ +#pragma once +#include +#include "typedef.hpp" +#include +#include +#include + + +/** + * @brief Helper function to load 16 bfloat16 values and convert to __m512 (float). + */ +inline __m512 load_bfloat16_to_m512(const bf16* ptr) { + // Load 16 bfloat16 values (32 bytes) into a __m256i + __m256i bf16_data = _mm256_loadu_si256(reinterpret_cast(ptr)); + + // Convert bfloat16 to float by shifting left 16 bits (bfloat16 is upper 16 bits of float) + __m512i shifted = _mm512_cvtepu16_epi32(bf16_data); + shifted = _mm512_slli_epi32(shifted, 16); + + return _mm512_castsi512_ps(shifted); +} + +/** + * @brief Helper function to store __m512 (float) as 16 bfloat16 values. + * Uses truncation (no rounding). + */ +inline void store_m512_to_bfloat16(bf16* ptr, __m512 data) { + // Convert float to bfloat16 by extracting upper 16 bits + __m512i int_data = _mm512_castps_si512(data); + __m512i shifted = _mm512_srli_epi32(int_data, 16); + __m256i bf16_data = _mm512_cvtepi32_epi16(shifted); + + _mm256_storeu_si256(reinterpret_cast<__m256i*>(ptr), bf16_data); +} + +/** + * @brief Helper function to store __m512 (float) as 16 bfloat16 values with rounding. + * Uses round-to-nearest-even for better accuracy. + */ +inline void store_m512_to_bfloat16_rne(bf16* ptr, __m512 data) { + // Convert float to bfloat16 with rounding to nearest even + __m512i int_data = _mm512_castps_si512(data); + + // Add 0x7FFF for round-to-nearest-even + __m512i rounding = _mm512_set1_epi32(0x7FFF); + __m512i rounded = _mm512_add_epi32(int_data, rounding); + + // Shift right by 16 to get bf16 in lower 16 bits + __m512i shifted = _mm512_srli_epi32(rounded, 16); + + // Pack to 16-bit values + __m256i bf16_data = _mm512_cvtepi32_epi16(shifted); + + _mm256_storeu_si256(reinterpret_cast<__m256i*>(ptr), bf16_data); +} + + + +// Fast, corrected AVX-512 exp approximation (single-precision). +// Notes: +// - Input x is clamped to [-88, 88] to avoid overflow/underflow. +// - Uses range reduction x = n*ln2 + r, where n is rounded to nearest int. +// - Uses a degree-5 polynomial for exp(r) evaluated with Horner + FMAs. +// - Constructs 2^n by writing the biased exponent field; the biased exponent +// is clamped to [0,255] as a safety measure. +// +// This is an approximation (not fully IEEE-754 accurate for all cases). +inline __m512 _mm512_exp_ps_corrected(__m512 x) { + // clamp x to a reasonable range to avoid overflow/underflow + const __m512 max_val = _mm512_set1_ps(88.0f); + const __m512 min_val = _mm512_set1_ps(-88.0f); + x = _mm512_min_ps(x, max_val); + x = _mm512_max_ps(x, min_val); + + // constants: 1/ln2 and split ln2 = ln2_hi + ln2_lo for extra precision + const __m512 ln2_inv = _mm512_set1_ps(1.44269504088896341f); // 1/ln(2) + const __m512 ln2_hi = _mm512_set1_ps(0.6931471824645996f); // hi part + const __m512 ln2_lo = _mm512_set1_ps(1.9082149292705877e-10f);// lo part + + // compute fx = x * (1/ln2) + __m512 fx = _mm512_mul_ps(x, ln2_inv); + + // round to nearest integer (using rounding intrinsic), storing integer-valued floats + fx = _mm512_roundscale_ps(fx, _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC); + + // convert to int32 (safe since fx holds integer values after rounding) + __m512i emm0 = _mm512_cvttps_epi32(fx); + + // convert back to float for range-reduction arithmetic + __m512 n_ps = _mm512_cvtepi32_ps(emm0); + + // r = x - n * ln2 (use fnmadd to compute c - a*b robustly) + // first r1 = x - n*ln2_hi + __m512 r = _mm512_fnmadd_ps(n_ps, ln2_hi, x); // r = x - n*ln2_hi + // then r = r - n*ln2_lo + r = _mm512_fnmadd_ps(n_ps, ln2_lo, r); // r = x - n*(ln2_hi + ln2_lo) + + // polynomial coefficients for exp(r) ~ 1 + r + r^2/2 + r^3/6 + r^4/24 + r^5/120 + const __m512 c5 = _mm512_set1_ps(0.008333333333333333f); // 1/120 + const __m512 c4 = _mm512_set1_ps(0.041666666666666664f); // 1/24 + const __m512 c3 = _mm512_set1_ps(0.16666666666666666f); // 1/6 + const __m512 c2 = _mm512_set1_ps(0.5f); // 1/2 + const __m512 c1 = _mm512_set1_ps(1.0f); + const __m512 one = _mm512_set1_ps(1.0f); + + // Horner evaluation using FMA: (((c5*r + c4)*r + c3)*r + c2)*r + c1 ; then final *r + 1 + __m512 y = _mm512_fmadd_ps(c5, r, c4); + y = _mm512_fmadd_ps(y, r, c3); + y = _mm512_fmadd_ps(y, r, c2); + y = _mm512_fmadd_ps(y, r, c1); + y = _mm512_fmadd_ps(y, r, one); // y now approximates exp(r) + + // Build 2^n by inserting biased exponent into float bits: + // biased = n + 127 + __m512i biased = _mm512_add_epi32(emm0, _mm512_set1_epi32(127)); + + // clamp biased exponent to [0,255] to avoid invalid bit patterns + biased = _mm512_max_epi32(biased, _mm512_set1_epi32(0)); + biased = _mm512_min_epi32(biased, _mm512_set1_epi32(255)); + + // shift into exponent position (bits 23..30) and reinterpret as float + biased = _mm512_slli_epi32(biased, 23); + __m512 pow2n = _mm512_castsi512_ps(biased); + + // final result: exp(x) ≈ exp(r) * 2^n + return _mm512_mul_ps(y, pow2n); +} + +// Fast AVX-512 log approximation (single-precision). +// Input x must be strictly positive. +inline __m512 _mm512_log_ps_approx(__m512 x) { + const __m512i inv_mant_mask = _mm512_set1_epi32(~0x7f800000); + const __m512i min_norm_pos = _mm512_set1_epi32(0x00800000); + const __m512i exponent_mask = _mm512_set1_epi32(0x7f800000); + const __m512 one = _mm512_set1_ps(1.0f); + + // Extract exponent + __m512i vx = _mm512_castps_si512(x); + __m512i emm0 = _mm512_srli_epi32(vx, 23); + emm0 = _mm512_sub_epi32(emm0, _mm512_set1_epi32(127)); + __m512 e = _mm512_cvtepi32_ps(emm0); + + // Extract mantissa and force exponent to 0 (which means range [1.0, 2.0)) + __m512i m_bits = _mm512_and_si512(vx, inv_mant_mask); + m_bits = _mm512_or_si512(m_bits, _mm512_set1_epi32(0x3f800000)); + __m512 m = _mm512_castsi512_ps(m_bits); + + // Map m from [1, 2) to a symmetric range using p = (m - 1) / (m + 1) + __m512 p1 = _mm512_sub_ps(m, one); + __m512 p2 = _mm512_add_ps(m, one); + __m512 p = _mm512_div_ps(p1, p2); + __m512 p_sq = _mm512_mul_ps(p, p); + + // Evaluate Taylor series for log((1+p)/(1-p)) = 2 * (p + p^3/3 + p^5/5 + p^7/7) + const __m512 c7 = _mm512_set1_ps(2.0f / 7.0f); + const __m512 c5 = _mm512_set1_ps(2.0f / 5.0f); + const __m512 c3 = _mm512_set1_ps(2.0f / 3.0f); + const __m512 c1 = _mm512_set1_ps(2.0f); + + __m512 res = _mm512_fmadd_ps(c7, p_sq, c5); + res = _mm512_fmadd_ps(res, p_sq, c3); + res = _mm512_fmadd_ps(res, p_sq, c1); + res = _mm512_mul_ps(res, p); + + // log(x) = res + e * ln(2) + const __m512 ln2 = _mm512_set1_ps(0.6931471805599453f); + return _mm512_fmadd_ps(e, ln2, res); +} + + +// AVX-512 GELU tanh-based, now using the corrected exp function +inline __m512 gelu_tanh_avx512_simd(__m512 gate_vec_fp32) { + // ---- GELU(gate) with tanh approximation: 0.5 * gate * (1 + tanh(sqrt(2/pi) * (gate + 0.044715 * gate^3))) ---- + const __m512 half = _mm512_set1_ps(0.5f); + const __m512 one = _mm512_set1_ps(1.0f); + const __m512 sqrt_2_pi = _mm512_set1_ps(0.7978845608f); // sqrt(2/pi) + const __m512 coeff = _mm512_set1_ps(0.044715f); + + // Compute gate^3 + __m512 gate_squared = _mm512_mul_ps(gate_vec_fp32, gate_vec_fp32); + __m512 gate_cubed = _mm512_mul_ps(gate_squared, gate_vec_fp32); + + // Compute gate + 0.044715 * gate^3 + __m512 inner_term = _mm512_fmadd_ps(coeff, gate_cubed, gate_vec_fp32); + + // Compute sqrt(2/pi) * (gate + 0.044715 * gate^3) + __m512 scaled_term = _mm512_mul_ps(sqrt_2_pi, inner_term); + + // Compute tanh using the corrected exp function: tanh(x) ≈ (exp(x) - exp(-x)) / (exp(x) + exp(-x)) + __m512 exp_pos = _mm512_exp_ps_corrected(scaled_term); + __m512 exp_neg = _mm512_exp_ps_corrected(_mm512_sub_ps(_mm512_setzero_ps(), scaled_term)); + + __m512 numerator = _mm512_sub_ps(exp_pos, exp_neg); + __m512 denominator = _mm512_add_ps(exp_pos, exp_neg); + __m512 tanh_approx = _mm512_div_ps(numerator, denominator); + + // Compute 1 + tanh(...) + __m512 one_plus_tanh = _mm512_add_ps(one, tanh_approx); + + // Compute 0.5 * gate * (1 + tanh(...)) + __m512 gelu = _mm512_mul_ps(half, _mm512_mul_ps(gate_vec_fp32, one_plus_tanh)); + return gelu; +} + + + + +// Vectorized gaussian function for 16 floats using AVX-512 exp approximation +inline __m512 gaussian_avx512(__m512 x, __m512 sigma) { + const __m512 one = _mm512_set1_ps(1.0f); + const __m512 two = _mm512_set1_ps(2.0f); + const __m512 zero = _mm512_setzero_ps(); + + // Check if sigma <= 0 + __mmask16 mask_zero_sigma = _mm512_cmp_ps_mask(sigma, zero, _CMP_LE_OQ); + + // Compute exp(-(x*x)/(2*sigma*sigma)) using fast AVX-512 approximation + __m512 x_sq = _mm512_mul_ps(x, x); + __m512 sigma_sq = _mm512_mul_ps(sigma, sigma); + __m512 two_sigma_sq = _mm512_mul_ps(two, sigma_sq); + + // Compute -(x*x)/(2*sigma*sigma) + __m512 neg_x_sq_over_2sigma_sq = _mm512_div_ps(_mm512_sub_ps(zero, x_sq), two_sigma_sq); + + // Apply fast exponential + __m512 exp_result = _mm512_exp_ps_corrected(neg_x_sq_over_2sigma_sq); + + // Return 1.0 if sigma <= 0, otherwise exp result + return _mm512_mask_blend_ps(mask_zero_sigma, exp_result, one); +} + +// Fast conversion from uint8 to float with normalization +inline void convert_uint8_to_float_avx512(const uint8_t* src, float* dst, size_t count) { + const size_t simd_count = count & ~15; // Process in chunks of 16 + + for (size_t i = 0; i < simd_count; i += 16) { + // Load 16 uint8 values + __m128i u8_vec = _mm_loadu_si128(reinterpret_cast(src + i)); + + // Convert to 32-bit integers + __m512i i32_vec = _mm512_cvtepu8_epi32(u8_vec); + + // Convert to float + __m512 f32_vec = _mm512_cvtepi32_ps(i32_vec); + + // Store result + _mm512_storeu_ps(dst + i, f32_vec); + } + + // Handle remaining elements + for (size_t i = simd_count; i < count; ++i) { + dst[i] = static_cast(src[i]); + } +} + +// Fast conversion from float to uint8 with clamping +inline void convert_float_to_uint8_avx512(const float* src, uint8_t* dst, size_t count) { + const __m512 zero = _mm512_setzero_ps(); + const __m512 max_val = _mm512_set1_ps(255.0f); + const size_t simd_count = count & ~15; // Process in chunks of 16 + + for (size_t i = 0; i < simd_count; i += 16) { + // Load 16 float values + __m512 f32_vec = _mm512_loadu_ps(src + i); + + // Round to nearest integer + f32_vec = _mm512_roundscale_ps(f32_vec, _MM_FROUND_TO_NEAREST_INT); + + // Clamp to [0, 255] + f32_vec = _mm512_max_ps(f32_vec, zero); + f32_vec = _mm512_min_ps(f32_vec, max_val); + + // Convert to 32-bit integers + __m512i i32_vec = _mm512_cvtps_epi32(f32_vec); + + // Pack to uint8 (with saturation) + __m128i u8_vec = _mm512_cvtusepi32_epi8(i32_vec); + + // Store result + _mm_storeu_si128(reinterpret_cast<__m128i*>(dst + i), u8_vec); + } + + // Handle remaining elements + for (size_t i = simd_count; i < count; ++i) { + dst[i] = static_cast(std::clamp(std::round(src[i]), 0.0f, 255.0f)); + } +} diff --git a/src/include/utils/debug_utils.hpp b/src/include/utils/debug_utils.hpp index 6106ae213..10a32d870 100644 --- a/src/include/utils/debug_utils.hpp +++ b/src/include/utils/debug_utils.hpp @@ -126,6 +126,9 @@ std::cerr << "\033[31m[" << header << "] " << oss.str() << "\033[0m" << std::endl; \ } while (0) +/// \brief Compile out a block unless DEBUG_LEVEL selects it. +#define DEBUG_BLOCK(level, operations) if (level <= DEBUG_LEVEL) { operations } + /// \brief header_print_g macro, in green color /// \param header the header of the message diff --git a/src/include/utils/error_measure.hpp b/src/include/utils/error_measure.hpp new file mode 100644 index 000000000..22e658d8e --- /dev/null +++ b/src/include/utils/error_measure.hpp @@ -0,0 +1,157 @@ +#ifndef __error_measure_cpu__ +#define __error_measure_cpu__ + +#include // Required for std::is_same_v +#include +#include // For std::min +template +float get_relativeL2(Ta* y, Tb* y_ref, int batch_size, int seq_len, int hidden_size, int seq_len_padded, int hidden_size_padded){ + // static_assert(std::is_same_v || std::is_same_v, + // "Error: T must be either float or std::bfloat16_t."); + float rmse = 0; + float ref_sum = 0; + for (int b = 0; b < batch_size; ++b) { + for (int s = 0; s < seq_len; ++s) { + for (int h = 0; h < hidden_size; ++h) { + long long y_idx = (long long)b * seq_len_padded * hidden_size_padded + (long long)s * hidden_size_padded + h; + long long y_ref_idx = (long long)b * seq_len * hidden_size + (long long)s * hidden_size + h; + ref_sum += (float)y_ref[y_ref_idx] * (float)y_ref[y_ref_idx]; + rmse += ((float)y[y_idx] - (float)y_ref[y_ref_idx]) * ((float)y[y_idx] - (float)y_ref[y_ref_idx]); + } + } + } + int y_size = batch_size * seq_len * hidden_size; + return sqrt(rmse / y_size) / sqrt(ref_sum / y_size); +} + +// y is [batch_size, seq_len_padded, hidden_size_padded] +// y_ref is [batch_size, seq_len, hidden_size] +template +float get_relativeL1(Ta* y, Tb *y_ref, int batch_size, int seq_len, int hidden_size, int seq_len_padded, int hidden_size_padded) +{ + float l1 = 0; + float ref_sum = 0; + for (int b = 0; b < batch_size; ++b) { + for (int s = 0; s < seq_len; ++s) { + for (int h = 0; h < hidden_size; ++h) { + long long y_idx = (long long)b * seq_len_padded * hidden_size_padded + (long long)s * hidden_size_padded + h; + long long y_ref_idx = (long long)b * seq_len * hidden_size + (long long)s * hidden_size + h; + ref_sum += abs((float)y_ref[y_ref_idx]); + l1 += abs((float)y[y_idx] - (float)y_ref[y_ref_idx]); + } + } + } + return l1 / ref_sum; +} + +template +float get_rmse(Ta *y, Tb *y_ref, + int batch_size, int seq_len, int hidden_size, + int seq_len_padded, int hidden_size_padded) +{ + double rmse = 0.0; + for (int b = 0; b < batch_size; ++b) { + for (int s = 0; s < seq_len; ++s) { + for (int h = 0; h < hidden_size; ++h) { + long long y_idx = (long long)b * seq_len_padded * hidden_size_padded + + (long long)s * hidden_size_padded + h; + long long y_ref_idx = (long long)b * seq_len * hidden_size + + (long long)s * hidden_size + h; + float dy = static_cast(y[y_idx]) - + static_cast(y_ref[y_ref_idx]); + rmse += dy * dy; + } + } + } + double y_size = (double)batch_size * seq_len * hidden_size; + return static_cast(sqrt(rmse / y_size)); +} + +// template +// float get_rmse(Ta *y, Tb *y_ref, int batch_size, int seq_len, int hidden_size, int seq_len_padded, int hidden_size_padded) +// { +// float rmse = 0; +// for (int b = 0; b < batch_size; ++b) { +// for (int s = 0; s < seq_len; ++s) { +// for (int h = 0; h < hidden_size; ++h) { +// long long y_idx = (long long)b * seq_len_padded * hidden_size_padded + (long long)s * hidden_size_padded + h; +// long long y_ref_idx = (long long)b * seq_len * hidden_size + (long long)s * hidden_size + h; +// rmse += ((float)y[y_idx] - (float)y_ref[y_ref_idx]) * ((float)y[y_idx] - (float)y_ref[y_ref_idx]); +// } +// } +// } +// int y_size = batch_size * seq_len * hidden_size; +// return sqrt(rmse / y_size); +// } + +template +float cal_rms_value(T* a, int a_size) { + float sum_sq = 0.0f; + for (int i = 0; i < a_size; ++i) { + sum_sq += static_cast(a[i]) * static_cast(a[i]); + } + return std::sqrt(sum_sq / a_size); +} +template +float get_cosine_similarity(Ta *y, Tb *y_ref, int batch_size, int seq_len, int hidden_size, int seq_len_padded, int hidden_size_padded) +{ + float dot_product = 0; + float norm_y = 0; + float norm_y_ref = 0; + for (int b = 0; b < batch_size; ++b) { + for (int s = 0; s < seq_len; ++s) { + for (int h = 0; h < hidden_size; ++h) { + long long y_idx = (long long)b * seq_len_padded * hidden_size_padded + (long long)s * hidden_size_padded + h; + long long y_ref_idx = (long long)b * seq_len * hidden_size + (long long)s * hidden_size + h; + dot_product += (float)y[y_idx] * (float)y_ref[y_ref_idx]; + norm_y += (float)y[y_idx] * (float)y[y_idx]; + norm_y_ref += (float)y_ref[y_ref_idx] * (float)y_ref[y_ref_idx]; + } + } + } + return dot_product / (sqrt(norm_y) * sqrt(norm_y_ref)); +} + + + +template +float get_max_abs_error(Ta *y, Tb *y_ref, + int batch_size, int seq_len, int hidden_size, + int seq_len_padded, int hidden_size_padded) +{ + float max_error = 0.0f; + for (int b = 0; b < batch_size; ++b) { + for (int s = 0; s < seq_len; ++s) { + for (int h = 0; h < hidden_size; ++h) { + long long y_idx = (long long)b * seq_len_padded * hidden_size_padded + + (long long)s * hidden_size_padded + h; + long long y_ref_idx = (long long)b * seq_len * hidden_size + + (long long)s * hidden_size + h; + float error = std::abs(static_cast(y[y_idx]) - + static_cast(y_ref[y_ref_idx])); + max_error = std::max(max_error, error); + } + } + } + return max_error; +} + +template +void print_error_metrics(T_a*a, T_b*b, int batch_size, int seq_len, int hidden_size, int seq_len_padded, int hidden_size_padded){ + + + float rmse = get_rmse(a, b, batch_size, seq_len, hidden_size, seq_len_padded, hidden_size_padded); + float relativeL1 = get_relativeL1(a, b, batch_size, seq_len, hidden_size, seq_len_padded, hidden_size_padded); + float relativeL2 = get_relativeL2(a, b, batch_size, seq_len, hidden_size, seq_len_padded, hidden_size_padded); + float cosine_similarity = get_cosine_similarity(a, b, batch_size, seq_len, hidden_size, seq_len_padded, hidden_size_padded); + float max_error = get_max_abs_error(a, b, batch_size, seq_len, hidden_size, seq_len_padded, hidden_size_padded); + + + + header_print("info", "Relative L1: " << relativeL1); + header_print("info", "Relative L2: " << relativeL2); + header_print("info", "Cosine similarity: " << cosine_similarity); + header_print("info", "RMSE: " << rmse << " | Max error: " << max_error); + +} +#endif \ No newline at end of file diff --git a/src/include/vision/norm.hpp b/src/include/vision/norm.hpp new file mode 100644 index 000000000..347175f04 --- /dev/null +++ b/src/include/vision/norm.hpp @@ -0,0 +1,48 @@ +#pragma once + + + +#include +#include +#include +#include +#include +#include "typedef.hpp" + + + + + +/** + * @brief High-precision Layer Normalization using scalar float32 operations (no AVX). + * + * Uses double precision for mean/variance accumulation and standard sqrt. + * This is slower but more accurate than AVX version - ideal for debugging. + */ +void layernorm_high_precision( + size_t N, + bf16* hidden_states, + bf16* weight, + bf16* bias, + float eps, + bf16* output); + + +/** + * @brief Parallel Layer Normalization across multiple sequences with OpenMP. + * + * Applies layer normalization to multiple sequences in parallel using up to 4 threads. + * Uses chunked distribution (static scheduling) for better cache locality. + * Automatically disables parallelization for small sequence counts. + */ +void layernorm_parallel( + int seq_len, + size_t hidden_dim, + size_t hidden_dim_padded, + bf16* hidden_states_base, + bf16* weight, + bf16* bias, + float eps, + bf16* output_base); + + diff --git a/src/include/weight_desc.hpp b/src/include/weight_desc.hpp new file mode 100644 index 000000000..7286d68a9 --- /dev/null +++ b/src/include/weight_desc.hpp @@ -0,0 +1,340 @@ +#ifndef __WEIGHT_DESC_HPP__ +#define __WEIGHT_DESC_HPP__ + +#include +#include +#include +#include +#include +#include +#include "buffer.hpp" +#include "utils/debug_utils.hpp" + +// Bit flags encoded into the dtype enum values below, describing the layout of +// a quantized weight: 8-bit vs 4-bit, presence of a zero-point, presence of bias. +#define IS_Q8_MASK 0x1 +#define HAS_ZP_MASK 0x2 +#define HAS_BIAS_MASK 0x4 + +const int FLM_TENSOR_MAX_DIMS = 4; + +// Quantization operates on 32x256 element blocks (the hardware tiling granularity). +constexpr int QXNX_ROW_BLOCK_SIZE = 32; +constexpr int QXNX_COL_BLOCK_SIZE = 256; + +// Scales/mins are shared by a group of 32 elements; q4_k additionally shares a +// bf16 S/M pair across a super-block of 256. +constexpr int GGML_GROUP_SIZE = 32; +constexpr int Q4K_SUPER_BLOCK_SIZE = 256; + +// Quantized dtypes (flm_q*) pack the flag bits above into their values, so a +// single comparison against flm_u8 separates quantized from plain dtypes. +typedef enum: uint8_t { + // Naming: q[b] — e.g. flm_q41b is 4-bit, has zero-point, has bias. + flm_q40 = 0, + flm_q40b = HAS_BIAS_MASK, + flm_q41 = HAS_ZP_MASK, + flm_q41b = HAS_ZP_MASK | HAS_BIAS_MASK, + flm_q80 = IS_Q8_MASK, + flm_q80b = IS_Q8_MASK | HAS_BIAS_MASK, + flm_q81 = IS_Q8_MASK | HAS_ZP_MASK, + flm_q81b = IS_Q8_MASK | HAS_ZP_MASK | HAS_BIAS_MASK, + // 4-bit with a uint8 scale and a uint8 min per 32-element group, re-fit by a + // bf16 S / bf16 M per 256-element super-block. Carries no flag bits: its + // layout is not expressible as "q4 plus an extra zero-point slice", so it is + // matched by identity everywhere instead of by mask. + flm_q4k, + // Plain (non-quantized) dtypes follow; all compare >= flm_u8. + flm_u8, + flm_i8, + flm_u16, + flm_i16, + flm_u32, + flm_i32, + flm_u64, + flm_i64, + flm_f32, + flm_bf16, + flm_unknown +} flm_dtype_t; + +/// \brief Returns whether a dtype is one of the quantized (flm_q*) types. +/// \param qtype The dtype to test. +/// \return True if quantized, false for plain dtypes. +inline bool is_quantize(flm_dtype_t qtype) { return qtype < flm_u8; } + +/// \brief Returns whether a dtype uses the q4_k (uint8 scale/min + bf16 S/M) layout. +/// \param qtype The dtype to test. +/// \return True for flm_q4k only. +inline bool is_q4_k(flm_dtype_t qtype) { return qtype == flm_q4k; } + +/// \brief Computes the on-device byte footprint of a quantized tensor. +/// +/// Data is laid out in fixed 512-byte slices; each 32x256 chunk contributes some +/// number of slices for its packed quant values + per-group scales, plus optional +/// zero-point/bias slices. q4_k is the exception: its chunk is not a whole number +/// of slices, so it is sized in bytes. +/// \param elems Total element count; must be a multiple of the 32x256 chunk size. +/// \param qtype The quantized dtype describing the layout. +/// \return The total byte size on device. +inline size_t get_quantization_byte_size(size_t elems, flm_dtype_t qtype) { + static constexpr size_t slice_size = 512; // Byte + static constexpr size_t chunk_elems = QXNX_ROW_BLOCK_SIZE * QXNX_COL_BLOCK_SIZE; + static constexpr size_t chunk_groups = chunk_elems / GGML_GROUP_SIZE; + static constexpr size_t chunk_supers = chunk_elems / Q4K_SUPER_BLOCK_SIZE; + assert(is_quantize(qtype)); // does not apply to other dtype + assert(elems % chunk_elems == 0); + + size_t chunks = elems / chunk_elems; + + if (is_q4_k(qtype)) { + // Half a byte per element, a uint8 scale and a uint8 min per group, and a + // bf16 S / bf16 M per super-block: 4736 B per chunk, i.e. 4.625 bits per + // weight. That is 9.25 slices, so this path counts bytes directly. + size_t chunk_bytes = chunk_elems / 2 // packed quants + + chunk_groups * 2 * sizeof(uint8_t) // scales + mins + + chunk_supers * 2 * sizeof(bf16); // S + M + return chunks * chunk_bytes; + } + + uint32_t slices = 0; + if ((qtype & IS_Q8_MASK) != 0) { + // 8-bit: 1 byte per element + one bf16 scale per 32-element group + slices += (chunk_elems + chunk_groups * sizeof(bf16)) / slice_size; // 16 quant + 1 scale + } + else { + // 4-bit: half a byte per element (hence / 2) + one bf16 scale per group + slices += (chunk_elems / 2 + chunk_groups * sizeof(bf16)) / slice_size; // / 2 for q4 + } + + // Zero-point and bias each occupy one additional slice per chunk when present. + if (qtype & HAS_ZP_MASK) { + slices += 1; + } + + if (qtype & HAS_BIAS_MASK) { + slices += 1; + } + + return chunks * slices * slice_size; +} + + +/// \brief Returns the per-element byte size of a plain dtype. +/// \param dtype The dtype to query. +/// \return Bytes per element; 1 for quantized/unknown types which have no fixed per-element size. +inline size_t get_dtype_byte(flm_dtype_t dtype) { + int dtype_size = -1; + switch (dtype) { + case flm_u8: + case flm_i8: + dtype_size= 1; + break; + case flm_u16: + case flm_i16: + case flm_bf16: + dtype_size = 2; + break; + case flm_u32: + case flm_i32: + case flm_f32: + dtype_size = 4; + break; + case flm_u64: + case flm_i64: + dtype_size = 8; + break; + case flm_unknown: + default: + dtype_size = 1; // quantized / unknown types have no fixed per-element size + break; + } + + return dtype_size; +} + + +/// \brief Fixed-rank (4D) tensor shape. +/// +/// Unused leading dimensions default to 1 so that elems() and comparisons work +/// regardless of the logical rank. +struct flm_shape_t { + std::array _data; + + /// \brief Constructs a shape with all dimensions set to 1. + flm_shape_t() { _data.fill(1); } + + /// \brief Constructs a shape from an initializer list, padding remaining dims with 1. + /// \param init Dimension values; entries beyond FLM_TENSOR_MAX_DIMS are ignored. + flm_shape_t(std::initializer_list init) { + _data.fill(1); + size_t i = 0; + for (int64_t v : init) { + if (i >= FLM_TENSOR_MAX_DIMS) break; + _data[i++] = v; + } + } + + /// \brief Accesses the dimension at the given index. + /// \param idx Dimension index in [0, FLM_TENSOR_MAX_DIMS). + /// \return Reference to the dimension value. + int64_t& operator[](size_t idx) { return _data[idx]; } + /// \brief Read-only access to the dimension at the given index. + /// \param idx Dimension index in [0, FLM_TENSOR_MAX_DIMS). + /// \return Const reference to the dimension value. + const int64_t& operator[](size_t idx) const { return _data[idx]; } + + auto begin() { return _data.begin(); } + auto end() { return _data.end(); } + auto begin() const { return _data.begin(); } + auto end() const { return _data.end(); } + /// \brief Returns the fixed rank of the shape. + /// \return Always FLM_TENSOR_MAX_DIMS. + constexpr size_t size() const { return FLM_TENSOR_MAX_DIMS; } + /// \brief Sets every dimension to the given value. + /// \param v Value to assign to all dimensions. + void fill(int64_t v) { _data.fill(v); } + /// \brief Computes the total number of elements. + /// \return Product of all four dimensions. + size_t elems() { return _data[0] * _data[1] * _data[2] * _data[3]; } + + /// \brief Equality comparison across all dimensions. + /// \param o Shape to compare against. + /// \return True if all dimensions match. + bool operator==(const flm_shape_t& o) const { return _data == o._data; } + /// \brief Inequality comparison across all dimensions. + /// \param o Shape to compare against. + /// \return True if any dimension differs. + bool operator!=(const flm_shape_t& o) const { return _data != o._data; } +}; + + +/// \brief Computes the byte size and block-unit shape of a quantized weight. +/// +/// Requires dims 0/1 to be exact multiples of the 256/32 block sizes. +/// \param dtype The quantized dtype. +/// \param shape The logical tensor shape. +/// \return A pair of (total byte size, shape rewritten with dims 0/1 divided by the block sizes). +inline std::pair quantize_rectify(flm_dtype_t& dtype, flm_shape_t& shape){ + // last 2 dims must satisfy 32x256 block + assert(shape[0] % QXNX_COL_BLOCK_SIZE == 0); + assert(shape[1] % QXNX_ROW_BLOCK_SIZE == 0); + + assert(shape[2] > 0); + assert(shape[3] > 0); + + size_t total_size = get_quantization_byte_size(shape.elems(), dtype); + + flm_shape_t new_shape = shape; + + new_shape[0] = shape[0] / QXNX_COL_BLOCK_SIZE; + new_shape[1] = shape[1] / QXNX_ROW_BLOCK_SIZE; + + return std::pair (total_size, new_shape); +} + +/// \brief Describes a single weight tensor: its shape, dtype, name, and location. +struct weight_desc_t { +public: + flm_shape_t shape; + flm_dtype_t dtype; + std::string name; + size_t offset; + + bool added; + bool loaded; + + /// \brief Default constructor; leaves fields uninitialized. + weight_desc_t() {} + + /// \brief Constructs a weight descriptor. + /// \param dtype The weight's dtype. + /// \param shape The weight's shape. + /// \param name The weight's name (may be a printf format string, see format_name). + weight_desc_t(flm_dtype_t dtype, flm_shape_t shape, std::string name) : + shape(shape), dtype(dtype), name(name), offset(0), added(false), loaded(false) {} + + /// \brief Marks the weight as registered (offset assigned). + void indp() { added = true; } + /// \brief Marks the weight's data as loaded. + void load() { loaded = true; } + + /// \brief Reports whether the weight is usable. + /// \return True once the weight has been both registered and loaded. + bool ready() { return added & loaded; } + + /// \brief Formats the name, treating it as a printf format string. + /// \param args Arguments substituted into the format string. + /// \return The formatted name. + template + std::string format_name(Args... args) { + int size = std::snprintf(nullptr, 0, name.c_str(), args...); + assert(size >= 0); + std::string out(size, '\0'); + std::snprintf(&out[0], size + 1, name.c_str(), args...); + return out; + } + + /// \brief Computes the weight's byte size. + /// \return Block-packed size for quantized dtypes, element-count * element-size otherwise. + size_t get_size() { + if (is_quantize(dtype)) { + auto new_size_shape = quantize_rectify(dtype, shape); + return new_size_shape.first; + } + else { + return shape.elems() * get_dtype_byte(dtype); + } + } + + /// \brief Produces a typed view into the shared parent buffer at this weight's offset. + /// \param parent_buffer The backing buffer holding all weights. + /// \return A buffer spanning this weight's region. + template + buffer locate_myself(bytes& parent_buffer){ + return buffer((T*)(parent_buffer.data() + offset), this->get_size() / sizeof(T)); + } +}; + +/// \brief Accumulates weights into a single contiguous region. +/// +/// Hands out an offset for each weight and tracks the running total size. +class weight_container{ + size_t total_size_counter; +public: + /// \brief Constructs an empty container with zero total size. + weight_container() { + total_size_counter = 0; + } + + /// \brief Reserves space for a weight of the given dtype/shape. + /// \param dtype The weight's dtype. + /// \param shape The weight's shape. + /// \return The offset at which the weight was placed. + size_t add_weight(flm_dtype_t& dtype, flm_shape_t& shape){ + size_t offset = total_size_counter; + if (is_quantize(dtype)) { + auto new_size_shape = quantize_rectify(dtype, shape); + total_size_counter += new_size_shape.first; + } + else { + total_size_counter += shape.elems() * get_dtype_byte(dtype); + } + return offset; + } + + /// \brief Reserves space for a weight and updates its descriptor in place. + /// \param weight The descriptor to place; its offset and added flag are set. + void add_weight(weight_desc_t& weight){ + weight.offset = this->add_weight(weight.dtype, weight.shape); + weight.added = true; + LOG_VERBOSE(1, "Added weight '" << weight.name << "' at offset " << weight.offset + << ", total size now " << total_size_counter); + } + + /// \brief Returns the total accumulated byte size of all added weights. + /// \return The running total size. + size_t get_size() { return total_size_counter; } +}; + +#endif \ No newline at end of file