MCPcopy Create free account
hub / github.com/MotrixLab/ViMoGen / T5RelativeEmbedding

Class T5RelativeEmbedding

models/transformer/wan/modules/t5.py:221–264  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

219
220
221class T5RelativeEmbedding(nn.Module):
222
223 def __init__(self, num_buckets, num_heads, bidirectional, max_dist=128):
224 super(T5RelativeEmbedding, self).__init__()
225 self.num_buckets = num_buckets
226 self.num_heads = num_heads
227 self.bidirectional = bidirectional
228 self.max_dist = max_dist
229
230 # layers
231 self.embedding = nn.Embedding(num_buckets, num_heads)
232
233 def forward(self, lq, lk):
234 device = self.embedding.weight.device
235 # rel_pos = torch.arange(lk).unsqueeze(0).to(device) - \
236 # torch.arange(lq).unsqueeze(1).to(device)
237 rel_pos = torch.arange(lk, device=device).unsqueeze(0) - \
238 torch.arange(lq, device=device).unsqueeze(1)
239 rel_pos = self._relative_position_bucket(rel_pos)
240 rel_pos_embeds = self.embedding(rel_pos)
241 rel_pos_embeds = rel_pos_embeds.permute(2, 0, 1).unsqueeze(
242 0) # [1, N, Lq, Lk]
243 return rel_pos_embeds.contiguous()
244
245 def _relative_position_bucket(self, rel_pos):
246 # preprocess
247 if self.bidirectional:
248 num_buckets = self.num_buckets // 2
249 rel_buckets = (rel_pos > 0).long() * num_buckets
250 rel_pos = torch.abs(rel_pos)
251 else:
252 num_buckets = self.num_buckets
253 rel_buckets = 0
254 rel_pos = -torch.min(rel_pos, torch.zeros_like(rel_pos))
255
256 # embeddings for small and large positions
257 max_exact = num_buckets // 2
258 rel_pos_large = max_exact + (torch.log(rel_pos.float() / max_exact) /
259 math.log(self.max_dist / max_exact) *
260 (num_buckets - max_exact)).long()
261 rel_pos_large = torch.min(
262 rel_pos_large, torch.full_like(rel_pos_large, num_buckets - 1))
263 rel_buckets += torch.where(rel_pos < max_exact, rel_pos, rel_pos_large)
264 return rel_buckets
265
266
267class T5Encoder(nn.Module):

Callers 4

__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected