@modular-prompt/driver 0.13.5 → 0.15.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 +117 -4
- package/dist/driver-registry/ai-service.d.ts +23 -1
- package/dist/driver-registry/ai-service.d.ts.map +1 -1
- package/dist/driver-registry/ai-service.js +44 -10
- package/dist/driver-registry/ai-service.js.map +1 -1
- package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
- package/dist/driver-registry/config-based-factory.js +16 -0
- package/dist/driver-registry/config-based-factory.js.map +1 -1
- package/dist/driver-registry/factory-helper.d.ts +2 -0
- package/dist/driver-registry/factory-helper.d.ts.map +1 -1
- package/dist/driver-registry/factory-helper.js +21 -3
- package/dist/driver-registry/factory-helper.js.map +1 -1
- package/dist/driver-registry/index.d.ts +2 -2
- package/dist/driver-registry/index.d.ts.map +1 -1
- package/dist/driver-registry/index.js +1 -1
- package/dist/driver-registry/index.js.map +1 -1
- package/dist/driver-registry/registry.d.ts.map +1 -1
- package/dist/driver-registry/registry.js +3 -1
- package/dist/driver-registry/registry.js.map +1 -1
- package/dist/driver-registry/types.d.ts +8 -2
- package/dist/driver-registry/types.d.ts.map +1 -1
- package/dist/index.d.ts +11 -2
- package/dist/index.d.ts.map +1 -1
- package/dist/index.js +11 -1
- package/dist/index.js.map +1 -1
- package/dist/local-inference/adapters.d.ts +66 -0
- package/dist/local-inference/adapters.d.ts.map +1 -0
- package/dist/local-inference/adapters.js +2 -0
- package/dist/local-inference/adapters.js.map +1 -0
- package/dist/local-inference/driver.d.ts +51 -0
- package/dist/local-inference/driver.d.ts.map +1 -0
- package/dist/local-inference/driver.js +309 -0
- package/dist/local-inference/driver.js.map +1 -0
- package/dist/local-inference/index.d.ts +22 -0
- package/dist/local-inference/index.d.ts.map +1 -0
- package/dist/local-inference/index.js +17 -0
- package/dist/local-inference/index.js.map +1 -0
- package/dist/local-inference/process-client.d.ts +50 -0
- package/dist/local-inference/process-client.d.ts.map +1 -0
- package/dist/local-inference/process-client.js +92 -0
- package/dist/local-inference/process-client.js.map +1 -0
- package/dist/local-inference/process-communication.d.ts +41 -0
- package/dist/local-inference/process-communication.d.ts.map +1 -0
- package/dist/{mlx-ml/process → local-inference}/process-communication.js +40 -47
- package/dist/local-inference/process-communication.js.map +1 -0
- package/dist/local-inference/process-port.d.ts +12 -0
- package/dist/local-inference/process-port.d.ts.map +1 -0
- package/dist/local-inference/process-port.js +2 -0
- package/dist/local-inference/process-port.js.map +1 -0
- package/dist/local-inference/prompt-utils.d.ts +6 -0
- package/dist/local-inference/prompt-utils.d.ts.map +1 -0
- package/dist/local-inference/prompt-utils.js +17 -0
- package/dist/local-inference/prompt-utils.js.map +1 -0
- package/dist/local-inference/protocol.d.ts +192 -0
- package/dist/local-inference/protocol.d.ts.map +1 -0
- package/dist/local-inference/protocol.js +2 -0
- package/dist/local-inference/protocol.js.map +1 -0
- package/dist/local-inference/queue-types.d.ts +54 -0
- package/dist/local-inference/queue-types.d.ts.map +1 -0
- package/dist/local-inference/queue-types.js +2 -0
- package/dist/local-inference/queue-types.js.map +1 -0
- package/dist/local-inference/request-queue.d.ts +36 -0
- package/dist/local-inference/request-queue.d.ts.map +1 -0
- package/dist/{mlx-ml/process/queue.js → local-inference/request-queue.js} +89 -56
- package/dist/local-inference/request-queue.js.map +1 -0
- package/dist/local-inference/stream-utils.d.ts +19 -0
- package/dist/local-inference/stream-utils.d.ts.map +1 -0
- package/dist/local-inference/stream-utils.js +76 -0
- package/dist/local-inference/stream-utils.js.map +1 -0
- package/dist/mlx-ml/mlx-cache-support.d.ts +23 -0
- package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -0
- package/dist/mlx-ml/mlx-cache-support.js +45 -0
- package/dist/mlx-ml/mlx-cache-support.js.map +1 -0
- package/dist/mlx-ml/mlx-driver.d.ts +20 -59
- package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
- package/dist/mlx-ml/mlx-driver.js +86 -415
- package/dist/mlx-ml/mlx-driver.js.map +1 -1
- package/dist/mlx-ml/mlx-local-inference-adapters.d.ts +3 -0
- package/dist/mlx-ml/mlx-local-inference-adapters.d.ts.map +1 -0
- package/dist/mlx-ml/mlx-local-inference-adapters.js +19 -0
- package/dist/mlx-ml/mlx-local-inference-adapters.js.map +1 -0
- package/dist/mlx-ml/mlx-options.d.ts +19 -0
- package/dist/mlx-ml/mlx-options.d.ts.map +1 -0
- package/dist/mlx-ml/mlx-options.js +30 -0
- package/dist/mlx-ml/mlx-options.js.map +1 -0
- package/dist/mlx-ml/process/index.d.ts +11 -8
- package/dist/mlx-ml/process/index.d.ts.map +1 -1
- package/dist/mlx-ml/process/index.js +75 -52
- package/dist/mlx-ml/process/index.js.map +1 -1
- package/dist/mlx-ml/process/model-specific.d.ts +2 -1
- package/dist/mlx-ml/process/model-specific.d.ts.map +1 -1
- package/dist/mlx-ml/process/model-specific.js.map +1 -1
- package/dist/mlx-ml/process/prompt-builder.d.ts +11 -0
- package/dist/mlx-ml/process/prompt-builder.d.ts.map +1 -0
- package/dist/mlx-ml/process/prompt-builder.js +51 -0
- package/dist/mlx-ml/process/prompt-builder.js.map +1 -0
- package/dist/mlx-ml/process/types.d.ts +15 -183
- package/dist/mlx-ml/process/types.d.ts.map +1 -1
- package/dist/mlx-ml/tool-call-parser/tool-formatter.js +2 -2
- package/dist/mlx-ml/tool-call-parser/tool-formatter.js.map +1 -1
- package/dist/mlx-ml/types.d.ts +2 -45
- package/dist/mlx-ml/types.d.ts.map +1 -1
- package/dist/models-config/index.d.ts +8 -0
- package/dist/models-config/index.d.ts.map +1 -0
- package/dist/models-config/index.js +7 -0
- package/dist/models-config/index.js.map +1 -0
- package/dist/models-config/loader.d.ts +20 -0
- package/dist/models-config/loader.d.ts.map +1 -0
- package/dist/models-config/loader.js +85 -0
- package/dist/models-config/loader.js.map +1 -0
- package/dist/models-config/paths.d.ts +7 -0
- package/dist/models-config/paths.d.ts.map +1 -0
- package/dist/models-config/paths.js +11 -0
- package/dist/models-config/paths.js.map +1 -0
- package/dist/models-config/resolve.d.ts +57 -0
- package/dist/models-config/resolve.d.ts.map +1 -0
- package/dist/models-config/resolve.js +187 -0
- package/dist/models-config/resolve.js.map +1 -0
- package/dist/models-config/types.d.ts +60 -0
- package/dist/models-config/types.d.ts.map +1 -0
- package/dist/models-config/types.js +5 -0
- package/dist/models-config/types.js.map +1 -0
- package/dist/pytorch/process/index.d.ts +35 -0
- package/dist/pytorch/process/index.d.ts.map +1 -0
- package/dist/pytorch/process/index.js +69 -0
- package/dist/pytorch/process/index.js.map +1 -0
- package/dist/pytorch/pytorch-driver.d.ts +35 -0
- package/dist/pytorch/pytorch-driver.d.ts.map +1 -0
- package/dist/pytorch/pytorch-driver.js +48 -0
- package/dist/pytorch/pytorch-driver.js.map +1 -0
- package/dist/pytorch/pytorch-local-inference-adapters.d.ts +3 -0
- package/dist/pytorch/pytorch-local-inference-adapters.d.ts.map +1 -0
- package/dist/pytorch/pytorch-local-inference-adapters.js +19 -0
- package/dist/pytorch/pytorch-local-inference-adapters.js.map +1 -0
- package/dist/pytorch/pytorch-options.d.ts +8 -0
- package/dist/pytorch/pytorch-options.d.ts.map +1 -0
- package/dist/pytorch/pytorch-options.js +21 -0
- package/dist/pytorch/pytorch-options.js.map +1 -0
- package/dist/query-logger.js +1 -1
- package/dist/query-logger.js.map +1 -1
- package/dist/query-utils.d.ts +29 -0
- package/dist/query-utils.d.ts.map +1 -0
- package/dist/query-utils.js +61 -0
- package/dist/query-utils.js.map +1 -0
- package/dist/runtime/check.d.ts +10 -0
- package/dist/runtime/check.d.ts.map +1 -0
- package/dist/runtime/check.js +24 -0
- package/dist/runtime/check.js.map +1 -0
- package/dist/runtime/index.d.ts +4 -0
- package/dist/runtime/index.d.ts.map +1 -0
- package/dist/runtime/index.js +4 -0
- package/dist/runtime/index.js.map +1 -0
- package/dist/runtime/manifest-core.d.mts +27 -0
- package/dist/runtime/manifest-core.d.mts.map +1 -0
- package/dist/runtime/manifest-core.mjs +68 -0
- package/dist/runtime/manifest-core.mjs.map +1 -0
- package/dist/runtime/manifest.d.ts +18 -0
- package/dist/runtime/manifest.d.ts.map +1 -0
- package/dist/runtime/manifest.js +9 -0
- package/dist/runtime/manifest.js.map +1 -0
- package/dist/runtime/paths-core.d.mts +17 -0
- package/dist/runtime/paths-core.d.mts.map +1 -0
- package/dist/runtime/paths-core.mjs +67 -0
- package/dist/runtime/paths-core.mjs.map +1 -0
- package/dist/runtime/paths.d.ts +12 -0
- package/dist/runtime/paths.d.ts.map +1 -0
- package/dist/runtime/paths.js +17 -0
- package/dist/runtime/paths.js.map +1 -0
- package/dist/types.d.ts +20 -1
- package/dist/types.d.ts.map +1 -1
- package/dist/types.js.map +1 -1
- package/package.json +8 -5
- package/scripts/download-model.js +25 -9
- package/scripts/runtime-cli.js +315 -0
- package/skills/driver-usage/SKILL.md +56 -1
- package/src/mlx-ml/python/__main__.py +43 -4
- package/src/mlx-ml/python/handlers/__init__.py +2 -1
- package/src/mlx-ml/python/handlers/cancel.py +53 -0
- package/src/mlx-ml/python/handlers/completion.py +3 -24
- package/src/mlx-ml/python/handlers/{chat.py → generate.py} +36 -106
- package/src/mlx-ml/python/handlers/render.py +40 -0
- package/src/mlx-ml/python/pyproject.toml +3 -2
- package/src/mlx-ml/python/server.py +35 -8
- package/src/mlx-ml/python/utils/template_render.py +80 -0
- package/src/mlx-ml/python/utils/token_utils.py +2 -2
- package/src/mlx-ml/python/uv.lock +549 -454
- package/src/pytorch/python/__main__.py +19 -0
- package/src/pytorch/python/backends/__init__.py +3 -0
- package/src/pytorch/python/backends/base.py +84 -0
- package/src/pytorch/python/backends/transformers_lm.py +127 -0
- package/src/pytorch/python/handlers/__init__.py +6 -0
- package/src/pytorch/python/handlers/cancel.py +53 -0
- package/src/pytorch/python/handlers/capabilities.py +6 -0
- package/src/pytorch/python/handlers/completion.py +15 -0
- package/src/pytorch/python/handlers/format_test.py +70 -0
- package/src/pytorch/python/handlers/generate.py +68 -0
- package/src/pytorch/python/handlers/render.py +40 -0
- package/src/pytorch/python/handlers/tokenize.py +63 -0
- package/src/pytorch/python/pyproject.toml +36 -0
- package/src/pytorch/python/server.py +140 -0
- package/src/pytorch/python/utils/__init__.py +0 -0
- package/src/pytorch/python/utils/chat_template_constraints.py +164 -0
- package/src/pytorch/python/utils/prompt_builder.py +54 -0
- package/src/pytorch/python/utils/template_render.py +80 -0
- package/src/pytorch/python/utils/token_utils.py +376 -0
- package/src/pytorch/python/uv.lock +694 -0
- package/dist/mlx-ml/process/process-communication.d.ts +0 -37
- package/dist/mlx-ml/process/process-communication.d.ts.map +0 -1
- package/dist/mlx-ml/process/process-communication.js.map +0 -1
- package/dist/mlx-ml/process/queue.d.ts +0 -33
- package/dist/mlx-ml/process/queue.d.ts.map +0 -1
- package/dist/mlx-ml/process/queue.js.map +0 -1
- package/scripts/setup-mlx.js +0 -53
|
@@ -0,0 +1,140 @@
|
|
|
1
|
+
"""JSON-RPC風サーバー: stdin/stdoutベースのリクエストディスパッチ"""
|
|
2
|
+
import json
|
|
3
|
+
import sys
|
|
4
|
+
|
|
5
|
+
from backends.base import ModelBackend
|
|
6
|
+
from handlers import handle_capabilities, handle_completion, handle_format_test, handle_generate, handle_render, handle_tokenize
|
|
7
|
+
from handlers.cancel import request_cancel, reset_cancel
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
MAX_READ_LINES = 10000
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def read():
|
|
14
|
+
lines = []
|
|
15
|
+
while True:
|
|
16
|
+
line = sys.stdin.readline()
|
|
17
|
+
if not line:
|
|
18
|
+
return None
|
|
19
|
+
lines.append(line)
|
|
20
|
+
if len(lines) > MAX_READ_LINES:
|
|
21
|
+
sys.stderr.write(f"Error: read buffer exceeded {MAX_READ_LINES} lines, discarding\n")
|
|
22
|
+
lines.clear()
|
|
23
|
+
continue
|
|
24
|
+
try:
|
|
25
|
+
return json.loads(''.join(lines))
|
|
26
|
+
except json.JSONDecodeError:
|
|
27
|
+
continue
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def _is_valid_generate_prompt(prompt) -> bool:
|
|
31
|
+
"""prompt は非空文字列または非空のトークン ID リスト"""
|
|
32
|
+
if isinstance(prompt, str):
|
|
33
|
+
return bool(prompt)
|
|
34
|
+
if isinstance(prompt, list):
|
|
35
|
+
return len(prompt) > 0
|
|
36
|
+
return False
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class Server:
|
|
40
|
+
def __init__(self, backend: ModelBackend, capabilities: dict):
|
|
41
|
+
self.backend = backend
|
|
42
|
+
self.capabilities = capabilities
|
|
43
|
+
|
|
44
|
+
def run(self):
|
|
45
|
+
while True:
|
|
46
|
+
req = read()
|
|
47
|
+
if req is None:
|
|
48
|
+
break
|
|
49
|
+
self._dispatch(req)
|
|
50
|
+
|
|
51
|
+
def _error_response(self, message: str) -> None:
|
|
52
|
+
sys.stderr.write(f"Error: {message}\n")
|
|
53
|
+
print(json.dumps({"error": message}), end='\0', flush=True)
|
|
54
|
+
|
|
55
|
+
def _dispatch(self, req: dict):
|
|
56
|
+
method = req.get('method')
|
|
57
|
+
if not method:
|
|
58
|
+
self._error_response("'method' field is required")
|
|
59
|
+
return
|
|
60
|
+
|
|
61
|
+
if method == 'cancel':
|
|
62
|
+
request_cancel()
|
|
63
|
+
return
|
|
64
|
+
|
|
65
|
+
reset_cancel()
|
|
66
|
+
|
|
67
|
+
try:
|
|
68
|
+
if method == 'capabilities':
|
|
69
|
+
handle_capabilities(self.capabilities)
|
|
70
|
+
|
|
71
|
+
elif method == 'format_test':
|
|
72
|
+
messages = req.get('messages')
|
|
73
|
+
if not messages:
|
|
74
|
+
self._error_response("'messages' field is required for format_test method")
|
|
75
|
+
return
|
|
76
|
+
handle_format_test(self.backend, self.capabilities, messages, req.get('options', {}), req.get('tools'))
|
|
77
|
+
|
|
78
|
+
elif method == 'tokenize':
|
|
79
|
+
messages = req.get('messages')
|
|
80
|
+
if messages is None:
|
|
81
|
+
self._error_response("'messages' field is required for tokenize method")
|
|
82
|
+
return
|
|
83
|
+
handle_tokenize(
|
|
84
|
+
self.backend, self.capabilities, messages,
|
|
85
|
+
tools=req.get('tools'),
|
|
86
|
+
reasoning_effort=req.get('reasoning_effort'),
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
elif method == 'cache_prefill':
|
|
90
|
+
self._error_response("cache_prefill is not supported by the PyTorch backend")
|
|
91
|
+
|
|
92
|
+
elif method == 'render':
|
|
93
|
+
messages = req.get('messages')
|
|
94
|
+
if not messages:
|
|
95
|
+
self._error_response("'messages' field is required for render method")
|
|
96
|
+
return
|
|
97
|
+
handle_render(
|
|
98
|
+
self.backend,
|
|
99
|
+
messages,
|
|
100
|
+
options=req.get('options', {}),
|
|
101
|
+
tools=req.get('tools'),
|
|
102
|
+
reasoning_effort=req.get('reasoning_effort'),
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
elif method == 'generate':
|
|
106
|
+
prompt = req.get('prompt')
|
|
107
|
+
if not _is_valid_generate_prompt(prompt):
|
|
108
|
+
self._error_response("'prompt' field is required for generate method")
|
|
109
|
+
return
|
|
110
|
+
images = req.get('images', [])
|
|
111
|
+
handle_generate(
|
|
112
|
+
self.backend,
|
|
113
|
+
prompt,
|
|
114
|
+
options=req.get('options', {}),
|
|
115
|
+
images=images if images else None,
|
|
116
|
+
max_image_size=req.get('maxImageSize', 768),
|
|
117
|
+
primer=req.get('primer'),
|
|
118
|
+
cache_path=req.get('cache_path'),
|
|
119
|
+
cache_trim_tokens=req.get('cache_trim_tokens'),
|
|
120
|
+
)
|
|
121
|
+
|
|
122
|
+
elif method == 'completion':
|
|
123
|
+
prompt = req.get('prompt')
|
|
124
|
+
if not prompt:
|
|
125
|
+
self._error_response("'prompt' field is required for completion method")
|
|
126
|
+
return
|
|
127
|
+
images = req.get('images', [])
|
|
128
|
+
handle_completion(
|
|
129
|
+
self.backend,
|
|
130
|
+
prompt,
|
|
131
|
+
options=req.get('options', {}),
|
|
132
|
+
images=images if images else None,
|
|
133
|
+
max_image_size=req.get('maxImageSize', 768),
|
|
134
|
+
)
|
|
135
|
+
|
|
136
|
+
else:
|
|
137
|
+
self._error_response(f"Unknown method '{method}'")
|
|
138
|
+
|
|
139
|
+
except Exception as e:
|
|
140
|
+
self._error_response(f"Error processing request: {e}")
|
|
File without changes
|
|
@@ -0,0 +1,164 @@
|
|
|
1
|
+
"""
|
|
2
|
+
チャットテンプレートの制約検出
|
|
3
|
+
|
|
4
|
+
tokenizerのapply_chat_templateを使用して、
|
|
5
|
+
モデルがサポートするメッセージパターンの制約を検出する。
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def detect_chat_restrictions(tokenizer) -> dict:
|
|
10
|
+
"""
|
|
11
|
+
チャットテンプレートの制約を検出
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
tokenizer: HuggingFace tokenizer (apply_chat_template対応)
|
|
15
|
+
|
|
16
|
+
Returns:
|
|
17
|
+
dict: chat_restrictions情報
|
|
18
|
+
{
|
|
19
|
+
"single_system_at_start": bool,
|
|
20
|
+
"max_system_messages": int,
|
|
21
|
+
"alternating_turns": bool,
|
|
22
|
+
"requires_user_last": bool,
|
|
23
|
+
"allow_empty_messages": bool
|
|
24
|
+
}
|
|
25
|
+
"""
|
|
26
|
+
if not hasattr(tokenizer, 'apply_chat_template'):
|
|
27
|
+
return None
|
|
28
|
+
|
|
29
|
+
# テストパターンを実行
|
|
30
|
+
test_results = {}
|
|
31
|
+
for pattern in _get_test_patterns():
|
|
32
|
+
try:
|
|
33
|
+
tokenizer.apply_chat_template(
|
|
34
|
+
pattern['messages'],
|
|
35
|
+
tokenize=False,
|
|
36
|
+
add_generation_prompt=False
|
|
37
|
+
)
|
|
38
|
+
test_results[pattern['name']] = {'success': True}
|
|
39
|
+
except Exception as e:
|
|
40
|
+
test_results[pattern['name']] = {'error': str(e)}
|
|
41
|
+
|
|
42
|
+
# テスト結果から制約を推論
|
|
43
|
+
return _infer_restrictions_from_results(test_results)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def _get_test_patterns():
|
|
47
|
+
"""テストパターンの定義"""
|
|
48
|
+
return [
|
|
49
|
+
# 基本パターン
|
|
50
|
+
{
|
|
51
|
+
'name': 'basic',
|
|
52
|
+
'messages': [
|
|
53
|
+
{'role': 'user', 'content': 'Hello'}
|
|
54
|
+
]
|
|
55
|
+
},
|
|
56
|
+
|
|
57
|
+
# システムメッセージ付き
|
|
58
|
+
{
|
|
59
|
+
'name': 'with-system',
|
|
60
|
+
'messages': [
|
|
61
|
+
{'role': 'system', 'content': 'You are a helpful assistant.'},
|
|
62
|
+
{'role': 'user', 'content': 'Hello'}
|
|
63
|
+
]
|
|
64
|
+
},
|
|
65
|
+
|
|
66
|
+
# 複数システムメッセージ
|
|
67
|
+
{
|
|
68
|
+
'name': 'multi-system',
|
|
69
|
+
'messages': [
|
|
70
|
+
{'role': 'system', 'content': 'First system message.'},
|
|
71
|
+
{'role': 'system', 'content': 'Second system message.'},
|
|
72
|
+
{'role': 'user', 'content': 'Hello'}
|
|
73
|
+
]
|
|
74
|
+
},
|
|
75
|
+
|
|
76
|
+
# 連続ユーザーメッセージ
|
|
77
|
+
{
|
|
78
|
+
'name': 'consecutive-user',
|
|
79
|
+
'messages': [
|
|
80
|
+
{'role': 'user', 'content': 'First question'},
|
|
81
|
+
{'role': 'user', 'content': 'Second question'}
|
|
82
|
+
]
|
|
83
|
+
},
|
|
84
|
+
|
|
85
|
+
# アシスタントで終わる
|
|
86
|
+
{
|
|
87
|
+
'name': 'assistant-last',
|
|
88
|
+
'messages': [
|
|
89
|
+
{'role': 'user', 'content': 'Hello'},
|
|
90
|
+
{'role': 'assistant', 'content': 'Hi there!'}
|
|
91
|
+
]
|
|
92
|
+
},
|
|
93
|
+
|
|
94
|
+
# 交互の会話
|
|
95
|
+
{
|
|
96
|
+
'name': 'alternating',
|
|
97
|
+
'messages': [
|
|
98
|
+
{'role': 'user', 'content': 'Question 1'},
|
|
99
|
+
{'role': 'assistant', 'content': 'Answer 1'},
|
|
100
|
+
{'role': 'user', 'content': 'Question 2'}
|
|
101
|
+
]
|
|
102
|
+
},
|
|
103
|
+
|
|
104
|
+
# 空メッセージ
|
|
105
|
+
{
|
|
106
|
+
'name': 'empty-message',
|
|
107
|
+
'messages': [
|
|
108
|
+
{'role': 'user', 'content': ''}
|
|
109
|
+
]
|
|
110
|
+
},
|
|
111
|
+
|
|
112
|
+
# システムメッセージが途中にある
|
|
113
|
+
{
|
|
114
|
+
'name': 'system-middle',
|
|
115
|
+
'messages': [
|
|
116
|
+
{'role': 'user', 'content': 'First'},
|
|
117
|
+
{'role': 'system', 'content': 'System in middle'},
|
|
118
|
+
{'role': 'user', 'content': 'Second'}
|
|
119
|
+
]
|
|
120
|
+
}
|
|
121
|
+
]
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
def _infer_restrictions_from_results(test_results: dict) -> dict:
|
|
125
|
+
"""
|
|
126
|
+
テスト結果から制約を推論
|
|
127
|
+
|
|
128
|
+
Args:
|
|
129
|
+
test_results: テストパターン名をキーとした結果の辞書
|
|
130
|
+
|
|
131
|
+
Returns:
|
|
132
|
+
dict: 検出された制約
|
|
133
|
+
"""
|
|
134
|
+
restrictions = {}
|
|
135
|
+
|
|
136
|
+
# システムメッセージの制約を検出
|
|
137
|
+
with_system = test_results.get('with-system')
|
|
138
|
+
multi_system = test_results.get('multi-system')
|
|
139
|
+
|
|
140
|
+
if with_system and 'error' in with_system:
|
|
141
|
+
# 単独のsystemメッセージもエラー → systemロール自体がサポートされていない
|
|
142
|
+
restrictions['max_system_messages'] = 0
|
|
143
|
+
elif multi_system and 'error' in multi_system:
|
|
144
|
+
# 複数はエラーだが単独は成功 → 最大1つまで
|
|
145
|
+
restrictions['single_system_at_start'] = True
|
|
146
|
+
restrictions['max_system_messages'] = 1
|
|
147
|
+
# それ以外(両方成功)→ max_system_messagesキーを設定しない(無制限)
|
|
148
|
+
|
|
149
|
+
# 連続ユーザーメッセージのテスト
|
|
150
|
+
consecutive_user = test_results.get('consecutive-user')
|
|
151
|
+
if consecutive_user and 'error' in consecutive_user:
|
|
152
|
+
restrictions['alternating_turns'] = True
|
|
153
|
+
|
|
154
|
+
# アシスタントで終わるテスト
|
|
155
|
+
assistant_last = test_results.get('assistant-last')
|
|
156
|
+
if assistant_last and 'error' in assistant_last:
|
|
157
|
+
restrictions['requires_user_last'] = True
|
|
158
|
+
|
|
159
|
+
# 空メッセージのテスト
|
|
160
|
+
empty_message = test_results.get('empty-message')
|
|
161
|
+
if empty_message and 'error' in empty_message:
|
|
162
|
+
restrictions['allow_empty_messages'] = False
|
|
163
|
+
|
|
164
|
+
return restrictions if restrictions else None
|
|
@@ -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
|