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 }