(
self,
query_dim: int,
cross_attention_dim: Optional[int] = None,
heads: int = 8,
dim_head: int = 64,
dropout: float = 0.0,
bias=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=False,
processor: Optional["AttnProcessor"] = None,
)
| 49 | """ |
| 50 | |
| 51 | def __init__( |
| 52 | self, |
| 53 | query_dim: int, |
| 54 | cross_attention_dim: Optional[int] = None, |
| 55 | heads: int = 8, |
| 56 | dim_head: int = 64, |
| 57 | dropout: float = 0.0, |
| 58 | bias=False, |
| 59 | upcast_attention: bool = False, |
| 60 | upcast_softmax: bool = False, |
| 61 | cross_attention_norm: Optional[str] = None, |
| 62 | cross_attention_norm_num_groups: int = 32, |
| 63 | added_kv_proj_dim: Optional[int] = None, |
| 64 | norm_num_groups: Optional[int] = None, |
| 65 | spatial_norm_dim: Optional[int] = None, |
| 66 | out_bias: bool = True, |
| 67 | scale_qk: bool = True, |
| 68 | only_cross_attention: bool = False, |
| 69 | eps: float = 1e-5, |
| 70 | rescale_output_factor: float = 1.0, |
| 71 | residual_connection: bool = False, |
| 72 | _from_deprecated_attn_block=False, |
| 73 | processor: Optional["AttnProcessor"] = None, |
| 74 | ): |
| 75 | super().__init__() |
| 76 | inner_dim = dim_head * heads |
| 77 | cross_attention_dim = cross_attention_dim if cross_attention_dim is not None else query_dim |
| 78 | self.upcast_attention = upcast_attention |
| 79 | self.upcast_softmax = upcast_softmax |
| 80 | self.rescale_output_factor = rescale_output_factor |
| 81 | self.residual_connection = residual_connection |
| 82 | |
| 83 | # we make use of this private variable to know whether this class is loaded |
| 84 | # with an deprecated state dict so that we can convert it on the fly |
| 85 | self._from_deprecated_attn_block = _from_deprecated_attn_block |
| 86 | |
| 87 | self.scale_qk = scale_qk |
| 88 | self.scale = dim_head**-0.5 if self.scale_qk else 1.0 |
| 89 | |
| 90 | self.heads = heads |
| 91 | # for slice_size > 0 the attention score computation |
| 92 | # is split across the batch axis to save memory |
| 93 | # You can set slice_size with `set_attention_slice` |
| 94 | self.sliceable_head_dim = heads |
| 95 | |
| 96 | self.added_kv_proj_dim = added_kv_proj_dim |
| 97 | self.only_cross_attention = only_cross_attention |
| 98 | |
| 99 | if self.added_kv_proj_dim is None and self.only_cross_attention: |
| 100 | raise ValueError( |
| 101 | "`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`." |
| 102 | ) |
| 103 | |
| 104 | if norm_num_groups is not None: |
| 105 | self.group_norm = nn.GroupNorm(num_channels=query_dim, num_groups=norm_num_groups, eps=eps, affine=True) |
| 106 | else: |
| 107 | self.group_norm = None |
| 108 |
no test coverage detected