Merge pull request #4 from sisl/options

Integrate options env and policy
This commit is contained in:
Arec Jamgochian
2022-02-20 20:20:28 -08:00
committed by GitHub
94 changed files with 1513 additions and 311 deletions

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

View File

@@ -13,4 +13,20 @@ python -m src.eval_main
# idm # idm
python -m src.eval_main --method=idm python -m src.eval_main --method=idm
# behavior cloning
python -m src.eval_main --method=bc --policy_file='checkpoints/bc-intersimple-setobs2.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}'
# GAIL
python -m src.eval_main --method=gail --policy_file='checkpoints/gail-intersimple-setobs2-03-02-22.pt' --env='NormalizedContinuousEvalEnv' --env_kwargs='{stop_on_collision:True}'
# options GAIL
python -m src.eval_main --method=ogail --policy_file='checkpoints/gail-options-setobs2-Feb15_18-49-05.pt' --env='NormalizedOptionsEvalEnv' --env_kwargs='{stop_on_collision:True}'
# options GAIL-PPO
python -m src.eval_main --method=ogail-ppo --policy_file='checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt' --env='NormalizedOptionsEvalEnv' --env_kwargs='{stop_on_collision:True}'
# SHAIL
python -m src.eval_main --method=sgail --policy_file='checkpoints/sgail-options-setobs2.pt' --env='NormalizedSafeOptionsEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}'
# SHAIL-PPO
python -m src.eval_main --method=sgail-ppo --policy_file='checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt' --env='NormalizedSafeOptionsEvalEnv' --env_kwargs='{stop_on_collision:True,max_episode_steps:1000}'

View File

