MCPcopy Create free account
hub / github.com/CompVis/zigma / CrossAttention

Class CrossAttention

model_zigma.py:95–135  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

93
94
95class CrossAttention(nn.Module):
96 def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0.0):
97 super().__init__()
98 inner_dim = dim_head * heads
99 context_dim = default(context_dim, query_dim)
100
101 self.scale = dim_head**-0.5
102 self.heads = heads
103
104 self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
105 self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
106 self.to_v = nn.Linear(context_dim, inner_dim, bias=False)
107
108 self.to_out = nn.Sequential(
109 nn.Linear(inner_dim, query_dim), nn.Dropout(dropout)
110 )
111
112 def forward(self, x, text, mask=None):
113 B, L, C = x.shape
114 q = self.to_q(x)
115 # text = default(text, x)
116 k = self.to_k(text)
117 v = self.to_v(text)
118
119 q, k, v = map(
120 lambda t: rearrange(t, "B L (H D) -> B H L D", H=self.heads), (q, k, v)
121 ) # B H L D
122 if ATTENTION_MODE == "flash":
123 x = torch.nn.functional.scaled_dot_product_attention(q, k, v)
124 x = einops.rearrange(x, "B H L D -> B L (H D)")
125 elif ATTENTION_MODE == "xformers":
126 x = xformers.ops.memory_efficient_attention(q, k, v)
127 x = einops.rearrange(x, "B L H D -> B L (H D)", H=self.heads)
128 elif ATTENTION_MODE == "math":
129 attn = (q @ k.transpose(-2, -1)) * self.scale
130 attn = attn.softmax(dim=-1)
131 attn = self.attn_drop(attn)
132 x = (attn @ v).transpose(1, 2).reshape(B, L, C)
133 else:
134 raise NotImplemented
135 return self.to_out(x)
136
137
138def drop_path(

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected