85 lines
2.7 KiB
Python
85 lines
2.7 KiB
Python
import json
|
|
import sys
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
import torch
|
|
from transformers import BertConfig, BertForSequenceClassification
|
|
|
|
|
|
MODULE_DIR = Path(__file__).resolve().parents[1]
|
|
sys.path.insert(0, str(MODULE_DIR))
|
|
|
|
import convert_coreml
|
|
from purpose_data import LABELS, DataError
|
|
|
|
|
|
class FixedShapeBertForCoreMLTests(unittest.TestCase):
|
|
def test_conversion_forward_matches_transformers(self):
|
|
torch.manual_seed(7)
|
|
config = BertConfig(
|
|
vocab_size=64,
|
|
hidden_size=16,
|
|
num_hidden_layers=1,
|
|
num_attention_heads=4,
|
|
intermediate_size=32,
|
|
max_position_embeddings=128,
|
|
type_vocab_size=2,
|
|
hidden_dropout_prob=0.0,
|
|
attention_probs_dropout_prob=0.0,
|
|
num_labels=len(LABELS),
|
|
)
|
|
model = BertForSequenceClassification(config).eval()
|
|
wrapper = convert_coreml.FixedShapeBertForCoreML(model).eval()
|
|
input_ids = torch.randint(0, config.vocab_size, (1, 128), dtype=torch.int32)
|
|
attention_mask = torch.zeros((1, 128), dtype=torch.int32)
|
|
attention_mask[:, :83] = 1
|
|
token_type_ids = torch.zeros((1, 128), dtype=torch.int32)
|
|
token_type_ids[:, 43:83] = 1
|
|
with torch.inference_mode():
|
|
reference = model(
|
|
input_ids=input_ids.long(),
|
|
attention_mask=attention_mask.long(),
|
|
token_type_ids=token_type_ids.long(),
|
|
).logits
|
|
candidate = wrapper(input_ids, attention_mask, token_type_ids)
|
|
torch.testing.assert_close(candidate, reference, rtol=1e-5, atol=2e-5)
|
|
|
|
traced = torch.jit.trace(
|
|
wrapper,
|
|
(input_ids, attention_mask, token_type_ids),
|
|
strict=True,
|
|
)
|
|
torch.testing.assert_close(
|
|
traced(input_ids, attention_mask, token_type_ids),
|
|
reference,
|
|
rtol=1e-5,
|
|
atol=2e-5,
|
|
)
|
|
|
|
|
|
class CheckpointConfigTests(unittest.TestCase):
|
|
def test_rejects_changed_label_order(self):
|
|
config = {
|
|
"model_type": "bert",
|
|
"hidden_size": 384,
|
|
"num_hidden_layers": 6,
|
|
"id2label": {
|
|
str(index): label
|
|
for index, label in enumerate(reversed(LABELS))
|
|
},
|
|
}
|
|
with tempfile.TemporaryDirectory() as temp:
|
|
model_dir = Path(temp)
|
|
(model_dir / "config.json").write_text(
|
|
json.dumps(config),
|
|
encoding="utf-8",
|
|
)
|
|
with self.assertRaisesRegex(DataError, "label order"):
|
|
convert_coreml._checkpoint_config(model_dir)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|