test_examples.py 6.66 KB
Newer Older
1
2
3
4
5
6
7
8
9
10
11
12
13
14
# coding=utf-8
# Copyright 2018 HuggingFace Inc..
#
# 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.
Aymeric Augustin's avatar
Aymeric Augustin committed
15

16
17

import argparse
18
import logging
19
import os
Aymeric Augustin's avatar
Aymeric Augustin committed
20
import sys
Aymeric Augustin's avatar
Aymeric Augustin committed
21
from unittest.mock import patch
Aymeric Augustin's avatar
Aymeric Augustin committed
22

Stas Bekman's avatar
Stas Bekman committed
23
24
import torch

25
26
from transformers.file_utils import is_apex_available
from transformers.testing_utils import TestCasePlus, torch_device
27

28
29
30
31
32
33
34
35
36
37
38
39

SRC_DIRS = [
    os.path.join(os.path.dirname(__file__), dirname)
    for dirname in ["text-generation", "text-classification", "language-modeling", "question-answering"]
]
sys.path.extend(SRC_DIRS)


if SRC_DIRS is not None:
    import run_generation
    import run_glue
    import run_language_modeling
40
    import run_pl_glue
41
    import run_squad
Aymeric Augustin's avatar
Aymeric Augustin committed
42

43

44
45
46
logging.basicConfig(level=logging.DEBUG)

logger = logging.getLogger()
47

48

49
50
def get_setup_file():
    parser = argparse.ArgumentParser()
51
    parser.add_argument("-f")
52
53
54
55
    args = parser.parse_args()
    return args.f


56
def is_cuda_and_apex_available():
57
58
59
60
    is_using_cuda = torch.cuda.is_available() and torch_device == "cuda"
    return is_using_cuda and is_apex_available()


61
class ExamplesTests(TestCasePlus):
62
63
64
65
    def test_run_glue(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

66
67
        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
68
            run_glue.py
69
            --model_name_or_path distilbert-base-uncased
70
            --data_dir ./tests/fixtures/tests_samples/MRPC/
71
72
            --output_dir {tmp_dir}
            --overwrite_output_dir
73
74
75
            --task_name mrpc
            --do_train
            --do_eval
76
77
            --per_device_train_batch_size=2
            --per_device_eval_batch_size=1
78
79
80
81
82
            --learning_rate=1e-4
            --max_steps=10
            --warmup_steps=2
            --seed=42
            --max_seq_length=128
83
84
85
86
87
            """
        output_dir = "./tests/fixtures/tests_samples/temp_dir_{}".format(hash(testargs))
        testargs += "--output_dir " + output_dir
        testargs = testargs.split()

88
        if is_cuda_and_apex_available():
89
            testargs.append("--fp16")
90

91
        with patch.object(sys, "argv", testargs):
92
            result = run_glue.main()
93
            del result["eval_loss"]
94
95
            for value in result.values():
                self.assertGreaterEqual(value, 0.75)
96

Stas Bekman's avatar
Stas Bekman committed
97
98
99
100
    def test_run_pl_glue(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

101
102
        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
Stas Bekman's avatar
Stas Bekman committed
103
104
105
            run_pl_glue.py
            --model_name_or_path bert-base-cased
            --data_dir ./tests/fixtures/tests_samples/MRPC/
106
            --output_dir {tmp_dir}
Stas Bekman's avatar
Stas Bekman committed
107
108
109
110
111
112
113
114
            --task mrpc
            --do_train
            --do_predict
            --train_batch_size=32
            --learning_rate=1e-4
            --num_train_epochs=1
            --seed=42
            --max_seq_length=128
115
            """.split()
Stas Bekman's avatar
Stas Bekman committed
116
        if torch.cuda.is_available():
117
            testargs += ["--gpus=1"]
118
        if is_cuda_and_apex_available():
119
            testargs.append("--fp16")
Stas Bekman's avatar
Stas Bekman committed
120
121
122
123
124
125
126
127
128
129
130
131
132
133

        with patch.object(sys, "argv", testargs):
            result = run_pl_glue.main()
            # for now just testing that the script can run to a completion
            self.assertGreater(result["acc"], 0.25)
            #
            # TODO: this fails on CI - doesn't get acc/f1>=0.75:
            #
            #     # remove all the various *loss* attributes
            #     result = {k: v for k, v in result.items() if "loss" not in k}
            #     for k, v in result.items():
            #         self.assertGreaterEqual(v, 0.75, f"({k})")
            #

Julien Chaumond's avatar
Julien Chaumond committed
134
135
136
137
    def test_run_language_modeling(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

138
139
        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
Julien Chaumond's avatar
Julien Chaumond committed
140
141
142
143
144
145
146
            run_language_modeling.py
            --model_name_or_path distilroberta-base
            --model_type roberta
            --mlm
            --line_by_line
            --train_data_file ./tests/fixtures/sample_text.txt
            --eval_data_file ./tests/fixtures/sample_text.txt
147
            --output_dir {tmp_dir}
Julien Chaumond's avatar
Julien Chaumond committed
148
149
150
151
            --overwrite_output_dir
            --do_train
            --do_eval
            --num_train_epochs=1
152
153
154
155
156
157
158
            """
        output_dir = "./tests/fixtures/tests_samples/temp_dir_{}".format(hash(testargs))
        testargs += "--output_dir " + output_dir
        testargs = testargs.split()

        if torch_device != "cuda":
            testargs.append("--no_cuda")
159

Julien Chaumond's avatar
Julien Chaumond committed
160
161
162
163
        with patch.object(sys, "argv", testargs):
            result = run_language_modeling.main()
            self.assertLess(result["perplexity"], 35)

164
165
166
167
    def test_run_squad(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

168
169
        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
170
            run_squad.py
171
172
            --model_type=distilbert
            --model_name_or_path=sshleifer/tiny-distilbert-base-cased-distilled-squad
173
            --data_dir=./tests/fixtures/tests_samples/SQUAD
174
175
            --output_dir {tmp_dir}
            --overwrite_output_dir
176
177
178
179
180
181
182
183
184
            --max_steps=10
            --warmup_steps=2
            --do_train
            --do_eval
            --version_2_with_negative
            --learning_rate=2e-4
            --per_gpu_train_batch_size=2
            --per_gpu_eval_batch_size=1
            --seed=42
185
186
        """.split()

187
        with patch.object(sys, "argv", testargs):
188
            result = run_squad.main()
189
190
            self.assertGreaterEqual(result["f1"], 25)
            self.assertGreaterEqual(result["exact"], 21)
191

192
193
194
195
    def test_generation(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

196
        testargs = ["run_generation.py", "--prompt=Hello", "--length=10", "--seed=42"]
197

198
        if is_cuda_and_apex_available():
199
200
201
202
203
204
            testargs.append("--fp16")

        model_type, model_name = (
            "--model_type=gpt2",
            "--model_name_or_path=sshleifer/tiny-gpt2",
        )
205
        with patch.object(sys, "argv", testargs + [model_type, model_name]):
206
            result = run_generation.main()
207
            self.assertGreaterEqual(len(result[0]), 10)