From 8f164334433ea838497df6fa4b8b5b04fb25e9ce Mon Sep 17 00:00:00 2001 From: Peter Shaw <46685926+drpetershaw@users.noreply.github.com> Date: Wed, 30 Jan 2019 12:26:35 +0000 Subject: [PATCH 1/6] Add files via upload --- 1_train_predictor.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/1_train_predictor.py b/1_train_predictor.py index e8aff59..bc575b4 100644 --- a/1_train_predictor.py +++ b/1_train_predictor.py @@ -56,6 +56,8 @@ help='save interval') parser.add_argument('--save_fig', action='store_true', help='save figure') +parser.add_argument('--path_save', type=str, default='./result/nyc_taxi', + help='file path to save data') parser.add_argument('--resume','-r', help='use checkpoint model parameters as initial parameters (default: False)', action="store_true") @@ -72,7 +74,7 @@ ############################################################################### # Load data ############################################################################### -TimeseriesData = preprocess_data.PickleDataLoad(data_type=args.data, filename=args.filename, +TimeseriesData = preprocess_data.PickleDataLoad(pathSave=args.path_save, data_type=args.data, filename=args.filename, augment_test_data=args.augment) train_dataset = TimeseriesData.batchify(args,TimeseriesData.trainData, args.batch_size) test_dataset = TimeseriesData.batchify(args,TimeseriesData.testData, args.eval_batch_size) @@ -167,7 +169,7 @@ def generate_output(args,epoch, model, gen_dataset, disp_uncertainty=True,startP plt.legend() plt.tight_layout() plt.text(startPoint-500+10, target.min(), 'Epoch: '+str(epoch),fontsize=15) - save_dir = Path('result',args.data,args.filename).with_suffix('').joinpath('fig_prediction') + save_dir = Path(args.path_save+'/result',args.data,args.filename).with_suffix('').joinpath('fig_prediction') save_dir.mkdir(parents=True,exist_ok=True) plt.savefig(save_dir.joinpath('fig_epoch'+str(epoch)).with_suffix('.png')) #plt.show() @@ -296,7 +298,7 @@ def evaluate(args, model, test_dataset): # Loop over epochs. if args.resume or args.pretrained: print("=> loading checkpoint ") - checkpoint = torch.load(Path('save', args.data, 'checkpoint', args.filename).with_suffix('.pth')) + checkpoint = torch.load(Path(args.path_save+'/save', args.data, 'checkpoint', args.filename).with_suffix('.pth')) args, start_epoch, best_val_loss = model.load_checkpoint(args,checkpoint,feature_dim) optimizer.load_state_dict((checkpoint['optimizer'])) del checkpoint From f3bc5e074b9a118a68e04aa686cfdc1d2ea6a544 Mon Sep 17 00:00:00 2001 From: Peter Shaw <46685926+drpetershaw@users.noreply.github.com> Date: Wed, 30 Jan 2019 12:29:15 +0000 Subject: [PATCH 2/6] Add files via upload --- model/model.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/model/model.py b/model/model.py index 10a66a8..322bb50 100644 --- a/model/model.py +++ b/model/model.py @@ -96,13 +96,13 @@ def repackage_hidden(self,h): def save_checkpoint(self,state, is_best): print("=> saving checkpoint ..") args = state['args'] - checkpoint_dir = Path('save',args.data,'checkpoint') + checkpoint_dir = Path(args.path_save+'/save',args.data,'checkpoint') checkpoint_dir.mkdir(parents=True,exist_ok=True) checkpoint = checkpoint_dir.joinpath(args.filename).with_suffix('.pth') torch.save(state, checkpoint) if is_best: - model_best_dir = Path('save',args.data,'model_best') + model_best_dir = Path(args.path_save+'/save',args.data,'model_best') model_best_dir.mkdir(parents=True,exist_ok=True) shutil.copyfile(checkpoint, model_best_dir.joinpath(args.filename).with_suffix('.pth')) From aa8098a4a803edb5e6e8dafd9bbeefb5df4eb288 Mon Sep 17 00:00:00 2001 From: Peter Shaw <46685926+drpetershaw@users.noreply.github.com> Date: Wed, 30 Jan 2019 12:30:38 +0000 Subject: [PATCH 3/6] Add files via upload --- 2_anomaly_detection.py | 10 ++++++---- preprocess_data.py | 6 +++--- 2 files changed, 9 insertions(+), 7 deletions(-) diff --git a/2_anomaly_detection.py b/2_anomaly_detection.py index 6d658b0..2e69851 100644 --- a/2_anomaly_detection.py +++ b/2_anomaly_detection.py @@ -21,6 +21,8 @@ help='filename of the dataset') parser.add_argument('--save_fig', action='store_true', help='save results as figures') +parser.add_argument('--path_save', type=str, default='./result/nyc_taxi', + help='file path to save data') parser.add_argument('--compensate', action='store_true', help='compensate anomaly score using anomaly score esimation') parser.add_argument('--beta', type=float, default=1.0, @@ -30,7 +32,7 @@ args_ = parser.parse_args() print('-' * 89) print("=> loading checkpoint ") -checkpoint = torch.load(str(Path('save',args_.data,'checkpoint',args_.filename).with_suffix('.pth'))) +checkpoint = torch.load(str(Path(args_.path_save+'/save',args_.data,'checkpoint',args_.filename).with_suffix('.pth'))) args = checkpoint['args'] args.prediction_window_size= args_.prediction_window_size args.beta = args_.beta @@ -46,7 +48,7 @@ ############################################################################### # Load data ############################################################################### -TimeseriesData = preprocess_data.PickleDataLoad(data_type=args.data,filename=args.filename, augment_test_data=False) +TimeseriesData = preprocess_data.PickleDataLoad(pathSave=args_.path_save, data_type=args.data,filename=args.filename, augment_test_data=False) train_dataset = TimeseriesData.batchify(args,TimeseriesData.trainData[:TimeseriesData.length], bsz=1) test_dataset = TimeseriesData.batchify(args,TimeseriesData.testData, bsz=1) @@ -139,7 +141,7 @@ if args.save_fig: - save_dir = Path('result',args.data,args.filename).with_suffix('').joinpath('fig_detection') + save_dir = Path(args_.path_save+'/result',args.data,args.filename).with_suffix('').joinpath('fig_detection') save_dir.mkdir(parents=True,exist_ok=True) plt.plot(precision.cpu().numpy(),label='precision') plt.plot(recall.cpu().numpy(),label='recall') @@ -191,7 +193,7 @@ print('=> saving the results as pickle extensions') -save_dir = Path('result', args.data, args.filename).with_suffix('') +save_dir = Path(args_.path_save+'/result', args.data, args.filename).with_suffix('') save_dir.mkdir(parents=True, exist_ok=True) pickle.dump(targets, open(str(save_dir.joinpath('target.pkl')),'wb')) pickle.dump(mean_predictions, open(str(save_dir.joinpath('mean_predictions.pkl')),'wb')) diff --git a/preprocess_data.py b/preprocess_data.py index 294ddc1..a8249b1 100644 --- a/preprocess_data.py +++ b/preprocess_data.py @@ -18,10 +18,10 @@ def reconstruct(seqData,mean,std): return seqData*std+mean class PickleDataLoad(object): - def __init__(self, data_type, filename, augment_test_data=True): + def __init__(self, pathSave, data_type, filename, augment_test_data=True): self.augment_test_data=augment_test_data - self.trainData, self.trainLabel = self.preprocessing(Path('dataset',data_type,'labeled','train',filename),train=True) - self.testData, self.testLabel = self.preprocessing(Path('dataset',data_type,'labeled','test',filename),train=False) + self.trainData, self.trainLabel = self.preprocessing(Path(pathSave+'/dataset',data_type,'labeled','train',filename),train=True) + self.testData, self.testLabel = self.preprocessing(Path(pathSave+'/dataset',data_type,'labeled','test',filename),train=False) def augmentation(self,data,label,noise_ratio=0.05,noise_interval=0.0005,max_length=100000): noiseSeq = torch.randn(data.size()) From 31e418588d2d0fd3162a1d898057d7f8c3acfc3d Mon Sep 17 00:00:00 2001 From: Peter Shaw <46685926+drpetershaw@users.noreply.github.com> Date: Wed, 30 Jan 2019 14:58:09 +0000 Subject: [PATCH 4/6] Update 1_train_predictor.py --- 1_train_predictor.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/1_train_predictor.py b/1_train_predictor.py index bc575b4..382470c 100644 --- a/1_train_predictor.py +++ b/1_train_predictor.py @@ -58,6 +58,8 @@ help='save figure') parser.add_argument('--path_save', type=str, default='./result/nyc_taxi', help='file path to save data') +parser.add_argument('--path_load', type=str, default='./result/nyc_taxi', + help='file path to load dataset from') parser.add_argument('--resume','-r', help='use checkpoint model parameters as initial parameters (default: False)', action="store_true") @@ -74,7 +76,7 @@ ############################################################################### # Load data ############################################################################### -TimeseriesData = preprocess_data.PickleDataLoad(pathSave=args.path_save, data_type=args.data, filename=args.filename, +TimeseriesData = preprocess_data.PickleDataLoad(pathVar=args.path_load, data_type=args.data, filename=args.filename, augment_test_data=args.augment) train_dataset = TimeseriesData.batchify(args,TimeseriesData.trainData, args.batch_size) test_dataset = TimeseriesData.batchify(args,TimeseriesData.testData, args.eval_batch_size) @@ -360,4 +362,4 @@ def evaluate(args, model, test_dataset): 'covs': covs } model.save_checkpoint(model_dictionary, True) -print('-' * 89) \ No newline at end of file +print('-' * 89) From 290357f7bf1172a980d71974b91f38a8d3183b8b Mon Sep 17 00:00:00 2001 From: Peter Shaw <46685926+drpetershaw@users.noreply.github.com> Date: Wed, 30 Jan 2019 14:59:38 +0000 Subject: [PATCH 5/6] Update preprocess_data.py --- preprocess_data.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/preprocess_data.py b/preprocess_data.py index a8249b1..1717c25 100644 --- a/preprocess_data.py +++ b/preprocess_data.py @@ -18,10 +18,10 @@ def reconstruct(seqData,mean,std): return seqData*std+mean class PickleDataLoad(object): - def __init__(self, pathSave, data_type, filename, augment_test_data=True): + def __init__(self, pathVar, data_type, filename, augment_test_data=True): self.augment_test_data=augment_test_data - self.trainData, self.trainLabel = self.preprocessing(Path(pathSave+'/dataset',data_type,'labeled','train',filename),train=True) - self.testData, self.testLabel = self.preprocessing(Path(pathSave+'/dataset',data_type,'labeled','test',filename),train=False) + self.trainData, self.trainLabel = self.preprocessing(Path(pathVar+'/dataset',data_type,'labeled','train',filename),train=True) + self.testData, self.testLabel = self.preprocessing(Path(pathVar+'/dataset',data_type,'labeled','test',filename),train=False) def augmentation(self,data,label,noise_ratio=0.05,noise_interval=0.0005,max_length=100000): noiseSeq = torch.randn(data.size()) From 1793c8eff1c2efeb25adde2ff6fc133689353efa Mon Sep 17 00:00:00 2001 From: Peter Shaw <46685926+drpetershaw@users.noreply.github.com> Date: Wed, 30 Jan 2019 15:21:48 +0000 Subject: [PATCH 6/6] Update 2_anomaly_detection.py --- 2_anomaly_detection.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/2_anomaly_detection.py b/2_anomaly_detection.py index 2e69851..3393cf9 100644 --- a/2_anomaly_detection.py +++ b/2_anomaly_detection.py @@ -23,6 +23,8 @@ help='save results as figures') parser.add_argument('--path_save', type=str, default='./result/nyc_taxi', help='file path to save data') +parser.add_argument('--path_load', type=str, default='./result/nyc_taxi', + help='file path to load dataset from') parser.add_argument('--compensate', action='store_true', help='compensate anomaly score using anomaly score esimation') parser.add_argument('--beta', type=float, default=1.0, @@ -48,7 +50,7 @@ ############################################################################### # Load data ############################################################################### -TimeseriesData = preprocess_data.PickleDataLoad(pathSave=args_.path_save, data_type=args.data,filename=args.filename, augment_test_data=False) +TimeseriesData = preprocess_data.PickleDataLoad(pathVar=args_.path_load, data_type=args.data,filename=args.filename, augment_test_data=False) train_dataset = TimeseriesData.batchify(args,TimeseriesData.trainData[:TimeseriesData.length], bsz=1) test_dataset = TimeseriesData.batchify(args,TimeseriesData.testData, bsz=1) @@ -204,4 +206,4 @@ pickle.dump(precisions, open(str(save_dir.joinpath('precision.pkl')),'wb')) pickle.dump(recalls, open(str(save_dir.joinpath('recall.pkl')),'wb')) pickle.dump(f_betas, open(str(save_dir.joinpath('f_beta.pkl')),'wb')) -print('-' * 89) \ No newline at end of file +print('-' * 89)