Repository navigation
Expand file tree
/
Copy pathruntime_meta.py
More file actions
238 lines (196 loc) · 8.56 KB
/
Copy pathruntime_meta.py
File metadata and controls
238 lines (196 loc) · 8.56 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.
"""Shared runtime helpers for the MLX LLM example runners.
Exports publish their limits as constant methods (``get_max_context_len``,
``get_max_seq_len``) so runners do not have to be told what a .pte
supports. This mirrors ``const_int`` in run_llm_hf.cpp.
Prompt handling (processor loading, chat templating, EOS lookup) lives here too:
it is driven by the checkpoint's own tokenizer config, so it is the same for
every model and every runner.
"""
import json
import logging
from typing import Optional, Sequence, Tuple
import torch
from executorch.extension.llm.export.model_metadata import (
MAX_CONTEXT_LEN_METHOD,
MAX_SEQ_LEN_METHOD,
)
logger = logging.getLogger(__name__)
def read_const_int(program, name: str) -> Optional[int]:
"""Read an int constant method from a loaded program, or None if absent."""
try:
result = program.load_method(name).execute([])
except Exception:
return None
if not result:
return None
value = result[0]
return int(value) if isinstance(value, int) else int(value.item())
def read_model_limits(program) -> Tuple[Optional[int], Optional[int]]:
"""Return (max_context_len, max_seq_len) as published by the export."""
return (
read_const_int(program, MAX_CONTEXT_LEN_METHOD),
read_const_int(program, MAX_SEQ_LEN_METHOD),
)
def chunked_prefill(
method,
input_ids: torch.Tensor,
chunk_size: int,
start_pos: int = 0,
concat_outputs: Sequence[int] = (),
):
"""Prefill ``input_ids`` in steps of at most ``chunk_size`` tokens.
Returns the final chunk's outputs, except that each index in
``concat_outputs`` is replaced by that output concatenated along dim 1
across every chunk — used for per-token outputs like tapped hidden states,
where the caller needs the whole prompt rather than the last step.
"""
if chunk_size < 1:
raise ValueError(f"chunk_size must be >= 1, got {chunk_size}")
seq_len = input_ids.shape[1]
collected = {i: [] for i in concat_outputs}
outputs = None
for off in range(0, seq_len, chunk_size):
end = min(off + chunk_size, seq_len)
tokens = input_ids[:, off:end].contiguous()
cache_position = torch.arange(
start_pos + off, start_pos + end, dtype=torch.long
)
outputs = method.execute([tokens, cache_position.contiguous()])
for i in collected:
collected[i].append(outputs[i])
if outputs is None:
raise ValueError("input_ids is empty; nothing to prefill")
outputs = list(outputs)
for i, parts in collected.items():
outputs[i] = torch.cat(parts, dim=1) if len(parts) > 1 else parts[0]
return outputs
def load_text_processor(
model_id: str,
revision: Optional[str] = None,
local_files_only: bool = False,
):
"""Load the tokenizer for ``model_id``, falling back to its processor.
Prefer AutoTokenizer for text-only prompting, even for checkpoints that also
ship an AutoProcessor. Some hybrid checkpoints (for example Gemma 4) expose
both, but the tokenizer path is the more stable interface for plain text
generation.
"""
from transformers import AutoProcessor, AutoTokenizer
logger.info(f"Loading tokenizer from HuggingFace: {model_id}...")
processor = None
try:
processor = AutoTokenizer.from_pretrained(
model_id, revision=revision, local_files_only=local_files_only
)
except Exception as exc:
logger.info(f"AutoTokenizer unavailable for {model_id}: {exc}")
if processor is None:
try:
candidate = AutoProcessor.from_pretrained(
model_id, revision=revision, local_files_only=local_files_only
)
if hasattr(candidate, "apply_chat_template") and hasattr(
candidate, "decode"
):
logger.info(f"Loaded processor from HuggingFace: {model_id}")
processor = candidate
except Exception as exc:
logger.info(f"AutoProcessor unavailable for {model_id}: {exc}")
if processor is None:
raise RuntimeError(f"Could not load tokenizer or processor for {model_id}")
_repair_chat_template(processor, model_id, local_files_only)
return processor
def _repair_chat_template(processor, model_id: str, local_files_only: bool) -> None:
"""Backfill chat_template from tokenizer_config.json when it did not load.
Some transformers versions expect a standalone chat_template.jinja and do
not fall back to the embedded tokenizer_config.json field.
"""
if getattr(processor, "chat_template", None) is not None:
return
try:
from pathlib import Path
from huggingface_hub import hf_hub_download
cfg_path = hf_hub_download(
model_id, "tokenizer_config.json", local_files_only=local_files_only
)
cfg = json.loads(Path(cfg_path).read_text())
except Exception as exc:
logger.info(f"Could not backfill chat_template for {model_id}: {exc}")
return
template = cfg.get("chat_template")
# Only a non-empty string is usable; the multi-template list form would
# render to nothing and surface much later as an empty prompt.
if isinstance(template, str) and template.strip():
processor.chat_template = template
logger.info(f"Backfilled chat_template for {model_id} from tokenizer_config")
else:
logger.info(
f"{model_id} tokenizer_config has no usable chat_template "
f"(got {type(template).__name__})"
)
def apply_chat_template(
processor, prompt: str, enable_thinking: bool = False
) -> torch.Tensor:
"""Render ``prompt`` through the checkpoint's chat template into input ids.
``enable_thinking`` is a template variable rather than a tokenizer argument;
checkpoints whose template does not declare it are retried without it.
"""
messages = [{"role": "user", "content": prompt}]
kwargs = {"add_generation_prompt": True, "tokenize": True, "return_tensors": "pt"}
try:
out = processor.apply_chat_template(
messages, enable_thinking=enable_thinking, **kwargs
)
except TypeError:
out = processor.apply_chat_template(messages, **kwargs)
# Different transformers versions return either a BatchEncoding or a tensor.
input_ids = out.input_ids if hasattr(out, "input_ids") else out
# A template that renders to nothing tokenizes to shape (1, 0) instead of
# raising, which would otherwise surface far downstream as an empty prefill.
if input_ids.numel() == 0:
try:
rendered = processor.apply_chat_template(
messages, add_generation_prompt=True, tokenize=False
)
detail = f"the template rendered {len(rendered)} characters"
except Exception as exc:
detail = f"re-rendering it also failed ({exc})"
raise ValueError(
f"Chat template produced no tokens for prompt {prompt!r}: {detail}. "
"The tokenizer likely loaded without a usable chat_template. "
"Re-run with --no-chat-template to bypass it."
)
return input_ids
def get_eos_token_ids(processor, model_id=None, local_files_only=False):
"""Collect every id that should stop generation.
A checkpoint can declare several: Qwen3 stops on both ``<|im_end|>`` and
``<|endoftext|>``, but only the former is ``tokenizer.eos_token_id``. The
rest live in generation_config, so pass ``model_id`` to pick them up.
"""
eos_ids = set()
candidate = getattr(processor, "eos_token_id", None)
if candidate is None:
candidate = getattr(getattr(processor, "tokenizer", None), "eos_token_id", None)
eos_ids.update(_as_id_set(candidate))
if model_id is not None:
try:
from transformers import GenerationConfig
generation_config = GenerationConfig.from_pretrained(
model_id, local_files_only=local_files_only
)
eos_ids.update(_as_id_set(generation_config.eos_token_id))
except Exception as exc:
logger.info(f"No generation_config eos ids for {model_id}: {exc}")
return eos_ids
def _as_id_set(value):
"""Normalize an int / list / None token-id field into a set of ints."""
if value is None:
return set()
if isinstance(value, int):
return {value}
return {int(v) for v in value if v is not None}