@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.
Files changed (192) hide show
  1. package/README.md +93 -10
  2. package/dist/cache-controller.d.ts +4 -0
  3. package/dist/cache-controller.d.ts.map +1 -1
  4. package/dist/driver-registry/config-based-factory.d.ts.map +1 -1
  5. package/dist/driver-registry/config-based-factory.js +6 -0
  6. package/dist/driver-registry/config-based-factory.js.map +1 -1
  7. package/dist/driver-registry/factory-helper.d.ts.map +1 -1
  8. package/dist/driver-registry/factory-helper.js +9 -2
  9. package/dist/driver-registry/factory-helper.js.map +1 -1
  10. package/dist/driver-registry/index.d.ts +1 -1
  11. package/dist/driver-registry/index.d.ts.map +1 -1
  12. package/dist/driver-registry/types.d.ts +15 -1
  13. package/dist/driver-registry/types.d.ts.map +1 -1
  14. package/dist/formatter/converter.d.ts.map +1 -1
  15. package/dist/formatter/converter.js +31 -2
  16. package/dist/formatter/converter.js.map +1 -1
  17. package/dist/index.d.ts +5 -3
  18. package/dist/index.d.ts.map +1 -1
  19. package/dist/index.js +5 -3
  20. package/dist/index.js.map +1 -1
  21. package/dist/local-inference/adapters.d.ts +6 -0
  22. package/dist/local-inference/adapters.d.ts.map +1 -1
  23. package/dist/local-inference/driver.d.ts.map +1 -1
  24. package/dist/local-inference/driver.js +45 -24
  25. package/dist/local-inference/driver.js.map +1 -1
  26. package/dist/local-inference/process-client.d.ts +4 -2
  27. package/dist/local-inference/process-client.d.ts.map +1 -1
  28. package/dist/local-inference/process-client.js +24 -8
  29. package/dist/local-inference/process-client.js.map +1 -1
  30. package/dist/local-inference/process-communication.d.ts +9 -2
  31. package/dist/local-inference/process-communication.d.ts.map +1 -1
  32. package/dist/local-inference/process-communication.js +37 -5
  33. package/dist/local-inference/process-communication.js.map +1 -1
  34. package/dist/local-inference/protocol.d.ts +4 -0
  35. package/dist/local-inference/protocol.d.ts.map +1 -1
  36. package/dist/local-inference/request-queue.d.ts +1 -1
  37. package/dist/local-inference/request-queue.d.ts.map +1 -1
  38. package/dist/local-inference/request-queue.js +26 -7
  39. package/dist/local-inference/request-queue.js.map +1 -1
  40. package/dist/local-inference/stream-utils.d.ts +6 -0
  41. package/dist/local-inference/stream-utils.d.ts.map +1 -1
  42. package/dist/local-inference/stream-utils.js.map +1 -1
  43. package/dist/mlx-ml/mlx-cache-controller.d.ts +9 -0
  44. package/dist/mlx-ml/mlx-cache-controller.d.ts.map +1 -1
  45. package/dist/mlx-ml/mlx-cache-controller.js +158 -33
  46. package/dist/mlx-ml/mlx-cache-controller.js.map +1 -1
  47. package/dist/mlx-ml/mlx-cache-support.d.ts +2 -2
  48. package/dist/mlx-ml/mlx-cache-support.d.ts.map +1 -1
  49. package/dist/mlx-ml/mlx-cache-support.js +8 -3
  50. package/dist/mlx-ml/mlx-cache-support.js.map +1 -1
  51. package/dist/mlx-ml/mlx-driver.d.ts +0 -1
  52. package/dist/mlx-ml/mlx-driver.d.ts.map +1 -1
  53. package/dist/mlx-ml/mlx-driver.js +1 -8
  54. package/dist/mlx-ml/mlx-driver.js.map +1 -1
  55. package/dist/mlx-ml/process/index.d.ts +1 -1
  56. package/dist/mlx-ml/process/index.d.ts.map +1 -1
  57. package/dist/mlx-ml/process/index.js +2 -2
  58. package/dist/mlx-ml/process/index.js.map +1 -1
  59. package/dist/models-config/index.d.ts +1 -1
  60. package/dist/models-config/index.d.ts.map +1 -1
  61. package/dist/models-config/index.js +1 -1
  62. package/dist/models-config/index.js.map +1 -1
  63. package/dist/models-config/resolve.d.ts +9 -1
  64. package/dist/models-config/resolve.d.ts.map +1 -1
  65. package/dist/models-config/resolve.js +94 -2
  66. package/dist/models-config/resolve.js.map +1 -1
  67. package/dist/models-config/types.d.ts +3 -1
  68. package/dist/models-config/types.d.ts.map +1 -1
  69. package/dist/pytorch/process/index.d.ts +4 -2
  70. package/dist/pytorch/process/index.d.ts.map +1 -1
  71. package/dist/pytorch/process/index.js +24 -7
  72. package/dist/pytorch/process/index.js.map +1 -1
  73. package/dist/pytorch/pytorch-cache-controller.d.ts +84 -0
  74. package/dist/pytorch/pytorch-cache-controller.d.ts.map +1 -0
  75. package/dist/pytorch/pytorch-cache-controller.js +742 -0
  76. package/dist/pytorch/pytorch-cache-controller.js.map +1 -0
  77. package/dist/pytorch/pytorch-cache-support.d.ts +23 -0
  78. package/dist/pytorch/pytorch-cache-support.d.ts.map +1 -0
  79. package/dist/pytorch/pytorch-cache-support.js +47 -0
  80. package/dist/pytorch/pytorch-cache-support.js.map +1 -0
  81. package/dist/pytorch/pytorch-driver.d.ts +8 -1
  82. package/dist/pytorch/pytorch-driver.d.ts.map +1 -1
  83. package/dist/pytorch/pytorch-driver.js +40 -0
  84. package/dist/pytorch/pytorch-driver.js.map +1 -1
  85. package/dist/runtime/check.d.ts.map +1 -1
  86. package/dist/runtime/check.js +8 -6
  87. package/dist/runtime/check.js.map +1 -1
  88. package/dist/runtime/index.d.ts +2 -2
  89. package/dist/runtime/index.d.ts.map +1 -1
  90. package/dist/runtime/index.js +2 -2
  91. package/dist/runtime/index.js.map +1 -1
  92. package/dist/runtime/manifest-core.d.mts +1 -0
  93. package/dist/runtime/manifest-core.mjs +1 -0
  94. package/dist/runtime/manifest-core.mjs.map +1 -1
  95. package/dist/runtime/manifest.d.ts +2 -0
  96. package/dist/runtime/manifest.d.ts.map +1 -1
  97. package/dist/runtime/manifest.js.map +1 -1
  98. package/dist/runtime/paths-core.d.mts +15 -1
  99. package/dist/runtime/paths-core.d.mts.map +1 -1
  100. package/dist/runtime/paths-core.mjs +50 -5
  101. package/dist/runtime/paths-core.mjs.map +1 -1
  102. package/dist/runtime/paths.d.ts +2 -2
  103. package/dist/runtime/paths.d.ts.map +1 -1
  104. package/dist/runtime/paths.js +2 -2
  105. package/dist/runtime/paths.js.map +1 -1
  106. package/dist/runtime/pytorch-template-core.d.mts +11 -0
  107. package/dist/runtime/pytorch-template-core.d.mts.map +1 -0
  108. package/dist/runtime/pytorch-template-core.mjs +54 -0
  109. package/dist/runtime/pytorch-template-core.mjs.map +1 -0
  110. package/dist/runtime/setup-commands-core.d.mts +3 -0
  111. package/dist/runtime/setup-commands-core.d.mts.map +1 -1
  112. package/dist/runtime/setup-commands-core.mjs +4 -0
  113. package/dist/runtime/setup-commands-core.mjs.map +1 -1
  114. package/dist/runtime/setup-commands.d.ts +1 -1
  115. package/dist/runtime/setup-commands.d.ts.map +1 -1
  116. package/dist/runtime/setup-commands.js +1 -1
  117. package/dist/runtime/setup-commands.js.map +1 -1
  118. package/docs/DRIVER_API.md +455 -0
  119. package/docs/LOCAL_MODEL_SETUP.md +765 -0
  120. package/docs/mlx-api-selection.md +301 -0
  121. package/package.json +9 -5
  122. package/scripts/runtime-cli.bin.test.ts +142 -0
  123. package/scripts/runtime-cli.js +305 -35
  124. package/scripts/runtime-cli.test.ts +163 -0
  125. package/src/mlx-ml/python/__main__.py +1 -1
  126. package/src/mlx-ml/python/backends/base.py +88 -18
  127. package/src/mlx-ml/python/backends/mlx_lm.py +28 -3
  128. package/src/mlx-ml/python/backends/mlx_vlm.py +679 -2
  129. package/src/mlx-ml/python/handlers/cache.py +4 -0
  130. package/src/mlx-ml/python/handlers/generate.py +33 -10
  131. package/src/mlx-ml/python/handlers/tokenize.py +1 -4
  132. package/src/mlx-ml/python/pyproject.toml +1 -1
  133. package/src/mlx-ml/python/server.py +2 -0
  134. package/src/mlx-ml/python/uv.lock +8 -8
  135. package/src/pytorch/templates/cpu-minimal/backends/base.py +139 -0
  136. package/src/pytorch/templates/cpu-minimal/backends/transformers_lm.py +1167 -0
  137. package/src/pytorch/{python → templates/cpu-minimal}/handlers/__init__.py +1 -0
  138. package/src/pytorch/templates/cpu-minimal/handlers/cache.py +88 -0
  139. package/src/pytorch/templates/cpu-minimal/handlers/generate.py +157 -0
  140. package/src/pytorch/{python → templates/cpu-minimal}/pyproject.toml +2 -2
  141. package/src/pytorch/{python → templates/cpu-minimal}/server.py +20 -2
  142. package/src/pytorch/templates/cpu-minimal/tests/test_cache_handler.py +284 -0
  143. package/src/pytorch/templates/cpu-minimal/tests/test_capabilities.py +14 -0
  144. package/src/pytorch/templates/cpu-minimal/tests/test_server.py +141 -0
  145. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_errors.py +89 -0
  146. package/src/pytorch/templates/cpu-minimal/tests/test_transformers_lm_cache.py +554 -0
  147. package/src/pytorch/templates/cpu-minimal/utils/__init__.py +0 -0
  148. package/src/pytorch/{python → templates/cpu-minimal}/utils/token_utils.py +2 -2
  149. package/src/pytorch/templates/cpu-minimal/utils/transformers_errors.py +54 -0
  150. package/src/pytorch/{python → templates/cpu-minimal}/uv.lock +149 -109
  151. package/src/pytorch/templates/cuda/__main__.py +19 -0
  152. package/src/pytorch/templates/cuda/backends/__init__.py +3 -0
  153. package/src/pytorch/{python → templates/cuda}/backends/base.py +54 -6
  154. package/src/pytorch/templates/cuda/backends/transformers_lm.py +379 -0
  155. package/src/pytorch/templates/cuda/handlers/__init__.py +7 -0
  156. package/src/pytorch/templates/cuda/handlers/cache.py +93 -0
  157. package/src/pytorch/templates/cuda/handlers/cancel.py +53 -0
  158. package/src/pytorch/templates/cuda/handlers/capabilities.py +6 -0
  159. package/src/pytorch/templates/cuda/handlers/completion.py +15 -0
  160. package/src/pytorch/templates/cuda/handlers/format_test.py +70 -0
  161. package/src/pytorch/templates/cuda/handlers/generate.py +152 -0
  162. package/src/pytorch/templates/cuda/handlers/render.py +40 -0
  163. package/src/pytorch/templates/cuda/handlers/tokenize.py +63 -0
  164. package/src/pytorch/templates/cuda/pyproject.toml +37 -0
  165. package/src/pytorch/templates/cuda/server.py +158 -0
  166. package/src/pytorch/templates/cuda/tests/__init__.py +0 -0
  167. package/src/pytorch/templates/cuda/tests/test_cache_handler.py +207 -0
  168. package/src/pytorch/templates/cuda/tests/test_capabilities.py +14 -0
  169. package/src/pytorch/templates/cuda/tests/test_server.py +145 -0
  170. package/src/pytorch/templates/cuda/tests/test_transformers_errors.py +89 -0
  171. package/src/pytorch/templates/cuda/tests/test_transformers_lm_cache.py +288 -0
  172. package/src/pytorch/templates/cuda/utils/__init__.py +0 -0
  173. package/src/pytorch/templates/cuda/utils/chat_template_constraints.py +164 -0
  174. package/src/pytorch/templates/cuda/utils/prompt_builder.py +54 -0
  175. package/src/pytorch/templates/cuda/utils/template_render.py +80 -0
  176. package/src/pytorch/templates/cuda/utils/token_utils.py +376 -0
  177. package/src/pytorch/templates/cuda/utils/transformers_errors.py +54 -0
  178. package/src/pytorch/templates/cuda/uv.lock +734 -0
  179. package/src/pytorch/python/backends/transformers_lm.py +0 -127
  180. package/src/pytorch/python/handlers/generate.py +0 -68
  181. /package/src/pytorch/{python → templates/cpu-minimal}/__main__.py +0 -0
  182. /package/src/pytorch/{python → templates/cpu-minimal}/backends/__init__.py +0 -0
  183. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/cancel.py +0 -0
  184. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/capabilities.py +0 -0
  185. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/completion.py +0 -0
  186. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/format_test.py +0 -0
  187. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/render.py +0 -0
  188. /package/src/pytorch/{python → templates/cpu-minimal}/handlers/tokenize.py +0 -0
  189. /package/src/pytorch/{python/utils → templates/cpu-minimal/tests}/__init__.py +0 -0
  190. /package/src/pytorch/{python → templates/cpu-minimal}/utils/chat_template_constraints.py +0 -0
  191. /package/src/pytorch/{python → templates/cpu-minimal}/utils/prompt_builder.py +0 -0
  192. /package/src/pytorch/{python → templates/cpu-minimal}/utils/template_render.py +0 -0
