pretrain logic and small fixes

This commit is contained in:
2026-08-07 14:32:35 +02:00
parent a1eb7b8573
commit 3e6e40f0c5
8 changed files with 324 additions and 125 deletions
+36 -19
View File
@@ -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(