test_examples.py 4.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
21
import sys
import unittest
Aymeric Augustin's avatar
Aymeric Augustin committed
22
from unittest.mock import patch
Aymeric Augustin's avatar
Aymeric Augustin committed
23

24
25
26
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:
    import run_generation
    import run_glue
    import run_language_modeling
    import run_squad
Aymeric Augustin's avatar
Aymeric Augustin committed
37

38

39
40
41
logging.basicConfig(level=logging.DEBUG)

logger = logging.getLogger()
42

43

44
45
def get_setup_file():
    parser = argparse.ArgumentParser()
46
    parser.add_argument("-f")
47
48
49
50
    args = parser.parse_args()
    return args.f


51
class ExamplesTests(unittest.TestCase):
52
53
54
55
    def test_run_glue(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

56
57
        testargs = """
            run_glue.py
58
            --model_name_or_path distilbert-base-uncased
59
60
61
62
63
            --data_dir ./tests/fixtures/tests_samples/MRPC/
            --task_name mrpc
            --do_train
            --do_eval
            --output_dir ./tests/fixtures/tests_samples/temp_dir
64
65
            --per_device_train_batch_size=2
            --per_device_eval_batch_size=1
66
67
68
69
70
71
72
73
            --learning_rate=1e-4
            --max_steps=10
            --warmup_steps=2
            --overwrite_output_dir
            --seed=42
            --max_seq_length=128
            """.split()
        with patch.object(sys, "argv", testargs):
74
            result = run_glue.main()
75
            del result["eval_loss"]
76
77
            for value in result.values():
                self.assertGreaterEqual(value, 0.75)
78

Julien Chaumond's avatar
Julien Chaumond committed
79
80
81
    def test_run_language_modeling(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)
82
        # TODO: switch to smaller model like sshleifer/tiny-distilroberta-base
Julien Chaumond's avatar
Julien Chaumond committed
83
84
85
86
87
88
89
90
91

        testargs = """
            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
92
            --output_dir ./tests/fixtures/tests_samples/temp_dir
Julien Chaumond's avatar
Julien Chaumond committed
93
94
95
96
97
98
99
100
101
102
            --overwrite_output_dir
            --do_train
            --do_eval
            --num_train_epochs=1
            --no_cuda
            """.split()
        with patch.object(sys, "argv", testargs):
            result = run_language_modeling.main()
            self.assertLess(result["perplexity"], 35)

103
104
105
106
    def test_run_squad(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

107
108
        testargs = """
            run_squad.py
109
110
            --model_type=distilbert
            --model_name_or_path=sshleifer/tiny-distilbert-base-cased-distilled-squad
111
112
113
114
115
116
117
118
119
120
121
122
123
124
            --data_dir=./tests/fixtures/tests_samples/SQUAD
            --output_dir=./tests/fixtures/tests_samples/temp_dir
            --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
            --overwrite_output_dir
            --seed=42
        """.split()
        with patch.object(sys, "argv", testargs):
125
            result = run_squad.main()
126
127
            self.assertGreaterEqual(result["f1"], 25)
            self.assertGreaterEqual(result["exact"], 21)
128

129
130
131
132
    def test_generation(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

133
        testargs = ["run_generation.py", "--prompt=Hello", "--length=10", "--seed=42"]
134
        model_type, model_name = ("--model_type=gpt2", "--model_name_or_path=sshleifer/tiny-gpt2")
135
        with patch.object(sys, "argv", testargs + [model_type, model_name]):
136
            result = run_generation.main()
137
            self.assertGreaterEqual(len(result[0]), 10)