Implement ValueDICE and some restructuring
This commit is contained in:
86
src/bc/bc.py
86
src/bc/bc.py
@@ -4,8 +4,9 @@ from torch.utils.data import DataLoader
|
||||
import pickle
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
|
||||
from src.policies import DeepSetsPolicy
|
||||
from src.policies import IntersimDeepSetsNet, IntersimPolicy, generate_transforms
|
||||
from src.util.transform import MinMaxScaler
|
||||
from src.util.nn_training import optimizer_factory
|
||||
from tqdm import tqdm
|
||||
import json5
|
||||
from ray import tune
|
||||
@@ -46,62 +47,21 @@ def bc_config(ray_config):
|
||||
}
|
||||
return config
|
||||
|
||||
class BehaviorCloningPolicy():
|
||||
class BehaviorCloningPolicy(IntersimPolicy):
|
||||
"""
|
||||
Class for (continuous) behavior cloning policy
|
||||
"""
|
||||
|
||||
def __init__(self, config: dict, transforms: dict={}):
|
||||
def __init__(self, config: dict, transforms: dict):
|
||||
"""
|
||||
Initialize BehaviorCloningPolicy
|
||||
Args:
|
||||
config (dict): configuration file to initialize DeepSetsPolicy with
|
||||
config (dict): configuration file to initialize IntersimDeepSetsNet with
|
||||
transforms (dict): dictionary of transforms to apply to different fields
|
||||
"""
|
||||
self._config = config
|
||||
self._transforms = transforms
|
||||
self._policy = DeepSetsPolicy(config)
|
||||
super(BehaviorCloningPolicy, self).__init__(config, transforms)
|
||||
self._policy = IntersimStateNet(config)
|
||||
|
||||
@property
|
||||
def transforms(self):
|
||||
return self._transforms
|
||||
|
||||
@transforms.setter
|
||||
def transforms(self, transforms):
|
||||
self._transforms=transforms
|
||||
|
||||
@property
|
||||
def policy(self):
|
||||
return self._policy
|
||||
|
||||
@policy.setter
|
||||
def policy(self, policy):
|
||||
self._policy = policy
|
||||
|
||||
def __call__(self, ob):
|
||||
|
||||
if 'action' in ob.keys():
|
||||
# extract state from dataloader samples
|
||||
pass
|
||||
else:
|
||||
# extract state from observation (using simulator)
|
||||
ob['path_x'] = ob['paths'][0]
|
||||
ob['path_y'] = ob['paths'][1]
|
||||
|
||||
# run observation through transforms
|
||||
for key in ['state', 'relative_state', 'path_x', 'path_y']:
|
||||
if key in self._transforms.keys():
|
||||
ob[key] = self._transforms[key].transform(ob[key])
|
||||
|
||||
# run transformed state through model
|
||||
action = self._policy(ob)
|
||||
assert action.ndim == 2, 'action has incorrect shape'
|
||||
|
||||
# untransform action
|
||||
if 'action' in self._transforms.keys():
|
||||
action = self._transforms['action'].inverse_transform(action)
|
||||
return action
|
||||
|
||||
@classmethod
|
||||
def load_model(cls, filestr: str, config: dict = None):
|
||||
"""
|
||||
@@ -141,24 +101,6 @@ class BehaviorCloningPolicy():
|
||||
pickle.dump(self._transforms, open(filestr+'_transforms.pkl', 'wb'))
|
||||
torch.save(self._policy.state_dict(), filestr+'_model.pt')
|
||||
|
||||
def generate_transforms(dataset):
|
||||
"""
|
||||
Generate transform dictionary from dataset
|
||||
Args:
|
||||
dataset (Dataset): dataset of demo observations and actions
|
||||
"""
|
||||
transforms = {
|
||||
'action': MinMaxScaler(),
|
||||
'state': MinMaxScaler(),
|
||||
'relative_state': MinMaxScaler(reduce_dim=2),
|
||||
'path_x': MinMaxScaler(reduce_dim=2),
|
||||
'path_y': MinMaxScaler(reduce_dim=2),
|
||||
}
|
||||
for key in transforms.keys():
|
||||
transforms[key].fit(dataset[:][key])
|
||||
|
||||
return transforms
|
||||
|
||||
def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
|
||||
|
||||
using_ray = kwargs.get('ray', False)
|
||||
@@ -169,9 +111,6 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
|
||||
loss_type = config['loss']
|
||||
train_epochs = config['train_epochs']
|
||||
train_batch_size = config['train_batch_size']
|
||||
optimizer_type = config['optim']['optimizer']
|
||||
learning_rate = config['optim']['lr']
|
||||
weight_decay = config['optim']['weight_decay']
|
||||
|
||||
cv_every = 1
|
||||
print_epoch_every = 1000
|
||||
@@ -179,12 +118,6 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
|
||||
checkpoint_every = 100
|
||||
cv_batch_size = 256 # doesn't matter
|
||||
|
||||
# generate transform from train_dataset
|
||||
transforms = generate_transforms(train_dataset)
|
||||
|
||||
# initialize policy
|
||||
policy.transforms = transforms
|
||||
|
||||
# training and testing dataloaders
|
||||
training_loader = DataLoader(train_dataset, batch_size=train_batch_size, shuffle=True)
|
||||
cv_loader = DataLoader(cv_dataset, batch_size=cv_batch_size, shuffle=True)
|
||||
@@ -200,10 +133,7 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
|
||||
loss_fn = nn.MSELoss(reduction='sum')
|
||||
else:
|
||||
raise NotImplementedError
|
||||
if optimizer_type == 'adam':
|
||||
optimizer = torch.optim.Adam(policy.parameters(), lr=learning_rate, weight_decay=weight_decay)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
optimizer = optimizer_factory(config['optim'], policy.parameters)
|
||||
|
||||
# generate tensorboard writer
|
||||
if not using_ray:
|
||||
|
||||
Reference in New Issue
Block a user