MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / AttnProcessor

Class AttnProcessor

architecture/attention_processor.py:1086–1155  ·  view source on GitHub ↗

r""" Default processor for performing attention-related computations.

Source from the content-addressed store, hash-verified

1084
1085
1086class 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)

Calls

no outgoing calls

Tested by

no test coverage detected