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():
|
def main():
|
||||||
# fmt: off
|
# fmt: off
|
||||||
parser = CommandLineParser(description="Generate an image based on a prompt.")
|
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_lora_arguments()
|
||||||
parser.add_image_generator_arguments(supports_metadata_config=True)
|
parser.add_image_generator_arguments(supports_metadata_config=True)
|
||||||
parser.add_image_to_image_arguments(required=False)
|
parser.add_image_to_image_arguments(required=False)
|
||||||
|
|||||||
@ -7,9 +7,9 @@ from mflux.ui.cli.parsers import CommandLineParser
|
|||||||
|
|
||||||
def main():
|
def main():
|
||||||
parser = CommandLineParser(description="Generate an image based on a prompt and a controlnet reference image.") # fmt: off
|
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_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_controlnet_arguments()
|
||||||
parser.add_output_arguments()
|
parser.add_output_arguments()
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|||||||
@ -4,7 +4,7 @@ from mflux.ui.cli.parsers import CommandLineParser
|
|||||||
|
|
||||||
def main():
|
def main():
|
||||||
parser = CommandLineParser(description="Save a quantized version of Flux.1 to disk.") # fmt: off
|
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()
|
parser.add_lora_arguments()
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
|||||||
@ -17,9 +17,9 @@ class CommandLineParser(argparse.ArgumentParser):
|
|||||||
self.supports_image_to_image = False
|
self.supports_image_to_image = False
|
||||||
self.supports_lora = 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":
|
if path_type == "load":
|
||||||
self.add_argument("--path", type=str, default=None, help="Local path for loading a model from disk")
|
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:
|
def parse_args(self, **kwargs) -> argparse.Namespace:
|
||||||
namespace = super().parse_args()
|
namespace = super().parse_args()
|
||||||
if hasattr(namespace, "path") and namespace.path is not None and namespace.model is None:
|
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):
|
if getattr(namespace, "config_from_metadata", False):
|
||||||
prior_gen_metadata = json.load(namespace.config_from_metadata.open("rt"))
|
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:
|
if namespace.prompt is None:
|
||||||
namespace.prompt = prior_gen_metadata.get("prompt", 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)
|
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)):
|
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
|
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:
|
if self.supports_image_generation and namespace.prompt is None:
|
||||||
# not supplied by CLI and not supplied by metadata config file
|
# 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:
|
def _create_mflux_generate_parser(with_controlnet=False) -> CommandLineParser:
|
||||||
parser = CommandLineParser(description="Generate an image based on a prompt.")
|
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_image_generator_arguments(supports_metadata_config=True)
|
||||||
parser.add_lora_arguments()
|
parser.add_lora_arguments()
|
||||||
parser.add_image_to_image_arguments(required=False)
|
parser.add_image_to_image_arguments(required=False)
|
||||||
@ -32,19 +32,19 @@ def mflux_generate_controlnet_parser() -> CommandLineParser:
|
|||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def mflux_save_parser() -> CommandLineParser:
|
def mflux_save_parser() -> CommandLineParser:
|
||||||
parser = CommandLineParser(description="Save a quantized version of Flux.1 to disk.") # fmt: off
|
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()
|
parser.add_lora_arguments()
|
||||||
return parser
|
return parser
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def mflux_generate_minimal_argv() -> list[str]:
|
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
|
@pytest.fixture
|
||||||
def mflux_generate_controlnet_minimal_argv() -> list[str]:
|
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
|
@pytest.fixture
|
||||||
@ -58,9 +58,9 @@ def temp_dir(tmp_path_factory) -> Path:
|
|||||||
def base_metadata_dict() -> dict:
|
def base_metadata_dict() -> dict:
|
||||||
return {
|
return {
|
||||||
"mflux_version": "0.4.0",
|
"mflux_version": "0.4.0",
|
||||||
"model": "schnell",
|
"model": "dev",
|
||||||
"seed": 42042,
|
"seed": 42042,
|
||||||
"steps": 4,
|
"steps": 14,
|
||||||
"guidance": None,
|
"guidance": None,
|
||||||
"precision": "mlx.core.bfloat16",
|
"precision": "mlx.core.bfloat16",
|
||||||
"quantize": None,
|
"quantize": None,
|
||||||
@ -71,6 +71,7 @@ def base_metadata_dict() -> dict:
|
|||||||
"init_image_strength": None,
|
"init_image_strength": None,
|
||||||
"controlnet_image": None,
|
"controlnet_image": None,
|
||||||
"controlnet_strength": 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)
|
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):
|
def test_prompt_arg(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir):
|
||||||
metadata_file = temp_dir / "prompt.json"
|
metadata_file = temp_dir / "prompt.json"
|
||||||
file_prompt = "origin of the universe"
|
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
|
base_metadata_dict["steps"] = 8
|
||||||
json.dump(base_metadata_dict, m, indent=4)
|
json.dump(base_metadata_dict, m, indent=4)
|
||||||
|
|
||||||
# test user default value
|
# test user default value for dev
|
||||||
with patch("sys.argv", mflux_generate_minimal_argv):
|
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()
|
args = mflux_generate_parser.parse_args()
|
||||||
assert args.steps == 4
|
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)
|
json.dump(base_metadata_dict, m, indent=4)
|
||||||
|
|
||||||
# test user default value
|
# 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()
|
args = mflux_generate_parser.parse_args()
|
||||||
assert args.lora_paths is None
|
assert args.lora_paths is None
|
||||||
assert args.lora_scales 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_paths == test_paths
|
||||||
assert args.lora_scales == [pytest.approx(0.3), pytest.approx(0.7)]
|
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"]
|
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
|
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()
|
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)
|
json.dump(base_metadata_dict, m, indent=4)
|
||||||
|
|
||||||
# test user default value
|
# 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()
|
args = mflux_generate_parser.parse_args()
|
||||||
assert args.init_image_path is None
|
assert args.init_image_path is None
|
||||||
assert args.init_image_strength == 0.4 # default
|
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()
|
args = mflux_generate_controlnet_parser.parse_args()
|
||||||
assert args.controlnet_image_path == test_path
|
assert args.controlnet_image_path == test_path
|
||||||
assert args.controlnet_strength == pytest.approx(0.48)
|
assert args.controlnet_strength == pytest.approx(0.48)
|
||||||
|
assert args.controlnet_save_canny is False
|
||||||
|
|
||||||
# test CLI override
|
# 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
|
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()
|
args = mflux_generate_controlnet_parser.parse_args()
|
||||||
assert args.controlnet_image_path == "/some/lora/2.safetensors"
|
assert args.controlnet_image_path == "/some/lora/2.safetensors"
|
||||||
assert args.controlnet_strength == pytest.approx(0.85)
|
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):
|
def test_save_args(mflux_save_parser):
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user