1
0
mirror of https://github.com/opencv/opencv.git synced 2026-07-31 08:13:04 +04:00
Files
opencv/3rdparty/mlas/patches/0002-platform-gemm-only.diff
Abhishek Gola bdf348c13a Merge pull request #28934 from abhishek-gola:mlas_gemm
Added MLAS third party module and integrated into GeMM path #28934

### Pull Request Readiness Checklist

See details at https://github.com/opencv/opencv/wiki/How_to_contribute#making-a-good-pull-request

- [x] I agree to contribute to the project under Apache 2 License.
- [x] To the best of my knowledge, the proposed patch is not based on a code under GPL or another license that is incompatible with OpenCV
- [x] The PR is proposed to the proper branch
- [x] There is a reference to the original bug report and related work
- [x] There is accuracy test, performance test and test data in opencv_extra repository, if applicable
      Patch to opencv_extra has the same branch name.
- [x] The feature is well documented and sample code can be built with the project CMake
2026-05-22 20:22:15 +03:00

175 lines
6.9 KiB
Diff

Patch against upstream onnxruntime/core/mlas/lib/platform.cpp
Base commit: 62f742f1aa0c3102745ed35e3d869eaee845b9ac (ORT v1.26.0)
Gate non-SGEMM dispatch behind MLAS_GEMM_ONLY so the SGEMM-only
subset can build without the rest of the MLAS sources (quantized
GEMM, conv, FP16 SoftMax, etc.). The original ctor is preserved
verbatim in the #else branch for clean re-vendoring.
Also gate the top-of-file erf_neon_fp16.h / gelu_neon_fp16.h
includes on !MLAS_GEMM_ONLY — they transitively pull in
fp16_common.h / softmax_kernel_neon.h, which we don't vendor.
The MLAS_GEMM_ONLY ctor additionally assigns ReduceMaximumF32Kernel
and ComputeSumExpF32Kernel to the portable compute.cpp fallbacks so
MlasFlashAttention works without per-arch softmax kernels.
--- a/3rdparty/mlas/lib/platform.cpp
+++ b/3rdparty/mlas/lib/platform.cpp
@@ -19,7 +19,7 @@
#ifdef MLAS_USE_SVE
#include "sve/mlasi_sve.h"
#endif
-#if defined(MLAS_NEON_INTRINSICS) && defined(MLAS_F16VEC_INTRINSICS_SUPPORTED)
+#if defined(MLAS_NEON_INTRINSICS) && defined(MLAS_F16VEC_INTRINSICS_SUPPORTED) && !defined(MLAS_GEMM_ONLY)
#include "erf_neon_fp16.h"
#include "gelu_neon_fp16.h"
#endif
@@ -288,6 +288,138 @@
};
#endif
+
+// =============================================================================
+// SGEMM-only constructor (vendor-local patch).
+//
+// When MLAS_GEMM_ONLY is defined, replace the original platform-init ctor
+// with a stripped-down version that only assigns the four (-ish) dispatch
+// fields read by sgemm.cpp:
+// - GemmFloatKernel
+// - KernelM1Routine (x86_64 only)
+// - KernelM1TransposeBRoutine (x86_64 only)
+// - TransposePackB16x4Routine (x86_64 / loongarch only)
+// Plus, on the SBGemm aarch64+linux path, the SBGemm batch overrides — but
+// those are nullptr-default and we don't enable SBGemm here.
+//
+// Also initializes the two softmax kernel pointers consumed by
+// flashattn.cpp (ReduceMaximumF32Kernel, ComputeSumExpF32Kernel) to the
+// portable fallbacks provided by compute.cpp. No SIMD-asm softmax kernels
+// are vendored — the flash-attention path uses the portable C++ rowmax /
+// sum-exp implementations.
+//
+// Every other dispatch field stays at its in-class default (most are
+// `= nullptr`). Calling any non-SGEMM / non-FlashAttention MLAS API in this
+// build is undefined.
+//
+// The original full ORT ctor is preserved unchanged below the #else for
+// future re-vendoring — drop MLAS_GEMM_ONLY to use it.
+// =============================================================================
+#ifdef MLAS_GEMM_ONLY
+MLAS_PLATFORM::MLAS_PLATFORM(void)
+{
+ // Portable softmax kernels (compute.cpp). flashattn.cpp dereferences these
+ // function pointers on the AMD64 / LARCH64 path; compute.cpp's
+ // MlasComputeSoftmax does the same on AMD64 / LARCH64 / SVE / RISCV64.
+ // Other paths call the symbols directly. Gates mirror the MLAS_PLATFORM
+ // member visibility in mlasi.h so we initialize the field wherever it
+ // exists — leaving it null would crash any future code that reads it via
+ // the struct on those targets.
+#if defined(MLAS_TARGET_LARCH64) || defined(MLAS_USE_SVE) || \
+ defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_RISCV64)
+ this->ReduceMaximumF32Kernel = MlasReduceMaximumF32Kernel;
+#endif
+#if defined(MLAS_USE_SVE) || defined(MLAS_TARGET_AMD64) || defined(MLAS_TARGET_RISCV64)
+ this->ComputeSumExpF32Kernel = MlasComputeSumExpF32Kernel;
+#endif
+
+ // The PreferredBufferAlignment field only exists on AMD64 (see
+ // MLAS_PLATFORM in mlasi.h). On other targets MlasGetPreferredBufferAlignment()
+ // returns MLAS_DEFAULT_PREFERRED_BUFFER_ALIGNMENT directly without
+ // consulting the struct.
+#if defined(MLAS_TARGET_AMD64)
+ this->PreferredBufferAlignment = MLAS_DEFAULT_PREFERRED_BUFFER_ALIGNMENT;
+#endif
+
+#if defined(MLAS_TARGET_AMD64_IX86)
+ // SSE2 baseline (every x86 since 2003).
+ this->GemmFloatKernel = MlasGemmFloatKernelSse;
+#if defined(MLAS_TARGET_AMD64)
+ this->TransposePackB16x4Routine = MlasSgemmTransposePackB16x4Sse;
+#endif
+
+ unsigned Cpuid1[4];
+#if defined(_WIN32)
+ __cpuid((int*)Cpuid1, 1);
+#else
+ __cpuid(1, Cpuid1[0], Cpuid1[1], Cpuid1[2], Cpuid1[3]);
+#endif
+ // AVX + OSXSAVE bits (matches the original ctor's checks).
+ if ((Cpuid1[2] & 0x18000000) == 0x18000000) {
+ uint64_t xcr0 = MlasReadExtendedControlRegister(_XCR_XFEATURE_ENABLED_MASK);
+ if ((xcr0 & 0x6) == 0x6) {
+ this->GemmFloatKernel = MlasGemmFloatKernelAvx;
+#if defined(MLAS_TARGET_AMD64)
+ this->KernelM1Routine = MlasSgemmKernelM1Avx;
+ this->KernelM1TransposeBRoutine = MlasSgemmKernelM1TransposeBAvx;
+ this->TransposePackB16x4Routine = MlasSgemmTransposePackB16x4Avx;
+#endif
+ unsigned Cpuid7[4];
+#if defined(_WIN32)
+ __cpuidex((int*)Cpuid7, 7, 0);
+#else
+ __cpuid_count(7, 0, Cpuid7[0], Cpuid7[1], Cpuid7[2], Cpuid7[3]);
+#endif
+ // AVX2 + FMA3.
+ if (((Cpuid1[2] & 0x1000) != 0) && ((Cpuid7[1] & 0x20) != 0)) {
+ this->GemmFloatKernel = MlasGemmFloatKernelFma3;
+ // AVX-512F + ZMM-state save.
+ if (((Cpuid7[1] & 0x10000) != 0) && ((xcr0 & 0xE0) == 0xE0)) {
+ this->GemmFloatKernel = MlasGemmFloatKernelAvx512F;
+ }
+ }
+ }
+ }
+#endif // MLAS_TARGET_AMD64_IX86
+
+#if defined(MLAS_TARGET_POWER)
+ // Default to the base SgemmKernelPower; the POWER10 detection branch in
+ // the original ctor is omitted because the POWER10 SgemmKernel symbol
+ // (MlasSgemmKernelPOWER10) is only present when -mcpu=power10 was
+ // detectable at configure time. CMake conditionally compiles it; the
+ // base kernel is always available.
+ this->GemmFloatKernel = MlasSgemmKernel;
+#endif
+
+#if defined(MLAS_TARGET_S390X)
+ this->GemmFloatKernel = MlasSgemmKernel;
+#endif
+
+#if defined(MLAS_TARGET_RISCV64)
+ this->GemmFloatKernel = nullptr;
+#if defined(MLAS_USE_RVV)
+ bool has_rvv = true;
+#if defined(__linux__)
+ has_rvv = (getauxval(AT_HWCAP) & COMPAT_HWCAP_ISA_V) != 0;
+#endif
+ if (has_rvv) {
+ this->GemmFloatKernel = MlasGemmFloatKernelRvv;
+ }
+#endif // MLAS_USE_RVV
+#endif // MLAS_TARGET_RISCV64
+
+#if defined(MLAS_TARGET_LARCH64)
+ // No fine-grained LSX/LASX detection here — pick LASX (256-bit) since
+ // the LoongArch64 spec requires it; LSX (128-bit) is the fallback.
+ this->GemmFloatKernel = MlasGemmFloatKernelLasx;
+ this->TransposePackB16x4Routine = MlasSgemmTransposePackB16x4Lasx;
+#endif
+
+ // ARM64 and WASM intentionally do nothing here — sgemm.cpp's #else branch
+ // calls MlasSgemmKernelZero / MlasSgemmKernelAdd directly without going
+ // through GetMlasPlatform().GemmFloatKernel.
+}
+#else // !MLAS_GEMM_ONLY
MLAS_PLATFORM::MLAS_PLATFORM(
void
)
@@ -909,6 +1041,7 @@
#endif // MLAS_TARGET_LARCH64
}
+#endif // MLAS_GEMM_ONLY
size_t
MLASCALL