MCPcopy Create free account
hub / github.com/TencentARC/BrushNet / AttnAddedKVProcessor

Class AttnAddedKVProcessor

src/diffusers/models/attention_processor.py:905–966  ·  view source on GitHub ↗

r""" Processor for performing attention-related computations with extra learnable key and value matrices for the text encoder.

Source from the content-addressed store, hash-verified

903
904
905class AttnAddedKVProcessor:
906 r"""
907 Processor for performing attention-related computations with extra learnable key and value matrices for the text
908 encoder.
909 """
910
911 def __call__(
912 self,
913 attn: Attention,
914 hidden_states: torch.FloatTensor,
915 encoder_hidden_states: Optional[torch.FloatTensor] = None,
916 attention_mask: Optional[torch.FloatTensor] = None,
917 scale: float = 1.0,
918 ) -> torch.Tensor:
919 residual = hidden_states
920
921 args = () if USE_PEFT_BACKEND else (scale,)
922
923 hidden_states = hidden_states.view(hidden_states.shape[0], hidden_states.shape[1], -1).transpose(1, 2)
924 batch_size, sequence_length, _ = hidden_states.shape
925
926 attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
927
928 if encoder_hidden_states is None:
929 encoder_hidden_states = hidden_states
930 elif attn.norm_cross:
931 encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
932
933 hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
934
935 query = attn.to_q(hidden_states, *args)
936 query = attn.head_to_batch_dim(query)
937
938 encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states, *args)
939 encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states, *args)
940 encoder_hidden_states_key_proj = attn.head_to_batch_dim(encoder_hidden_states_key_proj)
941 encoder_hidden_states_value_proj = attn.head_to_batch_dim(encoder_hidden_states_value_proj)
942
943 if not attn.only_cross_attention:
944 key = attn.to_k(hidden_states, *args)
945 value = attn.to_v(hidden_states, *args)
946 key = attn.head_to_batch_dim(key)
947 value = attn.head_to_batch_dim(value)
948 key = torch.cat([encoder_hidden_states_key_proj, key], dim=1)
949 value = torch.cat([encoder_hidden_states_value_proj, value], dim=1)
950 else:
951 key = encoder_hidden_states_key_proj
952 value = encoder_hidden_states_value_proj
953
954 attention_probs = attn.get_attention_scores(query, key, attention_mask)
955 hidden_states = torch.bmm(attention_probs, value)
956 hidden_states = attn.batch_to_head_dim(hidden_states)
957
958 # linear proj
959 hidden_states = attn.to_out[0](hidden_states, *args)
960 # dropout
961 hidden_states = attn.to_out[1](hidden_states)
962

Calls

no outgoing calls