| 189 | |
| 190 | |
| 191 | class 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, |