MCPcopy Create free account
hub / github.com/MLSysU/TD-Pipe / OPTDecoder

Class OPTDecoder

TD_Pipe/model_executor/models/opt.py:191–282  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

189
190
191class OPTDecoder(nn.Module):
192
193 def __init__(
194 self,
195 config: OPTConfig,
196 parallel_config: ParallelConfig,
197 attention_backend: str,
198 linear_method: Optional[LinearMethodBase] = None,
199 ):
200 super().__init__()
201 self.config = config
202 self.padding_idx = config.pad_token_id
203 self.max_target_positions = config.max_position_embeddings
204 self.vocab_size = config.vocab_size
205 self.parallel_config = parallel_config
206
207 self.embed_tokens = VocabParallelEmbedding(
208 config.vocab_size,
209 config.word_embed_proj_dim,
210 )
211 # Positional embeddings are replicated (not sharded).
212 self.embed_positions = OPTLearnedPositionalEmbedding(
213 config.max_position_embeddings, config.hidden_size)
214
215 if self.parallel_config.is_last:
216 # Project out & in will be replicated if they exist.
217 if config.word_embed_proj_dim != config.hidden_size:
218 self.project_out = ReplicatedLinear(
219 config.hidden_size,
220 config.word_embed_proj_dim,
221 bias=False,
222 linear_method=linear_method)
223 else:
224 self.project_out = None
225
226 if self.parallel_config.is_first:
227 if config.word_embed_proj_dim != config.hidden_size:
228 self.project_in = ReplicatedLinear(config.word_embed_proj_dim,
229 config.hidden_size,
230 bias=False,
231 linear_method=linear_method)
232 else:
233 self.project_in = None
234
235 if self.parallel_config.is_last:
236 # Note that the only purpose of `config._remove_final_layer_norm` is to
237 # keep backward compatibility with checkpoints that have been fine-tuned
238 # before transformers v4.20.1
239 # see https://github.com/facebookresearch/metaseq/pull/164
240 if config.do_layer_norm_before and not config._remove_final_layer_norm:
241 self.final_layer_norm = nn.LayerNorm(
242 config.hidden_size,
243 elementwise_affine=config.layer_norm_elementwise_affine)
244 else:
245 self.final_layer_norm = None
246
247 self.layers = nn.ModuleList()
248 for i in range(self.parallel_config.start,

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected