File size: 1,240 Bytes
57c2394 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 | #pragma once
#include <cstdint>
namespace orbitquant::cpu {
enum class ScalarKind : std::uint8_t {
Float32,
Float16,
BFloat16,
};
struct PackedMatmulArgs {
void *out;
void const *x;
std::uint8_t const *packed_weight_indices;
float const *row_norms;
float const *centroids;
float const *bias;
bool has_bias;
ScalarKind scalar_kind;
std::int64_t rows;
std::int64_t out_features;
std::int64_t in_features;
std::int64_t bits;
};
using PackedMatmulRangeFn = void (*)(
PackedMatmulArgs const &args,
std::int64_t out_start,
std::int64_t out_end);
void packed_matmul_scalar_range(
PackedMatmulArgs const &args,
std::int64_t out_start,
std::int64_t out_end);
bool packed_matmul_neon_available();
void packed_matmul_neon_range(
PackedMatmulArgs const &args,
std::int64_t out_start,
std::int64_t out_end);
bool packed_matmul_x86_avx2_available();
void packed_matmul_x86_avx2_range(
PackedMatmulArgs const &args,
std::int64_t out_start,
std::int64_t out_end);
bool packed_matmul_x86_avx512_available();
void packed_matmul_x86_avx512_range(
PackedMatmulArgs const &args,
std::int64_t out_start,
std::int64_t out_end);
} // namespace orbitquant::cpu
|