feat: V2/V2批量节点新增图片质量参数,日常走webp压缩,支持细分耗时统计
- 新增"图片质量"参数(日常/高清),日常时API请求Body携带image_compression=webp - V2批量节点同步该功能,与"图片格式"参数独立共存 - V2节点新增请求耗时/下载耗时/总耗时细分统计 - V2节点本地输出根据图片质量自动做格式转换(日常→JPEG,高清→PNG)
This commit is contained in:
@@ -72,6 +72,7 @@ class GeminiAsyncImageProvider(BaseAsyncImageProvider):
|
|||||||
resolution=resolution,
|
resolution=resolution,
|
||||||
enable_grounding=kwargs.get("enable_grounding", False),
|
enable_grounding=kwargs.get("enable_grounding", False),
|
||||||
enable_image_search=kwargs.get("enable_image_search", False),
|
enable_image_search=kwargs.get("enable_image_search", False),
|
||||||
|
image_compression=getattr(self, 'image_compression', None),
|
||||||
)
|
)
|
||||||
|
|
||||||
def extract_task_id(self, response: dict) -> str:
|
def extract_task_id(self, response: dict) -> str:
|
||||||
|
|||||||
@@ -186,6 +186,7 @@ class GeminiAPIClient(BaseAPIClient):
|
|||||||
resolution: str = "2K",
|
resolution: str = "2K",
|
||||||
enable_grounding: bool = False,
|
enable_grounding: bool = False,
|
||||||
enable_image_search: bool = False,
|
enable_image_search: bool = False,
|
||||||
|
image_compression: str = None,
|
||||||
**kwargs
|
**kwargs
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
@@ -282,6 +283,10 @@ class GeminiAPIClient(BaseAPIClient):
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# 添加图片压缩参数
|
||||||
|
if image_compression:
|
||||||
|
request_body["image_compression"] = image_compression
|
||||||
|
|
||||||
# 添加 Google Search Grounding 工具(如果启用)
|
# 添加 Google Search Grounding 工具(如果启用)
|
||||||
# 注意:enable_image_search=True 时会自动隐含 enable_grounding
|
# 注意:enable_image_search=True 时会自动隐含 enable_grounding
|
||||||
if enable_grounding or enable_image_search:
|
if enable_grounding or enable_image_search:
|
||||||
|
|||||||
+24
-2
@@ -139,6 +139,8 @@ class NanoBananaV2:
|
|||||||
for provider_extra in cls._PROVIDER_EXTRA_INPUTS.values():
|
for provider_extra in cls._PROVIDER_EXTRA_INPUTS.values():
|
||||||
optional.update(provider_extra)
|
optional.update(provider_extra)
|
||||||
|
|
||||||
|
optional["图片质量"] = (["日常", "高清"], {"default": "日常"})
|
||||||
|
|
||||||
for i in range(1, 10):
|
for i in range(1, 10):
|
||||||
optional[f"参考图{i}"] = ("IMAGE",)
|
optional[f"参考图{i}"] = ("IMAGE",)
|
||||||
|
|
||||||
@@ -407,6 +409,7 @@ class NanoBananaV2:
|
|||||||
on_progress(delta)
|
on_progress(delta)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
t_req_start = time.time()
|
||||||
task_id = await self._submit_one(
|
task_id = await self._submit_one(
|
||||||
session, provider, prompt, model,
|
session, provider, prompt, model,
|
||||||
resolution, aspect_ratio, input_images,
|
resolution, aspect_ratio, input_images,
|
||||||
@@ -415,11 +418,17 @@ class NanoBananaV2:
|
|||||||
response_data = await self._poll_one(
|
response_data = await self._poll_one(
|
||||||
session, provider, task_id, on_progress=_track_progress,
|
session, provider, task_id, on_progress=_track_progress,
|
||||||
)
|
)
|
||||||
|
request_time = time.time() - t_req_start
|
||||||
|
|
||||||
|
t_dl_start = time.time()
|
||||||
images_list = await provider.parse_result(response_data, session)
|
images_list = await provider.parse_result(response_data, session)
|
||||||
|
download_time = time.time() - t_dl_start
|
||||||
|
|
||||||
result["success"] = True
|
result["success"] = True
|
||||||
result["generated_count"] = len(images_list)
|
result["generated_count"] = len(images_list)
|
||||||
result["output_images"] = images_list
|
result["output_images"] = images_list
|
||||||
|
result["request_time"] = request_time
|
||||||
|
result["download_time"] = download_time
|
||||||
except InterruptProcessingException:
|
except InterruptProcessingException:
|
||||||
# 用户取消:补齐进度后向上传播,不吞掉
|
# 用户取消:补齐进度后向上传播,不吞掉
|
||||||
if contributed[0] < 1.0 and on_progress:
|
if contributed[0] < 1.0 and on_progress:
|
||||||
@@ -536,10 +545,13 @@ class NanoBananaV2:
|
|||||||
seed: int = kwargs.pop("seed", 0)
|
seed: int = kwargs.pop("seed", 0)
|
||||||
proxy_port: str = kwargs.pop("代理端口", "")
|
proxy_port: str = kwargs.pop("代理端口", "")
|
||||||
api_key_override: str = kwargs.pop("分组令牌", "")
|
api_key_override: str = kwargs.pop("分组令牌", "")
|
||||||
|
图片质量: str = kwargs.pop("图片质量", "日常")
|
||||||
|
image_format = "JPEG" if 图片质量 == "日常" else "PNG"
|
||||||
|
|
||||||
# 初始化 Provider
|
# 初始化 Provider
|
||||||
proxy_url = BaseAsyncImageProvider.build_proxy_url(proxy_port)
|
proxy_url = BaseAsyncImageProvider.build_proxy_url(proxy_port)
|
||||||
provider = self._get_provider(模型, proxy_url=proxy_url, api_key_override=api_key_override)
|
provider = self._get_provider(模型, proxy_url=proxy_url, api_key_override=api_key_override)
|
||||||
|
provider.image_compression = "webp" if 图片质量 == "日常" else None
|
||||||
|
|
||||||
if proxy_url:
|
if proxy_url:
|
||||||
print(f"{self.NODE_LABEL}: 已启用代理加速 -> {proxy_url}")
|
print(f"{self.NODE_LABEL}: 已启用代理加速 -> {proxy_url}")
|
||||||
@@ -644,7 +656,11 @@ class NanoBananaV2:
|
|||||||
error_msg = fr.get("error") or "未知错误"
|
error_msg = fr.get("error") or "未知错误"
|
||||||
print(f" FAIL #{idx}: {prompt_snippet}{'...' if len(prompt_snippet) >= 30 else ''} -> {error_msg}")
|
print(f" FAIL #{idx}: {prompt_snippet}{'...' if len(prompt_snippet) >= 30 else ''} -> {error_msg}")
|
||||||
else:
|
else:
|
||||||
print(f"{self.NODE_LABEL}: 完成!总耗时 {time_str}")
|
total_request = sum(r.get("request_time", 0) for r in results)
|
||||||
|
total_download = sum(r.get("download_time", 0) for r in results)
|
||||||
|
req_str = f"{total_request:.2f}s" if total_request >= 1 else f"{total_request:.3f}s"
|
||||||
|
dl_str = f"{total_download:.2f}s" if total_download >= 1 else f"{total_download:.3f}s"
|
||||||
|
print(f"{self.NODE_LABEL}: 完成!请求耗时 {req_str} | 下载耗时 {dl_str} | 总耗时 {time_str}")
|
||||||
|
|
||||||
# 收集输出图像
|
# 收集输出图像
|
||||||
output_images = []
|
output_images = []
|
||||||
@@ -661,6 +677,7 @@ class NanoBananaV2:
|
|||||||
f"所有任务均失败 ({len(failed)}/{len(results)}):\n{error_details}"
|
f"所有任务均失败 ({len(failed)}/{len(results)}):\n{error_details}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
output_images = [self._apply_image_format(img, image_format) for img in output_images]
|
||||||
output_tensor = _images_to_tensor_safe(output_images, self.NODE_LABEL)
|
output_tensor = _images_to_tensor_safe(output_images, self.NODE_LABEL)
|
||||||
|
|
||||||
import gc
|
import gc
|
||||||
@@ -721,7 +738,10 @@ class NanoBananaV2Batch(NanoBananaV2):
|
|||||||
if k in types["optional"]:
|
if k in types["optional"]:
|
||||||
new_optional[k] = types["optional"][k]
|
new_optional[k] = types["optional"][k]
|
||||||
|
|
||||||
# 2. 图片格式
|
# 2. 图片质量
|
||||||
|
new_optional["图片质量"] = (["日常", "高清"], {"default": "日常"})
|
||||||
|
|
||||||
|
# 2.5 图片格式
|
||||||
new_optional["图片格式"] = (["PNG", "JPEG"], {"default": "JPEG"})
|
new_optional["图片格式"] = (["PNG", "JPEG"], {"default": "JPEG"})
|
||||||
|
|
||||||
# 3. seed
|
# 3. seed
|
||||||
@@ -960,6 +980,7 @@ class NanoBananaV2Batch(NanoBananaV2):
|
|||||||
seed: int = kwargs.pop("seed", 0)
|
seed: int = kwargs.pop("seed", 0)
|
||||||
proxy_port: str = kwargs.pop("代理端口", "")
|
proxy_port: str = kwargs.pop("代理端口", "")
|
||||||
api_key_override: str = kwargs.pop("分组令牌", "")
|
api_key_override: str = kwargs.pop("分组令牌", "")
|
||||||
|
图片质量: str = kwargs.pop("图片质量", "日常")
|
||||||
image_format: str = kwargs.pop("图片格式", "JPEG")
|
image_format: str = kwargs.pop("图片格式", "JPEG")
|
||||||
save_path: str = kwargs.pop("图片保存路径(可选)", "").strip()
|
save_path: str = kwargs.pop("图片保存路径(可选)", "").strip()
|
||||||
命名规则选择: str = kwargs.pop("图片命名规则", "和图片同名")
|
命名规则选择: str = kwargs.pop("图片命名规则", "和图片同名")
|
||||||
@@ -970,6 +991,7 @@ class NanoBananaV2Batch(NanoBananaV2):
|
|||||||
|
|
||||||
proxy_url = BaseAsyncImageProvider.build_proxy_url(proxy_port)
|
proxy_url = BaseAsyncImageProvider.build_proxy_url(proxy_port)
|
||||||
provider = self._get_provider(模型, proxy_url=proxy_url, api_key_override=api_key_override)
|
provider = self._get_provider(模型, proxy_url=proxy_url, api_key_override=api_key_override)
|
||||||
|
provider.image_compression = "webp" if 图片质量 == "日常" else None
|
||||||
|
|
||||||
if proxy_url:
|
if proxy_url:
|
||||||
print(f"{self.NODE_LABEL}: 已启用代理加速 -> {proxy_url}")
|
print(f"{self.NODE_LABEL}: 已启用代理加速 -> {proxy_url}")
|
||||||
|
|||||||
Reference in New Issue
Block a user