Skip to content
Draft
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,13 @@ def setup_common_extension() -> CMakeExtension:
elif os.getenv("NVTE_FUSED_ATTN_CK") or os.getenv("NVTE_FUSED_ATTN"):
cmake_flags.append("-DUSE_FUSED_ATTN_CK=ON")

# AITER a4w4 (FP4) GEMM backend (gfx950-only; CMake disables it cleanly
# on other arches).
if int(os.getenv("NVTE_AITER_GEMM", "1"))==0:
cmake_flags.append("-DUSE_AITER_GEMM=OFF")
else:
cmake_flags.append("-DUSE_AITER_GEMM=ON")

if bool(int(os.getenv("NVTE_ENABLE_NVSHMEM", "0"))) and os.getenv("NVTE_ENABLE_ROCSHMEM") is None:
os.environ["NVTE_ENABLE_ROCSHMEM"] = '1'
os.environ["NVTE_ENABLE_NVSHMEM"] = '0'
Expand Down
20 changes: 20 additions & 0 deletions transformer_engine/common/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ option(USE_ROCM "Use ROCm" ON)
option(USE_FUSED_ATTN_AOTRITON "Use aotriton backend" ON)
option(USE_FUSED_ATTN_CK "Use ck backend" ON)
option(USE_HIPKITTENS_GEMM "Use HipKittens MXFP8 GEMM kernels" ON)
option(USE_AITER_GEMM "Use AITER a4w4 (FP4) GEMM backend" ON)
set(USE_CUDA OFF)

if (USE_ROCM)
Expand Down Expand Up @@ -336,6 +337,7 @@ if(USE_ROCM)
fused_attn_rocm/fused_attn_aotriton.cpp
fused_attn_rocm/fused_attn_ck.cpp
fused_attn_rocm/utils.cpp
gemm/aiter_a4w4_gemm.cpp
gemm/ck_grouped_gemm/ck_grouped_gemm.cpp
gemm/ck_grouped_gemm/ck_grouped_gemm_fp8.cpp
gemm/ck_grouped_gemm/ck_grouped_gemm_fp16.cpp
Expand Down Expand Up @@ -582,6 +584,20 @@ else() # USE_ROCM
add_subdirectory(gemm/kittens ${CMAKE_CURRENT_BINARY_DIR}/kittens)
endif()

# AITER a4w4 (FP4) GEMM is gfx950-only; skip cleanly on other arches so a
# default-on USE_AITER_GEMM does not break gfx942/gfx1250 builds. The NVTE
# C API in gemm/aiter_a4w4_gemm.cpp keeps stub symbols when disabled.
list(FIND CMAKE_HIP_ARCHITECTURES "gfx950" _aiter_gemm_gfx950_idx)
if(USE_AITER_GEMM AND NOT _aiter_gemm_gfx950_idx EQUAL -1)
set(__BUILD_AITER_GEMM TRUE)
add_subdirectory(aiter_gemm ${CMAKE_CURRENT_BINARY_DIR}/aiter_gemm)
else()
set(__BUILD_AITER_GEMM FALSE)
if(USE_AITER_GEMM)
message(STATUS "USE_AITER_GEMM requested but no gfx950 target present; a4w4 GEMM backend disabled.")
endif()
endif()

find_package(hip)
list(APPEND transformer_engine_LINKER_LIBS hip::host hip::device roctx64)
find_package(hiprtc)
Expand All @@ -600,6 +616,10 @@ else() # USE_ROCM
target_compile_definitions(transformer_engine PUBLIC USE_HIPKITTENS_GEMM)
list(APPEND transformer_engine_LINKER_LIBS kittens_gemm)
endif()
if(__BUILD_AITER_GEMM)
target_compile_definitions(transformer_engine PUBLIC USE_AITER_GEMM)
list(APPEND transformer_engine_LINKER_LIBS aiter_gemm)
endif()
target_link_libraries(transformer_engine PUBLIC ${transformer_engine_LINKER_LIBS})
endif()

Expand Down
127 changes: 127 additions & 0 deletions transformer_engine/common/aiter_gemm/CMakeLists.txt
Original file line number Diff line number Diff line change
@@ -0,0 +1,127 @@
# Copyright (c) 2026, Advanced Micro Devices, Inc. All rights reserved.
# SPDX-License-Identifier: MIT

