Skip to content
Draft
Show file tree
Hide file tree
Changes from 1 commit
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",
"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",
"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