Merge pull request #118 from anthonywu/run-n-seeds-auto-seeds
support multiple seeds batch generation
This commit is contained in:
commit
4bea7b9668
@ -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"`).
|
- **`--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.
|
- **`--height`** (optional, `int`, default: `1024`): Height of the output image in pixels.
|
||||||
|
|
||||||
|
|||||||
@ -1,4 +1,3 @@
|
|||||||
import time
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from mflux import Config, Flux1, ModelLookup, StopImageGenerationException
|
from mflux import Config, Flux1, ModelLookup, StopImageGenerationException
|
||||||
@ -25,23 +24,23 @@ def main():
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Generate an image
|
for seed_value in args.seed: # 1+ values: see argparser --seed and --auto-seeds
|
||||||
image = flux.generate_image(
|
# Generate an image for each seed value
|
||||||
seed=int(time.time()) if args.seed is None else args.seed,
|
image = flux.generate_image(
|
||||||
prompt=args.prompt,
|
seed=seed_value,
|
||||||
stepwise_output_dir=Path(args.stepwise_image_output_dir) if args.stepwise_image_output_dir else None,
|
prompt=args.prompt,
|
||||||
config=Config(
|
stepwise_output_dir=Path(args.stepwise_image_output_dir) if args.stepwise_image_output_dir else None,
|
||||||
num_inference_steps=args.steps,
|
config=Config(
|
||||||
height=args.height,
|
num_inference_steps=args.steps,
|
||||||
width=args.width,
|
height=args.height,
|
||||||
guidance=args.guidance,
|
width=args.width,
|
||||||
init_image_path=args.init_image_path,
|
guidance=args.guidance,
|
||||||
init_image_strength=args.init_image_strength,
|
init_image_path=args.init_image_path,
|
||||||
),
|
init_image_strength=args.init_image_strength,
|
||||||
)
|
),
|
||||||
|
)
|
||||||
# Save the image
|
# Save the image
|
||||||
image.save(path=args.output, export_json_metadata=args.metadata)
|
image.save(path=args.output.format(seed=seed_value), export_json_metadata=args.metadata)
|
||||||
except StopImageGenerationException as stop_exc:
|
except StopImageGenerationException as stop_exc:
|
||||||
print(stop_exc)
|
print(stop_exc)
|
||||||
|
|
||||||
|
|||||||
@ -1,4 +1,3 @@
|
|||||||
import time
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from mflux import ConfigControlnet, Flux1Controlnet, ModelLookup, StopImageGenerationException
|
from mflux import ConfigControlnet, Flux1Controlnet, ModelLookup, StopImageGenerationException
|
||||||
@ -24,25 +23,26 @@ def main():
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Generate an image
|
for seed_value in args.seed: # 1+ values: see argparser --seed and --auto-seeds
|
||||||
image = flux.generate_image(
|
# Generate an image for each seed value
|
||||||
seed=int(time.time()) if args.seed is None else args.seed,
|
image = flux.generate_image(
|
||||||
prompt=args.prompt,
|
seed=seed_value,
|
||||||
output=args.output,
|
prompt=args.prompt,
|
||||||
controlnet_image_path=args.controlnet_image_path,
|
output=args.output,
|
||||||
controlnet_save_canny=args.controlnet_save_canny,
|
controlnet_image_path=args.controlnet_image_path,
|
||||||
stepwise_output_dir=Path(args.stepwise_image_output_dir) if args.stepwise_image_output_dir else None,
|
controlnet_save_canny=args.controlnet_save_canny,
|
||||||
config=ConfigControlnet(
|
stepwise_output_dir=Path(args.stepwise_image_output_dir) if args.stepwise_image_output_dir else None,
|
||||||
num_inference_steps=args.steps,
|
config=ConfigControlnet(
|
||||||
height=args.height,
|
num_inference_steps=args.steps,
|
||||||
width=args.width,
|
height=args.height,
|
||||||
guidance=args.guidance,
|
width=args.width,
|
||||||
controlnet_strength=args.controlnet_strength,
|
guidance=args.guidance,
|
||||||
),
|
controlnet_strength=args.controlnet_strength,
|
||||||
)
|
),
|
||||||
|
)
|
||||||
|
|
||||||
# Save the image
|
# Save the image
|
||||||
image.save(path=args.output, export_json_metadata=args.metadata)
|
image.save(path=args.output.format(seed=seed_value), export_json_metadata=args.metadata)
|
||||||
except StopImageGenerationException as stop_exc:
|
except StopImageGenerationException as stop_exc:
|
||||||
print(stop_exc)
|
print(stop_exc)
|
||||||
|
|
||||||
|
|||||||
@ -1,5 +1,7 @@
|
|||||||
import argparse
|
import argparse
|
||||||
import json
|
import json
|
||||||
|
import random
|
||||||
|
import time
|
||||||
import typing as t
|
import typing as t
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
@ -56,7 +58,8 @@ class CommandLineParser(argparse.ArgumentParser):
|
|||||||
|
|
||||||
def add_image_generator_arguments(self, supports_metadata_config=False) -> None:
|
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("--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()
|
self._add_image_generator_common_arguments()
|
||||||
if supports_metadata_config:
|
if supports_metadata_config:
|
||||||
self.add_metadata_config()
|
self.add_metadata_config()
|
||||||
@ -121,8 +124,14 @@ class CommandLineParser(argparse.ArgumentParser):
|
|||||||
namespace.guidance = guidance_from_metadata
|
namespace.guidance = guidance_from_metadata
|
||||||
if namespace.quantize is None:
|
if namespace.quantize is None:
|
||||||
namespace.quantize = prior_gen_metadata.get("quantize", 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:
|
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:
|
if namespace.steps is None:
|
||||||
namespace.steps = prior_gen_metadata.get("steps", 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:
|
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.")
|
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:
|
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
|
||||||
self.error("--prompt argument required or 'prompt' required in metadata config file")
|
self.error("--prompt argument required or 'prompt' required in metadata config file")
|
||||||
|
|||||||
@ -1,4 +1,5 @@
|
|||||||
import json
|
import json
|
||||||
|
import random
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
@ -42,6 +43,11 @@ def mflux_generate_minimal_argv() -> list[str]:
|
|||||||
return ["mflux-generate", "--prompt", "meaning of life"]
|
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
|
@pytest.fixture
|
||||||
def mflux_generate_controlnet_minimal_argv() -> list[str]:
|
def mflux_generate_controlnet_minimal_argv() -> list[str]:
|
||||||
return ["mflux-generate-controlnet", "--prompt", "meaning of life, imitated"]
|
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
|
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"
|
metadata_file = temp_dir / "seed.json"
|
||||||
with metadata_file.open("wt") as m:
|
with metadata_file.open("wt") as m:
|
||||||
base_metadata_dict["seed"] = 24
|
base_metadata_dict["seed"] = 24
|
||||||
json.dump(base_metadata_dict, m, indent=4)
|
json.dump(base_metadata_dict, m, indent=4)
|
||||||
# test metadata config accepted
|
# 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()
|
args = mflux_generate_parser.parse_args()
|
||||||
assert args.seed == 24
|
assert args.seed == [24]
|
||||||
|
assert "_seed_{seed}" not in args.output
|
||||||
|
|
||||||
# test CLI override
|
# 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()
|
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
|
def test_steps_arg(mflux_generate_parser, mflux_generate_minimal_argv, base_metadata_dict, temp_dir): # fmt: off
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user