(
self,
query_dim: int,
cross_attention_dim: Optional[int] = None,
heads: int = 8,
dim_head: int = 64,
dropout: float = 0.0,
bias: bool = False,
upcast_attention: bool = False,
upcast_softmax: bool = False,
cross_attention_norm: Optional[str] = None,
cross_attention_norm_num_groups: int = 32,
added_kv_proj_dim: Optional[int] = None,
norm_num_groups: Optional[int] = None,
spatial_norm_dim: Optional[int] = None,
out_bias: bool = True,
scale_qk: bool = True,
only_cross_attention: bool = False,
eps: float = 1e-5,
rescale_output_factor: float = 1.0,
residual_connection: bool = False,
_from_deprecated_attn_block: bool = False,
processor: Optional["AttnProcessor"] = None,
updown=None,
)
| 74 | """ |
| 75 | |
| 76 | def __init__( |
| 77 | self, |
| 78 | query_dim: int, |
| 79 | cross_attention_dim: Optional[int] = None, |
| 80 | heads: int = 8, |
| 81 | dim_head: int = 64, |
| 82 | dropout: float = 0.0, |
| 83 | bias: bool = False, |
| 84 | upcast_attention: bool = False, |
| 85 | upcast_softmax: bool = False, |
| 86 | cross_attention_norm: Optional[str] = None, |
| 87 | cross_attention_norm_num_groups: int = 32, |
| 88 | added_kv_proj_dim: Optional[int] = None, |
| 89 | norm_num_groups: Optional[int] = None, |
| 90 | spatial_norm_dim: Optional[int] = None, |
| 91 | out_bias: bool = True, |
| 92 | scale_qk: bool = True, |
| 93 | only_cross_attention: bool = False, |
| 94 | eps: float = 1e-5, |
| 95 | rescale_output_factor: float = 1.0, |
| 96 | residual_connection: bool = False, |
| 97 | _from_deprecated_attn_block: bool = False, |
| 98 | processor: Optional["AttnProcessor"] = None, |
| 99 | updown=None, |
| 100 | ): |
| 101 | super().__init__() |
| 102 | self.inner_dim = dim_head * heads |
| 103 | self.cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim |
| 104 | self.upcast_attention = upcast_attention |
| 105 | self.upcast_softmax = upcast_softmax |
| 106 | self.rescale_output_factor = rescale_output_factor |
| 107 | self.residual_connection = residual_connection |
| 108 | self.dropout = dropout |
| 109 | |
| 110 | # we make use of this private variable to know whether this class is loaded |
| 111 | # with an deprecated state dict so that we can convert it on the fly |
| 112 | self._from_deprecated_attn_block = _from_deprecated_attn_block |
| 113 | |
| 114 | self.scale_qk = scale_qk |
| 115 | self.scale = dim_head**-0.5 if self.scale_qk else 1.0 |
| 116 | |
| 117 | self.heads = heads |
| 118 | # for slice_size > 0 the attention score computation |
| 119 | # is split across the batch axis to save memory |
| 120 | # You can set slice_size with `set_attention_slice` |
| 121 | self.sliceable_head_dim = heads |
| 122 | |
| 123 | self.added_kv_proj_dim = added_kv_proj_dim |
| 124 | self.only_cross_attention = only_cross_attention |
| 125 | self.updown = updown |
| 126 | |
| 127 | if self.added_kv_proj_dim is None and self.only_cross_attention: |
| 128 | raise ValueError( |
| 129 | "`only_cross_attention` can only be set to True if `added_kv_proj_dim` is not None. Make sure to set either `only_cross_attention=False` or define `added_kv_proj_dim`." |
| 130 | ) |
| 131 | |
| 132 | if norm_num_groups is not None: |
| 133 | self.group_norm = nn.GroupNorm(num_channels=query_dim, num_groups=norm_num_groups, eps=eps, affine=True) |
no test coverage detected