Download rtx/src/cuda/gemm.ptx from Snapkitty/sov-kernel-monster: direct link, hf CLI and curl.
- Browser
- Download file 8.82 kB
-
https://huggingface.co/Snapkitty/sov-kernel-monster/resolve/main/rtx/src/cuda/gemm.ptx
- Command line
-
hf download hf://Snapkitty/sov-kernel-monster/rtx/src/cuda/gemm.ptx
-
curl -L -o gemm.ptx https://huggingface.co/Snapkitty/sov-kernel-monster/resolve/main/rtx/src/cuda/gemm.ptx
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; | |
| } | |