Add explicit HYPIR restoration profile
This commit is contained in:
@@ -34,6 +34,7 @@ _SIZES = {
|
||||
}
|
||||
_DEFAULT_BASE_URL = "http://192.168.1.212:8081/v1"
|
||||
_DEFAULT_MODEL = "FLUX.2-klein-9B-fp8-beta"
|
||||
_DEFAULT_RESTORATION_MODEL = "HYPIR-SD2"
|
||||
_MAX_IMAGE_BYTES = 20 * 1024 * 1024
|
||||
|
||||
|
||||
@@ -47,6 +48,34 @@ def _model() -> str:
|
||||
return os.environ.get("ATHENA_IMAGE_MODEL", "").strip() or _DEFAULT_MODEL
|
||||
|
||||
|
||||
def _restoration_model() -> str:
|
||||
return (os.environ.get("ATHENA_RESTORATION_MODEL", "").strip()
|
||||
or _DEFAULT_RESTORATION_MODEL)
|
||||
|
||||
|
||||
def _restoration_profile_selected() -> bool:
|
||||
"""Use HYPIR only when Hermes explicitly selected the restoration model."""
|
||||
try:
|
||||
from hermes_cli.config import load_config
|
||||
config = load_config()
|
||||
model = config.get("model") if isinstance(config, dict) else None
|
||||
selected = model.get("default") if isinstance(model, dict) else None
|
||||
return isinstance(selected, str) and selected.casefold() in {
|
||||
"restauration", "restoration", "qwen-restoration",
|
||||
}
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _select_model(requested: object) -> str:
|
||||
"""Resolve an explicit image/profile selection; never inspect prompt text."""
|
||||
if _restoration_profile_selected():
|
||||
return _restoration_model()
|
||||
if isinstance(requested, str) and requested.strip() == _restoration_model():
|
||||
return _restoration_model()
|
||||
return _model()
|
||||
|
||||
|
||||
def _api_key() -> str:
|
||||
"""Prefer a scoped key; accept the existing router key for migration."""
|
||||
return (
|
||||
@@ -116,6 +145,12 @@ class AthenaLocalImageProvider(ImageGenProvider):
|
||||
"speed": "local",
|
||||
"strengths": "Private local generation and multi-reference editing",
|
||||
"price": "local / no cloud",
|
||||
}, {
|
||||
"id": _restoration_model(),
|
||||
"display": "HYPIR-SD2 Restoration on Athena",
|
||||
"speed": "local",
|
||||
"strengths": "Faithful denoise, deblur and photo restoration",
|
||||
"price": "local / no cloud",
|
||||
}]
|
||||
|
||||
def default_model(self) -> Optional[str]:
|
||||
@@ -178,6 +213,7 @@ class AthenaLocalImageProvider(ImageGenProvider):
|
||||
sources.append(image_url.strip())
|
||||
sources.extend(normalize_reference_images(reference_image_urls) or [])
|
||||
sources = sources[:4]
|
||||
model = _select_model(kwargs.get("model"))
|
||||
try:
|
||||
encoded_sources = [
|
||||
base64.b64encode(_load_private_image(source)).decode("ascii")
|
||||
@@ -204,6 +240,9 @@ class AthenaLocalImageProvider(ImageGenProvider):
|
||||
endpoint = "edits"
|
||||
request_data["image_b64"] = encoded_sources[0]
|
||||
request_data["reference_images_b64"] = encoded_sources[1:]
|
||||
if model == _restoration_model():
|
||||
request_data.update({"upscale": 1, "patch_size": 512,
|
||||
"stride": 256})
|
||||
request = urllib.request.Request(
|
||||
f"{base_url}/images/{endpoint}",
|
||||
data=json.dumps(request_data).encode("utf-8"), method="POST",
|
||||
@@ -241,7 +280,9 @@ class AthenaLocalImageProvider(ImageGenProvider):
|
||||
prompt=clean_prompt, aspect_ratio=aspect)
|
||||
|
||||
try:
|
||||
saved = save_b64_image(b64_data, prefix="athena_flux2")
|
||||
prefix = ("athena_hypir" if model == _restoration_model()
|
||||
else "athena_flux2")
|
||||
saved = save_b64_image(b64_data, prefix=prefix)
|
||||
except Exception as exc:
|
||||
return error_response(
|
||||
error=f"Generated image could not be saved: {exc}",
|
||||
|
||||
Reference in New Issue
Block a user