@@ -0,0 +1,145 @@
1
+ import pytest
2
+
3
+ pytest.importorskip("torch")
4
+
5
+ import json
6
+ from types import SimpleNamespace
7
+ from unittest.mock import MagicMock, patch
8
+
9
+
10
+ class _Tokenizer:
11
+ bos_token = "<bos>"
12
+ chat_template = "template"
13
+
14
+ def apply_chat_template(self, messages, **kwargs):
15
+ del messages, kwargs
16
+ return "prefix"
17
+
18
+
19
+ class _Backend:
20
+ model_kind = "lm"
21
+
22
+ def __init__(self):
23
+ self.tokenizer = _Tokenizer()
24
+ self.cache = None
25
+ self.generate_prompt = None
26
+ self.pending_cache_write_tokens = 0
27
+
28
+ def get_tokenizer(self):
29
+ return self.tokenizer
30
+
31
+ def cache_prefill(self, cache_path, prompt, **kwargs):
32
+ del prompt, kwargs
33
+ self.cache = object()
34
+ self.pending_cache_write_tokens = 2
35
+ return {
36
+ "cache_path": cache_path,
37
+ "token_count": 2,
38
+ "cache_write_tokens": 2,
39
+ }
40
+
41
+ def load_cache_from_file(self, cache_path, **kwargs):
42
+ del kwargs
43
+ return self.cache if cache_path == "memory://prefix" else None
44
+
45
+ def get_cache_offset(self, prompt_cache):
46
+ return 2 if prompt_cache is self.cache else 0
47
+
48
+ def consume_cache_write_tokens(self, cache_path):
49
+ del cache_path
50
+ result = self.pending_cache_write_tokens
51
+ self.pending_cache_write_tokens = 0
52
+ return result
53
+
54
+ def tokenize_prompt(self, prompt):
55
+ assert prompt == "prefix suffix"
56
+ return [1, 2, 3]
57
+
58
+ def stream_generate(self, prompt, options, images=None, prompt_cache=None):
59
+ del options, images
60
+ self.generate_prompt = prompt
61
+ assert prompt_cache is self.cache
62
+ yield SimpleNamespace(
63
+ text="ok",
64
+ prompt_tokens=3,
65
+ generation_tokens=1,
66
+ cache_read_tokens=2,
67
+ )
68
+
69
+
70
+ def _make_server():
71
+ from server import Server
72
+
73
+ backend = MagicMock()
74
+ return Server(backend, {"methods": ["capabilities"], "model_kind": "lm"})
75
+
76
+
77
+ def test_cache_prefill_dispatches_to_handler(capsys):
78
+ server = _make_server()
79
+
80
+ with patch("server.handle_cache_prefill") as mock_prefill:
81
+ server._dispatch(
82
+ {
83
+ "method": "cache_prefill",
84
+ "cache_path": "memory://prefix",
85
+ "messages": [{"role": "user", "content": "hello"}],
86
+ }
87
+ )
88
+
89
+ mock_prefill.assert_called_once()
90
+ assert mock_prefill.call_args.args[2:] == (
91
+ "memory://prefix",
92
+ [{"role": "user", "content": "hello"}],
93
+ )
94
+ assert capsys.readouterr().out == ""
95
+
96
+
97
+ def test_cache_prefill_requires_path_and_messages(capsys):
98
+ server = _make_server()
99
+
100
+ server._dispatch({"method": "cache_prefill"})
101
+
102
+ assert capsys.readouterr().out.endswith("\0")
103
+
104
+
105
+ def test_cache_prefill_then_generate_with_cache_via_server(capsys):
106
+ from server import Server
107
+
108
+ backend = _Backend()
109
+ server = Server(backend, {"methods": ["cache_prefill"], "model_kind": "lm"})
110
+
111
+ server._dispatch(
112
+ {
113
+ "method": "cache_prefill",
114
+ "cache_path": "memory://prefix",
115
+ "messages": [{"role": "user", "content": "hello"}],
116
+ }
117
+ )
118
+ prefill_output = capsys.readouterr().out
119
+ assert json.loads(prefill_output.split("\0", 1)[0]) == {
120
+ "cache_path": "memory://prefix",
121
+ "token_count": 2,
122
+ "cache_write_tokens": 2,
123
+ }
124
+
125
+ server._dispatch(
126
+ {
127
+ "method": "generate",
128
+ "prompt": "prefix suffix",
129
+ "options": {"max_tokens": 1},
130
+ "cache_path": "memory://prefix",
131
+ }
132
+ )
133
+ generate_output = capsys.readouterr().out
134
+ assert generate_output.startswith("ok")
135
+ meta = json.loads(
136
+ generate_output.split("\x1e__META__:", 1)[1].split("\0", 1)[0]
137
+ )
138
+ assert meta == {
139
+ "prompt_tokens": 3,
140
+ "generation_tokens": 1,
141
+ "cache_read_tokens": 2,
142
+ "cache_write_tokens": 2,
143
+ "cache_loaded": True,
144
+ }
145
+ assert backend.generate_prompt == [3]
@@ -0,0 +1,89 @@
1
+ from types import SimpleNamespace
2
+
3
+ import pytest
4
+
5
+ from utils.transformers_errors import (
6
+ MIN_TRANSFORMERS_VERSION,
7
+ extract_unsupported_model_type,
8
+ unsupported_model_type_error,
9
+ )
10
+
11
+
12
+ def test_extracts_model_type_from_transformers_value_error():
13
+ error = ValueError(
14
+ "The checkpoint has model type `qwen3_5` but Transformers does not "
15
+ "recognize this architecture."
16
+ )
17
+
18
+ assert extract_unsupported_model_type(error) == "qwen3_5"
19
+
20
+
21
+ def test_extracts_model_type_from_older_registry_key_error():
22
+ assert extract_unsupported_model_type(KeyError("qwen3_5")) == "qwen3_5"
23
+
24
+
25
+ def test_leaves_unrelated_errors_unchanged():
26
+ assert extract_unsupported_model_type(ValueError("invalid weights")) is None
27
+ assert extract_unsupported_model_type(KeyError("model_type")) is None
28
+
29
+
30
+ def test_error_includes_runtime_requirement_and_setup_guidance():
31
+ error = unsupported_model_type_error("Qwen/Qwen3.5-0.8B", "qwen3_5", "4.57.6")
32
+ message = str(error)
33
+
34
+ assert "qwen3_5" in message
35
+ assert "transformers 4.57.6" in message
36
+ assert f"transformers>={MIN_TRANSFORMERS_VERSION}" in message
37
+ assert "setup-pytorch" in message
38
+ assert "modular-prompt-runtime sync pytorch" in message
39
+
40
+
41
+ def test_backend_wraps_unknown_architecture_error(monkeypatch):
42
+ pytest.importorskip("torch")
43
+ import backends.transformers_lm as backend_module
44
+
45
+ tokenizer = SimpleNamespace(pad_token=None, eos_token="<eos>")
46
+ backend = backend_module.TransformersLmBackend(device="cpu")
47
+
48
+ def load_tokenizer(*args, **kwargs):
49
+ return tokenizer
50
+
51
+ def load_model(*args, **kwargs):
52
+ raise ValueError(
53
+ "The checkpoint has model type `qwen3_5` but Transformers does not "
54
+ "recognize this architecture."
55
+ )
56
+
57
+ monkeypatch.setattr(backend_module.AutoTokenizer, "from_pretrained", load_tokenizer)
58
+ monkeypatch.setattr(
59
+ backend_module.AutoModelForCausalLM,
60
+ "from_pretrained",
61
+ load_model,
62
+ )
63
+
64
+ with pytest.raises(RuntimeError, match="qwen3_5") as raised:
65
+ backend.load("Qwen/Qwen3.5-0.8B")
66
+
67
+ assert f"transformers {backend_module.transformers.__version__}" in str(raised.value)
68
+ assert "transformers>=5.14.0" in str(raised.value)
69
+
70
+
71
+ def test_backend_wraps_unknown_architecture_error_from_tokenizer(monkeypatch):
72
+ pytest.importorskip("torch")
73
+ import backends.transformers_lm as backend_module
74
+
75
+ backend = backend_module.TransformersLmBackend(device="cpu")
76
+
77
+ def load_tokenizer(*args, **kwargs):
78
+ raise ValueError(
79
+ "The checkpoint has model type `qwen3_5` but Transformers does not "
80
+ "recognize this architecture."
81
+ )
82
+
83
+ monkeypatch.setattr(backend_module.AutoTokenizer, "from_pretrained", load_tokenizer)
84
+
85
+ with pytest.raises(RuntimeError, match="qwen3_5") as raised:
86
+ backend.load("Qwen/Qwen3.5-0.8B")
87
+
88
+ assert f"transformers {backend_module.transformers.__version__}" in str(raised.value)
89
+ assert "transformers>=5.14.0" in str(raised.value)
@@ -0,0 +1,288 @@
1
+ import json
2
+ from types import SimpleNamespace
3
+
4
+ import pytest
5
+
6
+ torch = pytest.importorskip("torch")
7
+
8
+ from backends.transformers_lm import TransformersLmBackend
9
+ from handlers.generate import handle_generate
10
+ from tokenizers import Tokenizer
11
+ from tokenizers.models import WordLevel
12
+ from tokenizers.pre_tokenizers import Whitespace
13
+ from transformers import GPT2Config, GPT2LMHeadModel, PreTrainedTokenizerFast
14
+
15
+
16
+ class _Tokenizer:
17
+ bos_token = "<bos>"
18
+ pad_token = "<pad>"
19
+ eos_token = "<eos>"
20
+
21
+ def __init__(self):
22
+ self.prompts = {
23
+ "prefix": [1, 2],
24
+ "prefix suffix": [1, 2, 3],
25
+ "different suffix": [9, 8, 7],
26
+ }
27
+ self.token_text = {
28
+ 10: "a",
29
+ 11: "b",
30
+ 12: "c",
31
+ 13: "d",
32
+ }
33
+
34
+ def encode(self, prompt, add_special_tokens):
35
+ del add_special_tokens
36
+ return self.prompts[prompt]
37
+
38
+ def decode(self, token_ids, **kwargs):
39
+ del kwargs
40
+ if isinstance(token_ids, torch.Tensor):
41
+ token_ids = token_ids.flatten().tolist()
42
+ return "".join(self.token_text.get(int(token_id), "x") for token_id in token_ids)
43
+
44
+
45
+ class _Model:
46
+ def __init__(self, past_key_values, generated_texts=None, generated_token_batches=None):
47
+ self.past_key_values = past_key_values
48
+ self.generated_texts = generated_texts or ["generated"]
49
+ self.generated_token_batches = generated_token_batches
50
+ self.forward_calls = []
51
+ self.generate_calls = []
52
+
53
+ def __call__(self, **kwargs):
54
+ self.forward_calls.append(kwargs)
55
+ return SimpleNamespace(past_key_values=self.past_key_values)
56
+
57
+ def _emit_stream(self, streamer, input_ids, token_batches):
58
+ streamer.put(input_ids.cpu())
59
+ for token_batch in token_batches:
60
+ streamer.put(torch.tensor(token_batch, dtype=torch.long))
61
+ streamer.end()
62
+
63
+ def generate(self, **kwargs):
64
+ self.generate_calls.append(kwargs)
65
+ streamer = kwargs["streamer"]
66
+ if self.generated_token_batches is not None:
67
+ self._emit_stream(streamer, kwargs["input_ids"], self.generated_token_batches)
68
+ return
69
+ token_batches = [[10 + index] for index in range(len(self.generated_texts))]
70
+ self._emit_stream(streamer, kwargs["input_ids"], token_batches)
71
+
72
+
73
+ def _backend(generated_texts=None, generated_token_batches=None):
74
+ # Legacy tuple-shaped caches remain supported by Transformers and make the
75
+ # test independent of a specific Cache class implementation.
76
+ key = torch.zeros((1, 1, 2, 4))
77
+ value = torch.zeros((1, 1, 2, 4))
78
+ past_key_values = ((key, value),)
79
+
80
+ backend = TransformersLmBackend(device="cpu")
81
+ backend.tokenizer = _Tokenizer()
82
+ backend.model = _Model(past_key_values, generated_texts, generated_token_batches)
83
+ return backend, past_key_values
84
+
85
+
86
+ def _tiny_gpt2_backend():
87
+ vocab = {
88
+ "<pad>": 0,
89
+ "<unk>": 1,
90
+ "hello": 2,
91
+ "alpha": 3,
92
+ "beta": 4,
93
+ "gamma": 5,
94
+ }
95
+ tokenizer = Tokenizer(WordLevel(vocab=vocab, unk_token="<unk>"))
96
+ tokenizer.pre_tokenizer = Whitespace()
97
+ fast_tokenizer = PreTrainedTokenizerFast(
98
+ tokenizer_object=tokenizer,
99
+ unk_token="<unk>",
100
+ pad_token="<pad>",
101
+ )
102
+ config = GPT2Config(
103
+ vocab_size=len(vocab),
104
+ n_positions=32,
105
+ n_ctx=32,
106
+ n_embd=16,
107
+ n_layer=1,
108
+ n_head=1,
109
+ pad_token_id=vocab["<pad>"],
110
+ eos_token_id=None,
111
+ )
112
+ model = GPT2LMHeadModel(config)
113
+ # Make greedy output deterministic and non-special so the streamer emits
114
+ # several chunks without downloading a model in CI.
115
+ model.lm_head = torch.nn.Linear(config.n_embd, config.vocab_size, bias=True)
116
+ with torch.no_grad():
117
+ model.lm_head.weight.zero_()
118
+ model.lm_head.bias.fill_(-1)
119
+ model.lm_head.bias[vocab["alpha"]] = 1
120
+ model.eval()
121
+
122
+ backend = TransformersLmBackend()
123
+ backend.tokenizer = fast_tokenizer
124
+ backend.model = model
125
+ return backend
126
+
127
+
128
+ def test_cache_prefill_keeps_past_key_values_in_process():
129
+ backend, past_key_values = _backend()
130
+
131
+ result = backend.cache_prefill("memory://prefix", "prefix")
132
+
133
+ assert result == {
134
+ "cache_path": "memory://prefix",
135
+ "token_count": 2,
136
+ "cache_write_tokens": 2,
137
+ }
138
+ assert backend.load_cache_from_file("memory://prefix") is past_key_values
139
+ assert backend.consume_cache_write_tokens("memory://prefix") == 2
140
+ assert backend.consume_cache_write_tokens("memory://prefix") == 0
141
+ call = backend.model.forward_calls[0]
142
+ assert call["input_ids"].tolist() == [[1, 2]]
143
+ assert call["input_ids"].device == backend._device
144
+ assert call["use_cache"] is True
145
+
146
+
147
+ def test_load_cache_validates_prompt_prefix_without_reading_a_file():
148
+ backend, past_key_values = _backend()
149
+ backend.cache_prefill("memory://prefix", "prefix")
150
+
151
+ assert backend.load_cache_from_file(
152
+ "memory://prefix", prompt="prefix suffix"
153
+ ) is past_key_values
154
+ assert backend.load_cache_from_file(
155
+ "memory://prefix", prompt="different suffix"
156
+ ) is None
157
+ assert backend.load_cache_from_file("/tmp/not-created.cache") is None
158
+
159
+
160
+ def test_stream_generate_passes_cached_prefix_and_reports_usage():
161
+ backend, past_key_values = _backend()
162
+ backend.cache_prefill("memory://prefix", "prefix")
163
+
164
+ chunks = list(
165
+ backend.stream_generate(
166
+ [3],
167
+ {"max_tokens": 1, "temperature": 0},
168
+ prompt_cache=past_key_values,
169
+ )
170
+ )
171
+
172
+ assert chunks[-1].text == "a"
173
+ assert chunks[-1].generation_tokens == 1
174
+ call = backend.model.generate_calls[0]
175
+ assert call["input_ids"].tolist() == [[3]]
176
+ assert call["past_key_values"] is not past_key_values
177
+ assert torch.equal(call["past_key_values"][0][0], past_key_values[0][0])
178
+ assert call["cache_position"].tolist() == [2]
179
+ assert call["attention_mask"].tolist() == [[1, 1, 1]]
180
+ first_meta_chunk = next(chunk for chunk in chunks if chunk.prompt_tokens is not None)
181
+ assert first_meta_chunk.prompt_tokens == 3
182
+ assert first_meta_chunk.cache_read_tokens == 2
183
+
184
+
185
+ def test_stream_generate_reports_cumulative_generation_tokens_for_multiple_chunks():
186
+ backend, past_key_values = _backend(["a", "b", "c"])
187
+ backend.cache_prefill("memory://prefix", "prefix")
188
+
189
+ chunks = list(
190
+ backend.stream_generate(
191
+ [3],
192
+ {"max_tokens": 3, "temperature": 0},
193
+ prompt_cache=past_key_values,
194
+ )
195
+ )
196
+
197
+ assert chunks[-1].text == "abc"
198
+ assert chunks[-1].generation_tokens == 3
199
+ first_meta_chunk = next(chunk for chunk in chunks if chunk.prompt_tokens is not None)
200
+ assert first_meta_chunk.prompt_tokens == 3
201
+ assert first_meta_chunk.cache_read_tokens == 2
202
+
203
+
204
+ def test_stream_generate_counts_token_ids_not_empty_text_chunks(capsys):
205
+ backend, _ = _backend(generated_token_batches=[[10], [11], [12], [13]])
206
+
207
+ chunks = list(
208
+ backend.stream_generate(
209
+ "prefix suffix",
210
+ {"max_tokens": 4, "temperature": 0},
211
+ )
212
+ )
213
+
214
+ # TextIteratorStreamer emits an empty chunk while buffering each token and
215
+ # one final text chunk, so chunk count is five for four generated tokens.
216
+ assert len(chunks) == 5
217
+ assert chunks[0].text == ""
218
+ assert chunks[-1].text == "abcd"
219
+ assert chunks[-1].generation_tokens == 4
220
+
221
+ handle_generate(
222
+ backend,
223
+ "prefix suffix",
224
+ options={"max_tokens": 4, "temperature": 0},
225
+ )
226
+ output = capsys.readouterr().out
227
+ meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
228
+ assert meta["prompt_tokens"] == 3
229
+ assert meta["generation_tokens"] == 4
230
+
231
+
232
+ def test_generate_handler_reports_backend_prefill_write_usage(capsys):
233
+ backend, _ = _backend()
234
+ backend.cache_prefill("memory://prefix", "prefix")
235
+
236
+ handle_generate(
237
+ backend,
238
+ "prefix suffix",
239
+ options={"max_tokens": 1, "temperature": 0},
240
+ cache_path="memory://prefix",
241
+ )
242
+
243
+ output = capsys.readouterr().out
244
+ assert output.startswith("a")
245
+ meta = output.split("\x1e__META__:", 1)[1].split("\0", 1)[0]
246
+ assert json.loads(meta) == {
247
+ "prompt_tokens": 3,
248
+ "generation_tokens": 1,
249
+ "cache_read_tokens": 2,
250
+ "cache_write_tokens": 2,
251
+ "cache_loaded": True,
252
+ }
253
+
254
+
255
+ def test_generate_handler_preserves_usage_for_tiny_gpt2_multiple_chunks(capsys):
256
+ backend = _tiny_gpt2_backend()
257
+ options = {"max_tokens": 4, "temperature": 0}
258
+
259
+ expected_chunks = list(backend.stream_generate("hello", options))
260
+ # The real TextIteratorStreamer buffers the generated word, so the first
261
+ # text chunk is empty and the final chunk flushes all four generated
262
+ # tokens. Token usage must not be derived from this five-chunk sequence.
263
+ assert len(expected_chunks) == 5
264
+ assert expected_chunks[0].text == ""
265
+ expected_generation_tokens = expected_chunks[-1].generation_tokens
266
+ assert expected_generation_tokens == 4
267
+
268
+ handle_generate(backend, "hello", options=options)
269
+
270
+ output = capsys.readouterr().out
271
+ meta = json.loads(output.split("\x1e__META__:", 1)[1].split("\0", 1)[0])
272
+ assert meta["prompt_tokens"] == len(backend.tokenize_prompt("hello"))
273
+ assert meta["generation_tokens"] == expected_generation_tokens
274
+
275
+
276
+ def test_cache_prefill_rejects_phase_two_arguments():
277
+ backend, _ = _backend()
278
+
279
+ try:
280
+ backend.cache_prefill(
281
+ "memory://prefix",
282
+ "prefix",
283
+ base_cache_path="memory://base",
284
+ )
285
+ except ValueError as error:
286
+ assert "incremental prefill" in str(error)
287
+ else:
288
+ raise AssertionError("Phase 2 incremental prefill must be rejected")
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