Files
comfyui_o1key/nodes/grok_video.py
T
Jony ba920f2b66 Publish current ComfyUI O1Key code baseline
Replace the prior release tree with the current plugin, frontend, tests, and documentation. Document retired node IDs and the public Gitea update source.
2026-09-24 19:56:48 +08:00

312 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Lean ComfyUI nodes for O1Key Grok Imagine Video."""
import asyncio
import math
import os
import re
from typing import Dict
from ..clients.grok_video_client import GrokVideoClient
from ..utils.config import get_base_url_by_route
from ..utils.image_utils import tensor_to_pil
from ..utils.r2_uploader import upload_audio, upload_image, upload_video
try:
import folder_paths
except ImportError:
folder_paths = None
try:
from comfy.utils import ProgressBar
except ImportError:
ProgressBar = None
try:
from comfy_api.input_impl import VideoFromFile
except Exception:
try:
from comfy_api.latest import InputImpl
VideoFromFile = InputImpl.VideoFromFile
except Exception:
VideoFromFile = None
MODEL_OPTIONS = list(GrokVideoClient.MODEL_OPTIONS)
ASPECT_RATIO_OPTIONS = list(GrokVideoClient.ASPECT_RATIO_OPTIONS)
RESOLUTION_OPTIONS = list(GrokVideoClient.RESOLUTION_OPTIONS)
GENERATION_MODE_OPTIONS = ["文生视频", "图生视频", "参考生视频"]
EDIT_MODE_OPTIONS = ["编辑视频", "续写视频"]
IMAGE_INPUT_NAMES = [f"图片{i}" for i in range(1, 8)]
AUDIO_INPUT_NAMES = ["音频素材", "音频素材2", "音频素材3"]
def _get_output_dir() -> str:
if folder_paths is not None:
base_dir = folder_paths.get_temp_directory()
else:
plugin_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
base_dir = os.path.join(os.path.dirname(os.path.dirname(plugin_dir)), "temp")
output_dir = os.path.join(base_dir, "grok_video")
os.makedirs(output_dir, exist_ok=True)
return output_dir
def _single_pil_image(image_tensor, input_name: str):
if image_tensor is None:
return None
images = tensor_to_pil(image_tensor)
if len(images) != 1:
raise ValueError(f"{input_name} 只能连接 1 张图片,请拆分批次后再连接。")
return images[0].convert("RGB")
def _parse_voice_ids(value: object) -> list[str]:
voice_ids = [item.strip() for item in re.split(r"[,\n]", str(value or "")) if item.strip()]
if len(voice_ids) > 3:
raise ValueError("参考音色 ID 最多填写 3 个。")
return voice_ids
def _video_duration_seconds(video) -> float:
getter = getattr(video, "get_duration", None)
if not callable(getter):
raise ValueError("无法读取输入视频时长;请连接 ComfyUI 原生 VIDEO 输出。")
try:
duration = float(getter())
except Exception as exc:
raise ValueError("无法读取输入视频时长;请确认视频文件可以正常解码。") from exc
if not math.isfinite(duration) or duration <= 0:
raise ValueError("输入视频时长无效;请确认视频文件可以正常解码。")
return duration
def _progress_callback():
progress_bar = ProgressBar(100) if ProgressBar is not None else None
progress_value = [0]
def callback(progress: int, _status: str, _elapsed: float) -> None:
current = max(0, min(100, int(progress or 0)))
if progress_bar is not None and current > progress_value[0]:
progress_bar.update(current - progress_value[0])
progress_value[0] = current
return progress_bar, progress_value, callback
def _finish_video(result: Dict[str, object], progress_bar, progress_value):
if progress_bar is not None and progress_value[0] < 100:
progress_bar.update(100 - progress_value[0])
video_path = result["video_path"]
print(f"Grok Video:下载完成:{video_path}")
return (VideoFromFile(video_path),)
class O1keyGrokVideo:
"""文生、图生或参考素材生 Grok 视频。素材会自动上传为 URL。"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"生成模式": (GENERATION_MODE_OPTIONS, {"default": "文生视频"}),
"提示词": ("STRING", {"default": "", "multiline": True}),
"模型": (MODEL_OPTIONS, {"default": GrokVideoClient.DEFAULT_MODEL}),
"时长(秒)": ("INT", {"default": 8, "min": 1, "max": 15, "step": 1}),
"宽高比": (ASPECT_RATIO_OPTIONS, {"default": "16:9"}),
"分辨率": (RESOLUTION_OPTIONS, {"default": "480p"}),
"参考音色ID(逗号分隔)": ("STRING", {"default": ""}),
},
"optional": {
**{input_name: ("IMAGE",) for input_name in IMAGE_INPUT_NAMES},
**{input_name: ("AUDIO",) for input_name in AUDIO_INPUT_NAMES},
},
}
RETURN_TYPES = ("VIDEO",)
RETURN_NAMES = ("视频",)
FUNCTION = "generate"
CATEGORY = "comfyui_o1key/Video"
DESCRIPTION = (
"支持文生、图生和多参考素材生成。图生视频只连接图片 1;参考生视频最多使用 7 张图和 "
"3 个参考音频(AUDIO 或 voice_id 合计)。Grok 1.5 的文生/图生可选 1080p,多参考最高 720p。"
)
def generate(self, **kwargs):
if VideoFromFile is None:
raise RuntimeError("当前 ComfyUI 版本不支持 VideoFromFile,无法输出 VIDEO。")
mode = kwargs.get("生成模式", "文生视频")
model = kwargs.get("模型", GrokVideoClient.DEFAULT_MODEL)
prompt = (kwargs.get("提示词") or "").strip()
connected_images = [
(input_name, image)
for input_name in IMAGE_INPUT_NAMES
if (image := _single_pil_image(kwargs.get(input_name), input_name)) is not None
]
images = [image for _, image in connected_images]
audios = [kwargs.get(input_name) for input_name in AUDIO_INPUT_NAMES if kwargs.get(input_name) is not None]
voice_ids = _parse_voice_ids(kwargs.get("参考音色ID(逗号分隔)"))
duration = kwargs.get("时长(秒)", 8)
aspect_ratio = kwargs.get("宽高比", "16:9")
resolution = kwargs.get("分辨率", "480p")
if mode not in GENERATION_MODE_OPTIONS:
raise ValueError(f"不支持的生成模式:{mode}。")
if len(audios) + len(voice_ids) > 3:
raise ValueError("参考音频与参考音色 ID 合计最多 3 个。")
if mode == "文生视频":
if images or audios or voice_ids:
raise ValueError("文生视频不需要连接图像或音频素材。")
elif mode == "图生视频":
if len(images) != 1 or connected_images[0][0] != "图片1":
raise ValueError("图生视频需要在“图片 1”连接 1 张图片,其他图片端口请留空。")
if audios or voice_ids:
raise ValueError("图生视频不支持音频素材,请使用参考生视频。")
else:
if not images and not audios and not voice_ids:
raise ValueError("参考生视频至少需要连接图像素材或音频素材。")
placeholder_image = {"url": "https://example.invalid/image"} if mode == "图生视频" else None
placeholder_references = (
[{"url": f"https://example.invalid/reference-{index}"} for index in range(len(images))]
if mode == "参考生视频"
else []
)
placeholder_audios = (
[{"url": f"https://example.invalid/audio-{index}"} for index in range(len(audios))]
+ [{"voice_id": voice_id} for voice_id in voice_ids]
if mode == "参考生视频"
else []
)
# Validate every user-controlled field before temporary uploads or paid generation calls.
GrokVideoClient.build_video_body(
operation="generate",
prompt=prompt,
model=model,
duration=duration,
aspect_ratio=aspect_ratio,
resolution=resolution,
image=placeholder_image,
reference_images=placeholder_references,
reference_audios=placeholder_audios,
)
base_url = get_base_url_by_route()
client = GrokVideoClient(base_url=base_url)
async def upload_materials():
image_urls, audio_urls = await asyncio.gather(
asyncio.gather(*(upload_image(image, base_url=base_url) for image in images)),
asyncio.gather(*(upload_audio(audio, base_url=base_url) for audio in audios)),
)
return list(image_urls), list(audio_urls)
image_urls, audio_urls = client.run_async_in_thread(upload_materials())
if mode == "文生视频":
image = None
reference_images = []
reference_audios = []
elif mode == "图生视频":
image = {"url": image_urls[0]}
reference_images = []
reference_audios = []
else:
image = None
reference_images = [{"url": url} for url in image_urls]
reference_audios = [
*({"url": url} for url in audio_urls),
*({"voice_id": voice_id} for voice_id in voice_ids),
]
progress_bar, progress_value, callback = _progress_callback()
result = client.run_video_sync(
operation="generate",
prompt=prompt,
model=model,
duration=duration,
aspect_ratio=aspect_ratio,
resolution=resolution,
image=image,
reference_images=reference_images,
reference_audios=reference_audios,
output_dir=_get_output_dir(),
progress_callback=callback,
)
return _finish_video(result, progress_bar, progress_value)
class O1keyGrokVideoEdit:
"""编辑或续写 Grok 视频。输入 VIDEO 会自动上传为 URL。"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"操作": (EDIT_MODE_OPTIONS, {"default": "编辑视频"}),
"提示词": ("STRING", {"default": "", "multiline": True}),
"续写时长(秒)": ("INT", {"default": 6, "min": 2, "max": 10, "step": 1}),
"模型": (MODEL_OPTIONS, {"default": GrokVideoClient.DEFAULT_MODEL}),
},
"optional": {"视频素材": ("VIDEO",)},
}
RETURN_TYPES = ("VIDEO",)
RETURN_NAMES = ("视频",)
FUNCTION = "generate"
CATEGORY = "comfyui_o1key/Video"
DESCRIPTION = (
"编辑或续写视频。编辑输入最长 8.7 秒,并保留原时长和宽高比,输出最高 720p;"
"续写时长为 2–10 秒,输出总时长等于输入时长加续写时长。"
)
def generate(self, **kwargs):
if VideoFromFile is None:
raise RuntimeError("当前 ComfyUI 版本不支持 VideoFromFile,无法输出 VIDEO。")
video = kwargs.get("视频素材")
if video is None:
raise ValueError("请连接一个 VIDEO 类型的视频素材。")
selected_operation = kwargs.get("操作", "编辑视频")
if selected_operation not in EDIT_MODE_OPTIONS:
raise ValueError(f"不支持的 Grok 视频操作:{selected_operation}。")
operation = "edit" if selected_operation == "编辑视频" else "extend"
prompt = (kwargs.get("提示词") or "").strip()
model = kwargs.get("模型", GrokVideoClient.DEFAULT_MODEL)
duration = kwargs.get("续写时长(秒)", 6)
if operation == "edit" and _video_duration_seconds(video) > 8.7:
raise ValueError("Grok 视频编辑的输入视频不能超过 8.7 秒。")
GrokVideoClient.build_video_body(
operation=operation,
prompt=prompt,
model=model,
duration=duration,
video={"url": "https://example.invalid/video"},
)
base_url = get_base_url_by_route()
client = GrokVideoClient(base_url=base_url)
video_url = client.run_async_in_thread(upload_video(video, base_url=base_url))
progress_bar, progress_value, callback = _progress_callback()
result = client.run_video_sync(
operation=operation,
prompt=prompt,
model=model,
duration=duration,
video={"url": video_url},
output_dir=_get_output_dir(),
progress_callback=callback,
)
return _finish_video(result, progress_bar, progress_value)
NODE_CLASS_MAPPINGS = {
"O1keyGrokVideo": O1keyGrokVideo,
"O1keyGrokVideoEdit": O1keyGrokVideoEdit,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"O1keyGrokVideo": "Grok Video",
"O1keyGrokVideoEdit": "Grok Video Edit",
}