[MLAS] add q4 quantize and transpose kernel to support MatMulNBits QDQ fuse (#21054)

### Description

1. added kernel to quantize matmul B tensor to q4, and store in the same
shape as original tensor. scales and zero points are calculated as well.
scales and zero points have the same shape.
2. added kernel to transpose q4 B tensor to B tensor in MatMulNBits.
Scales and zero points are transposed as well.

####
Benchmark
<1024 x 4096 input, 64 quant block, 8 threads>: 
 - quantize: 23035923 ns
 - transpose: 718635 ns

<1024 x 4095 input, 64 quant block, 8 threads>: 
 - quantize: 26759319 ns
 - transpose: 1279064 ns

### Motivation and Context
The MatMulNbits tool chain current only supports converting a MatMul op
direct to MatMulNBits op. MatMulNbits op is not an ONNX standard op.
Therefore, we need the tool chain to support converting MatMul to Q/DQ
format, and later in the transform step converts DQ + MatMul to
MatMulNBits. The tensors stored in DQ are the quantized constants and
will be stored in the MatMulNBits.
This commit is contained in:
Jing Fang 2024-06-19 17:15:45 -07:00 committed by GitHub
parent 8448f31d90
commit 6817b013b9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 1258 additions and 146 deletions

View file

@ -358,3 +358,70 @@ MlasDequantizeBlockwise(
int columns,
MLAS_THREADPOOL* thread_pool
);
/**
* @brief Blockwise 2 bits or 4 bits quantization. After quantization, the weights and zero points
* are packed row-wise. In terms of the qbits type, dst and src have the same shape, and
* scales and zero_points have the same shape.
* columns must be multiple of 8 / qbits.
* @tparam Tin
* @tparam qbits number of bits used for quantization, 2 or 4
* @param src points to the floating point matrix, to be quantized, row major shape [rows, columns]
* @param scales points to the scales matrix, row major
* @param zero_points points to the zero_points matrix, row major
* @param dst points to the quantized matrix, shape [rows, columns] row major in qbits type.
* In uint8_t type, shape is [rows, columns * qbits / 8].
* @param columnwise true when quantize elements in a column, false when quantize elements in a row.
* @param rows
* @param columns
* @param quant_block_size number of elements in a quantize block
* @param thread_pool
*/
template <typename Tin, int qbits>
void
MlasQDQQuantizeBlockwise(
const Tin* src,
Tin* scales,
uint8_t* zero_points,
uint8_t* dst,
bool columnwise,
int rows,
int columns,
int quant_block_size,
MLAS_THREADPOOL* thread_pool
);
/**
* @brief Transpose blockwise quantized tensors. The src tensors are row major. src weights and zero
* points are packed row-wise. The dst tensors are column major. dst weights and zero points
* are packed column-wise.
* @tparam Tin
* @tparam qbits number of bits used for quantization, 2 or 4
* @param src_weights points to the quantized matrix, row major, shape [rows, columns] in qbits type.
* In uint8_t type, shape is [rows, columns * qbits / 8].
* @param src_scales points to the scales matrix, row major
* @param src_zero_points points to the zero_points matrix, row major. Packed row-wise.
* @param dst_weights points to the quantized matrix, column major. Packed column-wise.
* @param dst_scales points to the scales matrix, column major
* @param dst_zero_points points to the zero_points matrix, column major. Packed column-wise.
* @param columnwise true when quantize elements in a column, false when quantize elements in a row.
* @param rows
* @param columns
* @param quant_block_size number of elements in a quantize block
* @param thread_pool
*/
template <typename Tin, int qbits>
void
MlasQDQTransposeBlockwiseQuantized(
const uint8_t* src_weights,
const Tin* src_scales,
const uint8_t* src_zero_points,
uint8_t* dst_weights,
Tin* dst_scales,
uint8_t* dst_zero_points,
bool columnwise,
int rows,
int columns,
int quant_block_size,
MLAS_THREADPOOL* thread_pool
);

View file

@ -638,6 +638,669 @@ struct BlockwiseQuantizer {
}
};
/**
* @brief Blockwise quantization methods for QDQ format. Input tensor is quantized along column
* or row. Scales and zeros are calculated. Based on qbits, consecutive quantized elements
* in memory are packed together, which means the packing is along the row. Quantized data
* are stored in row major, so the output tensor reserves same shape, in terms of qbits type,
* as the input tensor.
* @tparam Tin source data type, e.g. fp32/fp16
* @tparam qbits number of bits in each quantized element
*/
template <typename Tin, int qbits>
struct BlockwiseQDQQuantizer;
template <typename Tin>
struct BlockwiseQDQQuantizer<Tin, 4> {
static MLAS_FORCEINLINE uint8_t GetElem(uint8_t val, int32_t idx)
{
return (val >> (idx << 2)) & 0xF;
}
static MLAS_FORCEINLINE uint8_t SetElem(uint8_t val, int32_t idx, uint8_t dst)
{
auto shift = idx << 2;
return ((val & 0xF) << shift) | (dst & (~(0xF << shift)));
}
static MLAS_FORCEINLINE uint8_t Pack(uint8_t v0, uint8_t v1)
{
return (v0 & 0xF) | ((v1 & 0xF) << 4);
}
// If src is row major, then dst is column major. Transpose:
// | src0: low 4 bit | src0: high 4 bit |
// | src1: low 4 bit | src1: high 4 bit |
// -->
// | dst0: low 4 bit | dst1: low 4 bit |
// | dst0: high 4 bit| dst1: high 4 bit |
// If src is column major, then dst is row major. Transpose:
// | src0: low 4 bit | src1: low 4 bit |
// | src0: high 4 bit| src1: high 4 bit |
// -->
// | dst0: low 4 bit | dst0: high 4 bit |
// | dst1: low 4 bit | dst1: high 4 bit |
static MLAS_FORCEINLINE void Transpose(uint8_t src0, uint8_t src1, uint8_t& dst0, uint8_t& dst1)
{
dst0 = (src0 & 0xF) | ((src1 & 0xF) << 4);
dst1 = ((src0 & 0xF0) >> 4) | (src1 & 0xF0);
}
static MLAS_FORCEINLINE uint8_t QuantizeV(Tin src, float reciprocal_scale, uint8_t zero_point)
{
return static_cast<uint8_t>(
std::clamp(
static_cast<int32_t>(
std::roundf(static_cast<float>(src) * reciprocal_scale)
) + static_cast<int32_t>(zero_point),
0,
BitsTraits<4>::kMax
)
);
}
/**
* @brief Quantize a matrix shape [rows, columns] row-wise. Scales and zero points are calculated.
* Quantized data are packed row-wise based on qbits. Quantized data are stored in row
* major, so the output tensor reserves the shape, in terms output type.
* Thread block is [1, quant_block_size * 2].
* @param src the source matrix, row major: [rows * columns]
* @param scales the scales of quantized blocks, row major layout with shape:
* [rows * ceil(columns / quant_block_size)]
* @param zero_points the zero points of quantized blocks, packed. Same shape as scales
* in terms of output type. In terms of uint8_t, the shape is:
* [ceil(rows * ceil(columns / quant_block_size) * qbits / 8)]
* @param dst the quantized weights, row major: [rows * columns] in terms of
* output type. In terms of uint8_t, the shape is: [ceil(rows * columns * qbits / 8]
* @param rows number of rows in the source matrix
* @param columns number of columns in the source matrix, must satisfy
* ceil(columns / quant_block_size) % 2 == 0, so in each thread block,
* zero points are packed into one byte.
* @param quant_block_size number of elements quantized together.
* @param thread_pool thread pool for parallel processing
*/
static void QuantizeRowWise(
const Tin* src,
Tin* scales,
uint8_t* zero_points,
uint8_t* dst,
int32_t rows,
int32_t columns,
int32_t quant_block_size,
MLAS_THREADPOOL* thread_pool
)
{
MLAS_UNREFERENCED_PARAMETER(src);
MLAS_UNREFERENCED_PARAMETER(scales);
MLAS_UNREFERENCED_PARAMETER(zero_points);
MLAS_UNREFERENCED_PARAMETER(dst);
MLAS_UNREFERENCED_PARAMETER(rows);
MLAS_UNREFERENCED_PARAMETER(columns);
MLAS_UNREFERENCED_PARAMETER(quant_block_size);
MLAS_UNREFERENCED_PARAMETER(thread_pool);
ORT_THROW("BlockwiseQDQQuantizer::BlockwiseQDQQuantizer is not implemented");
}
/**
* @brief Quantize a matrix shape [rows, columns] column-wise. Scales and zero points are calculated.
* Quantized data are packed row-wise based on qbits. Quantized data are stored in row major
* so the output tensor reserves the shape, in terms output type.
* @param src the source matrix, row major: [rows * columns]
* @param scales the scales of quantized blocks, row major with shape:
* [ceil(rows/quant_block_size) * columns]
* @param zero_points the zero points of quantized blocks, packed. Same shape as scales in terms
* of output type. In uint8_t, the shape is:
* [ceil(columns * ceil(rows / quant_block_size) * qbits / 8)]
* @param dst the quantized weights, row major: [rows * columns] in terms of output type.
* In uint8_t, the shape is: [ceil(rows * columns * qbits / 8]
* @param rows number of rows in the source matrix
* @param columns number of columns in the source matrix.
* @param quant_block_size number of rows/columns quantized together
* @param thread_pool thread pool for parallel processing
*/
static void QuantizeColumnWise(
const Tin* src,
Tin* scales,
uint8_t* zero_points,
uint8_t* dst,
int32_t rows,
int32_t columns,
int32_t quant_block_size,
MLAS_THREADPOOL* thread_pool
)
{
// Must avoid multiple thread write to a single byte, which means the starting index
// of a thread block must be even. To achieve that, we need to customize the thread
// block size based on the parity of columns.
if (columns & 1) {
QuantizeColumnWisePackUnaligned(
src, scales, zero_points, dst, rows, columns, quant_block_size, thread_pool
);
} else {
QuantizeColumnWisePackAligned(
src, scales, zero_points, dst, rows, columns, quant_block_size, thread_pool
);
}
}
/**
* @brief Transpose quantized tensors, which has been column-wise quantized, for use in MatMulNbits.
* Since both src tensor and dst tensor are packed, it's not needed to consider sign
* during the unpacking/packing in transpose.
* @param src_weights The quantized weights, row major: [rows, columns] in qbits type.
* In uint8_t, size of [ceil(rows * columns * qbits / 8)].
* @param src_scales [ceil(rows / quant_block_size), columns]
* @param src_zero_points [ceil(rows / quant_block_size), columns] in qbits type. In uint8_t, size of
* [ceil(ceil(rows / quant_block_size) * columns * qbits / 8 )].
* @param dst_weights the transposed quantized weights, column major. In uint8_t, the shape is
* [columns, ceil(rows / quant_block_size), ceil(quant_block_size * qbits / 8)]
* @param dst_scales [columns, ceil(rows / quant_block_size)]
* @param dst_zero_points [columns, ceil(ceil(rows / quant_block_size) * qbits / 8)] in uint8_t.
* @param rows number of src rows in qbits type.
* @param columns number of src columns in qbits type.
* @param quant_block_size number of elements quantized together
* @param thread_pool thread pool for parallel processing
*/
static void TransposeColumnWiseQuantized(
const uint8_t* src_weights,
const Tin* src_scales,
const uint8_t* src_zero_points,
uint8_t* dst_weights,
Tin* dst_scales,
uint8_t* dst_zero_points,
int32_t rows,
int32_t columns,
int32_t quant_block_size,
MLAS_THREADPOOL* thread_pool
)
{
// Must avoid multiple thread write to a single byte, which means the starting index
// of a thread block must be even. To achieve that, we need to customize the thread
// block size based on the parity of columns.
if (columns & 1) {
TransposeColumnWiseQuantizedPackUnaligned(
src_weights, src_scales, src_zero_points,
dst_weights, dst_scales, dst_zero_points,
rows, columns, quant_block_size, thread_pool
);
} else {
TransposeColumnWiseQuantizedPackAligned(
src_weights, src_scales, src_zero_points,
dst_weights, dst_scales, dst_zero_points,
rows, columns, quant_block_size, thread_pool
);
}
}
private:
static void QuantizeColumnWisePackAligned(
const Tin* src,
Tin* scales,
uint8_t* zero_points,
uint8_t* dst,
int32_t rows,
int32_t columns,
int32_t quant_block_size,
MLAS_THREADPOOL* thread_pool
)
{
ORT_ENFORCE(columns % 2 == 0, "Columns must be multiple of 2.");
// Thread block is [quant_block_size, thread_blk_size]. thread_blk_size % 2 == 0.
constexpr int32_t thread_blk_size = 128;
const auto num_row_thread_blk = (rows + quant_block_size - 1) / quant_block_size;
const auto num_col_thread_blk = (columns + thread_blk_size - 1) / thread_blk_size;
const auto num_thread_blk = num_row_thread_blk * num_col_thread_blk;
constexpr auto minf = std::numeric_limits<float>::lowest();
constexpr auto maxf = std::numeric_limits<float>::max();
MlasTryBatchParallel(
thread_pool, static_cast<ptrdiff_t>(num_thread_blk),
[&](ptrdiff_t thread_blk_idx) {
// !!warning!!: buffering the whole thread block
constexpr int32_t buffer_size = 128;
ORT_ENFORCE(buffer_size == thread_blk_size, "buffer size must be equal to thread block size.");
float reciprocal_scale_t[buffer_size];
uint8_t zp_t[buffer_size];
float vmin_t[buffer_size];
float vmax_t[buffer_size];
const int32_t row_thread_blk_idx = static_cast<int32_t>(thread_blk_idx / num_col_thread_blk);
const int32_t col_thread_blk_idx = static_cast<int32_t>(thread_blk_idx % num_col_thread_blk);
const int32_t row_idx = row_thread_blk_idx * quant_block_size;
const int32_t col_idx = col_thread_blk_idx * buffer_size;
const int32_t row_size = std::min(quant_block_size, rows - row_idx);
const int32_t col_size = std::min(buffer_size, columns - col_idx);
// input_idx, scale_idx, col_size are aligned to 2
auto input_idx = row_idx * columns + col_idx;
auto scale_idx = row_thread_blk_idx * columns + col_idx;
Tin scale0_tt, scale1_tt;
uint8_t v0_tt, v1_tt;
std::fill_n(vmin_t, buffer_size, maxf);
std::fill_n(vmax_t, buffer_size, minf);
// calculate min/max
for (int32_t j = 0, input_idx_t = input_idx; j < row_size; ++j, input_idx_t += columns) {
// TODO(fajin): use SIMD
for (int32_t i = 0; i < col_size; i += 2) {
auto v0 = static_cast<float>(src[input_idx_t + i]);
auto v1 = static_cast<float>(src[input_idx_t + i + 1]);
vmin_t[i] = std::min(vmin_t[i], v0);
vmax_t[i] = std::max(vmax_t[i], v0);
vmin_t[i + 1] = std::min(vmin_t[i + 1], v1);
vmax_t[i + 1] = std::max(vmax_t[i + 1], v1);
}
}
// calculate scale and zero point, and store
for (int32_t i = 0; i < col_size; i += 2) {
v0_tt = v1_tt = BitsTraits<4>::kMid;
if (zero_points) {
range2scalezp<Tin, 4>(vmin_t[i], vmax_t[i], scale0_tt, v0_tt);
range2scalezp<Tin, 4>(vmin_t[i + 1], vmax_t[i + 1], scale1_tt, v1_tt);
zero_points[(scale_idx + i) >> 1] = Pack(v0_tt, v1_tt);
} else {
range2scale<Tin, 4>(vmin_t[i], vmax_t[i], scale0_tt);
range2scale<Tin, 4>(vmin_t[i + 1], vmax_t[i + 1], scale1_tt);
}
scales[scale_idx + i] = scale0_tt;
scales[scale_idx + i + 1] = scale1_tt;
float scalef0 = static_cast<float>(scale0_tt);
reciprocal_scale_t[i] = scalef0 ? 1.0f / scalef0 : 0.0f;
zp_t[i] = v0_tt;
float scalef1 = static_cast<float>(scale1_tt);
reciprocal_scale_t[i + 1] = scalef1 ? 1.0f / scalef1 : 0.0f;
zp_t[i + 1] = v1_tt;
}
// quantize and pack
for (int32_t j = 0, input_idx_t = input_idx; j < row_size; ++j, input_idx_t += columns) {
// TODO(fajin): use SIMD
for (int32_t i = 0; i < col_size; i += 2) {
v0_tt = QuantizeV(src[input_idx_t + i], reciprocal_scale_t[i], zp_t[i]);
v1_tt = QuantizeV(src[input_idx_t + i + 1], reciprocal_scale_t[i + 1], zp_t[i + 1]);
dst[(input_idx_t + i) >> 1] = Pack(v0_tt, v1_tt);
}
}
}
);
}
static void QuantizeColumnWisePackUnaligned(
const Tin* src,
Tin* scales,
uint8_t* zero_points,
uint8_t* dst,
int32_t rows,
int32_t columns,
int32_t quant_block_size,
MLAS_THREADPOOL* thread_pool
)
{
// Thread block is [quant_block_size * 2, columns], so the packed bytes do not cross threads.
constexpr auto minf = std::numeric_limits<float>::lowest();
constexpr auto maxf = std::numeric_limits<float>::max();
auto row_thread_blk_size = quant_block_size * 2;
auto num_row_thread_blk = (rows + row_thread_blk_size - 1) / (row_thread_blk_size);
MlasTryBatchParallel(
thread_pool, static_cast<ptrdiff_t>(num_row_thread_blk),
[&](ptrdiff_t thread_blk_idx) {
constexpr int32_t buffer_size = 128;
float reciprocal_scale_t[buffer_size];
uint8_t zp_t[buffer_size];
float vmin_t[buffer_size];
float vmax_t[buffer_size];
auto row_thread_blk_idx = static_cast<int32_t>(thread_blk_idx);
int32_t row_idx = row_thread_blk_idx * row_thread_blk_size;
int32_t row_idx_end = std::min(row_thread_blk_size + row_idx, rows);
auto input_idx = row_idx * columns;
auto scale_idx = row_thread_blk_idx * 2 * columns;
Tin scale0_tt, scale1_tt;
uint8_t v0_tt, v1_tt;
for (; row_idx < row_idx_end; row_idx += quant_block_size) {
// per quant block row
auto quant_row_size = std::min(quant_block_size, row_idx_end - row_idx);
auto input_buffer_idx = input_idx;
auto scale_buffer_idx = scale_idx;
for (int32_t buffer_idx = 0; buffer_idx < columns; buffer_idx += buffer_size) {
// per buffer column
auto buffer_col_size = std::min(buffer_size, columns - buffer_idx);
std::fill_n(vmin_t, buffer_size, maxf);
std::fill_n(vmax_t, buffer_size, minf);
// calculate min/max of [quant block, buffer]
auto input_idx_t = input_buffer_idx;
for (int32_t j = 0; j < quant_row_size; ++j, input_idx_t += columns) {
// TODO(fajin): use SIMD
for (int32_t i = 0; i < buffer_col_size; ++i) {
auto v = static_cast<float>(src[input_idx_t + i]);
vmin_t[i] = std::min(vmin_t[i], v);
vmax_t[i] = std::max(vmax_t[i], v);
}
}
// calculate scale and zero point
auto scale_buffer_idx_end = scale_buffer_idx + buffer_col_size;
int32_t col_idx = 0;
// leading unailgned zero points
if (scale_buffer_idx & 1) {
v0_tt = BitsTraits<4>::kMid;
if (zero_points) {
range2scalezp<Tin, 4>(vmin_t[0], vmax_t[0], scale0_tt, v0_tt);
zero_points[scale_buffer_idx >> 1] = SetElem(
v0_tt, 1, zero_points[scale_buffer_idx >> 1]
);
} else {
range2scale<Tin, 4>(vmin_t[0], vmax_t[0], scale0_tt);
}
scales[scale_buffer_idx] = scale0_tt;
float scalef = static_cast<float>(scale0_tt);
reciprocal_scale_t[0] = scalef ? 1.0f / scalef : 0.0f;
zp_t[0] = v0_tt;
++col_idx;
++scale_buffer_idx;
}
// aligned zero points
for (; scale_buffer_idx < scale_buffer_idx_end - 1; col_idx += 2, scale_buffer_idx += 2) {
v0_tt = v1_tt = BitsTraits<4>::kMid;
if (zero_points) {
range2scalezp<Tin, 4>(vmin_t[col_idx], vmax_t[col_idx], scale0_tt, v0_tt);
range2scalezp<Tin, 4>(vmin_t[col_idx + 1], vmax_t[col_idx + 1], scale1_tt, v1_tt);
zero_points[scale_buffer_idx >> 1] = Pack(v0_tt, v1_tt);
} else {
range2scale<Tin, 4>(vmin_t[col_idx], vmax_t[col_idx], scale0_tt);
range2scale<Tin, 4>(vmin_t[col_idx + 1], vmax_t[col_idx + 1], scale1_tt);
}
scales[scale_buffer_idx] = scale0_tt;
scales[scale_buffer_idx + 1] = scale1_tt;
float scalef0 = static_cast<float>(scale0_tt);
reciprocal_scale_t[col_idx] = scalef0 ? 1.0f / scalef0 : 0.0f;
zp_t[col_idx] = v0_tt;
float scalef1 = static_cast<float>(scale1_tt);
reciprocal_scale_t[col_idx + 1] = scalef1 ? 1.0f / scalef1 : 0.0f;
zp_t[col_idx + 1] = v1_tt;
}
// tailing unaligned elements
if (scale_buffer_idx < scale_buffer_idx_end) {
v0_tt = BitsTraits<4>::kMid;
if (zero_points) {
range2scalezp<Tin, 4>(vmin_t[col_idx], vmax_t[col_idx], scale0_tt, v0_tt);
zero_points[scale_buffer_idx >> 1] = SetElem(
v0_tt, 0, zero_points[scale_buffer_idx >> 1]
);
} else {
range2scale<Tin, 4>(vmin_t[col_idx], vmax_t[col_idx], scale0_tt);
}
scales[scale_buffer_idx] = scale0_tt;
float scalef = static_cast<float>(scale0_tt);
reciprocal_scale_t[col_idx] = scalef ? 1.0f / scalef : 0.0f;
zp_t[col_idx] = v0_tt;
++scale_buffer_idx;
}
// quantize and pack
input_idx_t = input_buffer_idx;
for (int32_t j = 0; j < quant_row_size; ++j, input_idx_t += columns) {
auto input_idx_t_start = input_idx_t;
auto input_idx_t_end = input_idx_t + buffer_col_size;
col_idx = 0;
// leading unaligned output
if (input_idx_t_start & 1) {
v1_tt = QuantizeV(src[input_idx_t_start], reciprocal_scale_t[col_idx], zp_t[col_idx]);
dst[input_idx_t_start >> 1] = SetElem(v1_tt, 1, dst[input_idx_t_start >> 1]);
++col_idx;
++input_idx_t_start;
}
// aligned output
// TODO(fajin): use SIMD
for (; input_idx_t_start < input_idx_t_end - 1; col_idx += 2, input_idx_t_start += 2) {
v0_tt = QuantizeV(src[input_idx_t_start], reciprocal_scale_t[col_idx], zp_t[col_idx]);
v1_tt = QuantizeV(
src[input_idx_t_start + 1], reciprocal_scale_t[col_idx + 1], zp_t[col_idx + 1]
);
dst[input_idx_t_start >> 1] = Pack(v0_tt, v1_tt);
}
// tailing unaligned output
if (input_idx_t_start < input_idx_t_end) {
v0_tt = QuantizeV(src[input_idx_t_start], reciprocal_scale_t[col_idx], zp_t[col_idx]);
dst[input_idx_t_start >> 1] = SetElem(v0_tt, 0, dst[input_idx_t_start >> 1]);
}
}
input_buffer_idx += buffer_size;
}
input_idx += quant_block_size * columns;
scale_idx += columns;
}
}
);
}
static void TransposeColumnWiseQuantizedPackAligned(
const uint8_t* src_weights, // [rows, columns / 2]
const Tin* src_scales, // [ceil(rows / quant_block_size), columns]
const uint8_t* src_zero_points, // [ceil(rows / quant_block_size), columns / 2]
uint8_t* dst_weights, // [columns, ceil(rows / quant_block_size), ceil(quant_block_size / 2)]
Tin* dst_scales, // [columns, ceil(rows / quant_block_size)]
uint8_t* dst_zero_points, // [columns, ceil(ceil(rows / quant_block_size) / 2)]
int32_t rows,
int32_t columns,
int32_t quant_block_size,
MLAS_THREADPOOL* thread_pool
)
{
ORT_ENFORCE(columns % 2 == 0, "Columns must be multiple of 2");
auto row_quant_blk_num = (rows + quant_block_size - 1) / quant_block_size;
auto dst_bytes_per_quant_blk = (quant_block_size * 4 + 7) / 8;
// number of rows in transposed dst
auto dstT_num_row = row_quant_blk_num * dst_bytes_per_quant_blk;
auto packed_col_size = columns / 2;
// weight transpose thread block is [dst_bytes_per_quant_blk, 2] on dst_Transpose.
// Map to src it is [quant_block_size, 1]. Both in uint8_t.
auto num_thread_blk = row_quant_blk_num * packed_col_size;
MlasTryBatchParallel(
thread_pool, static_cast<ptrdiff_t>(num_thread_blk),
[&](ptrdiff_t thread_blk_idx) {
uint8_t src0_t, src1_t;
uint8_t dst0_t, dst1_t;
auto row_thread_blk_idx = static_cast<int32_t>(thread_blk_idx / packed_col_size);
auto col_thread_blk_idx = static_cast<int32_t>(thread_blk_idx % packed_col_size);
auto dstT_row_idx = row_thread_blk_idx * dst_bytes_per_quant_blk;
auto dstT_col_idx = col_thread_blk_idx * 2;
auto dst_idx = dstT_col_idx * dstT_num_row + dstT_row_idx;
auto src_row_idx = row_thread_blk_idx * quant_block_size;
auto src_row_end_idx = std::min(src_row_idx + quant_block_size, rows);
auto src_col_idx = col_thread_blk_idx;
auto src_idx = src_row_idx * packed_col_size + src_col_idx;
auto src_end_idx = src_row_end_idx * packed_col_size + src_col_idx;
for (; src_idx < src_end_idx - packed_col_size; ++dst_idx) {
src0_t = src_weights[src_idx];
src1_t = src_weights[src_idx + packed_col_size];
src_idx += packed_col_size + packed_col_size;
Transpose(src0_t, src1_t, dst0_t, dst1_t);
dst_weights[dst_idx] = dst0_t;
dst_weights[dst_idx + dstT_num_row] = dst1_t;
}
if (src_idx < src_end_idx) {
src0_t = src_weights[src_idx];
src1_t = 0;
Transpose(src0_t, src1_t, dst0_t, dst1_t);
dst_weights[dst_idx] = dst0_t;
dst_weights[dst_idx + dstT_num_row] = dst1_t;
}
}
);
// Transpose scales. Thread block is [row_quant_blk_num, 1] on dst_Transpose.
MlasTryBatchParallel(
thread_pool, static_cast<ptrdiff_t>(columns),
[&](ptrdiff_t thread_blk_idx) {
auto col_thread_blk_idx = static_cast<int32_t>(thread_blk_idx);
auto src_idx = col_thread_blk_idx;
auto dst_idx = col_thread_blk_idx * row_quant_blk_num;
for (int32_t i = 0; i < row_quant_blk_num; ++i, ++dst_idx, src_idx += columns) {
dst_scales[dst_idx] = src_scales[src_idx];
}
}
);
if (src_zero_points) {
// Transpose zero points. Thread block is [ceil(row_quant_blk_num / 2), 2]
// on dst_Transpose. Map to src it is [row_quant_blk_num, 1]. Both in uint8_t.
auto dst_zp_row_num = (row_quant_blk_num + 1) / 2;
MlasTryBatchParallel(
thread_pool, static_cast<ptrdiff_t>(packed_col_size),
[&](ptrdiff_t thread_blk_idx) {
uint8_t src0_t, src1_t;
uint8_t dst0_t, dst1_t;
auto col_thread_blk_idx = static_cast<int32_t>(thread_blk_idx);
auto src_idx = col_thread_blk_idx;
auto src_end_idx = row_quant_blk_num * packed_col_size + col_thread_blk_idx;
auto dst_idx = col_thread_blk_idx * 2 * dst_zp_row_num;
for (; src_idx < src_end_idx - packed_col_size; ++dst_idx) {
src0_t = src_zero_points[src_idx];
src1_t = src_zero_points[src_idx + packed_col_size];
Transpose(src0_t, src1_t, dst0_t, dst1_t);
dst_zero_points[dst_idx] = dst0_t;
dst_zero_points[dst_idx + dst_zp_row_num] = dst1_t;
src_idx += packed_col_size + packed_col_size;
}
if (src_idx < src_end_idx) {
src0_t = src_zero_points[src_idx];
src1_t = 0;
Transpose(src0_t, src1_t, dst0_t, dst1_t);
dst_zero_points[dst_idx] = dst0_t;
dst_zero_points[dst_idx + dst_zp_row_num] = dst1_t;
}
}
);
}
}
static void TransposeColumnWiseQuantizedPackUnaligned(
const uint8_t* src_weights, // size of [ceil(rows * columns / 2)]
const Tin* src_scales, // [ceil(rows / quant_block_size), columns]
const uint8_t* src_zero_points, // size of [ceil(ceil(rows / quant_block_size) * columns / 2)]
uint8_t *dst_weights, // [columns, ceil(rows / quant_block_size), ceil(quant_block_size / 2)]
Tin* dst_scales, // [columns, ceil(rows / quant_block_size)]
uint8_t* dst_zero_points, // [columns, ceil(ceil(rows / quant_block_size) / 2)]
int32_t rows,
int32_t columns,
int32_t quant_block_size,
MLAS_THREADPOOL* thread_pool)
{
auto row_quant_blk_num = (rows + quant_block_size - 1) / quant_block_size;
auto dst_bytes_per_quant_blk = (quant_block_size * 4 + 7) / 8;
// number of rows in transposed dst
auto dstT_num_row = row_quant_blk_num * dst_bytes_per_quant_blk;
// weight transpose thread block is [dst_bytes_per_quant_blk, 1] on dst_Transpose in uint8_t.
// Map to src it is [quant_block_size, 1] in int4.
auto num_thread_blk = row_quant_blk_num * columns;
MlasTryBatchParallel(
thread_pool, static_cast<ptrdiff_t>(num_thread_blk),
[&](ptrdiff_t thread_blk_idx) {
uint8_t src0_t, src1_t;
auto row_thread_blk_idx = static_cast<int32_t>(thread_blk_idx / columns);
auto col_thread_blk_idx = static_cast<int32_t>(thread_blk_idx % columns);
auto dstT_row_idx = row_thread_blk_idx * dst_bytes_per_quant_blk;
auto dst_idx = col_thread_blk_idx * dstT_num_row + dstT_row_idx;
auto src_row_idx = row_thread_blk_idx * quant_block_size;
auto src_row_end_idx = std::min(src_row_idx + quant_block_size, rows);
auto src_idx = src_row_idx * columns + col_thread_blk_idx;
auto src_end_idx = src_row_end_idx * columns + col_thread_blk_idx;
for (; src_idx < src_end_idx - columns; ++dst_idx) {
src0_t = GetElem(src_weights[src_idx >> 1], src_idx & 1);
src1_t = GetElem(src_weights[(src_idx + columns) >> 1], (src_idx + columns) & 1);
dst_weights[dst_idx] = (src0_t & 0xf) | ((src1_t & 0xf) << 4);
src_idx += columns + columns;
}
if (src_idx < src_end_idx) {
src0_t = GetElem(src_weights[src_idx >> 1], src_idx & 1);
dst_weights[dst_idx] = src0_t & 0xf;
}
}
);
// Transpose scales. Thread block is [row_quant_blk_num, 1] on dst_Transpose.
MlasTryBatchParallel(
thread_pool, static_cast<ptrdiff_t>(columns),
[&](ptrdiff_t thread_blk_idx) {
auto col_thread_blk_idx = static_cast<int32_t>(thread_blk_idx);
auto src_idx = col_thread_blk_idx;
auto dst_idx = col_thread_blk_idx * row_quant_blk_num;
for (int32_t i = 0; i < row_quant_blk_num; ++i, ++dst_idx, src_idx += columns) {
dst_scales[dst_idx] = src_scales[src_idx];
}
}
);
if (src_zero_points) {
// Transpose zero points. Thread block is [ceil(row_quant_blk_num / 2), 1] on dst_Transpose in uint8_t.
// Map to src it is [row_quant_blk_num, 1] in int4.
auto dst_zp_row_num = (row_quant_blk_num + 1) / 2;
MlasTryBatchParallel(
thread_pool, static_cast<ptrdiff_t>(columns),
[&](ptrdiff_t thread_blk_idx) {
uint8_t src0_t, src1_t;
auto col_thread_blk_idx = static_cast<int32_t>(thread_blk_idx);
auto src_idx = col_thread_blk_idx;
auto src_end_idx = row_quant_blk_num * columns + col_thread_blk_idx;
auto dst_idx = col_thread_blk_idx * dst_zp_row_num;
for (; src_idx < src_end_idx - columns; ++dst_idx) {
src0_t = GetElem(src_zero_points[src_idx >> 1], src_idx & 1);
src1_t = GetElem(src_zero_points[(src_idx + columns) >> 1], (src_idx + columns) & 1);
dst_zero_points[dst_idx] = (src0_t & 0xf) | ((src1_t & 0xf) << 4);
src_idx += columns + columns;
}
if (src_idx < src_end_idx) {
src0_t = GetElem(src_zero_points[src_idx >> 1], src_idx & 1);
dst_zero_points[dst_idx] = src0_t & 0xf;
}
}
);
}
}
};
template <typename T, int qbits>
void
@ -1068,8 +1731,7 @@ MlasDequantizeBlockwise(
}
}
template
void
template void
MlasDequantizeBlockwise<float, 4>(
float* dst,
const uint8_t* src,
@ -1080,4 +1742,111 @@ MlasDequantizeBlockwise<float, 4>(
int rows,
int columns,
MLAS_THREADPOOL* thread_pool
);
);
template <typename Tin, int qbits>
void
MlasQDQQuantizeBlockwise(
const Tin* src,
Tin* scales,
uint8_t* zero_points,
uint8_t* dst,
bool columnwise,
int rows,
int columns,
int quant_block_size,
MLAS_THREADPOOL* thread_pool
)
{
if (columnwise) {
BlockwiseQDQQuantizer<Tin, qbits>::QuantizeColumnWise(
src, scales, zero_points, dst, rows, columns, quant_block_size, thread_pool
);
} else {
BlockwiseQDQQuantizer<Tin, qbits>::QuantizeRowWise(
src, scales, zero_points, dst, rows, columns, quant_block_size, thread_pool
);
}
}
template void
MlasQDQQuantizeBlockwise<float, 4>(
const float* src,
float* scales,
uint8_t* zero_points,
uint8_t* dst,
bool columnwise,
int rows,
int columns,
int quant_block_size,
MLAS_THREADPOOL* thread_pool
);
template void
MlasQDQQuantizeBlockwise<MLAS_FP16, 4>(
const MLAS_FP16* src,
MLAS_FP16* scales,
uint8_t* zero_points,
uint8_t* dst,
bool columnwise,
int rows,
int columns,
int quant_block_size,
MLAS_THREADPOOL* thread_pool
);
template <typename Tin, int qbits>
void
MlasQDQTransposeBlockwiseQuantized(
const uint8_t* src_weights,
const Tin* src_scales,
const uint8_t* src_zero_points,
uint8_t* dst_weights,
Tin* dst_scales,
uint8_t* dst_zero_points,
bool columnwise,
int rows,
int columns,
int quant_block_size,
MLAS_THREADPOOL* thread_pool
)
{
if (columnwise) {
BlockwiseQDQQuantizer<Tin, qbits>::TransposeColumnWiseQuantized(
src_weights, src_scales, src_zero_points, dst_weights, dst_scales, dst_zero_points,
rows, columns, quant_block_size, thread_pool
);
} else {
ORT_THROW("Row-wise MlasQDQTransposeBlockwiseQuantized is not implemented");
}
}
template void
MlasQDQTransposeBlockwiseQuantized<float, 4>(
const uint8_t* src_weights,
const float* src_scales,
const uint8_t* src_zero_points,
uint8_t* dst_weights,
float* dst_scales,
uint8_t* dst_zero_points,
bool columnwise,
int rows,
int columns,
int quant_block_size,
MLAS_THREADPOOL* thread_pool
);
template void
MlasQDQTransposeBlockwiseQuantized<MLAS_FP16, 4>(
const uint8_t* src_weights,
const MLAS_FP16* src_scales,
const uint8_t* src_zero_points,
uint8_t* dst_weights,
MLAS_FP16* dst_scales,
uint8_t* dst_zero_points,
bool columnwise,
int rows,
int columns,
int quant_block_size,
MLAS_THREADPOOL* thread_pool
);

