Skip to content
Draft
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
4 changes: 4 additions & 0 deletions examples/rlhf/config/train.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -28,3 +28,7 @@ sys:
device: cuda # examples: cpu, cuda, cuda:0, cuda:1 etc., or try mps on macbooks
dtype: bfloat16 # float32, bfloat16, or float16, the latter will auto implement a GradScaler
compile: True # use PyTorch 2.0 to compile the model to be faster

hydra:
job:
chdir: true
4 changes: 4 additions & 0 deletions examples/rlhf/config/train_reward.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -30,3 +30,7 @@ sys:
device: cuda # examples: 'cpu', 'cuda', 'cuda:0', 'cuda:1' etc., or try 'mps' on macbooks
dtype: bfloat16 # 'float32', 'bfloat16', or 'float16', the latter will auto implement a GradScaler
compile: True # use PyTorch 2.0 to compile the model to be faster

hydra:
job:
chdir: true
4 changes: 4 additions & 0 deletions examples/rlhf/config/train_rlhf.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -39,3 +39,7 @@ sys:
ref_device: cuda:1 # device of reference model
dtype: bfloat16 # 'float32', 'bfloat16', or 'float16', the latter will auto implement a GradScaler
compile: False # use PyTorch 2.0 to compile the model to be faster

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion examples/rlhf/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,7 @@ def estimate_loss(model, dataloader):
return estimate_loss


@hydra.main(version_base="1.1", config_path="config", config_name="train")
@hydra.main(version_base="1.3", config_path="config", config_name="train")
def main(cfg):
loss_logger = get_file_logger("loss_logger", "transformer_loss_logger.log")

Expand Down
2 changes: 1 addition & 1 deletion examples/rlhf/train_reward.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,7 +45,7 @@ def estimate_loss(model, dataloader):
return estimate_loss


@hydra.main(version_base="1.1", config_path="config", config_name="train_reward")
@hydra.main(version_base="1.3", config_path="config", config_name="train_reward")
def main(cfg):
loss_logger = get_file_logger("loss_logger", "reward_loss_logger.log")

Expand Down
2 changes: 1 addition & 1 deletion examples/rlhf/train_rlhf.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@
)


@hydra.main(version_base="1.1", config_path="config", config_name="train_rlhf")
@hydra.main(version_base="1.3", config_path="config", config_name="train_rlhf")
def main(cfg):

