fix: preserve tokenizer eos token on merged saves (#5451)
Preserves the source tokenizer's eos_token in tokenizer_config.json after merged saves so runtimes such as vLLM read the correct stop token. Centralized inside the patched tokenizer save_pretrained so all save paths (merged_16bit, GGUF, torchao, push_to_hub) benefit, with filename_prefix support. Fixes #5386
This commit is contained in:
parent
0542dc0725
commit
4e9d772d36
2 changed files with 170 additions and 0 deletions
106
tests/saving/test_preserve_tokenizer_eos_token.py
Normal file
106
tests/saving/test_preserve_tokenizer_eos_token.py
Normal file
|
|
@ -0,0 +1,106 @@
|
|||
import ast
|
||||
import json
|
||||
import os
|
||||
import types
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class _Logger:
|
||||
def warning_once(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
|
||||
def _load_preserve_helper():
|
||||
source = Path(__file__).parents[2] / "unsloth" / "save.py"
|
||||
tree = ast.parse(source.read_text(encoding = "utf-8"))
|
||||
helper = next(
|
||||
node
|
||||
for node in tree.body
|
||||
if isinstance(node, ast.FunctionDef)
|
||||
and node.name == "_preserve_tokenizer_eos_token"
|
||||
)
|
||||
module = ast.Module(body = [helper], type_ignores = [])
|
||||
ast.fix_missing_locations(module)
|
||||
namespace = {"json": json, "os": os, "logger": _Logger()}
|
||||
exec(compile(module, str(source), "exec"), namespace)
|
||||
return namespace["_preserve_tokenizer_eos_token"]
|
||||
|
||||
|
||||
def test_preserve_tokenizer_eos_token_restores_gemma4_turn_token(tmp_path):
|
||||
preserve = _load_preserve_helper()
|
||||
tokenizer_config = tmp_path / "tokenizer_config.json"
|
||||
tokenizer_config.write_text(
|
||||
json.dumps({"eos_token": "<eos>", "other": True}),
|
||||
encoding = "utf-8",
|
||||
)
|
||||
tokenizer = types.SimpleNamespace(eos_token = "<turn|>")
|
||||
|
||||
preserve(tokenizer, tmp_path)
|
||||
|
||||
saved_config = json.loads(tokenizer_config.read_text(encoding = "utf-8"))
|
||||
assert saved_config["eos_token"] == "<turn|>"
|
||||
assert saved_config["other"] is True
|
||||
|
||||
|
||||
def test_preserve_tokenizer_eos_token_supports_processor_tokenizer(tmp_path):
|
||||
preserve = _load_preserve_helper()
|
||||
tokenizer_config = tmp_path / "tokenizer_config.json"
|
||||
tokenizer_config.write_text(json.dumps({"eos_token": "<eos>"}), encoding = "utf-8")
|
||||
processor = types.SimpleNamespace(
|
||||
tokenizer = types.SimpleNamespace(eos_token = "<turn|>")
|
||||
)
|
||||
|
||||
preserve(processor, tmp_path)
|
||||
|
||||
saved_config = json.loads(tokenizer_config.read_text(encoding = "utf-8"))
|
||||
assert saved_config["eos_token"] == "<turn|>"
|
||||
|
||||
|
||||
class _StringableToken:
|
||||
def __str__(self):
|
||||
return "<turn|>"
|
||||
|
||||
|
||||
def test_preserve_tokenizer_eos_token_serializes_stringable_tokens(tmp_path):
|
||||
preserve = _load_preserve_helper()
|
||||
tokenizer_config = tmp_path / "tokenizer_config.json"
|
||||
tokenizer_config.write_text(json.dumps({"eos_token": "<eos>"}), encoding = "utf-8")
|
||||
tokenizer = types.SimpleNamespace(eos_token = _StringableToken())
|
||||
|
||||
preserve(tokenizer, tmp_path)
|
||||
|
||||
saved_config = json.loads(tokenizer_config.read_text(encoding = "utf-8"))
|
||||
assert saved_config["eos_token"] == "<turn|>"
|
||||
|
||||
|
||||
def test_preserve_tokenizer_eos_token_supports_filename_prefix(tmp_path):
|
||||
preserve = _load_preserve_helper()
|
||||
prefixed_config = tmp_path / "adapter-tokenizer_config.json"
|
||||
prefixed_config.write_text(
|
||||
json.dumps({"eos_token": "<eos>", "other": True}),
|
||||
encoding = "utf-8",
|
||||
)
|
||||
tokenizer = types.SimpleNamespace(eos_token = "<turn|>")
|
||||
|
||||
preserve(tokenizer, tmp_path, filename_prefix = "adapter")
|
||||
|
||||
saved_config = json.loads(prefixed_config.read_text(encoding = "utf-8"))
|
||||
assert saved_config["eos_token"] == "<turn|>"
|
||||
assert saved_config["other"] is True
|
||||
# Unprefixed file must not be created as a side effect.
|
||||
assert not (tmp_path / "tokenizer_config.json").exists()
|
||||
|
||||
|
||||
def test_preserve_tokenizer_eos_token_filename_prefix_none_uses_default(tmp_path):
|
||||
preserve = _load_preserve_helper()
|
||||
default_config = tmp_path / "tokenizer_config.json"
|
||||
default_config.write_text(
|
||||
json.dumps({"eos_token": "<eos>"}),
|
||||
encoding = "utf-8",
|
||||
)
|
||||
tokenizer = types.SimpleNamespace(eos_token = "<turn|>")
|
||||
|
||||
preserve(tokenizer, tmp_path, filename_prefix = None)
|
||||
|
||||
saved_config = json.loads(default_config.read_text(encoding = "utf-8"))
|
||||
assert saved_config["eos_token"] == "<turn|>"
|
||||
|
|
@ -424,6 +424,60 @@ def fast_save_pickle(shard, name):
|
|||
return
|
||||
|
||||
|
||||
def _preserve_tokenizer_eos_token(tokenizer, save_directory, filename_prefix = None):
|
||||
"""Restore tokenizer_config.json eos_token from the tokenizer passed to save.
|
||||
|
||||
Some merge paths may re-save or mutate tokenizer metadata after the tokenizer
|
||||
is written. Gemma 4 instruct models use `<turn|>` as their chat EOS token;
|
||||
if tokenizer_config.json is reset to the raw base `<eos>` token, runtimes such
|
||||
as vLLM will not stop generation correctly. Keep the serialized metadata in
|
||||
sync with the source tokenizer without failing the save if the config is not
|
||||
present or cannot be edited.
|
||||
|
||||
`filename_prefix` mirrors the same argument on Transformers'
|
||||
`PreTrainedTokenizerBase.save_pretrained`: when provided, the tokenizer
|
||||
config is written as `{filename_prefix}-tokenizer_config.json` instead of
|
||||
`tokenizer_config.json`.
|
||||
"""
|
||||
if tokenizer is None or save_directory is None:
|
||||
return
|
||||
|
||||
source_tokenizer = (
|
||||
tokenizer.tokenizer if hasattr(tokenizer, "tokenizer") else tokenizer
|
||||
)
|
||||
eos_token = getattr(source_tokenizer, "eos_token", None)
|
||||
if eos_token is None and source_tokenizer is not tokenizer:
|
||||
eos_token = getattr(tokenizer, "eos_token", None)
|
||||
if eos_token is None:
|
||||
return
|
||||
eos_token = str(eos_token)
|
||||
|
||||
tokenizer_config_name = (
|
||||
f"{filename_prefix}-tokenizer_config.json"
|
||||
if filename_prefix
|
||||
else "tokenizer_config.json"
|
||||
)
|
||||
tokenizer_config = os.path.join(str(save_directory), tokenizer_config_name)
|
||||
if not os.path.isfile(tokenizer_config):
|
||||
return
|
||||
|
||||
try:
|
||||
with open(tokenizer_config, "r", encoding = "utf-8") as file:
|
||||
config = json.load(file)
|
||||
|
||||
if config.get("eos_token") == eos_token:
|
||||
return
|
||||
|
||||
config["eos_token"] = eos_token
|
||||
with open(tokenizer_config, "w", encoding = "utf-8") as file:
|
||||
json.dump(config, file, indent = 2, ensure_ascii = False)
|
||||
file.write("\n")
|
||||
except Exception as error:
|
||||
logger.warning_once(
|
||||
f"Unsloth: Could not preserve tokenizer eos_token in {tokenizer_config}: {error}"
|
||||
)
|
||||
|
||||
|
||||
@torch.inference_mode
|
||||
def unsloth_save_model(
|
||||
model,
|
||||
|
|
@ -980,6 +1034,11 @@ def unsloth_save_model(
|
|||
_tokenizer.padding_side = "left"
|
||||
|
||||
tokenizer.save_pretrained(**tokenizer_save_settings)
|
||||
_preserve_tokenizer_eos_token(
|
||||
tokenizer,
|
||||
tokenizer_save_settings["save_directory"],
|
||||
filename_prefix = tokenizer_save_settings.get("filename_prefix"),
|
||||
)
|
||||
|
||||
# Revert back padding side
|
||||
_tokenizer.padding_side = old_padding_side
|
||||
|
|
@ -3516,6 +3575,11 @@ def patch_saving_functions(model, vision = False):
|
|||
save_directory,
|
||||
token = kwargs.get("token", None),
|
||||
)
|
||||
_preserve_tokenizer_eos_token(
|
||||
self,
|
||||
save_directory,
|
||||
filename_prefix = filename_prefix,
|
||||
)
|
||||
if push_to_hub:
|
||||
push_kwargs = dict(kwargs)
|
||||
repo_id = push_kwargs.pop("repo_id", save_directory)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue