mirror of
https://github.com/comfyanonymous/ComfyUI.git
synced 2026-07-22 07:47:36 +08:00
Compare commits
49
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a7b63915dc | ||
|
|
5d8b3690ff | ||
|
|
515d73f6f2 | ||
|
|
a14cc66712 | ||
|
|
09d55ba54a | ||
|
|
c62c727aac | ||
|
|
71cf5a11f1 | ||
|
|
ee06f45ff3 | ||
|
|
08f0b15d60 | ||
|
|
cebe350e0e | ||
|
|
9807d2a743 | ||
|
|
88be6cd111 | ||
|
|
0ea2773122 | ||
|
|
589eb6e1d9 | ||
|
|
317317c98a | ||
|
|
c5d9d39d0c | ||
|
|
ea1e5786b8 | ||
|
|
78ca58cf29 | ||
|
|
5d0835deee | ||
|
|
dc582b0fe2 | ||
|
|
343afeb6d0 | ||
|
|
fe20a6483e | ||
|
|
e7b60beaf1 | ||
|
|
ce237f9a79 | ||
|
|
1d79c41064 | ||
|
|
ccadcaf4e3 | ||
|
|
2fc66a5e8d | ||
|
|
15aa401146 | ||
|
|
4c152491f0 | ||
|
|
72ce283d32 | ||
|
|
bd27ae82a1 | ||
|
|
9f9a9272ce | ||
|
|
a0b35f60be | ||
|
|
d0e71574a1 | ||
|
|
8a32df7108 | ||
|
|
03e6b57a86 | ||
|
|
391c628b8e | ||
|
|
790e6775da | ||
|
|
3b5483a225 | ||
|
|
d73483be0e | ||
|
|
acf8f95eb5 | ||
|
|
2bed0a5d73 | ||
|
|
96a3e0ccad | ||
|
|
ea583ae155 | ||
|
|
cbb0ddd87e | ||
|
|
b107b834f2 | ||
|
|
53da5f94e6 | ||
|
|
cca6687a0a | ||
|
|
6f4c742bc5 |
@@ -23,9 +23,9 @@ jobs:
|
||||
# SHA-pinned per zizmor `unpinned-uses: hash-pin`. Bump this SHA to pick up
|
||||
# upstream changes; keep `workflows_ref` matching so prompts/scripts load
|
||||
# from the same commit as the workflow definition.
|
||||
uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@964d5aad37cbfb57c5b23961d42c2fd85868bf1d # github-workflows main (964d5aa)
|
||||
uses: Comfy-Org/github-workflows/.github/workflows/cursor-review.yml@047ca48febe3a6647608ed2e0c4331b491cb9d6a # github-workflows#9
|
||||
with:
|
||||
workflows_ref: 964d5aad37cbfb57c5b23961d42c2fd85868bf1d
|
||||
workflows_ref: 047ca48febe3a6647608ed2e0c4331b491cb9d6a
|
||||
diff_excludes: >-
|
||||
:!**/.claude/**
|
||||
:!**/dist/**
|
||||
|
||||
@@ -19,9 +19,6 @@
|
||||
better to remove a broken feature path than keep a complicated partial fix.
|
||||
- Preserve existing APIs, node names, model-loading behavior, file layout, and
|
||||
workflow compatibility unless the change is explicitly about replacing them.
|
||||
- When compatibility is explicitly out of scope, remove compatibility-only
|
||||
aliases, duplicate nodes, legacy entry points, and preset wrappers instead of
|
||||
retaining parallel ways to perform the same operation.
|
||||
- Code must look hand-written for this repository. Changes that read like
|
||||
generic AI-generated code will be rejected automatically: unnecessary helper
|
||||
layers, vague names, boilerplate comments, defensive branches without a real
|
||||
@@ -99,13 +96,6 @@
|
||||
unless they are read by current code and change current behavior. Remove
|
||||
pass-through or stored-but-unused values instead of preserving upstream or
|
||||
deprecated API baggage.
|
||||
- Do not add a model-specific option to a shared helper when only one caller
|
||||
needs it. Keep one-off behavior at the model integration boundary, or extend
|
||||
the shared helper only when the option is a coherent reusable capability.
|
||||
- Implementations of shared model interfaces should accept the standard caller
|
||||
contract without model-specific rejection branches for optional capabilities
|
||||
they do not consume. Let supported behavior be determined by implementation
|
||||
paths that actually use those inputs.
|
||||
- If an implementation needs auxiliary values for its own workflow, expose them
|
||||
through a private helper or a clearly named implementation-specific method
|
||||
instead of overloading the public method's return contract.
|
||||
@@ -164,10 +154,6 @@
|
||||
`comfy-kitchen` helpers where they already solve the problem.
|
||||
- Use optimized comfy-kitchen ops in places where they improve performance
|
||||
without changing the expected dtype, device, memory, or interface behavior.
|
||||
- Prefer ComfyUI's shared optimized kernels and backend dispatchers over
|
||||
handwritten implementations of the same operation. Remove duplicate local
|
||||
kernels and adapt inputs to the shared operation's documented layout while
|
||||
preserving the model's original math and output contract.
|
||||
- All models should use the optimized attention function selected by ComfyUI.
|
||||
Treat optimized backend functions, dispatch helpers, and capability-selected
|
||||
callables as opaque. Higher-level code must not inspect function identity,
|
||||
@@ -190,12 +176,6 @@
|
||||
- Model detection code that inspects linear weight shapes should only use the
|
||||
first dimension. The second dimension may be half the original size for
|
||||
NVFP4 or other 4-bit quantized models.
|
||||
- A model-detection signature must guard every state-dict key it dereferences.
|
||||
Do not partially match a format and then raise an incidental `KeyError` while
|
||||
extracting its configuration.
|
||||
- Order model-detection checks from established or more-specific signatures to
|
||||
newer or broader signatures. Put a broad new detector near the generic
|
||||
fallback when giving it higher precedence could steal another model family.
|
||||
- Avoid adding `einops` usage in core inference code. Use native torch tensor
|
||||
ops such as `reshape`, `view`, `permute`, `transpose`, `flatten`, `unflatten`,
|
||||
`unsqueeze`, and `squeeze` instead.
|
||||
@@ -212,23 +192,11 @@
|
||||
methods for scalar or structural calculations.
|
||||
- Avoid unnecessary casts and transfers. Preserve the intended compute dtype,
|
||||
storage dtype, bias dtype, and original tensor shape metadata.
|
||||
- Do not cast the result of an optimized backend operation back to its input
|
||||
dtype unless that backend's documented result contract requires normalization.
|
||||
In particular, trust the selected optimized-attention implementation to honor
|
||||
its dtype contract.
|
||||
- Keep model-native latent layout handling inside the model or latent-format
|
||||
owner, not in helper nodes. Do not collapse, expand, pack, or unpack latent
|
||||
dimensions in nodes or other caller-side adapters just to satisfy a model
|
||||
forward; the model path should consume and return the native latent shape for
|
||||
that model family.
|
||||
- DiT models should accept latent dimensions that are not exact patch-size
|
||||
multiples. Use `comfy.ldm.common_dit.pad_to_patch_size` on every patchified
|
||||
target or reference input, then crop only the target output back to its
|
||||
original dimensions.
|
||||
- Avoid defensive shape and configuration checks that merely replace the clear
|
||||
failure from the tensor operation immediately below them. Add explicit
|
||||
validation only when it provides materially better context at a real boundary
|
||||
or prevents silent incorrect output.
|
||||
- Assume inputs to the main model forward are already in the compute dtype by
|
||||
default, except integer inputs such as some model timestep tensors. Do not add
|
||||
defensive or convenience casts in model code; it is better for invalid dtype
|
||||
@@ -292,15 +260,6 @@
|
||||
- Model implementations should add the minimal number of ComfyUI nodes required
|
||||
to run the model. Reuse existing nodes as much as possible; adapting the model
|
||||
to work with existing nodes is strongly preferred over creating new nodes.
|
||||
- Use `io.Autogrow` for a variable number of repeated inputs instead of a fixed
|
||||
series of numbered optional sockets. Set its minimum to zero when the model
|
||||
has a valid no-item path, and cap it only when the model has a real limit.
|
||||
- Mark inputs optional when execution has a valid path that does not read them.
|
||||
If one optional input is needed only to process another optional input, do not
|
||||
force users on the path that supplies neither to connect it.
|
||||
- Conditioning nodes should normally output conditioning only. Do not expose
|
||||
input or intermediate images as convenience outputs for downstream sizing or
|
||||
routing; use the existing image path or a dedicated image operation instead.
|
||||
- Nodes should output only values they own. Do not add pass-through outputs for
|
||||
workflow convenience unless the node is explicitly an output node. Existing
|
||||
models, latents, conditioning, or other inputs should flow directly to the
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
* @comfyanonymous @kosinkadink @guill @alexisrolland @rattus128 @kijai
|
||||
|
||||
/CODEOWNERS @comfyanonymous
|
||||
/AGENTS.md @comfyanonymous
|
||||
/.ci/ @comfyanonymous
|
||||
/.github/ @comfyanonymous
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
"""
|
||||
Download manager schema.
|
||||
|
||||
Adds the two tables that back the server-side model download manager:
|
||||
transient job/queue state (``downloads`` + per-segment ``download_segments``).
|
||||
|
||||
Revision ID: 0005_download_manager
|
||||
Revises: 0004_drop_tag_type
|
||||
Create Date: 2026-06-27
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "0005_download_manager"
|
||||
down_revision = "0004_drop_tag_type"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"downloads",
|
||||
sa.Column("id", sa.String(length=36), primary_key=True),
|
||||
sa.Column("url", sa.Text(), nullable=False),
|
||||
sa.Column("final_url", sa.Text(), nullable=True),
|
||||
sa.Column("model_id", sa.String(length=1024), nullable=False),
|
||||
sa.Column("dest_path", sa.Text(), nullable=False),
|
||||
sa.Column("temp_path", sa.Text(), nullable=False),
|
||||
sa.Column("status", sa.String(length=16), nullable=False),
|
||||
sa.Column("priority", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("total_bytes", sa.BigInteger(), nullable=True),
|
||||
sa.Column("bytes_done", sa.BigInteger(), nullable=False, server_default="0"),
|
||||
sa.Column("etag", sa.String(length=512), nullable=True),
|
||||
sa.Column("last_modified", sa.String(length=128), nullable=True),
|
||||
sa.Column(
|
||||
"accept_ranges", sa.Boolean(), nullable=False, server_default=sa.text("false")
|
||||
),
|
||||
sa.Column("expected_sha256", sa.String(length=64), nullable=True),
|
||||
sa.Column(
|
||||
"allow_any_extension",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
),
|
||||
sa.Column("attempts", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("error", sa.Text(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=False),
|
||||
sa.CheckConstraint("bytes_done >= 0", name="ck_downloads_bytes_done_nonneg"),
|
||||
sa.CheckConstraint(
|
||||
"total_bytes IS NULL OR total_bytes >= 0",
|
||||
name="ck_downloads_total_bytes_nonneg",
|
||||
),
|
||||
)
|
||||
op.create_index("ix_downloads_status", "downloads", ["status"])
|
||||
op.create_index("ix_downloads_priority", "downloads", ["priority"])
|
||||
op.create_index("ix_downloads_model_id", "downloads", ["model_id"])
|
||||
|
||||
op.create_table(
|
||||
"download_segments",
|
||||
sa.Column(
|
||||
"download_id",
|
||||
sa.String(length=36),
|
||||
sa.ForeignKey("downloads.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("idx", sa.Integer(), nullable=False),
|
||||
sa.Column("start_offset", sa.BigInteger(), nullable=False),
|
||||
sa.Column("end_offset", sa.BigInteger(), nullable=False),
|
||||
sa.Column("bytes_done", sa.BigInteger(), nullable=False, server_default="0"),
|
||||
sa.PrimaryKeyConstraint("download_id", "idx", name="pk_download_segments"),
|
||||
sa.CheckConstraint("bytes_done >= 0", name="ck_segments_bytes_done_nonneg"),
|
||||
sa.CheckConstraint("end_offset >= start_offset", name="ck_segments_range"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("download_segments")
|
||||
|
||||
op.drop_index("ix_downloads_model_id", table_name="downloads")
|
||||
op.drop_index("ix_downloads_priority", table_name="downloads")
|
||||
op.drop_index("ix_downloads_status", table_name="downloads")
|
||||
op.drop_table("downloads")
|
||||
+6
-4
@@ -4,7 +4,8 @@ import shutil
|
||||
from app.logger import log_startup_warning
|
||||
from utils.install_util import get_missing_requirements_message
|
||||
from filelock import FileLock, Timeout
|
||||
from comfy.cli_args import args
|
||||
# Import the module so tests that reload comfy.cli_args see the live object.
|
||||
import comfy.cli_args
|
||||
|
||||
_DB_AVAILABLE = False
|
||||
Session = None
|
||||
@@ -21,6 +22,7 @@ try:
|
||||
|
||||
from app.database.models import Base
|
||||
import app.assets.database.models # noqa: F401 — register models with Base.metadata
|
||||
import app.model_downloader.database.models # noqa: F401 — register models with Base.metadata
|
||||
|
||||
_DB_AVAILABLE = True
|
||||
except ImportError as e:
|
||||
@@ -57,13 +59,13 @@ def get_alembic_config():
|
||||
|
||||
config = Config(config_path)
|
||||
config.set_main_option("script_location", scripts_path)
|
||||
config.set_main_option("sqlalchemy.url", args.database_url)
|
||||
config.set_main_option("sqlalchemy.url", comfy.cli_args.args.database_url)
|
||||
|
||||
return config
|
||||
|
||||
|
||||
def get_db_path():
|
||||
url = args.database_url
|
||||
url = comfy.cli_args.args.database_url
|
||||
if url.startswith("sqlite:///"):
|
||||
return url.split("///")[1]
|
||||
else:
|
||||
@@ -97,7 +99,7 @@ def _is_memory_db(db_url):
|
||||
|
||||
|
||||
def init_db():
|
||||
db_url = args.database_url
|
||||
db_url = comfy.cli_args.args.database_url
|
||||
logging.debug(f"Database URL: {db_url}")
|
||||
|
||||
if _is_memory_db(db_url):
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
"""aiohttp routes for the download manager.
|
||||
|
||||
Endpoint surface (all under ``/api/download``), mirroring the response
|
||||
envelope used by ``app/assets/api/routes.py``:
|
||||
|
||||
POST /api/download/enqueue
|
||||
GET /api/download
|
||||
POST /api/download/availability
|
||||
POST /api/download/clear
|
||||
GET /api/download/auth
|
||||
POST /api/download/auth/{provider}/login
|
||||
POST /api/download/auth/{provider}/logout
|
||||
GET /api/download/{id}
|
||||
DELETE /api/download/{id}
|
||||
POST /api/download/{id}/pause
|
||||
POST /api/download/{id}/resume
|
||||
POST /api/download/{id}/cancel
|
||||
POST /api/download/{id}/priority
|
||||
|
||||
Note on ordering: the static ``auth`` routes are registered before the dynamic
|
||||
``/api/download/{id}`` route so a request to ``.../auth`` is not captured as
|
||||
``id == "auth"``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
from aiohttp import web
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from app.model_downloader.api import schemas_in, schemas_out
|
||||
from app.model_downloader.auth.oauth import LoginInProgress, OAuthNotConfigured
|
||||
from app.model_downloader.auth.providers import PROVIDERS
|
||||
from app.model_downloader.auth.store import AUTH_STORE
|
||||
from app.model_downloader.manager import DOWNLOAD_MANAGER, DownloadError
|
||||
|
||||
ROUTES = web.RouteTableDef()
|
||||
|
||||
|
||||
def register_routes(app: web.Application) -> None:
|
||||
"""Wire the download-manager routes into the running aiohttp app."""
|
||||
app.add_routes(ROUTES)
|
||||
|
||||
|
||||
# ----- envelope helpers (same shape as app/assets/api/routes.py) -----
|
||||
|
||||
|
||||
def _error(status: int, code: str, message: str, details: dict | None = None) -> web.Response:
|
||||
return web.json_response(
|
||||
{"error": {"code": code, "message": message, "details": details or {}}},
|
||||
status=status,
|
||||
)
|
||||
|
||||
|
||||
def _ok(payload, status: int = 200) -> web.Response:
|
||||
return web.json_response(payload, status=status)
|
||||
|
||||
|
||||
async def _parse(request: web.Request, model: type[BaseModel]):
|
||||
try:
|
||||
raw = await request.json()
|
||||
except json.JSONDecodeError:
|
||||
return _error(400, "INVALID_JSON", "Request body must be valid JSON.")
|
||||
try:
|
||||
return model.model_validate(raw)
|
||||
except ValidationError as ve:
|
||||
return _error(400, "INVALID_BODY", "Validation failed.", {"errors": json.loads(ve.json())})
|
||||
|
||||
|
||||
def _from_download_error(e: DownloadError) -> web.Response:
|
||||
return _error(e.http_status, e.code, e.message)
|
||||
|
||||
|
||||
# ----- downloads: collection + enqueue + availability -----
|
||||
|
||||
|
||||
@ROUTES.post("/api/download/enqueue")
|
||||
async def enqueue(request: web.Request) -> web.Response:
|
||||
parsed = await _parse(request, schemas_in.EnqueueRequest)
|
||||
if isinstance(parsed, web.Response):
|
||||
return parsed
|
||||
try:
|
||||
download_id = await DOWNLOAD_MANAGER.enqueue(
|
||||
parsed.url,
|
||||
parsed.model_id,
|
||||
priority=parsed.priority,
|
||||
expected_sha256=parsed.expected_sha256,
|
||||
allow_any_extension=parsed.allow_any_extension,
|
||||
)
|
||||
except DownloadError as e:
|
||||
return _from_download_error(e)
|
||||
return _ok({"download_id": download_id, "accepted": True}, status=202)
|
||||
|
||||
|
||||
@ROUTES.get("/api/download")
|
||||
async def list_downloads(request: web.Request) -> web.Response:
|
||||
return _ok({"downloads": await DOWNLOAD_MANAGER.list()})
|
||||
|
||||
|
||||
@ROUTES.post("/api/download/availability")
|
||||
async def availability(request: web.Request) -> web.Response:
|
||||
parsed = await _parse(request, schemas_in.AvailabilityRequest)
|
||||
if isinstance(parsed, web.Response):
|
||||
return parsed
|
||||
return _ok({"models": await DOWNLOAD_MANAGER.availability(parsed.models)})
|
||||
|
||||
|
||||
@ROUTES.post("/api/download/clear")
|
||||
async def clear(request: web.Request) -> web.Response:
|
||||
deleted = await DOWNLOAD_MANAGER.clear()
|
||||
return _ok({"deleted": deleted})
|
||||
|
||||
|
||||
# ----- auth (OAuth login + env-key status) — must precede /{id} -----
|
||||
|
||||
|
||||
@ROUTES.get("/api/download/auth")
|
||||
async def auth_status(request: web.Request) -> web.Response:
|
||||
return _ok({"providers": schemas_out.auth_status()})
|
||||
|
||||
|
||||
@ROUTES.post("/api/download/auth/{provider}/login")
|
||||
async def auth_login(request: web.Request) -> web.Response:
|
||||
provider = PROVIDERS.get(request.match_info["provider"])
|
||||
if provider is None:
|
||||
return _error(400, "UNKNOWN_PROVIDER", "No such auth provider.")
|
||||
try:
|
||||
authorize_url = await AUTH_STORE.begin_login(provider)
|
||||
except OAuthNotConfigured as e:
|
||||
return _error(400, "OAUTH_NOT_CONFIGURED", str(e))
|
||||
except LoginInProgress as e:
|
||||
return _error(409, "LOGIN_IN_PROGRESS", str(e))
|
||||
return _ok({"authorize_url": authorize_url})
|
||||
|
||||
|
||||
@ROUTES.post("/api/download/auth/{provider}/logout")
|
||||
async def auth_logout(request: web.Request) -> web.Response:
|
||||
provider = PROVIDERS.get(request.match_info["provider"])
|
||||
if provider is None:
|
||||
return _error(400, "UNKNOWN_PROVIDER", "No such auth provider.")
|
||||
AUTH_STORE.clear(provider.name)
|
||||
return _ok({"logged_out": True})
|
||||
|
||||
|
||||
# ----- single download by id (dynamic; registered last) -----
|
||||
|
||||
|
||||
@ROUTES.get("/api/download/{id}")
|
||||
async def get_download(request: web.Request) -> web.Response:
|
||||
view = await DOWNLOAD_MANAGER.status(request.match_info["id"])
|
||||
if view is None:
|
||||
return _error(404, "NOT_FOUND", "No such download.")
|
||||
return _ok(view)
|
||||
|
||||
|
||||
@ROUTES.delete("/api/download/{id}")
|
||||
async def delete_download(request: web.Request) -> web.Response:
|
||||
try:
|
||||
await DOWNLOAD_MANAGER.delete(request.match_info["id"])
|
||||
except DownloadError as e:
|
||||
return _from_download_error(e)
|
||||
return _ok({"deleted": True})
|
||||
|
||||
|
||||
@ROUTES.post("/api/download/{id}/pause")
|
||||
async def pause(request: web.Request) -> web.Response:
|
||||
try:
|
||||
await DOWNLOAD_MANAGER.pause(request.match_info["id"])
|
||||
except DownloadError as e:
|
||||
return _from_download_error(e)
|
||||
return _ok({"ok": True})
|
||||
|
||||
|
||||
@ROUTES.post("/api/download/{id}/resume")
|
||||
async def resume(request: web.Request) -> web.Response:
|
||||
try:
|
||||
await DOWNLOAD_MANAGER.resume(request.match_info["id"])
|
||||
except DownloadError as e:
|
||||
return _from_download_error(e)
|
||||
return _ok({"ok": True})
|
||||
|
||||
|
||||
@ROUTES.post("/api/download/{id}/cancel")
|
||||
async def cancel(request: web.Request) -> web.Response:
|
||||
try:
|
||||
await DOWNLOAD_MANAGER.cancel(request.match_info["id"])
|
||||
except DownloadError as e:
|
||||
return _from_download_error(e)
|
||||
return _ok({"ok": True})
|
||||
|
||||
|
||||
@ROUTES.post("/api/download/{id}/priority")
|
||||
async def set_priority(request: web.Request) -> web.Response:
|
||||
parsed = await _parse(request, schemas_in.PriorityRequest)
|
||||
if isinstance(parsed, web.Response):
|
||||
return parsed
|
||||
try:
|
||||
await DOWNLOAD_MANAGER.set_priority(request.match_info["id"], parsed.priority)
|
||||
except DownloadError as e:
|
||||
return _from_download_error(e)
|
||||
return _ok({"ok": True})
|
||||
@@ -0,0 +1,46 @@
|
||||
"""Request schemas for the download manager API.
|
||||
|
||||
Pydantic enforces shape at the boundary; handlers operate only on validated
|
||||
values past that point.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
|
||||
|
||||
class EnqueueRequest(BaseModel):
|
||||
url: str
|
||||
model_id: str
|
||||
priority: int = 0
|
||||
expected_sha256: Optional[str] = None
|
||||
allow_any_extension: bool = False
|
||||
|
||||
@field_validator("url")
|
||||
@classmethod
|
||||
def _strip_url(cls, v: str) -> str:
|
||||
return v.strip()
|
||||
|
||||
|
||||
class PriorityRequest(BaseModel):
|
||||
priority: int
|
||||
|
||||
|
||||
class AvailabilityRequest(BaseModel):
|
||||
"""``{model_id: url}`` — the URLs declared in the workflow JSON."""
|
||||
|
||||
models: dict[str, str] = Field(default_factory=dict)
|
||||
|
||||
@field_validator("models")
|
||||
@classmethod
|
||||
def _strip_urls(cls, v: dict[str, str]) -> dict[str, str]:
|
||||
return {k: url.strip() for k, url in v.items()}
|
||||
|
||||
|
||||
__all__ = [
|
||||
"EnqueueRequest",
|
||||
"PriorityRequest",
|
||||
"AvailabilityRequest",
|
||||
]
|
||||
@@ -0,0 +1,15 @@
|
||||
"""Response helpers for the download manager API.
|
||||
|
||||
The download/status read models are plain dicts produced by the manager. This
|
||||
module serializes the per-provider auth status (never a token) for the API.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from app.model_downloader.auth.providers import PROVIDERS
|
||||
from app.model_downloader.auth.store import AUTH_STORE
|
||||
|
||||
|
||||
def auth_status() -> list[dict]:
|
||||
"""Per-provider auth status — never includes a token."""
|
||||
return [AUTH_STORE.status(p) for p in PROVIDERS.values()]
|
||||
@@ -0,0 +1,269 @@
|
||||
"""Generic OAuth 2.0 PKCE engine + transient loopback callback server.
|
||||
|
||||
The flow, per provider:
|
||||
|
||||
1. :func:`start_login_flow` builds a PKCE challenge, binds a loopback callback
|
||||
server on ``127.0.0.1:<CALLBACK_PORT>`` at ``/callback/<provider>``, and
|
||||
returns the provider's authorize URL for the user to open.
|
||||
2. The provider redirects the browser back to the loopback URL with a ``code``
|
||||
and the ``state`` we generated. The server validates ``state``, exchanges the
|
||||
code for a :class:`Token`, hands it to the ``deliver`` sink, and tears down.
|
||||
3. If no callback arrives within :data:`_LOGIN_TIMEOUT`, the server tears down.
|
||||
|
||||
The callback runs on its own bare server, not ComfyUI's main server: the main
|
||||
server rejects cross-site navigations (``Sec-Fetch-Site: cross-site`` → 403),
|
||||
and an OAuth redirect from the provider is exactly such a navigation. The port
|
||||
is fixed because HuggingFace and Civitai require an exact registered
|
||||
``redirect_uri`` (port included); only one login runs at a time so the port
|
||||
never contends with itself.
|
||||
|
||||
Only public PKCE clients are supported (no client secret). All outbound calls
|
||||
go to the provider's own authorize/token endpoints, strictly user-initiated.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
import secrets
|
||||
import time
|
||||
from typing import Callable
|
||||
from urllib.parse import urlencode
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
from app.model_downloader.auth.providers import Provider
|
||||
from app.model_downloader.auth.token_store import Token
|
||||
from app.model_downloader.net.session import get_session, ssl_context
|
||||
|
||||
CALLBACK_HOST = "127.0.0.1"
|
||||
# Fixed loopback port for the OAuth redirect. Must match the redirect URI
|
||||
# registered on the provider's OAuth app; override in lockstep if you change it.
|
||||
CALLBACK_PORT = int(os.environ.get("COMFY_OAUTH_CALLBACK_PORT", "41954"))
|
||||
_LOGIN_TIMEOUT = 300.0 # seconds to wait for the browser callback
|
||||
|
||||
# The auth tab is opened by the frontend via window.open, so window.close() is
|
||||
# allowed here; the visible text is the fallback when the browser blocks it.
|
||||
_SUCCESS_HTML = (
|
||||
"<!doctype html><meta charset=utf-8><title>ComfyUI</title>"
|
||||
"<p>Login successful. You can close this window and return to ComfyUI.</p>"
|
||||
"<script>window.close()</script>"
|
||||
)
|
||||
|
||||
# Token sink: called with the provider name and the exchanged Token.
|
||||
TokenSink = Callable[[str, Token], None]
|
||||
|
||||
|
||||
class OAuthError(Exception):
|
||||
"""A user-facing OAuth failure."""
|
||||
|
||||
|
||||
class OAuthNotConfigured(OAuthError):
|
||||
"""The provider has no public client id configured."""
|
||||
|
||||
|
||||
class LoginInProgress(OAuthError):
|
||||
"""A login flow for this provider is already running."""
|
||||
|
||||
|
||||
def _b64url(data: bytes) -> str:
|
||||
return base64.urlsafe_b64encode(data).rstrip(b"=").decode("ascii")
|
||||
|
||||
|
||||
def _make_pkce() -> tuple[str, str]:
|
||||
"""Return ``(verifier, challenge)`` for the S256 PKCE method."""
|
||||
verifier = _b64url(secrets.token_bytes(32))
|
||||
challenge = _b64url(hashlib.sha256(verifier.encode("ascii")).digest())
|
||||
return verifier, challenge
|
||||
|
||||
|
||||
def build_authorize_url(
|
||||
provider: Provider, challenge: str, state: str, redirect_uri: str
|
||||
) -> str:
|
||||
params = {
|
||||
"response_type": "code",
|
||||
"client_id": provider.client_id,
|
||||
"redirect_uri": redirect_uri,
|
||||
"scope": provider.scope,
|
||||
"state": state,
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
}
|
||||
return f"{provider.authorize_url}?{urlencode(params)}"
|
||||
|
||||
|
||||
def _token_from_payload(payload: dict) -> Token:
|
||||
expires_in = payload.get("expires_in")
|
||||
expires_at = int(time.time()) + int(expires_in) if expires_in else 0
|
||||
return Token(
|
||||
access_token=payload["access_token"],
|
||||
refresh_token=payload.get("refresh_token"),
|
||||
expires_at=expires_at,
|
||||
token_type=payload.get("token_type", "Bearer"),
|
||||
scope=payload.get("scope"),
|
||||
)
|
||||
|
||||
|
||||
async def _post_token(provider: Provider, data: dict) -> Token:
|
||||
session = await get_session()
|
||||
resp = await session.post(
|
||||
provider.token_url,
|
||||
data=data,
|
||||
headers={"Accept": "application/json"},
|
||||
ssl=ssl_context(),
|
||||
)
|
||||
try:
|
||||
if resp.status != 200:
|
||||
body = await resp.text()
|
||||
raise OAuthError(
|
||||
f"{provider.name} token endpoint returned HTTP {resp.status}: {body[:200]}"
|
||||
)
|
||||
payload = await resp.json()
|
||||
finally:
|
||||
await resp.release()
|
||||
if "access_token" not in payload:
|
||||
raise OAuthError(f"{provider.name} token response missing access_token")
|
||||
return _token_from_payload(payload)
|
||||
|
||||
|
||||
async def exchange_code(
|
||||
provider: Provider, code: str, verifier: str, redirect_uri: str
|
||||
) -> Token:
|
||||
return await _post_token(
|
||||
provider,
|
||||
{
|
||||
"grant_type": "authorization_code",
|
||||
"code": code,
|
||||
"redirect_uri": redirect_uri,
|
||||
"client_id": provider.client_id,
|
||||
"code_verifier": verifier,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def refresh_access_token(provider: Provider, token: Token) -> Token:
|
||||
if not token.refresh_token:
|
||||
raise OAuthError(f"{provider.name} token is not refreshable")
|
||||
refreshed = await _post_token(
|
||||
provider,
|
||||
{
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": token.refresh_token,
|
||||
"client_id": provider.client_id,
|
||||
},
|
||||
)
|
||||
# Some providers omit a new refresh token on refresh; keep the old one.
|
||||
if refreshed.refresh_token is None:
|
||||
refreshed.refresh_token = token.refresh_token
|
||||
return refreshed
|
||||
|
||||
|
||||
class _LoginFlow:
|
||||
"""A single in-flight login: owns the loopback server and PKCE state."""
|
||||
|
||||
def __init__(self, provider: Provider, deliver: TokenSink) -> None:
|
||||
self.provider = provider
|
||||
self.deliver = deliver
|
||||
self.verifier, self.challenge = _make_pkce()
|
||||
self.state = secrets.token_urlsafe(24)
|
||||
self.redirect_uri = f"http://{CALLBACK_HOST}:{CALLBACK_PORT}/callback/{provider.name}"
|
||||
self._runner: web.AppRunner | None = None
|
||||
self._timeout_handle: asyncio.TimerHandle | None = None
|
||||
|
||||
async def start(self) -> str:
|
||||
app = web.Application()
|
||||
app.router.add_get("/callback/{provider}", self._handle_callback)
|
||||
self._runner = web.AppRunner(app)
|
||||
await self._runner.setup()
|
||||
site = web.TCPSite(self._runner, CALLBACK_HOST, CALLBACK_PORT, reuse_address=True)
|
||||
try:
|
||||
await site.start()
|
||||
except OSError as e:
|
||||
await self._runner.cleanup()
|
||||
self._runner = None
|
||||
raise OAuthError(f"could not bind callback port {CALLBACK_PORT}: {e}")
|
||||
loop = asyncio.get_running_loop()
|
||||
self._timeout_handle = loop.call_later(
|
||||
_LOGIN_TIMEOUT, lambda: asyncio.ensure_future(self._teardown())
|
||||
)
|
||||
return build_authorize_url(
|
||||
self.provider, self.challenge, self.state, self.redirect_uri
|
||||
)
|
||||
|
||||
async def _handle_callback(self, request: web.Request) -> web.Response:
|
||||
if request.match_info.get("provider") != self.provider.name:
|
||||
return web.Response(text="Unknown login.", content_type="text/plain", status=404)
|
||||
error = request.query.get("error")
|
||||
if error:
|
||||
asyncio.ensure_future(self._teardown())
|
||||
return web.Response(
|
||||
text=f"Login failed: {error}", content_type="text/plain", status=400
|
||||
)
|
||||
if request.query.get("state") != self.state:
|
||||
return web.Response(
|
||||
text="Login failed: state mismatch.",
|
||||
content_type="text/plain",
|
||||
status=400,
|
||||
)
|
||||
code = request.query.get("code")
|
||||
if not code:
|
||||
return web.Response(
|
||||
text="Login failed: no authorization code.",
|
||||
content_type="text/plain",
|
||||
status=400,
|
||||
)
|
||||
try:
|
||||
token = await exchange_code(
|
||||
self.provider, code, self.verifier, self.redirect_uri
|
||||
)
|
||||
except OAuthError as e:
|
||||
logging.warning("[model_downloader] %s login failed: %s", self.provider.name, e)
|
||||
asyncio.ensure_future(self._teardown())
|
||||
return web.Response(
|
||||
text=f"Login failed: {e}", content_type="text/plain", status=502
|
||||
)
|
||||
self.deliver(self.provider.name, token)
|
||||
asyncio.ensure_future(self._teardown())
|
||||
return web.Response(text=_SUCCESS_HTML, content_type="text/html")
|
||||
|
||||
async def _teardown(self) -> None:
|
||||
_ACTIVE.pop(self.provider.name, None)
|
||||
if self._timeout_handle is not None:
|
||||
self._timeout_handle.cancel()
|
||||
self._timeout_handle = None
|
||||
if self._runner is not None:
|
||||
try:
|
||||
await self._runner.cleanup()
|
||||
except Exception:
|
||||
logging.debug("[model_downloader] callback server cleanup error", exc_info=True)
|
||||
self._runner = None
|
||||
|
||||
|
||||
_ACTIVE: dict[str, _LoginFlow] = {}
|
||||
|
||||
|
||||
def login_in_progress(provider_name: str) -> bool:
|
||||
return provider_name in _ACTIVE
|
||||
|
||||
|
||||
async def start_login_flow(provider: Provider, deliver: TokenSink) -> str:
|
||||
"""Begin a login flow and return the authorize URL to open in a browser.
|
||||
|
||||
Binds the fixed-port loopback callback server; only one login may run at a
|
||||
time since that port is shared.
|
||||
"""
|
||||
if not provider.client_id:
|
||||
raise OAuthNotConfigured(
|
||||
f"OAuth app not configured for {provider.name}; set "
|
||||
f"{provider.client_id_env} or use an env API key."
|
||||
)
|
||||
if _ACTIVE:
|
||||
active = next(iter(_ACTIVE))
|
||||
raise LoginInProgress(f"A login for {active} is already in progress.")
|
||||
flow = _LoginFlow(provider, deliver)
|
||||
authorize_url = await flow.start()
|
||||
_ACTIVE[provider.name] = flow
|
||||
return authorize_url
|
||||
@@ -0,0 +1,87 @@
|
||||
"""Provider registry for download authentication.
|
||||
|
||||
A :class:`Provider` describes a hub that can authenticate downloads either from
|
||||
an environment API key or from an OAuth 2.0 access token. Both HuggingFace and
|
||||
Civitai are public PKCE clients, so no client secret is ever stored; the public
|
||||
``client_id`` is a placeholder overridable via env.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
|
||||
def normalize_host(host: str) -> str:
|
||||
"""Lowercase, strip port, IDNA-encode."""
|
||||
if not host:
|
||||
return ""
|
||||
host = host.strip()
|
||||
if "://" in host: # a full URL was pasted — extract just the host
|
||||
host = urlsplit(host).hostname or ""
|
||||
host = host.lower()
|
||||
if host.startswith("[") and "]" in host: # bracketed IPv6 literal
|
||||
host = host[1 : host.index("]")]
|
||||
elif host.count(":") == 1: # host:port (not IPv6)
|
||||
host = host.split(":", 1)[0]
|
||||
try:
|
||||
host = host.encode("idna").decode("ascii")
|
||||
except (UnicodeError, ValueError):
|
||||
pass
|
||||
return host
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Provider:
|
||||
name: str
|
||||
host: str
|
||||
authorize_url: str
|
||||
token_url: str
|
||||
scope: str
|
||||
# Env vars to try, in order, for a plain API key.
|
||||
env_keys: tuple[str, ...]
|
||||
# Env var overriding the public OAuth client id.
|
||||
client_id_env: str
|
||||
# Public PKCE client id. Empty means "not configured" until the env sets it.
|
||||
default_client_id: str = ""
|
||||
|
||||
@property
|
||||
def client_id(self) -> str:
|
||||
return os.environ.get(self.client_id_env, self.default_client_id) or ""
|
||||
|
||||
def env_token(self) -> str | None:
|
||||
for var in self.env_keys:
|
||||
token = os.environ.get(var)
|
||||
if token:
|
||||
return token
|
||||
return None
|
||||
|
||||
|
||||
PROVIDERS: dict[str, Provider] = {
|
||||
"huggingface": Provider(
|
||||
name="huggingface",
|
||||
host="huggingface.co",
|
||||
authorize_url="https://huggingface.co/oauth/authorize",
|
||||
token_url="https://huggingface.co/oauth/token",
|
||||
scope="openid read-repos gated-repos",
|
||||
env_keys=("HF_TOKEN", "HUGGING_FACE_HUB_TOKEN"),
|
||||
client_id_env="COMFY_HF_OAUTH_CLIENT_ID",
|
||||
),
|
||||
"civitai": Provider(
|
||||
name="civitai",
|
||||
host="civitai.com",
|
||||
authorize_url="https://auth.civitai.com/api/auth/oauth/authorize",
|
||||
token_url="https://auth.civitai.com/api/auth/oauth/token",
|
||||
scope="4", # ModelsRead; UserRead is auto-granted
|
||||
env_keys=("CIVITAI_API_TOKEN", "CIVITAI_API_KEY"),
|
||||
client_id_env="COMFY_CIVITAI_OAUTH_CLIENT_ID",
|
||||
),
|
||||
}
|
||||
|
||||
_HOST_TO_PROVIDER = {p.host: p for p in PROVIDERS.values()}
|
||||
|
||||
|
||||
def provider_for_host(host: str) -> Provider | None:
|
||||
"""Return the provider whose host exactly matches ``host`` (normalized)."""
|
||||
return _HOST_TO_PROVIDER.get(normalize_host(host))
|
||||
@@ -0,0 +1,42 @@
|
||||
"""Per-hop auth resolution (https only).
|
||||
|
||||
Recomputed from scratch on every redirect hop: a hop only gets a bearer token
|
||||
when *its own host* matches a configured provider, so a token bound to
|
||||
``huggingface.co`` is silently dropped when the request is redirected to a
|
||||
presigned CDN host — which is exactly what these hubs expect.
|
||||
|
||||
For a matching hop: env API key first, then the provider's OAuth access token
|
||||
(refreshed if expired), else no auth.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from app.model_downloader.auth.providers import provider_for_host
|
||||
from app.model_downloader.auth.store import AUTH_STORE
|
||||
|
||||
|
||||
@dataclass
|
||||
class RequestAuth:
|
||||
"""How to modify a single request to carry a bearer token."""
|
||||
|
||||
headers: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
|
||||
async def resolve_auth_for_hop(host: str, scheme: str) -> RequestAuth | None:
|
||||
"""Resolve the bearer token (if any) to attach for one request hop."""
|
||||
if scheme.lower() != "https":
|
||||
return None
|
||||
provider = provider_for_host(host)
|
||||
if provider is None:
|
||||
return None
|
||||
|
||||
token = provider.env_token()
|
||||
if token:
|
||||
return RequestAuth(headers={"Authorization": f"Bearer {token}"})
|
||||
|
||||
access = await AUTH_STORE.get_valid_token(provider)
|
||||
if access:
|
||||
return RequestAuth(headers={"Authorization": f"Bearer {access}"})
|
||||
return None
|
||||
@@ -0,0 +1,64 @@
|
||||
"""In-memory OAuth token cache over the on-disk token store.
|
||||
|
||||
:data:`AUTH_STORE` is the process singleton the resolver and API talk to. It
|
||||
lazily loads each provider's token from disk, refreshes an expired access token
|
||||
via its refresh token, and orchestrates the login flow (delegating the loopback
|
||||
callback server to :mod:`oauth`).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from app.model_downloader.auth import oauth, token_store
|
||||
from app.model_downloader.auth.providers import Provider
|
||||
from app.model_downloader.auth.token_store import Token
|
||||
|
||||
|
||||
class AuthStore:
|
||||
def __init__(self) -> None:
|
||||
# provider name -> Token, or None when known to be absent. A missing key
|
||||
# means "not yet loaded from disk".
|
||||
self._cache: dict[str, Token | None] = {}
|
||||
|
||||
def _load(self, name: str) -> Token | None:
|
||||
if name not in self._cache:
|
||||
self._cache[name] = token_store.load(name)
|
||||
return self._cache[name]
|
||||
|
||||
def set_token(self, name: str, token: Token) -> None:
|
||||
self._cache[name] = token
|
||||
token_store.save(name, token)
|
||||
|
||||
def clear(self, name: str) -> None:
|
||||
self._cache[name] = None
|
||||
token_store.delete(name)
|
||||
|
||||
async def get_valid_token(self, provider: Provider) -> str | None:
|
||||
"""Return a valid access token string for ``provider``, or ``None``.
|
||||
|
||||
Refreshes an expired token when a refresh token is available.
|
||||
"""
|
||||
token = self._load(provider.name)
|
||||
if token is None or not token.access_token:
|
||||
return None
|
||||
if token.is_expired():
|
||||
if not token.refresh_token:
|
||||
return None
|
||||
token = await oauth.refresh_access_token(provider, token)
|
||||
self.set_token(provider.name, token)
|
||||
return token.access_token
|
||||
|
||||
async def begin_login(self, provider: Provider) -> str:
|
||||
"""Start a login flow; returns the authorize URL to open in a browser."""
|
||||
return await oauth.start_login_flow(provider, self.set_token)
|
||||
|
||||
def status(self, provider: Provider) -> dict:
|
||||
token = self._load(provider.name)
|
||||
return {
|
||||
"provider": provider.name,
|
||||
"logged_in": token is not None and bool(token.access_token),
|
||||
"login_in_progress": oauth.login_in_progress(provider.name),
|
||||
"env_key_present": provider.env_token() is not None,
|
||||
}
|
||||
|
||||
|
||||
AUTH_STORE = AuthStore()
|
||||
@@ -0,0 +1,186 @@
|
||||
"""On-disk OAuth token persistence — one machine-bound blob per provider.
|
||||
|
||||
Tokens live under ``folder_paths.get_system_user_directory("download_auth")``,
|
||||
never in the SQLite DB. Each provider file is written ``0600`` and holds an
|
||||
opaque blob, not readable JSON: the token JSON is XORed with an HMAC-SHA256
|
||||
keystream whose key is derived from stable machine/install attributes plus a
|
||||
per-install random salt.
|
||||
|
||||
This is obfuscation, not confidentiality. It stops a token from being read by a
|
||||
human browsing files, grepped out of a backup, or lifted from a folder copied to
|
||||
another machine (the blob won't decrypt off its origin machine). It does not
|
||||
protect against code running inside this process (custom nodes) or an attacker
|
||||
who reads this source and recomputes the key.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import getpass
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import os
|
||||
import platform
|
||||
import secrets
|
||||
import time
|
||||
from dataclasses import asdict, dataclass
|
||||
|
||||
import folder_paths
|
||||
|
||||
_SALT_FILE = ".salt"
|
||||
_SALT_LEN = 32
|
||||
_NONCE_LEN = 16
|
||||
_PBKDF2_ITERS = 200_000
|
||||
|
||||
|
||||
@dataclass
|
||||
class Token:
|
||||
access_token: str
|
||||
refresh_token: str | None = None
|
||||
# Epoch seconds when the access token expires; 0 means "unknown / no expiry".
|
||||
expires_at: int = 0
|
||||
token_type: str = "Bearer"
|
||||
scope: str | None = None
|
||||
|
||||
def is_expired(self, skew: int = 60) -> bool:
|
||||
if not self.expires_at:
|
||||
return False
|
||||
return time.time() + skew >= self.expires_at
|
||||
|
||||
|
||||
def _auth_dir() -> str:
|
||||
path = folder_paths.get_system_user_directory("download_auth")
|
||||
os.makedirs(path, exist_ok=True)
|
||||
return path
|
||||
|
||||
|
||||
def _token_path(provider: str) -> str:
|
||||
return os.path.join(_auth_dir(), f"{provider}.bin")
|
||||
|
||||
|
||||
def _machine_id() -> bytes:
|
||||
"""A stable per-machine identifier, best-effort across platforms."""
|
||||
for path in ("/etc/machine-id", "/var/lib/dbus/machine-id"):
|
||||
try:
|
||||
with open(path, "rb") as f:
|
||||
return f.read().strip()
|
||||
except OSError:
|
||||
pass
|
||||
if os.name == "nt":
|
||||
import winreg
|
||||
|
||||
try:
|
||||
key = winreg.OpenKey(
|
||||
winreg.HKEY_LOCAL_MACHINE, r"SOFTWARE\Microsoft\Cryptography"
|
||||
)
|
||||
try:
|
||||
guid, _ = winreg.QueryValueEx(key, "MachineGuid")
|
||||
return str(guid).encode("utf-8")
|
||||
finally:
|
||||
winreg.CloseKey(key)
|
||||
except OSError:
|
||||
pass
|
||||
return platform.node().encode("utf-8")
|
||||
|
||||
|
||||
def _machine_material(auth_dir: str) -> bytes:
|
||||
try:
|
||||
user = getpass.getuser()
|
||||
except Exception:
|
||||
user = ""
|
||||
parts = (_machine_id(), platform.node().encode("utf-8"), user.encode("utf-8"), auth_dir.encode("utf-8"))
|
||||
return b"\x00".join(parts)
|
||||
|
||||
|
||||
def _load_or_create_salt(auth_dir: str) -> bytes | None:
|
||||
path = os.path.join(auth_dir, _SALT_FILE)
|
||||
try:
|
||||
with open(path, "rb") as f:
|
||||
salt = f.read()
|
||||
if len(salt) == _SALT_LEN:
|
||||
return salt
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except OSError:
|
||||
return None
|
||||
salt = secrets.token_bytes(_SALT_LEN)
|
||||
fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
|
||||
with os.fdopen(fd, "wb") as f:
|
||||
f.write(salt)
|
||||
os.chmod(path, 0o600)
|
||||
return salt
|
||||
|
||||
|
||||
def _derive_key(salt: bytes, auth_dir: str) -> bytes:
|
||||
return hashlib.pbkdf2_hmac(
|
||||
"sha256", _machine_material(auth_dir), salt, _PBKDF2_ITERS, dklen=32
|
||||
)
|
||||
|
||||
|
||||
def _keystream(key: bytes, nonce: bytes, n: int) -> bytes:
|
||||
out = bytearray()
|
||||
counter = 0
|
||||
while len(out) < n:
|
||||
out.extend(hmac.new(key, nonce + counter.to_bytes(8, "big"), hashlib.sha256).digest())
|
||||
counter += 1
|
||||
return bytes(out[:n])
|
||||
|
||||
|
||||
def _xor(data: bytes, stream: bytes) -> bytes:
|
||||
return bytes(a ^ b for a, b in zip(data, stream))
|
||||
|
||||
|
||||
def load(provider: str) -> Token | None:
|
||||
auth_dir = _auth_dir()
|
||||
try:
|
||||
with open(_token_path(provider), "rb") as f:
|
||||
blob = base64.b64decode(f.read())
|
||||
except FileNotFoundError:
|
||||
return None
|
||||
except (ValueError, OSError):
|
||||
return None
|
||||
salt = _load_or_create_salt(auth_dir)
|
||||
if salt is None or len(blob) <= _NONCE_LEN:
|
||||
return None
|
||||
nonce, ciphertext = blob[:_NONCE_LEN], blob[_NONCE_LEN:]
|
||||
key = _derive_key(salt, auth_dir)
|
||||
plaintext = _xor(ciphertext, _keystream(key, nonce, len(ciphertext)))
|
||||
# A wrong machine / corrupt file decrypts to garbage; treat as logged out.
|
||||
try:
|
||||
data = json.loads(plaintext)
|
||||
except ValueError:
|
||||
return None
|
||||
if not isinstance(data, dict) or "access_token" not in data:
|
||||
return None
|
||||
return Token(
|
||||
access_token=data.get("access_token", ""),
|
||||
refresh_token=data.get("refresh_token"),
|
||||
expires_at=int(data.get("expires_at", 0) or 0),
|
||||
token_type=data.get("token_type", "Bearer"),
|
||||
scope=data.get("scope"),
|
||||
)
|
||||
|
||||
|
||||
def save(provider: str, token: Token) -> None:
|
||||
auth_dir = _auth_dir()
|
||||
salt = _load_or_create_salt(auth_dir)
|
||||
if salt is None:
|
||||
return
|
||||
key = _derive_key(salt, auth_dir)
|
||||
nonce = secrets.token_bytes(_NONCE_LEN)
|
||||
plaintext = json.dumps(asdict(token)).encode("utf-8")
|
||||
ciphertext = _xor(plaintext, _keystream(key, nonce, len(plaintext)))
|
||||
blob = base64.b64encode(nonce + ciphertext)
|
||||
path = _token_path(provider)
|
||||
fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
|
||||
with os.fdopen(fd, "wb") as f:
|
||||
f.write(blob)
|
||||
os.chmod(path, 0o600)
|
||||
|
||||
|
||||
def delete(provider: str) -> None:
|
||||
try:
|
||||
os.remove(_token_path(provider))
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
@@ -0,0 +1,32 @@
|
||||
"""Shared constants for the download manager.
|
||||
|
||||
Status values are persisted as TEXT in the ``downloads`` table; keep them
|
||||
stable. The lifecycle is:
|
||||
|
||||
queued -> active -> verifying -> completed
|
||||
| |-> paused -> (resume) -> active
|
||||
| |-> failed (network, retryable) -> queued (backoff)
|
||||
|-> cancelled
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class DownloadStatus:
|
||||
QUEUED = "queued"
|
||||
ACTIVE = "active"
|
||||
PAUSED = "paused"
|
||||
VERIFYING = "verifying"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
#: States from which a worker is doing (or about to do) network I/O.
|
||||
LIVE = (QUEUED, ACTIVE, VERIFYING)
|
||||
#: Terminal states — the job will not transition again on its own.
|
||||
TERMINAL = (COMPLETED, FAILED, CANCELLED)
|
||||
|
||||
|
||||
# Default temp-file suffix. Distinctive so the startup orphan sweep only
|
||||
# removes files THIS subsystem created, never unrelated *.tmp files.
|
||||
TMP_SUFFIX = ".comfy-download.part"
|
||||
@@ -0,0 +1,125 @@
|
||||
"""SQLAlchemy models for the download manager.
|
||||
|
||||
Two tables:
|
||||
|
||||
- ``downloads`` one row per requested file (job + queue state).
|
||||
- ``download_segments`` per-segment byte progress, for segmented resume.
|
||||
|
||||
On completion a finished file is registered into the assets catalog;
|
||||
``downloads`` is kept only as job history.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import (
|
||||
BigInteger,
|
||||
Boolean,
|
||||
CheckConstraint,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
String,
|
||||
Text,
|
||||
)
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from app.database.models import Base
|
||||
|
||||
|
||||
def _uuid() -> str:
|
||||
return str(uuid.uuid4())
|
||||
|
||||
|
||||
def _now() -> int:
|
||||
return int(time.time())
|
||||
|
||||
|
||||
class Download(Base):
|
||||
__tablename__ = "downloads"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=_uuid)
|
||||
# Original requested URL and the final URL after validated redirects.
|
||||
url: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
final_url: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
# Canonical "<directory>/<filename>" identifier (resolved via folder_paths).
|
||||
model_id: Mapped[str] = mapped_column(String(1024), nullable=False)
|
||||
# Final on-disk location and the .part write target.
|
||||
dest_path: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
temp_path: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
|
||||
status: Mapped[str] = mapped_column(String(16), nullable=False)
|
||||
priority: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||
|
||||
total_bytes: Mapped[int | None] = mapped_column(BigInteger, nullable=True)
|
||||
bytes_done: Mapped[int] = mapped_column(BigInteger, nullable=False, default=0)
|
||||
|
||||
etag: Mapped[str | None] = mapped_column(String(512), nullable=True)
|
||||
last_modified: Mapped[str | None] = mapped_column(String(128), nullable=True)
|
||||
accept_ranges: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
|
||||
# Optional hub-provided checksum to verify against (NOT the dedup key).
|
||||
expected_sha256: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
|
||||
allow_any_extension: Mapped[bool] = mapped_column(
|
||||
Boolean, nullable=False, default=False
|
||||
)
|
||||
# How many retryable failures we have seen (for backoff capping).
|
||||
attempts: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||
|
||||
error: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
created_at: Mapped[int] = mapped_column(BigInteger, nullable=False, default=_now)
|
||||
updated_at: Mapped[int] = mapped_column(
|
||||
BigInteger, nullable=False, default=_now, onupdate=_now
|
||||
)
|
||||
|
||||
segments: Mapped[list[DownloadSegment]] = relationship(
|
||||
"DownloadSegment",
|
||||
back_populates="download",
|
||||
cascade="all,delete-orphan",
|
||||
passive_deletes=True,
|
||||
order_by="DownloadSegment.idx",
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_downloads_status", "status"),
|
||||
Index("ix_downloads_priority", "priority"),
|
||||
Index("ix_downloads_model_id", "model_id"),
|
||||
CheckConstraint("bytes_done >= 0", name="ck_downloads_bytes_done_nonneg"),
|
||||
CheckConstraint(
|
||||
"total_bytes IS NULL OR total_bytes >= 0",
|
||||
name="ck_downloads_total_bytes_nonneg",
|
||||
),
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Download id={self.id} model_id={self.model_id!r} status={self.status}>"
|
||||
|
||||
|
||||
class DownloadSegment(Base):
|
||||
__tablename__ = "download_segments"
|
||||
|
||||
download_id: Mapped[str] = mapped_column(
|
||||
String(36),
|
||||
ForeignKey("downloads.id", ondelete="CASCADE"),
|
||||
primary_key=True,
|
||||
)
|
||||
idx: Mapped[int] = mapped_column(Integer, primary_key=True)
|
||||
start_offset: Mapped[int] = mapped_column(BigInteger, nullable=False)
|
||||
end_offset: Mapped[int] = mapped_column(BigInteger, nullable=False)
|
||||
bytes_done: Mapped[int] = mapped_column(BigInteger, nullable=False, default=0)
|
||||
|
||||
download: Mapped[Download] = relationship("Download", back_populates="segments")
|
||||
|
||||
__table_args__ = (
|
||||
CheckConstraint("bytes_done >= 0", name="ck_segments_bytes_done_nonneg"),
|
||||
CheckConstraint("end_offset >= start_offset", name="ck_segments_range"),
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (
|
||||
f"<DownloadSegment {self.download_id}#{self.idx} "
|
||||
f"{self.start_offset}-{self.end_offset} done={self.bytes_done}>"
|
||||
)
|
||||
@@ -0,0 +1,177 @@
|
||||
"""Synchronous DB access for the download manager.
|
||||
|
||||
All functions open their own short-lived session via ``create_session`` and
|
||||
commit before returning, mirroring ``app/assets`` usage. They are blocking
|
||||
(SQLite) and should be called from async code through ``asyncio.to_thread``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import delete, select
|
||||
|
||||
from app.database.db import create_session
|
||||
from app.model_downloader.constants import DownloadStatus
|
||||
from app.model_downloader.database.models import Download, DownloadSegment
|
||||
|
||||
|
||||
# ----- downloads -----
|
||||
|
||||
|
||||
def insert_download(values: dict) -> None:
|
||||
with create_session() as session:
|
||||
session.add(Download(**values))
|
||||
session.commit()
|
||||
|
||||
|
||||
def get_download(download_id: str) -> Optional[Download]:
|
||||
with create_session() as session:
|
||||
row = session.get(Download, download_id)
|
||||
if row is not None:
|
||||
session.expunge_all()
|
||||
return row
|
||||
|
||||
|
||||
def list_downloads() -> list[Download]:
|
||||
with create_session() as session:
|
||||
rows = list(
|
||||
session.execute(
|
||||
select(Download).order_by(Download.created_at.desc())
|
||||
).scalars()
|
||||
)
|
||||
session.expunge_all()
|
||||
return rows
|
||||
|
||||
|
||||
def has_live_download_for_model(
|
||||
model_id: str, live_statuses: tuple[str, ...], exclude_id: Optional[str] = None
|
||||
) -> bool:
|
||||
with create_session() as session:
|
||||
stmt = select(Download.id).where(
|
||||
Download.model_id == model_id,
|
||||
Download.status.in_(live_statuses),
|
||||
).limit(1)
|
||||
if exclude_id is not None:
|
||||
stmt = stmt.where(Download.id != exclude_id)
|
||||
return session.execute(stmt).first() is not None
|
||||
|
||||
|
||||
def list_segments(download_id: str) -> list[DownloadSegment]:
|
||||
with create_session() as session:
|
||||
rows = list(
|
||||
session.execute(
|
||||
select(DownloadSegment)
|
||||
.where(DownloadSegment.download_id == download_id)
|
||||
.order_by(DownloadSegment.idx)
|
||||
).scalars()
|
||||
)
|
||||
session.expunge_all()
|
||||
return rows
|
||||
|
||||
|
||||
def update_download(download_id: str, **fields) -> None:
|
||||
if not fields:
|
||||
return
|
||||
fields.setdefault("updated_at", int(time.time()))
|
||||
with create_session() as session:
|
||||
row = session.get(Download, download_id)
|
||||
if row is None:
|
||||
return
|
||||
for key, value in fields.items():
|
||||
setattr(row, key, value)
|
||||
session.commit()
|
||||
|
||||
|
||||
def delete_download(download_id: str) -> None:
|
||||
with create_session() as session:
|
||||
row = session.get(Download, download_id)
|
||||
if row is not None:
|
||||
session.delete(row)
|
||||
session.commit()
|
||||
|
||||
|
||||
def delete_downloads(download_ids: list[str]) -> int:
|
||||
"""Delete many downloads in one transaction; returns the number removed.
|
||||
|
||||
Uses a bulk ``DELETE ... WHERE id IN (...)``. Segment rows are removed by
|
||||
the ``ON DELETE CASCADE`` foreign key (SQLite ``PRAGMA foreign_keys=ON`` is
|
||||
set in ``app/database/db.py``), so this stays consistent without loading the
|
||||
ORM relationship.
|
||||
"""
|
||||
if not download_ids:
|
||||
return 0
|
||||
with create_session() as session:
|
||||
result = session.execute(
|
||||
delete(Download).where(Download.id.in_(download_ids))
|
||||
)
|
||||
session.commit()
|
||||
return result.rowcount or 0
|
||||
|
||||
|
||||
def replace_segments(download_id: str, segments: list[dict]) -> None:
|
||||
"""Atomically replace the segment plan for a download."""
|
||||
with create_session() as session:
|
||||
session.query(DownloadSegment).filter(
|
||||
DownloadSegment.download_id == download_id
|
||||
).delete()
|
||||
for seg in segments:
|
||||
session.add(DownloadSegment(download_id=download_id, **seg))
|
||||
session.commit()
|
||||
|
||||
|
||||
def update_segment_progress(download_id: str, idx: int, bytes_done: int) -> None:
|
||||
with create_session() as session:
|
||||
row = session.get(DownloadSegment, {"download_id": download_id, "idx": idx})
|
||||
if row is None:
|
||||
return
|
||||
row.bytes_done = bytes_done
|
||||
session.commit()
|
||||
|
||||
|
||||
def list_queued_downloads() -> list[Download]:
|
||||
"""Queued rows ordered for admission (priority desc, then FIFO)."""
|
||||
with create_session() as session:
|
||||
rows = list(
|
||||
session.execute(
|
||||
select(Download)
|
||||
.where(Download.status == DownloadStatus.QUEUED)
|
||||
.order_by(Download.priority.desc(), Download.created_at.asc())
|
||||
).scalars()
|
||||
)
|
||||
session.expunge_all()
|
||||
return rows
|
||||
|
||||
|
||||
def reconcile_live_downloads() -> list[Download]:
|
||||
"""Reset any ``active``/``verifying`` rows left by a previous run.
|
||||
|
||||
On a clean restart there can be no live worker, so anything still marked
|
||||
live is stale. Move it back to ``queued`` (offsets are preserved on the
|
||||
segment rows) so the scheduler re-admits it. Returns the rows that should
|
||||
be re-queued by the scheduler (queued + paused).
|
||||
"""
|
||||
with create_session() as session:
|
||||
stale = list(
|
||||
session.execute(
|
||||
select(Download).where(
|
||||
Download.status.in_([DownloadStatus.ACTIVE, DownloadStatus.VERIFYING])
|
||||
)
|
||||
).scalars()
|
||||
)
|
||||
now = int(time.time())
|
||||
for row in stale:
|
||||
row.status = DownloadStatus.QUEUED
|
||||
row.updated_at = now
|
||||
session.commit()
|
||||
|
||||
resumable = list(
|
||||
session.execute(
|
||||
select(Download)
|
||||
.where(Download.status == DownloadStatus.QUEUED)
|
||||
.order_by(Download.priority.desc(), Download.created_at.asc())
|
||||
).scalars()
|
||||
)
|
||||
session.expunge_all()
|
||||
return resumable
|
||||
@@ -0,0 +1,611 @@
|
||||
"""The per-download worker.
|
||||
|
||||
One :class:`DownloadJob` drives a single file from probe to verified, cataloged
|
||||
completion. It supports cooperative pause / resume / cancel, segmented
|
||||
multi-connection transfer with positioned writes, and a verification gate
|
||||
(size + structural + optional sha256) before the atomic rename into place.
|
||||
|
||||
Control is cooperative: external callers flip ``_control`` via
|
||||
:meth:`request_pause` / :meth:`request_cancel`; segment loops observe it between
|
||||
chunks and raise, which unwinds cleanly and persists resume offsets.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Optional
|
||||
|
||||
from comfy.cli_args import args
|
||||
from app.model_downloader.constants import DownloadStatus
|
||||
from app.model_downloader.database import queries
|
||||
from app.model_downloader.engine.planner import (
|
||||
effective_segment_count,
|
||||
plan_segments,
|
||||
)
|
||||
from app.model_downloader.engine.writer import FileWriter
|
||||
from app.model_downloader.net.http import open_validated, redact_url
|
||||
from app.model_downloader.net.probe import gated_error_message, probe
|
||||
from app.model_downloader.verify import checksum, dedup, structural
|
||||
|
||||
_RETRYABLE_STATUSES = {408, 429, 500, 502, 503, 504}
|
||||
_PERSIST_INTERVAL = 2.0 # seconds between throttled progress persists
|
||||
|
||||
|
||||
class Paused(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class Cancelled(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class RemoteChanged(Exception):
|
||||
"""The remote file changed under a resume (got 200 where 206 expected)."""
|
||||
|
||||
|
||||
class RetryableError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class FatalError(Exception):
|
||||
"""Non-retryable: 4xx, checksum mismatch, structural failure, gated, etc."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class SegmentRuntime:
|
||||
idx: int
|
||||
start: int
|
||||
end: int # inclusive; may be -1 for unknown-size single stream
|
||||
bytes_done: int = 0
|
||||
|
||||
@property
|
||||
def length(self) -> int:
|
||||
return self.end - self.start + 1
|
||||
|
||||
|
||||
@dataclass
|
||||
class RuntimeState:
|
||||
download_id: str
|
||||
model_id: str
|
||||
url: str
|
||||
priority: int
|
||||
status: str
|
||||
total_bytes: Optional[int] = None
|
||||
bytes_done: int = 0
|
||||
error: Optional[str] = None
|
||||
segments: list[SegmentRuntime] = field(default_factory=list)
|
||||
started_at: float = field(default_factory=time.monotonic)
|
||||
_last_bytes: int = 0
|
||||
_last_time: float = field(default_factory=time.monotonic)
|
||||
speed_bps: float = 0.0
|
||||
|
||||
@property
|
||||
def progress(self) -> Optional[float]:
|
||||
if not self.total_bytes:
|
||||
return None
|
||||
return min(1.0, self.bytes_done / self.total_bytes)
|
||||
|
||||
@property
|
||||
def eta_seconds(self) -> Optional[float]:
|
||||
if not self.total_bytes or self.speed_bps <= 0:
|
||||
return None
|
||||
remaining = max(0, self.total_bytes - self.bytes_done)
|
||||
return remaining / self.speed_bps
|
||||
|
||||
|
||||
@dataclass
|
||||
class JobSpec:
|
||||
download_id: str
|
||||
url: str
|
||||
model_id: str
|
||||
dest_path: str
|
||||
temp_path: str
|
||||
priority: int = 0
|
||||
expected_sha256: Optional[str] = None
|
||||
allow_any_extension: bool = False
|
||||
etag: Optional[str] = None
|
||||
attempts: int = 0
|
||||
|
||||
|
||||
class DownloadJob:
|
||||
def __init__(
|
||||
self, spec: JobSpec, notify_cb: Optional[Callable[[str], None]] = None
|
||||
) -> None:
|
||||
self.spec = spec
|
||||
self._notify = notify_cb
|
||||
self._control = "run" # run | pause | cancel
|
||||
self.state = RuntimeState(
|
||||
download_id=spec.download_id,
|
||||
model_id=spec.model_id,
|
||||
url=spec.url,
|
||||
priority=spec.priority,
|
||||
status=DownloadStatus.QUEUED,
|
||||
)
|
||||
self._writer: Optional[FileWriter] = None
|
||||
self._etag: Optional[str] = spec.etag
|
||||
self._last_persist = 0.0
|
||||
|
||||
# ----- external control -----
|
||||
|
||||
def request_pause(self) -> None:
|
||||
if self._control == "run":
|
||||
self._control = "pause"
|
||||
|
||||
def request_cancel(self) -> None:
|
||||
self._control = "cancel"
|
||||
|
||||
def _check_control(self) -> None:
|
||||
if self._control == "cancel":
|
||||
raise Cancelled()
|
||||
if self._control == "pause":
|
||||
raise Paused()
|
||||
|
||||
# ----- lifecycle -----
|
||||
|
||||
async def run(self) -> str:
|
||||
"""Run to a terminal/paused state; returns the final status string."""
|
||||
await self._set_status(DownloadStatus.ACTIVE, error=None)
|
||||
try:
|
||||
pr = await self._probe_and_plan()
|
||||
await self._transfer(pr)
|
||||
await self._finalize()
|
||||
await self._set_status(DownloadStatus.COMPLETED)
|
||||
except Paused:
|
||||
await self._persist_progress(force=True)
|
||||
await self._set_status(DownloadStatus.PAUSED)
|
||||
except Cancelled:
|
||||
await self._close_writer()
|
||||
self._remove_temp()
|
||||
await self._set_status(DownloadStatus.CANCELLED)
|
||||
except RemoteChanged:
|
||||
await self._reset_for_restart()
|
||||
await self._set_status(
|
||||
DownloadStatus.QUEUED, error="remote file changed; restarting"
|
||||
)
|
||||
except RetryableError as e:
|
||||
await self._persist_progress(force=True)
|
||||
await self._set_status(DownloadStatus.QUEUED, error=str(e))
|
||||
except FatalError as e:
|
||||
await self._close_writer()
|
||||
self._remove_temp()
|
||||
await self._set_status(DownloadStatus.FAILED, error=str(e))
|
||||
except Exception as e: # unexpected -> treat as retryable
|
||||
logging.warning(
|
||||
"[model_downloader] %s unexpected error: %s",
|
||||
self.spec.model_id, e, exc_info=True,
|
||||
)
|
||||
await self._persist_progress(force=True)
|
||||
await self._set_status(DownloadStatus.QUEUED, error=f"{type(e).__name__}: {e}")
|
||||
finally:
|
||||
await self._close_writer()
|
||||
return self.state.status
|
||||
|
||||
# ----- probe + plan -----
|
||||
|
||||
async def _probe_and_plan(self):
|
||||
pr = await probe(self.spec.url)
|
||||
if not pr.ok:
|
||||
if pr.gated:
|
||||
raise FatalError(gated_error_message(self.spec.url, pr))
|
||||
if pr.status == 0 or pr.status in _RETRYABLE_STATUSES:
|
||||
raise RetryableError(pr.error or "probe failed")
|
||||
raise FatalError(pr.error or f"probe returned HTTP {pr.status}")
|
||||
|
||||
max_bytes = self._max_download_bytes()
|
||||
if max_bytes is not None and pr.total_bytes is not None and pr.total_bytes > max_bytes:
|
||||
raise FatalError(
|
||||
f"file size {pr.total_bytes} exceeds the maximum allowed "
|
||||
f"download size {max_bytes} (--download-max-bytes)"
|
||||
)
|
||||
|
||||
self._etag = pr.etag or self._etag
|
||||
self.state.total_bytes = pr.total_bytes
|
||||
await asyncio.to_thread(
|
||||
queries.update_download,
|
||||
self.spec.download_id,
|
||||
final_url=pr.final_url,
|
||||
total_bytes=pr.total_bytes,
|
||||
accept_ranges=pr.accept_ranges,
|
||||
etag=pr.etag,
|
||||
last_modified=pr.last_modified,
|
||||
)
|
||||
|
||||
seg_count = effective_segment_count(
|
||||
pr.total_bytes, pr.accept_ranges, max(1, args.download_segments)
|
||||
)
|
||||
existing = await asyncio.to_thread(queries.list_segments, self.spec.download_id)
|
||||
can_resume_segmented = (
|
||||
seg_count > 1
|
||||
and existing
|
||||
and pr.total_bytes is not None
|
||||
and existing[-1].end_offset == pr.total_bytes - 1
|
||||
)
|
||||
if can_resume_segmented and not self._segmented_part_valid(pr.total_bytes):
|
||||
# The persisted per-segment offsets describe bytes in a preallocated
|
||||
# .part that is now gone or the wrong size (e.g. the partial of a
|
||||
# failed download was swept on restart, or removed by a fatal
|
||||
# error). Trusting them would skip already-"complete" segments and
|
||||
# leave zero-filled holes. Discard the offsets and re-plan fresh.
|
||||
logging.info(
|
||||
"[model_downloader] %s discarding segmented resume offsets "
|
||||
"(preallocated .part missing or wrong size); restarting",
|
||||
self.spec.model_id,
|
||||
)
|
||||
self._remove_temp()
|
||||
await asyncio.to_thread(
|
||||
queries.replace_segments, self.spec.download_id, []
|
||||
)
|
||||
await asyncio.to_thread(
|
||||
queries.update_download, self.spec.download_id, bytes_done=0
|
||||
)
|
||||
existing = []
|
||||
can_resume_segmented = False
|
||||
|
||||
if can_resume_segmented:
|
||||
# Resume an existing segmented plan.
|
||||
self.state.segments = [
|
||||
SegmentRuntime(s.idx, s.start_offset, s.end_offset, s.bytes_done)
|
||||
for s in existing
|
||||
]
|
||||
elif seg_count > 1 and pr.total_bytes is not None:
|
||||
plans = plan_segments(pr.total_bytes, seg_count)
|
||||
await asyncio.to_thread(
|
||||
queries.replace_segments,
|
||||
self.spec.download_id,
|
||||
[
|
||||
{"idx": p.idx, "start_offset": p.start, "end_offset": p.end, "bytes_done": 0}
|
||||
for p in plans
|
||||
],
|
||||
)
|
||||
self.state.segments = [SegmentRuntime(p.idx, p.start, p.end, 0) for p in plans]
|
||||
else:
|
||||
# Single-stream: one logical segment; bytes_done tracked on the row.
|
||||
row = await asyncio.to_thread(queries.get_download, self.spec.download_id)
|
||||
resume_from = row.bytes_done if row else 0
|
||||
end = (pr.total_bytes - 1) if pr.total_bytes else -1
|
||||
# ``row.bytes_done`` may be the SUM of per-segment offsets from a
|
||||
# prior segmented run (a preallocated, non-contiguous .part). A
|
||||
# single-stream resume writes a contiguous prefix, so the offset is
|
||||
# only trustworthy when the on-disk file is exactly that many
|
||||
# contiguous bytes. This guards the case where a download that ran
|
||||
# segmented now resolves to one segment (server dropped
|
||||
# Accept-Ranges, or --download-segments was lowered between runs):
|
||||
# resuming over non-contiguous data would corrupt the output.
|
||||
if resume_from > 0 and not self._contiguous_prefix_valid(resume_from):
|
||||
logging.info(
|
||||
"[model_downloader] %s discarding untrusted resume offset "
|
||||
"%d (on-disk .part not a contiguous prefix); restarting",
|
||||
self.spec.model_id, resume_from,
|
||||
)
|
||||
resume_from = 0
|
||||
self._remove_temp()
|
||||
if await asyncio.to_thread(queries.list_segments, self.spec.download_id):
|
||||
await asyncio.to_thread(
|
||||
queries.replace_segments, self.spec.download_id, []
|
||||
)
|
||||
await asyncio.to_thread(
|
||||
queries.update_download, self.spec.download_id, bytes_done=0
|
||||
)
|
||||
self.state.segments = [SegmentRuntime(0, 0, end, resume_from)]
|
||||
self._recompute_bytes_done()
|
||||
return pr
|
||||
|
||||
# ----- transfer -----
|
||||
|
||||
async def _transfer(self, pr) -> None:
|
||||
self._writer = FileWriter(self.spec.temp_path)
|
||||
await self._writer.open()
|
||||
|
||||
segmented = len(self.state.segments) > 1
|
||||
if segmented and self.state.total_bytes:
|
||||
await self._writer.preallocate(self.state.total_bytes)
|
||||
await self._run_segmented()
|
||||
else:
|
||||
await self._run_single()
|
||||
|
||||
await self._writer.flush()
|
||||
|
||||
async def _run_segmented(self) -> None:
|
||||
pending = [
|
||||
asyncio.ensure_future(self._run_segment(seg))
|
||||
for seg in self.state.segments
|
||||
if seg.bytes_done < seg.length
|
||||
]
|
||||
if not pending:
|
||||
return
|
||||
done, not_done = await asyncio.wait(
|
||||
pending, return_when=asyncio.FIRST_EXCEPTION
|
||||
)
|
||||
first_exc: Optional[BaseException] = None
|
||||
for task in done:
|
||||
exc = task.exception()
|
||||
if exc is not None and first_exc is None:
|
||||
first_exc = exc
|
||||
if first_exc is not None:
|
||||
for task in not_done:
|
||||
task.cancel()
|
||||
await asyncio.gather(*not_done, return_exceptions=True)
|
||||
raise first_exc
|
||||
|
||||
async def _run_segment(self, seg: SegmentRuntime) -> None:
|
||||
offset = seg.start + seg.bytes_done
|
||||
headers = {
|
||||
"Range": f"bytes={offset}-{seg.end}",
|
||||
"Accept-Encoding": "identity",
|
||||
}
|
||||
if self._etag:
|
||||
headers["If-Range"] = self._etag
|
||||
async with open_validated(
|
||||
"GET", self.spec.url, headers=headers
|
||||
) as (resp, _final):
|
||||
if resp.status == 200:
|
||||
# Server ignored the range -> remote changed / no resume support.
|
||||
raise RemoteChanged()
|
||||
if resp.status not in (206,):
|
||||
self._raise_for_status(resp.status)
|
||||
async for chunk in resp.content.iter_chunked(args.download_chunk_size):
|
||||
self._check_control()
|
||||
# Never write past this segment's planned range: a
|
||||
# non-conforming 206 that returns more than the requested
|
||||
# bytes would otherwise overrun adjacent segments and the
|
||||
# preallocated file. Cap the write and abort on overflow.
|
||||
remaining = seg.length - seg.bytes_done
|
||||
if remaining <= 0:
|
||||
raise FatalError(
|
||||
f"segment {seg.idx}: server returned more than the "
|
||||
f"requested {seg.length} bytes"
|
||||
)
|
||||
overflow = len(chunk) > remaining
|
||||
if overflow:
|
||||
chunk = chunk[:remaining]
|
||||
await self._writer.write_at(offset, chunk)
|
||||
offset += len(chunk)
|
||||
seg.bytes_done += len(chunk)
|
||||
self._recompute_bytes_done()
|
||||
await self._persist_progress()
|
||||
if overflow:
|
||||
raise FatalError(
|
||||
f"segment {seg.idx}: server returned more than the "
|
||||
f"requested {seg.length} bytes"
|
||||
)
|
||||
|
||||
async def _run_single(self) -> None:
|
||||
seg = self.state.segments[0]
|
||||
offset = seg.bytes_done # resume from here for single-stream
|
||||
headers = {"Accept-Encoding": "identity"}
|
||||
if offset > 0:
|
||||
headers["Range"] = f"bytes={offset}-"
|
||||
if self._etag:
|
||||
headers["If-Range"] = self._etag
|
||||
async with open_validated(
|
||||
"GET", self.spec.url, headers=headers
|
||||
) as (resp, _final):
|
||||
if offset > 0 and resp.status == 200:
|
||||
# Resume not honoured -> start over from the beginning. Truncate
|
||||
# the existing partial so stale trailing bytes from the prior
|
||||
# attempt cannot survive past the new (possibly shorter) end.
|
||||
offset = 0
|
||||
seg.bytes_done = 0
|
||||
self.state.bytes_done = 0
|
||||
await self._writer.truncate(0)
|
||||
elif offset > 0 and resp.status != 206:
|
||||
self._raise_for_status(resp.status)
|
||||
elif offset == 0 and resp.status != 200:
|
||||
self._raise_for_status(resp.status)
|
||||
# Byte ceiling for this stream: the known total when the server
|
||||
# reported a size, otherwise the configured maximum download size.
|
||||
# Without a bound, a non-conforming response or an unknown-length
|
||||
# stream (end == -1) that never closes could fill the disk (DoS).
|
||||
limit = (seg.end + 1) if seg.end >= 0 else self._max_download_bytes()
|
||||
async for chunk in resp.content.iter_chunked(args.download_chunk_size):
|
||||
self._check_control()
|
||||
overflow = False
|
||||
if limit is not None:
|
||||
remaining = limit - offset
|
||||
if remaining <= 0:
|
||||
raise FatalError(
|
||||
f"download exceeded the maximum size {limit} bytes"
|
||||
)
|
||||
if len(chunk) > remaining:
|
||||
chunk = chunk[:remaining]
|
||||
overflow = True
|
||||
await self._writer.write_at(offset, chunk)
|
||||
offset += len(chunk)
|
||||
seg.bytes_done = offset
|
||||
self.state.bytes_done = offset
|
||||
await self._persist_progress()
|
||||
if overflow:
|
||||
raise FatalError(
|
||||
f"download exceeded the maximum size {limit} bytes"
|
||||
)
|
||||
|
||||
def _max_download_bytes(self) -> Optional[int]:
|
||||
"""Configured maximum download size in bytes, or ``None`` if disabled."""
|
||||
cap = getattr(args, "download_max_bytes", 0)
|
||||
return cap if cap and cap > 0 else None
|
||||
|
||||
def _raise_for_status(self, status: int) -> None:
|
||||
if status in (401, 403):
|
||||
raise FatalError(
|
||||
f"{redact_url(self.spec.url)} returned {status}; authenticate this "
|
||||
f"host via /api/download/auth or set its API key env var."
|
||||
)
|
||||
if status in _RETRYABLE_STATUSES:
|
||||
raise RetryableError(f"HTTP {status}")
|
||||
raise FatalError(f"unexpected HTTP {status}")
|
||||
|
||||
# ----- finalize / verify (PRD section 8.4) -----
|
||||
|
||||
async def _finalize(self) -> None:
|
||||
self._check_control()
|
||||
await self._close_writer()
|
||||
await self._set_status(DownloadStatus.VERIFYING)
|
||||
|
||||
total = self.state.total_bytes
|
||||
segmented = len(self.state.segments) > 1
|
||||
if segmented:
|
||||
# The .part was preallocated to total_bytes, so its on-disk size is
|
||||
# not evidence of completeness: a segment that ends short (truncated
|
||||
# 206 / server closes mid-range) leaves a zero-filled hole while the
|
||||
# file size still equals total. Verify each segment wrote its full
|
||||
# planned range, and trust the byte counter (== sum of segments)
|
||||
# rather than os.path.getsize for the total check.
|
||||
for seg in self.state.segments:
|
||||
if seg.bytes_done != seg.length:
|
||||
raise FatalError(
|
||||
f"segment {seg.idx} incomplete: wrote {seg.bytes_done} "
|
||||
f"of {seg.length} bytes"
|
||||
)
|
||||
observed = self.state.bytes_done
|
||||
else:
|
||||
# Single-stream writes a contiguous prefix, so the on-disk size is
|
||||
# an independent witness of how much actually landed.
|
||||
observed = os.path.getsize(self.spec.temp_path)
|
||||
if total is not None and observed != total:
|
||||
raise FatalError(
|
||||
f"size mismatch: wrote {observed} of {total} bytes"
|
||||
)
|
||||
|
||||
# Structural gate (cheap, no full read) then optional sha256 (full read).
|
||||
# Both failures are non-retryable (a truncated/corrupt or mismatched file
|
||||
# will not heal on retry), so surface them as FatalError rather than
|
||||
# letting the plain Exceptions fall through to the retryable handler.
|
||||
# ``temp_path`` carries the ``.part`` suffix; pass ``dest_path`` so the
|
||||
# structural check detects the real file format instead of skipping it.
|
||||
try:
|
||||
await asyncio.to_thread(
|
||||
structural.validate, self.spec.temp_path, self.spec.dest_path
|
||||
)
|
||||
if self.spec.expected_sha256:
|
||||
await asyncio.to_thread(
|
||||
checksum.verify_sha256,
|
||||
self.spec.temp_path,
|
||||
self.spec.expected_sha256,
|
||||
)
|
||||
except (structural.StructuralError, checksum.ChecksumError) as e:
|
||||
raise FatalError(str(e)) from e
|
||||
|
||||
os.makedirs(os.path.dirname(self.spec.dest_path), exist_ok=True)
|
||||
os.replace(self.spec.temp_path, self.spec.dest_path)
|
||||
logging.info(
|
||||
"[model_downloader] completed %s (%d bytes)",
|
||||
self.spec.model_id, observed,
|
||||
)
|
||||
# Catalog into the assets system (blake3 dedup identity). Best-effort.
|
||||
await dedup.register_completed(self.spec.dest_path)
|
||||
|
||||
# ----- helpers -----
|
||||
|
||||
def _recompute_bytes_done(self) -> None:
|
||||
self.state.bytes_done = sum(s.bytes_done for s in self.state.segments)
|
||||
now = time.monotonic()
|
||||
dt = now - self.state._last_time
|
||||
if dt >= 0.5:
|
||||
self.state.speed_bps = (self.state.bytes_done - self.state._last_bytes) / dt
|
||||
self.state._last_bytes = self.state.bytes_done
|
||||
self.state._last_time = now
|
||||
|
||||
async def _persist_progress(self, force: bool = False) -> None:
|
||||
# Both the DB write and the websocket notify are gated by the same
|
||||
# throttle: persisting hits SQLite, and notifying broadcasts to every
|
||||
# client, so doing either per-chunk (small --download-chunk-size or
|
||||
# many concurrent segments) would overwhelm both. Skip entirely inside
|
||||
# the window; the next persist (or a forced one) ships the latest bytes.
|
||||
now = time.monotonic()
|
||||
if not force and now - self._last_persist < _PERSIST_INTERVAL:
|
||||
return
|
||||
self._last_persist = now
|
||||
# SQLite is blocking; run it off the event loop per the queries module
|
||||
# contract so progress persists don't stall the web server.
|
||||
await asyncio.to_thread(self._write_progress)
|
||||
if self._notify:
|
||||
self._notify(self.spec.download_id)
|
||||
|
||||
def _write_progress(self) -> None:
|
||||
queries.update_download(self.spec.download_id, bytes_done=self.state.bytes_done)
|
||||
for seg in self.state.segments:
|
||||
if seg.end >= seg.start: # skip unknown-size sentinel
|
||||
queries.update_segment_progress(
|
||||
self.spec.download_id, seg.idx, seg.bytes_done
|
||||
)
|
||||
|
||||
async def _reset_for_restart(self) -> None:
|
||||
await self._close_writer()
|
||||
self._remove_temp()
|
||||
for seg in self.state.segments:
|
||||
seg.bytes_done = 0
|
||||
self.state.bytes_done = 0
|
||||
await asyncio.to_thread(
|
||||
queries.update_download, self.spec.download_id, bytes_done=0
|
||||
)
|
||||
if await asyncio.to_thread(queries.list_segments, self.spec.download_id):
|
||||
await asyncio.to_thread(
|
||||
queries.replace_segments, self.spec.download_id, []
|
||||
)
|
||||
|
||||
async def _close_writer(self) -> None:
|
||||
if self._writer is not None:
|
||||
try:
|
||||
await self._writer.close()
|
||||
except Exception:
|
||||
logging.debug("[model_downloader] writer close error", exc_info=True)
|
||||
self._writer = None
|
||||
|
||||
def _segmented_part_valid(self, total_bytes: int) -> bool:
|
||||
"""True when the temp file is the preallocated segmented ``.part``.
|
||||
|
||||
A segmented transfer preallocates the .part to ``total_bytes`` up front
|
||||
and tracks how much of each range landed via per-segment offsets. Those
|
||||
offsets are only trustworthy when the file they describe is still on
|
||||
disk at its full preallocated size. A missing file (swept after a
|
||||
failure, removed on a fatal error, deleted by hand) or a wrong-sized one
|
||||
means the persisted offsets no longer correspond to real bytes and must
|
||||
not be resumed over. Doing so would skip "complete" segments and leave
|
||||
zero-filled holes that pass the size-only verification gate.
|
||||
"""
|
||||
try:
|
||||
return os.path.getsize(self.spec.temp_path) == total_bytes
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
def _contiguous_prefix_valid(self, prefix_len: int) -> bool:
|
||||
"""True when the temp file is exactly ``prefix_len`` contiguous bytes.
|
||||
|
||||
Single-stream resume appends sequentially, so a valid resume point
|
||||
implies the .part size equals the persisted offset. A larger file (e.g.
|
||||
one preallocated to ``total_bytes`` by a previous segmented run) or a
|
||||
missing/short file means the persisted offset is not a trustworthy
|
||||
contiguous prefix and must not be resumed over.
|
||||
"""
|
||||
try:
|
||||
return os.path.getsize(self.spec.temp_path) == prefix_len
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
def _remove_temp(self) -> None:
|
||||
try:
|
||||
os.remove(self.spec.temp_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except OSError as e:
|
||||
logging.warning(
|
||||
"[model_downloader] could not remove %s: %s", self.spec.temp_path, e
|
||||
)
|
||||
|
||||
async def _set_status(self, status: str, error: Optional[str] = None) -> None:
|
||||
# ``error`` is authoritative: passing None clears any prior failure
|
||||
# text so transitions out of a failure state (retry/success) don't
|
||||
# leave stale messages on RuntimeState or in the persisted row.
|
||||
self.state.status = status
|
||||
self.state.error = error
|
||||
fields = {"status": status, "bytes_done": self.state.bytes_done, "error": error}
|
||||
if status == DownloadStatus.QUEUED:
|
||||
fields["attempts"] = self.spec.attempts + 1
|
||||
self.spec.attempts += 1
|
||||
await asyncio.to_thread(queries.update_download, self.spec.download_id, **fields)
|
||||
if self._notify:
|
||||
self._notify(self.spec.download_id)
|
||||
@@ -0,0 +1,51 @@
|
||||
"""Segment planning.
|
||||
|
||||
Split a known byte range into S roughly-equal segments, each fetched by its
|
||||
own coroutine with ``Range: bytes=start-end``. Falls back to a single segment
|
||||
when the server doesn't support ranges or the size is unknown/too small for
|
||||
segmentation to be worthwhile.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
# Below this size, the per-connection setup cost outweighs any parallelism.
|
||||
_MIN_SEGMENT_BYTES = 1 * 1024 * 1024
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SegmentPlan:
|
||||
idx: int
|
||||
start: int
|
||||
end: int # inclusive
|
||||
|
||||
@property
|
||||
def length(self) -> int:
|
||||
return self.end - self.start + 1
|
||||
|
||||
|
||||
def effective_segment_count(
|
||||
total_bytes: int | None, accept_ranges: bool, configured: int
|
||||
) -> int:
|
||||
"""How many segments to actually use for this file."""
|
||||
if not accept_ranges or total_bytes is None or total_bytes <= 0:
|
||||
return 1
|
||||
by_size = max(1, total_bytes // _MIN_SEGMENT_BYTES)
|
||||
return max(1, min(configured, by_size))
|
||||
|
||||
|
||||
def plan_segments(total_bytes: int, num_segments: int) -> list[SegmentPlan]:
|
||||
"""Return ``num_segments`` contiguous, inclusive byte ranges covering [0, total)."""
|
||||
if total_bytes <= 0 or num_segments <= 1:
|
||||
return [SegmentPlan(idx=0, start=0, end=max(0, total_bytes - 1))]
|
||||
base = total_bytes // num_segments
|
||||
plans: list[SegmentPlan] = []
|
||||
start = 0
|
||||
for i in range(num_segments):
|
||||
# Last segment soaks up the remainder.
|
||||
length = base if i < num_segments - 1 else total_bytes - start
|
||||
end = start + length - 1
|
||||
plans.append(SegmentPlan(idx=i, start=start, end=end))
|
||||
start = end + 1
|
||||
return plans
|
||||
@@ -0,0 +1,110 @@
|
||||
"""Positioned, off-loop file writes.
|
||||
|
||||
Network I/O stays on the event loop; every blocking disk op (preallocate,
|
||||
positioned write, fsync) is run in a bounded thread pool via
|
||||
``run_in_executor`` so downloads never stall inference or the web server.
|
||||
|
||||
A single file descriptor is opened for the whole download. Segments write to
|
||||
their own offsets with ``os.pwrite`` — which is offset-addressed and atomic
|
||||
per call, so concurrent segment writers need no extra locking. Per-chunk
|
||||
fsync is avoided; we fsync once at completion.
|
||||
|
||||
``os.pwrite`` is unavailable on Windows, so there we fall back to
|
||||
``os.lseek`` + ``os.write`` guarded by a per-writer lock (the seek/write pair
|
||||
is not atomic, so concurrent segment writers must be serialized).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Optional
|
||||
|
||||
# One shared, bounded pool for all download disk I/O.
|
||||
_EXECUTOR = ThreadPoolExecutor(max_workers=8, thread_name_prefix="dl-writer")
|
||||
|
||||
_HAS_PWRITE = hasattr(os, "pwrite")
|
||||
|
||||
# On Windows ``os.open`` defaults to text mode, which translates every ``\n``
|
||||
# byte into ``\r\n`` on write and corrupts binary payloads (the file grows by
|
||||
# one byte per 0x0A). ``O_BINARY`` disables that translation; it does not exist
|
||||
# on POSIX, where the default is already binary.
|
||||
_O_BINARY = getattr(os, "O_BINARY", 0)
|
||||
|
||||
|
||||
class FileWriter:
|
||||
"""Owns the ``.part`` file descriptor for one download."""
|
||||
|
||||
def __init__(self, path: str) -> None:
|
||||
self.path = path
|
||||
self._fd: Optional[int] = None
|
||||
# Serializes lseek+write on platforms without os.pwrite (Windows).
|
||||
self._seek_lock = threading.Lock()
|
||||
|
||||
def _open(self) -> None:
|
||||
os.makedirs(os.path.dirname(self.path), exist_ok=True)
|
||||
self._fd = os.open(self.path, os.O_RDWR | os.O_CREAT | _O_BINARY, 0o644)
|
||||
|
||||
async def open(self) -> None:
|
||||
await asyncio.get_running_loop().run_in_executor(_EXECUTOR, self._open)
|
||||
|
||||
async def preallocate(self, size: int) -> None:
|
||||
"""Grow the file to ``size`` so segments write to their offsets."""
|
||||
if self._fd is None or size <= 0:
|
||||
return
|
||||
await asyncio.get_running_loop().run_in_executor(
|
||||
_EXECUTOR, os.ftruncate, self._fd, size
|
||||
)
|
||||
|
||||
async def truncate(self, size: int = 0) -> None:
|
||||
"""Truncate the file to ``size`` bytes (default: empty it)."""
|
||||
if self._fd is None:
|
||||
return
|
||||
await asyncio.get_running_loop().run_in_executor(
|
||||
_EXECUTOR, os.ftruncate, self._fd, size
|
||||
)
|
||||
|
||||
def _pwrite_all(self, data: bytes, offset: int) -> None:
|
||||
"""A positioned write may write fewer bytes than requested (signal
|
||||
interruption, near-ENOSPC); loop until every byte lands so we never
|
||||
leave a gap while the caller advances by the full chunk length.
|
||||
|
||||
Uses ``os.pwrite`` where available (offset-addressed, atomic per call).
|
||||
On Windows it falls back to ``os.lseek`` + ``os.write`` under a lock,
|
||||
since that pair is not atomic across concurrent segment writers."""
|
||||
assert self._fd is not None, "writer not opened"
|
||||
view = memoryview(data)
|
||||
written = 0
|
||||
total = len(view)
|
||||
while written < total:
|
||||
if _HAS_PWRITE:
|
||||
n = os.pwrite(self._fd, view[written:], offset + written)
|
||||
else:
|
||||
with self._seek_lock:
|
||||
os.lseek(self._fd, offset + written, os.SEEK_SET)
|
||||
n = os.write(self._fd, view[written:])
|
||||
if n == 0:
|
||||
raise OSError(
|
||||
f"positioned write wrote 0 bytes at offset {offset + written} "
|
||||
f"({written}/{total} bytes written)"
|
||||
)
|
||||
written += n
|
||||
|
||||
async def write_at(self, offset: int, data: bytes) -> None:
|
||||
assert self._fd is not None, "writer not opened"
|
||||
await asyncio.get_running_loop().run_in_executor(
|
||||
_EXECUTOR, self._pwrite_all, data, offset
|
||||
)
|
||||
|
||||
async def flush(self) -> None:
|
||||
if self._fd is None:
|
||||
return
|
||||
await asyncio.get_running_loop().run_in_executor(_EXECUTOR, os.fsync, self._fd)
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._fd is None:
|
||||
return
|
||||
fd, self._fd = self._fd, None
|
||||
await asyncio.get_running_loop().run_in_executor(_EXECUTOR, os.close, fd)
|
||||
@@ -0,0 +1,444 @@
|
||||
"""Public facade for the download manager.
|
||||
|
||||
This is the only object the server imports. It validates requests, owns the
|
||||
:class:`Scheduler`, and exposes a small async API plus read models for status.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import uuid
|
||||
from typing import Callable, Optional
|
||||
|
||||
from app.model_downloader.constants import DownloadStatus
|
||||
from app.model_downloader.database import queries
|
||||
from app.model_downloader.net.probe import gated_error_message, probe
|
||||
from app.model_downloader.scheduler import SCHEDULER
|
||||
from app.model_downloader.security import paths
|
||||
from app.model_downloader.net.http import redact_url
|
||||
from app.model_downloader.security.allowlist import (
|
||||
ALLOWED_MODEL_EXTENSIONS,
|
||||
filename_extension,
|
||||
is_host_allowed_url,
|
||||
is_url_downloadable,
|
||||
url_path_extension,
|
||||
)
|
||||
from app.model_downloader.security.paths import InvalidModelId
|
||||
|
||||
# Non-terminal statuses: an existing row in one of these blocks a re-enqueue.
|
||||
_LIVE_STATUSES = (
|
||||
DownloadStatus.QUEUED,
|
||||
DownloadStatus.ACTIVE,
|
||||
DownloadStatus.PAUSED,
|
||||
DownloadStatus.VERIFYING,
|
||||
)
|
||||
|
||||
|
||||
class DownloadError(Exception):
|
||||
"""A user-facing error with a stable machine-readable code."""
|
||||
|
||||
def __init__(self, code: str, message: str, status: int = 400) -> None:
|
||||
super().__init__(message)
|
||||
self.code = code
|
||||
self.message = message
|
||||
self.http_status = status
|
||||
|
||||
|
||||
class DownloadManager:
|
||||
def __init__(self) -> None:
|
||||
self._scheduler = SCHEDULER
|
||||
self._notify_cb: Optional[Callable[[str], None]] = None
|
||||
# Serializes the "check for a live download, then write" critical section
|
||||
# per model_id. ``downloads`` has no uniqueness constraint on model_id
|
||||
# (history rows are kept), so without this two concurrent enqueue/resume
|
||||
# calls could both pass the live check and admit two jobs sharing one
|
||||
# temp/dest path. The manager is a process singleton over a local SQLite
|
||||
# DB, so an in-process lock is sufficient (and avoids a migration).
|
||||
self._model_locks: dict[str, asyncio.Lock] = {}
|
||||
|
||||
def set_notify(self, cb: Optional[Callable[[str], None]]) -> None:
|
||||
self._notify_cb = cb
|
||||
self._scheduler.set_notify(cb)
|
||||
|
||||
async def start(self) -> None:
|
||||
await self._scheduler.start()
|
||||
|
||||
# ----- enqueue -----
|
||||
|
||||
async def enqueue(
|
||||
self,
|
||||
url: str,
|
||||
model_id: str,
|
||||
*,
|
||||
priority: int = 0,
|
||||
expected_sha256: Optional[str] = None,
|
||||
allow_any_extension: bool = False,
|
||||
) -> str:
|
||||
# Coarse gate first: host/scheme must be allowlisted, and any extension
|
||||
# present in the URL path must be a known model type. A URL whose path
|
||||
# carries NO extension (e.g. Civitai's ``/api/download/models/<id>``) is
|
||||
# admitted here and its real extension is resolved from the network
|
||||
# below before the download is finally accepted.
|
||||
if allow_any_extension:
|
||||
if not is_host_allowed_url(url):
|
||||
raise DownloadError(
|
||||
"URL_NOT_ALLOWED",
|
||||
"URL is not on the download allowlist (host/scheme).",
|
||||
)
|
||||
elif not is_url_downloadable(url):
|
||||
raise DownloadError(
|
||||
"URL_NOT_ALLOWED",
|
||||
"URL is not on the download allowlist (host/scheme/extension).",
|
||||
)
|
||||
|
||||
# When the URL path has no extension, follow it to where it resolves and
|
||||
# adopt the real extension from the response, forcing the stored
|
||||
# filename to match. Skipped when the caller opted into any extension.
|
||||
if not allow_any_extension and url_path_extension(url) == "":
|
||||
resolved_ext = await self._resolve_extension(url)
|
||||
model_id = paths.apply_extension(model_id, resolved_ext)
|
||||
|
||||
try:
|
||||
paths.parse_model_id(model_id, allow_any_extension)
|
||||
dest_path, temp_path = paths.resolve_destination(model_id, allow_any_extension)
|
||||
except InvalidModelId as e:
|
||||
raise DownloadError("INVALID_MODEL_ID", str(e))
|
||||
|
||||
if await asyncio.to_thread(
|
||||
paths.resolve_existing, model_id, allow_any_extension
|
||||
):
|
||||
raise DownloadError(
|
||||
"ALREADY_AVAILABLE",
|
||||
f"Model already exists on disk: {model_id}",
|
||||
status=409,
|
||||
)
|
||||
download_id = str(uuid.uuid4())
|
||||
# Hold the per-model lock across the live check and the insert so a
|
||||
# concurrent enqueue/resume for the same model_id cannot interleave
|
||||
# between them and create a second job against the same temp/dest path.
|
||||
async with self._model_lock(model_id):
|
||||
if await self._has_live_download(model_id):
|
||||
raise DownloadError(
|
||||
"ALREADY_DOWNLOADING",
|
||||
f"A download for {model_id} is already in progress.",
|
||||
status=409,
|
||||
)
|
||||
await asyncio.to_thread(
|
||||
queries.insert_download,
|
||||
{
|
||||
"id": download_id,
|
||||
"url": url,
|
||||
"model_id": model_id,
|
||||
"dest_path": dest_path,
|
||||
"temp_path": temp_path,
|
||||
"status": DownloadStatus.QUEUED,
|
||||
"priority": priority,
|
||||
"expected_sha256": expected_sha256,
|
||||
"allow_any_extension": allow_any_extension,
|
||||
},
|
||||
)
|
||||
logging.info("[model_downloader] enqueued %s -> %s", redact_url(url), model_id)
|
||||
await self._scheduler.pump()
|
||||
return download_id
|
||||
|
||||
async def _resolve_extension(self, url: str) -> str:
|
||||
"""Follow ``url`` to its final response and return the real extension.
|
||||
|
||||
Used for allowlisted URLs whose path has no extension (e.g. Civitai
|
||||
download endpoints): the filename lives in the ``Content-Disposition``
|
||||
header or the post-redirect URL. Raises :class:`DownloadError` when the
|
||||
URL can't be resolved, needs authentication, or resolves to something
|
||||
that is not a known model file — so we never persist a bogus destination.
|
||||
"""
|
||||
pr = await probe(url)
|
||||
if not pr.ok:
|
||||
if pr.gated:
|
||||
raise DownloadError(
|
||||
"GATED_REPO" if pr.is_gated_repo else "CREDENTIALS_REQUIRED",
|
||||
gated_error_message(url, pr),
|
||||
status=401,
|
||||
)
|
||||
raise DownloadError(
|
||||
"URL_RESOLVE_FAILED",
|
||||
f"Could not resolve {redact_url(url)}: {pr.error or 'unknown error'}",
|
||||
status=502,
|
||||
)
|
||||
ext = filename_extension(pr.filename) if pr.filename else ""
|
||||
if ext not in ALLOWED_MODEL_EXTENSIONS:
|
||||
raise DownloadError(
|
||||
"URL_NOT_ALLOWED",
|
||||
f"URL resolves to {pr.filename or '<unknown>'!r}, which is not a "
|
||||
f"known model file type {ALLOWED_MODEL_EXTENSIONS}.",
|
||||
)
|
||||
return ext
|
||||
|
||||
def _model_lock(self, model_id: str) -> asyncio.Lock:
|
||||
# Lazily create one lock per model_id. There is no ``await`` between the
|
||||
# lookup and the insert, so under the single asyncio thread this is
|
||||
# atomic and cannot hand out two different locks for the same model_id.
|
||||
lock = self._model_locks.get(model_id)
|
||||
if lock is None:
|
||||
lock = asyncio.Lock()
|
||||
self._model_locks[model_id] = lock
|
||||
return lock
|
||||
|
||||
async def _has_live_download(
|
||||
self, model_id: str, *, exclude_id: Optional[str] = None
|
||||
) -> bool:
|
||||
return await asyncio.to_thread(
|
||||
queries.has_live_download_for_model, model_id, _LIVE_STATUSES, exclude_id
|
||||
)
|
||||
|
||||
# ----- control -----
|
||||
|
||||
async def pause(self, download_id: str) -> None:
|
||||
job = self._scheduler.get_job(download_id)
|
||||
if job is not None:
|
||||
job.request_pause()
|
||||
return
|
||||
row = await asyncio.to_thread(queries.get_download, download_id)
|
||||
if row is None:
|
||||
raise DownloadError("NOT_FOUND", "No such download.", status=404)
|
||||
if row.status == DownloadStatus.QUEUED:
|
||||
await asyncio.to_thread(
|
||||
queries.update_download, download_id, status=DownloadStatus.PAUSED
|
||||
)
|
||||
|
||||
async def resume(self, download_id: str) -> None:
|
||||
row = await asyncio.to_thread(queries.get_download, download_id)
|
||||
if row is None:
|
||||
raise DownloadError("NOT_FOUND", "No such download.", status=404)
|
||||
if row.status not in (DownloadStatus.PAUSED, DownloadStatus.FAILED):
|
||||
return
|
||||
# Re-queueing a paused/failed row must respect the single-live-per-model
|
||||
# invariant: another download (e.g. a newer enqueue) may already be live
|
||||
# for this model_id and would share this row's temp/dest path. Hold the
|
||||
# per-model lock across the check and the status flip, and exclude this
|
||||
# row itself (a paused row is already a "live" status).
|
||||
async with self._model_lock(row.model_id):
|
||||
if await self._has_live_download(row.model_id, exclude_id=download_id):
|
||||
raise DownloadError(
|
||||
"ALREADY_DOWNLOADING",
|
||||
f"A download for {row.model_id} is already in progress.",
|
||||
status=409,
|
||||
)
|
||||
await asyncio.to_thread(
|
||||
queries.update_download,
|
||||
download_id,
|
||||
status=DownloadStatus.QUEUED,
|
||||
error=None,
|
||||
)
|
||||
await self._scheduler.pump()
|
||||
|
||||
async def cancel(self, download_id: str) -> None:
|
||||
job = self._scheduler.get_job(download_id)
|
||||
if job is not None:
|
||||
job.request_cancel()
|
||||
return
|
||||
row = await asyncio.to_thread(queries.get_download, download_id)
|
||||
if row is None:
|
||||
raise DownloadError("NOT_FOUND", "No such download.", status=404)
|
||||
if row.status in _LIVE_STATUSES:
|
||||
try:
|
||||
os.remove(row.temp_path)
|
||||
except OSError:
|
||||
pass
|
||||
await asyncio.to_thread(
|
||||
queries.update_download, download_id, status=DownloadStatus.CANCELLED
|
||||
)
|
||||
|
||||
async def set_priority(self, download_id: str, priority: int) -> None:
|
||||
row = await asyncio.to_thread(queries.get_download, download_id)
|
||||
if row is None:
|
||||
raise DownloadError("NOT_FOUND", "No such download.", status=404)
|
||||
await asyncio.to_thread(
|
||||
queries.update_download, download_id, priority=priority
|
||||
)
|
||||
# Admission-order only; a higher priority is
|
||||
# picked up the next time a slot frees. Pump in case a slot is free now.
|
||||
await self._scheduler.pump()
|
||||
|
||||
async def delete(self, download_id: str) -> None:
|
||||
"""Delete a terminal download so it stays gone from history.
|
||||
|
||||
Refuses to delete a live download so a record is never removed out from
|
||||
under a running worker; cancel it first. Any leftover ``.part`` temp
|
||||
file (e.g. from a failed transfer) is removed, but the finished model
|
||||
file on disk is never touched.
|
||||
"""
|
||||
if self._scheduler.get_job(download_id) is not None:
|
||||
raise DownloadError(
|
||||
"DOWNLOAD_ACTIVE",
|
||||
"Cannot delete a download that is still in progress.",
|
||||
status=409,
|
||||
)
|
||||
row = await asyncio.to_thread(queries.get_download, download_id)
|
||||
if row is None:
|
||||
raise DownloadError("NOT_FOUND", "No such download.", status=404)
|
||||
if row.status in _LIVE_STATUSES:
|
||||
raise DownloadError(
|
||||
"DOWNLOAD_ACTIVE",
|
||||
"Cannot delete a download that is still in progress.",
|
||||
status=409,
|
||||
)
|
||||
|
||||
try:
|
||||
os.remove(row.temp_path)
|
||||
except OSError:
|
||||
pass
|
||||
await asyncio.to_thread(queries.delete_download, download_id)
|
||||
|
||||
async def clear(self) -> int:
|
||||
"""Delete all terminal downloads from history in one transaction.
|
||||
|
||||
Skips anything still live (queued/active/paused/verifying, or a running
|
||||
job) so an in-flight download is never removed out from under a worker.
|
||||
Finished model files on disk are never touched; only leftover ``.part``
|
||||
temp files from failed/cancelled transfers are removed. Returns the
|
||||
number of history rows deleted.
|
||||
"""
|
||||
|
||||
rows = await asyncio.to_thread(queries.list_downloads)
|
||||
deletable = [
|
||||
r
|
||||
for r in rows
|
||||
if r.status not in _LIVE_STATUSES
|
||||
and self._scheduler.get_job(r.id) is None
|
||||
]
|
||||
if not deletable:
|
||||
return 0
|
||||
for r in deletable:
|
||||
try:
|
||||
os.remove(r.temp_path)
|
||||
except OSError:
|
||||
pass
|
||||
return await asyncio.to_thread(
|
||||
queries.delete_downloads, [r.id for r in deletable]
|
||||
)
|
||||
|
||||
# ----- read models -----
|
||||
|
||||
def _view(self, row) -> dict:
|
||||
"""Combine the persisted row with live in-memory progress, if running."""
|
||||
job = self._scheduler.get_job(row.id)
|
||||
bytes_done = row.bytes_done
|
||||
total = row.total_bytes
|
||||
speed = None
|
||||
eta = None
|
||||
segments = None
|
||||
if job is not None:
|
||||
st = job.state
|
||||
bytes_done = st.bytes_done
|
||||
total = st.total_bytes if st.total_bytes is not None else total
|
||||
speed = st.speed_bps
|
||||
eta = st.eta_seconds
|
||||
segments = [
|
||||
{"idx": s.idx, "bytes_done": s.bytes_done, "length": s.length}
|
||||
for s in st.segments
|
||||
if s.end >= s.start
|
||||
]
|
||||
progress = (bytes_done / total) if total else None
|
||||
return {
|
||||
"download_id": row.id,
|
||||
"model_id": row.model_id,
|
||||
"url": redact_url(row.url),
|
||||
"status": row.status,
|
||||
"priority": row.priority,
|
||||
"total_bytes": total,
|
||||
"bytes_done": bytes_done,
|
||||
"progress": progress,
|
||||
"speed_bps": speed,
|
||||
"eta_seconds": eta,
|
||||
"segments": segments,
|
||||
"error": row.error,
|
||||
"created_at": row.created_at,
|
||||
"updated_at": row.updated_at,
|
||||
}
|
||||
|
||||
def _view_from_state(self, job) -> dict:
|
||||
"""Build a view purely from the live in-memory job state (no DB)."""
|
||||
st = job.state
|
||||
return {
|
||||
"download_id": st.download_id,
|
||||
"model_id": st.model_id,
|
||||
"url": redact_url(st.url),
|
||||
"status": st.status,
|
||||
"priority": st.priority,
|
||||
"total_bytes": st.total_bytes,
|
||||
"bytes_done": st.bytes_done,
|
||||
"progress": st.progress,
|
||||
"speed_bps": st.speed_bps,
|
||||
"eta_seconds": st.eta_seconds,
|
||||
"segments": [
|
||||
{"idx": s.idx, "bytes_done": s.bytes_done, "length": s.length}
|
||||
for s in st.segments
|
||||
if s.end >= s.start
|
||||
],
|
||||
"error": st.error,
|
||||
}
|
||||
|
||||
def status_sync(self, download_id: str) -> Optional[dict]:
|
||||
"""Synchronous status read for the websocket notify path.
|
||||
|
||||
Uses live in-memory state when the job is running (no DB round-trip on
|
||||
the hot path); falls back to a quick DB read otherwise.
|
||||
"""
|
||||
job = self._scheduler.get_job(download_id)
|
||||
if job is not None:
|
||||
return self._view_from_state(job)
|
||||
row = queries.get_download(download_id)
|
||||
return self._view(row) if row is not None else None
|
||||
|
||||
async def status(self, download_id: str) -> Optional[dict]:
|
||||
row = await asyncio.to_thread(queries.get_download, download_id)
|
||||
return self._view(row) if row is not None else None
|
||||
|
||||
async def list(self) -> list[dict]:
|
||||
rows = await asyncio.to_thread(queries.list_downloads)
|
||||
return [self._view(r) for r in rows]
|
||||
|
||||
async def availability(self, models: dict[str, str]) -> dict[str, dict]:
|
||||
"""Bulk per-id ``{state, progress, ...}`` for the frontend poll.
|
||||
|
||||
``state`` is ``available`` (on disk), ``downloading`` (live row), or
|
||||
``missing``. Cheap: a path lookup plus an in-memory/DB status check.
|
||||
"""
|
||||
rows = await asyncio.to_thread(queries.list_downloads)
|
||||
by_model: dict[str, object] = {}
|
||||
for r in rows:
|
||||
if r.status in _LIVE_STATUSES or r.model_id not in by_model:
|
||||
by_model[r.model_id] = r
|
||||
|
||||
# ``url_allowed`` mirrors the coarse enqueue gate (host/scheme + a
|
||||
# non-disallowed extension); URLs whose extension is only known after a
|
||||
# network resolve — e.g. Civitai download endpoints — report allowed.
|
||||
out: dict[str, dict] = {}
|
||||
for model_id, url in models.items():
|
||||
try:
|
||||
exists = await asyncio.to_thread(paths.resolve_existing, model_id)
|
||||
except InvalidModelId:
|
||||
out[model_id] = {"state": "missing", "url_allowed": is_url_downloadable(url)}
|
||||
continue
|
||||
if exists:
|
||||
out[model_id] = {"state": "available", "url_allowed": is_url_downloadable(url)}
|
||||
continue
|
||||
row = by_model.get(model_id)
|
||||
if row is not None and row.status in _LIVE_STATUSES:
|
||||
view = self._view(row)
|
||||
out[model_id] = {
|
||||
"state": "downloading",
|
||||
"url_allowed": is_url_downloadable(url),
|
||||
"download_id": view["download_id"],
|
||||
"progress": view["progress"],
|
||||
"bytes_done": view["bytes_done"],
|
||||
"total_bytes": view["total_bytes"],
|
||||
"speed_bps": view["speed_bps"],
|
||||
}
|
||||
else:
|
||||
out[model_id] = {"state": "missing", "url_allowed": is_url_downloadable(url)}
|
||||
return out
|
||||
|
||||
|
||||
DOWNLOAD_MANAGER = DownloadManager()
|
||||
@@ -0,0 +1,142 @@
|
||||
"""Manual, validated redirect-following request opener.
|
||||
|
||||
Automatic redirects are disabled. We follow hops ourselves
|
||||
so that on *every* hop we (a) re-validate scheme + reject credentials-in-URL,
|
||||
(b) recompute which auth — if any — applies to that hop's host, and (c) let the
|
||||
connector's resolver screen the IP. This is the single place that attaches a
|
||||
token, so it can never ride a redirect to a CDN host.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import AsyncIterator, Optional
|
||||
from urllib.parse import unquote, urljoin, urlsplit, urlunsplit
|
||||
|
||||
import aiohttp
|
||||
|
||||
from app.model_downloader.auth.resolver import resolve_auth_for_hop
|
||||
from app.model_downloader.net.session import get_session
|
||||
from app.model_downloader.security.ssrf import (
|
||||
MAX_REDIRECTS,
|
||||
SSRFError,
|
||||
check_redirect_hop,
|
||||
)
|
||||
|
||||
_REDIRECT_CODES = {301, 302, 303, 307, 308}
|
||||
DEFAULT_TIMEOUT = aiohttp.ClientTimeout(total=None, sock_connect=30, sock_read=120)
|
||||
|
||||
|
||||
def redact_url(url: str) -> str:
|
||||
"""Drop the query string so a query-scheme secret is never logged/stored."""
|
||||
try:
|
||||
parts = urlsplit(url)
|
||||
except ValueError:
|
||||
return "<unparseable-url>"
|
||||
return urlunsplit(parts._replace(query=""))
|
||||
|
||||
|
||||
_CD_FILENAME_STAR = re.compile(
|
||||
r"filename\*\s*=\s*[^']*'[^']*'([^;]+)", re.IGNORECASE
|
||||
)
|
||||
_CD_FILENAME_QUOTED = re.compile(r'filename\s*=\s*"([^"]+)"', re.IGNORECASE)
|
||||
_CD_FILENAME_BARE = re.compile(r"filename\s*=\s*([^;]+)", re.IGNORECASE)
|
||||
|
||||
|
||||
def filename_from_content_disposition(value: Optional[str]) -> Optional[str]:
|
||||
"""Extract the download filename from a ``Content-Disposition`` header.
|
||||
|
||||
Prefers the RFC 5987 ``filename*=`` form (percent-decoded) over the plain
|
||||
``filename=`` form. Any directory components in the value are stripped so a
|
||||
hostile header can only influence the *name*, never the target directory.
|
||||
Returns ``None`` when no filename is present.
|
||||
"""
|
||||
if not value:
|
||||
return None
|
||||
for pat, decode in (
|
||||
(_CD_FILENAME_STAR, True),
|
||||
(_CD_FILENAME_QUOTED, False),
|
||||
(_CD_FILENAME_BARE, False),
|
||||
):
|
||||
m = pat.search(value)
|
||||
if not m:
|
||||
continue
|
||||
raw = m.group(1).strip().strip('"')
|
||||
if decode:
|
||||
try:
|
||||
raw = unquote(raw)
|
||||
except Exception:
|
||||
pass
|
||||
name = raw.replace("\\", "/").rsplit("/", 1)[-1].strip()
|
||||
if name:
|
||||
return name
|
||||
return None
|
||||
|
||||
|
||||
async def _resolve_final_response(
|
||||
method: str,
|
||||
url: str,
|
||||
base_headers: dict[str, str],
|
||||
timeout: aiohttp.ClientTimeout,
|
||||
) -> tuple[aiohttp.ClientResponse, str]:
|
||||
"""Follow redirects manually until a non-redirect response.
|
||||
|
||||
Each intermediate redirect response is released before the next hop.
|
||||
Returns the final ``(response, final_url)``; the caller owns releasing it.
|
||||
"""
|
||||
session = await get_session()
|
||||
current = url
|
||||
hops = 0
|
||||
while True:
|
||||
check_redirect_hop(current, is_initial_url=(hops == 0))
|
||||
parts = urlsplit(current)
|
||||
auth = await resolve_auth_for_hop(parts.hostname or "", parts.scheme)
|
||||
req_headers = dict(base_headers)
|
||||
if auth is not None:
|
||||
req_headers.update(auth.headers)
|
||||
|
||||
resp = await session.request(
|
||||
method,
|
||||
current,
|
||||
allow_redirects=False,
|
||||
headers=req_headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
if resp.status in _REDIRECT_CODES and resp.headers.get("Location"):
|
||||
next_url = urljoin(str(resp.url), resp.headers["Location"])
|
||||
await resp.release()
|
||||
hops += 1
|
||||
if hops > MAX_REDIRECTS:
|
||||
raise SSRFError(
|
||||
f"too many redirects (> {MAX_REDIRECTS}) for {redact_url(url)}"
|
||||
)
|
||||
current = next_url
|
||||
continue
|
||||
return resp, redact_url(str(resp.url))
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def open_validated(
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
headers: Optional[dict[str, str]] = None,
|
||||
timeout: aiohttp.ClientTimeout = DEFAULT_TIMEOUT,
|
||||
) -> AsyncIterator[tuple[aiohttp.ClientResponse, str]]:
|
||||
"""Open ``method url`` following redirects manually and validated.
|
||||
|
||||
Yields ``(response, final_url)`` where ``final_url`` is redacted of any
|
||||
query string. The response is released automatically on exit.
|
||||
"""
|
||||
resp, final_url = await _resolve_final_response(
|
||||
method, url, dict(headers or {}), timeout
|
||||
)
|
||||
try:
|
||||
yield resp, final_url
|
||||
finally:
|
||||
try:
|
||||
await resp.release()
|
||||
except Exception: # pragma: no cover - best-effort cleanup
|
||||
logging.debug("[model_downloader] response release error", exc_info=True)
|
||||
@@ -0,0 +1,192 @@
|
||||
"""Pre-download probe.
|
||||
|
||||
Issues a tiny ranged GET (``Range: bytes=0-0``) — which doubles as a
|
||||
range-support test — to discover ``Content-Length``, ``Accept-Ranges``,
|
||||
``ETag``/``Last-Modified``, and the final post-redirect URL. For HuggingFace
|
||||
LFS files the true size also appears in the non-standard ``X-Linked-Size``
|
||||
header, which we read as a fallback.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
from urllib.parse import urlparse, urlsplit
|
||||
|
||||
import aiohttp
|
||||
|
||||
from app.model_downloader.net.http import (
|
||||
filename_from_content_disposition,
|
||||
open_validated,
|
||||
redact_url,
|
||||
)
|
||||
from app.model_downloader.net.session import parse_int_header
|
||||
|
||||
_PROBE_TIMEOUT = aiohttp.ClientTimeout(total=60, sock_connect=30, sock_read=30)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProbeResult:
|
||||
ok: bool
|
||||
status: int
|
||||
final_url: Optional[str] = None
|
||||
total_bytes: Optional[int] = None
|
||||
accept_ranges: bool = False
|
||||
etag: Optional[str] = None
|
||||
last_modified: Optional[str] = None
|
||||
gated: bool = False # 401/403 — needs (or has wrong) credentials
|
||||
error: Optional[str] = None
|
||||
# HuggingFace's ``X-Error-Code`` header (e.g. ``GatedRepo``,
|
||||
# ``RepoNotFound``) when the host reports one. Lets us tell "this repo is
|
||||
# gated — request access" apart from "you just need a token".
|
||||
error_code: Optional[str] = None
|
||||
# Filename the server intends this response to be saved as: the
|
||||
# ``Content-Disposition`` name if present, else the post-redirect URL's
|
||||
# basename. Used to resolve the real extension for URLs (e.g. Civitai's
|
||||
# ``/api/download`` endpoints) that carry no extension in their path.
|
||||
filename: Optional[str] = None
|
||||
|
||||
@property
|
||||
def is_gated_repo(self) -> bool:
|
||||
"""True when the host says the repo is gated (access must be granted).
|
||||
|
||||
Distinct from a plain missing/invalid token: even a valid credential
|
||||
won't help until the user accepts the model's terms on its page.
|
||||
"""
|
||||
return (self.error_code or "").lower() == "gatedrepo"
|
||||
|
||||
|
||||
def _error_detail(error_code: Optional[str], error_message: Optional[str]) -> str:
|
||||
"""Format the host's ``X-Error-Code``/``X-Error-Message`` for logs/messages."""
|
||||
detail = ": ".join(p.strip() for p in (error_code, error_message) if p and p.strip())
|
||||
return f" ({detail})" if detail else ""
|
||||
|
||||
|
||||
def _probe_failure_message(
|
||||
status: int, error_code: Optional[str], error_message: Optional[str]
|
||||
) -> str:
|
||||
msg = f"probe returned HTTP {status}{_error_detail(error_code, error_message)}"
|
||||
if status == 404:
|
||||
# HuggingFace returns 404 (not 403) for a private repo the current
|
||||
# credentials cannot see, so it is indistinguishable from a missing
|
||||
# file without the hint. Name both causes so the user can check the
|
||||
# URL or their access/token scope.
|
||||
msg += (
|
||||
" — the file may not exist, or it is private/gated and the "
|
||||
"credentials in use lack access to it"
|
||||
)
|
||||
return msg
|
||||
|
||||
|
||||
def _total_from_content_range(value: Optional[str]) -> Optional[int]:
|
||||
# "bytes 0-0/12345" -> 12345 ; "bytes 0-0/*" -> None
|
||||
if not value or "/" not in value:
|
||||
return None
|
||||
total = value.rsplit("/", 1)[1].strip()
|
||||
return parse_int_header(total)
|
||||
|
||||
|
||||
def _filename_from_response(
|
||||
content_disposition: Optional[str], final_url: Optional[str]
|
||||
) -> Optional[str]:
|
||||
name = filename_from_content_disposition(content_disposition)
|
||||
if name:
|
||||
return name
|
||||
if final_url:
|
||||
base = urlsplit(final_url).path.rsplit("/", 1)[-1]
|
||||
if base:
|
||||
return base
|
||||
return None
|
||||
|
||||
|
||||
async def probe(url: str) -> ProbeResult:
|
||||
"""Probe ``url`` and return discovered metadata, failing soft."""
|
||||
try:
|
||||
async with open_validated(
|
||||
"GET",
|
||||
url,
|
||||
headers={"Range": "bytes=0-0", "Accept-Encoding": "identity"},
|
||||
timeout=_PROBE_TIMEOUT,
|
||||
) as (resp, final_url):
|
||||
# HuggingFace (and some others) report the real reason in these
|
||||
# headers on any status, including 404 for a private/missing repo.
|
||||
error_code = resp.headers.get("X-Error-Code")
|
||||
error_message = resp.headers.get("X-Error-Message")
|
||||
if resp.status in (401, 403):
|
||||
logging.warning(
|
||||
"[model_downloader] probe %s -> HTTP %d%s",
|
||||
redact_url(final_url or url), resp.status,
|
||||
_error_detail(error_code, error_message),
|
||||
)
|
||||
return ProbeResult(
|
||||
ok=False, status=resp.status, final_url=final_url, gated=True,
|
||||
error_code=error_code,
|
||||
error=(
|
||||
error_message
|
||||
or f"host returned {resp.status} (authentication required)"
|
||||
),
|
||||
)
|
||||
if resp.status not in (200, 206):
|
||||
logging.warning(
|
||||
"[model_downloader] probe %s -> HTTP %d%s",
|
||||
redact_url(final_url or url), resp.status,
|
||||
_error_detail(error_code, error_message),
|
||||
)
|
||||
return ProbeResult(
|
||||
ok=False, status=resp.status, final_url=final_url,
|
||||
error_code=error_code,
|
||||
error=_probe_failure_message(resp.status, error_code, error_message),
|
||||
)
|
||||
|
||||
headers = resp.headers
|
||||
accept_ranges = False
|
||||
total: Optional[int] = None
|
||||
if resp.status == 206:
|
||||
accept_ranges = True
|
||||
total = _total_from_content_range(headers.get("Content-Range"))
|
||||
else: # 200: server ignored the range
|
||||
accept_ranges = headers.get("Accept-Ranges", "").lower() == "bytes"
|
||||
total = parse_int_header(headers.get("Content-Length"))
|
||||
|
||||
if total is None:
|
||||
total = parse_int_header(headers.get("X-Linked-Size"))
|
||||
|
||||
return ProbeResult(
|
||||
ok=True,
|
||||
status=resp.status,
|
||||
final_url=final_url,
|
||||
total_bytes=total,
|
||||
accept_ranges=accept_ranges,
|
||||
etag=headers.get("ETag"),
|
||||
last_modified=headers.get("Last-Modified"),
|
||||
filename=_filename_from_response(
|
||||
headers.get("Content-Disposition"), final_url
|
||||
),
|
||||
)
|
||||
except Exception as e: # network / SSRF / timeout
|
||||
host = urlparse(url).netloc or "<unknown>"
|
||||
logging.debug("[model_downloader] probe failed for %s: %s", host, type(e).__name__)
|
||||
return ProbeResult(ok=False, status=0, error="probe failed: network error")
|
||||
|
||||
|
||||
def gated_error_message(url: str, pr: ProbeResult) -> str:
|
||||
"""Build a user-facing message for a gated/auth-required probe result.
|
||||
|
||||
Distinguishes a *gated* repo (access must be requested/granted on the model
|
||||
page — a token alone is not enough) from a plain missing/invalid credential.
|
||||
"""
|
||||
redacted = redact_url(url)
|
||||
if pr.is_gated_repo:
|
||||
detail = (pr.error or "access is restricted").rstrip()
|
||||
if detail and not detail.endswith((".", "!", "?")):
|
||||
detail += "."
|
||||
return (
|
||||
f"{redacted} is a gated model — {detail} Request access on the model's "
|
||||
f"page, authenticate this host via /api/download/auth (or set its API "
|
||||
f"key env var), and retry."
|
||||
)
|
||||
return (
|
||||
f"{redacted} requires authentication. Authenticate this host via "
|
||||
f"/api/download/auth or set its API key env var, and retry."
|
||||
)
|
||||
@@ -0,0 +1,72 @@
|
||||
"""Lazily-created shared :class:`aiohttp.ClientSession`.
|
||||
|
||||
A single session reuses TLS handshakes and TCP connections across the probe
|
||||
and the many segment GETs to the same host (HuggingFace is the dominant
|
||||
case), which is a large speedup on cold connections and exactly the
|
||||
connection-reuse strategy that lets us match aria2c.
|
||||
|
||||
The connector uses :class:`ValidatingResolver` so every connection — initial
|
||||
or post-redirect — is screened for private/special-use IPs at connect time.
|
||||
TLS is pinned to certifi's CA bundle because the OS trust store is not wired
|
||||
up on some Python installs (python.org macOS, slim containers).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import ssl
|
||||
from typing import Optional
|
||||
|
||||
import aiohttp
|
||||
|
||||
try:
|
||||
import certifi
|
||||
_CA_FILE = certifi.where()
|
||||
except Exception: # pragma: no cover - certifi is a transitive dep of aiohttp
|
||||
_CA_FILE = None
|
||||
|
||||
from comfy.cli_args import args
|
||||
from app.model_downloader.security.ssrf import ValidatingResolver
|
||||
|
||||
_session: Optional[aiohttp.ClientSession] = None
|
||||
_lock = asyncio.Lock()
|
||||
|
||||
|
||||
def ssl_context() -> ssl.SSLContext:
|
||||
if _CA_FILE is not None:
|
||||
return ssl.create_default_context(cafile=_CA_FILE)
|
||||
return ssl.create_default_context()
|
||||
|
||||
|
||||
async def get_session() -> aiohttp.ClientSession:
|
||||
"""Return the shared session, creating it on first use."""
|
||||
global _session
|
||||
if _session is not None and not _session.closed:
|
||||
return _session
|
||||
async with _lock:
|
||||
if _session is None or _session.closed:
|
||||
connector = aiohttp.TCPConnector(
|
||||
limit_per_host=max(1, getattr(args, "download_max_connections_per_host", 16)),
|
||||
ssl=ssl_context(),
|
||||
resolver=ValidatingResolver(),
|
||||
)
|
||||
_session = aiohttp.ClientSession(connector=connector)
|
||||
return _session
|
||||
|
||||
|
||||
async def close_session() -> None:
|
||||
global _session
|
||||
if _session is not None and not _session.closed:
|
||||
await _session.close()
|
||||
_session = None
|
||||
|
||||
|
||||
def parse_int_header(value: Optional[str]) -> Optional[int]:
|
||||
"""Parse a non-negative integer header value, or None if bad/absent."""
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
n = int(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
return n if n >= 0 else None
|
||||
@@ -0,0 +1,176 @@
|
||||
"""Priority scheduler + lifecycle.
|
||||
|
||||
Owns the set of running jobs and admits queued downloads up to a global
|
||||
concurrency limit (K), highest priority first, FIFO within a priority. Runs
|
||||
entirely on the existing ComfyUI asyncio loop; blocking work (disk, hashing,
|
||||
DB) is offloaded by the job/writer layers.
|
||||
|
||||
On startup it reconciles DB vs. disk: ``active``/``verifying`` rows left by a
|
||||
previous run are reset to ``queued`` and resumed from persisted offsets, and
|
||||
orphaned ``.part`` files with no live download row are swept.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from typing import Callable, Optional
|
||||
|
||||
from comfy.cli_args import args
|
||||
from app.model_downloader.constants import DownloadStatus
|
||||
from app.model_downloader.database import queries
|
||||
from app.model_downloader.engine.job import DownloadJob, JobSpec
|
||||
from app.model_downloader.security import paths
|
||||
|
||||
# Backoff for retryable failures
|
||||
_BACKOFF_BASE = 2.0
|
||||
_BACKOFF_CAP = 300.0
|
||||
_MAX_ATTEMPTS = 6
|
||||
|
||||
|
||||
class Scheduler:
|
||||
def __init__(self) -> None:
|
||||
self._jobs: dict[str, DownloadJob] = {}
|
||||
self._tasks: dict[str, asyncio.Task] = {}
|
||||
self._backoff_until: dict[str, float] = {}
|
||||
self._pump_lock = asyncio.Lock()
|
||||
self._notify_cb: Optional[Callable[[str], None]] = None
|
||||
self._started = False
|
||||
|
||||
@property
|
||||
def max_active(self) -> int:
|
||||
return max(1, getattr(args, "download_max_active", 3))
|
||||
|
||||
def set_notify(self, cb: Optional[Callable[[str], None]]) -> None:
|
||||
self._notify_cb = cb
|
||||
|
||||
def get_job(self, download_id: str) -> Optional[DownloadJob]:
|
||||
return self._jobs.get(download_id)
|
||||
|
||||
def is_active(self, download_id: str) -> bool:
|
||||
return download_id in self._tasks
|
||||
|
||||
# ----- startup -----
|
||||
|
||||
async def start(self) -> None:
|
||||
if self._started:
|
||||
return
|
||||
self._started = True
|
||||
try:
|
||||
await asyncio.to_thread(queries.reconcile_live_downloads)
|
||||
await asyncio.to_thread(self._sweep_orphan_temp_files)
|
||||
except Exception as e:
|
||||
logging.warning("[model_downloader] startup reconcile failed: %s", e)
|
||||
await self.pump()
|
||||
|
||||
@staticmethod
|
||||
def _sweep_orphan_temp_files() -> None:
|
||||
"""Remove ``.part`` files not referenced by a resumable download row.
|
||||
|
||||
Resumable partials are preserved; only truly orphaned temp files from
|
||||
crashed runs are deleted. ``FAILED`` is included because
|
||||
:meth:`DownloadManager.resume` explicitly permits resuming a
|
||||
retry-exhausted failed row: deleting its partial here while the
|
||||
per-segment offsets survive in the DB would make the next resume
|
||||
preallocate a fresh sparse file, skip every "complete" segment, and
|
||||
leave zero-filled holes that pass the size-only verification gate.
|
||||
"""
|
||||
live = {
|
||||
row.temp_path
|
||||
for row in queries.list_downloads()
|
||||
if row.status
|
||||
in (
|
||||
DownloadStatus.QUEUED,
|
||||
DownloadStatus.PAUSED,
|
||||
DownloadStatus.FAILED,
|
||||
)
|
||||
}
|
||||
for path in paths.iter_all_tmp_paths():
|
||||
if path in live:
|
||||
continue
|
||||
try:
|
||||
os.remove(path)
|
||||
logging.info("[model_downloader] removed orphan temp file: %s", path)
|
||||
except OSError as e:
|
||||
logging.warning("[model_downloader] could not remove %s: %s", path, e)
|
||||
|
||||
# ----- admission -----
|
||||
|
||||
async def pump(self) -> None:
|
||||
async with self._pump_lock:
|
||||
slots = self.max_active - len(self._tasks)
|
||||
if slots <= 0:
|
||||
return
|
||||
now = time.monotonic()
|
||||
candidates = await asyncio.to_thread(queries.list_queued_downloads)
|
||||
for row in candidates:
|
||||
if slots <= 0:
|
||||
break
|
||||
if row.id in self._tasks:
|
||||
continue
|
||||
if self._backoff_until.get(row.id, 0.0) > now:
|
||||
continue
|
||||
self._admit(row)
|
||||
slots -= 1
|
||||
|
||||
def _admit(self, row) -> None:
|
||||
spec = JobSpec(
|
||||
download_id=row.id,
|
||||
url=row.url,
|
||||
model_id=row.model_id,
|
||||
dest_path=row.dest_path,
|
||||
temp_path=row.temp_path,
|
||||
priority=row.priority,
|
||||
expected_sha256=row.expected_sha256,
|
||||
allow_any_extension=row.allow_any_extension,
|
||||
etag=row.etag,
|
||||
attempts=row.attempts,
|
||||
)
|
||||
job = DownloadJob(spec, notify_cb=self._notify_cb)
|
||||
self._jobs[row.id] = job
|
||||
self._tasks[row.id] = asyncio.ensure_future(self._run_job(job))
|
||||
|
||||
async def _run_job(self, job: DownloadJob) -> None:
|
||||
download_id = job.spec.download_id
|
||||
status = DownloadStatus.FAILED
|
||||
try:
|
||||
status = await job.run()
|
||||
except Exception as e: # run() is defensive, but never let a task die silently
|
||||
logging.error("[model_downloader] job %s crashed: %s", download_id, e)
|
||||
queries.update_download(
|
||||
download_id,
|
||||
status=DownloadStatus.FAILED,
|
||||
error=f"internal error: {e}",
|
||||
)
|
||||
if self._notify_cb:
|
||||
self._notify_cb(download_id)
|
||||
finally:
|
||||
self._tasks.pop(download_id, None)
|
||||
self._jobs.pop(download_id, None)
|
||||
|
||||
if status == DownloadStatus.QUEUED:
|
||||
if job.spec.attempts >= _MAX_ATTEMPTS:
|
||||
queries.update_download(
|
||||
download_id,
|
||||
status=DownloadStatus.FAILED,
|
||||
error=f"giving up after {job.spec.attempts} attempts",
|
||||
)
|
||||
if self._notify_cb:
|
||||
self._notify_cb(download_id)
|
||||
else:
|
||||
delay = min(
|
||||
_BACKOFF_CAP, _BACKOFF_BASE ** job.spec.attempts
|
||||
) + random.uniform(0, 1.0)
|
||||
self._backoff_until[download_id] = time.monotonic() + delay
|
||||
asyncio.ensure_future(self._delayed_pump(delay))
|
||||
await self.pump()
|
||||
|
||||
async def _delayed_pump(self, delay: float) -> None:
|
||||
await asyncio.sleep(delay)
|
||||
await self.pump()
|
||||
|
||||
|
||||
SCHEDULER = Scheduler()
|
||||
@@ -0,0 +1,140 @@
|
||||
"""URL allowlist for server-side model fetches.
|
||||
|
||||
Default-deny. A URL is downloadable only when its parsed host + scheme are
|
||||
allowlisted AND (unless explicitly relaxed) its final filename ends in a
|
||||
known model extension.
|
||||
|
||||
The built-in host defaults mirror the frontend's ``isModelDownloadable``
|
||||
allowlist so the two flows agree on what is eligible; ``--download-allowed-hosts``
|
||||
extends it for self-hosted mirrors. Matching is done on ``urlparse().hostname``
|
||||
(never a raw string prefix) so userinfo tricks like
|
||||
``http://127.0.0.1@169.254.169.254/x.safetensors`` — whose real host is the
|
||||
metadata IP — cannot slip past.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
# host -> set of allowed schemes. Frontend parity (HuggingFace / Civitai /
|
||||
# localhost). Extra hosts from --download-allowed-hosts are https-only.
|
||||
_DEFAULT_ALLOWED_HOSTS: dict[str, set[str]] = {
|
||||
"huggingface.co": {"https"},
|
||||
"civitai.com": {"https"},
|
||||
"localhost": {"http", "https"},
|
||||
"127.0.0.1": {"http", "https"},
|
||||
}
|
||||
|
||||
# Hosts for which loopback addresses are intentionally permitted (the localhost
|
||||
# "download a local model" feature). Every other host's loopback resolution is
|
||||
# rejected by the SSRF resolver.
|
||||
LOOPBACK_HOSTS = frozenset({"localhost", "127.0.0.1", "::1"})
|
||||
|
||||
# Known model file extensions (frontend parity). Checked on the final filename.
|
||||
ALLOWED_MODEL_EXTENSIONS = (
|
||||
".safetensors",
|
||||
".sft",
|
||||
".ckpt",
|
||||
".pth",
|
||||
".pt",
|
||||
".gguf",
|
||||
".bin",
|
||||
)
|
||||
|
||||
|
||||
def _allowed_hosts() -> dict[str, set[str]]:
|
||||
hosts = {h: set(s) for h, s in _DEFAULT_ALLOWED_HOSTS.items()}
|
||||
for extra in getattr(args, "download_allowed_hosts", []) or []:
|
||||
host = extra.strip().lower()
|
||||
if host:
|
||||
hosts.setdefault(host, set()).add("https")
|
||||
return hosts
|
||||
|
||||
|
||||
def is_host_allowed(host: str | None, scheme: str | None) -> bool:
|
||||
"""True iff ``host`` is allowlisted for ``scheme``.
|
||||
|
||||
Used both for the initial URL and re-checked on every redirect hop,
|
||||
so a whitelisted URL cannot 30x into an off-list host.
|
||||
"""
|
||||
if not host or not scheme:
|
||||
return False
|
||||
allowed = _allowed_hosts().get(host.lower())
|
||||
return allowed is not None and scheme.lower() in allowed
|
||||
|
||||
|
||||
def has_allowed_extension(path: str, allow_any_extension: bool = False) -> bool:
|
||||
if allow_any_extension:
|
||||
return True
|
||||
return path.lower().endswith(ALLOWED_MODEL_EXTENSIONS)
|
||||
|
||||
|
||||
def filename_extension(name: str) -> str:
|
||||
"""Lowercased extension (including the leading dot) of a bare filename.
|
||||
|
||||
Returns ``""`` when there is no extension. A leading-dot name
|
||||
(``.safetensors``) is treated as having no extension (all stem), matching
|
||||
``os.path.splitext`` semantics so dotfiles aren't mistaken for typed files.
|
||||
"""
|
||||
base = name.replace("\\", "/").rsplit("/", 1)[-1]
|
||||
dot = base.rfind(".")
|
||||
if dot <= 0:
|
||||
return ""
|
||||
return base[dot:].lower()
|
||||
|
||||
|
||||
def is_allowed_extension_name(name: str) -> bool:
|
||||
"""True iff ``name`` ends in one of the known model extensions."""
|
||||
return name.lower().endswith(ALLOWED_MODEL_EXTENSIONS)
|
||||
|
||||
|
||||
def is_host_allowed_url(url: str) -> bool:
|
||||
"""True iff ``url`` parses and its host+scheme are allowlisted."""
|
||||
if not isinstance(url, str) or not url:
|
||||
return False
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
except ValueError:
|
||||
return False
|
||||
return is_host_allowed(parsed.hostname, parsed.scheme)
|
||||
|
||||
|
||||
def url_path_extension(url: str) -> str:
|
||||
"""Extension of the URL *path* basename (query ignored), or ``""``."""
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
except ValueError:
|
||||
return ""
|
||||
return filename_extension(parsed.path)
|
||||
|
||||
|
||||
def is_url_downloadable(url: str) -> bool:
|
||||
"""Coarse enqueue gate: host/scheme allowed and extension not disallowed.
|
||||
|
||||
Unlike :func:`is_url_allowed` (which demands a known extension *in the URL*),
|
||||
this also admits URLs whose path carries no extension at all — e.g. a Civitai
|
||||
``/api/download/models/<id>`` endpoint whose real filename only shows up in
|
||||
the redirect target / ``Content-Disposition``. The true extension is then
|
||||
resolved from the network and re-validated before the download is admitted.
|
||||
A path bearing an explicit *non-model* extension (``.zip``, ``.html``, ...)
|
||||
is still rejected here.
|
||||
"""
|
||||
if not is_host_allowed_url(url):
|
||||
return False
|
||||
ext = url_path_extension(url)
|
||||
return ext == "" or ext in ALLOWED_MODEL_EXTENSIONS
|
||||
|
||||
|
||||
def is_url_allowed(url: str, allow_any_extension: bool = False) -> bool:
|
||||
"""Check whether ``url`` is permitted as a server-side download source."""
|
||||
if not isinstance(url, str) or not url:
|
||||
return False
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
except ValueError:
|
||||
return False
|
||||
if not is_host_allowed(parsed.hostname, parsed.scheme):
|
||||
return False
|
||||
return has_allowed_extension(parsed.path, allow_any_extension)
|
||||
@@ -0,0 +1,132 @@
|
||||
"""Path resolution + traversal safety for downloads.
|
||||
|
||||
A ``model_id`` is a *relative destination path* of the form
|
||||
``<directory>/<filename>`` (e.g. ``loras/my_lora.safetensors``). This module
|
||||
turns one into an absolute on-disk path under one of ComfyUI's registered
|
||||
model folders, rejecting unknown folders, path traversal, and symlink escape.
|
||||
This is the only thing that composes destination paths, so the engine never
|
||||
touches user-supplied path strings directly.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
from typing import Iterator, Optional
|
||||
|
||||
import folder_paths
|
||||
|
||||
from app.model_downloader.constants import TMP_SUFFIX
|
||||
from app.model_downloader.security.allowlist import ALLOWED_MODEL_EXTENSIONS
|
||||
|
||||
# A model_id component is a single path segment of safe characters — no slashes,
|
||||
# no "..", no leading dots that could escape the target directory.
|
||||
_SEGMENT_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]*$")
|
||||
|
||||
|
||||
class InvalidModelId(ValueError):
|
||||
"""Raised when a model_id is malformed or names an unknown model folder."""
|
||||
|
||||
|
||||
def parse_model_id(model_id: str, allow_any_extension: bool = False) -> tuple[str, str]:
|
||||
"""Split ``<directory>/<filename>`` and validate both components.
|
||||
|
||||
Returns ``(directory, filename)``. Does not touch the filesystem.
|
||||
"""
|
||||
if not isinstance(model_id, str) or "/" not in model_id:
|
||||
raise InvalidModelId(
|
||||
f"model_id must be '<directory>/<filename>', got {model_id!r}"
|
||||
)
|
||||
directory, _, filename = model_id.partition("/")
|
||||
if "/" in filename or not directory or not filename:
|
||||
raise InvalidModelId(
|
||||
f"model_id must have exactly one '/' separator, got {model_id!r}"
|
||||
)
|
||||
if not _SEGMENT_RE.match(directory):
|
||||
raise InvalidModelId(f"invalid directory segment {directory!r}")
|
||||
if not _SEGMENT_RE.match(filename):
|
||||
raise InvalidModelId(f"invalid filename segment {filename!r}")
|
||||
if not allow_any_extension and not filename.lower().endswith(
|
||||
ALLOWED_MODEL_EXTENSIONS
|
||||
):
|
||||
raise InvalidModelId(
|
||||
f"filename must end with a known model extension "
|
||||
f"{ALLOWED_MODEL_EXTENSIONS}, got {filename!r}"
|
||||
)
|
||||
if directory not in folder_paths.folder_names_and_paths:
|
||||
raise InvalidModelId(f"unknown model folder {directory!r}")
|
||||
return directory, filename
|
||||
|
||||
|
||||
def apply_extension(model_id: str, ext: str) -> str:
|
||||
"""Return ``model_id`` with its filename forced to end in ``ext``.
|
||||
|
||||
``ext`` includes the leading dot (e.g. ``".safetensors"``). If the filename
|
||||
already ends in a *known model extension* it is replaced; otherwise ``ext``
|
||||
is appended (so ``loras/mymodel`` -> ``loras/mymodel.safetensors`` and
|
||||
``loras/mymodel.ckpt`` -> ``loras/mymodel.safetensors``). A filename with a
|
||||
non-model suffix (``my.model.v2``) is treated as an extensionless stem and
|
||||
``ext`` is appended. The directory part is left untouched; validation is
|
||||
still the caller's job via :func:`parse_model_id`.
|
||||
"""
|
||||
directory, sep, filename = model_id.partition("/")
|
||||
if not sep:
|
||||
return model_id # malformed; parse_model_id will reject it
|
||||
low = filename.lower()
|
||||
for known in ALLOWED_MODEL_EXTENSIONS:
|
||||
if low.endswith(known):
|
||||
filename = filename[: -len(known)]
|
||||
break
|
||||
return f"{directory}{sep}{filename}{ext}"
|
||||
|
||||
|
||||
def resolve_existing(model_id: str, allow_any_extension: bool = False) -> Optional[str]:
|
||||
"""Return the absolute path of an installed model, or None if missing.
|
||||
|
||||
Honours ``extra_model_paths.yaml`` transparently via ``get_full_path``.
|
||||
"""
|
||||
directory, filename = parse_model_id(model_id, allow_any_extension)
|
||||
return folder_paths.get_full_path(directory, filename)
|
||||
|
||||
|
||||
def resolve_destination(
|
||||
model_id: str, allow_any_extension: bool = False
|
||||
) -> tuple[str, str]:
|
||||
"""Return ``(final_path, temp_path)`` for a download.
|
||||
|
||||
Downloads land at the first registered path for the model's directory
|
||||
(the "primary" location). ``temp_path`` is a sibling ``.part`` file that
|
||||
is atomically renamed onto ``final_path`` on success. The result is
|
||||
asserted to stay within the registered root (defence in depth on top of
|
||||
the segment regex).
|
||||
"""
|
||||
directory, filename = parse_model_id(model_id, allow_any_extension)
|
||||
roots = folder_paths.get_folder_paths(directory)
|
||||
if not roots:
|
||||
raise InvalidModelId(f"no on-disk path registered for folder {directory!r}")
|
||||
root = os.path.realpath(roots[0])
|
||||
final_path = os.path.realpath(os.path.join(root, filename))
|
||||
if final_path != root and not final_path.startswith(root + os.sep):
|
||||
raise InvalidModelId(f"resolved path escapes model root: {model_id!r}")
|
||||
temp_path = f"{final_path}{TMP_SUFFIX}"
|
||||
return final_path, temp_path
|
||||
|
||||
|
||||
def iter_all_tmp_paths() -> Iterator[str]:
|
||||
"""Yield this subsystem's temp files under every registered model folder.
|
||||
|
||||
Matches only the distinctive ``TMP_SUFFIX`` so the startup orphan sweep
|
||||
can never delete temp files created by other tools.
|
||||
"""
|
||||
seen_roots: set[str] = set()
|
||||
for directory in list(folder_paths.folder_names_and_paths.keys()):
|
||||
for root in folder_paths.get_folder_paths(directory):
|
||||
if root in seen_roots or not os.path.isdir(root):
|
||||
continue
|
||||
seen_roots.add(root)
|
||||
try:
|
||||
for entry in os.scandir(root):
|
||||
if entry.is_file() and entry.name.endswith(TMP_SUFFIX):
|
||||
yield entry.path
|
||||
except OSError:
|
||||
continue
|
||||
@@ -0,0 +1,163 @@
|
||||
"""SSRF / exfiltration defenses.
|
||||
|
||||
Two cooperating layers:
|
||||
|
||||
1. :class:`ValidatingResolver` is installed on the shared connector. Every
|
||||
connection — the initial probe and every segment GET, including ones made
|
||||
after a redirect — resolves its host through this resolver, which rejects
|
||||
any address that lands on a private / special-use IP range. Because the
|
||||
resolve and the connect happen together inside the connector, there is no
|
||||
check-then-connect window for DNS rebinding to exploit.
|
||||
|
||||
2. :func:`check_redirect_hop` re-validates every hop. The host allowlist gates
|
||||
only the *initial* user-supplied URL (anti-SSRF for arbitrary input);
|
||||
legitimate downloads from allowlisted origins redirect to presigned CDN
|
||||
hosts that are deliberately NOT on the allowlist (HF ->
|
||||
``cdn-lfs*.huggingface.co``, Civitai -> signed Cloudflare/S3), so hops are
|
||||
instead screened for scheme, embedded credentials, and — via the resolver
|
||||
above — private IPs. Credentials are only ever attached when a hop's host
|
||||
exactly matches a stored credential, so they are dropped on the CDN hop.
|
||||
Loopback (the "download a local model" feature) is exempt from IP filtering
|
||||
only for the initial URL: a *redirect* may never target a loopback host or
|
||||
a blocked IP-literal, which the resolver alone can't enforce (it exempts
|
||||
loopback literals and never sees IP literals through DNS).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
import socket
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from aiohttp.abc import AbstractResolver
|
||||
from aiohttp.resolver import DefaultResolver
|
||||
|
||||
from app.model_downloader.security.allowlist import LOOPBACK_HOSTS
|
||||
|
||||
# Cap the redirect chain length a hop may use.
|
||||
MAX_REDIRECTS = 5
|
||||
|
||||
|
||||
class SSRFError(Exception):
|
||||
"""A hop failed an SSRF / allowlist check."""
|
||||
|
||||
|
||||
def is_scheme_allowed(scheme: str | None, host: str | None) -> bool:
|
||||
"""True iff ``scheme`` is permitted for ``host`` on a download hop.
|
||||
|
||||
https is always allowed; plain http only for loopback/approved dev hosts.
|
||||
"""
|
||||
if not scheme:
|
||||
return False
|
||||
scheme = scheme.lower()
|
||||
if scheme == "https":
|
||||
return True
|
||||
if scheme == "http":
|
||||
return bool(host) and host.lower() in LOOPBACK_HOSTS
|
||||
return False
|
||||
|
||||
|
||||
def is_blocked_ip(ip_str: str) -> bool:
|
||||
"""True for any address we refuse to connect to.
|
||||
|
||||
Covers loopback, link-local (incl. 169.254.169.254 cloud metadata),
|
||||
RFC1918 private ranges, unique-local (ULA), unspecified (0.0.0.0/::),
|
||||
multicast and other reserved ranges.
|
||||
"""
|
||||
try:
|
||||
ip = ipaddress.ip_address(ip_str)
|
||||
except ValueError:
|
||||
return True # unparseable -> refuse
|
||||
# On CPython before the gh-113171 fix (backported to 3.12.4/3.11.9/
|
||||
# 3.10.14/3.9.19) the is_* properties don't see through IPv4-mapped IPv6
|
||||
# (e.g. ::ffff:169.254.169.254), so resolve and re-check the embedded IPv4
|
||||
# to keep mapped metadata/private addresses from slipping past the filter.
|
||||
mapped = getattr(ip, "ipv4_mapped", None)
|
||||
if mapped is not None:
|
||||
ip = mapped
|
||||
return (
|
||||
ip.is_private
|
||||
or ip.is_loopback
|
||||
or ip.is_link_local
|
||||
or ip.is_multicast
|
||||
or ip.is_reserved
|
||||
or ip.is_unspecified
|
||||
)
|
||||
|
||||
|
||||
class ValidatingResolver(AbstractResolver):
|
||||
"""Delegating resolver that drops blocked IPs from every resolution.
|
||||
|
||||
If a hostname resolves only to blocked addresses, the connection fails
|
||||
closed with an :class:`OSError`, which aiohttp surfaces as a connection
|
||||
error to the caller.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._inner = DefaultResolver()
|
||||
|
||||
async def resolve(self, host, port=0, family=socket.AF_INET):
|
||||
infos = await self._inner.resolve(host, port, family)
|
||||
# localhost/127.0.0.1 are an explicit, opt-in allowlist feature.
|
||||
if isinstance(host, str) and host.lower() in LOOPBACK_HOSTS:
|
||||
return infos
|
||||
safe = [info for info in infos if not is_blocked_ip(info["host"])]
|
||||
if not safe:
|
||||
raise OSError(
|
||||
f"refusing to connect to {host!r}: resolves only to "
|
||||
f"private/special-use addresses"
|
||||
)
|
||||
return safe
|
||||
|
||||
async def close(self) -> None:
|
||||
await self._inner.close()
|
||||
|
||||
|
||||
def check_redirect_hop(url: str, *, is_initial_url: bool = False) -> str:
|
||||
"""Validate one hop's URL.
|
||||
|
||||
Returns the URL unchanged on success; raises :class:`SSRFError` otherwise.
|
||||
Requires https for external hosts (http only for loopback/approved dev
|
||||
hosts) and forbids credentials-in-URL. The host is NOT re-checked against
|
||||
the allowlist (CDN redirect targets are off-list by design); credential
|
||||
leakage is prevented by exact host matching at attach time, and the landing
|
||||
filename's extension is gated separately by the caller.
|
||||
|
||||
Loopback/blocked-IP screening: the connector's resolver filters resolvable
|
||||
hostnames but exempts literal loopback hosts (``localhost``/``127.0.0.1``/
|
||||
``::1``) and never sees IP literals through DNS. That loopback exemption is
|
||||
legitimate only for the *initial* user-supplied URL (``is_initial_url``);
|
||||
on a redirect hop we reject loopback hosts and any blocked IP-literal here,
|
||||
so a 30x can't steer a server-side GET at loopback/internal services.
|
||||
"""
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
except ValueError as e:
|
||||
raise SSRFError(f"unparseable redirect URL {url!r}: {e}") from e
|
||||
host = parsed.hostname
|
||||
if not host:
|
||||
raise SSRFError(f"redirect URL has no host: {url!r}")
|
||||
if not is_scheme_allowed(parsed.scheme, host):
|
||||
raise SSRFError(
|
||||
f"redirect to disallowed scheme {parsed.scheme!r} for host "
|
||||
f"{host!r} (https required for external hosts)"
|
||||
)
|
||||
if parsed.username or parsed.password:
|
||||
raise SSRFError("credentials-in-URL are not allowed")
|
||||
host_is_loopback = host.lower() in LOOPBACK_HOSTS
|
||||
if not is_initial_url and host_is_loopback:
|
||||
raise SSRFError(f"redirect to loopback host {host!r} is not allowed")
|
||||
# IP-literal targets never go through DNS, so the connector's resolver can't
|
||||
# screen them — check them directly. The only blocked IP allowed through is
|
||||
# a loopback literal on the initial URL (handled by the exemption above).
|
||||
try:
|
||||
ipaddress.ip_address(host)
|
||||
except ValueError:
|
||||
is_ip_literal = False
|
||||
else:
|
||||
is_ip_literal = True
|
||||
if is_ip_literal and is_blocked_ip(host) and not (
|
||||
is_initial_url and host_is_loopback
|
||||
):
|
||||
raise SSRFError(f"redirect to blocked internal address {host!r}")
|
||||
return url
|
||||
@@ -0,0 +1,49 @@
|
||||
"""Hub-checksum verification = SHA256.
|
||||
|
||||
Only used to confirm a download matches a *provided* ``expected_sha256``. It
|
||||
is NOT the dedup key (that is blake3, owned by the assets system). The full
|
||||
sequential read happens at most once, here, only when a checksum was supplied.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from typing import Callable, Optional
|
||||
|
||||
_CHUNK = 8 * 1024 * 1024
|
||||
|
||||
InterruptCheck = Callable[[], bool]
|
||||
|
||||
|
||||
class ChecksumError(Exception):
|
||||
"""The computed SHA256 did not match the expected value."""
|
||||
|
||||
|
||||
def sha256_file(path: str, interrupt_check: Optional[InterruptCheck] = None) -> Optional[str]:
|
||||
"""Stream the file and return its lowercase hex SHA256.
|
||||
|
||||
Returns ``None`` if interrupted via ``interrupt_check``.
|
||||
"""
|
||||
h = hashlib.sha256()
|
||||
with open(path, "rb") as f:
|
||||
while True:
|
||||
if interrupt_check is not None and interrupt_check():
|
||||
return None
|
||||
chunk = f.read(_CHUNK)
|
||||
if not chunk:
|
||||
break
|
||||
h.update(chunk)
|
||||
return h.hexdigest()
|
||||
|
||||
|
||||
def verify_sha256(
|
||||
path: str, expected: str, interrupt_check: Optional[InterruptCheck] = None
|
||||
) -> None:
|
||||
"""Raise :class:`ChecksumError` unless the file's SHA256 matches ``expected``."""
|
||||
actual = sha256_file(path, interrupt_check)
|
||||
if actual is None:
|
||||
return # interrupted; caller will re-verify on resume
|
||||
if actual.lower() != expected.lower():
|
||||
raise ChecksumError(
|
||||
f"sha256 mismatch: expected {expected.lower()}, got {actual.lower()}"
|
||||
)
|
||||
@@ -0,0 +1,53 @@
|
||||
"""Dedup + catalog handoff — reuse the assets system.
|
||||
|
||||
We do NOT build a parallel indexer. "Do I already have it?" is answered by
|
||||
``resolve_existing`` (path) at enqueue time and, where a hash is known, by the
|
||||
assets blake3 catalog. After a completed download we register the file
|
||||
through the assets ingest path so it is cataloged and (eventually) hashed by
|
||||
the existing enrichment worker.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
|
||||
def _register_sync(abs_path: str) -> Optional[str]:
|
||||
"""Register a finished file into the assets catalog. Returns asset hash."""
|
||||
try:
|
||||
from app.assets.services.ingest import register_file_in_place
|
||||
except Exception as e: # assets package import failure — non-fatal
|
||||
logging.debug("[model_downloader] assets ingest unavailable: %s", e)
|
||||
return None
|
||||
try:
|
||||
result = register_file_in_place(abs_path, name=os.path.basename(abs_path), tags=[])
|
||||
return result.asset.hash if result and result.asset else None
|
||||
except Exception as e:
|
||||
# The file is already safely on disk; cataloging is best-effort.
|
||||
logging.warning(
|
||||
"[model_downloader] could not register %s into assets catalog: %s",
|
||||
abs_path, e,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
async def register_completed(abs_path: str) -> Optional[str]:
|
||||
"""Catalog a completed download via the assets system (off the event loop)."""
|
||||
return await asyncio.to_thread(_register_sync, abs_path)
|
||||
|
||||
|
||||
def _find_by_hash_sync(blake3_hex: str) -> Optional[str]:
|
||||
try:
|
||||
from app.assets.services.asset_management import get_asset_by_hash
|
||||
except Exception:
|
||||
return None
|
||||
asset = get_asset_by_hash("blake3:" + blake3_hex)
|
||||
return asset.hash if asset is not None else None
|
||||
|
||||
|
||||
async def find_existing_by_hash(blake3_hex: str) -> Optional[str]:
|
||||
"""Pure DB lookup — never triggers hashing on the hot path."""
|
||||
return await asyncio.to_thread(_find_by_hash_sync, blake3_hex)
|
||||
@@ -0,0 +1,86 @@
|
||||
"""Cheap structural validation, no full read.
|
||||
|
||||
For ``.safetensors``/``.sft`` we parse the header (first few KB): it carries
|
||||
the tensor table and the byte length of the data region. We assert
|
||||
``file_size == 8 + header_len + data_region_len``. This detects truncation
|
||||
and most corruption for free, before any crypto hashing. Other extensions
|
||||
have no cheap structural check and pass through.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import struct
|
||||
from typing import Optional
|
||||
|
||||
_SAFETENSORS_EXTS = (".safetensors", ".sft")
|
||||
# A sane upper bound so a corrupt header length can't make us read gigabytes.
|
||||
_MAX_HEADER_BYTES = 100 * 1024 * 1024
|
||||
|
||||
|
||||
class StructuralError(Exception):
|
||||
"""The file failed its structural integrity check."""
|
||||
|
||||
|
||||
def validate(path: str, name_hint: Optional[str] = None) -> None:
|
||||
"""Validate the file at ``path``. Raises :class:`StructuralError` on failure.
|
||||
|
||||
The file format is detected from ``name_hint`` when provided, otherwise from
|
||||
``path``. Callers that download into a temp file with an opaque suffix (e.g.
|
||||
``*.comfy-download.part``) must pass the final destination name as
|
||||
``name_hint`` so the format check is not silently skipped.
|
||||
"""
|
||||
lower = (name_hint or path).lower()
|
||||
if lower.endswith(_SAFETENSORS_EXTS):
|
||||
_validate_safetensors(path)
|
||||
# No structural check for other formats; the size + (optional) checksum
|
||||
# gates in the engine cover those.
|
||||
|
||||
|
||||
def _validate_safetensors(path: str) -> None:
|
||||
file_size = os.path.getsize(path)
|
||||
if file_size < 8:
|
||||
raise StructuralError(f"file too small to be safetensors ({file_size} bytes)")
|
||||
with open(path, "rb") as f:
|
||||
header_len = struct.unpack("<Q", f.read(8))[0]
|
||||
if header_len <= 0 or header_len > _MAX_HEADER_BYTES:
|
||||
raise StructuralError(f"implausible safetensors header length {header_len}")
|
||||
if 8 + header_len > file_size:
|
||||
raise StructuralError("safetensors header extends past end of file")
|
||||
try:
|
||||
header = json.loads(f.read(header_len).decode("utf-8"))
|
||||
except (UnicodeDecodeError, json.JSONDecodeError) as e:
|
||||
raise StructuralError(f"safetensors header is not valid JSON: {e}") from e
|
||||
|
||||
if not isinstance(header, dict):
|
||||
raise StructuralError("safetensors header is not a JSON object")
|
||||
|
||||
data_len = 0
|
||||
for name, entry in header.items():
|
||||
if name == "__metadata__":
|
||||
continue
|
||||
if not isinstance(entry, dict) or "data_offsets" not in entry:
|
||||
raise StructuralError(f"tensor {name!r} missing data_offsets")
|
||||
offsets = entry["data_offsets"]
|
||||
if not (isinstance(offsets, list) and len(offsets) == 2):
|
||||
raise StructuralError(f"tensor {name!r} has malformed data_offsets")
|
||||
begin, end = offsets
|
||||
# bool is an int subclass; reject it explicitly to avoid True/False offsets.
|
||||
if (
|
||||
not isinstance(begin, int)
|
||||
or not isinstance(end, int)
|
||||
or isinstance(begin, bool)
|
||||
or isinstance(end, bool)
|
||||
or begin < 0
|
||||
or end < begin
|
||||
):
|
||||
raise StructuralError(f"tensor {name!r} has malformed data_offsets")
|
||||
data_len = max(data_len, end)
|
||||
|
||||
expected = 8 + header_len + data_len
|
||||
if file_size != expected:
|
||||
raise StructuralError(
|
||||
f"size mismatch: file is {file_size} bytes, header implies {expected} "
|
||||
f"(8 + {header_len} header + {data_len} data)"
|
||||
)
|
||||
@@ -33,6 +33,28 @@ class EnumAction(argparse.Action):
|
||||
setattr(namespace, self.dest, value)
|
||||
|
||||
|
||||
def _positive_int(value: str) -> int:
|
||||
"""argparse type that rejects zero and negative integers."""
|
||||
try:
|
||||
ivalue = int(value)
|
||||
except ValueError:
|
||||
raise argparse.ArgumentTypeError(f"{value!r} is not an integer")
|
||||
if ivalue <= 0:
|
||||
raise argparse.ArgumentTypeError(f"{value!r} must be a positive integer (> 0)")
|
||||
return ivalue
|
||||
|
||||
|
||||
def _non_negative_int(value: str) -> int:
|
||||
"""argparse type that rejects negatives but allows zero (a disable sentinel)."""
|
||||
try:
|
||||
ivalue = int(value)
|
||||
except ValueError:
|
||||
raise argparse.ArgumentTypeError(f"{value!r} is not an integer")
|
||||
if ivalue < 0:
|
||||
raise argparse.ArgumentTypeError(f"{value!r} must be a non-negative integer (>= 0)")
|
||||
return ivalue
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument("--listen", type=str, default="127.0.0.1", metavar="IP", nargs="?", const="0.0.0.0,::", help="Specify the IP address to listen on (default: 127.0.0.1). You can give a list of ip addresses by separating them with a comma like: 127.2.2.2,127.3.3.3 If --listen is provided without an argument, it defaults to 0.0.0.0,:: (listens on all ipv4 and ipv6)")
|
||||
@@ -246,6 +268,15 @@ parser.add_argument("--enable-asset-hashing", action="store_true", help="Compute
|
||||
parser.add_argument("--feature-flag", type=str, action='append', default=[], metavar="KEY[=VALUE]", help="Set a server feature flag. Use KEY=VALUE to set an explicit value, or bare KEY to set it to true. Can be specified multiple times. Boolean values (true/false) and numbers are auto-converted. Examples: --feature-flag show_signin_button=true or --feature-flag show_signin_button")
|
||||
parser.add_argument("--list-feature-flags", action="store_true", help="Print the registry of known CLI-settable feature flags as JSON and exit.")
|
||||
|
||||
# ----- Model download manager (PRD: docs/prd-download-manager.md) -----
|
||||
parser.add_argument("--download-segments", type=_positive_int, default=8, metavar="N", help="Number of parallel HTTP range segments per file for the model download manager (default: 8).")
|
||||
parser.add_argument("--download-max-active", type=_positive_int, default=3, metavar="N", help="Maximum number of model downloads running concurrently (default: 3).")
|
||||
parser.add_argument("--download-max-connections-per-host", type=_positive_int, default=16, metavar="N", help="Maximum simultaneous connections to a single host for the download manager (default: 16).")
|
||||
parser.add_argument("--download-chunk-size", type=_positive_int, default=4 * 1024 * 1024, metavar="BYTES", help="Read chunk size in bytes for the download manager (default: 4 MiB).")
|
||||
parser.add_argument("--download-max-bytes", type=_non_negative_int, default=1024 * 1024 * 1024 * 1024, metavar="BYTES", help="Maximum size in bytes of a single download; aborts transfers that exceed it (guards against malicious/non-conforming hosts filling the disk). Set to 0 to disable (default: 1 TiB).")
|
||||
parser.add_argument("--download-allowed-hosts", type=str, nargs="*", default=[], metavar="HOST", help="Additional hostnames to add to the download manager allowlist (https only). The built-in defaults always include huggingface.co and civitai.com.")
|
||||
parser.add_argument("--download-allow-any-extension", action="store_true", help="Allow the download manager to fetch files with any extension (default: only known model extensions like .safetensors).")
|
||||
|
||||
if comfy.options.args_parsing:
|
||||
args = parser.parse_args()
|
||||
else:
|
||||
|
||||
@@ -1,278 +0,0 @@
|
||||
import re
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
import comfy.ops
|
||||
import comfy.utils
|
||||
|
||||
|
||||
MODULE_PATTERN = re.compile(r"lllite_dit_blocks_(\d+)_(self_attn_[qkv]_proj|cross_attn_q_proj|mlp_layer1)$")
|
||||
|
||||
|
||||
def _group_norm(channels, device=None, dtype=None, operations=None):
|
||||
groups = 8
|
||||
while groups > 1 and channels % groups != 0:
|
||||
groups //= 2
|
||||
return operations.GroupNorm(groups, channels, device=device, dtype=dtype)
|
||||
|
||||
|
||||
class AnimaLLLiteResBlock(nn.Module):
|
||||
def __init__(self, channels, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
self.norm1 = _group_norm(channels, device=device, dtype=dtype, operations=operations)
|
||||
self.conv1 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype)
|
||||
self.norm2 = _group_norm(channels, device=device, dtype=dtype, operations=operations)
|
||||
self.conv2 = operations.Conv2d(channels, channels, kernel_size=3, padding=1, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x):
|
||||
h = self.conv1(F.silu(self.norm1(x)))
|
||||
h = self.conv2(F.silu(self.norm2(h)))
|
||||
return x + h
|
||||
|
||||
|
||||
class AnimaLLLiteASPP(nn.Module):
|
||||
def __init__(self, channels, dilations, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
branches = []
|
||||
for dilation in dilations:
|
||||
if dilation == 1:
|
||||
conv = operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype)
|
||||
else:
|
||||
conv = operations.Conv2d(channels, channels, kernel_size=3, padding=dilation, dilation=dilation, device=device, dtype=dtype)
|
||||
branches.append(nn.Sequential(conv, _group_norm(channels, device=device, dtype=dtype, operations=operations), nn.SiLU()))
|
||||
self.branches = nn.ModuleList(branches)
|
||||
self.global_pool = nn.AdaptiveAvgPool2d(1)
|
||||
self.global_conv = nn.Sequential(
|
||||
operations.Conv2d(channels, channels, kernel_size=1, device=device, dtype=dtype),
|
||||
_group_norm(channels, device=device, dtype=dtype, operations=operations),
|
||||
nn.SiLU(),
|
||||
)
|
||||
self.proj = nn.Sequential(
|
||||
operations.Conv2d(channels * (len(dilations) + 1), channels, kernel_size=1, device=device, dtype=dtype),
|
||||
_group_norm(channels, device=device, dtype=dtype, operations=operations),
|
||||
nn.SiLU(),
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
height, width = x.shape[-2:]
|
||||
outputs = [branch(x) for branch in self.branches]
|
||||
pooled = self.global_conv(self.global_pool(x))
|
||||
outputs.append(F.interpolate(pooled, size=(height, width), mode="bilinear", align_corners=False))
|
||||
return self.proj(torch.cat(outputs, dim=1))
|
||||
|
||||
|
||||
class AnimaLLLiteConditioning(nn.Module):
|
||||
def __init__(self, cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
half_dim = cond_dim // 2
|
||||
self.conv1 = operations.Conv2d(cond_in_channels, half_dim, kernel_size=4, stride=4, device=device, dtype=dtype)
|
||||
self.norm1 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations)
|
||||
self.conv2 = operations.Conv2d(half_dim, half_dim, kernel_size=3, padding=1, device=device, dtype=dtype)
|
||||
self.norm2 = _group_norm(half_dim, device=device, dtype=dtype, operations=operations)
|
||||
self.conv3 = operations.Conv2d(half_dim, cond_dim, kernel_size=4, stride=4, device=device, dtype=dtype)
|
||||
self.norm3 = _group_norm(cond_dim, device=device, dtype=dtype, operations=operations)
|
||||
self.resblocks = nn.ModuleList([
|
||||
AnimaLLLiteResBlock(cond_dim, device=device, dtype=dtype, operations=operations)
|
||||
for _ in range(cond_resblocks)
|
||||
])
|
||||
self.aspp = AnimaLLLiteASPP(cond_dim, aspp_dilations, device=device, dtype=dtype, operations=operations) if aspp_dilations else None
|
||||
self.proj = operations.Conv2d(cond_dim, cond_emb_dim, kernel_size=1, device=device, dtype=dtype)
|
||||
self.out_norm = operations.LayerNorm(cond_emb_dim, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.silu(self.norm1(self.conv1(x)))
|
||||
x = F.silu(self.norm2(self.conv2(x)))
|
||||
x = F.silu(self.norm3(self.conv3(x)))
|
||||
for block in self.resblocks:
|
||||
x = block(x)
|
||||
if self.aspp is not None:
|
||||
x = self.aspp(x)
|
||||
x = self.proj(x).flatten(2).transpose(1, 2).contiguous()
|
||||
return self.out_norm(x)
|
||||
|
||||
|
||||
class AnimaLLLiteModule(nn.Module):
|
||||
def __init__(self, in_dim, cond_emb_dim, mlp_dim, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
self.down = operations.Linear(in_dim, mlp_dim, device=device, dtype=dtype)
|
||||
self.mid = operations.Linear(mlp_dim + cond_emb_dim, mlp_dim, device=device, dtype=dtype)
|
||||
self.cond_to_film = operations.Linear(cond_emb_dim, 2 * mlp_dim, device=device, dtype=dtype)
|
||||
self.up = operations.Linear(mlp_dim, in_dim, device=device, dtype=dtype)
|
||||
self.depth_embed = nn.Parameter(torch.empty(cond_emb_dim, device=device, dtype=dtype), requires_grad=False)
|
||||
|
||||
def forward(self, x, cond_emb, strength):
|
||||
original_shape = x.shape
|
||||
if x.ndim == 5:
|
||||
x = x.flatten(1, 3)
|
||||
|
||||
if x.shape[0] != cond_emb.shape[0]:
|
||||
if x.shape[0] % cond_emb.shape[0] != 0:
|
||||
raise ValueError(f"Anima LLLite batch mismatch: model input batch {x.shape[0]}, control batch {cond_emb.shape[0]}")
|
||||
cond_emb = cond_emb.repeat(x.shape[0] // cond_emb.shape[0], 1, 1)
|
||||
if x.shape[1] != cond_emb.shape[1]:
|
||||
raise ValueError(f"Anima LLLite sequence mismatch: model input has {x.shape[1]} tokens, control has {cond_emb.shape[1]}")
|
||||
|
||||
cond_local = cond_emb + comfy.ops.cast_to_input(self.depth_embed, cond_emb)
|
||||
hidden = F.silu(self.down(x))
|
||||
gamma, beta = self.cond_to_film(cond_local).chunk(2, dim=-1)
|
||||
hidden = self.mid(torch.cat((cond_local, hidden), dim=-1))
|
||||
hidden = F.silu(hidden * (1 + gamma) + beta)
|
||||
x = x + self.up(hidden) * strength
|
||||
|
||||
if len(original_shape) == 5:
|
||||
x = x.reshape(original_shape)
|
||||
return x
|
||||
|
||||
|
||||
class AnimaLLLite(nn.Module):
|
||||
def __init__(self, state_dict, metadata, device=None, dtype=None, operations=None):
|
||||
super().__init__()
|
||||
metadata = metadata or {}
|
||||
version = metadata.get("lllite.version", "2")
|
||||
if version != "2":
|
||||
raise ValueError(f"Unsupported Anima LLLite version {version!r}; only named-key v2 checkpoints are supported")
|
||||
|
||||
module_names = sorted({key.split(".", 1)[0] for key in state_dict if key.startswith("lllite_dit_blocks_")})
|
||||
if not module_names:
|
||||
raise ValueError("Anima LLLite checkpoint has no lllite_dit_blocks_* modules")
|
||||
|
||||
cond_in_channels = state_dict["lllite_conditioning1.conv1.weight"].shape[1]
|
||||
cond_dim = state_dict["lllite_conditioning1.conv3.weight"].shape[0]
|
||||
cond_emb_dim = state_dict["lllite_conditioning1.proj.weight"].shape[0]
|
||||
resblock_ids = {int(key.split(".")[2]) for key in state_dict if key.startswith("lllite_conditioning1.resblocks.")}
|
||||
cond_resblocks = max(resblock_ids) + 1 if resblock_ids else 0
|
||||
use_aspp = any(key.startswith("lllite_conditioning1.aspp.") for key in state_dict)
|
||||
dilation_string = metadata.get("lllite.aspp_dilations", "1,2,4,8")
|
||||
aspp_dilations = tuple(int(value) for value in dilation_string.split(",") if value.strip()) if use_aspp else ()
|
||||
|
||||
self.cond_in_channels = cond_in_channels
|
||||
self.inpaint_masked_input = metadata.get("lllite.inpaint_masked_input", "false").lower() == "true"
|
||||
self.lllite_conditioning1 = AnimaLLLiteConditioning(
|
||||
cond_in_channels, cond_dim, cond_emb_dim, cond_resblocks, aspp_dilations,
|
||||
device=device, dtype=dtype, operations=operations,
|
||||
)
|
||||
|
||||
self.module_names = set()
|
||||
self.block_count = 0
|
||||
self.model_dim = None
|
||||
for name in module_names:
|
||||
match = MODULE_PATTERN.fullmatch(name)
|
||||
if match is None:
|
||||
raise ValueError(f"Unsupported Anima LLLite module name: {name}")
|
||||
down_shape = state_dict[f"{name}.down.weight"].shape
|
||||
mlp_dim, in_dim = down_shape
|
||||
module_cond_dim = state_dict[f"{name}.cond_to_film.weight"].shape[1]
|
||||
if module_cond_dim != cond_emb_dim:
|
||||
raise ValueError(f"Anima LLLite conditioning dimension mismatch in {name}: {module_cond_dim} != {cond_emb_dim}")
|
||||
if self.model_dim is None:
|
||||
self.model_dim = in_dim
|
||||
elif self.model_dim != in_dim:
|
||||
raise ValueError(f"Anima LLLite model dimension mismatch in {name}: {in_dim} != {self.model_dim}")
|
||||
self.add_module(name, AnimaLLLiteModule(in_dim, cond_emb_dim, mlp_dim, device=device, dtype=dtype, operations=operations))
|
||||
self.module_names.add(name)
|
||||
self.block_count = max(self.block_count, int(match.group(1)) + 1)
|
||||
|
||||
def encode_conditioning(self, image):
|
||||
return self.lllite_conditioning1(image)
|
||||
|
||||
def apply(self, x, cond_emb, block_index, target, strength):
|
||||
name = f"lllite_dit_blocks_{block_index}_{target}"
|
||||
if name not in self.module_names:
|
||||
return x
|
||||
return self.get_submodule(name)(x, cond_emb, strength)
|
||||
|
||||
|
||||
class AnimaLLLitePatch:
|
||||
def __init__(self, model_patch, image, mask, strength, sigma_start, sigma_end):
|
||||
self.model_patch = model_patch
|
||||
self.image = image
|
||||
self.mask = mask
|
||||
self.strength = strength
|
||||
self.sigma_start = sigma_start
|
||||
self.sigma_end = sigma_end
|
||||
|
||||
def __call__(self, args):
|
||||
x = args["x"]
|
||||
transformer_options = args["transformer_options"]
|
||||
if self.strength == 0.0:
|
||||
return args
|
||||
sigmas = transformer_options.get("sigmas")
|
||||
if sigmas is not None:
|
||||
sigma = float(sigmas.max().item())
|
||||
if not self.sigma_end <= sigma <= self.sigma_start:
|
||||
return args
|
||||
if x.shape[2] != 1:
|
||||
raise ValueError(f"Anima LLLite only supports T=1, got T={x.shape[2]}")
|
||||
|
||||
target_height = x.shape[-2] * 8
|
||||
target_width = x.shape[-1] * 8
|
||||
image = comfy.utils.common_upscale(
|
||||
self.image.movedim(-1, 1), target_width, target_height, "bicubic", crop="center"
|
||||
).clamp(0.0, 1.0)
|
||||
image = image.to(device=x.device, dtype=x.dtype) * 2.0 - 1.0
|
||||
|
||||
if self.model_patch.model.cond_in_channels == 4:
|
||||
mask = self.mask
|
||||
if mask.ndim == 3:
|
||||
mask = mask.unsqueeze(1)
|
||||
if mask.ndim != 4 or mask.shape[1] != 1:
|
||||
raise ValueError(f"Anima LLLite mask must have one channel, got shape {tuple(mask.shape)}")
|
||||
mask = comfy.utils.common_upscale(
|
||||
mask.float(), target_width, target_height, "nearest-exact", crop="center"
|
||||
)
|
||||
if mask.shape[0] != image.shape[0]:
|
||||
if image.shape[0] % mask.shape[0] != 0:
|
||||
raise ValueError(
|
||||
f"Anima LLLite mask batch {mask.shape[0]} cannot be broadcast to image batch {image.shape[0]}"
|
||||
)
|
||||
mask = mask.repeat(image.shape[0] // mask.shape[0], 1, 1, 1)
|
||||
mask = (mask >= 0.5).to(device=x.device, dtype=x.dtype)
|
||||
if self.model_patch.model.inpaint_masked_input:
|
||||
image = image * (mask < 0.5).to(image.dtype)
|
||||
image = torch.cat((image, mask * 2.0 - 1.0), dim=1)
|
||||
|
||||
cond_emb = self.model_patch.model.encode_conditioning(image)
|
||||
transformer_options["model_patch_data"][self] = cond_emb
|
||||
return args
|
||||
|
||||
def to(self, device_or_dtype):
|
||||
return self
|
||||
|
||||
def models(self):
|
||||
return [self.model_patch]
|
||||
|
||||
|
||||
class AnimaLLLiteAttentionPatch:
|
||||
def __init__(self, patch, targets):
|
||||
self.patch = patch
|
||||
self.targets = targets
|
||||
|
||||
def __call__(self, q, k, v, pe=None, attn_mask=None, extra_options=None):
|
||||
cond_emb = extra_options["model_patch_data"].get(self.patch)
|
||||
if cond_emb is None:
|
||||
return {"q": q, "k": k, "v": v, "pe": pe, "attn_mask": attn_mask}
|
||||
|
||||
block_index = extra_options["block_index"]
|
||||
values = {"q": q, "k": k, "v": v}
|
||||
for value_name, target in self.targets.items():
|
||||
values[value_name] = self.patch.model_patch.model.apply(
|
||||
values[value_name], cond_emb, block_index, target, self.patch.strength
|
||||
)
|
||||
|
||||
return {"q": values["q"], "k": values["k"], "v": values["v"], "pe": pe, "attn_mask": attn_mask}
|
||||
|
||||
|
||||
class AnimaLLLiteMLPPatch:
|
||||
def __init__(self, patch):
|
||||
self.patch = patch
|
||||
|
||||
def __call__(self, args):
|
||||
cond_emb = args["transformer_options"]["model_patch_data"].get(self.patch)
|
||||
if cond_emb is None:
|
||||
return args
|
||||
args["x"] = self.patch.model_patch.model.apply(
|
||||
args["x"], cond_emb, args["transformer_options"]["block_index"], "mlp_layer1", self.patch.strength
|
||||
)
|
||||
return args
|
||||
@@ -14,7 +14,6 @@ from torchvision import transforms
|
||||
import comfy.patcher_extension
|
||||
from comfy.ldm.modules.attention import optimized_attention
|
||||
import comfy.ldm.common_dit
|
||||
import comfy.ops
|
||||
import comfy.quant_ops
|
||||
|
||||
|
||||
@@ -149,29 +148,11 @@ class Attention(nn.Module):
|
||||
x: torch.Tensor,
|
||||
context: Optional[torch.Tensor] = None,
|
||||
rope_emb: Optional[torch.Tensor] = None,
|
||||
transformer_options: Optional[dict] = {},
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
q = self.q_proj(x)
|
||||
context = x if context is None else context
|
||||
q_input = x
|
||||
k_input = context
|
||||
v_input = context
|
||||
|
||||
transformer_patches = transformer_options.get("patches", {})
|
||||
patch_name = "attn1_patch" if self.is_selfattn else "attn2_patch"
|
||||
if patch_name in transformer_patches:
|
||||
extra_options = transformer_options.copy()
|
||||
extra_options["n_heads"] = self.n_heads
|
||||
extra_options["dim_head"] = self.head_dim
|
||||
for patch in transformer_patches[patch_name]:
|
||||
out = patch(q_input, k_input, v_input, pe=rope_emb, attn_mask=None, extra_options=extra_options)
|
||||
q_input = out.get("q", q_input)
|
||||
k_input = out.get("k", k_input)
|
||||
v_input = out.get("v", v_input)
|
||||
rope_emb = out.get("pe", rope_emb)
|
||||
|
||||
q = self.q_proj(q_input)
|
||||
k = self.k_proj(k_input)
|
||||
v = self.v_proj(v_input)
|
||||
k = self.k_proj(context)
|
||||
v = self.v_proj(context)
|
||||
q, k, v = map(
|
||||
lambda t: rearrange(t, "b ... (h d) -> b ... h d", h=self.n_heads, d=self.head_dim),
|
||||
(q, k, v),
|
||||
@@ -180,16 +161,11 @@ class Attention(nn.Module):
|
||||
def apply_norm_and_rotary_pos_emb(
|
||||
q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, rope_emb: Optional[torch.Tensor]
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
v = self.v_norm(v)
|
||||
if self.is_selfattn and rope_emb is not None: # only apply to self-attention!
|
||||
q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(self.q_norm, q, offloadable=True)
|
||||
k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(self.k_norm, k, offloadable=True)
|
||||
q, k = comfy.quant_ops.ck.rms_rope_split_half(q, k, rope_emb, q_scale, k_scale, self.q_norm.eps)
|
||||
comfy.ops.uncast_bias_weight(self.q_norm, q_scale, None, q_offload_stream)
|
||||
comfy.ops.uncast_bias_weight(self.k_norm, k_scale, None, k_offload_stream)
|
||||
else:
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
q, k = comfy.quant_ops.ck.apply_rope_split_half(q, k, rope_emb)
|
||||
return q, k, v
|
||||
|
||||
q, k, v = apply_norm_and_rotary_pos_emb(q, k, v, rope_emb)
|
||||
@@ -212,7 +188,7 @@ class Attention(nn.Module):
|
||||
x (Tensor): The query tensor of shape [B, Mq, K]
|
||||
context (Optional[Tensor]): The key tensor of shape [B, Mk, K] or use x as context [self attention] if None
|
||||
"""
|
||||
q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb, transformer_options=transformer_options)
|
||||
q, k, v = self.compute_qkv(x, context, rope_emb=rope_emb)
|
||||
return self.compute_attention(q, k, v, transformer_options=transformer_options)
|
||||
|
||||
|
||||
@@ -579,14 +555,8 @@ class Block(nn.Module):
|
||||
self.layer_norm_mlp,
|
||||
scale_mlp_B_T_1_1_D,
|
||||
shift_mlp_B_T_1_1_D,
|
||||
).to(compute_dtype)
|
||||
patches = transformer_options.get("patches", {})
|
||||
if "mlp_patch" in patches:
|
||||
args = {"x": normalized_x_B_T_H_W_D, "transformer_options": transformer_options}
|
||||
for patch in patches["mlp_patch"]:
|
||||
args = patch(args)
|
||||
normalized_x_B_T_H_W_D = args["x"]
|
||||
result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D)
|
||||
)
|
||||
result_B_T_H_W_D = self.mlp(normalized_x_B_T_H_W_D.to(compute_dtype))
|
||||
x_B_T_H_W_D = torch.addcmul(x_B_T_H_W_D, gate_mlp_B_T_1_1_D.to(residual_dtype), result_B_T_H_W_D.to(residual_dtype))
|
||||
return x_B_T_H_W_D
|
||||
|
||||
@@ -893,22 +863,11 @@ class MiniTrainDIT(nn.Module):
|
||||
x_B_T_H_W_D.shape == extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape
|
||||
), f"{x_B_T_H_W_D.shape} != {extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D.shape}"
|
||||
|
||||
transformer_options = kwargs.get("transformer_options", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
if "post_input" in patches:
|
||||
transformer_options = transformer_options.copy()
|
||||
transformer_options["model_patch_data"] = {}
|
||||
|
||||
if "post_input" in patches:
|
||||
for patch in patches["post_input"]:
|
||||
out = patch({"img": x_B_T_H_W_D, "x": x_B_C_T_H_W, "transformer_options": transformer_options})
|
||||
x_B_T_H_W_D = out["img"]
|
||||
|
||||
block_kwargs = {
|
||||
"rope_emb_L_1_1_D": rope_emb_L_1_1_D.unsqueeze(1).unsqueeze(0),
|
||||
"adaln_lora_B_T_3D": adaln_lora_B_T_3D,
|
||||
"extra_per_block_pos_emb": extra_pos_emb_B_T_H_W_D_or_T_H_W_B_D,
|
||||
"transformer_options": transformer_options,
|
||||
"transformer_options": kwargs.get("transformer_options", {}),
|
||||
}
|
||||
|
||||
# The residual stream for this model has large values. To make fp16 compute_dtype work, we keep the residual stream
|
||||
@@ -918,8 +877,7 @@ class MiniTrainDIT(nn.Module):
|
||||
if x_B_T_H_W_D.dtype == torch.float16:
|
||||
x_B_T_H_W_D = x_B_T_H_W_D.float()
|
||||
|
||||
for block_index, block in enumerate(self.blocks):
|
||||
transformer_options["block_index"] = block_index
|
||||
for block in self.blocks:
|
||||
x_B_T_H_W_D = block(
|
||||
x_B_T_H_W_D,
|
||||
t_embedding_B_T_D,
|
||||
|
||||
@@ -15,24 +15,24 @@ def make_two_pass_attention(ar_len: int, transformer_options=None):
|
||||
The AR pass goes through SDPA directand bypasses wrappers, it is only ~1% of T at typical edit sizes.
|
||||
"""
|
||||
|
||||
def two_pass_attention(q, k, v, heads, enable_gqa=False, **kwargs):
|
||||
def two_pass_attention(q, k, v, heads, **kwargs):
|
||||
B, H, T, D = q.shape
|
||||
|
||||
if T < k.shape[2]: # KV-cache hot path: Q is shorter than K/V (cached AR prefix is in K/V only), all fresh Q positions are in the gen region, single full-attention call
|
||||
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa)
|
||||
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options)
|
||||
elif ar_len >= T:
|
||||
out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa)
|
||||
out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True)
|
||||
elif ar_len <= 0:
|
||||
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa)
|
||||
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options)
|
||||
else:
|
||||
out_ar = comfy.ops.scaled_dot_product_attention(
|
||||
q[:, :, :ar_len], k[:, :, :ar_len], v[:, :, :ar_len],
|
||||
attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa,
|
||||
attn_mask=None, dropout_p=0.0, is_causal=True,
|
||||
)
|
||||
out_gen = optimized_attention(
|
||||
q[:, :, ar_len:], k, v, heads,
|
||||
mask=None, skip_reshape=True, skip_output_reshape=True,
|
||||
transformer_options=transformer_options, enable_gqa=enable_gqa,
|
||||
transformer_options=transformer_options,
|
||||
)
|
||||
out = torch.cat([out_ar, out_gen], dim=2)
|
||||
|
||||
|
||||
@@ -1,445 +0,0 @@
|
||||
# https://github.com/jdopensource/JoyAI-Image-Edit (Apache 2.0)
|
||||
import math
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import comfy_kitchen
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
import comfy.ldm.common_dit
|
||||
import comfy.ops
|
||||
import comfy.patcher_extension
|
||||
from comfy.ldm.lightricks.model import GELU_approx, PixArtAlphaTextProjection, TimestepEmbedding, Timesteps
|
||||
from comfy.ldm.modules.attention import optimized_attention
|
||||
|
||||
|
||||
class JoyImageModulate(nn.Module):
|
||||
def __init__(self, hidden_size: int, factor: int, dtype=None, device=None):
|
||||
super().__init__()
|
||||
self.factor = factor
|
||||
self.modulate_table = nn.Parameter(
|
||||
torch.empty(1, factor, hidden_size, dtype=dtype, device=device)
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> list:
|
||||
if x.ndim != 3:
|
||||
x = x.unsqueeze(1)
|
||||
table = comfy.ops.cast_to_input(self.modulate_table, x)
|
||||
return [o.squeeze(1) for o in (table + x).chunk(self.factor, dim=1)]
|
||||
|
||||
|
||||
class JoyImageFeedForward(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
inner_dim: int,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.net = nn.ModuleList([
|
||||
GELU_approx(dim, inner_dim, dtype=dtype, device=device, operations=operations),
|
||||
nn.Identity(),
|
||||
operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device),
|
||||
])
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
for module in self.net:
|
||||
x = module(x)
|
||||
return x
|
||||
|
||||
|
||||
class JoyImageAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
eps: float = 1e-6,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.img_attn_qkv = operations.Linear(dim, inner_dim * 3, bias=True, dtype=dtype, device=device)
|
||||
self.img_attn_q_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
|
||||
self.img_attn_k_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
|
||||
self.img_attn_proj = operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device)
|
||||
|
||||
self.txt_attn_qkv = operations.Linear(dim, inner_dim * 3, bias=True, dtype=dtype, device=device)
|
||||
self.txt_attn_q_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
|
||||
self.txt_attn_k_norm = operations.RMSNorm(attention_head_dim, eps=eps, dtype=dtype, device=device)
|
||||
self.txt_attn_proj = operations.Linear(inner_dim, dim, bias=True, dtype=dtype, device=device)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
img: torch.Tensor,
|
||||
txt: torch.Tensor,
|
||||
image_rotary_emb: torch.Tensor,
|
||||
transformer_options=None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
heads = self.num_attention_heads
|
||||
|
||||
img_q, img_k, img_v = self.img_attn_qkv(img).chunk(3, dim=-1)
|
||||
txt_q, txt_k, txt_v = self.txt_attn_qkv(txt).chunk(3, dim=-1)
|
||||
|
||||
img_q = img_q.unflatten(-1, (heads, -1))
|
||||
img_k = img_k.unflatten(-1, (heads, -1))
|
||||
img_v = img_v.unflatten(-1, (heads, -1))
|
||||
txt_q = txt_q.unflatten(-1, (heads, -1))
|
||||
txt_k = txt_k.unflatten(-1, (heads, -1))
|
||||
txt_v = txt_v.unflatten(-1, (heads, -1))
|
||||
|
||||
img_q = self.img_attn_q_norm(img_q)
|
||||
img_k = self.img_attn_k_norm(img_k)
|
||||
txt_q = self.txt_attn_q_norm(txt_q)
|
||||
txt_k = self.txt_attn_k_norm(txt_k)
|
||||
|
||||
img_q, img_k = comfy_kitchen.apply_rope(img_q, img_k, image_rotary_emb)
|
||||
|
||||
joint_q = torch.cat([img_q, txt_q], dim=1)
|
||||
joint_k = torch.cat([img_k, txt_k], dim=1)
|
||||
joint_v = torch.cat([img_v, txt_v], dim=1)
|
||||
|
||||
joint_q = joint_q.flatten(2, 3)
|
||||
joint_k = joint_k.flatten(2, 3)
|
||||
joint_v = joint_v.flatten(2, 3)
|
||||
|
||||
joint_out = optimized_attention(joint_q, joint_k, joint_v, heads=heads, transformer_options=transformer_options)
|
||||
|
||||
seq_img = img.shape[1]
|
||||
img_out = joint_out[:, :seq_img, :]
|
||||
txt_out = joint_out[:, seq_img:, :]
|
||||
|
||||
img_out = self.img_attn_proj(img_out)
|
||||
txt_out = self.txt_attn_proj(txt_out)
|
||||
return img_out, txt_out
|
||||
|
||||
|
||||
class JoyImageTransformerBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
mlp_width_ratio: float = 4.0,
|
||||
eps: float = 1e-6,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
mlp_hidden_dim = int(dim * mlp_width_ratio)
|
||||
|
||||
self.img_mod = JoyImageModulate(dim, factor=6, dtype=dtype, device=device)
|
||||
self.img_norm1 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
|
||||
self.img_norm2 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
|
||||
self.img_mlp = JoyImageFeedForward(dim, inner_dim=mlp_hidden_dim, dtype=dtype, device=device, operations=operations)
|
||||
|
||||
self.txt_mod = JoyImageModulate(dim, factor=6, dtype=dtype, device=device)
|
||||
self.txt_norm1 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
|
||||
self.txt_norm2 = operations.LayerNorm(dim, elementwise_affine=False, eps=eps, dtype=dtype, device=device)
|
||||
self.txt_mlp = JoyImageFeedForward(dim, inner_dim=mlp_hidden_dim, dtype=dtype, device=device, operations=operations)
|
||||
|
||||
self.attn = JoyImageAttention(
|
||||
dim=dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
eps=eps,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
image_rotary_emb: torch.Tensor,
|
||||
transformer_options=None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
(
|
||||
img_mod1_shift,
|
||||
img_mod1_scale,
|
||||
img_mod1_gate,
|
||||
img_mod2_shift,
|
||||
img_mod2_scale,
|
||||
img_mod2_gate,
|
||||
) = self.img_mod(temb)
|
||||
(
|
||||
txt_mod1_shift,
|
||||
txt_mod1_scale,
|
||||
txt_mod1_gate,
|
||||
txt_mod2_shift,
|
||||
txt_mod2_scale,
|
||||
txt_mod2_gate,
|
||||
) = self.txt_mod(temb)
|
||||
|
||||
img_normed = self.img_norm1(hidden_states)
|
||||
txt_normed = self.txt_norm1(encoder_hidden_states)
|
||||
img_modulated = img_normed * (1 + img_mod1_scale.unsqueeze(1)) + img_mod1_shift.unsqueeze(1)
|
||||
txt_modulated = txt_normed * (1 + txt_mod1_scale.unsqueeze(1)) + txt_mod1_shift.unsqueeze(1)
|
||||
|
||||
img_attn, txt_attn = self.attn(img_modulated, txt_modulated, image_rotary_emb, transformer_options=transformer_options)
|
||||
|
||||
hidden_states = hidden_states + img_attn * img_mod1_gate.unsqueeze(1)
|
||||
encoder_hidden_states = encoder_hidden_states + txt_attn * txt_mod1_gate.unsqueeze(1)
|
||||
|
||||
img_ffn_normed = self.img_norm2(hidden_states)
|
||||
txt_ffn_normed = self.txt_norm2(encoder_hidden_states)
|
||||
img_ffn_input = img_ffn_normed * (1 + img_mod2_scale.unsqueeze(1)) + img_mod2_shift.unsqueeze(1)
|
||||
txt_ffn_input = txt_ffn_normed * (1 + txt_mod2_scale.unsqueeze(1)) + txt_mod2_shift.unsqueeze(1)
|
||||
hidden_states = hidden_states + self.img_mlp(img_ffn_input) * img_mod2_gate.unsqueeze(1)
|
||||
encoder_hidden_states = encoder_hidden_states + self.txt_mlp(txt_ffn_input) * txt_mod2_gate.unsqueeze(1)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
class JoyImageTimeTextImageEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
time_freq_dim: int,
|
||||
time_proj_dim: int,
|
||||
text_embed_dim: int,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0)
|
||||
self.time_embedder = TimestepEmbedding(
|
||||
in_channels=time_freq_dim,
|
||||
time_embed_dim=dim,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations,
|
||||
)
|
||||
self.act_fn = nn.SiLU()
|
||||
self.time_proj = operations.Linear(dim, time_proj_dim, bias=True, dtype=dtype, device=device)
|
||||
self.text_embedder = PixArtAlphaTextProjection(
|
||||
text_embed_dim, dim, act_fn="gelu_tanh", dtype=dtype, device=device, operations=operations,
|
||||
)
|
||||
|
||||
def forward(self, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor):
|
||||
timestep = self.timesteps_proj(timestep)
|
||||
temb = self.time_embedder(timestep.to(dtype=encoder_hidden_states.dtype)).type_as(encoder_hidden_states)
|
||||
timestep_proj = self.time_proj(self.act_fn(temb))
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
return temb, timestep_proj, encoder_hidden_states
|
||||
|
||||
|
||||
class JoyImageTransformer3DModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: list = [1, 2, 2],
|
||||
in_channels: int = 16,
|
||||
out_channels: Optional[int] = None,
|
||||
hidden_size: int = 3072,
|
||||
num_attention_heads: int = 24,
|
||||
text_dim: int = 4096,
|
||||
mlp_width_ratio: float = 4.0,
|
||||
num_layers: int = 20,
|
||||
rope_dim_list: list = [16, 56, 56],
|
||||
theta: int = 256,
|
||||
image_model=None,
|
||||
dtype=None,
|
||||
device=None,
|
||||
operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dtype = dtype
|
||||
self.out_channels = out_channels or in_channels
|
||||
self.patch_size = list(patch_size)
|
||||
self.rope_dim_list = list(rope_dim_list)
|
||||
self.theta = theta
|
||||
|
||||
attention_head_dim = hidden_size // num_attention_heads
|
||||
|
||||
self.img_in = operations.Conv3d(
|
||||
in_channels,
|
||||
hidden_size,
|
||||
kernel_size=tuple(self.patch_size),
|
||||
stride=tuple(self.patch_size),
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
|
||||
self.condition_embedder = JoyImageTimeTextImageEmbedding(
|
||||
dim=hidden_size,
|
||||
time_freq_dim=256,
|
||||
time_proj_dim=hidden_size * 6,
|
||||
text_embed_dim=text_dim,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations,
|
||||
)
|
||||
|
||||
self.double_blocks = nn.ModuleList([
|
||||
JoyImageTransformerBlock(
|
||||
dim=hidden_size,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations,
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
])
|
||||
|
||||
self.norm_out = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device)
|
||||
self.proj_out = operations.Linear(
|
||||
hidden_size,
|
||||
self.out_channels * math.prod(self.patch_size),
|
||||
bias=True,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def _get_rotary_pos_embed_for_range(
|
||||
self,
|
||||
start: Tuple[int, int, int],
|
||||
stop: Tuple[int, int, int],
|
||||
device=None,
|
||||
) -> torch.Tensor:
|
||||
# 3D RoPE for the patch grid range [start, stop) over (t, h, w). Token order after
|
||||
# reshape(-1) is (t, h, w), matching the img_in Conv3d flatten.
|
||||
rope_dim_list = self.rope_dim_list
|
||||
|
||||
grids = [torch.arange(start[i], stop[i], dtype=torch.float32, device=device) for i in range(3)]
|
||||
mesh = torch.stack(torch.meshgrid(*grids, indexing="ij"), dim=0)
|
||||
|
||||
angles_parts = []
|
||||
for i, dim in enumerate(rope_dim_list):
|
||||
pos = mesh[i].reshape(-1)
|
||||
freqs = 1.0 / (self.theta ** (torch.arange(0, dim, 2, dtype=torch.float32, device=device)[: (dim // 2)] / dim))
|
||||
angles_parts.append(torch.outer(pos, freqs))
|
||||
|
||||
angles = torch.cat(angles_parts, dim=1)
|
||||
cos = angles.cos()
|
||||
sin = angles.sin()
|
||||
return torch.stack((cos, -sin, sin, cos), dim=-1).unflatten(-1, (2, 2))
|
||||
|
||||
def get_rotary_pos_embed_for_components(
|
||||
self,
|
||||
component_sizes,
|
||||
device=None,
|
||||
) -> torch.Tensor:
|
||||
# Per-component 3D RoPE. component_sizes is a list of (t, h, w) patch grid sizes in
|
||||
# sequence order [target, ref0, ref1, ...]; h/w restart at 0 for each component while t
|
||||
# continues from the running offset, giving every image its own temporal position band.
|
||||
freqs_parts = []
|
||||
t_offset = 0
|
||||
for (t, h, w) in component_sizes:
|
||||
freqs = self._get_rotary_pos_embed_for_range(
|
||||
start=(t_offset, 0, 0),
|
||||
stop=(t_offset + t, h, w),
|
||||
device=device,
|
||||
)
|
||||
freqs_parts.append(freqs)
|
||||
t_offset += t
|
||||
return torch.cat(freqs_parts, dim=0).unsqueeze(0).unsqueeze(2)
|
||||
|
||||
def unpatchify(self, x: torch.Tensor, t: int, h: int, w: int) -> torch.Tensor:
|
||||
c = self.out_channels
|
||||
pt, ph, pw = self.patch_size
|
||||
x = x.reshape(x.shape[0], t, h, w, pt, ph, pw, c)
|
||||
x = x.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
return x.reshape(x.shape[0], c, t * pt, h * ph, w * pw)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
context: torch.Tensor = None,
|
||||
ref_latents=None,
|
||||
control=None,
|
||||
transformer_options=None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
transformer_options = {} if transformer_options is None else transformer_options.copy()
|
||||
return comfy.patcher_extension.WrapperExecutor.new_class_executor(
|
||||
self._forward,
|
||||
self,
|
||||
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options)
|
||||
).execute(hidden_states, timestep, context, ref_latents, transformer_options, **kwargs)
|
||||
|
||||
def _forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
context: torch.Tensor,
|
||||
ref_latents=None,
|
||||
transformer_options=None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
pt, ph, pw = self.patch_size
|
||||
_, _, ot, oh, ow = hidden_states.shape
|
||||
|
||||
components = [hidden_states, *(ref_latents or [])]
|
||||
component_sizes = []
|
||||
img_tokens = []
|
||||
for comp in components:
|
||||
comp = comfy.ldm.common_dit.pad_to_patch_size(comp, self.patch_size)
|
||||
_, _, ct, ch, cw = comp.shape
|
||||
component_sizes.append((ct // pt, ch // ph, cw // pw))
|
||||
tokens = self.img_in(comp).flatten(2).transpose(1, 2) # (B, n_i, D)
|
||||
img_tokens.append(tokens)
|
||||
|
||||
img = torch.cat(img_tokens, dim=1)
|
||||
|
||||
_, vec, txt = self.condition_embedder(timestep, context)
|
||||
vec = vec.unflatten(1, (6, -1))
|
||||
|
||||
image_rotary_emb = self.get_rotary_pos_embed_for_components(
|
||||
component_sizes,
|
||||
device=hidden_states.device,
|
||||
)
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.double_blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
for i, block in enumerate(self.double_blocks):
|
||||
transformer_options["block_index"] = i
|
||||
if ("double_block", i) in blocks_replace:
|
||||
def block_wrap(args):
|
||||
out = {}
|
||||
out["img"], out["txt"] = block(
|
||||
hidden_states=args["img"],
|
||||
encoder_hidden_states=args["txt"],
|
||||
temb=args["vec"],
|
||||
image_rotary_emb=args["pe"],
|
||||
transformer_options=args.get("transformer_options"),
|
||||
)
|
||||
return out
|
||||
|
||||
out = blocks_replace[("double_block", i)]({"img": img,
|
||||
"txt": txt,
|
||||
"vec": vec,
|
||||
"pe": image_rotary_emb,
|
||||
"transformer_options": transformer_options},
|
||||
{"original_block": block_wrap})
|
||||
txt = out["txt"]
|
||||
img = out["img"]
|
||||
else:
|
||||
img, txt = block(
|
||||
hidden_states=img,
|
||||
encoder_hidden_states=txt,
|
||||
temb=vec,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
transformer_options=transformer_options,
|
||||
)
|
||||
|
||||
tt, th, tw = component_sizes[0]
|
||||
target_tokens = tt * th * tw
|
||||
img = img[:, :target_tokens, :]
|
||||
img = self.proj_out(self.norm_out(img))
|
||||
img = self.unpatchify(img, tt, th, tw)
|
||||
return img[:, :, :ot, :oh, :ow]
|
||||
+23
-124
@@ -15,7 +15,6 @@ from einops import rearrange
|
||||
import comfy.model_management
|
||||
import comfy.patcher_extension
|
||||
import comfy.ldm.common_dit
|
||||
import comfy.utils
|
||||
from comfy.ldm.flux.layers import EmbedND, timestep_embedding
|
||||
from comfy.ldm.flux.math import apply_rope
|
||||
from comfy.ldm.modules.attention import optimized_attention_masked
|
||||
@@ -74,20 +73,11 @@ class Attention(nn.Module):
|
||||
self.wo = operations.Linear(dim, dim, bias=bias, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x, freqs=None, mask=None, transformer_options={}):
|
||||
transformer_patches = transformer_options.get("patches", {})
|
||||
extra_options = transformer_options.copy()
|
||||
q, k, v, gate = self.wq(x), self.wk(x), self.wv(x), self.gate(x)
|
||||
q = rearrange(q, "B L (H D) -> B H L D", H=self.heads)
|
||||
k = rearrange(k, "B L (H D) -> B H L D", H=self.kvheads)
|
||||
v = rearrange(v, "B L (H D) -> B H L D", H=self.kvheads)
|
||||
q, k = self.qknorm(q, k)
|
||||
|
||||
if "block_index" in transformer_options and "attn1_patch" in transformer_patches:
|
||||
for p in transformer_patches["attn1_patch"]:
|
||||
out = p(q, k, v, pe=freqs, attn_mask=mask, extra_options=extra_options)
|
||||
q, k, v = out.get("q", q), out.get("k", k), out.get("v", v)
|
||||
freqs, mask = out.get("pe", freqs), out.get("attn_mask", mask)
|
||||
|
||||
if freqs is not None:
|
||||
q, k = apply_rope(q, k, freqs)
|
||||
if self.kvheads != self.heads:
|
||||
@@ -96,11 +86,6 @@ class Attention(nn.Module):
|
||||
v = v.repeat_interleave(rep, dim=1)
|
||||
out = optimized_attention_masked(q, k, v, self.heads, mask=mask, skip_reshape=True,
|
||||
transformer_options=transformer_options)
|
||||
|
||||
if "block_index" in transformer_options and "attn1_output_patch" in transformer_patches:
|
||||
for p in transformer_patches["attn1_output_patch"]:
|
||||
out = p(out, extra_options)
|
||||
|
||||
return self.wo(out * F.sigmoid(gate))
|
||||
|
||||
|
||||
@@ -173,44 +158,8 @@ class SingleStreamBlock(nn.Module):
|
||||
self.attn = Attention(features, heads, kvheads=kvheads, bias=bias, device=device, dtype=dtype, operations=operations)
|
||||
self.mlp = SwiGLU(features, multiplier, bias, device=device, dtype=dtype, operations=operations)
|
||||
|
||||
def forward(self, x, vec, freqs, mask=None, timestep_zero_index=None, transformer_options={}):
|
||||
def forward(self, x, vec, freqs, mask=None, transformer_options={}):
|
||||
prescale, preshift, pregate, postscale, postshift, postgate = self.mod(vec)
|
||||
if timestep_zero_index is not None:
|
||||
bs = x.shape[0]
|
||||
ref_prescale = prescale[bs:]
|
||||
ref_preshift = preshift[bs:]
|
||||
ref_pregate = pregate[bs:]
|
||||
ref_postscale = postscale[bs:]
|
||||
ref_postshift = postshift[bs:]
|
||||
ref_postgate = postgate[bs:]
|
||||
prescale = prescale[:bs]
|
||||
preshift = preshift[:bs]
|
||||
pregate = pregate[:bs]
|
||||
postscale = postscale[:bs]
|
||||
postshift = postshift[:bs]
|
||||
postgate = postgate[:bs]
|
||||
|
||||
pre = self.prenorm(x)
|
||||
pre[:, :timestep_zero_index].mul_(1 + prescale).add_(preshift)
|
||||
pre[:, timestep_zero_index:].mul_(1 + ref_prescale).add_(ref_preshift)
|
||||
attn = self.attn(pre, freqs, mask, transformer_options=transformer_options)
|
||||
del pre
|
||||
attn[:, :timestep_zero_index].mul_(pregate)
|
||||
attn[:, timestep_zero_index:].mul_(ref_pregate)
|
||||
x = x + attn
|
||||
del attn
|
||||
|
||||
post = self.postnorm(x)
|
||||
post[:, :timestep_zero_index].mul_(1 + postscale).add_(postshift)
|
||||
post[:, timestep_zero_index:].mul_(1 + ref_postscale).add_(ref_postshift)
|
||||
mlp = self.mlp(post)
|
||||
del post
|
||||
mlp[:, :timestep_zero_index].mul_(postgate)
|
||||
mlp[:, timestep_zero_index:].mul_(ref_postgate)
|
||||
x = x + mlp
|
||||
del mlp
|
||||
return x
|
||||
|
||||
x = x + pregate * self.attn((1 + prescale) * self.prenorm(x) + preshift, freqs, mask, transformer_options=transformer_options)
|
||||
x = x + postgate * self.mlp((1 + postscale) * self.postnorm(x) + postshift)
|
||||
return x
|
||||
@@ -232,7 +181,7 @@ class LastLayer(nn.Module):
|
||||
class SingleStreamDiT(nn.Module):
|
||||
def __init__(self, features=6144, tdim=256, txtdim=2560, heads=48, kvheads=12, multiplier=4,
|
||||
layers=28, patch=2, channels=16, bias=False, theta=1e3, txtlayers=12,
|
||||
txtheads=20, txtkvheads=20, default_ref_method=None, image_model=None,
|
||||
txtheads=20, txtkvheads=20, image_model=None,
|
||||
device=None, dtype=None, operations=None, **kwargs):
|
||||
super().__init__()
|
||||
self.dtype = dtype
|
||||
@@ -242,7 +191,6 @@ class SingleStreamDiT(nn.Module):
|
||||
self.heads = heads
|
||||
self.txtdim = txtdim
|
||||
self.txtlayers = txtlayers
|
||||
self.default_ref_method = default_ref_method
|
||||
|
||||
headdim = features // heads
|
||||
axes = [headdim - 12 * (headdim // 16), 6 * (headdim // 16), 6 * (headdim // 16)]
|
||||
@@ -273,110 +221,61 @@ class SingleStreamDiT(nn.Module):
|
||||
operations.Linear(features, features * 6, device=device, dtype=dtype),
|
||||
)
|
||||
|
||||
def forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs):
|
||||
def forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs):
|
||||
return comfy.patcher_extension.WrapperExecutor.new_class_executor(
|
||||
self._forward,
|
||||
self,
|
||||
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options),
|
||||
).execute(x, timesteps, context, attention_mask, ref_latents, transformer_options, **kwargs)
|
||||
).execute(x, timesteps, context, attention_mask, transformer_options, **kwargs)
|
||||
|
||||
def process_img(self, x, index=0):
|
||||
patch = self.patch
|
||||
x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch))
|
||||
h, w = x.shape[-2] // patch, x.shape[-1] // patch
|
||||
img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch)
|
||||
|
||||
img_ids = torch.zeros(h, w, 3, device=x.device, dtype=torch.float32)
|
||||
img_ids[..., 0] = index
|
||||
img_ids[..., 1] = torch.arange(h, device=x.device, dtype=torch.float32)[:, None]
|
||||
img_ids[..., 2] = torch.arange(w, device=x.device, dtype=torch.float32)[None, :]
|
||||
return img, img_ids.reshape(1, h * w, 3).repeat(x.shape[0], 1, 1), h, w
|
||||
|
||||
def _forward(self, x, timesteps, context, attention_mask=None, ref_latents=None, transformer_options={}, **kwargs):
|
||||
transformer_options = transformer_options.copy()
|
||||
def _forward(self, x, timesteps, context, attention_mask=None, transformer_options={}, **kwargs):
|
||||
temporal = x.ndim == 5
|
||||
if temporal:
|
||||
b5, c5, t5, h5, w5 = x.shape
|
||||
x = x.reshape(b5 * t5, c5, h5, w5)
|
||||
bs, _, h_orig, w_orig = x.shape
|
||||
bs, c, H_orig, W_orig = x.shape
|
||||
patch = self.patch
|
||||
# Pad the latent up to a multiple of patch (as Flux/Lumina/QwenImage do); crop back at the end.
|
||||
x = comfy.ldm.common_dit.pad_to_patch_size(x, (patch, patch))
|
||||
H, W = x.shape[-2], x.shape[-1]
|
||||
h_, w_ = H // patch, W // patch
|
||||
|
||||
# context arrives as (B, seq, txtlayers*txtdim); reshape to (B, txtlayers, seq, txtdim).
|
||||
context = self._unpack_context(context)
|
||||
|
||||
img, imgpos, h_, w_ = self.process_img(x)
|
||||
img_tokens = img.shape[1]
|
||||
timestep_zero_index = None
|
||||
ref_method = kwargs.get("ref_latents_method", self.default_ref_method)
|
||||
if ref_method is not None and ref_latents is not None and len(ref_latents) > 0:
|
||||
ref_tokens = []
|
||||
ref_pos = []
|
||||
ref_num_tokens = []
|
||||
for index, ref in enumerate(ref_latents, 1):
|
||||
if ref.ndim == 5:
|
||||
rb, rc, rt, rh5, rw5 = ref.shape
|
||||
ref = ref.reshape(rb * rt, rc, rh5, rw5)
|
||||
ref = comfy.utils.repeat_to_batch_size(ref, bs)
|
||||
kontext, kontext_ids, _, _ = self.process_img(ref, index=index)
|
||||
ref_tokens.append(kontext)
|
||||
ref_pos.append(kontext_ids)
|
||||
ref_num_tokens.append(kontext.shape[1])
|
||||
img = torch.cat([img] + ref_tokens, dim=1)
|
||||
imgpos = torch.cat([imgpos] + ref_pos, dim=1)
|
||||
del ref_tokens, ref_pos
|
||||
if ref_method == "index_timestep_zero":
|
||||
timestep_zero_index = img_tokens
|
||||
transformer_options["reference_image_num_tokens"] = ref_num_tokens
|
||||
|
||||
img = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch)
|
||||
img = self.first(img)
|
||||
|
||||
t = self.tmlp(timestep_embedding(timesteps, self.tdim).unsqueeze(1).to(img.dtype))
|
||||
tvec = self.tproj(t)
|
||||
if timestep_zero_index is not None:
|
||||
t0 = self.tmlp(timestep_embedding(torch.zeros_like(timesteps), self.tdim).unsqueeze(1).to(img.dtype))
|
||||
tvec = torch.cat((tvec, self.tproj(t0)), dim=0)
|
||||
|
||||
context = self.txtfusion(context, mask=None, transformer_options=transformer_options)
|
||||
context = self.txtmlp(context)
|
||||
|
||||
txtlen = context.shape[1]
|
||||
device = context.device
|
||||
txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32)
|
||||
|
||||
patches = transformer_options.get("patches", {})
|
||||
if "post_input" in patches:
|
||||
for p in patches["post_input"]:
|
||||
out = p({"img": img, "txt": context, "img_ids": imgpos, "txt_ids": txtpos, "transformer_options": transformer_options})
|
||||
img, context = out["img"], out["txt"]
|
||||
imgpos, txtpos = out["img_ids"], out["txt_ids"]
|
||||
|
||||
txtlen, imglen = context.shape[1], img.shape[1]
|
||||
combined = torch.cat((context, img), dim=1)
|
||||
del context, img
|
||||
if timestep_zero_index is not None:
|
||||
timestep_zero_index += txtlen
|
||||
|
||||
# Position ids: text at 0, image at (0, h_idx, w_idx).
|
||||
device = combined.device
|
||||
txtpos = torch.zeros(bs, txtlen, 3, device=device, dtype=torch.float32)
|
||||
imgids = torch.zeros(h_, w_, 3, device=device, dtype=torch.float32)
|
||||
imgids[..., 1] = torch.arange(h_, device=device, dtype=torch.float32)[:, None]
|
||||
imgids[..., 2] = torch.arange(w_, device=device, dtype=torch.float32)[None, :]
|
||||
imgpos = imgids.reshape(1, h_ * w_, 3).repeat(bs, 1, 1)
|
||||
pos = torch.cat((txtpos, imgpos), dim=1)
|
||||
del txtpos, imgpos
|
||||
|
||||
freqs = self.pe_embedder(pos)
|
||||
del pos
|
||||
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "single"
|
||||
transformer_options["img_slice"] = [txtlen, combined.shape[1]]
|
||||
for i, block in enumerate(self.blocks):
|
||||
transformer_options["block_index"] = i
|
||||
combined = block(combined, tvec, freqs, None, timestep_zero_index=timestep_zero_index, transformer_options=transformer_options)
|
||||
for block in self.blocks:
|
||||
combined = block(combined, tvec, freqs, None, transformer_options=transformer_options)
|
||||
|
||||
final = self.last(combined, t)
|
||||
del combined
|
||||
out = final[:, txtlen:txtlen + img_tokens, :]
|
||||
out = final[:, txtlen:txtlen + imglen, :]
|
||||
out = rearrange(out, "b (h w) (c ph pw) -> b c (h ph) (w pw)",
|
||||
h=h_, w=w_, ph=patch, pw=patch, c=self.channels)
|
||||
out = out[:, :, :h_orig, :w_orig] # crop padding back off
|
||||
out = out[:, :, :H_orig, :W_orig] # crop padding back off
|
||||
if temporal:
|
||||
out = out.reshape(b5, t5, self.channels, h_orig, w_orig).movedim(1, 2)
|
||||
out = out.reshape(b5, t5, self.channels, H_orig, W_orig).movedim(1, 2)
|
||||
return out
|
||||
|
||||
def _unpack_context(self, context):
|
||||
|
||||
@@ -552,7 +552,6 @@ class WanModel(torch.nn.Module):
|
||||
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
# embeddings
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
grid_sizes = x.shape[2:]
|
||||
transformer_options["grid_sizes"] = grid_sizes
|
||||
@@ -565,13 +564,11 @@ class WanModel(torch.nn.Module):
|
||||
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
|
||||
|
||||
full_ref = None
|
||||
img_offset = 0
|
||||
if self.ref_conv is not None:
|
||||
full_ref = kwargs.get("reference_latent", None)
|
||||
if full_ref is not None:
|
||||
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
|
||||
x = torch.concat((full_ref, x), dim=1)
|
||||
img_offset = full_ref.shape[1]
|
||||
|
||||
# In-context reference (Bernini)
|
||||
context_latents = kwargs.get("context_latents", None)
|
||||
@@ -592,7 +589,6 @@ class WanModel(torch.nn.Module):
|
||||
context_img_len = clip_fea.shape[-2]
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
@@ -608,11 +604,6 @@ class WanModel(torch.nn.Module):
|
||||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
# head
|
||||
x = self.head(x, e)
|
||||
|
||||
@@ -786,7 +777,6 @@ class VaceWanModel(WanModel):
|
||||
**kwargs,
|
||||
):
|
||||
# embeddings
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
grid_sizes = x.shape[2:]
|
||||
transformer_options["grid_sizes"] = grid_sizes
|
||||
@@ -817,7 +807,6 @@ class VaceWanModel(WanModel):
|
||||
x_orig = x
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
@@ -833,11 +822,6 @@ class VaceWanModel(WanModel):
|
||||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
ii = self.vace_layers_mapping.get(i, None)
|
||||
if ii is not None:
|
||||
for iii in range(len(c)):
|
||||
@@ -903,7 +887,6 @@ class CameraWanModel(WanModel):
|
||||
**kwargs,
|
||||
):
|
||||
# embeddings
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
if self.control_adapter is not None and camera_conditions is not None:
|
||||
x = x + self.control_adapter(camera_conditions).to(x.dtype)
|
||||
@@ -926,7 +909,6 @@ class CameraWanModel(WanModel):
|
||||
context_img_len = clip_fea.shape[-2]
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
@@ -942,11 +924,6 @@ class CameraWanModel(WanModel):
|
||||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
# head
|
||||
x = self.head(x, e)
|
||||
|
||||
@@ -1358,7 +1335,6 @@ class WanModel_S2V(WanModel):
|
||||
|
||||
# embeddings
|
||||
bs, _, time, height, width = x.shape
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
if control_video is not None:
|
||||
x = x + self.cond_encoder(control_video)
|
||||
@@ -1403,7 +1379,6 @@ class WanModel_S2V(WanModel):
|
||||
context = self.text_embedding(context)
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
@@ -1418,12 +1393,6 @@ class WanModel_S2V(WanModel):
|
||||
x = out["img"]
|
||||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
if audio_emb is not None:
|
||||
x = self.audio_injector(x, i, audio_emb, audio_emb_global, seq_len)
|
||||
# head
|
||||
@@ -1630,7 +1599,6 @@ class HumoWanModel(WanModel):
|
||||
bs, _, time, height, width = x.shape
|
||||
|
||||
# embeddings
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
grid_sizes = x.shape[2:]
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
@@ -1662,7 +1630,6 @@ class HumoWanModel(WanModel):
|
||||
audio = None
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
@@ -1678,11 +1645,6 @@ class HumoWanModel(WanModel):
|
||||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, audio=audio, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": 0, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
# head
|
||||
x = self.head(x, e)
|
||||
|
||||
@@ -1698,14 +1660,8 @@ class SCAILWanModel(WanModel):
|
||||
|
||||
def forward_orig(self, x, t, context, clip_fea=None, freqs=None, transformer_options={}, pose_latents=None, reference_latent=None, ref_mask_latents=None, sam_latents=None, **kwargs):
|
||||
|
||||
x_input = x
|
||||
|
||||
img_offset = 0
|
||||
if reference_latent is not None:
|
||||
x = torch.cat((reference_latent, x), dim=2)
|
||||
img_offset = (reference_latent.shape[2] // self.patch_size[0]) * \
|
||||
(reference_latent.shape[3] // self.patch_size[1]) * \
|
||||
(reference_latent.shape[4] // self.patch_size[2])
|
||||
|
||||
# embeddings
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
@@ -1741,7 +1697,6 @@ class SCAILWanModel(WanModel):
|
||||
context_img_len = clip_fea.shape[-2]
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
@@ -1757,11 +1712,6 @@ class SCAILWanModel(WanModel):
|
||||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
# head
|
||||
x = self.head(x, e)
|
||||
|
||||
|
||||
@@ -493,7 +493,6 @@ class AnimateWanModel(WanModel):
|
||||
**kwargs,
|
||||
):
|
||||
# embeddings
|
||||
x_input = x
|
||||
x = self.patch_embedding(x.float()).to(x.dtype)
|
||||
x, motion_vec = self.after_patch_embedding(x, pose_latents, face_pixel_values)
|
||||
grid_sizes = x.shape[2:]
|
||||
@@ -506,13 +505,11 @@ class AnimateWanModel(WanModel):
|
||||
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
|
||||
|
||||
full_ref = None
|
||||
img_offset = 0
|
||||
if self.ref_conv is not None:
|
||||
full_ref = kwargs.get("reference_latent", None)
|
||||
if full_ref is not None:
|
||||
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
|
||||
x = torch.concat((full_ref, x), dim=1)
|
||||
img_offset = full_ref.shape[1]
|
||||
|
||||
# context
|
||||
context = self.text_embedding(context)
|
||||
@@ -525,7 +522,6 @@ class AnimateWanModel(WanModel):
|
||||
context_img_len = clip_fea.shape[-2]
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
@@ -541,11 +537,6 @@ class AnimateWanModel(WanModel):
|
||||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
if i % 5 == 0 and motion_vec is not None:
|
||||
x = x + self.face_adapter.fuser_blocks[i // 5](x, motion_vec)
|
||||
|
||||
|
||||
@@ -111,7 +111,6 @@ class WanDancerModel(WanModel):
|
||||
|
||||
def forward_orig(self, x, t, context, clip_fea=None, clip_fea_ref=None, freqs=None, audio_embed=None, fps=30, audio_inject_scale=1.0, transformer_options={}, **kwargs):
|
||||
# embeddings
|
||||
x_input = x
|
||||
if int(fps + 0.5) != 30:
|
||||
x = self.patch_embedding_global(x.float()).to(x.dtype)
|
||||
else:
|
||||
@@ -129,13 +128,11 @@ class WanDancerModel(WanModel):
|
||||
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
|
||||
|
||||
full_ref = None
|
||||
img_offset = 0
|
||||
if self.ref_conv is not None: # model has the weight, but this wasn't used in the original pipeline
|
||||
full_ref = kwargs.get("reference_latent", None)
|
||||
if full_ref is not None:
|
||||
full_ref = self.ref_conv(full_ref).flatten(2).transpose(1, 2)
|
||||
x = torch.concat((full_ref, x), dim=1)
|
||||
img_offset = full_ref.shape[1]
|
||||
|
||||
# context
|
||||
context = self.text_embedding(context)
|
||||
@@ -166,7 +163,6 @@ class WanDancerModel(WanModel):
|
||||
context_img_len += clip_fea_ref.shape[-2]
|
||||
|
||||
patches_replace = transformer_options.get("patches_replace", {})
|
||||
patches = transformer_options.get("patches", {})
|
||||
blocks_replace = patches_replace.get("dit", {})
|
||||
transformer_options["total_blocks"] = len(self.blocks)
|
||||
transformer_options["block_type"] = "double"
|
||||
@@ -181,12 +177,6 @@ class WanDancerModel(WanModel):
|
||||
x = out["img"]
|
||||
else:
|
||||
x = block(x, e=e0, freqs=freqs, context=context, context_img_len=context_img_len, transformer_options=transformer_options)
|
||||
|
||||
if "double_block" in patches:
|
||||
for p in patches["double_block"]:
|
||||
out = p({"img": x, "x": x_input, "vec": e, "block_index": i, "img_offset": img_offset, "transformer_options": transformer_options})
|
||||
x = out["img"]
|
||||
|
||||
if audio_emb is not None:
|
||||
x = self.music_injector(x, i, audio_emb, audio_emb_global=None, seq_len=seq_len, scale=audio_inject_scale)
|
||||
|
||||
|
||||
@@ -1,149 +0,0 @@
|
||||
# Uni3C controlnet for Wan 2.1: https://github.com/ewrfcas/Uni3C
|
||||
# Converted from the original diffusers based implementation.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.ldm.flux.layers import EmbedND
|
||||
from .model import WanSelfAttention
|
||||
|
||||
|
||||
class Uni3CLayerNormZero(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
conditioning_dim,
|
||||
embedding_dim,
|
||||
eps=1e-5,
|
||||
device=None, dtype=None, operations=None
|
||||
):
|
||||
super().__init__()
|
||||
self.silu = nn.SiLU()
|
||||
self.linear = operations.Linear(conditioning_dim, 3 * embedding_dim, device=device, dtype=dtype)
|
||||
self.norm = operations.LayerNorm(embedding_dim, eps=eps, elementwise_affine=True, device=device, dtype=dtype)
|
||||
|
||||
def forward(self, x, temb):
|
||||
shift, scale, gate = self.linear(self.silu(temb)).chunk(3, dim=1)
|
||||
x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :]
|
||||
return x, gate[:, None, :]
|
||||
|
||||
|
||||
class Uni3CAttentionBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
ffn_dim,
|
||||
num_heads,
|
||||
time_embed_dim=5120,
|
||||
eps=1e-6,
|
||||
device=None, dtype=None, operations=None
|
||||
):
|
||||
super().__init__()
|
||||
operation_settings = {"operations": operations, "device": device, "dtype": dtype}
|
||||
self.norm1 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
|
||||
self.self_attn = WanSelfAttention(dim, num_heads, qk_norm=True, eps=eps, operation_settings=operation_settings)
|
||||
self.norm2 = Uni3CLayerNormZero(time_embed_dim, dim, device=device, dtype=dtype, operations=operations)
|
||||
self.ffn = nn.Sequential(
|
||||
operations.Linear(dim, ffn_dim, device=device, dtype=dtype), nn.GELU(approximate='tanh'),
|
||||
operations.Linear(ffn_dim, dim, device=device, dtype=dtype))
|
||||
|
||||
def forward(self, x, temb, freqs):
|
||||
norm_x, gate_msa = self.norm1(x, temb)
|
||||
x = x + gate_msa * self.self_attn(norm_x, freqs)
|
||||
norm_x, gate_ff = self.norm2(x, temb)
|
||||
x = x + gate_ff * self.ffn(norm_x)
|
||||
return x
|
||||
|
||||
|
||||
class MaskCamEmbed(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
add_channels=7,
|
||||
mid_channels=256,
|
||||
conv_out_dim=5120,
|
||||
device=None, dtype=None, operations=None
|
||||
):
|
||||
super().__init__()
|
||||
self.mask_padding = [0, 0, 0, 0, 3, 0] # first frame conditioning
|
||||
self.mask_proj = nn.Sequential(
|
||||
operations.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8), device=device, dtype=dtype),
|
||||
operations.GroupNorm(mid_channels // 8, mid_channels, device=device, dtype=dtype),
|
||||
nn.SiLU())
|
||||
self.mask_zero_proj = operations.Conv3d(mid_channels, conv_out_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2), device=device, dtype=dtype)
|
||||
|
||||
def forward(self, add_inputs):
|
||||
add_padded = torch.nn.functional.pad(add_inputs, self.mask_padding, mode="constant", value=0)
|
||||
add_embeds = self.mask_proj(add_padded)
|
||||
add_embeds = self.mask_zero_proj(add_embeds)
|
||||
add_embeds = add_embeds.flatten(2).transpose(1, 2)
|
||||
return add_embeds
|
||||
|
||||
|
||||
class WanUni3CControlnet(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels=36,
|
||||
conv_out_dim=5120,
|
||||
dim=1024,
|
||||
ffn_dim=8192,
|
||||
num_heads=16,
|
||||
num_layers=20,
|
||||
time_embed_dim=5120,
|
||||
out_proj_dim=5120,
|
||||
add_channels=7,
|
||||
mid_channels=256,
|
||||
device=None, dtype=None, operations=None
|
||||
):
|
||||
super().__init__()
|
||||
patch_size = (1, 2, 2)
|
||||
self.num_layers = num_layers
|
||||
|
||||
self.controlnet_patch_embedding = operations.Conv3d(
|
||||
in_channels, conv_out_dim, kernel_size=patch_size, stride=patch_size, device=device, dtype=torch.float32)
|
||||
self.controlnet_mask_embedding = MaskCamEmbed(add_channels, mid_channels, conv_out_dim, device=device, dtype=dtype, operations=operations)
|
||||
|
||||
if conv_out_dim != dim:
|
||||
self.proj_in = operations.Linear(conv_out_dim, dim, device=device, dtype=dtype)
|
||||
else:
|
||||
self.proj_in = nn.Identity()
|
||||
|
||||
self.controlnet_blocks = nn.ModuleList([
|
||||
Uni3CAttentionBlock(dim, ffn_dim, num_heads, time_embed_dim, device=device, dtype=dtype, operations=operations)
|
||||
for _ in range(num_layers)])
|
||||
self.proj_out = nn.ModuleList([
|
||||
operations.Linear(dim, out_proj_dim, device=device, dtype=dtype)
|
||||
for _ in range(num_layers)])
|
||||
|
||||
head_dim = dim // num_heads
|
||||
self.rope_embedder = EmbedND(dim=head_dim, theta=10000.0, axes_dim=[head_dim - 4 * (head_dim // 6), 2 * (head_dim // 6), 2 * (head_dim // 6)])
|
||||
|
||||
def rope_encode(self, t_len, h_len, w_len, device=None, dtype=None):
|
||||
img_ids = torch.zeros((t_len, h_len, w_len, 3), device=device, dtype=dtype)
|
||||
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.arange(t_len, device=device, dtype=dtype).reshape(-1, 1, 1)
|
||||
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.arange(h_len, device=device, dtype=dtype).reshape(1, -1, 1)
|
||||
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.arange(w_len, device=device, dtype=dtype).reshape(1, 1, -1)
|
||||
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
|
||||
freqs = self.rope_embedder(img_ids).movedim(1, 2)
|
||||
return freqs
|
||||
|
||||
def process_input(self, control_input, render_mask=None, camera_embedding=None):
|
||||
# render_mask/camera_embedding are the checkpoint's extra conditioning path, not wired up yet
|
||||
hidden = self.controlnet_patch_embedding(control_input.float()).to(control_input.dtype)
|
||||
t_len, h_len, w_len = hidden.shape[2:]
|
||||
freqs = self.rope_encode(t_len, h_len, w_len, device=hidden.device, dtype=hidden.dtype)
|
||||
hidden = hidden.flatten(2).transpose(1, 2)
|
||||
|
||||
add_inputs = None
|
||||
if camera_embedding is not None and render_mask is not None:
|
||||
add_inputs = torch.cat([render_mask, camera_embedding], dim=1)
|
||||
elif render_mask is not None:
|
||||
add_inputs = render_mask
|
||||
|
||||
if add_inputs is not None:
|
||||
hidden = hidden + self.controlnet_mask_embedding(add_inputs.to(hidden.dtype))
|
||||
|
||||
hidden = self.proj_in(hidden)
|
||||
return hidden, freqs
|
||||
|
||||
def forward_block(self, block_index, hidden, temb, freqs):
|
||||
hidden = self.controlnet_blocks[block_index](hidden, temb, freqs)
|
||||
residual = self.proj_out[block_index](hidden)
|
||||
return hidden, residual
|
||||
+6
-44
@@ -58,7 +58,6 @@ import comfy.ldm.omnigen.omnigen2
|
||||
import comfy.ldm.seedvr.model
|
||||
import comfy.ldm.boogu.model
|
||||
import comfy.ldm.qwen_image.model
|
||||
import comfy.ldm.joyimage.model
|
||||
import comfy.ldm.ideogram4.model
|
||||
import comfy.ldm.krea2.model
|
||||
import comfy.ldm.kandinsky5.model
|
||||
@@ -2024,11 +2023,11 @@ class WAN22_WanDancer(WAN21):
|
||||
|
||||
fps = kwargs.get("fps", None)
|
||||
if fps is not None:
|
||||
out['fps'] = comfy.conds.CONDConstant(fps)
|
||||
out['fps'] = comfy.conds.CONDRegular(torch.FloatTensor([fps]))
|
||||
|
||||
audio_inject_scale = kwargs.get("audio_inject_scale", None)
|
||||
if audio_inject_scale is not None:
|
||||
out['audio_inject_scale'] = comfy.conds.CONDConstant(audio_inject_scale)
|
||||
out['audio_inject_scale'] = comfy.conds.CONDRegular(torch.FloatTensor([audio_inject_scale]))
|
||||
return out
|
||||
|
||||
class Hunyuan3Dv2(BaseModel):
|
||||
@@ -2227,7 +2226,10 @@ class Omnigen2(BaseModel):
|
||||
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
|
||||
ref_latents = kwargs.get("reference_latents", None)
|
||||
if ref_latents is not None:
|
||||
out['ref_latents'] = comfy.conds.CONDList([self.process_latent_in(lat) for lat in ref_latents])
|
||||
latents = []
|
||||
for lat in ref_latents:
|
||||
latents.append(self.process_latent_in(lat))
|
||||
out['ref_latents'] = comfy.conds.CONDList(latents)
|
||||
return out
|
||||
|
||||
def extra_conds_shapes(self, **kwargs):
|
||||
@@ -2274,28 +2276,6 @@ class QwenImage(BaseModel):
|
||||
out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16])
|
||||
return out
|
||||
|
||||
class JoyImage(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.joyimage.model.JoyImageTransformer3DModel)
|
||||
self.memory_usage_factor_conds = ("ref_latents",)
|
||||
|
||||
def extra_conds(self, **kwargs):
|
||||
out = super().extra_conds(**kwargs)
|
||||
cross_attn = kwargs.get("cross_attn", None)
|
||||
if cross_attn is not None:
|
||||
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
|
||||
ref_latents = kwargs.get("reference_latents", None)
|
||||
if ref_latents is not None:
|
||||
out['ref_latents'] = comfy.conds.CONDList([self.process_latent_in(lat) for lat in ref_latents])
|
||||
return out
|
||||
|
||||
def extra_conds_shapes(self, **kwargs):
|
||||
out = {}
|
||||
ref_latents = kwargs.get("reference_latents", None)
|
||||
if ref_latents is not None:
|
||||
out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16])
|
||||
return out
|
||||
|
||||
class Ideogram4(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.ideogram4.model.Ideogram4Transformer2DModel)
|
||||
@@ -2314,30 +2294,12 @@ class Ideogram4(BaseModel):
|
||||
class Krea2(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.FLUX, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.krea2.model.SingleStreamDiT)
|
||||
self.memory_usage_factor_conds = ("ref_latents",)
|
||||
|
||||
def extra_conds(self, **kwargs):
|
||||
out = super().extra_conds(**kwargs)
|
||||
cross_attn = kwargs.get("cross_attn", None)
|
||||
if cross_attn is not None:
|
||||
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
|
||||
ref_latents = kwargs.get("reference_latents", None)
|
||||
if ref_latents is not None:
|
||||
latents = []
|
||||
for lat in ref_latents:
|
||||
latents.append(self.process_latent_in(lat))
|
||||
out['ref_latents'] = comfy.conds.CONDList(latents)
|
||||
|
||||
ref_latents_method = kwargs.get("reference_latents_method", None)
|
||||
if ref_latents_method is not None:
|
||||
out['ref_latents_method'] = comfy.conds.CONDConstant(ref_latents_method)
|
||||
return out
|
||||
|
||||
def extra_conds_shapes(self, **kwargs):
|
||||
out = {}
|
||||
ref_latents = kwargs.get("reference_latents", None)
|
||||
if ref_latents is not None:
|
||||
out['ref_latents'] = list([1, 16, sum(map(lambda a: math.prod(a.size()), ref_latents)) // 16])
|
||||
return out
|
||||
|
||||
class HunyuanImage21(BaseModel):
|
||||
|
||||
@@ -1058,25 +1058,6 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
|
||||
dit_config["image_model"] = "SAM31"
|
||||
return dit_config
|
||||
|
||||
if (
|
||||
'{}double_blocks.0.attn.img_attn_qkv.weight'.format(key_prefix) in state_dict_keys
|
||||
and '{}double_blocks.0.attn.img_attn_q_norm.weight'.format(key_prefix) in state_dict_keys
|
||||
and '{}condition_embedder.time_embedder.linear_1.weight'.format(key_prefix) in state_dict_keys
|
||||
and '{}img_in.weight'.format(key_prefix) in state_dict_keys
|
||||
and len(state_dict['{}img_in.weight'.format(key_prefix)].shape) == 5
|
||||
):
|
||||
img_in = state_dict['{}img_in.weight'.format(key_prefix)]
|
||||
head_dim = state_dict['{}double_blocks.0.attn.img_attn_q_norm.weight'.format(key_prefix)].shape[0]
|
||||
return {
|
||||
"image_model": "joyimage",
|
||||
"in_channels": img_in.shape[1],
|
||||
"hidden_size": img_in.shape[0],
|
||||
"patch_size": list(img_in.shape[2:]),
|
||||
"num_layers": count_blocks(state_dict_keys, '{}double_blocks.'.format(key_prefix) + '{}.'),
|
||||
"num_attention_heads": img_in.shape[0] // head_dim,
|
||||
"text_dim": 4096,
|
||||
}
|
||||
|
||||
if '{}input_blocks.0.0.weight'.format(key_prefix) not in state_dict_keys:
|
||||
return None
|
||||
|
||||
|
||||
@@ -473,7 +473,7 @@ except:
|
||||
|
||||
SUPPORT_FP8_OPS = args.supports_fp8_compute
|
||||
|
||||
AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1035", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"]
|
||||
AMD_RDNA2_AND_OLDER_ARCH = ["gfx1030", "gfx1031", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"]
|
||||
AMD_ENABLE_MIOPEN_ENV = 'COMFYUI_ENABLE_MIOPEN'
|
||||
|
||||
try:
|
||||
|
||||
+2
-13
@@ -76,7 +76,6 @@ import comfy.text_encoders.gemma4
|
||||
import comfy.text_encoders.cogvideo
|
||||
import comfy.text_encoders.sa3
|
||||
import comfy.text_encoders.gpt_oss
|
||||
import comfy.text_encoders.joyimage
|
||||
|
||||
import comfy.model_patcher
|
||||
import comfy.lora
|
||||
@@ -1378,7 +1377,6 @@ class CLIPType(Enum):
|
||||
IDEOGRAM4 = 30
|
||||
BOOGU = 31
|
||||
KREA2 = 32
|
||||
JOYIMAGE = 33
|
||||
|
||||
|
||||
|
||||
@@ -1434,7 +1432,6 @@ class TEModel(Enum):
|
||||
GPT_OSS_20B = 33
|
||||
QWEN3VL_4B = 34
|
||||
QWEN3VL_8B = 35
|
||||
GEMMA_4_12B = 36
|
||||
|
||||
|
||||
def detect_te_model(sd):
|
||||
@@ -1464,9 +1461,6 @@ def detect_te_model(sd):
|
||||
if 'model.layers.0.post_feedforward_layernorm.weight' in sd:
|
||||
if 'model.layers.59.self_attn.q_norm.weight' in sd:
|
||||
return TEModel.GEMMA_4_31B
|
||||
# Gemma4 12B Unified: 48 layers, encoder-free; global layers drop v_proj (attention_k_eq_v).
|
||||
if 'model.layers.47.self_attn.q_norm.weight' in sd and 'model.layers.5.self_attn.v_proj.weight' not in sd:
|
||||
return TEModel.GEMMA_4_12B
|
||||
if 'model.layers.41.self_attn.q_norm.weight' in sd and 'model.layers.47.self_attn.q_norm.weight' not in sd:
|
||||
return TEModel.GEMMA_4_E4B
|
||||
if 'model.layers.34.self_attn.q_norm.weight' in sd and 'model.layers.41.self_attn.q_norm.weight' not in sd:
|
||||
@@ -1622,11 +1616,10 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
|
||||
clip_target.clip = comfy.text_encoders.sa3.SAT5GemmaModel
|
||||
clip_target.tokenizer = comfy.text_encoders.sa3.SAT5GemmaTokenizer
|
||||
tokenizer_data["spiece_model"] = clip_data[0].get("spiece_model", None)
|
||||
elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B, TEModel.GEMMA_4_12B):
|
||||
elif te_model in (TEModel.GEMMA_4_E4B, TEModel.GEMMA_4_E2B, TEModel.GEMMA_4_31B):
|
||||
variant = {TEModel.GEMMA_4_E4B: comfy.text_encoders.gemma4.Gemma4_E4B,
|
||||
TEModel.GEMMA_4_E2B: comfy.text_encoders.gemma4.Gemma4_E2B,
|
||||
TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B,
|
||||
TEModel.GEMMA_4_12B: comfy.text_encoders.gemma4.Gemma4_12B}[te_model]
|
||||
TEModel.GEMMA_4_31B: comfy.text_encoders.gemma4.Gemma4_31B}[te_model]
|
||||
clip_target.clip = comfy.text_encoders.gemma4.gemma4_te(**llama_detect(clip_data), model_class=variant)
|
||||
clip_target.tokenizer = variant.tokenizer
|
||||
tokenizer_data["tokenizer_json"] = clip_data[0].get("tokenizer_json", None)
|
||||
@@ -1713,10 +1706,6 @@ def load_text_encoder_state_dicts(state_dicts=[], embedding_directory=None, clip
|
||||
clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."})
|
||||
clip_target.clip = comfy.text_encoders.krea2.te(**llama_detect(clip_data))
|
||||
clip_target.tokenizer = comfy.text_encoders.krea2.Krea2Tokenizer
|
||||
elif clip_type == CLIPType.JOYIMAGE and te_model == TEModel.QWEN3VL_8B: # JoyImageEdit: full Qwen3-VL-8B, edit-conditioning template + drop_idx.
|
||||
clip_data[0] = comfy.utils.state_dict_prefix_replace(clip_data[0], {"model.language_model.": "model.", "model.visual.": "visual.", "lm_head.": "model.lm_head."})
|
||||
clip_target.clip = comfy.text_encoders.joyimage.te(**llama_detect(clip_data))
|
||||
clip_target.tokenizer = comfy.text_encoders.joyimage.JoyImageTokenizer
|
||||
elif clip_type in (CLIPType.FLUX, CLIPType.FLUX2): # Flux2 Klein reuses the Qwen3-VL LM (3-layer tap -> 12288); visual unused.
|
||||
klein_model_type = "qwen3_8b" if te_model == TEModel.QWEN3VL_8B else "qwen3_4b"
|
||||
clip_target.clip = comfy.text_encoders.flux.klein_te(**llama_detect(clip_data), model_type=klein_model_type)
|
||||
|
||||
@@ -27,7 +27,6 @@ import comfy.text_encoders.z_image
|
||||
import comfy.text_encoders.ideogram4
|
||||
import comfy.text_encoders.boogu
|
||||
import comfy.text_encoders.krea2
|
||||
import comfy.text_encoders.joyimage
|
||||
import comfy.text_encoders.anima
|
||||
import comfy.text_encoders.ace15
|
||||
import comfy.text_encoders.longcat_image
|
||||
@@ -1912,38 +1911,6 @@ class QwenImage(supported_models_base.BASE):
|
||||
hunyuan_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen25_7b.transformer.".format(pref))
|
||||
return supported_models_base.ClipTarget(comfy.text_encoders.qwen_image.QwenImageTokenizer, comfy.text_encoders.qwen_image.te(**hunyuan_detect))
|
||||
|
||||
class JoyImage(supported_models_base.BASE):
|
||||
unet_config = {
|
||||
"image_model": "joyimage",
|
||||
}
|
||||
|
||||
sampling_settings = {
|
||||
"multiplier": 1000,
|
||||
"shift": 1.5,
|
||||
}
|
||||
|
||||
memory_usage_factor = 1.8
|
||||
|
||||
unet_extra_config = {
|
||||
"theta": 10000,
|
||||
"rope_dim_list": [16, 56, 56],
|
||||
}
|
||||
|
||||
latent_format = latent_formats.Wan21
|
||||
|
||||
supported_inference_dtypes = [torch.bfloat16, torch.float32]
|
||||
|
||||
vae_key_prefix = ["vae."]
|
||||
text_encoder_key_prefix = ["text_encoders."]
|
||||
|
||||
def get_model(self, state_dict, prefix="", device=None):
|
||||
return model_base.JoyImage(self, device=device)
|
||||
|
||||
def clip_target(self, state_dict={}):
|
||||
pref = self.text_encoder_key_prefix[0]
|
||||
qwen3vl_detect = comfy.text_encoders.hunyuan_video.llama_detect(state_dict, "{}qwen3vl.transformer.".format(pref))
|
||||
return supported_models_base.ClipTarget(comfy.text_encoders.joyimage.JoyImageTokenizer, comfy.text_encoders.joyimage.te(**qwen3vl_detect))
|
||||
|
||||
class HunyuanImage21(HunyuanVideo):
|
||||
unet_config = {
|
||||
"image_model": "hunyuan_video",
|
||||
@@ -2422,7 +2389,6 @@ models = [
|
||||
Omnigen2,
|
||||
Boogu,
|
||||
QwenImage,
|
||||
JoyImage,
|
||||
Ideogram4,
|
||||
Krea2,
|
||||
Flux2,
|
||||
|
||||
+44
-267
@@ -1,15 +1,11 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torchaudio.functional as AF
|
||||
import torchvision.transforms.functional as TVF
|
||||
import numpy as np
|
||||
from tokenizers import Tokenizer
|
||||
from dataclasses import dataclass
|
||||
import math
|
||||
|
||||
from comfy import sd1_clip
|
||||
import comfy.model_management
|
||||
import comfy.ops
|
||||
from comfy.ldm.modules.attention import optimized_attention_for_device
|
||||
from comfy.rmsnorm import rms_norm
|
||||
from comfy.text_encoders.llama import RMSNorm, MLP, BaseLlama, BaseGenerate, _make_scaled_embedding
|
||||
@@ -25,10 +21,6 @@ GEMMA4_VISION_CONFIG = {"hidden_size": 768, "image_size": 896, "intermediate_siz
|
||||
GEMMA4_VISION_31B_CONFIG = {"hidden_size": 1152, "image_size": 896, "intermediate_size": 4304, "num_attention_heads": 16, "num_hidden_layers": 27, "patch_size": 16, "head_dim": 72, "rms_norm_eps": 1e-6, "position_embedding_size": 10240, "pooling_kernel_size": 3}
|
||||
GEMMA4_AUDIO_CONFIG = {"hidden_size": 1024, "num_hidden_layers": 12, "num_attention_heads": 8, "intermediate_size": 4096, "conv_kernel_size": 5, "attention_chunk_size": 12, "attention_context_left": 13, "attention_context_right": 0, "attention_logit_cap": 50.0, "output_proj_dims": 1536, "rms_norm_eps": 1e-6, "residual_weight": 0.5}
|
||||
|
||||
# Encoder-free (gemma4_unified) multimodal embedders: raw patches/waveform projected directly into LM space.
|
||||
GEMMA4_UNIFIED_VISION_CONFIG = {"model_patch_size": 48, "patch_size": 16, "pooling_kernel_size": 3, "mm_embed_dim": 3840, "mm_posemb_size": 1120, "output_proj_dims": 3840, "rms_norm_eps": 1e-6}
|
||||
GEMMA4_UNIFIED_AUDIO_CONFIG = {"audio_samples_per_token": 640, "output_proj_dims": 640, "rms_norm_eps": 1e-6}
|
||||
|
||||
@dataclass
|
||||
class Gemma4Config:
|
||||
vocab_size: int = 262144
|
||||
@@ -43,9 +35,6 @@ class Gemma4Config:
|
||||
transformer_type: str = "gemma4"
|
||||
head_dim = 256
|
||||
global_head_dim = 512
|
||||
num_global_key_value_heads = None
|
||||
attention_k_eq_v = False
|
||||
vision_bidirectional = False
|
||||
rms_norm_add = False
|
||||
mlp_activation = "gelu_pytorch_tanh"
|
||||
qkv_bias = False
|
||||
@@ -62,7 +51,6 @@ class Gemma4Config:
|
||||
num_kv_shared_layers: int = 18
|
||||
use_double_wide_mlp: bool = False
|
||||
stop_tokens = [1, 50, 106]
|
||||
suppress_tokens = []
|
||||
vision_config = GEMMA4_VISION_CONFIG
|
||||
audio_config = GEMMA4_AUDIO_CONFIG
|
||||
mm_tokens_per_image = 280
|
||||
@@ -84,30 +72,12 @@ class Gemma4_31B_Config(Gemma4Config):
|
||||
num_hidden_layers: int = 60
|
||||
num_attention_heads: int = 32
|
||||
num_key_value_heads: int = 16
|
||||
vision_bidirectional = True
|
||||
sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
|
||||
hidden_size_per_layer_input: int = 0
|
||||
num_kv_shared_layers: int = 0
|
||||
audio_config = None
|
||||
vision_config = GEMMA4_VISION_31B_CONFIG
|
||||
|
||||
@dataclass
|
||||
class Gemma4_12B_Config(Gemma4Config):
|
||||
hidden_size: int = 3840
|
||||
intermediate_size: int = 15360
|
||||
num_hidden_layers: int = 48
|
||||
num_attention_heads: int = 16
|
||||
num_key_value_heads: int = 8
|
||||
num_global_key_value_heads = 1
|
||||
attention_k_eq_v = True
|
||||
vision_bidirectional = True
|
||||
sliding_attention = [1024, 1024, 1024, 1024, 1024, False]
|
||||
hidden_size_per_layer_input: int = 0
|
||||
num_kv_shared_layers: int = 0
|
||||
audio_config = GEMMA4_UNIFIED_AUDIO_CONFIG
|
||||
vision_config = GEMMA4_UNIFIED_VISION_CONFIG
|
||||
suppress_tokens = [258883, 258882]
|
||||
|
||||
|
||||
# unfused RoPE as addcmul_ RoPE diverges from reference code
|
||||
def _apply_rotary_pos_emb(x, freqs_cis):
|
||||
@@ -119,18 +89,17 @@ def _apply_rotary_pos_emb(x, freqs_cis):
|
||||
return out
|
||||
|
||||
class Gemma4Attention(nn.Module):
|
||||
def __init__(self, config, head_dim, num_kv_heads=None, k_eq_v=False, device=None, dtype=None, ops=None):
|
||||
def __init__(self, config, head_dim, device=None, dtype=None, ops=None):
|
||||
super().__init__()
|
||||
self.num_heads = config.num_attention_heads
|
||||
self.num_kv_heads = num_kv_heads if num_kv_heads is not None else config.num_key_value_heads
|
||||
self.num_kv_heads = config.num_key_value_heads
|
||||
self.hidden_size = config.hidden_size
|
||||
self.head_dim = head_dim
|
||||
self.inner_size = self.num_heads * head_dim
|
||||
|
||||
self.q_proj = ops.Linear(config.hidden_size, self.inner_size, bias=config.qkv_bias, device=device, dtype=dtype)
|
||||
self.k_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
|
||||
# k_eq_v: V reuses the K projection (no separate v_proj weight)
|
||||
self.v_proj = None if k_eq_v else ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
|
||||
self.v_proj = ops.Linear(config.hidden_size, self.num_kv_heads * head_dim, bias=config.qkv_bias, device=device, dtype=dtype)
|
||||
self.o_proj = ops.Linear(self.inner_size, config.hidden_size, bias=False, device=device, dtype=dtype)
|
||||
|
||||
self.q_norm = None
|
||||
@@ -164,10 +133,7 @@ class Gemma4Attention(nn.Module):
|
||||
shareable_kv = None
|
||||
else:
|
||||
xk = self.k_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
|
||||
if self.v_proj is not None:
|
||||
xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
|
||||
else:
|
||||
xv = xk # k_eq_v: V is the raw K projection (before k_norm/RoPE)
|
||||
xv = self.v_proj(hidden_states).view(batch_size, seq_length, self.num_kv_heads, self.head_dim)
|
||||
if self.k_norm is not None:
|
||||
xk = self.k_norm(xk)
|
||||
xv = rms_norm(xv)
|
||||
@@ -220,10 +186,7 @@ class TransformerBlockGemma4(nn.Module):
|
||||
|
||||
head_dim = config.head_dim if self.sliding_attention else config.global_head_dim
|
||||
|
||||
# k_eq_v only on global layers, which then use num_global_key_value_heads
|
||||
k_eq_v = config.attention_k_eq_v and not self.sliding_attention
|
||||
num_kv_heads = config.num_global_key_value_heads if k_eq_v else config.num_key_value_heads
|
||||
self.self_attn = Gemma4Attention(config, head_dim=head_dim, num_kv_heads=num_kv_heads, k_eq_v=k_eq_v, device=device, dtype=dtype, ops=ops)
|
||||
self.self_attn = Gemma4Attention(config, head_dim=head_dim, device=device, dtype=dtype, ops=ops)
|
||||
|
||||
num_kv_shared = config.num_kv_shared_layers
|
||||
first_kv_shared = config.num_hidden_layers - num_kv_shared
|
||||
@@ -240,9 +203,9 @@ class TransformerBlockGemma4(nn.Module):
|
||||
self.per_layer_input_gate = ops.Linear(config.hidden_size, self.hidden_size_per_layer_input, bias=False, device=device, dtype=dtype)
|
||||
self.per_layer_projection = ops.Linear(self.hidden_size_per_layer_input, config.hidden_size, bias=False, device=device, dtype=dtype)
|
||||
self.post_per_layer_input_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, device=device, dtype=dtype)
|
||||
|
||||
# layer_scalar exists on every gemma4 variant, independent of per-layer input
|
||||
self.register_buffer("layer_scalar", torch.empty(1, device=device, dtype=dtype))
|
||||
self.register_buffer("layer_scalar", torch.ones(1, device=device, dtype=dtype))
|
||||
else:
|
||||
self.layer_scalar = None
|
||||
|
||||
def forward(self, x, attention_mask=None, freqs_cis=None, past_key_value=None, per_layer_input=None, shared_kv=None):
|
||||
sliding_window = None
|
||||
@@ -281,7 +244,8 @@ class TransformerBlockGemma4(nn.Module):
|
||||
x = self.post_per_layer_input_norm(x)
|
||||
x = residual + x
|
||||
|
||||
x = x * comfy.ops.cast_to_input(self.layer_scalar, x)
|
||||
if self.layer_scalar is not None:
|
||||
x = x * self.layer_scalar
|
||||
|
||||
return x, present_key_value, shareable_kv
|
||||
|
||||
@@ -370,19 +334,6 @@ class Gemma4Transformer(nn.Module):
|
||||
causal_mask.masked_fill_(torch.ones_like(causal_mask, dtype=torch.bool).triu_(1), min_val)
|
||||
mask = mask + causal_mask if mask is not None else causal_mask
|
||||
|
||||
# Bidirectional attention within each image soft-token block (prefill only; text/audio stay causal).
|
||||
if self.config.vision_bidirectional and past_len == 0 and embeds_info:
|
||||
block_ids = torch.full((seq_len,), -1, dtype=torch.long, device=x.device)
|
||||
group = 0
|
||||
for info in embeds_info:
|
||||
if info.get("type") == "image":
|
||||
start = info["index"]
|
||||
block_ids[start:start + info["size"]] = group
|
||||
group += 1
|
||||
if group > 0:
|
||||
same_block = (block_ids[:, None] == block_ids[None, :]) & (block_ids[:, None] >= 0)
|
||||
mask = mask.masked_fill(same_block, 0.0)
|
||||
|
||||
# Per-layer inputs
|
||||
per_layer_inputs = None
|
||||
if self.hidden_size_per_layer_input:
|
||||
@@ -403,24 +354,8 @@ class Gemma4Transformer(nn.Module):
|
||||
shared_global_kv = None # KV from last non-shared global layer
|
||||
|
||||
intermediate = None
|
||||
all_intermediate = None
|
||||
only_layers = None
|
||||
if intermediate_output is not None:
|
||||
if isinstance(intermediate_output, list):
|
||||
all_intermediate = []
|
||||
only_layers = {len(self.layers) + layer if layer < 0 else layer for layer in intermediate_output}
|
||||
elif intermediate_output == "all":
|
||||
all_intermediate = []
|
||||
intermediate_output = None
|
||||
elif intermediate_output < 0:
|
||||
intermediate_output = len(self.layers) + intermediate_output
|
||||
|
||||
next_key_values = []
|
||||
for i, layer in enumerate(self.layers):
|
||||
if all_intermediate is not None:
|
||||
if only_layers is None or (i in only_layers):
|
||||
all_intermediate.append(x.unsqueeze(1).clone())
|
||||
|
||||
past_kv = past_key_values[i] if past_key_values is not None and len(past_key_values) > 0 else None
|
||||
|
||||
layer_kwargs = {}
|
||||
@@ -450,18 +385,7 @@ class Gemma4Transformer(nn.Module):
|
||||
if self.norm is not None:
|
||||
x = self.norm(x)
|
||||
|
||||
if all_intermediate is not None:
|
||||
if only_layers is None or (len(self.layers) in only_layers):
|
||||
all_intermediate.append(x.unsqueeze(1).clone())
|
||||
if len(all_intermediate) > 0:
|
||||
intermediate = torch.cat(all_intermediate, dim=1)
|
||||
|
||||
if intermediate is not None and final_layer_norm_intermediate and self.norm is not None:
|
||||
intermediate = self.norm(intermediate)
|
||||
|
||||
# Only hand back the KV cache when caching was actually requested; SDClipModel reads
|
||||
# outputs[2] as the pooled output.
|
||||
if past_key_values is not None and len(next_key_values) > 0:
|
||||
if len(next_key_values) > 0:
|
||||
return x, intermediate, next_key_values
|
||||
return x, intermediate
|
||||
|
||||
@@ -480,8 +404,6 @@ class Gemma4Base(BaseLlama, BaseGenerate, torch.nn.Module):
|
||||
cap = self.model.config.final_logit_softcapping
|
||||
if cap:
|
||||
logits = cap * torch.tanh(logits / cap)
|
||||
if self.model.config.suppress_tokens:
|
||||
logits[..., self.model.config.suppress_tokens] = torch.finfo(logits.dtype).min
|
||||
return logits
|
||||
|
||||
def init_kv_cache(self, batch, max_cache_len, device, execution_dtype):
|
||||
@@ -519,28 +441,6 @@ class Gemma4AudioMixin:
|
||||
return None, None
|
||||
|
||||
|
||||
class Gemma4UnifiedBase(Gemma4Base):
|
||||
"""Encoder-free multimodal Gemma4 (gemma4_unified, e.g. 12B): raw image patches and audio frames projected directly into LM space."""
|
||||
def _init_model(self, config, dtype, device, operations):
|
||||
self.num_layers = config.num_hidden_layers
|
||||
self.model = Gemma4Transformer(config, device=device, dtype=dtype, ops=operations)
|
||||
self.dtype = dtype
|
||||
self.vision_model = Gemma4UnifiedVisionEmbedder(config.vision_config, device=device, dtype=dtype, ops=operations)
|
||||
self.multi_modal_projector = Gemma4RMSNormProjector(config.vision_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
|
||||
self.audio_projector = Gemma4RMSNormProjector(config.audio_config["output_proj_dims"], config.hidden_size, dtype=dtype, device=device, ops=operations)
|
||||
|
||||
def preprocess_embed(self, embed, device):
|
||||
if embed["type"] == "image":
|
||||
pixels = embed.pop("data").movedim(-1, 1).to(device, dtype=self.dtype) # [B, H, W, C] -> [B, C, H, W], [0,1]
|
||||
patches, positions = self.vision_model.patchify(pixels)
|
||||
vision_out = self.vision_model(patches, positions)
|
||||
return self.multi_modal_projector(vision_out), None
|
||||
if embed["type"] == "audio":
|
||||
audio = embed.pop("data").to(device, dtype=self.dtype) # [1, T, audio_samples_per_token]
|
||||
return self.audio_projector(audio), None
|
||||
return None, None
|
||||
|
||||
|
||||
# Vision Encoder
|
||||
|
||||
def _compute_vision_2d_rope(head_dim, pixel_position_ids, theta=100.0, device=None):
|
||||
@@ -813,73 +713,6 @@ class Gemma4MultiModalProjector(Gemma4RMSNormProjector):
|
||||
super().__init__(config.vision_config["hidden_size"], config.hidden_size, dtype=dtype, device=device, ops=ops)
|
||||
|
||||
|
||||
# Encoder-free vision (gemma4_unified): raw merged pixel patches projected directly into LM space.
|
||||
|
||||
def _patches_merge(patches, positions_xy, length):
|
||||
patch_size = math.isqrt(patches.shape[-1] // 3)
|
||||
k = math.isqrt(patches.shape[-2] // length)
|
||||
batch = patches.shape[:-2]
|
||||
|
||||
max_x = positions_xy[..., 0].max(dim=-1, keepdim=True)[0] + 1
|
||||
kidx = torch.div(positions_xy, k, rounding_mode="floor")
|
||||
rem = torch.remainder(positions_xy, k)
|
||||
order = rem[..., 0] + rem[..., 1] * k + k * k * kidx[..., 0] + k * max_x * kidx[..., 1]
|
||||
perm = order.long().argsort(dim=-1)
|
||||
|
||||
merged = patches.gather(-2, perm.unsqueeze(-1).expand_as(patches))
|
||||
merged = merged.reshape(*batch, length, k, k, patch_size, patch_size, 3)
|
||||
merged = merged.permute(*range(len(batch)), -6, -5, -3, -4, -2, -1).reshape(*batch, length, (k * patch_size) ** 2 * 3)
|
||||
|
||||
pos = positions_xy.gather(-2, perm.unsqueeze(-1).expand_as(positions_xy))
|
||||
pad = (positions_xy == -1).all(dim=-1, keepdim=True)
|
||||
pos = torch.where(pad, positions_xy, pos).reshape(*batch, length, k * k, 2)
|
||||
pos = torch.div(pos, k, rounding_mode="floor").min(dim=-2)[0]
|
||||
return merged, pos
|
||||
|
||||
|
||||
class Gemma4UnifiedVisionEmbedder(nn.Module):
|
||||
"""Encoder-free patch embedder (LN -> Dense -> LN -> +2D posemb -> LN); projection to text space is the separate multi_modal_projector."""
|
||||
def __init__(self, config, device=None, dtype=None, ops=None):
|
||||
super().__init__()
|
||||
self.patch_size = config["patch_size"]
|
||||
self.pooling_kernel_size = config["pooling_kernel_size"]
|
||||
patch_dim = config["model_patch_size"] ** 2 * 3
|
||||
mm_embed_dim = config["mm_embed_dim"]
|
||||
self.patch_ln1 = ops.LayerNorm(patch_dim, device=device, dtype=dtype)
|
||||
self.patch_dense = ops.Linear(patch_dim, mm_embed_dim, device=device, dtype=dtype)
|
||||
self.patch_ln2 = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
|
||||
self.pos_embedding = nn.Parameter(torch.empty(config["mm_posemb_size"], 2, mm_embed_dim, device=device, dtype=dtype))
|
||||
self.pos_norm = ops.LayerNorm(mm_embed_dim, device=device, dtype=dtype)
|
||||
|
||||
def patchify(self, pixels):
|
||||
"""pixels: [B, C, H, W] in [0,1] -> merged patches [B, N, 6912], positions [B, N, 2]."""
|
||||
ps, k = self.patch_size, self.pooling_kernel_size
|
||||
out_patches, out_positions = [], []
|
||||
for img in pixels:
|
||||
ph, pw = img.shape[-2] // ps, img.shape[-1] // ps
|
||||
teacher = img.reshape(img.shape[0], ph, ps, pw, ps).permute(1, 3, 2, 4, 0).reshape(ph * pw, -1)
|
||||
grid = torch.meshgrid(torch.arange(pw, device=img.device), torch.arange(ph, device=img.device), indexing="xy")
|
||||
tpos = torch.stack(grid, dim=-1).reshape(teacher.shape[0], 2)
|
||||
n_model = teacher.shape[0] // (k * k)
|
||||
mp, mpos = _patches_merge(teacher.unsqueeze(0), tpos.unsqueeze(0), n_model)
|
||||
out_patches.append(mp.squeeze(0))
|
||||
out_positions.append(mpos.squeeze(0))
|
||||
return torch.stack(out_patches), torch.stack(out_positions)
|
||||
|
||||
def forward(self, pixel_values, image_position_ids):
|
||||
x = self.patch_ln1(pixel_values)
|
||||
x = self.patch_dense(x)
|
||||
x = self.patch_ln2(x)
|
||||
|
||||
clamped = image_position_ids.clamp(min=0).long()
|
||||
valid = (image_position_ids != -1).to(x.dtype).unsqueeze(-1)
|
||||
axes = torch.arange(2, device=image_position_ids.device)
|
||||
pos = comfy.model_management.cast_to_device(self.pos_embedding, x.device, x.dtype)
|
||||
pos_embs = (pos[clamped, axes] * valid).sum(-2)
|
||||
x = x + pos_embs
|
||||
return self.pos_norm(x)
|
||||
|
||||
|
||||
# Audio Encoder
|
||||
|
||||
class Gemma4AudioConvSubsampler(nn.Module):
|
||||
@@ -1157,30 +990,6 @@ class Gemma4AudioProjector(Gemma4RMSNormProjector):
|
||||
|
||||
# Tokenizer and Wrappers
|
||||
|
||||
def _get_aspect_ratio_preserving_size(height, width, patch_size, max_patches, pooling_kernel_size):
|
||||
target_px = max_patches * patch_size ** 2
|
||||
factor = math.sqrt(target_px / (height * width))
|
||||
side_mult = pooling_kernel_size * patch_size
|
||||
target_height = math.floor(factor * height / side_mult) * side_mult
|
||||
target_width = math.floor(factor * width / side_mult) * side_mult
|
||||
|
||||
if target_height == 0 and target_width == 0:
|
||||
raise ValueError(f"Attempting to resize to a 0 x 0 image. Resized height should be divisible by {side_mult}.")
|
||||
|
||||
max_side_length = (max_patches // pooling_kernel_size ** 2) * side_mult
|
||||
if target_height == 0:
|
||||
target_height = side_mult
|
||||
target_width = min(math.floor(width / height) * side_mult, max_side_length)
|
||||
elif target_width == 0:
|
||||
target_width = side_mult
|
||||
target_height = min(math.floor(height / width) * side_mult, max_side_length)
|
||||
|
||||
if target_height * target_width > target_px:
|
||||
raise ValueError(f"Resizing [{height}x{width}] to [{target_height}x{target_width}] exceeds the patch budget.")
|
||||
|
||||
return target_height, target_width
|
||||
|
||||
|
||||
class Gemma4_Tokenizer():
|
||||
tokenizer_json_data = None
|
||||
|
||||
@@ -1189,35 +998,25 @@ class Gemma4_Tokenizer():
|
||||
return {"tokenizer_json": self.tokenizer_json_data}
|
||||
return {}
|
||||
|
||||
def _audio_token_count(self, num_samples):
|
||||
# Default (E2B/E4B): mel frames after two stride-2 conv subsamples.
|
||||
_fl = 320 # int(round(16000 * 20.0 / 1000.0))
|
||||
_hl = 160 # int(round(16000 * 10.0 / 1000.0))
|
||||
_nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
|
||||
_t = _nmel
|
||||
for _ in range(2):
|
||||
_t = (_t + 2 - 3) // 2 + 1
|
||||
return min(_t, 750)
|
||||
|
||||
@staticmethod
|
||||
def _resample_16k(waveform, sample_rate):
|
||||
"""Mix to mono and resample to 16kHz. Kaiser params reproduce the reference (transformers
|
||||
load_audio -> librosa/soxr_hq) to ~1e-12 MSE using only torchaudio."""
|
||||
def _extract_mel_spectrogram(self, waveform, sample_rate):
|
||||
"""Extract 128-bin log mel spectrogram.
|
||||
Uses numpy for FFT/matmul/log to produce bit-identical results with reference code.
|
||||
"""
|
||||
# Mix to mono first, then resample to 16kHz
|
||||
if waveform.dim() > 1 and waveform.shape[0] > 1:
|
||||
waveform = waveform.mean(dim=0, keepdim=True)
|
||||
if waveform.dim() == 1:
|
||||
waveform = waveform.unsqueeze(0)
|
||||
audio = waveform.float()
|
||||
audio = waveform.squeeze(0).float().numpy()
|
||||
if sample_rate != 16000:
|
||||
audio = AF.resample(audio, sample_rate, 16000, resampling_method="sinc_interp_kaiser",
|
||||
lowpass_filter_width=121, rolloff=0.9568384289091556, beta=21.01531462440614)
|
||||
return audio.squeeze(0).contiguous()
|
||||
|
||||
def _extract_audio_features(self, waveform, sample_rate):
|
||||
"""Default (E2B/E4B): 128-bin log mel spectrogram for the conformer audio encoder.
|
||||
Uses numpy for FFT/matmul/log to produce bit-identical results with reference code.
|
||||
"""
|
||||
audio = self._resample_16k(waveform, sample_rate).numpy()
|
||||
# Use scipy's resample_poly with a high-quality FIR filter to get as close as possible to librosa's resampling (while still not full match)
|
||||
from scipy.signal import resample_poly, firwin
|
||||
from math import gcd
|
||||
g = gcd(sample_rate, 16000)
|
||||
up, down = 16000 // g, sample_rate // g
|
||||
L = max(up, down)
|
||||
h = firwin(160 * L + 1, 0.96 / L, window=('kaiser', 6.5))
|
||||
audio = resample_poly(audio, up, down, window=h).astype(np.float32)
|
||||
n = len(audio)
|
||||
|
||||
# Pad to multiple of 128, build sample-level mask
|
||||
@@ -1265,8 +1064,8 @@ class Gemma4_Tokenizer():
|
||||
if audio is not None:
|
||||
waveform = audio["waveform"].squeeze(0) if hasattr(audio, "__getitem__") else audio
|
||||
sample_rate = audio.get("sample_rate", 16000) if hasattr(audio, "get") else 16000
|
||||
feat, feat_mask = self._extract_audio_features(waveform, sample_rate)
|
||||
audio_features = [(feat.unsqueeze(0), feat_mask.unsqueeze(0))] # ([1, T, D], [1, T])
|
||||
mel, mel_mask = self._extract_mel_spectrogram(waveform, sample_rate)
|
||||
audio_features = [(mel.unsqueeze(0), mel_mask.unsqueeze(0))] # ([1, T, 128], [1, T])
|
||||
|
||||
# Process image/video frames
|
||||
is_video = video is not None
|
||||
@@ -1291,8 +1090,13 @@ class Gemma4_Tokenizer():
|
||||
pooling_k = 3
|
||||
max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280)
|
||||
max_patches = max_soft_tokens * pooling_k * pooling_k
|
||||
target_h, target_w = _get_aspect_ratio_preserving_size(h, w, patch_size, max_patches, pooling_k)
|
||||
target_px = max_patches * patch_size * patch_size
|
||||
factor = (target_px / (h * w)) ** 0.5
|
||||
side_mult = pooling_k * patch_size
|
||||
target_h = max(int(factor * h // side_mult) * side_mult, side_mult)
|
||||
target_w = max(int(factor * w // side_mult) * side_mult, side_mult)
|
||||
|
||||
import torchvision.transforms.functional as TVF
|
||||
for i in range(num_frames):
|
||||
# rescaling to match reference code
|
||||
s = (samples[i].clamp(0, 1) * 255).to(torch.uint8) # [C, H, W] uint8
|
||||
@@ -1311,7 +1115,7 @@ class Gemma4_Tokenizer():
|
||||
llama_text = llama_template.format(text)
|
||||
else:
|
||||
# Build template from modalities present
|
||||
system = "<|turn>system\n<|think|>\n<turn|>\n" if thinking else ""
|
||||
system = "<|turn>system\n<|think|><turn|>\n" if thinking else ""
|
||||
media = ""
|
||||
if len(images) > 0:
|
||||
if is_video:
|
||||
@@ -1331,11 +1135,15 @@ class Gemma4_Tokenizer():
|
||||
if len(audio_features) > 0:
|
||||
# Compute audio token count (always at 16kHz)
|
||||
num_samples = int(waveform.shape[-1] * 16000 / sample_rate) if sample_rate != 16000 else waveform.shape[-1]
|
||||
n_audio_tokens = self._audio_token_count(num_samples)
|
||||
_fl = 320 # int(round(16000 * 20.0 / 1000.0))
|
||||
_hl = 160 # int(round(16000 * 10.0 / 1000.0))
|
||||
_nmel = (num_samples + _fl // 2 - (_fl + 1)) // _hl + 1
|
||||
_t = _nmel
|
||||
for _ in range(2):
|
||||
_t = (_t + 2 - 3) // 2 + 1
|
||||
n_audio_tokens = min(_t, 750)
|
||||
media += "<|audio>" + "<|audio|>" * n_audio_tokens + "<audio|>"
|
||||
# Non-thinking mode primes an empty thought channel so the model answers directly.
|
||||
model_open = "" if thinking else "<|channel>thought\n<channel|>"
|
||||
llama_text = f"{system}<|turn>user\n{text}{media}<turn|>\n<|turn>model\n{model_open}"
|
||||
llama_text = f"{system}<|turn>user\n{media}{text}<turn|>\n<|turn>model\n"
|
||||
|
||||
text_tokens = super().tokenize_with_weights(llama_text, return_word_ids)
|
||||
|
||||
@@ -1370,6 +1178,7 @@ class Gemma4_Tokenizer():
|
||||
class _Gemma4Tokenizer:
|
||||
"""Tokenizer using the tokenizers (Gemma4 doesn't come with sentencepiece model)"""
|
||||
def __init__(self, tokenizer_json_bytes=None, **kwargs):
|
||||
from tokenizers import Tokenizer
|
||||
if isinstance(tokenizer_json_bytes, torch.Tensor):
|
||||
tokenizer_json_bytes = bytes(tokenizer_json_bytes.tolist())
|
||||
self.tokenizer = Tokenizer.from_str(tokenizer_json_bytes.decode("utf-8"))
|
||||
@@ -1415,30 +1224,6 @@ class Gemma4Tokenizer(sd1_clip.SD1Tokenizer):
|
||||
super().__init__(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data, name="gemma4", tokenizer=self.tokenizer_class)
|
||||
|
||||
|
||||
class Gemma4UnifiedSDTokenizer(Gemma4SDTokenizer):
|
||||
"""Encoder-free (gemma4_unified) audio: raw 16kHz waveform frames instead of mel spectrogram."""
|
||||
embedding_size = 3840
|
||||
|
||||
def _extract_audio_features(self, waveform, sample_rate):
|
||||
audio = self._resample_16k(waveform, sample_rate)
|
||||
spt = 640 # audio_samples_per_token (40ms at 16kHz)
|
||||
pad = (-audio.shape[0]) % spt
|
||||
if pad:
|
||||
audio = torch.nn.functional.pad(audio, (0, pad))
|
||||
num_tokens = audio.shape[0] // spt
|
||||
feats = audio[:num_tokens * spt].reshape(num_tokens, spt)
|
||||
feats = feats[:750] # audio_seq_length cap (matches reference truncation, ~30s)
|
||||
mask = torch.ones(feats.shape[0], dtype=torch.bool)
|
||||
return feats, mask
|
||||
|
||||
def _audio_token_count(self, num_samples):
|
||||
return min((num_samples + 639) // 640, 750)
|
||||
|
||||
|
||||
class Gemma4UnifiedTokenizer(Gemma4Tokenizer):
|
||||
tokenizer_class = Gemma4UnifiedSDTokenizer
|
||||
|
||||
|
||||
# Model wrappers
|
||||
class Gemma4Model(sd1_clip.SDClipModel):
|
||||
model_class = None
|
||||
@@ -1471,7 +1256,7 @@ class Gemma4Model(sd1_clip.SDClipModel):
|
||||
expanded_idx += 1
|
||||
initial_token_ids = [ids]
|
||||
input_ids = torch.tensor(initial_token_ids, device=self.execution_device)
|
||||
return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids, embeds_info=embeds_info)
|
||||
return self.transformer.generate(embeds, do_sample, max_length, temperature, top_k, top_p, min_p, repetition_penalty, seed, initial_tokens=initial_token_ids[0], presence_penalty=presence_penalty, initial_input_ids=input_ids)
|
||||
|
||||
|
||||
def gemma4_te(dtype_llama=None, llama_quantization_metadata=None, model_class=None):
|
||||
@@ -1511,11 +1296,3 @@ def _make_variant(config_cls):
|
||||
Gemma4_E4B = _make_variant(Gemma4Config)
|
||||
Gemma4_E2B = _make_variant(Gemma4_E2B_Config)
|
||||
Gemma4_31B = _make_variant(Gemma4_31B_Config)
|
||||
|
||||
|
||||
# Gemma4 12B Unified: encoder-free multimodal, distinct base/tokenizer (not via _make_variant).
|
||||
class Gemma4_12B(Gemma4UnifiedBase):
|
||||
def __init__(self, config_dict, dtype, device, operations):
|
||||
super().__init__()
|
||||
self._init_model(Gemma4_12B_Config(**config_dict), dtype, device, operations)
|
||||
Gemma4_12B.tokenizer = Gemma4UnifiedTokenizer
|
||||
|
||||
@@ -1,97 +0,0 @@
|
||||
import torch
|
||||
|
||||
from comfy import sd1_clip
|
||||
import comfy.text_encoders.qwen_vl
|
||||
from comfy.text_encoders.qwen3vl import Qwen3VL, Qwen3VLTokenizer
|
||||
|
||||
JOYIMAGE_VISION_BLOCK = "<|vision_start|><|image_pad|><|vision_end|>"
|
||||
JOYIMAGE_TEMPLATE_TEXT = (
|
||||
"<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, "
|
||||
"quantity, text, spatial relationships of the objects and background:<|im_end|>\n"
|
||||
"<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n"
|
||||
)
|
||||
JOYIMAGE_TEMPLATE_IMAGE = (
|
||||
"<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, "
|
||||
"quantity, text, spatial relationships of the objects and background:<|im_end|>\n"
|
||||
f"<|im_start|>user\n{JOYIMAGE_VISION_BLOCK}{{}}<|im_end|>\n<|im_start|>assistant\n"
|
||||
)
|
||||
# The DiT was trained without the leading system-prompt tokens.
|
||||
JOYIMAGE_DROP_IDX = 34
|
||||
PAD_TOKEN = 151643
|
||||
|
||||
|
||||
class Qwen3VL8B_JoyImage(Qwen3VL):
|
||||
model_type = "qwen3vl_8b"
|
||||
|
||||
def preprocess_embed(self, embed, device):
|
||||
if embed["type"] == "image":
|
||||
image, grid = comfy.text_encoders.qwen_vl.process_qwen2vl_images(
|
||||
embed["data"], min_pixels=65536, max_pixels=16777216, patch_size=16,
|
||||
image_mean=[0.5, 0.5, 0.5], image_std=[0.5, 0.5, 0.5],
|
||||
interpolation="bicubic",
|
||||
)
|
||||
merged, deepstack = self.visual(image.to(device, dtype=torch.float32), grid)
|
||||
return merged, {"grid": grid, "deepstack": deepstack}
|
||||
return None, None
|
||||
|
||||
|
||||
class JoyImageTokenizer(Qwen3VLTokenizer):
|
||||
def __init__(self, embedding_directory=None, tokenizer_data={}):
|
||||
super().__init__(
|
||||
embedding_directory=embedding_directory, tokenizer_data=tokenizer_data,
|
||||
model_type="qwen3vl_8b",
|
||||
)
|
||||
self.llama_template = JOYIMAGE_TEMPLATE_TEXT
|
||||
self.llama_template_images = JOYIMAGE_TEMPLATE_IMAGE
|
||||
|
||||
def tokenize_with_weights(self, text, return_word_ids=False, llama_template=None, images=None, **kwargs):
|
||||
kwargs.pop("thinking", None)
|
||||
return super().tokenize_with_weights(
|
||||
text, return_word_ids=return_word_ids, llama_template=llama_template,
|
||||
images=images or [], thinking=True, **kwargs,
|
||||
)
|
||||
|
||||
|
||||
class _JoyImageClipModel(sd1_clip.SDClipModel):
|
||||
def __init__(self, device="cpu", layer="hidden", layer_idx=-1, dtype=None,
|
||||
attention_mask=True, model_options={}):
|
||||
super().__init__(
|
||||
device=device, layer=layer, layer_idx=layer_idx, textmodel_json_config={},
|
||||
# JoyImage conditions on the pre-final-norm output of the last decoder layer.
|
||||
dtype=dtype, special_tokens={"pad": PAD_TOKEN}, layer_norm_hidden_state=False,
|
||||
model_class=Qwen3VL8B_JoyImage, enable_attention_masks=attention_mask,
|
||||
return_attention_masks=attention_mask, model_options=model_options,
|
||||
)
|
||||
|
||||
|
||||
class JoyImageTEModel(sd1_clip.SD1ClipModel):
|
||||
def __init__(self, device="cpu", dtype=None, model_options={}):
|
||||
super().__init__(
|
||||
device=device, dtype=dtype, name="qwen3vl_8b",
|
||||
clip_model=_JoyImageClipModel, model_options=model_options,
|
||||
)
|
||||
|
||||
def encode_token_weights(self, token_weight_pairs):
|
||||
out, pooled, extra = super().encode_token_weights(token_weight_pairs)
|
||||
if out.shape[1] <= JOYIMAGE_DROP_IDX:
|
||||
raise ValueError(
|
||||
f"JoyImageTEModel: encoded sequence length {out.shape[1]} is shorter "
|
||||
f"than drop_idx={JOYIMAGE_DROP_IDX}; the prompt did not include the "
|
||||
f"template prefix."
|
||||
)
|
||||
out = out[:, JOYIMAGE_DROP_IDX:]
|
||||
if "attention_mask" in extra:
|
||||
extra["attention_mask"] = extra["attention_mask"][:, JOYIMAGE_DROP_IDX:]
|
||||
return out, pooled, extra
|
||||
|
||||
|
||||
def te(dtype_llama=None, llama_quantization_metadata=None):
|
||||
class JoyImageTEModel_(JoyImageTEModel):
|
||||
def __init__(self, device="cpu", dtype=None, model_options={}):
|
||||
if llama_quantization_metadata is not None:
|
||||
model_options = model_options.copy()
|
||||
model_options["quantization_metadata"] = llama_quantization_metadata
|
||||
if dtype_llama is not None:
|
||||
dtype = dtype_llama
|
||||
super().__init__(device=device, dtype=dtype, model_options=model_options)
|
||||
return JoyImageTEModel_
|
||||
@@ -876,7 +876,7 @@ class BaseGenerate:
|
||||
torch.empty([batch, model_config.num_key_value_heads, max_cache_len, model_config.head_dim], device=device, dtype=execution_dtype), 0))
|
||||
return past_key_values
|
||||
|
||||
def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None, embeds_info=None):
|
||||
def generate(self, embeds=None, do_sample=True, max_length=256, temperature=1.0, top_k=50, top_p=0.9, min_p=0.0, repetition_penalty=1.0, seed=42, stop_tokens=None, initial_tokens=[], execution_dtype=None, min_tokens=0, presence_penalty=0.0, initial_input_ids=None, position_ids=None, deepstack_embeds=None, visual_pos_masks=None):
|
||||
device = embeds.device
|
||||
|
||||
if stop_tokens is None:
|
||||
@@ -911,7 +911,7 @@ class BaseGenerate:
|
||||
if step == 0 and deepstack_embeds is not None:
|
||||
extra["deepstack_embeds"] = deepstack_embeds
|
||||
extra["visual_pos_masks"] = visual_pos_masks
|
||||
x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra, embeds_info=(embeds_info if step == 0 else None))
|
||||
x, _, past_key_values = self.model.forward(None, embeds=embeds, attention_mask=None, past_key_values=past_key_values, input_ids=current_input_ids, position_ids=position_ids, **extra)
|
||||
logits = self.logits(x)[:, -1]
|
||||
next_token = self.sample_token(logits, temperature, top_k, top_p, min_p, repetition_penalty, initial_tokens + generated_token_ids, generator, do_sample=do_sample, presence_penalty=presence_penalty)
|
||||
token_id = next_token[0].item()
|
||||
|
||||
@@ -15,7 +15,6 @@ def process_qwen2vl_images(
|
||||
merge_size: int = 2,
|
||||
image_mean: list = None,
|
||||
image_std: list = None,
|
||||
interpolation: str = "bilinear",
|
||||
):
|
||||
if image_mean is None:
|
||||
image_mean = [0.48145466, 0.4578275, 0.40821073]
|
||||
@@ -48,9 +47,10 @@ def process_qwen2vl_images(
|
||||
img_resized = F.interpolate(
|
||||
img.unsqueeze(0),
|
||||
size=(h_bar, w_bar),
|
||||
mode=interpolation,
|
||||
mode='bilinear',
|
||||
align_corners=False
|
||||
).squeeze(0)
|
||||
|
||||
normalized = img_resized.clone()
|
||||
for c in range(3):
|
||||
normalized[c] = (img_resized[c] - image_mean[c]) / image_std[c]
|
||||
|
||||
@@ -105,6 +105,7 @@ _CORE_FEATURE_FLAGS: dict[str, Any] = {
|
||||
"extension": {"manager": {"supports_v4": True}},
|
||||
"node_replacements": True,
|
||||
"assets": args.enable_assets,
|
||||
"server_side_model_downloads": True,
|
||||
}
|
||||
|
||||
# CLI-provided flags cannot overwrite core flags
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
from av.container import InputContainer
|
||||
from av.subtitles.stream import SubtitleStream
|
||||
from av.video.reformatter import ColorRange
|
||||
from fractions import Fraction
|
||||
from typing import Optional
|
||||
from .._input import AudioInput, VideoInput
|
||||
@@ -10,7 +9,6 @@ import itertools
|
||||
import json
|
||||
import numpy as np
|
||||
import math
|
||||
import os
|
||||
import torch
|
||||
from .._util import VideoContainer, VideoCodec, VideoComponents
|
||||
import logging
|
||||
@@ -60,57 +58,6 @@ def video_stream_bit_depth(stream) -> int:
|
||||
return max(component.bits for component in stream.format.components)
|
||||
|
||||
|
||||
def last_decodable_audio_stream(container: InputContainer):
|
||||
"""Streams FFmpeg has no decoder for have no codec context, and decoding their
|
||||
packets crashes the process (e.g. APAC spatial-audio track in iPhone)."""
|
||||
stream = next(
|
||||
(s for s in reversed(container.streams.audio) if s.codec_context is not None),
|
||||
None,
|
||||
)
|
||||
if stream is None and len(container.streams.audio):
|
||||
logging.warning("No decodable audio stream found in video; ignoring audio.")
|
||||
return stream
|
||||
|
||||
|
||||
def probe_audio_params(container: InputContainer, audio_stream, max_packets: int = 200):
|
||||
"""Containers probed only up to a window (mpegts) leave audio codec parameters unset when
|
||||
audio starts beyond it; learn them by decoding ahead. The caller must seek back afterwards.
|
||||
Returns (sample_rate, channels), zeros when the stream never yields a decodable frame."""
|
||||
for i, packet in enumerate(container.demux(audio_stream)):
|
||||
try:
|
||||
frames = packet.decode()
|
||||
except av.error.FFmpegError:
|
||||
frames = ()
|
||||
if frames:
|
||||
return frames[0].sample_rate, frames[0].layout.nb_channels
|
||||
if i >= max_packets:
|
||||
break
|
||||
return 0, 0
|
||||
|
||||
|
||||
def write_output_metadata(container: InputContainer, output, metadata: dict | None):
|
||||
"""Copy the source container's metadata, then overlay the caller's tags."""
|
||||
for key, value in container.metadata.items():
|
||||
if metadata is None or key not in metadata:
|
||||
output.metadata[key] = value
|
||||
if metadata is not None:
|
||||
for key, value in metadata.items():
|
||||
output.metadata[key] = value if isinstance(value, str) else json.dumps(value)
|
||||
|
||||
|
||||
def mp4_output_open_kwargs(path: str | io.BytesIO, format: VideoContainer, codec: VideoCodec) -> dict:
|
||||
if format != VideoContainer.AUTO and format != VideoContainer.MP4:
|
||||
raise ValueError("Only MP4 format is supported for now")
|
||||
if codec != VideoCodec.AUTO and codec != VideoCodec.H264:
|
||||
raise ValueError("Only H264 codec is supported for now")
|
||||
open_kwargs = {"mode": "w", "options": {"movflags": "use_metadata_tags"}}
|
||||
if isinstance(format, VideoContainer) and format != VideoContainer.AUTO:
|
||||
open_kwargs["format"] = format.value
|
||||
elif isinstance(path, io.BytesIO):
|
||||
open_kwargs["format"] = "mp4" # no file extension to infer the format from
|
||||
return open_kwargs
|
||||
|
||||
|
||||
class VideoFromFile(VideoInput):
|
||||
"""
|
||||
Class representing video input from a file.
|
||||
@@ -245,10 +192,13 @@ class VideoFromFile(VideoInput):
|
||||
return estimated_frames
|
||||
|
||||
# 3. Last resort: decode frames and count them (streaming)
|
||||
start_time, duration = self.get_active_trim_window()
|
||||
if self.__start_time < 0:
|
||||
start_time = max(self._get_raw_duration() + self.__start_time, 0)
|
||||
else:
|
||||
start_time = self.__start_time
|
||||
frame_count = 1
|
||||
start_pts = int(start_time / video_stream.time_base)
|
||||
end_pts = int((start_time + duration) / video_stream.time_base)
|
||||
end_pts = int((start_time + self.__duration) / video_stream.time_base)
|
||||
container.seek(start_pts, stream=video_stream)
|
||||
frame_iterator = (
|
||||
container.decode(video_stream)
|
||||
@@ -303,14 +253,17 @@ class VideoFromFile(VideoInput):
|
||||
|
||||
def get_components_internal(self, container: InputContainer) -> VideoComponents:
|
||||
video_stream = self._get_first_video_stream(container)
|
||||
start_time, duration = self.get_active_trim_window()
|
||||
if self.__start_time < 0:
|
||||
start_time = max(self._get_raw_duration() + self.__start_time, 0)
|
||||
else:
|
||||
start_time = self.__start_time
|
||||
|
||||
# Get video frames
|
||||
frames = []
|
||||
audio_frames = []
|
||||
alphas = None
|
||||
start_pts = int(start_time / video_stream.time_base)
|
||||
end_pts = int((start_time + duration) / video_stream.time_base)
|
||||
end_pts = int((start_time + self.__duration) / video_stream.time_base)
|
||||
|
||||
if start_pts != 0:
|
||||
container.seek(start_pts, stream=video_stream)
|
||||
@@ -328,11 +281,18 @@ class VideoFromFile(VideoInput):
|
||||
video_done = False
|
||||
audio_done = True
|
||||
|
||||
audio_stream = last_decodable_audio_stream(container)
|
||||
# Use the last decodable audio stream. Streams FFmpeg has no decoder for have no codec context,
|
||||
# and decoding their packets crashes the process. (e.g. APAC spatial-audio track in iPhone)
|
||||
audio_stream = next(
|
||||
(s for s in reversed(container.streams.audio) if s.codec_context is not None),
|
||||
None,
|
||||
)
|
||||
if audio_stream is not None:
|
||||
streams += [audio_stream]
|
||||
resampler = av.audio.resampler.AudioResampler(format='fltp')
|
||||
audio_done = False
|
||||
elif len(container.streams.audio):
|
||||
logging.warning("No decodable audio stream found in video; ignoring audio.")
|
||||
|
||||
for packet in container.demux(*streams):
|
||||
if video_done and audio_done:
|
||||
@@ -345,7 +305,7 @@ class VideoFromFile(VideoInput):
|
||||
for frame in packet.decode():
|
||||
if frame.pts < start_pts:
|
||||
continue
|
||||
if duration and frame.pts >= end_pts:
|
||||
if self.__duration and frame.pts >= end_pts:
|
||||
video_done = True
|
||||
break
|
||||
|
||||
@@ -412,7 +372,7 @@ class VideoFromFile(VideoInput):
|
||||
map(resampler.resample, packet.decode())
|
||||
)
|
||||
for frame in aframes:
|
||||
if duration and frame.time > start_time + duration:
|
||||
if self.__duration and frame.time > start_time + self.__duration:
|
||||
audio_done = True
|
||||
break
|
||||
|
||||
@@ -434,8 +394,8 @@ class VideoFromFile(VideoInput):
|
||||
|
||||
if len(audio_frames) > 0:
|
||||
audio_data = np.concatenate(audio_frames, axis=1) # shape: (channels, total_samples)
|
||||
if duration:
|
||||
audio_data = audio_data[..., :int(duration * audio_stream.sample_rate)]
|
||||
if self.__duration:
|
||||
audio_data = audio_data[..., :int(self.__duration * audio_stream.sample_rate)]
|
||||
|
||||
audio_tensor = torch.from_numpy(audio_data).unsqueeze(0) # shape: (1, channels, total_samples)
|
||||
audio = AudioInput({
|
||||
@@ -481,14 +441,28 @@ class VideoFromFile(VideoInput):
|
||||
if not reuse_streams:
|
||||
if bit_depth is None:
|
||||
bit_depth = source_bit_depth
|
||||
return self._save_transcoded(container, path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth)
|
||||
components = self.get_components_internal(container)
|
||||
video = VideoFromComponents(components)
|
||||
return video.save_to(
|
||||
path, format=format, codec=codec, metadata=metadata, bit_depth=bit_depth,
|
||||
)
|
||||
|
||||
streams = container.streams
|
||||
|
||||
open_kwargs = get_open_write_kwargs(path, container_format, format)
|
||||
with av.open(path, **open_kwargs) as output_container:
|
||||
# Add metadata before writing any streams
|
||||
write_output_metadata(container, output_container, metadata)
|
||||
# Copy over the original metadata
|
||||
for key, value in container.metadata.items():
|
||||
if metadata is None or key not in metadata:
|
||||
output_container.metadata[key] = value
|
||||
|
||||
# Add our new metadata
|
||||
if metadata is not None:
|
||||
for key, value in metadata.items():
|
||||
if isinstance(value, str):
|
||||
output_container.metadata[key] = value
|
||||
else:
|
||||
output_container.metadata[key] = json.dumps(value)
|
||||
|
||||
# Add streams to the new container. Streams with no codec context cannot be used as an output template.
|
||||
stream_map = {}
|
||||
@@ -506,282 +480,6 @@ class VideoFromFile(VideoInput):
|
||||
packet.stream = stream_map[packet.stream]
|
||||
output_container.mux(packet)
|
||||
|
||||
def _save_transcoded(
|
||||
self,
|
||||
container: InputContainer,
|
||||
path: str | io.BytesIO,
|
||||
format: VideoContainer,
|
||||
codec: VideoCodec,
|
||||
metadata: dict | None,
|
||||
bit_depth: int,
|
||||
):
|
||||
"""Re-encode to H.264/AAC one frame at a time; peak memory does not scale with video length."""
|
||||
open_kwargs = mp4_output_open_kwargs(path, format, codec)
|
||||
video_stream = self._get_first_video_stream(container)
|
||||
start_time, duration = self.get_active_trim_window()
|
||||
start_pts = int(start_time / video_stream.time_base)
|
||||
end_pts = int((start_time + duration) / video_stream.time_base) if duration else None
|
||||
stream_end_pts = None
|
||||
if video_stream.duration is not None:
|
||||
stream_end_pts = (video_stream.start_time or 0) + video_stream.duration
|
||||
output_end_pts = end_pts
|
||||
if stream_end_pts is not None and (output_end_pts is None or stream_end_pts < output_end_pts):
|
||||
output_end_pts = stream_end_pts
|
||||
if start_pts != 0:
|
||||
container.seek(start_pts, stream=video_stream)
|
||||
|
||||
audio_stream = last_decodable_audio_stream(container)
|
||||
pix_fmt = "yuv420p10le" if bit_depth >= 10 else "yuv420p"
|
||||
rate = Fraction(video_stream.average_rate) if video_stream.average_rate else Fraction(1)
|
||||
|
||||
resampler = None
|
||||
sample_rate = 0
|
||||
audio_time_base = None
|
||||
duration_cap = None
|
||||
if audio_stream is not None:
|
||||
sample_rate = audio_stream.codec_context.sample_rate
|
||||
channels = audio_stream.codec_context.channels
|
||||
if not sample_rate:
|
||||
sample_rate, channels = probe_audio_params(container, audio_stream)
|
||||
container.seek(start_pts, stream=video_stream)
|
||||
if sample_rate:
|
||||
audio_stream.codec_context.flush_buffers()
|
||||
else:
|
||||
logging.warning("Audio stream parameters could not be determined; ignoring audio.")
|
||||
audio_stream = None
|
||||
if audio_stream is not None:
|
||||
audio_time_base = Fraction(1, sample_rate)
|
||||
layout = {1: "mono", 2: "stereo", 6: "5.1"}.get(channels, "stereo")
|
||||
resampler = av.audio.resampler.AudioResampler(format="fltp", layout=layout, rate=sample_rate)
|
||||
if duration:
|
||||
duration_cap = math.ceil(duration * sample_rate)
|
||||
|
||||
streams = [video_stream] if audio_stream is None else [video_stream, audio_stream]
|
||||
pts_step = max(1, int(round((1 / rate) / video_stream.time_base)))
|
||||
video_done = False
|
||||
audio_done = audio_stream is None
|
||||
video_pts_offset = None
|
||||
last_video_pts = None
|
||||
last_video_end = None
|
||||
# rebased pts -> true display duration: the mp4 muxer pads the last sample with 1/rate otherwise
|
||||
video_frame_durations = {}
|
||||
source_size = None
|
||||
rotation_k = 0
|
||||
rotation_filter = None
|
||||
audio_started = False
|
||||
samples_written = 0
|
||||
pending_audio = []
|
||||
# The output opens lazily on the first kept frame: it decides the geometry (90/270 rotation swaps dims),
|
||||
# and never seeking back keeps webm/mkv leading audio intact.
|
||||
output = None
|
||||
out_video = None
|
||||
out_audio = None
|
||||
|
||||
def audio_frame_from_ndarray(nd_planar):
|
||||
frame = av.AudioFrame.from_ndarray(np.ascontiguousarray(nd_planar), format="fltp", layout=layout)
|
||||
frame.sample_rate = sample_rate
|
||||
return frame
|
||||
|
||||
def drain_audio(final=False):
|
||||
# Audio may cover the pts span of the video written so far, capped by the requested duration
|
||||
nonlocal samples_written, audio_done
|
||||
if last_video_end is None:
|
||||
cap = 0
|
||||
else:
|
||||
cap = math.ceil(last_video_end * video_stream.time_base * sample_rate)
|
||||
if duration_cap is not None:
|
||||
cap = min(cap, duration_cap)
|
||||
while pending_audio and not audio_done:
|
||||
frame = pending_audio[0]
|
||||
if samples_written + frame.samples <= cap:
|
||||
frame.pts = samples_written
|
||||
frame.time_base = audio_time_base
|
||||
output.mux(out_audio.encode(frame))
|
||||
samples_written += frame.samples
|
||||
pending_audio.pop(0)
|
||||
continue
|
||||
if final:
|
||||
keep = frame.to_ndarray()[..., :cap - samples_written]
|
||||
if keep.shape[-1] > 0:
|
||||
tail = audio_frame_from_ndarray(keep)
|
||||
tail.pts = samples_written
|
||||
tail.time_base = audio_time_base
|
||||
output.mux(out_audio.encode(tail))
|
||||
samples_written += keep.shape[-1]
|
||||
pending_audio.clear()
|
||||
break
|
||||
if duration_cap is not None and samples_written >= duration_cap:
|
||||
audio_done = True
|
||||
return cap
|
||||
|
||||
try:
|
||||
for packet in container.demux(*streams):
|
||||
if video_done and audio_done:
|
||||
break
|
||||
|
||||
if packet.stream == video_stream and not video_done:
|
||||
try:
|
||||
frames = packet.decode()
|
||||
except av.error.InvalidDataError:
|
||||
logging.info("pyav decode error")
|
||||
continue
|
||||
for frame in frames:
|
||||
if frame.pts is not None and frame.pts < start_pts:
|
||||
continue
|
||||
if end_pts is not None and frame.pts is not None and frame.pts >= end_pts:
|
||||
video_done = True
|
||||
if last_video_pts is not None:
|
||||
# the source continues past the window: hold the last kept frame to the window end
|
||||
end_offset = video_pts_offset if video_pts_offset is not None else start_pts
|
||||
last_video_end = max(last_video_end, end_pts - end_offset)
|
||||
break
|
||||
# the source's true display duration of this frame; average_rate is not a
|
||||
# frame duration (sparse/VFR sources), so it is only the fallback
|
||||
frame_duration = frame.duration if frame.duration else pts_step
|
||||
if end_pts is not None and frame.pts is not None:
|
||||
frame_duration = min(frame_duration, end_pts - frame.pts)
|
||||
if output is None:
|
||||
rotation_k = int(round(frame.rotation // 90)) % 4 if frame.rotation else 0
|
||||
if rotation_k % 2:
|
||||
out_width, out_height = frame.height, frame.width
|
||||
else:
|
||||
out_width, out_height = frame.width, frame.height
|
||||
if out_width % 2 or out_height % 2:
|
||||
raise ValueError(f"H.264 output requires even dimensions, got {out_width}x{out_height}")
|
||||
source_size = (frame.width, frame.height)
|
||||
output = av.open(path, **open_kwargs)
|
||||
# Add metadata before writing any streams
|
||||
write_output_metadata(container, output, metadata)
|
||||
out_video = output.add_stream("h264", rate=rate)
|
||||
# no B-frames: reordering makes mp4 sample durations follow decode order,
|
||||
# so irregular-VFR spans and trim windows land wrong
|
||||
out_video.codec_context.max_b_frames = 0
|
||||
out_video.width = out_width
|
||||
out_video.height = out_height
|
||||
out_video.pix_fmt = pix_fmt
|
||||
# source pts pass through (rebased to 0), so variable frame rate survives
|
||||
out_video.codec_context.time_base = video_stream.time_base
|
||||
if audio_stream is not None:
|
||||
out_audio = output.add_stream("aac", rate=sample_rate, layout=layout)
|
||||
if (frame.width, frame.height) != source_size:
|
||||
# encoding would silently rescale the new geometry into the old one
|
||||
raise ValueError(
|
||||
f"Video resolution changes mid-stream "
|
||||
f"({source_size[0]}x{source_size[1]} -> {frame.width}x{frame.height}); cannot transcode"
|
||||
)
|
||||
if rotation_k:
|
||||
if rotation_filter is None:
|
||||
g = av.filter.Graph()
|
||||
g_src = g.add_buffer(width=frame.width, height=frame.height,
|
||||
format=frame.format.name, time_base=video_stream.time_base)
|
||||
tail = g_src
|
||||
for filter_name, filter_args in {1: [("transpose", "cclock")],
|
||||
2: [("hflip", None), ("vflip", None)],
|
||||
3: [("transpose", "clock")]}[rotation_k]:
|
||||
step = g.add(filter_name, filter_args)
|
||||
tail.link_to(step)
|
||||
tail = step
|
||||
g_sink = g.add("buffersink")
|
||||
tail.link_to(g_sink)
|
||||
g.configure()
|
||||
rotation_filter = (g_src, g_sink)
|
||||
rotation_filter[0].push(frame)
|
||||
frame = rotation_filter[1].pull()
|
||||
if frame.color_range == ColorRange.JPEG:
|
||||
# compress full-range sources (yuvj/MJPEG) to limited range
|
||||
frame = frame.reformat(format=pix_fmt, src_color_range="JPEG", dst_color_range="MPEG")
|
||||
else:
|
||||
frame = frame.reformat(format=pix_fmt)
|
||||
frame_output_end = None
|
||||
if frame.pts is not None:
|
||||
if video_pts_offset is None:
|
||||
video_pts_offset = frame.pts
|
||||
frame.pts -= video_pts_offset
|
||||
if output_end_pts is not None:
|
||||
frame_output_end = output_end_pts - video_pts_offset
|
||||
if frame.pts + frame_duration > frame_output_end:
|
||||
clamped_pts = frame_output_end - frame_duration
|
||||
if clamped_pts >= 0 and (last_video_pts is None or clamped_pts > last_video_pts):
|
||||
frame.pts = min(frame.pts, clamped_pts)
|
||||
elif frame.pts < frame_output_end:
|
||||
frame_duration = frame_output_end - frame.pts
|
||||
else:
|
||||
continue
|
||||
if frame.pts is None or (last_video_pts is not None and frame.pts <= last_video_pts):
|
||||
# broken sources emit missing/backward timestamps mid-stream, which the
|
||||
# muxer rejects; nudge them forward by one nominal frame interval
|
||||
frame.pts = 0 if last_video_pts is None else last_video_pts + pts_step
|
||||
if frame_output_end is not None and frame.pts + frame_duration > frame_output_end:
|
||||
if frame.pts >= frame_output_end:
|
||||
continue
|
||||
frame_duration = frame_output_end - frame.pts
|
||||
last_video_pts = frame.pts
|
||||
last_video_end = frame.pts + frame_duration
|
||||
video_frame_durations[frame.pts] = frame_duration
|
||||
# the decoded pict_type would force x264's frame types (intra-only
|
||||
# sources like MJPEG/ProRes would come out all-keyframe)
|
||||
frame.pict_type = 0
|
||||
for out_packet in out_video.encode(frame):
|
||||
out_packet.duration = video_frame_durations.pop(out_packet.pts, 0)
|
||||
output.mux(out_packet)
|
||||
drain_audio()
|
||||
|
||||
elif packet.stream == audio_stream and not audio_done:
|
||||
for resampled in itertools.chain.from_iterable(map(resampler.resample, packet.decode())):
|
||||
frame_start = None
|
||||
if resampled.pts is not None:
|
||||
# passthrough frames keep the source stream's time base
|
||||
tb = resampled.time_base if resampled.time_base else audio_time_base
|
||||
frame_start = float(resampled.pts * tb)
|
||||
if duration and not audio_started and frame_start >= start_time + duration:
|
||||
audio_done = True
|
||||
break
|
||||
if not audio_started:
|
||||
if frame_start is None:
|
||||
frame_start = 0.0
|
||||
to_skip = max(0, int((start_time - frame_start) * sample_rate))
|
||||
if to_skip >= resampled.samples:
|
||||
continue
|
||||
audio_started = True
|
||||
if duration and frame_start > start_time:
|
||||
duration_cap = min(duration_cap, math.ceil((start_time + duration - frame_start) * sample_rate))
|
||||
if to_skip:
|
||||
pending_audio.append(audio_frame_from_ndarray(resampled.to_ndarray()[..., to_skip:]))
|
||||
continue
|
||||
pending_audio.append(resampled)
|
||||
if video_done:
|
||||
# the video window is complete so the cap is final, but containers
|
||||
# that interleave audio behind video (fragmented mp4) still owe most
|
||||
# of it: stop only once the demuxed audio covers the cap
|
||||
cap = drain_audio()
|
||||
if pending_audio or samples_written >= cap:
|
||||
drain_audio(final=True)
|
||||
audio_done = True
|
||||
break
|
||||
|
||||
if output is None:
|
||||
raise ValueError(f"No decodable video frames found in file '{self.__file}'")
|
||||
if out_audio is not None and not audio_done:
|
||||
drain_audio(final=True)
|
||||
window_fill = last_video_end - last_video_pts if video_done and last_video_pts is not None else 0
|
||||
for out_packet in out_video.encode(None):
|
||||
duration = video_frame_durations.pop(out_packet.pts, 0)
|
||||
if out_packet.pts == last_video_pts:
|
||||
duration = max(duration, window_fill)
|
||||
out_packet.duration = duration
|
||||
output.mux(out_packet)
|
||||
if out_audio is not None:
|
||||
output.mux(out_audio.encode(None))
|
||||
except BaseException:
|
||||
if output is not None:
|
||||
output.close()
|
||||
if isinstance(path, (str, os.PathLike)) and os.path.exists(path):
|
||||
os.remove(path)
|
||||
raise
|
||||
else:
|
||||
if output is not None:
|
||||
output.close()
|
||||
|
||||
def _get_first_video_stream(self, container: InputContainer):
|
||||
if len(container.streams.video):
|
||||
return container.streams.video[0]
|
||||
@@ -829,12 +527,22 @@ class VideoFromComponents(VideoInput):
|
||||
bit_depth: int | None = None,
|
||||
):
|
||||
"""Save the video to a file path or BytesIO buffer."""
|
||||
open_kwargs = mp4_output_open_kwargs(path, format, codec)
|
||||
if format != VideoContainer.AUTO and format != VideoContainer.MP4:
|
||||
raise ValueError("Only MP4 format is supported for now")
|
||||
if codec != VideoCodec.AUTO and codec != VideoCodec.H264:
|
||||
raise ValueError("Only H264 codec is supported for now")
|
||||
# None means "use the depth this video was created with" (CreateVideo's choice).
|
||||
if bit_depth is None:
|
||||
bit_depth = self.__bit_depth
|
||||
is_10bit = bit_depth >= 10
|
||||
with av.open(path, **open_kwargs) as output:
|
||||
extra_kwargs = {}
|
||||
if isinstance(format, VideoContainer) and format != VideoContainer.AUTO:
|
||||
extra_kwargs["format"] = format.value
|
||||
elif isinstance(path, io.BytesIO):
|
||||
# BytesIO has no file extension, so av.open can't infer the format.
|
||||
# Default to mp4 since that's the only supported format anyway.
|
||||
extra_kwargs["format"] = "mp4"
|
||||
with av.open(path, mode='w', options={'movflags': 'use_metadata_tags'}, **extra_kwargs) as output:
|
||||
# Add metadata before writing any streams
|
||||
if metadata is not None:
|
||||
for key, value in metadata.items():
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from datetime import date
|
||||
from enum import Enum
|
||||
from typing import Any, Literal
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
@@ -242,60 +242,3 @@ class GeminiGenerateContentResponse(BaseModel):
|
||||
promptFeedback: GeminiPromptFeedback | None = Field(None)
|
||||
usageMetadata: GeminiUsageMetadata | None = Field(None)
|
||||
modelVersion: str | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionTextPart(BaseModel):
|
||||
type: Literal["text"] = "text"
|
||||
text: str = Field(...)
|
||||
|
||||
|
||||
class GeminiInteractionMediaPart(BaseModel):
|
||||
type: str = Field(..., description="One of: image, video, audio, document.")
|
||||
data: str | None = Field(None, description="Base64-encoded media bytes.")
|
||||
uri: str | None = Field(None, description="URI of the media, as an alternative to inline data.")
|
||||
mime_type: str | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionGenerationConfig(BaseModel):
|
||||
temperature: float | None = Field(None, ge=0.0, le=2.0)
|
||||
top_p: float | None = Field(None, ge=0.0, le=1.0)
|
||||
|
||||
|
||||
class GeminiInteractionRequest(BaseModel):
|
||||
model: str = Field(...)
|
||||
input: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = Field(...)
|
||||
generation_config: GeminiInteractionGenerationConfig | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionModalityTokens(BaseModel):
|
||||
modality: str | None = Field(None, description="One of: text, image, audio, video, document.")
|
||||
tokens: int | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionUsage(BaseModel):
|
||||
input_tokens_by_modality: list[GeminiInteractionModalityTokens] | None = Field(None)
|
||||
output_tokens_by_modality: list[GeminiInteractionModalityTokens] | None = Field(None)
|
||||
total_thought_tokens: int | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionContent(BaseModel):
|
||||
type: str | None = Field(None)
|
||||
text: str | None = Field(None)
|
||||
data: str | None = Field(None)
|
||||
uri: str | None = Field(None)
|
||||
mime_type: str | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteractionStep(BaseModel):
|
||||
type: str | None = Field(None)
|
||||
content: list[GeminiInteractionContent] | None = Field(None)
|
||||
|
||||
|
||||
class GeminiInteraction(BaseModel):
|
||||
id: str | None = Field(None)
|
||||
status: str | None = Field(
|
||||
None,
|
||||
description="One of: in_progress, requires_action, completed, failed, cancelled, incomplete.",
|
||||
)
|
||||
steps: list[GeminiInteractionStep] | None = Field(None)
|
||||
usage: GeminiInteractionUsage | None = Field(None)
|
||||
|
||||
@@ -1,452 +0,0 @@
|
||||
# (label, avatar_id, avatar_type, supported engines)
|
||||
HEYGEN_AVATAR_LOOKS: list[tuple[str, str, str, tuple[str, ...]]] = [
|
||||
(
|
||||
"Annie Lounge Standing Side",
|
||||
"Annie_Lounge_Standing_Side_public",
|
||||
"studio_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Yara Modern Lecture Hall",
|
||||
"fd6814ecc5e143cd899e615a80eaa2dc",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Brandon Business Sitting Front",
|
||||
"Brandon_Business_Sitting_Front_public",
|
||||
"studio_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Caroline Business Sitting Side",
|
||||
"Caroline_Business_Sitting_Side_public",
|
||||
"studio_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Ursula Lawyer Angle 4",
|
||||
"f7173d2bb8584c00bfec6905c5e9a492",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Sofia Corporate Presenter 01 Angle 3",
|
||||
"fe563971fd2d438e957372dac9e2be8c",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Seoyeon Health Nutrition Coach Angle 3",
|
||||
"fe3c5d5028d941398d064b8fc64a2dea",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Sanne Fitness Coach Angle 4",
|
||||
"d967f935a8bf4a0c8f0bccfd66c501d2",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
("Sander", "f5cd7b94056f495ca0610602d64a9aa3", "photo_avatar", ("avatar_v", "avatar_iv", "avatar_iii")),
|
||||
(
|
||||
"Rupert Personal Development Coach Angle 4",
|
||||
"f57b3e626adb4bc997b38f64884adce4",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Olivier Professor Angle 2",
|
||||
"f6659bbb094b459c87c967edbb9ee481",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Obi Health Nutrition Coach Angle 5",
|
||||
"f3dc2c38201d414382f506d2d8e8d029",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Matilda Modern Office Setting",
|
||||
"fda889ac354a440da8dbecc410981273",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Mateo Traditional Law Office",
|
||||
"ff172d6c499c4e47ba6fcc5de631e9fc",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Marlon Inviting Armchair Setting",
|
||||
"f5a57db099ab462daa3e7c604a05dacc",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Margaret Professor Angle 1",
|
||||
"fb472bc29ab04bcca576e3703978fecb",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Marek Therapy Coach Angle 3",
|
||||
"e197768703f1463a93dc25ada1f421fb",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Maeve Warm, Professional Setting",
|
||||
"faf66681d8cc48dc82c4283200b3e782",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
(
|
||||
"Lorenzo Professor Angle 5",
|
||||
"fc268dc244bb40d7a554663ce723dcf0",
|
||||
"photo_avatar",
|
||||
("avatar_v", "avatar_iv", "avatar_iii"),
|
||||
),
|
||||
("Luca", "Luca_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Bruce", "Bruce_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Nico", "Nico_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Lisa", "Lisa_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Sophie", "Sophie_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Aiko", "Aiko_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Rebecca (portrait)", "Rebecca_public", "studio_avatar", ("avatar_iii",)),
|
||||
("Daphne in Grey blazer (portrait)", "Daphne_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Bryce in Black t-shirt", "Bryce_public_5", "studio_avatar", ("avatar_iii",)),
|
||||
("Diora in White shirt", "Diora_public_3", "studio_avatar", ("avatar_iii",)),
|
||||
("Freja in White blazer", "Freja_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Albert in Blue blazer", "Albert_public_2", "studio_avatar", ("avatar_iii",)),
|
||||
("Emery in Red blazer", "Emery_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Minho in Blue shirt", "Minho_public_6", "studio_avatar", ("avatar_iii",)),
|
||||
("Aditya in Brown blazer", "Aditya_public_4", "studio_avatar", ("avatar_iii",)),
|
||||
("Nadim in Blue blazer", "Nadim_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Iker in Black blazer", "Iker_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Nour in Black blazer", "Nour_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Saskia in Blue blazer", "Saskia_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Lucien in Blue blazer", "Lucien_public_1", "studio_avatar", ("avatar_iii",)),
|
||||
("Esmond in Blue suit", "Esmond_public_3", "studio_avatar", ("avatar_iii",)),
|
||||
("Jinwoo in Blue suit", "Jinwoo_public_5", "studio_avatar", ("avatar_iii",)),
|
||||
("Annelore in Red sweater (portrait)", "Annelore_public_3", "studio_avatar", ("avatar_iii",)),
|
||||
("Bastien in Blue shirt", "Bastien_public_4", "studio_avatar", ("avatar_iii",)),
|
||||
("Zosia in Khaki blazer", "Zosia_public_3", "studio_avatar", ("avatar_iii",)),
|
||||
("Tahlia in Dark blue suit", "Tahlia_public_4", "studio_avatar", ("avatar_iii",)),
|
||||
]
|
||||
HEYGEN_AVATAR_OPTIONS = [x[0] for x in HEYGEN_AVATAR_LOOKS]
|
||||
HEYGEN_AVATAR_MAP = {x[0]: (x[1], x[2], x[3]) for x in HEYGEN_AVATAR_LOOKS}
|
||||
|
||||
# (label, voice_id) — Starfish-compatible voices for the TTS endpoint
|
||||
HEYGEN_VOICE_TTS: list[tuple[str, str]] = [
|
||||
("Chill Brian (English, male)", "d2f4f24783d04e22ab49ee8fdc3715e0"),
|
||||
("Zain (English, female)", "0047732240584155b1588455313e78ec"),
|
||||
("Narrator Mateo - Excited 🤩 (Spanish, male)", "0077225a877e457db4572ccaf245910b"),
|
||||
("Aria (English, female)", "007e1378fc454a9f976db570ba6164a7"),
|
||||
("Caryns (English, female)", "0082e70326864107823605db0d77f5e0"),
|
||||
("Klara (English, female)", "01209fdcd1c24a109c86dc24ee0f71c0"),
|
||||
("Bold Kasia - Excited 🤩 (Polish, female)", "015482a78b9a46ebae74bd0beb17765b"),
|
||||
("Shaun (English, male)", "01c42cddcfdc4665a57b8d89cba8ffc1"),
|
||||
("Senthil (English, male)", "01d674cfd32b4728a3fddd21b7e7d543"),
|
||||
("Cody (English, male)", "01f98ed43e6140349f47dbd37a416827"),
|
||||
("Saffron (English, female)", "0258bbc2cd8648cfa357adfb833f6d7b"),
|
||||
("Blanka - Lifelike (English, female)", "02880d1c6fd94b7799d91135581ed810"),
|
||||
("Rami (English, male)", "02d5366a90af4c7a87157808ff352e33"),
|
||||
("Rhodes (English, female)", "02dce0a169b3460084b6c914d18fb2c8"),
|
||||
("Michelle - Voice 1 (English, female)", "02X8sHnuxFpsq1caYWN0"),
|
||||
("Autumn - UGC 3 (English, female)", "03dca9ebfca441dba55fb14afa0791b7"),
|
||||
("Reassuring Rupert (English, male)", "03fcf8ecb0a94b6b94e9007edb7c35f8"),
|
||||
("Rose - UGC -2 (English, female)", "0495e14c2bd74eb3aeeef03583e0bce5"),
|
||||
("Derya - Lifelike - Broadcaster 🎙️ (English, female)", "04d0ae1d0af2489ca7d3bb402a39a890"),
|
||||
("Dynamic Derek (English, male)", "0516c2d857eb425c94e90b068241914e"),
|
||||
("Lotte (English, female)", "052fcfb83d1a4c2f8d0368c226fea4b9"),
|
||||
("Thanos - Broadcaster 🎙️ (English, male)", "054af44a167344d0af2722fdfef08d17"),
|
||||
("Marcia (English, female)", "05f19352e8f74b0392a8f411eba40de1"),
|
||||
("Camden (English, male)", "06468055edd4458aa131a1dfd813c1e9"),
|
||||
("Rumi (English, female)", "06672207805f41a9ad0af6797f8aa14b"),
|
||||
("Pippa (English, female)", "06b68c4dbb544935b9af984e80efa4fb"),
|
||||
("William Prescott - Broadcaster 🎙️ (English, male)", "06c816b952f14fa9b3a6c42aa151f731"),
|
||||
("Sammy (English, female)", "06e6facd99654b9dbb9308f67bf3a31c"),
|
||||
("Breezy Bagus (Indonesian, male)", "06e81a5d7c8b41818d3f0b38f7cf15a1"),
|
||||
("Ben (English, male)", "07ca39b243184dbcb82e7e0f0e524b21"),
|
||||
("Smooth Dev (English, male)", "07d2ba65847541feb97abc9b60181555"),
|
||||
("Daran inside booth (English, male)", "080d9383c0314056aef392892e009806"),
|
||||
("Peppy Stella (English, female)", "084760b4922a44599575c770070ec2d7"),
|
||||
("Silas (English, male)", "08f561403ec846dbbd8c691cc448f45a"),
|
||||
("Aditya (English, male)", "09c3d65e44e247dd8b78a97a903feb58"),
|
||||
("Christy (English, female)", "09d88c036bf449fa905900c08b235a37"),
|
||||
("Elio (English, male)", "0a0b38624ac64ec6afcd5842a977ca10"),
|
||||
("Luminous Laksh (Hindi, male)", "0adc547b76a5401c856274c379904eb7"),
|
||||
("Jeff (English, male)", "0add542e349f4ccaba6ecb3b7ced6034"),
|
||||
("Tahlia Brooks - Excited 🤩 (English, female)", "0b440d1ac2454d69a73302fc806522b1"),
|
||||
("Riya Mehta (Hindi, female)", "0b464b2f4e2249a4b5a05e60eaf41e7e"),
|
||||
("Ben Hart (English, male)", "0b47b5a637e944f9bfd49913999b344b"),
|
||||
("Skylar (English, female)", "0bbfbda5aa924a68a9d1da7b8496052a"),
|
||||
("Relaxed Reece (English, male)", "0c2151d538844c70a8b096de533f2828"),
|
||||
("Daniel (English, male)", "0c23804af39a4946ac6fda42bfff2738"),
|
||||
("Melani (English, female)", "0c54c6399ad64551a304e1a346677723"),
|
||||
("Clover (English, female)", "0ccb0bea067d4449ad367baeed7ea2e9"),
|
||||
("Pedro Lima - Serious 😐 (Portuguese, male)", "0d0e23e8170446e38b18a7380b2d30a8"),
|
||||
("Ana Carvalho (Portuguese, female)", "0d23c5b2f6004e909802a2e8bfcd52c2"),
|
||||
("Confident Connor - Excited 🤩 (English, male)", "0dd34c3eb79247238219eea35aeb58cd"),
|
||||
("Vibrant Victor (Spanish, male)", "1062976ea8bf42f4adc27c7e868b8fde"),
|
||||
("Young Olivier (French, male)", "1c5dc9a8f8cf4de0932f91d75f43a15d"),
|
||||
("Émile Noir (French, male)", "25a6a67280574d3da78e97b1935ebfc7"),
|
||||
("Steadfast Stefan (German, male)", "0eb85e6e8710473b82f7e88609ba3053"),
|
||||
("Deep Dieter (German, male)", "118949676b0a46629d1ad52981c3ef84"),
|
||||
("Serene Marco (Italian, male)", "72e922488a614041b5ab5f6ee07e3deb"),
|
||||
("Murmuring Matteo (Italian, male)", "755902b751654f30a6ef49e8bbcacfec"),
|
||||
("Gail in car (Multilingual, female)", "0214ac51f93e420f8711d568dcfbc50e"),
|
||||
("Daran outside walking (Multilingual, male)", "0ac81e725f4948dfa9638ceca216bcfa"),
|
||||
("BOB - Voice 1 (Chinese, unknown)", "dMkR1XwIkarpNqWUJLnX"),
|
||||
("Hakeem Hassan (Arabic, male)", "61a4359785664d01a59664ceb87ce6d4"),
|
||||
("Rami Idris (Arabic, male)", "a0bd2e5d41a74643be47ac75ca9171a2"),
|
||||
("Bold Kasia - Friendly 😊 (Polish, female)", "331624aec8b24a6c9287b8e16bdf54e8"),
|
||||
("Tranquil Tulin (Turkish, female)", "61646c861eb64e2d9036d8db51385356"),
|
||||
("Dynamic Derya (Turkish, female)", "664b73058b784aa89ddb2924c141d441"),
|
||||
("Quiet Dewa (Indonesian, male)", "1fa1193cf1d74f27ba58531c07ef9862"),
|
||||
("Cuong (Vietnamese, male)", "8af68d7ea38f4e7ca05cf46c3f7a590b"),
|
||||
]
|
||||
HEYGEN_VOICE_TTS_OPTIONS = [x[0] for x in HEYGEN_VOICE_TTS]
|
||||
HEYGEN_VOICE_TTS_MAP = dict(HEYGEN_VOICE_TTS)
|
||||
|
||||
# (label, voice_id) — top-ranked voices for video narration (any engine)
|
||||
HEYGEN_VOICE_GENERAL: list[tuple[str, str]] = [
|
||||
("Cassidy (English, female)", "16a09e4706f74997ba4ed05ea11470f6"),
|
||||
("Hope (English, female)", "42d00d4aac5441279d8536cd6b52c53c"),
|
||||
("Archer (English, male)", "453c20e1525a429080e2ad9e4b26f2cd"),
|
||||
("Brittney (English, female)", "4754e1ec667544b0bd18cdf4bec7d6a7"),
|
||||
("Mark (English, male)", "5d8c378ba8c3434586081a52ac368738"),
|
||||
("Andrew (English, male)", "6be73833ef9a4eb0aeee399b8fe9d62b"),
|
||||
("Spuds Oxley (English, male)", "76940a9adcd0490a9ce2cfe9a64a2664"),
|
||||
("Patrick (English, male)", "7e157ec62c9c45f1adca12faae72c86f"),
|
||||
("David Castlemore (English, male)", "828b59f834fd4c7188da322b6d9b6c75"),
|
||||
("Michael C (English, male)", "8661cd40d6c44c709e2d0031c0186ada"),
|
||||
("Adam Stone (English, male)", "88bb9ee1c81b466eb2a08fdde86d3619"),
|
||||
("Alex (English, male)", "897d6a9b2c844f56aa077238768fe10a"),
|
||||
("Monika Sogam (English, female)", "97dd67ab8ce242b6a9e7689cb00c6414"),
|
||||
("Jessica Anne Bogart (English, female)", "b966c31caf124c2a99f19ff1479c964f"),
|
||||
("John Doe (English, male)", "c4a8ceb7a2954500bc047fb092bcff3f"),
|
||||
("Ivy (English, female)", "cef3bc4e0a84424cafcde6f2cf466c97"),
|
||||
("Chill Brian (English, male)", "d2f4f24783d04e22ab49ee8fdc3715e0"),
|
||||
("Allison (English, female)", "f8c69e517f424cafaecde32dde57096b"),
|
||||
("Mia Starset (Norwegian, female)", "000466f8ac6d47a49f5743d50b3778de"),
|
||||
("William Shanks (Spanish, male)", "001248bb63f847888d37b766ee8b3a47"),
|
||||
("Zain (English, female)", "0047732240584155b1588455313e78ec"),
|
||||
("Jora Slobod (Romanian, male)", "00631519159a402ab5d8f719e51532bb"),
|
||||
("Narrator Mateo - Excited 🤩 (Spanish, male)", "0077225a877e457db4572ccaf245910b"),
|
||||
("Aria (English, female)", "007e1378fc454a9f976db570ba6164a7"),
|
||||
("Caryns (English, female)", "0082e70326864107823605db0d77f5e0"),
|
||||
("Klara (English, female)", "01209fdcd1c24a109c86dc24ee0f71c0"),
|
||||
("Son Tran (Vietnamese, male)", "0132f85950a94d11ba180f885101bf84"),
|
||||
("Bold Kasia - Excited 🤩 (Polish, female)", "015482a78b9a46ebae74bd0beb17765b"),
|
||||
("Marc Aurèle (French, male)", "018a94cf15574491a0bab7f6799ac15b"),
|
||||
("Shaun (English, male)", "01c42cddcfdc4665a57b8d89cba8ffc1"),
|
||||
("Senthil (English, male)", "01d674cfd32b4728a3fddd21b7e7d543"),
|
||||
("Cody (English, male)", "01f98ed43e6140349f47dbd37a416827"),
|
||||
("Saffron (English, female)", "0258bbc2cd8648cfa357adfb833f6d7b"),
|
||||
("Blanka - Lifelike (English, female)", "02880d1c6fd94b7799d91135581ed810"),
|
||||
("Rami (English, male)", "02d5366a90af4c7a87157808ff352e33"),
|
||||
("Rhodes (English, female)", "02dce0a169b3460084b6c914d18fb2c8"),
|
||||
("Michelle - Voice 1 (English, female)", "02X8sHnuxFpsq1caYWN0"),
|
||||
("Tuba (, female)", "034ca0c32b6542028748d6d365d90d6a"),
|
||||
("Autumn - UGC 3 (English, female)", "03dca9ebfca441dba55fb14afa0791b7"),
|
||||
("Reassuring Rupert (English, male)", "03fcf8ecb0a94b6b94e9007edb7c35f8"),
|
||||
]
|
||||
HEYGEN_VOICE_GENERAL_OPTIONS = [x[0] for x in HEYGEN_VOICE_GENERAL]
|
||||
HEYGEN_VOICE_GENERAL_MAP = dict(HEYGEN_VOICE_GENERAL)
|
||||
|
||||
HEYGEN_TRANSLATE_LANGUAGES = [
|
||||
"English",
|
||||
"Spanish",
|
||||
"Spanish (Spain)",
|
||||
"Spanish (Mexico)",
|
||||
"French",
|
||||
"French (France)",
|
||||
"German",
|
||||
"German (Germany)",
|
||||
"Portuguese",
|
||||
"Portuguese (Brazil)",
|
||||
"Italian",
|
||||
"Italian (Italy)",
|
||||
"Japanese",
|
||||
"Japanese (Japan)",
|
||||
"Korean",
|
||||
"Chinese (Mandarin, Simplified)",
|
||||
"Arabic",
|
||||
"Hindi",
|
||||
"Hindi (India)",
|
||||
"Russian",
|
||||
"Russian (Russia)",
|
||||
"Dutch",
|
||||
"Polish",
|
||||
"Turkish",
|
||||
"Indonesian",
|
||||
"Vietnamese",
|
||||
"Ukrainian",
|
||||
"Afrikaans (South Africa)",
|
||||
"Albanian (Albania)",
|
||||
"Amharic (Ethiopia)",
|
||||
"Arabic (Algeria)",
|
||||
"Arabic (Bahrain)",
|
||||
"Arabic (Egypt)",
|
||||
"Arabic (Iraq)",
|
||||
"Arabic (Jordan)",
|
||||
"Arabic (Kuwait)",
|
||||
"Arabic (Lebanon)",
|
||||
"Arabic (Libya)",
|
||||
"Arabic (Morocco)",
|
||||
"Arabic (Oman)",
|
||||
"Arabic (Qatar)",
|
||||
"Arabic (Saudi Arabia)",
|
||||
"Arabic (Syria)",
|
||||
"Arabic (Tunisia)",
|
||||
"Arabic (United Arab Emirates)",
|
||||
"Arabic (World)",
|
||||
"Arabic (Yemen)",
|
||||
"Armenian (Armenia)",
|
||||
"Azerbaijani (Latin, Azerbaijan)",
|
||||
"Bangla (Bangladesh)",
|
||||
"Basque",
|
||||
"Belarusian (Belarus)",
|
||||
"Bengali (India)",
|
||||
"Bosnian (Bosnia and Herzegovina)",
|
||||
"Bulgarian",
|
||||
"Bulgarian (Bulgaria)",
|
||||
"Burmese (Myanmar)",
|
||||
"Catalan",
|
||||
"Chinese (Cantonese, Traditional)",
|
||||
"Chinese (Jilu Mandarin, Simplified)",
|
||||
"Chinese (Northeastern Mandarin, Simplified)",
|
||||
"Chinese (Southwestern Mandarin, Simplified)",
|
||||
"Chinese (Taiwanese Mandarin, Traditional)",
|
||||
"Chinese (Wu, Simplified)",
|
||||
"Chinese (Zhongyuan Mandarin Henan, Simplified)",
|
||||
"Chinese (Zhongyuan Mandarin Shaanxi, Simplified)",
|
||||
"Croatian",
|
||||
"Croatian (Croatia)",
|
||||
"Czech",
|
||||
"Czech (Czechia)",
|
||||
"Danish",
|
||||
"Danish (Denmark)",
|
||||
"Dutch (Belgium)",
|
||||
"Dutch (Netherlands)",
|
||||
"English (Australia)",
|
||||
"English (Canada)",
|
||||
"English (Hong Kong SAR)",
|
||||
"English (India)",
|
||||
"English (Ireland)",
|
||||
"English (Kenya)",
|
||||
"English (New Zealand)",
|
||||
"English (Nigeria)",
|
||||
"English (Philippines)",
|
||||
"English (Singapore)",
|
||||
"English (South Africa)",
|
||||
"English (Tanzania)",
|
||||
"English (UK)",
|
||||
"English (United States)",
|
||||
"Estonian (Estonia)",
|
||||
"Filipino",
|
||||
"Filipino (Cebuano)",
|
||||
"Filipino (Philippines)",
|
||||
"Finnish",
|
||||
"Finnish (Finland)",
|
||||
"French (Belgium)",
|
||||
"French (Canada)",
|
||||
"French (Switzerland)",
|
||||
"Galician",
|
||||
"Georgian (Georgia)",
|
||||
"German (Austria)",
|
||||
"German (Switzerland)",
|
||||
"Greek",
|
||||
"Greek (Greece)",
|
||||
"Gujarati (India)",
|
||||
"Haitian Creole (Haiti)",
|
||||
"Hebrew (Israel)",
|
||||
"Hungarian (Hungary)",
|
||||
"Icelandic (Iceland)",
|
||||
"Indonesian (Indonesia)",
|
||||
"Irish (Ireland)",
|
||||
"Javanese (Latin, Indonesia)",
|
||||
"Kannada (India)",
|
||||
"Kazakh (Kazakhstan)",
|
||||
"Khmer (Cambodia)",
|
||||
"Konkani (India)",
|
||||
"Korean (Korea)",
|
||||
"Lao (Laos)",
|
||||
"Latin (Vatican City)",
|
||||
"Latvian (Latvia)",
|
||||
"Lithuanian (Lithuania)",
|
||||
"Luxembourgish (Luxembourg)",
|
||||
"Macedonian (North Macedonia)",
|
||||
"Maithili (India)",
|
||||
"Malagasy (Madagascar)",
|
||||
"Malay",
|
||||
"Malay (Malaysia)",
|
||||
"Malayalam (India)",
|
||||
"Maltese (Malta)",
|
||||
"Mandarin",
|
||||
"Marathi (India)",
|
||||
"Mongolian (Mongolia)",
|
||||
"Nepali (Nepal)",
|
||||
"Norwegian Bokmål (Norway)",
|
||||
"Norwegian Nynorsk (Norway)",
|
||||
"Odia (India)",
|
||||
"Pashto (Afghanistan)",
|
||||
"Persian (Iran)",
|
||||
"Polish (Poland)",
|
||||
"Portuguese (Portugal)",
|
||||
"Punjabi (India)",
|
||||
"Romanian",
|
||||
"Romanian (Romania)",
|
||||
"Serbian (Latin, Serbia)",
|
||||
"Sindhi (India)",
|
||||
"Sinhala (Sri Lanka)",
|
||||
"Slovak",
|
||||
"Slovak (Slovakia)",
|
||||
"Slovenian (Slovenia)",
|
||||
"Somali (Somalia)",
|
||||
"Spanish (Argentina)",
|
||||
"Spanish (Bolivia)",
|
||||
"Spanish (Chile)",
|
||||
"Spanish (Colombia)",
|
||||
"Spanish (Costa Rica)",
|
||||
"Spanish (Cuba)",
|
||||
"Spanish (Dominican Republic)",
|
||||
"Spanish (Ecuador)",
|
||||
"Spanish (El Salvador)",
|
||||
"Spanish (Equatorial Guinea)",
|
||||
"Spanish (Guatemala)",
|
||||
"Spanish (Honduras)",
|
||||
"Spanish (Latin America)",
|
||||
"Spanish (Nicaragua)",
|
||||
"Spanish (Panama)",
|
||||
"Spanish (Paraguay)",
|
||||
"Spanish (Peru)",
|
||||
"Spanish (Puerto Rico)",
|
||||
"Spanish (United States)",
|
||||
"Spanish (Uruguay)",
|
||||
"Spanish (Venezuela)",
|
||||
"Sundanese (Indonesia)",
|
||||
"Swahili (Kenya)",
|
||||
"Swahili (Tanzania)",
|
||||
"Swedish",
|
||||
"Swedish (Sweden)",
|
||||
"Tamil",
|
||||
"Tamil (India)",
|
||||
"Tamil (Malaysia)",
|
||||
"Tamil (Singapore)",
|
||||
"Tamil (Sri Lanka)",
|
||||
"Telugu (India)",
|
||||
"Thai (Thailand)",
|
||||
"Turkish (Türkiye)",
|
||||
"Ukrainian (Ukraine)",
|
||||
"Urdu (India)",
|
||||
"Urdu (Pakistan)",
|
||||
"Uzbek (Latin, Uzbekistan)",
|
||||
"Vietnamese (Vietnam)",
|
||||
"Welsh (United Kingdom)",
|
||||
"Zulu (South Africa)",
|
||||
]
|
||||
@@ -128,7 +128,7 @@ class OpenAIResponse(ModelResponseProperties, ResponseProperties):
|
||||
parallel_tool_calls: bool | None = Field(True)
|
||||
status: str | None = Field(
|
||||
None,
|
||||
description="One of `completed`, `failed`, `in_progress`, `incomplete`, `queued`, or `cancelled`.",
|
||||
description="One of `completed`, `failed`, `in_progress`, or `incomplete`.",
|
||||
)
|
||||
usage: ResponseUsage | None = Field(None)
|
||||
|
||||
|
||||
@@ -1,49 +0,0 @@
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class SyncInputItem(BaseModel):
|
||||
type: str = Field(..., description="Input kind: 'video', 'image' or 'audio'.")
|
||||
url: str = Field(...)
|
||||
|
||||
|
||||
class SyncActiveSpeakerDetection(BaseModel):
|
||||
auto_detect: bool | None = Field(
|
||||
None, description="Detect the active speaker automatically. Video input only; rejected for images."
|
||||
)
|
||||
frame_number: int | None = Field(
|
||||
None, description="Frame used for manual speaker selection. Must be 0 for image inputs."
|
||||
)
|
||||
coordinates: list[int] | None = Field(
|
||||
None, description="Pixel [x, y] of the speaker's face in the frame selected by frame_number."
|
||||
)
|
||||
|
||||
|
||||
class SyncGenerationOptions(BaseModel):
|
||||
sync_mode: str | None = Field(
|
||||
None,
|
||||
description="How to resolve an audio/video duration mismatch: "
|
||||
"cut_off, bounce, loop, silence or remap. Ignored for image inputs.",
|
||||
)
|
||||
i2v_prompt: str | None = Field(
|
||||
None, description="Motion prompt for image-to-video generation. Image input only."
|
||||
)
|
||||
active_speaker_detection: SyncActiveSpeakerDetection | None = Field(None)
|
||||
|
||||
|
||||
class SyncGenerationRequest(BaseModel):
|
||||
model: str = Field(..., description="Generation model, e.g. 'sync-3'.")
|
||||
input: list[SyncInputItem] = Field(
|
||||
..., description="Exactly one visual input (video or image) plus one audio input."
|
||||
)
|
||||
options: SyncGenerationOptions | None = Field(None)
|
||||
|
||||
|
||||
class SyncGeneration(BaseModel):
|
||||
"""Subset of the Generation object returned by POST /v2/generate and GET /v2/generate/{id}."""
|
||||
|
||||
id: str = Field(...)
|
||||
status: str = Field(..., description="PENDING | PROCESSING | COMPLETED | FAILED | REJECTED")
|
||||
outputUrl: str | None = Field(None)
|
||||
outputDuration: float | None = Field(None)
|
||||
error: str | None = Field(None, description="Human-readable failure message.")
|
||||
errorCode: str | None = Field(None, description="Stable machine-readable code from the GET /v2/errors catalog.")
|
||||
+48
-122
@@ -24,11 +24,6 @@ from comfy_api_nodes.apis.gemini import (
|
||||
GeminiImageGenerateContentRequest,
|
||||
GeminiImageGenerationConfig,
|
||||
GeminiInlineData,
|
||||
GeminiInteraction,
|
||||
GeminiInteractionGenerationConfig,
|
||||
GeminiInteractionMediaPart,
|
||||
GeminiInteractionRequest,
|
||||
GeminiInteractionTextPart,
|
||||
GeminiMimeType,
|
||||
GeminiPart,
|
||||
GeminiRole,
|
||||
@@ -56,11 +51,9 @@ from comfy_api_nodes.util import (
|
||||
)
|
||||
|
||||
GEMINI_BASE_ENDPOINT = "/proxy/vertexai/gemini"
|
||||
GEMINI_INTERACTIONS_ENDPOINT = "/proxy/gemini-interactions"
|
||||
GEMINI_MAX_INPUT_FILE_SIZE = 20 * 1024 * 1024 # 20 MB
|
||||
GEMINI_URL_INPUT_BUDGET = 10
|
||||
GEMINI_MAX_INLINE_BYTES = 18 * 1024 * 1024
|
||||
GEMINI_INTERACTIONS_MAX_INLINE_BYTES = 90 * 1024 * 1024 # the Interactions API rejects requests over ~100MiB
|
||||
GEMINI_IMAGE_SYS_PROMPT = (
|
||||
"You are an expert image-generation engine. You must ALWAYS produce an image.\n"
|
||||
"Interpret all user input—regardless of "
|
||||
@@ -238,10 +231,29 @@ async def get_image_from_response(response: GeminiGenerateContentResponse, thoug
|
||||
return torch.cat(image_tensors, dim=0)
|
||||
|
||||
|
||||
async def get_video_from_response(
|
||||
response: GeminiGenerateContentResponse, cls: type[IO.ComfyNode] | None = None
|
||||
) -> InputImpl.VideoFromFile:
|
||||
parts = get_parts_by_type(response, "video/*")
|
||||
for part in parts:
|
||||
if part.inlineData and part.inlineData.data:
|
||||
return InputImpl.VideoFromFile(BytesIO(base64.b64decode(part.inlineData.data)))
|
||||
if part.fileData and part.fileData.fileUri:
|
||||
return await download_url_to_video_output(part.fileData.fileUri, cls=cls)
|
||||
model_message = get_text_from_response(response).strip()
|
||||
if model_message:
|
||||
raise ValueError(f"Gemini did not generate a video. Model response: {model_message}")
|
||||
raise ValueError(
|
||||
"Gemini did not generate a video. Try rephrasing your prompt, "
|
||||
"shortening the requested duration, or reducing the number of input images/videos."
|
||||
)
|
||||
|
||||
|
||||
def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | None:
|
||||
if not response.modelVersion:
|
||||
return None
|
||||
# Define prices (Cost per 1,000,000 tokens), see https://cloud.google.com/vertex-ai/generative-ai/pricing
|
||||
output_video_tokens_price = 0.0
|
||||
if response.modelVersion == "gemini-2.5-pro":
|
||||
input_tokens_price = 1.25
|
||||
output_text_tokens_price = 10.0
|
||||
@@ -262,10 +274,6 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
|
||||
input_tokens_price = 0.25
|
||||
output_text_tokens_price = 1.50
|
||||
output_image_tokens_price = 0.0
|
||||
elif response.modelVersion == "gemini-3.5-flash":
|
||||
input_tokens_price = 1.50
|
||||
output_text_tokens_price = 9.0
|
||||
output_image_tokens_price = 0.0
|
||||
elif response.modelVersion in ("gemini-3-pro-image-preview", "gemini-3-pro-image"):
|
||||
input_tokens_price = 2
|
||||
output_text_tokens_price = 12.0
|
||||
@@ -278,6 +286,11 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
|
||||
input_tokens_price = 0.25
|
||||
output_text_tokens_price = 1.50
|
||||
output_image_tokens_price = 30.0
|
||||
elif response.modelVersion == "gemini-omni-flash-preview":
|
||||
input_tokens_price = 2.145
|
||||
output_text_tokens_price = 12.87
|
||||
output_image_tokens_price = 0.0
|
||||
output_video_tokens_price = 25.025
|
||||
else:
|
||||
return None
|
||||
final_price = response.usageMetadata.promptTokenCount * input_tokens_price
|
||||
@@ -285,6 +298,8 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
|
||||
for i in response.usageMetadata.candidatesTokensDetails:
|
||||
if i.modality == Modality.IMAGE:
|
||||
final_price += output_image_tokens_price * i.tokenCount # for Nano Banana models
|
||||
elif i.modality == Modality.VIDEO:
|
||||
final_price += output_video_tokens_price * i.tokenCount # for Omni Flash
|
||||
else:
|
||||
final_price += output_text_tokens_price * i.tokenCount
|
||||
if response.usageMetadata.thoughtsTokenCount:
|
||||
@@ -292,58 +307,6 @@ def calculate_tokens_price(response: GeminiGenerateContentResponse) -> float | N
|
||||
return final_price / 1_000_000.0
|
||||
|
||||
|
||||
def get_text_from_interaction(interaction: GeminiInteraction) -> str:
|
||||
"""Extract and concatenate all model output text from an Interactions API response."""
|
||||
texts = []
|
||||
for step in interaction.steps or []:
|
||||
if step.type != "model_output":
|
||||
continue
|
||||
for content in step.content or []:
|
||||
if content.type == "text" and content.text:
|
||||
texts.append(content.text)
|
||||
return "\n".join(texts)
|
||||
|
||||
|
||||
async def get_video_from_interaction(
|
||||
interaction: GeminiInteraction, cls: type[IO.ComfyNode] | None = None
|
||||
) -> InputImpl.VideoFromFile:
|
||||
for step in interaction.steps or []:
|
||||
if step.type != "model_output":
|
||||
continue
|
||||
for content in step.content or []:
|
||||
if content.type != "video":
|
||||
continue
|
||||
if content.data:
|
||||
return InputImpl.VideoFromFile(BytesIO(base64.b64decode(content.data)))
|
||||
if content.uri:
|
||||
return await download_url_to_video_output(content.uri, cls=cls)
|
||||
model_message = get_text_from_interaction(interaction).strip()
|
||||
if model_message:
|
||||
raise ValueError(f"Gemini did not generate a video. Model response: {model_message}")
|
||||
raise ValueError(
|
||||
"Gemini did not generate a video. Try rephrasing your prompt, "
|
||||
"shortening the requested duration, or reducing the number of input images/videos."
|
||||
)
|
||||
|
||||
|
||||
def calculate_interaction_tokens_price(interaction: GeminiInteraction) -> float | None:
|
||||
if interaction.usage is None:
|
||||
return None
|
||||
input_tokens_price = 1.5
|
||||
output_tokens_prices = {"text": 9.0, "video": 17.5}
|
||||
thoughts_tokens_price = 9.0
|
||||
final_price = 0.0
|
||||
for i in interaction.usage.input_tokens_by_modality or []:
|
||||
if i.tokens:
|
||||
final_price += input_tokens_price * i.tokens
|
||||
for i in interaction.usage.output_tokens_by_modality or []:
|
||||
if i.tokens and i.modality in output_tokens_prices:
|
||||
final_price += output_tokens_prices[i.modality] * i.tokens
|
||||
if interaction.usage.total_thought_tokens:
|
||||
final_price += thoughts_tokens_price * interaction.usage.total_thought_tokens
|
||||
return final_price / 1_000_000.0
|
||||
|
||||
|
||||
def create_video_parts(video_input: Input.Video) -> list[GeminiPart]:
|
||||
"""Convert a single video input to Gemini API compatible parts (inline MP4/H.264)."""
|
||||
base_64_string = video_to_base64_string(
|
||||
@@ -470,24 +433,14 @@ async def build_gemini_media_parts(
|
||||
part, nbytes = _media_inline_part(kind, payload)
|
||||
inline_bytes += nbytes
|
||||
if inline_bytes > max_inline_bytes:
|
||||
detail = f" after the first {url_budget} inputs are uploaded as URLs" if url_budget else ""
|
||||
raise ValueError(
|
||||
f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB{detail}). "
|
||||
"Reduce the number or size of attached media."
|
||||
f"Too much media to send inline (over {max_inline_bytes // (1024 * 1024)}MB after the first "
|
||||
f"{url_budget} inputs are uploaded as URLs). Reduce the number or size of attached media."
|
||||
)
|
||||
parts.append(part)
|
||||
return parts
|
||||
|
||||
|
||||
def to_interaction_media_part(part: GeminiPart) -> GeminiInteractionMediaPart:
|
||||
"""Convert a fileData/inlineData GeminiPart into an Interactions API media part."""
|
||||
if part.fileData:
|
||||
mime = part.fileData.mimeType.value
|
||||
return GeminiInteractionMediaPart(type=mime.split("/")[0], uri=part.fileData.fileUri, mime_type=mime)
|
||||
mime = part.inlineData.mimeType.value
|
||||
return GeminiInteractionMediaPart(type=mime.split("/")[0], data=part.inlineData.data, mime_type=mime)
|
||||
|
||||
|
||||
class GeminiNode(IO.ComfyNode):
|
||||
"""
|
||||
Node to generate text responses from a Gemini model.
|
||||
@@ -666,12 +619,11 @@ class GeminiNode(IO.ComfyNode):
|
||||
|
||||
GEMINI_V2_MODELS: dict[str, str] = {
|
||||
"Gemini 3.1 Pro": "gemini-3.1-pro-preview",
|
||||
"Gemini 3.5 Flash": "gemini-3.5-flash",
|
||||
"Gemini 3.1 Flash-Lite": "gemini-3.1-flash-lite-preview",
|
||||
}
|
||||
|
||||
|
||||
def _gemini_text_model_inputs(thinking_default: str, thinking_options: list[str] | None = None) -> list[Input]:
|
||||
def _gemini_text_model_inputs(thinking_default: str) -> list[Input]:
|
||||
"""Per-model inputs revealed by the model DynamicCombo (shared media + sampling controls)."""
|
||||
return [
|
||||
IO.Autogrow.Input(
|
||||
@@ -709,7 +661,7 @@ def _gemini_text_model_inputs(thinking_default: str, thinking_options: list[str]
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"thinking_level",
|
||||
options=thinking_options or ["LOW", "HIGH"],
|
||||
options=["LOW", "HIGH"],
|
||||
default=thinking_default,
|
||||
tooltip="How hard the model reasons internally before answering. "
|
||||
"HIGH improves quality on difficult tasks but costs more (thinking) tokens and is slower.",
|
||||
@@ -767,10 +719,6 @@ class GeminiNodeV2(IO.ComfyNode):
|
||||
IO.DynamicCombo.Input(
|
||||
"model",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"Gemini 3.5 Flash",
|
||||
_gemini_text_model_inputs("MEDIUM", ["MINIMAL", "LOW", "MEDIUM", "HIGH"]),
|
||||
),
|
||||
IO.DynamicCombo.Option("Gemini 3.1 Pro", _gemini_text_model_inputs("HIGH")),
|
||||
IO.DynamicCombo.Option("Gemini 3.1 Flash-Lite", _gemini_text_model_inputs("LOW")),
|
||||
],
|
||||
@@ -811,13 +759,7 @@ class GeminiNodeV2(IO.ComfyNode):
|
||||
"type": "list_usd",
|
||||
"usd": [0.00025, 0.0015],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "3.5 flash") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.0015, 0.009],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: {
|
||||
} : {
|
||||
"type": "list_usd",
|
||||
"usd": [0.002, 0.012],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
@@ -1191,9 +1133,7 @@ class GeminiImage2(IO.ComfyNode):
|
||||
) -> IO.NodeOutput:
|
||||
validate_string(prompt, strip_whitespace=True, min_length=1)
|
||||
if model == "Nano Banana 2 (Gemini 3.1 Flash Image)":
|
||||
model = "gemini-3.1-flash-image"
|
||||
elif model == "gemini-3-pro-image-preview":
|
||||
model = "gemini-3-pro-image"
|
||||
model = "gemini-3.1-flash-image-preview"
|
||||
|
||||
parts: list[GeminiPart] = [GeminiPart(text=prompt)]
|
||||
if images is not None:
|
||||
@@ -1567,7 +1507,7 @@ class GeminiNanoBanana2V2(IO.ComfyNode):
|
||||
validate_string(prompt, strip_whitespace=True, min_length=1)
|
||||
model_choice = model["model"]
|
||||
if model_choice == "Nano Banana 2 (Gemini 3.1 Flash Image)":
|
||||
model_id = "gemini-3.1-flash-image"
|
||||
model_id = "gemini-3.1-flash-image-preview"
|
||||
elif model_choice == "Nano Banana 2 Lite":
|
||||
model_id = "gemini-3.1-flash-lite-image"
|
||||
else:
|
||||
@@ -1719,7 +1659,7 @@ class GeminiVideoOmni(IO.ComfyNode):
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
expr='{"type":"usd","usd":0.101,"format":{"suffix":"/second","approximate":true}}'
|
||||
expr='{"type":"usd","usd":0.146,"format":{"suffix":"/second","approximate":true}}'
|
||||
),
|
||||
)
|
||||
|
||||
@@ -1738,41 +1678,27 @@ class GeminiVideoOmni(IO.ComfyNode):
|
||||
for video in videos:
|
||||
validate_video_duration(video, max_duration=10)
|
||||
|
||||
parts: list[GeminiInteractionTextPart | GeminiInteractionMediaPart] = []
|
||||
parts: list[GeminiPart] = []
|
||||
if images or videos:
|
||||
# The Interactions API accepts video only inline or as a Files API URI, not as an HTTP URL.
|
||||
media_parts = await build_gemini_media_parts(
|
||||
cls, [], [], videos, url_budget=0, max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES
|
||||
)
|
||||
video_inline_bytes = sum(len(p.inlineData.data) for p in media_parts)
|
||||
media_parts += await build_gemini_media_parts(
|
||||
cls, images, [], [], max_inline_bytes=GEMINI_INTERACTIONS_MAX_INLINE_BYTES - video_inline_bytes
|
||||
)
|
||||
parts.extend(to_interaction_media_part(p) for p in media_parts)
|
||||
parts.append(GeminiInteractionTextPart(text=prompt))
|
||||
interaction = await sync_op(
|
||||
parts.extend(await build_gemini_media_parts(cls, images, [], videos))
|
||||
parts.append(GeminiPart(text=prompt))
|
||||
response = await sync_op(
|
||||
cls,
|
||||
ApiEndpoint(path=GEMINI_INTERACTIONS_ENDPOINT, method="POST"),
|
||||
data=GeminiInteractionRequest(
|
||||
model=model_id,
|
||||
input=parts,
|
||||
generation_config=GeminiInteractionGenerationConfig(
|
||||
ApiEndpoint(path=f"{GEMINI_BASE_ENDPOINT}/{model_id}", method="POST"),
|
||||
data=GeminiGenerateContentRequest(
|
||||
contents=[GeminiContent(role=GeminiRole.user, parts=parts)],
|
||||
generationConfig=GeminiGenerationConfig(
|
||||
responseModalities=["TEXT", "VIDEO"],
|
||||
temperature=model.get("temperature", 1.0),
|
||||
top_p=model.get("top_p", 0.95),
|
||||
topP=model.get("top_p", 0.95),
|
||||
),
|
||||
),
|
||||
response_model=GeminiInteraction,
|
||||
price_extractor=calculate_interaction_tokens_price,
|
||||
response_model=GeminiGenerateContentResponse,
|
||||
price_extractor=calculate_tokens_price,
|
||||
)
|
||||
if interaction.status != "completed":
|
||||
model_message = get_text_from_interaction(interaction).strip()
|
||||
raise ValueError(
|
||||
f"Gemini interaction did not complete (status: {interaction.status})."
|
||||
+ (f" Model response: {model_message}" if model_message else "")
|
||||
)
|
||||
return IO.NodeOutput(
|
||||
await get_video_from_interaction(interaction, cls=cls),
|
||||
get_text_from_interaction(interaction),
|
||||
await get_video_from_response(response, cls=cls),
|
||||
get_text_from_response(response),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,799 +0,0 @@
|
||||
import uuid
|
||||
|
||||
import torch
|
||||
from typing_extensions import override
|
||||
|
||||
from comfy_api.latest import IO, ComfyExtension, Input
|
||||
from comfy_api_nodes.apis.heygen import (
|
||||
HEYGEN_AVATAR_MAP,
|
||||
HEYGEN_AVATAR_OPTIONS,
|
||||
HEYGEN_TRANSLATE_LANGUAGES,
|
||||
HEYGEN_VOICE_GENERAL_MAP,
|
||||
HEYGEN_VOICE_GENERAL_OPTIONS,
|
||||
HEYGEN_VOICE_TTS_MAP,
|
||||
HEYGEN_VOICE_TTS_OPTIONS,
|
||||
)
|
||||
from comfy_api_nodes.util import (
|
||||
ApiEndpoint,
|
||||
audio_bytes_to_audio_input,
|
||||
download_url_as_bytesio,
|
||||
download_url_to_image_tensor,
|
||||
download_url_to_video_output,
|
||||
downscale_image_tensor_by_max_side,
|
||||
get_number_of_images,
|
||||
poll_op_raw,
|
||||
sync_op_raw,
|
||||
upload_audio_to_comfyapi,
|
||||
upload_image_to_comfyapi,
|
||||
upload_images_to_comfyapi,
|
||||
upload_video_to_comfyapi,
|
||||
validate_string,
|
||||
)
|
||||
from server import PromptServer
|
||||
|
||||
_VIDEOS_PATH = "/proxy/heygen/v3/videos"
|
||||
_TRANSLATIONS_PATH = "/proxy/heygen/v3/video-translations"
|
||||
_SPEECH_PATH = "/proxy/heygen/v3/voices/speech"
|
||||
_AVATARS_PATH = "/proxy/heygen/v3/avatars"
|
||||
_LOOKS_PATH = "/proxy/heygen/v3/avatars/looks"
|
||||
|
||||
_DEFAULT_VOICE_OPTION = "(avatar's default voice)"
|
||||
|
||||
_AVATARS_BY_ENGINE = {
|
||||
e: [label for label, (_aid, _atype, engines) in HEYGEN_AVATAR_MAP.items() if e in engines]
|
||||
for e in ("avatar_iv", "avatar_iii", "avatar_v")
|
||||
}
|
||||
|
||||
|
||||
async def _apply_speech_source(cls: type[IO.ComfyNode], payload: dict, speech: dict, require_voice: bool) -> None:
|
||||
"""Fill script/audio speech fields of a /v3/videos payload from the DynamicCombo dict."""
|
||||
if speech["speech"] == "audio":
|
||||
payload["audio_url"] = await upload_audio_to_comfyapi(
|
||||
cls, speech["audio"], container_format="mp3", codec_name="libmp3lame", mime_type="audio/mpeg"
|
||||
)
|
||||
elif speech["speech"] == "script":
|
||||
validate_string(speech["text"], strip_whitespace=True, min_length=1, max_length=5000)
|
||||
payload["script"] = speech["text"]
|
||||
voice_id = speech.get("custom_voice_id", "").strip()
|
||||
if not voice_id and speech["voice"] != _DEFAULT_VOICE_OPTION:
|
||||
voice_id = HEYGEN_VOICE_GENERAL_MAP[speech["voice"]]
|
||||
if voice_id:
|
||||
payload["voice_id"] = voice_id
|
||||
elif require_voice:
|
||||
raise ValueError("A voice is required when driving the video with a text script.")
|
||||
speed = speech.get("voice_speed", 1.0)
|
||||
if speed != 1.0:
|
||||
payload["voice_settings"] = {"speed": round(speed, 2)}
|
||||
|
||||
|
||||
async def _create_and_poll_video(cls: type[IO.ComfyNode], payload: dict) -> dict:
|
||||
"""POST a /v3/videos payload, poll until terminal, and return the final video data."""
|
||||
created = await sync_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=_VIDEOS_PATH, method="POST", headers={"Idempotency-Key": uuid.uuid4().hex}),
|
||||
data=payload,
|
||||
)
|
||||
video_id = (created.get("data") or {}).get("video_id")
|
||||
if not video_id:
|
||||
raise ValueError(f"HeyGen did not return a video_id: {created}")
|
||||
final = await poll_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=f"{_VIDEOS_PATH}/{video_id}"),
|
||||
status_extractor=lambda r: (r.get("data") or {}).get("status"),
|
||||
queued_statuses=["pending", "waiting"],
|
||||
poll_interval=5.0,
|
||||
)
|
||||
data = final["data"]
|
||||
if not data.get("video_url"):
|
||||
raise ValueError(f"HeyGen returned no video_url for video {video_id}.")
|
||||
return data
|
||||
|
||||
|
||||
async def _resolve_avatar(
|
||||
cls: type[IO.ComfyNode], avatar_label: str, custom_avatar_id: str, engine_choice: str
|
||||
) -> tuple[str, str | None]:
|
||||
"""Resolve (avatar_id, engine_type) from the combo/override + engine widgets."""
|
||||
custom_avatar_id = custom_avatar_id.strip()
|
||||
if custom_avatar_id:
|
||||
look = (
|
||||
await sync_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=f"{_LOOKS_PATH}/{custom_avatar_id}"),
|
||||
final_label_on_success=None,
|
||||
)
|
||||
).get("data") or {}
|
||||
avatar_id = custom_avatar_id
|
||||
avatar_label = look.get("name") or custom_avatar_id
|
||||
supported = look.get("supported_api_engines") or []
|
||||
else:
|
||||
avatar_id, avatar_type, supported = HEYGEN_AVATAR_MAP[avatar_label]
|
||||
|
||||
if engine_choice == "auto":
|
||||
engine = next((e for e in ("avatar_iv", "avatar_iii", "avatar_v") if e in supported), None)
|
||||
else:
|
||||
engine = engine_choice
|
||||
if supported and engine not in supported:
|
||||
raise ValueError(
|
||||
f"Avatar '{avatar_label}' does not support the {engine} engine "
|
||||
f"(supported: {', '.join(supported)}). Set engine to 'auto' to pick "
|
||||
"a compatible engine automatically."
|
||||
)
|
||||
return avatar_id, engine
|
||||
|
||||
|
||||
class HeyGenTalkingPhotoNode(IO.ComfyNode):
|
||||
"""Animate a still image of a person into a lip-synced talking video."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="HeyGenTalkingPhotoNode",
|
||||
display_name="HeyGen Talking Photo",
|
||||
category="partner/video/HeyGen",
|
||||
description="Animate any image of a person into a lip-synced talking video "
|
||||
"(HeyGen Avatar IV). Drive it with a text script or your own audio.",
|
||||
inputs=[
|
||||
IO.Image.Input(
|
||||
"image",
|
||||
tooltip="Image of a person to animate. Downscaled automatically if larger than 2K.",
|
||||
),
|
||||
IO.DynamicCombo.Input(
|
||||
"speech",
|
||||
display_name="speech source",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"script",
|
||||
[
|
||||
IO.String.Input(
|
||||
"text",
|
||||
multiline=True,
|
||||
default="",
|
||||
tooltip="Text for the avatar to speak (up to 5000 characters). "
|
||||
"The generated speech must be at least 1 second long.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"voice",
|
||||
options=HEYGEN_VOICE_GENERAL_OPTIONS,
|
||||
tooltip="Voice for the script (HeyGen's most popular voices).",
|
||||
),
|
||||
IO.String.Input(
|
||||
"custom_voice_id",
|
||||
default="",
|
||||
optional=True,
|
||||
tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. "
|
||||
"Any voice from HeyGen's library (2000+) can be used.",
|
||||
),
|
||||
IO.Float.Input(
|
||||
"voice_speed",
|
||||
default=1.0,
|
||||
min=0.5,
|
||||
max=1.5,
|
||||
step=0.05,
|
||||
optional=True,
|
||||
tooltip="Speech speed multiplier.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"audio",
|
||||
[
|
||||
IO.Audio.Input(
|
||||
"audio",
|
||||
tooltip="Audio for the avatar to lip-sync, up to 10 minutes.",
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
tooltip="Drive the avatar with a text script (HeyGen text-to-speech) or your own audio.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"resolution",
|
||||
options=["720p", "1080p"],
|
||||
default="1080p",
|
||||
optional=True,
|
||||
tooltip="Output video resolution.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"aspect_ratio",
|
||||
options=["auto", "16:9", "9:16", "1:1", "4:5", "5:4"],
|
||||
default="auto",
|
||||
optional=True,
|
||||
tooltip="Output aspect ratio. 'auto' follows the input image.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"expressiveness",
|
||||
options=["low", "medium", "high"],
|
||||
default="low",
|
||||
optional=True,
|
||||
tooltip="How expressive the animated face and gestures are.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=42,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
optional=True,
|
||||
tooltip="Not sent to HeyGen; change it to force a re-run.",
|
||||
),
|
||||
],
|
||||
outputs=[IO.Video.Output()],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
expr="""{"type":"usd","usd":0.0715,"format":{"suffix":"/second"}}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
image: Input.Image,
|
||||
speech: dict,
|
||||
resolution: str = "1080p",
|
||||
aspect_ratio: str = "auto",
|
||||
expressiveness: str = "low",
|
||||
seed: int = 0,
|
||||
) -> IO.NodeOutput:
|
||||
image = downscale_image_tensor_by_max_side(image, max_side=2000)
|
||||
image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None)
|
||||
payload = {
|
||||
"type": "image",
|
||||
"image": {"type": "url", "url": image_url},
|
||||
"resolution": resolution,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"expressiveness": expressiveness,
|
||||
"title": "ComfyUI Talking Photo",
|
||||
}
|
||||
await _apply_speech_source(cls, payload, speech, require_voice=True)
|
||||
video = await _create_and_poll_video(cls, payload)
|
||||
return IO.NodeOutput(await download_url_to_video_output(video["video_url"]))
|
||||
|
||||
|
||||
class HeyGenAvatarVideoNode(IO.ComfyNode):
|
||||
"""Generate a presenter video from a HeyGen avatar look."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="HeyGenAvatarVideoNode",
|
||||
display_name="HeyGen Avatar Video",
|
||||
category="partner/video/HeyGen",
|
||||
description="Generate a talking-presenter video from a HeyGen avatar. "
|
||||
"Includes HeyGen's most popular public avatars; any look ID can be supplied as an override.",
|
||||
inputs=[
|
||||
IO.DynamicCombo.Input(
|
||||
"engine",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"auto",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"avatar",
|
||||
options=HEYGEN_AVATAR_OPTIONS,
|
||||
tooltip="Avatar look to present the video (curated from HeyGen's "
|
||||
"public library). The best engine the look supports is chosen "
|
||||
"automatically.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"avatar_iv",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"avatar",
|
||||
options=_AVATARS_BY_ENGINE["avatar_iv"],
|
||||
tooltip="Avatar looks that support the Avatar IV engine.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"avatar_iii",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"avatar",
|
||||
options=_AVATARS_BY_ENGINE["avatar_iii"],
|
||||
tooltip="Avatar looks that support the Avatar III engine.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"avatar_v",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"avatar",
|
||||
options=_AVATARS_BY_ENGINE["avatar_v"],
|
||||
tooltip="Avatar looks that support the Avatar V engine.",
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
tooltip="Rendering engine; each choice lists only the avatars that support it. "
|
||||
"'auto' offers every avatar and picks its best engine (Avatar IV preferred). "
|
||||
"Avatar V is highest fidelity, Avatar III is the most affordable.",
|
||||
),
|
||||
IO.String.Input(
|
||||
"custom_avatar_id",
|
||||
default="",
|
||||
optional=True,
|
||||
tooltip="Optional HeyGen avatar look ID. When set, overrides the avatar selected above. "
|
||||
"Any of HeyGen's 3000+ public looks (or your private avatars) can be used.",
|
||||
),
|
||||
IO.DynamicCombo.Input(
|
||||
"speech",
|
||||
display_name="speech source",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"script",
|
||||
[
|
||||
IO.String.Input(
|
||||
"text",
|
||||
multiline=True,
|
||||
default="",
|
||||
tooltip="Text for the avatar to speak (up to 5000 characters). "
|
||||
"The generated speech must be at least 1 second long.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"voice",
|
||||
options=[_DEFAULT_VOICE_OPTION] + HEYGEN_VOICE_GENERAL_OPTIONS,
|
||||
tooltip="Voice for the script. The default option uses the voice HeyGen assigned to the avatar.",
|
||||
),
|
||||
IO.String.Input(
|
||||
"custom_voice_id",
|
||||
default="",
|
||||
optional=True,
|
||||
tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. "
|
||||
"Any voice from HeyGen's library (2000+) can be used.",
|
||||
),
|
||||
IO.Float.Input(
|
||||
"voice_speed",
|
||||
default=1.0,
|
||||
min=0.5,
|
||||
max=1.5,
|
||||
step=0.05,
|
||||
optional=True,
|
||||
tooltip="Speech speed multiplier.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"audio",
|
||||
[
|
||||
IO.Audio.Input(
|
||||
"audio",
|
||||
tooltip="Audio for the avatar to lip-sync, up to 10 minutes.",
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
tooltip="Drive the avatar with a text script (HeyGen text-to-speech) or your own audio.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"resolution",
|
||||
options=["720p", "1080p"],
|
||||
default="1080p",
|
||||
optional=True,
|
||||
tooltip="Output video resolution.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"aspect_ratio",
|
||||
options=["auto", "16:9", "9:16", "1:1", "4:5", "5:4"],
|
||||
default="auto",
|
||||
optional=True,
|
||||
tooltip="Output aspect ratio. 'auto' follows the avatar's source footage.",
|
||||
),
|
||||
IO.String.Input(
|
||||
"background_color",
|
||||
default="",
|
||||
optional=True,
|
||||
tooltip="Optional solid background color as a hex code (e.g. '#00ff00'). "
|
||||
"Leave empty for the avatar's own background.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=42,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
optional=True,
|
||||
tooltip="Not sent to HeyGen; change it to force a re-run.",
|
||||
),
|
||||
],
|
||||
outputs=[IO.Video.Output()],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
depends_on=IO.PriceBadgeDepends(widgets=["engine"]),
|
||||
expr="""
|
||||
widgets.engine = "avatar_iii"
|
||||
? {"type":"range_usd","min_usd":0.023881,"max_usd":0.061919,"format":{"suffix":"/second"}}
|
||||
: widgets.engine = "avatar_v"
|
||||
? {"type":"usd","usd":0.095381,"format":{"suffix":"/second"}}
|
||||
: widgets.engine = "avatar_iv"
|
||||
? {"type":"range_usd","min_usd":0.0715,"max_usd":0.095381,"format":{"suffix":"/second"}}
|
||||
: {"type":"range_usd","min_usd":0.023881,"max_usd":0.095381,"format":{"suffix":"/second"}}
|
||||
""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
engine: dict,
|
||||
speech: dict,
|
||||
custom_avatar_id: str = "",
|
||||
resolution: str = "1080p",
|
||||
aspect_ratio: str = "auto",
|
||||
background_color: str = "",
|
||||
seed: int = 0,
|
||||
) -> IO.NodeOutput:
|
||||
avatar_id, engine_type = await _resolve_avatar(cls, engine["avatar"], custom_avatar_id, engine["engine"])
|
||||
payload = {
|
||||
"type": "avatar",
|
||||
"avatar_id": avatar_id,
|
||||
"resolution": resolution,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"title": "ComfyUI Avatar Video",
|
||||
}
|
||||
if engine_type:
|
||||
payload["engine"] = {"type": engine_type}
|
||||
background_color = background_color.strip()
|
||||
if background_color:
|
||||
if not background_color.startswith("#"):
|
||||
raise ValueError("background_color must be a hex color code like '#00ff00'.")
|
||||
payload["background"] = {"type": "color", "value": background_color}
|
||||
await _apply_speech_source(cls, payload, speech, require_voice=False)
|
||||
video = await _create_and_poll_video(cls, payload)
|
||||
return IO.NodeOutput(await download_url_to_video_output(video["video_url"]))
|
||||
|
||||
|
||||
class HeyGenCreateAvatarNode(IO.ComfyNode):
|
||||
"""Create a reusable HeyGen avatar from a photo or a text prompt."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="HeyGenCreateAvatarNode",
|
||||
display_name="HeyGen Create Avatar",
|
||||
category="partner/video/HeyGen",
|
||||
description="Create your own reusable HeyGen avatar from a photo of a person or "
|
||||
"from a text prompt (a generated character). Feed the resulting avatar_id into "
|
||||
"HeyGen Avatar Video's custom_avatar_id — and save the ID somewhere to reuse the "
|
||||
"avatar in future workflows.",
|
||||
inputs=[
|
||||
IO.DynamicCombo.Input(
|
||||
"source",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"prompt",
|
||||
[
|
||||
IO.String.Input(
|
||||
"prompt",
|
||||
multiline=True,
|
||||
default="",
|
||||
tooltip="Description of the avatar to generate (up to 1000 characters).",
|
||||
),
|
||||
IO.Autogrow.Input(
|
||||
"reference_images",
|
||||
template=IO.Autogrow.TemplateNames(
|
||||
IO.Image.Input("ref_image"),
|
||||
names=[f"ref_image_{i}" for i in range(1, 4)],
|
||||
min=0,
|
||||
),
|
||||
tooltip="Up to 3 reference images guiding the generated look.",
|
||||
),
|
||||
],
|
||||
),
|
||||
IO.DynamicCombo.Option(
|
||||
"photo",
|
||||
[
|
||||
IO.Image.Input(
|
||||
"identity_photo",
|
||||
tooltip="Photo of the person to turn into an avatar. "
|
||||
"Downscaled automatically if larger than 2K.",
|
||||
),
|
||||
],
|
||||
),
|
||||
],
|
||||
tooltip="Generate a new character from a text prompt, or create the avatar "
|
||||
"from a connected photo of a person.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
IO.String.Output(
|
||||
display_name="avatar_id",
|
||||
tooltip="Avatar look ID. Pass it to HeyGen Avatar Video's custom_avatar_id; "
|
||||
"save it to reuse the avatar later.",
|
||||
),
|
||||
IO.Image.Output(display_name="preview"),
|
||||
],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
expr="""{"type":"usd","usd":1.43}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
source: dict,
|
||||
) -> IO.NodeOutput:
|
||||
payload: dict = {"name": "ComfyUI Avatar"}
|
||||
if source["source"] == "photo":
|
||||
image = downscale_image_tensor_by_max_side(source["identity_photo"], max_side=2000)
|
||||
image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None)
|
||||
payload["type"] = "photo"
|
||||
payload["file"] = {"type": "url", "url": image_url}
|
||||
else:
|
||||
validate_string(source["prompt"], strip_whitespace=True, min_length=1, max_length=1000)
|
||||
payload["type"] = "prompt"
|
||||
payload["prompt"] = source["prompt"]
|
||||
ref_tensors = [t for t in (source.get("reference_images") or {}).values() if t is not None]
|
||||
if ref_tensors:
|
||||
n_images = sum(get_number_of_images(t) for t in ref_tensors)
|
||||
if n_images > 3:
|
||||
raise ValueError(f"HeyGen accepts at most 3 reference images; got {n_images}.")
|
||||
scaled = [downscale_image_tensor_by_max_side(t, max_side=2000) for t in ref_tensors]
|
||||
ref_urls = await upload_images_to_comfyapi(
|
||||
cls, scaled, max_images=3, mime_type="image/png", total_pixels=None
|
||||
)
|
||||
payload["reference_images"] = [{"type": "url", "url": u} for u in ref_urls]
|
||||
created = await sync_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=_AVATARS_PATH, method="POST"),
|
||||
data=payload,
|
||||
)
|
||||
look_id = ((created.get("data") or {}).get("avatar_item") or {}).get("id")
|
||||
if not look_id:
|
||||
raise ValueError(f"HeyGen did not return an avatar: {created}")
|
||||
final = await poll_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=f"{_LOOKS_PATH}/{look_id}"),
|
||||
# A missing status means the look needed no training and is ready.
|
||||
status_extractor=lambda r: (r.get("data") or {}).get("status") or "completed",
|
||||
failed_statuses=["failed", "pending_consent"],
|
||||
poll_interval=5.0,
|
||||
)
|
||||
data = final["data"]
|
||||
if data.get("preview_image_url"):
|
||||
preview = await download_url_to_image_tensor(data["preview_image_url"])
|
||||
else:
|
||||
preview = torch.zeros(1, 64, 64, 3)
|
||||
PromptServer.instance.send_progress_text(
|
||||
f"Please save the avatar_id for reuse.\n\navatar_id: {look_id}",
|
||||
cls.hidden.unique_id,
|
||||
)
|
||||
return IO.NodeOutput(look_id, preview)
|
||||
|
||||
|
||||
class HeyGenVideoTranslateNode(IO.ComfyNode):
|
||||
"""Translate a spoken video into another language with voice cloning and lip sync."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="HeyGenVideoTranslateNode",
|
||||
display_name="HeyGen Video Translate",
|
||||
category="partner/video/HeyGen",
|
||||
description="Translate a spoken video into another language. Clones the original "
|
||||
"speaker's voice and re-animates the mouth to match the translated speech.",
|
||||
inputs=[
|
||||
IO.Video.Input(
|
||||
"video",
|
||||
tooltip="Video with speech to translate.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"output_language",
|
||||
options=HEYGEN_TRANSLATE_LANGUAGES,
|
||||
tooltip="Target language for the translated video.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"mode",
|
||||
options=["speed", "precision"],
|
||||
default="speed",
|
||||
tooltip="'speed' is faster; 'precision' produces higher-quality lip sync at twice the price.",
|
||||
),
|
||||
IO.Boolean.Input(
|
||||
"translate_audio_only",
|
||||
default=False,
|
||||
optional=True,
|
||||
tooltip="Only swap the audio track, keeping the original mouth movements (no lip sync).",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_count",
|
||||
default=0,
|
||||
min=0,
|
||||
max=10,
|
||||
optional=True,
|
||||
tooltip="Number of speakers in the video. 0 = detect automatically.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=42,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
optional=True,
|
||||
tooltip="Not sent to HeyGen; change it to force a re-run.",
|
||||
),
|
||||
],
|
||||
outputs=[IO.Video.Output()],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
depends_on=IO.PriceBadgeDepends(widgets=["mode"]),
|
||||
expr="""{"type":"usd","usd": widgets.mode = "precision" ? 0.095381 : 0.047619,"""
|
||||
""""format":{"suffix":"/second"}}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
video: Input.Video,
|
||||
output_language: str,
|
||||
mode: str,
|
||||
translate_audio_only: bool = False,
|
||||
speaker_count: int = 0,
|
||||
seed: int = 0,
|
||||
) -> IO.NodeOutput:
|
||||
video_url = await upload_video_to_comfyapi(cls, video)
|
||||
payload = {
|
||||
"video": {"type": "url", "url": video_url},
|
||||
"output_languages": [output_language],
|
||||
"mode": mode,
|
||||
"translate_audio_only": translate_audio_only,
|
||||
"title": "ComfyUI Video Translate",
|
||||
}
|
||||
if speaker_count > 0:
|
||||
payload["speaker_num"] = speaker_count
|
||||
created = await sync_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=_TRANSLATIONS_PATH, method="POST"),
|
||||
data=payload,
|
||||
)
|
||||
translation_ids = (created.get("data") or {}).get("video_translation_ids") or []
|
||||
if not translation_ids:
|
||||
raise ValueError(f"HeyGen did not return a translation ID: {created}")
|
||||
final = await poll_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=f"{_TRANSLATIONS_PATH}/{translation_ids[0]}"),
|
||||
status_extractor=lambda r: (r.get("data") or {}).get("status"),
|
||||
queued_statuses=["pending"],
|
||||
poll_interval=5.0,
|
||||
)
|
||||
data = final["data"]
|
||||
if not data.get("video_url"):
|
||||
raise ValueError(f"HeyGen returned no video_url for translation {translation_ids[0]}.")
|
||||
return IO.NodeOutput(await download_url_to_video_output(data["video_url"]))
|
||||
|
||||
|
||||
class HeyGenTextToSpeechNode(IO.ComfyNode):
|
||||
"""Synthesize speech audio from text with HeyGen's Starfish TTS engine."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="HeyGenTextToSpeechNode",
|
||||
display_name="HeyGen Text to Speech",
|
||||
category="partner/audio/HeyGen",
|
||||
description="Generate speech audio from text using HeyGen's Starfish TTS engine. "
|
||||
"Includes HeyGen's most popular voices across 17 languages.",
|
||||
inputs=[
|
||||
IO.String.Input(
|
||||
"text",
|
||||
multiline=True,
|
||||
default="",
|
||||
tooltip="Text to synthesize (up to 5000 characters). The generated speech "
|
||||
"must be at least 1 second long.",
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"voice",
|
||||
options=HEYGEN_VOICE_TTS_OPTIONS,
|
||||
tooltip="Voice to use (curated from HeyGen's most popular Starfish-compatible voices).",
|
||||
),
|
||||
IO.String.Input(
|
||||
"custom_voice_id",
|
||||
default="",
|
||||
optional=True,
|
||||
tooltip="Optional HeyGen voice ID. When set, overrides the voice selected above. "
|
||||
"The voice must support the Starfish engine.",
|
||||
),
|
||||
IO.Float.Input(
|
||||
"speed",
|
||||
default=1.0,
|
||||
min=0.5,
|
||||
max=2.0,
|
||||
step=0.05,
|
||||
optional=True,
|
||||
tooltip="Speech speed multiplier.",
|
||||
),
|
||||
IO.Boolean.Input(
|
||||
"ssml",
|
||||
default=False,
|
||||
optional=True,
|
||||
tooltip="Treat the text as SSML markup (for pauses, emphasis, and pronunciation control).",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=42,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
optional=True,
|
||||
tooltip="Not sent to HeyGen; change it to force a re-run.",
|
||||
),
|
||||
],
|
||||
outputs=[IO.Audio.Output()],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
expr="""{"type":"usd","usd":0.00095381,"format":{"approximate":true,"suffix":"/second"}}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
text: str,
|
||||
voice: str,
|
||||
custom_voice_id: str = "",
|
||||
speed: float = 1.0,
|
||||
ssml: bool = False,
|
||||
seed: int = 0,
|
||||
) -> IO.NodeOutput:
|
||||
validate_string(text, strip_whitespace=True, min_length=1, max_length=5000)
|
||||
payload = {
|
||||
"text": text,
|
||||
"voice_id": custom_voice_id.strip() or HEYGEN_VOICE_TTS_MAP[voice],
|
||||
"speed": round(speed, 2),
|
||||
}
|
||||
if ssml:
|
||||
payload["input_type"] = "ssml"
|
||||
response = await sync_op_raw(
|
||||
cls,
|
||||
ApiEndpoint(path=_SPEECH_PATH, method="POST"),
|
||||
data=payload,
|
||||
)
|
||||
audio_url = (response.get("data") or {}).get("audio_url")
|
||||
if not audio_url:
|
||||
raise ValueError(f"HeyGen did not return an audio_url: {response}")
|
||||
audio_bytes = await download_url_as_bytesio(audio_url)
|
||||
return IO.NodeOutput(audio_bytes_to_audio_input(audio_bytes.getvalue()))
|
||||
|
||||
|
||||
class HeyGenExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[IO.ComfyNode]]:
|
||||
return [
|
||||
HeyGenTalkingPhotoNode,
|
||||
HeyGenAvatarVideoNode,
|
||||
HeyGenCreateAvatarNode,
|
||||
HeyGenVideoTranslateNode,
|
||||
HeyGenTextToSpeechNode,
|
||||
]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> HeyGenExtension:
|
||||
return HeyGenExtension()
|
||||
@@ -41,9 +41,6 @@ STARTING_POINT_ID_PATTERN = r"<starting_point_id:(.*)>"
|
||||
|
||||
|
||||
class SupportedOpenAIModel(str, Enum):
|
||||
gpt_5_6_sol = "gpt-5.6-sol"
|
||||
gpt_5_6_terra = "gpt-5.6-terra"
|
||||
gpt_5_6_luna = "gpt-5.6-luna"
|
||||
gpt_5_5_pro = "gpt-5.5-pro"
|
||||
gpt_5_5 = "gpt-5.5"
|
||||
gpt_5 = "gpt-5"
|
||||
@@ -1066,21 +1063,6 @@ class OpenAIChatNode(IO.ComfyNode):
|
||||
"usd": [0.002, 0.008],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "gpt-5.6-terra") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.0025, 0.015],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "gpt-5.6-luna") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.001, 0.006],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "gpt-5.6") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.005, 0.03],
|
||||
"format": { "approximate": true, "separator": "-", "suffix": " per 1K tokens" }
|
||||
}
|
||||
: $contains($m, "gpt-5.5-pro") ? {
|
||||
"type": "list_usd",
|
||||
"usd": [0.03, 0.18],
|
||||
|
||||
@@ -1,391 +0,0 @@
|
||||
from typing_extensions import override
|
||||
|
||||
from comfy_api.latest import IO, ComfyExtension, Input
|
||||
from comfy_api_nodes.apis.sync_so import (
|
||||
SyncActiveSpeakerDetection,
|
||||
SyncGeneration,
|
||||
SyncGenerationOptions,
|
||||
SyncGenerationRequest,
|
||||
SyncInputItem,
|
||||
)
|
||||
from comfy_api_nodes.util import (
|
||||
ApiEndpoint,
|
||||
download_url_to_video_output,
|
||||
downscale_image_tensor,
|
||||
downscale_image_tensor_by_max_side,
|
||||
get_image_dimensions,
|
||||
get_number_of_images,
|
||||
poll_op,
|
||||
sync_op,
|
||||
upload_audio_to_comfyapi,
|
||||
upload_image_to_comfyapi,
|
||||
upload_video_to_comfyapi,
|
||||
validate_audio_duration,
|
||||
)
|
||||
|
||||
|
||||
class SyncLipSyncNode(IO.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="SyncLipSyncNode",
|
||||
display_name="sync.so Lip Sync",
|
||||
category="partner/video/sync.so",
|
||||
description=(
|
||||
"Re-sync mouth movement in a video to new speech audio using sync.so. "
|
||||
"Handles close-ups, profiles and obstructions automatically while preserving "
|
||||
"the speaker's expression. Cost scales with output duration."
|
||||
),
|
||||
inputs=[
|
||||
IO.Video.Input(
|
||||
"video",
|
||||
tooltip="Footage of the speaker to re-sync. Up to 4K (4096x2160); "
|
||||
"a constant frame rate of 24/25/30 fps works best.",
|
||||
),
|
||||
IO.Audio.Input(
|
||||
"audio",
|
||||
tooltip="Speech audio to sync the mouth to.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=42,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
tooltip="Seed controls whether the node should re-run; "
|
||||
"results are non-deterministic regardless of seed.",
|
||||
),
|
||||
IO.DynamicCombo.Input(
|
||||
"model",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"sync-3",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"sync_mode",
|
||||
options=["bounce", "cut_off", "loop", "silence", "remap"],
|
||||
default="bounce",
|
||||
tooltip=(
|
||||
"How to handle a duration mismatch between video and audio; "
|
||||
"this also sets the output length. "
|
||||
"bounce: video plays forward then backward until the audio ends "
|
||||
"(output = audio length). "
|
||||
"loop: video restarts until the audio ends (output = audio length). "
|
||||
"remap: video is time-stretched to match the audio (output = audio length). "
|
||||
"cut_off: the longer track is trimmed (output = shorter length). "
|
||||
"silence: nothing is trimmed; the shorter track is padded "
|
||||
"(output = longer length)."
|
||||
),
|
||||
),
|
||||
IO.Combo.Input(
|
||||
"speaker_selection",
|
||||
options=["default", "auto-detect", "coordinates"],
|
||||
default="default",
|
||||
tooltip=(
|
||||
"Which face to lipsync when several people are visible. "
|
||||
"default: let the model decide. "
|
||||
"auto-detect: detect and follow the active speaker. "
|
||||
"coordinates: target the face at pixel (speaker_x, speaker_y) "
|
||||
"in the frame chosen by speaker_frame."
|
||||
),
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_frame",
|
||||
default=0,
|
||||
min=0,
|
||||
max=1_000_000,
|
||||
advanced=True,
|
||||
tooltip="Video frame used to locate the speaker. "
|
||||
"Only used when speaker_selection is 'coordinates'.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_x",
|
||||
default=0,
|
||||
min=0,
|
||||
max=4096,
|
||||
advanced=True,
|
||||
tooltip="X pixel coordinate of the speaker's face. "
|
||||
"Only used when speaker_selection is 'coordinates'.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_y",
|
||||
default=0,
|
||||
min=0,
|
||||
max=4096,
|
||||
advanced=True,
|
||||
tooltip="Y pixel coordinate of the speaker's face. "
|
||||
"Only used when speaker_selection is 'coordinates'.",
|
||||
),
|
||||
],
|
||||
)
|
||||
],
|
||||
tooltip="sync.so generation model.",
|
||||
),
|
||||
],
|
||||
outputs=[IO.Video.Output()],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
expr="""{"type":"usd","usd":0.19019,"format":{"approximate":true,"suffix":"/second"}}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
video: Input.Video,
|
||||
audio: Input.Audio,
|
||||
seed: int,
|
||||
model: dict,
|
||||
) -> IO.NodeOutput:
|
||||
try:
|
||||
width, height = video.get_dimensions()
|
||||
except Exception:
|
||||
width = height = None
|
||||
if width and height and (max(width, height) > 4096 or width * height > 4096 * 2160):
|
||||
raise ValueError(
|
||||
f"sync.so rejects videos above 4K (4096x2160); got {width}x{height}. Downscale the video first."
|
||||
)
|
||||
validate_audio_duration(audio, max_duration=600)
|
||||
|
||||
if model["speaker_selection"] == "auto-detect":
|
||||
speaker_detection = SyncActiveSpeakerDetection(auto_detect=True)
|
||||
elif model["speaker_selection"] == "coordinates":
|
||||
speaker_detection = SyncActiveSpeakerDetection(
|
||||
frame_number=model["speaker_frame"],
|
||||
coordinates=[model["speaker_x"], model["speaker_y"]],
|
||||
)
|
||||
else:
|
||||
speaker_detection = None
|
||||
|
||||
video_url = await upload_video_to_comfyapi(cls, video, max_duration=600)
|
||||
audio_url = await upload_audio_to_comfyapi(cls, audio)
|
||||
|
||||
generation = await sync_op(
|
||||
cls,
|
||||
ApiEndpoint(path="/proxy/synclabs/v2/generate", method="POST"),
|
||||
response_model=SyncGeneration,
|
||||
data=SyncGenerationRequest(
|
||||
model=model["model"],
|
||||
input=[
|
||||
SyncInputItem(type="video", url=video_url),
|
||||
SyncInputItem(type="audio", url=audio_url),
|
||||
],
|
||||
options=SyncGenerationOptions(
|
||||
sync_mode=model["sync_mode"],
|
||||
active_speaker_detection=speaker_detection,
|
||||
),
|
||||
),
|
||||
)
|
||||
generation = await poll_op(
|
||||
cls,
|
||||
ApiEndpoint(path=f"/proxy/synclabs/v2/generate/{generation.id}"),
|
||||
response_model=SyncGeneration,
|
||||
status_extractor=lambda g: g.status,
|
||||
completed_statuses=["COMPLETED", "FAILED", "REJECTED"],
|
||||
failed_statuses=[],
|
||||
queued_statuses=["PENDING"],
|
||||
poll_interval=10.0,
|
||||
)
|
||||
if generation.status != "COMPLETED":
|
||||
code = f" [{generation.errorCode}]" if generation.errorCode else ""
|
||||
raise ValueError(
|
||||
f"sync.so generation {generation.status.lower()}{code}: "
|
||||
f"{generation.error or 'no error details provided'}"
|
||||
)
|
||||
if not generation.outputUrl:
|
||||
raise ValueError("sync.so generation completed but no output URL was returned.")
|
||||
return IO.NodeOutput(await download_url_to_video_output(generation.outputUrl))
|
||||
|
||||
|
||||
class SyncTalkingImageNode(IO.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> IO.Schema:
|
||||
return IO.Schema(
|
||||
node_id="SyncTalkingImageNode",
|
||||
display_name="sync.so Talking Image",
|
||||
category="partner/video/sync.so",
|
||||
description=(
|
||||
"Animate a still portrait into a talking video driven by speech audio, "
|
||||
"using sync.so's sync-3 model. The output duration matches the audio. "
|
||||
"Cost scales with output duration."
|
||||
),
|
||||
inputs=[
|
||||
IO.Image.Input(
|
||||
"image",
|
||||
tooltip="A single image with a clearly visible face, up to 4K (4096x2160).",
|
||||
),
|
||||
IO.Audio.Input(
|
||||
"audio",
|
||||
tooltip="Speech audio driving the talking video; the output duration matches it. "
|
||||
"Chain any TTS node here to drive the animation from text.",
|
||||
),
|
||||
IO.String.Input(
|
||||
"prompt",
|
||||
multiline=True,
|
||||
default="",
|
||||
tooltip="Optional guidance for how the portrait comes to life, e.g. "
|
||||
"'make the subject smile and look at the camera'. "
|
||||
"Leave empty for natural talking motion.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"seed",
|
||||
default=0,
|
||||
min=0,
|
||||
max=2147483647,
|
||||
control_after_generate=True,
|
||||
tooltip="Seed controls whether the node should re-run; "
|
||||
"results are non-deterministic regardless of seed.",
|
||||
),
|
||||
IO.DynamicCombo.Input(
|
||||
"model",
|
||||
options=[
|
||||
IO.DynamicCombo.Option(
|
||||
"sync-3",
|
||||
[
|
||||
IO.Combo.Input(
|
||||
"speaker_selection",
|
||||
options=["default", "coordinates"],
|
||||
default="default",
|
||||
tooltip=(
|
||||
"Which face to animate when several people are visible. "
|
||||
"default: let the model decide. "
|
||||
"coordinates: target the face at pixel (speaker_x, speaker_y) "
|
||||
"in the image. Auto-detection is not supported for images."
|
||||
),
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_x",
|
||||
default=0,
|
||||
min=0,
|
||||
max=4096,
|
||||
advanced=True,
|
||||
tooltip="X pixel coordinate of the speaker's face. "
|
||||
"Only used when speaker_selection is 'coordinates'.",
|
||||
),
|
||||
IO.Int.Input(
|
||||
"speaker_y",
|
||||
default=0,
|
||||
min=0,
|
||||
max=4096,
|
||||
advanced=True,
|
||||
tooltip="Y pixel coordinate of the speaker's face. "
|
||||
"Only used when speaker_selection is 'coordinates'.",
|
||||
),
|
||||
IO.Boolean.Input(
|
||||
"auto_downscale",
|
||||
default=True,
|
||||
advanced=True,
|
||||
tooltip="Automatically downscale the image if it exceeds the 4K "
|
||||
"(4096x2160) input limit; speaker coordinates are scaled to match. "
|
||||
"When disabled, an oversized image raises an error instead.",
|
||||
),
|
||||
],
|
||||
)
|
||||
],
|
||||
tooltip="sync.so generation model. Image input is exclusive to sync-3.",
|
||||
),
|
||||
],
|
||||
outputs=[IO.Video.Output()],
|
||||
hidden=[
|
||||
IO.Hidden.auth_token_comfy_org,
|
||||
IO.Hidden.api_key_comfy_org,
|
||||
IO.Hidden.unique_id,
|
||||
],
|
||||
is_api_node=True,
|
||||
price_badge=IO.PriceBadge(
|
||||
expr="""{"type":"usd","usd":0.19019,"format":{"approximate":true,"suffix":"/second"}}""",
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def execute(
|
||||
cls,
|
||||
image: Input.Image,
|
||||
audio: Input.Audio,
|
||||
prompt: str,
|
||||
seed: int,
|
||||
model: dict,
|
||||
) -> IO.NodeOutput:
|
||||
if get_number_of_images(image) != 1:
|
||||
raise ValueError("Exactly one image is required; got a batch. Pick one frame first.")
|
||||
validate_audio_duration(audio, max_duration=600)
|
||||
|
||||
height, width = get_image_dimensions(image)
|
||||
speaker_x, speaker_y = model["speaker_x"], model["speaker_y"]
|
||||
if max(width, height) > 4096 or width * height > 4096 * 2160:
|
||||
if not model["auto_downscale"]:
|
||||
raise ValueError(
|
||||
f"sync.so rejects images above 4K (4096x2160); got {width}x{height}. "
|
||||
"Downscale the image first or enable auto_downscale."
|
||||
)
|
||||
image = downscale_image_tensor(image, total_pixels=4096 * 2160)
|
||||
image = downscale_image_tensor_by_max_side(image, max_side=4096)
|
||||
new_height, new_width = get_image_dimensions(image)
|
||||
# speaker coordinates are given in the original image's pixel space
|
||||
speaker_x = min(new_width - 1, round(speaker_x * new_width / width))
|
||||
speaker_y = min(new_height - 1, round(speaker_y * new_height / height))
|
||||
|
||||
if model["speaker_selection"] == "coordinates":
|
||||
speaker_detection = SyncActiveSpeakerDetection(
|
||||
frame_number=0, # images have a single frame; auto_detect is rejected by the API
|
||||
coordinates=[speaker_x, speaker_y],
|
||||
)
|
||||
else:
|
||||
speaker_detection = None
|
||||
|
||||
image_url = await upload_image_to_comfyapi(cls, image, mime_type="image/png", total_pixels=None)
|
||||
audio_url = await upload_audio_to_comfyapi(cls, audio)
|
||||
|
||||
generation = await sync_op(
|
||||
cls,
|
||||
ApiEndpoint(path="/proxy/synclabs/v2/generate", method="POST"),
|
||||
response_model=SyncGeneration,
|
||||
data=SyncGenerationRequest(
|
||||
model=model["model"],
|
||||
input=[
|
||||
SyncInputItem(type="image", url=image_url),
|
||||
SyncInputItem(type="audio", url=audio_url),
|
||||
],
|
||||
options=SyncGenerationOptions(
|
||||
i2v_prompt=prompt.strip() or None,
|
||||
active_speaker_detection=speaker_detection,
|
||||
),
|
||||
),
|
||||
)
|
||||
generation = await poll_op(
|
||||
cls,
|
||||
ApiEndpoint(path=f"/proxy/synclabs/v2/generate/{generation.id}"),
|
||||
response_model=SyncGeneration,
|
||||
status_extractor=lambda g: g.status,
|
||||
completed_statuses=["COMPLETED", "FAILED", "REJECTED"],
|
||||
failed_statuses=[],
|
||||
queued_statuses=["PENDING"],
|
||||
poll_interval=10.0,
|
||||
)
|
||||
if generation.status != "COMPLETED":
|
||||
code = f" [{generation.errorCode}]" if generation.errorCode else ""
|
||||
raise ValueError(
|
||||
f"sync.so generation {generation.status.lower()}{code}: "
|
||||
f"{generation.error or 'no error details provided'}"
|
||||
)
|
||||
if not generation.outputUrl:
|
||||
raise ValueError("sync.so generation completed but no output URL was returned.")
|
||||
return IO.NodeOutput(await download_url_to_video_output(generation.outputUrl))
|
||||
|
||||
|
||||
class SyncExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[IO.ComfyNode]]:
|
||||
return [
|
||||
SyncLipSyncNode,
|
||||
SyncTalkingImageNode,
|
||||
]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> SyncExtension:
|
||||
return SyncExtension()
|
||||
@@ -15,8 +15,6 @@ from comfy.comfy_api_env import normalize_comfy_api_base
|
||||
from comfy.deploy_environment import get_deploy_environment
|
||||
from comfy.model_management import processing_interrupted
|
||||
from comfy_api.latest import IO
|
||||
from comfy_execution.utils import get_executing_context
|
||||
from comfyui_version import __version__ as comfyui_version
|
||||
|
||||
from .common_exceptions import ProcessingInterrupted
|
||||
|
||||
@@ -58,16 +56,11 @@ def get_comfy_api_headers(node_cls: type[IO.ComfyNode]) -> dict[str, str]:
|
||||
relative/cloud URLs resolved against ``default_base_url()``; because the result
|
||||
includes auth, callers must not attach it to arbitrary absolute/presigned URLs.
|
||||
"""
|
||||
headers = {
|
||||
return {
|
||||
**get_auth_header(node_cls),
|
||||
"Comfy-Env": get_deploy_environment(),
|
||||
"Comfy-Usage-Source": get_usage_source(node_cls),
|
||||
"Comfy-Core-Version": comfyui_version,
|
||||
}
|
||||
ctx = get_executing_context()
|
||||
if ctx is not None:
|
||||
headers["Comfy-Job-Id"] = ctx.prompt_id
|
||||
return headers
|
||||
|
||||
|
||||
def default_base_url() -> str:
|
||||
|
||||
@@ -2,7 +2,6 @@ import logging
|
||||
import os
|
||||
import json
|
||||
|
||||
import av
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
@@ -10,7 +9,7 @@ from typing_extensions import override
|
||||
|
||||
import folder_paths
|
||||
import node_helpers
|
||||
from comfy_api.latest import ComfyExtension, io, Input, InputImpl, Types
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
|
||||
|
||||
def load_and_process_images(image_files, input_dir):
|
||||
@@ -43,38 +42,6 @@ def load_and_process_images(image_files, input_dir):
|
||||
return output_images
|
||||
|
||||
|
||||
VALID_VIDEO_EXTENSIONS = [".mp4", ".avi", ".mov", ".webm", ".mkv", ".flv"]
|
||||
|
||||
|
||||
def _decode_selected_frames(video: Input.Video, indices: list[int]) -> Input.Video:
|
||||
"""Decode only the requested frame indices from a video.
|
||||
|
||||
Opens the underlying container once, decodes frames in presentation order,
|
||||
keeps only the ones whose index is in ``indices``, and returns the result
|
||||
wrapped in a VideoFromComponents so it still satisfies the VideoInput
|
||||
contract for downstream nodes.
|
||||
"""
|
||||
indices_sorted = sorted(set(indices))
|
||||
max_idx = indices_sorted[-1]
|
||||
source = video.get_stream_source()
|
||||
|
||||
frames_by_idx: dict[int, torch.Tensor] = {}
|
||||
with av.open(source, mode="r") as container:
|
||||
stream = container.streams.video[0]
|
||||
wanted = set(indices_sorted)
|
||||
for frame_idx, frame in enumerate(container.decode(stream)):
|
||||
if frame_idx in wanted:
|
||||
img = frame.to_ndarray(format="rgb24")
|
||||
frames_by_idx[frame_idx] = torch.from_numpy(img.copy()).float() / 255.0
|
||||
if frame_idx >= max_idx:
|
||||
break
|
||||
|
||||
stacked = torch.stack([frames_by_idx[i] for i in indices])
|
||||
return InputImpl.VideoFromComponents(
|
||||
Types.VideoComponents(images=stacked, frame_rate=video.get_frame_rate())
|
||||
)
|
||||
|
||||
|
||||
class LoadImageDataSetFromFolderNode(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
@@ -190,116 +157,6 @@ class LoadImageTextDataSetFromFolderNode(io.ComfyNode):
|
||||
return io.NodeOutput(output_tensor, captions)
|
||||
|
||||
|
||||
class LoadVideoDataSetFromFolderNode(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LoadVideoDataSetFromFolder",
|
||||
search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"],
|
||||
display_name="Load Video (from Folder)",
|
||||
category="video",
|
||||
description="Load a dataset of videos from a specified folder and return a list of videos. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.Combo.Input(
|
||||
"folder",
|
||||
options=folder_paths.get_input_subfolders(),
|
||||
tooltip="The folder containing video files.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(
|
||||
display_name="videos",
|
||||
is_output_list=True,
|
||||
tooltip="Lazy video references; frames are decoded only when needed downstream.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, folder):
|
||||
sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder)
|
||||
video_files = sorted([
|
||||
f for f in os.listdir(sub_input_dir)
|
||||
if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS)
|
||||
])
|
||||
|
||||
if not video_files:
|
||||
raise ValueError(f"No video files found in {sub_input_dir}")
|
||||
|
||||
videos = [InputImpl.VideoFromFile(os.path.join(sub_input_dir, f)) for f in video_files]
|
||||
logging.info(f"Loaded {len(videos)} lazy video references from {sub_input_dir}")
|
||||
return io.NodeOutput(videos)
|
||||
|
||||
|
||||
class LoadVideoTextDataSetFromFolderNode(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="LoadVideoTextDataSetFromFolder",
|
||||
search_aliases=["load folder", "load from folder", "load dataset", "load videos", "import dataset"],
|
||||
display_name="Load Video-Text (from Folder)",
|
||||
category="video",
|
||||
description="Load a dataset of pairs of videos and text captions from a specified folder and return them as a list. Supported formats: MP4, AVI, MOV, WEBM, MKV, FLV.",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.Combo.Input(
|
||||
"folder",
|
||||
options=folder_paths.get_input_subfolders(),
|
||||
tooltip="The folder containing video files and .txt captions.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(
|
||||
display_name="videos",
|
||||
is_output_list=True,
|
||||
tooltip="Lazy video references; frames are decoded only when needed downstream.",
|
||||
),
|
||||
io.String.Output(
|
||||
display_name="texts",
|
||||
is_output_list=True,
|
||||
tooltip="List of text captions.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, folder):
|
||||
sub_input_dir = os.path.join(folder_paths.get_input_directory(), folder)
|
||||
|
||||
video_files = []
|
||||
for item in sorted(os.listdir(sub_input_dir)):
|
||||
path = os.path.join(sub_input_dir, item)
|
||||
if any(item.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS):
|
||||
video_files.append(path)
|
||||
elif os.path.isdir(path):
|
||||
# Support kohya-ss/sd-scripts folder structure: {repeat}_{desc}/
|
||||
repeat = 1
|
||||
if item.split("_")[0].isdigit():
|
||||
repeat = int(item.split("_")[0])
|
||||
video_files.extend([
|
||||
os.path.join(path, f)
|
||||
for f in sorted(os.listdir(path))
|
||||
if any(f.lower().endswith(ext) for ext in VALID_VIDEO_EXTENSIONS)
|
||||
] * repeat)
|
||||
|
||||
if not video_files:
|
||||
raise ValueError(f"No video files found in {sub_input_dir}")
|
||||
|
||||
captions = []
|
||||
for vf in video_files:
|
||||
caption_path = os.path.splitext(vf)[0] + ".txt"
|
||||
if os.path.exists(caption_path):
|
||||
with open(caption_path, "r", encoding="utf-8") as f:
|
||||
captions.append(f.read().strip())
|
||||
else:
|
||||
captions.append("")
|
||||
|
||||
videos = [InputImpl.VideoFromFile(vf) for vf in video_files]
|
||||
logging.info(f"Loaded {len(videos)} lazy video references with captions from {sub_input_dir}")
|
||||
return io.NodeOutput(videos, captions)
|
||||
|
||||
|
||||
def save_images_to_folder(image_list, output_dir, prefix="image", overwrite=True):
|
||||
"""Utility function to save a list of image tensors to disk.
|
||||
|
||||
@@ -613,15 +470,7 @@ class ImageProcessingNode(io.ComfyNode):
|
||||
|
||||
@classmethod
|
||||
def execute(cls, images, **kwargs):
|
||||
"""Execute the node. Routes to _process or _group_process based on mode.
|
||||
|
||||
For individual processing (_process), automatically handles multi-frame
|
||||
inputs (video tensors [T, H, W, C]) by applying _process per-frame and
|
||||
concatenating the results. This allows all spatial transform nodes to
|
||||
work with video without modification. Nodes that natively handle batched
|
||||
tensors (e.g. pure tensor math) can set per_frame_process = False to
|
||||
skip the per-frame loop.
|
||||
"""
|
||||
"""Execute the node. Routes to _process or _group_process based on mode."""
|
||||
is_group = cls._detect_processing_mode()
|
||||
|
||||
if is_group:
|
||||
@@ -640,16 +489,7 @@ class ImageProcessingNode(io.ComfyNode):
|
||||
result = cls._group_process(images, **params)
|
||||
else:
|
||||
# Individual processing: images is single item, call _process
|
||||
# Auto-loop over frames for multi-frame inputs (video [T, H, W, C])
|
||||
# so that PIL-based spatial transforms work per-frame automatically.
|
||||
if images.shape[0] > 1 and getattr(cls, 'per_frame_process', True):
|
||||
results = []
|
||||
for i in range(images.shape[0]):
|
||||
frame_result = cls._process(images[i:i + 1], **params)
|
||||
results.append(frame_result)
|
||||
result = torch.cat(results, dim=0)
|
||||
else:
|
||||
result = cls._process(images, **params)
|
||||
result = cls._process(images, **params)
|
||||
|
||||
return io.NodeOutput(result)
|
||||
|
||||
@@ -963,7 +803,6 @@ class NormalizeImagesNode(ImageProcessingNode):
|
||||
display_name = "Normalize Image Colors"
|
||||
category = "image/color"
|
||||
description = "Normalize images using mean and standard deviation."
|
||||
per_frame_process = False # Pure tensor math, handles any batch size
|
||||
extra_inputs = [
|
||||
io.Float.Input(
|
||||
"mean",
|
||||
@@ -994,7 +833,6 @@ class AdjustBrightnessNode(ImageProcessingNode):
|
||||
display_name = "Adjust Brightness"
|
||||
category="image/adjustments"
|
||||
description = "Adjust the brightness of an image."
|
||||
per_frame_process = False # Pure tensor math, handles any batch size
|
||||
extra_inputs = [
|
||||
io.Float.Input(
|
||||
"factor",
|
||||
@@ -1016,7 +854,6 @@ class AdjustContrastNode(ImageProcessingNode):
|
||||
display_name = "Adjust Contrast"
|
||||
category="image/adjustments"
|
||||
description = "Adjust the contrast of an image."
|
||||
per_frame_process = False # Pure tensor math, handles any batch size
|
||||
extra_inputs = [
|
||||
io.Float.Input(
|
||||
"factor",
|
||||
@@ -1098,261 +935,6 @@ class ShuffleImageTextDatasetNode(io.ComfyNode):
|
||||
return io.NodeOutput(shuffled_images, shuffled_texts)
|
||||
|
||||
|
||||
# ========== Video Processing Nodes ==========
|
||||
|
||||
|
||||
class VideoFrameSampleNode(io.ComfyNode):
|
||||
"""Sample a fixed number of frames from a video using various strategies.
|
||||
|
||||
For contiguous strategies ("head"/"tail") the result is a fully lazy
|
||||
VideoInput (no frames decoded). For non-contiguous strategies
|
||||
("uniform"/"random") only the selected indices are decoded.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="VideoFrameSample",
|
||||
search_aliases=["sample frames", "extract frames"],
|
||||
display_name="Sample Video Frame",
|
||||
category="video",
|
||||
description="Sample a fixed number of frames from a video using various strategies.",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.Video.Input("video", tooltip="Input video."),
|
||||
io.Int.Input(
|
||||
"num_frames",
|
||||
default=16,
|
||||
min=1,
|
||||
max=9999,
|
||||
tooltip="Number of frames to sample.",
|
||||
),
|
||||
io.Combo.Input(
|
||||
"strategy",
|
||||
options=["uniform", "head", "tail", "random"],
|
||||
default="uniform",
|
||||
tooltip="uniform: evenly spaced, head: first N, tail: last N, random: random sorted.",
|
||||
),
|
||||
io.Int.Input(
|
||||
"seed",
|
||||
default=0,
|
||||
min=0,
|
||||
max=0xFFFFFFFFFFFFFFFF,
|
||||
tooltip="Random seed (only used with 'random' strategy).",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(display_name="video", tooltip="Sampled video."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, video, num_frames, strategy, seed):
|
||||
total_frames = video.get_frame_count()
|
||||
num_frames = min(num_frames, total_frames)
|
||||
fps = float(video.get_frame_rate())
|
||||
|
||||
if strategy == "head":
|
||||
return io.NodeOutput(
|
||||
video.as_trimmed(0.0, num_frames / fps, strict_duration=False)
|
||||
)
|
||||
if strategy == "tail":
|
||||
start_t = (total_frames - num_frames) / fps
|
||||
return io.NodeOutput(
|
||||
video.as_trimmed(start_t, num_frames / fps, strict_duration=False)
|
||||
)
|
||||
|
||||
if strategy == "uniform":
|
||||
if num_frames == 1:
|
||||
indices = [total_frames // 2]
|
||||
else:
|
||||
indices = [round(i * (total_frames - 1) / (num_frames - 1)) for i in range(num_frames)]
|
||||
elif strategy == "random":
|
||||
rng = np.random.RandomState(seed % (2**32 - 1))
|
||||
indices = sorted(rng.choice(total_frames, size=num_frames, replace=False).tolist())
|
||||
else:
|
||||
raise ValueError(f"Unknown strategy: {strategy}")
|
||||
|
||||
return io.NodeOutput(_decode_selected_frames(video, indices))
|
||||
|
||||
|
||||
class VideoTemporalCropNode(io.ComfyNode):
|
||||
"""Crop a continuous range of frames from a video (fully lazy)."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="VideoTemporalCrop",
|
||||
search_aliases=["crop", "crop video", "temporal crop", "truncate video"],
|
||||
display_name="Crop Video (Temporal)",
|
||||
category="video/transform",
|
||||
description="Crop a continuous range of frames from a video.",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.Video.Input("video", tooltip="Input video."),
|
||||
io.Int.Input(
|
||||
"start_frame",
|
||||
default=0,
|
||||
min=0,
|
||||
max=99999,
|
||||
tooltip="Starting frame index.",
|
||||
),
|
||||
io.Int.Input(
|
||||
"length",
|
||||
default=16,
|
||||
min=1,
|
||||
max=99999,
|
||||
tooltip="Number of frames to keep.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, video, start_frame, length):
|
||||
total_frames = video.get_frame_count()
|
||||
fps = float(video.get_frame_rate())
|
||||
start_frame = min(start_frame, max(total_frames - 1, 0))
|
||||
length = min(length, total_frames - start_frame)
|
||||
return io.NodeOutput(
|
||||
video.as_trimmed(start_frame / fps, length / fps, strict_duration=False)
|
||||
)
|
||||
|
||||
|
||||
class VideoRandomTemporalCropNode(io.ComfyNode):
|
||||
"""Randomly crop a continuous range of frames from a video (fully lazy)."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="VideoRandomTemporalCrop",
|
||||
search_aliases=["crop", "crop video", "temporal crop", "truncate video", "random crop"],
|
||||
display_name="Crop Video (Temporal Random)",
|
||||
category="video/transform",
|
||||
description="Randomly crop a continuous range of frames from a video.",
|
||||
is_experimental=True,
|
||||
inputs=[
|
||||
io.Video.Input("video", tooltip="Input video."),
|
||||
io.Int.Input(
|
||||
"length",
|
||||
default=16,
|
||||
min=1,
|
||||
max=99999,
|
||||
tooltip="Number of frames to keep.",
|
||||
),
|
||||
io.Int.Input(
|
||||
"seed",
|
||||
default=0,
|
||||
min=0,
|
||||
max=0xFFFFFFFFFFFFFFFF,
|
||||
tooltip="Random seed.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(display_name="video", tooltip="Cropped video (lazy)."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, video, length, seed):
|
||||
total_frames = video.get_frame_count()
|
||||
fps = float(video.get_frame_rate())
|
||||
length = min(length, total_frames)
|
||||
max_start = total_frames - length
|
||||
rng = np.random.RandomState(seed % (2**32 - 1))
|
||||
start = rng.randint(0, max_start + 1) if max_start > 0 else 0
|
||||
return io.NodeOutput(
|
||||
video.as_trimmed(start / fps, length / fps, strict_duration=False)
|
||||
)
|
||||
|
||||
|
||||
class ShuffleVideoDatasetNode(io.ComfyNode):
|
||||
"""Randomly shuffle the order of videos in the dataset."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ShuffleVideoDataset",
|
||||
search_aliases=["shuffle", "randomize", "mix"],
|
||||
display_name="Shuffle Videos List",
|
||||
category="video/batch",
|
||||
description="Randomly shuffle the order of videos in a list.",
|
||||
is_experimental=True,
|
||||
is_input_list=True,
|
||||
inputs=[
|
||||
io.Video.Input("videos", tooltip="List of videos to shuffle."),
|
||||
io.Int.Input(
|
||||
"seed", default=0, min=0, max=0xFFFFFFFFFFFFFFFF, tooltip="Random seed."
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(
|
||||
display_name="videos",
|
||||
is_output_list=True,
|
||||
tooltip="Shuffled videos",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, videos, seed):
|
||||
seed = seed[0] if isinstance(seed, list) else seed
|
||||
np.random.seed(seed % (2**32 - 1))
|
||||
indices = np.random.permutation(len(videos))
|
||||
return io.NodeOutput([videos[i] for i in indices])
|
||||
|
||||
|
||||
class ShuffleVideoTextDatasetNode(io.ComfyNode):
|
||||
"""Shuffle videos and their captions together, preserving pairs."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="ShuffleVideoTextDataset",
|
||||
search_aliases=["shuffle", "randomize", "mix"],
|
||||
display_name="Shuffle Pairs of Video-Text",
|
||||
category="dataset/video",
|
||||
description="Randomly shuffle the order of pairs of video-text in a list.",
|
||||
is_experimental=True,
|
||||
is_input_list=True,
|
||||
inputs=[
|
||||
io.Video.Input("videos", tooltip="List of videos to shuffle."),
|
||||
io.String.Input("texts", tooltip="List of texts to shuffle."),
|
||||
io.Int.Input(
|
||||
"seed",
|
||||
default=0,
|
||||
min=0,
|
||||
max=0xFFFFFFFFFFFFFFFF,
|
||||
tooltip="Random seed.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Video.Output(
|
||||
display_name="videos",
|
||||
is_output_list=True,
|
||||
tooltip="Shuffled videos",
|
||||
),
|
||||
io.String.Output(
|
||||
display_name="texts",
|
||||
is_output_list=True,
|
||||
tooltip="Shuffled texts",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, videos, texts, seed):
|
||||
seed = seed[0] if isinstance(seed, list) else seed
|
||||
np.random.seed(seed % (2**32 - 1))
|
||||
indices = np.random.permutation(len(videos))
|
||||
return io.NodeOutput(
|
||||
[videos[i] for i in indices],
|
||||
[texts[i] for i in indices],
|
||||
)
|
||||
|
||||
|
||||
# ========== Text Transform Nodes ==========
|
||||
|
||||
|
||||
@@ -2026,10 +1608,7 @@ class DatasetExtension(ComfyExtension):
|
||||
LoadImageTextDataSetFromFolderNode,
|
||||
SaveImageDataSetToFolderNode,
|
||||
SaveImageTextDataSetToFolderNode,
|
||||
# Video data loading nodes
|
||||
LoadVideoDataSetFromFolderNode,
|
||||
LoadVideoTextDataSetFromFolderNode,
|
||||
# Image transform nodes (auto-handle video via per-frame processing)
|
||||
# Image transform nodes
|
||||
ResizeImagesByShorterEdgeNode,
|
||||
ResizeImagesByLongerEdgeNode,
|
||||
CenterCropImagesNode,
|
||||
@@ -2039,12 +1618,6 @@ class DatasetExtension(ComfyExtension):
|
||||
AdjustContrastNode,
|
||||
ShuffleDatasetNode,
|
||||
ShuffleImageTextDatasetNode,
|
||||
# Video processing nodes (lazy VideoInput in/out)
|
||||
VideoFrameSampleNode,
|
||||
VideoTemporalCropNode,
|
||||
VideoRandomTemporalCropNode,
|
||||
ShuffleVideoDatasetNode,
|
||||
ShuffleVideoTextDatasetNode,
|
||||
# Text transform nodes
|
||||
TextToLowercaseNode,
|
||||
TextToUppercaseNode,
|
||||
|
||||
@@ -10,7 +10,7 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
|
||||
Apply frequency-dependent scaling to an image tensor using Fourier transforms.
|
||||
|
||||
Parameters:
|
||||
x: Input tensor of shape (..., H, W)
|
||||
x: Input tensor of shape (B, C, H, W)
|
||||
scale_low: Scaling factor for low-frequency components (default: 1.0)
|
||||
scale_high: Scaling factor for high-frequency components (default: 1.5)
|
||||
freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20)
|
||||
@@ -31,8 +31,8 @@ def Fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
|
||||
# Initialize mask with high-frequency scaling factor
|
||||
mask = torch.ones(x_freq.shape, device=device) * scale_high
|
||||
m = mask
|
||||
for d in range(2):
|
||||
dim = len(x_freq.shape) - 2 + d
|
||||
for d in range(len(x_freq.shape) - 2):
|
||||
dim = d + 2
|
||||
cc = x_freq.shape[dim] // 2
|
||||
f_c = min(freq_cutoff, cc)
|
||||
m = m.narrow(dim, cc - f_c, f_c * 2)
|
||||
|
||||
@@ -844,18 +844,15 @@ class ImageMergeTileList(IO.ComfyNode):
|
||||
# Format specifications
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Maps (file_format, bit_depth, num_channels) -> (quantization scale, numpy dtype,
|
||||
# av frame pix_fmt, stream pix_fmt). Keeps the encode path declarative instead of branchy.
|
||||
# Maps (file_format, bit_depth, has_alpha) -> (numpy dtype scale, av pixel format,
|
||||
# stream pix_fmt). Keeps the encode path declarative instead of branchy.
|
||||
_FORMAT_SPECS = {
|
||||
("png", "8-bit", 1): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "gray", "stream_fmt": "gray"},
|
||||
("png", "8-bit", 3): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgb24", "stream_fmt": "rgb24"},
|
||||
("png", "8-bit", 4): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgba", "stream_fmt": "rgba"},
|
||||
("png", "16-bit", 1): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "gray16le", "stream_fmt": "gray16be"},
|
||||
("png", "16-bit", 3): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgb48le", "stream_fmt": "rgb48be"},
|
||||
("png", "16-bit", 4): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgba64le", "stream_fmt": "rgba64be"},
|
||||
("exr", "32-bit float", 1): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "grayf32le", "stream_fmt": "grayf32le"},
|
||||
("exr", "32-bit float", 3): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrpf32le", "stream_fmt": "gbrpf32le"},
|
||||
("exr", "32-bit float", 4): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrapf32le", "stream_fmt": "gbrapf32le"},
|
||||
("png", "8-bit", False): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgb24", "stream_fmt": "rgb24"},
|
||||
("png", "8-bit", True): {"scale": 255.0, "dtype": np.uint8, "frame_fmt": "rgba", "stream_fmt": "rgba"},
|
||||
("png", "16-bit", False): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgb48le", "stream_fmt": "rgb48be"},
|
||||
("png", "16-bit", True): {"scale": 65535.0, "dtype": np.uint16, "frame_fmt": "rgba64le", "stream_fmt": "rgba64be"},
|
||||
("exr", "32-bit float", False): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrpf32le", "stream_fmt": "gbrpf32le"},
|
||||
("exr", "32-bit float", True): {"scale": 1.0, "dtype": np.float32, "frame_fmt": "gbrapf32le", "stream_fmt": "gbrapf32le"},
|
||||
}
|
||||
|
||||
|
||||
@@ -894,11 +891,10 @@ def hlg_to_linear(t: torch.Tensor) -> torch.Tensor:
|
||||
return torch.cat([hlg_to_linear(rgb), alpha], dim=-1)
|
||||
|
||||
# Piecewise: sqrt branch below 0.5, log branch above.
|
||||
# Clamp the log branch at the 0.5 branch point (not above it) so the
|
||||
# unselected lane stays finite in exp() without altering selected values;
|
||||
# Clamp inside the log branch so negative / out-of-range values don't blow up;
|
||||
# values above 1.0 are allowed and extrapolate naturally.
|
||||
low = (t ** 2) / 3.0
|
||||
high = (torch.exp((t.clamp(min=0.5) - _HLG_C) / _HLG_A) + _HLG_B) / 12.0
|
||||
high = (torch.exp((t.clamp(min=_HLG_C) - _HLG_C) / _HLG_A) + _HLG_B) / 12.0
|
||||
return torch.where(t <= 0.5, low, high)
|
||||
|
||||
|
||||
@@ -1091,8 +1087,7 @@ def _encode_image(
|
||||
bit_depth: str,
|
||||
colorspace: str,
|
||||
) -> bytes:
|
||||
"""Encode a single HxWxC (or channel-less HxW grayscale) tensor to PNG or
|
||||
EXR bytes in memory. Grayscale is written as single-channel PNG / Y-only EXR.
|
||||
"""Encode a single HxWxC tensor to PNG or EXR bytes in memory.
|
||||
|
||||
For EXR the input is interpreted according to `colorspace` and converted
|
||||
to scene-linear (EXR's convention) before writing:
|
||||
@@ -1106,16 +1101,10 @@ def _encode_image(
|
||||
For PNG, colorspace selection does not modify pixels — PNG is delivered
|
||||
sRGB-encoded and there is no PNG path for wide-gamut HDR in this node.
|
||||
"""
|
||||
if img_tensor.ndim == 2:
|
||||
img_tensor = img_tensor.unsqueeze(-1) # Some nodes emit grayscale as (H, W) with no channel dim, mask-style.
|
||||
height, width, num_channels = img_tensor.shape
|
||||
has_alpha = num_channels == 4
|
||||
|
||||
spec = _FORMAT_SPECS.get((file_format, bit_depth, num_channels))
|
||||
if spec is None:
|
||||
raise ValueError(
|
||||
f"No {file_format}/{bit_depth} encoder for {num_channels}-channel images: "
|
||||
"supported channel counts are 1 (grayscale), 3 (RGB) and 4 (RGBA)."
|
||||
)
|
||||
spec = _FORMAT_SPECS[(file_format, bit_depth, has_alpha)]
|
||||
|
||||
if spec["dtype"] == np.float32:
|
||||
# EXR path: preserve full range, no clamp.
|
||||
|
||||
@@ -1,102 +0,0 @@
|
||||
from typing_extensions import override
|
||||
|
||||
import comfy.utils
|
||||
import node_helpers
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
|
||||
|
||||
# fmt: off
|
||||
BUCKETS_1024 = [
|
||||
(512, 1792), (512, 1856), (512, 1920), (512, 1984), (512, 2048),
|
||||
(576, 1600), (576, 1664), (576, 1728), (576, 1792),
|
||||
(640, 1472), (640, 1536), (640, 1600),
|
||||
(704, 1344), (704, 1408), (704, 1472),
|
||||
(768, 1216), (768, 1280), (768, 1344),
|
||||
(832, 1152), (832, 1216),
|
||||
(896, 1088), (896, 1152),
|
||||
(960, 1024), (960, 1088),
|
||||
(1024, 960), (1024, 1024),
|
||||
(1088, 896), (1088, 960),
|
||||
(1152, 832), (1152, 896),
|
||||
(1216, 768), (1216, 832),
|
||||
(1280, 768),
|
||||
(1344, 704), (1344, 768),
|
||||
(1408, 704),
|
||||
(1472, 640), (1472, 704),
|
||||
(1536, 640),
|
||||
(1600, 576), (1600, 640),
|
||||
(1664, 576),
|
||||
(1728, 576),
|
||||
(1792, 512), (1792, 576),
|
||||
(1856, 512),
|
||||
(1920, 512),
|
||||
(1984, 512),
|
||||
(2048, 512),
|
||||
]
|
||||
# fmt: on
|
||||
|
||||
|
||||
def _find_best_bucket(height: int, width: int) -> tuple[int, int]:
|
||||
target_ratio = height / width
|
||||
return min(BUCKETS_1024, key=lambda hw: abs(hw[0] / hw[1] - target_ratio))
|
||||
|
||||
|
||||
def _resize_reference(image):
|
||||
if image.shape[0] != 1:
|
||||
raise ValueError("JoyImage reference inputs must contain one image each")
|
||||
samples = image.movedim(-1, 1)
|
||||
bucket_h, bucket_w = _find_best_bucket(samples.shape[2], samples.shape[3])
|
||||
resized = comfy.utils.common_upscale(samples, bucket_w, bucket_h, "bilinear", "center")
|
||||
return resized.movedim(1, -1)[:, :, :, :3]
|
||||
|
||||
|
||||
def _encode(clip, prompt, vae, images):
|
||||
resized_images = [_resize_reference(image) for image in images]
|
||||
conditioning = clip.encode_from_tokens_scheduled(clip.tokenize(prompt, images=resized_images))
|
||||
if vae is not None and resized_images:
|
||||
ref_latents = [vae.encode(image) for image in resized_images]
|
||||
conditioning = node_helpers.conditioning_set_values(
|
||||
conditioning, {"reference_latents": ref_latents}, append=True,
|
||||
)
|
||||
return conditioning
|
||||
|
||||
|
||||
class TextEncodeJoyImageEdit(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
image_template = io.Autogrow.TemplatePrefix(
|
||||
io.Image.Input("image"),
|
||||
prefix="image",
|
||||
min=0,
|
||||
max=6,
|
||||
)
|
||||
return io.Schema(
|
||||
node_id="TextEncodeJoyImageEdit",
|
||||
category="model/conditioning/joyimage",
|
||||
inputs=[
|
||||
io.Clip.Input("clip"),
|
||||
io.String.Input("prompt", multiline=True, dynamic_prompts=True),
|
||||
io.Vae.Input("vae", optional=True),
|
||||
io.Autogrow.Input("images", template=image_template, optional=True),
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output(),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, clip, prompt, vae=None, images: io.Autogrow.Type = None) -> io.NodeOutput:
|
||||
images = images or {}
|
||||
return io.NodeOutput(_encode(clip, prompt, vae, list(images.values())))
|
||||
|
||||
|
||||
class JoyImageExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [
|
||||
TextEncodeJoyImageEdit,
|
||||
]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> JoyImageExtension:
|
||||
return JoyImageExtension()
|
||||
@@ -174,9 +174,8 @@ class Preview3DAdvanced(IO.ComfyNode):
|
||||
filename = f"preview3d_advanced_{uuid.uuid4().hex}.{model_3d.format}"
|
||||
model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename))
|
||||
|
||||
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
|
||||
camera_info_input = kwargs.get("camera_info", None)
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info']
|
||||
model_3d_info_input = kwargs.get("model_3d_info", None)
|
||||
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
|
||||
return IO.NodeOutput(
|
||||
@@ -244,9 +243,8 @@ class PreviewGaussianSplat(IO.ComfyNode):
|
||||
filename = f"preview_splat_{uuid.uuid4().hex}.{model_3d.format}"
|
||||
model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename))
|
||||
|
||||
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
|
||||
camera_info_input = kwargs.get("camera_info", None)
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info']
|
||||
model_3d_info_input = kwargs.get("model_3d_info", None)
|
||||
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
|
||||
return IO.NodeOutput(
|
||||
@@ -305,9 +303,8 @@ class PreviewPointCloud(IO.ComfyNode):
|
||||
filename = f"preview_pointcloud_{uuid.uuid4().hex}.{model_3d.format}"
|
||||
model_3d.save_to(os.path.join(folder_paths.get_temp_directory(), filename))
|
||||
|
||||
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
|
||||
camera_info_input = kwargs.get("camera_info", None)
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info']
|
||||
model_3d_info_input = kwargs.get("model_3d_info", None)
|
||||
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
|
||||
return IO.NodeOutput(
|
||||
@@ -378,9 +375,8 @@ class Load3DAdvanced(IO.ComfyNode):
|
||||
file_3d = None
|
||||
if model_file and model_file != "none":
|
||||
file_3d = Types.File3D(folder_paths.get_annotated_filepath(model_file))
|
||||
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
|
||||
model_3d_info = viewport_state.get('model_3d_info', [])
|
||||
return IO.NodeOutput(file_3d, model_3d_info, viewport_state.get('camera_info'), width, height)
|
||||
return IO.NodeOutput(file_3d, model_3d_info, viewport_state['camera_info'], width, height)
|
||||
|
||||
|
||||
class Load3DExtension(ComfyExtension):
|
||||
|
||||
@@ -8,8 +8,6 @@ import comfy.ldm.common_dit
|
||||
import comfy.latent_formats
|
||||
import comfy.ldm.lumina.controlnet
|
||||
import comfy.ldm.supir.supir_modules
|
||||
import comfy.ldm.anima.lllite
|
||||
import comfy.ldm.wan.uni3c
|
||||
from comfy.ldm.wan.model_multitalk import WanMultiTalkAttentionBlock, MultiTalkAudioProjModel
|
||||
from comfy_api.latest import io
|
||||
from comfy.ldm.supir.supir_patch import SUPIRPatch
|
||||
@@ -238,12 +236,10 @@ class ModelPatchLoader:
|
||||
|
||||
def load_model_patch(self, name):
|
||||
model_patch_path = folder_paths.get_full_path_or_raise("model_patches", name)
|
||||
sd, metadata = comfy.utils.load_torch_file(model_patch_path, safe_load=True, return_metadata=True)
|
||||
sd = comfy.utils.load_torch_file(model_patch_path, safe_load=True)
|
||||
dtype = comfy.utils.weight_dtype(sd)
|
||||
|
||||
if 'lllite_conditioning1.conv1.weight' in sd:
|
||||
model = comfy.ldm.anima.lllite.AnimaLLLite(sd, metadata, device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast)
|
||||
elif 'controlnet_blocks.0.y_rms.weight' in sd:
|
||||
if 'controlnet_blocks.0.y_rms.weight' in sd:
|
||||
additional_in_dim = sd["img_in.weight"].shape[1] - 64
|
||||
model = QwenImageBlockWiseControlNet(additional_in_dim=additional_in_dim, device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast)
|
||||
elif 'feature_embedder.mid_layer_norm.bias' in sd:
|
||||
@@ -265,37 +261,6 @@ class ModelPatchLoader:
|
||||
if torch.count_nonzero(ref_weight) == 0:
|
||||
config['broken'] = True
|
||||
model = comfy.ldm.lumina.controlnet.ZImage_Control(device=comfy.model_management.unet_offload_device(), dtype=dtype, operations=comfy.ops.manual_cast, **config)
|
||||
elif 'controlnet_patch_embedding.weight' in sd: # Uni3C controlnet for Wan
|
||||
attn_key_replace = {".self_attn.to_q.": ".self_attn.q.",
|
||||
".self_attn.to_k.": ".self_attn.k.",
|
||||
".self_attn.to_v.": ".self_attn.v.",
|
||||
".self_attn.to_out.0.": ".self_attn.o."}
|
||||
converted_sd = {}
|
||||
for k, w in sd.items():
|
||||
for r, rr in attn_key_replace.items():
|
||||
k = k.replace(r, rr)
|
||||
converted_sd[k] = w
|
||||
sd = converted_sd
|
||||
|
||||
num_layers = sum(1 for k in sd if k.startswith("proj_out.") and k.endswith(".weight"))
|
||||
conv_out_dim = sd["controlnet_patch_embedding.weight"].shape[0]
|
||||
if "proj_in.weight" in sd:
|
||||
dim = sd["proj_in.weight"].shape[0]
|
||||
else:
|
||||
dim = conv_out_dim
|
||||
model = comfy.ldm.wan.uni3c.WanUni3CControlnet(
|
||||
in_channels=sd["controlnet_patch_embedding.weight"].shape[1],
|
||||
conv_out_dim=conv_out_dim,
|
||||
dim=dim,
|
||||
ffn_dim=sd["controlnet_blocks.0.ffn.0.bias"].shape[0],
|
||||
num_layers=num_layers,
|
||||
time_embed_dim=sd["controlnet_blocks.0.norm1.linear.weight"].shape[1],
|
||||
out_proj_dim=sd["proj_out.0.weight"].shape[0],
|
||||
add_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[1],
|
||||
mid_channels=sd["controlnet_mask_embedding.mask_proj.0.weight"].shape[0],
|
||||
device=comfy.model_management.unet_offload_device(),
|
||||
dtype=dtype,
|
||||
operations=comfy.ops.manual_cast)
|
||||
elif "audio_proj.proj1.weight" in sd:
|
||||
model = MultiTalkModelPatch(
|
||||
audio_window=5, context_tokens=32, vae_scale=4,
|
||||
@@ -331,50 +296,6 @@ class ModelPatchLoader:
|
||||
return (model_patcher,)
|
||||
|
||||
|
||||
class AnimaLLLiteApply:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"model": ("MODEL",),
|
||||
"model_patch": ("MODEL_PATCH",),
|
||||
"image": ("IMAGE",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
},
|
||||
"optional": {"mask": ("MASK",),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "apply_patch"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
CATEGORY = "model_patches/anima"
|
||||
|
||||
def apply_patch(self, model, model_patch, image, strength, start_percent, end_percent, mask=None):
|
||||
image = image[..., :3]
|
||||
|
||||
if model_patch.model.cond_in_channels == 4 and mask is None:
|
||||
mask = torch.zeros_like(image[..., 0])
|
||||
elif model_patch.model.cond_in_channels != 4:
|
||||
mask = None
|
||||
|
||||
model_sampling = model.get_model_object("model_sampling")
|
||||
sigma_start = float(model_sampling.percent_to_sigma(start_percent))
|
||||
sigma_end = float(model_sampling.percent_to_sigma(end_percent))
|
||||
patch = comfy.ldm.anima.lllite.AnimaLLLitePatch(model_patch, image, mask, strength, sigma_start, sigma_end)
|
||||
model_patched = model.clone()
|
||||
model_patched.set_model_post_input_patch(patch)
|
||||
model_patched.set_model_attn1_patch(comfy.ldm.anima.lllite.AnimaLLLiteAttentionPatch(
|
||||
patch,
|
||||
{"q": "self_attn_q_proj", "k": "self_attn_k_proj", "v": "self_attn_v_proj"},
|
||||
))
|
||||
model_patched.set_model_attn2_patch(comfy.ldm.anima.lllite.AnimaLLLiteAttentionPatch(
|
||||
patch,
|
||||
{"q": "cross_attn_q_proj"},
|
||||
))
|
||||
model_patched.set_model_patch(comfy.ldm.anima.lllite.AnimaLLLiteMLPPatch(patch), "mlp_patch")
|
||||
return (model_patched,)
|
||||
|
||||
|
||||
class DiffSynthCnetPatch:
|
||||
def __init__(self, model_patch, vae, image, strength, mask=None):
|
||||
self.model_patch = model_patch
|
||||
@@ -593,150 +514,6 @@ class ZImageFunControlnet(QwenImageDiffsynthControlnet):
|
||||
|
||||
CATEGORY = "model/patch/z-image"
|
||||
|
||||
class WanUni3CCnetPatch:
|
||||
def __init__(self, model_patch, render_video, vae, latent_format, strength, sigma_start, sigma_end):
|
||||
self.model_patch = model_patch
|
||||
self.render_video = render_video
|
||||
self.vae = vae
|
||||
self.latent_format = latent_format
|
||||
self.strength = strength
|
||||
self.sigma_start = sigma_start
|
||||
self.sigma_end = sigma_end
|
||||
self.prepared_render = None
|
||||
self.temp_data = None
|
||||
|
||||
def encode_render_video(self, target_latent_shape):
|
||||
t_len, h_len, w_len = target_latent_shape
|
||||
temporal_compression = self.vae.temporal_compression_decode() or 1
|
||||
spatial_compression = self.vae.spacial_compression_encode()
|
||||
target_frames = (t_len - 1) * temporal_compression + 1
|
||||
target_height = h_len * spatial_compression
|
||||
target_width = w_len * spatial_compression
|
||||
|
||||
frames = self.render_video
|
||||
if frames.shape[0] > target_frames:
|
||||
frames = frames[:target_frames]
|
||||
elif frames.shape[0] < target_frames:
|
||||
last_frame = frames[-1:].expand(target_frames - frames.shape[0], -1, -1, -1)
|
||||
frames = torch.cat([frames, last_frame], dim=0)
|
||||
|
||||
if frames.shape[1] != target_height or frames.shape[2] != target_width:
|
||||
frames = comfy.utils.common_upscale(frames.movedim(-1, 1), target_width, target_height, "bilinear", "center").movedim(1, -1)
|
||||
|
||||
loaded_models = comfy.model_management.loaded_models(only_currently_used=True)
|
||||
render_latent = self.vae.encode(frames)
|
||||
comfy.model_management.load_models_gpu(loaded_models)
|
||||
return self.latent_format.process_in(render_latent)
|
||||
|
||||
def build_controlnet_input(self, x, dtype, samples_per_cond):
|
||||
# first 20 channels of the model input: noise latent + I2V mask (zero padded for T2V)
|
||||
hidden = x[:samples_per_cond, :20].to(dtype)
|
||||
if hidden.shape[1] < 20:
|
||||
pad_shape = list(hidden.shape)
|
||||
pad_shape[1] = 20 - hidden.shape[1]
|
||||
hidden = torch.cat([hidden, torch.zeros(pad_shape, dtype=hidden.dtype, device=hidden.device)], dim=1)
|
||||
|
||||
render = self.prepared_render
|
||||
if render is None or render.shape[2:] != hidden.shape[2:]:
|
||||
render = self.encode_render_video(hidden.shape[2:])
|
||||
render = render.to(device=hidden.device, dtype=dtype)
|
||||
self.prepared_render = render
|
||||
if render.shape[0] != hidden.shape[0]:
|
||||
render = render.expand(hidden.shape[0], -1, -1, -1, -1)
|
||||
return torch.cat([hidden, render], dim=1)
|
||||
|
||||
def __call__(self, kwargs):
|
||||
img = kwargs.get("img")
|
||||
block_index = kwargs.get("block_index")
|
||||
transformer_options = kwargs.get("transformer_options", {})
|
||||
|
||||
if block_index == 0:
|
||||
self.temp_data = None
|
||||
active = True
|
||||
sigmas = transformer_options.get("sigmas", None)
|
||||
if sigmas is not None:
|
||||
sigma = sigmas[0].item()
|
||||
if sigma > self.sigma_start or sigma < self.sigma_end:
|
||||
active = False
|
||||
if active:
|
||||
x = kwargs.get("x")
|
||||
# cond and uncond chunks share latents, so we can reuse residuals
|
||||
num_conds = len(transformer_options.get("cond_or_uncond", [0]))
|
||||
samples_per_cond = x.shape[0]
|
||||
if num_conds > 0 and x.shape[0] % num_conds == 0:
|
||||
samples_per_cond = x.shape[0] // num_conds
|
||||
temb = kwargs.get("vec")[:samples_per_cond]
|
||||
if temb.ndim == 3:
|
||||
temb = temb[:, 0]
|
||||
model = self.model_patch.model
|
||||
controlnet_input = self.build_controlnet_input(x, img.dtype, samples_per_cond)
|
||||
hidden, freqs = model.process_input(controlnet_input)
|
||||
self.temp_data = (hidden, temb.to(img.dtype), freqs)
|
||||
|
||||
num_layers = self.model_patch.model.num_layers
|
||||
if self.temp_data is not None and block_index < num_layers:
|
||||
hidden, temb, freqs = self.temp_data
|
||||
hidden, residual = self.model_patch.model.forward_block(block_index, hidden, temb, freqs)
|
||||
residual = residual.to(img.dtype) * self.strength
|
||||
if residual.shape[0] != img.shape[0]:
|
||||
residual = residual.repeat(img.shape[0] // residual.shape[0], 1, 1)
|
||||
img_offset = kwargs.get("img_offset", 0)
|
||||
img[:, img_offset:img_offset + residual.shape[1]] += residual
|
||||
if block_index >= num_layers - 1:
|
||||
self.temp_data = None
|
||||
else:
|
||||
self.temp_data = (hidden, temb, freqs)
|
||||
|
||||
return kwargs
|
||||
|
||||
def to(self, device_or_dtype):
|
||||
if isinstance(device_or_dtype, torch.device):
|
||||
if self.prepared_render is not None:
|
||||
self.prepared_render = self.prepared_render.to(device_or_dtype)
|
||||
self.temp_data = None
|
||||
return self
|
||||
|
||||
def models(self):
|
||||
return [self.model_patch]
|
||||
|
||||
|
||||
class WanUni3CControlnetApply:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "model": ("MODEL",),
|
||||
"model_patch": ("MODEL_PATCH",),
|
||||
"vae": ("VAE",),
|
||||
"render_video": ("IMAGE", {"tooltip": "The guidance video rendered from the camera trajectory, most commonly warped point cloud renders of the input image."}),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
}}
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "apply_patch"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
CATEGORY = "model/patch/wan"
|
||||
|
||||
def apply_patch(self, model, model_patch, vae, render_video, strength, start_percent, end_percent):
|
||||
if not isinstance(model_patch.model, comfy.ldm.wan.uni3c.WanUni3CControlnet):
|
||||
raise ValueError("The connected model patch is not a Uni3C ControlNet.")
|
||||
cnet_dim = model_patch.model.controlnet_blocks[0].norm1.linear.in_features
|
||||
model_dim = getattr(model.get_model_object("diffusion_model"), "dim", None)
|
||||
if model_dim is None:
|
||||
raise ValueError("The Uni3C ControlNet only works with Wan models.")
|
||||
if model_dim != cnet_dim:
|
||||
raise ValueError("This Uni3C ControlNet expects a Wan model with dim {}, the loaded model has dim {}.".format(cnet_dim, model_dim))
|
||||
|
||||
model_patched = model.clone()
|
||||
model_sampling = model.get_model_object("model_sampling")
|
||||
sigma_start = model_sampling.percent_to_sigma(start_percent)
|
||||
sigma_end = model_sampling.percent_to_sigma(end_percent)
|
||||
latent_format = model.get_model_object("latent_format")
|
||||
patch = WanUni3CCnetPatch(model_patch, render_video[:, :, :, :3], vae, latent_format, strength, sigma_start, sigma_end)
|
||||
model_patched.set_model_double_block_patch(patch)
|
||||
return (model_patched,)
|
||||
|
||||
|
||||
class UsoStyleProjectorPatch:
|
||||
def __init__(self, model_patch, encoded_image):
|
||||
self.model_patch = model_patch
|
||||
@@ -895,18 +672,14 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ModelPatchLoader": ModelPatchLoader,
|
||||
"QwenImageDiffsynthControlnet": QwenImageDiffsynthControlnet,
|
||||
"ZImageFunControlnet": ZImageFunControlnet,
|
||||
"WanUni3CControlnetApply": WanUni3CControlnetApply,
|
||||
"USOStyleReference": USOStyleReference,
|
||||
"SUPIRApply": SUPIRApply,
|
||||
"AnimaLLLiteApply": AnimaLLLiteApply,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ModelPatchLoader": "Load Model Patch",
|
||||
"QwenImageDiffsynthControlnet": "Apply Qwen Image DiffSynth ControlNet",
|
||||
"ZImageFunControlnet": "Apply Z-Image Fun ControlNet",
|
||||
"WanUni3CControlnetApply": "Apply Wan Uni3C ControlNet",
|
||||
"USOStyleReference": "Apply USO Style Reference",
|
||||
"SUPIRApply": "Apply SUPIR Patch",
|
||||
"AnimaLLLiteApply": "Apply Anima LLLite",
|
||||
}
|
||||
|
||||
@@ -418,9 +418,8 @@ def _save_file3d_to_output(model_3d: Types.File3D, filename_prefix: str) -> str:
|
||||
|
||||
def execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs) -> IO.NodeOutput:
|
||||
model_file = _save_file3d_to_output(model_3d, filename_prefix)
|
||||
viewport_state = viewport_state if isinstance(viewport_state, dict) else {}
|
||||
camera_info_input = kwargs.get("camera_info", None)
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('camera_info')
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state['camera_info']
|
||||
model_3d_info_input = kwargs.get("model_3d_info", None)
|
||||
model_3d_info = model_3d_info_input if model_3d_info_input is not None else viewport_state.get('model_3d_info', [])
|
||||
return IO.NodeOutput(
|
||||
|
||||
@@ -920,11 +920,10 @@ def _run_training_loop(
|
||||
"""
|
||||
sigmas = torch.tensor(range(num_images))
|
||||
noise = comfy_extras.nodes_custom_sampler.Noise_RandomNoise(seed)
|
||||
ndim = latents[0].ndim
|
||||
|
||||
if bucket_mode:
|
||||
# Use first bucket's first latent as dummy for guider
|
||||
dummy_latent = latents[0][:1].repeat(num_images, *[1]*(ndim-1))
|
||||
dummy_latent = latents[0][:1].repeat(num_images, 1, 1, 1)
|
||||
guider.sample(
|
||||
noise.generate_noise({"samples": dummy_latent}),
|
||||
dummy_latent,
|
||||
@@ -934,7 +933,7 @@ def _run_training_loop(
|
||||
)
|
||||
elif multi_res:
|
||||
# use first latent as dummy latent if multi_res
|
||||
latents = latents[0].repeat(num_images, *[1]*(ndim-1))
|
||||
latents = latents[0].repeat(num_images, 1, 1, 1)
|
||||
guider.sample(
|
||||
noise.generate_noise({"samples": latents}),
|
||||
latents,
|
||||
|
||||
+1
-1
@@ -1,3 +1,3 @@
|
||||
# This file is automatically generated by the build process when version is
|
||||
# updated in pyproject.toml.
|
||||
__version__ = "0.28.0"
|
||||
__version__ = "0.27.0"
|
||||
|
||||
+2
-2
@@ -426,12 +426,12 @@ def _is_intermediate_output(dynprompt, node_id):
|
||||
|
||||
|
||||
def _send_cached_ui(server, node_id, display_node_id, cached, prompt_id, ui_outputs):
|
||||
if cached.ui is not None:
|
||||
ui_outputs[node_id] = cached.ui
|
||||
if server.client_id is None:
|
||||
return
|
||||
cached_ui = cached.ui or {}
|
||||
server.send_sync("executed", { "node": node_id, "display_node": display_node_id, "output": cached_ui.get("output", None), "prompt_id": prompt_id }, server.client_id)
|
||||
if cached.ui is not None:
|
||||
ui_outputs[node_id] = cached.ui
|
||||
|
||||
async def execute(server, dynprompt, caches, current_item, extra_data, executed, prompt_id, execution_list, pending_subgraph_results, pending_async_nodes, ui_outputs):
|
||||
unique_id = current_item
|
||||
|
||||
@@ -992,7 +992,7 @@ class CLIPLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "clip_name": (folder_paths.get_filename_list("text_encoders"), ),
|
||||
"type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2", "joyimage"], ),
|
||||
"type": (["stable_diffusion", "stable_cascade", "sd3", "stable_audio", "mochi", "ltxv", "pixart", "cosmos", "lumina2", "wan", "hidream", "chroma", "ace", "omnigen2", "qwen_image", "hunyuan_image", "flux2", "ovis", "longcat_image", "cogvideox", "lens", "pixeldit", "ideogram4", "boogu", "krea2"], ),
|
||||
},
|
||||
"optional": {
|
||||
"device": (["default", "cpu"], {"advanced": True}),
|
||||
@@ -1002,7 +1002,7 @@ class CLIPLoader:
|
||||
|
||||
CATEGORY = "model/loaders"
|
||||
|
||||
DESCRIPTION = "Recipes:\nsd: clip-l\nstable cascade: clip-g\nsd3: t5 xxl / clip-g / clip-l\nstable audio: t5 base\nmochi: t5 xxl\ncogvideox: t5 xxl (226-token padding)\ncosmos: old t5 xxl\nlumina2: gemma 2 2B\nwan: umt5 xxl\nhidream: llama-3.1 (Recommend) or t5\nomnigen2: qwen vl 2.5 3B\njoyimage: qwen3-vl 8B\nlens: gpt-oss-20b\npixeldit: gemma 2 2B elm"
|
||||
DESCRIPTION = "Recipes:\nsd: clip-l\nstable cascade: clip-g\nsd3: t5 xxl / clip-g / clip-l\nstable audio: t5 base\nmochi: t5 xxl\ncogvideox: t5 xxl (226-token padding)\ncosmos: old t5 xxl\nlumina2: gemma 2 2B\nwan: umt5 xxl\nhidream: llama-3.1 (Recommend) or t5\nomnigen2: qwen vl 2.5 3B\nlens: gpt-oss-20b\npixeldit: gemma 2 2B elm"
|
||||
|
||||
def load_clip(self, clip_name, type="stable_diffusion", device="default"):
|
||||
clip_type = getattr(comfy.sd.CLIPType, type.upper(), comfy.sd.CLIPType.STABLE_DIFFUSION)
|
||||
@@ -2462,7 +2462,6 @@ async def init_builtin_extra_nodes():
|
||||
"nodes_seedvr.py",
|
||||
"nodes_context_windows.py",
|
||||
"nodes_qwen.py",
|
||||
"nodes_joyimage.py",
|
||||
"nodes_boogu.py",
|
||||
"nodes_chroma_radiance.py",
|
||||
"nodes_pid.py",
|
||||
|
||||
+499
-74
@@ -7,18 +7,18 @@ components:
|
||||
description: Timestamp when the asset was created
|
||||
format: date-time
|
||||
type: string
|
||||
display_name:
|
||||
description: Display name of the asset. Mirrors name for backwards compatibility.
|
||||
nullable: true
|
||||
type: string
|
||||
file_path:
|
||||
description: Relative path in global-namespace-root form (e.g. "models/checkpoints/flux.safetensors")
|
||||
nullable: true
|
||||
type: string
|
||||
hash:
|
||||
description: Blake3 hash of the asset content.
|
||||
pattern: ^blake3:[a-f0-9]{64}$
|
||||
type: string
|
||||
loader_path:
|
||||
description: The value a loader consumes to load this asset. Null when no loader can resolve the file.
|
||||
nullable: true
|
||||
type: string
|
||||
display_name:
|
||||
description: Human-facing label for the asset. Not unique.
|
||||
nullable: true
|
||||
type: string
|
||||
id:
|
||||
description: Unique identifier for the asset
|
||||
format: uuid
|
||||
@@ -144,14 +144,6 @@ components:
|
||||
AssetUpdated:
|
||||
description: Response returned when an existing asset is successfully updated.
|
||||
properties:
|
||||
display_name:
|
||||
description: Display name of the asset. Mirrors name for backwards compatibility.
|
||||
nullable: true
|
||||
type: string
|
||||
file_path:
|
||||
description: Relative path in global-namespace-root form (e.g. "models/checkpoints/flux.safetensors")
|
||||
nullable: true
|
||||
type: string
|
||||
hash:
|
||||
description: Blake3 hash of the asset content.
|
||||
pattern: ^blake3:[a-f0-9]{64}$
|
||||
@@ -230,6 +222,89 @@ components:
|
||||
- base_version
|
||||
- workflow_json
|
||||
type: object
|
||||
DownloadEnqueueRequest:
|
||||
description: Request body for enqueuing a server-side model download.
|
||||
properties:
|
||||
allow_any_extension:
|
||||
default: false
|
||||
description: Permit a non-model file extension (default only allows known model extensions).
|
||||
type: boolean
|
||||
expected_sha256:
|
||||
description: Optional hub-provided SHA256 to verify the completed file against (fail-closed).
|
||||
nullable: true
|
||||
type: string
|
||||
model_id:
|
||||
description: Destination as "<directory>/<filename>", resolving to a registered model folder (e.g. "loras/my_lora.safetensors").
|
||||
type: string
|
||||
priority:
|
||||
default: 0
|
||||
description: Scheduling priority; higher is admitted first.
|
||||
type: integer
|
||||
url:
|
||||
description: Source URL; must be on the allowlist (host + scheme + extension).
|
||||
type: string
|
||||
required:
|
||||
- url
|
||||
- model_id
|
||||
type: object
|
||||
DownloadStatus:
|
||||
description: Current state and live progress of a single download.
|
||||
properties:
|
||||
bytes_done:
|
||||
type: integer
|
||||
created_at:
|
||||
type: integer
|
||||
download_id:
|
||||
format: uuid
|
||||
type: string
|
||||
error:
|
||||
nullable: true
|
||||
type: string
|
||||
eta_seconds:
|
||||
nullable: true
|
||||
type: number
|
||||
model_id:
|
||||
type: string
|
||||
priority:
|
||||
type: integer
|
||||
progress:
|
||||
description: Fraction in [0,1]; null until total size is known.
|
||||
nullable: true
|
||||
type: number
|
||||
segments:
|
||||
description: Per-segment progress (segmented downloads only).
|
||||
items:
|
||||
properties:
|
||||
bytes_done:
|
||||
type: integer
|
||||
idx:
|
||||
type: integer
|
||||
length:
|
||||
type: integer
|
||||
type: object
|
||||
nullable: true
|
||||
type: array
|
||||
speed_bps:
|
||||
nullable: true
|
||||
type: number
|
||||
status:
|
||||
enum:
|
||||
- queued
|
||||
- active
|
||||
- paused
|
||||
- verifying
|
||||
- completed
|
||||
- failed
|
||||
- cancelled
|
||||
type: string
|
||||
total_bytes:
|
||||
nullable: true
|
||||
type: integer
|
||||
updated_at:
|
||||
type: integer
|
||||
url:
|
||||
type: string
|
||||
type: object
|
||||
ErrorResponse:
|
||||
description: Standard error response with a machine-readable code and human-readable message.
|
||||
properties:
|
||||
@@ -511,6 +586,27 @@ components:
|
||||
required:
|
||||
- history
|
||||
type: object
|
||||
DownloadAuthProvider:
|
||||
description: Per-provider download-auth status. Never includes a token.
|
||||
properties:
|
||||
env_key_present:
|
||||
description: Whether an API key for this provider is set via an environment variable.
|
||||
type: boolean
|
||||
logged_in:
|
||||
description: Whether a stored OAuth token exists for this provider.
|
||||
type: boolean
|
||||
login_in_progress:
|
||||
description: Whether an OAuth login flow is currently awaiting the browser callback.
|
||||
type: boolean
|
||||
provider:
|
||||
description: Provider name (e.g. "huggingface", "civitai").
|
||||
type: string
|
||||
required:
|
||||
- provider
|
||||
- logged_in
|
||||
- login_in_progress
|
||||
- env_key_present
|
||||
type: object
|
||||
JobCancelResponse:
|
||||
description: Response for POST /api/jobs/{job_id}/cancel. Returned on both fresh cancels and idempotent no-ops.
|
||||
properties:
|
||||
@@ -530,10 +626,6 @@ components:
|
||||
description: Job creation timestamp (Unix timestamp in milliseconds)
|
||||
format: int64
|
||||
type: integer
|
||||
execution_end_time:
|
||||
description: Workflow execution completion timestamp (Unix milliseconds, only present for terminal states)
|
||||
format: int64
|
||||
type: integer
|
||||
execution_error:
|
||||
allOf:
|
||||
- $ref: '#/components/schemas/ExecutionError'
|
||||
@@ -542,10 +634,6 @@ components:
|
||||
additionalProperties: true
|
||||
description: Node-level execution metadata (only for terminal states)
|
||||
type: object
|
||||
execution_start_time:
|
||||
description: Workflow execution start timestamp (Unix milliseconds, only present once execution has started)
|
||||
format: int64
|
||||
type: integer
|
||||
execution_status:
|
||||
additionalProperties: true
|
||||
description: ComfyUI execution status and timeline (only for terminal states)
|
||||
@@ -578,12 +666,6 @@ components:
|
||||
description: Last update timestamp (Unix timestamp in milliseconds)
|
||||
format: int64
|
||||
type: integer
|
||||
user_id:
|
||||
description: |
|
||||
ID of the user that owns this job (see the `workspace_id`
|
||||
description above for why this is always the caller's own id
|
||||
on a successful response).
|
||||
type: string
|
||||
workflow:
|
||||
additionalProperties: true
|
||||
description: |
|
||||
@@ -597,18 +679,6 @@ components:
|
||||
workflow_id:
|
||||
description: UUID identifying the workflow graph definition
|
||||
type: string
|
||||
workspace_id:
|
||||
description: |
|
||||
ID of the workspace that owns this job. A successful (200)
|
||||
response from this operation is only ever returned for the
|
||||
caller's own job (see this operation's ownership-scoped
|
||||
query), so this is always the caller's own workspace —
|
||||
consumers that also need to correlate this job to its
|
||||
live-progress broadcast channel (workspace+user scoped; see
|
||||
the internal common/gateways/broadcast package) can use this
|
||||
value directly rather than resolving their own identity a
|
||||
second way.
|
||||
type: string
|
||||
required:
|
||||
- id
|
||||
- status
|
||||
@@ -809,6 +879,14 @@ components:
|
||||
ModelFolder:
|
||||
description: Represents a folder containing models
|
||||
properties:
|
||||
extensions:
|
||||
description: The folder's registered file-extension allowlist. An empty array means the folder accepts any extension (match-all).
|
||||
example:
|
||||
- .ckpt
|
||||
- .safetensors
|
||||
items:
|
||||
type: string
|
||||
type: array
|
||||
folders:
|
||||
description: List of paths where models of this type are stored
|
||||
example:
|
||||
@@ -1591,13 +1669,7 @@ paths:
|
||||
schema:
|
||||
default: true
|
||||
type: boolean
|
||||
- description: |
|
||||
Filter assets by content hash, in the canonical `blake3:<hex>`
|
||||
form. Matches regardless of which of this asset store's two
|
||||
internal hash storage formats the matching row was written
|
||||
under (the canonical form used by from-hash-created references,
|
||||
or the raw `<hex>.<ext>`/bare `<hex>` storage key used by direct
|
||||
uploads) — both represent the same content hash.
|
||||
- description: Filter assets by exact content hash.
|
||||
in: query
|
||||
name: hash
|
||||
schema:
|
||||
@@ -1676,7 +1748,7 @@ paths:
|
||||
format: uuid
|
||||
type: string
|
||||
tags:
|
||||
description: JSON-encoded array of freeform tag strings, e.g. '["models","checkpoint"]'. Common types include "models", "input", "output", and "temp", but any tag can be used in any order.
|
||||
description: JSON-encoded array of tag strings. For new byte uploads, include exactly one destination role (`input`, `output`, or `models`); `models` uploads also require exactly one `model_type:<folder_name>` tag. Extra tags are stored as labels and do not create path components.
|
||||
type: string
|
||||
user_metadata:
|
||||
description: Custom JSON metadata as a string
|
||||
@@ -1861,7 +1933,7 @@ paths:
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/AssetUpdated'
|
||||
$ref: '#/components/schemas/Asset'
|
||||
description: Asset updated successfully
|
||||
"400":
|
||||
content:
|
||||
@@ -2382,6 +2454,377 @@ paths:
|
||||
summary: Get tag histogram for filtered assets
|
||||
tags:
|
||||
- file
|
||||
/api/download:
|
||||
get:
|
||||
description: List all known downloads (queued, active, paused, and terminal) with live progress.
|
||||
operationId: listDownloads
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
downloads:
|
||||
items:
|
||||
$ref: '#/components/schemas/DownloadStatus'
|
||||
type: array
|
||||
type: object
|
||||
description: List of downloads
|
||||
summary: List downloads
|
||||
tags:
|
||||
- download
|
||||
/api/download/availability:
|
||||
post:
|
||||
description: |
|
||||
Bulk per-id availability for a set of model_ids declared in a workflow.
|
||||
Returns whether each model is available on disk, currently downloading
|
||||
(with progress), or missing, plus whether its URL is on the allowlist.
|
||||
operationId: getModelsAvailability
|
||||
requestBody:
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
models:
|
||||
additionalProperties:
|
||||
type: string
|
||||
description: Map of "<directory>/<filename>" model_id to its declared source URL.
|
||||
type: object
|
||||
type: object
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
models:
|
||||
additionalProperties: true
|
||||
type: object
|
||||
type: object
|
||||
description: Per-id availability map
|
||||
summary: Bulk model availability + status
|
||||
tags:
|
||||
- download
|
||||
/api/download/clear:
|
||||
post:
|
||||
description: |
|
||||
Delete all terminal downloads (completed, failed, cancelled) from history
|
||||
in one transaction, so the cleared history persists across reloads. Live
|
||||
downloads (queued, active, paused, verifying) are skipped. Finished model
|
||||
files on disk are never removed; only leftover .part temp files are cleaned up.
|
||||
operationId: clearDownloads
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
deleted:
|
||||
description: Number of history rows removed.
|
||||
type: integer
|
||||
type: object
|
||||
description: History cleared
|
||||
summary: Clear terminal downloads from history
|
||||
tags:
|
||||
- download
|
||||
/api/download/auth:
|
||||
get:
|
||||
description: Per-provider download-auth status (OAuth login state + env-key presence). Never returns a token.
|
||||
operationId: getDownloadAuth
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
providers:
|
||||
items:
|
||||
$ref: '#/components/schemas/DownloadAuthProvider'
|
||||
type: array
|
||||
type: object
|
||||
description: Per-provider auth status
|
||||
summary: Get download auth status
|
||||
tags:
|
||||
- download
|
||||
/api/download/auth/{provider}/login:
|
||||
post:
|
||||
description: |
|
||||
Start an OAuth 2.0 PKCE login for a provider. Spawns a transient loopback
|
||||
callback server and returns the authorize URL to open in a browser.
|
||||
operationId: loginDownloadAuth
|
||||
parameters:
|
||||
- in: path
|
||||
name: provider
|
||||
required: true
|
||||
schema:
|
||||
type: string
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
authorize_url:
|
||||
type: string
|
||||
type: object
|
||||
description: Authorize URL to open in a browser
|
||||
"400":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Unknown provider or OAuth app not configured
|
||||
"409":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: A login for this provider is already in progress
|
||||
summary: Start OAuth login for a provider
|
||||
tags:
|
||||
- download
|
||||
/api/download/auth/{provider}/logout:
|
||||
post:
|
||||
description: Clear the stored OAuth token (memory + on-disk file) for a provider.
|
||||
operationId: logoutDownloadAuth
|
||||
parameters:
|
||||
- in: path
|
||||
name: provider
|
||||
required: true
|
||||
schema:
|
||||
type: string
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
logged_out:
|
||||
type: boolean
|
||||
type: object
|
||||
description: Token cleared
|
||||
"400":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Unknown provider
|
||||
summary: Log out a provider
|
||||
tags:
|
||||
- download
|
||||
/api/download/enqueue:
|
||||
post:
|
||||
description: |
|
||||
Enqueue a server-side model download. The URL must be on the allowlist
|
||||
(host + scheme + extension) and the model_id must be "<directory>/<filename>"
|
||||
resolving to a registered model folder. Returns immediately; track progress
|
||||
via GET /api/download/{id} or the "download_progress" websocket event.
|
||||
operationId: enqueueDownload
|
||||
requestBody:
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/DownloadEnqueueRequest'
|
||||
responses:
|
||||
"202":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
accepted:
|
||||
type: boolean
|
||||
download_id:
|
||||
format: uuid
|
||||
type: string
|
||||
type: object
|
||||
description: Download accepted and queued
|
||||
"400":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Invalid request (bad URL, model_id, or not allowlisted)
|
||||
"409":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Already on disk or already downloading
|
||||
summary: Enqueue a model download
|
||||
tags:
|
||||
- download
|
||||
/api/download/{id}:
|
||||
delete:
|
||||
description: |
|
||||
Delete a single terminal download from history so it stays gone across
|
||||
reloads. Refuses (409) to delete a live download (queued, active, paused,
|
||||
verifying) — cancel it first. The finished model file on disk is never
|
||||
removed; only a leftover .part temp file is cleaned up.
|
||||
operationId: deleteDownload
|
||||
parameters:
|
||||
- in: path
|
||||
name: id
|
||||
required: true
|
||||
schema:
|
||||
format: uuid
|
||||
type: string
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
deleted:
|
||||
type: boolean
|
||||
type: object
|
||||
description: Deleted
|
||||
"404":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: No such download
|
||||
"409":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Download is still in progress
|
||||
summary: Delete a download from history
|
||||
tags:
|
||||
- download
|
||||
get:
|
||||
description: Get the current status + progress of a single download.
|
||||
operationId: getDownloadStatus
|
||||
parameters:
|
||||
- in: path
|
||||
name: id
|
||||
required: true
|
||||
schema:
|
||||
format: uuid
|
||||
type: string
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/DownloadStatus'
|
||||
description: Download status
|
||||
"404":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: No such download
|
||||
summary: Get download status
|
||||
tags:
|
||||
- download
|
||||
/api/download/{id}/cancel:
|
||||
post:
|
||||
description: Cancel a download. The partial file is removed.
|
||||
operationId: cancelDownload
|
||||
parameters:
|
||||
- in: path
|
||||
name: id
|
||||
required: true
|
||||
schema:
|
||||
format: uuid
|
||||
type: string
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
ok:
|
||||
type: boolean
|
||||
type: object
|
||||
description: Cancelled
|
||||
summary: Cancel a download
|
||||
tags:
|
||||
- download
|
||||
/api/download/{id}/pause:
|
||||
post:
|
||||
description: Pause a download. The partial file and per-segment offsets are retained for resume.
|
||||
operationId: pauseDownload
|
||||
parameters:
|
||||
- in: path
|
||||
name: id
|
||||
required: true
|
||||
schema:
|
||||
format: uuid
|
||||
type: string
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
ok:
|
||||
type: boolean
|
||||
type: object
|
||||
description: Paused
|
||||
summary: Pause a download
|
||||
tags:
|
||||
- download
|
||||
/api/download/{id}/priority:
|
||||
post:
|
||||
description: Set a download's scheduling priority. Higher priority is admitted first when a slot frees.
|
||||
operationId: setDownloadPriority
|
||||
parameters:
|
||||
- in: path
|
||||
name: id
|
||||
required: true
|
||||
schema:
|
||||
format: uuid
|
||||
type: string
|
||||
requestBody:
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
priority:
|
||||
type: integer
|
||||
required:
|
||||
- priority
|
||||
type: object
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
ok:
|
||||
type: boolean
|
||||
type: object
|
||||
description: Priority updated
|
||||
summary: Set download priority
|
||||
tags:
|
||||
- download
|
||||
/api/download/{id}/resume:
|
||||
post:
|
||||
description: Resume a paused (or failed) download from its persisted offsets.
|
||||
operationId: resumeDownload
|
||||
parameters:
|
||||
- in: path
|
||||
name: id
|
||||
required: true
|
||||
schema:
|
||||
format: uuid
|
||||
type: string
|
||||
responses:
|
||||
"200":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
properties:
|
||||
ok:
|
||||
type: boolean
|
||||
type: object
|
||||
description: Resumed
|
||||
summary: Resume a download
|
||||
tags:
|
||||
- download
|
||||
/api/embeddings:
|
||||
get:
|
||||
description: Returns the list of text-encoder embeddings available on disk.
|
||||
@@ -2496,29 +2939,15 @@ paths:
|
||||
schema:
|
||||
additionalProperties: true
|
||||
properties:
|
||||
free_tier_balance:
|
||||
description: Free-tier job allowance for an authenticated non-paid (FREE-tier) user in the rollout. Absent for paid users and unauthenticated requests. Synthesized from config before a grant row exists so a brand-new user still sees their full allowance.
|
||||
properties:
|
||||
allowance:
|
||||
description: Total free jobs granted for the current period
|
||||
type: integer
|
||||
remaining:
|
||||
description: Free jobs remaining (allowance - used, floored at 0)
|
||||
type: integer
|
||||
used:
|
||||
description: Free jobs consumed so far
|
||||
type: integer
|
||||
required:
|
||||
- allowance
|
||||
- used
|
||||
- remaining
|
||||
type: object
|
||||
max_upload_size:
|
||||
description: Maximum upload size in bytes
|
||||
type: integer
|
||||
supports_preview_metadata:
|
||||
description: Whether the server supports preview metadata
|
||||
type: boolean
|
||||
supports_model_type_tags:
|
||||
description: Whether the server supports namespaced model type asset tags
|
||||
type: boolean
|
||||
type: object
|
||||
description: Success
|
||||
headers:
|
||||
@@ -3346,12 +3775,6 @@ paths:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Invalid request parameters
|
||||
"401":
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/ErrorResponse'
|
||||
description: Unauthorized - Authentication required
|
||||
"500":
|
||||
content:
|
||||
application/json:
|
||||
@@ -5158,3 +5581,5 @@ tags:
|
||||
name: queue
|
||||
- description: Job lifecycle queries
|
||||
name: job
|
||||
- description: Model download management
|
||||
name: download
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "ComfyUI"
|
||||
version = "0.28.0"
|
||||
version = "0.27.0"
|
||||
readme = "README.md"
|
||||
license = { file = "LICENSE" }
|
||||
requires-python = ">=3.10"
|
||||
|
||||
+4
-4
@@ -1,6 +1,6 @@
|
||||
comfyui-frontend-package==1.45.21
|
||||
comfyui-workflow-templates==0.11.12
|
||||
comfyui-embedded-docs==0.5.8
|
||||
comfyui-frontend-package==1.45.20
|
||||
comfyui-workflow-templates==0.11.6
|
||||
comfyui-embedded-docs==0.5.7
|
||||
torch
|
||||
torchsde
|
||||
torchvision
|
||||
@@ -22,7 +22,7 @@ alembic
|
||||
SQLAlchemy>=2.0.0
|
||||
filelock
|
||||
av>=16.0.0
|
||||
comfy-kitchen==0.2.22
|
||||
comfy-kitchen==0.2.18
|
||||
comfy-aimdo==0.4.10
|
||||
requests
|
||||
simpleeval>=1.0.0
|
||||
|
||||
@@ -46,6 +46,8 @@ from app.frontend_management import FrontendManager, parse_version
|
||||
from comfy_api.internal import _ComfyNodeInternal
|
||||
from app.assets.seeder import asset_seeder
|
||||
from app.assets.api.routes import register_assets_routes
|
||||
from app.model_downloader.api.routes import register_routes as register_model_downloader_routes
|
||||
from app.model_downloader.manager import DOWNLOAD_MANAGER
|
||||
from app.assets.services.ingest import register_file_in_place
|
||||
from app.assets.services.path_utils import get_known_subfolder_tags
|
||||
from app.assets.services.asset_management import resolve_hash_to_path
|
||||
@@ -259,6 +261,7 @@ class PromptServer():
|
||||
else:
|
||||
register_assets_routes(self.app)
|
||||
asset_seeder.disable()
|
||||
register_model_downloader_routes(self.app)
|
||||
routes = web.RouteTableDef()
|
||||
self.routes = routes
|
||||
self.last_node_id = None
|
||||
@@ -1206,6 +1209,29 @@ class PromptServer():
|
||||
async def setup(self):
|
||||
timeout = aiohttp.ClientTimeout(total=None) # no timeout
|
||||
self.client_session = aiohttp.ClientSession(timeout=timeout)
|
||||
await self._setup_model_downloader()
|
||||
|
||||
async def _setup_model_downloader(self):
|
||||
"""Start the download manager: push progress over the websocket and
|
||||
resume any downloads interrupted by a previous run."""
|
||||
def _notify(download_id: str) -> None:
|
||||
try:
|
||||
view = DOWNLOAD_MANAGER.status_sync(download_id)
|
||||
if view is not None:
|
||||
# Drop the url field before broadcasting: the redacted URL
|
||||
# (scheme + host + path) should not leak to every connected
|
||||
# websocket client. download_id / model_id are sufficient to
|
||||
# correlate progress on the frontend.
|
||||
broadcast = {k: v for k, v in view.items() if k != "url"}
|
||||
self.send_sync("download_progress", broadcast)
|
||||
except Exception:
|
||||
logging.debug("download progress notify failed", exc_info=True)
|
||||
|
||||
DOWNLOAD_MANAGER.set_notify(_notify)
|
||||
try:
|
||||
await DOWNLOAD_MANAGER.start()
|
||||
except Exception as e:
|
||||
logging.warning("Failed to start model download manager: %s", e)
|
||||
|
||||
def add_routes(self):
|
||||
self.user_manager.add_routes(self.routes)
|
||||
|
||||
@@ -2,12 +2,11 @@ import pytest
|
||||
import torch
|
||||
import tempfile
|
||||
import os
|
||||
import sys
|
||||
import av
|
||||
import io
|
||||
from fractions import Fraction
|
||||
from comfy_api.input_impl.video_types import VideoFromFile, VideoFromComponents
|
||||
from comfy_api.util.video_types import VideoComponents, VideoContainer, VideoCodec
|
||||
from comfy_api.util.video_types import VideoComponents
|
||||
from comfy_api.input.basic_types import AudioInput
|
||||
from av.error import InvalidDataError
|
||||
|
||||
@@ -238,526 +237,3 @@ def test_duration_consistency(video_components):
|
||||
manual_duration = float(components.images.shape[0] / components.frame_rate)
|
||||
|
||||
assert duration == pytest.approx(manual_duration)
|
||||
|
||||
|
||||
def create_transcode_source(
|
||||
width=64, height=64, frames=30, fps=30, audio_streams=1, undecodable_audio=0, rotation=False,
|
||||
container_format="mov", audio_codec="pcm_s16le",
|
||||
):
|
||||
"""Create a temp video that save_to must transcode (mpeg4 video, so codec != h264).
|
||||
|
||||
``undecodable_audio`` trailing PCM streams get their fourcc corrupted so no decoder exists
|
||||
(``codec_context is None``), like the APAC track in iPhone spatial-audio recordings.
|
||||
``rotation`` patches a 90-degree display matrix into the video track header.
|
||||
"""
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format=container_format) as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=fps)
|
||||
video_stream.width = width
|
||||
video_stream.height = height
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio = []
|
||||
for _ in range(audio_streams + undecodable_audio):
|
||||
stream = container.add_stream(audio_codec, rate=44100)
|
||||
stream.sample_rate = 44100
|
||||
audio.append(stream)
|
||||
|
||||
for i in range(frames):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((height, width, 3), (i * 7) % 256, dtype=torch.uint8).numpy(),
|
||||
format="rgb24",
|
||||
)
|
||||
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
|
||||
# write audio in 1024-sample frames, like real decoders produce, so the
|
||||
# per-frame skip/cap logic in the transcode path actually runs
|
||||
for stream in audio:
|
||||
for offset in range(0, 44100 * frames // fps, 1024):
|
||||
n = min(1024, 44100 * frames // fps - offset)
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(1, n, dtype=torch.int16).numpy(), format="s16", layout="mono"
|
||||
)
|
||||
audio_frame.sample_rate = 44100
|
||||
audio_frame.pts = offset
|
||||
container.mux(stream.encode(audio_frame))
|
||||
for stream in [video_stream, *audio]:
|
||||
container.mux(stream.encode(None))
|
||||
|
||||
data = bytearray(buffer.getvalue())
|
||||
end = len(data)
|
||||
for _ in range(undecodable_audio):
|
||||
end = data.rindex(b"sowt", 0, end)
|
||||
data[end:end + 4] = b"Xpac"
|
||||
if rotation:
|
||||
# the 3x3 display matrix sits 40 bytes into the version-0 tkhd payload; first tkhd
|
||||
# inside moov = video track (search from moov so mdat bytes can't false-match)
|
||||
matrix_offset = data.index(b"tkhd", data.rindex(b"moov")) + 4 + 40
|
||||
values = [0, 1 << 16, 0, -(1 << 16), 0, 0, 0, 0, 1 << 30]
|
||||
data[matrix_offset:matrix_offset + 36] = b"".join(v.to_bytes(4, "big", signed=True) for v in values)
|
||||
|
||||
tmp = tempfile.NamedTemporaryFile(suffix=f".{container_format}", delete=False)
|
||||
tmp.write(bytes(data))
|
||||
tmp.close()
|
||||
return tmp.name
|
||||
|
||||
|
||||
def transcode_and_probe(video):
|
||||
buffer = io.BytesIO()
|
||||
video.save_to(buffer, format=VideoContainer.MP4, codec=VideoCodec.H264)
|
||||
buffer.seek(0)
|
||||
with av.open(buffer) as container:
|
||||
video_stream = container.streams.video[0]
|
||||
audio_stream = container.streams.audio[0] if container.streams.audio else None
|
||||
frames = 0
|
||||
first_pts = None
|
||||
for packet in container.demux(video_stream):
|
||||
for frame in packet.decode():
|
||||
if first_pts is None:
|
||||
first_pts = frame.pts
|
||||
frames += 1
|
||||
return {
|
||||
"codec": video_stream.codec_context.name,
|
||||
"width": video_stream.codec_context.width,
|
||||
"height": video_stream.codec_context.height,
|
||||
"frames": frames,
|
||||
"first_pts": first_pts,
|
||||
"video_seconds": float(video_stream.duration * video_stream.time_base) if video_stream.duration else None,
|
||||
"audio_seconds": float(audio_stream.duration * audio_stream.time_base)
|
||||
if audio_stream and audio_stream.duration else None,
|
||||
"audio_codecs": [s.codec_context.name for s in container.streams.audio],
|
||||
}
|
||||
|
||||
|
||||
def test_save_to_transcode_streams_without_buffering_frames():
|
||||
"""Transcoding must not decode the whole video into memory first (~2 GiB for this source)"""
|
||||
resource = pytest.importorskip("resource") # no getrusage on Windows
|
||||
rss_scale = 1 if sys.platform == "darwin" else 1024 # ru_maxrss: bytes on macOS, KiB elsewhere
|
||||
# ru_maxrss is a lifetime peak: a heavier test running earlier would shrink the measured
|
||||
# delta and quietly defang this canary, so keep this source the biggest thing in the suite
|
||||
file_path = create_transcode_source(width=640, height=480, frames=300)
|
||||
try:
|
||||
rss_before = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * rss_scale
|
||||
result = transcode_and_probe(VideoFromFile(file_path))
|
||||
rss_delta = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * rss_scale - rss_before
|
||||
|
||||
assert result["codec"] == "h264"
|
||||
assert result["frames"] == 300
|
||||
assert rss_delta < 500 * 2**20, f"transcode buffered frames in RAM (peak grew {rss_delta / 2**20:.0f} MiB)"
|
||||
finally:
|
||||
os.unlink(file_path)
|
||||
|
||||
|
||||
def test_save_to_transcode_honors_trim_window():
|
||||
"""start_time/duration trim applies to both video and audio on the streaming path"""
|
||||
file_path = create_transcode_source(frames=90) # 3s @ 30fps
|
||||
try:
|
||||
result = transcode_and_probe(VideoFromFile(file_path, start_time=1, duration=1))
|
||||
assert result["frames"] == pytest.approx(30, abs=2)
|
||||
assert result["first_pts"] == 0 # trimmed output is rebased to start at zero
|
||||
assert result["video_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
finally:
|
||||
os.unlink(file_path)
|
||||
|
||||
|
||||
def test_save_to_transcode_keeps_audio_of_sparse_video():
|
||||
"""Audio that runs ahead of a sparse video track (slideshows, timelapses) must be
|
||||
kept in full — it is only clamped to the video's end, never to the video cursor."""
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=30)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio_stream = container.add_stream("aac", rate=48000, layout="stereo")
|
||||
for t in (0, 30, 60): # 3 frames spread over 60 seconds
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), t * 4, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
).reformat(format="yuv420p")
|
||||
frame.pts = t * 15360
|
||||
frame.time_base = Fraction(1, 15360)
|
||||
container.mux(video_stream.encode(frame))
|
||||
container.mux(video_stream.encode(None))
|
||||
for offset in range(0, 48000 * 60, 1024):
|
||||
n = min(1024, 48000 * 60 - offset)
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(2, n, dtype=torch.float32).numpy(), format="fltp", layout="stereo"
|
||||
)
|
||||
audio_frame.sample_rate = 48000
|
||||
audio_frame.pts = offset
|
||||
audio_frame.time_base = Fraction(1, 48000)
|
||||
container.mux(audio_stream.encode(audio_frame))
|
||||
container.mux(audio_stream.encode(None))
|
||||
|
||||
buffer.seek(0)
|
||||
result = transcode_and_probe(VideoFromFile(buffer))
|
||||
assert result["audio_seconds"] == pytest.approx(60.0, abs=1.0)
|
||||
|
||||
|
||||
def test_save_to_transcode_vfr_audio_covers_video_span():
|
||||
"""A trim window in the sparse region of a VFR file keeps audio for the true pts span
|
||||
of the kept frames. Deriving the span as frames/average_rate undercuts it badly: the
|
||||
average is dominated by the dense region (and can be plain wrong on MediaRecorder files)."""
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=30)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio_stream = container.add_stream("aac", rate=48000, layout="stereo")
|
||||
# 10 frames inside the first second, then one every 1.25 s
|
||||
for i, t in enumerate([x / 10 for x in range(10)] + [1.0, 2.25, 3.5, 4.75]):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), (i * 16) % 256, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
).reformat(format="yuv420p")
|
||||
frame.pts = int(t * 15360)
|
||||
frame.time_base = Fraction(1, 15360)
|
||||
container.mux(video_stream.encode(frame))
|
||||
container.mux(video_stream.encode(None))
|
||||
for offset in range(0, 48000 * 6, 1024):
|
||||
n = min(1024, 48000 * 6 - offset)
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(2, n, dtype=torch.float32).numpy(), format="fltp", layout="stereo"
|
||||
)
|
||||
audio_frame.sample_rate = 48000
|
||||
audio_frame.pts = offset
|
||||
audio_frame.time_base = Fraction(1, 48000)
|
||||
container.mux(audio_stream.encode(audio_frame))
|
||||
container.mux(audio_stream.encode(None))
|
||||
|
||||
buffer.seek(0)
|
||||
result = transcode_and_probe(VideoFromFile(buffer, start_time=1, duration=5))
|
||||
# kept frames: 1.0/2.25/3.5/4.75 s -> rebased span 3.75 s + one nominal interval
|
||||
assert result["frames"] == 4
|
||||
assert result["audio_seconds"] == pytest.approx(4.0, abs=0.45)
|
||||
|
||||
|
||||
def test_save_to_transcode_trims_audio_in_stream_time_base_units():
|
||||
"""Matroska audio timestamps tick in 1/1000, not 1/sample_rate; trim and audio timing
|
||||
must convert through the frame's time base instead of assuming sample units. AAC audio,
|
||||
because it decodes straight to the encoder's format and hits the resampler passthrough
|
||||
that keeps the source time base on the frames."""
|
||||
file_path = create_transcode_source(frames=90, container_format="matroska", audio_codec="aac")
|
||||
try:
|
||||
result = transcode_and_probe(VideoFromFile(file_path, start_time=1, duration=1))
|
||||
assert result["audio_codecs"] == ["aac"]
|
||||
assert result["video_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
finally:
|
||||
os.unlink(file_path)
|
||||
|
||||
|
||||
def test_save_to_transcode_learns_unprobed_audio_params():
|
||||
"""mpegts is only probed a few seconds deep at open, so an audio stream whose first
|
||||
packet comes later (live captures where audio kicks in late) still has sample_rate 0
|
||||
when the transcode starts; the parameters must be learned from the stream itself."""
|
||||
sample_rate, fps, video_seconds, audio_start = 48000, 30, 13, 12
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mpegts") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=fps)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
|
||||
for i in range(video_seconds * fps):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
)
|
||||
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
|
||||
for offset in range(0, (video_seconds - audio_start) * sample_rate, 1024):
|
||||
n = min(1024, (video_seconds - audio_start) * sample_rate - offset)
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(1, n, dtype=torch.float32).numpy(), format="fltp", layout="mono"
|
||||
)
|
||||
audio_frame.sample_rate = sample_rate
|
||||
audio_frame.pts = audio_start * sample_rate + offset
|
||||
container.mux(audio_stream.encode(audio_frame))
|
||||
for stream in (video_stream, audio_stream):
|
||||
container.mux(stream.encode(None))
|
||||
|
||||
buffer.seek(0)
|
||||
with av.open(buffer) as container:
|
||||
# the scenario requires unprobed parameters; if a future FFmpeg probes deeper,
|
||||
# push audio_start/video_seconds further out to restore it
|
||||
assert container.streams.audio[0].codec_context.sample_rate == 0
|
||||
result = transcode_and_probe(VideoFromFile(buffer))
|
||||
assert result["frames"] == video_seconds * fps
|
||||
assert result["audio_codecs"] == ["aac"]
|
||||
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
|
||||
buffer.seek(0)
|
||||
trimmed_before_audio = transcode_and_probe(VideoFromFile(buffer, duration=1))
|
||||
assert trimmed_before_audio["frames"] == fps
|
||||
assert trimmed_before_audio["audio_codecs"] == []
|
||||
assert trimmed_before_audio["audio_seconds"] is None
|
||||
|
||||
buffer.seek(0)
|
||||
trimmed_crossing_audio = transcode_and_probe(VideoFromFile(buffer, start_time=11.5, duration=1))
|
||||
assert trimmed_crossing_audio["frames"] == fps
|
||||
assert trimmed_crossing_audio["audio_codecs"] == ["aac"]
|
||||
assert trimmed_crossing_audio["video_seconds"] == pytest.approx(1.0, abs=0.05)
|
||||
assert trimmed_crossing_audio["audio_seconds"] == pytest.approx(0.5, abs=0.1)
|
||||
|
||||
|
||||
def test_save_to_transcode_trimmed_fragmented_mp4_keeps_audio():
|
||||
"""Fragmented mp4 (MediaRecorder, DASH/HLS-derived files) delivers audio well behind
|
||||
video, so when the trim window's last video frame arrives the audio demuxed so far
|
||||
does not cover the window yet; the transcode must keep demuxing audio until it does
|
||||
instead of finalizing on the first audio frame it sees afterwards."""
|
||||
sample_rate, fps, seconds = 48000, 30, 6
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4", options={"movflags": "frag_keyframe+empty_moov"}) as container:
|
||||
video_stream = container.add_stream("h264", rate=fps)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
|
||||
next_audio_pts = 0
|
||||
for i in range(seconds * fps):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
)
|
||||
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
|
||||
while next_audio_pts / sample_rate <= i / fps: # feed audio alongside, like a live pipeline
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(1, 1024, dtype=torch.float32).numpy(), format="fltp", layout="mono"
|
||||
)
|
||||
audio_frame.sample_rate = sample_rate
|
||||
audio_frame.pts = next_audio_pts
|
||||
container.mux(audio_stream.encode(audio_frame))
|
||||
next_audio_pts += 1024
|
||||
for stream in (video_stream, audio_stream):
|
||||
container.mux(stream.encode(None))
|
||||
|
||||
result = transcode_and_probe(VideoFromFile(buffer, start_time=0.5, duration=1.0))
|
||||
assert result["video_seconds"] == pytest.approx(1.0, abs=0.05)
|
||||
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.05)
|
||||
|
||||
|
||||
def test_save_to_transcode_sparse_video_keeps_true_duration():
|
||||
"""average_rate is not a frame duration: a 3-frame video spanning 60 s averages
|
||||
0.05 fps, and padding the last frame with 1/average_rate used to extend the
|
||||
output — and the audio kept with it — about 20 s past the source span."""
|
||||
sample_rate = 48000
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=30)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
|
||||
for i, second in enumerate((0, 30, 60)):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), i * 80, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
).reformat(format="yuv420p")
|
||||
frame.pts = second * 30
|
||||
frame.time_base = Fraction(1, 30)
|
||||
container.mux(video_stream.encode(frame))
|
||||
for offset in range(0, 90 * sample_rate, 1024):
|
||||
n = min(1024, 90 * sample_rate - offset)
|
||||
audio_frame = av.AudioFrame.from_ndarray(
|
||||
torch.zeros(1, n, dtype=torch.float32).numpy(), format="fltp", layout="mono"
|
||||
)
|
||||
audio_frame.sample_rate = sample_rate
|
||||
audio_frame.pts = offset
|
||||
container.mux(audio_stream.encode(audio_frame))
|
||||
for stream in (video_stream, audio_stream):
|
||||
container.mux(stream.encode(None))
|
||||
|
||||
result = transcode_and_probe(VideoFromFile(buffer))
|
||||
assert result["frames"] == 3
|
||||
# the last frame keeps its true stts duration (1/30 s), not 1/average_rate (~20 s)
|
||||
assert result["video_seconds"] == pytest.approx(60.03, abs=0.05)
|
||||
assert result["audio_seconds"] == pytest.approx(60.03, abs=0.1)
|
||||
|
||||
trimmed = transcode_and_probe(VideoFromFile(buffer, duration=45))
|
||||
assert trimmed["frames"] == 2
|
||||
# a kept frame whose source duration crosses the window end is clamped to it
|
||||
assert trimmed["video_seconds"] == pytest.approx(45.0, abs=0.05)
|
||||
assert trimmed["audio_seconds"] == pytest.approx(45.0, abs=0.1)
|
||||
|
||||
|
||||
def test_save_to_transcode_clamps_final_pts_to_declared_stream_duration():
|
||||
"""Some iPhone MOVs report a video stream duration that ends before the final
|
||||
decoded frame's nominal duration. A transcode must not turn that trailing
|
||||
timestamp quirk into an extra frame interval compared to the source/remux path."""
|
||||
fps = 30
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=fps)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
for i, pts in enumerate([*range(31), 32]):
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
|
||||
).reformat(format="yuv420p")
|
||||
frame.pts = pts
|
||||
frame.time_base = Fraction(1, fps)
|
||||
container.mux(video_stream.encode(frame))
|
||||
container.mux(video_stream.encode(None))
|
||||
|
||||
class _StreamProxy:
|
||||
def __init__(self, stream, duration):
|
||||
self._stream = stream
|
||||
self.duration = duration
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._stream, name)
|
||||
|
||||
class _StreamsProxy:
|
||||
def __init__(self, video_stream):
|
||||
self.video = [video_stream]
|
||||
self.audio = []
|
||||
|
||||
class _PacketProxy:
|
||||
def __init__(self, packet, stream):
|
||||
self._packet = packet
|
||||
self.stream = stream
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._packet, name)
|
||||
|
||||
class _ContainerProxy:
|
||||
def __init__(self, container, stream):
|
||||
self._container = container
|
||||
self._stream = stream
|
||||
self.streams = _StreamsProxy(stream)
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._container, name)
|
||||
|
||||
def demux(self, *streams):
|
||||
for packet in self._container.demux(self._stream._stream):
|
||||
yield _PacketProxy(packet, self._stream)
|
||||
|
||||
buffer.seek(0)
|
||||
output = io.BytesIO()
|
||||
with av.open(buffer) as container:
|
||||
real_stream = container.streams.video[0]
|
||||
declared_duration = 32 * int(round((1 / fps) / real_stream.time_base))
|
||||
stream = _StreamProxy(real_stream, declared_duration)
|
||||
VideoFromFile(buffer)._save_transcoded(
|
||||
_ContainerProxy(container, stream), output, VideoContainer.MP4, VideoCodec.H264, None, 8
|
||||
)
|
||||
|
||||
output.seek(0)
|
||||
with av.open(output) as container:
|
||||
video_stream = container.streams.video[0]
|
||||
frames = [f for p in container.demux(video_stream) for f in p.decode()]
|
||||
assert len(frames) == 32
|
||||
assert float(video_stream.duration * video_stream.time_base) == pytest.approx(32 / fps, abs=0.01)
|
||||
assert float(frames[-1].pts * frames[-1].time_base) == pytest.approx(31 / fps, abs=0.01)
|
||||
|
||||
|
||||
def test_save_to_transcode_irregular_vfr_keeps_span():
|
||||
"""B-frames reorder packets, and mp4 sample durations follow decode order: the dts
|
||||
timeline ends before the pts timeline, so an irregular-VFR source's tail holds fell
|
||||
out of the container (this 20.23 s span used to come out as 15.27 s, and the 10 s
|
||||
trim as 6.03 s). The transcode encodes without B-frames so every sample keeps its
|
||||
true display duration."""
|
||||
durations = [1, 1, 60, 1, 1, 120, 1, 180, 1, 1, 150, 90] # 1/30 s ticks, span 20.2333 s
|
||||
generator = torch.Generator().manual_seed(7)
|
||||
buffer = io.BytesIO()
|
||||
with av.open(buffer, mode="w", format="mp4") as container:
|
||||
video_stream = container.add_stream("mpeg4", rate=30)
|
||||
video_stream.width = video_stream.height = 64
|
||||
video_stream.pix_fmt = "yuv420p"
|
||||
pts = 0
|
||||
for duration in durations:
|
||||
# textured frames, so an encoder with default settings has B-frames to gain from
|
||||
frame = av.VideoFrame.from_ndarray(
|
||||
torch.randint(0, 255, (64, 64, 3), generator=generator, dtype=torch.uint8).numpy(),
|
||||
format="rgb24",
|
||||
).reformat(format="yuv420p")
|
||||
frame.pts = pts
|
||||
frame.time_base = Fraction(1, 30)
|
||||
pts += duration
|
||||
for packet in video_stream.encode(frame):
|
||||
packet.duration = duration # exact stts in the source
|
||||
container.mux(packet)
|
||||
container.mux(video_stream.encode(None))
|
||||
|
||||
result = transcode_and_probe(VideoFromFile(buffer))
|
||||
assert result["frames"] == len(durations)
|
||||
assert result["video_seconds"] == pytest.approx(sum(durations) / 30, abs=0.05)
|
||||
|
||||
trimmed = transcode_and_probe(VideoFromFile(buffer, duration=10))
|
||||
assert trimmed["frames"] == 8 # frames at 12.167 s+ fall outside the window
|
||||
assert trimmed["video_seconds"] == pytest.approx(10.0, abs=0.05)
|
||||
|
||||
|
||||
def test_save_to_transcode_trim_survives_missing_leading_pts():
|
||||
"""A trim should survive pts-less kept frames followed by a real-pts frame past the window."""
|
||||
nulled_frames = 0
|
||||
|
||||
class _PacketProxy:
|
||||
def __init__(self, packet):
|
||||
self._packet = packet
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._packet, name)
|
||||
|
||||
@property
|
||||
def stream(self):
|
||||
return self._packet.stream
|
||||
|
||||
def decode(self):
|
||||
nonlocal nulled_frames
|
||||
frames = self._packet.decode()
|
||||
for frame in frames:
|
||||
if nulled_frames < 2:
|
||||
frame.pts = None
|
||||
nulled_frames += 1
|
||||
return frames
|
||||
|
||||
class _ContainerProxy:
|
||||
def __init__(self, real):
|
||||
self._real = real
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._real, name)
|
||||
|
||||
def demux(self, *streams):
|
||||
for packet in self._real.demux(*streams):
|
||||
yield _PacketProxy(packet)
|
||||
|
||||
file_path = create_transcode_source(frames=10, audio_streams=0)
|
||||
try:
|
||||
buffer = io.BytesIO()
|
||||
with av.open(file_path) as container:
|
||||
# 0.05 s window: both pts-less frames are kept (synthesized pts 0 and 512),
|
||||
# and the first real-pts frame (1024 ticks) already lies past end_pts (768)
|
||||
VideoFromFile(file_path, duration=0.05)._save_transcoded(
|
||||
_ContainerProxy(container), buffer, VideoContainer.MP4, VideoCodec.H264, None, 8
|
||||
)
|
||||
assert nulled_frames == 2
|
||||
buffer.seek(0)
|
||||
with av.open(buffer) as container:
|
||||
video_stream = container.streams.video[0]
|
||||
frames = [f for p in container.demux(video_stream) for f in p.decode()]
|
||||
assert len(frames) == 2
|
||||
assert float(video_stream.duration * video_stream.time_base) == pytest.approx(2 / 30, abs=0.01)
|
||||
finally:
|
||||
os.unlink(file_path)
|
||||
|
||||
|
||||
def test_save_to_transcode_bakes_rotation():
|
||||
"""A 90-degree display-matrix rotation swaps the output dimensions (portrait video)"""
|
||||
file_path = create_transcode_source(width=64, height=32, rotation=True)
|
||||
try:
|
||||
result = transcode_and_probe(VideoFromFile(file_path))
|
||||
assert (result["width"], result["height"]) == (32, 64)
|
||||
assert result["frames"] == 30
|
||||
finally:
|
||||
os.unlink(file_path)
|
||||
|
||||
|
||||
def test_save_to_transcode_skips_undecodable_audio():
|
||||
"""Streaming transcode keeps the decodable audio track and drops undecodable ones;
|
||||
with no decodable audio at all the output is video-only instead of crashing."""
|
||||
mixed = all_bad = None
|
||||
try:
|
||||
mixed = create_transcode_source(audio_streams=1, undecodable_audio=1)
|
||||
all_bad = create_transcode_source(audio_streams=0, undecodable_audio=2)
|
||||
result = transcode_and_probe(VideoFromFile(mixed))
|
||||
assert result["audio_codecs"] == ["aac"]
|
||||
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
|
||||
assert transcode_and_probe(VideoFromFile(all_bad))["audio_codecs"] == []
|
||||
finally:
|
||||
for path in (mixed, all_bad):
|
||||
if path:
|
||||
os.unlink(path)
|
||||
|
||||
@@ -112,17 +112,6 @@ def _make_pid_v1_5_sd(latent_proj_channels=16):
|
||||
return sd
|
||||
|
||||
|
||||
def _make_joyimage_edit_plus_sd():
|
||||
sd = {
|
||||
"img_in.weight": torch.empty(4096, 16, 1, 2, 2, device="meta"),
|
||||
"condition_embedder.time_embedder.linear_1.weight": torch.empty(1, device="meta"),
|
||||
"double_blocks.0.attn.img_attn_q_norm.weight": torch.empty(128, device="meta"),
|
||||
}
|
||||
for i in range(40):
|
||||
sd[f"double_blocks.{i}.attn.img_attn_qkv.weight"] = torch.empty(1, device="meta")
|
||||
return sd
|
||||
|
||||
|
||||
def _add_model_diffusion_prefix(sd):
|
||||
return {f"model.diffusion_model.{k}": v for k, v in sd.items()}
|
||||
|
||||
@@ -269,26 +258,6 @@ class TestModelDetection:
|
||||
assert processed["pixel_blocks.0.adaLN_modulation_msa.bias"].shape == (12288,)
|
||||
assert processed["pixel_blocks.0.adaLN_modulation_mlp.bias"].shape == (12288,)
|
||||
|
||||
def test_joyimage_edit_plus_detection(self):
|
||||
sd = _make_joyimage_edit_plus_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config == {
|
||||
"image_model": "joyimage",
|
||||
"in_channels": 16,
|
||||
"hidden_size": 4096,
|
||||
"patch_size": [1, 2, 2],
|
||||
"num_layers": 40,
|
||||
"num_attention_heads": 32,
|
||||
"text_dim": 4096,
|
||||
}
|
||||
assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "JoyImage"
|
||||
|
||||
def test_incomplete_joyimage_signature_is_not_detected(self):
|
||||
sd = _make_joyimage_edit_plus_sd()
|
||||
del sd["double_blocks.0.attn.img_attn_q_norm.weight"]
|
||||
assert detect_unet_config(sd, "") is None
|
||||
|
||||
def test_unet_config_and_required_keys_combination_is_unique(self):
|
||||
"""Each model in the registry must have a unique combination of
|
||||
``unet_config`` and ``required_keys``. If two models share the same
|
||||
|
||||
@@ -0,0 +1,90 @@
|
||||
"""Shared fixtures for the model download manager tests.
|
||||
|
||||
These run in-process (no ComfyUI subprocess): a file-backed SQLite DB is
|
||||
initialized once, a temp model folder is registered with ``folder_paths``, and
|
||||
the shared aiohttp session is reset between tests so each async test gets a
|
||||
session bound to its own event loop.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _drain_scheduler_tasks(scheduler) -> None:
|
||||
"""Cancel and await live scheduler tasks so none outlive the test.
|
||||
|
||||
Uses the actual task handles rather than only clearing ``_tasks``: each
|
||||
per-test event loop is created by ``asyncio.run``, so a task left behind by
|
||||
a crashed/aborted test would otherwise keep its coroutine alive. We cancel
|
||||
every live task and, when its loop is still usable, run it to completion to
|
||||
let the cancellation propagate before dropping the reference.
|
||||
"""
|
||||
for task in list(scheduler._tasks.values()):
|
||||
if task is None:
|
||||
continue
|
||||
loop = task.get_loop()
|
||||
if task.done() or loop.is_closed():
|
||||
continue
|
||||
task.cancel()
|
||||
if not loop.is_running():
|
||||
try:
|
||||
loop.run_until_complete(asyncio.gather(task, return_exceptions=True))
|
||||
except Exception:
|
||||
pass
|
||||
scheduler._tasks.clear()
|
||||
|
||||
|
||||
@pytest.fixture(scope="session", autouse=True)
|
||||
def _init_db():
|
||||
import app.database.db as db
|
||||
from comfy.cli_args import args
|
||||
|
||||
fd, db_path = tempfile.mkstemp(suffix="-dlmgr-test.sqlite3")
|
||||
os.close(fd)
|
||||
args.database_url = f"sqlite:///{db_path}"
|
||||
db.init_db()
|
||||
yield
|
||||
try:
|
||||
os.remove(db_path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_runtime():
|
||||
"""Reset module singletons that hold event-loop-bound or cross-test state."""
|
||||
import app.model_downloader.net.session as ns
|
||||
from app.model_downloader.scheduler import SCHEDULER
|
||||
|
||||
ns._session = None
|
||||
_drain_scheduler_tasks(SCHEDULER)
|
||||
SCHEDULER._jobs.clear()
|
||||
SCHEDULER._backoff_until.clear()
|
||||
SCHEDULER._started = False
|
||||
yield
|
||||
_drain_scheduler_tasks(SCHEDULER)
|
||||
ns._session = None
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def model_root(tmp_path):
|
||||
"""Register a temp 'loras' model folder and return its absolute path."""
|
||||
import folder_paths
|
||||
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
saved = folder_paths.folder_names_and_paths.get("loras")
|
||||
folder_paths.folder_names_and_paths["loras"] = (
|
||||
[str(root)],
|
||||
{".safetensors", ".sft", ".ckpt", ".pt", ".pth"},
|
||||
)
|
||||
yield str(root)
|
||||
if saved is not None:
|
||||
folder_paths.folder_names_and_paths["loras"] = saved
|
||||
else:
|
||||
folder_paths.folder_names_and_paths.pop("loras", None)
|
||||
@@ -0,0 +1,387 @@
|
||||
"""Unit tests for download authentication.
|
||||
|
||||
Covers env-key resolution, OAuth token resolution/refresh/expiry, provider
|
||||
host matching (including the per-hop drop on a CDN host), and the auth routes.
|
||||
Async tests are driven via ``asyncio.run`` so no pytest-asyncio plugin is needed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from urllib.parse import parse_qs, urlsplit
|
||||
|
||||
import pytest
|
||||
from aiohttp.test_utils import make_mocked_request
|
||||
|
||||
from app.model_downloader.api import routes
|
||||
from app.model_downloader.auth import oauth, token_store
|
||||
from app.model_downloader.auth.providers import PROVIDERS, provider_for_host
|
||||
from app.model_downloader.auth.resolver import resolve_auth_for_hop
|
||||
from app.model_downloader.auth.store import AUTH_STORE
|
||||
from app.model_downloader.auth.token_store import Token
|
||||
|
||||
_HF_ENV = ("HF_TOKEN", "HUGGING_FACE_HUB_TOKEN")
|
||||
_CIVITAI_ENV = ("CIVITAI_API_TOKEN", "CIVITAI_API_KEY")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def auth_tmp(monkeypatch, tmp_path):
|
||||
"""Isolate the on-disk token store and clear the in-memory cache."""
|
||||
d = tmp_path / "download_auth"
|
||||
d.mkdir()
|
||||
monkeypatch.setattr(token_store, "_auth_dir", lambda: str(d))
|
||||
AUTH_STORE._cache.clear()
|
||||
yield
|
||||
AUTH_STORE._cache.clear()
|
||||
|
||||
|
||||
def _clear_env(monkeypatch, *names):
|
||||
for name in names:
|
||||
monkeypatch.delenv(name, raising=False)
|
||||
|
||||
|
||||
# ----- provider host matching -----
|
||||
|
||||
|
||||
def test_provider_for_host():
|
||||
assert provider_for_host("HuggingFace.co:443").name == "huggingface"
|
||||
assert provider_for_host("civitai.com").name == "civitai"
|
||||
# sibling CDN hosts must not match — this is what drops the token on redirect
|
||||
assert provider_for_host("cdn-lfs.huggingface.co") is None
|
||||
assert provider_for_host("cas-bridge.xethub.hf.co") is None
|
||||
assert provider_for_host("example.com") is None
|
||||
|
||||
|
||||
# ----- env-key resolution -----
|
||||
|
||||
|
||||
def test_env_key_resolution_hf(monkeypatch, auth_tmp):
|
||||
monkeypatch.setenv("HF_TOKEN", "hf_env")
|
||||
|
||||
async def _run():
|
||||
auth = await resolve_auth_for_hop("huggingface.co", "https")
|
||||
assert auth is not None
|
||||
assert auth.headers["Authorization"] == "Bearer hf_env"
|
||||
# never over http, never on a CDN redirect host
|
||||
assert await resolve_auth_for_hop("huggingface.co", "http") is None
|
||||
assert await resolve_auth_for_hop("cdn-lfs.huggingface.co", "https") is None
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_env_key_resolution_civitai_secondary_var(monkeypatch, auth_tmp):
|
||||
_clear_env(monkeypatch, "CIVITAI_API_TOKEN")
|
||||
monkeypatch.setenv("CIVITAI_API_KEY", "civ_env")
|
||||
|
||||
async def _run():
|
||||
auth = await resolve_auth_for_hop("civitai.com", "https")
|
||||
assert auth is not None
|
||||
assert auth.headers["Authorization"] == "Bearer civ_env"
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_env_key_takes_precedence_over_oauth(monkeypatch, auth_tmp):
|
||||
monkeypatch.setenv("HF_TOKEN", "hf_env")
|
||||
|
||||
async def _run():
|
||||
AUTH_STORE.set_token("huggingface", Token(access_token="oauth_acc"))
|
||||
auth = await resolve_auth_for_hop("huggingface.co", "https")
|
||||
assert auth.headers["Authorization"] == "Bearer hf_env"
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
# ----- OAuth token resolution / refresh / expiry -----
|
||||
|
||||
|
||||
def test_oauth_token_resolution(monkeypatch, auth_tmp):
|
||||
_clear_env(monkeypatch, *_HF_ENV)
|
||||
|
||||
async def _run():
|
||||
AUTH_STORE.set_token("huggingface", Token(access_token="acc", expires_at=0))
|
||||
auth = await resolve_auth_for_hop("huggingface.co", "https")
|
||||
assert auth is not None
|
||||
assert auth.headers["Authorization"] == "Bearer acc"
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_oauth_refresh_on_expiry(monkeypatch, auth_tmp):
|
||||
_clear_env(monkeypatch, *_HF_ENV)
|
||||
|
||||
async def fake_refresh(provider, tok):
|
||||
return Token(
|
||||
access_token="new_acc",
|
||||
refresh_token="r2",
|
||||
expires_at=int(time.time()) + 3600,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(oauth, "refresh_access_token", fake_refresh)
|
||||
|
||||
async def _run():
|
||||
AUTH_STORE.set_token(
|
||||
"huggingface",
|
||||
Token(access_token="old", refresh_token="r1", expires_at=1),
|
||||
)
|
||||
access = await AUTH_STORE.get_valid_token(PROVIDERS["huggingface"])
|
||||
assert access == "new_acc"
|
||||
# the refreshed token is persisted (cache + disk)
|
||||
assert token_store.load("huggingface").access_token == "new_acc"
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_oauth_expired_without_refresh_returns_none(monkeypatch, auth_tmp):
|
||||
_clear_env(monkeypatch, *_HF_ENV)
|
||||
|
||||
async def _run():
|
||||
AUTH_STORE.set_token(
|
||||
"huggingface",
|
||||
Token(access_token="old", refresh_token=None, expires_at=1),
|
||||
)
|
||||
assert await AUTH_STORE.get_valid_token(PROVIDERS["huggingface"]) is None
|
||||
assert await resolve_auth_for_hop("huggingface.co", "https") is None
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_no_auth_when_nothing_configured(monkeypatch, auth_tmp):
|
||||
_clear_env(monkeypatch, *_HF_ENV, *_CIVITAI_ENV)
|
||||
|
||||
async def _run():
|
||||
assert await resolve_auth_for_hop("huggingface.co", "https") is None
|
||||
assert await resolve_auth_for_hop("example.com", "https") is None
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
# ----- auth routes -----
|
||||
|
||||
|
||||
def test_auth_status_route(monkeypatch, auth_tmp):
|
||||
_clear_env(monkeypatch, *_HF_ENV, *_CIVITAI_ENV)
|
||||
|
||||
async def _run():
|
||||
resp = await routes.auth_status(make_mocked_request("GET", "/api/download/auth"))
|
||||
data = json.loads(resp.body)
|
||||
by_name = {p["provider"]: p for p in data["providers"]}
|
||||
assert set(by_name) == {"huggingface", "civitai"}
|
||||
assert by_name["huggingface"]["logged_in"] is False
|
||||
assert by_name["huggingface"]["env_key_present"] is False
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_login_unconfigured_returns_400(monkeypatch, auth_tmp):
|
||||
_clear_env(monkeypatch, "COMFY_HF_OAUTH_CLIENT_ID")
|
||||
|
||||
async def _run():
|
||||
req = make_mocked_request(
|
||||
"POST", "/api/download/auth/huggingface/login",
|
||||
match_info={"provider": "huggingface"},
|
||||
)
|
||||
resp = await routes.auth_login(req)
|
||||
assert resp.status == 400
|
||||
assert json.loads(resp.body)["error"]["code"] == "OAUTH_NOT_CONFIGURED"
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_login_unknown_provider_returns_400(auth_tmp):
|
||||
async def _run():
|
||||
req = make_mocked_request(
|
||||
"POST", "/api/download/auth/nope/login",
|
||||
match_info={"provider": "nope"},
|
||||
)
|
||||
resp = await routes.auth_login(req)
|
||||
assert resp.status == 400
|
||||
assert json.loads(resp.body)["error"]["code"] == "UNKNOWN_PROVIDER"
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_login_start_and_in_progress(monkeypatch, auth_tmp):
|
||||
monkeypatch.setenv("COMFY_HF_OAUTH_CLIENT_ID", "test-client")
|
||||
|
||||
async def _run():
|
||||
try:
|
||||
req = make_mocked_request(
|
||||
"POST", "/api/download/auth/huggingface/login",
|
||||
match_info={"provider": "huggingface"},
|
||||
)
|
||||
resp = await routes.auth_login(req)
|
||||
assert resp.status == 200
|
||||
url = json.loads(resp.body)["authorize_url"]
|
||||
assert url.startswith("https://huggingface.co/oauth/authorize?")
|
||||
assert "code_challenge=" in url and "code_challenge_method=S256" in url
|
||||
assert "client_id=test-client" in url
|
||||
# a second concurrent login is rejected
|
||||
resp2 = await routes.auth_login(
|
||||
make_mocked_request(
|
||||
"POST", "/api/download/auth/huggingface/login",
|
||||
match_info={"provider": "huggingface"},
|
||||
)
|
||||
)
|
||||
assert resp2.status == 409
|
||||
finally:
|
||||
flow = oauth._ACTIVE.get("huggingface")
|
||||
if flow is not None:
|
||||
await flow._teardown()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def _redirect_uri(authorize_url: str) -> str:
|
||||
return parse_qs(urlsplit(authorize_url).query)["redirect_uri"][0]
|
||||
|
||||
|
||||
def test_callback_uri_is_fixed_loopback_port(monkeypatch, auth_tmp):
|
||||
"""The redirect targets the fixed loopback port, ignoring the request Host.
|
||||
|
||||
In dev the UI is served by Vite on a different port that proxies ``/api``,
|
||||
so the login request arrives with the frontend port in ``Host``; the
|
||||
callback URI must stay pinned to the registered loopback port.
|
||||
"""
|
||||
monkeypatch.setenv("COMFY_HF_OAUTH_CLIENT_ID", "hf-client")
|
||||
|
||||
async def _run():
|
||||
try:
|
||||
req = make_mocked_request(
|
||||
"POST", "/api/download/auth/huggingface/login",
|
||||
headers={"Host": "localhost:5173"},
|
||||
match_info={"provider": "huggingface"},
|
||||
)
|
||||
resp = await routes.auth_login(req)
|
||||
redirect = _redirect_uri(json.loads(resp.body)["authorize_url"])
|
||||
assert redirect == f"http://127.0.0.1:{oauth.CALLBACK_PORT}/callback/huggingface"
|
||||
finally:
|
||||
flow = oauth._ACTIVE.get("huggingface")
|
||||
if flow is not None:
|
||||
await flow._teardown()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_second_login_rejected_single_flight(monkeypatch, auth_tmp):
|
||||
"""Only one login runs at a time — the shared callback port is single-flight."""
|
||||
monkeypatch.setenv("COMFY_HF_OAUTH_CLIENT_ID", "hf-client")
|
||||
monkeypatch.setenv("COMFY_CIVITAI_OAUTH_CLIENT_ID", "civ-client")
|
||||
|
||||
async def _run():
|
||||
try:
|
||||
hf = await routes.auth_login(
|
||||
make_mocked_request(
|
||||
"POST", "/api/download/auth/huggingface/login",
|
||||
match_info={"provider": "huggingface"},
|
||||
)
|
||||
)
|
||||
assert hf.status == 200
|
||||
# a different provider can't start while HF holds the callback port
|
||||
civ = await routes.auth_login(
|
||||
make_mocked_request(
|
||||
"POST", "/api/download/auth/civitai/login",
|
||||
match_info={"provider": "civitai"},
|
||||
)
|
||||
)
|
||||
assert civ.status == 409
|
||||
assert json.loads(civ.body)["error"]["code"] == "LOGIN_IN_PROGRESS"
|
||||
assert oauth.login_in_progress("huggingface")
|
||||
assert not oauth.login_in_progress("civitai")
|
||||
finally:
|
||||
for name in ("huggingface", "civitai"):
|
||||
flow = oauth._ACTIVE.get(name)
|
||||
if flow is not None:
|
||||
await flow._teardown()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_callback_completes_login(monkeypatch, auth_tmp):
|
||||
"""The loopback callback exchanges the code, stores the token, tears down."""
|
||||
monkeypatch.setenv("COMFY_HF_OAUTH_CLIENT_ID", "hf-client")
|
||||
_clear_env(monkeypatch, *_HF_ENV)
|
||||
|
||||
async def fake_exchange(provider, code, verifier, redirect_uri):
|
||||
assert code == "the-code"
|
||||
return Token(access_token="acc_from_callback")
|
||||
|
||||
monkeypatch.setattr(oauth, "exchange_code", fake_exchange)
|
||||
|
||||
async def _run():
|
||||
try:
|
||||
login = await routes.auth_login(
|
||||
make_mocked_request(
|
||||
"POST", "/api/download/auth/huggingface/login",
|
||||
match_info={"provider": "huggingface"},
|
||||
)
|
||||
)
|
||||
assert login.status == 200
|
||||
flow = oauth._ACTIVE["huggingface"]
|
||||
cb = await flow._handle_callback(
|
||||
make_mocked_request(
|
||||
"GET",
|
||||
f"/callback/huggingface?state={flow.state}&code=the-code",
|
||||
match_info={"provider": "huggingface"},
|
||||
)
|
||||
)
|
||||
assert cb.status == 200
|
||||
assert "window.close" in cb.text
|
||||
assert AUTH_STORE.status(PROVIDERS["huggingface"])["logged_in"] is True
|
||||
# the handler schedules its own teardown; let it run
|
||||
await asyncio.sleep(0.05)
|
||||
assert not oauth.login_in_progress("huggingface")
|
||||
finally:
|
||||
flow = oauth._ACTIVE.get("huggingface")
|
||||
if flow is not None:
|
||||
await flow._teardown()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_callback_rejects_bad_state(monkeypatch, auth_tmp):
|
||||
monkeypatch.setenv("COMFY_HF_OAUTH_CLIENT_ID", "hf-client")
|
||||
|
||||
async def _run():
|
||||
try:
|
||||
await routes.auth_login(
|
||||
make_mocked_request(
|
||||
"POST", "/api/download/auth/huggingface/login",
|
||||
match_info={"provider": "huggingface"},
|
||||
)
|
||||
)
|
||||
flow = oauth._ACTIVE["huggingface"]
|
||||
cb = await flow._handle_callback(
|
||||
make_mocked_request(
|
||||
"GET",
|
||||
"/callback/huggingface?state=wrong&code=x",
|
||||
match_info={"provider": "huggingface"},
|
||||
)
|
||||
)
|
||||
assert cb.status == 400
|
||||
# the flow stays pending so a genuine callback can still arrive
|
||||
assert oauth.login_in_progress("huggingface")
|
||||
finally:
|
||||
flow = oauth._ACTIVE.get("huggingface")
|
||||
if flow is not None:
|
||||
await flow._teardown()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_logout_route_clears_token(auth_tmp):
|
||||
async def _run():
|
||||
AUTH_STORE.set_token("civitai", Token(access_token="x"))
|
||||
assert AUTH_STORE.status(PROVIDERS["civitai"])["logged_in"] is True
|
||||
req = make_mocked_request(
|
||||
"POST", "/api/download/auth/civitai/logout",
|
||||
match_info={"provider": "civitai"},
|
||||
)
|
||||
resp = await routes.auth_logout(req)
|
||||
assert resp.status == 200
|
||||
assert json.loads(resp.body)["logged_out"] is True
|
||||
assert AUTH_STORE.status(PROVIDERS["civitai"])["logged_in"] is False
|
||||
|
||||
asyncio.run(_run())
|
||||
@@ -0,0 +1,136 @@
|
||||
"""Unit tests for ``DownloadManager.delete`` and ``DownloadManager.clear``.
|
||||
|
||||
Deleting a terminal row must remove it from history for good (so it does not
|
||||
reappear on the next ``list``), leave live rows untouched, and clean up any
|
||||
leftover ``.part`` temp file without touching the finished model file.
|
||||
|
||||
``clear()`` is the bulk variant: it removes all terminal rows atomically, skips
|
||||
live ones, and returns the count of rows deleted.
|
||||
|
||||
Async methods are driven via ``asyncio.run`` so no pytest-asyncio plugin is
|
||||
required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
from app.model_downloader.constants import DownloadStatus
|
||||
from app.model_downloader.database import queries
|
||||
from app.model_downloader.manager import DOWNLOAD_MANAGER, DownloadError
|
||||
|
||||
|
||||
def _insert(download_id: str, status: str, *, temp_path: str = "/tmp/none.part") -> None:
|
||||
queries.insert_download(
|
||||
{
|
||||
"id": download_id,
|
||||
"url": "https://huggingface.co/org/model.safetensors",
|
||||
"model_id": "loras/model.safetensors",
|
||||
"dest_path": "/tmp/model.safetensors",
|
||||
"temp_path": temp_path,
|
||||
"status": status,
|
||||
"priority": 0,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_delete_removes_terminal_row_from_history():
|
||||
_insert("done", DownloadStatus.COMPLETED)
|
||||
|
||||
asyncio.run(DOWNLOAD_MANAGER.delete("done"))
|
||||
|
||||
assert queries.get_download("done") is None
|
||||
|
||||
|
||||
def test_delete_refuses_live_row():
|
||||
_insert("live", DownloadStatus.QUEUED)
|
||||
|
||||
with pytest.raises(DownloadError) as excinfo:
|
||||
asyncio.run(DOWNLOAD_MANAGER.delete("live"))
|
||||
|
||||
assert excinfo.value.code == "DOWNLOAD_ACTIVE"
|
||||
assert queries.get_download("live") is not None
|
||||
|
||||
|
||||
def test_delete_missing_row_raises_not_found():
|
||||
with pytest.raises(DownloadError) as excinfo:
|
||||
asyncio.run(DOWNLOAD_MANAGER.delete("nope"))
|
||||
|
||||
assert excinfo.value.code == "NOT_FOUND"
|
||||
|
||||
|
||||
def test_delete_removes_leftover_temp_file(tmp_path):
|
||||
partial = tmp_path / "model.safetensors.part"
|
||||
partial.write_bytes(b"partial")
|
||||
_insert("failed", DownloadStatus.FAILED, temp_path=str(partial))
|
||||
|
||||
asyncio.run(DOWNLOAD_MANAGER.delete("failed"))
|
||||
|
||||
assert not os.path.exists(partial)
|
||||
assert queries.get_download("failed") is None
|
||||
|
||||
|
||||
# ----- clear -----
|
||||
|
||||
|
||||
def test_clear_removes_all_terminal_rows():
|
||||
_insert("c-done", DownloadStatus.COMPLETED)
|
||||
_insert("c-fail", DownloadStatus.FAILED)
|
||||
_insert("c-canc", DownloadStatus.CANCELLED)
|
||||
|
||||
deleted = asyncio.run(DOWNLOAD_MANAGER.clear())
|
||||
|
||||
assert deleted == 3
|
||||
assert queries.get_download("c-done") is None
|
||||
assert queries.get_download("c-fail") is None
|
||||
assert queries.get_download("c-canc") is None
|
||||
|
||||
|
||||
def test_clear_skips_live_rows():
|
||||
_insert("cl-queued", DownloadStatus.QUEUED)
|
||||
_insert("cl-paused", DownloadStatus.PAUSED)
|
||||
_insert("cl-done", DownloadStatus.COMPLETED)
|
||||
|
||||
deleted = asyncio.run(DOWNLOAD_MANAGER.clear())
|
||||
|
||||
assert deleted == 1
|
||||
assert queries.get_download("cl-queued") is not None
|
||||
assert queries.get_download("cl-paused") is not None
|
||||
assert queries.get_download("cl-done") is None
|
||||
|
||||
|
||||
def test_clear_returns_zero_when_nothing_to_delete():
|
||||
_insert("cl-only-live", DownloadStatus.QUEUED)
|
||||
|
||||
deleted = asyncio.run(DOWNLOAD_MANAGER.clear())
|
||||
|
||||
assert deleted == 0
|
||||
assert queries.get_download("cl-only-live") is not None
|
||||
|
||||
|
||||
def test_clear_removes_leftover_temp_files(tmp_path):
|
||||
partial = tmp_path / "clear_partial.part"
|
||||
partial.write_bytes(b"partial data")
|
||||
finished = tmp_path / "finished.safetensors"
|
||||
finished.write_bytes(b"real model weights")
|
||||
|
||||
_insert("cl-part", DownloadStatus.FAILED, temp_path=str(partial))
|
||||
# The finished file is not the temp_path; temp_path for a completed download
|
||||
# no longer exists (already renamed), so use a non-existent path here to
|
||||
# verify clear() tolerates a missing temp file without raising.
|
||||
_insert("cl-comp", DownloadStatus.COMPLETED, temp_path=str(tmp_path / "gone.part"))
|
||||
|
||||
asyncio.run(DOWNLOAD_MANAGER.clear())
|
||||
|
||||
# Leftover .part from the failed download is cleaned up.
|
||||
assert not partial.exists()
|
||||
# Finished model file is never touched.
|
||||
assert finished.exists()
|
||||
|
||||
|
||||
def test_clear_empty_db_returns_zero():
|
||||
deleted = asyncio.run(DOWNLOAD_MANAGER.clear())
|
||||
assert deleted == 0
|
||||
@@ -0,0 +1,637 @@
|
||||
"""Integration tests for the download engine against a local aiohttp server.
|
||||
|
||||
Covers single-stream and segmented transfers, deterministic resume from a
|
||||
partial file, and cancel rollback. Async tests are driven via ``asyncio.run``
|
||||
so no pytest-asyncio plugin is required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import struct
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
from aiohttp import web
|
||||
|
||||
from comfy.cli_args import args
|
||||
from app.model_downloader.constants import DownloadStatus
|
||||
from app.model_downloader.database import queries
|
||||
from app.model_downloader.engine.job import DownloadJob, JobSpec
|
||||
from app.model_downloader.net.session import close_session
|
||||
from app.model_downloader.security import paths
|
||||
|
||||
PAYLOAD_ETAG = '"v1"'
|
||||
|
||||
|
||||
def _payload(n: int) -> bytes:
|
||||
return bytes((i * 37 + 11) % 256 for i in range(n))
|
||||
|
||||
|
||||
def _safetensors_payload(total: int) -> bytes:
|
||||
"""A structurally valid ``.safetensors`` blob of exactly ``total`` bytes.
|
||||
|
||||
Success-path tests download to ``.safetensors`` destinations, which the
|
||||
engine now structurally validates before the atomic rename, so their
|
||||
payloads must parse as real safetensors (header length + JSON header +
|
||||
data region whose size matches the declared ``data_offsets``).
|
||||
"""
|
||||
def _header(data_len: int) -> bytes:
|
||||
return json.dumps(
|
||||
{"w": {"dtype": "U8", "shape": [data_len], "data_offsets": [0, data_len]}}
|
||||
).encode("utf-8")
|
||||
|
||||
# The header's byte length depends on the digit count of ``data_len``, so
|
||||
# iterate until ``total == 8 + len(header) + data_len`` is self-consistent.
|
||||
data_len = total - 8 - len(_header(total))
|
||||
for _ in range(8):
|
||||
header = _header(data_len)
|
||||
new_data_len = total - 8 - len(header)
|
||||
if new_data_len == data_len:
|
||||
break
|
||||
data_len = new_data_len
|
||||
assert data_len >= 0, "total too small for a safetensors payload"
|
||||
header = _header(data_len)
|
||||
body = bytes((i * 37 + 11) % 256 for i in range(data_len))
|
||||
return struct.pack("<Q", len(header)) + header + body
|
||||
|
||||
|
||||
def _range_handler(payload: bytes):
|
||||
async def handler(request: web.Request) -> web.Response:
|
||||
rng = request.headers.get("Range")
|
||||
if rng:
|
||||
spec = rng.split("=", 1)[1]
|
||||
s, _, e = spec.partition("-")
|
||||
start = int(s)
|
||||
end = int(e) if e else len(payload) - 1
|
||||
chunk = payload[start : end + 1]
|
||||
return web.Response(
|
||||
status=206,
|
||||
body=chunk,
|
||||
headers={
|
||||
"Content-Range": f"bytes {start}-{end}/{len(payload)}",
|
||||
"Accept-Ranges": "bytes",
|
||||
"ETag": PAYLOAD_ETAG,
|
||||
},
|
||||
)
|
||||
return web.Response(
|
||||
status=200, body=payload, headers={"Accept-Ranges": "bytes", "ETag": PAYLOAD_ETAG}
|
||||
)
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _content_disposition_handler(payload: bytes, filename: str):
|
||||
"""A range-capable server that only reveals its filename via a header.
|
||||
|
||||
Models a Civitai-style ``/api/download/...`` endpoint: the URL path has no
|
||||
extension, and the real filename (hence extension) lives in the response
|
||||
``Content-Disposition`` header.
|
||||
"""
|
||||
|
||||
async def handler(request: web.Request) -> web.Response:
|
||||
headers = {
|
||||
"Accept-Ranges": "bytes",
|
||||
"ETag": PAYLOAD_ETAG,
|
||||
"Content-Disposition": f'attachment; filename="{filename}"',
|
||||
}
|
||||
rng = request.headers.get("Range")
|
||||
if rng:
|
||||
spec = rng.split("=", 1)[1]
|
||||
s, _, e = spec.partition("-")
|
||||
start = int(s)
|
||||
end = int(e) if e else len(payload) - 1
|
||||
chunk = payload[start : end + 1]
|
||||
return web.Response(
|
||||
status=206,
|
||||
body=chunk,
|
||||
headers={**headers, "Content-Range": f"bytes {start}-{end}/{len(payload)}"},
|
||||
)
|
||||
return web.Response(status=200, body=payload, headers=headers)
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _noranges_handler(payload: bytes):
|
||||
async def handler(request: web.Request) -> web.Response:
|
||||
# Always full body, never advertises Accept-Ranges -> single-stream.
|
||||
return web.Response(status=200, body=payload)
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _slow_handler(payload: bytes, chunk: int = 16384, delay: float = 0.01):
|
||||
async def handler(request: web.Request) -> web.StreamResponse:
|
||||
resp = web.StreamResponse(
|
||||
status=200, headers={"Content-Length": str(len(payload))}
|
||||
)
|
||||
await resp.prepare(request)
|
||||
for i in range(0, len(payload), chunk):
|
||||
await resp.write(payload[i : i + chunk])
|
||||
await asyncio.sleep(delay)
|
||||
await resp.write_eof()
|
||||
return resp
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _overflow_range_handler(payload: bytes, extra: int = 256 * 1024):
|
||||
"""A non-conforming 206 server that returns MORE than the requested range."""
|
||||
|
||||
async def handler(request: web.Request) -> web.Response:
|
||||
rng = request.headers.get("Range")
|
||||
if rng:
|
||||
spec = rng.split("=", 1)[1]
|
||||
s, _, e = spec.partition("-")
|
||||
start = int(s)
|
||||
end = int(e) if e else len(payload) - 1
|
||||
# Maliciously overrun: append extra bytes past the requested end.
|
||||
body = payload[start : end + 1] + bytes(extra)
|
||||
return web.Response(
|
||||
status=206,
|
||||
body=body,
|
||||
headers={
|
||||
"Content-Range": f"bytes {start}-{end}/{len(payload)}",
|
||||
"Accept-Ranges": "bytes",
|
||||
"ETag": PAYLOAD_ETAG,
|
||||
},
|
||||
)
|
||||
return web.Response(
|
||||
status=200, body=payload, headers={"Accept-Ranges": "bytes", "ETag": PAYLOAD_ETAG}
|
||||
)
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _short_range_handler(payload: bytes, drop: int = 64 * 1024):
|
||||
"""A 206 server that returns fewer bytes than requested for later segments.
|
||||
|
||||
Simulates a server cleanly closing a range connection early. The response
|
||||
is internally consistent (Content-Length matches the short body), so the
|
||||
client sees no error and the segment just ends short, leaving a zero-filled
|
||||
hole in the preallocated file.
|
||||
"""
|
||||
|
||||
async def handler(request: web.Request) -> web.Response:
|
||||
rng = request.headers.get("Range")
|
||||
if rng:
|
||||
spec = rng.split("=", 1)[1]
|
||||
s, _, e = spec.partition("-")
|
||||
start = int(s)
|
||||
end = int(e) if e else len(payload) - 1
|
||||
chunk = payload[start : end + 1]
|
||||
if start > 0 and len(chunk) > drop:
|
||||
chunk = chunk[:-drop] # truncate a non-first segment
|
||||
return web.Response(
|
||||
status=206,
|
||||
body=chunk,
|
||||
headers={
|
||||
"Content-Range": f"bytes {start}-{end}/{len(payload)}",
|
||||
"Accept-Ranges": "bytes",
|
||||
"ETag": PAYLOAD_ETAG,
|
||||
},
|
||||
)
|
||||
return web.Response(
|
||||
status=200, body=payload, headers={"Accept-Ranges": "bytes", "ETag": PAYLOAD_ETAG}
|
||||
)
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _unbounded_handler(total: int, chunk: int = 16384):
|
||||
"""A 200 stream with no Content-Length / Accept-Ranges (unknown length)."""
|
||||
|
||||
async def handler(request: web.Request) -> web.StreamResponse:
|
||||
resp = web.StreamResponse(status=200)
|
||||
await resp.prepare(request)
|
||||
sent = 0
|
||||
while sent < total:
|
||||
await resp.write(bytes(min(chunk, total - sent)))
|
||||
sent += chunk
|
||||
await resp.write_eof()
|
||||
return resp
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
async def _serve(handler):
|
||||
app = web.Application()
|
||||
app.router.add_route("*", "/{name:.*}", handler)
|
||||
runner = web.AppRunner(app)
|
||||
await runner.setup()
|
||||
site = web.TCPSite(runner, "127.0.0.1", 0)
|
||||
await site.start()
|
||||
port = site._server.sockets[0].getsockname()[1]
|
||||
return runner, port
|
||||
|
||||
|
||||
def _insert(model_id: str, url: str, status: str = DownloadStatus.QUEUED) -> tuple[str, str, str]:
|
||||
final_path, temp_path = paths.resolve_destination(model_id)
|
||||
download_id = str(uuid.uuid4())
|
||||
queries.insert_download(
|
||||
{
|
||||
"id": download_id,
|
||||
"url": url,
|
||||
"model_id": model_id,
|
||||
"dest_path": final_path,
|
||||
"temp_path": temp_path,
|
||||
"status": status,
|
||||
}
|
||||
)
|
||||
return download_id, final_path, temp_path
|
||||
|
||||
|
||||
# ----- single-stream -----
|
||||
|
||||
|
||||
def test_single_stream_download(model_root):
|
||||
payload = _safetensors_payload(300_000)
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
runner, port = await _serve(_noranges_handler(payload))
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/model.safetensors"
|
||||
did, final_path, _temp = _insert("loras/single.safetensors", url)
|
||||
job = DownloadJob(JobSpec(
|
||||
download_id=did, url=url, model_id="loras/single.safetensors",
|
||||
dest_path=final_path, temp_path=_temp,
|
||||
))
|
||||
status = await job.run()
|
||||
assert status == DownloadStatus.COMPLETED, queries.get_download(did).error
|
||||
assert os.path.exists(final_path)
|
||||
assert open(final_path, "rb").read() == payload
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
# ----- segmented -----
|
||||
|
||||
|
||||
def test_segmented_download(model_root):
|
||||
payload = _safetensors_payload(4 * 1024 * 1024) # 4 MiB -> multiple segments
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
runner, port = await _serve(_range_handler(payload))
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/model.safetensors"
|
||||
did, final_path, temp = _insert("loras/seg.safetensors", url)
|
||||
job = DownloadJob(JobSpec(
|
||||
download_id=did, url=url, model_id="loras/seg.safetensors",
|
||||
dest_path=final_path, temp_path=temp,
|
||||
))
|
||||
status = await job.run()
|
||||
assert status == DownloadStatus.COMPLETED, queries.get_download(did).error
|
||||
assert open(final_path, "rb").read() == payload
|
||||
# More than one segment row was planned.
|
||||
assert len(queries.list_segments(did)) > 1
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
# ----- deterministic resume from a partial file -----
|
||||
|
||||
|
||||
def test_resume_from_partial(model_root):
|
||||
payload = _safetensors_payload(512 * 1024) # < 1 MiB -> single segment
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
runner, port = await _serve(_range_handler(payload))
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/model.safetensors"
|
||||
did, final_path, temp = _insert("loras/resume.safetensors", url)
|
||||
# Simulate a prior partial: first 200 KiB already written, offset persisted.
|
||||
prefix = 200 * 1024
|
||||
os.makedirs(os.path.dirname(temp), exist_ok=True)
|
||||
with open(temp, "wb") as f:
|
||||
f.write(payload[:prefix])
|
||||
queries.update_download(did, bytes_done=prefix, etag=PAYLOAD_ETAG)
|
||||
|
||||
job = DownloadJob(JobSpec(
|
||||
download_id=did, url=url, model_id="loras/resume.safetensors",
|
||||
dest_path=final_path, temp_path=temp, etag=PAYLOAD_ETAG,
|
||||
))
|
||||
status = await job.run()
|
||||
assert status == DownloadStatus.COMPLETED, queries.get_download(did).error
|
||||
assert open(final_path, "rb").read() == payload
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
# ----- cancel rollback -----
|
||||
|
||||
|
||||
def test_cancel_rollback(model_root, monkeypatch):
|
||||
monkeypatch.setattr(args, "download_chunk_size", 16384, raising=False)
|
||||
payload = _payload(1024 * 1024)
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
runner, port = await _serve(_slow_handler(payload))
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/model.safetensors"
|
||||
did, final_path, temp = _insert("loras/cancel.safetensors", url)
|
||||
job = DownloadJob(JobSpec(
|
||||
download_id=did, url=url, model_id="loras/cancel.safetensors",
|
||||
dest_path=final_path, temp_path=temp,
|
||||
))
|
||||
task = asyncio.ensure_future(job.run())
|
||||
# Wait until some bytes have been written, then cancel.
|
||||
for _ in range(200):
|
||||
await asyncio.sleep(0.01)
|
||||
if job.state.bytes_done > 0:
|
||||
break
|
||||
job.request_cancel()
|
||||
status = await task
|
||||
assert status == DownloadStatus.CANCELLED
|
||||
assert not os.path.exists(temp)
|
||||
assert not os.path.exists(final_path)
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
# ----- size-bound enforcement (malicious / non-conforming hosts) -----
|
||||
|
||||
|
||||
def test_segment_overflow_aborts(model_root):
|
||||
"""A 206 returning more than the requested range must not overrun."""
|
||||
payload = _payload(4 * 1024 * 1024) # large enough to segment
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
runner, port = await _serve(_overflow_range_handler(payload))
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/model.safetensors"
|
||||
did, final_path, temp = _insert("loras/overflow.safetensors", url)
|
||||
job = DownloadJob(JobSpec(
|
||||
download_id=did, url=url, model_id="loras/overflow.safetensors",
|
||||
dest_path=final_path, temp_path=temp,
|
||||
))
|
||||
status = await job.run()
|
||||
assert status == DownloadStatus.FAILED
|
||||
assert not os.path.exists(final_path)
|
||||
assert not os.path.exists(temp)
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_short_segment_fails_closed(model_root):
|
||||
"""A segment that ends short must fail, not be accepted as complete.
|
||||
|
||||
The file is preallocated to total_bytes, so the on-disk size still equals
|
||||
total even with a zero-filled hole; completeness must be judged per-segment.
|
||||
"""
|
||||
payload = _safetensors_payload(4 * 1024 * 1024) # large enough to segment
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
runner, port = await _serve(_short_range_handler(payload))
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/model.safetensors"
|
||||
did, final_path, temp = _insert("loras/short.safetensors", url)
|
||||
job = DownloadJob(JobSpec(
|
||||
download_id=did, url=url, model_id="loras/short.safetensors",
|
||||
dest_path=final_path, temp_path=temp,
|
||||
))
|
||||
status = await job.run()
|
||||
assert status == DownloadStatus.FAILED, queries.get_download(did).error
|
||||
assert "incomplete" in (queries.get_download(did).error or "")
|
||||
assert not os.path.exists(final_path)
|
||||
assert not os.path.exists(temp)
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_structural_validation_rejects_corrupt(model_root):
|
||||
"""A correctly sized but structurally invalid file fails closed (not retried).
|
||||
|
||||
Regression for the dead structural gate: validation must key off the
|
||||
destination extension, not the ``.part`` temp suffix.
|
||||
"""
|
||||
payload = _payload(300_000) # right size, but not a valid safetensors blob
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
runner, port = await _serve(_noranges_handler(payload))
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/model.safetensors"
|
||||
did, final_path, temp = _insert("loras/corrupt.safetensors", url)
|
||||
job = DownloadJob(JobSpec(
|
||||
download_id=did, url=url, model_id="loras/corrupt.safetensors",
|
||||
dest_path=final_path, temp_path=temp,
|
||||
))
|
||||
status = await job.run()
|
||||
assert status == DownloadStatus.FAILED, queries.get_download(did).error
|
||||
assert not os.path.exists(final_path)
|
||||
assert not os.path.exists(temp)
|
||||
# Failed closed at first attempt, not re-queued as retryable.
|
||||
assert queries.get_download(did).attempts == 0
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_rejects_oversized_known_download(model_root, monkeypatch):
|
||||
"""A file whose advertised size exceeds the cap is rejected at probe."""
|
||||
monkeypatch.setattr(args, "download_max_bytes", 100_000, raising=False)
|
||||
payload = _payload(300_000)
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
runner, port = await _serve(_noranges_handler(payload))
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/model.safetensors"
|
||||
did, final_path, temp = _insert("loras/toobig.safetensors", url)
|
||||
job = DownloadJob(JobSpec(
|
||||
download_id=did, url=url, model_id="loras/toobig.safetensors",
|
||||
dest_path=final_path, temp_path=temp,
|
||||
))
|
||||
status = await job.run()
|
||||
assert status == DownloadStatus.FAILED
|
||||
assert not os.path.exists(final_path)
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_unknown_length_capped_by_max_bytes(model_root, monkeypatch):
|
||||
"""An unbounded unknown-length stream is capped by --download-max-bytes."""
|
||||
monkeypatch.setattr(args, "download_max_bytes", 100_000, raising=False)
|
||||
monkeypatch.setattr(args, "download_chunk_size", 16384, raising=False)
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
runner, port = await _serve(_unbounded_handler(2 * 1024 * 1024))
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/model.safetensors"
|
||||
did, final_path, temp = _insert("loras/unbounded.safetensors", url)
|
||||
job = DownloadJob(JobSpec(
|
||||
download_id=did, url=url, model_id="loras/unbounded.safetensors",
|
||||
dest_path=final_path, temp_path=temp,
|
||||
))
|
||||
status = await job.run()
|
||||
assert status == DownloadStatus.FAILED
|
||||
assert not os.path.exists(final_path)
|
||||
assert not os.path.exists(temp)
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
# ----- manager + scheduler end-to-end -----
|
||||
|
||||
|
||||
def test_manager_enqueue_to_completion(model_root):
|
||||
payload = _safetensors_payload(2 * 1024 * 1024)
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
from app.model_downloader.manager import DOWNLOAD_MANAGER
|
||||
|
||||
runner, port = await _serve(_range_handler(payload))
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/model.safetensors"
|
||||
did = await DOWNLOAD_MANAGER.enqueue(url, "loras/e2e.safetensors")
|
||||
# Wait for completion.
|
||||
final_path, _ = paths.resolve_destination("loras/e2e.safetensors")
|
||||
for _ in range(500):
|
||||
await asyncio.sleep(0.02)
|
||||
row = queries.get_download(did)
|
||||
if row.status in DownloadStatus.TERMINAL:
|
||||
break
|
||||
row = queries.get_download(did)
|
||||
assert row.status == DownloadStatus.COMPLETED, row.error
|
||||
assert open(final_path, "rb").read() == payload
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_manager_rejects_disallowed_url(model_root):
|
||||
async def _run():
|
||||
from app.model_downloader.manager import DOWNLOAD_MANAGER, DownloadError
|
||||
|
||||
with pytest.raises(DownloadError) as ei:
|
||||
await DOWNLOAD_MANAGER.enqueue(
|
||||
"https://evil.example.com/x.safetensors", "loras/bad.safetensors"
|
||||
)
|
||||
assert ei.value.code == "URL_NOT_ALLOWED"
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_manager_resolves_extensionless_url(model_root):
|
||||
"""An allowlisted URL with no extension in its path is resolved from the
|
||||
response, and the stored file adopts the resolved extension."""
|
||||
payload = _safetensors_payload(1 * 1024 * 1024)
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
from app.model_downloader.manager import DOWNLOAD_MANAGER
|
||||
|
||||
runner, port = await _serve(
|
||||
_content_disposition_handler(payload, "RealModel.safetensors")
|
||||
)
|
||||
try:
|
||||
# No extension in the path (Civitai-style) and none in the model_id.
|
||||
url = f"http://127.0.0.1:{port}/api/download/models/12345"
|
||||
did = await DOWNLOAD_MANAGER.enqueue(url, "loras/my_civitai_model")
|
||||
|
||||
row = queries.get_download(did)
|
||||
# The resolved extension was appended to the model_id + destination.
|
||||
assert row.model_id == "loras/my_civitai_model.safetensors"
|
||||
assert row.dest_path.endswith("my_civitai_model.safetensors")
|
||||
|
||||
final_path, _ = paths.resolve_destination(
|
||||
"loras/my_civitai_model.safetensors"
|
||||
)
|
||||
for _ in range(500):
|
||||
await asyncio.sleep(0.02)
|
||||
row = queries.get_download(did)
|
||||
if row.status in DownloadStatus.TERMINAL:
|
||||
break
|
||||
row = queries.get_download(did)
|
||||
assert row.status == DownloadStatus.COMPLETED, row.error
|
||||
assert open(final_path, "rb").read() == payload
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_manager_overrides_extension_from_resolution(model_root):
|
||||
"""A model_id carrying a different known extension is corrected to match
|
||||
the resolved URL's extension."""
|
||||
payload = _safetensors_payload(256 * 1024)
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
from app.model_downloader.manager import DOWNLOAD_MANAGER
|
||||
|
||||
runner, port = await _serve(
|
||||
_content_disposition_handler(payload, "weights.safetensors")
|
||||
)
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/api/download/models/777"
|
||||
# Caller guessed .ckpt; resolution says .safetensors -> corrected.
|
||||
did = await DOWNLOAD_MANAGER.enqueue(url, "loras/guessed.ckpt")
|
||||
row = queries.get_download(did)
|
||||
assert row.model_id == "loras/guessed.safetensors"
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
|
||||
def test_manager_rejects_non_model_resolution(model_root):
|
||||
"""A URL that resolves to a non-model file is rejected, not downloaded."""
|
||||
|
||||
async def _run():
|
||||
await close_session()
|
||||
from app.model_downloader.manager import DOWNLOAD_MANAGER, DownloadError
|
||||
|
||||
runner, port = await _serve(
|
||||
_content_disposition_handler(b"not a model", "installer.zip")
|
||||
)
|
||||
try:
|
||||
url = f"http://127.0.0.1:{port}/api/download/models/999"
|
||||
with pytest.raises(DownloadError) as ei:
|
||||
await DOWNLOAD_MANAGER.enqueue(url, "loras/whatever")
|
||||
assert ei.value.code == "URL_NOT_ALLOWED"
|
||||
finally:
|
||||
await runner.cleanup()
|
||||
await close_session()
|
||||
|
||||
asyncio.run(_run())
|
||||
@@ -0,0 +1,81 @@
|
||||
"""Unit tests for the segment planner and structural safetensors validation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import struct
|
||||
|
||||
import pytest
|
||||
|
||||
from app.model_downloader.engine.planner import (
|
||||
effective_segment_count,
|
||||
plan_segments,
|
||||
)
|
||||
from app.model_downloader.verify import structural
|
||||
|
||||
|
||||
# ----- planner -----
|
||||
|
||||
|
||||
def test_plan_segments_covers_full_range_contiguously():
|
||||
total = 1000
|
||||
plans = plan_segments(total, 4)
|
||||
assert len(plans) == 4
|
||||
assert plans[0].start == 0
|
||||
assert plans[-1].end == total - 1
|
||||
# contiguous, no gaps/overlaps
|
||||
for a, b in zip(plans, plans[1:]):
|
||||
assert b.start == a.end + 1
|
||||
assert sum(p.length for p in plans) == total
|
||||
|
||||
|
||||
def test_effective_segment_count_falls_back_to_single():
|
||||
# No range support -> single
|
||||
assert effective_segment_count(10_000_000, False, 8) == 1
|
||||
# Unknown size -> single
|
||||
assert effective_segment_count(None, True, 8) == 1
|
||||
# Tiny file -> fewer segments than configured
|
||||
assert effective_segment_count(1024, True, 8) == 1
|
||||
# Large file with range support -> configured count
|
||||
assert effective_segment_count(1_000_000_000, True, 8) == 8
|
||||
|
||||
|
||||
# ----- structural -----
|
||||
|
||||
|
||||
def _make_safetensors(tensor_data_len: int, *, corrupt_size: bool = False) -> bytes:
|
||||
header = {"t": {"dtype": "F32", "shape": [tensor_data_len], "data_offsets": [0, tensor_data_len]}}
|
||||
header_bytes = json.dumps(header).encode("utf-8")
|
||||
body = b"\x00" * tensor_data_len
|
||||
if corrupt_size:
|
||||
body = body[:-1] # truncate one byte
|
||||
return struct.pack("<Q", len(header_bytes)) + header_bytes + body
|
||||
|
||||
|
||||
def test_structural_valid_safetensors(tmp_path):
|
||||
p = tmp_path / "ok.safetensors"
|
||||
p.write_bytes(_make_safetensors(256))
|
||||
structural.validate(str(p)) # no raise
|
||||
|
||||
|
||||
def test_structural_detects_truncation(tmp_path):
|
||||
p = tmp_path / "bad.safetensors"
|
||||
p.write_bytes(_make_safetensors(256, corrupt_size=True))
|
||||
with pytest.raises(structural.StructuralError):
|
||||
structural.validate(str(p))
|
||||
|
||||
|
||||
def test_structural_skips_unknown_extension(tmp_path):
|
||||
p = tmp_path / "weights.bin"
|
||||
p.write_bytes(b"anything")
|
||||
structural.validate(str(p)) # no structural check, no raise
|
||||
|
||||
|
||||
def test_structural_detects_truncation_via_name_hint(tmp_path):
|
||||
# The downloader validates the opaque temp file (a ``.part`` path) but keys
|
||||
# the format check off the final destination name via ``name_hint``, so
|
||||
# truncation must still be detected instead of silently skipped.
|
||||
p = tmp_path / "bad.comfy-download.part"
|
||||
p.write_bytes(_make_safetensors(256, corrupt_size=True))
|
||||
with pytest.raises(structural.StructuralError):
|
||||
structural.validate(str(p), name_hint="model.safetensors")
|
||||
@@ -0,0 +1,231 @@
|
||||
"""Unit tests for the security layer: allowlist, SSRF checks, path safety."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from app.model_downloader.security import allowlist, paths
|
||||
from app.model_downloader.security.ssrf import (
|
||||
SSRFError,
|
||||
check_redirect_hop,
|
||||
is_blocked_ip,
|
||||
)
|
||||
|
||||
|
||||
# ----- allowlist -----
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url,allowed",
|
||||
[
|
||||
("https://huggingface.co/org/repo/resolve/main/model.safetensors", True),
|
||||
("https://civitai.com/api/download/x/model.safetensors", True),
|
||||
("http://localhost/model.safetensors", True),
|
||||
# off-list host
|
||||
("https://evil.example.com/model.safetensors", False),
|
||||
# http to a non-loopback allowlisted host is not permitted (https only)
|
||||
("http://huggingface.co/org/repo/resolve/main/model.safetensors", False),
|
||||
# bad extension on an allowed host
|
||||
("https://huggingface.co/org/repo/resolve/main/config.json", False),
|
||||
# userinfo trick: real host is the metadata IP, not 127.0.0.1
|
||||
("http://127.0.0.1@169.254.169.254/x.safetensors", False),
|
||||
],
|
||||
)
|
||||
def test_is_url_allowed(url, allowed):
|
||||
assert allowlist.is_url_allowed(url) is allowed
|
||||
|
||||
|
||||
def test_allow_any_extension_relaxes_extension_only():
|
||||
url = "https://huggingface.co/org/repo/resolve/main/weights.bin"
|
||||
assert allowlist.is_url_allowed(url) is True # .bin is in the known set
|
||||
odd = "https://huggingface.co/org/repo/resolve/main/weights.zip"
|
||||
assert allowlist.is_url_allowed(odd) is False
|
||||
assert allowlist.is_url_allowed(odd, allow_any_extension=True) is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url,downloadable",
|
||||
[
|
||||
# known model extension in the path -> allowed
|
||||
("https://civitai.com/x/model.safetensors", True),
|
||||
# no extension in the path (Civitai download API) -> allowed, resolved later
|
||||
("https://civitai.com/api/download/models/3031464?fileId=2910346", True),
|
||||
("https://civitai.com/api/download/models/3031464", True),
|
||||
# explicit non-model extension -> rejected even on an allowed host
|
||||
("https://civitai.com/api/download/models/thing.zip", False),
|
||||
("https://huggingface.co/org/repo/resolve/main/config.json", False),
|
||||
# off-list host is never downloadable
|
||||
("https://evil.example.com/api/download/models/1", False),
|
||||
# http to a non-loopback allowlisted host is not permitted
|
||||
("http://civitai.com/api/download/models/1", False),
|
||||
],
|
||||
)
|
||||
def test_is_url_downloadable(url, downloadable):
|
||||
assert allowlist.is_url_downloadable(url) is downloadable
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"name,ext",
|
||||
[
|
||||
("model.safetensors", ".safetensors"),
|
||||
("model.SAFETENSORS", ".safetensors"),
|
||||
("archive.tar.gz", ".gz"),
|
||||
("noext", ""),
|
||||
(".safetensors", ""), # leading-dot dotfile -> no extension
|
||||
("a/b/c/model.ckpt", ".ckpt"),
|
||||
],
|
||||
)
|
||||
def test_filename_extension(name, ext):
|
||||
assert allowlist.filename_extension(name) == ext
|
||||
|
||||
|
||||
# ----- SSRF: blocked IPs -----
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"ip,blocked",
|
||||
[
|
||||
("169.254.169.254", True), # cloud metadata / link-local
|
||||
("127.0.0.1", True),
|
||||
("10.0.0.5", True),
|
||||
("192.168.1.1", True),
|
||||
("172.16.0.1", True),
|
||||
("::1", True),
|
||||
("0.0.0.0", True),
|
||||
# IPv4-mapped IPv6: must see through the mapping even on CPython
|
||||
# versions predating the gh-113171 is_* property fix.
|
||||
("::ffff:169.254.169.254", True), # mapped cloud metadata
|
||||
("::ffff:127.0.0.1", True), # mapped loopback
|
||||
("::ffff:10.0.0.1", True), # mapped RFC1918
|
||||
("::ffff:8.8.8.8", False), # mapped public address stays allowed
|
||||
("8.8.8.8", False),
|
||||
("1.1.1.1", False),
|
||||
("not-an-ip", True), # unparseable -> refuse
|
||||
],
|
||||
)
|
||||
def test_is_blocked_ip(ip, blocked):
|
||||
assert is_blocked_ip(ip) is blocked
|
||||
|
||||
|
||||
# ----- SSRF: redirect hop validation -----
|
||||
|
||||
|
||||
def test_check_redirect_hop_rejects_bad_scheme_and_userinfo():
|
||||
with pytest.raises(SSRFError):
|
||||
check_redirect_hop("ftp://huggingface.co/x.safetensors")
|
||||
with pytest.raises(SSRFError):
|
||||
check_redirect_hop("https://user:pass@cdn.example.com/x")
|
||||
# A CDN host that is NOT on the allowlist is allowed as a redirect target
|
||||
# (private-IP protection is the resolver's job; credential leak is prevented
|
||||
# by exact host matching).
|
||||
assert check_redirect_hop("https://cdn-lfs.huggingface.co/abc") is not None
|
||||
|
||||
|
||||
def test_check_redirect_hop_http_only_for_loopback():
|
||||
# Plain http to an external host is rejected (no plaintext downgrade).
|
||||
with pytest.raises(SSRFError):
|
||||
check_redirect_hop("http://cdn-lfs.huggingface.co/abc")
|
||||
# http is honored for loopback only on the initial user-supplied URL (the
|
||||
# "download a local model" feature).
|
||||
assert (
|
||||
check_redirect_hop("http://localhost/x.safetensors", is_initial_url=True)
|
||||
is not None
|
||||
)
|
||||
assert (
|
||||
check_redirect_hop("http://127.0.0.1/x.safetensors", is_initial_url=True)
|
||||
is not None
|
||||
)
|
||||
|
||||
|
||||
def test_check_redirect_hop_blocks_loopback_and_ip_literals_on_redirect():
|
||||
# A redirect (is_initial_url=False, the default) must never reach loopback,
|
||||
# whether by hostname or by IP literal, nor any other internal IP literal.
|
||||
for target in (
|
||||
"http://localhost/x.safetensors",
|
||||
"http://127.0.0.1/x.safetensors",
|
||||
"https://[::1]/x.safetensors",
|
||||
"https://169.254.169.254/x.safetensors", # cloud metadata
|
||||
"https://10.0.0.5/x.safetensors", # RFC1918
|
||||
):
|
||||
with pytest.raises(SSRFError):
|
||||
check_redirect_hop(target)
|
||||
# Off-allowlist public CDN hosts (hostnames) remain valid redirect targets;
|
||||
# their resolved IPs are screened by the connector's resolver.
|
||||
assert check_redirect_hop("https://cdn-lfs.huggingface.co/abc") is not None
|
||||
|
||||
|
||||
# ----- path safety -----
|
||||
|
||||
|
||||
def test_parse_model_id_valid(model_root):
|
||||
directory, filename = paths.parse_model_id("loras/my_lora.safetensors")
|
||||
assert directory == "loras"
|
||||
assert filename == "my_lora.safetensors"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_id",
|
||||
[
|
||||
"loras/../etc/passwd.safetensors", # traversal
|
||||
"loras/sub/dir.safetensors", # nested
|
||||
"unknownfolder/x.safetensors", # unknown folder
|
||||
"loras/model.txt", # bad extension
|
||||
"noslash.safetensors", # missing directory
|
||||
"loras/", # empty filename
|
||||
],
|
||||
)
|
||||
def test_parse_model_id_rejects(model_root, model_id):
|
||||
with pytest.raises(paths.InvalidModelId):
|
||||
paths.parse_model_id(model_id)
|
||||
|
||||
|
||||
def test_resolve_destination_stays_in_root(model_root):
|
||||
final_path, temp_path = paths.resolve_destination("loras/x.safetensors")
|
||||
assert final_path.startswith(model_root)
|
||||
assert temp_path.startswith(model_root)
|
||||
assert temp_path != final_path
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_id,ext,expected",
|
||||
[
|
||||
# no extension -> append the resolved one
|
||||
("loras/my_civitai_model", ".safetensors", "loras/my_civitai_model.safetensors"),
|
||||
# different known extension -> replace it
|
||||
("loras/mymodel.ckpt", ".safetensors", "loras/mymodel.safetensors"),
|
||||
# same extension -> unchanged
|
||||
("loras/mymodel.safetensors", ".safetensors", "loras/mymodel.safetensors"),
|
||||
# non-model suffix is treated as a stem, extension appended
|
||||
("loras/my.model.v2", ".safetensors", "loras/my.model.v2.safetensors"),
|
||||
# malformed (no slash) is returned untouched for parse_model_id to reject
|
||||
("noslash", ".safetensors", "noslash"),
|
||||
],
|
||||
)
|
||||
def test_apply_extension(model_id, ext, expected):
|
||||
assert paths.apply_extension(model_id, ext) == expected
|
||||
|
||||
|
||||
# ----- Content-Disposition filename parsing -----
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"header,expected",
|
||||
[
|
||||
('attachment; filename="model.safetensors"', "model.safetensors"),
|
||||
("attachment; filename=model.ckpt", "model.ckpt"),
|
||||
# RFC 5987 form is preferred and percent-decoded
|
||||
(
|
||||
"attachment; filename=\"fallback.bin\"; filename*=UTF-8''my%20model.safetensors",
|
||||
"my model.safetensors",
|
||||
),
|
||||
# directory components in a hostile header are stripped to the basename
|
||||
('attachment; filename="../../etc/passwd"', "passwd"),
|
||||
('attachment; filename="a\\\\b\\\\model.pt"', "model.pt"),
|
||||
("inline", None),
|
||||
(None, None),
|
||||
],
|
||||
)
|
||||
def test_filename_from_content_disposition(header, expected):
|
||||
from app.model_downloader.net.http import filename_from_content_disposition
|
||||
|
||||
assert filename_from_content_disposition(header) == expected
|
||||
@@ -818,30 +818,6 @@ class TestExecution:
|
||||
except urllib.error.HTTPError:
|
||||
pass # Expected behavior
|
||||
|
||||
def test_cached_outputs_in_job_without_client_id(self, client: ComfyClient, builder: GraphBuilder):
|
||||
g = builder
|
||||
image = g.node("StubImage", content="BLACK", height=32, width=32, batch_size=1)
|
||||
output = g.node("SaveImage", images=image.out(0))
|
||||
|
||||
# Prime the cache with a normal run.
|
||||
client.run(g)
|
||||
|
||||
# Resubmit anonymously (no client_id) so output nodes are cache hits with no websocket client.
|
||||
data = json.dumps({"prompt": g.finalize()}).encode('utf-8')
|
||||
req = urllib.request.Request(f"http://{client.server_address}/prompt", data=data)
|
||||
prompt_id = json.loads(urllib.request.urlopen(req).read())['prompt_id']
|
||||
|
||||
for _ in range(100):
|
||||
job = client.get_job(prompt_id)
|
||||
if job is not None and job['status'] not in ('pending', 'in_progress'):
|
||||
break
|
||||
time.sleep(0.1)
|
||||
else:
|
||||
raise AssertionError("Prompt did not complete in time")
|
||||
|
||||
assert job['status'] == 'completed'
|
||||
assert output.id in job['outputs'], "Cached outputs must appear in job outputs without a client_id"
|
||||
|
||||
def _create_history_item(self, client, builder):
|
||||
g = GraphBuilder(prefix="offset_test")
|
||||
input_node = g.node(
|
||||
|
||||
Reference in New Issue
Block a user