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:
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user