update: do not require model arg, add controlnet_save_canny, fix bugs, add tests

This commit is contained in:
Anthony Wu 2024-10-26 20:22:32 -07:00
parent ffaf562a8c
commit 98d7ef5751
5 changed files with 83 additions and 20 deletions

View File

@ -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)

View File

@ -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()

View File

@ -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()

View File

@ -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

View 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):