# ============ Retrieve config ============ #
Expand Down
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -51,7 +51,7 @@ dependencies = [
dev = [
"gym",
"gymnasium",
"hydra-core>=1.1",
"hydra-core>=1.3,<1.4",
"hydra-submitit-launcher",
"imageio",
"moviepy",
Expand Down Expand Up @@ -121,7 +121,7 @@ utils = [
"tensorboard",
"wandb",
"tqdm",
"hydra-core>=1.1",
"hydra-core>=1.3,<1.4",
"hydra-submitit-launcher",
"pre-commit",
"autoflake",
Expand Down
2 changes: 1 addition & 1 deletion sota-implementations/a2c/a2c_atari.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
torch.set_float32_matmul_precision("high")


@hydra.main(config_path="", config_name="config_atari", version_base="1.1")
@hydra.main(config_path="", config_name="config_atari", version_base="1.3")
def main(cfg: DictConfig): # noqa: F821

from copy import deepcopy
Expand Down
2 changes: 1 addition & 1 deletion sota-implementations/a2c/a2c_mujoco.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
torch.set_float32_matmul_precision("high")


@hydra.main(config_path="", config_name="config_mujoco", version_base="1.1")
@hydra.main(config_path="", config_name="config_mujoco", version_base="1.3")
def main(cfg: DictConfig): # noqa: F821

from copy import deepcopy
Expand Down
4 changes: 4 additions & 0 deletions sota-implementations/a2c/config_atari.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -41,3 +41,7 @@ compile:
compile: False
compile_mode:
cudagraphs: False

hydra:
job:
chdir: true
4 changes: 4 additions & 0 deletions sota-implementations/a2c/config_mujoco.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -37,3 +37,7 @@ compile:
compile: False
compile_mode: default
cudagraphs: False

hydra:
job:
chdir: true
4 changes: 4 additions & 0 deletions sota-implementations/a2c_trainer/config/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -137,3 +137,7 @@ trainer:
optim_steps_per_batch: null
num_epochs: 1
async_collection: false

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion sota-implementations/a2c_trainer/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
from torchrl.trainers.algorithms.configs import A2CTrainerConfig # noqa: F401


@hydra.main(config_path="config", config_name="config", version_base="1.1")
@hydra.main(config_path="config", config_name="config", version_base="1.3")
def main(cfg):
def print_reward(td):
torchrl.logger.info(f"reward: {td['next', 'reward'].mean(): 4.4f}")
Expand Down
2 changes: 1 addition & 1 deletion sota-implementations/cql/cql_offline.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@
torch.set_float32_matmul_precision("high")


@hydra.main(config_path="", config_name="offline_config", version_base="1.1")
@hydra.main(config_path="", config_name="offline_config", version_base="1.3")
def main(cfg: DictConfig): # noqa: F821
# Create logger
exp_name = generate_exp_name("CQL-offline", cfg.logger.exp_name)
Expand Down
2 changes: 1 addition & 1 deletion sota-implementations/cql/cql_online.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@
torch.set_float32_matmul_precision("high")


@hydra.main(version_base="1.1", config_path="", config_name="online_config")
@hydra.main(version_base="1.3", config_path="", config_name="online_config")
def main(cfg: DictConfig): # noqa: F821
# Create logger
exp_name = generate_exp_name("CQL-online", cfg.logger.exp_name)
Expand Down
2 changes: 1 addition & 1 deletion sota-implementations/cql/discrete_cql_offline.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@
torch.set_float32_matmul_precision("high")


@hydra.main(version_base="1.1", config_path="", config_name="discrete_offline_config")
@hydra.main(version_base="1.3", config_path="", config_name="discrete_offline_config")
def main(cfg): # noqa: F821
device = (
torch.device(cfg.optim.device) if cfg.optim.device else get_available_device()
Expand Down
2 changes: 1 addition & 1 deletion sota-implementations/cql/discrete_cql_online.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@
torch.set_float32_matmul_precision("high")


@hydra.main(version_base="1.1", config_path="", config_name="discrete_online_config")
@hydra.main(version_base="1.3", config_path="", config_name="discrete_online_config")
def main(cfg: DictConfig): # noqa: F821
device = (
torch.device(cfg.optim.device) if cfg.optim.device else get_available_device()
Expand Down
4 changes: 4 additions & 0 deletions sota-implementations/cql/discrete_offline_config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -64,3 +64,7 @@ compile:
compile: False
compile_mode:
cudagraphs: False

hydra:
job:
chdir: true
4 changes: 4 additions & 0 deletions sota-implementations/cql/discrete_online_config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -62,3 +62,7 @@ compile:
compile: False
compile_mode:
cudagraphs: False

hydra:
job:
chdir: true
4 changes: 4 additions & 0 deletions sota-implementations/cql/offline_config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -59,3 +59,7 @@ compile:
compile: False
compile_mode:
cudagraphs: False

hydra:
job:
chdir: true
4 changes: 4 additions & 0 deletions sota-implementations/cql/online_config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -71,3 +71,7 @@ compile:
compile: False
compile_mode:
cudagraphs: False

hydra:
job:
chdir: true
4 changes: 4 additions & 0 deletions sota-implementations/cql_trainer/config/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -144,3 +144,7 @@ trainer:
save_trainer_file: null
optim_steps_per_batch: 64
async_collection: false

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion sota-implementations/cql_trainer/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
from torchrl.trainers.algorithms.configs import * # noqa: F401, F403


@hydra.main(config_path="config", config_name="config", version_base="1.1")
@hydra.main(config_path="config", config_name="config", version_base="1.3")
def main(cfg):
trainer = hydra.utils.instantiate(cfg.trainer)
trainer.train()
Expand Down
4 changes: 4 additions & 0 deletions sota-implementations/crossq/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -61,3 +61,7 @@ logger:
exp_name: ${env.name}_CrossQ
mode: online
eval_iter: 25000

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion sota-implementations/crossq/crossq.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@
torch.set_float32_matmul_precision("high")


@hydra.main(version_base="1.1", config_path=".", config_name="config")
@hydra.main(version_base="1.3", config_path=".", config_name="config")
def main(cfg: DictConfig): # noqa: F821
device = (
torch.device(cfg.network.device)
Expand Down
4 changes: 4 additions & 0 deletions sota-implementations/ddpg/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -55,3 +55,7 @@ logger:
eval_iter: 25000
video: False
num_eval_envs: 1

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion sota-implementations/ddpg/ddpg.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@
)


@hydra.main(version_base="1.1", config_path="", config_name="config")
@hydra.main(version_base="1.3", config_path="", config_name="config")
def main(cfg: DictConfig): # noqa: F821
device = (
torch.device(cfg.optim.device) if cfg.optim.device else get_available_device()
Expand Down
4 changes: 4 additions & 0 deletions sota-implementations/ddpg_trainer/config/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -147,3 +147,7 @@ trainer:
save_trainer_file: null
optim_steps_per_batch: 64
async_collection: false

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion sota-implementations/ddpg_trainer/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
from torchrl.trainers.algorithms.configs import * # noqa: F401, F403


@hydra.main(config_path="config", config_name="config", version_base="1.1")
@hydra.main(config_path="config", config_name="config", version_base="1.3")
def main(cfg):
trainer = hydra.utils.instantiate(cfg.trainer)
trainer.train()
Expand Down
2 changes: 1 addition & 1 deletion sota-implementations/decision_transformer/dt.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@
)


@hydra.main(config_path="", config_name="dt_config", version_base="1.1")
@hydra.main(config_path="", config_name="dt_config", version_base="1.3")
def main(cfg: DictConfig): # noqa: F821
set_gym_backend(cfg.env.backend).set()

Expand Down
4 changes: 4 additions & 0 deletions sota-implementations/decision_transformer/dt_config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -71,3 +71,7 @@ transformer:
n_positions: 1024
resid_pdrop: 0.1
attn_pdrop: 0.1

hydra:
job:
chdir: true
4 changes: 4 additions & 0 deletions sota-implementations/decision_transformer/odt_config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -72,3 +72,7 @@ transformer:
n_positions: 1024
resid_pdrop: 0.1
attn_pdrop: 0.1

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion sota-implementations/decision_transformer/online_dt.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@
)


@hydra.main(config_path="", config_name="odt_config", version_base="1.1")
@hydra.main(config_path="", config_name="odt_config", version_base="1.3")
def main(cfg: DictConfig): # noqa: F821
set_gym_backend(cfg.env.backend).set()

Expand Down
4 changes: 4 additions & 0 deletions sota-implementations/diffusion_bc/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -45,3 +45,7 @@ compile:
compile: False
compile_mode:
cudagraphs: False

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion sota-implementations/diffusion_bc/diffusion_bc.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@
)


@hydra.main(config_path="", config_name="config", version_base="1.1")
@hydra.main(config_path="", config_name="config", version_base="1.3")
def main(cfg: DictConfig): # noqa: F821
set_gym_backend(cfg.env.library).set()

Expand Down
4 changes: 4 additions & 0 deletions sota-implementations/discrete_sac/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -58,3 +58,7 @@ logger:
mode: online
eval_iter: 5000
video: False

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion sota-implementations/discrete_sac/discrete_sac.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@
)


@hydra.main(version_base="1.1", config_path="", config_name="config")
@hydra.main(version_base="1.3", config_path="", config_name="config")
def main(cfg: DictConfig): # noqa: F821
device = cfg.network.device
if device in ("", None):
Expand Down
4 changes: 4 additions & 0 deletions sota-implementations/dqn/config_atari.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -45,3 +45,7 @@ compile:
compile: False
compile_mode: default
cudagraphs: False

hydra:
job:
chdir: true
4 changes: 4 additions & 0 deletions sota-implementations/dqn/config_cartpole.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -52,3 +52,7 @@ compile:
compile: False
compile_mode:
cudagraphs: False

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion sota-implementations/dqn/dqn_atari.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@
torch.set_float32_matmul_precision("high")


@hydra.main(config_path="", config_name="config_atari", version_base="1.1")
@hydra.main(config_path="", config_name="config_atari", version_base="1.3")
def main(cfg: DictConfig): # noqa: F821

device = torch.device(cfg.device) if cfg.device else get_available_device()
Expand Down
2 changes: 1 addition & 1 deletion sota-implementations/dqn/dqn_cartpole.py
Original file line number Diff line number Diff line change
Expand Up @@ -94,7 +94,7 @@ def _resume_checkpoint(
)


@hydra.main(config_path="", config_name="config_cartpole", version_base="1.1")
@hydra.main(config_path="", config_name="config_cartpole", version_base="1.3")
def main(cfg: DictConfig):

device = torch.device(cfg.device) if cfg.device else get_available_device()
Expand Down
4 changes: 4 additions & 0 deletions sota-implementations/dqn_trainer/config/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -122,3 +122,7 @@ trainer:
eps_end: 0.05
annealing_num_steps: 250_000
async_collection: false

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion sota-implementations/dqn_trainer/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
from torchrl.trainers.algorithms.configs import * # noqa: F401, F403


@hydra.main(config_path="config", config_name="config", version_base="1.1")
@hydra.main(config_path="config", config_name="config", version_base="1.3")
def main(cfg):
trainer = hydra.utils.instantiate(cfg.trainer)
trainer.train()
Expand Down
4 changes: 4 additions & 0 deletions sota-implementations/dreamer/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -129,3 +129,7 @@ profiling:
trace_file: collector_trace_{worker_idx}.json
# Override init_random_frames when collector profiling is enabled (0 = skip random frames phase)
init_random_frames_override: 0

hydra:
job:
chdir: true
4 changes: 4 additions & 0 deletions sota-implementations/dreamer/config_isaac.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -84,3 +84,7 @@ logger:
# Enable async eval with viewport rendering (requires >= 3 GPUs).
video: True
video_skip: 1

hydra:
job:
chdir: true
2 changes: 1 addition & 1 deletion sota-implementations/dreamer/dreamer.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,7 @@
from torchrl.record.loggers import generate_exp_name, get_logger


@hydra.main(version_base="1.1", config_path="", config_name="config")
@hydra.main(version_base="1.3", config_path="", config_name="config")
def main(cfg: DictConfig): # noqa: F821
# cfg = correct_for_frame_skip(cfg)

Expand Down
Loading
Loading