changing how tune reporting works so the scheduler doesnt break if itcant find cv loss. also fixed bug in config structure that was rendering impossible policies
This commit is contained in:
40
src/bc/bc.py
40
src/bc/bc.py
@@ -20,8 +20,11 @@ def bc_config(ray_config):
|
|||||||
'hidden_dim': ray_config['deepsets_phi_hidden_dim']
|
'hidden_dim': ray_config['deepsets_phi_hidden_dim']
|
||||||
},
|
},
|
||||||
'latent_dim': ray_config['deepsets_latent_dim'],
|
'latent_dim': ray_config['deepsets_latent_dim'],
|
||||||
'rho': {'hidden_n': 0, 'hidden_dim': 10},
|
'rho': {
|
||||||
'output_dim': 0
|
'hidden_n': ray_config['deepsets_rho_hidden_n'],
|
||||||
|
'hidden_dim': ray_config['deepsets_rho_hidden_dim']
|
||||||
|
},
|
||||||
|
'output_dim': ray_config['deepsets_output_dim']
|
||||||
},
|
},
|
||||||
'path_encoder': {'input_dim': 40, 'hidden_n': 0, 'hidden_dim': 0, 'output_dim': 0},
|
'path_encoder': {'input_dim': 40, 'hidden_n': 0, 'hidden_dim': 0, 'output_dim': 0},
|
||||||
'head': {
|
'head': {
|
||||||
@@ -36,14 +39,13 @@ def bc_config(ray_config):
|
|||||||
'lr':ray_config['lr'],
|
'lr':ray_config['lr'],
|
||||||
'weight_decay':ray_config['weight_decay']
|
'weight_decay':ray_config['weight_decay']
|
||||||
},
|
},
|
||||||
'train_epochs': 100,
|
'train_epochs': 50,
|
||||||
'train_batch_size': ray_config['train_batch_size'],
|
'train_batch_size': ray_config['train_batch_size'],
|
||||||
'loss': ray_config['loss'],
|
'loss': ray_config['loss'],
|
||||||
|
|
||||||
}
|
}
|
||||||
return config
|
return config
|
||||||
|
|
||||||
|
|
||||||
class BehaviorCloningPolicy():
|
class BehaviorCloningPolicy():
|
||||||
"""
|
"""
|
||||||
Class for (continuous) behavior cloning policy
|
Class for (continuous) behavior cloning policy
|
||||||
@@ -160,6 +162,8 @@ def generate_transforms(dataset):
|
|||||||
def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
|
def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
|
||||||
|
|
||||||
using_ray = kwargs.get('ray', False)
|
using_ray = kwargs.get('ray', False)
|
||||||
|
if using_ray:
|
||||||
|
print('using ray')
|
||||||
|
|
||||||
# hyperparams
|
# hyperparams
|
||||||
loss_type = config['loss']
|
loss_type = config['loss']
|
||||||
@@ -169,9 +173,9 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
|
|||||||
learning_rate = config['optim']['lr']
|
learning_rate = config['optim']['lr']
|
||||||
weight_decay = config['optim']['weight_decay']
|
weight_decay = config['optim']['weight_decay']
|
||||||
|
|
||||||
cv_every = 5
|
cv_every = 1
|
||||||
print_epoch_every = 1000
|
print_epoch_every = 1000
|
||||||
print_cv_every = 1000
|
print_cv_every = 5
|
||||||
checkpoint_every = 100
|
checkpoint_every = 100
|
||||||
cv_batch_size = 256 # doesn't matter
|
cv_batch_size = 256 # doesn't matter
|
||||||
|
|
||||||
@@ -207,6 +211,10 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
|
|||||||
|
|
||||||
for i in tqdm(range(train_epochs)):
|
for i in tqdm(range(train_epochs)):
|
||||||
|
|
||||||
|
# save model checkpoints
|
||||||
|
if i % checkpoint_every == 0:
|
||||||
|
policy.save_model(filestr + '_epoch%04i'%(i) )
|
||||||
|
|
||||||
# train
|
# train
|
||||||
epoch_loss = 0
|
epoch_loss = 0
|
||||||
for (batch_idx, batch) in enumerate(training_loader):
|
for (batch_idx, batch) in enumerate(training_loader):
|
||||||
@@ -222,12 +230,6 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
|
|||||||
|
|
||||||
epoch_loss += loss.item() / len(train_dataset)
|
epoch_loss += loss.item() / len(train_dataset)
|
||||||
|
|
||||||
# Write epoch loss
|
|
||||||
if using_ray:
|
|
||||||
tune.report(training_loss=epoch_loss, training_iteration=i)
|
|
||||||
else:
|
|
||||||
writer.add_scalar('training loss',epoch_loss, i)
|
|
||||||
|
|
||||||
# if i % print_epoch_every == 0:
|
# if i % print_epoch_every == 0:
|
||||||
# print('Epoch: {}, Training Loss: {}'.format(i, epoch_loss))
|
# print('Epoch: {}, Training Loss: {}'.format(i, epoch_loss))
|
||||||
|
|
||||||
@@ -241,16 +243,20 @@ def train(config, policy, train_dataset, cv_dataset, filestr, **kwargs):
|
|||||||
loss = cv_loss_fn(pred_action, batch['action'])
|
loss = cv_loss_fn(pred_action, batch['action'])
|
||||||
cv_loss += loss.item() / len(cv_dataset)
|
cv_loss += loss.item() / len(cv_dataset)
|
||||||
|
|
||||||
if using_ray:
|
|
||||||
tune.report(cv_loss=cv_loss, cv_epoch=i)
|
# Write epoch loss
|
||||||
|
if using_ray:
|
||||||
|
if i % cv_every == 0:
|
||||||
|
tune.report(training_loss=epoch_loss, cv_loss=cv_loss, training_iteration=i+1)
|
||||||
else:
|
else:
|
||||||
|
tune.report(training_loss=epoch_loss, training_iteration=i+1)
|
||||||
|
else:
|
||||||
|
writer.add_scalar('training loss',epoch_loss, i)
|
||||||
|
if i % cv_every == 0:
|
||||||
writer.add_scalar('cv loss', cv_loss, i)
|
writer.add_scalar('cv loss', cv_loss, i)
|
||||||
|
|
||||||
# if i % print_cv_every == 0:
|
# if i % print_cv_every == 0:
|
||||||
# print('Epoch: {}, CV Loss: {}'.format(i, cv_loss))
|
# print('Epoch: {}, CV Loss: {}'.format(i, cv_loss))
|
||||||
|
|
||||||
# save model checkpoints
|
|
||||||
if i % checkpoint_every == 0:
|
|
||||||
policy.save_model(filestr + '_epoch%04i'%(i) )
|
|
||||||
|
|
||||||
policy.save_model(filestr)
|
policy.save_model(filestr)
|
||||||
|
|||||||
19
src/main.py
19
src/main.py
@@ -146,12 +146,6 @@ def parse_args():
|
|||||||
}
|
}
|
||||||
return kwargs
|
return kwargs
|
||||||
|
|
||||||
def main_wrapper(**kwargs):
|
|
||||||
if kwargs['all_runs']:
|
|
||||||
pass
|
|
||||||
else:
|
|
||||||
main(**kwargs)
|
|
||||||
|
|
||||||
def get_full_config(ray_config:dict, method:str)->dict:
|
def get_full_config(ray_config:dict, method:str)->dict:
|
||||||
"""
|
"""
|
||||||
Get full model configuration from ray config and method string
|
Get full model configuration from ray config and method string
|
||||||
@@ -186,7 +180,7 @@ def get_ray_config(method:str)->dict:
|
|||||||
"deepsets_rho_hidden_n": tune.choice([0,1,2]),
|
"deepsets_rho_hidden_n": tune.choice([0,1,2]),
|
||||||
"deepsets_rho_hidden_dim": tune.choice([16,32,64]),
|
"deepsets_rho_hidden_dim": tune.choice([16,32,64]),
|
||||||
"deepsets_output_dim": tune.choice([8,16,32,64]),
|
"deepsets_output_dim": tune.choice([8,16,32,64]),
|
||||||
"head_hidden_n": tune.choice([0,1,2]),
|
"head_hidden_n": tune.choice([1,2,3]),
|
||||||
"head_hidden_dim": tune.choice([16,32,64]),
|
"head_hidden_dim": tune.choice([16,32,64]),
|
||||||
"head_final_activation": tune.choice(['sigmoid', None]),
|
"head_final_activation": tune.choice(['sigmoid', None]),
|
||||||
}
|
}
|
||||||
@@ -216,8 +210,12 @@ if __name__ == '__main__':
|
|||||||
main(full_config, filestr='exp', datadir=datadir, ray=True, **kwargs)
|
main(full_config, filestr='exp', datadir=datadir, ray=True, **kwargs)
|
||||||
|
|
||||||
# set up ray tune
|
# set up ray tune
|
||||||
|
import ray
|
||||||
from ray import tune
|
from ray import tune
|
||||||
from ray.tune.schedulers import ASHAScheduler
|
from ray.tune.schedulers import ASHAScheduler
|
||||||
|
ray.shutdown()
|
||||||
|
ray.init(log_to_driver=False)
|
||||||
|
|
||||||
datadir = os.path.abspath('./expert_data')
|
datadir = os.path.abspath('./expert_data')
|
||||||
ray_config = get_ray_config(kwargs['method'])
|
ray_config = get_ray_config(kwargs['method'])
|
||||||
custom_scheduler = ASHAScheduler(
|
custom_scheduler = ASHAScheduler(
|
||||||
@@ -230,8 +228,9 @@ if __name__ == '__main__':
|
|||||||
config=ray_config,
|
config=ray_config,
|
||||||
scheduler=custom_scheduler,
|
scheduler=custom_scheduler,
|
||||||
local_dir=outdir,
|
local_dir=outdir,
|
||||||
resources_per_trial={"cpu": 2},
|
#resources_per_trial={"cpu": 2},
|
||||||
num_samples=20,
|
time_budget_s=45*60,
|
||||||
|
num_samples=30,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise Exception('No valid config found')
|
raise Exception('No valid config found')
|
||||||
@@ -240,4 +239,4 @@ if __name__ == '__main__':
|
|||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user