Compare commits

...
Author SHA1 Message Date
Morgan Funtowicz f545956ff5 Use the recommended nonzero(..., as_tuple=False) overload.
Signed-off-by: Morgan Funtowicz <morgan@huggingface.co>
2020-11-16 16:29:29 +01:00
Morgan Funtowicz f85f993442 Do not use .T and prefer .t() to be able to create a record in the exported graph.
Signed-off-by: Morgan Funtowicz <morgan@huggingface.co>
2020-11-16 16:17:26 +01:00
Morgan Funtowicz 18b9a8ebf8 Do not use item() as torch.unique might return a vector.
Signed-off-by: Morgan Funtowicz <morgan@huggingface.co>
2020-11-16 15:54:13 +01:00
Morgan Funtowicz 2dc31b8f96 Ensure output shape match expected shape
Signed-off-by: Morgan Funtowicz <morgan@huggingface.co>
2020-11-16 15:50:54 +01:00
Morgan Funtowicz 642a139696 Attempt to gather the latest eos_token representation in an ONNX compatible way
Signed-off-by: Morgan Funtowicz <morgan@huggingface.co>
2020-11-16 15:28:22 +01:00
Morgan Funtowicz adf58572bb Avoid using python len()
Signed-off-by: Morgan Funtowicz <morgan@huggingface.co>
2020-11-16 11:47:48 +01:00
+6 -3
View File
@@ -1199,10 +1199,13 @@ class BartForSequenceClassification(PretrainedBartModel):
)
x = outputs[0] # last hidden state
eos_mask = input_ids.eq(self.config.eos_token_id)
if len(torch.unique(eos_mask.sum(1))) > 1:
if torch.unique(eos_mask.sum(1)).size(-1) > 1:
raise ValueError("All examples must have the same number of <eos> tokens.")
sentence_representation = x[eos_mask, :].view(x.size(0), -1, x.size(-1))[:, -1, :]
logits = self.classification_head(sentence_representation)
# Attempt to gather the latest eos_token representation in an ONNX compatible way
# (-1: remove the batch indexes from the vector, -1: Take the last occurrence of eos_token)
sentence_representation = torch.index_select(x, 1, eos_mask.nonzero(as_tuple=False).t()[-1, -1])
logits = self.classification_head(sentence_representation).squeeze(1)
loss = None
if labels is not None: