test_examples.py 7.49 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

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:
Sylvain Gugger's avatar
Sylvain Gugger committed
37
    import run_clm
38
39
40
    import run_generation
    import run_glue
    import run_language_modeling
41
    import run_pl_glue
42
    import run_squad
Aymeric Augustin's avatar
Aymeric Augustin committed
43

44

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

logger = logging.getLogger()
48

49

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


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


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

67
68
        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
69
            run_glue.py
70
            --model_name_or_path distilbert-base-uncased
71
72
            --output_dir {tmp_dir}
            --overwrite_output_dir
Sylvain Gugger's avatar
Sylvain Gugger committed
73
74
            --train_file ./tests/fixtures/tests_samples/MRPC/train.csv
            --validation_file ./tests/fixtures/tests_samples/MRPC/dev.csv
75
76
            --do_train
            --do_eval
77
78
            --per_device_train_batch_size=2
            --per_device_eval_batch_size=1
79
80
81
82
83
            --learning_rate=1e-4
            --max_steps=10
            --warmup_steps=2
            --seed=42
            --max_seq_length=128
84
            """.split()
85

86
        if is_cuda_and_apex_available():
87
            testargs.append("--fp16")
88

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

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

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

        with patch.object(sys, "argv", testargs):
120
121
            result = run_pl_glue.main()[0]
            # for now just testing that the script can run to completion
Stas Bekman's avatar
Stas Bekman committed
122
123
124
125
126
127
128
129
130
131
            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})")
            #

Sylvain Gugger's avatar
Sylvain Gugger committed
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
    def test_run_clm(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
            run_clm.py
            --model_name_or_path distilgpt2
            --train_file ./tests/fixtures/sample_text.txt
            --validation_file ./tests/fixtures/sample_text.txt
            --do_train
            --do_eval
            --block_size 128
            --per_device_train_batch_size 5
            --per_device_eval_batch_size 5
            --num_train_epochs 2
            --output_dir {tmp_dir}
            --overwrite_output_dir
            --prediction_loss_only
            """.split()

        if torch.cuda.device_count() > 1:
            # Skipping because there are not enough batches to train the model + would need a drop_last to work.
            return

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

        with patch.object(sys, "argv", testargs):
            result = run_clm.main()
            self.assertLess(result["perplexity"], 100)

Julien Chaumond's avatar
Julien Chaumond committed
164
165
166
167
    def test_run_language_modeling(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

168
169
        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
Julien Chaumond's avatar
Julien Chaumond committed
170
171
172
173
174
175
176
            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
177
            --output_dir {tmp_dir}
Julien Chaumond's avatar
Julien Chaumond committed
178
179
180
181
            --overwrite_output_dir
            --do_train
            --do_eval
            --num_train_epochs=1
182
            """.split()
183
184
185

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

Julien Chaumond's avatar
Julien Chaumond committed
187
188
        with patch.object(sys, "argv", testargs):
            result = run_language_modeling.main()
189
            self.assertLess(result["perplexity"], 42)
Julien Chaumond's avatar
Julien Chaumond committed
190

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

195
196
        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
197
            run_squad.py
198
199
            --model_type=distilbert
            --model_name_or_path=sshleifer/tiny-distilbert-base-cased-distilled-squad
200
            --data_dir=./tests/fixtures/tests_samples/SQUAD
201
202
            --output_dir {tmp_dir}
            --overwrite_output_dir
203
204
205
206
207
208
209
210
211
            --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
212
213
        """.split()

214
        with patch.object(sys, "argv", testargs):
215
            result = run_squad.main()
216
217
            self.assertGreaterEqual(result["f1"], 25)
            self.assertGreaterEqual(result["exact"], 21)
218

219
220
221
222
    def test_generation(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

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

225
        if is_cuda_and_apex_available():
226
227
228
229
230
231
            testargs.append("--fp16")

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