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

Method forward

model/modeling_glm.py:107–135  ·  view source on GitHub ↗
(self, input_ids, position_ids, attention_mask, *mems, return_memory=False, detach_memory=True,
                prompt_pos=None)

Source from the content-addressed store, hash-verified

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
138class EncoderDecoder(torch.nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected