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},
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");
32 "`CutlassScaledMm` requires float16 or bfloat16 output");
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 "
45 "`CutlassScaledMm` matrix strides must cover their logical dimensions");
46 assert(
k_ % 16 == 0 &&
n_ % 16 == 0 &&
lda_ % 16 == 0 &&
ldb_ % 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_};
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();
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");
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 "
80 assert(same_device_as_a(*bias) &&
81 "`CutlassScaledMm` bias must be on the output device");