MCPcopy Create free account
hub / github.com/Yasoz/DiffTraj / AttnBlock

Class AttnBlock

utils/module.py:39–91  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

37
38
39class AttnBlock(nn.Module):
40 def __init__(self, in_channels):
41 super().__init__()
42 self.in_channels = in_channels
43
44 self.norm = Normalize(in_channels)
45 self.q = torch.nn.Conv2d(in_channels,
46 in_channels,
47 kernel_size=1,
48 stride=1,
49 padding=0)
50 self.k = torch.nn.Conv2d(in_channels,
51 in_channels,
52 kernel_size=1,
53 stride=1,
54 padding=0)
55 self.v = torch.nn.Conv2d(in_channels,
56 in_channels,
57 kernel_size=1,
58 stride=1,
59 padding=0)
60 self.proj_out = torch.nn.Conv2d(in_channels,
61 in_channels,
62 kernel_size=1,
63 stride=1,
64 padding=0)
65
66 def forward(self, x):
67 h_ = x
68 h_ = self.norm(h_)
69 q = self.q(h_)
70 k = self.k(h_)
71 v = self.v(h_)
72
73 # compute attention
74 b, c, h, w = q.shape
75 q = q.reshape(b, c, h * w)
76 q = q.permute(0, 2, 1) # b,hw,c
77 k = k.reshape(b, c, h * w) # b,c,hw
78 w_ = torch.bmm(q, k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
79 w_ = w_ * (int(c)**(-0.5))
80 w_ = torch.nn.functional.softmax(w_, dim=2)
81
82 # attend to values
83 v = v.reshape(b, c, h * w)
84 w_ = w_.permute(0, 2, 1) # b,hw,hw (first hw of k, second of q)
85 # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
86 h_ = torch.bmm(v, w_)
87 h_ = h_.reshape(b, c, h, w)
88
89 h_ = self.proj_out(h_)
90
91 return x + h_

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected