Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
attention.py89 linesDownload Raw Back to models
1import torch2from einops import rearrange3 4 5def low_version_attention(query, key, value, attn_bias=None):6    scale = 1 / query.shape[-1] ** 0.57    query = query * scale8    attn = torch.matmul(query, key.transpose(-2, -1))9    if attn_bias is not None:10        attn = attn + attn_bias11    attn = attn.softmax(-1)12    return attn @ value13 14 15class Attention(torch.nn.Module):16 17    def __init__(self, q_dim, num_heads, head_dim, kv_dim=None, bias_q=False, bias_kv=False, bias_out=False):18        super().__init__()19        dim_inner = head_dim * num_heads20        kv_dim = kv_dim if kv_dim is not None else q_dim21        self.num_heads = num_heads22        self.head_dim = head_dim23 24        self.to_q = torch.nn.Linear(q_dim, dim_inner, bias=bias_q)25        self.to_k = torch.nn.Linear(kv_dim, dim_inner, bias=bias_kv)26        self.to_v = torch.nn.Linear(kv_dim, dim_inner, bias=bias_kv)27        self.to_out = torch.nn.Linear(dim_inner, q_dim, bias=bias_out)28 29    def interact_with_ipadapter(self, hidden_states, q, ip_k, ip_v, scale=1.0):30        batch_size = q.shape[0]31        ip_k = ip_k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)32        ip_v = ip_v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)33        ip_hidden_states = torch.nn.functional.scaled_dot_product_attention(q, ip_k, ip_v)34        hidden_states = hidden_states + scale * ip_hidden_states35        return hidden_states36 37    def torch_forward(self, hidden_states, encoder_hidden_states=None, attn_mask=None, ipadapter_kwargs=None, qkv_preprocessor=None):38        if encoder_hidden_states is None:39            encoder_hidden_states = hidden_states40 41        batch_size = encoder_hidden_states.shape[0]42 43        q = self.to_q(hidden_states)44        k = self.to_k(encoder_hidden_states)45        v = self.to_v(encoder_hidden_states)46 47        q = q.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)48        k = k.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)49        v = v.view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2)50 51        if qkv_preprocessor is not None:52            q, k, v = qkv_preprocessor(q, k, v)53 54        hidden_states = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)55        if ipadapter_kwargs is not None:56            hidden_states = self.interact_with_ipadapter(hidden_states, q, **ipadapter_kwargs)57        hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, self.num_heads * self.head_dim)58        hidden_states = hidden_states.to(q.dtype)59 60        hidden_states = self.to_out(hidden_states)61 62        return hidden_states63    64    def xformers_forward(self, hidden_states, encoder_hidden_states=None, attn_mask=None):65        if encoder_hidden_states is None:66            encoder_hidden_states = hidden_states67 68        q = self.to_q(hidden_states)69        k = self.to_k(encoder_hidden_states)70        v = self.to_v(encoder_hidden_states)71 72        q = rearrange(q, "b f (n d) -> (b n) f d", n=self.num_heads)73        k = rearrange(k, "b f (n d) -> (b n) f d", n=self.num_heads)74        v = rearrange(v, "b f (n d) -> (b n) f d", n=self.num_heads)75 76        if attn_mask is not None:77            hidden_states = low_version_attention(q, k, v, attn_bias=attn_mask)78        else:79            import xformers.ops as xops80            hidden_states = xops.memory_efficient_attention(q, k, v)81        hidden_states = rearrange(hidden_states, "(b n) f d -> b f (n d)", n=self.num_heads)82 83        hidden_states = hidden_states.to(q.dtype)84        hidden_states = self.to_out(hidden_states)85 86        return hidden_states87 88    def forward(self, hidden_states, encoder_hidden_states=None, attn_mask=None, ipadapter_kwargs=None, qkv_preprocessor=None):89        return self.torch_forward(hidden_states, encoder_hidden_states=encoder_hidden_states, attn_mask=attn_mask, ipadapter_kwargs=ipadapter_kwargs, qkv_preprocessor=qkv_preprocessor)