[diffusion] Fix native text-encoder loading for T5/UMT5 encoder-decoder models (#27432)
Co-authored-by: xiaoyu.zhang <xiaoyu.zhang@radixark.net>
This commit is contained in:
co-authored by
xiaoyu.zhang
parent
5bf7dd8e4a
commit
6c2770149b
+42
-2
@@ -9,7 +9,6 @@ import torch
|
|||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
from torch import nn
|
from torch import nn
|
||||||
from torch.distributed import init_device_mesh
|
from torch.distributed import init_device_mesh
|
||||||
from transformers import AutoModel
|
|
||||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models import EncoderConfig, ModelConfig
|
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
|
1 if component_model_path.rstrip("/").endswith("text_encoder_2") else 0
|
||||||
)
|
)
|
||||||
encoder_dtype = server_args.pipeline_config.text_encoder_precisions[encoder_idx]
|
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,
|
component_model_path,
|
||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
revision=server_args.revision,
|
revision=server_args.revision,
|
||||||
torch_dtype=PRECISION_TO_TYPE[encoder_dtype],
|
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(
|
def _prepare_weights(
|
||||||
self,
|
self,
|
||||||
model_name_or_path: str,
|
model_name_or_path: str,
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user