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 }