| 1208 | |
| 1209 | |
| 1210 | class ChatGLMForSequenceClassification(ChatGLMPreTrainedModel): |
| 1211 | def __init__(self, config: ChatGLMConfig, empty_init=True, device=None): |
| 1212 | super().__init__(config) |
| 1213 | |
| 1214 | self.num_labels = config.num_labels |
| 1215 | self.transformer = ChatGLMModel(config, empty_init=empty_init, device=device) |
| 1216 | |
| 1217 | self.classifier_head = nn.Linear(config.hidden_size, config.num_labels, bias=True, dtype=torch.half) |
| 1218 | if config.classifier_dropout is not None: |
| 1219 | self.dropout = nn.Dropout(config.classifier_dropout) |
| 1220 | else: |
| 1221 | self.dropout = None |
| 1222 | self.config = config |
| 1223 | |
| 1224 | if self.config.quantization_bit: |
| 1225 | self.quantize(self.config.quantization_bit, empty_init=True) |
| 1226 | |
| 1227 | def forward( |
| 1228 | self, |
| 1229 | input_ids: Optional[torch.LongTensor] = None, |
| 1230 | position_ids: Optional[torch.LongTensor] = None, |
| 1231 | attention_mask: Optional[torch.Tensor] = None, |
| 1232 | full_attention_mask: Optional[torch.Tensor] = None, |
| 1233 | past_key_values: Optional[Tuple[Tuple[torch.Tensor, torch.Tensor], ...]] = None, |
| 1234 | inputs_embeds: Optional[torch.LongTensor] = None, |
| 1235 | labels: Optional[torch.LongTensor] = None, |
| 1236 | use_cache: Optional[bool] = None, |
| 1237 | output_hidden_states: Optional[bool] = None, |
| 1238 | return_dict: Optional[bool] = None, |
| 1239 | ) -> Union[Tuple[torch.Tensor, ...], SequenceClassifierOutputWithPast]: |
| 1240 | return_dict = return_dict if return_dict is not None else self.config.use_return_dict |
| 1241 | |
| 1242 | transformer_outputs = self.transformer( |
| 1243 | input_ids=input_ids, |
| 1244 | position_ids=position_ids, |
| 1245 | attention_mask=attention_mask, |
| 1246 | full_attention_mask=full_attention_mask, |
| 1247 | past_key_values=past_key_values, |
| 1248 | inputs_embeds=inputs_embeds, |
| 1249 | use_cache=use_cache, |
| 1250 | output_hidden_states=output_hidden_states, |
| 1251 | return_dict=return_dict, |
| 1252 | ) |
| 1253 | |
| 1254 | hidden_states = transformer_outputs[0] |
| 1255 | pooled_hidden_states = hidden_states[-1] |
| 1256 | if self.dropout is not None: |
| 1257 | pooled_hidden_states = self.dropout(pooled_hidden_states) |
| 1258 | logits = self.classifier_head(pooled_hidden_states) |
| 1259 | |
| 1260 | loss = None |
| 1261 | if labels is not None: |
| 1262 | labels.to(logits) |
| 1263 | if self.config.problem_type is None: |
| 1264 | if self.num_labels == 1: |
| 1265 | self.config.problem_type = "regression" |
| 1266 | elif self.num_labels > 1 and (labels.dtype == torch.long or labels.dtype == torch.int): |
| 1267 | self.config.problem_type = "single_label_classification" |
nothing calls this directly
no outgoing calls
no test coverage detected