Skip to content

Commit 31e61cb

Browse files
authored
Merge pull request #325 from RobotControlStack/juelg/bump-agents
bump(examples): vlagents interface refactor
2 parents 3366833 + 61de399 commit 31e61cb

4 files changed

Lines changed: 48 additions & 44 deletions

File tree

‎examples/inference/README.md‎

Lines changed: 9 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -12,13 +12,12 @@ Before starting `franka.py`, make sure a `vlagents` policy server is already run
1212

1313
The policy server setup and supported launch commands are documented in:
1414

15-
- [RobotControlStack/vlagents](https://github.com/RobotControlStack/vlagents)
16-
- [vlagents/README.md](../../vlagents/README.md)
15+
- [vlagents](https://github.com/RobotControlStack/vlagents)
1716

1817
Typical server startup looks like:
1918

2019
```shell
21-
python -m vlagents start-server lerobot --port 20000 --host 0.0.0.0 --kwargs '{"policy_name": "act", "checkpoint_path": "<path to pretrained_model>", "n_action_steps": 1}'
20+
uv run python -m vlagents start-server lerobot --port 20000 --host 0.0.0.0 --kwargs '{"policy_name": "act", "checkpoint_path": "<path to pretrained_model>"}'
2221
```
2322

2423
For other policies such as `pi05` or `xvla`, use the matching startup command from the `vlagents` README and make sure the values in `franka.json` point at that server.
@@ -31,12 +30,13 @@ For other policies such as `pi05` or `xvla`, use the matching startup command fr
3130
- `vlagents_port`: Port exposed by the policy server.
3231
- `vlagents_model`: Agent id passed to `vlagents`, for example `lerobot`.
3332
- `instruction`: Natural-language task instruction sent to the policy on reset.
34-
- `robot_keys`: Robot ordering used to pack observations and unpack actions. The script assumes one 8-value action block per robot in this order: `7` joint values plus `1` gripper value.
33+
- `robot_keys`: Robot names expected in each returned action dictionary and used to construct per-robot observations.
3534
- `jpeg_encoding`: Whether observations are sent to the policy server using JPEG-compressed images.
3635
- `on_same_machine`: Set this according to whether the policy server runs on the same machine as the control process.
36+
- `image_size`: Client-side `(width, height)` resize applied before JPEG or shared-memory transport; defaults to `[224, 224]`. Set it to `null` to retain native resolution.
3737
- `fps`: Control loop target frequency used by the local rate limiter.
3838
- `record_path`: Output directory used when recording episodes.
39-
- `n_action_steps`: If `null`, the script requests one action per control step. If set to an integer greater than `0`, the script buffers that many actions from each policy response chunk.
39+
- `n_action_steps`: Local action-chunk execution horizon. If `null`, the script requests and executes one action per control step. If set to a positive integer, it buffers up to that many actions from each policy response chunk.
4040
- `max_rel_mov_joints`: Maximum allowed relative joint movement per step when running in joint control mode.
4141
- `max_rel_mov_cart`: Maximum allowed relative Cartesian translation and rotation per step when running in Cartesian modes.
4242

@@ -57,23 +57,17 @@ When [franka.py](franka.py) is running, it waits for keyboard input on stdin. Th
5757

5858
The script translates RCS observations to the `vlagents` `Obs` format as follows:
5959

60-
- Every camera frame in `obs["frames"]` is converted to RGB and resized to `224x224`.
61-
- State is built by iterating through `robot_keys` in order and concatenating each robot's `joints` and `gripper` values.
60+
- Camera frames are passed to `RemoteAgent` at native resolution; the client resizes them to `image_size` before JPEG or shared-memory transport.
61+
- Each robot gets a `SingleObs` containing the shared camera set plus its own joints and gripper state.
6262

63-
Action decoding is also order-dependent:
64-
65-
- For each robot in `robot_keys`, the script reads `8` values from the policy action vector.
66-
- Values `0:7` become the robot joint command.
67-
- Value `7:8` becomes the robot gripper command.
68-
69-
That means `robot_keys` must match the policy's expected robot ordering exactly.
63+
Action chunks contain one action dictionary per environment step. For each robot, the script forwards `SingleAct.action` as the joint command and `SingleAct.gripper` as the gripper command. The action dictionary must include every configured `robot_key`.
7064

7165
## Running
7266

7367
After the policy server is up and `franka.json` is configured, run:
7468

7569
```shell
76-
python examples/inference/franka.py
70+
uv run python examples/inference/franka.py
7771
```
7872

7973
If the policy server is unreachable, the script will keep retrying connection until it becomes available or you exit.

‎examples/inference/franka.json‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
],
1010
"jpeg_encoding": true,
1111
"on_same_machine": false,
12+
"image_size": [224, 224],
1213
"fps": 30,
1314
"record_path": "inference_recordings_bin_sort_duobench_xvla_bin_sort_real_2026-05-20_23-25-47_040000",
1415
"n_action_steps": 30,

‎examples/inference/franka.py‎

Lines changed: 37 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,6 @@
1010

1111
import gymnasium as gym
1212
import numpy as np
13-
from PIL import Image
1413
from rcs._core.common import BaseCameraConfig, RobotPlatform
1514
from rcs._core.sim import SimConfig
1615
from rcs.envs.base import ControlMode, RelativeTo
@@ -20,7 +19,7 @@
2019

2120
# from rcs_duobench.tasks.bin_sort import BinSortEnvConfig
2221
from vlagents.client import RemoteAgent
23-
from vlagents.policies import Act, Obs
22+
from vlagents.policies.interface import Obs, SingleAct, SingleObs
2423

2524
import rcs
2625

@@ -103,6 +102,7 @@ class InferenceConfig:
103102
robot_keys: list[str] = field(default_factory=lambda: ["left", "right"])
104103
jpeg_encoding: bool = True
105104
on_same_machine: bool = False
105+
image_size: tuple[int, int] | None = (224, 224)
106106
fps: int = FPS
107107
record_path: str = RECORD_PATH
108108
n_action_steps: int | None = None
@@ -126,7 +126,12 @@ def __init__(self, env: gym.Env, cfg: InferenceConfig):
126126
self._command_queue: Queue[str] = Queue()
127127
self._shutdown_requested = threading.Event()
128128
self.remote_agent = RemoteAgent(
129-
cfg.vlagents_host, cfg.vlagents_port, cfg.vlagents_model, cfg.on_same_machine, cfg.jpeg_encoding
129+
cfg.vlagents_host,
130+
cfg.vlagents_port,
131+
cfg.vlagents_model,
132+
cfg.on_same_machine,
133+
cfg.jpeg_encoding,
134+
cfg.image_size,
130135
)
131136
self.frame_rate = SimpleFrameRate(self._cfg.fps)
132137
self._action_buffer = []
@@ -164,41 +169,45 @@ def _drain_commands(self) -> tuple[bool, bool, bool, bool, bool]:
164169
return start_requested, record_requested, success_requested, stop_requested, reload_requested
165170

166171
def obs_rcs2agents(self, obs: dict, info: dict | None = None) -> Obs:
167-
cameras = {}
168-
for frame in obs["frames"]:
169-
cameras[frame] = obs["frames"][frame]["rgb"]["data"]
170-
cameras[frame] = np.array(Image.fromarray(cameras[frame]).resize((224, 224), Image.Resampling.BILINEAR))
172+
cameras = {frame: obs["frames"][frame]["rgb"]["data"] for frame in obs["frames"]}
171173

172-
state = []
174+
obs_by_robot = {}
173175
for robot in self._cfg.robot_keys:
174-
# TODO: currently hardcoded for joints
175-
state.append(obs[robot]["joints"])
176-
state.append(obs[robot]["gripper"])
176+
obs_by_robot[robot] = SingleObs(
177+
cameras=copy.deepcopy(cameras),
178+
joints=np.asarray(obs[robot]["joints"], dtype=np.float32),
179+
gripper=float(obs[robot]["gripper"]),
180+
xyzrpy=np.asarray(obs[robot]["xyzrpy"], dtype=np.float32) if "xyzrpy" in obs[robot] else None,
181+
tquat=np.asarray(obs[robot]["tquat"], dtype=np.float32) if "tquat" in obs[robot] else None,
182+
# info=copy.deepcopy(info) if info is not None else {},
183+
)
177184

178-
return Obs(cameras=cameras, gripper=None, info=info, state=np.concatenate(state))
185+
return Obs(obs=obs_by_robot, language_instruction=self._cfg.instruction)
179186

180-
def act(self, obs_dict) -> None:
181-
done = False
187+
def act(self, obs_dict: Obs) -> dict[str, SingleAct]:
182188
if self._cfg.n_action_steps is None:
183-
return self.remote_agent.act(obs_dict)
189+
action_chunk = self.remote_agent.act(obs_dict).acts
190+
if not action_chunk:
191+
message = "Received empty action chunk from policy"
192+
raise ValueError(message)
193+
return action_chunk[0]
184194
if len(self._action_buffer) == 0:
185195
action = self.remote_agent.act(obs_dict)
186-
selected_action = action.action[: self._cfg.n_action_steps]
187-
self._action_buffer = selected_action.tolist()
188-
done = action.done
196+
selected_action = action.acts[: self._cfg.n_action_steps]
197+
self._action_buffer = list(selected_action)
189198
if RELATIVETO == RelativeTo.CONFIGURED_ORIGIN:
190199
for robot in self.env.get_wrapper_attr("envs"):
191200
self.env.get_wrapper_attr("envs")[robot].get_wrapper_attr("set_origin_to_current")()
192-
act = self._action_buffer.pop(0)
193-
return Act(action=act, done=done)
201+
return self._action_buffer.pop(0)
194202

195-
def action_agents2rcs(self, action: Act) -> dict[str, Any]:
203+
def action_agents2rcs(self, action: dict[str, SingleAct]) -> dict[str, Any]:
196204
act = {}
197-
for idx, robot in enumerate(self._cfg.robot_keys):
198-
# TODO: this is currently hard coded for franka joints
199-
act[robot] = {}
200-
act[robot]["joints"] = action.action[idx * 8 : idx * 8 + 7]
201-
act[robot]["gripper"] = action.action[idx * 8 + 7 : idx * 8 + 8]
205+
for robot in self._cfg.robot_keys:
206+
robot_action = action[robot]
207+
act[robot] = {
208+
"joints": np.asarray(robot_action.action, dtype=np.float32),
209+
"gripper": np.asarray([robot_action.gripper], dtype=np.float32),
210+
}
202211
return act
203212

204213
def loop(self):
@@ -222,6 +231,7 @@ def loop(self):
222231
model=self._cfg.vlagents_model,
223232
on_same_machine=self._cfg.on_same_machine,
224233
jpeg_encoding=self._cfg.jpeg_encoding,
234+
image_size=self._cfg.image_size,
225235
)
226236
logger.info(
227237
"reloaded config from %s with host=%s port=%s model=%s",
@@ -269,14 +279,13 @@ def loop(self):
269279
if record_requested:
270280
self.env.start_record()
271281
logger.info("starting episode%s", " with recording" if record_requested else "")
272-
self.remote_agent.reset(copy.deepcopy(obs_dict), instruction=self._cfg.instruction)
273282
self._episode_running = True
274283
else:
275284
sleep(0.05)
276285
continue
277286

278287
action = self.act(copy.deepcopy(obs_dict))
279-
if action.done:
288+
if any(robot_action.done for robot_action in action.values()):
280289
logger.info("done issued by agent, resetting environment")
281290
obs, _ = self.env.reset()
282291
obs_dict = self.obs_rcs2agents(obs)
Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
vlagents @ git+https://github.com/RobotControlStack/vlagents.git@lerobot
1+
vlagents==0.3.0

0 commit comments

Comments
 (0)