Merge pull request #4 from sisl/options
Integrate options env and policy
This commit is contained in:
BIN
checkpoints/bc-intersimple-setobs2.pt
Normal file
BIN
checkpoints/bc-intersimple-setobs2.pt
Normal file
Binary file not shown.
BIN
checkpoints/gail-intersimple-setobs2-03-02-22.pt
Normal file
BIN
checkpoints/gail-intersimple-setobs2-03-02-22.pt
Normal file
Binary file not shown.
BIN
checkpoints/gail-options-setobs2-15-02-2022.pt
Normal file
BIN
checkpoints/gail-options-setobs2-15-02-2022.pt
Normal file
Binary file not shown.
BIN
checkpoints/gail-options-setobs2-Feb15_18-49-05.pt
Normal file
BIN
checkpoints/gail-options-setobs2-Feb15_18-49-05.pt
Normal file
Binary file not shown.
BIN
checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt
Normal file
BIN
checkpoints/gail-ppo-options-setobs2-Feb15_22-05-38.pt
Normal file
Binary file not shown.
BIN
checkpoints/sgail-options-setobs2.pt
Normal file
BIN
checkpoints/sgail-options-setobs2.pt
Normal file
Binary file not shown.
BIN
checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt
Normal file
BIN
checkpoints/sgail-ppo-options-setobs2-17-02-2022.pt
Normal file
Binary file not shown.
BIN
checkpoints/wgail-options-setobs2-Feb16_01-06-27.pt
Normal file
BIN
checkpoints/wgail-options-setobs2-Feb16_01-06-27.pt
Normal file
Binary file not shown.
BIN
checkpoints/wgail-ppo-options-setobs2-Feb16_04-02-56.pt
Normal file
BIN
checkpoints/wgail-ppo-options-setobs2-Feb16_04-02-56.pt
Normal file
Binary file not shown.
@@ -13,4 +13,20 @@ python -m src.eval_main
|
||||
# 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}'
|
||||
|
||||
232
scratch/etienne/trpo/.gitignore
vendored
232
scratch/etienne/trpo/.gitignore
vendored
@@ -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
|
||||
@@ -26,7 +26,7 @@ torch.save(policy.state_dict(), 'bc-intersimple-setobs2.pt')
|
||||
# %%
|
||||
import numpy as np
|
||||
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.intersimple import speed_reward
|
||||
import functools
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Minobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Minobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -8,7 +8,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Minobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Minobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
from core.reparam_module import ReparamPolicy
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Setobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -8,7 +8,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Setobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
from core.reparam_module import ReparamPolicy
|
||||
@@ -8,7 +8,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Setobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
from core.reparam_module import ReparamPolicy
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
|
||||
envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat(
|
||||
n_rays=5,
|
||||
@@ -8,7 +8,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs
|
||||
import numpy as np
|
||||
from options.options import OptionsEnv
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -8,7 +8,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
import numpy as np
|
||||
from options.options import OptionsEnv
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -9,7 +9,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
import numpy as np
|
||||
from options.options import OptionsEnv
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -56,6 +56,11 @@ 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'gail-options-setobs2-{epoch}.pt')
|
||||
torch.save(value.state_dict(), f'gail-options-setobs2-value-{epoch}.pt')
|
||||
|
||||
value, policy = gail(
|
||||
env_fn=env_fn,
|
||||
expert_data=expert_data,
|
||||
@@ -66,7 +71,7 @@ value, policy = gail(
|
||||
value=value,
|
||||
v_opt=v_opt,
|
||||
v_iters=1000,
|
||||
epochs=200,
|
||||
epochs=300,
|
||||
rollout_episodes=60,
|
||||
rollout_steps=60,
|
||||
gamma=0.99,
|
||||
@@ -75,6 +80,7 @@ value, policy = gail(
|
||||
backtrack_coeff=0.8,
|
||||
backtrack_iters=10,
|
||||
logger=SummaryWriter(comment='gail-options-setobs2'),
|
||||
callback=callback,
|
||||
)
|
||||
|
||||
torch.save(policy.state_dict(), 'gail-options-setobs2.pt')
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Minobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Minobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -8,7 +8,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Setobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
from core.reparam_module import ReparamPolicy
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
|
||||
envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat(
|
||||
n_rays=5,
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Minobs
|
||||
import numpy as np
|
||||
from options.options import OptionsEnv
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
import numpy as np
|
||||
from options.options import OptionsEnv
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -8,7 +8,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
import numpy as np
|
||||
from options.options import OptionsEnv
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -57,6 +57,11 @@ 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'gail-ppo-options-setobs2-{epoch}.pt')
|
||||
torch.save(value.state_dict(), f'gail-ppo-options-setobs2-value-{epoch}.pt')
|
||||
|
||||
value, policy = gail_ppo(
|
||||
env_fn=env_fn,
|
||||
expert_data=expert_data,
|
||||
@@ -76,6 +81,7 @@ value, policy = gail_ppo(
|
||||
pi_opt=pi_opt,
|
||||
pi_iters=100,
|
||||
logger=SummaryWriter(comment='gail-ppo-options-setobs2'),
|
||||
callback=callback,
|
||||
)
|
||||
|
||||
torch.save(policy.state_dict(), 'gail-ppo-options-setobs2.pt')
|
||||
@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
from intersim.expert import NormalizedIntersimpleExpert
|
||||
from wrappers import CollisionPenaltyWrapper, Minobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Minobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
from intersim.expert import NormalizedIntersimpleExpert
|
||||
from wrappers import CollisionPenaltyWrapper, Minobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Minobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
from intersim.expert import NormalizedIntersimpleExpert
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
from intersim.expert import NormalizedIntersimpleExpert
|
||||
from wrappers import CollisionPenaltyWrapper, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Setobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
from intersim.expert import NormalizedIntersimpleExpert
|
||||
from wrappers import CollisionPenaltyWrapper, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Setobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -4,7 +4,7 @@ from core.sampling import rollout_sb3
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
from intersim.expert import NormalizedIntersimpleExpert
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
|
||||
env = CollisionPenaltyWrapper(IntersimpleLidarFlat(
|
||||
n_rays=5,
|
||||
@@ -9,7 +9,7 @@ import torch.optim
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
from wrappers import Minobs
|
||||
from util.wrappers import Minobs
|
||||
|
||||
obs_min = np.array([
|
||||
[-1000, -1000, 0, -np.pi, -1e-1, 0.],
|
||||
@@ -9,7 +9,7 @@ import torch.optim
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
from wrappers import Minobs
|
||||
from util.wrappers import Minobs
|
||||
|
||||
obs_min = np.array([
|
||||
[-1000, -1000, 0, -np.pi, -1e-1, 0.],
|
||||
@@ -9,9 +9,9 @@ from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
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
|
||||
|
||||
obs_min = np.array([
|
||||
3
scratch/etienne/trpo/experiments/requirements.txt
Normal file
3
scratch/etienne/trpo/experiments/requirements.txt
Normal file
@@ -0,0 +1,3 @@
|
||||
torch
|
||||
stable-baselines3
|
||||
gym
|
||||
111
scratch/etienne/trpo/experiments/sgail-options-setobs2.py
Normal file
111
scratch/etienne/trpo/experiments/sgail-options-setobs2.py
Normal 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()
|
||||
|
||||
# %%
|
||||
110
scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py
Normal file
110
scratch/etienne/trpo/experiments/sgail-ppo-options-setobs2.py
Normal 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()
|
||||
# %%
|
||||
@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
from core.reparam_module import ReparamPolicy
|
||||
|
||||
from wrappers import Minobs
|
||||
from util.wrappers import Minobs
|
||||
|
||||
obs_min = np.array([
|
||||
[-1000, -1000, 0, -np.pi, -1e-1, 0.],
|
||||
@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
from core.reparam_module import ReparamPolicy
|
||||
|
||||
from wrappers import Minobs
|
||||
from util.wrappers import Minobs
|
||||
|
||||
obs_min = np.array([
|
||||
[-1000, -1000, 0, -np.pi, -1e-1, 0.],
|
||||
@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
from core.reparam_module import ReparamPolicy
|
||||
|
||||
from wrappers import Setobs
|
||||
from util.wrappers import Setobs
|
||||
|
||||
obs_min = np.array([
|
||||
[-1000, -1000, 0, -np.pi, -1e-1, 0.],
|
||||
@@ -10,10 +10,10 @@ from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
from core.reparam_module import ReparamPolicy
|
||||
|
||||
from wrappers import Setobs
|
||||
from util.wrappers import Setobs
|
||||
|
||||
obs_min = np.array([
|
||||
[-1000, -1000, 0, -np.pi, -1e-1, 0.],
|
||||
@@ -9,10 +9,10 @@ from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
import numpy as np
|
||||
from wrappers import CollisionPenaltyWrapper, TransformObservation
|
||||
from util.wrappers import CollisionPenaltyWrapper, TransformObservation
|
||||
from core.reparam_module import ReparamPolicy
|
||||
|
||||
from wrappers import Minobs
|
||||
from util.wrappers import Minobs
|
||||
from options.options import OptionsEnv
|
||||
|
||||
obs_min = np.array([
|
||||
346
scratch/etienne/trpo/experiments/vec-env.ipynb
Normal file
346
scratch/etienne/trpo/experiments/vec-env.ipynb
Normal 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
|
||||
}
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Minobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Minobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Minobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Minobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Setobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
|
||||
envs = [CollisionPenaltyWrapper(IntersimpleLidarFlat(
|
||||
n_rays=5,
|
||||
@@ -9,7 +9,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
import numpy as np
|
||||
from options.options import OptionsEnv
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -9,7 +9,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
import numpy as np
|
||||
from options.options import OptionsEnv
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -50,7 +50,7 @@ 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-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 = Buffer(*expert_data)
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Minobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Minobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, Setobs
|
||||
import numpy as np
|
||||
from gym.wrappers import TransformObservation
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -7,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
import numpy as np
|
||||
from options.options import OptionsEnv
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -1,4 +1,3 @@
|
||||
# %%
|
||||
import gym
|
||||
from options.options import gail_ppo, Buffer
|
||||
from core.value import SetValue
|
||||
@@ -8,7 +7,7 @@ import torch.optim
|
||||
from intersim.envs import IntersimpleLidarFlatRandom
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
from wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
from util.wrappers import CollisionPenaltyWrapper, TransformObservation, Setobs
|
||||
import numpy as np
|
||||
from options.options import OptionsEnv
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
@@ -51,12 +50,11 @@ 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-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 = Buffer(*expert_data)
|
||||
|
||||
# %%
|
||||
value, policy = gail_ppo(
|
||||
env_fn=env_fn,
|
||||
expert_data=expert_data,
|
||||
@@ -7,7 +7,7 @@ from intersim.envs import IntersimpleLidarFlat
|
||||
from intersim.envs.intersimple import speed_reward
|
||||
import functools
|
||||
import torch
|
||||
from wrappers import CollisionPenaltyWrapper
|
||||
from util.wrappers import CollisionPenaltyWrapper
|
||||
|
||||
model = PPO.load('sb3-ppo-intersimple')
|
||||
env = CollisionPenaltyWrapper(IntersimpleLidarFlat(
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
from core.reparam_module import ReparamPolicy
|
||||
from core.sampling import rollout
|
||||
from core.trpo import trpo_step
|
||||
from core.ppo import ppo_step
|
||||
from src.core.reparam_module import ReparamPolicy
|
||||
from src.core.sampling import rollout
|
||||
from src.core.trpo import trpo_step
|
||||
from src.core.ppo import ppo_step
|
||||
from tqdm import tqdm
|
||||
|
||||
class TerminalLogger:
|
||||
@@ -15,8 +15,13 @@ class BasePolicy(nn.Module):
|
||||
def sample(self, dist):
|
||||
return self.torch_dist(dist).sample()
|
||||
|
||||
def predict(self, states):
|
||||
return self.sample(self.forward(states))
|
||||
def predict(self, observations, state=None, episode_start=None, deterministic=True):
|
||||
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):
|
||||
return self.torch_dist(dist).log_prob(actions)
|
||||
@@ -59,6 +64,14 @@ class DiscretePolicy(BasePolicy):
|
||||
def torch_dist(self, 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):
|
||||
|
||||
def forward(self, states):
|
||||
@@ -1,6 +1,6 @@
|
||||
import torch
|
||||
from core.sampling import rollout
|
||||
from core.value_estimation import gae
|
||||
from src.core.sampling import rollout
|
||||
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):
|
||||
|
||||
@@ -160,3 +160,6 @@ class ReparamPolicy(ReparamModule):
|
||||
|
||||
def predict(self, *args, **kwargs):
|
||||
return self.module.predict(*args, **kwargs)
|
||||
|
||||
def unsafe_probability_mass(self, *args, **kwargs):
|
||||
return self.module.unsafe_probability_mass(*args, **kwargs)
|
||||
@@ -1,8 +1,8 @@
|
||||
import torch
|
||||
from core.reparam_module import ReparamPolicy
|
||||
from core.sampling import rollout
|
||||
from core.value_estimation import gae
|
||||
from core.optimization import conjugate_gradient, line_search
|
||||
from src.core.reparam_module import ReparamPolicy
|
||||
from src.core.sampling import rollout
|
||||
from src.core.value_estimation import gae
|
||||
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):
|
||||
|
||||
@@ -8,6 +8,11 @@ from src.baselines import IDMRulePolicy
|
||||
from src.evaluation import IntersimpleEvaluation
|
||||
import src.gail.options as options_envs
|
||||
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
|
||||
import torch
|
||||
@@ -33,15 +38,44 @@ def load_policy(method:str,
|
||||
if method == 'idm':
|
||||
policy = IDMRulePolicy(env, **policy_kwargs)
|
||||
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':
|
||||
policy = sb3.PPO.load(policy_file)
|
||||
raise NotImplementedError
|
||||
policy = SetPolicy(env.action_space.shape[-1])
|
||||
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':
|
||||
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':
|
||||
policy = sb3.PPO.load(policy_file)
|
||||
raise NotImplementedError
|
||||
policy = SetMaskedDiscretePolicy(env.action_space.n)
|
||||
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:
|
||||
raise NotImplementedError
|
||||
return policy
|
||||
@@ -158,6 +192,8 @@ def evaluate_policy(locations:List[Tuple[int,int]],
|
||||
"""
|
||||
envs_dict = dict(intersim.envs.intersimple.__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)
|
||||
|
||||
# iterate through vehicles
|
||||
|
||||
@@ -6,6 +6,8 @@ from typing import Callable, Dict, Optional
|
||||
import os
|
||||
import pickle
|
||||
from tqdm import tqdm
|
||||
from src.options.envs import OptionsEnv
|
||||
from src.util.wrappers import OptionsTimeLimit
|
||||
|
||||
class IntersimpleEvaluation:
|
||||
"""
|
||||
@@ -34,6 +36,7 @@ class IntersimpleEvaluation:
|
||||
self.env = eval_env
|
||||
self.n_episodes = eval_env.nv
|
||||
self.use_pbar = use_pbar
|
||||
self.is_options_env = isinstance(self.env, (OptionsEnv, OptionsTimeLimit))
|
||||
|
||||
# metrics present on every step of every episode
|
||||
self.metric_keys_all = ['v_all', 'a_all', 'col_all']
|
||||
@@ -85,11 +88,14 @@ class IntersimpleEvaluation:
|
||||
if self.use_pbar:
|
||||
self.pbar = tqdm(total=self.n_episodes)
|
||||
|
||||
if self.is_options_env:
|
||||
print('Evaluating an options environment')
|
||||
|
||||
evaluate_policy(
|
||||
policy,
|
||||
self.env,
|
||||
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
|
||||
)
|
||||
if self.use_pbar:
|
||||
@@ -100,6 +106,13 @@ class IntersimpleEvaluation:
|
||||
self.save(filestr)
|
||||
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):
|
||||
"""
|
||||
Callback run in evaluate_policy after taking an action and receiving an observation
|
||||
@@ -110,8 +123,11 @@ class IntersimpleEvaluation:
|
||||
done = local_vars['done']
|
||||
_agent = info['agent']
|
||||
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
|
||||
self._metrics['v_all'][_agent].append(info['prev_state'][_agent,2].item())
|
||||
self._metrics['a_all'][_agent].append(info['action_taken'][_agent,0].item())
|
||||
|
||||
111
src/options/envs.py
Normal file
111
src/options/envs.py
Normal 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
|
||||
@@ -18,7 +18,7 @@ class OptionsRollout:
|
||||
|
||||
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()):
|
||||
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 = 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)
|
||||
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()):
|
||||
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])
|
||||
@@ -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)
|
||||
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):
|
||||
@@ -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))))
|
||||
|
||||
states[:, 0] = torch.tensor(env.reset()).clone().detach()
|
||||
dones[:, 0] = False
|
||||
|
||||
for s in tqdm(range(max_steps_per_episode), 'Rollout'):
|
||||
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)
|
||||
states[:, s + 1] = torch.tensor(o).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_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)
|
||||
actions[k] = u
|
||||
rewards[k] = r
|
||||
env_done[k+1] = d
|
||||
env_done[k] = d
|
||||
infos.append(i)
|
||||
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)
|
||||
hl_obs = ll_obs[ll_steps]
|
||||
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 = {
|
||||
'll': {
|
||||
'observations': ll_obs,
|
||||
185
src/safe_options/collisions.py
Normal file
185
src/safe_options/collisions.py
Normal 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
260
src/safe_options/options.py
Normal 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)
|
||||
44
src/safe_options/policy.py
Normal file
44
src/safe_options/policy.py
Normal 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)
|
||||
107
src/safe_options/policy_gradient.py
Normal file
107
src/safe_options/policy_gradient.py
Normal 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
|
||||
54
src/safe_options/test_options.py
Normal file
54
src/safe_options/test_options.py
Normal 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
|
||||
@@ -9,6 +9,10 @@ class TransformObservation(gym.wrappers.TransformObservation):
|
||||
def __getattr__(self, name):
|
||||
return getattr(self.env, name)
|
||||
|
||||
class OptionsTimeLimit(gym.wrappers.TimeLimit):
|
||||
def __getattr__(self, name):
|
||||
return getattr(self.env, name)
|
||||
|
||||
class CollisionPenaltyWrapper(Wrapper):
|
||||
|
||||
def __init__(self, env, collision_distance, collision_penalty, *args, **kwargs):
|
||||
Reference in New Issue
Block a user