refactor: 从data 模块分离tokenizer

This commit is contained in:
2026-04-04 16:12:58 +08:00
parent b531232a9b
commit bd9741dc5f
16 changed files with 108 additions and 92 deletions
+2 -67
View File
@@ -1,78 +1,13 @@
from dataclasses import dataclass
from typing import Dict, Generator, List, Optional, Tuple, Union
from typing import Generator, List, Optional, Tuple, Union
import torch
from jinja2 import Template
from torch import Tensor
from astrai.config.param_config import ModelParameter
from astrai.factory import BaseFactory
from astrai.inference.core import EmbeddingEncoderCore, GeneratorCore, KVCacheManager
HistoryType = List[Tuple[str, str]]
MessageType = Dict[str, str]
# Predefined chat templates using jinja2
CHAT_TEMPLATES: Dict[str, str] = {
"chatml": """{%- if system_prompt -%}
<im▁start>system
{{ system_prompt }}<im▁end>
{%- endif -%}
{%- for message in messages -%}
<im▁start>{{ message['role'] }}
{{ message['content'] }}<im▁end>
{%- endfor -%}
<im▁start>assistant
""",
}
def build_prompt(
query: str,
system_prompt: Optional[str] = None,
history: Optional[HistoryType] = None,
template: Optional[str] = None,
) -> str:
"""Build prompt using jinja2 template for query and history.
Args:
query (str): query string.
system_prompt (Optional[str]): system prompt string.
history (Optional[HistoryType]): history list of query and response.
template (Optional[str]): jinja2 template string. If None, uses default chatml template.
Returns:
str: prompt string formatted according to the template.
Example:
# Use default template
prompt = build_prompt(query="Hello", history=[...])
# Use custom template
custom_template = '''
{%- for msg in messages -%}
{{ msg['role'] }}: {{ msg['content'] }}
{%- endfor -%}
'''
prompt = build_prompt(query="Hello", template=custom_template)
"""
# Convert history to message format
messages: List[MessageType] = []
if history:
for user_msg, assistant_msg in history:
messages.append({"role": "user", "content": user_msg})
messages.append({"role": "assistant", "content": assistant_msg})
messages.append({"role": "user", "content": query})
# Use provided template or default chatml template
template_str = template if template is not None else CHAT_TEMPLATES["chatml"]
# Render template
jinja_template = Template(template_str)
return jinja_template.render(
messages=messages,
system_prompt=system_prompt,
)
from astrai.tokenizer.chat_template import HistoryType, build_prompt
def pad_sequence(ids_list: List[List[int]], pad_id: int) -> Tuple[List[List[int]], int]: