🪫 Battery Saver - manage power over long batch operations (#189)

Co-authored-by: Anthony Wu <pls-file-gh-issue@users.noreply.github.com>
Co-authored-by: filipstrand <strand.filip@gmail.com>
This commit is contained in:
Anthony Wu 2025-05-19 08:13:41 -07:00 committed by GitHub
parent 4b56633e6c
commit d20f2833da
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
12 changed files with 339 additions and 82 deletions

View File

@ -199,6 +199,8 @@ mflux-generate --model dev --prompt "Luxury food photograph" --steps 25 --seed 2
- **`--low-ram`** (optional): Reduces GPU memory usage by limiting the MLX cache size and releasing text encoders and transformer components after use (single image generation only). While this may slightly decrease performance, it helps prevent system memory swapping to disk, allowing image generation on systems with limited RAM.
- **`--battery-percentage-stop-limit`** or **`-B`** (optional, `int`, default: `5`): On Mac laptops powered by battery, automatically stops image generation when battery percentage reaches this threshold. Prevents your Mac from shutting down and becoming unresponsive during long generation sessions.
- **`--lora-name`** (optional, `str`, default: `None`): The name of the LoRA to download from Hugging Face.
- **`--lora-repo-id`** (optional, `str`, default: `"ali-vilab/In-Context-LoRA"`): The Hugging Face repository ID for LoRAs.
@ -1275,6 +1277,7 @@ See `uv run tools/rename_images.py --help` for full CLI usage help.
- shortcut for dev model: `alias mflux-dev='mflux-generate --model dev'`
- shortcut for schnell model *and* always save metadata: `alias mflux-schnell='mflux-generate --model schnell --metadata'`
- For systems with limited memory, use the `--low-ram` flag to reduce memory usage by constraining the MLX cache size and releasing components after use
- On battery-powered Macs, use `--battery-percentage-stop-limit` (or `-B`) to prevent your laptop from shutting down during long generation sessions
- When generating multiple images with different seeds, use `--seed` with multiple values or `--auto-seeds` to automatically generate a series of random seeds
- Use `--stepwise-image-output-dir` to save intermediate images at each denoising step, which can be useful for debugging or creating animations of the generation process

View File

@ -0,0 +1,63 @@
import json
import logging
import re
import subprocess
from mflux.callbacks.callback import BeforeLoopCallback
from mflux.error.exceptions import StopImageGenerationException
PMSET_AC_POWER_STATUS = "Now drawing from 'AC Power'"
PMSET_BATT_STATUS_PATTERN = r"InternalBattery-.+?(\d+)%"
logger = logging.getLogger(__name__)
def _get_machine_model() -> str:
"""Get the Mac machine model using system_profiler."""
try:
result = subprocess.run(
["system_profiler", "-json", "SPHardwareDataType"], capture_output=True, text=True, check=True
)
data = json.loads(result.stdout)
return data["SPHardwareDataType"][0]["machine_model"]
except (subprocess.CalledProcessError, json.JSONDecodeError, IndexError, KeyError) as e:
logger.warning(f"Cannot determine machine model via 'system_profiler -json SPHardwareDataType': {e}")
return "Unknown"
MACHINE_MODEL = _get_machine_model()
# assumption: all Apple Silicon models powered by battery are "MacBook"s
MACHINE_IS_BATTERY_POWERED = "MacBook" in MACHINE_MODEL
def get_battery_percentage() -> int | None:
"""Get the current battery percentage of a battery-powered Mac.
Returns None if Mac is not a battery-powered machine."""
if not MACHINE_IS_BATTERY_POWERED:
return None
percentage = None
try:
# running the subprocess would be expensive in a tight loop
# but in mflux use case, we would call this only once every
# few minutes due to N-minutes-long generation times
result = subprocess.run(["pmset", "-g", "batt"], capture_output=True, text=True, check=True)
if PMSET_AC_POWER_STATUS not in result.stdout:
if match := re.search(PMSET_BATT_STATUS_PATTERN, result.stdout):
percentage = int(match.group(1))
except (subprocess.CalledProcessError, TypeError) as e:
logger.warning(
f"Cannot read battery percentage via 'pmset -g batt': {e}. Battery saver functionality is disabled and the program will continue running."
)
return percentage
class BatterySaver(BeforeLoopCallback):
def __init__(self, battery_percentage_stop_limit=10):
self.limit = battery_percentage_stop_limit
def call_before_loop(self, **kwargs) -> None:
current_pct: int | None = get_battery_percentage()
if current_pct is not None and current_pct <= self.limit:
raise StopImageGenerationException(f"Battery below {self.limit}% threshold: {current_pct}%")

View File

@ -1,5 +1,8 @@
from argparse import Namespace
from mflux import Config, Flux1, ModelConfig, StopImageGenerationException
from mflux.callbacks.callback_registry import CallbackRegistry
from mflux.callbacks.instances.battery_saver import BatterySaver
from mflux.callbacks.instances.memory_saver import MemorySaver
from mflux.callbacks.instances.stepwise_handler import StepwiseHandler
from mflux.error.exceptions import PromptFileReadError
@ -27,19 +30,8 @@ def main():
lora_scales=args.lora_scales,
)
# 2. Register the optional callbacks
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
# 2. Register callbacks
memory_saver = _register_callbacks(args=args, flux=flux)
try:
for seed in args.seed:
@ -65,5 +57,27 @@ def main():
print(memory_saver.memory_stats())
def _register_callbacks(args: Namespace, flux: Flux1) -> MemorySaver | None:
# Battery saver
battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit)
CallbackRegistry.register_before_loop(battery_saver)
# Stepwise Handler
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
# Memory Saver
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
return memory_saver
if __name__ == "__main__":
main()

View File

@ -1,5 +1,8 @@
from argparse import Namespace
from mflux import Config, Flux1Controlnet, ModelConfig, StopImageGenerationException
from mflux.callbacks.callback_registry import CallbackRegistry
from mflux.callbacks.instances.battery_saver import BatterySaver
from mflux.callbacks.instances.canny_saver import CannyImageSaver
from mflux.callbacks.instances.memory_saver import MemorySaver
from mflux.callbacks.instances.stepwise_handler import StepwiseHandler
@ -28,21 +31,8 @@ def main():
lora_scales=args.lora_scales,
)
# 2. Register the optional callbacks
if args.controlnet_save_canny:
CallbackRegistry.register_before_loop(CannyImageSaver(path=args.output))
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
# 2. Register callbacks
memory_saver = _register_callbacks(args=args, flux=flux)
try:
for seed in args.seed:
@ -69,5 +59,32 @@ def main():
print(memory_saver.memory_stats())
def _register_callbacks(args: Namespace, flux: Flux1Controlnet) -> MemorySaver | None:
# Battery saver
battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit)
CallbackRegistry.register_before_loop(battery_saver)
# Canny Image Saver
if args.controlnet_save_canny:
canny_image_saver = CannyImageSaver(path=args.output)
CallbackRegistry.register_before_loop(canny_image_saver)
# Stepwise Handler
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
# Memory Saver
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
return memory_saver
if __name__ == "__main__":
main()

View File

@ -1,5 +1,8 @@
from argparse import Namespace
from mflux import Config, StopImageGenerationException
from mflux.callbacks.callback_registry import CallbackRegistry
from mflux.callbacks.instances.battery_saver import BatterySaver
from mflux.callbacks.instances.depth_saver import DepthImageSaver
from mflux.callbacks.instances.memory_saver import MemorySaver
from mflux.callbacks.instances.stepwise_handler import StepwiseHandler
@ -32,21 +35,8 @@ def main():
lora_scales=args.lora_scales,
)
# 2. Register the optional callbacks
if args.save_depth_map:
CallbackRegistry.register_before_loop(DepthImageSaver(path=args.output))
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
# 2. Register callbacks
memory_saver = _register_callbacks(args=args, flux=flux)
try:
for seed in args.seed:
@ -73,5 +63,32 @@ def main():
print(memory_saver.memory_stats())
def _register_callbacks(args: Namespace, flux: Flux1Depth) -> MemorySaver | None:
# Battery saver
battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit)
CallbackRegistry.register_before_loop(battery_saver)
# Depth Image Saver
if args.save_depth_map:
depth_image_saver = DepthImageSaver(path=args.output)
CallbackRegistry.register_before_loop(depth_image_saver)
# Stepwise Handler
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
# Memory Saver
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
return memory_saver
if __name__ == "__main__":
main()

View File

@ -1,5 +1,8 @@
from argparse import Namespace
from mflux import Config, StopImageGenerationException
from mflux.callbacks.callback_registry import CallbackRegistry
from mflux.callbacks.instances.battery_saver import BatterySaver
from mflux.callbacks.instances.memory_saver import MemorySaver
from mflux.callbacks.instances.stepwise_handler import StepwiseHandler
from mflux.error.exceptions import PromptFileReadError
@ -31,19 +34,8 @@ def main():
lora_scales=args.lora_scales,
)
# 2. Register the optional callbacks
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
# 2. Register callbacks
memory_saver = _register_callbacks(args=args, flux=flux)
try:
for seed in args.seed:
@ -70,5 +62,27 @@ def main():
print(memory_saver.memory_stats())
def _register_callbacks(args: Namespace, flux: Flux1Fill) -> MemorySaver | None:
# Battery saver
battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit)
CallbackRegistry.register_before_loop(battery_saver)
# Stepwise Handler
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
# Memory Saver
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
return memory_saver
if __name__ == "__main__":
main()

View File

@ -1,7 +1,9 @@
from argparse import Namespace
from pathlib import Path
from mflux import Config, StopImageGenerationException
from mflux.callbacks.callback_registry import CallbackRegistry
from mflux.callbacks.instances.battery_saver import BatterySaver
from mflux.callbacks.instances.memory_saver import MemorySaver
from mflux.callbacks.instances.stepwise_handler import StepwiseHandler
from mflux.community.in_context_lora.flux_in_context_lora import Flux1InContextLoRA
@ -39,19 +41,8 @@ def main():
lora_scales=args.lora_scales,
)
# 2. Register the optional callbacks
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
# 2. Register callbacks
memory_saver = _register_callbacks(args=args, flux=flux)
try:
for seed in args.seed:
@ -79,5 +70,27 @@ def main():
print(memory_saver.memory_stats())
def _register_callbacks(args: Namespace, flux: Flux1InContextLoRA) -> MemorySaver | None:
# Battery saver
battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit)
CallbackRegistry.register_before_loop(battery_saver)
# Stepwise Handler
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
# Memory Saver
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
return memory_saver
if __name__ == "__main__":
main()

View File

@ -1,7 +1,10 @@
from argparse import Namespace
from pathlib import Path
from mflux import Config, ModelConfig, StopImageGenerationException
from mflux.callbacks.callback_registry import CallbackRegistry
from mflux.callbacks.instances.battery_saver import BatterySaver
from mflux.callbacks.instances.memory_saver import MemorySaver
from mflux.callbacks.instances.stepwise_handler import StepwiseHandler
from mflux.error.exceptions import PromptFileReadError
@ -36,19 +39,8 @@ def main():
lora_scales=args.lora_scales,
)
# 2. Register the optional callbacks
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
# 2. Register callbacks
memory_saver = _register_callbacks(args=args, flux=flux)
try:
for seed in args.seed:
@ -75,6 +67,28 @@ def main():
print(memory_saver.memory_stats())
def _register_callbacks(args: Namespace, flux: Flux1Redux) -> MemorySaver | None:
# Battery saver
battery_saver = BatterySaver(battery_percentage_stop_limit=args.battery_percentage_stop_limit)
CallbackRegistry.register_before_loop(battery_saver)
# Stepwise Handler
if args.stepwise_image_output_dir:
handler = StepwiseHandler(flux=flux, output_dir=args.stepwise_image_output_dir)
CallbackRegistry.register_before_loop(handler)
CallbackRegistry.register_in_loop(handler)
CallbackRegistry.register_interrupt(handler)
# Memory Saver
memory_saver = None
if args.low_ram:
memory_saver = MemorySaver(flux=flux, keep_transformer=len(args.seed) > 1)
CallbackRegistry.register_before_loop(memory_saver)
CallbackRegistry.register_in_loop(memory_saver)
CallbackRegistry.register_after_loop(memory_saver)
return memory_saver
def _validate_redux_image_strengths(
redux_image_paths: list[Path],
redux_image_strengths: list[float] | None,

View File

@ -42,6 +42,7 @@ class CommandLineParser(argparse.ArgumentParser):
self.require_model_arg = True
def add_general_arguments(self) -> None:
self.add_argument("--battery-percentage-stop-limit", "-B", type=lambda v: max(min(int(v), 99), 1), default=ui_defaults.BATTERY_PERCENTAGE_STOP_LIMIT, help=f"On Macs powered by battery, stop image generation when battery reaches this percentage. Default: {ui_defaults.BATTERY_PERCENTAGE_STOP_LIMIT}")
self.add_argument("--low-ram", action="store_true", help="Enable low-RAM mode to reduce memory usage (may impact performance).")
def add_model_arguments(self, path_type: t.Literal["load", "save"] = "load", require_model_arg: bool = True) -> None:

View File

@ -1,3 +1,4 @@
BATTERY_PERCENTAGE_STOP_LIMIT = 5
CONTROLNET_STRENGTH = 0.4
GUIDANCE_SCALE = 3.5
HEIGHT, WIDTH = 1024, 1024

View File

View File

@ -0,0 +1,100 @@
from unittest.mock import MagicMock, patch
import pytest
from mflux.callbacks.instances.battery_saver import BatterySaver, get_battery_percentage
from mflux.error.exceptions import StopImageGenerationException
def test_get_battery_percentage_success():
"""Test that the battery percentage is correctly extracted when subprocess returns a valid output."""
with patch("subprocess.run") as mock_run:
# Set up mock to return a sample battery status output
mock_result = MagicMock()
mock_result.stdout = "Now drawing from 'Battery Power'\n -InternalBattery-0 (id=1234567) 42%;\n"
mock_run.return_value = mock_result
# Call the function
percentage = get_battery_percentage()
# Verify subprocess.run was called with the correct arguments
mock_run.assert_called_once_with(["pmset", "-g", "batt"], capture_output=True, text=True, check=True)
# Assert the function properly extracted the battery percentage
assert percentage == 42
def test_get_battery_percentage_while_charging():
"""Test that the function returns None when the output doesn't match the expected pattern."""
with patch("subprocess.run") as mock_run:
# Set up mock to return an output that doesn't match the expected pattern
mock_result = MagicMock()
mock_result.stdout = "Now drawing from 'AC Power'"
mock_run.return_value = mock_result
# Call the function
percentage = get_battery_percentage()
# Assert the function returns None when no match is found
assert percentage is None
def test_battery_saver_below_limit():
"""Test that BatterySaver raises an exception when the battery is below the limit."""
with patch("mflux.callbacks.instances.battery_saver.get_battery_percentage") as mock_get:
# Configure mock to return a battery percentage below the limit
mock_get.return_value = 5
# Create a BatterySaver instance with a limit of 10%
battery_saver = BatterySaver(battery_percentage_stop_limit=10)
# Assert that calling before_loop raises StopImageGenerationException
with pytest.raises(StopImageGenerationException) as excinfo:
battery_saver.call_before_loop()
# Check the exception message contains the correct limit and percentage
assert "10%" in str(excinfo.value)
assert "5%" in str(excinfo.value)
def test_battery_saver_above_limit():
"""Test that BatterySaver does not raise an exception when the battery is above the limit."""
with patch("mflux.callbacks.instances.battery_saver.get_battery_percentage") as mock_get:
# Configure mock to return a battery percentage above the limit
mock_get.return_value = 20
# Create a BatterySaver instance with a limit of 10%
battery_saver = BatterySaver(battery_percentage_stop_limit=10)
# Assert that calling before_loop does not raise an exception
battery_saver.call_before_loop()
def test_battery_saver_none_percentage():
"""Test that BatterySaver does not raise an exception when percentage is None."""
with patch("mflux.callbacks.instances.battery_saver.get_battery_percentage") as mock_get:
# Configure mock to return None (e.g., on non-battery systems)
mock_get.return_value = None
# Create a BatterySaver instance
battery_saver = BatterySaver()
# Assert that calling before_loop does not raise an exception when percentage is None
battery_saver.call_before_loop()
def test_battery_saver_equal_to_limit():
"""Test that BatterySaver raises an exception when the battery equals the limit."""
with patch("mflux.callbacks.instances.battery_saver.get_battery_percentage") as mock_get:
# Configure mock to return a battery percentage equal to the limit
mock_get.return_value = 10
# Create a BatterySaver instance with a limit of 10%
battery_saver = BatterySaver(battery_percentage_stop_limit=10)
# Assert that calling before_loop raises StopImageGenerationException
with pytest.raises(StopImageGenerationException) as excinfo:
battery_saver.call_before_loop()
# Check the exception message contains the correct limit and percentage
assert "10%" in str(excinfo.value)