Skip to content

Commit cd1f515

Browse files
committed
Add LSX
1 parent ac78a76 commit cd1f515

9 files changed

Lines changed: 594 additions & 3 deletions

File tree

.github/workflows/codegen.yml

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
name: codegen
22
on:
3-
- pull_request
3+
- push
44

55
jobs:
66
install-pspamm:
@@ -65,12 +65,14 @@ jobs:
6565
- rvv256
6666
- rvv512
6767
- rvv1024
68+
- lsx128
69+
- lsx256
6870
steps:
6971
- name: apt-get
7072
run: |
7173
set -euo pipefail
7274
sudo apt-get update
73-
sudo apt-get install g++-aarch64-linux-gnu g++-riscv64-linux-gnu qemu-user-static
75+
sudo apt-get install g++-aarch64-linux-gnu g++-riscv64-linux-gnu g++-14-loongarch64-linux-gnu qemu-user-static
7476
7577
- name: setup-python
7678
uses: actions/setup-python@v4
Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,29 @@
1+
class Max:
2+
@classmethod
3+
def getBlocksize(cls, m, n, bk, v_size, prec):
4+
bm = v_size
5+
bn = 1
6+
maxval = 0
7+
8+
for i in range(v_size, m+1, v_size):
9+
for j in range(1, n+1):
10+
# can be replaced by cls.LSX_condition_extended here
11+
# (but that seemed to be slower in the end)
12+
if cls.LSX_condition(i, j, bk, v_size):
13+
if i*j > maxval and (cls.LSX_condition(i, j, bk, v_size) or j > 1):
14+
maxval = i*j
15+
bm = i
16+
bn = j
17+
18+
while cls.LSX_condition(bm, bn, bk+1, v_size):
19+
bk += 1
20+
21+
return (bm, bn, bk)
22+
23+
@classmethod
24+
def LSX_condition(cls, bm, bn, bk, v_size):
25+
# ceiling division
26+
vm = -(bm // -v_size)
27+
return (bn + bk) * vm + bn * bk <= 32
28+
29+
Default = Max
Lines changed: 236 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,236 @@
1+
from pspamm.cursors import *
2+
3+
from pspamm.codegen.architectures.lsx.operands import *
4+
from pspamm.codegen.ast import *
5+
from pspamm.codegen.sugar import *
6+
from pspamm.codegen.generator import *
7+
from pspamm.codegen.precision import *
8+
from pspamm.codegen.regcache import *
9+
10+
class Generator(AbstractGenerator):
11+
template = """
12+
void {funcName} (const {real_type}* A, const {real_type}* B, {real_type}* C, {real_type} alpha, {real_type} beta, {real_type} const* prefetch) {{
13+
__asm__ __volatile__(
14+
{body_text}
15+
: : {args} : {clobbered});
16+
17+
#ifndef NDEBUG
18+
#ifdef _OPENMP
19+
#pragma omp atomic
20+
#endif
21+
pspamm_num_total_flops += {flop};
22+
#endif
23+
}}
24+
"""
25+
v_len = 2
26+
27+
def get_v_size(self):
28+
return (16 // self.precision.size()) * self.v_len
29+
30+
def get_template(self):
31+
return Generator.template
32+
33+
def use_broadcast(self):
34+
return True
35+
36+
def has_masks(self):
37+
return False
38+
39+
def init_mask(self, m, bm, v_size, tempreg, maskregs):
40+
return block("")
41+
42+
def make_argument_load(self, starting_regs, prefetch):
43+
asm = block("Load arguments")
44+
asm.add(ld(InputOperand(f'0', 'm', 'A'), starting_regs[0], False))
45+
asm.add(ld(InputOperand(f'1', 'm', 'B'), starting_regs[1], False))
46+
asm.add(ld(InputOperand(f'2', 'm', 'C'), starting_regs[2], False))
47+
asm.add(ld(InputOperand(f'3', 'm', 'alpha'), starting_regs[3], False))
48+
asm.add(ld(InputOperand(f'4', 'm', 'beta'), starting_regs[4], False))
49+
if prefetch:
50+
asm.add(ld(InputOperand(f'5', 'm', 'prefetch'), starting_regs[5], False))
51+
return asm
52+
53+
def make_reg_blocks(self, bm:int, bn:int, bk:int, v_size:int, nnz:int, m:int, n:int, k:int, prefetch: str):
54+
assert(bm % v_size == 0)
55+
vm = self.ceil_div(bm, v_size)
56+
57+
assert (bn + bk) * vm + bn * bk <= 32
58+
59+
vmm = {
60+
1: vr,
61+
2: xr
62+
}[self.v_len]
63+
64+
A_regs = Matrix([[vmm(vm*c + r) for c in range(bk)] for r in range(vm)])
65+
Aoffset = vm*bk
66+
67+
B_regs = Matrix([[vmm(Aoffset + bn * r + c) for c in range(bn)] for r in range(bk)])
68+
C_regs = Matrix([[vmm(32 - vm*bn + vm*c + r) for c in range(bn)]
69+
for r in range(vm)])
70+
71+
b_reg = Aoffset
72+
alpha_reg = [vmm(b_reg)] * 2
73+
beta_reg = [vmm(b_reg + 1)] * 2
74+
75+
starting_regs = [r(10), r(11), r(12), r(13), r(14), r(6), r(5)]
76+
77+
additional_regs = [r(15), r(16), r(17), r(31), r(7)]
78+
79+
loop_regs = [r(28), r(29), r(30)]
80+
81+
prefetch_reg = prefetch == 'BL2viaC'
82+
83+
return A_regs, B_regs, C_regs, starting_regs, alpha_reg, beta_reg, loop_regs, additional_regs, [], prefetch_reg
84+
85+
def make_scaling_offsets(self,
86+
additional_regs: List[Register],
87+
nnz: int
88+
) -> Block:
89+
return block("")
90+
91+
def init_block(self, size):
92+
return block("")
93+
94+
def move_register_block(self,
95+
cursor: Cursor,
96+
cursor_ptr: CursorLocation,
97+
block_offset: Coords,
98+
registers: Matrix[Register],
99+
v_size: int,
100+
additional_regs,
101+
mask: Matrix[bool] = None,
102+
store: bool = False,
103+
prefetching: str = None,
104+
load_offset: int = 0,
105+
pf_cursor: Cursor = None,
106+
pf_cursor_ptr: CursorLocation = None,
107+
temp = None
108+
) -> Block:
109+
110+
rows, cols = registers.shape
111+
action = "Store" if store else "Load"
112+
asm = block(f"{action} {cursor.name} register block @ {block_offset}")
113+
114+
max_offs = 2047
115+
cur11 = 0
116+
117+
for ic in range(cols):
118+
for ir in range(rows):
119+
if (mask is None) or (mask[ir,ic]):
120+
all_coords = [Coords(down=ir*v_size+i,right=ic) for i in range(v_size)]
121+
has_nonzero = [cursor.has_nonzero_cell(cursor_ptr, block_offset, offset) for offset in all_coords]
122+
if all(has_nonzero):
123+
cell_offset = all_coords[0]
124+
addr, comment = cursor.look(cursor_ptr, block_offset, cell_offset)
125+
addr.disp += self.precision.size() * load_offset
126+
needsmove = False
127+
if addr.disp > max_offs:
128+
moved = addr.disp - cur11
129+
if moved > 0 and moved <= max_offs:
130+
addr.disp = moved
131+
else:
132+
asm.add(add(addr.disp, additional_regs[0], "", addr.base))
133+
cur11 = addr.disp
134+
addr.disp = 0
135+
needsmove = True
136+
137+
addr.base = additional_regs[0]
138+
if store:
139+
asm.add(st(registers[ir,ic], addr, True, comment))
140+
if prefetching == 'BL2viaC' and pf_cursor is not None:
141+
addr, comment = pf_cursor.look(pf_cursor_ptr, block_offset, cell_offset)
142+
addr.disp += self.precision.size() * load_offset
143+
if addr.disp > max_offs:
144+
moved = addr.disp - cur11
145+
if needsmove:
146+
asm.add(add(addr.disp, additional_regs[3], "", addr.base))
147+
addr.disp = 0
148+
else:
149+
addr.disp = moved
150+
addr.base = additional_regs[3]
151+
asm.add(prefetch(addr, closeness="L2"))
152+
else:
153+
asm.add(ld(addr, registers[ir,ic], True, comment))
154+
elif any(has_nonzero):
155+
raise NotImplementedError("Element-wise sparsity in A is not yet fully implemented.")
156+
return asm
157+
158+
def make_zero_block(self, registers: Matrix[Register], additional_regs) -> Block:
159+
160+
rows, cols = registers.shape
161+
asm = block("zero registers")
162+
163+
for ic in range(cols):
164+
for ir in range(rows):
165+
asm.add(mov(0, registers[ir,ic], True))
166+
167+
return asm
168+
169+
170+
def make_microkernel(self,
171+
A: Cursor,
172+
B: Cursor,
173+
A_ptr: CursorLocation,
174+
B_ptr: CursorLocation,
175+
A_regs: Matrix[Register],
176+
B_regs,
177+
C_regs: Matrix[Register],
178+
v_size:int,
179+
additional_regs,
180+
to_A_block: Coords = Coords(),
181+
to_B_block: Coords = Coords(),
182+
sub: bool = False
183+
) -> Block:
184+
185+
""" make_microkernel generates a GEMM microkernel for two blocks using the outer-product formulation.
186+
It is responsible for loading and unloading the A block,
187+
It does not assume that the A or B cursors point to the start of the block.
188+
Instead, the coordinates to the start of the block are passed separately.
189+
It does not modify any cursor pointers.
190+
"""
191+
asm = block("Block GEMM microkernel")
192+
bm,bk,aidx,apattern = A.get_block(A_ptr, to_A_block)
193+
bk,bn,bidx,bpattern = B.get_block(B_ptr, to_B_block)
194+
assert(bm % v_size == 0)
195+
196+
mask = sparse_mask(A_regs, A, A_ptr, to_A_block, B, B_ptr, to_B_block, v_size)
197+
asm.add(self.move_register_block(A, A_ptr, to_A_block, A_regs, v_size, additional_regs, mask, store=False, temp=B_regs[0,0]))
198+
199+
Vm = self.ceil_div(bm, v_size)
200+
cur11 = 0
201+
max_offs = 2047
202+
203+
bs = []
204+
for Vmi in range(Vm):
205+
for bni in range(bn): # inside this n-block
206+
for bki in range(bk): # inside this k-block
207+
to_bcell = Coords(down=bki, right=bni)
208+
to_acell = Coords(down=Vmi*v_size, right=bki)
209+
if B.has_nonzero_cell(B_ptr, to_B_block, to_bcell) and A.has_nonzero_cell(A_ptr, to_A_block, to_acell):
210+
B_cell_addr, B_comment = B.look(B_ptr, to_B_block, to_bcell)
211+
if B_regs[bki, bni] not in bs:
212+
# max_offs is the maximum allowed immediate offset when using ld1rd/ld1rw to broadcast a scalar value
213+
if B_cell_addr.disp > max_offs:
214+
moved = B_cell_addr.disp - cur11
215+
if moved > 0 and moved <= max_offs:
216+
B_cell_addr.disp = moved
217+
else:
218+
asm.add(add(B_cell_addr.disp, additional_regs[0], "", B_cell_addr.base))
219+
cur11 = B_cell_addr.disp
220+
B_cell_addr.disp = 0
221+
222+
B_cell_addr.base = additional_regs[0]
223+
224+
asm.add(bcst(B_cell_addr, B_regs[bki, bni], B_comment))
225+
bs.append(B_regs[bki, bni])
226+
227+
for bki in range(bk): # inside this k-block
228+
for Vmi in range(Vm):
229+
for bni in range(bn): # inside this n-block
230+
to_bcell = Coords(down=bki, right=bni)
231+
to_acell = Coords(down=Vmi*v_size, right=bki)
232+
if B.has_nonzero_cell(B_ptr, to_B_block, to_bcell) and A.has_nonzero_cell(A_ptr, to_A_block, to_acell):
233+
_, B_comment = B.look(B_ptr, to_B_block, to_bcell)
234+
comment = f"C[{Vmi*v_size}:{Vmi*v_size+v_size},{bni}] += A[{Vmi*v_size}:{Vmi*v_size+v_size},{bki}]*{B_comment}"
235+
asm.add(fma(B_regs[bki, bni], A_regs[Vmi, bki], C_regs[Vmi, bni], comment=comment, bcast=None, sub=sub))
236+
return asm

0 commit comments

Comments
 (0)