r""" Default processor for performing attention-related computations.
| 1084 | |
| 1085 | |
| 1086 | class AttnProcessor: |
| 1087 | r""" |
| 1088 | Default processor for performing attention-related computations. |
| 1089 | """ |
| 1090 | |
| 1091 | def __call__( |
| 1092 | self, |
| 1093 | attn: Attention, |
| 1094 | hidden_states: torch.Tensor, |
| 1095 | encoder_hidden_states: Optional[torch.Tensor] = None, |
| 1096 | attention_mask: Optional[torch.Tensor] = None, |
| 1097 | temb: Optional[torch.Tensor] = None, |
| 1098 | *args, |
| 1099 | **kwargs, |
| 1100 | ) -> torch.Tensor: |
| 1101 | if len(args) > 0 or kwargs.get("scale", None) is not None: |
| 1102 | deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." |
| 1103 | deprecate("scale", "1.0.0", deprecation_message) |
| 1104 | |
| 1105 | residual = hidden_states |
| 1106 | |
| 1107 | if attn.spatial_norm is not None: |
| 1108 | hidden_states = attn.spatial_norm(hidden_states, temb) |
| 1109 | |
| 1110 | input_ndim = hidden_states.ndim |
| 1111 | |
| 1112 | if input_ndim == 4: |
| 1113 | batch_size, channel, height, width = hidden_states.shape |
| 1114 | hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) |
| 1115 | |
| 1116 | batch_size, sequence_length, _ = ( |
| 1117 | hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape |
| 1118 | ) |
| 1119 | attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) |
| 1120 | |
| 1121 | if attn.group_norm is not None: |
| 1122 | hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) |
| 1123 | |
| 1124 | query = attn.to_q(hidden_states) |
| 1125 | |
| 1126 | if encoder_hidden_states is None: |
| 1127 | encoder_hidden_states = hidden_states |
| 1128 | elif attn.norm_cross: |
| 1129 | encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) |
| 1130 | |
| 1131 | key = attn.to_k(encoder_hidden_states) |
| 1132 | value = attn.to_v(encoder_hidden_states) |
| 1133 | |
| 1134 | query = attn.head_to_batch_dim(query) |
| 1135 | key = attn.head_to_batch_dim(key) |
| 1136 | value = attn.head_to_batch_dim(value) |
| 1137 | |
| 1138 | attention_probs = attn.get_attention_scores(query, key, attention_mask) |
| 1139 | hidden_states = torch.bmm(attention_probs, value) |
| 1140 | hidden_states = attn.batch_to_head_dim(hidden_states) |
| 1141 | |
| 1142 | # linear proj |
| 1143 | hidden_states = attn.to_out[0](hidden_states) |
no outgoing calls
no test coverage detected