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

Class Transformer

distributed/FSDP2/model.py:100–134  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

98# A toy transformer model, partly inspired by the nanoGPT model:
99# https://github.com/karpathy/nanoGPT.
100class Transformer(nn.Module):
101 def __init__(self, args: ModelArgs):
102 super().__init__()
103 assert args.vocab_size is not None
104 assert args.max_seq_len is not None
105 self.model_args = args
106 self.max_seq_len = args.max_seq_len
107 self.tok_embeddings = nn.Embedding(args.vocab_size, args.dim)
108 self.pos_embeddings = nn.Embedding(args.max_seq_len, args.dim)
109 self.dropout = nn.Dropout(args.dropout_p)
110 self.layers = nn.ModuleList()
111 for _ in range(args.n_layers):
112 self.layers.append(TransformerBlock(args))
113 self.norm = nn.LayerNorm(args.dim)
114 self.output = nn.Linear(args.dim, args.vocab_size, bias=False)
115
116 def forward(self, tokens):
117 _bsz, seq_len = tokens.size()
118 assert seq_len <= self.max_seq_len
119 h = self.tok_embeddings(tokens)
120 pos = torch.arange(0, seq_len, device=tokens.device)
121 p = self.pos_embeddings(pos) # positional embeddings of shape (seq_len, dim)
122 h = h + p
123 h = self.dropout(h)
124 for layer in self.layers:
125 h = layer(h)
126 h = self.norm(h)
127 output = self.output(h).float()
128 return output
129
130 def reset_parameters(self):
131 self.tok_embeddings.reset_parameters()
132 self.pos_embeddings.reset_parameters()
133 self.norm.reset_parameters()
134 self.output.reset_parameters()

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected