code update

This commit is contained in:
Alex Blank
2025-05-19 22:54:53 +02:00
parent 58305effdb
commit cb7896900b
7 changed files with 238 additions and 51 deletions
+8 -3
View File
@@ -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