1 /*
   2  * Copyright (c) 2024, 2026, 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.backend.ffi;
  26 
  27 import optkl.FuncOpParams;
  28 import optkl.ParamVar;
  29 import optkl.codebuilders.CodeBuilder;
  30 
  31 import jdk.incubator.code.*;
  32 import jdk.incubator.code.dialect.core.CoreOp;
  33 import jdk.incubator.code.dialect.java.JavaOp;
  34 import jdk.incubator.code.dialect.java.JavaType;
  35 
  36 import java.lang.foreign.MemoryLayout;
  37 import java.lang.invoke.MethodHandles;
  38 import java.util.ArrayList;
  39 import java.util.HashMap;
  40 import java.util.List;
  41 import java.util.Map;
  42 import java.util.stream.Stream;
  43 
  44 import static optkl.OpHelper.FieldAccess.fieldAccess;
  45 import static optkl.OpHelper.Invoke;
  46 
  47 import static optkl.OpHelper.Invoke.invoke;
  48 
  49 
  50 public class PTXHATKernelBuilder extends CodeBuilder<PTXHATKernelBuilder> {
  51 
  52     Map<Value, PTXRegister> varToRegMap;
  53     List<String> paramNames;
  54     List<Block.Parameter> paramObjects;
  55     Map<Field, PTXRegister> fieldToRegMap;
  56 
  57     HashMap<PTXRegister.Type, Integer> ordinalMap;
  58 
  59     PTXRegister returnReg;
  60     private int addressSize;
  61 
  62     public enum Field {
  63         NTID_X ("ntid.x", false),
  64         CTAID_X ("ctaid.x", false),
  65         TID_X ("tid.x", false),
  66         NTID_Y ("ntid.y", false),
  67         CTAID_Y ("ctaid.y", false),
  68         TID_Y ("tid.y", false),
  69         KC_X ("x", false),
  70         KC_Y ("y", false),
  71         KC_ADDR("kc", true),
  72         KC_MAXX ("maxX", false);
  73 
  74         private final String name;
  75         private final boolean destination;
  76 
  77         Field(String name, boolean destination) {
  78             this.name = name;
  79             this.destination = destination;
  80         }
  81         public String toString() {
  82             return this.name;
  83         }
  84         public boolean isDestination() {return this.destination;}
  85     }
  86 
  87     public PTXHATKernelBuilder(int addressSize) {
  88         varToRegMap = new HashMap<>();
  89         paramNames = new ArrayList<>();
  90         fieldToRegMap = new HashMap<>();
  91         paramObjects = new ArrayList<>();
  92         ordinalMap = new HashMap<>();
  93         this.addressSize = addressSize;
  94     }
  95 
  96     public PTXHATKernelBuilder() {
  97         this(32);
  98     }
  99 
 100     public void ptxHeader(int major, int minor, String target, int addressSize) {
 101         this.addressSize = addressSize;
 102         version().sp().major(major).dot().minor(minor).nl();
 103         target().sp().target(target).nl();
 104         addressSize().sp().size(addressSize);
 105     }
 106 
 107     public void functionHeader(String funcName, boolean entry, CodeType yieldType) {
 108         if (entry) {
 109             visible().sp().entry().sp();
 110         } else {
 111             func().sp();
 112         }
 113         if (!yieldType.toString().equals("void")) {
 114             returnReg = new PTXRegister(getOrdinal(getResultType(yieldType)), getResultType(yieldType));
 115             returnReg.name("%retReg");
 116             oparen().dot().param().sp().paramType(yieldType);
 117             sp().regName(returnReg).cparen().sp();
 118         }
 119         funcName(funcName);
 120     }
 121 
 122     public PTXHATKernelBuilder parameters(List<FuncOpParams.Info> infoList) {
 123         paren(_ ->
 124                 nl()
 125                         .commaNlSeparated(
 126                         infoList,
 127                         info -> {
 128                             ptxIndent().dot().param().sp().paramType(info.javaType);
 129                             sp().regName(info.varOp.varName());
 130                             paramNames.add(info.varOp.varName());
 131                         }
 132                         ).nl()).nl();
 133         return this;
 134     }
 135 
 136     public void blockBody(MethodHandles.Lookup lookup,Block block, Stream<Op> ops) {
 137         if (block.index() == 0) {
 138             for (Block.Parameter p : block.parameters()) {
 139                 ptxIndent().ld().dot().param();
 140                 resultType(p.type(), false).ptxIndent().sp();
 141                 reg(p, getResultType(p.type())).csp().osbrace().regName(paramNames.get(p.index())).csbrace().semicolon().nl();
 142                 paramObjects.add(p);
 143             }
 144         }
 145         nl();
 146         block(block);
 147         colon().nl();
 148         ops.forEach(op -> {
 149             if (invoke(lookup,op) instanceof Invoke invoke && !invoke.isMappableIface()) {
 150                 ptxIndent().convert(lookup,op).nl();
 151             } else {
 152                 ptxIndent().convert(lookup,op).semicolon().nl();
 153             }
 154         });
 155     }
 156 
 157     public void ptxRegisterDecl() {
 158         for (PTXRegister.Type t : ordinalMap.keySet()) {
 159             ptxIndent().reg().sp();
 160             if (t.equals(PTXRegister.Type.U32)) {
 161                 b32();
 162             } else if (t.equals(PTXRegister.Type.U64)) {
 163                 b64();
 164             } else {
 165                 dot().regType(t);
 166             }
 167             ptxIndent().regTypePrefix(t).oabrace().intVal(ordinalMap.get(t)).cabrace().semicolon().nl();
 168         }
 169         nl();
 170     }
 171 
 172     public void functionPrologue() {
 173         obrace().nl();
 174     }
 175 
 176     public void functionEpilogue() {
 177         cbrace();
 178     }
 179 
 180 
 181     public PTXHATKernelBuilder convert(MethodHandles.Lookup lookup,Op op) {
 182         switch (op) {
 183             case JavaOp.FieldAccessOp.FieldLoadOp $ -> fieldLoad(lookup,$);
 184             case JavaOp.FieldAccessOp.FieldStoreOp $ -> fieldStore($);
 185             case JavaOp.BinaryOp $ -> binaryOperation($);
 186             case JavaOp.CompareOp $ -> compareOperation($);
 187             case JavaOp.ConvOp $ -> conv($);
 188             case CoreOp.ConstantOp $ -> constant($);
 189             case CoreOp.YieldOp $ -> javaYield($);
 190             case JavaOp.InvokeOp $ -> methodCall(invoke(lookup,$));
 191             case CoreOp.VarOp $ when ParamVar.of($) != null -> varFuncDeclaration($);
 192             case CoreOp.VarOp $ -> varDeclaration($);
 193             case CoreOp.ReturnOp $ -> ret($);
 194             case JavaOp.BreakOp $ -> javaBreak($);
 195             case CoreOp.BranchOp $ -> branch($);
 196             case CoreOp.ConditionalBranchOp $ -> condBranch($);
 197             case JavaOp.NegOp $ -> neg($);
 198             case PTXPtrOp $ -> ptxPtr($);
 199             default -> throw new IllegalStateException("op translation doesn't exist");
 200         }
 201         return this;
 202     }
 203 /*
 204     private void hatThreadOp(HATThreadOp threadOp) {
 205         switch (threadOp) {
 206             case HATThreadOp.HAT_GI.HAT_GIX $ -> hatGix($);
 207             case HATThreadOp.HAT_GI.HAT_GIY $ -> hatGiy($);
 208             case HATThreadOp.HAT_GS.HAT_GSX $ -> hatGsx($);
 209             case HATThreadOp.HAT_GS.HAT_GSY $ -> hatGsy($);
 210             default -> throw new IllegalStateException("thread op translation doesn't exist");
 211         }
 212     }
 213 
 214     public void hatGix(HATThreadOp.HAT_GI.HAT_GIX op) {
 215         ensureThreadXRegs();
 216         if (!fieldToRegMap.containsKey(Field.KC_X)) {
 217             mad().lo().s32().sp().fieldReg(Field.KC_X).csp().fieldReg(Field.CTAID_X)
 218                     .csp().fieldReg(Field.NTID_X).csp().fieldReg(Field.TID_X).ptxNl();
 219         }
 220         mov().u32().sp().resultReg(op, PTXRegister.Type.U32).csp().fieldReg(Field.KC_X);
 221     }
 222 
 223     public void hatGsx(HATThreadOp.HAT_GS.HAT_GSX op) {
 224         ensureKcAddr();
 225         ld().global().u32().sp().resultReg(op, PTXRegister.Type.U32).csp()
 226                 .address(fieldToRegMap.get(Field.KC_ADDR).name(), 16);
 227     }
 228 
 229     public void hatGiy(HATThreadOp.HAT_GI.HAT_GIY op) {
 230         ensureThreadYRegs();
 231         if (!fieldToRegMap.containsKey(Field.KC_Y)) {
 232             mad().lo().s32().sp().fieldReg(Field.KC_Y).csp().fieldReg(Field.CTAID_Y)
 233                     .csp().fieldReg(Field.NTID_Y).csp().fieldReg(Field.TID_Y).ptxNl();
 234         }
 235         mov().u32().sp().resultReg(op, PTXRegister.Type.U32).csp().fieldReg(Field.KC_Y);
 236     }
 237 
 238     public void hatGsy(HATThreadOp.HAT_GS.HAT_GSY op) {
 239         ensureKcAddr();
 240         ld().global().u32().sp().resultReg(op, PTXRegister.Type.U32).csp()
 241                 .address(fieldToRegMap.get(Field.KC_ADDR).name(), 20);
 242     } */
 243 
 244     private void ensureKcAddr() {
 245         if (!fieldToRegMap.containsKey(Field.KC_ADDR)) {
 246             cvta().to().global().size().sp().fieldReg(Field.KC_ADDR).csp()
 247                     .reg(paramObjects.get(paramNames.indexOf(Field.KC_ADDR.toString())), addressType()).ptxNl();
 248         }
 249     }
 250 
 251     private void ensureThreadXRegs() {
 252         if (!fieldToRegMap.containsKey(Field.NTID_X)) {
 253             mov().u32().sp().fieldReg(Field.NTID_X).csp().percent().regName(Field.NTID_X.toString()).ptxNl();
 254             mov().u32().sp().fieldReg(Field.CTAID_X).csp().percent().regName(Field.CTAID_X.toString()).ptxNl();
 255             mov().u32().sp().fieldReg(Field.TID_X).csp().percent().regName(Field.TID_X.toString()).ptxNl();
 256         }
 257     }
 258 
 259     private void ensureThreadYRegs() {
 260         if (!fieldToRegMap.containsKey(Field.NTID_Y)) {
 261             mov().u32().sp().fieldReg(Field.NTID_Y).csp().percent().regName(Field.NTID_Y.toString()).ptxNl();
 262             mov().u32().sp().fieldReg(Field.CTAID_Y).csp().percent().regName(Field.CTAID_Y.toString()).ptxNl();
 263             mov().u32().sp().fieldReg(Field.TID_Y).csp().percent().regName(Field.TID_Y.toString()).ptxNl();
 264         }
 265     }
 266 
 267     public void ptxPtr(PTXPtrOp op) {
 268         PTXRegister source;
 269         int offset = (int) op.boundSchema.groupLayout().byteOffset(MemoryLayout.PathElement.groupElement(op.fieldName));
 270 
 271         if (op.fieldName.equals("array")) {
 272             source = new PTXRegister(incrOrdinal(addressType()), addressType());
 273             addKeyword().s64().sp().regName(source).csp().reg(op.operands().get(0)).csp().reg(op.operands().get(1)).ptxNl();
 274         } else {
 275             source = getReg(op.operands().getFirst());
 276         }
 277 
 278         if (op.resultType.toString().equals("void")) {
 279             st().global().dot().regType(op.operands().getLast()).sp().address(source.name(), offset).csp().reg(op.operands().getLast());
 280         } else {
 281             ld().global().resultType(op.resultType(), true).sp().reg(op.result(), getResultType(op.resultType())).csp().address(source.name(), offset);
 282         }
 283     }
 284 
 285     public void fieldLoad(MethodHandles.Lookup lookup,JavaOp.FieldAccessOp.FieldLoadOp fieldLoadOp) {
 286 
 287         var fieldAccess = fieldAccess(lookup,fieldLoadOp);
 288         if (fieldAccess.named(Field.KC_X.toString())) {
 289             if (!fieldToRegMap.containsKey(Field.KC_X)) {
 290                 loadKcX(fieldLoadOp.result());
 291             } else {
 292                 mov().u32().sp().resultReg(fieldLoadOp, PTXRegister.Type.U32).csp().fieldReg(Field.KC_X);
 293             }
 294         } else if (fieldAccess.named(Field.KC_MAXX.toString())) {
 295             if (!fieldToRegMap.containsKey(Field.KC_X)) {
 296                 loadKcX(fieldLoadOp.operands().getFirst());
 297             }
 298             ld().global().u32().sp().fieldReg(Field.KC_MAXX, fieldLoadOp.result()).csp()
 299                     .address(fieldToRegMap.get(Field.KC_ADDR).name(), 4);
 300         } else {
 301             ld().global().u32().sp().resultReg(fieldLoadOp, PTXRegister.Type.U64).csp().reg(fieldLoadOp.operands().getFirst());
 302         }
 303     }
 304 
 305     public void loadKcX(Value value) {
 306         cvta().to().global().size().sp().fieldReg(Field.KC_ADDR).csp()
 307                 .reg(paramObjects.get(paramNames.indexOf(Field.KC_ADDR.toString())), addressType()).ptxNl();
 308         mov().u32().sp().fieldReg(Field.NTID_X).csp().percent().regName(Field.NTID_X.toString()).ptxNl();
 309         mov().u32().sp().fieldReg(Field.CTAID_X).csp().percent().regName(Field.CTAID_X.toString()).ptxNl();
 310         mov().u32().sp().fieldReg(Field.TID_X).csp().percent().regName(Field.TID_X.toString()).ptxNl();
 311         mad().lo().s32().sp().fieldReg(Field.KC_X, value).csp().fieldReg(Field.CTAID_X)
 312                 .csp().fieldReg(Field.NTID_X).csp().fieldReg(Field.TID_X).ptxNl();
 313         st().global().u32().sp().address(fieldToRegMap.get(Field.KC_ADDR).name()).csp().fieldReg(Field.KC_X);
 314     }
 315 
 316     public void fieldStore(JavaOp.FieldAccessOp.FieldStoreOp op) {
 317         // TODO: fix
 318         st().global().u64().sp().resultReg(op, PTXRegister.Type.U64).csp().reg(op.operands().getFirst());
 319     }
 320     // this might be duplication of CodeBuilder symbol....
 321 @Override public
 322 PTXHATKernelBuilder symbol(Op op) {
 323         return switch (op) {
 324             case JavaOp.ModOp _ -> remKw();
 325             case JavaOp.MulOp _ -> mulKeword();
 326             case JavaOp.DivOp _ -> divKeyword();
 327             case JavaOp.AddOp _ -> addKeyword();
 328             case JavaOp.SubOp _ -> subKeyword();
 329             case JavaOp.LtOp _ -> ltKeyword();
 330             case JavaOp.GtOp _ -> gtKeyword();
 331             case JavaOp.LeOp _ -> le();
 332             case JavaOp.GeOp _ -> ge();
 333             case JavaOp.NeqOp _ -> neKeyword();
 334             case JavaOp.EqOp _ -> eqKeyword();
 335             case JavaOp.OrOp _ -> or();
 336             case JavaOp.AndOp _ -> and();
 337             case JavaOp.XorOp _ -> xor();
 338             case JavaOp.LshlOp _ -> shl();
 339             case JavaOp.AshrOp _, JavaOp.LshrOp _ -> shr();
 340             default -> throw new IllegalStateException("Unexpected value");
 341         };
 342     }
 343 
 344     public void binaryOperation(JavaOp.BinaryOp op) {
 345         symbol(op);
 346         if (getResultType(op.resultType()).getBasicType().equals(PTXRegister.Type.BasicType.FLOATING)
 347                 && (op instanceof JavaOp.DivOp || op instanceof JavaOp.MulOp)) {
 348             rn();
 349         } else if (!getResultType(op.resultType()).getBasicType().equals(PTXRegister.Type.BasicType.FLOATING)
 350                 && op instanceof JavaOp.MulOp) {
 351             lo();
 352         }
 353         resultType(op.resultType(), true).sp();
 354         resultReg(op, getResultType(op.resultType()));
 355         csp();
 356         reg(op.operands().getFirst());
 357         csp();
 358         reg(op.operands().get(1));
 359     }
 360 
 361     public void compareOperation(JavaOp.CompareOp op) {
 362         setp().dot();
 363         symbol(op).resultType(op.operands().getFirst().type(), true).sp();
 364         resultReg(op, PTXRegister.Type.PREDICATE);
 365         csp();
 366         reg(op.operands().getFirst());
 367         csp();
 368         reg(op.operands().get(1));
 369     }
 370 
 371     public void conv(JavaOp.ConvOp op) {
 372         if (op.resultType().equals(JavaType.LONG)) {
 373             if (isIndex(op)) {
 374                 mulKeword().wide().s32().sp().resultReg(op, PTXRegister.Type.U64).csp()
 375                         .reg(op.operands().getFirst()).csp().intVal(4);
 376             } else {
 377                 cvt().u64().dot().regType(op.operands().getFirst()).sp()
 378                         .resultReg(op, PTXRegister.Type.U64).csp().reg(op.operands().getFirst()).ptxNl();
 379             }
 380         } else if (op.resultType().equals(JavaType.FLOAT)) {
 381             cvt().rn().f32().dot().regType(op.operands().getFirst()).sp()
 382                     .resultReg(op, PTXRegister.Type.F32).csp().reg(op.operands().getFirst());
 383         } else if (op.resultType().equals(JavaType.DOUBLE)) {
 384             cvt();
 385             if (op.operands().getFirst().type().equals(JavaType.INT)) {
 386                 rn();
 387             }
 388             f64().dot().regType(op.operands().getFirst()).sp()
 389                     .resultReg(op, PTXRegister.Type.F64).csp().reg(op.operands().getFirst());
 390         } else if (op.resultType().equals(JavaType.INT)) {
 391             cvt();
 392             if (op.operands().getFirst().type().equals(JavaType.DOUBLE) || op.operands().getFirst().type().equals(JavaType.FLOAT)) {
 393                 rzi();
 394             } else {
 395                 rn();
 396             }
 397             s32().dot().regType(op.operands().getFirst()).sp()
 398                     .resultReg(op, PTXRegister.Type.S32).csp().reg(op.operands().getFirst());
 399         } else {
 400             cvt().rn().s32().dot().regType(op.operands().getFirst()).sp()
 401                     .resultReg(op, PTXRegister.Type.S32).csp().reg(op.operands().getFirst());
 402         }
 403     }
 404 
 405 
 406 
 407 
 408 
 409     public static class PTXRegister {
 410         private String name;
 411         private final Type type;
 412 
 413         public enum Type {
 414             S8 (8, BasicType.SIGNED, "s8", "%s"),
 415             S16 (16, BasicType.SIGNED, "s16", "%s"),
 416             S32 (32, BasicType.SIGNED, "s32", "%s"),
 417             S64 (64, BasicType.SIGNED, "s64", "%sd"),
 418             U8 (8, BasicType.UNSIGNED, "u8", "%r"),
 419             U16 (16, BasicType.UNSIGNED, "u16", "%r"),
 420             U32 (32, BasicType.UNSIGNED, "u32", "%r"),
 421             U64 (64, BasicType.UNSIGNED, "u64", "%rd"),
 422             F16 (16, BasicType.FLOATING, "f16", "%f"),
 423             F16X2 (16, BasicType.FLOATING, "f16", "%f"),
 424             F32 (32, BasicType.FLOATING, "f32", "%f"),
 425             F64 (64, BasicType.FLOATING, "f64", "%fd"),
 426             B8 (8, BasicType.BIT, "b8", "%b"),
 427             B16 (16, BasicType.BIT, "b16", "%b"),
 428             B32 (32, BasicType.BIT, "b32", "%b"),
 429             B64 (64, BasicType.BIT, "b64", "%bd"),
 430             B128 (128, BasicType.BIT, "b128", "%b"),
 431             PREDICATE (1, BasicType.PREDICATE, "pred", "%p");
 432 
 433             public enum BasicType {
 434                 SIGNED,
 435                 UNSIGNED,
 436                 FLOATING,
 437                 BIT,
 438                 PREDICATE
 439             }
 440 
 441             private final int size;
 442             private final BasicType basicType;
 443             private final String name;
 444             private final String regPrefix;
 445 
 446             Type(int size, BasicType type, String name, String regPrefix) {
 447                 this.size = size;
 448                 this.basicType = type;
 449                 this.name = name;
 450                 this.regPrefix = regPrefix;
 451             }
 452 
 453             public int getSize() {
 454                 return this.size;
 455             }
 456 
 457             public BasicType getBasicType() {
 458                 return this.basicType;
 459             }
 460 
 461             public String getName() {
 462                 return this.name;
 463             }
 464 
 465             public String getRegPrefix() {
 466                 return this.regPrefix;
 467             }
 468         }
 469 
 470         public PTXRegister(int num, Type type) {
 471             this.type = type;
 472             this.name = type.regPrefix + num;
 473         }
 474 
 475         public String name() {
 476             return this.name;
 477         }
 478 
 479         public void name(String name) {
 480             this.name = name;
 481         }
 482 
 483         public Type type() {
 484             return this.type;
 485         }
 486     }
 487 
 488 
 489     private boolean isIndex(JavaOp.ConvOp op) {
 490         for (Op.Result r : op.result().uses()) {
 491             if (r.op() instanceof PTXPtrOp) return true;
 492         }
 493         return false;
 494     }
 495 
 496     public void constant(CoreOp.ConstantOp op) {
 497         mov().resultType(op.resultType(), false).sp().resultReg(op, getResultType(op.resultType())).csp();
 498         if (op.resultType().toString().equals("float")) {
 499             if (op.value().toString().equals("0.0")) {
 500                 floatVal("00000000");
 501             } else {
 502                 floatVal(Integer.toHexString(Float.floatToIntBits(Float.parseFloat(op.value().toString()))).toUpperCase());
 503             }
 504         } else {
 505             constant(op.value().toString());
 506         }
 507     }
 508 
 509     public void javaYield(CoreOp.YieldOp op) {
 510         exit();
 511     }
 512 
 513     // S32Array and S32Array2D functions can be deleted after schema is done
 514     public void methodCall(Invoke invoke) {
 515        // Invoke invoke = Invoke.invokeOpHelper(MethodHandles.lookup(),invokeOp);
 516         switch (invoke.op().invokeReference().toString()) {
 517             // S32Array functions
 518             case "hat.buffer.S32Array::array(long)int" -> {
 519                 PTXRegister temp = new PTXRegister(incrOrdinal(addressType()), addressType());
 520                 addKeyword().s64().sp().regName(temp).csp().reg(invoke.op().operands().getFirst()).csp().reg(invoke.op().operands().get(1)).ptxNl();
 521                 ld().global().u32().sp().resultReg(invoke.op(), PTXRegister.Type.U32).csp().address(temp.name(), 4);
 522             }
 523             case "hat.buffer.S32Array::array(long, int)void" -> {
 524                 PTXRegister temp = new PTXRegister(incrOrdinal(addressType()), addressType());
 525                 addKeyword().s64().sp().regName(temp).csp().reg(invoke.op().operands().getFirst()).csp().reg(invoke.op().operands().get(1)).ptxNl();
 526                 st().global().u32().sp().address(temp.name(), 4).csp().reg(invoke.op().operands().get(2));
 527             }
 528             case "hat.buffer.S32Array::length()int" -> {
 529                 ld().global().u32().sp().resultReg(invoke.op(), PTXRegister.Type.U32).csp().address(getReg(invoke.op().operands().getFirst()).name());
 530             }
 531             // S32Array2D functions
 532             case "hat.buffer.S32Array2D::array(long, int)void" -> {
 533                 PTXRegister temp = new PTXRegister(incrOrdinal(addressType()), addressType());
 534                 addKeyword().s64().sp().regName(temp).csp().reg(invoke.op().operands().getFirst()).csp().reg(invoke.op().operands().get(1)).ptxNl();
 535                 st().global().u32().sp().address(temp.name(), 8).csp().reg(invoke.op().operands().get(2));
 536             }
 537             case "hat.buffer.S32Array2D::width()int" -> {
 538                 ld().global().u32().sp().resultReg(invoke.op(), PTXRegister.Type.U32).csp().address(getReg(invoke.op().operands().getFirst()).name());
 539             }
 540             case "hat.buffer.S32Array2D::height()int" -> {
 541                 ld().global().u32().sp().resultReg(invoke.op(), PTXRegister.Type.U32).csp().address(getReg(invoke.op().operands().getFirst()).name(), 4);
 542             }
 543             // Java Math function
 544             case "java.lang.Math::sqrt(double)double" -> {
 545                 sqrt().rn().f64().sp().resultReg(invoke.op(), PTXRegister.Type.F64).csp().reg(invoke.op().operands().getFirst()).semicolon();
 546             }
 547             default -> {
 548                 obrace().nl().ptxIndent();
 549                 for (int i = 0; i < invoke.op().operands().size(); i++) {
 550                     dot().param().sp().paramType(invoke.op().operands().get(i).type()).sp().param().intVal(i).ptxNl();
 551                     st().dot().param().paramType(invoke.op().operands().get(i).type()).sp().osbrace().param().intVal(i).csbrace().csp().reg(invoke.op().operands().get(i)).ptxNl();
 552                 }
 553                 dot().param().sp().paramType(invoke.op().resultType()).sp().retVal().ptxNl();
 554                 call().uni().sp().oparen().retVal().cparen().csp().id(invoke.name()).csp();
 555                 final int[] counter = {0};
 556                 paren(_ ->
 557                         commaSpaceSeparated(
 558                                 invoke.op().operands(),
 559                                 _ -> param().intVal(counter[0]++)
 560                         )
 561                 ).ptxNl();
 562                 ld().dot().param().paramType(invoke.op().resultType()).sp().resultReg(invoke.op(), getResultType(invoke.op().resultType())).csp().osbrace().retVal().csbrace();
 563                 ptxNl().cbrace();
 564             }
 565         }
 566     }
 567 
 568     public void varDeclaration(CoreOp.VarOp op) {
 569         ld().dot().param().resultType(op.resultType(), false).sp().resultReg(op, addressType()).csp().reg(op.operands().getFirst());
 570     }
 571 
 572     public void varFuncDeclaration(CoreOp.VarOp op) {
 573         ld().dot().param().resultType(op.resultType(), false).sp().resultReg(op, addressType()).csp().reg(op.operands().getFirst());
 574     }
 575 
 576     public void ret(CoreOp.ReturnOp op) {
 577         if (!op.operands().isEmpty()) {
 578             st().dot().param();
 579             if (returnReg.type().equals(PTXRegister.Type.U32)) {
 580                 b32();
 581             } else if (returnReg.type().equals(PTXRegister.Type.U64)) {
 582                 b64();
 583             } else {
 584                 dot().regType(returnReg.type());
 585             }
 586             sp().osbrace().regName(returnReg).csbrace().csp().reg(op.operands().getFirst()).ptxNl();
 587         }
 588         ret();
 589     }
 590 
 591     public void javaBreak(JavaOp.BreakOp op) {
 592         brkpt();
 593     }
 594 
 595     public void branch(CoreOp.BranchOp op) {
 596         loadBlockParams(op.successors().getFirst());
 597         bra().sp().block(op.successors().getFirst().targetBlock());
 598     }
 599 
 600     public void condBranch(CoreOp.ConditionalBranchOp op) {
 601         loadBlockParams(op.successors().getFirst());
 602         loadBlockParams(op.successors().getLast());
 603         at().reg(op.operands().getFirst()).sp()
 604                 .bra().sp().block(op.successors().getFirst().targetBlock()).ptxNl();
 605         bra().sp().block(op.successors().getLast().targetBlock());
 606     }
 607 
 608     public void neg(JavaOp.NegOp op) {
 609         neg().resultType(op.resultType(), true).sp().reg(op.result(), getResultType(op.resultType())).csp().reg(op.operands().getFirst());
 610     }
 611 
 612     /*
 613      * Helper functions for printing blocks and variables
 614      */
 615 
 616     public void loadBlockParams(Block.Reference block) {
 617         for (int i = 0; i < block.arguments().size(); i++) {
 618             Block.Parameter p = block.targetBlock().parameters().get(i);
 619             mov().resultType(p.type(), false).sp().reg(p, getResultType(p.type()))
 620                     .csp().reg(block.arguments().get(i)).ptxNl();
 621         }
 622     }
 623 
 624     public PTXHATKernelBuilder block(Block block) {
 625         return type("block_").intVal(block.index());
 626     }
 627 
 628     public PTXHATKernelBuilder fieldReg(Field ref) {
 629         if (fieldToRegMap.containsKey(ref)) {
 630             return regName(fieldToRegMap.get(ref));
 631         }
 632         if (ref.isDestination()) {
 633             fieldToRegMap.putIfAbsent(ref, new PTXRegister(incrOrdinal(addressType()), addressType()));
 634         } else {
 635             fieldToRegMap.putIfAbsent(ref, new PTXRegister(incrOrdinal(PTXRegister.Type.U32), PTXRegister.Type.U32));
 636         }
 637         return regName(fieldToRegMap.get(ref));
 638     }
 639 
 640     public PTXHATKernelBuilder fieldReg(Field ref, Value value) {
 641         if (fieldToRegMap.containsKey(ref)) {
 642             return regName(fieldToRegMap.get(ref));
 643         }
 644         if (ref.isDestination()) {
 645             fieldToRegMap.putIfAbsent(ref, new PTXRegister(getOrdinal(addressType()), addressType()));
 646             return reg(value, addressType());
 647         } else {
 648             fieldToRegMap.putIfAbsent(ref, new PTXRegister(getOrdinal(PTXRegister.Type.U32), PTXRegister.Type.U32));
 649             return reg(value, PTXRegister.Type.U32);
 650         }
 651     }
 652 
 653     public Field getFieldObj(String fieldName) {
 654         for (Field f : fieldToRegMap.keySet()) {
 655             if (f.toString().equals(fieldName)) return f;
 656         }
 657         throw new IllegalStateException("no existing field");
 658     }
 659 
 660     public PTXHATKernelBuilder resultReg(Op op, PTXRegister.Type type) {
 661         return id(addReg(op.result(), type));
 662     }
 663 
 664     public PTXHATKernelBuilder reg(Value val, PTXRegister.Type type) {
 665         if (varToRegMap.containsKey(val)) {
 666             return regName(getReg(val));
 667         } else {
 668             return id(addReg(val, type));
 669         }
 670     }
 671 
 672     public PTXHATKernelBuilder reg(Value val) {
 673         return regName(getReg(val));
 674     }
 675 
 676     public PTXRegister getReg(Value val) {
 677         if (varToRegMap.get(val) == null && val instanceof Op.Result result && result.op() instanceof JavaOp.FieldAccessOp.FieldLoadOp fieldLoadOp) {
 678             return fieldToRegMap.get(getFieldObj(fieldLoadOp.fieldReference().name()));
 679         }
 680         if (varToRegMap.containsKey(val)) {
 681             return varToRegMap.get(val);
 682         } else {
 683             throw new IllegalStateException("var to reg mapping doesn't exist");
 684         }
 685     }
 686 
 687     public String addReg(Value val, PTXRegister.Type type) {
 688         if (varToRegMap.containsKey(val)) {
 689             return varToRegMap.get(val).name();
 690         }
 691         varToRegMap.put(val, new PTXRegister(incrOrdinal(type), type));
 692         return varToRegMap.get(val).name();
 693     }
 694 
 695     public Integer getOrdinal(PTXRegister.Type type) {
 696         ordinalMap.putIfAbsent(type, 1);
 697         return ordinalMap.get(type);
 698     }
 699 
 700     public Integer incrOrdinal(PTXRegister.Type type) {
 701         ordinalMap.putIfAbsent(type, 1);
 702         int out = ordinalMap.get(type);
 703         ordinalMap.put(type, out + 1);
 704         return out;
 705     }
 706 
 707     public PTXHATKernelBuilder size() {
 708         return (addressSize == 32) ? u32() : u64();
 709     }
 710 
 711     public PTXRegister.Type addressType() {
 712         return (addressSize == 32) ? PTXRegister.Type.U32 : PTXRegister.Type.U64;
 713     }
 714 
 715     public PTXHATKernelBuilder resultType(CodeType type, boolean signedResult) {
 716         PTXRegister.Type res = getResultType(type);
 717         if (signedResult && (res == PTXRegister.Type.U32)) return s32();
 718         return dot().type(getResultType(type).getName());
 719     }
 720 
 721     public PTXHATKernelBuilder paramType(CodeType type) {
 722         PTXRegister.Type res = getResultType(type);
 723         if (res == PTXRegister.Type.U32) return b32();
 724         if (res == PTXRegister.Type.U64) return b64();
 725         return dot().type(getResultType(type).getName());
 726     }
 727 
 728     public PTXRegister.Type getResultType(CodeType type) {
 729         switch (type.toString()) {
 730             case "float" -> {
 731                 return PTXRegister.Type.F32;
 732             }
 733             case "double" -> {
 734                 return PTXRegister.Type.F64;
 735             }
 736             case "int" -> {
 737                 return PTXRegister.Type.U32;
 738             }
 739             case "boolean" -> {
 740                 return PTXRegister.Type.PREDICATE;
 741             }
 742             default -> {
 743                 return PTXRegister.Type.U64;
 744             }
 745         }
 746     }
 747 
 748     /*
 749      * Basic CodeBuilder functions
 750      */
 751 
 752     // used for parameter list
 753     // prints out items separated by a comma then new line
 754     // Don't know why this was overriding with the same code grf.
 755    /* @Override
 756     public <I> PTXHATKernelBuilder commaNlSeparated(Iterable<I> iterable, Consumer<I> c) {
 757         StreamCounter.of(iterable, (counter, t) -> {
 758             if (counter.isNotFirst()) {
 759                 comma().nl();
 760             }
 761             c.accept(t);
 762         });
 763         return self();
 764     }
 765 */
 766     public PTXHATKernelBuilder address(String address) {
 767         return osbrace().constant(address).csbrace();
 768     }
 769 
 770     public PTXHATKernelBuilder address(String address, int offset) {
 771         osbrace().constant(address);
 772         if (offset == 0) {
 773             return csbrace();
 774         } else if (offset > 0) {
 775             plus();
 776         }
 777         return intVal(offset).csbrace();
 778     }
 779 
 780     public PTXHATKernelBuilder ptxNl() {
 781         return semicolon().nl().ptxIndent();
 782     }
 783 
 784 
 785     public PTXHATKernelBuilder param() {
 786         return keyword("param");
 787     }
 788 
 789     public PTXHATKernelBuilder global() {
 790         return dot().keyword("global");
 791     }
 792 
 793     public PTXHATKernelBuilder rn() {
 794         return dot().keyword("rn");
 795     }
 796 
 797     public PTXHATKernelBuilder rm() {
 798         return dot().keyword("rm");
 799     }
 800 
 801     public PTXHATKernelBuilder rzi() {
 802         return dot().keyword("rzi");
 803     }
 804 
 805     public PTXHATKernelBuilder to() {
 806         return dot().keyword("to");
 807     }
 808 
 809     public PTXHATKernelBuilder lo() {
 810         return dot().keyword("lo");
 811     }
 812 
 813     public PTXHATKernelBuilder wide() {
 814         return dot().keyword("wide");
 815     }
 816 
 817     public PTXHATKernelBuilder uni() {
 818         return dot().keyword("uni");
 819     }
 820 
 821     public PTXHATKernelBuilder sat() {
 822         return dot().keyword("sat");
 823     }
 824 
 825     public PTXHATKernelBuilder ftz() {
 826         return dot().keyword("ftz");
 827     }
 828 
 829     public PTXHATKernelBuilder approx() {
 830         return dot().keyword("approx");
 831     }
 832 
 833     public PTXHATKernelBuilder mov() {
 834         return keyword("mov");
 835     }
 836 
 837     public PTXHATKernelBuilder setp() {
 838         return keyword("setp");
 839     }
 840 
 841     public PTXHATKernelBuilder selp() {
 842         return keyword("selp");
 843     }
 844 
 845     public PTXHATKernelBuilder ld() {
 846         return keyword("ld");
 847     }
 848 
 849     public PTXHATKernelBuilder st() {
 850         return keyword("st");
 851     }
 852 
 853     public PTXHATKernelBuilder cvt() {
 854         return keyword("cvt");
 855     }
 856 
 857     public PTXHATKernelBuilder bra() {
 858         return keyword("bra");
 859     }
 860 
 861     public PTXHATKernelBuilder ret() {
 862         return keyword("ret");
 863     }
 864 
 865     public PTXHATKernelBuilder remKw() {
 866         return keyword("rem");
 867     }
 868 
 869     public PTXHATKernelBuilder mulKeword() {
 870         return keyword("mul");
 871     }
 872 
 873     public PTXHATKernelBuilder divKeyword() {
 874         return keyword("div");
 875     }
 876 
 877     public PTXHATKernelBuilder rcp() {
 878         return keyword("rcp");
 879     }
 880 
 881     public PTXHATKernelBuilder addKeyword() {
 882         return keyword("add");
 883     }
 884 
 885     public PTXHATKernelBuilder subKeyword() {
 886         return keyword("sub");
 887     }
 888 
 889     public PTXHATKernelBuilder ltKeyword() {
 890         return keyword("lt");
 891     }
 892 
 893     public PTXHATKernelBuilder gtKeyword() {
 894         return keyword("gt");
 895     }
 896 
 897     public PTXHATKernelBuilder le() {
 898         return keyword("le");
 899     }
 900 
 901     public PTXHATKernelBuilder ge() {
 902         return keyword("ge");
 903     }
 904 
 905     public PTXHATKernelBuilder geu() {
 906         return keyword("geu");
 907     }
 908 
 909     public PTXHATKernelBuilder neKeyword() {
 910         return keyword("ne");
 911     }
 912 
 913     public PTXHATKernelBuilder eqKeyword() {
 914         return keyword("eq");
 915     }
 916 
 917     public PTXHATKernelBuilder xor() {
 918         return keyword("xor");
 919     }
 920 
 921     public PTXHATKernelBuilder or() {
 922         return keyword("or");
 923     }
 924 
 925     public PTXHATKernelBuilder and() {
 926         return keyword("and");
 927     }
 928 
 929     public PTXHATKernelBuilder cvta() {
 930         return keyword("cvta");
 931     }
 932 
 933     public PTXHATKernelBuilder mad() {
 934         return keyword("mad");
 935     }
 936 
 937     public PTXHATKernelBuilder fma() {
 938         return keyword("fma");
 939     }
 940 
 941     public PTXHATKernelBuilder sqrt() {
 942         return keyword("sqrt");
 943     }
 944 
 945     public PTXHATKernelBuilder abs() {
 946         return keyword("abs");
 947     }
 948 
 949     public PTXHATKernelBuilder ex2() {
 950         return keyword("ex2");
 951     }
 952 
 953     public PTXHATKernelBuilder shl() {
 954         return keyword("shl");
 955     }
 956 
 957     public PTXHATKernelBuilder shr() {
 958         return keyword("shr");
 959     }
 960 
 961     public PTXHATKernelBuilder neg() {
 962         return keyword("neg");
 963     }
 964 
 965     public PTXHATKernelBuilder call() {
 966         return keyword("call");
 967     }
 968 
 969     public PTXHATKernelBuilder exit() {
 970         return keyword("exit");
 971     }
 972 
 973     public PTXHATKernelBuilder brkpt() {
 974         return keyword("brkpt");
 975     }
 976 
 977     public PTXHATKernelBuilder ptxIndent() {
 978         return sp().sp().sp().sp();
 979     }
 980 
 981     public PTXHATKernelBuilder u32() {
 982         return dot().type(PTXRegister.Type.U32.getName());
 983     }
 984 
 985     public PTXHATKernelBuilder s32() {
 986         return dot().type(PTXRegister.Type.S32.getName());
 987     }
 988 
 989     public PTXHATKernelBuilder f32() {
 990         return dot().type(PTXRegister.Type.F32.getName());
 991     }
 992 
 993     public PTXHATKernelBuilder b32() {
 994         return dot().type(PTXRegister.Type.B32.getName());
 995     }
 996 
 997     public PTXHATKernelBuilder u64() {
 998         return dot().type(PTXRegister.Type.U64.getName());
 999     }
1000 
1001     public PTXHATKernelBuilder s64() {
1002         return dot().type(PTXRegister.Type.S64.getName());
1003     }
1004 
1005     public PTXHATKernelBuilder f64() {
1006         return dot().type(PTXRegister.Type.F64.getName());
1007     }
1008 
1009     public PTXHATKernelBuilder b64() {
1010         return dot().type(PTXRegister.Type.B64.getName());
1011     }
1012 
1013     public PTXHATKernelBuilder version() {
1014         return dot().keyword("version");
1015     }
1016 
1017     public PTXHATKernelBuilder target() {
1018         return dot().keyword("target");
1019     }
1020 
1021     public PTXHATKernelBuilder addressSize() {
1022         return dot().keyword("address_size");
1023     }
1024 
1025     public PTXHATKernelBuilder major(int major) {
1026         return intVal(major);
1027     }
1028 
1029     public PTXHATKernelBuilder minor(int minor) {
1030         return intVal(minor);
1031     }
1032 
1033     public PTXHATKernelBuilder target(String target) {
1034         return keyword(target);
1035     }
1036 
1037     public PTXHATKernelBuilder size(int addressSize) {
1038         return intVal(addressSize);
1039     }
1040 
1041 
1042 
1043     public PTXHATKernelBuilder visible() {
1044         return dot().keyword("visible");
1045     }
1046 
1047     public PTXHATKernelBuilder entry() {
1048         return dot().keyword("entry");
1049     }
1050 
1051     public PTXHATKernelBuilder func() {
1052         return dot().keyword("func");
1053     }
1054 
1055     public PTXHATKernelBuilder oabrace() {
1056         return symbol("<");
1057     }
1058 
1059     public PTXHATKernelBuilder cabrace() {
1060         return symbol(">");
1061     }
1062 
1063     public PTXHATKernelBuilder regName(PTXRegister reg) {
1064         return id(reg.name());
1065     }
1066 
1067     public PTXHATKernelBuilder regName(String regName) {
1068         return id(regName);
1069     }
1070 
1071     public PTXHATKernelBuilder regType(Value val) {
1072         return keyword(getReg(val).type().getName());
1073     }
1074 
1075     public PTXHATKernelBuilder regType(PTXRegister.Type t) {
1076         return keyword(t.getName());
1077     }
1078 
1079     public PTXHATKernelBuilder regTypePrefix(PTXRegister.Type t) {
1080         return keyword(t.getRegPrefix());
1081     }
1082 
1083     public PTXHATKernelBuilder reg() {
1084         return dot().keyword("reg");
1085     }
1086 
1087     public PTXHATKernelBuilder retVal() {
1088         return keyword("retval");
1089     }
1090 
1091     public PTXHATKernelBuilder intVal(int i) {
1092         return constant(String.valueOf(i));
1093     }
1094 
1095     public PTXHATKernelBuilder floatVal(String s) {
1096         return constant("0f").constant(s);
1097     }
1098 
1099     public PTXHATKernelBuilder doubleVal(String s) {
1100         return constant("0d").constant(s);
1101     }
1102 }