feat: support SFT directly from JSONL without dataset_config.json
- JsonlStore falls back to built-in messages config when no config file found and tokenizer_path is provided - DatasetFactory.load forwards tokenizer_path to store for SFT/SEQ+jsonl - assistant turns train, other roles masked, position_ids doc_reset
This commit is contained in:
@@ -346,7 +346,15 @@ 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)
|
||||||
else:
|
else:
|
||||||
store.load(load_path, **kwargs)
|
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)
|
||||||
|
|
||||||
return cls.create(train_type, store=store)
|
return cls.create(train_type, store=store)
|
||||||
|
|
||||||
|
|||||||
+42
-14
@@ -56,6 +56,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.preprocessing.transform import TokenizeTransform
|
||||||
from astrai.serialization import (
|
from astrai.serialization import (
|
||||||
@@ -545,18 +546,29 @@ class JsonlSource:
|
|||||||
|
|
||||||
@StoreFactory.register("jsonl")
|
@StoreFactory.register("jsonl")
|
||||||
class JsonlStore(Store, Streamable, Recordable):
|
class JsonlStore(Store, Streamable, Recordable):
|
||||||
"""JSONL reader with two tokenisation modes.
|
"""JSONL reader with eager/lazy tokenisation modes.
|
||||||
|
|
||||||
A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl``
|
A JSONL dataset is a ``.jsonl`` file or a directory of ``*.jsonl``
|
||||||
files plus (optionally) a ``dataset_config.json`` describing the
|
files plus (optionally) a ``dataset_config.json`` describing the
|
||||||
tokenization pipeline.
|
tokenization pipeline.
|
||||||
|
|
||||||
Two modes, selected at :meth:`load` time:
|
Three ways to supply an eager transform (first match wins):
|
||||||
|
|
||||||
- **Eager** (default): applies a :class:`TokenizeTransform` to every
|
- **Explicit** (``transform=``): caller-built
|
||||||
record at load time and registers per-key tensors via
|
:class:`TokenizeTransform` applied eagerly.
|
||||||
``_normalize``. Both ``fetch`` (stream) and ``fetch_record``
|
- **Config file**: ``dataset_config.json`` alongside the ``*.jsonl``
|
||||||
(record) work.
|
files — loaded via :meth:`TokenizeTransform.from_config_file`.
|
||||||
|
- **Default messages** (``tokenizer_path=`` given, no config file):
|
||||||
|
a built-in chatml config that tokenises the ``messages`` field,
|
||||||
|
masking every role except ``assistant`` (loss on assistant only).
|
||||||
|
Lets SFT/SEQ train straight from a chat-style JSONL directory
|
||||||
|
without a hand-written config.
|
||||||
|
|
||||||
|
Two tokenisation modes, selected at :meth:`load` time:
|
||||||
|
|
||||||
|
- **Eager** (default): applies the transform to every record at load
|
||||||
|
time and registers per-key tensors via ``_normalize``. Both
|
||||||
|
``fetch`` (stream) and ``fetch_record`` (record) work.
|
||||||
- **Lazy** (``processor=fn`` passed): keeps raw records and defers
|
- **Lazy** (``processor=fn`` passed): keeps raw records and defers
|
||||||
tokenisation to ``fetch_record``. Only record access works —
|
tokenisation to ``fetch_record``. Only record access works —
|
||||||
``len(store)`` returns ``num_records``; stream primitives raise.
|
``len(store)`` returns ``num_records``; stream primitives raise.
|
||||||
@@ -565,6 +577,16 @@ class JsonlStore(Store, Streamable, Recordable):
|
|||||||
CONFIG_NAME = "dataset_config.json"
|
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,
|
||||||
@@ -587,14 +609,20 @@ class JsonlStore(Store, Streamable, Recordable):
|
|||||||
if transform is None:
|
if transform is None:
|
||||||
root = Path(path)
|
root = Path(path)
|
||||||
config_path = root / self.CONFIG_NAME if root.is_dir() else None
|
config_path = root / self.CONFIG_NAME if root.is_dir() else None
|
||||||
if config_path is None or not config_path.exists():
|
if config_path is not None and config_path.exists():
|
||||||
raise FileNotFoundError(
|
transform = TokenizeTransform.from_config_file(str(config_path))
|
||||||
f"JSONL dataset config not found. Expected "
|
else:
|
||||||
f"{self.CONFIG_NAME} alongside *.jsonl files, pass an "
|
tokenizer_path = kwargs.get("tokenizer_path")
|
||||||
f"explicit transform, or pass processor= for lazy "
|
if not tokenizer_path:
|
||||||
f"on-the-fly tokenisation."
|
raise FileNotFoundError(
|
||||||
)
|
f"JSONL dataset config not found. Expected "
|
||||||
transform = TokenizeTransform.from_config_file(str(config_path))
|
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)
|
||||||
|
|||||||
@@ -654,6 +654,92 @@ def test_jsonl_store_sft(base_test_env):
|
|||||||
assert item["loss_mask"].dtype == torch.bool
|
assert item["loss_mask"].dtype == torch.bool
|
||||||
|
|
||||||
|
|
||||||
|
def test_sft_jsonl_default_messages_config(base_test_env):
|
||||||
|
"""SFT loads a chat-style JSONL dir with no dataset_config.json.
|
||||||
|
|
||||||
|
Falls back to the built-in messages config: every role except
|
||||||
|
``assistant`` is masked, loss on assistant only.
|
||||||
|
"""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
tokenizer = base_test_env["tokenizer"]
|
||||||
|
tokenizer.set_chat_template(
|
||||||
|
"{% for message in messages %}{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
|
||||||
|
)
|
||||||
|
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
|
||||||
|
|
||||||
|
data_dir = os.path.join(test_dir, "jsonl_data")
|
||||||
|
os.makedirs(data_dir, exist_ok=True)
|
||||||
|
records = [
|
||||||
|
{
|
||||||
|
"messages": [
|
||||||
|
{"role": "system", "content": "sys"},
|
||||||
|
{"role": "user", "content": "hi"},
|
||||||
|
{"role": "assistant", "content": "hello"},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"messages": [
|
||||||
|
{"role": "user", "content": "bye"},
|
||||||
|
{"role": "assistant", "content": "see you"},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
]
|
||||||
|
with open(os.path.join(data_dir, "data.jsonl"), "w", encoding="utf-8") as f:
|
||||||
|
for record in records:
|
||||||
|
f.write(json.dumps(record, ensure_ascii=False) + "\n")
|
||||||
|
|
||||||
|
dataset = DatasetFactory.load(
|
||||||
|
"sft", data_dir, window_size=8, tokenizer_path=tokenizer_path
|
||||||
|
)
|
||||||
|
assert "sequence" in dataset.keys
|
||||||
|
assert "loss_mask" in dataset.keys
|
||||||
|
assert "position_ids" in dataset.keys
|
||||||
|
assert len(dataset) > 0
|
||||||
|
item = dataset[0]
|
||||||
|
assert "input_ids" in item
|
||||||
|
assert "target_ids" in item
|
||||||
|
assert "loss_mask" in item
|
||||||
|
assert "position_ids" in item
|
||||||
|
assert item["loss_mask"].dtype == torch.bool
|
||||||
|
|
||||||
|
|
||||||
|
def test_sft_jsonl_explicit_config_takes_priority(base_test_env):
|
||||||
|
"""When dataset_config.json exists, it overrides the default messages config."""
|
||||||
|
test_dir = base_test_env["test_dir"]
|
||||||
|
tokenizer = base_test_env["tokenizer"]
|
||||||
|
tokenizer.set_chat_template(
|
||||||
|
"{% for message in messages %}{{ message['role'] }}:{{ message['content'] }}\n{% endfor %}"
|
||||||
|
)
|
||||||
|
tokenizer_path = _save_test_tokenizer(test_dir, tokenizer)
|
||||||
|
|
||||||
|
data_dir = _write_jsonl_dataset(
|
||||||
|
test_dir,
|
||||||
|
tokenizer_path,
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"messages": [
|
||||||
|
{"role": "user", "content": "hi"},
|
||||||
|
{"role": "assistant", "content": "hello"},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
],
|
||||||
|
config_overrides={
|
||||||
|
"input": {
|
||||||
|
"sections": [{"field": "messages", "action": "$role", "template": True}]
|
||||||
|
},
|
||||||
|
"mask": {"user": "mask", "assistant": "train"},
|
||||||
|
"mask_default": "mask",
|
||||||
|
"preprocessing": {"max_seq_len": 128},
|
||||||
|
"output": {"position_ids_mode": "doc_reset"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
dataset = DatasetFactory.load(
|
||||||
|
"sft", data_dir, window_size=8, tokenizer_path=tokenizer_path
|
||||||
|
)
|
||||||
|
assert "sequence" in dataset.keys
|
||||||
|
assert "loss_mask" in dataset.keys
|
||||||
|
|
||||||
|
|
||||||
def test_jsonl_store_pipeline_config_roundtrip(base_test_env):
|
def test_jsonl_store_pipeline_config_roundtrip(base_test_env):
|
||||||
test_dir = base_test_env["test_dir"]
|
test_dir = base_test_env["test_dir"]
|
||||||
config_path = os.path.join(test_dir, "dataset_config.json")
|
config_path = os.path.join(test_dir, "dataset_config.json")
|
||||||
|
|||||||
Reference in New Issue
Block a user