(self, input_ids, position_ids, attention_mask, *mems, return_memory=False, detach_memory=True,
prompt_pos=None)
| 105 | print_rank_0(log_str) |
| 106 | |
| 107 | def forward(self, input_ids, position_ids, attention_mask, *mems, return_memory=False, detach_memory=True, |
| 108 | prompt_pos=None): |
| 109 | # Embeddings. |
| 110 | batch_size = input_ids.size(0) |
| 111 | words_embeddings = self.word_embeddings(input_ids) |
| 112 | embeddings = words_embeddings |
| 113 | if prompt_pos is not None: |
| 114 | embeddings = embeddings.clone() |
| 115 | prompt_embeds = self.prompt_spell() |
| 116 | batch_index = torch.arange(batch_size, device=input_ids.device).unsqueeze(1) |
| 117 | embeddings[batch_index, prompt_pos] = prompt_embeds |
| 118 | # Transformer. |
| 119 | transformer_output = self.transformer(embeddings, position_ids, attention_mask, mems, |
| 120 | return_memory=return_memory, detach_memory=detach_memory) |
| 121 | logits, hidden_layers = transformer_output |
| 122 | outputs = hidden_layers |
| 123 | |
| 124 | if self.output_predict: |
| 125 | # Parallel logits. |
| 126 | logits_parallel = mpu.copy_to_model_parallel_region( |
| 127 | logits) |
| 128 | logits_parallel = F.linear(logits_parallel, self.word_embeddings.weight) |
| 129 | |
| 130 | if self.parallel_output: |
| 131 | return (logits_parallel, *outputs) |
| 132 | |
| 133 | return (mpu.gather_from_model_parallel_region(logits_parallel), *outputs) |
| 134 | else: |
| 135 | return (logits, *outputs) |
| 136 | |
| 137 | |
| 138 | class EncoderDecoder(torch.nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected