refactor : Pipeline 去除去重,ids 重命名为 sequence,泛型透传

- 移除 Pipeline 内置去重逻辑及 dedup_signature 工具函数
- 删除 ProcessingConfig.deduplicate 字段
- builder 返回 'sequence' 替代 'ids',与 dataset 层统一
- pipeline 纯透传,泛型处理任意 key 补齐默认值
This commit is contained in:
2026-05-31 15:14:27 +08:00
parent 14f83cbdac
commit 01ce1fb9e3
6 changed files with 32 additions and 52 deletions
+18 -27
View File
@@ -15,7 +15,7 @@ from astrai.preprocessing.builder import (
MaskBuilderFactory,
SectionedMaskBuilder,
)
from astrai.preprocessing.pipeline import Pipeline, dedup_signature, filter_by_length
from astrai.preprocessing.pipeline import Pipeline, filter_by_length
from astrai.tokenize import AutoTokenizer
_SPECIAL_TOKENS_CONFIG = {
@@ -199,16 +199,16 @@ class TestChatMaskBuilder:
}
result = builder.build(item, config, chat_tokenizer)
assert result is not None
assert "ids" in result
assert "sequence" in result
assert "loss_mask" in result
assert len(result["ids"]) == len(result["loss_mask"])
assert len(result["sequence"]) == len(result["loss_mask"])
ids = chat_tokenizer.decode(result["ids"], skip_special_tokens=False)
ids = chat_tokenizer.decode(result["sequence"], skip_special_tokens=False)
assert "system" in ids.lower() or "<|im_start|>system" in ids
assert "assistant" in ids.lower() or "<|im_start|>assistant" in ids
total = len(result["ids"])
total = len(result["sequence"])
trained = sum(result["loss_mask"])
assert trained > 0, "At least assistant tokens should be trained"
assert trained < total, "System and user tokens should be masked"
@@ -224,7 +224,7 @@ class TestChatMaskBuilder:
}
result = builder.build(item, config, chat_tokenizer)
mask = result["loss_mask"]
ids = result["ids"]
ids = result["sequence"]
assert len(ids) == len(mask)
@@ -266,7 +266,7 @@ class TestChatMaskBuilder:
]
}
result = builder.build(item, config, chat_tokenizer)
assert sum(result["loss_mask"]) == len(result["ids"]) - 1
assert sum(result["loss_mask"]) == len(result["sequence"]) - 1
def test_empty_messages_returns_none(self, chat_tokenizer):
config = make_chat_config()
@@ -311,8 +311,8 @@ class TestChatMaskBuilder:
]
}
result = builder.build(item, config, chat_tokenizer)
assert len(result["ids"]) <= 10
assert len(result["loss_mask"]) == len(result["ids"])
assert len(result["sequence"]) <= 10
assert len(result["loss_mask"]) == len(result["sequence"])
class TestInstructionMaskBuilder:
@@ -322,7 +322,7 @@ class TestInstructionMaskBuilder:
item = {"prompt": "Translate to French: Hello", "response": "Bonjour"}
result = builder.build(item, config, test_tokenizer)
assert result is not None
assert len(result["ids"]) == len(result["loss_mask"])
assert len(result["sequence"]) == len(result["loss_mask"])
def test_prompt_masked_response_trained(self, test_tokenizer):
config = make_instruction_config()
@@ -330,7 +330,7 @@ class TestInstructionMaskBuilder:
item = {"prompt": "hello", "response": "world"}
result = builder.build(item, config, test_tokenizer)
mask = result["loss_mask"]
ids = result["ids"]
ids = result["sequence"]
prompt_ids = test_tokenizer.encode("hello", add_special_tokens=True)
response_ids = test_tokenizer.encode("world", add_special_tokens=False)
@@ -359,7 +359,7 @@ class TestInstructionMaskBuilder:
item = {"prompt": "hello", "response": "world"}
result = builder.build(item, config, test_tokenizer)
mask = result["loss_mask"]
ids = result["ids"]
ids = result["sequence"]
prompt_ids = test_tokenizer.encode("hello", add_special_tokens=True)
p_len = min(len(prompt_ids), len(ids))
@@ -373,8 +373,8 @@ class TestTextMaskBuilder:
item = {"text": "Hello world. This is a test document."}
result = builder.build(item, config, test_tokenizer)
assert result is not None
assert "ids" in result
assert len(result["ids"]) > 0
assert "sequence" in result
assert len(result["sequence"]) > 0
assert "loss_mask" not in result
def test_empty_text_returns_none(self, test_tokenizer):
@@ -399,7 +399,7 @@ class TestTextMaskBuilder:
builder = SectionedMaskBuilder()
item = {"text": "This is a very long text that should be truncated"}
result = builder.build(item, config, test_tokenizer)
assert len(result["ids"]) <= 3
assert len(result["sequence"]) <= 3
class TestPipeline:
@@ -446,7 +446,7 @@ class TestPipeline:
input=InputConfig(sections=_CHAT_SECTIONS),
mask={"system": "mask", "user": "mask", "assistant": "train"},
mask_default="mask",
preprocessing=ProcessingConfig(max_seq_len=2048, deduplicate=True),
preprocessing=ProcessingConfig(max_seq_len=2048),
output=OutputConfig(storage_format="bin", domain_key=None),
)
@@ -505,9 +505,7 @@ class TestPipeline:
config = PipelineConfig(
input=InputConfig(sections=_TEXT_SECTIONS),
preprocessing=ProcessingConfig(
max_seq_len=2048, min_chars=10, deduplicate=True
),
preprocessing=ProcessingConfig(max_seq_len=2048, min_chars=10),
output=OutputConfig(storage_format="bin"),
)
@@ -648,13 +646,6 @@ class TestUtility:
assert not filter_by_length("x" * 100, max_len=50)
assert filter_by_length("just right", min_len=5, max_len=20)
def test_dedup_signature(self):
a = {"key": "value", "number": 1}
b = {"number": 1, "key": "value"}
assert dedup_signature(a) == dedup_signature(b)
c = {"key": "different"}
assert dedup_signature(a) != dedup_signature(c)
class TestSectionedMaskBuilder:
def test_sectioned_chat(self, chat_tokenizer):
@@ -673,7 +664,7 @@ class TestSectionedMaskBuilder:
}
result = builder.build(item, config, chat_tokenizer)
assert result is not None
assert len(result["ids"]) == len(result["loss_mask"])
assert len(result["sequence"]) == len(result["loss_mask"])
assert sum(result["loss_mask"]) > 0
assert 0 in result["loss_mask"]