fix(chat_templates): check find() return value before slicing on placeholders (#5763)
* fix(chat_templates): check find() return value before slicing on placeholders
Two places in `construct_chat_template()` use `str.find()` for sentinel
placeholders (`{INPUT}` / `{OUTPUT}`) without checking the -1 return:
1. The `except:` fallback (around line 2464) computes
`chat_template[chat_template.find("{OUTPUT}") + len("{OUTPUT}"):]`.
If the template has no `{OUTPUT}` marker, `find()` returns -1 and the
slice starts at offset 7 (`-1 + len("{OUTPUT}")`), producing garbage
that's then `re.escape`-d and fed back into the template-recovery
regex. The user sees a confusing `IndexError` on
`response_part = response_part[0]` instead of the real problem.
2. The final trim before returning (`input_part[:input_part.find("{INPUT}")]`
and the matching `{OUTPUT}` line) silently drops the last character
when the placeholder is missing — `find()` returns -1, and `[:-1]`
slices everything except the last character, returning a corrupted
template prefix to the caller.
Replace both with an explicit `-1` check that raises a clear
`RuntimeError` naming the missing placeholder, matching the existing
guard pattern from #4923 (`try_fix_tokenizer`).
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
* fix(chat_templates): also guard {INPUT} and fallback regex/separator paths
Builds on the {OUTPUT} / final-trim guards in this branch by closing
the three remaining ways the except-block fallback in
construct_chat_template() can still raise a confusing IndexError or
AttributeError on malformed templates:
1. Validate both {INPUT} and {OUTPUT} before deriving `ending`. The
regex two lines later (`{INPUT} + ending + ...`) still produced an
empty list and crashed on `response_part[0]` if {INPUT} was missing.
2. Guard the regex no-match case. Some templates contain both
placeholders but not in a recoverable two-example shape, in which
case `re.findall` returns an empty list and `[0]` raises.
3. Initialize `found = None` before the separator-search loop and
raise if the loop never sets it. Previously, if the first
iteration's `re.finditer` was empty the loop broke without binding
`found`, and `found.group(1)` raised AttributeError on the stale
int left over from the outer rfind loop.
Rephrase the final-trim error messages from internal variable names
("input_part") to user-facing wording ("instruction section") and
include a bounded (200-char) excerpt of the offending content so the
error is debuggable without being unbounded.
Add tests/python/test_construct_chat_template_validation.py covering
each failure mode with a fake tokenizer (no HF_TOKEN, no model
download, CPU-only).
* [pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
---------
Co-authored-by: Claude Opus 4.7 <noreply@anthropic.com>
Co-authored-by: Daniel Han <danielhanchen@gmail.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
parent
748aa1c482
commit
af6504f900
2 changed files with 115 additions and 3 deletions
77
tests/python/test_construct_chat_template_validation.py
Normal file
77
tests/python/test_construct_chat_template_validation.py
Normal file
|
|
@ -0,0 +1,77 @@
|
||||||
|
"""Negative-path validation tests for unsloth.chat_templates.construct_chat_template.
|
||||||
|
|
||||||
|
Regression coverage for the str.find() / regex no-match guards added in
|
||||||
|
PR #5763 follow-up: missing placeholders or unrecoverable two-example
|
||||||
|
structures must raise RuntimeError with a clear message, not IndexError
|
||||||
|
or AttributeError, and must never silently drop the last character via
|
||||||
|
s[:-1].
|
||||||
|
|
||||||
|
Uses a minimal fake tokenizer so the cases run on CPU-only CI without
|
||||||
|
HF_TOKEN and without downloading a gated model. The validation paths
|
||||||
|
exercised here fail before construct_chat_template reaches any heavy
|
||||||
|
tokenizer interaction, so the stub stays small.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from unsloth.chat_templates import construct_chat_template
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeTokenizer:
|
||||||
|
"""Minimum surface construct_chat_template touches before the
|
||||||
|
validation guards fire."""
|
||||||
|
|
||||||
|
name_or_path = "fake/tokenizer"
|
||||||
|
eos_token = "</s>"
|
||||||
|
|
||||||
|
def get_vocab(self):
|
||||||
|
return {"</s>": 0}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"template, expected_in_message",
|
||||||
|
[
|
||||||
|
("only {INPUT} here, no output marker", "{OUTPUT}"),
|
||||||
|
("only {OUTPUT} here, no input marker", "{INPUT}"),
|
||||||
|
("neither sentinel here, just literal text", "{INPUT}"),
|
||||||
|
("neither sentinel here, just literal text", "{OUTPUT}"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_missing_placeholder_in_chat_template_raises(template, expected_in_message):
|
||||||
|
with pytest.raises(RuntimeError) as exc_info:
|
||||||
|
construct_chat_template(
|
||||||
|
tokenizer = _FakeTokenizer(),
|
||||||
|
chat_template = template,
|
||||||
|
extra_eos_tokens = ["</s>"],
|
||||||
|
)
|
||||||
|
assert expected_in_message in str(exc_info.value)
|
||||||
|
|
||||||
|
|
||||||
|
def test_single_pair_template_raises_clear_error_not_attribute_error():
|
||||||
|
"""One {INPUT}/{OUTPUT} pair (rather than the required two) used to
|
||||||
|
crash with AttributeError on `found.group(1)` after the for-loop
|
||||||
|
broke without setting `found`. Must raise RuntimeError now."""
|
||||||
|
template = "user: {INPUT}\nassistant: {OUTPUT}\n"
|
||||||
|
with pytest.raises(RuntimeError):
|
||||||
|
construct_chat_template(
|
||||||
|
tokenizer = _FakeTokenizer(),
|
||||||
|
chat_template = template,
|
||||||
|
extra_eos_tokens = ["</s>"],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_error_message_excerpt_is_bounded():
|
||||||
|
"""Error messages must include a bounded excerpt of the offending
|
||||||
|
template, not dump arbitrarily large content into the traceback."""
|
||||||
|
huge = ("garbage " * 5000) + "{INPUT}" # ~40 KB, missing {OUTPUT}
|
||||||
|
with pytest.raises(RuntimeError) as exc_info:
|
||||||
|
construct_chat_template(
|
||||||
|
tokenizer = _FakeTokenizer(),
|
||||||
|
chat_template = huge,
|
||||||
|
extra_eos_tokens = ["</s>"],
|
||||||
|
)
|
||||||
|
msg = str(exc_info.value)
|
||||||
|
# Excerpt is repr-quoted and capped; total message should stay well
|
||||||
|
# under the template length.
|
||||||
|
assert len(msg) < 1000
|
||||||
|
assert "{OUTPUT}" in msg
|
||||||
|
|
@ -2461,17 +2461,40 @@ extra_eos_tokens = None,
|
||||||
f"{left_changed}"
|
f"{left_changed}"
|
||||||
)
|
)
|
||||||
except:
|
except:
|
||||||
ending = chat_template[chat_template.find("{OUTPUT}") + len("{OUTPUT}"):]
|
output_pos = chat_template.find("{OUTPUT}")
|
||||||
|
input_pos = chat_template.find("{INPUT}")
|
||||||
|
if output_pos == -1 or input_pos == -1:
|
||||||
|
missing = []
|
||||||
|
if input_pos == -1: missing.append("{INPUT}")
|
||||||
|
if output_pos == -1: missing.append("{OUTPUT}")
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Unsloth: chat_template must contain {' and '.join(missing)} "
|
||||||
|
f"placeholder(s). Got: {chat_template[:200]!r}"
|
||||||
|
)
|
||||||
|
ending = chat_template[output_pos + len("{OUTPUT}"):]
|
||||||
|
|
||||||
ending = re.escape(ending)
|
ending = re.escape(ending)
|
||||||
find_text = "{INPUT}" + ending + "(.+?{OUTPUT}" + ending + ")"
|
find_text = "{INPUT}" + ending + "(.+?{OUTPUT}" + ending + ")"
|
||||||
response_part = re.findall(find_text, chat_template, flags = re.DOTALL | re.MULTILINE)
|
response_part = re.findall(find_text, chat_template, flags = re.DOTALL | re.MULTILINE)
|
||||||
|
if len(response_part) == 0:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Unsloth: Could not recover a two-example structure from chat_template. "
|
||||||
|
"Provide exactly two {INPUT}/{OUTPUT} pairs (and optionally {SYSTEM}). "
|
||||||
|
f"Got: {chat_template[:200]!r}"
|
||||||
|
)
|
||||||
response_part = response_part[0]
|
response_part = response_part[0]
|
||||||
|
|
||||||
|
found = None
|
||||||
for j in range(1, len(response_part)):
|
for j in range(1, len(response_part)):
|
||||||
try_find = re.escape(response_part[:j])
|
try_find = re.escape(response_part[:j])
|
||||||
try: found = next(re.finditer("(" + try_find + ").+?\\{INPUT\\}", chat_template, flags = re.DOTALL | re.MULTILINE))
|
try: found = next(re.finditer("(" + try_find + ").+?\\{INPUT\\}", chat_template, flags = re.DOTALL | re.MULTILINE))
|
||||||
except: break
|
except: break
|
||||||
|
if found is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Unsloth: Could not locate a separator between examples in chat_template. "
|
||||||
|
"Provide exactly two {INPUT}/{OUTPUT} pairs (and optionally {SYSTEM}). "
|
||||||
|
f"Got: {chat_template[:200]!r}"
|
||||||
|
)
|
||||||
separator = found.group(1)
|
separator = found.group(1)
|
||||||
|
|
||||||
response_start = chat_template.find(response_part)
|
response_start = chat_template.find(response_part)
|
||||||
|
|
@ -2607,8 +2630,20 @@ extra_eos_tokens = None,
|
||||||
jinja_template = "{{ bos_token }}" + jinja_template
|
jinja_template = "{{ bos_token }}" + jinja_template
|
||||||
|
|
||||||
# Get instruction and output parts for train_on_inputs = False
|
# Get instruction and output parts for train_on_inputs = False
|
||||||
input_part = input_part [:input_part .find("{INPUT}")]
|
input_idx = input_part .find("{INPUT}")
|
||||||
output_part = output_part[:output_part.find("{OUTPUT}")]
|
output_idx = output_part.find("{OUTPUT}")
|
||||||
|
if input_idx == -1:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Unsloth: The instruction section of the template must contain the "
|
||||||
|
f"'{{INPUT}}' placeholder. Section: {input_part[:200]!r}"
|
||||||
|
)
|
||||||
|
if output_idx == -1:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Unsloth: The response section of the template must contain the "
|
||||||
|
f"'{{OUTPUT}}' placeholder. Section: {output_part[:200]!r}"
|
||||||
|
)
|
||||||
|
input_part = input_part [:input_idx ]
|
||||||
|
output_part = output_part[:output_idx]
|
||||||
return modelfile, jinja_template, input_part, output_part
|
return modelfile, jinja_template, input_part, output_part
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue