MCPcopy Create free account
hub / github.com/Francis-Rings/StableAnimator / forward

Method forward

animation/modules/id_encoder.py:28–58  ·  view source on GitHub ↗

Args: x (torch.Tensor): image features shape (b, n1, D) latent (torch.Tensor): latent features shape (b, n2, D)

(self, x, latents)

Source from the content-addressed store, hash-verified

26 self.to_out = nn.Linear(inner_dim, dim, bias=False)
27
28 def forward(self, x, latents):
29 """
30 Args:
31 x (torch.Tensor): image features
32 shape (b, n1, D)
33 latent (torch.Tensor): latent features
34 shape (b, n2, D)
35 """
36
37 x = self.norm1(x)
38 latents = self.norm2(latents)
39
40 b, l, _ = latents.shape
41
42 q = self.to_q(latents)
43 kv_input = torch.cat((x, latents), dim=-2)
44 k, v = self.to_kv(kv_input).chunk(2, dim=-1)
45
46 q = reshape_tensor(q, self.heads)
47 k = reshape_tensor(k, self.heads)
48 v = reshape_tensor(v, self.heads)
49
50 # attention
51 scale = 1 / math.sqrt(math.sqrt(self.dim_head))
52 weight = (q * scale) @ (k * scale).transpose(-2, -1)
53 weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
54 out = weight @ v
55
56 out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
57
58 return self.to_out(out)
59
60def FeedForward(dim, mult=4):
61 inner_dim = int(dim * mult)

Callers

nothing calls this directly

Calls 1

reshape_tensorFunction · 0.85

Tested by

no test coverage detected