1use cubecl_core::{
2 self as cubecl,
3 cmma::{MatrixIdent, MatrixShape},
4 ir::{
5 AddressType, ContextExt, ElemType, FloatKind, IntKind, UIntKind,
6 dialect::matrix::{LdMatrixOp, MmaManualOp, MmaManualScaledOp, StMatrixOp},
7 interfaces::{ScalarType, TypedExt},
8 types::VectorType,
9 },
10 prelude::*,
11};
12use pliron::{
13 builtin::types::{IntegerType, Signedness},
14 context::Context,
15 derive::op_interface_impl,
16 r#type::{Type, Typed, type_cast},
17 value::Value,
18};
19
20use crate::{
21 cuda::ptx::generic_to_shared,
22 shared::{lowering::LowerOp, ty::TypedExtCPP},
23 target::Cuda,
24};
25
26define_scalar!(A);
28define_scalar!(B);
29define_scalar!(CD);
30define_scalar!(S);
31
32define_size!(NA);
34define_size!(NB);
35define_size!(NCD);
36define_size!(NS);
37
38define_scalar!(RegAB);
39define_scalar!(RegCD);
40
41#[cube]
42fn mma(
43 frag_a: Vector<RegAB, NA>,
44 frag_b: Vector<RegAB, NB>,
45 frag_c: Vector<RegCD, NCD>,
46 #[comptime] k: usize,
47 #[comptime] kind: &str,
48) -> Vector<RegCD, NCD> {
49 let a_ty = ptx_mma_ty::<A>().comptime();
50 let b_ty = ptx_mma_ty::<B>().comptime();
51 let cd_ty = ptx_mma_ty::<CD>().comptime();
52
53 let out: Vector<RegCD, NCD>;
54 gpu_asm!(
55 "mma.sync.aligned.m16n8k{k}.row.col{kind}.{cd_ty}.{a_ty}.{b_ty}.{cd_ty} {d}, {a}, {b}, {c};",
56 a = in(_) frag_a, b = in(_) frag_b, c = in(_) frag_c, d = out(_) out,
57 options(nomem),
58 );
59 out
60}
61
62#[cube]
63fn scaled_mma(
64 frag_a: Vector<RegAB, NA>,
65 frag_b: Vector<RegAB, NB>,
66 frag_c: Vector<RegCD, NCD>,
67 scale_a: Vector<S, NS>,
68 scale_b: Vector<S, NS>,
69 #[comptime] k: usize,
70 #[comptime] scales_factor: usize,
71) -> Vector<RegCD, NCD> {
72 let a_ty = ptx_mma_ty::<A>().comptime();
73 let b_ty = ptx_mma_ty::<B>().comptime();
74 let cd_ty = ptx_mma_ty::<CD>().comptime();
75 let scale_ty = ptx_scale_ty::<S>().comptime();
76
77 let kind = comptime![match scales_factor {
78 1 => "mxf8f6f4",
79 2 | 4 => "mxf4nvf4",
80 _ => unreachable!(),
81 }];
82
83 let out: Vector<RegCD, NCD>;
84 gpu_asm!(
85 "mma.sync.aligned.m16n8k{k}.row.col.kind::{kind}.block_scale.scale_vec::{scales_factor}X",
86 ".{cd_ty}.{a_ty}.{b_ty}.{cd_ty}.{scale_ty} {d}, {a}, {b}, {c}, ",
87 "{scale_a}, {{0, 0}}, {scale_b}, {{0, 0}};",
88 a = in(_) frag_a, b = in(_) frag_b, c = in(_) frag_c, d = out(_) out,
89 scale_a = in(_) u32::reinterpret(scale_a), scale_b = in(_) u32::reinterpret(scale_b),
90 options(nomem),
91 );
92 out
93}
94
95#[cube]
96fn ldmatrix(
97 row: *const u32,
98 #[comptime] num: usize,
99 #[comptime] transpose: &str,
100 #[comptime] open: &str,
101 #[comptime] close: &str,
102) -> Vector<u32, NCD> {
103 let row_addr = generic_to_shared::<u32>(row);
104
105 let out: Vector<u32, NCD>;
106 gpu_asm!(
107 "ldmatrix.sync.aligned.m8n8.x{num}{transpose}.shared::cta.b16 {open}{out}{close}, [{addr}];",
108 out = out(_) out, addr = mem_in(_) row_addr, options(explicit_mem),
109 );
110 out
111}
112
113#[cube]
114fn stmatrix(
115 value: Vector<u32, NCD>,
116 row: *const u32,
117 #[comptime] num: usize,
118 #[comptime] transpose: &str,
119) {
120 let row_addr = generic_to_shared::<u32>(row);
121
122 gpu_asm!(
123 "stmatrix.sync.aligned.m8n8.x{num}{transpose}.shared::cta.b16 [{addr}], {val};",
124 addr = mem_out(_) row_addr, val = in(_) value, options(explicit_mem),
125 );
126}
127
128#[op_interface_impl]
129impl LowerOp<Cuda> for MmaManualOp {
130 fn lower(&self, scope: &Scope) -> Vec<Value> {
131 let ctx = scope.ctx();
132 let frag_a = self.registers_a(ctx);
133 let frag_b = self.registers_b(ctx);
134 let frag_c = self.registers_c(ctx);
135 let frag_d = self.registers_d(ctx);
136 let shape = self.shape(ctx).0;
137
138 let kind = if frag_a.element_ty(ctx).is_fp8_fp6_fp4(ctx)
139 || frag_b.element_ty(ctx).is_fp8_fp6_fp4(ctx)
140 {
141 ".kind::f8f6f4"
142 } else {
143 ""
144 };
145
146 let (frag_a, frag_b, frag_c) = frags_as_vectors(scope, shape, frag_a, frag_b, frag_c);
147
148 let frag_out = mma::expand(
149 scope,
150 frag_a.into(),
151 frag_b.into(),
152 frag_c.into(),
153 shape.k,
154 kind,
155 );
156 let frag_out = reinterpret_value(scope, frag_out.read_value(scope), frag_d.unwrap_ptr(ctx));
157 assign::expand_element(scope, frag_out.into(), frag_d.into());
158 vec![]
159 }
160}
161
162#[op_interface_impl]
163impl LowerOp<Cuda> for MmaManualScaledOp {
164 fn lower(&self, scope: &Scope) -> Vec<Value> {
165 let ctx = scope.ctx();
166 let frag_a = self.registers_a(ctx);
167 let frag_b = self.registers_b(ctx);
168 let frag_c = self.registers_c(ctx);
169 let frag_d = self.registers_d(ctx);
170 let scales_a = self.scales_a(ctx);
171 let scales_b = self.scales_b(ctx);
172 let scales_factor = self.scales_factor(ctx).0;
173 let shape = self.shape(ctx).0;
174
175 scope.register_value_type::<S, NS>(scales_a);
176
177 let (frag_a, frag_b, frag_c) = frags_as_vectors(scope, shape, frag_a, frag_b, frag_c);
178
179 let frag_out = scaled_mma::expand(
180 scope,
181 frag_a.into(),
182 frag_b.into(),
183 frag_c.into(),
184 scales_a.into(),
185 scales_b.into(),
186 shape.k,
187 scales_factor,
188 );
189 let frag_out = reinterpret_value(scope, frag_out.read_value(scope), frag_d.unwrap_ptr(ctx));
190 assign::expand_element(scope, frag_out.into(), frag_d.into());
191 vec![]
192 }
193}
194
195#[op_interface_impl]
196impl LowerOp<Cuda> for LdMatrixOp {
197 fn lower(&self, scope: &Scope) -> Vec<Value> {
198 let ctx = scope.ctx();
199 let row_ptr = self.ptr(ctx).into();
200 let out_arr = self.out_arr(ctx);
201 let factor = self.factor(ctx).0;
202 let trans = if self.transpose(ctx).0 { ".trans" } else { "" };
203 let (open, close) = if factor == 1 { ("{", "}") } else { ("", "") };
204
205 scope.register_size::<NCD>(factor);
206
207 let frag_out = ldmatrix::expand(scope, &row_ptr, factor, trans, open, close);
208 let frag_out =
209 reinterpret_value(scope, frag_out.read_value(scope), out_arr.unwrap_ptr(ctx));
210 assign::expand_element(scope, frag_out.into(), out_arr.into());
211 vec![]
212 }
213}
214
215#[op_interface_impl]
216impl LowerOp<Cuda> for StMatrixOp {
217 fn lower(&self, scope: &Scope) -> Vec<Value> {
218 let ctx = scope.ctx();
219 let row_ptr = self.destination(ctx).into();
220 let factor = self.factor(ctx).0;
221 let trans = if self.transpose(ctx).0 { ".trans" } else { "" };
222
223 let u32 = IntegerType::get(ctx, 32, Signedness::Unsigned).to_handle();
224 let vec_ty = VectorType::get(ctx, u32, factor);
225 let value = reinterpret_value(scope, self.registers(ctx), vec_ty.to_handle()).into();
226 scope.register_size::<NCD>(factor);
227
228 stmatrix::expand(scope, value, &row_ptr, factor, trans);
229 vec![]
230 }
231}
232
233fn num_elems(shape: MatrixShape) -> (usize, usize, usize) {
234 let a_elems = shape.num_elems(MatrixIdent::A) / 32;
235 let b_elems = shape.num_elems(MatrixIdent::B) / 32;
236 let c_elems = shape.num_elems(MatrixIdent::Accumulator) / 32;
237 (a_elems, b_elems, c_elems)
238}
239
240fn num_regs(
241 ctx: &Context,
242 shape: MatrixShape,
243 frag_a: Value,
244 frag_b: Value,
245 frag_c: Value,
246) -> (usize, usize, usize) {
247 let (a_elems, b_elems, c_elems) = num_elems(shape);
248 let a_regs = a_elems / (32 / frag_a.unpacked_size_bits(ctx));
249 let b_regs = b_elems / (32 / frag_b.unpacked_size_bits(ctx));
250 let c_regs = c_elems / (32 / frag_c.unpacked_size_bits(ctx));
251 (a_regs, b_regs, c_regs)
252}
253
254fn frags_as_vectors(
255 scope: &Scope,
256 shape: MatrixShape,
257 frag_a: Value,
258 frag_b: Value,
259 frag_c: Value,
260) -> (Value, Value, Value) {
261 let ctx = scope.ctx();
262
263 scope.register_value_type::<A, ()>(frag_a.element_ty(ctx));
264 scope.register_value_type::<B, ()>(frag_b.element_ty(ctx));
265 scope.register_value_type::<CD, ()>(frag_c.element_ty(ctx));
266
267 let (a_regs, b_regs, c_regs) = num_regs(ctx, shape, frag_a, frag_b, frag_c);
268
269 scope.register_size::<NA>(a_regs);
270 scope.register_size::<NB>(b_regs);
271 scope.register_size::<NCD>(c_regs);
272
273 let reg_ty_ab = reg_ty(ctx, frag_a).to_type(ctx);
274 let reg_ty_cd = reg_ty(ctx, frag_c).to_type(ctx);
275
276 scope.register_type::<RegAB>(reg_ty(ctx, frag_a));
277 scope.register_type::<RegCD>(reg_ty(ctx, frag_c));
278
279 let ty_a = VectorType::get(ctx, reg_ty_ab, a_regs).to_handle();
280 let ty_b = VectorType::get(ctx, reg_ty_ab, b_regs).to_handle();
281 let ty_c = VectorType::get(ctx, reg_ty_cd, c_regs).to_handle();
282
283 let frag_a = reinterpret_value(scope, frag_a, ty_a);
284 let frag_b = reinterpret_value(scope, frag_b, ty_b);
285 let frag_c = reinterpret_value(scope, frag_c, ty_c);
286
287 (frag_a, frag_b, frag_c)
288}
289
290fn reg_ty(ctx: &Context, frag: impl Typed) -> ElemType {
291 if frag.element_ty(ctx).scalar_ty(ctx).is_float32(ctx) {
292 FloatKind::F32.into()
293 } else {
294 UIntKind::U32.into()
295 }
296}
297
298#[cube]
299pub fn ptx_mma_ty<T: Scalar>() -> comptime_type!(&'static str) {
300 intrinsic!(|scope| {
301 match T::elem_type(scope) {
302 ElemType::Index => match scope.ctx().address_type() {
303 AddressType::U32 => "u32",
304 AddressType::U64 => "u64",
305 },
306 ElemType::Float(kind) => match kind {
307 FloatKind::E2M1 => "e2m1",
308 FloatKind::E2M1x2 => "e2m1",
309 FloatKind::E2M3 => "e2m3",
310 FloatKind::E3M2 => "e3m2",
311 FloatKind::E4M3 => "e4m3",
312 FloatKind::E5M2 => "e5m2",
313 FloatKind::UE8M0 => "ue8m0",
314 FloatKind::F16 => "f16",
315 FloatKind::BF16 => "bf16",
316 FloatKind::Flex32 | FloatKind::F32 => "f32",
317 FloatKind::TF32 => "tf32",
318 FloatKind::F64 => "f64",
319 },
320 ElemType::Int(kind) => match kind {
321 IntKind::I8 => "s8",
322 IntKind::I16 => "s16",
323 IntKind::I32 => "s32",
324 IntKind::I64 => "s64",
325 },
326 ElemType::UInt(kind) => match kind {
327 UIntKind::U8 => "u8",
328 UIntKind::U16 => "u16",
329 UIntKind::U32 => "u32",
330 UIntKind::U64 => "u64",
331 },
332 ElemType::Complex(_) => panic!("Complex values aren't supported by PTX MMA"),
333 ElemType::Bool => "b1",
334 }
335 })
336}
337
338#[cube]
339pub fn ptx_scale_ty<T: Scalar>() -> comptime_type!(&'static str) {
340 intrinsic!(|scope| {
341 match T::elem_type(scope) {
342 ElemType::Float(FloatKind::UE8M0) => "ue8m0",
343 ElemType::Float(FloatKind::E4M3) => "ue4m3",
344 _ => panic!("Unsupported scales type"),
345 }
346 })
347}
348
349pub fn mma_ty(ctx: &Context, elem: &dyn Type) -> &'static str {
350 let elem = type_cast::<dyn ScalarType>(elem).unwrap();
351 match elem.elem_type(ctx) {
352 ElemType::Index => match ctx.address_type() {
353 AddressType::U32 => "u32",
354 AddressType::U64 => "u64",
355 },
356 ElemType::Float(kind) => match kind {
357 FloatKind::E2M1 => "e2m1",
358 FloatKind::E2M1x2 => "e2m1",
359 FloatKind::E2M3 => "e2m3",
360 FloatKind::E3M2 => "e3m2",
361 FloatKind::E4M3 => "e4m3",
362 FloatKind::E5M2 => "e5m2",
363 FloatKind::UE8M0 => "ue8m0",
364 FloatKind::F16 => "f16",
365 FloatKind::BF16 => "bf16",
366 FloatKind::Flex32 | FloatKind::F32 => "f32",
367 FloatKind::TF32 => "tf32",
368 FloatKind::F64 => "f64",
369 },
370 ElemType::Int(kind) => match kind {
371 IntKind::I8 => "s8",
372 IntKind::I16 => "s16",
373 IntKind::I32 => "s32",
374 IntKind::I64 => "s64",
375 },
376 ElemType::UInt(kind) => match kind {
377 UIntKind::U8 => "u8",
378 UIntKind::U16 => "u16",
379 UIntKind::U32 => "u32",
380 UIntKind::U64 => "u64",
381 },
382 ElemType::Complex(_) => panic!("Complex values aren't supported by PTX MMA"),
383 ElemType::Bool => "b1",
384 }
385}