(self, x, context=None, training=False, mask=None)
| 229 | change_checkpoint(self.model.diffusion_model) |
| 230 | |
| 231 | def new_forward(self, x, context=None, training=False, mask=None): |
| 232 | h = self.heads |
| 233 | crossattn = False |
| 234 | if context is not None: |
| 235 | crossattn = True |
| 236 | q = self.to_q(x) |
| 237 | context = default(context, x) |
| 238 | k = self.to_k(context) |
| 239 | v = self.to_v(context) |
| 240 | if crossattn: |
| 241 | modifier = torch.ones_like(k) |
| 242 | modifier[:, :1, :] = modifier[:, :1, :]*0. |
| 243 | k = modifier*k + (1-modifier)*k.detach() |
| 244 | v = modifier*v + (1-modifier)*v.detach() |
| 245 | |
| 246 | q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v)) |
| 247 | sim = einsum('b i d, b j d -> b i j', q, k) * self.scale |
| 248 | attn = sim.softmax(dim=-1) |
| 249 | |
| 250 | if crossattn and (attention_store_background_mask or attention_store_step): |
| 251 | if attn16: |
| 252 | if attn.shape[1] == 256: |
| 253 | controller(attn) |
| 254 | else: |
| 255 | if attn.shape[1] in store_list: |
| 256 | controller(attn) |
| 257 | elif crossattn and check_token_attn: |
| 258 | if attn.shape[1] == 64: |
| 259 | place = 'mid' |
| 260 | else: |
| 261 | place = 'down' |
| 262 | controller(attn, crossattn, place) |
| 263 | |
| 264 | out = einsum('b i j, b j d -> b i d', attn, v) |
| 265 | out = rearrange(out, '(b h) n d -> b n (h d)', h=h) |
| 266 | return self.to_out(out) |
| 267 | cross_att_count = 0 |
| 268 | def change_forward(model, count): |
| 269 | for layer in model.children(): |
nothing calls this directly
no outgoing calls
no test coverage detected