Skip to main content

cubecl_cpp/cuda/ptx/
mma.rs

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
26// Types
27define_scalar!(A);
28define_scalar!(B);
29define_scalar!(CD);
30define_scalar!(S);
31
32// Regs per thread
33define_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}