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.
  8  *
  9  * This code is distributed in the hope that it will be useful, but WITHOUT
 10  * ANY WARRANTY; without even the implied warranty of MERCHANTABILITY or
 11  * FITNESS FOR A PARTICULAR PURPOSE.  See the GNU General Public License
 12  * version 2 for more details (a copy is included in the LICENSE file that
 13  * accompanied this code).
 14  *
 15  * You should have received a copy of the GNU General Public License version
 16  * 2 along with this work; if not, write to the Free Software Foundation,
 17  * Inc., 51 Franklin St, Fifth Floor, Boston, MA 02110-1301 USA.
 18  *
 19  * Please contact Oracle, 500 Oracle Parkway, Redwood Shores, CA 94065 USA
 20  * or visit www.oracle.com if you need additional information or have any
 21  * questions.
 22  */
 23 
 24 package oracle.code.onnx.lift;
 25 
 26 import java.lang.reflect.Method;
 27 import java.lang.foreign.ValueLayout;
 28 import java.util.List;
 29 import java.util.Map;
 30 import java.util.Optional;
 31 import java.util.SequencedMap;
 32 import java.util.function.Function;
 33 import java.util.stream.Collectors;
 34 import java.util.stream.IntStream;
 35 import java.util.stream.LongStream;
 36 import jdk.incubator.code.Block;
 37 import jdk.incubator.code.CodeItem;
 38 import jdk.incubator.code.Op;
 39 import jdk.incubator.code.CodeType;
 40 import jdk.incubator.code.Value;
 41 import jdk.incubator.code.dialect.core.CoreOp;
 42 import jdk.incubator.code.dialect.java.JavaType;
 43 import jdk.incubator.code.extern.ExternalizedOp;
 44 import oracle.code.onnx.OnnxOperators;
 45 import oracle.code.onnx.Tensor;
 46 import oracle.code.onnx.ir.OnnxOp;
 47 import oracle.code.onnx.ir.OnnxType;
 48 import oracle.code.onnx.proto.OnnxModel;
 49 
 50 final class JavaTemplate {
 51 
 52     private static final String WEIGHT_FIELD_TEMPLATE = """
 53             final %s %s = load("%s", %dL, %d, %s%s);
 54         """;
 55 
 56     private static final String TEMPLATE = """
 57         import java.io.IOException;
 58         import java.io.RandomAccessFile;
 59         import java.lang.foreign.Arena;
 60         import java.lang.foreign.MemorySegment;
 61         import java.nio.channels.FileChannel;
 62         import java.util.List;
 63         import jdk.incubator.code.Reflect;
 64         import oracle.code.onnx.Tensor;
 65 
 66         import static java.util.Optional.*;
 67         import static oracle.code.onnx.OnnxOperators.*;
 68         import static oracle.code.onnx.Tensor.ElementType.*;
 69 
 70         public class %s {
 71 
 72             final Arena arena = Arena.ofAuto();
 73 
 74             MemorySegment mmap(String pathname, long offset, long length) {
 75                 try (var f = new RandomAccessFile(pathname, "r")) {
 76                     return f.getChannel().map(FileChannel.MapMode.READ_ONLY, offset, length < 0 ? f.length() : length, arena);
 77                 } catch (IOException e) {
 78                     throw new RuntimeException(e);
 79                 }
 80             }
 81 
 82             <T> Tensor<T> load(String path, long offset, long length, Tensor.ElementType type, long... shape) {
 83                 return new Tensor<>(arena, mmap(path, offset, length), type, shape);
 84             }
 85 
 86         %s
 87             @Reflect
 88             public Object mainGraph(
 89         %s) {%s
 90             }
 91         }
 92         """;
 93 
 94     static String toJava(OnnxLift.LiftedModelWrapper model, String className) {
 95         Block entryBlock = model.func().bodies().getFirst().entryBlock();
 96         List<Block.Parameter> parameters = entryBlock.parameters();
 97         Function<CodeItem, String> namer = OnnxLift.namer(model.names());
 98         parameters.forEach(namer::apply); // initialize namer with all parameters first
 99 
100         return TEMPLATE.formatted(
101                 className,
102                 weightFields(namer, parameters, model.weights()),
103                 parameters(namer, parameters, model.weights()),
104                 body(namer, entryBlock.ops()));
105     }
106 
107     private static String weightFields(Function<CodeItem, String> namer, List<Block.Parameter> parameters, List<OnnxModel.TensorProto> weights) {
108         StringBuilder out = new StringBuilder();
109         List<jdk.incubator.code.Block.Parameter> weightParams = parameters.subList(parameters.size() - weights.size(), parameters.size());
110         Map<String, oracle.code.onnx.proto.OnnxModel.TensorProto> wMap = weights.stream().collect(Collectors.toUnmodifiableMap(OnnxModel.TensorProto::name, Function.identity()));
111         for (int i = 0; i < weightParams.size(); i++) {
112             Block.Parameter wp = weightParams.get(i);
113             OnnxModel.TensorProto w = wMap.get(namer.apply(wp));
114             String name = OnnxLift.toJavaName(w.name());
115             long[] dims = OnnxLift.joinLongArray(w.dims());
116             String location;
117             long offset;
118             int length;
119             if (w.externalData() instanceof List<OnnxModel.StringStringEntryProto> ssep) {
120                 var map = ssep.stream().collect(Collectors.toUnmodifiableMap(OnnxModel.StringStringEntryProto::key, OnnxModel.StringStringEntryProto::value));
121                 location = map.get("location");
122                 offset = Long.parseLong(map.get("offset"));
123                 length = Integer.parseInt(map.get("length"));
124             } else {
125                 location = name;
126                 offset = 0;
127                 length = -1;
128             }
129             out.append(WEIGHT_FIELD_TEMPLATE.formatted(toJavaType(wp.type()), name, location, offset, length, Tensor.ElementType.fromOnnxId(w.dataType()).name(), dims.length > 0 ? (", " + longJoin(dims)) : ""));
130         }
131         return out.toString();
132     }
133 
134     private static String parameters(Function<CodeItem, String> namer, List<Block.Parameter> parameters, List<OnnxModel.TensorProto> weights) {
135         StringBuilder out = new StringBuilder();
136         int realParamsSize = parameters.size() - weights.size();
137         for (int i = 0; i < realParamsSize; i++) {
138             if (i > 0) {
139                 out.append(",\n");
140             }
141             Block.Parameter param = parameters.get(i);
142             out.append("            ").append(toJavaType(param.type())).append(' ').append(OnnxLift.toJavaName(namer.apply(param)));
143         }
144         return out.toString();
145     }
146 
147     private static String body(Function<CodeItem, String> namer, List<Op> ops) {
148         StringBuilder out = new StringBuilder();
149         for (jdk.incubator.code.Op op : ops) {
150             if (!(op instanceof CoreOp.TupleLoadOp)) {
151                 // lazy tupple loads
152                 out.append("\n        ");
153                 if (!op.resultType().equals(JavaType.VOID)) {
154                     out.append(toJavaType(op.resultType())).append(' ').append(OnnxLift.toJavaName(namer.apply(op.result()))).append(" = ");
155                 }
156                 switch (op) {
157                     case OnnxOp oo -> {
158                         String opName = oo.externalizeOpName();
159                         out.append(opName.substring(opName.lastIndexOf('.') + 1)).append('(');
160                         OnnxOp.OnnxSchema schema = getSchema(oo);
161                         SequencedMap<OnnxOp.OnnxParameter, Object> inputs = oo.onnxInputs();
162                         boolean first = true;
163                         for (OnnxOp.OnnxParameter oi : schema.inputs()) {
164                             if (first) {
165                                 first = false;
166                             } else {
167                                 out.append(", ");
168                             }
169                             out.append(toJava(namer, inputs.get(oi)));
170                         }
171                         Map<String, Object> attrs = oo.onnxAttributes();
172                         for (OnnxOp.OnnxAttribute oa : schema.attributes()) {
173                             if (first) {
174                                 first = false;
175                             } else {
176                                 out.append(", ");
177                             }
178                             Object a = attrs.get(oa.name());
179                             if (a == null) {
180                                 out.append("empty()");
181                             } else if (oa.isOptional()) {
182                                 out.append("of(").append(toString(a)).append(')');
183                             } else {
184                                 out.append(toString(a));
185                             }
186                         }
187                         out.append(");");
188                     }
189                     case CoreOp.TupleOp to -> {
190                         out.append("List.of(");
191                         boolean first = true;
192                         for (jdk.incubator.code.Value te : to.operands()) {
193                             if (first) {
194                                 first = false;
195                             } else {
196                                 out.append(", ");
197                             }
198                             out.append(toJava(namer, te));
199                         }
200                         out.append(");");
201                     }
202                     case CoreOp.ReturnOp ro -> {
203                         out.append("return ").append(toJava(namer, ro.operands().getFirst())).append(';');
204                     }
205                     default -> throw new UnsupportedOperationException(op.toText());
206                 }
207             } else {
208                 namer.apply(op.result());
209             }
210         }
211         return out.toString();
212     }
213 
214     private static String toString(Object o) {
215         return switch (o) {
216             case long[] la -> newArray(la);
217             case float[] fa -> newArray(fa);
218             case Long l -> l.toString() + "L";
219             case Float f -> f.toString() + "F";
220             case String s -> "\"" + s + "\"";
221             case Tensor t -> "Tensor.ofShape(" + newArray(t.shape()) + ", " + toString(getData(t)) + ")";
222             default -> o.toString();
223         };
224     }
225 
226     private static Object getData(Tensor t) {
227         return switch (t.elementType()) {
228             case FLOAT -> t.data().toArray(ValueLayout.JAVA_FLOAT);
229             case INT64 -> t.data().toArray(ValueLayout.JAVA_LONG);
230             default -> throw new UnsupportedOperationException(t.elementType().name() + " tensor type");
231         };
232     }
233 
234     private static String newArray(long[] la) {
235         for (long l : la) {
236             if (l != 0l) {
237                 return "new long[] {" + longJoin(la) + "}";
238             }
239         }
240         return "new long[" + la.length + "]";
241     }
242 
243     private static String longJoin(long[] la) {
244         return LongStream.of(la).mapToObj(d -> String.valueOf(d) + "L").collect(Collectors.joining(", "));
245     }
246 
247     private static String newArray(float[] fa) {
248         for (float f : fa) {
249             if (f != 0f) {
250                 return IntStream.range(0, fa.length).mapToObj(i -> String.valueOf(fa[i]) + "F").collect(Collectors.joining(", ", "new float[] {", "}"));
251             }
252         }
253         return "new float[" + fa.length + "]";
254     }
255 
256     private static String tupleAccessor(Value tuple, int componentIndex) {
257         if (tuple instanceof Op.Result or && or.op() instanceof OnnxOp oo) {
258             String mName = oo.externalizeOpName();
259             mName = mName.substring(mName.lastIndexOf('.') + 1);
260             for (Method m : OnnxOperators.class.getMethods()) {
261                 if (m.getName().equals(mName)) {
262                     return m.getReturnType().getRecordComponents()[componentIndex].getAccessor().getName() + "()";
263                 }
264             }
265             throw new IllegalStateException(mName);
266         }
267         return "get(" + componentIndex + ")"; // fallback to List
268     }
269 
270     private static String toJava(Function<CodeItem, String> namer, Object value) {
271         return switch (value) {
272             case Optional o when o.isEmpty() -> "empty()";
273             case Optional o -> "of(" + toJava(namer, o.get()) + ")";
274             case List l -> "List.of(" + l.stream().map(le -> toJava(namer, le)).collect(Collectors.joining(", ")) + ")";
275             case Op.Result or when or.op() instanceof CoreOp.TupleLoadOp tlo -> OnnxLift.toJavaName(namer.apply(tlo.operands().getFirst())) + '.' + tupleAccessor(tlo.operands().getFirst(), tlo.index());
276             case Value v -> OnnxLift.toJavaName(namer.apply(v));
277             default -> throw new UnsupportedOperationException(value.toString());
278         };
279     }
280 
281     private static OnnxOp.OnnxSchema getSchema(OnnxOp oo) {
282         try {
283             return (OnnxOp.OnnxSchema) oo.getClass().getDeclaredField("SCHEMA").get(null);
284         } catch (ReflectiveOperationException ex) {
285             throw new RuntimeException(ex);
286         }
287     }
288 
289     private static String toJavaType(CodeType t) {
290         return switch (t) {
291             case OnnxType.TensorType tt ->
292                 "Tensor<" + switch (tt.eType()) {
293                     case OnnxType.Float32Type _ -> "Float";
294                     case OnnxType.Int64Type _ -> "Long";
295                     case OnnxType.Int32Type _ -> "Integer";
296                     case OnnxType.UInt8Type _ -> "Byte";
297                     case OnnxType.BoolType _ -> "Boolean";
298                     case OnnxType.StringType _ -> "String";
299                     default -> throw new UnsupportedOperationException(t.toString());
300                 } + ">";
301             default -> "var";
302         };
303     }
304 }