MCPcopy Create free account
hub / github.com/pytorch/examples / forward

Method forward

distributed/FSDP2/model.py:32–51  ·  view source on GitHub ↗
(self, x)

Source from the content-addressed store, hash-verified

30 self.wo = nn.Linear(args.dim, args.dim, bias=False)
31
32 def forward(self, x):
33 bsz, seq_len, _ = x.size()
34 queries, keys, values = self.wq(x), self.wk(x), self.wv(x)
35 queries = queries.view(bsz, seq_len, self.n_heads, self.head_dim)
36 keys = keys.view(bsz, seq_len, self.n_heads, self.head_dim)
37 values = values.view(bsz, seq_len, self.n_heads, self.head_dim)
38
39 queries = queries.transpose(1, 2) # (bsz, n_heads, seq_len, head_dim)
40 keys = keys.transpose(1, 2) # (bsz, n_heads, seq_len, head_dim)
41 values = values.transpose(1, 2) # (bsz, n_heads, seq_len, head_dim)
42
43 output = F.scaled_dot_product_attention(
44 queries,
45 keys,
46 values,
47 None,
48 self.dropout_p if self.training else 0,
49 )
50 output = output.transpose(1, 2).contiguous().view(bsz, seq_len, -1)
51 return self.resid_dropout(self.wo(output))
52
53 def reset_parameters(self):
54 self.wq.reset_parameters()

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected