"vscode:/vscode.git/clone" did not exist on "6a5cde6ae248ef4466b90692dac4913611dbcba1"
test_modeling_flax_auto.py 4.16 KB
Newer Older
Sylvain Gugger's avatar
Sylvain Gugger committed
1
2
3
4
5
6
7
8
9
10
11
12
13
14
# Copyright 2020 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

15
16
17
import unittest

from transformers import AutoConfig, AutoTokenizer, BertConfig, TensorType, is_flax_available
18
from transformers.testing_utils import DUMMY_UNKNOWN_IDENTIFIER, require_flax, slow
19
20
21
22


if is_flax_available():
    import jax
23

Sylvain Gugger's avatar
Sylvain Gugger committed
24
25
26
    from transformers.models.auto.modeling_flax_auto import FlaxAutoModel
    from transformers.models.bert.modeling_flax_bert import FlaxBertModel
    from transformers.models.roberta.modeling_flax_roberta import FlaxRobertaModel
27
28
29
30
31
32


@require_flax
class FlaxAutoModelTest(unittest.TestCase):
    @slow
    def test_bert_from_pretrained(self):
33
        for model_name in ["google-bert/bert-base-cased", "google-bert/bert-large-uncased"]:
34
35
36
37
38
39
40
41
42
43
44
            with self.subTest(model_name):
                config = AutoConfig.from_pretrained(model_name)
                self.assertIsNotNone(config)
                self.assertIsInstance(config, BertConfig)

                model = FlaxAutoModel.from_pretrained(model_name)
                self.assertIsNotNone(model)
                self.assertIsInstance(model, FlaxBertModel)

    @slow
    def test_roberta_from_pretrained(self):
45
        for model_name in ["FacebookAI/roberta-base", "FacebookAI/roberta-large"]:
46
47
48
49
50
51
52
53
54
55
56
            with self.subTest(model_name):
                config = AutoConfig.from_pretrained(model_name)
                self.assertIsNotNone(config)
                self.assertIsInstance(config, BertConfig)

                model = FlaxAutoModel.from_pretrained(model_name)
                self.assertIsNotNone(model)
                self.assertIsInstance(model, FlaxRobertaModel)

    @slow
    def test_bert_jax_jit(self):
57
        for model_name in ["google-bert/bert-base-cased", "google-bert/bert-large-uncased"]:
58
59
60
61
62
63
64
65
66
67
68
69
            tokenizer = AutoTokenizer.from_pretrained(model_name)
            model = FlaxBertModel.from_pretrained(model_name)
            tokens = tokenizer("Do you support jax jitted function?", return_tensors=TensorType.JAX)

            @jax.jit
            def eval(**kwargs):
                return model(**kwargs)

            eval(**tokens).block_until_ready()

    @slow
    def test_roberta_jax_jit(self):
70
        for model_name in ["FacebookAI/roberta-base", "FacebookAI/roberta-large"]:
71
72
73
74
75
76
77
78
79
            tokenizer = AutoTokenizer.from_pretrained(model_name)
            model = FlaxRobertaModel.from_pretrained(model_name)
            tokens = tokenizer("Do you support jax jitted function?", return_tensors=TensorType.JAX)

            @jax.jit
            def eval(**kwargs):
                return model(**kwargs)

            eval(**tokens).block_until_ready()
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102

    def test_repo_not_found(self):
        with self.assertRaisesRegex(
            EnvironmentError, "bert-base is not a local folder and is not a valid model identifier"
        ):
            _ = FlaxAutoModel.from_pretrained("bert-base")

    def test_revision_not_found(self):
        with self.assertRaisesRegex(
            EnvironmentError, r"aaaaaa is not a valid git identifier \(branch name, tag name or commit id\)"
        ):
            _ = FlaxAutoModel.from_pretrained(DUMMY_UNKNOWN_IDENTIFIER, revision="aaaaaa")

    def test_model_file_not_found(self):
        with self.assertRaisesRegex(
            EnvironmentError,
            "hf-internal-testing/config-no-model does not appear to have a file named flax_model.msgpack",
        ):
            _ = FlaxAutoModel.from_pretrained("hf-internal-testing/config-no-model")

    def test_model_from_pt_suggestion(self):
        with self.assertRaisesRegex(EnvironmentError, "Use `from_pt=True` to load this model"):
            _ = FlaxAutoModel.from_pretrained("hf-internal-testing/tiny-bert-pt-only")