From af01fab7d5502530d274b06c6960315b17821f34 Mon Sep 17 00:00:00 2001 From: Anthony Wu Date: Tue, 21 Jan 2025 18:14:24 -0800 Subject: [PATCH] support multiple seeds batch generation --- README.md | 6 ++-- src/mflux/generate.py | 35 +++++++++--------- src/mflux/generate_controlnet.py | 38 ++++++++++---------- src/mflux/ui/cli/parsers.py | 27 ++++++++++++-- tests/arg_parser/test_cli_argparser.py | 50 +++++++++++++++++++++++--- 5 files changed, 110 insertions(+), 46 deletions(-) diff --git a/README.md b/README.md index 9538358..4fb6194 100644 --- a/README.md +++ b/README.md @@ -155,9 +155,11 @@ mflux-generate --model dev --prompt "Luxury food photograph" --steps 25 --seed 2 - **`--model`** or **`-m`** (required, `str`): Model to use for generation (`"schnell"` or `"dev"`). -- **`--output`** (optional, `str`, default: `"image.png"`): Output image filename. +- **`--output`** (optional, `str`, default: `"image.png"`): Output image filename. If `--seed` `--auto-seeds` establishes N > 1 seed values, the "stem" of the output file name will automatically append `_seed_{value}`. -- **`--seed`** (optional, `int`, default: `None`): Seed for random number generation. Default is time-based. +- **`--seed`** (optional, repeatable `int` args, default: `None`): 1 or more seeds for random number generation. e.g. `--seed 42` or `--seed 123 456 789`. Default is a single time-based value. + +- **`--auto-seeds`** (optional, `int`, default: `None`): Auto generate N random Seeds in a series of image generations. Superceded by `--seed` arg and `seed` values in `--config-from-metadata` files. - **`--height`** (optional, `int`, default: `1024`): Height of the output image in pixels. diff --git a/src/mflux/generate.py b/src/mflux/generate.py index 7d07b5d..5579f24 100644 --- a/src/mflux/generate.py +++ b/src/mflux/generate.py @@ -1,4 +1,3 @@ -import time from pathlib import Path from mflux import Config, Flux1, ModelLookup, StopImageGenerationException @@ -25,23 +24,23 @@ def main(): ) try: - # Generate an image - image = flux.generate_image( - seed=int(time.time()) if args.seed is None else args.seed, - prompt=args.prompt, - stepwise_output_dir=Path(args.stepwise_image_output_dir) if args.stepwise_image_output_dir else None, - config=Config( - num_inference_steps=args.steps, - height=args.height, - width=args.width, - guidance=args.guidance, - init_image_path=args.init_image_path, - init_image_strength=args.init_image_strength, - ), - ) - - # Save the image - image.save(path=args.output, export_json_metadata=args.metadata) + for seed_value in args.seed: # 1+ values: see argparser --seed and --auto-seeds + # Generate an image for each seed value + image = flux.generate_image( + seed=seed_value, + prompt=args.prompt, + stepwise_output_dir=Path(args.stepwise_image_output_dir) if args.stepwise_image_output_dir else None, + config=Config( + num_inference_steps=args.steps, + height=args.height, + width=args.width, + guidance=args.guidance, + init_image_path=args.init_image_path, + init_image_strength=args.init_image_strength, + ), + ) + # Save the image + image.save(path=args.output.format(seed=seed_value), export_json_metadata=args.metadata) except StopImageGenerationException as stop_exc: print(stop_exc) diff --git a/src/mflux/generate_controlnet.py b/src/mflux/generate_controlnet.py index ca04e90..c7ddf30 100644 --- a/src/mflux/generate_controlnet.py +++ b/src/mflux/generate_controlnet.py @@ -1,4 +1,3 @@ -import time from pathlib import Path from mflux import ConfigControlnet, Flux1Controlnet, ModelLookup, StopImageGenerationException @@ -24,25 +23,26 @@ def main(): ) try: - # Generate an image - image = flux.generate_image( - seed=int(time.time()) if args.seed is None else args.seed, - prompt=args.prompt, - output=args.output, - controlnet_image_path=args.controlnet_image_path, - controlnet_save_canny=args.controlnet_save_canny, - stepwise_output_dir=Path(args.stepwise_image_output_dir) if args.stepwise_image_output_dir else None, - config=ConfigControlnet( - num_inference_steps=args.steps, - height=args.height, - width=args.width, - guidance=args.guidance, - controlnet_strength=args.controlnet_strength, - ), - ) + for seed_value in args.seed: # 1+ values: see argparser --seed and --auto-seeds + # Generate an image for each seed value + image = flux.generate_image( + seed=seed_value, + prompt=args.prompt, + output=args.output, + controlnet_image_path=args.controlnet_image_path, + controlnet_save_canny=args.controlnet_save_canny, + stepwise_output_dir=Path(args.stepwise_image_output_dir) if args.stepwise_image_output_dir else None, + config=ConfigControlnet( + num_inference_steps=args.steps, + height=args.height, + width=args.width, + guidance=args.guidance, + controlnet_strength=args.controlnet_strength, + ), + ) - # Save the image - image.save(path=args.output, export_json_metadata=args.metadata) + # Save the image + image.save(path=args.output.format(seed=seed_value), export_json_metadata=args.metadata) except StopImageGenerationException as stop_exc: print(stop_exc) diff --git a/src/mflux/ui/cli/parsers.py b/src/mflux/ui/cli/parsers.py index 7710e60..92a1bb9 100644 --- a/src/mflux/ui/cli/parsers.py +++ b/src/mflux/ui/cli/parsers.py @@ -1,5 +1,7 @@ import argparse import json +import random +import time import typing as t from pathlib import Path @@ -56,7 +58,8 @@ class CommandLineParser(argparse.ArgumentParser): def add_image_generator_arguments(self, supports_metadata_config=False) -> None: self.add_argument("--prompt", type=str, required=(not supports_metadata_config), default=None, help="The textual description of the image to generate.") - self.add_argument("--seed", type=int, default=None, help="Entropy Seed (Default is time-based random-seed)") + self.add_argument("--seed", type=int, default=None, nargs='+', help="Specify 1+ Entropy Seeds (Default is 1 time-based random-seed)") + self.add_argument("--auto-seeds", type=int, default=-1, help="Auto generate N Entropy Seeds (random ints between 0 and 1 billion") self._add_image_generator_common_arguments() if supports_metadata_config: self.add_metadata_config() @@ -121,8 +124,14 @@ class CommandLineParser(argparse.ArgumentParser): namespace.guidance = guidance_from_metadata if namespace.quantize is None: namespace.quantize = prior_gen_metadata.get("quantize", None) + seed_from_metadata = prior_gen_metadata.get("seed", None) + if namespace.seed is None and seed_from_metadata is not None: + namespace.seed = [seed_from_metadata] + if namespace.seed is None: - namespace.seed = prior_gen_metadata.get("seed", None) + # not passed by user, not populated by metadata + namespace.seed = [int(time.time())] + if namespace.steps is None: namespace.steps = prior_gen_metadata.get("steps", None) @@ -157,6 +166,20 @@ class CommandLineParser(argparse.ArgumentParser): if namespace.model is None and not has_training_args: self.error("--model / -m must be provided, or 'model' must be specified in the config file.") + if self.supports_image_generation and namespace.seed is None and namespace.auto_seeds > 0: + # choose N int seeds in the range of 0 < value < 1 billion + namespace.seed = [random.randint(0, int(1e7)) for _ in range(namespace.auto_seeds)] + + if self.supports_image_generation and namespace.seed is None: + # final default: did not obtain seed from metadata, --seed, or --auto-seeds + namespace.seed = [int(time.time())] + + if self.supports_image_generation and len(namespace.seed) > 1: + # auto append seed-$value to output names for multi image generations + # e.g. output.png -> output_seed_101.png output_seed_102.png, etc + output_path = Path(namespace.output) + namespace.output = str(output_path.with_stem(output_path.stem + "_seed_{seed}")) + if self.supports_image_generation and namespace.prompt is None: # not supplied by CLI and not supplied by metadata config file self.error("--prompt argument required or 'prompt' required in metadata config file") diff --git a/tests/arg_parser/test_cli_argparser.py b/tests/arg_parser/test_cli_argparser.py index 1558a3b..8ea9d9a 100644 --- a/tests/arg_parser/test_cli_argparser.py +++ b/tests/arg_parser/test_cli_argparser.py @@ -1,4 +1,5 @@ import json +import random from pathlib import Path from unittest.mock import patch @@ -42,6 +43,11 @@ def mflux_generate_minimal_argv() -> list[str]: return ["mflux-generate", "--prompt", "meaning of life"] +@pytest.fixture +def mflux_generate_minimal_model_argv() -> list[str]: + return ["mflux-generate", "--prompt", "meaning of life", "--model", "dev"] + + @pytest.fixture def mflux_generate_controlnet_minimal_argv() -> list[str]: return ["mflux-generate-controlnet", "--prompt", "meaning of life, imitated"] @@ -182,19 +188,53 @@ def test_quantize_arg(mflux_generate_parser, mflux_generate_minimal_argv, base_m assert args.quantize == 8 -def test_seed_arg(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir): # fmt: off +def test_seed_arg(mflux_generate_parser, mflux_generate_minimal_model_argv, base_metadata_dict, temp_dir): # fmt: off metadata_file = temp_dir / "seed.json" with metadata_file.open("wt") as m: base_metadata_dict["seed"] = 24 json.dump(base_metadata_dict, m, indent=4) # test metadata config accepted - with patch('sys.argv', mflux_generate_minimal_argv + ['--config-from-metadata', metadata_file.as_posix()]): # fmt: off + with patch('sys.argv', mflux_generate_minimal_model_argv + ['--config-from-metadata', metadata_file.as_posix()]): # fmt: off args = mflux_generate_parser.parse_args() - assert args.seed == 24 + assert args.seed == [24] + assert "_seed_{seed}" not in args.output + # test CLI override - with patch('sys.argv', mflux_generate_minimal_argv + ['--seed', '2424', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off + with patch('sys.argv', mflux_generate_minimal_model_argv + ['--seed', '2424', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off args = mflux_generate_parser.parse_args() - assert args.seed == 2424 + # --seed arg overrides metadata + assert args.seed == [2424] + assert "_seed_{seed}" not in args.output + + with patch('sys.argv', mflux_generate_minimal_model_argv + ['--seed', '2424', '4848', '9696']): # fmt: off + args = mflux_generate_parser.parse_args() + assert args.seed == [2424, 4848, 9696] + assert "_seed_{seed}" in args.output + + with patch('sys.argv', mflux_generate_minimal_model_argv + ['--auto-seeds', '5', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off + args = mflux_generate_parser.parse_args() + # auto-seeds defers to value from metadata, and is ignored + assert len(args.seed) == 1 + assert args.seed == [24] + assert "_seed_{seed}" not in args.output + + +def test_auto_seeds_arg(mflux_generate_parser, mflux_generate_minimal_model_argv): + with patch("sys.argv", mflux_generate_minimal_model_argv + ["--seed", "24", "48", "--auto-seeds", "5"]): + args = mflux_generate_parser.parse_args() + # auto-seeds defers to explicit values of --seed + assert len(args.seed) == 2 + assert args.seed == [24, 48] + assert "_seed_{seed}" in args.output + + for _ in range(0, 10): + random_auto_seed_count = random.randint(0, 100) + with patch("sys.argv", mflux_generate_minimal_model_argv + ["--auto-seeds", str(random_auto_seed_count)]): + args = mflux_generate_parser.parse_args() + assert len(set(args.seed)) == random_auto_seed_count + assert "_seed_{seed}" in args.output + for _ in args.seed: + assert isinstance(_, int) def test_steps_arg(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir): # fmt: off