making test case for typing bug and fixing some small typing errors in bc

This commit is contained in:
Arec
2021-07-22 08:23:46 -07:00
parent 827a8e7172
commit 91d052445e
6 changed files with 42 additions and 17 deletions

View File

@@ -22,7 +22,7 @@ class BehaviorCloningPolicy():
""" """
self._config = config self._config = config
self._transforms = transforms self._transforms = transforms
self._policy_model = DeepSetsPolicy(config["ego_state"], config["deepsets"], config["path_encoder"], config["head"]) self._policy = DeepSetsPolicy(config["ego_state"], config["deepsets"], config["path_encoder"], config["head"])
@property @property
def transforms(self): def transforms(self):
@@ -32,9 +32,17 @@ class BehaviorCloningPolicy():
def transforms(self, transforms): def transforms(self, transforms):
self._transforms=transforms self._transforms=transforms
@property
def policy(self):
return self._policy
@policy.setter
def policy(self, policy):
self._policy = policy
def __call__(self, ob): def __call__(self, ob):
if 'actions' in ob.keys(): if 'action' in ob.keys():
# extract state from dataloader samples # extract state from dataloader samples
pass pass
else: else:
@@ -48,7 +56,7 @@ class BehaviorCloningPolicy():
ob[key] = self._transforms[key].transform(ob[key]) ob[key] = self._transforms[key].transform(ob[key])
# run transformed state through model # run transformed state through model
action = self._policy_model(ob) action = self._policy(ob)
assert action.ndim == 2, 'action has incorrect shape' assert action.ndim == 2, 'action has incorrect shape'
# untransform action # untransform action
@@ -68,14 +76,14 @@ class BehaviorCloningPolicy():
""" """
transforms = pickle.load(open(filestr+'_transforms.pkl', 'rb')) transforms = pickle.load(open(filestr+'_transforms.pkl', 'rb'))
model = cls(config, transforms=transforms) model = cls(config, transforms=transforms)
model._policy_model.load_state_dict(torch.load(filestr+'_model.pt')) model._policy.load_state_dict(torch.load(filestr+'_model.pt'))
return model return model
def eval(self): def eval(self):
self._policy_model.eval() self._policy.eval()
def parameters(self): def parameters(self):
return self._policy_model.parameters() return self._policy.parameters()
def save_model(self, filestr): def save_model(self, filestr):
""" """
@@ -84,7 +92,7 @@ class BehaviorCloningPolicy():
filestr (str): string prefix to save model to filestr (str): string prefix to save model to
""" """
pickle.dump(self._transforms, open(filestr+'_transforms.pkl', 'wb')) pickle.dump(self._transforms, open(filestr+'_transforms.pkl', 'wb'))
torch.save(self._policy_model.state_dict(), filestr+'_model.pt') torch.save(self._policy.state_dict(), filestr+'_model.pt')
def generate_transforms(dataset): def generate_transforms(dataset):
""" """
@@ -112,7 +120,7 @@ def train(train_dataset, cv_dataset, policy, filestr, **kwargs):
train_batch_size = 64 train_batch_size = 64
cv_batch_size = 256 # doesn't matter cv_batch_size = 256 # doesn't matter
learning_rate = 1e-3 learning_rate = 1e-3
weight_decay=0.1 weight_decay = 0.1
# generate transform from train_dataset # generate transform from train_dataset
transforms = generate_transforms(train_dataset) transforms = generate_transforms(train_dataset)
@@ -124,16 +132,19 @@ def train(train_dataset, cv_dataset, policy, filestr, **kwargs):
training_loader = DataLoader(train_dataset, batch_size=train_batch_size, shuffle=True) training_loader = DataLoader(train_dataset, batch_size=train_batch_size, shuffle=True)
cv_loader = DataLoader(cv_dataset, batch_size=cv_batch_size, shuffle=True) cv_loader = DataLoader(cv_dataset, batch_size=cv_batch_size, shuffle=True)
# change policy dtype
policy.policy = policy.policy.type(train_dataset[0]['state'].dtype)
# generate loss function, optimizer # generate loss function, optimizer
loss_fn = nn.HuberLoss(reduction='sum') loss_fn = nn.HuberLoss(reduction='sum')
import pdb
pdb.set_trace()
optimizer = torch.optim.Adam(policy.parameters(), lr=learning_rate, weight_decay=weight_decay) optimizer = torch.optim.Adam(policy.parameters(), lr=learning_rate, weight_decay=weight_decay)
pickle.dump(train_dataset[150:160], open(filestr+'_test_batch.pkl', 'wb'))
for i in train_epochs: policy.save_model(filestr)
for i in range(train_epochs):
epoch_loss = 0 epoch_loss = 0
for (batch_idx, batch) in enumerate(training_loader): for (batch_idx, batch) in enumerate(training_loader):
# sample mini-batch and run through policy # sample mini-batch and run through policy
pred_action = policy(batch) pred_action = policy(batch)
loss = loss_fn(pred_action, batch['action']) loss = loss_fn(pred_action, batch['action'])

View File

@@ -43,11 +43,11 @@ class InteractionDatasetSingleAgent(Dataset):
for t in range(T): for t in range(T):
nni = ~torch.isnan(observations[t]['state'][:,0]) nni = ~torch.isnan(observations[t]['state'][:,0])
max_nv = max(max_nv,nni.count_nonzero()) max_nv = max(max_nv,nni.count_nonzero())
self.raw_data['state'].append(observations[t]['state'][nni]) self.raw_data['state'].append(observations[t]['state'][nni].float())
self.raw_data['relative_state'].append(observations[t]['relative_state'][nni.nonzero(),nni.nonzero()]) self.raw_data['relative_state'].append(observations[t]['relative_state'][nni.nonzero(),nni.nonzero()].float())
self.raw_data['action'].append(actions[t][nni]) self.raw_data['action'].append(actions[t][nni].float())
self.raw_data['path_x'].append(observations[t]['paths'][0][nni]) self.raw_data['path_x'].append(observations[t]['paths'][0][nni].float())
self.raw_data['path_y'].append(observations[t]['paths'][1][nni]) self.raw_data['path_y'].append(observations[t]['paths'][1][nni].float())
# cat lists # cat lists
self.raw_data['state'] = torch.cat(self.raw_data['state']) self.raw_data['state'] = torch.cat(self.raw_data['state'])

Binary file not shown.

Binary file not shown.

Binary file not shown.

View File

@@ -0,0 +1,14 @@
import torch
import pickle
import json5
from src.bc import BehaviorCloningPolicy
config_path = "config/networks.json5"
with open(config_path, 'r') as cfg:
config = json5.load(cfg)
filestr = 'tests/policies/base_'
batch = pickle.load(open(filestr+'_test_batch.pkl', 'rb'))
policy = BehaviorCloningPolicy.load_model(config, filestr)
action = policy(batch)