refactor(data): 重构MmapFileHandler类并改进数据加载机制

This commit is contained in:
2026-01-11 19:37:28 +08:00
parent 9dab96c31f
commit 7dfa5cc0ac
3 changed files with 42 additions and 58 deletions
+2 -2
View File
@@ -4,7 +4,7 @@ import bisect
from abc import ABC, abstractmethod
from torch import Tensor
from torch.utils.data import Dataset
from khaosz.data.mmap import MmapFileHander
from khaosz.data.mmap import MmapFileHandler
from typing import Callable, List, Dict, Literal, Optional, Union
Seg = List[Tensor]
@@ -74,7 +74,7 @@ class BaseDataset(Dataset, ABC):
self.total_samples = None
def load(self, load_path: str):
self.segments, self.total_samples = MmapFileHander.load(load_path)
self.segments, self.total_samples = MmapFileHandler.load(load_path)
self.fetcher = MultiSegmentFetcher(self.segments)
def get_index(self, index: int) -> int:
+22 -36
View File
@@ -5,14 +5,14 @@ import torch
from torch import Tensor
from typing import List, Dict, Tuple
class MmapFileHander:
class MmapFileHandler:
"""
json metadata like this:
```
[
{"file_name": "file1.bin", "size": 1000, "dtype": "float32", "key": "key1"},
{"file_name": "file2.bin", "size": 2000, "dtype": "float32", "key": "key2"}
{"file_name": "file1.bin", "size": 1000, "key": "key1"},
{"file_name": "file2.bin", "size": 2000, "key": "key2"}
...
]
```
@@ -20,29 +20,20 @@ class MmapFileHander:
```
folder_path:
- file_mapper.json
- metadata.json
- file1.bin
- file2.bin
...
```
"""
DTYPE_MAP = {
"float32": torch.float32,
"float64": torch.float64,
"int32": torch.int32,
"int64": torch.int64,
"bool": torch.bool,
}
REVERSE_DTYPE_MAP = {v: k for k, v in DTYPE_MAP.items()}
META_DATA = "metadata.json"
@staticmethod
def load(root_path: str, shared: bool=True) -> Tuple[Dict[str, List[Tensor]], int]:
metadata_list = []
mmap_shared_group: Dict[str, List[Tensor]] = {}
tensor_group: Dict[str, List[Tensor]] = {}
file_mapper_path = os.path.join(root_path, "file_mapper.json")
file_mapper_path = os.path.join(root_path, MmapFileHandler.META_DATA)
if not os.path.exists(file_mapper_path):
raise FileNotFoundError(f"File mapper not found: {file_mapper_path}")
@@ -50,25 +41,20 @@ class MmapFileHander:
metadata_list = json.load(f)
for metadata in metadata_list:
file_path = os.path.join(root_path, metadata["file_name"])
if not os.path.exists(file_path):
raise FileNotFoundError(f"Binary data file not found: {file_path}")
size = metadata["size"]
dtype = MmapFileHander.DTYPE_MAP[metadata["dtype"]]
segment_key = metadata["key"]
mmap_tensor = torch.from_file(file_path, shared=shared, size=size, dtype=dtype)
if segment_key not in mmap_shared_group:
mmap_shared_group[segment_key] = []
file_key = metadata["key"]
file_name = metadata["file_name"]
file_path = os.path.join(root_path, file_name)
elm = torch.load(file_path, map_location="cpu", mmap=shared)
mmap_shared_group[segment_key].append(mmap_tensor)
if file_key not in tensor_group:
tensor_group[file_key] = []
tensor_group[file_key].append(elm)
num_samples = sum(metadata["size"] for metadata in metadata_list)
num_keys = max(len(set(metadata['key'] for metadata in metadata_list)), 1)
sample_per_key = num_samples // num_keys
return mmap_shared_group, sample_per_key
return tensor_group, sample_per_key
@staticmethod
def save(save_path: str, mmap_shared_group: Dict[str, List[Tensor]]) -> None:
@@ -79,18 +65,18 @@ class MmapFileHander:
for idx, tensor in enumerate(segment_tensors):
try:
with open(os.path.join(save_path, f"{segment_key}_{idx}.bin"), "wb") as f:
f.write(tensor.cpu().numpy().tobytes())
with open(os.path.join(save_path, f"{segment_key}_{idx}.pt"), "wb") as f:
torch.save(tensor.contiguous().cpu(), f)
except Exception as e:
raise RuntimeError(f"Error saving tensor: {e}")
metadata_list.append({
"file_name": f"{segment_key}_{idx}.bin",
"file_name": f"{segment_key}_{idx}.pt",
"size": tensor.numel(),
"dtype": MmapFileHander.REVERSE_DTYPE_MAP[tensor.dtype],
"key": segment_key
})
metadata_path = os.path.join(save_path, "file_mapper.json")
metadata_path = os.path.join(save_path, MmapFileHandler.META_DATA)
with open(metadata_path, "w") as f:
json.dump(metadata_list, f)
json.dump(metadata_list, f)