mirror of
https://github.com/wassname/ray.git
synced 2026-09-10 12:38:43 +08:00
[RLlib] DQN torch version. (#7597)
* Fix. * Rollback. * WIP. * WIP. * WIP. * WIP. * WIP. * WIP. * WIP. * WIP. * Fix. * Fix. * Fix. * Fix. * Fix. * WIP. * WIP. * Fix. * Test case fixes. * Test case fixes and LINT. * Test case fixes and LINT. * Rollback. * WIP. * WIP. * Test case fixes. * Fix. * Fix. * Fix. * Add regression test for DQN w/ param noise. * Fixes and LINT. * Fixes and LINT. * Fixes and LINT. * Fixes and LINT. * Fixes and LINT. * Comment * Regression test case. * WIP. * WIP. * LINT. * LINT. * WIP. * Fix. * Fix. * Fix. * LINT. * Fix (SAC does currently not support eager). * Fix. * WIP. * LINT. * Update rllib/evaluation/sampler.py Co-Authored-By: Eric Liang <ekhliang@gmail.com> * Update rllib/evaluation/sampler.py Co-Authored-By: Eric Liang <ekhliang@gmail.com> * Update rllib/utils/exploration/exploration.py Co-Authored-By: Eric Liang <ekhliang@gmail.com> * Update rllib/utils/exploration/exploration.py Co-Authored-By: Eric Liang <ekhliang@gmail.com> * WIP. * WIP. * Fix. * LINT. * LINT. * Fix and LINT. * WIP. * WIP. * WIP. * WIP. * Fix. * LINT. * Fix. * Fix and LINT. * Update rllib/utils/exploration/exploration.py * Update rllib/policy/dynamic_tf_policy.py Co-Authored-By: Eric Liang <ekhliang@gmail.com> * Update rllib/policy/dynamic_tf_policy.py Co-Authored-By: Eric Liang <ekhliang@gmail.com> * Update rllib/policy/dynamic_tf_policy.py Co-Authored-By: Eric Liang <ekhliang@gmail.com> * Fixes. * WIP. * LINT. * Fixes and LINT. * LINT and fixes. * LINT. * Move action_dist back into torch extra_action_out_fn and LINT. * Working SimpleQ learning cartpole on both torch AND tf. * Working Rainbow learning cartpole on tf. * Working Rainbow learning cartpole on tf. * WIP. * LINT. * LINT. * Update docs and add torch to APEX test. * LINT. * Fix. * LINT. * Fix. * Fix. * Fix and docstrings. * Fix broken RLlib tests in master. * Split BAZEL learning tests into cartpole and pendulum (reached the 60min barrier). * Fix error_outputs option in BAZEL for RLlib regression tests. * Fix. * Tune param-noise tests. * LINT. * Fix. * Fix. * test * test * test * Fix. * Fix. * WIP. * WIP. * WIP. * WIP. * LINT. * WIP. Co-authored-by: Eric Liang <ekhliang@gmail.com>
This commit is contained in:
@@ -135,16 +135,26 @@ def try_import_torch(error=False):
|
||||
return None, nn
|
||||
|
||||
|
||||
def get_variable(value, framework="tf", tf_name="unnamed-variable"):
|
||||
def get_variable(value,
|
||||
framework="tf",
|
||||
trainable=False,
|
||||
tf_name="unnamed-variable",
|
||||
torch_tensor=False):
|
||||
"""
|
||||
Args:
|
||||
value (any): The initial value to use. In the non-tf case, this will
|
||||
be returned as is.
|
||||
framework (str): One of "tf", "torch", or None.
|
||||
tf_name (str): An optional name for the variable. Only for tf.
|
||||
trainable (bool): Whether the generated variable should be
|
||||
trainable (tf)/require_grad (torch) or not (default: False).
|
||||
tf_name (str): For framework="tf": An optional name for the
|
||||
tf.Variable.
|
||||
torch_tensor (bool): For framework="torch": Whether to actually create
|
||||
a torch.tensor, or just a python value (default).
|
||||
|
||||
Returns:
|
||||
any: A framework-specific variable (tf.Variable or python primitive).
|
||||
any: A framework-specific variable (tf.Variable, torch.tensor, or
|
||||
python primitive).
|
||||
"""
|
||||
if framework == "tf":
|
||||
import tensorflow as tf
|
||||
@@ -153,7 +163,12 @@ def get_variable(value, framework="tf", tf_name="unnamed-variable"):
|
||||
if isinstance(value, float) else tf.int32
|
||||
if isinstance(value, int) else None)
|
||||
return tf.compat.v1.get_variable(
|
||||
tf_name, initializer=value, dtype=dtype)
|
||||
tf_name, initializer=value, dtype=dtype, trainable=trainable)
|
||||
elif framework == "torch" and torch_tensor is True:
|
||||
import torch
|
||||
var_ = torch.from_numpy(value)
|
||||
var_.requires_grad = trainable
|
||||
return var_
|
||||
# torch or None: Return python primitive.
|
||||
return value
|
||||
|
||||
|
||||
Reference in New Issue
Block a user