(v, local_devices_to_use)
| 59 | |
| 60 | |
| 61 | def _data_pmap(v, local_devices_to_use): |
| 62 | # devices = jax.local_devices() |
| 63 | # v = v.reshape((local_devices_to_use, v.shape[0] // local_devices_to_use,) + v.shape[1:]) |
| 64 | # v = [jnp.array(v[i]) for (i, device) in enumerate(devices)] |
| 65 | # return jax.device_put_sharded(v, devices) |
| 66 | return v.reshape((local_devices_to_use, v.shape[0]//local_devices_to_use,) + v.shape[1:]) |
| 67 | |
| 68 | |
| 69 | def train(environment: envs.Env, |