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

Method forward

model/downstream.py:131–145  ·  view source on GitHub ↗
(self, input_ids, position_ids, attention_mask, target_ids=None, logit_mask=None, prompt_pos=None)

Source from the content-addressed store, hash-verified

129 return self.model.named_parameters(prefix=prefix, recurse=recurse)
130
131 def forward(self, input_ids, position_ids, attention_mask, target_ids=None, logit_mask=None, prompt_pos=None):
132 if target_ids is None:
133 return self.model(input_ids, position_ids, attention_mask)
134 assert len(input_ids.shape) == 2
135 outputs, *mems = self.model(input_ids, position_ids, attention_mask, prompt_pos=prompt_pos)
136 batch_ids = torch.arange(outputs.size(0), dtype=attention_mask.dtype, device=attention_mask.device)
137 target_logits = outputs[batch_ids, attention_mask]
138 if self.take_softmax:
139 target_prob = torch.nn.functional.log_softmax(target_logits, dim=-1)
140 else:
141 target_prob = target_logits
142 batch_ids = batch_ids.unsqueeze(1).expand_as(target_ids)
143 output = target_prob[batch_ids, target_ids]
144
145 return (output, target_logits, *mems)
146
147
148class GLMForSequenceClassification(torch.nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected