diff --git a/tests/saving/test_preserve_tokenizer_eos_token.py b/tests/saving/test_preserve_tokenizer_eos_token.py new file mode 100644 index 0000000000..6e2f8c7f9d --- /dev/null +++ b/tests/saving/test_preserve_tokenizer_eos_token.py @@ -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": "", "other": True}), + encoding = "utf-8", + ) + tokenizer = types.SimpleNamespace(eos_token = "") + + preserve(tokenizer, tmp_path) + + saved_config = json.loads(tokenizer_config.read_text(encoding = "utf-8")) + assert saved_config["eos_token"] == "" + 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": ""}), encoding = "utf-8") + processor = types.SimpleNamespace( + tokenizer = types.SimpleNamespace(eos_token = "") + ) + + preserve(processor, tmp_path) + + saved_config = json.loads(tokenizer_config.read_text(encoding = "utf-8")) + assert saved_config["eos_token"] == "" + + +class _StringableToken: + def __str__(self): + return "" + + +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": ""}), 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"] == "" + + +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": "", "other": True}), + encoding = "utf-8", + ) + tokenizer = types.SimpleNamespace(eos_token = "") + + preserve(tokenizer, tmp_path, filename_prefix = "adapter") + + saved_config = json.loads(prefixed_config.read_text(encoding = "utf-8")) + assert saved_config["eos_token"] == "" + 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": ""}), + encoding = "utf-8", + ) + tokenizer = types.SimpleNamespace(eos_token = "") + + preserve(tokenizer, tmp_path, filename_prefix = None) + + saved_config = json.loads(default_config.read_text(encoding = "utf-8")) + assert saved_config["eos_token"] == "" diff --git a/unsloth/save.py b/unsloth/save.py index d67a1a0550..3628c468e3 100644 --- a/unsloth/save.py +++ b/unsloth/save.py @@ -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 `` as their chat EOS token; + if tokenizer_config.json is reset to the raw base `` 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)