update: do not require model arg, add controlnet_save_canny, fix bugs, add tests
This commit is contained in:
parent
ffaf562a8c
commit
98d7ef5751
@ -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)
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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()
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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):
|
||||
|
||||
Loading…
Reference in New Issue
Block a user