@modular-prompt/driver 0.16.0 → 0.17.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +93 -10
- package/dist/cache-controller.d.ts +4 -0
- package/dist/cache-controller.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.js +6 -0
- package/dist/driver-registry/config-based-factory.js.map +1 -1
- package/dist/driver-registry/factory-helper.d.ts.map +1 -1
- package/dist/driver-registry/factory-helper.js +9 -2
- package/dist/driver-registry/factory-helper.js.map +1 -1
- package/dist/driver-registry/index.d.ts +1 -1
- package/dist/driver-registry/index.d.ts.map +1 -1
- package/dist/driver-registry/types.d.ts +15 -1
- package/dist/driver-registry/types.d.ts.map +1 -1
- package/dist/formatter/converter.d.ts.map +1 -1
- package/dist/formatter/converter.js +31 -2
- package/dist/formatter/converter.js.map +1 -1
- package/dist/index.d.ts +5 -3
- package/dist/index.d.ts.map +1 -1
- package/dist/index.js +5 -3
- package/dist/index.js.map +1 -1
- package/dist/local-inference/adapters.d.ts +6 -0
- package/dist/local-inference/adapters.d.ts.map +1 -1
- package/dist/local-inference/driver.d.ts.map +1 -1
- package/dist/local-inference/driver.js +45 -24
- package/dist/local-inference/driver.js.map +1 -1
- package/dist/local-inference/process-client.d.ts +4 -2
- package/dist/local-inference/process-client.d.ts.map +1 -1
- package/dist/local-inference/process-client.js +24 -8
- package/dist/local-inference/process-client.js.map +1 -1
- package/dist/local-inference/process-communication.d.ts +9 -2
- package/dist/local-inference/process-communication.d.ts.map +1 -1
- package/dist/local-inference/process-communication.js +37 -5
- package/dist/local-inference/process-communication.js.map +1 -1
- package/dist/local-inference/protocol.d.ts +4 -0
- package/dist/local-inference/protocol.d.ts.map +1 -1
- package/dist/local-inference/request-queue.d.ts +1 -1
- package/dist/local-inference/request-queue.d.ts.map +1 -1
- package/dist/local-inference/request-queue.js +26 -7
- package/dist/local-inference/request-queue.js.map +1 -1
- package/dist/local-inference/stream-utils.d.ts +6 -0
- package/dist/local-inference/stream-utils.d.ts.map +1 -1
- package/dist/local-inference/stream-utils.js.map +1 -1
- package/dist/mlx-ml/mlx-cache-controller.d.ts +9 -0
- package/dist/mlx-ml/mlx-cache-controller.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-cache-controller.js +158 -33
- package/dist/mlx-ml/mlx-cache-controller.js.map +1 -1
- package/dist/mlx-ml/mlx-cache-support.d.ts +2 -2
- package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-cache-support.js +8 -3
- package/dist/mlx-ml/mlx-cache-support.js.map +1 -1
- package/dist/mlx-ml/mlx-driver.d.ts +0 -1
- package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-driver.js +1 -8
- package/dist/mlx-ml/mlx-driver.js.map +1 -1
- package/dist/mlx-ml/process/index.d.ts +1 -1
- package/dist/mlx-ml/process/index.d.ts.map +1 -1
- package/dist/mlx-ml/process/index.js +2 -2
- package/dist/mlx-ml/process/index.js.map +1 -1
- package/dist/models-config/index.d.ts +1 -1
- package/dist/models-config/index.d.ts.map +1 -1
- package/dist/models-config/index.js +1 -1
- package/dist/models-config/index.js.map +1 -1
- package/dist/models-config/resolve.d.ts +9 -1
- package/dist/models-config/resolve.d.ts.map +1 -1
- package/dist/models-config/resolve.js +94 -2
- package/dist/models-config/resolve.js.map +1 -1
- package/dist/models-config/types.d.ts +3 -1
- package/dist/models-config/types.d.ts.map +1 -1
- package/dist/pytorch/process/index.d.ts +4 -2
- package/dist/pytorch/process/index.d.ts.map +1 -1
- package/dist/pytorch/process/index.js +24 -7
- package/dist/pytorch/process/index.js.map +1 -1
- package/dist/pytorch/pytorch-cache-controller.d.ts +84 -0
- package/dist/pytorch/pytorch-cache-controller.d.ts.map +1 -0
- package/dist/pytorch/pytorch-cache-controller.js +742 -0
- package/dist/pytorch/pytorch-cache-controller.js.map +1 -0
- package/dist/pytorch/pytorch-cache-support.d.ts +23 -0
- package/dist/pytorch/pytorch-cache-support.d.ts.map +1 -0
- package/dist/pytorch/pytorch-cache-support.js +47 -0
- package/dist/pytorch/pytorch-cache-support.js.map +1 -0
- package/dist/pytorch/pytorch-driver.d.ts +8 -1
- package/dist/pytorch/pytorch-driver.d.ts.map +1 -1
- package/dist/pytorch/pytorch-driver.js +40 -0
- package/dist/pytorch/pytorch-driver.js.map +1 -1
- package/dist/runtime/check.d.ts.map +1 -1
- package/dist/runtime/check.js +8 -6
- package/dist/runtime/check.js.map +1 -1
- package/dist/runtime/index.d.ts +2 -2
- package/dist/runtime/index.d.ts.map +1 -1
- package/dist/runtime/index.js +2 -2
- package/dist/runtime/index.js.map +1 -1
- package/dist/runtime/manifest-core.d.mts +1 -0
- package/dist/runtime/manifest-core.mjs +1 -0
- package/dist/runtime/manifest-core.mjs.map +1 -1
- package/dist/runtime/manifest.d.ts +2 -0
- package/dist/runtime/manifest.d.ts.map +1 -1
- package/dist/runtime/manifest.js.map +1 -1
- package/dist/runtime/paths-core.d.mts +15 -1
- package/dist/runtime/paths-core.d.mts.map +1 -1
- package/dist/runtime/paths-core.mjs +50 -5
- package/dist/runtime/paths-core.mjs.map +1 -1
- package/dist/runtime/paths.d.ts +2 -2
- package/dist/runtime/paths.d.ts.map +1 -1
- package/dist/runtime/paths.js +2 -2
- package/dist/runtime/paths.js.map +1 -1
- package/dist/runtime/pytorch-template-core.d.mts +11 -0
- package/dist/runtime/pytorch-template-core.d.mts.map +1 -0
- package/dist/runtime/pytorch-template-core.mjs +54 -0
- package/dist/runtime/pytorch-template-core.mjs.map +1 -0
- package/dist/runtime/setup-commands-core.d.mts +3 -0
- package/dist/runtime/setup-commands-core.d.mts.map +1 -1
- package/dist/runtime/setup-commands-core.mjs +4 -0
- package/dist/runtime/setup-commands-core.mjs.map +1 -1
- package/dist/runtime/setup-commands.d.ts +1 -1
- package/dist/runtime/setup-commands.d.ts.map +1 -1
- package/dist/runtime/setup-commands.js +1 -1
- package/dist/runtime/setup-commands.js.map +1 -1
- package/docs/DRIVER_API.md +455 -0
- package/docs/LOCAL_MODEL_SETUP.md +765 -0
- package/docs/mlx-api-selection.md +301 -0
- package/package.json +9 -5
- package/scripts/runtime-cli.bin.test.ts +142 -0
- package/scripts/runtime-cli.js +305 -35
- package/scripts/runtime-cli.test.ts +163 -0
- package/src/mlx-ml/python/__main__.py +1 -1
- package/src/mlx-ml/python/backends/base.py +88 -18
- package/src/mlx-ml/python/backends/mlx_lm.py +28 -3
- package/src/mlx-ml/python/backends/mlx_vlm.py +679 -2
- package/src/mlx-ml/python/handlers/cache.py +4 -0
- package/src/mlx-ml/python/handlers/generate.py +33 -10
- package/src/mlx-ml/python/handlers/tokenize.py +1 -4
- package/src/mlx-ml/python/pyproject.toml +1 -1
- package/src/mlx-ml/python/server.py +2 -0
- package/src/mlx-ml/python/uv.lock +8 -8
- package/src/pytorch/templates/cpu-minimal/backends/base.py +139 -0
- package/src/pytorch/templates/cpu-minimal/backends/transformers_lm.py +1167 -0
- package/src/pytorch/{python → templates/cpu-minimal}/handlers/__init__.py +1 -0
- package/src/pytorch/templates/cpu-minimal/handlers/cache.py +88 -0
- package/src/pytorch/templates/cpu-minimal/handlers/generate.py +157 -0
- package/src/pytorch/{python → templates/cpu-minimal}/pyproject.toml +2 -2
- package/src/pytorch/{python → templates/cpu-minimal}/server.py +20 -2
- package/src/pytorch/templates/cpu-minimal/tests/test_cache_handler.py +284 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_capabilities.py +14 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_server.py +141 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_transformers_errors.py +89 -0
- package/src/pytorch/templates/cpu-minimal/tests/test_transformers_lm_cache.py +554 -0
- package/src/pytorch/templates/cpu-minimal/utils/__init__.py +0 -0
- package/src/pytorch/{python → templates/cpu-minimal}/utils/token_utils.py +2 -2
- package/src/pytorch/templates/cpu-minimal/utils/transformers_errors.py +54 -0
- package/src/pytorch/{python → templates/cpu-minimal}/uv.lock +149 -109
- package/src/pytorch/templates/cuda/__main__.py +19 -0
- package/src/pytorch/templates/cuda/backends/__init__.py +3 -0
- package/src/pytorch/{python → templates/cuda}/backends/base.py +54 -6
- package/src/pytorch/templates/cuda/backends/transformers_lm.py +379 -0
- package/src/pytorch/templates/cuda/handlers/__init__.py +7 -0
- package/src/pytorch/templates/cuda/handlers/cache.py +93 -0
- package/src/pytorch/templates/cuda/handlers/cancel.py +53 -0
- package/src/pytorch/templates/cuda/handlers/capabilities.py +6 -0
- package/src/pytorch/templates/cuda/handlers/completion.py +15 -0
- package/src/pytorch/templates/cuda/handlers/format_test.py +70 -0
- package/src/pytorch/templates/cuda/handlers/generate.py +152 -0
- package/src/pytorch/templates/cuda/handlers/render.py +40 -0
- package/src/pytorch/templates/cuda/handlers/tokenize.py +63 -0
- package/src/pytorch/templates/cuda/pyproject.toml +37 -0
- package/src/pytorch/templates/cuda/server.py +158 -0
- package/src/pytorch/templates/cuda/tests/__init__.py +0 -0
- package/src/pytorch/templates/cuda/tests/test_cache_handler.py +207 -0
- package/src/pytorch/templates/cuda/tests/test_capabilities.py +14 -0
- package/src/pytorch/templates/cuda/tests/test_server.py +145 -0
- package/src/pytorch/templates/cuda/tests/test_transformers_errors.py +89 -0
- package/src/pytorch/templates/cuda/tests/test_transformers_lm_cache.py +288 -0
- package/src/pytorch/templates/cuda/utils/__init__.py +0 -0
- package/src/pytorch/templates/cuda/utils/chat_template_constraints.py +164 -0
- package/src/pytorch/templates/cuda/utils/prompt_builder.py +54 -0
- package/src/pytorch/templates/cuda/utils/template_render.py +80 -0
- package/src/pytorch/templates/cuda/utils/token_utils.py +376 -0
- package/src/pytorch/templates/cuda/utils/transformers_errors.py +54 -0
- package/src/pytorch/templates/cuda/uv.lock +734 -0
- package/src/pytorch/python/backends/transformers_lm.py +0 -127
- package/src/pytorch/python/handlers/generate.py +0 -68
- /package/src/pytorch/{python → templates/cpu-minimal}/__main__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/backends/__init__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/cancel.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/capabilities.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/completion.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/format_test.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/render.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/handlers/tokenize.py +0 -0
- /package/src/pytorch/{python/utils → templates/cpu-minimal/tests}/__init__.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/utils/chat_template_constraints.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/utils/prompt_builder.py +0 -0
- /package/src/pytorch/{python → templates/cpu-minimal}/utils/template_render.py +0 -0
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
"""プロンプト生成ユーティリティ"""
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
def supports_chat_template(tokenizer) -> bool:
|
|
5
|
+
return (hasattr(tokenizer, 'apply_chat_template') and
|
|
6
|
+
hasattr(tokenizer, 'chat_template') and
|
|
7
|
+
tokenizer.chat_template is not None)
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def generate_merged_prompt(messages, capabilities):
|
|
11
|
+
"""apply_chat_templateがない場合のプロンプト生成"""
|
|
12
|
+
prompt_parts = []
|
|
13
|
+
special_tokens = capabilities.get('special_tokens', {})
|
|
14
|
+
|
|
15
|
+
for msg in messages:
|
|
16
|
+
role = msg['role']
|
|
17
|
+
role_upper = role.upper()
|
|
18
|
+
|
|
19
|
+
role_token = special_tokens.get(role)
|
|
20
|
+
|
|
21
|
+
if role_token and isinstance(role_token, dict) and 'start' in role_token:
|
|
22
|
+
start_token = role_token['start']['text']
|
|
23
|
+
end_token = role_token['end']['text']
|
|
24
|
+
prompt_parts.extend([
|
|
25
|
+
start_token,
|
|
26
|
+
msg['content'].strip(),
|
|
27
|
+
end_token,
|
|
28
|
+
''
|
|
29
|
+
])
|
|
30
|
+
else:
|
|
31
|
+
block_token = None
|
|
32
|
+
for candidate in ['block', 'context', 'quote', 'section']:
|
|
33
|
+
token = special_tokens.get(candidate)
|
|
34
|
+
if token and isinstance(token, dict) and 'start' in token:
|
|
35
|
+
block_token = token
|
|
36
|
+
break
|
|
37
|
+
|
|
38
|
+
if block_token:
|
|
39
|
+
start_token = block_token['start']['text']
|
|
40
|
+
end_token = block_token['end']['text']
|
|
41
|
+
prompt_parts.extend([
|
|
42
|
+
f'{start_token}{role_upper}:\n{msg["content"].strip()}',
|
|
43
|
+
end_token,
|
|
44
|
+
''
|
|
45
|
+
])
|
|
46
|
+
else:
|
|
47
|
+
prompt_parts.extend([
|
|
48
|
+
f'<!-- begin of {role_upper} -->',
|
|
49
|
+
msg['content'].strip(),
|
|
50
|
+
f'<!-- end of {role_upper} -->',
|
|
51
|
+
''
|
|
52
|
+
])
|
|
53
|
+
|
|
54
|
+
return '\n'.join(prompt_parts[:-1])
|
|
@@ -0,0 +1,80 @@
|
|
|
1
|
+
"""apply_chat_template によるプロンプト整形(推論なし)"""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from utils.prompt_builder import supports_chat_template
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def apply_chat_template_prompt(
|
|
9
|
+
backend,
|
|
10
|
+
messages: list,
|
|
11
|
+
*,
|
|
12
|
+
primer: str | None = None,
|
|
13
|
+
tools: list | None = None,
|
|
14
|
+
reasoning_effort: str | None = None,
|
|
15
|
+
trust_remote_code: bool | None = None,
|
|
16
|
+
) -> str | list:
|
|
17
|
+
"""messages を chat template で整形して prompt を返す"""
|
|
18
|
+
tokenizer = backend.get_tokenizer()
|
|
19
|
+
|
|
20
|
+
if not supports_chat_template(tokenizer):
|
|
21
|
+
raise ValueError("chat_template_not_available")
|
|
22
|
+
|
|
23
|
+
add_generation_prompt = True
|
|
24
|
+
fmt_messages = list(messages)
|
|
25
|
+
if primer is not None:
|
|
26
|
+
fmt_messages.append({"role": "assistant", "content": primer})
|
|
27
|
+
add_generation_prompt = False
|
|
28
|
+
|
|
29
|
+
if backend.supports_vision():
|
|
30
|
+
try:
|
|
31
|
+
prompt = tokenizer.apply_chat_template(
|
|
32
|
+
fmt_messages,
|
|
33
|
+
tools=tools,
|
|
34
|
+
add_generation_prompt=add_generation_prompt,
|
|
35
|
+
tokenize=False,
|
|
36
|
+
)
|
|
37
|
+
except TypeError:
|
|
38
|
+
prompt = tokenizer.apply_chat_template(
|
|
39
|
+
fmt_messages,
|
|
40
|
+
add_generation_prompt=add_generation_prompt,
|
|
41
|
+
tokenize=False,
|
|
42
|
+
)
|
|
43
|
+
else:
|
|
44
|
+
extra_kwargs: dict = {}
|
|
45
|
+
if tools is not None:
|
|
46
|
+
extra_kwargs["tools"] = tools
|
|
47
|
+
if reasoning_effort is not None:
|
|
48
|
+
extra_kwargs["reasoning_effort"] = reasoning_effort
|
|
49
|
+
if trust_remote_code is not None:
|
|
50
|
+
extra_kwargs["trust_remote_code"] = trust_remote_code
|
|
51
|
+
|
|
52
|
+
try:
|
|
53
|
+
prompt = tokenizer.apply_chat_template(
|
|
54
|
+
fmt_messages,
|
|
55
|
+
add_generation_prompt=add_generation_prompt,
|
|
56
|
+
tokenize=False,
|
|
57
|
+
**extra_kwargs,
|
|
58
|
+
)
|
|
59
|
+
except TypeError:
|
|
60
|
+
try:
|
|
61
|
+
fallback_kwargs: dict = {}
|
|
62
|
+
if tools is not None:
|
|
63
|
+
fallback_kwargs["tools"] = tools
|
|
64
|
+
prompt = tokenizer.apply_chat_template(
|
|
65
|
+
fmt_messages,
|
|
66
|
+
add_generation_prompt=add_generation_prompt,
|
|
67
|
+
tokenize=False,
|
|
68
|
+
**fallback_kwargs,
|
|
69
|
+
)
|
|
70
|
+
except TypeError:
|
|
71
|
+
prompt = tokenizer.apply_chat_template(
|
|
72
|
+
fmt_messages,
|
|
73
|
+
add_generation_prompt=add_generation_prompt,
|
|
74
|
+
tokenize=False,
|
|
75
|
+
)
|
|
76
|
+
|
|
77
|
+
if primer is not None and isinstance(prompt, str):
|
|
78
|
+
prompt = primer.join(prompt.split(primer)[0:-1]) + primer
|
|
79
|
+
|
|
80
|
+
return prompt
|
|
@@ -0,0 +1,376 @@
|
|
|
1
|
+
"""
|
|
2
|
+
トークン関連のユーティリティ関数
|
|
3
|
+
"""
|
|
4
|
+
from utils.chat_template_constraints import detect_chat_restrictions
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def is_eod_token(response, tokenizer):
|
|
8
|
+
"""
|
|
9
|
+
レスポンスがEODトークンかどうかを判定する
|
|
10
|
+
|
|
11
|
+
Args:
|
|
12
|
+
response: stream_generateからのレスポンス
|
|
13
|
+
tokenizer: tokenizerオブジェクト(必須)
|
|
14
|
+
|
|
15
|
+
Returns:
|
|
16
|
+
bool: EODトークンの場合True
|
|
17
|
+
"""
|
|
18
|
+
# 1. finish_reasonによる終了判定(MLX-LMの標準的な方法)
|
|
19
|
+
if hasattr(response, 'finish_reason') and response.finish_reason == 'stop':
|
|
20
|
+
return True
|
|
21
|
+
|
|
22
|
+
# 2. response.tokenによる終了トークン判定
|
|
23
|
+
if hasattr(response, 'token'):
|
|
24
|
+
token = response.token
|
|
25
|
+
|
|
26
|
+
# special_tokens_mapとadded_tokens_encoderから終了トークンを取得
|
|
27
|
+
end_token_ids = []
|
|
28
|
+
|
|
29
|
+
# special_tokens_mapから標準的な終了トークンを取得
|
|
30
|
+
if hasattr(tokenizer, 'special_tokens_map') and hasattr(tokenizer, 'added_tokens_encoder'):
|
|
31
|
+
special_map = tokenizer.special_tokens_map
|
|
32
|
+
added_encoder = tokenizer.added_tokens_encoder
|
|
33
|
+
|
|
34
|
+
# EOSトークン
|
|
35
|
+
eos_token_str = special_map.get('eos_token')
|
|
36
|
+
if eos_token_str and eos_token_str in added_encoder:
|
|
37
|
+
end_token_ids.append(added_encoder[eos_token_str])
|
|
38
|
+
|
|
39
|
+
# その他の終了関連トークン
|
|
40
|
+
end_related_keys = ['eoi_token'] # end_of_image
|
|
41
|
+
for key in end_related_keys:
|
|
42
|
+
token_str = special_map.get(key)
|
|
43
|
+
if token_str and token_str in added_encoder:
|
|
44
|
+
end_token_ids.append(added_encoder[token_str])
|
|
45
|
+
|
|
46
|
+
# added_tokens_encoderから直接取得(会話終了トークンなど)
|
|
47
|
+
if hasattr(tokenizer, 'added_tokens_encoder'):
|
|
48
|
+
added_encoder = tokenizer.added_tokens_encoder
|
|
49
|
+
conversation_end_tokens = ['<end_of_turn>']
|
|
50
|
+
for token_str in conversation_end_tokens:
|
|
51
|
+
token_id = added_encoder.get(token_str)
|
|
52
|
+
if token_id is not None:
|
|
53
|
+
end_token_ids.append(token_id)
|
|
54
|
+
|
|
55
|
+
# 重複を除去してチェック
|
|
56
|
+
if token in set(end_token_ids):
|
|
57
|
+
return True
|
|
58
|
+
|
|
59
|
+
return False
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def get_special_tokens(tokenizer):
|
|
63
|
+
"""
|
|
64
|
+
tokenizerから特殊トークンを取得する
|
|
65
|
+
|
|
66
|
+
Returns:
|
|
67
|
+
dict: special_tokens情報
|
|
68
|
+
"""
|
|
69
|
+
special_tokens = {}
|
|
70
|
+
|
|
71
|
+
# 標準的なspecial tokens(tokenizerに定義されているもの)
|
|
72
|
+
# VLM processorではこれらの属性がない場合があるためgetattr使用
|
|
73
|
+
standard_tokens = {
|
|
74
|
+
"eod": getattr(tokenizer, "eos_token", None), # End of Document/Sequence
|
|
75
|
+
"bos": getattr(tokenizer, "bos_token", None), # Beginning of Sequence
|
|
76
|
+
"unk": getattr(tokenizer, "unk_token", None), # Unknown token
|
|
77
|
+
"pad": getattr(tokenizer, "pad_token", None), # Padding token
|
|
78
|
+
}
|
|
79
|
+
|
|
80
|
+
for name, token in standard_tokens.items():
|
|
81
|
+
if token is not None:
|
|
82
|
+
token_id = getattr(tokenizer, f"{name}_token_id", None)
|
|
83
|
+
if token_id is not None:
|
|
84
|
+
special_tokens[name] = {"text": token, "id": token_id}
|
|
85
|
+
|
|
86
|
+
# ペアトークン(存在する場合のみ)
|
|
87
|
+
pair_tokens = {
|
|
88
|
+
# ChatML基本形式
|
|
89
|
+
"system": ("<|system|>", "<|/system|>"),
|
|
90
|
+
"user": ("<|user|>", "<|/user|>"),
|
|
91
|
+
"assistant": ("<|assistant|>", "<|/assistant|>"),
|
|
92
|
+
|
|
93
|
+
# フォーマット・構造化
|
|
94
|
+
"code": ("<|code_start|>", "<|code_end|>"),
|
|
95
|
+
"python": ("<|python|>", "<|/python|>"),
|
|
96
|
+
"javascript": ("<|javascript|>", "<|/javascript|>"),
|
|
97
|
+
"bash": ("<|bash|>", "<|/bash|>"),
|
|
98
|
+
"quote": ("<|quote|>", "<|/quote|>"),
|
|
99
|
+
"ref": ("<|ref|>", "<|/ref|>"),
|
|
100
|
+
"citation": ("<|citation|>", "<|/citation|>"),
|
|
101
|
+
"table": ("<|table|>", "<|/table|>"),
|
|
102
|
+
"heading": ("<|heading|>", "<|/heading|>"),
|
|
103
|
+
|
|
104
|
+
# メディア・リッチコンテンツ
|
|
105
|
+
"image": ("<|image|>", "<|/image|>"),
|
|
106
|
+
"audio": ("<|audio|>", "<|/audio|>"),
|
|
107
|
+
"video": ("<|video|>", "<|/video|>"),
|
|
108
|
+
|
|
109
|
+
# 機能・制御
|
|
110
|
+
"tool_call": ("<|tool_call|>", "<|/tool_call|>"),
|
|
111
|
+
"function": ("<|function|>", "<|/function|>"),
|
|
112
|
+
"api": ("<|api|>", "<|/api|>"),
|
|
113
|
+
"search": ("<|search|>", "<|/search|>"),
|
|
114
|
+
"knowledge": ("<|knowledge|>", "<|/knowledge|>"),
|
|
115
|
+
"context": ("<|context|>", "<|/context|>"),
|
|
116
|
+
|
|
117
|
+
# 思考・推論
|
|
118
|
+
"thinking": ("<|thinking|>", "</thinking>"),
|
|
119
|
+
"reasoning": ("<|reasoning|>", "<|/reasoning|>"),
|
|
120
|
+
"scratchpad": ("<|scratchpad|>", "<|/scratchpad|>"),
|
|
121
|
+
"analysis": ("<|analysis|>", "<|/analysis|>"),
|
|
122
|
+
"summary": ("<|summary|>", "<|/summary|>"),
|
|
123
|
+
"explanation": ("<|explanation|>", "<|/explanation|>"),
|
|
124
|
+
|
|
125
|
+
# tool_call バリエーション(追加)
|
|
126
|
+
"tool_call_explicit": ("<|tool_call_start|>", "<|tool_call_end|>"),
|
|
127
|
+
"tool_call_xml": ("<tool_call>", "</tool_call>"),
|
|
128
|
+
"tool_calls_section": ("<|tool_calls_section_begin|>", "<|tool_calls_section_end|>"),
|
|
129
|
+
"function_call_tags": ("<start_function_call>", "<end_function_call>"),
|
|
130
|
+
"longcat_tool_call": ("<longcat_tool_call>", "</longcat_tool_call>"),
|
|
131
|
+
"minimax_tool_call": ("<minimax:tool_call>", "</minimax:tool_call>"),
|
|
132
|
+
}
|
|
133
|
+
|
|
134
|
+
# 単体トークン(存在する場合のみ)
|
|
135
|
+
single_tokens = {
|
|
136
|
+
# Fill-in-the-Middle
|
|
137
|
+
"fim_prefix": "<|fim_prefix|>",
|
|
138
|
+
"fim_middle": "<|fim_middle|>",
|
|
139
|
+
"fim_suffix": "<|fim_suffix|>",
|
|
140
|
+
|
|
141
|
+
# リスト・構造
|
|
142
|
+
"list_item": "<|list_item|>",
|
|
143
|
+
|
|
144
|
+
# メディア単体
|
|
145
|
+
"vision": "<|vision|>",
|
|
146
|
+
|
|
147
|
+
# 一般的なマークダウン風
|
|
148
|
+
"code_inline": "`",
|
|
149
|
+
"code_block_start": "```",
|
|
150
|
+
"code_block_end": "```",
|
|
151
|
+
|
|
152
|
+
# ツール関連の単体トークン(追加)
|
|
153
|
+
"tool_calls_marker": "[TOOL_CALLS]",
|
|
154
|
+
# Harmony形式のcallトークン(tool_call_endとは異なる用途)
|
|
155
|
+
"harmony_call": "<|call|>",
|
|
156
|
+
}
|
|
157
|
+
|
|
158
|
+
# VLM processorではconvert_tokens_to_idsがない場合がある
|
|
159
|
+
if not hasattr(tokenizer, 'convert_tokens_to_ids'):
|
|
160
|
+
return special_tokens
|
|
161
|
+
|
|
162
|
+
unk_token_id = getattr(tokenizer, "unk_token_id", None)
|
|
163
|
+
|
|
164
|
+
# ペアトークンの処理
|
|
165
|
+
for name, (start_token, end_token) in pair_tokens.items():
|
|
166
|
+
start_id = tokenizer.convert_tokens_to_ids(start_token)
|
|
167
|
+
end_id = tokenizer.convert_tokens_to_ids(end_token)
|
|
168
|
+
|
|
169
|
+
# unk_tokenでない場合のみ追加
|
|
170
|
+
if start_id != unk_token_id and end_id != unk_token_id:
|
|
171
|
+
special_tokens[name] = {
|
|
172
|
+
"start": {"text": start_token, "id": start_id},
|
|
173
|
+
"end": {"text": end_token, "id": end_id}
|
|
174
|
+
}
|
|
175
|
+
|
|
176
|
+
# 単体トークンの処理
|
|
177
|
+
for name, token_text in single_tokens.items():
|
|
178
|
+
token_id = tokenizer.convert_tokens_to_ids(token_text)
|
|
179
|
+
|
|
180
|
+
# unk_tokenでない場合のみ追加
|
|
181
|
+
if token_id != unk_token_id:
|
|
182
|
+
special_tokens[name] = {"text": token_text, "id": token_id}
|
|
183
|
+
|
|
184
|
+
return special_tokens
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
def detect_tool_call_format(tokenizer):
|
|
188
|
+
"""tokenizer設定からtool call/resultのデリミタを検出する
|
|
189
|
+
|
|
190
|
+
tokenizerアクセスが必要な情報のみを抽出する:
|
|
191
|
+
- tool_parser_type: tokenizer_config由来のパーサー種別文字列
|
|
192
|
+
- chat_templateテキストからのデリミタパターン検出
|
|
193
|
+
|
|
194
|
+
パーサー種別→デリミタのマッピングはTS側(detector.ts)に一元化。
|
|
195
|
+
Python側はtokenizerから生の情報を抽出して渡す役割に専念する。
|
|
196
|
+
|
|
197
|
+
Returns:
|
|
198
|
+
dict or None: {
|
|
199
|
+
"tool_parser_type": str, # tokenizer_configのtool_parser_type
|
|
200
|
+
"call_start": str, # chat_templateから検出した開始デリミタ
|
|
201
|
+
"call_end": str, # chat_templateから検出した終了デリミタ
|
|
202
|
+
"response_start": str, # tool responseの開始デリミタ(検出時)
|
|
203
|
+
"response_end": str, # tool responseの終了デリミタ(検出時)
|
|
204
|
+
} or None
|
|
205
|
+
"""
|
|
206
|
+
import re
|
|
207
|
+
|
|
208
|
+
tool_parser_type = None
|
|
209
|
+
if hasattr(tokenizer, 'init_kwargs'):
|
|
210
|
+
tool_parser_type = tokenizer.init_kwargs.get('tool_parser_type')
|
|
211
|
+
|
|
212
|
+
template = getattr(tokenizer, 'chat_template', None)
|
|
213
|
+
if not template and hasattr(tokenizer, 'init_kwargs'):
|
|
214
|
+
template = tokenizer.init_kwargs.get('chat_template', '')
|
|
215
|
+
|
|
216
|
+
if not tool_parser_type and not template:
|
|
217
|
+
return None
|
|
218
|
+
|
|
219
|
+
result = {}
|
|
220
|
+
if tool_parser_type:
|
|
221
|
+
result["tool_parser_type"] = tool_parser_type
|
|
222
|
+
|
|
223
|
+
if template:
|
|
224
|
+
tool_call_patterns = [
|
|
225
|
+
(r'<\|?tool_call\|?>', r'</tool_call>|<\|/tool_call\|>|<tool_call\|>'),
|
|
226
|
+
(r'<\|tool_call_start\|>', r'<\|tool_call_end\|>'),
|
|
227
|
+
(r'<start_function_call>', r'<end_function_call>'),
|
|
228
|
+
(r'<\|tool_calls_section_begin\|>', r'<\|tool_calls_section_end\|>'),
|
|
229
|
+
(r'<longcat_tool_call>', r'</longcat_tool_call>'),
|
|
230
|
+
(r'<minimax:tool_call>', r'</minimax:tool_call>'),
|
|
231
|
+
]
|
|
232
|
+
|
|
233
|
+
for start_pattern, end_pattern in tool_call_patterns:
|
|
234
|
+
start_match = re.search(start_pattern, template)
|
|
235
|
+
end_match = re.search(end_pattern, template)
|
|
236
|
+
if start_match and end_match:
|
|
237
|
+
result["call_start"] = start_match.group(0)
|
|
238
|
+
result["call_end"] = end_match.group(0)
|
|
239
|
+
break
|
|
240
|
+
|
|
241
|
+
if "call_start" not in result:
|
|
242
|
+
has_functions = re.search(r'"functions\."', template)
|
|
243
|
+
has_call = re.search(r'<\|call\|>', template)
|
|
244
|
+
if has_functions and has_call:
|
|
245
|
+
result["tool_parser_type"] = "harmony"
|
|
246
|
+
result["call_start"] = "to=functions."
|
|
247
|
+
result["call_end"] = "<|call|>"
|
|
248
|
+
|
|
249
|
+
if "call_start" not in result:
|
|
250
|
+
mistral_match = re.search(r'\[TOOL_CALLS\]', template)
|
|
251
|
+
if mistral_match:
|
|
252
|
+
result["call_start"] = "[TOOL_CALLS]"
|
|
253
|
+
result["call_end"] = ""
|
|
254
|
+
|
|
255
|
+
resp_tags = re.findall(r'<[|/]?tool_response[|]?>', template)
|
|
256
|
+
if len(resp_tags) >= 2:
|
|
257
|
+
open_tags = [t for t in resp_tags if '/' not in t]
|
|
258
|
+
close_tags = [t for t in resp_tags if '/' in t]
|
|
259
|
+
if open_tags and close_tags:
|
|
260
|
+
result["response_start"] = open_tags[0]
|
|
261
|
+
result["response_end"] = close_tags[0]
|
|
262
|
+
|
|
263
|
+
return result if result else None
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
def get_chat_template_info(tokenizer):
|
|
267
|
+
"""チャットテンプレートの詳細情報を取得"""
|
|
268
|
+
if not hasattr(tokenizer, 'apply_chat_template'):
|
|
269
|
+
return None
|
|
270
|
+
|
|
271
|
+
template_info = {
|
|
272
|
+
"supported_roles": [],
|
|
273
|
+
"preview": None,
|
|
274
|
+
"constraints": {}
|
|
275
|
+
}
|
|
276
|
+
|
|
277
|
+
# tool callフォーマット検出
|
|
278
|
+
tool_format = detect_tool_call_format(tokenizer)
|
|
279
|
+
if tool_format:
|
|
280
|
+
template_info["tool_call_format"] = tool_format
|
|
281
|
+
|
|
282
|
+
# サポートされるroleを検査
|
|
283
|
+
test_roles = ["system", "user", "assistant", "tool", "function"]
|
|
284
|
+
for role in test_roles:
|
|
285
|
+
test_msg = [{"role": role, "content": "test"}]
|
|
286
|
+
try:
|
|
287
|
+
tokenizer.apply_chat_template(test_msg, tokenize=False, add_generation_prompt=False)
|
|
288
|
+
template_info["supported_roles"].append(role)
|
|
289
|
+
except:
|
|
290
|
+
continue
|
|
291
|
+
|
|
292
|
+
# プレビュー生成
|
|
293
|
+
if template_info["supported_roles"]:
|
|
294
|
+
sample_messages = []
|
|
295
|
+
if "system" in template_info["supported_roles"]:
|
|
296
|
+
sample_messages.append({"role": "system", "content": "You are a helpful assistant."})
|
|
297
|
+
if "user" in template_info["supported_roles"]:
|
|
298
|
+
sample_messages.append({"role": "user", "content": "Hello!"})
|
|
299
|
+
if "assistant" in template_info["supported_roles"]:
|
|
300
|
+
sample_messages.append({"role": "assistant", "content": "Hi there!"})
|
|
301
|
+
|
|
302
|
+
try:
|
|
303
|
+
template_info["preview"] = tokenizer.apply_chat_template(
|
|
304
|
+
sample_messages,
|
|
305
|
+
tokenize=False,
|
|
306
|
+
add_generation_prompt=False
|
|
307
|
+
)
|
|
308
|
+
except Exception as e:
|
|
309
|
+
template_info["preview"] = f"Preview error: {e}"
|
|
310
|
+
|
|
311
|
+
return template_info
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
def get_tokenizer_features(tokenizer):
|
|
315
|
+
"""
|
|
316
|
+
tokenizerの機能情報を取得する
|
|
317
|
+
|
|
318
|
+
Returns:
|
|
319
|
+
dict: features情報
|
|
320
|
+
"""
|
|
321
|
+
# apply_chat_templateメソッドの存在だけでなく、テンプレートが実際に設定されているか確認
|
|
322
|
+
has_chat_template = hasattr(tokenizer, 'apply_chat_template') and bool(getattr(tokenizer, 'chat_template', None))
|
|
323
|
+
features = {
|
|
324
|
+
"apply_chat_template": has_chat_template,
|
|
325
|
+
"vocab_size": getattr(tokenizer, 'vocab_size', None),
|
|
326
|
+
"model_max_length": getattr(tokenizer, 'model_max_length', None)
|
|
327
|
+
}
|
|
328
|
+
|
|
329
|
+
# チャットテンプレート情報を追加
|
|
330
|
+
chat_template_info = get_chat_template_info(tokenizer)
|
|
331
|
+
if chat_template_info:
|
|
332
|
+
features["chat_template"] = chat_template_info
|
|
333
|
+
|
|
334
|
+
return features
|
|
335
|
+
|
|
336
|
+
|
|
337
|
+
def get_capabilities(tokenizer):
|
|
338
|
+
"""
|
|
339
|
+
tokenizerの全機能情報を取得する(capabilities API用)
|
|
340
|
+
|
|
341
|
+
Returns:
|
|
342
|
+
dict: capabilities情報
|
|
343
|
+
"""
|
|
344
|
+
# 基本メソッド
|
|
345
|
+
methods = ["capabilities", "completion", "generate", "format_test", "cache_prefill"]
|
|
346
|
+
|
|
347
|
+
# apply_chat_templateがある場合はchatメソッドを追加
|
|
348
|
+
if hasattr(tokenizer, 'apply_chat_template'):
|
|
349
|
+
methods.append("render")
|
|
350
|
+
|
|
351
|
+
capabilities = {
|
|
352
|
+
"methods": methods,
|
|
353
|
+
"special_tokens": get_special_tokens(tokenizer),
|
|
354
|
+
"features": get_tokenizer_features(tokenizer)
|
|
355
|
+
}
|
|
356
|
+
|
|
357
|
+
# tool_call_formatの情報をspecial_tokensに反映(補完)
|
|
358
|
+
features = capabilities.get("features", {})
|
|
359
|
+
chat_template = features.get("chat_template")
|
|
360
|
+
if chat_template:
|
|
361
|
+
tcf = chat_template.get("tool_call_format")
|
|
362
|
+
if tcf and tcf.get("call_start") and "tool_call" not in capabilities["special_tokens"]:
|
|
363
|
+
call_start = tcf["call_start"]
|
|
364
|
+
call_end = tcf.get("call_end", "")
|
|
365
|
+
if call_end: # ペアがある場合のみ
|
|
366
|
+
capabilities["special_tokens"]["tool_call"] = {
|
|
367
|
+
"start": {"text": call_start, "id": -1},
|
|
368
|
+
"end": {"text": call_end, "id": -1}
|
|
369
|
+
}
|
|
370
|
+
|
|
371
|
+
# チャット制約を検出して追加
|
|
372
|
+
chat_restrictions = detect_chat_restrictions(tokenizer)
|
|
373
|
+
if chat_restrictions:
|
|
374
|
+
capabilities["chat_restrictions"] = chat_restrictions
|
|
375
|
+
|
|
376
|
+
return capabilities
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
"""Helpful errors for model architectures not supported by the runtime."""
|
|
2
|
+
|
|
3
|
+
import re
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
MIN_TRANSFORMERS_VERSION = "5.14.0"
|
|
7
|
+
|
|
8
|
+
_MODEL_TYPE_MESSAGE = re.compile(
|
|
9
|
+
r"model type [`'\"](?P<model_type>[A-Za-z0-9_.-]+)[`'\"]",
|
|
10
|
+
re.IGNORECASE,
|
|
11
|
+
)
|
|
12
|
+
_MODEL_TYPE_KEY = re.compile(r"[A-Za-z][A-Za-z0-9_.-]*")
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def extract_unsupported_model_type(error: BaseException) -> str | None:
|
|
16
|
+
"""Extract a model type from Transformers' unknown-architecture errors.
|
|
17
|
+
|
|
18
|
+
Transformers normally raises ``ValueError`` with a message containing the
|
|
19
|
+
model type. Older releases can leak the registry ``KeyError`` instead,
|
|
20
|
+
so handle that form as well while leaving unrelated errors untouched.
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
match = _MODEL_TYPE_MESSAGE.search(str(error))
|
|
24
|
+
if match:
|
|
25
|
+
return match.group("model_type")
|
|
26
|
+
|
|
27
|
+
if isinstance(error, KeyError) and len(error.args) == 1:
|
|
28
|
+
candidate = error.args[0]
|
|
29
|
+
if (
|
|
30
|
+
isinstance(candidate, str)
|
|
31
|
+
and candidate != "model_type"
|
|
32
|
+
and _MODEL_TYPE_KEY.fullmatch(candidate)
|
|
33
|
+
):
|
|
34
|
+
return candidate
|
|
35
|
+
|
|
36
|
+
return None
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def unsupported_model_type_error(
|
|
40
|
+
model_name: str,
|
|
41
|
+
model_type: str,
|
|
42
|
+
transformers_version: str,
|
|
43
|
+
) -> RuntimeError:
|
|
44
|
+
"""Build the actionable load error shown to PyTorch runtime users."""
|
|
45
|
+
|
|
46
|
+
return RuntimeError(
|
|
47
|
+
f"Cannot load model '{model_name}': Transformers does not recognize "
|
|
48
|
+
f"model_type '{model_type}'. The PyTorch runtime is using transformers "
|
|
49
|
+
f"{transformers_version}, but this model requires transformers>="
|
|
50
|
+
f"{MIN_TRANSFORMERS_VERSION}. Run `setup-pytorch` again after updating "
|
|
51
|
+
"@modular-prompt/driver. If the runtime already has a Python project, "
|
|
52
|
+
"update its transformers requirement and run "
|
|
53
|
+
"`modular-prompt-runtime sync pytorch`."
|
|
54
|
+
)
|