#include "core/mlas/inc/mlas.h"
#include "core/mlas/inc/mlas_qnbit.h"
#include <cstddef>
#include <cstring>
#include <functional>
#include <new>
extern "C" {
typedef void (*mlas_task_fn)(void* task_ctx, std::ptrdiff_t tid);
typedef void (*mlas_parallel_for_fn)(
void* rust_ctx,
std::ptrdiff_t iterations,
mlas_task_fn task,
void* task_ctx);
typedef int (*mlas_max_threads_fn)(void* rust_ctx);
}
namespace {
mlas_parallel_for_fn g_parallel_for = nullptr;
mlas_max_threads_fn g_max_threads = nullptr;
void* g_rust_ctx = nullptr;
}
extern "C" void mlas_set_threading(
mlas_parallel_for_fn parallel_for,
mlas_max_threads_fn max_threads,
void* rust_ctx)
{
g_parallel_for = parallel_for;
g_max_threads = max_threads;
g_rust_ctx = rust_ctx;
}
extern "C" int MlasStandaloneMaxThreads()
{
if (g_max_threads != nullptr) {
int n = g_max_threads(g_rust_ctx);
return n > 0 ? n : 1;
}
return 1;
}
namespace {
void mlas_std_function_trampoline(void* task_ctx, std::ptrdiff_t tid)
{
(*static_cast<const std::function<void(std::ptrdiff_t)>*>(task_ctx))(tid);
}
}
extern "C" void MlasStandaloneParallelFor(std::ptrdiff_t iterations, void* work)
{
const auto& fn = *static_cast<const std::function<void(std::ptrdiff_t)>*>(work);
if (g_parallel_for != nullptr && iterations > 1) {
g_parallel_for(
g_rust_ctx,
iterations,
&mlas_std_function_trampoline,
const_cast<void*>(static_cast<const void*>(&fn)));
} else {
for (std::ptrdiff_t tid = 0; tid < iterations; ++tid) {
fn(tid);
}
}
}
extern "C" void mlas_sgemm(
int transA, int transB,
size_t M,
size_t N,
size_t K,
float alpha,
const float* A,
size_t lda,
const float* B,
size_t ldb,
float beta,
float* C,
size_t ldc)
{
MLAS_SGEMM_DATA_PARAMS data;
data.A = A;
data.lda = lda;
data.B = B;
data.ldb = ldb;
data.C = C;
data.ldc = ldc;
data.alpha = alpha;
data.beta = beta;
data.BIsPacked = false;
MlasGemmBatch(
transA ? CblasTrans : CblasNoTrans,
transB ? CblasTrans : CblasNoTrans,
M, N, K,
&data, 1,
nullptr,
nullptr);
}
extern "C" size_t mlas_sgemm_pack_b_size(int transA, int transB, size_t N, size_t K)
{
return MlasGemmPackBSize(
transA ? CblasTrans : CblasNoTrans,
transB ? CblasTrans : CblasNoTrans,
N, K, nullptr);
}
extern "C" void mlas_sgemm_pack_b(
int transA, int transB, size_t N, size_t K,
const float* B, size_t ldb, void* packed_b)
{
MlasGemmPackB(
transA ? CblasTrans : CblasNoTrans,
transB ? CblasTrans : CblasNoTrans,
N, K, B, ldb, packed_b, nullptr);
}
extern "C" int mlas_qnbit_gemm_available(size_t bits, size_t blk_len, int comp_type)
{
return MlasIsQNBitGemmAvailable(
bits, blk_len, static_cast<MLAS_QNBIT_GEMM_COMPUTE_TYPE>(comp_type))
? 1
: 0;
}
extern "C" size_t mlas_qnbit_gemm_pack_b_size(
size_t n, size_t k, size_t bits, size_t blk_len, int has_zp, int comp_type)
{
return MlasQNBitGemmPackQuantBDataSize(
n, k, bits, blk_len, has_zp != 0,
static_cast<MLAS_QNBIT_GEMM_COMPUTE_TYPE>(comp_type),
nullptr);
}
extern "C" void mlas_qnbit_gemm_pack_b(
size_t n,
size_t k,
size_t bits,
size_t blk_len,
int comp_type,
const void* quant_b_data,
void* packed_b,
const void* quant_b_scale,
int has_zp,
const void* quant_b_zero_point)
{
MlasQNBitGemmPackQuantBData(
n, k, bits, blk_len,
static_cast<MLAS_QNBIT_GEMM_COMPUTE_TYPE>(comp_type),
quant_b_data,
packed_b,
quant_b_scale,
has_zp != 0,
quant_b_zero_point,
nullptr,
nullptr);
}
extern "C" size_t mlas_qnbit_gemm_workspace_size(
size_t m, size_t n, size_t k, size_t bits, size_t blk_len, int has_zp, int comp_type)
{
return MlasQNBitGemmBatchWorkspaceSize(
m, n, k, 1, bits, blk_len, has_zp != 0,
static_cast<MLAS_QNBIT_GEMM_COMPUTE_TYPE>(comp_type),
nullptr);
}
extern "C" void mlas_qnbit_gemm(
size_t m,
size_t n,
size_t k,
size_t bits,
size_t blk_len,
int comp_type,
const float* a,
size_t lda,
const void* packed_b,
const float* quant_b_scale,
int has_zp,
const void* quant_b_zero_point,
const float* bias,
float* c,
size_t ldc,
void* workspace,
int multithread)
{
MLAS_QNBIT_GEMM_DATA_PARAMS<float> params;
params.A = a;
params.lda = lda;
params.Bias = bias;
params.C = c;
params.ldc = ldc;
const auto ct = static_cast<MLAS_QNBIT_GEMM_COMPUTE_TYPE>(comp_type);
if (ct == SQNBIT_CompInt8) {
params.QuantBDataWorkspace = packed_b;
} else {
params.PackedQuantBData = static_cast<const std::byte*>(packed_b);
params.QuantBScale = quant_b_scale;
params.QuantBZeroPoint = has_zp != 0 ? quant_b_zero_point : nullptr;
}
MLAS_THREADPOOL* thread_pool =
multithread != 0 ? reinterpret_cast<MLAS_THREADPOOL*>(1) : nullptr;
MlasQNBitGemmBatch<float>(
m, n, k, 1, bits, blk_len, ct, ¶ms, workspace, thread_pool,
nullptr);
}
extern "C" void mlas_sgemm_packed(
int transA,
int transB,
size_t M,
size_t N,
size_t K,
float alpha,
const float* A,
size_t lda,
const void* packed_b,
float beta,
float* C,
size_t ldc)
{
MLAS_SGEMM_DATA_PARAMS data;
data.A = A;
data.lda = lda;
data.B = reinterpret_cast<const float*>(packed_b);
data.ldb = 0;
data.C = C;
data.ldc = ldc;
data.alpha = alpha;
data.beta = beta;
data.BIsPacked = true;
MlasGemmBatch(
transA ? CblasTrans : CblasNoTrans,
transB ? CblasTrans : CblasNoTrans,
M, N, K,
&data, 1,
nullptr,
nullptr);
}
extern "C" void mlas_compute_logistic(
const float* input,
float* output,
size_t n)
{
MlasComputeLogistic(input, output, n);
}
extern "C" void mlas_compute_silu(
const float* input,
float* output,
size_t n)
{
MlasComputeSilu(input, output, n);
}
extern "C" void mlas_eltwise_add(
const float* left,
const float* right,
float* output,
size_t n)
{
MlasEltwiseAdd<float>(left, right, output, n);
}
extern "C" void mlas_compute_activation(
int kind,
float minimum,
float maximum,
const float* input,
float* output,
size_t n)
{
if (input != output) {
std::memcpy(output, input, n * sizeof(float));
}
MLAS_ACTIVATION activation{};
activation.ActivationKind = static_cast<MLAS_ACTIVATION_KIND>(kind);
activation.Parameters.Clip.minimum = minimum;
activation.Parameters.Clip.maximum = maximum;
MlasActivation(&activation, output, nullptr, 1, n, n);
}
namespace {
struct mlas_conv_plan {
MLAS_ACTIVATION activation{};
MLAS_CONV_PARAMETERS parameters{};
};
}
extern "C" void* mlas_conv_prepare(
size_t dimensions,
size_t batch_count,
size_t group_count,
size_t input_channels_per_group,
const int64_t* input_shape,
const int64_t* kernel_shape,
const int64_t* dilation_shape,
const int64_t* padding,
const int64_t* stride_shape,
const int64_t* output_shape,
size_t filter_count_per_group,
size_t* working_buffer_elements)
{
mlas_conv_plan* plan = nullptr;
try {
plan = new mlas_conv_plan();
plan->activation.ActivationKind = MlasIdentityActivation;
MLAS_THREADPOOL* thread_pool = reinterpret_cast<MLAS_THREADPOOL*>(1);
MlasConvPrepare(
&plan->parameters,
dimensions,
batch_count,
group_count,
input_channels_per_group,
input_shape,
kernel_shape,
dilation_shape,
padding,
stride_shape,
output_shape,
filter_count_per_group,
&plan->activation,
working_buffer_elements,
false,
0.0f,
thread_pool);
return plan;
} catch (...) {
delete plan;
return nullptr;
}
}
extern "C" void mlas_conv_run(
const void* opaque_plan,
const float* input,
const float* filter,
const float* bias,
float* working_buffer,
float* output)
{
const auto* plan = static_cast<const mlas_conv_plan*>(opaque_plan);
MLAS_THREADPOOL* thread_pool = reinterpret_cast<MLAS_THREADPOOL*>(1);
MlasConv(
&plan->parameters,
input,
filter,
bias,
working_buffer,
output,
thread_pool);
}
extern "C" void mlas_conv_plan_destroy(void* opaque_plan)
{
delete static_cast<mlas_conv_plan*>(opaque_plan);
}
extern "C" size_t mlas_nchwc_block_size()
{
return MlasNchwcGetBlockSize();
}
extern "C" void mlas_nchwc_reorder_input_nchw(
const float* source,
float* dest,
size_t channels,
size_t input_size)
{
MlasReorderInputNchw(source, dest, channels, input_size);
}
extern "C" void mlas_nchwc_reorder_output_nchw(
const int64_t* output_shape,
const float* source,
float* dest)
{
MlasReorderOutputNchw(output_shape, source, dest, nullptr);
}
extern "C" void mlas_nchwc_reorder_filter_bibo(
const int64_t* filter_shape,
const float* source,
float* dest)
{
MlasReorderFilterOIHWBiBo(filter_shape, source, dest);
}
extern "C" void mlas_nchwc_reorder_filter_bo(
const int64_t* filter_shape,
const float* source,
float* dest)
{
MlasReorderFilterOIHWBo(filter_shape, source, dest);
}
extern "C" void mlas_nchwc_conv(
const int64_t* input_shape,
const int64_t* kernel_shape,
const int64_t* dilation_shape,
const int64_t* padding,
const int64_t* stride_shape,
const int64_t* output_shape,
size_t group_count,
const float* input,
const float* filter,
const float* bias,
float* output,
int activation_kind,
float activation_value0,
float activation_value1,
int zero_mode)
{
MLAS_ACTIVATION activation{};
activation.ActivationKind = static_cast<MLAS_ACTIVATION_KIND>(activation_kind);
activation.Parameters.Values[0] = activation_value0;
activation.Parameters.Values[1] = activation_value1;
MLAS_THREADPOOL* thread_pool = reinterpret_cast<MLAS_THREADPOOL*>(1);
MlasNchwcConv(
input_shape,
kernel_shape,
dilation_shape,
padding,
stride_shape,
output_shape,
group_count,
input,
filter,
bias,
output,
&activation,
zero_mode != 0,
thread_pool,
nullptr,
false);
}
extern "C" void mlas_pool(
int kind,
size_t dimensions,
const int64_t* input_shape,
const int64_t* kernel_shape,
const int64_t* padding,
const int64_t* stride_shape,
const int64_t* output_shape,
const float* input,
float* output)
{
MLAS_THREADPOOL* thread_pool = reinterpret_cast<MLAS_THREADPOOL*>(1);
MlasPool(
static_cast<MLAS_POOLING_KIND>(kind),
dimensions,
input_shape,
kernel_shape,
padding,
stride_shape,
output_shape,
input,
output,
thread_pool);
}
extern "C" void mlas_nchwc_pool(
int kind,
const int64_t* input_shape,
const int64_t* kernel_shape,
const int64_t* dilation_shape,
const int64_t* padding,
const int64_t* stride_shape,
const int64_t* output_shape,
const float* input,
float* output)
{
MLAS_THREADPOOL* thread_pool = reinterpret_cast<MLAS_THREADPOOL*>(1);
MlasNchwcPool(
static_cast<MLAS_POOLING_KIND>(kind),
input_shape,
kernel_shape,
dilation_shape,
padding,
stride_shape,
output_shape,
input,
output,
thread_pool);
}