MCPcopy Create free account
hub / github.com/aigc3d/LHM / __init__

Method __init__

engine/BiRefNet/dataset.py:35–112  ·  view source on GitHub ↗
(self, datasets, image_size, is_train=True)

Source from the content-addressed store, hash-verified

33
34class MyData(data.Dataset):
35 def __init__(self, datasets, image_size, is_train=True):
36 self.size_train = image_size
37 self.size_test = image_size
38 self.keep_size = not config.size
39 self.data_size = config.size
40 self.is_train = is_train
41 self.load_all = config.load_all
42 self.device = config.device
43 valid_extensions = [".png", ".jpg", ".PNG", ".JPG", ".JPEG"]
44
45 if self.is_train and config.auxiliary_classification:
46 self.cls_name2id = {
47 _name: _id for _id, _name in enumerate(class_labels_TR_sorted)
48 }
49 self.transform_image = transforms.Compose(
50 [
51 transforms.Resize(self.data_size[::-1]),
52 transforms.ToTensor(),
53 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
54 ][self.load_all or self.keep_size :]
55 )
56 self.transform_label = transforms.Compose(
57 [
58 transforms.Resize(self.data_size[::-1]),
59 transforms.ToTensor(),
60 ][self.load_all or self.keep_size :]
61 )
62 dataset_root = os.path.join(config.data_root_dir, config.task)
63 # datasets can be a list of different datasets for training on combined sets.
64 self.image_paths = []
65 for dataset in datasets.split("+"):
66 image_root = os.path.join(dataset_root, dataset, "im")
67 self.image_paths += [
68 os.path.join(image_root, p)
69 for p in os.listdir(image_root)
70 if any(p.endswith(ext) for ext in valid_extensions)
71 ]
72 self.label_paths = []
73 for p in self.image_paths:
74 for ext in valid_extensions:
75 ## 'im' and 'gt' may need modifying
76 p_gt = p.replace("/im/", "/gt/")[: -(len(p.split(".")[-1]) + 1)] + ext
77 file_exists = False
78 if os.path.exists(p_gt):
79 self.label_paths.append(p_gt)
80 file_exists = True
81 break
82 if not file_exists:
83 print("Not exists:", p_gt)
84
85 if len(self.label_paths) != len(self.image_paths):
86 set_image_paths = set(
87 [os.path.splitext(p.split(os.sep)[-1])[0] for p in self.image_paths]
88 )
89 set_label_paths = set(
90 [os.path.splitext(p.split(os.sep)[-1])[0] for p in self.label_paths]
91 )
92 print("Path diff:", set_image_paths - set_label_paths)

Callers

nothing calls this directly

Calls 2

printFunction · 0.85
path_to_imageFunction · 0.85

Tested by

no test coverage detected