View file

@ -552,53 +552,80 @@ struct BlockedQuantizeLinear<float, TOut, 2> {
std::ptrdiff_t N, const std::ptrdiff_t quant_block_size,
const std::ptrdiff_t thread_block_size, bool saturate) {
ORT_UNUSED_PARAMETER(saturate);
// to avoid a byte being writen from mutiple threads, use 2 * N as thread block
ORT_UNUSED_PARAMETER(thread_block_size);
constexpr auto low = static_cast<int32_t>(TOut::min_val);
constexpr auto high = static_cast<int32_t>(TOut::max_val);
const auto num_thread_block_N = (N + thread_block_size - 1) / thread_block_size;
const auto num_thread_block = M * K * num_thread_block_N;
const TensorOpCost unit_cost{static_cast<double>(thread_block_size * sizeof(float) * 2),
static_cast<double>(thread_block_size * sizeof(typename TOut::UnpackedType)),
static_cast<double>(thread_block_size) * 2.0};
auto KN = K * N;
auto num_quant_block_KN = (K + quant_block_size - 1) / quant_block_size * N;
const auto num_thread_block_KN = K * num_thread_block_N;
auto size_thread_block = 2 * N;
auto num_thread_block = (M * K + 1) / 2;
auto num_quant_block_K = (K + quant_block_size - 1) / quant_block_size;
auto num_quant_block_KN = num_quant_block_K * N;
auto MK = M * K;
const TensorOpCost unit_cost{static_cast<double>(size_thread_block * sizeof(float) * 2),
static_cast<double>(size_thread_block * sizeof(typename TOut::UnpackedType)),
static_cast<double>(size_thread_block) * 2.0};
concurrency::ThreadPool::TryParallelFor(
thread_pool,
num_thread_block,
unit_cost,
[&](std::ptrdiff_t begin, std::ptrdiff_t end) {
auto m = begin / num_thread_block_KN, k = begin % num_thread_block_KN / num_thread_block_N;
auto n_blk = begin % num_thread_block_N, n = n_blk * thread_block_size;
auto output_idx = m * KN + k * N + n;
auto quant_param_idx = m * num_quant_block_KN + k / quant_block_size * N;
auto quant_param_idx_t = quant_param_idx + n;
begin <<= 1, end = std::min(end << 1, MK);
auto output_idx = begin * N;
auto m = begin / K, k = begin % K;
auto zp_idx = m * num_quant_block_KN + k / quant_block_size * N;
for (; begin < end; ++begin) {
auto n_end = std::min(N, n + thread_block_size);
// TODO(fajin): 1> use SIMD, 2> set block to quant_block_size * thread_block_size
// TODO(fajin): process 2 elements at a time
for (; n < n_end; ++n, ++output_idx, ++quant_param_idx_t) {
// TODO(fajin): perf difference
auto zp_idx_t = zp_idx;
auto output_idx_end = output_idx + N;
// leading unaligned output
if (output_idx & 1) {
auto zp = zero_point
? static_cast<int32_t>(zero_point[quant_param_idx_t >> 1].GetElem(quant_param_idx_t & 1))
? static_cast<int32_t>(zero_point[zp_idx_t >> 1].GetElem(zp_idx_t & 1))
: 0;
auto sc = scale[quant_param_idx_t];
auto sc = scale[zp_idx_t];
auto v = std::clamp(static_cast<int32_t>(std::nearbyint(input[output_idx] / sc)) + zp, low, high);
output[output_idx >> 1].SetElem(output_idx & 1, static_cast<typename TOut::UnpackedType>(v));
output[output_idx >> 1].SetElem(1, static_cast<typename TOut::UnpackedType>(v));
++output_idx;
++zp_idx_t;
}
if (n == N) {
n = 0;
++k;
if (k == K) {
k = 0;
quant_param_idx += N;
} else if (k % quant_block_size == 0) {
quant_param_idx += N;
}
// TODO(fajin): use SIMD
// aligned output
auto output_t = reinterpret_cast<typename TOut::UnpackedType*>(output);
for (; output_idx < output_idx_end - 1; output_idx += 2, zp_idx_t += 2) {
auto zp0 = zero_point
? static_cast<int32_t>(zero_point[zp_idx_t >> 1].GetElem(zp_idx_t & 1))
: 0;
auto zp1 = zero_point
? static_cast<int32_t>(zero_point[(zp_idx_t + 1) >> 1].GetElem((zp_idx_t + 1) & 1))
: 0;
auto sc0 = scale[zp_idx_t];
auto sc1 = scale[zp_idx_t + 1];
auto v0 = std::clamp(static_cast<int32_t>(std::nearbyint(input[output_idx] / sc0)) + zp0, low, high);
auto v1 = std::clamp(static_cast<int32_t>(std::nearbyint(input[output_idx + 1] / sc1)) + zp1, low, high);
output_t[output_idx >> 1] = static_cast<typename TOut::UnpackedType>((v0 & 0xF) | ((v1 & 0xF) << 4));
}
quant_param_idx_t = quant_param_idx;
// tailing unaligned output
if (output_idx < output_idx_end) {
auto zp = zero_point
? static_cast<int32_t>(zero_point[zp_idx_t >> 1].GetElem(zp_idx_t & 1))
: 0;
auto sc = scale[zp_idx_t];
auto v = std::clamp(static_cast<int32_t>(std::nearbyint(input[output_idx] / sc)) + zp, low, high);
output[output_idx >> 1].SetElem(0, static_cast<typename TOut::UnpackedType>(v));
++output_idx;
}
++k;
if (k == K) {
k = 0;
zp_idx += N;
} else if (k % quant_block_size == 0) {
zp_idx += N;
}
}
});
@ -610,53 +637,59 @@ struct BlockedQuantizeLinear<float, TOut, 2> {
ORT_UNUSED_PARAMETER(saturate);
constexpr auto low = static_cast<int32_t>(TOut::min_val);
constexpr auto high = static_cast<int32_t>(TOut::max_val);
// quant block size is used as thread block size
const auto num_thread_block_K = (K + quant_block_size - 1) / quant_block_size;
const auto num_thread_block = num_thread_block_K * M;
const TensorOpCost unit_cost{static_cast<double>(quant_block_size * sizeof(float)),
static_cast<double>(quant_block_size * sizeof(typename TOut ::UnpackedType)),
static_cast<double>(quant_block_size) * 2.0};
// to avoid a byte being writen from mutiple threads, use 2 * K as thread block
auto size_thread_block = 2 * K;
auto quant_block_num_K = (K + quant_block_size - 1) / quant_block_size;
auto num_thread_block = (M + 1) / 2;
TensorOpCost unit_cost{static_cast<double>(size_thread_block * sizeof(float)),
static_cast<double>(size_thread_block * sizeof(typename TOut ::UnpackedType)),
static_cast<double>(size_thread_block) * 2.0};
concurrency::ThreadPool::TryParallelFor(
thread_pool,
num_thread_block,
unit_cost,
[&](std::ptrdiff_t begin, std::ptrdiff_t end) {
auto m = begin / num_thread_block_K, k_blk = begin % num_thread_block_K, k = k_blk * quant_block_size;
auto output_idx = m * K + k;
begin <<= 1, end = std::min(end << 1, M);
auto output_idx = begin * K;
auto zp_idx = begin * quant_block_num_K;
for (; begin < end; ++begin) {
auto zp = zero_point ? static_cast<int32_t>(zero_point[begin >> 1].GetElem(begin & 1)) : 0;
auto sc = scale[begin];
size_t output_idx_end = std::min(K - k, quant_block_size) + output_idx;
size_t out_start = output_idx, out_end = output_idx_end;
for (; begin < end; ++begin, output_idx += K) {
auto output_row_idx_start = output_idx;
auto output_row_idx_end = output_row_idx_start + K;
if (out_start & 1) {
auto v = std::clamp(static_cast<int32_t>(std::nearbyint(input[out_start] / sc)) + zp, low, high);
output[out_start >> 1].SetElem(1, static_cast<typename TOut::UnpackedType>(v));
++out_start;
for (; output_row_idx_start < output_row_idx_end; output_row_idx_start += quant_block_size, ++zp_idx) {
auto zp = zero_point ? static_cast<int32_t>(zero_point[zp_idx >> 1].GetElem(zp_idx & 1)) : 0;
auto sc = scale[zp_idx];
size_t out_start = output_row_idx_start;
size_t out_end = std::min(output_row_idx_start + quant_block_size, output_row_idx_end);
if (out_start & 1) {
auto v = std::clamp(static_cast<int32_t>(std::nearbyint(input[out_start] / sc)) + zp, low, high);
output[out_start >> 1].SetElem(1, static_cast<typename TOut::UnpackedType>(v));
++out_start;
}
if (out_end & 1) {
--out_end;
auto v = std::clamp(static_cast<int32_t>(std::nearbyint(input[out_end] / sc)) + zp, low, high);
output[out_end >> 1].SetElem(0, static_cast<typename TOut::UnpackedType>(v));
}
if constexpr (std::is_same<TOut, Int4x2>::value) {
MlasQuantizeLinearS4(input + out_start, reinterpret_cast<uint8_t*>(&(output[out_start >> 1])),
out_end - out_start, sc, static_cast<int8_t>(zp));
} else {
MlasQuantizeLinearU4(input + out_start, reinterpret_cast<uint8_t*>(&(output[out_start >> 1])),
out_end - out_start, sc, static_cast<int8_t>(zp));
}
}
if (out_end & 1) {
--out_end;
auto v = std::clamp(static_cast<int32_t>(std::nearbyint(input[out_end] / sc)) + zp, low, high);
output[out_end >> 1].SetElem(0, static_cast<typename TOut::UnpackedType>(v));
}
if constexpr (std::is_same<TOut, Int4x2>::value) {
MlasQuantizeLinearS4(input + out_start, reinterpret_cast<uint8_t*>(&(output[out_start >> 1])),
out_end - out_start, sc, static_cast<int8_t>(zp));
} else {
MlasQuantizeLinearU4(input + out_start, reinterpret_cast<uint8_t*>(&(output[out_start >> 1])),
out_end - out_start, sc, static_cast<int8_t>(zp));
}
output_idx = output_idx_end;
k = output_idx % K;
}
});
}
};
// Bug(fajin): the same byte in output / zero_point must not be written by different threads, otherwise
// the result is undefined. This is not handled in the current implementation.
template <typename TOut>
struct BlockedQuantizeLinear<MLFloat16, TOut, 2> {
static void opNotLastAxis(concurrency::ThreadPool* thread_pool, const MLFloat16* input, const MLFloat16* scale,
@ -664,54 +697,84 @@ struct BlockedQuantizeLinear<MLFloat16, TOut, 2> {
std::ptrdiff_t N, const std::ptrdiff_t quant_block_size,
const std::ptrdiff_t thread_block_size, bool saturate) {
ORT_UNUSED_PARAMETER(saturate);
// to avoid a byte being writen from mutiple threads, use 2 * N as thread block
ORT_UNUSED_PARAMETER(thread_block_size);
constexpr auto low = static_cast<int32_t>(TOut::min_val);
constexpr auto high = static_cast<int32_t>(TOut::max_val);
const auto num_thread_block_N = (N + thread_block_size - 1) / thread_block_size;
const auto num_thread_block = M * K * num_thread_block_N;
const TensorOpCost unit_cost{static_cast<double>(thread_block_size * sizeof(MLFloat16) * 2),
static_cast<double>(thread_block_size * sizeof(typename TOut::UnpackedType)),
static_cast<double>(thread_block_size) * 2.0};
auto KN = K * N;
auto num_quant_block_KN = (K + quant_block_size - 1) / quant_block_size * N;
const auto num_thread_block_KN = K * num_thread_block_N;
auto size_thread_block = 2 * N;
auto num_thread_block = (M * K + 1) / 2;
auto num_quant_block_K = (K + quant_block_size - 1) / quant_block_size;
auto num_quant_block_KN = num_quant_block_K * N;
auto MK = M * K;
const TensorOpCost unit_cost{static_cast<double>(size_thread_block * sizeof(float) * 2),
static_cast<double>(size_thread_block * sizeof(typename TOut::UnpackedType)),
static_cast<double>(size_thread_block) * 2.0};
concurrency::ThreadPool::TryParallelFor(
thread_pool,
num_thread_block,
unit_cost,
[&](std::ptrdiff_t begin, std::ptrdiff_t end) {
auto m = begin / num_thread_block_KN, k = begin % num_thread_block_KN / num_thread_block_N;
auto n_blk = begin % num_thread_block_N, n = n_blk * thread_block_size;
auto output_idx = m * KN + k * N + n;
auto quant_param_idx = m * num_quant_block_KN + k / quant_block_size * N;
auto quant_param_idx_t = quant_param_idx + n;
begin <<= 1, end = std::min(end << 1, MK);
auto output_idx = begin * N;
auto m = begin / K, k = begin % K;
auto zp_idx = m * num_quant_block_KN + k / quant_block_size * N;
for (; begin < end; ++begin) {
auto n_end = std::min(N, n + thread_block_size);
// TODO(fajin): 1> use SIMD, 2> set block to quant_block_size * thread_block_size
// TODO(fajin): process 2 elements at a time
for (; n < n_end; ++n, ++output_idx, ++quant_param_idx_t) {
// TODO(fajin): perf difference
auto zp_idx_t = zp_idx;
auto output_idx_end = output_idx + N;
// leading unaligned output
if (output_idx & 1) {
auto zp = zero_point
? static_cast<int32_t>(zero_point[quant_param_idx_t >> 1].GetElem(quant_param_idx_t & 1))
? static_cast<int32_t>(zero_point[zp_idx_t >> 1].GetElem(zp_idx_t & 1))
: 0;
auto sc = scale[quant_param_idx_t].ToFloat();
auto v = std::clamp(static_cast<int32_t>(std::nearbyint(input[output_idx].ToFloat() / sc)) + zp,
low, high);
output[output_idx >> 1].SetElem(output_idx & 1, static_cast<typename TOut::UnpackedType>(v));
auto sc = scale[zp_idx_t].ToFloat();
auto v = std::clamp(
static_cast<int32_t>(std::nearbyint(input[output_idx].ToFloat() / sc)) + zp, low, high);
output[output_idx >> 1].SetElem(1, static_cast<typename TOut::UnpackedType>(v));
++output_idx;
++zp_idx_t;
}
if (n == N) {
n = 0;
++k;
if (k == K) {
k = 0;
quant_param_idx += N;
} else if (k % quant_block_size == 0) {
quant_param_idx += N;
}
// TODO(fajin): use SIMD
// aligned output
auto output_t = reinterpret_cast<typename TOut::UnpackedType*>(output);
for (; output_idx < output_idx_end - 1; output_idx += 2, zp_idx_t += 2) {
auto zp0 = zero_point
? static_cast<int32_t>(zero_point[zp_idx_t >> 1].GetElem(zp_idx_t & 1))
: 0;
auto zp1 = zero_point
? static_cast<int32_t>(zero_point[(zp_idx_t + 1) >> 1].GetElem((zp_idx_t + 1) & 1))
: 0;
auto sc0 = scale[zp_idx_t].ToFloat();
auto sc1 = scale[zp_idx_t + 1].ToFloat();
auto v0 = std::clamp(
static_cast<int32_t>(std::nearbyint(input[output_idx].ToFloat() / sc0)) + zp0, low, high);
auto v1 = std::clamp(
static_cast<int32_t>(std::nearbyint(input[output_idx + 1].ToFloat() / sc1)) + zp1, low, high);
output_t[output_idx >> 1] = static_cast<typename TOut::UnpackedType>((v0 & 0xF) | ((v1 & 0xF) << 4));
}
quant_param_idx_t = quant_param_idx;
// tailing unaligned output
if (output_idx < output_idx_end) {
auto zp = zero_point
? static_cast<int32_t>(zero_point[zp_idx_t >> 1].GetElem(zp_idx_t & 1))
: 0;
auto sc = scale[zp_idx_t].ToFloat();
auto v = std::clamp(
static_cast<int32_t>(std::nearbyint(input[output_idx].ToFloat() / sc)) + zp, low, high);
output[output_idx >> 1].SetElem(0, static_cast<typename TOut::UnpackedType>(v));
++output_idx;
}
++k;
if (k == K) {
k = 0;
zp_idx += N;
} else if (k % quant_block_size == 0) {
zp_idx += N;
}
}
});
@ -723,32 +786,55 @@ struct BlockedQuantizeLinear<MLFloat16, TOut, 2> {
ORT_UNUSED_PARAMETER(saturate);
constexpr auto low = static_cast<int32_t>(TOut::min_val);
constexpr auto high = static_cast<int32_t>(TOut::max_val);
// quant block size is used as thread block size
const auto num_thread_block_K = (K + quant_block_size - 1) / quant_block_size;
const auto num_thread_block = num_thread_block_K * M;
const TensorOpCost unit_cost{static_cast<double>(quant_block_size * sizeof(MLFloat16)),
static_cast<double>(quant_block_size * sizeof(typename TOut::UnpackedType)),
static_cast<double>(quant_block_size) * 2.0};
// to avoid a byte being writen from mutiple threads, use 2 * K as thread block
auto size_thread_block = 2 * K;
auto quant_block_num_K = (K + quant_block_size - 1) / quant_block_size;
auto num_thread_block = (M + 1) / 2;
TensorOpCost unit_cost{static_cast<double>(size_thread_block * sizeof(float)),
static_cast<double>(size_thread_block * sizeof(typename TOut ::UnpackedType)),
static_cast<double>(size_thread_block) * 2.0};
concurrency::ThreadPool::TryParallelFor(
thread_pool,
num_thread_block,
unit_cost,
[&](std::ptrdiff_t begin, std::ptrdiff_t end) {
auto m = begin / num_thread_block_K, k_blk = begin % num_thread_block_K, k = k_blk * quant_block_size;
auto output_idx = m * K + k;
begin <<= 1, end = std::min(end << 1, M);
auto output_idx = begin * K;
auto zp_idx = begin * quant_block_num_K;
for (; begin < end; ++begin) {
// each thread block is also a quantization block
auto zp = zero_point ? static_cast<int32_t>(zero_point[begin >> 1].GetElem(begin & 1)) : 0;
auto sc = scale[begin].ToFloat();
auto output_idx_end = std::min(K - k, quant_block_size) + output_idx;
for (; output_idx < output_idx_end; ++output_idx) {
auto v = std::clamp(static_cast<int32_t>(std::nearbyint(input[output_idx].ToFloat() / sc)) + zp,
low, high);
output[output_idx >> 1].SetElem(output_idx & 1, static_cast<typename TOut::UnpackedType>(v));
for (; begin < end; ++begin, output_idx += K) {
auto output_row_idx_start = output_idx;
auto output_row_idx_end = output_row_idx_start + K;
for (; output_row_idx_start < output_row_idx_end; output_row_idx_start += quant_block_size, ++zp_idx) {
auto zp = zero_point ? static_cast<int32_t>(zero_point[zp_idx >> 1].GetElem(zp_idx & 1)) : 0;
auto sc = scale[zp_idx].ToFloat();
size_t out_start = output_row_idx_start;
size_t out_end = std::min(output_row_idx_start + quant_block_size, output_row_idx_end);
if (out_start & 1) {
auto v = std::clamp(
static_cast<int32_t>(std::nearbyint(input[out_start].ToFloat() / sc)) + zp, low, high);
output[out_start >> 1].SetElem(1, static_cast<typename TOut::UnpackedType>(v));
++out_start;
}
if (out_end & 1) {
--out_end;
auto v = std::clamp(
static_cast<int32_t>(std::nearbyint(input[out_end].ToFloat() / sc)) + zp, low, high);
output[out_end >> 1].SetElem(0, static_cast<typename TOut::UnpackedType>(v));
}
auto output_t = reinterpret_cast<typename TOut::UnpackedType*>(output);
for (; out_start < out_end; out_start += 2) {
auto v0 = std::clamp(
static_cast<int32_t>(std::nearbyint(input[out_start].ToFloat() / sc)) + zp, low, high);
auto v1 = std::clamp(
static_cast<int32_t>(std::nearbyint(input[out_start + 1].ToFloat() / sc)) + zp, low, high);
output_t[out_start >> 1] = static_cast<typename TOut::UnpackedType>((v0 & 0xF) | ((v1 & 0xF) << 4));
}
}
k = output_idx % K;
}
});
}

View file

@ -66,6 +66,37 @@ void QuantizeMatMul4BitsBlockwise(
tp.get());
}
template <typename T>
void QuantizeQDQMatMul4BitsBlockwise(
py::array_t<uint8_t> dst, // shape: [K, N / 2]
py::array_t<T> src, // shape: [K, N]
py::array_t<T> scale, // shape: [block_per_K, N]
py::array_t<uint8_t> zero_points, // shape: [block_per_K, N / 2]
int32_t quant_block_size,
int32_t N,
int32_t K,
bool is_symmetric) {
OrtThreadPoolParams to;
auto tp = concurrency::CreateThreadPool(&onnxruntime::Env::Default(), to,
concurrency::ThreadPoolType::INTRA_OP);
py::buffer_info dst_buf = dst.request();
py::buffer_info src_buf = src.request();
py::buffer_info scale_buf = scale.request();
py::buffer_info zp_buf = zero_points.request();
MlasQDQQuantizeBlockwise<T, 4>(
reinterpret_cast<const T*>(src_buf.ptr),
reinterpret_cast<T*>(scale_buf.ptr),
is_symmetric ? nullptr : reinterpret_cast<uint8_t*>(zp_buf.ptr),
reinterpret_cast<uint8_t*>(dst_buf.ptr),
true,
K,
N,
quant_block_size,
tp.get());
}
template <typename T>
void QuantizeMatMulBnb4Blockwise(
py::array_t<uint8_t> dst,
@ -99,6 +130,8 @@ void CreateQuantPybindModule(py::module& m) {
m.def("quantize_matmul_4bits", &QuantizeMatMul4BitsBlockwise<MLFloat16>);
m.def("quantize_matmul_bnb4", &QuantizeMatMulBnb4Blockwise<float>);
m.def("quantize_matmul_bnb4", &QuantizeMatMulBnb4Blockwise<MLFloat16>);
m.def("quantize_qdq_matmul_4bits", &QuantizeQDQMatMul4BitsBlockwise<float>);
m.def("quantize_qdq_matmul_4bits", &QuantizeQDQMatMul4BitsBlockwise<MLFloat16>);
}
} // namespace python

View file

@ -0,0 +1,118 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include <stdexcept>
#include <numeric>
#include "core/mlas/inc/mlas_q4.h"
#include "test/mlas/bench/bench_util.h"
#include "core/util/thread_utils.h"
static void BM_QDQBlockwiseQuantizer_QuantizeColumnwise(benchmark::State& state) {
int M = state.range(0);
int N = state.range(1);
int quant_block_size = state.range(2);
int threads = state.range(3);
size_t scale_size = (M + quant_block_size - 1) / quant_block_size * N;
auto src = RandomVectorUniform(M * N, -16.0f, 14.0f);
auto scales = std::vector<float>(scale_size);
auto zero_points = std::vector<uint8_t>((scale_size + 1) / 2);
auto dst = std::vector<uint8_t>((M * N + 1) / 2);
OrtThreadPoolParams tpo;
tpo.thread_pool_size = static_cast<int>(threads);
tpo.auto_set_affinity = true;
std::unique_ptr<onnxruntime::concurrency::ThreadPool> tp(
onnxruntime::concurrency::CreateThreadPool(&onnxruntime::Env::Default(),
tpo, onnxruntime::concurrency::ThreadPoolType::INTRA_OP));
for (auto _ : state) {
benchmark::DoNotOptimize(dst.data());
MlasQDQQuantizeBlockwise<float, 4>(
src.data(), scales.data(), zero_points.data(), dst.data(),
true, M, N, quant_block_size, tp.get());
benchmark::ClobberMemory();
}
}
static void BM_MlasQuantizeBlockwise(benchmark::State& state) {
int M = state.range(0);
int N = state.range(1);
int quant_block_size = state.range(2);
int threads = state.range(3);
size_t scale_size = (M + quant_block_size - 1) / quant_block_size * N;
auto src = RandomVectorUniform(M * N, -16.0f, 14.0f);
auto scales = std::vector<float>(scale_size);
auto zero_points = std::vector<uint8_t>((scale_size + 1) / 2);
auto dst = std::vector<uint8_t>((M * N + 1) / 2);
OrtThreadPoolParams tpo;
tpo.thread_pool_size = static_cast<int>(threads);
tpo.auto_set_affinity = true;
std::unique_ptr<onnxruntime::concurrency::ThreadPool> tp(
onnxruntime::concurrency::CreateThreadPool(&onnxruntime::Env::Default(),
tpo, onnxruntime::concurrency::ThreadPoolType::INTRA_OP));
for (auto _ : state) {
benchmark::DoNotOptimize(dst.data());
MlasQuantizeBlockwise<float, 4>(
dst.data(), scales.data(), zero_points.data(), src.data(),
quant_block_size, true, M, N, N, tp.get());
benchmark::ClobberMemory();
}
}
static void BM_QDQBlockwiseQuantizer_TransposeColumnwise(benchmark::State& state) {
int M = state.range(0);
int N = state.range(1);
int quant_block_size = state.range(2);
int threads = state.range(3);
int quant_num_M = (M + quant_block_size - 1) / quant_block_size;
int blob_size = (quant_block_size + 1) / 2;
size_t scale_size = quant_num_M * N;
auto scales = RandomVectorUniform<float>(scale_size, -16.0f, 14.0f);
auto zero_points = RandomVectorUniform<uint8_t>(static_cast<size_t>((scale_size + 1) / 2), 0, 255);
auto dst = RandomVectorUniform<uint8_t>(static_cast<size_t>((M * N + 1) / 2), 0, 255);
auto scales_T = std::vector<float>(scale_size);
auto zero_points_T = std::vector<uint8_t>(((quant_num_M + 1) / 2) * N);
auto dst_T = std::vector<uint8_t>(quant_num_M * blob_size * N);
OrtThreadPoolParams tpo;
tpo.thread_pool_size = static_cast<int>(threads);
tpo.auto_set_affinity = true;
std::unique_ptr<onnxruntime::concurrency::ThreadPool> tp(
onnxruntime::concurrency::CreateThreadPool(&onnxruntime::Env::Default(),
tpo, onnxruntime::concurrency::ThreadPoolType::INTRA_OP));
for (auto _ : state) {
benchmark::DoNotOptimize(dst.data());
MlasQDQTransposeBlockwiseQuantized<float, 4>(
dst.data(), scales.data(), zero_points.data(), dst_T.data(), scales_T.data(), zero_points_T.data(),
true, M, N, quant_block_size, tp.get());
benchmark::ClobberMemory();
}
}
BENCHMARK(BM_QDQBlockwiseQuantizer_QuantizeColumnwise)
->UseRealTime()
->Apply([](benchmark::internal::Benchmark* b) {
b->ArgNames({"M", "N", "quant_block_size", "threads"});
b->ArgsProduct({{1024, 4096}, {4096, 4095}, {64, 128}, {8}});
});
BENCHMARK(BM_MlasQuantizeBlockwise)
->UseRealTime()
->Apply([](benchmark::internal::Benchmark* b) {
b->ArgNames({"M", "N", "quant_block_size", "threads"});
b->ArgsProduct({{1024, 4096}, {4096, 4095}, {64, 128}, {8}});
});
BENCHMARK(BM_QDQBlockwiseQuantizer_TransposeColumnwise)
->UseRealTime()
->Apply([](benchmark::internal::Benchmark* b) {
b->ArgNames({"M", "N", "quant_block_size", "threads"});
b->ArgsProduct({{1024, 4096}, {4096, 4095}, {64, 128}, {2, 8, 16}});
});

View file

@ -29,10 +29,18 @@ class MlasBlockwiseQdqTest : public MlasTestBase {
MatrixGuardBuffer<uint8_t> OutputElements;
MatrixGuardBuffer<float> OutputScales;
MatrixGuardBuffer<uint8_t> OutputOffsets;
MatrixGuardBuffer<uint8_t> QDQOutputElements;
MatrixGuardBuffer<float> QDQOutputScales;
MatrixGuardBuffer<uint8_t> QDQOutputOffsets;
MatrixGuardBuffer<uint8_t> QDQTransposedOutputElements;
MatrixGuardBuffer<float> QDQTransposedOutputScales;
MatrixGuardBuffer<uint8_t> QDQTransposedOutputOffsets;
void Test(int rows, int columns, int block_size, bool columnwise, bool symmetric) {
float* dequant_buf = FpBuf.GetBuffer(rows * columns, true);
float* transposed = FpBuf2.GetBuffer(rows * columns, true);
size_t scale_size = (rows + block_size - 1) / block_size * columns;
size_t zp_size = (scale_size + 1) / 2;
MLAS_THREADPOOL* threadpool_ptr = GetMlasThreadPool();
@ -49,6 +57,8 @@ class MlasBlockwiseQdqTest : public MlasTestBase {
q_data_size_in_bytes, q_scale_size, &q_zp_size_in_bytes);
uint8_t* elements = InputElements.GetBuffer(q_data_size_in_bytes, true);
uint8_t* qdq_weights = QDQOutputElements.GetBuffer((rows * columns + 1) / 2, true);
uint8_t* qdq_weights_T = QDQTransposedOutputElements.GetBuffer(q_data_size_in_bytes, true);
int v = 7;
for (int c = 0; c < columns; c++) {
@ -75,7 +85,11 @@ class MlasBlockwiseQdqTest : public MlasTestBase {
}
float* scales = InputScales.GetBuffer(q_scale_size);
float* qdq_scales = QDQOutputScales.GetBuffer(scale_size);
float* qdq_scales_T = QDQTransposedOutputScales.GetBuffer(q_scale_size);
uint8_t* zp = symmetric ? nullptr : InputOffsets.GetBuffer(q_zp_size_in_bytes, true);
uint8_t* qdq_zp = symmetric ? nullptr : QDQOutputOffsets.GetBuffer(zp_size, true);
uint8_t* qdq_zp_T = symmetric ? nullptr : QDQTransposedOutputOffsets.GetBuffer(q_zp_size_in_bytes, true);
if (zp) {
for (int c = 0; c < meta_cols; c++) {
for (int r = 0; r < meta_rows; r += 2) {
@ -112,16 +126,37 @@ class MlasBlockwiseQdqTest : public MlasTestBase {
MlasQuantizeBlockwise<float, 4>(o_elements, o_scales, o_zp, transposed, block_size,
columnwise, rows, columns, columns, threadpool_ptr);
if (columnwise) {
MlasQDQQuantizeBlockwise<float, 4>(
transposed, qdq_scales, qdq_zp, qdq_weights,
true, rows, columns, block_size, threadpool_ptr);
MlasQDQTransposeBlockwiseQuantized<float, 4>(
qdq_weights, qdq_scales, qdq_zp, qdq_weights_T, qdq_scales_T, qdq_zp_T,
true, rows, columns, block_size, threadpool_ptr);
}
for (int c = 0; c < columns; c++) {
for (int r = 0; r < rows; r += 2) {
int idx = c * q_rows + r / 2;
ASSERT_EQ(o_elements[idx] & 0xf, elements[idx] & 0xf)
<< ", index=[" << r << "x" << c << "], shape=[" << rows << "x" << columns
<< "] block: " << block_size << ", symmetric: " << symmetric << ", columnwise: " << columnwise;
if (columnwise) {
ASSERT_EQ(qdq_weights_T[idx] & 0xf, elements[idx] & 0xf)
<< ", index=[" << r << "x" << c << "], shape=[" << rows << "x" << columns
<< "] block: " << block_size << ", symmetric: " << symmetric << ", columnwise: " << columnwise;
}
if (r + 1 < rows) {
ASSERT_EQ(o_elements[idx] >> 4, elements[idx] >> 4)
<< ", index=[" << r + 1 << "x" << c << "], shape=[" << rows << "x" << columns
<< "] block: " << block_size << ", symmetric: " << symmetric << ", columnwise: " << columnwise;
if (columnwise) {
ASSERT_EQ(qdq_weights_T[idx] >> 4, elements[idx] >> 4)
<< ", index=[" << r + 1 << "x" << c << "], shape=[" << rows << "x" << columns
<< "] block: " << block_size << ", symmetric: " << symmetric << ", columnwise: " << columnwise;
}
}
}
}
@ -132,6 +167,12 @@ class MlasBlockwiseQdqTest : public MlasTestBase {
ASSERT_EQ(o_scales[idx], scales[idx])
<< ", index=" << r << "x" << c << ", shape=[" << rows << "x" << columns
<< "] block: " << block_size << ", symmetric: " << symmetric << ", columnwise: " << columnwise;
if (columnwise) {
ASSERT_EQ(qdq_scales_T[idx], scales[idx])
<< ", index=" << r << "x" << c << ", shape=[" << rows << "x" << columns
<< "] block: " << block_size << ", symmetric: " << symmetric << ", columnwise: " << columnwise;
}
}
}
@ -142,10 +183,20 @@ class MlasBlockwiseQdqTest : public MlasTestBase {
ASSERT_EQ(o_zp[idx] & 0xf, zp[idx] & 0xf)
<< ", index=" << r << "x" << c << ", shape=[" << rows << "x" << columns
<< "] block: " << block_size << ", symmetric: " << symmetric << ", columnwise: " << columnwise;
if (columnwise) {
ASSERT_EQ(qdq_zp_T[idx] & 0xf, zp[idx] & 0xf)
<< ", index=" << r << "x" << c << ", shape=[" << rows << "x" << columns
<< "] block: " << block_size << ", symmetric: " << symmetric << ", columnwise: " << columnwise;
}
if (r + 1 < meta_rows) {
ASSERT_EQ(o_zp[idx] >> 4, zp[idx] >> 4)
<< ", index=" << r + 1 << "x" << c << ", shape=[" << rows << "x" << columns
<< "] block: " << block_size << ", symmetric: " << symmetric << ", columnwise: " << columnwise;
if (columnwise) {
ASSERT_EQ(qdq_zp_T[idx] >> 4, zp[idx] >> 4)
<< ", index=" << r + 1 << "x" << c << ", shape=[" << rows << "x" << columns
<< "] block: " << block_size << ", symmetric: " << symmetric << ", columnwise: " << columnwise;
}
}
}
}

View file

@ -82,12 +82,11 @@ BENCHMARK(BM_Quantize)
static void BM_BlockedQuantize_NotLastAxis(benchmark::State& state) {
using Int4 = onnxruntime::Int4x2;
using UnpackedType = Int4::UnpackedType;
const std::ptrdiff_t M[] = {96, 192, 192};
const std::ptrdiff_t N[] = {2048, 2048, 4096};
const int64_t size_idx = state.range(0);
const int64_t threads = state.range(1);
const int64_t M = state.range(0);
const int64_t N = state.range(1);
const int64_t block_size = state.range(2);
size_t batch_size = M[size_idx] * N[size_idx];
const int64_t threads = state.range(3);
size_t batch_size = M * N;
size_t quant_block_size = 64;
size_t scale_size = batch_size / quant_block_size;
@ -108,7 +107,7 @@ static void BM_BlockedQuantize_NotLastAxis(benchmark::State& state) {
benchmark::DoNotOptimize(a_data_quant);
onnxruntime::BlockedQuantizeLinear<float, Int4, 2>::opNotLastAxis(
tp.get(), a_data, scale, reinterpret_cast<Int4*>(zero_point), reinterpret_cast<Int4*>(a_data_quant),
1, M[size_idx], N[size_idx], static_cast<std::ptrdiff_t>(quant_block_size),
1, M, N, static_cast<std::ptrdiff_t>(quant_block_size),
static_cast<std::ptrdiff_t>(block_size), true);
benchmark::ClobberMemory();
}
@ -121,12 +120,11 @@ static void BM_BlockedQuantize_NotLastAxis(benchmark::State& state) {
static void BM_BlockedQuantize_LastAxis(benchmark::State& state) {
using Int4 = onnxruntime::Int4x2;
using UnpackedType = Int4::UnpackedType;
const std::ptrdiff_t M[] = {96, 192, 192};
const std::ptrdiff_t N[] = {2048, 2048, 4096};
const int64_t size_idx = state.range(0);
const int64_t threads = state.range(1);
const int64_t M = state.range(0);
const int64_t N = state.range(1);
const int64_t quant_block_size = state.range(2);
size_t batch_size = M[size_idx] * N[size_idx];
const int64_t threads = state.range(3);
size_t batch_size = M * N;
size_t scale_size = batch_size / quant_block_size;
float* a_data = GenerateArrayWithRandomValue<float>(batch_size, -16, 14);
@ -146,7 +144,7 @@ static void BM_BlockedQuantize_LastAxis(benchmark::State& state) {
benchmark::DoNotOptimize(a_data_quant);
onnxruntime::BlockedQuantizeLinear<float, Int4, 2>::opLastAxis(
tp.get(), a_data, scale, reinterpret_cast<Int4*>(zero_point), reinterpret_cast<Int4*>(a_data_quant),
M[size_idx], N[size_idx], static_cast<std::ptrdiff_t>(quant_block_size), true);
M, N, static_cast<std::ptrdiff_t>(quant_block_size), true);
benchmark::ClobberMemory();
}
aligned_free(a_data_quant);
@ -159,24 +157,14 @@ BENCHMARK(BM_BlockedQuantize_NotLastAxis)
->UseRealTime()
->Unit(benchmark::TimeUnit::kNanosecond)
->Apply([](benchmark::internal::Benchmark* b) {
for (int size_idx : {0, 1, 2}) {
for (int thread : {2, 4, 8}) {
for (int block_size : {64, 128}) {
b->Args({size_idx, thread, block_size});
}
}
}
b->ArgNames({"M", "N", "block_size", "threads"});
b->ArgsProduct({{1024, 4096}, {4096}, {128}, {2, 8}});
});
BENCHMARK(BM_BlockedQuantize_LastAxis)
->UseRealTime()
->Unit(benchmark::TimeUnit::kNanosecond)
->Apply([](benchmark::internal::Benchmark* b) {
for (int size_idx : {0, 1, 2}) {
for (int thread : {2, 4, 8}) {
for (int quant_block_size : {16, 64, 256}) {
b->Args({size_idx, thread, quant_block_size});
}
}
}
b->ArgNames({"M", "N", "quant_block_size", "threads"});
b->ArgsProduct({{1024, 4096}, {4096}, {64, 128}, {2, 8}});
});