diff --git a/pipelines_ootd/transformer_garm_2d.py b/pipelines_ootd/transformer_garm_2d.py index 6f4d987..2834cc6 100644 --- a/pipelines_ootd/transformer_garm_2d.py +++ b/pipelines_ootd/transformer_garm_2d.py @@ -26,7 +26,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config from diffusers.models.embeddings import ImagePositionalEmbeddings from diffusers.utils import USE_PEFT_BACKEND, BaseOutput, deprecate # from diffusers.models.attention import BasicTransformerBlock -from diffusers.models.embeddings import CaptionProjection, PatchEmbed +from diffusers.models.embeddings import PixArtAlphaTextProjection, PatchEmbed from diffusers.models.lora import LoRACompatibleConv, LoRACompatibleLinear from diffusers.models.modeling_utils import ModelMixin from diffusers.models.normalization import AdaLayerNormSingle @@ -237,7 +237,7 @@ class Transformer2DModel(ModelMixin, ConfigMixin): self.caption_projection = None if caption_channels is not None: - self.caption_projection = CaptionProjection(in_features=caption_channels, hidden_size=inner_dim) + self.caption_projection = PixArtAlphaTextProjection(in_features=caption_channels, hidden_size=inner_dim) self.gradient_checkpointing = False diff --git a/pipelines_ootd/transformer_vton_2d.py b/pipelines_ootd/transformer_vton_2d.py index 276ee5b..98e15fe 100644 --- a/pipelines_ootd/transformer_vton_2d.py +++ b/pipelines_ootd/transformer_vton_2d.py @@ -26,7 +26,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config from diffusers.models.embeddings import ImagePositionalEmbeddings from diffusers.utils import USE_PEFT_BACKEND, BaseOutput, deprecate # from diffusers.models.attention import BasicTransformerBlock -from diffusers.models.embeddings import CaptionProjection, PatchEmbed +from diffusers.models.embeddings import PixArtAlphaTextProjection, PatchEmbed from diffusers.models.lora import LoRACompatibleConv, LoRACompatibleLinear from diffusers.models.modeling_utils import ModelMixin from diffusers.models.normalization import AdaLayerNormSingle @@ -237,7 +237,7 @@ class Transformer2DModel(ModelMixin, ConfigMixin): self.caption_projection = None if caption_channels is not None: - self.caption_projection = CaptionProjection(in_features=caption_channels, hidden_size=inner_dim) + self.caption_projection = PixArtAlphaTextProjection(in_features=caption_channels, hidden_size=inner_dim) self.gradient_checkpointing = False diff --git a/pipelines_ootd/unet_garm_2d_condition.py b/pipelines_ootd/unet_garm_2d_condition.py index 0b28a3a..5c4188a 100644 --- a/pipelines_ootd/unet_garm_2d_condition.py +++ b/pipelines_ootd/unet_garm_2d_condition.py @@ -44,7 +44,7 @@ from diffusers.models.embeddings import ( ImageHintTimeEmbedding, ImageProjection, ImageTimeEmbedding, - PositionNet, + GLIGENTextBoundingboxProjection, TextImageProjection, TextImageTimeEmbedding, TextTimeEmbedding, @@ -624,7 +624,7 @@ class UNetGarm2DConditionModel(ModelMixin, ConfigMixin, UNet2DConditionLoadersMi positive_len = cross_attention_dim[0] feature_type = "text-only" if attention_type == "gated" else "text-image" - self.position_net = PositionNet( + self.position_net = GLIGENTextBoundingboxProjection( positive_len=positive_len, out_dim=cross_attention_dim, feature_type=feature_type ) diff --git a/pipelines_ootd/unet_vton_2d_condition.py b/pipelines_ootd/unet_vton_2d_condition.py index 1e3a3fe..7b49fc5 100644 --- a/pipelines_ootd/unet_vton_2d_condition.py +++ b/pipelines_ootd/unet_vton_2d_condition.py @@ -44,7 +44,7 @@ from diffusers.models.embeddings import ( ImageHintTimeEmbedding, ImageProjection, ImageTimeEmbedding, - PositionNet, + GLIGENTextBoundingboxProjection, TextImageProjection, TextImageTimeEmbedding, TextTimeEmbedding, @@ -624,7 +624,7 @@ class UNetVton2DConditionModel(ModelMixin, ConfigMixin, UNet2DConditionLoadersMi positive_len = cross_attention_dim[0] feature_type = "text-only" if attention_type == "gated" else "text-image" - self.position_net = PositionNet( + self.position_net = GLIGENTextBoundingboxProjection( positive_len=positive_len, out_dim=cross_attention_dim, feature_type=feature_type ) diff --git a/requirements.txt b/requirements.txt index 1349b8c..45de0bc 100644 --- a/requirements.txt +++ b/requirements.txt @@ -5,7 +5,7 @@ scipy scikit-image opencv-python pillow -diffusers==0.24.0 +diffusers transformers accelerate matplotlib