refactor: split mask builder by single/multi output

- Extract SingleOutputMaskBuilder for SFT and pretrain configs
- Extract MultiOutputMaskBuilder for DPO and GRPO configs
- Keep SectionedMaskBuilder as backward-compatible facade
- Register "single" and "multi" names in MaskBuilderFactory
- Add parity and rejection tests for concrete builders
This commit is contained in:
2026-07-08 21:18:34 +08:00
parent c8567a6f65
commit 841a582b28
4 changed files with 130 additions and 39 deletions
+15 -1
View File
@@ -10,7 +10,11 @@ from astrai.config.preprocess_config import (
PipelineConfig,
ProcessingConfig,
)
from astrai.preprocessing.builder import SectionedMaskBuilder
from astrai.preprocessing.builder import (
MultiOutputMaskBuilder,
SectionedMaskBuilder,
SingleOutputMaskBuilder,
)
from astrai.tokenize import AutoTokenizer
_SPECIAL_TOKENS_CONFIG = {
@@ -210,6 +214,16 @@ def builder():
return SectionedMaskBuilder()
@pytest.fixture
def single_builder():
return SingleOutputMaskBuilder()
@pytest.fixture
def multi_builder():
return MultiOutputMaskBuilder()
@pytest.fixture
def tokenizer_dir(temp_dir, test_tokenizer):
d = os.path.join(temp_dir, "tok")
+66 -2
View File
@@ -8,7 +8,9 @@ from astrai.config.preprocess_config import (
)
from astrai.preprocessing.builder import (
MaskBuilderFactory,
MultiOutputMaskBuilder,
SectionedMaskBuilder,
SingleOutputMaskBuilder,
)
from tests.data.conftest import (
_CHAT_SECTIONS,
@@ -272,12 +274,18 @@ def test_sectioned_text_too_short(test_tokenizer, builder):
def test_factory_registered():
names = MaskBuilderFactory.list_registered()
assert "single" in names
assert "multi" in names
assert "sectioned" in names
def test_factory_create():
builder_obj = MaskBuilderFactory.create("sectioned")
assert isinstance(builder_obj, SectionedMaskBuilder)
single = MaskBuilderFactory.create("single")
assert isinstance(single, SingleOutputMaskBuilder)
multi = MaskBuilderFactory.create("multi")
assert isinstance(multi, MultiOutputMaskBuilder)
sectioned = MaskBuilderFactory.create("sectioned")
assert isinstance(sectioned, SectionedMaskBuilder)
def test_dpo_chat_basic(chat_tokenizer, builder):
@@ -367,3 +375,59 @@ def test_grpo_single_reward(chat_tokenizer, builder):
}
result = builder.build(item, config, chat_tokenizer)
assert result["rewards"] == [0.9]
def test_single_builder_matches_facade(chat_tokenizer, builder, single_builder):
config = make_chat_config()
item = {
"messages": [
{"role": "user", "content": "What is 2+2?"},
{"role": "assistant", "content": "4"},
]
}
facade_result = builder.build(item, config, chat_tokenizer)
single_result = single_builder.build(item, config, chat_tokenizer)
assert single_result == facade_result
def test_single_builder_rejects_multi_config(chat_tokenizer, single_builder):
config = make_dpo_chat_config()
item = {
"chosen": [
{"role": "user", "content": "What is 2+2?"},
{"role": "assistant", "content": "4"},
],
"rejected": [
{"role": "user", "content": "What is 2+2?"},
{"role": "assistant", "content": "5"},
],
}
assert single_builder.build(item, config, chat_tokenizer) is None
def test_multi_builder_matches_facade(chat_tokenizer, builder, multi_builder):
config = make_dpo_chat_config()
item = {
"chosen": [
{"role": "user", "content": "What is 2+2?"},
{"role": "assistant", "content": "4"},
],
"rejected": [
{"role": "user", "content": "What is 2+2?"},
{"role": "assistant", "content": "5"},
],
}
facade_result = builder.build(item, config, chat_tokenizer)
multi_result = multi_builder.build(item, config, chat_tokenizer)
assert multi_result == facade_result
def test_multi_builder_rejects_single_config(chat_tokenizer, multi_builder):
config = make_chat_config()
item = {
"messages": [
{"role": "user", "content": "What is 2+2?"},
{"role": "assistant", "content": "4"},
]
}
assert multi_builder.build(item, config, chat_tokenizer) is None