Fix processor config loading for object-storage model paths (#31211)
Co-authored-by: Alex Nails <alex.nails@radixark.ai>
This commit is contained in:
co-authored by
Alex Nails
parent
980acd6eca
commit
dec0836302
@@ -153,6 +153,8 @@ def get_processor(
|
|||||||
|
|
||||||
revision = kwargs.pop("revision", tokenizer_revision)
|
revision = kwargs.pop("revision", tokenizer_revision)
|
||||||
tokenizer_name = resolve_runai_obj_uri(tokenizer_name)
|
tokenizer_name = resolve_runai_obj_uri(tokenizer_name)
|
||||||
|
if model_name is not None:
|
||||||
|
model_name = resolve_runai_obj_uri(model_name)
|
||||||
|
|
||||||
if is_mistral_model(tokenizer_name):
|
if is_mistral_model(tokenizer_name):
|
||||||
config = load_mistral_config(
|
config = load_mistral_config(
|
||||||
|
|||||||
@@ -8,11 +8,12 @@ import inspect
|
|||||||
import tempfile
|
import tempfile
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
from transformers import PretrainedConfig
|
from transformers import PretrainedConfig
|
||||||
from transformers.image_processing_utils import BaseImageProcessor
|
from transformers.image_processing_utils import BaseImageProcessor
|
||||||
|
|
||||||
|
import sglang.srt.utils.hf_transformers.processor as processor_utils
|
||||||
from sglang.srt.utils import hf_transformers_patches
|
from sglang.srt.utils import hf_transformers_patches
|
||||||
from sglang.srt.utils.hf_transformers.common import (
|
from sglang.srt.utils.hf_transformers.common import (
|
||||||
_is_deepseek_ocr2_model,
|
_is_deepseek_ocr2_model,
|
||||||
@@ -31,6 +32,45 @@ from sglang.test.ci.ci_register import register_cpu_ci
|
|||||||
register_cpu_ci(est_time=6, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=6, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# get_processor
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class TestGetProcessor(unittest.TestCase):
|
||||||
|
def test_resolves_model_name_before_loading_config(self):
|
||||||
|
remote_model = "s3://bucket/model"
|
||||||
|
local_model = "/cache/model"
|
||||||
|
config = SimpleNamespace(model_type="clip", auto_map={})
|
||||||
|
loaded_processor = MagicMock()
|
||||||
|
loaded_processor.tokenizer.chat_template = "template"
|
||||||
|
auto_config = MagicMock()
|
||||||
|
auto_config.from_pretrained.return_value = config
|
||||||
|
auto_processor = MagicMock()
|
||||||
|
auto_processor.from_pretrained.return_value = loaded_processor
|
||||||
|
|
||||||
|
def resolve_uri(path):
|
||||||
|
return local_model if path == remote_model else path
|
||||||
|
|
||||||
|
with patch.multiple(
|
||||||
|
processor_utils,
|
||||||
|
resolve_runai_obj_uri=MagicMock(side_effect=resolve_uri),
|
||||||
|
AutoConfig=auto_config,
|
||||||
|
AutoProcessor=auto_processor,
|
||||||
|
):
|
||||||
|
processor = processor_utils.get_processor(
|
||||||
|
"local-tokenizer",
|
||||||
|
model_name=remote_model,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIs(processor, loaded_processor)
|
||||||
|
auto_config.from_pretrained.assert_called_once_with(
|
||||||
|
local_model,
|
||||||
|
trust_remote_code=False,
|
||||||
|
revision=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# _patch_image_processor_kwargs
|
# _patch_image_processor_kwargs
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user