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:
Arec
2021-07-28 09:55:11 -07:00
parent 5a7090a21c
commit 09f77e0587
2 changed files with 32 additions and 27 deletions

View File

@@ -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)

View File

@@ -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')