@modular-prompt/driver 0.14.0 → 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.
Files changed (205) hide show
  1. package/README.md +58 -5
  2. package/dist/driver-registry/ai-service.d.ts +23 -1
  3. package/dist/driver-registry/ai-service.d.ts.map +1 -1
  4. package/dist/driver-registry/ai-service.js +44 -10
  5. package/dist/driver-registry/ai-service.js.map +1 -1
  6. package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
  7. package/dist/driver-registry/config-based-factory.js +16 -0
  8. package/dist/driver-registry/config-based-factory.js.map +1 -1
  9. package/dist/driver-registry/factory-helper.d.ts +2 -0
  10. package/dist/driver-registry/factory-helper.d.ts.map +1 -1
  11. package/dist/driver-registry/factory-helper.js +21 -3
  12. package/dist/driver-registry/factory-helper.js.map +1 -1
  13. package/dist/driver-registry/index.d.ts +2 -2
  14. package/dist/driver-registry/index.d.ts.map +1 -1
  15. package/dist/driver-registry/index.js +1 -1
  16. package/dist/driver-registry/index.js.map +1 -1
  17. package/dist/driver-registry/registry.d.ts.map +1 -1
  18. package/dist/driver-registry/registry.js +3 -1
  19. package/dist/driver-registry/registry.js.map +1 -1
  20. package/dist/driver-registry/types.d.ts +8 -2
  21. package/dist/driver-registry/types.d.ts.map +1 -1
  22. package/dist/index.d.ts +10 -2
  23. package/dist/index.d.ts.map +1 -1
  24. package/dist/index.js +9 -1
  25. package/dist/index.js.map +1 -1
  26. package/dist/local-inference/adapters.d.ts +66 -0
  27. package/dist/local-inference/adapters.d.ts.map +1 -0
  28. package/dist/local-inference/adapters.js +2 -0
  29. package/dist/local-inference/adapters.js.map +1 -0
  30. package/dist/local-inference/driver.d.ts +51 -0
  31. package/dist/local-inference/driver.d.ts.map +1 -0
  32. package/dist/local-inference/driver.js +309 -0
  33. package/dist/local-inference/driver.js.map +1 -0
  34. package/dist/local-inference/index.d.ts +22 -0
  35. package/dist/local-inference/index.d.ts.map +1 -0
  36. package/dist/local-inference/index.js +17 -0
  37. package/dist/local-inference/index.js.map +1 -0
  38. package/dist/local-inference/process-client.d.ts +50 -0
  39. package/dist/local-inference/process-client.d.ts.map +1 -0
  40. package/dist/local-inference/process-client.js +92 -0
  41. package/dist/local-inference/process-client.js.map +1 -0
  42. package/dist/local-inference/process-communication.d.ts +41 -0
  43. package/dist/local-inference/process-communication.d.ts.map +1 -0
  44. package/dist/{mlx-ml/process → local-inference}/process-communication.js +25 -61
  45. package/dist/local-inference/process-communication.js.map +1 -0
  46. package/dist/local-inference/process-port.d.ts +12 -0
  47. package/dist/local-inference/process-port.d.ts.map +1 -0
  48. package/dist/local-inference/process-port.js +2 -0
  49. package/dist/local-inference/process-port.js.map +1 -0
  50. package/dist/local-inference/prompt-utils.d.ts +6 -0
  51. package/dist/local-inference/prompt-utils.d.ts.map +1 -0
  52. package/dist/local-inference/prompt-utils.js +17 -0
  53. package/dist/local-inference/prompt-utils.js.map +1 -0
  54. package/dist/local-inference/protocol.d.ts +192 -0
  55. package/dist/local-inference/protocol.d.ts.map +1 -0
  56. package/dist/local-inference/protocol.js +2 -0
  57. package/dist/local-inference/protocol.js.map +1 -0
  58. package/dist/local-inference/queue-types.d.ts +54 -0
  59. package/dist/local-inference/queue-types.d.ts.map +1 -0
  60. package/dist/local-inference/queue-types.js +2 -0
  61. package/dist/local-inference/queue-types.js.map +1 -0
  62. package/dist/local-inference/request-queue.d.ts +36 -0
  63. package/dist/local-inference/request-queue.d.ts.map +1 -0
  64. package/dist/{mlx-ml/process/queue.js → local-inference/request-queue.js} +83 -56
  65. package/dist/local-inference/request-queue.js.map +1 -0
  66. package/dist/local-inference/stream-utils.d.ts +19 -0
  67. package/dist/local-inference/stream-utils.d.ts.map +1 -0
  68. package/dist/local-inference/stream-utils.js +76 -0
  69. package/dist/local-inference/stream-utils.js.map +1 -0
  70. package/dist/mlx-ml/mlx-cache-support.d.ts +23 -0
  71. package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -0
  72. package/dist/mlx-ml/mlx-cache-support.js +45 -0
  73. package/dist/mlx-ml/mlx-cache-support.js.map +1 -0
  74. package/dist/mlx-ml/mlx-driver.d.ts +20 -59
  75. package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
  76. package/dist/mlx-ml/mlx-driver.js +86 -460
  77. package/dist/mlx-ml/mlx-driver.js.map +1 -1
  78. package/dist/mlx-ml/mlx-local-inference-adapters.d.ts +3 -0
  79. package/dist/mlx-ml/mlx-local-inference-adapters.d.ts.map +1 -0
  80. package/dist/mlx-ml/mlx-local-inference-adapters.js +19 -0
  81. package/dist/mlx-ml/mlx-local-inference-adapters.js.map +1 -0
  82. package/dist/mlx-ml/mlx-options.d.ts +19 -0
  83. package/dist/mlx-ml/mlx-options.d.ts.map +1 -0
  84. package/dist/mlx-ml/mlx-options.js +30 -0
  85. package/dist/mlx-ml/mlx-options.js.map +1 -0
  86. package/dist/mlx-ml/process/index.d.ts +10 -8
  87. package/dist/mlx-ml/process/index.d.ts.map +1 -1
  88. package/dist/mlx-ml/process/index.js +73 -54
  89. package/dist/mlx-ml/process/index.js.map +1 -1
  90. package/dist/mlx-ml/process/model-specific.d.ts +2 -1
  91. package/dist/mlx-ml/process/model-specific.d.ts.map +1 -1
  92. package/dist/mlx-ml/process/model-specific.js.map +1 -1
  93. package/dist/mlx-ml/process/prompt-builder.d.ts +11 -0
  94. package/dist/mlx-ml/process/prompt-builder.d.ts.map +1 -0
  95. package/dist/mlx-ml/process/prompt-builder.js +51 -0
  96. package/dist/mlx-ml/process/prompt-builder.js.map +1 -0
  97. package/dist/mlx-ml/process/types.d.ts +15 -183
  98. package/dist/mlx-ml/process/types.d.ts.map +1 -1
  99. package/dist/mlx-ml/types.d.ts +2 -45
  100. package/dist/mlx-ml/types.d.ts.map +1 -1
  101. package/dist/models-config/index.d.ts +8 -0
  102. package/dist/models-config/index.d.ts.map +1 -0
  103. package/dist/models-config/index.js +7 -0
  104. package/dist/models-config/index.js.map +1 -0
  105. package/dist/models-config/loader.d.ts +20 -0
  106. package/dist/models-config/loader.d.ts.map +1 -0
  107. package/dist/models-config/loader.js +85 -0
  108. package/dist/models-config/loader.js.map +1 -0
  109. package/dist/models-config/paths.d.ts +7 -0
  110. package/dist/models-config/paths.d.ts.map +1 -0
  111. package/dist/models-config/paths.js +11 -0
  112. package/dist/models-config/paths.js.map +1 -0
  113. package/dist/models-config/resolve.d.ts +57 -0
  114. package/dist/models-config/resolve.d.ts.map +1 -0
  115. package/dist/models-config/resolve.js +187 -0
  116. package/dist/models-config/resolve.js.map +1 -0
  117. package/dist/models-config/types.d.ts +60 -0
  118. package/dist/models-config/types.d.ts.map +1 -0
  119. package/dist/models-config/types.js +5 -0
  120. package/dist/models-config/types.js.map +1 -0
  121. package/dist/pytorch/process/index.d.ts +35 -0
  122. package/dist/pytorch/process/index.d.ts.map +1 -0
  123. package/dist/pytorch/process/index.js +69 -0
  124. package/dist/pytorch/process/index.js.map +1 -0
  125. package/dist/pytorch/pytorch-driver.d.ts +35 -0
  126. package/dist/pytorch/pytorch-driver.d.ts.map +1 -0
  127. package/dist/pytorch/pytorch-driver.js +48 -0
  128. package/dist/pytorch/pytorch-driver.js.map +1 -0
  129. package/dist/pytorch/pytorch-local-inference-adapters.d.ts +3 -0
  130. package/dist/pytorch/pytorch-local-inference-adapters.d.ts.map +1 -0
  131. package/dist/pytorch/pytorch-local-inference-adapters.js +19 -0
  132. package/dist/pytorch/pytorch-local-inference-adapters.js.map +1 -0
  133. package/dist/pytorch/pytorch-options.d.ts +8 -0
  134. package/dist/pytorch/pytorch-options.d.ts.map +1 -0
  135. package/dist/pytorch/pytorch-options.js +21 -0
  136. package/dist/pytorch/pytorch-options.js.map +1 -0
  137. package/dist/query-logger.js +1 -1
  138. package/dist/query-logger.js.map +1 -1
  139. package/dist/runtime/check.d.ts +10 -0
  140. package/dist/runtime/check.d.ts.map +1 -0
  141. package/dist/runtime/check.js +24 -0
  142. package/dist/runtime/check.js.map +1 -0
  143. package/dist/runtime/index.d.ts +4 -0
  144. package/dist/runtime/index.d.ts.map +1 -0
  145. package/dist/runtime/index.js +4 -0
  146. package/dist/runtime/index.js.map +1 -0
  147. package/dist/runtime/manifest-core.d.mts +27 -0
  148. package/dist/runtime/manifest-core.d.mts.map +1 -0
  149. package/dist/runtime/manifest-core.mjs +68 -0
  150. package/dist/runtime/manifest-core.mjs.map +1 -0
  151. package/dist/runtime/manifest.d.ts +18 -0
  152. package/dist/runtime/manifest.d.ts.map +1 -0
  153. package/dist/runtime/manifest.js +9 -0
  154. package/dist/runtime/manifest.js.map +1 -0
  155. package/dist/runtime/paths-core.d.mts +17 -0
  156. package/dist/runtime/paths-core.d.mts.map +1 -0
  157. package/dist/runtime/paths-core.mjs +67 -0
  158. package/dist/runtime/paths-core.mjs.map +1 -0
  159. package/dist/runtime/paths.d.ts +12 -0
  160. package/dist/runtime/paths.d.ts.map +1 -0
  161. package/dist/runtime/paths.js +17 -0
  162. package/dist/runtime/paths.js.map +1 -0
  163. package/dist/types.d.ts +11 -1
  164. package/dist/types.d.ts.map +1 -1
  165. package/dist/types.js.map +1 -1
  166. package/package.json +10 -7
  167. package/scripts/download-model.js +25 -9
  168. package/scripts/runtime-cli.js +315 -0
  169. package/src/mlx-ml/python/__main__.py +43 -4
  170. package/src/mlx-ml/python/handlers/__init__.py +2 -1
  171. package/src/mlx-ml/python/handlers/completion.py +3 -27
  172. package/src/mlx-ml/python/handlers/{chat.py → generate.py} +33 -106
  173. package/src/mlx-ml/python/handlers/render.py +40 -0
  174. package/src/mlx-ml/python/pyproject.toml +3 -2
  175. package/src/mlx-ml/python/server.py +28 -8
  176. package/src/mlx-ml/python/utils/template_render.py +80 -0
  177. package/src/mlx-ml/python/utils/token_utils.py +2 -2
  178. package/src/mlx-ml/python/uv.lock +549 -454
  179. package/src/pytorch/python/__main__.py +19 -0
  180. package/src/pytorch/python/backends/__init__.py +3 -0
  181. package/src/pytorch/python/backends/base.py +84 -0
  182. package/src/pytorch/python/backends/transformers_lm.py +127 -0
  183. package/src/pytorch/python/handlers/__init__.py +6 -0
  184. package/src/pytorch/python/handlers/cancel.py +53 -0
  185. package/src/pytorch/python/handlers/capabilities.py +6 -0
  186. package/src/pytorch/python/handlers/completion.py +15 -0
  187. package/src/pytorch/python/handlers/format_test.py +70 -0
  188. package/src/pytorch/python/handlers/generate.py +68 -0
  189. package/src/pytorch/python/handlers/render.py +40 -0
  190. package/src/pytorch/python/handlers/tokenize.py +63 -0
  191. package/src/pytorch/python/pyproject.toml +36 -0
  192. package/src/pytorch/python/server.py +140 -0
  193. package/src/pytorch/python/utils/__init__.py +0 -0
  194. package/src/pytorch/python/utils/chat_template_constraints.py +164 -0
  195. package/src/pytorch/python/utils/prompt_builder.py +54 -0
  196. package/src/pytorch/python/utils/template_render.py +80 -0
  197. package/src/pytorch/python/utils/token_utils.py +376 -0
  198. package/src/pytorch/python/uv.lock +694 -0
  199. package/dist/mlx-ml/process/process-communication.d.ts +0 -45
  200. package/dist/mlx-ml/process/process-communication.d.ts.map +0 -1
  201. package/dist/mlx-ml/process/process-communication.js.map +0 -1
  202. package/dist/mlx-ml/process/queue.d.ts +0 -35
  203. package/dist/mlx-ml/process/queue.d.ts.map +0 -1
  204. package/dist/mlx-ml/process/queue.js.map +0 -1
  205. 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