refactor: 重构流水线架构,添加Pipeline抽象并拆分IOHandler
This commit is contained in:
+53
-23
@@ -7,11 +7,11 @@ import torch
|
||||
import h5py
|
||||
from pathlib import Path
|
||||
|
||||
from pipeline.io import IOHandler
|
||||
from pipeline.io import FileScanner, HDF5Handler
|
||||
|
||||
|
||||
class TestIOHandler:
|
||||
def test_fetch_files_in_directory(self):
|
||||
class TestFileScanner:
|
||||
def test_scan_files_in_directory(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
Path(tmpdir, "file1.txt").touch()
|
||||
Path(tmpdir, "file2.txt").touch()
|
||||
@@ -19,42 +19,61 @@ class TestIOHandler:
|
||||
os.makedirs(subdir)
|
||||
Path(subdir, "file3.txt").touch()
|
||||
|
||||
files = IOHandler.fetch_files(tmpdir)
|
||||
files = FileScanner.scan(tmpdir)
|
||||
assert len(files) == 3
|
||||
|
||||
def test_fetch_files_empty_directory(self):
|
||||
def test_scan_empty_directory(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
assert IOHandler.fetch_files(tmpdir) == []
|
||||
assert FileScanner.scan(tmpdir) == []
|
||||
|
||||
def test_fetch_folders_in_directory(self):
|
||||
def test_scan_with_suffix_filter(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
Path(tmpdir, "file1.txt").touch()
|
||||
Path(tmpdir, "file2.json").touch()
|
||||
|
||||
txt_files = FileScanner.scan(tmpdir, suffix=".txt")
|
||||
assert len(txt_files) == 1
|
||||
assert txt_files[0].endswith(".txt")
|
||||
|
||||
def test_scan_folders_in_directory(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
os.makedirs(os.path.join(tmpdir, "folder1"))
|
||||
os.makedirs(os.path.join(tmpdir, "folder2"))
|
||||
os.makedirs(os.path.join(tmpdir, "folder1", "nested"))
|
||||
|
||||
folders = IOHandler.fetch_folders(tmpdir)
|
||||
folders = FileScanner.scan_folders(tmpdir)
|
||||
assert len(folders) == 3
|
||||
|
||||
def test_fetch_folders_with_filter(self):
|
||||
def test_scan_folders_with_filter(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
os.makedirs(os.path.join(tmpdir, "folder1"))
|
||||
os.makedirs(os.path.join(tmpdir, "folder2"))
|
||||
|
||||
folders = IOHandler.fetch_folders(
|
||||
folders = FileScanner.scan_folders(
|
||||
tmpdir, filter_func=lambda x: "folder1" in x
|
||||
)
|
||||
assert len(folders) == 1
|
||||
|
||||
def test_save_and_load_h5(self):
|
||||
def test_group_by_extension(self):
|
||||
files = ["/path/file1.txt", "/path/file2.txt", "/path/file3.json"]
|
||||
groups = FileScanner.group_by_extension(files)
|
||||
assert ".txt" in groups
|
||||
assert ".json" in groups
|
||||
assert len(groups[".txt"]) == 2
|
||||
assert len(groups[".json"]) == 1
|
||||
|
||||
|
||||
class TestHDF5Handler:
|
||||
def test_save_and_load(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
tensor_group = {
|
||||
"sequence": [torch.tensor([1, 2, 3], dtype=torch.int32)],
|
||||
"labels": [torch.tensor([4, 5], dtype=torch.int32)],
|
||||
}
|
||||
IOHandler.save_h5(tmpdir, "test", tensor_group)
|
||||
h5_path = HDF5Handler.save(tmpdir, "test", tensor_group)
|
||||
|
||||
assert os.path.exists(os.path.join(tmpdir, "test.h5"))
|
||||
loaded = IOHandler.load_h5(tmpdir, share_memory=False)
|
||||
assert os.path.exists(h5_path)
|
||||
loaded = HDF5Handler.load(h5_path, share_memory=False)
|
||||
assert "sequence" in loaded
|
||||
assert "labels" in loaded
|
||||
assert torch.equal(
|
||||
@@ -64,14 +83,14 @@ class TestIOHandler:
|
||||
loaded["labels"][0], torch.tensor([4, 5], dtype=torch.int32)
|
||||
)
|
||||
|
||||
def test_save_h5_creates_directory(self):
|
||||
def test_save_creates_directory(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
output_dir = os.path.join(tmpdir, "nested", "output")
|
||||
IOHandler.save_h5(output_dir, "test", {"data": [torch.tensor([1, 2, 3])]})
|
||||
HDF5Handler.save(output_dir, "test", {"data": [torch.tensor([1, 2, 3])]})
|
||||
assert os.path.exists(output_dir)
|
||||
assert os.path.exists(os.path.join(output_dir, "test.h5"))
|
||||
|
||||
def test_load_h5_multiple_files(self):
|
||||
def test_load_directory_with_multiple_files(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
for i, data in enumerate([[1, 2, 3], [4, 5, 6]]):
|
||||
h5_path = os.path.join(tmpdir, f"file{i}.h5")
|
||||
@@ -79,10 +98,10 @@ class TestIOHandler:
|
||||
grp = f.create_group("data")
|
||||
grp.create_dataset("data_0", data=data)
|
||||
|
||||
loaded = IOHandler.load_h5(tmpdir, share_memory=False)
|
||||
loaded = HDF5Handler.load(tmpdir, share_memory=False)
|
||||
assert len(loaded["data"]) == 2
|
||||
|
||||
def test_load_h5_with_rglob(self):
|
||||
def test_load_directory_with_nested_files(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
subdir = os.path.join(tmpdir, "subdir")
|
||||
os.makedirs(subdir)
|
||||
@@ -91,11 +110,11 @@ class TestIOHandler:
|
||||
grp = f.create_group("test")
|
||||
grp.create_dataset("data_0", data=[1, 2])
|
||||
|
||||
loaded = IOHandler.load_h5(tmpdir, share_memory=False)
|
||||
loaded = HDF5Handler.load(tmpdir, share_memory=False)
|
||||
assert "test" in loaded
|
||||
assert len(loaded["test"]) == 1
|
||||
|
||||
def test_save_h5_multiple_tensors_per_key(self):
|
||||
def test_save_multiple_tensors_per_key(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
tensor_group = {
|
||||
"batch": [
|
||||
@@ -104,6 +123,17 @@ class TestIOHandler:
|
||||
torch.tensor([6]),
|
||||
],
|
||||
}
|
||||
IOHandler.save_h5(tmpdir, "multi", tensor_group)
|
||||
loaded = IOHandler.load_h5(tmpdir, share_memory=False)
|
||||
HDF5Handler.save(tmpdir, "multi", tensor_group)
|
||||
loaded = HDF5Handler.load(tmpdir, share_memory=False)
|
||||
assert len(loaded["batch"]) == 3
|
||||
|
||||
def test_get_metadata(self):
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
tensor_group = {
|
||||
"data": [torch.tensor([1, 2, 3]) for _ in range(5)],
|
||||
}
|
||||
HDF5Handler.save(tmpdir, "meta", tensor_group)
|
||||
|
||||
h5_path = os.path.join(tmpdir, "meta.h5")
|
||||
metadata = HDF5Handler.get_metadata(h5_path)
|
||||
assert metadata["data"] == 5
|
||||
|
||||
Reference in New Issue
Block a user