refactor: break JsonlStore→preprocessing circular dependency
- move JSONL transform auto-creation from JsonlStore.load to DatasetFactory.load via _build_jsonl_transform helper - remove TokenizeTransform and PipelineConfig imports from storage module - JsonlStore.load now requires explicit transform= for eager mode - DatasetFactory.load remains the public API with identical convenience behavior
This commit is contained in:
@@ -25,20 +25,50 @@ function (pure ``record -> Dict[str, Tensor]``) is forwarded to
|
|||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from functools import partial
|
from functools import partial
|
||||||
|
from pathlib import Path
|
||||||
from typing import Callable, Dict, List, Optional
|
from typing import Callable, Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from torch.utils.data import Dataset
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
|
from astrai.config.preprocess_config import PipelineConfig
|
||||||
from astrai.dataset.storage import (
|
from astrai.dataset.storage import (
|
||||||
Store,
|
Store,
|
||||||
StoreFactory,
|
StoreFactory,
|
||||||
detect_format,
|
detect_format,
|
||||||
)
|
)
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.preprocessing.transform import TokenizeTransform
|
||||||
from astrai.tokenize import AutoTokenizer
|
from astrai.tokenize import AutoTokenizer
|
||||||
|
|
||||||
|
_DEFAULT_MESSAGES_CONFIG = {
|
||||||
|
"version": 1,
|
||||||
|
"input": {"sections": [{"field": "messages", "action": "$role", "template": True}]},
|
||||||
|
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
|
||||||
|
"mask_default": "mask",
|
||||||
|
"output": {"position_ids_mode": "doc_reset"},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _build_jsonl_transform(
|
||||||
|
path: str, tokenizer_path: Optional[str] = None
|
||||||
|
) -> Optional["TokenizeTransform"]:
|
||||||
|
"""Auto-build a TokenizeTransform for JSONL eager loading.
|
||||||
|
|
||||||
|
Reads ``dataset_config.json`` from the data dir if present, or
|
||||||
|
falls back to the built-in chatml SFT config when *tokenizer_path*
|
||||||
|
is provided.
|
||||||
|
"""
|
||||||
|
root = Path(path)
|
||||||
|
config_path = root / "dataset_config.json" if root.is_dir() else None
|
||||||
|
if config_path is not None and config_path.exists():
|
||||||
|
return TokenizeTransform.from_config_file(str(config_path))
|
||||||
|
if tokenizer_path:
|
||||||
|
config = PipelineConfig.from_dict(_DEFAULT_MESSAGES_CONFIG)
|
||||||
|
return TokenizeTransform(config, tokenizer_path)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
def dpo_tokenize(
|
def dpo_tokenize(
|
||||||
record: dict,
|
record: dict,
|
||||||
@@ -349,16 +379,18 @@ class DatasetFactory(BaseFactory["BaseDataset"]):
|
|||||||
)
|
)
|
||||||
if processor is not None:
|
if processor is not None:
|
||||||
store.load(load_path, processor=processor, **kwargs)
|
store.load(load_path, processor=processor, **kwargs)
|
||||||
|
elif storage_type == "jsonl":
|
||||||
|
transform = _build_jsonl_transform(load_path, tokenizer_path)
|
||||||
|
if transform is None:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"JSONL dataset config not found. Expected "
|
||||||
|
f"dataset_config.json alongside *.jsonl files, pass "
|
||||||
|
f"tokenizer_path= for the built-in messages config, or "
|
||||||
|
f"use processor= for lazy on-the-fly tokenisation."
|
||||||
|
)
|
||||||
|
store.load(load_path, transform=transform, **kwargs)
|
||||||
else:
|
else:
|
||||||
load_kwargs = dict(kwargs)
|
store.load(load_path, **kwargs)
|
||||||
if (
|
|
||||||
tokenizer_path is not None
|
|
||||||
and storage_type == "jsonl"
|
|
||||||
and train_type in ("seq", "sft")
|
|
||||||
and "tokenizer_path" not in load_kwargs
|
|
||||||
):
|
|
||||||
load_kwargs["tokenizer_path"] = tokenizer_path
|
|
||||||
store.load(load_path, **load_kwargs)
|
|
||||||
|
|
||||||
return cls.create(train_type, store=store)
|
return cls.create(train_type, store=store)
|
||||||
|
|
||||||
|
|||||||
@@ -55,9 +55,7 @@ from typing import Callable, Dict, List, Optional, Tuple, Union
|
|||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.config.preprocess_config import PipelineConfig
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.preprocessing.transform import TokenizeTransform
|
|
||||||
from astrai.serialization import (
|
from astrai.serialization import (
|
||||||
load_bin,
|
load_bin,
|
||||||
load_bin_offsets,
|
load_bin_offsets,
|
||||||
@@ -536,19 +534,8 @@ class JsonlStore(Store, Streamable, Recordable):
|
|||||||
``len(store)`` returns ``num_records``; stream primitives raise.
|
``len(store)`` returns ``num_records``; stream primitives raise.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
CONFIG_NAME = "dataset_config.json"
|
|
||||||
segments_are_records = True
|
segments_are_records = True
|
||||||
|
|
||||||
_DEFAULT_MESSAGES_CONFIG = {
|
|
||||||
"version": 1,
|
|
||||||
"input": {
|
|
||||||
"sections": [{"field": "messages", "action": "$role", "template": True}]
|
|
||||||
},
|
|
||||||
"mask": {"system": "mask", "user": "mask", "assistant": "train"},
|
|
||||||
"mask_default": "mask",
|
|
||||||
"output": {"position_ids_mode": "doc_reset"},
|
|
||||||
}
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
window_size: int = 0,
|
window_size: int = 0,
|
||||||
@@ -569,22 +556,10 @@ class JsonlStore(Store, Streamable, Recordable):
|
|||||||
return
|
return
|
||||||
|
|
||||||
if transform is None:
|
if transform is None:
|
||||||
root = Path(path)
|
raise ValueError(
|
||||||
config_path = root / self.CONFIG_NAME if root.is_dir() else None
|
"JsonlStore eager mode requires transform=. "
|
||||||
if config_path is not None and config_path.exists():
|
"Use DatasetFactory.load() which auto-constructs it."
|
||||||
transform = TokenizeTransform.from_config_file(str(config_path))
|
)
|
||||||
else:
|
|
||||||
tokenizer_path = kwargs.get("tokenizer_path")
|
|
||||||
if not tokenizer_path:
|
|
||||||
raise FileNotFoundError(
|
|
||||||
f"JSONL dataset config not found. Expected "
|
|
||||||
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
|
|
||||||
f"explicit transform, pass processor= for lazy "
|
|
||||||
f"on-the-fly tokenisation, or pass tokenizer_path= to "
|
|
||||||
f"use the built-in messages config."
|
|
||||||
)
|
|
||||||
config = PipelineConfig.from_dict(self._DEFAULT_MESSAGES_CONFIG)
|
|
||||||
transform = TokenizeTransform(config, tokenizer_path)
|
|
||||||
|
|
||||||
transformed = transform.apply(records)
|
transformed = transform.apply(records)
|
||||||
self._normalize(transformed)
|
self._normalize(transformed)
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import torch
|
|||||||
from astrai.dataset.dataset import (
|
from astrai.dataset.dataset import (
|
||||||
DatasetFactory,
|
DatasetFactory,
|
||||||
GRPODataset,
|
GRPODataset,
|
||||||
|
_build_jsonl_transform,
|
||||||
dpo_tokenize,
|
dpo_tokenize,
|
||||||
grpo_collate_fn,
|
grpo_collate_fn,
|
||||||
)
|
)
|
||||||
@@ -525,7 +526,7 @@ def test_json_store_seq(base_test_env):
|
|||||||
)
|
)
|
||||||
|
|
||||||
store = StoreFactory.create("jsonl")
|
store = StoreFactory.create("jsonl")
|
||||||
store.load(data_dir)
|
store.load(data_dir, transform=_build_jsonl_transform(data_dir))
|
||||||
assert len(store) > 0
|
assert len(store) > 0
|
||||||
assert "sequence" in store.keys
|
assert "sequence" in store.keys
|
||||||
|
|
||||||
@@ -580,7 +581,7 @@ def test_json_store_no_tokenizer_path(base_test_env):
|
|||||||
json.dump(config, f, ensure_ascii=False, indent=2)
|
json.dump(config, f, ensure_ascii=False, indent=2)
|
||||||
|
|
||||||
store = StoreFactory.create("jsonl")
|
store = StoreFactory.create("jsonl")
|
||||||
store.load(data_dir)
|
store.load(data_dir, transform=_build_jsonl_transform(data_dir))
|
||||||
assert len(store) > 0
|
assert len(store) > 0
|
||||||
assert "sequence" in store.keys
|
assert "sequence" in store.keys
|
||||||
assert "loss_mask" in store.keys
|
assert "loss_mask" in store.keys
|
||||||
@@ -597,7 +598,7 @@ def test_jsonl_store_seq(base_test_env):
|
|||||||
)
|
)
|
||||||
|
|
||||||
store = StoreFactory.create("jsonl")
|
store = StoreFactory.create("jsonl")
|
||||||
store.load(data_dir)
|
store.load(data_dir, transform=_build_jsonl_transform(data_dir))
|
||||||
assert len(store) > 0
|
assert len(store) > 0
|
||||||
assert "sequence" in store.keys
|
assert "sequence" in store.keys
|
||||||
|
|
||||||
@@ -638,7 +639,7 @@ def test_jsonl_store_sft(base_test_env):
|
|||||||
)
|
)
|
||||||
|
|
||||||
store = StoreFactory.create("jsonl")
|
store = StoreFactory.create("jsonl")
|
||||||
store.load(data_dir)
|
store.load(data_dir, transform=_build_jsonl_transform(data_dir))
|
||||||
assert "sequence" in store.keys
|
assert "sequence" in store.keys
|
||||||
assert "loss_mask" in store.keys
|
assert "loss_mask" in store.keys
|
||||||
assert "position_ids" in store.keys
|
assert "position_ids" in store.keys
|
||||||
@@ -1085,7 +1086,7 @@ def test_jsonl_store_eager_len_returns_token_count(base_test_env):
|
|||||||
)
|
)
|
||||||
|
|
||||||
store = JsonlStore()
|
store = JsonlStore()
|
||||||
store.load(data_dir)
|
store.load(data_dir, transform=_build_jsonl_transform(data_dir))
|
||||||
|
|
||||||
assert store.num_records == 2
|
assert store.num_records == 2
|
||||||
assert len(store.keys) > 0
|
assert len(store.keys) > 0
|
||||||
|
|||||||
Reference in New Issue
Block a user