"third_party/mio/.gitkeep" did not exist on "37c494a74a267c551c947640476fb7eb248ec950"
test_run_glue_with_pabee.py 1.42 KB
Newer Older
1
2
3
4
5
6
import argparse
import logging
import sys
from unittest.mock import patch

import run_glue_with_pabee
7
from transformers.testing_utils import TestCasePlus, require_torch_non_multi_gpu_but_fix_me
8
9
10
11
12
13
14
15
16
17
18
19
20
21


logging.basicConfig(level=logging.DEBUG)

logger = logging.getLogger()


def get_setup_file():
    parser = argparse.ArgumentParser()
    parser.add_argument("-f")
    args = parser.parse_args()
    return args.f


22
class PabeeTests(TestCasePlus):
23
    @require_torch_non_multi_gpu_but_fix_me
24
25
26
27
    def test_run_glue(self):
        stream_handler = logging.StreamHandler(sys.stdout)
        logger.addHandler(stream_handler)

28
29
        tmp_dir = self.get_auto_remove_tmp_dir()
        testargs = f"""
30
31
32
33
            run_glue_with_pabee.py
            --model_type albert
            --model_name_or_path albert-base-v2
            --data_dir ./tests/fixtures/tests_samples/MRPC/
34
35
            --output_dir {tmp_dir}
            --overwrite_output_dir
36
37
38
39
40
41
42
43
44
45
            --task_name mrpc
            --do_train
            --do_eval
            --per_gpu_train_batch_size=2
            --per_gpu_eval_batch_size=1
            --learning_rate=2e-5
            --max_steps=50
            --warmup_steps=2
            --seed=42
            --max_seq_length=128
46
47
            """.split()

48
49
50
51
        with patch.object(sys, "argv", testargs):
            result = run_glue_with_pabee.main()
            for value in result.values():
                self.assertGreaterEqual(value, 0.75)