mirror of
https://github.com/wassname/curl.git
synced 2026-09-09 11:20:39 +08:00
video and logging
This commit is contained in:
@@ -1,5 +1,7 @@
|
||||
Adapt to Apple gym
|
||||
|
||||
UPTO: env.render returns nothing... make it return last obs. This should have
|
||||
video work
|
||||
|
||||
TODO:
|
||||
- apple
|
||||
|
||||
+23
-14
@@ -4,7 +4,7 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import copy
|
||||
import math
|
||||
|
||||
from pathlib import Path
|
||||
import utils
|
||||
from encoder import make_encoder
|
||||
|
||||
@@ -466,24 +466,33 @@ class CurlSacAgent(object):
|
||||
obs_anchor, obs_pos = cpc_kwargs["obs_anchor"], cpc_kwargs["obs_pos"]
|
||||
self.update_cpc(obs_anchor, obs_pos,cpc_kwargs, L, step)
|
||||
|
||||
def save(self, model_dir, step):
|
||||
torch.save(
|
||||
self.actor.state_dict(), '%s/actor_%s.pt' % (model_dir, step)
|
||||
)
|
||||
torch.save(
|
||||
self.critic.state_dict(), '%s/critic_%s.pt' % (model_dir, step)
|
||||
)
|
||||
# def save(self, model_dir, step):
|
||||
# torch.save(
|
||||
# self.actor.state_dict(), '%s/actor_%s.pt' % (model_dir, step)
|
||||
# )
|
||||
# torch.save(
|
||||
# self.critic.state_dict(), '%s/critic_%s.pt' % (model_dir, step)
|
||||
# )
|
||||
|
||||
def save_curl(self, model_dir, step):
|
||||
torch.save(
|
||||
self.CURL.state_dict(), '%s/curl_%s.pt' % (model_dir, step)
|
||||
)
|
||||
|
||||
def load(self, model_dir, step):
|
||||
self.actor.load_state_dict(
|
||||
torch.load('%s/actor_%s.pt' % (model_dir, step))
|
||||
)
|
||||
self.critic.load_state_dict(
|
||||
torch.load('%s/critic_%s.pt' % (model_dir, step))
|
||||
# def load(self, model_dir, step):
|
||||
# self.actor.load_state_dict(
|
||||
# torch.load('%s/actor_%s.pt' % (model_dir, step))
|
||||
# )
|
||||
# self.critic.load_state_dict(
|
||||
# torch.load('%s/critic_%s.pt' % (model_dir, step))
|
||||
# )
|
||||
|
||||
def load_curl(self, model_dir, step=None):
|
||||
if step is None:
|
||||
steps = sorted([int(f.stem.split('_')[1]) for f in Path(model_dir).glob("curl_*.pt")])
|
||||
step = steps[-1]
|
||||
print(f'auto loading steps step {step}')
|
||||
self.CURL.load_state_dict(
|
||||
torch.load('%s/curl_%s.pt' % (model_dir, step))
|
||||
)
|
||||
|
||||
@@ -1,44 +0,0 @@
|
||||
{
|
||||
"action_repeat": 1,
|
||||
"actor_beta": 0.9,
|
||||
"actor_log_std_max": 2,
|
||||
"actor_log_std_min": -10,
|
||||
"actor_lr": 0.001,
|
||||
"actor_update_freq": 2,
|
||||
"agent": "curl_sac",
|
||||
"alpha_beta": 0.5,
|
||||
"alpha_lr": 0.0001,
|
||||
"batch_size": 32,
|
||||
"critic_beta": 0.9,
|
||||
"critic_lr": 0.001,
|
||||
"critic_target_update_freq": 2,
|
||||
"critic_tau": 0.01,
|
||||
"curl_latent_dim": 128,
|
||||
"detach_encoder": false,
|
||||
"discount": 0.99,
|
||||
"domain_name": "ApplePick-v0",
|
||||
"encoder_feature_dim": 50,
|
||||
"encoder_lr": 0.001,
|
||||
"encoder_tau": 0.05,
|
||||
"encoder_type": "mixed",
|
||||
"eval_freq": 1000,
|
||||
"frame_stack": 3,
|
||||
"hidden_dim": 1024,
|
||||
"image_size": 84,
|
||||
"init_steps": 1000,
|
||||
"init_temperature": 0.1,
|
||||
"log_interval": 100,
|
||||
"num_eval_episodes": 10,
|
||||
"num_filters": 32,
|
||||
"num_layers": 4,
|
||||
"num_train_steps": 1000000,
|
||||
"pre_transform_image_size": 124,
|
||||
"render": false,
|
||||
"replay_buffer_capacity": 10000,
|
||||
"save_buffer": false,
|
||||
"save_model": false,
|
||||
"save_tb": true,
|
||||
"save_video": false,
|
||||
"seed": 1,
|
||||
"work_dir": "./runs/ApplePick-v0-02-18-im84-b32-s1-mixed"
|
||||
}
|
||||
BIN
Binary file not shown.
Executable
+16
@@ -0,0 +1,16 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
ulimit -S -m 55000000
|
||||
ulimit -S -v 55000000
|
||||
python=/home/wassname/anaconda/envs/diygym4/bin/python
|
||||
# DEBUG
|
||||
CUDA_VISIBLE_DEVICES=1 $python \
|
||||
-m pdb -c continue \
|
||||
train.py \
|
||||
--batch_size 64 \
|
||||
--init_steps 256 \
|
||||
--eval_freq 64 \
|
||||
--replay_buffer_capacity 4000 \
|
||||
--render \
|
||||
--load auto \
|
||||
--num_train_steps 128
|
||||
+26
-2
@@ -1,6 +1,30 @@
|
||||
#!/bin/bash
|
||||
set -e
|
||||
CUDA_VISIBLE_DEVICES=1 /home/wassname/anaconda/envs/diygym4/bin/python \
|
||||
ulimit -S -m 55000000
|
||||
ulimit -S -v 55000000
|
||||
python=/home/wassname/anaconda/envs/diygym4/bin/python
|
||||
CUDA_VISIBLE_DEVICES=0
|
||||
# DEBUG
|
||||
# $python \
|
||||
# -m pdb -c continue \
|
||||
# train.py \
|
||||
# --save_tb \
|
||||
# --save_model \
|
||||
# --save_video \
|
||||
# --batch_size 64 \
|
||||
# --init_steps 64 \
|
||||
# --eval_freq 64 \
|
||||
# --replay_buffer_capacity 4000 \
|
||||
# --render
|
||||
|
||||
$python \
|
||||
-m pdb -c continue \
|
||||
train.py \
|
||||
--save_tb --init_steps 10
|
||||
--save_tb \
|
||||
--save_model \
|
||||
--save_video \
|
||||
--batch_size 512 \
|
||||
--init_steps 10000 \
|
||||
--replay_buffer_capacity 40000 \
|
||||
# --load auto
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from typing import DefaultDict
|
||||
import numpy as np
|
||||
import torch
|
||||
import argparse
|
||||
@@ -8,6 +9,8 @@ import sys
|
||||
import random
|
||||
import time
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
|
||||
# import dmc2gym
|
||||
import copy
|
||||
@@ -35,15 +38,15 @@ def parse_args():
|
||||
parser.add_argument("--frame_stack", default=3, type=int)
|
||||
parser.add_argument("--render", action="store_true")
|
||||
# replay buffer
|
||||
parser.add_argument("--replay_buffer_capacity", default=30000, type=int)
|
||||
parser.add_argument("--replay_buffer_capacity", default=50000, type=int)
|
||||
# train
|
||||
parser.add_argument("--agent", default="curl_sac", type=str)
|
||||
parser.add_argument("--init_steps", default=1000, type=int)
|
||||
parser.add_argument("--init_steps", default=10000, type=int)
|
||||
parser.add_argument("--num_train_steps", default=3000000, type=int)
|
||||
parser.add_argument("--batch_size", default=32, type=int)
|
||||
parser.add_argument("--hidden_dim", default=1024, type=int)
|
||||
# eval
|
||||
parser.add_argument("--eval_freq", default=10000, type=int)
|
||||
parser.add_argument("--eval_freq", default=2000, type=int)
|
||||
parser.add_argument("--num_eval_episodes", default=4, type=int)
|
||||
# critic
|
||||
parser.add_argument("--critic_lr", default=1e-3, type=float)
|
||||
@@ -72,18 +75,26 @@ def parse_args():
|
||||
parser.add_argument("--alpha_lr", default=1e-4, type=float)
|
||||
parser.add_argument("--alpha_beta", default=0.5, type=float)
|
||||
# misc
|
||||
parser.add_argument("--seed", default=1, type=int)
|
||||
parser.add_argument("--seed", default=-1, type=int)
|
||||
parser.add_argument("--work_dir", default="./runs", type=str)
|
||||
parser.add_argument("--save_tb", default=False, action="store_true")
|
||||
parser.add_argument("--save_buffer", default=False, action="store_true")
|
||||
parser.add_argument("--save_video", default=False, action="store_true")
|
||||
parser.add_argument("--save_model", default=False, action="store_true")
|
||||
parser.add_argument("--detach_encoder", default=False, action="store_true")
|
||||
parser.add_argument("--load", type=str)
|
||||
|
||||
parser.add_argument("--log_interval", default=100, type=int)
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
keys_to_monitor=[
|
||||
'env_reward/apple_pick/tree/min_fruit_dist_reward',
|
||||
'env_reward/apple_pick/tree/gripping_fruit_reward',
|
||||
# 'env_reward/apple_pick/tree/force_tree_reward',
|
||||
# 'env_reward/apple_pick/tree/force_fruit_reward',
|
||||
'env_obs/apple_pick/tree/picks'
|
||||
]
|
||||
|
||||
def evaluate(env, agent, video, num_episodes, L, step, args):
|
||||
all_ep_rewards = []
|
||||
@@ -91,11 +102,12 @@ def evaluate(env, agent, video, num_episodes, L, step, args):
|
||||
def run_eval_loop(sample_stochastically=True):
|
||||
start_time = time.time()
|
||||
prefix = "stochastic_" if sample_stochastically else ""
|
||||
for i in range(num_episodes):
|
||||
for i in tqdm(range(num_episodes), desc='eval', unit='ep'):
|
||||
obs = env.reset()
|
||||
video.init(enabled=(i == 0))
|
||||
done = False
|
||||
episode_reward = 0
|
||||
episode_info = defaultdict(int)
|
||||
while not done:
|
||||
# center crop image
|
||||
if args.encoder_type == "mixed":
|
||||
@@ -108,18 +120,14 @@ def evaluate(env, agent, video, num_episodes, L, step, args):
|
||||
else:
|
||||
action = agent.select_action(obs)
|
||||
obs, reward, done, info = env.step(action)
|
||||
# keys_to_monitor=[
|
||||
# 'env_reward/apple_pick/tree/min_fruit_dist_reward',
|
||||
# 'env_reward/apple_pick/tree/gripping_fruit_reward',
|
||||
# # 'env_reward/apple_pick/tree/force_tree_reward',
|
||||
# # 'env_reward/apple_pick/tree/force_fruit_reward',
|
||||
# 'env_obs/apple_pick/tree/picks'
|
||||
# ]
|
||||
# for k in keys_to_monitor:
|
||||
# L.log("eval/" + prefix + "episode_reward", episode_reward, step)
|
||||
video.record(env)
|
||||
|
||||
for k in keys_to_monitor:
|
||||
episode_info[k] += info[k]
|
||||
video.record(env, yaw=i)
|
||||
episode_reward += reward
|
||||
|
||||
for k in keys_to_monitor:
|
||||
L.log("eval/" + prefix + k, np.sum(episode_info[k]), step)
|
||||
video.save("%d.mp4" % step)
|
||||
L.log("eval/" + prefix + "episode_reward", episode_reward, step)
|
||||
all_ep_rewards.append(episode_reward)
|
||||
@@ -169,12 +177,19 @@ def make_agent(obs_shape, action_shape, args, device):
|
||||
|
||||
|
||||
def main():
|
||||
import logging
|
||||
from rich.logging import RichHandler
|
||||
logging.basicConfig(level=logging.INFO, handlers=[RichHandler(rich_tracebacks=True, markup=True)])
|
||||
|
||||
args = parse_args()
|
||||
if args.seed == -1:
|
||||
args.__dict__["seed"] = np.random.randint(1, 1000000)
|
||||
print('seed', args.__dict__["seed"])
|
||||
print(args)
|
||||
utils.set_seed_everywhere(args.seed)
|
||||
|
||||
env = gym.make(args.domain_name, render=args.render)
|
||||
print(env)
|
||||
# TODO action repeat wrapper?
|
||||
|
||||
env.seed(args.seed)
|
||||
@@ -189,6 +204,13 @@ def main():
|
||||
keys=["img"],
|
||||
)
|
||||
|
||||
if args.load =='auto':
|
||||
load_dirs = Path(args.work_dir).glob('*/model/curl*.pt')
|
||||
load_dirs = sorted(set([str(d.parent) for d in load_dirs]))
|
||||
print('load_dirs', load_dirs)
|
||||
args.load = str(load_dirs[-1])
|
||||
print('auto load', load_dirs)
|
||||
|
||||
# make directory
|
||||
ts = time.gmtime()
|
||||
ts = time.strftime("%m-%d", ts)
|
||||
@@ -207,6 +229,7 @@ def main():
|
||||
+ args.encoder_type
|
||||
)
|
||||
args.work_dir = args.work_dir + "/" + exp_name
|
||||
print('work_dir', args.work_dir)
|
||||
|
||||
utils.make_dir(args.work_dir)
|
||||
video_dir = utils.make_dir(os.path.join(args.work_dir, "video"))
|
||||
@@ -219,9 +242,10 @@ def main():
|
||||
json.dump(vars(args), f, sort_keys=True, indent=4)
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
print(f'device {device}')
|
||||
|
||||
# shapes
|
||||
action_shape = env.action_space.shape
|
||||
|
||||
img = env.observation_space.sample()["img"]
|
||||
img_aug = utils.center_crop_image(img, args.image_size)
|
||||
obs_shape = {"img": img_aug.shape, "state": env.observation_space["state"].shape}
|
||||
@@ -238,6 +262,8 @@ def main():
|
||||
agent = make_agent(
|
||||
obs_shape=obs_shape, action_shape=action_shape, args=args, device=device
|
||||
)
|
||||
if args.load is not None:
|
||||
agent.load_curl(args.load)
|
||||
|
||||
|
||||
# summarize
|
||||
@@ -259,12 +285,13 @@ def main():
|
||||
L = Logger(args.work_dir, use_tb=args.save_tb)
|
||||
|
||||
episode, episode_reward, done = 0, 0, True
|
||||
episode_info = defaultdict(int)
|
||||
start_time = time.time()
|
||||
|
||||
for step in tqdm(range(args.num_train_steps), desc="train", unit="step", mininterval=360):
|
||||
# evaluate agent periodically
|
||||
|
||||
if (step % args.eval_freq == 0) and (step > args.eval_freq):
|
||||
if (step % args.eval_freq == 0) and (step >= args.eval_freq):
|
||||
L.log("eval/episode", episode, step)
|
||||
evaluate(env, agent, video, args.num_eval_episodes, L, step, args)
|
||||
if args.save_model:
|
||||
@@ -281,12 +308,16 @@ def main():
|
||||
if step % args.log_interval == 0:
|
||||
L.log("train/episode_reward", episode_reward, step)
|
||||
|
||||
for k in keys_to_monitor:
|
||||
L.log("train/episode_info" + k, np.sum(episode_info[k]), step)
|
||||
|
||||
obs = env.reset()
|
||||
assert env.observation_space.contains(
|
||||
obs
|
||||
), f"obs should be in space. ob={obs} space={env.observation_space}"
|
||||
done = False
|
||||
episode_reward = 0
|
||||
episode_info = defaultdict(int)
|
||||
episode_step = 0
|
||||
episode += 1
|
||||
if step % args.log_interval == 0:
|
||||
@@ -302,11 +333,12 @@ def main():
|
||||
action
|
||||
), f"obs should be in space. ob={action} space={env.action_space}"
|
||||
|
||||
# run training update
|
||||
if step >= args.init_steps:
|
||||
num_updates = 1
|
||||
for _ in range(num_updates):
|
||||
agent.update(replay_buffer, L, step)
|
||||
if step % 10 ==0:
|
||||
# run training update
|
||||
if step >= args.init_steps:
|
||||
num_updates = 1
|
||||
for _ in range(num_updates):
|
||||
agent.update(replay_buffer, L, step)
|
||||
|
||||
next_obs, reward, done, info = env.step(action)
|
||||
|
||||
@@ -314,6 +346,8 @@ def main():
|
||||
done_bool = 0 if episode_step + 1 == env._max_episode_steps else float(done)
|
||||
episode_reward += reward
|
||||
replay_buffer.add(obs, action, reward, next_obs, done_bool)
|
||||
for k in keys_to_monitor:
|
||||
episode_info[k] += info[k]
|
||||
|
||||
obs = next_obs
|
||||
episode_step += 1
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import imageio
|
||||
import os
|
||||
import numpy as np
|
||||
|
||||
import logging
|
||||
|
||||
class VideoRecorder(object):
|
||||
def __init__(self, dir_name, height=256, width=256, camera_id=0, fps=30):
|
||||
@@ -16,14 +16,14 @@ class VideoRecorder(object):
|
||||
self.frames = []
|
||||
self.enabled = self.dir_name is not None and enabled
|
||||
|
||||
def record(self, env):
|
||||
def record(self, env, yaw=0):
|
||||
if self.enabled:
|
||||
try:
|
||||
frame = env.render(
|
||||
mode='rgb_array',
|
||||
height=self.height,
|
||||
width=self.width,
|
||||
camera_id=self.camera_id
|
||||
yaw=yaw
|
||||
)
|
||||
except:
|
||||
frame = env.render(
|
||||
@@ -36,3 +36,5 @@ class VideoRecorder(object):
|
||||
if self.enabled:
|
||||
path = os.path.join(self.dir_name, file_name)
|
||||
imageio.mimsave(path, self.frames, fps=self.fps)
|
||||
|
||||
logging.info(f"saved {len(self.frames)} frames to video: '{path}'")
|
||||
|
||||
Reference in New Issue
Block a user