mirror of
https://github.com/wassname/kair_algorithms_draft.git
synced 2026-09-09 11:25:10 +08:00
Merge branch 'master' into feat/demo_refactoring
This commit is contained in:
@@ -8,6 +8,7 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.optim as optim
|
||||
from config.agent.lunarlander_continuous_v2.utils import LunarLanderContinuousHER
|
||||
|
||||
from algorithms.common.networks.mlp import MLP, FlattenMLP, TanhGaussianDistParams
|
||||
from algorithms.sac.agent import Agent
|
||||
@@ -38,6 +39,10 @@ hyper_params = {
|
||||
"VF_HIDDEN_SIZES": [256, 256],
|
||||
"QF_HIDDEN_SIZES": [256, 256],
|
||||
},
|
||||
# HER
|
||||
"USE_HER": True,
|
||||
"SUCCESS_SCORE": 250.0,
|
||||
"DESIRED_STATES_FROM_DEMO": True,
|
||||
}
|
||||
|
||||
|
||||
@@ -52,6 +57,9 @@ def get(env, args):
|
||||
state_dim = env.observation_space.shape[0]
|
||||
action_dim = env.action_space.shape[0]
|
||||
|
||||
if hyper_params["USE_HER"]:
|
||||
state_dim *= 2
|
||||
|
||||
hidden_sizes_actor = hyper_params["NETWORK"]["ACTOR_HIDDEN_SIZES"]
|
||||
hidden_sizes_vf = hyper_params["NETWORK"]["VF_HIDDEN_SIZES"]
|
||||
hidden_sizes_qf = hyper_params["NETWORK"]["QF_HIDDEN_SIZES"]
|
||||
@@ -107,5 +115,8 @@ def get(env, args):
|
||||
models = (actor, vf, vf_target, qf_1, qf_2)
|
||||
optims = (actor_optim, vf_optim, qf_1_optim, qf_2_optim)
|
||||
|
||||
# HER
|
||||
her = LunarLanderContinuousHER() if hyper_params["USE_HER"] else None
|
||||
|
||||
# create an agent
|
||||
return Agent(env, args, hyper_params, models, optims, target_entropy)
|
||||
return Agent(env, args, hyper_params, models, optims, target_entropy, her)
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""Utils for examples on LunarLanderContinuous-v2.
|
||||
- Author: Kh Kim
|
||||
- Contact: kh.kim@medipixel.io
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
|
||||
from algorithms.common.abstract.her import HER
|
||||
from algorithms.common.abstract.reward_fn import RewardFn
|
||||
|
||||
|
||||
class L1DistanceRewardFn(RewardFn):
|
||||
def __call__(self, transition, goal_state):
|
||||
"""L1 Distance reward function."""
|
||||
next_state = transition[3]
|
||||
eps = 1e-6
|
||||
if np.abs(next_state - goal_state).sum() < eps:
|
||||
return np.float64(0.0)
|
||||
else:
|
||||
return np.float64(-1.0)
|
||||
|
||||
|
||||
class LunarLanderContinuousHER(HER):
|
||||
"""HER for LunarLanderContinuous-v2 environment.
|
||||
Attributes:
|
||||
demo_goal_indices (np.ndarray): indices about goal of demo list
|
||||
desired_states (np.ndarray): desired states from demonstration
|
||||
"""
|
||||
|
||||
def __init__(self, reward_func=L1DistanceRewardFn):
|
||||
"""Initialization."""
|
||||
HER.__init__(self, reward_func=reward_func)
|
||||
|
||||
# pylint: disable=attribute-defined-outside-init
|
||||
def fetch_desired_states_from_demo(self, demo):
|
||||
"""Return desired goal states from demonstration data."""
|
||||
np_demo = np.array(demo)
|
||||
self.demo_goal_indices = np.where(np_demo[:, 4])[0]
|
||||
self.desired_states = np_demo[self.demo_goal_indices][:, 0]
|
||||
|
||||
def get_desired_state(self, *args):
|
||||
"""Sample one of the desired states."""
|
||||
return np.random.choice(self.desired_states, 1).item()
|
||||
|
||||
def _get_final_state(self, transition):
|
||||
"""Get final state from transitions for making HER transitions."""
|
||||
return transition[0]
|
||||
|
||||
def generate_demo_transitions(self, demo):
|
||||
"""Return generated demo transitions for HER."""
|
||||
new_demo = list()
|
||||
|
||||
# generate demo transitions
|
||||
prev_idx = 0
|
||||
for idx in self.demo_goal_indices:
|
||||
demo_final_state = self._get_final_state(demo[idx])
|
||||
transitions = [demo[i] for i in range(prev_idx, idx + 1)]
|
||||
prev_idx = idx + 1
|
||||
|
||||
transitions = self.generate_transitions(
|
||||
transitions, demo_final_state, 0, is_demo=True
|
||||
)
|
||||
|
||||
new_demo.extend(transitions)
|
||||
|
||||
return new_demo
|
||||
@@ -2,7 +2,6 @@ from math import pi
|
||||
|
||||
from geometry_msgs.msg import Quaternion
|
||||
|
||||
|
||||
config = {
|
||||
"ENV_NAME": "OpenManipulatorReacher",
|
||||
"MAX_EPISODE_STEPS": 100,
|
||||
|
||||
Reference in New Issue
Block a user