cmake_minimum_required(VERSION 3.21)
set(CMAKE_CXX_STANDARD 17)
project(aiter_gemm LANGUAGES HIP CXX)

set(AITER_GEMM_INSTALL_DIR "${CMAKE_INSTALL_PREFIX}/transformer_engine/lib")

# a4w4 kernels are gfx950-only (ASM f4gemm .co blobs ship for gfx950). Drop
# unsupported arches; a runtime guard covers dispatch.
set(__AG_SUPPORTED_ARCHS "gfx950")
set(__AG_ARCHS)
foreach(__arch ${CMAKE_HIP_ARCHITECTURES})
if(__arch IN_LIST __AG_SUPPORTED_ARCHS)
list(APPEND __AG_ARCHS ${__arch})
else()
message(WARNING "[aiter_gemm] Skipping unsupported arch ${__arch} for a4w4 GEMM build.")
endif()
endforeach()
if(NOT __AG_ARCHS)
message(FATAL_ERROR
"[aiter_gemm] No supported architectures (need one of: ${__AG_SUPPORTED_ARCHS}). "
"Re-run the build with NVTE_AITER_GEMM=0 to disable the a4w4 GEMM backend.")
endif()

set(__QOLA_DIR "${CMAKE_CURRENT_LIST_DIR}/../../../3rdparty/QoLA")
set(__AITER_SOURCE_DIR "${CMAKE_CURRENT_BINARY_DIR}/qola/third_party/aiter")
if(DEFINED ENV{NVTE_AITER_SOURCE_DIR} AND NOT $ENV{NVTE_AITER_SOURCE_DIR} STREQUAL "")
set(__AITER_SOURCE_DIR $ENV{NVTE_AITER_SOURCE_DIR})
message(STATUS "[aiter_gemm] Using AITER source from NVTE_AITER_SOURCE_DIR=${__AITER_SOURCE_DIR}. Disable AITER checkout.")
set(__SKIP_AITER_CHECKOUT TRUE)
else()
set(__SKIP_AITER_CHECKOUT FALSE)
endif()
set(AITER_INCLUDE_DIR "${__AITER_SOURCE_DIR}/csrc/include")

if(NOT Python_EXECUTABLE)
find_package(Python COMPONENTS Interpreter QUIET)
endif()

# Resolve the manifest-pinned AITER commit (defines AITER_SHA).
set(__QOLA_MANIFEST "${CMAKE_CURRENT_LIST_DIR}/qola_manifest.toml")
set_property(DIRECTORY APPEND PROPERTY CMAKE_CONFIGURE_DEPENDS "${__QOLA_MANIFEST}")
file(STRINGS "${__QOLA_MANIFEST}" __AITER_COMMIT_LINES
REGEX "^[ \t]*aiter_commit[ \t]*=[ \t]*\"[^\"]+\"")
list(LENGTH __AITER_COMMIT_LINES __AITER_COMMIT_COUNT)
if(NOT __AITER_COMMIT_COUNT EQUAL 1)
message(FATAL_ERROR
"Expected exactly one 'aiter_commit = \"...\"' line in ${__QOLA_MANIFEST}.")
endif()
list(GET __AITER_COMMIT_LINES 0 __AITER_COMMIT_LINE)
string(REGEX MATCH "\"([^\"]+)\"" _UNUSED "${__AITER_COMMIT_LINE}")
set(AITER_SHA "${CMAKE_MATCH_1}")

if(Python_EXECUTABLE AND NOT __SKIP_AITER_CHECKOUT)
execute_process(
COMMAND sh -c
"PYTHONPATH=\"${__QOLA_DIR}:$PYTHONPATH\" '${Python_EXECUTABLE}' -m qola.cli checkout \
--manifest '${__QOLA_MANIFEST}' \
--aiter-root '${__AITER_SOURCE_DIR}'"
RESULT_VARIABLE AITER_CHECKOUT_RESULT
OUTPUT_VARIABLE AITER_CHECKOUT_OUTPUT
ERROR_VARIABLE AITER_CHECKOUT_ERROR)
if(NOT AITER_CHECKOUT_RESULT EQUAL 0)
message(FATAL_ERROR
"[aiter_gemm] Failed to sync AITER source at ${__AITER_SOURCE_DIR} to ${AITER_SHA}.\n"
"${AITER_CHECKOUT_OUTPUT}\n${AITER_CHECKOUT_ERROR}")
endif()
message(STATUS "[aiter_gemm] Synced ${__AITER_SOURCE_DIR} to ${AITER_SHA}")
endif()

