Skip to content

Commit 6f82fbf

Browse files
mstrathmanclaude
andcommitted
feat(onnx): GPU execution providers (CUDA/TensorRT, fp16/int8), gated
Wires the GPU serving path behind -DSQLITE_PREDICT_ONNX_GPU (make loadable-onnx-gpu): - CUDA and TensorRT execution providers via proper provider-options (Create/Update/Release), not the previous NULL stub. TensorRT honors trt_fp16_enable when precision=fp16. Appending fails loud if the EP is not in the onnxruntime build — never a silent drop to CPU. - fp16/int8 precision allowed only in the GPU build, and only with a cuda/tensorrt device (pairing fp16 with cpu/coreml is rejected rather than silently computing fp32). - The provider-options symbols are in every onnxruntime C API, so the GPU build compiles and links against the CPU onnxruntime. CI now compile- checks it (make loadable-onnx-gpu); real GPU execution is validated on a dedicated GPU job (needs onnxruntime-gpu + a GPU runner). Verified locally against the CPU onnxruntime: the GPU build compiles and links clean, device=cuda/tensorrt fail loud with RUNTIME_UNAVAILABLE (no crash), fp16+cpu is rejected, and the regular onnx suite stays 124 green. The receipt already records device+precision, so GPU results are honestly distinguishable from the deterministic CPU path. README/ARCHITECTURE/ CHANGELOG updated. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Signed-off-by: mstrathman <matthew.strathman@gmail.com>
1 parent 870449e commit 6f82fbf

6 files changed

Lines changed: 96 additions & 20 deletions

File tree

‎.github/workflows/ci.yml‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -127,3 +127,5 @@ jobs:
127127
run: make test-onnx
128128
- name: ASan + LSan soak of the onnx backend
129129
run: make test-asan-onnx
130+
- name: Compile-check the GPU build (CUDA/TensorRT wiring)
131+
run: make loadable-onnx-gpu

‎ARCHITECTURE.md‎

Lines changed: 10 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -53,10 +53,16 @@ precision), run query rows in batches, and select the execution provider
5353
explicitly, erasing no failure into a silent CPU fallback. Weights are pinned
5454
by content hash, so a receipt records exactly which bytes ran, and the
5555
in-context receipt anchors the training rows too, so mutating the context
56-
breaks replay. Both run on CPU today; the GPU execution providers are the
57-
next layer, validated on the gated GPU CI job. Even so, the `benchmarks/`
58-
numbers are why the default answer for a teacher's accuracy is usually
59-
distillation to a small vector student.
56+
breaks replay. Both run on CPU. A GPU build (`make loadable-onnx-gpu`,
57+
`-DSQLITE_PREDICT_ONNX_GPU`) wires the CUDA and TensorRT providers and
58+
fp16/int8 precision; the provider-options symbols are in every onnxruntime
59+
C API, so it compiles and links against the CPU onnxruntime for a CI
60+
compile-check, while real GPU execution is validated on a dedicated GPU job.
61+
The receipt records the execution provider and precision, which is what
62+
makes GPU results honestly distinguishable from the deterministic CPU path.
63+
Even so, the `benchmarks/` numbers (and the TabFM→ONNX eval in
64+
`benchmarks/results/tabfm-onnx.md`) are why the default answer for a
65+
teacher's accuracy is usually distillation to a small vector student.
6066

6167
## Receipts, anchoring, and replay
6268

‎CHANGELOG.md‎

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -21,8 +21,10 @@ project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
2121
execution-provider selection. The `vector` layout serves a self-contained
2222
model (a distilled student or classifier/regressor); the `in_context`
2323
layout serves a teacher that ingests the `train_query` rows as context
24-
each call (TabFM-shaped), anchoring those rows in the receipt. The default
25-
build stays zero-dependency.
24+
each call (TabFM-shaped), anchoring those rows in the receipt. A GPU build
25+
(`make loadable-onnx-gpu`) adds the CUDA and TensorRT execution providers
26+
and fp16/int8 precision, compile-checked in CI and validated on a
27+
dedicated GPU job. The default build stays zero-dependency.
2628
- `predict_register(model_id, config_json)` to register an external model,
2729
pinning its weights by content hash. `_predict_models` gains `weights_uri`
2830
and `io_spec`; receipts record the execution provider and precision.

‎Makefile‎

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -78,6 +78,17 @@ loadable-onnx: $(prefix) vendor/sqlite3ext.h sqlite-predict.h $(ONNX_OBJS)
7878
$(ONNX_CFLAGS) $(CFLAGS) $(ONNX_OBJS) -o $(TARGET_LOADABLE) \
7979
$(LDFLAGS) $(ONNX_LDFLAGS)
8080

