1#ifndef INFINI_OPS_BASE_GPTQ_MARLIN_REPACK_H_
2#define INFINI_OPS_BASE_GPTQ_MARLIN_REPACK_H_
17 const int64_t size_k,
const int64_t size_n,
18 const int64_t num_bits,
const bool is_a_8bit,
Tensor out)
31 Validate(b_q_weight, perm, out);
35 const int64_t size_k,
const int64_t size_n,
36 const int64_t num_bits,
const bool is_a_8bit,
41 const int64_t size_k,
const int64_t size_n,
42 const int64_t num_bits,
const bool is_a_8bit,
46 "`GptqMarlinRepack` attributes changed after descriptor creation");
48 const std::equal_to<Tensor> same_metadata;
52 "`GptqMarlinRepack` tensor metadata differs from its descriptor");
76 void Validate(
const Tensor b_q_weight,
const Tensor perm,
79 "`GptqMarlinRepack` requires `num_bits` to be 4 or 8");
82 "`GptqMarlinRepack` requires positive `size_k` divisible by 16 "
83 "and `size_n` divisible by 64");
85 "`GptqMarlinRepack` requires `size_k` divisible by 32 for A8 "
87 assert(
size_k_ <= std::numeric_limits<int>::max() &&
88 size_n_ <= std::numeric_limits<int>::max() &&
89 size_n_ <= std::numeric_limits<int64_t>::max() / 16 &&
90 "`GptqMarlinRepack` dimensions exceed CUDA kernel limits");
92 assert(b_q_weight.ndim() == 2 &&
95 b_q_weight.size(1) ==
static_cast<Tensor::Size
>(
size_n_) &&
96 "`GptqMarlinRepack` requires `b_q_weight` shape "
97 "[`size_k / pack_factor`, `size_n`]");
98 assert(b_q_weight.dtype() == DataType::kInt32 &&
99 b_q_weight.IsContiguous() &&
100 "`GptqMarlinRepack` requires contiguous int32 `b_q_weight`");
102 assert(perm.ndim() == 1 &&
104 "`GptqMarlinRepack` requires empty `perm` or shape [`size_k`]");
105 assert(perm.dtype() == DataType::kInt32 && perm.IsContiguous() &&
106 "`GptqMarlinRepack` requires contiguous int32 `perm`");
108 "`GptqMarlinRepack` does not support `perm` for A8 layouts");
110 const Tensor::Shape expected_out_shape{
111 static_cast<Tensor::Size
>(
size_k_ / 16),
113 assert(out.shape() == expected_out_shape &&
114 "`GptqMarlinRepack` output shape is incorrect");
115 assert(out.dtype() == DataType::kInt32 && out.IsContiguous() &&
116 "`GptqMarlinRepack` requires contiguous int32 output");
118 const auto same_device = [&](
const Tensor tensor) {
119 return tensor.device().type() == b_q_weight.device().type() &&
120 tensor.device().index() == b_q_weight.device().index();
122 assert(same_device(perm) && same_device(out) &&
123 "`GptqMarlinRepack` requires all tensors on the same device");
Definition gptq_marlin_repack.h:14
Tensor out_metadata_
Definition gptq_marlin_repack.h:59
void ValidateCallMetadata(const Tensor b_q_weight, const Tensor perm, 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 gptq_marlin_repack.h:40
bool is_a_8bit_
Definition gptq_marlin_repack.h:69
virtual void operator()(const Tensor b_q_weight, const Tensor perm, const int64_t size_k, const int64_t size_n, const int64_t num_bits, const bool is_a_8bit, Tensor out) const =0
bool has_perm_
Definition gptq_marlin_repack.h:71
int device_index_
Definition gptq_marlin_repack.h:73
int64_t num_bits_
Definition gptq_marlin_repack.h:65
GptqMarlinRepack(const Tensor b_q_weight, const Tensor perm, const int64_t size_k, const int64_t size_n, const int64_t num_bits, const bool is_a_8bit, Tensor out)
Definition gptq_marlin_repack.h:16
int64_t size_n_
Definition gptq_marlin_repack.h:63
Tensor b_q_weight_metadata_
Definition gptq_marlin_repack.h:55
int64_t pack_factor_
Definition gptq_marlin_repack.h:67
Tensor perm_metadata_
Definition gptq_marlin_repack.h:57
int64_t size_k_
Definition gptq_marlin_repack.h:61
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8