diff --git a/src/mflux/generate.py b/src/mflux/generate.py index 3334d1d..236b6b0 100644 --- a/src/mflux/generate.py +++ b/src/mflux/generate.py @@ -8,7 +8,7 @@ from mflux.ui.cli.parsers import CommandLineParser def main(): # fmt: off parser = CommandLineParser(description="Generate an image based on a prompt.") - parser.add_model_arguments() + parser.add_model_arguments(require_model_arg=False) parser.add_lora_arguments() parser.add_image_generator_arguments(supports_metadata_config=True) parser.add_image_to_image_arguments(required=False) diff --git a/src/mflux/generate_controlnet.py b/src/mflux/generate_controlnet.py index 122ed04..acd7443 100644 --- a/src/mflux/generate_controlnet.py +++ b/src/mflux/generate_controlnet.py @@ -7,9 +7,9 @@ from mflux.ui.cli.parsers import CommandLineParser def main(): parser = CommandLineParser(description="Generate an image based on a prompt and a controlnet reference image.") # fmt: off - parser.add_model_arguments() + parser.add_model_arguments(require_model_arg=True) parser.add_lora_arguments() - parser.add_image_generator_arguments(supports_metadata_config=True) + parser.add_image_generator_arguments(supports_metadata_config=False) parser.add_controlnet_arguments() parser.add_output_arguments() args = parser.parse_args() diff --git a/src/mflux/save.py b/src/mflux/save.py index db473ac..21c0b57 100644 --- a/src/mflux/save.py +++ b/src/mflux/save.py @@ -4,7 +4,7 @@ from mflux.ui.cli.parsers import CommandLineParser def main(): parser = CommandLineParser(description="Save a quantized version of Flux.1 to disk.") # fmt: off - parser.add_model_arguments(path_type="save") + parser.add_model_arguments(path_type="save", require_model_arg=True) parser.add_lora_arguments() args = parser.parse_args() diff --git a/src/mflux/ui/cli/parsers.py b/src/mflux/ui/cli/parsers.py index 3dad873..c6da341 100644 --- a/src/mflux/ui/cli/parsers.py +++ b/src/mflux/ui/cli/parsers.py @@ -17,9 +17,9 @@ class CommandLineParser(argparse.ArgumentParser): self.supports_image_to_image = False self.supports_lora = False - def add_model_arguments(self, path_type: t.Literal["load", "save"] = "load") -> None: + def add_model_arguments(self, path_type: t.Literal["load", "save"] = "load", require_model_arg: bool = True) -> None: - self.add_argument("--model", "-m", type=str, required=True, choices=ui_defaults.MODEL_CHOICES, help=f"The model to use ({' or '.join(ui_defaults.MODEL_CHOICES)}).") + self.add_argument("--model", "-m", type=str, required=require_model_arg, choices=ui_defaults.MODEL_CHOICES, help=f"The model to use ({' or '.join(ui_defaults.MODEL_CHOICES)}).") if path_type == "load": self.add_argument("--path", type=str, default=None, help="Local path for loading a model from disk") @@ -74,11 +74,15 @@ class CommandLineParser(argparse.ArgumentParser): def parse_args(self, **kwargs) -> argparse.Namespace: namespace = super().parse_args() if hasattr(namespace, "path") and namespace.path is not None and namespace.model is None: - namespace.error("--model must be specified when using --path") + self.error("--model must be specified when using --path") if getattr(namespace, "config_from_metadata", False): prior_gen_metadata = json.load(namespace.config_from_metadata.open("rt")) + if namespace.model is None: + # when not provided by CLI flag, find it in the config file + namespace.model = prior_gen_metadata.get("model", None) + if namespace.prompt is None: namespace.prompt = prior_gen_metadata.get("prompt", None) @@ -118,8 +122,11 @@ class CommandLineParser(argparse.ArgumentParser): namespace.controlnet_image_path = prior_gen_metadata.get("controlnet_image_path", None) if namespace.controlnet_strength == self.get_default("controlnet_strength") and (cnet_strength_from_metadata := prior_gen_metadata.get("controlnet_strength", None)): namespace.controlnet_strength = cnet_strength_from_metadata + if namespace.controlnet_save_canny == self.get_default("controlnet_save_canny") and (cnet_canny_from_metadata := prior_gen_metadata.get("controlnet_save_canny", None)): + namespace.controlnet_save_canny = cnet_canny_from_metadata - + if namespace.model is None: + self.error("--model / -m must be provided, or 'model' must be specified in the config file.") if self.supports_image_generation and namespace.prompt is None: # not supplied by CLI and not supplied by metadata config file diff --git a/tests/test_cli_argparser.py b/tests/test_cli_argparser.py index d7a5698..0a043e1 100644 --- a/tests/test_cli_argparser.py +++ b/tests/test_cli_argparser.py @@ -9,7 +9,7 @@ from mflux.ui.cli.parsers import CommandLineParser def _create_mflux_generate_parser(with_controlnet=False) -> CommandLineParser: parser = CommandLineParser(description="Generate an image based on a prompt.") - parser.add_model_arguments() + parser.add_model_arguments(require_model_arg=False) parser.add_image_generator_arguments(supports_metadata_config=True) parser.add_lora_arguments() parser.add_image_to_image_arguments(required=False) @@ -32,19 +32,19 @@ def mflux_generate_controlnet_parser() -> CommandLineParser: @pytest.fixture def mflux_save_parser() -> CommandLineParser: parser = CommandLineParser(description="Save a quantized version of Flux.1 to disk.") # fmt: off - parser.add_model_arguments(path_type="save") + parser.add_model_arguments(path_type="save", require_model_arg=True) parser.add_lora_arguments() return parser @pytest.fixture def mflux_generate_minimal_argv() -> list[str]: - return ["mflux-generate", "--model", "schnell", "--prompt", "meaning of life"] + return ["mflux-generate", "--prompt", "meaning of life"] @pytest.fixture def mflux_generate_controlnet_minimal_argv() -> list[str]: - return ["mflux-generate-controlnet", "--model", "dev", "--prompt", "meaning of life, imitated"] + return ["mflux-generate-controlnet", "--prompt", "meaning of life, imitated"] @pytest.fixture @@ -58,9 +58,9 @@ def temp_dir(tmp_path_factory) -> Path: def base_metadata_dict() -> dict: return { "mflux_version": "0.4.0", - "model": "schnell", + "model": "dev", "seed": 42042, - "steps": 4, + "steps": 14, "guidance": None, "precision": "mlx.core.bfloat16", "quantize": None, @@ -71,6 +71,7 @@ def base_metadata_dict() -> dict: "init_image_strength": None, "controlnet_image": None, "controlnet_strength": None, + "controlnet_save_canny": False, } @@ -80,6 +81,39 @@ def test_model_path_requires_model_arg(mflux_generate_parser): assert pytest.raises(SystemExit, mflux_generate_parser.parse_args) +def test_model_arg_not_in_file(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir): + metadata_file = temp_dir / "model.json" + with metadata_file.open("wt") as m: + del base_metadata_dict["model"] + json.dump(base_metadata_dict, m, indent=4) + # test model arg not provided in either flag or file + with patch('sys.argv', mflux_generate_minimal_argv + ['--config-from-metadata', metadata_file.as_posix()]): # fmt: off + pytest.raises(SystemExit, mflux_generate_parser.parse_args) + # test value read from flag + with patch('sys.argv', mflux_generate_minimal_argv + ['--model', 'dev', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off + args = mflux_generate_parser.parse_args() + assert args.model == "dev" + # test value read from flag + with patch('sys.argv', mflux_generate_minimal_argv + ['--model', 'schnell', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off + args = mflux_generate_parser.parse_args() + assert args.model == "schnell" + + +def test_model_arg_in_file(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir): + metadata_file = temp_dir / "model.json" + with metadata_file.open("wt") as m: + base_metadata_dict["model"] = "dev" + json.dump(base_metadata_dict, m, indent=4) + # test value read from file + with patch('sys.argv', mflux_generate_minimal_argv + ['--config-from-metadata', metadata_file.as_posix()]): # fmt: off + args = mflux_generate_parser.parse_args() + assert args.model == "dev" + # test value read from flag, overrides value from file + with patch('sys.argv', mflux_generate_minimal_argv + ['--model', 'schnell', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off + args = mflux_generate_parser.parse_args() + assert args.model == "schnell" + + def test_prompt_arg(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir): metadata_file = temp_dir / "prompt.json" file_prompt = "origin of the universe" @@ -148,8 +182,13 @@ def test_steps_arg(mflux_generate_parser, mflux_generate_minimal_argv, base_meta base_metadata_dict["steps"] = 8 json.dump(base_metadata_dict, m, indent=4) - # test user default value - with patch("sys.argv", mflux_generate_minimal_argv): + # test user default value for dev + with patch("sys.argv", mflux_generate_minimal_argv + ["--model", "dev"]): + args = mflux_generate_parser.parse_args() + assert args.steps == 14 + + # test user default value for schnell + with patch("sys.argv", mflux_generate_minimal_argv + ["--model", "schnell"]): args = mflux_generate_parser.parse_args() assert args.steps == 4 @@ -173,7 +212,7 @@ def test_lora_args(mflux_generate_parser, mflux_generate_minimal_argv, base_meta json.dump(base_metadata_dict, m, indent=4) # test user default value - with patch("sys.argv", mflux_generate_minimal_argv): + with patch("sys.argv", mflux_generate_minimal_argv + ["-m", "schnell"]): args = mflux_generate_parser.parse_args() assert args.lora_paths is None assert args.lora_scales is None @@ -184,7 +223,7 @@ def test_lora_args(mflux_generate_parser, mflux_generate_minimal_argv, base_meta assert args.lora_paths == test_paths assert args.lora_scales == [pytest.approx(0.3), pytest.approx(0.7)] - # test CLI override + # test CLI override that merges CLI loras and config file loras new_loras = ["--lora-paths", "/some/lora/3.safetensors", "/some/lora/4.safetensors", "--lora-scales", "0.1", "0.9"] with patch('sys.argv', mflux_generate_minimal_argv + new_loras + ['--config-from-metadata', metadata_file.as_posix()]): # fmt: off args = mflux_generate_parser.parse_args() @@ -202,7 +241,7 @@ def test_image_to_image_args(mflux_generate_parser, mflux_generate_minimal_argv, json.dump(base_metadata_dict, m, indent=4) # test user default value - with patch("sys.argv", mflux_generate_minimal_argv): + with patch("sys.argv", mflux_generate_minimal_argv + ["-m", "dev"]): args = mflux_generate_parser.parse_args() assert args.init_image_path is None assert args.init_image_strength == 0.4 # default @@ -239,13 +278,30 @@ def test_controlnet_args(mflux_generate_controlnet_parser, mflux_generate_contro args = mflux_generate_controlnet_parser.parse_args() assert args.controlnet_image_path == test_path assert args.controlnet_strength == pytest.approx(0.48) + assert args.controlnet_save_canny is False # test CLI override - override_cnet = ["--controlnet-image-path", "/some/lora/2.safetensors", "--controlnet-strength", "0.85"] + override_cnet = [ + "--controlnet-image-path", + "/some/lora/2.safetensors", + "--controlnet-strength", + "0.85", + "--controlnet-save-canny", + ] with patch('sys.argv', mflux_generate_controlnet_minimal_argv + override_cnet + ['--config-from-metadata', metadata_file.as_posix()]): # fmt: off args = mflux_generate_controlnet_parser.parse_args() assert args.controlnet_image_path == "/some/lora/2.safetensors" assert args.controlnet_strength == pytest.approx(0.85) + assert args.controlnet_save_canny is True + + # test controlnet_save_canny is False when not specified + with metadata_file.open("wt") as m: + del base_metadata_dict["controlnet_save_canny"] + json.dump(base_metadata_dict, m, indent=4) + + with patch('sys.argv', mflux_generate_controlnet_minimal_argv + ['--config-from-metadata', metadata_file.as_posix()]): # fmt: off + args = mflux_generate_controlnet_parser.parse_args() + assert args.controlnet_save_canny is False def test_save_args(mflux_save_parser):