81+
# GPU-enabled onnx build: adds the CUDA and TensorRT execution providers
82+
# (fp16/int8 precision) behind -DSQLITE_PREDICT_ONNX_GPU. Running it needs
83+
# an onnxruntime-gpu install; the provider-options symbols are in the C API
84+
# of every onnxruntime build, so this compiles and links against the CPU
85+
# onnxruntime for a CI compile-check. Real GPU execution is validated on the
86+
# gated GPU job.
87+
loadable-onnx-gpu: $(prefix) vendor/sqlite3ext.h sqlite-predict.h $(ONNX_OBJS)
88+
$(CC) -fPIC -shared -std=c99 -Wall -Wextra -Ivendor/ -I./ -O3 \
89+
$(ONNX_CFLAGS) -DSQLITE_PREDICT_ONNX_GPU $(CFLAGS) $(ONNX_OBJS) \
90+
-o $(TARGET_LOADABLE) $(LDFLAGS) $(ONNX_LDFLAGS)
91+
8192
debug: $(prefix) vendor/sqlite3ext.h sqlite-predict.h $(OBJS)
8293
$(CC) -fPIC -shared -std=c99 -Wall -Wextra -Ivendor/ -I./ -g -O0 -DSQLITE_PREDICT_DEBUG $(CFLAGS) $(OBJS) -o $(TARGET_LOADABLE) $(LDFLAGS)
8394

@@ -176,5 +187,5 @@ clean:
176187
format:
177188
clang-format -i sqlite-predict.c predict-*.c
178189

179-
.PHONY: loadable loadable-onnx debug test test-loadable test-onnx \
180-
test-asan-onnx clean format
190+
.PHONY: loadable loadable-onnx loadable-onnx-gpu debug test test-loadable \
191+
test-onnx test-asan-onnx clean format

‎README.md‎

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -90,8 +90,12 @@ serves two shapes, with a cached session and batched inference:
9090
on each call and labels the `apply_query` rows against them, the way
9191
TabFM works.
9292

93-
Both run on CPU today. Running a teacher on a GPU (CUDA/TensorRT, fp16) is
94-
the remaining step. The default build never links onnxruntime.
93+
Both run on CPU today. A GPU build (`make loadable-onnx-gpu`) adds the CUDA
94+
and TensorRT execution providers and fp16/int8 precision; it needs an
95+
onnxruntime-gpu install, is compile-checked in CI, and its GPU execution is
96+
validated on a dedicated GPU job. Provider selection is explicit and fails
97+
loud: asking for `cuda` on a build without it errors, never a silent drop to
98+
CPU. The default build never links onnxruntime at all.
9599

96100
## Installing
97101

@@ -110,7 +114,9 @@ build. Requires a C99 compiler. Then, from any SQLite client:
110114

111115
For the ONNX serving path, install onnxruntime (macOS: `brew install
112116
onnxruntime`; Linux: extract an [onnxruntime release][ort] and point
113-
`ONNXRUNTIME_PREFIX` at it) and build `make loadable-onnx`.
117+
`ONNXRUNTIME_PREFIX` at it) and build `make loadable-onnx`. For GPU
118+
execution, point `ONNXRUNTIME_PREFIX` at an onnxruntime-gpu install and
119+
build `make loadable-onnx-gpu`.
114120

115121
[ort]: https://github.com/microsoft/onnxruntime/releases
116122

‎predict-onnx.c‎

