Fix bug: model.eval() is back now
This commit is contained in:
+2
-2
@@ -152,7 +152,7 @@ if __name__ == "__main__":
|
|||||||
model_version = conf.model.experiment_name.split('_v')[-1][:6]
|
model_version = conf.model.experiment_name.split('_v')[-1][:6]
|
||||||
model = get_double_photon_model_class(model_version)().cuda()
|
model = get_double_photon_model_class(model_version)().cuda()
|
||||||
model.load_state_dict(torch.load(f'{conf.model.base_dir}/{conf.model.experiment_name}/Models/{conf.model.name}', weights_only=True))
|
model.load_state_dict(torch.load(f'{conf.model.base_dir}/{conf.model.experiment_name}/Models/{conf.model.name}', weights_only=True))
|
||||||
# model.eval()
|
model.eval()
|
||||||
|
|
||||||
### data loading
|
### data loading
|
||||||
files_list = get_files_list(conf)
|
files_list = get_files_list(conf)
|
||||||
@@ -160,7 +160,7 @@ if __name__ == "__main__":
|
|||||||
BinningFactor = conf.inference.binning_factor
|
BinningFactor = conf.inference.binning_factor
|
||||||
numberOfAugOps = conf.inference.num_aug_ops
|
numberOfAugOps = conf.inference.num_aug_ops
|
||||||
flag_normalize = conf.data.normalize
|
flag_normalize = conf.data.normalize
|
||||||
nChunks = np.ceil(len(files_list) / 16).astype(int)
|
nChunks = np.ceil(len(files_list) / conf.inference.chunk_size).astype(int)
|
||||||
|
|
||||||
ml_super_frame = np.zeros((NY*BinningFactor, NX*BinningFactor), dtype=np.int32)
|
ml_super_frame = np.zeros((NY*BinningFactor, NX*BinningFactor), dtype=np.int32)
|
||||||
count_frame = np.zeros((NY, NX), dtype=np.int32)
|
count_frame = np.zeros((NY, NX), dtype=np.int32)
|
||||||
|
|||||||
Reference in New Issue
Block a user