| .version 8.0
|
| .target sm_89
|
| .address_size 64
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| .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];
|
|
|
|
|
| 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;
|
| mul.lo.u32 %r11, %r8, 8;
|
|
|
|
|
|
|
| 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;
|
|
|
|
|
| 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;
|
|
|
|
|
| shr.u32 %r15, %r9, 2;
|
| and.b32 %r16, %r9, 3;
|
| add.u32 %r17, %r10, %r15;
|
| add.u32 %r18, %r17, 8;
|
| mul.lo.u32 %r19, %r16, 2;
|
| add.u32 %r20, %r11, %r19;
|
| add.u32 %r21, %r20, 1;
|
| add.u32 %r22, %r11, %r15;
|
|
|
| cvt.u64.u32 %rd3, %r3;
|
| cvt.u64.u32 %rd4, %r4;
|
| cvt.u64.u32 %rd5, %r5;
|
|
|
|
|
| cvt.u64.u32 %rd6, %r17;
|
| cvt.u64.u32 %rd7, %r18;
|
| mul.lo.u64 %rd10, %rd6, %rd3;
|
| mul.lo.u64 %rd11, %rd7, %rd3;
|
|
|
|
|
|
|
| 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;
|
|
|
|
|
|
|
|
|
| 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};
|
|
|
|
|
|
|
| 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;
|
|
|
| mov.u32 %r11, %ctaid.x;
|
| mov.u32 %r12, %ntid.x;
|
| mov.u32 %r13, %tid.x;
|
| mad.lo.u32 %r14, %r11, %r12, %r13;
|
|
|
| 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;
|
|
|
|
|
| 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;
|
|
|
|
|
| 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;
|
| }
|
|
|