Lines changed: 58 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -266,12 +266,43 @@ static int onnx_build_session(const predict0_model_row *model,
266266
#endif
267267
} else if (strcmp(device, "cuda") == 0 || strcmp(device, "tensorrt") == 0) {
268268
#ifdef SQLITE_PREDICT_ONNX_GPU
269-
OrtStatus *ep =
270-
strcmp(device, "cuda") == 0
271-
? g_ort->SessionOptionsAppendExecutionProvider_CUDA_V2(so, NULL)
272-
: g_ort->SessionOptionsAppendExecutionProvider_TensorRT_V2(so,
273-
NULL);
274-
BUILD_CHECK(ep, PREDICT_ERR_RUNTIME_UNAVAILABLE, "append GPU EP");
269+
/* Real provider-options wiring. Appending fails loud if the EP is not in
270+
* this onnxruntime build (i.e. onnxruntime-gpu is required). Compiled
271+
* against the CPU headers for the CI compile-check; exercised for real on
272+
* the gated GPU job. */
273+
int fp16 = opts->precision && strcmp(opts->precision, "fp16") == 0;
274+
if (strcmp(device, "cuda") == 0) {
275+
OrtCUDAProviderOptionsV2 *cu = NULL;
276+
BUILD_CHECK(g_ort->CreateCUDAProviderOptions(&cu),
277+
PREDICT_ERR_RUNTIME_UNAVAILABLE, "CreateCUDAProviderOptions");
278+
/* CUDA EP runs the model's own dtype; fp16 comes from an fp16 model,
279+
* not an EP flag. precision is recorded in the receipt regardless. */
280+
OrtStatus *ap =
281+
g_ort->SessionOptionsAppendExecutionProvider_CUDA_V2(so, cu);
282+
g_ort->ReleaseCUDAProviderOptions(cu);
283+
BUILD_CHECK(ap, PREDICT_ERR_RUNTIME_UNAVAILABLE, "append CUDA EP");
284+
} else {
285+
OrtTensorRTProviderOptionsV2 *trt = NULL;
286+
BUILD_CHECK(g_ort->CreateTensorRTProviderOptions(&trt),
287+
PREDICT_ERR_RUNTIME_UNAVAILABLE,
288+
"CreateTensorRTProviderOptions");
289+
if (fp16) {
290+
const char *keys[] = {"trt_fp16_enable"};
291+
const char *vals[] = {"1"};
292+
OrtStatus *up =
293+
g_ort->UpdateTensorRTProviderOptions(trt, keys, vals, 1);
294+
if (up) {
295+
g_ort->ReleaseTensorRTProviderOptions(trt);
296+
rc = onnx_fail(up, PREDICT_ERR_RUNTIME_UNAVAILABLE,
297+
"trt_fp16_enable", errmsg);
298+
goto fail;
299+
}
300+
}
301+
OrtStatus *ap =
302+
g_ort->SessionOptionsAppendExecutionProvider_TensorRT_V2(so, trt);
303+
g_ort->ReleaseTensorRTProviderOptions(trt);
304+
BUILD_CHECK(ap, PREDICT_ERR_RUNTIME_UNAVAILABLE, "append TensorRT EP");
305+
}
275306
#else
276307
/* Honest state: the GPU execution providers are validated on the gated
277308
* GPU CI job and compiled only into the GPU build. This CPU build does
@@ -1153,13 +1184,31 @@ int predict0_onnx_predict(sqlite3 *db, const char *model_id,
11531184
return rc;
11541185

11551186
const char *precision = opts->precision ? opts->precision : "fp32";
1156-
if (strcmp(precision, "fp32") != 0) {
1187+
#ifdef SQLITE_PREDICT_ONNX_GPU
1188+
int prec_ok = strcmp(precision, "fp32") == 0 ||
1189+
strcmp(precision, "fp16") == 0 ||
1190+
strcmp(precision, "int8") == 0;
1191+
#else
1192+
int prec_ok = strcmp(precision, "fp32") == 0;
1193+
#endif
1194+
if (!prec_ok) {
11571195
*errmsg = sqlite3_mprintf(
1158-
"%s: precision '%s' is not in this build (fp32 only; fp16/int8 land"
1159-
" with the GPU path)",
1196+
"%s: precision '%s' is not available in this build (fp32 only;"
1197+
" fp16/int8 need the GPU build, loadable-onnx-gpu)",
11601198
PREDICT_ERR_RUNTIME_UNAVAILABLE, precision);
11611199
return SQLITE_ERROR;
11621200
}
1201+
/* fp16/int8 only take effect on a GPU device; pairing them with cpu/coreml
1202+
* would silently compute fp32, so reject it rather than no-op quietly. */
1203+
if (strcmp(precision, "fp32") != 0) {
1204+
const char *dev = opts->device ? opts->device : "cpu";
1205+
if (strcmp(dev, "cuda") != 0 && strcmp(dev, "tensorrt") != 0) {
1206+
*errmsg = sqlite3_mprintf(
1207+
"%s: precision '%s' requires a cuda or tensorrt device", PREDICT_ERR_OPTIONS,
1208+
precision);
1209+
return SQLITE_ERROR;
1210+
}
1211+
}
11631212
if (!license_ok(model->license, opts->accept_license)) {
11641213
*errmsg = sqlite3_mprintf(
11651214
"%s: model license '%s' requires accept_license to match",

0 commit comments

Comments
 (0)