-
Notifications
You must be signed in to change notification settings - Fork 956
Expand file tree
/
Copy pathGpuBridge.java
More file actions
85 lines (71 loc) · 2.38 KB
/
Copy pathGpuBridge.java
File metadata and controls
85 lines (71 loc) · 2.38 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
import io.btrace.core.annotations.BTrace;
import io.btrace.core.annotations.Duration;
import io.btrace.core.annotations.Injected;
import io.btrace.core.annotations.Kind;
import io.btrace.core.annotations.Location;
import io.btrace.core.annotations.OnEvent;
import io.btrace.core.annotations.OnMethod;
import io.btrace.core.annotations.OnTimer;
import io.btrace.gpu.GpuBridgeService;
import static io.btrace.core.BTraceUtils.*;
/**
* Traces GPU model inference via ONNX Runtime and DJL (Deep Java Library).
* Tracks inference latency, batch sizes, and model load times.
*
* <p>Attach to a JVM running ONNX or DJL inference:
* <pre>
* btrace <pid> GpuBridge.java
* </pre>
*/
@BTrace
public class GpuBridge {
@Injected
private static GpuBridgeService gpu;
// ==================== ONNX Runtime ====================
@OnMethod(
clazz = "ai.onnxruntime.OrtSession",
method = "run",
location = @Location(Kind.RETURN))
public static void onOnnxInference(@Duration long dur) {
gpu.recordInference("onnx", "session", dur);
}
@OnMethod(
clazz = "ai.onnxruntime.OrtSession",
method = "<init>",
location = @Location(Kind.RETURN))
public static void onOnnxModelLoad(@Duration long dur) {
gpu.recordModelLoad("onnx", "session", dur);
}
// ==================== DJL (Deep Java Library) ====================
@OnMethod(
clazz = "/ai\\.djl\\.inference\\.Predictor/",
method = "predict",
location = @Location(Kind.RETURN))
public static void onDjlPredict(@Duration long dur) {
gpu.recordInference("djl", "predictor", dur);
}
@OnMethod(
clazz = "/ai\\.djl\\.repository\\.zoo\\.ModelZoo/",
method = "loadModel",
location = @Location(Kind.RETURN))
public static void onDjlModelLoad(@Duration long dur) {
gpu.recordModelLoad("djl", "model-zoo", dur);
}
// ==================== TensorFlow Java ====================
@OnMethod(
clazz = "/org\\.tensorflow\\.Session/",
method = "run",
location = @Location(Kind.RETURN))
public static void onTensorFlowRun(@Duration long dur) {
gpu.recordInference("tensorflow", "session", dur);
}
// ==================== Periodic summary ====================
@OnTimer(30000)
public static void periodicSummary() {
println(gpu.getSummary());
}
@OnEvent("summary")
public static void onDemandSummary() {
println(gpu.getSummary());
}
}