refactor : Storage 层重构为 Store,移除 Fetcher 中间层,支持多段数据与显式长度
- 合并 BaseStorage + MultiSegmentFetcher + BaseSegmentFetcher 三层为 Store ABC - Store._data 直接持有 Dict[str, List[Tensor]],不做强制拼接避免 OOM - _fetch_key 统一用 bisect 跨段切片,单段多段同一路径 - _length 显式存储(min total across keys),__len__ 返回 O(1) - MmapStore/H5Store/JSONStore 统一走 _normalize() 注册分段并预计算累积长度 - 所有 I/O 函数 (save_h5/load_h5/json_to_bin 等) 保持不变
This commit is contained in:
@@ -8,8 +8,8 @@ from torch import Tensor
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from astrai.dataset.storage import (
|
||||
BaseStorage,
|
||||
StorageFactory,
|
||||
Store,
|
||||
StoreFactory,
|
||||
detect_format,
|
||||
)
|
||||
from astrai.factory import BaseFactory
|
||||
@@ -26,7 +26,7 @@ class BaseDataset(Dataset, ABC):
|
||||
super().__init__()
|
||||
self.window_size = window_size
|
||||
self.stride = stride
|
||||
self.storage: Optional[BaseStorage] = None
|
||||
self.storage: Optional[Store] = None
|
||||
|
||||
@property
|
||||
def required_keys(self) -> List[str]:
|
||||
@@ -65,7 +65,7 @@ class BaseDataset(Dataset, ABC):
|
||||
"""
|
||||
if storage_type is None:
|
||||
storage_type = detect_format(load_path)
|
||||
self.storage = StorageFactory.create(storage_type)
|
||||
self.storage = StoreFactory.create(storage_type)
|
||||
self._load_path = load_path
|
||||
self.storage.load(load_path, tokenizer=tokenizer)
|
||||
self._validate_keys()
|
||||
|
||||
Reference in New Issue
Block a user