//
// Copyright (c) 2026 Lukasz Szpakowski
//
// This Source Code Form is subject to the terms of the Mozilla Public
// License, v. 2.0. If a copy of the MPL was not distributed with this
// file, You can obtain one at https://mozilla.org/MPL/2.0/.
//
.version 7.4
.target sm_52
.address_size 64
.visible .entry ptx_mul_a_b(
.param .u64 a,
.param .u64 b,
.param .u64 c,
.param .u64 n,
.param .u64 m,
.param .u64 l)
{
.reg .b64 %aik; // 2
.reg .b64 %bkj; // 2 + 2 = 4
.reg .b64 %cij; // 4 + 2 = 6
.reg .b64 %tmpa; // 6 + 2 = 8
.reg .b64 %tmpa2; // 8 + 2 = 12
.reg .f32 %ar<4>; // 12 + 4 = 16
.reg .f32 %br<4>; // 16 + 4 = 20
.reg .f32 %cr<16>; // 20 + 16 = 36
.reg .b32 %n; // 36 + 1 = 37
.reg .b32 %m; // 37 + 1 = 38
.reg .b32 %l; // 38 + 1 = 39
.reg .b32 %i; // 39 + 1 = 40
.reg .b32 %j; // 40 + 1 = 41
.reg .b32 %k; // 41 + 1 = 42
.reg .b32 %asti; // 42 + 1 = 43
.reg .b32 %bstj; // 43 + 1 = 44
.reg .b32 %ti; // 44 + 1 = 45
.reg .b32 %tj; // 45 + 1 = 46
.reg .b32 %tk; // 46 + 1 = 47
.reg .b32 %tmp; // 47 + 1 = 48
.reg .b32 %tmp2; // 48 + 1 = 49
.reg .pred %p; // 49 + 1 = 50
.reg .pred %p2; // 50 + 1 = 51
.shared .align 16 .u8 as[16384];
.shared .align 16 .u8 bs[16384];
ld.param.u64 %aik, [a];
ld.param.u64 %bkj, [b];
ld.param.u64 %cij, [c];
ld.param.u64 %tmpa, [n];
cvt.u32.u64 %n, %tmpa;
ld.param.u64 %tmpa, [m];
cvt.u32.u64 %m, %tmpa;
ld.param.u64 %tmpa, [l];
cvt.u32.u64 %l, %tmpa;
cvta.to.global.u64 %aik, %aik;
cvta.to.global.u64 %bkj, %bkj;
cvta.to.global.u64 %cij, %cij;
// i
mov.u32 %tmp, %ntid.y;
mov.u32 %tmp2, %ctaid.y;
mul.lo.u32 %i, %tmp, %tmp2;
mov.u32 %tmp, %tid.y;
add.u32 %i, %i, %tmp;
shl.b32 %i, %i, 2;
// j
mov.u32 %tmp, %ntid.x;
mov.u32 %tmp2, %ctaid.x;
mul.lo.u32 %j, %tmp, %tmp2;
mov.u32 %tmp, %tid.x;
add.u32 %j, %j, %tmp;
shl.b32 %j, %j, 2;
mov.u32 %ti, %tid.y; // ti and ik
mov.u32 %tj, %tid.x; // tj and jk
// a[i][k]
mul.wide.u32 %tmpa, %l, %i;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %aik, %aik, %tmpa;
// b[k][j]
cvt.u64.u32 %tmpa, %j;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %bkj, %bkj, %tmpa;
// c[i][j]
mul.wide.u32 %tmpa, %m, %i;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %cij, %cij, %tmpa;
cvt.u64.u32 %tmpa, %j;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %cij, %cij, %tmpa;
// as[ti]
mov.u32 %asti, as;
mov.u32 %tmp, %ti;
shl.b32 %tmp, %tmp, 5 + 4;
add.u32 %asti, %asti, %tmp;
// bs[tj]
mov.u32 %bstj, bs;
mov.u32 %tmp, %tj;
shl.b32 %tmp, %tmp, 5 + 4;
add.u32 %bstj, %bstj, %tmp;
// cr
mov.f32 %cr0, 0.0;
mov.f32 %cr1, 0.0;
mov.f32 %cr2, 0.0;
mov.f32 %cr3, 0.0;
mov.f32 %cr4, 0.0;
mov.f32 %cr5, 0.0;
mov.f32 %cr6, 0.0;
mov.f32 %cr7, 0.0;
mov.f32 %cr8, 0.0;
mov.f32 %cr9, 0.0;
mov.f32 %cr10, 0.0;
mov.f32 %cr11, 0.0;
mov.f32 %cr12, 0.0;
mov.f32 %cr13, 0.0;
mov.f32 %cr14, 0.0;
mov.f32 %cr15, 0.0;
mov.u32 %k, 0;
loop: setp.ge.u32 %p, %k, %l;
@%p bra eloop;
// k + tj < l
mov.u32 %tmp, %k;
add.u32 %tmp, %tmp, %tj;
setp.lt.u32 %p, %tmp, %l;
// a[i][k + tj]
mov.u64 %tmpa, %aik;
cvt.u64.u32 %tmpa2, %tj;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 0 < n
mov.u32 %tmp, %i;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar0, 0.0; // ar0 = 0.0
@%p2 ld.global.f32 %ar0, [%tmpa]; // ar0 = a[i + 0][k + tj]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 1 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar1, 0.0; // ar1 = 0.0
@%p2 ld.global.f32 %ar1, [%tmpa]; // ar1 = a[i + 1][k + tj]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 2 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar2, 0.0; // ar2 = 0.0
@%p2 ld.global.f32 %ar2, [%tmpa]; // ar2 = a[i + 2][k + tj]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 3 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar3, 0.0; // ar3 = 0.0
@%p2 ld.global.f32 %ar3, [%tmpa]; // ar3 = a[i + 3][k + tj]
//cvt.u64.u32 %tmpa2, %l;
//shl.b64 %tmpa2, %tmpa2, 2;
//add.u64 %tmpa, %tmpa, %tmpa2;
// as[ti][tj + ik]
mov.u32 %tmp, %asti;
mov.u32 %tmp2, %tj;
add.u32 %tmp2, %tmp2, %ti;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
st.shared.v4.f32 [%tmp], { %ar0, %ar1, %ar2, %ar3 };
// k + ti < l
mov.u32 %tmp, %k;
add.u32 %tmp, %tmp, %ti;
setp.lt.u32 %p, %tmp, %l;
// b[k + ti][j]
mov.u64 %tmpa, %bkj;
mul.wide.u32 %tmpa2, %m, %ti;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br0, 0.0; // br0 = 0.0
@%p2 ld.global.f32 %br0, [%tmpa]; // br0 = b[k + ti][j + 0]
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br1, 0.0; // br1 = 0.0
@%p2 ld.global.f32 %br1, [%tmpa + (1 << 2)]; // br1 = b[k + ti][j + 1]
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br2, 0.0; // br2 = 0.0
@%p2 ld.global.f32 %br2, [%tmpa + (2 << 2)]; // br2 = b[k + ti][j + 2]
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br3, 0.0; // br3 = 0.0
@%p2 ld.global.f32 %br3, [%tmpa + (3 << 2)]; // br3 = b[k + ti][j + 3]
// bs[tj][ti + jk]
mov.u32 %tmp, %bstj;
mov.u32 %tmp2, %ti;
add.u32 %tmp2, %tmp2, %tj;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
st.shared.v4.f32 [%tmp], { %br0, %br1, %br2, %br3 };
bar.sync 0;
mov.u32 %tk, 0;
loop2: setp.ge.u32 %p, %tk, 32;
@%p bra eloop2;
// as[ti][tk + ik]
mov.u32 %tmp, %asti;
mov.u32 %tmp2, %tk;
add.u32 %tmp2, %tmp2, %ti;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
ld.shared.v4.f32 { %ar0, %ar1, %ar2, %ar3 }, [%tmp];
// bs[tj][ti + jk]
mov.u32 %tmp, %bstj;
mov.u32 %tmp2, %tk;
add.u32 %tmp2, %tmp2, %tj;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
ld.shared.v4.f32 { %br0, %br1, %br2, %br3 }, [%tmp];
fma.rn.f32 %cr0, %ar0, %br0, %cr0;
fma.rn.f32 %cr1, %ar0, %br1, %cr1;
fma.rn.f32 %cr2, %ar0, %br2, %cr2;
fma.rn.f32 %cr3, %ar0, %br3, %cr3;
fma.rn.f32 %cr4, %ar1, %br0, %cr4;
fma.rn.f32 %cr5, %ar1, %br1, %cr5;
fma.rn.f32 %cr6, %ar1, %br2, %cr6;
fma.rn.f32 %cr7, %ar1, %br3, %cr7;
fma.rn.f32 %cr8, %ar2, %br0, %cr8;
fma.rn.f32 %cr9, %ar2, %br1, %cr9;
fma.rn.f32 %cr10, %ar2, %br2, %cr10;
fma.rn.f32 %cr11, %ar2, %br3, %cr11;
fma.rn.f32 %cr12, %ar3, %br0, %cr12;
fma.rn.f32 %cr13, %ar3, %br1, %cr13;
fma.rn.f32 %cr14, %ar3, %br2, %cr14;
fma.rn.f32 %cr15, %ar3, %br3, %cr15;
add.u32 %tk, %tk, 1;
bra loop2;
eloop2: bar.sync 0;
add.u64 %aik, %aik, 32 << 2;
cvt.u64.u32 %tmpa, %m;
shl.b64 %tmpa, %tmpa, 5 + 2;
add.u64 %bkj, %bkj, %tmpa;
add.u32 %k, %k, 32;
bra loop;
eloop: mov.u64 %tmpa, %cij;
// i + 0 < n
mov.u32 %tmp, %i;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr0; // c[i + 0][j + 0] = cr0
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr1; // c[i + 0][j + 1] = cr1
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr2; // c[i + 0][j + 2] = cr2
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr3; // c[i + 0][j + 3] = cr3
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 1 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr4; // c[i + 1][j + 0] = cr4
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr5; // c[i + 1][j + 1] = cr5
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr6; // c[i + 1][j + 2] = cr6
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr7; // c[i + 1][j + 3] = cr7
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 2 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr8; // c[i + 2][j + 0] = cr8
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr9; // c[i + 2][j + 1] = cr9
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr10;// c[i + 2][j + 2] = cr10
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr11;// c[i + 2][j + 3] = cr11
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 3 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr12; // c[i + 3][j + 0] = cr12
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr13;// c[i + 3][j + 1] = cr13
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr14;// c[i + 3][j + 2] = cr14
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr15;// c[i + 3][j + 3] = cr15
//cvt.u64.u32 %tmpa2, %m;
//shl.b64 %tmpa2, %tmpa2, 2;
//add.u64 %tmpa, %tmpa, %tmpa2;
ret;
}
.visible .entry ptx_mul_at_b(
.param .u64 a,
.param .u64 b,
.param .u64 c,
.param .u64 n,
.param .u64 m,
.param .u64 l)
{
.reg .b64 %aki; // 2
.reg .b64 %bkj; // 2 + 2 = 4
.reg .b64 %cij; // 4 + 2 = 6
.reg .b64 %tmpa; // 6 + 2 = 8
.reg .b64 %tmpa2; // 8 + 2 = 12
.reg .f32 %ar<4>; // 12 + 4 = 16
.reg .f32 %br<4>; // 16 + 4 = 20
.reg .f32 %cr<16>; // 20 + 16 = 36
.reg .b32 %n; // 36 + 1 = 37
.reg .b32 %m; // 37 + 1 = 38
.reg .b32 %l; // 38 + 1 = 39
.reg .b32 %i; // 39 + 1 = 40
.reg .b32 %j; // 40 + 1 = 41
.reg .b32 %k; // 41 + 1 = 42
.reg .b32 %asti; // 42 + 1 = 43
.reg .b32 %bstj; // 43 + 1 = 44
.reg .b32 %ti; // 44 + 1 = 45
.reg .b32 %tj; // 45 + 1 = 46
.reg .b32 %tk; // 46 + 1 = 47
.reg .b32 %tmp; // 47 + 1 = 48
.reg .b32 %tmp2; // 48 + 1 = 49
.reg .pred %p; // 49 + 1 = 50
.reg .pred %p2; // 50 + 1 = 51
.shared .align 16 .u8 as[16384];
.shared .align 16 .u8 bs[16384];
ld.param.u64 %aki, [a];
ld.param.u64 %bkj, [b];
ld.param.u64 %cij, [c];
ld.param.u64 %tmpa, [n];
cvt.u32.u64 %n, %tmpa;
ld.param.u64 %tmpa, [m];
cvt.u32.u64 %m, %tmpa;
ld.param.u64 %tmpa, [l];
cvt.u32.u64 %l, %tmpa;
cvta.to.global.u64 %aki, %aki;
cvta.to.global.u64 %bkj, %bkj;
cvta.to.global.u64 %cij, %cij;
// i
mov.u32 %tmp, %ntid.x;
mov.u32 %tmp2, %ctaid.x;
mul.lo.u32 %i, %tmp, %tmp2;
mov.u32 %tmp, %tid.x;
add.u32 %i, %i, %tmp;
shl.b32 %i, %i, 2;
// j
mov.u32 %tmp, %ntid.y;
mov.u32 %tmp2, %ctaid.y;
mul.lo.u32 %j, %tmp, %tmp2;
mov.u32 %tmp, %tid.y;
add.u32 %j, %j, %tmp;
shl.b32 %j, %j, 2;
mov.u32 %ti, %tid.x; // ti and ik
mov.u32 %tj, %tid.y; // tj and jk
// a[k][i]
cvt.u64.u32 %tmpa, %i;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %aki, %aki, %tmpa;
// b[k][j]
cvt.u64.u32 %tmpa, %j;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %bkj, %bkj, %tmpa;
// c[i][j]
mul.wide.u32 %tmpa, %m, %i;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %cij, %cij, %tmpa;
cvt.u64.u32 %tmpa, %j;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %cij, %cij, %tmpa;
// as[ti]
mov.u32 %asti, as;
mov.u32 %tmp, %ti;
shl.b32 %tmp, %tmp, 5 + 4;
add.u32 %asti, %asti, %tmp;
// bs[tj]
mov.u32 %bstj, bs;
mov.u32 %tmp, %tj;
shl.b32 %tmp, %tmp, 5 + 4;
add.u32 %bstj, %bstj, %tmp;
// cr
mov.f32 %cr0, 0.0;
mov.f32 %cr1, 0.0;
mov.f32 %cr2, 0.0;
mov.f32 %cr3, 0.0;
mov.f32 %cr4, 0.0;
mov.f32 %cr5, 0.0;
mov.f32 %cr6, 0.0;
mov.f32 %cr7, 0.0;
mov.f32 %cr8, 0.0;
mov.f32 %cr9, 0.0;
mov.f32 %cr10, 0.0;
mov.f32 %cr11, 0.0;
mov.f32 %cr12, 0.0;
mov.f32 %cr13, 0.0;
mov.f32 %cr14, 0.0;
mov.f32 %cr15, 0.0;
mov.u32 %k, 0;
loop: setp.ge.u32 %p, %k, %l;
@%p bra eloop;
// k + tj < l
mov.u32 %tmp, %k;
add.u32 %tmp, %tmp, %tj;
setp.lt.u32 %p, %tmp, %l;
// a[k + tj][i]
mov.u64 %tmpa, %aki;
mul.wide.u32 %tmpa2, %n, %tj;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 0 < n
mov.u32 %tmp, %i;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar0, 0.0; // ar0 = 0.0
@%p2 ld.global.f32 %ar0, [%tmpa]; // ar0 = a[k + tj][i + 0]
// i + 1 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar1, 0.0; // ar1 = 0.0
@%p2 ld.global.f32 %ar1, [%tmpa + (1 << 2)]; // ar1 = a[k + tj][i + 1]
// i + 2 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar2, 0.0; // ar2 = 0.0
@%p2 ld.global.f32 %ar2, [%tmpa + (2 << 2)]; // ar2 = a[k + tj][i + 2]
// i + 3 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar3, 0.0; // ar3 = 0.0
@%p2 ld.global.f32 %ar3, [%tmpa + (3 << 2)]; // ar3 = a[k + tj][i + 3]
// as[ti][tj + ik]
mov.u32 %tmp, %asti;
mov.u32 %tmp2, %tj;
add.u32 %tmp2, %tmp2, %ti;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
st.shared.v4.f32 [%tmp], { %ar0, %ar1, %ar2, %ar3 };
// k + ti < l
mov.u32 %tmp, %k;
add.u32 %tmp, %tmp, %ti;
setp.lt.u32 %p, %tmp, %l;
// b[k + ti][j]
mov.u64 %tmpa, %bkj;
mul.wide.u32 %tmpa2, %m, %ti;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br0, 0.0; // br0 = 0.0
@%p2 ld.global.f32 %br0, [%tmpa]; // br0 = b[k + ti][j + 0]
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br1, 0.0; // br1 = 0.0
@%p2 ld.global.f32 %br1, [%tmpa + (1 << 2)]; // br1 = b[k + ti][j + 1]
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br2, 0.0; // br2 = 0.0
@%p2 ld.global.f32 %br2, [%tmpa + (2 << 2)]; // br2 = b[k + ti][j + 2]
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br3, 0.0; // br3 = 0.0
@%p2 ld.global.f32 %br3, [%tmpa + (3 << 2)]; // br3 = b[k + ti][j + 3]
// bs[tj][ti + jk]
mov.u32 %tmp, %bstj;
mov.u32 %tmp2, %ti;
add.u32 %tmp2, %tmp2, %tj;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
st.shared.v4.f32 [%tmp], { %br0, %br1, %br2, %br3 };
bar.sync 0;
mov.u32 %tk, 0;
loop2: setp.ge.u32 %p, %tk, 32;
@%p bra eloop2;
// as[ti][tk + ik]
mov.u32 %tmp, %asti;
mov.u32 %tmp2, %tk;
add.u32 %tmp2, %tmp2, %ti;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
ld.shared.v4.f32 { %ar0, %ar1, %ar2, %ar3 }, [%tmp];
// bs[tj][ti + jk]
mov.u32 %tmp, %bstj;
mov.u32 %tmp2, %tk;
add.u32 %tmp2, %tmp2, %tj;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
ld.shared.v4.f32 { %br0, %br1, %br2, %br3 }, [%tmp];
fma.rn.f32 %cr0, %ar0, %br0, %cr0;
fma.rn.f32 %cr1, %ar0, %br1, %cr1;
fma.rn.f32 %cr2, %ar0, %br2, %cr2;
fma.rn.f32 %cr3, %ar0, %br3, %cr3;
fma.rn.f32 %cr4, %ar1, %br0, %cr4;
fma.rn.f32 %cr5, %ar1, %br1, %cr5;
fma.rn.f32 %cr6, %ar1, %br2, %cr6;
fma.rn.f32 %cr7, %ar1, %br3, %cr7;
fma.rn.f32 %cr8, %ar2, %br0, %cr8;
fma.rn.f32 %cr9, %ar2, %br1, %cr9;
fma.rn.f32 %cr10, %ar2, %br2, %cr10;
fma.rn.f32 %cr11, %ar2, %br3, %cr11;
fma.rn.f32 %cr12, %ar3, %br0, %cr12;
fma.rn.f32 %cr13, %ar3, %br1, %cr13;
fma.rn.f32 %cr14, %ar3, %br2, %cr14;
fma.rn.f32 %cr15, %ar3, %br3, %cr15;
add.u32 %tk, %tk, 1;
bra loop2;
eloop2: bar.sync 0;
cvt.u64.u32 %tmpa, %n;
shl.b64 %tmpa, %tmpa, 5 + 2;
add.u64 %aki, %aki, %tmpa;
cvt.u64.u32 %tmpa, %m;
shl.b64 %tmpa, %tmpa, 5 + 2;
add.u64 %bkj, %bkj, %tmpa;
add.u32 %k, %k, 32;
bra loop;
eloop: mov.u64 %tmpa, %cij;
// i + 0 < n
mov.u32 %tmp, %i;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr0; // c[i + 0][j + 0] = cr0
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr1; // c[i + 0][j + 1] = cr1
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr2; // c[i + 0][j + 2] = cr2
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr3; // c[i + 0][j + 3] = cr3
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 1 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr4; // c[i + 1][j + 0] = cr4
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr5; // c[i + 1][j + 1] = cr5
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr6; // c[i + 1][j + 2] = cr6
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr7; // c[i + 1][j + 3] = cr7
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 2 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr8; // c[i + 2][j + 0] = cr8
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr9; // c[i + 2][j + 1] = cr9
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr10;// c[i + 2][j + 2] = cr10
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr11;// c[i + 2][j + 3] = cr11
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 3 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr12; // c[i + 3][j + 0] = cr12
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr13;// c[i + 3][j + 1] = cr13
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr14;// c[i + 3][j + 2] = cr14
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr15;// c[i + 3][j + 3] = cr15
//cvt.u64.u32 %tmpa2, %m;
//shl.b64 %tmpa2, %tmpa2, 2;
//add.u64 %tmpa, %tmpa, %tmpa2;
ret;
}
.visible .entry ptx_mul_a_bt(
.param .u64 a,
.param .u64 b,
.param .u64 c,
.param .u64 n,
.param .u64 m,
.param .u64 l)
{
.reg .b64 %aik; // 2
.reg .b64 %bjk; // 2 + 2 = 4
.reg .b64 %cij; // 4 + 2 = 6
.reg .b64 %tmpa; // 6 + 2 = 8
.reg .b64 %tmpa2; // 8 + 2 = 12
.reg .f32 %ar<4>; // 12 + 4 = 16
.reg .f32 %br<4>; // 16 + 4 = 20
.reg .f32 %cr<16>; // 20 + 16 = 36
.reg .b32 %n; // 36 + 1 = 37
.reg .b32 %m; // 37 + 1 = 38
.reg .b32 %l; // 38 + 1 = 39
.reg .b32 %i; // 39 + 1 = 40
.reg .b32 %j; // 40 + 1 = 41
.reg .b32 %k; // 41 + 1 = 42
.reg .b32 %asti; // 42 + 1 = 43
.reg .b32 %bstj; // 43 + 1 = 44
.reg .b32 %ti; // 44 + 1 = 45
.reg .b32 %tj; // 45 + 1 = 46
.reg .b32 %tk; // 46 + 1 = 47
.reg .b32 %tmp; // 47 + 1 = 48
.reg .b32 %tmp2; // 48 + 1 = 49
.reg .pred %p; // 49 + 1 = 50
.reg .pred %p2; // 50 + 1 = 51
.shared .align 16 .u8 as[16384];
.shared .align 16 .u8 bs[16384];
ld.param.u64 %aik, [a];
ld.param.u64 %bjk, [b];
ld.param.u64 %cij, [c];
ld.param.u64 %tmpa, [n];
cvt.u32.u64 %n, %tmpa;
ld.param.u64 %tmpa, [m];
cvt.u32.u64 %m, %tmpa;
ld.param.u64 %tmpa, [l];
cvt.u32.u64 %l, %tmpa;
cvta.to.global.u64 %aik, %aik;
cvta.to.global.u64 %bjk, %bjk;
cvta.to.global.u64 %cij, %cij;
// i
mov.u32 %tmp, %ntid.y;
mov.u32 %tmp2, %ctaid.y;
mul.lo.u32 %i, %tmp, %tmp2;
mov.u32 %tmp, %tid.y;
add.u32 %i, %i, %tmp;
shl.b32 %i, %i, 2;
// j
mov.u32 %tmp, %ntid.x;
mov.u32 %tmp2, %ctaid.x;
mul.lo.u32 %j, %tmp, %tmp2;
mov.u32 %tmp, %tid.x;
add.u32 %j, %j, %tmp;
shl.b32 %j, %j, 2;
mov.u32 %ti, %tid.y; // ti and ik
mov.u32 %tj, %tid.x; // tj and jk
// a[i][k]
mul.wide.u32 %tmpa, %l, %i;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %aik, %aik, %tmpa;
// b[j][k]
mul.wide.u32 %tmpa, %l, %j;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %bjk, %bjk, %tmpa;
// c[i][j]
mul.wide.u32 %tmpa, %m, %i;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %cij, %cij, %tmpa;
cvt.u64.u32 %tmpa, %j;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %cij, %cij, %tmpa;
// as[ti]
mov.u32 %asti, as;
mov.u32 %tmp, %ti;
shl.b32 %tmp, %tmp, 5 + 4;
add.u32 %asti, %asti, %tmp;
// bs[tj]
mov.u32 %bstj, bs;
mov.u32 %tmp, %tj;
shl.b32 %tmp, %tmp, 5 + 4;
add.u32 %bstj, %bstj, %tmp;
// cr
mov.f32 %cr0, 0.0;
mov.f32 %cr1, 0.0;
mov.f32 %cr2, 0.0;
mov.f32 %cr3, 0.0;
mov.f32 %cr4, 0.0;
mov.f32 %cr5, 0.0;
mov.f32 %cr6, 0.0;
mov.f32 %cr7, 0.0;
mov.f32 %cr8, 0.0;
mov.f32 %cr9, 0.0;
mov.f32 %cr10, 0.0;
mov.f32 %cr11, 0.0;
mov.f32 %cr12, 0.0;
mov.f32 %cr13, 0.0;
mov.f32 %cr14, 0.0;
mov.f32 %cr15, 0.0;
mov.u32 %k, 0;
loop: setp.ge.u32 %p, %k, %l;
@%p bra eloop;
// k + tj < l
mov.u32 %tmp, %k;
add.u32 %tmp, %tmp, %tj;
setp.lt.u32 %p, %tmp, %l;
// a[i][k + tj]
mov.u64 %tmpa, %aik;
cvt.u64.u32 %tmpa2, %tj;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 0 < n
mov.u32 %tmp, %i;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar0, 0.0; // ar0 = 0.0
@%p2 ld.global.f32 %ar0, [%tmpa]; // ar0 = a[i + 0][k + tj]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 1 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar1, 0.0; // ar1 = 0.0
@%p2 ld.global.f32 %ar1, [%tmpa]; // ar1 = a[i + 1][k + tj]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 2 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar2, 0.0; // ar2 = 0.0
@%p2 ld.global.f32 %ar2, [%tmpa]; // ar2 = a[i + 2][k + tj]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 3 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar3, 0.0; // ar3 = 0.0
@%p2 ld.global.f32 %ar3, [%tmpa]; // ar3 = a[i + 3][k + tj]
//cvt.u64.u32 %tmpa2, %l;
//shl.b64 %tmpa2, %tmpa2, 2;
//add.u64 %tmpa, %tmpa, %tmpa2;
// as[ti][tj + ik]
mov.u32 %tmp, %asti;
mov.u32 %tmp2, %tj;
add.u32 %tmp2, %tmp2, %ti;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
st.shared.v4.f32 [%tmp], { %ar0, %ar1, %ar2, %ar3 };
// k + ti < l
mov.u32 %tmp, %k;
add.u32 %tmp, %tmp, %ti;
setp.lt.u32 %p, %tmp, %l;
// b[j][k + ti]
mov.u64 %tmpa, %bjk;
cvt.u64.u32 %tmpa2, %ti;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br0, 0.0; // br0 = 0.0
@%p2 ld.global.f32 %br0, [%tmpa]; // br0 = b[j + 0][k + ti]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br1, 0.0; // br1 = 0.0
@%p2 ld.global.f32 %br1, [%tmpa]; // br1 = b[j + 1][k + ti]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br2, 0.0; // br2 = 0.0
@%p2 ld.global.f32 %br2, [%tmpa]; // br2 = b[j + 2][k + ti]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br3, 0.0; // br3 = 0.0
@%p2 ld.global.f32 %br3, [%tmpa]; // br3 = b[j + 3][k + ti]
//cvt.u64.u32 %tmpa2, %l;
//shl.b64 %tmpa2, %tmpa2, 2;
//add.u64 %tmpa, %tmpa, %tmpa2;
// bs[tj][ti + jk]
mov.u32 %tmp, %bstj;
mov.u32 %tmp2, %ti;
add.u32 %tmp2, %tmp2, %tj;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
st.shared.v4.f32 [%tmp], { %br0, %br1, %br2, %br3 };
bar.sync 0;
mov.u32 %tk, 0;
loop2: setp.ge.u32 %p, %tk, 32;
@%p bra eloop2;
// as[ti][tk + ik]
mov.u32 %tmp, %asti;
mov.u32 %tmp2, %tk;
add.u32 %tmp2, %tmp2, %ti;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
ld.shared.v4.f32 { %ar0, %ar1, %ar2, %ar3 }, [%tmp];
// bs[tj][ti + jk]
mov.u32 %tmp, %bstj;
mov.u32 %tmp2, %tk;
add.u32 %tmp2, %tmp2, %tj;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
ld.shared.v4.f32 { %br0, %br1, %br2, %br3 }, [%tmp];
fma.rn.f32 %cr0, %ar0, %br0, %cr0;
fma.rn.f32 %cr1, %ar0, %br1, %cr1;
fma.rn.f32 %cr2, %ar0, %br2, %cr2;
fma.rn.f32 %cr3, %ar0, %br3, %cr3;
fma.rn.f32 %cr4, %ar1, %br0, %cr4;
fma.rn.f32 %cr5, %ar1, %br1, %cr5;
fma.rn.f32 %cr6, %ar1, %br2, %cr6;
fma.rn.f32 %cr7, %ar1, %br3, %cr7;
fma.rn.f32 %cr8, %ar2, %br0, %cr8;
fma.rn.f32 %cr9, %ar2, %br1, %cr9;
fma.rn.f32 %cr10, %ar2, %br2, %cr10;
fma.rn.f32 %cr11, %ar2, %br3, %cr11;
fma.rn.f32 %cr12, %ar3, %br0, %cr12;
fma.rn.f32 %cr13, %ar3, %br1, %cr13;
fma.rn.f32 %cr14, %ar3, %br2, %cr14;
fma.rn.f32 %cr15, %ar3, %br3, %cr15;
add.u32 %tk, %tk, 1;
bra loop2;
eloop2: bar.sync 0;
add.u64 %aik, %aik, 32 << 2;
add.u64 %bjk, %bjk, 32 << 2;
add.u32 %k, %k, 32;
bra loop;
eloop: mov.u64 %tmpa, %cij;
// i + 0 < n
mov.u32 %tmp, %i;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr0; // c[i + 0][j + 0] = cr0
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr1; // c[i + 0][j + 1] = cr1
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr2; // c[i + 0][j + 2] = cr2
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr3; // c[i + 0][j + 3] = cr3
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 1 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr4; // c[i + 1][j + 0] = cr4
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr5; // c[i + 1][j + 1] = cr5
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr6; // c[i + 1][j + 2] = cr6
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr7; // c[i + 1][j + 3] = cr7
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 2 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr8; // c[i + 2][j + 0] = cr8
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr9; // c[i + 2][j + 1] = cr9
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr10;// c[i + 2][j + 2] = cr10
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr11;// c[i + 2][j + 3] = cr11
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 3 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr12; // c[i + 3][j + 0] = cr12
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr13;// c[i + 3][j + 1] = cr13
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr14;// c[i + 3][j + 2] = cr14
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr15;// c[i + 3][j + 3] = cr15
//cvt.u64.u32 %tmpa2, %m;
//shl.b64 %tmpa2, %tmpa2, 2;
//add.u64 %tmpa, %tmpa, %tmpa2;
ret;
}
.visible .entry ptx_mul_at_bt(
.param .u64 a,
.param .u64 b,
.param .u64 c,
.param .u64 n,
.param .u64 m,
.param .u64 l)
{
.reg .b64 %aki; // 2
.reg .b64 %bjk; // 2 + 2 = 4
.reg .b64 %cij; // 4 + 2 = 6
.reg .b64 %tmpa; // 6 + 2 = 8
.reg .b64 %tmpa2; // 8 + 2 = 12
.reg .f32 %ar<4>; // 12 + 4 = 16
.reg .f32 %br<4>; // 16 + 4 = 20
.reg .f32 %cr<16>; // 20 + 16 = 36
.reg .b32 %n; // 36 + 1 = 37
.reg .b32 %m; // 37 + 1 = 38
.reg .b32 %l; // 38 + 1 = 39
.reg .b32 %i; // 39 + 1 = 40
.reg .b32 %j; // 40 + 1 = 41
.reg .b32 %k; // 41 + 1 = 42
.reg .b32 %asti; // 42 + 1 = 43
.reg .b32 %bstj; // 43 + 1 = 44
.reg .b32 %ti; // 44 + 1 = 45
.reg .b32 %tj; // 45 + 1 = 46
.reg .b32 %tk; // 46 + 1 = 47
.reg .b32 %tmp; // 47 + 1 = 48
.reg .b32 %tmp2; // 48 + 1 = 49
.reg .pred %p; // 49 + 1 = 50
.reg .pred %p2; // 50 + 1 = 51
.shared .align 16 .u8 as[16384];
.shared .align 16 .u8 bs[16384];
ld.param.u64 %aki, [a];
ld.param.u64 %bjk, [b];
ld.param.u64 %cij, [c];
ld.param.u64 %tmpa, [n];
cvt.u32.u64 %n, %tmpa;
ld.param.u64 %tmpa, [m];
cvt.u32.u64 %m, %tmpa;
ld.param.u64 %tmpa, [l];
cvt.u32.u64 %l, %tmpa;
cvta.to.global.u64 %aki, %aki;
cvta.to.global.u64 %bjk, %bjk;
cvta.to.global.u64 %cij, %cij;
// i
mov.u32 %tmp, %ntid.x;
mov.u32 %tmp2, %ctaid.x;
mul.lo.u32 %i, %tmp, %tmp2;
mov.u32 %tmp, %tid.x;
add.u32 %i, %i, %tmp;
shl.b32 %i, %i, 2;
// j
mov.u32 %tmp, %ntid.y;
mov.u32 %tmp2, %ctaid.y;
mul.lo.u32 %j, %tmp, %tmp2;
mov.u32 %tmp, %tid.y;
add.u32 %j, %j, %tmp;
shl.b32 %j, %j, 2;
mov.u32 %ti, %tid.x; // ti and ik
mov.u32 %tj, %tid.y; // tj and jk
// a[k][i]
cvt.u64.u32 %tmpa, %i;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %aki, %aki, %tmpa;
// b[j][k]
mul.wide.u32 %tmpa, %l, %j;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %bjk, %bjk, %tmpa;
// c[i][j]
mul.wide.u32 %tmpa, %m, %i;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %cij, %cij, %tmpa;
cvt.u64.u32 %tmpa, %j;
shl.b64 %tmpa, %tmpa, 2;
add.u64 %cij, %cij, %tmpa;
// as[ti]
mov.u32 %asti, as;
mov.u32 %tmp, %ti;
shl.b32 %tmp, %tmp, 5 + 4;
add.u32 %asti, %asti, %tmp;
// bs[tj]
mov.u32 %bstj, bs;
mov.u32 %tmp, %tj;
shl.b32 %tmp, %tmp, 5 + 4;
add.u32 %bstj, %bstj, %tmp;
// cr
mov.f32 %cr0, 0.0;
mov.f32 %cr1, 0.0;
mov.f32 %cr2, 0.0;
mov.f32 %cr3, 0.0;
mov.f32 %cr4, 0.0;
mov.f32 %cr5, 0.0;
mov.f32 %cr6, 0.0;
mov.f32 %cr7, 0.0;
mov.f32 %cr8, 0.0;
mov.f32 %cr9, 0.0;
mov.f32 %cr10, 0.0;
mov.f32 %cr11, 0.0;
mov.f32 %cr12, 0.0;
mov.f32 %cr13, 0.0;
mov.f32 %cr14, 0.0;
mov.f32 %cr15, 0.0;
mov.u32 %k, 0;
loop: setp.ge.u32 %p, %k, %l;
@%p bra eloop;
// k + tj < l
mov.u32 %tmp, %k;
add.u32 %tmp, %tmp, %tj;
setp.lt.u32 %p, %tmp, %l;
// a[k + tj][i]
mov.u64 %tmpa, %aki;
mul.wide.u32 %tmpa2, %n, %tj;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 0 < n
mov.u32 %tmp, %i;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar0, 0.0; // ar0 = 0.0
@%p2 ld.global.f32 %ar0, [%tmpa]; // ar0 = a[k + tj][i + 0]
// i + 1 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar1, 0.0; // ar1 = 0.0
@%p2 ld.global.f32 %ar1, [%tmpa + (1 << 2)]; // ar1 = a[k + tj][i + 1]
// i + 2 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar2, 0.0; // ar2 = 0.0
@%p2 ld.global.f32 %ar2, [%tmpa + (2 << 2)]; // ar2 = a[k + tj][i + 2]
// i + 3 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p2, %tmp, %n;
and.pred %p2, %p2, %p;
mov.f32 %ar3, 0.0; // ar3 = 0.0
@%p2 ld.global.f32 %ar3, [%tmpa + (3 << 2)]; // ar3 = a[k + tj][i + 3]
// as[ti][tj + ik]
mov.u32 %tmp, %asti;
mov.u32 %tmp2, %tj;
add.u32 %tmp2, %tmp2, %ti;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
st.shared.v4.f32 [%tmp], { %ar0, %ar1, %ar2, %ar3 };
// k + ti < l
mov.u32 %tmp, %k;
add.u32 %tmp, %tmp, %ti;
setp.lt.u32 %p, %tmp, %l;
// b[j][k + ti]
mov.u64 %tmpa, %bjk;
cvt.u64.u32 %tmpa2, %ti;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br0, 0.0; // br0 = 0.0
@%p2 ld.global.f32 %br0, [%tmpa]; // br0 = b[j + 0][k + ti]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br1, 0.0; // br1 = 0.0
@%p2 ld.global.f32 %br1, [%tmpa]; // br1 = b[j + 1][k + ti]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br2, 0.0; // br2 = 0.0
@%p2 ld.global.f32 %br2, [%tmpa]; // br2 = b[j + 2][k + ti]
cvt.u64.u32 %tmpa2, %l;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
mov.f32 %br3, 0.0; // br3 = 0.0
@%p2 ld.global.f32 %br3, [%tmpa]; // br3 = b[j + 3][k + ti]
//cvt.u64.u32 %tmpa2, %l;
//shl.b64 %tmpa2, %tmpa2, 2;
//add.u64 %tmpa, %tmpa, %tmpa2;
// bs[tj][ti + jk]
mov.u32 %tmp, %bstj;
mov.u32 %tmp2, %ti;
add.u32 %tmp2, %tmp2, %tj;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
st.shared.v4.f32 [%tmp], { %br0, %br1, %br2, %br3 };
bar.sync 0;
mov.u32 %tk, 0;
loop2: setp.ge.u32 %p, %tk, 32;
@%p bra eloop2;
// as[ti][tk + ik]
mov.u32 %tmp, %asti;
mov.u32 %tmp2, %tk;
add.u32 %tmp2, %tmp2, %ti;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
ld.shared.v4.f32 { %ar0, %ar1, %ar2, %ar3 }, [%tmp];
// bs[tj][ti + jk]
mov.u32 %tmp, %bstj;
mov.u32 %tmp2, %tk;
add.u32 %tmp2, %tmp2, %tj;
and.b32 %tmp2, %tmp2, 31;
shl.b32 %tmp2, %tmp2, 4;
add.u32 %tmp, %tmp, %tmp2;
ld.shared.v4.f32 { %br0, %br1, %br2, %br3 }, [%tmp];
fma.rn.f32 %cr0, %ar0, %br0, %cr0;
fma.rn.f32 %cr1, %ar0, %br1, %cr1;
fma.rn.f32 %cr2, %ar0, %br2, %cr2;
fma.rn.f32 %cr3, %ar0, %br3, %cr3;
fma.rn.f32 %cr4, %ar1, %br0, %cr4;
fma.rn.f32 %cr5, %ar1, %br1, %cr5;
fma.rn.f32 %cr6, %ar1, %br2, %cr6;
fma.rn.f32 %cr7, %ar1, %br3, %cr7;
fma.rn.f32 %cr8, %ar2, %br0, %cr8;
fma.rn.f32 %cr9, %ar2, %br1, %cr9;
fma.rn.f32 %cr10, %ar2, %br2, %cr10;
fma.rn.f32 %cr11, %ar2, %br3, %cr11;
fma.rn.f32 %cr12, %ar3, %br0, %cr12;
fma.rn.f32 %cr13, %ar3, %br1, %cr13;
fma.rn.f32 %cr14, %ar3, %br2, %cr14;
fma.rn.f32 %cr15, %ar3, %br3, %cr15;
add.u32 %tk, %tk, 1;
bra loop2;
eloop2: bar.sync 0;
cvt.u64.u32 %tmpa, %n;
shl.b64 %tmpa, %tmpa, 5 + 2;
add.u64 %aki, %aki, %tmpa;
add.u64 %bjk, %bjk, 32 << 2;
add.u32 %k, %k, 32;
bra loop;
eloop: mov.u64 %tmpa, %cij;
// i + 0 < n
mov.u32 %tmp, %i;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr0; // c[i + 0][j + 0] = cr0
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr1; // c[i + 0][j + 1] = cr1
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr2; // c[i + 0][j + 2] = cr2
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr3; // c[i + 0][j + 3] = cr3
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 1 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr4; // c[i + 1][j + 0] = cr4
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr5; // c[i + 1][j + 1] = cr5
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr6; // c[i + 1][j + 2] = cr6
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr7; // c[i + 1][j + 3] = cr7
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 2 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr8; // c[i + 2][j + 0] = cr8
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr9; // c[i + 2][j + 1] = cr9
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr10;// c[i + 2][j + 2] = cr10
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr11;// c[i + 2][j + 3] = cr11
cvt.u64.u32 %tmpa2, %m;
shl.b64 %tmpa2, %tmpa2, 2;
add.u64 %tmpa, %tmpa, %tmpa2;
// i + 3 < n
mov.u32 %tmp, %i;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p, %tmp, %n;
// j + 0 < m
mov.u32 %tmp, %j;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa], %cr12; // c[i + 3][j + 0] = cr12
// j + 1 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 1;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (1 << 2)], %cr13;// c[i + 3][j + 1] = cr13
// j + 2 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 2;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (2 << 2)], %cr14;// c[i + 3][j + 2] = cr14
// j + 3 < m
mov.u32 %tmp, %j;
add.u32 %tmp, %tmp, 3;
setp.lt.u32 %p2, %tmp, %m;
and.pred %p2, %p2, %p;
@%p2 st.global.f32 [%tmpa + (3 << 2)], %cr15;// c[i + 3][j + 3] = cr15
//cvt.u64.u32 %tmpa2, %m;
//shl.b64 %tmpa2, %tmpa2, 2;
//add.u64 %tmpa, %tmpa, %tmpa2;
ret;
}