SNAPKITTYWEST's picture
chore: push full sov-kernel-monster content from local build
9425aed verified
Raw
History Blame Contribute Delete
8.82 kB
.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;
}