MCPcopy Create free account
hub / github.com/LeapLabTHU/DAT / __init__

Method __init__

models/dat_blocks.py:132–216  ·  view source on GitHub ↗
(
        self, q_size, kv_size, n_heads, n_head_channels, n_groups,
        attn_drop, proj_drop, stride, 
        offset_range_factor, use_pe, dwc_pe,
        no_off, fixed_pe, ksize, log_cpb
    )

Source from the content-addressed store, hash-verified

130class DAttentionBaseline(nn.Module):
131
132 def __init__(
133 self, q_size, kv_size, n_heads, n_head_channels, n_groups,
134 attn_drop, proj_drop, stride,
135 offset_range_factor, use_pe, dwc_pe,
136 no_off, fixed_pe, ksize, log_cpb
137 ):
138
139 super().__init__()
140 self.dwc_pe = dwc_pe
141 self.n_head_channels = n_head_channels
142 self.scale = self.n_head_channels ** -0.5
143 self.n_heads = n_heads
144 self.q_h, self.q_w = q_size
145 # self.kv_h, self.kv_w = kv_size
146 self.kv_h, self.kv_w = self.q_h // stride, self.q_w // stride
147 self.nc = n_head_channels * n_heads
148 self.n_groups = n_groups
149 self.n_group_channels = self.nc // self.n_groups
150 self.n_group_heads = self.n_heads // self.n_groups
151 self.use_pe = use_pe
152 self.fixed_pe = fixed_pe
153 self.no_off = no_off
154 self.offset_range_factor = offset_range_factor
155 self.ksize = ksize
156 self.log_cpb = log_cpb
157 self.stride = stride
158 kk = self.ksize
159 pad_size = kk // 2 if kk != stride else 0
160
161 self.conv_offset = nn.Sequential(
162 nn.Conv2d(self.n_group_channels, self.n_group_channels, kk, stride, pad_size, groups=self.n_group_channels),
163 LayerNormProxy(self.n_group_channels),
164 nn.GELU(),
165 nn.Conv2d(self.n_group_channels, 2, 1, 1, 0, bias=False)
166 )
167 if self.no_off:
168 for m in self.conv_offset.parameters():
169 m.requires_grad_(False)
170
171 self.proj_q = nn.Conv2d(
172 self.nc, self.nc,
173 kernel_size=1, stride=1, padding=0
174 )
175
176 self.proj_k = nn.Conv2d(
177 self.nc, self.nc,
178 kernel_size=1, stride=1, padding=0
179 )
180
181 self.proj_v = nn.Conv2d(
182 self.nc, self.nc,
183 kernel_size=1, stride=1, padding=0
184 )
185
186 self.proj_out = nn.Conv2d(
187 self.nc, self.nc,
188 kernel_size=1, stride=1, padding=0
189 )

Callers

nothing calls this directly

Calls 2

LayerNormProxyClass · 0.85
__init__Method · 0.45

Tested by

no test coverage detected