if(NOT EXISTS "${AITER_INCLUDE_DIR}")
message(FATAL_ERROR
"[aiter_gemm] Could not find AITER API at ${AITER_INCLUDE_DIR}.")
endif()

# Obtain the torch-free kernel libs: prebuilt bypass or build from source (QoLA).
set(__AITER_GEMM_LIB_PATH "")
set(__QOLA_INCLUDE_DIR "")
if(DEFINED ENV{AITER_GEMM_PATH})
message(STATUS "[aiter_gemm] Using prebuilt libs from AITER_GEMM_PATH=$ENV{AITER_GEMM_PATH}")
set(__AITER_GEMM_LIB_PATH "$ENV{AITER_GEMM_PATH}/lib")
set(__QOLA_INCLUDE_DIR "$ENV{AITER_GEMM_PATH}/include")
else()
list(JOIN __AG_ARCHS ";" GPU_ARCHS_STR)
set(__QOLA_BUILD_DIR "${CMAKE_CURRENT_BINARY_DIR}/qola")
message(STATUS "[aiter_gemm] Building a4w4 GEMM kernels for ${GPU_ARCHS_STR} via QoLA.")
execute_process(
COMMAND ${CMAKE_COMMAND} -E env "PYTHONPATH=${__QOLA_DIR}:$ENV{PYTHONPATH}"
${Python_EXECUTABLE} -m qola.cli build
--manifest ${__QOLA_MANIFEST}
--aiter-root ${__AITER_SOURCE_DIR}
--output-dir ${__QOLA_BUILD_DIR}
--arch "${GPU_ARCHS_STR}"
--skip-checkout
RESULT_VARIABLE QOLA_BUILD_RESULT)
if(NOT QOLA_BUILD_RESULT EQUAL 0)
message(FATAL_ERROR "[aiter_gemm] QoLA build failed.")
endif()
set(__AITER_GEMM_LIB_PATH "${__QOLA_BUILD_DIR}/lib")
set(__QOLA_INCLUDE_DIR "${__QOLA_BUILD_DIR}/include")
endif()

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The QoLA build runs via execute_process at CMake configure time. CMAKE_CONFIGURE_DEPENDS on qola_manifest.toml (line 44) covers manifest edits, but a modified AITER source tree (e.g. NVTE_AITER_SOURCE_DIR pointing at a local checkout that the developer is iterating on) will not trigger a rebuild — only re-running cmake .. will. The aiter_gemm shared library also has no add_dependencies link back to the QoLA output, so .sos that get replaced under ${__QOLA_BUILD_DIR}/lib do not cause relinking.

For a first integration this is likely acceptable, but a follow-up switching to ExternalProject_Add (or a custom target with DEPENDS/BYPRODUCTS on the te_libgemm_a4w4_*.sos) would make iterative development on gfx950 far less error-prone. Worth at least a TODO comment.


if(NOT EXISTS "${__QOLA_INCLUDE_DIR}/qola_config.h")
message(FATAL_ERROR "[aiter_gemm] Could not find QoLA public headers at ${__QOLA_INCLUDE_DIR}.")
endif()

add_library(aiter_gemm SHARED src/aiter_gemm_a4w4.cpp)
target_include_directories(aiter_gemm PUBLIC "${CMAKE_CURRENT_SOURCE_DIR}/include")
target_include_directories(aiter_gemm PRIVATE ${AITER_INCLUDE_DIR} ${__QOLA_INCLUDE_DIR})

find_package(hip)
target_link_directories(aiter_gemm PUBLIC ${__AITER_GEMM_LIB_PATH})
target_link_libraries(aiter_gemm PUBLIC
hip::host hip::device
-l:te_libgemm_a4w4_blockscale.so
-l:te_libgemm_a4w4_asm.so)
set_target_properties(aiter_gemm PROPERTIES INSTALL_RPATH "$ORIGIN")

if(NOT "${__AITER_GEMM_LIB_PATH}" STREQUAL "${AITER_GEMM_INSTALL_DIR}")
install(FILES
${__AITER_GEMM_LIB_PATH}/te_libgemm_a4w4_blockscale.so
${__AITER_GEMM_LIB_PATH}/te_libgemm_a4w4_asm.so
DESTINATION ${AITER_GEMM_INSTALL_DIR})
endif()
install(TARGETS aiter_gemm DESTINATION ${AITER_GEMM_INSTALL_DIR})
Original file line number Diff line number Diff line change
@@ -0,0 +1,82 @@
/*************************************************************************
* Copyright (c) 2026, Advanced Micro Devices, Inc. All rights reserved.
*
* License for AMD contributions = MIT. See LICENSE for more information
************************************************************************/

