163 lines
5.6 KiB
Python
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
|