test_examples.py 5.8 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.testing_utils import TestCasePlus

27
28
29
30
31
32
33
34
35
36
37

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
Stas Bekman's avatar
Stas Bekman committed
38
    import run_pl_glue
39
40
    import run_language_modeling
    import run_squad
Aymeric Augustin's avatar
Aymeric Augustin committed
41

42

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

logger = logging.getLogger()
46

47

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


55
class ExamplesTests(TestCasePlus):
56
57
58
59
    def test_run_glue(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

60
61
        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
62
            run_glue.py
63
            --model_name_or_path distilbert-base-uncased
64
            --data_dir ./tests/fixtures/tests_samples/MRPC/
65
66
            --output_dir {tmp_dir}
            --overwrite_output_dir
67
68
69
            --task_name mrpc
            --do_train
            --do_eval
70
71
            --per_device_train_batch_size=2
            --per_device_eval_batch_size=1
72
73
74
75
76
            --learning_rate=1e-4
            --max_steps=10
            --warmup_steps=2
            --seed=42
            --max_seq_length=128
77
78
            """.split()

79
        with patch.object(sys, "argv", testargs):
80
            result = run_glue.main()
81
            del result["eval_loss"]
82
83
            for value in result.values():
                self.assertGreaterEqual(value, 0.75)
84

Stas Bekman's avatar
Stas Bekman committed
85
86
87
88
    def test_run_pl_glue(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

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

        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
120
121
122
123
    def test_run_language_modeling(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

124
125
        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
Julien Chaumond's avatar
Julien Chaumond committed
126
127
128
129
130
131
132
            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
133
            --output_dir {tmp_dir}
Julien Chaumond's avatar
Julien Chaumond committed
134
135
136
137
138
            --overwrite_output_dir
            --do_train
            --do_eval
            --num_train_epochs=1
            --no_cuda
139
140
            """.split()

Julien Chaumond's avatar
Julien Chaumond committed
141
142
143
144
        with patch.object(sys, "argv", testargs):
            result = run_language_modeling.main()
            self.assertLess(result["perplexity"], 35)

145
146
147
148
    def test_run_squad(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

149
150
        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
151
            run_squad.py
152
153
            --model_type=distilbert
            --model_name_or_path=sshleifer/tiny-distilbert-base-cased-distilled-squad
154
            --data_dir=./tests/fixtures/tests_samples/SQUAD
155
156
            --output_dir {tmp_dir}
            --overwrite_output_dir
157
158
159
160
161
162
163
164
165
            --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
166
167
        """.split()

168
        with patch.object(sys, "argv", testargs):
169
            result = run_squad.main()
170
171
            self.assertGreaterEqual(result["f1"], 25)
            self.assertGreaterEqual(result["exact"], 21)
172

173
174
175
176
    def test_generation(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

177
        testargs = ["run_generation.py", "--prompt=Hello", "--length=10", "--seed=42"]
178
        model_type, model_name = ("--model_type=gpt2", "--model_name_or_path=sshleifer/tiny-gpt2")
179
        with patch.object(sys, "argv", testargs + [model_type, model_name]):
180
            result = run_generation.main()
181
            self.assertGreaterEqual(len(result[0]), 10)