MCPcopy Create free account
hub / github.com/FreedomIntelligence/CMB / ChatGLMForSequenceClassification

Class ChatGLMForSequenceClassification

workers/chatglm3_modeling.py:1210–1294  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1208
1209
1210class 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"

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected