MCPcopy Create free account
hub / github.com/THUDM/GLM / forward

Method forward

model/modeling_glm.py:192–210  ·  view source on GitHub ↗
(self, source_ids, target_ids, source_position_ids, target_position_ids, source_mask, target_mask)

Source from the content-addressed store, hash-verified

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
213def glm_get_params_for_weight_decay_optimization(module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected