/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
src/backends/cpu/matmul.hpp
53 строки
1 KB
kolkir
revert compiler compatibility to gcc12
03 мар 2025, 23:52
03 мар 2025, 23:52
26a9274
Код
Авторство
О чём код?
#pragma once #include <cblas.h> #include <adept/print.hpp> #include <adept/shape.hpp> namespace adept { template <typename T> void gemm(CBLAS_TRANSPOSE trans_a, CBLAS_TRANSPOSE trans_b, index_t m, index_t n, index_t k, T alpha, const T* a, index_t lda, const T* b, index_t ldb, T beta, T* c, index_t ldc, CBLAS_ORDER order = CblasRowMajor) { if constexpr (std::is_same_v<T, float32_t>) { cblas_sgemm(order, trans_a, trans_b, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc); } else if constexpr (std::is_same_v<T, float64_t>) { cblas_dgemm(order, trans_a, trans_b, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc); } else { // We don't use static_assert to make code compilable THROW_ERROR("gemm doesn't support {}", to_dtype<T>()); } } template <typename T> void matmul(const Shape& a_shape, const T* a, const Shape& b_shape, const T* b, T* c) { auto m = a_shape.dim(0); auto k = a_shape.dim(1); auto n = b_shape.dim(1); if (k != b_shape.dim(0)) { THROW_ERROR("matmul got incompatible matrix shapes ", a_shape, " and ", b_shape); } auto lda = k; auto ldb = n; auto ldc = n; gemm(CblasNoTrans, CblasNoTrans, m, n, k, static_cast<T>(1), a, lda, b, ldb, static_cast<T>(1), c, ldc); } } // namespace adept