se3labs 0.0.0a0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- se3labs-0.0.0a0/.gitignore +92 -0
- se3labs-0.0.0a0/PKG-INFO +74 -0
- se3labs-0.0.0a0/README.md +52 -0
- se3labs-0.0.0a0/pyproject.toml +45 -0
- se3labs-0.0.0a0/src/se3labs/__init__.py +10 -0
- se3labs-0.0.0a0/src/se3labs/__main__.py +7 -0
- se3labs-0.0.0a0/src/se3labs/cli.py +156 -0
- se3labs-0.0.0a0/src/se3labs/eval/__init__.py +1 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/__init__.py +9 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/__main__.py +7 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/client.py +170 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/proto/__init__.py +3 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/proto/job_status.proto +31 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/proto/job_status_pb2.py +44 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/proto/job_status_pb2.pyi +38 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/proto/job_status_pb2_grpc.py +98 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/proto/policy_session.proto +319 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/proto/policy_session_pb2.py +114 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/proto/policy_session_pb2.pyi +441 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/proto/policy_session_pb2_grpc.py +110 -0
- se3labs-0.0.0a0/src/se3labs/eval/policy/reference.py +74 -0
|
@@ -0,0 +1,92 @@
|
|
|
1
|
+
# Data
|
|
2
|
+
*.rrd
|
|
3
|
+
|
|
4
|
+
# Prerequisites
|
|
5
|
+
*.d
|
|
6
|
+
|
|
7
|
+
# Compiled Object files
|
|
8
|
+
*.slo
|
|
9
|
+
*.lo
|
|
10
|
+
*.o
|
|
11
|
+
*.obj
|
|
12
|
+
|
|
13
|
+
# Precompiled Headers
|
|
14
|
+
*.gch
|
|
15
|
+
*.pch
|
|
16
|
+
|
|
17
|
+
# Linker files
|
|
18
|
+
*.ilk
|
|
19
|
+
|
|
20
|
+
# Debugger Files
|
|
21
|
+
*.pdb
|
|
22
|
+
|
|
23
|
+
# Compiled Dynamic libraries
|
|
24
|
+
*.so
|
|
25
|
+
*.dylib
|
|
26
|
+
*.dll
|
|
27
|
+
*.so.*
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
# Fortran module files
|
|
31
|
+
*.mod
|
|
32
|
+
*.smod
|
|
33
|
+
|
|
34
|
+
# Compiled Static libraries
|
|
35
|
+
*.lai
|
|
36
|
+
*.la
|
|
37
|
+
*.a
|
|
38
|
+
*.lib
|
|
39
|
+
|
|
40
|
+
# Executables
|
|
41
|
+
*.exe
|
|
42
|
+
*.out
|
|
43
|
+
*.app
|
|
44
|
+
|
|
45
|
+
# Build directories
|
|
46
|
+
build/
|
|
47
|
+
Build/
|
|
48
|
+
build-*/
|
|
49
|
+
|
|
50
|
+
# CMake generated files
|
|
51
|
+
CMakeFiles/
|
|
52
|
+
CMakeCache.txt
|
|
53
|
+
cmake_install.cmake
|
|
54
|
+
Makefile
|
|
55
|
+
install_manifest.txt
|
|
56
|
+
compile_commands.json
|
|
57
|
+
|
|
58
|
+
# Temporary files
|
|
59
|
+
*.tmp
|
|
60
|
+
*.log
|
|
61
|
+
*.bak
|
|
62
|
+
*.swp
|
|
63
|
+
|
|
64
|
+
# Python
|
|
65
|
+
__pycache__/
|
|
66
|
+
*.py[cod]
|
|
67
|
+
*.egg-info/
|
|
68
|
+
.venv/
|
|
69
|
+
venv/
|
|
70
|
+
dist/
|
|
71
|
+
.pytest_cache/
|
|
72
|
+
.ruff_cache/
|
|
73
|
+
|
|
74
|
+
# vcpkg
|
|
75
|
+
vcpkg_installed/
|
|
76
|
+
|
|
77
|
+
# debug information files
|
|
78
|
+
*.dwo
|
|
79
|
+
|
|
80
|
+
# test output & cache
|
|
81
|
+
Testing/
|
|
82
|
+
.cache/
|
|
83
|
+
|
|
84
|
+
# Documentation site (doc/public, Docusaurus)
|
|
85
|
+
doc/public/node_modules/
|
|
86
|
+
doc/public/build/
|
|
87
|
+
doc/public/.docusaurus/
|
|
88
|
+
doc/public/.cache-loader/
|
|
89
|
+
|
|
90
|
+
# Customer package build output (sdk/)
|
|
91
|
+
sdk/dist/
|
|
92
|
+
sdk/build/
|
se3labs-0.0.0a0/PKG-INFO
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
1
|
+
Metadata-Version: 2.5
|
|
2
|
+
Name: se3labs
|
|
3
|
+
Version: 0.0.0a0
|
|
4
|
+
Summary: SE3 Labs client library: evaluate your robot policy on an SE3 Labs station without shipping model code.
|
|
5
|
+
Project-URL: Documentation, https://se3labs.github.io/YAM-Station/
|
|
6
|
+
Project-URL: Source, https://github.com/se3labs/YAM-Station
|
|
7
|
+
Author: SE3 Labs
|
|
8
|
+
License-Expression: Apache-2.0
|
|
9
|
+
Keywords: grpc,policy-evaluation,robotics
|
|
10
|
+
Classifier: Development Status :: 3 - Alpha
|
|
11
|
+
Classifier: Intended Audience :: Science/Research
|
|
12
|
+
Classifier: Programming Language :: Python :: 3
|
|
13
|
+
Classifier: Programming Language :: Python :: 3.10
|
|
14
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
15
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
16
|
+
Classifier: Programming Language :: Python :: 3.13
|
|
17
|
+
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
|
18
|
+
Requires-Python: >=3.10
|
|
19
|
+
Requires-Dist: grpcio>=1.75
|
|
20
|
+
Requires-Dist: protobuf>=6
|
|
21
|
+
Description-Content-Type: text/markdown
|
|
22
|
+
|
|
23
|
+
# se3labs
|
|
24
|
+
|
|
25
|
+
Client library for evaluating a robot policy on an SE3 Labs station. Your
|
|
26
|
+
model runs on your machine; the station streams observations to it over one
|
|
27
|
+
gRPC stream, executes the actions it returns under the station's safety
|
|
28
|
+
layer, and records everything. Nothing here touches hardware: the only
|
|
29
|
+
dependencies are `grpcio` and `protobuf`.
|
|
30
|
+
|
|
31
|
+
```bash
|
|
32
|
+
pip install se3labs
|
|
33
|
+
```
|
|
34
|
+
|
|
35
|
+
Run a reference policy against the job id your SE3 Labs contact gives you,
|
|
36
|
+
and poll the job's progress:
|
|
37
|
+
|
|
38
|
+
```bash
|
|
39
|
+
se3labs eval run --job job_… --address relay.example.com:7443 --policy hold
|
|
40
|
+
se3labs eval status --job job_… --address relay.example.com:7443
|
|
41
|
+
```
|
|
42
|
+
|
|
43
|
+
Wrap your own model:
|
|
44
|
+
|
|
45
|
+
```python
|
|
46
|
+
from se3labs.eval.policy import client
|
|
47
|
+
from se3labs.eval.policy.proto import policy_session_pb2 as pb
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class MyPolicy:
|
|
51
|
+
def manifest(self) -> pb.ClientManifest:
|
|
52
|
+
return client.default_manifest(model="my-model", version="1", max_horizon=8)
|
|
53
|
+
|
|
54
|
+
def reset(self, reset: pb.Reset) -> None:
|
|
55
|
+
pass
|
|
56
|
+
|
|
57
|
+
def act(self, observe: pb.Observe) -> list[pb.Action]:
|
|
58
|
+
targets = {a.name: (list(a.joint_pos_rad), None) for a in observe.observation.arms}
|
|
59
|
+
return [client.joint_action(targets) for _ in range(8)]
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
client.serve_policy("job_…", "relay.example.com:7443", MyPolicy())
|
|
63
|
+
```
|
|
64
|
+
|
|
65
|
+
Or, with that class saved in `my_policy.py`, without writing the last line:
|
|
66
|
+
|
|
67
|
+
```bash
|
|
68
|
+
se3labs eval run --job job_… --address relay.example.com:7443 --policy my_policy:MyPolicy
|
|
69
|
+
```
|
|
70
|
+
|
|
71
|
+
Full documentation: https://se3labs.github.io/YAM-Station/
|
|
72
|
+
|
|
73
|
+
Protocol version: `se3labs.eval.policy.PROTOCOL_VERSION`. A release that
|
|
74
|
+
changes it cannot talk to a station expecting the old one.
|
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
# se3labs
|
|
2
|
+
|
|
3
|
+
Client library for evaluating a robot policy on an SE3 Labs station. Your
|
|
4
|
+
model runs on your machine; the station streams observations to it over one
|
|
5
|
+
gRPC stream, executes the actions it returns under the station's safety
|
|
6
|
+
layer, and records everything. Nothing here touches hardware: the only
|
|
7
|
+
dependencies are `grpcio` and `protobuf`.
|
|
8
|
+
|
|
9
|
+
```bash
|
|
10
|
+
pip install se3labs
|
|
11
|
+
```
|
|
12
|
+
|
|
13
|
+
Run a reference policy against the job id your SE3 Labs contact gives you,
|
|
14
|
+
and poll the job's progress:
|
|
15
|
+
|
|
16
|
+
```bash
|
|
17
|
+
se3labs eval run --job job_… --address relay.example.com:7443 --policy hold
|
|
18
|
+
se3labs eval status --job job_… --address relay.example.com:7443
|
|
19
|
+
```
|
|
20
|
+
|
|
21
|
+
Wrap your own model:
|
|
22
|
+
|
|
23
|
+
```python
|
|
24
|
+
from se3labs.eval.policy import client
|
|
25
|
+
from se3labs.eval.policy.proto import policy_session_pb2 as pb
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class MyPolicy:
|
|
29
|
+
def manifest(self) -> pb.ClientManifest:
|
|
30
|
+
return client.default_manifest(model="my-model", version="1", max_horizon=8)
|
|
31
|
+
|
|
32
|
+
def reset(self, reset: pb.Reset) -> None:
|
|
33
|
+
pass
|
|
34
|
+
|
|
35
|
+
def act(self, observe: pb.Observe) -> list[pb.Action]:
|
|
36
|
+
targets = {a.name: (list(a.joint_pos_rad), None) for a in observe.observation.arms}
|
|
37
|
+
return [client.joint_action(targets) for _ in range(8)]
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
client.serve_policy("job_…", "relay.example.com:7443", MyPolicy())
|
|
41
|
+
```
|
|
42
|
+
|
|
43
|
+
Or, with that class saved in `my_policy.py`, without writing the last line:
|
|
44
|
+
|
|
45
|
+
```bash
|
|
46
|
+
se3labs eval run --job job_… --address relay.example.com:7443 --policy my_policy:MyPolicy
|
|
47
|
+
```
|
|
48
|
+
|
|
49
|
+
Full documentation: https://se3labs.github.io/YAM-Station/
|
|
50
|
+
|
|
51
|
+
Protocol version: `se3labs.eval.policy.PROTOCOL_VERSION`. A release that
|
|
52
|
+
changes it cannot talk to a station expecting the old one.
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
[project]
|
|
2
|
+
name = "se3labs"
|
|
3
|
+
# PEP 440 spelling of 0.0.0-alpha. Bump on every release; PyPI never
|
|
4
|
+
# accepts the same version twice.
|
|
5
|
+
version = "0.0.0a0"
|
|
6
|
+
description = "SE3 Labs client library: evaluate your robot policy on an SE3 Labs station without shipping model code."
|
|
7
|
+
readme = "README.md"
|
|
8
|
+
requires-python = ">=3.10"
|
|
9
|
+
license = "Apache-2.0"
|
|
10
|
+
authors = [{ name = "SE3 Labs" }]
|
|
11
|
+
keywords = ["robotics", "policy-evaluation", "grpc"]
|
|
12
|
+
classifiers = [
|
|
13
|
+
"Development Status :: 3 - Alpha",
|
|
14
|
+
"Intended Audience :: Science/Research",
|
|
15
|
+
"Programming Language :: Python :: 3",
|
|
16
|
+
"Programming Language :: Python :: 3.10",
|
|
17
|
+
"Programming Language :: Python :: 3.11",
|
|
18
|
+
"Programming Language :: Python :: 3.12",
|
|
19
|
+
"Programming Language :: Python :: 3.13",
|
|
20
|
+
"Topic :: Scientific/Engineering :: Artificial Intelligence",
|
|
21
|
+
]
|
|
22
|
+
# The generated *_pb2 modules check the protobuf runtime at import and
|
|
23
|
+
# refuse one older than the protoc they came from; the floor tracks what
|
|
24
|
+
# tool/gen_proto.sh was last run with.
|
|
25
|
+
dependencies = [
|
|
26
|
+
"grpcio>=1.75",
|
|
27
|
+
"protobuf>=6",
|
|
28
|
+
]
|
|
29
|
+
|
|
30
|
+
[project.scripts]
|
|
31
|
+
se3labs = "se3labs.cli:main"
|
|
32
|
+
|
|
33
|
+
[project.urls]
|
|
34
|
+
Documentation = "https://se3labs.github.io/YAM-Station/"
|
|
35
|
+
Source = "https://github.com/se3labs/YAM-Station"
|
|
36
|
+
|
|
37
|
+
[build-system]
|
|
38
|
+
requires = ["hatchling"]
|
|
39
|
+
build-backend = "hatchling.build"
|
|
40
|
+
|
|
41
|
+
[tool.hatch.build.targets.wheel]
|
|
42
|
+
packages = ["src/se3labs"]
|
|
43
|
+
|
|
44
|
+
[tool.hatch.build.targets.sdist]
|
|
45
|
+
include = ["src/se3labs", "README.md", "pyproject.toml"]
|
|
@@ -0,0 +1,10 @@
|
|
|
1
|
+
"""SE3 Labs customer-facing packages. Hardware-free: nothing here imports a
|
|
2
|
+
robot driver. ``se3labs.eval.policy`` is the client for evaluating a policy
|
|
3
|
+
on an SE3 Labs station."""
|
|
4
|
+
|
|
5
|
+
from importlib.metadata import PackageNotFoundError, version as _version
|
|
6
|
+
|
|
7
|
+
try:
|
|
8
|
+
__version__ = _version("se3labs")
|
|
9
|
+
except PackageNotFoundError: # source tree without an install
|
|
10
|
+
__version__ = "0+unknown"
|
|
@@ -0,0 +1,156 @@
|
|
|
1
|
+
"""``se3labs``: the command line of the customer package.
|
|
2
|
+
|
|
3
|
+
se3labs eval run --job JOB --address HOST:PORT [--policy hold|random|module:Attr]
|
|
4
|
+
se3labs eval status --job JOB --address HOST:PORT
|
|
5
|
+
se3labs --version
|
|
6
|
+
|
|
7
|
+
Plain argparse so the package keeps its two runtime dependencies.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
import argparse
|
|
13
|
+
import importlib
|
|
14
|
+
import json
|
|
15
|
+
import sys
|
|
16
|
+
from collections.abc import Sequence
|
|
17
|
+
from importlib.metadata import PackageNotFoundError, version
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def _version() -> str:
|
|
21
|
+
try:
|
|
22
|
+
return version("se3labs")
|
|
23
|
+
except PackageNotFoundError: # running from a source tree without an install
|
|
24
|
+
return "0+unknown"
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def _parse_image(spec: str) -> tuple[str, int, int]:
|
|
28
|
+
"""``cam_chest:640x480`` → ("cam_chest", 640, 480); size defaults to 640x480."""
|
|
29
|
+
camera, _, size = spec.partition(":")
|
|
30
|
+
if not camera:
|
|
31
|
+
raise argparse.ArgumentTypeError(f"image spec needs a camera name: {spec!r}")
|
|
32
|
+
w, _, h = size.partition("x")
|
|
33
|
+
try:
|
|
34
|
+
return camera, int(w or 640), int(h or 480)
|
|
35
|
+
except ValueError:
|
|
36
|
+
raise argparse.ArgumentTypeError(f"image size must be WxH, got {size!r}") from None
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def load_policy(spec: str, *, horizon: int, images: Sequence[tuple[str, int, int]]):
|
|
40
|
+
"""A policy object from its spec: a reference name, or ``module:Attr``
|
|
41
|
+
naming a class (instantiated with no arguments) or an instance."""
|
|
42
|
+
from se3labs.eval.policy import reference
|
|
43
|
+
|
|
44
|
+
if spec == "hold":
|
|
45
|
+
return reference.HoldStill(horizon=horizon, images=images)
|
|
46
|
+
if spec == "random":
|
|
47
|
+
return reference.RandomWalk(horizon=horizon)
|
|
48
|
+
module_name, sep, attr = spec.partition(":")
|
|
49
|
+
if not sep or not module_name or not attr:
|
|
50
|
+
raise SystemExit(f"--policy must be 'hold', 'random', or 'module:Attr'; got {spec!r}")
|
|
51
|
+
sys.path.insert(0, "") # the working directory, as `python -m` would
|
|
52
|
+
try:
|
|
53
|
+
obj = getattr(importlib.import_module(module_name), attr)
|
|
54
|
+
except (ImportError, AttributeError) as exc:
|
|
55
|
+
raise SystemExit(f"cannot load policy {spec!r}: {exc}") from None
|
|
56
|
+
return obj() if isinstance(obj, type) else obj
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def cmd_eval_run(args: argparse.Namespace) -> int:
|
|
60
|
+
from se3labs.eval.policy import client
|
|
61
|
+
|
|
62
|
+
policy = load_policy(args.policy, horizon=args.horizon, images=args.image)
|
|
63
|
+
try:
|
|
64
|
+
report = client.serve_policy(args.job, args.address, policy, secure=args.secure)
|
|
65
|
+
except client.SessionRejected as exc:
|
|
66
|
+
print(f"rejected: {exc}", file=sys.stderr)
|
|
67
|
+
return 2
|
|
68
|
+
print(f"session closed: {report.close_reason or 'stream ended'}; {report.episodes} resets, {report.steps} steps")
|
|
69
|
+
if report.errors:
|
|
70
|
+
print("policy errors:", *report.errors[:5], sep="\n ", file=sys.stderr)
|
|
71
|
+
return 1
|
|
72
|
+
return 0
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def cmd_eval_status(args: argparse.Namespace) -> int:
|
|
76
|
+
import grpc
|
|
77
|
+
|
|
78
|
+
from se3labs.eval.policy.proto import job_status_pb2, job_status_pb2_grpc
|
|
79
|
+
|
|
80
|
+
channel = (
|
|
81
|
+
grpc.secure_channel(args.address, grpc.ssl_channel_credentials())
|
|
82
|
+
if args.secure
|
|
83
|
+
else grpc.insecure_channel(args.address)
|
|
84
|
+
)
|
|
85
|
+
with channel:
|
|
86
|
+
try:
|
|
87
|
+
status = job_status_pb2_grpc.JobStatusStub(channel).Get(
|
|
88
|
+
job_status_pb2.StatusRequest(job_id=args.job), timeout=args.timeout
|
|
89
|
+
)
|
|
90
|
+
except grpc.RpcError as exc:
|
|
91
|
+
print(f"{exc.code().name}: {exc.details()}", file=sys.stderr)
|
|
92
|
+
return 2
|
|
93
|
+
if args.json:
|
|
94
|
+
print(json.dumps({
|
|
95
|
+
"job_id": status.job_id,
|
|
96
|
+
"state": status.state,
|
|
97
|
+
"episode_index": status.episode_index,
|
|
98
|
+
"episodes_total": status.episodes_total,
|
|
99
|
+
"outcomes": dict(status.outcomes),
|
|
100
|
+
"updated_unix_ms": status.updated_unix_ms,
|
|
101
|
+
"message": status.message,
|
|
102
|
+
}))
|
|
103
|
+
else:
|
|
104
|
+
outcomes = ", ".join(f"{k}={v}" for k, v in sorted(status.outcomes.items())) or "none yet"
|
|
105
|
+
print(f"{status.job_id}: {status.state} episode {status.episode_index}/{status.episodes_total} {outcomes}")
|
|
106
|
+
if status.message:
|
|
107
|
+
print(f" {status.message}")
|
|
108
|
+
return 0
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def build_parser() -> argparse.ArgumentParser:
|
|
112
|
+
parser = argparse.ArgumentParser(prog="se3labs", description="SE3 Labs command line.")
|
|
113
|
+
parser.add_argument("--version", action="version", version=f"%(prog)s {_version()}")
|
|
114
|
+
groups = parser.add_subparsers(dest="group", required=True, metavar="<group>")
|
|
115
|
+
|
|
116
|
+
ev = groups.add_parser("eval", help="policy evaluation on an SE3 Labs station")
|
|
117
|
+
cmds = ev.add_subparsers(dest="command", required=True, metavar="<command>")
|
|
118
|
+
|
|
119
|
+
def connection(p: argparse.ArgumentParser) -> None:
|
|
120
|
+
p.add_argument("--job", required=True, help="job id from SE3 Labs")
|
|
121
|
+
p.add_argument("--address", required=True, help="relay or station host:port")
|
|
122
|
+
p.add_argument("--secure", action="store_true", help="TLS to the address")
|
|
123
|
+
|
|
124
|
+
run = cmds.add_parser("run", help="connect a policy to a job and run every episode")
|
|
125
|
+
connection(run)
|
|
126
|
+
run.add_argument(
|
|
127
|
+
"--policy",
|
|
128
|
+
default="hold",
|
|
129
|
+
help="'hold' or 'random' (reference policies), or module:Attr naming your policy class or instance (default: hold)",
|
|
130
|
+
)
|
|
131
|
+
run.add_argument("--horizon", type=int, default=4, help="actions per chunk for the reference policies (default: 4)")
|
|
132
|
+
run.add_argument(
|
|
133
|
+
"--image",
|
|
134
|
+
action="append",
|
|
135
|
+
default=[],
|
|
136
|
+
type=_parse_image,
|
|
137
|
+
metavar="CAMERA[:WxH]",
|
|
138
|
+
help="request a camera stream, e.g. cam_chest:640x480; repeatable",
|
|
139
|
+
)
|
|
140
|
+
run.set_defaults(func=cmd_eval_run)
|
|
141
|
+
|
|
142
|
+
status = cmds.add_parser("status", help="poll a job's progress from the relay")
|
|
143
|
+
connection(status)
|
|
144
|
+
status.add_argument("--timeout", type=float, default=10.0, help="seconds to wait for the relay (default: 10)")
|
|
145
|
+
status.add_argument("--json", action="store_true", help="print the status as one JSON object")
|
|
146
|
+
status.set_defaults(func=cmd_eval_status)
|
|
147
|
+
return parser
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
def main(argv: Sequence[str] | None = None) -> int:
|
|
151
|
+
args = build_parser().parse_args(argv)
|
|
152
|
+
return args.func(args)
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
if __name__ == "__main__":
|
|
156
|
+
sys.exit(main())
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""Policy evaluation on an SE3 Labs station, seen from the customer's side."""
|
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
"""The customer-facing policy protocol and, later, the client SDK.
|
|
2
|
+
|
|
3
|
+
Importable without hardware dependencies: nothing under this package may
|
|
4
|
+
import the driver, i2rt or MuJoCo. The wire contract lives in
|
|
5
|
+
:mod:`se3labs.eval.policy.proto`; see doc/01-remote-policy-evaluation.md §6.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
#: Bumped on any wire change a client could not ignore. Sent in ``Hello``.
|
|
9
|
+
PROTOCOL_VERSION = 1
|
|
@@ -0,0 +1,170 @@
|
|
|
1
|
+
"""The customer-side client: wrap a policy object and speak the stream.
|
|
2
|
+
|
|
3
|
+
from se3labs.eval.policy import client, reference
|
|
4
|
+
client.serve_policy("job_…", "relay.example.com:7443", reference.HoldStill())
|
|
5
|
+
|
|
6
|
+
A :class:`Policy` answers three things: its manifest (what it wants and
|
|
7
|
+
can do), a reset per episode, and actions per observation. The client
|
|
8
|
+
handles the handshake, echoes every identifier, and returns a report when
|
|
9
|
+
the station closes the session. It needs only ``grpcio`` and ``protobuf``.
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
from __future__ import annotations
|
|
13
|
+
|
|
14
|
+
import queue
|
|
15
|
+
import threading
|
|
16
|
+
from collections.abc import Iterator, Mapping, Sequence
|
|
17
|
+
from dataclasses import dataclass, field
|
|
18
|
+
from typing import Protocol, runtime_checkable
|
|
19
|
+
|
|
20
|
+
import grpc
|
|
21
|
+
|
|
22
|
+
from se3labs.eval.policy import PROTOCOL_VERSION
|
|
23
|
+
from se3labs.eval.policy.proto import policy_session_pb2 as pb
|
|
24
|
+
from se3labs.eval.policy.proto import policy_session_pb2_grpc as rpc
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@runtime_checkable
|
|
28
|
+
class Policy(Protocol):
|
|
29
|
+
def manifest(self) -> pb.ClientManifest: ...
|
|
30
|
+
|
|
31
|
+
def reset(self, reset: pb.Reset) -> None: ...
|
|
32
|
+
|
|
33
|
+
def act(self, observe: pb.Observe) -> Sequence[pb.Action]:
|
|
34
|
+
"""1 to ``negotiated.max_horizon`` actions at the negotiated dt."""
|
|
35
|
+
...
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class SessionRejected(RuntimeError):
|
|
39
|
+
def __init__(self, reject: pb.Reject) -> None:
|
|
40
|
+
super().__init__(f"{pb.RejectCode.Name(reject.code)}: {reject.message}")
|
|
41
|
+
self.reject = reject
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
@dataclass
|
|
45
|
+
class SessionReport:
|
|
46
|
+
welcome: pb.Welcome | None = None
|
|
47
|
+
episodes: int = 0
|
|
48
|
+
steps: int = 0
|
|
49
|
+
close_reason: str = ""
|
|
50
|
+
errors: list[str] = field(default_factory=list)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def joint_action(targets: Mapping[str, tuple[Sequence[float], float | None]]) -> pb.Action:
|
|
54
|
+
"""One time step: arm name → (joint positions in rad, gripper or None)."""
|
|
55
|
+
action = pb.Action()
|
|
56
|
+
for arm, (joints, gripper) in targets.items():
|
|
57
|
+
jp = pb.JointPositionAction(arm=arm, joint_pos_rad=list(joints))
|
|
58
|
+
if gripper is not None:
|
|
59
|
+
jp.gripper_pos = gripper
|
|
60
|
+
action.arms.add(joint_position=jp)
|
|
61
|
+
return action
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def default_manifest(
|
|
65
|
+
*,
|
|
66
|
+
model: str = "reference",
|
|
67
|
+
version: str = "0",
|
|
68
|
+
images: Sequence[tuple[str, int, int]] = (),
|
|
69
|
+
max_horizon: int = 8,
|
|
70
|
+
stateful: bool = False,
|
|
71
|
+
) -> pb.ClientManifest:
|
|
72
|
+
manifest = pb.ClientManifest(
|
|
73
|
+
model=pb.ModelIdentity(name=model, version=version),
|
|
74
|
+
action_schemas=[pb.ACTION_SCHEMA_JOINT_POSITION_V1],
|
|
75
|
+
max_horizon=max_horizon,
|
|
76
|
+
stateful=stateful,
|
|
77
|
+
)
|
|
78
|
+
for camera, w, h in images:
|
|
79
|
+
manifest.observation.images.append(
|
|
80
|
+
pb.ImageRequest(camera=camera, max_width=w, max_height=h, encoding=pb.IMAGE_ENCODING_JPEG)
|
|
81
|
+
)
|
|
82
|
+
return manifest
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
class _Outbox:
|
|
86
|
+
"""Client→server messages as the iterator gRPC consumes."""
|
|
87
|
+
|
|
88
|
+
def __init__(self) -> None:
|
|
89
|
+
self._q: queue.Queue[pb.ClientMsg | None] = queue.Queue()
|
|
90
|
+
|
|
91
|
+
def put(self, msg: pb.ClientMsg) -> None:
|
|
92
|
+
self._q.put(msg)
|
|
93
|
+
|
|
94
|
+
def close(self) -> None:
|
|
95
|
+
self._q.put(None)
|
|
96
|
+
|
|
97
|
+
def __iter__(self) -> Iterator[pb.ClientMsg]:
|
|
98
|
+
return iter(self._q.get, None)
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def run_session(
|
|
102
|
+
stream: Iterator[pb.ServerMsg],
|
|
103
|
+
outbox: _Outbox,
|
|
104
|
+
job_id: str,
|
|
105
|
+
policy: Policy,
|
|
106
|
+
) -> SessionReport:
|
|
107
|
+
"""Drive one session over an already-open stream. Split from the
|
|
108
|
+
transport so the relay and runner tests can exercise it in-process."""
|
|
109
|
+
report = SessionReport()
|
|
110
|
+
outbox.put(pb.ClientMsg(hello=pb.Hello(job_id=job_id, protocol_version=PROTOCOL_VERSION, manifest=policy.manifest())))
|
|
111
|
+
try:
|
|
112
|
+
for msg in stream:
|
|
113
|
+
kind = msg.WhichOneof("msg")
|
|
114
|
+
if kind == "welcome":
|
|
115
|
+
report.welcome = msg.welcome
|
|
116
|
+
elif kind == "reject":
|
|
117
|
+
raise SessionRejected(msg.reject)
|
|
118
|
+
elif kind == "reset":
|
|
119
|
+
policy.reset(msg.reset)
|
|
120
|
+
report.episodes += 1
|
|
121
|
+
outbox.put(pb.ClientMsg(reset_ack=pb.ResetAck(episode_id=msg.reset.episode_id)))
|
|
122
|
+
elif kind == "observe":
|
|
123
|
+
o = msg.observe
|
|
124
|
+
try:
|
|
125
|
+
actions = list(policy.act(o))
|
|
126
|
+
except Exception as exc: # noqa: BLE001 - a policy bug is reported, not fatal
|
|
127
|
+
report.errors.append(f"step {o.step_id}: {exc!r}")
|
|
128
|
+
actions = []
|
|
129
|
+
chunk = pb.Chunk(episode_id=o.episode_id, step_id=o.step_id, request_id=o.request_id, actions=actions)
|
|
130
|
+
outbox.put(pb.ClientMsg(chunk=chunk))
|
|
131
|
+
report.steps += 1
|
|
132
|
+
elif kind == "close":
|
|
133
|
+
report.close_reason = msg.close.reason
|
|
134
|
+
break
|
|
135
|
+
finally:
|
|
136
|
+
outbox.close()
|
|
137
|
+
return report
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def serve_policy(
|
|
141
|
+
job_id: str,
|
|
142
|
+
address: str,
|
|
143
|
+
policy: Policy,
|
|
144
|
+
*,
|
|
145
|
+
secure: bool = False,
|
|
146
|
+
timeout_s: float | None = None,
|
|
147
|
+
) -> SessionReport:
|
|
148
|
+
"""Connect to ``address`` (the relay, or a station directly), run the
|
|
149
|
+
whole session for ``job_id``, and return the report."""
|
|
150
|
+
channel = grpc.secure_channel(address, grpc.ssl_channel_credentials()) if secure else grpc.insecure_channel(address)
|
|
151
|
+
try:
|
|
152
|
+
grpc.channel_ready_future(channel).result(timeout=timeout_s or 30.0)
|
|
153
|
+
stub = rpc.PolicySessionStub(channel)
|
|
154
|
+
outbox = _Outbox()
|
|
155
|
+
stream = stub.Run(iter(outbox), timeout=timeout_s)
|
|
156
|
+
return run_session(stream, outbox, job_id, policy)
|
|
157
|
+
finally:
|
|
158
|
+
channel.close()
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
def serve_policy_in_thread(job_id: str, address: str, policy: Policy) -> tuple[threading.Thread, list[SessionReport]]:
|
|
162
|
+
"""Test helper: the session on a thread, its report appended when done."""
|
|
163
|
+
reports: list[SessionReport] = []
|
|
164
|
+
|
|
165
|
+
def run() -> None:
|
|
166
|
+
reports.append(serve_policy(job_id, address, policy))
|
|
167
|
+
|
|
168
|
+
t = threading.Thread(target=run, name=f"policy-{job_id}", daemon=True)
|
|
169
|
+
t.start()
|
|
170
|
+
return t, reports
|
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
// Job progress, as a customer may poll it from the relay. The station
|
|
2
|
+
// pushes one Status per state change; the relay caches the latest.
|
|
3
|
+
//
|
|
4
|
+
// Regenerate the Python after editing: tool/gen_proto.sh (in the station
|
|
5
|
+
// repo). Design: doc/01-remote-policy-evaluation.md §9.
|
|
6
|
+
|
|
7
|
+
syntax = "proto3";
|
|
8
|
+
|
|
9
|
+
package se3labs.eval.policy.v1;
|
|
10
|
+
|
|
11
|
+
service JobStatus {
|
|
12
|
+
// The latest status the station reported for a job id.
|
|
13
|
+
rpc Get(StatusRequest) returns (Status);
|
|
14
|
+
}
|
|
15
|
+
|
|
16
|
+
message StatusRequest {
|
|
17
|
+
string job_id = 1;
|
|
18
|
+
}
|
|
19
|
+
|
|
20
|
+
message Status {
|
|
21
|
+
string job_id = 1;
|
|
22
|
+
// CREATED, WAITING_POLICY, PREFLIGHT, RUNNING, PAUSED, FINALIZING,
|
|
23
|
+
// SUCCEEDED, FAILED, CANCELLED.
|
|
24
|
+
string state = 2;
|
|
25
|
+
uint32 episode_index = 3;
|
|
26
|
+
uint32 episodes_total = 4;
|
|
27
|
+
// Count per outcome: success, failure, indeterminate, invalidated.
|
|
28
|
+
map<string, uint32> outcomes = 5;
|
|
29
|
+
int64 updated_unix_ms = 6;
|
|
30
|
+
string message = 7;
|
|
31
|
+
}
|