refactor: remove H5 storage backend in favor of mmap bin

- Remove H5Store, H5Writer, save_h5/load_h5 and h5py dependency
- MmapStore (bin) is the sole pre-tokenized storage backend
- Move setup_logging after imports to fix E402 in __init__.py
- Clean up unused imports across test files
- Move inline test imports to file top
This commit is contained in:
2026-07-29 12:50:27 +08:00
parent c2b04d8458
commit 115192c67c
17 changed files with 91 additions and 209 deletions
+42 -46
View File
@@ -7,18 +7,24 @@ import pytest
import torch
from astrai.config.preprocess_config import PipelineConfig
from astrai.dataset.dataset import DatasetFactory, dpo_tokenize
from astrai.dataset.dataset import (
DatasetFactory,
GRPODataset,
dpo_tokenize,
grpo_collate_fn,
)
from astrai.dataset.storage import (
H5Store,
JsonlStore,
MmapStore,
StoreFactory,
detect_format,
)
from astrai.preprocessing.builder import SectionedMaskBuilder
from astrai.serialization import (
load_bin,
save_bin,
save_h5,
)
from tests.data.conftest import make_grpo_no_template_config
def _rand_seq(length, vocab=1000):
@@ -70,7 +76,7 @@ def _make_seq_dataset(
):
if data is None:
data = {"sequence": [_rand_seq(seq_length)]}
save_h5(test_dir, name, data)
save_bin(test_dir, data)
return DatasetFactory.load(
train_type,
test_dir,
@@ -83,12 +89,15 @@ def test_dataset_loader_random_paths(base_test_env):
"""Test dataset loader with multiple random paths"""
test_dir = base_test_env["test_dir"]
loaded_dataset = None
num_files = np.random.randint(2, 5)
for i in range(num_files):
seq_length = np.random.randint(200, 400)
dummy_data = {"sequence": [_rand_seq(seq_length) for _ in range(10)]}
sub_dir = os.path.join(test_dir, f"sub_{i}")
os.makedirs(sub_dir, exist_ok=True)
loaded_dataset = _make_seq_dataset(
test_dir, f"data_{i}", seq_length, data=dummy_data
sub_dir, f"data_{i}", seq_length, data=dummy_data
)
assert loaded_dataset is not None
assert len(loaded_dataset) > 0
@@ -113,8 +122,15 @@ def test_dpo_strategy_with_random_data(base_test_env):
"chosen_mask": [torch.ones(seq_length, dtype=torch.bool)],
"rejected_mask": [torch.ones(seq_length, dtype=torch.bool)],
}
dpo_dataset = _make_seq_dataset(
test_dir, "dpo_data", seq_length, train_type="dpo", data=dummy_data
save_bin(
test_dir,
dummy_data,
record_keys=["chosen", "rejected", "chosen_mask", "rejected_mask"],
)
dpo_dataset = DatasetFactory.load(
train_type="dpo",
load_path=test_dir,
window_size=0,
)
assert dpo_dataset is not None
@@ -196,24 +212,20 @@ def test_dataset_too_short_for_window(base_test_env):
def test_unloaded_sample_window_raises():
"""Store.sample_window before load raises RuntimeError."""
from astrai.dataset.storage import H5Store
store = H5Store(window_size=64, stride=64)
store = MmapStore(window_size=64, stride=64)
with pytest.raises(IndexError, match="Data too short"):
store.sample_window(0)
def test_unloaded_dataset_len():
"""__len__ on a store with no data returns 0."""
from astrai.dataset.storage import H5Store
store = H5Store(window_size=64, stride=64)
store = MmapStore(window_size=64, stride=64)
assert len(store) == 0
def test_store_unloaded_len():
"""Unloaded Store has __len__ == 0"""
store = H5Store()
store = MmapStore()
assert len(store) == 0
assert store.keys == []
@@ -227,7 +239,7 @@ def test_store_fetch_begin_equals_end(base_test_env):
def test_store_fetch_before_load():
"""Store.fetch before load raises RuntimeError"""
store = H5Store()
store = MmapStore()
with pytest.raises(RuntimeError, match="not loaded"):
store.fetch(0, 10, "sequence")
@@ -255,9 +267,7 @@ def test_create_store_invalid_type():
def test_store_multi_segment_concat(base_test_env):
"""Multi-segment H5 data is concatenated into single tensor at load time"""
import os
"""Multi-segment data is concatenated into single tensor at load time"""
test_dir = base_test_env["test_dir"]
data_dir = os.path.join(test_dir, "multi_seg")
os.makedirs(data_dir, exist_ok=True)
@@ -267,9 +277,9 @@ def test_store_multi_segment_concat(base_test_env):
torch.tensor([4, 5, 6, 7]),
torch.tensor([8, 9]),
]
save_h5(data_dir, "data", {"sequence": segs})
save_bin(data_dir, {"sequence": segs})
store = StoreFactory.create("h5")
store = StoreFactory.create("bin")
store.load(data_dir)
assert store.token_count == 9
result = store.fetch(2, 7, "sequence")
@@ -321,7 +331,7 @@ def test_mmap_dataset_load(base_test_env):
def test_normalize_empty_key():
"""_normalize with empty tensor list does not crash."""
store = H5Store()
store = MmapStore()
store._normalize({"sequence": []})
assert len(store) == 0
assert store.num_records == 0 # empty key forces num_records=0
@@ -330,7 +340,7 @@ def test_normalize_empty_key():
def test_normalize_mixed_empty_key():
"""_normalize with empty + non-empty keys returns min=0 records."""
store = H5Store()
store = MmapStore()
store._normalize({"sequence": [torch.tensor([1, 2, 3])], "loss_mask": []})
assert len(store) == 0
assert store.num_records == 0
@@ -340,8 +350,6 @@ def test_normalize_mixed_empty_key():
def test_grpo_dataset_dtype(base_test_env):
"""GRPO dataset returns correct dtypes for per-record structured data."""
from astrai.dataset.dataset import GRPODataset
G = 4
store = type(
"FakeStore",
@@ -373,8 +381,6 @@ def test_grpo_dataset_dtype(base_test_env):
def test_grpo_dataset_load(base_test_env):
"""GRPO dataset loads record-structured data with per-response boundaries."""
from astrai.dataset.dataset import GRPODataset
G = 3
prompt_len = 8
resp_lens = [5, 7, 4]
@@ -430,15 +436,14 @@ def test_detect_format_bin_dir(base_test_env):
def test_store_fetch_multi_key(base_test_env):
test_dir = base_test_env["test_dir"]
save_h5(
save_bin(
test_dir,
"multi_key",
{
"sequence": [torch.randint(0, 100, (100,), dtype=torch.int64)],
"loss_mask": [torch.ones(100, dtype=torch.int64)],
},
)
store = StoreFactory.create("h5")
store = StoreFactory.create("bin")
store.load(test_dir)
result = store.fetch(10, 20, ["sequence", "loss_mask"])
assert isinstance(result, dict)
@@ -448,8 +453,8 @@ def test_store_fetch_multi_key(base_test_env):
def test_store_fetch_out_of_bounds(base_test_env):
test_dir = base_test_env["test_dir"]
save_h5(test_dir, "bounds", {"sequence": [torch.randint(0, 100, (50,))]})
store = StoreFactory.create("h5")
save_bin(test_dir, {"sequence": [torch.randint(0, 100, (50,))]})
store = StoreFactory.create("bin")
store.load(test_dir)
with pytest.raises(ValueError, match="out of bounds"):
store.fetch(-1, 10, "sequence")
@@ -461,7 +466,7 @@ def test_store_fetch_out_of_bounds(base_test_env):
def test_dataset_load_explicit_storage_type(base_test_env):
test_dir = base_test_env["test_dir"]
dataset = _make_seq_dataset(test_dir, "explicit", storage_type="h5")
dataset = _make_seq_dataset(test_dir, "explicit", storage_type="bin")
assert len(dataset) > 0
assert dataset.token_count == 200
@@ -820,9 +825,6 @@ def _write_grpo_jsonl(test_dir, tokenizer_path, records):
def test_grpo_builder_preserves_response_boundaries(base_test_env):
"""MultiOutputMaskBuilder with list_field returns List[List[int]] for responses."""
from astrai.preprocessing.builder import SectionedMaskBuilder
from tests.data.conftest import make_grpo_no_template_config
tokenizer = base_test_env["tokenizer"]
_save_test_tokenizer(base_test_env["test_dir"], tokenizer)
@@ -863,8 +865,6 @@ def test_grpo_builder_preserves_response_boundaries(base_test_env):
def test_grpo_end_to_end_jsonl(base_test_env):
"""Full GRPO pipeline: JSONL → JsonlStore → GRPODataset → collate_fn."""
from astrai.dataset.dataset import grpo_collate_fn
test_dir = base_test_env["test_dir"]
tokenizer = base_test_env["tokenizer"]
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
@@ -913,8 +913,6 @@ def test_grpo_end_to_end_jsonl(base_test_env):
def test_grpo_collate_variable_lengths():
"""collate_fn pads variable-length responses to [B, G, R_max]."""
from astrai.dataset.dataset import grpo_collate_fn
batch = [
{
"prompts": torch.tensor([1, 2, 3]),
@@ -954,8 +952,6 @@ def test_grpo_collate_variable_lengths():
def test_grpo_multiple_records(base_test_env):
"""GRPODataset loads multiple records with correct structure."""
from astrai.dataset.dataset import GRPODataset
G = 4
n_records = 5
@@ -1128,8 +1124,8 @@ def test_jsonl_store_eager_len_returns_token_count(base_test_env):
assert len(store.keys) > 0
def test_h5_store_dual_mode(base_test_env):
"""H5Store supports both fetch (stream) and fetch_record (record).
def test_mmap_store_dual_mode(base_test_env):
"""MmapStore supports both fetch (stream) and fetch_record (record).
No window configured → ``len(store)`` reflects the record count
(2). ``token_count`` retains the legacy stream length (128), and
@@ -1143,9 +1139,9 @@ def test_h5_store_dual_mode(base_test_env):
"chosen": [_rand_seq(seq_length), _rand_seq(seq_length)],
"rejected": [_rand_seq(seq_length), _rand_seq(seq_length)],
}
save_h5(test_dir, "dpo_data", dummy_data)
save_bin(test_dir, dummy_data, record_keys=["chosen", "rejected"])
store = H5Store()
store = MmapStore()
store.load(test_dir)
assert store.token_count == seq_length * 2
@@ -1160,7 +1156,7 @@ def test_h5_store_dual_mode(base_test_env):
# Window-configured view of the same data uses stream sample count:
# token_count=128, window_size=64 → num_samples = (128-1-64)//64 + 1 = 1
stream_view = H5Store(window_size=seq_length, stride=seq_length)
stream_view = MmapStore(window_size=seq_length, stride=seq_length)
stream_view.load(test_dir)
assert len(stream_view) == 1
+2 -2
View File
@@ -30,7 +30,7 @@ def test_from_dict_flat():
"mask": {"system": "mask", "assistant": "train"},
"mask_default": "mask",
"preprocessing": {"max_seq_len": 1024},
"output": {"storage_format": "h5"},
"output": {"storage_format": "bin"},
}
config = PipelineConfig.from_dict(data)
assert config.input.sections == [
@@ -38,7 +38,7 @@ def test_from_dict_flat():
]
assert config.mask == {"system": "mask", "assistant": "train"}
assert config.preprocessing.max_seq_len == 1024
assert config.output.storage_format == "h5"
assert config.output.storage_format == "bin"
def test_to_dict_roundtrip():
+1 -1
View File
@@ -2,7 +2,7 @@
from unittest.mock import MagicMock
from astrai.inference import STOP, Task, TaskManager, TaskStatus
from astrai.inference import Task, TaskManager, TaskStatus
def _make_mock_tokenizer():
-1
View File
@@ -1,7 +1,6 @@
import json
import os
import pytest
import safetensors.torch as st
import torch
-2
View File
@@ -1,5 +1,3 @@
import os
import pytest
import torch
+1 -1
View File
@@ -3,7 +3,7 @@ import torch
from astrai.model.transformer import AutoRegressiveLM
from astrai.trainer.strategy import GRPOStrategy
from tests.helpers import FakeExecutor, make_frozen, make_model, make_rollout_config
from tests.helpers import FakeExecutor, make_frozen, make_model
def _make_batch(