Renaming: Remove 'init' prefix for a more general interface
This commit is contained in:
parent
552eb6a3bb
commit
5ec7ab5b41
10
README.md
10
README.md
@ -185,9 +185,9 @@ mflux-generate --model dev --prompt "Luxury food photograph" --steps 25 --seed 2
|
|||||||
|
|
||||||
- **`--controlnet-save-canny`** (optional, bool, default: False): If set, saves the Canny edge detection reference image used by ControlNet.
|
- **`--controlnet-save-canny`** (optional, bool, default: False): If set, saves the Canny edge detection reference image used by ControlNet.
|
||||||
|
|
||||||
- **`--init-image-path`** (optional, `str`, default: `None`): Local path to the initial image for image-to-image generation.
|
- **`--image-path`** (optional, `str`, default: `None`): Local path to the initial image for image-to-image generation.
|
||||||
|
|
||||||
- **`--init-image-strength`** (optional, `float`, default: `0.4`): Controls how strongly the initial image influences the output image. A value of `0.0` means no influence. (Default is `0.4`)
|
- **`--image-strength`** (optional, `float`, default: `0.4`): Controls how strongly the initial image influences the output image. A value of `0.0` means no influence. (Default is `0.4`)
|
||||||
|
|
||||||
- **`--config-from-metadata`** or **`-C`** (optional, `str`): [EXPERIMENTAL] Path to a prior file saved via `--metadata`, or a compatible handcrafted config file adhering to the expected args schema.
|
- **`--config-from-metadata`** or **`-C`** (optional, `str`): [EXPERIMENTAL] Path to a prior file saved via `--metadata`, or a compatible handcrafted config file adhering to the expected args schema.
|
||||||
|
|
||||||
@ -514,15 +514,15 @@ processed a bit differently, which is why we require this structure above.*
|
|||||||
### 🎨 Image-to-Image
|
### 🎨 Image-to-Image
|
||||||
|
|
||||||
One way to condition the image generation is by starting from an existing image and let MFLUX produce new variations.
|
One way to condition the image generation is by starting from an existing image and let MFLUX produce new variations.
|
||||||
Use the `--init-image-path` flag to specify the reference image, and the `--init-image-strength` to control how much the reference
|
Use the `--image-path` flag to specify the reference image, and the `--image-strength` to control how much the reference
|
||||||
image should guide the generation. For example, given the reference image below, the following command produced the first
|
image should guide the generation. For example, given the reference image below, the following command produced the first
|
||||||
image using the [Sketching](https://civitai.com/models/803456/sketching?modelVersionId=898364) LoRA:
|
image using the [Sketching](https://civitai.com/models/803456/sketching?modelVersionId=898364) LoRA:
|
||||||
|
|
||||||
```sh
|
```sh
|
||||||
mflux-generate \
|
mflux-generate \
|
||||||
--prompt "sketching of an Eiffel architecture, masterpiece, best quality. The site is lit by lighting professionals, creating a subtle illumination effect. Ink on paper with very fine touches with colored markers, (shadings:1.1), loose lines, Schematic, Conceptual, Abstract, Gestural. Quick sketches to explore ideas and concepts." \
|
--prompt "sketching of an Eiffel architecture, masterpiece, best quality. The site is lit by lighting professionals, creating a subtle illumination effect. Ink on paper with very fine touches with colored markers, (shadings:1.1), loose lines, Schematic, Conceptual, Abstract, Gestural. Quick sketches to explore ideas and concepts." \
|
||||||
--init-image-path "reference.png" \
|
--image-path "reference.png" \
|
||||||
--init-image-strength 0.3 \
|
--image-strength 0.3 \
|
||||||
--lora-paths Architectural_Sketching.safetensors \
|
--lora-paths Architectural_Sketching.safetensors \
|
||||||
--lora-scales 1.0 \
|
--lora-scales 1.0 \
|
||||||
--model dev \
|
--model dev \
|
||||||
|
|||||||
@ -15,8 +15,8 @@ class Config:
|
|||||||
width: int = 1024,
|
width: int = 1024,
|
||||||
height: int = 1024,
|
height: int = 1024,
|
||||||
guidance: float = 4.0,
|
guidance: float = 4.0,
|
||||||
init_image_path: Path | None = None,
|
image_path: Path | None = None,
|
||||||
init_image_strength: float | None = None,
|
image_strength: float | None = None,
|
||||||
controlnet_strength: float | None = None,
|
controlnet_strength: float | None = None,
|
||||||
):
|
):
|
||||||
if width % 16 != 0 or height % 16 != 0:
|
if width % 16 != 0 or height % 16 != 0:
|
||||||
@ -25,6 +25,6 @@ class Config:
|
|||||||
self.height = 16 * (height // 16)
|
self.height = 16 * (height // 16)
|
||||||
self.num_inference_steps = num_inference_steps
|
self.num_inference_steps = num_inference_steps
|
||||||
self.guidance = guidance
|
self.guidance = guidance
|
||||||
self.init_image_path = init_image_path
|
self.image_path = image_path
|
||||||
self.init_image_strength = init_image_strength
|
self.image_strength = image_strength
|
||||||
self.controlnet_strength = controlnet_strength
|
self.controlnet_strength = controlnet_strength
|
||||||
|
|||||||
@ -44,12 +44,12 @@ class RuntimeConfig:
|
|||||||
return self.model_config.num_train_steps
|
return self.model_config.num_train_steps
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def init_image_path(self) -> str:
|
def image_path(self) -> str:
|
||||||
return self.config.init_image_path
|
return self.config.image_path
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def init_image_strength(self) -> float:
|
def image_strength(self) -> float | None:
|
||||||
return self.config.init_image_strength
|
return self.config.image_strength
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def init_time_step(self) -> int:
|
def init_time_step(self) -> int:
|
||||||
|
|||||||
@ -62,7 +62,7 @@ class Flux1(nn.Module):
|
|||||||
vae=self.vae,
|
vae=self.vae,
|
||||||
sigmas=config.sigmas,
|
sigmas=config.sigmas,
|
||||||
init_time_step=config.init_time_step,
|
init_time_step=config.init_time_step,
|
||||||
init_image_path=config.init_image_path,
|
image_path=config.image_path,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@ -141,8 +141,8 @@ class Flux1(nn.Module):
|
|||||||
quantization=self.bits,
|
quantization=self.bits,
|
||||||
lora_paths=self.lora_paths,
|
lora_paths=self.lora_paths,
|
||||||
lora_scales=self.lora_scales,
|
lora_scales=self.lora_scales,
|
||||||
init_image_path=config.init_image_path,
|
image_path=config.image_path,
|
||||||
init_image_strength=config.init_image_strength,
|
image_strength=config.image_strength,
|
||||||
generation_time=time_steps.format_dict["elapsed"],
|
generation_time=time_steps.format_dict["elapsed"],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -50,8 +50,8 @@ def main():
|
|||||||
height=args.height,
|
height=args.height,
|
||||||
width=args.width,
|
width=args.width,
|
||||||
guidance=args.guidance,
|
guidance=args.guidance,
|
||||||
init_image_path=args.init_image_path,
|
image_path=args.image_path,
|
||||||
init_image_strength=args.init_image_strength,
|
image_strength=args.image_strength,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
# 4. Save the image
|
# 4. Save the image
|
||||||
|
|||||||
@ -11,12 +11,12 @@ class Img2Img:
|
|||||||
vae: VAE,
|
vae: VAE,
|
||||||
sigmas: mx.array,
|
sigmas: mx.array,
|
||||||
init_time_step: int,
|
init_time_step: int,
|
||||||
init_image_path: int,
|
image_path: int,
|
||||||
):
|
):
|
||||||
self.vae = vae
|
self.vae = vae
|
||||||
self.sigmas = sigmas
|
self.sigmas = sigmas
|
||||||
self.init_time_step = init_time_step
|
self.init_time_step = init_time_step
|
||||||
self.init_image_path = init_image_path
|
self.image_path = image_path
|
||||||
|
|
||||||
|
|
||||||
class LatentCreator:
|
class LatentCreator:
|
||||||
@ -39,7 +39,7 @@ class LatentCreator:
|
|||||||
img2img: Img2Img,
|
img2img: Img2Img,
|
||||||
) -> mx.array:
|
) -> mx.array:
|
||||||
# 0. Determine type of image generation
|
# 0. Determine type of image generation
|
||||||
is_text2img = img2img.init_image_path is None
|
is_text2img = img2img.image_path is None
|
||||||
|
|
||||||
if is_text2img:
|
if is_text2img:
|
||||||
# 1. Create the pure noise
|
# 1. Create the pure noise
|
||||||
|
|||||||
@ -25,8 +25,8 @@ class GeneratedImage:
|
|||||||
lora_scales: list[float],
|
lora_scales: list[float],
|
||||||
controlnet_image_path: str | pathlib.Path | None = None,
|
controlnet_image_path: str | pathlib.Path | None = None,
|
||||||
controlnet_strength: float | None = None,
|
controlnet_strength: float | None = None,
|
||||||
init_image_path: str | pathlib.Path | None = None,
|
image_path: str | pathlib.Path | None = None,
|
||||||
init_image_strength: float | None = None,
|
image_strength: float | None = None,
|
||||||
):
|
):
|
||||||
self.image = image
|
self.image = image
|
||||||
self.model_config = model_config
|
self.model_config = model_config
|
||||||
@ -41,8 +41,8 @@ class GeneratedImage:
|
|||||||
self.lora_scales = lora_scales
|
self.lora_scales = lora_scales
|
||||||
self.controlnet_image_path = controlnet_image_path
|
self.controlnet_image_path = controlnet_image_path
|
||||||
self.controlnet_strength = controlnet_strength
|
self.controlnet_strength = controlnet_strength
|
||||||
self.init_image_path = init_image_path
|
self.image_path = image_path
|
||||||
self.init_image_strength = init_image_strength
|
self.image_strength = image_strength
|
||||||
|
|
||||||
def save(
|
def save(
|
||||||
self,
|
self,
|
||||||
@ -72,8 +72,8 @@ class GeneratedImage:
|
|||||||
"generation_time_seconds": round(self.generation_time, 2),
|
"generation_time_seconds": round(self.generation_time, 2),
|
||||||
"lora_paths": [str(p) for p in self.lora_paths] if self.lora_paths else None,
|
"lora_paths": [str(p) for p in self.lora_paths] if self.lora_paths else None,
|
||||||
"lora_scales": [round(scale, 2) for scale in self.lora_scales] if self.lora_scales else None,
|
"lora_scales": [round(scale, 2) for scale in self.lora_scales] if self.lora_scales else None,
|
||||||
"init_image_path": str(self.init_image_path) if self.init_image_path else None,
|
"image_path": str(self.image_path) if self.image_path else None,
|
||||||
"init_image_strength": self.init_image_strength if self.init_image_path else None,
|
"image_strength": self.image_strength if self.image_path else None,
|
||||||
"controlnet_image_path": str(self.controlnet_image_path) if self.controlnet_image_path else None,
|
"controlnet_image_path": str(self.controlnet_image_path) if self.controlnet_image_path else None,
|
||||||
"controlnet_strength": round(self.controlnet_strength, 2) if self.controlnet_strength else None,
|
"controlnet_strength": round(self.controlnet_strength, 2) if self.controlnet_strength else None,
|
||||||
"prompt": self.prompt,
|
"prompt": self.prompt,
|
||||||
|
|||||||
@ -26,8 +26,8 @@ class ImageUtil:
|
|||||||
lora_paths: list[str],
|
lora_paths: list[str],
|
||||||
lora_scales: list[float],
|
lora_scales: list[float],
|
||||||
controlnet_image_path: str | None = None,
|
controlnet_image_path: str | None = None,
|
||||||
init_image_path: str | None = None,
|
image_path: str | None = None,
|
||||||
init_image_strength: float | None = None,
|
image_strength: float | None = None,
|
||||||
) -> GeneratedImage:
|
) -> GeneratedImage:
|
||||||
normalized = ImageUtil._denormalize(decoded_latents)
|
normalized = ImageUtil._denormalize(decoded_latents)
|
||||||
normalized_numpy = ImageUtil._to_numpy(normalized)
|
normalized_numpy = ImageUtil._to_numpy(normalized)
|
||||||
@ -44,8 +44,8 @@ class ImageUtil:
|
|||||||
generation_time=generation_time,
|
generation_time=generation_time,
|
||||||
lora_paths=lora_paths,
|
lora_paths=lora_paths,
|
||||||
lora_scales=lora_scales,
|
lora_scales=lora_scales,
|
||||||
init_image_path=init_image_path,
|
image_path=image_path,
|
||||||
init_image_strength=init_image_strength,
|
image_strength=image_strength,
|
||||||
controlnet_image_path=controlnet_image_path,
|
controlnet_image_path=controlnet_image_path,
|
||||||
controlnet_strength=config.controlnet_strength,
|
controlnet_strength=config.controlnet_strength,
|
||||||
)
|
)
|
||||||
|
|||||||
@ -69,8 +69,8 @@ class CommandLineParser(argparse.ArgumentParser):
|
|||||||
|
|
||||||
def add_image_to_image_arguments(self, required=False) -> None:
|
def add_image_to_image_arguments(self, required=False) -> None:
|
||||||
self.supports_image_to_image = True
|
self.supports_image_to_image = True
|
||||||
self.add_argument("--init-image-path", type=Path, required=required, default=None, help="Local path to init image")
|
self.add_argument("--image-path", type=Path, required=required, default=None, help="Local path to init image")
|
||||||
self.add_argument("--init-image-strength", type=float, required=False, default=ui_defaults.INIT_IMAGE_STRENGTH, help=f"Controls how strongly the init image influences the output image. A value of 0.0 means no influence. (Default is {ui_defaults.INIT_IMAGE_STRENGTH})")
|
self.add_argument("--image-strength", type=float, required=False, default=ui_defaults.IMAGE_STRENGTH, help=f"Controls how strongly the init image influences the output image. A value of 0.0 means no influence. (Default is {ui_defaults.IMAGE_STRENGTH})")
|
||||||
|
|
||||||
def add_batch_image_generator_arguments(self) -> None:
|
def add_batch_image_generator_arguments(self) -> None:
|
||||||
self.add_argument("--prompts-file", type=Path, required=True, default=argparse.SUPPRESS, help="Local path for a file that holds a batch of prompts.")
|
self.add_argument("--prompts-file", type=Path, required=True, default=argparse.SUPPRESS, help="Local path for a file that holds a batch of prompts.")
|
||||||
@ -152,10 +152,10 @@ class CommandLineParser(argparse.ArgumentParser):
|
|||||||
namespace.lora_scales = prior_gen_metadata.get("lora_scales", []) + namespace.lora_scales
|
namespace.lora_scales = prior_gen_metadata.get("lora_scales", []) + namespace.lora_scales
|
||||||
|
|
||||||
if self.supports_image_to_image:
|
if self.supports_image_to_image:
|
||||||
if namespace.init_image_path is None:
|
if namespace.image_path is None:
|
||||||
namespace.init_image_path = prior_gen_metadata.get("init_image_path", None)
|
namespace.image_path = prior_gen_metadata.get("image_path", None)
|
||||||
if namespace.init_image_strength == self.get_default("init_image_strength") and (init_img_strength_from_metadata := prior_gen_metadata.get("init_image_strength", None)):
|
if namespace.image_strength == self.get_default("image_strength") and (img_strength_from_metadata := prior_gen_metadata.get("image_strength", None)):
|
||||||
namespace.init_image_strength = init_img_strength_from_metadata
|
namespace.image_strength = img_strength_from_metadata
|
||||||
|
|
||||||
if self.supports_controlnet:
|
if self.supports_controlnet:
|
||||||
if namespace.controlnet_image_path is None:
|
if namespace.controlnet_image_path is None:
|
||||||
|
|||||||
@ -1,7 +1,7 @@
|
|||||||
CONTROLNET_STRENGTH = 0.4
|
CONTROLNET_STRENGTH = 0.4
|
||||||
GUIDANCE_SCALE = 3.5
|
GUIDANCE_SCALE = 3.5
|
||||||
HEIGHT, WIDTH = 1024, 1024
|
HEIGHT, WIDTH = 1024, 1024
|
||||||
INIT_IMAGE_STRENGTH = 0.4 # for image-to-image init_image
|
IMAGE_STRENGTH = 0.4
|
||||||
MODEL_CHOICES = ["dev", "schnell"]
|
MODEL_CHOICES = ["dev", "schnell"]
|
||||||
MODEL_INFERENCE_STEPS = {
|
MODEL_INFERENCE_STEPS = {
|
||||||
"dev": 14,
|
"dev": 14,
|
||||||
|
|||||||
@ -75,8 +75,8 @@ def base_metadata_dict() -> dict:
|
|||||||
"generation_time_seconds": 42.0,
|
"generation_time_seconds": 42.0,
|
||||||
"lora_paths": None,
|
"lora_paths": None,
|
||||||
"lora_scales": None,
|
"lora_scales": None,
|
||||||
"init_image": None,
|
"image": None,
|
||||||
"init_image_strength": None,
|
"image_strength": None,
|
||||||
"controlnet_image": None,
|
"controlnet_image": None,
|
||||||
"controlnet_strength": None,
|
"controlnet_strength": None,
|
||||||
"controlnet_save_canny": False,
|
"controlnet_save_canny": False,
|
||||||
@ -300,32 +300,32 @@ def test_image_to_image_args(mflux_generate_parser, mflux_generate_minimal_argv,
|
|||||||
metadata_file = temp_dir / "image_to_image.json"
|
metadata_file = temp_dir / "image_to_image.json"
|
||||||
test_path = "/some/awesome/image.png"
|
test_path = "/some/awesome/image.png"
|
||||||
with metadata_file.open("wt") as m:
|
with metadata_file.open("wt") as m:
|
||||||
base_metadata_dict["init_image_path"] = test_path
|
base_metadata_dict["image_path"] = test_path
|
||||||
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 + ["-m", "dev"]):
|
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.image_path is None
|
||||||
assert args.init_image_strength == 0.4 # default
|
assert args.image_strength == 0.4 # default
|
||||||
|
|
||||||
# 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_argv + ['--config-from-metadata', metadata_file.as_posix()]): # fmt: off
|
||||||
args = mflux_generate_parser.parse_args()
|
args = mflux_generate_parser.parse_args()
|
||||||
assert args.init_image_path == test_path
|
assert args.image_path == test_path
|
||||||
assert args.init_image_strength == 0.4 # default
|
assert args.image_strength == 0.4 # default
|
||||||
|
|
||||||
# test strength override
|
# test strength override
|
||||||
with patch('sys.argv', mflux_generate_minimal_argv + ['--init-image-strength', '0.7', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off
|
with patch('sys.argv', mflux_generate_minimal_argv + ['--image-strength', '0.7', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off
|
||||||
args = mflux_generate_parser.parse_args()
|
args = mflux_generate_parser.parse_args()
|
||||||
assert args.init_image_path == test_path
|
assert args.image_path == test_path
|
||||||
assert args.init_image_strength == 0.7
|
assert args.image_strength == 0.7
|
||||||
|
|
||||||
# test image path override
|
# test image path override
|
||||||
with patch('sys.argv', mflux_generate_minimal_argv + ['--init-image-path', '/some/better/image.png', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off
|
with patch('sys.argv', mflux_generate_minimal_argv + ['--image-path', '/some/better/image.png', '--config-from-metadata', metadata_file.as_posix()]): # fmt: off
|
||||||
args = mflux_generate_parser.parse_args()
|
args = mflux_generate_parser.parse_args()
|
||||||
assert args.init_image_path == Path("/some/better/image.png")
|
assert args.image_path == Path("/some/better/image.png")
|
||||||
assert args.init_image_strength == 0.4 # default
|
assert args.image_strength == 0.4 # default
|
||||||
|
|
||||||
|
|
||||||
def test_controlnet_args(mflux_generate_controlnet_parser, mflux_generate_controlnet_minimal_argv, base_metadata_dict, temp_dir): # fmt: off
|
def test_controlnet_args(mflux_generate_controlnet_parser, mflux_generate_controlnet_minimal_argv, base_metadata_dict, temp_dir): # fmt: off
|
||||||
|
|||||||
@ -18,8 +18,8 @@ class ImageGeneratorTestHelper:
|
|||||||
seed: int,
|
seed: int,
|
||||||
height: int = None,
|
height: int = None,
|
||||||
width: int = None,
|
width: int = None,
|
||||||
init_image_path: str | None = None,
|
image_path: str | None = None,
|
||||||
init_image_strength: float | None = None,
|
image_strength: float | None = None,
|
||||||
lora_paths: list[str] | None = None,
|
lora_paths: list[str] | None = None,
|
||||||
lora_scales: list[float] | None = None,
|
lora_scales: list[float] | None = None,
|
||||||
):
|
):
|
||||||
@ -42,8 +42,8 @@ class ImageGeneratorTestHelper:
|
|||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
config=Config(
|
config=Config(
|
||||||
num_inference_steps=steps,
|
num_inference_steps=steps,
|
||||||
init_image_path=ImageGeneratorTestHelper.resolve_path(init_image_path),
|
image_path=ImageGeneratorTestHelper.resolve_path(image_path),
|
||||||
init_image_strength=init_image_strength,
|
image_strength=image_strength,
|
||||||
height=height,
|
height=height,
|
||||||
width=width,
|
width=width,
|
||||||
),
|
),
|
||||||
|
|||||||
@ -60,8 +60,8 @@ class TestImageGenerator:
|
|||||||
def test_image_generation_dev_image_to_image(self):
|
def test_image_generation_dev_image_to_image(self):
|
||||||
ImageGeneratorTestHelper.assert_matches_reference_image(
|
ImageGeneratorTestHelper.assert_matches_reference_image(
|
||||||
reference_image_path="reference_dev_image_to_image_result.png",
|
reference_image_path="reference_dev_image_to_image_result.png",
|
||||||
init_image_path="reference_dev_lora.png",
|
image_path="reference_dev_lora.png",
|
||||||
init_image_strength=0.4,
|
image_strength=0.4,
|
||||||
output_image_path=TestImageGenerator.OUTPUT_IMAGE_FILENAME,
|
output_image_path=TestImageGenerator.OUTPUT_IMAGE_FILENAME,
|
||||||
model_config=ModelConfig.dev(),
|
model_config=ModelConfig.dev(),
|
||||||
steps=8,
|
steps=8,
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user