InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
gptq_marlin_repack.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_GPTQ_MARLIN_REPACK_H_
2#define INFINI_OPS_BASE_GPTQ_MARLIN_REPACK_H_
3
4#include <cassert>
5#include <cstdint>
6#include <functional>
7#include <limits>
8
9#include "operator.h"
10
11namespace infini::ops {
12
13// Aligned with vLLM's low-level `gptq_marlin_repack` operator.
14class GptqMarlinRepack : public Operator<GptqMarlinRepack> {
15 public:
16 GptqMarlinRepack(const Tensor b_q_weight, const Tensor perm,
17 const int64_t size_k, const int64_t size_n,
18 const int64_t num_bits, const bool is_a_8bit, Tensor out)
19 : b_q_weight_metadata_{b_q_weight},
20 perm_metadata_{perm},
21 out_metadata_{out},
22 size_k_{size_k},
23 size_n_{size_n},
24 num_bits_{num_bits},
25 pack_factor_{num_bits == 4 ? 8
26 : num_bits == 8 ? 4
27 : 0},
28 is_a_8bit_{is_a_8bit},
29 has_perm_{perm.numel() != 0},
30 device_index_{b_q_weight.device().index()} {
31 Validate(b_q_weight, perm, out);
32 }
33
34 virtual void operator()(const Tensor b_q_weight, const Tensor perm,
35 const int64_t size_k, const int64_t size_n,
36 const int64_t num_bits, const bool is_a_8bit,
37 Tensor out) const = 0;
38
39 protected:
40 void ValidateCallMetadata(const Tensor b_q_weight, const Tensor perm,
41 const int64_t size_k, const int64_t size_n,
42 const int64_t num_bits, const bool is_a_8bit,
43 const Tensor out) const {
44 assert(size_k == size_k_ && size_n == size_n_ && num_bits == num_bits_ &&
45 is_a_8bit == is_a_8bit_ &&
46 "`GptqMarlinRepack` attributes changed after descriptor creation");
47
48 const std::equal_to<Tensor> same_metadata;
49 assert(same_metadata(b_q_weight_metadata_, b_q_weight) &&
50 same_metadata(perm_metadata_, perm) &&
51 same_metadata(out_metadata_, out) &&
52 "`GptqMarlinRepack` tensor metadata differs from its descriptor");
53 }
54
56
58
60
61 int64_t size_k_{0};
62
63 int64_t size_n_{0};
64
65 int64_t num_bits_{0};
66
67 int64_t pack_factor_{0};
68
69 bool is_a_8bit_{false};
70
71 bool has_perm_{false};
72
74
75 private:
76 void Validate(const Tensor b_q_weight, const Tensor perm,
77 const Tensor out) const {
78 assert((num_bits_ == 4 || num_bits_ == 8) &&
79 "`GptqMarlinRepack` requires `num_bits` to be 4 or 8");
80 assert(size_k_ > 0 && size_n_ > 0 && size_k_ % 16 == 0 &&
81 size_n_ % 64 == 0 &&
82 "`GptqMarlinRepack` requires positive `size_k` divisible by 16 "
83 "and `size_n` divisible by 64");
84 assert((!is_a_8bit_ || size_k_ % 32 == 0) &&
85 "`GptqMarlinRepack` requires `size_k` divisible by 32 for A8 "
86 "layouts");
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");
91
92 assert(b_q_weight.ndim() == 2 &&
93 b_q_weight.size(0) ==
94 static_cast<Tensor::Size>(size_k_ / pack_factor_) &&
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`");
101
102 assert(perm.ndim() == 1 &&
103 (!has_perm_ || perm.numel() == static_cast<Tensor::Size>(size_k_)) &&
104 "`GptqMarlinRepack` requires empty `perm` or shape [`size_k`]");
105 assert(perm.dtype() == DataType::kInt32 && perm.IsContiguous() &&
106 "`GptqMarlinRepack` requires contiguous int32 `perm`");
107 assert((!is_a_8bit_ || !has_perm_) &&
108 "`GptqMarlinRepack` does not support `perm` for A8 layouts");
109
110 const Tensor::Shape expected_out_shape{
111 static_cast<Tensor::Size>(size_k_ / 16),
112 static_cast<Tensor::Size>(size_n_ * 16 / pack_factor_)};
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");
117
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();
121 };
122 assert(same_device(perm) && same_device(out) &&
123 "`GptqMarlinRepack` requires all tensors on the same device");
124 }
125};
126
127} // namespace infini::ops
128
129#endif // INFINI_OPS_BASE_GPTQ_MARLIN_REPACK_H_
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