| 50 | |
| 51 | class ExamplesTests(unittest.TestCase): |
| 52 | def test_run_glue(self): |
| 53 | stream_handler = logging.StreamHandler(sys.stdout) |
| 54 | logger.addHandler(stream_handler) |
| 55 | |
| 56 | testargs = """ |
| 57 | run_glue.py |
| 58 | --model_name_or_path distilbert-base-uncased |
| 59 | --data_dir ./tests/fixtures/tests_samples/MRPC/ |
| 60 | --task_name mrpc |
| 61 | --do_train |
| 62 | --do_eval |
| 63 | --output_dir ./tests/fixtures/tests_samples/temp_dir |
| 64 | --per_device_train_batch_size=2 |
| 65 | --per_device_eval_batch_size=1 |
| 66 | --learning_rate=1e-4 |
| 67 | --max_steps=10 |
| 68 | --warmup_steps=2 |
| 69 | --overwrite_output_dir |
| 70 | --seed=42 |
| 71 | --max_seq_length=128 |
| 72 | """.split() |
| 73 | with patch.object(sys, "argv", testargs): |
| 74 | result = run_glue.main() |
| 75 | del result["eval_loss"] |
| 76 | for value in result.values(): |
| 77 | self.assertGreaterEqual(value, 0.75) |
| 78 | |
| 79 | def test_run_language_modeling(self): |
| 80 | stream_handler = logging.StreamHandler(sys.stdout) |