mirror of
https://github.com/jcreek/CosmicClash.git
synced 2026-09-15 10:22:38 +00:00
feat(*): Add --n-steps and --batch-size flags to train.py, applied on resume as well
This commit is contained in:
+16
-4
@@ -40,6 +40,8 @@ def parse_args():
|
|||||||
parser.add_argument("--seed", type=int, default=0)
|
parser.add_argument("--seed", type=int, default=0)
|
||||||
parser.add_argument("--resume", default=None, help="Checkpoint .zip to resume from")
|
parser.add_argument("--resume", default=None, help="Checkpoint .zip to resume from")
|
||||||
parser.add_argument("--ent-coef", type=float, default=0.0001, help="Entropy bonus coefficient (applied on resume too)")
|
parser.add_argument("--ent-coef", type=float, default=0.0001, help="Entropy bonus coefficient (applied on resume too)")
|
||||||
|
parser.add_argument("--n-steps", type=int, default=256, help="Rollout length per env between updates (applied on resume too)")
|
||||||
|
parser.add_argument("--batch-size", type=int, default=256, help="PPO minibatch size (applied on resume too)")
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--reset-std",
|
"--reset-std",
|
||||||
type=float,
|
type=float,
|
||||||
@@ -74,8 +76,18 @@ def main():
|
|||||||
env = VecMonitor(env)
|
env = VecMonitor(env)
|
||||||
|
|
||||||
if args.resume:
|
if args.resume:
|
||||||
model = PPO.load(args.resume, env=env, tensorboard_log=str(log_dir), ent_coef=args.ent_coef)
|
model = PPO.load(
|
||||||
print(f"Resumed from {args.resume} at {model.num_timesteps} timesteps (ent_coef={args.ent_coef})")
|
args.resume,
|
||||||
|
env=env,
|
||||||
|
tensorboard_log=str(log_dir),
|
||||||
|
ent_coef=args.ent_coef,
|
||||||
|
n_steps=args.n_steps,
|
||||||
|
batch_size=args.batch_size,
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"Resumed from {args.resume} at {model.num_timesteps} timesteps "
|
||||||
|
f"(ent_coef={args.ent_coef}, n_steps={args.n_steps}, batch_size={args.batch_size})"
|
||||||
|
)
|
||||||
if args.reset_std is not None:
|
if args.reset_std is not None:
|
||||||
import math
|
import math
|
||||||
|
|
||||||
@@ -90,8 +102,8 @@ def main():
|
|||||||
env,
|
env,
|
||||||
verbose=1,
|
verbose=1,
|
||||||
ent_coef=args.ent_coef,
|
ent_coef=args.ent_coef,
|
||||||
n_steps=256,
|
n_steps=args.n_steps,
|
||||||
batch_size=256,
|
batch_size=args.batch_size,
|
||||||
learning_rate=3e-4,
|
learning_rate=3e-4,
|
||||||
tensorboard_log=str(log_dir),
|
tensorboard_log=str(log_dir),
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user