Merge remote-tracking branch 'origin/main' into studio-diffusion-images

This commit is contained in:
Daniel Han-Chen 2026-05-25 14:43:31 +00:00
commit 1f5f13c986
2 changed files with 115 additions and 3 deletions

View 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

View file

@ -2461,17 +2461,40 @@ extra_eos_tokens = None,
f"{left_changed}"
)
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)
find_text = "{INPUT}" + ending + "(.+?{OUTPUT}" + ending + ")"
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]
found = None
for j in range(1, len(response_part)):
try_find = re.escape(response_part[:j])
try: found = next(re.finditer("(" + try_find + ").+?\\{INPUT\\}", chat_template, flags = re.DOTALL | re.MULTILINE))
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)
response_start = chat_template.find(response_part)
@ -2607,8 +2630,20 @@ extra_eos_tokens = None,
jinja_template = "{{ bos_token }}" + jinja_template
# Get instruction and output parts for train_on_inputs = False
input_part = input_part [:input_part .find("{INPUT}")]
output_part = output_part[:output_part.find("{OUTPUT}")]
input_idx = input_part .find("{INPUT}")
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