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 }