Files
ycyy-ComfyUI-YCYY-API/modelscope/modelscope_image_node.py
T

387 lines
13 KiB
Python

import json
import time
from io import BytesIO
from typing import Dict, List, Optional, Tuple
import requests
import torch
from PIL import Image
from comfy_api.latest import io
try:
from ..utils.config_utils import (
DEFAULT_MODELSCOPE_IMAGE_MODELS,
get_config_section,
get_modelscope_image_api_config,
get_modelscope_image_apis,
)
from ..utils.image_utils import pil_to_tensor, tensor_to_base64_string
except (ImportError, ValueError):
from utils.config_utils import (
DEFAULT_MODELSCOPE_IMAGE_MODELS,
get_config_section,
get_modelscope_image_api_config,
get_modelscope_image_apis,
)
from utils.image_utils import pil_to_tensor, tensor_to_base64_string
try:
from aiohttp import web
from server import PromptServer
@PromptServer.instance.routes.get("/ycyy/modelscope-image/apis/all")
async def get_all_modelscope_image_apis(request):
try:
return web.json_response([
{"api-name": item["api-name"], "models": item["models"]}
for item in get_modelscope_image_apis()
])
except Exception as exc:
return web.json_response({"error": str(exc)}, status=500)
except Exception:
pass
DEFAULT_MODELS = list(DEFAULT_MODELSCOPE_IMAGE_MODELS)
class ModelScopeImage(io.ComfyNode):
"""Generate or edit an image through the asynchronous ModelScope API."""
@classmethod
def _load_models_from_config(cls, api_name: Optional[str] = None) -> List[str]:
try:
apis = get_modelscope_image_apis()
if api_name:
for item in apis:
if item["api-name"] == api_name:
return item.get("models") or list(DEFAULT_MODELS)
models = list(dict.fromkeys(
model for item in apis for model in item.get("models", [])
))
return models or list(DEFAULT_MODELS)
except Exception:
return list(DEFAULT_MODELS)
@classmethod
def _load_config_credentials(
cls,
api_name: Optional[str] = None,
config_options: Optional[dict] = None,
) -> Tuple[str, str, int]:
try:
api_config = get_modelscope_image_api_config(api_name)
except Exception:
apis = get_modelscope_image_apis()
api_config = apis[0] if apis else {
"base_url": "https://api-inference.modelscope.cn",
"api_key": "",
"timeout": 300,
}
base_url = str(
api_config.get("base_url") or "https://api-inference.modelscope.cn"
).strip()
api_key = str(api_config.get("api_key") or "").strip()
timeout = api_config.get("timeout", 300)
config_options = config_options or {}
override_base_url = str(config_options.get("base_url") or "").strip()
override_api_key = str(config_options.get("api_key") or "").strip()
if override_base_url:
base_url = override_base_url
if override_api_key:
api_key = override_api_key
if config_options.get("timeout"):
timeout = config_options["timeout"]
try:
timeout = int(timeout)
except (TypeError, ValueError):
timeout = 300
if timeout <= 0:
timeout = 300
if not base_url:
raise ValueError("ModelScope base_url cannot be empty")
if not api_key:
raise ValueError(
"ModelScope API key not found. Please provide an api_key in "
"config.json ('modelscope-image') or via API Config Options."
)
return base_url, api_key, timeout
@classmethod
def _get_proxy_config(
cls, proxy_options: Optional[dict] = None
) -> Optional[Dict[str, str]]:
if proxy_options is not None:
if not proxy_options.get("enable", False):
return None
proxies = {
key: proxy_options[key].strip()
for key in ("http", "https")
if isinstance(proxy_options.get(key), str)
and proxy_options[key].strip()
}
return proxies or None
try:
proxy_config = get_config_section("proxy") or {}
if not proxy_config.get("enable", False):
return None
proxies = {
key: proxy_config[key]
for key in ("http", "https")
if proxy_config.get(key)
}
return proxies or None
except Exception:
return None
@classmethod
def define_schema(cls) -> io.Schema:
try:
apis = get_modelscope_image_apis()
api_names = [item["api-name"] for item in apis]
models = list(dict.fromkeys(
model for item in apis for model in item.get("models", [])
))
except Exception:
api_names = ["default"]
models = list(DEFAULT_MODELS)
if not api_names:
api_names = ["default"]
if not models:
models = list(DEFAULT_MODELS)
return io.Schema(
node_id="YCYY_ModelScope_Image_API",
display_name="ModelScope Image API",
category="YCYY/API/image",
inputs=[
io.String.Input(
id="prompt",
multiline=True,
tooltip="Prompt used to generate or edit the image.",
),
io.String.Input(
id="negative_prompt",
multiline=True,
tooltip="Negative prompt.",
),
io.Combo.Input(
id="api_name",
options=api_names,
default=api_names[0],
tooltip="Select a ModelScope image API name.",
),
io.Combo.Input(
id="model",
options=models,
default=models[0],
tooltip="Select a model from the chosen API name.",
),
io.Int.Input(id="width", min=64, max=2048, default=1024, step=8),
io.Int.Input(id="height", min=64, max=2048, default=1024, step=8),
io.Int.Input(id="steps", min=1, max=100, default=30, step=1),
io.Float.Input(
id="guidance", min=1.5, max=20, default=3.5, step=0.1
),
io.Int.Input(
id="seed",
min=0,
max=2147483647,
default=0,
control_after_generate=True,
),
io.Image.Input(
id="image",
optional=True,
tooltip="Optional source image. Connect it to edit an image; leave it disconnected to generate one.",
),
io.AnyType.Input(
id="config_options",
optional=True,
tooltip="Optional configuration override from YCYY API Config Options.",
),
io.AnyType.Input(
id="proxy_options",
optional=True,
tooltip="Optional proxy configuration override from YCYY API Proxy Options.",
),
],
outputs=[io.Image.Output(), io.String.Output()],
description="Generate or edit images through the ModelScope API. Connecting an image enables edit mode.",
)
@classmethod
def execute(
cls,
prompt,
negative_prompt,
model,
width,
height,
steps,
guidance,
seed,
api_name=None,
image=None,
config_options=None,
proxy_options=None,
) -> io.NodeOutput:
if not prompt or not prompt.strip():
raise ValueError("prompt cannot be empty")
base_url, api_key, timeout = cls._load_config_credentials(
api_name=api_name, config_options=config_options
)
return cls._request_image(
base_url=base_url,
api_key=api_key,
prompt=prompt,
negative_prompt=negative_prompt,
model=model,
width=width,
height=height,
steps=steps,
guidance=guidance,
seed=seed,
image=image,
timeout=timeout,
proxies=cls._get_proxy_config(proxy_options),
)
@classmethod
def _request_image(
cls,
base_url,
api_key,
prompt,
negative_prompt,
model,
width,
height,
steps,
guidance,
seed,
image,
timeout,
proxies,
) -> io.NodeOutput:
service_root = cls._normalize_base_url(base_url)
api_url = f"{service_root}/v1/images/generations"
headers = {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
"X-ModelScope-Async-Mode": "true",
}
mode = "edit" if image is not None else "generation"
payload = {
"model": model,
"prompt": prompt,
"size": f"{width}x{height}",
"steps": steps,
"guidance": guidance,
"seed": seed,
}
if negative_prompt:
payload["negative_prompt"] = negative_prompt
if image is not None:
image_data = tensor_to_base64_string(image)
payload["image_url"] = f"data:image/png;base64,{image_data}"
try:
response = requests.post(
api_url,
headers=headers,
json=payload,
timeout=timeout,
proxies=proxies,
)
if response.status_code != 200:
raise RuntimeError(f"HTTP {response.status_code}: {response.text}")
task_id = response.json().get("task_id")
if not task_id:
raise RuntimeError("ModelScope response did not contain task_id")
output_image_url, task_data = cls._wait_for_task(
service_root, api_key, task_id, timeout, proxies
)
output_response = requests.get(
output_image_url, timeout=timeout, proxies=proxies
)
output_response.raise_for_status()
result_image = Image.open(BytesIO(output_response.content)).convert("RGB")
result_info = {
"success": True,
"message": f"Image {mode} success.",
"mode": mode,
"model": model,
"task_id": task_id,
"image_url": output_image_url,
"task": task_data,
}
return io.NodeOutput(
pil_to_tensor(result_image),
json.dumps(result_info, ensure_ascii=False),
)
except Exception as error:
raise RuntimeError(
json.dumps(
{
"success": False,
"mode": mode,
"message": f"ModelScope image {mode} failed: {error}",
},
ensure_ascii=False,
)
) from error
@classmethod
def _wait_for_task(cls, service_root, api_key, task_id, timeout, proxies):
deadline = time.monotonic() + timeout
while time.monotonic() < deadline:
response = requests.get(
f"{service_root}/v1/tasks/{task_id}",
headers={
"Authorization": f"Bearer {api_key}",
"X-ModelScope-Task-Type": "image_generation",
},
timeout=timeout,
proxies=proxies,
)
if response.status_code != 200:
raise RuntimeError(
f"Task query HTTP {response.status_code}: {response.text}"
)
data = response.json()
status = data.get("task_status")
if status == "SUCCEED":
output_images = data.get("output_images") or []
if output_images:
return output_images[0], data
raise RuntimeError("Task succeeded without output image")
if status == "FAILED":
raise RuntimeError(data.get("message") or "Image task failed")
remaining = deadline - time.monotonic()
if remaining > 0:
time.sleep(min(5, remaining))
raise TimeoutError("Timed out waiting for ModelScope image task")
@classmethod
def _normalize_base_url(cls, base_url: str) -> str:
clean_url = base_url.strip().rstrip("/")
for suffix in ("/v1/images/generations", "/v1/images", "/v1"):
if clean_url.endswith(suffix):
clean_url = clean_url[:-len(suffix)]
break
if not clean_url:
raise ValueError("ModelScope base_url cannot be empty")
return clean_url
@classmethod
def _create_empty_image(cls):
return torch.zeros(1, 512, 512, 3, dtype=torch.float32)