@@ -1,232 +0,0 @@
PyTorch-Reparam-Module
cg.ipynb
vec-env.ipynb
*.zip
*.pt
*.mp4
*.pkl
runs/
# Created by https://www.toptal.com/developers/gitignore/api/linux,macos,python,visualstudiocode
# Edit at https://www.toptal.com/developers/gitignore?templates=linux,macos,python,visualstudiocode
### Linux ###
*~
# temporary files which can be created if a process still has a handle open of a deleted file
.fuse_hidden*
# KDE directory preferences
.directory
# Linux trash folder which might appear on any partition or disk
.Trash-*
# .nfs files are created when an open file is removed but is still being accessed
.nfs*
### macOS ###
# General
.DS_Store
.AppleDouble
.LSOverride
# Icon must end with two \r
Icon
# Thumbnails
._*
# Files that might appear in the root of a volume
.DocumentRevisions-V100
.fseventsd
.Spotlight-V100
.TemporaryItems
.Trashes
.VolumeIcon.icns
.com.apple.timemachine.donotpresent
# Directories potentially created on remote AFP share
.AppleDB
.AppleDesktop
Network Trash Folder
Temporary Items
.apdisk
### Python ###
# Byte-compiled / optimized / DLL files
__pycache__/
*.py[cod]
*$py.class
# C extensions
*.so
# Distribution / packaging
.Python
build/
develop-eggs/
dist/
downloads/
eggs/
.eggs/
lib/
lib64/
parts/
sdist/
var/
wheels/
share/python-wheels/
*.egg-info/
.installed.cfg
*.egg
MANIFEST
# PyInstaller
# Usually these files are written by a python script from a template
# before PyInstaller builds the exe, so as to inject date/other infos into it.
*.manifest
*.spec
# Installer logs
pip-log.txt
pip-delete-this-directory.txt
# Unit test / coverage reports
htmlcov/
.tox/
.nox/
.coverage
.coverage.*
.cache
nosetests.xml
coverage.xml
*.cover
*.py,cover
.hypothesis/
.pytest_cache/
cover/
# Translations
*.mo
*.pot
# Django stuff:
*.log
local_settings.py
db.sqlite3
db.sqlite3-journal
# Flask stuff:
instance/
.webassets-cache
# Scrapy stuff:
.scrapy
# Sphinx documentation
docs/_build/
# PyBuilder
.pybuilder/
target/
# Jupyter Notebook
.ipynb_checkpoints
# IPython
profile_default/
ipython_config.py
# pyenv
# For a library or package, you might want to ignore these files since the code is
# intended to run in multiple environments; otherwise, check them in:
# .python-version
# pipenv
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
# However, in case of collaboration, if having platform-specific dependencies or dependencies
# having no cross-platform support, pipenv may install dependencies that don't work, or not
# install all needed dependencies.
#Pipfile.lock
# poetry
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
# This is especially recommended for binary packages to ensure reproducibility, and is more
# commonly ignored for libraries.
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
#poetry.lock
# PEP 582; used by e.g. github.com/David-OConnor/pyflow
__pypackages__/
# Celery stuff
celerybeat-schedule
celerybeat.pid
# SageMath parsed files
*.sage.py
# Environments
.env
.venv
env/
venv/
ENV/
env.bak/
venv.bak/
# Spyder project settings
.spyderproject
.spyproject
# Rope project settings
.ropeproject
# mkdocs documentation
/site
# mypy
.mypy_cache/
.dmypy.json
dmypy.json
# Pyre type checker
.pyre/
# pytype static type analyzer
.pytype/
# Cython debug symbols
cython_debug/
# PyCharm
# JetBrains specific template is maintainted in a separate JetBrains.gitignore that can
# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
# and can be added to the global gitignore or merged into this file. For a more nuclear
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
#.idea/
### VisualStudioCode ###
.vscode/*
!.vscode/settings.json
!.vscode/tasks.json
!.vscode/launch.json
!.vscode/extensions.json
!.vscode/*.code-snippets
# Local History for Visual Studio Code
.history/
# Built Visual Studio Code Extensions
*.vsix
### VisualStudioCode Patch ###
# Ignore all local history of files
.history
.ionide
# Support for Project snippet scope
# End of https://www.toptal.com/developers/gitignore/api/linux,macos,python,visualstudiocode

View File

@@ -26,7 +26,7 @@ torch.save(policy.state_dict(), 'bc-intersimple-setobs2.pt')
# %% # %%
import numpy as np import numpy as np
from core.policy import SetPolicy from core.policy import SetPolicy
from wrappers import Setobs, TransformObservation, CollisionPenaltyWrapper from util.wrappers import Setobs, TransformObservation, CollisionPenaltyWrapper
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat(
n_rays=5, n_rays=5,

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -9,7 +9,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter
@@ -56,6 +56,11 @@ expert_data = torch.load('intersimple-expert-data-setobs2.pt')
expert_data = Buffer(*expert_data) expert_data = Buffer(*expert_data)
# %% # %%
def callback(epoch, value, policy):
if not epoch % 10:
torch.save(policy.state_dict(), f'gail-options-setobs2-{epoch}.pt')
torch.save(value.state_dict(), f'gail-options-setobs2-value-{epoch}.pt')
value, policy = gail( value, policy = gail(
env_fn=env_fn, env_fn=env_fn,
expert_data=expert_data, expert_data=expert_data,
@@ -66,7 +71,7 @@ value, policy = gail(
value=value, value=value,
v_opt=v_opt, v_opt=v_opt,
v_iters=1000, v_iters=1000,
epochs=200, epochs=300,
rollout_episodes=60, rollout_episodes=60,
rollout_steps=60, rollout_steps=60,
gamma=0.99, gamma=0.99,
@@ -75,6 +80,7 @@ value, policy = gail(
backtrack_coeff=0.8, backtrack_coeff=0.8,
backtrack_iters=10, backtrack_iters=10,
logger=SummaryWriter(comment='gail-options-setobs2'), logger=SummaryWriter(comment='gail-options-setobs2'),
callback=callback,
) )
torch.save(policy.state_dict(), 'gail-options-setobs2.pt') torch.save(policy.state_dict(), 'gail-options-setobs2.pt')

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat(
n_rays=5, n_rays=5,

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -8,7 +8,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter
@@ -57,6 +57,11 @@ expert_data = torch.load('intersimple-expert-data-setobs2.pt')
expert_data = Buffer(*expert_data) expert_data = Buffer(*expert_data)
# %% # %%
def callback(epoch, value, policy):
if not epoch % 10:
torch.save(policy.state_dict(), f'gail-ppo-options-setobs2-{epoch}.pt')
torch.save(value.state_dict(), f'gail-ppo-options-setobs2-value-{epoch}.pt')
value, policy = gail_ppo( value, policy = gail_ppo(
env_fn=env_fn, env_fn=env_fn,
expert_data=expert_data, expert_data=expert_data,
@@ -76,6 +81,7 @@ value, policy = gail_ppo(
pi_opt=pi_opt, pi_opt=pi_opt,
pi_iters=100, pi_iters=100,
logger=SummaryWriter(comment='gail-ppo-options-setobs2'), logger=SummaryWriter(comment='gail-ppo-options-setobs2'),
callback=callback,
) )
torch.save(policy.state_dict(), 'gail-ppo-options-setobs2.pt') torch.save(policy.state_dict(), 'gail-ppo-options-setobs2.pt')

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
from intersim.expert import NormalizedIntersimpleExpert from intersim.expert import NormalizedIntersimpleExpert
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
env = CollisionPenaltyWrapper(IntersimpleLidarFlat( env = CollisionPenaltyWrapper(IntersimpleLidarFlat(
n_rays=5, n_rays=5,

View File

@@ -9,7 +9,7 @@ import torch.optim
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import Minobs from util.wrappers import Minobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -9,7 +9,7 @@ import torch.optim
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import Minobs from util.wrappers import Minobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -9,9 +9,9 @@ from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np import numpy as np
from wrappers import CollisionPenaltyWrapper, TransformObservation from util.wrappers import CollisionPenaltyWrapper, TransformObservation
from wrappers import Minobs from util.wrappers import Minobs
from options.options import OptionsEnv from options.options import OptionsEnv
obs_min = np.array([ obs_min = np.array([

View File

@@ -0,0 +1,3 @@
torch
stable-baselines3
gym

View File

@@ -0,0 +1,111 @@
# %%
import sys
sys.path.append('../../../../')
import gym
from src.safe_options.options import gail
from src.core.gail import Buffer
from src.core.value import SetValue
from src.safe_options.policy import SetMaskedDiscretePolicy
from src.core.discriminator import DeepsetDiscriminator
import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward
import functools
from src.util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np
from src.safe_options.options import SafeOptionsEnv
from torch.utils.tensorboard import SummaryWriter
from src.core.reparam_module import ReparamPolicy
obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
]).reshape(-1)
obs_max = np.array([
[1000, 1000, 20, np.pi, 1e-1, 0.],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
]).reshape(-1)
envs = [SafeOptionsEnv(Setobs(
TransformObservation(CollisionPenaltyWrapper(IntersimpleLidarFlatRandom(
n_rays=5,
reward=functools.partial(
speed_reward,
collision_penalty=0
),
stop_on_collision=True,
), collision_distance=6, collision_penalty=100), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method='circle', abort_unsafe_collision_method='circle') for _ in range(60)]
env_fn = lambda i: envs[i]
policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n)
value = SetValue()
v_opt = torch.optim.Adam(value.parameters(), lr=1e-4)
discriminator = DeepsetDiscriminator()
disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4)
expert_data = torch.load('intersimple-expert-data-setobs2.pt')
expert_data = Buffer(*expert_data)
# %%
def callback(epoch, value, policy):
if not epoch % 10:
torch.save(policy.state_dict(), f'sgail-options-setobs2-{epoch}.pt')
torch.save(value.state_dict(), f'sgail-options-setobs2-value-{epoch}.pt')
value, policy = gail(
env_fn=env_fn,
expert_data=expert_data,
discriminator=discriminator,
disc_opt=disc_opt,
disc_iters=100,
policy=policy,
value=value,
v_opt=v_opt,
v_iters=1000,
epochs=300,
rollout_episodes=60,
rollout_steps=60,
gamma=0.99,
gae_lambda=0.9,
delta=0.01,
backtrack_coeff=0.8,
backtrack_iters=10,
logger=SummaryWriter(comment='sgail-options-setobs2'),
callback=callback,
)
torch.save(policy.state_dict(), 'sgail-options-setobs2.pt')
# %%
policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n)
policy(torch.zeros(env_fn(0).observation_space['observation'].shape), torch.zeros(env_fn(0).observation_space['safe_actions'].shape))
policy = ReparamPolicy(policy)
policy.load_state_dict(torch.load('sgail-options-setobs2.pt'))
env = env_fn(0)
obs = env.reset()
env.render(mode='post')
for i in range(300):
action = policy.sample(policy(
torch.tensor(obs['observation'], dtype=torch.float32),
torch.tensor(obs['safe_actions'], dtype=torch.float32),
))
obs, reward, done, _ = env.step(action, render_mode='post')
print('step', i, 'reward', reward)
if done:
break
env.close()
# %%

View File

@@ -0,0 +1,110 @@
# %%
import sys
sys.path.append('../../../../')
import gym
from src.safe_options.options import gail_ppo, Buffer
from src.core.value import SetValue
from src.safe_options.policy import SetMaskedDiscretePolicy
from src.core.discriminator import DeepsetDiscriminator
import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward
import functools
from src.util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np
from src.safe_options.options import SafeOptionsEnv
from torch.utils.tensorboard import SummaryWriter
obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
]).reshape(-1)
obs_max = np.array([
[1000, 1000, 20, np.pi, 1e-1, 0.],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
]).reshape(-1)
envs = [SafeOptionsEnv(Setobs(
TransformObservation(CollisionPenaltyWrapper(IntersimpleLidarFlatRandom(
n_rays=5,
reward=functools.partial(
speed_reward,
collision_penalty=0
),
stop_on_collision=True,
), collision_distance=6, collision_penalty=100), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method='circle', abort_unsafe_collision_method='circle') for _ in range(60)]
env_fn = lambda i: envs[i]
policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n)
pi_opt = torch.optim.Adam(policy.parameters(), lr=3e-4)
value = SetValue()
v_opt = torch.optim.Adam(value.parameters(), lr=1e-3)
discriminator = DeepsetDiscriminator()
disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4)
expert_data = torch.load('intersimple-expert-data-setobs2.pt')
expert_data = Buffer(*expert_data)
# %%
def callback(epoch, value, policy):
if not epoch % 10:
torch.save(policy.state_dict(), f'sgail-ppo-options-setobs2-{epoch}.pt')
torch.save(value.state_dict(), f'sgail-ppo-options-setobs2-value-{epoch}.pt')
value, policy = gail_ppo(
env_fn=env_fn,
expert_data=expert_data,
discriminator=discriminator,
disc_opt=disc_opt,
disc_iters=100,
policy=policy,
value=value,
v_opt=v_opt,
v_iters=1000,
epochs=200,
rollout_episodes=60,
rollout_steps=60,
gamma=0.99,
gae_lambda=0.9,
clip_ratio=0.2,
pi_opt=pi_opt,
pi_iters=100,
logger=SummaryWriter(comment='sgail-ppo-options-setobs2'),
callback=callback,
)
torch.save(policy.state_dict(), 'sgail-ppo-options-setobs2.pt')
# %%
policy = SetMaskedDiscretePolicy(env_fn(0).action_space.n)
policy(torch.zeros(env_fn(0).observation_space['observation'].shape), torch.zeros(env_fn(0).observation_space['safe_actions'].shape))
policy.load_state_dict(torch.load('sgail-ppo-options-setobs2.pt'))
env = env_fn(0)
obs = env.reset()
env.render(mode='post')
for i in range(300):
action = policy.sample(policy(
torch.tensor(obs['observation'], dtype=torch.float32),
torch.tensor(obs['safe_actions'], dtype=torch.float32),
))
obs, reward, done, _ = env.step(action, render_mode='post')
print('step', i, 'reward', reward, 'safe actions', obs['safe_actions'])
if done:
break
env.close()
# %%

View File

@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy
from wrappers import Minobs from util.wrappers import Minobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy
from wrappers import Minobs from util.wrappers import Minobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy
from wrappers import Setobs from util.wrappers import Setobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy
from wrappers import Setobs from util.wrappers import Setobs
obs_min = np.array([ obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.], [-1000, -1000, 0, -np.pi, -1e-1, 0.],

View File

@@ -9,10 +9,10 @@ from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
import numpy as np import numpy as np
from wrappers import CollisionPenaltyWrapper, TransformObservation from util.wrappers import CollisionPenaltyWrapper, TransformObservation
from core.reparam_module import ReparamPolicy from core.reparam_module import ReparamPolicy
from wrappers import Minobs from util.wrappers import Minobs
from options.options import OptionsEnv from options.options import OptionsEnv
obs_min = np.array([ obs_min = np.array([

View File

@@ -0,0 +1,346 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [],
"source": [
"from stable_baselines3.common.env_util import make_vec_env\n",
"import numpy as np"
]
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {},
"outputs": [],
"source": [
"env = make_vec_env('Pendulum-v0', n_envs=6)"
]
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"(6, 3)"
]
},
"execution_count": 5,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"obs = env.reset()\n",
"obs.shape"
]
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"1\n",
"2\n",
"3\n",
"4\n",
"5\n",
"6\n",
"7\n",
"8\n",
"9\n",
"10\n",
"11\n",
"12\n",
"13\n",
"14\n",
"15\n",
"16\n",
"17\n",
"18\n",
"19\n",
"20\n",
"21\n",
"22\n",
"23\n",
"24\n",
"25\n",
"26\n",
"27\n",
"28\n",
"29\n",
"30\n",
"31\n",
"32\n",
"33\n",
"34\n",
"35\n",
"36\n",
"37\n",
"38\n",
"39\n",
"40\n",
"41\n",
"42\n",
"43\n",
"44\n",
"45\n",
"46\n",
"47\n",
"48\n",
"49\n",
"50\n",
"51\n",
"52\n",
"53\n",
"54\n",
"55\n",
"56\n",
"57\n",
"58\n",
"59\n",
"60\n",
"61\n",
"62\n",
"63\n",
"64\n",
"65\n",
"66\n",
"67\n",
"68\n",
"69\n",
"70\n",
"71\n",
"72\n",
"73\n",
"74\n",
"75\n",
"76\n",
"77\n",
"78\n",
"79\n",
"80\n",
"81\n",
"82\n",
"83\n",
"84\n",
"85\n",
"86\n",
"87\n",
"88\n",
"89\n",
"90\n",
"91\n",
"92\n",
"93\n",
"94\n",
"95\n",
"96\n",
"97\n",
"98\n",
"99\n",
"100\n",
"101\n",
"102\n",
"103\n",
"104\n",
"105\n",
"106\n",
"107\n",
"108\n",
"109\n",
"110\n",
"111\n",
"112\n",
"113\n",
"114\n",
"115\n",
"116\n",
"117\n",
"118\n",
"119\n",
"120\n",
"121\n",
"122\n",
"123\n",
"124\n",
"125\n",
"126\n",
"127\n",
"128\n",
"129\n",
"130\n",
"131\n",
"132\n",
"133\n",
"134\n",
"135\n",
"136\n",
"137\n",
"138\n",
"139\n",
"140\n",
"141\n",
"142\n",
"143\n",
"144\n",
"145\n",
"146\n",
"147\n",
"148\n",
"149\n",
"150\n",
"151\n",
"152\n",
"153\n",
"154\n",
"155\n",
"156\n",
"157\n",
"158\n",
"159\n",
"160\n",
"161\n",
"162\n",
"163\n",
"164\n",
"165\n",
"166\n",
"167\n",
"168\n",
"169\n",
"170\n",
"171\n",
"172\n",
"173\n",
"174\n",
"175\n",
"176\n",
"177\n",
"178\n",
"179\n",
"180\n",
"181\n",
"182\n",
"183\n",
"184\n",
"185\n",
"186\n",
"187\n",
"188\n",
"189\n",
"190\n",
"191\n",
"192\n",
"193\n",
"194\n",
"195\n",
"196\n",
"197\n",
"198\n",
"199\n",
"200\n"
]
}
],
"source": [
"dones = [False]\n",
"i = 0\n",
"while not any(dones):\n",
" i += 1\n",
" print(i)\n",
" _, _, dones, _ = env.step(np.zeros((6, 1)))"
]
},
{
"cell_type": "code",
"execution_count": 7,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"array([ True, True, True, True, True, True])"
]
},
"execution_count": 7,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"dones"
]
},
{
"cell_type": "code",
"execution_count": 8,
"metadata": {},
"outputs": [],
"source": [
"_, _, dones, _ = env.step(np.zeros((6, 1)))"
]
},
{
"cell_type": "code",
"execution_count": 9,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"array([False, False, False, False, False, False])"
]
},
"execution_count": 9,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"dones"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": []
}
],
"metadata": {
"interpreter": {
"hash": "6c7a4ac80dd345f83235e10baa3acc437d966916e1cc075a45b91bb9cc030938"
},
"kernelspec": {
"display_name": "Python 3.9.7 64-bit ('.venv': venv)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.9.7"
},
"orig_nbformat": 4
},
"nbformat": 4,
"nbformat_minor": 2
}

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat( envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat(
n_rays=5, n_rays=5,

View File

@@ -9,7 +9,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -9,7 +9,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter
@@ -50,7 +50,7 @@ value = SetValue()
v_opt = torch.optim.Adam(value.parameters(), lr=1e-4) v_opt = torch.optim.Adam(value.parameters(), lr=1e-4)
discriminator = DeepsetDiscriminator() discriminator = DeepsetDiscriminator()
disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-3) disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4)
expert_data = torch.load('intersimple-expert-data-setobs2.pt') expert_data = torch.load('intersimple-expert-data-setobs2.pt')
expert_data = Buffer(*expert_data) expert_data = Buffer(*expert_data)

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Minobs from util.wrappers import CollisionPenaltyWrapper, Minobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, Setobs from util.wrappers import CollisionPenaltyWrapper, Setobs
import numpy as np import numpy as np
from gym.wrappers import TransformObservation from gym.wrappers import TransformObservation
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -7,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlat from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter

View File

@@ -1,4 +1,3 @@
# %%
import gym import gym
from options.options import gail_ppo, Buffer from options.options import gail_ppo, Buffer
from core.value import SetValue from core.value import SetValue
@@ -8,7 +7,7 @@ import torch.optim
from intersim.envs import IntersimpleLidarFlatRandom from intersim.envs import IntersimpleLidarFlatRandom
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
import numpy as np import numpy as np
from options.options import OptionsEnv from options.options import OptionsEnv
from torch.utils.tensorboard import SummaryWriter from torch.utils.tensorboard import SummaryWriter
@@ -51,12 +50,11 @@ value = SetValue()
v_opt = torch.optim.Adam(value.parameters(), lr=1e-3) v_opt = torch.optim.Adam(value.parameters(), lr=1e-3)
discriminator = DeepsetDiscriminator() discriminator = DeepsetDiscriminator()
disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-3) disc_opt = torch.optim.Adam(discriminator.parameters(), lr=1e-3, weight_decay=1e-4)
expert_data = torch.load('intersimple-expert-data-setobs2.pt') expert_data = torch.load('intersimple-expert-data-setobs2.pt')
expert_data = Buffer(*expert_data) expert_data = Buffer(*expert_data)
# %%
value, policy = gail_ppo( value, policy = gail_ppo(
env_fn=env_fn, env_fn=env_fn,
expert_data=expert_data, expert_data=expert_data,

View File

@@ -7,7 +7,7 @@ from intersim.envs import IntersimpleLidarFlat
from intersim.envs.intersimple import speed_reward from intersim.envs.intersimple import speed_reward
import functools import functools
import torch import torch
from wrappers import CollisionPenaltyWrapper from util.wrappers import CollisionPenaltyWrapper
model = PPO.load('sb3-ppo-intersimple') model = PPO.load('sb3-ppo-intersimple')
env = CollisionPenaltyWrapper(IntersimpleLidarFlat( env = CollisionPenaltyWrapper(IntersimpleLidarFlat(

View File

@@ -1,10 +1,10 @@
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
from dataclasses import dataclass from dataclasses import dataclass
from core.reparam_module import ReparamPolicy from src.core.reparam_module import ReparamPolicy
from core.sampling import rollout from src.core.sampling import rollout
from core.trpo import trpo_step from src.core.trpo import trpo_step
from core.ppo import ppo_step from src.core.ppo import ppo_step
from tqdm import tqdm from tqdm import tqdm
class TerminalLogger: class TerminalLogger:

View File

@@ -15,8 +15,13 @@ class BasePolicy(nn.Module):
def sample(self, dist): def sample(self, dist):
return self.torch_dist(dist).sample() return self.torch_dist(dist).sample()
def predict(self, states): def predict(self, observations, state=None, episode_start=None, deterministic=True):
return self.sample(self.forward(states)) observations = torch.tensor(observations)
if deterministic:
actions = self.forward(observations)[..., :self.action_dim]
else:
actions = self.sample(self.forward(observations))
return actions, None
def log_prob(self, dist, actions): def log_prob(self, dist, actions):
return self.torch_dist(dist).log_prob(actions) return self.torch_dist(dist).log_prob(actions)
@@ -59,6 +64,14 @@ class DiscretePolicy(BasePolicy):
def torch_dist(self, dist): def torch_dist(self, dist):
return Categorical(logits=dist) return Categorical(logits=dist)
def predict(self, observations, state=None, episode_start=None, deterministic=True):
observations = torch.tensor(observations)
if deterministic:
_, actions = self.forward(observations).max(-1)
else:
actions = self.sample(self.forward(observations))
return actions, None
class SetPolicy(Policy): class SetPolicy(Policy):
def forward(self, states): def forward(self, states):

View File

@@ -1,6 +1,6 @@
import torch import torch
from core.sampling import rollout from src.core.sampling import rollout
from core.value_estimation import gae from src.core.value_estimation import gae
def ppo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, clip_ratio, pi_opt, pi_iters, v_opt, v_iters, target_kl=None, max_grad_norm=None): def ppo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, clip_ratio, pi_opt, pi_iters, v_opt, v_iters, target_kl=None, max_grad_norm=None):

View File

@@ -160,3 +160,6 @@ class ReparamPolicy(ReparamModule):
def predict(self, *args, **kwargs): def predict(self, *args, **kwargs):
return self.module.predict(*args, **kwargs) return self.module.predict(*args, **kwargs)
def unsafe_probability_mass(self, *args, **kwargs):
return self.module.unsafe_probability_mass(*args, **kwargs)

View File

@@ -1,8 +1,8 @@
import torch import torch
from core.reparam_module import ReparamPolicy from src.core.reparam_module import ReparamPolicy
from core.sampling import rollout from src.core.sampling import rollout
from core.value_estimation import gae from src.core.value_estimation import gae
from core.optimization import conjugate_gradient, line_search from src.core.optimization import conjugate_gradient, line_search
def trpo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1): def trpo(env_fn, value, policy, epochs, rollout_episodes, rollout_steps, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1):

View File

@@ -8,6 +8,11 @@ from src.baselines import IDMRulePolicy
from src.evaluation import IntersimpleEvaluation from src.evaluation import IntersimpleEvaluation
import src.gail.options as options_envs import src.gail.options as options_envs
from src.evaluation.metrics import divergence, visualize_distribution from src.evaluation.metrics import divergence, visualize_distribution
from src.core.policy import SetPolicy, SetDiscretePolicy
from src.core.reparam_module import ReparamPolicy
from src.options import envs as options_envs2
from src.safe_options.policy import SetMaskedDiscretePolicy
from src.safe_options import options as options_envs3
from typing import Optional, List, Dict, Tuple from typing import Optional, List, Dict, Tuple
import torch import torch
@@ -33,15 +38,44 @@ def load_policy(method:str,
if method == 'idm': if method == 'idm':
policy = IDMRulePolicy(env, **policy_kwargs) policy = IDMRulePolicy(env, **policy_kwargs)
elif method == 'bc': elif method == 'bc':
raise NotImplementedError policy = SetPolicy(env.action_space.shape[-1])
policy.load_state_dict(torch.load(policy_file))
policy.eval()
elif method == 'gail': elif method == 'gail':
policy = sb3.PPO.load(policy_file) policy = SetPolicy(env.action_space.shape[-1])
raise NotImplementedError policy(torch.zeros(env.observation_space.shape))
policy = ReparamPolicy(policy)
policy.load_state_dict(torch.load(policy_file))
policy.eval()
elif method == 'gail-ppo':
policy = SetPolicy(env.action_space.shape[-1])
policy.load_state_dict(torch.load(policy_file))
policy.eval()
elif method == 'rail': elif method == 'rail':
raise NotImplementedError raise NotImplementedError
elif method == 'ogail':
policy = SetDiscretePolicy(env.action_space.n)
policy(torch.zeros(env.observation_space.shape))
policy = ReparamPolicy(policy)
policy.load_state_dict(torch.load(policy_file))
policy.eval()
elif method == 'ogail-ppo':
policy = SetDiscretePolicy(env.action_space.n)
policy.load_state_dict(torch.load(policy_file))
policy.eval()
elif method == 'sgail': elif method == 'sgail':
policy = sb3.PPO.load(policy_file) policy = SetMaskedDiscretePolicy(env.action_space.n)
raise NotImplementedError policy(
torch.zeros(env.observation_space['observation'].shape),
torch.zeros(env.observation_space['safe_actions'].shape)
)
policy = ReparamPolicy(policy)
policy.load_state_dict(torch.load(policy_file))
policy.eval()
elif method == 'sgail-ppo':
policy = SetMaskedDiscretePolicy(env.action_space.n)
policy.load_state_dict(torch.load(policy_file))
policy.eval()
else: else:
raise NotImplementedError raise NotImplementedError
return policy return policy
@@ -158,6 +192,8 @@ def evaluate_policy(locations:List[Tuple[int,int]],
""" """
envs_dict = dict(intersim.envs.intersimple.__dict__) envs_dict = dict(intersim.envs.intersimple.__dict__)
envs_dict.update(dict(options_envs.__dict__)) envs_dict.update(dict(options_envs.__dict__))
envs_dict.update(dict(options_envs2.__dict__))
envs_dict.update(dict(options_envs3.__dict__))
policy_metrics = [None]* len(locations) policy_metrics = [None]* len(locations)
# iterate through vehicles # iterate through vehicles

View File

@@ -6,6 +6,8 @@ from typing import Callable, Dict, Optional
import os import os
import pickle import pickle
from tqdm import tqdm from tqdm import tqdm
from src.options.envs import OptionsEnv
from src.util.wrappers import OptionsTimeLimit
class IntersimpleEvaluation: class IntersimpleEvaluation:
""" """
@@ -34,6 +36,7 @@ class IntersimpleEvaluation:
self.env = eval_env self.env = eval_env
self.n_episodes = eval_env.nv self.n_episodes = eval_env.nv
self.use_pbar = use_pbar self.use_pbar = use_pbar
self.is_options_env = isinstance(self.env, (OptionsEnv, OptionsTimeLimit))
# metrics present on every step of every episode # metrics present on every step of every episode
self.metric_keys_all = ['v_all', 'a_all', 'col_all'] self.metric_keys_all = ['v_all', 'a_all', 'col_all']
@@ -85,11 +88,14 @@ class IntersimpleEvaluation:
if self.use_pbar: if self.use_pbar:
self.pbar = tqdm(total=self.n_episodes) self.pbar = tqdm(total=self.n_episodes)
if self.is_options_env:
print('Evaluating an options environment')
evaluate_policy( evaluate_policy(
policy, policy,
self.env, self.env,
n_eval_episodes=self.n_episodes, n_eval_episodes=self.n_episodes,
callback=self.evaluate_policy_callback, callback=self.evaluate_options_policy_callback if self.is_options_env else self.evaluate_policy_callback,
return_episode_rewards=False return_episode_rewards=False
) )
if self.use_pbar: if self.use_pbar:
@@ -100,6 +106,13 @@ class IntersimpleEvaluation:
self.save(filestr) self.save(filestr)
return self._metrics return self._metrics
def evaluate_options_policy_callback(self, local_vars, global_vars):
infos = local_vars['info']['ll']['infos']
dones = local_vars['info']['ll']['env_done']
agents = [info['agent'] for info in infos]
for info, done, agent in zip(infos, dones, agents):
self.eval_policy_step(info, done, agent)
def evaluate_policy_callback(self, local_vars, global_vars): def evaluate_policy_callback(self, local_vars, global_vars):
""" """
Callback run in evaluate_policy after taking an action and receiving an observation Callback run in evaluate_policy after taking an action and receiving an observation
@@ -110,8 +123,11 @@ class IntersimpleEvaluation:
done = local_vars['done'] done = local_vars['done']
_agent = info['agent'] _agent = info['agent']
env = local_vars['env'].envs[venv_i] env = local_vars['env'].envs[venv_i]
assert isinstance(env, Intersimple) # assert isinstance(env, Intersimple)
self.eval_policy_step(info, done, _agent)
def eval_policy_step(self, info, done, _agent):
# Increase collision counter if episode terminated with a collision # Increase collision counter if episode terminated with a collision
self._metrics['v_all'][_agent].append(info['prev_state'][_agent,2].item()) self._metrics['v_all'][_agent].append(info['prev_state'][_agent,2].item())
self._metrics['a_all'][_agent].append(info['action_taken'][_agent,0].item()) self._metrics['a_all'][_agent].append(info['action_taken'][_agent,0].item())

111
src/options/envs.py Normal file
View File

@@ -0,0 +1,111 @@
import gym
import numpy as np
from src.util.wrappers import Wrapper, Setobs, TransformObservation
from intersim.envs import IntersimpleLidarFlatIncrementingAgent
obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
]).reshape(-1)
obs_max = np.array([
[1000, 1000, 20, np.pi, 1e-1, 0.],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
]).reshape(-1)
def NormalizedOptionsEvalEnv(**kwargs):
return OptionsEnv(Setobs(
TransformObservation(IntersimpleLidarFlatIncrementingAgent(
n_rays=5,
**kwargs,
), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5)])
def NormalizedContinuousEvalEnv(**kwargs):
return Setobs(
TransformObservation(IntersimpleLidarFlatIncrementingAgent(
n_rays=5,
**kwargs,
), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
)
class OptionsEnv(Wrapper):
def __init__(self, env, options):
super().__init__(env)
self.ll_action_space = env.action_space
self.options = options
self.action_space = gym.spaces.Discrete(len(options))
self.max_plan_length = max(t for _, t in options)
def plan(self, option):
target_v, t = option
current_v = self.env._env.state[self.env._agent, 1].item()
dt = self.env._env._dt
a = (target_v - current_v) / (t * dt)
a = self.env._normalize(a)
a = a * np.ones((t,))
a += 0.01 * np.random.randn(*a.shape)
a = np.clip(a, self.ll_action_space.low, self.ll_action_space.high)
return a
def execute_plan(self, obs, option, render_mode=None):
observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape))
actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape))
rewards = np.zeros((self.max_plan_length + 1,))
env_done = np.ones((self.max_plan_length + 1,), dtype=bool)
plan_done = np.ones((self.max_plan_length + 1,), dtype=bool)
infos = []
observations[0] = obs
env_done[0] = False
for k, u in enumerate(self.plan(option)):
plan_done[k] = False
o, r, d, i = super().step(u)
actions[k] = u
rewards[k] = r
env_done[k] = d
infos.append(i)
observations[k+1] = o
if render_mode is not None:
self.env.render(render_mode)
if d:
break
n_steps = k + 1
return observations, actions, rewards, env_done, plan_done, infos, n_steps
def step(self, action, render_mode=None):
a = int(action)
assert a == action
ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode)
hl_obs = ll_obs[ll_steps]
hl_reward = (ll_rewards * ~ll_plan_done).sum().item()
hl_done = ll_env_done[ll_steps-1].item()
hl_infos = {
'll': {
'observations': ll_obs,
'actions': ll_actions,
'rewards': ll_rewards,
'env_done': ll_env_done,
'plan_done': ll_plan_done,
'infos': ll_infos,
'steps': ll_steps,
}
}
self.last_obs = hl_obs
return hl_obs, hl_reward, hl_done, hl_infos
def reset(self, *args, **kwargs):
self.last_obs = super().reset(*args, **kwargs)
return self.last_obs

View File

@@ -18,7 +18,7 @@ class OptionsRollout:
def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value,
v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma,
gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None):
policy(torch.zeros(env_fn(0).observation_space.shape)) policy(torch.zeros(env_fn(0).observation_space.shape))
policy = ReparamPolicy(policy) policy = ReparamPolicy(policy)
@@ -49,11 +49,14 @@ def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value
value, policy = trpo_step(value, policy, generator_data.hl.states, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping) value, policy = trpo_step(value, policy, generator_data.hl.states, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping)
expert_data = roll_buffer(expert_data, shifts=-3, dims=0) expert_data = roll_buffer(expert_data, shifts=-3, dims=0)
if callback is not None:
callback(epoch, value, policy)
return value, policy return value, policy
def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value, def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value,
v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma, v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma,
gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger()): gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None):
logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0]) logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0])
logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0]) logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0])
@@ -81,6 +84,9 @@ def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, v
value, policy = ppo_step(value, policy, generator_data.hl.states, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm) value, policy = ppo_step(value, policy, generator_data.hl.states, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm)
expert_data = roll_buffer(expert_data, shifts=-3, dims=0) expert_data = roll_buffer(expert_data, shifts=-3, dims=0)
if callback is not None:
callback(epoch, value, policy)
return value, policy return value, policy
def rollout(env_fn, policy, n_episodes, max_steps_per_episode): def rollout(env_fn, policy, n_episodes, max_steps_per_episode):
@@ -99,7 +105,6 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode):
env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes)))) env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes))))
states[:, 0] = torch.tensor(env.reset()).clone().detach() states[:, 0] = torch.tensor(env.reset()).clone().detach()
dones[:, 0] = False
for s in tqdm(range(max_steps_per_episode), 'Rollout'): for s in tqdm(range(max_steps_per_episode), 'Rollout'):
actions[:, s] = policy.sample(policy(states[:, s])).clone().detach() actions[:, s] = policy.sample(policy(states[:, s])).clone().detach()
@@ -111,7 +116,7 @@ def rollout(env_fn, policy, n_episodes, max_steps_per_episode):
o, r, d, info = env.step(clipped_actions) o, r, d, info = env.step(clipped_actions)
states[:, s + 1] = torch.tensor(o).clone().detach() states[:, s + 1] = torch.tensor(o).clone().detach()
rewards[:, s] = torch.tensor(r).clone().detach() rewards[:, s] = torch.tensor(r).clone().detach()
dones[:, s + 1] = torch.tensor(d).clone().detach() dones[:, s] = torch.tensor(d).clone().detach()
ll_states[:, s] = torch.from_numpy(np.stack([i['ll']['observations'] for i in info])).clone().detach() ll_states[:, s] = torch.from_numpy(np.stack([i['ll']['observations'] for i in info])).clone().detach()
ll_actions[:, s] = torch.from_numpy(np.stack([i['ll']['actions'] for i in info])).clone().detach() ll_actions[:, s] = torch.from_numpy(np.stack([i['ll']['actions'] for i in info])).clone().detach()
@@ -162,7 +167,7 @@ class OptionsEnv(gym.Wrapper):
o, r, d, i = super().step(u) o, r, d, i = super().step(u)
actions[k] = u actions[k] = u
rewards[k] = r rewards[k] = r
env_done[k+1] = d env_done[k] = d
infos.append(i) infos.append(i)
observations[k+1] = o observations[k+1] = o
@@ -181,7 +186,7 @@ class OptionsEnv(gym.Wrapper):
ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode) ll_obs, ll_actions, ll_rewards, ll_env_done, ll_plan_done, ll_infos, ll_steps = self.execute_plan(self.last_obs, self.options[a], render_mode)
hl_obs = ll_obs[ll_steps] hl_obs = ll_obs[ll_steps]
hl_reward = (ll_rewards * ~ll_plan_done).sum().item() hl_reward = (ll_rewards * ~ll_plan_done).sum().item()
hl_done = ll_env_done[ll_steps].item() hl_done = ll_env_done[ll_steps-1].item()
hl_infos = { hl_infos = {
'll': { 'll': {
'observations': ll_obs, 'observations': ll_obs,

View File

@@ -0,0 +1,185 @@
import torch
import numpy as np
from intersim.collisions import state_to_polygon
def safety_plan(env, plan):
return np.concatenate((plan, np.array(5 * [env._env._min_acc])), axis=0)
def available_actions(env, options):
"""Return mask of available actions given current `env` state."""
plans = [generate_plan(env, i, options) for i, _ in enumerate(options)]
# is emergency braking still possible?
plans = list(map(lambda p: safety_plan(env, p), plans))
T = max(len(p) for p in plans)
plans = [np.pad(p, ((0, T-len(p)),), constant_values=np.nan) for p in plans]
plans = np.stack(plans, axis=0)
valid = feasible(env, plans)
return valid
def target_velocity_plan(current_v: float, target_v: float, t: int, dt: float):
"""Smoothly target a velocity in a given number of steps"""
# for now, constant acceleration
a = (target_v - current_v) / (t * dt)
return a*np.ones((t,))
def generate_plan(env, i, options):
"""Generate input profile for high-level action `i`."""
assert i < len(options), "Invalid option index {i}"
target_v, t = options[i]
current_v = env._env.state[env._agent, 1].item() # extract from env
plan = target_velocity_plan(current_v, target_v, t, env._env._dt)
assert len(plan) == t, "incorrect plan length"
return plan
def feasible(env, plan, method='exact'):
"""Check if input profile is feasible given current `env` state."""
# zero pad plan - Take (B, T) or (T,) np plan and convert it to (B, T, nv, 1) torch.Tensor
plan = torch.tensor(plan)
plan = plan.reshape(-1, plan.shape[-1])
full_plan = torch.zeros(*plan.shape, env._env._nv, 1)
full_plan[:, :, env._agent, 0] = plan
# check_future_collisions_fast takes in B-list and outputs (B,) bool tensor
if method=='circle':
valid = check_future_collisions_fast(env, full_plan)
elif method=='ncircles':
valid = check_future_collisions_ncircles(env, full_plan)
elif method=='exact':
valid = check_future_collisions_exact(env, full_plan)
else:
raise NotImplementedError('Invalid collision-checking method')
return valid
def check_future_collisions_ncircles(env, actions, n_circles:int=2):
"""Checks whether `env._agent` would collide with other agents assuming `actions` as input.
Vehicles are (over-)approximated by multiple circles.
Args:
env (gym.Env): current environment state
actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles
Returns:
feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free
"""
assert n_circles >= 2
B, (T, nv, _) = len(actions), actions[0].shape
states = env._env.propagate_action_profile_vectorized(actions)
assert states.shape == (B, T, nv, 5)
centers = states[:, :, :, :2]
psi = states[:, :, :, 3]
lon = torch.stack([psi.cos(), psi.sin()],dim=-1) # (B, T, nv, 2)
# offset between [-env._env.lengths+env._env.widths/2, env._env.lengths/2-env._env.widths/2]
back = (-env._env._lengths/2+env._env._widths/2).unsqueeze(-1) # (nv, 1)
length = (env._env._lengths-env._env._widths).unsqueeze(-1) # (nv, 1)
diff_d = back + length*(torch.arange(n_circles)/(n_circles-1)).unsqueeze(0) # (nv, n_circles)
assert diff_d.shape == (nv, n_circles)
offsets = diff_d[None, None, :, :, None] * lon[:, :, :, None, :]
assert offsets.shape == (B, T, nv, n_circles, 2)
expanded_centers=centers.unsqueeze(-2) + offsets #(B, T, nv, n_circles, 2)
assert expanded_centers.shape == (B, T, nv, n_circles, 2)
agent_centers = expanded_centers[:,:,env._agent:env._agent+1,:,:] #(B, T, 1, n_circles, 2)
ds = expanded_centers.reshape((B, T, nv*n_circles, 1, 2)) - agent_centers #(B, T, nv*nc,1, 2) - (B, T, 1, nc, 2) = (B, T, nv*nc, nc, 2)
distance = (ds**2).sum(-1).sqrt().reshape((B, T, nv, n_circles, n_circles)) # (B, T, nv, nc, nc)
distance = torch.where(distance.isnan(), np.inf*torch.ones_like(distance), distance) # only collide with spawned agents
distance[:, :, env._agent] = np.inf # cannot collide with itself
assert distance.shape == (B, T, nv, n_circles, n_circles)
radius = env._env._widths*np.sqrt(2) / 2
min_distance = radius[env._agent] + radius
min_distance = min_distance[None, None, :, None, None]
assert min_distance.shape == (1, 1, nv, 1, 1)
return (distance > min_distance).all(-1).all(-1).all(-1).all(-1)
def check_future_collisions_circle(env, actions):
"""Compute collision information for circular vehicle approximations
Args:
env (gym.Env): current environment state
actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles
Returns:
states (torch.Tensor): tensor of shape (B, T, nv, 5) of future states based on the action profiles
collision_tensor (torch.Tensor): tensor of shape (B, T, nv) of bools indicating which plan collides with which vehicles in which time frame
false: colliding, true: not colliding
"""
B, (T, nv, _) = len(actions), actions[0].shape
states = env._env.propagate_action_profile_vectorized(actions)
assert states.shape == (B, T, nv, 5)
distance = ((states[:, :, :, :2] - states[:, :, env._agent:env._agent+1, :2])**2).sum(-1).sqrt()
distance = torch.where(distance.isnan(), np.inf*torch.ones_like(distance), distance) # only collide with spawned agents
distance[:, :, env._agent] = np.inf # cannot collide with itself
assert distance.shape == (B, T, nv)
radius = (env._env._lengths**2 + env._env._widths**2).sqrt() / 2
min_distance = radius[env._agent] + radius
min_distance = min_distance.unsqueeze(0).unsqueeze(0)
assert min_distance.shape == (1, 1, nv)
collision_tensor = distance > min_distance
assert collision_tensor.shape == (B, T, nv)
return states, collision_tensor
def check_future_collisions_fast(env, actions):
"""Checks whether `env._agent` would collide with other agents assuming `actions` as input.
Vehicles are (over-)approximated by single circles.
Args:
env (gym.Env): current environment state
actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles
Returns:
feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free
"""
_, collision_tensor = check_future_collisions_circle(env, actions)
return collision_tensor.all(-1).all(-1)
def check_future_collisions_exact(env, actions):
"""
Checks whether `env._agent` would collide with other agents assuming `actions` as input.
Args:
env (gym.Env): current environment state
actions (list of torch.Tensor): list of B (T, nv, adims) T-length action profiles
Returns:
feasible (torch.Tensor): tensor of shape (B,) indicating whether the respective action profiles are collision-free
"""
# First check with simple circle collision check
states, collision_tensor = check_future_collisions_circle(env, actions)
(B, T, nv, _) = states.shape
# For those that have colliding circles, check exactly
colliding_mask = ~collision_tensor
ego_states = states[:, :, env._agent:env._agent+1, :].expand(states.shape)
assert ego_states.shape == states.shape
# get dimensions
lengths = env._env._lengths.expand(states.shape[:3])
widths = env._env._widths.expand(states.shape[:3])
ego_lengths = lengths[:, :, env._agent:env._agent+1].expand(lengths.shape)
ego_widths = widths[:, :, env._agent:env._agent+1].expand(widths.shape)
assert lengths.shape == widths.shape == ego_lengths.shape == ego_widths.shape == (B, T, nv)
# For every collision instance between ego and other vehicle, check whether rectangles intersect
exact_collisions = torch.zeros_like(collision_tensor[colliding_mask])
for i, (ego_state, ego_length, ego_width, other_state, other_length, other_width) in enumerate(zip(
ego_states[colliding_mask], ego_lengths[colliding_mask], ego_widths[colliding_mask],
states[colliding_mask], lengths[colliding_mask], widths[colliding_mask]
)):
assert ego_state.shape == other_state.shape == (5,)
assert ego_length.shape == ego_width.shape == other_length.shape == other_width.shape == ()
p_ego = state_to_polygon(ego_state, ego_length, ego_width)
p_other = state_to_polygon(other_state, other_length, other_width)
exact_collisions[i] = p_ego.intersects(p_other)
collision_tensor[colliding_mask] = ~exact_collisions
return collision_tensor.all(-1).all(-1)

260
src/safe_options/options.py Normal file
View File

@@ -0,0 +1,260 @@
import gym
import numpy as np
import torch
from stable_baselines3.common.vec_env import DummyVecEnv as VecEnv
from src.core.reparam_module import ReparamPolicy
from tqdm import tqdm
from src.core.gail import train_discriminator, roll_buffer, TerminalLogger
from dataclasses import dataclass
from src.safe_options.policy_gradient import trpo_step, ppo_step
import torch.nn.functional as F
from src.options.envs import OptionsEnv
from src.safe_options.collisions import feasible
from intersim.envs import IntersimpleLidarFlatIncrementingAgent
from src.util.wrappers import OptionsTimeLimit, Setobs, TransformObservation
@dataclass
class Buffer:
states: torch.Tensor
actions: torch.Tensor
rewards: torch.Tensor
dones: torch.Tensor
@dataclass
class HLBuffer:
states: torch.Tensor
safe_actions: torch.Tensor
actions: torch.Tensor
rewards: torch.Tensor
dones: torch.Tensor
@dataclass
class OptionsRollout:
hl: HLBuffer
ll: Buffer
def gail(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value,
v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma,
gae_lambda, delta, backtrack_coeff, backtrack_iters, cg_iters=10, cg_damping=0.1, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None):
policy(torch.zeros(env_fn(0).observation_space['observation'].shape), torch.zeros(env_fn(0).observation_space['safe_actions'].shape))
policy = ReparamPolicy(policy)
logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0])
logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0])
for epoch in tqdm(range(epochs)):
hl_data, ll_data = rollout(env_fn, policy, rollout_episodes, rollout_steps)
generator_data = OptionsRollout(HLBuffer(*hl_data), Buffer(*ll_data))
generator_data.ll.actions += 0.1 * torch.randn_like(generator_data.ll.actions)
logger.add_scalar('gen/mean_episode_length', (~generator_data.ll.dones).sum() / generator_data.ll.states.shape[0], epoch)
logger.add_scalar('gen/mean_reward_per_episode', generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0], epoch)
logger.add_scalar('gen/unsafe_probability_mass', policy.unsafe_probability_mass(policy(generator_data.hl.states[~generator_data.hl.dones], generator_data.hl.safe_actions[~generator_data.hl.dones])).mean(), epoch)
discriminator, loss = train_discriminator(expert_data, generator_data.ll, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c)
if wasserstein:
generator_data.ll.rewards = discriminator(generator_data.ll.states, generator_data.ll.actions)
else:
generator_data.ll.rewards = -F.logsigmoid(discriminator(generator_data.ll.states, generator_data.ll.actions))
logger.add_scalar('disc/final_loss', loss, epoch)
logger.add_scalar('disc/mean_reward_per_episode', generator_data.ll.rewards[~generator_data.ll.dones].sum() / generator_data.ll.states.shape[0], epoch)
#assert generator_data.ll.rewards.shape == generator_data.ll.dones.shape
generator_data.hl.rewards = torch.where(~generator_data.ll.dones, generator_data.ll.rewards, torch.tensor(0.)).sum(-1)
value, policy = trpo_step(value, policy, generator_data.hl.states, generator_data.hl.safe_actions, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters, cg_damping)
expert_data = roll_buffer(expert_data, shifts=-3, dims=0)
if callback is not None:
callback(epoch, value, policy)
return value, policy
def gail_ppo(env_fn, expert_data, discriminator, disc_opt, disc_iters, policy, value,
v_opt, v_iters, epochs, rollout_episodes, rollout_steps, gamma,
gae_lambda, clip_ratio, pi_opt, pi_iters, target_kl=None, max_grad_norm=None, wasserstein=False, wasserstein_c=None, logger=TerminalLogger(), callback=None):
logger.add_scalar('expert/mean_episode_length', (~expert_data.dones).sum() / expert_data.states.shape[0])
logger.add_scalar('expert/mean_reward_per_episode', expert_data.rewards[~expert_data.dones].sum() / expert_data.states.shape[0])
for epoch in range(epochs):
hl_data, ll_data = rollout(env_fn, policy, rollout_episodes, rollout_steps)
generator_data = OptionsRollout(HLBuffer(*hl_data), Buffer(*ll_data))
generator_data.ll.actions += 0.1 * torch.randn_like(generator_data.ll.actions)
logger.add_scalar('gen/mean_episode_length', (~generator_data.ll.dones).sum() / generator_data.ll.states.shape[0], epoch)
logger.add_scalar('gen/mean_reward_per_episode', generator_data.hl.rewards[~generator_data.hl.dones].sum() / generator_data.hl.states.shape[0], epoch)
logger.add_scalar('gen/unsafe_probability_mass', policy.unsafe_probability_mass(policy(generator_data.hl.states[~generator_data.hl.dones], generator_data.hl.safe_actions[~generator_data.hl.dones])).mean(), epoch)
discriminator, loss = train_discriminator(expert_data, generator_data.ll, discriminator, disc_opt, disc_iters, wasserstein, wasserstein_c)
if wasserstein:
generator_data.ll.rewards = discriminator(generator_data.ll.states, generator_data.ll.actions)
else:
generator_data.ll.rewards = -F.logsigmoid(discriminator(generator_data.ll.states, generator_data.ll.actions))
logger.add_scalar('disc/final_loss', loss, epoch)
logger.add_scalar('disc/mean_reward_per_episode', generator_data.ll.rewards[~generator_data.ll.dones].sum() / generator_data.ll.states.shape[0], epoch)
#assert generator_data.ll.rewards.shape == generator_data.ll.dones.shape
generator_data.hl.rewards = torch.where(~generator_data.ll.dones, generator_data.ll.rewards, torch.tensor(0.)).sum(-1)
value, policy = ppo_step(value, policy, generator_data.hl.states, generator_data.hl.safe_actions, generator_data.hl.actions, generator_data.hl.rewards, generator_data.hl.dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm)
expert_data = roll_buffer(expert_data, shifts=-3, dims=0)
if callback is not None:
callback(epoch, value, policy)
return value, policy
def rollout(env_fn, policy, n_episodes, max_steps_per_episode):
env = env_fn(0)
states = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space['observation'].shape)
safe_actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.observation_space['safe_actions'].shape)
actions = torch.zeros(n_episodes, max_steps_per_episode + 1, *env.action_space.shape)
rewards = torch.zeros(n_episodes, max_steps_per_episode + 1)
dones = torch.ones(n_episodes, max_steps_per_episode + 1, dtype=bool)
ll_states = torch.zeros(n_episodes, max_steps_per_episode, env.max_plan_length + 1, *env.observation_space['observation'].shape)
ll_actions = torch.zeros(n_episodes, max_steps_per_episode, env.max_plan_length + 1, *env.ll_action_space.shape)
ll_rewards = torch.zeros(n_episodes, max_steps_per_episode, env.max_plan_length + 1)
ll_dones = torch.ones(n_episodes, max_steps_per_episode, env.max_plan_length + 1, dtype=bool)
env = VecEnv(list(map(lambda i: (lambda: env_fn(i)), range(n_episodes))))
obs = env.reset()
states[:, 0] = torch.tensor(obs['observation']).clone().detach()
safe_actions[:, 0] = torch.tensor(obs['safe_actions']).clone().detach()
dones[:, 0] = False
for s in tqdm(range(max_steps_per_episode), 'Rollout'):
actions[:, s] = policy.sample(policy(states[:, s], safe_actions[:, s])).clone().detach()
clipped_actions = actions[:, s]
if isinstance(env.action_space, gym.spaces.Box):
clipped_actions = torch.clamp(clipped_actions, torch.from_numpy(env.action_space.low), torch.from_numpy(env.action_space.high))
o, r, d, info = env.step(clipped_actions)
states[:, s + 1] = torch.tensor(o['observation']).clone().detach()
safe_actions[:, s + 1] = torch.tensor(o['safe_actions']).clone().detach()
rewards[:, s] = torch.tensor(r).clone().detach()
dones[:, s + 1] = torch.tensor(d).clone().detach()
ll_states[:, s] = torch.from_numpy(np.stack([i['ll']['observations'] for i in info])).clone().detach()
ll_actions[:, s] = torch.from_numpy(np.stack([i['ll']['actions'] for i in info])).clone().detach()
ll_rewards[:, s] = torch.from_numpy(np.stack([i['ll']['rewards'] for i in info])).clone().detach()
ll_dones[:, s] = torch.from_numpy(np.stack([i['ll']['plan_done'] for i in info])).clone().detach()
dones = dones.cumsum(1) > 0
states = states[:, :max_steps_per_episode]
safe_actions = safe_actions[:, :max_steps_per_episode]
actions = actions[:, :max_steps_per_episode]
rewards = rewards[:, :max_steps_per_episode]
dones = dones[:, :max_steps_per_episode]
return (states, safe_actions, actions, rewards, dones), (ll_states, ll_actions, ll_rewards, ll_dones)
class SafeOptionsEnv(OptionsEnv):
def __init__(self, env, options, safe_actions_collision_method=None, abort_unsafe_collision_method=None):
super().__init__(env, options)
self.safe_actions_collision_method = safe_actions_collision_method
self.abort_unsafe_collision_method = abort_unsafe_collision_method
self.observation_space = gym.spaces.Dict({
'observation': self.observation_space,
'safe_actions': gym.spaces.Box(low=0., high=1., shape=(self.action_space.n,)),
})
def safe_actions(self):
if self.safe_actions_collision_method is None:
return np.ones(len(self.options), dtype=bool)
plans = [self.plan(o) for o in self.options]
plans = np.stack(plans)
safe = feasible(self.env, plans, method=self.safe_actions_collision_method)
if not safe.any():
# action 0 is considered safe fallback
safe[0] = True
return safe
def reset(self, *args, **kwargs):
obs = super().reset(*args, **kwargs)
obs = {
'observation': obs,
'safe_actions': self.safe_actions(),
}
return obs
def step(self, action, render_mode=None):
obs, reward, done, info = super().step(action, render_mode)
obs = {
'observation': obs,
'safe_actions': self.safe_actions(),
}
return obs, reward, done, info
def execute_plan(self, obs, option, render_mode=None):
observations = np.zeros((self.max_plan_length + 1, *self.env.observation_space.shape))
actions = np.zeros((self.max_plan_length + 1, *self.ll_action_space.shape))
rewards = np.zeros((self.max_plan_length + 1,))
env_done = np.ones((self.max_plan_length + 1,), dtype=bool)
plan_done = np.ones((self.max_plan_length + 1,), dtype=bool)
infos = []
plan = self.plan(option)
observations[0] = obs
env_done[0] = False
for k, u in enumerate(plan):
plan_done[k] = False
o, r, d, i = self.env.step(u)
actions[k] = u
rewards[k] = r
env_done[k] = d
infos.append(i)
observations[k+1] = o
if render_mode is not None:
self.env.render(render_mode)
if d:
break
if self.abort_unsafe_collision_method is not None and \
not feasible(self.env, plan[k:], method=self.abort_unsafe_collision_method):
break
n_steps = k + 1
return observations, actions, rewards, env_done, plan_done, infos, n_steps
obs_min = np.array([
[-1000, -1000, 0, -np.pi, -1e-1, 0.],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
[0, -np.pi, -20, -20, -np.pi, -1e-1],
]).reshape(-1)
obs_max = np.array([
[1000, 1000, 20, np.pi, 1e-1, 0.],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
[50, np.pi, 20, 20, np.pi, 1e-1],
]).reshape(-1)
def NormalizedSafeOptionsEvalEnv(max_episode_steps=float('inf'), safe_actions_collision_method=None, abort_unsafe_collision_method=None, **kwargs):
return OptionsTimeLimit(SafeOptionsEnv(Setobs(
TransformObservation(IntersimpleLidarFlatIncrementingAgent(
n_rays=5,
**kwargs,
), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method=safe_actions_collision_method, abort_unsafe_collision_method=abort_unsafe_collision_method), max_episode_steps=max_episode_steps)

View File

@@ -0,0 +1,44 @@
import torch
import torch.nn as nn
from torch.distributions import Categorical
from torch.distributions.kl import kl_divergence
from src.core.policy import SetDiscretePolicy
class SetMaskedDiscretePolicy(SetDiscretePolicy):
def forward(self, observation, safe_actions):
return torch.cat((super().forward(observation), safe_actions), -1)
def torch_dist(self, dist):
logits = dist[..., :self.action_dim]
z = dist[..., self.action_dim:]
a = super().torch_dist(logits).probs
return Categorical(probs=a*z)
def unsafe_probability_mass(self, dist):
logits = dist[..., :self.action_dim]
z = dist[..., self.action_dim:]
a = super().torch_dist(logits).probs
return (a * (1 - z)).sum(-1)
def predict(self, observations, state=None, episode_start=None, deterministic=True):
observation = torch.tensor(observations['observation'])
safe_actions = torch.tensor(observations['safe_actions'])
if deterministic:
_, actions = self.forward(observation, safe_actions).max(-1)
else:
actions = self.sample(self.forward(observation, safe_actions))
return actions, None
# def torch_dist_nomask(self, dist):
# print('no mask logprob')
# logits = dist[..., :self.action_dim]
# return super().torch_dist(logits)
# def log_prob(self, dist, actions):
# return self.torch_dist_nomask(dist).log_prob(actions)
# def kl_divergence(self, dist1, dist2):
# d1 = self.torch_dist_nomask(dist1)
# d2 = self.torch_dist_nomask(dist2)
# return kl_divergence(d1, d2)

View File

@@ -0,0 +1,107 @@
import torch
from src.core.value_estimation import gae
from src.core.optimization import conjugate_gradient, line_search
def trpo_step(value, policy, states, safe_actions, actions, rewards, dones, gamma, gae_lambda, delta, backtrack_coeff, backtrack_iters, v_opt, v_iters, cg_iters=10, cg_damping=0.1):
states = states.detach()
actions = actions.detach()
rewards = rewards.detach()
dones = dones.detach()
advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda)
advantages = advantages.detach()
returns = returns.detach()
# update value function
for _ in range(v_iters):
v_opt.zero_grad()
value_loss = (value(states) - returns).pow(2)[valid].mean()
value_loss.backward()
v_opt.step()
# compute policy gradient
plogprob = policy.log_prob(policy(states, safe_actions), actions)
surrogate_advantage = (plogprob * advantages)[valid].sum() / states.shape[0]
g = torch.cat(torch.autograd.grad(surrogate_advantage, policy.flat_param)).detach()
def Hx(x):
kl = policy.kl_divergence(policy(states, safe_actions), policy(states, safe_actions).detach())[valid].mean()
dKL = torch.cat(torch.autograd.grad(kl, policy.flat_param, create_graph=True))
H_x = torch.cat(torch.autograd.grad(dKL.T @ x, policy.flat_param)).detach()
return H_x + cg_damping * x
x = conjugate_gradient(Hx, g, cg_iters)
npg = torch.sqrt(2 * delta / (x.T @ Hx(x))) * x
# perform line search
def L(theta):
rplogprob = policy.log_prob(policy(states, safe_actions, flat_param=theta), actions)
return ((rplogprob - plogprob.detach()).exp() * advantages)[valid].sum() / advantages.shape[0]
condition = lambda theta: policy.kl_divergence(policy(states, safe_actions, flat_param=theta), policy(states, safe_actions))[valid].mean() < delta
x0 = policy.flat_param
g0 = torch.cat(torch.autograd.grad(L(x0), x0))
theta = line_search(L, x0, npg, g0, backtrack_coeff, condition, max_steps=backtrack_iters)
# update policy parameters
with torch.no_grad():
policy.flat_param.copy_(theta)
return value, policy
def ppo_step(value, policy, states, safe_actions, actions, rewards, dones, clip_ratio, gamma, gae_lambda, pi_opt, pi_iters, v_opt, v_iters, target_kl, max_grad_norm):
states = states.detach()
actions = actions.detach()
rewards = rewards.detach()
dones = dones.detach()
advantages, returns, valid = gae(states, rewards, value(states), dones, gamma, gae_lambda)
advantages = advantages.detach()
returns = returns.detach()
# update value function
for _ in range(v_iters):
v_opt.zero_grad()
value_loss = (value(states) - returns).pow(2)[valid].mean()
value_loss.backward()
v_opt.step()
# update policy
old_dist = policy(states, safe_actions).detach()
old_logprob = policy.log_prob(old_dist, actions).detach()
def g(advantages, clip_ratio):
return torch.where(advantages >= 0, (1 + clip_ratio) * advantages, (1 - clip_ratio) * advantages)
def L(states, actions, advantages, clip_ratio):
return torch.minimum(
(policy.log_prob(policy(states, safe_actions), actions) - old_logprob).exp() * advantages,
g(advantages, clip_ratio)
)[valid].mean()
for _ in range(pi_iters):
pi_opt.zero_grad()
ppo_loss = -L(states, actions, advantages, clip_ratio)
ppo_loss.backward()
if max_grad_norm:
torch.nn.utils.clip_grad_norm(policy.parameters(), max_grad_norm)
pi_opt.step()
kl = policy.kl_divergence(policy(states, safe_actions), old_dist)[valid].mean()
if target_kl and kl > target_kl:
break
print('KL', kl.item())
return value, policy

View File

@@ -0,0 +1,54 @@
from intersim.envs import IntersimpleLidarFlat
from options import OptionsEnv
import gym
import numpy as np
def test_obs_shape():
options = [(0, 5), (5, 5), (10, 5)]
env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options)
assert env.reset().shape == (36,)
def test_act_space():
options = [(0, 5), (5, 5), (10, 5)]
env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options)
assert env.action_space == gym.spaces.Discrete(3)
def test_plan():
options = [(0, 5), (5, 5), (10, 5)]
env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options)
env.reset()
plan = env.plan(options[0])
assert np.allclose(plan, -13.998268127441406 * np.ones((5,)))
def test_plan2():
options = [(0, 5), (5, 5), (10, 5)]
env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options)
obs = env.reset()
states, actions, rewards, dones, plan_done, infos, n_steps = env.execute_plan(obs, options[0])
assert states.shape == (6, 36)
assert rewards.shape == (6,)
assert dones.shape == (6,)
assert len(infos) == 5
def test_step():
options = [(0, 5), (5, 5), (10, 5)]
env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options)
env.reset()
obs, reward, done, _ = env.step(0)
assert obs.shape == (36,)
assert reward == 5.0
assert done == False
def test_ll_step():
options = [(0, 5), (5, 5), (10, 5)]
env = OptionsEnv(IntersimpleLidarFlat(n_rays=5), options)
env.reset()
_, _, _, info = env.step(0)
assert info['ll']['observations'].shape == (6, 36)
assert info['ll']['actions'].shape == (6, 1)
assert info['ll']['rewards'].shape == (6,)
assert info['ll']['env_done'].shape == (6,)
assert info['ll']['plan_done'].shape == (6,)
assert info['ll']['plan_done'][5] == True
assert info['ll']['steps'] == 5
assert len(info['ll']['infos']) == 5

View File

@@ -9,6 +9,10 @@ class TransformObservation(gym.wrappers.TransformObservation):
def __getattr__(self, name): def __getattr__(self, name):
return getattr(self.env, name) return getattr(self.env, name)
class OptionsTimeLimit(gym.wrappers.TimeLimit):
def __getattr__(self, name):
return getattr(self.env, name)
class CollisionPenaltyWrapper(Wrapper): class CollisionPenaltyWrapper(Wrapper):
def __init__(self, env, collision_distance, collision_penalty, *args, **kwargs): def __init__(self, env, collision_distance, collision_penalty, *args, **kwargs):