MCPcopy Create free account
hub / github.com/AIS-SNU/Smart-Infinity / forward

Method forward

DeepSpeedExample/megatron/model/bert_model.py:135–171  ·  view source on GitHub ↗
(self, input_ids, attention_mask,
                tokentype_ids=None, lm_labels=None)

Source from the content-addressed store, hash-verified

133 self._binary_head_key = 'binary_head'
134
135 def forward(self, input_ids, attention_mask,
136 tokentype_ids=None, lm_labels=None):
137
138 extended_attention_mask = bert_extended_attention_mask(attention_mask)
139 position_ids = bert_position_ids(input_ids)
140
141 if self.add_binary_head:
142 lm_output, pooled_output = self.language_model(
143 input_ids,
144 position_ids,
145 extended_attention_mask,
146 tokentype_ids=tokentype_ids)
147 else:
148 lm_output = self.language_model(
149 input_ids,
150 position_ids,
151 extended_attention_mask,
152 tokentype_ids=tokentype_ids)
153
154 # Output.
155 lm_logits = self.lm_head(
156 lm_output, self.language_model.embedding.word_embeddings.weight)
157
158 binary_logits = None
159 if self.add_binary_head:
160 binary_logits = self.binary_head(pooled_output)
161
162 if lm_labels is None:
163 return lm_logits, binary_logits
164 else:
165 if self.fp16_lm_cross_entropy:
166 assert lm_logits.dtype == torch.half
167 lm_loss = mpu.vocab_parallel_cross_entropy(lm_logits, lm_labels)
168 else:
169 lm_loss = mpu.vocab_parallel_cross_entropy(lm_logits.float(),
170 lm_labels)
171 return lm_loss, binary_logits
172
173
174 def state_dict_for_save_checkpoint(self, destination=None, prefix='',

Callers

nothing calls this directly

Calls 2

bert_position_idsFunction · 0.85

Tested by

no test coverage detected