From 09218c6699d4071d943ceb2eaacf7a4ab31e6f3d Mon Sep 17 00:00:00 2001 From: Takuma Mori Date: Fri, 20 Jan 2023 21:43:34 +0900 Subject: [PATCH 01/10] allow passing op to xFormers attention original code by @patil-suraj huggingface/diffusers@ae0cc0b71f28c0f2c5c27026b18f1bea98b505f1 --- src/diffusers/models/attention.py | 8 +++++--- src/diffusers/models/cross_attention.py | 11 +++++++---- src/diffusers/models/modeling_utils.py | 8 ++++---- src/diffusers/pipelines/pipeline_utils.py | 10 +++++----- 4 files changed, 21 insertions(+), 16 deletions(-) diff --git a/src/diffusers/models/attention.py b/src/diffusers/models/attention.py index 08263875d0c2..4cf9ce72f73b 100644 --- a/src/diffusers/models/attention.py +++ b/src/diffusers/models/attention.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. import math -from typing import Optional +from typing import Callable, Optional import torch import torch.nn.functional as F @@ -72,6 +72,7 @@ def __init__( self.proj_attn = nn.Linear(channels, channels, 1) self._use_memory_efficient_attention_xformers = False + self._attention_op = None def reshape_heads_to_batch_dim(self, tensor): batch_size, seq_len, dim = tensor.shape @@ -87,7 +88,7 @@ def reshape_batch_dim_to_heads(self, tensor): tensor = tensor.permute(0, 2, 1, 3).reshape(batch_size // head_size, seq_len, dim * head_size) return tensor - def set_use_memory_efficient_attention_xformers(self, use_memory_efficient_attention_xformers: bool): + def set_use_memory_efficient_attention_xformers(self, use_memory_efficient_attention_xformers: bool, attention_op: Optional[Callable] = None): if use_memory_efficient_attention_xformers: if not is_xformers_available(): raise ModuleNotFoundError( @@ -113,6 +114,7 @@ def set_use_memory_efficient_attention_xformers(self, use_memory_efficient_atten except Exception as e: raise e self._use_memory_efficient_attention_xformers = use_memory_efficient_attention_xformers + self._attention_op = attention_op def forward(self, hidden_states): residual = hidden_states @@ -136,7 +138,7 @@ def forward(self, hidden_states): if self._use_memory_efficient_attention_xformers: # Memory efficient attention - hidden_states = xformers.ops.memory_efficient_attention(query_proj, key_proj, value_proj, attn_bias=None) + hidden_states = xformers.ops.memory_efficient_attention(query_proj, key_proj, value_proj, attn_bias=None, op=self._attention_op) hidden_states = hidden_states.to(query_proj.dtype) else: attention_scores = torch.baddbmm( diff --git a/src/diffusers/models/cross_attention.py b/src/diffusers/models/cross_attention.py index d4da50c23f66..b80daa44d32d 100644 --- a/src/diffusers/models/cross_attention.py +++ b/src/diffusers/models/cross_attention.py @@ -11,7 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -from typing import Optional, Union +from typing import Optional, Union, Callable import torch import torch.nn.functional as F @@ -93,7 +93,7 @@ def __init__( processor = processor if processor is not None else CrossAttnProcessor() self.set_processor(processor) - def set_use_memory_efficient_attention_xformers(self, use_memory_efficient_attention_xformers: bool): + def set_use_memory_efficient_attention_xformers(self, use_memory_efficient_attention_xformers: bool, attention_op: Optional[Callable] = None): if use_memory_efficient_attention_xformers: if self.added_kv_proj_dim is not None: # TODO(Anton, Patrick, Suraj, William) - currently xformers doesn't work for UnCLIP @@ -127,7 +127,7 @@ def set_use_memory_efficient_attention_xformers(self, use_memory_efficient_atten except Exception as e: raise e - processor = XFormersCrossAttnProcessor() + processor = XFormersCrossAttnProcessor(attention_op=attention_op) else: processor = CrossAttnProcessor() @@ -351,6 +351,9 @@ def __call__(self, attn: CrossAttention, hidden_states, encoder_hidden_states=No class XFormersCrossAttnProcessor: + def __init__(self, attention_op: Optional[Callable] = None): + self.attention_op = attention_op + def __call__(self, attn: CrossAttention, hidden_states, encoder_hidden_states=None, attention_mask=None): batch_size, sequence_length, _ = hidden_states.shape @@ -366,7 +369,7 @@ def __call__(self, attn: CrossAttention, hidden_states, encoder_hidden_states=No key = attn.head_to_batch_dim(key).contiguous() value = attn.head_to_batch_dim(value).contiguous() - hidden_states = xformers.ops.memory_efficient_attention(query, key, value, attn_bias=attention_mask) + hidden_states = xformers.ops.memory_efficient_attention(query, key, value, attn_bias=attention_mask, op=self.attention_op) hidden_states = hidden_states.to(query.dtype) hidden_states = attn.batch_to_head_dim(hidden_states) diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index afe5689fdb24..a44427dd4e05 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -190,13 +190,13 @@ def disable_gradient_checkpointing(self): if self._supports_gradient_checkpointing: self.apply(partial(self._set_gradient_checkpointing, value=False)) - def set_use_memory_efficient_attention_xformers(self, valid: bool) -> None: + def set_use_memory_efficient_attention_xformers(self, valid: bool, attention_op: Optional[Callable] = None) -> None: # Recursively walk through all the children. # Any children which exposes the set_use_memory_efficient_attention_xformers method # gets the message def fn_recursive_set_mem_eff(module: torch.nn.Module): if hasattr(module, "set_use_memory_efficient_attention_xformers"): - module.set_use_memory_efficient_attention_xformers(valid) + module.set_use_memory_efficient_attention_xformers(valid, attention_op) for child in module.children(): fn_recursive_set_mem_eff(child) @@ -205,7 +205,7 @@ def fn_recursive_set_mem_eff(module: torch.nn.Module): if isinstance(module, torch.nn.Module): fn_recursive_set_mem_eff(module) - def enable_xformers_memory_efficient_attention(self): + def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Callable] = None): r""" Enable memory efficient attention as implemented in xformers. @@ -215,7 +215,7 @@ def enable_xformers_memory_efficient_attention(self): Warning: When Memory Efficient Attention and Sliced attention are both enabled, the Memory Efficient Attention is used. """ - self.set_use_memory_efficient_attention_xformers(True) + self.set_use_memory_efficient_attention_xformers(True, attention_op) def disable_xformers_memory_efficient_attention(self): r""" diff --git a/src/diffusers/pipelines/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index ea28ac875f81..b9c93580929b 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -19,7 +19,7 @@ import os from dataclasses import dataclass from pathlib import Path -from typing import Any, Dict, List, Optional, Union +from typing import Any, Dict, List, Optional, Union, Callable import numpy as np import torch @@ -838,7 +838,7 @@ def progress_bar(self, iterable=None, total=None): def set_progress_bar_config(self, **kwargs): self._progress_bar_config = kwargs - def enable_xformers_memory_efficient_attention(self): + def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Callable] = None): r""" Enable memory efficient attention as implemented in xformers. @@ -848,7 +848,7 @@ def enable_xformers_memory_efficient_attention(self): Warning: When Memory Efficient Attention and Sliced attention are both enabled, the Memory Efficient Attention is used. """ - self.set_use_memory_efficient_attention_xformers(True) + self.set_use_memory_efficient_attention_xformers(True, attention_op) def disable_xformers_memory_efficient_attention(self): r""" @@ -856,13 +856,13 @@ def disable_xformers_memory_efficient_attention(self): """ self.set_use_memory_efficient_attention_xformers(False) - def set_use_memory_efficient_attention_xformers(self, valid: bool) -> None: + def set_use_memory_efficient_attention_xformers(self, valid: bool, attention_op: Optional[Callable] = None) -> None: # Recursively walk through all the children. # Any children which exposes the set_use_memory_efficient_attention_xformers method # gets the message def fn_recursive_set_mem_eff(module: torch.nn.Module): if hasattr(module, "set_use_memory_efficient_attention_xformers"): - module.set_use_memory_efficient_attention_xformers(valid) + module.set_use_memory_efficient_attention_xformers(valid, attention_op) for child in module.children(): fn_recursive_set_mem_eff(child) From 29555049b6b72594a1aea88fa2565770b8320308 Mon Sep 17 00:00:00 2001 From: Takuma Mori Date: Fri, 20 Jan 2023 21:46:24 +0900 Subject: [PATCH 02/10] correct style by `make style` --- src/diffusers/models/attention.py | 8 ++++++-- src/diffusers/models/cross_attention.py | 10 +++++++--- src/diffusers/models/modeling_utils.py | 4 +++- src/diffusers/pipelines/pipeline_utils.py | 6 ++++-- 4 files changed, 20 insertions(+), 8 deletions(-) diff --git a/src/diffusers/models/attention.py b/src/diffusers/models/attention.py index 4cf9ce72f73b..b5acd6f4f900 100644 --- a/src/diffusers/models/attention.py +++ b/src/diffusers/models/attention.py @@ -88,7 +88,9 @@ def reshape_batch_dim_to_heads(self, tensor): tensor = tensor.permute(0, 2, 1, 3).reshape(batch_size // head_size, seq_len, dim * head_size) return tensor - def set_use_memory_efficient_attention_xformers(self, use_memory_efficient_attention_xformers: bool, attention_op: Optional[Callable] = None): + def set_use_memory_efficient_attention_xformers( + self, use_memory_efficient_attention_xformers: bool, attention_op: Optional[Callable] = None + ): if use_memory_efficient_attention_xformers: if not is_xformers_available(): raise ModuleNotFoundError( @@ -138,7 +140,9 @@ def forward(self, hidden_states): if self._use_memory_efficient_attention_xformers: # Memory efficient attention - hidden_states = xformers.ops.memory_efficient_attention(query_proj, key_proj, value_proj, attn_bias=None, op=self._attention_op) + hidden_states = xformers.ops.memory_efficient_attention( + query_proj, key_proj, value_proj, attn_bias=None, op=self._attention_op + ) hidden_states = hidden_states.to(query_proj.dtype) else: attention_scores = torch.baddbmm( diff --git a/src/diffusers/models/cross_attention.py b/src/diffusers/models/cross_attention.py index b80daa44d32d..7dda30fbdadc 100644 --- a/src/diffusers/models/cross_attention.py +++ b/src/diffusers/models/cross_attention.py @@ -11,7 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -from typing import Optional, Union, Callable +from typing import Callable, Optional, Union import torch import torch.nn.functional as F @@ -93,7 +93,9 @@ def __init__( processor = processor if processor is not None else CrossAttnProcessor() self.set_processor(processor) - def set_use_memory_efficient_attention_xformers(self, use_memory_efficient_attention_xformers: bool, attention_op: Optional[Callable] = None): + def set_use_memory_efficient_attention_xformers( + self, use_memory_efficient_attention_xformers: bool, attention_op: Optional[Callable] = None + ): if use_memory_efficient_attention_xformers: if self.added_kv_proj_dim is not None: # TODO(Anton, Patrick, Suraj, William) - currently xformers doesn't work for UnCLIP @@ -369,7 +371,9 @@ def __call__(self, attn: CrossAttention, hidden_states, encoder_hidden_states=No key = attn.head_to_batch_dim(key).contiguous() value = attn.head_to_batch_dim(value).contiguous() - hidden_states = xformers.ops.memory_efficient_attention(query, key, value, attn_bias=attention_mask, op=self.attention_op) + hidden_states = xformers.ops.memory_efficient_attention( + query, key, value, attn_bias=attention_mask, op=self.attention_op + ) hidden_states = hidden_states.to(query.dtype) hidden_states = attn.batch_to_head_dim(hidden_states) diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index a44427dd4e05..9aeb764a79ff 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -190,7 +190,9 @@ def disable_gradient_checkpointing(self): if self._supports_gradient_checkpointing: self.apply(partial(self._set_gradient_checkpointing, value=False)) - def set_use_memory_efficient_attention_xformers(self, valid: bool, attention_op: Optional[Callable] = None) -> None: + def set_use_memory_efficient_attention_xformers( + self, valid: bool, attention_op: Optional[Callable] = None + ) -> None: # Recursively walk through all the children. # Any children which exposes the set_use_memory_efficient_attention_xformers method # gets the message diff --git a/src/diffusers/pipelines/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index b9c93580929b..0ca387512200 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -19,7 +19,7 @@ import os from dataclasses import dataclass from pathlib import Path -from typing import Any, Dict, List, Optional, Union, Callable +from typing import Any, Callable, Dict, List, Optional, Union import numpy as np import torch @@ -856,7 +856,9 @@ def disable_xformers_memory_efficient_attention(self): """ self.set_use_memory_efficient_attention_xformers(False) - def set_use_memory_efficient_attention_xformers(self, valid: bool, attention_op: Optional[Callable] = None) -> None: + def set_use_memory_efficient_attention_xformers( + self, valid: bool, attention_op: Optional[Callable] = None + ) -> None: # Recursively walk through all the children. # Any children which exposes the set_use_memory_efficient_attention_xformers method # gets the message From b57d983c60ad081f58fdcf484ed87f03f2af8f32 Mon Sep 17 00:00:00 2001 From: Takuma Mori Date: Fri, 20 Jan 2023 23:15:47 +0900 Subject: [PATCH 03/10] add attention_op arg documents --- src/diffusers/models/modeling_utils.py | 6 ++++++ src/diffusers/pipelines/pipeline_utils.py | 6 ++++++ 2 files changed, 12 insertions(+) diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index 9aeb764a79ff..a6ada2c4a19b 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -216,6 +216,12 @@ def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Call Warning: When Memory Efficient Attention and Sliced attention are both enabled, the Memory Efficient Attention is used. + + Parameters: + attention_op (`Callable`, *optional*): + Override the default `None` operator for use as `op` argument to the + [`memory_efficient_attention()`](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.memory_efficient_attention) + function of xFormers. """ self.set_use_memory_efficient_attention_xformers(True, attention_op) diff --git a/src/diffusers/pipelines/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index 0ca387512200..7cf32fc6c945 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -847,6 +847,12 @@ def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Call Warning: When Memory Efficient Attention and Sliced attention are both enabled, the Memory Efficient Attention is used. + + Parameters: + attention_op (`Callable`, *optional*): + Override the default `None` operator for use as `op` argument to the + [`memory_efficient_attention()`](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.memory_efficient_attention) + function of xFormers. """ self.set_use_memory_efficient_attention_xformers(True, attention_op) From 425f3f23ba8f65182dd1c0208b874d967a853d8c Mon Sep 17 00:00:00 2001 From: Takuma Mori Date: Mon, 23 Jan 2023 23:04:46 +0900 Subject: [PATCH 04/10] add usage example to docstring Co-authored-by: Patrick von Platen --- src/diffusers/pipelines/pipeline_utils.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/src/diffusers/pipelines/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index 7cf32fc6c945..67a95e7642a5 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -853,6 +853,15 @@ def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Call Override the default `None` operator for use as `op` argument to the [`memory_efficient_attention()`](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.memory_efficient_attention) function of xFormers. + + Examples: + ```py + from diffusers import DiffusionPipeline + from xformers import ... # some attention op + + pipe = DiffusionPipeline.from_pretrained("CompVis/stable-diffusion-v1-4") + pipe.enable_xformers_memory_efficient_attention(attention_op=...) + ``` """ self.set_use_memory_efficient_attention_xformers(True, attention_op) From e783dc0aed59363ad01ebe5f6591838090a69b42 Mon Sep 17 00:00:00 2001 From: Takuma Mori Date: Mon, 23 Jan 2023 23:05:02 +0900 Subject: [PATCH 05/10] add usage example to docstring Co-authored-by: Patrick von Platen --- src/diffusers/models/modeling_utils.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index a6ada2c4a19b..e763283c1b40 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -222,6 +222,15 @@ def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Call Override the default `None` operator for use as `op` argument to the [`memory_efficient_attention()`](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.memory_efficient_attention) function of xFormers. + + Examples: + ```py + from diffusers import UNet2DConditionModel + from xformers import ... # some attention op + + model = UNet2DConditionModel.from_pretrained("CompVis/stable-diffusion-v1-4", subfolder="unet") + model.enable_xformers_memory_efficient_attention(attention_op=...) + ``` """ self.set_use_memory_efficient_attention_xformers(True, attention_op) From 9796e0be150032e71a62663831138732b9f770e4 Mon Sep 17 00:00:00 2001 From: Takuma Mori Date: Mon, 23 Jan 2023 23:11:53 +0900 Subject: [PATCH 06/10] code style correction by `make style` --- src/diffusers/models/modeling_utils.py | 4 ++-- src/diffusers/pipelines/pipeline_utils.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index e763283c1b40..c375418b44ef 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -222,12 +222,12 @@ def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Call Override the default `None` operator for use as `op` argument to the [`memory_efficient_attention()`](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.memory_efficient_attention) function of xFormers. - + Examples: ```py from diffusers import UNet2DConditionModel from xformers import ... # some attention op - + model = UNet2DConditionModel.from_pretrained("CompVis/stable-diffusion-v1-4", subfolder="unet") model.enable_xformers_memory_efficient_attention(attention_op=...) ``` diff --git a/src/diffusers/pipelines/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index 67a95e7642a5..40b91c259758 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -853,12 +853,12 @@ def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Call Override the default `None` operator for use as `op` argument to the [`memory_efficient_attention()`](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.memory_efficient_attention) function of xFormers. - + Examples: ```py from diffusers import DiffusionPipeline from xformers import ... # some attention op - + pipe = DiffusionPipeline.from_pretrained("CompVis/stable-diffusion-v1-4") pipe.enable_xformers_memory_efficient_attention(attention_op=...) ``` From 410be66e3fabaa3460de40929ee96a2c31d1e76c Mon Sep 17 00:00:00 2001 From: Takuma Mori Date: Tue, 24 Jan 2023 00:32:52 +0900 Subject: [PATCH 07/10] Update docstring code to a valid python example Co-authored-by: Suraj Patil --- src/diffusers/models/modeling_utils.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index c375418b44ef..b76e46588765 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -224,12 +224,13 @@ def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Call function of xFormers. Examples: + ```py - from diffusers import UNet2DConditionModel - from xformers import ... # some attention op + >>> from diffusers import UNet2DConditionModel + >>> from xformers.ops import MemoryEfficientAttentionFlashAttentionOp - model = UNet2DConditionModel.from_pretrained("CompVis/stable-diffusion-v1-4", subfolder="unet") - model.enable_xformers_memory_efficient_attention(attention_op=...) + >>> model = UNet2DConditionModel.from_pretrained("CompVis/stable-diffusion-v1-4", subfolder="unet").to("cuda") + >>> model.enable_xformers_memory_efficient_attention(attention_op=MemoryEfficientAttentionFlashAttentionOp) ``` """ self.set_use_memory_efficient_attention_xformers(True, attention_op) From 1a5a352ead83d238374cfa3e975e59b1ba97c220 Mon Sep 17 00:00:00 2001 From: Takuma Mori Date: Tue, 24 Jan 2023 00:33:02 +0900 Subject: [PATCH 08/10] Update docstring code to a valid python example Co-authored-by: Suraj Patil --- src/diffusers/pipelines/pipeline_utils.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/src/diffusers/pipelines/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index 40b91c259758..dd7d9b0508de 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -855,12 +855,13 @@ def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Call function of xFormers. Examples: + ```py - from diffusers import DiffusionPipeline - from xformers import ... # some attention op + >>> from diffusers import DiffusionPipeline + >>> from xformers.ops import MemoryEfficientAttentionFlashAttentionOp - pipe = DiffusionPipeline.from_pretrained("CompVis/stable-diffusion-v1-4") - pipe.enable_xformers_memory_efficient_attention(attention_op=...) + >>> pipe = DiffusionPipeline.from_pretrained("CompVis/stable-diffusion-v1-4").to("cuda") + >>> pipe.enable_xformers_memory_efficient_attention(attention_op=MemoryEfficientAttentionFlashAttentionOp) ``` """ self.set_use_memory_efficient_attention_xformers(True, attention_op) From 02a6f7e3c85fe190e3e3b10ea72e450dcaf219e8 Mon Sep 17 00:00:00 2001 From: Takuma Mori Date: Tue, 24 Jan 2023 00:35:21 +0900 Subject: [PATCH 09/10] style correction by `make style` --- src/diffusers/models/modeling_utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index b76e46588765..68751380e406 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -230,7 +230,7 @@ def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Call >>> from xformers.ops import MemoryEfficientAttentionFlashAttentionOp >>> model = UNet2DConditionModel.from_pretrained("CompVis/stable-diffusion-v1-4", subfolder="unet").to("cuda") - >>> model.enable_xformers_memory_efficient_attention(attention_op=MemoryEfficientAttentionFlashAttentionOp) + >>> model.enable_xformers_memory_efficient_attention(attention_op=MemoryEfficientAttentionFlashAttentionOp) ``` """ self.set_use_memory_efficient_attention_xformers(True, attention_op) From 1f05d4de21f4a869a2c18fb1ffd00eeaf750e822 Mon Sep 17 00:00:00 2001 From: Takuma Mori Date: Wed, 25 Jan 2023 01:06:36 +0900 Subject: [PATCH 10/10] Update code exmaple to fully functional --- src/diffusers/models/modeling_utils.py | 6 +++++- src/diffusers/pipelines/pipeline_utils.py | 6 +++++- 2 files changed, 10 insertions(+), 2 deletions(-) diff --git a/src/diffusers/models/modeling_utils.py b/src/diffusers/models/modeling_utils.py index 68751380e406..ceddf70b2cc8 100644 --- a/src/diffusers/models/modeling_utils.py +++ b/src/diffusers/models/modeling_utils.py @@ -226,10 +226,14 @@ def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Call Examples: ```py + >>> import torch >>> from diffusers import UNet2DConditionModel >>> from xformers.ops import MemoryEfficientAttentionFlashAttentionOp - >>> model = UNet2DConditionModel.from_pretrained("CompVis/stable-diffusion-v1-4", subfolder="unet").to("cuda") + >>> model = UNet2DConditionModel.from_pretrained( + ... "stabilityai/stable-diffusion-2-1", subfolder="unet", torch_dtype=torch.float16 + ... ) + >>> model = model.to("cuda") >>> model.enable_xformers_memory_efficient_attention(attention_op=MemoryEfficientAttentionFlashAttentionOp) ``` """ diff --git a/src/diffusers/pipelines/pipeline_utils.py b/src/diffusers/pipelines/pipeline_utils.py index 50660786022b..1c7d2c41a91c 100644 --- a/src/diffusers/pipelines/pipeline_utils.py +++ b/src/diffusers/pipelines/pipeline_utils.py @@ -861,11 +861,15 @@ def enable_xformers_memory_efficient_attention(self, attention_op: Optional[Call Examples: ```py + >>> import torch >>> from diffusers import DiffusionPipeline >>> from xformers.ops import MemoryEfficientAttentionFlashAttentionOp - >>> pipe = DiffusionPipeline.from_pretrained("CompVis/stable-diffusion-v1-4").to("cuda") + >>> pipe = DiffusionPipeline.from_pretrained("stabilityai/stable-diffusion-2-1", torch_dtype=torch.float16) + >>> pipe = pipe.to("cuda") >>> pipe.enable_xformers_memory_efficient_attention(attention_op=MemoryEfficientAttentionFlashAttentionOp) + >>> # Workaround for not accepting attention shape using VAE for Flash Attention + >>> pipe.vae.enable_xformers_memory_efficient_attention(attention_op=None) ``` """ self.set_use_memory_efficient_attention_xformers(True, attention_op)