#ifndef AITER_GEMM_H
#define AITER_GEMM_H

#include <cstdint>
#include <hip/hip_runtime.h>

// TE-side wrapper around AITER's torch-free a4w4 (FP4 x FP4) GEMM kernels.
//
// This header is intentionally free of any AITER / QoLA headers so that
// libtransformer_engine.so can consume it without pulling in the AITER kernel
// headers. The translation to AITER's aiter_tensor_t POD lives in the .cpp.
//
// Kernel selection (tuned-CSV lookup) and weight/scale pre-shuffling are the
// caller's responsibility -- these entry points are thin executors that take a
// resolved kernel name and already-shuffled inputs.
namespace aiter_gemm {

// Mirrors AiterDtype in AITER's aiter_enum.h (only the subset a4w4 needs).
enum class DType {
fp4x2 = 0, /*!< two packed FP4 (E2M1) values per byte */
e8m0 = 1, /*!< 8-bit exponent-only microscaling factor (1 byte) */
bf16 = 2,
fp16 = 3,
fp32 = 4,
u8 = 5,
i8 = 6,
};

// Lightweight tensor descriptor (raw device pointer + layout). The caller owns
// the storage; the descriptor must outlive the call but not the storage.
struct TensorDesc {
void* ptr = nullptr;
int ndim = 0;
int64_t shape[8] = {0};
int64_t strides[8] = {0};
DType dtype = DType::fp4x2;
int device_id = 0;
};

// CK blockscale a4w4 GEMM: Y = XQ @ WQ^T with per-1x32 microscaling.
// XQ [M, K/2] fp4x2
// WQ [N, K/2] fp4x2
// x_scale [M, K/32] e8m0
// w_scale [N, K/32] e8m0
// Y [M, N] bf16 / fp16 (output, pre-allocated)
// `kernel_name` empty -> default heuristic; non-empty must exist in the
// compiled registry. Returns hipSuccess on success, hipErrorUnknown on a
// kernel-side failure (message logged).
hipError_t gemm_a4w4_blockscale(const TensorDesc& XQ,
const TensorDesc& WQ,
const TensorDesc& x_scale,
const TensorDesc& w_scale,
const TensorDesc& Y,
int split_k,
const char* kernel_name,
hipStream_t stream);

// ASM (f4gemm) a4w4 GEMM: D = alpha*A*B + beta*C.
// A/B/scales/out layout as above; `bias` may be null.
// `kernel_name` empty -> ASM heuristic.
hipError_t gemm_a4w4_asm(const TensorDesc& A,
const TensorDesc& B,
const TensorDesc& a_scale,
const TensorDesc& b_scale,
const TensorDesc& out,
const TensorDesc* bias,
const char* kernel_name,
float alpha,
float beta,
int bpreshuffle,
int log2_k_split,
hipStream_t stream);

} // namespace aiter_gemm

#endif // AITER_GEMM_H
25 changes: 25 additions & 0 deletions transformer_engine/common/aiter_gemm/qola_manifest.toml
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
# Copyright (c) 2026, Advanced Micro Devices, Inc. All rights reserved.
# SPDX-License-Identifier: MIT
#
# QoLA consumer manifest for TE's AITER a4w4 (FP4) GEMM integration.
# Builds the torch-free CK-blockscale and ASM kernels as standalone
# te_libgemm_a4w4_*.so shared libraries (cpp_itfs mode).

[qola]
# NB: with NVTE_AITER_SOURCE_DIR set the build skips `qola checkout` and uses
# that tree directly; this pin documents the intended AITER commit for the
# reproducible (checkout) path.
aiter_commit = "773c8f67b9bfd99be95700a527d887afb4d8ba6a"
namespace = "te"
rocm_versions = ["7.2"]

[build]
architectures = ["gfx950"]

[[modules]]
name = "libgemm_a4w4_blockscale"
mode = "cpp_itfs"

[[modules]]
name = "libgemm_a4w4_asm"
mode = "cpp_itfs"
Loading
Loading