pretrain logic and small fixes
This commit is contained in:
+36
-19
@@ -57,6 +57,12 @@ def main():
|
||||
parser.add_argument("--save-dir", type=str, default="ml/checkpoints", help="Directory for model checkpoints")
|
||||
parser.add_argument("--save-freq", type=int, default=50_000, help="Checkpoint save frequency (steps)")
|
||||
parser.add_argument("--gui", action="store_true", help="Enable PyBullet 3D visual GUI rendering")
|
||||
parser.add_argument(
|
||||
"--pretrained-model",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to pre-trained base model checkpoint (.zip) to start PPO training from"
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
os.makedirs(args.log_dir, exist_ok=True)
|
||||
@@ -74,25 +80,36 @@ def main():
|
||||
]
|
||||
vec_env = SubprocVecEnv(env_fns)
|
||||
|
||||
# Initialize PPO Policy Hyperparameters
|
||||
model = PPO(
|
||||
policy="MlpPolicy",
|
||||
env=vec_env,
|
||||
learning_rate=1e-4,
|
||||
n_steps=256,
|
||||
batch_size=256,
|
||||
n_epochs=10,
|
||||
gamma=0.99,
|
||||
gae_lambda=0.95,
|
||||
clip_range=0.2,
|
||||
ent_coef=0.03,
|
||||
target_kl=0.05,
|
||||
vf_coef=0.5,
|
||||
max_grad_norm=0.5,
|
||||
verbose=1,
|
||||
tensorboard_log=args.log_dir,
|
||||
device="cpu",
|
||||
)
|
||||
# Initialize or Load PPO Policy Model
|
||||
if args.pretrained_model and os.path.exists(args.pretrained_model):
|
||||
print(f"[Train] Loading pre-trained base knowledge from: {args.pretrained_model}")
|
||||
model = PPO.load(
|
||||
args.pretrained_model,
|
||||
env=vec_env,
|
||||
learning_rate=1e-4, # Lower learning rate so RL fine-tunes without destroying base gait
|
||||
tensorboard_log=args.log_dir,
|
||||
device="cpu",
|
||||
)
|
||||
else:
|
||||
print("[Train] No base model provided. Starting training from scratch...")
|
||||
model = PPO(
|
||||
policy="MlpPolicy",
|
||||
env=vec_env,
|
||||
learning_rate=1e-4,
|
||||
n_steps=256,
|
||||
batch_size=256,
|
||||
n_epochs=10,
|
||||
gamma=0.99,
|
||||
gae_lambda=0.95,
|
||||
clip_range=0.2,
|
||||
ent_coef=0.01,
|
||||
target_kl=0.05,
|
||||
vf_coef=0.5,
|
||||
max_grad_norm=0.5,
|
||||
verbose=1,
|
||||
tensorboard_log=args.log_dir,
|
||||
device="cpu",
|
||||
)
|
||||
|
||||
# Setup Callbacks with ppo<number> naming
|
||||
checkpoint_callback = CheckpointCallback(
|
||||
|
||||
Reference in New Issue
Block a user