.version 8.0 .target sm_89 .address_size 64 // C = A * B + C // // A, B, and C contain IEEE-754 binary16 values in row-major order. // Accumulation is performed in binary32 and the result is rounded to // binary16 when stored. // // gemm_f16_f32_accum: // grid = (M / 16, N / 8, 1) // block = (32, 1, 1) // Requires complete 16x8 output tiles and K divisible by 16. // // gemm_f16_f32_accum_scalar: // grid = (ceil(N / 16), ceil(M / 16), 1) // block = (16, 16, 1) // Handles arbitrary positive dimensions and all boundary tiles. .visible .entry gemm_f16_f32_accum( .param .u64 A_ptr, .param .u64 B_ptr, .param .u64 C_ptr, .param .u32 M, .param .u32 N, .param .u32 K, .param .u32 lda, .param .u32 ldb, .param .u32 ldc, .param .u32 power_state ) { .reg .pred %p<8>; .reg .b16 %h<4>; .reg .b32 %a<4>; .reg .b32 %b<2>; .reg .u32 %r<32>; .reg .u64 %rd<40>; .reg .f32 %f<4>; ld.param.u64 %rd0, [A_ptr]; ld.param.u64 %rd1, [B_ptr]; ld.param.u64 %rd2, [C_ptr]; ld.param.u32 %r0, [M]; ld.param.u32 %r1, [N]; ld.param.u32 %r2, [K]; ld.param.u32 %r3, [lda]; ld.param.u32 %r4, [ldb]; ld.param.u32 %r5, [ldc]; ld.param.u32 %r6, [power_state]; // No matrix memory may be touched while the scheduler is throttled. setp.ne.u32 %p0, %r6, 0; @%p0 ret; mov.u32 %r7, %ctaid.x; mov.u32 %r8, %ctaid.y; mov.u32 %r9, %tid.x; mul.lo.u32 %r10, %r7, 16; // tile row mul.lo.u32 %r11, %r8, 8; // tile column // This entry point accepts complete tiles only. These guards also make // an accidental direct launch fail closed instead of reading a tail. add.u32 %r12, %r10, 15; add.u32 %r13, %r11, 7; and.b32 %r14, %r2, 15; setp.ge.u32 %p1, %r12, %r0; setp.ge.u32 %p2, %r13, %r1; setp.ne.u32 %p3, %r14, 0; or.pred %p4, %p1, %p2; or.pred %p4, %p4, %p3; @%p4 ret; // Leading dimensions must cover their logical row widths. setp.lt.u32 %p1, %r3, %r2; setp.lt.u32 %p2, %r4, %r1; setp.lt.u32 %p3, %r5, %r1; or.pred %p4, %p1, %p2; or.pred %p4, %p4, %p3; @%p4 ret; // Fragment coordinates for mma.m16n8k16. shr.u32 %r15, %r9, 2; // lane group, 0..7 and.b32 %r16, %r9, 3; // lane in group, 0..3 add.u32 %r17, %r10, %r15; // accumulator row 0 add.u32 %r18, %r17, 8; // accumulator row 1 mul.lo.u32 %r19, %r16, 2; add.u32 %r20, %r11, %r19; // accumulator column 0 add.u32 %r21, %r20, 1; // accumulator column 1 add.u32 %r22, %r11, %r15; // B fragment column cvt.u64.u32 %rd3, %r3; cvt.u64.u32 %rd4, %r4; cvt.u64.u32 %rd5, %r5; // Preserve row bases, in elements, for the K loop. cvt.u64.u32 %rd6, %r17; cvt.u64.u32 %rd7, %r18; mul.lo.u64 %rd10, %rd6, %rd3; mul.lo.u64 %rd11, %rd7, %rd3; // Load the four C accumulator elements using the documented D fragment // layout: (row0,col0), (row0,col1), (row1,col0), (row1,col1). mul.lo.u64 %rd12, %rd6, %rd5; cvt.u64.u32 %rd8, %r20; add.u64 %rd12, %rd12, %rd8; shl.b64 %rd12, %rd12, 1; add.u64 %rd30, %rd2, %rd12; ld.global.u16 %h0, [%rd30]; ld.global.u16 %h1, [%rd30+2]; cvt.f32.f16 %f0, %h0; cvt.f32.f16 %f1, %h1; mul.lo.u64 %rd13, %rd7, %rd5; add.u64 %rd13, %rd13, %rd8; shl.b64 %rd13, %rd13, 1; add.u64 %rd31, %rd2, %rd13; ld.global.u16 %h0, [%rd31]; ld.global.u16 %h1, [%rd31+2]; cvt.f32.f16 %f2, %h0; cvt.f32.f16 %f3, %h1; mov.u32 %r23, 0; GEMM_MMA_K_LOOP: setp.ge.u32 %p0, %r23, %r2; @%p0 bra GEMM_MMA_STORE; // A fragment element mapping: // a0: row0, k+[0:7] pair; a1: row1, k+[0:7] pair // a2: row0, k+[8:15] pair; a3: row1, k+[8:15] pair cvt.u64.u32 %rd14, %r23; cvt.u64.u32 %rd15, %r19; add.u64 %rd16, %rd14, %rd15; add.u64 %rd17, %rd10, %rd16; shl.b64 %rd17, %rd17, 1; add.u64 %rd17, %rd0, %rd17; ld.global.u16 %h0, [%rd17]; ld.global.u16 %h1, [%rd17+2]; mov.b32 %a0, {%h0, %h1}; add.u64 %rd18, %rd11, %rd16; shl.b64 %rd18, %rd18, 1; add.u64 %rd18, %rd0, %rd18; ld.global.u16 %h0, [%rd18]; ld.global.u16 %h1, [%rd18+2]; mov.b32 %a1, {%h0, %h1}; add.u64 %rd19, %rd16, 8; add.u64 %rd20, %rd10, %rd19; shl.b64 %rd20, %rd20, 1; add.u64 %rd20, %rd0, %rd20; ld.global.u16 %h0, [%rd20]; ld.global.u16 %h1, [%rd20+2]; mov.b32 %a2, {%h0, %h1}; add.u64 %rd21, %rd11, %rd19; shl.b64 %rd21, %rd21, 1; add.u64 %rd21, %rd0, %rd21; ld.global.u16 %h0, [%rd21]; ld.global.u16 %h1, [%rd21+2]; mov.b32 %a3, {%h0, %h1}; // B is row-major in memory but the MMA B operand is logically column // major. Each lane gathers its four values before packing f16x2 regs. cvt.u64.u32 %rd22, %r22; mul.lo.u64 %rd23, %rd16, %rd4; add.u64 %rd23, %rd23, %rd22; shl.b64 %rd23, %rd23, 1; add.u64 %rd23, %rd1, %rd23; ld.global.u16 %h0, [%rd23]; add.u64 %rd24, %rd16, 1; mul.lo.u64 %rd24, %rd24, %rd4; add.u64 %rd24, %rd24, %rd22; shl.b64 %rd24, %rd24, 1; add.u64 %rd24, %rd1, %rd24; ld.global.u16 %h1, [%rd24]; mov.b32 %b0, {%h0, %h1}; add.u64 %rd25, %rd16, 8; mul.lo.u64 %rd26, %rd25, %rd4; add.u64 %rd26, %rd26, %rd22; shl.b64 %rd26, %rd26, 1; add.u64 %rd26, %rd1, %rd26; ld.global.u16 %h0, [%rd26]; add.u64 %rd27, %rd25, 1; mul.lo.u64 %rd27, %rd27, %rd4; add.u64 %rd27, %rd27, %rd22; shl.b64 %rd27, %rd27, 1; add.u64 %rd27, %rd1, %rd27; ld.global.u16 %h1, [%rd27]; mov.b32 %b1, {%h0, %h1}; mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 {%f0, %f1, %f2, %f3}, {%a0, %a1, %a2, %a3}, {%b0, %b1}, {%f0, %f1, %f2, %f3}; add.u32 %r23, %r23, 16; bra GEMM_MMA_K_LOOP; GEMM_MMA_STORE: cvt.rn.f16.f32 %h0, %f0; cvt.rn.f16.f32 %h1, %f1; st.global.u16 [%rd30], %h0; st.global.u16 [%rd30+2], %h1; cvt.rn.f16.f32 %h0, %f2; cvt.rn.f16.f32 %h1, %f3; st.global.u16 [%rd31], %h0; st.global.u16 [%rd31+2], %h1; ret; } .visible .entry gemm_f16_f32_accum_scalar( .param .u64 A_ptr, .param .u64 B_ptr, .param .u64 C_ptr, .param .u32 M, .param .u32 N, .param .u32 K, .param .u32 lda, .param .u32 ldb, .param .u32 ldc, .param .u32 power_state ) { .reg .pred %p<6>; .reg .b16 %h<3>; .reg .u32 %r<20>; .reg .u64 %rd<24>; .reg .f32 %f<4>; ld.param.u64 %rd0, [A_ptr]; ld.param.u64 %rd1, [B_ptr]; ld.param.u64 %rd2, [C_ptr]; ld.param.u32 %r0, [M]; ld.param.u32 %r1, [N]; ld.param.u32 %r2, [K]; ld.param.u32 %r3, [lda]; ld.param.u32 %r4, [ldb]; ld.param.u32 %r5, [ldc]; ld.param.u32 %r6, [power_state]; setp.ne.u32 %p0, %r6, 0; @%p0 ret; mov.u32 %r7, %ctaid.y; mov.u32 %r8, %ntid.y; mov.u32 %r9, %tid.y; mad.lo.u32 %r10, %r7, %r8, %r9; // row mov.u32 %r11, %ctaid.x; mov.u32 %r12, %ntid.x; mov.u32 %r13, %tid.x; mad.lo.u32 %r14, %r11, %r12, %r13; // column setp.ge.u32 %p1, %r10, %r0; setp.ge.u32 %p2, %r14, %r1; or.pred %p3, %p1, %p2; @%p3 ret; setp.lt.u32 %p1, %r3, %r2; setp.lt.u32 %p2, %r4, %r1; setp.lt.u32 %p3, %r5, %r1; or.pred %p4, %p1, %p2; or.pred %p4, %p4, %p3; @%p4 ret; cvt.u64.u32 %rd3, %r3; cvt.u64.u32 %rd4, %r4; cvt.u64.u32 %rd5, %r5; cvt.u64.u32 %rd6, %r10; cvt.u64.u32 %rd7, %r14; // Initialize the f32 accumulator with C[row, column]. mul.lo.u64 %rd8, %rd6, %rd5; add.u64 %rd8, %rd8, %rd7; shl.b64 %rd8, %rd8, 1; add.u64 %rd9, %rd2, %rd8; ld.global.u16 %h0, [%rd9]; cvt.f32.f16 %f0, %h0; // Keep A's row base in element units. mul.lo.u64 %rd10, %rd6, %rd3; mov.u32 %r15, 0; GEMM_SCALAR_K_LOOP: setp.ge.u32 %p0, %r15, %r2; @%p0 bra GEMM_SCALAR_STORE; cvt.u64.u32 %rd11, %r15; add.u64 %rd12, %rd10, %rd11; shl.b64 %rd12, %rd12, 1; add.u64 %rd12, %rd0, %rd12; ld.global.u16 %h0, [%rd12]; cvt.f32.f16 %f1, %h0; mul.lo.u64 %rd13, %rd11, %rd4; add.u64 %rd13, %rd13, %rd7; shl.b64 %rd13, %rd13, 1; add.u64 %rd13, %rd1, %rd13; ld.global.u16 %h1, [%rd13]; cvt.f32.f16 %f2, %h1; fma.rn.f32 %f0, %f1, %f2, %f0; add.u32 %r15, %r15, 1; bra GEMM_SCALAR_K_LOOP; GEMM_SCALAR_STORE: cvt.rn.f16.f32 %h2, %f0; st.global.u16 [%rd9], %h2; ret; }