LongformerForTokenClassification (#4638)

This commit is contained in:
Suraj Patil
2020-05-28 16:18:18 +05:30
committed by GitHub
parent 3cc2c2a150
commit e444648a30
4 changed files with 122 additions and 0 deletions

View File

@@ -30,6 +30,7 @@ if is_torch_available():
LongformerModel,
LongformerForMaskedLM,
LongformerForSequenceClassification,
LongformerForTokenClassification,
LongformerForQuestionAnswering,
)
@@ -212,6 +213,21 @@ class LongformerModelTester(object):
self.parent.assertListEqual(list(result["logits"].size()), [self.batch_size, self.num_labels])
self.check_loss_output(result)
def create_and_check_longformer_for_token_classification(
self, config, input_ids, token_type_ids, input_mask, sequence_labels, token_labels, choice_labels
):
config.num_labels = self.num_labels
model = LongformerForTokenClassification(config=config)
model.to(torch_device)
model.eval()
loss, logits = model(input_ids, attention_mask=input_mask, token_type_ids=token_type_ids, labels=token_labels)
result = {
"loss": loss,
"logits": logits,
}
self.parent.assertListEqual(list(result["logits"].size()), [self.batch_size, self.seq_length, self.num_labels])
self.check_loss_output(result)
def prepare_config_and_inputs_for_common(self):
config_and_inputs = self.prepare_config_and_inputs()
(
@@ -278,6 +294,10 @@ class LongformerModelTest(ModelTesterMixin, unittest.TestCase):
config_and_inputs = self.model_tester.prepare_config_and_inputs()
self.model_tester.create_and_check_longformer_for_sequence_classification(*config_and_inputs)
def test_for_token_classification(self):
config_and_inputs = self.model_tester.prepare_config_and_inputs()
self.model_tester.create_and_check_longformer_for_token_classification(*config_and_inputs)
class LongformerModelIntegrationTest(unittest.TestCase):
@slow