feat: add record-mode to Store for DPO/GRPO

- Store gains fetch_record/num_records alongside stream fetch/__len__
- save_bin/load_bin support per-record offsets via record_keys param
- H5Store/MmapStore/JsonlStore all support dual stream+record access
- DPODataset/GRPODataset use fetch_record, no cross-record concat
- dpo_collate_fn + collate_fn wired through TrainConfig
- fixes attention context leakage in DPO from windowed concatenation
This commit is contained in:
2026-07-18 21:02:29 +08:00
parent 28886e4241
commit a74e5b91a3
12 changed files with 334 additions and 92 deletions
+13 -3
View File
@@ -56,6 +56,13 @@ def _write_jsonl_dataset(test_dir, tokenizer_path, records, config_overrides=Non
return data_dir
def _fake_fetch_record(self, idx, keys):
"""FakeStore.fetch_record matching real Store semantics."""
if isinstance(keys, str):
return self._data[keys][idx]
return {k: self._data[k][idx] for k in keys}
def _make_seq_dataset(
test_dir, name="data", seq_length=200, train_type="seq", data=None, **load_kwargs
):
@@ -338,6 +345,7 @@ def test_grpo_dataset_dtype(base_test_env):
(),
{
"keys": ["prompts", "responses", "masks", "rewards"],
"num_records": 1,
"_data": {
"prompts": [torch.randint(0, 100, (10,), dtype=torch.int32)],
"responses": [
@@ -346,9 +354,9 @@ def test_grpo_dataset_dtype(base_test_env):
"masks": [[torch.ones(5, dtype=torch.int32) for _ in range(G)]],
"rewards": [torch.rand(G, dtype=torch.float32)],
},
"fetch_record": _fake_fetch_record,
},
)()
dataset._build_records()
item = dataset[0]
assert item["prompts"].dtype == torch.long
@@ -371,15 +379,16 @@ def test_grpo_dataset_load(base_test_env):
(),
{
"keys": ["prompts", "responses", "masks", "rewards"],
"num_records": 1,
"_data": {
"prompts": [torch.randint(0, 100, (prompt_len,))],
"responses": [[torch.randint(0, 100, (rl,)) for rl in resp_lens]],
"masks": [[torch.ones(rl, dtype=torch.int64) for rl in resp_lens]],
"rewards": [torch.tensor([0.9, 0.3, 0.7], dtype=torch.float32)],
},
"fetch_record": _fake_fetch_record,
},
)()
dataset._build_records()
assert len(dataset) == 1
item = dataset[0]
@@ -864,6 +873,7 @@ def test_grpo_multiple_records(base_test_env):
(),
{
"keys": ["prompts", "responses", "masks", "rewards"],
"num_records": n_records,
"_data": {
"prompts": [torch.randint(0, 100, (10,)) for _ in range(n_records)],
"responses": dummy_responses,
@@ -875,9 +885,9 @@ def test_grpo_multiple_records(base_test_env):
torch.rand(G, dtype=torch.float32) for _ in range(n_records)
],
},
"fetch_record": _fake_fetch_record,
},
)()
dataset._build_records()
assert len(dataset) == n_records
+6 -6
View File
@@ -1,4 +1,4 @@
from astrai.dataset import ResumableDistributedSampler
from astrai.dataset import RDSampler
def test_random_sampler_consistency(random_dataset):
@@ -6,8 +6,8 @@ def test_random_sampler_consistency(random_dataset):
dataset = random_dataset
# Create two samplers with same seed
sampler1 = ResumableDistributedSampler(dataset, seed=42)
sampler2 = ResumableDistributedSampler(dataset, seed=42)
sampler1 = RDSampler(dataset, seed=42)
sampler2 = RDSampler(dataset, seed=42)
indices1 = list(iter(sampler1))
indices2 = list(iter(sampler2))
@@ -20,8 +20,8 @@ def test_random_sampler_different_seeds(random_dataset):
dataset = random_dataset
# Create two samplers with different seeds
sampler1 = ResumableDistributedSampler(dataset, seed=42)
sampler2 = ResumableDistributedSampler(dataset, seed=123)
sampler1 = RDSampler(dataset, seed=42)
sampler2 = RDSampler(dataset, seed=123)
indices1 = list(iter(sampler1))
indices2 = list(iter(sampler2))
@@ -35,7 +35,7 @@ def test_sampler_across_epochs(random_dataset):
dataset = random_dataset
n = len(dataset)
sampler = ResumableDistributedSampler(dataset, seed=42)
sampler = RDSampler(dataset, seed=42)
# Get indices for first epoch
epoch1_indices = list(iter(sampler))