refactor: split Store into StreamStore and RecordStore

- StreamStore: fetch(begin, end, key) for stream access (SEQ/SFT)
- RecordStore: mixin with fetch_record(i, key) for record access
- H5Store/MmapStore/JsonlStore now dual-inherit both (C3 MRO)
- JsonlStore supports lazy mode via processor= (no TokenizeTransform)
- RecordDataset base class holds processor, DPO/GRPO simplified
- dpo_tokenize pure function for on-the-fly JSONL tokenisation
- DatasetFactory builds processor for jsonl+record datasets
- train.py passes tokenizer_path=param_path uniformly
- progress: len(dataset) returns sample count (stream=windows, record=records)
- json no longer auto-detected as jsonl format
This commit is contained in:
2026-07-18 23:04:31 +08:00
parent b33250dc28
commit b133fc9c07
5 changed files with 461 additions and 277 deletions
+200 -64
View File
@@ -1,18 +1,93 @@
"""Dataset implementations with factory pattern for training."""
from abc import ABC, abstractmethod
from typing import Dict, List, Optional
from functools import partial
from typing import Callable, Dict, List, Optional
import torch
from torch import Tensor
from torch.utils.data import Dataset
from astrai.dataset.storage import (
RecordStore,
Store,
StoreFactory,
StreamStore,
detect_format,
)
from astrai.factory import BaseFactory
from astrai.tokenize import AutoTokenizer
def _to_tensor(value: list, dtype: Optional[torch.dtype] = None) -> Tensor:
if dtype is not None:
return torch.tensor(value, dtype=dtype)
if value and isinstance(value[0], bool):
return torch.tensor(value, dtype=torch.bool)
return torch.tensor(value, dtype=torch.int32)
def dpo_tokenize(
record: dict,
tokenizer,
max_len: int = 2048,
pad_id: int = 2,
) -> Optional[dict]:
"""Tokenize one DPO record into chosen/rejected + masks.
Pure processor function (HF ``datasets.map`` style):
``record -> dict_of_lists``. Each value is a flat list of ints/bools.
No packing, no ``position_ids`` — DPO sequences are independent and
the model defaults to ``arange(0, seq_len)``.
"""
inp = record.get("input")
chosen_text = record.get("chosen")
rejected_text = record.get("rejected")
if inp is None or chosen_text is None or rejected_text is None:
return None
in_ids = tokenizer.encode(inp, add_special_tokens=True)
ch_ids = tokenizer.encode(chosen_text, add_special_tokens=False)
re_ids = tokenizer.encode(rejected_text, add_special_tokens=False)
full_ch = (in_ids + ch_ids)[:max_len]
full_re = (in_ids + re_ids)[:max_len]
max_record_len = max(len(full_ch), len(full_re))
ch_pad = full_ch + [pad_id] * (max_record_len - len(full_ch))
re_pad = full_re + [pad_id] * (max_record_len - len(full_re))
ch_mask = [0] * len(in_ids) + [1] * len(ch_ids)
ch_mask = ch_mask[:max_len]
ch_mask += [0] * (max_record_len - len(ch_mask))
re_mask = [0] * len(in_ids) + [1] * len(re_ids)
re_mask = re_mask[:max_len]
re_mask += [0] * (max_record_len - len(re_mask))
return {
"chosen": ch_pad,
"rejected": re_pad,
"chosen_mask": ch_mask,
"rejected_mask": re_mask,
}
def dpo_processor(
record: dict,
tokenizer,
max_len: int = 2048,
) -> Dict[str, Tensor]:
"""DPO processor: wraps :func:`dpo_tokenize` and returns tensors."""
result = dpo_tokenize(record, tokenizer, max_len=max_len)
if result is None:
raise ValueError(f"Malformed DPO record: {list(record.keys())}")
return {
"chosen": torch.tensor(result["chosen"], dtype=torch.int32),
"rejected": torch.tensor(result["rejected"], dtype=torch.int32),
"chosen_mask": torch.tensor(result["chosen_mask"], dtype=torch.bool),
"rejected_mask": torch.tensor(result["rejected_mask"], dtype=torch.bool),
}
def dpo_collate_fn(batch: List[Dict[str, Tensor]]) -> Dict[str, Tensor]:
@@ -206,6 +281,63 @@ class BaseDataset(Dataset, ABC):
return (total - 1 - self.window_size) // self.stride + 1
class RecordDataset(BaseDataset):
"""Base class for record-structured datasets (DPO/GRPO).
Each sample is an independent record — no windowing, stride, or
cross-record concatenation. ``__len__`` returns the record count
so progress bars advance per-record.
A *processor* (pure ``record -> Dict[str, Tensor]`` function) may be
supplied for lazy on-the-fly tokenisation of raw JSONL. The
processor is forwarded to ``JsonlStore`` and applied per access;
pre-tokenised backends (H5/bin) ignore it.
"""
def __init__(
self,
window_size: int = 0,
stride: int = 0,
processor: Optional[Callable[[dict], Dict[str, Tensor]]] = None,
**kwargs,
):
super().__init__(window_size=window_size, stride=stride or window_size)
self.processor = processor
def load(self, load_path: str, storage_type: Optional[str] = None, **kwargs):
"""Load data from *load_path*.
Args:
load_path: Path to data file or directory.
storage_type: Force backend ("h5"/"bin"/"jsonl") or None for
auto-detection.
**kwargs: Forwarded to ``store.load()``. When the backend is
JSONL and a processor was set, it is passed as
``processor=`` for lazy tokenisation.
"""
if storage_type is None:
storage_type = detect_format(load_path)
self.storage = StoreFactory.create(storage_type, **kwargs)
self._load_path = load_path
if self.processor is not None:
self.storage.load(load_path, processor=self.processor, **kwargs)
else:
self.storage.load(load_path, **kwargs)
self._validate_keys()
def __len__(self) -> int:
if self.storage is None:
return 0
return self.storage.num_records
@property
def count(self) -> int:
if self.storage is None:
return 0
return self.storage.num_records
class DatasetFactory(BaseFactory["BaseDataset"]):
"""Factory class for creating dataset instances.
@@ -229,6 +361,8 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
window_size: int,
stride: Optional[int] = None,
storage_type: Optional[str] = None,
tokenizer_path: Optional[str] = None,
max_len: int = 2048,
**kwargs,
) -> "BaseDataset":
"""Create and load a dataset in one step.
@@ -239,6 +373,11 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
window_size: Window size for data sampling
stride: Stride between consecutive samples (default: same as window_size)
storage_type: Storage type ("h5", "bin", "jsonl") or None for auto-detection
tokenizer_path: Path to tokenizer. Used to build an on-the-fly
processor when loading raw JSONL with a record dataset
(DPO/GRPO). Ignored for pre-tokenised backends (H5/bin)
and for stream datasets (SEQ/SFT).
max_len: Max sequence length for the processor.
**kwargs: Extra arguments forwarded to ``dataset.load()``.
Returns:
@@ -247,11 +386,57 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
if stride is None:
stride = window_size
dataset = cls.create(train_type, window_size, stride)
if storage_type is None:
storage_type = detect_format(load_path)
processor = cls._maybe_build_processor(
train_type, storage_type, tokenizer_path, max_len
)
dataset = cls.create(train_type, window_size, stride, processor=processor)
dataset.load(load_path, storage_type=storage_type, **kwargs)
return dataset
@classmethod
def from_store(
cls,
train_type: str,
store: Store,
window_size: int = 0,
stride: Optional[int] = None,
) -> "BaseDataset":
"""Create a dataset bound to an already-loaded store.
The caller is responsible for constructing and loading the store
(including any processor). The dataset simply wraps it.
"""
if stride is None:
stride = window_size
dataset = cls.create(train_type, window_size, stride)
dataset.storage = store
return dataset
@staticmethod
def _maybe_build_processor(
train_type: str,
storage_type: str,
tokenizer_path: Optional[str],
max_len: int,
) -> Optional[Callable[[dict], Dict[str, Tensor]]]:
"""Build an on-the-fly tokenisation processor if applicable.
Only raw JSONL + record datasets (DPO/GRPO) need a processor;
pre-tokenised backends (H5/bin) and stream datasets (SEQ/SFT)
return ``None`` so no tokenizer is loaded.
"""
if tokenizer_path is None or storage_type != "jsonl":
return None
if train_type == "dpo":
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
return None
@DatasetFactory.register("seq")
class SEQDataset(BaseDataset):
@@ -301,7 +486,7 @@ class SFTDataset(BaseDataset):
@DatasetFactory.register("dpo")
class DPODataset(BaseDataset):
class DPODataset(RecordDataset):
"""Record-structured dataset for Direct Preference Optimization.
Each sample is one preference pair (chosen + rejected) and is an
@@ -309,41 +494,21 @@ class DPODataset(BaseDataset):
concatenation. This keeps each sequence self-contained so attention
never leaks across preference pairs.
Delegates record access to ``Store.fetch_record``, which works with
any storage backend (H5 per-record datasets, bin+offsets memmap, or
JSONL on-the-fly tokenization).
"""
Two loading paths (handled by :class:`RecordDataset`):
def __init__(self, window_size: int = 0, stride: int = 0, **kwargs):
super().__init__(window_size=window_size, stride=stride or window_size)
- **Pre-tokenized** (H5/bin): ``load(path)`` reads per-record tensors,
``__getitem__`` returns them directly.
- **Raw JSONL** (``tokenizer_path=...``): builds a lazy processor via
:func:`dpo_processor` that tokenises on the fly — no packing, no
``position_ids``.
"""
@property
def required_keys(self) -> List[str]:
return ["chosen", "rejected", "chosen_mask", "rejected_mask"]
def load(self, load_path: str, storage_type: Optional[str] = None, **kwargs):
if storage_type is None:
storage_type = detect_format(load_path)
self.storage = StoreFactory.create(storage_type, **kwargs)
self._load_path = load_path
self.storage.load(load_path, **kwargs)
self._validate_keys()
def _validate_keys(self):
actual_keys = set(self.storage.keys)
missing = [k for k in self.required_keys if k not in actual_keys]
if missing:
raise KeyError(
f"DPODataset requires keys {self.required_keys}, "
f"but storage only has {sorted(actual_keys)}. Missing: {missing}"
)
@property
def count(self) -> int:
return self.storage.num_records
def __len__(self) -> int:
return self.storage.num_records
def make_processor(self, tokenizer, max_len: int):
return partial(dpo_processor, tokenizer=tokenizer, max_len=max_len)
def __getitem__(self, index: int) -> Dict[str, Tensor]:
return {
@@ -361,13 +526,11 @@ class DPODataset(BaseDataset):
@DatasetFactory.register("grpo")
class GRPODataset(BaseDataset):
class GRPODataset(RecordDataset):
"""Dataset for offline Group Relative Policy Optimization.
Unlike the window-based datasets (SEQ/SFT/DPO), GRPO data is
record-structured: each sample is one prompt with its group of
responses and scalar rewards. There is no windowing or stride —
every record is an independent training unit.
Each sample is one prompt with its group of responses and scalar
rewards — an independent training unit with no windowing or stride.
Expected storage layout (produced by JsonlStore or pre-tokenized):
@@ -377,37 +540,10 @@ class GRPODataset(BaseDataset):
- ``rewards``: List[Tensor] — one 1-D float tensor (len G) per record
"""
def __init__(self, window_size: int = 0, stride: int = 0, **kwargs):
super().__init__(window_size=window_size, stride=stride or window_size)
@property
def required_keys(self) -> List[str]:
return ["prompts", "responses", "masks", "rewards"]
def load(self, load_path: str, storage_type: Optional[str] = None, **kwargs):
if storage_type is None:
storage_type = detect_format(load_path)
self.storage = StoreFactory.create(storage_type, **kwargs)
self._load_path = load_path
self.storage.load(load_path, **kwargs)
self._validate_keys()
def _validate_keys(self):
actual_keys = set(self.storage.keys)
missing = [k for k in self.required_keys if k not in actual_keys]
if missing:
raise KeyError(
f"GRPODataset requires keys {self.required_keys}, "
f"but storage only has {sorted(actual_keys)}. Missing: {missing}"
)
@property
def count(self) -> int:
return self.storage.num_records
def __len__(self) -> int:
return self.storage.num_records
def __getitem__(self, index: int) -> Dict[str, Tensor]:
prompts = self.storage.fetch_record(index, "prompts")
responses = self.storage.fetch_record(index, "responses")