Fix a sneaky reference to compute_loss in the tests

This commit is contained in:
matt
2022-01-18 17:33:38 +00:00
parent 979ca24e39
commit 2085f20901

View File

@@ -1064,7 +1064,7 @@ class TFModelTesterMixin:
config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
for model_class in self.all_model_classes:
model = model_class(config)
if getattr(model, "compute_loss", None):
if getattr(model, "hf_compute_loss", None):
# The number of elements in the loss should be the same as the number of elements in the label
prepared_for_class = self._prepare_for_class(inputs_dict.copy(), model_class, return_labels=True)
added_label = prepared_for_class[