Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 32 additions & 13 deletions ss2r/algorithms/sac/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
from brax.training import acting
from brax.training.acme import running_statistics
from brax.training.replay_buffers import ReplayBuffer
from brax.training.types import Params, PRNGKey
from brax.training.types import Params, PRNGKey, Transition

from ss2r.algorithms.sac.go1_sac_to_onnx import convert_policy_to_onnx
from ss2r.algorithms.sac.types import CollectDataFn, ReplayBufferState, float16
Expand All @@ -18,23 +18,30 @@ def get_collection_fn(cfg):
if cfg.agent.data_collection.name == "step":
return collect_single_step
elif cfg.agent.data_collection.name == "episodic":
fn = (
lambda env,

def generate_episodic_unroll(
env,
env_state,
make_policy_fn,
policy_params,
key,
extra_fields: generate_unroll(
episode_length,
extra_fields,
):
env_state, transitions = acting.generate_unroll(
env,
env_state,
make_policy_fn,
policy_params,
make_policy_fn(policy_params),
key,
cfg.training.episode_length,
episode_length,
extra_fields,
)
)
return make_collection_fn(fn)
transitions = jax.tree.map(
lambda x: x.reshape(-1, *x.shape[2:]), transitions
)
return env_state, transitions

return make_collection_fn(generate_episodic_unroll)
elif cfg.agent.data_collection.name == "hardware":
data_collection_cfg = cfg.agent.data_collection
if "Go1" in cfg.environment.task_name:
Expand All @@ -46,6 +53,7 @@ def get_collection_fn(cfg):
orchestrator = OnlineEpisodeOrchestrator(
policy_translate_fn,
cfg.training.episode_length,
go1_postprocess_data,
data_collection_cfg.address,
)
return make_collection_fn(orchestrator.request_data)
Expand Down Expand Up @@ -104,10 +112,6 @@ def collect_data(
key,
extra_fields=extra_fields,
)
if transitions.reward.ndim == 2:
transitions = jax.tree.map(
lambda x: x.reshape(-1, *x.shape[2:]), transitions
)
normalizer_params = running_statistics.update(
normalizer_params, transitions.observation
)
Expand All @@ -124,3 +128,18 @@ def make_go1_policy(make_policy_fn, params, cfg):
del make_policy_fn
proto_model = convert_policy_to_onnx(params, cfg, 12, 48)
return proto_model.SerializeToString()


def go1_postprocess_data(raw_data, extra_fields):
observation, action, reward, next_observation, done, info = raw_data
state_extras = {x: info[x] for x in extra_fields}
policy_extras = {}
transitions = Transition(
observation=observation,
action=action,
reward=reward,
discount=1 - done,
next_observation=next_observation,
extras={"policy_extras": policy_extras, "state_extras": state_extras},
)
return transitions
52 changes: 37 additions & 15 deletions ss2r/rl/online.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@

import cloudpickle as pickle
import jax
import jax.numpy as jnp
import zmq
from brax import envs
from brax.training import acting
Expand All @@ -11,12 +12,16 @@

from ss2r.rl.types import MakePolicyFn

_REQUEST_TIMEOUT = 120000
_REQUEST_RETRIES = 120


class OnlineEpisodeOrchestrator:
def __init__(
self,
translate_policy_to_binary_fn,
num_steps,
data_postprocess_fn=lambda x, y: x,
address="tcp://localhost:5555",
):
"""Orchestrator for requesting episodes over ZMQ, with optional SSH reverse tunnel.
Expand All @@ -40,6 +45,7 @@ def __init__(
SSH target (e.g., 'user@host[:port]') to tunnel through.
"""
self._translate_policy_to_binary_fn = translate_policy_to_binary_fn
self._data_postprocess_fn = data_postprocess_fn
self.num_steps = num_steps
self._address = address

Expand All @@ -53,37 +59,53 @@ def request_data(
*,
extra_fields: Sequence[str],
) -> Tuple[envs.State, Transition]:
dummy_transitions = acting.generate_unroll(
dummy_transitions = acting.actor_step(
env,
env_state,
make_policy_fn(policy_params),
key,
self.num_steps,
extra_fields,
)[1]
dummy_transitions = jax.tree.map(lambda x: x.squeeze(1), dummy_transitions)
dummy_transitions = jax.tree.map(
lambda x: jnp.tile(x, (self.num_steps,) + (1,) * (x.ndim - 1)),
dummy_transitions,
)
transitions = io_callback(
functools.partial(self._send_request, make_policy_fn),
functools.partial(self._send_request, make_policy_fn, extra_fields),
dummy_transitions,
policy_params,
ordered=True,
)
state_extras = {x: transitions.extras["state_extras"][x] for x in extra_fields}
transitions.extras["state_extras"] = state_extras
return env_state, transitions

def _send_request(self, make_policy_fn, policy_params):
def _send_request(self, make_policy_fn, extra_fields, policy_params):
"""Implements a lazy pirate client reliability pattern"""
policy_bytes = self._translate_policy_to_binary_fn(
make_policy_fn, policy_params
)
with zmq.Context() as ctx:
with ctx.socket(zmq.REQ) as socket:
socket.connect(self._address)
while True:
socket = ctx.socket(zmq.REQ)
socket.connect(self._address)
retries_left = _REQUEST_RETRIES
print("Requesting data...")
# Send data
socket.send(pickle.dumps((policy_bytes, self.num_steps)))
while True:
if (socket.poll(_REQUEST_TIMEOUT) & zmq.POLLIN) != 0:
# Receive response
raw_data = pickle.loads(socket.recv())
transitions = self._data_postprocess_fn(raw_data, extra_fields)
print(f"Received {len(transitions.reward)} transitions...")
return transitions
else:
retries_left -= 1
print(f"Request timed out, {retries_left} retries left...")
socket.setsockopt(zmq.LINGER, 0)
socket.close()
if retries_left == 0:
raise RuntimeError("Request timed out.")
print("Retrying...")
socket = ctx.socket(zmq.REQ)
socket.connect(self._address)
print("Requesting data...")
# Send data
socket.send(pickle.dumps((policy_bytes, self.num_steps)))
# Receive response
success, transitions = pickle.loads(socket.recv())
if success:
return transitions
83 changes: 0 additions & 83 deletions tests/test_online_orchestrator.py

This file was deleted.