diff --git a/astrai/dataset/dataset.py b/astrai/dataset/dataset.py index 6bcff35..08a4566 100644 --- a/astrai/dataset/dataset.py +++ b/astrai/dataset/dataset.py @@ -25,20 +25,50 @@ function (pure ``record -> Dict[str, Tensor]``) is forwarded to from abc import ABC, abstractmethod from functools import partial +from pathlib import Path from typing import Callable, Dict, List, Optional import torch from torch import Tensor from torch.utils.data import Dataset +from astrai.config.preprocess_config import PipelineConfig from astrai.dataset.storage import ( Store, StoreFactory, detect_format, ) from astrai.factory import BaseFactory +from astrai.preprocessing.transform import TokenizeTransform 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( record: dict, @@ -349,16 +379,18 @@ class DatasetFactory(BaseFactory["BaseDataset"]): ) if processor is not None: 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: - load_kwargs = dict(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) + store.load(load_path, **kwargs) return cls.create(train_type, store=store) diff --git a/astrai/dataset/storage.py b/astrai/dataset/storage.py index 698f913..51b7bd8 100644 --- a/astrai/dataset/storage.py +++ b/astrai/dataset/storage.py @@ -55,9 +55,7 @@ from typing import Callable, Dict, List, Optional, Tuple, Union import torch from torch import Tensor -from astrai.config.preprocess_config import PipelineConfig from astrai.factory import BaseFactory -from astrai.preprocessing.transform import TokenizeTransform from astrai.serialization import ( load_bin, load_bin_offsets, @@ -536,19 +534,8 @@ class JsonlStore(Store, Streamable, Recordable): ``len(store)`` returns ``num_records``; stream primitives raise. """ - CONFIG_NAME = "dataset_config.json" 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__( self, window_size: int = 0, @@ -569,22 +556,10 @@ class JsonlStore(Store, Streamable, Recordable): return if transform is None: - root = Path(path) - config_path = root / self.CONFIG_NAME if root.is_dir() else None - if config_path is not None and config_path.exists(): - 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) + raise ValueError( + "JsonlStore eager mode requires transform=. " + "Use DatasetFactory.load() which auto-constructs it." + ) transformed = transform.apply(records) self._normalize(transformed) diff --git a/tests/data/test_dataset.py b/tests/data/test_dataset.py index 59300f9..b2b74e7 100644 --- a/tests/data/test_dataset.py +++ b/tests/data/test_dataset.py @@ -9,6 +9,7 @@ import torch from astrai.dataset.dataset import ( DatasetFactory, GRPODataset, + _build_jsonl_transform, dpo_tokenize, grpo_collate_fn, ) @@ -525,7 +526,7 @@ def test_json_store_seq(base_test_env): ) store = StoreFactory.create("jsonl") - store.load(data_dir) + store.load(data_dir, transform=_build_jsonl_transform(data_dir)) assert len(store) > 0 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) store = StoreFactory.create("jsonl") - store.load(data_dir) + store.load(data_dir, transform=_build_jsonl_transform(data_dir)) assert len(store) > 0 assert "sequence" 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.load(data_dir) + store.load(data_dir, transform=_build_jsonl_transform(data_dir)) assert len(store) > 0 assert "sequence" in store.keys @@ -638,7 +639,7 @@ def test_jsonl_store_sft(base_test_env): ) store = StoreFactory.create("jsonl") - store.load(data_dir) + store.load(data_dir, transform=_build_jsonl_transform(data_dir)) assert "sequence" in store.keys assert "loss_mask" 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.load(data_dir) + store.load(data_dir, transform=_build_jsonl_transform(data_dir)) assert store.num_records == 2 assert len(store.keys) > 0