InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
cutlass_scaled_mm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_CUTLASS_SCALED_MM_H_
2#define INFINI_OPS_BASE_CUTLASS_SCALED_MM_H_
3
4#include <cassert>
5#include <limits>
6#include <optional>
7
8#include "operator.h"
9
10namespace infini::ops {
11
12class CutlassScaledMm : public Operator<CutlassScaledMm> {
13 public:
14 CutlassScaledMm(const Tensor a, const Tensor b, const Tensor scale_a,
15 const Tensor scale_b, std::optional<Tensor> bias,
16 const DataType out_dtype, Tensor out)
17 : m_{a.ndim() == 2 ? a.size(0) : 0},
18 n_{b.ndim() == 2 ? b.size(1) : 0},
19 k_{a.ndim() == 2 ? a.size(1) : 0},
20 lda_{a.ndim() == 2 ? a.stride(0) : 0},
21 ldb_{b.ndim() == 2 ? b.stride(1) : 0},
22 ldo_{out.ndim() == 2 ? out.stride(0) : 0},
23 scale_a_size_{scale_a.numel()},
24 scale_b_size_{scale_b.numel()},
25 out_dtype_{out_dtype} {
26 assert(a.ndim() == 2 && b.ndim() == 2 && out.ndim() == 2 &&
27 "`CutlassScaledMm` requires 2D matrices");
28 assert(a.dtype() == DataType::kInt8 && b.dtype() == DataType::kInt8 &&
29 "`CutlassScaledMm` requires int8 matrix inputs");
30 assert((out_dtype_ == DataType::kFloat16 ||
31 out_dtype_ == DataType::kBFloat16) &&
32 "`CutlassScaledMm` requires float16 or bfloat16 output");
33 assert(out.dtype() == out_dtype_ &&
34 "`CutlassScaledMm` requires `out_dtype` to match the output dtype");
35 assert(a.size(1) == b.size(0) && out.size(0) == a.size(0) &&
36 out.size(1) == b.size(1) &&
37 "`CutlassScaledMm` matrix shapes are incompatible");
38 assert(m_ > 0 && n_ > 0 && k_ > 0 &&
39 "`CutlassScaledMm` requires non-empty matrices");
40 assert(a.stride(1) == 1 && b.stride(0) == 1 && out.stride(1) == 1 &&
41 "`CutlassScaledMm` requires row-major `a` and `out` and "
42 "column-major `b`");
43 assert(
44 lda_ >= k_ && ldb_ >= k_ && ldo_ >= n_ &&
45 "`CutlassScaledMm` matrix strides must cover their logical dimensions");
46 assert(k_ % 16 == 0 && n_ % 16 == 0 && lda_ % 16 == 0 && ldb_ % 16 == 0 &&
47 ldo_ % 16 == 0 &&
48 "`CutlassScaledMm` requires 16-element aligned matrix dimensions");
49 assert(scale_a.dtype() == DataType::kFloat32 &&
50 scale_b.dtype() == DataType::kFloat32 &&
51 "`CutlassScaledMm` requires float32 scales");
52 assert(scale_a.IsContiguous() && scale_b.IsContiguous() &&
53 "`CutlassScaledMm` requires contiguous scales");
54 const auto scale_a_is_per_token{
55 scale_a.ndim() == 2 && scale_a.size(0) == m_ && scale_a.size(1) == 1};
56 const auto scale_b_is_per_channel{
57 scale_b.ndim() == 2 && scale_b.size(0) == 1 && scale_b.size(1) == n_};
58 assert(
59 (scale_a_size_ == 1 || scale_a_is_per_token) &&
60 (scale_b_size_ == 1 || scale_b_is_per_channel) &&
61 "`CutlassScaledMm` scales must be scalar, per-token, or per-channel");
62 const auto same_device_as_a = [&](const Tensor tensor) {
63 return tensor.device().type() == a.device().type() &&
64 tensor.device().index() == a.device().index();
65 };
66 assert(same_device_as_a(b) && same_device_as_a(scale_a) &&
67 same_device_as_a(scale_b) && same_device_as_a(out) &&
68 "`CutlassScaledMm` tensors must be on the same device");
69 assert(m_ <= std::numeric_limits<int>::max() &&
70 n_ <= std::numeric_limits<int>::max() &&
71 k_ <= std::numeric_limits<int>::max() &&
72 "`CutlassScaledMm` matrix dimensions exceed CUTLASS limits");
73
74 if (bias) {
75 assert(bias->ndim() == 1 && bias->numel() == n_ &&
76 "`CutlassScaledMm` bias must have shape `[n]`");
77 assert(bias->dtype() == out_dtype_ && bias->IsContiguous() &&
78 "`CutlassScaledMm` bias must match the output dtype and be "
79 "contiguous");
80 assert(same_device_as_a(*bias) &&
81 "`CutlassScaledMm` bias must be on the output device");
82 }
83 }
84
85 virtual void operator()(const Tensor a, const Tensor b, const Tensor scale_a,
86 const Tensor scale_b, std::optional<Tensor> bias,
87 const DataType out_dtype, Tensor out) const = 0;
88
89 protected:
90 Tensor::Size m_{0};
91
92 Tensor::Size n_{0};
93
94 Tensor::Size k_{0};
95
96 Tensor::Stride lda_{0};
97
98 Tensor::Stride ldb_{0};
99
100 Tensor::Stride ldo_{0};
101
102 Tensor::Size scale_a_size_{0};
103
104 Tensor::Size scale_b_size_{0};
105
106 DataType out_dtype_;
107};
108
109} // namespace infini::ops
110
111#endif // INFINI_OPS_BASE_CUTLASS_SCALED_MM_H_
Definition cutlass_scaled_mm.h:12
virtual void operator()(const Tensor a, const Tensor b, const Tensor scale_a, const Tensor scale_b, std::optional< Tensor > bias, const DataType out_dtype, Tensor out) const =0
Tensor::Stride ldb_
Definition cutlass_scaled_mm.h:98
Tensor::Size m_
Definition cutlass_scaled_mm.h:90
Tensor::Size k_
Definition cutlass_scaled_mm.h:94
Tensor::Stride ldo_
Definition cutlass_scaled_mm.h:100
Tensor::Stride lda_
Definition cutlass_scaled_mm.h:96
Tensor::Size scale_b_size_
Definition cutlass_scaled_mm.h:104
CutlassScaledMm(const Tensor a, const Tensor b, const Tensor scale_a, const Tensor scale_b, std::optional< Tensor > bias, const DataType out_dtype, Tensor out)
Definition cutlass_scaled_mm.h:14
Tensor::Size scale_a_size_
Definition cutlass_scaled_mm.h:102
DataType out_dtype_
Definition cutlass_scaled_mm.h:106
Tensor::Size n_
Definition cutlass_scaled_mm.h:92
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8