|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
import os |
|
import shutil |
|
import unittest |
|
from unittest.mock import patch |
|
|
|
from transformers.testing_utils import CaptureStd, is_pt_tf_cross_test, require_torch |
|
|
|
|
|
class CLITest(unittest.TestCase): |
|
@patch("sys.argv", ["fakeprogrampath", "env"]) |
|
def test_cli_env(self): |
|
|
|
import transformers.commands.transformers_cli |
|
|
|
with CaptureStd() as cs: |
|
transformers.commands.transformers_cli.main() |
|
self.assertIn("Python version", cs.out) |
|
self.assertIn("Platform", cs.out) |
|
self.assertIn("Using distributed or parallel set-up in script?", cs.out) |
|
|
|
@is_pt_tf_cross_test |
|
@patch( |
|
"sys.argv", ["fakeprogrampath", "pt-to-tf", "--model-name", "hf-internal-testing/tiny-random-gptj", "--no-pr"] |
|
) |
|
def test_cli_pt_to_tf(self): |
|
import transformers.commands.transformers_cli |
|
|
|
shutil.rmtree("/tmp/hf-internal-testing/tiny-random-gptj", ignore_errors=True) |
|
transformers.commands.transformers_cli.main() |
|
|
|
self.assertTrue(os.path.exists("/tmp/hf-internal-testing/tiny-random-gptj/tf_model.h5")) |
|
|
|
@require_torch |
|
@patch("sys.argv", ["fakeprogrampath", "download", "hf-internal-testing/tiny-random-gptj", "--cache-dir", "/tmp"]) |
|
def test_cli_download(self): |
|
import transformers.commands.transformers_cli |
|
|
|
|
|
shutil.rmtree("/tmp/models--hf-internal-testing--tiny-random-gptj", ignore_errors=True) |
|
|
|
|
|
transformers.commands.transformers_cli.main() |
|
|
|
|
|
self.assertTrue(os.path.exists("/tmp/models--hf-internal-testing--tiny-random-gptj/blobs")) |
|
self.assertTrue(os.path.exists("/tmp/models--hf-internal-testing--tiny-random-gptj/refs")) |
|
self.assertTrue(os.path.exists("/tmp/models--hf-internal-testing--tiny-random-gptj/snapshots")) |
|
|
|
@require_torch |
|
@patch( |
|
"sys.argv", |
|
[ |
|
"fakeprogrampath", |
|
"download", |
|
"hf-internal-testing/test_dynamic_model_with_tokenizer", |
|
"--trust-remote-code", |
|
"--cache-dir", |
|
"/tmp", |
|
], |
|
) |
|
def test_cli_download_trust_remote(self): |
|
import transformers.commands.transformers_cli |
|
|
|
|
|
shutil.rmtree("/tmp/models--hf-internal-testing--test_dynamic_model_with_tokenizer", ignore_errors=True) |
|
|
|
|
|
transformers.commands.transformers_cli.main() |
|
|
|
|
|
self.assertTrue(os.path.exists("/tmp/models--hf-internal-testing--test_dynamic_model_with_tokenizer/blobs")) |
|
self.assertTrue(os.path.exists("/tmp/models--hf-internal-testing--test_dynamic_model_with_tokenizer/refs")) |
|
self.assertTrue( |
|
os.path.exists("/tmp/models--hf-internal-testing--test_dynamic_model_with_tokenizer/snapshots") |
|
) |
|
|