updating main testing function to use configs and seeds, finishing first pass at behavior cloning policy and training loop. not yet tested
This commit is contained in:
92
src/bc/bc.py
92
src/bc/bc.py
@@ -1,19 +1,29 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data import DataLoader, RandomSampler
|
||||
import pickle
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
|
||||
from src.policies import
|
||||
from src.policies import DeepSetsPolicy
|
||||
from src.util.transform import SciKitMinMaxScaler
|
||||
import json5
|
||||
|
||||
class BehaviorCloningPolicy():
|
||||
"""
|
||||
Class for (continuous) behavior cloning policy
|
||||
"""
|
||||
|
||||
def __init__(self, transforms={}, **kwargs):
|
||||
def __init__(self, config: dict transforms: dict={}):
|
||||
"""
|
||||
Initialize BehaviorCloningPolicy
|
||||
Args:
|
||||
config (dict): configuration file to initialize DeepSetsPolicy with
|
||||
transforms (dict): dictionary of transforms to apply to different fields
|
||||
"""
|
||||
self._config = config
|
||||
self._transforms = transforms
|
||||
self._policy_model = PolicyModel(**kwargs)
|
||||
|
||||
self._policy_model = DeepSetsPolicy(config["ego_state"], config["deepsets"], config["path_encoder"], config["head"])
|
||||
|
||||
@property
|
||||
def transforms(self):
|
||||
return self._transforms
|
||||
@@ -35,31 +45,38 @@ class BehaviorCloningPolicy():
|
||||
# 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](ob[key])
|
||||
ob[key] = self._transforms[key].transform(ob[key])
|
||||
|
||||
# run transformed state through model
|
||||
action = self._policy_model(ob)
|
||||
assert action.ndim == 2, 'action has incorrect shape'
|
||||
|
||||
# untransform action
|
||||
|
||||
pass
|
||||
if 'action' in self._transforms.keys():
|
||||
action = self._transforms['action'].inverse_transform(action)
|
||||
return action
|
||||
|
||||
@classmethod
|
||||
def load_model(cls, filestr, **kwargs):
|
||||
def load_model(cls, config: dict, filestr: str):
|
||||
"""
|
||||
Load a model from a file prefix
|
||||
Args:
|
||||
config (dict): configuration dict to set up model
|
||||
filestr (str): string prefix to load model from
|
||||
Returns
|
||||
model (BehaviorCloningPolicy): loaded model
|
||||
"""
|
||||
transforms = pickle.load(open(filestr+'_transforms.pkl', 'rb'))
|
||||
model = cls(transforms=transforms, **kwargs)
|
||||
model = cls(config, transforms=transforms)
|
||||
model._policy_model.load_state_dict(torch.load(filestr+'_model.pt'))
|
||||
return model
|
||||
|
||||
def eval(self):
|
||||
self._policy_model.eval()
|
||||
|
||||
def parameters(self):
|
||||
return self._policy_model.parameters()
|
||||
|
||||
def save_model(self, filestr):
|
||||
"""
|
||||
Save transforms and state_dict to a location specificed by filestr
|
||||
@@ -72,41 +89,72 @@ class BehaviorCloningPolicy():
|
||||
def generate_transforms(dataset):
|
||||
"""
|
||||
Generate transform dictionary from dataset
|
||||
Args:
|
||||
dataset (Dataset): dataset of demo observations and actions
|
||||
"""
|
||||
pass
|
||||
transforms = {
|
||||
'action': SciKitMinMaxScaler()
|
||||
'state': SciKitMinMaxScaler()
|
||||
'relative_state': SciKitMinMaxScaler(reduce_dim=2)
|
||||
'path_x': SciKitMinMaxScaler(reduce_dim=2)
|
||||
'path_y': SciKitMinMaxScaler(reduce_dim=2)
|
||||
}
|
||||
for key in transforms.keys():
|
||||
transforms[key].fit(dataset[:][key])
|
||||
|
||||
def train(train_dataset, cv_dataset, policy_class, filestr=filestr, **kwargs):
|
||||
return transforms
|
||||
|
||||
def train(train_dataset, cv_dataset, policy, filestr, **kwargs):
|
||||
|
||||
# hyperparams
|
||||
train_epochs = 10000
|
||||
cv_every = 100
|
||||
train_batch_size = 64
|
||||
cv_batch_size = 256
|
||||
cv_batch_size = 256 # doesn't matter
|
||||
learning_rate = 1e-3
|
||||
weight_decay=0.1
|
||||
|
||||
# generate transform from train_dataset
|
||||
transforms = generate_transforms(train_dataset)
|
||||
|
||||
# initialize policy
|
||||
policy = BehaviorCloningPolicy(transforms=transforms, **kwargs)
|
||||
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)
|
||||
|
||||
# generate loss function, optimizer
|
||||
loss_fn = nn.HuberLoss(reduction='sum')
|
||||
optimizer = torch.optim.Adam(policy.parameters(), lr=learning_rate, weight_decay=weight_decay)
|
||||
|
||||
for i in train_epochs:
|
||||
|
||||
# sample mini-batch and run through policy
|
||||
pred_action = policy(batch)
|
||||
epoch_loss = 0
|
||||
for (batch_idx, batch) in enumerate(training_loader):
|
||||
# sample mini-batch and run through policy
|
||||
pred_action = policy(batch)
|
||||
loss = loss_fn(pred_action, batch['action'])
|
||||
|
||||
# compute loss and step optimizer
|
||||
loss.backwards()
|
||||
optimizer.step()
|
||||
# compute loss and step optimizer
|
||||
optimizer.zero_grad()
|
||||
loss.backwards()
|
||||
optimizer.step()
|
||||
|
||||
epoch_loss += loss.item() / len(train_dataset)
|
||||
|
||||
# Write epoch loss
|
||||
if i % 10 == 0:
|
||||
print('Epoch: {}, Training Loss: {}'.format(i, epoch_loss))
|
||||
|
||||
# measure L2 on every cv epoch
|
||||
# measure cv loss
|
||||
if i % cv_every == 0:
|
||||
pass
|
||||
|
||||
with torch.no_grad():
|
||||
cv_loss = 0.
|
||||
for (batch_idx, batch) in enumerate(cv_loader):
|
||||
pred_action = policy(batch)
|
||||
loss = loss_fn(pred_action, batch['action'])
|
||||
cv_loss += loss.item() / len(cv_dataset)
|
||||
print('Epoch: {}, CV Loss: {}'.format(i, cv_loss))
|
||||
|
||||
policy.save_model(filestr)
|
||||
|
||||
Reference in New Issue
Block a user