code update
This commit is contained in:
@@ -479,6 +479,10 @@ class LMDBIterableDataset(IterableDataset):
|
||||
self.key_subset = key_subset
|
||||
# reset length
|
||||
self.len = None
|
||||
# reset lmdb env
|
||||
if self.lmdb_env is not None:
|
||||
self.lmdb_env.close()
|
||||
self.lmdb_env = None
|
||||
|
||||
def get_length_of_data_subset(self, key_set: list[str]):
|
||||
"""
|
||||
@@ -522,12 +526,13 @@ class LMDBIterableDataset(IterableDataset):
|
||||
return num_steps
|
||||
|
||||
def __iter__(self):
|
||||
self.init_lmdb_env()
|
||||
|
||||
# if no key subset is set, use all keys
|
||||
if self.key_subset is None:
|
||||
self.key_subset = self.lmdb_keys
|
||||
|
||||
self.init_lmdb_env()
|
||||
|
||||
random.shuffle(self.key_subset)
|
||||
|
||||
batch = list()
|
||||
@@ -563,12 +568,12 @@ class LMDBIterableDataset(IterableDataset):
|
||||
meminit=False)
|
||||
|
||||
def __len__(self):
|
||||
self.init_lmdb_env()
|
||||
|
||||
# if no key subset is set, use all keys
|
||||
if self.key_subset is None:
|
||||
self.key_subset = self.lmdb_keys
|
||||
|
||||
self.init_lmdb_env()
|
||||
|
||||
if self.len is None:
|
||||
num_steps = self.get_length_of_data_subset(self.key_subset)
|
||||
self.len = num_steps
|
||||
|
||||
Reference in New Issue
Block a user