(self, source_ids, target_ids, source_position_ids, target_position_ids, source_mask, target_mask)
| 190 | use_decoder_layer=True) |
| 191 | |
| 192 | def forward(self, source_ids, target_ids, source_position_ids, target_position_ids, source_mask, target_mask): |
| 193 | # Embeddings. |
| 194 | source_embeddings = self.word_embeddings(source_ids) |
| 195 | target_embeddings = self.word_embeddings(target_ids) |
| 196 | |
| 197 | # Transformer. |
| 198 | encoder_output, _ = self.encoder(source_embeddings, source_position_ids, source_mask) |
| 199 | decoder_output, _ = self.decoder(target_embeddings, target_position_ids, target_mask) |
| 200 | if self.output_predict: |
| 201 | # Parallel logits. |
| 202 | output_parallel = mpu.copy_to_model_parallel_region(decoder_output) |
| 203 | logits_parallel = F.linear(output_parallel, self.word_embeddings.weight) |
| 204 | |
| 205 | if self.parallel_output: |
| 206 | return (logits_parallel,) |
| 207 | |
| 208 | return (mpu.gather_from_model_parallel_region(logits_parallel),) |
| 209 | else: |
| 210 | return (decoder_output,) |
| 211 | |
| 212 | |
| 213 | def glm_get_params_for_weight_decay_optimization(module): |
nothing calls this directly
no outgoing calls
no test coverage detected