Serving GeoSAM
gbx.models.serving takes a loaded GeoSAM model from a notebook call to a queryable GPU
endpoint: wrap it as an MLflow pyfunc, register it to the Unity Catalog model registry
("Unity Gateway"), stand up a GPU Model Serving endpoint, and query it. Every step is a
thin wrapper over MLflow / Databricks Model Serving — mlflow itself is imported lazily
inside each function, so databricks.labs.gbx.models.serving imports cleanly even where
mlflow is not installed.
The served pyfunc
build_geosam_pyfunc(*, model_type="vit_h", weights=None) returns an MLflow
PythonModel whose predict method:
- base64-decodes a single
image_b64input string into raw image/GeoTIFF bytes, - calls
gbx.models.runner.segment_rasteron those bytes (via therunnermodule reference, so the served model always calls the currentsegment_raster— including under test monkeypatching), and - returns
{"geojson": <FeatureCollection string>}— a GeoJSONFeatureCollectionwith oneFeatureper segmented object (geometryfrom each row's WKBgeom,properties.labelandproperties.scorefrom the corresponding columns).
from databricks.labs.gbx.models import serving
pyfunc = serving.build_geosam_pyfunc(model_type="vit_h")
The input/output contract is plain data — one string column in
(serving.PREDICT_INPUT_COLUMNS, image_b64:string), one string column out
(serving.PREDICT_OUTPUT_COLUMNS, geojson:string) — and geosam_signature() converts
that contract into a real MLflow ModelSignature lazily, for use at registration time.
Registering to Unity Catalog — usage affiliation
model_version = serving.register_to_unity_gateway(
pyfunc,
"catalog.schema.geosam_model",
profile="oauth-fe",
)
register_to_unity_gateway(pyfunc, name, *, profile, signature=None, pip_reqs=None) logs
pyfunc to the Unity Catalog model registry under the three-level UC name name
(catalog.schema.model) and returns the resulting model version as a string.
profile selects the DATABRICKS_CONFIG_PROFILE used for the registry call — this
function never auto-selects a profile; the caller chooses it explicitly.
pip_reqs defaults to serving.DEFAULT_PIP_REQS:
DEFAULT_PIP_REQS = (
"geobrix[models_gpu_env5] @ file:///Volumes/.../geobrix-0.5.2-py3-none-any.whl",
"segment-geospatial",
)
This is a deliberate usage-affiliation choice: the served model's logged environment
carries the GeoBrix light wheel alongside segment-geospatial, not just the raw
segmentation backend. Override pip_reqs= with the actual staged wheel path/version for
your workspace.
Creating a GPU endpoint
endpoint = serving.create_endpoint(
"catalog.schema.geosam_model",
model_version,
profile="oauth-fe",
endpoint_name="geosam-endpoint",
workload_type="GPU_MEDIUM",
scale_to_zero=True,
)
create_endpoint(name, model_version, *, profile, endpoint_name, workload_type="GPU_MEDIUM", scale_to_zero=True)
assembles a served-entity/traffic config for UC model name@model_version and creates a
Model Serving endpoint named endpoint_name. It polls until the endpoint is genuinely
ready — both state.ready == "READY" and state.config_update == "NOT_UPDATING" —
since state.ready alone can read READY mid version-swap while the previous version is
still the one actually serving traffic.
create_endpoint is create-only — it calls the deployment client's create API, not
an update. Rolling a new model_version onto an endpoint that already exists needs a
separate update call (not implemented in this module); calling create_endpoint again
against an existing endpoint_name is not the way to deploy a new version.
Scale-to-zero vs. provisioned
scale_to_zero=True(default) — the endpoint's GPU capacity scales down to zero when idle and cold-starts on the next request. Lowest cost for intermittent or demo/dev usage; the first request after an idle period pays a cold-start latency penalty (loading the model onto a GPU).scale_to_zero=False(provisioned) — GPU capacity stays warm continuously. Predictable low latency, at continuous GPU cost for the workload type — appropriate for latency-sensitive or steady-traffic serving.
workload_type follows the standard Databricks GPU workload-type names (e.g.
"GPU_MEDIUM"); pick the size that matches your model and expected traffic.
Querying the endpoint
result = serving.query("geosam-endpoint", raw_image_bytes, profile="oauth-fe")
query(endpoint_name, image, *, profile) base64-encodes image (raw image/GeoTIFF
bytes) and calls the running endpoint via the classical-ML dataframe_records request
shape the single-string-column signature expects — the same shape
build_geosam_pyfunc's predict handles server-side. It returns the endpoint's raw
prediction response (an MLflow deployments-client dict), with the geojson string inside
predictions.
Cost and auth notes
- Auth is always explicit. Every function in this module takes a
profileargument and setsDATABRICKS_CONFIG_PROFILEfrom it for that call — nothing here auto-selects a profile. - GPU serving is not free while warm. A provisioned (
scale_to_zero=False) GPU endpoint accrues cost continuously, independent of query volume. Preferscale_to_zero=Trueunless you have steady traffic that justifies keeping the model warm. - Registration ships GeoBrix into the endpoint's environment. Because
DEFAULT_PIP_REQSincludes the GeoBrix wheel, the endpoint's serving container installs GeoBrix at startup — factor that into cold-start time for a scale-to-zero endpoint. mlflowis not part of the light CI environment. All ofmlflow's imports in this module are lazy and scoped to the function that needs it; only code paths that callregister_to_unity_gateway,create_endpoint, orquery— orbuild_geosam_pyfunc/geosam_signature— requiremlflowto be installed.
See also
- GeoSAM —
load_geosam+segment_raster, the model this page serves. - Geospatial Models overview — the load → run → serve lifecycle.
- Orthomosaic example —
03_segment's upsize step registers and serves GeoSAM this same way, then queries the endpoint with a cropped window of the published orthomosaic COG.