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:
Anmol Mishra 2026-05-17 19:14:23 +05:30 committed by GitHub
commit 4e9d772d36
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 170 additions and 0 deletions

View 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|>"

View file

@ -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)