InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
awq_marlin_repack.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_AWQ_MARLIN_REPACK_H_
2#define INFINI_OPS_BASE_AWQ_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 `awq_marlin_repack` operator.
14class AwqMarlinRepack : public Operator<AwqMarlinRepack> {
15 public:
16 AwqMarlinRepack(const Tensor b_q_weight, const int64_t size_k,
17 const int64_t size_n, const int64_t num_bits,
18 const bool is_a_8bit, Tensor out)
19 : b_q_weight_metadata_{b_q_weight},
20 out_metadata_{out},
21 size_k_{size_k},
22 size_n_{size_n},
23 num_bits_{num_bits},
24 pack_factor_{num_bits == 4 ? 8
25 : num_bits == 8 ? 4
26 : 0},
27 is_a_8bit_{is_a_8bit},
28 device_index_{b_q_weight.device().index()} {
29 Validate(b_q_weight, out);
30 }
31
32 virtual void operator()(const Tensor b_q_weight, const int64_t size_k,
33 const int64_t size_n, const int64_t num_bits,
34 const bool is_a_8bit, Tensor out) const = 0;
35
36 protected:
37 void ValidateCallMetadata(const Tensor b_q_weight, const int64_t size_k,
38 const int64_t size_n, const int64_t num_bits,
39 const bool is_a_8bit, const Tensor out) const {
40 assert(size_k == size_k_ && size_n == size_n_ && num_bits == num_bits_ &&
41 is_a_8bit == is_a_8bit_ &&
42 "`AwqMarlinRepack` attributes changed after descriptor creation");
43
44 const std::equal_to<Tensor> same_metadata;
45 assert(same_metadata(b_q_weight_metadata_, b_q_weight) &&
46 same_metadata(out_metadata_, out) &&
47 "`AwqMarlinRepack` tensor metadata differs from its descriptor");
48 }
49
51
53
54 int64_t size_k_{0};
55
56 int64_t size_n_{0};
57
58 int64_t num_bits_{0};
59
60 int64_t pack_factor_{0};
61
62 bool is_a_8bit_{false};
63
65
66 private:
67 void Validate(const Tensor b_q_weight, const Tensor out) const {
68 assert((num_bits_ == 4 || num_bits_ == 8) &&
69 "`AwqMarlinRepack` requires `num_bits` to be 4 or 8");
70 assert(size_k_ > 0 && size_n_ > 0 && size_k_ % 16 == 0 &&
71 size_n_ % 64 == 0 &&
72 "`AwqMarlinRepack` requires positive `size_k` divisible by 16 "
73 "and `size_n` divisible by 64");
74 assert((!is_a_8bit_ || size_k_ % 32 == 0) &&
75 "`AwqMarlinRepack` requires `size_k` divisible by 32 for A8 "
76 "layouts");
77 assert(size_k_ <= std::numeric_limits<int>::max() &&
78 size_n_ <= std::numeric_limits<int>::max() &&
79 size_k_ <=
80 std::numeric_limits<int>::max() / (size_n_ / pack_factor_) &&
81 "`AwqMarlinRepack` dimensions exceed CUDA kernel limits");
82
83 assert(b_q_weight.ndim() == 2 &&
84 b_q_weight.size(0) == static_cast<Tensor::Size>(size_k_) &&
85 b_q_weight.size(1) ==
86 static_cast<Tensor::Size>(size_n_ / pack_factor_) &&
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`");
92
93 const Tensor::Shape expected_out_shape{
94 static_cast<Tensor::Size>(size_k_ / 16),
95 static_cast<Tensor::Size>(size_n_ * 16 / pack_factor_)};
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");
103 }
104};
105
106} // namespace infini::ops
107
108#endif // INFINI_OPS_BASE_AWQ_MARLIN_REPACK_H_
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