Merge pull request #118 from anthonywu/run-n-seeds-auto-seeds

support multiple seeds batch generation
This commit is contained in:
Filip Strand 2025-01-22 19:58:51 +01:00 committed by GitHub
commit 4bea7b9668
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
5 changed files with 110 additions and 46 deletions

View File

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

View File

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

View File

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

View File

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

View File

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