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()