diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py index 85c10eec4..3346eb72b 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/text_encoder_loader.py @@ -9,7 +9,6 @@ import torch import torch.distributed as dist from torch import nn from torch.distributed import init_device_mesh -from transformers import AutoModel from transformers.utils import SAFE_WEIGHTS_INDEX_NAME from sglang.multimodal_gen.configs.models import EncoderConfig, ModelConfig @@ -109,13 +108,54 @@ class TextEncoderLoader(ComponentLoader): 1 if component_model_path.rstrip("/").endswith("text_encoder_2") else 0 ) encoder_dtype = server_args.pipeline_config.text_encoder_precisions[encoder_idx] - return AutoModel.from_pretrained( + transformers_model_class = self._resolve_transformers_text_encoder_class( + component_model_path, server_args + ) + return transformers_model_class.from_pretrained( component_model_path, trust_remote_code=server_args.trust_remote_code, revision=server_args.revision, torch_dtype=PRECISION_TO_TYPE[encoder_dtype], ) + @staticmethod + def _resolve_transformers_text_encoder_class(component_model_path, server_args): + """Resolve the concrete transformers class for a text encoder. + + AutoModel maps encoder-decoder model types (e.g. T5/UMT5) to full + seq2seq classes, whose forward expects decoder inputs and raises when + the module is used purely as a text encoder. For such checkpoints, + prefer the encoder-only class from the config architectures or map the + full seq2seq architecture to its encoder-only counterpart. Encoders that + are not encoder-decoder keep using AutoModel unchanged. + """ + import transformers + from transformers import AutoConfig, AutoModel + + try: + config = AutoConfig.from_pretrained( + component_model_path, + trust_remote_code=server_args.trust_remote_code, + revision=server_args.revision, + ) + except Exception: + return AutoModel + if getattr(config, "is_encoder_decoder", False): + encoder_only_map = { + "T5Model": "T5EncoderModel", + "T5ForConditionalGeneration": "T5EncoderModel", + "UMT5Model": "UMT5EncoderModel", + "UMT5ForConditionalGeneration": "UMT5EncoderModel", + "MT5Model": "MT5EncoderModel", + "MT5ForConditionalGeneration": "MT5EncoderModel", + } + for arch in getattr(config, "architectures", None) or []: + encoder_arch = encoder_only_map.get(arch, arch) + transformers_model_class = getattr(transformers, encoder_arch, None) + if isinstance(transformers_model_class, type): + return transformers_model_class + return AutoModel + def _prepare_weights( self, model_name_or_path: str, diff --git a/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py b/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py new file mode 100644 index 000000000..4c04bffa3 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_text_encoder_loader.py @@ -0,0 +1,83 @@ +import unittest +from types import SimpleNamespace +from unittest import mock + +import transformers + +from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import ( + TextEncoderLoader, +) + + +class TestTextEncoderClassResolution(unittest.TestCase): + """load_native must not load encoder-decoder text encoders via AutoModel. + + AutoModel maps T5/UMT5 model types to the full seq2seq class + (T5Model/UMT5Model), whose forward needs decoder inputs and raises when the + module is used purely as a text encoder. + """ + + server_args = SimpleNamespace(trust_remote_code=False, revision=None) + + def _resolve(self, is_encoder_decoder, architectures): + config = SimpleNamespace( + is_encoder_decoder=is_encoder_decoder, architectures=architectures + ) + with mock.patch.object( + transformers.AutoConfig, "from_pretrained", return_value=config + ): + return TextEncoderLoader._resolve_transformers_text_encoder_class( + "dummy/path", self.server_args + ) + + def test_umt5_encoder_decoder_uses_encoder_only_class(self): + self.assertIs( + self._resolve(True, ["UMT5EncoderModel"]), transformers.UMT5EncoderModel + ) + self.assertIs(self._resolve(True, ["UMT5Model"]), transformers.UMT5EncoderModel) + self.assertIs( + self._resolve(True, ["UMT5ForConditionalGeneration"]), + transformers.UMT5EncoderModel, + ) + + def test_t5_encoder_decoder_uses_encoder_only_class(self): + self.assertIs( + self._resolve(True, ["T5EncoderModel"]), transformers.T5EncoderModel + ) + self.assertIs(self._resolve(True, ["T5Model"]), transformers.T5EncoderModel) + self.assertIs( + self._resolve(True, ["T5ForConditionalGeneration"]), + transformers.T5EncoderModel, + ) + + def test_mt5_encoder_decoder_uses_encoder_only_class(self): + self.assertIs( + self._resolve(True, ["MT5EncoderModel"]), transformers.MT5EncoderModel + ) + self.assertIs(self._resolve(True, ["MT5Model"]), transformers.MT5EncoderModel) + self.assertIs( + self._resolve(True, ["MT5ForConditionalGeneration"]), + transformers.MT5EncoderModel, + ) + + def test_non_encoder_decoder_keeps_automodel(self): + # e.g. CLIP/Mistral/Qwen text encoders are not encoder-decoder. + self.assertIs(self._resolve(False, ["CLIPTextModel"]), transformers.AutoModel) + + def test_unknown_architecture_falls_back_to_automodel(self): + self.assertIs(self._resolve(True, ["NotARealClass"]), transformers.AutoModel) + + def test_config_load_failure_falls_back_to_automodel(self): + with mock.patch.object( + transformers.AutoConfig, + "from_pretrained", + side_effect=OSError("no config"), + ): + cls = TextEncoderLoader._resolve_transformers_text_encoder_class( + "dummy/path", self.server_args + ) + self.assertIs(cls, transformers.AutoModel) + + +if __name__ == "__main__": + unittest.main()