diff --git a/BUILD.bazel b/BUILD.bazel index cb93ac75..239178a8 100644 --- a/BUILD.bazel +++ b/BUILD.bazel @@ -34,6 +34,17 @@ config_setting( define_values = {"gemma_onednn_brgemm": "1"}, ) +# To enable the OneDNN matmul-primitive backend (threadpool runtime), build with: +# bazel build --define gemma_onednn_matmul=1 ... +# This is mutually exclusive with gemma_onednn_brgemm: passing both --defines +# makes the select()s below match two conditions at once, which Bazel rejects +# with an "ambiguous select" error -- the intended failure mode, mirroring the +# CMake FATAL_ERROR and the #error guard in ops/onednn_matmul.h. +config_setting( + name = "gemma_onednn_matmul", + define_values = {"gemma_onednn_matmul": "1"}, +) + cc_library( name = "basics", srcs = ["util/basics.cc"], @@ -363,12 +374,16 @@ cc_library( hdrs = [ "ops/brgemm.h", "ops/matmul.h", + "ops/onednn_matmul.h", ], defines = select({ ":gemma_onednn_brgemm": [ "GEMMA_ONEDNN_BRGEMM=1", "DNNL_EXPERIMENTAL_UKERNEL", ], + ":gemma_onednn_matmul": [ + "GEMMA_ONEDNN_MATMUL=1", + ], "//conditions:default": [], }), deps = [ @@ -384,6 +399,7 @@ cc_library( "@highway//:profiler", ] + select({ ":gemma_onednn_brgemm": ["@onednn"], + ":gemma_onednn_matmul": ["@onednn_tp//:onednn"], "//conditions:default": [], }), ) @@ -395,6 +411,7 @@ cc_library( textual_hdrs = [ "ops/brgemm-inl.h", "ops/matmul-inl.h", + "ops/onednn_matmul-inl.h", ], deps = [ ":allocator", @@ -412,6 +429,7 @@ cc_library( "@highway//:profiler", ] + select({ ":gemma_onednn_brgemm": ["@onednn"], + ":gemma_onednn_matmul": ["@onednn_tp//:onednn"], "//conditions:default": [], }), ) @@ -434,6 +452,7 @@ cc_library( "ops/brgemm-inl.h", "ops/matmul_static-inl.h", "ops/matmul-inl.h", + "ops/onednn_matmul-inl.h", ], deps = [ ":allocator", @@ -450,6 +469,7 @@ cc_library( "@highway//:timer", ] + select({ ":gemma_onednn_brgemm": ["@onednn"], + ":gemma_onednn_matmul": ["@onednn_tp//:onednn"], "//conditions:default": [], }), ) diff --git a/CMakeLists.txt b/CMakeLists.txt index 2bc3820a..23e78f19 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -33,6 +33,19 @@ set(CMAKE_EXPORT_COMPILE_COMMANDS ON) # Enable with: cmake -DGEMMA_ONEDNN_BRGEMM=ON ... option(GEMMA_ONEDNN_BRGEMM "Enable OneDNN BRGeMM micro-kernel for MatMul (x86-64)" OFF) +# Optional: OneDNN matmul primitive via the threadpool runtime (x86-64 only). +# Enable with: cmake -DGEMMA_ONEDNN_MATMUL=ON ... +option(GEMMA_ONEDNN_MATMUL "Enable OneDNN matmul primitive (threadpool runtime) for MatMul (x86-64)" OFF) + +# oneDNN's CPU runtime (SEQ vs THREADPOOL) is a whole-library compile-time +# choice, so the two oneDNN backends cannot share one oneDNN build. +if(GEMMA_ONEDNN_BRGEMM AND GEMMA_ONEDNN_MATMUL) + message(FATAL_ERROR + "GEMMA_ONEDNN_BRGEMM and GEMMA_ONEDNN_MATMUL are mutually exclusive: " + "oneDNN's CPU runtime (SEQ for BRGeMM vs THREADPOOL for the matmul " + "primitive) is a whole-library compile-time choice. Enable only one.") +endif() + if(EMSCRIPTEN) add_compile_options("-sMEMORY64") add_compile_options("-msimd128") @@ -119,6 +132,24 @@ if(GEMMA_ONEDNN_BRGEMM) message(STATUS "OneDNN BRGeMM micro-kernel support enabled") endif() +# OneDNN matmul primitive support via the threadpool runtime (optional, x86-64). +# Same oneDNN source as the BRGeMM backend, but built with the THREADPOOL CPU +# runtime instead of SEQ so oneDNN can call back into gemma.cpp's thread pool. +if(GEMMA_ONEDNN_MATMUL) + set(DNNL_BUILD_TESTS OFF CACHE BOOL "" FORCE) + set(DNNL_BUILD_EXAMPLES OFF CACHE BOOL "" FORCE) + set(DNNL_CPU_RUNTIME "THREADPOOL" CACHE STRING "" FORCE) + set(DNNL_GPU_RUNTIME "NONE" CACHE STRING "" FORCE) + set(DNNL_LIBRARY_TYPE "STATIC" CACHE STRING "" FORCE) + FetchContent_Declare(onednn + GIT_REPOSITORY https://github.com/uxlfoundation/oneDNN.git + GIT_TAG v3.11 + EXCLUDE_FROM_ALL + ) + FetchContent_MakeAvailable(onednn) + message(STATUS "OneDNN matmul primitive (threadpool runtime) support enabled") +endif() + # Base source files set(SOURCES compression/compress-inl.h @@ -193,6 +224,8 @@ set(SOURCES ops/matmul.h ops/brgemm.h ops/brgemm-inl.h + ops/onednn_matmul.h + ops/onednn_matmul-inl.h ops/ops-inl.h ops/ops.h ops/sum-inl.h @@ -255,6 +288,10 @@ if(GEMMA_ONEDNN_BRGEMM) target_compile_definitions(libgemma PUBLIC GEMMA_ONEDNN_BRGEMM=1 DNNL_EXPERIMENTAL_UKERNEL) target_link_libraries(libgemma dnnl) endif() +if(GEMMA_ONEDNN_MATMUL) + target_compile_definitions(libgemma PUBLIC GEMMA_ONEDNN_MATMUL=1) + target_link_libraries(libgemma dnnl) +endif() install(TARGETS libgemma DESTINATION lib) # Shared library target for C# interop @@ -283,6 +320,10 @@ if(GEMMA_ONEDNN_BRGEMM) target_compile_definitions(gemma_shared PUBLIC GEMMA_ONEDNN_BRGEMM=1 DNNL_EXPERIMENTAL_UKERNEL) target_link_libraries(gemma_shared PRIVATE dnnl) endif() +if(GEMMA_ONEDNN_MATMUL) + target_compile_definitions(gemma_shared PUBLIC GEMMA_ONEDNN_MATMUL=1) + target_link_libraries(gemma_shared PRIVATE dnnl) +endif() install(TARGETS gemma_shared DESTINATION lib) install(FILES gemma/c_api.h DESTINATION include/gemma) install(FILES gemma/GemmaInterop.cs DESTINATION include/gemma) diff --git a/MODULE.bazel b/MODULE.bazel index 00a47331..1c001090 100644 --- a/MODULE.bazel +++ b/MODULE.bazel @@ -36,6 +36,24 @@ http_archive( ], ) +# Same oneDNN v3.11 source as @onednn, but built for the THREADPOOL CPU runtime +# (the matmul-primitive backend, GEMMA_ONEDNN_MATMUL) rather than SEQ (BRGeMM). +# oneDNN's CPU runtime is a whole-library compile-time choice, so the two cannot +# share one build; this is a separate repo with a different build_file. Both +# build_files are one call to onednn_targets() in bazel/onednn.bzl with a +# different cpu_runtime. The tarball and sha256 are identical, so Bazel's +# download cache is shared and only the arm selected by //:gemma_onednn_matmul +# actually compiles. +http_archive( + name = "onednn_tp", + build_file = "@//bazel:onednn_threadpool.BUILD", + sha256 = "04df98b18300daf6c3aa7cc2d5e7ce8a8f430fed1787151daed0254d8dd4e64e", + strip_prefix = "oneDNN-3.11", + urls = [ + "https://github.com/uxlfoundation/oneDNN/archive/refs/tags/v3.11.tar.gz", + ], +) + http_archive( name = "com_google_absl_py", sha256 = "8a3d0830e4eb4f66c4fa907c06edf6ce1c719ced811a12e26d9d3162f8471758", diff --git a/bazel/onednn.BUILD b/bazel/onednn.BUILD index 0cbd436d..93b85921 100644 --- a/bazel/onednn.BUILD +++ b/bazel/onednn.BUILD @@ -1,227 +1,11 @@ -load("@bazel_skylib//rules:expand_template.bzl", "expand_template") +# oneDNN built for the SEQ CPU runtime, used by the BRGeMM ukernel backend +# (GEMMA_ONEDNN_BRGEMM); gemma.cpp owns all parallelism, so oneDNN itself must be +# single-threaded. +# +# Same sources as the THREADPOOL build in bazel/onednn_threadpool.BUILD; see +# bazel/onednn.bzl for exactly what the runtime changes and why the two cannot +# share one build. -exports_files(["LICENSE"]) +load("@gemma//bazel:onednn.bzl", "onednn_targets") -expand_template( - name = "dnnl_config_h", - out = "include/oneapi/dnnl/dnnl_config.h", - substitutions = { - "#cmakedefine DNNL_EXPERIMENTAL_UKERNEL": "#define DNNL_EXPERIMENTAL_UKERNEL 1", - "#cmakedefine DNNL_SAFE_RBP": "#undef DNNL_SAFE_RBP", - "#cmakedefine DNNL_CPU_THREADING_RUNTIME DNNL_RUNTIME_${DNNL_CPU_THREADING_RUNTIME}": "#define DNNL_CPU_THREADING_RUNTIME DNNL_RUNTIME_SEQ", - "#cmakedefine DNNL_CPU_RUNTIME DNNL_RUNTIME_${DNNL_CPU_RUNTIME}": "#define DNNL_CPU_RUNTIME DNNL_RUNTIME_SEQ", - "#cmakedefine DNNL_DISABLE_GPU_REF_KERNELS": "#define DNNL_DISABLE_GPU_REF_KERNELS", - "#cmakedefine DNNL_GPU_RUNTIME DNNL_RUNTIME_${DNNL_GPU_RUNTIME}": "#define DNNL_GPU_RUNTIME DNNL_RUNTIME_NONE", - "#cmakedefine DNNL_GPU_VENDOR DNNL_VENDOR_${DNNL_GPU_VENDOR}": "#define DNNL_GPU_VENDOR DNNL_VENDOR_NONE", - "#cmakedefine DNNL_USE_RT_OBJECTS_IN_PRIMITIVE_CACHE": "#undef DNNL_USE_RT_OBJECTS_IN_PRIMITIVE_CACHE", - "#cmakedefine DNNL_WITH_SYCL": "#undef DNNL_WITH_SYCL", - "#cmakedefine DNNL_WITH_LEVEL_ZERO": "#undef DNNL_WITH_LEVEL_ZERO", - "#cmakedefine DNNL_SYCL_CUDA": "#undef DNNL_SYCL_CUDA", - "#cmakedefine DNNL_SYCL_GENERIC": "#undef DNNL_SYCL_GENERIC", - "#cmakedefine DNNL_SYCL_HIP": "#undef DNNL_SYCL_HIP", - "#cmakedefine DNNL_ENABLE_STACK_CHECKER": "#undef DNNL_ENABLE_STACK_CHECKER", - "#cmakedefine ONEDNN_BUILD_GRAPH": "#define ONEDNN_BUILD_GRAPH", - "#cmakedefine DNNL_EXPERIMENTAL_SPARSE": "#undef DNNL_EXPERIMENTAL_SPARSE", - "#cmakedefine DNNL_EXPERIMENTAL_LOGGING": "#undef DNNL_EXPERIMENTAL_LOGGING", - "#cmakedefine DNNL_EXPERIMENTAL_PROFILING": "#undef DNNL_EXPERIMENTAL_PROFILING", - "#cmakedefine DNNL_EXPERIMENTAL_SYCL_KERNEL_COMPILER": "#undef DNNL_EXPERIMENTAL_SYCL_KERNEL_COMPILER", - "#cmakedefine DNNL_EXPERIMENTAL": "#undef DNNL_EXPERIMENTAL", - "#cmakedefine01 BUILD_TRAINING": "#define BUILD_TRAINING 1", - "#cmakedefine01 BUILD_INFERENCE": "#define BUILD_INFERENCE 0", - "#cmakedefine01 BUILD_PRIMITIVE_ALL": "#define BUILD_PRIMITIVE_ALL 1", - "#cmakedefine01 BUILD_BATCH_NORMALIZATION": "#define BUILD_BATCH_NORMALIZATION 0", - "#cmakedefine01 BUILD_BINARY": "#define BUILD_BINARY 0", - "#cmakedefine01 BUILD_CONCAT": "#define BUILD_CONCAT 0", - "#cmakedefine01 BUILD_CONVOLUTION": "#define BUILD_CONVOLUTION 0", - "#cmakedefine01 BUILD_DECONVOLUTION": "#define BUILD_DECONVOLUTION 0", - "#cmakedefine01 BUILD_ELTWISE": "#define BUILD_ELTWISE 0", - "#cmakedefine01 BUILD_GEMM_KERNELS_ALL": "#define BUILD_GEMM_KERNELS_ALL 1", - "#cmakedefine01 BUILD_GEMM_KERNELS_NONE": "#define BUILD_GEMM_KERNELS_NONE 0", - "#cmakedefine01 BUILD_GEMM_SSE41": "#define BUILD_GEMM_SSE41 1", - "#cmakedefine01 BUILD_GEMM_AVX2": "#define BUILD_GEMM_AVX2 1", - "#cmakedefine01 BUILD_GEMM_AVX512": "#define BUILD_GEMM_AVX512 1", - "#cmakedefine01 BUILD_GROUP_NORMALIZATION": "#define BUILD_GROUP_NORMALIZATION 1", - "#cmakedefine01 BUILD_INNER_PRODUCT": "#define BUILD_INNER_PRODUCT 0", - "#cmakedefine01 BUILD_LAYER_NORMALIZATION": "#define BUILD_LAYER_NORMALIZATION 0", - "#cmakedefine01 BUILD_LRN": "#define BUILD_LRN 0", - "#cmakedefine01 BUILD_MATMUL": "#define BUILD_MATMUL 0", - "#cmakedefine01 BUILD_POOLING": "#define BUILD_POOLING 0", - "#cmakedefine01 BUILD_PRELU": "#define BUILD_PRELU 0", - "#cmakedefine01 BUILD_REDUCTION": "#define BUILD_REDUCTION 0", - "#cmakedefine01 BUILD_REORDER": "#define BUILD_REORDER 0", - "#cmakedefine01 BUILD_RESAMPLING": "#define BUILD_RESAMPLING 0", - "#cmakedefine01 BUILD_RNN": "#define BUILD_RNN 0", - "#cmakedefine01 BUILD_SHUFFLE": "#define BUILD_SHUFFLE 0", - "#cmakedefine01 BUILD_SOFTMAX": "#define BUILD_SOFTMAX 0", - "#cmakedefine01 BUILD_SUM": "#define BUILD_SUM 0", - "#cmakedefine01 BUILD_PRIMITIVE_CPU_ISA_ALL": "#define BUILD_PRIMITIVE_CPU_ISA_ALL 1", - "#cmakedefine01 BUILD_SSE41": "#define BUILD_SSE41 0", - "#cmakedefine01 BUILD_AVX2": "#define BUILD_AVX2 0", - "#cmakedefine01 BUILD_AVX512": "#define BUILD_AVX512 0", - "#cmakedefine01 BUILD_AMX": "#define BUILD_AMX 0", - "#cmakedefine01 BUILD_PRIMITIVE_GPU_ISA_ALL": "#define BUILD_PRIMITIVE_GPU_ISA_ALL 0", - "#cmakedefine01 BUILD_XE2": "#define BUILD_XE2 0", - "#cmakedefine01 BUILD_XELP": "#define BUILD_XELP 0", - "#cmakedefine01 BUILD_XEHPG": "#define BUILD_XEHPG 0", - "#cmakedefine01 BUILD_XEHPC": "#define BUILD_XEHPC 0", - "#cmakedefine01 BUILD_XEHP": "#define BUILD_XEHP 0", - "#cmakedefine01 BUILD_SDPA": "#define BUILD_SDPA 1", - "#cmakedefine01 BUILD_XE3": "#define BUILD_XE3 0", - }, - template = "include/oneapi/dnnl/dnnl_config.h.in", -) - -expand_template( - name = "dnnl_version_h", - out = "include/oneapi/dnnl/dnnl_version.h", - substitutions = { - "@DNNL_VERSION_MAJOR@": "3", - "@DNNL_VERSION_MINOR@": "11", - "@DNNL_VERSION_PATCH@": "0", - }, - template = "include/oneapi/dnnl/dnnl_version.h.in", -) - -expand_template( - name = "dnnl_version_hash_h", - out = "include/oneapi/dnnl/dnnl_version_hash.h", - substitutions = { - "@DNNL_VERSION_HASH@": "fc6151651a4577beae5ffac5a4132e75d39e1409", - }, - template = "include/oneapi/dnnl/dnnl_version_hash.h.in", -) - -cc_library( - name = "onednn_autogen", - srcs = glob(["src/cpu/x64/gemm/**/*_kern_autogen*.cpp"]), - copts = [ - "-O1", - "-U_FORTIFY_SOURCE", - "-fexceptions", - "-UUSE_MKL", - "-UUSE_CBLAS", - "-DDNNL_ENABLE_MAX_CPU_ISA", - "-DDNNL_ENABLE_ITT_TASKS", - "-DDNNL_ENABLE_GRAPH_DUMP", - "-DDNNL_EXPERIMENTAL_UKERNEL", - ], - includes = [ - "include", - "src", - "src/common", - "src/cpu", - "src/cpu/gemm", - "src/graph", - "third_party", - "third_party/ittnotify", - "third_party/xbyak", - ], - textual_hdrs = glob([ - "include/**/*", - "src/common/*.hpp", - "src/cpu/*.hpp", - "src/cpu/**/*.hpp", - "src/cpu/jit_utils/**/*.hpp", - "src/graph/interface/*.hpp", - "src/graph/backend/*.hpp", - "src/graph/backend/dnnl/*.hpp", - "src/graph/backend/dnnl/executables/*.hpp", - "src/graph/backend/fake/*.hpp", - "src/graph/backend/dnnl/passes/*.hpp", - "src/graph/backend/dnnl/patterns/*.hpp", - "src/graph/backend/dnnl/kernels/*.hpp", - "src/graph/utils/*.hpp", - "src/graph/utils/pm/*.hpp", - "third_party/ittnotify/**/*.h", - "third_party/spdlog/**/*.h", - "third_party/xbyak/*.h", - ]) + [ - ":dnnl_config_h", - ":dnnl_version_h", - ":dnnl_version_hash_h", - ], - visibility = ["//visibility:public"], -) - -cc_library( - name = "onednn", - srcs = glob( - [ - "src/common/*.cpp", - "src/cpu/*.cpp", - "src/cpu/**/*.cpp", - "src/cpu/jit_utils/**/*.cpp", - "src/cpu/x64/**/*.cpp", - "src/graph/interface/*.cpp", - "src/graph/backend/*.cpp", - "src/graph/backend/dnnl/*.cpp", - "src/graph/backend/dnnl/executables/*.cpp", - "src/graph/backend/fake/*.cpp", - "src/graph/backend/dnnl/passes/*.cpp", - "src/graph/backend/dnnl/patterns/*.cpp", - "src/graph/backend/dnnl/kernels/*.cpp", - "src/graph/utils/*.cpp", - "src/graph/utils/pm/*.cpp", - "third_party/ittnotify/*.c", - ], - exclude = [ - "src/cpu/aarch64/**", - "src/cpu/rv64/**", - "src/cpu/ppc64/**", - "src/cpu/s390x/**", - "src/cpu/x64/gemm/**/*_kern_autogen.cpp", - "src/cpu/sycl/**", - ], - ), - copts = [ - "-fexceptions", - "-UUSE_MKL", - "-UUSE_CBLAS", - "-DDNNL_ENABLE_MAX_CPU_ISA", - "-DDNNL_ENABLE_ITT_TASKS", - "-DDNNL_ENABLE_GRAPH_DUMP", - "-DDNNL_EXPERIMENTAL_UKERNEL", - ], - includes = [ - "include", - "src", - "src/common", - "src/cpu", - "src/cpu/gemm", - "src/graph", - "third_party", - "third_party/ittnotify", - "third_party/xbyak", - ], - linkopts = [ - "-lrt", - "-Wl,--allow-multiple-definition", - ], - textual_hdrs = glob([ - "include/**/*", - "src/common/*.hpp", - "src/cpu/*.hpp", - "src/cpu/**/*.hpp", - "src/cpu/jit_utils/**/*.hpp", - "src/graph/interface/*.hpp", - "src/graph/backend/*.hpp", - "src/graph/backend/dnnl/*.hpp", - "src/graph/backend/fake/*.hpp", - "src/graph/backend/dnnl/passes/*.hpp", - "src/graph/backend/dnnl/patterns/*.hpp", - "src/graph/backend/dnnl/kernels/*.hpp", - "src/graph/utils/*.hpp", - "src/graph/utils/pm/*.hpp", - "third_party/ittnotify/**/*.h", - "third_party/spdlog/**/*.h", - "third_party/xbyak/*.h", - ]) + [ - ":dnnl_config_h", - ":dnnl_version_h", - ":dnnl_version_hash_h", - ], - visibility = ["//visibility:public"], - deps = [ - ":onednn_autogen", - ], -) +onednn_targets(cpu_runtime = "SEQ") diff --git a/bazel/onednn.bzl b/bazel/onednn.bzl new file mode 100644 index 00000000..9e6942db --- /dev/null +++ b/bazel/onednn.bzl @@ -0,0 +1,284 @@ +# Shared Bazel definitions for building oneDNN from source. +# +# oneDNN's CPU runtime is a whole-library compile-time choice, and gemma.cpp has +# two oneDNN backends that need different ones: +# +# * SEQ -- ops/brgemm.h (GEMMA_ONEDNN_BRGEMM) drives oneDNN's low-level +# BRGeMM ukernel API; gemma.cpp owns all parallelism, so oneDNN +# itself must be single-threaded. +# * THREADPOOL -- ops/onednn_matmul.h (GEMMA_ONEDNN_MATMUL) uses the high-level +# dnnl::matmul primitive, which parallelizes internally by +# calling back into gemma.cpp's thread pool via the adapter. +# +# One oneDNN build cannot be both, so MODULE.bazel declares two http_archives +# over the identical tarball (@onednn and @onednn_tp, sharing the download +# cache). Their BUILD files are the same except for five generated config lines +# and one copt, so both just call onednn_targets() below with a different +# cpu_runtime. Only the arm selected by //:gemma_onednn_brgemm or +# //:gemma_onednn_matmul is ever compiled. + +load("@bazel_skylib//rules:expand_template.bzl", "expand_template") + +_VERSION_MAJOR = "3" + +_VERSION_MINOR = "11" + +_VERSION_PATCH = "0" + +_VERSION_HASH = "fc6151651a4577beae5ffac5a4132e75d39e1409" + +_INCLUDES = [ + "include", + "src", + "src/common", + "src/cpu", + "src/cpu/gemm", + "src/graph", + "third_party", + "third_party/ittnotify", + "third_party/xbyak", +] + +_BASE_COPTS = [ + "-fexceptions", + "-UUSE_MKL", + "-UUSE_CBLAS", + "-DDNNL_ENABLE_MAX_CPU_ISA", + "-DDNNL_ENABLE_ITT_TASKS", + "-DDNNL_ENABLE_GRAPH_DUMP", +] + +_TEXTUAL_HDR_PATTERNS = [ + "include/**/*", + "src/common/*.hpp", + "src/cpu/*.hpp", + "src/cpu/**/*.hpp", + "src/cpu/jit_utils/**/*.hpp", + "src/graph/interface/*.hpp", + "src/graph/backend/*.hpp", + "src/graph/backend/dnnl/*.hpp", + "src/graph/backend/fake/*.hpp", + "src/graph/backend/dnnl/passes/*.hpp", + "src/graph/backend/dnnl/patterns/*.hpp", + "src/graph/backend/dnnl/kernels/*.hpp", + "src/graph/utils/*.hpp", + "src/graph/utils/pm/*.hpp", + "third_party/ittnotify/**/*.h", + "third_party/spdlog/**/*.h", + "third_party/xbyak/*.h", +] + +_GENERATED_HDRS = [ + ":dnnl_config_h", + ":dnnl_version_h", + ":dnnl_version_hash_h", +] + +# Substitutions for include/oneapi/dnnl/dnnl_config.h.in that are the same for +# every runtime. The runtime-dependent ones are added by _config_substitutions. +# +# ORDER MATTERS. expand_template applies substitutions in dict order, and these +# keys are not mutually exclusive: "#cmakedefine DNNL_EXPERIMENTAL" also matches +# the start of "#cmakedefine DNNL_EXPERIMENTAL_SPARSE" and of the _UKERNEL line +# that _config_substitutions adds. A more specific key must come before any key +# it starts with, otherwise the general rule rewrites the line first and the +# specific one silently never matches. _check_substitution_order enforces this. +_COMMON_CONFIG_SUBSTITUTIONS = { + "#cmakedefine DNNL_SAFE_RBP": "#undef DNNL_SAFE_RBP", + "#cmakedefine DNNL_DISABLE_GPU_REF_KERNELS": "#define DNNL_DISABLE_GPU_REF_KERNELS", + "#cmakedefine DNNL_GPU_RUNTIME DNNL_RUNTIME_${DNNL_GPU_RUNTIME}": "#define DNNL_GPU_RUNTIME DNNL_RUNTIME_NONE", + "#cmakedefine DNNL_GPU_VENDOR DNNL_VENDOR_${DNNL_GPU_VENDOR}": "#define DNNL_GPU_VENDOR DNNL_VENDOR_NONE", + "#cmakedefine DNNL_USE_RT_OBJECTS_IN_PRIMITIVE_CACHE": "#undef DNNL_USE_RT_OBJECTS_IN_PRIMITIVE_CACHE", + "#cmakedefine DNNL_WITH_SYCL": "#undef DNNL_WITH_SYCL", + "#cmakedefine DNNL_WITH_LEVEL_ZERO": "#undef DNNL_WITH_LEVEL_ZERO", + "#cmakedefine DNNL_SYCL_CUDA": "#undef DNNL_SYCL_CUDA", + "#cmakedefine DNNL_SYCL_GENERIC": "#undef DNNL_SYCL_GENERIC", + "#cmakedefine DNNL_SYCL_HIP": "#undef DNNL_SYCL_HIP", + "#cmakedefine DNNL_ENABLE_STACK_CHECKER": "#undef DNNL_ENABLE_STACK_CHECKER", + "#cmakedefine ONEDNN_BUILD_GRAPH": "#define ONEDNN_BUILD_GRAPH", + "#cmakedefine DNNL_EXPERIMENTAL_SPARSE": "#undef DNNL_EXPERIMENTAL_SPARSE", + "#cmakedefine DNNL_EXPERIMENTAL_LOGGING": "#undef DNNL_EXPERIMENTAL_LOGGING", + "#cmakedefine DNNL_EXPERIMENTAL_PROFILING": "#undef DNNL_EXPERIMENTAL_PROFILING", + "#cmakedefine DNNL_EXPERIMENTAL_SYCL_KERNEL_COMPILER": "#undef DNNL_EXPERIMENTAL_SYCL_KERNEL_COMPILER", + "#cmakedefine DNNL_EXPERIMENTAL": "#undef DNNL_EXPERIMENTAL", + "#cmakedefine01 BUILD_TRAINING": "#define BUILD_TRAINING 1", + "#cmakedefine01 BUILD_INFERENCE": "#define BUILD_INFERENCE 0", + "#cmakedefine01 BUILD_PRIMITIVE_ALL": "#define BUILD_PRIMITIVE_ALL 1", + "#cmakedefine01 BUILD_BATCH_NORMALIZATION": "#define BUILD_BATCH_NORMALIZATION 0", + "#cmakedefine01 BUILD_BINARY": "#define BUILD_BINARY 0", + "#cmakedefine01 BUILD_CONCAT": "#define BUILD_CONCAT 0", + "#cmakedefine01 BUILD_CONVOLUTION": "#define BUILD_CONVOLUTION 0", + "#cmakedefine01 BUILD_DECONVOLUTION": "#define BUILD_DECONVOLUTION 0", + "#cmakedefine01 BUILD_ELTWISE": "#define BUILD_ELTWISE 0", + "#cmakedefine01 BUILD_GEMM_KERNELS_ALL": "#define BUILD_GEMM_KERNELS_ALL 1", + "#cmakedefine01 BUILD_GEMM_KERNELS_NONE": "#define BUILD_GEMM_KERNELS_NONE 0", + "#cmakedefine01 BUILD_GEMM_SSE41": "#define BUILD_GEMM_SSE41 1", + "#cmakedefine01 BUILD_GEMM_AVX2": "#define BUILD_GEMM_AVX2 1", + "#cmakedefine01 BUILD_GEMM_AVX512": "#define BUILD_GEMM_AVX512 1", + "#cmakedefine01 BUILD_GROUP_NORMALIZATION": "#define BUILD_GROUP_NORMALIZATION 1", + "#cmakedefine01 BUILD_INNER_PRODUCT": "#define BUILD_INNER_PRODUCT 0", + "#cmakedefine01 BUILD_LAYER_NORMALIZATION": "#define BUILD_LAYER_NORMALIZATION 0", + "#cmakedefine01 BUILD_LRN": "#define BUILD_LRN 0", + "#cmakedefine01 BUILD_POOLING": "#define BUILD_POOLING 0", + "#cmakedefine01 BUILD_PRELU": "#define BUILD_PRELU 0", + "#cmakedefine01 BUILD_REDUCTION": "#define BUILD_REDUCTION 0", + "#cmakedefine01 BUILD_RESAMPLING": "#define BUILD_RESAMPLING 0", + "#cmakedefine01 BUILD_RNN": "#define BUILD_RNN 0", + "#cmakedefine01 BUILD_SHUFFLE": "#define BUILD_SHUFFLE 0", + "#cmakedefine01 BUILD_SOFTMAX": "#define BUILD_SOFTMAX 0", + "#cmakedefine01 BUILD_SUM": "#define BUILD_SUM 0", + "#cmakedefine01 BUILD_PRIMITIVE_CPU_ISA_ALL": "#define BUILD_PRIMITIVE_CPU_ISA_ALL 1", + "#cmakedefine01 BUILD_SSE41": "#define BUILD_SSE41 0", + "#cmakedefine01 BUILD_AVX2": "#define BUILD_AVX2 0", + "#cmakedefine01 BUILD_AVX512": "#define BUILD_AVX512 0", + "#cmakedefine01 BUILD_AMX": "#define BUILD_AMX 0", + "#cmakedefine01 BUILD_PRIMITIVE_GPU_ISA_ALL": "#define BUILD_PRIMITIVE_GPU_ISA_ALL 0", + "#cmakedefine01 BUILD_XE2": "#define BUILD_XE2 0", + "#cmakedefine01 BUILD_XELP": "#define BUILD_XELP 0", + "#cmakedefine01 BUILD_XEHPG": "#define BUILD_XEHPG 0", + "#cmakedefine01 BUILD_XEHPC": "#define BUILD_XEHPC 0", + "#cmakedefine01 BUILD_XEHP": "#define BUILD_XEHP 0", + "#cmakedefine01 BUILD_SDPA": "#define BUILD_SDPA 1", + "#cmakedefine01 BUILD_XE3": "#define BUILD_XE3 0", +} + +def _check_substitution_order(substitutions): + """Fails if a general key precedes a more specific one it would shadow.""" + keys = substitutions.keys() + for i in range(len(keys)): + for j in range(i + 1, len(keys)): + if keys[j].startswith(keys[i]): + fail("dnnl_config.h substitution %r precedes %r, which it starts " % (keys[i], keys[j]) + + "with, so the second would never match. Move it earlier.") + return substitutions + +def _config_substitutions(cpu_runtime): + """Returns the dnnl_config.h.in substitutions for one CPU runtime.""" + + # The ukernel API is what the BRGeMM backend calls, and the matmul primitive + # plus its weights reorder are what the primitive backend calls. Enabling + # only what a backend needs keeps the other backend's code out of the build. + ukernel = cpu_runtime == "SEQ" + primitive = cpu_runtime == "THREADPOOL" + + # Runtime-dependent keys go first so that "#cmakedefine DNNL_EXPERIMENTAL" + # in the common set cannot shadow the _UKERNEL line. See the comment there. + substitutions = { + "#cmakedefine DNNL_EXPERIMENTAL_UKERNEL": ( + "#define DNNL_EXPERIMENTAL_UKERNEL 1" if ukernel else "#undef DNNL_EXPERIMENTAL_UKERNEL" + ), + "#cmakedefine DNNL_CPU_THREADING_RUNTIME DNNL_RUNTIME_${DNNL_CPU_THREADING_RUNTIME}": ( + "#define DNNL_CPU_THREADING_RUNTIME DNNL_RUNTIME_" + cpu_runtime + ), + "#cmakedefine DNNL_CPU_RUNTIME DNNL_RUNTIME_${DNNL_CPU_RUNTIME}": ( + "#define DNNL_CPU_RUNTIME DNNL_RUNTIME_" + cpu_runtime + ), + # BUILD_PRIMITIVE_ALL is 1, which already registers every primitive, so + # these two are redundant today. They are pinned on explicitly for the + # primitive backend because if a future trim sets BUILD_PRIMITIVE_ALL to + # 0 without flipping them, every DoMatMul_OneDnn call throws + # "unimplemented" and silently falls back to the stock path with zero + # speedup -- the single easiest mistake to make here. + "#cmakedefine01 BUILD_MATMUL": "#define BUILD_MATMUL " + ("1" if primitive else "0"), + "#cmakedefine01 BUILD_REORDER": "#define BUILD_REORDER " + ("1" if primitive else "0"), + } + substitutions.update(_COMMON_CONFIG_SUBSTITUTIONS) + return _check_substitution_order(substitutions) + +def onednn_targets(cpu_runtime): + """Declares the oneDNN targets for one CPU runtime. + + Called from bazel/onednn.BUILD (SEQ) and bazel/onednn_threadpool.BUILD + (THREADPOOL), each of which is the build_file of an http_archive in + MODULE.bazel. Declares :onednn (the library to depend on) and + :onednn_autogen, plus the three generated config headers. + + Args: + cpu_runtime: "SEQ" or "THREADPOOL"; sets DNNL_CPU_RUNTIME and + DNNL_CPU_THREADING_RUNTIME and selects which API is enabled. + """ + if cpu_runtime not in ("SEQ", "THREADPOOL"): + fail("onednn_targets: cpu_runtime must be \"SEQ\" or \"THREADPOOL\", got " + repr(cpu_runtime)) + + native.exports_files(["LICENSE"]) + + expand_template( + name = "dnnl_config_h", + out = "include/oneapi/dnnl/dnnl_config.h", + substitutions = _config_substitutions(cpu_runtime), + template = "include/oneapi/dnnl/dnnl_config.h.in", + ) + + expand_template( + name = "dnnl_version_h", + out = "include/oneapi/dnnl/dnnl_version.h", + substitutions = { + "@DNNL_VERSION_MAJOR@": _VERSION_MAJOR, + "@DNNL_VERSION_MINOR@": _VERSION_MINOR, + "@DNNL_VERSION_PATCH@": _VERSION_PATCH, + }, + template = "include/oneapi/dnnl/dnnl_version.h.in", + ) + + expand_template( + name = "dnnl_version_hash_h", + out = "include/oneapi/dnnl/dnnl_version_hash.h", + substitutions = {"@DNNL_VERSION_HASH@": _VERSION_HASH}, + template = "include/oneapi/dnnl/dnnl_version_hash.h.in", + ) + + # The ukernel API needs its -D on the command line too, not just in the + # generated config header. + copts = _BASE_COPTS + (["-DDNNL_EXPERIMENTAL_UKERNEL"] if cpu_runtime == "SEQ" else []) + + native.cc_library( + name = "onednn_autogen", + srcs = native.glob(["src/cpu/x64/gemm/**/*_kern_autogen*.cpp"]), + copts = ["-O1", "-U_FORTIFY_SOURCE"] + copts, + includes = _INCLUDES, + textual_hdrs = native.glob( + _TEXTUAL_HDR_PATTERNS + ["src/graph/backend/dnnl/executables/*.hpp"], + ) + _GENERATED_HDRS, + visibility = ["//visibility:public"], + ) + + native.cc_library( + name = "onednn", + srcs = native.glob( + [ + "src/common/*.cpp", + "src/cpu/*.cpp", + "src/cpu/**/*.cpp", + "src/cpu/jit_utils/**/*.cpp", + "src/cpu/x64/**/*.cpp", + "src/graph/interface/*.cpp", + "src/graph/backend/*.cpp", + "src/graph/backend/dnnl/*.cpp", + "src/graph/backend/dnnl/executables/*.cpp", + "src/graph/backend/fake/*.cpp", + "src/graph/backend/dnnl/passes/*.cpp", + "src/graph/backend/dnnl/patterns/*.cpp", + "src/graph/backend/dnnl/kernels/*.cpp", + "src/graph/utils/*.cpp", + "src/graph/utils/pm/*.cpp", + "third_party/ittnotify/*.c", + ], + exclude = [ + "src/cpu/aarch64/**", + "src/cpu/rv64/**", + "src/cpu/ppc64/**", + "src/cpu/s390x/**", + "src/cpu/x64/gemm/**/*_kern_autogen.cpp", + "src/cpu/sycl/**", + ], + ), + copts = copts, + includes = _INCLUDES, + linkopts = [ + "-lrt", + "-Wl,--allow-multiple-definition", + ], + textual_hdrs = native.glob(_TEXTUAL_HDR_PATTERNS) + _GENERATED_HDRS, + visibility = ["//visibility:public"], + deps = [":onednn_autogen"], + ) diff --git a/bazel/onednn_threadpool.BUILD b/bazel/onednn_threadpool.BUILD new file mode 100644 index 00000000..59734d9b --- /dev/null +++ b/bazel/onednn_threadpool.BUILD @@ -0,0 +1,10 @@ +# oneDNN built for the THREADPOOL CPU runtime, used by the matmul-primitive +# backend (GEMMA_ONEDNN_MATMUL). oneDNN parallelizes by calling back into +# gemma.cpp's thread pool via the adapter in ops/onednn_matmul.h. +# +# Same sources as the SEQ build in bazel/onednn.BUILD; see bazel/onednn.bzl for +# exactly what the runtime changes and why the two cannot share one build. + +load("@gemma//bazel:onednn.bzl", "onednn_targets") + +onednn_targets(cpu_runtime = "THREADPOOL") diff --git a/ops/bench_matmul.cc b/ops/bench_matmul.cc index e9432276..09b65e49 100644 --- a/ops/bench_matmul.cc +++ b/ops/bench_matmul.cc @@ -121,6 +121,9 @@ void BenchMatMul(size_t M, size_t K, size_t N, bool add, MatMulEnv& env) { double keep = 0.0; MMPerKey* per_key; +#if GEMMA_ONEDNN_MATMUL + bool first_onednn_call = true; +#endif // Until enough samples collected *after* autotuning finished: while (times.size() < num_samples) { const double t0 = hwy::platform::Now(); @@ -133,6 +136,14 @@ void BenchMatMul(size_t M, size_t K, size_t N, bool add, MatMulEnv& env) { bool done = per_key->autotune.Best(); #if GEMMA_ONEDNN_BRGEMM done = done || per_key->brgemm_autotune.Best(); +#endif +#if GEMMA_ONEDNN_MATMUL + // oneDNN has no autotune sweep; exclude only the first (JIT + weight + // reorder) call, whose cost is amortized by the primitive/weights caches. + if (per_key->onednn_built) { + done = done || !first_onednn_call; + first_onednn_call = false; + } #endif if (done) times.push_back(elapsed); } diff --git a/ops/matmul-inl.h b/ops/matmul-inl.h index d6c4d8f8..78469145 100644 --- a/ops/matmul-inl.h +++ b/ops/matmul-inl.h @@ -44,6 +44,9 @@ #if GEMMA_ONEDNN_BRGEMM #include "ops/brgemm-inl.h" #endif // GEMMA_ONEDNN_BRGEMM +#if GEMMA_ONEDNN_MATMUL +#include "ops/onednn_matmul-inl.h" +#endif // GEMMA_ONEDNN_MATMUL HWY_BEFORE_NAMESPACE(); namespace gcpp { @@ -1187,6 +1190,22 @@ HWY_NOINLINE MMPerKey* MatMul(const MatPtrT& A, const MatPtrT& B, } // if constexpr BF16/float #endif // GEMMA_ONEDNN_BRGEMM +#if GEMMA_ONEDNN_MATMUL + // OneDNN matmul-primitive path for BF16xBF16 via the threadpool runtime. + // M == 1 was showing worse performance with OneDNN. + if constexpr (IsBF16() && IsBF16()) { + if (M > 1) { + const float scale = A.Scale() * B.Scale(); + if (DoMatMul_OneDnn(A, B, C_rows, M, K, N, scale, add, env.ctx, + cluster_idx)) { + per_key.onednn_built = true; + return &per_key; + } + // oneDNN path failed/declined; fall through to standard matmul. + } + } +#endif // GEMMA_ONEDNN_MATMUL + // (Also auto-tunes, hence outside the timed section to prevent interference.) const StridedViewBF A_view = MMDecompress::MaybeDecompressA(A, per_key.autotune_par_a, env, options); diff --git a/ops/matmul.h b/ops/matmul.h index 130e8615..e249f81f 100644 --- a/ops/matmul.h +++ b/ops/matmul.h @@ -25,6 +25,7 @@ // IWYU pragma: begin_exports #include "ops/brgemm.h" // BRGeMMConfig, GEMMA_ONEDNN_BRGEMM +#include "ops/onednn_matmul.h" // GEMMA_ONEDNN_MATMUL #include "ops/gilbert.h" #include "util/basics.h" #include "util/mat.h" @@ -732,6 +733,11 @@ struct MMPerKey { #if GEMMA_ONEDNN_BRGEMM MMAutoTune brgemm_autotune; #endif // GEMMA_ONEDNN_BRGEMM +#if GEMMA_ONEDNN_MATMUL + // Set true once the oneDNN matmul-primitive path handled this shape, so the + // benchmark can exclude the first (JIT + weight-reorder) call from timings. + bool onednn_built = false; +#endif // GEMMA_ONEDNN_MATMUL }; // Stores state shared across MatMul calls. Non-copyable. `ctx` must outlive diff --git a/ops/onednn_matmul-inl.h b/ops/onednn_matmul-inl.h new file mode 100644 index 00000000..be394963 --- /dev/null +++ b/ops/onednn_matmul-inl.h @@ -0,0 +1,205 @@ +// Copyright 2026 DeepMind Technologies Limited. +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// OneDNN matmul-primitive dispatch for BF16 MatMul via the threadpool runtime. +// See ops/onednn_matmul.h for the engine, adapter, and reordered weight cache. + +#include +#include + +#include +#include + +#include "ops/matmul.h" +#include "ops/onednn_matmul.h" +#include "util/mat.h" +#include "util/threading_context.h" +#include "util/zones.h" +#include "hwy/base.h" + +// Include guard for (potentially) SIMD code. +#if defined(THIRD_PARTY_GEMMA_CPP_ONEDNN_MATMUL_TOGGLE) == \ + defined(HWY_TARGET_TOGGLE) +#ifdef THIRD_PARTY_GEMMA_CPP_ONEDNN_MATMUL_TOGGLE +#undef THIRD_PARTY_GEMMA_CPP_ONEDNN_MATMUL_TOGGLE +#else +#define THIRD_PARTY_GEMMA_CPP_ONEDNN_MATMUL_TOGGLE +#endif + +#include "hwy/highway.h" + +HWY_BEFORE_NAMESPACE(); +namespace gcpp { +namespace HWY_NAMESPACE { + +#if GEMMA_ONEDNN_MATMUL + +// Thread-local byte buffer for oneDNN's user-managed scratchpad. +inline std::vector& GetOneDnnScratchpad() { + static thread_local std::vector scratch; + return scratch; +} + +// oneDNN data type for the C element type (only f32/bf16 outputs are used). +template +constexpr dnnl::memory::data_type OneDnnDstType() { + if constexpr (hwy::IsSame()) { + return dnnl::memory::data_type::f32; + } else { + static_assert(IsBF16(), "OneDnn matmul path expects f32 or bf16 C."); + return dnnl::memory::data_type::bf16; + } +} + +// Computes C[M,N] = scale * (A[M,K] * B[N,K]^T) (+ add) using the dnnl::matmul +// primitive, parallelized via the threadpool adapter. Scale and the optional +// per-column bias are fused into the primitive, which writes final values +// Returns false (with no writes to C) on any failure and +// allows the caller to fall back to the stock path. +template +static HWY_NOINLINE bool DoMatMul_OneDnn(const MatPtrT& A, + const MatPtrT& B, RowPtrs C, + size_t M, size_t K, size_t N, + float scale, + const float* add, + ThreadingContext& ctx, + size_t cluster_idx) { + static_assert(IsBF16() && IsBF16(), + "OneDnn matmul path expects BF16 A and B."); + try { + using dt = dnnl::memory::data_type; + using dims = dnnl::memory::dims; + dnnl::engine& engine = OneDnnEngine(); + + const hwy::pool::Caller caller = + ctx.pool_callers.Get(Callers::kOneDnnMatMul); + HwyThreadPoolAdapter adapter(ctx, cluster_idx, caller); + // Must precede primitive create/exec; see comment on the helper. + SetOneDnnMaxConcurrency(adapter.get_num_threads()); + dnnl::stream stream = + dnnl::threadpool_interop::make_stream(engine, &adapter); + + const int64_t Mi = static_cast(M); + const int64_t Ki = static_cast(K); + const int64_t Ni = static_cast(N); + const int64_t lda = static_cast(A.Stride()); + const int64_t ldb = static_cast(B.Stride()); + + // Build oneDNN primitive once + // JIT-compiled kernel will be used from the 2nd call onward. + + // A: [M,K] bf16, actual leading dim lda (handles non-packed A directly). + const dnnl::memory::desc src_md({Mi, Ki}, dt::bf16, dims{lda, 1}); + + // C: [M,N], written directly by the primitive. RowPtrs permits arbitrary + // per-row pointers, but oneDNN can only address a single leading dim, so + // verify the rows are regularly strided and bail to the stock path if not. + TC* const C0 = C.Row(0); + const ptrdiff_t ldc = (M > 1) ? (C.Row(1) - C0) : static_cast(N); + if (ldc < static_cast(N)) return false; + for (size_t m = 1; m < M; ++m) { + if (C.Row(m) != C0 + ldc * static_cast(m)) return false; + } + const dnnl::memory::desc dst_md({Mi, Ni}, OneDnnDstType(), + dims{static_cast(ldc), 1}); + + // B: logical [K,N] bf16, format_tag::any so oneDNN picks its best layout. + const dnnl::memory::desc wei_any( + {Ki, Ni}, dt::bf16, dnnl::memory::format_tag::any); + // Optional per-column bias (the `add`), broadcast over rows. Empty desc = + // no bias. + const dnnl::memory::desc bias_md = + add ? dnnl::memory::desc({1, Ni}, dt::f32, dims{Ni, 1}) + : dnnl::memory::desc(); + + dnnl::primitive_attr attr; + // Fuse the scalar product scale + attr.set_scales_mask(DNNL_ARG_WEIGHTS, 0); + // switch from default library-managed scratchpad to user-managed scratchpad + attr.set_scratchpad_mode(dnnl::scratchpad_mode::user); + const dnnl::matmul::primitive_desc pd(engine, src_md, wei_any, bias_md, + dst_md, attr); + const dnnl::memory::desc weights_md = pd.weights_desc(); + dnnl::matmul prim(pd); + + // Weights cache: B reordered into the kernel's layout, keyed on the B + // pointer alone. Reorder happens once per distinct B. The key carries no + // shape, but oneDNN's layout can be M-dependent and running against a + // wrongly-packed B is silently wrong, so we reuse an entry only if its + // actual layout equals the one this primitive wants. Large shapes like + // those seen in gemma are M-independent. + const uintptr_t B_ptr = reinterpret_cast(B.Row(0)); + const OneDnnWeightsKey w_key{B_ptr}; + auto& w_cache = GetOneDnnWeightsCache(); + auto w_it = w_cache.find(w_key); + const bool needs_reorder = + w_it == w_cache.end() || w_it->second.packed.get_desc() != weights_md; + if (needs_reorder) { + const dnnl::memory::desc user_wei_md({Ki, Ni}, dt::bf16, dims{1, ldb}); + dnnl::memory user_wei(user_wei_md, engine, + const_cast(B.Row(0))); + OneDnnWeightsEntry we; + we.packed = dnnl::memory(weights_md, engine); + dnnl::reorder(user_wei, we.packed) + .execute(stream, user_wei, we.packed); + stream.wait(); + if (w_it == w_cache.end()) { + w_it = w_cache.emplace(w_key, std::move(we)).first; + } else { + w_it->second = std::move(we); + } + } + OneDnnWeightsEntry& we = w_it->second; + + dnnl::memory src_mem(src_md, engine, const_cast(A.Row(0))); + dnnl::memory dst_mem(dst_md, engine, C0); + dnnl::memory scale_mem({{1}, dt::f32, dnnl::memory::format_tag::x}, engine, + &scale); + + // User-managed scratchpad + const dnnl::memory::desc scratchpad_md = pd.scratchpad_desc(); + const size_t scratchpad_size = scratchpad_md.get_size(); + std::vector& sp_buf = GetOneDnnScratchpad(); + if (sp_buf.size() < scratchpad_size) { + sp_buf.resize(scratchpad_size ? scratchpad_size : 1); + } + dnnl::memory scratchpad_mem(scratchpad_md, engine, sp_buf.data()); + + std::unordered_map args{ + {DNNL_ARG_SRC, src_mem}, + {DNNL_ARG_WEIGHTS, we.packed}, + {DNNL_ARG_DST, dst_mem}, + {DNNL_ARG_ATTR_SCALES | DNNL_ARG_WEIGHTS, scale_mem}, + {DNNL_ARG_SCRATCHPAD, scratchpad_mem}}; + if (add) { + args.emplace(DNNL_ARG_BIAS, + dnnl::memory(bias_md, engine, const_cast(add))); + } + prim.execute(stream, args); + stream.wait(); + return true; + } catch (...) { + return false; + } +} + +#endif // GEMMA_ONEDNN_MATMUL + +// NOLINTNEXTLINE(google-readability-namespace-comments) +} // namespace HWY_NAMESPACE +} // namespace gcpp +HWY_AFTER_NAMESPACE(); + +#endif // NOLINT (include guard) diff --git a/ops/onednn_matmul.h b/ops/onednn_matmul.h new file mode 100644 index 00000000..36eda862 --- /dev/null +++ b/ops/onednn_matmul.h @@ -0,0 +1,154 @@ +// Copyright 2026 DeepMind Technologies Limited. +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// OneDNN matmul-primitive integration for MatMul via the threadpool runtime. +// Enabled at compile time via GEMMA_ONEDNN_MATMUL=1 (Bazel: --define +// gemma_onednn_matmul=1). +// +// This is a different integration than ops/brgemm.h (GEMMA_ONEDNN_BRGEMM): +// - BRGeMM uses oneDNN's low-level ukernel API with oneDNN built SEQ; gemma.cpp +// drives all parallelism. +// - Here we use the high-level dnnl::matmul primitive with oneDNN built for the +// THREADPOOL runtime, so oneDNN picks the kernel and parallelizes *internally* +// by calling back into gemma.cpp's hwy::ThreadPool through the adapter below. + +#ifndef THIRD_PARTY_GEMMA_CPP_OPS_ONEDNN_MATMUL_H_ +#define THIRD_PARTY_GEMMA_CPP_OPS_ONEDNN_MATMUL_H_ + +#include +#include + +// opt-in +#ifndef GEMMA_ONEDNN_MATMUL +#define GEMMA_ONEDNN_MATMUL 0 +#endif // GEMMA_ONEDNN_MATMUL + +#if GEMMA_ONEDNN_MATMUL && GEMMA_ONEDNN_BRGEMM +#error \ + "GEMMA_ONEDNN_MATMUL and GEMMA_ONEDNN_BRGEMM are mutually exclusive: " \ + "oneDNN's CPU runtime (SEQ for BRGeMM vs THREADPOOL for the matmul " \ + "primitive) is a whole-library compile-time choice. Enable only one." +#endif + +#if GEMMA_ONEDNN_MATMUL + +#include +#include + +#include "oneapi/dnnl/dnnl.hpp" +#include "oneapi/dnnl/dnnl_threadpool.hpp" +#include "oneapi/dnnl/dnnl_threadpool_iface.hpp" +#include "util/threading_context.h" +#include "util/zones.h" +#include "hwy/base.h" +#include "hwy/contrib/thread_pool/thread_pool.h" + +namespace gcpp { + +inline dnnl::engine& OneDnnEngine() { + static dnnl::engine engine(dnnl::engine::kind::cpu, 0); + return engine; +} + +inline void SetOneDnnMaxConcurrency(int num_threads) { + dnnl_threadpool_interop_set_max_concurrency(num_threads); +} + +// Adapts one gemma.cpp cluster ThreadPool to oneDNN's threadpool interface. +class HwyThreadPoolAdapter : public dnnl::threadpool_interop::threadpool_iface { + public: + HwyThreadPoolAdapter(ThreadingContext& ctx, size_t cluster_idx, + hwy::pool::Caller caller) + : pool_(ctx.pools.Cluster(cluster_idx)), caller_(caller) {} + + int get_num_threads() const override { + return static_cast(pool_.NumWorkers()); + } + + // True when the calling thread is already executing inside our parallel_for. + // oneDNN uses this to avoid nesting parallel regions + bool get_in_parallel() const override { return InParallel(); } + + // Synchronous flags (0): parallel_for must block. + uint64_t get_flags() const override { return 0; } + + void parallel_for(int n, + const std::function& fn) override { + if (n <= 0) return; + // Run inline if there is nothing to fan out to, or if we are already inside + // a parallel region on this pinned pool + if (n == 1 || pool_.NumWorkers() <= 1 || InParallel()) { + for (int i = 0; i < n; ++i) fn(i, n); + return; + } + pool_.Run(0, static_cast(n), caller_, + [&](uint64_t task, size_t /*worker*/) { + InParallelGuard guard; + fn(static_cast(task), n); + }); + } + + void wait() override {} + + private: + static bool& InParallel() { + static thread_local bool in_parallel = false; + return in_parallel; + } + + struct InParallelGuard { + InParallelGuard() { InParallel() = true; } + ~InParallelGuard() { InParallel() = false; } + }; + + hwy::ThreadPool& pool_; + hwy::pool::Caller caller_; +}; + +// --------------------------------------------------------------------------- +// Reordered-weights cache keyed by B pointer. +// The caller separately checks the stored layout +// against the one the primitive wants. +struct OneDnnWeightsKey { + uintptr_t B_ptr; + bool operator==(const OneDnnWeightsKey& o) const { + return B_ptr == o.B_ptr; + } +}; + +struct OneDnnWeightsKeyHash { + size_t operator()(const OneDnnWeightsKey& k) const { + size_t h = 14695981039346656037ULL; + h = (h ^ k.B_ptr) * 1099511628211ULL; + return h; + } +}; + +struct OneDnnWeightsEntry { + dnnl::memory packed; +}; + +inline auto& GetOneDnnWeightsCache() { + static std::unordered_map + cache; + return cache; +} + +} // namespace gcpp + +#endif // GEMMA_ONEDNN_MATMUL + +#endif // THIRD_PARTY_GEMMA_CPP_OPS_ONEDNN_MATMUL_H_ diff --git a/util/zones.cc b/util/zones.cc index 388f5fd6..2c63f61c 100644 --- a/util/zones.cc +++ b/util/zones.cc @@ -189,6 +189,8 @@ const char* CallerName(Callers caller) { return "MoE.ComputeAllExpertOutputs"; case Callers::kMoEWeightedSumOfExperts: return "MoE.WeightedSumOfExperts"; + case Callers::kOneDnnMatMul: + return "OneDnnMatMul"; case Callers::kOpsAddFromBatched: return "Ops.AddFromBatched"; case Callers::kOpsGroupedRMSNormBatched: diff --git a/util/zones.h b/util/zones.h index 9252e0d1..01adb005 100644 --- a/util/zones.h +++ b/util/zones.h @@ -107,6 +107,7 @@ enum class Callers { // Keep sorted kMoEChooseExperts, kMoEComputeAllExpertOutputs, kMoEWeightedSumOfExperts, + kOneDnnMatMul, kOpsAddFromBatched, kOpsGroupedRMSNormBatched, kOpsGroupedRMSNormInplaceBatched,