Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 8 additions & 4 deletions 1_train_predictor.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,10 @@
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('--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")
Expand All @@ -72,7 +76,7 @@
###############################################################################
# Load data
###############################################################################
TimeseriesData = preprocess_data.PickleDataLoad(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)
Expand Down Expand Up @@ -167,7 +171,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()
Expand Down Expand Up @@ -296,7 +300,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
Expand Down Expand Up @@ -358,4 +362,4 @@ def evaluate(args, model, test_dataset):
'covs': covs
}
model.save_checkpoint(model_dictionary, True)
print('-' * 89)
print('-' * 89)
14 changes: 9 additions & 5 deletions 2_anomaly_detection.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,10 @@
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('--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,
Expand All @@ -30,7 +34,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
Expand All @@ -46,7 +50,7 @@
###############################################################################
# Load data
###############################################################################
TimeseriesData = preprocess_data.PickleDataLoad(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)

Expand Down Expand Up @@ -139,7 +143,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')
Expand Down Expand Up @@ -191,7 +195,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'))
Expand All @@ -202,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)
print('-' * 89)
4 changes: 2 additions & 2 deletions model/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'))
Expand Down
6 changes: 3 additions & 3 deletions preprocess_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -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, pathVar, 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(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())
Expand Down