instadeepai / instadeepai/flashbax
[BUG] Replay buffer on CPU is much slower than the performance in doc.
- Dominant language
- Python
- Stars
- 283
- Forks
- 22
- PR merge metrics
- No merged PRs in 30d
Description
### Describe the bug
I am trying to put the buffer on CPU to alleviate the memory shortage on GPU, but encountered a large speed degradation.
As shown in `readme.md`, the speed on CPU should be much faster than that on GPU or TPU. However in my case the speed on CPU is much slower than both.
### To Reproduce
This is a minimal script to test the speed.
```Python
import os
# os.environ["JAX_PLATFORM_NAME"] = "cpu"
from types import SimpleNamespace
import jax
import jax.numpy as jnp
from flax import struct
import timeit
import flashbax as fbx
from flashbax.buffers.prioritised_trajectory_buffer import PrioritisedTrajectoryBufferState
config = SimpleNamespace()
config.num_envs = 16
config.batch_size = 256
config.obs_n_stack = 1
config.num_unroll_steps = 5
config.buffer_size = 2000000
config.priority_prob_alpha = 0.5
config.max_moves = 100
config.td_steps = 5
config.start_transitions = 400
config.obj_size = 100
config.device = jax.default_backend()
print(f"jax backend: {jax.default_backend()}")
@struct.dataclass
class BufferState:
_state: PrioritisedTrajectoryBufferState
class ReplayBuffer:
def __init__(self, config):
self.buffer = fbx.make_prioritised_trajectory_buffer(
add_batch_size=config.num_envs,
sample_batch_size=config.batch_size,
sample_sequence_length=config.obs_n_stack + config.num_unroll_steps + config.td_steps,
period=1,
min_length_time_axis=(config.start_transitions + config.num_envs - 1) // config.num_envs,
max_length_time_axis=(config.buffer_size + config.num_envs - 1) // config.num_envs,
priority_exponent=config.priority_prob_alpha,
device=config.device,
)
self.config = config
def init(self):
xx = jnp.zeros((config.obj_size,), dtype=jnp.float32)
_state = self.buffer.init(xx)
init_state = BufferState(
_state=_state,
)
return init_state
def add(self, state: BufferState, batched_trajs) -> BufferState:
_state = self.buffer.add(state._state, batched_trajs)
return BufferState(
_state=_state
)
# initialize training components
buffer = ReplayBuffer(config)
buffer_state = buffer.init()
def buf_add(buffer_state):
his = jnp.ones((config.num_envs, config.max_moves, config.obj_size), dtype=jnp.float32)
buffer_state = buffer.add(buffer_state, his)
return buffer_state
jit_buf_add = jax.jit(buf_add)
jit_buf_add_donate = jax.jit(buf_add, donate_argnums=(0,))
# warmup
buffer_state = jit_buf_add(buffer_state)
jax.block_until_ready(buffer_state)
buffer_state = jit_buf_add_donate(buffer_state)
jax.block_until_ready(buffer_state)
tot = 10
start_time = timeit.default_timer()
for i in range(tot):
buffer_state = jit_buf_add(buffer_state)
jax.block_until_ready(buffer_state)
execution_time = timeit.default_timer() - start_time
print('jit_buf_add execution_time (sec):', execution_time / tot)
start_time = timeit.default_timer()
for i in range(tot):
buffer_state = jit_buf_add_donate(buffer_state)
jax.block_until_ready(buffer_state)
execution_time = timeit.default_timer() - start_time
print('jit_buf_add_donate execution_time (sec):', execution_time / tot)
```
The output is:
```
$ python buf.py
jax backend: cpu
jit_buf_add execution_time (sec): 0.4437830951996148
jit_buf_add_donate execution_time (sec): 0.20866907988674938
$ python buf.py
jax backend: gpu
jit_buf_add execution_time (sec): 0.002522601559758186
jit_buf_add_donate execution_time (sec): 0.0025197524111717938
```
### Context (Environment)
```
brax 0.10.3 pypi_0 pypi
flashbax 0.1.3 pypi_0 pypi
flax 0.10.7 pypi_0 pypi
gymnax 0.0.9 pypi_0 pypi
jax 0.6.2 pypi_0 pypi
jax-cuda12-pjrt 0.6.2 pypi_0 pypi
jax-cuda12-plugin 0.6.2 pypi_0 pypi
jaxlib 0.6.2 pypi_0 pypi
jaxmarl 0.0.7 pypi_0 pypi
jaxopt 0.8.5 pypi_0 pypi
optax 0.2.5 pypi_0 pypi
orbax-checkpoint 0.11.14 pypi_0 pypi
```
Contributor guide
Assessment
This issue has not been assessed yet.