fix(hpc): overriding args
This commit is contained in:
parent
d33dc32506
commit
24ae7a2399
2 changed files with 22 additions and 5 deletions
|
|
@ -5,7 +5,7 @@ seed: 0
|
||||||
track: true # Test WandB integration
|
track: true # Test WandB integration
|
||||||
capture_video: false # No rendering for smoke test
|
capture_video: false # No rendering for smoke test
|
||||||
save_model: true # Test the end-of-training save routine
|
save_model: true # Test the end-of-training save routine
|
||||||
num_envs: 4
|
num_envs: 512
|
||||||
total_timesteps: 500
|
total_timesteps: 65536
|
||||||
num_steps: 64
|
num_steps: 128
|
||||||
cuda: true
|
cuda: true
|
||||||
|
|
|
||||||
21
src/train.py
21
src/train.py
|
|
@ -1,4 +1,6 @@
|
||||||
|
import datetime
|
||||||
import random
|
import random
|
||||||
|
import yaml
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
import time
|
import time
|
||||||
|
|
@ -332,7 +334,7 @@ def train(args: PPOArgs):
|
||||||
sps = int(global_step / (time.time() - start_time))
|
sps = int(global_step / (time.time() - start_time))
|
||||||
remaining_steps = args.total_timesteps - global_step
|
remaining_steps = args.total_timesteps - global_step
|
||||||
eta_seconds = int(remaining_steps / sps) if sps > 0 else 0
|
eta_seconds = int(remaining_steps / sps) if sps > 0 else 0
|
||||||
eta_str = str(time.timedelta(seconds=eta_seconds))
|
eta_str = str(datetime.timedelta(seconds=eta_seconds))
|
||||||
|
|
||||||
print(
|
print(
|
||||||
f"Iteration {iteration}/{args.num_iterations} | "
|
f"Iteration {iteration}/{args.num_iterations} | "
|
||||||
|
|
@ -371,7 +373,22 @@ def train(args: PPOArgs):
|
||||||
|
|
||||||
|
|
||||||
def main() -> None:
|
def main() -> None:
|
||||||
args = tyro.cli(PPOArgs)
|
temp_args = tyro.cli(PPOArgs)
|
||||||
|
|
||||||
|
if temp_args.env_config_path is not None:
|
||||||
|
with open(temp_args.env_config_path, "r") as f:
|
||||||
|
config = yaml.safe_load(f)
|
||||||
|
if config:
|
||||||
|
# parse PPOArgs with defaults from yaml.
|
||||||
|
for key, value in config.items():
|
||||||
|
if hasattr(temp_args, key):
|
||||||
|
setattr(temp_args, key, value)
|
||||||
|
|
||||||
|
# Re-parse CLI to ensure they OVERRIDE the yaml
|
||||||
|
args = tyro.cli(PPOArgs, default=temp_args)
|
||||||
|
else:
|
||||||
|
args = temp_args
|
||||||
|
|
||||||
train(args)
|
train(args)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
Reference in a new issue