Files
AareLC_train/src/core/inference_client.py
T
2026-04-14 16:07:31 +02:00

163 lines
5.6 KiB
Python

"""Client for ML inference server communication."""
import requests
import cv2
import json
import numpy as np
class InferenceClient:
"""Handles communication with ML inference server."""
def __init__(self, server_url):
"""
Initialize inference client.
Args:
server_url: Base URL of inference server (e.g., "http://mx-ml.psi.ch:8002")
"""
self.server_url = server_url
def predict(self, image, timeout=10):
"""
Send image to server for prediction.
Args:
image: OpenCV image (BGR format)
timeout: Request timeout in seconds
Returns:
List of detection results, or None on error
Raises:
requests.RequestException: On connection errors
ValueError: On invalid response
"""
_, buf = cv2.imencode('.jpg', image)
url = f"{self.server_url}/predict/"
files = {'file': ('image.jpg', buf.tobytes(), 'image/jpeg')}
response = requests.post(url, files=files, timeout=timeout)
if response.status_code == 200:
data = response.json()
return data.get('results', [])
else:
raise ValueError(f"Server returned status {response.status_code}")
def segment_anything(self, image, x, y, timeout=15):
"""
Request a high-precision mask from the SAM backend.
"""
_, buf = cv2.imencode('.jpg', image)
# Use trailing slash to match the server's /sam/ route exactly
url = f"{self.server_url.rstrip('/')}/sam/"
params = {"point_x": int(x), "point_y": int(y)}
files = {'file': ('image.jpg', buf.tobytes(), 'image/jpeg')}
response = requests.post(url, files=files, params=params, timeout=timeout)
if response.status_code == 200:
return response.json()
else:
raise ValueError(f"SAM error: {response.status_code} - {response.text}")
def segment_anything_box(self, image, x1, y1, x2, y2, timeout=15):
"""
Request a high-precision mask from the SAM backend using a bounding box prompt.
"""
_, buf = cv2.imencode('.jpg', image)
url = f"{self.server_url.rstrip('/')}/sam_box/"
params = {
"x1": int(x1),
"y1": int(y1),
"x2": int(x2),
"y2": int(y2)
}
files = {'file': ('image.jpg', buf.tobytes(), 'image/jpeg')}
response = requests.post(url, files=files, params=params, timeout=timeout)
if response.status_code == 200:
return response.json()
else:
raise ValueError(f"SAM Box error: {response.status_code} - {response.text}")
def segment_multi_box(self, image, boxes, timeout=30):
"""
Request masks for multiple boxes in one call.
boxes: List of [x1, y1, x2, y2]
"""
_, buf = cv2.imencode('.jpg', image)
url = f"{self.server_url.rstrip('/')}/sam_multi_box/"
params = {"boxes_json": json.dumps(boxes)}
files = {'file': ('image.jpg', buf.tobytes(), 'image/jpeg')}
response = requests.post(url, files=files, params=params, timeout=timeout)
if response.status_code == 200:
return response.json().get('results', [])
else:
raise ValueError(f"SAM Multi Box error: {response.status_code}")
def predict_depth(self, image, timeout=20, colorize=True, depth_model=None, depth_max_side=0):
"""
Request depth map for an image.
Returns:
OpenCV image array decoded from PNG response.
"""
_, buf = cv2.imencode('.jpg', image)
url = f"{self.server_url.rstrip('/')}/depth/"
params = {"colorize": "true" if colorize else "false"}
if depth_model:
params["depth_model"] = str(depth_model)
if isinstance(depth_max_side, int) and depth_max_side > 0:
params["depth_max_side"] = str(depth_max_side)
files = {'file': ('image.jpg', buf.tobytes(), 'image/jpeg')}
response = requests.post(url, files=files, params=params, timeout=timeout)
if response.status_code != 200:
raise ValueError(f"Depth error: {response.status_code} - {response.text}")
if "image/png" not in response.headers.get("content-type", ""):
try:
err = response.json().get("error", response.text)
except Exception:
err = response.text
raise ValueError(f"Depth error: {err}")
depth_png = cv2.imdecode(np.frombuffer(response.content, np.uint8), cv2.IMREAD_UNCHANGED)
if depth_png is None:
raise ValueError("Depth response decode failed.")
return depth_png
def send_training_data(self, image, metadata, timeout=10):
"""
Send annotated image to training pipeline.
Args:
image: OpenCV image (BGR format)
metadata: Dict containing params and detections
timeout: Request timeout in seconds
Returns:
Response object
Raises:
requests.RequestException: On connection errors
"""
_, buf = cv2.imencode('.jpg', image)
url = f"{self.server_url}/train/"
files = {'file': ('img.jpg', buf.tobytes(), 'image/jpeg')}
data = {"metadata": json.dumps(metadata)}
response = requests.post(url, files=files, data=data, timeout=timeout)
response.raise_for_status()
return response
def set_server_url(self, url):
"""Update server URL."""
self.server_url = url
def get_server_url(self):
"""Get current server URL."""
return self.server_url