Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion swift/cli/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,7 +61,7 @@ def parse_yaml_args(argv):
for k, v in config.items():
config_argv.append(f'--{k}')
if isinstance(v, list):
config_argv += v
config_argv += [str(i) for i in v]
else:
if isinstance(v, dict):
v = json.dumps(v, ensure_ascii=False)
Expand Down
61 changes: 61 additions & 0 deletions tests/utils/test_cli_config_args.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
import os
import tempfile
import unittest

from swift.cli.main import parse_yaml_args


class TestYamlConfigArgv(unittest.TestCase):
"""`swift sft config.yaml` (and the JSON form) rewrites argv in place, and the launcher then
prints and execs it as one command line, so every value has to arrive as a string."""

def setUp(self):
self._dir = tempfile.TemporaryDirectory()
self.addCleanup(self._dir.cleanup)
# parse_yaml_args records the config path for the entry point that runs after it
self._saved_config = os.environ.pop('SWIFT_CONFIG_FILE', None)
self.addCleanup(self._restore_config)

def _restore_config(self):
os.environ.pop('SWIFT_CONFIG_FILE', None)
if self._saved_config is not None:
os.environ['SWIFT_CONFIG_FILE'] = self._saved_config

def parse(self, name, content):
path = os.path.join(self._dir.name, name)
with open(path, 'w', encoding='utf-8') as f:
f.write(content)
argv = [path]
parse_yaml_args(argv)
return argv

def assert_launchable(self, argv):
try:
' '.join(argv)
except TypeError as exc:
self.fail(f'the launcher cannot build a command line from {argv!r}: {exc}')

def test_float_list_is_passed_as_strings(self):
# e.g. `interleave_prob: [0.5, 0.5]`, declared Optional[List[float]] in DataArguments
argv = self.parse('sft.yaml', 'dataset:\n- a\n- b\ninterleave_prob: [0.5, 0.5]\n')
self.assert_launchable(argv)
self.assertEqual(argv[-3:], ['--interleave_prob', '0.5', '0.5'])

def test_int_list_is_passed_as_strings(self):
# e.g. `data_range: [0, 2]`, declared List[int] in SamplingArguments
argv = self.parse('sample.json', '{"num_samples": 10, "data_range": [0, 2]}')
self.assert_launchable(argv)
self.assertEqual(argv, ['--num_samples', '10', '--data_range', '0', '2'])

def test_string_list_is_unchanged(self):
argv = self.parse('datasets.yaml', 'dataset:\n- a#100\n- b#100\n')
self.assertEqual(argv, ['--dataset', 'a#100', 'b#100'])

def test_dict_value_is_still_serialized(self):
argv = self.parse('engine.yaml', 'model: m\nvllm_engine_kwargs:\n gpu_memory_utilization: 0.8\n')
self.assert_launchable(argv)
self.assertEqual(argv, ['--model', 'm', '--vllm_engine_kwargs', '{"gpu_memory_utilization": 0.8}'])


if __name__ == '__main__':
unittest.main()
Loading