reafactor: 重构项目

This commit is contained in:
2026-03-30 20:58:51 +08:00
parent f67bad0d8b
commit 35963bcb08
29 changed files with 1395 additions and 1234 deletions
+118
View File
@@ -0,0 +1,118 @@
"""Tests for strategy module."""
import pytest
from pipeline.strategies import (
PromptStrategy,
ChatMLStrategy,
AlpacaStrategy,
StrategyFactory,
)
from pipeline.tokenizer import BpeTokenizer
class DummyTokenizer:
def encode(self, text: str, add_special_tokens: bool = False):
return [ord(c) for c in text]
class DummyStrategy(PromptStrategy):
def __init__(self, tokenizer):
super().__init__(tokenizer)
@property
def name(self) -> str:
return "dummy"
def assemble_prompt(self, query_tokens):
prefix = self._encode_format("Q: ")
return prefix + query_tokens
def assemble_response(self, response_tokens):
suffix = self._encode_format("<eos>")
return response_tokens + suffix
def _decode(tokens):
return "".join(chr(t) for t in tokens)
class TestChatMLStrategy:
def test_name(self):
assert ChatMLStrategy(DummyTokenizer()).name == "chatml"
def test_assemble_prompt(self):
tk = DummyTokenizer()
strategy = ChatMLStrategy(tk)
query_tokens = tk.encode("hello")
prompt = strategy.assemble_prompt(query_tokens)
text = _decode(prompt)
assert "<|im_start|>user" in text
assert "hello" in text
assert "<|im_start|>assistant" in text
def test_assemble_response(self):
tk = DummyTokenizer()
strategy = ChatMLStrategy(tk)
response_tokens = tk.encode("world")
response = strategy.assemble_response(response_tokens)
text = _decode(response)
assert "world" in text
assert "<|im_end|>" in text
assert "<eos>" in text
def test_prompt_ends_with_assistant_start(self):
tk = DummyTokenizer()
strategy = ChatMLStrategy(tk)
prompt = strategy.assemble_prompt(tk.encode("hi"))
# prompt 末尾应该是 assistant_start 的 token ids
assert prompt[-len(strategy._assistant_start_ids):] == strategy._assistant_start_ids
class TestAlpacaStrategy:
def test_name(self):
assert AlpacaStrategy(DummyTokenizer()).name == "alpaca"
def test_assemble_prompt(self):
tk = DummyTokenizer()
strategy = AlpacaStrategy(tk)
query_tokens = tk.encode("hello")
prompt = strategy.assemble_prompt(query_tokens)
text = _decode(prompt)
assert "### Instruction:" in text
assert "hello" in text
assert "### Response:" in text
def test_assemble_response(self):
tk = DummyTokenizer()
strategy = AlpacaStrategy(tk)
response_tokens = tk.encode("world")
response = strategy.assemble_response(response_tokens)
text = _decode(response)
assert "world" in text
assert "<eos>" in text
class TestStrategyFactory:
def test_create_chatml(self):
tk = DummyTokenizer()
assert isinstance(StrategyFactory.create("chatml", tk), ChatMLStrategy)
def test_create_alpaca(self):
tk = DummyTokenizer()
assert isinstance(StrategyFactory.create("alpaca", tk), AlpacaStrategy)
def test_create_invalid_raises_error(self):
with pytest.raises(ValueError, match="Unknown strategy"):
StrategyFactory.create("invalid_strategy", DummyTokenizer())
def test_register_and_create(self):
StrategyFactory.register("dummy")(DummyStrategy)
tk = DummyTokenizer()
strategy = StrategyFactory.create("dummy", tk)
assert isinstance(strategy, DummyStrategy)
assert strategy.name == "dummy"
def test_available_strategies(self):
strategies = StrategyFactory.available_types()
assert "chatml" in strategies
assert "alpaca" in strategies