making test case for typing bug and fixing some small typing errors in bc
This commit is contained in:
35
src/bc/bc.py
35
src/bc/bc.py
@@ -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'])
|
||||||
|
|||||||
@@ -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'])
|
||||||
|
|||||||
BIN
tests/policies/base__model.pt
Normal file
BIN
tests/policies/base__model.pt
Normal file
Binary file not shown.
BIN
tests/policies/base__test_batch.pkl
Normal file
BIN
tests/policies/base__test_batch.pkl
Normal file
Binary file not shown.
BIN
tests/policies/base__transforms.pkl
Normal file
BIN
tests/policies/base__transforms.pkl
Normal file
Binary file not shown.
14
tests/policies/test_forward_pass.py
Normal file
14
tests/policies/test_forward_pass.py
Normal 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)
|
||||||
Reference in New Issue
Block a user