Instructions to use ezhoureal/aura_style with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use ezhoureal/aura_style with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline from diffusers.utils import load_image # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("black-forest-labs/FLUX.2-dev,black-forest-labs/FLUX.2-klein-4B,stabilityai/stable-diffusion-3.5-medium", dtype=torch.bfloat16, device_map="cuda") pipe.load_lora_weights("ezhoureal/aura_style") prompt = "Turn this cat into a dog" input_image = load_image("https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/cat.png") image = pipe(image=input_image, prompt=prompt).images[0] - Inference
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- Draw Things
| #!/usr/bin/env python3 | |
| from __future__ import annotations | |
| import argparse | |
| import importlib.util | |
| import json | |
| import re | |
| import sys | |
| import time | |
| from pathlib import Path | |
| from typing import Any | |
| import torch | |
| from safetensors.torch import load_file, save_file | |
| REPO_ROOT = Path(__file__).resolve().parents[1] | |
| DEFAULT_LORA = REPO_ROOT / "fal_flux2_edit_lora" / "pytorch_lora_weights.safetensors" | |
| DEFAULT_OUTPUT_DIR = REPO_ROOT / "outputs" / "local_flux2_edit_inference" | |
| DEFAULT_MODEL = "diffusers/FLUX.2-dev-bnb-4bit" | |
| DEFAULT_PROMPT = ( | |
| "Transform this photorealistic image into the trained radiant aura style: smooth colorful " | |
| "gradients, ethereal haze, subtle contour lighting, and a refined cinematic glow. Preserve the " | |
| "subject identity, composition, pose, silhouette, camera framing, and important details." | |
| ) | |
| SUPPORTED_IMAGE_SUFFIXES = {".avif", ".bmp", ".jpeg", ".jpg", ".png", ".webp"} | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser( | |
| description="Run local FLUX.2 image editing with the fal-trained LoRA, or validate it offline." | |
| ) | |
| parser.add_argument( | |
| "input_dir", | |
| nargs="?", | |
| type=Path, | |
| help="Directory containing photorealistic input images.", | |
| ) | |
| parser.add_argument("--prompt", default=DEFAULT_PROMPT, help="Edit prompt.") | |
| parser.add_argument("--lora", type=Path, default=DEFAULT_LORA, help="Input LoRA safetensors file.") | |
| parser.add_argument( | |
| "--converted-lora", | |
| type=Path, | |
| default=None, | |
| help="Optional path for a converted diffusers-format LoRA safetensors file.", | |
| ) | |
| parser.add_argument( | |
| "--model", | |
| default=DEFAULT_MODEL, | |
| help="Local path or Hugging Face model id. Defaults to the 4-bit FLUX.2-dev diffusers repo.", | |
| ) | |
| parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR) | |
| parser.add_argument("--output-path", type=Path, default=None) | |
| parser.add_argument( | |
| "--batch-size", | |
| type=int, | |
| default=1, | |
| help=( | |
| "Number of input images to edit per pipeline call. Keep this low on <24GB VRAM; " | |
| "try 2 first, then increase if memory allows." | |
| ), | |
| ) | |
| parser.add_argument("--height", type=int, default=1024) | |
| parser.add_argument("--width", type=int, default=1024) | |
| parser.add_argument("--num-inference-steps", type=int, default=28) | |
| parser.add_argument("--guidance-scale", type=float, default=2.5) | |
| parser.add_argument("--lora-scale", type=float, default=1.0) | |
| parser.add_argument("--seed", type=int, default=None) | |
| parser.add_argument( | |
| "--torch-dtype", | |
| choices=("auto", "float32", "float16", "bfloat16"), | |
| default="bfloat16", | |
| help="Pipeline dtype. Use bfloat16 on modern NVIDIA GPUs.", | |
| ) | |
| parser.add_argument( | |
| "--device", | |
| default=None, | |
| help="Torch device. Defaults to cuda if available, otherwise cpu.", | |
| ) | |
| parser.add_argument( | |
| "--device-map", | |
| default=None, | |
| help='Optional diffusers/accelerate device map, for example "balanced".', | |
| ) | |
| parser.add_argument( | |
| "--local-files-only", | |
| action="store_true", | |
| help="Do not download model files from Hugging Face.", | |
| ) | |
| parser.add_argument( | |
| "--check-only", | |
| action="store_true", | |
| help="Validate/convert LoRA against the default Flux2Transformer2DModel shape without loading the base model.", | |
| ) | |
| return parser.parse_args() | |
| def dtype_from_arg(value: str) -> torch.dtype | str: | |
| if value == "auto": | |
| return "auto" | |
| return { | |
| "float32": torch.float32, | |
| "float16": torch.float16, | |
| "bfloat16": torch.bfloat16, | |
| }[value] | |
| def require_module(import_name: str, install_name: str | None = None) -> None: | |
| if importlib.util.find_spec(import_name) is None: | |
| package = install_name or import_name | |
| raise RuntimeError(f"Missing required package `{package}`. Install it with `uv add {package}`.") | |
| def uses_4bit_model(model: str) -> bool: | |
| return "bnb-4bit" in model.lower() or "4bit" in model.lower() | |
| def preflight_environment(args: argparse.Namespace) -> None: | |
| if args.check_only: | |
| return | |
| require_module("google.protobuf", "protobuf") | |
| device_name = args.device or ("cuda" if torch.cuda.is_available() else "cpu") | |
| if uses_4bit_model(args.model): | |
| require_module("bitsandbytes") | |
| if device_name == "cpu" or not torch.cuda.is_available(): | |
| raise RuntimeError( | |
| "The 4-bit FLUX.2 model needs a CUDA GPU with bitsandbytes. " | |
| "This environment does not expose CUDA to PyTorch." | |
| ) | |
| def convert_fal_key(key: str, tensor: torch.Tensor) -> dict[str, torch.Tensor]: | |
| prefix = "base_model.model." | |
| if not key.startswith(prefix): | |
| return {key: tensor} | |
| body = key.removeprefix(prefix) | |
| suffix = ".lora_A.weight" if body.endswith(".lora_A.weight") else ".lora_B.weight" | |
| base = body.removesuffix(suffix) | |
| simple_map = { | |
| "img_in": "x_embedder", | |
| "txt_in": "context_embedder", | |
| "time_in.in_layer": "time_guidance_embed.timestep_embedder.linear_1", | |
| "time_in.out_layer": "time_guidance_embed.timestep_embedder.linear_2", | |
| "guidance_in.in_layer": "time_guidance_embed.guidance_embedder.linear_1", | |
| "guidance_in.out_layer": "time_guidance_embed.guidance_embedder.linear_2", | |
| "double_stream_modulation_img.lin": "double_stream_modulation_img.linear", | |
| "double_stream_modulation_txt.lin": "double_stream_modulation_txt.linear", | |
| "single_stream_modulation.lin": "single_stream_modulation.linear", | |
| "final_layer.linear": "proj_out", | |
| } | |
| if base in simple_map: | |
| return {f"transformer.{simple_map[base]}{suffix}": tensor} | |
| double_match = re.fullmatch(r"double_blocks\.(\d+)\.(img_attn|txt_attn)\.(qkv|proj)", base) | |
| if double_match: | |
| block, stream, layer = double_match.groups() | |
| stem = f"transformer.transformer_blocks.{block}.attn" | |
| if layer == "proj": | |
| target = "to_out.0" if stream == "img_attn" else "to_add_out" | |
| return {f"{stem}.{target}{suffix}": tensor} | |
| targets = ( | |
| ("to_q", "to_k", "to_v") | |
| if stream == "img_attn" | |
| else ("add_q_proj", "add_k_proj", "add_v_proj") | |
| ) | |
| if suffix == ".lora_A.weight": | |
| return {f"{stem}.{target}{suffix}": tensor.clone() for target in targets} | |
| chunks = tensor.chunk(3, dim=0) | |
| return {f"{stem}.{target}{suffix}": chunk.contiguous() for target, chunk in zip(targets, chunks)} | |
| single_match = re.fullmatch(r"single_blocks\.(\d+)\.(linear1|linear2)", base) | |
| if single_match: | |
| block, layer = single_match.groups() | |
| target = "to_qkv_mlp_proj" if layer == "linear1" else "to_out" | |
| return {f"transformer.single_transformer_blocks.{block}.attn.{target}{suffix}": tensor} | |
| raise ValueError(f"Unsupported fal LoRA key: {key}") | |
| def convert_fal_lora_to_diffusers(input_path: Path, output_path: Path) -> dict[str, Any]: | |
| state = load_file(input_path) | |
| converted: dict[str, torch.Tensor] = {} | |
| for key, tensor in state.items(): | |
| for new_key, new_tensor in convert_fal_key(key, tensor).items(): | |
| if new_key in converted: | |
| raise ValueError(f"Duplicate converted LoRA key: {new_key}") | |
| converted[new_key] = new_tensor | |
| output_path.parent.mkdir(parents=True, exist_ok=True) | |
| save_file(converted, output_path, metadata={"format": "pt"}) | |
| return { | |
| "input_keys": len(state), | |
| "converted_keys": len(converted), | |
| "input_bytes": input_path.stat().st_size, | |
| "converted_bytes": output_path.stat().st_size, | |
| } | |
| def expected_linear_shapes() -> dict[str, tuple[int, ...]]: | |
| from accelerate import init_empty_weights | |
| from diffusers import Flux2Transformer2DModel | |
| with init_empty_weights(): | |
| model = Flux2Transformer2DModel() | |
| return { | |
| f"transformer.{name}": tuple(module.weight.shape) | |
| for name, module in model.named_modules() | |
| if module.__class__.__name__ == "Linear" | |
| } | |
| def validate_converted_lora(path: Path) -> dict[str, Any]: | |
| state = load_file(path) | |
| shapes = expected_linear_shapes() | |
| missing_targets = [] | |
| bad_shapes = [] | |
| ranks = set() | |
| for key, tensor in state.items(): | |
| if key.endswith(".lora_A.weight"): | |
| target = key.removesuffix(".lora_A.weight") | |
| ranks.add(tensor.shape[0]) | |
| expected = shapes.get(target) | |
| if expected is None: | |
| missing_targets.append(target) | |
| elif tuple(tensor.shape[1:]) != (expected[1],): | |
| bad_shapes.append((key, tuple(tensor.shape), expected)) | |
| elif key.endswith(".lora_B.weight"): | |
| target = key.removesuffix(".lora_B.weight") | |
| ranks.add(tensor.shape[1]) | |
| expected = shapes.get(target) | |
| if expected is None: | |
| missing_targets.append(target) | |
| elif tuple(tensor.shape[:1]) != (expected[0],): | |
| bad_shapes.append((key, tuple(tensor.shape), expected)) | |
| else: | |
| missing_targets.append(key) | |
| return { | |
| "keys": len(state), | |
| "target_modules": len({key.rsplit(".lora_", 1)[0] for key in state}), | |
| "ranks": sorted(ranks), | |
| "missing_targets": sorted(set(missing_targets)), | |
| "bad_shapes": bad_shapes, | |
| "valid": not missing_targets and not bad_shapes, | |
| } | |
| def load_flux2_pipeline(args: argparse.Namespace, dtype: torch.dtype | str, device_name: str): | |
| from diffusers import Flux2Pipeline | |
| if uses_4bit_model(args.model) and device_name.startswith("cuda"): | |
| from diffusers import AutoModel | |
| from transformers import Mistral3ForConditionalGeneration | |
| print("Loading 4-bit FLUX.2 with local text encoder on CPU and model CPU offload.", flush=True) | |
| text_encoder = Mistral3ForConditionalGeneration.from_pretrained( | |
| args.model, | |
| subfolder="text_encoder", | |
| torch_dtype=dtype, | |
| device_map="cpu", | |
| local_files_only=args.local_files_only, | |
| ) | |
| transformer = AutoModel.from_pretrained( | |
| args.model, | |
| subfolder="transformer", | |
| torch_dtype=dtype, | |
| device_map="cpu", | |
| local_files_only=args.local_files_only, | |
| ) | |
| pipe = Flux2Pipeline.from_pretrained( | |
| args.model, | |
| text_encoder=text_encoder, | |
| transformer=transformer, | |
| torch_dtype=dtype, | |
| local_files_only=args.local_files_only, | |
| ) | |
| pipe.enable_model_cpu_offload() | |
| return pipe | |
| load_kwargs: dict[str, Any] = { | |
| "torch_dtype": dtype, | |
| "local_files_only": args.local_files_only, | |
| } | |
| if args.device_map is not None: | |
| load_kwargs["device_map"] = args.device_map | |
| elif device_name.startswith("cuda"): | |
| load_kwargs["device_map"] = device_name | |
| pipe = Flux2Pipeline.from_pretrained(args.model, **load_kwargs) | |
| if "device_map" not in load_kwargs: | |
| pipe.to(device_name) | |
| return pipe | |
| def batched(values: list[Path], batch_size: int) -> list[list[Path]]: | |
| return [values[index : index + batch_size] for index in range(0, len(values), batch_size)] | |
| def discover_input_images(input_dir: Path) -> list[Path]: | |
| return sorted( | |
| ( | |
| path | |
| for path in input_dir.iterdir() | |
| if path.is_file() and path.suffix.lower() in SUPPORTED_IMAGE_SUFFIXES | |
| ), | |
| key=lambda path: path.name.lower(), | |
| ) | |
| def output_paths_for_inputs(args: argparse.Namespace, input_images: list[Path]) -> list[Path]: | |
| if args.output_path is None: | |
| return [ | |
| args.output_dir / f"{input_path.stem}-flux2-local-stylized.png" | |
| for input_path in input_images | |
| ] | |
| if args.output_path.suffix: | |
| stem = args.output_path.with_suffix("") | |
| suffix = args.output_path.suffix | |
| return [ | |
| stem.with_name(f"{stem.name}-{index:04d}{suffix}") | |
| for index, _input_path in enumerate(input_images, start=1) | |
| ] | |
| return [ | |
| args.output_path / f"{input_path.stem}-flux2-local-stylized.png" | |
| for input_path in input_images | |
| ] | |
| def generators_for_batch( | |
| seed: int | None, | |
| device_name: str, | |
| *, | |
| start_index: int, | |
| batch_size: int, | |
| ) -> torch.Generator | list[torch.Generator] | None: | |
| if seed is None: | |
| return None | |
| if batch_size == 1: | |
| return torch.Generator(device=device_name).manual_seed(seed + start_index) | |
| return [ | |
| torch.Generator(device=device_name).manual_seed(seed + start_index + index) | |
| for index in range(batch_size) | |
| ] | |
| def run_inference(args: argparse.Namespace, lora_path: Path) -> list[Path]: | |
| from diffusers.utils import load_image | |
| device_name = args.device or ("cuda" if torch.cuda.is_available() else "cpu") | |
| dtype = dtype_from_arg(args.torch_dtype) | |
| input_paths = discover_input_images(args.input_dir.expanduser().resolve()) | |
| output_paths = output_paths_for_inputs(args, input_paths) | |
| print( | |
| f"Found {len(input_paths)} input image(s). Processing in batches of {args.batch_size}.", | |
| flush=True, | |
| ) | |
| print("Using text encoder mode: local", flush=True) | |
| pipe = load_flux2_pipeline(args, dtype, device_name) | |
| pipe.load_lora_weights(str(lora_path), adapter_name="aura") | |
| pipe.set_adapters(["aura"], adapter_weights=[args.lora_scale]) | |
| for start_index, batch_paths in enumerate(batched(input_paths, args.batch_size)): | |
| batch_offset = start_index * args.batch_size | |
| input_images = [load_image(str(input_path)) for input_path in batch_paths] | |
| image_arg: Any = input_images[0] if len(input_images) == 1 else input_images | |
| prompt_arg: Any = args.prompt if len(input_images) == 1 else [args.prompt] * len(batch_paths) | |
| call_kwargs: dict[str, Any] = { | |
| "image": image_arg, | |
| "height": args.height, | |
| "width": args.width, | |
| "num_inference_steps": args.num_inference_steps, | |
| "guidance_scale": args.guidance_scale, | |
| "generator": generators_for_batch( | |
| args.seed, | |
| device_name, | |
| start_index=batch_offset, | |
| batch_size=len(batch_paths), | |
| ), | |
| "prompt": prompt_arg, | |
| } | |
| images = pipe(**call_kwargs).images | |
| if len(images) != len(batch_paths): | |
| raise RuntimeError(f"Expected {len(batch_paths)} outputs from pipeline, received {len(images)}.") | |
| for image, output_path in zip(images, output_paths[batch_offset : batch_offset + len(images)]): | |
| output_path.parent.mkdir(parents=True, exist_ok=True) | |
| image.save(output_path) | |
| return output_paths | |
| def main() -> int: | |
| args = parse_args() | |
| lora_path = args.lora.expanduser().resolve() | |
| if not lora_path.exists(): | |
| print(f"LoRA file does not exist: {lora_path}", file=sys.stderr) | |
| return 1 | |
| if not args.check_only: | |
| if args.input_dir is None: | |
| print("input_dir is required unless --check-only is set.", file=sys.stderr) | |
| return 1 | |
| if args.batch_size < 1: | |
| print("--batch-size must be at least 1.", file=sys.stderr) | |
| return 1 | |
| input_dir = args.input_dir.expanduser().resolve() | |
| if not input_dir.exists(): | |
| print(f"Input directory does not exist: {input_dir}", file=sys.stderr) | |
| return 1 | |
| if not input_dir.is_dir(): | |
| print(f"Input path is not a directory: {input_dir}", file=sys.stderr) | |
| return 1 | |
| input_images = discover_input_images(input_dir) | |
| if not input_images: | |
| print( | |
| f"No supported images found in {input_dir}. " | |
| f"Supported extensions: {', '.join(sorted(SUPPORTED_IMAGE_SUFFIXES))}.", | |
| file=sys.stderr, | |
| ) | |
| return 1 | |
| output_dir = args.output_dir.expanduser().resolve() | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| converted_path = ( | |
| args.converted_lora.expanduser().resolve() | |
| if args.converted_lora | |
| else output_dir / "pytorch_lora_weights.diffusers.safetensors" | |
| ) | |
| try: | |
| conversion = convert_fal_lora_to_diffusers(lora_path, converted_path) | |
| validation = validate_converted_lora(converted_path) | |
| except Exception as exc: | |
| print(str(exc), file=sys.stderr) | |
| return 1 | |
| report: dict[str, Any] = { | |
| "model": args.model, | |
| "text_encoder_mode": "local", | |
| "lora": str(lora_path), | |
| "converted_lora": str(converted_path), | |
| "conversion": conversion, | |
| "validation": validation, | |
| } | |
| print(json.dumps(report, indent=2, default=str)) | |
| if not validation["valid"]: | |
| print("Converted LoRA did not validate against Flux2Transformer2DModel.", file=sys.stderr) | |
| return 1 | |
| if args.check_only: | |
| return 0 | |
| try: | |
| preflight_environment(args) | |
| except RuntimeError as exc: | |
| print(str(exc), file=sys.stderr) | |
| return 1 | |
| started_at = time.time() | |
| try: | |
| image_paths = run_inference(args, converted_path) | |
| except Exception as exc: | |
| print(f"Local inference failed after {time.time() - started_at:.1f}s: {exc}", file=sys.stderr) | |
| return 1 | |
| for image_path in image_paths: | |
| print(f"Saved local FLUX.2 edit output: {image_path}") | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |