1 /*
2 * Copyright (c) 2025, Oracle and/or its affiliates. All rights reserved.
3 * DO NOT ALTER OR REMOVE COPYRIGHT NOTICES OR THIS FILE HEADER.
4 *
5 * This code is free software; you can redistribute it and/or modify it
6 * under the terms of the GNU General Public License version 2 only, as
7 * published by the Free Software Foundation. Oracle designates this
8 * particular file as subject to the "Classpath" exception as provided
9 * by Oracle in the LICENSE file that accompanied this code.
10 *
11 * This code is distributed in the hope that it will be useful, but WITHOUT
12 * ANY WARRANTY; without even the implied warranty of MERCHANTABILITY or
13 * FITNESS FOR A PARTICULAR PURPOSE. See the GNU General Public License
14 * version 2 for more details (a copy is included in the LICENSE file that
15 * accompanied this code).
16 *
17 * You should have received a copy of the GNU General Public License version
18 * 2 along with this work; if not, write to the Free Software Foundation,
19 * Inc., 51 Franklin St, Fifth Floor, Boston, MA 02110-1301 USA.
20 *
21 * Please contact Oracle, 500 Oracle Parkway, Redwood Shores, CA 94065 USA
22 * or visit www.oracle.com if you need additional information or have any
23 * questions.
24 */
25 package hat.test;
26
27 import hat.Accelerator;
28 import hat.ComputeContext;
29 import hat.NDRange;
30
31 import static hat.KernelContext.*;
32 import hat.backend.Backend;
33 import hat.types.BF16;
34 import hat.buffer.BF16Array;
35 import hat.types.F16;
36 import hat.buffer.F16Array;
37 import hat.buffer.F32Array;
38 import hat.buffer.F32ArrayPadded;
39 import hat.types.Float4;
40 import hat.device.DeviceSchema;
41 import hat.device.NonMappableIface;
42 import hat.test.annotation.HatTest;
43 import hat.test.exceptions.HATAssertionError;
44 import hat.test.exceptions.HATAsserts;
45 import hat.test.exceptions.HATExpectedPrecisionError;
46 import jdk.incubator.code.Reflect;
47
48 import java.lang.invoke.MethodHandles;
49 import java.util.Random;
50
51 public class TestMatMul {
52
53 private static final int SIZE = 256;
54
55 @Reflect
56 public static void matrixMultiplyKernel2D( F32Array matrixA, F32Array matrixB, F32Array matrixC, int size) {
57 if (GIX() < GSX()) {
58 if (GIY() < GSY()) {
59 float acc = 0.0f;
60 for (int k = 0; k < size; k++) {
61 acc += (matrixA.array(GIX() * size + k) * matrixB.array(k * size + GIY()));
62 }
63 matrixC.array(GIX() * size + GIY(), acc);
64 }
65 }
66 }
67
68 @Reflect
69 public static void matrixMultiplyKernel2DLI( F32Array matrixA, F32Array matrixB, F32Array matrixC, int size) {
70 if (GIX() < GSX()) {
71 if (GIY() < GSY()) {
72 float acc = 0.0f;
73 for (int k = 0; k < size; k++) {
74 acc += (matrixA.array(GIY() * size + k) * matrixB.array(k * size + GIX()));
75 }
76 matrixC.array(GIY() * size + GIX(), acc);
77 }
78 }
79 }
80
81 @Reflect
82 public static void matrixMultiplyKernel2DLIF16( F16Array matrixA, F16Array matrixB, F16Array matrixC, int size) {
83 if (GIX() < GSX()) {
84 if (GIY() < GSY()) {
85 F16 acc = F16.of(0.0f);
86 for (int k = 0; k < size; k++) {
87 F16 valA = matrixA.array(GIY() * size + k);
88 F16 valB = matrixB.array(k * size + GIX());
89 F16 valc = F16.mul(valA, valB);
90 acc = F16.add(acc, valc);
91 }
92 F16 resultC = matrixC.array(GIY() * size + GIX());
93 resultC.value(acc.value());
94 }
95 }
96 }
97
98 private interface MyLocalArrayFixedSize extends NonMappableIface {
99 void array(long index, float value);
100
101 float array(long index);
102
103 DeviceSchema<MyLocalArrayFixedSize> deviceSchema = DeviceSchema.of(MyLocalArrayFixedSize.class,
104 myPrivateArray -> myPrivateArray
105 .array("array", 256));
106
107 static MyLocalArrayFixedSize create(Accelerator accelerator) {
108 return null;
109 }
110
111 static MyLocalArrayFixedSize createLocal() {
112 return null;
113 }
114 }
115
116 @Reflect
117 public static void matrixMultiplyKernel2DTiling( F32Array matrixA, F32Array matrixB, F32Array matrixC, int size) {
118
119 final int tileSize = 16;
120 MyLocalArrayFixedSize tileA = MyLocalArrayFixedSize.createLocal();
121 MyLocalArrayFixedSize tileB = MyLocalArrayFixedSize.createLocal();
122
123 int groupIndexX = BIX();
124 int groupIndexY = BIY();
125 int localIdx = LIX();
126 int localIdy = LIY();
127
128 // we identify the row and column
129 int row = groupIndexY * tileSize + localIdy;
130 int col = groupIndexX * tileSize + localIdx;
131
132 // Compute matrix-vector and accumulate the result over the tiles
133 float sum = 0.0f;
134 for (int tile = 0; tile < (size / tileSize); tile++) {
135 // Copy from global to shared memory
136 tileA.array((long) localIdy * tileSize + localIdx, matrixA.array((long) row * size + tile * tileSize + localIdx));
137 tileB.array((long) localIdy * tileSize + localIdx, matrixB.array((tile * tileSize + localIdy) * size + col));
138
139 // Apply a barrier for the local group: we need to guarantee that all threads that belong
140 // to the same group reach this point before doing the partial reduction
141 barrier();
142
143 // compute partial reductions over the tile
144 for (int k = 0; k < tileSize; k++) {
145 sum += (tileA.array((long) localIdy * tileSize + k) * tileB.array(k * tileSize + localIdx));
146 }
147
148 // A new local barrier for all threads that belong to the same group before loading a new tile into
149 // share memory. With the following barrier, we can ensure that all threads within the same workgroup
150 // finished the compute for the partial reduction
151 barrier();
152 }
153
154 // copy result from shared memory to global memory
155 matrixC.array((long) row * size + col, sum);
156 }
157
158 @Reflect
159 public static float compute( F32Array matrixA, F32Array matrixB, int size, int j) {
160 float acc = 0.0f;
161 for (int k = 0; k < size; k++) {
162 acc += (matrixA.array(GIX() * size + k) * matrixB.array(k * size + j));
163 }
164 return acc;
165 }
166
167 @Reflect
168 public static void matrixMultiplyKernel1D( F32Array matrixA, F32Array matrixB, F32Array matrixC, int size) {
169 if (GIX() < GSX()) {
170 for (int j = 0; j < size; j++) {
171 float acc = 0.0f;
172 for (int k = 0; k < size; k++) {
173 acc += (matrixA.array(GIX() * size + k) * matrixB.array(k * size + j));
174 }
175 matrixC.array(GIX() * size + j, acc);
176 }
177 }
178 }
179
180 @Reflect
181 public static void matrixMultiplyKernel1DWithFunctionCalls( F32Array matrixA, F32Array matrixB, F32Array matrixC, int size) {
182 if (GIX() < GSX()) {
183 for (int j = 0; j < size; j++) {
184 float acc = compute( matrixA, matrixB, size, j);
185 matrixC.array(GIX() * size + j, acc);
186 }
187 }
188 }
189
190 @Reflect
191 public static void matrixMultiply1D( ComputeContext cc, F32Array matrixA, F32Array matrixB, F32Array matrixC, int globalSize) {
192 cc.dispatchKernel(NDRange.of1D(globalSize,16),
193 ()-> matrixMultiplyKernel1D( matrixA, matrixB, matrixC, globalSize)
194 );
195 }
196
197 final static int BLOCK_SIZE = 16;
198
199 @Reflect
200 public static void matrixMultiply1DWithFunctionCalls( ComputeContext cc, F32Array matrixA, F32Array matrixB, F32Array matrixC, int size) {
201 cc.dispatchKernel(NDRange.of1D(size),
202 ()-> matrixMultiplyKernel1DWithFunctionCalls( matrixA, matrixB, matrixC, size)
203 );
204 }
205
206 @Reflect
207 public static void matrixMultiply2D( ComputeContext cc, F32Array matrixA, F32Array matrixB, F32Array matrixC, int globalSize) {
208 cc.dispatchKernel(NDRange.of2D(globalSize, globalSize,BLOCK_SIZE, BLOCK_SIZE),
209 ()-> matrixMultiplyKernel2D( matrixA, matrixB, matrixC, globalSize)
210 );
211 }
212
213 @Reflect
214 public static void matrixMultiply2DLI( ComputeContext cc, F32Array matrixA, F32Array matrixB, F32Array matrixC, int globalSize) {
215 cc.dispatchKernel(NDRange.of2D(globalSize, globalSize,BLOCK_SIZE, BLOCK_SIZE),
216 ()-> matrixMultiplyKernel2DLI( matrixA, matrixB, matrixC, globalSize)
217 );
218 }
219
220 @Reflect
221 public static void matrixMultiply2DLIF16( ComputeContext cc, F16Array matrixA, F16Array matrixB, F16Array matrixC, int globalSize) {
222 cc.dispatchKernel(NDRange.of2D(globalSize, globalSize, BLOCK_SIZE, BLOCK_SIZE),
223 ()-> matrixMultiplyKernel2DLIF16( matrixA, matrixB, matrixC, globalSize)
224 );
225 }
226
227 @Reflect
228 public static void matrixMultiply2DTiling( ComputeContext cc, F32Array matrixA, F32Array matrixB, F32Array matrixC, int globalSize) {
229 cc.dispatchKernel(NDRange.of2D(globalSize, globalSize, BLOCK_SIZE, BLOCK_SIZE),
230 ()-> matrixMultiplyKernel2DTiling( matrixA, matrixB, matrixC, globalSize)
231 );
232 }
233
234 private static void runSequential(F16Array matrixA, F16Array matrixB, F16Array matrixC, final int size) {
235 for (int i = 0; i < size; i++) {
236 for (int j = 0; j < size; j++) {
237 F16 sum = F16.of(0.0f);
238 for (int k = 0; k < size; k++) {
239 F16 a = matrixA.array((long) i * size + k);
240 F16 b = matrixB.array((long) k * size + j);
241 sum = F16.add(sum, F16.mul(a, b));
242 }
243 matrixC.array((long) i * size + j).value(sum.value());
244 }
245 }
246 }
247
248 private static void runSequential(F32Array matrixA, F32Array matrixB, F32Array matrixC, final int size) {
249 for (int i = 0; i < size; i++) {
250 for (int j = 0; j < size; j++) {
251 float sum = 0;
252 for (int k = 0; k < size; k++) {
253 float a = matrixA.array((long) i * size + k);
254 float b = matrixB.array((long) k * size + j);
255 sum += a * b;
256 }
257 matrixC.array((long) i * size + j, sum);
258 }
259 }
260 }
261
262 private static void runSequential(F32ArrayPadded matrixA, F32ArrayPadded matrixB, F32ArrayPadded matrixC, final int size) {
263 for (int i = 0; i < size; i++) {
264 for (int j = 0; j < size; j++) {
265 float sum = 0;
266 for (int k = 0; k < size; k++) {
267 float a = matrixA.array((long) i * size + k);
268 float b = matrixB.array((long) k * size + j);
269 sum += a * b;
270 }
271 matrixC.array((long) i * size + j, sum);
272 }
273 }
274 }
275
276 private static void runSequential(BF16Array matrixA, BF16Array matrixB, BF16Array matrixC, final int size) {
277 for (int i = 0; i < size; i++) {
278 for (int j = 0; j < size; j++) {
279 BF16 sum = BF16.of(0.0f);
280 for (int k = 0; k < size; k++) {
281 BF16 a = matrixA.array((long) i * size + k);
282 BF16 b = matrixB.array((long) k * size + j);
283 sum = BF16.add(sum, BF16.mul(a, b));
284 }
285 matrixC.array((long) i * size + j).value(sum.value());
286 }
287 }
288 }
289
290 @HatTest
291 @Reflect
292 public void testMatrixMultiply1D() {
293 var lookup = MethodHandles.lookup();
294 var accelerator = new Accelerator(lookup, Backend.FIRST);
295
296 final int size = SIZE;
297 var matrixA = F32Array.create(accelerator, size * size);
298 var matrixB = F32Array.create(accelerator, size * size);
299
300 // Matrix for the results
301 var matrixC = F32Array.create(accelerator, size * size);
302 var resultSeq = F32Array.create(accelerator, size * size);
303
304 // Initialize matrices (A and B have the same size)
305 Random r = new Random(19);
306
307 for (int j = 0; j < matrixA.length(); j++) {
308 matrixA.array(j, r.nextFloat());
309 matrixB.array(j, r.nextFloat());
310 }
311
312 accelerator.compute(cc ->
313 TestMatMul.matrixMultiply1D(cc, matrixA, matrixB, matrixC, size));
314
315 // Run Seq for reference
316 runSequential(matrixA, matrixB, resultSeq, size);
317
318 for (int j = 0; j < size; j++) {
319 for (int i = 0; i < size; i++) {
320 HATAsserts.assertEquals(resultSeq.array(i * size + j), matrixC.array(i * size + j), 0.01f);
321 }
322 }
323 }
324
325 @HatTest
326 @Reflect
327 public void testMatrixMultiply1DWithFunctionCalls() {
328 var lookup = MethodHandles.lookup();
329 var accelerator = new Accelerator(lookup, Backend.FIRST);
330
331 final int size = SIZE;
332 var matrixA = F32Array.create(accelerator, size * size);
333 var matrixB = F32Array.create(accelerator, size * size);
334
335 // Matrix for the results
336 var matrixC = F32Array.create(accelerator, size * size);
337 var resultSeq = F32Array.create(accelerator, size * size);
338
339 // Initialize matrices (A and B have the same size)
340 Random r = new Random(19);
341
342 for (int j = 0; j < matrixA.length(); j++) {
343 matrixA.array(j, r.nextFloat());
344 matrixB.array(j, r.nextFloat());
345 }
346
347 accelerator.compute(cc ->
348 TestMatMul.matrixMultiply1DWithFunctionCalls(cc, matrixA, matrixB, matrixC, size));
349
350 // Run Seq for reference
351 runSequential(matrixA, matrixB, resultSeq, size);
352
353 for (int j = 0; j < size; j++) {
354 for (int i = 0; i < size; i++) {
355 HATAsserts.assertEquals(resultSeq.array(i * size + j), matrixC.array(i * size + j), 0.01f);
356 }
357 }
358 }
359
360
361 @HatTest
362 @Reflect
363 public void testMatrixMultiply2D() {
364 var lookup = MethodHandles.lookup();
365 var accelerator = new Accelerator(lookup, Backend.FIRST);
366
367 final int size = SIZE;
368 var matrixA = F32Array.create(accelerator, size * size);
369 var matrixB = F32Array.create(accelerator, size * size);
370
371 // Matrix for the results
372 var matrixC = F32Array.create(accelerator, size * size);
373 var resultSeq = F32Array.create(accelerator, size * size);
374
375 // Initialize matrices (A and B have the same size)
376 Random r = new Random(19);
377
378 for (int j = 0; j < matrixA.length(); j++) {
379 matrixA.array(j, r.nextFloat());
380 matrixB.array(j, r.nextFloat());
381 }
382
383 accelerator.compute(cc ->
384 TestMatMul.matrixMultiply2D(cc, matrixA, matrixB, matrixC, size));
385
386 // Run Seq for reference
387 runSequential(matrixA, matrixB, resultSeq, size);
388
389 for (int j = 0; j < size; j++) {
390 for (int i = 0; i < size; i++) {
391 HATAsserts.assertEquals(resultSeq.array(i * size + j), matrixC.array(i * size + j), 0.01f);
392 }
393 }
394 }
395
396 @HatTest
397 @Reflect
398 public void testMatrixMultiply2DLI() {
399 var lookup = MethodHandles.lookup();
400 var accelerator = new Accelerator(lookup, Backend.FIRST);
401
402 final int size = SIZE;
403 var matrixA = F32Array.create(accelerator, size * size);
404 var matrixB = F32Array.create(accelerator, size * size);
405
406 // Matrix for the results
407 var matrixC = F32Array.create(accelerator, size * size);
408 var resultSeq = F32Array.create(accelerator, size * size);
409
410 // Initialize matrices (A and B have the same size)
411 Random r = new Random(19);
412
413 for (int j = 0; j < matrixA.length(); j++) {
414 matrixA.array(j, r.nextFloat());
415 matrixB.array(j, r.nextFloat());
416 }
417
418 accelerator.compute(cc ->
419 TestMatMul.matrixMultiply2DLI(cc, matrixA, matrixB, matrixC, size));
420
421 // Run Seq for reference
422 runSequential(matrixA, matrixB, resultSeq, size);
423
424 for (int j = 0; j < size; j++) {
425 for (int i = 0; i < size; i++) {
426 HATAsserts.assertEquals(resultSeq.array(i * size + j), matrixC.array(i * size + j), 0.01f);
427 }
428 }
429 }
430
431 @HatTest
432 @Reflect
433 public void testMatrixMultiply2DLIF16() {
434 var lookup = MethodHandles.lookup();
435 var accelerator = new Accelerator(lookup, Backend.FIRST);
436
437 final int size = SIZE;
438 var matrixA = F16Array.create(accelerator, size * size);
439 var matrixB = F16Array.create(accelerator, size * size);
440
441 // Matrix for the results
442 var matrixC = F16Array.create(accelerator, size * size);
443 var resultSeq = F16Array.create(accelerator, size * size);
444
445 // Initialize matrices (A and B have the same size)
446 Random r = new Random(19);
447
448 for (int j = 0; j < matrixA.length(); j++) {
449 matrixA.array(j).value(F16.floatToF16(r.nextFloat()).value());
450 matrixB.array(j).value(F16.floatToF16(r.nextFloat()).value());
451 }
452
453 accelerator.compute(cc ->
454 TestMatMul.matrixMultiply2DLIF16(cc, matrixA, matrixB, matrixC, size));
455
456 // Run Seq for reference
457 runSequential(matrixA, matrixB, resultSeq, size);
458
459 for (int j = 0; j < size; j++) {
460 for (int i = 0; i < size; i++) {
461 try {
462 HATAsserts.assertEquals(
463 Float.float16ToFloat(resultSeq.array(i * size + j).value()),
464 Float.float16ToFloat(matrixC.array(i * size + j).value()),
465 0.01f);
466 } catch (HATAssertionError hatAssertionError) {
467 throw new HATExpectedPrecisionError(hatAssertionError.getMessage());
468 }
469 }
470 }
471 }
472
473 @HatTest
474 @Reflect
475 public void testMatrixMultiply2DTiling() {
476 var lookup = MethodHandles.lookup();
477 var accelerator = new Accelerator(lookup, Backend.FIRST);
478
479 final int size = SIZE;
480 var matrixA = F32Array.create(accelerator, size * size);
481 var matrixB = F32Array.create(accelerator, size * size);
482
483 // Matrix for the results
484 var matrixC = F32Array.create(accelerator, size * size);
485 var resultSeq = F32Array.create(accelerator, size * size);
486
487 // Initialize matrices (A and B have the same size)
488 Random r = new Random(19);
489
490 for (int j = 0; j < matrixA.length(); j++) {
491 matrixA.array(j, r.nextFloat());
492 matrixB.array(j, r.nextFloat());
493 }
494
495 accelerator.compute(cc ->
496 TestMatMul.matrixMultiply2DTiling(cc, matrixA, matrixB, matrixC, size));
497
498 // Run Seq for reference
499 runSequential(matrixA, matrixB, resultSeq, size);
500
501 for (int j = 0; j < size; j++) {
502 for (int i = 0; i < size; i++) {
503 HATAsserts.assertEquals(resultSeq.array(i * size + j), matrixC.array(i * size + j), 0.01f);
504 }
505 }
506 }
507
508 private interface SharedMemory extends NonMappableIface {
509 void array(long index, float value);
510
511 float array(long index);
512
513 DeviceSchema<SharedMemory> deviceSchema = DeviceSchema.of(SharedMemory.class,
514 arr -> arr.array("array", 1024));
515
516 static SharedMemory create(Accelerator accelerator) {
517 return null;
518 }
519
520 static SharedMemory createLocal() {
521 return null;
522 }
523
524 default void storeFloat4View(Float4 float4, int index) {
525 }
526 }
527
528 private interface PrivateArray extends NonMappableIface {
529 void array(long index, float value);
530
531 float array(long index);
532
533 DeviceSchema<PrivateArray> deviceSchema = DeviceSchema.of(PrivateArray.class,
534 arr -> arr.array("array", 16));
535
536 static PrivateArray create(Accelerator accelerator) {
537 return null;
538 }
539
540 static PrivateArray createPrivate() {
541 return null;
542 }
543 }
544
545 private interface FlatPrivate extends NonMappableIface {
546 void array(long index, float value);
547
548 float array(long index);
549
550 DeviceSchema<FlatPrivate> deviceSchema = DeviceSchema.of(FlatPrivate.class,
551 arr -> arr.array("array", 4));
552
553 static FlatPrivate create(Accelerator accelerator) {
554 return null;
555 }
556
557 static FlatPrivate createPrivate() {
558 return null;
559 }
560 }
561
562 // Code ported from the HAT example module.
563 @Reflect
564 public static void matrixMultiplyKernel2DRegisterTiling( F32Array matrixA, F32Array matrixB, F32Array matrixC, int size) {
565
566 // Configuration for the kernel: Keep in mind that if you change the following parameters,
567 // also change the scheduling (global and local work sizes).
568 final int BM = 64;
569 final int BN = 64;
570 final int BK = 16;
571 final int TM = 4;
572 final int TN = 4;
573
574 int bx = BIX();
575 int by = BIY();
576
577 int totalResultsBlockTile = BM * BN;
578 final int numThreadsBlockTile = totalResultsBlockTile / (TM * TN);
579
580 final int linearLocalId = LIY() * LSX() + LIX();
581 final int threadCol = LIX();
582 final int threadRow = LIY();
583
584 SharedMemory tileA = SharedMemory.createLocal();
585 SharedMemory tileB = SharedMemory.createLocal();
586
587 int aFrom = by * BM * size;
588 int bFrom = bx * BN;
589 int v = bx * BN;
590 int cFrom = (by * BM * size) + (v);
591
592 final int innerRowA = linearLocalId / BK;
593 final int innerColA = linearLocalId % BK;
594
595 final int strideA = numThreadsBlockTile / BK;
596 final int innerRowB = linearLocalId / BN;
597 final int innerColB = linearLocalId % BN;
598
599 int strideB = numThreadsBlockTile / BN;
600
601 // Declarations of the arrays in private memory to perform register tiling
602 PrivateArray threadResults = PrivateArray.createPrivate();
603 FlatPrivate regM = FlatPrivate.createPrivate();
604 FlatPrivate regN = FlatPrivate.createPrivate();
605
606 // initialize values
607 for (int i = 0; i < (TN * TN); i++) {
608 threadResults.array(i, 0.0f);
609 }
610
611 // Each thread loops over the tiles
612 for (int bkIdx = 0; bkIdx < size; bkIdx += BK) {
613
614 // A) Load data into shared memory for array A
615 for (int loadOffset = 0; loadOffset < BM; loadOffset += strideA) {
616 tileA.array((innerRowA + loadOffset) * BK + innerColA,
617 matrixA.array(((innerRowA + loadOffset) * size + innerColA) + aFrom));
618 }
619
620 // B) Load data matrixB into shared memory for array B
621 for (int loadOffset = 0; loadOffset < BK; loadOffset += strideB) {
622 tileB.array((innerRowB + loadOffset) * BN + innerColB,
623 matrixB.array(((innerRowB + loadOffset) * size + innerColB) + bFrom));
624 }
625 barrier();
626
627 aFrom += (BK);
628 int f = BK * size;
629 bFrom += f;
630
631 // Per-thread, we load the data from the shared memory into register for both
632 // array A and array B (matrix A and B), and then perform the reduction within
633 // the small region in private memory.
634 for (int dotIdx = 0; dotIdx < BK; dotIdx++) {
635 // block into registers
636 for (int i = 0; i < TM; i++) {
637 regM.array(i, tileA.array((threadRow * TM + i) * BK + dotIdx));
638 }
639 for (int i = 0; i < TN; i++) {
640 regN.array(i, tileB.array(dotIdx * BN + threadCol * TN + i));
641 }
642 for (int resIdxM = 0; resIdxM < TM; resIdxM++) {
643 for (int resIdxN = 0; resIdxN < TN; resIdxN++) {
644 float val = regM.array(resIdxM) * regN.array(resIdxN);
645 float acc = threadResults.array(resIdxM * TN + resIdxN);
646 acc += val;
647 threadResults.array((resIdxM * TN + resIdxN), (acc));
648 }
649 }
650 }
651 barrier();
652 }
653
654 // Finally, we store the results of the reductions for the whole 2D register block into global memory.
655 // Essentially, each thread compute a small block of TM * TN sub-block size.
656 for (int resIdxM = 0; resIdxM < TM; resIdxM++) {
657 for (int resIdxN = 0; resIdxN < TN; resIdxN++) {
658 float value = threadResults.array(resIdxM * TN + resIdxN);
659 matrixC.array((((threadRow * TM + resIdxM) * size + threadCol * TN + resIdxN) + (cFrom)), value);
660 }
661 }
662 }
663
664 // Code ported from the HAT example module.
665 @Reflect
666 public static void matrixMultiplyKernel2DRegisterTilingVectorized( F32ArrayPadded matrixA, F32ArrayPadded matrixB, F32ArrayPadded matrixC, int size) {
667
668 // Configuration for the kernel: Keep in mind that if you change the following parameters,
669 // also change the scheduling (global and local work sizes).
670 // final int M = size;
671 final int N = size;
672 final int K = size;
673 final int BM = 64;
674 final int BN = 64;
675 final int BK = 16;
676 final int TM = 4;
677 final int TN = 4;
678
679 int bx = BIX();
680 int by = BIY();
681
682 final int linearLocalId = LIY() * LSX() + LIX();
683 final int threadCol = LIX();
684 final int threadRow = LIY();
685
686 SharedMemory tileA = SharedMemory.createLocal();
687 SharedMemory tileB = SharedMemory.createLocal();
688
689 int aFrom = by * BM * size;
690 int bFrom = bx * BN;
691 int v = bx * BN;
692 int cFrom = (by * BM * size) + (v);
693
694 final int innerRowA = linearLocalId / (BK / 4);
695 final int innerColA = linearLocalId % (BK / 4);
696 final int innerRowB = linearLocalId / (BN / 4);
697 final int innerColB = linearLocalId % (BN / 4);
698
699 // Declarations of the arrays in private memory to perform register tiling
700 PrivateArray threadResults = PrivateArray.createPrivate();
701 FlatPrivate regM = FlatPrivate.createPrivate();
702 FlatPrivate regN = FlatPrivate.createPrivate();
703
704 // initialize values
705 for (int i = 0; i < (TN * TN); i++) {
706 threadResults.array(i, 0.0f);
707 }
708
709 final int extraCols = 0;
710
711 // Each thread loops over the tiles
712 for (int bkIdx = 0; bkIdx < size; bkIdx += BK) {
713
714 Float4 loadA = matrixA.float4View((innerRowA * K + innerColA * 4) + aFrom);
715 tileA.array((innerColA * 4 + 0) * BM + innerRowA, loadA.x());
716 tileA.array((innerColA * 4 + 1) * BM + innerRowA, loadA.y());
717 tileA.array((innerColA * 4 + 2) * BM + innerRowA, loadA.z());
718 tileA.array((innerColA * 4 + 3) * BM + innerRowA, loadA.w());
719
720 Float4 loadB = matrixB.float4View((innerRowB * N + innerColB * 4) + bFrom);
721 tileB.array(innerRowB * (BN + extraCols) + innerColB * 4 + 0, loadB.x());
722 tileB.array(innerRowB * (BN + extraCols) + innerColB * 4 + 1, loadB.y());
723 tileB.array(innerRowB * (BN + extraCols) + innerColB * 4 + 2, loadB.z());
724 tileB.array(innerRowB * (BN + extraCols) + innerColB * 4 + 3, loadB.w());
725
726 barrier();
727
728 aFrom += (BK);
729 int f = BK * size;
730 bFrom += f;
731
732 // Per-thread, we load the data from the shared memory into register for both
733 // array A and array B (matrix A and B), and then perform the reduction within
734 // the small region in private memory.
735 for (int dotIdx = 0; dotIdx < BK; dotIdx++) {
736 // block into registers
737 for (int i = 0; i < TM; i++) {
738 regM.array(i, tileA.array(dotIdx * BM + threadRow * TM + i));
739 }
740 for (int i = 0; i < TN; i++) {
741 regN.array(i, tileB.array(dotIdx * (BN + extraCols) + threadCol * TN + i));
742 }
743 for (int resIdxM = 0; resIdxM < TM; resIdxM++) {
744 for (int resIdxN = 0; resIdxN < TN; resIdxN++) {
745 float val = regM.array(resIdxM) * regN.array(resIdxN);
746 float acc = threadResults.array(resIdxM * TN + resIdxN);
747 acc += val;
748 threadResults.array((resIdxM * TN + resIdxN), (acc));
749 }
750 }
751 }
752 barrier();
753 }
754
755 // Finally, we store the results of the reductions for the whole 2D register block into global memory.
756 // Essentially, each thread compute a small block of TM * TN sub-block size.
757 for (int resIdxM = 0; resIdxM < TM; resIdxM++) {
758 for (int resIdxN = 0; resIdxN < TN; resIdxN++) {
759 float value = threadResults.array(resIdxM * TN + resIdxN);
760 matrixC.array((((threadRow * TM + resIdxM) * N + threadCol * TN + resIdxN) + (cFrom)), value);
761 }
762 }
763 }
764
765 @Reflect
766 public static void matrixMultiply2DRegisterTiling( ComputeContext cc, F32Array matrixA, F32Array matrixB, F32Array matrixC, final int size) {
767 cc.dispatchKernel(NDRange.of2D(256, 256,16, 16),
768 ()-> matrixMultiplyKernel2DRegisterTiling( matrixA, matrixB, matrixC, size)
769 );
770 }
771
772 @Reflect
773 public static void matrixMultiply2DRegisterTilingVectorized( ComputeContext cc, F32ArrayPadded matrixA, F32ArrayPadded matrixB, F32ArrayPadded matrixC, final int size) {
774 cc.dispatchKernel(NDRange.of2D(256, 256,16, 16),
775 ()-> matrixMultiplyKernel2DRegisterTilingVectorized( matrixA, matrixB, matrixC, size)
776 );
777 }
778
779 @HatTest
780 @Reflect
781 public void testMatMul2DRegisterTiling() {
782 var lookup = MethodHandles.lookup();
783 var accelerator = new Accelerator(lookup, Backend.FIRST);
784
785 final int size = 1024;
786 var matrixA = F32Array.create(accelerator, size * size);
787 var matrixB = F32Array.create(accelerator, size * size);
788
789 // Matrix for the results
790 var matrixC = F32Array.create(accelerator, size * size);
791 var resultSeq = F32Array.create(accelerator, size * size);
792
793 // Initialize matrices (A and B have the same size)
794 Random r = new Random(19);
795
796 for (int j = 0; j < matrixA.length(); j++) {
797 matrixA.array(j, r.nextFloat());
798 matrixB.array(j, r.nextFloat());
799 }
800
801 accelerator.compute(cc ->
802 TestMatMul.matrixMultiply2DRegisterTiling(cc, matrixA, matrixB, matrixC, size));
803
804 // Run Seq for reference
805 runSequential(matrixA, matrixB, resultSeq, size);
806
807 for (int j = 0; j < size; j++) {
808 for (int i = 0; i < size; i++) {
809 HATAsserts.assertEquals(resultSeq.array(i * size + j), matrixC.array(i * size + j), 0.01f);
810 }
811 }
812 }
813
814 @HatTest
815 @Reflect
816 public void testMatMul2DRegisterTilingVectorized() {
817 var lookup = MethodHandles.lookup();
818 var accelerator = new Accelerator(lookup, Backend.FIRST);
819
820 final int size = 1024;
821 var matrixA = F32ArrayPadded.create(accelerator, size * size);
822 var matrixB = F32ArrayPadded.create(accelerator, size * size);
823
824 // Matrix for the results
825 var matrixC = F32ArrayPadded.create(accelerator, size * size);
826 var resultSeq = F32ArrayPadded.create(accelerator, size * size);
827
828 // Initialize matrices (A and B have the same size)
829 Random r = new Random(19);
830
831 for (int j = 0; j < matrixA.length(); j++) {
832 matrixA.array(j, r.nextFloat());
833 matrixB.array(j, r.nextFloat());
834 }
835
836 accelerator.compute(cc ->
837 TestMatMul.matrixMultiply2DRegisterTilingVectorized(cc, matrixA, matrixB, matrixC, size));
838
839 // Run Seq for reference
840 runSequential(matrixA, matrixB, resultSeq, size);
841
842 for (int j = 0; j < size; j++) {
843 for (int i = 0; i < size; i++) {
844 HATAsserts.assertEquals(resultSeq.array(i * size + j), matrixC.array(i * size + j), 0.01f);
845 }
846 }
847 }
848
849 private interface SharedMemoryHalf extends NonMappableIface {
850 F16 array(int index);
851
852 DeviceSchema<SharedMemoryHalf> deviceSchema = DeviceSchema.of(SharedMemoryHalf.class, arr ->
853 arr.array("array", 1024, half -> half.field("value"))
854 );
855
856 static SharedMemoryHalf create(Accelerator accelerator) {
857 return null;
858 }
859
860 static SharedMemoryHalf createLocal() {
861 return null;
862 }
863 }
864
865 private interface PrivateArrayHalf extends NonMappableIface {
866 F16 array(int index);
867
868 DeviceSchema<PrivateArrayHalf> deviceSchema = DeviceSchema.of(PrivateArrayHalf.class, arr ->
869 arr.array("array", 16, half -> half.field("value"))
870 );
871
872 static PrivateArrayHalf create(Accelerator accelerator) {
873 return null;
874 }
875
876 static PrivateArrayHalf createPrivate() {
877 return null;
878 }
879 }
880
881 private interface FlatPrivateHalf extends NonMappableIface {
882 F16 array(int index);
883
884 DeviceSchema<FlatPrivateHalf> deviceSchema = DeviceSchema.of(FlatPrivateHalf.class, arr ->
885 arr.array("array", 4, half -> half.field("value"))
886 );
887
888 static FlatPrivateHalf create(Accelerator accelerator) {
889 return null;
890 }
891
892 static FlatPrivateHalf createPrivate() {
893 return null;
894 }
895 }
896
897 // Taking from the HAT Examples module
898 @Reflect
899 public static void matrixMultiplyKernel2DRegisterTilingHalf( F16Array matrixA, F16Array matrixB, F16Array matrixC, int size) {
900 final int BM = 64;
901 final int BN = 64;
902 final int BK = 16;
903 final int TM = 4;
904 final int TN = 4;
905
906 int bx = BIX();
907 int by = BIY();
908
909 int totalResultsBlockTile = BM * BN;
910 final int numThreadsBlockTile = totalResultsBlockTile / (TM * TN);
911
912 final int linearLocalId = LIY() * LSX() + LIX();
913 final int threadCol = LIX();
914 final int threadRow = LIY();
915
916 SharedMemoryHalf tileA = SharedMemoryHalf.createLocal();
917 SharedMemoryHalf tileB = SharedMemoryHalf.createLocal();
918
919 int aFrom = by * BM * size;
920 int bFrom = bx * BN;
921 int v = bx * BN;
922 int cFrom = (by * BM * size) + (v);
923
924 final int innerRowA = linearLocalId / BK;
925 final int innerColA = linearLocalId % BK;
926
927 final int strideA = numThreadsBlockTile / BK;
928 final int innerRowB = linearLocalId / BN;
929 final int innerColB = linearLocalId % BN;
930
931 int strideB = numThreadsBlockTile / BN;
932
933 PrivateArrayHalf threadResults = PrivateArrayHalf.createPrivate();
934 FlatPrivateHalf regM = FlatPrivateHalf.createPrivate();
935 FlatPrivateHalf regN = FlatPrivateHalf.createPrivate();
936
937 for (int i = 0; i < (TN * TN); i++) {
938 F16 init = F16.of(0.0f);
939 threadResults.array(i).value(init.value());
940 }
941
942 for (int bkIdx = 0; bkIdx < size; bkIdx += BK) {
943 for (int loadOffset = 0; loadOffset < BM; loadOffset += strideA) {
944 F16 ha = matrixA.array(((innerRowA + loadOffset) * size + innerColA) + aFrom);
945 tileA.array((innerRowA + loadOffset) * BK + innerColA).value(ha.value());
946 }
947 for (int loadOffset = 0; loadOffset < BK; loadOffset += strideB) {
948 F16 hb = matrixB.array(((innerRowB + loadOffset) * size + innerColB) + bFrom);
949 tileB.array((innerRowB + loadOffset) * BN + innerColB).value(hb.value());
950 }
951 barrier();
952
953 aFrom += (BK);
954 int f = BK * size;
955 bFrom += f;
956
957 for (int dotIdx = 0; dotIdx < BK; dotIdx++) {
958 for (int i = 0; i < TM; i++) {
959 F16 ha = tileA.array((threadRow * TM + i) * BK + dotIdx);
960 regM.array(i).value(ha.value());
961 }
962 for (int i = 0; i < TN; i++) {
963 F16 hb = tileB.array(dotIdx * BN + threadCol * TN + i);
964 regN.array(i).value(hb.value());
965 }
966 for (int resIdxM = 0; resIdxM < TM; resIdxM++) {
967 for (int resIdxN = 0; resIdxN < TN; resIdxN++) {
968 F16 privA = regM.array(resIdxM);
969 F16 privB = regN.array(resIdxN);
970 F16 mul = F16.mul(privA, privB);
971 F16 acc = threadResults.array(resIdxM * TN + resIdxN);
972 acc = F16.add(acc, mul);
973 threadResults.array((resIdxM * TN + resIdxN)).value(acc.value());
974 }
975 }
976 }
977 barrier();
978 }
979 for (int resIdxM = 0; resIdxM < TM; resIdxM++) {
980 for (int resIdxN = 0; resIdxN < TN; resIdxN++) {
981 F16 result = threadResults.array(resIdxM * TN + resIdxN);
982 matrixC.array((((threadRow * TM + resIdxM) * size + threadCol * TN + resIdxN) + (cFrom))).value(result.value());
983 }
984 }
985 }
986
987 private interface SharedMemoryBfloat16 extends NonMappableIface {
988 BF16 array(int index);
989
990 DeviceSchema<SharedMemoryBfloat16> deviceSchema = DeviceSchema.of(SharedMemoryBfloat16.class, arr ->
991 arr.array("array", 1024, half -> half.field("value"))
992 );
993
994 static SharedMemoryBfloat16 create(Accelerator accelerator) {
995 return null;
996 }
997
998 static SharedMemoryBfloat16 createLocal() {
999 return null;
1000 }
1001 }
1002
1003 private interface PrivateArrayBfloat16 extends NonMappableIface {
1004 BF16 array(int index);
1005
1006 DeviceSchema<PrivateArrayBfloat16> deviceSchema = DeviceSchema.of(PrivateArrayBfloat16.class, arr ->
1007 arr.array("array", 16, half -> half.field("value"))
1008 );
1009
1010 static PrivateArrayBfloat16 create(Accelerator accelerator) {
1011 return null;
1012 }
1013
1014 static PrivateArrayBfloat16 createPrivate() {
1015 return null;
1016 }
1017 }
1018
1019 private interface FlatPrivateBfloat16 extends NonMappableIface {
1020 BF16 array(int index);
1021
1022 DeviceSchema<FlatPrivateBfloat16> deviceSchema = DeviceSchema.of(FlatPrivateBfloat16.class, arr ->
1023 arr.array("array", 4,half -> half.field("value"))
1024 );
1025
1026 static FlatPrivateBfloat16 create(Accelerator accelerator) {
1027 return null;
1028 }
1029
1030 static FlatPrivateBfloat16 createPrivate() {
1031 return null;
1032 }
1033 }
1034
1035 @Reflect
1036 public static void matrixMultiplyKernel2DRegisterTilingBFloat16( BF16Array matrixA, BF16Array matrixB, BF16Array matrixC, int size) {
1037 final int BM = 64;
1038 final int BN = 64;
1039 final int BK = 16;
1040 final int TM = 4;
1041 final int TN = 4;
1042
1043 int bx = BIX();
1044 int by = BIY();
1045
1046 int totalResultsBlockTile = BM * BN;
1047 final int numThreadsBlockTile = totalResultsBlockTile / (TM * TN);
1048
1049 final int linearLocalId = LIY() * LSX() + LIX();
1050 final int threadCol = LIX();
1051 final int threadRow = LIY();
1052
1053 SharedMemoryBfloat16 tileA = SharedMemoryBfloat16.createLocal();
1054 SharedMemoryBfloat16 tileB = SharedMemoryBfloat16.createLocal();
1055
1056 int aFrom = by * BM * size;
1057 int bFrom = bx * BN;
1058 int v = bx * BN;
1059 int cFrom = (by * BM * size) + (v);
1060
1061 final int innerRowA = linearLocalId / BK;
1062 final int innerColA = linearLocalId % BK;
1063
1064 final int strideA = numThreadsBlockTile / BK;
1065 final int innerRowB = linearLocalId / BN;
1066 final int innerColB = linearLocalId % BN;
1067
1068 int strideB = numThreadsBlockTile / BN;
1069
1070 PrivateArrayBfloat16 threadResults = PrivateArrayBfloat16.createPrivate();
1071 FlatPrivateBfloat16 regM = FlatPrivateBfloat16.createPrivate();
1072 FlatPrivateBfloat16 regN = FlatPrivateBfloat16.createPrivate();
1073
1074 for (int i = 0; i < (TN * TN); i++) {
1075 BF16 init = BF16.of(0.0f);
1076 threadResults.array(i).value(init.value());
1077 }
1078
1079 for (int bkIdx = 0; bkIdx < size; bkIdx += BK) {
1080 for (int loadOffset = 0; loadOffset < BM; loadOffset += strideA) {
1081 BF16 ha = matrixA.array(((innerRowA + loadOffset) * size + innerColA) + aFrom);
1082 tileA.array((innerRowA + loadOffset) * BK + innerColA).value(ha.value());
1083 }
1084 for (int loadOffset = 0; loadOffset < BK; loadOffset += strideB) {
1085 BF16 hb = matrixB.array(((innerRowB + loadOffset) * size + innerColB) + bFrom);
1086 tileB.array((innerRowB + loadOffset) * BN + innerColB).value(hb.value());
1087 }
1088 barrier();
1089
1090 aFrom += (BK);
1091 int f = BK * size;
1092 bFrom += f;
1093
1094 for (int dotIdx = 0; dotIdx < BK; dotIdx++) {
1095 for (int i = 0; i < TM; i++) {
1096 BF16 ha = tileA.array((threadRow * TM + i) * BK + dotIdx);
1097 regM.array(i).value(ha.value());
1098 }
1099 for (int i = 0; i < TN; i++) {
1100 BF16 hb = tileB.array(dotIdx * BN + threadCol * TN + i);
1101 regN.array(i).value(hb.value());
1102 }
1103 for (int resIdxM = 0; resIdxM < TM; resIdxM++) {
1104 for (int resIdxN = 0; resIdxN < TN; resIdxN++) {
1105 BF16 privA = regM.array(resIdxM);
1106 BF16 privB = regN.array(resIdxN);
1107 BF16 mul = BF16.mul(privA, privB);
1108 BF16 acc = threadResults.array(resIdxM * TN + resIdxN);
1109 acc = BF16.add(acc, mul);
1110 threadResults.array((resIdxM * TN + resIdxN)).value(acc.value());
1111 }
1112 }
1113 }
1114 barrier();
1115 }
1116 for (int resIdxM = 0; resIdxM < TM; resIdxM++) {
1117 for (int resIdxN = 0; resIdxN < TN; resIdxN++) {
1118 BF16 result = threadResults.array(resIdxM * TN + resIdxN);
1119 matrixC.array((((threadRow * TM + resIdxM) * size + threadCol * TN + resIdxN) + (cFrom))).value(result.value());
1120 }
1121 }
1122 }
1123
1124 @Reflect
1125 public static void matrixMultiply2DRegisterTilingHalf( ComputeContext cc, F16Array matrixA, F16Array matrixB, F16Array matrixC, int globalSize) {
1126 cc.dispatchKernel(NDRange.of2D(256, 256, 16, 16),
1127 ()-> matrixMultiplyKernel2DRegisterTilingHalf( matrixA, matrixB, matrixC, globalSize)
1128 );
1129 }
1130
1131 @Reflect
1132 public static void matrixMultiply2DRegisterTilingBFloat16( ComputeContext cc, BF16Array matrixA, BF16Array matrixB, BF16Array matrixC, int globalSize) {
1133 cc.dispatchKernel(NDRange.of2D(256, 256, 16, 16),
1134 ()-> matrixMultiplyKernel2DRegisterTilingBFloat16( matrixA, matrixB, matrixC, globalSize)
1135 );
1136 }
1137
1138 @HatTest
1139 @Reflect
1140 public void matrixMultiply2DRegisterTilingHalf() {
1141 var lookup = MethodHandles.lookup();
1142 var accelerator = new Accelerator(lookup, Backend.FIRST);
1143
1144 final int size = 1024;
1145 var matrixA = F16Array.create(accelerator, size * size);
1146 var matrixB = F16Array.create(accelerator, size * size);
1147
1148 // Matrix for the results
1149 var matrixC = F16Array.create(accelerator, size * size);
1150 var resultSeq = F16Array.create(accelerator, size * size);
1151
1152 // Initialize matrices (A and B have the same size)
1153 Random r = new Random(19);
1154 for (int j = 0; j < matrixA.length(); j++) {
1155 matrixA.array(j).value(F16.floatToF16(r.nextFloat()).value());
1156 matrixB.array(j).value(F16.floatToF16(r.nextFloat()).value());
1157 }
1158
1159 accelerator.compute(cc ->
1160 TestMatMul.matrixMultiply2DRegisterTilingHalf(cc, matrixA, matrixB, matrixC, size));
1161
1162 // Run Seq for reference
1163 runSequential(matrixA, matrixB, resultSeq, size);
1164
1165 for (int i = 0; i < size; i++) {
1166 for (int j = 0; j < size; j++) {
1167 try {
1168 HATAsserts.assertEquals(F16.f16ToFloat(resultSeq.array(i * size + j)),
1169 F16.f16ToFloat(matrixC.array(i * size + j)),
1170 0.01f);
1171 } catch (HATAssertionError hatAssertionError) {
1172 throw new HATExpectedPrecisionError(hatAssertionError.getMessage());
1173 }
1174 }
1175 }
1176 }
1177
1178 @HatTest
1179 @Reflect
1180 public void matrixMultiply2DRegisterTilingBFloat16() {
1181 var lookup = MethodHandles.lookup();
1182 var accelerator = new Accelerator(lookup, Backend.FIRST);
1183
1184 final int size = 1024;
1185 var matrixA = BF16Array.create(accelerator, size * size);
1186 var matrixB = BF16Array.create(accelerator, size * size);
1187
1188 // Matrix for the results
1189 var matrixC = BF16Array.create(accelerator, size * size);
1190 var resultSeq = BF16Array.create(accelerator, size * size);
1191
1192 // Initialize matrices (A and B have the same size)
1193 Random r = new Random(19);
1194 for (int j = 0; j < matrixA.length(); j++) {
1195 matrixA.array(j).value(BF16.float2bfloat16(r.nextFloat()).value());
1196 matrixB.array(j).value(BF16.float2bfloat16(r.nextFloat()).value());
1197 }
1198
1199 accelerator.compute(cc ->
1200 TestMatMul.matrixMultiply2DRegisterTilingBFloat16(cc, matrixA, matrixB, matrixC, size));
1201
1202 // Run Seq for reference
1203 runSequential(matrixA, matrixB, resultSeq, size);
1204
1205 for (int i = 0; i < size; i++) {
1206 for (int j = 0; j < size; j++) {
1207 try {
1208 HATAsserts.assertEquals(BF16.bfloat162float(resultSeq.array(i * size + j)),
1209 BF16.bfloat162float(matrixC.array(i * size + j)),
1210 0.01f);
1211 } catch (HATAssertionError hatAssertionError) {
1212 throw new HATExpectedPrecisionError(hatAssertionError.getMessage());
1213 }
1214 }
1215 }
1216 }
1217 }