|
| 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