1#ifndef INFINI_OPS_BASE_AWQ_MARLIN_REPACK_H_
2#define INFINI_OPS_BASE_AWQ_MARLIN_REPACK_H_
17 const int64_t size_n,
const int64_t num_bits,
18 const bool is_a_8bit,
Tensor out)
29 Validate(b_q_weight, out);
33 const int64_t size_n,
const int64_t num_bits,
34 const bool is_a_8bit,
Tensor out)
const = 0;
38 const int64_t size_n,
const int64_t num_bits,
39 const bool is_a_8bit,
const Tensor out)
const {
42 "`AwqMarlinRepack` attributes changed after descriptor creation");
44 const std::equal_to<Tensor> same_metadata;
47 "`AwqMarlinRepack` tensor metadata differs from its descriptor");
67 void Validate(
const Tensor b_q_weight,
const Tensor out)
const {
69 "`AwqMarlinRepack` requires `num_bits` to be 4 or 8");
72 "`AwqMarlinRepack` requires positive `size_k` divisible by 16 "
73 "and `size_n` divisible by 64");
75 "`AwqMarlinRepack` requires `size_k` divisible by 32 for A8 "
77 assert(
size_k_ <= std::numeric_limits<int>::max() &&
78 size_n_ <= std::numeric_limits<int>::max() &&
81 "`AwqMarlinRepack` dimensions exceed CUDA kernel limits");
83 assert(b_q_weight.ndim() == 2 &&
84 b_q_weight.size(0) ==
static_cast<Tensor::Size
>(
size_k_) &&
87 "`AwqMarlinRepack` requires `b_q_weight` shape "
88 "[`size_k`, `size_n / pack_factor`]");
89 assert(b_q_weight.dtype() == DataType::kInt32 &&
90 b_q_weight.IsContiguous() &&
91 "`AwqMarlinRepack` requires contiguous int32 `b_q_weight`");
93 const Tensor::Shape expected_out_shape{
94 static_cast<Tensor::Size
>(
size_k_ / 16),
96 assert(out.shape() == expected_out_shape &&
97 "`AwqMarlinRepack` output shape is incorrect");
98 assert(out.dtype() == DataType::kInt32 && out.IsContiguous() &&
99 "`AwqMarlinRepack` requires contiguous int32 output");
100 assert(out.device().type() == b_q_weight.device().type() &&
101 out.device().index() == b_q_weight.device().index() &&
102 "`AwqMarlinRepack` requires input and output on the same device");
Definition awq_marlin_repack.h:14
int64_t size_n_
Definition awq_marlin_repack.h:56
void ValidateCallMetadata(const Tensor b_q_weight, const int64_t size_k, const int64_t size_n, const int64_t num_bits, const bool is_a_8bit, const Tensor out) const
Definition awq_marlin_repack.h:37
int64_t size_k_
Definition awq_marlin_repack.h:54
bool is_a_8bit_
Definition awq_marlin_repack.h:62
AwqMarlinRepack(const Tensor b_q_weight, const int64_t size_k, const int64_t size_n, const int64_t num_bits, const bool is_a_8bit, Tensor out)
Definition awq_marlin_repack.h:16
Tensor out_metadata_
Definition awq_marlin_repack.h:52
int64_t pack_factor_
Definition awq_marlin_repack.h:60
virtual void operator()(const Tensor b_q_weight, const int64_t size_k, const int64_t size_n, const int64_t num_bits, const bool is_a_8bit, Tensor out) const =0
int64_t num_bits_
Definition awq_marlin_repack.h:58
int device_index_
Definition awq_marlin_repack.h:64
Tensor b_q_weight_metadata_
Definition awq_marlin_repack.h:50
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8