@@ -320,7 +320,11 @@ def get_ds(fn, trn, val, tfms, test=None, **kwargs):
320320 fn (val [0 ], val [1 ], tfms [0 ], ** kwargs ) # aug
321321 ]
322322 if test is not None :
323- test_lbls = np .zeros ((len (test ),1 ))
323+ if isinstance (test , tuple ):
324+ test_lbls = test [1 ]
325+ test = test [0 ]
326+ else :
327+ test_lbls = np .zeros ((len (test ),1 ))
324328 res += [
325329 fn (test , test_lbls , tfms [1 ], ** kwargs ), # test
326330 fn (test , test_lbls , tfms [0 ], ** kwargs ) # test_aug
@@ -350,7 +354,7 @@ def from_arrays(cls, path, trn, val, bs=64, tfms=(None,None), classes=None, num_
350354 return cls (path , datasets , bs , num_workers , classes = classes )
351355
352356 @classmethod
353- def from_paths (cls , path , bs = 64 , tfms = (None ,None ), trn_name = 'train' , val_name = 'valid' , test_name = None , num_workers = 8 ):
357+ def from_paths (cls , path , bs = 64 , tfms = (None ,None ), trn_name = 'train' , val_name = 'valid' , test_name = None , test_with_labels = False , num_workers = 8 ):
354358 """ Read in images and their labels given as sub-folder names
355359
356360 Arguments:
@@ -368,8 +372,10 @@ def from_paths(cls, path, bs=64, tfms=(None,None), trn_name='train', val_name='v
368372 assert isinstance (tfms [0 ], Transforms ) and isinstance (tfms [1 ], Transforms ), \
369373 "please provide transformations for your train and validation sets"
370374 trn ,val = [folder_source (path , o ) for o in (trn_name , val_name )]
371- test_fnames = read_dir (path , test_name ) if test_name else None
372- datasets = cls .get_ds (FilesIndexArrayDataset , trn , val , tfms , path = path , test = test_fnames )
375+ if test_name :
376+ test = folder_source (path , test_name ) if test_with_labels else read_dir (path , test_name )
377+ else : test = None
378+ datasets = cls .get_ds (FilesIndexArrayDataset , trn , val , tfms , path = path , test = test )
373379 return cls (path , datasets , bs , num_workers , classes = trn [2 ])
374380
375381 @classmethod
0 commit comments