mirror of
https://github.com/comfyanonymous/ComfyUI.git
synced 2026-07-22 07:47:36 +08:00
Compare commits
36
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
67b8d23c2f | ||
|
|
0aecac867d | ||
|
|
80acfcf0fe | ||
|
|
c35a622acd | ||
|
|
da2608926e | ||
|
|
5658a68a87 | ||
|
|
5bb831a3f5 | ||
|
|
5697b97017 | ||
|
|
ec0e8b3447 | ||
|
|
8deaa4d911 | ||
|
|
b58f829b57 | ||
|
|
917faef771 | ||
|
|
8b099de36a | ||
|
|
69ea58697b | ||
|
|
f3a36e7484 | ||
|
|
92ddf07ba1 | ||
|
|
1f51e146a8 | ||
|
|
5976ee37cd | ||
|
|
328144ce24 | ||
|
|
8310b0e0db | ||
|
|
94fa08223e | ||
|
|
1377a2f729 | ||
|
|
206b9245dc | ||
|
|
89ecc5cf8c | ||
|
|
8e2e54e2b8 | ||
|
|
e2a6e30d89 | ||
|
|
099522f85b | ||
|
|
62e025a4f3 | ||
|
|
b7a648ca20 | ||
|
|
73e84d5ec8 | ||
|
|
1ea724339c | ||
|
|
412aaab0e2 | ||
|
|
04a30fb375 | ||
|
|
b35819712e | ||
|
|
d0008a8958 | ||
|
|
55a15f87ce |
@@ -32,9 +32,11 @@ jobs:
|
||||
PR_NUMBER: ${{ github.event.pull_request.number || github.event.issue.number }}
|
||||
PR_AUTHOR: ${{ github.event.pull_request.user.login || github.event.issue.user.login }}
|
||||
BASE_ALLOWLIST: action@github.com,actions-user,ampagent,claude,comfy-pr-bot,GitHub Action,github-actions,github-actions[bot],Glary Bot,Glary-Bot,*[bot]
|
||||
# For each commit emit the GitHub login when the author/committer email resolves to a GitHub account
|
||||
# otherwise fall back to the raw git name.
|
||||
run: |
|
||||
others=$(gh api "repos/${{ github.repository }}/pulls/${PR_NUMBER}/commits" --paginate \
|
||||
--jq '.[] | (.author.login // empty), (.committer.login // empty)' \
|
||||
--jq '.[] | (.author.login // .commit.author.name // empty), (.committer.login // .commit.committer.name // empty)' \
|
||||
| sort -u | grep -vix "${PR_AUTHOR}" | paste -sd, -)
|
||||
if [ -n "$others" ]; then
|
||||
echo "allowlist=${BASE_ALLOWLIST},${others}" >> "$GITHUB_OUTPUT"
|
||||
@@ -43,7 +45,7 @@ jobs:
|
||||
fi
|
||||
|
||||
- name: CLA Assistant
|
||||
# Run on PR events, on "recheck" comment, or when someone posts the exact signing phrase.
|
||||
# Run on PR events, on "recheck" comment, or when someone posts the signing phrase.
|
||||
# IMPORTANT: this phrase must match `custom-pr-sign-comment` below.
|
||||
if: >
|
||||
github.event_name == 'pull_request_target' ||
|
||||
|
||||
@@ -0,0 +1,107 @@
|
||||
"""
|
||||
Allow case-sensitive tag names.
|
||||
|
||||
Revision ID: 0005_allow_case_sensitive_tags
|
||||
Revises: 0004_drop_tag_type
|
||||
Create Date: 2026-06-16
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "0005_allow_case_sensitive_tags"
|
||||
down_revision = "0004_drop_tag_type"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
if bind.dialect.name == "sqlite":
|
||||
# SQLite cannot ALTER/DROP CHECK constraints. Recreate the small tag
|
||||
# vocabulary table without the lowercase constraint while preserving
|
||||
# existing tag names.
|
||||
op.execute("PRAGMA foreign_keys=OFF")
|
||||
try:
|
||||
op.execute(
|
||||
"CREATE TABLE tags_new ("
|
||||
"name VARCHAR(512) NOT NULL, "
|
||||
"CONSTRAINT pk_tags PRIMARY KEY (name)"
|
||||
")"
|
||||
)
|
||||
op.execute("INSERT INTO tags_new(name) SELECT name FROM tags")
|
||||
op.execute("DROP TABLE tags")
|
||||
op.execute("ALTER TABLE tags_new RENAME TO tags")
|
||||
finally:
|
||||
op.execute("PRAGMA foreign_keys=ON")
|
||||
return
|
||||
|
||||
op.drop_constraint("ck_tags_ck_tags_lowercase", "tags", type_="check")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Existing mixed-case tags cannot satisfy the old constraint. Lowercase them
|
||||
# before restoring it, merging duplicate vocabulary/link rows that collide.
|
||||
bind = op.get_bind()
|
||||
|
||||
tag_names = [row[0] for row in bind.execute(sa.text("SELECT name FROM tags"))]
|
||||
existing_names = set(tag_names)
|
||||
lowercase_names = sorted({name.lower() for name in tag_names})
|
||||
missing_lowercase_rows = [
|
||||
{"name": name} for name in lowercase_names if name not in existing_names
|
||||
]
|
||||
if missing_lowercase_rows:
|
||||
bind.execute(sa.text("INSERT INTO tags(name) VALUES (:name)"), missing_lowercase_rows)
|
||||
|
||||
link_rows = bind.execute(
|
||||
sa.text(
|
||||
"SELECT asset_reference_id, tag_name, origin, added_at "
|
||||
"FROM asset_reference_tags "
|
||||
"ORDER BY asset_reference_id, tag_name"
|
||||
)
|
||||
).mappings()
|
||||
deduped_links = {}
|
||||
for row in link_rows:
|
||||
key = (row["asset_reference_id"], row["tag_name"].lower())
|
||||
deduped_links.setdefault(
|
||||
key,
|
||||
{
|
||||
"asset_reference_id": row["asset_reference_id"],
|
||||
"tag_name": row["tag_name"].lower(),
|
||||
"origin": row["origin"],
|
||||
"added_at": row["added_at"],
|
||||
},
|
||||
)
|
||||
|
||||
op.execute("DELETE FROM asset_reference_tags")
|
||||
if deduped_links:
|
||||
bind.execute(
|
||||
sa.text(
|
||||
"INSERT INTO asset_reference_tags "
|
||||
"(asset_reference_id, tag_name, origin, added_at) "
|
||||
"VALUES (:asset_reference_id, :tag_name, :origin, :added_at)"
|
||||
),
|
||||
list(deduped_links.values()),
|
||||
)
|
||||
op.execute("DELETE FROM tags WHERE name != lower(name)")
|
||||
|
||||
if bind.dialect.name == "sqlite":
|
||||
op.execute("PRAGMA foreign_keys=OFF")
|
||||
try:
|
||||
op.execute(
|
||||
"CREATE TABLE tags_new ("
|
||||
"name VARCHAR(512) NOT NULL, "
|
||||
"CONSTRAINT pk_tags PRIMARY KEY (name), "
|
||||
"CONSTRAINT ck_tags_lowercase CHECK (name = lower(name))"
|
||||
")"
|
||||
)
|
||||
op.execute("INSERT INTO tags_new(name) SELECT name FROM tags")
|
||||
op.execute("DROP TABLE tags")
|
||||
op.execute("ALTER TABLE tags_new RENAME TO tags")
|
||||
finally:
|
||||
op.execute("PRAGMA foreign_keys=ON")
|
||||
return
|
||||
|
||||
op.create_check_constraint(
|
||||
"ck_tags_ck_tags_lowercase", "tags", "name = lower(name)"
|
||||
)
|
||||
@@ -0,0 +1,30 @@
|
||||
"""
|
||||
Add loader_path column to asset_references.
|
||||
|
||||
Stores the in-root loader path (path relative to the storage root with the
|
||||
top-level model category dropped) derived from file_path at scan/ingest time,
|
||||
so the assets API can return it without re-resolving against every registered
|
||||
model-folder base on every request.
|
||||
|
||||
Revision ID: 0006_add_loader_path
|
||||
Revises: 0005_allow_case_sensitive_tags
|
||||
Create Date: 2026-07-02
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "0006_add_loader_path"
|
||||
down_revision = "0005_allow_case_sensitive_tags"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
with op.batch_alter_table("asset_references") as batch_op:
|
||||
batch_op.add_column(sa.Column("loader_path", sa.Text(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
with op.batch_alter_table("asset_references") as batch_op:
|
||||
batch_op.drop_column("loader_path")
|
||||
+10
-12
@@ -40,6 +40,7 @@ from app.assets.services import (
|
||||
upload_from_temp_path,
|
||||
)
|
||||
from app.assets.services.cursor import InvalidCursorError
|
||||
from app.assets.services.path_utils import compute_display_name
|
||||
from app.assets.services.tagging import list_tag_histogram
|
||||
|
||||
ROUTES = web.RouteTableDef()
|
||||
@@ -161,11 +162,19 @@ def _build_asset_response(result: schemas.AssetDetailResult | schemas.UploadResu
|
||||
preview_url = None
|
||||
else:
|
||||
preview_url = _build_preview_url_from_view(result.tags, result.ref.user_metadata)
|
||||
if result.ref.file_path:
|
||||
display_name = compute_display_name(result.ref.file_path)
|
||||
# In-root loader path (model category dropped): what model loaders consume.
|
||||
loader_path = result.ref.loader_path
|
||||
else:
|
||||
display_name, loader_path = None, None
|
||||
asset_content_hash = result.asset.hash if result.asset else None
|
||||
return schemas_out.Asset(
|
||||
id=result.ref.id,
|
||||
name=result.ref.name,
|
||||
hash=asset_content_hash,
|
||||
loader_path=loader_path,
|
||||
display_name=display_name,
|
||||
asset_hash=asset_content_hash,
|
||||
size=int(result.asset.size_bytes) if result.asset else None,
|
||||
mime_type=result.asset.mime_type if result.asset else None,
|
||||
@@ -419,17 +428,6 @@ async def upload_asset(request: web.Request) -> web.Response:
|
||||
400, "INVALID_BODY", f"Validation failed: {ve.json()}"
|
||||
)
|
||||
|
||||
if spec.tags and spec.tags[0] == "models":
|
||||
if (
|
||||
len(spec.tags) < 2
|
||||
or spec.tags[1] not in folder_paths.folder_names_and_paths
|
||||
):
|
||||
delete_temp_file_if_exists(parsed.tmp_path)
|
||||
category = spec.tags[1] if len(spec.tags) >= 2 else ""
|
||||
return _build_error_response(
|
||||
400, "INVALID_BODY", f"unknown models category '{category}'"
|
||||
)
|
||||
|
||||
try:
|
||||
# Fast path: hash exists, create AssetReference without writing anything
|
||||
if spec.hash and parsed.provided_hash_exists is True:
|
||||
@@ -473,7 +471,7 @@ async def upload_asset(request: web.Request) -> web.Response:
|
||||
return _build_error_response(400, e.code, str(e))
|
||||
except ValueError as e:
|
||||
delete_temp_file_if_exists(parsed.tmp_path)
|
||||
return _build_error_response(400, "BAD_REQUEST", str(e))
|
||||
return _build_error_response(400, "INVALID_BODY", str(e))
|
||||
except HashMismatchError as e:
|
||||
delete_temp_file_if_exists(parsed.tmp_path)
|
||||
return _build_error_response(400, "HASH_MISMATCH", str(e))
|
||||
|
||||
@@ -140,7 +140,7 @@ class CreateFromHashBody(BaseModel):
|
||||
if v is None:
|
||||
return []
|
||||
if isinstance(v, list):
|
||||
out = [str(t).strip().lower() for t in v if str(t).strip()]
|
||||
out = [str(t).strip() for t in v if str(t).strip()]
|
||||
seen = set()
|
||||
dedup = []
|
||||
for t in out:
|
||||
@@ -149,7 +149,7 @@ class CreateFromHashBody(BaseModel):
|
||||
dedup.append(t)
|
||||
return dedup
|
||||
if isinstance(v, str):
|
||||
return [t.strip().lower() for t in v.split(",") if t.strip()]
|
||||
return list(dict.fromkeys(t.strip() for t in v.split(",") if t.strip()))
|
||||
return []
|
||||
|
||||
|
||||
@@ -206,7 +206,7 @@ class TagsListQuery(BaseModel):
|
||||
if v is None:
|
||||
return v
|
||||
v = v.strip()
|
||||
return v.lower() or None
|
||||
return v or None
|
||||
|
||||
|
||||
class TagsAdd(BaseModel):
|
||||
@@ -220,7 +220,7 @@ class TagsAdd(BaseModel):
|
||||
for t in v:
|
||||
if not isinstance(t, str):
|
||||
raise TypeError("tags must be strings")
|
||||
tnorm = t.strip().lower()
|
||||
tnorm = t.strip()
|
||||
if tnorm:
|
||||
out.append(tnorm)
|
||||
seen = set()
|
||||
@@ -239,8 +239,8 @@ class TagsRemove(TagsAdd):
|
||||
class UploadAssetSpec(BaseModel):
|
||||
"""Upload Asset operation.
|
||||
|
||||
- tags: optional list; if provided, first is root ('models'|'input'|'output');
|
||||
if root == 'models', second must be a valid category
|
||||
- tags: labels plus one destination role ('models'|'input'|'output') for new bytes;
|
||||
if role == 'models', exactly one model_type:<folder_name> tag is required
|
||||
- name: display name
|
||||
- user_metadata: arbitrary JSON object (optional)
|
||||
- hash: optional canonical 'blake3:<hex>' for validation / fast-path
|
||||
@@ -309,7 +309,7 @@ class UploadAssetSpec(BaseModel):
|
||||
norm = []
|
||||
seen = set()
|
||||
for t in items:
|
||||
tnorm = str(t).strip().lower()
|
||||
tnorm = str(t).strip()
|
||||
if tnorm and tnorm not in seen:
|
||||
seen.add(tnorm)
|
||||
norm.append(tnorm)
|
||||
@@ -335,14 +335,4 @@ class UploadAssetSpec(BaseModel):
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_order(self):
|
||||
if not self.tags:
|
||||
raise ValueError("at least one tag is required for uploads")
|
||||
root = self.tags[0]
|
||||
if root not in {"models", "input", "output"}:
|
||||
raise ValueError("first tag must be one of: models, input, output")
|
||||
if root == "models":
|
||||
if len(self.tags) < 2:
|
||||
raise ValueError(
|
||||
"models uploads require a category tag as the second tag"
|
||||
)
|
||||
return self
|
||||
|
||||
@@ -9,8 +9,20 @@ class Asset(BaseModel):
|
||||
``id`` here is the AssetReference id, not the content-addressed Asset id."""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
name: str = Field(
|
||||
...,
|
||||
deprecated=True,
|
||||
description="Reference label, often caller-provided or derived from the filename. Deprecated for storage path/display semantics; use `loader_path` and `display_name` when present.",
|
||||
)
|
||||
hash: str | None = None
|
||||
loader_path: str | None = Field(
|
||||
default=None,
|
||||
description="The value a loader consumes to load this asset. `None` when no loader can resolve the file.",
|
||||
)
|
||||
display_name: str | None = Field(
|
||||
default=None,
|
||||
description="Human-facing label for the asset. Not unique.",
|
||||
)
|
||||
asset_hash: str | None = None
|
||||
size: int | None = None
|
||||
mime_type: str | None = None
|
||||
|
||||
@@ -140,7 +140,6 @@ async def parse_multipart_upload(
|
||||
provided_mime_type = ((await field.text()) or "").strip() or None
|
||||
elif fname == "preview_id":
|
||||
provided_preview_id = ((await field.text()) or "").strip() or None
|
||||
|
||||
if not file_present and not (provided_hash and provided_hash_exists):
|
||||
raise UploadError(
|
||||
400, "MISSING_FILE", "Form must include a 'file' part or a known 'hash'."
|
||||
|
||||
@@ -76,6 +76,8 @@ class AssetReference(Base):
|
||||
|
||||
# Cache state fields (from former AssetCacheState)
|
||||
file_path: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
# In-root loader path derived from file_path at scan/ingest time.
|
||||
loader_path: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
mtime_ns: Mapped[int | None] = mapped_column(BigInteger, nullable=True)
|
||||
needs_verify: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
is_missing: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||
|
||||
@@ -650,6 +650,7 @@ def upsert_reference(
|
||||
name: str,
|
||||
mtime_ns: int,
|
||||
owner_id: str = "",
|
||||
loader_path: str | None = None,
|
||||
) -> tuple[bool, bool]:
|
||||
"""Upsert a reference by file_path. Returns (created, updated).
|
||||
|
||||
@@ -659,6 +660,7 @@ def upsert_reference(
|
||||
vals = {
|
||||
"asset_id": asset_id,
|
||||
"file_path": file_path,
|
||||
"loader_path": loader_path,
|
||||
"name": name,
|
||||
"owner_id": owner_id,
|
||||
"mtime_ns": int(mtime_ns),
|
||||
@@ -686,13 +688,14 @@ def upsert_reference(
|
||||
AssetReference.asset_id != asset_id,
|
||||
AssetReference.mtime_ns.is_(None),
|
||||
AssetReference.mtime_ns != int(mtime_ns),
|
||||
AssetReference.loader_path.is_distinct_from(loader_path),
|
||||
AssetReference.is_missing == True, # noqa: E712
|
||||
AssetReference.deleted_at.isnot(None),
|
||||
)
|
||||
)
|
||||
.values(
|
||||
asset_id=asset_id, mtime_ns=int(mtime_ns), is_missing=False,
|
||||
deleted_at=None, updated_at=now,
|
||||
asset_id=asset_id, mtime_ns=int(mtime_ns), loader_path=loader_path,
|
||||
is_missing=False, deleted_at=None, updated_at=now,
|
||||
)
|
||||
)
|
||||
res2 = session.execute(upd)
|
||||
|
||||
@@ -265,6 +265,8 @@ def list_tags_with_usage(
|
||||
order: str = "count_desc",
|
||||
owner_id: str = "",
|
||||
) -> tuple[list[tuple[str, str, int]], int]:
|
||||
prefix_filter = prefix.strip() if prefix else ""
|
||||
|
||||
counts_sq = (
|
||||
select(
|
||||
AssetReferenceTag.tag_name.label("tag_name"),
|
||||
@@ -293,9 +295,8 @@ def list_tags_with_usage(
|
||||
.join(counts_sq, counts_sq.c.tag_name == Tag.name, isouter=True)
|
||||
)
|
||||
|
||||
if prefix:
|
||||
escaped, esc = escape_sql_like_string(prefix.strip().lower())
|
||||
q = q.where(Tag.name.like(escaped + "%", escape=esc))
|
||||
if prefix_filter:
|
||||
q = q.where(func.substr(Tag.name, 1, len(prefix_filter)) == prefix_filter)
|
||||
|
||||
if not include_zero:
|
||||
q = q.where(func.coalesce(counts_sq.c.cnt, 0) > 0)
|
||||
@@ -306,9 +307,8 @@ def list_tags_with_usage(
|
||||
q = q.order_by(func.coalesce(counts_sq.c.cnt, 0).desc(), Tag.name.asc())
|
||||
|
||||
total_q = select(func.count()).select_from(Tag)
|
||||
if prefix:
|
||||
escaped, esc = escape_sql_like_string(prefix.strip().lower())
|
||||
total_q = total_q.where(Tag.name.like(escaped + "%", escape=esc))
|
||||
if prefix_filter:
|
||||
total_q = total_q.where(func.substr(Tag.name, 1, len(prefix_filter)) == prefix_filter)
|
||||
if not include_zero:
|
||||
visible_tags_sq = (
|
||||
select(AssetReferenceTag.tag_name)
|
||||
|
||||
@@ -41,10 +41,10 @@ def get_utc_now() -> datetime:
|
||||
def normalize_tags(tags: list[str] | None) -> list[str]:
|
||||
"""
|
||||
Normalize a list of tags by:
|
||||
- Stripping whitespace and converting to lowercase.
|
||||
- Removing duplicates.
|
||||
- Stripping whitespace.
|
||||
- Removing exact duplicates while preserving order and case.
|
||||
"""
|
||||
return list(dict.fromkeys(t.strip().lower() for t in (tags or []) if (t or "").strip()))
|
||||
return list(dict.fromkeys(t.strip() for t in (tags or []) if (t or "").strip()))
|
||||
|
||||
|
||||
def validate_blake3_hash(s: str) -> str:
|
||||
|
||||
@@ -36,7 +36,7 @@ from app.assets.services.hashing import HashCheckpoint, compute_blake3_hash
|
||||
from app.assets.services.image_dimensions import extract_image_dimensions
|
||||
from app.assets.services.metadata_extract import extract_file_metadata
|
||||
from app.assets.services.path_utils import (
|
||||
compute_relative_filename,
|
||||
compute_loader_path,
|
||||
get_comfy_models_folders,
|
||||
get_name_and_tags_from_asset_path,
|
||||
)
|
||||
@@ -63,7 +63,7 @@ RootType = Literal["models", "input", "output"]
|
||||
def get_prefixes_for_root(root: RootType) -> list[str]:
|
||||
if root == "models":
|
||||
bases: list[str] = []
|
||||
for _bucket, paths in get_comfy_models_folders():
|
||||
for _bucket, paths, _exts in get_comfy_models_folders():
|
||||
bases.extend(paths)
|
||||
return [os.path.abspath(p) for p in bases]
|
||||
if root == "input":
|
||||
@@ -81,7 +81,7 @@ def get_all_known_prefixes() -> list[str]:
|
||||
|
||||
def collect_models_files() -> list[str]:
|
||||
out: list[str] = []
|
||||
for folder_name, bases in get_comfy_models_folders():
|
||||
for folder_name, bases, _exts in get_comfy_models_folders():
|
||||
rel_files = folder_paths.get_filename_list(folder_name) or []
|
||||
for rel_path in rel_files:
|
||||
if not all(is_visible(part) for part in Path(rel_path).parts):
|
||||
@@ -308,7 +308,7 @@ def build_asset_specs(
|
||||
if not stat_p.st_size:
|
||||
continue
|
||||
name, tags = get_name_and_tags_from_asset_path(abs_p)
|
||||
rel_fname = compute_relative_filename(abs_p)
|
||||
rel_fname = compute_loader_path(abs_p)
|
||||
|
||||
# Extract metadata (tier 1: filesystem, tier 2: safetensors header)
|
||||
metadata = None
|
||||
@@ -430,7 +430,7 @@ def enrich_asset(
|
||||
return new_level
|
||||
|
||||
initial_mtime_ns = get_mtime_ns(stat_p)
|
||||
rel_fname = compute_relative_filename(file_path)
|
||||
rel_fname = compute_loader_path(file_path)
|
||||
mime_type: str | None = None
|
||||
metadata = None
|
||||
|
||||
|
||||
@@ -38,7 +38,7 @@ from app.assets.database.queries import (
|
||||
update_reference_updated_at,
|
||||
)
|
||||
from app.assets.helpers import select_best_live_path
|
||||
from app.assets.services.path_utils import compute_relative_filename
|
||||
from app.assets.services.path_utils import compute_loader_path
|
||||
from app.assets.services.schemas import (
|
||||
AssetData,
|
||||
AssetDetailResult,
|
||||
@@ -91,7 +91,7 @@ def update_asset_metadata(
|
||||
update_reference_name(session, reference_id=reference_id, name=name)
|
||||
touched = True
|
||||
|
||||
computed_filename = compute_relative_filename(ref.file_path) if ref.file_path else None
|
||||
computed_filename = compute_loader_path(ref.file_path) if ref.file_path else None
|
||||
|
||||
new_meta: dict | None = None
|
||||
if user_metadata is not None:
|
||||
|
||||
@@ -56,6 +56,7 @@ class ReferenceRow(TypedDict):
|
||||
id: str
|
||||
asset_id: str
|
||||
file_path: str
|
||||
loader_path: str | None
|
||||
mtime_ns: int
|
||||
owner_id: str
|
||||
name: str
|
||||
@@ -134,6 +135,14 @@ def batch_insert_seed_assets(
|
||||
|
||||
for spec in specs:
|
||||
absolute_path = os.path.abspath(spec["abs_path"])
|
||||
existing_asset_id = path_to_asset_id.get(absolute_path)
|
||||
if existing_asset_id is not None:
|
||||
existing_tags = asset_id_to_ref_data[existing_asset_id]["tags"]
|
||||
asset_id_to_ref_data[existing_asset_id]["tags"] = list(
|
||||
dict.fromkeys([*existing_tags, *spec["tags"]])
|
||||
)
|
||||
continue
|
||||
|
||||
asset_id = str(uuid.uuid4())
|
||||
reference_id = str(uuid.uuid4())
|
||||
absolute_path_list.append(absolute_path)
|
||||
@@ -164,6 +173,8 @@ def batch_insert_seed_assets(
|
||||
"id": reference_id,
|
||||
"asset_id": asset_id,
|
||||
"file_path": absolute_path,
|
||||
# spec["fname"] is compute_loader_path(abs_path) from build_asset_specs.
|
||||
"loader_path": spec["fname"],
|
||||
"mtime_ns": spec["mtime_ns"],
|
||||
"owner_id": owner_id,
|
||||
"name": spec["info_name"],
|
||||
|
||||
@@ -33,8 +33,9 @@ from app.assets.services.bulk_ingest import batch_insert_seed_assets
|
||||
from app.assets.services.file_utils import get_size_and_mtime_ns
|
||||
from app.assets.services.image_dimensions import extract_image_dimensions
|
||||
from app.assets.services.path_utils import (
|
||||
compute_relative_filename,
|
||||
compute_loader_path,
|
||||
get_name_and_tags_from_asset_path,
|
||||
get_path_derived_tags_from_path,
|
||||
resolve_destination_from_tags,
|
||||
validate_path_within_base,
|
||||
)
|
||||
@@ -91,6 +92,7 @@ def _ingest_file_from_path(
|
||||
name=info_name or os.path.basename(locator),
|
||||
mtime_ns=mtime_ns,
|
||||
owner_id=owner_id,
|
||||
loader_path=compute_loader_path(locator),
|
||||
)
|
||||
|
||||
# Get the reference we just created/updated
|
||||
@@ -101,17 +103,32 @@ def _ingest_file_from_path(
|
||||
if preview_id and ref.preview_id != preview_id:
|
||||
ref.preview_id = preview_id
|
||||
|
||||
norm = normalize_tags(list(tags))
|
||||
if norm:
|
||||
try:
|
||||
backend_tags = get_path_derived_tags_from_path(locator)
|
||||
except ValueError:
|
||||
backend_tags = []
|
||||
caller_tags = normalize_tags(tags)
|
||||
backend_tags = normalize_tags(backend_tags)
|
||||
all_tags = normalize_tags([*caller_tags, *backend_tags])
|
||||
if all_tags:
|
||||
if require_existing_tags:
|
||||
validate_tags_exist(session, norm)
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=reference_id,
|
||||
tags=norm,
|
||||
origin=tag_origin,
|
||||
create_if_missing=not require_existing_tags,
|
||||
)
|
||||
validate_tags_exist(session, all_tags)
|
||||
if backend_tags:
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=reference_id,
|
||||
tags=backend_tags,
|
||||
origin="automatic",
|
||||
create_if_missing=not require_existing_tags,
|
||||
)
|
||||
if caller_tags:
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=reference_id,
|
||||
tags=caller_tags,
|
||||
origin=tag_origin,
|
||||
create_if_missing=not require_existing_tags,
|
||||
)
|
||||
|
||||
_update_metadata_with_filename(
|
||||
session,
|
||||
@@ -228,7 +245,7 @@ def ingest_existing_file(
|
||||
"mtime_ns": mtime_ns,
|
||||
"info_name": name,
|
||||
"tags": tags,
|
||||
"fname": os.path.basename(abs_path),
|
||||
"fname": compute_loader_path(abs_path),
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": mime_type,
|
||||
@@ -288,7 +305,7 @@ def _register_existing_asset(
|
||||
return result
|
||||
|
||||
new_meta = dict(user_metadata)
|
||||
computed_filename = compute_relative_filename(ref.file_path) if ref.file_path else None
|
||||
computed_filename = compute_loader_path(ref.file_path) if ref.file_path else None
|
||||
if computed_filename:
|
||||
new_meta["filename"] = computed_filename
|
||||
|
||||
@@ -335,7 +352,7 @@ def _update_metadata_with_filename(
|
||||
current_metadata: dict | None,
|
||||
user_metadata: dict[str, Any],
|
||||
) -> None:
|
||||
computed_filename = compute_relative_filename(file_path) if file_path else None
|
||||
computed_filename = compute_loader_path(file_path) if file_path else None
|
||||
|
||||
current_meta = current_metadata or {}
|
||||
new_meta = dict(current_meta)
|
||||
@@ -474,6 +491,10 @@ def upload_from_temp_path(
|
||||
existing = get_asset_by_hash(session, asset_hash=asset_hash)
|
||||
|
||||
if existing is not None:
|
||||
# Once content is already known, duplicate byte uploads are treated as
|
||||
# reference-only creation. Request tags are labels only here: do not
|
||||
# require upload destination tags, do not move bytes, and do not
|
||||
# synthesize path-derived classification or uploaded provenance.
|
||||
with contextlib.suppress(Exception):
|
||||
if temp_path and os.path.exists(temp_path):
|
||||
os.remove(temp_path)
|
||||
@@ -535,7 +556,7 @@ def upload_from_temp_path(
|
||||
owner_id=owner_id,
|
||||
preview_id=preview_id,
|
||||
user_metadata=user_metadata or {},
|
||||
tags=tags,
|
||||
tags=[*(tags or []), "uploaded"],
|
||||
tag_origin="manual",
|
||||
require_existing_tags=False,
|
||||
)
|
||||
@@ -569,15 +590,19 @@ def register_file_in_place(
|
||||
) -> UploadResult:
|
||||
"""Register an already-saved file in the asset database without moving it.
|
||||
|
||||
Tags are derived from the filesystem path (root category + subfolder names),
|
||||
merged with any caller-provided tags, matching the behavior of the scanner.
|
||||
This helper is used by upload paths that have already written bytes before
|
||||
registering the file, so it records the same ``uploaded`` tag as the
|
||||
multipart byte-upload path.
|
||||
|
||||
Tags are derived from trusted filesystem classification and merged with any
|
||||
caller-provided tags, matching the behavior of the scanner.
|
||||
If the path is not under a known root, only the caller-provided tags are used.
|
||||
"""
|
||||
try:
|
||||
_, path_tags = get_name_and_tags_from_asset_path(abs_path)
|
||||
except ValueError:
|
||||
path_tags = []
|
||||
merged_tags = normalize_tags([*path_tags, *tags])
|
||||
merged_tags = normalize_tags([*path_tags, *tags, "uploaded"])
|
||||
|
||||
try:
|
||||
digest, _ = hashing.compute_blake3_hash(abs_path)
|
||||
|
||||
@@ -3,59 +3,66 @@ from pathlib import Path
|
||||
from typing import Literal
|
||||
|
||||
import folder_paths
|
||||
from app.assets.helpers import normalize_tags
|
||||
|
||||
|
||||
_NON_MODEL_FOLDER_NAMES = frozenset({"custom_nodes"})
|
||||
_NON_MODEL_FOLDER_NAMES = frozenset({"configs", "custom_nodes"})
|
||||
_KNOWN_SUBFOLDER_TAGS = frozenset({"3d", "pasted", "painter", "threed", "webcam"})
|
||||
|
||||
|
||||
def get_comfy_models_folders() -> list[tuple[str, list[str]]]:
|
||||
"""Build list of (folder_name, base_paths[]) for all model locations.
|
||||
def get_comfy_models_folders() -> list[tuple[str, list[str], set[str]]]:
|
||||
"""Build list of (folder_name, base_paths[], extensions) for all model locations.
|
||||
|
||||
Includes every category registered in folder_names_and_paths,
|
||||
regardless of whether its paths are under the main models_dir,
|
||||
but excludes non-model entries like custom_nodes.
|
||||
but excludes non-model entries like configs and custom_nodes.
|
||||
|
||||
An empty extensions set means the category accepts any extension,
|
||||
matching folder_paths.filter_files_extensions semantics.
|
||||
"""
|
||||
targets: list[tuple[str, list[str]]] = []
|
||||
targets: list[tuple[str, list[str], set[str]]] = []
|
||||
for name, values in folder_paths.folder_names_and_paths.items():
|
||||
if name in _NON_MODEL_FOLDER_NAMES:
|
||||
continue
|
||||
paths, _exts = values[0], values[1]
|
||||
paths, exts = values[0], values[1]
|
||||
if paths:
|
||||
targets.append((name, paths))
|
||||
targets.append((name, paths, set(exts)))
|
||||
return targets
|
||||
|
||||
|
||||
def resolve_destination_from_tags(tags: list[str]) -> tuple[str, list[str]]:
|
||||
"""Validates and maps tags -> (base_dir, subdirs_for_fs)"""
|
||||
if not tags:
|
||||
raise ValueError("tags must not be empty")
|
||||
root = tags[0].lower()
|
||||
"""Validates and maps upload routing tags -> (base_dir, subdirs_for_fs).
|
||||
|
||||
The request tags are only used to choose the write destination. Extra tags
|
||||
remain labels; they do not become path components or trusted classification.
|
||||
"""
|
||||
destination_roles = [t for t in tags if t in {"input", "models", "output"}]
|
||||
if len(destination_roles) != 1:
|
||||
raise ValueError("uploads require exactly one destination role: input, models, or output")
|
||||
|
||||
root = destination_roles[0]
|
||||
if root == "models":
|
||||
if len(tags) < 2:
|
||||
raise ValueError("at least two tags required for model asset")
|
||||
model_type_tags = [t for t in tags if t.startswith("model_type:")]
|
||||
if len(model_type_tags) != 1:
|
||||
raise ValueError("models uploads require exactly one model_type:<folder_name> tag")
|
||||
folder_name = model_type_tags[0].split(":", 1)[1]
|
||||
if not folder_name:
|
||||
raise ValueError("models uploads require exactly one model_type:<folder_name> tag")
|
||||
model_folder_paths = {
|
||||
name: paths for name, paths, _exts in get_comfy_models_folders()
|
||||
}
|
||||
try:
|
||||
bases = folder_paths.folder_names_and_paths[tags[1]][0]
|
||||
bases = model_folder_paths[folder_name]
|
||||
except KeyError:
|
||||
raise ValueError(f"unknown model category '{tags[1]}'")
|
||||
raise ValueError(f"unknown model category '{folder_name}'")
|
||||
if not bases:
|
||||
raise ValueError(f"no base path configured for category '{tags[1]}'")
|
||||
raise ValueError(f"no base path configured for category '{folder_name}'")
|
||||
base_dir = os.path.abspath(bases[0])
|
||||
raw_subdirs = tags[2:]
|
||||
elif root == "input":
|
||||
base_dir = os.path.abspath(folder_paths.get_input_directory())
|
||||
raw_subdirs = tags[1:]
|
||||
elif root == "output":
|
||||
base_dir = os.path.abspath(folder_paths.get_output_directory())
|
||||
raw_subdirs = tags[1:]
|
||||
else:
|
||||
raise ValueError(f"unknown root tag '{tags[0]}'; expected 'models', 'input', or 'output'")
|
||||
_sep_chars = frozenset(("/", "\\", os.sep))
|
||||
for i in raw_subdirs:
|
||||
if i in (".", "..") or _sep_chars & set(i):
|
||||
raise ValueError("invalid path component in tags")
|
||||
base_dir = os.path.abspath(folder_paths.get_output_directory())
|
||||
|
||||
return base_dir, raw_subdirs if raw_subdirs else []
|
||||
return base_dir, []
|
||||
|
||||
|
||||
def validate_path_within_base(candidate: str, base: str) -> None:
|
||||
@@ -65,14 +72,79 @@ def validate_path_within_base(candidate: str, base: str) -> None:
|
||||
raise ValueError("destination escapes base directory")
|
||||
|
||||
|
||||
def compute_relative_filename(file_path: str) -> str | None:
|
||||
def _compute_relative_path(child: str, parent: str) -> str:
|
||||
rel = os.path.relpath(os.path.abspath(child), os.path.abspath(parent))
|
||||
if rel == ".":
|
||||
return ""
|
||||
return rel.replace(os.sep, "/")
|
||||
|
||||
|
||||
def _is_relative_to(child: str, parent: str) -> bool:
|
||||
return Path(os.path.abspath(child)).is_relative_to(os.path.abspath(parent))
|
||||
|
||||
|
||||
def compute_asset_response_paths(file_path: str) -> tuple[str, str | None] | None:
|
||||
"""Return (logical_path, display_name) for a file path.
|
||||
|
||||
``logical_path`` is the internal namespaced storage locator (e.g.
|
||||
``models/checkpoints/foo/bar.safetensors``); ``display_name`` is the
|
||||
human-facing label below that namespace, served on Asset responses. These
|
||||
are storage locators, not model-loader namespaces. Registered model-folder
|
||||
membership is represented by backend tags such as
|
||||
``model_type:<folder_name>``; these paths only use known storage roots.
|
||||
"""
|
||||
Return the model's path relative to the last well-known folder (the model category),
|
||||
using forward slashes, eg:
|
||||
fp_abs = os.path.abspath(file_path)
|
||||
candidates: list[tuple[int, int, str, str]] = []
|
||||
|
||||
for order, (namespace, base) in enumerate(
|
||||
(
|
||||
("input", folder_paths.get_input_directory()),
|
||||
("output", folder_paths.get_output_directory()),
|
||||
("temp", folder_paths.get_temp_directory()),
|
||||
("models", getattr(folder_paths, "models_dir", "")),
|
||||
)
|
||||
):
|
||||
if not base:
|
||||
continue
|
||||
base_abs = os.path.abspath(base)
|
||||
if _is_relative_to(fp_abs, base_abs):
|
||||
candidates.append((len(base_abs), -order, namespace, base_abs))
|
||||
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
_base_len, _order, namespace, base = max(candidates)
|
||||
rel = _compute_relative_path(fp_abs, base)
|
||||
public_path = f"{namespace}/{rel}" if rel else namespace
|
||||
return public_path, rel or None
|
||||
|
||||
|
||||
def compute_display_name(file_path: str) -> str | None:
|
||||
"""Return the asset's `display_name`, or None for unknown paths."""
|
||||
result = compute_asset_response_paths(file_path)
|
||||
return result[1] if result else None
|
||||
|
||||
|
||||
def compute_logical_path(file_path: str) -> str | None:
|
||||
"""Return the internal namespaced storage locator, or None for unknown paths."""
|
||||
result = compute_asset_response_paths(file_path)
|
||||
return result[0] if result else None
|
||||
|
||||
|
||||
def compute_loader_path(file_path: str) -> str | None:
|
||||
"""
|
||||
Return the asset's in-root loader path: the path relative to the last
|
||||
well-known folder (the model category), using forward slashes, eg:
|
||||
/.../models/checkpoints/flux/123/flux.safetensors -> "flux/123/flux.safetensors"
|
||||
/.../models/text_encoders/clip_g.safetensors -> "clip_g.safetensors"
|
||||
|
||||
For non-model paths, returns None.
|
||||
This is the value model loaders consume (the model category is dropped). It
|
||||
is persisted as ``AssetReference.loader_path`` and served as the public
|
||||
Asset response `loader_path` field. The human-facing `display_name` comes
|
||||
from compute_asset_response_paths().
|
||||
|
||||
For input/output/temp paths the full path relative to that root is returned.
|
||||
For paths outside any known root, returns None.
|
||||
"""
|
||||
try:
|
||||
root_category, rel_path = get_asset_category_and_relative_path(file_path)
|
||||
@@ -116,9 +188,10 @@ def get_asset_category_and_relative_path(
|
||||
def _compute_relative(child: str, parent: str) -> str:
|
||||
# Normalize relative path, stripping any leading ".." components
|
||||
# by anchoring to root (os.sep) then computing relpath back from it.
|
||||
return os.path.relpath(
|
||||
rel = os.path.relpath(
|
||||
os.path.join(os.sep, os.path.relpath(child, parent)), os.sep
|
||||
)
|
||||
return "" if rel == "." else rel.replace(os.sep, "/")
|
||||
|
||||
# 1) input
|
||||
input_base = os.path.abspath(folder_paths.get_input_directory())
|
||||
@@ -136,8 +209,14 @@ def get_asset_category_and_relative_path(
|
||||
return "temp", _compute_relative(fp_abs, temp_base)
|
||||
|
||||
# 4) models (check deepest matching base to avoid ambiguity)
|
||||
ext = os.path.splitext(fp_abs)[1].lower()
|
||||
best: tuple[int, str, str] | None = None # (base_len, bucket, rel_inside_bucket)
|
||||
for bucket, bases in get_comfy_models_folders():
|
||||
for bucket, bases, extensions in get_comfy_models_folders():
|
||||
# A bucket only lists files within its extension set (empty set
|
||||
# accepts any extension), so a bucket that cannot load the file
|
||||
# must not contribute a loader path.
|
||||
if extensions and ext not in extensions:
|
||||
continue
|
||||
for b in bases:
|
||||
base_abs = os.path.abspath(b)
|
||||
if not _check_is_within(fp_abs, base_abs):
|
||||
@@ -149,25 +228,111 @@ def get_asset_category_and_relative_path(
|
||||
if best is not None:
|
||||
_, bucket, rel_inside = best
|
||||
combined = os.path.join(bucket, rel_inside)
|
||||
return "models", os.path.relpath(os.path.join(os.sep, combined), os.sep)
|
||||
normalized = os.path.relpath(os.path.join(os.sep, combined), os.sep)
|
||||
return "models", normalized.replace(os.sep, "/")
|
||||
|
||||
raise ValueError(
|
||||
f"Path is not within input, output, temp, or configured model bases: {file_path}"
|
||||
)
|
||||
|
||||
|
||||
def get_backend_system_tags_from_path(path: str) -> list[str]:
|
||||
"""Return trusted backend tags derived from current filesystem facts.
|
||||
|
||||
The returned tags are only the backend-generated system tags: ``models``,
|
||||
``model_type:<folder_name>``, ``input``, ``output``, and ``temp``. Model
|
||||
type tags are based on registered folder names, not path components.
|
||||
|
||||
A ``model_type:<folder_name>`` tag is only emitted when the file's
|
||||
extension is accepted by that folder's registered extension set, so
|
||||
categories sharing a base directory tag only the files they can
|
||||
actually load. Files under a model base whose extension matches no
|
||||
category still get the ``models`` tag.
|
||||
"""
|
||||
fp_abs = os.path.abspath(path)
|
||||
fp_path = Path(fp_abs)
|
||||
tags: list[str] = []
|
||||
|
||||
def _add(tag: str) -> None:
|
||||
if tag not in tags:
|
||||
tags.append(tag)
|
||||
|
||||
for role, base in (
|
||||
("input", folder_paths.get_input_directory()),
|
||||
("output", folder_paths.get_output_directory()),
|
||||
("temp", folder_paths.get_temp_directory()),
|
||||
):
|
||||
if fp_path.is_relative_to(os.path.abspath(base)):
|
||||
_add(role)
|
||||
|
||||
ext = os.path.splitext(fp_abs)[1].lower()
|
||||
model_types: list[str] = []
|
||||
under_models_base = False
|
||||
for folder_name, bases, extensions in get_comfy_models_folders():
|
||||
for base in bases:
|
||||
if fp_path.is_relative_to(os.path.abspath(base)):
|
||||
under_models_base = True
|
||||
# Empty set accepts any extension, matching
|
||||
# folder_paths.filter_files_extensions semantics.
|
||||
if not extensions or ext in extensions:
|
||||
model_types.append(folder_name)
|
||||
break
|
||||
|
||||
if under_models_base:
|
||||
_add("models")
|
||||
for folder_name in model_types:
|
||||
_add(f"model_type:{folder_name}")
|
||||
|
||||
if not tags:
|
||||
raise ValueError(
|
||||
f"Path is not within input, output, temp, or configured model bases: {path}"
|
||||
)
|
||||
return tags
|
||||
|
||||
|
||||
def get_known_subfolder_tags(subfolder: str | None) -> list[str]:
|
||||
"""Return tags for known UI/input subfolder names."""
|
||||
if subfolder in _KNOWN_SUBFOLDER_TAGS:
|
||||
return [subfolder]
|
||||
return []
|
||||
|
||||
|
||||
def get_known_input_subfolder_tags_from_path(path: str) -> list[str]:
|
||||
"""Return known input-layout tags for files in canonical input subfolders.
|
||||
|
||||
These are compatibility tags for current UI-origin input directories such as
|
||||
``pasted`` and ``webcam``. They are intentionally narrow: only files directly
|
||||
inside a known top-level input directory receive the matching tag.
|
||||
"""
|
||||
fp_abs = os.path.abspath(path)
|
||||
input_base = os.path.abspath(folder_paths.get_input_directory())
|
||||
if not Path(fp_abs).is_relative_to(input_base):
|
||||
return []
|
||||
|
||||
rel = os.path.relpath(fp_abs, input_base)
|
||||
parts = Path(rel).parts
|
||||
if len(parts) == 2:
|
||||
return get_known_subfolder_tags(parts[0])
|
||||
return []
|
||||
|
||||
|
||||
def get_path_derived_tags_from_path(path: str) -> list[str]:
|
||||
"""Return all backend-derived tags for an asset path."""
|
||||
tags = get_backend_system_tags_from_path(path)
|
||||
for tag in get_known_input_subfolder_tags_from_path(path):
|
||||
if tag not in tags:
|
||||
tags.append(tag)
|
||||
return tags
|
||||
|
||||
|
||||
def get_name_and_tags_from_asset_path(file_path: str) -> tuple[str, list[str]]:
|
||||
"""Return (name, tags) derived from a filesystem path.
|
||||
|
||||
- name: base filename with extension
|
||||
- tags: [root_category] + parent folder names in order
|
||||
- tags: backend-derived tags from root/model classification and known input
|
||||
subfolder layout conventions
|
||||
|
||||
Raises:
|
||||
ValueError: path does not belong to any known root.
|
||||
"""
|
||||
root_category, some_path = get_asset_category_and_relative_path(file_path)
|
||||
p = Path(some_path)
|
||||
parent_parts = [
|
||||
part for part in p.parent.parts if part not in (".", "..", p.anchor)
|
||||
]
|
||||
return p.name, list(dict.fromkeys(normalize_tags([root_category, *parent_parts])))
|
||||
return Path(file_path).name, get_path_derived_tags_from_path(file_path)
|
||||
|
||||
@@ -25,6 +25,7 @@ class ReferenceData:
|
||||
preview_id: str | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
loader_path: str | None = None
|
||||
system_metadata: dict[str, Any] | None = None
|
||||
job_id: str | None = None
|
||||
last_access_time: datetime | None = None
|
||||
@@ -93,6 +94,7 @@ def extract_reference_data(ref: AssetReference) -> ReferenceData:
|
||||
id=ref.id,
|
||||
name=ref.name,
|
||||
file_path=ref.file_path,
|
||||
loader_path=ref.loader_path,
|
||||
user_metadata=ref.user_metadata,
|
||||
preview_id=ref.preview_id,
|
||||
system_metadata=ref.system_metadata,
|
||||
|
||||
@@ -35,7 +35,11 @@ class ModelFileManager:
|
||||
for folder in model_types:
|
||||
if folder in folder_black_list:
|
||||
continue
|
||||
output_folders.append({"name": folder, "folders": folder_paths.get_folder_paths(folder)})
|
||||
output_folders.append({
|
||||
"name": folder,
|
||||
"folders": folder_paths.get_folder_paths(folder),
|
||||
"extensions": sorted(folder_paths.folder_names_and_paths[folder][1]),
|
||||
})
|
||||
return web.json_response(output_folders)
|
||||
|
||||
# NOTE: This is an experiment to replace `/models/{folder}`
|
||||
|
||||
@@ -92,6 +92,7 @@ parser.add_argument("--directml", type=int, nargs="?", metavar="DIRECTML_DEVICE"
|
||||
parser.add_argument("--oneapi-device-selector", type=str, default=None, metavar="SELECTOR_STRING", help="Sets the oneAPI device(s) this instance will use.")
|
||||
parser.add_argument("--supports-fp8-compute", action="store_true", help="ComfyUI will act like if the device supports fp8 compute.")
|
||||
parser.add_argument("--enable-triton-backend", action="store_true", help="ComfyUI will enable the use of Triton backend in comfy-kitchen. Is disabled at launch by default.")
|
||||
parser.add_argument("--disable-triton-backend", action="store_true", help="Force-disable the comfy-kitchen Triton backend, overriding the automatic ROCm/AMD default and --enable-triton-backend.")
|
||||
|
||||
class LatentPreviewMethod(enum.Enum):
|
||||
NoPreviews = "none"
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
"""Runtime config the frontend reads from /features to follow --comfy-api-base.
|
||||
|
||||
For a non-prod comfy.org backend (staging or an ephemeral preview env), "/features" exposes the api and
|
||||
platform base so the frontend talks to it without a rebuild, plus the Firebase environment it should use.
|
||||
Prod bases are left alone and keep their build-time defaults.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
_STAGING_API_HOST = "stagingapi.comfy.org"
|
||||
_TESTENV_HOST_SUFFIX = ".testenvs.comfy.org"
|
||||
_STAGING_PLATFORM_BASE_URL = "https://stagingplatform.comfy.org"
|
||||
|
||||
|
||||
def _is_staging_tier(host: str) -> bool:
|
||||
return host == _STAGING_API_HOST or host.endswith(_TESTENV_HOST_SUFFIX)
|
||||
|
||||
|
||||
def normalize_comfy_api_base(url: str) -> str:
|
||||
"""Rewrite a testenv's friendly main host to its comfy-api '-registry' sibling."""
|
||||
parsed = urlparse(url)
|
||||
host = parsed.hostname or ""
|
||||
if not host.endswith(_TESTENV_HOST_SUFFIX):
|
||||
return url
|
||||
label = host[: -len(_TESTENV_HOST_SUFFIX)]
|
||||
if label.endswith("-registry"):
|
||||
return url
|
||||
return f"{parsed.scheme or 'https'}://{label}-registry{_TESTENV_HOST_SUFFIX}"
|
||||
|
||||
|
||||
def environment_overrides_for_base(base_url: str) -> dict[str, Any] | None:
|
||||
"""The /features overrides for a staging-tier base, or None for prod."""
|
||||
if not _is_staging_tier(urlparse(base_url).hostname or ""):
|
||||
return None
|
||||
return {
|
||||
"comfy_api_base_url": normalize_comfy_api_base(base_url).rstrip("/"),
|
||||
"comfy_platform_base_url": _STAGING_PLATFORM_BASE_URL,
|
||||
"firebase_env": "dev",
|
||||
}
|
||||
|
||||
|
||||
def get_environment_overrides() -> dict[str, Any] | None:
|
||||
return environment_overrides_for_base(getattr(args, "comfy_api_base", "") or "")
|
||||
@@ -779,6 +779,10 @@ class ACEAudio(LatentFormat):
|
||||
latent_channels = 8
|
||||
latent_dimensions = 2
|
||||
|
||||
class SeedVR2(LatentFormat):
|
||||
latent_channels = 16
|
||||
latent_dimensions = 3
|
||||
|
||||
class ACEAudio15(LatentFormat):
|
||||
latent_channels = 64
|
||||
latent_dimensions = 1
|
||||
|
||||
@@ -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, **kwargs):
|
||||
def two_pass_attention(q, k, v, heads, enable_gqa=False, **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)
|
||||
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa)
|
||||
elif ar_len >= T:
|
||||
out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True)
|
||||
out = comfy.ops.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa)
|
||||
elif ar_len <= 0:
|
||||
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options)
|
||||
out = optimized_attention(q, k, v, heads, mask=None, skip_reshape=True, skip_output_reshape=True, transformer_options=transformer_options, enable_gqa=enable_gqa)
|
||||
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,
|
||||
attn_mask=None, dropout_p=0.0, is_causal=True, enable_gqa=enable_gqa,
|
||||
)
|
||||
out_gen = optimized_attention(
|
||||
q[:, :, ar_len:], k, v, heads,
|
||||
mask=None, skip_reshape=True, skip_output_reshape=True,
|
||||
transformer_options=transformer_options,
|
||||
transformer_options=transformer_options, enable_gqa=enable_gqa,
|
||||
)
|
||||
out = torch.cat([out_ar, out_gen], dim=2)
|
||||
|
||||
|
||||
@@ -709,7 +709,7 @@ def attention3_sage(q, k, v, heads, mask=None, attn_precision=None, skip_reshape
|
||||
return out
|
||||
|
||||
try:
|
||||
@torch.library.custom_op("flash_attention::flash_attn", mutates_args=())
|
||||
@torch.library.custom_op("comfy::flash_attn", mutates_args=())
|
||||
def flash_attn_wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
|
||||
dropout_p: float = 0.0, causal: bool = False, softmax_scale: float = -1.0) -> torch.Tensor:
|
||||
softmax_scale_arg = None if softmax_scale == -1.0 else softmax_scale
|
||||
|
||||
@@ -22,7 +22,7 @@ def torch_cat_if_needed(xl, dim):
|
||||
else:
|
||||
return None
|
||||
|
||||
def get_timestep_embedding(timesteps, embedding_dim):
|
||||
def get_timestep_embedding(timesteps, embedding_dim, flip_sin_to_cos=False, downscale_freq_shift=1):
|
||||
"""
|
||||
This matches the implementation in Denoising Diffusion Probabilistic Models:
|
||||
From Fairseq.
|
||||
@@ -33,11 +33,13 @@ def get_timestep_embedding(timesteps, embedding_dim):
|
||||
assert len(timesteps.shape) == 1
|
||||
|
||||
half_dim = embedding_dim // 2
|
||||
emb = math.log(10000) / (half_dim - 1)
|
||||
emb = math.log(10000) / (half_dim - downscale_freq_shift)
|
||||
emb = torch.exp(torch.arange(half_dim, dtype=torch.float32) * -emb)
|
||||
emb = emb.to(device=timesteps.device)
|
||||
emb = timesteps.float()[:, None] * emb[None, :]
|
||||
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1)
|
||||
if flip_sin_to_cos:
|
||||
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
|
||||
if embedding_dim % 2 == 1: # zero pad
|
||||
emb = torch.nn.functional.pad(emb, (0,1,0,0))
|
||||
return emb
|
||||
|
||||
@@ -197,6 +197,9 @@ class PixDiT_T2I(nn.Module):
|
||||
"""Hook for subclasses to inject per-block state into the patch stream (e.g. PiD's LQ gate)."""
|
||||
return s
|
||||
|
||||
def _pre_pixel_blocks(self, s, **kwargs):
|
||||
return s
|
||||
|
||||
def _forward(self, x, timesteps, context=None, attention_mask=None, transformer_options={}, **kwargs):
|
||||
H_orig, W_orig = x.shape[2], x.shape[3]
|
||||
x = comfy.ldm.common_dit.pad_to_patch_size(x, (self.patch_size, self.patch_size))
|
||||
@@ -226,6 +229,7 @@ class PixDiT_T2I(nn.Module):
|
||||
s, y_emb = blk(s, y_emb, condition, pos_img, pos_txt, None, transformer_options=transformer_options)
|
||||
s = F.silu(t_emb + s)
|
||||
|
||||
s = self._pre_pixel_blocks(s, **kwargs)
|
||||
s_cond = s.view(B * L, self.hidden_size)
|
||||
x_pixels = self.pixel_embedder(x, patch_size=self.patch_size)
|
||||
for blk in self.pixel_blocks:
|
||||
|
||||
+50
-14
@@ -13,15 +13,15 @@ from .model import PixDiT_T2I
|
||||
from .modules import precompute_freqs_cis_2d
|
||||
|
||||
|
||||
class SigmaAwareGatePerTokenPerDim(nn.Module):
|
||||
class SigmaAwareGate(nn.Module):
|
||||
"""gate = sigmoid(content_proj(cat[x, lq]) - exp(log_alpha) * sigma); out = x + gate * lq.
|
||||
|
||||
Trained init gives ~0.88 gate at sigma=0, ~0.05 at sigma=1.
|
||||
"""
|
||||
|
||||
def __init__(self, dim: int, dtype=None, device=None, operations=None):
|
||||
def __init__(self, dim: int, per_token: bool = False, dtype=None, device=None, operations=None):
|
||||
super().__init__()
|
||||
self.content_proj = operations.Linear(dim * 2, dim, dtype=dtype, device=device)
|
||||
self.content_proj = operations.Linear(dim * 2, 1 if per_token else dim, dtype=dtype, device=device)
|
||||
self.log_alpha = nn.Parameter(torch.empty((), dtype=dtype, device=device))
|
||||
|
||||
def forward(self, x: torch.Tensor, lq: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor:
|
||||
@@ -36,15 +36,15 @@ class SigmaAwareGatePerTokenPerDim(nn.Module):
|
||||
class ResBlock(nn.Module):
|
||||
"""Pre-activation ResNet block: GN -> SiLU -> Conv -> GN -> SiLU -> Conv + skip."""
|
||||
|
||||
def __init__(self, channels: int, num_groups: int = 4, dtype=None, device=None, operations=None):
|
||||
def __init__(self, channels: int, num_groups: int = 4, conv_padding_mode: str = "zeros", dtype=None, device=None, operations=None):
|
||||
super().__init__()
|
||||
self.block = nn.Sequential(
|
||||
operations.GroupNorm(num_groups, channels, dtype=dtype, device=device),
|
||||
nn.SiLU(),
|
||||
operations.Conv2d(channels, channels, kernel_size=3, padding=1, dtype=dtype, device=device),
|
||||
operations.Conv2d(channels, channels, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device),
|
||||
operations.GroupNorm(num_groups, channels, dtype=dtype, device=device),
|
||||
nn.SiLU(),
|
||||
operations.Conv2d(channels, channels, kernel_size=3, padding=1, dtype=dtype, device=device),
|
||||
operations.Conv2d(channels, channels, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device),
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
@@ -62,9 +62,13 @@ class LQProjection2D(nn.Module):
|
||||
patch_size: int = 16,
|
||||
sr_scale: int = 4,
|
||||
latent_spatial_down_factor: int = 8,
|
||||
latent_unpatchify_factor: int = 1,
|
||||
num_res_blocks: int = 4,
|
||||
num_outputs: int = 7,
|
||||
interval: int = 2,
|
||||
conv_padding_mode: str = "zeros",
|
||||
gate_per_token: bool = False,
|
||||
pit_output: bool = False,
|
||||
dtype=None, device=None, operations=None,
|
||||
):
|
||||
super().__init__()
|
||||
@@ -74,34 +78,38 @@ class LQProjection2D(nn.Module):
|
||||
self.patch_size = patch_size
|
||||
self.sr_scale = sr_scale
|
||||
self.latent_spatial_down_factor = latent_spatial_down_factor
|
||||
self.latent_unpatchify_factor = latent_unpatchify_factor
|
||||
self.num_outputs = num_outputs
|
||||
self.interval = interval
|
||||
|
||||
z_to_patch_ratio = (sr_scale * latent_spatial_down_factor) / patch_size
|
||||
effective_latent_channels = latent_channels // (latent_unpatchify_factor * latent_unpatchify_factor)
|
||||
effective_spatial_down_factor = latent_spatial_down_factor // latent_unpatchify_factor
|
||||
z_to_patch_ratio = (sr_scale * effective_spatial_down_factor) / patch_size
|
||||
self.z_to_patch_ratio = z_to_patch_ratio
|
||||
if z_to_patch_ratio >= 1:
|
||||
self.latent_fold_factor = 0
|
||||
latent_proj_in_ch = latent_channels
|
||||
latent_proj_in_ch = effective_latent_channels
|
||||
else:
|
||||
fold_factor = int(1 / z_to_patch_ratio)
|
||||
assert fold_factor * z_to_patch_ratio == 1.0
|
||||
self.latent_fold_factor = fold_factor
|
||||
latent_proj_in_ch = latent_channels * fold_factor * fold_factor
|
||||
latent_proj_in_ch = effective_latent_channels * fold_factor * fold_factor
|
||||
|
||||
layers = [
|
||||
operations.Conv2d(latent_proj_in_ch, hidden_dim, kernel_size=3, padding=1, dtype=dtype, device=device),
|
||||
operations.Conv2d(latent_proj_in_ch, hidden_dim, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device),
|
||||
nn.SiLU(),
|
||||
operations.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1, dtype=dtype, device=device),
|
||||
operations.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1, padding_mode=conv_padding_mode, dtype=dtype, device=device),
|
||||
]
|
||||
for _ in range(num_res_blocks):
|
||||
layers.append(ResBlock(hidden_dim, dtype=dtype, device=device, operations=operations))
|
||||
layers.append(ResBlock(hidden_dim, conv_padding_mode=conv_padding_mode, dtype=dtype, device=device, operations=operations))
|
||||
self.latent_proj = nn.Sequential(*layers)
|
||||
|
||||
self.output_heads = nn.ModuleList(
|
||||
[operations.Linear(hidden_dim, out_dim, dtype=dtype, device=device) for _ in range(num_outputs)]
|
||||
)
|
||||
self.pit_head = operations.Linear(hidden_dim, out_dim, dtype=dtype, device=device) if pit_output else None
|
||||
self.gate_modules = nn.ModuleList(
|
||||
[SigmaAwareGatePerTokenPerDim(out_dim, dtype=dtype, device=device, operations=operations)
|
||||
[SigmaAwareGate(out_dim, per_token=gate_per_token, dtype=dtype, device=device, operations=operations)
|
||||
for _ in range(num_outputs)]
|
||||
)
|
||||
|
||||
@@ -115,6 +123,11 @@ class LQProjection2D(nn.Module):
|
||||
return self.gate_modules[out_idx](x, lq_feature, sigma)
|
||||
|
||||
def _align_latent_to_patch_grid(self, lq_latent: torch.Tensor, pH: int, pW: int) -> torch.Tensor:
|
||||
f = self.latent_unpatchify_factor
|
||||
if f > 1:
|
||||
B, C, H, W = lq_latent.shape
|
||||
lq_latent = lq_latent.reshape(B, C // (f * f), f, f, H, W)
|
||||
lq_latent = lq_latent.permute(0, 1, 4, 2, 5, 3).reshape(B, C // (f * f), H * f, W * f)
|
||||
B, z_dim = lq_latent.shape[:2]
|
||||
if self.z_to_patch_ratio >= 1:
|
||||
if lq_latent.shape[2] != pH or lq_latent.shape[3] != pW:
|
||||
@@ -134,7 +147,10 @@ class LQProjection2D(nn.Module):
|
||||
feat = self._align_latent_to_patch_grid(lq_latent, target_pH, target_pW)
|
||||
B, C, H, W = feat.shape
|
||||
tokens = feat.permute(0, 2, 3, 1).contiguous().view(B, H * W, C)
|
||||
return [head(tokens) for head in self.output_heads]
|
||||
outputs = [head(tokens) for head in self.output_heads]
|
||||
if self.pit_head is not None:
|
||||
outputs.append(self.pit_head(tokens))
|
||||
return outputs
|
||||
|
||||
|
||||
class PidNet(PixDiT_T2I):
|
||||
@@ -148,6 +164,10 @@ class PidNet(PixDiT_T2I):
|
||||
lq_interval: int = 2,
|
||||
sr_scale: int = 4,
|
||||
latent_spatial_down_factor: int = 8,
|
||||
lq_latent_unpatchify_factor: int = 1,
|
||||
lq_conv_padding_mode: str = "zeros",
|
||||
lq_gate_per_token: bool = False,
|
||||
pit_lq_inject: bool = False,
|
||||
rope_ref_h: int = 1024, # NTK ref resolution in PIXEL units: 1024px / patch=16 -> grid_ref=64.
|
||||
rope_ref_w: int = 1024,
|
||||
image_model=None,
|
||||
@@ -165,6 +185,8 @@ class PidNet(PixDiT_T2I):
|
||||
for blk in self.pixel_blocks:
|
||||
blk._rope_fn = _pit_rope_fn
|
||||
|
||||
self.pit_lq_inject = pit_lq_inject
|
||||
|
||||
num_lq_outputs = (self.patch_depth + lq_interval - 1) // lq_interval
|
||||
self.lq_proj = LQProjection2D(
|
||||
latent_channels=lq_latent_channels,
|
||||
@@ -173,13 +195,20 @@ class PidNet(PixDiT_T2I):
|
||||
patch_size=self.patch_size,
|
||||
sr_scale=sr_scale,
|
||||
latent_spatial_down_factor=latent_spatial_down_factor,
|
||||
latent_unpatchify_factor=lq_latent_unpatchify_factor,
|
||||
num_res_blocks=lq_num_res_blocks,
|
||||
num_outputs=num_lq_outputs,
|
||||
interval=lq_interval,
|
||||
conv_padding_mode=lq_conv_padding_mode,
|
||||
gate_per_token=lq_gate_per_token,
|
||||
pit_output=pit_lq_inject,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations,
|
||||
)
|
||||
self.pit_lq_gate = SigmaAwareGate(
|
||||
self.hidden_size, per_token=lq_gate_per_token, dtype=dtype, device=device, operations=operations
|
||||
) if pit_lq_inject else None
|
||||
|
||||
def _fetch_patch_pos(self, height, width, device, dtype, **rope_opts):
|
||||
return precompute_freqs_cis_2d(
|
||||
@@ -197,6 +226,11 @@ class PidNet(PixDiT_T2I):
|
||||
return s
|
||||
return self.lq_proj.gate(s, pid_lq_features[out_idx], pid_degrade_sigma, out_idx)
|
||||
|
||||
def _pre_pixel_blocks(self, s, pid_pit_lq_feature=None, pid_degrade_sigma=None, **kwargs):
|
||||
if pid_pit_lq_feature is None:
|
||||
return s
|
||||
return self.pit_lq_gate(s, pid_pit_lq_feature, pid_degrade_sigma)
|
||||
|
||||
def _forward(self, x, timesteps, context=None, attention_mask=None, transformer_options={}, lq_latent=None, degrade_sigma=None, **kwargs):
|
||||
if lq_latent is None:
|
||||
raise ValueError("PidNet requires lq_latent — attach via PiDConditioning")
|
||||
@@ -216,12 +250,14 @@ class PidNet(PixDiT_T2I):
|
||||
degrade_sigma = degrade_sigma.expand(B).contiguous()
|
||||
|
||||
lq_features = self.lq_proj(lq_latent=lq_latent.to(x), target_pH=Hs, target_pW=Ws)
|
||||
pit_lq_feature = lq_features.pop() if self.pit_lq_inject else None
|
||||
|
||||
return super()._forward(
|
||||
x, timesteps,
|
||||
context=context, attention_mask=attention_mask,
|
||||
transformer_options=transformer_options,
|
||||
pid_lq_features=lq_features,
|
||||
pid_pit_lq_feature=pit_lq_feature,
|
||||
pid_degrade_sigma=degrade_sigma,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
import torch
|
||||
|
||||
from comfy.ldm.modules import attention as _attention
|
||||
|
||||
|
||||
def _var_attention_qkv(q, k, v, heads, skip_reshape):
|
||||
if skip_reshape:
|
||||
return q, k, v, q.shape[-1]
|
||||
total_tokens, embed_dim = q.shape
|
||||
head_dim = embed_dim // heads
|
||||
return (
|
||||
q.view(total_tokens, heads, head_dim),
|
||||
k.view(k.shape[0], heads, head_dim),
|
||||
v.view(v.shape[0], heads, head_dim),
|
||||
head_dim,
|
||||
)
|
||||
|
||||
|
||||
def _var_attention_output(out, heads, head_dim, skip_output_reshape):
|
||||
if skip_output_reshape:
|
||||
return out
|
||||
return out.reshape(-1, heads * head_dim)
|
||||
|
||||
|
||||
def var_attention_optimized_split(q, k, v, heads, cu_seqlens_q, cu_seqlens_k, *args, skip_reshape=False, skip_output_reshape=False, **kwargs):
|
||||
q, k, v, head_dim = _var_attention_qkv(q, k, v, heads, skip_reshape)
|
||||
|
||||
q_split_indices = cu_seqlens_q[1:-1]
|
||||
k_split_indices = cu_seqlens_k[1:-1]
|
||||
if k.shape[0] != v.shape[0]:
|
||||
raise ValueError("cu_seqlens_k does not match v token count")
|
||||
|
||||
q_splits = torch.tensor_split(q, q_split_indices, dim=0)
|
||||
k_splits = torch.tensor_split(k, k_split_indices, dim=0)
|
||||
v_splits = torch.tensor_split(v, k_split_indices, dim=0)
|
||||
if len(q_splits) != len(k_splits) or len(q_splits) != len(v_splits):
|
||||
raise ValueError("cu_seqlens_q and cu_seqlens_k must describe the same sequence count")
|
||||
|
||||
out = []
|
||||
for q_i, k_i, v_i in zip(q_splits, k_splits, v_splits):
|
||||
q_i = q_i.permute(1, 0, 2).unsqueeze(0)
|
||||
k_i = k_i.permute(1, 0, 2).unsqueeze(0)
|
||||
v_i = v_i.permute(1, 0, 2).unsqueeze(0)
|
||||
out_i = _attention.optimized_attention(q_i, k_i, v_i, heads, skip_reshape=True, skip_output_reshape=True)
|
||||
out.append(out_i.squeeze(0).permute(1, 0, 2))
|
||||
|
||||
out = torch.cat(out, dim=0)
|
||||
return _var_attention_output(out, heads, head_dim, skip_output_reshape)
|
||||
|
||||
|
||||
optimized_var_attention = var_attention_optimized_split
|
||||
@@ -0,0 +1,301 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
|
||||
from comfy.ldm.seedvr.constants import (
|
||||
CIELAB_DELTA,
|
||||
CIELAB_KAPPA,
|
||||
D65_WHITE_X,
|
||||
D65_WHITE_Z,
|
||||
WAVELET_DECOMP_LEVELS,
|
||||
)
|
||||
|
||||
|
||||
def wavelet_blur(image: Tensor, radius):
|
||||
max_safe_radius = max(1, min(image.shape[-2:]) // 8)
|
||||
if radius > max_safe_radius:
|
||||
radius = max_safe_radius
|
||||
|
||||
num_channels = image.shape[1]
|
||||
|
||||
kernel_vals = [
|
||||
[0.0625, 0.125, 0.0625],
|
||||
[0.125, 0.25, 0.125],
|
||||
[0.0625, 0.125, 0.0625],
|
||||
]
|
||||
kernel = torch.tensor(kernel_vals, dtype=image.dtype, device=image.device)
|
||||
kernel = kernel[None, None].repeat(num_channels, 1, 1, 1)
|
||||
|
||||
image = F.pad(image, (radius, radius, radius, radius), mode='replicate')
|
||||
output = F.conv2d(image, kernel, groups=num_channels, dilation=radius)
|
||||
|
||||
return output
|
||||
|
||||
def wavelet_decomposition(image: Tensor, levels: int = WAVELET_DECOMP_LEVELS):
|
||||
high_freq = torch.zeros_like(image)
|
||||
|
||||
for i in range(levels):
|
||||
radius = 2 ** i
|
||||
low_freq = wavelet_blur(image, radius)
|
||||
high_freq.add_(image).sub_(low_freq)
|
||||
image = low_freq
|
||||
|
||||
return high_freq, low_freq
|
||||
|
||||
def wavelet_reconstruction(content_feat: Tensor, style_feat: Tensor) -> Tensor:
|
||||
|
||||
if content_feat.shape != style_feat.shape:
|
||||
if len(content_feat.shape) >= 3:
|
||||
style_feat = F.interpolate(
|
||||
style_feat,
|
||||
size=content_feat.shape[-2:],
|
||||
mode='bilinear',
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
content_high_freq, content_low_freq = wavelet_decomposition(content_feat)
|
||||
del content_low_freq
|
||||
|
||||
style_high_freq, style_low_freq = wavelet_decomposition(style_feat)
|
||||
del style_high_freq
|
||||
|
||||
if content_high_freq.shape != style_low_freq.shape:
|
||||
style_low_freq = F.interpolate(
|
||||
style_low_freq,
|
||||
size=content_high_freq.shape[-2:],
|
||||
mode='bilinear',
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
content_high_freq.add_(style_low_freq)
|
||||
|
||||
return content_high_freq.clamp_(-1.0, 1.0)
|
||||
|
||||
def _histogram_matching_channel(source: Tensor, reference: Tensor) -> Tensor:
|
||||
original_shape = source.shape
|
||||
|
||||
source_flat = source.flatten()
|
||||
reference_flat = reference.flatten()
|
||||
|
||||
source_sorted, source_indices = torch.sort(source_flat)
|
||||
reference_sorted, _ = torch.sort(reference_flat)
|
||||
del reference_flat
|
||||
|
||||
n_source = len(source_sorted)
|
||||
n_reference = len(reference_sorted)
|
||||
|
||||
if n_source == n_reference:
|
||||
matched_sorted = reference_sorted
|
||||
else:
|
||||
source_quantiles = torch.linspace(0, 1, n_source, device=source.device)
|
||||
ref_indices = (source_quantiles * (n_reference - 1)).long()
|
||||
ref_indices.clamp_(0, n_reference - 1)
|
||||
matched_sorted = reference_sorted[ref_indices]
|
||||
del source_quantiles, ref_indices, reference_sorted
|
||||
|
||||
del source_sorted, source_flat
|
||||
|
||||
inverse_indices = torch.argsort(source_indices)
|
||||
del source_indices
|
||||
matched_flat = matched_sorted[inverse_indices]
|
||||
del matched_sorted, inverse_indices
|
||||
|
||||
return matched_flat.reshape(original_shape)
|
||||
|
||||
def _lab_to_rgb_batch(lab: Tensor, matrix_inv: Tensor, epsilon: float, kappa: float) -> Tensor:
|
||||
L, a, b = lab[:, 0], lab[:, 1], lab[:, 2]
|
||||
|
||||
fy = (L + 16.0) / 116.0
|
||||
fx = a.div(500.0).add_(fy)
|
||||
fz = fy - b / 200.0
|
||||
del L, a, b
|
||||
|
||||
x = torch.where(
|
||||
fx > epsilon,
|
||||
torch.pow(fx, 3.0),
|
||||
fx.mul(116.0).sub_(16.0).div_(kappa)
|
||||
)
|
||||
y = torch.where(
|
||||
fy > epsilon,
|
||||
torch.pow(fy, 3.0),
|
||||
fy.mul(116.0).sub_(16.0).div_(kappa)
|
||||
)
|
||||
z = torch.where(
|
||||
fz > epsilon,
|
||||
torch.pow(fz, 3.0),
|
||||
fz.mul(116.0).sub_(16.0).div_(kappa)
|
||||
)
|
||||
del fx, fy, fz
|
||||
|
||||
x.mul_(D65_WHITE_X)
|
||||
z.mul_(D65_WHITE_Z)
|
||||
|
||||
xyz = torch.stack([x, y, z], dim=1)
|
||||
del x, y, z
|
||||
|
||||
B, _, H, W = xyz.shape
|
||||
xyz_flat = xyz.permute(0, 2, 3, 1).reshape(-1, 3)
|
||||
del xyz
|
||||
|
||||
xyz_flat = xyz_flat.to(dtype=matrix_inv.dtype)
|
||||
rgb_linear_flat = torch.matmul(xyz_flat, matrix_inv.T)
|
||||
del xyz_flat
|
||||
|
||||
rgb_linear = rgb_linear_flat.reshape(B, H, W, 3).permute(0, 3, 1, 2)
|
||||
del rgb_linear_flat
|
||||
|
||||
mask = rgb_linear > 0.0031308
|
||||
rgb = torch.where(
|
||||
mask,
|
||||
torch.pow(torch.clamp(rgb_linear, min=0.0), 1.0 / 2.4).mul_(1.055).sub_(0.055),
|
||||
rgb_linear * 12.92
|
||||
)
|
||||
del mask, rgb_linear
|
||||
|
||||
return torch.clamp(rgb, 0.0, 1.0)
|
||||
|
||||
def _rgb_to_lab_batch(rgb: Tensor, matrix: Tensor, epsilon: float, kappa: float) -> Tensor:
|
||||
mask = rgb > 0.04045
|
||||
rgb_linear = torch.where(
|
||||
mask,
|
||||
torch.pow((rgb + 0.055) / 1.055, 2.4),
|
||||
rgb / 12.92
|
||||
)
|
||||
del mask
|
||||
|
||||
B, _, H, W = rgb_linear.shape
|
||||
rgb_flat = rgb_linear.permute(0, 2, 3, 1).reshape(-1, 3)
|
||||
del rgb_linear
|
||||
|
||||
rgb_flat = rgb_flat.to(dtype=matrix.dtype)
|
||||
xyz_flat = torch.matmul(rgb_flat, matrix.T)
|
||||
del rgb_flat
|
||||
|
||||
xyz = xyz_flat.reshape(B, H, W, 3).permute(0, 3, 1, 2)
|
||||
del xyz_flat
|
||||
|
||||
xyz[:, 0].div_(D65_WHITE_X)
|
||||
xyz[:, 2].div_(D65_WHITE_Z)
|
||||
|
||||
epsilon_cubed = epsilon ** 3
|
||||
mask = xyz > epsilon_cubed
|
||||
f_xyz = torch.where(
|
||||
mask,
|
||||
torch.pow(xyz, 1.0 / 3.0),
|
||||
xyz.mul(kappa).add_(16.0).div_(116.0)
|
||||
)
|
||||
del xyz, mask
|
||||
|
||||
L = f_xyz[:, 1].mul(116.0).sub_(16.0)
|
||||
a = (f_xyz[:, 0] - f_xyz[:, 1]).mul_(500.0)
|
||||
b = (f_xyz[:, 1] - f_xyz[:, 2]).mul_(200.0)
|
||||
del f_xyz
|
||||
|
||||
return torch.stack([L, a, b], dim=1)
|
||||
|
||||
def lab_color_transfer(
|
||||
content_feat: Tensor,
|
||||
style_feat: Tensor,
|
||||
luminance_weight: float = 0.8
|
||||
) -> Tensor:
|
||||
content_feat = wavelet_reconstruction(content_feat, style_feat)
|
||||
|
||||
if content_feat.shape != style_feat.shape:
|
||||
style_feat = F.interpolate(
|
||||
style_feat,
|
||||
size=content_feat.shape[-2:],
|
||||
mode='bilinear',
|
||||
align_corners=False
|
||||
)
|
||||
|
||||
device = content_feat.device
|
||||
original_dtype = content_feat.dtype
|
||||
content_feat = content_feat.float()
|
||||
style_feat = style_feat.float()
|
||||
|
||||
rgb_to_xyz_matrix = torch.tensor([
|
||||
[0.4124564, 0.3575761, 0.1804375],
|
||||
[0.2126729, 0.7151522, 0.0721750],
|
||||
[0.0193339, 0.1191920, 0.9503041]
|
||||
], dtype=torch.float32, device=device)
|
||||
|
||||
xyz_to_rgb_matrix = torch.tensor([
|
||||
[ 3.2404542, -1.5371385, -0.4985314],
|
||||
[-0.9692660, 1.8760108, 0.0415560],
|
||||
[ 0.0556434, -0.2040259, 1.0572252]
|
||||
], dtype=torch.float32, device=device)
|
||||
|
||||
epsilon = CIELAB_DELTA
|
||||
kappa = CIELAB_KAPPA
|
||||
|
||||
content_feat.add_(1.0).mul_(0.5).clamp_(0.0, 1.0)
|
||||
style_feat.add_(1.0).mul_(0.5).clamp_(0.0, 1.0)
|
||||
|
||||
content_lab = _rgb_to_lab_batch(content_feat, rgb_to_xyz_matrix, epsilon, kappa)
|
||||
del content_feat
|
||||
|
||||
style_lab = _rgb_to_lab_batch(style_feat, rgb_to_xyz_matrix, epsilon, kappa)
|
||||
del style_feat, rgb_to_xyz_matrix
|
||||
|
||||
matched_a = _histogram_matching_channel(content_lab[:, 1], style_lab[:, 1])
|
||||
matched_b = _histogram_matching_channel(content_lab[:, 2], style_lab[:, 2])
|
||||
|
||||
if luminance_weight < 1.0:
|
||||
matched_L = _histogram_matching_channel(content_lab[:, 0], style_lab[:, 0])
|
||||
result_L = content_lab[:, 0].mul(luminance_weight).add_(matched_L.mul(1.0 - luminance_weight))
|
||||
del matched_L
|
||||
else:
|
||||
result_L = content_lab[:, 0]
|
||||
|
||||
del content_lab, style_lab
|
||||
|
||||
result_lab = torch.stack([result_L, matched_a, matched_b], dim=1)
|
||||
del result_L, matched_a, matched_b
|
||||
|
||||
result_rgb = _lab_to_rgb_batch(result_lab, xyz_to_rgb_matrix, epsilon, kappa)
|
||||
del result_lab, xyz_to_rgb_matrix
|
||||
|
||||
result = result_rgb.mul_(2.0).sub_(1.0)
|
||||
del result_rgb
|
||||
|
||||
result = result.to(original_dtype)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def wavelet_color_transfer(content_feat: Tensor, style_feat: Tensor) -> Tensor:
|
||||
return wavelet_reconstruction(content_feat, style_feat)
|
||||
|
||||
|
||||
def adain_color_transfer(content_feat: Tensor, style_feat: Tensor, eps: float = 1e-5) -> Tensor:
|
||||
if content_feat.shape != style_feat.shape:
|
||||
style_feat = F.interpolate(
|
||||
style_feat,
|
||||
size=content_feat.shape[-2:],
|
||||
mode='bilinear',
|
||||
align_corners=False,
|
||||
)
|
||||
|
||||
original_dtype = content_feat.dtype
|
||||
content_feat = content_feat.float()
|
||||
style_feat = style_feat.float()
|
||||
|
||||
b, c = content_feat.shape[:2]
|
||||
content_flat = content_feat.reshape(b, c, -1)
|
||||
style_flat = style_feat.reshape(b, c, -1)
|
||||
|
||||
content_mean = content_flat.mean(dim=2).reshape(b, c, 1, 1)
|
||||
content_std = (content_flat.var(dim=2, correction=0) + eps).sqrt().reshape(b, c, 1, 1)
|
||||
style_mean = style_flat.mean(dim=2).reshape(b, c, 1, 1)
|
||||
style_std = (style_flat.var(dim=2, correction=0) + eps).sqrt().reshape(b, c, 1, 1)
|
||||
del content_flat, style_flat
|
||||
|
||||
normalized = (content_feat - content_mean) / content_std
|
||||
del content_mean, content_std
|
||||
result = normalized * style_std + style_mean
|
||||
del normalized, style_mean, style_std
|
||||
|
||||
result = result.clamp_(-1.0, 1.0)
|
||||
if result.dtype != original_dtype:
|
||||
result = result.to(original_dtype)
|
||||
return result
|
||||
@@ -0,0 +1,48 @@
|
||||
"""SeedVR2 constants."""
|
||||
|
||||
# Temporal chunk-size law: the sampler's activation wall is linear in
|
||||
# T_latent * pixel area (17-cell resolution sweep + T bisection, RTX 5090, 3b fp16):
|
||||
# max_latent_frames = (free_GiB - RESERVED - K*SIGMA) / (GIB_PER_MPX_FRAME * megapixels)
|
||||
# RESERVED covers model staging plus fixed CUDA/torch overhead; SIGMA is the measured
|
||||
# run-to-run spread of the wall; K=4 trades ~10% smaller chunks for ~1e-5 OOM odds.
|
||||
SEEDVR2_CHUNK_GIB_PER_MPX_FRAME = 0.55
|
||||
SEEDVR2_CHUNK_RESERVED_GIB = 8.5
|
||||
SEEDVR2_CHUNK_SIGMA_GIB = 0.55
|
||||
SEEDVR2_CHUNK_SIGMA_K = 4
|
||||
|
||||
SEEDVR2_7B_VID_DIM = 3072
|
||||
SEEDVR2_OOM_BACKOFF_DIVISOR = 2
|
||||
SEEDVR2_DTYPE_BYTES_FLOOR = 4
|
||||
SEEDVR2_7B_MLP_CHUNK = 8192
|
||||
SEEDVR2_ROPE_PARTIAL_CHUNK_TOKENS = 4096 # partial-RoPE application token-chunk.
|
||||
SEEDVR2_LATENT_CHANNELS = 16
|
||||
|
||||
SEEDVR2_COLOR_MEM_HEADROOM = 0.75
|
||||
SEEDVR2_LAB_SCALE_MULTIPLIER = 13
|
||||
SEEDVR2_WAVELET_SCALE_MULTIPLIER = 10 # per-frame byte multiplier, wavelet path.
|
||||
SEEDVR2_ADAIN_SCALE_MULTIPLIER = 6
|
||||
|
||||
BYTEDANCE_VAE_SCALING_FACTOR = 0.9152 # configs_3b/main.yaml:57.
|
||||
BYTEDANCE_VAE_SHIFTING_FACTOR = 0.0
|
||||
BYTEDANCE_VAE_CONV_MEM_GIB = 0.5
|
||||
BYTEDANCE_VAE_NORM_MEM_GIB = 0.5
|
||||
BYTEDANCE_LOGVAR_CLAMP_MIN = -30.0 # video_vae_v3/modules/types.py:28.
|
||||
BYTEDANCE_LOGVAR_CLAMP_MAX = 20.0 # video_vae_v3/modules/types.py:28.
|
||||
BYTEDANCE_GN_CHUNKS_FP16 = 4 # causal_inflation_lib.py:351 (GroupNorm chunk count, fp16).
|
||||
BYTEDANCE_GN_CHUNKS_FP32 = 2 # causal_inflation_lib.py:351 (GroupNorm chunk count, fp32).
|
||||
BYTEDANCE_BLOCK_OUT_CHANNELS = (128, 256, 512, 512) # s8_c16_t4_inflation_sd3.yaml:7-11.
|
||||
BYTEDANCE_SLICING_SAMPLE_MIN = 4 # s8_c16_t4_inflation_sd3.yaml:22 (slicing_sample_min_size).
|
||||
BYTEDANCE_VAE_TEMPORAL_DOWNSAMPLE = 4 # infer.py:230 (temporal_downsample_factor); the 4n+1 factor.
|
||||
BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE = 8 # infer.py:231 (spatial_downsample_factor).
|
||||
BYTEDANCE_720P_REF_AREA = 45 * 80 # dit_v2/window.py:32 (720p reference area for window scaling).
|
||||
BYTEDANCE_MAX_TEMPORAL_WINDOW = 30 # dit_v2/window.py:35 (max temporal window frames).
|
||||
BYTEDANCE_ROPE_MAX_FREQ = 256 # dit_v2/rope.py:31 (pixel-RoPE max frequency).
|
||||
BYTEDANCE_SINUSOIDAL_DIM = 256 # dit_3b/nadit.py:120 (timestep sinusoidal embed dim).
|
||||
|
||||
ROPE_THETA = 10000 # RoPE base; Su et al., "RoFormer", arXiv:2104.09864.
|
||||
|
||||
CIELAB_DELTA = 6.0 / 29.0 # CIE 15 (delta).
|
||||
CIELAB_KAPPA = (29.0 / 3.0) ** 3 # CIE 15 (kappa).
|
||||
D65_WHITE_X = 0.95047 # CIE D65 standard illuminant Xn (Yn = 1).
|
||||
D65_WHITE_Z = 1.08883 # CIE D65 standard illuminant Zn.
|
||||
WAVELET_DECOMP_LEVELS = 5 # wavelet color-fix decomposition depth (GIMP/Krita; StableSR).
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -55,6 +55,7 @@ import comfy.ldm.pixeldit.model
|
||||
import comfy.ldm.pixeldit.pid
|
||||
import comfy.ldm.ace.model
|
||||
import comfy.ldm.omnigen.omnigen2
|
||||
import comfy.ldm.seedvr.model
|
||||
import comfy.ldm.boogu.model
|
||||
import comfy.ldm.qwen_image.model
|
||||
import comfy.ldm.ideogram4.model
|
||||
@@ -932,6 +933,17 @@ class HunyuanDiT(BaseModel):
|
||||
out['image_meta_size'] = comfy.conds.CONDRegular(torch.FloatTensor([[height, width, target_height, target_width, 0, 0]]))
|
||||
return out
|
||||
|
||||
class SeedVR2(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.seedvr.model.NaDiT)
|
||||
|
||||
def extra_conds(self, **kwargs):
|
||||
out = super().extra_conds(**kwargs)
|
||||
condition = kwargs.get("condition", None)
|
||||
if condition is not None:
|
||||
out["condition"] = comfy.conds.CONDRegular(condition)
|
||||
return out
|
||||
|
||||
class PixArt(BaseModel):
|
||||
def __init__(self, model_config, model_type=ModelType.EPS, device=None):
|
||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.pixart.pixartms.PixArtMS)
|
||||
|
||||
@@ -470,15 +470,46 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
|
||||
# PiD (Pixel Diffusion Decoder). Must check BEFORE plain PixelDiT_T2I.
|
||||
_lq_w_key = '{}lq_proj.latent_proj.0.weight'.format(key_prefix)
|
||||
if _lq_w_key in state_dict_keys:
|
||||
in_ch = int(state_dict[_lq_w_key].shape[1])
|
||||
latent_proj_in_channels = int(state_dict[_lq_w_key].shape[1])
|
||||
hidden_dim = int(state_dict[_lq_w_key].shape[0])
|
||||
_gate_prefix = '{}lq_proj.gate_modules.'.format(key_prefix)
|
||||
num_gates = len({k[len(_gate_prefix):].split('.')[0]
|
||||
for k in state_dict_keys if k.startswith(_gate_prefix)})
|
||||
pid_v1_5 = '{}lq_proj.pit_head.weight'.format(key_prefix) in state_dict_keys
|
||||
dit_config = {"image_model": "pid",
|
||||
"lq_latent_channels": in_ch,
|
||||
"latent_spatial_down_factor": 16 if in_ch >= 64 else 8}
|
||||
"lq_hidden_dim": hidden_dim}
|
||||
if num_gates > 0:
|
||||
dit_config["lq_interval"] = (14 + num_gates - 1) // num_gates
|
||||
if pid_v1_5:
|
||||
pid_v1_5_variants = {
|
||||
16: { # Flux and QwenImage
|
||||
"lq_latent_channels": 16,
|
||||
"latent_spatial_down_factor": 8,
|
||||
"lq_latent_unpatchify_factor": 1,
|
||||
},
|
||||
32: { # Flux2 after 2x latent unpatchify
|
||||
"lq_latent_channels": 128,
|
||||
"latent_spatial_down_factor": 16,
|
||||
"lq_latent_unpatchify_factor": 2,
|
||||
},
|
||||
}
|
||||
variant = pid_v1_5_variants.get(latent_proj_in_channels)
|
||||
if variant is None:
|
||||
raise ValueError(f"Unsupported PiD v1.5 latent projection with {latent_proj_in_channels} input channels")
|
||||
gate_weight = state_dict['{}lq_proj.gate_modules.0.content_proj.weight'.format(key_prefix)]
|
||||
dit_config.update(variant)
|
||||
dit_config.update({
|
||||
"lq_conv_padding_mode": "replicate",
|
||||
"lq_gate_per_token": gate_weight.shape[0] == 1,
|
||||
"pit_lq_inject": True,
|
||||
"rope_ref_h": 2048,
|
||||
"rope_ref_w": 2048,
|
||||
})
|
||||
else:
|
||||
dit_config.update({
|
||||
"lq_latent_channels": latent_proj_in_channels,
|
||||
"latent_spatial_down_factor": 16 if latent_proj_in_channels >= 64 else 8,
|
||||
})
|
||||
return dit_config
|
||||
|
||||
if '{}core.pixel_embedder.proj.weight'.format(key_prefix) in state_dict_keys: # PixelDiT T2I
|
||||
@@ -598,6 +629,44 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
|
||||
|
||||
return dit_config
|
||||
|
||||
seedvr2_7b_separate_key = "{}blocks.35.mlp.vid.proj_out.weight".format(key_prefix)
|
||||
if seedvr2_7b_separate_key in state_dict_keys and state_dict[seedvr2_7b_separate_key].shape[0] == 3072: # seedvr2 7b
|
||||
dit_config = {}
|
||||
dit_config["image_model"] = "seedvr2"
|
||||
dit_config["vid_dim"] = 3072
|
||||
dit_config["heads"] = 24
|
||||
dit_config["num_layers"] = 36
|
||||
# This checkpoint uses separate vid/txt MMModule keys in every block.
|
||||
dit_config["mm_layers"] = 36
|
||||
dit_config["norm_eps"] = 1e-5
|
||||
dit_config["rope_type"] = "rope3d"
|
||||
dit_config["rope_dim"] = 64
|
||||
dit_config["mlp_type"] = "normal"
|
||||
return dit_config
|
||||
if "{}blocks.35.mlp.all.proj_in_gate.weight".format(key_prefix) in state_dict_keys: # seedvr2 7b
|
||||
dit_config = {}
|
||||
dit_config["image_model"] = "seedvr2"
|
||||
dit_config["vid_dim"] = 3072
|
||||
dit_config["heads"] = 24
|
||||
dit_config["num_layers"] = 36
|
||||
# This checkpoint uses shared all.* MMModule keys after the initial blocks.
|
||||
dit_config["mm_layers"] = 10
|
||||
dit_config["norm_eps"] = 1e-5
|
||||
dit_config["rope_type"] = "rope3d"
|
||||
dit_config["rope_dim"] = 64
|
||||
dit_config["mlp_type"] = "swiglu"
|
||||
return dit_config
|
||||
if "{}blocks.31.mlp.all.proj_in_gate.weight".format(key_prefix) in state_dict_keys: # seedvr2 3b
|
||||
dit_config = {}
|
||||
dit_config["image_model"] = "seedvr2"
|
||||
dit_config["vid_dim"] = 2560
|
||||
dit_config["heads"] = 20
|
||||
dit_config["num_layers"] = 32
|
||||
dit_config["norm_eps"] = 1.0e-05
|
||||
dit_config["mlp_type"] = "swiglu"
|
||||
dit_config["vid_out_norm"] = True
|
||||
return dit_config
|
||||
|
||||
if '{}head.modulation'.format(key_prefix) in state_dict_keys: # Wan 2.1
|
||||
dit_config = {}
|
||||
dit_config["image_model"] = "wan2.1"
|
||||
@@ -1119,9 +1188,10 @@ def detect_unet_config(state_dict, key_prefix, metadata=None):
|
||||
|
||||
return unet_config
|
||||
|
||||
def model_config_from_unet_config(unet_config, state_dict=None):
|
||||
|
||||
def model_config_from_unet_config(unet_config, state_dict=None, unet_key_prefix=""):
|
||||
for model_config in comfy.supported_models.models:
|
||||
if model_config.matches(unet_config, state_dict):
|
||||
if model_config.matches(unet_config, state_dict, unet_key_prefix=unet_key_prefix):
|
||||
return model_config(unet_config)
|
||||
|
||||
logging.error("no match {}".format(unet_config))
|
||||
@@ -1131,7 +1201,7 @@ def model_config_from_unet(state_dict, unet_key_prefix, use_base_if_no_match=Fal
|
||||
unet_config = detect_unet_config(state_dict, unet_key_prefix, metadata=metadata)
|
||||
if unet_config is None:
|
||||
return None
|
||||
model_config = model_config_from_unet_config(unet_config, state_dict)
|
||||
model_config = model_config_from_unet_config(unet_config, state_dict, unet_key_prefix)
|
||||
if model_config is None and use_base_if_no_match:
|
||||
model_config = comfy.supported_models_base.BASE(unet_config)
|
||||
|
||||
|
||||
@@ -616,6 +616,8 @@ PIN_PRESSURE_HYSTERESIS = 256 * 1024 * 1024
|
||||
#Freeing registerables on pressure does imply a GPU sync, so go big on
|
||||
#the hysteresis so each expensive sync gives us back a good chunk.
|
||||
REGISTERABLE_PIN_HYSTERESIS = 2048 * 1024 * 1024
|
||||
WINDOWS_PIN_EVICTION_SWAP_PERCENT = 5.0
|
||||
WINDOWS_PIN_EVICTION_EMERGENCY_AVAILABLE = 512 * 1024 ** 2
|
||||
|
||||
def module_size(module):
|
||||
module_mem = 0
|
||||
@@ -642,6 +644,15 @@ def free_pins(size, evict_active=False):
|
||||
size -= freed
|
||||
return freed_total
|
||||
|
||||
def should_free_pins_for_ram_pressure(shortfall):
|
||||
if shortfall <= 0:
|
||||
return False
|
||||
if not WINDOWS:
|
||||
return True
|
||||
if psutil.virtual_memory().available < WINDOWS_PIN_EVICTION_EMERGENCY_AVAILABLE:
|
||||
return True
|
||||
return psutil.swap_memory().percent >= WINDOWS_PIN_EVICTION_SWAP_PERCENT
|
||||
|
||||
def ensure_pin_budget(size, evict_active=False):
|
||||
if args.high_ram:
|
||||
return True
|
||||
|
||||
+33
-7
@@ -1104,6 +1104,21 @@ def _load_quantized_module(module, super_load, state_dict, prefix, local_metadat
|
||||
scales["convrot_groupsize"] = int(
|
||||
layer_conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256))
|
||||
)
|
||||
elif module.quant_format == "convrot_w4a4":
|
||||
scale = pop_scale("weight_scale")
|
||||
if scale is None:
|
||||
raise ValueError(f"Missing ConvRot W4A4 weight scale for layer {layer_name}")
|
||||
params_conf = layer_conf.get("params", {})
|
||||
if not isinstance(params_conf, dict):
|
||||
params_conf = {}
|
||||
scales = {
|
||||
"scale": scale,
|
||||
"convrot_groupsize": int(
|
||||
layer_conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256))
|
||||
),
|
||||
"quant_group_size": 64,
|
||||
"linear_dtype": layer_conf.get("linear_dtype", params_conf.get("linear_dtype", "int4")),
|
||||
}
|
||||
else:
|
||||
raise ValueError(f"Unsupported quantization format: {module.quant_format}")
|
||||
|
||||
@@ -1150,6 +1165,11 @@ def _quantized_weight_state_dict(module, sd, prefix, extra_quant_conf=None, extr
|
||||
if module.quant_format == "int8_tensorwise" and getattr(params, "convrot", False):
|
||||
quant_conf["convrot"] = True
|
||||
quant_conf["convrot_groupsize"] = getattr(params, "convrot_groupsize", 256)
|
||||
elif module.quant_format == "convrot_w4a4":
|
||||
quant_conf["convrot_groupsize"] = getattr(params, "convrot_groupsize", 256)
|
||||
linear_dtype = getattr(params, "linear_dtype", "int4")
|
||||
if linear_dtype != "int4":
|
||||
quant_conf["linear_dtype"] = linear_dtype
|
||||
if extra_quant_conf:
|
||||
quant_conf.update(extra_quant_conf)
|
||||
sd[f"{prefix}comfy_quant"] = torch.tensor(list(json.dumps(quant_conf).encode("utf-8")), dtype=torch.uint8)
|
||||
@@ -1237,7 +1257,7 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
|
||||
run_every_op()
|
||||
|
||||
input_shape = input.shape
|
||||
reshaped_3d = False
|
||||
reshaped_nd = False
|
||||
#If cast needs to apply lora, it should be done in the compute dtype
|
||||
compute_dtype = input.dtype
|
||||
|
||||
@@ -1274,12 +1294,12 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
|
||||
# Inference path (unchanged)
|
||||
if _use_quantized and quantize_input:
|
||||
|
||||
# Reshape 3D tensors to 2D for quantization (needed for NVFP4 and others)
|
||||
input_reshaped = input.reshape(-1, input_shape[2]) if input.ndim == 3 else input
|
||||
# Reshape >=3D tensors to 2D for quantization (needed for NVFP4 and others)
|
||||
input_reshaped = input.reshape(-1, input_shape[-1]) if input.ndim >= 3 else input
|
||||
|
||||
# Fall back to non-quantized for non-2D tensors
|
||||
if input_reshaped.ndim == 2:
|
||||
reshaped_3d = input.ndim == 3
|
||||
reshaped_nd = input.ndim >= 3
|
||||
# dtype is now implicit in the layout class
|
||||
scale = getattr(self, 'input_scale', None)
|
||||
if scale is not None:
|
||||
@@ -1294,9 +1314,9 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
|
||||
weight_only_quant=weight_only_quant,
|
||||
)
|
||||
|
||||
# Reshape output back to 3D if input was 3D
|
||||
if reshaped_3d:
|
||||
output = output.reshape((input_shape[0], input_shape[1], self.weight.shape[0]))
|
||||
# Reshape output back to original rank if input was >2D
|
||||
if reshaped_nd:
|
||||
output = output.reshape((*input_shape[:-1], self.weight.shape[0]))
|
||||
|
||||
return output
|
||||
|
||||
@@ -1430,6 +1450,12 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec
|
||||
}
|
||||
if hasattr(params, "block_scale"): # NVFP4
|
||||
kwargs["block_scale"] = params.block_scale[i]
|
||||
if hasattr(params, "quant_group_size"):
|
||||
kwargs["quant_group_size"] = params.quant_group_size
|
||||
if hasattr(params, "convrot_groupsize"):
|
||||
kwargs["convrot_groupsize"] = params.convrot_groupsize
|
||||
if hasattr(params, "linear_dtype"):
|
||||
kwargs["linear_dtype"] = params.linear_dtype
|
||||
return QuantizedTensor(weight._qdata[i], weight._layout_cls, type(params)(**kwargs))
|
||||
|
||||
def state_dict(self, *args, destination=None, prefix="", **kwargs):
|
||||
|
||||
+44
-2
@@ -3,6 +3,22 @@ import logging
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
|
||||
def _rocm_kitchen_arch_supported():
|
||||
"""comfy-kitchen's INT8 Triton kernels compile tl.dot to matrix-core instructions.
|
||||
RDNA3/3.5/4 (gfx11xx/gfx12xx) have WMMA and CDNA (gfx9xx) has MFMA; RDNA1/RDNA2
|
||||
(gfx10xx) have neither, so the INT8 path hangs the GPU there. Gates the automatic
|
||||
ROCm default so those cards stay on the eager fallback (an explicit
|
||||
--enable-triton-backend still forces it on any arch)."""
|
||||
try:
|
||||
arch = torch.cuda.get_device_properties(torch.cuda.current_device()).gcnArchName.split(":")[0]
|
||||
except Exception:
|
||||
return False
|
||||
if arch.startswith(("gfx11", "gfx12")):
|
||||
return True
|
||||
return arch in ("gfx908", "gfx90a", "gfx940", "gfx941", "gfx942", "gfx950")
|
||||
|
||||
|
||||
try:
|
||||
import comfy_kitchen as ck
|
||||
from comfy_kitchen.tensor import (
|
||||
@@ -10,6 +26,7 @@ try:
|
||||
QuantizedLayout,
|
||||
TensorCoreFP8Layout as _CKFp8Layout,
|
||||
TensorCoreNVFP4Layout as _CKNvfp4Layout,
|
||||
TensorCoreConvRotW4A4Layout as _CKTensorCoreConvRotW4A4Layout,
|
||||
TensorWiseINT8Layout as _CKTensorWiseINT8Layout,
|
||||
register_layout_op,
|
||||
register_layout_class,
|
||||
@@ -24,10 +41,22 @@ try:
|
||||
ck.registry.disable("cuda")
|
||||
logging.warning("WARNING: You need pytorch with cu130 or higher to use optimized CUDA operations.")
|
||||
|
||||
if args.enable_triton_backend:
|
||||
# On ROCm/AMD the CUDA backend is unavailable, so Triton is the only accelerated
|
||||
# comfy-kitchen backend. Enable it by default there, but only on Triton >= 3.7 AND a
|
||||
# matrix-core GPU (RDNA3+ WMMA gfx11xx/gfx12xx, CDNA MFMA gfx9xx). RDNA1/RDNA2
|
||||
# (gfx10xx) have no WMMA -> the INT8 tl.dot path hangs the GPU, so they stay eager.
|
||||
# older Triton lacks libdevice.rint on the HIP backend and hard-crashes the INT8 path.
|
||||
if args.disable_triton_backend:
|
||||
ck.registry.disable("triton")
|
||||
elif args.enable_triton_backend: # or (torch.version.hip is not None and _rocm_kitchen_arch_supported()):
|
||||
try:
|
||||
import triton
|
||||
logging.info("Found triton %s. Enabling comfy-kitchen triton backend.", triton.__version__)
|
||||
triton_version = tuple(int(v) for v in triton.__version__.split(".")[:2])
|
||||
if args.enable_triton_backend or triton_version >= (3, 7):
|
||||
logging.info("Found triton %s. Enabling comfy-kitchen triton backend.", triton.__version__)
|
||||
else:
|
||||
logging.info("Triton %s is too old for the ROCm INT8 path (needs >= 3.7); comfy-kitchen triton backend disabled.", triton.__version__)
|
||||
ck.registry.disable("triton")
|
||||
except ImportError as e:
|
||||
logging.error(f"Failed to import triton, Error: {e}, the comfy-kitchen triton backend will not be available.")
|
||||
ck.registry.disable("triton")
|
||||
@@ -51,6 +80,9 @@ except ImportError as e:
|
||||
class _CKTensorWiseINT8Layout:
|
||||
pass
|
||||
|
||||
class _CKTensorCoreConvRotW4A4Layout:
|
||||
pass
|
||||
|
||||
def register_layout_class(name, cls):
|
||||
pass
|
||||
|
||||
@@ -179,6 +211,7 @@ class TensorCoreFP8E5M2Layout(_TensorCoreFP8LayoutBase):
|
||||
# Backward compatibility alias - default to E4M3
|
||||
TensorCoreFP8Layout = TensorCoreFP8E4M3Layout
|
||||
TensorWiseINT8Layout = _CKTensorWiseINT8Layout
|
||||
TensorCoreConvRotW4A4Layout = _CKTensorCoreConvRotW4A4Layout
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
@@ -190,6 +223,7 @@ register_layout_class("TensorCoreFP8E4M3Layout", TensorCoreFP8E4M3Layout)
|
||||
register_layout_class("TensorCoreFP8E5M2Layout", TensorCoreFP8E5M2Layout)
|
||||
register_layout_class("TensorCoreNVFP4Layout", TensorCoreNVFP4Layout)
|
||||
register_layout_class("TensorWiseINT8Layout", _CKTensorWiseINT8Layout)
|
||||
register_layout_class("TensorCoreConvRotW4A4Layout", _CKTensorCoreConvRotW4A4Layout)
|
||||
if _CK_MXFP8_AVAILABLE:
|
||||
register_layout_class("TensorCoreMXFP8Layout", TensorCoreMXFP8Layout)
|
||||
|
||||
@@ -227,6 +261,13 @@ QUANT_ALGOS["int8_tensorwise"] = {
|
||||
"quantize_input": False,
|
||||
}
|
||||
|
||||
QUANT_ALGOS["convrot_w4a4"] = {
|
||||
"storage_t": torch.int8,
|
||||
"parameters": {"weight_scale"},
|
||||
"comfy_tensor_layout": "TensorCoreConvRotW4A4Layout",
|
||||
"quantize_input": False,
|
||||
}
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Re-exports for backward compatibility
|
||||
@@ -239,6 +280,7 @@ __all__ = [
|
||||
"TensorCoreFP8E4M3Layout",
|
||||
"TensorCoreFP8E5M2Layout",
|
||||
"TensorCoreNVFP4Layout",
|
||||
"TensorCoreConvRotW4A4Layout",
|
||||
"TensorWiseINT8Layout",
|
||||
"QUANT_ALGOS",
|
||||
"register_layout_op",
|
||||
|
||||
+83
-19
@@ -16,6 +16,7 @@ import comfy.ldm.cosmos.vae
|
||||
import comfy.ldm.wan.vae
|
||||
import comfy.ldm.wan.vae2_2
|
||||
import comfy.ldm.hunyuan3d.vae
|
||||
import comfy.ldm.seedvr.vae
|
||||
import comfy.ldm.triposplat.vae
|
||||
import comfy.ldm.ace.vae.music_dcae_pipeline
|
||||
import comfy.ldm.cogvideo.vae
|
||||
@@ -473,7 +474,8 @@ class CLIP:
|
||||
|
||||
class VAE:
|
||||
def __init__(self, sd=None, device=None, config=None, dtype=None, metadata=None):
|
||||
if 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys(): #diffusers format
|
||||
is_seedvr2_vae = "decoder.up_blocks.2.upsamplers.0.upscale_conv.weight" in sd
|
||||
if not is_seedvr2_vae and 'decoder.up_blocks.0.resnets.0.norm1.weight' in sd.keys(): #diffusers format
|
||||
sd = diffusers_convert.convert_vae_state_dict(sd)
|
||||
|
||||
if model_management.is_amd():
|
||||
@@ -500,6 +502,8 @@ class VAE:
|
||||
self.upscale_index_formula = None
|
||||
self.extra_1d_channel = None
|
||||
self.crop_input = True
|
||||
self.handles_tiling = False
|
||||
self.format_encoded = None
|
||||
|
||||
self.audio_sample_rate = 44100
|
||||
|
||||
@@ -546,6 +550,22 @@ class VAE:
|
||||
self.first_stage_model = StageC_coder()
|
||||
self.downscale_ratio = 32
|
||||
self.latent_channels = 16
|
||||
elif "decoder.up_blocks.2.upsamplers.0.upscale_conv.weight" in sd: # seedvr2
|
||||
self.first_stage_model = comfy.ldm.seedvr.vae.VideoAutoencoderKLWrapper()
|
||||
self.latent_channels = comfy.ldm.seedvr.vae.SEEDVR2_LATENT_CHANNELS
|
||||
self.latent_dim = 3
|
||||
self.disable_offload = True
|
||||
self.memory_used_decode = lambda shape, dtype: self.first_stage_model.comfy_memory_used_decode(shape)
|
||||
self.memory_used_encode = lambda shape, dtype: (max(shape[2], 5) * shape[3] * shape[4] * 64) * model_management.dtype_size(dtype)
|
||||
self.working_dtypes = [torch.float16, torch.bfloat16, torch.float32]
|
||||
self.handles_tiling = True
|
||||
self.format_encoded = self.first_stage_model.comfy_format_encoded
|
||||
self.downscale_ratio = (lambda a: max(0, math.floor((a + 3) / 4)), 8, 8)
|
||||
self.downscale_index_formula = (4, 8, 8)
|
||||
self.upscale_ratio = (lambda a: max(0, a * 4 - 3), 8, 8)
|
||||
self.upscale_index_formula = (4, 8, 8)
|
||||
self.process_input = lambda image: image * 2.0 - 1.0
|
||||
self.crop_input = False
|
||||
elif "decoder.conv_in.weight" in sd:
|
||||
if sd['decoder.conv_in.weight'].shape[1] == 64:
|
||||
ddconfig = {"block_out_channels": [128, 256, 512, 512, 1024, 1024], "in_channels": 3, "out_channels": 3, "num_res_blocks": 2, "ffactor_spatial": 32, "downsample_match_channel": True, "upsample_match_channel": True}
|
||||
@@ -1012,6 +1032,10 @@ class VAE:
|
||||
decode_fn = lambda a: self.first_stage_model.decode(a.to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype())
|
||||
return self.process_output(comfy.utils.tiled_scale_multidim(samples, decode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.upscale_ratio, out_channels=self.output_channels, index_formulas=self.upscale_index_formula, output_device=self.output_device))
|
||||
|
||||
def _decode_tiled_owned(self, samples, **kwargs):
|
||||
out = self.first_stage_model.decode_tiled(samples.to(self.vae_dtype).to(self.device), **kwargs)
|
||||
return self.process_output(out.to(device=self.output_device, dtype=self.vae_output_dtype(), copy=True))
|
||||
|
||||
def encode_tiled_(self, pixel_samples, tile_x=512, tile_y=512, overlap = 64):
|
||||
steps = pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x, tile_y, overlap)
|
||||
steps += pixel_samples.shape[0] * comfy.utils.get_tiled_scale_steps(pixel_samples.shape[3], pixel_samples.shape[2], tile_x // 2, tile_y * 2, overlap)
|
||||
@@ -1048,6 +1072,25 @@ class VAE:
|
||||
encode_fn = lambda a: self.first_stage_model.encode((self.process_input(a)).to(self.vae_dtype).to(self.device)).to(dtype=self.vae_output_dtype())
|
||||
return comfy.utils.tiled_scale_multidim(samples, encode_fn, tile=(tile_t, tile_x, tile_y), overlap=overlap, upscale_amount=self.downscale_ratio, out_channels=self.latent_channels, downscale=True, index_formulas=self.downscale_index_formula, output_device=self.output_device)
|
||||
|
||||
def _encode_tiled_owned(self, pixel_samples, **kwargs):
|
||||
x = self.process_input(pixel_samples).to(self.vae_dtype).to(self.device)
|
||||
out = self.first_stage_model.encode_tiled(x, **kwargs)
|
||||
return out.to(device=self.output_device, dtype=self.vae_output_dtype())
|
||||
|
||||
def _owned_tiled_args(self, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None):
|
||||
args = {}
|
||||
if tile_x is not None:
|
||||
args["tile_x"] = tile_x
|
||||
if tile_y is not None:
|
||||
args["tile_y"] = tile_y
|
||||
if overlap is not None:
|
||||
args["overlap"] = overlap
|
||||
if tile_t is not None:
|
||||
args["tile_t"] = tile_t
|
||||
if overlap_t is not None:
|
||||
args["overlap_t"] = overlap_t
|
||||
return args
|
||||
|
||||
def decode(self, samples_in, vae_options={}):
|
||||
self.throw_exception_if_invalid()
|
||||
pixel_samples = None
|
||||
@@ -1095,11 +1138,19 @@ class VAE:
|
||||
if dims == 1 or self.extra_1d_channel is not None:
|
||||
pixel_samples = self.decode_tiled_1d(samples_in)
|
||||
elif dims == 2:
|
||||
pixel_samples = self.decode_tiled_(samples_in)
|
||||
if self.handles_tiling:
|
||||
tile = 256 // self.spacial_compression_decode()
|
||||
overlap = tile // 4
|
||||
pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap)
|
||||
else:
|
||||
pixel_samples = self.decode_tiled_(samples_in)
|
||||
elif dims == 3:
|
||||
tile = 256 // self.spacial_compression_decode()
|
||||
overlap = tile // 4
|
||||
pixel_samples = self.decode_tiled_3d(samples_in, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap))
|
||||
if self.handles_tiling:
|
||||
pixel_samples = self._decode_tiled_owned(samples_in, tile_x=tile, tile_y=tile, overlap=overlap)
|
||||
else:
|
||||
pixel_samples = self.decode_tiled_3d(samples_in, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap))
|
||||
|
||||
pixel_samples = pixel_samples.to(self.output_device).movedim(1,-1)
|
||||
return pixel_samples
|
||||
@@ -1118,7 +1169,9 @@ class VAE:
|
||||
args["overlap"] = overlap
|
||||
|
||||
with model_management.cuda_device_context(self.device):
|
||||
if dims == 1 or self.extra_1d_channel is not None:
|
||||
if self.handles_tiling and dims in (2, 3):
|
||||
output = self._decode_tiled_owned(samples, **self._owned_tiled_args(tile_x, tile_y, overlap, tile_t, overlap_t))
|
||||
elif dims == 1 or self.extra_1d_channel is not None:
|
||||
args.pop("tile_y")
|
||||
output = self.decode_tiled_1d(samples, **args)
|
||||
elif dims == 2:
|
||||
@@ -1179,12 +1232,17 @@ class VAE:
|
||||
if self.latent_dim == 3:
|
||||
tile = 256
|
||||
overlap = tile // 4
|
||||
samples = self.encode_tiled_3d(pixel_samples, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap))
|
||||
if self.handles_tiling:
|
||||
samples = self._encode_tiled_owned(pixel_samples, tile_x=tile, tile_y=tile, overlap=overlap)
|
||||
else:
|
||||
samples = self.encode_tiled_3d(pixel_samples, tile_x=tile, tile_y=tile, overlap=(1, overlap, overlap))
|
||||
elif self.latent_dim == 1 or self.extra_1d_channel is not None:
|
||||
samples = self.encode_tiled_1d(pixel_samples)
|
||||
else:
|
||||
samples = self.encode_tiled_(pixel_samples)
|
||||
|
||||
if self.format_encoded is not None:
|
||||
samples = self.format_encoded(samples)
|
||||
return samples
|
||||
|
||||
def encode_tiled(self, pixel_samples, tile_x=None, tile_y=None, overlap=None, tile_t=None, overlap_t=None):
|
||||
@@ -1192,7 +1250,7 @@ class VAE:
|
||||
pixel_samples = self.vae_encode_crop_pixels(pixel_samples)
|
||||
dims = self.latent_dim
|
||||
pixel_samples = pixel_samples.movedim(-1, 1)
|
||||
if dims == 3:
|
||||
if dims == 3 and pixel_samples.ndim < 5:
|
||||
if not self.not_video:
|
||||
pixel_samples = pixel_samples.movedim(1, 0).unsqueeze(0)
|
||||
else:
|
||||
@@ -1216,21 +1274,27 @@ class VAE:
|
||||
elif dims == 2:
|
||||
samples = self.encode_tiled_(pixel_samples, **args)
|
||||
elif dims == 3:
|
||||
if tile_t is not None:
|
||||
tile_t_latent = max(2, self.downscale_ratio[0](tile_t))
|
||||
if self.handles_tiling:
|
||||
samples = self._encode_tiled_owned(pixel_samples, **self._owned_tiled_args(tile_x, tile_y, overlap, tile_t, overlap_t))
|
||||
else:
|
||||
tile_t_latent = 9999
|
||||
args["tile_t"] = self.upscale_ratio[0](tile_t_latent)
|
||||
if tile_t is not None:
|
||||
tile_t_latent = max(2, self.downscale_ratio[0](tile_t))
|
||||
else:
|
||||
tile_t_latent = 9999
|
||||
args["tile_t"] = self.upscale_ratio[0](tile_t_latent)
|
||||
|
||||
if overlap_t is None:
|
||||
args["overlap"] = (1, overlap, overlap)
|
||||
else:
|
||||
args["overlap"] = (self.upscale_ratio[0](max(1, min(tile_t_latent // 2, self.downscale_ratio[0](overlap_t)))), overlap, overlap)
|
||||
maximum = pixel_samples.shape[2]
|
||||
maximum = self.upscale_ratio[0](self.downscale_ratio[0](maximum))
|
||||
spatial_overlap = overlap if overlap is not None else 64
|
||||
if overlap_t is None:
|
||||
args["overlap"] = (1, spatial_overlap, spatial_overlap)
|
||||
else:
|
||||
args["overlap"] = (self.upscale_ratio[0](max(1, min(tile_t_latent // 2, self.downscale_ratio[0](overlap_t)))), spatial_overlap, spatial_overlap)
|
||||
maximum = pixel_samples.shape[2]
|
||||
maximum = self.upscale_ratio[0](self.downscale_ratio[0](maximum))
|
||||
|
||||
samples = self.encode_tiled_3d(pixel_samples[:,:,:maximum], **args)
|
||||
samples = self.encode_tiled_3d(pixel_samples[:,:,:maximum], **args)
|
||||
|
||||
if self.format_encoded is not None:
|
||||
samples = self.format_encoded(samples)
|
||||
return samples
|
||||
|
||||
def get_sd(self):
|
||||
@@ -1898,7 +1962,7 @@ def load_state_dict_guess_config(sd, output_vae=True, output_clip=True, output_c
|
||||
manual_cast_dtype = model_management.unet_manual_cast(None, load_device, model_config.supported_inference_dtypes)
|
||||
else:
|
||||
manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device, model_config.supported_inference_dtypes)
|
||||
model_config.set_inference_dtype(unet_dtype, manual_cast_dtype)
|
||||
model_config.set_inference_dtype(unet_dtype, manual_cast_dtype, device=load_device)
|
||||
|
||||
if model_config.clip_vision_prefix is not None:
|
||||
if output_clipvision:
|
||||
@@ -2039,7 +2103,7 @@ def load_diffusion_model_state_dict(sd, model_options={}, metadata=None, disable
|
||||
manual_cast_dtype = model_management.unet_manual_cast(None, load_device, model_config.supported_inference_dtypes)
|
||||
else:
|
||||
manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device, model_config.supported_inference_dtypes)
|
||||
model_config.set_inference_dtype(unet_dtype, manual_cast_dtype)
|
||||
model_config.set_inference_dtype(unet_dtype, manual_cast_dtype, device=load_device)
|
||||
|
||||
if custom_operations is not None:
|
||||
model_config.custom_operations = custom_operations
|
||||
|
||||
@@ -1685,6 +1685,40 @@ class Chroma(supported_models_base.BASE):
|
||||
t5_detect = comfy.text_encoders.sd3_clip.t5_xxl_detect(state_dict, "{}t5xxl.transformer.".format(pref))
|
||||
return supported_models_base.ClipTarget(comfy.text_encoders.pixart_t5.PixArtTokenizer, comfy.text_encoders.pixart_t5.pixart_te(**t5_detect))
|
||||
|
||||
class SeedVR2(supported_models_base.BASE):
|
||||
unet_config = {
|
||||
"image_model": "seedvr2"
|
||||
}
|
||||
unet_extra_config = {}
|
||||
required_keys = {
|
||||
"{}positive_conditioning",
|
||||
"{}negative_conditioning",
|
||||
}
|
||||
latent_format = comfy.latent_formats.SeedVR2
|
||||
|
||||
vae_key_prefix = ["vae."]
|
||||
text_encoder_key_prefix = ["text_encoders."]
|
||||
supported_inference_dtypes = [torch.bfloat16, torch.float16, torch.float32]
|
||||
sampling_settings = {
|
||||
"shift": 1.0,
|
||||
}
|
||||
|
||||
def set_inference_dtype(self, dtype, manual_cast_dtype, device=None):
|
||||
if (
|
||||
dtype == torch.float16
|
||||
and manual_cast_dtype is None
|
||||
and comfy.model_management.should_use_bf16(device)
|
||||
):
|
||||
manual_cast_dtype = torch.bfloat16
|
||||
super().set_inference_dtype(dtype, manual_cast_dtype, device=device)
|
||||
|
||||
def get_model(self, state_dict, prefix="", device=None):
|
||||
out = model_base.SeedVR2(self, device=device)
|
||||
return out
|
||||
|
||||
def clip_target(self, state_dict={}):
|
||||
return None
|
||||
|
||||
class ChromaRadiance(Chroma):
|
||||
unet_config = {
|
||||
"image_model": "chroma_radiance",
|
||||
@@ -2348,6 +2382,7 @@ models = [
|
||||
HiDream,
|
||||
HiDreamO1,
|
||||
Chroma,
|
||||
SeedVR2,
|
||||
ChromaRadiance,
|
||||
ACEStep,
|
||||
ACEStep15,
|
||||
|
||||
@@ -54,13 +54,13 @@ class BASE:
|
||||
optimizations = {"fp8": False}
|
||||
|
||||
@classmethod
|
||||
def matches(s, unet_config, state_dict=None):
|
||||
def matches(s, unet_config, state_dict=None, unet_key_prefix=""):
|
||||
for k in s.unet_config:
|
||||
if k not in unet_config or s.unet_config[k] != unet_config[k]:
|
||||
return False
|
||||
if state_dict is not None:
|
||||
for k in s.required_keys:
|
||||
if k not in state_dict:
|
||||
if k.format(unet_key_prefix) not in state_dict:
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -115,7 +115,7 @@ class BASE:
|
||||
replace_prefix = {"": self.vae_key_prefix[0]}
|
||||
return utils.state_dict_prefix_replace(state_dict, replace_prefix)
|
||||
|
||||
def set_inference_dtype(self, dtype, manual_cast_dtype):
|
||||
def set_inference_dtype(self, dtype, manual_cast_dtype, device=None):
|
||||
self.unet_config['dtype'] = dtype
|
||||
self.manual_cast_dtype = manual_cast_dtype
|
||||
|
||||
|
||||
@@ -1088,7 +1088,7 @@ class Gemma4_Tokenizer():
|
||||
h, w = samples.shape[2], samples.shape[3]
|
||||
patch_size = 16
|
||||
pooling_k = 3
|
||||
max_soft_tokens = 70 if is_video else 280 # video uses smaller token budget per frame
|
||||
max_soft_tokens = kwargs.get("max_soft_tokens", 70 if is_video else 280)
|
||||
max_patches = max_soft_tokens * pooling_k * pooling_k
|
||||
target_px = max_patches * patch_size * patch_size
|
||||
factor = (target_px / (h * w)) ** 0.5
|
||||
|
||||
@@ -90,6 +90,27 @@ class Qwen3VL(BaseLlama, BaseQwen3, BaseGenerate, torch.nn.Module):
|
||||
deepstack = [torch.cat([deepstack[i], ds[i]], dim=0) for i in range(len(ds))]
|
||||
return position_ids, visual_pos_masks, deepstack
|
||||
|
||||
def forward(self, input_ids, attention_mask=None, embeds=None, num_tokens=None, intermediate_output=None, final_layer_norm_intermediate=True, dtype=None, embeds_info=[], **kwargs):
|
||||
position_ids = kwargs.pop("position_ids", None)
|
||||
visual_pos_masks = kwargs.pop("visual_pos_masks", None)
|
||||
deepstack_embeds = kwargs.pop("deepstack_embeds", None)
|
||||
if embeds is not None and position_ids is None:
|
||||
position_ids, visual_pos_masks, deepstack_embeds = self.build_image_inputs(embeds, embeds_info)
|
||||
return self.model(
|
||||
input_ids,
|
||||
attention_mask=attention_mask,
|
||||
embeds=embeds,
|
||||
num_tokens=num_tokens,
|
||||
intermediate_output=intermediate_output,
|
||||
final_layer_norm_intermediate=final_layer_norm_intermediate,
|
||||
dtype=dtype,
|
||||
position_ids=position_ids,
|
||||
embeds_info=embeds_info,
|
||||
visual_pos_masks=visual_pos_masks,
|
||||
deepstack_embeds=deepstack_embeds,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
def _make_qwen3vl_model(model_type):
|
||||
class Qwen3VL_(Qwen3VL):
|
||||
|
||||
@@ -100,6 +100,8 @@ def _parse_cli_feature_flags() -> dict[str, Any]:
|
||||
# Default server capabilities
|
||||
_CORE_FEATURE_FLAGS: dict[str, Any] = {
|
||||
"supports_preview_metadata": True,
|
||||
"supports_node_failure_policy": True,
|
||||
"supports_model_type_tags": True,
|
||||
"max_upload_size": args.max_upload_size * 1024 * 1024, # Convert MB to bytes
|
||||
"extension": {"manager": {"supports_v4": True}},
|
||||
"node_replacements": True,
|
||||
|
||||
@@ -17,6 +17,10 @@ class Seedream4Options(BaseModel):
|
||||
max_images: int = Field(15)
|
||||
|
||||
|
||||
class Seedream5OptimizePromptOptions(BaseModel):
|
||||
thinking: Literal["auto", "enabled", "disabled"] = Field(...)
|
||||
|
||||
|
||||
class Seedream4TaskCreationRequest(BaseModel):
|
||||
model: str = Field(...)
|
||||
prompt: str = Field(...)
|
||||
@@ -28,6 +32,7 @@ class Seedream4TaskCreationRequest(BaseModel):
|
||||
sequential_image_generation_options: Seedream4Options | None = Field(Seedream4Options(max_images=15))
|
||||
watermark: bool = Field(False)
|
||||
output_format: str | None = None
|
||||
optimize_prompt_options: Seedream5OptimizePromptOptions | None = None
|
||||
|
||||
|
||||
class ImageTaskCreationResponse(BaseModel):
|
||||
|
||||
@@ -77,6 +77,7 @@ class To3DUVTaskRequest(BaseModel):
|
||||
|
||||
class To3DPartTaskRequest(BaseModel):
|
||||
File: TaskFile3DInput = Field(...)
|
||||
EnableStagedGeneration: bool | None = Field(None)
|
||||
|
||||
|
||||
class TextureEditImageInfo(BaseModel):
|
||||
|
||||
@@ -34,6 +34,7 @@ from comfy_api_nodes.apis.bytedance import (
|
||||
SeedanceVirtualLibraryCreateAssetRequest,
|
||||
Seedream4Options,
|
||||
Seedream4TaskCreationRequest,
|
||||
Seedream5OptimizePromptOptions,
|
||||
TaskAudioContent,
|
||||
TaskAudioContentUrl,
|
||||
TaskCreationResponse,
|
||||
@@ -875,6 +876,17 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode):
|
||||
tooltip='Whether to add an "AI generated" watermark to the image.',
|
||||
advanced=True,
|
||||
),
|
||||
IO.Boolean.Input(
|
||||
"thinking",
|
||||
default=True,
|
||||
tooltip=(
|
||||
"Enable the model's prompt-optimization reasoning ('thinking') for better adherence. "
|
||||
"Can substantially increase generation time — notably on Seedream 5.0 Pro. "
|
||||
"Can only be disabled for text-to-image (not when reference images are provided)."
|
||||
),
|
||||
optional=True,
|
||||
advanced=True,
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
IO.Image.Output(),
|
||||
@@ -920,6 +932,7 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode):
|
||||
model: dict,
|
||||
seed: int = 0,
|
||||
watermark: bool = False,
|
||||
thinking: bool = True,
|
||||
) -> IO.NodeOutput:
|
||||
validate_string(prompt, strip_whitespace=True, min_length=1)
|
||||
model_id = SEEDREAM_MODELS[model["model"]]
|
||||
@@ -979,6 +992,10 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode):
|
||||
raise ValueError(
|
||||
"The maximum number of generated images plus the number of reference images cannot exceed 15."
|
||||
)
|
||||
if not thinking and n_input_images > 0:
|
||||
raise ValueError(
|
||||
"'thinking' can only be disabled for text-to-image; enable it when using reference images."
|
||||
)
|
||||
|
||||
reference_images_urls: list[str] = []
|
||||
if image_tensors:
|
||||
@@ -992,6 +1009,9 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode):
|
||||
wait_label="Uploading reference images",
|
||||
)
|
||||
|
||||
optimize_prompt_options = None
|
||||
if n_input_images == 0:
|
||||
optimize_prompt_options = Seedream5OptimizePromptOptions(thinking="enabled" if thinking else "disabled")
|
||||
response = await sync_op(
|
||||
cls,
|
||||
ApiEndpoint(path=BYTEPLUS_IMAGE_ENDPOINT, method="POST"),
|
||||
@@ -1005,6 +1025,7 @@ class ByteDanceSeedreamNodeV2(IO.ComfyNode):
|
||||
sequential_image_generation=None if is_pro else sequential_image_generation,
|
||||
sequential_image_generation_options=None if is_pro else Seedream4Options(max_images=max_images),
|
||||
watermark=watermark,
|
||||
optimize_prompt_options=optimize_prompt_options,
|
||||
),
|
||||
)
|
||||
if len(response.data) == 1:
|
||||
|
||||
@@ -1133,7 +1133,9 @@ 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-preview"
|
||||
model = "gemini-3.1-flash-image"
|
||||
elif model == "gemini-3-pro-image-preview":
|
||||
model = "gemini-3-pro-image"
|
||||
|
||||
parts: list[GeminiPart] = [GeminiPart(text=prompt)]
|
||||
if images is not None:
|
||||
@@ -1507,7 +1509,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-preview"
|
||||
model_id = "gemini-3.1-flash-image"
|
||||
elif model_choice == "Nano Banana 2 Lite":
|
||||
model_id = "gemini-3.1-flash-lite-image"
|
||||
else:
|
||||
|
||||
@@ -642,6 +642,7 @@ class Tencent3DPartNode(IO.ComfyNode):
|
||||
response_model=To3DProTaskCreateResponse,
|
||||
data=To3DPartTaskRequest(
|
||||
File=TaskFile3DInput(Type=file_format.upper(), Url=model_url),
|
||||
EnableStagedGeneration=True,
|
||||
),
|
||||
is_rate_limited=_is_tencent_rate_limited,
|
||||
)
|
||||
|
||||
@@ -11,9 +11,11 @@ from io import BytesIO
|
||||
from yarl import URL
|
||||
|
||||
from comfy.cli_args import args
|
||||
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 comfyui_version import __version__ as comfyui_version
|
||||
|
||||
from .common_exceptions import ProcessingInterrupted
|
||||
|
||||
@@ -59,11 +61,12 @@ def get_comfy_api_headers(node_cls: type[IO.ComfyNode]) -> dict[str, str]:
|
||||
**get_auth_header(node_cls),
|
||||
"Comfy-Env": get_deploy_environment(),
|
||||
"Comfy-Usage-Source": get_usage_source(node_cls),
|
||||
"Comfy-Core-Version": comfyui_version,
|
||||
}
|
||||
|
||||
|
||||
def default_base_url() -> str:
|
||||
return getattr(args, "comfy_api_base", "https://api.comfy.org")
|
||||
return normalize_comfy_api_base(getattr(args, "comfy_api_base", "https://api.comfy.org"))
|
||||
|
||||
|
||||
async def sleep_with_interrupt(
|
||||
|
||||
@@ -503,6 +503,8 @@ RAM_CACHE_DEFAULT_RAM_USAGE = 0.05
|
||||
|
||||
RAM_CACHE_OLD_WORKFLOW_OOM_MULTIPLIER = 1.3
|
||||
|
||||
RAM_CACHE_LARGE_INTERMEDIATE = 512 * 1024 ** 2
|
||||
|
||||
|
||||
def all_outputs_dynamic(outputs):
|
||||
if outputs is None:
|
||||
@@ -517,7 +519,6 @@ def all_outputs_dynamic(outputs):
|
||||
|
||||
return True
|
||||
|
||||
|
||||
class RAMPressureCache(LRUCache):
|
||||
|
||||
def __init__(self, key_class, enable_providers=False):
|
||||
@@ -539,9 +540,9 @@ class RAMPressureCache(LRUCache):
|
||||
self.timestamps[self.cache_key_set.get_data_key(node_id)] = time.time()
|
||||
super().set_local(node_id, value)
|
||||
|
||||
def ram_release(self, target, free_active=False):
|
||||
def ram_release(self, target, free_active=False, min_entry_size=0):
|
||||
if psutil.virtual_memory().available >= target:
|
||||
return
|
||||
return 0
|
||||
|
||||
clean_list = []
|
||||
|
||||
@@ -555,8 +556,9 @@ class RAMPressureCache(LRUCache):
|
||||
oom_score = RAM_CACHE_OLD_WORKFLOW_OOM_MULTIPLIER ** (self.generation - self.used_generation[key])
|
||||
|
||||
ram_usage = RAM_CACHE_DEFAULT_RAM_USAGE
|
||||
oom_ram_usage = ram_usage
|
||||
def scan_list_for_ram_usage(outputs):
|
||||
nonlocal ram_usage
|
||||
nonlocal ram_usage, oom_ram_usage
|
||||
if outputs is None:
|
||||
return
|
||||
for output in outputs:
|
||||
@@ -564,19 +566,26 @@ class RAMPressureCache(LRUCache):
|
||||
scan_list_for_ram_usage(output)
|
||||
elif isinstance(output, torch.Tensor) and output.device.type == 'cpu':
|
||||
ram_usage += output.numel() * output.element_size()
|
||||
oom_ram_usage += output.numel() * output.element_size()
|
||||
elif isinstance(output, ModelPatcher) and self.used_generation[key] != self.generation:
|
||||
#old ModelPatchers are the first to go
|
||||
ram_usage = 1e30
|
||||
oom_ram_usage = 1e30
|
||||
scan_list_for_ram_usage(cache_entry.outputs)
|
||||
|
||||
oom_score *= ram_usage
|
||||
if ram_usage < min_entry_size:
|
||||
continue
|
||||
|
||||
oom_score *= oom_ram_usage
|
||||
#In the case where we have no information on the node ram usage at all,
|
||||
#break OOM score ties on the last touch timestamp (pure LRU)
|
||||
bisect.insort(clean_list, (oom_score, self.timestamps[key], key))
|
||||
bisect.insort(clean_list, (oom_score, self.timestamps[key], key, ram_usage))
|
||||
|
||||
freed = 0
|
||||
while psutil.virtual_memory().available < target and clean_list:
|
||||
_, _, key = clean_list.pop()
|
||||
_, _, key, ram_usage = clean_list.pop()
|
||||
del self.cache[key]
|
||||
self.used_generation.pop(key, None)
|
||||
self.timestamps.pop(key, None)
|
||||
self.children.pop(key, None)
|
||||
freed += ram_usage
|
||||
return freed
|
||||
|
||||
@@ -3,11 +3,12 @@ from typing import Type, Literal
|
||||
import nodes
|
||||
import asyncio
|
||||
import inspect
|
||||
from comfy_execution.graph_utils import is_link, ExecutionBlocker
|
||||
from comfy_execution.graph_utils import is_link, ExecutionBlocker, ExecutionFailureBlocker
|
||||
from comfy.comfy_types.node_typing import ComfyNodeABC, InputTypeDict, InputTypeOptions
|
||||
|
||||
# NOTE: ExecutionBlocker code got moved to graph_utils.py to prevent torch being imported too soon during unit tests
|
||||
ExecutionBlocker = ExecutionBlocker
|
||||
ExecutionFailureBlocker = ExecutionFailureBlocker
|
||||
|
||||
class DependencyCycleError(Exception):
|
||||
pass
|
||||
@@ -201,19 +202,28 @@ class ExecutionList(TopologicalSort):
|
||||
self.staged_node_id = None
|
||||
self.execution_cache = {}
|
||||
self.execution_cache_listeners = {}
|
||||
self.transient_cache = {}
|
||||
self.failure_tainted_parents = set()
|
||||
|
||||
def is_cached(self, node_id):
|
||||
return self.output_cache.get_local(node_id) is not None
|
||||
return node_id in self.transient_cache or self.output_cache.get_local(node_id) is not None
|
||||
|
||||
def _get_cache_value(self, node_id):
|
||||
if node_id in self.transient_cache:
|
||||
return self.transient_cache[node_id]
|
||||
return self.output_cache.get_local(node_id)
|
||||
|
||||
def cache_link(self, from_node_id, to_node_id):
|
||||
if to_node_id not in self.execution_cache:
|
||||
self.execution_cache[to_node_id] = {}
|
||||
self.execution_cache[to_node_id][from_node_id] = self.output_cache.get_local(from_node_id)
|
||||
self.execution_cache[to_node_id][from_node_id] = self._get_cache_value(from_node_id)
|
||||
if from_node_id not in self.execution_cache_listeners:
|
||||
self.execution_cache_listeners[from_node_id] = set()
|
||||
self.execution_cache_listeners[from_node_id].add(to_node_id)
|
||||
|
||||
def get_cache(self, from_node_id, to_node_id):
|
||||
if from_node_id in self.transient_cache:
|
||||
return self.transient_cache[from_node_id]
|
||||
if to_node_id not in self.execution_cache:
|
||||
return None
|
||||
value = self.execution_cache[to_node_id].get(from_node_id)
|
||||
@@ -223,7 +233,9 @@ class ExecutionList(TopologicalSort):
|
||||
self.output_cache.set_local(from_node_id, value)
|
||||
return value
|
||||
|
||||
def cache_update(self, node_id, value):
|
||||
def cache_update(self, node_id, value, transient=False):
|
||||
if transient:
|
||||
self.transient_cache[node_id] = value
|
||||
if node_id in self.execution_cache_listeners:
|
||||
for to_node_id in self.execution_cache_listeners[node_id]:
|
||||
if to_node_id in self.execution_cache:
|
||||
@@ -233,6 +245,25 @@ class ExecutionList(TopologicalSort):
|
||||
super().add_strong_link(from_node_id, from_socket, to_node_id)
|
||||
self.cache_link(from_node_id, to_node_id)
|
||||
|
||||
def add_completion_link(self, from_node_id, to_node_id):
|
||||
# Block to_node_id until from_node_id finishes, without consuming any of its output sockets.
|
||||
if not self.is_cached(from_node_id):
|
||||
self.add_node(from_node_id)
|
||||
if to_node_id not in self.blocking[from_node_id]:
|
||||
self.blocking[from_node_id][to_node_id] = {}
|
||||
self.blockCount[to_node_id] += 1
|
||||
|
||||
def mark_failure_tainted(self, node_id):
|
||||
# Taint all ephemeral ancestors of a failed or failure-blocked node so dynamically-expanded parents are
|
||||
# never cached as reusable when part of their expansion did not complete.
|
||||
parent_id = self.dynprompt.get_parent_node_id(node_id)
|
||||
while parent_id is not None and parent_id not in self.failure_tainted_parents:
|
||||
self.failure_tainted_parents.add(parent_id)
|
||||
parent_id = self.dynprompt.get_parent_node_id(parent_id)
|
||||
|
||||
def is_failure_tainted(self, node_id):
|
||||
return node_id in self.failure_tainted_parents
|
||||
|
||||
async def stage_node_execution(self):
|
||||
assert self.staged_node_id is None
|
||||
if self.is_empty():
|
||||
@@ -323,7 +354,9 @@ class ExecutionList(TopologicalSort):
|
||||
blocked_by = { node_id: {} for node_id in self.pendingNodes }
|
||||
for from_node_id in self.blocking:
|
||||
for to_node_id in self.blocking[from_node_id]:
|
||||
if True in self.blocking[from_node_id][to_node_id].values():
|
||||
# Strong links have a True socket entry; completion links have no socket entries at all.
|
||||
sockets = self.blocking[from_node_id][to_node_id]
|
||||
if len(sockets) == 0 or True in sockets.values():
|
||||
blocked_by[to_node_id][from_node_id] = True
|
||||
to_remove = [node_id for node_id in blocked_by if len(blocked_by[node_id]) == 0]
|
||||
while len(to_remove) > 0:
|
||||
|
||||
@@ -153,3 +153,9 @@ class ExecutionBlocker:
|
||||
"""
|
||||
def __init__(self, message):
|
||||
self.message = message
|
||||
|
||||
|
||||
class ExecutionFailureBlocker(ExecutionBlocker):
|
||||
def __init__(self, node_id):
|
||||
super().__init__(None)
|
||||
self.node_id = node_id
|
||||
|
||||
+33
-4
@@ -56,6 +56,9 @@ PREVIEWABLE_MEDIA_TYPES = frozenset({'images', 'video', 'audio', '3d', 'text'})
|
||||
# 3D file extensions for preview fallback (no dedicated media_type exists)
|
||||
THREE_D_EXTENSIONS = frozenset({'.obj', '.fbx', '.gltf', '.glb', '.usdz'})
|
||||
|
||||
# Text file extensions for preview fallback (the formats SaveText can produce)
|
||||
TEXT_EXTENSIONS = frozenset({'.txt', '.md', '.json'})
|
||||
|
||||
|
||||
def has_3d_extension(filename: str) -> bool:
|
||||
lower = filename.lower()
|
||||
@@ -143,9 +146,10 @@ def is_previewable(media_type: str, item: dict) -> bool:
|
||||
Maintains backwards compatibility with existing logic.
|
||||
|
||||
Priority:
|
||||
1. media_type is 'images', 'video', 'audio', or '3d'
|
||||
1. media_type is 'images', 'video', 'audio', '3d', or 'text'
|
||||
2. format field starts with 'video/' or 'audio/'
|
||||
3. filename has a 3D extension (.obj, .fbx, .gltf, .glb, .usdz)
|
||||
4. filename has a text extension (.txt, .md, .json, ...)
|
||||
"""
|
||||
if media_type in PREVIEWABLE_MEDIA_TYPES:
|
||||
return True
|
||||
@@ -156,10 +160,12 @@ def is_previewable(media_type: str, item: dict) -> bool:
|
||||
if fmt and (fmt.startswith('video/') or fmt.startswith('audio/')):
|
||||
return True
|
||||
|
||||
# Check for 3D files by extension
|
||||
# Check for 3D and text files by extension
|
||||
filename = item.get('filename', '').lower()
|
||||
if any(filename.endswith(ext) for ext in THREE_D_EXTENSIONS):
|
||||
return True
|
||||
if any(filename.endswith(ext) for ext in TEXT_EXTENSIONS):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
@@ -198,10 +204,14 @@ def normalize_history_item(prompt_id: str, history_item: dict, include_outputs:
|
||||
outputs_count, preview_output = get_outputs_summary(outputs)
|
||||
|
||||
execution_error = None
|
||||
execution_errors = []
|
||||
execution_start_time = None
|
||||
execution_end_time = None
|
||||
execution_success = None
|
||||
was_interrupted = False
|
||||
execution_summary = {}
|
||||
if status_info:
|
||||
execution_summary = status_info.get('execution_summary') or {}
|
||||
messages = status_info.get('messages', [])
|
||||
for entry in messages:
|
||||
if isinstance(entry, (list, tuple)) and len(entry) >= 2:
|
||||
@@ -211,10 +221,22 @@ def normalize_history_item(prompt_id: str, history_item: dict, include_outputs:
|
||||
execution_start_time = event_data.get('timestamp')
|
||||
elif event_name in ('execution_success', 'execution_error', 'execution_interrupted'):
|
||||
execution_end_time = event_data.get('timestamp')
|
||||
if event_name == 'execution_error':
|
||||
if event_name == 'execution_success':
|
||||
execution_success = event_data
|
||||
elif event_name == 'execution_error':
|
||||
execution_error = event_data
|
||||
elif event_name == 'execution_interrupted':
|
||||
was_interrupted = True
|
||||
elif event_name == 'execution_node_error':
|
||||
execution_errors.append(event_data)
|
||||
|
||||
completion_status = execution_summary.get('completion_status')
|
||||
if completion_status is None and execution_success is not None:
|
||||
completion_status = execution_success.get('completion_status', 'success')
|
||||
if completion_status is None and status_str == 'success':
|
||||
completion_status = 'success'
|
||||
has_errors = execution_summary.get('has_errors', bool(execution_errors)) if completion_status is not None else None
|
||||
execution_error_count = execution_summary.get('execution_error_count', len(execution_errors)) if completion_status is not None else None
|
||||
|
||||
if status_str == 'success':
|
||||
status = JobStatus.COMPLETED
|
||||
@@ -231,6 +253,9 @@ def normalize_history_item(prompt_id: str, history_item: dict, include_outputs:
|
||||
'execution_start_time': execution_start_time,
|
||||
'execution_end_time': execution_end_time,
|
||||
'execution_error': execution_error,
|
||||
'completion_status': completion_status,
|
||||
'has_errors': has_errors,
|
||||
'execution_error_count': execution_error_count,
|
||||
'outputs_count': outputs_count,
|
||||
'preview_output': preview_output,
|
||||
'workflow_id': workflow_id,
|
||||
@@ -239,6 +264,7 @@ def normalize_history_item(prompt_id: str, history_item: dict, include_outputs:
|
||||
if include_outputs:
|
||||
job['outputs'] = normalize_outputs(outputs)
|
||||
job['execution_status'] = status_info
|
||||
job['execution_errors'] = execution_errors
|
||||
job['workflow'] = {
|
||||
'prompt': prompt,
|
||||
'extra_data': extra_data,
|
||||
@@ -255,6 +281,10 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]:
|
||||
Preview priority (matching frontend):
|
||||
1. type="output" with previewable media
|
||||
2. Any previewable media
|
||||
|
||||
Text content entries (strings under 'text') are preview-only metadata,
|
||||
matching the frontend's METADATA_KEYS: they can serve as the fallback
|
||||
preview but are not counted as outputs.
|
||||
"""
|
||||
count = 0
|
||||
preview_output = None
|
||||
@@ -275,7 +305,6 @@ def get_outputs_summary(outputs: dict) -> tuple[int, Optional[dict]]:
|
||||
if normalized is None:
|
||||
# Not a 3D file string — check for text preview
|
||||
if media_type == 'text':
|
||||
count += 1
|
||||
if preview_output is None:
|
||||
if isinstance(item, tuple):
|
||||
text_value = item[0] if item else ''
|
||||
|
||||
@@ -17,6 +17,7 @@ class NodeState(Enum):
|
||||
Running = "running"
|
||||
Finished = "finished"
|
||||
Error = "error"
|
||||
Blocked = "blocked"
|
||||
|
||||
|
||||
class NodeProgressState(TypedDict):
|
||||
@@ -301,17 +302,25 @@ class ProgressRegistry:
|
||||
node_id, value, max_value, entry, self.prompt_id, image
|
||||
)
|
||||
|
||||
def finish_progress(self, node_id: str) -> None:
|
||||
"""Finish progress tracking for a node"""
|
||||
def _finish_progress(self, node_id: str, state: NodeState) -> None:
|
||||
entry = self.ensure_entry(node_id)
|
||||
entry["state"] = NodeState.Finished
|
||||
entry["state"] = state
|
||||
entry["value"] = entry["max"]
|
||||
|
||||
# Notify all enabled handlers
|
||||
for handler in self.handlers.values():
|
||||
if handler.enabled:
|
||||
handler.finish_handler(node_id, entry, self.prompt_id)
|
||||
|
||||
def finish_progress(self, node_id: str) -> None:
|
||||
"""Finish progress tracking for a node"""
|
||||
self._finish_progress(node_id, NodeState.Finished)
|
||||
|
||||
def error_progress(self, node_id: str) -> None:
|
||||
self._finish_progress(node_id, NodeState.Error)
|
||||
|
||||
def block_progress(self, node_id: str) -> None:
|
||||
self._finish_progress(node_id, NodeState.Blocked)
|
||||
|
||||
def reset_handlers(self) -> None:
|
||||
"""Reset all handlers"""
|
||||
for handler in self.handlers.values():
|
||||
|
||||
@@ -298,6 +298,7 @@ class PreviewAudio(IO.ComfyNode):
|
||||
search_aliases=["play audio"],
|
||||
display_name="Preview Audio",
|
||||
category="audio",
|
||||
description="Preview the audio without saving it to the ComfyUI output directory.",
|
||||
inputs=[
|
||||
IO.Audio.Input("audio"),
|
||||
],
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import json
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageDraw, ImageEnhance, ImageFont
|
||||
@@ -166,6 +168,111 @@ def boxes_to_regions(boxes, width: int, height: int) -> list:
|
||||
return regions
|
||||
|
||||
|
||||
def normalize_incoming_boxes(bboxes) -> list:
|
||||
if isinstance(bboxes, dict):
|
||||
frame = [bboxes]
|
||||
elif not isinstance(bboxes, list) or not bboxes:
|
||||
frame = []
|
||||
elif isinstance(bboxes[0], dict):
|
||||
frame = bboxes
|
||||
else:
|
||||
frame = bboxes[0] if isinstance(bboxes[0], list) else []
|
||||
boxes = []
|
||||
for box in frame:
|
||||
if not isinstance(box, dict):
|
||||
continue
|
||||
norm = {
|
||||
"x": box.get("x", 0),
|
||||
"y": box.get("y", 0),
|
||||
"width": box.get("width", 0),
|
||||
"height": box.get("height", 0),
|
||||
}
|
||||
meta = box.get("metadata")
|
||||
if isinstance(meta, dict):
|
||||
norm["metadata"] = meta
|
||||
boxes.append(norm)
|
||||
return boxes
|
||||
|
||||
|
||||
def _looks_like_element(box: dict) -> bool:
|
||||
bbox = box.get("bbox")
|
||||
return isinstance(bbox, (list, tuple)) and len(bbox) == 4
|
||||
|
||||
|
||||
def _looks_like_bbox(box: dict) -> bool:
|
||||
return all(key in box for key in ("x", "y", "width", "height"))
|
||||
|
||||
|
||||
def elements_to_boxes(elements: list, width: int, height: int) -> list:
|
||||
boxes = []
|
||||
for element in elements:
|
||||
if not isinstance(element, dict):
|
||||
continue
|
||||
bbox = element.get("bbox")
|
||||
if not (isinstance(bbox, (list, tuple)) and len(bbox) == 4):
|
||||
raise ValueError("bboxes element is missing a valid 'bbox' [ymin, xmin, ymax, xmax]")
|
||||
try:
|
||||
ymin, xmin, ymax, xmax = (float(v) / 1000.0 for v in bbox)
|
||||
except (TypeError, ValueError):
|
||||
raise ValueError("bboxes element 'bbox' must contain four numbers")
|
||||
etype = "text" if element.get("type") == "text" else "obj"
|
||||
boxes.append({
|
||||
"x": round(min(xmin, xmax) * width),
|
||||
"y": round(min(ymin, ymax) * height),
|
||||
"width": round(abs(xmax - xmin) * width),
|
||||
"height": round(abs(ymax - ymin) * height),
|
||||
"metadata": {
|
||||
"type": etype,
|
||||
"text": element.get("text", "") if etype == "text" else "",
|
||||
"desc": element.get("desc", ""),
|
||||
"palette": element.get("color_palette", []) or [],
|
||||
},
|
||||
})
|
||||
return boxes
|
||||
|
||||
|
||||
def boxes_from_input(data, width: int, height: int) -> list:
|
||||
if data is None:
|
||||
return []
|
||||
if isinstance(data, str):
|
||||
text = data.strip()
|
||||
if not text:
|
||||
return []
|
||||
try:
|
||||
data = json.loads(text)
|
||||
except (ValueError, TypeError) as exc:
|
||||
raise ValueError(f"bboxes string input is not valid JSON: {exc}") from exc
|
||||
if isinstance(data, dict):
|
||||
if _looks_like_element(data):
|
||||
return elements_to_boxes([data], width, height)
|
||||
if _looks_like_bbox(data):
|
||||
return normalize_incoming_boxes(data)
|
||||
raise ValueError(
|
||||
"bboxes dict must be a bounding box (x, y, width, height) or an element (with a 'bbox')"
|
||||
)
|
||||
if not isinstance(data, list):
|
||||
raise ValueError(
|
||||
"bboxes input must be bounding boxes, elements, or a JSON string, "
|
||||
f"got {type(data).__name__}"
|
||||
)
|
||||
if not data:
|
||||
return []
|
||||
first = data[0]
|
||||
if isinstance(first, list):
|
||||
return normalize_incoming_boxes(data)
|
||||
if isinstance(first, dict):
|
||||
if _looks_like_element(first):
|
||||
return elements_to_boxes(data, width, height)
|
||||
if _looks_like_bbox(first):
|
||||
return normalize_incoming_boxes(data)
|
||||
raise ValueError(
|
||||
"bboxes items must be bounding boxes (x, y, width, height) or elements (with a 'bbox')"
|
||||
)
|
||||
raise ValueError(
|
||||
f"bboxes list must contain bounding boxes or elements, got {type(first).__name__}"
|
||||
)
|
||||
|
||||
|
||||
def _norm_bbox(region: dict) -> list[int]:
|
||||
def grid(value: float) -> int:
|
||||
return max(0, min(1000, round(value * 1000)))
|
||||
@@ -217,29 +324,48 @@ class CreateBoundingBoxes(io.ComfyNode):
|
||||
optional=True,
|
||||
tooltip="Optional image used as background in the canvas and preview.",
|
||||
),
|
||||
io.MultiType.Input(
|
||||
"bboxes",
|
||||
[io.BoundingBox, io.Array, io.String],
|
||||
optional=True,
|
||||
tooltip="Bounding boxes, elements, or a JSON string to initialize the canvas. A new upstream value initializes the canvas; edits made on the canvas take priority and are kept until the upstream value changes again.",
|
||||
),
|
||||
io.Int.Input("width", default=1024, min=64, max=16384, step=16,
|
||||
tooltip="Width of the canvas and the pixel grid for the bounding boxes."),
|
||||
io.Int.Input("height", default=1024, min=64, max=16384, step=16,
|
||||
tooltip="Height of the canvas and the pixel grid for the bounding boxes."),
|
||||
editor_state,
|
||||
io.BoundingBoxes.Input(
|
||||
"last_incoming",
|
||||
optional=True,
|
||||
tooltip="Internal state managed by the canvas: the upstream bboxes value that last initialized it. Leave empty to re-initialize the canvas from the bboxes input on the next run.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
io.Image.Output(display_name="preview"),
|
||||
io.BoundingBox.Output(display_name="bboxes"),
|
||||
io.Array.Output(display_name="elements"),
|
||||
],
|
||||
is_output_node=True,
|
||||
is_experimental=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, width, height, editor_state=None, background=None) -> io.NodeOutput:
|
||||
regions = boxes_to_regions(editor_state, width, height)
|
||||
def execute(cls, width, height, editor_state=None, last_incoming=None, background=None, bboxes=None) -> io.NodeOutput:
|
||||
incoming = boxes_from_input(bboxes, width, height)
|
||||
applied = last_incoming if isinstance(last_incoming, list) else []
|
||||
upstream_changed = bool(incoming) and incoming != applied
|
||||
source = incoming if upstream_changed else (editor_state or [])
|
||||
regions = boxes_to_regions(source, width, height)
|
||||
preview = render_preview(regions, width, height, _bg_from_image(background))
|
||||
ui = {"dims": [width, height]}
|
||||
if incoming:
|
||||
ui["input_bboxes"] = incoming
|
||||
return io.NodeOutput(
|
||||
preview,
|
||||
fractions_to_bbox_frame(regions, width, height),
|
||||
build_elements(regions),
|
||||
ui={"dims": [width, height]},
|
||||
ui=ui,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -844,15 +844,18 @@ class ImageMergeTileList(IO.ComfyNode):
|
||||
# Format specifications
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Maps (file_format, bit_depth, has_alpha) -> (numpy dtype scale, av pixel format,
|
||||
# stream pix_fmt). Keeps the encode path declarative instead of branchy.
|
||||
# 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.
|
||||
_FORMAT_SPECS = {
|
||||
("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"},
|
||||
("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"},
|
||||
}
|
||||
|
||||
|
||||
@@ -891,10 +894,11 @@ 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 inside the log branch so negative / out-of-range values don't blow up;
|
||||
# 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;
|
||||
# values above 1.0 are allowed and extrapolate naturally.
|
||||
low = (t ** 2) / 3.0
|
||||
high = (torch.exp((t.clamp(min=_HLG_C) - _HLG_C) / _HLG_A) + _HLG_B) / 12.0
|
||||
high = (torch.exp((t.clamp(min=0.5) - _HLG_C) / _HLG_A) + _HLG_B) / 12.0
|
||||
return torch.where(t <= 0.5, low, high)
|
||||
|
||||
|
||||
@@ -1087,7 +1091,8 @@ def _encode_image(
|
||||
bit_depth: str,
|
||||
colorspace: str,
|
||||
) -> bytes:
|
||||
"""Encode a single HxWxC tensor to PNG or EXR bytes in memory.
|
||||
"""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.
|
||||
|
||||
For EXR the input is interpreted according to `colorspace` and converted
|
||||
to scene-linear (EXR's convention) before writing:
|
||||
@@ -1101,10 +1106,16 @@ 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[(file_format, bit_depth, has_alpha)]
|
||||
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)."
|
||||
)
|
||||
|
||||
if spec["dtype"] == np.float32:
|
||||
# EXR path: preserve full range, no clamp.
|
||||
|
||||
@@ -61,14 +61,10 @@ class Load3D(IO.ComfyNode):
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model_file, image, **kwargs) -> IO.NodeOutput:
|
||||
image_path = folder_paths.get_annotated_filepath(image['image'])
|
||||
mask_path = folder_paths.get_annotated_filepath(image['mask'])
|
||||
normal_path = folder_paths.get_annotated_filepath(image['normal'])
|
||||
|
||||
load_image_node = nodes.LoadImage()
|
||||
output_image, ignore_mask = load_image_node.load_image(image=image_path)
|
||||
ignore_image, output_mask = load_image_node.load_image(image=mask_path)
|
||||
normal_image, ignore_mask2 = load_image_node.load_image(image=normal_path)
|
||||
output_image, ignore_mask = load_image_node.load_image(image=image['image'])
|
||||
ignore_image, output_mask = load_image_node.load_image(image=image['mask'])
|
||||
normal_image, ignore_mask2 = load_image_node.load_image(image=image['normal'])
|
||||
|
||||
video = None
|
||||
|
||||
@@ -96,6 +92,7 @@ class Preview3D(IO.ComfyNode):
|
||||
search_aliases=["view mesh", "3d viewer"],
|
||||
display_name="Preview 3D & Animation",
|
||||
category="3d",
|
||||
description="Preview a 3D model file without saving it to the ComfyUI output directory.",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
inputs=[
|
||||
@@ -140,6 +137,7 @@ class Preview3DAdvanced(IO.ComfyNode):
|
||||
display_name="Preview 3D (Advanced)",
|
||||
search_aliases=["preview 3d", "3d viewer", "view mesh", "frame 3d", "3d camera output"],
|
||||
category="3d",
|
||||
description="Preview a 3D model file without saving it to the ComfyUI output directory.",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
inputs=[
|
||||
@@ -176,8 +174,9 @@ 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['camera_info']
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('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(
|
||||
@@ -197,6 +196,7 @@ class PreviewGaussianSplat(IO.ComfyNode):
|
||||
node_id="PreviewGaussianSplat",
|
||||
display_name="Preview Splat",
|
||||
category="3d",
|
||||
description="Preview a gaussian splat 3D file without saving it to the ComfyUI output directory.",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
search_aliases=[
|
||||
@@ -244,8 +244,9 @@ 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['camera_info']
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('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(
|
||||
@@ -265,6 +266,7 @@ class PreviewPointCloud(IO.ComfyNode):
|
||||
node_id="PreviewPointCloud",
|
||||
display_name="Preview Point Cloud",
|
||||
category="3d",
|
||||
description="Preview a point cloud 3D file without saving it to the ComfyUI output directory.",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
search_aliases=[
|
||||
@@ -303,8 +305,9 @@ 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['camera_info']
|
||||
camera_info = camera_info_input if camera_info_input is not None else viewport_state.get('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(
|
||||
@@ -375,8 +378,9 @@ 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['camera_info'], width, height)
|
||||
return IO.NodeOutput(file_3d, model_3d_info, viewport_state.get('camera_info'), width, height)
|
||||
|
||||
|
||||
class Load3DExtension(ComfyExtension):
|
||||
|
||||
@@ -419,17 +419,18 @@ class MaskPreview(IO.ComfyNode):
|
||||
search_aliases=["show mask", "view mask", "inspect mask", "debug mask"],
|
||||
display_name="Preview Mask",
|
||||
category="image/mask",
|
||||
description="Saves the input images to your ComfyUI output directory.",
|
||||
description="Preview the masks without saving them to the ComfyUI output directory.",
|
||||
inputs=[
|
||||
IO.Mask.Input("mask"),
|
||||
],
|
||||
hidden=[IO.Hidden.prompt, IO.Hidden.extra_pnginfo],
|
||||
is_output_node=True,
|
||||
outputs=[IO.Mask.Output(display_name="mask")]
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, mask, filename_prefix="ComfyUI") -> IO.NodeOutput:
|
||||
return IO.NodeOutput(ui=UI.PreviewMask(mask))
|
||||
return IO.NodeOutput(mask, ui=UI.PreviewMask(mask))
|
||||
|
||||
|
||||
class MaskExtension(ComfyExtension):
|
||||
|
||||
@@ -18,6 +18,7 @@ class PreviewAny():
|
||||
|
||||
CATEGORY = "utilities"
|
||||
SEARCH_ALIASES = ["show output", "inspect", "debug", "print value", "show text"]
|
||||
DESCRIPTION = "Preview any input value as text."
|
||||
|
||||
def main(self, source=None):
|
||||
torch.set_printoptions(edgeitems=6)
|
||||
|
||||
@@ -10,11 +10,10 @@ class String(io.ComfyNode):
|
||||
return io.Schema(
|
||||
node_id="PrimitiveString",
|
||||
search_aliases=["text", "string", "text box", "prompt"],
|
||||
display_name="Text String (DEPRECATED)",
|
||||
display_name="Text",
|
||||
category="utilities/primitive",
|
||||
inputs=[io.String.Input("value")],
|
||||
outputs=[io.String.Output()],
|
||||
is_deprecated=True
|
||||
outputs=[io.String.Output()]
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -28,7 +27,7 @@ class StringMultiline(io.ComfyNode):
|
||||
return io.Schema(
|
||||
node_id="PrimitiveStringMultiline",
|
||||
search_aliases=["text", "string", "text multiline", "string multiline", "text box", "prompt"],
|
||||
display_name="Input Text",
|
||||
display_name="Text (Multiline)",
|
||||
category="utilities/primitive",
|
||||
essentials_category="Basics",
|
||||
inputs=[io.String.Input("value", multiline=True)],
|
||||
|
||||
@@ -13,7 +13,7 @@ from typing_extensions import override
|
||||
|
||||
import folder_paths
|
||||
from comfy.cli_args import args
|
||||
from comfy_api.latest import ComfyExtension, IO, Types
|
||||
from comfy_api.latest import ComfyExtension, IO, Types, UI
|
||||
|
||||
|
||||
def pack_variable_mesh_batch(vertices, faces, colors=None, uvs=None, texture=None, unlit=False):
|
||||
@@ -406,10 +406,165 @@ class SaveGLB(IO.ComfyNode):
|
||||
return IO.NodeOutput(ui={"3d": results})
|
||||
|
||||
|
||||
def _save_file3d_to_output(model_3d: Types.File3D, filename_prefix: str) -> str:
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
|
||||
filename_prefix, folder_paths.get_output_directory()
|
||||
)
|
||||
ext = model_3d.format or "glb"
|
||||
saved_filename = f"{filename}_{counter:05}.{ext}"
|
||||
model_3d.save_to(os.path.join(full_output_folder, saved_filename))
|
||||
return f"{subfolder}/{saved_filename}" if subfolder else saved_filename
|
||||
|
||||
|
||||
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')
|
||||
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(
|
||||
model_3d,
|
||||
model_3d_info,
|
||||
camera_info,
|
||||
width,
|
||||
height,
|
||||
ui=UI.PreviewUI3DAdvanced(model_file, camera_info, model_3d_info),
|
||||
)
|
||||
|
||||
|
||||
class Save3DAdvanced(IO.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return IO.Schema(
|
||||
node_id="Save3DAdvanced",
|
||||
display_name="Save 3D (Advanced)",
|
||||
search_aliases=["save 3d", "export 3d model", "save mesh advanced"],
|
||||
category="3d",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
inputs=[
|
||||
IO.MultiType.Input(
|
||||
"model_3d",
|
||||
types=[
|
||||
IO.File3DGLB,
|
||||
IO.File3DGLTF,
|
||||
IO.File3DFBX,
|
||||
IO.File3DOBJ,
|
||||
IO.File3DSTL,
|
||||
IO.File3DUSDZ,
|
||||
IO.File3DAny,
|
||||
],
|
||||
tooltip="3D model file from an upstream 3D node.",
|
||||
),
|
||||
IO.String.Input("filename_prefix", default="3d/ComfyUI"),
|
||||
IO.Load3D.Input("viewport_state"),
|
||||
IO.Load3DModelInfo.Input("model_3d_info", optional=True, advanced=True),
|
||||
IO.Load3DCamera.Input("camera_info", optional=True, advanced=True),
|
||||
IO.Int.Input("width", default=1024, min=1, max=4096, step=1),
|
||||
IO.Int.Input("height", default=1024, min=1, max=4096, step=1),
|
||||
],
|
||||
outputs=[
|
||||
IO.File3DAny.Output(display_name="model_3d"),
|
||||
IO.Load3DModelInfo.Output(display_name="model_3d_info"),
|
||||
IO.Load3DCamera.Output(display_name="camera_info"),
|
||||
IO.Int.Output(display_name="width"),
|
||||
IO.Int.Output(display_name="height"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model_3d: Types.File3D, viewport_state, width: int, height: int, filename_prefix: str, **kwargs) -> IO.NodeOutput:
|
||||
return execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs)
|
||||
|
||||
|
||||
class SaveGaussianSplat(IO.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return IO.Schema(
|
||||
node_id="SaveGaussianSplat",
|
||||
display_name="Save Splat",
|
||||
search_aliases=["save splat", "save gaussian splat", "export gaussian", "export splat"],
|
||||
category="3d",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
inputs=[
|
||||
IO.MultiType.Input(
|
||||
"model_3d",
|
||||
types=[
|
||||
IO.File3DSplatAny,
|
||||
IO.File3DPLY,
|
||||
IO.File3DSPLAT,
|
||||
IO.File3DSPZ,
|
||||
IO.File3DKSPLAT,
|
||||
],
|
||||
tooltip="A gaussian splat 3D file.",
|
||||
),
|
||||
IO.String.Input("filename_prefix", default="3d/ComfyUI"),
|
||||
IO.Load3D.Input("viewport_state"),
|
||||
IO.Load3DModelInfo.Input("model_3d_info", optional=True, advanced=True),
|
||||
IO.Load3DCamera.Input("camera_info", optional=True, advanced=True),
|
||||
IO.Int.Input("width", default=1024, min=1, max=4096, step=1),
|
||||
IO.Int.Input("height", default=1024, min=1, max=4096, step=1),
|
||||
],
|
||||
outputs=[
|
||||
IO.File3DSplatAny.Output(display_name="model_3d"),
|
||||
IO.Load3DModelInfo.Output(display_name="model_3d_info"),
|
||||
IO.Load3DCamera.Output(display_name="camera_info"),
|
||||
IO.Int.Output(display_name="width"),
|
||||
IO.Int.Output(display_name="height"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model_3d: Types.File3D, viewport_state, width: int, height: int, filename_prefix: str, **kwargs) -> IO.NodeOutput:
|
||||
return execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs)
|
||||
|
||||
|
||||
class SavePointCloud(IO.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return IO.Schema(
|
||||
node_id="SavePointCloud",
|
||||
display_name="Save Point Cloud",
|
||||
search_aliases=["save point cloud", "save pointcloud", "export point cloud"],
|
||||
category="3d",
|
||||
is_experimental=True,
|
||||
is_output_node=True,
|
||||
inputs=[
|
||||
IO.MultiType.Input(
|
||||
"model_3d",
|
||||
types=[
|
||||
IO.File3DPointCloudAny,
|
||||
IO.File3DPLY,
|
||||
],
|
||||
tooltip="Point cloud file (.ply)",
|
||||
),
|
||||
IO.String.Input("filename_prefix", default="3d/ComfyUI"),
|
||||
IO.Load3D.Input("viewport_state"),
|
||||
IO.Load3DModelInfo.Input("model_3d_info", optional=True, advanced=True),
|
||||
IO.Load3DCamera.Input("camera_info", optional=True, advanced=True),
|
||||
IO.Int.Input("width", default=1024, min=1, max=4096, step=1),
|
||||
IO.Int.Input("height", default=1024, min=1, max=4096, step=1),
|
||||
],
|
||||
outputs=[
|
||||
IO.File3DPointCloudAny.Output(display_name="model_3d"),
|
||||
IO.Load3DModelInfo.Output(display_name="model_3d_info"),
|
||||
IO.Load3DCamera.Output(display_name="camera_info"),
|
||||
IO.Int.Output(display_name="width"),
|
||||
IO.Int.Output(display_name="height"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model_3d: Types.File3D, viewport_state, width: int, height: int, filename_prefix: str, **kwargs) -> IO.NodeOutput:
|
||||
return execute_save_3d_advanced(model_3d, viewport_state, width, height, filename_prefix, kwargs)
|
||||
|
||||
|
||||
class Save3DExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[IO.ComfyNode]]:
|
||||
return [SaveGLB]
|
||||
return [SaveGLB, Save3DAdvanced, SaveGaussianSplat, SavePointCloud]
|
||||
|
||||
|
||||
async def comfy_entrypoint() -> Save3DExtension:
|
||||
|
||||
@@ -0,0 +1,614 @@
|
||||
import logging
|
||||
|
||||
from typing_extensions import override
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
import torch
|
||||
|
||||
import comfy.model_management
|
||||
from comfy.ldm.seedvr.color_fix import (
|
||||
adain_color_transfer,
|
||||
lab_color_transfer,
|
||||
wavelet_color_transfer,
|
||||
)
|
||||
from comfy.ldm.seedvr.constants import (
|
||||
BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE,
|
||||
SEEDVR2_ADAIN_SCALE_MULTIPLIER,
|
||||
SEEDVR2_CHUNK_GIB_PER_MPX_FRAME,
|
||||
SEEDVR2_CHUNK_RESERVED_GIB,
|
||||
SEEDVR2_CHUNK_SIGMA_GIB,
|
||||
SEEDVR2_CHUNK_SIGMA_K,
|
||||
SEEDVR2_COLOR_MEM_HEADROOM,
|
||||
SEEDVR2_DTYPE_BYTES_FLOOR,
|
||||
SEEDVR2_LAB_SCALE_MULTIPLIER,
|
||||
SEEDVR2_LATENT_CHANNELS,
|
||||
SEEDVR2_OOM_BACKOFF_DIVISOR,
|
||||
SEEDVR2_WAVELET_SCALE_MULTIPLIER,
|
||||
)
|
||||
|
||||
from torchvision.transforms import functional as TVF
|
||||
from torchvision.transforms.functional import InterpolationMode
|
||||
|
||||
|
||||
_SEEDVR2_INVALID_MODEL_MSG_PREFIX = "SeedVR2Conditioning: model object does not match expected SeedVR2 structure"
|
||||
_ATTR_MISSING = object()
|
||||
|
||||
|
||||
def _resolve_seedvr2_diffusion_model(model):
|
||||
inner = getattr(model, "model", _ATTR_MISSING)
|
||||
if inner is _ATTR_MISSING:
|
||||
raise RuntimeError(
|
||||
f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: input has no 'model' attribute "
|
||||
f"(got type {type(model).__name__})."
|
||||
)
|
||||
if inner is None:
|
||||
raise RuntimeError(
|
||||
f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: input.model is None "
|
||||
f"(input type {type(model).__name__})."
|
||||
)
|
||||
diffusion_model = getattr(inner, "diffusion_model", _ATTR_MISSING)
|
||||
if diffusion_model is _ATTR_MISSING:
|
||||
raise RuntimeError(
|
||||
f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: 'model.model' has no "
|
||||
f"'diffusion_model' attribute (got type {type(inner).__name__})."
|
||||
)
|
||||
if diffusion_model is None:
|
||||
raise RuntimeError(
|
||||
f"{_SEEDVR2_INVALID_MODEL_MSG_PREFIX}: 'model.model.diffusion_model' "
|
||||
f"is None (model.model type {type(inner).__name__})."
|
||||
)
|
||||
return diffusion_model
|
||||
|
||||
|
||||
def div_pad(image, factor):
|
||||
height_factor, width_factor = factor
|
||||
height, width = image.shape[-2:]
|
||||
|
||||
pad_height = (height_factor - (height % height_factor)) % height_factor
|
||||
pad_width = (width_factor - (width % width_factor)) % width_factor
|
||||
|
||||
if pad_height == 0 and pad_width == 0:
|
||||
return image
|
||||
|
||||
padding = (0, pad_width, 0, pad_height)
|
||||
return torch.nn.functional.pad(image, padding, mode='constant', value=0.0)
|
||||
|
||||
def cut_videos(videos):
|
||||
t = videos.size(1)
|
||||
if t < 1:
|
||||
raise ValueError("SeedVR2Preprocess expected at least one frame.")
|
||||
if t == 1:
|
||||
return videos
|
||||
if t <= 4:
|
||||
padding = videos[:, -1:].repeat(1, 4 - t + 1, 1, 1, 1)
|
||||
return torch.cat([videos, padding], dim=1)
|
||||
if (t - 1) % 4 == 0:
|
||||
return videos
|
||||
padding = videos[:, -1:].repeat(1, 4 - ((t - 1) % 4), 1, 1, 1)
|
||||
videos = torch.cat([videos, padding], dim=1)
|
||||
if (videos.size(1) - 1) % 4 != 0:
|
||||
raise ValueError(f"SeedVR2Preprocess failed to pad video length to 4n+1; got {videos.size(1)} frames.")
|
||||
return videos
|
||||
|
||||
def _seedvr2_input_shorter_edge(images, node_name):
|
||||
if images.dim() == 4:
|
||||
return min(images.shape[1], images.shape[2])
|
||||
if images.dim() == 5:
|
||||
return min(images.shape[2], images.shape[3])
|
||||
raise ValueError(
|
||||
f"{node_name}: expected 4-D or 5-D IMAGE tensor, "
|
||||
f"got shape {tuple(images.shape)}"
|
||||
)
|
||||
|
||||
|
||||
def _seedvr2_pad(images, upscaled_shorter_edge, node_name):
|
||||
if upscaled_shorter_edge < 2:
|
||||
raise ValueError(
|
||||
f"{node_name}: input shorter edge must be at least 2 pixels; "
|
||||
f"got {upscaled_shorter_edge}."
|
||||
)
|
||||
if images.shape[-1] > 3:
|
||||
images = images[..., :3]
|
||||
if images.dim() == 4:
|
||||
# Comfy video components arrive as a 4-D IMAGE frame sequence:
|
||||
# (frames, H, W, C). SeedVR2 consumes that as one video.
|
||||
images = images.unsqueeze(0)
|
||||
elif images.dim() != 5:
|
||||
raise ValueError(
|
||||
f"{node_name}: expected 4-D or 5-D IMAGE tensor, "
|
||||
f"got shape {tuple(images.shape)}"
|
||||
)
|
||||
images = images.permute(0, 1, 4, 2, 3)
|
||||
|
||||
b, t, c, h, w = images.shape
|
||||
images = images.reshape(b * t, c, h, w)
|
||||
|
||||
images = torch.clamp(images, 0.0, 1.0)
|
||||
images = div_pad(images, (16, 16))
|
||||
_, _, new_h, new_w = images.shape
|
||||
|
||||
images = images.reshape(b, t, c, new_h, new_w)
|
||||
images = cut_videos(images)
|
||||
images_bthwc = images.permute(0, 1, 3, 4, 2).contiguous()
|
||||
|
||||
return io.NodeOutput(images_bthwc)
|
||||
|
||||
|
||||
class SeedVR2Preprocess(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SeedVR2Preprocess",
|
||||
display_name="Pre-Process SeedVR2 Input",
|
||||
category="image/pre-processors",
|
||||
description="Pad a resized image for SeedVR2 model. Alpha channel is dropped. The node Post-Process SeedVR2 Output re-applies it from the original resized image.",
|
||||
search_aliases=["seedvr2", "upscale", "video upscale", "pad", "preprocess"],
|
||||
inputs=[
|
||||
io.Image.Input("resized_images", tooltip="The resized image to process."),
|
||||
],
|
||||
outputs=[
|
||||
io.Image.Output("images", tooltip="The padded image for VAE encoding."),
|
||||
]
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, resized_images):
|
||||
upscaled_shorter_edge = _seedvr2_input_shorter_edge(resized_images, "SeedVR2Preprocess")
|
||||
return _seedvr2_pad(
|
||||
resized_images, upscaled_shorter_edge, "SeedVR2Preprocess",
|
||||
)
|
||||
|
||||
|
||||
class SeedVR2PostProcessing(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SeedVR2PostProcessing",
|
||||
display_name="Post-Process SeedVR2 Output",
|
||||
category="image/post-processors",
|
||||
description="Align the generated image with the original resized image and apply color correction.",
|
||||
search_aliases=["seedvr2", "upscale", "color correction", "color match", "postprocess"],
|
||||
inputs=[
|
||||
io.Image.Input("images", tooltip="The generated image to process."),
|
||||
io.Image.Input("original_resized_images", tooltip="The original resized image before pre-processing, used as reference."),
|
||||
io.Combo.Input("color_correction_method", options=["lab", "wavelet", "adain", "none"], default="lab", tooltip="Method to match the generated image colors to the original image. lab: transfer color in CIELAB space, preserving detail (most faithful). wavelet: transfer low-frequency color, keeping upscaled high-frequency detail. adain: match per-channel mean/std (fastest, global tint). none: skip color transfer (geometry alignment only)."),
|
||||
],
|
||||
outputs=[io.Image.Output(display_name="images", tooltip="The aligned, color-corrected image.")],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, images, original_resized_images, color_correction_method):
|
||||
alpha_input = None
|
||||
if original_resized_images.shape[-1] == 4:
|
||||
alpha_input = original_resized_images[..., 3:4]
|
||||
original_resized_images = original_resized_images[..., :3]
|
||||
decoded_5d, decoded_was_4d = cls._as_bthwc(images)
|
||||
reference_full, _ = cls._as_bthwc(original_resized_images)
|
||||
decoded_5d = cls._restore_reference_batch_time(decoded_5d, reference_full)
|
||||
|
||||
b = min(decoded_5d.shape[0], reference_full.shape[0])
|
||||
t = min(decoded_5d.shape[1], reference_full.shape[1])
|
||||
reference_h = reference_full.shape[2]
|
||||
reference_w = reference_full.shape[3]
|
||||
|
||||
decoded_5d = decoded_5d[:b, :t, :, :, :]
|
||||
target_h = min(decoded_5d.shape[2], reference_h)
|
||||
target_w = min(decoded_5d.shape[3], reference_w)
|
||||
decoded_5d = decoded_5d[:, :, :target_h, :target_w, :]
|
||||
if color_correction_method in ("lab", "wavelet", "adain"):
|
||||
reference_5d = reference_full[:b, :t, :, :, :]
|
||||
reference_5d = cls._resize_reference(reference_5d, target_h, target_w)
|
||||
output_device = decoded_5d.device
|
||||
decoded_raw = cls._to_seedvr2_raw(decoded_5d)
|
||||
reference_raw = cls._to_seedvr2_raw(reference_5d)
|
||||
decoded_flat = decoded_raw.permute(0, 1, 4, 2, 3).reshape(b * t, decoded_raw.shape[4], target_h, target_w)
|
||||
reference_flat = reference_raw.permute(0, 1, 4, 2, 3).reshape(b * t, reference_raw.shape[4], target_h, target_w)
|
||||
output = cls._color_transfer_chunked(
|
||||
decoded_flat, reference_flat, output_device, color_correction_method,
|
||||
)
|
||||
output = output.reshape(b, t, output.shape[1], output.shape[2], output.shape[3]).permute(0, 1, 3, 4, 2)
|
||||
output = output.add(1.0).div(2.0).clamp(0.0, 1.0)
|
||||
elif color_correction_method == "none":
|
||||
output = decoded_5d
|
||||
else:
|
||||
raise ValueError(f"SeedVR2PostProcessing: unknown color_correction_method {color_correction_method!r}")
|
||||
|
||||
if alpha_input is not None:
|
||||
alpha_5d, _ = cls._as_bthwc(alpha_input)
|
||||
alpha_5d = alpha_5d[:output.shape[0], :output.shape[1], :output.shape[2], :output.shape[3], :]
|
||||
output = torch.cat([output, alpha_5d.to(dtype=output.dtype, device=output.device)], dim=-1)
|
||||
h2 = output.shape[-3] - (output.shape[-3] % 2)
|
||||
w2 = output.shape[-2] - (output.shape[-2] % 2)
|
||||
output = output[:, :, :h2, :w2, :]
|
||||
if decoded_was_4d:
|
||||
output = output.reshape(-1, output.shape[-3], output.shape[-2], output.shape[-1])
|
||||
return io.NodeOutput(output)
|
||||
|
||||
@staticmethod
|
||||
def _as_bthwc(images):
|
||||
if images.ndim == 4:
|
||||
return images.unsqueeze(0), True
|
||||
if images.ndim == 5:
|
||||
return images, False
|
||||
raise ValueError(
|
||||
f"SeedVR2PostProcessing: expected 4-D or 5-D IMAGE tensor, got shape {tuple(images.shape)}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _restore_reference_batch_time(decoded, reference):
|
||||
if decoded.shape[0] != 1:
|
||||
return decoded
|
||||
ref_b, ref_t = reference.shape[:2]
|
||||
if ref_b < 1 or decoded.shape[1] % ref_b != 0:
|
||||
return decoded
|
||||
decoded_t = decoded.shape[1] // ref_b
|
||||
if decoded_t < ref_t:
|
||||
return decoded
|
||||
return decoded.reshape(ref_b, decoded_t, decoded.shape[2], decoded.shape[3], decoded.shape[4])
|
||||
|
||||
@staticmethod
|
||||
def _to_seedvr2_raw(images):
|
||||
return images.mul(2.0).sub(1.0)
|
||||
|
||||
@staticmethod
|
||||
def _color_transfer_on_vae_device(decoded_flat, reference_flat, output_device, transfer_fn):
|
||||
color_device = comfy.model_management.vae_device()
|
||||
decoded_flat = decoded_flat.to(device=color_device)
|
||||
reference_flat = reference_flat.to(device=color_device)
|
||||
output = transfer_fn(decoded_flat, reference_flat)
|
||||
return output.to(device=output_device)
|
||||
|
||||
@staticmethod
|
||||
def _lab_color_transfer_on_vae_device(decoded_flat, reference_flat, output_device):
|
||||
color_device = comfy.model_management.vae_device()
|
||||
result = None
|
||||
for start in range(decoded_flat.shape[0]):
|
||||
decoded_frame = decoded_flat[start:start + 1].to(device=color_device).clone()
|
||||
reference_frame = reference_flat[start:start + 1].to(device=color_device).clone()
|
||||
output = lab_color_transfer(decoded_frame, reference_frame).to(device=output_device)
|
||||
if result is None:
|
||||
result = torch.empty(
|
||||
(decoded_flat.shape[0],) + tuple(output.shape[1:]),
|
||||
device=output_device,
|
||||
dtype=output.dtype,
|
||||
)
|
||||
result[start:start + 1].copy_(output)
|
||||
if result is None:
|
||||
raise ValueError("SeedVR2PostProcessing: LAB color correction requires at least one frame.")
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def _color_transfer_chunked(cls, decoded_flat, reference_flat, output_device, color_correction_method):
|
||||
chunk_size = cls._estimate_color_correction_chunk_size(decoded_flat, color_correction_method)
|
||||
while True:
|
||||
try:
|
||||
return cls._run_color_transfer_chunks(
|
||||
decoded_flat, reference_flat, output_device, color_correction_method, chunk_size,
|
||||
)
|
||||
except Exception as e:
|
||||
comfy.model_management.raise_non_oom(e)
|
||||
if chunk_size <= 1:
|
||||
raise RuntimeError(
|
||||
"SeedVR2PostProcessing: color correction OOM at one frame; "
|
||||
f"color_correction_method={color_correction_method}, shape={tuple(decoded_flat.shape)}."
|
||||
) from e
|
||||
chunk_size = max(1, chunk_size // SEEDVR2_OOM_BACKOFF_DIVISOR)
|
||||
|
||||
@classmethod
|
||||
def _run_color_transfer_chunks(cls, decoded_flat, reference_flat, output_device, color_correction_method, chunk_size):
|
||||
result = None
|
||||
for start in range(0, decoded_flat.shape[0], chunk_size):
|
||||
end = min(start + chunk_size, decoded_flat.shape[0])
|
||||
decoded_chunk = decoded_flat[start:end]
|
||||
reference_chunk = reference_flat[start:end]
|
||||
if color_correction_method == "lab":
|
||||
output = cls._lab_color_transfer_on_vae_device(decoded_chunk, reference_chunk, output_device)
|
||||
elif color_correction_method == "wavelet":
|
||||
output = cls._color_transfer_on_vae_device(
|
||||
decoded_chunk, reference_chunk, output_device, wavelet_color_transfer,
|
||||
)
|
||||
else:
|
||||
output = cls._color_transfer_on_vae_device(
|
||||
decoded_chunk, reference_chunk, output_device, adain_color_transfer,
|
||||
)
|
||||
if result is None:
|
||||
result = torch.empty(
|
||||
(decoded_flat.shape[0],) + tuple(output.shape[1:]),
|
||||
device=output_device,
|
||||
dtype=output.dtype,
|
||||
)
|
||||
result[start:end].copy_(output)
|
||||
if result is None:
|
||||
raise ValueError("SeedVR2PostProcessing: color correction requires at least one frame.")
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def _estimate_color_correction_chunk_size(cls, decoded_flat, color_correction_method):
|
||||
multiplier = cls._color_correction_memory_multiplier(color_correction_method)
|
||||
frames = decoded_flat.shape[0]
|
||||
_, channels, height, width = decoded_flat.shape
|
||||
dtype_bytes = max(decoded_flat.element_size(), SEEDVR2_DTYPE_BYTES_FLOOR)
|
||||
bytes_per_frame = height * width * channels * dtype_bytes * multiplier
|
||||
if bytes_per_frame <= 0:
|
||||
return frames
|
||||
color_device = comfy.model_management.vae_device()
|
||||
free_memory = comfy.model_management.get_free_memory(color_device)
|
||||
chunk_size = int((free_memory * SEEDVR2_COLOR_MEM_HEADROOM) // bytes_per_frame)
|
||||
return max(1, min(frames, chunk_size))
|
||||
|
||||
@staticmethod
|
||||
def _color_correction_memory_multiplier(color_correction_method):
|
||||
if color_correction_method == "lab":
|
||||
return SEEDVR2_LAB_SCALE_MULTIPLIER
|
||||
if color_correction_method == "wavelet":
|
||||
return SEEDVR2_WAVELET_SCALE_MULTIPLIER
|
||||
if color_correction_method == "adain":
|
||||
return SEEDVR2_ADAIN_SCALE_MULTIPLIER
|
||||
raise ValueError(f"SeedVR2PostProcessing: unknown color_correction_method {color_correction_method!r}")
|
||||
|
||||
@staticmethod
|
||||
def _resize_reference(reference, height, width):
|
||||
if reference.shape[2] == height and reference.shape[3] == width:
|
||||
return reference
|
||||
b, t = reference.shape[:2]
|
||||
reference_flat = reference.permute(0, 1, 4, 2, 3).reshape(b * t, reference.shape[4], reference.shape[2], reference.shape[3])
|
||||
resized = TVF.resize(
|
||||
reference_flat,
|
||||
size=(height, width),
|
||||
interpolation=InterpolationMode.BICUBIC,
|
||||
antialias=not (isinstance(reference_flat, torch.Tensor) and reference_flat.device.type == "mps"),
|
||||
)
|
||||
return resized.reshape(b, t, resized.shape[1], height, width).permute(0, 1, 3, 4, 2)
|
||||
|
||||
|
||||
class SeedVR2Conditioning(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SeedVR2Conditioning",
|
||||
display_name="Apply SeedVR2 Conditioning",
|
||||
category="model/conditioning",
|
||||
description="Build SeedVR2 positive/negative conditioning from a VAE latent.",
|
||||
search_aliases=["seedvr2", "upscale", "conditioning"],
|
||||
inputs=[
|
||||
io.Model.Input("model", tooltip="The SeedVR2 model."),
|
||||
io.Latent.Input("vae_conditioning", display_name="latent"),
|
||||
],
|
||||
outputs=[
|
||||
io.Conditioning.Output(display_name="positive", tooltip="The positive conditioning for sampling."),
|
||||
io.Conditioning.Output(display_name="negative", tooltip="The negative conditioning for sampling."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, model, vae_conditioning) -> io.NodeOutput:
|
||||
|
||||
vae_conditioning = vae_conditioning["samples"]
|
||||
if vae_conditioning.ndim != 5:
|
||||
raise ValueError(
|
||||
"SeedVR2Conditioning expects a 5-D VAE latent in Comfy "
|
||||
f"channel-first layout; got shape {tuple(vae_conditioning.shape)}."
|
||||
)
|
||||
if vae_conditioning.shape[1] != SEEDVR2_LATENT_CHANNELS:
|
||||
if vae_conditioning.shape[-1] == SEEDVR2_LATENT_CHANNELS:
|
||||
raise ValueError(
|
||||
"SeedVR2Conditioning expects SeedVR2 VAE latents in Comfy "
|
||||
f"channel-first layout (B, {SEEDVR2_LATENT_CHANNELS}, T, H, W); "
|
||||
f"got channel-last shape {tuple(vae_conditioning.shape)}."
|
||||
)
|
||||
raise ValueError(
|
||||
"SeedVR2Conditioning expects SeedVR2 VAE latents with "
|
||||
f"{SEEDVR2_LATENT_CHANNELS} channels; got shape {tuple(vae_conditioning.shape)}."
|
||||
)
|
||||
vae_conditioning = vae_conditioning.movedim(1, -1).contiguous()
|
||||
model = _resolve_seedvr2_diffusion_model(model)
|
||||
pos_cond = model.positive_conditioning
|
||||
neg_cond = model.negative_conditioning
|
||||
|
||||
mask = vae_conditioning.new_ones(vae_conditioning.shape[:-1] + (1,))
|
||||
condition = torch.cat((vae_conditioning, mask), dim=-1)
|
||||
condition = condition.movedim(-1, 1)
|
||||
|
||||
negative = [[neg_cond.unsqueeze(0), {"condition": condition}]]
|
||||
positive = [[pos_cond.unsqueeze(0), {"condition": condition}]]
|
||||
|
||||
return io.NodeOutput(positive, negative)
|
||||
|
||||
def _seedvr2_chunk_crossfade_weights(overlap, device, dtype):
|
||||
"""Descending previous-chunk weights across the overlap (next chunk gets ``1 - w``): a Hann fade over the middle third, flat shoulders on the outer thirds."""
|
||||
ramp = torch.linspace(0.0, 1.0, steps=overlap, device=device, dtype=dtype)
|
||||
ramp = ((ramp - 1.0 / 3.0) / (1.0 / 3.0)).clamp(0.0, 1.0)
|
||||
return 0.5 + 0.5 * torch.cos(torch.pi * ramp)
|
||||
|
||||
|
||||
class SeedVR2TemporalChunk(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SeedVR2TemporalChunk",
|
||||
display_name="Split SeedVR2 Latent",
|
||||
category="model/latent/batch",
|
||||
description="Split a SeedVR2 video latent into overlapping temporal chunks small enough to sample one at a time within VRAM, wiring latents outputs to both Apply SeedVR2 Conditioning and the sampler latent input before recombining with Merge SeedVR2 Latents.",
|
||||
search_aliases=["seedvr2", "split", "chunk", "temporal", "video upscale", "rebatch"],
|
||||
inputs=[
|
||||
io.Latent.Input("latent", tooltip="The VAE-encoded SeedVR2 latent to split."),
|
||||
io.Int.Input("temporal_overlap", default=0, min=0, max=16384,
|
||||
tooltip="Latent frames shared between adjacent chunks and crossfaded at merge; 0 = no overlap."),
|
||||
io.DynamicCombo.Input("chunking_mode",
|
||||
tooltip="manual = use frames_per_chunk exactly; auto = predict the largest chunk that fits free VRAM.",
|
||||
options=[
|
||||
io.DynamicCombo.Option("auto", []),
|
||||
io.DynamicCombo.Option("manual", [
|
||||
io.Int.Input("frames_per_chunk", default=21, min=1, max=16384, step=4,
|
||||
tooltip="Pixel frames per temporal chunk (4n+1: 1, 5, 9, 13, ...)."),
|
||||
]),
|
||||
]),
|
||||
],
|
||||
outputs=[
|
||||
io.Latent.Output(display_name="latents", is_output_list=True,
|
||||
tooltip="The temporal chunks in sequence order."),
|
||||
io.Int.Output(display_name="temporal_overlap",
|
||||
tooltip="The effective latent-frame overlap between adjacent chunks, for Merge SeedVR2 Latents."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, latent, temporal_overlap, chunking_mode) -> io.NodeOutput:
|
||||
samples = latent["samples"]
|
||||
if samples.ndim != 5:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalChunk: expected a 5-D video latent (B, C, T, H, W); "
|
||||
f"got shape {tuple(samples.shape)}."
|
||||
)
|
||||
if samples.shape[1] != SEEDVR2_LATENT_CHANNELS:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalChunk: expected {SEEDVR2_LATENT_CHANNELS} latent channels; "
|
||||
f"got shape {tuple(samples.shape)}."
|
||||
)
|
||||
if temporal_overlap < 0:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalChunk: temporal_overlap must be >= 0; got {temporal_overlap}."
|
||||
)
|
||||
mode = chunking_mode["chunking_mode"]
|
||||
if mode not in ("auto", "manual"):
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalChunk: chunking_mode must be 'auto' or 'manual'; "
|
||||
f"got {mode!r}."
|
||||
)
|
||||
t_latent = samples.shape[2]
|
||||
t_pixel = 4 * (t_latent - 1) + 1
|
||||
|
||||
if mode == "auto":
|
||||
free_gb = comfy.model_management.get_free_memory(
|
||||
comfy.model_management.get_torch_device()) / (1024 ** 3)
|
||||
mpx_per_frame = (samples.shape[0] * samples.shape[3] * samples.shape[4]) * (BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE ** 2) / 1e6
|
||||
budget_gb = free_gb - SEEDVR2_CHUNK_RESERVED_GIB - SEEDVR2_CHUNK_SIGMA_K * SEEDVR2_CHUNK_SIGMA_GIB
|
||||
chunk_latent_max = max(1, int(budget_gb / (SEEDVR2_CHUNK_GIB_PER_MPX_FRAME * mpx_per_frame)))
|
||||
frames_per_chunk = min(4 * (chunk_latent_max - 1) + 1, t_pixel)
|
||||
logging.info(
|
||||
"SeedVR2TemporalChunk auto: free=%.2fGiB, %.2fMpx -> frames_per_chunk=%d (t_pixel=%d).",
|
||||
free_gb, mpx_per_frame, frames_per_chunk, t_pixel,
|
||||
)
|
||||
else:
|
||||
frames_per_chunk = chunking_mode["frames_per_chunk"]
|
||||
if frames_per_chunk < 1 or (frames_per_chunk - 1) % 4 != 0:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalChunk: frames_per_chunk must be a 4n+1 pixel-frame count "
|
||||
f"(1, 5, 9, 13, 17, 21, ...); got {frames_per_chunk}."
|
||||
)
|
||||
|
||||
if t_pixel <= frames_per_chunk:
|
||||
return io.NodeOutput([latent], 0)
|
||||
|
||||
chunk_latent = (frames_per_chunk - 1) // 4 + 1
|
||||
temporal_overlap = min(temporal_overlap, chunk_latent - 1)
|
||||
step = chunk_latent - temporal_overlap
|
||||
|
||||
chunks = []
|
||||
for start in range(0, t_latent, step):
|
||||
end = min(start + chunk_latent, t_latent)
|
||||
chunk = latent.copy()
|
||||
chunk["samples"] = samples[:, :, start:end].contiguous()
|
||||
chunks.append(chunk)
|
||||
if end >= t_latent:
|
||||
break
|
||||
return io.NodeOutput(chunks, temporal_overlap)
|
||||
|
||||
|
||||
class SeedVR2TemporalMerge(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SeedVR2TemporalMerge",
|
||||
display_name="Merge SeedVR2 Latents",
|
||||
category="model/latent/batch",
|
||||
is_input_list=True,
|
||||
description="Recombine sampled SeedVR2 latent temporal chunks into one latent, crossfading each overlap with a Hann window sized by the temporal_overlap wired from Split SeedVR2 Latent.",
|
||||
search_aliases=["seedvr2", "merge", "temporal", "hann", "crossfade"],
|
||||
inputs=[
|
||||
io.Latent.Input("latents", tooltip="The sampled temporal chunks in sequence order."),
|
||||
io.Int.Input("temporal_overlap", default=0, min=0, max=16384, force_input=True,
|
||||
tooltip="The temporal_overlap output of Split SeedVR2 Latent. 0 = plain concatenation."),
|
||||
],
|
||||
outputs=[
|
||||
io.Latent.Output(display_name="latent", tooltip="The recombined full-length latent."),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, latents, temporal_overlap) -> io.NodeOutput:
|
||||
temporal_overlap = temporal_overlap[0]
|
||||
if temporal_overlap < 0:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalMerge: temporal_overlap must be >= 0; got {temporal_overlap}."
|
||||
)
|
||||
chunks = [entry["samples"] for entry in latents]
|
||||
first = chunks[0]
|
||||
if first.ndim != 5:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalMerge: expected 5-D video latents (B, C, T, H, W); "
|
||||
f"chunk 0 has shape {tuple(first.shape)}."
|
||||
)
|
||||
for i, chunk in enumerate(chunks[1:], start=1):
|
||||
if chunk.shape[:2] != first.shape[:2] or chunk.shape[3:] != first.shape[3:]:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalMerge: chunk {i} shape {tuple(chunk.shape)} does not "
|
||||
f"match chunk 0 shape {tuple(first.shape)} outside the temporal axis."
|
||||
)
|
||||
if i < len(chunks) - 1 and chunk.shape[2] != first.shape[2]:
|
||||
raise ValueError(
|
||||
f"SeedVR2TemporalMerge: chunk {i} has {chunk.shape[2]} latent frames but "
|
||||
f"chunk 0 has {first.shape[2]}; only the final chunk may be shorter."
|
||||
)
|
||||
|
||||
out = latents[0].copy()
|
||||
out.pop("noise_mask", None)
|
||||
|
||||
if len(chunks) == 1:
|
||||
out["samples"] = first
|
||||
return io.NodeOutput(out)
|
||||
if temporal_overlap == 0:
|
||||
out["samples"] = torch.cat(chunks, dim=2)
|
||||
return io.NodeOutput(out)
|
||||
|
||||
chunk_latent = first.shape[2]
|
||||
step = chunk_latent - min(temporal_overlap, chunk_latent - 1)
|
||||
t_total = step * (len(chunks) - 1) + chunks[-1].shape[2]
|
||||
b, c, _, h, w = first.shape
|
||||
merged = torch.empty((b, c, t_total, h, w), device=first.device, dtype=first.dtype)
|
||||
|
||||
merged[:, :, :chunk_latent] = first
|
||||
filled = chunk_latent
|
||||
for i, chunk in enumerate(chunks[1:], start=1):
|
||||
start = i * step
|
||||
end = start + chunk.shape[2]
|
||||
# Crossfade width is bounded by the previous fill frontier and by a runt
|
||||
# final chunk shorter than the configured overlap.
|
||||
fade = min(filled - start, chunk.shape[2])
|
||||
if fade > 0:
|
||||
w_prev = _seedvr2_chunk_crossfade_weights(
|
||||
fade, chunk.device, chunk.dtype).view(1, 1, fade, 1, 1)
|
||||
merged[:, :, start:start + fade] = (
|
||||
merged[:, :, start:start + fade] * w_prev + chunk[:, :, :fade] * (1.0 - w_prev)
|
||||
)
|
||||
merged[:, :, start + fade:end] = chunk[:, :, fade:]
|
||||
else:
|
||||
merged[:, :, start:end] = chunk
|
||||
filled = end
|
||||
|
||||
out["samples"] = merged
|
||||
return io.NodeOutput(out)
|
||||
|
||||
|
||||
class SeedVRExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [
|
||||
SeedVR2Conditioning,
|
||||
SeedVR2Preprocess,
|
||||
SeedVR2PostProcessing,
|
||||
SeedVR2TemporalChunk,
|
||||
SeedVR2TemporalMerge,
|
||||
]
|
||||
|
||||
async def comfy_entrypoint() -> SeedVRExtension:
|
||||
return SeedVRExtension()
|
||||
@@ -0,0 +1,71 @@
|
||||
import os
|
||||
import json
|
||||
from typing_extensions import override
|
||||
from comfy_api.latest import io, ComfyExtension, ui
|
||||
import folder_paths
|
||||
|
||||
|
||||
class SaveTextNode(io.ComfyNode):
|
||||
"""Save text content to .txt, .md, or .json."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="SaveText",
|
||||
search_aliases=["save text", "write text", "export text"],
|
||||
display_name="Save Text",
|
||||
category="text",
|
||||
description="Save text content to a file in the output directory.",
|
||||
inputs=[
|
||||
io.String.Input("text", force_input=True),
|
||||
io.String.Input("filename_prefix", default="ComfyUI"),
|
||||
io.Combo.Input("format", options=["txt", "md", "json"], default="txt"),
|
||||
],
|
||||
outputs=[io.String.Output(display_name="text")],
|
||||
is_output_node=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, text, filename_prefix, format):
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
|
||||
filename_prefix,
|
||||
folder_paths.get_output_directory(),
|
||||
1,
|
||||
1,
|
||||
)
|
||||
|
||||
file = f"{filename}_{counter:05}.{format}"
|
||||
filepath = os.path.join(full_output_folder, file)
|
||||
|
||||
if format == "json":
|
||||
# tries to pretty print otherwise saves normally
|
||||
try:
|
||||
data = json.loads(text)
|
||||
with open(filepath, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, indent=2, ensure_ascii=False)
|
||||
except json.JSONDecodeError:
|
||||
with open(filepath, "w", encoding="utf-8") as f:
|
||||
f.write(text)
|
||||
else:
|
||||
with open(filepath, "w", encoding="utf-8") as f:
|
||||
f.write(text)
|
||||
|
||||
return io.NodeOutput(
|
||||
text,
|
||||
ui={
|
||||
"text": (text,),
|
||||
"files": [
|
||||
ui.SavedResult(file, subfolder, io.FolderType.output)
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
class TextExtension(ComfyExtension):
|
||||
@override
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return [
|
||||
SaveTextNode
|
||||
]
|
||||
|
||||
async def comfy_entrypoint() -> TextExtension:
|
||||
return TextExtension()
|
||||
@@ -81,7 +81,7 @@ class SaveVideo(io.ComfyNode):
|
||||
display_name="Save Video",
|
||||
category="video",
|
||||
essentials_category="Basics",
|
||||
description="Saves the input images to your ComfyUI output directory.",
|
||||
description="Saves the input videos to your ComfyUI output directory.",
|
||||
inputs=[
|
||||
io.Video.Input("video", tooltip="The video to save."),
|
||||
io.String.Input("filename_prefix", default="video/ComfyUI", tooltip="The prefix for the file to save. This may include formatting information such as %date:yyyy-MM-dd% or %Empty Latent Image.width% to include values from nodes."),
|
||||
|
||||
+222
-27
@@ -29,11 +29,16 @@ from comfy_execution.caching import (
|
||||
HierarchicalCache,
|
||||
LRUCache,
|
||||
RAMPressureCache,
|
||||
RAM_CACHE_LARGE_INTERMEDIATE,
|
||||
)
|
||||
from comfy_execution.graph import (
|
||||
DependencyCycleError,
|
||||
DynamicPrompt,
|
||||
ExecutionBlocker,
|
||||
ExecutionFailureBlocker,
|
||||
ExecutionList,
|
||||
NodeInputError,
|
||||
NodeNotFoundError,
|
||||
get_input_info,
|
||||
)
|
||||
from comfy_execution.graph_utils import GraphBuilder, is_link
|
||||
@@ -50,6 +55,16 @@ class ExecutionResult(Enum):
|
||||
SUCCESS = 0
|
||||
FAILURE = 1
|
||||
PENDING = 2
|
||||
BLOCKED = 3
|
||||
|
||||
|
||||
NODE_FAILURE_POLICY_FAIL_FAST = "fail_fast"
|
||||
NODE_FAILURE_POLICY_CONTINUE_INDEPENDENT = "continue_independent"
|
||||
NODE_FAILURE_POLICIES = frozenset({
|
||||
NODE_FAILURE_POLICY_FAIL_FAST,
|
||||
NODE_FAILURE_POLICY_CONTINUE_INDEPENDENT,
|
||||
})
|
||||
NODE_FAILURE_POLICY_EXTRA_DATA_KEY = "_node_failure_policy"
|
||||
|
||||
class DuplicateNodeError(Exception):
|
||||
pass
|
||||
@@ -103,6 +118,47 @@ class CacheEntry(NamedTuple):
|
||||
outputs: list
|
||||
|
||||
|
||||
def _failure_blocker_state(value):
|
||||
if isinstance(value, ExecutionFailureBlocker):
|
||||
return True, False
|
||||
if isinstance(value, dict):
|
||||
values = value.values()
|
||||
elif isinstance(value, (list, tuple)):
|
||||
values = value
|
||||
else:
|
||||
return False, True
|
||||
|
||||
has_failure_blocker = False
|
||||
has_normal_value = False
|
||||
for child in values:
|
||||
child_failure, child_normal = _failure_blocker_state(child)
|
||||
has_failure_blocker = has_failure_blocker or child_failure
|
||||
has_normal_value = has_normal_value or child_normal
|
||||
if has_failure_blocker and has_normal_value:
|
||||
break
|
||||
return has_failure_blocker, has_normal_value
|
||||
|
||||
|
||||
def _tag_node_raised(ex):
|
||||
# Marks an exception as raised by node code (vs execution machinery). Best effort: an
|
||||
# exception class rejecting attribute assignment must not mask the original error.
|
||||
try:
|
||||
ex._node_raised = True
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _is_recoverable_node_failure(ex):
|
||||
if isinstance(ex, (
|
||||
comfy.model_management.InterruptProcessingException,
|
||||
DependencyCycleError,
|
||||
NodeInputError,
|
||||
NodeNotFoundError,
|
||||
)):
|
||||
return False
|
||||
return not comfy.model_management.is_oom(ex)
|
||||
|
||||
|
||||
class CacheType(Enum):
|
||||
CLASSIC = 0
|
||||
LRU = 1
|
||||
@@ -151,7 +207,7 @@ class CacheSet:
|
||||
}
|
||||
return result
|
||||
|
||||
SENSITIVE_EXTRA_DATA_KEYS = ("auth_token_comfy_org", "api_key_comfy_org")
|
||||
SENSITIVE_EXTRA_DATA_KEYS = ("auth_token_comfy_org", "api_key_comfy_org", NODE_FAILURE_POLICY_EXTRA_DATA_KEY)
|
||||
|
||||
def get_input_data(inputs, class_def, unique_id, execution_list=None, dynprompt=None, extra_data={}):
|
||||
is_v3 = issubclass(class_def, _ComfyNodeInternal)
|
||||
@@ -254,16 +310,22 @@ async def _async_map_node_over_list(prompt_id, unique_id, obj, input_data_all, f
|
||||
async def process_inputs(inputs, index=None, input_is_list=False):
|
||||
if allow_interrupt:
|
||||
nodes.before_node_execution()
|
||||
execution_block = None
|
||||
# Prefer a runtime-failure blocker over a plain user blocker so failure
|
||||
# taint and retry semantics are not masked by input ordering.
|
||||
blocked_value = None
|
||||
for k, v in inputs.items():
|
||||
if input_is_list:
|
||||
for e in v:
|
||||
if isinstance(e, ExecutionBlocker):
|
||||
v = e
|
||||
values = v if input_is_list else (v,)
|
||||
for e in values:
|
||||
if isinstance(e, ExecutionBlocker):
|
||||
if blocked_value is None or isinstance(e, ExecutionFailureBlocker):
|
||||
blocked_value = e
|
||||
if isinstance(blocked_value, ExecutionFailureBlocker):
|
||||
break
|
||||
if isinstance(v, ExecutionBlocker):
|
||||
execution_block = execution_block_cb(v) if execution_block_cb else v
|
||||
if isinstance(blocked_value, ExecutionFailureBlocker):
|
||||
break
|
||||
execution_block = None
|
||||
if blocked_value is not None:
|
||||
execution_block = execution_block_cb(blocked_value) if execution_block_cb else blocked_value
|
||||
if execution_block is None:
|
||||
if pre_execute_cb is not None and index is not None:
|
||||
pre_execute_cb(index)
|
||||
@@ -289,7 +351,11 @@ async def _async_map_node_over_list(prompt_id, unique_id, obj, input_data_all, f
|
||||
if inspect.iscoroutinefunction(f):
|
||||
async def async_wrapper(f, prompt_id, unique_id, list_index, args):
|
||||
with CurrentNodeContext(prompt_id, unique_id, list_index):
|
||||
return await f(**args)
|
||||
try:
|
||||
return await f(**args)
|
||||
except Exception as ex:
|
||||
_tag_node_raised(ex)
|
||||
raise
|
||||
task = asyncio.create_task(async_wrapper(f, prompt_id, unique_id, index, args=inputs))
|
||||
# Give the task a chance to execute without yielding
|
||||
await asyncio.sleep(0)
|
||||
@@ -300,7 +366,11 @@ async def _async_map_node_over_list(prompt_id, unique_id, obj, input_data_all, f
|
||||
results.append(task)
|
||||
else:
|
||||
with CurrentNodeContext(prompt_id, unique_id, index):
|
||||
result = f(**inputs)
|
||||
try:
|
||||
result = f(**inputs)
|
||||
except Exception as ex:
|
||||
_tag_node_raised(ex)
|
||||
raise
|
||||
results.append(result)
|
||||
else:
|
||||
results.append(execution_block)
|
||||
@@ -447,11 +517,16 @@ async def execute(server, dynprompt, caches, current_item, extra_data, executed,
|
||||
execution_list.cache_update(unique_id, cached)
|
||||
return (ExecutionResult.SUCCESS, None, None)
|
||||
|
||||
continue_on_failure = extra_data.get(NODE_FAILURE_POLICY_EXTRA_DATA_KEY) == NODE_FAILURE_POLICY_CONTINUE_INDEPENDENT
|
||||
input_data_all = None
|
||||
failure_blocked_invocations = 0
|
||||
successful_invocations = 0
|
||||
resumed_subgraph = False
|
||||
try:
|
||||
if unique_id in pending_async_nodes:
|
||||
pending_results, failure_blocked_invocations, successful_invocations = pending_async_nodes[unique_id]
|
||||
results = []
|
||||
for r in pending_async_nodes[unique_id]:
|
||||
for r in pending_results:
|
||||
if isinstance(r, asyncio.Task):
|
||||
try:
|
||||
results.append(r.result())
|
||||
@@ -464,7 +539,8 @@ async def execute(server, dynprompt, caches, current_item, extra_data, executed,
|
||||
del pending_async_nodes[unique_id]
|
||||
output_data, output_ui, has_subgraph = get_output_from_returns(results, class_def)
|
||||
elif unique_id in pending_subgraph_results:
|
||||
cached_results = pending_subgraph_results[unique_id]
|
||||
cached_results, failure_blocked_invocations = pending_subgraph_results[unique_id]
|
||||
resumed_subgraph = True
|
||||
resolved_outputs = []
|
||||
for is_subgraph, result in cached_results:
|
||||
if not is_subgraph:
|
||||
@@ -517,6 +593,9 @@ async def execute(server, dynprompt, caches, current_item, extra_data, executed,
|
||||
return (ExecutionResult.PENDING, None, None)
|
||||
|
||||
def execution_block_cb(block):
|
||||
nonlocal failure_blocked_invocations
|
||||
if isinstance(block, ExecutionFailureBlocker):
|
||||
failure_blocked_invocations += 1
|
||||
if block.message is not None:
|
||||
mes = {
|
||||
"prompt_id": prompt_id,
|
||||
@@ -535,6 +614,8 @@ async def execute(server, dynprompt, caches, current_item, extra_data, executed,
|
||||
else:
|
||||
return block
|
||||
def pre_execute_cb(call_index):
|
||||
nonlocal successful_invocations
|
||||
successful_invocations += 1
|
||||
# TODO - How to handle this with async functions without contextvars (which requires Python 3.12)?
|
||||
GraphBuilder.set_default_prefix(unique_id, call_index, 0)
|
||||
|
||||
@@ -549,7 +630,7 @@ async def execute(server, dynprompt, caches, current_item, extra_data, executed,
|
||||
comfy_aimdo.model_vbar.vbars_reset_watermark_limits()
|
||||
|
||||
if has_pending_tasks:
|
||||
pending_async_nodes[unique_id] = output_data
|
||||
pending_async_nodes[unique_id] = (output_data, failure_blocked_invocations, successful_invocations)
|
||||
unblock = execution_list.add_external_block(unique_id)
|
||||
async def await_completion():
|
||||
tasks = [x for x in output_data if isinstance(x, asyncio.Task)]
|
||||
@@ -604,14 +685,26 @@ async def execute(server, dynprompt, caches, current_item, extra_data, executed,
|
||||
for node_id in new_output_ids:
|
||||
execution_list.add_node(node_id)
|
||||
execution_list.cache_link(node_id, unique_id)
|
||||
if continue_on_failure:
|
||||
# Wait for expanded output nodes before finishing, so a failure inside the expansion taints
|
||||
# this node before its outputs can be cached as reusable.
|
||||
execution_list.add_completion_link(node_id, unique_id)
|
||||
for link in new_output_links:
|
||||
execution_list.add_strong_link(link[0], link[1], unique_id)
|
||||
pending_subgraph_results[unique_id] = cached_outputs
|
||||
pending_subgraph_results[unique_id] = (cached_outputs, failure_blocked_invocations)
|
||||
return (ExecutionResult.PENDING, None, None)
|
||||
|
||||
cache_entry = CacheEntry(ui=ui_outputs.get(unique_id), outputs=output_data)
|
||||
execution_list.cache_update(unique_id, cache_entry)
|
||||
await caches.outputs.set(unique_id, cache_entry)
|
||||
if continue_on_failure:
|
||||
has_failure_blocker, has_normal_output = _failure_blocker_state(output_data)
|
||||
else:
|
||||
has_failure_blocker, has_normal_output = False, True
|
||||
failure_tainted = has_failure_blocker or failure_blocked_invocations > 0 or execution_list.is_failure_tainted(unique_id)
|
||||
if failure_tainted:
|
||||
execution_list.mark_failure_tainted(unique_id)
|
||||
execution_list.cache_update(unique_id, cache_entry, transient=failure_tainted)
|
||||
if not failure_tainted:
|
||||
await caches.outputs.set(unique_id, cache_entry)
|
||||
|
||||
except comfy.model_management.InterruptProcessingException as iex:
|
||||
logging.info("Processing interrupted")
|
||||
@@ -648,11 +741,23 @@ async def execute(server, dynprompt, caches, current_item, extra_data, executed,
|
||||
"exception_message": "{}\n{}".format(ex, tips),
|
||||
"exception_type": exception_type,
|
||||
"traceback": traceback.format_tb(tb),
|
||||
"current_inputs": input_data_formatted
|
||||
"current_inputs": input_data_formatted,
|
||||
"node_raised": getattr(ex, "_node_raised", False),
|
||||
}
|
||||
|
||||
return (ExecutionResult.FAILURE, error_details, ex)
|
||||
|
||||
if resumed_subgraph:
|
||||
fully_failure_blocked = has_failure_blocker and not has_normal_output
|
||||
else:
|
||||
fully_failure_blocked = successful_invocations == 0 and (
|
||||
failure_blocked_invocations > 0
|
||||
or (has_failure_blocker and not has_normal_output)
|
||||
)
|
||||
if fully_failure_blocked:
|
||||
get_progress_state().block_progress(unique_id)
|
||||
return (ExecutionResult.BLOCKED, None, None)
|
||||
|
||||
get_progress_state().finish_progress(unique_id)
|
||||
executed.add(unique_id)
|
||||
|
||||
@@ -669,6 +774,7 @@ class PromptExecutor:
|
||||
self.caches = CacheSet(cache_type=self.cache_type, cache_args=self.cache_args)
|
||||
self.status_messages = []
|
||||
self.success = True
|
||||
self.execution_summary = None
|
||||
|
||||
def add_message(self, event, data: dict, broadcast: bool):
|
||||
data = {
|
||||
@@ -707,6 +813,21 @@ class PromptExecutor:
|
||||
}
|
||||
self.add_message("execution_error", mes, broadcast=False)
|
||||
|
||||
def handle_node_execution_error(self, prompt_id, prompt, current_outputs, executed, error):
|
||||
node_id = error["node_id"]
|
||||
mes = {
|
||||
"prompt_id": prompt_id,
|
||||
"node_id": node_id,
|
||||
"node_type": prompt[node_id]["class_type"],
|
||||
"executed": list(executed),
|
||||
"exception_message": error["exception_message"],
|
||||
"exception_type": error["exception_type"],
|
||||
"traceback": error["traceback"],
|
||||
"current_inputs": error["current_inputs"],
|
||||
"current_outputs": list(current_outputs),
|
||||
}
|
||||
self.add_message("execution_node_error", mes, broadcast=False)
|
||||
|
||||
def _notify_prompt_lifecycle(self, event: str, prompt_id: str):
|
||||
if not _has_cache_providers():
|
||||
return
|
||||
@@ -727,6 +848,8 @@ class PromptExecutor:
|
||||
set_preview_method(extra_data.get("preview_method"))
|
||||
|
||||
nodes.interrupt_processing(False)
|
||||
self.success = True
|
||||
self.execution_summary = None
|
||||
|
||||
if "client_id" in extra_data:
|
||||
self.server.client_id = extra_data["client_id"]
|
||||
@@ -771,35 +894,75 @@ class PromptExecutor:
|
||||
executed = set()
|
||||
execution_list = ExecutionList(dynamic_prompt, self.caches.outputs)
|
||||
current_outputs = self.caches.outputs.all_node_ids()
|
||||
output_targets = set(execute_outputs)
|
||||
failed_node_ids = set()
|
||||
blocked_node_ids = set()
|
||||
blocked_output_node_ids = set()
|
||||
successful_output_node_ids = set()
|
||||
node_failures = []
|
||||
continue_independent = extra_data.get(
|
||||
NODE_FAILURE_POLICY_EXTRA_DATA_KEY,
|
||||
NODE_FAILURE_POLICY_FAIL_FAST,
|
||||
) == NODE_FAILURE_POLICY_CONTINUE_INDEPENDENT
|
||||
for node_id in list(execute_outputs):
|
||||
execution_list.add_node(node_id)
|
||||
|
||||
while not execution_list.is_empty():
|
||||
node_id, error, ex = await execution_list.stage_node_execution()
|
||||
if error is not None:
|
||||
self.success = False
|
||||
self.handle_execution_error(prompt_id, dynamic_prompt.original_prompt, current_outputs, executed, error, ex)
|
||||
break
|
||||
|
||||
assert node_id is not None, "Node ID should not be None at this point"
|
||||
result, error, ex = await execute(self.server, dynamic_prompt, self.caches, node_id, extra_data, executed, prompt_id, execution_list, pending_subgraph_results, pending_async_nodes, ui_node_outputs)
|
||||
self.success = result != ExecutionResult.FAILURE
|
||||
if result == ExecutionResult.FAILURE:
|
||||
self.handle_execution_error(prompt_id, dynamic_prompt.original_prompt, current_outputs, executed, error, ex)
|
||||
break
|
||||
if continue_independent and error.get("node_raised") and _is_recoverable_node_failure(ex):
|
||||
self.handle_node_execution_error(prompt_id, dynamic_prompt.original_prompt, current_outputs, executed, error)
|
||||
real_node_id = error["node_id"]
|
||||
failed_node_ids.add(real_node_id)
|
||||
node_failures.append((error, ex))
|
||||
blocker = ExecutionFailureBlocker(real_node_id)
|
||||
class_type = dynamic_prompt.get_node(node_id)["class_type"]
|
||||
class_def = nodes.NODE_CLASS_MAPPINGS[class_type]
|
||||
cache_entry = CacheEntry(
|
||||
ui=None,
|
||||
outputs=[[blocker] for _ in class_def.RETURN_TYPES],
|
||||
)
|
||||
execution_list.cache_update(node_id, cache_entry, transient=True)
|
||||
execution_list.mark_failure_tainted(node_id)
|
||||
get_progress_state().error_progress(node_id)
|
||||
execution_list.complete_node_execution()
|
||||
else:
|
||||
self.success = False
|
||||
self.handle_execution_error(prompt_id, dynamic_prompt.original_prompt, current_outputs, executed, error, ex)
|
||||
break
|
||||
elif result == ExecutionResult.PENDING:
|
||||
execution_list.unstage_node_execution()
|
||||
elif result == ExecutionResult.BLOCKED:
|
||||
real_node_id = dynamic_prompt.get_real_node_id(node_id)
|
||||
blocked_node_ids.add(real_node_id)
|
||||
if node_id in output_targets:
|
||||
blocked_output_node_ids.add(real_node_id)
|
||||
execution_list.complete_node_execution()
|
||||
else: # result == ExecutionResult.SUCCESS:
|
||||
if node_id in output_targets:
|
||||
successful_output_node_ids.add(dynamic_prompt.get_real_node_id(node_id))
|
||||
execution_list.complete_node_execution()
|
||||
|
||||
if self.cache_type == CacheType.RAM_PRESSURE:
|
||||
ram_release_callback(ram_inactive_headroom)
|
||||
ram_shortfall = ram_headroom - psutil.virtual_memory().available
|
||||
freed = comfy.model_management.free_pins(ram_shortfall + 512 * (1024 ** 2))
|
||||
if freed < ram_shortfall:
|
||||
if freed > 64 * (1024 ** 2):
|
||||
# AIMDO MEM_DECOMMIT can outrun psutil.available catching up.
|
||||
time.sleep(0.05)
|
||||
ram_release_callback(ram_headroom, free_active=True)
|
||||
if ram_shortfall > 0:
|
||||
freed = ram_release_callback(ram_headroom, free_active=True, min_entry_size=RAM_CACHE_LARGE_INTERMEDIATE)
|
||||
ram_shortfall -= freed
|
||||
if comfy.model_management.should_free_pins_for_ram_pressure(ram_shortfall):
|
||||
freed = comfy.model_management.free_pins(ram_shortfall + 512 * (1024 ** 2))
|
||||
if freed < ram_shortfall:
|
||||
if freed > 64 * (1024 ** 2):
|
||||
# AIMDO MEM_DECOMMIT can outrun psutil.available catching up.
|
||||
time.sleep(0.05)
|
||||
ram_release_callback(ram_headroom, free_active=True)
|
||||
else:
|
||||
# Only execute when the while-loop ends without break
|
||||
# Send cached UI for intermediate output nodes that weren't executed
|
||||
@@ -812,7 +975,36 @@ class PromptExecutor:
|
||||
if cached is not None:
|
||||
display_node_id = dynamic_prompt.get_display_node_id(node_id)
|
||||
_send_cached_ui(self.server, node_id, display_node_id, cached, prompt_id, ui_node_outputs)
|
||||
self.add_message("execution_success", { "prompt_id": prompt_id }, broadcast=False)
|
||||
|
||||
if node_failures:
|
||||
self.execution_summary = {
|
||||
"has_errors": True,
|
||||
"execution_error_count": len(node_failures),
|
||||
"failed_node_ids": sorted(failed_node_ids)[:100],
|
||||
"blocked_node_ids": sorted(blocked_node_ids)[:100],
|
||||
"blocked_output_node_ids": sorted(blocked_output_node_ids)[:100],
|
||||
"successful_output_node_ids": sorted(successful_output_node_ids)[:100],
|
||||
}
|
||||
if successful_output_node_ids:
|
||||
self.execution_summary["completion_status"] = "partial_success"
|
||||
self.add_message(
|
||||
"execution_success",
|
||||
{"prompt_id": prompt_id, **self.execution_summary},
|
||||
broadcast=False,
|
||||
)
|
||||
else:
|
||||
self.success = False
|
||||
last_error, last_ex = node_failures[-1]
|
||||
self.handle_execution_error(
|
||||
prompt_id,
|
||||
dynamic_prompt.original_prompt,
|
||||
current_outputs,
|
||||
executed,
|
||||
last_error,
|
||||
last_ex,
|
||||
)
|
||||
else:
|
||||
self.add_message("execution_success", { "prompt_id": prompt_id }, broadcast=False)
|
||||
|
||||
ui_outputs = {}
|
||||
meta_outputs = {}
|
||||
@@ -1270,6 +1462,7 @@ class PromptQueue:
|
||||
status_str: Literal['success', 'error']
|
||||
completed: bool
|
||||
messages: List[str]
|
||||
execution_summary: Optional[dict] = None
|
||||
|
||||
def task_done(self, item_id, history_result,
|
||||
status: Optional['PromptQueue.ExecutionStatus'], process_item=None):
|
||||
@@ -1281,6 +1474,8 @@ class PromptQueue:
|
||||
status_dict: Optional[dict] = None
|
||||
if status is not None:
|
||||
status_dict = copy.deepcopy(status._asdict())
|
||||
if status_dict.get("execution_summary") is None:
|
||||
del status_dict["execution_summary"]
|
||||
|
||||
if process_item is not None:
|
||||
prompt = process_item(prompt)
|
||||
|
||||
@@ -366,7 +366,8 @@ def prompt_worker(q, server_instance):
|
||||
status=execution.PromptQueue.ExecutionStatus(
|
||||
status_str='success' if e.success else 'error',
|
||||
completed=e.success,
|
||||
messages=e.status_messages), process_item=remove_sensitive)
|
||||
messages=e.status_messages,
|
||||
execution_summary=e.execution_summary), process_item=remove_sensitive)
|
||||
if server_instance.client_id is not None:
|
||||
server_instance.send_sync("executing", {"node": None, "prompt_id": prompt_id}, server_instance.client_id)
|
||||
|
||||
|
||||
@@ -1709,6 +1709,7 @@ class PreviewImage(SaveImage):
|
||||
self.compress_level = 1
|
||||
|
||||
SEARCH_ALIASES = ["preview", "preview image", "show image", "view image", "display image", "image viewer"]
|
||||
DESCRIPTION = "Preview the images without saving them to the ComfyUI output directory."
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -2458,6 +2459,7 @@ async def init_builtin_extra_nodes():
|
||||
"nodes_camera_trajectory.py",
|
||||
"nodes_edit_model.py",
|
||||
"nodes_tcfg.py",
|
||||
"nodes_seedvr.py",
|
||||
"nodes_context_windows.py",
|
||||
"nodes_qwen.py",
|
||||
"nodes_boogu.py",
|
||||
@@ -2503,6 +2505,7 @@ async def init_builtin_extra_nodes():
|
||||
"nodes_triposplat.py",
|
||||
"nodes_depth_anything_3.py",
|
||||
"nodes_seed.py",
|
||||
"nodes_text.py",
|
||||
]
|
||||
|
||||
import_failed = []
|
||||
|
||||
@@ -922,6 +922,13 @@ components:
|
||||
number:
|
||||
description: Priority number for the queue (lower numbers have higher priority)
|
||||
type: number
|
||||
node_failure_policy:
|
||||
default: fail_fast
|
||||
description: Controls whether a runtime node failure terminates the prompt or only blocks dependent nodes
|
||||
enum:
|
||||
- fail_fast
|
||||
- continue_independent
|
||||
type: string
|
||||
partial_execution_targets:
|
||||
description: List of node names to execute
|
||||
items:
|
||||
@@ -3297,6 +3304,12 @@ 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:
|
||||
|
||||
+3
-3
@@ -1,6 +1,6 @@
|
||||
comfyui-frontend-package==1.45.20
|
||||
comfyui-workflow-templates==0.11.6
|
||||
comfyui-embedded-docs==0.5.7
|
||||
comfyui-workflow-templates==0.11.9
|
||||
comfyui-embedded-docs==0.5.8
|
||||
torch
|
||||
torchsde
|
||||
torchvision
|
||||
@@ -22,7 +22,7 @@ alembic
|
||||
SQLAlchemy>=2.0.0
|
||||
filelock
|
||||
av>=16.0.0
|
||||
comfy-kitchen==0.2.16
|
||||
comfy-kitchen==0.2.19
|
||||
comfy-aimdo==0.4.10
|
||||
requests
|
||||
simpleeval>=1.0.0
|
||||
|
||||
@@ -39,6 +39,7 @@ from comfy.deploy_environment import get_deploy_environment
|
||||
import comfy.utils
|
||||
import comfy.model_management
|
||||
from comfy_api import feature_flags
|
||||
from comfy.comfy_api_env import get_environment_overrides
|
||||
import node_helpers
|
||||
from comfyui_version import __version__
|
||||
from app.frontend_management import FrontendManager, parse_version
|
||||
@@ -46,6 +47,7 @@ from comfy_api.internal import _ComfyNodeInternal
|
||||
from app.assets.seeder import asset_seeder
|
||||
from app.assets.api.routes import register_assets_routes
|
||||
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
|
||||
|
||||
from app.user_manager import UserManager
|
||||
@@ -441,7 +443,9 @@ class PromptServer():
|
||||
if args.enable_assets:
|
||||
try:
|
||||
tag = image_upload_type if image_upload_type in ("input", "output") else "input"
|
||||
result = register_file_in_place(abs_path=filepath, name=filename, tags=[tag])
|
||||
tags = [tag]
|
||||
tags.extend(get_known_subfolder_tags(subfolder))
|
||||
result = register_file_in_place(abs_path=filepath, name=filename, tags=tags)
|
||||
resp["asset"] = {
|
||||
"id": result.ref.id,
|
||||
"name": result.ref.name,
|
||||
@@ -724,7 +728,11 @@ class PromptServer():
|
||||
|
||||
@routes.get("/features")
|
||||
async def get_features(request):
|
||||
return web.json_response(feature_flags.get_server_features())
|
||||
features = feature_flags.get_server_features()
|
||||
overrides = get_environment_overrides()
|
||||
if overrides:
|
||||
features.update(overrides)
|
||||
return web.json_response(features)
|
||||
|
||||
@routes.get("/prompt")
|
||||
async def get_prompt(request):
|
||||
@@ -1089,6 +1097,19 @@ class PromptServer():
|
||||
if "partial_execution_targets" in json_data:
|
||||
partial_execution_targets = json_data["partial_execution_targets"]
|
||||
|
||||
node_failure_policy = json_data.get(
|
||||
"node_failure_policy",
|
||||
execution.NODE_FAILURE_POLICY_FAIL_FAST,
|
||||
)
|
||||
if not isinstance(node_failure_policy, str) or node_failure_policy not in execution.NODE_FAILURE_POLICIES:
|
||||
error = {
|
||||
"type": "invalid_node_failure_policy",
|
||||
"message": "node_failure_policy must be 'fail_fast' or 'continue_independent'",
|
||||
"details": f"Invalid node_failure_policy: {node_failure_policy!r}",
|
||||
"extra_info": {},
|
||||
}
|
||||
return web.json_response({"error": error, "node_errors": {}}, status=400)
|
||||
|
||||
self.node_replace_manager.apply_replacements(prompt)
|
||||
|
||||
valid = await execution.validate_prompt(prompt_id, prompt, partial_execution_targets)
|
||||
@@ -1096,6 +1117,10 @@ class PromptServer():
|
||||
if "extra_data" in json_data:
|
||||
extra_data = json_data["extra_data"]
|
||||
|
||||
extra_data.pop(execution.NODE_FAILURE_POLICY_EXTRA_DATA_KEY, None)
|
||||
if node_failure_policy == execution.NODE_FAILURE_POLICY_CONTINUE_INDEPENDENT:
|
||||
extra_data[execution.NODE_FAILURE_POLICY_EXTRA_DATA_KEY] = node_failure_policy
|
||||
|
||||
if "client_id" in json_data:
|
||||
extra_data["client_id"] = json_data["client_id"]
|
||||
|
||||
|
||||
@@ -24,6 +24,28 @@ def app(model_manager):
|
||||
app.add_routes(routes)
|
||||
return app
|
||||
|
||||
async def test_get_model_folders_includes_registered_extensions(aiohttp_client, app, tmp_path):
|
||||
"""Folders expose their registered extension set verbatim; an empty list
|
||||
means match-all (filter_files_extensions semantics)."""
|
||||
with patch('folder_paths.folder_names_and_paths', {
|
||||
'test_checkpoints': ([str(tmp_path)], {'.safetensors', '.ckpt'}),
|
||||
'test_configs': ([str(tmp_path)], ['.yaml']),
|
||||
'test_match_all': ([str(tmp_path)], set()),
|
||||
'configs': ([str(tmp_path)], ['.yaml']),
|
||||
}):
|
||||
client = await aiohttp_client(app)
|
||||
response = await client.get('/experiment/models')
|
||||
|
||||
assert response.status == 200
|
||||
folders = {f['name']: f for f in await response.json()}
|
||||
|
||||
assert 'configs' not in folders # blocklisted
|
||||
assert folders['test_checkpoints']['folders'] == [str(tmp_path)]
|
||||
assert folders['test_checkpoints']['extensions'] == ['.ckpt', '.safetensors']
|
||||
assert folders['test_configs']['extensions'] == ['.yaml']
|
||||
# Match-all registrations are exposed honestly, not substituted.
|
||||
assert folders['test_match_all']['extensions'] == []
|
||||
|
||||
async def test_get_model_preview_safetensors(aiohttp_client, app, tmp_path):
|
||||
img = Image.new('RGB', (100, 100), 'white')
|
||||
img_byte_arr = BytesIO()
|
||||
|
||||
@@ -8,6 +8,7 @@ upgrade/downgrade for 0003+.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sqlite3
|
||||
|
||||
import pytest
|
||||
from alembic import command
|
||||
@@ -30,6 +31,12 @@ def _make_config(db_path: str) -> Config:
|
||||
return cfg
|
||||
|
||||
|
||||
def _sqlite_path(cfg: Config) -> str:
|
||||
url = cfg.get_main_option("sqlalchemy.url")
|
||||
assert url is not None and url.startswith("sqlite:///")
|
||||
return url.removeprefix("sqlite:///")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def migration_db(tmp_path):
|
||||
"""Yield an alembic Config pre-upgraded to the baseline revision."""
|
||||
@@ -55,3 +62,26 @@ def test_upgrade_downgrade_cycle(migration_db):
|
||||
command.upgrade(migration_db, "head")
|
||||
command.downgrade(migration_db, _BASELINE)
|
||||
command.upgrade(migration_db, "head")
|
||||
|
||||
|
||||
def test_case_sensitive_tags_downgrade_normalizes_existing_tags(migration_db):
|
||||
"""Downgrading 0005 folds mixed-case tag vocabulary before restoring CHECK."""
|
||||
command.upgrade(migration_db, "0005_allow_case_sensitive_tags")
|
||||
|
||||
db_path = _sqlite_path(migration_db)
|
||||
with sqlite3.connect(db_path) as conn:
|
||||
conn.execute("INSERT INTO tags(name) VALUES (?)", ("NewTag",))
|
||||
conn.execute("INSERT INTO tags(name) VALUES (?)", ("newtag",))
|
||||
conn.execute("INSERT INTO tags(name) VALUES (?)", ("model_type:LLM",))
|
||||
|
||||
command.downgrade(migration_db, "0004_drop_tag_type")
|
||||
|
||||
with sqlite3.connect(db_path) as conn:
|
||||
tags = {row[0] for row in conn.execute("SELECT name FROM tags")}
|
||||
assert "newtag" in tags
|
||||
assert "model_type:llm" in tags
|
||||
assert "NewTag" not in tags
|
||||
assert "model_type:LLM" not in tags
|
||||
|
||||
with pytest.raises(sqlite3.IntegrityError):
|
||||
conn.execute("INSERT INTO tags(name) VALUES (?)", ("Upper",))
|
||||
|
||||
@@ -234,7 +234,7 @@ def seeded_asset(request: pytest.FixtureRequest, http: requests.Session, api_bas
|
||||
p = getattr(request, "param", {}) or {}
|
||||
tags: Optional[list[str]] = p.get("tags")
|
||||
if tags is None:
|
||||
tags = ["models", "checkpoints", "unit-tests", "alpha"]
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "alpha"]
|
||||
meta = {"purpose": "test", "epoch": 1, "flags": ["x", "y"], "nullable": None}
|
||||
# Unique content per test so the seed always creates a fresh asset (201).
|
||||
# Delete is now always a soft delete, so content from a prior test survives
|
||||
|
||||
@@ -133,6 +133,66 @@ class TestListReferencesPage:
|
||||
assert total == 1
|
||||
assert refs[0].name == "tagged"
|
||||
|
||||
def test_include_tags_filter_ands_persisted_model_tags(self, session: Session):
|
||||
asset = _make_asset(session, "hash-model-tags")
|
||||
checkpoint = _make_reference(session, asset, name="checkpoint")
|
||||
lora = _make_reference(session, asset, name="lora")
|
||||
input_ref = _make_reference(session, asset, name="input")
|
||||
ensure_tags_exist(
|
||||
session,
|
||||
["models", "model_type:checkpoints", "model_type:loras", "unit-tests"],
|
||||
)
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=checkpoint.id,
|
||||
tags=["models", "model_type:checkpoints", "unit-tests"],
|
||||
origin="automatic",
|
||||
)
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=lora.id,
|
||||
tags=["models", "model_type:loras", "unit-tests"],
|
||||
origin="automatic",
|
||||
)
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=input_ref.id,
|
||||
tags=["unit-tests"],
|
||||
)
|
||||
session.commit()
|
||||
|
||||
refs, _, total = list_references_page(
|
||||
session,
|
||||
include_tags=["models", "model_type:checkpoints", "unit-tests"],
|
||||
)
|
||||
|
||||
assert total == 1
|
||||
assert refs[0].id == checkpoint.id
|
||||
|
||||
def test_include_tags_filter_preserves_model_type_case(self, session: Session):
|
||||
asset = _make_asset(session, "hash-model-case")
|
||||
ref = _make_reference(session, asset, name="llm")
|
||||
ensure_tags_exist(session, ["models", "model_type:LLM"])
|
||||
add_tags_to_reference(
|
||||
session,
|
||||
reference_id=ref.id,
|
||||
tags=["models", "model_type:LLM"],
|
||||
origin="automatic",
|
||||
)
|
||||
session.commit()
|
||||
|
||||
refs, _, total = list_references_page(
|
||||
session, include_tags=["models", "model_type:LLM"]
|
||||
)
|
||||
refs_lower, _, total_lower = list_references_page(
|
||||
session, include_tags=["models", "model_type:llm"]
|
||||
)
|
||||
|
||||
assert total == 1
|
||||
assert refs[0].id == ref.id
|
||||
assert total_lower == 0
|
||||
assert refs_lower == []
|
||||
|
||||
def test_exclude_tags_filter(self, session: Session):
|
||||
asset = _make_asset(session, "hash1")
|
||||
_make_reference(session, asset, name="keep")
|
||||
|
||||
@@ -176,6 +176,39 @@ class TestUpsertReference:
|
||||
ref = session.query(AssetReference).filter_by(file_path=file_path).one()
|
||||
assert ref.mtime_ns == final_mtime
|
||||
|
||||
def test_upsert_refreshes_loader_path_on_existing_reference(self, session: Session):
|
||||
"""Re-ingesting an existing reference writes the loader_path computed
|
||||
by that ingest, healing NULL or stale values even when nothing else
|
||||
about the row changed."""
|
||||
asset = _make_asset(session, "hash1")
|
||||
file_path = "/models/checkpoints/sub/model.safetensors"
|
||||
|
||||
upsert_reference(
|
||||
session, asset_id=asset.id, file_path=file_path, name="model",
|
||||
mtime_ns=100, loader_path=None,
|
||||
)
|
||||
session.commit()
|
||||
|
||||
created, updated = upsert_reference(
|
||||
session, asset_id=asset.id, file_path=file_path, name="model",
|
||||
mtime_ns=100, loader_path="sub/model.safetensors",
|
||||
)
|
||||
session.commit()
|
||||
|
||||
assert created is False
|
||||
assert updated is True
|
||||
ref = session.query(AssetReference).filter_by(file_path=file_path).one()
|
||||
assert ref.loader_path == "sub/model.safetensors"
|
||||
|
||||
# Identical loader_path is a no-op, not a spurious update.
|
||||
created, updated = upsert_reference(
|
||||
session, asset_id=asset.id, file_path=file_path, name="model",
|
||||
mtime_ns=100, loader_path="sub/model.safetensors",
|
||||
)
|
||||
session.commit()
|
||||
assert created is False
|
||||
assert updated is False
|
||||
|
||||
def test_upsert_restores_missing_reference(self, session: Session):
|
||||
"""Upserting a reference that was marked missing should restore it."""
|
||||
asset = _make_asset(session, "hash1")
|
||||
|
||||
@@ -58,7 +58,7 @@ class TestEnsureTagsExist:
|
||||
session.commit()
|
||||
|
||||
tags = session.query(Tag).all()
|
||||
assert {t.name for t in tags} == {"alpha", "beta"}
|
||||
assert {t.name for t in tags} == {"ALPHA", "Beta", "alpha"}
|
||||
|
||||
def test_empty_list_is_noop(self, session: Session):
|
||||
ensure_tags_exist(session, [])
|
||||
@@ -258,6 +258,16 @@ class TestListTagsWithUsage:
|
||||
tag_names = {name for name, _ in rows}
|
||||
assert tag_names == {"alpha", "alphabet"}
|
||||
|
||||
def test_prefix_filter_is_case_sensitive(self, session: Session):
|
||||
ensure_tags_exist(session, ["model_type:LLM", "model_type:llm"])
|
||||
session.commit()
|
||||
|
||||
rows, total = list_tags_with_usage(session, prefix="model_type:L")
|
||||
|
||||
tag_names = {name for name, _ in rows}
|
||||
assert tag_names == {"model_type:LLM"}
|
||||
assert total == 1
|
||||
|
||||
def test_order_by_name(self, session: Session):
|
||||
ensure_tags_exist(session, ["zebra", "alpha", "middle"])
|
||||
session.commit()
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
"""Tests for how _build_asset_response derives the response `loader_path`.
|
||||
|
||||
Guards the persist-and-read contract: the response reads the stored
|
||||
`loader_path` verbatim, with no read-time recomputation. Like tags, the
|
||||
value is a seed-time derivative healed by the scan lifecycle.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from app.assets.api.routes import _build_asset_response
|
||||
from app.assets.services.schemas import AssetDetailResult, ReferenceData
|
||||
|
||||
_TS = datetime(2024, 1, 1, 0, 0, 0)
|
||||
|
||||
|
||||
def _make_result(
|
||||
*, file_path: str | None, loader_path: str | None
|
||||
) -> AssetDetailResult:
|
||||
ref = ReferenceData(
|
||||
id="ref-1",
|
||||
name="model.safetensors",
|
||||
file_path=file_path,
|
||||
loader_path=loader_path,
|
||||
user_metadata=None,
|
||||
preview_id=None,
|
||||
created_at=_TS,
|
||||
updated_at=_TS,
|
||||
last_access_time=_TS,
|
||||
)
|
||||
return AssetDetailResult(ref=ref, asset=None, tags=[])
|
||||
|
||||
|
||||
def test_uses_persisted_loader_path_without_recomputing():
|
||||
"""A stored loader_path is returned verbatim, not re-derived from file_path.
|
||||
|
||||
The sentinel value could never be produced by compute_loader_path for this
|
||||
file_path, so seeing it in the response proves the stored column is read.
|
||||
"""
|
||||
result = _make_result(
|
||||
file_path="/unmatched/root/model.safetensors",
|
||||
loader_path="SENTINEL/stored.safetensors",
|
||||
)
|
||||
|
||||
resp = _build_asset_response(result)
|
||||
|
||||
assert resp.loader_path == "SENTINEL/stored.safetensors"
|
||||
|
||||
|
||||
def test_null_stored_loader_path_is_served_as_null(tmp_path: Path):
|
||||
"""No read-time recomputation: a NULL column is served as null even when
|
||||
the path would resolve."""
|
||||
models = tmp_path / "models"
|
||||
ckpt = models / "checkpoints"
|
||||
ckpt.mkdir(parents=True)
|
||||
f = ckpt / "bar.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch("app.assets.services.path_utils.folder_paths") as mock_fp, patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(ckpt)], {".safetensors"})],
|
||||
):
|
||||
mock_fp.get_input_directory.return_value = str(tmp_path / "in")
|
||||
mock_fp.get_output_directory.return_value = str(tmp_path / "out")
|
||||
mock_fp.get_temp_directory.return_value = str(tmp_path / "tmp")
|
||||
mock_fp.models_dir = str(models)
|
||||
|
||||
result = _make_result(file_path=str(f), loader_path=None)
|
||||
resp = _build_asset_response(result)
|
||||
|
||||
assert resp.loader_path is None
|
||||
assert resp.display_name == "checkpoints/bar.safetensors"
|
||||
|
||||
|
||||
def test_all_path_fields_null_without_file_path():
|
||||
"""API-created / hash-only references (no file_path) expose no paths."""
|
||||
result = _make_result(file_path=None, loader_path=None)
|
||||
|
||||
resp = _build_asset_response(result)
|
||||
|
||||
assert resp.loader_path is None
|
||||
assert resp.display_name is None
|
||||
@@ -1,10 +1,14 @@
|
||||
"""Tests for bulk ingest services."""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.assets.database.models import Asset, AssetReference
|
||||
from app.assets.database.queries import get_reference_tags
|
||||
from app.assets.scanner import build_asset_specs
|
||||
from app.assets.services.bulk_ingest import SeedAssetSpec, batch_insert_seed_assets
|
||||
|
||||
|
||||
@@ -101,6 +105,184 @@ class TestBatchInsertSeedAssets:
|
||||
asset = session.query(Asset).filter_by(id=ref.asset_id).first()
|
||||
assert asset.mime_type == expected_mime, f"Expected {expected_mime} for {filename}, got {asset.mime_type}"
|
||||
|
||||
def test_duplicate_paths_merge_tags_before_insert(
|
||||
self, session: Session, temp_dir: Path
|
||||
):
|
||||
"""Overlapping model-folder registrations can emit the same path twice."""
|
||||
file_path = temp_dir / "shared.safetensors"
|
||||
file_path.write_bytes(b"shared model")
|
||||
|
||||
specs: list[SeedAssetSpec] = [
|
||||
{
|
||||
"abs_path": str(file_path),
|
||||
"size_bytes": 12,
|
||||
"mtime_ns": 1234567890000000000,
|
||||
"info_name": "Shared Model",
|
||||
"tags": ["models", "model_type:checkpoints"],
|
||||
"fname": "shared.safetensors",
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": "application/safetensors",
|
||||
},
|
||||
{
|
||||
"abs_path": str(file_path),
|
||||
"size_bytes": 12,
|
||||
"mtime_ns": 1234567890000000000,
|
||||
"info_name": "Shared Model",
|
||||
"tags": ["models", "model_type:diffusion_models"],
|
||||
"fname": "shared.safetensors",
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": "application/safetensors",
|
||||
},
|
||||
]
|
||||
|
||||
result = batch_insert_seed_assets(session, specs=specs, owner_id="")
|
||||
|
||||
assert result.inserted_refs == 1
|
||||
assert result.won_paths == 1
|
||||
refs = session.query(AssetReference).all()
|
||||
assert len(refs) == 1
|
||||
assert set(get_reference_tags(session, reference_id=refs[0].id)) == {
|
||||
"models",
|
||||
"model_type:checkpoints",
|
||||
"model_type:diffusion_models",
|
||||
}
|
||||
|
||||
def test_duplicate_paths_are_merged_after_abspath_normalization(
|
||||
self, session: Session, temp_dir: Path, monkeypatch
|
||||
):
|
||||
"""The scanner may emit equivalent paths with different spelling."""
|
||||
file_path = temp_dir / "same-file.safetensors"
|
||||
file_path.write_bytes(b"shared model")
|
||||
monkeypatch.chdir(temp_dir)
|
||||
relative_path = file_path.name
|
||||
absolute_path = os.path.abspath(relative_path)
|
||||
|
||||
specs: list[SeedAssetSpec] = [
|
||||
{
|
||||
"abs_path": relative_path,
|
||||
"size_bytes": 12,
|
||||
"mtime_ns": 1234567890000000000,
|
||||
"info_name": "Shared Model",
|
||||
"tags": ["models", "model_type:checkpoints"],
|
||||
"fname": "same-file.safetensors",
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": "application/safetensors",
|
||||
},
|
||||
{
|
||||
"abs_path": absolute_path,
|
||||
"size_bytes": 12,
|
||||
"mtime_ns": 1234567890000000000,
|
||||
"info_name": "Shared Model",
|
||||
"tags": ["models", "model_type:diffusion_models"],
|
||||
"fname": "same-file.safetensors",
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": "application/safetensors",
|
||||
},
|
||||
]
|
||||
|
||||
result = batch_insert_seed_assets(session, specs=specs, owner_id="")
|
||||
|
||||
assert result.inserted_refs == 1
|
||||
assert result.won_paths == 1
|
||||
refs = session.query(AssetReference).all()
|
||||
assert len(refs) == 1
|
||||
assert refs[0].file_path == absolute_path
|
||||
# loader_path is persisted from the spec's fname (compute_loader_path).
|
||||
assert refs[0].loader_path == "same-file.safetensors"
|
||||
assert set(get_reference_tags(session, reference_id=refs[0].id)) == {
|
||||
"models",
|
||||
"model_type:checkpoints",
|
||||
"model_type:diffusion_models",
|
||||
}
|
||||
|
||||
def test_scanner_duplicate_shared_model_paths_keep_all_model_type_tags(
|
||||
self, session: Session, temp_dir: Path
|
||||
):
|
||||
"""Shared extra model roots make scanner collection emit duplicate paths."""
|
||||
shared_root = temp_dir / "shared"
|
||||
input_dir = temp_dir / "input"
|
||||
output_dir = temp_dir / "output"
|
||||
temp_root = temp_dir / "temp"
|
||||
for directory in (shared_root, input_dir, output_dir, temp_root):
|
||||
directory.mkdir()
|
||||
file_path = shared_root / "dual_use_model.safetensors"
|
||||
file_path.write_bytes(b"shared model")
|
||||
|
||||
with (
|
||||
patch("app.assets.services.path_utils.folder_paths") as mock_fp,
|
||||
patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
("checkpoints", [str(shared_root)], {".safetensors"}),
|
||||
("diffusion_models", [str(shared_root)], {".safetensors"}),
|
||||
],
|
||||
),
|
||||
):
|
||||
mock_fp.get_input_directory.return_value = str(input_dir)
|
||||
mock_fp.get_output_directory.return_value = str(output_dir)
|
||||
mock_fp.get_temp_directory.return_value = str(temp_root)
|
||||
|
||||
specs, tag_pool, skipped = build_asset_specs(
|
||||
paths=[str(file_path), str(file_path)],
|
||||
existing_paths=set(),
|
||||
enable_metadata_extraction=False,
|
||||
compute_hashes=False,
|
||||
)
|
||||
|
||||
assert skipped == 0
|
||||
assert len(specs) == 2
|
||||
assert tag_pool == {
|
||||
"models",
|
||||
"model_type:checkpoints",
|
||||
"model_type:diffusion_models",
|
||||
}
|
||||
|
||||
result = batch_insert_seed_assets(session, specs=specs, owner_id="")
|
||||
|
||||
assert result.inserted_refs == 1
|
||||
assert result.won_paths == 1
|
||||
refs = session.query(AssetReference).all()
|
||||
assert len(refs) == 1
|
||||
assert set(get_reference_tags(session, reference_id=refs[0].id)) == {
|
||||
"models",
|
||||
"model_type:checkpoints",
|
||||
"model_type:diffusion_models",
|
||||
}
|
||||
|
||||
def test_loader_path_persisted_as_null_when_fname_is_none(
|
||||
self, session: Session, temp_dir: Path
|
||||
):
|
||||
"""A file with no in-root loader path (fname=None, e.g. an orphan under
|
||||
models_root) persists loader_path as NULL rather than a synthesized value."""
|
||||
file_path = temp_dir / "orphan.bin"
|
||||
file_path.write_bytes(b"x")
|
||||
|
||||
specs: list[SeedAssetSpec] = [
|
||||
{
|
||||
"abs_path": str(file_path),
|
||||
"size_bytes": 1,
|
||||
"mtime_ns": 1234567890000000000,
|
||||
"info_name": "orphan.bin",
|
||||
"tags": [],
|
||||
"fname": None,
|
||||
"metadata": None,
|
||||
"hash": None,
|
||||
"mime_type": None,
|
||||
}
|
||||
]
|
||||
|
||||
result = batch_insert_seed_assets(session, specs=specs, owner_id="")
|
||||
|
||||
assert result.inserted_refs == 1
|
||||
refs = session.query(AssetReference).all()
|
||||
assert len(refs) == 1
|
||||
assert refs[0].file_path == str(file_path)
|
||||
assert refs[0].loader_path is None
|
||||
|
||||
|
||||
class TestMetadataExtraction:
|
||||
def test_extracts_mime_type_for_model_files(self, temp_dir: Path):
|
||||
|
||||
@@ -94,6 +94,47 @@ class TestIngestFileFromPath:
|
||||
ref_tags = get_reference_tags(session, reference_id=result.reference_id)
|
||||
assert set(ref_tags) == {"models", "checkpoints"}
|
||||
|
||||
def test_path_derived_tags_use_automatic_origin(
|
||||
self, mock_create_session, temp_dir: Path, session: Session
|
||||
):
|
||||
input_dir = temp_dir / "input"
|
||||
output_dir = temp_dir / "output"
|
||||
temp_root = temp_dir / "temp"
|
||||
for directory in (input_dir, output_dir, temp_root):
|
||||
directory.mkdir()
|
||||
file_path = input_dir / "pasted" / "tagged.png"
|
||||
file_path.parent.mkdir()
|
||||
file_path.write_bytes(b"data")
|
||||
|
||||
with (
|
||||
patch("app.assets.services.path_utils.folder_paths") as mock_fp,
|
||||
patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[],
|
||||
),
|
||||
):
|
||||
mock_fp.get_input_directory.return_value = str(input_dir)
|
||||
mock_fp.get_output_directory.return_value = str(output_dir)
|
||||
mock_fp.get_temp_directory.return_value = str(temp_root)
|
||||
|
||||
result = _ingest_file_from_path(
|
||||
abs_path=str(file_path),
|
||||
asset_hash="blake3:pathorigin",
|
||||
size_bytes=4,
|
||||
mtime_ns=1234567890000000000,
|
||||
info_name="Tagged Asset",
|
||||
tags=["input", "manual-label"],
|
||||
)
|
||||
|
||||
assert result.reference_id is not None
|
||||
links = session.query(AssetReferenceTag).filter_by(
|
||||
asset_reference_id=result.reference_id
|
||||
)
|
||||
origin_by_tag = {link.tag_name: link.origin for link in links}
|
||||
assert origin_by_tag["input"] == "automatic"
|
||||
assert origin_by_tag["pasted"] == "automatic"
|
||||
assert origin_by_tag["manual-label"] == "manual"
|
||||
|
||||
def test_idempotent_upsert(self, mock_create_session, temp_dir: Path, session: Session):
|
||||
file_path = temp_dir / "dup.bin"
|
||||
file_path.write_bytes(b"content")
|
||||
@@ -288,6 +329,45 @@ class TestIngestExistingFileTagFK:
|
||||
assert "output" in ref_tag_names
|
||||
|
||||
|
||||
class TestIngestExistingFileLoaderPath:
|
||||
"""Outputs saved into a subfolder must persist the subfolder-qualified
|
||||
loader path, not the bare basename (regression: spec["fname"] was
|
||||
os.path.basename)."""
|
||||
|
||||
def test_subfoldered_output_persists_relative_loader_path(
|
||||
self, mock_create_session, temp_dir: Path, session: Session
|
||||
):
|
||||
input_dir = temp_dir / "input"
|
||||
output_dir = temp_dir / "output"
|
||||
temp_root = temp_dir / "temp"
|
||||
for directory in (input_dir, output_dir, temp_root):
|
||||
directory.mkdir()
|
||||
file_path = output_dir / "sub" / "img_00001_.png"
|
||||
file_path.parent.mkdir()
|
||||
file_path.write_bytes(b"image data")
|
||||
|
||||
with (
|
||||
patch("app.assets.services.path_utils.folder_paths") as mock_fp,
|
||||
patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[],
|
||||
),
|
||||
):
|
||||
mock_fp.get_input_directory.return_value = str(input_dir)
|
||||
mock_fp.get_output_directory.return_value = str(output_dir)
|
||||
mock_fp.get_temp_directory.return_value = str(temp_root)
|
||||
|
||||
assert ingest_existing_file(abs_path=str(file_path)) is True
|
||||
|
||||
ref = (
|
||||
session.query(AssetReference)
|
||||
.filter_by(file_path=str(file_path))
|
||||
.one()
|
||||
)
|
||||
assert ref.loader_path == "sub/img_00001_.png"
|
||||
assert (ref.user_metadata or {}).get("filename") == "sub/img_00001_.png"
|
||||
|
||||
|
||||
class TestIngestImageDimensions:
|
||||
"""system_metadata should carry {kind, width, height} for image assets."""
|
||||
|
||||
|
||||
@@ -6,7 +6,16 @@ from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from app.assets.services.path_utils import get_asset_category_and_relative_path
|
||||
from app.assets.services.path_utils import (
|
||||
compute_display_name,
|
||||
compute_loader_path,
|
||||
compute_logical_path,
|
||||
get_asset_category_and_relative_path,
|
||||
get_known_input_subfolder_tags_from_path,
|
||||
get_known_subfolder_tags,
|
||||
get_name_and_tags_from_asset_path,
|
||||
resolve_destination_from_tags,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -17,7 +26,8 @@ def fake_dirs():
|
||||
input_dir = root_path / "input"
|
||||
output_dir = root_path / "output"
|
||||
temp_dir = root_path / "temp"
|
||||
models_dir = root_path / "models" / "checkpoints"
|
||||
models_root = root_path / "models"
|
||||
models_dir = models_root / "checkpoints"
|
||||
for d in (input_dir, output_dir, temp_dir, models_dir):
|
||||
d.mkdir(parents=True)
|
||||
|
||||
@@ -25,15 +35,17 @@ def fake_dirs():
|
||||
mock_fp.get_input_directory.return_value = str(input_dir)
|
||||
mock_fp.get_output_directory.return_value = str(output_dir)
|
||||
mock_fp.get_temp_directory.return_value = str(temp_dir)
|
||||
mock_fp.models_dir = str(models_root)
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(models_dir)])],
|
||||
return_value=[("checkpoints", [str(models_dir)], {".safetensors"})],
|
||||
):
|
||||
yield {
|
||||
"input": input_dir,
|
||||
"output": output_dir,
|
||||
"temp": temp_dir,
|
||||
"models_root": models_root,
|
||||
"models": models_dir,
|
||||
}
|
||||
|
||||
@@ -76,6 +88,538 @@ class TestGetAssetCategoryAndRelativePath:
|
||||
cat, rel = get_asset_category_and_relative_path(str(f))
|
||||
assert cat == "models"
|
||||
|
||||
def test_model_path_tags_include_registered_model_type_only(self, fake_dirs):
|
||||
f = fake_dirs["models"] / "subdir" / "model.safetensors"
|
||||
f.parent.mkdir()
|
||||
f.touch()
|
||||
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
assert "checkpoints" not in tags
|
||||
assert "subdir" not in tags
|
||||
|
||||
def test_model_type_preserves_registered_folder_case(self, fake_dirs):
|
||||
llm_dir = fake_dirs["models"].parent / "LLM"
|
||||
llm_dir.mkdir()
|
||||
f = llm_dir / "model.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("LLM", [str(llm_dir)], {".safetensors"})],
|
||||
):
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:LLM" in tags
|
||||
assert "model_type:llm" not in tags
|
||||
|
||||
def test_path_components_do_not_create_model_type_tags(self, fake_dirs):
|
||||
f = fake_dirs["models"] / "loras" / "model.safetensors"
|
||||
f.parent.mkdir()
|
||||
f.touch()
|
||||
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
assert "loras" not in tags
|
||||
assert "model_type:loras" not in tags
|
||||
|
||||
def test_shared_root_returns_all_matching_model_type_tags(self, fake_dirs):
|
||||
shared_root = fake_dirs["models"].parent / "shared"
|
||||
shared_root.mkdir()
|
||||
f = shared_root / "foo.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
("checkpoints", [str(shared_root)], {".safetensors"}),
|
||||
("loras", [str(shared_root)], {".safetensors"}),
|
||||
],
|
||||
):
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
assert "model_type:loras" in tags
|
||||
|
||||
def test_shared_root_model_type_tags_respect_bucket_extensions(self, fake_dirs):
|
||||
"""Buckets sharing a base dir only tag files matching their extensions."""
|
||||
shared_root = fake_dirs["models"].parent / "unet"
|
||||
shared_root.mkdir()
|
||||
safetensors_file = shared_root / "wan.safetensors"
|
||||
gguf_file = shared_root / "wan.gguf"
|
||||
safetensors_file.touch()
|
||||
gguf_file.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
("diffusion_models", [str(shared_root)], {".safetensors"}),
|
||||
("unet_gguf", [str(shared_root)], {".gguf"}),
|
||||
],
|
||||
):
|
||||
_name, safetensors_tags = get_name_and_tags_from_asset_path(str(safetensors_file))
|
||||
_name, gguf_tags = get_name_and_tags_from_asset_path(str(gguf_file))
|
||||
|
||||
assert "model_type:diffusion_models" in safetensors_tags
|
||||
assert "model_type:unet_gguf" not in safetensors_tags
|
||||
assert "model_type:unet_gguf" in gguf_tags
|
||||
assert "model_type:diffusion_models" not in gguf_tags
|
||||
|
||||
def test_empty_extension_set_tags_any_extension(self, fake_dirs):
|
||||
"""Custom buckets registered without extensions accept every file."""
|
||||
custom_root = fake_dirs["models"].parent / "custom_bucket"
|
||||
custom_root.mkdir()
|
||||
f = custom_root / "weights.bin"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("custom_bucket", [str(custom_root)], set())],
|
||||
):
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:custom_bucket" in tags
|
||||
|
||||
def test_no_extension_match_keeps_models_tag_without_model_type(self, fake_dirs):
|
||||
f = fake_dirs["models"] / "notes.txt"
|
||||
f.touch()
|
||||
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert not any(tag.startswith("model_type:") for tag in tags)
|
||||
|
||||
def test_output_backed_registered_folder_gets_model_and_output_tags(self, fake_dirs):
|
||||
output_checkpoints_dir = fake_dirs["output"] / "checkpoints"
|
||||
output_checkpoints_dir.mkdir()
|
||||
f = output_checkpoints_dir / "saved.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(output_checkpoints_dir)], {".safetensors"})],
|
||||
):
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
assert "output" in tags
|
||||
|
||||
def test_temp_path_tags_include_temp_not_output_or_preview(self, fake_dirs):
|
||||
f = fake_dirs["temp"] / "preview.png"
|
||||
f.touch()
|
||||
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
|
||||
assert "temp" in tags
|
||||
assert "output" not in tags
|
||||
assert "preview:true" not in tags
|
||||
|
||||
def test_known_subfolder_tags_are_centralized(self):
|
||||
assert get_known_subfolder_tags("pasted") == ["pasted"]
|
||||
assert get_known_subfolder_tags("arbitrary") == []
|
||||
|
||||
def test_known_input_subfolder_tags_are_path_derived_for_direct_children(self, fake_dirs):
|
||||
f = fake_dirs["input"] / "pasted" / "image.png"
|
||||
f.parent.mkdir()
|
||||
f.touch()
|
||||
|
||||
assert get_known_input_subfolder_tags_from_path(str(f)) == ["pasted"]
|
||||
|
||||
_name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert "input" in tags
|
||||
assert "pasted" in tags
|
||||
|
||||
def test_known_input_subfolder_tags_do_not_apply_to_nested_or_other_roots(self, fake_dirs):
|
||||
nested = fake_dirs["input"] / "pasted" / "session" / "image.png"
|
||||
output = fake_dirs["output"] / "pasted" / "image.png"
|
||||
for path in (nested, output):
|
||||
path.parent.mkdir(parents=True)
|
||||
path.touch()
|
||||
|
||||
assert get_known_input_subfolder_tags_from_path(str(nested)) == []
|
||||
assert get_known_input_subfolder_tags_from_path(str(output)) == []
|
||||
|
||||
def test_unknown_path_raises(self, fake_dirs):
|
||||
with pytest.raises(ValueError, match="not within"):
|
||||
get_asset_category_and_relative_path("/some/random/path.png")
|
||||
|
||||
|
||||
class TestResponseStoragePaths:
|
||||
def test_input_file_path_and_display_name_include_subfolder(self, fake_dirs):
|
||||
sub = fake_dirs["input"] / "some" / "folder"
|
||||
sub.mkdir(parents=True)
|
||||
f = sub / "image.png"
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "input/some/folder/image.png"
|
||||
assert compute_display_name(str(f)) == "some/folder/image.png"
|
||||
|
||||
def test_output_file_path_and_display_name_include_subfolder(self, fake_dirs):
|
||||
sub = fake_dirs["output"] / "renders"
|
||||
sub.mkdir()
|
||||
f = sub / "ComfyUI_00001_.png"
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "output/renders/ComfyUI_00001_.png"
|
||||
assert compute_display_name(str(f)) == "renders/ComfyUI_00001_.png"
|
||||
|
||||
def test_temp_file_path_and_display_name(self, fake_dirs):
|
||||
f = fake_dirs["temp"] / "preview.png"
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "temp/preview.png"
|
||||
assert compute_display_name(str(f)) == "preview.png"
|
||||
|
||||
def test_exact_storage_root_has_no_display_name(self, fake_dirs):
|
||||
assert compute_logical_path(str(fake_dirs["input"])) == "input"
|
||||
assert compute_display_name(str(fake_dirs["input"])) is None
|
||||
|
||||
def test_longest_matching_builtin_root_wins(self, fake_dirs, tmp_path: Path):
|
||||
nested_output = fake_dirs["input"] / "nested-output"
|
||||
nested_output.mkdir()
|
||||
f = nested_output / "image.png"
|
||||
f.touch()
|
||||
|
||||
with patch("app.assets.services.path_utils.folder_paths") as mock_fp:
|
||||
mock_fp.get_input_directory.return_value = str(fake_dirs["input"])
|
||||
mock_fp.get_output_directory.return_value = str(nested_output)
|
||||
mock_fp.get_temp_directory.return_value = str(tmp_path / "temp")
|
||||
mock_fp.models_dir = str(fake_dirs["models_root"])
|
||||
|
||||
assert compute_logical_path(str(f)) == "output/image.png"
|
||||
assert compute_display_name(str(f)) == "image.png"
|
||||
|
||||
def test_model_file_path_is_relative_to_physical_models_root(self, fake_dirs):
|
||||
sub = fake_dirs["models"] / "flux"
|
||||
sub.mkdir()
|
||||
f = sub / "model.safetensors"
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "models/checkpoints/flux/model.safetensors"
|
||||
assert compute_display_name(str(f)) == "checkpoints/flux/model.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "model.safetensors"
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
assert "checkpoints" not in tags
|
||||
assert "flux" not in tags
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"folder_name",
|
||||
["checkpoints", "clip", "vae", "diffusion_models", "loras"],
|
||||
)
|
||||
def test_output_model_folder_uses_output_storage_file_path(self, fake_dirs, folder_name):
|
||||
output_model_dir = fake_dirs["output"] / folder_name
|
||||
output_model_dir.mkdir(exist_ok=True)
|
||||
default_model_dir = fake_dirs["models_root"] / folder_name
|
||||
default_model_dir.mkdir(exist_ok=True)
|
||||
f = output_model_dir / "saved.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
(folder_name, [str(default_model_dir), str(output_model_dir)], {".safetensors"})
|
||||
],
|
||||
):
|
||||
assert compute_logical_path(str(f)) == f"output/{folder_name}/saved.safetensors"
|
||||
assert compute_display_name(str(f)) == f"{folder_name}/saved.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "saved.safetensors"
|
||||
assert "output" in tags
|
||||
assert "models" in tags
|
||||
assert f"model_type:{folder_name}" in tags
|
||||
assert folder_name not in tags
|
||||
|
||||
def test_output_model_subfolder_uses_output_storage_file_path(self, fake_dirs):
|
||||
folder_name = "loras"
|
||||
output_model_dir = fake_dirs["output"] / folder_name
|
||||
subdir = output_model_dir / "experiments"
|
||||
subdir.mkdir(parents=True)
|
||||
f = subdir / "my_lora.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[(folder_name, [str(output_model_dir)], {".safetensors"})],
|
||||
):
|
||||
assert (
|
||||
compute_logical_path(str(f))
|
||||
== "output/loras/experiments/my_lora.safetensors"
|
||||
)
|
||||
assert compute_display_name(str(f)) == "loras/experiments/my_lora.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "my_lora.safetensors"
|
||||
assert "output" in tags
|
||||
assert "models" in tags
|
||||
assert "model_type:loras" in tags
|
||||
assert "loras" not in tags
|
||||
assert "experiments" not in tags
|
||||
|
||||
def test_external_model_folder_without_provenance_has_no_file_path(self, tmp_path: Path):
|
||||
external_checkpoints_dir = tmp_path / "external" / "not_named_like_category"
|
||||
external_checkpoints_dir.mkdir(parents=True)
|
||||
f = external_checkpoints_dir / "external.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(external_checkpoints_dir)], {".safetensors"})],
|
||||
):
|
||||
assert compute_logical_path(str(f)) is None
|
||||
assert compute_display_name(str(f)) is None
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "external.safetensors"
|
||||
assert "models" in tags
|
||||
assert "model_type:checkpoints" in tags
|
||||
|
||||
def test_same_relative_model_file_under_multiple_external_roots_has_no_storage_file_path(
|
||||
self, tmp_path: Path
|
||||
):
|
||||
foo_dir = tmp_path / "foo"
|
||||
bar_dir = tmp_path / "bar"
|
||||
foo_dir.mkdir()
|
||||
bar_dir.mkdir()
|
||||
foo_file = foo_dir / "baz.safetensors"
|
||||
bar_file = bar_dir / "baz.safetensors"
|
||||
foo_file.touch()
|
||||
bar_file.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(foo_dir), str(bar_dir)], {".safetensors"})],
|
||||
):
|
||||
assert compute_logical_path(str(foo_file)) is None
|
||||
assert compute_logical_path(str(bar_file)) is None
|
||||
assert compute_display_name(str(foo_file)) is None
|
||||
assert compute_display_name(str(bar_file)) is None
|
||||
|
||||
def test_output_clip_folder_uses_output_storage_and_text_encoder_tag(self, fake_dirs):
|
||||
output_clip_dir = fake_dirs["output"] / "clip"
|
||||
output_clip_dir.mkdir()
|
||||
f = output_clip_dir / "clip_l.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("text_encoders", [str(output_clip_dir)], {".safetensors"})],
|
||||
):
|
||||
assert compute_logical_path(str(f)) == "output/clip/clip_l.safetensors"
|
||||
assert compute_display_name(str(f)) == "clip/clip_l.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "clip_l.safetensors"
|
||||
assert "output" in tags
|
||||
assert "models" in tags
|
||||
assert "model_type:text_encoders" in tags
|
||||
assert "clip" not in tags
|
||||
|
||||
def test_physical_unet_folder_uses_storage_path_and_diffusion_models_tag(self, fake_dirs):
|
||||
unet_dir = fake_dirs["models_root"] / "unet"
|
||||
diffusion_models_dir = fake_dirs["models_root"] / "diffusion_models"
|
||||
unet_dir.mkdir()
|
||||
diffusion_models_dir.mkdir()
|
||||
f = unet_dir / "wan.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
("diffusion_models", [str(unet_dir), str(diffusion_models_dir)], {".safetensors"})
|
||||
],
|
||||
):
|
||||
assert compute_logical_path(str(f)) == "models/unet/wan.safetensors"
|
||||
assert compute_display_name(str(f)) == "unet/wan.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "wan.safetensors"
|
||||
assert "models" in tags
|
||||
assert "model_type:diffusion_models" in tags
|
||||
assert "unet" not in tags
|
||||
|
||||
def test_unregistered_file_under_physical_models_root_still_has_storage_file_path(self, fake_dirs):
|
||||
f = fake_dirs["models_root"] / "not_registered" / "orphan.bin"
|
||||
f.parent.mkdir()
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "models/not_registered/orphan.bin"
|
||||
assert compute_display_name(str(f)) == "not_registered/orphan.bin"
|
||||
|
||||
def test_output_checkpoint_folder_without_registration_has_only_output_tag(self, fake_dirs):
|
||||
f = fake_dirs["output"] / "checkpoints" / "saved.safetensors"
|
||||
f.parent.mkdir(exist_ok=True)
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[],
|
||||
):
|
||||
assert compute_logical_path(str(f)) == "output/checkpoints/saved.safetensors"
|
||||
assert compute_display_name(str(f)) == "checkpoints/saved.safetensors"
|
||||
|
||||
name, tags = get_name_and_tags_from_asset_path(str(f))
|
||||
assert name == "saved.safetensors"
|
||||
assert "output" in tags
|
||||
assert "models" not in tags
|
||||
assert not any(tag.startswith("model_type:") for tag in tags)
|
||||
|
||||
def test_unknown_path_returns_none(self):
|
||||
assert compute_logical_path("/some/random/path.png") is None
|
||||
assert compute_display_name("/some/random/path.png") is None
|
||||
|
||||
|
||||
class TestLoaderPath:
|
||||
"""In-root loader path: relative to the storage root, model category dropped."""
|
||||
|
||||
def test_model_loader_path_drops_category(self, fake_dirs):
|
||||
sub = fake_dirs["models"] / "flux"
|
||||
sub.mkdir()
|
||||
f = sub / "model.safetensors"
|
||||
f.touch()
|
||||
|
||||
# logical_path keeps the category, file_path (loader) drops it
|
||||
assert compute_logical_path(str(f)) == "models/checkpoints/flux/model.safetensors"
|
||||
assert compute_loader_path(str(f)) == "flux/model.safetensors"
|
||||
|
||||
def test_model_loader_path_flat_file(self, fake_dirs):
|
||||
f = fake_dirs["models"] / "model.safetensors"
|
||||
f.touch()
|
||||
|
||||
assert compute_loader_path(str(f)) == "model.safetensors"
|
||||
|
||||
def test_input_loader_path_keeps_subfolders(self, fake_dirs):
|
||||
sub = fake_dirs["input"] / "some" / "folder"
|
||||
sub.mkdir(parents=True)
|
||||
f = sub / "image.png"
|
||||
f.touch()
|
||||
|
||||
assert compute_loader_path(str(f)) == "some/folder/image.png"
|
||||
|
||||
def test_temp_loader_path(self, fake_dirs):
|
||||
f = fake_dirs["temp"] / "preview.png"
|
||||
f.touch()
|
||||
|
||||
assert compute_loader_path(str(f)) == "preview.png"
|
||||
|
||||
def test_unregistered_file_under_models_root_has_no_loader_path(self, fake_dirs):
|
||||
# Under models_root but not within any registered category base.
|
||||
f = fake_dirs["models_root"] / "not_registered" / "orphan.bin"
|
||||
f.parent.mkdir()
|
||||
f.touch()
|
||||
|
||||
# It still has a namespaced logical_path, but no loader path.
|
||||
assert compute_logical_path(str(f)) == "models/not_registered/orphan.bin"
|
||||
assert compute_loader_path(str(f)) is None
|
||||
|
||||
def test_extension_mismatch_in_registered_bucket_has_no_loader_path(self, fake_dirs):
|
||||
# Inside a registered bucket, but the bucket's extension set cannot
|
||||
# load it: no model_type tag, and no loader path either.
|
||||
f = fake_dirs["models"] / "notes.txt"
|
||||
f.touch()
|
||||
|
||||
assert compute_logical_path(str(f)) == "models/checkpoints/notes.txt"
|
||||
assert compute_loader_path(str(f)) is None
|
||||
|
||||
def test_shared_base_loader_path_uses_extension_matching_bucket(self, fake_dirs):
|
||||
shared_root = fake_dirs["models"].parent / "unet"
|
||||
shared_root.mkdir()
|
||||
f = shared_root / "wan.gguf"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[
|
||||
("diffusion_models", [str(shared_root)], {".safetensors"}),
|
||||
("unet_gguf", [str(shared_root)], {".gguf"}),
|
||||
],
|
||||
):
|
||||
assert compute_loader_path(str(f)) == "wan.gguf"
|
||||
|
||||
def test_match_all_bucket_provides_loader_path_for_any_extension(self, fake_dirs):
|
||||
custom_root = fake_dirs["models"].parent / "custom_bucket"
|
||||
custom_root.mkdir()
|
||||
f = custom_root / "weights.bin"
|
||||
f.touch()
|
||||
|
||||
with patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("custom_bucket", [str(custom_root)], set())],
|
||||
):
|
||||
assert compute_loader_path(str(f)) == "weights.bin"
|
||||
|
||||
def test_extra_path_model_has_loader_path_but_no_logical_path(self, tmp_path: Path):
|
||||
"""Registered category base outside models_dir (extra_model_paths style).
|
||||
|
||||
Loadable, so loader_path resolves; but it is not under any canonical
|
||||
storage root, so logical_path/display_name are None. This asymmetry is
|
||||
intentional: loader_path resolves every registered model-folder base,
|
||||
logical_path only resolves the canonical storage roots.
|
||||
"""
|
||||
extra = tmp_path / "extra_ckpts"
|
||||
extra.mkdir()
|
||||
f = extra / "foo.safetensors"
|
||||
f.touch()
|
||||
|
||||
with patch("app.assets.services.path_utils.folder_paths") as mock_fp, patch(
|
||||
"app.assets.services.path_utils.get_comfy_models_folders",
|
||||
return_value=[("checkpoints", [str(extra)], {".safetensors"})],
|
||||
):
|
||||
mock_fp.get_input_directory.return_value = str(tmp_path / "in")
|
||||
mock_fp.get_output_directory.return_value = str(tmp_path / "out")
|
||||
mock_fp.get_temp_directory.return_value = str(tmp_path / "tmp")
|
||||
mock_fp.models_dir = str(tmp_path / "models") # extra is NOT under this
|
||||
|
||||
assert compute_loader_path(str(f)) == "foo.safetensors"
|
||||
assert compute_logical_path(str(f)) is None
|
||||
assert compute_display_name(str(f)) is None
|
||||
|
||||
def test_unknown_path_returns_none(self):
|
||||
assert compute_loader_path("/some/random/path.png") is None
|
||||
|
||||
|
||||
class TestResolveDestinationFromTags:
|
||||
def test_extra_tags_are_not_path_components(self, fake_dirs):
|
||||
base_dir, subdirs = resolve_destination_from_tags(["input", "unit-tests", "foo"])
|
||||
|
||||
assert base_dir == os.path.abspath(fake_dirs["input"])
|
||||
assert subdirs == []
|
||||
|
||||
def test_model_upload_rejects_non_writable_registered_folders(self):
|
||||
with tempfile.TemporaryDirectory() as root:
|
||||
root_path = Path(root)
|
||||
checkpoints_dir = root_path / "models" / "checkpoints"
|
||||
configs_dir = root_path / "models" / "configs"
|
||||
custom_nodes_dir = root_path / "custom_nodes"
|
||||
for path in (checkpoints_dir, configs_dir, custom_nodes_dir):
|
||||
path.mkdir(parents=True)
|
||||
|
||||
with patch("app.assets.services.path_utils.folder_paths") as mock_fp:
|
||||
mock_fp.folder_names_and_paths = {
|
||||
"checkpoints": ([str(checkpoints_dir)], set()),
|
||||
"configs": ([str(configs_dir)], set()),
|
||||
"custom_nodes": ([str(custom_nodes_dir)], set()),
|
||||
}
|
||||
|
||||
base_dir, subdirs = resolve_destination_from_tags(
|
||||
["models", "model_type:checkpoints"]
|
||||
)
|
||||
assert base_dir == os.path.abspath(checkpoints_dir)
|
||||
assert subdirs == []
|
||||
|
||||
for folder_name in ("configs", "custom_nodes"):
|
||||
with pytest.raises(ValueError, match="unknown model category"):
|
||||
resolve_destination_from_tags(
|
||||
["models", f"model_type:{folder_name}"]
|
||||
)
|
||||
|
||||
@@ -19,7 +19,8 @@ def test_seed_asset_removed_when_file_is_deleted(
|
||||
"""Asset without hash (seed) whose file disappears:
|
||||
after triggering sync_seed_assets, Asset + AssetInfo disappear.
|
||||
"""
|
||||
# Create a file directly under input/unit-tests/<case> so tags include "unit-tests"
|
||||
# Create a file directly under input/unit-tests/<case>. Backend tags only
|
||||
# classify the root; nested path components are not exposed as tags.
|
||||
case_dir = comfy_tmp_base_dir / root / "unit-tests" / "syncseed"
|
||||
case_dir.mkdir(parents=True, exist_ok=True)
|
||||
name = f"seed_{uuid.uuid4().hex[:8]}.bin"
|
||||
@@ -32,7 +33,7 @@ def test_seed_asset_removed_when_file_is_deleted(
|
||||
# Verify it is visible via API and carries no hash (seed)
|
||||
r1 = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": "unit-tests,syncseed", "name_contains": name},
|
||||
params={"include_tags": root, "name_contains": name},
|
||||
timeout=120,
|
||||
)
|
||||
body1 = r1.json()
|
||||
@@ -54,7 +55,7 @@ def test_seed_asset_removed_when_file_is_deleted(
|
||||
# It should disappear (AssetInfo and seed Asset gone)
|
||||
r2 = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": "unit-tests,syncseed", "name_contains": name},
|
||||
params={"include_tags": root, "name_contains": name},
|
||||
timeout=120,
|
||||
)
|
||||
body2 = r2.json()
|
||||
@@ -132,7 +133,7 @@ def test_hashed_asset_two_asset_infos_both_get_missing(
|
||||
second_id = b2["id"]
|
||||
|
||||
# Remove the single underlying file
|
||||
p = comfy_tmp_base_dir / "input" / "unit-tests" / "multiinfo" / get_asset_filename(b2["asset_hash"], ".png")
|
||||
p = comfy_tmp_base_dir / "input" / get_asset_filename(created["asset_hash"], ".png")
|
||||
assert p.exists()
|
||||
p.unlink()
|
||||
|
||||
@@ -250,8 +251,7 @@ def test_missing_tag_clears_on_fastpass_when_mtime_and_size_match(
|
||||
|
||||
a = asset_factory(name, [root, "unit-tests", scope], {}, data)
|
||||
aid = a["id"]
|
||||
base = comfy_tmp_base_dir / root / "unit-tests" / scope
|
||||
p = base / get_asset_filename(a["asset_hash"], ".bin")
|
||||
p = comfy_tmp_base_dir / root / get_asset_filename(a["asset_hash"], ".bin")
|
||||
st0 = p.stat()
|
||||
orig_mtime_ns = getattr(st0, "st_mtime_ns", int(st0.st_mtime * 1_000_000_000))
|
||||
|
||||
|
||||
@@ -290,7 +290,7 @@ def test_metadata_filename_is_set_for_seed_asset_without_hash(
|
||||
|
||||
r1 = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": f"unit-tests,{scope}", "name_contains": name},
|
||||
params={"include_tags": root, "name_contains": name},
|
||||
timeout=120,
|
||||
)
|
||||
body = r1.json()
|
||||
|
||||
@@ -23,7 +23,7 @@ def test_download_svg_forced_to_attachment(http: requests.Session, api_base: str
|
||||
svg = b'<svg xmlns="http://www.w3.org/2000/svg"><script>alert(1)</script></svg>'
|
||||
files = {"file": ("evil.svg", svg, "image/svg+xml")}
|
||||
form_data = {
|
||||
"tags": json.dumps(["models", "checkpoints", "unit-tests", "svgxss"]),
|
||||
"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "svgxss"]),
|
||||
"name": "evil.svg",
|
||||
}
|
||||
up = http.post(api_base + "/api/assets", files=files, data=form_data, timeout=120)
|
||||
@@ -131,7 +131,7 @@ def test_download_chooses_existing_state_and_updates_access_time(
|
||||
assert t1 > t0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("seeded_asset", [{"tags": ["models", "checkpoints"]}], indirect=True)
|
||||
@pytest.mark.parametrize("seeded_asset", [{"tags": ["models", "model_type:checkpoints"]}], indirect=True)
|
||||
def test_download_missing_file_returns_404(
|
||||
http: requests.Session, api_base: str, comfy_tmp_base_dir: Path, seeded_asset: dict
|
||||
):
|
||||
|
||||
@@ -13,7 +13,7 @@ def _seed(asset_factory, make_asset_bytes, count: int, tag: str) -> list[str]:
|
||||
for n in names:
|
||||
asset_factory(
|
||||
n,
|
||||
["models", "checkpoints", "unit-tests", tag],
|
||||
["models", "model_type:checkpoints", "unit-tests", tag],
|
||||
{},
|
||||
make_asset_bytes(n, size=2048),
|
||||
)
|
||||
@@ -208,7 +208,7 @@ def test_cursor_walks_for_non_name_sorts(sort_field, http: requests.Session, api
|
||||
names = []
|
||||
for i in range(4):
|
||||
n = f"cursor_{sort_field}_{i:02d}.safetensors"
|
||||
asset_factory(n, ["models", "checkpoints", "unit-tests", f"cursor-{sort_field}"], {}, make_asset_bytes(n, size=2048 + i))
|
||||
asset_factory(n, ["models", "model_type:checkpoints", "unit-tests", f"cursor-{sort_field}"], {}, make_asset_bytes(n, size=2048 + i))
|
||||
names.append(n)
|
||||
|
||||
params = {
|
||||
|
||||
@@ -11,7 +11,7 @@ def test_list_assets_paging_and_sort(http: requests.Session, api_base: str, asse
|
||||
for n in names:
|
||||
asset_factory(
|
||||
n,
|
||||
["models", "checkpoints", "unit-tests", "paging"],
|
||||
["models", "model_type:checkpoints", "unit-tests", "paging"],
|
||||
{"epoch": 1},
|
||||
make_asset_bytes(n, size=2048),
|
||||
)
|
||||
@@ -45,8 +45,8 @@ def test_list_assets_paging_and_sort(http: requests.Session, api_base: str, asse
|
||||
|
||||
|
||||
def test_list_assets_include_exclude_and_name_contains(http: requests.Session, api_base: str, asset_factory):
|
||||
a = asset_factory("inc_a.safetensors", ["models", "checkpoints", "unit-tests", "alpha"], {}, b"X" * 1024)
|
||||
b = asset_factory("inc_b.safetensors", ["models", "checkpoints", "unit-tests", "beta"], {}, b"Y" * 1024)
|
||||
a = asset_factory("inc_a.safetensors", ["models", "model_type:checkpoints", "unit-tests", "alpha"], {}, b"X" * 1024)
|
||||
b = asset_factory("inc_b.safetensors", ["models", "model_type:checkpoints", "unit-tests", "beta"], {}, b"Y" * 1024)
|
||||
|
||||
r = http.get(
|
||||
api_base + "/api/assets",
|
||||
@@ -81,7 +81,7 @@ def test_list_assets_include_exclude_and_name_contains(http: requests.Session, a
|
||||
|
||||
|
||||
def test_list_assets_sort_by_size_both_orders(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-size"]
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-size"]
|
||||
n1, n2, n3 = "sz1.safetensors", "sz2.safetensors", "sz3.safetensors"
|
||||
asset_factory(n1, t, {}, make_asset_bytes(n1, 1024))
|
||||
asset_factory(n2, t, {}, make_asset_bytes(n2, 2048))
|
||||
@@ -108,7 +108,7 @@ def test_list_assets_sort_by_size_both_orders(http, api_base, asset_factory, mak
|
||||
|
||||
|
||||
def test_list_assets_sort_by_updated_at_desc(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-upd"]
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-upd"]
|
||||
a1 = asset_factory("upd_a.safetensors", t, {}, make_asset_bytes("upd_a", 1200))
|
||||
a2 = asset_factory("upd_b.safetensors", t, {}, make_asset_bytes("upd_b", 1200))
|
||||
|
||||
@@ -131,7 +131,7 @@ def test_list_assets_sort_by_updated_at_desc(http, api_base, asset_factory, make
|
||||
|
||||
|
||||
def test_list_assets_sort_by_last_access_time_desc(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-access"]
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-access"]
|
||||
asset_factory("acc_a.safetensors", t, {}, make_asset_bytes("acc_a", 1100))
|
||||
time.sleep(0.02)
|
||||
a2 = asset_factory("acc_b.safetensors", t, {}, make_asset_bytes("acc_b", 1100))
|
||||
@@ -154,14 +154,14 @@ def test_list_assets_sort_by_last_access_time_desc(http, api_base, asset_factory
|
||||
|
||||
|
||||
def test_list_assets_include_tags_variants_and_case(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-include"]
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-include"]
|
||||
a = asset_factory("incvar_alpha.safetensors", [*t, "alpha"], {}, make_asset_bytes("iva"))
|
||||
asset_factory("incvar_beta.safetensors", [*t, "beta"], {}, make_asset_bytes("ivb"))
|
||||
|
||||
# CSV + case-insensitive
|
||||
# CSV tag filters are whitespace-trimmed and case-sensitive.
|
||||
r1 = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": "UNIT-TESTS,LF-INCLUDE,alpha"},
|
||||
params={"include_tags": "unit-tests,lf-include,alpha"},
|
||||
timeout=120,
|
||||
)
|
||||
b1 = r1.json()
|
||||
@@ -196,14 +196,14 @@ def test_list_assets_include_tags_variants_and_case(http, api_base, asset_factor
|
||||
|
||||
|
||||
def test_list_assets_exclude_tags_dedup_and_case(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-exclude"]
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-exclude"]
|
||||
a = asset_factory("ex_a_alpha.safetensors", [*t, "alpha"], {}, make_asset_bytes("exa", 900))
|
||||
asset_factory("ex_b_beta.safetensors", [*t, "beta"], {}, make_asset_bytes("exb", 900))
|
||||
|
||||
# Exclude uppercase should work
|
||||
# Exclude filters are case-sensitive.
|
||||
r1 = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": "unit-tests,lf-exclude", "exclude_tags": "BETA"},
|
||||
params={"include_tags": "unit-tests,lf-exclude", "exclude_tags": "beta"},
|
||||
timeout=120,
|
||||
)
|
||||
b1 = r1.json()
|
||||
@@ -225,7 +225,7 @@ def test_list_assets_exclude_tags_dedup_and_case(http, api_base, asset_factory,
|
||||
|
||||
|
||||
def test_list_assets_name_contains_case_and_specials(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-name"]
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-name"]
|
||||
a1 = asset_factory("CaseMix.SAFE", t, {}, make_asset_bytes("cm", 800))
|
||||
a2 = asset_factory("case-other.safetensors", t, {}, make_asset_bytes("co", 800))
|
||||
|
||||
@@ -261,7 +261,7 @@ def test_list_assets_name_contains_case_and_specials(http, api_base, asset_facto
|
||||
|
||||
|
||||
def test_list_assets_offset_beyond_total_and_limit_boundary(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "checkpoints", "unit-tests", "lf-pagelimits"]
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "lf-pagelimits"]
|
||||
asset_factory("pl1.safetensors", t, {}, make_asset_bytes("pl1", 600))
|
||||
asset_factory("pl2.safetensors", t, {}, make_asset_bytes("pl2", 600))
|
||||
asset_factory("pl3.safetensors", t, {}, make_asset_bytes("pl3", 600))
|
||||
@@ -319,7 +319,7 @@ def test_list_assets_name_contains_literal_underscore(
|
||||
- foobar.safetensors (must NOT match)
|
||||
"""
|
||||
scope = f"lf-underscore-{uuid.uuid4().hex[:6]}"
|
||||
tags = ["models", "checkpoints", "unit-tests", scope]
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", scope]
|
||||
|
||||
a = asset_factory("foo_bar.safetensors", tags, {}, make_asset_bytes("a", 700))
|
||||
b = asset_factory("fooxbar.safetensors", tags, {}, make_asset_bytes("b", 700))
|
||||
|
||||
@@ -5,7 +5,7 @@ def test_meta_and_across_keys_and_types(
|
||||
http, api_base: str, asset_factory, make_asset_bytes
|
||||
):
|
||||
name = "mf_and_mix.safetensors"
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-and"]
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-and"]
|
||||
meta = {"purpose": "mix", "epoch": 1, "active": True, "score": 1.23}
|
||||
asset_factory(name, tags, meta, make_asset_bytes(name, 4096))
|
||||
|
||||
@@ -41,7 +41,7 @@ def test_meta_and_across_keys_and_types(
|
||||
|
||||
def test_meta_type_strictness_int_vs_str_and_bool(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_types.safetensors"
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-types"]
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-types"]
|
||||
meta = {"epoch": 1, "active": True}
|
||||
asset_factory(name, tags, meta, make_asset_bytes(name))
|
||||
|
||||
@@ -95,7 +95,7 @@ def test_meta_type_strictness_int_vs_str_and_bool(http, api_base, asset_factory,
|
||||
|
||||
def test_meta_any_of_list_of_scalars(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_list_scalars.safetensors"
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-list"]
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-list"]
|
||||
meta = {"flags": ["red", "green"]}
|
||||
asset_factory(name, tags, meta, make_asset_bytes(name, 3000))
|
||||
|
||||
@@ -134,7 +134,7 @@ def test_meta_none_semantics_missing_or_null_and_any_of_with_none(
|
||||
http, api_base, asset_factory, make_asset_bytes
|
||||
):
|
||||
# a1: key missing; a2: explicit null; a3: concrete value
|
||||
t = ["models", "checkpoints", "unit-tests", "mf-none"]
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "mf-none"]
|
||||
a1 = asset_factory("mf_none_missing.safetensors", t, {"x": 1}, make_asset_bytes("a1"))
|
||||
a2 = asset_factory("mf_none_null.safetensors", t, {"maybe": None}, make_asset_bytes("a2"))
|
||||
a3 = asset_factory("mf_none_value.safetensors", t, {"maybe": "x"}, make_asset_bytes("a3"))
|
||||
@@ -166,7 +166,7 @@ def test_meta_none_semantics_missing_or_null_and_any_of_with_none(
|
||||
|
||||
def test_meta_nested_json_object_equality(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_nested_json.safetensors"
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-nested"]
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-nested"]
|
||||
cfg = {"optimizer": "adam", "lr": 0.001, "schedule": {"type": "cosine", "warmup": 100}}
|
||||
asset_factory(name, tags, {"config": cfg}, make_asset_bytes(name, 2200))
|
||||
|
||||
@@ -197,7 +197,7 @@ def test_meta_nested_json_object_equality(http, api_base, asset_factory, make_as
|
||||
|
||||
def test_meta_list_of_objects_any_of(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_list_objects.safetensors"
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-objlist"]
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-objlist"]
|
||||
transforms = [{"type": "crop", "size": 128}, {"type": "flip", "p": 0.5}]
|
||||
asset_factory(name, tags, {"transforms": transforms}, make_asset_bytes(name, 2048))
|
||||
|
||||
@@ -228,7 +228,7 @@ def test_meta_list_of_objects_any_of(http, api_base, asset_factory, make_asset_b
|
||||
|
||||
def test_meta_with_special_and_unicode_keys(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_keys_unicode.safetensors"
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-keys"]
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-keys"]
|
||||
meta = {
|
||||
"weird.key": "v1",
|
||||
"path/like": 7,
|
||||
@@ -259,7 +259,7 @@ def test_meta_with_special_and_unicode_keys(http, api_base, asset_factory, make_
|
||||
|
||||
|
||||
def test_meta_with_zero_and_boolean_lists(http, api_base, asset_factory, make_asset_bytes):
|
||||
t = ["models", "checkpoints", "unit-tests", "mf-zero-bool"]
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "mf-zero-bool"]
|
||||
a0 = asset_factory("mf_zero_count.safetensors", t, {"count": 0}, make_asset_bytes("z", 1025))
|
||||
a1 = asset_factory("mf_bool_list.safetensors", t, {"choices": [True, False]}, make_asset_bytes("b", 1026))
|
||||
|
||||
@@ -286,7 +286,7 @@ def test_meta_with_zero_and_boolean_lists(http, api_base, asset_factory, make_as
|
||||
|
||||
def test_meta_mixed_list_types_and_strictness(http, api_base, asset_factory, make_asset_bytes):
|
||||
name = "mf_mixed_list.safetensors"
|
||||
tags = ["models", "checkpoints", "unit-tests", "mf-mixed"]
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "mf-mixed"]
|
||||
meta = {"mix": ["1", 1, True, None]}
|
||||
asset_factory(name, tags, meta, make_asset_bytes(name, 1999))
|
||||
|
||||
@@ -311,7 +311,7 @@ def test_meta_mixed_list_types_and_strictness(http, api_base, asset_factory, mak
|
||||
|
||||
def test_meta_unknown_key_and_none_behavior_with_scope_tags(http, api_base, asset_factory, make_asset_bytes):
|
||||
# Use a unique scope tag to avoid interference
|
||||
t = ["models", "checkpoints", "unit-tests", "mf-unknown-scope"]
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "mf-unknown-scope"]
|
||||
x = asset_factory("mf_unknown_a.safetensors", t, {"k1": 1}, make_asset_bytes("ua"))
|
||||
y = asset_factory("mf_unknown_b.safetensors", t, {"k2": 2}, make_asset_bytes("ub"))
|
||||
|
||||
@@ -340,13 +340,13 @@ def test_meta_with_tags_include_exclude_and_name_contains(http, api_base, asset_
|
||||
# alpha matches epoch=1; beta has epoch=2
|
||||
a = asset_factory(
|
||||
"mf_tag_alpha.safetensors",
|
||||
["models", "checkpoints", "unit-tests", "mf-tag", "alpha"],
|
||||
["models", "model_type:checkpoints", "unit-tests", "mf-tag", "alpha"],
|
||||
{"epoch": 1},
|
||||
make_asset_bytes("alpha"),
|
||||
)
|
||||
b = asset_factory(
|
||||
"mf_tag_beta.safetensors",
|
||||
["models", "checkpoints", "unit-tests", "mf-tag", "beta"],
|
||||
["models", "model_type:checkpoints", "unit-tests", "mf-tag", "beta"],
|
||||
{"epoch": 2},
|
||||
make_asset_bytes("beta"),
|
||||
)
|
||||
@@ -367,7 +367,7 @@ def test_meta_with_tags_include_exclude_and_name_contains(http, api_base, asset_
|
||||
|
||||
def test_meta_sort_and_paging_under_filter(http, api_base, asset_factory, make_asset_bytes):
|
||||
# Three assets in same scope with different sizes and a common filter key
|
||||
t = ["models", "checkpoints", "unit-tests", "mf-sort"]
|
||||
t = ["models", "model_type:checkpoints", "unit-tests", "mf-sort"]
|
||||
n1, n2, n3 = "mf_sort_1.safetensors", "mf_sort_2.safetensors", "mf_sort_3.safetensors"
|
||||
asset_factory(n1, t, {"group": "g"}, make_asset_bytes(n1, 1024))
|
||||
asset_factory(n2, t, {"group": "g"}, make_asset_bytes(n2, 2048))
|
||||
|
||||
@@ -29,7 +29,7 @@ def create_seed_file(comfy_tmp_base_dir: Path):
|
||||
def find_asset(http: requests.Session, api_base: str):
|
||||
"""Query API for assets matching scope and optional name."""
|
||||
def _find(scope: str, name: str | None = None) -> list[dict]:
|
||||
params = {"include_tags": f"unit-tests,{scope}"}
|
||||
params = {"limit": "500"}
|
||||
if name:
|
||||
params["name_contains"] = name
|
||||
r = http.get(f"{api_base}/api/assets", params=params, timeout=120)
|
||||
@@ -91,7 +91,7 @@ def test_hashed_asset_not_pruned_when_file_missing(
|
||||
data = make_asset_bytes("test", 2048)
|
||||
a = asset_factory("test.bin", ["input", "unit-tests", scope], {}, data)
|
||||
|
||||
path = comfy_tmp_base_dir / "input" / "unit-tests" / scope / get_asset_filename(a["asset_hash"], ".bin")
|
||||
path = comfy_tmp_base_dir / "input" / get_asset_filename(a["asset_hash"], ".bin")
|
||||
path.unlink()
|
||||
|
||||
trigger_sync_seed_assets(http, api_base)
|
||||
@@ -108,18 +108,20 @@ def test_prune_across_multiple_roots(
|
||||
):
|
||||
"""Prune correctly handles assets across input and output roots."""
|
||||
scope = f"multi-{uuid.uuid4().hex[:6]}"
|
||||
input_fp = create_seed_file("input", scope, "input.bin")
|
||||
create_seed_file("output", scope, "output.bin")
|
||||
input_name = f"{scope}-input.bin"
|
||||
output_name = f"{scope}-output.bin"
|
||||
input_fp = create_seed_file("input", scope, input_name)
|
||||
create_seed_file("output", scope, output_name)
|
||||
|
||||
trigger_sync_seed_assets(http, api_base)
|
||||
assert len(find_asset(scope)) == 2
|
||||
assert find_asset(scope, input_name)
|
||||
assert find_asset(scope, output_name)
|
||||
|
||||
input_fp.unlink()
|
||||
trigger_sync_seed_assets(http, api_base)
|
||||
|
||||
remaining = find_asset(scope)
|
||||
assert len(remaining) == 1
|
||||
assert remaining[0]["name"] == "output.bin"
|
||||
assert not find_asset(scope, input_name)
|
||||
assert find_asset(scope, output_name)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dirname", ["100%_done", "my_folder_name", "has spaces"])
|
||||
|
||||
@@ -10,9 +10,9 @@ def test_tags_present(http: requests.Session, api_base: str, seeded_asset: dict)
|
||||
body1 = r1.json()
|
||||
assert r1.status_code == 200
|
||||
names = [t["name"] for t in body1["tags"]]
|
||||
# A few system tags from migration should exist:
|
||||
# A few selected contract tags should exist.
|
||||
assert "models" in names
|
||||
assert "checkpoints" in names
|
||||
assert "model_type:checkpoints" in names
|
||||
|
||||
# Only used tags before we add anything new from this test cycle
|
||||
r2 = http.get(api_base + "/api/tags", params={"include_zero": "false"}, timeout=120)
|
||||
@@ -21,7 +21,7 @@ def test_tags_present(http: requests.Session, api_base: str, seeded_asset: dict)
|
||||
# We already seeded one asset via fixture, so used tags must be non-empty
|
||||
used_names = [t["name"] for t in body2["tags"]]
|
||||
assert "models" in used_names
|
||||
assert "checkpoints" in used_names
|
||||
assert "model_type:checkpoints" in used_names
|
||||
|
||||
# Prefix filter should refine the list
|
||||
r3 = http.get(api_base + "/api/tags", params={"include_zero": "false", "prefix": "uni"}, timeout=120)
|
||||
@@ -45,7 +45,7 @@ def test_tags_empty_usage(http: requests.Session, api_base: str, asset_factory,
|
||||
body1 = r1.json()
|
||||
assert r1.status_code == 200
|
||||
names = [t["name"] for t in body1["tags"]]
|
||||
assert "models" in names and "checkpoints" in names
|
||||
assert "models" in names and "model_type:checkpoints" in names
|
||||
|
||||
# Create a short-lived asset under input with a unique custom tag
|
||||
scope = f"tags-empty-usage-{uuid.uuid4().hex[:6]}"
|
||||
@@ -89,28 +89,28 @@ def test_tags_empty_usage(http: requests.Session, api_base: str, asset_factory,
|
||||
def test_add_and_remove_tags(http: requests.Session, api_base: str, seeded_asset: dict):
|
||||
aid = seeded_asset["id"]
|
||||
|
||||
# Add tags with duplicates and mixed case
|
||||
payload_add = {"tags": ["NewTag", "unit-tests", "newtag", "BETA"]}
|
||||
# Add tags with duplicates while preserving source case.
|
||||
payload_add = {"tags": ["NewTag", "unit-tests", "NewTag", "BETA"]}
|
||||
r1 = http.post(f"{api_base}/api/assets/{aid}/tags", json=payload_add, timeout=120)
|
||||
b1 = r1.json()
|
||||
assert r1.status_code == 200, b1
|
||||
# normalized, deduplicated; 'unit-tests' was already present from the seed
|
||||
assert set(b1["added"]) == {"newtag", "beta"}
|
||||
# stripped, deduplicated; 'unit-tests' was already present from the seed
|
||||
assert set(b1["added"]) == {"NewTag", "BETA"}
|
||||
assert set(b1["already_present"]) == {"unit-tests"}
|
||||
assert "newtag" in b1["total_tags"] and "beta" in b1["total_tags"]
|
||||
assert "NewTag" in b1["total_tags"] and "BETA" in b1["total_tags"]
|
||||
|
||||
rg = http.get(f"{api_base}/api/assets/{aid}", timeout=120)
|
||||
g = rg.json()
|
||||
assert rg.status_code == 200
|
||||
tags_now = set(g["tags"])
|
||||
assert {"newtag", "beta"}.issubset(tags_now)
|
||||
assert {"NewTag", "BETA"}.issubset(tags_now)
|
||||
|
||||
# Remove a tag and a non-existent tag
|
||||
payload_del = {"tags": ["newtag", "does-not-exist"]}
|
||||
payload_del = {"tags": ["NewTag", "does-not-exist"]}
|
||||
r2 = http.delete(f"{api_base}/api/assets/{aid}/tags", json=payload_del, timeout=120)
|
||||
b2 = r2.json()
|
||||
assert r2.status_code == 200
|
||||
assert set(b2["removed"]) == {"newtag"}
|
||||
assert set(b2["removed"]) == {"NewTag"}
|
||||
assert set(b2["not_present"]) == {"does-not-exist"}
|
||||
|
||||
# Verify remaining tags after deletion
|
||||
@@ -118,8 +118,44 @@ def test_add_and_remove_tags(http: requests.Session, api_base: str, seeded_asset
|
||||
g2 = rg2.json()
|
||||
assert rg2.status_code == 200
|
||||
tags_later = set(g2["tags"])
|
||||
assert "newtag" not in tags_later
|
||||
assert "beta" in tags_later # still present
|
||||
assert "NewTag" not in tags_later
|
||||
assert "BETA" in tags_later # still present
|
||||
|
||||
|
||||
def test_add_system_looking_tags_allowed_as_labels(
|
||||
http: requests.Session, api_base: str, seeded_asset: dict
|
||||
):
|
||||
aid = seeded_asset["id"]
|
||||
|
||||
response = http.post(
|
||||
f"{api_base}/api/assets/{aid}/tags",
|
||||
json={
|
||||
"tags": [
|
||||
"models",
|
||||
"model_type:manual",
|
||||
"model:true",
|
||||
"models:foo",
|
||||
"input:true",
|
||||
"output:true",
|
||||
"uploaded:true",
|
||||
"temp:true",
|
||||
"temporary",
|
||||
]
|
||||
},
|
||||
timeout=120,
|
||||
)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 200, body
|
||||
assert "models" in body["total_tags"]
|
||||
assert "model_type:manual" in body["total_tags"]
|
||||
assert "model:true" in body["total_tags"]
|
||||
assert "models:foo" in body["total_tags"]
|
||||
assert "input:true" in body["total_tags"]
|
||||
assert "output:true" in body["total_tags"]
|
||||
assert "uploaded:true" in body["total_tags"]
|
||||
assert "temp:true" in body["total_tags"]
|
||||
assert "temporary" in body["total_tags"]
|
||||
|
||||
|
||||
def test_tags_list_order_and_prefix(http: requests.Session, api_base: str, seeded_asset: dict):
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
import json
|
||||
import uuid
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from pathlib import Path
|
||||
|
||||
import requests
|
||||
import pytest
|
||||
|
||||
from app.assets.api.schemas_in import UploadAssetSpec
|
||||
from app.assets.api.schemas_out import Asset, AssetCreated
|
||||
from helpers import get_asset_filename
|
||||
|
||||
|
||||
def test_asset_created_inherits_hash_field():
|
||||
@@ -20,9 +23,18 @@ def test_asset_created_inherits_hash_field():
|
||||
assert AssetCreated.model_fields["hash"].annotation == Asset.model_fields["hash"].annotation
|
||||
|
||||
|
||||
def test_upload_asset_spec_ignores_subfolder_field():
|
||||
spec = UploadAssetSpec.model_validate(
|
||||
{"tags": ["input"], "subfolder": "pasted", "name": "image.png"}
|
||||
)
|
||||
|
||||
assert "subfolder" not in UploadAssetSpec.model_fields
|
||||
assert not hasattr(spec, "subfolder")
|
||||
|
||||
|
||||
def test_upload_ok_duplicate_reference(http: requests.Session, api_base: str, make_asset_bytes):
|
||||
name = "dup_a.safetensors"
|
||||
tags = ["models", "checkpoints", "unit-tests", "alpha"]
|
||||
tags = ["models", "model_type:checkpoints", "unit-tests", "alpha"]
|
||||
meta = {"purpose": "dup"}
|
||||
data = make_asset_bytes(name)
|
||||
files = {"file": (name, data, "application/octet-stream")}
|
||||
@@ -43,6 +55,8 @@ def test_upload_ok_duplicate_reference(http: requests.Session, api_base: str, ma
|
||||
assert a2["asset_hash"] == a1["asset_hash"]
|
||||
assert a2["hash"] == a1["hash"]
|
||||
assert a2["id"] != a1["id"] # new reference with same content
|
||||
assert a2.get("loader_path") is None
|
||||
assert a2.get("display_name") is None
|
||||
|
||||
# Third upload with the same data but different name also creates new AssetReference
|
||||
files = {"file": (name, data, "application/octet-stream")}
|
||||
@@ -53,12 +67,14 @@ def test_upload_ok_duplicate_reference(http: requests.Session, api_base: str, ma
|
||||
assert a3["asset_hash"] == a1["asset_hash"]
|
||||
assert a3["id"] != a1["id"]
|
||||
assert a3["id"] != a2["id"]
|
||||
assert a3.get("loader_path") is None
|
||||
assert a3.get("display_name") is None
|
||||
|
||||
|
||||
def test_upload_fastpath_from_existing_hash_no_file(http: requests.Session, api_base: str):
|
||||
# Seed a small file first
|
||||
name = "fastpath_seed.safetensors"
|
||||
tags = ["models", "checkpoints", "unit-tests"]
|
||||
tags = ["input", "unit-tests"]
|
||||
meta = {}
|
||||
files = {"file": (name, b"B" * 1024, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(tags), "name": name, "user_metadata": json.dumps(meta)}
|
||||
@@ -69,9 +85,10 @@ def test_upload_fastpath_from_existing_hash_no_file(http: requests.Session, api_
|
||||
assert b1["hash"] == h
|
||||
|
||||
# Now POST /api/assets with only hash and no file
|
||||
hash_only_tags = ["models", "checkpoints", "unit-tests", "hash-labels"]
|
||||
files = [
|
||||
("hash", (None, h)),
|
||||
("tags", (None, json.dumps(tags))),
|
||||
("tags", (None, json.dumps(hash_only_tags))),
|
||||
("name", (None, "fastpath_copy.safetensors")),
|
||||
("user_metadata", (None, json.dumps({"purpose": "copy"}))),
|
||||
]
|
||||
@@ -81,6 +98,53 @@ def test_upload_fastpath_from_existing_hash_no_file(http: requests.Session, api_
|
||||
assert b2["created_new"] is False
|
||||
assert b2["asset_hash"] == h
|
||||
assert b2["hash"] == h
|
||||
assert "models" in b2["tags"]
|
||||
assert "checkpoints" in b2["tags"]
|
||||
assert "uploaded" not in b2["tags"]
|
||||
assert not any(tag.startswith("model_type:") for tag in b2["tags"])
|
||||
assert b2.get("loader_path") is None
|
||||
assert b2.get("display_name") is None
|
||||
|
||||
rg = http.get(f"{api_base}/api/assets/{b2['id']}", timeout=120)
|
||||
detail = rg.json()
|
||||
assert rg.status_code == 200, detail
|
||||
assert detail.get("loader_path") is None
|
||||
assert detail.get("display_name") is None
|
||||
|
||||
|
||||
def test_create_from_hash_with_model_tags_does_not_synthesize_loader_path(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
seed_name = "from_hash_seed.safetensors"
|
||||
seed_tags = ["models", "model_type:checkpoints", "unit-tests"]
|
||||
files = {"file": (seed_name, b"D" * 1024, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(seed_tags),
|
||||
"name": seed_name,
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
seed_r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
seed = seed_r.json()
|
||||
assert seed_r.status_code == 201, seed
|
||||
|
||||
payload = {
|
||||
"hash": seed["asset_hash"],
|
||||
"name": "from_hash_copy.safetensors",
|
||||
"tags": ["models", "model_type:checkpoints", "unit-tests", "spoofed"],
|
||||
}
|
||||
created_r = http.post(api_base + "/api/assets/from-hash", json=payload, timeout=120)
|
||||
created = created_r.json()
|
||||
assert created_r.status_code == 201, created
|
||||
assert created["created_new"] is False
|
||||
assert created["asset_hash"] == seed["asset_hash"]
|
||||
assert created.get("loader_path") is None
|
||||
assert created.get("display_name") is None
|
||||
|
||||
detail_r = http.get(f"{api_base}/api/assets/{created['id']}", timeout=120)
|
||||
detail = detail_r.json()
|
||||
assert detail_r.status_code == 200, detail
|
||||
assert detail.get("loader_path") is None
|
||||
assert detail.get("display_name") is None
|
||||
|
||||
|
||||
def test_upload_fastpath_with_known_hash_and_file(
|
||||
@@ -88,7 +152,7 @@ def test_upload_fastpath_with_known_hash_and_file(
|
||||
):
|
||||
# Seed
|
||||
files = {"file": ("seed.safetensors", b"C" * 128, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "fp"]), "name": "seed.safetensors", "user_metadata": json.dumps({})}
|
||||
form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "fp"]), "name": "seed.safetensors", "user_metadata": json.dumps({})}
|
||||
r1 = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
b1 = r1.json()
|
||||
assert r1.status_code == 201, b1
|
||||
@@ -104,11 +168,49 @@ def test_upload_fastpath_with_known_hash_and_file(
|
||||
assert b2["created_new"] is False
|
||||
assert b2["asset_hash"] == h
|
||||
assert b2["hash"] == h
|
||||
assert "checkpoints" in b2["tags"]
|
||||
assert "uploaded" not in b2["tags"]
|
||||
assert not any(tag == "model_type:checkpoints" for tag in b2["tags"])
|
||||
|
||||
|
||||
def test_duplicate_byte_upload_is_reference_only_and_does_not_need_destination(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
data = b"duplicate-reference-only" * 64
|
||||
seed_files = {"file": ("duplicate-seed.bin", data, "application/octet-stream")}
|
||||
seed_form = {
|
||||
"tags": json.dumps(["input", "unit-tests", "duplicate-seed"]),
|
||||
"name": "duplicate-seed.bin",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
seed_response = http.post(api_base + "/api/assets", data=seed_form, files=seed_files, timeout=120)
|
||||
seed = seed_response.json()
|
||||
assert seed_response.status_code == 201, seed
|
||||
|
||||
duplicate_files = {"file": ("duplicate-copy.bin", data, "application/octet-stream")}
|
||||
duplicate_form = {
|
||||
"tags": json.dumps(["not-a-destination", "unit-tests", "duplicate-copy"]),
|
||||
"name": "duplicate-copy.bin",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
duplicate_response = http.post(
|
||||
api_base + "/api/assets", data=duplicate_form, files=duplicate_files, timeout=120
|
||||
)
|
||||
duplicate = duplicate_response.json()
|
||||
|
||||
assert duplicate_response.status_code == 200, duplicate
|
||||
assert duplicate["created_new"] is False
|
||||
assert duplicate["asset_hash"] == seed["asset_hash"]
|
||||
assert "not-a-destination" in duplicate["tags"]
|
||||
assert "uploaded" not in duplicate["tags"]
|
||||
assert "input" not in duplicate["tags"]
|
||||
assert duplicate.get("loader_path") is None
|
||||
assert duplicate.get("display_name") is None
|
||||
|
||||
|
||||
def test_upload_multiple_tags_fields_are_merged(http: requests.Session, api_base: str):
|
||||
data = [
|
||||
("tags", "models,checkpoints"),
|
||||
("tags", "models,model_type:checkpoints"),
|
||||
("tags", json.dumps(["unit-tests", "alpha"])),
|
||||
("name", "merge.safetensors"),
|
||||
("user_metadata", json.dumps({"u": 1})),
|
||||
@@ -124,7 +226,71 @@ def test_upload_multiple_tags_fields_are_merged(http: requests.Session, api_base
|
||||
detail = rg.json()
|
||||
assert rg.status_code == 200, detail
|
||||
tags = set(detail["tags"])
|
||||
assert {"models", "checkpoints", "unit-tests", "alpha"}.issubset(tags)
|
||||
assert {"models", "model_type:checkpoints", "unit-tests", "alpha"}.issubset(tags)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
(
|
||||
"tags",
|
||||
"extension",
|
||||
"expected_display_prefix",
|
||||
),
|
||||
[
|
||||
(["input", "unit-tests"], ".png", ""),
|
||||
(
|
||||
["models", "model_type:checkpoints", "unit-tests"],
|
||||
".safetensors",
|
||||
"checkpoints/",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_upload_response_includes_loader_path_and_display_name(
|
||||
tags: list[str],
|
||||
extension: str,
|
||||
expected_display_prefix: str,
|
||||
http: requests.Session,
|
||||
api_base: str,
|
||||
make_asset_bytes,
|
||||
):
|
||||
scope = f"response-paths-{uuid.uuid4().hex[:6]}"
|
||||
scoped_tags = [*tags, scope]
|
||||
name = f"asset_response_path{extension}"
|
||||
|
||||
files = {"file": (name, make_asset_bytes(name, 1024), "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(scoped_tags),
|
||||
"name": name,
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
created_r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
created = created_r.json()
|
||||
assert created_r.status_code in (200, 201), created
|
||||
stored_filename = get_asset_filename(created["asset_hash"], extension)
|
||||
expected_suffix = stored_filename
|
||||
expected_display_name = f"{expected_display_prefix}{expected_suffix}"
|
||||
# In-root loader path: model category dropped, no subfolders here -> just the filename.
|
||||
expected_loader_path = expected_suffix
|
||||
|
||||
assert created["loader_path"] == expected_loader_path
|
||||
assert created["display_name"] == expected_display_name
|
||||
assert "logical_path" not in created
|
||||
|
||||
detail_r = http.get(f"{api_base}/api/assets/{created['id']}", timeout=120)
|
||||
detail = detail_r.json()
|
||||
assert detail_r.status_code == 200, detail
|
||||
assert detail["loader_path"] == expected_loader_path
|
||||
assert detail["display_name"] == expected_display_name
|
||||
|
||||
list_r = http.get(
|
||||
api_base + "/api/assets",
|
||||
params={"include_tags": f"unit-tests,{scope}", "limit": "50"},
|
||||
timeout=120,
|
||||
)
|
||||
listed = list_r.json()
|
||||
assert list_r.status_code == 200, listed
|
||||
match = next(a for a in listed["assets"] if a["id"] == created["id"])
|
||||
assert match["loader_path"] == expected_loader_path
|
||||
assert match["display_name"] == expected_display_name
|
||||
|
||||
|
||||
@pytest.mark.parametrize("root", ["input", "output"])
|
||||
@@ -192,16 +358,55 @@ def test_create_from_hash_endpoint_404(http: requests.Session, api_base: str):
|
||||
assert body["error"]["code"] == "ASSET_NOT_FOUND"
|
||||
|
||||
|
||||
def test_create_from_hash_accepts_arbitrary_system_looking_tags(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
files = {"file": ("hash-seed.bin", b"hash-seed" * 64, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(["input", "unit-tests", "hash-seed"]),
|
||||
"name": "hash-seed.bin",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
seed_response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
seed = seed_response.json()
|
||||
assert seed_response.status_code == 201, seed
|
||||
|
||||
response = http.post(
|
||||
api_base + "/api/assets/from-hash",
|
||||
json={
|
||||
"hash": seed["asset_hash"],
|
||||
"name": "hash-copy.bin",
|
||||
"tags": [
|
||||
"models",
|
||||
"model:true",
|
||||
"models:foo",
|
||||
"temporary:true",
|
||||
"unit-tests",
|
||||
"hash-copy",
|
||||
],
|
||||
},
|
||||
timeout=120,
|
||||
)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 201, body
|
||||
assert "models" in body["tags"]
|
||||
assert "model:true" in body["tags"]
|
||||
assert "models:foo" in body["tags"]
|
||||
assert "temporary:true" in body["tags"]
|
||||
assert "uploaded" not in body["tags"]
|
||||
|
||||
|
||||
def test_upload_zero_byte_rejected(http: requests.Session, api_base: str):
|
||||
files = {"file": ("empty.safetensors", b"", "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "edge"]), "name": "empty.safetensors", "user_metadata": json.dumps({})}
|
||||
form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "edge"]), "name": "empty.safetensors", "user_metadata": json.dumps({})}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = r.json()
|
||||
assert r.status_code == 400
|
||||
assert body["error"]["code"] == "EMPTY_UPLOAD"
|
||||
|
||||
|
||||
def test_upload_invalid_root_tag_rejected(http: requests.Session, api_base: str):
|
||||
def test_upload_rejects_arbitrary_labels_without_required_destination_role(http: requests.Session, api_base: str):
|
||||
files = {"file": ("badroot.bin", b"A" * 64, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["not-a-root", "whatever"]), "name": "badroot.bin", "user_metadata": json.dumps({})}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
@@ -212,7 +417,7 @@ def test_upload_invalid_root_tag_rejected(http: requests.Session, api_base: str)
|
||||
|
||||
def test_upload_user_metadata_must_be_json(http: requests.Session, api_base: str):
|
||||
files = {"file": ("badmeta.bin", b"A" * 128, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "edge"]), "name": "badmeta.bin", "user_metadata": "{not json}"}
|
||||
form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "edge"]), "name": "badmeta.bin", "user_metadata": "{not json}"}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = r.json()
|
||||
assert r.status_code == 400
|
||||
@@ -228,7 +433,7 @@ def test_upload_requires_multipart(http: requests.Session, api_base: str):
|
||||
|
||||
def test_upload_missing_file_and_hash(http: requests.Session, api_base: str):
|
||||
files = [
|
||||
("tags", (None, json.dumps(["models", "checkpoints", "unit-tests"]))),
|
||||
("tags", (None, json.dumps(["models", "model_type:checkpoints", "unit-tests"]))),
|
||||
("name", (None, "x.safetensors")),
|
||||
]
|
||||
r = http.post(api_base + "/api/assets", files=files, timeout=120)
|
||||
@@ -237,17 +442,33 @@ def test_upload_missing_file_and_hash(http: requests.Session, api_base: str):
|
||||
assert body["error"]["code"] == "MISSING_FILE"
|
||||
|
||||
|
||||
def test_upload_models_unknown_category(http: requests.Session, api_base: str):
|
||||
def test_upload_models_unknown_model_type(http: requests.Session, api_base: str):
|
||||
files = {"file": ("m.safetensors", b"A" * 128, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models", "no_such_category", "unit-tests"]), "name": "m.safetensors"}
|
||||
form = {"tags": json.dumps(["models", "model_type:no_such_category", "unit-tests"]), "name": "m.safetensors"}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = r.json()
|
||||
assert r.status_code == 400
|
||||
assert r.status_code == 400, body
|
||||
assert body["error"]["code"] == "INVALID_BODY"
|
||||
assert body["error"]["message"].startswith("unknown models category")
|
||||
|
||||
|
||||
def test_upload_models_requires_category(http: requests.Session, api_base: str):
|
||||
@pytest.mark.parametrize("model_type", ["configs", "custom_nodes"])
|
||||
def test_upload_models_rejects_non_model_registered_folder(
|
||||
model_type: str, http: requests.Session, api_base: str
|
||||
):
|
||||
files = {"file": ("not-a-model.py", b"A" * 128, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(["models", f"model_type:{model_type}", "unit-tests"]),
|
||||
"name": "not-a-model.py",
|
||||
}
|
||||
|
||||
response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 400, body
|
||||
assert body["error"]["code"] == "INVALID_BODY"
|
||||
|
||||
|
||||
def test_upload_models_requires_model_type(http: requests.Session, api_base: str):
|
||||
files = {"file": ("nocat.safetensors", b"A" * 64, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models"]), "name": "nocat.safetensors", "user_metadata": json.dumps({})}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
@@ -256,13 +477,152 @@ def test_upload_models_requires_category(http: requests.Session, api_base: str):
|
||||
assert body["error"]["code"] == "INVALID_BODY"
|
||||
|
||||
|
||||
def test_upload_tags_traversal_guard(http: requests.Session, api_base: str):
|
||||
def test_upload_extra_tags_are_labels_not_path_components(http: requests.Session, api_base: str):
|
||||
files = {"file": ("evil.safetensors", b"A" * 256, "application/octet-stream")}
|
||||
form = {"tags": json.dumps(["models", "checkpoints", "unit-tests", "..", "zzz"]), "name": "evil.safetensors"}
|
||||
form = {"tags": json.dumps(["models", "model_type:checkpoints", "unit-tests", "..", "zzz"]), "name": "evil.safetensors"}
|
||||
r = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = r.json()
|
||||
assert r.status_code == 400
|
||||
assert body["error"]["code"] in ("BAD_REQUEST", "INVALID_BODY")
|
||||
assert r.status_code == 201, body
|
||||
assert ".." in body["tags"]
|
||||
assert "zzz" in body["tags"]
|
||||
assert "models" in body["tags"]
|
||||
assert "model_type:checkpoints" in body["tags"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("subfolder", "expected_tag", "unexpected_tags"),
|
||||
[
|
||||
("custom/session", None, {"custom", "session"}),
|
||||
("pasted", "pasted", set()),
|
||||
],
|
||||
)
|
||||
def test_upload_image_accepts_arbitrary_subfolder_but_only_known_values_become_tags(
|
||||
http: requests.Session,
|
||||
api_base: str,
|
||||
comfy_tmp_base_dir: Path,
|
||||
subfolder: str,
|
||||
expected_tag: str | None,
|
||||
unexpected_tags: set[str],
|
||||
):
|
||||
name = f"upload-image-{uuid.uuid4().hex}.png"
|
||||
files = {"image": (name, b"image-upload" * 64, "image/png")}
|
||||
form = {"type": "input", "subfolder": subfolder}
|
||||
|
||||
response = http.post(api_base + "/upload/image", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 200, body
|
||||
assert body["subfolder"] == subfolder
|
||||
assert (comfy_tmp_base_dir / "input" / subfolder / body["name"]).exists()
|
||||
|
||||
asset = body["asset"]
|
||||
tags = set(asset["tags"])
|
||||
assert "input" in tags
|
||||
assert "uploaded" in tags
|
||||
if expected_tag:
|
||||
assert expected_tag in tags
|
||||
assert tags.isdisjoint(unexpected_tags)
|
||||
|
||||
|
||||
def test_multipart_upload_accepts_system_looking_extra_labels(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
files = {"file": ("relaxed-labels.bin", b"relaxed" * 64, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(
|
||||
[
|
||||
"input",
|
||||
"unit-tests",
|
||||
"model:true",
|
||||
"models:foo",
|
||||
"temporary",
|
||||
"uploaded:true",
|
||||
]
|
||||
),
|
||||
"name": "relaxed-labels.bin",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 201, body
|
||||
assert "input" in body["tags"]
|
||||
assert "model:true" in body["tags"]
|
||||
assert "models:foo" in body["tags"]
|
||||
assert "temporary" in body["tags"]
|
||||
assert "uploaded:true" in body["tags"]
|
||||
|
||||
|
||||
def test_multipart_upload_rejects_ambiguous_destination_roles(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
files = {"file": ("ambiguous.bin", b"ambiguous" * 64, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(["input", "output", "unit-tests"]),
|
||||
"name": "ambiguous.bin",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 400, body
|
||||
assert body["error"]["code"] == "INVALID_BODY"
|
||||
|
||||
|
||||
def test_multipart_upload_rejects_multiple_model_types_for_models_destination(
|
||||
http: requests.Session, api_base: str
|
||||
):
|
||||
files = {"file": ("ambiguous-model.safetensors", b"ambiguous-model" * 64, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(
|
||||
["models", "model_type:checkpoints", "model_type:loras", "unit-tests"]
|
||||
),
|
||||
"name": "ambiguous-model.safetensors",
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 400, body
|
||||
assert body["error"]["code"] == "INVALID_BODY"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("tags", "expected_root", "extension"),
|
||||
[
|
||||
(["input", "unit-tests", "upload-location-input"], "input", ".bin"),
|
||||
(["output", "unit-tests", "upload-location-output"], "output", ".bin"),
|
||||
(
|
||||
["models", "model_type:checkpoints", "unit-tests", "upload-location-model"],
|
||||
"models/checkpoints",
|
||||
".safetensors",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_multipart_upload_role_selects_write_location(
|
||||
http: requests.Session,
|
||||
api_base: str,
|
||||
comfy_tmp_base_dir: Path,
|
||||
tags: list[str],
|
||||
expected_root: str,
|
||||
extension: str,
|
||||
):
|
||||
role = next(tag for tag in tags if tag in {"input", "models", "output"})
|
||||
name = f"{role}-role-upload{extension}"
|
||||
files = {"file": (name, f"{role}-role-bytes".encode() * 64, "application/octet-stream")}
|
||||
form = {
|
||||
"tags": json.dumps(tags),
|
||||
"name": name,
|
||||
"user_metadata": json.dumps({}),
|
||||
}
|
||||
|
||||
response = http.post(api_base + "/api/assets", data=form, files=files, timeout=120)
|
||||
body = response.json()
|
||||
|
||||
assert response.status_code == 201, body
|
||||
stored_name = get_asset_filename(body["asset_hash"], extension)
|
||||
expected_disk_path = comfy_tmp_base_dir / expected_root / stored_name
|
||||
assert expected_disk_path.exists()
|
||||
|
||||
|
||||
def test_upload_empty_tags_rejected(http: requests.Session, api_base: str):
|
||||
|
||||
@@ -0,0 +1,186 @@
|
||||
"""SeedVR2 conditioning node regression tests."""
|
||||
|
||||
import importlib
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
from comfy.ldm.seedvr.constants import SEEDVR2_LATENT_CHANNELS
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
|
||||
_SENTINEL = object()
|
||||
_TARGETS = (
|
||||
("comfy.model_management", "comfy"),
|
||||
("comfy_extras.nodes_seedvr", "comfy_extras"),
|
||||
)
|
||||
|
||||
|
||||
def _import_nodes_seedvr_isolated():
|
||||
"""Import comfy_extras.nodes_seedvr with comfy.model_management mocked."""
|
||||
priors = []
|
||||
for mod_name, parent_name in _TARGETS:
|
||||
prior_mod = sys.modules.get(mod_name, _SENTINEL)
|
||||
parent = sys.modules.get(parent_name)
|
||||
attr = mod_name.split(".")[-1]
|
||||
prior_attr = (
|
||||
getattr(parent, attr, _SENTINEL) if parent is not None else _SENTINEL
|
||||
)
|
||||
priors.append((mod_name, parent_name, attr, prior_mod, prior_attr))
|
||||
|
||||
mock_mm = MagicMock()
|
||||
for fn in (
|
||||
"xformers_enabled", "xformers_enabled_vae",
|
||||
"pytorch_attention_enabled", "pytorch_attention_enabled_vae",
|
||||
"sage_attention_enabled", "flash_attention_enabled",
|
||||
"is_intel_xpu",
|
||||
):
|
||||
getattr(mock_mm, fn).return_value = False
|
||||
tv = torch.version.__version__.split(".")
|
||||
mock_mm.torch_version_numeric = (int(tv[0]), int(tv[1]))
|
||||
mock_mm.WINDOWS = False
|
||||
sys.modules["comfy.model_management"] = mock_mm
|
||||
if sys.modules.get("comfy") is None:
|
||||
importlib.import_module("comfy")
|
||||
comfy_pkg = sys.modules.get("comfy")
|
||||
if comfy_pkg is not None:
|
||||
setattr(comfy_pkg, "model_management", mock_mm)
|
||||
nodes_seedvr = sys.modules.get("comfy_extras.nodes_seedvr") or (
|
||||
importlib.import_module("comfy_extras.nodes_seedvr")
|
||||
)
|
||||
|
||||
def _restore():
|
||||
for mod_name, parent_name, attr, prior_mod, prior_attr in priors:
|
||||
if prior_mod is _SENTINEL:
|
||||
sys.modules.pop(mod_name, None)
|
||||
else:
|
||||
sys.modules[mod_name] = prior_mod
|
||||
parent = sys.modules.get(parent_name)
|
||||
if parent is None:
|
||||
continue
|
||||
if prior_attr is _SENTINEL:
|
||||
if hasattr(parent, attr):
|
||||
delattr(parent, attr)
|
||||
else:
|
||||
setattr(parent, attr, prior_attr)
|
||||
|
||||
return nodes_seedvr, _restore
|
||||
|
||||
|
||||
class _Rope(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.freqs = nn.Parameter(torch.zeros(4))
|
||||
|
||||
|
||||
class _Block(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.rope = _Rope()
|
||||
|
||||
|
||||
class _DiffusionModel(nn.Module):
|
||||
def __init__(self, n_blocks=3, conditioning_dtype=torch.float32):
|
||||
super().__init__()
|
||||
self.blocks = nn.ModuleList([_Block() for _ in range(n_blocks)])
|
||||
self.register_buffer("positive_conditioning", torch.ones((2, 4), dtype=conditioning_dtype))
|
||||
self.register_buffer("negative_conditioning", torch.zeros((3, 4), dtype=conditioning_dtype))
|
||||
|
||||
|
||||
class _ModelInner:
|
||||
def __init__(self, diffusion_model):
|
||||
self.diffusion_model = diffusion_model
|
||||
|
||||
|
||||
class _ModelPatcher:
|
||||
def __init__(self, diffusion_model):
|
||||
self.model = _ModelInner(diffusion_model)
|
||||
|
||||
|
||||
def test_seedvr2_conditioning_schema_exposes_conditioning_outputs():
|
||||
nodes_seedvr, restore = _import_nodes_seedvr_isolated()
|
||||
try:
|
||||
schema = nodes_seedvr.SeedVR2Conditioning.define_schema()
|
||||
assert [input_item.id for input_item in schema.inputs] == [
|
||||
"model",
|
||||
"vae_conditioning",
|
||||
]
|
||||
assert schema.inputs[1].display_name == "latent"
|
||||
assert [output.display_name for output in schema.outputs] == [
|
||||
"positive",
|
||||
"negative",
|
||||
]
|
||||
finally:
|
||||
restore()
|
||||
|
||||
|
||||
def test_seedvr2_conditioning_rejects_wrong_latent_channels():
|
||||
nodes_seedvr, restore = _import_nodes_seedvr_isolated()
|
||||
try:
|
||||
patcher = _ModelPatcher(_DiffusionModel())
|
||||
vae_conditioning = {"samples": torch.zeros(1, 8, 2, 2, 2)}
|
||||
|
||||
with pytest.raises(ValueError, match=f"{SEEDVR2_LATENT_CHANNELS} channels"):
|
||||
nodes_seedvr.SeedVR2Conditioning.execute(patcher, vae_conditioning)
|
||||
finally:
|
||||
restore()
|
||||
|
||||
|
||||
def test_seedvr2_conditioning_returns_conditioning_deterministically():
|
||||
nodes_seedvr, restore = _import_nodes_seedvr_isolated()
|
||||
try:
|
||||
diffusion_model = _DiffusionModel()
|
||||
patcher = _ModelPatcher(diffusion_model)
|
||||
samples = torch.arange(
|
||||
1,
|
||||
1 + SEEDVR2_LATENT_CHANNELS * 3 * 2 * 2,
|
||||
dtype=torch.float32,
|
||||
).reshape(1, SEEDVR2_LATENT_CHANNELS, 3, 2, 2)
|
||||
vae_conditioning = {"samples": samples}
|
||||
|
||||
first_positive, first_negative = (
|
||||
nodes_seedvr.SeedVR2Conditioning.execute(
|
||||
patcher,
|
||||
vae_conditioning,
|
||||
)
|
||||
)
|
||||
second_positive, second_negative = (
|
||||
nodes_seedvr.SeedVR2Conditioning.execute(
|
||||
patcher,
|
||||
vae_conditioning,
|
||||
)
|
||||
)
|
||||
|
||||
channel_last = samples.movedim(1, -1).contiguous()
|
||||
expected_condition = torch.cat(
|
||||
[
|
||||
channel_last,
|
||||
torch.ones((*channel_last.shape[:-1], 1)),
|
||||
],
|
||||
dim=-1,
|
||||
).movedim(-1, 1)
|
||||
|
||||
assert torch.equal(
|
||||
first_positive[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
assert torch.equal(
|
||||
second_positive[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
assert torch.equal(
|
||||
first_negative[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
assert torch.equal(
|
||||
second_negative[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
finally:
|
||||
restore()
|
||||
@@ -0,0 +1,55 @@
|
||||
import importlib
|
||||
import inspect
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
|
||||
def test_seedvr_node_signature_matches_schema():
|
||||
mock_mm = MagicMock()
|
||||
mock_mm.xformers_enabled.return_value = False
|
||||
mock_mm.xformers_enabled_vae.return_value = False
|
||||
mock_mm.sage_attention_enabled.return_value = False
|
||||
mock_mm.flash_attention_enabled.return_value = False
|
||||
|
||||
sentinel = object()
|
||||
prior_cpu = cli_args.cpu
|
||||
cli_args.cpu = True
|
||||
prior_module = sys.modules.get("comfy_extras.nodes_seedvr", sentinel)
|
||||
comfy_pkg = sys.modules.get("comfy")
|
||||
prior_mm_attr = getattr(comfy_pkg, "model_management", sentinel) if comfy_pkg else sentinel
|
||||
|
||||
with patch.dict(sys.modules, {"comfy.model_management": mock_mm}):
|
||||
if comfy_pkg is not None:
|
||||
setattr(comfy_pkg, "model_management", mock_mm)
|
||||
sys.modules.pop("comfy_extras.nodes_seedvr", None)
|
||||
try:
|
||||
nodes_seedvr = importlib.import_module("comfy_extras.nodes_seedvr")
|
||||
for node_cls in (nodes_seedvr.SeedVR2Preprocess, nodes_seedvr.SeedVR2PostProcessing, nodes_seedvr.SeedVR2Conditioning):
|
||||
schema_ids = [i.id for i in node_cls.define_schema().inputs]
|
||||
exec_params = [
|
||||
p for p in inspect.signature(node_cls.execute).parameters.keys()
|
||||
if p != "cls"
|
||||
]
|
||||
assert schema_ids == exec_params, (
|
||||
f"{node_cls.__name__} schema/execute drift: "
|
||||
f"schema_ids={schema_ids}, exec_params={exec_params}"
|
||||
)
|
||||
finally:
|
||||
cli_args.cpu = prior_cpu
|
||||
if prior_module is sentinel:
|
||||
sys.modules.pop("comfy_extras.nodes_seedvr", None)
|
||||
else:
|
||||
sys.modules["comfy_extras.nodes_seedvr"] = prior_module
|
||||
if comfy_pkg is not None:
|
||||
if prior_mm_attr is sentinel:
|
||||
if hasattr(comfy_pkg, "model_management"):
|
||||
delattr(comfy_pkg, "model_management")
|
||||
else:
|
||||
setattr(comfy_pkg, "model_management", prior_mm_attr)
|
||||
@@ -0,0 +1,51 @@
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
from comfy_extras import nodes_seedvr # noqa: E402
|
||||
|
||||
|
||||
def _schema_ids(items):
|
||||
return [item.id for item in items]
|
||||
|
||||
|
||||
def test_seedvr2_post_processing_schema():
|
||||
schema = nodes_seedvr.SeedVR2PostProcessing.define_schema()
|
||||
|
||||
assert _schema_ids(schema.inputs) == ["images", "original_resized_images", "color_correction_method"]
|
||||
assert schema.inputs[2].options == ["lab", "wavelet", "adain", "none"]
|
||||
assert schema.inputs[2].default == "lab"
|
||||
assert schema.outputs[0].get_io_type() == "IMAGE"
|
||||
|
||||
|
||||
def test_seedvr2_post_processing_oom_error_uses_color_correction_method(monkeypatch):
|
||||
decoded = torch.full((1, 3, 4, 4), 0.25)
|
||||
reference = torch.full((1, 3, 4, 4), 0.75)
|
||||
|
||||
def _lab(content, style):
|
||||
raise torch.cuda.OutOfMemoryError("CUDA out of memory")
|
||||
|
||||
monkeypatch.setattr(nodes_seedvr.comfy.model_management, "vae_device", lambda: torch.device("cpu"))
|
||||
monkeypatch.setattr(nodes_seedvr.comfy.model_management, "get_free_memory", lambda device: 1_000_000)
|
||||
|
||||
with patch.object(nodes_seedvr, "lab_color_transfer", _lab):
|
||||
with pytest.raises(RuntimeError) as excinfo:
|
||||
nodes_seedvr.SeedVR2PostProcessing._color_transfer_chunked(
|
||||
decoded, reference, torch.device("cpu"), "lab",
|
||||
)
|
||||
assert "color_correction_method=lab" in str(excinfo.value)
|
||||
assert " method=lab" not in str(excinfo.value)
|
||||
|
||||
|
||||
def test_seedvr2_post_processing_unknown_color_correction_method_raises():
|
||||
decoded = torch.zeros(1, 2, 4, 4, 3)
|
||||
original = torch.zeros(1, 2, 4, 4, 3)
|
||||
with pytest.raises(ValueError) as excinfo:
|
||||
nodes_seedvr.SeedVR2PostProcessing.execute(decoded, original, "bogus")
|
||||
assert "color_correction_method" in str(excinfo.value)
|
||||
@@ -0,0 +1,77 @@
|
||||
"""SeedVR2 temporal chunk/merge node regression tests."""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
from comfy.ldm.seedvr.constants import (
|
||||
BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE,
|
||||
SEEDVR2_CHUNK_GIB_PER_MPX_FRAME,
|
||||
SEEDVR2_CHUNK_RESERVED_GIB,
|
||||
SEEDVR2_CHUNK_SIGMA_GIB,
|
||||
SEEDVR2_CHUNK_SIGMA_K,
|
||||
SEEDVR2_LATENT_CHANNELS,
|
||||
)
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
import comfy.model_management # noqa: E402
|
||||
from comfy_extras.nodes_seedvr import SeedVR2TemporalChunk, SeedVR2TemporalMerge, _seedvr2_chunk_crossfade_weights # noqa: E402
|
||||
|
||||
def _latent(t_latent, h=8, w=8, b=1):
|
||||
g = torch.Generator().manual_seed(7)
|
||||
return {"samples": torch.randn(b, SEEDVR2_LATENT_CHANNELS, t_latent, h, w, generator=g)}
|
||||
|
||||
def _split(latent, frames_per_chunk, temporal_overlap, chunking_mode="manual"):
|
||||
combo = {"chunking_mode": chunking_mode}
|
||||
if chunking_mode != "auto":
|
||||
combo["frames_per_chunk"] = frames_per_chunk
|
||||
return SeedVR2TemporalChunk.execute(latent, temporal_overlap, combo).args
|
||||
|
||||
def _merge(chunks, temporal_overlap):
|
||||
return SeedVR2TemporalMerge.execute(chunks, [temporal_overlap]).args[0]
|
||||
|
||||
def test_chunk_temporal_windows_and_validation():
|
||||
with pytest.raises(ValueError, match="4n\\+1"):
|
||||
_split(_latent(9), 20, 0)
|
||||
with pytest.raises(ValueError, match="5-D"):
|
||||
_split({"samples": torch.zeros(1, SEEDVR2_LATENT_CHANNELS * 9, 8, 8)}, 21, 0)
|
||||
with pytest.raises(ValueError, match="chunking_mode"):
|
||||
_split(_latent(13), 21, 0, "adaptive")
|
||||
latent = _latent(13)
|
||||
chunks, overlap = _split(latent, 21, 2) # chunk_latent=6, step=4 -> [0:6], [4:10], [8:13]
|
||||
assert overlap == 2 and [c["samples"].shape[2] for c in chunks] == [6, 6, 5]
|
||||
assert all(torch.equal(c["samples"], latent["samples"][:, :, s:e]) for c, (s, e) in zip(chunks, [(0, 6), (4, 10), (8, 13)]))
|
||||
assert len(_split(_latent(13), 21, 999)[0]) == 8 # overlap clamps to chunk_latent-1 -> step=1
|
||||
assert (r := _split(_latent(5), 21, 3)) and len(r[0]) == 1 and r[1] == 0 # t_pixel <= 21: passthrough
|
||||
|
||||
def test_chunk_auto_mode_applies_vram_law(monkeypatch):
|
||||
mpx_per_frame = (32 * 32) * (BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE ** 2) / 1e6
|
||||
free_gb = (
|
||||
SEEDVR2_CHUNK_RESERVED_GIB
|
||||
+ SEEDVR2_CHUNK_SIGMA_K * SEEDVR2_CHUNK_SIGMA_GIB
|
||||
+ 5.1 * SEEDVR2_CHUNK_GIB_PER_MPX_FRAME * mpx_per_frame
|
||||
)
|
||||
monkeypatch.setattr(comfy.model_management, "get_free_memory", lambda dev=None: free_gb * (1024 ** 3))
|
||||
assert [c["samples"].shape[2] for c in _split(_latent(13, h=32, w=32), 1, 0, "auto")[0]] == [5, 5, 3]
|
||||
assert _split(_latent(13, h=32, w=32, b=2), 1, 0, "auto")[0][0]["samples"].shape[2] == 2 # batch halves the chunk
|
||||
|
||||
def test_merge_crossfade_and_reassembly():
|
||||
latent = _latent(13)
|
||||
latent["noise_mask"] = torch.rand(1, 1, 13, 8, 8)
|
||||
latent["batch_index"] = [0]
|
||||
merged = _merge(_split(latent, 21, 0)[0], 0)
|
||||
assert torch.equal(merged["samples"], latent["samples"])
|
||||
assert "noise_mask" not in merged and merged["batch_index"] == [0]
|
||||
assert torch.allclose(_merge(_split(latent, 21, 3)[0], 3)["samples"], latent["samples"], atol=1e-6)
|
||||
w = _seedvr2_chunk_crossfade_weights(3, merged["samples"].device, merged["samples"].dtype)
|
||||
assert w[0] == 1.0 and w[-1] == 0.0 and torch.all(w[:-1] >= w[1:])
|
||||
ones, zeros = {"samples": torch.ones(1, SEEDVR2_LATENT_CHANNELS, 6, 8, 8)}, {"samples": torch.zeros(1, SEEDVR2_LATENT_CHANNELS, 6, 8, 8)}
|
||||
fused = _merge([ones, zeros], 3)["samples"] # overlap equals w: prev fades out, next fades in
|
||||
assert torch.equal(fused[:, :, 3:6], w.view(1, 1, 3, 1, 1).expand(1, SEEDVR2_LATENT_CHANNELS, 3, 8, 8))
|
||||
assert torch.equal(fused[:, :, :3], ones["samples"][:, :, :3]) and torch.equal(fused[:, :, 6:], zeros["samples"][:, :, :3])
|
||||
short = _split(latent, 21, 2)[0]
|
||||
short[0]["samples"] = short[0]["samples"][:, :, :4]
|
||||
with pytest.raises(ValueError, match="only the final chunk may be shorter"):
|
||||
_merge(short, 2)
|
||||
@@ -15,7 +15,7 @@ if not has_gpu():
|
||||
args.cpu = True
|
||||
|
||||
from comfy import ops
|
||||
from comfy.quant_ops import QuantizedTensor
|
||||
from comfy.quant_ops import QUANT_ALGOS, QuantizedTensor
|
||||
import comfy.utils
|
||||
|
||||
|
||||
@@ -283,7 +283,59 @@ class TestMixedPrecisionOps(unittest.TestCase):
|
||||
saved = model.state_dict()
|
||||
saved_conf = json.loads(saved["layer.comfy_quant"].numpy().tobytes())
|
||||
self.assertTrue(saved_conf["convrot"])
|
||||
|
||||
def test_convrot_w4a4_loads_into_params(self):
|
||||
"""ConvRot W4A4 checkpoints must load as the dedicated kitchen layout."""
|
||||
if "convrot_w4a4" not in QUANT_ALGOS:
|
||||
self.skipTest("comfy_kitchen does not provide ConvRot W4A4")
|
||||
|
||||
torch.manual_seed(456)
|
||||
layer_quant_config = {
|
||||
"layer": {
|
||||
"format": "convrot_w4a4",
|
||||
"convrot_groupsize": 256,
|
||||
"linear_dtype": "int8",
|
||||
}
|
||||
}
|
||||
weight = torch.randn(16, 256, dtype=torch.bfloat16)
|
||||
bias = torch.randn(16, dtype=torch.bfloat16)
|
||||
q_weight = QuantizedTensor.from_float(
|
||||
weight,
|
||||
"TensorCoreConvRotW4A4Layout",
|
||||
convrot_groupsize=256,
|
||||
quant_group_size=64,
|
||||
)
|
||||
state_dict = {
|
||||
"layer.weight": q_weight._qdata,
|
||||
"layer.bias": bias,
|
||||
"layer.weight_scale": q_weight._params.scale,
|
||||
}
|
||||
|
||||
state_dict, _ = comfy.utils.convert_old_quants(
|
||||
state_dict,
|
||||
metadata={"_quantization_metadata": json.dumps({"layers": layer_quant_config})},
|
||||
)
|
||||
model = torch.nn.Module()
|
||||
model.layer = ops.mixed_precision_ops({}).Linear(256, 16, device="cpu", dtype=torch.bfloat16)
|
||||
model.load_state_dict(state_dict, strict=False)
|
||||
|
||||
self.assertIsInstance(model.layer.weight, QuantizedTensor)
|
||||
self.assertEqual(model.layer.weight._layout_cls, "TensorCoreConvRotW4A4Layout")
|
||||
self.assertEqual(model.layer.weight._params.convrot_groupsize, 256)
|
||||
self.assertEqual(model.layer.weight._params.quant_group_size, 64)
|
||||
self.assertEqual(model.layer.weight._params.linear_dtype, "int8")
|
||||
|
||||
input_tensor = torch.randn(4, 256, dtype=torch.bfloat16)
|
||||
loaded_out = model.layer(input_tensor)
|
||||
ref_out = torch.nn.functional.linear(input_tensor, q_weight, bias)
|
||||
self.assertTrue(torch.equal(loaded_out, ref_out))
|
||||
|
||||
saved = model.state_dict()
|
||||
saved_conf = json.loads(saved["layer.comfy_quant"].numpy().tobytes())
|
||||
self.assertEqual(saved_conf["format"], "convrot_w4a4")
|
||||
self.assertEqual(saved_conf["convrot_groupsize"], 256)
|
||||
self.assertEqual(saved_conf["linear_dtype"], "int8")
|
||||
self.assertNotIn("quant_group_size", saved_conf)
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -2,7 +2,7 @@ from collections import defaultdict
|
||||
|
||||
import torch
|
||||
|
||||
from comfy.model_detection import detect_unet_config, model_config_from_unet_config
|
||||
from comfy.model_detection import detect_unet_config, model_config_from_unet, model_config_from_unet_config
|
||||
import comfy.supported_models
|
||||
|
||||
|
||||
@@ -73,6 +73,49 @@ def _make_flux_schnell_comfyui_sd():
|
||||
return sd
|
||||
|
||||
|
||||
def _make_seedvr2_7b_separate_mm_sd():
|
||||
return {
|
||||
"blocks.35.mlp.vid.proj_out.weight": torch.empty(3072, 1),
|
||||
"positive_conditioning": torch.empty(58, 5120),
|
||||
"negative_conditioning": torch.empty(64, 5120),
|
||||
}
|
||||
|
||||
|
||||
def _make_seedvr2_7b_shared_mm_sd():
|
||||
return {
|
||||
"blocks.35.mlp.all.proj_in_gate.weight": torch.empty(1, 1),
|
||||
"positive_conditioning": torch.empty(58, 5120),
|
||||
"negative_conditioning": torch.empty(64, 5120),
|
||||
}
|
||||
|
||||
|
||||
def _make_seedvr2_3b_shared_mm_sd():
|
||||
return {
|
||||
"blocks.31.mlp.all.proj_in_gate.weight": torch.empty(1, 1),
|
||||
"positive_conditioning": torch.empty(58, 5120),
|
||||
"negative_conditioning": torch.empty(64, 5120),
|
||||
}
|
||||
|
||||
|
||||
def _make_pid_v1_5_sd(latent_proj_channels=16):
|
||||
sd = {
|
||||
"pixel_embedder.proj.weight": torch.empty(16, 3, device="meta"),
|
||||
"lq_proj.latent_proj.0.weight": torch.empty(1024, latent_proj_channels, 3, 3, device="meta"),
|
||||
"lq_proj.pit_head.weight": torch.empty(1536, 1024, device="meta"),
|
||||
"lq_proj.gate_modules.0.content_proj.weight": torch.empty(1, 3072, device="meta"),
|
||||
"pixel_blocks.0.attn.q_norm.weight": torch.empty(72, device="meta"),
|
||||
"pixel_blocks.0.adaLN_modulation.0.weight": torch.empty(24576, 1536, device="meta"),
|
||||
"pixel_blocks.0.adaLN_modulation.0.bias": torch.empty(24576, device="meta"),
|
||||
}
|
||||
for i in range(7):
|
||||
sd[f"lq_proj.gate_modules.{i}.log_alpha"] = torch.empty((), device="meta")
|
||||
return sd
|
||||
|
||||
|
||||
def _add_model_diffusion_prefix(sd):
|
||||
return {f"model.diffusion_model.{k}": v for k, v in sd.items()}
|
||||
|
||||
|
||||
class TestModelDetection:
|
||||
"""Verify that first-match model detection selects the correct model
|
||||
based on list ordering and unet_config specificity."""
|
||||
@@ -125,6 +168,96 @@ class TestModelDetection:
|
||||
assert model_config is not None
|
||||
assert type(model_config).__name__ == "FluxSchnell"
|
||||
|
||||
def test_seedvr2_7b_separate_mm_detection_config(self):
|
||||
sd = _make_seedvr2_7b_separate_mm_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config is not None
|
||||
assert unet_config["image_model"] == "seedvr2"
|
||||
assert unet_config["vid_dim"] == 3072
|
||||
assert unet_config["heads"] == 24
|
||||
assert unet_config["num_layers"] == 36
|
||||
assert unet_config["mm_layers"] == 36
|
||||
assert unet_config["mlp_type"] == "normal"
|
||||
assert unet_config["rope_type"] == "rope3d"
|
||||
assert unet_config["rope_dim"] == 64
|
||||
|
||||
def test_seedvr2_7b_shared_mm_detection_config(self):
|
||||
sd = _make_seedvr2_7b_shared_mm_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config is not None
|
||||
assert unet_config["image_model"] == "seedvr2"
|
||||
assert unet_config["vid_dim"] == 3072
|
||||
assert unet_config["heads"] == 24
|
||||
assert unet_config["num_layers"] == 36
|
||||
assert unet_config["mm_layers"] == 10
|
||||
assert unet_config["mlp_type"] == "swiglu"
|
||||
assert unet_config["rope_type"] == "rope3d"
|
||||
assert unet_config["rope_dim"] == 64
|
||||
|
||||
def test_seedvr2_3b_shared_mm_detection_config(self):
|
||||
sd = _make_seedvr2_3b_shared_mm_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config is not None
|
||||
assert unet_config["image_model"] == "seedvr2"
|
||||
assert unet_config["vid_dim"] == 2560
|
||||
assert unet_config["heads"] == 20
|
||||
assert unet_config["num_layers"] == 32
|
||||
assert unet_config["mlp_type"] == "swiglu"
|
||||
|
||||
def test_seedvr2_model_match_requires_conditioning_tensors(self):
|
||||
sd = _make_seedvr2_7b_shared_mm_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "SeedVR2"
|
||||
|
||||
del sd["positive_conditioning"]
|
||||
assert model_config_from_unet_config(unet_config, sd) is None
|
||||
|
||||
def test_seedvr2_model_match_accepts_full_checkpoint_prefix(self):
|
||||
sd = _add_model_diffusion_prefix(_make_seedvr2_7b_shared_mm_sd())
|
||||
|
||||
assert type(model_config_from_unet(sd, "model.diffusion_model.")).__name__ == "SeedVR2"
|
||||
|
||||
def test_pid_v1_5_detection(self):
|
||||
sd = _make_pid_v1_5_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config == {
|
||||
"image_model": "pid",
|
||||
"lq_latent_channels": 16,
|
||||
"lq_hidden_dim": 1024,
|
||||
"latent_spatial_down_factor": 8,
|
||||
"lq_interval": 2,
|
||||
"lq_latent_unpatchify_factor": 1,
|
||||
"lq_conv_padding_mode": "replicate",
|
||||
"lq_gate_per_token": True,
|
||||
"pit_lq_inject": True,
|
||||
"rope_ref_h": 2048,
|
||||
"rope_ref_w": 2048,
|
||||
}
|
||||
assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "PiD"
|
||||
|
||||
def test_pid_v1_5_flux2_detection(self):
|
||||
unet_config = detect_unet_config(_make_pid_v1_5_sd(latent_proj_channels=32), "")
|
||||
|
||||
assert unet_config["lq_latent_channels"] == 128
|
||||
assert unet_config["latent_spatial_down_factor"] == 16
|
||||
assert unet_config["lq_latent_unpatchify_factor"] == 2
|
||||
|
||||
def test_pid_v1_5_pixel_adaln_conversion(self):
|
||||
sd = _make_pid_v1_5_sd()
|
||||
model_config = model_config_from_unet_config(detect_unet_config(sd, ""), sd)
|
||||
processed = model_config.process_unet_state_dict(sd)
|
||||
|
||||
assert processed["pixel_blocks.0.attn.q_norm.weight"].shape == (72,)
|
||||
assert processed["pixel_blocks.0.adaLN_modulation_msa.weight"].shape == (12288, 1536)
|
||||
assert processed["pixel_blocks.0.adaLN_modulation_mlp.weight"].shape == (12288, 1536)
|
||||
assert processed["pixel_blocks.0.adaLN_modulation_msa.bias"].shape == (12288,)
|
||||
assert processed["pixel_blocks.0.adaLN_modulation_mlp.bias"].shape == (12288,)
|
||||
|
||||
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,74 @@
|
||||
"""Regression tests for the SeedVR2 VAE forward return contract."""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
from comfy.ldm.seedvr.vae import SEEDVR2_LATENT_CHANNELS, VideoAutoencoderKL # noqa: E402
|
||||
|
||||
|
||||
_LATENT_SHAPE = (1, SEEDVR2_LATENT_CHANNELS, 2, 2, 2)
|
||||
_DECODED_SHAPE = (1, 3, 5, 16, 16)
|
||||
_INPUT_ENCODE_SHAPE = (1, 3, 5, 16, 16)
|
||||
_INPUT_DECODE_SHAPE = _LATENT_SHAPE
|
||||
|
||||
|
||||
class _StubVAE(VideoAutoencoderKL):
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self._encode_out = torch.zeros(*_LATENT_SHAPE)
|
||||
self._decode_out = torch.zeros(*_DECODED_SHAPE)
|
||||
|
||||
def encode(self, x, return_dict=True):
|
||||
return self._encode_out
|
||||
|
||||
def decode_(self, z, return_dict=True):
|
||||
return self._decode_out
|
||||
|
||||
|
||||
def test_forward_encode_returns_tensor():
|
||||
vae = _StubVAE()
|
||||
x = torch.zeros(*_INPUT_ENCODE_SHAPE)
|
||||
result = vae.forward(x, mode="encode")
|
||||
assert type(result) is torch.Tensor
|
||||
assert result.shape == torch.Size(_LATENT_SHAPE)
|
||||
|
||||
|
||||
def test_forward_decode_returns_tensor():
|
||||
vae = _StubVAE()
|
||||
z = torch.zeros(*_INPUT_DECODE_SHAPE)
|
||||
result = vae.forward(z, mode="decode")
|
||||
assert type(result) is torch.Tensor
|
||||
assert result.shape == torch.Size(_DECODED_SHAPE)
|
||||
|
||||
|
||||
class _TupleReturningStubVAE(VideoAutoencoderKL):
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self._encode_tensor = torch.zeros(*_LATENT_SHAPE)
|
||||
self._decode_tensor = torch.zeros(*_DECODED_SHAPE)
|
||||
|
||||
def encode(self, x, return_dict=True):
|
||||
return (self._encode_tensor,)
|
||||
|
||||
def decode_(self, z, return_dict=True):
|
||||
return (self._decode_tensor,)
|
||||
|
||||
|
||||
def test_forward_all_unwraps_one_tuple_at_each_step():
|
||||
vae = _TupleReturningStubVAE()
|
||||
x = torch.zeros(*_INPUT_ENCODE_SHAPE)
|
||||
result = vae.forward(x, mode="all")
|
||||
assert type(result) is torch.Tensor
|
||||
assert result.shape == torch.Size(_DECODED_SHAPE)
|
||||
|
||||
|
||||
def test_forward_rejects_unknown_mode():
|
||||
vae = _StubVAE()
|
||||
with pytest.raises(ValueError, match="Unknown SeedVR2 VAE forward mode"):
|
||||
vae.forward(torch.zeros(*_INPUT_ENCODE_SHAPE), mode="bogus")
|
||||
@@ -0,0 +1,79 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
import comfy.sd
|
||||
import comfy.supported_models
|
||||
import comfy.ldm.seedvr.model as seedvr_model
|
||||
import comfy.ldm.seedvr.vae as seedvr_vae
|
||||
|
||||
|
||||
def test_seedvr2_fp16_manual_cast_only_for_bf16_device(monkeypatch):
|
||||
bf16_device = object()
|
||||
fp16_device = object()
|
||||
|
||||
monkeypatch.setattr(
|
||||
comfy.supported_models.comfy.model_management,
|
||||
"should_use_bf16",
|
||||
lambda device=None: device is bf16_device,
|
||||
)
|
||||
|
||||
bf16_config = comfy.supported_models.SeedVR2({"image_model": "seedvr2"})
|
||||
bf16_config.set_inference_dtype(torch.float16, None, device=bf16_device)
|
||||
assert bf16_config.manual_cast_dtype is torch.bfloat16
|
||||
|
||||
fp16_config = comfy.supported_models.SeedVR2({"image_model": "seedvr2"})
|
||||
fp16_config.set_inference_dtype(torch.float16, None, device=fp16_device)
|
||||
assert fp16_config.manual_cast_dtype is None
|
||||
|
||||
|
||||
def test_seedvr2_text_conditioning_accepts_cfg1_single_branch():
|
||||
context = torch.arange(6, dtype=torch.float32).reshape(1, 3, 2)
|
||||
|
||||
txt, txt_shape = seedvr_model.NaDiT._resolve_text_conditioning(object(), context, [0])
|
||||
|
||||
torch.testing.assert_close(txt, context.squeeze(0))
|
||||
torch.testing.assert_close(txt_shape, torch.tensor([[3]], device=context.device))
|
||||
|
||||
|
||||
def test_seedvr2_vae_decode_memory_covers_full_frame_lab_transfer():
|
||||
wrapper = seedvr_vae.VideoAutoencoderKLWrapper.__new__(seedvr_vae.VideoAutoencoderKLWrapper)
|
||||
latent_channels = seedvr_vae.SEEDVR2_LATENT_CHANNELS
|
||||
estimate = wrapper.comfy_memory_used_decode((1, latent_channels, 26, 120, 160))
|
||||
old_estimate = latent_channels * 120 * 160 * (4 * 8 * 8) * 2
|
||||
|
||||
assert estimate == 101 * 960 * 1280 * 160
|
||||
assert estimate > 15 * 1024 ** 3
|
||||
assert estimate > old_estimate * 100
|
||||
|
||||
|
||||
def test_seedvr2_vae_encode_preserves_compute_dtype(monkeypatch):
|
||||
wrapper = seedvr_vae.VideoAutoencoderKLWrapper.__new__(seedvr_vae.VideoAutoencoderKLWrapper)
|
||||
nn.Module.__init__(wrapper)
|
||||
wrapper._dummy = nn.Parameter(torch.empty(1, dtype=torch.float16))
|
||||
input_dtype = None
|
||||
|
||||
def encode(self, x):
|
||||
nonlocal input_dtype
|
||||
input_dtype = x.dtype
|
||||
return x
|
||||
|
||||
monkeypatch.setattr(seedvr_vae.VideoAutoencoderKL, "encode", encode)
|
||||
|
||||
x = torch.zeros((1, 3, 1, 8, 8), dtype=torch.float32)
|
||||
wrapper._encode_with_raw_latent(x)
|
||||
|
||||
assert input_dtype == torch.float32
|
||||
|
||||
|
||||
def test_seedvr2_vae_ops_cast_weights_to_compute_dtype():
|
||||
attention = seedvr_vae.Attention(query_dim=4, heads=1, dim_head=4).to(torch.float16)
|
||||
hidden_states = torch.zeros((1, 2, 4), dtype=torch.float32)
|
||||
|
||||
output = attention(hidden_states)
|
||||
|
||||
assert output.dtype == torch.float32
|
||||
@@ -0,0 +1,169 @@
|
||||
"""SeedVR2 internals regression tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
args.cpu = True
|
||||
|
||||
import comfy.ldm.seedvr.model as seedvr_model # noqa: E402
|
||||
import comfy.ldm.seedvr.vae as vae_mod # noqa: E402
|
||||
import comfy.ldm.modules.attention as attention # noqa: E402
|
||||
import comfy.ops as comfy_ops # noqa: E402
|
||||
from comfy.ldm.seedvr.vae import ( # noqa: E402
|
||||
causal_norm_wrapper,
|
||||
set_norm_limit,
|
||||
)
|
||||
from comfy.ldm.seedvr.attention import var_attention_optimized_split # noqa: E402
|
||||
|
||||
|
||||
_NUM_CHANNELS = 8
|
||||
_NUM_GROUPS = 4
|
||||
_TENSOR_SHAPE = (1, 8, 2, 4, 4)
|
||||
|
||||
_GROUPNORM_SUBCLASSES = [
|
||||
pytest.param(comfy_ops.disable_weight_init.GroupNorm, id="disable_weight_init"),
|
||||
pytest.param(comfy_ops.manual_cast.GroupNorm, id="manual_cast"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("groupnorm_cls", _GROUPNORM_SUBCLASSES)
|
||||
def test_seedvr_groupnorm_low_limit_uses_chunked_groupnorm_path(groupnorm_cls):
|
||||
real_group_norm = vae_mod.F.group_norm
|
||||
set_norm_limit(1e-9)
|
||||
try:
|
||||
gn = groupnorm_cls(num_channels=_NUM_CHANNELS, num_groups=_NUM_GROUPS)
|
||||
gn.eval()
|
||||
|
||||
forward_hook_calls = []
|
||||
|
||||
def _hook(module, inputs, output):
|
||||
forward_hook_calls.append(tuple(inputs[0].shape))
|
||||
|
||||
spy_calls = []
|
||||
|
||||
def _group_norm_spy(input_tensor, num_groups_arg, *args, **kwargs):
|
||||
spy_calls.append({"num_groups": int(num_groups_arg)})
|
||||
return real_group_norm(input_tensor, num_groups_arg, *args, **kwargs)
|
||||
|
||||
handle = gn.register_forward_hook(_hook)
|
||||
try:
|
||||
with patch.object(vae_mod.F, "group_norm", side_effect=_group_norm_spy):
|
||||
out_tensor = causal_norm_wrapper(gn, torch.randn(*_TENSOR_SHAPE))
|
||||
finally:
|
||||
handle.remove()
|
||||
|
||||
full_calls = len(forward_hook_calls)
|
||||
chunked_calls = sum(1 for entry in spy_calls if entry["num_groups"] < _NUM_GROUPS)
|
||||
|
||||
assert tuple(int(s) for s in out_tensor.shape) == _TENSOR_SHAPE
|
||||
assert full_calls == 0, (
|
||||
f"low-limit GroupNorm gate must NOT take the full-forward path; got full_calls={full_calls}"
|
||||
)
|
||||
assert chunked_calls > 0, (
|
||||
f"low-limit GroupNorm gate must take the chunked path; got chunked_calls={chunked_calls}"
|
||||
)
|
||||
finally:
|
||||
set_norm_limit(None)
|
||||
|
||||
|
||||
def test_seedvr2_7b_swin_attention_forward_uses_optimized_var_attention(monkeypatch):
|
||||
dim = 8
|
||||
heads = 2
|
||||
head_dim = 4
|
||||
attn = seedvr_model.NaSwinAttention(
|
||||
vid_dim=dim,
|
||||
txt_dim=dim,
|
||||
heads=heads,
|
||||
head_dim=head_dim,
|
||||
qk_bias=False,
|
||||
qk_norm=comfy_ops.disable_weight_init.RMSNorm,
|
||||
qk_norm_eps=1e-6,
|
||||
rope_type=None,
|
||||
rope_dim=head_dim,
|
||||
shared_weights=False,
|
||||
window=(2, 1, 1),
|
||||
window_method="720pwin_by_size_bysize",
|
||||
version=True,
|
||||
device="cpu",
|
||||
dtype=torch.float32,
|
||||
operations=comfy_ops.disable_weight_init,
|
||||
)
|
||||
generator = torch.Generator(device="cpu").manual_seed(11)
|
||||
vid = torch.randn(8, dim, generator=generator)
|
||||
txt = torch.randn(3, dim, generator=generator)
|
||||
vid_shape = torch.tensor([[2, 2, 2]], dtype=torch.long)
|
||||
txt_shape = torch.tensor([[3]], dtype=torch.long)
|
||||
calls = []
|
||||
|
||||
def fake_optimized_var_attention(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return kwargs["q"]
|
||||
|
||||
monkeypatch.setattr(seedvr_model, "optimized_var_attention", fake_optimized_var_attention)
|
||||
|
||||
vid_out, txt_out = attn(vid, txt, vid_shape, txt_shape, seedvr_model.Cache(disable=True))
|
||||
|
||||
assert tuple(vid_out.shape) == (8, dim)
|
||||
assert tuple(txt_out.shape) == (3, dim)
|
||||
assert len(calls) == 1
|
||||
call = calls[0]
|
||||
assert tuple(call["q"].shape) == (14, heads, head_dim)
|
||||
assert tuple(call["k"].shape) == (14, heads, head_dim)
|
||||
assert tuple(call["v"].shape) == (14, heads, head_dim)
|
||||
assert call["heads"] == heads
|
||||
assert call["skip_reshape"] is True
|
||||
assert call["skip_output_reshape"] is True
|
||||
assert call["cu_seqlens_q"] == [0, 7, 14]
|
||||
assert call["cu_seqlens_k"] == [0, 7, 14]
|
||||
|
||||
|
||||
def test_var_attention_optimized_split_calls_dense_backend_per_window(monkeypatch):
|
||||
heads = 2
|
||||
head_dim = 3
|
||||
q = torch.arange(30, dtype=torch.float32).reshape(5, heads, head_dim)
|
||||
k = q + 100
|
||||
v = q + 200
|
||||
cu = [0, 2, 5]
|
||||
calls = []
|
||||
|
||||
def fake_optimized_attention(q_arg, k_arg, v_arg, heads_arg, **kwargs):
|
||||
calls.append(
|
||||
{
|
||||
"q_shape": tuple(q_arg.shape),
|
||||
"k_shape": tuple(k_arg.shape),
|
||||
"v_shape": tuple(v_arg.shape),
|
||||
"heads": heads_arg,
|
||||
"kwargs": kwargs,
|
||||
}
|
||||
)
|
||||
return q_arg + v_arg
|
||||
|
||||
monkeypatch.setattr(attention, "optimized_attention", fake_optimized_attention)
|
||||
|
||||
out = var_attention_optimized_split(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
heads,
|
||||
cu,
|
||||
cu,
|
||||
skip_reshape=True,
|
||||
skip_output_reshape=True,
|
||||
)
|
||||
|
||||
assert tuple(out.shape) == (5, heads, head_dim)
|
||||
assert len(calls) == 2
|
||||
assert calls[0]["q_shape"] == (1, heads, 2, head_dim)
|
||||
assert calls[1]["q_shape"] == (1, heads, 3, head_dim)
|
||||
assert all(call["heads"] == heads for call in calls)
|
||||
assert all(call["kwargs"]["skip_reshape"] is True for call in calls)
|
||||
assert all(call["kwargs"]["skip_output_reshape"] is True for call in calls)
|
||||
torch.testing.assert_close(out, q + v, rtol=0, atol=0)
|
||||
|
||||
@@ -0,0 +1,320 @@
|
||||
"""SeedVR2 model, latent-format, and VAE graph regression tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
args.cpu = True
|
||||
|
||||
import comfy # noqa: E402
|
||||
import comfy.latent_formats # noqa: E402
|
||||
import comfy.ldm.seedvr.model as seedvr_model # noqa: E402
|
||||
import comfy.ldm.seedvr.vae as seedvr_vae_mod # noqa: E402
|
||||
import comfy.model_management # noqa: E402
|
||||
import comfy.ops as comfy_ops # noqa: E402
|
||||
import comfy.sample # noqa: E402
|
||||
import comfy.sd as sd_mod # noqa: E402
|
||||
import nodes as nodes_mod # noqa: E402
|
||||
from comfy.ldm.seedvr.model import NaDiT # noqa: E402
|
||||
|
||||
|
||||
_LATENT_CHANNELS = seedvr_vae_mod.SEEDVR2_LATENT_CHANNELS
|
||||
|
||||
|
||||
def _make_standin(positive_conditioning):
|
||||
class _StandIn(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.register_buffer(
|
||||
"positive_conditioning", positive_conditioning
|
||||
)
|
||||
|
||||
_resolve_text_conditioning = NaDiT._resolve_text_conditioning
|
||||
|
||||
return _StandIn()
|
||||
|
||||
|
||||
class _StubModule(nn.Module):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
|
||||
def _capture_last_layer_flags(monkeypatch, vid_dim: int, txt_in_dim: int) -> list[bool]:
|
||||
flags = []
|
||||
|
||||
class _Block(_StubModule):
|
||||
def __init__(self, *args, **kwargs):
|
||||
flags.append(kwargs["is_last_layer"])
|
||||
super().__init__()
|
||||
|
||||
monkeypatch.setattr(seedvr_model, "NaPatchIn", _StubModule)
|
||||
monkeypatch.setattr(seedvr_model, "NaPatchOut", _StubModule)
|
||||
monkeypatch.setattr(seedvr_model, "TimeEmbedding", _StubModule)
|
||||
monkeypatch.setattr(seedvr_model, "NaMMSRTransformerBlock", _Block)
|
||||
|
||||
seedvr_model.NaDiT(
|
||||
norm_eps=1e-5,
|
||||
num_layers=4,
|
||||
mlp_type="normal",
|
||||
vid_dim=vid_dim,
|
||||
txt_in_dim=txt_in_dim,
|
||||
heads=24,
|
||||
mm_layers=3,
|
||||
operations=comfy_ops.disable_weight_init,
|
||||
)
|
||||
|
||||
return flags
|
||||
|
||||
|
||||
class _Model:
|
||||
def __init__(self, latent_format):
|
||||
self._latent_format = latent_format
|
||||
|
||||
def get_model_object(self, name):
|
||||
assert name == "latent_format"
|
||||
return self._latent_format
|
||||
|
||||
|
||||
class _Patcher:
|
||||
def get_free_memory(self, device):
|
||||
return 1024 * 1024 * 1024
|
||||
|
||||
|
||||
class _EncodeWrapper(seedvr_vae_mod.VideoAutoencoderKLWrapper):
|
||||
def __init__(self, encoded):
|
||||
nn.Module.__init__(self)
|
||||
self.encoded = encoded
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.seen = []
|
||||
|
||||
def encode(self, x):
|
||||
self.seen.append(tuple(x.shape))
|
||||
return self.encoded.to(device=x.device, dtype=x.dtype)
|
||||
|
||||
|
||||
class _DecodeWrapper(seedvr_vae_mod.VideoAutoencoderKLWrapper):
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.calls = []
|
||||
|
||||
def decode(self, z, seedvr2_tiling=None):
|
||||
self.calls.append({"shape": tuple(z.shape), "seedvr2_tiling": seedvr2_tiling})
|
||||
if z.ndim == 4:
|
||||
b, tc, h, w = z.shape
|
||||
t = tc // _LATENT_CHANNELS
|
||||
else:
|
||||
b, _, t, h, w = z.shape
|
||||
return torch.zeros(b, 3, t, h * 8, w * 8, dtype=z.dtype, device=z.device)
|
||||
|
||||
|
||||
def test_seedvr2_wrapper_public_encode_returns_tensor(monkeypatch):
|
||||
raw_latent = torch.full((1, _LATENT_CHANNELS, 1, 4, 5), 2.0)
|
||||
seen_shapes = []
|
||||
|
||||
def base_encode(self, x):
|
||||
seen_shapes.append(tuple(x.shape))
|
||||
return raw_latent.to(device=x.device, dtype=x.dtype)
|
||||
|
||||
monkeypatch.setattr(seedvr_vae_mod.VideoAutoencoderKL, "encode", base_encode)
|
||||
|
||||
vae = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(seedvr_vae_mod.VideoAutoencoderKLWrapper)
|
||||
nn.Module.__init__(vae)
|
||||
vae._dummy = nn.Parameter(torch.zeros((), dtype=torch.float32))
|
||||
|
||||
latent = vae.encode(torch.zeros(1, 3, 32, 40))
|
||||
|
||||
assert type(latent) is torch.Tensor
|
||||
assert tuple(latent.shape) == (1, _LATENT_CHANNELS, 4, 5)
|
||||
assert seen_shapes == [(1, 3, 1, 32, 40)]
|
||||
|
||||
|
||||
def test_seedvr2_wrapper_private_encode_helper_keeps_raw_latent(monkeypatch):
|
||||
raw_latent = torch.full((1, _LATENT_CHANNELS, 1, 4, 5), 3.0)
|
||||
|
||||
def base_encode(self, x):
|
||||
return raw_latent.to(device=x.device, dtype=x.dtype)
|
||||
|
||||
monkeypatch.setattr(seedvr_vae_mod.VideoAutoencoderKL, "encode", base_encode)
|
||||
|
||||
vae = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(seedvr_vae_mod.VideoAutoencoderKLWrapper)
|
||||
nn.Module.__init__(vae)
|
||||
vae._dummy = nn.Parameter(torch.zeros((), dtype=torch.float32))
|
||||
|
||||
latent, raw = vae._encode_with_raw_latent(torch.zeros(1, 3, 32, 40))
|
||||
|
||||
assert tuple(latent.shape) == (1, _LATENT_CHANNELS, 4, 5)
|
||||
assert tuple(raw.shape) == (1, _LATENT_CHANNELS, 1, 4, 5)
|
||||
assert torch.equal(raw, raw_latent)
|
||||
|
||||
|
||||
def _make_vae(wrapper):
|
||||
vae = sd_mod.VAE.__new__(sd_mod.VAE)
|
||||
vae.first_stage_model = wrapper
|
||||
vae.device = torch.device("cpu")
|
||||
vae.output_device = torch.device("cpu")
|
||||
vae.vae_dtype = torch.float32
|
||||
vae.latent_channels = _LATENT_CHANNELS
|
||||
vae.latent_dim = 3
|
||||
vae.downscale_ratio = (lambda a: max(0, (a + 3) // 4), 8, 8)
|
||||
vae.upscale_ratio = (lambda a: max(0, a * 4 - 3), 8, 8)
|
||||
vae.output_channels = 3
|
||||
vae.disable_offload = True
|
||||
vae.extra_1d_channel = None
|
||||
vae.crop_input = False
|
||||
vae.not_video = False
|
||||
vae.handles_tiling = isinstance(wrapper, seedvr_vae_mod.VideoAutoencoderKLWrapper)
|
||||
vae.format_encoded = wrapper.comfy_format_encoded
|
||||
vae.patcher = _Patcher()
|
||||
vae.process_input = lambda image: image
|
||||
vae.process_output = lambda image: image.add(1.0).div(2.0).clamp(0.0, 1.0)
|
||||
vae.vae_output_dtype = lambda: torch.float32
|
||||
vae.memory_used_encode = lambda shape, dtype: 1
|
||||
vae.memory_used_decode = lambda shape, dtype: 1
|
||||
vae.throw_exception_if_invalid = lambda: None
|
||||
vae.vae_encode_crop_pixels = lambda pixels: pixels
|
||||
vae.spacial_compression_decode = lambda: 8
|
||||
vae.temporal_compression_decode = lambda: 4
|
||||
return vae
|
||||
|
||||
|
||||
def test_missing_context_falls_back_to_positive_buffer():
|
||||
pos_buffer = torch.full((58, 5120), 7.0)
|
||||
standin = _make_standin(pos_buffer)
|
||||
txt, txt_shape = standin._resolve_text_conditioning(None)
|
||||
assert txt.shape == (58, 5120)
|
||||
assert (txt == 7.0).all(), (
|
||||
"fallback path must use the positive_conditioning buffer "
|
||||
"verbatim, not a zero tensor"
|
||||
)
|
||||
assert txt_shape.shape == (1, 1)
|
||||
assert txt_shape[0, 0].item() == 58
|
||||
|
||||
|
||||
def test_seedvr2_7b_keeps_final_block_text_path(monkeypatch):
|
||||
assert _capture_last_layer_flags(monkeypatch, vid_dim=3072, txt_in_dim=3072) == [
|
||||
False,
|
||||
False,
|
||||
False,
|
||||
False,
|
||||
]
|
||||
|
||||
|
||||
def test_seedvr2_7b_rope3d_matches_wrapper_oracle():
|
||||
rope = seedvr_model.get_na_rope("rope3d", dim=64)
|
||||
generator = torch.Generator(device="cpu").manual_seed(0)
|
||||
q = torch.randn(4, 2, 128, generator=generator)
|
||||
k = torch.randn(4, 2, 128, generator=generator)
|
||||
shape = torch.tensor([[1, 2, 2]], dtype=torch.long)
|
||||
freqs = rope.get_axial_freqs(1, 2, 2).reshape(4, -1)
|
||||
|
||||
expected_q = seedvr_model._apply_seedvr2_rotary_emb(
|
||||
freqs,
|
||||
q.permute(1, 0, 2).float(),
|
||||
).to(q.dtype).permute(1, 0, 2)
|
||||
expected_k = seedvr_model._apply_seedvr2_rotary_emb(
|
||||
freqs,
|
||||
k.permute(1, 0, 2).float(),
|
||||
).to(k.dtype).permute(1, 0, 2)
|
||||
|
||||
actual_q, actual_k = rope(q.clone(), k.clone(), shape, seedvr_model.Cache(disable=True))
|
||||
|
||||
torch.testing.assert_close(actual_q, expected_q, rtol=0, atol=0)
|
||||
torch.testing.assert_close(actual_k, expected_k, rtol=0, atol=0)
|
||||
|
||||
|
||||
def test_seedvr2_forward_requires_conditioning_latents():
|
||||
model = NaDiT.__new__(NaDiT)
|
||||
x = torch.zeros(1, _LATENT_CHANNELS, 1, 4, 5)
|
||||
|
||||
with pytest.raises(ValueError, match="requires conditioning latents"):
|
||||
NaDiT.forward(model, x, timestep=torch.tensor([1.0]), context=None)
|
||||
|
||||
|
||||
def test_seedvr2_latent_format_uses_native_video_latent_shape():
|
||||
latent_format = comfy.latent_formats.SeedVR2()
|
||||
latent_image = torch.zeros(1, 1, 4, 5)
|
||||
|
||||
fixed = comfy.sample.fix_empty_latent_channels(_Model(latent_format), latent_image)
|
||||
|
||||
assert latent_format.latent_channels == _LATENT_CHANNELS
|
||||
assert latent_format.latent_dimensions == 3
|
||||
assert fixed.shape == (1, _LATENT_CHANNELS, 1, 4, 5)
|
||||
|
||||
|
||||
def test_seedvr2_model_requires_native_5d_latent():
|
||||
latent = torch.zeros(1, _LATENT_CHANNELS, 2, 4, 5)
|
||||
assert NaDiT._check_seedvr2_video_latent(latent, _LATENT_CHANNELS, "latent") is latent
|
||||
|
||||
with pytest.raises(ValueError, match="5-D native latent"):
|
||||
NaDiT._check_seedvr2_video_latent(torch.zeros(1, _LATENT_CHANNELS * 2, 4, 5), _LATENT_CHANNELS, "latent")
|
||||
|
||||
|
||||
def test_seedvr2_encode_and_encode_tiled_preserve_native_latent_contract(monkeypatch):
|
||||
monkeypatch.setattr(sd_mod.model_management, "load_models_gpu", lambda *a, **k: None)
|
||||
|
||||
encoded = torch.full((1, _LATENT_CHANNELS, 2, 4, 5), 2.0)
|
||||
vae = _make_vae(_EncodeWrapper(encoded))
|
||||
pixels = torch.zeros(1, 5, 32, 40, 3)
|
||||
|
||||
node_output = nodes_mod.VAEEncode().encode(vae, pixels)[0]
|
||||
node_latent = node_output["samples"]
|
||||
assert set(node_output) == {"samples"}
|
||||
assert tuple(node_latent.shape) == (1, _LATENT_CHANNELS, 2, 4, 5)
|
||||
assert node_latent.dtype == torch.float32
|
||||
assert node_latent.stride()[-1] == 1
|
||||
assert torch.equal(node_latent, torch.full_like(node_latent, 2.0 * seedvr_vae_mod.BYTEDANCE_VAE_SCALING_FACTOR))
|
||||
|
||||
tiled = torch.full((1, _LATENT_CHANNELS, 2, 4, 5), 3.0)
|
||||
monkeypatch.setattr(seedvr_vae_mod, "tiled_vae", MagicMock(return_value=tiled))
|
||||
tiled_output = nodes_mod.VAEEncodeTiled().encode(
|
||||
vae,
|
||||
pixels,
|
||||
tile_size=512,
|
||||
overlap=64,
|
||||
temporal_size=16,
|
||||
temporal_overlap=4,
|
||||
)[0]
|
||||
tiled_latent = tiled_output["samples"]
|
||||
assert set(tiled_output) == {"samples"}
|
||||
assert tuple(tiled_latent.shape) == (1, _LATENT_CHANNELS, 2, 4, 5)
|
||||
assert tiled_latent.dtype == torch.float32
|
||||
assert torch.equal(tiled_latent, torch.full_like(tiled_latent, 3.0 * seedvr_vae_mod.BYTEDANCE_VAE_SCALING_FACTOR))
|
||||
|
||||
|
||||
def test_vaedecode_tiled_spatial_applies_temporal_discarded(monkeypatch):
|
||||
monkeypatch.setattr(sd_mod.model_management, "load_models_gpu", lambda *a, **k: None)
|
||||
vae = _make_vae(_DecodeWrapper())
|
||||
|
||||
nodes_mod.VAEDecodeTiled().decode(
|
||||
vae,
|
||||
{"samples": torch.zeros(1, _LATENT_CHANNELS, 2, 4, 5)},
|
||||
tile_size=512,
|
||||
overlap=64,
|
||||
temporal_size=16,
|
||||
temporal_overlap=4,
|
||||
)
|
||||
|
||||
# Spatial inputs flow through; temporal inputs are discarded as public tiling
|
||||
# knobs, but SeedVR2's internal MemoryState causal slicing is left intact.
|
||||
assert vae.first_stage_model.calls == [
|
||||
{
|
||||
"shape": (1, _LATENT_CHANNELS, 2, 4, 5),
|
||||
"seedvr2_tiling": {
|
||||
"enable_tiling": True,
|
||||
"tile_size": (512, 512),
|
||||
"tile_overlap": (64, 64),
|
||||
"temporal_size": None,
|
||||
"temporal_overlap": None,
|
||||
},
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,94 @@
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
import comfy.ldm.seedvr.vae as vae_mod # noqa: E402
|
||||
from comfy_extras import nodes_seedvr # noqa: E402
|
||||
|
||||
|
||||
_LATENT_CHANNELS = vae_mod.SEEDVR2_LATENT_CHANNELS
|
||||
|
||||
|
||||
def _make_wrapper() -> vae_mod.VideoAutoencoderKLWrapper:
|
||||
wrapper = vae_mod.VideoAutoencoderKLWrapper.__new__(
|
||||
vae_mod.VideoAutoencoderKLWrapper
|
||||
)
|
||||
nn.Module.__init__(wrapper)
|
||||
return wrapper
|
||||
|
||||
|
||||
def _fingerprint_decode_(self, z, return_dict=True):
|
||||
b = int(z.shape[0])
|
||||
t = int(z.shape[2])
|
||||
h = int(z.shape[3])
|
||||
w = int(z.shape[4])
|
||||
out = torch.empty(b, 3, t, h * 8, w * 8)
|
||||
for batch_idx in range(b):
|
||||
out[batch_idx].fill_(float(batch_idx + 1))
|
||||
return out
|
||||
|
||||
|
||||
def _decode_with_patches(wrapper, z):
|
||||
with patch.object(vae_mod.VideoAutoencoderKL, "decode_", _fingerprint_decode_):
|
||||
return wrapper.decode(z)
|
||||
|
||||
|
||||
def test_decode_b2_t3_multi_frame_batch_unchanged():
|
||||
wrapper = _make_wrapper()
|
||||
|
||||
out = _decode_with_patches(wrapper, torch.zeros(2, _LATENT_CHANNELS * 3, 2, 2))
|
||||
|
||||
assert tuple(out.shape) == (2, 3, 3, 16, 16)
|
||||
|
||||
|
||||
class _Wrapper(vae_mod.VideoAutoencoderKLWrapper):
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self.calls = []
|
||||
|
||||
def parameters(self):
|
||||
return iter([torch.nn.Parameter(torch.zeros(()))])
|
||||
|
||||
def _decode_stub(self, latent):
|
||||
self.calls.append(tuple(latent.shape))
|
||||
return torch.zeros(latent.shape[0], 3, latent.shape[2], latent.shape[3] * 8, latent.shape[4] * 8)
|
||||
|
||||
|
||||
def test_seedvr2_wrapper_decode_accepts_5d_channel_first_latents_without_preprocessor_state():
|
||||
wrapper = _Wrapper()
|
||||
|
||||
with patch.object(vae_mod.VideoAutoencoderKL, "decode_", _decode_stub):
|
||||
out = wrapper.decode(torch.zeros(1, _LATENT_CHANNELS, 2, 4, 5))
|
||||
|
||||
assert tuple(out.shape) == (1, 3, 2, 32, 40)
|
||||
assert wrapper.calls == [(1, _LATENT_CHANNELS, 2, 4, 5)]
|
||||
|
||||
|
||||
def test_seedvr2_wrapper_decode_rejects_wrong_rank_latents():
|
||||
wrapper = _Wrapper()
|
||||
|
||||
with pytest.raises(RuntimeError, match=r"latent input must be 4-D collapsed .* or 5-D"):
|
||||
wrapper.decode(torch.zeros(1, _LATENT_CHANNELS, 4))
|
||||
|
||||
|
||||
def _t_padded(t_in: int) -> int:
|
||||
if t_in == 1:
|
||||
return 1
|
||||
if t_in <= 4:
|
||||
return 5
|
||||
if (t_in - 1) % 4 == 0:
|
||||
return t_in
|
||||
return t_in + (4 - ((t_in - 1) % 4))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("t_in", [1, 5, 9])
|
||||
def test_t_padded_matches_cut_videos(t_in):
|
||||
dummy = torch.zeros(1, t_in, 1, 1, 1)
|
||||
assert nodes_seedvr.cut_videos(dummy).shape[1] == _t_padded(t_in)
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user